diff --git a/.github/actions/detect-changes/action.yml b/.github/actions/detect-changes/action.yml index ade05ba124d6d..e176ff67e3416 100644 --- a/.github/actions/detect-changes/action.yml +++ b/.github/actions/detect-changes/action.yml @@ -51,6 +51,12 @@ outputs: rust: description: Run `cargo test` for the Tauri bootstrap installer. value: ${{ steps.classify.outputs.rust }} + context_continuity: + description: Run exact paired native owner and Governor qualification. + value: ${{ steps.classify.outputs.context_continuity }} + current_owner_integration: + description: Run canonical owner to Ares consumer qualification. + value: ${{ steps.classify.outputs.current_owner_integration }} mcp_catalog: description: Require MCP catalog security review label. value: ${{ steps.classify.outputs.mcp_catalog }} @@ -105,7 +111,7 @@ runs: if CHANGED="$(gh api \ --paginate \ "repos/${REPO}/compare/${BASE_SHA}...${HEAD_SHA}" \ - --jq '.files[]?.filename')"; then + --jq '.files[]? | .filename, (.previous_filename // empty)')"; then break fi if [ "$i" = 3 ]; then @@ -118,6 +124,15 @@ runs: done fi + # A capped compare cannot prove either native qualification irrelevant. + # Renames emit both paths, so this threshold may also conservatively + # run all lanes before the 300-file API limit is reached. + if [ "$EVENT_NAME" = "pull_request" ] && [ -n "$CHANGED" ] && \ + [ "$(printf '%s\n' "$CHANGED" | wc -l)" -ge 300 ]; then + echo "::warning::compare may be capped — failing open (all lanes run)" + CHANGED="" + fi + echo "Changed files:" printf '%s\n' "${CHANGED:-(none)}" printf '%s\n' "${CHANGED:-}" | python3 scripts/ci/classify_changes.py diff --git a/.github/workflows/ci.yaml b/.github/workflows/ci.yaml index 397fd0e2db063..a8f6f5e8c6ca0 100644 --- a/.github/workflows/ci.yaml +++ b/.github/workflows/ci.yaml @@ -58,6 +58,8 @@ jobs: mcp_catalog: ${{ steps.classify.outputs.mcp_catalog }} ci_review: ${{ steps.classify.outputs.ci_review }} ci_review_files: ${{ steps.classify.outputs.ci_review_files }} + context_continuity: ${{ steps.classify.outputs.context_continuity }} + current_owner_integration: ${{ steps.classify.outputs.current_owner_integration }} event_name: ${{ github.event_name }} steps: - uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 @@ -118,6 +120,24 @@ jobs: if: needs.detect.outputs.rust == 'true' uses: ./.github/workflows/rust-tests.yml + # Preserve the existing native owners and exact pairing. PR-controlled + # qualification receives only read access, without inherited secrets. + context-continuity: + name: Context continuity qualification + needs: detect + if: needs.detect.outputs.context_continuity == 'true' + permissions: + contents: read + uses: ./.github/workflows/context-continuity-qualification.yml + + current-owner-integration: + name: Current owner integration + needs: detect + if: needs.detect.outputs.current_owner_integration == 'true' + permissions: + contents: read + uses: ./.github/workflows/current-owner-integration.yml + e2e-desktop: name: Desktop E2E needs: detect @@ -234,6 +254,8 @@ jobs: - js-tests - installer-tests - rust-tests + - context-continuity + - current-owner-integration - e2e-desktop - docs-site - history-check @@ -290,9 +312,13 @@ jobs: - name: Lint GitHub Actions workflows # The disabled Desktop lane is an explicit CI-01 applicability exception. - # Lint the planner workflow enforced by this aggregate. Other workflows - # own their existing shellcheck debt and are outside this bounded gate. - run: ./actionlint -ignore 'constant expression "false" in condition' .github/workflows/ci.yaml + # Lint the planner and native qualifications enforced by this aggregate. + # Other workflows retain their existing independent lint ownership. + run: >- + ./actionlint -ignore 'constant expression "false" in condition' + .github/workflows/ci.yaml + .github/workflows/context-continuity-qualification.yml + .github/workflows/current-owner-integration.yml - name: Evaluate job results id: evaluate diff --git a/.github/workflows/context-continuity-qualification.yml b/.github/workflows/context-continuity-qualification.yml index fe8430a78728d..a73a8552df2c1 100644 --- a/.github/workflows/context-continuity-qualification.yml +++ b/.github/workflows/context-continuity-qualification.yml @@ -1,118 +1,16 @@ name: Context continuity qualification on: + workflow_call: + outputs: + native_external_owner_result: + description: Result of the exact paired native owner qualification. + value: ${{ jobs.native-external-owner.outputs.qualification_result }} + focused_tests_result: + description: Result of the pinned Governor and focused Ares qualification. + value: ${{ jobs.focused-tests.outputs.qualification_result }} push: - branches: [main, feat/context-continuity-v4-20260923] - pull_request: - paths: - - '.github/workflows/context-continuity-qualification.yml' - - 'ares_runtime/continuity/**' - - 'tests/ares_runtime/test_continuity*' - - 'tests/ares_runtime/test_context_rebase_state.py' - - 'tests/ares_runtime/test_context_controller_credentials.py' - - 'tests/ares_runtime/test_context_authority_binding.py' - - 'tests/ares_runtime/test_context_native*.py' - - 'hermes_state_context_authority.py' - - 'hermes_cli/context_authority.py' - - 'tests/ares_runtime/test_managed_calls.py' - - 'tests/hermes_cli/test_goals.py' - - 'tests/hermes_cli/test_goal_lifecycle_contract.py' - - 'tests/test_model_tools.py' - - 'tests/tools/test_code_execution.py' - - 'tests/run_agent/test_message_sequence_repair.py' - - 'tests/run_agent/test_tool_executor_contextvar_propagation.py' - - 'tests/hermes_cli/test_heartbeat.py' - - 'tests/hermes_cli/test_loops.py' - - 'tests/test_run_checkpoint_import_boundaries.py' - - 'tests/test_run_checkpoint*.py' - - 'tests/test_run_task_custody.py' - - 'tests/test_run_custody_cold_recovery.py' - - 'tests/test_run_custody_ready_recovery.py' - - 'tests/test_run_custody_obligations.py' - - 'tests/gateway/test_context_input.py' - - 'tests/gateway/test_context_input_recovery.py' - - 'tests/ares_runtime/test_context_input_lifetime.py' - - 'tests/cli/test_quick_commands.py' - - 'tests/test_estop.py' - - 'tests/ares_runtime/test_permit_readback.py' - - 'tests/test_ares_collaboration.py' - - 'ares_runtime/collaboration.py' - - 'tests/gateway/test_pre_gateway_dispatch.py' - - 'tests/gateway/test_restart_resume_pending.py' - - 'tests/gateway/test_multiplex_session_db_profile_scope.py' - - 'tests/gateway/test_multiplex_adapter_registry.py' - - 'tests/gateway/test_adapter_startup_secret_scope.py' - - 'tests/gateway/test_startup_connect_parallel.py' - - 'tests/gateway/test_run_progress_topics.py' - - 'tests/gateway/test_42039_duplicate_user_message.py' - - 'tests/gateway/test_internal_event_never_interrupts_busy_session.py' - - 'tests/gateway/test_multiplex_busy_input_mode.py' - - 'tests/gateway/test_platform_base.py' - - 'tests/gateway/test_base_topic_sessions.py' - - 'tests/gateway/test_profile_routing.py' - - 'tests/gateway/test_busy_session_auth_bypass.py' - - 'tests/gateway/test_busy_session_ack.py' - - 'gateway/context_input.py' - - 'gateway/context_input_recovery.py' - - 'gateway/turn_context.py' - - 'gateway/platforms/base.py' - - 'gateway/run.py' - - 'tests/state/test_message_copy_mapping.py' - - 'tests/state/test_message_row_publication.py' - - 'tests/test_compression_watermark_commit.py' - - 'tests/state/test_todo_compaction.py' - - 'tests/run_agent/test_in_place_compaction.py' - - 'tests/agent/test_micro_compaction.py' - - 'tests/hermes_state/test_append_messages_batch.py' - - 'tests/tui_gateway/test_run_checkpoint_claim_rpc.py' - - 'tests/tui_gateway/test_goal_command.py' - - 'tests/test_turn_run_custody.py' - - 'tests/tui_gateway/test_inline_rpc_gil_starvation.py' - - 'tests/tui_gateway/test_kanban_notify_poller.py' - - 'tests/test_tui_gateway_server.py' - - 'tests/cli/test_cli_goal_interrupt.py' - - 'tests/cli/test_cli_async_delegation_delivery.py' - - 'tests/gateway/test_goal_resume_restart.py' - - 'tests/agent/test_synthetic_turn_display_kind.py' - - 'tests/agent/test_turn_context.py' - - 'tests/run_agent/test_run_agent.py' - - 'tests/run_agent/test_1630_context_overflow_loop.py' - - 'tests/run_agent/test_compression_lock_defer.py' - - 'agent/agent_init.py' - - 'agent/agent_runtime_helpers.py' - - 'agent/tool_executor.py' - - 'model_tools.py' - - 'agent/conversation_loop.py' - - 'agent/chat_completion_helpers.py' - - 'agent/codex_runtime.py' - - 'agent/run_checkpoint_custody.py' - - 'agent/turn_context.py' - - 'agent/context_input.py' - - 'cli.py' - - 'hermes_cli/goals.py' - - 'hermes_cli/heartbeat.py' - - 'hermes_cli/loops.py' - - 'hermes_state.py' - - 'hermes_state_common.py' - - 'hermes_state_continuity.py' - - 'hermes_state_inbox.py' - - 'hermes_state_input_turns.py' - - 'agent/turn_finalizer.py' - - 'hermes_state_runs.py' - - 'plugins/context_engine/_context_governor/**' - - 'scripts/run_checkpoint_context.py' - - 'scripts/run_checkpoint_claim.py' - - 'scripts/run_checkpoint_resume.py' - - 'tui_gateway/run_checkpoint_rpc.py' - - 'scripts/run_tests.sh' - - 'tests/run_agent/test_streaming.py' - - 'tests/run_agent/test_run_agent_codex_responses.py' - - 'tests/run_agent/test_codex_sdk_transform_bypass.py' - - 'tests/plugins/test_context_governor*.py' - - 'tui_gateway/server.py' - - 'tui_gateway/methods_prompt.py' - - 'tui_gateway/compute_host.py' - - 'docs/context-continuity/**' + branches: [feat/context-continuity-v4-20260923] permissions: contents: read @@ -123,6 +21,8 @@ concurrency: jobs: native-external-owner: + outputs: + qualification_result: ${{ job.status }} name: Native external owner paired source runs-on: ubuntu-latest timeout-minutes: 45 @@ -209,6 +109,8 @@ jobs: HERMES_PYTHON="$PWD/.venv/bin/python" HERMES_TEST_WORKERS=1 HERMES_TEST_FILE_RETRIES=0 \ scripts/run_tests.sh tests/ares_runtime/test_context_native_paired.py -- -q focused-tests: + outputs: + qualification_result: ${{ job.status }} runs-on: ubuntu-latest timeout-minutes: 20 steps: diff --git a/.github/workflows/current-owner-integration.yml b/.github/workflows/current-owner-integration.yml index db2ec2229aa1b..37060edaacb2e 100644 --- a/.github/workflows/current-owner-integration.yml +++ b/.github/workflows/current-owner-integration.yml @@ -1,37 +1,19 @@ name: Current owner integration on: - push: - branches: [main] - paths: - - '.github/workflows/current-owner-integration.yml' - - 'ares_runtime/collaboration.py' - - 'ares_runtime/governed_context.py' - - 'ares_runtime/__init__.py' - - 'tests/owner_integration/**' - - 'tests/test_ares_collaboration.py' - - 'tests/ares_runtime/test_governed_context_materialization.py' - - 'tests/ares_runtime/test_memory_witness_v2.py' - - 'tests/ares_runtime/test_policy_basis_v2.py' - - 'tests/ares_runtime/fixtures/profile_runtime_v2_owner.json' - pull_request: - paths: - - '.github/workflows/current-owner-integration.yml' - - 'ares_runtime/collaboration.py' - - 'ares_runtime/governed_context.py' - - 'ares_runtime/__init__.py' - - 'tests/owner_integration/**' - - 'tests/test_ares_collaboration.py' - - 'tests/ares_runtime/test_governed_context_materialization.py' - - 'tests/ares_runtime/test_memory_witness_v2.py' - - 'tests/ares_runtime/test_policy_basis_v2.py' - - 'tests/ares_runtime/fixtures/profile_runtime_v2_owner.json' + workflow_call: + outputs: + profile_runtime_consumer_result: + description: Result of the canonical owner to Ares consumer qualification. + value: ${{ jobs.profile-runtime-consumer.outputs.qualification_result }} permissions: contents: read jobs: profile-runtime-consumer: + outputs: + qualification_result: ${{ job.status }} name: Current profile-runtime owner to Ares runs-on: ubuntu-latest timeout-minutes: 25 diff --git a/agent/agent_init.py b/agent/agent_init.py index 06b16b733eb76..2e992886de6ea 100644 --- a/agent/agent_init.py +++ b/agent/agent_init.py @@ -2106,6 +2106,8 @@ def init_agent( _agent_section = _agent_cfg.get("agent", {}) if not isinstance(_agent_section, dict): _agent_section = {} + from agent.transports.ri_llm import configure_ri_pipeline + configure_ri_pipeline(agent, _agent_cfg) agent._tool_use_enforcement = _agent_section.get("tool_use_enforcement", "auto") # Execution-discipline guidance gate: "auto" (default — matches @@ -2821,7 +2823,16 @@ def _parse_prune_int(raw, default): except Exception: pass - if _engine_name != "compressor": + if _engine_name == "ri-context-governor": + # Ares explicitly selects a certified, receipt-preserving engine. + # Readiness failure must not select Hermes' LLM compressor instead. + from plugins.context_engine import ( + ContextEngineActivationError, + load_context_engine_strict, + ) + + _selected_engine = load_context_engine_strict(_engine_name) + elif _engine_name != "compressor": # Try loading from plugins/context_engine// try: from plugins.context_engine import load_context_engine @@ -2905,6 +2916,7 @@ def _parse_prune_int(raw, default): api_key=getattr(agent, "api_key", ""), provider=agent.provider, api_mode=agent.api_mode, + max_tokens=agent.max_tokens, threshold_percent=compression_threshold, ) if not agent.quiet_mode: @@ -2937,8 +2949,12 @@ def _parse_prune_int(raw, default): if callable(_bind_session_state): try: _bind_session_state(session_db=session_db, session_id=agent.session_id) - except Exception: - pass + except Exception as _ce_bind_err: + if _engine_name == "ri-context-governor": + raise ContextEngineActivationError( + f"configured context engine '{_engine_name}' failed its " + f"session binding: {_ce_bind_err}" + ) from _ce_bind_err agent.compression_enabled = compression_enabled agent.compression_in_place = compression_in_place agent.context_rebase_enabled = compression_context_rebase @@ -3045,6 +3061,7 @@ def _parse_prune_int(raw, default): agent.enabled_toolsets is None or "context_engine" in agent.enabled_toolsets ) + and "context_engine" not in (agent.disabled_toolsets or []) ): _existing_tool_names = { t.get("function", {}).get("name") @@ -3075,6 +3092,11 @@ def _parse_prune_int(raw, default): # Notify context engine of session start if hasattr(agent, "context_compressor") and agent.context_compressor: + _session_start_context = ( + {"session_db": session_db} + if _engine_name == "ri-context-governor" + else {} + ) try: agent.context_compressor.on_session_start( agent.session_id, @@ -3083,8 +3105,14 @@ def _parse_prune_int(raw, default): model=agent.model, context_length=getattr(agent.context_compressor, "context_length", 0), conversation_id=getattr(agent, "_gateway_session_key", None), + **_session_start_context, ) except Exception as _ce_err: + if _engine_name == "ri-context-governor": + raise ContextEngineActivationError( + f"configured context engine '{_engine_name}' failed its " + f"session start: {_ce_err}" + ) from _ce_err _ra().logger.debug("Context engine on_session_start: %s", _ce_err) agent._subdirectory_hints = SubdirectoryHintTracker( diff --git a/agent/agent_runtime_helpers.py b/agent/agent_runtime_helpers.py index be22ac87665c3..16e82fb718afb 100644 --- a/agent/agent_runtime_helpers.py +++ b/agent/agent_runtime_helpers.py @@ -1508,6 +1508,64 @@ def drop_thinking_only_and_merge_users( +def _snapshot_context_handoff(agent): + """Return a rollback for host state until the context engine accepts a route. + + Keep engine admission and durable state with their existing owners. This + captures only reversible host fields; fallback index/cooldown progress must + survive refusal so rejected candidates are not retried indefinitely. + """ + missing = object() + snapshot = { + name: getattr(agent, name, missing) + for name in ( + "model", "provider", "requested_provider", "base_url", "api_mode", + "api_key", "client", "_anthropic_client", "_anthropic_api_key", + "_anthropic_base_url", "_is_anthropic_oauth", "_client_kwargs", + "_credential_pool", "_credential_pool_entry_id", + "_config_context_length", "_custom_providers", "_reasoning_echo_flag", + "_use_prompt_caching", "_use_native_cache_layout", "_transport_cache", + "_fallback_activated", + ) + } + # Preserve cache/kwargs identity as well as contents: clearing a live dict + # must not poison the rollback target or a reference held by its caller. + dict_contents = { + name: dict(value) for name, value in snapshot.items() + if name in {"_client_kwargs", "_transport_cache"} and isinstance(value, dict) + } + + def restore(): + rejected_clients = ( + getattr(agent, "client", None), + getattr(agent, "_anthropic_client", None), + ) + for name, value in snapshot.items(): + if value is missing: + if hasattr(agent, name): + delattr(agent, name) + else: + if name in dict_contents: + value.clear() + value.update(dict_contents[name]) + setattr(agent, name, value) + # Retire only newly attached clients, never the restored shared ones. + # Socket shutdown defers FD release to GC instead of hard-closing a + # shared client from a thread whose earlier request may still unwind. + preserved = (snapshot["client"], snapshot["_anthropic_client"]) + retired = set() + for client in rejected_clients: + if client is None or any(client is old for old in preserved) or id(client) in retired: + continue + retired.add(id(client)) + try: + agent._retire_shared_openai_client(client, reason="context_handoff_refused") + except Exception: + logger.debug("Rejected context handoff client retirement failed", exc_info=True) + + return restore + + def restore_primary_runtime(agent) -> bool: """Restore the primary runtime at the start of a new turn. @@ -1622,6 +1680,8 @@ def restore_primary_runtime(agent) -> bool: provider_fallback_active = bool( getattr(agent, "_provider_fallback_active", False) ) + restore_handoff = _snapshot_context_handoff(agent) + handoff_accepted = False try: # ── Core runtime state ── agent.model = rt["model"] @@ -1679,14 +1739,23 @@ def restore_primary_runtime(agent) -> bool: # ── Restore context engine state ── cc = agent.context_compressor - cc.update_model( - model=rt["compressor_model"], - context_length=rt["compressor_context_length"], - base_url=rt["compressor_base_url"], - api_key=rt["compressor_api_key"], - provider=rt["compressor_provider"], - api_mode=rt.get("compressor_api_mode", ""), - ) + from agent.auxiliary_client import _update_compressor_model + + if cc is not None: + # A host with no engine has no engine state to restore. Every + # actual engine still requires its saved destination and must + # accept it before the host handoff can commit. + _update_compressor_model( + cc, + model=rt["compressor_model"], + context_length=rt["compressor_context_length"], + base_url=rt["compressor_base_url"], + api_key=rt["compressor_api_key"], + provider=rt["compressor_provider"], + api_mode=rt.get("compressor_api_mode", ""), + max_tokens=getattr(agent, "max_tokens", None), + ) + handoff_accepted = True # ── Rebind and re-select the primary credential pool ── # A cross-provider fallback attaches the fallback provider's pool. The @@ -1824,6 +1893,8 @@ def restore_primary_runtime(agent) -> bool: pass return True except Exception as e: + if not handoff_accepted: + restore_handoff() logger.warning("Failed to restore primary runtime: %s", e) return False @@ -2768,56 +2839,10 @@ def switch_model(agent, new_model, new_provider, api_key='', base_url='', api_mo old_model = agent.model old_provider = agent.provider - # ── Snapshot all fields the swap+rebuild can mutate ── - # If the rebuild raises (bad API key, network error, build_anthropic_client - # failure, etc.) we restore these atomically so the agent isn't left with a - # new model/provider name paired with the OLD client — that mismatch causes - # HTTP 400s like "claude-sonnet-4-6 is not supported on openai-codex" on the - # next turn. Callers in cli.py / gateway/run.py / tui_gateway/server.py - # catch the re-raised exception and show the user a warning; without this - # rollback the warning is misleading because the swap partially succeeded. - # Use a sentinel so we can distinguish "attribute was unset" from - # "attribute was None" and skip the restore for genuinely-missing - # attributes (tests construct bare agents via __new__ without all fields). - _MISSING = object() - _snapshot = { - name: getattr(agent, name, _MISSING) - for name in ( - "model", - "provider", - "requested_provider", - "base_url", - "api_mode", - "api_key", - "client", - "_anthropic_client", - "_anthropic_api_key", - "_anthropic_base_url", - "_is_anthropic_oauth", - "_config_context_length", - "_reasoning_echo_flag", - ) - } - # _client_kwargs is a dict — snapshot a shallow copy so mutating the - # live dict doesn't poison the rollback target. - _snapshot["_client_kwargs"] = dict(getattr(agent, "_client_kwargs", {}) or {}) - # Snapshot the credential pool reference so a failed client rebuild can - # restore the original pool (issue #52727: pool reload is part of this - # switch and must be reversible on rollback). - _snapshot["_credential_pool"] = getattr(agent, "_credential_pool", _MISSING) - _snapshot["_credential_pool_entry_id"] = getattr( - agent, "_credential_pool_entry_id", _MISSING - ) - - def _restore_snapshot() -> None: - for _name, _value in _snapshot.items(): - if _value is _MISSING: - # Attribute did not exist before the swap — don't fabricate it. - continue - try: - setattr(agent, _name, _value) - except Exception: # noqa: BLE001 - pass + # Callers retain this agent after a failed /model swap and rely on its + # previous route surviving. Include cache/pool state and preserve that + # contract until the context engine has accepted the destination budget. + restore_handoff = _snapshot_context_handoff(agent) try: # Clear the per-config context_length override so the new model's @@ -2902,7 +2927,7 @@ def _restore_snapshot() -> None: # moa://local placeholder → HTTP 404 → fallback to a reference # model. Pin chat_completions here so the primary call always goes # through MoAClient.chat.completions, matching agent_init.py. - agent.api_mode = "chat_completions" + api_mode = agent.api_mode = "chat_completions" agent.api_key = api_key or "moa-virtual-provider" agent.base_url = "moa://local" agent._client_kwargs = {} @@ -2990,102 +3015,108 @@ def _restore_snapshot() -> None: # caller's exception handler can surface a meaningful warning. The # exception is re-raised; cli.py / gateway/run.py / tui_gateway catch # it and print "Agent swap failed; change applied to next session". - _restore_snapshot() + restore_handoff() raise - # ── LM Studio: preload before probing context length ── - _sm_custom_providers = None try: - from hermes_cli.config import ( - get_compatible_custom_providers, - get_custom_provider_context_length, - load_config, - ) + # ── LM Studio: preload before probing context length ── + _sm_custom_providers = None + try: + from hermes_cli.config import ( + get_compatible_custom_providers, + get_custom_provider_context_length, + load_config, + ) - _sm_cfg = load_config() - _sm_custom_providers = get_compatible_custom_providers(_sm_cfg) - _destination_context_intent = get_custom_provider_context_length( - model=agent.model, - base_url=agent.base_url, - custom_providers=_sm_custom_providers, + _sm_cfg = load_config() + _sm_custom_providers = get_compatible_custom_providers(_sm_cfg) + _destination_context_intent = get_custom_provider_context_length( + model=agent.model, + base_url=agent.base_url, + custom_providers=_sm_custom_providers, + ) + except Exception: + _destination_context_intent = None + agent._config_context_length = _destination_context_intent + _runtime_context_length = agent._ensure_lmstudio_runtime_loaded( + _destination_context_intent ) - except Exception: - _destination_context_intent = None - agent._config_context_length = _destination_context_intent - _runtime_context_length = agent._ensure_lmstudio_runtime_loaded( - _destination_context_intent - ) - if agent._lmstudio_load_was_unverified(_runtime_context_length): - logger.warning( - "LM Studio model activation was rejected or completed without a " - "verifiable active context length during model switch; continuing " - "with configured context" + if agent._lmstudio_load_was_unverified(_runtime_context_length): + logger.warning( + "LM Studio model activation was rejected or completed without a " + "verifiable active context length during model switch; continuing " + "with configured context" + ) + _effective_context_length = agent._effective_lmstudio_context_length( + _destination_context_intent, + _runtime_context_length, ) - _effective_context_length = agent._effective_lmstudio_context_length( - _destination_context_intent, - _runtime_context_length, - ) - # ── Re-evaluate prompt caching ── - # Refresh the custom-provider snapshot from the config just loaded above - # so the per-model ``prompt_caching`` capability lookup sees the same - # live list the context-length resolution used — without this, a flag - # added to config.yaml after session start is invisible to a /model - # switch (the policy would read the stale init-time snapshot). - if _sm_custom_providers is not None: - agent._custom_providers = _sm_custom_providers - agent._use_prompt_caching, agent._use_native_cache_layout = ( - agent._anthropic_prompt_cache_policy( - provider=new_provider, - base_url=agent.base_url, - api_mode=api_mode, - model=new_model, + # ── Re-evaluate prompt caching ── + # Refresh the custom-provider snapshot from the config just loaded above + # so the per-model ``prompt_caching`` capability lookup sees the same + # live list the context-length resolution used — without this, a flag + # added to config.yaml after session start is invisible to a /model + # switch (the policy would read the stale init-time snapshot). + if _sm_custom_providers is not None: + agent._custom_providers = _sm_custom_providers + agent._use_prompt_caching, agent._use_native_cache_layout = ( + agent._anthropic_prompt_cache_policy( + provider=new_provider, + base_url=agent.base_url, + api_mode=api_mode, + model=new_model, + ) ) - ) - # ── Update context compressor ── - if hasattr(agent, "context_compressor") and agent.context_compressor: - from agent.model_metadata import get_model_context_length - if _sm_custom_providers is None: - try: - from hermes_cli.config import get_compatible_custom_providers, load_config - _sm_custom_providers = get_compatible_custom_providers(load_config()) - except Exception: - _sm_custom_providers = None - # ``agent.api_key`` may be a callable (Azure Foundry Entra ID - # token provider). ``get_model_context_length`` expects a - # string for its live-probe paths; for Foundry the context - # length normally resolves via config or static catalogs and - # never hits a probe, but coerce to empty string defensively. - _ctx_api_key = agent.api_key if isinstance(agent.api_key, str) else "" - new_context_length = get_model_context_length( - agent.model, - base_url=agent.base_url, - api_key=_ctx_api_key, - provider=agent.provider, - config_context_length=_effective_context_length, - custom_providers=_sm_custom_providers, - ) - # Forward the per-model resolved threshold (Codex gpt-5.x autoraise - # included) to engines that accept threshold_percent; the built-in - # compressor re-resolves internally and is skipped by the guard. - from agent.auxiliary_client import ( - _effective_compression_threshold_percent, - _update_compressor_model, - ) + # ── Update context compressor ── + if hasattr(agent, "context_compressor") and agent.context_compressor: + from agent.model_metadata import get_model_context_length + if _sm_custom_providers is None: + try: + from hermes_cli.config import get_compatible_custom_providers, load_config + _sm_custom_providers = get_compatible_custom_providers(load_config()) + except Exception: + _sm_custom_providers = None + # ``agent.api_key`` may be a callable (Azure Foundry Entra ID + # token provider). ``get_model_context_length`` expects a + # string for its live-probe paths; for Foundry the context + # length normally resolves via config or static catalogs and + # never hits a probe, but coerce to empty string defensively. + _ctx_api_key = agent.api_key if isinstance(agent.api_key, str) else "" + new_context_length = get_model_context_length( + agent.model, + base_url=agent.base_url, + api_key=_ctx_api_key, + provider=agent.provider, + config_context_length=_effective_context_length, + custom_providers=_sm_custom_providers, + ) + # Forward the per-model resolved threshold (Codex gpt-5.x autoraise + # included) to engines that accept threshold_percent; the built-in + # compressor re-resolves internally and is skipped by the guard. + from agent.auxiliary_client import ( + _effective_compression_threshold_percent, + _update_compressor_model, + ) - _update_compressor_model( - agent.context_compressor, - model=agent.model, - context_length=new_context_length, - base_url=agent.base_url, - api_key=agent.api_key, # context_compressor forwards to call_llm; callable preserved - provider=agent.provider, - api_mode=agent.api_mode, - threshold_percent=_effective_compression_threshold_percent( - agent.model, agent.provider - ), - ) + _update_compressor_model( + agent.context_compressor, + model=agent.model, + context_length=new_context_length, + base_url=agent.base_url, + api_key=agent.api_key, # context_compressor forwards to call_llm; callable preserved + provider=agent.provider, + api_mode=agent.api_mode, + max_tokens=getattr(agent, "max_tokens", None), + threshold_percent=_effective_compression_threshold_percent( + agent.model, agent.provider + ), + ) + + except Exception: + restore_handoff() + raise # ── Re-resolve reasoning_config from per-model override ── # The new model may have a different reasoning_effort override. Re-read diff --git a/agent/auxiliary_client.py b/agent/auxiliary_client.py index a7b26a08a703c..bc1b0167a601c 100644 --- a/agent/auxiliary_client.py +++ b/agent/auxiliary_client.py @@ -836,15 +836,17 @@ def _update_compressor_model( api_key: Any = "", provider: str = "", api_mode: str = "", + max_tokens: Optional[int] = None, threshold_percent: Optional[float] = None, ) -> None: - """Call ``update_model``, forwarding ``threshold_percent`` when supported. + """Forward route budget fields accepted by the engine's ``update_model``. The built-in ContextCompressor re-resolves its threshold internally and does not accept the kwarg; external engines (e.g. ri-context-governor) accept it so the resolved host threshold (including the Codex gpt-5.x - autoraise) reaches their trigger. A signature guard keeps every engine - on its own contract. + autoraise) reaches their trigger. The response reserve is also forwarded + when explicitly supported, including None to clear an old reservation. + A signature guard keeps every engine on its own contract. """ _kwargs = { "model": model, @@ -854,13 +856,15 @@ def _update_compressor_model( "provider": provider, "api_mode": api_mode, } - if threshold_percent is not None: - try: - import inspect + try: + import inspect - _params = inspect.signature(compressor.update_model).parameters - except (TypeError, ValueError): - _params = {} + _params = inspect.signature(compressor.update_model).parameters + except (TypeError, ValueError): + _params = {} + if "max_tokens" in _params: + _kwargs["max_tokens"] = max_tokens + if threshold_percent is not None: if "threshold_percent" in _params: _kwargs["threshold_percent"] = threshold_percent compressor.update_model(**_kwargs) diff --git a/agent/chat_completion_helpers.py b/agent/chat_completion_helpers.py index 39a72462d3306..866ceea41eb28 100644 --- a/agent/chat_completion_helpers.py +++ b/agent/chat_completion_helpers.py @@ -979,6 +979,9 @@ def _dispatch_nonstreaming_api_request(agent, api_kwargs: dict, *, make_client): raise return normalize_converse_response(raw_response) if agent.provider == "moa": + # Explicit native selection must refuse before MoA facade effects. + if agent.api_mode == "chat_completions" and _should_use_ri_pipeline(agent, api_kwargs): + return ri_pipeline_chat_completion(agent, api_kwargs) # MoA is a virtual chat-completions provider backed by the # in-process MoAClient facade. Do not rebuild a request-local # OpenAI client from the virtual runtime metadata. @@ -999,7 +1002,7 @@ def _dispatch_nonstreaming_api_request(agent, api_kwargs: dict, *, make_client): if ( agent.api_mode == "chat_completions" and api_kwargs.get("stream") is not True - and _should_use_ri_pipeline(agent) + and _should_use_ri_pipeline(agent, api_kwargs) ): logger.debug( "Using llm-pipeline transport for non-streaming chat completion call " @@ -2472,7 +2475,9 @@ def _fallback_reason_text(reason: "FailoverReason | None") -> str: return str(value or reason or "provider failure").replace("_", " ") -def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool: +def try_activate_fallback( + agent, reason: "FailoverReason | None" = None, *, _continuing_chain: bool = False +) -> bool: """Switch to the next fallback model/provider in the chain. Called when the current model is failing after retries. Swaps the @@ -2484,7 +2489,10 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool auth resolution and client construction — no duplicated provider→key mappings. """ - if reason in {FailoverReason.rate_limit, FailoverReason.billing, FailoverReason.upstream_rate_limit}: + # Restoring the primary after an inadmissible candidate must not turn the + # internal chain walk into another originating provider failure. Separate + # public calls still arm backoff even when the chain is already exhausted. + if not _continuing_chain and reason in {FailoverReason.rate_limit, FailoverReason.billing, FailoverReason.upstream_rate_limit}: # Only start cooldown when leaving the primary provider. If we're # already on a fallback and chain-switching, the primary wasn't the # source of the 429 so the cooldown should not be reset/extended. @@ -2532,11 +2540,11 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool agent._unavailable_fallback_keys = unavailable if fb_key in unavailable: logger.debug("Fallback skip: %s previously marked unavailable", fb_key) - return agent._try_activate_fallback(reason) + return try_activate_fallback(agent, reason, _continuing_chain=True) fb_provider = (fb.get("provider") or "").strip().lower() fb_model = (fb.get("model") or "").strip() if not fb_provider or not fb_model: - return agent._try_activate_fallback(reason) # skip invalid, try next + return try_activate_fallback(agent, reason, _continuing_chain=True) # skip invalid, try next local_skip_reason = _fallback_entry_unavailable_without_network(agent, fb) if local_skip_reason: @@ -2547,7 +2555,7 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool fb_model, local_skip_reason, ) - return agent._try_activate_fallback(reason) + return try_activate_fallback(agent, reason, _continuing_chain=True) # Skip entries that resolve to the same backend that just failed — # falling back to it loops the failure. Identity semantics (which axes @@ -2572,11 +2580,15 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool "as the current one (%s)", fb_provider, fb_model, current_ident.base_url or current_ident.provider, ) - return agent._try_activate_fallback(reason) + return try_activate_fallback(agent, reason, _continuing_chain=True) # Use centralized router for client construction. # raw_codex=True because the main agent needs direct responses.stream() # access for Codex providers. + from agent.agent_runtime_helpers import _snapshot_context_handoff + + restore_handoff = _snapshot_context_handoff(agent) + handoff_accepted = False try: from agent.auxiliary_client import resolve_provider_client # Pass base_url and api_key from fallback config so custom @@ -2629,7 +2641,7 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool "Fallback to %s failed: provider not configured", fb_provider) unavailable.add(fb_key) - return agent._try_activate_fallback(reason) # try next in chain + return try_activate_fallback(agent, reason, _continuing_chain=True) # try next in chain try: from hermes_cli.model_normalize import normalize_model_for_provider @@ -2844,10 +2856,12 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool api_key=getattr(agent, "api_key", ""), # callable preserved → call_llm provider=agent.provider, api_mode=agent.api_mode, + max_tokens=getattr(agent, "max_tokens", None), threshold_percent=_effective_compression_threshold_percent( agent.model, agent.provider ), ) + handoff_accepted = True # Re-resolve reasoning_config for the new fallback model (Closes #21256). # Shared chokepoint: per-model override > global reasoning_effort @@ -2906,10 +2920,12 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool _reset_stale_streak(agent) return True except Exception as e: + if not handoff_accepted: + restore_handoff() if fb_provider == "nous": unavailable.add(fb_key) logger.error("Failed to activate fallback %s: %s", fb_model, e) - return agent._try_activate_fallback(reason) # try next in chain + return try_activate_fallback(agent, reason, _continuing_chain=True) # try next in chain @@ -3413,7 +3429,7 @@ def _emit_stream_end(*, final_text: str, finished: bool, error: str | None) -> N # branch below — routing through the _interruptible_api_call method keeps the # outer loop's per-request retry/refresh seam intact. if should_use_direct_api_call(agent): - if agent.api_mode == "chat_completions" and _should_use_ri_pipeline(agent): + if agent.api_mode == "chat_completions" and _should_use_ri_pipeline(agent, api_kwargs): _nonstreaming_args = dict(api_kwargs) _nonstreaming_args["stream"] = False return agent._interruptible_api_call(_nonstreaming_args) @@ -3442,7 +3458,7 @@ def _emit_stream_end(*, final_text: str, finished: bool, error: str | None) -> N # any provider-side stream behavior. if ( agent.api_mode == "chat_completions" - and _should_use_ri_pipeline(agent) + and _should_use_ri_pipeline(agent, api_kwargs) ): _streaming_disabled_kwargs = dict(api_kwargs) _streaming_disabled_kwargs["stream"] = False diff --git a/agent/conversation_loop.py b/agent/conversation_loop.py index 62da760f752e4..46c973c93ba51 100644 --- a/agent/conversation_loop.py +++ b/agent/conversation_loop.py @@ -2493,6 +2493,10 @@ def _context_rebase_stopped_result(result, *, calls=None): if effective_system: api_messages = [{"role": "system", "content": effective_system}] + api_messages + if agent.provider == "moa": + from agent.transports.ri_llm import _should_use_ri_pipeline, RiTransportUnsupported + if _should_use_ri_pipeline(agent): + raise RiTransportUnsupported("native binding supports only ollama-launch") if moa_config: try: from agent.message_content import flatten_message_text as _flatten_mt @@ -2668,6 +2672,9 @@ def _context_rebase_stopped_result(result, *, calls=None): # request later without running the advisors a second time. _moa_prepared_request = None if agent.provider == "moa": + from agent.transports.ri_llm import _should_use_ri_pipeline, RiTransportUnsupported + if _should_use_ri_pipeline(agent): + raise RiTransportUnsupported("native binding supports only ollama-launch") _moa_completions = getattr(getattr(agent.client, "chat", None), "completions", None) if pending_moa_prepared_request is not None: _rebase_moa_request = getattr(_moa_completions, "rebase_prepared_request", None) @@ -6476,21 +6483,22 @@ def _record_physical(call, kwargs): # exists; otherwise "trying fallback..." is a lie and the # session looks like it's recovering when it's about to # abort silently (#35314, #17446). - if agent._has_pending_fallback(): + if not (classified.error_context or {}).get("native_transport_refusal") and agent._has_pending_fallback(): if classified.reason == FailoverReason.content_policy_blocked: agent._buffer_status("⚠️ Provider safety filter blocked this request — trying fallback...") elif classified.reason == FailoverReason.ssl_cert_verification: agent._buffer_status("⚠️ TLS certificate verification failed — trying fallback...") else: agent._buffer_status(f"⚠️ Non-retryable error (HTTP {status_code}) — trying fallback...") - if agent._try_activate_fallback(): - active_system_prompt = _sync_failover_system_message( - agent, api_messages, active_system_prompt) - retry_count = 0 - compression_attempts = 0 - _retry.primary_recovery_attempted = False - _retry.restart_with_rebuilt_messages = True - break + if not (classified.error_context or {}).get("native_transport_refusal"): + if agent._try_activate_fallback(): + active_system_prompt = _sync_failover_system_message( + agent, api_messages, active_system_prompt) + retry_count = 0 + compression_attempts = 0 + _retry.primary_recovery_attempted = False + _retry.restart_with_rebuilt_messages = True + break if api_kwargs is not None: agent._dump_api_request_debug( api_kwargs, reason="non_retryable_client_error", error=api_error, diff --git a/agent/error_classifier.py b/agent/error_classifier.py index 39e8ed5c5278e..9ac98d7b5d238 100644 --- a/agent/error_classifier.py +++ b/agent/error_classifier.py @@ -853,6 +853,8 @@ def classify_api_error( provider_lower = (provider or "").strip().lower() model_lower = (model or "").strip().lower() + from agent.transports.ri_llm import RiTransportUnsupported + def _result(reason: FailoverReason, **overrides) -> ClassifiedError: defaults = { "reason": reason, @@ -862,8 +864,20 @@ def _result(reason: FailoverReason, **overrides) -> ClassifiedError: "message": _extract_message(error, body), } defaults.update(overrides) + if not isinstance(error, RiTransportUnsupported) and defaults.get("error_context"): + defaults["error_context"] = dict(defaults["error_context"]) + defaults["error_context"].pop("native_transport_refusal", None) return ClassifiedError(**defaults) + # Local native selection is a request-contract refusal, not a provider error. + # Recognize the actual exception type before plugins or text heuristics. + if isinstance(error, RiTransportUnsupported): + return _result( + FailoverReason.format_error, + retryable=False, + error_context={"native_transport_refusal": True}, + ) + # ── 0. Plugin classifiers (first valid result wins) ───────────── # # Consulted BEFORE the built-in pipeline so a provider plugin can both diff --git a/agent/lsp/__init__.py b/agent/lsp/__init__.py index 7819162dd4598..b4370757b6f92 100644 --- a/agent/lsp/__init__.py +++ b/agent/lsp/__init__.py @@ -29,6 +29,8 @@ from __future__ import annotations import atexit +import hashlib +import json import logging import threading from typing import Optional @@ -37,13 +39,13 @@ logger = logging.getLogger("agent.lsp") -_service: Optional[LSPService] = None +_services: dict[str, tuple[tuple[str, str], LSPService]] = {} _atexit_registered = False _service_lock = threading.Lock() def get_service() -> Optional[LSPService]: - """Return the process-wide LSP service singleton, or None when disabled. + """Return the current profile's generation-bound service, or None. The service is created lazily on first call. ``None`` is returned when LSP is disabled in config, when no workspace can be detected, @@ -54,13 +56,41 @@ def get_service() -> Optional[LSPService]: CLI or gateway session doesn't leak pyright/gopls/etc. processes when it terminates. """ - global _service, _atexit_registered - if _service is not None: - return _service if _service.is_active() else None + global _atexit_registered + from agent.secret_scope import build_profile_env_boundary + from hermes_cli.config import load_config_readonly + from hermes_constants import get_hermes_home, get_process_hermes_home, hermes_home_key + + home = get_hermes_home() + owner = hermes_home_key(home) + try: + config = load_config_readonly() + lsp_config = config.get("lsp") or {} + if not isinstance(lsp_config, dict): + lsp_config = {} + if not bool(lsp_config.get("enabled", True)): + shutdown_service(profile_home=home) + return None + boundary = build_profile_env_boundary(get_process_hermes_home(), home) + signature = ( + boundary.target_generation, + hashlib.sha256(json.dumps(lsp_config, sort_keys=True).encode()).hexdigest(), + ) + except Exception: + shutdown_service(profile_home=home) + return None + with _service_lock: - if _service is not None: - return _service if _service.is_active() else None - _service = LSPService.create_from_config() + entry = _services.get(owner) + if entry is not None and entry[0] == signature: + service = entry[1] + return service if service.is_active() else None + if entry is not None: + _services.pop(owner) + entry[1].shutdown() + service = LSPService.create_from_config(config=config, profile_boundary=boundary) + if service is not None: + _services[owner] = (signature, service) if not _atexit_registered: # ``atexit`` handlers run in LIFO order on normal Python # exit and on SystemExit, but NOT on os._exit() or @@ -74,21 +104,22 @@ def get_service() -> Optional[LSPService]: # stdout buffers drain. atexit.register(_atexit_shutdown) _atexit_registered = True - return _service if (_service is not None and _service.is_active()) else None + return service if (service is not None and service.is_active()) else None -def shutdown_service() -> None: - """Tear down the LSP service if one was started. +def shutdown_service(*, profile_home=None) -> None: + """Tear down only the selected/current profile's service. Safe to call multiple times; safe to call when no service was created. """ - global _service + from hermes_constants import get_hermes_home, hermes_home_key + + owner = hermes_home_key(profile_home if profile_home is not None else get_hermes_home()) with _service_lock: - svc = _service - _service = None - if svc is not None: + entry = _services.pop(owner, None) + if entry is not None: try: - svc.shutdown() + entry[1].shutdown() except Exception as e: # noqa: BLE001 logger.debug("LSP shutdown error: %s", e) @@ -98,7 +129,14 @@ def _atexit_shutdown() -> None: atexit fires the user has already seen the agent's final output — a noisy shutdown line on top of that is just clutter.""" try: - shutdown_service() + with _service_lock: + services = [entry[1] for entry in _services.values()] + _services.clear() + for service in services: + try: + service.shutdown() + except Exception: + logger.debug("atexit LSP service shutdown failed", exc_info=True) except Exception as e: # noqa: BLE001 logger.debug("atexit LSP shutdown failed: %s", e) diff --git a/agent/lsp/client.py b/agent/lsp/client.py index 36c272102761f..0dd7d7bcbe3be 100644 --- a/agent/lsp/client.py +++ b/agent/lsp/client.py @@ -194,7 +194,17 @@ def __init__( cwd: Optional[str] = None, initialization_options: Optional[Dict[str, Any]] = None, seed_diagnostics_on_first_push: bool = False, + profile_boundary=None, ) -> None: + from agent.secret_scope import build_profile_env_boundary + from hermes_constants import get_hermes_home, get_process_hermes_home + + boundary = profile_boundary or build_profile_env_boundary( + get_process_hermes_home(), get_hermes_home(), + ) + self._profile_home = boundary.target_home + self._source_home = boundary.source_home + self._profile_generation = boundary.target_generation self.server_id = server_id self.workspace_root = workspace_root self._command = list(command) @@ -305,9 +315,20 @@ def _win_wrap_cmd(cmd: List[str]) -> List[str]: return cmd async def _spawn(self) -> None: - env = dict(os.environ) - if self._env: - env.update(self._env) + from agent.secret_scope import build_profile_env_boundary + from tools.environments.local import build_subprocess_env, hermes_subprocess_env + + boundary = build_profile_env_boundary(self._source_home, self._profile_home) + if boundary.target_generation != self._profile_generation: + raise LSPProtocolError("LSP profile authority changed; refusing spawn") + # First remove launch-profile authority using the non-model child + # policy. Then admit explicit target-authored server env through the + # existing non-model filter; it must not become source-profile input. + base = hermes_subprocess_env(profile_boundary=boundary) + env = build_subprocess_env( + base=base, extra=self._env, profile_home=self._profile_home, + source_profile_home=self._profile_home, enforce_profile_boundary=True, + ) cmd = self._command if sys.platform == "win32": diff --git a/agent/lsp/install.py b/agent/lsp/install.py index fc9bea59307b1..a0cde4a7b7dca 100644 --- a/agent/lsp/install.py +++ b/agent/lsp/install.py @@ -112,8 +112,8 @@ } -_install_locks: Dict[str, threading.Lock] = {} -_install_results: Dict[str, Optional[str]] = {} +_install_locks: Dict[tuple[str, str], threading.Lock] = {} +_install_results: Dict[tuple[str, str], Optional[str]] = {} _install_lock_meta = threading.Lock() _WINDOWS_WRAPPER_SUFFIXES = (".cmd", ".exe", ".bat") @@ -161,12 +161,12 @@ def _existing_binary(name: str) -> Optional[str]: return None -def _get_lock(pkg: str) -> threading.Lock: +def _get_lock(key: tuple[str, str]) -> threading.Lock: with _install_lock_meta: - lock = _install_locks.get(pkg) + lock = _install_locks.get(key) if lock is None: lock = threading.Lock() - _install_locks[pkg] = lock + _install_locks[key] = lock return lock @@ -177,7 +177,7 @@ def try_install(pkg: str, strategy: str = "auto") -> Optional[str]: ``manual``/``off`` mode, this function only probes for an existing binary and returns ``None`` if not found. - The install is cached per-package — a second call returns the + The install is cached per-profile/package — a second call returns the same path (or ``None``) without reinstalling. Concurrent calls are serialized. """ @@ -188,16 +188,19 @@ def try_install(pkg: str, strategy: str = "auto") -> Optional[str]: bin_name = recipe.get("bin", pkg) return _existing_binary(bin_name) - if pkg in _install_results: - return _install_results[pkg] + from hermes_constants import get_hermes_home, hermes_home_key - lock = _get_lock(pkg) + key = (hermes_home_key(get_hermes_home()), pkg) + if key in _install_results: + return _install_results[key] + + lock = _get_lock(key) with lock: # Double-check after acquiring lock. - if pkg in _install_results: - return _install_results[pkg] + if key in _install_results: + return _install_results[key] result = _do_install(pkg) - _install_results[pkg] = result + _install_results[key] = result return result diff --git a/agent/lsp/manager.py b/agent/lsp/manager.py index 7ba1b914f74c3..5a77b71ba2dfc 100644 --- a/agent/lsp/manager.py +++ b/agent/lsp/manager.py @@ -25,7 +25,7 @@ ``beforeFileEdited`` / ``getNewDiagnostics`` pattern, except wired to the local LSP layer instead of MCP IDE RPC. -The service is **off by default** — call :meth:`is_active` to check +The service is **enabled by default** — call :meth:`is_active` to check whether it's actually doing anything. When LSP is disabled in config, when no git workspace can be detected, when all configured servers are missing binaries and auto-install is off, ``is_active`` @@ -156,7 +156,17 @@ def __init__( init_overrides: Optional[Dict[str, Dict[str, Any]]] = None, disabled_servers: Optional[List[str]] = None, idle_timeout: float = DEFAULT_IDLE_TIMEOUT, + profile_boundary=None, ) -> None: + from agent.secret_scope import build_profile_env_boundary + from hermes_constants import get_hermes_home, get_process_hermes_home + + boundary = profile_boundary or build_profile_env_boundary( + get_process_hermes_home(), get_hermes_home(), + ) + self._profile_home = boundary.target_home + self._source_home = boundary.source_home + self._profile_generation = boundary.target_generation self._enabled = enabled self._wait_mode = wait_mode if wait_mode in {"document", "full"} else "document" self._wait_timeout = wait_timeout @@ -189,7 +199,7 @@ def __init__( self._loop.run(self._start_idle_reaper(), timeout=2.0) @classmethod - def create_from_config(cls) -> Optional["LSPService"]: + def create_from_config(cls, *, config=None, profile_boundary=None) -> Optional["LSPService"]: """Build a service from ``hermes_cli.config`` settings. Returns ``None`` if the config can't be loaded. The service @@ -197,7 +207,7 @@ def create_from_config(cls) -> Optional["LSPService"]: """ try: from hermes_cli.config import load_config_readonly - cfg = load_config_readonly() + cfg = load_config_readonly() if config is None else config except Exception as e: # noqa: BLE001 logger.debug("LSP config load failed: %s", e) return None @@ -251,6 +261,7 @@ def create_from_config(cls) -> Optional["LSPService"]: init_overrides=init_overrides, disabled_servers=disabled, idle_timeout=idle_timeout, + profile_boundary=profile_boundary, ) # ------------------------------------------------------------------ @@ -471,6 +482,7 @@ def shutdown(self) -> None: except Exception as e: # noqa: BLE001 logger.debug("LSP shutdown error: %s", e) self._loop.stop() + self._enabled = False clear_cache() # ------------------------------------------------------------------ @@ -533,6 +545,17 @@ async def _current_diags_async(self, file_path: str) -> List[Dict[str, Any]]: return list(client.diagnostics_for(file_path, fresh_only=True)) async def _get_or_spawn(self, file_path: str) -> Optional[LSPClient]: + if not self._enabled: + return None + from agent.secret_scope import build_profile_env_boundary + from hermes_constants import set_hermes_home_override, reset_hermes_home_override + + try: + boundary = build_profile_env_boundary(self._source_home, self._profile_home) + except Exception: + return None + if boundary.target_generation != self._profile_generation: + return None srv = find_server_for_file(file_path) if srv is None: return None @@ -579,7 +602,11 @@ async def _get_or_spawn(self, file_path: str) -> Optional[LSPClient]: env_overrides=self._env_overrides, init_overrides=self._init_overrides, ) - spec = srv.build_spawn(per_server_root, ctx) + home_token = set_hermes_home_override(self._profile_home) + try: + spec = srv.build_spawn(per_server_root, ctx) + finally: + reset_hermes_home_override(home_token) if spec is None: # ``build_spawn`` returns None when the binary can't be # located (auto-install disabled, manual-only server, @@ -597,6 +624,7 @@ async def _get_or_spawn(self, file_path: str) -> Optional[LSPClient]: cwd=spec.cwd, initialization_options=spec.initialization_options, seed_diagnostics_on_first_push=spec.seed_diagnostics_on_first_push or srv.seed_first_push, + profile_boundary=boundary, ) try: await client.start() diff --git a/agent/secret_scope.py b/agent/secret_scope.py index 81b0652d0ed95..d521d705f96fa 100644 --- a/agent/secret_scope.py +++ b/agent/secret_scope.py @@ -662,6 +662,15 @@ def build_profile_secret_scope( continue secrets[key] = value + # An explicitly admitted root underlay belongs to this profile's grant + # history. Keep that provenance if inheritance is later revoked while an + # old value remains in the launch environment. Never classify unrelated + # shell exports or unadmitted root namespaces as profile authority. + if inherited: + history_key = str(home.resolve()) + with _PROFILE_OWNED_NAME_HISTORY_LOCK: + _PROFILE_OWNED_NAME_HISTORY.setdefault(history_key, set()).update(inherited) + return _immutable_scope( secrets, profile_home=home, @@ -773,8 +782,9 @@ def get_profile_owned_secret_names( ) -> frozenset[str]: """Return exact secret names owned by one profile, without reading values. - The profile's dotenv files and the external-source provenance snapshot are - the ownership sources. Ordinary shell exports are intentionally excluded: + The profile's dotenv files, external-source provenance snapshot and + explicitly admitted root underlay are the ownership sources. + Ordinary shell exports are intentionally excluded: they are user/process state, not profile-owned credentials. """ home = Path(hermes_home) @@ -796,6 +806,12 @@ def get_profile_owned_secret_names( ) observed_names = set(op_snapshot.data) observed_names.update(env_snapshot.data) + observed_names.update( + _root_profile_fallback_secrets( + home, + fail_closed_external=fail_closed_external, + ) + ) observed_names.update( _profile_external_secret_values( home, diff --git a/agent/transports/chat_completions.py b/agent/transports/chat_completions.py index bec0f9a82b0a8..39d8f473d19ac 100644 --- a/agent/transports/chat_completions.py +++ b/agent/transports/chat_completions.py @@ -891,7 +891,9 @@ def normalize_response(self, response: Any, **kwargs) -> NormalizedResponse: _fr = getattr(choice, "finish_reason", None) if isinstance(_fr, int): _fr = str(_fr) - finish_reason = _fr or "stop" + from agent.transports.ri_llm import RiCompletionResponse + # Only the internal native result attests that finish metadata is absent. + finish_reason = _fr if isinstance(response, RiCompletionResponse) else (_fr or "stop") tool_calls = None message_tool_calls = getattr(msg, "tool_calls", None) diff --git a/agent/transports/ri_llm.py b/agent/transports/ri_llm.py index 9042ef347daeb..c68f0c9eb0de7 100644 --- a/agent/transports/ri_llm.py +++ b/agent/transports/ri_llm.py @@ -107,25 +107,122 @@ def __repr__(self) -> str: return f"RiPipeline(url={self.url}, model={self.model}, {status})" -# ── Phase 2: RiChatCompletionsTransport (universal) ─────────────── +# ── Phase 2: RiChatCompletionsTransport ───────────────────────── # -# Plugs into the chat_completion_helpers dispatch. Active by default when -# the native extension is available and not runtime-disabled. Provider- -# agnostic — works for any OpenAI-compatible provider. -# OpenAI-compatible provider. Set HERMES_RI_PIPELINE_PROVIDERS to -# a comma-separated whitelist to restrict (e.g. 'ollama-launch,deepseek'). -# Falls through to the stock httpx/openai path on any error. - -import json as _json +# Plugs into chat_completion_helpers for request shapes that the native +# Ollama binding can preserve. Default-incompatible requests retain their SDK +# route. Explicit native selection refuses unsupported requests before effects. + +import math as _math import os as _os from types import SimpleNamespace as _SimpleNamespace +from urllib.parse import urlsplit as _urlsplit + + +class RiTransportUnsupported(ValueError): + """The selected native binding cannot preserve this request's contract.""" + + def __init__(self, reason: str): + self.code = "RI_PIPELINE_REQUEST_UNSUPPORTED" + self.reason = reason + super().__init__(f"{self.code}: {reason}") + + +class RiCompletionResponse(_SimpleNamespace): + """Internal raw-text binding result with no reported finish or usage.""" + + +def configure_ri_pipeline(agent, config: dict) -> None: + """Hydrate transport selection from the slash command's canonical owner.""" + from hermes_cli.llm_pipeline_switch import get_current_state + + enabled, providers = get_current_state(config) + agent._ri_pipeline_enabled = enabled + agent._ri_pipeline_providers = providers + section = config.get("agent", {}) if isinstance(config, dict) else {} + selection = section.get("llm_pipeline", {}) if isinstance(section, dict) else {} + agent._ri_pipeline_explicit = bool(providers) or ( + isinstance(selection, dict) and selection.get("enabled") is True + ) + + +def _native_text_request(agent, api_kwargs: dict): + """Qualify the prompt-only, unauthenticated Ollama binding before effects. + + The pinned Python API accepts one prompt, one system string and LlmConfig. + It has no chat history, tools, media, headers, auth or reported-usage API. + """ + if str(getattr(agent, "provider", "")).strip().lower() != "ollama-launch": + raise RiTransportUnsupported("native binding supports only ollama-launch") + base_url = getattr(agent, "base_url", None) + if not isinstance(base_url, str): + raise RiTransportUnsupported("missing endpoint") + try: + endpoint = _urlsplit(base_url) + endpoint.port # Validate the port without making a connection. + valid_endpoint = ( + endpoint.scheme in {"http", "https"} and endpoint.hostname + and endpoint.username is None and endpoint.password is None + and not endpoint.query and not endpoint.fragment + and endpoint.path.rstrip("/") in {"", "/v1", "/api"} + ) + except ValueError: + valid_endpoint = False + if not valid_endpoint: + raise RiTransportUnsupported("endpoint cannot be preserved by native binding") + # Never invoke a credential supplier during selection/qualification. + api_key = getattr(agent, "api_key", None) + if api_key not in (None, "", "no-key-required", "ollama"): + raise RiTransportUnsupported("authenticated route requires a per-call native auth API") + client_options = getattr(agent, "_client_kwargs", {}) + if type(client_options) is not dict or set(client_options) - {"api_key", "base_url"}: + raise RiTransportUnsupported("canonical client options unavailable in native binding") + if "base_url" in client_options and client_options["base_url"] != base_url: + raise RiTransportUnsupported("canonical endpoint differs from native endpoint") + if "api_key" in client_options and client_options["api_key"] != api_key: + raise RiTransportUnsupported("canonical credential differs from native credential") + if type(api_kwargs) is not dict: + raise RiTransportUnsupported("invalid request") + allowed = {"model", "messages", "temperature", "max_tokens", "stream"} + if set(api_kwargs) - allowed: + raise RiTransportUnsupported("request fields unavailable in native binding") + messages = api_kwargs.get("messages") + if not isinstance(messages, list) or len(messages) not in {1, 2}: + raise RiTransportUnsupported("chat history unavailable in native binding") + roles = [item.get("role") if isinstance(item, dict) else None for item in messages] + if roles not in (["user"], ["system", "user"]): + raise RiTransportUnsupported("message roles unavailable in native binding") + if any(set(item) != {"role", "content"} or not isinstance(item["content"], str) for item in messages): + raise RiTransportUnsupported("message metadata or media unavailable in native binding") + system = messages[0]["content"] if len(messages) == 2 else None + if system == "": + raise RiTransportUnsupported("empty system role unavailable in native binding") + if "{input}" in messages[-1]["content"]: + raise RiTransportUnsupported("native prompt template would alter literal input placeholder") + model = api_kwargs.get("model", getattr(agent, "model", None)) + if not isinstance(model, str) or not model.strip(): + raise RiTransportUnsupported("invalid model") + if "temperature" not in api_kwargs or "max_tokens" not in api_kwargs: + raise RiTransportUnsupported("omitted generation defaults unavailable in native binding") + temperature = api_kwargs["temperature"] + if (isinstance(temperature, bool) or not isinstance(temperature, (int, float)) + or not _math.isfinite(temperature)): + raise RiTransportUnsupported("invalid temperature") + max_tokens = api_kwargs["max_tokens"] + if type(max_tokens) is not int or not 0 < max_tokens <= 2**32 - 1: + raise RiTransportUnsupported("invalid max_tokens") + if type(api_kwargs.get("stream", False)) is not bool: + raise RiTransportUnsupported("invalid stream setting") + return base_url, model, messages[-1]["content"], system, RiLlmConfig( + temperature=temperature, max_tokens=max_tokens, + ) -def _should_use_ri_pipeline(agent) -> bool: +def _should_use_ri_pipeline(agent, api_kwargs: dict | None = None) -> bool: """Return True when the RiPipeline fast path should be used. - Active by default when the native extension is available. - Provider-agnostic — works for any OpenAI-compatible provider. + Default availability is restricted to representable Ollama requests. + An explicit native selection must pass the typed dispatch qualification. Set HERMES_RI_PIPELINE=0 to disable, or HERMES_RI_PIPELINE_PROVIDERS to a comma-separated whitelist (e.g. 'ollama-launch,deepseek'). If no env whitelist is set, agent._ri_pipeline_enabled and @@ -139,17 +236,33 @@ def _should_use_ri_pipeline(agent) -> bool: if not bool(getattr(agent, "_ri_pipeline_enabled", True)): return False + explicit = bool(getattr(agent, "_ri_pipeline_explicit", False)) + explicit = explicit or _os.environ.get("HERMES_RI_PIPELINE") == "1" + provider = str(getattr(agent, "provider", "")).strip().lower() whitelist = _os.environ.get("HERMES_RI_PIPELINE_PROVIDERS") if whitelist: allowed = _normalize_ri_pipeline_provider_list(whitelist) - return str(getattr(agent, "provider", "")).strip().lower() in allowed - - config_whitelist = _normalize_ri_pipeline_provider_list( - getattr(agent, "_ri_pipeline_providers", []) - ) - if config_whitelist: - return str(getattr(agent, "provider", "")).strip().lower() in config_whitelist - + if provider not in allowed: + return False + explicit = True + else: + config_whitelist = _normalize_ri_pipeline_provider_list( + getattr(agent, "_ri_pipeline_providers", []) + ) + if config_whitelist: + if provider not in config_whitelist: + return False + explicit = True + # Default availability never selects an incompatible wire protocol. + if not explicit and provider != "ollama-launch": + return False + if api_kwargs is not None: + try: + _native_text_request(agent, api_kwargs) + except RiTransportUnsupported: + # Explicit selection reaches the typed refusal in dispatch. Ordinary + # unsupported requests retain their existing same-provider SDK path. + return explicit return True @@ -178,119 +291,28 @@ def ri_pipeline_chat_completion(agent, api_kwargs: dict): and returns an OpenAI-compatible response namespace so the rest of the agent loop is unchanged. """ - model = api_kwargs.get("model", agent.model) - messages = api_kwargs.get("messages", []) - base_url = getattr(agent, "base_url", "http://localhost:11434/v1") - - # Inject API key into environment for the Rust pipeline. - # The Rust OpenAiBackend reads OPENAI_API_KEY from the environment. - _prev_key = _os.environ.get("OPENAI_API_KEY") - agent_api_key = getattr(agent, "api_key", None) - if callable(agent_api_key): - try: - agent_api_key = agent_api_key() - except Exception: - agent_api_key = None - if agent_api_key and isinstance(agent_api_key, str) and agent_api_key.strip(): - _os.environ["OPENAI_API_KEY"] = agent_api_key - try: - return _ri_chat_completion_impl( - agent, api_kwargs, model, messages, base_url - ) - finally: - if _prev_key is not None: - _os.environ["OPENAI_API_KEY"] = _prev_key - elif "OPENAI_API_KEY" in _os.environ: - del _os.environ["OPENAI_API_KEY"] - - -def _ri_chat_completion_impl(agent, api_kwargs, model, messages, base_url): - # Build a text prompt from the messages list (basic: system + user + assistant) - system_prompt = "" - prompt_parts = [] - for msg in messages: - role = msg.get("role", "user") - content = msg.get("content", "") - if isinstance(content, list): - # Multimodal content: extract text parts only - text_parts = [p.get("text", "") for p in content if isinstance(p, dict) and p.get("type") == "text"] - content = " ".join(text_parts) - if not isinstance(content, str): - content = str(content) - if role == "system": - system_prompt = content - elif role == "user": - prompt_parts.append(f"User: {content}") - elif role == "assistant": - prompt_parts.append(f"Assistant: {content}") - elif role == "tool": - prompt_parts.append(f"Tool output: {content}") - - full_prompt = "\n".join(prompt_parts) - - # Try to extract tool schemas for structured output - tools = api_kwargs.get("tools") or api_kwargs.get("functions") - use_json_mode = bool(tools and not api_kwargs.get("stream")) - - pipe = RiPipeline(base_url, model) + base_url, model, prompt, system, config = _native_text_request(agent, api_kwargs) + return _ri_chat_completion_impl(model, prompt, system, base_url, config) + + +def _ri_chat_completion_impl(model, prompt, system, base_url, config): + pipe = RiPipeline(base_url, model, config=config) if not pipe.available: raise RuntimeError("RiPipeline native extension not available") - if use_json_mode: - # Tool-calling: use call_structured with a JSON schema for tool_choice - tool_names = [t.get("function", {}).get("name", "tool") for t in tools] - json_schema = _json.dumps({ - "type": "object", - "properties": { - "tool": {"type": "string", "enum": tool_names}, - "arguments": {"type": "object"}, - }, - "required": ["tool", "arguments"], - }) - raw = pipe.call_structured(full_prompt, json_schema, system=system_prompt or None) - else: - raw = pipe.call(full_prompt, system=system_prompt or None) - - # Parse tool calls if present (basic JSON extraction) - tool_calls = None - content = raw - if use_json_mode and raw.strip(): - try: - parsed = _json.loads(raw) - tool_name = parsed.get("tool", "") - tool_args = parsed.get("arguments", {}) - if tool_name: - import uuid as _uuid - tool_calls = [{ - "id": f"call_{_uuid.uuid4().hex[:8]}", - "type": "function", - "function": { - "name": tool_name, - "arguments": _json.dumps(tool_args), - }, - }] - content = None # Tool call — no text content - except _json.JSONDecodeError: - pass - - # Approximate token counts (rough character-based estimate) - prompt_chars = len(full_prompt) + len(system_prompt) - completion_chars = len(raw) + raw = pipe.call(prompt, system=system) # Build response namespace matching OpenAI shape message = _SimpleNamespace( role="assistant", - content=content, - tool_calls=[_SimpleNamespace(**tc) for tc in tool_calls] if tool_calls else None, + content=raw, + tool_calls=None, ) choice = _SimpleNamespace( index=0, message=message, - finish_reason="tool_calls" if tool_calls else "stop", - ) - usage = _SimpleNamespace( - prompt_tokens=max(1, prompt_chars // 4), - completion_tokens=max(1, completion_chars // 4), - total_tokens=max(2, (prompt_chars + completion_chars) // 4), + finish_reason=None, ) - return _SimpleNamespace(choices=[choice], usage=usage, model=model) + # The binding returns raw text only. Estimates cannot attest provider usage, + # billing, compaction effectiveness or a known-fitting request baseline. + return RiCompletionResponse(choices=[choice], usage=None, model=model) diff --git a/agent/verification_evidence.py b/agent/verification_evidence.py index de08a6e87db91..64124b2d039ee 100644 --- a/agent/verification_evidence.py +++ b/agent/verification_evidence.py @@ -311,6 +311,17 @@ def _equivalent_needles(needle: list[str]) -> list[list[str]]: return candidates +def _verification_cwd_is_attributable( + segments: list[_ShellSegment], match_index: int +) -> bool: + """Only the first segment is bound to the caller's pre-command cwd. + + An earlier arbitrary shell command may change cwd. A final cwd or a + successful shell status cannot establish where a later verifier ran. + """ + return bool(segments) and match_index == 0 + + def _find_canonical_match( command: str, canonical_commands: list[str], @@ -328,6 +339,7 @@ def _find_canonical_match( for candidate in _equivalent_needles(needle): if ( candidate_tokens[:len(candidate)] == candidate + and _verification_cwd_is_attributable(segments, index) and _exit_status_is_attributable(segments, index, exit_code) ): return canonical, candidate_tokens[len(candidate):] @@ -361,8 +373,54 @@ def _looks_like_target(arg: str) -> bool: ) -def _scope_for_args(args: list[str]) -> str: - return "targeted" if any(_looks_like_target(arg) for arg in args) else "full" +_NEUTRAL_VERIFICATION_OPTIONS = { + "-q": False, "-v": False, "-s": False, "--quiet": False, + "--verbose": False, + "--no-header": False, "--no-summary": False, "--disable-warnings": False, + "-n": True, "--numprocesses": True, "-j": True, "--jobs": True, + "--file-timeout": True, "--file-retries": True, "--color": True, + "--tb": True, "--show-capture": True, "--junitxml": True, + "-r": True, +} + + +def _scope_for_args( + args: list[str], *, canonical: str = "", prefix_args: list[str] | None = None +) -> str: + """Label selection intent conservatively; never prove executed coverage.""" + for prefix in prefix_args or []: + name, separator, _ = prefix.partition("=") + if separator and name not in { + "CI", "HERMES_TEST_WORKERS", "HERMES_TEST_FILE_TIMEOUT", "HERMES_TEST_FILE_RETRIES" + }: + return "targeted" + # Option meanings belong to their runner, not to a shared spelling. + # In particular Go's -run is not pytest's attached reporting option -r. + known_options = canonical in {"pytest", "scripts/run_tests.sh"} + index = 0 + while index < len(args): + arg = args[index] + if not arg or arg == "--": + index += 1 + continue + if _looks_like_target(arg): + return "targeted" + option, separator, _ = arg.partition("=") + if known_options and option in _NEUTRAL_VERIFICATION_OPTIONS: + if _NEUTRAL_VERIFICATION_OPTIONS[option] and not separator: + if index + 1 >= len(args) or args[index + 1].startswith("-"): + return "targeted" + index += 1 + elif known_options and len(arg) > 1 and arg.startswith("-") and set(arg[1:]) <= {"q", "v", "s"}: + pass + elif known_options and any(arg.startswith(short) and len(arg) > 2 for short in ("-n", "-j", "-r")): + pass + else: + # Explicit selectors and unknown arguments cannot establish an + # unrestricted suite. This includes attached -k/-m and --slice. + return "targeted" + index += 1 + return "full" def _is_under_temp_dir(token: str) -> bool: @@ -431,8 +489,10 @@ def _find_ad_hoc_match( segments = _split_shell_segments(command, posix=posix) for index, segment in enumerate(segments): trailing_args = _ad_hoc_script_args(segment.tokens, root) - if trailing_args is not None and _exit_status_is_attributable( - segments, index, exit_code + if ( + trailing_args is not None + and _verification_cwd_is_attributable(segments, index) + and _exit_status_is_attributable(segments, index, exit_code) ): return trailing_args return None @@ -546,11 +606,17 @@ def classify_verification_command( return None canonical, trailing_args = match + # Keep explicit selection assignments that matching strips from the + # first segment. They are command-bound inputs, unlike inherited env. + first_tokens = _split_shell_segments(command)[0].tokens + prefix_length = len(first_tokens) - len(_strip_command_prefix(first_tokens)) return VerificationEvidence( command=command, canonical_command=canonical, kind="ad_hoc" if is_ad_hoc else _kind_for_command(canonical), - scope="targeted" if is_ad_hoc else _scope_for_args(trailing_args), + scope="targeted" if is_ad_hoc else _scope_for_args( + trailing_args, canonical=canonical, prefix_args=first_tokens[:prefix_length] + ), status="passed" if int(exit_code) == 0 else "failed", exit_code=int(exit_code), cwd=str(Path(cwd or ".").resolve()), diff --git a/apps/desktop/e2e/boot.spec.ts b/apps/desktop/e2e/boot.spec.ts index 19f2d81fa4530..6397a83faa368 100644 --- a/apps/desktop/e2e/boot.spec.ts +++ b/apps/desktop/e2e/boot.spec.ts @@ -32,9 +32,8 @@ test.afterAll(async () => { }) test.describe('dev-mode boot with mock backend', () => { - test('window opens with Hermes title', async () => { - const title = await fixture!.page.title() - expect(title).toContain('Hermes') + test('window opens with Ares title', async () => { + await expect(fixture!.page).toHaveTitle('Ares') }) test('renderer mounts and shows DOM content', async () => { diff --git a/apps/desktop/e2e/fixtures.ts b/apps/desktop/e2e/fixtures.ts index 787be421886b7..81be78c6fc1de 100644 --- a/apps/desktop/e2e/fixtures.ts +++ b/apps/desktop/e2e/fixtures.ts @@ -668,9 +668,13 @@ export async function waitForAppReady(fixture: MockBackendFixture | NoProviderFi return w ? w.isVisible() : false }).catch(() => false) - if (visible) {break} + if (visible) { + return + } await page.waitForTimeout(500) } + + throw new Error(`Electron window did not become visible within ${timeoutMs}ms`) } } diff --git a/apps/desktop/package.json b/apps/desktop/package.json index b1a4dd792e46c..73c1b821c6e58 100644 --- a/apps/desktop/package.json +++ b/apps/desktop/package.json @@ -157,7 +157,7 @@ "bippy": "0.5.43", "concurrently": "10.0.4", "cross-env": "10.1.0", - "electron": "41.10.3", + "electron": "41.10.6", "electron-builder": "^26.8.1", "esbuild": "^0.28.1", "eslint": "^9.39.4", @@ -176,7 +176,7 @@ "wait-on": "^9.0.5" }, "build": { - "electronVersion": "41.10.3", + "electronVersion": "41.10.6", "appId": "com.recursiveintell.ares", "productName": "Ares", "executableName": "Ares", diff --git a/apps/desktop/src/api/models.ts b/apps/desktop/src/api/models.ts index b484c3f2faa5e..2d12ef561817b 100644 --- a/apps/desktop/src/api/models.ts +++ b/apps/desktop/src/api/models.ts @@ -50,7 +50,10 @@ export function getGlobalModelOptions( return window.hermesDesktop.api({ ...capabilityScoped(profile), - ...(profile && typeof profile === 'object' ? { connectionId: profile.connectionId || 'local' } : {}), + // Explicit null preserves the captured legacy route, which may be remote. + ...(profile && typeof profile === 'object' && profile.connectionId !== null + ? { connectionId: profile.connectionId || 'local' } + : {}), path: params.size > 0 ? `/api/model/options?${params.toString()}` : '/api/model/options', timeoutMs: STARTUP_REQUEST_TIMEOUT_MS }) diff --git a/apps/desktop/src/app/hooks/use-composer-model-owner.ts b/apps/desktop/src/app/hooks/use-composer-model-owner.ts new file mode 100644 index 0000000000000..ede992fe406c4 --- /dev/null +++ b/apps/desktop/src/app/hooks/use-composer-model-owner.ts @@ -0,0 +1,98 @@ +import { useStore } from '@nanostores/react' +import { useCallback, useLayoutEffect, useRef, useState } from 'react' + +import { + captureDraftComposerOwner, + captureModelRequestOwner, + composerOwnerKey, + type ComposerSelectionOwner +} from '@/app/session/hooks/composer-model-selection-owner' +import { $activeGatewayProfile, $newChatConnectionId, $newChatProfile, $newChatRoute } from '@/store/profile' +import { $connection } from '@/store/session' + +// These atoms publish rehomes; canonical capture still owns route precedence. +// Do not infer model authority from the mirrored model/provider atoms here. +function useModelOwnerChanges() { + useStore($connection) + useStore($activeGatewayProfile) + useStore($newChatConnectionId) + useStore($newChatProfile) + useStore($newChatRoute) +} + +export function useModelRequestOwner(scopeProfile?: string): ComposerSelectionOwner { + // Canonical capture reads imperative stores; it must run on every rehome. + 'use no memo' + + useModelOwnerChanges() + + return captureModelRequestOwner(scopeProfile) +} + +export function useDraftComposerOwner(): ComposerSelectionOwner { + 'use no memo' + + useModelOwnerChanges() + + return captureDraftComposerOwner() +} + +/** Form lifecycle only: a batched round trip still needs a fresh form. */ +export function useModelFormKey(owner: ComposerSelectionOwner, scopeProfile?: string): string { + const [revision, setRevision] = useState(0) + useLayoutEffect(() => { + let previousKey = composerOwnerKey(captureModelRequestOwner(scopeProfile)) + + const rehome = () => { + const nextKey = composerOwnerKey(captureModelRequestOwner(scopeProfile)) + + if (nextKey !== previousKey) { + previousKey = nextKey + setRevision(value => value + 1) + } + } + + const unlisten = [$connection, $activeGatewayProfile, $newChatConnectionId, $newChatProfile, $newChatRoute] + .map(store => store.listen(rehome)) + + return () => { unlisten.forEach(stop => stop()) } + }, [scopeProfile]) + + return JSON.stringify([composerOwnerKey(owner), revision]) +} + +/** A remounted form cannot regain permission when the user returns A→B→A. */ +export function useModelOwnerIsCurrent(owner: ComposerSelectionOwner, scopeProfile?: string): () => boolean { + const ownerKey = composerOwnerKey(owner) + const mounted = useRef(true) + useLayoutEffect(() => { + mounted.current = true + + const invalidate = () => { + if (composerOwnerKey(captureModelRequestOwner(scopeProfile)) !== ownerKey) { + mounted.current = false + } + } + + // Observe the transition itself, even when React batches A→B→A into one + // paint. This lease cannot become valid again before a fresh form mount. + const unlisten = [$connection, $activeGatewayProfile, $newChatConnectionId, $newChatProfile, $newChatRoute] + .map(store => store.listen(invalidate)) + + return () => { + mounted.current = false + unlisten.forEach(stop => stop()) + } + }, [ownerKey, scopeProfile]) + + return useCallback( + () => mounted.current && composerOwnerKey(captureModelRequestOwner(scopeProfile)) === ownerKey, + [ownerKey, scopeProfile] + ) +} + +export function requireCurrentModelOwner(isCurrent: () => boolean): void { + if (!isCurrent()) { + throw new Error('Model settings target changed. Reopen the selection before saving.') + } +} diff --git a/apps/desktop/src/app/model-picker-overlay.test.tsx b/apps/desktop/src/app/model-picker-overlay.test.tsx new file mode 100644 index 0000000000000..b0c3448bdb10f --- /dev/null +++ b/apps/desktop/src/app/model-picker-overlay.test.tsx @@ -0,0 +1,132 @@ +import { QueryClient, QueryClientProvider } from '@tanstack/react-query' +import { act, cleanup, render, screen, waitFor } from '@testing-library/react' +import { afterEach, beforeAll, beforeEach, expect, it, vi } from 'vitest' + +import { setApiRequestConnection, setApiRequestProfile } from '@/api/client' +import type * as Gateway from '@/store/gateway' +import { $newChatConnectionId, $newChatProfile, $newChatRoute } from '@/store/profile' +import { + $activeSessionId, $currentModel, $currentProvider, $gatewayState, $modelPickerOpen, $sessions, + _resetComposerModelSelectionsForTests, _resetSessionOwnerHintsForTests, captureComposerModelSelection, + recordComposerModelSelection, setComposerModelSelectionOwner, setSessionOwnerHint +} from '@/store/session' +import { knownOwnerForSession } from '@/store/session-states' + +import { ModelPickerOverlay } from './model-picker-overlay' + +const calls = vi.hoisted(() => ({ agent: vi.fn(), profile: vi.fn(), rest: vi.fn() })) +vi.mock('@/store/gateway', async original => ({ + ...await original(), + requestGatewayForAgent: (...args: unknown[]) => calls.agent(...args), + requestGatewayForProfile: (...args: unknown[]) => calls.profile(...args) +})) +vi.mock('@/hermes', () => ({ getGlobalModelOptions: (...args: unknown[]) => calls.rest(...args), setApiRequestProfile: vi.fn() })) +beforeAll(() => { + Element.prototype.scrollIntoView = vi.fn() + vi.stubGlobal('ResizeObserver', class { observe() {} unobserve() {} disconnect() {} }) +}) +const owner = { connectionId: 'source-b', profile: 'same-name', targetProfile: 'backend-b' } +const options = { model: 'catalog-b', provider: 'custom:b', providers: [{ name: 'B', slug: 'custom:b', models: ['catalog-b', 'pinned-b'] }] } + +beforeEach(() => { + vi.clearAllMocks() + calls.agent.mockReset().mockResolvedValue(options) + calls.profile.mockReset().mockResolvedValue(options) + calls.rest.mockReset().mockResolvedValue(options) + _resetComposerModelSelectionsForTests() + _resetSessionOwnerHintsForTests() + $sessions.set([]) + $activeSessionId.set(null) + $gatewayState.set('open') + $modelPickerOpen.set(true) + $currentModel.set('ambient-a') + $currentProvider.set('provider-a') + setApiRequestConnection('source-a') + setApiRequestProfile('default') + $newChatRoute.set(owner) + $newChatProfile.set(owner.profile) + $newChatConnectionId.set(owner.connectionId) + setComposerModelSelectionOwner(owner) +}) +afterEach(() => { cleanup(); setApiRequestConnection(null); setApiRequestProfile('default'); $newChatRoute.set(null); $newChatProfile.set(null); $newChatConnectionId.set(null); $modelPickerOpen.set(false); $gatewayState.set('idle'); $sessions.set([]); _resetSessionOwnerHintsForTests() }) + +function mount() { + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }) + const ambient = vi.fn(async () => options) + const select = vi.fn() + const view = render() + + return { client, view, ambient, select } +} + +it('uses the captured fresh draft source and backend target in the actual overlay catalog', async () => { + const { ambient } = mount() + await screen.findByText('catalog-b') + expect(calls.agent).toHaveBeenCalledWith('source-b', 'same-name', 'model.options', { profile: 'backend-b', explicit_only: true }) + expect(ambient).not.toHaveBeenCalled() +}) + +it('does not mark a stale catalog default current when a valid scalar receipt owns the draft', async () => { + recordComposerModelSelection(captureComposerModelSelection(owner), { model: 'scalar-b', provider: '', source: 'default' }) + $currentModel.set('scalar-b') + $currentProvider.set('') + const { select } = mount() + const row = await screen.findByText('catalog-b') + expect(row.closest('[cmdk-item]')?.className).not.toContain('bg-primary text-primary-foreground') + expect(select).not.toHaveBeenCalled() +}) + +it('preserves a deliberate draft pin through a late catalog reply', async () => { + let resolve!: (value: typeof options) => void + calls.agent.mockReturnValueOnce(new Promise(r => { resolve = r })) + mount() + await waitFor(() => expect(calls.agent).toHaveBeenCalledOnce()) + await act(async () => { + recordComposerModelSelection(captureComposerModelSelection(owner), { model: 'pinned-b', provider: 'custom:b', source: 'manual' }) + $currentModel.set('pinned-b') + $currentProvider.set('custom:b') + resolve(options) + }) + const row = await screen.findByText('pinned-b') + expect(row.closest('[cmdk-item]')?.className).toContain('bg-primary text-primary-foreground') +}) + +it.each([ + ['legacy', 'rejected'], ['legacy', 'empty'], ['local', 'rejected'], ['local', 'empty'] +] as const)('preserves a live %s owner through %s RPC recovery in the overlay', async (source, rpcUnavailable) => { + const runtimeId = 'owned-runtime' + const connectionId = source === 'legacy' ? null : 'local' + $sessions.set([{ id: runtimeId, profile: 'default' }] as never) + + if (connectionId) { + setSessionOwnerHint(runtimeId, { connectionId, profile: 'default', mode: 'local' }) + } + + expect(knownOwnerForSession(runtimeId)).toEqual(connectionId ? { connectionId, profile: 'default', mode: 'local' } : 'default') + $activeSessionId.set(runtimeId) + $currentModel.set('') + $currentProvider.set('') + const rpc = source === 'legacy' ? calls.profile : calls.agent + + if (rpcUnavailable === 'empty') { + rpc.mockResolvedValue({ providers: [] }) + } else { + rpc.mockRejectedValue(new Error('Offline catalog RPC')) + } + + const model = source === 'legacy' ? 'legacy-a' : 'local-b' + const expected = { ...options, model, providers: [{ ...options.providers[0], models: [model] }] } + calls.rest.mockImplementation(async (_opts, scope) => scope.connectionId === connectionId ? expected : options) + const { client, ambient } = mount() + await waitFor(() => expect(calls.rest).toHaveBeenCalledWith({ explicitOnly: true }, { connectionId, profile: 'default' })) + await screen.findByText(model) + expect(ambient).not.toHaveBeenCalled() + + if (source === 'legacy') { + expect(calls.profile).toHaveBeenCalledWith('default', 'model.options', { profile: 'default', session_id: runtimeId, explicit_only: true }, undefined, undefined) + } else { + expect(calls.agent).toHaveBeenCalledWith('local', 'default', 'model.options', { profile: 'default', session_id: runtimeId, explicit_only: true }) + } + + expect(client.getQueryData(['model-options', source === 'legacy' ? 'default' : 'local::default', runtimeId])).toEqual(expected) +}) diff --git a/apps/desktop/src/app/model-picker-overlay.tsx b/apps/desktop/src/app/model-picker-overlay.tsx index 626327b75edbb..e9670636c4126 100644 --- a/apps/desktop/src/app/model-picker-overlay.tsx +++ b/apps/desktop/src/app/model-picker-overlay.tsx @@ -4,17 +4,21 @@ import type { ModelSelection } from '@/app/shell/model-menu-panel' import { ModelPickerDialog } from '@/components/model-picker' import type { HermesGateway } from '@/hermes' import { useStoreSelector } from '@/lib/use-session-slice' +import { requestGatewayForAgent } from '@/store/gateway' import { $activeSessionId, $currentModel, $currentProvider, $gatewayState, $modelPickerOpen, + getComposerModelSelection, setModelPickerOpen } from '@/store/session' import { knownOwnerForSession, requestForOwnedSession } from '@/store/session-states' import { $focusedRuntimeId, $focusedSessionState } from '@/store/session-states' +import { useDraftComposerOwner } from './hooks/use-composer-model-owner' + interface ModelPickerOverlayProps { gateway?: HermesGateway onSelect: (selection: ModelSelection) => void @@ -36,6 +40,7 @@ export function ModelPickerOverlay({ gateway, onSelect, profile }: ModelPickerOv const focusedProvider = useStoreSelector($focusedSessionState, state => state?.provider ?? null) const gatewayOpen = useStore($gatewayState) === 'open' const open = useStore($modelPickerOpen) + const draftOwner = useDraftComposerOwner() // Prefer the focused tile's runtime when the overlay opens from a tile that // lacked a live menu (gateway closed → fallback path). @@ -48,20 +53,24 @@ export function ModelPickerOverlay({ gateway, onSelect, profile }: ModelPickerOv } const owner = knownOwnerForSession(sessionId) + const draftSelection = !sessionId ? getComposerModelSelection(draftOwner) : null const ownerProfile = typeof owner === 'string' ? owner : (owner?.targetProfile || owner?.profile) - const ownerConnection = owner && typeof owner === 'object' ? owner.connectionId : owner ? 'local' : undefined + const ownerConnection = owner && typeof owner === 'object' ? owner.connectionId : owner ? null : undefined return ( onSelect({ ...selection, sessionId })} open={open} - profile={ownerProfile || profile} - request={gateway && sessionId ? (method, params) => requestForOwnedSession(sessionId, gateway.request.bind(gateway), method, params) : undefined} + profile={ownerProfile || (!sessionId ? draftOwner.targetProfile || draftOwner.profile : profile)} + request={!sessionId + ? (method, params) => requestGatewayForAgent(draftOwner.connectionId, draftOwner.profile, method, params) + : gateway ? (method, params) => requestForOwnedSession(sessionId, gateway.request.bind(gateway), method, params) : undefined} + selectionIsAuthoritative={Boolean(draftSelection)} sessionId={sessionId} /> ) diff --git a/apps/desktop/src/app/session/hooks/model-selection-composition.test.tsx b/apps/desktop/src/app/session/hooks/model-selection-composition.test.tsx index d65205fc5f940..242895b0ef57b 100644 --- a/apps/desktop/src/app/session/hooks/model-selection-composition.test.tsx +++ b/apps/desktop/src/app/session/hooks/model-selection-composition.test.tsx @@ -513,6 +513,14 @@ describe('actual producer/store/admission composition', () => { }) it('refuses the actual Settings confirmation retry after A changes to B with the same logical profile', async () => { + vi.mocked(getGlobalModelOptions).mockImplementation(async (_opts, scope) => { + const connectionId = scope && typeof scope === 'object' ? scope.connectionId : getApiRequestConnection() + const pair = connectionId === b.connectionId ? defaultB : defaultA + + return { + providers: [{ slug: pair.provider, name: 'Offline owner', authenticated: true, models: [pair.model] }] + } + }) const value = await setup(1) await seed(value) const assignmentSources: (string | null)[] = [] diff --git a/apps/desktop/src/app/settings/custom-endpoints-settings.test.tsx b/apps/desktop/src/app/settings/custom-endpoints-settings.test.tsx index 5477ab4008f72..dd9a48d9ffcc4 100644 --- a/apps/desktop/src/app/settings/custom-endpoints-settings.test.tsx +++ b/apps/desktop/src/app/settings/custom-endpoints-settings.test.tsx @@ -1,10 +1,10 @@ -import { cleanup, fireEvent, render, screen, waitFor } from '@testing-library/react' +import { act, cleanup, fireEvent, render, screen, waitFor } from '@testing-library/react' import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' -import { setApiRequestConnection, setApiRequestProfile } from '@/api/client' +import { getApiRequestConnection, setApiRequestConnection, setApiRequestProfile } from '@/api/client' import { confirm } from '@/store/confirm' import { $activeGatewayProfile, $newChatProfile, $newChatRoute } from '@/store/profile' -import { _resetComposerModelSelectionsForTests } from '@/store/session' +import { $connection, _resetComposerModelSelectionsForTests } from '@/store/session' import type { CustomEndpoint, CustomEndpointsResponse } from '@/types/hermes' import { deferred } from '../../test/deferred' @@ -57,6 +57,7 @@ beforeEach(() => { deleteCustomEndpoint.mockReset() vi.mocked(confirm).mockReset().mockResolvedValue(true) _resetComposerModelSelectionsForTests() + $connection.set(null) $newChatRoute.set(owner) $newChatProfile.set(owner.profile) $activeGatewayProfile.set(owner.profile) @@ -82,6 +83,92 @@ async function renderSettings(changed = vi.fn()) { } describe('custom endpoint save ownership', () => { + it('clears source A form authority and key draft when the same profile rehomes to B and back', async () => { + const second = { ...endpoint, name: 'Source B endpoint', model: 'model-b', models: ['model-b'], base_url: 'https://b.invalid/v1' } + + const rehome = (connectionId: string) => { + setApiRequestConnection(connectionId) + $connection.set({ connectionId } as never) + $newChatRoute.set({ connectionId, profile: owner.profile, targetProfile: connectionId === 'source-a' ? 'backend-a' : 'backend-b' }) + } + + getCustomEndpoints.mockImplementation(async () => ({ endpoints: [getApiRequestConnection() === 'source-b' ? second : endpoint] })) + saveCustomEndpoint.mockResolvedValue({ id: second.id, endpoints: [second] }) + await renderSettings() + fireEvent.change(screen.getByPlaceholderText('Leave blank to keep current key'), { target: { value: 'fake-source-a-key' } }) + await act(async () => rehome('source-b')) + await screen.findByDisplayValue(second.base_url) + expect((screen.getByPlaceholderText('Leave blank to keep current key') as HTMLInputElement).value).toBe('') + fireEvent.click(screen.getByRole('button', { name: 'Save' })) + await waitFor(() => expect(saveCustomEndpoint).toHaveBeenCalledWith(expect.objectContaining({ + base_url: second.base_url, model: second.model, api_key: undefined + }))) + await act(async () => rehome('source-a')) + await screen.findByDisplayValue(endpoint.base_url) + }) + + it('rejects a delayed Delete confirmation from an unmounted owner even after A → B → A', async () => { + const confirmation = deferred() + vi.mocked(confirm).mockReturnValueOnce(confirmation.promise) + deleteCustomEndpoint.mockResolvedValue({ endpoints: [] }) + await renderSettings() + fireEvent.click(screen.getByRole('button', { name: 'Delete endpoint' })) + await waitFor(() => expect(confirm).toHaveBeenCalledOnce()) + + const rehome = (connectionId: string) => { + setApiRequestConnection(connectionId) + $connection.set({ connectionId } as never) + $newChatRoute.set({ connectionId, profile: owner.profile, targetProfile: connectionId === 'source-a' ? 'backend-a' : 'backend-b' }) + } + + await act(async () => rehome('source-b')) + await screen.findByRole('button', { name: 'Save' }) + await act(async () => rehome('source-a')) + await screen.findByRole('button', { name: 'Save' }) + await act(async () => confirmation.resolve(true)) + expect(deleteCustomEndpoint).not.toHaveBeenCalled() + }) + + it('rejects a stale Delete confirmation when React batches the source A → B → A round trip', async () => { + const confirmation = deferred() + vi.mocked(confirm).mockReturnValueOnce(confirmation.promise) + deleteCustomEndpoint.mockResolvedValue({ endpoints: [] }) + saveCustomEndpoint.mockResolvedValue(savedResponse()) + await renderSettings() + fireEvent.click(screen.getByRole('button', { name: 'Delete endpoint' })) + await waitFor(() => expect(confirm).toHaveBeenCalledOnce()) + await act(async () => { + setApiRequestConnection('source-b') + $connection.set({ connectionId: 'source-b' } as never) + setApiRequestConnection(owner.connectionId) + $connection.set({ connectionId: owner.connectionId } as never) + confirmation.resolve(true) + }) + expect(deleteCustomEndpoint).not.toHaveBeenCalled() + expect(notifyError).toHaveBeenCalledWith(expect.objectContaining({ message: expect.stringMatching(/target changed/i) }), 'Delete failed') + await waitFor(() => expect(getCustomEndpoints).toHaveBeenCalledTimes(2)) + await screen.findByDisplayValue(endpoint.base_url) + fireEvent.click(await screen.findByRole('button', { name: 'Save' })) + await waitFor(() => expect(saveCustomEndpoint).toHaveBeenCalledOnce()) + }) + + it('discards a late source A inventory after the settings owner has rehomed to B', async () => { + const pending = deferred() + const second = { ...endpoint, name: 'Source B endpoint', base_url: 'https://b.invalid/v1' } + getCustomEndpoints.mockReturnValueOnce(pending.promise).mockResolvedValue({ endpoints: [second] }) + const { CustomEndpointsSettings } = await import('./custom-endpoints-settings') + render() + await waitFor(() => expect(getCustomEndpoints).toHaveBeenCalledOnce()) + await act(async () => { + setApiRequestConnection('source-b') + $connection.set({ connectionId: 'source-b' } as never) + $newChatRoute.set({ connectionId: 'source-b', profile: owner.profile }) + }) + await act(async () => pending.resolve({ endpoints: [endpoint], current: savedResponse().current })) + await screen.findByDisplayValue(second.base_url) + expect(screen.queryByDisplayValue(endpoint.base_url)).toBeNull() + }) + it('carries the original owner through a delayed save callback', async () => { const pending = deferred() saveCustomEndpoint.mockReturnValueOnce(pending.promise) diff --git a/apps/desktop/src/app/settings/custom-endpoints-settings.tsx b/apps/desktop/src/app/settings/custom-endpoints-settings.tsx index 35e9254578564..e4d9252555636 100644 --- a/apps/desktop/src/app/settings/custom-endpoints-settings.tsx +++ b/apps/desktop/src/app/settings/custom-endpoints-settings.tsx @@ -1,7 +1,7 @@ import { useEffect, useRef, useState } from 'react' import { beginMainModelSave, ownsMainModelSave } from '@/app/session/hooks/composer-model-selection-owner' -import type { OnMainModelChanged } from '@/app/session/hooks/composer-model-selection-owner' +import type { ComposerSelectionOwner, OnMainModelChanged } from '@/app/session/hooks/composer-model-selection-owner' import { Button } from '@/components/ui/button' import { Checkbox } from '@/components/ui/checkbox' import { Input } from '@/components/ui/input' @@ -19,6 +19,8 @@ import { confirm } from '@/store/confirm' import { notify, notifyError } from '@/store/notifications' import type { CustomEndpoint, CustomEndpointUpdate } from '@/types/hermes' +import { requireCurrentModelOwner, useModelFormKey, useModelOwnerIsCurrent, useModelRequestOwner } from '../hooks/use-composer-model-owner' + import { EmptyState, Pill, SectionHeading, SettingsContent, SettingsSkeleton } from './primitives' interface CustomEndpointsSettingsProps { @@ -78,6 +80,14 @@ function toPayload(form: EndpointForm, models?: string[]): CustomEndpointUpdate } export function CustomEndpointsSettings({ onConfigSaved, onMainModelChanged }: CustomEndpointsSettingsProps) { + const owner = useModelRequestOwner() + const formKey = useModelFormKey(owner) + + return +} + +function CustomEndpointsForm({ onConfigSaved, onMainModelChanged, owner }: CustomEndpointsSettingsProps & { owner: ComposerSelectionOwner }) { + const isOwnerCurrent = useModelOwnerIsCurrent(owner) const [loading, setLoading] = useState(true) const [saving, setSaving] = useState(false) const [testing, setTesting] = useState(false) @@ -91,8 +101,15 @@ export function CustomEndpointsSettings({ onConfigSaved, onMainModelChanged }: C const modelWritePendingRef = useRef(false) async function refresh() { + if (!isOwnerCurrent()) { + return + } + const data = await getCustomEndpoints() - setEndpoints(data.endpoints) + + if (isOwnerCurrent()) { + setEndpoints(data.endpoints) + } } useEffect(() => { @@ -102,7 +119,7 @@ export function CustomEndpointsSettings({ onConfigSaved, onMainModelChanged }: C try { const data = await getCustomEndpoints() - if (cancelled) { + if (cancelled || !isOwnerCurrent()) { return } @@ -114,9 +131,11 @@ export function CustomEndpointsSettings({ onConfigSaved, onMainModelChanged }: C setDiscoveredModels(current.models) } } catch (err) { - notifyError(err, 'Could not load custom endpoints') + if (!cancelled && isOwnerCurrent()) { + notifyError(err, 'Could not load custom endpoints') + } } finally { - if (!cancelled) { + if (!cancelled && isOwnerCurrent()) { setLoading(false) } } @@ -127,7 +146,7 @@ export function CustomEndpointsSettings({ onConfigSaved, onMainModelChanged }: C return () => { cancelled = true } - }, []) + }, [isOwnerCurrent]) async function handleSave() { if (modelWritePendingRef.current) { @@ -135,9 +154,10 @@ export function CustomEndpointsSettings({ onConfigSaved, onMainModelChanged }: C } modelWritePendingRef.current = true - const origin = beginMainModelSave() try { + requireCurrentModelOwner(isOwnerCurrent) + const origin = beginMainModelSave() setSaving(true) const response = await saveCustomEndpoint(toPayload(form, discoveredModels)) @@ -145,20 +165,26 @@ export function CustomEndpointsSettings({ onConfigSaved, onMainModelChanged }: C return } - setEndpoints(response.endpoints) const saved = response.endpoints.find(endpoint => endpoint.id === response.id) + if (saved && saved.is_current) { + onMainModelChanged?.({ ...origin, provider: saved.id, model: saved.model }) + } + + onConfigSaved?.() + + if (!isOwnerCurrent()) { + return + } + + setEndpoints(response.endpoints) + if (saved) { setForm(formFromEndpoint(saved)) setDiscoveredModels(saved.models) } - if (saved && saved.is_current) { - onMainModelChanged?.({ ...origin, provider: saved.id, model: saved.model }) - } - triggerHaptic('success') - onConfigSaved?.() notify({ kind: 'success', message: 'Custom endpoint saved.' }) } catch (err) { notifyError(err, 'Save failed') @@ -170,8 +196,14 @@ export function CustomEndpointsSettings({ onConfigSaved, onMainModelChanged }: C async function handleValidate() { try { + requireCurrentModelOwner(isOwnerCurrent) setTesting(true) const response = await validateCustomEndpoint(toPayload(form)) + + if (!isOwnerCurrent()) { + return + } + setDiscoveredModels(response.models) if (response.ok) { @@ -204,9 +236,10 @@ export function CustomEndpointsSettings({ onConfigSaved, onMainModelChanged }: C } modelWritePendingRef.current = true - const origin = beginMainModelSave() try { + requireCurrentModelOwner(isOwnerCurrent) + const origin = beginMainModelSave() setActivating(endpoint.id) const response = await activateCustomEndpoint(endpoint.id) @@ -218,6 +251,11 @@ export function CustomEndpointsSettings({ onConfigSaved, onMainModelChanged }: C // the committed default from the composer or report activation failure. onConfigSaved?.() onMainModelChanged?.({ ...origin, provider: response.provider, model: response.model }) + + if (!isOwnerCurrent()) { + return + } + triggerHaptic('success') try { @@ -250,8 +288,15 @@ export function CustomEndpointsSettings({ onConfigSaved, onMainModelChanged }: C modelWritePendingRef.current = true try { + requireCurrentModelOwner(isOwnerCurrent) setDeleting(endpoint.id) const response = await deleteCustomEndpoint(endpoint.id) + onConfigSaved?.() + + if (!isOwnerCurrent()) { + return + } + setEndpoints(response.endpoints) if (form.id === endpoint.id) { @@ -259,7 +304,6 @@ export function CustomEndpointsSettings({ onConfigSaved, onMainModelChanged }: C setDiscoveredModels([]) } - onConfigSaved?.() triggerHaptic('success') } catch (err) { notifyError(err, 'Delete failed') diff --git a/apps/desktop/src/app/settings/model-settings.test.tsx b/apps/desktop/src/app/settings/model-settings.test.tsx index 8d80c37523c7f..125df1f7c67a9 100644 --- a/apps/desktop/src/app/settings/model-settings.test.tsx +++ b/apps/desktop/src/app/settings/model-settings.test.tsx @@ -4,8 +4,9 @@ import { MemoryRouter } from 'react-router' import { afterEach, beforeAll, beforeEach, describe, expect, it, vi } from 'vitest' import { getApiRequestConnection, setApiRequestConnection, setApiRequestProfile } from '@/api/client' +import { getHermesConfigRecord as readConfigOverBridge, saveHermesConfig as saveConfigOverBridge } from '@/api/config' import { $activeGatewayProfile, $newChatProfile, $newChatRoute } from '@/store/profile' -import { _resetComposerModelSelectionsForTests } from '@/store/session' +import { $connection, _resetComposerModelSelectionsForTests } from '@/store/session' import { deferred } from '../../test/deferred' @@ -39,13 +40,13 @@ vi.mock('@/hermes', () => ({ getAuxiliaryModels: (profile?: null | string) => getAuxiliaryModels(profile), getApiRequestProfile: () => 'default', getMoaModels: (profile?: null | string) => getMoaModels(profile), - profileScopeKey: (scope?: null | string) => (scope ?? '').trim() || 'default', + profileScopeKey: (scope?: null | string | { connectionId?: string | null; profile?: string | null }) => typeof scope === 'object' && scope ? `${scope.connectionId || 'local'}::${scope.profile || 'default'}` : (scope ?? '').trim() || 'default', setModelAssignment: (body: unknown) => setModelAssignment(body), getRecommendedDefaultModel: (slug: string) => getRecommendedDefaultModel(slug), saveMoaModels: (body: unknown) => saveMoaModels(body), setEnvVar: (key: string, value: string) => setEnvVar(key, value), - getHermesConfigRecord: () => getHermesConfigRecord(), - saveHermesConfig: (config: unknown) => saveHermesConfig(config), + getHermesConfigRecord: (profile?: null | string) => getHermesConfigRecord(profile), + saveHermesConfig: (config: unknown, profile?: null | string) => saveHermesConfig(config, profile), setApiRequestProfile: () => {} })) @@ -63,6 +64,7 @@ vi.mock('../hooks/use-on-profile-switch', () => ({ beforeEach(() => { _resetComposerModelSelectionsForTests() + $connection.set(null) $newChatRoute.set(null) $newChatProfile.set(null) $activeGatewayProfile.set('default') @@ -114,6 +116,77 @@ async function renderModelSettings(scopeProfile?: string, onMainModelChanged = v } describe('ModelSettings profile scope', () => { + it('reloads the form for same-named source owners A → B → A before saving', async () => { + const rehome = (connectionId: string) => { + setApiRequestConnection(connectionId) + $connection.set({ connectionId } as never) + $newChatRoute.set({ connectionId, profile: 'default' }) + } + + rehome('source-a') + getGlobalModelInfo.mockImplementation(async () => ({ + provider: getApiRequestConnection() === 'source-b' ? 'custom:b' : 'nous', + model: getApiRequestConnection() === 'source-b' ? 'model-b' : 'hermes-4' + })) + getGlobalModelOptions.mockImplementation(async () => ({ + providers: getApiRequestConnection() === 'source-b' + ? [{ name: 'Source B', slug: 'custom:b', models: ['model-b'], authenticated: true, api_url: 'https://b.invalid/v1' }] + : [{ name: 'Nous', slug: 'nous', models: ['hermes-4'], authenticated: true }] + })) + await renderModelSettings() + await screen.findByRole('button', { name: 'Apply' }) + await act(async () => rehome('source-b')) + await waitFor(() => expect(screen.getAllByRole('combobox')[0].textContent).toContain('Source B')) + fireEvent.click(screen.getByRole('button', { name: 'Apply' })) + await waitFor(() => expect(setModelAssignment).toHaveBeenCalledWith(expect.objectContaining({ + provider: 'custom:b', model: 'model-b', base_url: 'https://b.invalid/v1' + }))) + await act(async () => rehome('source-a')) + await waitFor(() => expect(screen.getAllByRole('combobox')[0].textContent).toContain('Nous')) + }) + + it.each(['source-b', 'local'])('keeps config GET/PUT routing and record authority together on %s', async destination => { + const previousBridge = window.hermesDesktop + + const api = vi.fn(async (request: { connectionId?: string; method?: string; body?: unknown }) => + request.method === 'PUT' ? { ok: true } : { + fixture_source: request.connectionId, + agent: { reasoning_effort: 'medium', service_tier: 'normal' }, + memory: { enabled: false }, governor: { semantic_memory_enabled: true } + }) + + window.hermesDesktop = { api } as never + getHermesConfigRecord.mockImplementation(readConfigOverBridge) + saveHermesConfig.mockImplementation(saveConfigOverBridge) + + const rehome = (connectionId: string) => { + setApiRequestConnection(connectionId) + $connection.set({ connectionId } as never) + $newChatRoute.set({ connectionId, profile: 'default' }) + } + + try { + rehome('source-a') + await renderModelSettings() + await screen.findByRole('switch') + await act(async () => rehome(destination)) + await waitFor(() => expect(api).toHaveBeenCalledWith(expect.objectContaining({ connectionId: destination, path: '/api/config' }))) + fireEvent.click(await screen.findByRole('switch')) + await waitFor(() => expect(api).toHaveBeenCalledWith(expect.objectContaining({ + connectionId: destination, method: 'PUT', body: { config: expect.objectContaining({ + fixture_source: destination, memory: { enabled: false }, governor: { semantic_memory_enabled: true } + }) + } + }))) + expect(getHermesConfigRecord).toHaveBeenCalledWith(undefined) + expect(saveHermesConfig).toHaveBeenCalledWith(expect.objectContaining({ fixture_source: destination }), undefined) + } finally { + window.hermesDesktop = previousBridge + getHermesConfigRecord.mockReset() + saveHermesConfig.mockReset() + } + }) + it('carries the original source/profile/target through an asynchronous save callback', async () => { const owner = { connectionId: 'source-a', profile: 'specialist', targetProfile: 'backend-a' } $newChatRoute.set(owner) @@ -128,7 +201,10 @@ describe('ModelSettings profile scope', () => { fireEvent.click(await screen.findByRole('button', { name: 'Apply' })) await waitFor(() => expect(setModelAssignment).toHaveBeenCalledOnce()) const originalSource = getApiRequestConnection() - setApiRequestConnection('source-b') + await act(async () => { + setApiRequestConnection('source-b') + $connection.set({ connectionId: 'source-b' } as never) + }) pending.resolve({ ok: true, provider: 'nous', model: 'hermes-4', gateway_tools: [] }) await waitFor(() => expect(changed).toHaveBeenCalledOnce()) expect(changed).toHaveBeenCalledWith( @@ -345,7 +421,8 @@ describe('ModelSettings', () => { await waitFor(() => expect(saveHermesConfig).toHaveBeenCalledWith( - expect.objectContaining({ agent: expect.objectContaining({ service_tier: 'fast' }) }) + expect.objectContaining({ agent: expect.objectContaining({ service_tier: 'fast' }) }), + undefined ) ) }) diff --git a/apps/desktop/src/app/settings/model-settings.tsx b/apps/desktop/src/app/settings/model-settings.tsx index c26af6000e3ab..34b0e3b692a7f 100644 --- a/apps/desktop/src/app/settings/model-settings.tsx +++ b/apps/desktop/src/app/settings/model-settings.tsx @@ -1,11 +1,13 @@ +import { useQuery, useQueryClient } from '@tanstack/react-query' import { useCallback, useEffect, useMemo, useRef, useState } from 'react' import { beginMainModelSave, + composerOwnerKey, isMainModelSaveOriginCurrent, ownsMainModelSave } from '@/app/session/hooks/composer-model-selection-owner' -import type { OnMainModelChanged } from '@/app/session/hooks/composer-model-selection-owner' +import type { ComposerSelectionOwner, OnMainModelChanged } from '@/app/session/hooks/composer-model-selection-owner' import { Button } from '@/components/ui/button' import { Input } from '@/components/ui/input' import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from '@/components/ui/select' @@ -15,6 +17,7 @@ import { getAuxiliaryModels, getGlobalModelInfo, getGlobalModelOptions, + getHermesConfigRecord, getMoaModels, getRecommendedDefaultModel, saveHermesConfig, @@ -24,6 +27,7 @@ import { } from '@/hermes' import type { AuxiliaryModelsResponse, + HermesConfigRecord, MoaConfigResponse, MoaModelSlot, ModelOptionProvider, @@ -37,7 +41,8 @@ import { setMainModelAssignment } from '@/store/cron-model-impact' import { notifyError } from '@/store/notifications' import { startManualLocalEndpoint, startManualOnboarding, startManualProviderOAuth } from '@/store/onboarding' -import { hermesConfigCacheWriter, invalidateHermesConfig, useHermesConfigRecord } from '../hooks/use-config-record' +import { requireCurrentModelOwner, useModelFormKey, useModelOwnerIsCurrent, useModelRequestOwner } from '../hooks/use-composer-model-owner' +import { HERMES_CONFIG_KEY } from '../hooks/use-config-record' import { useOnProfileSwitch } from '../hooks/use-on-profile-switch' import { CONTROL_TEXT } from './constants' @@ -195,6 +200,22 @@ interface ModelSettingsProps { } export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSettingsProps) { + const owner = useModelRequestOwner(scopeProfile) + const formKey = useModelFormKey(owner, scopeProfile) + + return +} + +function ModelSettingsForm({ onMainModelChanged, owner, scopeProfile }: ModelSettingsProps & { owner: ComposerSelectionOwner }) { + const isOwnerCurrent = useModelOwnerIsCurrent(owner, scopeProfile) + const ownerKey = composerOwnerKey(owner) + const queryClient = useQueryClient() + + const configKey = useMemo( + () => [...HERMES_CONFIG_KEY, ownerKey] as const, + [ownerKey] + ) + const { t } = useI18n() const m = t.settings.model const [loading, setLoading] = useState(true) @@ -207,10 +228,26 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting const [moa, setMoa] = useState(null) const [selectedMoaPreset, setSelectedMoaPreset] = useState('') const [newMoaPresetName, setNewMoaPresetName] = useState('') + // agent.* defaults round-trip through the shared config cache (read → write // back the whole record), so a save here shows in the MCP/model surfaces. - const { data: config } = useHermesConfigRecord(scopeProfile) - const setConfig = useMemo(() => hermesConfigCacheWriter(scopeProfile), [scopeProfile]) + const { data: config } = useQuery({ + queryKey: configKey, + queryFn: () => { + requireCurrentModelOwner(isOwnerCurrent) + + // Keep GET and PUT on the same legacy routing seam. In particular an + // ambient explicit 'local' tag must bypass legacy remote overrides. + return getHermesConfigRecord(scopeProfile) + }, + staleTime: 0 + }) + + const setConfig = useCallback( + (value: HermesConfigRecord) => queryClient.setQueryData(configKey, value), + [configKey, queryClient] + ) + const [applying, setApplying] = useState(false) const [editingAuxTask, setEditingAuxTask] = useState(null) const [auxDraft, setAuxDraft] = useState<{ model: string; provider: string }>({ model: '', provider: '' }) @@ -232,11 +269,15 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting // Every profile-scoped async here captures this and bails before writing back, // so a request in flight when the user switches profiles can't paint profile - // A's models/providers into profile B (or fire onMainModelChanged for A). + // A's models/providers into profile B. Confirmed saves still report origin. const profileEpoch = useRef(0) const refresh = useCallback( async ({ replaceSelection = false }: { replaceSelection?: boolean } = {}) => { + if (!isOwnerCurrent()) { + return + } + const epoch = profileEpoch.current setLoading(true) setError('') @@ -249,7 +290,7 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting getMoaModels(scopeProfile).catch(() => null) ]) - if (profileEpoch.current !== epoch) { + if (profileEpoch.current !== epoch || !isOwnerCurrent()) { return } @@ -273,18 +314,18 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting // The config record loads via its own shared query; a model switch can // change it server-side (aux slots), so nudge that cache to refetch. - void invalidateHermesConfig(scopeProfile) + void queryClient.invalidateQueries({ queryKey: configKey }) } catch (err) { - if (profileEpoch.current === epoch) { + if (profileEpoch.current === epoch && isOwnerCurrent()) { setError(err instanceof Error ? err.message : String(err)) } } finally { - if (profileEpoch.current === epoch) { + if (profileEpoch.current === epoch && isOwnerCurrent()) { setLoading(false) } } }, - [scopeProfile] + [configKey, isOwnerCurrent, queryClient, scopeProfile] ) useEffect(() => { @@ -406,20 +447,24 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting } moaSaveTimer.current = window.setTimeout(() => { + if (!isOwnerCurrent()) { + return + } + void saveMoaModels(next, scopeProfile) .then(saved => { - if (moaSaveGeneration.current === generation) { + if (moaSaveGeneration.current === generation && isOwnerCurrent()) { setMoa(saved) } }) .catch(err => { - if (moaSaveGeneration.current === generation) { + if (moaSaveGeneration.current === generation && isOwnerCurrent()) { setError(err instanceof Error ? err.message : String(err)) } }) }, 600) }, - [scopeProfile] + [isOwnerCurrent, scopeProfile] ) const updateMoaPreset = useCallback( @@ -475,9 +520,10 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting setError('') try { + requireCurrentModelOwner(isOwnerCurrent) const saved = await saveMoaModels(next, scopeProfile) - if (profileEpoch.current !== epoch) { + if (profileEpoch.current !== epoch || !isOwnerCurrent()) { return } @@ -488,7 +534,7 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting setApplying(false) } }, - [scopeProfile] + [isOwnerCurrent, scopeProfile] ) const auxiliaryTaskLabel = useCallback((key: string) => m.tasks[key]?.label ?? key, [m.tasks]) @@ -544,16 +590,18 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting const prev = config const next = setNested(config, key, value) - setConfig(next) try { + requireCurrentModelOwner(isOwnerCurrent) + setConfig(next) await saveHermesConfig(next, scopeProfile) + void queryClient.invalidateQueries({ queryKey: HERMES_CONFIG_KEY }) } catch (err) { setConfig(prev) notifyError(err, m.defaultsFailed) } }, - [config, m.defaultsFailed, scopeProfile, setConfig] + [config, isOwnerCurrent, m.defaultsFailed, queryClient, scopeProfile, setConfig] ) // Paste an API key for the selected `api_key` provider, persist it, then @@ -572,7 +620,13 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting setError('') try { + requireCurrentModelOwner(isOwnerCurrent) await setEnvVar(keyEnv, apiKeyDraft.trim(), scopeProfile) + + if (!isOwnerCurrent()) { + return + } + setApiKeyDraft('') // Pick a sensible default for the freshly-activated provider (mirrors @@ -587,9 +641,13 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting nextModel = '' } + if (!isOwnerCurrent()) { + return + } + const options = await getGlobalModelOptions(undefined, scopeProfile) - if (profileEpoch.current !== epoch) { + if (profileEpoch.current !== epoch || !isOwnerCurrent()) { return } @@ -602,7 +660,7 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting } finally { setActivating(false) } - }, [apiKeyDraft, scopeProfile, selectedProviderRow]) + }, [apiKeyDraft, isOwnerCurrent, scopeProfile, selectedProviderRow]) // OAuth / external providers can't be activated with a pasted key — hand off // to the shared onboarding flow scoped to this provider's real sign-in. The @@ -613,7 +671,7 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting const rowSlug = selectedProviderRow?.slug.trim() ?? '' const slug = rowSlug || selectedProvider.trim() - if (!slug) { + if (!slug || !isOwnerCurrent()) { return } @@ -628,19 +686,20 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting // provider picker instead of deep-linking an unknown or stale slug. startManualOnboarding() } - }, [selectedProvider, selectedProviderRow]) + }, [isOwnerCurrent, selectedProvider, selectedProviderRow]) const applyMainModel = useCallback(async () => { if (!selectedProvider || !selectedModel) { return } - const epoch = profileEpoch.current - const origin = beginMainModelSave(scopeProfile) setApplying(true) setError('') try { + requireCurrentModelOwner(isOwnerCurrent) + const origin = beginMainModelSave(scopeProfile) + const result = await setMainModelAssignment( { model: selectedModel, @@ -648,29 +707,32 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting ...(selectedProviderRow?.api_url ? { base_url: selectedProviderRow.api_url } : {}) }, scopeProfile, - { ownsOrigin: () => isMainModelSaveOriginCurrent(origin, scopeProfile) } + { ownsOrigin: () => isOwnerCurrent() && isMainModelSaveOriginCurrent(origin, scopeProfile) } ) - if (profileEpoch.current !== epoch || !ownsMainModelSave(origin)) { + if (!ownsMainModelSave(origin)) { return } const provider = result.provider || selectedProvider const model = result.model || selectedModel - setMainModel({ provider, model }) - setSwitchStaleAux(result.stale_aux ?? []) - // The callback carries origin scope; controls decide whether this // owner may paint the foreground draft or only its own default cache. onMainModelChanged?.({ ...origin, provider, model }) + if (!isOwnerCurrent()) { + return + } + + setMainModel({ provider, model }) + setSwitchStaleAux(result.stale_aux ?? []) await refresh() } catch (err) { setError(err instanceof Error ? err.message : String(err)) } finally { setApplying(false) } - }, [onMainModelChanged, refresh, scopeProfile, selectedModel, selectedProvider, selectedProviderRow]) + }, [isOwnerCurrent, onMainModelChanged, refresh, scopeProfile, selectedModel, selectedProvider, selectedProviderRow]) // Sibling of the applyMainModel endpoint passthrough (#65254): auxiliary // assignments targeting a user-defined provider must carry that provider's @@ -696,6 +758,7 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting setError('') try { + requireCurrentModelOwner(isOwnerCurrent) await setModelAssignment( { model: mainModel.model, @@ -713,7 +776,7 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting setApplying(false) } }, - [endpointForProvider, mainModel, refresh, scopeProfile] + [endpointForProvider, isOwnerCurrent, mainModel, refresh, scopeProfile] ) const applyAuxiliaryDraft = useCallback( @@ -726,6 +789,7 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting setError('') try { + requireCurrentModelOwner(isOwnerCurrent) await setModelAssignment( { model: auxDraft.model, @@ -736,6 +800,11 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting }, scopeProfile ) + + if (!isOwnerCurrent()) { + return + } + setEditingAuxTask(null) await refresh() } catch (err) { @@ -744,7 +813,7 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting setApplying(false) } }, - [auxDraft, endpointForProvider, refresh, scopeProfile] + [auxDraft, endpointForProvider, isOwnerCurrent, refresh, scopeProfile] ) const beginAuxiliaryEdit = useCallback( @@ -770,6 +839,7 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting setError('') try { + requireCurrentModelOwner(isOwnerCurrent) await setModelAssignment( { model: mainModel.model, @@ -779,6 +849,11 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting }, scopeProfile ) + + if (!isOwnerCurrent()) { + return + } + setSwitchStaleAux([]) await refresh() } catch (err) { @@ -786,7 +861,7 @@ export function ModelSettings({ onMainModelChanged, scopeProfile }: ModelSetting } finally { setApplying(false) } - }, [mainModel, refresh, scopeProfile]) + }, [isOwnerCurrent, mainModel, refresh, scopeProfile]) if (loading && !mainModel) { return diff --git a/apps/desktop/src/app/shell/model-menu-panel-owner.test.tsx b/apps/desktop/src/app/shell/model-menu-panel-owner.test.tsx new file mode 100644 index 0000000000000..66cfdec29e369 --- /dev/null +++ b/apps/desktop/src/app/shell/model-menu-panel-owner.test.tsx @@ -0,0 +1,129 @@ +import { QueryClient, QueryClientProvider } from '@tanstack/react-query' +import { act, cleanup, render, screen, waitFor } from '@testing-library/react' +import { afterEach, beforeEach, expect, it, vi } from 'vitest' + +import { setApiRequestConnection, setApiRequestProfile } from '@/api/client' +import type * as Gateway from '@/store/gateway' +import { $activeGatewayProfile, $newChatConnectionId, $newChatProfile, $newChatRoute } from '@/store/profile' +import { + $activeSessionId, $currentModel, $currentProvider, $sessions, _resetComposerModelSelectionsForTests, + _resetSessionOwnerHintsForTests, captureComposerModelSelection, recordComposerModelSelection, + setComposerModelSelectionOwner, setSessionOwnerHint +} from '@/store/session' +import { knownOwnerForSession } from '@/store/session-states' +import type { ModelOptionsResponse } from '@/types/hermes' + +import { deferred } from '../../test/deferred' + +import { ModelMenuPanel } from './model-menu-panel' + +const calls = vi.hoisted(() => ({ agent: vi.fn(), rest: vi.fn() })) +vi.mock('@/store/gateway', async original => ({ + ...await original(), + requestGatewayForAgent: (...args: unknown[]) => calls.agent(...args) +})) +vi.mock('@/hermes', () => ({ getGlobalModelOptions: (...args: unknown[]) => calls.rest(...args), setApiRequestProfile: vi.fn() })) +// Inspect the real producer's current pair without involving dropdown portals. +vi.mock('./model-catalog-menu', () => ({ + ModelMenuCloseContext: {}, + ModelCatalogMenu: ({ controller }: { controller: { current: { model: string; provider: string } } }) => + {controller.current.model}|{controller.current.provider} +})) + +const ownerB = { connectionId: 'source-b', profile: 'same-name', targetProfile: 'backend-b' } +const options = (model: string): ModelOptionsResponse => ({ model, provider: 'custom:b', providers: [{ name: 'B', slug: 'custom:b', models: [model] }] }) + +beforeEach(() => { + vi.clearAllMocks() + calls.agent.mockReset().mockResolvedValue(options('catalog-b')) + calls.rest.mockReset().mockResolvedValue(options('rest-b')) + _resetComposerModelSelectionsForTests() + _resetSessionOwnerHintsForTests() + $sessions.set([]) + $activeSessionId.set(null) + $currentModel.set('ambient-a') + $currentProvider.set('provider-a') + setApiRequestConnection('source-a') + setApiRequestProfile('default') + $activeGatewayProfile.set('default') + $newChatRoute.set(ownerB) + $newChatProfile.set(ownerB.profile) + $newChatConnectionId.set(ownerB.connectionId) + setComposerModelSelectionOwner(ownerB) +}) +afterEach(() => { cleanup(); setApiRequestConnection(null); setApiRequestProfile('default'); $newChatRoute.set(null); $newChatProfile.set(null); $newChatConnectionId.set(null); $sessions.set([]); _resetSessionOwnerHintsForTests() }) + +function mount(rpcUnavailable?: 'empty' | 'rejected') { + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }) + const ambient = vi.fn(async () => options('ambient-catalog-a')) + + if (rpcUnavailable === 'empty') { + ambient.mockResolvedValue({ providers: [] }) + } else if (rpcUnavailable === 'rejected') { + ambient.mockRejectedValue(new Error('Offline catalog RPC')) + } + + const view = render() + + return { client, ambient, view } +} + +it('routes a fresh pending/failed B draft catalog to B while the foreground gateway stays on A', async () => { + mount() + await waitFor(() => expect(calls.agent).toHaveBeenCalledWith('source-b', 'same-name', 'model.options', { profile: 'backend-b', explicit_only: true })) + await waitFor(() => expect(screen.getByTestId('current').textContent).toBe('catalog-b|custom:b')) +}) + +it('retains a valid scalar draft receipt beside a stale complete catalog instead of borrowing ambient A', async () => { + recordComposerModelSelection(captureComposerModelSelection(ownerB), { model: 'scalar-b', provider: '', source: 'default' }) + mount() + await waitFor(() => expect(screen.getByTestId('current').textContent).toBe('scalar-b|')) +}) + +it('keeps late B catalog replies in B cache after a B → A draft rehome', async () => { + const pending = deferred() + calls.agent.mockReturnValueOnce(pending.promise).mockResolvedValue(options('catalog-a')) + const { client } = mount() + await waitFor(() => expect(calls.agent).toHaveBeenCalledOnce()) + await act(async () => { + const ownerA = { connectionId: 'source-a', profile: 'same-name', targetProfile: 'backend-a' } + $newChatRoute.set(ownerA) + $newChatConnectionId.set(ownerA.connectionId) + setComposerModelSelectionOwner(ownerA) + }) + await waitFor(() => expect(screen.getByTestId('current').textContent).toBe('catalog-a|custom:b')) + await act(async () => pending.resolve(options('late-b'))) + expect(screen.getByTestId('current').textContent).toBe('catalog-a|custom:b') + expect(client.getQueryData(['model-options', 'source-b::backend-b', 'global'])).toEqual(options('late-b')) +}) + +it('recovers a failed B RPC through the same captured B REST target', async () => { + calls.agent.mockRejectedValueOnce(new Error('B route unavailable')) + mount() + await waitFor(() => expect(calls.rest).toHaveBeenCalledWith({ explicitOnly: true }, { connectionId: 'source-b', profile: 'backend-b' })) + await waitFor(() => expect(screen.getByTestId('current').textContent).toBe('rest-b|custom:b')) +}) + +it.each([ + ['legacy', 'rejected'], ['legacy', 'empty'], ['local', 'rejected'], ['local', 'empty'] +] as const)('preserves a live %s owner through %s RPC recovery', async (source, rpcUnavailable) => { + const runtimeId = 'owned-runtime' + const connectionId = source === 'legacy' ? null : 'local' + $sessions.set([{ id: runtimeId, profile: 'default' }] as never) + + if (connectionId) { + setSessionOwnerHint(runtimeId, { connectionId, profile: 'default', mode: 'local' }) + } + + expect(knownOwnerForSession(runtimeId)).toEqual(connectionId ? { connectionId, profile: 'default', mode: 'local' } : 'default') + $activeSessionId.set(runtimeId) + $currentModel.set('') + $currentProvider.set('') + calls.rest.mockImplementation(async (_opts, scope) => options(scope.connectionId === 'local' ? 'local-b' : 'legacy-a')) + const { client, ambient } = mount(rpcUnavailable) + await waitFor(() => expect(calls.rest).toHaveBeenCalledWith({ explicitOnly: true }, { connectionId, profile: 'default' })) + const model = source === 'legacy' ? 'legacy-a' : 'local-b' + await waitFor(() => expect(screen.getByTestId('current').textContent).toBe(`${model}|custom:b`)) + expect(ambient).toHaveBeenCalledWith('model.options', { profile: 'default', session_id: runtimeId, explicit_only: true }) + expect(client.getQueryData(['model-options', source === 'legacy' ? 'default' : 'local::default', runtimeId])).toEqual(options(model)) +}) diff --git a/apps/desktop/src/app/shell/model-menu-panel.test.tsx b/apps/desktop/src/app/shell/model-menu-panel.test.tsx index ae9d1d4be77d9..d964d105816a6 100644 --- a/apps/desktop/src/app/shell/model-menu-panel.test.tsx +++ b/apps/desktop/src/app/shell/model-menu-panel.test.tsx @@ -86,15 +86,17 @@ function renderPanel(onSelectModel = vi.fn()) { it.each(['legacy-local', 'remote-target'])('loads the known session owner catalog while ambient source differs (%s)', async kind => { setApiRequestConnection('ambient-remote') + if (kind === 'legacy-local') { $sessions.set([{ id: 'runtime-1', profile: 'local-specialist' }] as never) } else { setSessionOwnerHint('runtime-1', { connectionId: 'owner-remote', profile: 'logical', targetProfile: 'backend-target', mode: 'remote' }) } + try { renderPanel() await vi.waitFor(() => expect(getGlobalModelOptions).toHaveBeenCalledWith({ explicitOnly: true }, { - connectionId: kind === 'legacy-local' ? 'local' : 'owner-remote', + connectionId: kind === 'legacy-local' ? null : 'owner-remote', profile: kind === 'legacy-local' ? 'local-specialist' : 'backend-target' })) } finally { diff --git a/apps/desktop/src/app/shell/model-menu-panel.tsx b/apps/desktop/src/app/shell/model-menu-panel.tsx index 2a84da56097a1..00e1a8dbb8786 100644 --- a/apps/desktop/src/app/shell/model-menu-panel.tsx +++ b/apps/desktop/src/app/shell/model-menu-panel.tsx @@ -12,7 +12,7 @@ import { modelOptionsQueryKey, requestModelOptions, selectionUnavailable } from import { currentPickerSelection } from '@/lib/model-status-label' import { DEFAULT_REASONING_EFFORT } from '@/lib/reasoning-effort' import { cn } from '@/lib/utils' -import { activeGatewayConnectionId } from '@/store/gateway' +import { activeGatewayConnectionId, requestGatewayForAgent } from '@/store/gateway' import { $modelPresets, applyModelPreset, modelPresetKey, setModelPreset } from '@/store/model-presets' import { $visibleModels } from '@/store/model-visibility' import { notifyError } from '@/store/notifications' @@ -23,6 +23,7 @@ import { $defaultReasoningEffort, $selectedStoredSessionId, beginRuntimeOptionIntent, + getComposerModelSelection, markComposerSelectionManual, ownsRuntimeOptionIntent, setCurrentFastMode, @@ -31,6 +32,8 @@ import { import { $sessionStates, knownOwnerForSession, sessionTileDelegate } from '@/store/session-states' import type { ModelOptionsResponse } from '@/types/hermes' +import { useDraftComposerOwner } from '../hooks/use-composer-model-owner' + import { ModelCatalogMenu, type ModelMenuController } from './model-catalog-menu' export { ModelMenuCloseContext } from './model-catalog-menu' @@ -66,8 +69,13 @@ export function ModelMenuPanel({ gateway, onSelectModel, profile = 'default', re const view = useSessionView() const activeSessionId = useStore(view.$runtimeId) const owner = knownOwnerForSession(activeSessionId) - const connectionId = owner && typeof owner === 'object' ? owner.connectionId : owner ? 'local' : activeGatewayConnectionId() - const catalogProfile = (typeof owner === 'string' ? owner : (owner?.targetProfile || owner?.profile)) || profile + const draftOwner = useDraftComposerOwner() + const connectionId = owner && typeof owner === 'object' ? owner.connectionId : owner ? null : !activeSessionId ? draftOwner.connectionId : activeGatewayConnectionId() + const catalogProfile = (typeof owner === 'string' ? owner : (owner?.targetProfile || owner?.profile)) || (!activeSessionId ? draftOwner.targetProfile || draftOwner.profile : profile) + + const catalogRequest = !activeSessionId + ? (method: string, params?: Record) => requestGatewayForAgent(draftOwner.connectionId, draftOwner.profile, method, params) + : requestGateway const unconfirmedOptions = useStore( useMemo( @@ -82,6 +90,7 @@ export function ModelMenuPanel({ gateway, onSelectModel, profile = 'default', re const currentFastMode = useStore(view.$fast) const currentModel = useStore(view.$model) const currentProvider = useStore(view.$provider) + const draftSelection = !activeSessionId ? getComposerModelSelection(draftOwner) : null const currentReasoningEffort = useStore(view.$reasoningEffort) const modelPresets = useStore($modelPresets) const defaultEffort = useStore($defaultReasoningEffort) || DEFAULT_REASONING_EFFORT @@ -96,11 +105,13 @@ export function ModelMenuPanel({ gateway, onSelectModel, profile = 'default', re const modelOptions = useQuery({ queryKey: modelOptionsQueryKey(catalogProfile, activeSessionId, connectionId), queryFn: (): Promise => - requestModelOptions({ connectionId, gateway, profile: catalogProfile, request: requestGateway, sessionId: activeSessionId }) + requestModelOptions({ connectionId, gateway, profile: catalogProfile, request: catalogRequest, sessionId: activeSessionId }) }) const { model: optionsModel, provider: optionsProvider } = currentPickerSelection( - { model: currentModel, provider: currentProvider }, + activeSessionId + ? { model: currentModel, provider: currentProvider } + : { model: draftSelection?.model || '', provider: draftSelection?.provider || '', authoritative: Boolean(draftSelection) }, modelOptions.data ) @@ -123,7 +134,7 @@ export function ModelMenuPanel({ gateway, onSelectModel, profile = 'default', re gateway, profile: catalogProfile, refresh: true, - request: requestGateway, + request: catalogRequest, sessionId: activeSessionId }) @@ -329,7 +340,7 @@ export function ModelMenuPanel({ gateway, onSelectModel, profile = 'default', re gateway={gateway} includeMoa profile={catalogProfile} - request={requestGateway} + request={catalogRequest} sessionId={activeSessionId} /> ) diff --git a/apps/desktop/src/components/model-picker.test.tsx b/apps/desktop/src/components/model-picker.test.tsx index f6bf30e3696d6..5b646e4d06cbe 100644 --- a/apps/desktop/src/components/model-picker.test.tsx +++ b/apps/desktop/src/components/model-picker.test.tsx @@ -17,14 +17,28 @@ afterEach(() => { cleanup(); vi.clearAllMocks() }) const payload = (model: string) => ({ providers: [{ models: [model], name: 'Target provider', slug: 'target' }] }) +it('retains a confirmed scalar selection instead of marking a stale catalog model current', async () => { + vi.mocked(getGlobalModelOptions).mockResolvedValue({ ...payload('catalog-model'), model: 'catalog-model', provider: 'target' }) + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }) + + const view = render( + + ) + + const row = await view.findByText('catalog-model') + expect(row.closest('[cmdk-item]')?.className).not.toContain('bg-primary text-primary-foreground') +}) + it('sends the modal target profile and explains a missing session selection without changing it', async () => { vi.mocked(getGlobalModelOptions).mockResolvedValue(payload('model-b')) const select = vi.fn() const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }) + const view = render( ) + await view.findByText(/selected provider or model is unavailable in this profile/i) await view.findByText(/session keeps its own model selection/i) expect(getGlobalModelOptions).toHaveBeenCalledWith({ explicitOnly: true }, { connectionId: 'owner', profile: 'specialist' }) @@ -37,10 +51,12 @@ it('keeps a late catalog reply in its original source and profile cache', async .mockImplementationOnce(() => new Promise(resolve => { resolveA = resolve })) .mockResolvedValueOnce(payload('model-b')) const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }) + const picker = (connectionId: string, profile: string) => + const view = render(picker('source-a', 'target-a')) await vi.waitFor(() => expect(getGlobalModelOptions).toHaveBeenCalledTimes(1)) view.rerender(picker('source-b', 'target-b')) @@ -54,9 +70,11 @@ it('visibility and picker subscribers share the same target catalog instead of a const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }) const request = vi.fn(async (_method: string, params?: Record) => payload(params?.profile === 'target' ? 'target-model' : 'source-model')) const gateway = { request } as never + const view = render( ) + await vi.waitFor(() => expect(client.getQueryData(modelOptionsQueryKey('target'))).toEqual(payload('target-model'))) view.rerender( void profile?: string /** @@ -48,6 +50,7 @@ export function ModelPickerDialog({ sessionId, currentModel, currentProvider, + selectionIsAuthoritative = false, onSelect, profile = 'default', contentClassName @@ -70,7 +73,7 @@ export function ModelPickerDialog({ const providers = modelOptions.data?.providers ?? [] const { model: optionsModel, provider: optionsProvider } = currentPickerSelection( - { model: currentModel, provider: currentProvider }, + { model: currentModel, provider: currentProvider, authoritative: selectionIsAuthoritative }, modelOptions.data ) diff --git a/apps/desktop/src/lib/model-options.test.ts b/apps/desktop/src/lib/model-options.test.ts index fbdc4e35a3d35..2efb140bb17be 100644 --- a/apps/desktop/src/lib/model-options.test.ts +++ b/apps/desktop/src/lib/model-options.test.ts @@ -1,7 +1,16 @@ -import { afterEach, describe, expect, it, vi } from 'vitest' +import { QueryClient, QueryClientProvider } from '@tanstack/react-query' +import { cleanup, render, screen, waitFor } from '@testing-library/react' +import { createElement } from 'react' +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' import { setApiRequestConnection, setApiRequestProfile } from '@/api/client' +import { getGlobalModelOptions as readGlobalModelOptions } from '@/api/models' +import { ModelPickerDialog } from '@/components/model-picker' +import type { HermesApiRequest } from '@/global' import { getGlobalModelOptions } from '@/hermes' +import type { ModelOptionsResponse } from '@/types/hermes' + +import { deferred } from '../test/deferred' import { modelOptionsQueryKey, @@ -15,6 +24,24 @@ vi.mock('@/hermes', () => ({ getGlobalModelOptions: vi.fn(() => Promise.resolve(globalOptions)) })) +// Electron helpers have a separate compile project. Load their real pure +// exports for this bridge regression without pulling that project into the +// renderer's strict typecheck or changing Electron source/compiler settings. +const { apiRequestRegistryConnectionId, resolveProfileApiRequest } = await vi.importActual<{ + apiRequestRegistryConnectionId: (request: HermesApiRequest) => null | string + resolveProfileApiRequest: (profile: unknown, path: string, opts: Record) => { backendProfile: null | string } +}>('../../electron/connection-config') + +const { normalizeRegistry, resolvedConnectionId, resolveRegistryLocalRoute } = await vi.importActual<{ + normalizeRegistry: (input: unknown) => unknown + resolvedConnectionId: (registry: unknown, descriptor: Record) => null | string + resolveRegistryLocalRoute: (profile: unknown, opts: { globalRemote: boolean }) => { delegate: boolean; poolKey: string } +}>('../../electron/connection-registry') + +const { resolveDesktopRemoteRoute } = await vi.importActual<{ + resolveDesktopRemoteRoute: (input: Record) => null | { connectionId?: string; kind: string; source: string; url?: string } +}>('../../electron/desktop-remote-route') + describe('requestModelOptions', () => { afterEach(() => { vi.clearAllMocks() @@ -60,7 +87,7 @@ describe('requestModelOptions', () => { provider: 'hermes-local' }) - expect(getGlobalModelOptions).toHaveBeenCalledWith({ explicitOnly: true }, { connectionId: 'local', profile: 'default' }) + expect(getGlobalModelOptions).toHaveBeenCalledWith({ explicitOnly: true }, { connectionId: null, profile: 'default' }) }) it('recovers through profile-scoped REST when the gateway catalog request fails', async () => { @@ -79,7 +106,7 @@ describe('requestModelOptions', () => { await expect(requestModelOptions({ gateway: gateway as never, sessionId: 'session-1' })).resolves.toEqual( restPayload ) - expect(getGlobalModelOptions).toHaveBeenCalledWith({ explicitOnly: true }, { connectionId: 'local', profile: 'default' }) + expect(getGlobalModelOptions).toHaveBeenCalledWith({ explicitOnly: true }, { connectionId: null, profile: 'default' }) }) it('preserves the gateway error when its REST recovery path also fails', async () => { @@ -117,13 +144,13 @@ describe('requestModelOptions', () => { refresh: true, session_id: 'session-1' }) - expect(getGlobalModelOptions).toHaveBeenCalledWith({ explicitOnly: true, refresh: true }, { connectionId: 'local', profile: 'default' }) + expect(getGlobalModelOptions).toHaveBeenCalledWith({ explicitOnly: true, refresh: true }, { connectionId: null, profile: 'default' }) }) it('falls back to REST when no gateway is connected', async () => { await requestModelOptions({ refresh: true }) - expect(getGlobalModelOptions).toHaveBeenCalledWith({ explicitOnly: true, refresh: true }, { connectionId: 'local', profile: 'default' }) + expect(getGlobalModelOptions).toHaveBeenCalledWith({ explicitOnly: true, refresh: true }, { connectionId: null, profile: 'default' }) }) it('prefers an owner-routed request over the ambient gateway socket', async () => { @@ -168,7 +195,7 @@ describe('requestModelOptions', () => { vi.mocked(getGlobalModelOptions).mockResolvedValueOnce(restPayload) await expect(requestModelOptions({ profile: 'berry', request, sessionId: 'tile-1' })).resolves.toEqual(restPayload) - expect(getGlobalModelOptions).toHaveBeenCalledWith({ explicitOnly: true }, { connectionId: 'local', profile: 'berry' }) + expect(getGlobalModelOptions).toHaveBeenCalledWith({ explicitOnly: true }, { connectionId: null, profile: 'berry' }) }) it('freezes source and target before a late gateway failure', async () => { @@ -205,7 +232,7 @@ describe('modelOptionsQueryKey', () => { it('isolates identically named profiles and sessions on different sources', () => { expect(modelOptionsQueryKey('target', 'session', 'remote-a')).not.toEqual(modelOptionsQueryKey('target', 'session', 'remote-b')) - expect(modelOptionsQueryKey('target', 'session', 'local')).toEqual(['model-options', 'target', 'session']) + expect(modelOptionsQueryKey('target', 'session', 'local')).toEqual(['model-options', 'local::target', 'session']) }) }) @@ -226,3 +253,154 @@ describe('selectionUnavailable', () => { expect(selectionUnavailable(undefined, 'ollama-launch', 'model-a')).toBe(false) }) }) + +describe('legacy-null and explicit-local catalog authority', () => { + const env = { url: 'https://legacy-a.invalid', token: 'inert-routing-fixture' } + + const registry = normalizeRegistry({ + version: 2, + primary: 'local', + connections: [{ id: 'local', kind: 'local', label: 'This device' }] + }) + + const catalog = (model: string): ModelOptionsResponse => ({ + model, + provider: model + '-provider', + providers: [{ slug: model + '-provider', name: model + ' provider', models: [model] }] + }) + + const routed: { request: HermesApiRequest; target: string }[] = [] + + // main's API branch uses this real tag resolver. Its legacy branch uses + // resolveProfileApiRequest/resolveDesktopRemoteRoute; registry local uses + // resolveRegistryLocalRoute. Exercise those pure decisions without main, + // spawning, network, or a live connection/configuration. + async function bridge(request: HermesApiRequest): Promise { + expect(request.path).toMatch(/^\/api\/model\/options\?/) + const connectionId = apiRequestRegistryConnectionId(request) + + if (connectionId === null) { + const route = resolveProfileApiRequest(request.profile, request.path, { + primaryProfile: 'default', globalRemote: true + }) + + const remote = resolveDesktopRemoteRoute({ + config: { mode: 'local' }, env, profile: route.backendProfile, registry + }) + + expect(route.backendProfile).toBeNull() + expect(remote).toMatchObject({ kind: 'remote', source: 'env', url: env.url }) + expect(remote?.connectionId).toBeUndefined() + expect(resolvedConnectionId(registry, { + mode: 'remote', remoteKind: 'url', baseUrl: env.url, token: env.token, authMode: 'token' + })).toBeNull() + routed.push({ request, target: env.url }) + + return catalog('model-legacy-a') as T + } + + expect(connectionId).toBe('local') + const local = resolveRegistryLocalRoute(request.profile, { globalRemote: Boolean(env.url) }) + expect(local).toEqual({ delegate: false, poolKey: 'conn:local::default' }) + routed.push({ request, target: local.poolKey }) + + return catalog('model-local-b') as T + } + + beforeEach(() => { + routed.length = 0 + setApiRequestConnection(null) + setApiRequestProfile('default') + vi.mocked(getGlobalModelOptions).mockReset().mockImplementation(readGlobalModelOptions) + vi.stubGlobal('hermesDesktop', { api: vi.fn(bridge) }) + }) + + afterEach(() => { + cleanup() + vi.restoreAllMocks() + vi.unstubAllGlobals() + vi.mocked(getGlobalModelOptions).mockReset().mockResolvedValue(globalOptions) + setApiRequestConnection(null) + setApiRequestProfile(null) + }) + + it('distinguishes keys for physically separate legacy environment remote A and forced-local B', async () => { + const remote = resolveDesktopRemoteRoute({ config: { mode: 'local' }, env, profile: 'default', registry }) + expect(remote).toMatchObject({ kind: 'remote', source: 'env', url: env.url }) + expect(remote?.connectionId).toBeUndefined() + expect(resolveRegistryLocalRoute('default', { globalRemote: Boolean(env.url) })).toEqual({ + delegate: false, poolKey: 'conn:local::default' + }) + expect(modelOptionsQueryKey('default', null, null)).not.toEqual(modelOptionsQueryKey('default', null, 'local')) + }) + + it.each(['rejected', 'empty'])('recovers a %s legacy RPC through the real API and legacy bridge route', async failure => { + const request = vi.fn(async () => { + if (failure === 'rejected') { + throw new Error('inert legacy RPC failure') + } + + return { providers: [] } + }) + + await expect(requestModelOptions({ connectionId: null, profile: 'default', request: request as never })).resolves.toEqual(catalog('model-legacy-a')) + expect(request).toHaveBeenCalledWith('model.options', { explicit_only: true, profile: 'default' }) + expect(routed).toEqual([{ request: expect.objectContaining({ profile: 'default' }), target: env.url }]) + expect(routed[0].request).not.toHaveProperty('connectionId') + }) + + it('keeps forced-local B recovery registry-pinned while the legacy primary resolves to remote A', async () => { + const request = vi.fn(async () => { throw new Error('inert local RPC failure') }) + await expect(requestModelOptions({ connectionId: 'local', profile: 'default', request })).resolves.toEqual(catalog('model-local-b')) + expect(routed).toEqual([{ + request: expect.objectContaining({ connectionId: 'local', profile: 'default' }), target: 'conn:local::default' + }]) + }) + + it('retains captured legacy A when its RPC fails after the foreground moves to named source C', async () => { + const pending = deferred() + const request = vi.fn(() => pending.promise) + const result = requestModelOptions({ request: request as never }) + setApiRequestConnection('source-c') + setApiRequestProfile('other-profile') + pending.reject(new Error('inert late legacy RPC failure')) + await expect(result).resolves.toEqual(catalog('model-legacy-a')) + expect(routed[0]).toMatchObject({ request: { profile: 'default' }, target: env.url }) + expect(routed[0].request).not.toHaveProperty('connectionId') + }) + + it('keeps a late legacy A catalog in its own cache without painting the foreground local B picker', async () => { + Element.prototype.scrollIntoView = vi.fn() + vi.stubGlobal('ResizeObserver', class { observe() {} unobserve() {} disconnect() {} }) + const pending = deferred() + const requestA = vi.fn(() => pending.promise) + const requestB = vi.fn(async () => { throw new Error('inert local RPC failure') }) + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }) + + const picker = (connectionId: null | string) => createElement(QueryClientProvider, { client }, + createElement(ModelPickerDialog, { + connectionId, profile: 'default', request: (connectionId === null ? requestA : requestB) as never, + open: true, onOpenChange: vi.fn(), onSelect: vi.fn(), currentModel: '', currentProvider: '' + })) + + const view = render(picker(null)) + + try { + await waitFor(() => expect(requestA).toHaveBeenCalledOnce()) + setApiRequestConnection('local') + view.rerender(picker('local')) + await screen.findByText('model-local-b') + setApiRequestConnection('source-c') + setApiRequestProfile('other-profile') + pending.reject(new Error('inert late legacy RPC failure')) + await waitFor(() => expect(client.getQueryData(modelOptionsQueryKey('default', null, null))).toMatchObject({ model: 'model-legacy-a' })) + expect(client.getQueryData(modelOptionsQueryKey('default', null, 'local'))).toMatchObject({ model: 'model-local-b' }) + expect(screen.queryByText('model-legacy-a')).toBeNull() + expect(screen.getByText('model-local-b')).toBeTruthy() + } finally { + pending.resolve({ providers: [] }) + view.unmount() + client.clear() + } + }) +}) diff --git a/apps/desktop/src/lib/model-options.ts b/apps/desktop/src/lib/model-options.ts index 7b76cf9860c6a..7eb0c9b0e5b24 100644 --- a/apps/desktop/src/lib/model-options.ts +++ b/apps/desktop/src/lib/model-options.ts @@ -42,7 +42,7 @@ export function modelOptionsQueryKey( ) { const profileKey = (profile ?? '').trim() || 'default' - const sourceKey = connectionId && connectionId !== 'local' ? `${connectionId}::${profileKey}` : profileKey + const sourceKey = connectionId ? `${connectionId}::${profileKey}` : profileKey return ['model-options', sourceKey, sessionId || 'global'] as const } @@ -57,6 +57,7 @@ function restModelOptions( profile: ProfileScope ): Promise { const opts = { explicitOnly, ...(refresh ? { refresh: true } : {}) } + return getGlobalModelOptions(opts, profile) } @@ -71,7 +72,7 @@ export async function requestModelOptions({ }: ModelOptionsRequest): Promise { // Capture the owner before either async leg; foreground source changes must // not redirect a late REST recovery into another profile or connection. - const scope = { connectionId: connectionId || 'local', profile: profile ?? getApiRequestProfile() ?? 'default' } + const scope = { connectionId, profile: profile ?? getApiRequestProfile() ?? 'default' } const dispatch = request ?? (gateway ? gateway.request.bind(gateway) : null) if (dispatch) { diff --git a/apps/desktop/src/lib/model-status-label.test.ts b/apps/desktop/src/lib/model-status-label.test.ts index d1ad06fb73ee2..22743a5edb1dd 100644 --- a/apps/desktop/src/lib/model-status-label.test.ts +++ b/apps/desktop/src/lib/model-status-label.test.ts @@ -66,6 +66,14 @@ describe('model-status-label', () => { expect(currentPickerSelection({ model: 'opus', provider: '' }, options)).toEqual(options) }) + it('keeps an authoritative scalar receipt without inventing its provider from a stale catalog', () => { + expect(currentPickerSelection({ model: 'scalar', provider: '', authoritative: true }, options)).toEqual({ model: 'scalar', provider: '' }) + }) + + it('still hydrates an empty pair even when a caller marks it authoritative', () => { + expect(currentPickerSelection({ model: '', provider: '', authoritative: true }, options)).toEqual(options) + }) + it('falls back to the store while options are still loading', () => { expect(currentPickerSelection(store, undefined)).toEqual(store) }) diff --git a/apps/desktop/src/lib/model-status-label.ts b/apps/desktop/src/lib/model-status-label.ts index 5dde20746f8e7..0cf97b9a50e13 100644 --- a/apps/desktop/src/lib/model-status-label.ts +++ b/apps/desktop/src/lib/model-status-label.ts @@ -3,10 +3,11 @@ import { DEFAULT_REASONING_EFFORT, reasoningEffortLabel } from '@/lib/reasoning- /** Which model/provider pair a picker should mark "current". SessionView state * also drives the composer label, so a complete pair there wins over an older * `model.options` response. During initial hydration (or pre-session startup), - * options remain the fallback. Pick one complete pair before mixing fields so + * options remain the fallback. A valid owned scalar receipt is authoritative + * even before backend provider resolution. Pick one pair before mixing fields so * a model is never shown under a different provider. */ export function currentPickerSelection( - store: { model: string; provider: string }, + store: { model: string; provider: string; authoritative?: boolean }, options?: { model?: string; provider?: string } ): { model: string; provider: string } { const storeSelection = { @@ -19,7 +20,7 @@ export function currentPickerSelection( provider: String(options?.provider || '') } - if (storeSelection.model && storeSelection.provider) { + if (storeSelection.model && (storeSelection.provider || store.authoritative)) { return storeSelection } diff --git a/apps/desktop/src/plugins/hermes-bots/plugin.js b/apps/desktop/src/plugins/hermes-bots/plugin.js index 519e205382020..a217c6cb4ebbd 100644 --- a/apps/desktop/src/plugins/hermes-bots/plugin.js +++ b/apps/desktop/src/plugins/hermes-bots/plugin.js @@ -529,6 +529,7 @@ const GROUP_ACTIVITY_LABELS = { capped: 'turn stopped at the round/message cap', delivered: 'delivered a late reply', held: 'is held (stopped by you) — @mention it or say resume to release', + stopping: 'held the room — interruption was applied; waiting for turn retirement', 'stop-unconfirmed': 'held the room — interruption is unconfirmed; Stop can retry', stopped: 'stopped the room — remaining turns are held until resumed' } @@ -549,6 +550,7 @@ const GROUP_ACTIVITY_GLYPHS = { capped: 'debug-step-over', delivered: 'mail-read', held: 'debug-pause', + stopping: 'debug-stop', 'stop-unconfirmed': 'error', stopped: 'debug-stop' } @@ -1285,6 +1287,7 @@ async function pullGroupChatServerState(connectionId = groupChatSyncConnectionId preserveRooms: pending?.changedRooms || [], deletedRooms: pending?.deletedRooms || [] }) + rehomeGroupRuntimeOwners(merged) $groupChats.set(merged) await persistGroupChatRooms(merged) return true @@ -1376,6 +1379,7 @@ async function flushGroupChatServerSync(connectionId) { preserveRooms: pending?.changedRooms || [], deletedRooms: pending?.deletedRooms || [] }) + rehomeGroupRuntimeOwners(mergedRooms) $groupChats.set(mergedRooms) await persistGroupChatRooms(mergedRooms) } @@ -1412,6 +1416,7 @@ async function flushGroupChatServerSync(connectionId) { preserveRooms: pending?.changedRooms || [], deletedRooms: pending?.deletedRooms || [] }) + rehomeGroupRuntimeOwners(mergedRooms) $groupChats.set(mergedRooms) await persistGroupChatRooms(mergedRooms) } @@ -7285,9 +7290,9 @@ function setGroupChatImage(group, image) { }) } -/** Rename a group chat. The group's NAME is its identity everywhere — the - * room-map key, each local member's ui_meta membership list, and derived - * state — so a rename re-keys all of them. Member gateway sessions are kept +/** Rename a group chat. The group's name keys its presentation — the + * room map, each local member's ui_meta membership list, and derived + * state — so a rename re-keys them while retaining its room lifetime. Member gateway sessions are kept * as-is: stored sids keep resuming, so no history is lost. The room's * immutable roomId (the member-session title) is preserved across the * rename, so even a member whose sid is later lost falls back to the same @@ -7333,6 +7338,7 @@ async function renameGroupChat(oldName, newName, members) { all[next] = room } + rehomeGroupRuntimeOwners(all) $groupChats.set(all) const needs = { ...$groupNeedsYou.get() } @@ -7343,13 +7349,13 @@ async function renameGroupChat(oldName, newName, members) { $groupNeedsYou.set(needs) } - // Mirrored clarify cards key by group name; drop the old room's — the - // next poll re-mirrors any still-blocking question under the new name. - clearGroupClarify(oldName) + // Runtime owners/cards follow the same immutable lifetime above. A display + // rename neither replaces the accepted occurrence nor erases its question. // Local memberships: swap the name inside each member's canonical groups // list (syncs cross-machine via ui_meta). Remote members' seating lives in // the room record we just moved. + const metadataPersistence = [] for (const member of members || []) { if (!member?.name) { continue @@ -7358,11 +7364,20 @@ async function renameGroupChat(oldName, newName, members) { const meta = botRosterMeta(member, $botMeta.get()) || {} const groups = [...new Set(botGroups(meta).map(g => (g === oldName ? next : g)))] - await saveBotMeta(member, { groups, group: groups[0] || null }) + // saveBotMeta applies its local atom before yielding. Start every member's + // transition now so a later rename cannot leave old-name writes in this loop. + metadataPersistence.push(saveBotMeta(member, { groups, group: groups[0] || null })) } // Persist the re-keyed map (updateGroupChat writes the whole durable map). - updateGroupChat(next, r => r, { sync: false }) + const renamedRoom = updateGroupChat(next, r => { + // Idle legacy/metadata-only rooms need the same existing lifetime token + // as a runtime coordinator before this operation first yields. + if (!r.roomId && !r.coordinationId) r.coordinationId = groupChatEntryId() + return r + }, { sync: false }) + const lifetime = { group: next, roomId: renamedRoom.roomId || null, + roomToken: renamedRoom.coordinationId || null } // A rename is one revisioned state transition: the new identity is updated // and the old identity is tombstoned together, so cold hydration cannot // merge the pre-rename room back into the roster. @@ -7381,13 +7396,18 @@ async function renameGroupChat(oldName, newName, members) { openGroupChat(next) } + // Room, membership, persistence and navigation intent are all committed + // before waiting for metadata acknowledgement. No captured name is written + // by this continuation after an overlapping local/remote rename or disband. + await Promise.all(metadataPersistence) + // Same convergence as disband: drop the pre-rename roster snapshot so the // old name can't linger anywhere the fence doesn't cover. if (typeof queryClient !== 'undefined' && queryClient?.invalidateQueries) { queryClient.invalidateQueries({ queryKey: ROSTER_KEY }) } - return next + return groupNameForLifetime(lifetime) } function groupChatEntryId() { @@ -7532,7 +7552,7 @@ async function ensureGroupChatSession(group, member, requestMember = member, occ const stored = res.session_key || known if (stored) { - updateGroupChat(group, current => { + updateGroupChat(occurrence?.group || group, current => { if (occurrence && !groupOccurrenceCurrent(occurrence)) return current current.sessions = { ...(current.sessions || {}), [key]: stored } current.sessionOwners = { ...(current.sessionOwners || {}), [key]: groupSessionOwner(requestMember) } @@ -7560,7 +7580,7 @@ async function ensureGroupChatSession(group, member, requestMember = member, occ const stored = created?.stored_session_id || null if (stored) { - updateGroupChat(group, r => { + updateGroupChat(occurrence?.group || group, r => { if (occurrence && !groupOccurrenceCurrent(occurrence)) return r r.sessions = { ...(r.sessions || {}), [key]: stored } r.sessionOwners = { ...(r.sessionOwners || {}), [key]: groupSessionOwner(requestMember) } @@ -7673,8 +7693,21 @@ async function retainGroupTurnRoute(member) { * exactly once more. Returns the runtime id the submit actually landed on so * the poll loop keeps a live fallback target. */ async function submitGroupTurnPrompt(member, runtime, stored, text, occurrence, canSubmit) { + const submit = async target => { + if (occurrence) occurrence.admissionRefused = false + try { + return await requestForBot(member, 'prompt.submit', { session_id: target, text }) + } catch (error) { + // The gateway rejects a missing runtime before prompt admission. A + // cancelled remint/probe after that refusal owns no accepted turn. + if (occurrence && (error?.code === 4001 || error?.code === 4090 || error?.code === 'POOL_CAPACITY_EXCEEDED')) { + occurrence.admissionRefused = true + } + throw error + } + } try { - const ack = await requestForBot(member, 'prompt.submit', { session_id: runtime, text }) + const ack = await submit(runtime) return { runtime, acceptedTurn: groupAcceptedTurn(ack?.accepted_turn, runtime) } } catch (error) { @@ -7710,7 +7743,7 @@ async function submitGroupTurnPrompt(member, runtime, stored, text, occurrence, if (occurrence?.cancelled) await interruptGroupOccurrence(occurrence) throw error } - const ack = await requestForBot(member, 'prompt.submit', { session_id: fresh, text }) + const ack = await submit(fresh) return { runtime: fresh, acceptedTurn: groupAcceptedTurn(ack?.accepted_turn, fresh) } } @@ -7845,7 +7878,7 @@ async function waitForGroupTurnCollector(group, memberKey, collector, occurrence const epoch = $groupChats.get()[group]?.epoch || 0 let timer, unbind const cancelled = () => { - const room = $groupChats.get()[group] + const room = $groupChats.get()[occurrence?.group || group] return !room || room.tombstone || Boolean(room.holds?.[memberKey]) || (occurrence?.userIds ? !groupOccurrenceCanSubmit(occurrence) : (room.epoch || 0) !== epoch) } @@ -7874,6 +7907,10 @@ function groupTurnMarkerBlocksDispatch(room, memberKey, coordinator) { function consumeGroupTurnMarker(group, memberKey, marker, published = false) { if ($groupChats.get()[group]?.stranded?.[memberKey] !== marker) return false const owned = [...(groupRoomCoordinators.get(group)?.occurrences || [])].find(o => o.id === marker?.occurrence_id) + // session.interrupt targets the runtime's current turn when applied. Even + // exact predecessor terminal proof cannot let a successor reuse that runtime + // while an older hot generation or a cold marker's control is still pending. + if (groupOccurrenceHasPendingInterrupt(owned) || interruptingGroupTurnMarkers.has(marker)) return false const coldDrive = owned ? null : restoreGroupDrive(group, marker) let consumed = false updateGroupChat(group, room => { @@ -7901,7 +7938,8 @@ function consumeGroupTurnMarker(group, memberKey, marker, published = false) { } function groupTurnMarkerIntentIsCurrent(room, marker) { - if (!room || room.tombstone || room.holds?.[marker.delivery.member_key]) return false + if (!room || room.tombstone || marker.stop_requested || marker.hold_requested || + room.holds?.[marker.delivery.member_key]) return false const newerUser = groupTurnHasNewerUser(room, marker.thread, marker.user_ids, marker.anchor_id, marker.input_version) return shouldCommitMemberTurn(marker.epoch ?? 0, room.epoch || 0, newerUser, marker.input_version !== undefined || Array.isArray(marker.user_ids)) @@ -8004,7 +8042,7 @@ function syncGroupClarify(group, member, state, requestMember = member) { return true } -/** Drop every mirrored clarify belonging to `group` (disband/rename). */ +/** Drop every mirrored clarify belonging to `group` (disband/Stop). */ function clearGroupClarify(group) { const all = $groupClarify.get() const next = {} @@ -8098,11 +8136,15 @@ async function respondGroupClarify(entry, member, answers, occurrence, fence) { requireCurrentGroupQuestion(entry) attempted = true if (entry.kind === 'approval') { - await requestForBot(member, 'approval.respond', { + const result = await requestForBot(member, 'approval.respond', { session_id: entry.sessionId || undefined, request_id: entry.requestId, choice: typeof answers === 'string' && answers ? answers : 'deny' }) + if (!Number.isSafeInteger(result?.resolved) || result.resolved <= 0) { + throw new Error(result?.resolved === 0 ? 'The approval is no longer pending.' + : 'The approval acknowledgement is unconfirmed.') + } } else if (entry.questions && entry.questions.length) { for (const question of entry.questions) { const qid = question?.qid ?? question?.id @@ -8164,6 +8206,49 @@ const groupRoomCoordinators = new Map() const groupRuntimeSessionOwners = new Map() const GROUP_CHAT_PARALLEL_CEILING = 4 +function groupNameForLifetime(owner, rooms = $groupChats.get()) { + const matches = room => room && !room.tombstone && + (room.roomId || null) === owner.roomId && (room.coordinationId || null) === owner.roomToken + if (matches(rooms[owner.group])) return owner.group + // Legacy owners also need their minted coordination token. Never follow a + // display-name reuse or an unidentifiable room into a replacement lifetime. + if (!owner.roomId && !owner.roomToken) return null + const names = Object.keys(rooms).filter(name => matches(rooms[name])) + return names.length === 1 ? names[0] : null +} + +function rehomeGroupRuntimeOwners(rooms) { + for (const [oldName, coordinator] of [...groupRoomCoordinators]) { + const next = groupNameForLifetime(coordinator, rooms) + if (!next || next === oldName) continue + if (groupRoomCoordinators.get(oldName) === coordinator) groupRoomCoordinators.delete(oldName) + coordinator.group = next + groupRoomCoordinators.set(next, coordinator) + for (const occurrence of coordinator.occurrences) occurrence.group = next + for (const drive of coordinator.capturedDrives.values()) drive.group = next + } + const cards = {}, needs = { ...$groupNeedsYou.get() }, activity = { ...$groupActivity.get() } + let cardsChanged = false, needsChanged = false, activityChanged = false + for (const entry of Object.values($groupClarify.get())) { + const next = groupNameForLifetime({ group: entry.group, roomId: entry.roomId, roomToken: entry.roomToken }, rooms) + if (next && next !== entry.group) { + entry.group = next // preserve pending response/card identity + cardsChanged = true + } + cards[`${entry.group}::${entry.memberKey}`] = entry + } + for (const [oldName, room] of Object.entries($groupChats.get())) { + const next = groupNameForLifetime({ group: oldName, roomId: room.roomId || null, + roomToken: room.coordinationId || null }, rooms) + if (!next || next === oldName) continue + if (oldName in needs) { needs[next] = needs[oldName]; delete needs[oldName]; needsChanged = true } + if (oldName in activity) { activity[next] = activity[oldName]; delete activity[oldName]; activityChanged = true } + } + if (cardsChanged) $groupClarify.set(cards) + if (needsChanged) $groupNeedsYou.set(needs) + if (activityChanged) $groupActivity.set(activity) +} + function groupRoomCoordinator(group) { let room = $groupChats.get()[group] || {} if (!room.roomId && !room.coordinationId) { @@ -8197,7 +8282,7 @@ function registerGroupOccurrence(group, member, thread, deliveryResult = {}, pre ready: new Promise(resolve => { ready = resolve }) } occurrence.settle = settle occurrence.markReady = ready - occurrence.memberLock = groupSourceSessionKey(captured, `room:${room.roomId || group}`) + occurrence.memberLock = groupSourceSessionKey(captured, `room:${room.roomId || room.coordinationId || group}`) occurrence.inputEndId = room.log?.length ? groupChatSyncEntryKey(room.log.at(-1)) : null occurrence.inputVersion = room.threadInputVersions?.[occurrence.thread] || 0 occurrence.userIds = (room.log || []).filter(e => groupIsUserInstruction(e) && groupThreadOf(e) === occurrence.thread).map(groupChatSyncEntryKey) @@ -8207,11 +8292,16 @@ function registerGroupOccurrence(group, member, thread, deliveryResult = {}, pre return occurrence } -function groupOccurrenceCurrent(occurrence) { +function groupOccurrenceRoomIsCurrent(occurrence) { const room = $groupChats.get()[occurrence.group] - return !occurrence.cancelled && room && !room.tombstone && + return room && !room.tombstone && (room.roomId || null) === occurrence.roomId && - (room.coordinationId || null) === occurrence.roomToken && !room.holds?.[occurrence.memberKey] + (room.coordinationId || null) === occurrence.roomToken +} + +function groupOccurrenceCurrent(occurrence) { + return !occurrence.cancelled && groupOccurrenceRoomIsCurrent(occurrence) && + !$groupChats.get()[occurrence.group].holds?.[occurrence.memberKey] } function groupOccurrenceCanSubmit(occurrence) { @@ -8228,7 +8318,7 @@ function paintGroupOccurrences(coordinator) { (current.coordinationId || null) !== coordinator.roomToken) return const turns = [...coordinator.occurrences].filter(o => !o.released).map(o => ({ id: o.id, memberKey: o.memberKey, member: o.captured.member.name, - phase: o.cancelled ? (groupOccurrenceStopConfirmed(o) ? 'stopping' : 'stop-unconfirmed') : o.phase, thread: o.thread })) + phase: o.cancelled ? (groupOccurrenceInterruptApplied(o) ? 'stopping' : 'stop-unconfirmed') : o.phase, thread: o.thread })) updateGroupChat(coordinator.group, room => { room.turns = turns room.turn = turns.find(o => o.phase === 'running' || o.phase === 'starting')?.member || null @@ -8237,16 +8327,26 @@ function paintGroupOccurrences(coordinator) { } function groupInterruptConfirmed(reply) { + // This confirms application of cancellation, not execution retirement. return reply?.status === 'interrupted' } function interruptStoppedGroupMarker(group, memberKey, marker, target) { const prior = interruptingGroupTurnMarkers.get(marker) if (prior) return prior + const room = $groupChats.get()[group] + const lifetime = { group, roomId: room?.roomId || null, roomToken: room?.coordinationId || null } const pending = Promise.resolve().then(() => requestForBot(target, 'session.interrupt', { session_id: marker.delivery.accepted_turn.session_id })).then(reply => { const confirmed = groupInterruptConfirmed(reply) - if (confirmed) consumeGroupTurnMarker(group, memberKey, marker) + const currentGroup = groupNameForLifetime(lifetime) + if (confirmed && currentGroup) updateGroupChat(currentGroup, room => { + if (room.stranded?.[memberKey] === marker) { + marker.interrupt_applied = true + room.stranded = { ...room.stranded } + } + return room + }, { sync: false }) return confirmed }, () => false).finally(() => { if (interruptingGroupTurnMarkers.get(marker) === pending) interruptingGroupTurnMarkers.delete(marker) @@ -8256,10 +8356,19 @@ function interruptStoppedGroupMarker(group, memberKey, marker, target) { } function groupOccurrenceStopConfirmed(occurrence) { - if (occurrence.submissionPending || occurrence.answerPromise) return false + if (occurrence.preparationPending || occurrence.submissionPending || occurrence.answerPromise || groupOccurrenceHasPendingInterrupt(occurrence)) return false if (occurrence.terminalObserved) return true if (!occurrence.submitAttempted) return Boolean(occurrence.collectorDone) - return occurrence.interrupts.get(`${occurrence.runtime}::${occurrence.admissionVersion || 0}`)?.confirmed === true + return false // A matching interrupt ACK still needs exact terminal evidence. +} + +function groupOccurrenceHasPendingInterrupt(occurrence) { + return [...(occurrence?.interrupts?.values() || [])].some(attempt => attempt.pending) +} + +function groupOccurrenceInterruptApplied(occurrence) { + return !occurrence.preparationPending && !occurrence.submissionPending && !occurrence.answerPromise && + occurrence.interrupts.get(`${occurrence.runtime}::${occurrence.admissionVersion || 0}`)?.confirmed === true } function markGroupOccurrenceStop(occurrence) { @@ -8300,7 +8409,7 @@ async function interruptGroupOccurrence(occurrence, runtime = occurrence.runtime function finishGroupOccurrence(occurrence) { const c = occurrence.coordinator - if (occurrence.released) return + if (occurrence.released || groupOccurrenceHasPendingInterrupt(occurrence)) return occurrence.released = true occurrence.releaseLease?.() if (c.members.get(occurrence.memberLock) === occurrence) c.members.delete(occurrence.memberLock) @@ -8399,7 +8508,12 @@ function retainUnresolvedGroupOccurrence(occurrence) { return false } markGroupOccurrenceStop(occurrence) - } else if (!groupOccurrenceCurrent(occurrence) || marker?.occurrence_id !== occurrence.id) return false + if (groupOccurrenceRoomIsCurrent(occurrence) && groupTurnDeliveryKey(marker?.delivery)) { + // Late admission/collector settlement may supply the first exact receipt + // after Stop. Reconcile it without making control completion wait on a read. + void harvestStrandedGroupReply(occurrence.group, occurrence.captured.member).catch(() => undefined) + } + } else if (!groupOccurrenceRoomIsCurrent(occurrence) || marker?.occurrence_id !== occurrence.id) return false // A timeout/unavailable projection is not capacity evidence. Waiting has // already surrendered only its worker; running/unknown keeps that worker. occurrence.markReady() @@ -8408,6 +8522,7 @@ function retainUnresolvedGroupOccurrence(occurrence) { } async function executeGroupOccurrence(occurrence) { + occurrence.preparationPending = true try { const release = await retainGroupTurnRoute(occurrence.captured.requestMember) let released = false @@ -8419,6 +8534,7 @@ async function executeGroupOccurrence(occurrence) { return await runGroupChatMemberTurnLeased(occurrence.group, occurrence.captured, occurrence.prompt, occurrence.thread, occurrence.images, occurrence.deliveryResult, occurrence) } finally { + occurrence.preparationPending = false // Stop owns any interrupt through its acknowledgement, including a runtime // that became available during acquisition. Successors cannot start yet. if (occurrence.answerPromise) await occurrence.answerPromise.catch(() => undefined) @@ -8462,6 +8578,7 @@ async function runGroupChatMemberTurnLeased(group, captured, prompt, thread, ima return null } while (true) { + group = occurrence?.group || group const current = $groupChats.get()[group] || {} const prior = current.stranded?.[memberKey] const priorCollector = prior && collectingGroupTurnMarkers.get(prior) @@ -8480,6 +8597,7 @@ async function runGroupChatMemberTurnLeased(group, captured, prompt, thread, ima // Another waiter may have reserved the member. Recheck synchronously // before any session preparation, attachments, or prompt admission. } + group = occurrence?.group || group const roomAtDispatch = $groupChats.get()[group] || {} const dispatchEpoch = roomAtDispatch.epoch || 0 let marker = { delivery: { accepted_turn: null, member_key: memberKey, owner }, @@ -8499,6 +8617,7 @@ async function runGroupChatMemberTurnLeased(group, captured, prompt, thread, ima let submitAttempted = false try { const prepared = await ensureGroupChatSession(group, member, requestMember, occurrence) + group = occurrence?.group || group let { runtime } = prepared const { stored } = prepared if (occurrence) { @@ -8513,6 +8632,7 @@ async function runGroupChatMemberTurnLeased(group, captured, prompt, thread, ima if (occurrence.cancelled) await interruptGroupOccurrence(occurrence) } const beforeSubmit = () => { + group = occurrence?.group || group const room = $groupChats.get()[group] || {} return (!occurrence || groupOccurrenceCanSubmit(occurrence)) && room.stranded?.[memberKey] === marker && !room.tombstone && !((room.epoch || 0) !== dispatchEpoch && room.holds?.[memberKey]) @@ -8522,6 +8642,7 @@ async function runGroupChatMemberTurnLeased(group, captured, prompt, thread, ima return discarded() } runtime = await requireGroupTurnProtocol(requestMember, prepared) + group = occurrence?.group || group if (occurrence && runtime !== occurrence.runtime) { const sessionLock = groupSourceSessionKey(captured, runtime) const prior = groupRuntimeSessionOwners.get(sessionLock) @@ -8541,7 +8662,10 @@ async function runGroupChatMemberTurnLeased(group, captured, prompt, thread, ima consumeGroupTurnMarker(group, memberKey, marker) return discarded() } - if (occurrence) { occurrence.phase = 'running'; paintGroupOccurrences(occurrence.coordinator) } + if (occurrence) { + occurrence.preparationPending = false + occurrence.phase = 'running'; paintGroupOccurrences(occurrence.coordinator) + } recordGroupActivity(group, { kind: 'working', member: member.name, thread }) const fileRefs = [] for (const img of Array.isArray(images) ? images : []) { @@ -8583,6 +8707,7 @@ async function runGroupChatMemberTurnLeased(group, captured, prompt, thread, ima if (occurrence.cancelled) { markGroupOccurrenceStop(occurrence); await interruptGroupOccurrence(occurrence) } } } + group = occurrence?.group || group const previous = marker marker = { ...marker, runtime: submitted.runtime, delivery: { ...marker.delivery, accepted_turn: submitted.acceptedTurn } } @@ -8605,15 +8730,11 @@ async function runGroupChatMemberTurnLeased(group, captured, prompt, thread, ima let progress = 'working' while (Date.now() < deadline) { await new Promise(resolve => setTimeout(resolve, GROUP_TURN_POLL_MS)) + group = occurrence?.group || group const roomDuringPoll = $groupChats.get()[group] || {} if (roomDuringPoll.stranded?.[memberKey] !== marker) return discarded() - if (occurrence && !groupOccurrenceCurrent(occurrence)) { + if (occurrence && (occurrence.cancelled || !groupOccurrenceRoomIsCurrent(occurrence))) { if (occurrence.cancelled) markGroupOccurrenceStop(occurrence) - else consumeGroupTurnMarker(group, memberKey, marker) - return discarded() - } - if ((roomDuringPoll.epoch || 0) !== dispatchEpoch && roomDuringPoll.holds?.[memberKey]) { - consumeGroupTurnMarker(group, memberKey, marker) return discarded() } let state @@ -8624,28 +8745,21 @@ async function runGroupChatMemberTurnLeased(group, captured, prompt, thread, ima session_id: submitted.acceptedTurn.session_id, profile: member.name, accepted_turn: submitted.acceptedTurn }) } catch (error) { + group = occurrence?.group || group const roomAfterError = $groupChats.get()[group] || {} if (roomAfterError.stranded?.[memberKey] !== marker) return discarded() if (occurrence?.cancelled) { markGroupOccurrenceStop(occurrence); return discarded() } - if ((roomAfterError.epoch || 0) !== dispatchEpoch && roomAfterError.holds?.[memberKey]) { - consumeGroupTurnMarker(group, memberKey, marker) - return discarded() - } if (!groupTurnMarkerIntentIsCurrent(roomAfterError, marker)) return discarded() // Observation failure grants no replay or resume authority. Surface it // now and retain the accepted receipt for an explicit later harvest. throw groupTurnOutcomeError({ state: 'unavailable', reason: `Could not observe member turn: ${error?.message || 'gateway poll failed'}` }) } + group = occurrence?.group || group const roomAfterResume = $groupChats.get()[group] || {} if (roomAfterResume.stranded?.[memberKey] !== marker) return discarded() - if (occurrence && !groupOccurrenceCurrent(occurrence)) { + if (occurrence && (occurrence.cancelled || !groupOccurrenceRoomIsCurrent(occurrence))) { if (occurrence.cancelled) markGroupOccurrenceStop(occurrence) - else consumeGroupTurnMarker(group, memberKey, marker) - return discarded() - } - if ((roomAfterResume.epoch || 0) !== dispatchEpoch && roomAfterResume.holds?.[memberKey]) { - consumeGroupTurnMarker(group, memberKey, marker) return discarded() } const outcome = readGroupTurnOutcome(state, marker.delivery) @@ -8689,7 +8803,9 @@ async function runGroupChatMemberTurnLeased(group, captured, prompt, thread, ima } paintGroupOccurrences(occurrence.coordinator) } else if (occurrence) { - if (occurrence.phase === 'waiting') await reserveGroupResumeWorker(occurrence) + if (occurrence.phase === 'waiting' && !['complete', 'error', 'interrupted'].includes(outcome.state)) { + await reserveGroupResumeWorker(occurrence) + } if (occurrence.resumeFence === resumeFenceAtPoll && responseAcknowledgedAtPoll) occurrence.resumeFence = null paintGroupOccurrences(occurrence.coordinator) } @@ -8713,8 +8829,9 @@ async function runGroupChatMemberTurnLeased(group, captured, prompt, thread, ima // Observation expiry cannot invalidate a retained unavailable child card. return null } catch (error) { - if (!submitAttempted || error?.code === 4090 || error?.code === 'POOL_CAPACITY_EXCEEDED') { - if (error?.code === 4090 || error?.code === 'POOL_CAPACITY_EXCEEDED') { + group = occurrence?.group || group + if (!submitAttempted || occurrence?.admissionRefused || error?.code === 4090 || error?.code === 'POOL_CAPACITY_EXCEEDED') { + if (occurrence?.admissionRefused || error?.code === 4090 || error?.code === 'POOL_CAPACITY_EXCEEDED') { if (occurrence) occurrence.submitAttempted = false error.data = { ...error.data, outcomeState: 'admission-refused', reason: error.message } } @@ -8742,9 +8859,11 @@ async function runGroupChatMemberTurnLeased(group, captured, prompt, thread, ima async function harvestStrandedGroupReply(group, member) { const memberKey = groupMemberKey(member) const room = $groupChats.get()[group] || {} + const lifetime = { group, roomId: room.roomId || null, roomToken: room.coordinationId || null } const marker = room.stranded?.[memberKey] - if (marker === undefined || collectingGroupTurnMarkers.has(marker) || - (room.holds?.[memberKey] && !marker?.stop_requested)) return + const reconciliationOnly = marker?.stop_requested || marker?.hold_requested + if (marker === undefined || (collectingGroupTurnMarkers.has(marker) && !reconciliationOnly) || + (room.holds?.[memberKey] && !reconciliationOnly)) return if (marker?.room_token && marker.room_token !== room.coordinationId) { reportUnavailableGroupTurn(group, member, marker, 'Legacy room lifetime is unavailable after reload or replacement') return @@ -8772,23 +8891,25 @@ async function harvestStrandedGroupReply(group, member) { session_id: marker.delivery.accepted_turn.session_id, profile: captured.member.name, accepted_turn: marker.delivery.accepted_turn }) } catch (error) { + group = groupNameForLifetime(lifetime) || group const current = $groupChats.get()[group] if (current?.stranded?.[memberKey] !== marker || !groupTurnMarkerIntentIsCurrent(current, marker)) return reportUnavailableGroupTurn(group, member, marker, `Could not observe member turn: ${error?.message || 'gateway poll failed'}`) return } + group = groupNameForLifetime(lifetime) || group if ($groupChats.get()[group]?.stranded?.[memberKey] !== marker || ownedOccurrence?.answerPromise) return const outcome = readGroupTurnOutcome(state, marker.delivery) const freshResumeRead = !ownedOccurrence?.resumeFence || (ownedOccurrence.resumeFence === resumeFenceAtRead && responseAcknowledgedAtRead) - if (marker.stop_requested || ownedOccurrence?.cancelled) { - // A stopped receipt is reconciliation only, including after reload. No + if (reconciliationOnly || ownedOccurrence?.cancelled) { + // A stopped/held receipt is reconciliation only, including after reload. No // reply, attention, input consumption, clarify card or prompt replay. if (['complete', 'error', 'interrupted'].includes(outcome.state) && !ownedOccurrence?.submissionPending && freshResumeRead) { if (ownedOccurrence) ownedOccurrence.terminalObserved = true - consumeGroupTurnMarker(group, memberKey, marker) + if (consumeGroupTurnMarker(group, memberKey, marker) && ownedOccurrence) settleGroupPublication(ownedOccurrence) } else if (ownedOccurrence && outcome.state === 'waiting' && !ownedOccurrence.submissionPending && freshResumeRead) { ownedOccurrence.resumeFence = null @@ -9107,6 +9228,7 @@ function groupTurnHasNewerUser(room, thread, userIds, anchorId, inputVersion) { * settle, so an old cleanup can never interrupt a replacement occurrence. */ async function stopGroupThread(group, thread, members = null) { const room = $groupChats.get()[group] || {} + const lifetime = { group, roomId: room.roomId || null, roomToken: room.coordinationId || null } const coordinator = groupRoomCoordinators.get(group) const occurrences = coordinator && coordinator.roomId === (room.roomId || null) && coordinator.roomToken === (room.coordinationId || null) ? [...coordinator.occurrences] : [] @@ -9141,7 +9263,7 @@ async function stopGroupThread(group, thread, members = null) { } if (coordinator) pumpGroupOccurrences(coordinator) const interrupts = occurrences.map(occurrence => interruptGroupOccurrence(occurrence, occurrence.runtime, true)) - let coldUnconfirmed = 0 + let legacyUnconfirmed = 0 // A cold reload has no runtime registry: accepted durable receipts still // identify exact runtime and source, never a display-name roster guess. for (const [memberKey, marker] of Object.entries(room.stranded || {})) { @@ -9155,11 +9277,10 @@ async function stopGroupThread(group, thread, members = null) { if (r.stranded?.[memberKey] === marker) { marker.stop_requested = true; r.stranded = { ...r.stranded } } return r }, { sync: false }) - if (!groupTurnDeliveryKey(marker?.delivery)) { coldUnconfirmed++; continue } + if (!groupTurnDeliveryKey(marker?.delivery)) continue const owner = marker.delivery.owner const target = Object.freeze({ ...owner, ...(owner.route ? { route: Object.freeze({ ...owner.route }) } : {}) }) - interrupts.push(interruptStoppedGroupMarker(group, memberKey, marker, target) - .then(confirmed => { if (!confirmed) coldUnconfirmed++ })) + interrupts.push(interruptStoppedGroupMarker(group, memberKey, marker, target)) } // Legacy transient room.turn has no acceptance proof. Retain its old Stop // behavior only when the name is unambiguous and no occurrences exist. @@ -9169,30 +9290,67 @@ async function stopGroupThread(group, thread, members = null) { const member = candidates[0] const sid = room.sessions?.[groupMemberKey(member)] if (sid) interrupts.push(Promise.resolve().then(() => requestForBot(member, 'session.interrupt', { session_id: sid })) - .then(reply => { if (!groupInterruptConfirmed(reply)) coldUnconfirmed++ }, () => { coldUnconfirmed++ })) + .then(() => { legacyUnconfirmed++ }, () => { legacyUnconfirmed++ })) } } for (const interrupt of interrupts) await interrupt + const missingLifetime = () => { + const pending = Math.max(1, occurrences.filter(o => !o.released && !groupOccurrenceStopConfirmed(o)).length) + return { status: 'unconfirmed', unconfirmed: pending, pending } + } + group = groupNameForLifetime(lifetime) + if (!group) return missingLifetime() + // The ACK only applied cancellation. The existing exact-turn collector owns + // retirement; do not infer a vacant worker from a successful control write. + for (const [memberKey, marker] of Object.entries($groupChats.get()[group]?.stranded || {})) { + if (!marker?.stop_requested || !groupTurnDeliveryKey(marker.delivery)) continue + if ($groupChats.get()[group]?.stranded?.[memberKey] !== marker) continue + const occurrence = occurrences.find(o => o.id === marker.occurrence_id) + const owner = marker.delivery.owner + const target = occurrence?.captured.member || Object.freeze({ ...owner, + ...(owner.route ? { route: Object.freeze({ ...owner.route }) } : {}) }) + // A slow/unavailable poll is observation, not part of the interrupt ACK. + // Its exact terminal may retire custody later through the same collector. + void harvestStrandedGroupReply(group, target).catch(() => undefined) + } for (const occurrence of occurrences) { - const marker = $groupChats.get()[group]?.stranded?.[occurrence.memberKey] if (occurrence.collectorDone) { if (occurrence.answerPromise) await occurrence.answerPromise.catch(() => undefined) await interruptGroupOccurrence(occurrence) + group = groupNameForLifetime(lifetime) + if (!group) return missingLifetime() + const settledMarker = $groupChats.get()[group]?.stranded?.[occurrence.memberKey] + if (settledMarker?.occurrence_id === occurrence.id && groupTurnDeliveryKey(settledMarker.delivery)) { + // A late human response advanced the generation after the first read + // was fenced. Observe only after its producer and final interrupt settle. + void harvestStrandedGroupReply(group, occurrence.captured.member).catch(() => undefined) + } if (groupOccurrenceStopConfirmed(occurrence)) { + const marker = $groupChats.get()[group]?.stranded?.[occurrence.memberKey] if (marker?.occurrence_id === occurrence.id) consumeGroupTurnMarker(group, occurrence.memberKey, marker) finishGroupOccurrence(occurrence) } } } - const unconfirmed = coldUnconfirmed + occurrences.filter(o => !o.released && !groupOccurrenceStopConfirmed(o)).length + const pendingMarkers = Object.values($groupChats.get()[group]?.stranded || {}).filter(marker => marker?.stop_requested) + const unconfirmed = legacyUnconfirmed + pendingMarkers.filter(marker => { + const occurrence = occurrences.find(o => o.id === marker.occurrence_id) + return occurrence ? !groupOccurrenceInterruptApplied(occurrence) : !marker.interrupt_applied + }).length + occurrences.filter(o => !o.released && !groupOccurrenceStopConfirmed(o) && + !pendingMarkers.some(marker => marker.occurrence_id === o.id) && !groupOccurrenceInterruptApplied(o)).length + const pending = pendingMarkers.length + occurrences.filter(o => !o.released && !groupOccurrenceStopConfirmed(o) && + !pendingMarkers.some(marker => marker.occurrence_id === o.id)).length if (coordinator) paintGroupOccurrences(coordinator) - const result = { status: unconfirmed ? 'unconfirmed' : 'stopped', unconfirmed } + const result = { status: unconfirmed ? 'unconfirmed' : pending ? 'stopping' : 'stopped', unconfirmed, pending } const current = $groupChats.get()[group] if (current && (current.roomId || null) === (room.roomId || null) && (current.coordinationId || null) === (room.coordinationId || null)) { - recordGroupActivity(group, { kind: unconfirmed ? 'stop-unconfirmed' : 'stopped', + recordGroupActivity(group, { kind: unconfirmed ? 'stop-unconfirmed' : pending ? 'stopping' : 'stopped', member: 'You', thread: thread || null, epoch: stoppedRoom.epoch }) } + if (pending && typeof window !== 'undefined') { + void harvestStrandedUntilSettled(group, occurrences.map(o => o.captured.member).concat(roster), thread) + } return result } @@ -9280,6 +9438,10 @@ function captureGroupDrive(group, members, thread) { /** A frozen round, never a raw Promise.all over roster members. The room * coordinator is the sole start owner; each completion publishes immediately. */ function runGroupChatRounds(group, members, thread, capturedDrive) { + if (capturedDrive) { + group = groupNameForLifetime({ ...capturedDrive, group }) + if (!group) return Promise.resolve() // deletion/replacement cannot revive a captured drive + } const drive = capturedDrive || captureGroupDrive(group, members, thread) const coordinator = groupRoomCoordinator(group) const key = groupDriveKey(drive) @@ -9295,6 +9457,7 @@ function runGroupChatRounds(group, members, thread, capturedDrive) { const running = Promise.resolve().then(() => driveFrozenGroupRounds(group, members, drive, coordinator)) coordinator.drives.set(key, running) void running.finally(() => { + group = coordinator.group if (coordinator.drives.get(key) === running) coordinator.drives.delete(key) drive.running = false if (!Object.values($groupChats.get()[group]?.stranded || {}).some(m => m?.drive_key === key)) { @@ -9340,6 +9503,7 @@ async function driveFrozenGroupRounds(group, members, drive, coordinator) { const { thread } = drive members = drive.members.map(captured => captured.member) const sameRoom = () => { + group = coordinator.group const room = $groupChats.get()[group] return room && !room.tombstone && (room.roomId || null) === drive.roomId && (room.coordinationId || null) === drive.roomToken @@ -9421,6 +9585,7 @@ async function driveFrozenGroupRounds(group, members, drive, coordinator) { job.phase = 'queued' const complete = runGroupChatMemberTurn(group, job.captured.member, job.prompt, thread, job.images, job.deliveryResult, job).then(reply => { + if (!sameRoom()) return const room = $groupChats.get()[group] if (job.cancelled) return // Stop owns the confirmation/unconfirmed activity. if (!sameRoom() || job.deliveryResult.discarded || @@ -9441,6 +9606,7 @@ async function driveFrozenGroupRounds(group, members, drive, coordinator) { if (job.wasWaiting) scheduleGroupDriveContinuation(group, drive) } }, error => { + if (!sameRoom()) return const room = $groupChats.get()[group] if (!sameRoom() || job.cancelled || !shouldCommitMemberTurn(job.epoch, room.epoch || 0, groupTurnHasNewerUser(room, thread, job.userIds, job.inputEndId, job.inputVersion), true)) return @@ -9539,10 +9705,15 @@ async function driveFrozenGroupRounds(group, members, drive, coordinator) { async function harvestStrandedUntilSettled(group, members, thread) { const HARVEST_INTERVAL_MS = 5000 const HARVEST_MAX_TRIES = 60 + const initial = $groupChats.get()[group] + if (!initial) return + const lifetime = { group, roomId: initial.roomId || null, roomToken: initial.coordinationId || null } for (let attempt = 0; attempt < HARVEST_MAX_TRIES; attempt++) { await new Promise(resolve => window.setTimeout(resolve, HARVEST_INTERVAL_MS)) + group = groupNameForLifetime(lifetime) + if (!group) return const room = $groupChats.get()[group] if (!room || room.running) { @@ -9610,29 +9781,35 @@ function sendToGroupChat(group, members, text, thread, images) { { at: sent?.at, byMessageId: sent?.id, thread: target }, members.map(member => groupMemberKey(member)) ) + for (const [memberKey, marker] of Object.entries(room.stranded || {})) { + if (room.holds?.[memberKey] && marker && typeof marker === 'object') { + // A future-turn hold suppresses this result; it is not death evidence. + // Keep the live collector/CAS identity and persist reconciliation intent. + marker.hold_requested = true + room.stranded = { ...room.stranded } + } + } return room }) recordGroupActivity(group, { kind: 'queued', member: 'You', thread: target }) const capturedDrive = captureGroupDrive(group, members, target) + const clearFailedDrive = () => { + const currentGroup = groupNameForLifetime({ ...capturedDrive, group }) + if (!currentGroup) return + updateGroupChat(currentGroup, room => { + if ((room.epoch || 0) === capturedDrive.epoch) room.running = false + return room + }) + } if (!wasRunning) { - void runGroupChatRounds(group, members, target, capturedDrive).catch(() => { - updateGroupChat(group, r => { - if ((r.roomId || null) === capturedDrive.roomId && (r.epoch || 0) === capturedDrive.epoch) r.running = false - return r - }) - }) + void runGroupChatRounds(group, members, target, capturedDrive).catch(clearFailedDrive) } else { // A loop is live; it bails at its next boundary. Chain the fresh loop // after a short settle so exactly one drive owns the room. setTimeout(() => { - void runGroupChatRounds(group, members, target, capturedDrive).catch(() => { - updateGroupChat(group, r => { - if ((r.roomId || null) === capturedDrive.roomId && (r.epoch || 0) === capturedDrive.epoch) r.running = false - return r - }) - }) + void runGroupChatRounds(group, members, target, capturedDrive).catch(clearFailedDrive) }, 250) } @@ -13498,7 +13675,11 @@ function GroupImageControls({ image, onImage, seedName, seedMembers }) { * the room record. Both apply on Save so a cancelled dialog changes nothing. */ function GroupChatSettingsDialog({ group, members, open, onClose, onRenamed }) { const rooms = useValue($groupChats) - const current = (rooms[group] || {}).image || null + const displayedRoom = rooms[group] + const displayedLifetime = { group, roomId: displayedRoom?.roomId || null, + roomToken: displayedRoom?.coordinationId || null } + const metadataOnlyGroup = !displayedRoom && knownGroups($botMeta.get()).includes(group) + const current = (displayedRoom || {}).image || null const [name, setName] = useState(group) const [image, setImage] = useState(current) @@ -13511,11 +13692,31 @@ function GroupChatSettingsDialog({ group, members, open, onClose, onRenamed }) { }, [open, group]) const save = async () => { - const finalName = await renameGroupChat(group, name, members) + let savingGroup = group + if (displayedRoom) { + savingGroup = groupNameForLifetime(displayedLifetime) + // A button rendered for a deleted/replaced room cannot adopt today's + // same-name record. Unidentified legacy rooms require the same object + // until their existing canonical lifetime token is minted below. + if (!savingGroup || (!displayedLifetime.roomId && !displayedLifetime.roomToken && + $groupChats.get()[group] !== displayedRoom)) return + } else if (!metadataOnlyGroup || $groupChats.get()[group] || !knownGroups($botMeta.get()).includes(group)) { + return + } + const original = $groupChats.get()[savingGroup] + if (original?.tombstone) return + const savingRoom = original?.roomId || original?.coordinationId ? original : updateGroupChat(savingGroup, room => { + room.coordinationId = groupChatEntryId() + return room + }, { sync: false }) + const lifetime = { group: savingGroup, roomId: savingRoom.roomId || null, roomToken: savingRoom.coordinationId || null } + const renamed = await renameGroupChat(savingGroup, name, members) - if (finalName === null) { - return // collision — dialog stays open for a different name + if (renamed === null) { + return // collision or removed lifetime — the dialog cannot update another room } + const finalName = groupNameForLifetime(lifetime) + if (finalName === null) return if (image !== current) { setGroupChatImage(finalName, image) @@ -14009,7 +14210,6 @@ function GroupMentionInput({ members, onChange, onSubmitDraft, value, ...inputPr * (once/session/always/deny) as buttons — no free text; approvals are a * closed choice. Answer sends via the member's own source. */ function GroupClarifyCard({ entry, members }) { - const { group } = entry const isApproval = entry.kind === 'approval' const member = members.find(m => groupMemberKey(m) === entry.memberKey) || members.find(m => m.name === entry.member) const [drafts, setDrafts] = useState({}) @@ -14062,7 +14262,10 @@ function GroupClarifyCard({ entry, members }) { : questions .map(q => (questions.length > 1 ? `${q.question}: ${answerFor(q)}` : answerFor(q))) .join('\n') - appendGroupChatEntry(group, { kind: 'user', name: 'You' }, summary, entry.thread || 'legacy', + const currentRoom = $groupChats.get()[entry.group] + if (!currentRoom || currentRoom.tombstone || (currentRoom.roomId || null) !== entry.roomId || + (currentRoom.coordinationId || null) !== entry.roomToken) return + appendGroupChatEntry(entry.group, { kind: 'user', name: 'You' }, summary, entry.thread || 'legacy', undefined, undefined, entry.receipt?.delivery) } catch (err) { host.notify({ kind: 'error', message: `Could not send the answer to @${botHandle(entry.member, member)}: ${err?.message || err}` }) @@ -14587,7 +14790,9 @@ function GroupChatWorkspace({ group, members, onBack, visible = true }) { const result = await stopGroupThread(group, latestActivity?.thread || null, memberDescriptors()) host.notify(result.status === 'stopped' ? { kind: 'success', message: `Stopped ${group} — remaining turns are held until you resume` } - : { kind: 'info', message: `Held ${group} — ${result.unconfirmed} interruption(s) are unconfirmed. Stop can retry.` }) + : result.status === 'stopping' + ? { kind: 'info', message: `Stopping ${group} — waiting for the remaining turns to finish.` } + : { kind: 'info', message: `Held ${group} — ${result.unconfirmed} interruption(s) are unconfirmed. Stop can retry.` }) } const activityPanel = jsxs('div', { @@ -16733,7 +16938,7 @@ const groupTurnRuntime = { groupMemberKey, updateGroupChat, groupBlockedMembers, GroupBlockedNotice, CreateGroupChatDialog, createFreshGroupChat, groupComposerDraftKey, groupComposerDraftSnapshot, updateGroupComposerDraft, GroupChatWorkspace, GroupClarifyCard, - groupRoomCanStop, + groupRoomCanStop, renameGroupChat, pullGroupChatServerState, GroupChatSettingsDialog, $groupChatWorkspace, bindGroupTurnPorts(ports) { ({ Date, setTimeout, clearTimeout, setInterval, clearInterval, document } = createGroupTurnPorts(ports)) }, @@ -17339,5 +17544,3 @@ export default { }) } } - - diff --git a/apps/desktop/src/plugins/hermes-bots/tests/group-chat.test.mjs b/apps/desktop/src/plugins/hermes-bots/tests/group-chat.test.mjs index d646bc6b7a18a..50b2c73887128 100644 --- a/apps/desktop/src/plugins/hermes-bots/tests/group-chat.test.mjs +++ b/apps/desktop/src/plugins/hermes-bots/tests/group-chat.test.mjs @@ -193,7 +193,7 @@ function load(turnScript, { busyUntilResumeCall, clarifyUntilResumeCall, approva } if (method === 'approval.respond') { approvalResponds.push({ ...params }) - return { resolved: true } + return { resolved: 1 } } return {} }, diff --git a/apps/desktop/src/plugins/hermes-bots/tests/group-lifecycle-boundaries.test.mjs b/apps/desktop/src/plugins/hermes-bots/tests/group-lifecycle-boundaries.test.mjs new file mode 100644 index 0000000000000..f8e78e0106762 --- /dev/null +++ b/apps/desktop/src/plugins/hermes-bots/tests/group-lifecycle-boundaries.test.mjs @@ -0,0 +1,580 @@ +import assert from 'node:assert/strict' +import test from 'node:test' +import { harness, members, deferred, flush, drive } from './stop-custody-harness.mjs' + +// Offline only: the shipped plugin runs against fake gateway/profile ports and +// a bounded fake clock. No provider, backend process or installed state is used. +const bounded = (name, body) => test(name, { timeout: 10000 }, body) +const clone = value => JSON.parse(JSON.stringify(value)) +const receipt = (h, member = h.roster[0], name = 'Room') => + h.gc.$groupChats.get()[name]?.stranded?.[h.gc.groupMemberKey(member)] +const coordinator = (h, name = 'Room') => h.gc.groupRoomCoordinators.get(name) +const posts = (h, name = 'Room') => h.gc.$groupChats.get()[name]?.log.filter(e => e.from.kind === 'member') || [] +const sessionFor = (h, profile = 'bot1') => [...h.sessions.values()].find(s => s.profile === profile) +const cards = h => Object.values(h.gc.$groupClarify.get()) + +async function settle(h, pending, name = 'Room') { + for (let i = 0; i < 30 && coordinator(h, name)?.occurrences.size; i++) { + for (const session of h.sessions.values()) { session.state = 'complete'; session.pending = null } + await h.advance() + for (const member of h.roster) await h.gc.harvestStrandedGroupReply(name, member) + } + if (pending) await pending +} + +for (const terminal of ['complete', 'error', 'interrupted']) { + bounded(`GC-TERM: saturated waiting collector retires ${terminal} without a worker reservation`, async () => { + const h = await harness(members(6)), pending = drive(h) + await flush() + assert.equal(h.rpc('prompt.submit').length, 4) + const parked = sessionFor(h), accepted = clone(parked.ref) + parked.state = 'waiting'; parked.pending = { request_id: 'question-1', question: 'Choose?' } + await h.advance() + assert.equal(h.rpc('prompt.submit').length, 5, 'the fifth member uses the positively waiting slot') + assert.equal(coordinator(h).active, 4) + assert.equal(coordinator(h).queue.length, 1) + parked.state = terminal; parked.pending = null; parked.text = 'TERMINAL_WITH_ALL_SLOTS_BUSY' + await h.advance(); await flush() + assert.equal(receipt(h), undefined, 'exact terminal does not wait behind the four busy workers') + assert.equal(h.leases.find(l => l.route.targetProfile === 'bot1').releases, 1) + assert.equal(coordinator(h).active, 4) + assert.equal(h.rpc('prompt.submit').length, 5, 'terminal collection does not manufacture a sixth worker') + assert.deepEqual(parked.ref, accepted) + assert.equal(posts(h).filter(e => e.text === 'TERMINAL_WITH_ALL_SLOTS_BUSY').length, terminal === 'complete' ? 1 : 0) + await settle(h, pending) + assert.equal(coordinator(h).active, 0) + assert.ok(h.leases.every(l => l.releases === 1)) + }) +} + +bounded('GC-HOLD: typed hold retains running custody and release never replays or publishes the held turn', async () => { + const h = await harness(members(1)), pending = drive(h) + await flush() + const first = sessionFor(h), accepted = clone(first.ref), lease = h.leases[0] + h.gc.sendToGroupChat('Room', h.roster, '@bot1 stop', 'hold-thread') + await h.advance(); await pending + assert.deepEqual(receipt(h)?.delivery.accepted_turn, accepted) + assert.equal(lease.releases, 0, 'a future-turn hold is not terminal evidence') + assert.equal(coordinator(h).active, 1) + assert.equal(h.rpc('session.interrupt').length, 0) + assert.equal(posts(h).length, 0) + h.gc.sendToGroupChat('Room', h.roster, '@bot1 continue with NEW_INPUT', 'new-thread') + await h.advance() + assert.equal(h.rpc('prompt.submit').length, 1, 'accepted work owns the member until exact retirement') + assert.equal(lease.releases, 0) + first.state = 'complete'; first.pending = null; first.text = 'HELD_OLD_RESULT' + await h.gc.harvestStrandedGroupReply('Room', h.roster[0]); await flush(); await h.advance() + assert.equal(posts(h).filter(e => e.text === 'HELD_OLD_RESULT').length, 0) + assert.equal(lease.releases, 1) + // An instruction shown as blocked requires a fresh explicit drive. It is + // not silently replayed when an unknown old turn later becomes terminal. + void h.gc.runGroupChatRounds('Room', h.roster, 'new-thread') + await flush(); await h.advance() + assert.equal(h.rpc('prompt.submit').length, 2, 'the new instruction admits once after retirement') + assert.ok(h.rpc('prompt.submit')[1].params.text.includes('NEW_INPUT')) + await settle(h) +}) + +for (const unavailable of [false, true]) { + bounded(`GC-ACK: application ACK retains hot capacity while outcome is ${unavailable ? 'unavailable' : 'running'}`, async () => { + const options = { interruptReply: { status: 'interrupted' }, unavailable } + const h = await harness(members(6), options), pending = drive(h) + await flush() + const accepted = [...h.sessions.values()].map(s => clone(s.ref)) + const result = await h.gc.stopGroupThread('Room', 't1', h.roster) + await h.advance(); await pending + assert.equal(result.status, 'stopping') + assert.equal(result.unconfirmed, 0) + assert.equal(Object.keys(h.room().stranded).length, 4) + assert.equal(coordinator(h).active, 4) + assert.equal(h.activeLeases(), 4) + assert.ok(h.leases.every(l => l.releases === 0)) + const repeats = await h.gc.stopGroupThread('Room', 't1', h.roster) + assert.equal(repeats.status, 'stopping') + assert.equal(h.rpc('session.interrupt').length, 4, 'confirmed application need not repeat the control write') + h.gc.sendToGroupChat('Room', h.roster, '@all resume NEW_INPUT', 't-new') + await h.advance() + assert.equal(h.rpc('prompt.submit').length, 4, 'all four unknown workers still occupy capacity') + options.unavailable = false + for (const session of h.sessions.values()) { session.state = 'interrupted'; session.pending = null } + for (const member of h.roster) await h.gc.harvestStrandedGroupReply('Room', member) + await h.advance() + assert.ok(h.rpc('prompt.submit').length > 4, 'exact terminal releases capacity for the new instruction') + assert.equal(posts(h).length, 0, 'Stop reconciliation never publishes the old result') + for (const ref of accepted) assert.ok(h.rpc('session.turn.poll').some(c => + c.params.accepted_turn?.request_id === ref.request_id && c.params.accepted_turn?.host_boot_id === ref.host_boot_id)) + await settle(h) + assert.equal(coordinator(h).active, 0) + assert.ok(h.leases.every(l => l.releases === 1)) + }) +} + +for (const cold of [false, true]) { + bounded(`GC-ACK: ${cold ? 'cold' : 'hot'} lost ACK retries the captured generation and still awaits terminal`, async () => { + const options = { interruptError: true, interruptReply: { status: 'interrupted' } } + const hot = await harness(members(1), cold ? {} : options), pending = drive(hot) + await flush() + let h = hot + if (cold) { + h = await harness(hot.roster, options) + for (const [id, session] of hot.sessions) h.sessions.set(id, clone(session)) + const rooms = clone(hot.gc.durableGroupChatRooms()) + rooms.Room.running = false; rooms.Room.turns = []; rooms.Room.turn = null + h.gc.$groupChats.set(rooms) + } + const ref = clone(receipt(h).delivery.accepted_turn) + const first = await h.gc.stopGroupThread('Room', 't1', h.roster) + await h.advance() + assert.equal(first.status, 'unconfirmed') + assert.deepEqual(receipt(h)?.delivery.accepted_turn, ref) + options.interruptError = false + const second = await h.gc.stopGroupThread('Room', 't1', h.roster) + assert.equal(second.status, 'stopping') + assert.deepEqual(h.rpc('session.interrupt').map(c => c.params.session_id), [ref.session_id, ref.session_id]) + assert.deepEqual(receipt(h)?.delivery.accepted_turn, ref) + sessionFor(h).state = 'interrupted'; sessionFor(h).pending = null + await h.gc.harvestStrandedGroupReply('Room', h.roster[0]); await h.advance() + assert.equal(receipt(h), undefined) + assert.equal(posts(h).length, 0) + if (!cold) assert.equal(h.activeLeases(), 0) + await settle(hot, pending) + }) +} + +function renderer() { + const states = []; let cursor = 0 + return { react: { + useState: initial => { + const index = cursor++ + if (!(index in states)) states[index] = typeof initial === 'function' ? initial() : initial + return [states[index], value => { states[index] = typeof value === 'function' ? value(states[index]) : value }] + }, useRef: current => ({ current }), useEffect: () => {}, useMemo: factory => factory() + }, sdk: { useValue: atom => atom.get(), cn: (...values) => values.filter(Boolean).join(' '), + profileColor: () => '#000000', relativeTime: () => '' }, + render: (component, props, fresh = false) => { + cursor = 0; if (fresh) states.length = 0 + return component(props) + } } +} +function nodes(tree) { + const result = [] + const walk = node => { + if (Array.isArray(node)) node.forEach(walk) + else if (node && typeof node === 'object') { result.push(node); walk(node.props?.children) } + } + walk(tree); return result +} +function preparedAnswer(h, entry, approvalChoice = 'once') { + const props = { entry, members: h.roster } + let tree = h.ui.render(h.gc.GroupClarifyCard, props, true) + if (entry.kind === 'approval') nodes(tree).find(n => n.props?.children === approvalChoice && n.props?.onClick).props.onClick() + else nodes(tree).find(n => n.props?.['aria-label'] === `Answer @${entry.member}`).props.onChange({ target: { value: 'ANSWER' } }) + tree = h.ui.render(h.gc.GroupClarifyCard, props) + const button = nodes(tree).find(n => ['Answer', 'Respond'].includes(n.props?.children) && n.props?.onClick) + assert.equal(button.props.disabled, false) + return () => button.props.onClick() +} +async function park(kind = 'approval', options = {}) { + const ui = renderer(), notices = [], originalProjection = options.resumeProjection + Object.assign(options, ui, { onNotify: notice => notices.push(notice), + resumeProjection: (session, method, projection) => { + if (session.approval && session.state === 'waiting') projection.pending_approval = session.approval + return originalProjection?.(session, method, projection) || projection + } }) + const h = await harness(members(1), options), pending = drive(h) + h.ui = ui; h.notices = notices; await flush() + const session = sessionFor(h) + session.state = 'waiting' + if (kind === 'approval') session.approval = { request_id: 'approval-1', command: 'echo offline', choices: ['once', 'deny'] } + else session.pending = { request_id: 'clarify-1', question: 'Choose?' } + await h.advance(); await pending + assert.equal(cards(h).length, 1) + return { h, session, options } +} + +for (const approvalResult of [{ resolved: 0 }, {}, { resolved: '1' }, { resolved: -1 }]) { + bounded(`GC-APPROVAL: ${JSON.stringify(approvalResult)} produces no successful echo`, async () => { + const { h, session } = await park('approval', { approvalResult }) + const accepted = clone(session.ref), log = clone(h.room().log), entry = cards(h)[0] + await preparedAnswer(h, entry)(); await flush() + assert.deepEqual(h.room().log, log) + assert.equal(cards(h)[0], entry, 'a negative or unknown acknowledgement keeps the card') + assert.equal(h.notices.filter(n => n.kind === 'error').length, 1) + assert.deepEqual(receipt(h)?.delivery.accepted_turn, accepted) + assert.equal(h.rpc('prompt.submit').length, 1) + await settle(h) + }) +} +for (const choice of ['once', 'deny']) { + bounded(`GC-APPROVAL: resolved=1 confirms ${choice} and delivers only the exact accepted turn`, async () => { + const { h, session } = await park('approval', { approvalResult: { resolved: 1 } }) + const accepted = clone(session.ref) + session.text = `AFTER_${choice}` + await preparedAnswer(h, cards(h)[0], choice)(); await h.advance() + assert.equal(cards(h).length, 0) + assert.equal(h.room().log.filter(e => e.from.kind === 'user' && e.answer_to).length, 1) + assert.equal(posts(h).filter(e => e.text === `AFTER_${choice}`).length, 1) + assert.deepEqual(session.ref, accepted) + assert.equal(h.rpc('prompt.submit').length, 1) + await settle(h) + }) +} +bounded('GC-APPROVAL: lost response ACK retains card, receipt and honest error', async () => { + const { h } = await park('approval', { answerError: () => true }) + const entry = cards(h)[0], log = clone(h.room().log) + await preparedAnswer(h, entry)(); await flush() + assert.deepEqual(h.room().log, log) + assert.equal(cards(h)[0], entry) + assert.equal(h.notices.filter(n => n.kind === 'error').length, 1) + assert.ok(receipt(h)) + await settle(h) +}) + +for (const legacy of [false, true]) { + bounded(`GC-RENAME: ${legacy ? 'legacy token' : 'room ID'} retains running owner through repeated local renames`, async t => { + const h = await harness(members(1)) + if (!h.gc.renameGroupChat) return t.skip('Exact baseline lacks the rename test port; final source must run this gate') + if (legacy) delete h.room().roomId + const pending = drive(h); await flush() + const owner = coordinator(h), session = sessionFor(h), accepted = clone(session.ref), title = session.title + await h.gc.renameGroupChat('Room', 'Renamed', h.roster) + await h.gc.renameGroupChat('Renamed', 'RenamedAgain', h.roster) + assert.equal(coordinator(h, 'RenamedAgain'), owner) + assert.equal(coordinator(h), undefined) + assert.equal([...owner.occurrences][0].group, 'RenamedAgain') + assert.deepEqual(receipt(h, h.roster[0], 'RenamedAgain')?.delivery.accepted_turn, accepted) + session.state = 'complete'; session.text = 'FINAL_AFTER_RENAME' + await h.advance(); await pending + assert.equal(posts(h, 'RenamedAgain').filter(e => e.text === 'FINAL_AFTER_RENAME').length, 1) + assert.equal(h.gc.$groupChats.get().Room, undefined) + assert.equal(h.gc.$groupChats.get().Renamed, undefined) + assert.equal(h.gc.$groupChats.get().RenamedAgain.running, false) + assert.equal(session.title, title) + assert.equal(h.rpc('session.create').length, 1) + assert.equal(h.rpc('prompt.submit').length, 1) + await settle(h, null, 'RenamedAgain') + }) +} +bounded('GC-RENAME: pending card response and echo follow the same lifetime across rename', async t => { + const answerAckGate = deferred(), { h, session } = await park('clarify', { answerAckGate }) + if (!h.gc.renameGroupChat) return t.skip('Exact baseline lacks the rename test port; final source must run this gate') + const entry = cards(h)[0], owner = coordinator(h), accepted = clone(session.ref) + session.text = 'FINAL_AFTER_RENAMED_ANSWER' + const answering = preparedAnswer(h, entry)(); await flush() + await h.gc.renameGroupChat('Room', 'AnsweredRoom', h.roster) + assert.equal(cards(h)[0], entry) + assert.equal(entry.group, 'AnsweredRoom') + answerAckGate.resolve(); await answering; await h.advance() + const room = h.gc.$groupChats.get().AnsweredRoom + assert.equal(h.gc.$groupChats.get().Room, undefined) + assert.equal(room.log.filter(e => e.answer_to).length, 1) + assert.equal(posts(h, 'AnsweredRoom').filter(e => e.text === 'FINAL_AFTER_RENAMED_ANSWER').length, 1) + assert.equal(coordinator(h, 'AnsweredRoom'), owner) + assert.deepEqual(session.ref, accepted) + assert.equal(h.rpc('prompt.submit').length, 1) + await settle(h, null, 'AnsweredRoom') +}) +bounded('GC-RENAME: in-flight Stop follows rename and retains custody until exact terminal', async t => { + const interruptGate = deferred(), h = await harness(members(1), { interruptGate, interruptReply: { status: 'interrupted' } }) + if (!h.gc.renameGroupChat) return t.skip('Exact baseline lacks the rename test port; final source must run this gate') + const pending = drive(h); await flush() + const accepted = clone(sessionFor(h).ref) + const stopping = h.gc.stopGroupThread('Room', 't1', h.roster); await flush() + await h.gc.renameGroupChat('Room', 'StoppedRoom', h.roster) + interruptGate.resolve(); const result = await stopping; await h.advance(); await pending + assert.equal(result.status, 'stopping') + assert.deepEqual(receipt(h, h.roster[0], 'StoppedRoom')?.delivery.accepted_turn, accepted) + assert.equal(h.activeLeases(), 1) + sessionFor(h).state = 'interrupted' + await h.gc.harvestStrandedGroupReply('StoppedRoom', h.roster[0]); await flush() + assert.equal(receipt(h, h.roster[0], 'StoppedRoom'), undefined) + assert.equal(h.activeLeases(), 0) + assert.equal(h.gc.$groupChats.get().Room, undefined) + assert.equal(posts(h, 'StoppedRoom').length, 0) +}) +bounded('GC-RENAME: remote receive rehomes running owner before room listeners see the new name', async t => { + const options = { rpcResponse: (_route, method) => method === 'profiles.list' + ? { profiles: [{ name: 'default', ui_meta: { 'hermes-bots-groups': options.remote } }] } + : method === 'profiles.configure' ? { applied: { ui_meta: true } } : undefined } + const h = await harness(members(1), options) + if (!h.gc.pullGroupChatServerState || !h.gc.renameGroupChat) return t.skip('Exact baseline lacks the rename test port; final source must run this gate') + const pending = drive(h); await flush() + const owner = coordinator(h), session = sessionFor(h), accepted = clone(session.ref) + const snapshot = h.gc.groupChatSyncSnapshot(), [roomKey, sourceRoom] = Object.entries(snapshot.rooms)[0] + const projected = clone(sourceRoom) + assert.ok(projected) + options.remote = { ...snapshot, rooms: { [roomKey]: { ...projected, name: 'RemoteRoom', revision: projected.revision + 10 } } } + // This receive starts after local writes have settled. An intentionally + // pending local edit has separate preserveRooms authority over the name. + h.gc.stopGroupChatServerSync() + await h.gc.pullGroupChatServerState('local') + assert.equal(coordinator(h, 'RemoteRoom'), owner) + assert.equal(h.gc.$groupChats.get().Room, undefined) + assert.deepEqual(receipt(h, h.roster[0], 'RemoteRoom')?.delivery.accepted_turn, accepted) + session.state = 'complete'; session.text = 'REMOTE_RENAMED_FINAL' + await h.advance(); await pending + assert.equal(posts(h, 'RemoteRoom').filter(e => e.text === 'REMOTE_RENAMED_FINAL').length, 1) + assert.equal(h.gc.$groupChats.get().RemoteRoom.running, false) + assert.equal(h.rpc('prompt.submit').length, 1) + await settle(h, null, 'RemoteRoom') +}) +bounded('GC-RENAME: same display name with a new room lifetime cannot adopt the old final', async () => { + const h = await harness(members(1)), pending = drive(h) + await flush() + const session = sessionFor(h) + h.gc.$groupChats.set({ Room: { ...clone(h.room()), roomId: 'replacement-room', epoch: 1, running: false, + sessions: {}, sessionOwners: {}, stranded: {}, log: [], turns: [], turn: null } }) + session.state = 'complete'; session.text = 'OLD_LIFETIME_FINAL' + await h.advance(); await pending + assert.equal(posts(h).length, 0) + assert.equal(h.room().roomId, 'replacement-room') + assert.equal(h.rpc('prompt.submit').length, 1) +}) +bounded('GC-RENAME: late Stop ACK cannot poll or update a replacement room at the same name', async () => { + const interruptGate = deferred(), h = await harness(members(1), { interruptGate, interruptReply: { status: 'interrupted' } }) + const pending = drive(h); await flush() + const stopping = h.gc.stopGroupThread('Room', 't1', h.roster); await flush() + const before = h.rpc('session.turn.poll').length + const replacement = { ...clone(h.room()), roomId: 'new-lifetime', running: true, epoch: 20, + stranded: {}, sessions: {}, sessionOwners: {}, holds: {}, log: [], turns: [], turn: null } + h.gc.$groupChats.set({ Room: replacement }) + interruptGate.resolve(); const result = await stopping; await h.advance(); await pending + assert.equal(result.status, 'unconfirmed') + assert.equal(h.rpc('session.turn.poll').length, before) + assert.equal(h.room().roomId, 'new-lifetime') + assert.equal(h.room().epoch, 20) + assert.equal(h.room().running, true) + assert.equal(h.room().log.length, 0) +}) +bounded('GC-HOLD: waiting hold retains the exact question custody while suppressing stale Answer', async () => { + const { h, session } = await park('clarify') + const accepted = clone(session.ref), staleAnswer = preparedAnswer(h, cards(h)[0]) + assert.equal(coordinator(h).active, 0) + h.gc.sendToGroupChat('Room', h.roster, '@bot1 pause', 'hold-thread') + await h.advance(); await staleAnswer(); await flush() + assert.equal(h.rpc('clarify.respond').length, 0) + assert.deepEqual(receipt(h)?.delivery.accepted_turn, accepted) + assert.equal(h.activeLeases(), 1) + assert.equal(coordinator(h).active, 0, 'verified waiting keeps no execution worker') + assert.equal(h.rpc('session.interrupt').length, 0) + session.state = 'complete'; session.pending = null; session.text = 'HELD_QUESTION_FINAL' + await h.gc.harvestStrandedGroupReply('Room', h.roster[0]); await h.advance() + assert.equal(receipt(h), undefined) + assert.equal(h.activeLeases(), 0) + assert.equal(posts(h).length, 0) +}) +bounded('GC-ACK: unavailable timeout and repeated Stop retain four occupied workers', async () => { + const options = { unavailable: true, interruptReply: { status: 'interrupted' } } + const h = await harness(members(6), options), pending = drive(h) + await flush(); await h.advance(20 * 60000) + assert.equal(h.rpc('prompt.submit').length, 4) + assert.equal(coordinator(h).active, 4) + assert.equal(h.activeLeases(), 4) + const first = await h.gc.stopGroupThread('Room', 't1', h.roster); await h.advance(); await pending + assert.equal(first.status, 'stopping') + const second = await h.gc.stopGroupThread('Room', 't1', h.roster) + assert.equal(second.status, 'stopping') + assert.equal(Object.keys(h.room().stranded).length, 4) + assert.equal(coordinator(h).active, 4) + assert.equal(h.rpc('prompt.submit').length, 4) + options.unavailable = false + for (const session of h.sessions.values()) session.state = 'interrupted' + for (const member of h.roster) await h.gc.harvestStrandedGroupReply('Room', member) + assert.equal(h.activeLeases(), 0) + assert.equal(coordinator(h).active, 0) + assert.equal(posts(h).length, 0) +}) + +for (const cold of [false, true]) { + for (const failed of [false, true]) { + bounded(`GC-ACK-02: ${cold ? 'cold' : 'parked hot'} terminal retains custody through ${failed ? 'rejected' : 'applied'} pending interrupt`, async () => { + const interruptGate = deferred(), options = { unavailable: true, interruptGate, interruptError: failed } + const hot = await harness(members(1), cold ? { unavailable: true } : options), initial = drive(hot) + await flush(); await hot.advance(); await initial + const oldOwner = [...coordinator(hot).occurrences][0] + assert.equal(oldOwner.collectorDone, true, 'the collector has parked on unavailable outcome') + let h = hot + if (cold) { + h = await harness(hot.roster, options) + h.gc.$groupChats.set(clone(hot.gc.durableGroupChatRooms())) + for (const [id, session] of hot.sessions) h.sessions.set(id, clone(session)) + } + const marker = receipt(h), accepted = clone(marker.delivery.accepted_turn) + const stopping = h.gc.stopGroupThread('Room', 't1', h.roster) + const repeated = h.gc.stopGroupThread('Room', 't1', h.roster) + await flush() + options.unavailable = false + const session = sessionFor(h) + session.state = 'complete'; session.pending = null; session.text = 'OLD_STOPPED_FINAL' + h.gc.sendToGroupChat('Room', h.roster, '@all resume REPLACEMENT_INPUT', 'replacement-thread') + await h.advance(250) + await h.gc.harvestStrandedGroupReply('Room', h.roster[0]); await flush() + assert.equal(receipt(h), marker, 'exact terminal cannot erase a pending session control target') + assert.equal(h.rpc('prompt.submit').length, cold ? 0 : 1, 'no successor may reuse the runtime before control settles') + assert.equal(h.rpc('session.interrupt').length, 1, 'concurrent Stop shares the exact pending request') + assert.deepEqual(session.ref, accepted) + if (!cold) { + assert.equal(oldOwner.released, undefined) + assert.equal(coordinator(h).members.get(oldOwner.memberLock), oldOwner) + assert.equal(h.gc.groupRuntimeSessionOwners.get(oldOwner.sessionLock), oldOwner) + assert.equal(h.leases[0].releases, 0) + assert.equal(coordinator(h).active, 1) + assert.equal(oldOwner.reservation, true, 'pending retirement keeps its publication reservation') + } + interruptGate.resolve() + await stopping; await repeated; await flush() + await h.gc.harvestStrandedGroupReply('Room', h.roster[0]); await h.advance() + // A cold blocked instruction can require a fresh explicit drive. That is + // intentional; it cannot replay the uncertain predecessor automatically. + if (h.rpc('prompt.submit').length === (cold ? 0 : 1)) { + void h.gc.runGroupChatRounds('Room', h.roster, 'replacement-thread') + await flush(); await h.advance() + } + assert.equal(h.rpc('prompt.submit').length, cold ? 1 : 2) + assert.ok(h.rpc('prompt.submit').at(-1).params.text.includes('REPLACEMENT_INPUT')) + assert.equal(h.rpc('session.interrupt').length, 1, 'no delayed cleanup targets the admitted replacement') + assert.equal(posts(h).filter(e => e.text === 'OLD_STOPPED_FINAL').length, 0) + if (!cold) assert.equal(h.leases[0].releases, 1) + await settle(h) + await settle(hot) + }) + } +} + +async function delayedRenameHarness(extra = {}) { + const metadataGate = deferred(), ui = renderer(), options = { ...ui, ...extra, + rpcResponse: async (route, method, params) => { + const supplied = await extra.rpcResponse?.(route, method, params) + if (supplied !== undefined) return supplied + if (method === 'profiles.configure') { + if (params.name === 'bot1' && params.ui_meta?.['hermes-bots']?.groups?.includes('Renamed')) { + await metadataGate.promise + } + return { applied: { ui_meta: true } } + } + return undefined + } } + const h = await harness(members(2), options) + h.ui = ui + // Seed memberships through the actual creation handler, then use the same + // room/metadata owners as the real rename and settings actions. + h.gc.$groupChats.set({}) + assert.equal(h.gc.createFreshGroupChat('Room', h.roster), 'Room') + await flush() + h.gc.updateGroupChat('Room', room => { room.log = [clone(h.input)]; return room }, { sync: false }) + h.gc.$groupChatWorkspace.set('Room') + return { h, metadataGate, options } +} + +bounded('GC-RENAME-02: overlapping local renames cannot resurrect a name or stale later-member membership', async () => { + const { h, metadataGate } = await delayedRenameHarness() + const roomId = h.room().roomId + const first = h.gc.renameGroupChat('Room', 'Renamed', h.roster); await flush() + const second = await h.gc.renameGroupChat('Renamed', 'RenamedAgain', h.roster) + metadataGate.resolve() + assert.equal(await first, 'RenamedAgain', 'the late return names the same current lifetime') + assert.equal(second, 'RenamedAgain') + assert.equal(h.gc.$groupChats.get().Room, undefined) + assert.equal(h.gc.$groupChats.get().Renamed, undefined) + assert.equal(h.gc.$groupChats.get().RenamedAgain.roomId, roomId) + assert.equal(h.gc.$groupChatWorkspace.get(), 'RenamedAgain') + for (const member of h.roster) { + const last = h.rpc('profiles.configure').filter(c => c.params.name === member.name).at(-1) + assert.deepEqual(last.params.ui_meta['hermes-bots'].groups, ['RenamedAgain']) + } + assert.equal(h.rpc('session.create').length, 0) + assert.equal(h.rpc('prompt.submit').length, 0) +}) + +function prepareSettings(h, renamed = []) { + let closed = 0 + const props = { group: 'Room', members: h.roster, open: true, onClose: () => { closed++ }, onRenamed: name => renamed.push(name) } + let tree = h.ui.render(h.gc.GroupChatSettingsDialog, props, true) + nodes(tree).find(n => n.props?.['aria-label'] === 'Group name').props.onChange({ target: { value: 'Renamed' } }) + nodes(tree).find(n => n.props?.onImage).props.onImage('OFFLINE_ROOM_IMAGE') + tree = h.ui.render(h.gc.GroupChatSettingsDialog, props) + const click = nodes(tree).find(n => n.props?.children === 'Save' && n.props?.onClick).props.onClick + return { click, renamed, closed: () => closed } +} +function submitSettings(h, renamed = []) { + const settings = prepareSettings(h, renamed) + settings.click() + return settings +} + +bounded('GC-RENAME-02: the actual settings image and callback follow an overlapping rename', async () => { + const { h, metadataGate } = await delayedRenameHarness(), settings = submitSettings(h) + await flush() + await h.gc.renameGroupChat('Renamed', 'RenamedAgain', h.roster) + metadataGate.resolve(); await flush() + assert.equal(h.gc.$groupChats.get().Renamed, undefined) + assert.equal(h.gc.$groupChats.get().RenamedAgain.image, 'OFFLINE_ROOM_IMAGE') + assert.deepEqual(settings.renamed, ['RenamedAgain']) + assert.equal(settings.closed(), 1) +}) + +bounded('GC-RENAME-02: remote rename during delayed member persistence keeps the current lifetime only', async () => { + let remote + const { h, metadataGate } = await delayedRenameHarness({ rpcResponse: (_route, method) => method === 'profiles.list' + ? { profiles: [{ name: 'default', ui_meta: { 'hermes-bots-groups': remote } }] } : undefined }) + const first = h.gc.renameGroupChat('Room', 'Renamed', h.roster); await flush() + const snapshot = h.gc.groupChatSyncSnapshot(), [key, room] = Object.entries(snapshot.rooms)[0] + remote = { ...snapshot, rooms: { [key]: { ...clone(room), name: 'RemoteRenamed', revision: room.revision + 10 } } } + // Clear only the fake mirror's pending edit, matching a settled/received + // remote revision. The real receive handler retains its existing CAS rule. + h.gc.stopGroupChatServerSync() + await h.gc.pullGroupChatServerState('local') + metadataGate.resolve() + assert.equal(await first, 'RemoteRenamed') + assert.equal(h.gc.$groupChats.get().Room, undefined) + assert.equal(h.gc.$groupChats.get().Renamed, undefined) + assert.equal(h.gc.$groupChats.get().RemoteRenamed.roomId, room.roomId) + assert.equal(h.rpc('prompt.submit').length, 0) +}) + +for (const replaced of [false, true]) { + bounded(`GC-RENAME-02: delayed actual settings cannot update a ${replaced ? 'replacement' : 'deleted'} room`, async () => { + const { h, metadataGate } = await delayedRenameHarness(), settings = submitSettings(h) + await flush() + const replacement = { ...clone(h.gc.$groupChats.get().Renamed), roomId: 'replacement-lifetime', + log: [], sessions: {}, sessionOwners: {}, stranded: {}, image: 'REPLACEMENT_IMAGE' } + h.gc.$groupChats.set(replaced ? { Renamed: replacement } : {}) + metadataGate.resolve(); await flush() + assert.equal(h.gc.$groupChats.get().Room, undefined) + assert.deepEqual(settings.renamed, []) + assert.equal(settings.closed(), 0) + if (replaced) { + assert.equal(h.gc.$groupChats.get().Renamed, replacement) + assert.equal(replacement.image, 'REPLACEMENT_IMAGE') + } else assert.equal(h.gc.$groupChats.get().Renamed, undefined) + }) +} + +for (const replaced of [false, true]) { + bounded(`GC-RENAME-02: an already rendered Save cannot adopt a ${replaced ? 'replacement' : 'deleted'} lifetime`, async () => { + const { h, metadataGate } = await delayedRenameHarness(), settings = prepareSettings(h) + const replacement = { ...clone(h.room()), roomId: 'replacement-before-click', image: 'REPLACEMENT_IMAGE', log: [] } + const before = h.rpc('profiles.configure').length + h.gc.$groupChats.set(replaced ? { Room: replacement } : {}) + settings.click(); metadataGate.resolve(); await flush() + assert.deepEqual(h.gc.$groupChats.get(), replaced ? { Room: replacement } : {}) + assert.equal(h.rpc('profiles.configure').length, before, 'a stale button cannot rename members of a vanished room') + assert.deepEqual(settings.renamed, []) + assert.equal(settings.closed(), 0) + }) +} + +bounded('GC-RENAME-02: actual settings still create and rename a legitimate metadata-only group', async () => { + const { h, metadataGate } = await delayedRenameHarness() + h.gc.$groupChats.set({}) + const settings = submitSettings(h); await flush() + metadataGate.resolve(); await flush() + const room = h.gc.$groupChats.get().Renamed + assert.ok(room) + assert.ok(room.coordinationId, 'the existing legacy lifetime token fences the save') + assert.equal(room.image, 'OFFLINE_ROOM_IMAGE') + assert.equal(h.gc.$groupChats.get().Room, undefined) + assert.deepEqual(settings.renamed, ['Renamed']) + assert.equal(settings.closed(), 1) + assert.equal(h.rpc('prompt.submit').length, 0) +}) diff --git a/apps/desktop/src/plugins/hermes-bots/tests/group-parallel.test.mjs b/apps/desktop/src/plugins/hermes-bots/tests/group-parallel.test.mjs index 35bfe3a108c2b..e2871f592f5d5 100644 --- a/apps/desktop/src/plugins/hermes-bots/tests/group-parallel.test.mjs +++ b/apps/desktop/src/plugins/hermes-bots/tests/group-parallel.test.mjs @@ -191,6 +191,11 @@ for (const expired of [false, true]) { assert.equal(h.rpc('prompt.submit').length, 5, 'sixth stays queued while four may be running') assert.equal(h.leases.find(l => l.route.profile === 'bot1').releases, 0) await h.gc.stopGroupThread('Room', 't1', h.roster); await h.advance(); await pending; await flush() + if (options.unavailable) { + assert.ok(h.leases.some(l => l.releases === 0), 'an interrupt ACK does not retire an unavailable accepted turn') + options.unavailable = false + for (const member of h.roster) await h.gc.harvestStrandedGroupReply('Room', member) + } assert.ok(h.leases.every(l => l.releases === 1)) }) } diff --git a/apps/desktop/src/plugins/hermes-bots/tests/group-protocol-custody.test.mjs b/apps/desktop/src/plugins/hermes-bots/tests/group-protocol-custody.test.mjs index 9c27b8f5cc492..c90b0cd722719 100644 --- a/apps/desktop/src/plugins/hermes-bots/tests/group-protocol-custody.test.mjs +++ b/apps/desktop/src/plugins/hermes-bots/tests/group-protocol-custody.test.mjs @@ -24,11 +24,15 @@ test('capability resume recreation transfers custody to the admitted runtime', a await flush() assert.equal(submittedRuntime, `${oldRuntime}-recreated`) assert.equal(h.rpc('prompt.submit').length, 1) - assert.equal((await h.gc.stopGroupThread('Room', 't1', h.roster)).status, 'stopped') + const stopping = await h.gc.stopGroupThread('Room', 't1', h.roster) + assert.equal(stopping.status, 'stopping', 'interrupt ACK precedes exact terminal retirement') + assert.equal(stopping.pending, 1) await h.advance() await pending assert.equal(h.rpc('session.interrupt').at(-1).params.session_id, submittedRuntime) assert.equal(h.gc.groupRuntimeSessionOwners.size, 0) + assert.equal(h.activeLeases(), 0) + assert.equal((await h.gc.stopGroupThread('Room', 't1', h.roster)).status, 'stopped') }) test('Stop during capability resume interrupts the recreated runtime without admission', async () => { @@ -126,10 +130,12 @@ test('retry remint collision preserves destination custody and never interrupts assert.equal(h.gc.groupRuntimeSessionOwners.get(destinationKey), occupant) assert.equal(h.rpc('session.interrupt').filter(c => c.params.session_id === occupiedRuntime).length, 0) h.gc.groupRuntimeSessionOwners.delete(destinationKey) - const retained = [...h.gc.groupRuntimeSessionOwners.values()] - assert.equal(retained.length, 1, 'failed original admission retains its unresolved source custody') - assert.equal(retained[0], originalOccurrence) - assert.equal(retained[0].runtime, originalRuntime) - assert.equal(retained[0].sessionLock, originalLock) - assert.equal(h.gc.groupRuntimeSessionOwners.get(originalLock), originalOccurrence) + assert.equal(h.gc.groupRuntimeSessionOwners.size, 0, 'the proven refusal releases only its unused source custody') + assert.equal(h.gc.groupRuntimeSessionOwners.get(originalLock), undefined) + assert.equal(originalOccurrence.released, true) + assert.equal(originalOccurrence.runtime, originalRuntime) + assert.equal(originalOccurrence.sessionLock, originalLock) + assert.equal(h.activeLeases(), 0) + assert.equal(h.leases[0].releases, 1) + assert.equal(h.gc.groupRoomCoordinators.get('Room').active, 0) }) diff --git a/apps/desktop/src/plugins/hermes-bots/tests/group-stop-thread.test.mjs b/apps/desktop/src/plugins/hermes-bots/tests/group-stop-thread.test.mjs index cfc16a188af51..8aade666d9b67 100644 --- a/apps/desktop/src/plugins/hermes-bots/tests/group-stop-thread.test.mjs +++ b/apps/desktop/src/plugins/hermes-bots/tests/group-stop-thread.test.mjs @@ -224,21 +224,25 @@ test('stopGroupThread with nobody on turn stops the room without any interrupt R assert.equal(gc.$groupChats.get().Room.epoch, 4) }) -test('stopGroupThread records a stopped activity event visible in the CURRENT run', async () => { +test('stopGroupThread records an honest legacy stop-unconfirmed event in the CURRENT run', async () => { const gc = load() seedRoom(gc) await gc.stopGroupThread('Room', 't1', MEMBERS) const events = gc.currentGroupActivity('Room') - const stopped = events.find(event => event.kind === 'stopped') + const stopped = events.find(event => event.kind === 'stop-unconfirmed') assert.ok(stopped, 'stopped event is tagged with the POST-bump epoch, so it survives the epoch filter') assert.equal(stopped.member, 'You') assert.equal(stopped.thread, 't1') // The label comes from the shared GROUP_ACTIVITY_LABELS map (the plugin's // label pattern) — and stays plain English, no hardcoded localized text. - assert.ok(gc.GROUP_ACTIVITY_LABELS.stopped) - assert.ok(gc.GROUP_ACTIVITY_GLYPHS.stopped) + assert.ok(gc.GROUP_ACTIVITY_LABELS['stop-unconfirmed']) + assert.ok(gc.GROUP_ACTIVITY_GLYPHS['stop-unconfirmed']) + const idle = load() + seedRoom(idle, { turn: null }) + assert.equal((await idle.stopGroupThread('Room', 't1', MEMBERS)).status, 'stopped') + assert.ok(idle.currentGroupActivity('Room').some(event => event.kind === 'stopped' && event.epoch === 4)) }) test('stopGroupThread falls back to the durable room roster when called without members', async () => { diff --git a/apps/desktop/src/plugins/hermes-bots/tests/group-turn-lease.test.mjs b/apps/desktop/src/plugins/hermes-bots/tests/group-turn-lease.test.mjs index 6dccbe76a2058..60e754d2a261d 100644 --- a/apps/desktop/src/plugins/hermes-bots/tests/group-turn-lease.test.mjs +++ b/apps/desktop/src/plugins/hermes-bots/tests/group-turn-lease.test.mjs @@ -267,15 +267,21 @@ test('the per-turn lease is released after the turn — refcount returns to zero assert.equal(gc.stats().disposals, 1) }) -test('uncertain submit failure retains the per-turn lease until captured Stop', async () => { +test('uncertain submit failure retains the per-turn lease after a captured Stop ACK', async () => { const fatal = new Error('backend exploded') const gc = load({ failEverySubmitWith: fatal }) await assert.rejects(() => gc.runGroupChatMemberTurn('Room', ROUTED_MEMBER, 'hi', 't1', [])) assert.equal(gc.stats().refcount, 1, 'unknown acceptance is not vacant capacity') - await gc.stopGroupThread('Room', 't1', [ROUTED_MEMBER]) - assert.equal(gc.stats().refcount, 0) + const first = await gc.stopGroupThread('Room', 't1', [ROUTED_MEMBER]) + assert.equal(first.status, 'stopping') + assert.equal(first.pending, 1) + assert.equal(gc.stats().refcount, 1, 'ACK cannot prove retirement without admission identity') + const second = await gc.stopGroupThread('Room', 't1', [ROUTED_MEMBER]) + assert.equal(second.status, 'stopping') + assert.equal(gc.stats().refcount, 1) + assert.equal(gc.stats().submits, 1, 'unknown acceptance never replays user text') }) test('hosts without retainProfile still run the turn (feature detection)', async () => { diff --git a/apps/desktop/src/plugins/hermes-bots/tests/group-turn-outcomes.test.mjs b/apps/desktop/src/plugins/hermes-bots/tests/group-turn-outcomes.test.mjs index 66691229b61f3..05cde7594a7b9 100644 --- a/apps/desktop/src/plugins/hermes-bots/tests/group-turn-outcomes.test.mjs +++ b/apps/desktop/src/plugins/hermes-bots/tests/group-turn-outcomes.test.mjs @@ -436,10 +436,18 @@ test('explicit stop exits the collector and never collects its late complete out let stopped = false const h = await harness({ alpha: [async (s, gc) => { if (!stopped) { stopped = true; await gc.stopGroupThread('Room', 'thread-1', [ALPHA, BETA]) } - return { turn_outcomes: wire(s.ref, 'running', [final('late candidate')]) } - }] }) + return { turn_outcomes: wire(s.ref, s.terminalAfterStop ? 'complete' : 'running', [final('late candidate')]) } + }] }, { connectionId: 'pc' }) assert.equal(await run(h), null) + const marker = room(h).stranded.alpha + assert.equal(marker.stop_requested, true) + assert.deepEqual(marker.delivery.accepted_turn, h.sessions.get('alpha').ref) + assert.equal(h.releases(), 0, 'running custody outlives the interrupt ACK') + h.sessions.get('alpha').terminalAfterStop = true + await h.gc.harvestStrandedGroupReply('Room', ALPHA) + await new Promise(resolve => setImmediate(resolve)) assert.equal(room(h).stranded.alpha, undefined) + assert.equal(h.releases(), 1, 'the exact terminal read releases custody after control settles') assert.equal(posts(h).length, 0) assert.equal(h.rpc('prompt.submit').length, 1) }) @@ -716,12 +724,22 @@ test('rejected read-only poll is visible immediately and next member advances on }) test('stopped turn discards an in-flight poll rejection without stale failure publication', async () => { - const h = await harness({ alpha: [async (_s, gc) => { - await gc.stopGroupThread('Room', 'thread-1', [ALPHA]) + let stopped = false + const h = await harness({ alpha: [async (s, gc) => { + if (!stopped) { stopped = true; await gc.stopGroupThread('Room', 'thread-1', [ALPHA]) } + if (s.terminalAfterStop) return { turn_outcomes: wire(s.ref, 'complete', [final('late rejected candidate')]) } throw new Error('late lost observer') - }] }, { members: [ALPHA] }) + }] }, { members: [ALPHA], connectionId: 'pc' }) assert.equal(await run(h), null) + assert.equal(room(h).stranded.alpha.stop_requested, true) + assert.deepEqual(room(h).stranded.alpha.delivery.accepted_turn, h.sessions.get('alpha').ref) + assert.equal(h.rpc('prompt.submit').length, 1) + assert.equal(h.releases(), 0, 'a rejected observation is not terminal evidence') + h.sessions.get('alpha').terminalAfterStop = true + await h.gc.harvestStrandedGroupReply('Room', ALPHA) + await new Promise(resolve => setImmediate(resolve)) assert.equal(room(h).stranded.alpha, undefined) + assert.equal(h.releases(), 1) assert.equal(posts(h).length, 0) assert.equal(h.gc.currentGroupActivity('Room').filter(e => e.kind === 'unavailable').length, 0) }) diff --git a/apps/desktop/src/plugins/hermes-bots/tests/pr118-postmerge.test.mjs b/apps/desktop/src/plugins/hermes-bots/tests/pr118-postmerge.test.mjs index f7ec4c18992d0..4950349ab0493 100644 --- a/apps/desktop/src/plugins/hermes-bots/tests/pr118-postmerge.test.mjs +++ b/apps/desktop/src/plugins/hermes-bots/tests/pr118-postmerge.test.mjs @@ -185,6 +185,7 @@ for (const failsFirst of [false, true]) { assert.ok(receipts(cold)[0].stop_requested) options.interruptError = false await cold.gc.stopGroupThread('Room', 't1', cold.roster) + await flush() // control completion precedes exact terminal observation assert.deepEqual(cold.rpc('session.interrupt').map(call => call.params.session_id), [accepted.session_id, accepted.session_id]) assert.equal(receipts(cold).length, 0) } @@ -267,6 +268,52 @@ boundedTest('actual UI exposes cold Stop, retains it after failed interrupt and await hot.gc.stopGroupThread('Room', 't1', hot.roster) }) +for (const cold of [false, true]) { + boundedTest(`actual Stop button projects pending retirement honestly after applied ACK; cold=${cold}`, async () => { + const notices = [], options = { interruptReply: { status: 'interrupted' }, onNotify: value => notices.push(value) } + const hot = await uiHarness(members(1), cold ? {} : options), pending = drive(hot) + await flush() + const h = cold ? await reload(hot, options) : hot + const accepted = clone(receipts(h)[0].delivery.accepted_turn), button = h.stopButton() + assert.ok(button) + button.props.onClick(); await flush() + assert.equal(notices.length, 1) + assert.equal(notices[0].kind, 'info') + assert.match(notices[0].message, /^Stopping Room/) + assert.match(notices[0].message, /waiting for the remaining turns to finish/) + assert.doesNotMatch(notices[0].message, /unconfirmed|Stop can retry/) + assert.deepEqual(receipts(h)[0].delivery.accepted_turn, accepted) + assert.equal(h.rpc('session.interrupt').length, 1) + assert.equal(h.rpc('prompt.submit').length, cold ? 0 : 1) + if (!cold) { + assert.equal(h.activeLeases(), 1) + assert.equal(h.gc.groupRoomCoordinators.get('Room').active, 1) + } + const session = [...h.sessions.values()][0] + session.state = 'interrupted'; session.pending = null + await h.gc.harvestStrandedGroupReply('Room', h.roster[0]); await h.advance() + assert.equal(receipts(h).length, 0) + if (cold) await hot.gc.stopGroupThread('Room', 't1', hot.roster) + await hot.advance(); await pending + }) +} + +boundedTest('actual Stop button keeps failed interruption wording and retry custody', async () => { + const notices = [], h = await uiHarness(members(1), { interruptError: true, onNotify: value => notices.push(value) }) + const pending = drive(h); await flush() + const accepted = clone(receipts(h)[0].delivery.accepted_turn) + h.stopButton().props.onClick(); await flush(); await h.advance(); await pending + assert.equal(notices.length, 1) + assert.equal(notices[0].kind, 'info') + assert.match(notices[0].message, /1 interruption\(s\) are unconfirmed.*Stop can retry/) + assert.deepEqual(receipts(h)[0].delivery.accepted_turn, accepted) + assert.equal(h.activeLeases(), 1) + assert.ok(h.stopButton()) + h.finish('bot1', '', 'interrupted') + await h.gc.harvestStrandedGroupReply('Room', h.roster[0]) + assert.equal(h.activeLeases(), 0) +}) + boundedTest('legacy unknown receipt does not manufacture an interrupt target or Stop affordance', async () => { const h = await uiHarness() h.gc.updateGroupChat('Room', room => ({ ...room, running: false, turns: [], @@ -379,6 +426,7 @@ boundedTest(`unresolved waiting reservations retain capacity; ${invalidation} ca const secondAdmitted = new Set() const h = await uiHarness(members(6), { resumeProjection: (session, method, projection) => { if (method !== 'session.turn.poll' || !session.ref) return projection + if (session.state === 'interrupted') return projection // preserve exact terminal after fake Stop if (session.submits === 2) secondAdmitted.add(session.profile) session.state = session.submits === 1 ? 'complete' : 'waiting' session.text = `FIRST_${session.profile} @all` diff --git a/apps/desktop/src/plugins/hermes-bots/tests/pr127-ci-boundaries.test.mjs b/apps/desktop/src/plugins/hermes-bots/tests/pr127-ci-boundaries.test.mjs new file mode 100644 index 0000000000000..9457018451ad8 --- /dev/null +++ b/apps/desktop/src/plugins/hermes-bots/tests/pr127-ci-boundaries.test.mjs @@ -0,0 +1,81 @@ +import assert from 'node:assert/strict' +import test from 'node:test' +import { harness, members, deferred, flush, drive } from './stop-custody-harness.mjs' + +test('PR127: cancelled proven refusal retains preparation custody, then releases without admission', async () => { + const gate = deferred() + const h = await harness(members(1), { + submitError: s => s.submits === 1 ? Object.assign(new Error('session not found'), { code: 4001 }) : null, + beforeResume(s, sessions) { + if (s.submits !== 1) return + sessions.delete(s.runtime); s.runtime += '-retry'; sessions.set(s.runtime, s) + }, + resumeProjection(s, method, projection) { + if (method !== 'session.resume' || s.submits !== 1) return projection + const lazy = { ...projection }; delete lazy.turn_outcomes; return lazy + }, + async capabilityPoll() { + await gate.promise + throw Object.assign(new Error('session_id and full accepted_turn identity required'), { code: 4006 }) + } + }) + const pending = drive(h) + await flush() + const coordinator = h.gc.groupRoomCoordinators.get('Room') + const stopped = await h.gc.stopGroupThread('Room', 't1', h.roster) + assert.equal(stopped.status, 'unconfirmed') + assert.equal(coordinator.active, 1) + assert.equal(h.activeLeases(), 1) + assert.equal(h.gc.groupRuntimeSessionOwners.size, 1) + assert.equal(h.rpc('prompt.submit').length, 1) + gate.resolve(); await flush(); await h.advance(); await pending + assert.equal(h.rpc('prompt.submit').length, 1) + assert.equal(h.gc.groupRuntimeSessionOwners.size, 0) + assert.equal(coordinator.active, 0) + assert.equal(h.activeLeases(), 0) + assert.equal(h.leases[0].releases, 1) + assert.equal(Object.keys(h.room().stranded).length, 0) + assert.equal((await h.gc.stopGroupThread('Room', 't1', h.roster)).status, 'stopped') +}) + +test('PR127: a refused first submit never excuses unknown acceptance of its retry', async () => { + const h = await harness(members(1), { + submitError: s => s.submits === 1 + ? Object.assign(new Error('session not found'), { code: 4001 }) + : new Error('lost retry admission acknowledgement') + }) + await drive(h) + assert.equal(h.rpc('prompt.submit').length, 2) + assert.equal(h.activeLeases(), 1) + assert.equal(h.gc.groupRoomCoordinators.get('Room').active, 1) + const first = await h.gc.stopGroupThread('Room', 't1', h.roster) + assert.equal(first.status, 'stopping') + assert.equal(first.pending, 1) + assert.equal(h.activeLeases(), 1) + assert.equal(h.gc.groupRuntimeSessionOwners.size, 1) + const second = await h.gc.stopGroupThread('Room', 't1', h.roster) + assert.equal(second.status, 'stopping') + assert.equal(h.rpc('prompt.submit').length, 2, 'Stop never retries uncertain text') + assert.equal(h.leases[0].releases, 0) +}) + +test('PR127: four unknown retry admissions continue to block a six-member room', async () => { + const h = await harness(members(6), { + submitError: s => s.submits === 1 + ? Object.assign(new Error('session not found'), { code: 4001 }) + : new Error('lost retry acknowledgement') + }) + const pending = drive(h) + await flush() + assert.equal(h.rpc('prompt.submit').length, 8) + assert.equal(h.gc.groupRoomCoordinators.get('Room').active, 4) + assert.equal(h.gc.groupRoomCoordinators.get('Room').queue.length, 2) + assert.equal(h.activeLeases(), 4) + const stopped = await h.gc.stopGroupThread('Room', 't1', h.roster) + await h.advance(); await pending + assert.equal(stopped.status, 'stopping') + assert.equal(stopped.pending, 4) + assert.equal(h.gc.groupRoomCoordinators.get('Room').active, 4) + assert.equal(h.activeLeases(), 4) + assert.equal(h.rpc('prompt.submit').length, 8) +}) diff --git a/apps/desktop/src/plugins/hermes-bots/tests/stop-custody-harness.mjs b/apps/desktop/src/plugins/hermes-bots/tests/stop-custody-harness.mjs index caf570c4a3c1e..ada7cbac2a585 100644 --- a/apps/desktop/src/plugins/hermes-bots/tests/stop-custody-harness.mjs +++ b/apps/desktop/src/plugins/hermes-bots/tests/stop-custody-harness.mjs @@ -17,6 +17,8 @@ async function harness(roster = members(3), options = {}) { const sourceKey = (route, profile) => `${route?.connectionId || connection}::${route?.targetProfile || profile}` const handle = async (route, method, params) => { calls.push({ route: route && { ...route }, method, params: { ...params }, at: now }) + const supplied = await options.rpcResponse?.(route, method, params) + if (supplied !== undefined) return supplied const key = sourceKey(route, params.profile) if (method === 'session.create') { if (options.createGate) await options.createGate.promise @@ -53,8 +55,10 @@ async function harness(roster = members(3), options = {}) { if (method === 'clarify.respond' || method === 'approval.respond') { if (options.answerError?.()) throw new Error('response transport failed') if (options.answerGate) await options.answerGate.promise - const acknowledgement = method === 'clarify.respond' ? (options.clarifyResult ?? { status: 'ok' }) : {} + const acknowledgement = method === 'clarify.respond' + ? (options.clarifyResult ?? { status: 'ok' }) : (options.approvalResult ?? { resolved: 1 }) if (method === 'clarify.respond' && acknowledgement?.status !== 'ok') return acknowledgement + if (method === 'approval.respond' && !(Number.isSafeInteger(acknowledgement?.resolved) && acknowledgement.resolved > 0)) return acknowledgement if (method === 'clarify.respond' && session.pending?.questions?.length) { session.questionAnswers ||= new Set() session.questionAnswers.add(params.question_id) @@ -104,7 +108,7 @@ async function harness(roster = members(3), options = {}) { return () => { lease.releases++; activeLeases-- } }, state: { profile: atom('default'), gateway: atom(null), connectionId: { get: () => connection, listen: () => () => undefined } }, - notify: () => undefined, notifyError: () => undefined } + notify: notice => options.onNotify?.(notice), notifyError: () => undefined } }) gc.stopGroupChatServerSync() gc.bindGroupTurnTestStorage({ get: key => clone(storage.get(key) ?? null), set: (key, value) => storage.set(key, clone(value)) }) diff --git a/apps/desktop/src/plugins/hermes-bots/tests/stop-custody.test.mjs b/apps/desktop/src/plugins/hermes-bots/tests/stop-custody.test.mjs index 309f7634c0506..1d5f983734437 100644 --- a/apps/desktop/src/plugins/hermes-bots/tests/stop-custody.test.mjs +++ b/apps/desktop/src/plugins/hermes-bots/tests/stop-custody.test.mjs @@ -42,7 +42,7 @@ for (const mode of ['running', 'waiting', 'expired']) { if (mode === 'expired') await h.advance(21 * 60 * 1000) const count = mode === 'waiting' ? 6 : 4, workers = mode === 'waiting' ? 2 : 4 const result = await settleStop(h, pending) - assert.deepEqual(result, { status: 'unconfirmed', unconfirmed: count }) + assert.deepEqual(result, { status: 'unconfirmed', unconfirmed: count, pending: count }) assertCustody(h, count, workers) assert.equal(h.rpc('session.interrupt').length, count) assert.equal(Object.keys(h.gc.$groupClarify.get()).length, 0) @@ -56,7 +56,10 @@ for (const mode of ['running', 'waiting', 'expired']) { assert.deepEqual(h.rpc('session.interrupt').slice(count).map(c => [c.route, c.params.session_id]), firstTargets) assertCustody(h, count, workers) options.interruptError = false - assert.deepEqual(await h.gc.stopGroupThread('Room', 't1', h.roster), { status: 'stopped', unconfirmed: 0 }) + const stopped = await h.gc.stopGroupThread('Room', 't1', h.roster) + await flush() + assert.equal(stopped.unconfirmed, 0) + assert.ok(['stopping', 'stopped'].includes(stopped.status)) assertReleased(h) await h.gc.stopGroupThread('Room', 't1', h.roster) assert.equal(h.rpc('session.interrupt').length, count * 3) @@ -73,16 +76,18 @@ for (const reply of [{}, { interrupted: false }, { interrupted: true }, { status assert.equal([...h.sessions.values()][0].state, 'running') options.interruptReply = undefined await h.gc.stopGroupThread('Room', 't1', h.roster) + await flush() assertReleased(h) }) } -test('lost ACK after backend applied interrupt retains custody until exact interrupted outcome', async () => { +test('lost ACK after backend applied interrupt can retire only from exact interrupted outcome', async () => { const h = await harness(members(1), { interruptBehavior: session => { session.state = 'interrupted'; throw new Error('ACK was lost after apply') } }), pending = drive(h); await flush() - assert.equal((await settleStop(h, pending)).status, 'unconfirmed') - assertCustody(h, 1, 1) + await settleStop(h, pending) + assert.equal(h.rpc('session.turn.poll').at(-1).params.accepted_turn.request_id, + [...h.sessions.values()][0].ref.request_id, 'retirement used the accepted request despite a lost ACK') await h.gc.harvestStrandedGroupReply('Room', h.roster[0]) assertReleased(h) assert.equal(h.posts().length, 0) @@ -129,6 +134,7 @@ test('exact waiting reconciliation releases only worker, preserves stopped recei assert.equal(Object.keys(h.gc.$groupClarify.get()).length, 0) options.interruptError = false await h.gc.stopGroupThread('Room', 't1', h.roster) + await flush() assertReleased(h) }) @@ -202,7 +208,9 @@ test('late submit rejection after earlier interrupt ACK requires new-generation assert.equal(h.leases[0].releases, 0) options.interruptBehavior = undefined await h.gc.stopGroupThread('Room', 't1', h.roster) - assertReleased(h) + assert.equal(receipts(h).length, 1, 'a later ACK cannot reconstruct the lost accepted identity') + assert.equal(coordinator(h).active, 1) + assert.equal(h.leases[0].releases, 0) }) test('Stop during late session acquisition cannot submit; lease closes once after producer settles', async () => { @@ -267,7 +275,12 @@ test('reload Stop failures retain exact receipt, coalesce concurrent retry and l const a = cold.gc.stopGroupThread('Room', 't1', cold.roster), b = cold.gc.stopGroupThread('Room', 't1', cold.roster) await flush() assert.equal(cold.rpc('session.interrupt').length, 2, 'two concurrent retries issue one exact RPC') - gate.resolve(); assert.equal((await a).status, 'stopped'); assert.equal((await b).status, 'stopped') + gate.resolve() + for (const result of [await a, await b]) { + assert.equal(result.unconfirmed, 0) + assert.ok(['stopping', 'stopped'].includes(result.status)) + } + await flush() assert.equal(receipts(cold).length, 0) assert.equal(cold.rpc('prompt.submit').length, 0) }) diff --git a/ares_runtime/authority.py b/ares_runtime/authority.py index 89ebf1cc36bd4..ff93fe953bcaa 100644 --- a/ares_runtime/authority.py +++ b/ares_runtime/authority.py @@ -39,6 +39,11 @@ def _normalized_time(value: Any, field: str) -> str: return instant.astimezone(timezone.utc).isoformat().replace("+00:00", "Z") +def _time_instant(value: str) -> datetime: + """Compare normalized bounds as instants, preserving their wire spelling.""" + return datetime.fromisoformat(value.replace("Z", "+00:00")) + + def normalize_scope(raw: Any) -> dict[str, Any]: """Canonicalize a scope record; reject unknown fields and bad types.""" if not isinstance(raw, Mapping): @@ -64,7 +69,12 @@ def normalize_scope(raw: Any) -> dict[str, Any]: if not value: raise ContractError("INVALID_TIME_SCOPE") normalized_time = {bound: _normalized_time(value[bound], "time." + bound) for bound in ("not_before", "not_after") if bound in value} - if "not_before" in normalized_time and "not_after" in normalized_time and normalized_time["not_before"] > normalized_time["not_after"]: + if ( + "not_before" in normalized_time + and "not_after" in normalized_time + and _time_instant(normalized_time["not_before"]) + > _time_instant(normalized_time["not_after"]) + ): raise ContractError("INVALID_TIME_SCOPE") normalized[field] = normalized_time else: @@ -108,9 +118,15 @@ def is_subset_scope(subset: Mapping[str, Any], superset: Mapping[str, Any]) -> b return False elif field == "time": sa, sb = a[field], b[field] - if "not_before" in sb and ("not_before" not in sa or sa["not_before"] < sb["not_before"]): + if "not_before" in sb and ( + "not_before" not in sa + or _time_instant(sa["not_before"]) < _time_instant(sb["not_before"]) + ): return False - if "not_after" in sb and ("not_after" not in sa or sa["not_after"] > sb["not_after"]): + if "not_after" in sb and ( + "not_after" not in sa + or _time_instant(sa["not_after"]) > _time_instant(sb["not_after"]) + ): return False else: if a[field] != b[field]: @@ -168,8 +184,9 @@ def attenuate(self, child_scope: Mapping[str, Any], *, child_generation: int, ch child_uses = inherited.get("use_count", 1) if "use_count" in self._scope and self._charged_count + self._delegated_count + child_uses > self._scope["use_count"]: raise ContractError("USE_COUNT_EXHAUSTED") + child = AuthorityScopeV1(scope=inherited, generation=child_generation, holder=child_holder) self._delegated_count += child_uses - return AuthorityScopeV1(scope=inherited, generation=child_generation, holder=child_holder) + return child def reserve(self, *, consumption_ref: str, args_digest: str, target_ref: str | None = None) -> dict[str, Any]: """Open a reservation against finite remaining use.""" @@ -185,24 +202,35 @@ def reserve(self, *, consumption_ref: str, args_digest: str, target_ref: str | N target_ref = _require_str(target_ref, "INVALID_TARGET_REF") if "target" in self._scope and target_ref != self._scope["target"]: raise ContractError("TARGET_OUTSIDE_SCOPE") - self._settlements[consumption_ref] = { + record = { "consumption_ref": consumption_ref, "state": "reserved", "args_digest": args_digest, "target_ref": target_ref, } - self._open_count += 1 - self._charged_count += 1 - return self._receipt(consumption_ref) + open_count = self._open_count + 1 + charged_count = self._charged_count + 1 + receipt = self._build_receipt(record, open_count=open_count, charged_count=charged_count) + self._settlements[consumption_ref] = record + self._open_count = open_count + self._charged_count = charged_count + return receipt def commit(self, consumption_ref: str, *, effect_receipt_digest: str) -> dict[str, Any]: + effect_receipt_digest = _require_str(effect_receipt_digest, "INVALID_EFFECT_RECEIPT_DIGEST") + if not effect_receipt_digest.strip(): + raise ContractError("INVALID_EFFECT_RECEIPT_DIGEST") return self._settle(consumption_ref, "committed", effect_receipt_digest=effect_receipt_digest) def release(self, consumption_ref: str, *, reason: str = "operator_release") -> dict[str, Any]: + reason = _require_str(reason, "INVALID_REASON") + if not reason.strip(): + raise ContractError("INVALID_REASON") return self._settle(consumption_ref, "released", reason=reason) def mark_indeterminate(self, consumption_ref: str, *, reason: str) -> dict[str, Any]: - if not _require_str(reason, "INVALID_REASON"): + reason = _require_str(reason, "INVALID_REASON") + if not reason.strip(): raise ContractError("INVALID_REASON") return self._settle(consumption_ref, "indeterminate", reason=reason) @@ -213,24 +241,34 @@ def _settle(self, consumption_ref: str, state: str, **extra: str) -> dict[str, A raise ContractError("UNKNOWN_CONSUMPTION_REF") if record["state"] != "reserved": raise ContractError("ALREADY_SETTLED") - record["state"] = state - record.update(extra) - self._open_count -= 1 - if state == "released": - self._charged_count -= 1 - return self._receipt(consumption_ref) + candidate = dict(record) + candidate["state"] = state + candidate.update(extra) + open_count = self._open_count - 1 + charged_count = self._charged_count - int(state == "released") + receipt = self._build_receipt(candidate, open_count=open_count, charged_count=charged_count) + self._settlements[consumption_ref] = candidate + self._open_count = open_count + self._charged_count = charged_count + return receipt def _receipt(self, consumption_ref: str) -> dict[str, Any]: - record = self._settlements[consumption_ref] + return self._build_receipt( + self._settlements[consumption_ref], + open_count=self._open_count, + charged_count=self._charged_count, + ) + + def _build_receipt(self, record: Mapping[str, Any], *, open_count: int, charged_count: int) -> dict[str, Any]: receipt = { "schema": SCHEMA_VERSION, "scope_fingerprint": self.fingerprint(), "generation": self._generation, "holder": self._holder, "record": dict(record), - "open_reservations": self._open_count, - "charged_total": self._charged_count, - "consumed_total": self._charged_count, + "open_reservations": open_count, + "charged_total": charged_count, + "consumed_total": charged_count, "delegated_total": self._delegated_count, } receipt["receipt_digest"] = digest(receipt) diff --git a/ares_runtime/collaboration.py b/ares_runtime/collaboration.py index 7887c09be4e59..3820e861927e8 100644 --- a/ares_runtime/collaboration.py +++ b/ares_runtime/collaboration.py @@ -3112,6 +3112,7 @@ def project( _check_ref(event_ref, "source_event_refs") if not source_event_exists(event_ref): raise ContractError("MISSING_SOURCE_EVENT", event_ref) + prior = None if previous_projection is not None: prior = ( previous_projection.to_dict() @@ -3123,10 +3124,6 @@ def project( or prior.get("closure_profile") != closure_profile ): raise ContractError("PROJECTION_LINEAGE_MISMATCH") - if prior.get("state") == "closed" and not set(normalized_events).difference( - prior.get("source_event_refs", []) - ): - raise ContractError("REOPEN_REQUIRES_NEW_EVIDENCE") unknown = set(flags) - self.ALLOWED_FLAGS if unknown: raise ContractError("UNKNOWN_DIVERGENCE_FLAG", sorted(unknown)[0]) @@ -3140,7 +3137,7 @@ def project( if not unsatisfied and not flags else ("quarantined" if "AMBIGUOUS_EFFECT" in flags else "evidence_pending") ) - return make_artifact( + projection = make_artifact( "closure", { "mission_ref": mission_ref, @@ -3154,6 +3151,16 @@ def project( "divergence_flags": sorted(set(flags)), }, ) + if ( + prior is not None + and prior.get("state") == "closed" + and not set(normalized_events).difference(prior.get("source_event_refs", [])) + and projection.to_dict() != prior + ): + # Replaying identical owner evidence is a rebuild, not a reopen. + # Changed gates, flags or provenance still require a new event. + raise ContractError("REOPEN_REQUIRES_NEW_EVIDENCE") + return projection BASELINE_DEFINITIONS: Mapping[str, Mapping[str, str]] = { diff --git a/ares_runtime/continuity/budget.py b/ares_runtime/continuity/budget.py index c9ccf276bd3d3..62db89b3e9cb5 100644 --- a/ares_runtime/continuity/budget.py +++ b/ares_runtime/continuity/budget.py @@ -142,20 +142,54 @@ def final_request_upper_bound(*, route_ref: str, payload: dict[str, Any]) -> Inp "timeout", "extra_headers", "headers", }} - def check(value, depth=0): + # Only recognized tool-definition positions contain JSON Schema data. + # A schema-looking object elsewhere remains transport input. The schema + # bytes still participate in depth/JSON validation, counting and hashing. + schema_roots = set() + tools = body.get("tools") + if isinstance(tools, list): + for index, tool in enumerate(tools): + if type(tool) is not dict: + continue + if tool.get("type") == "function": + function = tool.get("function") + if (type(function) is dict and isinstance(function.get("name"), str) + and function["name"] and type(function.get("parameters")) is dict): + schema_roots.add(("tools", index, "function", "parameters")) + elif (isinstance(tool.get("name"), str) and tool["name"] + and type(tool.get("parameters")) is dict): + schema_roots.add(("tools", index, "parameters")) + elif (tool.get("type") in (None, "custom") + and isinstance(tool.get("name"), str) and tool["name"] + and type(tool.get("input_schema")) is dict): + schema_roots.add(("tools", index, "input_schema")) + tool_config = body.get("toolConfig") + if type(tool_config) is dict and isinstance(tool_config.get("tools"), list): + for index, tool in enumerate(tool_config["tools"]): + spec = tool.get("toolSpec") if type(tool) is dict else None + inputs = spec.get("inputSchema") if type(spec) is dict else None + if (type(spec) is dict and isinstance(spec.get("name"), str) and spec["name"] + and type(inputs) is dict and type(inputs.get("json")) is dict): + schema_roots.add(("toolConfig", "tools", index, "toolSpec", "inputSchema", "json")) + + def check(value, depth=0, path=(), schema=False): if depth > 64: raise BudgetError("INVALID_FINAL_PAYLOAD") + schema = schema or path in schema_roots if isinstance(value, dict): - if value.get("type") in { + kind = value.get("type") + if not schema and isinstance(kind, (dict, list)): + raise BudgetError("INVALID_FINAL_PAYLOAD") + if not schema and ((isinstance(kind, str) and kind in { "image", "image_url", "input_image", "input_file", "file", "audio", "input_audio", "compaction", "computer_screenshot", - } or any(key in value for key in ("encrypted_content", "previous_response_id", "conversation")): + }) or any(key in value for key in ("encrypted_content", "previous_response_id", "conversation"))): raise BudgetError("FINAL_PAYLOAD_OPAQUE_ACCOUNTING_UNQUALIFIED") - for item in value.values(): - check(item, depth + 1) + for key, item in value.items(): + check(item, depth + 1, path + (key,), schema) elif isinstance(value, list): - for item in value: - check(item, depth + 1) + for index, item in enumerate(value): + check(item, depth + 1, path + (index,), schema) check(body) try: diff --git a/ares_runtime/local_runtime.py b/ares_runtime/local_runtime.py index a6d5d6cec11fc..06e05a3e639db 100644 --- a/ares_runtime/local_runtime.py +++ b/ares_runtime/local_runtime.py @@ -36,6 +36,47 @@ class AresLocalRuntimeError(RuntimeError): """Raised when the explicit local-runtime contract is not satisfied.""" +def _prepare_gateway_stop_marker() -> None: + """Publish stop intent through the gateway owner in its process home.""" + from gateway import status + + pid_path = status._get_pid_path() + try: + identity = status.get_running_pid_identity_strict(pid_path) + except (OSError, RuntimeError) as exc: + raise AresLocalRuntimeError( + f"Ares gateway stop identity is ambiguous: {exc}" + ) from exc + if identity is None: + return + home = pid_path.parent + records = ( + status._read_pid_record(pid_path), + status._read_gateway_lock_record(status._get_gateway_lock_path(pid_path)), + ) + for record in records: + recorded_home = record.get("hermes_home") + if ( + not isinstance(recorded_home, str) + or not recorded_home.strip() + or not status._same_hermes_home(recorded_home, home) + or not status._record_matches_live_gateway_pid( + record, identity[0], expected_home=home + ) + ): + raise AresLocalRuntimeError("Ares gateway stop identity has an ambiguous home") + try: + current = status.get_running_pid_identity_strict(pid_path) + except (OSError, RuntimeError) as exc: + raise AresLocalRuntimeError( + f"Ares gateway stop identity changed: {exc}" + ) from exc + if current != identity: + raise AresLocalRuntimeError("Ares gateway stop identity changed before marker publication") + if not status.write_planned_stop_marker(identity[0]): + raise AresLocalRuntimeError("Ares gateway planned-stop marker could not be written") + + def _desktop_launch_arguments( executable: Path, *, @@ -939,6 +980,9 @@ def _materialize_upstream_candidate( self._require_complete_release( self._release_source(candidate_revision), desktop=desktop ) + # The verified release is reused; the fresh staging checkout + # remains owned by this invocation and must not accumulate. + shutil.rmtree(staging) return candidate_revision self._build_runtime(source, desktop=desktop) self._atomic_json( @@ -1162,8 +1206,38 @@ def setup( raise return revision, seeded + def _require_base_update_recipe(self) -> None: + """Refuse recipes the ordinary updater cannot reproduce.""" + + current = self._release_from_link(self.paths.current_link, "current") + if current is None: + return + record = current[1] / ".venv" / "share" / "ares-full-install.json" + try: + # The website recipe includes SDK overrides even without native + # enhancements. Its record is unsupported by this base builder; + # malformed records and dangling links must not mean "base". + record.lstat() + except FileNotFoundError: + return + except OSError as exc: + raise AresLocalRuntimeError( + "Cannot inspect installed recipe metadata; ares update cannot " + "establish a supported base recipe. Current and previous " + "releases are preserved. Reconcile the installation recipe " + "before retrying." + ) from exc + raise AresLocalRuntimeError( + "ares update cannot preserve the recorded installer recipe. " + "Current and previous releases are preserved. Rerun the official " + "full-distribution installer for a new Ares source revision. " + "Dependency-only changes at the same revision require recipe-aware " + "update support; keep the complete release selected." + ) + def update(self, *, desktop: bool) -> tuple[str, bool]: with self.locked(): + self._require_base_update_recipe() config = self._read_config() remote = str(config["remote"]) branch = str(config["branch"]) @@ -1224,6 +1298,15 @@ def rollback(self) -> str: try: self._atomic_link(self.paths.current_link, previous[1]) self._atomic_link(self.paths.previous_link, current[1]) + if self.paths.unit_path.exists(): + self._systemctl("restart", "ares-gateway.service") + time.sleep(1) + if not self._systemctl( + "is-active", "--quiet", "ares-gateway.service", required=False + ): + raise AresLocalRuntimeError( + "Ares gateway did not remain active after rollback" + ) except Exception: try: self._restore_release_pair(current, previous) @@ -1231,19 +1314,16 @@ def rollback(self) -> str: raise AresLocalRuntimeError( "Ares rollback failed and the prior release pointers could not be restored" ) from restore_exc + if self.paths.unit_path.exists(): + try: + self._systemctl( + "restart", "ares-gateway.service", required=False + ) + except Exception: + # Recovery is best effort; preserve the rollback error + # after restoring the authoritative pointer pair. + pass raise - if self.paths.unit_path.exists(): - self._systemctl("restart", "ares-gateway.service") - time.sleep(1) - if not self._systemctl( - "is-active", "--quiet", "ares-gateway.service", required=False - ): - self._atomic_link(self.paths.current_link, current[1]) - self._atomic_link(self.paths.previous_link, previous[1]) - self._systemctl("restart", "ares-gateway.service", required=False) - raise AresLocalRuntimeError( - "Ares gateway did not remain active after rollback" - ) return previous[0] @staticmethod @@ -1455,12 +1535,49 @@ def tui(self, arguments: Sequence[str]) -> None: def chat(self, arguments: Sequence[str]) -> None: self._exec_hermes(arguments) + def _prepare_gateway_stop(self) -> None: + # Gateway identity/marker paths deliberately ignore task-local home + # overrides. Scope their canonical owner in a child without changing + # the caller's process environment or active profile context. + _, source = self.active_release() + python = self._python_for(source) + environment = self._agent_environment() + environment["HERMES_HOME"] = str(self.paths.agent_home.expanduser().resolve()) + environment["PYTHONDONTWRITEBYTECODE"] = "1" + try: + completed = subprocess.run( + [ + str(python), + "-c", + "from ares_runtime.local_runtime import _prepare_gateway_stop_marker; " + "_prepare_gateway_stop_marker()", + ], + cwd=source, + env=environment, + text=True, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + check=False, + timeout=10, + ) + except (OSError, subprocess.TimeoutExpired) as exc: + raise AresLocalRuntimeError( + f"Ares gateway stop preparation failed: {exc}" + ) from exc + if completed.returncode: + detail = (completed.stderr or completed.stdout).strip() + raise AresLocalRuntimeError( + "Ares gateway stop preparation failed" + + (f": {detail}" if detail else "") + ) + def gateway(self, action: str) -> None: if action == "foreground": self._exec_hermes(["gateway"]) if action == "start": self._systemctl("enable", "--now", "ares-gateway.service") elif action == "stop": + self._prepare_gateway_stop() self._systemctl("disable", "--now", "ares-gateway.service") elif action == "restart": self._systemctl("restart", "ares-gateway.service") diff --git a/ares_runtime/specialist_dispatch.py b/ares_runtime/specialist_dispatch.py index 9a3506dac95ae..a1de2614e64ca 100644 --- a/ares_runtime/specialist_dispatch.py +++ b/ares_runtime/specialist_dispatch.py @@ -383,7 +383,16 @@ def _run_profile_worker(profile_id: str, request: ExplicitDispatchRequest, recei profile_home = root / "profiles" / profile_id if not profile_home.is_dir(): return {"outcome": "runner_failed", "exit_code": None, "error_type": "PROFILE_HOME_MISSING"} - environment = os.environ.copy() + from agent.secret_scope import build_profile_env_boundary + from hermes_constants import get_process_hermes_home + from tools.environments.local import hermes_subprocess_env + + boundary = build_profile_env_boundary( + source_home=get_process_hermes_home(), target_home=profile_home, + ) + environment = hermes_subprocess_env( + inherit_credentials=True, profile_boundary=boundary, + ) environment.update({"HERMES_HOME": str(profile_home), "ARES_MANAGED_RUNTIME": "1", "HERMES_SESSION_SOURCE": "cli"}) command = [sys.executable, "-m", "hermes_cli.main", "--in", str(request.workspace), "-z", request.brief] process = subprocess.Popen(command, cwd=request.workspace, env=environment, stdout=subprocess.PIPE, stderr=subprocess.PIPE) diff --git a/cron/jobs.py b/cron/jobs.py index 32d54a22c7cac..5c06e4bf4d7e4 100644 --- a/cron/jobs.py +++ b/cron/jobs.py @@ -3100,6 +3100,24 @@ def _machine_id() -> str: return f"{host}:{os.getpid()}" +@dataclass(frozen=True) +class UnstartedFireClaimReceipt: + """Ephemeral fence for a claim whose runner is proven unstarted. + + Presence flags distinguish missing fields from explicit nulls. Values + are detached snapshots; this descriptor is never stored in jobs.json. + """ + + job_id: str + owner: str + prior_next_present: bool + prior_next_value: Any + claimed_next_present: bool + claimed_next_value: Any + claimed_schedule_present: bool + claimed_schedule_value: Any + + def claim_job_for_fire( job_id: str, *, @@ -3124,7 +3142,8 @@ def _claim_job_for_fire_locked( claim_ttl_seconds: int = 300, force: bool = False, return_job: bool = False, -) -> Union[bool, Dict[str, Any]]: + return_unstarted_receipt: bool = False, +) -> Union[bool, Dict[str, Any], Tuple[Dict[str, Any], UnstartedFireClaimReceipt]]: """Atomically claim a job for a single external 'fire' (multi-machine at-most-once). Returns True iff THIS caller won the claim. @@ -3177,6 +3196,9 @@ def _claim_job_for_fire_locked( return False # someone holds a fresh claim except Exception: pass # malformed claim → overwrite + if return_unstarted_receipt: + prior_next_present = "next_run_at" in job + prior_next_value = copy.deepcopy(job.get("next_run_at")) if force: job["enabled"] = True job["state"] = "scheduled" @@ -3192,11 +3214,113 @@ def _claim_job_for_fire_locked( nxt = compute_next_run(job["schedule"], now.isoformat()) if nxt: job["next_run_at"] = nxt + if return_unstarted_receipt: + claimed = copy.deepcopy(job) + receipt = UnstartedFireClaimReceipt( + job_id=job_id, + owner=owner, + prior_next_present=prior_next_present, + prior_next_value=prior_next_value, + claimed_next_present="next_run_at" in job, + claimed_next_value=copy.deepcopy(job.get("next_run_at")), + claimed_schedule_present="schedule" in job, + claimed_schedule_value=copy.deepcopy(job.get("schedule")), + ) + # Prepare both snapshots before publication so copying cannot + # turn a confirmed claim into a post-save descriptor failure. + save_jobs(jobs) + return claimed, receipt save_jobs(jobs) return copy.deepcopy(job) if return_job else True return False +def claim_job_for_fire_with_unstarted_receipt( + job_id: str, + *, + claim_ttl_seconds: int = 300, +) -> Optional[Tuple[Dict[str, Any], UnstartedFireClaimReceipt]]: + """Claim without forced resume and return an ephemeral release fence. + + Ordinary claim callers keep their existing bool/dict contract. This + opt-in path is for dispatch admission that can prove work never started. + """ + with _fire_job_lock(job_id) as acquired: + if not acquired: + return None + result = _claim_job_for_fire_locked( + job_id, + claim_ttl_seconds=claim_ttl_seconds, + force=False, + return_job=True, + return_unstarted_receipt=True, + ) + return result if isinstance(result, tuple) else None + + +def release_unstarted_fire_claim( + receipt: UnstartedFireClaimReceipt, +) -> Dict[str, Any]: + """Release only a positively unstarted claim; never an uncertain submit. + + Owner and field fences preserve concurrent edits. A receipt is not a + general cancellation API and does not authorize restoring a whole job. + """ + if not isinstance(receipt, UnstartedFireClaimReceipt): + return {"released": False, "status": "conflict", "reason": "invalid_receipt"} + try: + with _fire_job_lock(receipt.job_id) as acquired: + if not acquired: + return {"released": False, "status": "lock_unavailable"} + return _release_unstarted_fire_claim_locked(receipt) + except Exception as exc: + # Lock/load failures cannot certify the claim's current disposition. + return { + "released": False, "status": "release_unavailable", + "claim_release": "uncertain", "error": str(exc), + } + + +def _release_unstarted_fire_claim_locked( + receipt: UnstartedFireClaimReceipt, +) -> Dict[str, Any]: + """Release under the existing fire-job then jobs-store lock order.""" + with _jobs_lock(): + jobs = load_jobs() + for job in jobs: + if job.get("id") != receipt.job_id: + continue + claim = job.get("fire_claim") + if not isinstance(claim, dict) or claim.get("by") != receipt.owner: + return {"released": False, "status": "conflict", "reason": "owner_changed"} + if ( + ("schedule" in job) != receipt.claimed_schedule_present + or job.get("schedule") != receipt.claimed_schedule_value + ): + return {"released": False, "status": "conflict", "reason": "schedule_changed"} + if ( + ("next_run_at" in job) != receipt.claimed_next_present + or job.get("next_run_at") != receipt.claimed_next_value + ): + return {"released": False, "status": "conflict", "reason": "next_run_changed"} + job.pop("fire_claim", None) + if receipt.prior_next_present: + job["next_run_at"] = copy.deepcopy(receipt.prior_next_value) + else: + job.pop("next_run_at", None) + try: + save_jobs(jobs) + except Exception as exc: + # A save may raise after publication. Do not certify release + # or claim retention and never compensate by running the job. + return { + "released": False, "status": "write_uncertain", + "claim_release": "uncertain", "error": str(exc), + } + return {"released": True, "status": "released"} + return {"released": False, "status": "missing"} + + # Completed one-shot job records are retained in jobs.json (final status + # delivery error stay inspectable via `cronjob list`) instead of being deleted # at completion, then pruned by _sweep_completed_oneshots once they age out. diff --git a/cron/scheduler.py b/cron/scheduler.py index a96b4175aefaf..ddd960784c17c 100644 --- a/cron/scheduler.py +++ b/cron/scheduler.py @@ -2664,12 +2664,32 @@ def _deliver_to_bot_chat(job: dict, content: str, profile: str) -> Optional[str] except Exception: return "bot-chat delivery failed: hermes CLI not resolvable" - env = os.environ.copy() + try: + from agent.secret_scope import ( + build_profile_env_boundary, build_profile_secret_scope, + reset_secret_scope, set_secret_scope, + ) + from hermes_constants import get_hermes_home, get_process_hermes_home + from hermes_cli.profiles import resolve_profile_env + from tools.environments.local import hermes_subprocess_env + + target_home = Path(resolve_profile_env(profile)) if profile else get_hermes_home() + scope_token = set_secret_scope(build_profile_secret_scope( + target_home, fail_closed_external=True, + )) + try: + boundary = build_profile_env_boundary( + source_home=get_process_hermes_home(), target_home=target_home, + ) + env = hermes_subprocess_env( + inherit_credentials=True, profile_boundary=boundary, + ) + finally: + reset_secret_scope(scope_token) + except Exception: + return "bot-chat delivery failed: profile authority unavailable" if profile: argv += ["-p", profile] - # -p owns profile resolution in the child; a leftover HERMES_HOME - # from THIS scheduler's profile must not shadow it. - env.pop("HERMES_HOME", None) # The prefix tells the receiving bot this is scheduled output, not the # human typing — mirrors the Bot Mode sender-attribution convention. diff --git a/gateway/authz_mixin.py b/gateway/authz_mixin.py index 6dc33534f9244..cfbd426e59dd3 100644 --- a/gateway/authz_mixin.py +++ b/gateway/authz_mixin.py @@ -506,7 +506,14 @@ def _is_user_authorized( } if getattr(source, "is_bot", False): allow_bots_var = platform_allow_bots_map.get(source.platform) - if allow_bots_var and _platform_gate_env(allow_bots_var, "none").lower().strip() in {"mentions", "all"}: + allow_bots = _platform_gate_env(allow_bots_var, "") if allow_bots_var else "" + if not allow_bots and source.platform == Platform.FEISHU: + # The existing adapter resolver refuses a default-profile + # fallback for secondary sources. Share its resolved target + # YAML/env admission policy instead of borrowing process env. + adapter = self._authorization_adapter(source.platform, adapter_profile) + allow_bots = str(getattr(adapter, "_allow_bots", "none")) + if allow_bots.lower().strip() in {"mentions", "all"}: return True if not user_id: diff --git a/gateway/session.py b/gateway/session.py index d7770188613a3..2ca7f11e276fa 100644 --- a/gateway/session.py +++ b/gateway/session.py @@ -97,6 +97,36 @@ def _hash_chat_id(value: str) -> str: from utils import atomic_replace from agent.turn_context import extract_api_content_sidecar + +def _append_transcript_message_to_db(db, session_id: str, message: Dict[str, Any]) -> None: + """Persist the supported transcript fields through the existing DB owner.""" + db.append_message( + session_id=session_id, + role=message.get("role", "unknown"), + content=message.get("content"), + tool_name=message.get("tool_name"), + tool_calls=message.get("tool_calls"), + tool_call_id=message.get("tool_call_id"), + reasoning=message.get("reasoning") if message.get("role") == "assistant" else None, + reasoning_content=message.get("reasoning_content") if message.get("role") == "assistant" else None, + reasoning_details=message.get("reasoning_details") if message.get("role") == "assistant" else None, + codex_reasoning_items=message.get("codex_reasoning_items") if message.get("role") == "assistant" else None, + codex_message_items=message.get("codex_message_items") if message.get("role") == "assistant" else None, + platform_message_id=(message.get("platform_message_id") or message.get("message_id")), + observed=bool(message.get("observed")), + timestamp=message.get("timestamp"), + # api_content sidecar: the exact bytes sent to the API for + # this message (prompt-cache-stable replay). Must survive + # any gateway-side persistence path or the next turn's + # replay diverges at this row. + api_content=extract_api_content_sidecar(message), + # Presentation typing (e.g. "internal_notification" for + # self-injected async-delegation/background notification turns, + # #82888). DB-only; stripped from provider-bound payloads. + display_kind=message.get("display_kind"), + display_metadata=message.get("display_metadata"), + ) + # Session keys/ids flow into filesystem paths downstream (e.g. # ``sessions_dir / f"{session_id}.json"`` in hermes_state, request-dump # filenames in agent_runtime_helpers). Any value that could escape the @@ -3897,32 +3927,7 @@ def _drain_spooled_drops(self, session_id: str) -> None: def _append_transcript_message(self, session_id: str, message: Dict[str, Any]) -> None: """Write one transcript row. Caller handles retry queuing.""" - self._db.append_message( - session_id=session_id, - role=message.get("role", "unknown"), - content=message.get("content"), - tool_name=message.get("tool_name"), - tool_calls=message.get("tool_calls"), - tool_call_id=message.get("tool_call_id"), - reasoning=message.get("reasoning") if message.get("role") == "assistant" else None, - reasoning_content=message.get("reasoning_content") if message.get("role") == "assistant" else None, - reasoning_details=message.get("reasoning_details") if message.get("role") == "assistant" else None, - codex_reasoning_items=message.get("codex_reasoning_items") if message.get("role") == "assistant" else None, - codex_message_items=message.get("codex_message_items") if message.get("role") == "assistant" else None, - platform_message_id=(message.get("platform_message_id") or message.get("message_id")), - observed=bool(message.get("observed")), - timestamp=message.get("timestamp"), - # api_content sidecar: the exact bytes sent to the API for - # this message (prompt-cache-stable replay). Must survive - # any gateway-side persistence path or the next turn's - # replay diverges at this row. - api_content=extract_api_content_sidecar(message), - # Presentation typing (e.g. "internal_notification" for - # self-injected async-delegation/background notification turns, - # #82888). DB-only; stripped from provider-bound payloads. - display_kind=message.get("display_kind"), - display_metadata=message.get("display_metadata"), - ) + _append_transcript_message_to_db(self._db, session_id, message) # Maximum in-memory pending messages per session before dropping the # oldest. Prevents unbounded growth when the DB is persistently broken. diff --git a/gateway/shutdown_flush.py b/gateway/shutdown_flush.py index a727dd975a011..ce18e7840438d 100644 --- a/gateway/shutdown_flush.py +++ b/gateway/shutdown_flush.py @@ -347,12 +347,13 @@ def _close_owned_db() -> None: path, ) continue - session_db.append_message( - session_id=spooled_sid, - role=message.get("role", "unknown"), - content=message.get("content") or "", - timestamp=message.get("timestamp") or payload.get("ts"), - ) + from gateway.session import _append_transcript_message_to_db + + recovered_message = dict(message) + # Preserve the established recovery coercion and timestamp fallback. + recovered_message["content"] = message.get("content") or "" + recovered_message["timestamp"] = message.get("timestamp") or payload.get("ts") + _append_transcript_message_to_db(session_db, spooled_sid, recovered_message) recovered += 1 path.unlink(missing_ok=True) continue diff --git a/hermes_cli/backup.py b/hermes_cli/backup.py index 6a3377fdbba3a..843eb846a9736 100644 --- a/hermes_cli/backup.py +++ b/hermes_cli/backup.py @@ -1101,12 +1101,10 @@ def _extract_member_atomically( ``atomic_replace`` rather than a bare ``os.replace``: it resolves a symlinked target first, so a deployment that links ``config.yaml`` into a dotfiles repo keeps the link instead of having it silently swapped for a - regular file (GitHub #16743), and it falls back to copy/fsync/unlink on - ``EXDEV``/``EBUSY`` for cross-device and bind-mount installs. That - fallback uses ``shutil.copyfile``, which does truncate in place, so on the - cross-device path the guarantee above degrades to today's behaviour rather - than improving on it; closing that belongs in ``utils.atomic_replace``, - where every atomic writer in the repo would benefit, not here. + regular file (GitHub #16743). On EXDEV the shared helper restages beside + the resolved destination before publication. Busy or persistently + contended destinations fail without an in-place overwrite, preserving + the previous complete file for retry. Permission bits *and* ownership are carried across the replace so routing through mkstemp does not change the file the caller would otherwise have @@ -1156,7 +1154,7 @@ def _extract_member_atomically( if mode is not None: # Apply the mode to the temp file BEFORE the replace so the # target never transits through mkstemp's 0600, and so - # ``atomic_replace``'s EXDEV/EBUSY ``shutil.copystat`` fallback + # ``atomic_replace``'s EXDEV restaging ``shutil.copystat`` # copies the intended bits rather than 0600. fchmod is # Unix-only; Windows takes the path-based chmod. if hasattr(os, "fchmod"): @@ -1328,7 +1326,10 @@ def run_import(args) -> None: # Summary print() - print(f"Import complete: {restored} files restored in {elapsed:.1f}s") + if errors: + print(f"Import incomplete: {restored} files restored, {len(errors)} failed in {elapsed:.1f}s") + else: + print(f"Import complete: {restored} files restored in {elapsed:.1f}s") print(f" Target: {display_hermes_home()}") if restored_external: @@ -1338,7 +1339,7 @@ def run_import(args) -> None: ) if errors: - print(f"\n Warnings ({len(errors)} files skipped):") + print(f"\n Failures ({len(errors)} files not fully restored):") for e in errors[:10]: print(e) if len(errors) > 10: @@ -1354,6 +1355,10 @@ def run_import(args) -> None: if len(skipped_runtime) > 10: print(f" ... and {len(skipped_runtime) - 10} more") + if errors: + # A partial restore must not activate mixed old/new configuration. + raise SystemExit(1) + # Post-import: restore profile wrapper scripts profiles_dir = hermes_root / "profiles" restored_profiles = [] diff --git a/hermes_cli/logs.py b/hermes_cli/logs.py index a214d52c8a5f3..b4fa75a607b8f 100644 --- a/hermes_cli/logs.py +++ b/hermes_cli/logs.py @@ -19,10 +19,15 @@ hermes logs --since 30m -f # follow, starting 30 min ago """ +import io +import json +import os import re +import stat import sys import time from datetime import datetime, timedelta +from collections import deque from pathlib import Path from typing import Optional, Sequence @@ -35,8 +40,7 @@ "gateway": "gateway.log", "gui": "gui.log", "desktop": "desktop.log", - # Every stdio MCP subprocess's stderr (tools/mcp_tool.py redirects it - # here, with per-server session markers) — the "MCP output channel". + # Legacy raw stderr and the index of profile-owned stdio attempt files. "mcp": "mcp-stderr.log", } @@ -215,7 +219,8 @@ def tail_log( # Read and display the tail try: - lines = _read_tail(log_path, num_lines, has_filters=has_filters, + read_tail = _read_mcp_tail if log_name == "mcp" else _read_tail + lines = read_tail(log_path, num_lines, has_filters=has_filters, min_level=min_level, session_filter=session, since=since_dt, component_prefixes=component_prefixes) except PermissionError: @@ -247,12 +252,257 @@ def tail_log( # Follow mode — poll for new content try: - _follow_log(log_path, min_level=min_level, session_filter=session, + follow_log = _follow_mcp_log if log_name == "mcp" else _follow_log + follow_log(log_path, min_level=min_level, session_filter=session, since=since_dt, component_prefixes=component_prefixes) except KeyboardInterrupt: print("\n--- stopped ---") +_MCP_MAX_ATTEMPTS = 128 +_MCP_INDEX_LINES = 2000 +_MCP_TAIL_BYTES = 1048576 +_MCP_MAX_PREVIEW_LINES = 10000 +_MCP_FOLLOW_BYTES = 65536 +_MCP_MAX_LINE_BYTES = 65536 + + +class _MCPPreviewNotice(str): + """Reader metadata, distinct from child content subject to log filters.""" + + +def _mcp_open_regular(path: Path): + return open(path, "rb", opener=lambda name, flags: os.open( + name, flags | getattr(os, "O_NOFOLLOW", 0) | getattr(os, "O_NONBLOCK", 0) + )) + + +def _mcp_oversized_line(index: bool) -> str: + source = "index" if index else "stderr" + return _MCPPreviewNotice(f"MCP {source} line exceeds MCP preview limit; full content remains in the log file.\n") + + +def _mcp_tail_rows(path: Path, n: int, *, index: bool = False) -> list: + """A finite byte window; the generic log reader is unchanged.""" + if n <= 0: + return [] + rows = deque(maxlen=min(n, _MCP_MAX_PREVIEW_LINES)) + try: + with _mcp_open_regular(path) as stream: + info = os.fstat(stream.fileno()) + if not stat.S_ISREG(info.st_mode): + return [] + start = max(0, info.st_size - _MCP_TAIL_BYTES) + stream.seek(start) + data = stream.read(_MCP_TAIL_BYTES) + if start: + fragment, separator, data = data.partition(b"\n") + if not separator or len(fragment) > _MCP_MAX_LINE_BYTES: + rows.append(_mcp_oversized_line(index)) + for row in io.BytesIO(data): + payload = row.removesuffix(b"\n").removesuffix(b"\r") + if len(payload) > _MCP_MAX_LINE_BYTES: + rows.append(_mcp_oversized_line(index)) + else: + decoded = row.decode("utf-8", errors="replace") + # An EOF preview is a display record, even before the child + # finishes its line. The stored bytes remain unchanged. + rows.append(decoded if decoded.endswith("\n") else decoded + "\n") + except OSError: + return [] + return list(rows) + + +def _mcp_control_record(line: str) -> Optional[dict]: + try: + record = json.loads(line) + except (ValueError, TypeError): + return None + return record if isinstance(record, dict) and record.get("kind") == "mcp.stdio.attempt" else None + + +def _mcp_attempt(record: dict, index: Path) -> Optional[dict]: + """Resolve only this profile's UUID-named capture, never an arbitrary path.""" + attempt = record.get("attempt_id") + if (not isinstance(attempt, str) or re.fullmatch(r"[0-9a-f]{32}", attempt) is None + or not isinstance(record.get("server"), str) + or type(record.get("parent_pid")) is not int or record["parent_pid"] <= 0 + or record.get("destination") != "file" + or not isinstance(record.get("config_home"), str) + or not isinstance(record.get("stderr_path"), str)): + return None + try: + home = index.parent.parent.resolve() + expected = home / "logs" / "mcp-stderr" / f"{attempt}.log" + if (Path(record["config_home"]).resolve() != home + or Path(record["stderr_path"]).resolve() != expected + or not expected.is_file()): + return None + except (OSError, ValueError, RuntimeError): + return None + return {"attempt_id": attempt, "server": record["server"], "parent_pid": record["parent_pid"], + "config_home": str(home), "stderr_path": str(expected), "destination": "file"} + + +def _mcp_child_line(line: str, record: dict, filters: dict) -> Optional[str]: + control = _mcp_control_record(line) + if control is not None and control.get("attempt_id") == record["attempt_id"]: + return None + owner = " ".join(f"{key}={json.dumps(value)}" for key, value in ( + ("profile", record["config_home"]), ("server", record["server"]), + ("attempt", record["attempt_id"]), + )) + rendered = f"[mcp {owner}] {line}" + session = filters.get("session_filter") + if session is not None and session not in rendered: + return None + if isinstance(line, _MCPPreviewNotice): + return rendered + if not _matches_filters(line, **{**filters, "session_filter": None}): + return None + return rendered + + +def _read_mcp_tail(path: Path, num_lines: int, *, has_filters: bool = False, **filters) -> list: + """Read actual stderr in attempt order; legacy rows keep their own format.""" + if num_lines <= 0: + return [] + requested = num_lines + num_lines = min(num_lines, _MCP_MAX_PREVIEW_LINES) + result = [] + remaining = _MCP_TAIL_BYTES + exhausted = False + seen = set() + + def append(line): + nonlocal remaining, exhausted + if len(line) > remaining: + exhausted = True + return + result.append(line) + remaining -= len(line) + + for line in reversed(_mcp_tail_rows(path, max(num_lines * 20, _MCP_INDEX_LINES), index=True)): + control = _mcp_control_record(line) + if control is None: + if isinstance(line, _MCPPreviewNotice) or _matches_filters(line, **filters): + append(line) + else: + record = _mcp_attempt(control, path) + if record is None or record["attempt_id"] in seen or len(seen) >= _MCP_MAX_ATTEMPTS: + continue + seen.add(record["attempt_id"]) + rows = _mcp_tail_rows(Path(record["stderr_path"]), + max(num_lines * 20, _MCP_INDEX_LINES) if has_filters else num_lines + 4) + for row in reversed(rows): + rendered = _mcp_child_line(row, record, filters) + if rendered is not None: + append(rendered) + if exhausted or len(result) >= num_lines: + break + if exhausted or len(result) >= num_lines: + break + lines = list(reversed(result)) + if exhausted or requested > _MCP_MAX_PREVIEW_LINES: + lines.insert(0, "MCP preview output limit reached; full content remains in the log files.\n") + return lines + + +def _mcp_cursor(path: Path): + try: + with _mcp_open_regular(path) as stream: + info = os.fstat(stream.fileno()) + if not stat.S_ISREG(info.st_mode): + return None + # Seed only the bounded final fragment. If it later completes, + # display its complete new record; never replay complete history. + start = max(0, info.st_size - _MCP_MAX_LINE_BYTES - 1) + stream.seek(start) + recent = stream.read(min(info.st_size, _MCP_MAX_LINE_BYTES + 1)) + pending = recent.rsplit(b"\n", 1)[-1] + dropping = len(pending) > _MCP_MAX_LINE_BYTES + return ((info.st_dev, info.st_ino), info.st_size, b"" if dropping else pending, dropping) + except OSError: + return None + + +def _mcp_read_chunk(path: Path, cursor, *, index: bool = False): + """Frame bounded binary chunks; retain fragments, never file descriptors.""" + try: + with _mcp_open_regular(path) as stream: + info = os.fstat(stream.fileno()) + if not stat.S_ISREG(info.st_mode): + return [], cursor + identity = (info.st_dev, info.st_ino) + retained = cursor is not None and cursor[0] == identity and cursor[1] <= info.st_size + offset, pending, dropping = cursor[1:] if retained else (0, b"", False) + stream.seek(offset) + chunk = stream.read(_MCP_FOLLOW_BYTES) + offset = stream.tell() + data = pending + chunk + lines = [] + if dropping: + _fragment, separator, data = data.partition(b"\n") + if not separator: + return [], (identity, offset, b"", True) + parts = data.split(b"\n") + for row in parts[:-1]: + lines.append(_mcp_oversized_line(index) if len(row) > _MCP_MAX_LINE_BYTES + else row.decode("utf-8", errors="replace") + "\n") + pending = parts[-1] + dropping = len(pending) > _MCP_MAX_LINE_BYTES + if dropping: + lines.append(_mcp_oversized_line(index)) + pending = b"" + return lines, (identity, offset, pending, dropping) + except OSError: + return [], cursor + + +def _follow_mcp_log(path: Path, **filters) -> None: + records = {} + + def remember(line): + control = _mcp_control_record(line) + record = _mcp_attempt(control, path) if control is not None else None + if record is not None: + records[record["attempt_id"]] = record + if len(records) > _MCP_MAX_ATTEMPTS: + records.pop(next(iter(records))) + print("MCP follow preview retains the latest 128 attempts; older captures remain in the log files.") + + for line in _mcp_tail_rows(path, _MCP_INDEX_LINES, index=True): + remember(line) + positions = {key: _mcp_cursor(Path(record["stderr_path"])) for key, record in records.items()} + for key, cursor in positions.items(): + if cursor is not None and cursor[3]: + rendered = _mcp_child_line(_mcp_oversized_line(False), records[key], filters) + if rendered is not None: + print(rendered, end="") + index_cursor = _mcp_cursor(path) + if index_cursor is not None and index_cursor[3]: + print(_mcp_oversized_line(True), end="") + while True: + lines, index_cursor = _mcp_read_chunk(path, index_cursor, index=True) + for line in lines: + if _mcp_control_record(line) is not None: + remember(line) + elif isinstance(line, _MCPPreviewNotice) or _matches_filters(line, **filters): + print(line, end="") + sys.stdout.flush() + for key, record in records.items(): + if _mcp_attempt(record, path) is None: + continue + lines, positions[key] = _mcp_read_chunk(Path(record["stderr_path"]), positions.get(key)) + for line in lines: + rendered = _mcp_child_line(line, record, filters) + if rendered is not None: + print(rendered, end="") + sys.stdout.flush() + positions = {key: positions[key] for key in records if key in positions} + time.sleep(0.3) + + def _read_tail( path: Path, num_lines: int, diff --git a/nix/checks.nix b/nix/checks.nix index 227f57d5e51e3..1e778b969830c 100644 --- a/nix/checks.nix +++ b/nix/checks.nix @@ -1160,7 +1160,7 @@ json.dump(sorted(leaf_paths(DEFAULT_CONFIG)), sys.stdout, indent=2) # Case-insensitive: the message names the managing system as the # identifier it is keyed by, and the display form is not the # property under test here. - echo "$OUTPUT" | grep -qi "managed by nixos" || (echo "FAIL: $label not guarded"; echo "$OUTPUT"; exit 1) + grep -qi "managed by nixos" <<< "$OUTPUT" || (echo "FAIL: $label not guarded"; echo "$OUTPUT"; exit 1) echo "PASS: $label blocked in managed mode" } diff --git a/package-lock.json b/package-lock.json index e9bb2269813fa..6f180ce34dd50 100644 --- a/package-lock.json +++ b/package-lock.json @@ -238,7 +238,7 @@ "bippy": "0.5.43", "concurrently": "10.0.4", "cross-env": "10.1.0", - "electron": "41.10.3", + "electron": "41.10.6", "electron-builder": "^26.8.1", "esbuild": "^0.28.1", "eslint": "^9.39.4", @@ -11136,9 +11136,9 @@ } }, "node_modules/electron": { - "version": "41.10.3", - "resolved": "https://registry.npmjs.org/electron/-/electron-41.10.3.tgz", - "integrity": "sha512-MJuSODPw8siv/I8JjhctW/cS/XNldwI4gLRyyWZx6QkoZJUDgbEvitp7IVOnGrHENTQb6Udo+zMpKhFnhlIhdg==", + "version": "41.10.6", + "resolved": "https://registry.npmjs.org/electron/-/electron-41.10.6.tgz", + "integrity": "sha512-NEB7QnVpT2p8S9U5jOCPAtCyhl1WjRT/ODvJP6ec1/dZRUmf/EjYVOBb/dthCmJh4bISqpAroGJ/xRRxX7VwJw==", "dev": true, "hasInstallScript": true, "license": "MIT", diff --git a/package.json b/package.json index 52040eb809a82..b4af6417c3899 100644 --- a/package.json +++ b/package.json @@ -71,7 +71,7 @@ "esbuild@0.28.1": true, "esbuild@0.28.2": true, "node-pty@1.1.0": true, - "electron@41.10.3": true, + "electron@41.10.6": true, "fsevents@2.3.2": true, "fsevents@2.3.3": true } diff --git a/plugins/context_engine/_context_governor/__init__.py b/plugins/context_engine/_context_governor/__init__.py index 9c155ca3f1aae..592ddc2ffcc09 100644 --- a/plugins/context_engine/_context_governor/__init__.py +++ b/plugins/context_engine/_context_governor/__init__.py @@ -205,6 +205,7 @@ def __init__( self.last_real_prompt_tokens = 0 self.last_compression_rough_tokens = 0 self.awaiting_real_usage_after_compression = False + self._pending_request_rough_tokens = 0 # Anti-thrashing state self._ineffective_compression_count = 0 @@ -351,6 +352,8 @@ def __deepcopy__(self, memo): clone._ineffective_compression_count = self._ineffective_compression_count clone._last_compression_savings_pct = self._last_compression_savings_pct clone._set_defer_baseline(self.last_rough_tokens_when_real_prompt_fit) + # A clone has not sent the original owner's pending request. + clone._pending_request_rough_tokens = 0 clone._previous_summary = self._previous_summary clone._summary_mode = self._summary_mode clone._summary_model = self._summary_model @@ -622,6 +625,19 @@ def update_model( protect_first_n: int | None = None, protect_last_n: int | None = None, ) -> None: + resolved_context_length = int(context_length or 0) + output_reserve = int(max_tokens) if max_tokens and int(max_tokens) > 0 else None + effective_window = resolved_context_length - (output_reserve or 0) + if resolved_context_length > 0 and effective_window <= 0: + raise ValueError( + "context-governor response reservation leaves no input budget " + f"(context_length={resolved_context_length}, max_tokens={output_reserve})" + ) + pressure_fields = ( + "model", "base_url", "api_key", "provider", "api_mode", + "context_length", "max_tokens", "threshold_tokens", + ) + previous_pressure_route = tuple(getattr(self, key, None) for key in pressure_fields) # Persist the active agent route. The optional summary-specific fields # override these only when explicitly configured. self.model = str(model or "") @@ -635,17 +651,31 @@ def update_model( self.protect_first_n = int(protect_first_n) if protect_last_n is not None: self.protect_last_n = int(protect_last_n) - self.context_length = int(context_length or 0) - self.max_tokens = ( - int(max_tokens) if max_tokens and int(max_tokens) > 0 else None - ) + self.context_length = resolved_context_length + self.max_tokens = output_reserve # Account for output reservation in effective input budget - effective_window = self.context_length - (self.max_tokens or 0) - if effective_window <= 0: - effective_window = self.context_length self.threshold_tokens = ( - int(effective_window * self.threshold_percent) if effective_window else 0 + int(effective_window * self.threshold_percent) if effective_window > 0 else 0 ) + if previous_pressure_route != tuple(getattr(self, key, None) for key in pressure_fields): + # Token evidence from another route/budget cannot qualify this + # request. Keep telemetry, but invalidate pressure projections. + self._set_defer_baseline(0) + self._pending_request_rough_tokens = 0 + self.last_compression_rough_tokens = 0 + self.awaiting_real_usage_after_compression = False + + def note_request_rough_estimate(self, rough_tokens: int) -> None: + """Record host request pressure for its next matching usage update. + + This is an advisory rough estimate, not final wire/native accounting. + A rebuilt request replaces the pending estimate; only its positive + provider usage can establish a new fitting pair. + """ + try: + self._pending_request_rough_tokens = max(0, int(rough_tokens)) + except (TypeError, ValueError, OverflowError): + self._pending_request_rough_tokens = 0 def update_from_response(self, usage: Dict[str, Any]) -> None: self.last_prompt_tokens = int( @@ -658,19 +688,19 @@ def update_from_response(self, usage: Dict[str, Any]) -> None: usage.get("total_tokens") or (self.last_prompt_tokens + self.last_completion_tokens) ) + request_rough = self._pending_request_rough_tokens + self._pending_request_rough_tokens = 0 # Mirror the built-in contract: last_real_prompt_tokens tracks the # most recent non-zero provider-reported prompt count, separate from # last_prompt_tokens (which can be -1 after a deferred preflight). if self.last_prompt_tokens > 0: self.last_real_prompt_tokens = self.last_prompt_tokens - if self.last_prompt_tokens < self.threshold_tokens: - if ( - self.awaiting_real_usage_after_compression - and self.last_compression_rough_tokens > 0 - ): - self._set_defer_baseline(self.last_compression_rough_tokens) - else: - self._set_defer_baseline(0) + # Update the pair atomically. An unmatched positive reading must + # not be projected against an unrelated older rough anchor; + # compaction diagnostics alone are not a request measurement. + self._set_defer_baseline( + request_rough if self.last_prompt_tokens < self.threshold_tokens else 0 + ) self.awaiting_real_usage_after_compression = False def should_compress(self, prompt_tokens: int = None) -> bool: @@ -727,8 +757,9 @@ def should_defer_preflight_to_real_usage(self, rough_tokens: int) -> bool: Mirrors the built-in ContextCompressor contract. The rough preflight estimator includes tool/schema overhead and can overestimate immediately - after compaction; provider-reported real usage is a better signal only - for a bounded growth window. Once rough growth exceeds tolerance, + after compaction; a matching rough/real pair can project cumulative + growth below the trigger. The projection is advisory, not a qualified + tokenizer bound. Once projected growth reaches the trigger, preflight must compress again instead of letting the session creep toward the hard context limit. """ @@ -741,9 +772,9 @@ def should_defer_preflight_to_real_usage(self, rough_tokens: int) -> bool: # At emergency pressure, fall through so should_compress() forces a # safety retry before the provider sees the oversized request. emergency_threshold = self._emergency_pressure_threshold() - if self.awaiting_real_usage_after_compression and ( - emergency_threshold <= 0 or rough_tokens < emergency_threshold - ): + if emergency_threshold > 0 and rough_tokens >= emergency_threshold: + return False + if self.awaiting_real_usage_after_compression: return True # Futility deferral is a normal-band optimization, not a safety gate. # At emergency pressure the caller must reach should_compress(), whose @@ -751,26 +782,19 @@ def should_defer_preflight_to_real_usage(self, rough_tokens: int) -> bool: # the request. Keeping this unconditional used to make the governor # look completely dead after two low-yield passes: both automatic # preflight paths returned here before should_compress() was evaluated. - if self._ineffective_compression_count >= 2 and ( - emergency_threshold <= 0 or rough_tokens < emergency_threshold - ): + if self._ineffective_compression_count >= 2: return True if self.last_real_prompt_tokens <= 0: return False if self.last_real_prompt_tokens >= self.threshold_tokens: return False - baseline = ( - self.last_rough_tokens_when_real_prompt_fit - or self.last_compression_rough_tokens - ) + baseline = self.last_rough_tokens_when_real_prompt_fit if baseline <= 0: return False growth = max(0, rough_tokens - baseline) - tolerated = max(4096, int(self.threshold_tokens * 0.05)) - if growth > tolerated: - return False - self._set_defer_baseline(max(baseline, rough_tokens)) - return True + # Only a matching positive usage update refreshes the pair. Repeated + # pressure checks or usage-less growth cannot ratchet this anchor. + return self.last_real_prompt_tokens + growth < self.threshold_tokens def has_content_to_compress(self, messages: List[Dict[str, Any]]) -> bool: non_system = [ @@ -904,6 +928,7 @@ def on_session_reset(self) -> None: self.last_real_prompt_tokens = 0 self.last_compression_rough_tokens = 0 self.awaiting_real_usage_after_compression = False + self._pending_request_rough_tokens = 0 self._ineffective_compression_count = 0 self._last_compression_savings_pct = 100.0 self._set_defer_baseline(0) @@ -3823,17 +3848,30 @@ def _request_min_net_savings_tokens(self) -> int | None: return None def _target_tokens(self, current_tokens: int | None) -> int: + target = None explicit = self._policy.get("token_budget") try: if explicit is not None and int(explicit) > 0: - return max(512, int(explicit)) + target = max(512, int(explicit)) except (TypeError, ValueError): pass - if self.context_length: - return max(512, int(self.context_length * 0.20)) - if current_tokens: - return max(512, int(current_tokens * 0.20)) - return 8000 + if target is None: + if self.context_length: + target = max(512, int(self.context_length * 0.20)) + elif current_tokens: + target = max(512, int(current_tokens * 0.20)) + else: + target = 8000 + if self.context_length > 0: + input_window = self.context_length - (self.max_tokens or 0) + if input_window <= 0: + raise ValueError("context-governor response reservation leaves no input budget") + # Configured policy remains the requested target. The active route + # is a hard ceiling, including when its input window is below the + # adapter's usual 512-token target floor. Rust still owns admission + # and may refuse a target that cannot preserve protected content. + target = min(target, input_window) + return target def _run_json( self, diff --git a/plugins/platforms/email/adapter.py b/plugins/platforms/email/adapter.py index 89ead8a82a610..62253ae66b069 100644 --- a/plugins/platforms/email/adapter.py +++ b/plugins/platforms/email/adapter.py @@ -996,7 +996,7 @@ def _allow_all_senders() -> bool: truthy = {"true", "1", "yes"} return ( _get_secret("EMAIL_ALLOW_ALL_USERS", "").strip().lower() in truthy - or os.getenv("GATEWAY_ALLOW_ALL_USERS", "").strip().lower() in truthy + or _get_secret("GATEWAY_ALLOW_ALL_USERS", "").strip().lower() in truthy ) @staticmethod @@ -1011,7 +1011,7 @@ def _allowlist_in_effect() -> bool: """ return bool( _get_secret("EMAIL_ALLOWED_USERS", "").strip() - or os.getenv("GATEWAY_ALLOWED_USERS", "").strip() + or _get_secret("GATEWAY_ALLOWED_USERS", "").strip() ) async def _dispatch_message(self, msg_data: Dict[str, Any]) -> None: @@ -1035,7 +1035,7 @@ async def _dispatch_message(self, msg_data: Dict[str, Any]) -> None: allowed_raw = _get_secret("EMAIL_ALLOWED_USERS", "").strip() if not allowed_raw: if _get_secret("EMAIL_ALLOW_ALL_USERS", "").strip().lower() not in {"true", "1", "yes"} and ( - os.getenv("GATEWAY_ALLOW_ALL_USERS", "").strip().lower() not in {"true", "1", "yes"} + _get_secret("GATEWAY_ALLOW_ALL_USERS", "").strip().lower() not in {"true", "1", "yes"} ): logger.debug( "[Email] Dropping sender at dispatch — EMAIL_ALLOWED_USERS is unset " diff --git a/plugins/platforms/feishu/adapter.py b/plugins/platforms/feishu/adapter.py index 42a70e79fb5ec..99ab0ad0729ba 100644 --- a/plugins/platforms/feishu/adapter.py +++ b/plugins/platforms/feishu/adapter.py @@ -1589,9 +1589,12 @@ def _load_settings(extra: Dict[str, Any]) -> FeishuAdapterSettings: # Default group policy (for groups not in group_rules) default_group_policy = str(extra.get("default_group_policy", "")).strip().lower() - # Env-only so adapter and gateway auth bypass share one source; yaml - # feishu.allow_bots is bridged to this env var at config load. - allow_bots = os.getenv("FEISHU_ALLOW_BOTS", "none").strip().lower() + # Target env takes precedence over target PlatformConfig.extra. The + # YAML hook seeds extra without mutating the process environment. + allow_bots = str( + _get_scoped_secret("FEISHU_ALLOW_BOTS", "") + or extra.get("allow_bots", "none") + ).strip().lower() if allow_bots not in {"none", "mentions", "all"}: logger.warning( "[Feishu] Unknown allow_bots=%r, falling back to 'none'. Valid: none, mentions, all.", @@ -1600,25 +1603,25 @@ def _load_settings(extra: Dict[str, Any]) -> FeishuAdapterSettings: allow_bots = "none" return FeishuAdapterSettings( - app_id=str(extra.get("app_id") or os.getenv("FEISHU_APP_ID", "")).strip(), + app_id=str(extra.get("app_id") or _get_scoped_secret("FEISHU_APP_ID", "")).strip(), app_secret=str(extra.get("app_secret") or _get_scoped_secret("FEISHU_APP_SECRET", "")).strip(), - domain_name=str(extra.get("domain") or os.getenv("FEISHU_DOMAIN", "feishu")).strip().lower(), + domain_name=str(extra.get("domain") or _get_scoped_secret("FEISHU_DOMAIN", "feishu")).strip().lower(), connection_mode=str( - extra.get("connection_mode") or os.getenv("FEISHU_CONNECTION_MODE", "websocket") + extra.get("connection_mode") or _get_scoped_secret("FEISHU_CONNECTION_MODE", "websocket") ).strip().lower(), encrypt_key=str(extra.get("encrypt_key") or _get_scoped_secret("FEISHU_ENCRYPT_KEY", "")).strip(), verification_token=str( extra.get("verification_token") or _get_scoped_secret("FEISHU_VERIFICATION_TOKEN", "") ).strip(), - group_policy=os.getenv("FEISHU_GROUP_POLICY", "allowlist").strip().lower(), + group_policy=_get_scoped_secret("FEISHU_GROUP_POLICY", "allowlist").strip().lower(), allowed_group_users=frozenset( item.strip() - for item in os.getenv("FEISHU_ALLOWED_USERS", "").split(",") + for item in _get_scoped_secret("FEISHU_ALLOWED_USERS", "").split(",") if item.strip() ), - bot_open_id=os.getenv("FEISHU_BOT_OPEN_ID", "").strip(), - bot_user_id=os.getenv("FEISHU_BOT_USER_ID", "").strip(), - bot_name=os.getenv("FEISHU_BOT_NAME", "").strip(), + bot_open_id=_get_scoped_secret("FEISHU_BOT_OPEN_ID", "").strip(), + bot_user_id=_get_scoped_secret("FEISHU_BOT_USER_ID", "").strip(), + bot_name=_get_scoped_secret("FEISHU_BOT_NAME", "").strip(), dedup_cache_size=max( 32, env_int("HERMES_FEISHU_DEDUP_CACHE_SIZE", _DEFAULT_DEDUP_CACHE_SIZE), @@ -1641,13 +1644,13 @@ def _load_settings(extra: Dict[str, Any]) -> FeishuAdapterSettings: "HERMES_FEISHU_MEDIA_BATCH_DELAY_SECONDS", _DEFAULT_MEDIA_BATCH_DELAY_SECONDS ), webhook_host=str( - extra.get("webhook_host") or os.getenv("FEISHU_WEBHOOK_HOST", _DEFAULT_WEBHOOK_HOST) + extra.get("webhook_host") or _get_scoped_secret("FEISHU_WEBHOOK_HOST", _DEFAULT_WEBHOOK_HOST) ).strip(), webhook_port=int( - extra.get("webhook_port") or os.getenv("FEISHU_WEBHOOK_PORT", str(_DEFAULT_WEBHOOK_PORT)) + extra.get("webhook_port") or _get_scoped_secret("FEISHU_WEBHOOK_PORT", str(_DEFAULT_WEBHOOK_PORT)) ), webhook_path=( - str(extra.get("webhook_path") or os.getenv("FEISHU_WEBHOOK_PATH", _DEFAULT_WEBHOOK_PATH)).strip() + str(extra.get("webhook_path") or _get_scoped_secret("FEISHU_WEBHOOK_PATH", _DEFAULT_WEBHOOK_PATH)).strip() or _DEFAULT_WEBHOOK_PATH ), ws_reconnect_nonce=_coerce_required_int(extra.get("ws_reconnect_nonce"), default=30, min_value=0), @@ -1659,7 +1662,7 @@ def _load_settings(extra: Dict[str, Any]) -> FeishuAdapterSettings: group_rules=group_rules, allow_bots=allow_bots, require_mention=_to_boolean( - extra.get("require_mention", os.getenv("FEISHU_REQUIRE_MENTION", "true")) + extra.get("require_mention", _get_scoped_secret("FEISHU_REQUIRE_MENTION", "true")) ), ) @@ -4373,9 +4376,9 @@ def _admit(self, sender: Any, message: Any) -> Optional[RejectReason]: return "bot_not_mentioned" if not is_group: - if os.getenv("FEISHU_ALLOW_ALL_USERS", "").strip().lower() in {"true", "1", "yes"}: + if _get_scoped_secret("FEISHU_ALLOW_ALL_USERS", "").strip().lower() in {"true", "1", "yes"}: return None - if os.getenv("GATEWAY_ALLOW_ALL_USERS", "").strip().lower() in {"true", "1", "yes"}: + if _get_scoped_secret("GATEWAY_ALLOW_ALL_USERS", "").strip().lower() in {"true", "1", "yes"}: return None # Empty FEISHU_ALLOWED_USERS is the pairing-mode default from setup: # forward DMs to gateway intake so the pairing handshake can run. @@ -5851,14 +5854,9 @@ def interactive_setup() -> None: def _apply_yaml_config(yaml_cfg: dict, feishu_cfg: dict) -> dict | None: - """Translate config.yaml feishu: keys into FEISHU_* env vars. - - Implements the apply_yaml_config_fn contract (#24849). Mirrors the legacy - feishu_cfg block from gateway/config.py::load_gateway_config() (allow_bots). - Env vars take precedence over YAML. Returns None — flows through env. - """ - if "allow_bots" in feishu_cfg and not os.getenv("FEISHU_ALLOW_BOTS"): - os.environ["FEISHU_ALLOW_BOTS"] = str(feishu_cfg["allow_bots"]).lower() + """Seed target PlatformConfig.extra; scoped env retains precedence.""" + if "allow_bots" in feishu_cfg: + return {"allow_bots": str(feishu_cfg["allow_bots"]).lower()} return None diff --git a/scripts/ci/classify_changes.py b/scripts/ci/classify_changes.py index 7f608d0af9cd2..dcd999e45a821 100644 --- a/scripts/ci/classify_changes.py +++ b/scripts/ci/classify_changes.py @@ -29,6 +29,8 @@ lives under ``apps/``, so without this lane a Rust change matched ``frontend`` and only the TypeScript matrix ran. * ``mcp_catalog`` — bundled MCP catalog / installer review. +* ``context_continuity`` — exact paired native owner and Governor qualification. +* ``current_owner_integration`` — canonical owner to Ares consumer qualification. Docker is not a lane — it builds on push-to-main and release only, never per-PR. @@ -56,6 +58,7 @@ from __future__ import annotations import json +from fnmatch import fnmatchcase import os import sys @@ -116,6 +119,134 @@ _RUST_PATHS = ("apps/bootstrap-installer/src-tauri/",) _RUST_FILENAMES = {"Cargo.toml", "Cargo.lock"} +# Native qualification applicability belongs to this classifier. These positive +# patterns preserve the former standalone PR trigger coverage. fnmatchcase may +# conservatively include deeper paths; it must never narrow a required lane. +_CONTEXT_CONTINUITY_PATHS = ( + '.github/workflows/context-continuity-qualification.yml', + 'ares_runtime/continuity/**', + 'tests/ares_runtime/test_continuity*', + 'tests/ares_runtime/test_context_rebase_state.py', + 'tests/ares_runtime/test_context_controller_credentials.py', + 'tests/ares_runtime/test_context_authority_binding.py', + 'tests/ares_runtime/test_context_native*.py', + 'hermes_state_context_authority.py', + 'hermes_cli/context_authority.py', + 'tests/ares_runtime/test_managed_calls.py', + 'tests/hermes_cli/test_goals.py', + 'tests/hermes_cli/test_goal_lifecycle_contract.py', + 'tests/test_model_tools.py', + 'tests/tools/test_code_execution.py', + 'tests/run_agent/test_message_sequence_repair.py', + 'tests/run_agent/test_tool_executor_contextvar_propagation.py', + 'tests/hermes_cli/test_heartbeat.py', + 'tests/hermes_cli/test_loops.py', + 'tests/test_run_checkpoint_import_boundaries.py', + 'tests/test_run_checkpoint*.py', + 'tests/test_run_task_custody.py', + 'tests/test_run_custody_cold_recovery.py', + 'tests/test_run_custody_ready_recovery.py', + 'tests/test_run_custody_obligations.py', + 'tests/gateway/test_context_input.py', + 'tests/gateway/test_context_input_recovery.py', + 'tests/ares_runtime/test_context_input_lifetime.py', + 'tests/cli/test_quick_commands.py', + 'tests/test_estop.py', + 'tests/ares_runtime/test_permit_readback.py', + 'tests/test_ares_collaboration.py', + 'ares_runtime/collaboration.py', + 'tests/gateway/test_pre_gateway_dispatch.py', + 'tests/gateway/test_restart_resume_pending.py', + 'tests/gateway/test_multiplex_session_db_profile_scope.py', + 'tests/gateway/test_multiplex_adapter_registry.py', + 'tests/gateway/test_adapter_startup_secret_scope.py', + 'tests/gateway/test_startup_connect_parallel.py', + 'tests/gateway/test_run_progress_topics.py', + 'tests/gateway/test_42039_duplicate_user_message.py', + 'tests/gateway/test_internal_event_never_interrupts_busy_session.py', + 'tests/gateway/test_multiplex_busy_input_mode.py', + 'tests/gateway/test_platform_base.py', + 'tests/gateway/test_base_topic_sessions.py', + 'tests/gateway/test_profile_routing.py', + 'tests/gateway/test_busy_session_auth_bypass.py', + 'tests/gateway/test_busy_session_ack.py', + 'gateway/context_input.py', + 'gateway/context_input_recovery.py', + 'gateway/turn_context.py', + 'gateway/platforms/base.py', + 'gateway/run.py', + 'tests/state/test_message_copy_mapping.py', + 'tests/state/test_message_row_publication.py', + 'tests/test_compression_watermark_commit.py', + 'tests/state/test_todo_compaction.py', + 'tests/run_agent/test_in_place_compaction.py', + 'tests/agent/test_micro_compaction.py', + 'tests/hermes_state/test_append_messages_batch.py', + 'tests/tui_gateway/test_run_checkpoint_claim_rpc.py', + 'tests/tui_gateway/test_goal_command.py', + 'tests/test_turn_run_custody.py', + 'tests/tui_gateway/test_inline_rpc_gil_starvation.py', + 'tests/tui_gateway/test_kanban_notify_poller.py', + 'tests/test_tui_gateway_server.py', + 'tests/cli/test_cli_goal_interrupt.py', + 'tests/cli/test_cli_async_delegation_delivery.py', + 'tests/gateway/test_goal_resume_restart.py', + 'tests/agent/test_synthetic_turn_display_kind.py', + 'tests/agent/test_turn_context.py', + 'tests/run_agent/test_run_agent.py', + 'tests/run_agent/test_1630_context_overflow_loop.py', + 'tests/run_agent/test_compression_lock_defer.py', + 'agent/agent_init.py', + 'agent/agent_runtime_helpers.py', + 'agent/tool_executor.py', + 'model_tools.py', + 'agent/conversation_loop.py', + 'agent/chat_completion_helpers.py', + 'agent/codex_runtime.py', + 'agent/run_checkpoint_custody.py', + 'agent/turn_context.py', + 'agent/context_input.py', + 'cli.py', + 'hermes_cli/goals.py', + 'hermes_cli/heartbeat.py', + 'hermes_cli/loops.py', + 'hermes_state.py', + 'hermes_state_common.py', + 'hermes_state_continuity.py', + 'hermes_state_inbox.py', + 'hermes_state_input_turns.py', + 'agent/turn_finalizer.py', + 'hermes_state_runs.py', + 'plugins/context_engine/_context_governor/**', + 'scripts/run_checkpoint_context.py', + 'scripts/run_checkpoint_claim.py', + 'scripts/run_checkpoint_resume.py', + 'tui_gateway/run_checkpoint_rpc.py', + 'scripts/run_tests.sh', + 'tests/run_agent/test_streaming.py', + 'tests/run_agent/test_run_agent_codex_responses.py', + 'tests/run_agent/test_codex_sdk_transform_bypass.py', + 'tests/plugins/test_context_governor*.py', + 'tui_gateway/server.py', + 'tui_gateway/methods_prompt.py', + 'tui_gateway/compute_host.py', + 'docs/context-continuity/**', +) +_CURRENT_OWNER_INTEGRATION_PATHS = ( + '.github/workflows/current-owner-integration.yml', + 'ares_runtime/collaboration.py', + 'ares_runtime/governed_context.py', + 'ares_runtime/__init__.py', + 'tests/owner_integration/**', + 'tests/test_ares_collaboration.py', + 'tests/ares_runtime/test_governed_context_materialization.py', + 'tests/ares_runtime/test_memory_witness_v2.py', + 'tests/ares_runtime/test_policy_basis_v2.py', + 'tests/ares_runtime/fixtures/profile_runtime_v2_owner.json', +) +_QUALIFICATION_INFRA_PATHS = ("scripts/ci/", "tests/ci/") +_QUALIFICATION_INFRA_FILES = {"scripts/run_tests.sh", "scripts/run_tests_parallel.py"} + def _is_docs(p: str) -> bool: if p.startswith(("skills/", "optional-skills/")): return False @@ -183,6 +314,10 @@ def ci_review_files(files: list[str]) -> list[str]: return sorted({f.strip() for f in files if f.strip() and _is_ci_review(f.strip())}) +def _qualification_path(p: str, patterns: tuple[str, ...]) -> bool: + return any(fnmatchcase(p, pattern) for pattern in patterns) + + def classify(files: list[str]) -> dict[str, bool]: """Map changed paths to ``{lane: should_run}``.""" files = [f.strip() for f in files if f.strip()] @@ -208,6 +343,10 @@ def classify(files: list[str]) -> dict[str, bool]: "rust": any(_is_rust(f) for f in files), "mcp_catalog": any(_is_mcp_catalog(f) for f in files), "ci_review": any(_is_ci_review(f) for f in files), + "context_continuity": any(_qualification_path(f, _CONTEXT_CONTINUITY_PATHS) for f in files), + "current_owner_integration": any( + _qualification_path(f, _CURRENT_OWNER_INTEGRATION_PATHS) for f in files + ), "nix": python_prod or frontend or any(_is_nix(f) for f in files) } if not files or any(f.startswith(".github/") for f in files): @@ -225,8 +364,13 @@ def classify(files: list[str]) -> dict[str, bool]: ret["rust"] = True ret["nix"] = True ret["ci_review"] = True + ret["context_continuity"] = True + ret["current_owner_integration"] = True # explicitly skip mcp catalog here. it's not needed unless those files are modified. + if any(f.startswith(_QUALIFICATION_INFRA_PATHS) or f in _QUALIFICATION_INFRA_FILES for f in files): + ret["context_continuity"] = True + ret["current_owner_integration"] = True return ret diff --git a/scripts/ci/evaluate_required_checks.py b/scripts/ci/evaluate_required_checks.py index 02c1eb1593826..8454e4f20c755 100644 --- a/scripts/ci/evaluate_required_checks.py +++ b/scripts/ci/evaluate_required_checks.py @@ -30,6 +30,8 @@ "mcp_catalog", "ci_review", "ci_review_files", + "context_continuity", + "current_owner_integration", ) CLASSIFIER_BOOLEAN_KEYS = frozenset(CLASSIFIER_KEYS) - {"ci_review_files"} @@ -45,6 +47,8 @@ "js-tests", "installer-tests", "rust-tests", + "context-continuity", + "current-owner-integration", "e2e-desktop", "docs-site", "history-check", @@ -63,6 +67,11 @@ # It is an explicit, reviewable exception, not a general skipped-is-green rule. DISABLED_JOBS = frozenset({"e2e-desktop"}) VALID_RESULTS = frozenset({"success", "failure", "cancelled", "skipped"}) +# Reusable-call success cannot substitute for each mandatory inner owner job. +QUALIFICATION_RESULTS = { + "context-continuity": ("native_external_owner_result", "focused_tests_result"), + "current-owner-integration": ("profile_runtime_consumer_result",), +} def _failure(message: str) -> dict[str, Any]: @@ -114,6 +123,19 @@ def _output_bool(needs: Mapping[str, Any], job: str, key: str) -> bool: return value == "true" +def _qualification_failures(needs: Mapping[str, Any], job: str) -> list[str]: + outputs = needs[job].get("outputs") + if not isinstance(outputs, Mapping): + return [f"qualification job {job} outputs must be an object"] + failures = [] + for key in QUALIFICATION_RESULTS[job]: + if key not in outputs: + failures.append(f"qualification job {job} output missing: {key}") + elif outputs[key] != "success": + failures.append(f"qualification job {job} output {key} must succeed, got {outputs[key]!r}") + return failures + + def _applicable(event_name: str, flags: Mapping[str, bool], needs: Mapping[str, Any], job: str) -> bool: if job in DISABLED_JOBS or job in {"detect", "infographic-check", "osv-scanner"}: return True @@ -125,6 +147,10 @@ def _applicable(event_name: str, flags: Mapping[str, bool], needs: Mapping[str, return flags["installer"] if job == "rust-tests": return flags["rust"] + if job == "context-continuity": + return event_name != "pull_request" or flags["context_continuity"] + if job == "current-owner-integration": + return event_name != "pull_request" or flags["current_owner_integration"] if job == "docs-site": return flags["site"] if job == "history-check": @@ -215,6 +241,8 @@ def evaluate( required_jobs.append(job) if result != "success": failures.append(f"required job {job} must succeed, got {result}") + elif job in QUALIFICATION_RESULTS: + failures.extend(_qualification_failures(needs_map, job)) elif result != "skipped": failures.append(f"non-applicable job {job} must be explicitly skipped, got {result}") diff --git a/scripts/npm_failure_diagnostics.py b/scripts/npm_failure_diagnostics.py index 661b7e2a4d5be..3cb71ca74cb6e 100644 --- a/scripts/npm_failure_diagnostics.py +++ b/scripts/npm_failure_diagnostics.py @@ -11,7 +11,7 @@ import tempfile PREFIX = "HERMES_NPM_DIAGNOSTIC " -SCHEMA = "npm-install-diagnostic/v1" +SCHEMA = "npm-install-diagnostic/v2" LIMIT = 256 * 1024 CAUSES = { "ERESOLVE": "Dependency resolution conflict", @@ -20,6 +20,7 @@ "EPERM": "Operation not permitted", "ENOSPC": "Insufficient disk space", "ENOENT": "Required file or executable missing", + "ENOTEMPTY": "Filesystem directory is not empty", "EINTEGRITY": "Package integrity verification failed", "E401": "Registry authentication required or rejected", "E403": "Registry access forbidden", @@ -34,11 +35,18 @@ "SELF_SIGNED_CERT_IN_CHAIN": "TLS certificate chain is not trusted", "DEPTH_ZERO_SELF_SIGNED_CERT": "TLS peer certificate is self-signed", "UNABLE_TO_VERIFY_LEAF_SIGNATURE": "TLS certificate issuer cannot be verified", + "UNABLE_TO_GET_ISSUER_CERT": "TLS certificate issuer unavailable", "UNABLE_TO_GET_ISSUER_CERT_LOCALLY": "TLS issuer unavailable in configured trust", "ERR_TLS_CERT_ALTNAME_INVALID": "TLS certificate hostname mismatch", "ELIFECYCLE": "Package lifecycle script failed", } CODE_LINE = re.compile(r"^npm (?:ERR!|error) code ([A-Z0-9_]+)\s*$", re.MULTILINE) +OUTPUT_DETAILS = { + "empty": "no non-whitespace output in bounded tail", + "unrecognized-code": "only unrecognized code lines in bounded tail", + "non-code": "non-code output in bounded tail", + "recognized-code": "Allowlisted npm error codes only", +} def bounded_tail(path: Path) -> tuple[str, bool]: @@ -53,7 +61,8 @@ def bounded_tail(path: Path) -> tuple[str, bool]: return stream.read(LIMIT).decode("utf-8", "replace"), info.st_size > LIMIT -def record(stage: str, exit_code: int, elapsed: int, codes: list[str], truncated: bool) -> dict: +def record(stage: str, exit_code: int, elapsed: int, codes: list[str], truncated: bool, + *, output_kind: str) -> dict: if stage not in ("root", "tui") or type(exit_code) is not int or not 1 <= exit_code <= 255: raise ValueError("invalid status") if type(elapsed) is not int or not 0 <= elapsed <= 86400 or type(truncated) is not bool: @@ -61,17 +70,26 @@ def record(stage: str, exit_code: int, elapsed: int, codes: list[str], truncated if not isinstance(codes, list) or any(not isinstance(code, str) or code not in CAUSES for code in codes): raise ValueError("invalid codes") codes = sorted(set(codes)) + if not isinstance(output_kind, str) or output_kind not in OUTPUT_DETAILS: + raise ValueError("invalid output kind") + if (output_kind == "recognized-code") != bool(codes): + raise ValueError("inconsistent output kind") return {"schema": SCHEMA, "stage": stage, "exit_code": exit_code, "elapsed_seconds": elapsed, "timeout_status": exit_code == 124, "npm_codes": codes, "safe_causes": [CAUSES[code] for code in codes], - "output_truncated": truncated, - "detail": "Raw npm output intentionally omitted; no recognized error code" if not codes else "Allowlisted npm error codes only"} + "output_truncated": truncated, "output_kind": output_kind, + "detail": ("Raw npm output intentionally omitted; no recognized error code; " + OUTPUT_DETAILS[output_kind]) + if not codes else OUTPUT_DETAILS[output_kind]} def summarize(path: Path, stage: str, exit_code: int, elapsed: int) -> dict: text, truncated = bounded_tail(path) - codes = [code for code in CODE_LINE.findall(text) if code in CAUSES] - return record(stage, exit_code, elapsed, codes, truncated) + code_lines = CODE_LINE.findall(text) + codes = [code for code in code_lines if code in CAUSES] + # Describe retained output shape, never a cause inferred from process status. + output_kind = ("recognized-code" if codes else "unrecognized-code" if code_lines + else "non-code" if text.strip() else "empty") + return record(stage, exit_code, elapsed, codes, truncated, output_kind=output_kind) def collect(path: Path) -> list[dict]: @@ -88,7 +106,7 @@ def collect(path: Path) -> list[dict]: if value.get("schema") != SCHEMA: continue rebuilt = record(value["stage"], value["exit_code"], value["elapsed_seconds"], - value["npm_codes"], value["output_truncated"]) + value["npm_codes"], value["output_truncated"], output_kind=value["output_kind"]) if value != rebuilt: continue records.append(rebuilt) diff --git a/tests/agent/lsp/test_lifecycle.py b/tests/agent/lsp/test_lifecycle.py index f38934082cfe6..5829024211eb1 100644 --- a/tests/agent/lsp/test_lifecycle.py +++ b/tests/agent/lsp/test_lifecycle.py @@ -1,4 +1,4 @@ -"""Tests for service-singleton lifecycle: atexit handler, idempotent shutdown. +"""Tests for profile service lifecycle: atexit handler, idempotent shutdown. These cover the exit-cleanup behavior added to plug the language-server process leak — without the atexit hook, ``hermes chat`` exits while @@ -15,18 +15,14 @@ @pytest.fixture(autouse=True) -def _reset_singleton(): - """Force a clean module state before each test. - - Tests in this file share process-global state (the lazy - singleton + atexit registration flag); reset both before and - after every test so order doesn't matter. - """ - lsp_module._service = None - lsp_module._atexit_registered = False +def _reset_services(monkeypatch): + """Isolate the production registry and drain only this test's fake services.""" + with lsp_module._service_lock: + monkeypatch.setattr(lsp_module, "_services", {}) + monkeypatch.setattr(lsp_module, "_atexit_registered", False) yield - lsp_module._service = None - lsp_module._atexit_registered = False + lsp_module._atexit_shutdown() + assert lsp_module._services == {} def test_get_service_registers_atexit_handler_once(monkeypatch): @@ -36,8 +32,21 @@ def test_get_service_registers_atexit_handler_once(monkeypatch): twice — harmless but wasteful).""" fake_svc = MagicMock() fake_svc.is_active.return_value = True + creations = [] + + def create(cls, *, config, profile_boundary): + from agent.secret_scope import ProfileEnvBoundary + from hermes_constants import get_hermes_home + + assert isinstance(config, dict) + assert isinstance(profile_boundary, ProfileEnvBoundary) + assert profile_boundary.target_home == get_hermes_home().resolve() + assert profile_boundary.target_generation + creations.append((config, profile_boundary)) + return fake_svc + monkeypatch.setattr( - lsp_module.LSPService, "create_from_config", classmethod(lambda cls: fake_svc) + lsp_module.LSPService, "create_from_config", classmethod(create) ) registrations = [] @@ -54,6 +63,8 @@ def fake_register(fn): assert a is fake_svc assert b is fake_svc assert c is fake_svc + assert len(creations) == 1 + assert len(lsp_module._services) == 1 assert len(registrations) == 1 # The registered callable must be our internal shutdown wrapper. assert registrations[0] is lsp_module._atexit_shutdown @@ -61,31 +72,78 @@ def fake_register(fn): -def test_atexit_shutdown_swallows_exceptions(monkeypatch): +def test_atexit_shutdown_swallows_exceptions_and_drains_other_services(tmp_path): + from hermes_constants import hermes_home_key + + failing = MagicMock() + sibling = MagicMock() + registry_empty_at_shutdown = [] + def boom(): + registry_empty_at_shutdown.append(lsp_module._services == {}) raise RuntimeError("server already dead") - monkeypatch.setattr(lsp_module, "shutdown_service", boom) - # Must not raise. + failing.shutdown.side_effect = boom + with lsp_module._service_lock: + lsp_module._services[hermes_home_key(tmp_path / "first")] = (("g1", "cfg"), failing) + lsp_module._services[hermes_home_key(tmp_path / "second")] = (("g2", "cfg"), sibling) lsp_module._atexit_shutdown() + failing.shutdown.assert_called_once_with() + sibling.shutdown.assert_called_once_with() + assert registry_empty_at_shutdown == [True] + assert lsp_module._services == {} + lsp_module._atexit_shutdown() + assert failing.shutdown.call_count == 1 + assert sibling.shutdown.call_count == 1 -def test_shutdown_service_idempotent(monkeypatch): +def test_shutdown_service_idempotent(monkeypatch, tmp_path): """Calling shutdown twice must be safe — first call cleans up, second call no-ops (nothing to shut down).""" fake_svc = MagicMock() fake_svc.is_active.return_value = True fake_svc.shutdown = MagicMock() + sibling = MagicMock() + sibling.is_active.return_value = True + creations = [] + + def create(cls, *, config, profile_boundary): + from agent.secret_scope import ProfileEnvBoundary + from hermes_constants import get_hermes_home + + assert isinstance(config, dict) + assert isinstance(profile_boundary, ProfileEnvBoundary) + assert profile_boundary.target_home == get_hermes_home().resolve() + assert profile_boundary.target_generation + creations.append(profile_boundary.identity) + return fake_svc if len(creations) == 1 else sibling + monkeypatch.setattr( - lsp_module.LSPService, "create_from_config", classmethod(lambda cls: fake_svc) + lsp_module.LSPService, "create_from_config", classmethod(create) ) monkeypatch.setattr(atexit, "register", lambda fn: None) - lsp_module.get_service() + from hermes_constants import get_hermes_home, hermes_home_key + + first_home = get_hermes_home() + assert lsp_module.get_service() is fake_svc + sibling_home = tmp_path / "sibling" + sibling_home.mkdir() + with monkeypatch.context() as sibling_context: + sibling_context.setenv("HERMES_HOME", str(sibling_home)) + assert lsp_module.get_service() is sibling + assert len(creations) == 2 lsp_module.shutdown_service() lsp_module.shutdown_service() # must not raise assert fake_svc.shutdown.call_count == 1 + assert hermes_home_key(first_home) not in lsp_module._services + assert lsp_module._services[hermes_home_key(sibling_home)][1] is sibling + sibling.shutdown.assert_not_called() + lsp_module.shutdown_service(profile_home=sibling_home) + lsp_module.shutdown_service(profile_home=sibling_home) + sibling.shutdown.assert_called_once_with() + assert lsp_module._services == {} diff --git a/tests/agent/test_chat_completion_helpers_ri.py b/tests/agent/test_chat_completion_helpers_ri.py index a69ce7c159ec3..b239a8be198b4 100644 --- a/tests/agent/test_chat_completion_helpers_ri.py +++ b/tests/agent/test_chat_completion_helpers_ri.py @@ -26,7 +26,7 @@ def test_dispatch_nonstreaming_uses_ri_when_enabled(): ) assert result is response - can_use_ri.assert_called_once_with(agent) + can_use_ri.assert_called_once_with(agent, api_kwargs) ri_call.assert_called_once_with(agent, api_kwargs) @@ -45,7 +45,7 @@ def test_dispatch_nonstreaming_reverts_to_openai_when_ri_disabled(): ) assert result is response - can_use_ri.assert_called_once_with(agent) + can_use_ri.assert_called_once_with(agent, api_kwargs) ri_call.assert_not_called() mock_client.chat.completions.create.assert_called_once_with(**api_kwargs) @@ -135,8 +135,8 @@ def test_interruptible_streaming_codex_path_not_intercepted_by_ri_gate(): assert agent._interruptible_api_call.call_args.args[0]["stream"] is True -def test_should_use_ri_pipeline_defaults_to_native_plus_config_path(monkeypatch): - """Use RiPipeline when native is available and provider filters allow it.""" +def test_should_use_ri_pipeline_default_incompatible_provider_retains_sdk(monkeypatch): + """Default native availability must preserve provider wire compatibility.""" agent = SimpleNamespace( provider="openrouter", _ri_pipeline_enabled=True, @@ -146,7 +146,7 @@ def test_should_use_ri_pipeline_defaults_to_native_plus_config_path(monkeypatch) with monkeypatch.context() as cm: cm.delenv("HERMES_RI_PIPELINE", raising=False) cm.delenv("HERMES_RI_PIPELINE_PROVIDERS", raising=False) - assert _should_use_ri_pipeline(agent) is True + assert _should_use_ri_pipeline(agent) is False def test_should_use_ri_pipeline_honors_agent_disable_config(monkeypatch): diff --git a/tests/agent/test_context_governor_pressure_contract.py b/tests/agent/test_context_governor_pressure_contract.py new file mode 100644 index 0000000000000..df94a32cd5078 --- /dev/null +++ b/tests/agent/test_context_governor_pressure_contract.py @@ -0,0 +1,160 @@ +"""Pressure evidence must come from a matching request/response, never a predicate.""" +import copy +import json +from unittest.mock import patch + +import pytest + +from plugins.context_engine._context_governor import ContextGovernorEngine + + +@pytest.fixture +def governor(tmp_path): + with patch("hermes_cli.config.load_config", return_value={}): + engine = ContextGovernorEngine( + binary=str(tmp_path / "no-native-governor"), + store_dir=str(tmp_path / "governor"), + ) + engine.update_model("test-model", 100_000, max_tokens=10_000, + provider="test", api_mode="chat_completions", + threshold_percent=0.5) + return engine + + +def seed_pair(engine, rough=46_000, real=20_000): + # The hostile audit's public-state seed, not a native receipt claim. + engine.last_real_prompt_tokens = real + engine.last_rough_tokens_when_real_prompt_fit = rough + + +def observe(engine, rough, real): + engine.note_request_rough_estimate(rough) + engine.update_from_response({"prompt_tokens": real}) + + +def test_gradual_missing_usage_does_not_ratchet(governor): + seed_pair(governor) + decisions = [] + for rough in range(49_000, 118_001, 3_000): + decisions.append(governor.should_defer_preflight_to_real_usage(rough)) + governor.update_from_response({}) + print(json.dumps({"rough_final": 118_000, "all_deferred": all(decisions), + "anchor": governor.last_rough_tokens_when_real_prompt_fit})) + assert not all(decisions) + assert governor.last_rough_tokens_when_real_prompt_fit == 46_000 + + +def test_repeated_predicate_is_pure(governor): + seed_pair(governor) + decisions = [governor.should_defer_preflight_to_real_usage(49_000) for _ in range(8)] + assert all(decisions) + assert governor.last_rough_tokens_when_real_prompt_fit == 46_000 + + +@pytest.mark.parametrize("awaiting,ineffective", [(False, 0), (True, 0), (False, 2), (True, 2)]) +@pytest.mark.parametrize("rough", [81_000, 90_000, 100_000]) +def test_every_deferral_branch_yields_at_emergency(governor, awaiting, ineffective, rough): + seed_pair(governor, rough=rough - 1000) + governor.awaiting_real_usage_after_compression = awaiting + governor._ineffective_compression_count = ineffective + assert not governor.should_defer_preflight_to_real_usage(rough) + assert governor.should_compress(rough) + + +def test_matching_fresh_usage_replaces_pair_and_not_predicate(governor): + observe(governor, 60_000, 15_000) + assert governor.last_rough_tokens_when_real_prompt_fit == 60_000 + assert governor.last_real_prompt_tokens == 15_000 + assert governor.should_defer_preflight_to_real_usage(70_000) + assert governor.last_rough_tokens_when_real_prompt_fit == 60_000 + observe(governor, 70_000, 25_000) + assert governor.last_rough_tokens_when_real_prompt_fit == 70_000 + assert governor.last_real_prompt_tokens == 25_000 + + +@pytest.mark.parametrize("usage", [{}, {"prompt_tokens": 0}]) +def test_missing_usage_preserves_old_pair_but_cannot_pair_a_later_unmatched_reading(governor, usage): + observe(governor, 46_000, 20_000) + governor.note_request_rough_estimate(50_000) + governor.update_from_response(usage) + assert governor.last_rough_tokens_when_real_prompt_fit == 46_000 + assert governor.last_real_prompt_tokens == 20_000 + governor.update_from_response({"prompt_tokens": 22_000}) + assert governor.last_rough_tokens_when_real_prompt_fit == 0 + assert not governor.should_defer_preflight_to_real_usage(52_000) + + +def test_compaction_diagnostic_does_not_mint_a_matching_pair(governor): + governor.last_compression_rough_tokens = 60_000 + governor.awaiting_real_usage_after_compression = True + governor.update_from_response({"prompt_tokens": 20_000}) + assert governor.last_rough_tokens_when_real_prompt_fit == 0 + assert not governor.should_defer_preflight_to_real_usage(62_000) + + +def test_latest_request_note_replaces_unsent_or_missing_usage_request(governor): + governor.note_request_rough_estimate(70_000) + governor.note_request_rough_estimate(50_000) + governor.update_from_response({"input_tokens": 20_000}) + assert governor.last_rough_tokens_when_real_prompt_fit == 50_000 + + +def test_known_noisy_estimate_retains_runway_until_cumulative_growth(governor): + observe(governor, 60_000, 30_000) + assert governor.should_defer_preflight_to_real_usage(72_000) + assert not governor.should_defer_preflight_to_real_usage(75_000) + assert governor.last_rough_tokens_when_real_prompt_fit == 60_000 + + +def test_one_post_compaction_measurement_opportunity_is_bounded(governor): + governor.awaiting_real_usage_after_compression = True + assert governor.should_defer_preflight_to_real_usage(50_000) + governor.update_from_response({}) + assert not governor.should_defer_preflight_to_real_usage(50_000) + + +def test_pressured_real_usage_clears_fit(governor): + observe(governor, 60_000, 45_000) + assert governor.last_rough_tokens_when_real_prompt_fit == 0 + assert not governor.should_defer_preflight_to_real_usage(60_000) + + +def test_clone_and_reset_do_not_transfer_pending_request(governor): + observe(governor, 60_000, 20_000) + governor.note_request_rough_estimate(70_000) + with patch("hermes_cli.config.load_config", return_value={}): + clone = copy.deepcopy(governor) + assert clone.last_rough_tokens_when_real_prompt_fit == 60_000 + clone.update_from_response({"prompt_tokens": 25_000}) + assert clone.last_rough_tokens_when_real_prompt_fit == 0 + assert governor.last_rough_tokens_when_real_prompt_fit == 60_000 + governor.on_session_reset() + governor.update_from_response({"prompt_tokens": 25_000}) + assert governor.last_rough_tokens_when_real_prompt_fit == 0 + + +@pytest.mark.parametrize("changed", [ + {"model": "other"}, {"provider": "other"}, {"api_mode": "codex_responses"}, + {"base_url": "https://other.invalid"}, {"context_length": 80_000}, + {"max_tokens": 20_000}, {"threshold_percent": 0.6}, +]) +def test_route_or_budget_rebinding_invalidates_old_pressure_evidence(governor, changed): + observe(governor, 60_000, 20_000) + governor.note_request_rough_estimate(70_000) + values = dict(model="test-model", context_length=100_000, max_tokens=10_000, + provider="test", api_mode="chat_completions", base_url="", + threshold_percent=0.5) + values.update(changed) + governor.update_model(**values) + assert governor.last_rough_tokens_when_real_prompt_fit == 0 + governor.update_from_response({"prompt_tokens": 20_000}) + assert governor.last_rough_tokens_when_real_prompt_fit == 0 + + +def test_same_route_update_preserves_a_matching_pending_request(governor): + governor.note_request_rough_estimate(60_000) + governor.update_model("test-model", 100_000, max_tokens=10_000, + provider="test", api_mode="chat_completions", + threshold_percent=0.5) + governor.update_from_response({"prompt_tokens": 20_000}) + assert governor.last_rough_tokens_when_real_prompt_fit == 60_000 diff --git a/tests/agent/test_context_governor_route_budget.py b/tests/agent/test_context_governor_route_budget.py new file mode 100644 index 0000000000000..4c98ee20fc619 --- /dev/null +++ b/tests/agent/test_context_governor_route_budget.py @@ -0,0 +1,87 @@ +"""Route admission limits and signature-compatible model budget handoffs.""" + +from unittest.mock import patch + +import pytest + +from agent.auxiliary_client import _update_compressor_model +from plugins.context_engine._context_governor import ContextGovernorEngine + + +@pytest.fixture +def governor(tmp_path): + with patch("hermes_cli.config.load_config", return_value={}): + return ContextGovernorEngine( + binary=str(tmp_path / "fake-governor"), + store_dir=str(tmp_path / "governor-store"), + ) + + +@pytest.mark.parametrize( + ("policy_budget", "window", "reserve", "expected"), + [ + (128_000, 64_000, 4096, 59_904), + (128_000, 64_000, None, 64_000), + (8000, 64_000, 4096, 8000), + (128_000, 256, 128, 128), + (None, 64_000, 4096, 12_800), + (None, 256, 128, 128), + (128_000, 0, 4096, 128_000), + (128_000, 1_050_000, 4096, 128_000), + (1_050_000, 1_050_000, None, 1_050_000), + (1_050_000, 1_048_576, 4096, 1_044_480), + ], +) +def test_policy_target_cannot_exceed_known_input_window( + governor, policy_budget, window, reserve, expected +): + governor._policy["token_budget"] = policy_budget + governor.update_model("test-model", window, max_tokens=reserve) + assert governor._target_tokens(100_000) == expected + if window: + assert expected <= window - (reserve or 0) + else: + assert governor.threshold_tokens == 0 + + +@pytest.mark.parametrize("reserve", [64_000, 64_001]) +def test_exhausted_input_budget_refuses_update_without_rebinding(governor, reserve): + governor.update_model("valid-model", 64_000, max_tokens=4096) + with pytest.raises(ValueError, match="no input budget"): + governor.update_model("invalid-model", 64_000, max_tokens=reserve) + assert governor.model == "valid-model" + assert governor.max_tokens == 4096 + assert governor.context_length == 64_000 + + +def test_handoff_forwards_and_clears_the_engine_response_reserve(governor): + governor._policy["token_budget"] = 128_000 + for reserve, target in [(4096, 59_904), (None, 64_000)]: + _update_compressor_model( + governor, + model="test-model", + context_length=64_000, + max_tokens=reserve, + threshold_percent=0.75, + ) + assert governor.max_tokens == reserve + assert governor._target_tokens(100_000) == target + + +def test_legacy_engine_receives_only_its_declared_update_api(): + calls = [] + + class LegacyEngine: + def update_model( + self, model, context_length, base_url="", api_key="", provider="", api_mode="" + ): + calls.append((model, context_length)) + + _update_compressor_model( + LegacyEngine(), + model="legacy-model", + context_length=64_000, + max_tokens=4096, + threshold_percent=0.75, + ) + assert calls == [("legacy-model", 64_000)] diff --git a/tests/agent/test_message_sanitization_policy.py b/tests/agent/test_message_sanitization_policy.py index 683d2d8df0f6e..8bae2f5b56f5c 100644 --- a/tests/agent/test_message_sanitization_policy.py +++ b/tests/agent/test_message_sanitization_policy.py @@ -455,7 +455,7 @@ def test_restore_primary_reverts_flag(self): agent._reasoning_echo_flag = False from agent.agent_runtime_helpers import restore_primary_runtime - restore_primary_runtime(agent) + assert restore_primary_runtime(agent) is True # Flag should be restored from snapshot assert agent._reasoning_echo_flag is True diff --git a/tests/agent/test_moa_switch_api_mode.py b/tests/agent/test_moa_switch_api_mode.py index 4d1b048bdc018..743461c13714f 100644 --- a/tests/agent/test_moa_switch_api_mode.py +++ b/tests/agent/test_moa_switch_api_mode.py @@ -16,23 +16,26 @@ from __future__ import annotations -import types +from unittest.mock import Mock import pytest -def _make_fake_agent(): - """A minimal stand-in carrying only the attributes switch_model touches.""" - agent = types.SimpleNamespace() +def _make_fake_agent(primary_api_mode): + """Use the real host methods with inert client and no context engine.""" + from run_agent import AIAgent + + agent = object.__new__(AIAgent) agent.model = "minimax-m3" agent.provider = "opencode-go" - agent.api_mode = "anthropic_messages" + agent.api_mode = primary_api_mode agent.api_key = "old-key" agent.base_url = "https://old.example/v1" agent.client = object() agent._client_kwargs = {"base_url": "https://old.example/v1"} agent._config_context_length = 123456 agent._transport_cache = {} + agent.context_compressor = None agent.quiet_mode = True # switch_model re-reads reasoning_echo for the incoming model as part of the # core field swap, before the moa branch runs. On a real AIAgent this is a @@ -40,6 +43,10 @@ def _make_fake_agent(): # this test asserts on. agent._reasoning_echo_flag = False agent._read_reasoning_echo_from_config = lambda: False + if primary_api_mode == "anthropic_messages": + agent._anthropic_api_key = "old-anthropic-key" + agent._anthropic_base_url = "https://old.example" + agent._is_anthropic_oauth = False return agent @@ -47,7 +54,8 @@ def _make_fake_agent(): "incoming_api_mode", ["codex_responses", "anthropic_messages", "chat_completions", ""], ) -def test_switch_to_moa_pins_chat_completions(monkeypatch, incoming_api_mode): +@pytest.mark.parametrize("primary_api_mode", ["chat_completions", "anthropic_messages"]) +def test_switch_to_moa_pins_chat_completions(monkeypatch, incoming_api_mode, primary_api_mode): """Switching to provider=moa must force api_mode=chat_completions. No matter what transport the resolver/aggregator implies for the preset, @@ -57,27 +65,23 @@ def test_switch_to_moa_pins_chat_completions(monkeypatch, incoming_api_mode): """ from agent import agent_runtime_helpers as arh - # Neutralize the post-swap machinery that needs a real AIAgent (credential - # pool reload, context-compressor refresh, primary-runtime bookkeeping). - # We only assert the api_mode invariant set in the moa client-build branch. - monkeypatch.setattr(arh, "load_pool", lambda *a, **k: None, raising=False) + # Keep profile/config and credential reads offline while exercising the + # real host helpers and lazy MoA facade through a complete successful swap. + monkeypatch.setattr("agent.credential_pool.load_pool", lambda *a, **k: None) + monkeypatch.setattr("hermes_cli.config.load_config", lambda: {}) + monkeypatch.setattr("hermes_cli.config.load_config_readonly", lambda: {}) - agent = _make_fake_agent() - try: - arh.switch_model( - agent, - new_model="frontier", - new_provider="moa", - api_key="moa-virtual-provider", - base_url="moa://local", - api_mode=incoming_api_mode, - ) - except Exception: - # switch_model does post-swap work (compressor, pool, runtime) that may - # raise against a fake agent. The runtime-field swap — including the - # api_mode pin in the moa branch — happens before any of that, so the - # invariant we care about is already set even if a later step blew up. - pass + agent = _make_fake_agent(primary_api_mode) + cache_policy = Mock(wraps=agent._anthropic_prompt_cache_policy) + monkeypatch.setattr(agent, "_anthropic_prompt_cache_policy", cache_policy) + arh.switch_model( + agent, + new_model="frontier", + new_provider="moa", + api_key="moa-virtual-provider", + base_url="moa://local", + api_mode=incoming_api_mode, + ) assert agent.provider == "moa" assert agent.base_url == "moa://local" @@ -88,3 +92,8 @@ def test_switch_to_moa_pins_chat_completions(monkeypatch, incoming_api_mode): ) # The MoAClient facade should be installed as the client. assert type(agent.client).__name__ == "MoAClient" + assert cache_policy.call_args.kwargs["api_mode"] == agent.api_mode + assert agent._primary_runtime["provider"] == agent.provider + assert agent._primary_runtime["api_mode"] == agent.api_mode + assert "anthropic_api_key" not in agent._primary_runtime + assert "anthropic_base_url" not in agent._primary_runtime diff --git a/tests/agent/test_prompt_cache_ttl_propagation.py b/tests/agent/test_prompt_cache_ttl_propagation.py index f2ab8d97f349a..fd6099997b6b7 100644 --- a/tests/agent/test_prompt_cache_ttl_propagation.py +++ b/tests/agent/test_prompt_cache_ttl_propagation.py @@ -13,6 +13,192 @@ import ast import inspect +import copy +from contextlib import ExitStack +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest +from tests.run_agent.test_run_agent import agent as canonical_minimal_agent + + +@pytest.fixture +def minimal_agent(request): + # Canonical constructor, with its metadata warm-up made inert as well. + with patch("agent.agent_init.fetch_model_metadata", return_value={}): + return request.getfixturevalue("canonical_minimal_agent") + + +class TestFailoverConversationBehavior: + """Full-loop witnesses with canonical real-agent construction and inert LLMs.""" + + @staticmethod + def response(text): + return SimpleNamespace(id="inert", model="inert", usage=None, + choices=[SimpleNamespace(finish_reason="stop", message=SimpleNamespace( + content=text, tool_calls=None, reasoning=None, reasoning_content=None))]) + + @staticmethod + def prepare(item): + item._cached_system_prompt = "Model: primary\nProvider: openrouter\nAnswer the current question." + item.model = "g2-primary" + item.provider = "openrouter" + item._use_prompt_caching = False + item.save_trajectories = False + item.context_rebase_enabled = False + item.platform = "cron" + item._api_max_retries = 3 + item.max_iterations = 4 + item.tools = [] + item._valid_tool_names = set() + item._cached_system_prompt_static = None + item.compression_checkpoint_required = False + + @staticmethod + def external_boundaries(stack, item): + # Keep request building, preflight, dispatch and fallback real. Only + # unrelated persistence/display and external client creation are inert. + for name in ("_persist_session", "_save_trajectory", "_cleanup_task_resources", + "_try_refresh_env_client_credentials", "_should_start_quiet_spinner"): + stack.enter_context(patch.object(item, name, return_value=False)) + stack.enter_context(patch.object(item, "_create_request_openai_client", side_effect=lambda **_:item.client)) + stack.enter_context(patch.object(item, "_create_openai_client", side_effect=lambda *_, **__:item.client)) + stack.enter_context(patch("agent.auxiliary_client.get_model_context_length", return_value=64000)) + stack.enter_context(patch("agent.model_metadata.fetch_model_metadata", return_value={})) + + def test_smaller_window_fallback_rebuilds_and_compacts_before_dispatch(self, minimal_agent): + from agent import conversation_loop as loop + from agent.model_metadata import estimate_request_tokens_rough + item = minimal_agent + self.prepare(item) + item.compression_enabled = True + item.compression_in_place = False + compressor = item.context_compressor + compressor.protect_first_n = 1 + compressor.protect_last_n = 2 + compressor.update_model(model=item.model, context_length=256000, + base_url=item.base_url, api_key=item.api_key, provider=item.provider, + api_mode=item.api_mode, max_tokens=1024) + history = [{"role":"user" if i % 2 == 0 else "assistant", + "content":f"historical row {i}: " + "long evidence " * 1800} + for i in range(20)] + before = copy.deepcopy(history) + events, requests = [], [] + primary, fallback = MagicMock(), MagicMock() + fallback.base_url = "https://fallback.invalid/v1" + primary.chat.completions.create.side_effect = lambda **kw: ( + events.append(("dispatch", "primary")), requests.append(copy.deepcopy(kw)), + (_ for _ in ()).throw(ValueError("ordinary SDK validation failure")))[-1] + def complete(**kw): + events.append(("dispatch", "fallback")) + requests.append(copy.deepcopy(kw)) + return self.response("fallback completed") + fallback.chat.completions.create.side_effect = complete + item.client = primary + item._fallback_chain = [dict(provider="openai", model="g2-fallback", api_key="inert-test-key", + base_url=str(fallback.base_url), api_mode="chat_completions")] + item._fallback_index = 0 + def summary(**kw): + events.append(("summary", item.model)) + return self.response("Task Snapshot: historical evidence reviewed.\nPending Asks: answer the current question.") + real_estimate = loop.estimate_messages_tokens_rough + def measure(*args, **kwargs): + result = real_estimate(*args, **kwargs) + events.append(("preflight", item.model, compressor.context_length, result)) + return result + real_build = item._build_api_kwargs + def build(*args, **kwargs): + events.append(("build", item.model)) + return real_build(*args, **kwargs) + with ExitStack() as stack: + self.external_boundaries(stack, item) + stack.enter_context(patch("agent.auxiliary_client.resolve_provider_client", return_value=(fallback,"g2-fallback"))) + stack.enter_context(patch("agent.model_metadata.get_model_context_length", return_value=64000)) + stack.enter_context(patch("agent.context_compressor.call_llm", side_effect=summary)) + stack.enter_context(patch("agent.auxiliary_client.call_llm", side_effect=summary)) + stack.enter_context(patch.object(loop, "estimate_messages_tokens_rough", side_effect=measure)) + stack.enter_context(patch.object(item, "_build_api_kwargs", side_effect=build)) + activation = stack.enter_context(patch.object(item, "_try_activate_fallback", wraps=item._try_activate_fallback)) + result = item.run_conversation("Answer this current question.", conversation_history=history) + assert result["final_response"] == "fallback completed", (result, events) + assert activation.call_count == 1 + assert len(requests) == 2 + assert compressor.context_length == 64000 + first = events.index(("dispatch", "primary")) + last = events.index(("dispatch", "fallback")) + pressure = [i for i,event in enumerate(events) if event[0] == "preflight" and event[1] == "g2-fallback"] + summaries = [i for i,event in enumerate(events) if event[0] == "summary"] + builds = [i for i,event in enumerate(events) if event == ("build", "g2-fallback")] + assert pressure and summaries and builds, events + assert first < pressure[0] < summaries[0] < builds[-1] < last, events + assert any(events[i][3] >= compressor.threshold_tokens for i in pressure), events + assert estimate_request_tokens_rough(requests[1]["messages"], tools=requests[1].get("tools")) < compressor.threshold_tokens + assert requests[0]["messages"] != requests[1]["messages"] + assert history == before, "request rebuilding mutated the caller's history" + + def test_typed_native_refusal_has_no_provider_or_fallback_attempt(self, minimal_agent): + from agent.transports import ri_llm + from agent import chat_completion_helpers as helpers + item = minimal_agent + self.prepare(item) + item.compression_enabled = False + item.provider = "ollama-launch" + item.api_key = "no-key-required" + item.base_url = "http://inert.invalid/v1" + item._client_kwargs = {"api_key":item.api_key, "base_url":item.base_url} + item._fallback_chain = [dict(provider="openai", model="unused", api_key="inert-test-key")] + ri_llm.configure_ri_pipeline(item, {"agent":{"llm_pipeline":{"enabled":True}}}) + with ExitStack() as stack: + self.external_boundaries(stack, item) + unused_client = MagicMock() + unused_client.base_url = "https://unused.invalid/v1" + stack.enter_context(patch("agent.auxiliary_client.resolve_provider_client", return_value=(unused_client,"unused"))) + stack.enter_context(patch("agent.model_metadata.get_model_context_length", return_value=64000)) + native_pipeline, native_config = MagicMock(), MagicMock() + stack.enter_context(patch.multiple(ri_llm, create=True, + _NATIVE_AVAILABLE=True, _NativePipeline=native_pipeline, LlmConfig=native_config)) + selected = stack.enter_context(patch.object(helpers, "ri_pipeline_chat_completion", wraps=helpers.ri_pipeline_chat_completion)) + activation = stack.enter_context(patch.object(item, "_try_activate_fallback", wraps=item._try_activate_fallback)) + result = item.run_conversation("Question requiring unsupported chat history.", + conversation_history=[{"role":"user","content":"prior"},{"role":"assistant","content":"answer"}]) + assert "RI_PIPELINE_REQUEST_UNSUPPORTED" in result.get("error", ""), result + assert result.get("failed") is True + assert selected.call_count == 1 + activation.assert_not_called() + item.client.chat.completions.create.assert_not_called() + native_pipeline.assert_not_called() + native_config.assert_not_called() + unused_client.chat.completions.create.assert_not_called() + + def test_exhausted_real_preflight_prevents_dispatch(self, minimal_agent): + from agent.context_compressor import ContextCompressor + from ares_runtime.continuity.runtime import ContextDispatchError + item = minimal_agent + self.prepare(item) + item.compression_enabled = True + # Real compressor policy with a history too short to summarize. Its + # real no-progress result must stop admission before provider dispatch. + compressor = item.context_compressor + assert isinstance(compressor, ContextCompressor) + compressor.update_model(model=item.model, context_length=64000, + base_url=item.base_url, api_key=item.api_key, provider=item.provider, + api_mode=item.api_mode, max_tokens=1024) + history = [{"role":"user","content":"historical evidence " * 17000}, + {"role":"assistant","content":"historical response " * 17000}] + item.client.chat.completions.create.return_value = self.response("unexpected admission") + with ExitStack() as stack: + self.external_boundaries(stack, item) + compression = stack.enter_context(patch.object(item, "_compress_context", wraps=item._compress_context)) + summary = stack.enter_context(patch("agent.context_compressor.call_llm", side_effect=AssertionError("short history must not call summarizer"))) + build = stack.enter_context(patch.object(item, "_build_api_kwargs", wraps=item._build_api_kwargs)) + dispatch = stack.enter_context(patch.object(item, "_interruptible_api_call", wraps=item._interruptible_api_call)) + with pytest.raises(ContextDispatchError, match="CONTEXT_PREFLIGHT_EXHAUSTED"): + item.run_conversation("Answer the current question.", conversation_history=history) + assert compression.call_count >= 1 + summary.assert_not_called() + build.assert_not_called() + dispatch.assert_not_called() + item.client.chat.completions.create.assert_not_called() def _collect_cache_controls(obj): diff --git a/tests/agent/test_ri_llm_contract.py b/tests/agent/test_ri_llm_contract.py new file mode 100644 index 0000000000000..9ed5b76d9700a --- /dev/null +++ b/tests/agent/test_ri_llm_contract.py @@ -0,0 +1,563 @@ +"""The native transport must preserve or refuse the admitted request.""" +import concurrent.futures +import math +import os +import subprocess +import sys +import textwrap +from types import SimpleNamespace +import unittest +from unittest.mock import patch +from pathlib import Path + +from agent.transports import ri_llm as owner + + +class OptionalNativeImportContract(unittest.TestCase): + def _fresh_dispatch(self, state): + script = textwrap.dedent(''' + import importlib.abc + import importlib.util + import json + from pathlib import Path + import sys + from types import SimpleNamespace + import unittest + + state, root, contract_path = sys.argv[1:] + sys.path.insert(0, root) + attempted, native_calls, sdk_calls = [], [], [] + class ControlledConfig: + def __init__(self, **values): self.values = values + class ControlledPipeline: + def __init__(self, url, model, *, config): + native_calls.append((url, model, config.values)) + def call(self, prompt, *, system=None, config=None): + return "controlled native answer" + class OptionalLoader(importlib.abc.MetaPathFinder, importlib.abc.Loader): + def find_spec(self, fullname, path=None, target=None): + if fullname == "llm_pipeline" or fullname.startswith("llm_pipeline."): + attempted.append(fullname) + if state == "absent": + raise ModuleNotFoundError("controlled package absence", name=fullname) + return importlib.util.spec_from_loader(fullname, self, + is_package=(fullname == "llm_pipeline")) + return None + def create_module(self, spec): return None + def exec_module(self, module): + if module.__name__ == "llm_pipeline._native": + if state == "import_error": + raise ImportError("controlled extension loader failure") + module.LlmConfig, module.Pipeline = ControlledConfig, ControlledPipeline + def no_network(event, args): + if event in {"socket.connect", "socket.sendto", "socket.getaddrinfo", "os.system"}: + raise AssertionError("external effect denied: " + event) + sys.addaudithook(no_network) + sys.meta_path.insert(0, OptionalLoader()) + from agent.transports import ri_llm as fresh + from agent.chat_completion_helpers import _dispatch_nonstreaming_api_request + assert attempted, "native import boundary was not exercised" + assert fresh._NATIVE_AVAILABLE is (state == "present") + assert hasattr(fresh, "_NativePipeline") is (state == "present") + assert hasattr(fresh, "LlmConfig") is (state == "present") + def sdk_create(**payload): + sdk_calls.append(payload) + return SimpleNamespace(id="inert-sdk") + client = SimpleNamespace(chat=SimpleNamespace(completions=SimpleNamespace(create=sdk_create))) + item = SimpleNamespace(provider="ollama-launch", api_mode="chat_completions", + base_url="http://inert.invalid/v1", model="inert", api_key="no-key-required", + _ri_pipeline_enabled=True, _ri_pipeline_explicit=False, _ri_pipeline_providers=[]) + payload = dict(model="inert", messages=[{"role":"user", "content":"hi"}], temperature=0, max_tokens=7) + def dispatch(data): + return _dispatch_nonstreaming_api_request(item, data, make_client=lambda *_:client) + if state == "present": + assert fresh._should_use_ri_pipeline(item, payload) is True + assert dispatch(dict(payload)).choices[0].message.content == "controlled native answer" + assert len(native_calls) == 1 and sdk_calls == [] + unsupported = dict(payload, tools=[{"type":"function", "function":{"name":"inert"}}]) + assert fresh._should_use_ri_pipeline(item, unsupported) is False + assert dispatch(dict(unsupported)).id == "inert-sdk" + assert len(native_calls) == 1 and len(sdk_calls) == 1 + item._ri_pipeline_explicit = True + try: dispatch(dict(unsupported)) + except ValueError as error: assert "RI_PIPELINE_REQUEST_UNSUPPORTED" in str(error) + else: raise AssertionError("explicit unsupported request did not refuse") + assert len(native_calls) == 1 and len(sdk_calls) == 1 + # Exercise the original 27 methods after a real successful + # controlled import, then verify fixture restoration. + spec = importlib.util.spec_from_file_location("fresh_contract", contract_path) + contract = importlib.util.module_from_spec(spec) + spec.loader.exec_module(contract) + suite = unittest.defaultTestLoader.loadTestsFromTestCase(contract.TransportContract) + result = unittest.TextTestRunner(verbosity=0).run(suite) + assert result.wasSuccessful() and result.testsRun == 27 + assert fresh._NativePipeline is ControlledPipeline + assert fresh.LlmConfig is ControlledConfig and fresh._NATIVE_AVAILABLE is True + else: + for explicit in (False, True): + item._ri_pipeline_explicit = explicit + assert fresh._should_use_ri_pipeline(item, payload) is False + assert dispatch(dict(payload)).id == "inert-sdk" + assert native_calls == [] and len(sdk_calls) == 2 + assert not hasattr(fresh, "_NativePipeline") and not hasattr(fresh, "LlmConfig") + print(json.dumps(dict(state=state, import_attempts=attempted, + native_dispatches=len(native_calls), sdk_dispatches=len(sdk_calls)))) + ''') + result = subprocess.run( + [sys.executable, "-I", "-B", "-c", script, state, + str(Path(owner.__file__).resolve().parents[2]), __file__], + env={"PATH": os.defpath, "PYTHONDONTWRITEBYTECODE": "1", + "HERMES_HOME": os.environ.get("HERMES_HOME", "")}, + capture_output=True, text=True, timeout=20, + ) + self.assertEqual(result.returncode, 0, result.stdout + result.stderr) + + def test_fresh_absent_import_dispatches_sdk_with_zero_native_calls(self): + self._fresh_dispatch("absent") + + def test_fresh_present_exports_qualify_actual_dispatch_and_existing_contract(self): + self._fresh_dispatch("present") + + def test_fresh_extension_import_error_dispatches_sdk_with_zero_native_calls(self): + self._fresh_dispatch("import_error") + + def test_setup_failure_after_native_patch_restores_environment_and_symbols(self): + environment = dict(os.environ) + missing = object() + fields = ("_NATIVE_AVAILABLE", "_NativePipeline", "LlmConfig") + originals = {name:getattr(owner, name, missing) for name in fields} + case = TransportContract("test_default_unknown_provider_retains_sdk_route") + real_enter = case.enterContext + entries = [] + def enter_then_fail(context): + value = real_enter(context) + entries.append(context) + if len(entries) == 2: + self.assertIs(owner._NativePipeline, NativePipeline) + self.assertIs(owner.LlmConfig, NativeConfig) + self.assertTrue(owner._NATIVE_AVAILABLE) + raise RuntimeError("fault after native patch entry") + return value + case.enterContext = enter_then_fail + result = unittest.TestResult() + case.run(result) + self.assertEqual(len(entries), 2) + self.assertEqual(len(result.errors), 1) + self.assertIn("fault after native patch entry", result.errors[0][1]) + self.assertEqual(dict(os.environ), environment) + for name, previous in originals.items(): + if previous is missing: + self.assertFalse(hasattr(owner, name), name) + else: + self.assertIs(getattr(owner, name), previous) + + def test_cold_import_without_optional_native_retains_sdk(self): + script = textwrap.dedent(''' + import importlib.util + import sys + from types import SimpleNamespace + + attempted = [] + class DeniedOptionalImport: + def find_spec(self, fullname, path=None, target=None): + if fullname == "llm_pipeline" or fullname.startswith("llm_pipeline."): + attempted.append(fullname) + raise ModuleNotFoundError("controlled optional-native absence", name=fullname) + return None + + sys.meta_path.insert(0, DeniedOptionalImport()) + spec = importlib.util.spec_from_file_location("cold_ri_llm", sys.argv[1]) + fresh = importlib.util.module_from_spec(spec) + spec.loader.exec_module(fresh) + assert attempted, "optional import was not exercised" + assert fresh._NATIVE_AVAILABLE is False + assert not hasattr(fresh, "_NativePipeline") + assert not hasattr(fresh, "LlmConfig") + for explicit in (False, True): + item = SimpleNamespace(provider="ollama-launch", _ri_pipeline_enabled=True, + _ri_pipeline_explicit=explicit, _ri_pipeline_providers=[]) + assert fresh._should_use_ri_pipeline(item, {}) is False + pipeline = fresh.RiPipeline("http://inert.invalid", "inert") + assert pipeline.available is False + assert fresh.RiLlmConfig()._to_native() is None + try: + pipeline.call("inert") + except RuntimeError as error: + assert "not installed" in str(error) + else: + raise AssertionError("unavailable native pipeline accepted a call") + assert not hasattr(fresh, "_NativePipeline") + assert not hasattr(fresh, "LlmConfig") + ''') + result = subprocess.run( + [sys.executable, "-I", "-B", "-c", script, owner.__file__], + env={"PATH": os.defpath, "PYTHONDONTWRITEBYTECODE": "1"}, + capture_output=True, text=True, timeout=10, + ) + self.assertEqual(result.returncode, 0, result.stdout + result.stderr) + + +class NativeConfig: + def __init__(self, **values): + self.values = values + + +class NativePipeline: + calls = [] + + def __init__(self, url, model, *, config): + self.trace = dict(url=url, model=model, config=config.values) + self.calls.append(self.trace) + + def call(self, prompt, *, system=None, config=None): + self.trace.update(prompt=prompt, system=system) + return "plain answer" + + def call_structured(self, *args, **kwargs): + raise AssertionError("unsupported structured tool request reached native") + + +def agent(provider="ollama-launch", **overrides): + values = dict(provider=provider, model="llama", api_mode="chat_completions", + api_key="no-key-required", base_url="http://localhost:11434/v1", + _interrupt_requested=False, platform="cron") + values.update(overrides) + return SimpleNamespace(**values) + + +def request(**overrides): + values = dict(model="chosen-model", messages=[{"role":"user", "content":"raw user"}], + temperature=0, max_tokens=7) + values.update(overrides) + return values + + +class TransportContract(unittest.TestCase): + def setUp(self): + self.enterContext(patch.dict(os.environ, {}, clear=True)) + self.enterContext(patch.multiple( + owner, create=True, _NATIVE_AVAILABLE=True, + _NativePipeline=NativePipeline, LlmConfig=NativeConfig, + )) + NativePipeline.calls = [] + + def qualify(self, item, config): + owner.configure_ri_pipeline(item, config) + return item + + def test_supported_ollama_preserves_parameters_content_and_unknown_usage(self): + item = agent() + payload = request(messages=[{"role":"system", "content":"exact system"}, + {"role":"user", "content":"raw user"}]) + result = owner.ri_pipeline_chat_completion(item, payload) + self.assertEqual(NativePipeline.calls, [dict(url=item.base_url, model="chosen-model", + config=dict(temperature=0, max_tokens=7, thinking=False, json_mode=False), + prompt="raw user", system="exact system")]) + self.assertIsNone(result.usage) + self.assertEqual(result.model, "chosen-model") + self.assertEqual(result.choices[0].message.content, "plain answer") + self.assertIsNone(result.choices[0].message.tool_calls) + self.assertIsNone(result.choices[0].finish_reason) + + def test_configured_off_beats_available_native_and_env_whitelist(self): + item = self.qualify(agent(), {"agent":{"llm_pipeline":{"enabled":False}}}) + os.environ["HERMES_RI_PIPELINE_PROVIDERS"] = "ollama-launch" + os.environ["HERMES_RI_PIPELINE"] = "1" + self.assertFalse(owner._should_use_ri_pipeline(item, request())) + self.assertEqual(NativePipeline.calls, []) + + def test_configured_whitelist_does_not_admit_another_provider(self): + item = self.qualify(agent("openrouter"), {"agent":{"llm_pipeline":{"providers":["ollama-launch"]}}}) + self.assertFalse(owner._should_use_ri_pipeline(item, request())) + + def test_env_whitelist_precedes_config_and_selected_incompatible_refuses(self): + item = self.qualify(agent("openrouter"), {"agent":{"llm_pipeline":{"providers":["ollama-launch"]}}}) + os.environ["HERMES_RI_PIPELINE_PROVIDERS"] = "openrouter" + self.assertTrue(owner._should_use_ri_pipeline(item, request())) + with self.assertRaisesRegex(ValueError, "RI_PIPELINE_REQUEST_UNSUPPORTED"): + owner.ri_pipeline_chat_completion(item, request()) + self.assertEqual(NativePipeline.calls, []) + + def test_env_off_overrides_explicit_selection(self): + item = self.qualify(agent(), {"agent":{"llm_pipeline":{"enabled":True}}}) + os.environ["HERMES_RI_PIPELINE"] = "0" + self.assertFalse(owner._should_use_ri_pipeline(item, request())) + + def test_default_unknown_provider_retains_sdk_route(self): + for provider in ("openrouter", "openai", "unknown-provider", ""): + with self.subTest(provider=provider): + item = self.qualify(agent(provider), {}) + self.assertFalse(owner._should_use_ri_pipeline(item, request())) + + def test_native_unavailable_retains_sdk_even_when_explicit(self): + item = self.qualify(agent(), {"agent":{"llm_pipeline":{"enabled":True}}}) + with patch.object(owner, "_NATIVE_AVAILABLE", False): + self.assertFalse(owner._should_use_ri_pipeline(item, request())) + + def test_default_unsupported_shapes_retain_sdk(self): + item = self.qualify(agent(), {}) + for payload in self.unsupported(): + with self.subTest(payload=payload): + self.assertFalse(owner._should_use_ri_pipeline(item, payload)) + + @staticmethod + def unsupported(): + return [ + request(tools=[{"type":"function", "function":{"name":"search"}}], tool_choice="none"), + request(messages=[{"role":"user", "content":[{"type":"text", "text":"hi"}, + {"type":"image_url", "image_url":{"url":"https://inert.invalid"}}]}]), + request(messages=[{"role":"developer", "content":"policy"}, {"role":"user", "content":"hi"}]), + request(messages=[{"role":"user", "content":"first"}, {"role":"assistant", "content":"prior"}, + {"role":"user", "content":"next"}]), + request(messages=[{"role":"assistant", "content":None, "tool_calls":[]}, + {"role":"tool", "tool_call_id":"old", "content":"result"}]), + request(timeout=10), request(extra_headers={"X-Test":"yes"}), + request(extra_body={"thinking":True}), request(response_format={"type":"json_object"}), + request(messages=[{"role":"system", "content":""},{"role":"user", "content":"hi"}]), + request(messages=[{"role":"user", "content":"explain literal {input}"}]), + ] + + def test_explicit_unsupported_shapes_refuse_before_native_effects(self): + item = self.qualify(agent(), {"agent":{"llm_pipeline":{"enabled":True}}}) + for payload in self.unsupported(): + with self.subTest(payload=payload): + self.assertTrue(owner._should_use_ri_pipeline(item, payload)) + with self.assertRaisesRegex(ValueError, "RI_PIPELINE_REQUEST_UNSUPPORTED"): + owner.ri_pipeline_chat_completion(item, payload) + self.assertEqual(NativePipeline.calls, []) + + def test_authenticated_and_callable_credentials_refuse_without_evaluation(self): + def forbidden_key(): + self.fail("credential callable was evaluated") + for key in ("SYNTHETIC-KEY", forbidden_key): + with self.subTest(key=type(key).__name__): + item = self.qualify(agent(api_key=key), {"agent":{"llm_pipeline":{"enabled":True}}}) + with self.assertRaisesRegex(ValueError, "RI_PIPELINE_REQUEST_UNSUPPORTED"): + owner.ri_pipeline_chat_completion(item, request()) + self.assertEqual(NativePipeline.calls, []) + + def test_invalid_generation_settings_and_endpoint_refuse(self): + for payload in (request(max_tokens=0), request(max_tokens=True), request(max_tokens=1.5), + request(temperature=math.nan), request(temperature=True), request(stream=1)): + with self.subTest(payload=payload): + with self.assertRaisesRegex(ValueError, "RI_PIPELINE_REQUEST_UNSUPPORTED"): + owner.ri_pipeline_chat_completion(agent(), payload) + for url in ("https://inert.invalid/custom", "https://user:synthetic@inert.invalid/v1", + "https://inert.invalid/v1?key=synthetic", "http://localhost:bad/v1", "not-a-url"): + with self.subTest(url=url): + with self.assertRaisesRegex(ValueError, "RI_PIPELINE_REQUEST_UNSUPPORTED"): + owner.ri_pipeline_chat_completion(agent(base_url=url), request()) + self.assertEqual(NativePipeline.calls, []) + + def test_concurrent_supported_calls_leave_process_credentials_untouched(self): + os.environ["OPENAI_API_KEY"] = "SYNTHETIC-SENTINEL" + with concurrent.futures.ThreadPoolExecutor(max_workers=2) as pool: + results = list(pool.map(lambda _: owner.ri_pipeline_chat_completion(agent(), request()), range(2))) + self.assertEqual(os.environ["OPENAI_API_KEY"], "SYNTHETIC-SENTINEL") + self.assertTrue(all(item.usage is None for item in results)) + + def test_actual_downstream_usage_normalizer_retains_missing_raw_usage(self): + from agent.usage_pricing import normalize_usage + response = owner.ri_pipeline_chat_completion(agent(), request()) + usage = normalize_usage(response.usage, provider="ollama-launch", api_mode="chat_completions") + self.assertIsNone(usage.raw_usage) + self.assertEqual(usage.prompt_tokens, 0) + self.assertFalse(bool(response.usage)) + + def test_native_unknown_finish_survives_actual_normalizer_and_message_builder(self): + from agent.transports.chat_completions import ChatCompletionsTransport + from agent import chat_completion_helpers as host + item = agent() + response = owner.ri_pipeline_chat_completion(item, request()) + normalized = ChatCompletionsTransport().normalize_response(response) + self.assertIsNone(normalized.finish_reason) + self.assertIsNone(normalized.usage) + item._extract_reasoning = lambda *_:None + item._strip_think_blocks = lambda text:text + # Content redaction is outside this metadata regression; no secret read. + redactor = SimpleNamespace(redact_sensitive_text=lambda text:text) + with patch.dict(sys.modules, {"agent.redact":redactor}): + persisted = host.build_assistant_message(item, normalized, normalized.finish_reason) + self.assertIsNone(persisted["finish_reason"]) + self.assertEqual(persisted["content"], "plain answer") + + def test_sdk_missing_finish_and_text_marker_keep_legacy_stop_default(self): + from agent.transports.chat_completions import ChatCompletionsTransport + for text in ("plain answer", "RiCompletionResponse unknown native finish"): + response = SimpleNamespace(choices=[SimpleNamespace( + message=SimpleNamespace(content=text, tool_calls=None), finish_reason=None)], + usage=None, native=True, _ri_finish_unknown=True) + normalized = ChatCompletionsTransport().normalize_response(response) + self.assertEqual(normalized.finish_reason, "stop") + self.assertIsNone(normalized.usage) + + def test_omitted_generation_settings_keep_sdk_or_refuse_explicit_native(self): + for key in ("temperature", "max_tokens"): + payload = request() + del payload[key] + ordinary = self.qualify(agent(), {}) + self.assertFalse(owner._should_use_ri_pipeline(ordinary, payload)) + selected = self.qualify(agent(), {"agent":{"llm_pipeline":{"enabled":True}}}) + self.assertTrue(owner._should_use_ri_pipeline(selected, payload)) + with self.assertRaises(owner.RiTransportUnsupported): + owner.ri_pipeline_chat_completion(selected, payload) + self.assertEqual(NativePipeline.calls, []) + + def test_canonical_client_overrides_keep_sdk_or_refuse_before_native(self): + for options in ({"default_headers":{"X-Inert":"required"}}, + {"http_client":object()}, {"timeout":120}, + {"default_query":{"api-version":"synthetic"}}, + {"api_key":"SYNTHETIC-OVERRIDE"}, + {"base_url":"http://localhost:11435/v1"}): + with self.subTest(option=next(iter(options))): + ordinary = self.qualify(agent(_client_kwargs=options), {}) + self.assertFalse(owner._should_use_ri_pipeline(ordinary, request())) + selected = self.qualify(agent(_client_kwargs=options), + {"agent":{"llm_pipeline":{"enabled":True}}}) + with self.assertRaises(owner.RiTransportUnsupported): + owner.ri_pipeline_chat_completion(selected, request()) + self.assertEqual(NativePipeline.calls, []) + + def test_actual_dispatch_keeps_omissions_and_client_options_on_sdk(self): + from agent import chat_completion_helpers as host + item = self.qualify(agent(_client_kwargs={"default_headers":{"X-Inert":"required"}}), {}) + payload = request() + del payload["max_tokens"] + sdk_calls = [] + sdk = SimpleNamespace(chat=SimpleNamespace(completions=SimpleNamespace( + create=lambda **kwargs:sdk_calls.append(kwargs) or "sdk-response"))) + self.assertEqual(host._dispatch_nonstreaming_api_request(item, payload, + make_client=lambda *_:sdk), "sdk-response") + self.assertEqual(sdk_calls, [payload]) + self.assertNotIn("max_tokens", payload) + self.assertEqual(item._client_kwargs, {"default_headers":{"X-Inert":"required"}}) + self.assertEqual(NativePipeline.calls, []) + + def test_actual_host_dispatch_retains_sdk_and_refuses_explicit_selection(self): + from agent import chat_completion_helpers as host + payload = request(tools=[]) + sdk_calls = [] + sdk = SimpleNamespace(chat=SimpleNamespace(completions=SimpleNamespace( + create=lambda **kwargs: sdk_calls.append(kwargs) or "sdk-response"))) + ordinary = self.qualify(agent("openrouter"), {}) + self.assertEqual(host._dispatch_nonstreaming_api_request(ordinary, payload, make_client=lambda *_:sdk), "sdk-response") + self.assertEqual(sdk_calls, [payload]) + selected = self.qualify(agent(), {"agent":{"llm_pipeline":{"enabled":True}}}) + with self.assertRaisesRegex(ValueError, "RI_PIPELINE_REQUEST_UNSUPPORTED"): + host._dispatch_nonstreaming_api_request(selected, payload, make_client=lambda *_:self.fail("SDK fallback")) + self.assertEqual(NativePipeline.calls, []) + + def test_actual_host_nonstream_and_cron_stream_use_supported_native(self): + from agent import chat_completion_helpers as host + item = self.qualify(agent(), {}) + item._interruptible_api_call = lambda kwargs: host._dispatch_nonstreaming_api_request( + item, kwargs, make_client=lambda *_:self.fail("unexpected SDK dispatch")) + response = host.interruptible_streaming_api_call(item, request(stream=True)) + self.assertEqual(response.choices[0].message.content, "plain answer") + self.assertIsNone(response.usage) + + def moa_receiver(self, config): + trace = [] + response = SimpleNamespace(id="moa-sdk-response") + + def prepare(**kwargs): + trace.append(("prepare", kwargs)) + + def create(**kwargs): + prepare(**kwargs) + trace.append(("create", kwargs)) + return response + + item = self.qualify(agent("moa"), config) + item.client = SimpleNamespace(chat=SimpleNamespace( + completions=SimpleNamespace(prepare=prepare, create=create))) + return item, trace, response + + def test_actual_host_selected_moa_refuses_before_prepare_or_create(self): + from agent import chat_completion_helpers as host + for config in ({"agent":{"llm_pipeline":{"enabled":True}}}, + {"agent":{"llm_pipeline":{"providers":["moa"]}}}): + for streamed in (False, True): + with self.subTest(config=config, streamed=streamed): + item, trace, _ = self.moa_receiver(config) + with self.assertRaises(owner.RiTransportUnsupported): + host._dispatch_nonstreaming_api_request(item, request(stream=streamed), + make_client=lambda *_:self.fail("unexpected client construction")) + self.assertEqual(trace, []) + self.assertEqual(NativePipeline.calls, []) + + def test_actual_host_default_moa_preserves_facade_prepare_and_create(self): + from agent import chat_completion_helpers as host + item, trace, response = self.moa_receiver({}) + payload = request(_moa_prepared_request="inert-prepared") + result = host._dispatch_nonstreaming_api_request(item, payload, + make_client=lambda *_:self.fail("MoA facade was replaced")) + self.assertIs(result, response) + self.assertEqual(trace, [("prepare",payload), ("create",payload)]) + self.assertEqual(payload["_moa_prepared_request"], "inert-prepared") + self.assertEqual(NativePipeline.calls, []) + + def test_actual_host_selected_moa_cron_stream_refuses_before_effects(self): + from agent import chat_completion_helpers as host + item, trace, _ = self.moa_receiver({"agent":{"llm_pipeline":{"enabled":True}}}) + item._interruptible_api_call = lambda kwargs: host._dispatch_nonstreaming_api_request( + item, kwargs, make_client=lambda *_:self.fail("unexpected client construction")) + with self.assertRaises(owner.RiTransportUnsupported): + host.interruptible_streaming_api_call(item, request(stream=True)) + self.assertEqual(trace, []) + self.assertEqual(NativePipeline.calls, []) + + def test_actual_host_default_moa_cron_stream_keeps_facade_route(self): + from agent import chat_completion_helpers as host + item, trace, response = self.moa_receiver({}) + item._interruptible_api_call = lambda kwargs: host._dispatch_nonstreaming_api_request( + item, kwargs, make_client=lambda *_:self.fail("MoA facade was replaced")) + payload = request(stream=True) + # MoA keeps its established worker; exercise its canonical dispatch seam + # without launching the real interrupt/stream worker in this offline test. + self.assertFalse(host.should_use_direct_api_call(item)) + self.assertFalse(owner._should_use_ri_pipeline(item, payload)) + result = host._dispatch_nonstreaming_api_request(item, payload, + make_client=lambda *_:self.fail("MoA facade was replaced")) + self.assertIs(result, response) + self.assertEqual(trace, [("prepare",payload), ("create",payload)]) + self.assertEqual(NativePipeline.calls, []) + + def test_actual_error_classifier_does_not_offer_retry_for_typed_refusal(self): + from agent.error_classifier import classify_api_error + item = self.qualify(agent(), {"agent":{"llm_pipeline":{"enabled":True}}}) + try: + owner.ri_pipeline_chat_completion(item, request(tools=[])) + except ValueError as error: + verdict = classify_api_error(error, provider="ollama-launch") + self.assertFalse(verdict.retryable) + self.assertFalse(verdict.should_fallback) + self.assertFalse(verdict.should_rotate_credential) + self.assertFalse(verdict.should_compress) + self.assertTrue(verdict.error_context.get("native_transport_refusal")) + else: + self.fail("unsupported request did not refuse") + + def test_provider_text_and_unknown_errors_keep_existing_retry_behavior(self): + from agent.error_classifier import classify_api_error + for error in (RuntimeError("RI_PIPELINE_REQUEST_UNSUPPORTED"), + ValueError("request fields unavailable in native binding")): + with self.subTest(error=type(error).__name__): + verdict = classify_api_error(error, provider="ollama-launch") + self.assertTrue(verdict.retryable) + self.assertFalse(verdict.error_context.get("native_transport_refusal")) + + def test_plugin_cannot_attach_native_refusal_marker_to_provider_error(self): + from agent.error_classifier import classify_api_error, FailoverReason + plugin = SimpleNamespace(get_plugin_error_classification=lambda **kwargs: { + "reason":FailoverReason.unknown, "retryable":True, + "error_context":{"native_transport_refusal":True, "plugin_note":"retained"}, + }) + with patch.dict(sys.modules, {"hermes_cli.plugins":plugin}): + verdict = classify_api_error(RuntimeError("RI_PIPELINE_REQUEST_UNSUPPORTED")) + self.assertTrue(verdict.retryable) + self.assertEqual(verdict.error_context, {"plugin_note":"retained"}) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/agent/test_secret_scope.py b/tests/agent/test_secret_scope.py index de62ec10212d2..c53d475087889 100644 --- a/tests/agent/test_secret_scope.py +++ b/tests/agent/test_secret_scope.py @@ -223,11 +223,13 @@ def test_build_profile_secret_scope_includes_home_external_secrets( (tmp_path / ".env").write_text("XIAOMI_API_KEY=placeholder\n") from hermes_cli import env_loader - home_key = str(tmp_path.resolve()) - monkeypatch.setitem( - env_loader._SECRET_SOURCE_VALUES_BY_HOME, - home_key, - {"XIAOMI_API_KEY": "sk-from-bitwarden"}, + # Seed the authoritative current generation. A legacy value-only + # projection cannot replace a retained/stale typed snapshot when + # pytest reuses the temp home after removing a passing fixture. + env_loader._record_external_secret_snapshot( + tmp_path, + data={"XIAOMI_API_KEY": "sk-from-bitwarden"}, + status="ready", ) assert ss.build_profile_secret_scope(tmp_path) == { diff --git a/tests/agent/test_verification_evidence.py b/tests/agent/test_verification_evidence.py index aae41eda91e40..6031e272b1f76 100644 --- a/tests/agent/test_verification_evidence.py +++ b/tests/agent/test_verification_evidence.py @@ -109,7 +109,7 @@ def test_masking_shell_control_is_not_verification_evidence( assert evidence is None -@pytest.mark.parametrize("command", ["prepare && pytest", "pytest && report"]) +@pytest.mark.parametrize("command", ["pytest && report", "pytest && cd elsewhere && cd back"]) def test_successful_and_chain_preserves_passing_evidence( tmp_path, monkeypatch, command ): @@ -127,9 +127,9 @@ def test_successful_and_chain_preserves_passing_evidence( assert evidence.status == "passed" -@pytest.mark.parametrize("exit_code, expected", [(0, "passed"), (1, "failed")]) -def test_final_verifier_after_sequence_owns_shell_exit_status( - tmp_path, monkeypatch, exit_code, expected +@pytest.mark.parametrize("exit_code", [0, 1]) +def test_final_verifier_after_sequence_has_unbound_workspace( + tmp_path, monkeypatch, exit_code ): monkeypatch.setenv("HERMES_HOME", str(tmp_path / ".hermes")) _python_project(tmp_path) @@ -141,8 +141,114 @@ def test_final_verifier_after_sequence_owns_shell_exit_status( exit_code=exit_code, ) + assert evidence is None + + +@pytest.mark.parametrize( + "command", + ["cd elsewhere && pytest", "cd elsewhere && pytest && cd back", "prepare && pytest", "prepare\npytest"], +) +def test_preceding_shell_command_cannot_establish_original_workspace_freshness( + tmp_path, monkeypatch, command +): + monkeypatch.setenv("HERMES_HOME", str(tmp_path / ".hermes")) + _python_project(tmp_path) + assert classify_verification_command(command, cwd=tmp_path, session_id="s1", exit_code=0) is None + + +@pytest.mark.parametrize("prefix", ["cd elsewhere && ", "prepare; ", "prepare && "]) +def test_ad_hoc_verifier_after_shell_command_has_unbound_workspace(tmp_path, monkeypatch, prefix): + monkeypatch.setenv("HERMES_HOME", str(tmp_path / ".hermes")) + (tmp_path / "package.json").write_text("{}", encoding="utf-8") + script = Path(tempfile.gettempdir()) / f"hermes-verify-cwd-{tmp_path.name}.py" + assert classify_verification_command( + f"{prefix}python {script}", cwd=tmp_path, session_id="s1", exit_code=0 + ) is None + + +@pytest.mark.parametrize( + "command", + [ + "pytest -k smoke", "pytest -ksmoke", "pytest -k=smoke", + "pytest -m smoke", "pytest -msmoke", "pytest -m=smoke", "pytest -qk smoke", + "pytest --deselect=smoke", "pytest --lf", "pytest --unknown-selection=smoke", + "HERMES_TEST_SLICE=1/8 pytest", "env HERMES_TEST_PATHS=unit pytest", + "PYTEST_ADDOPTS='-k smoke' pytest", + "TEST_FILTER=smoke pytest", "pytest time", "pytest root=unit", + "pytest tests/unit.py::test_case[value=smoke]", + "pytest -- -", + ], +) +def test_selection_and_ambiguous_arguments_are_targeted(tmp_path, monkeypatch, command): + monkeypatch.setenv("HERMES_HOME", str(tmp_path / ".hermes")) + _python_project(tmp_path) + evidence = classify_verification_command(command, cwd=tmp_path, session_id="s1", exit_code=0) + assert evidence is not None + assert evidence.scope == "targeted" + + +@pytest.mark.parametrize( + "arguments", + ["--slice=1/8", "--slice 1/8", "--files=unit", "--paths=unit"], +) +def test_runner_selection_arguments_are_targeted(tmp_path, monkeypatch, arguments): + monkeypatch.setenv("HERMES_HOME", str(tmp_path / ".hermes")) + _node_project(tmp_path) + evidence = classify_verification_command( + f"scripts/run_tests.sh {arguments}", cwd=tmp_path, session_id="s1", exit_code=0 + ) + assert evidence is not None + assert evidence.scope == "targeted" + + +def test_go_run_selector_is_not_pytest_reporting_option(tmp_path, monkeypatch): + from agent import coding_context + + monkeypatch.setattr(coding_context, "project_facts_for", lambda cwd: { + "root": str(tmp_path), "verifyCommands": ["go test"], + }) + evidence = classify_verification_command( + "go test -run=Smoke", cwd=tmp_path, session_id="s1", exit_code=0 + ) + assert evidence is not None + assert evidence.scope == "targeted" + + +def test_make_test_selection_assignment_is_not_a_shell_prefix(tmp_path, monkeypatch): + from agent import coding_context + + monkeypatch.setattr(coding_context, "project_facts_for", lambda cwd: { + "root": str(tmp_path), "verifyCommands": ["make test"], + }) + evidence = classify_verification_command( + "make test TEST=smoke", cwd=tmp_path, session_id="s1", exit_code=0 + ) + assert evidence is not None + assert evidence.scope == "targeted" + + +@pytest.mark.parametrize("arguments", ["-q", "-vv", "-n 2", "--tb=short", "--junitxml=reports/results.xml"]) +def test_reporting_and_worker_options_preserve_unrestricted_selection(tmp_path, monkeypatch, arguments): + monkeypatch.setenv("HERMES_HOME", str(tmp_path / ".hermes")) + _python_project(tmp_path) + evidence = classify_verification_command(f"pytest {arguments}", cwd=tmp_path, session_id="s1", exit_code=0) assert evidence is not None - assert evidence.status == expected + assert evidence.scope == "full" + + +def test_other_workspace_verifier_cannot_clear_initial_workspace_stale_state(tmp_path, monkeypatch): + monkeypatch.setenv("HERMES_HOME", str(tmp_path / ".hermes")) + project_a, project_b = tmp_path / "a", tmp_path / "b" + for project in (project_a, project_b): + project.mkdir() + _python_project(project) + record_terminal_result(command="pytest", cwd=project_a, session_id="conversation", exit_code=0) + mark_workspace_edited(session_id="conversation", cwd=project_a, paths=[str(project_a / "changed.py")]) + assert record_terminal_result( + command=f"cd {project_b} && pytest && cd {project_a}", + cwd=project_a, session_id="conversation", exit_code=0, + ) is None + assert verification_status(session_id="conversation", cwd=project_a)["status"] == "stale" def test_quoted_shell_operator_remains_a_verifier_argument(tmp_path, monkeypatch): diff --git a/tests/agent/test_verification_stop.py b/tests/agent/test_verification_stop.py index f0d221a50d025..672088f6e3abe 100644 --- a/tests/agent/test_verification_stop.py +++ b/tests/agent/test_verification_stop.py @@ -8,6 +8,7 @@ from agent.verification_evidence import ( mark_workspace_edited, record_terminal_result, + verification_status, ) from agent.verification_stop import ( build_verify_on_stop_nudge, @@ -149,6 +150,50 @@ def test_nudge_checks_all_edited_workspaces(tmp_path, monkeypatch): assert "fresh passing verification evidence" in nudge +@pytest.mark.parametrize("operation", ["update-delete", "move"]) +def test_patch_producer_stales_every_workspace_before_stop(tmp_path, monkeypatch, operation): + """Real producer/state/consumer; patch bytes stop at an inert backend.""" + from types import SimpleNamespace + from tools import file_tools + + monkeypatch.setenv("HERMES_HOME", str(tmp_path / ".hermes")) + project_a, project_b = tmp_path / "a", tmp_path / "b" + for project in (project_a, project_b): + _make_project(project) + (project / "src").mkdir() + (project / "src" / "app.ts").write_text("old\n", encoding="utf-8") + record_terminal_result(command="pnpm test", cwd=project, session_id="conversation", exit_code=0) + changed_a = str(project_a / "src" / "app.ts") + changed_b = str(project_b / "src" / "app.ts") + headers = ( + f"*** Move File: {changed_a} -> {changed_b}" + if operation == "move" + else f"*** Update File: {changed_a}\n@@\n-old\n+new\n*** Delete File: {changed_b}" + ) + backend = SimpleNamespace(patch_v4a=lambda patch: SimpleNamespace(to_dict=lambda: {"success": True})) + monkeypatch.setattr(file_tools, "_get_file_ops", lambda task_id: backend) + result = json.loads(file_tools.patch_tool( + mode="patch", patch=f"*** Begin Patch\n{headers}\n*** End Patch\n", + task_id="turn", session_id="conversation", + )) + assert "error" not in result + assert set(result["files_modified"]) == {changed_a, changed_b} + for project in (project_a, project_b): + status = verification_status(session_id="conversation", cwd=project) + assert status["status"] == "stale" + assert status["changed_paths"] == [str(project / "src" / "app.ts")] + assert verification_status(session_id="turn", cwd=project)["status"] == "unverified" + # Refresh only A. The real stop consumer must still notice stale B. + record_terminal_result(command="pnpm test", cwd=project_a, session_id="conversation", exit_code=0) + assert build_verify_on_stop_nudge( + session_id="conversation", changed_paths=[changed_a, changed_b] + ) is not None + record_terminal_result(command="pnpm test", cwd=project_b, session_id="conversation", exit_code=0) + assert build_verify_on_stop_nudge( + session_id="conversation", changed_paths=[changed_a, changed_b] + ) is None + + diff --git a/tests/ares_runtime/test_authority_failure_atomicity.py b/tests/ares_runtime/test_authority_failure_atomicity.py new file mode 100644 index 0000000000000..cf807dad9ee1e --- /dev/null +++ b/tests/ares_runtime/test_authority_failure_atomicity.py @@ -0,0 +1,151 @@ +"""Pure authority failure atomicity and temporal composition controls.""" +import copy +import unittest +from ares_runtime import authority as AUTH + +def state(authority): + return copy.deepcopy(authority.__dict__) + +class AuthorityProbes(unittest.TestCase): + def parent(self, uses=2): + return AUTH.AuthorityScopeV1(scope={'tool': 'write_file', 'target': 'path:/fixture', 'use_count': uses}, generation=1, holder='holder:root') + + def reserved(self): + parent = self.parent() + parent.reserve(consumption_ref='consume:1', args_digest='sha256:args', target_ref='path:/fixture') + return parent + + def test_rejected_child_keeps_every_parent_field(self): + for holder in ['', 7, []]: + with self.subTest(holder=holder): + parent = self.parent(1) + before = state(parent) + with self.assertRaises(AUTH.ContractError) as raised: + parent.attenuate({'use_count': 1}, child_generation=2, child_holder=holder) + self.assertEqual(raised.exception.code, 'INVALID_HOLDER') + self.assertEqual(state(parent), before) + child = parent.attenuate({'use_count': 1}, child_generation=2, child_holder='holder:child') + self.assertTrue(child.subset_witness(parent)['contained']) + + def test_valid_child_spends_once_and_cannot_duplicate(self): + parent = self.parent(1) + child = parent.attenuate({'use_count': 1}, child_generation=2) + self.assertTrue(child.is_subset(parent)) + self.assertEqual(parent._delegated_count, 1) + with self.assertRaises(AUTH.ContractError): + parent.attenuate({'use_count': 1}, child_generation=2) + + def test_rejected_generation_and_scope_preserve_parent(self): + parent = self.parent(1) + before = state(parent) + for scope, generation in [({}, 1), ({'target': 'path:/other'}, 2), ({'use_count': 2}, 2)]: + with self.subTest(scope=scope, generation=generation): + with self.assertRaises(AUTH.ContractError): + parent.attenuate(scope, child_generation=generation) + self.assertEqual(state(parent), before) + + def test_invalid_commit_evidence_does_not_settle(self): + for value in ['', ' ', None, 7, {}, float('nan'), float('inf')]: + with self.subTest(value=repr(value)): + parent = self.reserved() + before = state(parent) + with self.assertRaises(AUTH.ContractError) as raised: + parent.commit('consume:1', effect_receipt_digest=value) + self.assertEqual(raised.exception.code, 'INVALID_EFFECT_RECEIPT_DIGEST') + self.assertEqual(state(parent), before) + receipt = parent.commit('consume:1', effect_receipt_digest='sha256:effect') + self.assertEqual(receipt['record']['state'], 'committed') + + def test_invalid_release_reason_does_not_settle(self): + for operation in ['release', 'mark_indeterminate']: + for value in ['', ' ', None, float('nan')]: + with self.subTest(operation=operation, value=repr(value)): + parent = self.reserved() + before = state(parent) + with self.assertRaises(AUTH.ContractError): + getattr(parent, operation)('consume:1', reason=value) + self.assertEqual(state(parent), before) + + def test_serializer_failure_leaves_settlement_reserved(self): + for operation in ['commit', 'release', 'mark_indeterminate']: + with self.subTest(operation=operation): + parent = self.reserved() + before = state(parent) + original = AUTH.digest + def fail(value): + if isinstance(value, dict) and 'record' in value: + raise ValueError('injected final receipt serializer failure') + return original(value) + AUTH.digest = fail + try: + with self.assertRaises(ValueError): + if operation == 'commit': + parent.commit('consume:1', effect_receipt_digest='sha256:effect') + else: + getattr(parent, operation)('consume:1', reason='fixture') + finally: + AUTH.digest = original + self.assertEqual(state(parent), before) + parent.commit('consume:1', effect_receipt_digest='sha256:effect') + + def test_serializer_failure_leaves_reservation_unopened(self): + parent = self.parent() + before = state(parent) + original = AUTH.digest + def fail(value): + if isinstance(value, dict) and 'record' in value: + raise ValueError('injected final receipt serializer failure') + return original(value) + AUTH.digest = fail + try: + with self.assertRaises(ValueError): + parent.reserve(consumption_ref='consume:1', args_digest='sha256:args', target_ref='path:/fixture') + finally: + AUTH.digest = original + self.assertEqual(state(parent), before) + + def test_valid_settlements_preserve_accounting_and_digest(self): + for operation, expected, charged in [('commit', 'committed', 1), ('release', 'released', 0), ('mark_indeterminate', 'indeterminate', 1)]: + with self.subTest(operation=operation): + parent = self.reserved() + if operation == 'commit': + receipt = parent.commit('consume:1', effect_receipt_digest='sha256:effect') + else: + receipt = getattr(parent, operation)('consume:1', reason='fixture') + self.assertEqual(receipt['record']['state'], expected) + self.assertEqual(parent.settlement('consume:1'), receipt['record']) + self.assertEqual(receipt['open_reservations'], 0) + self.assertEqual(receipt['charged_total'], charged) + self.assertEqual(receipt['consumed_total'], charged) + self.assertEqual(receipt['receipt_digest'], AUTH.digest({k:v for k,v in receipt.items() if k != 'receipt_digest'})) + before = state(parent) + with self.assertRaises(AUTH.ContractError): + parent.commit('consume:1', effect_receipt_digest='sha256:other') + self.assertEqual(state(parent), before) + + def test_duplicate_and_unknown_refs_are_refused(self): + parent = self.reserved() + before = state(parent) + with self.assertRaises(AUTH.ContractError): + parent.reserve(consumption_ref='consume:1', args_digest='sha256:args', target_ref='path:/fixture') + with self.assertRaises(AUTH.ContractError): + parent.commit('consume:unknown', effect_receipt_digest='sha256:effect') + self.assertEqual(state(parent), before) + + def test_pr125_temporal_subset_controls(self): + whole = '2026-10-04T00:00:00Z' + fraction = '2026-10-04T00:00:00.100000Z' + for bound, parent, child, contained in [ + ('not_after', whole, fraction, False), ('not_after', fraction, whole, True), + ('not_before', fraction, whole, False), ('not_before', whole, fraction, True), + ('not_after', whole, '2026-10-04T01:00:00.000000+01:00', True), + ('not_before', fraction, '2026-10-04T01:00:00.100000+01:00', True), + ]: + with self.subTest(bound=bound, parent=parent, child=child): + self.assertEqual(AUTH.is_subset_scope({'time': {bound: child}}, {'time': {bound: parent}}), contained) + + def test_pr125_wire_fingerprint_control(self): + scope = {'time': {'not_after': '2026-10-04T00:00:00Z'}} + equivalent = {'time': {'not_after': '2026-10-04T01:00:00.000000+01:00'}} + self.assertEqual(AUTH.normalize_scope(scope), scope) + self.assertEqual(AUTH.scope_fingerprint(scope), AUTH.scope_fingerprint(equivalent)) diff --git a/tests/ares_runtime/test_authority_fractional_time.py b/tests/ares_runtime/test_authority_fractional_time.py new file mode 100644 index 0000000000000..f602281a38981 --- /dev/null +++ b/tests/ares_runtime/test_authority_fractional_time.py @@ -0,0 +1,84 @@ +"""Temporal attenuation compares UTC instants without changing wire digests.""" + +import pytest + +from ares_runtime.authority import ( + AuthorityScopeV1, + ContractError, + is_subset_scope, + normalize_scope, + scope_fingerprint, +) + + +WHOLE = "2026-10-04T00:00:00Z" +FRACTION = "2026-10-04T00:00:00.100000Z" + + +@pytest.mark.parametrize( + "start,end", + [ + (WHOLE, FRACTION), + ("2026-10-04T01:00:00+01:00", FRACTION), + (FRACTION, "2026-10-04T00:00:00.200000Z"), + (WHOLE, "2026-10-04T00:00:00.000000Z"), + ], +) +def test_interval_accepts_ordered_same_second_instants(start, end): + scope = normalize_scope({"time": {"not_before": start, "not_after": end}}) + assert set(scope["time"]) == {"not_before", "not_after"} + + +@pytest.mark.parametrize( + "start,end", + [ + (FRACTION, WHOLE), + (FRACTION, "2026-10-04T01:00:00+01:00"), + ("2026-10-04T00:00:00.200000Z", FRACTION), + ], +) +def test_interval_rejects_reversed_same_second_instants(start, end): + with pytest.raises(ContractError, match="INVALID_TIME_SCOPE"): + normalize_scope({"time": {"not_before": start, "not_after": end}}) + + +@pytest.mark.parametrize( + "bound,parent,child,contained", + [ + ("not_after", WHOLE, FRACTION, False), + ("not_after", FRACTION, WHOLE, True), + ("not_before", FRACTION, WHOLE, False), + ("not_before", WHOLE, FRACTION, True), + ("not_after", WHOLE, "2026-10-04T01:00:00.000000+01:00", True), + ("not_before", FRACTION, "2026-10-04T01:00:00.100000+01:00", True), + ], +) +def test_same_second_subset_respects_each_bound(bound, parent, child, contained): + assert is_subset_scope( + {"time": {bound: child}}, {"time": {bound: parent}} + ) is contained + + +@pytest.mark.parametrize( + "bound,parent_time,child_time", + [("not_after", WHOLE, FRACTION), ("not_before", FRACTION, WHOLE)], +) +def test_rejected_temporal_widening_does_not_spend_delegation( + bound, parent_time, child_time +): + parent = AuthorityScopeV1( + scope={"tool": "read_file", "time": {bound: parent_time}, "use_count": 1}, + generation=1, + ) + with pytest.raises(ContractError, match="ATTENUATION_ESCALATION"): + parent.attenuate({"time": {bound: child_time}}, child_generation=2) + child = parent.attenuate({"time": {bound: parent_time}}, child_generation=2) + assert child.subset_witness(parent)["contained"] is True + + +def test_existing_canonical_timestamp_and_fingerprint_are_preserved(): + scope = {"time": {"not_after": WHOLE}} + equivalent = {"time": {"not_after": "2026-10-04T01:00:00.000000+01:00"}} + assert normalize_scope(scope) == scope + assert normalize_scope(equivalent) == scope + assert scope_fingerprint(equivalent) == scope_fingerprint(scope) diff --git a/tests/ares_runtime/test_closure_projection_rebuild.py b/tests/ares_runtime/test_closure_projection_rebuild.py new file mode 100644 index 0000000000000..41045a3623b37 --- /dev/null +++ b/tests/ares_runtime/test_closure_projection_rebuild.py @@ -0,0 +1,92 @@ +"""Rebuilding closed evidence is idempotent; reopening still needs evidence.""" + +import pytest + +from ares_runtime.collaboration import ClosureProjector, ContractError + + +def project(projector, *, gates=None, events=("event:close",), flags=(), prior=None): + return projector.project( + "mission:rebuild", + "engineering", + {"test": True} if gates is None else gates, + source_event_refs=events, + source_event_exists=lambda _ref: True, + flags=flags, + previous_projection=prior, + ) + + +@pytest.mark.parametrize("prior_form", ["artifact", "mapping"]) +def test_identical_closed_projection_rebuilds_same_artifact(prior_form): + projector = ClosureProjector() + closed = project(projector) + prior = closed if prior_form == "artifact" else closed.to_dict() + rebuilt = project(projector, prior=prior) + assert rebuilt.canonical_bytes() == closed.canonical_bytes() + assert rebuilt.artifact_digest == closed.artifact_digest + + +def test_rebuild_normalizes_duplicate_event_order_without_new_evidence(): + projector = ClosureProjector() + closed = project(projector, events=("event:a", "event:b")) + rebuilt = project(projector, events=("event:b", "event:a", "event:a"), prior=closed) + assert rebuilt.canonical_bytes() == closed.canonical_bytes() + + +@pytest.mark.parametrize( + "gates,flags", + [ + ({"test": False}, ()), + ({"test": True}, ("LEDGER_AHEAD_OF_UI",)), + ({"test": True}, ("AMBIGUOUS_EFFECT",)), + ], +) +def test_departure_from_closed_state_without_new_evidence_remains_rejected(gates, flags): + projector = ClosureProjector() + closed = project(projector) + with pytest.raises(ContractError, match="REOPEN_REQUIRES_NEW_EVIDENCE"): + project(projector, gates=gates, flags=flags, prior=closed) + + +@pytest.mark.parametrize("flags,state", [((), "evidence_pending"), (("AMBIGUOUS_EFFECT",), "quarantined")]) +def test_new_source_event_allows_derived_reopening(flags, state): + projector = ClosureProjector() + closed = project(projector) + reopened = project( + projector, + gates={"test": False}, + events=("event:close", "event:failure"), + flags=flags, + prior=closed, + ) + assert reopened.to_dict()["state"] == state + + +def test_same_evidence_closed_rebuild_still_requires_matching_lineage(): + projector = ClosureProjector() + closed = project(projector).to_dict() + closed["mission_ref"] = "mission:other" + with pytest.raises(ContractError, match="PROJECTION_LINEAGE_MISMATCH"): + project(projector, prior=closed) + + +@pytest.mark.parametrize("change", ["gate", "event"]) +def test_changed_closed_projection_without_new_evidence_is_still_rejected(change): + projector = ClosureProjector() + closed = project(projector, events=("event:a", "event:b")) + gates = {"test": True, "other_check": True} if change == "gate" else None + events = ("event:a",) if change == "event" else ("event:a", "event:b") + with pytest.raises(ContractError, match="REOPEN_REQUIRES_NEW_EVIDENCE"): + project(projector, gates=gates, events=events, prior=closed) + + +def test_identical_rebuild_rechecks_source_event_existence(): + projector = ClosureProjector() + closed = project(projector) + with pytest.raises(ContractError, match="MISSING_SOURCE_EVENT"): + projector.project( + "mission:rebuild", "engineering", {"test": True}, + source_event_refs=["event:close"], source_event_exists=lambda _ref: False, + previous_projection=closed, + ) diff --git a/tests/ares_runtime/test_continuity_input.py b/tests/ares_runtime/test_continuity_input.py index 5fc1fe7485e4a..61912f378f7ce 100644 --- a/tests/ares_runtime/test_continuity_input.py +++ b/tests/ares_runtime/test_continuity_input.py @@ -186,6 +186,7 @@ def test_tui_busy_ack_is_durable_and_preserves_distinct_occurrences(db, monkeypa session = _session(agent=SimpleNamespace(_session_db=db, session_id="s", platform="cli", context_rebase_enabled=True), session_key="s", running=True) session["inflight_turn"] = {"user": "Identical words"} + monkeypatch.setitem(server._sessions, "s", session) monkeypatch.setattr(server, "_sess_nowait", lambda *a: (session, None)) monkeypatch.setattr(server, "_ensure_active_session_slot", lambda *a: None) monkeypatch.setattr(server, "_load_dashboard_process_isolation_config", lambda: {}) @@ -229,6 +230,7 @@ def test_tui_rejected_submit_never_enters_durable_inbox(db, monkeypatch, refusal session = _session(agent=SimpleNamespace(_session_db=db, session_id="s", platform="cli", context_rebase_enabled=True), session_key="s", running=refusal == "busy_confirm") session["lazy"] = refusal == "watch" + monkeypatch.setitem(server._sessions, "s", session) monkeypatch.setattr(server, "_sess_nowait", lambda *a: (session, None)) monkeypatch.setattr(server, "_ensure_active_session_slot", lambda *a: None) monkeypatch.setattr(server, "_load_dashboard_process_isolation_config", lambda: {}) diff --git a/tests/ares_runtime/test_final_context_payload_semantics.py b/tests/ares_runtime/test_final_context_payload_semantics.py new file mode 100644 index 0000000000000..aca62b59f6817 --- /dev/null +++ b/tests/ares_runtime/test_final_context_payload_semantics.py @@ -0,0 +1,214 @@ +"""Final qualification counts schema data without accepting opaque transport input.""" +from copy import deepcopy +import hashlib +import json +from types import SimpleNamespace + +import pytest + +from ares_runtime.continuity.budget import BudgetError, CountMethod, final_request_upper_bound +from ares_runtime.continuity.runtime import ( + ContextDispatchError, admit_final_context_dispatch, context_dispatch_payload_digest, + context_dispatch_route_identity, +) +from hermes_state import SessionDB + + +NAMES = ("conversation", "previous_response_id", "encrypted_content") +MODES = ("chat_completions", "codex_responses", "anthropic_messages", "bedrock_converse") + + +def schema(name): + return {"type": "object", "properties": {name: {"type": "string"}}, "required": [name]} + + +def payload(mode, name="query", definition=None): + definition = schema(name) if definition is None else definition + fn = {"name": "search_history", "description": "Plain text lookup", "parameters": definition} + if mode == "codex_responses": + return {"model": "test-model", "max_output_tokens": 1000, + "input": [{"role": "user", "content": "hi"}], + "tools": [{"type": "function", **fn}]} + if mode == "anthropic_messages": + return {"model": "test-model", "max_tokens": 1000, + "messages": [{"role": "user", "content": "hi"}], + "tools": [{"name": fn["name"], "description": fn["description"], + "input_schema": definition}]} + if mode == "bedrock_converse": + return {"modelId": "test-model", "inferenceConfig": {"maxTokens": 1000}, + "messages": [{"role": "user", "content": [{"text": "hi"}]}], + "toolConfig": {"tools": [{"toolSpec": { + "name": fn["name"], "description": fn["description"], + "inputSchema": {"json": definition}}}]}} + return {"model": "test-model", "max_tokens": 1000, + "messages": [{"role": "user", "content": "hi"}], + "tools": [{"type": "function", "function": fn}]} + + +@pytest.mark.parametrize("mode", MODES) +@pytest.mark.parametrize("name", NAMES) +def test_valid_schema_parameter_names_count_every_byte(mode, name): + body = payload(mode, name) + result = final_request_upper_bound(route_ref="test:route", payload=body) + raw = json.dumps(body, ensure_ascii=False, sort_keys=True, + separators=(",", ":"), allow_nan=False).encode() + assert result.tokens == len(raw) + 8192 + 64 + assert result.method is CountMethod.QUALIFIED_UPPER_BOUND + assert result.payload_digest == "sha256:" + hashlib.sha256(raw).hexdigest() + + +@pytest.mark.parametrize("mode", MODES) +def test_nested_schema_types_enums_defaults_and_examples_are_data(mode): + definition = {"type": "object", "properties": { + "type": {"type": ["string", "null"]}, + "nested": {"type": "object", "properties": { + "conversation": {"type": "string", "description": "previous_response_id"}, + "previous_response_id": {"enum": [{"encrypted_content": "example"}]}, + "encrypted_content": {"default": {"type": "input_image", "image_url": "example"}}, + }}, + }, "examples": [{"type": "compaction", "encrypted_content": "example"}]} + original = final_request_upper_bound(route_ref="r", payload=payload(mode, definition=definition)) + changed = deepcopy(definition) + changed["description"] = "Every added schema byte remains counted." + newer = final_request_upper_bound(route_ref="r", payload=payload(mode, definition=changed)) + assert newer.tokens > original.tokens + assert newer.payload_digest != original.payload_digest + + +@pytest.mark.parametrize("name", NAMES) +@pytest.mark.parametrize("location", ["root", "extension", "tool_sibling", "schema_lookalike"]) +def test_opaque_names_outside_genuine_schema_locations_still_refuse(name, location): + body = payload("chat_completions") + if location == "root": + body[name] = "server-state" + elif location == "extension": + body["extra_body"] = {"vendor_extension": {name: "server-state"}} + elif location == "tool_sibling": + body["tools"][0]["function"][name] = "server-state" + else: + body["extension"] = {"tools": [{"type": "function", "function": { + "name": "fake", "parameters": schema(name)}}]} + with pytest.raises(BudgetError, match="FINAL_PAYLOAD_OPAQUE_ACCOUNTING_UNQUALIFIED"): + final_request_upper_bound(route_ref="r", payload=body) + + +@pytest.mark.parametrize("mode", MODES) +def test_media_sibling_of_schema_still_refuses(mode): + body = payload(mode) + if mode == "bedrock_converse": + body["toolConfig"]["tools"][0]["toolSpec"]["inputSchema"]["opaque"] = { + "type": "input_image"} + else: + body["tools"][0]["opaque"] = {"type": "input_image"} + with pytest.raises(BudgetError, match="FINAL_PAYLOAD_OPAQUE_ACCOUNTING_UNQUALIFIED"): + final_request_upper_bound(route_ref="r", payload=body) + + +@pytest.mark.parametrize("kind", [ + "image", "image_url", "input_image", "input_file", "file", "audio", + "input_audio", "compaction", "computer_screenshot", +]) +def test_actual_media_and_native_items_still_refuse(kind): + body = payload("codex_responses") + body["input"] = [{"type": kind}] + with pytest.raises(BudgetError, match="FINAL_PAYLOAD_OPAQUE_ACCOUNTING_UNQUALIFIED"): + final_request_upper_bound(route_ref="r", payload=body) + + +@pytest.mark.parametrize("malformed", [object(), float("nan")]) +def test_schema_data_does_not_skip_serialization_validation(malformed): + with pytest.raises(BudgetError, match="INVALID_FINAL_PAYLOAD"): + final_request_upper_bound(route_ref="r", + payload=payload("chat_completions", definition={"default": malformed})) + + +def test_schema_data_does_not_skip_depth_validation(): + definition = {} + for _ in range(65): + definition = {"properties": {"x": definition}} + with pytest.raises(BudgetError, match="INVALID_FINAL_PAYLOAD"): + final_request_upper_bound(route_ref="r", payload=payload("chat_completions", definition=definition)) + + +def test_schema_lookalike_with_wrong_tool_kind_is_not_exempt(): + body = payload("codex_responses", "conversation") + body["tools"][0]["type"] = "unknown-server-tool" + with pytest.raises(BudgetError, match="FINAL_PAYLOAD_OPAQUE_ACCOUNTING_UNQUALIFIED"): + final_request_upper_bound(route_ref="r", payload=body) + + +def test_schema_alias_at_a_content_location_is_still_opaque(): + shared = {"type": "object", "properties": {"conversation": {"type": "string"}}} + body = payload("codex_responses", definition=shared) + body["input"] = [shared] + with pytest.raises(BudgetError, match="FINAL_PAYLOAD_OPAQUE_ACCOUNTING_UNQUALIFIED"): + final_request_upper_bound(route_ref="r", payload=body) + + +@pytest.mark.parametrize("kind", [["input_image"], {"input_image": True}]) +def test_malformed_transport_discriminator_refuses_with_typed_error(kind): + body = payload("codex_responses") + body["input"] = [{"type": kind}] + with pytest.raises(BudgetError, match="INVALID_FINAL_PAYLOAD"): + final_request_upper_bound(route_ref="r", payload=body) + + +@pytest.mark.parametrize("kind", [7, True]) +def test_scalar_extension_type_metadata_retains_prior_counting(kind): + body = payload("chat_completions") + body["extra_body"] = {"vendor_metadata": {"type": kind}} + assert final_request_upper_bound(route_ref="r", payload=body).tokens > 8192 + + +@pytest.fixture +def bound_agent(tmp_path): + db = SessionDB(db_path=tmp_path / "context.db") + db.create_session("s", source="cli") + db.append_message("s", "user", "hi") + assert db.try_acquire_session_turn_lease("s", "holder", ttl_seconds=300) + agent = SimpleNamespace(provider="test", model="test-model", api_mode="chat_completions", + base_url="", client=object(), max_tokens=1000, + context_compressor=SimpleNamespace(context_length=100_000), + _session_db=db, session_id="s", _active_session_turn_lease_holder="holder") + yield agent + db.close() + + +@pytest.mark.parametrize("mode", MODES) +@pytest.mark.parametrize("name", NAMES) +def test_real_final_admission_and_materialization_share_schema_count(bound_agent, mode, name): + bound_agent.api_mode = mode + body = payload(mode, name) + snapshot = bound_agent._session_db.read_context_rebase_snapshot("s") + digest = context_dispatch_payload_digest(body) + admission = admit_final_context_dispatch(bound_agent, snapshot, body, + attempt_id="attempt", materialization_digest=digest, + route_identity=context_dispatch_route_identity(bound_agent)) + assert admission["payload_digest"] == digest + + +@pytest.mark.parametrize("changed,code", [ + ({"previous_response_id": "opaque"}, "FINAL_PAYLOAD_OPAQUE_ACCOUNTING_UNQUALIFIED"), + ({"extra_body": {"tools": []}}, "CONTEXT_DISPATCH_WIRE_OVERRIDE_UNQUALIFIED"), + ({"max_tokens": 0}, "CONTEXT_DISPATCH_OUTPUT_BUDGET_UNQUALIFIED"), + ({"model": "other"}, "CONTEXT_DISPATCH_ROUTE_CHANGED"), +]) +def test_public_admission_retains_negative_controls(bound_agent, changed, code): + body = payload("chat_completions") + body.update(changed) + snapshot = bound_agent._session_db.read_context_rebase_snapshot("s") + with pytest.raises(ContextDispatchError, match=code): + admit_final_context_dispatch(bound_agent, snapshot, body, attempt_id="refused") + + +def test_public_admission_preserves_overflow_digest_and_no_snapshot_modes(bound_agent): + body = payload("chat_completions") + snapshot = bound_agent._session_db.read_context_rebase_snapshot("s") + with pytest.raises(ContextDispatchError, match="CONTEXT_DISPATCH_MATERIALIZATION_CHANGED"): + admit_final_context_dispatch(bound_agent, snapshot, body, attempt_id="digest", + materialization_digest="sha256:" + "0" * 64) + bound_agent.context_compressor.context_length = 9000 + with pytest.raises(ContextDispatchError, match="CONTEXT_DISPATCH_FINAL_PAYLOAD_TOO_LARGE"): + admit_final_context_dispatch(bound_agent, snapshot, body, attempt_id="overflow") + assert admit_final_context_dispatch(bound_agent, None, {"previous_response_id": "opaque"}, + attempt_id="ordinary") is None diff --git a/tests/ares_runtime/test_local_runtime.py b/tests/ares_runtime/test_local_runtime.py index 2f44a94605ee8..9a5265b2bc8dd 100644 --- a/tests/ares_runtime/test_local_runtime.py +++ b/tests/ares_runtime/test_local_runtime.py @@ -32,6 +32,176 @@ def _runtime(tmp_path: Path) -> AresLocalRuntime: ) +@pytest.mark.parametrize("failure", ["restart", "health", "recovery_restart"]) +def test_rollback_service_failure_restores_the_exact_pointer_pair( + tmp_path: Path, monkeypatch, failure: str +) -> None: + runtime = _runtime(tmp_path) + previous_source = _release(runtime, "a" * 40) + current_source = _release(runtime, "b" * 40) + runtime._activate("a" * 40) + runtime._activate("b" * 40) + runtime.paths.unit_path.parent.mkdir(parents=True) + runtime.paths.unit_path.write_text("inert service fixture", encoding="utf-8") + calls: list[tuple[tuple[str, ...], bool]] = [] + + def systemctl(*args: str, required: bool = True) -> bool: + calls.append((args, required)) + if args[0] == "restart": + if required and failure != "health": + raise AresLocalRuntimeError("injected rollback restart failure") + if not required and failure == "recovery_restart": + raise OSError("injected recovery restart failure") + return not (failure == "health" and args[0] == "is-active") + + monkeypatch.setattr(runtime, "_systemctl", systemctl) + monkeypatch.setattr("ares_runtime.local_runtime.time.sleep", lambda _seconds: None) + error = ( + "did not remain active" if failure == "health" else "rollback restart failure" + ) + with pytest.raises(AresLocalRuntimeError, match=error): + runtime.rollback() + + assert runtime.active_release() == ("b" * 40, current_source.resolve()) + assert runtime.previous_release() == ("a" * 40, previous_source.resolve()) + assert calls[-1] == (("restart", "ares-gateway.service"), False) + + +@pytest.mark.parametrize("gateway", [False, True]) +def test_rollback_success_selects_previous_with_and_without_gateway( + tmp_path: Path, monkeypatch, gateway: bool +) -> None: + runtime = _runtime(tmp_path) + previous_source = _release(runtime, "a" * 40) + current_source = _release(runtime, "b" * 40) + runtime._activate("a" * 40) + runtime._activate("b" * 40) + if gateway: + runtime.paths.unit_path.parent.mkdir(parents=True) + runtime.paths.unit_path.write_text("inert service fixture", encoding="utf-8") + calls: list[tuple[str, ...]] = [] + monkeypatch.setattr( + runtime, "_systemctl", lambda *args, **_kw: calls.append(args) or True + ) + monkeypatch.setattr("ares_runtime.local_runtime.time.sleep", lambda _seconds: None) + + assert runtime.rollback() == "a" * 40 + assert runtime.active_release() == ("a" * 40, previous_source.resolve()) + assert runtime.previous_release() == ("b" * 40, current_source.resolve()) + assert bool(calls) is gateway + + +def test_rollback_compensation_failure_is_reported_without_recovery_claim( + tmp_path: Path, monkeypatch +) -> None: + runtime = _runtime(tmp_path) + _release(runtime, "a" * 40) + _release(runtime, "b" * 40) + runtime._activate("a" * 40) + runtime._activate("b" * 40) + runtime.paths.unit_path.parent.mkdir(parents=True) + runtime.paths.unit_path.write_text("inert service fixture", encoding="utf-8") + monkeypatch.setattr( + runtime, + "_systemctl", + lambda *_args, **_kw: (_ for _ in ()).throw( + AresLocalRuntimeError("injected rollback restart failure") + ), + ) + monkeypatch.setattr( + runtime, + "_restore_release_pair", + lambda *_args: (_ for _ in ()).throw(OSError("injected restore failure")), + ) + + with pytest.raises( + AresLocalRuntimeError, match="prior release pointers could not be restored" + ) as error: + runtime.rollback() + assert isinstance(error.value.__cause__, OSError) + assert str(error.value.__cause__) == "injected restore failure" + + +@pytest.mark.parametrize("cleanup_failure", [False, True]) +def test_upstream_candidate_reuse_removes_owned_staging( + tmp_path: Path, monkeypatch, cleanup_failure: bool +) -> None: + runtime = _runtime(tmp_path) + runtime._ensure_layout() + upstream, downstream = "c" * 40, "d" * 40 + installed = _release(runtime, upstream) + python = runtime._python_for(installed) + python.parent.mkdir(parents=True) + python.write_text("inert interpreter fixture", encoding="utf-8") + runtime._atomic_json( + runtime._release_dir(upstream) / "release.json", + {"upstream_revision": upstream, "downstream_revision": downstream}, + ) + prior_source = _release(runtime, "a" * 40) + runtime._activate("a" * 40) + runtime._activate(upstream) + marker = installed / "immutable-marker" + marker.write_bytes(b"existing immutable release") + + def run(args, **_kwargs): + if list(args[:2]) == ["git", "clone"]: + staging_source = Path(args[-1]) + staging_source.mkdir(parents=True) + (staging_source / "clone-fixture").write_bytes( + b"owned tiny staging fixture" + ) + return SimpleNamespace(returncode=0, stdout="", stderr="") + + monkeypatch.setattr(runtime, "_run", run) + monkeypatch.setattr( + runtime, + "_git_output", + lambda _source, *args: ( + upstream if args == ("rev-parse", "FETCH_HEAD") else "0" * 40 + ), + ) + monkeypatch.setattr( + subprocess, + "run", + lambda *_args, **_kw: SimpleNamespace(returncode=0, stdout="", stderr=""), + ) + monkeypatch.setattr( + runtime, + "_build_runtime", + lambda *_args, **_kw: (_ for _ in ()).throw( + AssertionError("existing release must never rebuild") + ), + ) + if cleanup_failure: + monkeypatch.setattr( + "ares_runtime.local_runtime.shutil.rmtree", + lambda path: (_ for _ in ()).throw(OSError("injected cleanup failure")), + ) + + def reuse(): + return runtime._materialize_upstream_candidate( + downstream_remote="inert-downstream", + downstream_revision=downstream, + upstream_remote="inert-upstream", + upstream_branch="main", + upstream_revision=upstream, + desktop=False, + ) + + if cleanup_failure: + with pytest.raises( + AresLocalRuntimeError, match="upstream candidate cleanup failed" + ): + reuse() + else: + assert reuse() == upstream + assert reuse() == upstream + assert not list(runtime.paths.staging_dir.iterdir()) + assert marker.read_bytes() == b"existing immutable release" + assert runtime.active_release() == (upstream, installed.resolve()) + assert runtime.previous_release() == ("a" * 40, prior_source.resolve()) + + def _release(runtime: AresLocalRuntime, revision: str) -> Path: source = runtime.paths.releases_dir / revision / "source" source.mkdir(parents=True) @@ -61,6 +231,225 @@ def _repository(path: Path) -> Path: return path +@pytest.fixture +def gateway_stop_case(tmp_path: Path, monkeypatch): + from gateway import status + + runtime = _runtime(tmp_path) + source = _release(runtime, "a" * 40) + python = source / ".venv" / "bin" / "python" + python.parent.mkdir(parents=True) + python.touch() + runtime._activate("a" * 40) + home = runtime.paths.agent_home + home.mkdir() + pid, start = 987654, 1200 + record = { + "pid": pid, + "start_time": start, + "kind": "hermes-gateway", + "argv": ["python", "-m", "hermes_cli.main", "gateway", "run"], + "hermes_home": str(home), + } + for name in ("gateway.pid", "gateway.lock"): + (home / name).write_text(json.dumps(record)) + monkeypatch.setenv("HERMES_HOME", str(home)) + monkeypatch.setattr(status, "_is_gateway_runtime_lock_active_strict", lambda _path: True) + monkeypatch.setattr(status, "_pid_exists", lambda candidate: candidate == pid) + monkeypatch.setattr(status, "_get_process_start_time", lambda _pid: start) + monkeypatch.setattr(status, "_read_process_cmdline", lambda _pid: " ".join(record["argv"])) + return SimpleNamespace(runtime=runtime, source=source, home=home, + status=status, pid=pid, start=start, record=record) + + +def _gateway_stop_child(case, monkeypatch, events): + import ares_runtime.local_runtime as local_runtime + + def run(command, **kwargs): + assert command[:2] == [str(case.source / ".venv/bin/python"), "-c"] + assert kwargs["cwd"] == case.source + assert kwargs["env"]["HERMES_HOME"] == str(case.home.resolve()) + assert kwargs["timeout"] == 10 + events.append("prepare") + # Emulate the child's process environment without launching a process. + with monkeypatch.context() as child: + child.setenv("HERMES_HOME", kwargs["env"]["HERMES_HOME"]) + try: + exec(command[-1], {}) + except AresLocalRuntimeError as exc: + return SimpleNamespace(returncode=1, stdout="", stderr=str(exc)) + return SimpleNamespace(returncode=0, stdout="", stderr="") + + monkeypatch.setattr(local_runtime.subprocess, "run", run) + monkeypatch.setattr(case.runtime, "_systemctl", + lambda *args, **kwargs: events.append(args) or True) + + +def test_gateway_stop_publishes_verified_marker_before_service_stop(gateway_stop_case, monkeypatch): + case = gateway_stop_case + events = [] + _gateway_stop_child(case, monkeypatch, events) + case.runtime.gateway("stop") + + assert events == ["prepare", ("disable", "--now", "ares-gateway.service")] + marker = json.loads((case.home / ".gateway-planned-stop.json").read_text()) + assert (marker["target_pid"], marker["target_start_time"]) == (case.pid, case.start) + with monkeypatch.context() as consumer: + consumer.setattr(case.status.os, "getpid", lambda: case.pid) + assert case.status.consume_planned_stop_marker_for_self() is True + + +def test_gateway_stop_keeps_isolated_unit_routing(gateway_stop_case, monkeypatch): + from dataclasses import replace + import ares_runtime.local_runtime as local_runtime + + case = gateway_stop_case + case.runtime = AresLocalRuntime(replace(case.runtime.paths, unit_path=case.home / "offline-stop.service")) + events = [] + _gateway_stop_child(case, monkeypatch, events) + child_run = local_runtime.subprocess.run + monkeypatch.setattr(case.runtime, "_systemctl", AresLocalRuntime._systemctl.__get__(case.runtime)) + monkeypatch.setattr(local_runtime.shutil, "which", lambda _name: "/usr/bin/systemctl") + + def run(command, **kwargs): + if command[0] != "systemctl": + return child_run(command, **kwargs) + assert (case.home / ".gateway-planned-stop.json").is_file() + events.append(command) + return SimpleNamespace(returncode=0, stdout="", stderr="") + + monkeypatch.setattr(local_runtime.subprocess, "run", run) + case.runtime.gateway("stop") + assert events == ["prepare", ["systemctl", "--user", "disable", "--now", "offline-stop.service"]] + + +@pytest.mark.parametrize("boundary", ["stale_start", "dead_pid", "wrong_pid", "wrong_role"]) +def test_gateway_stop_rejects_stale_or_disagreeing_identity(gateway_stop_case, monkeypatch, boundary): + case = gateway_stop_case + if boundary == "stale_start": + monkeypatch.setattr(case.status, "_get_process_start_time", lambda _pid: case.start + 1) + elif boundary == "dead_pid": + monkeypatch.setattr(case.status, "_pid_exists", lambda _pid: False) + elif boundary == "wrong_pid": + record = dict(case.record, pid=case.pid + 1) + (case.home / "gateway.lock").write_text(json.dumps(record)) + else: + monkeypatch.setattr(case.status, "_read_process_cmdline", lambda _pid: "python -m hermes_cli.main serve") + events = [] + _gateway_stop_child(case, monkeypatch, events) + with pytest.raises(AresLocalRuntimeError, match="identity"): + case.runtime.gateway("stop") + assert events == ["prepare"] + assert not (case.home / ".gateway-planned-stop.json").exists() + + +@pytest.mark.parametrize("record_name", ["gateway.pid", "gateway.lock", "live_profile"]) +def test_gateway_stop_rejects_foreign_home(gateway_stop_case, monkeypatch, tmp_path, record_name): + case = gateway_stop_case + if record_name == "live_profile": + monkeypatch.setattr(case.status, "_read_process_cmdline", lambda _pid: "python -m hermes_cli.main --profile other gateway run") + else: + record = dict(case.record, hermes_home=str(tmp_path / "foreign-home")) + (case.home / record_name).write_text(json.dumps(record)) + events = [] + _gateway_stop_child(case, monkeypatch, events) + with pytest.raises(AresLocalRuntimeError, match="home"): + case.runtime.gateway("stop") + assert events == ["prepare"] + assert not (case.home / ".gateway-planned-stop.json").exists() + + +def test_gateway_stop_rejects_identity_change_before_publication(gateway_stop_case, monkeypatch): + case = gateway_stop_case + identities = iter(((case.pid, case.start), (case.pid, case.start + 1))) + monkeypatch.setattr(case.status, "get_running_pid_identity_strict", lambda _path: next(identities)) + events = [] + _gateway_stop_child(case, monkeypatch, events) + with pytest.raises(AresLocalRuntimeError, match="identity"): + case.runtime.gateway("stop") + assert events == ["prepare"] + assert not (case.home / ".gateway-planned-stop.json").exists() + + +def test_gateway_stop_aborts_when_marker_publication_fails(gateway_stop_case, monkeypatch): + case = gateway_stop_case + monkeypatch.setattr(case.status, "write_planned_stop_marker", lambda _pid: False) + events = [] + _gateway_stop_child(case, monkeypatch, events) + with pytest.raises(AresLocalRuntimeError, match="marker"): + case.runtime.gateway("stop") + assert events == ["prepare"] + + +def test_gateway_stop_without_a_runtime_owner_does_not_create_marker(gateway_stop_case, monkeypatch): + case = gateway_stop_case + (case.home / "gateway.pid").unlink() + (case.home / "gateway.lock").unlink() + events = [] + _gateway_stop_child(case, monkeypatch, events) + case.runtime.gateway("stop") + assert events == ["prepare", ("disable", "--now", "ares-gateway.service")] + assert not (case.home / ".gateway-planned-stop.json").exists() + + +def test_gateway_stop_preserves_caller_environment_and_profile_context(gateway_stop_case, monkeypatch, tmp_path): + from hermes_constants import get_hermes_home, reset_hermes_home_override, set_hermes_home_override + + case = gateway_stop_case + ambient_home = tmp_path / "ambient-home" + monkeypatch.setenv("HERMES_HOME", str(ambient_home)) + monkeypatch.setenv("PYTHONPATH", "ambient-python-path") + profile = tmp_path / "context-profile" + token = set_hermes_home_override(profile) + before = dict(os.environ) + try: + events = [] + _gateway_stop_child(case, monkeypatch, events) + case.runtime.gateway("stop") + assert dict(os.environ) == before + assert get_hermes_home() == profile + assert (case.home / ".gateway-planned-stop.json").is_file() + assert not (ambient_home / ".gateway-planned-stop.json").exists() + assert not (profile / ".gateway-planned-stop.json").exists() + finally: + reset_hermes_home_override(token) + + +@pytest.mark.parametrize("failure", [OSError("synthetic launch failure"), subprocess.TimeoutExpired("synthetic", 10)]) +def test_gateway_stop_preparation_failure_never_calls_service(gateway_stop_case, monkeypatch, failure): + case = gateway_stop_case + before = dict(os.environ) + monkeypatch.setattr("ares_runtime.local_runtime.subprocess.run", + lambda *args, **kwargs: (_ for _ in ()).throw(failure)) + calls = [] + monkeypatch.setattr(case.runtime, "_systemctl", lambda *args, **kwargs: calls.append(args)) + with pytest.raises(AresLocalRuntimeError, match="stop"): + case.runtime.gateway("stop") + assert calls == [] + assert dict(os.environ) == before + + +@pytest.mark.parametrize("mismatch", ["stale", "pid", "start", "home"]) +def test_gateway_stop_marker_owner_rejects_stale_or_wrong_consumer(gateway_stop_case, monkeypatch, tmp_path, mismatch): + import ares_runtime.local_runtime as local_runtime + + case = gateway_stop_case + local_runtime._prepare_gateway_stop_marker() + path = case.home / ".gateway-planned-stop.json" + marker = json.loads(path.read_text()) + if mismatch == "stale": + marker["written_at"] = "2000-01-01T00:00:00+00:00" + elif mismatch == "pid": + marker["target_pid"] += 1 + elif mismatch == "start": + marker["target_start_time"] += 1 + else: + marker["target_hermes_home"] = str(tmp_path / "foreign-consumer") + path.write_text(json.dumps(marker)) + monkeypatch.setattr(case.status.os, "getpid", lambda: case.pid) + assert case.status.consume_planned_stop_marker_for_self() is False + + def test_current_link_is_the_only_active_runtime_pointer(tmp_path: Path) -> None: runtime = _runtime(tmp_path) first = "a" * 40 @@ -261,6 +650,152 @@ def test_update_activates_only_the_verified_upstream_candidate( assert runtime.update(desktop=False) == (candidate_revision, False) +@pytest.fixture +def recipe_update_case(tmp_path: Path): + runtime = _runtime(tmp_path) + first = _release(runtime, "a" * 40) + current = _release(runtime, "b" * 40) + candidate = _release(runtime, "c" * 40) + runtime._activate("a" * 40) + runtime._activate("b" * 40) + runtime._write_config( + remote="fixture-downstream", branch="main", + upstream_remote="fixture-upstream", upstream_branch="main", + ) + home = runtime.paths.agent_home + home.mkdir() + profile = home / "profiles" / "fixture" + profile.mkdir(parents=True) + for name, data in { + "config.yaml": "context: {engine: ri-context-governor}\n", + "auth.json": '{"fixture":"inert-not-a-credential"}\n', + ".env": "FIXTURE_ONLY=unchanged\n", + }.items(): + (profile / name).write_text(data) + receipt = home / "install-receipts" / "latest.json" + receipt.parent.mkdir() + receipt.write_text('{"fixture":"retained"}\n') + + def snapshot(): + return ( + runtime.paths.current_link.readlink(), + runtime.paths.previous_link.readlink(), + runtime.paths.config_path.read_bytes(), + {p.name: p.read_bytes() for p in profile.iterdir()}, + receipt.read_bytes(), + ) + + return SimpleNamespace(runtime=runtime, current=current, candidate=candidate, + record=current / ".venv/share/ares-full-install.json", + snapshot=snapshot) + + +def _forbid_recipe_update_effects(case, monkeypatch): + calls = [] + for name in ("_read_config", "_remote_revision", "_materialize_upstream_candidate", + "_activate", "_write_config", "_install_gateway_unit", "_systemctl"): + def forbidden(*args, _name=name, **kwargs): + calls.append(_name) + raise AssertionError(f"Recipe admission must precede {_name}") + monkeypatch.setattr(case.runtime, name, forbidden) + return calls + + +@pytest.mark.parametrize("kind", ["full", "sdk_only", "malformed", "unknown", "directory", "dangling_link"]) +def test_recipe_admission_refuses_before_update_effects(recipe_update_case, monkeypatch, kind): + case = recipe_update_case + case.record.parent.mkdir(parents=True) + if kind == "directory": + case.record.mkdir() + elif kind == "dangling_link": + case.record.symlink_to("missing-recipe") + else: + records = { + "full": json.dumps({"inputs": {"recipe_version": "3", "enhancements": True, + "desktop": True, "sdk_dependencies": ["mcp==2.2.0"]}}), + "sdk_only": json.dumps({"inputs": {"recipe_version": "3", "enhancements": False, + "desktop": False, "sdk_dependencies": ["mcp==2.2.0"]}}), + "malformed": "{incomplete", + "unknown": json.dumps({"future_recipe": "unrecognized"}), + } + case.record.write_text(records[kind]) + before = case.snapshot() + calls = _forbid_recipe_update_effects(case, monkeypatch) + with pytest.raises(AresLocalRuntimeError, match="recorded installer recipe") as error: + case.runtime.update(desktop=False) + assert "full-distribution installer" in str(error.value) + assert "same revision" in str(error.value) + assert calls == [] + assert case.snapshot() == before + assert not case.runtime.paths.staging_dir.exists() + + +def test_recipe_admission_uninspectable_metadata_fails_closed(recipe_update_case, monkeypatch): + case = recipe_update_case + original = Path.lstat + def lstat(path, *args, **kwargs): + if path == case.record: + raise PermissionError("fixture metadata inspection denied") + return original(path, *args, **kwargs) + monkeypatch.setattr(Path, "lstat", lstat) + before = case.snapshot() + calls = _forbid_recipe_update_effects(case, monkeypatch) + with pytest.raises(AresLocalRuntimeError, match="recipe metadata"): + case.runtime.update(desktop=False) + assert calls == [] + assert case.snapshot() == before + + +def test_recipe_admission_installed_cli_reports_preservation(recipe_update_case, monkeypatch, capsys): + import ares_runtime.local_runtime as local_runtime + case = recipe_update_case + case.record.parent.mkdir(parents=True) + case.record.write_text(json.dumps({"inputs": {"enhancements": True}})) + before = case.snapshot() + calls = _forbid_recipe_update_effects(case, monkeypatch) + monkeypatch.setattr(local_runtime, "AresLocalRuntime", lambda: case.runtime) + with pytest.raises(SystemExit) as error: + local_runtime.main(["update", "--no-desktop"]) + assert error.value.code == 1 + assert "recorded installer recipe" in capsys.readouterr().err + assert calls == [] + assert case.snapshot() == before + + +@pytest.mark.parametrize("current_tuple", [False, True]) +def test_recipe_admission_preserves_base_update_and_noop(recipe_update_case, monkeypatch, current_tuple): + case = recipe_update_case + runtime = case.runtime + remote_calls = [] + def resolve(remote, branch): + remote_calls.append((remote, branch)) + return "d" * 40 if remote == "fixture-downstream" else "e" * 40 + monkeypatch.setattr(runtime, "_remote_revision", resolve) + effects = [] + monkeypatch.setattr(runtime, "_materialize_upstream_candidate", + lambda **kwargs: effects.append(kwargs) or "c" * 40) + monkeypatch.setattr(runtime, "_systemctl", lambda *args, **kwargs: pytest.fail("No service in base fixture")) + if current_tuple: + runtime._atomic_json(case.current.parent / "release.json", { + "downstream_revision": "d" * 40, "upstream_revision": "e" * 40, + "upstream_remote": "fixture-upstream", "upstream_branch": "main", + }) + before = case.snapshot() + result = runtime.update(desktop=False) + assert remote_calls == [("fixture-downstream", "main"), ("fixture-upstream", "main")] + if current_tuple: + assert result == ("b" * 40, False) + assert effects == [] + assert case.snapshot() == before + else: + assert result == ("c" * 40, True) + assert len(effects) == 1 + assert effects[0]["desktop"] is False + assert runtime.active_release() == ("c" * 40, case.candidate.resolve()) + assert runtime.previous_release() == ("b" * 40, case.current.resolve()) + assert case.snapshot()[2:] == before[2:] + + def test_upstream_candidate_conflict_never_publishes_a_release( tmp_path: Path, monkeypatch ) -> None: diff --git a/tests/ci/test_classify_changes.py b/tests/ci/test_classify_changes.py index 4902cbafc28b6..b59a23b015ec8 100644 --- a/tests/ci/test_classify_changes.py +++ b/tests/ci/test_classify_changes.py @@ -9,6 +9,9 @@ import importlib.util import re +import shutil +import subprocess +import sys from pathlib import Path import pytest @@ -38,10 +41,12 @@ "rust": True, "mcp_catalog": False, "ci_review": True, + "context_continuity": True, + "current_owner_integration": True, } -def _lanes(python=False, frontend=False, site=False, scan=False, deps=False, uv_lock=False, npm_lock=False, installer=False, rust=False, mcp_catalog=False, docker_meta=False, ci_review=False, python_prod=None, nix=None, docker=None) -> dict[str, bool]: +def _lanes(python=False, frontend=False, site=False, scan=False, deps=False, uv_lock=False, npm_lock=False, installer=False, rust=False, mcp_catalog=False, docker_meta=False, ci_review=False, python_prod=None, nix=None, docker=None, context_continuity=False, current_owner_integration=False) -> dict[str, bool]: # python_prod tracks python except for tests-only diffs; default it to # python so the majority of cases don't need to spell it out. # @@ -66,6 +71,8 @@ def _lanes(python=False, frontend=False, site=False, scan=False, deps=False, uv_ "rust": rust, "mcp_catalog": mcp_catalog, "ci_review": ci_review, + "context_continuity": context_continuity, + "current_owner_integration": current_owner_integration, } @@ -173,7 +180,7 @@ def _lanes(python=False, frontend=False, site=False, scan=False, deps=False, uv_ # real failures, so it keeps the conservative full lane set. "test runner script → python_prod stays on": ( ["scripts/run_tests_parallel.py"], - _lanes(python=True, scan=True), + _lanes(python=True, scan=True, context_continuity=True, current_owner_integration=True), ), # Supply-chain lanes ".pth file → scan": (["evil.pth"], _lanes(python=True, scan=True)), @@ -300,3 +307,127 @@ def test_ci_review_files_returns_only_sensitive_paths_sorted_and_unique(): ".github/workflows/ci.yml", "apps/desktop/eslint.config.mjs", ] + + +@pytest.mark.parametrize("path,continuity,owner", [ + ("ares_runtime/continuity/runtime.py", True, False), + ("tests/ares_runtime/test_continuity_rollover.py", True, False), + ("tests/ares_runtime/test_context_native_paired.py", True, False), + ("plugins/context_engine/_context_governor/__init__.py", True, False), + ("gateway/run.py", True, False), + ("tui_gateway/server.py", True, False), + ("docs/context-continuity/native-external-owner.json", True, False), + ("docs/context-continuity/native-external-owner.patch", True, False), + ("ares_runtime/collaboration.py", True, True), + ("tests/test_ares_collaboration.py", True, True), + ("ares_runtime/governed_context.py", False, True), + ("ares_runtime/__init__.py", False, True), + ("tests/owner_integration/profile_runtime_fixture.rs", False, True), + ("tests/owner_integration/semantic_memory_revocation_owner.rs", False, True), + ("tests/ares_runtime/fixtures/profile_runtime_v2_owner.json", False, True), + ("tests/ares_runtime/test_memory_witness_v2.py", False, True), + ("tests/ares_runtime/test_policy_basis_v2.py", False, True), + ("README.md", False, False), + ("apps/desktop/src/app.tsx", False, False), +]) +def test_native_qualification_path_coverage(path, continuity, owner): + values = classify([path]) + assert values["context_continuity"] is continuity + assert values["current_owner_integration"] is owner + + +@pytest.mark.parametrize("paths", [ + [], + ["", " "], + [".github/workflows/ci.yaml"], + [".github/actions/detect-changes/action.yml"], + ["scripts/ci/classify_changes.py"], + ["scripts/ci/evaluate_required_checks.py"], + ["tests/ci/test_required_checks.py"], + ["scripts/run_tests.sh"], + ["scripts/run_tests_parallel.py"], +]) +def test_uncertain_diff_and_qualification_infrastructure_select_both_native_owners(paths): + values = classify(paths) + assert values["context_continuity"] + assert values["current_owner_integration"] + + +def test_mixed_paths_select_each_owner_without_mutating_paths(): + paths = ["README.md", "ares_runtime/continuity/runtime.py", "tests/owner_integration/fixture.rs"] + original = list(paths) + values = classify(paths) + assert values["context_continuity"] and values["current_owner_integration"] + assert paths == original + + +@pytest.mark.linux_only +@pytest.mark.parametrize("scenario", ["rename", "cap300", "belowcap299", "failure", "push", "workflow_dispatch"]) +def test_composite_action_preserves_native_applicability_with_fake_compare(tmp_path, scenario): + # Execute the actual composite bash configuration, with an inert compare + # provider. Real jq applies the requested projection to our JSON fixture. + import json + import yaml + + jq = shutil.which("jq") + assert jq is not None, "jq is required for the compare-projection contract" + action = yaml.load( + (_REPO / ".github/actions/detect-changes/action.yml").read_text(encoding="utf-8"), + Loader=yaml.BaseLoader, + ) + bin_dir = tmp_path / "bin" + bin_dir.mkdir() + fixture = tmp_path / "compare.json" + call_log = tmp_path / "calls" + output = tmp_path / "outputs" + files = [{"filename": "README.md"}] + if scenario == "rename": + files[0]["previous_filename"] = "ares_runtime/collaboration.py" + elif scenario in {"cap300", "belowcap299"}: + count = 300 if scenario == "cap300" else 299 + files = [{"filename": f"docs/file-{number}.md"} for number in range(count)] + fixture.write_text(json.dumps({"files": files}), encoding="utf-8") + gh = bin_dir / "gh" + gh.write_text( + '#!/bin/bash\n' + 'printf "call\\n" >> "$FAKE_CALL_LOG"\n' + 'if [ "$FAKE_SCENARIO" = failure ]; then exit 1; fi\n' + 'while [ "$#" -gt 0 ]; do\n' + ' if [ "$1" = --jq ]; then shift; exec "$FAKE_JQ" -r "$1" "$FAKE_COMPARE"; fi\n' + ' shift\n' + 'done\n' + 'exit 2\n', + encoding="utf-8", + ) + gh.chmod(0o700) + sleep = bin_dir / "sleep" + sleep.write_text("#!/bin/sh\nexit 0\n", encoding="utf-8") + sleep.chmod(0o700) + python = bin_dir / "python3" + python.symlink_to(sys.executable) + event = scenario if scenario in {"push", "workflow_dispatch"} else "pull_request" + env = { + "PATH": f"{bin_dir}:/usr/bin:/bin", + "EVENT_NAME": event, + "REPO": "inert/fixture", + "BASE_SHA": "base", + "HEAD_SHA": "head", + "GH_TOKEN": "inert", + "GITHUB_OUTPUT": str(output), + "FAKE_COMPARE": str(fixture), + "FAKE_CALL_LOG": str(call_log), + "FAKE_SCENARIO": scenario, + "FAKE_JQ": jq, + "PYTHONDONTWRITEBYTECODE": "1", + } + result = subprocess.run( + ["/bin/bash", "-c", action["runs"]["steps"][0]["run"]], + cwd=_REPO, env=env, capture_output=True, text=True, timeout=10, + ) + assert result.returncode == 0, result.stderr + emitted = dict(line.split("=", 1) for line in output.read_text(encoding="utf-8").splitlines()) + expected = "false" if scenario == "belowcap299" else "true" + assert emitted["context_continuity"] == expected + assert emitted["current_owner_integration"] == expected + calls = call_log.read_text(encoding="utf-8").splitlines() if call_log.exists() else [] + assert len(calls) == (0 if event != "pull_request" else 3 if scenario == "failure" else 1) diff --git a/tests/ci/test_required_checks.py b/tests/ci/test_required_checks.py index d40a13bc674bf..144f2e6dcd53b 100644 --- a/tests/ci/test_required_checks.py +++ b/tests/ci/test_required_checks.py @@ -23,6 +23,8 @@ "mcp_catalog", "ci_review", "ci_review_files", + "context_continuity", + "current_owner_integration", ) @@ -33,6 +35,12 @@ def classifier(**overrides: str) -> dict[str, str]: return result +NATIVE_CALLS = ( + ("context-continuity", "context_continuity", ("native_external_owner_result", "focused_tests_result")), + ("current-owner-integration", "current_owner_integration", ("profile_runtime_consumer_result",)), +) + + def jobs_for(classifier_values: dict[str, str], *, event: str = "pull_request") -> dict[str, dict[str, object]]: result: dict[str, dict[str, object]] = {"detect": {"result": "success"}} conditional = { @@ -71,6 +79,12 @@ def jobs_for(classifier_values: dict[str, str], *, event: str = "pull_request") # explicit applicability exception, not a blanket skipped-is-green rule. result["e2e-desktop"] = {"result": "skipped"} result["osv-scanner"] = {"result": "success"} + for job, lane, outputs in NATIVE_CALLS: + applies = event != "pull_request" or classifier_values[lane] == "true" + result[job] = { + "result": "success" if applies else "skipped", + "outputs": {key: "success" for key in outputs} if applies else {}, + } return result @@ -179,3 +193,135 @@ def test_missing_or_malformed_critical_findings_fails_closed(outputs): result = run(values, jobs=jobs) assert result["status"] == "FAIL" assert any("critical_findings" in failure for failure in result["failures"]) + + +@pytest.mark.parametrize("lanes", [ + {"context_continuity": "true"}, + {"current_owner_integration": "true"}, + {"context_continuity": "true", "current_owner_integration": "true"}, +]) +def test_applicable_native_call_and_each_inner_owner_success_pass(lanes): + values = classifier(**lanes) + result = run(values) + assert result["status"] == "PASS" + for job, lane, _ in NATIVE_CALLS: + assert (job in result["required_jobs"]) == (values[lane] == "true") + + +def test_applicable_native_calls_missing_from_existing_graph_fail_closed(): + values = classifier(context_continuity="true", current_owner_integration="true") + jobs = jobs_for(values) + for job, _, _ in NATIVE_CALLS: + jobs.pop(job) + result = run(values, jobs=jobs) + assert result["status"] == "FAIL" + assert all(any(job in failure for failure in result["failures"]) for job, _, _ in NATIVE_CALLS) + + +@pytest.mark.parametrize("job,lane,outputs", NATIVE_CALLS) +@pytest.mark.parametrize("bad_result", ["failure", "cancelled", "skipped", "neutral", "", None, True]) +def test_applicable_native_caller_non_success_fails_closed(job, lane, outputs, bad_result): + values = classifier(**{lane: "true"}) + jobs = jobs_for(values) + jobs[job]["result"] = bad_result + result = run(values, jobs=jobs) + assert result["status"] == "FAIL" + assert any(job in failure for failure in result["failures"]) + + +@pytest.mark.parametrize("job,lane,outputs", NATIVE_CALLS) +@pytest.mark.parametrize("shape", ["missing", "null", "no-result", "wrong-result-key"]) +def test_applicable_native_caller_missing_or_malformed_fails_closed(job, lane, outputs, shape): + values = classifier(**{lane: "true"}) + jobs = jobs_for(values) + if shape == "missing": + jobs.pop(job) + elif shape == "null": + jobs[job] = None + elif shape == "no-result": + jobs[job] = {} + else: + jobs[job] = {"conclusion": "success"} + assert run(values, jobs=jobs)["status"] == "FAIL" + + +@pytest.mark.parametrize("job,lane,key", [ + (job, lane, key) for job, lane, outputs in NATIVE_CALLS for key in outputs +]) +@pytest.mark.parametrize("bad_result", ["failure", "cancelled", "skipped", "neutral", "", None, True, {}, []]) +def test_successful_call_cannot_hide_non_success_inner_owner(job, lane, key, bad_result): + values = classifier(**{lane: "true"}) + jobs = jobs_for(values) + jobs[job]["outputs"][key] = bad_result + result = run(values, jobs=jobs) + assert result["status"] == "FAIL" + assert any(key in failure for failure in result["failures"]) + + +@pytest.mark.parametrize("job,lane,key", [ + (job, lane, key) for job, lane, outputs in NATIVE_CALLS for key in outputs +]) +def test_successful_call_cannot_hide_missing_inner_owner_result(job, lane, key): + values = classifier(**{lane: "true"}) + jobs = jobs_for(values) + del jobs[job]["outputs"][key] + assert run(values, jobs=jobs)["status"] == "FAIL" + + +@pytest.mark.parametrize("job,lane,outputs", NATIVE_CALLS) +@pytest.mark.parametrize("shape", [None, "", [], {}]) +def test_applicable_call_requires_native_output_object_and_bindings(job, lane, outputs, shape): + values = classifier(**{lane: "true"}) + jobs = jobs_for(values) + jobs[job]["outputs"] = shape + assert run(values, jobs=jobs)["status"] == "FAIL" + + +@pytest.mark.parametrize("job,lane,outputs", NATIVE_CALLS) +def test_non_applicable_native_call_must_be_present_and_skipped(job, lane, outputs): + values = classifier() + jobs = jobs_for(values) + assert jobs[job]["result"] == "skipped" + jobs[job].pop("outputs") + assert run(values, jobs=jobs)["status"] == "PASS" + jobs.pop(job) + assert run(values, jobs=jobs)["status"] == "FAIL" + + +@pytest.mark.parametrize("job,lane,outputs", NATIVE_CALLS) +@pytest.mark.parametrize("bad_result", ["success", "failure", "cancelled"]) +def test_non_applicable_native_call_rejects_unexpected_result(job, lane, outputs, bad_result): + values = classifier() + jobs = jobs_for(values) + jobs[job]["result"] = bad_result + assert run(values, jobs=jobs)["status"] == "FAIL" + + +@pytest.mark.parametrize("event", ["push", "workflow_dispatch"]) +@pytest.mark.parametrize("job,lane,outputs", NATIVE_CALLS) +def test_postmerge_and_dispatch_require_native_owners_even_with_false_flags(event, job, lane, outputs): + values = classifier() + jobs = jobs_for(values, event=event) + assert run(values, event=event, jobs=jobs)["status"] == "PASS" + jobs[job] = {"result": "skipped"} + assert run(values, event=event, jobs=jobs)["status"] == "FAIL" + + +@pytest.mark.parametrize("lane", ["context_continuity", "current_owner_integration"]) +@pytest.mark.parametrize("bad_value", ["missing", "", "unknown", None, True]) +def test_native_classifier_missing_or_malformed_cannot_authorize_skip(lane, bad_value): + values = classifier() + jobs = jobs_for(values) + if bad_value == "missing": + del values[lane] + else: + values[lane] = bad_value + assert run(values, jobs=jobs)["status"] == "FAIL" + + +@pytest.mark.parametrize("result", ["failure", "cancelled", "skipped"]) +def test_failed_detect_cannot_authorize_native_optional_results(result): + values = classifier() + jobs = jobs_for(values) + jobs["detect"]["result"] = result + assert run(values, jobs=jobs)["status"] == "FAIL" diff --git a/tests/ci/test_workflow_contract.py b/tests/ci/test_workflow_contract.py index 97f8a6557ec32..98f3248a92b16 100644 --- a/tests/ci/test_workflow_contract.py +++ b/tests/ci/test_workflow_contract.py @@ -50,10 +50,11 @@ def test_aggregate_runs_checksum_pinned_actionlint_before_evaluator(): assert "ACTIONLINT_SHA256: 900919a84f2229bac68ca9cd4103ea297abc35e9689ebb842c6e34a3d1b01b0a" in block assert "archive=\"actionlint_${ACTIONLINT_VERSION}_linux_amd64.tar.gz\"" in block assert "Lint GitHub Actions workflows" in block - assert ( - "run: ./actionlint -ignore 'constant expression \"false\" in condition' " - ".github/workflows/ci.yaml" - ) in block + lint = next(step for step in _workflow()["jobs"]["all-checks-pass"]["steps"] + if step.get("name") == "Lint GitHub Actions workflows") + assert ".github/workflows/ci.yaml" in lint["run"] + assert ".github/workflows/context-continuity-qualification.yml" in lint["run"] + assert ".github/workflows/current-owner-integration.yml" in lint["run"] assert "for attempt in range(1, 4):" in block assert "time.sleep(attempt * 2)" in block assert "scripts/ci/evaluate_required_checks.py" in block @@ -77,3 +78,59 @@ def test_aggregate_has_always_and_complete_known_need_set(): assert re.search(r"^ if: always\(\)\s*$", block, re.MULTILINE) needs = set(re.findall(r"^ - ([a-z0-9_-]+)\s*$", block, re.MULTILINE)) assert KNOWN_JOBS == needs + + + +def _workflow(path=".github/workflows/ci.yaml"): + import yaml + # BaseLoader preserves GitHub's "on" and literal false as scalar text. + return yaml.load((ROOT / path).read_text(encoding="utf-8"), Loader=yaml.BaseLoader) + + +def test_native_callers_reuse_owners_with_read_only_tokens_and_no_secrets(): + jobs = _workflow()["jobs"] + for caller, lane, workflow in [ + ("context-continuity", "context_continuity", "context-continuity-qualification.yml"), + ("current-owner-integration", "current_owner_integration", "current-owner-integration.yml"), + ]: + job = jobs[caller] + assert job["needs"] == "detect" + assert job["if"] == f"needs.detect.outputs.{lane} == 'true'" + assert job["uses"] == f"./.github/workflows/{workflow}" + assert job["permissions"] == {"contents": "read"} + assert "secrets" not in job + assert job.get("continue-on-error", "false") == "false" + assert lane in jobs["detect"]["outputs"] + assert caller in jobs["all-checks-pass"]["needs"] + + +def test_native_workflow_outputs_bind_each_unconditional_owner_job(): + for path, bindings in [ + ("context-continuity-qualification.yml", { + "native_external_owner_result": "native-external-owner", + "focused_tests_result": "focused-tests", + }), + ("current-owner-integration.yml", { + "profile_runtime_consumer_result": "profile-runtime-consumer", + }), + ]: + workflow = _workflow(f".github/workflows/{path}") + assert workflow["permissions"] == {"contents": "read"} + assert "pull_request" not in workflow["on"] + outputs = workflow["on"]["workflow_call"]["outputs"] + for output, job_id in bindings.items(): + assert outputs[output]["value"] == "${{ jobs." + job_id + ".outputs.qualification_result }}" + job = workflow["jobs"][job_id] + assert job["outputs"]["qualification_result"] == "${{ job.status }}" + assert "if" not in job + assert job.get("continue-on-error", "false") == "false" + assert all(step.get("continue-on-error", "false") == "false" for step in job["steps"]) + + +def test_native_main_and_pr_orchestration_has_one_owner(): + ci = _workflow() + assert ci["on"]["push"]["branches"] == ["main"] + continuity = _workflow(".github/workflows/context-continuity-qualification.yml") + assert continuity["on"]["push"]["branches"] == ["feat/context-continuity-v4-20260923"] + owner = _workflow(".github/workflows/current-owner-integration.yml") + assert set(owner["on"]) == {"workflow_call"} diff --git a/tests/cron/test_claim_job_for_fire.py b/tests/cron/test_claim_job_for_fire.py index fa0f7b39b3495..456e3db808381 100644 --- a/tests/cron/test_claim_job_for_fire.py +++ b/tests/cron/test_claim_job_for_fire.py @@ -215,3 +215,324 @@ def test_fire_claim_fence_rejects_stale_owner(temp_home): with fire_claim_fence(job["id"], expected_owner="stale") as owns_claim: assert owns_claim is False + + +# SD03B controls: real claim/release owners, dictionary-only persistence and +# inert locks. Select this class alone; the preceding legacy file tests use +# task-local filesystem stores and include a threaded fence control. +@pytest.fixture +def sd03b_memory_jobs(monkeypatch): + from contextlib import contextmanager + from copy import deepcopy + from datetime import datetime, timezone + from types import SimpleNamespace + + from cron import jobs as subject + + store = SimpleNamespace( + subject=subject, + jobs=[{ + "id": "inert-job", + "name": "inert", + "prompt": "inert request", + "enabled": True, + "state": "scheduled", + "schedule": {"kind": "interval", "seconds": 300}, + "next_run_at": "2026-10-04T12:00:00+00:00", + "repeat": {"times": 3, "completed": 1}, + "last_status": "ok", + "last_run_at": "2026-10-04T11:00:00+00:00", + }], + trace=[], + fire_lock_allowed=True, + save_failure=None, + saves=0, + ) + now = datetime(2026, 10, 4, 12, 0, tzinfo=timezone.utc) + sequence = [0] + + @contextmanager + def fire_lock(job_id): + store.trace.append(("fire_enter", job_id)) + try: + yield store.fire_lock_allowed + finally: + store.trace.append(("fire_exit", job_id)) + + @contextmanager + def jobs_lock(): + store.trace.append(("jobs_enter",)) + try: + yield + finally: + store.trace.append(("jobs_exit",)) + + def load(): + store.trace.append(("load",)) + return deepcopy(store.jobs) + + def save(records, **kwargs): + store.trace.append(("save",)) + store.saves += 1 + if store.save_failure == "before": + raise OSError("inert save refusal before publication") + store.jobs = deepcopy(records) + if store.save_failure == "after": + raise OSError("inert save uncertainty after publication") + + def next_run(schedule, reference): + store.trace.append(("compute", deepcopy(schedule), reference)) + return "2026-10-04T12:05:00+00:00" + + def uuid4(): + sequence[0] += 1 + return SimpleNamespace(hex=f"inert-acquisition-{sequence[0]}") + + def forbidden_accounting(*args, **kwargs): + pytest.fail("unstarted disposition must not call run accounting") + + monkeypatch.setattr(subject, "_fire_job_lock", fire_lock) + monkeypatch.setattr(subject, "_jobs_lock", jobs_lock) + monkeypatch.setattr(subject, "load_jobs", load) + monkeypatch.setattr(subject, "save_jobs", save) + monkeypatch.setattr(subject, "_hermes_now", lambda: now) + monkeypatch.setattr(subject, "compute_next_run", next_run) + monkeypatch.setattr(subject, "_machine_id", lambda: "inert-owner") + monkeypatch.setattr(subject, "uuid", SimpleNamespace(uuid4=uuid4)) + monkeypatch.setattr(subject, "mark_job_run", forbidden_accounting) + monkeypatch.setattr(subject, "_mark_job_run_locked", forbidden_accounting) + return store + + +def _sd03b_api(store, name): + # Missing baseline API is an assertion in the selected test, never a + # collection-time import or attribute error. + api = getattr(store.subject, name, None) + assert callable(api), f"SD03B owner API is absent: {name}" + return api + + +def _sd03b_claim(store): + result = _sd03b_api( + store, "claim_job_for_fire_with_unstarted_receipt", + )("inert-job") + assert isinstance(result, tuple) and len(result) == 2 + claimed, receipt = result + assert isinstance(claimed, dict) + return claimed, receipt + + +def _sd03b_release(store, receipt): + result = _sd03b_api(store, "release_unstarted_fire_claim")(receipt) + assert isinstance(result, dict) + assert isinstance(result.get("released"), bool) + return result + + +class TestSD03BUnstartedFireClaim: + def test_legacy_bool_dict_and_force_contract(self, sd03b_memory_jobs): + from inspect import signature + + store = sd03b_memory_jobs + claim = store.subject.claim_job_for_fire + assert tuple(signature(claim).parameters) == ( + "job_id", "claim_ttl_seconds", "force", "return_job", + ) + assert claim("missing") is False + assert claim("inert-job") is True + assert claim("inert-job") is False + snapshot = claim("inert-job", claim_ttl_seconds=0, return_job=True) + assert isinstance(snapshot, dict) + assert snapshot["fire_claim"]["by"] == store.jobs[0]["fire_claim"]["by"] + snapshot["repeat"]["completed"] = 99 + assert store.jobs[0]["repeat"]["completed"] == 1 + store.jobs[0].pop("fire_claim") + store.jobs[0].update({ + "enabled": False, "state": "paused", "paused_at": "paused", + "paused_reason": "operator", + }) + assert claim("inert-job") is False + assert claim("inert-job", force=True) is True + assert store.jobs[0]["enabled"] is True + assert store.jobs[0]["state"] == "scheduled" + assert store.jobs[0]["paused_at"] is None + assert store.jobs[0]["paused_reason"] is None + + @pytest.mark.parametrize("prior", ["absent", "null", "value"]) + def test_receipt_is_ephemeral_frozen_and_captures_prior_field( + self, sd03b_memory_jobs, prior, + ): + from copy import deepcopy + from dataclasses import FrozenInstanceError, is_dataclass + + store = sd03b_memory_jobs + if prior == "absent": + store.jobs[0].pop("next_run_at") + elif prior == "null": + store.jobs[0]["next_run_at"] = None + before = deepcopy(store.jobs[0]) + claimed, receipt = _sd03b_claim(store) + receipt_type = getattr(store.subject, "UnstartedFireClaimReceipt", None) + assert receipt_type is not None and isinstance(receipt, receipt_type) + assert is_dataclass(receipt) + assert receipt.job_id == "inert-job" + assert receipt.owner == claimed["fire_claim"]["by"] + assert receipt.prior_next_present is (prior != "absent") + assert receipt.prior_next_value == before.get("next_run_at") + assert receipt.claimed_next_present is True + assert receipt.claimed_next_value == claimed["next_run_at"] + assert receipt.claimed_schedule_present is True + assert receipt.claimed_schedule_value == before["schedule"] + with pytest.raises(FrozenInstanceError): + receipt.owner = "cannot-rewrite" + assert set(store.jobs[0]) == set(before) | {"fire_claim", "next_run_at"} + assert claimed == store.jobs[0] + claimed["schedule"]["seconds"] = 999 + assert receipt.claimed_schedule_value["seconds"] == 300 + assert store.jobs[0]["schedule"]["seconds"] == 300 + names = [entry[0] for entry in store.trace] + assert names.index("fire_enter") < names.index("jobs_enter") + assert names.index("jobs_enter") < names.index("load") + assert names.index("save") < names.index("jobs_exit") + assert names.index("jobs_exit") < names.index("fire_exit") + + def test_private_opt_in_returns_snapshot_and_receipt(self, sd03b_memory_jobs): + from dataclasses import is_dataclass + + store = sd03b_memory_jobs + _sd03b_api(store, "claim_job_for_fire_with_unstarted_receipt") + claimed, receipt = store.subject._claim_job_for_fire_locked( + "inert-job", return_job=True, return_unstarted_receipt=True, + ) + assert isinstance(claimed, dict) and is_dataclass(receipt) + assert receipt.owner == claimed["fire_claim"]["by"] + assert set(store.jobs[0]) == set(claimed) + assert "unstarted_receipt" not in claimed + + @pytest.mark.parametrize("prior", ["absent", "null", "value"]) + def test_matching_release_restores_only_claim_and_next( + self, sd03b_memory_jobs, prior, + ): + from copy import deepcopy + + store = sd03b_memory_jobs + if prior == "absent": + store.jobs[0].pop("next_run_at") + elif prior == "null": + store.jobs[0]["next_run_at"] = None + before = deepcopy(store.jobs[0]) + _, receipt = _sd03b_claim(store) + store.trace.clear() + result = _sd03b_release(store, receipt) + assert result["released"] is True + assert result["status"] == "released" + assert store.jobs[0] == before + assert [entry[0] for entry in store.trace] == [ + "fire_enter", "jobs_enter", "load", "save", "jobs_exit", "fire_exit", + ] + # An old receipt cannot release a later acquisition or double-release. + store.trace.clear() + assert _sd03b_release(store, receipt)["released"] is False + assert store.jobs[0] == before + assert not any(entry[0] == "save" for entry in store.trace) + + @pytest.mark.parametrize("edit", [ + "owner", "next_value", "next_absent", "next_null", + "schedule", "schedule_absent", "deleted", + ]) + def test_release_conflict_preserves_current_job(self, sd03b_memory_jobs, edit): + from copy import deepcopy + + store = sd03b_memory_jobs + _, receipt = _sd03b_claim(store) + if edit == "owner": + store.jobs[0]["fire_claim"]["by"] = "replacement-owner" + elif edit == "next_value": + store.jobs[0]["next_run_at"] = "operator-edited-next" + elif edit == "next_absent": + store.jobs[0].pop("next_run_at") + elif edit == "next_null": + store.jobs[0]["next_run_at"] = None + elif edit == "schedule": + store.jobs[0]["schedule"]["seconds"] = 600 + elif edit == "schedule_absent": + store.jobs[0].pop("schedule") + else: + store.jobs.clear() + current = deepcopy(store.jobs) + store.trace.clear() + result = _sd03b_release(store, receipt) + assert result["released"] is False + assert result["status"] != "released" + assert store.jobs == current + assert not any(entry[0] == "save" for entry in store.trace) + + def test_release_preserves_pause_and_unrelated_edits(self, sd03b_memory_jobs): + from copy import deepcopy + + store = sd03b_memory_jobs + old_next = store.jobs[0]["next_run_at"] + _, receipt = _sd03b_claim(store) + store.jobs[0].update({ + "enabled": False, "state": "paused", "paused_at": "operator-pause", + "paused_reason": "operator", "name": "edited-name", + "prompt": "edited-prompt", "last_status": "edited-status", + "repeat": {"times": 7, "completed": 2}, + }) + expected = deepcopy(store.jobs[0]) + expected.pop("fire_claim") + expected["next_run_at"] = old_next + assert _sd03b_release(store, receipt)["released"] is True + assert store.jobs[0] == expected + + def test_claim_and_release_lock_refusal_are_inert(self, sd03b_memory_jobs): + from copy import deepcopy + + store = sd03b_memory_jobs + wrapper = _sd03b_api(store, "claim_job_for_fire_with_unstarted_receipt") + before = deepcopy(store.jobs) + store.fire_lock_allowed = False + assert wrapper("inert-job") is None + assert store.jobs == before + assert [entry[0] for entry in store.trace] == ["fire_enter", "fire_exit"] + store.fire_lock_allowed = True + _, receipt = _sd03b_claim(store) + current = deepcopy(store.jobs) + store.trace.clear() + store.fire_lock_allowed = False + result = _sd03b_release(store, receipt) + assert result["released"] is False + assert store.jobs == current + assert [entry[0] for entry in store.trace] == ["fire_enter", "fire_exit"] + + @pytest.mark.parametrize("failure", ["before", "after"]) + def test_release_save_failure_reports_uncertainty(self, sd03b_memory_jobs, failure): + from copy import deepcopy + + store = sd03b_memory_jobs + original = deepcopy(store.jobs[0]) + _, receipt = _sd03b_claim(store) + claimed = deepcopy(store.jobs[0]) + store.trace.clear() + store.save_failure = failure + result = _sd03b_release(store, receipt) + assert result["released"] is False + assert result.get("claim_release") == "uncertain" + assert result.get("error") + assert store.jobs[0] == (claimed if failure == "before" else original) + assert [entry[0] for entry in store.trace] == [ + "fire_enter", "jobs_enter", "load", "save", "jobs_exit", "fire_exit", + ] + + def test_missing_paused_and_terminal_opt_in_claims_refuse(self, sd03b_memory_jobs): + from copy import deepcopy + + store = sd03b_memory_jobs + wrapper = _sd03b_api(store, "claim_job_for_fire_with_unstarted_receipt") + assert wrapper("missing") is None + for state in ("paused", "completed", "error"): + store.jobs[0]["state"] = state + before = deepcopy(store.jobs) + assert wrapper("inert-job") is None + assert store.jobs == before + assert store.saves == 0 diff --git a/tests/cron/test_cron_bot_chat_delivery.py b/tests/cron/test_cron_bot_chat_delivery.py index 92ebadde49035..2bc5ecd4c20c5 100644 --- a/tests/cron/test_cron_bot_chat_delivery.py +++ b/tests/cron/test_cron_bot_chat_delivery.py @@ -144,7 +144,45 @@ def fake_run(argv, **kwargs): assert not any("the output" in str(a) for a in argv) -def test_deliver_named_profile_uses_p_flag_and_clears_home(): +@pytest.fixture +def named_delivery_profiles(tmp_path, monkeypatch): + from hermes_cli.env_loader import _record_external_secret_snapshot + + root = tmp_path / "root" + source = root / "profiles" / "source" + target = root / "profiles" / "research" + source.mkdir(parents=True) + target.mkdir() + (root / ".env").write_text("ROOT_ONLY_TOKEN=synthetic-root\n") + (source / ".env").write_text( + "OPENAI_API_KEY=synthetic-source\n" + "GROQ_API_KEY=synthetic-source-only\n" + "SOURCE_CUSTOM=synthetic-custom\n" + ) + (target / ".env").write_text("OPENAI_API_KEY=synthetic-target\n") + for home in (source, target): + (home / "config.yaml").write_text( + "security:\n inherit_root_credentials: " + + ("true" if home == source else "false") + "\n" + ) + _record_external_secret_snapshot( + source, data={"SOURCE_EXTERNAL_TOKEN": "synthetic-source-external"}, status="ready" + ) + _record_external_secret_snapshot( + target, data={"ANTHROPIC_API_KEY": "synthetic-target-only"}, status="ready" + ) + monkeypatch.setenv("HERMES_HOME", str(source)) + monkeypatch.setenv("OPENAI_API_KEY", "synthetic-source") + monkeypatch.setenv("GROQ_API_KEY", "synthetic-source-only") + monkeypatch.setenv("SOURCE_CUSTOM", "synthetic-custom") + monkeypatch.setenv("SOURCE_EXTERNAL_TOKEN", "synthetic-source-external") + monkeypatch.setenv("ROOT_ONLY_TOKEN", "synthetic-root") + monkeypatch.setenv("SHELL_CONTROL", "synthetic-unowned") + return source, target + + +def test_deliver_named_profile_uses_p_flag_and_target_authority(named_delivery_profiles): + _, target = named_delivery_profiles calls = {} def fake_run(argv, **kwargs): @@ -153,15 +191,42 @@ def fake_run(argv, **kwargs): return _completed() with mock.patch.object(sched.subprocess, "run", side_effect=fake_run), \ - mock.patch.object(sched.shutil, "which", return_value="/usr/bin/hermes"), \ - mock.patch.dict(sched.os.environ, {"HERMES_HOME": "/tmp/other-profile"}): + mock.patch.object(sched.shutil, "which", return_value="/usr/bin/hermes"): err = _deliver_to_bot_chat({"id": "j1", "name": "n"}, "out", "research") assert err is None argv = calls["argv"] assert argv[1:3] == ["-p", "research"] - # -p owns resolution; the scheduler's own HERMES_HOME must not leak in. - assert "HERMES_HOME" not in calls["kwargs"]["env"] + env = calls["kwargs"]["env"] + assert env["HERMES_HOME"] == str(target) + assert env["OPENAI_API_KEY"] == "synthetic-target" + assert env["ANTHROPIC_API_KEY"] == "synthetic-target-only" + for name in ("GROQ_API_KEY", "SOURCE_CUSTOM", "SOURCE_EXTERNAL_TOKEN", "ROOT_ONLY_TOKEN"): + assert name not in env + assert env["SHELL_CONTROL"] == "synthetic-unowned" + + +def test_deliver_named_profile_missing_authority_starts_no_child(tmp_path, monkeypatch): + source = tmp_path / "root" / "profiles" / "source" + source.mkdir(parents=True) + monkeypatch.setenv("HERMES_HOME", str(source)) + with mock.patch.object(sched.subprocess, "run") as child, \ + mock.patch.object(sched.shutil, "which", return_value="/usr/bin/hermes"): + err = _deliver_to_bot_chat({"id": "j1"}, "out", "research") + assert err == "bot-chat delivery failed: profile authority unavailable" + child.assert_not_called() + + +def test_deliver_named_profile_failed_snapshot_starts_no_child(named_delivery_profiles): + from hermes_cli.env_loader import _record_external_secret_snapshot + + _, target = named_delivery_profiles + _record_external_secret_snapshot(target, data={}, status="failed", error_kind="synthetic") + with mock.patch.object(sched.subprocess, "run") as child, \ + mock.patch.object(sched.shutil, "which", return_value="/usr/bin/hermes"): + err = _deliver_to_bot_chat({"id": "j1"}, "out", "research") + assert err == "bot-chat delivery failed: profile authority unavailable" + child.assert_not_called() def test_deliver_failure_returns_error_string(): diff --git a/tests/gateway/test_pending_queue_spool.py b/tests/gateway/test_pending_queue_spool.py index cec51cb2780fe..7d2ac987b8447 100644 --- a/tests/gateway/test_pending_queue_spool.py +++ b/tests/gateway/test_pending_queue_spool.py @@ -221,3 +221,157 @@ def test_roundtrip_order(self, spool_home): assert remaining == 0 assert seen == ["c0", "c1", "c2"] assert _spool_files(spool_home) == [] + + +# SD05: expected kwargs below are independently specified contract fixtures, +# never computed by the production transcript mapper. +_SD05_SUPPORTED_FIELD_CASES = [ + pytest.param( + { + "role": "assistant", + "content": None, + "tool_name": "fixture_tool", + "tool_calls": [{"id": "call-fixture-1", "type": "function", "function": {"name": "fixture_tool", "arguments": "{\"x\":1}"}}], + "tool_call_id": "assistant-call-id", + "reasoning": "assistant reasoning", + "reasoning_content": "reasoning content", + "reasoning_details": [{"type": "text", "text": "detail"}], + "codex_reasoning_items": [{"type": "reasoning", "id": "r1"}], + "codex_message_items": [{"type": "message", "id": "m1"}], + "platform_message_id": "platform-assistant", + "message_id": "ignored-fallback", + "observed": 1, + "timestamp": 0, + "api_content": "", + "display_kind": "internal_notification", + "display_metadata": {"source": "fixture", "ordinal": 1}, + }, + { + "session_id": "sd05-session", + "role": "assistant", + "content": None, + "tool_name": "fixture_tool", + "tool_calls": [{"id": "call-fixture-1", "type": "function", "function": {"name": "fixture_tool", "arguments": "{\"x\":1}"}}], + "tool_call_id": "assistant-call-id", + "reasoning": "assistant reasoning", + "reasoning_content": "reasoning content", + "reasoning_details": [{"type": "text", "text": "detail"}], + "codex_reasoning_items": [{"type": "reasoning", "id": "r1"}], + "codex_message_items": [{"type": "message", "id": "m1"}], + "platform_message_id": "platform-assistant", + "observed": True, + "timestamp": 0, + "api_content": "", + "display_kind": "internal_notification", + "display_metadata": {"source": "fixture", "ordinal": 1}, + }, + id="assistant", + ), + pytest.param( + { + "role": "tool", + "content": "tool output", + "tool_name": "fixture_tool", + "tool_call_id": "call-fixture-1", + "reasoning": "must not persist", + "reasoning_content": "must not persist", + "reasoning_details": [{"text": "must not persist"}], + "codex_reasoning_items": [{"id": "must-not-persist"}], + "codex_message_items": [{"id": "must-not-persist"}], + "platform_message_id": "", + "message_id": "platform-tool-fallback", + "observed": [], + "timestamp": 50.0, + "api_content": {"not": "a string"}, + "display_kind": "tool_result", + "display_metadata": {"source": "fixture", "ordinal": 2}, + }, + { + "session_id": "sd05-session", + "role": "tool", + "content": "tool output", + "tool_name": "fixture_tool", + "tool_calls": None, + "tool_call_id": "call-fixture-1", + "reasoning": None, + "reasoning_content": None, + "reasoning_details": None, + "codex_reasoning_items": None, + "codex_message_items": None, + "platform_message_id": "platform-tool-fallback", + "observed": False, + "timestamp": 50.0, + "api_content": None, + "display_kind": "tool_result", + "display_metadata": {"source": "fixture", "ordinal": 2}, + }, + id="tool", + ), + pytest.param( + { + "role": "user", + "content": "user message", + "reasoning": "must not persist", + "reasoning_content": "must not persist", + "reasoning_details": [{"text": "must not persist"}], + "codex_reasoning_items": [{"id": "must-not-persist"}], + "codex_message_items": [{"id": "must-not-persist"}], + "message_id": "platform-user-fallback", + "observed": False, + "timestamp": None, + "api_content": "user message\n\nfixture context", + "display_kind": "user_input", + "display_metadata": {"source": "fixture", "ordinal": 3}, + }, + { + "session_id": "sd05-session", + "role": "user", + "content": "user message", + "tool_name": None, + "tool_calls": None, + "tool_call_id": None, + "reasoning": None, + "reasoning_content": None, + "reasoning_details": None, + "codex_reasoning_items": None, + "codex_message_items": None, + "platform_message_id": "platform-user-fallback", + "observed": False, + "timestamp": None, + "api_content": "user message\n\nfixture context", + "display_kind": "user_input", + "display_metadata": {"source": "fixture", "ordinal": 3}, + }, + id="user", + ), +] + + +class _Sd05RecordingDb: + """A supplied recording object; no SessionDB constructor or live storage.""" + + def __init__(self): + self.rows = [] + + def append_message(self, **kwargs): + self.rows.append(kwargs) + + +@pytest.mark.parametrize("message, expected", _SD05_SUPPORTED_FIELD_CASES) +def test_append_transcript_message_preserves_full_message_fields(message, expected): + """The extracted mapper must leave live forwarding unchanged.""" + before = json.dumps(message, sort_keys=True) + db = _Sd05RecordingDb() + store = _make_store(db) + store._append_transcript_message("sd05-session", message) + assert db.rows == [expected] + assert json.dumps(message, sort_keys=True) == before + + +@pytest.mark.parametrize("content", [None, "", [], {}], ids=["none", "empty-string", "empty-list", "empty-dict"]) +def test_append_transcript_message_keeps_live_falsy_content_and_zero_timestamp(content): + message = {"role": "assistant", "content": content, "timestamp": 0} + db = _Sd05RecordingDb() + _make_store(db)._append_transcript_message("sd05-session", message) + assert db.rows[0]["content"] == content + assert db.rows[0]["timestamp"] == 0 diff --git a/tests/gateway/test_shutdown_flush.py b/tests/gateway/test_shutdown_flush.py index fe67ab7168235..9856cdb4ea5c7 100644 --- a/tests/gateway/test_shutdown_flush.py +++ b/tests/gateway/test_shutdown_flush.py @@ -168,3 +168,257 @@ def fake_get_hermes_home(): assert result == tmp_path / "pending_messages" + + +# SD05: expected kwargs below are independently specified contract fixtures, +# never computed by the production transcript mapper. +_SD05_SUPPORTED_FIELD_CASES = [ + pytest.param( + { + "role": "assistant", + "content": None, + "tool_name": "fixture_tool", + "tool_calls": [{"id": "call-fixture-1", "type": "function", "function": {"name": "fixture_tool", "arguments": "{\"x\":1}"}}], + "tool_call_id": "assistant-call-id", + "reasoning": "assistant reasoning", + "reasoning_content": "reasoning content", + "reasoning_details": [{"type": "text", "text": "detail"}], + "codex_reasoning_items": [{"type": "reasoning", "id": "r1"}], + "codex_message_items": [{"type": "message", "id": "m1"}], + "platform_message_id": "platform-assistant", + "message_id": "ignored-fallback", + "observed": 1, + "timestamp": 0, + "api_content": "", + "display_kind": "internal_notification", + "display_metadata": {"source": "fixture", "ordinal": 1}, + }, + { + "session_id": "sd05-session", + "role": "assistant", + "content": "", + "tool_name": "fixture_tool", + "tool_calls": [{"id": "call-fixture-1", "type": "function", "function": {"name": "fixture_tool", "arguments": "{\"x\":1}"}}], + "tool_call_id": "assistant-call-id", + "reasoning": "assistant reasoning", + "reasoning_content": "reasoning content", + "reasoning_details": [{"type": "text", "text": "detail"}], + "codex_reasoning_items": [{"type": "reasoning", "id": "r1"}], + "codex_message_items": [{"type": "message", "id": "m1"}], + "platform_message_id": "platform-assistant", + "observed": True, + "timestamp": 1234567, + "api_content": "", + "display_kind": "internal_notification", + "display_metadata": {"source": "fixture", "ordinal": 1}, + }, + id="assistant", + ), + pytest.param( + { + "role": "tool", + "content": "tool output", + "tool_name": "fixture_tool", + "tool_call_id": "call-fixture-1", + "reasoning": "must not persist", + "reasoning_content": "must not persist", + "reasoning_details": [{"text": "must not persist"}], + "codex_reasoning_items": [{"id": "must-not-persist"}], + "codex_message_items": [{"id": "must-not-persist"}], + "platform_message_id": "", + "message_id": "platform-tool-fallback", + "observed": [], + "timestamp": 50.0, + "api_content": {"not": "a string"}, + "display_kind": "tool_result", + "display_metadata": {"source": "fixture", "ordinal": 2}, + }, + { + "session_id": "sd05-session", + "role": "tool", + "content": "tool output", + "tool_name": "fixture_tool", + "tool_calls": None, + "tool_call_id": "call-fixture-1", + "reasoning": None, + "reasoning_content": None, + "reasoning_details": None, + "codex_reasoning_items": None, + "codex_message_items": None, + "platform_message_id": "platform-tool-fallback", + "observed": False, + "timestamp": 50.0, + "api_content": None, + "display_kind": "tool_result", + "display_metadata": {"source": "fixture", "ordinal": 2}, + }, + id="tool", + ), + pytest.param( + { + "role": "user", + "content": "user message", + "reasoning": "must not persist", + "reasoning_content": "must not persist", + "reasoning_details": [{"text": "must not persist"}], + "codex_reasoning_items": [{"id": "must-not-persist"}], + "codex_message_items": [{"id": "must-not-persist"}], + "message_id": "platform-user-fallback", + "observed": False, + "timestamp": None, + "api_content": "user message\n\nfixture context", + "display_kind": "user_input", + "display_metadata": {"source": "fixture", "ordinal": 3}, + }, + { + "session_id": "sd05-session", + "role": "user", + "content": "user message", + "tool_name": None, + "tool_calls": None, + "tool_call_id": None, + "reasoning": None, + "reasoning_content": None, + "reasoning_details": None, + "codex_reasoning_items": None, + "codex_message_items": None, + "platform_message_id": "platform-user-fallback", + "observed": False, + "timestamp": 1234567, + "api_content": "user message\n\nfixture context", + "display_kind": "user_input", + "display_metadata": {"source": "fixture", "ordinal": 3}, + }, + id="user", + ), +] + + +class _Sd05RecordingDb: + """Only record supplied kwargs; never construct or open a real DB.""" + + def __init__(self): + self.rows = [] + + def append_message(self, **kwargs): + self.rows.append(kwargs) + + +def _sd05_write_cap_spool(tmp_path, monkeypatch, message, session_id="sd05-session"): + flush_dir = _make_flush_dir(tmp_path) + monkeypatch.setattr("gateway.shutdown_flush._get_flush_dir", lambda: flush_dir) + payload = { + "session_key": session_id, + "reason": "transcript_cap_drop", + "ts": 1234567, + "data": {"session_id": session_id, "message": message}, + } + original = json.dumps(payload, sort_keys=True, indent=2).encode("utf-8") + assert len(original) < 4096 + path = flush_dir / "sd05-cap.json" + path.write_bytes(original) + return path, original + + +@pytest.mark.parametrize("message, expected", _SD05_SUPPORTED_FIELD_CASES) +def test_recover_transcript_cap_drop_preserves_full_message_fields( + tmp_path, monkeypatch, message, expected +): + """Startup restores supported metadata with existing content/time fallbacks.""" + path, original = _sd05_write_cap_spool(tmp_path, monkeypatch, message) + + class AppendBeforeUnlinkDb(_Sd05RecordingDb): + def append_message(self, **kwargs): + assert path.exists() + assert path.read_bytes() == original + super().append_message(**kwargs) + + db = AppendBeforeUnlinkDb() + assert recover_pending_to_db(db) == 1 + assert db.rows == [expected] + assert not path.exists() + + +@pytest.mark.parametrize( + "message", + [ + pytest.param({"role": "user", "content": "message"}, id="missing"), + pytest.param({"role": "user", "content": "message", "timestamp": None}, id="none"), + pytest.param({"role": "user", "content": "message", "timestamp": 0}, id="zero"), + ], +) +def test_recover_transcript_cap_drop_uses_payload_timestamp( + tmp_path, monkeypatch, message +): + path, _original = _sd05_write_cap_spool(tmp_path, monkeypatch, message) + db = _Sd05RecordingDb() + assert recover_pending_to_db(db) == 1 + assert db.rows == [{ + "session_id": "sd05-session", + "role": "user", + "content": "message", + "tool_name": None, + "tool_calls": None, + "tool_call_id": None, + "reasoning": None, + "reasoning_content": None, + "reasoning_details": None, + "codex_reasoning_items": None, + "codex_message_items": None, + "platform_message_id": None, + "observed": False, + "timestamp": 1234567, + "api_content": None, + "display_kind": None, + "display_metadata": None, + }] + assert not path.exists() + + +@pytest.mark.parametrize("content", [None, "", [], {}], ids=["none", "empty-string", "empty-list", "empty-dict"]) +def test_recover_transcript_cap_drop_keeps_falsy_content_compatibility( + tmp_path, monkeypatch, content +): + path, _original = _sd05_write_cap_spool( + tmp_path, monkeypatch, {"role": "assistant", "content": content, "timestamp": 9} + ) + db = _Sd05RecordingDb() + assert recover_pending_to_db(db) == 1 + assert db.rows[0]["content"] == "" + assert db.rows[0]["timestamp"] == 9 + assert not path.exists() + + +def test_recover_transcript_cap_drop_failure_preserves_spool_bytes(tmp_path, monkeypatch): + message = {"role": "tool", "content": "result", "tool_call_id": "call-fixture-1", "api_content": "exact bytes"} + path, original = _sd05_write_cap_spool(tmp_path, monkeypatch, message) + + class FailingRecordingDb(_Sd05RecordingDb): + def append_message(self, **kwargs): + assert path.read_bytes() == original + super().append_message(**kwargs) + raise RuntimeError("sd05 inert append failure") + + db = FailingRecordingDb() + # Preserve the current BaseException-before-Exception propagation behavior. + with pytest.raises(RuntimeError, match="sd05 inert append failure"): + recover_pending_to_db(db) + assert len(db.rows) == 1 + assert path.read_bytes() == original + + +@pytest.mark.parametrize( + "session_id, message", + [ + pytest.param("", {"role": "user", "content": "kept"}, id="missing-session"), + pytest.param("sd05-session", ["invalid message"], id="non-dict-message"), + ], +) +def test_recover_transcript_cap_drop_invalid_payload_retains_spool( + tmp_path, monkeypatch, session_id, message +): + path, original = _sd05_write_cap_spool(tmp_path, monkeypatch, message, session_id) + db = _Sd05RecordingDb() + assert recover_pending_to_db(db) == 0 + assert db.rows == [] + assert path.read_bytes() == original diff --git a/tests/hermes_cli/test_backup.py b/tests/hermes_cli/test_backup.py index 04369783ac2b6..6b2750e1811a9 100644 --- a/tests/hermes_cli/test_backup.py +++ b/tests/hermes_cli/test_backup.py @@ -309,6 +309,70 @@ def test_state_db_passes(self, tmp_path): # --------------------------------------------------------------------------- class TestImport: + + def _import_effect_spies(self, monkeypatch, hermes_home): + import builtins + import sys + from types import ModuleType + from hermes_cli import backup as backup_mod + calls = {key: [] for key in ("wrappers", "service_checks", "service_starts", "provider_paths", "provider_imports")} + profiles = ModuleType("hermes_cli.profiles") + profiles.create_wrapper_script = lambda name: calls["wrappers"].append(name) or hermes_home / "unused-wrapper" + profiles.check_alias_collision = lambda name: None + profiles._is_wrapper_dir_in_path = lambda: True + profiles._get_wrapper_dir = lambda: hermes_home / "unused-wrappers" + monkeypatch.setitem(sys.modules, "hermes_cli.profiles", profiles) + gateway = ModuleType("hermes_cli.gateway") + gateway._is_service_running = lambda: calls["service_checks"].append(True) or False + gateway.ensure_gateway_service = lambda **kwargs: calls["service_starts"].append(kwargs) or True + monkeypatch.setitem(sys.modules, "hermes_cli.gateway", gateway) + monkeypatch.setattr(backup_mod, "_collect_memory_provider_external_paths", lambda: calls["provider_paths"].append(True) or []) + original_import = builtins.__import__ + def guarded_import(name, globals=None, locals=None, fromlist=(), level=0): + if name == "plugins.memory" or name.startswith("plugins.memory."): + calls["provider_imports"].append(name) + raise AssertionError("provider import during restore") + return original_import(name, globals, locals, fromlist, level) + monkeypatch.setattr(builtins, "__import__", guarded_import) + return calls + + @pytest.mark.parametrize("case", ["partial", "profile", "all-failed"]) + def test_import_enospc_reports_incomplete_and_exits_before_activation(self, tmp_path, monkeypatch, capsys, case): + from hermes_cli.backup import run_import + home = tmp_path / "inert-home" + home.mkdir() + failed = "profiles/coder/config.yaml" if case == "profile" else "config.yaml" + target = home / failed + target.parent.mkdir(parents=True, exist_ok=True) + old = b"model: original\n" + target.write_bytes(old) + monkeypatch.setenv("HERMES_HOME", str(home)) + monkeypatch.setattr(Path, "home", lambda: tmp_path) + members = {"config.yaml": "model: restored\n"} + if case == "partial": + members["notes.txt"] = "inert note\n" + elif case == "profile": + members[failed] = "model: restored-profile\n" + zipped = tmp_path / "backup.zip" + self._make_backup_zip(zipped, members) + calls = self._import_effect_spies(monkeypatch, home) + _break_member(monkeypatch, failed) + code = None + try: + run_import(Namespace(zipfile=str(zipped), force=True)) + except SystemExit as exc: + code = exc.code + output = capsys.readouterr().out + print({"pre": old, "post": target.read_bytes(), "exit": code, "calls": calls, "output": output}) + assert target.read_bytes() == old + assert list(target.parent.glob(".config.yaml.*")) == [] + assert code == 1 + assert all(not values for values in calls.values()) + restored = 0 if case == "all-failed" else 1 + assert f"Import incomplete: {restored} files restored, 1 failed" in output + assert "Import complete:" not in output + assert "Done. Your Hermes configuration has been restored." not in output + def _make_backup_zip(self, zip_path: Path, files: dict[str, str | bytes]) -> None: """Create a test zip with given files.""" with zipfile.ZipFile(zip_path, "w") as zf: @@ -769,17 +833,19 @@ def test_failed_member_leaves_existing_file_intact(self, tmp_path, monkeypatch): """A dying member must not destroy the file it was replacing.""" hermes_home = tmp_path / ".hermes" hermes_home.mkdir() - original = "model: original\napi_key: keep-me\n" + original = "model: original\nnote: keep-me\n" (hermes_home / "config.yaml").write_text(original) monkeypatch.setenv("HERMES_HOME", str(hermes_home)) monkeypatch.setattr(Path, "home", lambda: tmp_path) zip_path = tmp_path / "backup.zip" - self._zip(zip_path, {"config.yaml": "model: replacement\n", "state.db": ""}) + self._zip(zip_path, {"config.yaml": "model: replacement\n"}) _break_member(monkeypatch, "config.yaml") from hermes_cli.backup import run_import - run_import(Namespace(zipfile=str(zip_path), force=True)) + with pytest.raises(SystemExit) as caught: + run_import(Namespace(zipfile=str(zip_path), force=True)) + assert caught.value.code == 1 # Pre-fix this file is 0 bytes: the truncate landed, the write did not. assert (hermes_home / "config.yaml").read_text() == original @@ -808,7 +874,9 @@ def test_failed_external_member_leaves_existing_file_intact(self, tmp_path, monk _break_member(monkeypatch, "_external/.honcho/config.json") from hermes_cli.backup import run_import - run_import(Namespace(zipfile=str(zip_path), force=True)) + with pytest.raises(SystemExit) as caught: + run_import(Namespace(zipfile=str(zip_path), force=True)) + assert caught.value.code == 1 assert (honcho / "config.json").read_text() == original assert list(honcho.glob(".config.json.*")) == [] diff --git a/tests/hermes_cli/test_backup_atomic_publication.py b/tests/hermes_cli/test_backup_atomic_publication.py new file mode 100644 index 0000000000000..e3a5881b73c96 --- /dev/null +++ b/tests/hermes_cli/test_backup_atomic_publication.py @@ -0,0 +1,75 @@ +"""Direct import-publication coupling with inert ZIP members and errno.""" +import errno +import io +import os +import stat +import zipfile +from pathlib import Path + +import pytest +from hermes_cli import backup + +def archive(): + raw = io.BytesIO() + with zipfile.ZipFile(raw, "w") as zipped: + zipped.writestr("config.yaml", b"NEW-complete\n") + raw.seek(0) + return zipfile.ZipFile(raw) + +@pytest.mark.parametrize("kind", ["exdev", "ebusy"]) +def test_extract_fallback_failure_preserves_old_target_and_symlink(tmp_path, monkeypatch, kind): + tmp_path = tmp_path / "publication" + tmp_path.mkdir() + target, link = tmp_path / "real.txt", tmp_path / "config.yaml" + old = b"OLD-complete\n" + target.write_bytes(old) + link.symlink_to(target) + monkeypatch.setattr("utils.os.replace", lambda *args: (_ for _ in ()).throw(OSError(errno.EXDEV if kind == "exdev" else errno.EBUSY, "injected rename"))) + def copyfile(src, dst, **kwargs): + Path(dst).write_bytes(b"NEW") + raise OSError(errno.ENOSPC, "injected copy") + import shutil + original_copyfileobj = shutil.copyfileobj + def copyfileobj(src, dst, *args, **kwargs): + if isinstance(src, zipfile.ZipExtFile): + return original_copyfileobj(src, dst, *args, **kwargs) + dst.write(b"NEW") + raise OSError(errno.ENOSPC, "injected copy") + monkeypatch.setattr("utils.shutil.copyfile", copyfile) + monkeypatch.setattr("utils.shutil.copyfileobj", copyfileobj) + with archive() as zipped, pytest.raises(OSError): + backup._extract_member_atomically(zipped, "config.yaml", link, 0o640) + print({"pre": old, "post": target.read_bytes(), "kind": kind}) + assert target.read_bytes() == old and link.is_symlink() + assert sorted(p.name for p in tmp_path.iterdir()) == ["config.yaml", "real.txt"] + +@pytest.mark.parametrize("publication", ["rename", "exdev"]) +def test_extract_success_preserves_mode_and_owner_order(tmp_path, monkeypatch, publication): + tmp_path = tmp_path / "publication" + tmp_path.mkdir() + target = tmp_path / "config.yaml" + target.write_bytes(b"OLD-complete\n") + target.chmod(0o6640) + if publication == "exdev": + real_replace = os.replace + renames = [] + def replace(src, dst): + renames.append((src, dst)) + if len(renames) == 1: + raise OSError(errno.EXDEV, "injected first rename") + return real_replace(src, dst) + monkeypatch.setattr("utils.os.replace", replace) + calls = [] + monkeypatch.setattr(backup, "_preserve_file_owner", lambda path: (123, 456)) + monkeypatch.setattr(backup, "_restore_file_owner", lambda path, owner: calls.append(("owner", owner))) + restore_mode = backup._restore_file_mode + def mode(path, value): + calls.append(("mode", value)) + restore_mode(path, value) + monkeypatch.setattr(backup, "_restore_file_mode", mode) + with archive() as zipped: + backup._extract_member_atomically(zipped, "config.yaml", target, 0o600) + assert target.read_bytes() == b"NEW-complete\n" + assert stat.S_IMODE(target.stat().st_mode) == 0o640 + assert calls == [("owner", (123, 456)), ("mode", 0o640)] + assert sorted(p.name for p in tmp_path.iterdir()) == ["config.yaml"] diff --git a/tests/hermes_cli/test_mcp_logs_readers.py b/tests/hermes_cli/test_mcp_logs_readers.py new file mode 100644 index 0000000000000..80d046870b4ca --- /dev/null +++ b/tests/hermes_cli/test_mcp_logs_readers.py @@ -0,0 +1,431 @@ +"""Actual MCP log-reader integration with temporary stderr files only.""" + +import json +from datetime import datetime, timedelta +from pathlib import Path + +import pytest + + +@pytest.fixture +def runtime(tmp_path, monkeypatch): + import hermes_cli.logs as logs + import tools.mcp_tool as mcp_tool + from hermes_constants import reset_hermes_home_override, set_hermes_home_override + + home = tmp_path / "profile" + monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setattr(mcp_tool, "_mcp_stderr_log_files", {}) + token = set_hermes_home_override(str(home)) + try: + yield logs, mcp_tool, home + finally: + reset_hermes_home_override(token) + for fh in mcp_tool._mcp_stderr_log_files.values(): + fh.close() + + +def capture(runtime, text, server="cea_graph"): + _logs, mcp_tool, _home = runtime + record = mcp_tool._begin_stdio_diagnostic(server) + record["stream"].write(text) + record["stream"].flush() + record["status"] = "closed" + mcp_tool._finish_stdio_diagnostic(record) + return record + + +def manual_record(home, attempt="a" * 32, server="cea_graph", text="child payload\n"): + path = home / "logs" / "mcp-stderr" / f"{attempt}.log" + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(text) + record = { + "kind": "mcp.stdio.attempt", "attempt_id": attempt, "server": server, + "config_home": str(home), "parent_pid": 123, "stderr_path": str(path), + "destination": "file", "phase": "transport", "status": "starting", + } + return record + + +def write_index(home, records, legacy=""): + path = home / "logs" / "mcp-stderr.log" + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(legacy + "".join(json.dumps(record) + "\n" for record in records)) + return path + + +def test_actual_writer_and_reader_show_stderr_for_only_the_selected_profile( + runtime, tmp_path, capsys +): + from hermes_constants import reset_hermes_home_override, set_hermes_home_override + + logs, _mcp_tool, home = runtime + first = capture(runtime, "first alpha stderr\n") + token = set_hermes_home_override(str(tmp_path / "beta")) + try: + capture(runtime, "beta stderr must stay separate\n") + finally: + reset_hermes_home_override(token) + second = capture(runtime, "second alpha stderr\n", server="pilot_bridge") + logs.tail_log("mcp", num_lines=50) + output = capsys.readouterr().out + assert "first alpha stderr" in output + assert "second alpha stderr" in output + assert "beta stderr must stay separate" not in output + assert first["attempt_id"] in output and second["attempt_id"] in output + assert "cea_graph" in output and "pilot_bridge" in output + assert str(home) in output + # The reader retrieves actual content, without duplicating it in the index. + assert "first alpha stderr" not in (home / "logs" / "mcp-stderr.log").read_text() + assert '"kind": "mcp.stdio.attempt"' not in output + + +def test_tail_limit_applies_to_child_content_and_preserves_raw_json(runtime, capsys): + logs, _mcp_tool, home = runtime + record = manual_record(home, text='older\n{"kind":"child.event","message":"newest"}\n') + write_index(home, [record]) + logs.tail_log("mcp", num_lines=1) + output = capsys.readouterr().out + assert '"message":"newest"' in output + assert "older" not in output + assert record["attempt_id"] in output + + +def test_filters_use_child_timestamp_level_component_and_attempt(runtime, capsys): + logs, _mcp_tool, home = runtime + now = datetime.now() + old = (now - timedelta(hours=3)).strftime("%Y-%m-%d %H:%M:%S") + recent = now.strftime("%Y-%m-%d %H:%M:%S") + record = manual_record(home, text=( + f"{old} ERROR tools.reader: stale payload\n" + f"{recent} INFO tools.reader: lower level\n" + f"{recent} ERROR gateway.reader: wrong component\n" + f"{recent} ERROR tools.reader: matching child payload\n" + )) + write_index(home, [record]) + logs.tail_log("mcp", level="WARNING", since="1h", component="tools", session=record["attempt_id"]) + output = capsys.readouterr().out + assert "matching child payload" in output + assert "stale payload" not in output + assert "lower level" not in output + assert "wrong component" not in output + + +def test_legacy_stderr_and_new_attempts_are_both_readable(runtime, capsys): + logs, _mcp_tool, home = runtime + record = manual_record(home, text="attributed child payload\n") + write_index(home, [record], legacy='legacy raw stderr\n{"kind":"old.server","message":"legacy JSON"}\n') + logs.tail_log("mcp", num_lines=10) + output = capsys.readouterr().out + assert "legacy raw stderr" in output + assert '"message":"legacy JSON"' in output + assert "attributed child payload" in output + + +@pytest.mark.parametrize("payload", ["last unterminated child payload", "progress\r"]) +def test_final_unterminated_child_line_stays_separate_from_metadata(runtime, capsys, payload): + logs, _mcp_tool, _home = runtime + capture(runtime, payload) + logs.tail_log("mcp", num_lines=1) + output = capsys.readouterr().out + assert payload in output + assert '"kind": "mcp.stdio.attempt"' not in output + + +def test_tail_separates_partial_rows_from_two_active_attempts(runtime, capsys): + logs, mcp_tool, _home = runtime + first = mcp_tool._begin_stdio_diagnostic("cea_graph") + first["stream"].write("alpha partial") + first["stream"].flush() + second = mcp_tool._begin_stdio_diagnostic("pilot_bridge") + second["stream"].write("beta complete\n") + second["stream"].flush() + try: + logs.tail_log("mcp", num_lines=20) + assert Path(first["stderr_path"]).read_bytes().endswith(b"alpha partial") + finally: + for record in (second, first): + record["status"] = "closed" + mcp_tool._finish_stdio_diagnostic(record) + output = capsys.readouterr().out + assert "alpha partial\n" in output + assert "beta complete\n" in output + + +def test_tail_to_follow_separates_the_preview_and_completes_its_partial_line(runtime, monkeypatch, capsys): + logs, mcp_tool, _home = runtime + record = mcp_tool._begin_stdio_diagnostic("cea_graph") + record["stream"].write("alpha partial") + record["stream"].flush() + iterations = 0 + + def advance(_seconds): + nonlocal iterations + iterations += 1 + if iterations == 1: + record["stream"].write(" remainder\nfresh complete child\n") + record["stream"].flush() + else: + raise KeyboardInterrupt + + monkeypatch.setattr(logs.time, "sleep", advance) + try: + logs.tail_log("mcp", num_lines=20, follow=True) + finally: + record["status"] = "closed" + mcp_tool._finish_stdio_diagnostic(record) + output = capsys.readouterr().out + assert "alpha partial\n" in output + assert output.count("alpha partial remainder\n") == 1 + assert output.count("fresh complete child\n") == 1 + + +@pytest.mark.parametrize("invalid", ["wrong_profile", "outside_path", "symlink", "bad_id", "discarded", "missing", "invalid_pid"]) +def test_index_cannot_read_an_unowned_or_unavailable_file(runtime, tmp_path, capsys, invalid): + logs, _mcp_tool, home = runtime + record = manual_record(home) + foreign = tmp_path / "other-profile" / "foreign.log" + foreign.parent.mkdir() + foreign.write_text("FOREIGN_CONTENT_MUST_NOT_BE_READ\n") + if invalid == "wrong_profile": + record["config_home"] = str(foreign.parent) + elif invalid == "outside_path": + record["stderr_path"] = str(foreign) + elif invalid == "symlink": + path = Path(record["stderr_path"]) + path.unlink() + path.symlink_to(foreign) + elif invalid == "bad_id": + record["attempt_id"] = "../foreign" + elif invalid == "discarded": + record["destination"] = "discarded" + elif invalid == "missing": + Path(record["stderr_path"]).unlink() + elif invalid == "invalid_pid": + record["parent_pid"] = True + write_index(home, [record]) + logs.tail_log("mcp", num_lines=10) + output = capsys.readouterr().out + assert "FOREIGN_CONTENT_MUST_NOT_BE_READ" not in output + assert "child payload" not in output + + +def test_follow_reads_existing_and_new_attempts_without_replaying_history( + runtime, monkeypatch, capsys +): + logs, mcp_tool, home = runtime + active = mcp_tool._begin_stdio_diagnostic("cea_graph") + active["stream"].write("historical child payload\n") + active["stream"].flush() + iterations = 0 + + def advance(_seconds): + nonlocal iterations + iterations += 1 + if iterations == 1: + active["stream"].write("fresh active stderr\n") + active["stream"].flush() + capture(runtime, "fresh next-attempt stderr\n", server="pilot_bridge") + else: + raise KeyboardInterrupt + + monkeypatch.setattr(logs.time, "sleep", advance) + try: + logs.tail_log("mcp", num_lines=0, follow=True) + finally: + active["status"] = "closed" + mcp_tool._finish_stdio_diagnostic(active) + output = capsys.readouterr().out + assert "historical child payload" not in output + assert output.count("fresh active stderr") == 1 + assert output.count("fresh next-attempt stderr") == 1 + assert active["attempt_id"] in output + assert str(home) in output + + +def test_follow_preserves_legacy_appends_and_handles_attempt_replacement( + runtime, monkeypatch, capsys +): + logs, _mcp_tool, home = runtime + record = manual_record(home, text="existing child payload\n") + index = write_index(home, [record], legacy="legacy history\n") + iterations = 0 + + def advance(_seconds): + nonlocal iterations + iterations += 1 + if iterations == 1: + with index.open("a") as stream: + stream.write("legacy fresh stderr\n") + path = Path(record["stderr_path"]) + replacement = path.with_suffix(".new") + replacement.write_text("replacement child stderr\n") + replacement.replace(path) + else: + raise KeyboardInterrupt + + monkeypatch.setattr(logs.time, "sleep", advance) + logs.tail_log("mcp", num_lines=0, follow=True) + output = capsys.readouterr().out + assert output.count("legacy fresh stderr") == 1 + assert output.count("replacement child stderr") == 1 + assert "existing child payload" not in output + + +def test_follow_preserves_an_index_record_across_the_read_budget(runtime, monkeypatch, capsys): + logs, _mcp_tool, home = runtime + index = write_index(home, []) + record = manual_record(home, text="boundary child stderr\n") + iterations = 0 + + def advance(_seconds): + nonlocal iterations + iterations += 1 + if iterations == 1: + with index.open("a") as stream: + stream.write("x" * (65536 - 32) + "\n" + json.dumps(record) + "\n") + elif iterations > 2: + raise KeyboardInterrupt + + monkeypatch.setattr(logs.time, "sleep", advance) + logs.tail_log("mcp", num_lines=0, follow=True) + output = capsys.readouterr().out + assert output.count("boundary child stderr") == 1 + assert '"kind": "mcp.stdio.attempt"' not in output + + +@pytest.mark.parametrize("filtered", [False, True]) +def test_follow_frames_a_child_line_written_across_polls(runtime, monkeypatch, capsys, filtered): + logs, _mcp_tool, home = runtime + record = manual_record(home, text="") + write_index(home, [record]) + path = Path(record["stderr_path"]) + stamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S") + first = f"{stamp} ERROR tools.reader: split child " + iterations = 0 + + def advance(_seconds): + nonlocal iterations + iterations += 1 + if iterations == 1: + with path.open("a") as stream: + stream.write(first) + elif iterations == 2: + assert "split child" not in capsys.readouterr().out + with path.open("a") as stream: + stream.write("payload\n") + else: + raise KeyboardInterrupt + + monkeypatch.setattr(logs.time, "sleep", advance) + options = {"level": "WARNING", "component": "tools", "since": "1h"} if filtered else {} + logs.tail_log("mcp", num_lines=0, follow=True, **options) + output = capsys.readouterr().out + assert output.count(first + "payload\n") == 1 + assert output.count("[mcp ") == 1 + + +def test_follow_frames_a_control_record_written_across_polls(runtime, monkeypatch, capsys): + logs, _mcp_tool, home = runtime + index = write_index(home, []) + record = manual_record(home, text="split-index child payload\n") + encoded = json.dumps(record) + iterations = 0 + + def advance(_seconds): + nonlocal iterations + iterations += 1 + if iterations == 1: + with index.open("a") as stream: + stream.write(encoded[:40]) + elif iterations == 2: + assert '"kind"' not in capsys.readouterr().out + with index.open("a") as stream: + stream.write(encoded[40:] + "\n") + else: + raise KeyboardInterrupt + + monkeypatch.setattr(logs.time, "sleep", advance) + logs.tail_log("mcp", num_lines=0, follow=True) + output = capsys.readouterr().out + assert output.count("split-index child payload") == 1 + assert '"kind": "mcp.stdio.attempt"' not in output + + +@pytest.mark.parametrize("filtered", [False, True]) +def test_follow_advances_past_an_oversized_line_with_a_visible_notice(runtime, monkeypatch, capsys, filtered): + logs, _mcp_tool, home = runtime + monkeypatch.setattr(logs, "_MCP_FOLLOW_BYTES", 512, raising=False) + monkeypatch.setattr(logs, "_MCP_MAX_LINE_BYTES", 2048, raising=False) + record = manual_record(home, text="") + write_index(home, [record]) + path = Path(record["stderr_path"]) + stamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S") + iterations = 0 + + def advance(_seconds): + nonlocal iterations + iterations += 1 + if iterations == 1: + with path.open("a") as stream: + stream.write("x" * 4097 + f"\n{stamp} ERROR tools.reader: after oversized child line\n") + elif iterations >= 12: + raise KeyboardInterrupt + + monkeypatch.setattr(logs.time, "sleep", advance) + options = {"level": "WARNING", "component": "tools", "since": "1h"} if filtered else {} + logs.tail_log("mcp", num_lines=0, follow=True, **options) + output = capsys.readouterr().out + assert output.count("exceeds MCP preview limit") == 1 + assert output.count("after oversized child line") == 1 + + +@pytest.mark.parametrize("filtered", [False, True]) +def test_tail_has_a_finite_window_and_keeps_recent_content(runtime, monkeypatch, capsys, filtered): + logs, _mcp_tool, home = runtime + monkeypatch.setattr(logs, "_MCP_TAIL_BYTES", 2048, raising=False) + monkeypatch.setattr(logs, "_MCP_MAX_LINE_BYTES", 1024, raising=False) + stamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S") + record = manual_record(home, text="x" * 4096 + f"\n{stamp} ERROR tools.reader: recent child payload\n") + write_index(home, [record]) + options = {"level": "WARNING", "component": "tools", "since": "1h"} if filtered else {} + logs.tail_log("mcp", num_lines=10, **options) + output = capsys.readouterr().out + assert "recent child payload" in output + assert "exceeds MCP preview limit" in output + assert "x" * 2048 not in output + + +def test_follow_discards_a_pending_fragment_on_observed_truncation(runtime, monkeypatch, capsys): + logs, _mcp_tool, home = runtime + record = manual_record(home, text="") + write_index(home, [record]) + path = Path(record["stderr_path"]) + iterations = 0 + + def advance(_seconds): + nonlocal iterations + iterations += 1 + if iterations == 1: + path.write_text("stale fragment" * 100) + elif iterations == 2: + path.write_text("replacement complete\n") + else: + raise KeyboardInterrupt + + monkeypatch.setattr(logs.time, "sleep", advance) + logs.tail_log("mcp", num_lines=0, follow=True) + output = capsys.readouterr().out + assert "stale fragment" not in output + assert output.count("replacement complete\n") == 1 + + +def test_non_mcp_reader_keeps_its_existing_behavior(runtime, capsys): + logs, _mcp_tool, home = runtime + path = home / "logs" / "agent.log" + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text("older agent row\nnewest agent row\n") + logs.tail_log("agent", num_lines=1) + output = capsys.readouterr().out + assert "newest agent row" in output and "older agent row" not in output + assert "[mcp " not in output diff --git a/tests/run_agent/test_ares_context_handoff_refusal.py b/tests/run_agent/test_ares_context_handoff_refusal.py new file mode 100644 index 0000000000000..0dfccf810822e --- /dev/null +++ b/tests/run_agent/test_ares_context_handoff_refusal.py @@ -0,0 +1,274 @@ +"""An Ares budget refusal must leave the owning host on its previous route. + +These follow-up counterexamples use real AIAgent handoffs and the real governor +budget validator. Provider clients, metadata, activation and native commands are +fakes; no conversation or provider request runs. +""" + +from unittest.mock import MagicMock, patch + +import pytest + +from agent.error_classifier import FailoverReason +from tests.run_agent.test_ares_context_initialization import _host_init, governor + + +_HOST_FIELDS = ( + "model", "provider", "requested_provider", "base_url", "_base_url_lower", + "api_mode", "api_key", "client", "_anthropic_client", "_anthropic_api_key", + "_anthropic_base_url", "_is_anthropic_oauth", "_client_kwargs", + "_credential_pool", "_credential_pool_entry_id", "_config_context_length", + "_custom_providers", "_reasoning_echo_flag", "_use_prompt_caching", + "_use_native_cache_layout", "_transport_cache", "_fallback_activated", + "_primary_runtime", "_cached_system_prompt", "reasoning_config", + "_provider_fallback_active", "_provider_fallback_route", + "_pending_fallback_notice", +) +_ENGINE_FIELDS = ( + "model", "provider", "base_url", "api_key", "api_mode", "context_length", + "max_tokens", "threshold_tokens", "threshold_percent", "protect_first_n", + "protect_last_n", "session_id", "_lineage_session_id", "_session_db", + "_pending_admission", "last_receipt_id", "compression_count", +) + + +@pytest.fixture(autouse=True) +def _fake_catalog_prewarm(): + # Initializer prewarming uses an imported alias in its background thread. + # Both metadata entry points stay fake, in addition to the worker guard. + with ( + patch("agent.agent_init.fetch_model_metadata", return_value={}), + patch("agent.model_metadata.fetch_model_metadata", return_value={}), + ): + yield + + +def _state(obj, fields): + values = {} + for name in fields: + if not hasattr(obj, name): + continue + value = getattr(obj, name) + if isinstance(value, dict): + value = dict(value) + elif isinstance(value, list): + value = list(value) + values[name] = value + return values + + +def _sentinel_runtime(agent): + # Make silent collateral changes observable, including in-place cache + # clearing and re-reading a custom-provider snapshot during a switch. + agent._transport_cache = {"previous-route": object()} + agent._cached_system_prompt = "stable previous prompt" + agent._credential_pool = MagicMock(provider=agent.provider) + agent._credential_pool_entry_id = "previous-credential" + agent._custom_providers = [{"name": "previous-route"}] + agent._use_prompt_caching = True + agent._use_native_cache_layout = True + agent._config_context_length = 64_000 + + +def _fallback_client(): + client = MagicMock(name="FakeFallbackClient") + client.base_url = "https://api.openai.com/v1" + client.api_key = "fake-fallback-key" + return client + + +@pytest.mark.parametrize("api_mode", ["chat_completions", "anthropic_messages"]) +def test_refused_manual_switch_preserves_host_and_engine(governor, api_mode): + with _host_init(governor) as create: + agent = create(max_tokens=4096) + _sentinel_runtime(agent) + before_host = _state(agent, _HOST_FIELDS) + before_engine = _state(governor, _ENGINE_FIELDS) + previous_client = agent.client + previous_pool = agent._credential_pool + previous_cache = agent._transport_cache + new_client = MagicMock(name="RefusedClient") + retire = MagicMock() + with ( + patch("agent.model_metadata.get_model_context_length", return_value=4096), + patch("agent.credential_pool.load_pool", return_value=None), + patch("agent.anthropic_adapter.build_anthropic_client", return_value=new_client), + patch.object(agent, "_create_openai_client", return_value=new_client), + patch.object(agent, "_retire_shared_openai_client", retire), + ): + with pytest.raises(ValueError, match="no input budget"): + agent.switch_model( + "refused-model", "openai", api_key="fake-new-key", + base_url="https://api.openai.com/v1", api_mode=api_mode, + ) + assert _state(agent, _HOST_FIELDS) == before_host + assert _state(governor, _ENGINE_FIELDS) == before_engine + assert agent.client is previous_client + assert agent._credential_pool is previous_pool + assert agent._transport_cache is previous_cache + retire.assert_called_once_with(new_client, reason="context_handoff_refused") + previous_client.close.assert_not_called() + new_client.close.assert_not_called() + new_client.chat.completions.create.assert_not_called() + new_client.messages.create.assert_not_called() + + +@pytest.mark.parametrize("already_on_fallback", [False, True]) +def test_refused_exhausted_fallback_preserves_active_runtime(governor, already_on_fallback): + with _host_init(governor) as create: + agent = create(max_tokens=4096) + if already_on_fallback: + agent._fallback_chain = [{"provider": "openai", "model": "previous-fallback"}] + with ( + patch("agent.auxiliary_client.resolve_provider_client", return_value=(_fallback_client(), None)), + patch("agent.model_metadata.get_model_context_length", return_value=32_000), + patch("agent.credential_pool.load_pool", return_value=None), + ): + assert agent._try_activate_fallback() is True + agent._fallback_index = 0 + _sentinel_runtime(agent) + agent._fallback_chain = [{"provider": "openai", "model": "refused-fallback"}] + agent._fallback_model = agent._fallback_chain[0] + before_host = _state(agent, _HOST_FIELDS) + before_engine = _state(governor, _ENGINE_FIELDS) + fallback_client = _fallback_client() + with ( + patch("agent.auxiliary_client.resolve_provider_client", return_value=(fallback_client, None)), + patch("agent.model_metadata.get_model_context_length", return_value=4096), + patch("agent.credential_pool.load_pool", return_value=None), + ): + assert agent._try_activate_fallback() is False + assert _state(agent, _HOST_FIELDS) == before_host + assert _state(governor, _ENGINE_FIELDS) == before_engine + # Preserve intentional chain progress and exhaustion cooldown. Rolling + # those back would retry the same rejected route indefinitely. + assert agent._fallback_index == 1 + assert agent._rate_limited_until > 0 + fallback_client.responses.create.assert_not_called() + + +def test_refused_fallback_advances_to_a_valid_candidate(governor): + with _host_init(governor) as create: + agent = create(max_tokens=4096) + agent._fallback_chain = [ + {"provider": "openai", "model": "refused-fallback"}, + {"provider": "openai", "model": "valid-fallback"}, + ] + clients = [_fallback_client(), _fallback_client()] + with ( + patch("agent.auxiliary_client.resolve_provider_client", side_effect=[(c, None) for c in clients]), + patch("agent.model_metadata.get_model_context_length", side_effect=[4096, 32_000]), + patch("agent.credential_pool.load_pool", return_value=None), + ): + assert agent._try_activate_fallback() is True + assert agent.model == governor.model == "valid-fallback" + assert agent.client is clients[1] + assert governor.context_length == 32_000 + assert governor.max_tokens == 4096 + assert agent._fallback_index == 2 + assert agent._provider_fallback_route == ("valid-fallback", "openai") + assert len(agent._pending_fallback_notice) == 1 + assert "test-model via openrouter" in agent._pending_fallback_notice[0] + assert "using valid-fallback" in agent._pending_fallback_notice[0] + + +@pytest.mark.parametrize( + "reason", [FailoverReason.rate_limit, FailoverReason.billing, FailoverReason.upstream_rate_limit] +) +@pytest.mark.parametrize("has_valid_tail", [False, True]) +def test_one_primary_failure_arms_backoff_once_through_refused_routes( + governor, reason, has_valid_tail +): + with _host_init(governor) as create: + agent = create(max_tokens=4096) + agent._fallback_chain = [{"provider": "openai", "model": "refused-fallback"}] + if has_valid_tail: + agent._fallback_chain.append({"provider": "openai", "model": "valid-fallback"}) + clients = [_fallback_client() for _ in agent._fallback_chain] + windows = [4096, 32_000] if has_valid_tail else [4096] + before_host = _state(agent, _HOST_FIELDS) + before_engine = _state(governor, _ENGINE_FIELDS) + with ( + patch("agent.chat_completion_helpers.time.monotonic", return_value=1000.0), + patch("agent.auxiliary_client.resolve_provider_client", side_effect=[(c, None) for c in clients]), + patch("agent.model_metadata.get_model_context_length", side_effect=windows), + patch("agent.credential_pool.load_pool", return_value=None), + ): + assert agent._try_activate_fallback(reason) is has_valid_tail + assert agent._rate_limit_backoff_count == 1 + assert agent._rate_limited_until == 1060.0 + if has_valid_tail: + assert agent.model == governor.model == "valid-fallback" + assert agent._fallback_index == 2 + else: + assert _state(agent, _HOST_FIELDS) == before_host + assert _state(governor, _ENGINE_FIELDS) == before_engine + assert agent._fallback_index == 1 + # A distinct originating primary failure still advances the + # backoff even with an already exhausted fallback chain. + assert agent._try_activate_fallback(reason) is False + assert agent._rate_limit_backoff_count == 2 + assert agent._rate_limited_until == 1120.0 + + +def test_refused_primary_restore_preserves_fallback_then_allows_retry(governor): + with _host_init(governor) as create: + agent = create(max_tokens=4096) + agent._fallback_chain = [{"provider": "openai", "model": "valid-fallback"}] + agent._fallback_model = agent._fallback_chain[0] + fallback_client = _fallback_client() + with ( + patch("agent.auxiliary_client.resolve_provider_client", return_value=(fallback_client, None)), + patch("agent.model_metadata.get_model_context_length", return_value=32_000), + patch("agent.credential_pool.load_pool", return_value=None), + ): + assert agent._try_activate_fallback() is True + _sentinel_runtime(agent) + before_host = _state(agent, _HOST_FIELDS) + before_engine = _state(governor, _ENGINE_FIELDS) + original_window = agent._primary_runtime["compressor_context_length"] + agent._primary_runtime["compressor_context_length"] = 4096 + # The saved intended destination may be invalid after reserve changes; + # the active fallback must remain internally coherent on refusal. + before_host["_primary_runtime"]["compressor_context_length"] = 4096 + assert agent._restore_primary_runtime() is False + assert _state(agent, _HOST_FIELDS) == before_host + assert _state(governor, _ENGINE_FIELDS) == before_engine + assert agent.client is fallback_client + agent._primary_runtime["compressor_context_length"] = original_window + assert agent._restore_primary_runtime() is True + assert agent.model == governor.model == "test-model" + assert agent._fallback_activated is False + assert agent._provider_fallback_active is False + assert governor.context_length == original_window + assert governor.max_tokens == 4096 + + +@pytest.mark.parametrize("missing_field", ["compressor_model", "compressor_context_length"]) +@pytest.mark.parametrize("falsey_engine", [False, True]) +def test_active_engine_requires_restore_snapshot_fields( + governor, monkeypatch, missing_field, falsey_engine +): + with _host_init(governor) as create: + agent = create(max_tokens=4096) + agent._fallback_chain = [{"provider": "openai", "model": "valid-fallback"}] + with ( + patch("agent.auxiliary_client.resolve_provider_client", return_value=(_fallback_client(), None)), + patch("agent.model_metadata.get_model_context_length", return_value=32_000), + patch("agent.credential_pool.load_pool", return_value=None), + ): + assert agent._try_activate_fallback() is True + if falsey_engine: + monkeypatch.setattr(type(governor), "__bool__", lambda self: False, raising=False) + saved_value = agent._primary_runtime.pop(missing_field) + before_host = _state(agent, _HOST_FIELDS) + before_engine = _state(governor, _ENGINE_FIELDS) + previous_client = agent.client + assert agent._restore_primary_runtime() is False + assert _state(agent, _HOST_FIELDS) == before_host + assert _state(governor, _ENGINE_FIELDS) == before_engine + assert agent.client is previous_client + agent._primary_runtime[missing_field] = saved_value + assert agent._restore_primary_runtime() is True + assert agent.model == governor.model == "test-model" + assert agent._fallback_activated is False diff --git a/tests/run_agent/test_ares_context_initialization.py b/tests/run_agent/test_ares_context_initialization.py new file mode 100644 index 0000000000000..71ef990d02861 --- /dev/null +++ b/tests/run_agent/test_ares_context_initialization.py @@ -0,0 +1,229 @@ +"""Ares engine selection and resumed lineage through the real host initializer. + +All provider clients, activation probes and governor commands are fakes. The +session ancestry is read from an isolated real SessionDB, not an adapter stub. +""" + +import importlib +from contextlib import ExitStack, contextmanager +from unittest.mock import MagicMock, patch + +import pytest + +from hermes_state import SessionDB +from plugins.context_engine import ContextEngineActivationError +from agent.context_compressor import ContextCompressor + + +@pytest.fixture +def governor(tmp_path): + engine_type = importlib.import_module( + "plugins.context_engine.ri-context-governor" + ).RiContextGovernorEngine + with patch("hermes_cli.config.load_config", return_value={}): + engine = engine_type( + binary=str(tmp_path / "fake-governor"), + store_dir=str(tmp_path / "governor-store"), + ) + engine.probe_activation = MagicMock(return_value={"verified": True}) + engine._run_certified_json = MagicMock(return_value=[]) + engine._load_prior_session_context = MagicMock() + return engine + + +@contextmanager +def _host_init(engine, *, engine_name="ri-context-governor"): + cfg = { + "context": {"engine": engine_name}, + "compression": {"threshold": 0.75}, + "agent": {"environment_probe": False}, + } + with ExitStack() as stack: + for target, value in ( + ("hermes_cli.config.load_config", cfg), + ("hermes_cli.config.load_config_readonly", cfg), + ("plugins.context_engine._load_engine_from_dir", engine), + ("agent.model_metadata.get_model_context_length", 64_000), + ("agent.context_compressor.get_model_context_length", 64_000), + ("run_agent.get_tool_definitions", []), + ("run_agent.check_toolset_requirements", {}), + ): + stack.enter_context(patch(target, return_value=value)) + stack.enter_context(patch("run_agent.OpenAI")) + if engine_name != "ri-context-governor": + stack.enter_context( + patch("plugins.context_engine.load_context_engine", return_value=engine) + ) + from run_agent import AIAgent + + def create(**kwargs): + options = dict( + model="test-model", + provider="openrouter", + api_key="fake-provider-key", + base_url="https://provider.invalid/v1", + quiet_mode=True, + skip_context_files=True, + skip_memory=True, + skip_background_review=True, + ) + options.update(kwargs) + return AIAgent(**options) + + yield create + + +def test_explicit_ares_selection_refuses_missing_engine(): + with ( + _host_init(None) as create, + patch("agent.agent_init.ContextCompressor", wraps=ContextCompressor) as stock, + patch("hermes_cli.plugins.get_plugin_context_engine") as optional, + ): + with pytest.raises(ContextEngineActivationError, match="instantiated"): + create() + stock.assert_not_called() + optional.assert_not_called() + + +@pytest.mark.parametrize( + ("stage", "method", "message"), + [ + ("probe", "probe_activation", "activation probe"), + ("binding", "bind_session_state", "session binding"), + ("startup", "on_session_start", "session start"), + ], +) +def test_explicit_ares_selection_surfaces_startup_failure( + governor, stage, method, message +): + setattr(governor, method, MagicMock(side_effect=RuntimeError(f"{stage} failed"))) + with ( + _host_init(governor) as create, + patch("agent.agent_init.ContextCompressor") as stock, + ): + with pytest.raises(ContextEngineActivationError, match=message) as failure: + create() + assert isinstance(failure.value.__cause__, RuntimeError) + stock.assert_not_called() + + +def test_explicit_ares_selection_probes_before_session_binding(governor): + events = [] + governor.probe_activation.side_effect = lambda: events.append("probe") + original_bind = governor.bind_session_state + + def bind(**kwargs): + events.append("bind") + return original_bind(**kwargs) + + governor.bind_session_state = bind + with _host_init(governor) as create: + agent = create(session_id="new-session", max_tokens=4096) + assert agent.context_compressor is governor + assert events[0:2] == ["probe", "bind"] + governor.probe_activation.assert_called_once() + assert governor.max_tokens == 4096 + assert governor.threshold_tokens == int((64_000 - 4096) * 0.75) + + +@pytest.mark.parametrize( + ("end_reason", "model_config", "expected_root"), + [ + ("compression", {}, "root"), + ("compression", {"_branched_from": "root"}, "tip"), + ("compression", {"_delegate_from": "root"}, "tip"), + ( + "context_rebase", + { + "_context_rebase_from": "root", + "_context_rebase_transition": "transition-fixture", + "_context_epoch": 1, + }, + "root", + ), + ("context_rebase", {}, "tip"), + ], +) +def test_session_start_keeps_the_sessiondb_lineage( + governor, tmp_path, end_reason, model_config, expected_root +): + db = SessionDB(tmp_path / "session-state.db") + try: + db.create_session("root", "cli") + db.end_session("root", end_reason) + db.create_session( + "tip", "cli", parent_session_id="root", model_config=model_config + ) + with _host_init(governor) as create: + create(session_id="tip", session_db=db) + assert governor.session_id == "tip" + assert governor._session_db is db + assert governor._governor_session_id() == expected_root + finally: + db.close() + + +def test_optional_engine_missing_still_uses_stock_compressor(): + with ( + _host_init(None, engine_name="optional-engine") as create, + patch("hermes_cli.plugins.get_plugin_context_engine", return_value=None), + ): + agent = create() + from agent.context_compressor import ContextCompressor + + assert isinstance(agent.context_compressor, ContextCompressor) + + +def test_optional_engine_startup_remains_best_effort(governor): + governor.probe_activation.side_effect = AssertionError("must not strict-probe") + governor.bind_session_state = MagicMock(side_effect=RuntimeError("optional bind")) + governor.on_session_start = MagicMock(side_effect=RuntimeError("optional start")) + with _host_init(governor, engine_name="optional-engine") as create: + agent = create() + assert agent.context_compressor is governor + governor.probe_activation.assert_not_called() + assert "session_db" not in governor.on_session_start.call_args.kwargs + + +def test_manual_model_switch_keeps_the_output_reserve(governor): + governor._policy["token_budget"] = 128_000 + with _host_init(governor) as create: + agent = create(max_tokens=4096) + with patch("agent.model_metadata.get_model_context_length", return_value=32_000): + agent.switch_model( + "switched-model", + "openrouter", + api_key="fake-provider-key", + base_url="https://provider.invalid/v1", + ) + assert governor.model == "switched-model" + assert governor.max_tokens == 4096 + assert governor.context_length == 32_000 + assert governor._target_tokens(100_000) == 27_904 + + +def test_fallback_and_primary_restore_keep_the_output_reserve(governor): + governor._policy["token_budget"] = 128_000 + fallback_client = MagicMock() + fallback_client.base_url = "https://api.openai.com/v1" + fallback_client.api_key = "fake-fallback-key" + with _host_init(governor) as create: + agent = create(max_tokens=4096) + agent._fallback_chain = [{"provider": "openai", "model": "fallback-model"}] + agent._fallback_model = agent._fallback_chain[0] + with ( + patch( + "agent.auxiliary_client.resolve_provider_client", + return_value=(fallback_client, None), + ), + patch("agent.model_metadata.get_model_context_length", return_value=32_000), + ): + assert agent._try_activate_fallback() is True + assert governor.model == "fallback-model" + assert governor.max_tokens == 4096 + assert governor._target_tokens(100_000) == 27_904 + assert agent._restore_primary_runtime() is True + assert governor.model == "test-model" + assert governor.max_tokens == 4096 + assert governor.context_length == 64_000 + assert governor._target_tokens(100_000) == 59_904 diff --git a/tests/run_agent/test_context_pressure_final_admission.py b/tests/run_agent/test_context_pressure_final_admission.py new file mode 100644 index 0000000000000..3b8b34e642b2e --- /dev/null +++ b/tests/run_agent/test_context_pressure_final_admission.py @@ -0,0 +1,112 @@ +"""Real host gates with actual Governor pressure state; SDK and compaction effects inert.""" +import json +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest + +from plugins.context_engine._context_governor import ContextGovernorEngine +from run_agent import AIAgent + + +def response(prompt=None, tool=False, index=0): + calls = [SimpleNamespace(id=f"call-{index}", type="function", function=SimpleNamespace( + name="lookup", arguments=json.dumps({"query": f"step-{index}"})))] if tool else None + msg = SimpleNamespace(content=None if tool else "done", reasoning_content=None, + reasoning=None, tool_calls=calls) + usage = None if prompt is None else SimpleNamespace( + prompt_tokens=prompt, completion_tokens=1, total_tokens=prompt + 1) + return SimpleNamespace(choices=[SimpleNamespace(message=msg, + finish_reason="tool_calls" if tool else "stop")], model="test-model", usage=usage) + + +@pytest.fixture +def agent(tmp_path): + tools = [{"type": "function", "function": {"name": "lookup", + "description": "inert lookup", "parameters": {"type": "object", + "properties": {"conversation": {"type": "string"}}}}}] + with (patch("run_agent.get_tool_definitions", return_value=tools), + patch("run_agent.check_toolset_requirements", return_value={}), + patch("tools.env_probe.warm_environment_probe_async"), + patch("run_agent.OpenAI"), patch("hermes_cli.config.load_config", return_value={})): + value = AIAgent(api_key="test-key-1234567890", base_url="https://example.invalid", + quiet_mode=True, skip_context_files=True, skip_memory=True, + max_iterations=20) + engine = ContextGovernorEngine(binary=str(tmp_path / "no-native"), + store_dir=str(tmp_path / "governor")) + engine.update_model("test-model", 100_000, max_tokens=10_000, + provider="test", api_mode="chat_completions", threshold_percent=0.5) + engine.last_real_prompt_tokens = 20_000 + engine.last_rough_tokens_when_real_prompt_fit = 46_000 + engine.last_compression_rough_tokens = 46_000 + value.context_compressor = engine + value.model = "test-model" + value.max_tokens = 10_000 + value.client = MagicMock() + value._cached_system_prompt = "Helpful." + value._use_prompt_caching = False + value._disable_streaming = True + value.save_trajectories = False + value.compression_enabled = True + value.tool_delay = 0 + value._environment_probe = False + value.context_rebase_enabled = False + return value + + +def run(agent, texts, history=None, *, responses=None, tool_result=None): + compactions = [] + provider_payloads = [] + replies = iter(responses or [response() for _ in texts]) + def inert_create(**kwargs): + provider_payloads.append(kwargs) + return next(replies) + def inert_compact(messages, system, **kwargs): + compactions.append(kwargs.get("approx_tokens")) + engine = agent.context_compressor + engine.last_prompt_tokens = -1 + engine.awaiting_real_usage_after_compression = True + compacted = [dict(m, content="[summary]") + if len(str(m.get("content") or "")) > 5000 else m for m in messages] + return compacted, "Helpful." + agent.client.chat.completions.create.side_effect = inert_create + with (patch.object(agent, "_compress_context", side_effect=inert_compact), + patch.object(agent, "_persist_session"), patch.object(agent, "_save_trajectory"), + patch.object(agent, "_cleanup_task_resources"), + patch("run_agent.handle_function_call", return_value=json.dumps({"result": tool_result or "ok"}))): + results = [] + for text in texts: + result = agent.run_conversation(text, conversation_history=history) + results.append(result) + history = result["messages"] + return results, compactions, provider_payloads + + +def test_actual_plain_host_compacts_gradual_no_usage_growth(agent): + history = [{"role": "user", "content": "h" * 180_000}, + {"role": "assistant", "content": "previous answer"}] + results, compactions, calls = run(agent, ["x" * 12_000] * 20, history) + print(json.dumps({"turns": len(results), "provider_calls": len(calls), + "compactions": len(compactions)})) + assert all(not r["failed"] for r in results) + assert len(calls) == 20 + assert compactions + assert agent.context_compressor.last_rough_tokens_when_real_prompt_fit == 46_000 + + +@pytest.mark.parametrize("prompt", [None, 20_000]) +def test_actual_tool_host_rebuilds_request_with_usage_or_missing_usage(agent, prompt): + history = [{"role": "user", "content": "h" * 180_000}, + {"role": "assistant", "content": "previous answer"}] + results, compactions, calls = run(agent, ["look up"], history, responses=[ + *[response(prompt, tool=True, index=i) for i in range(5)], + response(prompt)], tool_result="t" * 35_000) + assert not results[0]["failed"] + assert len(calls) == 6 + assert compactions + # Late/custom schema remains on the actual request; counting separately + # is not evidence that a native/provider request was authorized. + assert calls[-1]["tools"][0]["function"]["parameters"]["properties"]["conversation"] + if prompt: + assert agent.context_compressor.last_rough_tokens_when_real_prompt_fit > 0 + assert agent.context_compressor.last_real_prompt_tokens == prompt diff --git a/tests/run_agent/test_plugin_context_engine_init.py b/tests/run_agent/test_plugin_context_engine_init.py index 4961ec2cb2a1f..253197f8d7f51 100644 --- a/tests/run_agent/test_plugin_context_engine_init.py +++ b/tests/run_agent/test_plugin_context_engine_init.py @@ -6,6 +6,8 @@ from unittest.mock import MagicMock, patch +import pytest + from agent.context_engine import ContextEngine @@ -221,3 +223,80 @@ def test_codex_gpt55_autoraise_still_applies_to_builtin_compressor(): assert agent._compression_warning and "85%" in agent._compression_warning +def _context_policy_agent(engine, *, enabled, disabled): + """Run the real initializer; only provider and schema-source seams are inert.""" + cfg = {"context": {"engine": "stub"}, "agent": {}} + base_schema = {"type": "function", "function": { + "name": "policy_keep", "description": "Inert schema witness", + "parameters": {"type": "object", "properties": {}}, + }} + with ( + patch("hermes_cli.config.load_config", return_value=cfg), + patch("hermes_cli.config.load_config_readonly", return_value=cfg), + patch("plugins.context_engine.load_context_engine", return_value=engine), + patch("agent.model_metadata.get_model_context_length", return_value=204_800), + patch("run_agent.get_tool_definitions", side_effect=lambda **kw: [base_schema]), + patch("run_agent.check_toolset_requirements", return_value={}), + patch("run_agent.OpenAI"), + ): + from run_agent import AIAgent + + return AIAgent( + model="inert-context-policy-model", provider="custom", + api_key="inert-test-key", base_url="http://127.0.0.1:1/v1", + enabled_toolsets=enabled, disabled_toolsets=disabled, + quiet_mode=True, skip_context_files=True, skip_memory=True, + ) + + +@pytest.mark.parametrize("enabled", [None, ["context_engine"], [], ["coding"]]) +@pytest.mark.parametrize("disabled", [None, [], ["context_engine"], ["unrelated"]]) +def test_context_engine_disabled_policy_at_real_initializer(enabled, disabled): + engine = _ToolEngine() + engine.update_model = MagicMock(wraps=engine.update_model) + engine.on_session_start = MagicMock() + agent = _context_policy_agent(engine, enabled=enabled, disabled=disabled) + allowed = (enabled is None or "context_engine" in enabled) and "context_engine" not in (disabled or []) + + names = {tool["function"]["name"] for tool in agent.tools} + assert ("stub_recover" in names) is allowed + assert ("stub_recover" in agent.valid_tool_names) is allowed + assert agent._context_engine_tool_names == ({"stub_recover"} if allowed else set()) + assert names == ({"policy_keep", "stub_recover"} if allowed else {"policy_keep"}) + # Disabling model tools does not disable or remove the context engine. + assert agent.context_compressor is engine + engine.update_model.assert_called_once() + engine.on_session_start.assert_called_once() + assert engine.context_length == 204_800 + messages = [{"role": "user", "content": "INERT-CONTEXT-POLICY"}] + assert engine.compress(messages) == messages + + +@pytest.mark.parametrize("enabled", [None, ["context_engine"]]) +def test_real_initialized_engine_policy_survives_repeated_refresh(monkeypatch, enabled): + """Late rebuilds remove stale schemas/routing; re-enable preserves semantics.""" + import model_tools + from tools import mcp_tool + + engine = _ToolEngine() + engine.on_session_start = MagicMock() + agent = _context_policy_agent(engine, enabled=enabled, disabled=[]) + keep = {"type": "function", "function": { + "name": "policy_keep", "description": "", "parameters": {}, + }} + monkeypatch.setattr(model_tools, "get_tool_definitions", lambda **kw: [keep]) + assert agent._context_engine_tool_names == {"stub_recover"} + + for _ in range(2): + mcp_tool.refresh_agent_mcp_tools(agent, disabled_override=["context_engine"]) + assert agent.valid_tool_names == {"policy_keep"} + assert agent._context_engine_tool_names == set() + assert [tool["function"]["name"] for tool in agent.tools] == ["policy_keep"] + added = mcp_tool.refresh_agent_mcp_tools(agent, disabled_override=[]) + assert added == {"stub_recover"} + assert agent.valid_tool_names == {"policy_keep", "stub_recover"} + assert agent._context_engine_tool_names == {"stub_recover"} + assert sum(tool["function"]["name"] == "stub_recover" for tool in agent.tools) == 1 + assert agent.context_compressor is engine + engine.on_session_start.assert_called_once() + diff --git a/tests/test_atomic_replace_symlinks.py b/tests/test_atomic_replace_symlinks.py index 42c85568892d3..f7bb46672df7c 100644 --- a/tests/test_atomic_replace_symlinks.py +++ b/tests/test_atomic_replace_symlinks.py @@ -224,23 +224,20 @@ def test_atomic_replace_broken_symlink_creates_target(tmp_path: Path) -> None: -def test_atomic_replace_copy_fallback_preserves_symlink( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -) -> None: +def test_atomic_replace_copy_fallback_preserves_symlink(tmp_path, monkeypatch): real = tmp_path / "real.yaml" link = tmp_path / "link.yaml" real.write_text("old\n", encoding="utf-8") link.symlink_to(real) tmp = _write_tmp(tmp_path, "new\n") - - def fail_replace(src: str, dst: str) -> None: - raise OSError(errno.EXDEV, os.strerror(errno.EXDEV), src, None, dst) - - monkeypatch.setattr("utils.os.replace", fail_replace) - + replace = os.replace + def initial_exdev(src, dst): + if Path(src) == tmp: + raise OSError(errno.EXDEV, "cross-device") + return replace(src, dst) + monkeypatch.setattr("utils.os.replace", initial_exdev) assert Path(atomic_replace(tmp, link)) == real - assert link.is_symlink() - assert real.read_text(encoding="utf-8") == "new\n" + assert link.is_symlink() and real.read_text(encoding="utf-8") == "new\n" assert not tmp.exists() @@ -313,37 +310,20 @@ def fast_replace_retries(monkeypatch: pytest.MonkeyPatch) -> None: @pytest.mark.parametrize("winerror", [5, 32, 33]) -def test_contended_rename_retries_then_rewrites_in_place( - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, - fast_replace_retries: None, - winerror: int, -) -> None: - """A target held for the whole call: the rename is retried the full - budget, then the in-place rewrite lands the write anyway. - - winerror 5 is the code the reported bug actually produces; 32 and 33 are - the sibling contention codes. All three must recover. - """ - import utils as utils_mod - - target = tmp_path / "gateway_state.json" +def test_contended_rename_exhaustion_preserves_existing_file(tmp_path, monkeypatch, fast_replace_retries, winerror): + import utils + target = tmp_path / "state.json" target.write_text("old", encoding="utf-8") tmp = _write_tmp(tmp_path, "new") - - attempts = [] - - def always_contended(src: str, dst: str) -> None: - attempts.append(src) - raise _sharing_error(winerror) - - monkeypatch.setattr("utils.os.replace", always_contended) + denial = _sharing_error(winerror) + replace = MagicMock(side_effect=denial) + monkeypatch.setattr("utils.os.replace", replace) monkeypatch.setattr("utils._IS_WINDOWS", True) - - assert Path(atomic_replace(tmp, target)) == target - assert len(attempts) == 1 + utils_mod._REPLACE_RETRY_ATTEMPTS - assert target.read_text(encoding="utf-8") == "new" - assert not tmp.exists() + with pytest.raises(PermissionError) as caught: + atomic_replace(tmp, target) + assert caught.value is denial + assert replace.call_count == 1 + utils._REPLACE_RETRY_ATTEMPTS + assert target.read_text(encoding="utf-8") == "old" and tmp.exists() def test_contended_rename_retry_wins_keeps_write_atomic( @@ -369,7 +349,7 @@ def forbid(*_args: object, **_kw: object) -> None: monkeypatch.setattr("utils.os.replace", contended_twice) monkeypatch.setattr("utils._IS_WINDOWS", True) - monkeypatch.setattr("utils._rewrite_in_place", forbid) + monkeypatch.setattr("utils._copy_fallback", forbid) monkeypatch.setattr("utils.shutil.copyfile", forbid) assert Path(atomic_replace(tmp, target)) == target @@ -378,60 +358,39 @@ def forbid(*_args: object, **_kw: object) -> None: assert not tmp.exists() -def test_genuine_denial_propagates_after_budget( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch, fast_replace_retries: None -) -> None: - """A real ACL denial reports the same winerror as contention, so it is - not classified up front — it exhausts the budget, fails the in-place - rewrite too, and surfaces to the caller instead of being swallowed.""" +def test_genuine_denial_propagates_after_budget(tmp_path, monkeypatch, fast_replace_retries): target = tmp_path / "denied.json" target.write_text("old", encoding="utf-8") tmp = _write_tmp(tmp_path, "new") - - monkeypatch.setattr( - "utils.os.replace", MagicMock(side_effect=_sharing_error(5)) - ) + denial = _sharing_error(5) + monkeypatch.setattr("utils.os.replace", MagicMock(side_effect=denial)) monkeypatch.setattr("utils._IS_WINDOWS", True) - - denial = PermissionError(errno.EACCES, "access is denied") - - def cannot_open(*_args: object, **_kw: object) -> None: - raise denial - - monkeypatch.setattr("utils.os.open", cannot_open) - + open_target = MagicMock(side_effect=AssertionError("must not open target for rewriting")) + monkeypatch.setattr("utils.os.open", open_target) with pytest.raises(PermissionError) as caught: atomic_replace(tmp, target) - - assert caught.value is denial - assert target.read_text(encoding="utf-8") == "old" - assert tmp.exists(), "the pending write must survive for the caller" + assert caught.value is denial and not open_target.called + assert target.read_text(encoding="utf-8") == "old" and tmp.exists() -def test_contended_retry_switching_to_exdev_uses_copy_fallback( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch, fast_replace_retries: None -) -> None: - """A retry that turns into EXDEV must stop consuming the sharing budget - and take the copy fallback — EXDEV never clears on retry.""" +def test_contended_retry_switching_to_exdev_uses_copy_fallback(tmp_path, monkeypatch, fast_replace_retries): target = tmp_path / "target.json" target.write_text("old", encoding="utf-8") tmp = _write_tmp(tmp_path, "new") - - replace = MagicMock( - side_effect=[_sharing_error(5), OSError(errno.EXDEV, "cross-device")] - ) + real_replace = os.replace + calls = [] + def replace(src, dst): + calls.append((src, dst)) + if len(calls) == 1: + raise _sharing_error(5) + if len(calls) == 2: + raise OSError(errno.EXDEV, "cross-device") + return real_replace(src, dst) monkeypatch.setattr("utils.os.replace", replace) monkeypatch.setattr("utils._IS_WINDOWS", True) - - def forbid_rewrite(*_a: object, **_k: object) -> None: - raise AssertionError("EXDEV must use the copy fallback, not a rewrite") - - monkeypatch.setattr("utils._rewrite_in_place", forbid_rewrite) - assert Path(atomic_replace(tmp, target)) == target - assert replace.call_count == 2 - assert target.read_text(encoding="utf-8") == "new" - assert not tmp.exists() + assert len(calls) == 3 + assert target.read_text(encoding="utf-8") == "new" and not tmp.exists() def test_non_contended_oserror_propagates_without_retry( @@ -479,70 +438,42 @@ def test_posix_eacces_propagates_without_retry( assert tmp.exists() -def test_in_place_rewrite_never_exposes_a_truncated_file( - tmp_path: Path, -) -> None: - """The in-place rewrite must not truncate-then-fill: a concurrent reader - can observe a 0-byte file during shutil.copyfile, which for auth.json - means an empty credential store. Shrinking writes must also not leave - trailing bytes from the previous, longer content. - """ - import utils as utils_mod - - target = tmp_path / "auth.json" - observed: list[int] = [] - - target.write_text("A" * 5000, encoding="utf-8") - tmp = _write_tmp(tmp_path, "B" * 5000) - utils_mod._rewrite_in_place(str(tmp), str(target)) - observed.append(len(target.read_text(encoding="utf-8"))) - assert target.read_text(encoding="utf-8") == "B" * 5000 - assert not tmp.exists() - - # Shrinking rewrite: ftruncate must drop the tail. - tmp = _write_tmp(tmp_path, "C" * 10) - utils_mod._rewrite_in_place(str(tmp), str(target)) - assert target.read_text(encoding="utf-8") == "C" * 10 - assert observed == [5000] +def test_busy_target_does_not_overwrite_existing_file(tmp_path, monkeypatch): + target = tmp_path / "existing.txt" + target.write_text("old", encoding="utf-8") + tmp = _write_tmp(tmp_path, "new") + monkeypatch.setattr("utils.os.replace", MagicMock(side_effect=OSError(errno.EBUSY, "busy"))) + with pytest.raises(OSError) as caught: + atomic_replace(tmp, target) + assert caught.value.errno == errno.EBUSY + assert target.read_text(encoding="utf-8") == "old" and tmp.exists() -def test_symlinked_target_survives_a_contended_rename( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch, fast_replace_retries: None -) -> None: - """The #16743 invariant must hold on the contended path too: a symlinked - config.yaml stays a symlink when the rewrite fallback runs.""" +def test_symlinked_target_survives_a_contended_rename(tmp_path, monkeypatch, fast_replace_retries): real = tmp_path / "real.yaml" link = tmp_path / "config.yaml" real.write_text("old\n", encoding="utf-8") link.symlink_to(real) tmp = _write_tmp(tmp_path, "new\n") - - monkeypatch.setattr( - "utils.os.replace", MagicMock(side_effect=_sharing_error(5)) - ) + monkeypatch.setattr("utils.os.replace", MagicMock(side_effect=_sharing_error(5))) monkeypatch.setattr("utils._IS_WINDOWS", True) - - assert Path(atomic_replace(tmp, link)) == real - assert link.is_symlink(), "symlink must survive the rewrite fallback" - assert real.read_text(encoding="utf-8") == "new\n" - assert not tmp.exists() + with pytest.raises(PermissionError): + atomic_replace(tmp, link) + assert link.is_symlink() and real.read_text(encoding="utf-8") == "old\n" and tmp.exists() # ── native Windows: real contended handles ──────────────────────────────── @pytest.mark.windows_only -def test_windows_real_held_read_handle_lands_the_write(tmp_path: Path) -> None: - """The reported bug, end to end against a real held handle.""" +def test_windows_real_held_read_handle_preserves_existing_file(tmp_path): target = tmp_path / "gateway_state.json" target.write_text('{"active_agents": 1}', encoding="utf-8") tmp = _write_tmp(tmp_path, '{"active_agents": 2}') - - with open(target, "r", encoding="utf-8"): - assert Path(atomic_replace(tmp, target)) == target - - assert json.loads(target.read_text(encoding="utf-8")) == {"active_agents": 2} - assert not tmp.exists() + with open(target, "r", encoding="utf-8"), pytest.raises(OSError): + atomic_replace(tmp, target) + assert json.loads(target.read_text(encoding="utf-8")) == {"active_agents": 1} + assert tmp.exists() @pytest.mark.windows_only @@ -566,20 +497,14 @@ def test_windows_real_held_handle_reports_access_denied(tmp_path: Path) -> None: @pytest.mark.windows_only -def test_windows_atomic_json_write_with_concurrent_reader( - tmp_path: Path, -) -> None: - """End-to-end gateway_state.json scenario through atomic_json_write: - the write lands and no .tmp file is orphaned.""" +def test_windows_atomic_json_write_with_concurrent_reader(tmp_path): target = tmp_path / "gateway_state.json" atomic_json_write(target, {"active_agents": 1}) - - with open(target, "r", encoding="utf-8"): + with open(target, "r", encoding="utf-8"), pytest.raises(OSError): atomic_json_write(target, {"active_agents": 2}) - - assert json.loads(target.read_text(encoding="utf-8")) == {"active_agents": 2} + assert json.loads(target.read_text(encoding="utf-8")) == {"active_agents": 1} leftovers = list(tmp_path.glob("*.tmp")) + list(tmp_path.glob(".*.tmp")) - assert leftovers == [], f"orphaned temp files: {leftovers}" + assert leftovers == [] @pytest.mark.windows_only @@ -597,3 +522,133 @@ def test_windows_readonly_target_still_raises(tmp_path: Path) -> None: assert target.read_text(encoding="utf-8") == "old" finally: subprocess.run(["attrib", "-R", str(target)], capture_output=True) + + +# Ordinary publication failures: tiny payloads and injected errno only. +def _publication_sentinels(tmp_path): + tmp_path = tmp_path / "publication" + tmp_path.mkdir() + target = tmp_path / "real.txt" + link = tmp_path / "linked.txt" + old, new = b"OLD-complete\n", b"NEW-complete\n" + target.write_bytes(old) + link.symlink_to(target) + incoming = tmp_path / "incoming.tmp" + incoming.write_bytes(new) + return target, link, incoming, old, new + + +def _fail_copy_after_prefix(monkeypatch): + def copyfile(src, dst, **kwargs): + with open(src, "rb") as source, open(dst, "wb") as sink: + sink.write(source.read(3)) + raise OSError(errno.ENOSPC, "injected short copy") + + def copyfileobj(source, sink, *args, **kwargs): + sink.write(source.read(3)) + raise OSError(errno.ENOSPC, "injected short copy") + + monkeypatch.setattr("utils.shutil.copyfile", copyfile) + monkeypatch.setattr("utils.shutil.copyfileobj", copyfileobj) + + +def test_exdev_short_copy_enospc_preserves_existing_bytes(tmp_path, monkeypatch): + target, link, incoming, old, new = _publication_sentinels(tmp_path) + monkeypatch.setattr("utils.os.replace", lambda *args: (_ for _ in ()).throw(OSError(errno.EXDEV, "injected EXDEV"))) + _fail_copy_after_prefix(monkeypatch) + with pytest.raises(OSError) as caught: + atomic_replace(incoming, link) + assert caught.value.errno == errno.ENOSPC + print({"pre": old, "post": target.read_bytes(), "incoming": incoming.read_bytes()}) + assert target.read_bytes() == old + assert link.is_symlink() and incoming.read_bytes() == new + assert sorted(p.name for p in target.parent.iterdir()) == ["incoming.tmp", "linked.txt", "real.txt"] + + +def test_ebusy_refuses_before_destination_write(tmp_path, monkeypatch): + target, link, incoming, old, new = _publication_sentinels(tmp_path) + calls = [] + import utils + real_copy = utils.shutil.copyfile + def observed_copy(src, dst, **kwargs): + calls.append(str(dst)) + return real_copy(src, dst, **kwargs) + monkeypatch.setattr("utils.shutil.copyfile", observed_copy) + monkeypatch.setattr("utils.os.replace", lambda *args: (_ for _ in ()).throw(OSError(errno.EBUSY, "injected busy"))) + error = None + try: + atomic_replace(incoming, link) + except OSError as exc: + error = exc + print({"pre": old, "post": target.read_bytes(), "destination_copies": calls}) + assert error is not None and error.errno == errno.EBUSY + assert calls == [] + assert target.read_bytes() == old and incoming.read_bytes() == new and link.is_symlink() + + +@pytest.mark.parametrize("winerror", [5, 32, 33]) +def test_persistent_windows_contention_preserves_existing_bytes(tmp_path, monkeypatch, winerror): + import utils + target, link, incoming, old, new = _publication_sentinels(tmp_path) + writes = [] + real_write = os.write + def short_write(fd, data): + writes.append(bytes(data)) + if len(writes) == 1: + return real_write(fd, data[:3]) + raise OSError(errno.ENOSPC, "injected short write") + rename = MagicMock(side_effect=_sharing_error(winerror)) + monkeypatch.setattr("utils._IS_WINDOWS", True) + monkeypatch.setattr("utils.os.replace", rename) + monkeypatch.setattr("utils.os.write", short_write) + monkeypatch.setattr("utils.time.sleep", lambda delay: None) + with pytest.raises(OSError) as caught: + atomic_replace(incoming, link) + print({"pre": old, "post": target.read_bytes(), "writes": writes}) + assert target.read_bytes() == old + assert getattr(caught.value, "winerror", None) == winerror + assert rename.call_count == 1 + utils._REPLACE_RETRY_ATTEMPTS + assert writes == [] and incoming.read_bytes() == new and link.is_symlink() + + +def test_exdev_restaging_publishes_from_resolved_target_parent(tmp_path, monkeypatch): + target, link, incoming, old, new = _publication_sentinels(tmp_path) + import utils + real_replace = os.replace + calls = [] + def replace(src, dst): + calls.append((Path(src), Path(dst))) + if len(calls) == 1: + raise OSError(errno.EXDEV, "injected EXDEV") + assert Path(src).parent == target.parent + assert Path(src) != incoming and Path(dst) == target + return real_replace(src, dst) + monkeypatch.setattr("utils.os.replace", replace) + assert Path(atomic_replace(incoming, link)) == target + assert len(calls) == 2 + assert target.read_bytes() == new and link.is_symlink() and not incoming.exists() + assert sorted(p.name for p in target.parent.iterdir()) == ["linked.txt", "real.txt"] + + +@pytest.mark.parametrize("failure", ["copy", "fsync", "rename"]) +def test_exdev_staging_failure_preserves_existing_bytes(tmp_path, monkeypatch, failure): + target, link, incoming, old, new = _publication_sentinels(tmp_path) + calls = [] + def replace(src, dst): + calls.append((src, dst)) + raise OSError(errno.EXDEV if len(calls) == 1 else errno.ENOSPC, "injected rename") + monkeypatch.setattr("utils.os.replace", replace) + if failure == "copy": + _fail_copy_after_prefix(monkeypatch) + elif failure == "fsync": + monkeypatch.setattr("utils.os.fsync", lambda fd: (_ for _ in ()).throw(OSError(errno.ENOSPC, "injected fsync"))) + error = None + try: + atomic_replace(incoming, link) + except OSError as exc: + error = exc + print({"failure": failure, "pre": old, "post": target.read_bytes()}) + assert error is not None and error.errno == errno.ENOSPC + assert target.read_bytes() == old + assert incoming.read_bytes() == new and link.is_symlink() + assert sorted(p.name for p in target.parent.iterdir()) == ["incoming.tmp", "linked.txt", "real.txt"] diff --git a/tests/test_npm_failure_diagnostics.py b/tests/test_npm_failure_diagnostics.py index f444df55f4584..e67a63949cc6c 100644 --- a/tests/test_npm_failure_diagnostics.py +++ b/tests/test_npm_failure_diagnostics.py @@ -19,7 +19,7 @@ def test_bounded_tail_and_known_codes_only(tmp_path): def test_collector_revalidates_instead_of_copying_strings(tmp_path): - safe = api["record"]("tui", 1, 0, ["E401"], False) + safe = api["record"]("tui", 1, 0, ["E401"], False, output_kind="recognized-code") evil = {**safe, "safe_causes": ["Authorization: Bearer secret-token"]} extra = {**safe, "raw_output": "private-token"} log = tmp_path / "transcript" @@ -28,7 +28,7 @@ def test_collector_revalidates_instead_of_copying_strings(tmp_path): def test_collector_record_count_bound(tmp_path): - safe = api["record"]("root", 1, 0, [], False) + safe = api["record"]("root", 1, 0, [], False, output_kind="empty") log = tmp_path / "transcript" log.write_text((api["PREFIX"] + json.dumps(safe) + "\n") * 50) assert len(api["collect"](log)) == 4 @@ -37,7 +37,7 @@ def test_collector_record_count_bound(tmp_path): @pytest.mark.parametrize("bad", [None, [], {}, "secret", True, -1, 256]) def test_invalid_status_rejected(bad): with pytest.raises(ValueError): - api["record"]("root", bad, 0, [], False) + api["record"]("root", bad, 0, [], False, output_kind="empty") @pytest.mark.linux_only @@ -80,7 +80,7 @@ def test_two_invocations_under_reused_parent_have_distinct_current_artifacts(tmp run = Path(made.stdout.strip()) directories.append(run) assert run.parent == parent - value = api["record"]("root", code, 2, ["ERESOLVE"], False) + value = api["record"]("root", code, 2, ["ERESOLVE"], False, output_kind="recognized-code") transcript = run / "reinstall.log" transcript.write_text(api["PREFIX"] + json.dumps(value) + "\n") subprocess.run([sys.executable, api["__file__"], "collect", "--input", str(transcript), @@ -89,3 +89,70 @@ def test_two_invocations_under_reused_parent_have_distinct_current_artifacts(tmp assert directories[0] != directories[1] assert json.loads((directories[0] / "npm-diagnostics-reinstall.json").read_text())[0]["exit_code"] == 37 assert json.loads((directories[1] / "npm-diagnostics-reinstall.json").read_text())[0]["exit_code"] == 124 + + +@pytest.mark.parametrize("text,kind,codes", [ + ("", "empty", []), + (" \n\t", "empty", []), + ("npm error code EUNRECOGNIZED\nAuthorization: sentinel-private-token", "unrecognized-code", []), + ("npm ERR! code EUNRECOGNIZED\n", "unrecognized-code", []), + ("SSL EOF occurred\nAuthorization: sentinel-private-token", "non-code", []), + ("npm error code ENOTEMPTY\n", "recognized-code", ["ENOTEMPTY"]), + ("npm ERR! code UNABLE_TO_GET_ISSUER_CERT\n", "recognized-code", ["UNABLE_TO_GET_ISSUER_CERT"]), + ("npm error code ERESOLVE\nnpm error code EUNRECOGNIZED\n", "recognized-code", ["ERESOLVE"]), +]) +def test_bounded_output_classification_does_not_expose_unknown_text(tmp_path, text, kind, codes): + log = tmp_path / "raw" + log.write_text(text) + value = api["summarize"](log, "root", 217, 2) + assert value["schema"] == "npm-install-diagnostic/v2" + assert value["output_kind"] == kind + assert value["npm_codes"] == codes + assert value["exit_code"] == 217 + assert value["timeout_status"] is False + assert "sentinel" not in json.dumps(value) + assert "EUNRECOGNIZED" not in json.dumps(value) + assert len(json.dumps(value)) < 1024 + + +@pytest.mark.parametrize("status", [1, 217, 124]) +def test_exit_status_alone_never_classifies_a_filesystem_or_tls_cause(tmp_path, status): + log = tmp_path / "raw" + log.write_text("") + value = api["summarize"](log, "root", status, 2) + assert value["output_kind"] == "empty" + assert value["npm_codes"] == [] + assert value["safe_causes"] == [] + assert value["exit_code"] == status + assert value["timeout_status"] == (status == 124) + + +def test_collector_rebuilds_and_validates_classification(tmp_path): + safe = api["record"]("root", 217, 2, [], False, output_kind="unrecognized-code") + invalid = [{**safe, "output_kind": "private-token"}, + {**safe, "output_kind": "recognized-code"}, + {**safe, "detail": "Authorization: secret-token"}, + {**safe, "schema": "npm-install-diagnostic/v1"}] + missing = {key: value for key, value in safe.items() if key != "output_kind"} + log = tmp_path / "transcript" + log.write_text("\n".join(api["PREFIX"] + json.dumps(value) for value in [*invalid, missing, safe])) + assert api["collect"](log) == [safe] + + +@pytest.mark.parametrize("kind,codes", [("empty", ["ERESOLVE"]), + ("recognized-code", []), + ("unrecognized-code", ["ERESOLVE"]), + ("non-code", ["ERESOLVE"]), + ("secret-token", []), (None, []), (True, [])]) +def test_invalid_output_classification_rejected(kind, codes): + with pytest.raises(ValueError): + api["record"]("root", 1, 0, codes, False, output_kind=kind) + + +def test_classification_refers_to_the_bounded_tail_only(tmp_path): + log = tmp_path / "raw" + log.write_text("npm error code ENOTEMPTY\n" + " " * (api["LIMIT"] + 10)) + value = api["summarize"](log, "root", 217, 2) + assert value["output_truncated"] is True + assert value["output_kind"] == "empty" + assert value["npm_codes"] == [] diff --git a/tests/test_tui_gateway_queue_on_busy.py b/tests/test_tui_gateway_queue_on_busy.py index 86edf8aa598c2..1b48bf49d7720 100644 --- a/tests/test_tui_gateway_queue_on_busy.py +++ b/tests/test_tui_gateway_queue_on_busy.py @@ -604,13 +604,13 @@ def test_busy_image_prompts_keep_b_and_c_attachments_in_submission_order(monkeyp "drain-b", "sid", "B", - {"image_paths": ["/tmp/b.png"], "queued_prompt_generation": 0}, + {"image_paths": ["/tmp/b.png"], "queued_prompt_generation": 0, "turn_transport": None}, ), ( "drain-c", "sid", "C", - {"image_paths": ["/tmp/c.png"], "queued_prompt_generation": 0}, + {"image_paths": ["/tmp/c.png"], "queued_prompt_generation": 0, "turn_transport": None}, ), ] @@ -741,15 +741,21 @@ def test_drain_does_not_clear_stop_after_its_final_generation_check(monkeypatch) class _Agent: clear_calls = 0 + interrupt_calls = 0 + def clear_interrupt(self): self.clear_calls += 1 + def interrupt(self): + self.interrupt_calls += 1 + agent = _Agent() session = _session(agent=agent, queued_prompt={"text": "B", "transport": None}) + monkeypatch.setitem(server._sessions, "sid", session) original_run = server._run_prompt_submit def stop_before_run(*args, **kwargs): - session["_queued_prompt_generation"] = 1 + assert server._interrupt_session_turn("sid", session) is False return original_run(*args, **kwargs) monkeypatch.setattr(server, "_session_uses_compute_host", lambda _session: False) @@ -757,7 +763,11 @@ def stop_before_run(*args, **kwargs): assert server._drain_queued_prompt("r1", "sid", session) is True assert agent.clear_calls == 0 + assert agent.interrupt_calls == 1 assert session["running"] is False + assert session["_turn_cancel_requested"] is True + assert session["_queued_prompt_generation"] == session["_last_stop_queue_generation"] == 1 + assert session["queued_prompt"] is None def test_drain_continues_with_later_queued_prompt_after_dispatch_failure(monkeypatch): @@ -775,6 +785,7 @@ def _run(_rid, _sid, session, text, **_kwargs): queued_prompts=[{"text": "next", "image_paths": ["/tmp/next.png"], "transport": None}], ) + monkeypatch.setitem(server._sessions, "sid", session) assert server._drain_queued_prompt("r1", "sid", session) is True assert calls == ["broken", "next"] assert session["queued_prompt"] is None diff --git a/tests/test_tui_gateway_server.py b/tests/test_tui_gateway_server.py index 77222157a18d3..fd373b408566e 100644 --- a/tests/test_tui_gateway_server.py +++ b/tests/test_tui_gateway_server.py @@ -3659,9 +3659,9 @@ def test_slow_resume_hydration_degrades_instead_of_killing_agent_init(monkeypatc with caplog.at_level("WARNING", logger="tui_gateway.server"): outcome = server._await_resume_history(session, sid, "hydration-degrade-key") assert outcome == "degraded" - assert session["resume_hydrating"] is False + assert session["resume_hydrating"] is True assert session["history"] == [] - assert event.is_set() + assert not event.is_set() statuses = [payload["status"] for name, payload in events if name == "session.resume_progress"] assert statuses == ["slow", "degraded_timeout"] @@ -13884,7 +13884,8 @@ def test_mirror_slash_side_effects_allowed_when_idle(monkeypatch): applied = {"model": False} - def _fake_apply_model(sid, session, arg): + def _fake_apply_model(sid, session, arg, *, explicit_model_intent=False): + assert explicit_model_intent is True applied["model"] = True return {"value": arg, "warning": ""} @@ -15997,6 +15998,8 @@ def test_session_activate_returns_inflight_stream_before_completion(monkeypatch) started = threading.Event() release = threading.Event() done = threading.Event() + allow_settle = threading.Event() + turn_thread = None class _Agent: model = "model-live" @@ -16025,6 +16028,7 @@ def run_conversation(self, prompt, conversation_history=None, stream_callback=No def _emit(event, sid, payload=None): if event == "message.complete": done.set() + assert allow_settle.wait(2), "test did not release terminal publication" monkeypatch.setattr(server, "_emit", _emit) @@ -16038,6 +16042,7 @@ def _emit(event, sid, payload=None): ) assert submit["result"]["status"] == "streaming" assert started.wait(2), "fake model did not stream before activation" + turn_thread = server._sessions["sid-live"]["_run_thread"] resp = server.handle_request( { @@ -16060,6 +16065,21 @@ def _emit(event, sid, payload=None): release.set() assert done.wait(2), "fake model turn did not complete" + # Bubble completion retains this admission until guarded lifecycle and + # once-model cleanup finish; it does not admit a successor by itself. + terminal = server.handle_request({ + "id": "activate-terminal", "method": "session.activate", + "params": {"session_id": "sid-live"}, + }) + assert terminal["result"]["inflight"] == { + "assistant": "partial answer", "streaming": False, + "user": "write a long answer", + } + assert terminal["result"]["turn_started_at"] == turn_started_at + assert server._sessions["sid-live"]["running"] is True + allow_settle.set() + turn_thread.join(timeout=2) + assert not turn_thread.is_alive(), "fake model turn did not settle" completed = server.handle_request( { "id": "activate-done", @@ -16075,7 +16095,9 @@ def _emit(event, sid, payload=None): ] finally: release.set() - done.wait(2) + allow_settle.set() + if turn_thread is not None: + turn_thread.join(timeout=2) server._sessions.pop("sid-live", None) diff --git a/tests/tools/test_async_delegation_retention.py b/tests/tools/test_async_delegation_retention.py new file mode 100644 index 0000000000000..95a88b59ef85a --- /dev/null +++ b/tests/tools/test_async_delegation_retention.py @@ -0,0 +1,367 @@ +"""Durable receipt pruning against a tiny in-memory ledger; no live DB.""" +import json +import sqlite3 +from contextlib import contextmanager + +import pytest +from tools import async_delegation as subject + +NOW = 200000.0 + +@pytest.fixture +def ledger(monkeypatch): + conn = sqlite3.connect(":memory:") + conn.execute("""CREATE TABLE async_delegations ( + delegation_id TEXT PRIMARY KEY, origin_session TEXT, state TEXT, + dispatched_at REAL, completed_at REAL, updated_at REAL, result_json TEXT, + delivery_state TEXT, delivery_attempts INTEGER, origin_session_id TEXT, + delivery_claim TEXT, event_json TEXT)""") + @contextmanager + def transaction(): + with conn: + yield conn + def forbid(*args, **kwargs): + raise AssertionError("live database/worker seam must not run") + monkeypatch.setattr(subject, "_transaction", transaction) + monkeypatch.setattr(subject, "_connect", forbid) + monkeypatch.setattr(subject, "_initialize_schema", forbid) + monkeypatch.setattr(subject, "_get_executor", forbid) + monkeypatch.setattr(subject, "recover_abandoned_delegations", lambda: 0) + monkeypatch.setattr(subject.time, "time", lambda: NOW) + yield conn + conn.close() + +def insert(conn, identifier, *, state="completed", delivery="pending", claim=None, age=0, compact=False): + event = {"delegation_id": identifier, "type": "async_delegation", "summary": "tiny receipt"} + result = {"summary": "tiny result"} + if compact: + event, result = {}, {} + conn.execute("INSERT INTO async_delegations VALUES (?,?,?,?,?,?,?,?,?,?,?,?)", ( + identifier, "" if compact else "inert-origin", state, NOW - age, NOW - age, + NOW - age, json.dumps(result), delivery, 0, "" if compact else "inert-session", claim, json.dumps(event))) + conn.commit() + +def snapshot(conn): + return conn.execute("SELECT * FROM async_delegations ORDER BY delegation_id").fetchall() + +@pytest.mark.parametrize("count", [51, 1001]) +def test_history_cap_preserves_fresh_pending_completions(ledger, count): + for i in range(count): + insert(ledger, f"p-{i:04}", age=i, compact=count > 51) + before = snapshot(ledger) + subject._prune_durable_records() + after = snapshot(ledger) + print({"pre_count": len(before), "post_count": len(after), "lost": [r for r in before if r not in after]}) + assert after == before + +def test_history_prune_preserves_claimed_and_live_rows(ledger): + old = subject._DURABLE_RETENTION_SECONDS + 1 + cases = [ + ("pending", "completed", "pending", None), + ("claimed-delivered", "completed", "delivered", "claim"), + ("claimed-dropped", "completed", "dropped", "claim"), + ("empty-claim", "completed", "delivered", ""), + ("running", "running", "delivered", None), + ("finalizing", "finalizing", "delivered", None), + ("unknown-disposition", "completed", "unknown", None), + ] + for identifier, state, delivery, claim in cases: + insert(ledger, identifier, state=state, delivery=delivery, claim=claim, age=old) + before = snapshot(ledger) + subject._prune_durable_records() + assert snapshot(ledger) == before + +def test_history_cap_prunes_only_unclaimed_terminal_history(ledger): + for i in range(40): + insert(ledger, f"delivered-{i:03}", delivery="delivered", age=100 - i) + for i in range(20): + insert(ledger, f"dropped-{i:03}", delivery="dropped", age=500 - i) + insert(ledger, "protected-pending") + insert(ledger, "protected-claim", delivery="delivered", claim="c") + protected = [r for r in snapshot(ledger) if r[0].startswith("protected-")] + subject._prune_durable_records() + after = snapshot(ledger) + assert len(after) == 52 + assert [r for r in after if r[0].startswith("protected-")] == protected + assert not any(r[0] == "delivered-000" for r in after) + assert all(any(r[0] == f"dropped-{i:03}" for r in after) for i in range(20)) + +def test_delivered_age_expiry_preserves_other_receipts(ledger): + old = subject._DURABLE_RETENTION_SECONDS + 1 + insert(ledger, "expired", delivery="delivered", age=old) + insert(ledger, "old-dropped", delivery="dropped", age=old) + insert(ledger, "old-pending", age=old) + insert(ledger, "old-claimed", delivery="delivered", claim="c", age=old) + insert(ledger, "old-running", state="running", delivery="delivered", age=old) + before = [r for r in snapshot(ledger) if r[0] != "expired"] + subject._prune_durable_records() + assert snapshot(ledger) == before + +def test_pending_payload_survives_prune_lookup_and_replay(ledger): + for i in range(51): + insert(ledger, f"pending-{i:03}", age=i) + before = snapshot(ledger) + subject._prune_durable_records() + for row in before: + receipt = subject.get_durable_delegation(row[0]) + assert receipt is not None + assert receipt["result"] == json.loads(row[6]) + assert receipt["delivery_state"] == "pending" + class Queue: + def __init__(self): + self.events = [] + def put(self, event): + self.events.append(event) + queue = Queue() + assert subject.restore_undelivered_completions(queue) == 51 + expected = [dict(json.loads(row[11]), restored=True) for row in before] + assert sorted(queue.events, key=lambda event: event["delegation_id"]) == expected + assert snapshot(ledger) == before + + +# SD03B controls deliberately avoid workers, providers and production state.db. +@pytest.fixture +def admission_ledger(monkeypatch): + from gateway import status as process_status + conn = sqlite3.connect(":memory:") + conn.execute("""CREATE TABLE async_delegations ( + delegation_id TEXT PRIMARY KEY, origin_session TEXT, + origin_ui_session_id TEXT, parent_session_id TEXT, state TEXT, + dispatched_at REAL, completed_at REAL, updated_at REAL, + event_json TEXT, result_json TEXT, delivery_state TEXT, + delivery_attempts INTEGER, delivered_at REAL, owner_pid INTEGER, + owner_started_at INTEGER, task_json TEXT, delivery_claim TEXT, + delivery_claimed_at REAL, origin_session_id TEXT)""") + calls = {"executor": 0, "submit": 0, "monitor": 0, "runner": 0} + captured = [] + trace = [] + conn.set_trace_callback(trace.append) + + @contextmanager + def transaction(): + with conn: + yield conn + + def forbid(*args, **kwargs): + raise AssertionError("live DB/provider boundary must not run") + + class RecordingExecutor: + def submit(self, callback): + calls["submit"] += 1 + captured.append(callback) + return object() + + executor = RecordingExecutor() + def get_executor(_limit): + calls["executor"] += 1 + return executor + + monkeypatch.setattr(subject, "_transaction", transaction) + monkeypatch.setattr(subject, "_connect", forbid) + monkeypatch.setattr(subject, "_initialize_schema", forbid) + monkeypatch.setattr(subject, "_records", {}) + monkeypatch.setattr(subject, "_new_delegation_id", lambda: "sd03b-new") + monkeypatch.setattr(subject, "_capture_routing_origin", lambda: {}) + monkeypatch.setattr(subject, "_get_executor", get_executor) + monkeypatch.setattr(subject, "_ensure_stale_monitor", lambda: calls.__setitem__("monitor", calls["monitor"] + 1)) + monkeypatch.setattr(subject, "_prune_durable_records", lambda: None) + monkeypatch.setattr(subject.time, "time", lambda: NOW) + monkeypatch.setattr(process_status, "get_process_start_time", lambda _pid: 42) + yield conn, calls, executor, captured, trace, transaction + conn.close() + + +def _sd03b_seed(conn, count, mixed=False): + rows = [] + for i in range(count): + state = ("running", "finalizing", "completed")[i % 3] if mixed else "completed" + claim = "held" if mixed and i % 2 else None + rows.append((f"p-{i:04}", "fixture", state, NOW, NOW, + "pending", claim, "{}", "{}")) + conn.executemany("""INSERT INTO async_delegations + (delegation_id, origin_session, state, dispatched_at, updated_at, + delivery_state, delivery_claim, event_json, result_json) + VALUES (?,?,?,?,?,?,?,?,?)""", rows) + conn.commit() + + +def _sd03b_dispatch(calls, batch=False, **overrides): + def runner(): + calls["runner"] += 1 + raise AssertionError("inert executor must not run a provider") + kwargs = dict(session_key="fixture", runner=runner, context=None, + toolsets=None, role="leaf", model=None, + progress_fn=lambda: (0, False)) + kwargs.update(overrides) + if batch: + return subject.dispatch_async_delegation_batch(goals=["inert"], **kwargs) + return subject.dispatch_async_delegation(goal="inert", **kwargs) + + +@pytest.mark.parametrize("batch", [False, True], ids=["single", "batch"]) +@pytest.mark.parametrize("count", [999, 1000]) +def test_sd03b_pending_admission_boundary(admission_ledger, batch, count): + conn, calls, _executor, _captured, _trace, _transaction = admission_ledger + assert subject._MAX_DURABLE_PENDING == 1000 + _sd03b_seed(conn, count) + before = snapshot(conn) + result = _sd03b_dispatch(calls, batch=batch) + if count == 999: + assert result["status"] == "dispatched" + assert conn.execute("SELECT COUNT(*) FROM async_delegations WHERE delivery_state='pending'").fetchone()[0] == 1000 + assert calls == {"executor": 1, "submit": 1, "monitor": 1, "runner": 0} + else: + assert result["status"] == "rejected" + assert result["error_code"] == "durable_backlog_full" + assert result["execution_started"] is False + assert snapshot(conn) == before + assert calls == {"executor": 0, "submit": 0, "monitor": 0, "runner": 0} + assert subject._records == {} + + +def test_sd03b_pending_active_and_claimed_rows_reserve_capacity(admission_ledger): + conn, calls, *_rest = admission_ledger + _sd03b_seed(conn, 1000, mixed=True) + before = snapshot(conn) + result = _sd03b_dispatch(calls) + assert result["status"] == "rejected" + assert snapshot(conn) == before + assert not any(calls.values()) + + +def test_sd03b_disposition_frees_one_admission_slot(admission_ledger): + conn, calls, *_rest = admission_ledger + _sd03b_seed(conn, 1000) + assert subject.mark_completion_delivered("p-0000") + delivered = conn.execute("SELECT * FROM async_delegations WHERE delegation_id='p-0000'").fetchone() + assert _sd03b_dispatch(calls)["status"] == "dispatched" + assert conn.execute("SELECT * FROM async_delegations WHERE delegation_id='p-0000'").fetchone() == delivered + assert conn.execute("SELECT COUNT(*) FROM async_delegations WHERE delivery_state='pending'").fetchone()[0] == 1000 + + +def test_sd03b_count_and_insert_use_immediate_write_transaction(admission_ledger): + conn, calls, _executor, _captured, trace, _transaction = admission_ledger + trace.clear() + assert _sd03b_dispatch(calls)["status"] == "dispatched" + normalized = [line.upper().strip() for line in trace] + begin = next(i for i, line in enumerate(normalized) if line == "BEGIN IMMEDIATE") + count = next(i for i, line in enumerate(normalized) if line.startswith("SELECT COUNT(*)")) + insert_at = next(i for i, line in enumerate(normalized) if line.startswith("INSERT INTO")) + commit = next(i for i, line in enumerate(normalized) if line == "COMMIT") + assert begin < count < insert_at < commit + + +@pytest.mark.parametrize("batch", [False, True], ids=["single", "batch"]) +def test_sd03b_storage_failure_removes_provisional_record(admission_ledger, monkeypatch, batch): + conn, calls, *_rest = admission_ledger + @contextmanager + def full_storage(): + raise sqlite3.OperationalError("inert database full") + yield conn + monkeypatch.setattr(subject, "_transaction", full_storage) + result = _sd03b_dispatch(calls, batch=batch) + assert result["status"] == "rejected" + assert result["error_code"] == "durable_storage_unavailable" + assert result["execution_started"] is False + assert subject._records == {} + assert not any(calls.values()) + + +@pytest.mark.parametrize("batch", [False, True], ids=["single", "batch"]) +def test_sd03b_commit_failure_rolls_back_before_execution(admission_ledger, monkeypatch, batch): + conn, calls, *_rest = admission_ledger + @contextmanager + def failed_commit(): + try: + yield conn + raise sqlite3.OperationalError("inert commit failure") + finally: + conn.rollback() + monkeypatch.setattr(subject, "_transaction", failed_commit) + result = _sd03b_dispatch(calls, batch=batch) + assert result["status"] == "rejected" + assert result["execution_started"] is False + assert subject._records == {} + assert snapshot(conn) == [] + assert not any(calls.values()) + + +def test_sd03b_duplicate_id_preserves_durable_receipt(admission_ledger): + conn, _calls, *_rest = admission_ledger + _sd03b_seed(conn, 1) + before = snapshot(conn) + record = {"delegation_id": "p-0000", "session_key": "other", "dispatched_at": NOW, "goal": "new"} + with pytest.raises(sqlite3.IntegrityError): + subject._persist_dispatch(record) + assert snapshot(conn) == before + + +def test_sd03b_duplicate_memory_id_preserves_existing_record(admission_ledger): + conn, calls, *_rest = admission_ledger + original = {"status": "completed", "result": {"summary": "retained"}} + subject._records["sd03b-new"] = original + result = _sd03b_dispatch(calls) + assert result["status"] == "rejected" + assert subject._records["sd03b-new"] is original + assert snapshot(conn) == [] + assert not any(calls.values()) + + +def test_sd03b_pruning_failure_keeps_committed_admission(admission_ledger, monkeypatch): + conn, calls, *_rest = admission_ledger + def failed_prune(): + raise OSError(28, "inert ENOSPC") + monkeypatch.setattr(subject, "_prune_durable_records", failed_prune) + assert _sd03b_dispatch(calls)["status"] == "dispatched" + assert conn.execute("SELECT COUNT(*) FROM async_delegations").fetchone()[0] == 1 + assert calls["submit"] == 1 and calls["runner"] == 0 + + +@pytest.mark.parametrize("batch", [False, True], ids=["single", "batch"]) +def test_sd03b_executor_lookup_failure_releases_unstarted_reservation(admission_ledger, monkeypatch, batch): + conn, calls, *_rest = admission_ledger + def unavailable(_limit): + raise RuntimeError("inert executor unavailable before submission") + monkeypatch.setattr(subject, "_get_executor", unavailable) + result = _sd03b_dispatch(calls, batch=batch) + assert result["status"] == "rejected" + assert result["execution_started"] is False + assert result["durable_reservation_released"] is True + assert subject._records == {} and snapshot(conn) == [] + assert calls["submit"] == calls["runner"] == 0 + + +@pytest.mark.parametrize("batch", [False, True], ids=["single", "batch"]) +def test_sd03b_ambiguous_submit_retains_record_without_replay(admission_ledger, monkeypatch, batch): + conn, calls, executor, captured, *_rest = admission_ledger + def enqueue_then_raise(callback): + calls["submit"] += 1 + captured.append(callback) + raise RuntimeError("inert submission outcome unknown") + monkeypatch.setattr(executor, "submit", enqueue_then_raise) + result = _sd03b_dispatch(calls, batch=batch) + assert result["status"] == "dispatch_uncertain" + assert result["error_code"] == "scheduling_uncertain" + assert result["execution_started"] is None + assert result["delegation_id"] == "sd03b-new" + assert "sd03b-new" in subject._records + assert conn.execute("SELECT COUNT(*) FROM async_delegations").fetchone()[0] == 1 + assert len(captured) == 1 and calls["runner"] == 0 + + +@pytest.mark.parametrize("protected", ["claim", "event", "completion", "owner", "timestamp"]) +def test_sd03b_release_never_deletes_nonmatching_or_protected_rows(admission_ledger, protected): + conn, _calls, *_rest = admission_ledger + record = {"delegation_id": "owned", "session_key": "fixture", "dispatched_at": NOW} + reservation = subject._persist_dispatch(record) + changes = { + "claim": ("delivery_claim", "held"), "event": ("event_json", "{}"), + "completion": ("completed_at", NOW), "owner": ("owner_started_at", 43), + "timestamp": ("dispatched_at", NOW + 1), + } + column, value = changes[protected] + conn.execute(f"UPDATE async_delegations SET {column}=? WHERE delegation_id='owned'", (value,)) + conn.commit() + before = snapshot(conn) + assert subject._delete_durable_delegation("owned", reservation=reservation) is False + assert snapshot(conn) == before diff --git a/tests/tools/test_cronjob_run_background.py b/tests/tools/test_cronjob_run_background.py index a35d4c4a8fbca..d23bd8cdd7ca1 100644 --- a/tests/tools/test_cronjob_run_background.py +++ b/tests/tools/test_cronjob_run_background.py @@ -14,6 +14,7 @@ """ import json import threading +from contextlib import contextmanager from unittest.mock import patch from tools.cronjob_tools import ( @@ -55,6 +56,62 @@ def _cm(): return _cm() +def _claimed_with_receipt(job_id): + return {**_job(job_id), "fire_claim": {"by": "bg-owner"}}, object() + + +@contextmanager +def _background_runtime(): + """Real admission, worker and queue with an inert ledger and joined futures.""" + import sqlite3 + import time + + from tools import async_delegation as owner + + conn = sqlite3.connect(":memory:", check_same_thread=False) + owner._initialize_schema(conn) + transaction_lock = threading.RLock() + + @contextmanager + def transaction(): + with transaction_lock, conn: + yield conn + + executor = owner._DaemonThreadPoolExecutor( + max_workers=1, thread_name_prefix="cron-test-delegate" + ) + original_submit = executor.submit + futures = [] + + def submit(*args, **kwargs): + future = original_submit(*args, **kwargs) + futures.append(future) + return future + + with patch.object(owner, "_transaction", transaction), \ + patch.object(owner, "_connect", side_effect=AssertionError("live ledger forbidden")), \ + patch.object(owner, "_records", {}), \ + patch.object(owner, "_get_executor", return_value=executor), \ + patch.object(executor, "submit", side_effect=submit), \ + patch("cron.executions.recover_interrupted_executions", return_value=0), \ + patch("tools.cronjob_tools._latest_job_output_excerpt", return_value=None), \ + patch("tools.cronjob_tools._notify_provider_jobs_changed_safe"), \ + patch("tools.delegate_tool._get_max_async_children", return_value=1): + joined = False + try: + yield + finally: + deadline = time.monotonic() + 5.0 + try: + for future in futures: + future.result(timeout=max(0.0, deadline - time.monotonic())) + joined = True + finally: + executor.shutdown(wait=False) + if joined: + conn.close() + + class TestBackgroundDispatch: def test_dispatches_and_returns_handle_immediately(self): """With a routable session, run claims sync then dispatches async.""" @@ -67,23 +124,23 @@ def slow_run_one_job(job, **kw): return True with _bound_session_key(): - with patch("tools.cronjob_tools.claim_job_for_fire", side_effect=lambda jid, **kw: {**_job(jid), "fire_claim": {"by": "bg-owner"}}) as m_claim, \ + with patch("cron.jobs.claim_job_for_fire_with_unstarted_receipt", side_effect=_claimed_with_receipt) as m_claim, \ patch("cron.scheduler.run_one_job", side_effect=slow_run_one_job), \ patch("tools.cronjob_tools.get_job", - return_value={"last_status": "ok", "last_error": None}): - res = _try_dispatch_background_run(_job('job-bg-01')) - - try: - # Returned BEFORE the job finished — that's the whole point. - assert res is not None - assert res["claimed"] is True - assert res["dispatched"] is True - assert res["delegation_id"] - m_claim.assert_called_once_with("job-bg-01", return_job=True) - # The job actually starts on the daemon executor. - assert run_started.wait(timeout=5.0), "job never started in background" - finally: - run_release.set() + return_value={"last_status": "ok", "last_error": None}), \ + _background_runtime(): + try: + res = _try_dispatch_background_run(_job('job-bg-01')) + # Returned while the actual worker is still gated. + assert res is not None + assert res["claimed"] is True + assert res["dispatched"] is True + assert res["delegation_id"] + m_claim.assert_called_once_with("job-bg-01") + assert run_started.wait(timeout=5.0), "job never started in background" + assert not run_release.is_set() + finally: + run_release.set() def test_completion_event_reaches_shared_queue(self): """The finished run pushes a type='async_delegation' event carrying @@ -95,11 +152,12 @@ def test_completion_event_reaches_shared_queue(self): # The runner executes on a daemon thread — the patches must stay # active until the completion event lands, so poll INSIDE the blocks. with _bound_session_key("agent:main:telegram:dm:777"): - with patch("tools.cronjob_tools.claim_job_for_fire", side_effect=lambda jid, **kw: {**_job(jid), "fire_claim": {"by": "bg-owner"}}), \ + with patch("cron.jobs.claim_job_for_fire_with_unstarted_receipt", side_effect=_claimed_with_receipt), \ patch("cron.scheduler.run_one_job", return_value=True), \ patch("tools.cronjob_tools.get_job", return_value={"last_status": "ok", "last_error": None, - "next_run_at": "2026-08-07T09:00:00"}): + "next_run_at": "2026-08-07T09:00:00"}), \ + _background_runtime(): res = _try_dispatch_background_run(_job('job-bg-02')) assert res["dispatched"] is True @@ -128,11 +186,12 @@ def test_failed_run_reports_error_status_in_event(self): from tools.process_registry import process_registry with _bound_session_key("agent:main:telegram:dm:778"): - with patch("tools.cronjob_tools.claim_job_for_fire", side_effect=lambda jid, **kw: {**_job(jid), "fire_claim": {"by": "bg-owner"}}), \ + with patch("cron.jobs.claim_job_for_fire_with_unstarted_receipt", side_effect=_claimed_with_receipt), \ patch("cron.scheduler.run_one_job", return_value=True), \ patch("tools.cronjob_tools.get_job", return_value={"last_status": "error", - "last_error": "provider exploded"}): + "last_error": "provider exploded"}), \ + _background_runtime(): res = _try_dispatch_background_run(_job('job-bg-03')) assert res["dispatched"] is True @@ -156,13 +215,14 @@ def test_claim_lost_reports_immediately_without_dispatch(self): """Paused/already-firing jobs report in the tool response, not as a delayed completion event.""" with _bound_session_key(): - with patch("tools.cronjob_tools.claim_job_for_fire", return_value=False), \ + with patch("cron.jobs.claim_job_for_fire_with_unstarted_receipt", return_value=None) as m_claim, \ patch("tools.cronjob_tools.get_job", return_value={**_JOB, "enabled": False}), \ patch("tools.async_delegation.dispatch_async_delegation") as m_disp: res = _try_dispatch_background_run(_job('job-bg-04')) assert res["claimed"] is False assert "paused/disabled" in res["error"] + m_claim.assert_called_once_with("job-bg-04") m_disp.assert_not_called() @@ -180,19 +240,19 @@ def test_async_delivery_unsupported_falls_back_to_sync(self): res = _try_dispatch_background_run(_job('job-bg-06')) assert res is None - def test_pool_at_capacity_runs_inline(self): - """A rejected dispatch must not strand the already-taken claim.""" - with _bound_session_key(): - with patch("tools.cronjob_tools.claim_job_for_fire", side_effect=lambda jid, **kw: {**_job(jid), "fire_claim": {"by": "bg-owner"}}), \ - patch("tools.async_delegation.dispatch_async_delegation", - return_value={"status": "rejected", "error": "capacity"}), \ - patch("cron.scheduler.run_one_job", return_value=True) as m_run, \ - patch("tools.cronjob_tools.get_job", - return_value={"last_status": "ok", "last_error": None}): - res = _try_dispatch_background_run(_job('job-bg-07')) + def test_pool_at_capacity_runs_inline(self, sd03b_cron_state): + """Typed pool refusal keeps exactly one inline run with its claim owner.""" + state = sd03b_cron_state + state.dispatch.return_value = { + "status": "rejected", "error": "capacity", + "error_code": "pool_capacity", "execution_started": False, + } + res = _try_dispatch_background_run(state.job) assert res["dispatched"] is False assert res["success"] is True - m_run.assert_called_once() # ran inline on this thread + state.run.assert_called_once_with(state.claimed_job, extra_prompt=None) + state.release.assert_not_called() + state.mark.assert_not_called() class TestInFlightDedupe: @@ -245,7 +305,7 @@ def test_background_dispatch_reports_running_job_immediately(self): assert sched.try_register_running_job("job-bg-10") try: with _bound_session_key(): - with patch("tools.cronjob_tools.claim_job_for_fire") as m_claim, \ + with patch("cron.jobs.claim_job_for_fire_with_unstarted_receipt") as m_claim, \ patch("tools.async_delegation.dispatch_async_delegation") as m_disp: res = _try_dispatch_background_run(_job('job-bg-10')) assert res["claimed"] is False @@ -278,11 +338,12 @@ def test_run_action_returns_background_note(self): """cronjob(action='run') surfaces the handle + do-not-wait note.""" with _bound_session_key(): with patch("tools.cronjob_tools.resolve_job_ref", return_value=_job('job-bg-12')), \ - patch("tools.cronjob_tools.claim_job_for_fire", side_effect=lambda jid, **kw: {**_job(jid), "fire_claim": {"by": "bg-owner"}}), \ + patch("cron.jobs.claim_job_for_fire_with_unstarted_receipt", side_effect=_claimed_with_receipt), \ patch("cron.scheduler.run_one_job", return_value=True), \ patch("tools.cronjob_tools.get_job", return_value={"id": "job-bg-12", "name": "bg run", - "last_status": "ok", "last_error": None}): + "last_status": "ok", "last_error": None}), \ + _background_runtime(): out = json.loads(cronjob(action="run", job_id="job-bg-12")) assert out["success"] is True @@ -306,3 +367,381 @@ def test_run_action_sync_path_unchanged_without_session(self): assert out["job"]["execution_success"] is True m_claim.assert_called_once_with("job-bg-13", return_job=True) m_run.assert_called_once() + + +# SD03B fixtures replace every execution/claim/routing/config dependency with +# in-memory recording objects. No runner, provider, DB or service is started. +import sys +from types import ModuleType, SimpleNamespace +from unittest.mock import Mock + +import pytest + + +@pytest.fixture +def sd03b_cron_state(monkeypatch): + import tools.cronjob_tools as caller + import tools.async_delegation as async_owner + import tools.delegate_tool as delegate_owner + import tools.approval as approval + import cron.jobs as job_owner + + scheduler = ModuleType("cron.scheduler") + scheduler.get_running_job_ids = Mock(return_value=set()) + scheduler.run_one_job = Mock(side_effect=AssertionError("real cron runner forbidden")) + scheduler.try_register_running_job = Mock(return_value=True) + scheduler.release_running_job = Mock() + monkeypatch.setitem(sys.modules, "cron.scheduler", scheduler) + executions = ModuleType("cron.executions") + executions.recover_interrupted_executions = Mock(return_value=0) + monkeypatch.setitem(sys.modules, "cron.executions", executions) + session_context = ModuleType("gateway.session_context") + session_context.async_delivery_supported = Mock(return_value=True) + session_context.get_session_env = Mock(return_value="") + monkeypatch.setitem(sys.modules, "gateway.session_context", session_context) + monkeypatch.setattr(approval, "get_current_session_key", lambda **_kw: "sd03b-session") + monkeypatch.setattr(async_owner, "_current_origin_session_id", lambda: "sd03b-parent") + monkeypatch.setattr(delegate_owner, "_get_max_async_children", lambda: 1) + + job = _job("sd03b-job") + claimed_job = {**job, "fire_claim": {"by": "sd03b-fire-owner", "at": "2026-10-04T00:00:00Z"}} + receipt = object() # An ephemeral exact object; never serialize or persist it. + old_claim = Mock(return_value=claimed_job) + receipt_claim = Mock(return_value=(claimed_job, receipt)) + release = Mock(return_value={"status": "released"}) + # Both surfaces are present so baseline and candidate enter the same + # inert path, while candidate is required to release the exact receipt. + for owner in (caller, job_owner): + monkeypatch.setattr(owner, "claim_job_for_fire", old_claim) + monkeypatch.setattr(owner, "claim_job_for_fire_with_unstarted_receipt", receipt_claim, raising=False) + monkeypatch.setattr(owner, "release_unstarted_fire_claim", release, raising=False) + dispatch = Mock() + monkeypatch.setattr(async_owner, "dispatch_async_delegation", dispatch) + run = Mock(return_value={"claimed": True, "success": True, "error": None}) + mark = Mock() + monkeypatch.setattr(caller, "_run_claimed_job", run) + monkeypatch.setattr(caller, "mark_job_run", mark) + monkeypatch.setattr(caller, "get_job", Mock(return_value=claimed_job)) + monkeypatch.setattr(caller, "_latest_job_output_excerpt", Mock(return_value=None)) + monkeypatch.setattr(caller, "_notify_provider_jobs_changed_safe", Mock()) + return SimpleNamespace( + caller=caller, job=job, claimed_job=claimed_job, receipt=receipt, + old_claim=old_claim, receipt_claim=receipt_claim, release=release, + dispatch=dispatch, run=run, mark=mark, scheduler=scheduler, + reclaim=executions.recover_interrupted_executions, + ) + + +def _sd03b_assert_exact_receipt_released(state): + state.release.assert_called_once() + args, kwargs = state.release.call_args + assert any(value is state.receipt for value in (*args, *kwargs.values())) + + +@pytest.mark.parametrize("reservation_released", [False, True]) +@pytest.mark.parametrize("error_code", ["durable_backlog_full", "durable_storage_unavailable"]) +def test_background_storage_refusal_releases_unstarted_claim_without_inline( + sd03b_cron_state, error_code, reservation_released +): + state = sd03b_cron_state + state.dispatch.return_value = { + "status": "rejected", "error_code": error_code, + "execution_started": False, "error": "inert storage refusal", + "delegation_id": "sd03b-refused-reservation", + "durable_reservation_released": reservation_released, + } + result = state.caller._try_dispatch_background_run(state.job) + state.run.assert_not_called() + state.scheduler.run_one_job.assert_not_called() + state.mark.assert_not_called() + _sd03b_assert_exact_receipt_released(state) + assert result["status"] == "rejected" + assert result["error_code"] == error_code + assert result["execution_started"] is False + assert result["delegation_id"] == "sd03b-refused-reservation" + assert result["durable_reservation_released"] is reservation_released + assert result["dispatched"] is False + assert result["success"] is False + assert result["claim_release_status"] == "released" + assert result["claim_release_confirmed"] is True + + +@pytest.mark.parametrize("release_status", ["conflict", "missing", "write_uncertain"]) +def test_background_storage_refusal_reports_unconfirmed_claim_release( + sd03b_cron_state, release_status +): + state = sd03b_cron_state + state.dispatch.return_value = { + "status": "rejected", "error_code": "durable_storage_unavailable", + "execution_started": False, "error": "inert storage refusal", + } + state.release.return_value = {"status": release_status} + result = state.caller._try_dispatch_background_run(state.job) + state.run.assert_not_called() + state.scheduler.run_one_job.assert_not_called() + state.mark.assert_not_called() + _sd03b_assert_exact_receipt_released(state) + assert result["status"] == "rejected" + assert result["error_code"] == "durable_storage_unavailable" + assert result["execution_started"] is False + assert result["success"] is False + assert result["claim_release_status"] == release_status + assert result["claim_release_confirmed"] is False + + +def test_background_dispatch_uncertain_keeps_claim_without_inline(sd03b_cron_state): + state = sd03b_cron_state + state.dispatch.return_value = { + "status": "dispatch_uncertain", "error_code": "scheduling_uncertain", + "execution_started": None, "delegation_id": "sd03b-uncertain", + "error": "inert uncertain submit", + } + result = state.caller._try_dispatch_background_run(state.job) + state.run.assert_not_called() + state.scheduler.run_one_job.assert_not_called() + state.release.assert_not_called() + state.mark.assert_not_called() + assert result["status"] == "dispatch_uncertain" + assert result["error_code"] == "scheduling_uncertain" + assert result["execution_started"] is None + assert result["delegation_id"] == "sd03b-uncertain" + + +def test_background_pool_capacity_runs_owner_bearing_claim_once(sd03b_cron_state): + state = sd03b_cron_state + state.dispatch.return_value = { + "status": "rejected", "error_code": "pool_capacity", + "execution_started": False, "error": "inert pool capacity", + } + result = state.caller._try_dispatch_background_run(state.job, extra_prompt="inert context") + state.run.assert_called_once_with(state.claimed_job, extra_prompt="inert context") + state.release.assert_not_called() + state.mark.assert_not_called() + state.scheduler.run_one_job.assert_not_called() + assert result["dispatched"] is False + assert result["success"] is True + + +# SD03B public response fixtures: the real wrapper is exercised with every +# execution, claim, store, scanner and notification dependency made inert. +@pytest.fixture +def sd03b_public_state(monkeypatch): + import tools.cronjob_tools as caller + + job = {"id": "sd03b-public-job", "name": "inert public job"} + refreshed_job = {"id": "sd03b-public-job", "name": "inert refreshed job"} + resolve = Mock(return_value=job) + read = Mock(return_value=refreshed_job) + # This small fixed view is independent of the production formatter/schema. + format_job = Mock(side_effect=lambda _job: {"id": "formatted-job", "view": "inert"}) + scanner = Mock(return_value=None) + background = Mock() + sync = Mock(return_value={"claimed": True, "success": True, "error": None}) + notify = Mock() + for name, value in ( + ("resolve_job_ref", resolve), ("get_job", read), + ("_format_job", format_job), ("_scan_cron_prompt", scanner), + ("_try_dispatch_background_run", background), + ("_execute_job_now", sync), ("_notify_provider_jobs_changed_safe", notify), + ): + monkeypatch.setattr(caller, name, value) + guards = {} + for name in ( + "claim_job_for_fire", "claim_job_for_fire_with_unstarted_receipt", + "release_unstarted_fire_claim", "mark_job_run", "pause_job", + "resume_job", "remove_job", "update_job", "list_jobs", + "parse_schedule", "_run_claimed_job", "_latest_job_output_excerpt", + "_origin_from_env", "_gateway_liveness_notice", + "_validate_cron_script_path", "_validate_cron_base_url", + "_validate_bot_chat_deliver", "_resolve_cron_context_deliver", + ): + guard = Mock(side_effect=AssertionError("public wrapper crossed inert guard: " + name)) + monkeypatch.setattr(caller, name, guard, raising=False) + guards[name] = guard + return SimpleNamespace( + caller=caller, job=job, refreshed_job=refreshed_job, + resolve=resolve, read=read, format_job=format_job, scanner=scanner, + background=background, sync=sync, notify=notify, guards=guards, + ) + + +def _sd03b_public_call(state, action="run"): + return json.loads(state.caller.cronjob( + action=action, job_id="inert requested name", + session_id="sd03b-public-session", prompt="inert per-run context", + )) + + +def _sd03b_public_assert_inert_route(state, *, read_after_run, notifications): + state.resolve.assert_called_once_with("inert requested name") + state.scanner.assert_called_once_with("inert per-run context") + state.background.assert_called_once_with( + state.job, session_id="sd03b-public-session", extra_prompt="inert per-run context", + ) + if read_after_run: + state.read.assert_called_once_with("sd03b-public-job") + state.format_job.assert_called_once_with(state.refreshed_job) + else: + state.read.assert_not_called() + state.format_job.assert_called_once_with(state.job) + assert state.notify.call_count == notifications + if notifications: + state.notify.assert_called_once_with() + for guard in state.guards.values(): + guard.assert_not_called() + + +@pytest.mark.parametrize("error_code", ["durable_backlog_full", "durable_storage_unavailable"]) +@pytest.mark.parametrize("reservation_released", [False, True]) +@pytest.mark.parametrize("release_status,release_confirmed", [ + ("released", True), ("write_uncertain", False), +]) +def test_public_storage_refusal_preserves_dispatch_evidence( + sd03b_public_state, error_code, reservation_released, release_status, release_confirmed +): + state = sd03b_public_state + state.background.return_value = { + "claimed": True, "dispatched": False, "success": False, + "status": "rejected", "error_code": error_code, + "execution_started": False, "error": "inert storage refusal", + "delegation_id": "sd03b-public-refused-reservation", + "durable_reservation_released": reservation_released, + "claim_release_status": release_status, + "claim_release_confirmed": release_confirmed, + } + result = _sd03b_public_call(state) + assert result == { + "success": False, + "job": {"id": "formatted-job", "view": "inert", "executed": False}, + "status": "rejected", "error_code": error_code, + "execution_started": False, "error": "inert storage refusal", + "delegation_id": "sd03b-public-refused-reservation", + "durable_reservation_released": reservation_released, + "claim_release_status": release_status, + "claim_release_confirmed": release_confirmed, + } + state.sync.assert_not_called() + _sd03b_public_assert_inert_route(state, read_after_run=False, notifications=0) + + +def test_public_storage_refusal_does_not_invent_optional_fields(sd03b_public_state): + state = sd03b_public_state + state.background.return_value = { + "claimed": True, "dispatched": False, "success": False, + "status": "rejected", "error_code": "executor_unavailable", + "execution_started": False, "error": "inert unavailable executor", + } + result = _sd03b_public_call(state) + assert result == { + "success": False, + "job": {"id": "formatted-job", "view": "inert", "executed": False}, + "status": "rejected", "error_code": "executor_unavailable", + "execution_started": False, "error": "inert unavailable executor", + } + for key in ( + "delegation_id", "durable_reservation_released", + "claim_release_status", "claim_release_confirmed", + ): + assert key not in result + state.sync.assert_not_called() + _sd03b_public_assert_inert_route(state, read_after_run=False, notifications=0) + + +@pytest.mark.parametrize("action", ["run", "run_now", "trigger"]) +def test_public_dispatch_uncertain_preserves_handle_and_no_replay_note(sd03b_public_state, action): + state = sd03b_public_state + state.background.return_value = { + "claimed": True, "dispatched": False, "success": False, + "status": "dispatch_uncertain", "error_code": "scheduling_uncertain", + "execution_started": None, "error": "inert uncertain submit", + "delegation_id": "sd03b-public-uncertain", + "durable_reservation_released": False, + } + result = _sd03b_public_call(state, action) + assert result == { + "success": False, + "job": {"id": "formatted-job", "view": "inert", "executed": None}, + "status": "dispatch_uncertain", "error_code": "scheduling_uncertain", + "execution_started": None, "error": "inert uncertain submit", + "delegation_id": "sd03b-public-uncertain", + "durable_reservation_released": False, + "note": ( + "Work may already be running. Keep the delegation handle; do not retry " + "or run inline while submission remains uncertain." + ), + } + state.sync.assert_not_called() + _sd03b_public_assert_inert_route(state, read_after_run=False, notifications=0) + + +def test_public_confirmed_dispatch_keeps_existing_payload(sd03b_public_state): + state = sd03b_public_state + state.background.return_value = { + "claimed": True, "dispatched": True, "delegation_id": "sd03b-public-confirmed", + } + result = _sd03b_public_call(state) + assert result == { + "success": True, + "job": { + "id": "formatted-job", "view": "inert", "executed": True, + "execution_mode": "background", "delegation_id": "sd03b-public-confirmed", + }, + "note": ( + "The job is running in the background. You and the user can keep working; " + "its outcome re-enters the conversation as a new message when it finishes. " + "Do not wait or poll — just continue." + ), + } + state.sync.assert_not_called() + _sd03b_public_assert_inert_route(state, read_after_run=True, notifications=1) + + +def test_public_typed_pool_fallback_keeps_existing_terminal_payload(sd03b_public_state): + state = sd03b_public_state + # The core helper has already executed the permitted pool-capacity fallback. + state.background.return_value = { + "claimed": True, "dispatched": False, "success": True, "error": None, + } + result = _sd03b_public_call(state) + assert result == { + "success": True, + "job": { + "id": "formatted-job", "view": "inert", "executed": True, + "execution_success": True, + }, + } + state.sync.assert_not_called() + _sd03b_public_assert_inert_route(state, read_after_run=True, notifications=1) + + +def test_public_background_unsupported_executes_inert_sync_once(sd03b_public_state): + state = sd03b_public_state + state.background.return_value = None + result = _sd03b_public_call(state) + assert result == { + "success": True, + "job": { + "id": "formatted-job", "view": "inert", "executed": True, + "execution_success": True, + }, + } + state.sync.assert_called_once_with(state.job, extra_prompt="inert per-run context") + _sd03b_public_assert_inert_route(state, read_after_run=True, notifications=1) + + +def test_public_claim_lost_keeps_existing_skipped_payload(sd03b_public_state): + state = sd03b_public_state + state.background.return_value = { + "claimed": False, "dispatched": False, "success": False, + "error": "inert claim lost", + } + result = _sd03b_public_call(state) + assert result == { + "success": True, + "job": { + "id": "formatted-job", "view": "inert", "executed": False, + "execution_success": False, "execution_skipped": "inert claim lost", + }, + } + state.sync.assert_not_called() + _sd03b_public_assert_inert_route(state, read_after_run=True, notifications=0) diff --git a/tests/tools/test_delegate_storage_backpressure.py b/tests/tools/test_delegate_storage_backpressure.py new file mode 100644 index 0000000000000..c6794c58b7116 --- /dev/null +++ b/tests/tools/test_delegate_storage_backpressure.py @@ -0,0 +1,183 @@ +"""SD03B inert delegate caller gates for durable dispatch backpressure. + +Only supplied recording children and callbacks are used. Dispatch never invokes +its runner; the typed pool case alone exercises one inert synchronous callback. +""" +import json +import sys +from types import ModuleType, SimpleNamespace +from unittest.mock import Mock + +import pytest + + +class _Sd03bChild: + def __init__(self, identifier, pending_steer): + self._subagent_id = identifier + self._delegate_role = "leaf" + self.session_id = identifier + "-session" + self.close = Mock() + self._drain_pending_steer = Mock(return_value=pending_steer) + self.interrupt = Mock() + self.tool_progress_callback = None + + +@pytest.fixture +def sd03b_delegate_state(monkeypatch): + import tools.delegate_tool as caller + import tools.async_delegation as async_owner + import tools.approval as approval + + live = ModuleType("tools.delegation_live_log") + writer = SimpleNamespace(path="/inert/sd03b/task-0.log", finalize=Mock()) + live.create_live_transcripts = Mock(return_value=("sd03b-delegation", [writer], [writer.path])) + live.update_manifest_statuses = Mock() + live.wrap_progress_callback = lambda callback, _writer: callback + monkeypatch.setitem(sys.modules, "tools.delegation_live_log", live) + schemas = ModuleType("tools.delegation_output_schema") + schemas.coerce_output_schema = lambda _raw: (None, None) + monkeypatch.setitem(sys.modules, "tools.delegation_output_schema", schemas) + session_context = ModuleType("gateway.session_context") + session_context.async_delivery_supported = Mock(return_value=True) + session_context.get_session_env = Mock(return_value="") + monkeypatch.setitem(sys.modules, "gateway.session_context", session_context) + monkeypatch.setattr(approval, "get_current_session_key", lambda **_kw: "sd03b-session") + monkeypatch.setattr(async_owner, "_current_origin_session_id", lambda: "sd03b-parent") + monkeypatch.setattr(caller, "_capture_gateway_steer_authority", lambda _sid: (None, None)) + monkeypatch.setattr(caller, "is_spawn_paused", lambda: False) + monkeypatch.setattr(caller, "_get_max_spawn_depth", lambda: 3) + monkeypatch.setattr(caller, "_get_max_concurrent_children", lambda: 1) + monkeypatch.setattr(caller, "_get_max_async_children", lambda: 1) + monkeypatch.setattr(caller, "_load_config", lambda: {"max_iterations": 1}) + credentials = { + "model": "sd03b-inert", "provider": None, "base_url": None, + "api_key": None, "api_mode": None, "command": None, "args": None, + } + monkeypatch.setattr(caller, "_resolve_delegation_credentials", lambda *_a, **_kw: credentials) + monkeypatch.setattr(caller, "_finalize_child_results", Mock()) + monkeypatch.setattr(caller, "_emit_parent_console", Mock()) + + missed_steer = "Keep the exact accepted steer.\nSecond line." + child = _Sd03bChild("sd03b-child", missed_steer) + unrelated = _Sd03bChild("sd03b-already-running", "unrelated accepted steer") + parent = SimpleNamespace( + _delegate_depth=0, session_id="sd03b-parent", _interrupt_requested=False, + _active_children=[unrelated], _active_children_lock=None, + ) + unrelated_record = { + "subagent_id": unrelated._subagent_id, "agent": unrelated, + "accepting_steer": True, "status": "running", "goal": "already running", + } + registry = {unrelated._subagent_id: unrelated_record} + monkeypatch.setattr(caller, "_active_subagents", registry) + monkeypatch.setattr(caller, "_recent_subagents", {}) + + def build_child(**_kwargs): + parent._active_children.append(child) + caller._register_subagent({ + "subagent_id": child._subagent_id, "agent": child, + "accepting_steer": True, "status": "running", "goal": "inert child", + "owner_agent_session_id": parent.session_id, + }) + return child + + build = Mock(side_effect=build_child) + monkeypatch.setattr(caller, "_build_child_preserving_parent_tools", build) + close_steering = Mock(wraps=caller._close_subagent_steering) + unregister = Mock(wraps=caller._unregister_subagent) + monkeypatch.setattr(caller, "_close_subagent_steering", close_steering) + monkeypatch.setattr(caller, "_unregister_subagent", unregister) + run = Mock(return_value={ + "task_index": 0, "status": "completed", "summary": "inert once", + "api_calls": 0, "duration_seconds": 0, + }) + monkeypatch.setattr(caller, "_run_single_child", run) + dispatch = Mock() + monkeypatch.setattr(async_owner, "dispatch_async_delegation_batch", dispatch) + return SimpleNamespace( + caller=caller, parent=parent, child=child, unrelated=unrelated, + unrelated_record=unrelated_record, registry=registry, + missed_steer=missed_steer, build=build, run=run, dispatch=dispatch, + close_steering=close_steering, unregister=unregister, writer=writer, live=live, + ) + + +@pytest.mark.parametrize("reservation_released", [False, True]) +@pytest.mark.parametrize("error_code", ["durable_backlog_full", "durable_storage_unavailable"]) +def test_delegate_storage_refusal_closes_only_unstarted_child_and_preserves_steer( + sd03b_delegate_state, error_code, reservation_released +): + state = sd03b_delegate_state + state.dispatch.return_value = { + "status": "rejected", "error_code": error_code, + "execution_started": False, "error": "inert storage refusal", + "delegation_id": "sd03b-refused-reservation", + "durable_reservation_released": reservation_released, + } + result = json.loads(state.caller.delegate_task( + goal="inert child", background=True, parent_agent=state.parent + )) + state.run.assert_not_called() + state.close_steering.assert_called_once_with(state.child._subagent_id, state.child) + state.child._drain_pending_steer.assert_called_once_with() + state.child.close.assert_called_once_with() + state.unrelated.close.assert_not_called() + state.unrelated._drain_pending_steer.assert_not_called() + assert state.registry[state.unrelated._subagent_id] is state.unrelated_record + assert state.unrelated_record["accepting_steer"] is True + assert state.unrelated in state.parent._active_children + assert result["status"] == "rejected" + assert result["error_code"] == error_code + assert result["execution_started"] is False + assert result["delegation_id"] == "sd03b-refused-reservation" + assert result["durable_reservation_released"] is reservation_released + assert result["results"][0]["missed_steer"] == state.missed_steer + state.writer.finalize.assert_called_once() + state.live.update_manifest_statuses.assert_called_once() + manifest_id, entries = state.live.update_manifest_statuses.call_args.args + assert manifest_id == "sd03b-delegation" + assert entries[0]["status"] in {"rejected", "error", "cancelled"} + assert entries[0]["missed_steer"] == state.missed_steer + + +def test_delegate_dispatch_uncertain_keeps_children_open_without_sync(sd03b_delegate_state): + state = sd03b_delegate_state + state.dispatch.return_value = { + "status": "dispatch_uncertain", "error_code": "scheduling_uncertain", + "execution_started": None, "delegation_id": "sd03b-delegation", + "error": "inert uncertain submit", + } + result = json.loads(state.caller.delegate_task( + goal="inert child", background=True, parent_agent=state.parent + )) + state.run.assert_not_called() + state.close_steering.assert_not_called() + state.unregister.assert_not_called() + state.child.close.assert_not_called() + state.child._drain_pending_steer.assert_not_called() + state.child.interrupt.assert_not_called() + state.writer.finalize.assert_not_called() + state.live.update_manifest_statuses.assert_not_called() + assert state.registry[state.child._subagent_id]["agent"] is state.child + assert state.registry[state.child._subagent_id]["accepting_steer"] is True + assert callable(state.dispatch.call_args.kwargs["runner"]) + assert result["status"] == "dispatch_uncertain" + assert result["error_code"] == "scheduling_uncertain" + assert result["execution_started"] is None + assert result["delegation_id"] == "sd03b-delegation" + + +def test_delegate_typed_pool_capacity_runs_single_inert_child_once(sd03b_delegate_state): + state = sd03b_delegate_state + state.dispatch.return_value = { + "status": "rejected", "error_code": "pool_capacity", + "execution_started": False, "error": "inert pool capacity", + } + result = json.loads(state.caller.delegate_task( + goal="inert child", background=True, parent_agent=state.parent + )) + state.run.assert_called_once() + assert state.run.call_args.args[:3] == (0, "inert child", state.child) + assert result["results"][0]["status"] == "completed" + assert result["results"][0]["summary"] == "inert once" + assert "SYNCHRONOUSLY" in result["note"] diff --git a/tests/tools/test_mcp_parked_self_probe.py b/tests/tools/test_mcp_parked_self_probe.py index 61a1a6578e341..9883921efbcab 100644 --- a/tests/tools/test_mcp_parked_self_probe.py +++ b/tests/tools/test_mcp_parked_self_probe.py @@ -68,6 +68,16 @@ async def _fast_sleep(_delay, *a, **kw): "revived_registration": 0, } + def _register(name, server, config): + # Production discovery registers an owned revival before readiness. + assert name == "srv" + assert mcp_tool._servers.get(name) is server + assert not server._ready.is_set() + state["revived_registration"] += 1 + return ["srv__tool"] + + monkeypatch.setattr(mcp_tool, "_register_server_tools", _register) + async def _scenario(): class _Task(MCPServerTask): def _is_http(self): @@ -77,11 +87,6 @@ def _deregister_tools(self): state["deregistered"] += 1 self._registered_tool_names = [] - def _register_discovered_tools_if_needed(self): - if self._ready.is_set() and not self._registered_tool_names: - state["revived_registration"] += 1 - self._registered_tool_names = ["srv__tool"] - async def _run_stdio(self, config): state["transport_calls"] += 1 if state["transport_calls"] == 1: @@ -95,12 +100,21 @@ async def _run_stdio(self, config): raise RuntimeError("backend still down") # Backend recovered: establish a session and park in the # lifecycle wait like the real transport does. - self.session = object() - self._register_discovered_tools_if_needed() + assert not self._ready.is_set() + self.session = SimpleNamespace( + list_tools=AsyncMock( + return_value=SimpleNamespace(tools=[SimpleNamespace(name="tool")]), + ) + ) + # Match _run_stdio: discover/publish first, signal ready after. + await self._discover_tools() + assert not self._ready.is_set() + self._ready.set() await self._wait_for_lifecycle_event() task = _Task("srv") task._registered_tool_names = ["srv__tool"] + monkeypatch.setitem(mcp_tool._servers, task.name, task) run_task = asyncio.ensure_future(task.run({"command": "x"})) @@ -117,7 +131,7 @@ async def _run_stdio(self, config): state["backend_up"] = True for _ in range(200): await _real_sleep(0.01) - if task.session is not None: + if task.session is not None and task._ready.is_set(): break assert task.session is not None, ( @@ -127,6 +141,9 @@ async def _run_stdio(self, config): assert state["revived_registration"] >= 1, ( "revived server did not re-register its tools" ) + assert task._ready.is_set(), "revived server never completed discovery" + assert task._registered_tool_names == ["srv__tool"] + assert [tool.name for tool in task._tools] == ["tool"] task._shutdown_event.set() task._reconnect_event.set() diff --git a/tests/tools/test_mcp_readiness_status.py b/tests/tools/test_mcp_readiness_status.py new file mode 100644 index 0000000000000..a2b478ecc2643 --- /dev/null +++ b/tests/tools/test_mcp_readiness_status.py @@ -0,0 +1,83 @@ +"""Status boundaries only; no MCP processes, RPCs or provider calls.""" + +import pytest + + +@pytest.fixture +def status_runtime(monkeypatch): + import tools.mcp_tool as mcp_tool + + monkeypatch.setattr(mcp_tool, "_servers", {}) + monkeypatch.setattr(mcp_tool, "_server_connecting", set()) + monkeypatch.setattr(mcp_tool, "_server_connect_errors", {}) + monkeypatch.setattr( + mcp_tool, "_load_mcp_config", lambda: {"cea_graph": {"command": "inert"}} + ) + server = mcp_tool.MCPServerTask("cea_graph") + mcp_tool._servers["cea_graph"] = server + return mcp_tool, server + + +def test_session_before_discovery_is_connecting(status_runtime): + mcp_tool, server = status_runtime + # The actual stdio path assigns this before awaiting discovery. A retained + # task can be reconnecting without appearing in _server_connecting. + server.session = object() + status = mcp_tool.get_mcp_status()[0] + assert status["status"] == "connecting" + assert status["connected"] is False + assert status["tools"] == 0 + + +def test_ready_session_with_no_tools_is_still_connected(status_runtime): + mcp_tool, server = status_runtime + server.session = object() + server._ready.set() + # Resource/prompt-only servers legitimately have zero tool definitions. + status = mcp_tool.get_mcp_status()[0] + assert status["status"] == "connected" + assert status["connected"] is True + assert status["tools"] == 0 + + +def test_ready_failure_is_not_connection_success(status_runtime): + mcp_tool, server = status_runtime + server.session = object() + server._ready.set() # start() also wakes its waiter on failure. + server._error = RuntimeError("Connection closed") + status = mcp_tool.get_mcp_status()[0] + assert status["status"] == "failed" + assert status["connected"] is False + assert "Connection closed" in status["error"] + + +def test_retained_failure_reports_task_error_without_discovery_map(status_runtime): + mcp_tool, server = status_runtime + server._error = RuntimeError("Connection closed") + status = mcp_tool.get_mcp_status()[0] + assert status["status"] == "failed" + assert status["connected"] is False + + +def test_recovered_ready_session_overrides_old_discovery_error(status_runtime): + mcp_tool, server = status_runtime + server.session = object() + server._ready.set() + mcp_tool._server_connect_errors["cea_graph"] = "older attempt failed" + assert mcp_tool.get_mcp_status()[0]["status"] == "connected" + + +def test_ready_event_without_session_is_not_connected(status_runtime): + mcp_tool, server = status_runtime + server._ready.set() + assert mcp_tool.get_mcp_status()[0]["connected"] is False + + +def test_disabled_unstarted_server_keeps_disabled_status(status_runtime, monkeypatch): + mcp_tool, _server = status_runtime + mcp_tool._servers.clear() + monkeypatch.setattr( + mcp_tool, "_load_mcp_config", + lambda: {"cea_graph": {"command": "inert", "enabled": False}}, + ) + assert mcp_tool.get_mcp_status()[0]["status"] == "disabled" diff --git a/tests/tools/test_mcp_stdio_diagnostics.py b/tests/tools/test_mcp_stdio_diagnostics.py new file mode 100644 index 0000000000000..1243d7077068a --- /dev/null +++ b/tests/tools/test_mcp_stdio_diagnostics.py @@ -0,0 +1,239 @@ +"""Actual diagnostic/stdio orchestration with inert transport seams only.""" + +import asyncio +import json +import os +from contextlib import asynccontextmanager +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest + + +@pytest.fixture +def diagnostics(tmp_path, monkeypatch): + import tools.mcp_tool as mcp_tool + from hermes_constants import reset_hermes_home_override, set_hermes_home_override + + monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setattr(mcp_tool, "_mcp_stderr_log_files", {}, raising=False) + if hasattr(mcp_tool, "_mcp_stderr_log_fh"): + monkeypatch.setattr(mcp_tool, "_mcp_stderr_log_fh", None) + token = set_hermes_home_override(str(tmp_path / "profile")) + try: + yield mcp_tool + finally: + reset_hermes_home_override(token) + for fh in mcp_tool._mcp_stderr_log_files.values(): + fh.close() + legacy = getattr(mcp_tool, "_mcp_stderr_log_fh", None) + if legacy is not None: + legacy.close() + + +def test_existing_index_seam_follows_the_current_config_home(diagnostics, tmp_path): + from hermes_constants import reset_hermes_home_override, set_hermes_home_override + + for name in ("alpha", "beta"): + token = set_hermes_home_override(str(tmp_path / name)) + try: + fh = diagnostics._get_mcp_stderr_log() + fh.write(name + "\n") + fh.flush() + finally: + reset_hermes_home_override(token) + assert (tmp_path / "alpha" / "logs" / "mcp-stderr.log").read_text() == "alpha\n" + assert (tmp_path / "beta" / "logs" / "mcp-stderr.log").read_text() == "beta\n" + + +def test_profile_and_attempt_output_remain_distinct(diagnostics, tmp_path): + from hermes_constants import reset_hermes_home_override, set_hermes_home_override + + mcp_tool = diagnostics + records = [] + for home, text in [("alpha", "first"), ("beta", "second"), ("alpha", "third")]: + token = set_hermes_home_override(str(tmp_path / home)) + try: + capture = mcp_tool._begin_stdio_diagnostic("cea_graph") + capture["stream"].write(text + "\n") + capture["status"] = "closed" + mcp_tool._finish_stdio_diagnostic(capture) + assert capture["stream"].closed + records.append(capture) + finally: + reset_hermes_home_override(token) + assert len({record["stderr_path"] for record in records}) == 3 + assert len({record["attempt_id"] for record in records}) == 3 + for record, text in zip(records, ["first", "second", "third"]): + path = Path(record["stderr_path"]) + assert path.parent.parent.parent == Path(record["config_home"]) + content = path.read_text() + assert text + "\n" in content + assert all(other + "\n" not in content for other in {"first", "second", "third"} - {text}) + if os.name == "posix": + assert path.stat().st_mode & 0o777 == 0o600 + for home, expected in [("alpha", 2), ("beta", 1)]: + index = (tmp_path / home / "logs" / "mcp-stderr.log").read_text().splitlines() + entries = [json.loads(line) for line in index] + assert len(entries) == expected + assert all(entry["config_home"] == str(tmp_path / home) for entry in entries) + assert mcp_tool._mcp_stdio_diagnostic.get() is None + + +def test_nested_attempts_restore_context_and_close_only_their_stream(diagnostics): + mcp_tool = diagnostics + outer = mcp_tool._begin_stdio_diagnostic("cea_graph") + inner = mcp_tool._begin_stdio_diagnostic("pilot_bridge") + inner["status"] = "closed" + mcp_tool._finish_stdio_diagnostic(inner) + assert inner["stream"].closed + assert not outer["stream"].closed + assert mcp_tool._mcp_stdio_diagnostic.get() is outer + outer["status"] = "closed" + mcp_tool._finish_stdio_diagnostic(outer) + assert mcp_tool._mcp_stdio_diagnostic.get() is None + + +def test_missing_capture_reports_discarded_without_retaining_error_text( + diagnostics, monkeypatch, caplog +): + import builtins + + real_open = builtins.open + + def unavailable(path, *args, **kwargs): + if str(path) != os.devnull: + raise OSError("MUST_NOT_RETAIN_SECRET") + return real_open(path, *args, **kwargs) + + monkeypatch.setattr(builtins, "open", unavailable) + capture = diagnostics._begin_stdio_diagnostic("cea_graph") + capture["status"] = "closed" + diagnostics._finish_stdio_diagnostic(capture) + assert capture["destination"] == "discarded" + assert "discarded" in caplog.text + assert "MUST_NOT_RETAIN_SECRET" not in caplog.text + assert capture["stream"].closed + assert diagnostics._mcp_stdio_diagnostic.get() is None + + +@pytest.mark.parametrize("phase", ["negotiate", "discover_tools"]) +@pytest.mark.parametrize("error_type", [RuntimeError, asyncio.CancelledError]) +def test_actual_stdio_failure_preserves_error_and_phase_without_live_io( + diagnostics, monkeypatch, caplog, phase, error_type +): + import tools.osv_check + + mcp_tool = diagnostics + monkeypatch.setattr(mcp_tool, "_ensure_mcp_sdk", lambda: True) + monkeypatch.setattr(mcp_tool, "StdioServerParameters", SimpleNamespace, raising=False) + monkeypatch.setattr(mcp_tool, "_resolve_stdio_command", lambda command, env: (command, env)) + monkeypatch.setattr(mcp_tool, "_wrap_command_with_watchdog", lambda command, args: (command, args)) + monkeypatch.setattr(mcp_tool, "_kill_orphaned_mcp_children", lambda: None) + monkeypatch.setattr(mcp_tool, "_snapshot_child_pids", lambda: set()) + monkeypatch.setattr(mcp_tool, "_MCP_NOTIFICATION_TYPES", False) + monkeypatch.setattr(mcp_tool, "_MCP_LOGGING_CALLBACK_SUPPORTED", False) + monkeypatch.setattr(tools.osv_check, "check_package_for_malware", lambda *_args: None) + + @asynccontextmanager + async def inert_stdio(_params, *, errlog): + errlog.write("inert child stderr\n") + yield None, None + + @asynccontextmanager + async def inert_session(*_args, **_kwargs): + yield object() + + monkeypatch.setattr(mcp_tool, "stdio_client", inert_stdio, raising=False) + monkeypatch.setattr(mcp_tool, "ClientSession", inert_session, raising=False) + server = mcp_tool.MCPServerTask("cea_graph") + monkeypatch.setattr(mcp_tool, "_servers", {"cea_graph": server}) + monkeypatch.setattr(mcp_tool, "_server_connecting", set()) + monkeypatch.setattr(mcp_tool, "_server_connect_errors", {}) + monkeypatch.setattr(mcp_tool, "_load_mcp_config", lambda: {"cea_graph": {"command": "inert"}}) + original = error_type("MUST_NOT_RETAIN_SECRET") + + async def discover(_server): + # Exercise get_mcp_status during the actual retained-task window: + # session assigned, tools not discovered, readiness not yet signalled. + assert mcp_tool.get_mcp_status()[0]["status"] == "connecting" + raise original + + monkeypatch.setattr( + mcp_tool.MCPServerTask, + "_negotiate_session", + AsyncMock( + side_effect=original if phase == "negotiate" else None, + return_value=SimpleNamespace(), + ), + ) + monkeypatch.setattr(mcp_tool.MCPServerTask, "_discover_tools", discover) + with pytest.raises(error_type) as caught: + asyncio.run(server._run_stdio({"command": "inert-never-executed"})) + assert caught.value is original + assert not server._ready.is_set() + assert mcp_tool._mcp_stdio_diagnostic.get() is None + records = [json.loads(record.getMessage().split(": ", 1)[1]) for record in caplog.records + if record.getMessage().startswith("MCP stdio attempt ended: ")] + assert len(records) == 1 + record = records[0] + assert record["phase"] == phase + assert record["status"] == "failed" + assert record["exception_type"] == error_type.__name__ + assert record["destination"] == "file" + assert "MUST_NOT_RETAIN_SECRET" not in caplog.text + assert "MUST_NOT_RETAIN_SECRET" not in Path(record["stderr_path"]).read_text() + assert "inert child stderr" in Path(record["stderr_path"]).read_text() + + +def test_actual_run_retry_clears_readiness_before_second_stdio_discovery(diagnostics, monkeypatch): + import tools.osv_check + + mcp_tool = diagnostics + monkeypatch.setattr(mcp_tool, "_ensure_mcp_sdk", lambda: True) + monkeypatch.setattr(mcp_tool, "StdioServerParameters", SimpleNamespace, raising=False) + monkeypatch.setattr(mcp_tool, "_resolve_stdio_command", lambda command, env: (command, env)) + monkeypatch.setattr(mcp_tool, "_wrap_command_with_watchdog", lambda command, args: (command, args)) + monkeypatch.setattr(mcp_tool, "_kill_orphaned_mcp_children", lambda: None) + monkeypatch.setattr(mcp_tool, "_snapshot_child_pids", lambda: set()) + monkeypatch.setattr(mcp_tool, "_MCP_NOTIFICATION_TYPES", False) + monkeypatch.setattr(mcp_tool, "_MCP_LOGGING_CALLBACK_SUPPORTED", False) + monkeypatch.setattr(tools.osv_check, "check_package_for_malware", lambda *_args: None) + monkeypatch.setattr(mcp_tool.asyncio, "sleep", AsyncMock()) + + @asynccontextmanager + async def inert_stdio(_params, *, errlog): + errlog.write("inert retry stderr\n") + yield None, None + + @asynccontextmanager + async def inert_session(*_args, **_kwargs): + yield object() + + monkeypatch.setattr(mcp_tool, "stdio_client", inert_stdio, raising=False) + monkeypatch.setattr(mcp_tool, "ClientSession", inert_session, raising=False) + monkeypatch.setattr(mcp_tool.MCPServerTask, "_negotiate_session", AsyncMock(return_value=SimpleNamespace())) + server = mcp_tool.MCPServerTask("cea_graph") + monkeypatch.setattr(mcp_tool, "_servers", {"cea_graph": server}) + monkeypatch.setattr(mcp_tool, "_server_connecting", set()) + monkeypatch.setattr(mcp_tool, "_server_connect_errors", {}) + monkeypatch.setattr(mcp_tool, "_load_mcp_config", lambda: {"cea_graph": {"command": "inert"}}) + statuses = [] + + async def discover(_server): + statuses.append(mcp_tool.get_mcp_status()[0]["status"]) + if len(statuses) == 2: + server._shutdown_event.set() + + async def lifecycle(_server): + if len(statuses) == 1: + raise BrokenPipeError("inert transient failure after first readiness") + return "shutdown" + + monkeypatch.setattr(mcp_tool.MCPServerTask, "_discover_tools", discover) + monkeypatch.setattr(mcp_tool.MCPServerTask, "_wait_for_lifecycle_event", lifecycle) + asyncio.run(server.run({"command": "inert-never-executed", "sampling": {"enabled": False}, + "elicitation": {"enabled": False}})) + assert statuses == ["connecting", "connecting"] + assert server._ever_connected diff --git a/tests/tools/test_refresh_agent_mcp_tools.py b/tests/tools/test_refresh_agent_mcp_tools.py index b1aa95e12f560..4fbe8735f23a3 100644 --- a/tests/tools/test_refresh_agent_mcp_tools.py +++ b/tests/tools/test_refresh_agent_mcp_tools.py @@ -11,6 +11,8 @@ import threading import types +import pytest + from tools import mcp_tool @@ -224,3 +226,45 @@ def test_wait_returns_instantly_when_no_discovery_thread(monkeypatch): t0 = time.time() mcp_startup.wait_for_mcp_discovery() assert time.time() - t0 < 0.2 # never blocks on the bound when nothing's pending + + +@pytest.mark.parametrize("enabled", [None, ["context_engine"], [], ["coding"]]) +@pytest.mark.parametrize("disabled", [None, [], ["context_engine"], ["unrelated"]]) +def test_refresh_applies_context_engine_disabled_policy(monkeypatch, enabled, disabled): + """Actual refresh publishes only permitted late schemas and routing names.""" + calls = [] + schema = {"name": "policy_context_recover", "description": "Inert context witness", "parameters": {}} + agent = _agent(["policy_keep", "policy_context_recover"], enabled=enabled, disabled=disabled) + engine = types.SimpleNamespace(get_tool_schemas=lambda: calls.append("schemas") or [schema, schema]) + agent.context_compressor = engine + agent._context_engine_tool_names = {"policy_context_recover"} + + import model_tools + monkeypatch.setattr(model_tools, "get_tool_definitions", lambda **kw: [_tool("policy_keep"), _tool("mcp_policy_late")]) + added = mcp_tool.refresh_agent_mcp_tools(agent) + allowed = (enabled is None or "context_engine" in enabled) and "context_engine" not in (disabled or []) + names = {tool["function"]["name"] for tool in agent.tools} + expected = {"policy_keep", "mcp_policy_late"} | ({"policy_context_recover"} if allowed else set()) + assert names == agent.valid_tool_names == expected + assert added == {"mcp_policy_late"} + assert agent._context_engine_tool_names == ({"policy_context_recover"} if allowed else set()) + assert calls == (["schemas"] if allowed else []) + assert agent.context_compressor is engine + assert sum(tool["function"]["name"] == "policy_context_recover" for tool in agent.tools) == int(allowed) + + +def test_reinjection_does_not_claim_existing_registry_context_name(): + """Enabled context tools retain dedup and existing registry dispatch ownership.""" + agent = _agent([], enabled=["context_engine"], disabled=[]) + agent.context_compressor = types.SimpleNamespace(get_tool_schemas=lambda: [ + {"name": "registry_owned_context", "description": "", "parameters": {}}, + {"name": "engine_owned_context", "description": "", "parameters": {}}, + ]) + staged = [_tool("registry_owned_context")] + names = {"registry_owned_context"} + assert mcp_tool._reinject_post_build_tools(agent, staged, names) == {"engine_owned_context"} + assert names == {"registry_owned_context", "engine_owned_context"} + assert len(staged) == 2 + # The reinjector stages locals without touching the currently published pair. + assert agent.tools == [] + assert agent.valid_tool_names == set() diff --git a/tests/tui_gateway/test_goal_command.py b/tests/tui_gateway/test_goal_command.py index 35dffb2b95804..c3ff9cfb9c132 100644 --- a/tests/tui_gateway/test_goal_command.py +++ b/tests/tui_gateway/test_goal_command.py @@ -375,6 +375,7 @@ def run_conversation(message, **_kwargs): ) session = _turn_session(agent, session_key) session_holder["session"] = session + monkeypatch.setitem(server._sessions, "sid", session) server._run_prompt_submit("rid", "sid", session, "initial work") diff --git a/tests/tui_gateway/test_host_targeted_stop.py b/tests/tui_gateway/test_host_targeted_stop.py index f04e2a0d915b2..d224c0e267bc4 100644 --- a/tests/tui_gateway/test_host_targeted_stop.py +++ b/tests/tui_gateway/test_host_targeted_stop.py @@ -82,13 +82,23 @@ def handoff(sid, **kwargs): def test_cancelled_successor_cannot_clear_interrupt_or_start(monkeypatch): - session = {"running": True, "_turn_cancel_requested": True, - "history_lock": threading.Lock(), - "agent": types.SimpleNamespace(clear_interrupt=lambda: pytest.fail("cancel latch cleared"))} + interrupted = [] + session = {"running": True, "history_lock": threading.Lock(), + "agent": types.SimpleNamespace( + interrupt=lambda: interrupted.append(True), + clear_interrupt=lambda: pytest.fail("cancel latch cleared"))} + monkeypatch.setitem(server._sessions, "s", session) + monkeypatch.setattr(server, "_session_uses_compute_host", lambda _session: False) monkeypatch.setattr(server, "_start_inflight_turn", lambda *a: pytest.fail("cancelled turn admitted")) + assert server._interrupt_session_turn("s", session) is False + stopped_generation = session["_queued_prompt_generation"] assert server._run_prompt_submit("A", "s", session, "follow-up", display_kind="internal_notification") is False + assert interrupted == [True] assert session["running"] is False + assert session["inflight_turn"] is None + assert session["_turn_cancel_requested"] is True + assert session["_queued_prompt_generation"] == session["_last_stop_queue_generation"] == stopped_generation def test_real_host_rejects_stale_stop_then_interrupts_matching_request(tmp_path, monkeypatch): diff --git a/tests/tui_gateway/test_model_config_composition.py b/tests/tui_gateway/test_model_config_composition.py new file mode 100644 index 0000000000000..cf29006f43e7a --- /dev/null +++ b/tests/tui_gateway/test_model_config_composition.py @@ -0,0 +1,253 @@ +"""Real gateway configuration boundaries; external clients remain inert.""" +import copy +import threading +from types import SimpleNamespace + +import pytest + +from agent.context_engine import ContextEngine +from hermes_cli.config import load_config +from tests.hermes_cli.test_model_switch_route_contract import routes # noqa: F401 +from tui_gateway import server +from tests.tui_gateway.test_failed_turn_retention import _session, turn_env # noqa: F401 + +_REAL_MODEL_SYNC = server._sync_agent_model_with_config + + +class InertEngine(ContextEngine): + @property + def name(self): + return "composition-inert" + + def update_from_response(self, usage): + pass + + def should_compress(self, prompt_tokens=None): + return False + + def compress(self, messages, current_tokens=None): + return messages + + +@pytest.fixture +def world(routes, monkeypatch): + routes.config["context"] = {"engine": "composition-inert"} + routes.config["agent"] = {"reasoning_overrides": {"chosen-model": "high"}} + routes.write() + monkeypatch.setenv("HERMES_IGNORE_RULES", "1") + monkeypatch.setattr(server, "_load_cfg", load_config) + monkeypatch.setattr(server, "_load_enabled_toolsets", lambda *a: []) + monkeypatch.setattr(server, "_load_fallback_model", lambda: []) + monkeypatch.setattr(server, "_load_service_tier", lambda: None) + monkeypatch.setattr(server, "_load_provider_routing", lambda: {}) + monkeypatch.setattr(server, "_get_db", lambda: None) + monkeypatch.setattr("hermes_cli.mcp_startup.wait_for_mcp_discovery", lambda: None) + monkeypatch.setattr("tui_gateway.entry.wait_for_mcp_discovery", lambda: None) + monkeypatch.setattr("plugins.context_engine.load_context_engine", lambda *a, **k: InertEngine()) + monkeypatch.setattr("agent.model_metadata.get_model_context_length", lambda *a, **k: 128_000) + monkeypatch.setattr("run_agent.get_tool_definitions", lambda *a, **k: []) + monkeypatch.setattr("run_agent.check_toolset_requirements", lambda *a, **k: {}) + monkeypatch.setattr("run_agent.OpenAI", lambda **k: SimpleNamespace(kwargs=k)) + monkeypatch.setattr("hermes_cli.timeouts.get_provider_request_timeout", lambda *a, **k: None) + monkeypatch.setattr(server, "_restart_slash_worker", lambda *a: None) + monkeypatch.setattr(server, "_persist_live_session_runtime", lambda *a: None) + monkeypatch.setattr(server, "_persist_live_session_system_prompt", lambda *a: None) + monkeypatch.setattr(server, "_probe_credentials", lambda *a: None) + monkeypatch.setattr(server, "_git_branch_for_cwd", lambda *a: None) + monkeypatch.setattr(server, "_project_info_for_cwd", lambda *a: None) + monkeypatch.setattr("hermes_cli.banner.get_available_skills", lambda: {}) + monkeypatch.setattr("hermes_cli.banner.get_update_result", lambda **k: None) + monkeypatch.setattr("tools.mcp_tool.get_mcp_status", lambda: []) + agents = [] + + def build(override=None): + value = server._make_agent("selected", "inert-session", model_override=override) + value._create_openai_client = lambda kwargs, **k: SimpleNamespace(kwargs=kwargs.copy()) + agents.append(value) + return value + + yield SimpleNamespace(routes=routes, build=build) + server._sessions.clear() + for value in agents: + db = getattr(value, "_session_db", None) + if db is not None and getattr(value, "_owns_session_db", False): + db.close() + + +def outgoing(value): + return value._build_api_kwargs([{"role": "user", "content": "inert input"}], []) + + +def selected_session(value, monkeypatch): + session = {"agent": value, "history": [], "history_lock": threading.Lock(), + "model_override": {"model": "model-a", "provider": "endpoint-a"}, + "model_verified_for": ("endpoint-a", "model-a")} + monkeypatch.setattr(server, "_sessions", {"selected": session}) + monkeypatch.setattr(server, "_emit", lambda *a, **k: None) + return session + + +@pytest.mark.parametrize("actual", [None, {"enabled": False}, {"enabled": True, "effort": "medium"}]) +def test_once_restores_actual_reasoning_and_primary(world, monkeypatch, actual): + value = world.build() + # A is normally initialized, without the preparatory switch used by older + # runtime fixtures. A manual session value may differ from primary state. + assert "reasoning_config" not in value._primary_runtime + assert value.reasoning_config is None + value.reasoning_config = copy.deepcopy(actual) + session = selected_session(value, monkeypatch) + pin = copy.deepcopy(session["model_override"]) + result = server._apply_model_switch("selected", session, + "chosen-model --provider endpoint-b --once", confirm_expensive_model=True, + persist_override=False) + assert result["scope"] == "once" + assert value.reasoning_config == {"enabled": True, "effort": "high"} + server._restore_agent_model_runtime(value, session.pop("one_turn_model_restore")) + assert (value.model, value.base_url) == ("model-a", "http://a.invalid/v1") + assert value.reasoning_config == actual + assert value._primary_runtime["reasoning_config"] == actual + assert session["model_override"] == pin + # A subsequent primary retry must preserve the restored session setting. + value._fallback_activated = True + assert value._restore_primary_runtime() + assert value.reasoning_config == actual + + +def test_snapshot_copies_actual_manual_reasoning(world): + value = world.build() + value.reasoning_config = {"enabled": True, "effort": "medium"} + value._primary_runtime["reasoning_config"] = {"enabled": True, "effort": "low"} + snapshot = server._snapshot_agent_model_runtime(value) + value.reasoning_config["effort"] = "high" + assert snapshot["reasoning_config"] == {"enabled": True, "effort": "medium"} + assert snapshot["primary_runtime"]["reasoning_config"] == {"enabled": True, "effort": "medium"} + + +@pytest.mark.parametrize("setting,expected", [ + (None, None), ("none", {"enabled": False}), + ("medium", {"enabled": True, "effort": "medium"}), +]) +@pytest.mark.parametrize("outcome", ["complete", "returned_error", "raised_error", "interrupted"]) +def test_real_pinned_once_turn_restores_session_reasoning(world, monkeypatch, setting, expected, outcome): + value = world.build() + assert value.reasoning_config is None + # Initialize the real agent before this fixture substitutes inline turn + # threads; the shared Thread substitution would deadlock QueueListener. + turn_env.__wrapped__(monkeypatch, world.routes.home) + session = selected_session(value, monkeypatch) + session.update(_session(value, session_key=value.session_id, running=True, + profile_home=str(world.routes.home), model_override=session["model_override"])) + monkeypatch.setattr(server, "_sync_agent_model_with_config", _REAL_MODEL_SYNC) + monkeypatch.setattr(server, "_sync_agent_compression_with_config", lambda *a: None) + monkeypatch.setattr(server, "_voice_tts_enabled", lambda: False) + monkeypatch.setattr(server, "_load_interim_assistant_messages", lambda: False) + monkeypatch.setattr(server, "_start_usage_ticker", lambda *a: (threading.Event(), SimpleNamespace(join=lambda *a, **k: None))) + monkeypatch.setattr(server, "_notify_session_boundary", lambda *a: None) + monkeypatch.setattr(server, "record_turn_start", lambda *a, **k: None) + monkeypatch.setattr(server, "_retire_turn_marker", lambda *a: None) + if setting is not None: + ack = server._methods["config.set"]("reasoning-request", { + "key": "reasoning", "value": setting, "session_id": "selected", "scope": "session"}) + assert "error" not in ack + assert value.reasoning_config == expected + pin = copy.deepcopy(session["model_override"]) + original_config = (world.routes.home / "config.yaml").read_bytes() + selected = server._apply_model_switch("selected", session, + "chosen-model --provider endpoint-b --once", confirm_expensive_model=True, + persist_override=False) + assert selected["scope"] == "once" + observations = [] + + def offline(*a, **k): + observations.append((value.model, copy.deepcopy(value.reasoning_config))) + assert "one_turn_model_restore" not in session + if outcome == "raised_error": + raise RuntimeError("inert conversation failure") + if outcome == "returned_error": + return {"error": "inert returned failure", "final_response": "", "completed": False, "api_calls": 0} + return {"final_response": "inert", "completed": False, "api_calls": 0, + "interrupted": outcome == "interrupted"} + + monkeypatch.setattr(value, "run_conversation", offline) + server._run_prompt_submit("once-request", "selected", session, "inert input") + assert observations == [("chosen-model", {"enabled": True, "effort": "high"})] + assert (value.model, value.base_url) == ("model-a", "http://a.invalid/v1") + assert value.reasoning_config == expected + assert value._primary_runtime["reasoning_config"] == expected + assert session["model_override"] == pin + assert not session.get("_one_turn_model_runtime") + assert (world.routes.home / "config.yaml").read_bytes() == original_config + + +def test_failed_route_restore_does_not_apply_saved_reasoning(world): + value = world.build() + value.reasoning_config = None + snapshot = server._snapshot_agent_model_runtime(value) + value.reasoning_config = {"enabled": True, "effort": "high"} + value._restore_primary_runtime = lambda: False + + def refuse(**kw): + raise ValueError("inert refused route") + + value.switch_model = refuse + with pytest.raises(ValueError, match="inert refused route"): + server._restore_agent_model_runtime(value, snapshot) + assert value.reasoning_config == {"enabled": True, "effort": "high"} + + +def test_legacy_snapshot_without_reasoning_keeps_legacy_restore(world): + value = world.build() + snapshot = server._snapshot_agent_model_runtime(value) + snapshot.pop("reasoning_config", None) + snapshot["primary_runtime"].pop("reasoning_config", None) + value.reasoning_config = {"enabled": True, "effort": "high"} + server._restore_agent_model_runtime(value, snapshot) + assert value.reasoning_config == {"enabled": True, "effort": "high"} + + +@pytest.mark.parametrize("global_cap,provider_cap,expected", [ + (None, 7, 7), (11, 7, 11), ("11", 7, 11), + (None, None, None), (None, 0, None), (None, -1, None), + (None, "7", None), (None, True, None), + (0, 7, None), (False, 7, None), ("invalid", 7, None), +]) +def test_actual_factory_provider_global_cap_precedence(world, global_cap, provider_cap, expected): + cfg = world.routes.config + cfg["model"]["max_tokens"] = global_cap + cfg["providers"]["endpoint-a"]["max_output_tokens"] = provider_cap + world.routes.write() + from hermes_cli.runtime_provider import resolve_runtime_provider + resolved = resolve_runtime_provider(requested="endpoint-a", target_model="model-a") + expected_resolved = provider_cap if isinstance(provider_cap, int) and provider_cap > 0 else None + assert resolved.get("max_output_tokens") == expected_resolved, resolved + value = world.build() + assert value.max_tokens == expected + kwargs = outgoing(value) + if expected is not None: + for key, amount in value._max_tokens_param(expected).items(): + assert kwargs[key] == amount + + +def test_fresh_pinned_provider_caps_remain_independent(world): + cfg = world.routes.config + cfg["providers"]["endpoint-a"]["max_output_tokens"] = 7 + cfg["providers"]["endpoint-b"]["max_output_tokens"] = 13 + world.routes.write() + a = world.build({"model": "model-a", "provider": "endpoint-a"}) + b = world.build({"model": "model-b", "provider": "endpoint-b"}) + assert (a.max_tokens, b.max_tokens) == (7, 13) + assert (a.base_url, b.base_url) == ("http://a.invalid/v1", "http://b.invalid/v1") + + +def test_clear_global_then_provider_cap_at_fresh_construction(world): + cfg = world.routes.config + cfg["model"]["max_tokens"] = 11 + cfg["providers"]["endpoint-a"]["max_output_tokens"] = 7 + world.routes.write() + assert world.build().max_tokens == 11 + cfg["model"].pop("max_tokens") + world.routes.write() + assert world.build().max_tokens == 7 + cfg["providers"]["endpoint-a"].pop("max_output_tokens") + world.routes.write() + assert world.build().max_tokens is None diff --git a/tests/tui_gateway/test_model_intent_admission_order.py b/tests/tui_gateway/test_model_intent_admission_order.py new file mode 100644 index 0000000000000..70e3c2139c674 --- /dev/null +++ b/tests/tui_gateway/test_model_intent_admission_order.py @@ -0,0 +1,535 @@ +"""Model intent ordering through real gateway admission and inert runtimes.""" +import copy +import errno +import threading +from concurrent.futures import ThreadPoolExecutor +from types import SimpleNamespace + +import pytest + +from hermes_cli import model_switch as ms +from hermes_state import SessionDB +from tests.hermes_cli.test_model_switch_route_contract import routes # noqa: F401 +from tests.hermes_cli.test_model_switch_runtime_boundary import agent, assert_route # noqa: F401 +from tests.tui_gateway.test_failed_turn_retention import turn_env # noqa: F401 +from tests.tui_gateway.test_model_once_turn_runtime_owner import live_turn # noqa: F401 +from tests.tui_gateway.test_model_switch_target_session_contract import gateway # noqa: F401 +from tests.tui_gateway.test_config_dispatch_responsiveness import Transport +from tui_gateway import server +from tui_gateway.transport import StdioTransport + +REAL_THREAD = threading.Thread +REAL_CONFIG_SYNC = server._sync_agent_model_with_config +REAL_RESTART = server._restart_slash_worker + + +@pytest.fixture +def intent_turn(live_turn, routes, monkeypatch): + session, events = live_turn + session['running'] = False + session['transport'] = Transport() + routes.config['providers']['endpoint-c'] = { + 'name': 'Endpoint C', 'base_url': 'http://c.invalid/v1', + 'api_key': 'synthetic-key-c', 'default_model': 'model-c'} + routes.write() + db = SessionDB(routes.home / 'state.db') + db.create_session(session['session_key'], source='tui') + session['agent']._session_db = db + monkeypatch.setattr(server, '_db', db) + monkeypatch.setattr(server, '_session_uses_compute_host', lambda *a, **kw: False) + monkeypatch.setattr(server, '_voice_mode_enabled', lambda: False) + monkeypatch.setattr(server, '_ensure_active_session_slot', lambda *a: None) + monkeypatch.setattr(server, '_sync_bot_capabilities', lambda *a: None) + yield session, events + db.close() + + +def select(model, provider, *, once=False, confirmed=True): + raw = f'{model} --provider {provider} ' + ('--once' if once else '--session') + return {'id': 'pick-' + model, 'method': 'config.set', 'params': { + 'session_id': 'selected', 'key': 'model', 'value': raw, + 'confirm_expensive_model': confirmed}} + + +def run_turn(session, monkeypatch, rid='offline-turn'): + observations = [] + value = session['agent'] + + def offline(*a, **kw): + info = server._session_info(value, session) + observations.append((value.model, value.provider, info['model'], info['provider'])) + return {'final_response': 'offline', 'completed': False, 'api_calls': 0} + + monkeypatch.setattr(value, 'run_conversation', offline) + session['running'] = True + server._run_prompt_submit(rid, 'selected', session, 'offline input') + session['_run_thread'].join(5) + assert not session['_run_thread'].is_alive() + return observations + + +@pytest.mark.parametrize('pause', ['admitted', 'executing', 'steered']) +def test_slow_idle_selection_cannot_repoint_admitted_send(intent_turn, monkeypatch, pause): + session, _ = intent_turn + value = session['agent'] + transport = session['transport'] + entered, resolve_release = threading.Event(), threading.Event() + turn_entered, turn_release = threading.Event(), threading.Event() + real_switch = ms.switch_model + samples = {} + monkeypatch.setattr(server.threading, 'Thread', REAL_THREAD) + + def slow(*a, **kw): + entered.set() + assert resolve_release.wait(5) + return real_switch(*a, **kw) + + def ready(*a, **kw): + if pause != 'executing': + turn_entered.set() + assert turn_release.wait(5) + return None + + def offline(*a, **kw): + samples['start'] = (value.model, value.provider, value.client) + if pause == 'executing': + turn_entered.set() + assert turn_release.wait(5) + samples['end'] = (value.model, value.provider, value.client) + return {'final_response': 'offline', 'completed': False, 'api_calls': 0} + + monkeypatch.setattr(ms, 'switch_model', slow) + monkeypatch.setattr(server, '_wait_agent_for_prompt', ready) + monkeypatch.setattr(value, 'run_conversation', offline) + original = (value.model, value.provider, value.client) + with ThreadPoolExecutor(max_workers=1) as pool: + monkeypatch.setattr(server, '_pool', pool) + assert server.dispatch(select('model-b', 'endpoint-b'), transport) is None + try: + assert entered.wait(5) + ack = server.dispatch({'id': 'send-a', 'method': 'prompt.submit', 'params': { + 'session_id': 'selected', 'text': 'offline input'}}, transport) + assert not ack.get('error'), ack + assert ack['result']['status'] == 'streaming' + assert turn_entered.wait(5) + assert session['running'] + resolve_release.set() + assert transport.done.wait(5) + if pause == 'steered': + admitted_snapshot = session['inflight_turn'] + correction = server.dispatch({'id': 'steer-a', 'method': 'session.steer', + 'params': {'session_id': 'selected', 'text': 'offline correction'}}, transport) + assert correction['result']['status'] == 'queued' + assert session['inflight_turn'] is not admitted_snapshot + assert session['inflight_turn']['started_at'] is admitted_snapshot['started_at'] + finally: + resolve_release.set() + turn_release.set() + thread = session.get('_run_thread') + while thread is not None: + thread.join(5) + assert not thread.is_alive() + successor = session.get('_run_thread') + if successor is thread: + break + thread = successor + assert not session['_run_thread'].is_alive() + assert samples == {'start': original, 'end': original} + result = transport.frames[0]['result'] + assert result['deferred'] is True and result['confirm_required'] is False + assert session['pending_model_switch']['display_model'] == 'model-b' + assert_route(value, 'http://a.invalid/v1', 'synthetic-key-a') + # The intent becomes eligible only after this admission has settled. + monkeypatch.setattr(ms, 'switch_model', real_switch) + assert run_turn(session, monkeypatch, 'next-b') == [ + ('model-b', 'endpoint-b', 'model-b', 'endpoint-b')] + assert_route(value, 'http://b.invalid/v1', 'synthetic-key-b') + + +def test_live_swap_and_once_publication_share_admission_lock(intent_turn, monkeypatch): + session, _ = intent_turn + value = session['agent'] + real_switch = value.switch_model + lock_observations = [] + + def switch(*a, **kw): + acquired = session['history_lock'].acquire(blocking=False) + if acquired: + session['history_lock'].release() + lock_observations.append(acquired) + return real_switch(*a, **kw) + + monkeypatch.setattr(value, 'switch_model', switch) + response = server.handle_request(select('model-b', 'endpoint-b', once=True)) + assert not response.get('error'), response + assert lock_observations == [False] + snapshot, runtime = server._consume_one_turn_model_runtime(session, value) + assert snapshot['model'] == 'model-a' and runtime['active'] + assert_route(value, 'http://b.invalid/v1', 'synthetic-key-b') + + +def test_new_idle_choice_retires_older_deferred_choice(intent_turn, monkeypatch, routes): + session, _ = intent_turn + other_pending = {'raw': 'other --session', 'display_model': 'other'} + server._sessions['other'] = {'pending_model_switch': other_pending} + before = (routes.home / 'config.yaml').read_bytes() + session['running'] = True + queued = server.handle_request(select('model-b', 'endpoint-b')) + assert queued['result']['deferred'] is True + session['running'] = False + chosen = server.handle_request(select('model-c', 'endpoint-c')) + assert not chosen.get('error'), chosen + assert chosen['result']['value'] == 'model-c' + assert 'pending_model_switch' not in session + info = server._session_info(session['agent'], session) + assert (info['model'], info['provider']) == ('model-c', 'endpoint-c') + assert run_turn(session, monkeypatch) == [('model-c', 'endpoint-c', 'model-c', 'endpoint-c')] + assert_route(session['agent'], 'http://c.invalid/v1', 'synthetic-key-c') + assert server._sessions['other']['pending_model_switch'] is other_pending + assert (routes.home / 'config.yaml').read_bytes() == before + + +@pytest.mark.parametrize('failure', ['disabled', 'client', 'unconfirmed']) +def test_unsuccessful_idle_choice_preserves_deferred_intent(intent_turn, routes, monkeypatch, failure): + session, _ = intent_turn + value = session['agent'] + session['running'] = True + assert server.handle_request(select('model-b', 'endpoint-b'))['result']['deferred'] + pending = session['pending_model_switch'] + session['running'] = False + old_client = value.client + if failure == 'disabled': + routes.config['providers']['endpoint-c']['enabled'] = False + routes.write() + elif failure == 'client': + create = value._create_openai_client + + def fail(kwargs, **kw): + if kwargs['base_url'] == 'http://c.invalid/v1': + raise RuntimeError('offline client construction failed') + return create(kwargs, **kw) + + monkeypatch.setattr(value, '_create_openai_client', fail) + else: + monkeypatch.setattr('hermes_cli.model_selection_guards.combined_selection_warning', + lambda *a, **kw: SimpleNamespace(message='offline consent required')) + response = server.handle_request(select('model-c', 'endpoint-c', confirmed=failure != 'unconfirmed')) + if failure == 'unconfirmed': + assert response['result']['confirm_required'] is True + else: + assert response['error']['code'] == 5001 + assert session['pending_model_switch'] is pending + assert value.client is old_client + assert_route(value, 'http://a.invalid/v1', 'synthetic-key-a') + assert session['history_lock'].acquire(blocking=False) + session['history_lock'].release() + assert run_turn(session, monkeypatch) == [('model-b', 'endpoint-b', 'model-b', 'endpoint-b')] + + +def test_replacing_unused_once_returns_to_original_runtime(intent_turn, monkeypatch, routes): + session, _ = intent_turn + pin = copy.deepcopy(session['model_override']) + before = (routes.home / 'config.yaml').read_bytes() + assert not server.handle_request(select('model-b', 'endpoint-b', once=True)).get('error') + original_restore = session['one_turn_model_restore'] + original_lease = session['_one_turn_model_runtime'] + assert not server.handle_request(select('model-c', 'endpoint-c', once=True)).get('error') + assert session['one_turn_model_restore'] is original_restore + assert session['_one_turn_model_runtime'] is not original_lease + assert run_turn(session, monkeypatch) == [('model-c', 'endpoint-c', 'model-c', 'endpoint-c')] + value = session['agent'] + assert (value.model, value.provider) == ('model-a', 'endpoint-a') + assert_route(value, 'http://a.invalid/v1', 'synthetic-key-a') + assert session['model_override'] == pin + info = server._session_info(value, session) + assert (info['model'], info['provider']) == ('model-a', 'endpoint-a') + assert not session.get('_one_turn_model_runtime') + assert (routes.home / 'config.yaml').read_bytes() == before + + +def test_durable_choice_supersedes_unused_once_baseline(intent_turn, monkeypatch): + session, _ = intent_turn + assert not server.handle_request(select('model-b', 'endpoint-b', once=True)).get('error') + assert not server.handle_request(select('model-c', 'endpoint-c')).get('error') + assert not session.get('one_turn_model_restore') and not session.get('_one_turn_model_runtime') + assert run_turn(session, monkeypatch) == [('model-c', 'endpoint-c', 'model-c', 'endpoint-c')] + assert session['model_override']['model'] == 'model-c' + assert_route(session['agent'], 'http://c.invalid/v1', 'synthetic-key-c') + + +@pytest.mark.parametrize('newer_choice', [False, True]) +def test_publication_error_does_not_discard_or_resurrect_old_intent(intent_turn, monkeypatch, newer_choice): + session, _ = intent_turn + session['running'] = True + assert server.handle_request(select('model-b', 'endpoint-b'))['result']['deferred'] + pending = session['pending_model_switch'] + session['running'] = False + collect = server._emit + + class FailingStream: + def write(self, text): + raise OSError(errno.EIO, 'offline event publication failed') + + failed_transport = StdioTransport(lambda: FailingStream(), threading.Lock()) + + def publish(name, sid, payload=None): + if name == 'session.info' and payload['model'] == 'model-c': + if newer_choice: + server._mirror_slash_side_effects('selected', session, + '/model model-d --provider endpoint-c --session') + return failed_transport.write(server._event_frame(name, sid, payload)) + return collect(name, sid, payload) + + monkeypatch.setattr(server, '_emit', publish) + response = server.handle_request(select('model-c', 'endpoint-c')) + monkeypatch.setattr(server, '_emit', collect) + assert response['error']['code'] == 5001 + assert session['history_lock'].acquire(blocking=False) + session['history_lock'].release() + if newer_choice: + assert not session.get('pending_model_switch') + assert session['model_override']['model'] == 'model-d' + assert run_turn(session, monkeypatch) == [('model-d', 'endpoint-c', 'model-d', 'endpoint-c')] + else: + assert session['pending_model_switch'] is pending + # The suffix error does not claim to roll back the already working C + # client. B remains the previously acknowledged next-turn intent. + assert session['agent'].model == 'model-c' + assert run_turn(session, monkeypatch) == [('model-b', 'endpoint-b', 'model-b', 'endpoint-b')] + + +def test_internal_one_shot_restore_keeps_user_deferred_pick(intent_turn, monkeypatch): + session, _ = intent_turn + restore = {'override': session['model_override'], 'model': 'model-a', 'provider': 'endpoint-a'} + # Use an inert custom runtime for the one-shot; its actual finally branch + # restores through the same helper used by the MoA one-shot producer. + server._apply_model_switch('selected', session, 'model-c --provider endpoint-c --session', + confirm_expensive_model=True, persist_override=False) + session['moa_one_shot_restore'] = restore + pending = [] + + def offline(*a, **kw): + assert server.handle_request(select('model-b', 'endpoint-b'))['result']['deferred'] + pending.append(session['pending_model_switch']) + return {'final_response': 'offline', 'completed': False, 'api_calls': 0} + + monkeypatch.setattr(session['agent'], 'run_conversation', offline) + session['running'] = True + server._run_prompt_submit('internal-one-shot', 'selected', session, 'offline input') + session['_run_thread'].join(5) + assert not session['_run_thread'].is_alive() + assert session['agent'].model == 'model-a' + assert session['pending_model_switch'] is pending[0] + assert run_turn(session, monkeypatch) == [('model-b', 'endpoint-b', 'model-b', 'endpoint-b')] + + +def test_publication_error_after_once_completion_keeps_old_intent(intent_turn, monkeypatch): + session, _ = intent_turn + session['running'] = True + assert server.handle_request(select('model-b', 'endpoint-b'))['result']['deferred'] + pending = session['pending_model_switch'] + session['running'] = False + collect = server._emit + + class FailingStream: + def write(self, text): + raise OSError(errno.EIO, 'offline once publication failed') + + failed_transport = StdioTransport(lambda: FailingStream(), threading.Lock()) + + def publish(name, sid, payload=None): + if name == 'session.info' and payload['model'] == 'model-c': + monkeypatch.setattr(server, '_emit', collect) + assert run_turn(session, monkeypatch, 'during-once-publication') == [ + ('model-c', 'endpoint-c', 'model-c', 'endpoint-c')] + assert session['agent'].model == 'model-a' + assert not session.get('_one_turn_model_runtime') + restored = server._session_info(session['agent'], session) + assert (restored['model'], restored['provider']) == ('model-a', 'endpoint-a') + return failed_transport.write(server._event_frame(name, sid, payload)) + return collect(name, sid, payload) + + monkeypatch.setattr(server, '_emit', publish) + response = server.handle_request(select('model-c', 'endpoint-c', once=True)) + monkeypatch.setattr(server, '_emit', collect) + assert response['error']['code'] == 5001 + assert session['pending_model_switch'] is pending + assert_route(session['agent'], 'http://a.invalid/v1', 'synthetic-key-a') + assert run_turn(session, monkeypatch) == [('model-b', 'endpoint-b', 'model-b', 'endpoint-b')] + + +@pytest.mark.parametrize('ordering', ['once_success', 'c_failure_first', 'd_failure_first']) +def test_publication_error_after_newer_once_completion_does_not_restore_old_intent(intent_turn, monkeypatch, ordering): + session, _ = intent_turn + transport = session['transport'] + session['running'] = True + assert server.handle_request(select('model-b', 'endpoint-b'))['result']['deferred'] + old_pending = session['pending_model_switch'] + session['running'] = False + entered, release = threading.Event(), threading.Event() + turn_entered, turn_release = threading.Event(), threading.Event() + turn_samples = [] + publisher = [] + newer_publisher = threading.get_ident() + collect = server._emit + monkeypatch.setattr(server.threading, 'Thread', REAL_THREAD) + + class InertWorker: + def __init__(self, block=False): + self._closed = False + self.block = block + + def close(self): + if self._closed: + return + self._closed = True + if self.block: + publisher.append(threading.get_ident()) + entered.set() + assert release.wait(5) + + class FailingStream: + def write(self, text): + raise OSError(errno.EIO, 'offline late publication failed') + + failed_transport = StdioTransport(lambda: FailingStream(), threading.Lock()) + initial_worker = InertWorker(block=True) + session['slash_worker'] = initial_worker + monkeypatch.setattr(server, '_SlashWorker', lambda *a, **kw: InertWorker()) + monkeypatch.setattr(server, '_restart_slash_worker', REAL_RESTART) + + def publish(name, sid, payload=None): + if name == 'session.info' and threading.get_ident() in publisher: + return failed_transport.write(server._event_frame(name, sid, payload)) + if (name == 'session.info' and payload['model'] == 'model-d' + and threading.get_ident() == newer_publisher and ordering != 'once_success'): + if ordering == 'c_failure_first': + release.set() + assert transport.done.wait(5) + return failed_transport.write(server._event_frame(name, sid, payload)) + return collect(name, sid, payload) + + monkeypatch.setattr(server, '_emit', publish) + with ThreadPoolExecutor(max_workers=1) as pool: + monkeypatch.setattr(server, '_pool', pool) + assert server.dispatch(select('model-c', 'endpoint-c'), transport) is None + try: + assert entered.wait(5) and initial_worker._closed + pin = session['model_override'] + if ordering == 'once_success': + server._mirror_slash_side_effects('selected', session, + '/model model-d --provider endpoint-c --once') + assert session['agent'].model == 'model-d' + assert run_turn(session, monkeypatch, 'newer-once') == [ + ('model-d', 'endpoint-c', 'model-d', 'endpoint-c')] + assert session['agent'].model == 'model-c' + assert session['model_override'] is pin + assert not session.get('_one_turn_model_runtime') + else: + output = server._mirror_slash_side_effects('selected', session, + '/model model-d --provider endpoint-c --session') + assert 'publication failed' in output + if ordering == 'd_failure_first': + # C still owns an unfinished suffix: its old B cannot + # become eligible merely because D released its claim. + server._apply_pending_model_switch('selected', session) + assert session['agent'].model == 'model-d' + value = session['agent'] + admitted_runtime = (value.model, value.provider, value.client) + + def ready(*a, **kw): + turn_entered.set() + assert turn_release.wait(5) + return None + + def offline(*a, **kw): + turn_samples.append((value.model, value.provider, value.client)) + return {'final_response': 'offline', 'completed': False, 'api_calls': 0} + + monkeypatch.setattr(server, '_wait_agent_for_prompt', ready) + monkeypatch.setattr(value, 'run_conversation', offline) + ack = server.dispatch({'id': 'send-d-between-failures', 'method': 'prompt.submit', + 'params': {'session_id': 'selected', 'text': 'offline input'}}, transport) + assert not ack.get('error'), ack + assert ack['result']['status'] == 'streaming' + assert turn_entered.wait(5) and session['running'] + release.set() + assert transport.done.wait(5) + assert old_pending['after_inflight_turn'] is session['inflight_turn'] + finally: + release.set() + turn_release.set() + thread = session.get('_run_thread') + while thread is not None: + thread.join(5) + assert not thread.is_alive() + successor = session.get('_run_thread') + if successor is thread: + break + thread = successor + assert transport.done.wait(5) + monkeypatch.setattr(server, '_emit', collect) + assert transport.frames[0]['error']['code'] == 5001 + if ordering == 'once_success': + assert not session.get('pending_model_switch') + assert run_turn(session, monkeypatch) == [('model-c', 'endpoint-c', 'model-c', 'endpoint-c')] + else: + assert session['pending_model_switch'] is old_pending + if ordering == 'd_failure_first': + assert turn_samples == [admitted_runtime] + assert_route(session['agent'], 'http://c.invalid/v1', 'synthetic-key-c') + monkeypatch.setattr(server, '_wait_agent_for_prompt', lambda *a, **kw: None) + assert run_turn(session, monkeypatch) == [('model-b', 'endpoint-b', 'model-b', 'endpoint-b')] + + +def test_config_adoption_keeps_pick_fenced_after_current_admission(intent_turn, routes, monkeypatch): + session, _ = intent_turn + session.pop('model_override') + routes.config['model'].update(default='model-c', provider='endpoint-c') + routes.write() + monkeypatch.setattr(server, '_sync_agent_model_with_config', REAL_CONFIG_SYNC) + session['running'] = True + with session['history_lock']: + server._start_inflight_turn(session, 'already admitted') + assert server.handle_request(select('model-b', 'endpoint-b'))['result']['deferred'] + pending = session['pending_model_switch'] + assert run_turn(session, monkeypatch) == [('model-c', 'endpoint-c', 'model-b', 'endpoint-b')] + assert session['pending_model_switch'] is pending + assert run_turn(session, monkeypatch) == [('model-b', 'endpoint-b', 'model-b', 'endpoint-b')] + + +@pytest.mark.parametrize('replacement', ['record', 'agent', 'transport']) +def test_delayed_idle_selection_refuses_changed_owner(intent_turn, monkeypatch, replacement): + session, events = intent_turn + value = session['agent'] + before = copy.deepcopy(session['model_override']) + entered, release = threading.Event(), threading.Event() + transport = session['transport'] + real_switch = ms.switch_model + monkeypatch.setattr(server.threading, 'Thread', REAL_THREAD) + + def slow(*a, **kw): + entered.set() + assert release.wait(5) + return real_switch(*a, **kw) + + monkeypatch.setattr(ms, 'switch_model', slow) + with ThreadPoolExecutor(max_workers=1) as pool: + monkeypatch.setattr(server, '_pool', pool) + assert server.dispatch(select('model-b', 'endpoint-b'), transport) is None + try: + assert entered.wait(5) + if replacement == 'record': + server._sessions['selected'] = {'agent': SimpleNamespace(model='replacement')} + elif replacement == 'agent': + session['agent'] = SimpleNamespace(model='replacement') + else: + session['transport'] = Transport() + finally: + release.set() + assert transport.done.wait(5) + assert transport.frames[0]['error']['code'] == 5001 + assert_route(value, 'http://a.invalid/v1', 'synthetic-key-a') + assert session['model_override'] == before + assert not session.get('_config_set_pending') + assert not events diff --git a/tests/tui_gateway/test_model_once_turn_runtime_owner.py b/tests/tui_gateway/test_model_once_turn_runtime_owner.py index db60bb745e61b..0d00d2856622e 100644 --- a/tests/tui_gateway/test_model_once_turn_runtime_owner.py +++ b/tests/tui_gateway/test_model_once_turn_runtime_owner.py @@ -104,10 +104,11 @@ def test_cancel_before_turn_does_not_consume_queued_once(live_turn, monkeypatch, @pytest.mark.parametrize('replacement_kind', ['agent', 'session']) -def test_retired_once_owner_does_not_restore_replacement(live_turn, monkeypatch, replacement_kind): +def test_lost_once_owner_retains_unproven_A_without_restoring_replacement(live_turn, monkeypatch, replacement_kind): session, _ = live_turn value = session['agent'] once(session) + lease = session['_one_turn_model_runtime'] replacement = _make_agent_openrouter() replacement.quiet_mode = True replacement._create_openai_client = lambda kwargs, **kw: SimpleNamespace(kwargs=kwargs.copy()) @@ -136,7 +137,10 @@ def offline(*a, **kw): server._run_prompt_submit('once-request', 'selected', session, 'offline input') assert (replacement.model, replacement.provider) == ('replacement-model', 'replacement-provider') assert_route(replacement, 'http://replacement.invalid/v1', 'synthetic-replacement-key') - assert not session.get('_one_turn_model_runtime') + assert session['_one_turn_model_runtime'] is lease + assert lease['restore_snapshot']['model'] == 'model-a' + if replacement_kind == 'agent': + assert lease['restore_failed'] and not lease['active'] assert not restore_side_effects @@ -206,6 +210,7 @@ def fail(*a, **kw): def test_stale_queued_once_owner_fails_before_dispatch(live_turn, monkeypatch): session, events = live_turn once(session) + lease = session['_one_turn_model_runtime'] replacement = _make_agent_openrouter() replacement.quiet_mode = True session['agent'] = replacement @@ -213,8 +218,8 @@ def test_stale_queued_once_owner_fails_before_dispatch(live_turn, monkeypatch): monkeypatch.setattr(replacement, 'run_conversation', lambda *a, **kw: dispatched.append(True)) server._run_prompt_submit('stale-once-request', 'selected', session, 'offline input') assert not dispatched - assert not session.get('_one_turn_model_runtime') - assert not session.get('one_turn_model_restore') + assert session['_one_turn_model_runtime'] is lease and lease['restore_failed'] + assert session['one_turn_model_restore'] is lease['restore_snapshot'] assert any(name == 'message.complete' and payload.get('status') == 'error' for name, _, payload in events) diff --git a/tests/tui_gateway/test_once_restore_Stop_only.py b/tests/tui_gateway/test_once_restore_Stop_only.py new file mode 100644 index 0000000000000..a5056d981d4ae --- /dev/null +++ b/tests/tui_gateway/test_once_restore_Stop_only.py @@ -0,0 +1,138 @@ +"""Actual Stop once settlement and later-owner controls through real gateway paths.""" +from unittest.mock import Mock + +import pytest + +from tests.tui_gateway.test_once_restore_owner_races import ( # noqa: F401 + agent, gateway, intent_turn, live_turn, routes, turn_env, select, submit, + initialize_inert_agent_interrupt_owner, REAL_EMIT, +) +from tui_gateway import server + + +def stop(): + response = server._methods['session.interrupt']('stop', {'session_id': 'selected'}) + assert not response.get('error'), response + + +def pick_once(session): + assert not server.handle_request(select('model-b', 'endpoint-b', once=True)).get('error') + return session['_one_turn_model_runtime'] + + +@pytest.mark.parametrize('admitted', [False, True]) +def test_actual_Stop_only_restores_A_and_publishes_owned_idle_metadata(intent_turn, monkeypatch, admitted): + session, _ = intent_turn + value = session['agent']; sink = session['transport'] + lease = pick_once(session) + calls = []; lock_observations = [] + real_restore = server._restore_agent_model_runtime + def restore(*a): + calls.append(a[1]['model']); return real_restore(*a) + monkeypatch.setattr(server, '_restore_agent_model_runtime', restore) + real_write = sink.write + def write(frame): + p = frame.get('params', {}) + if p.get('type') == 'session.info' and not session['running'] and p['payload']['model'] == 'model-a': + acquired = session['history_lock'].acquire(blocking=False) + lock_observations.append(acquired) + if acquired: session['history_lock'].release() + return real_write(frame) + monkeypatch.setattr(sink, 'write', write) + monkeypatch.setattr(server, '_emit', REAL_EMIT) + monkeypatch.setattr(server, '_wait_agent_for_prompt', lambda *a: None) + def offline(*a, persist_user_event_id=None, **kw): + stop() + return {'final_response': 'offline', 'interrupted': True, 'completed': False, 'api_calls': 0} + monkeypatch.setattr(value, 'run_conversation', offline) + if admitted: + response = submit('first'); assert not response.get('error'), response + else: + session['running'] = True + server._run_prompt_submit('old', 'selected', session, 'first') + assert calls == ['model-a'] and value.model == 'model-a' + assert session.get('_one_turn_model_runtime') is None and not lease.get('restore_failed') + assert not session['running'] and session.get('inflight_turn') is None + assert session['_turn_cancel_requested'] and lock_observations == [True] + if admitted: + assert session['_turn_outcomes'].turns[-1]['state'] == 'interrupted' + + +@pytest.mark.parametrize('failure', ['restore', 'publication']) +def test_actual_Stop_restore_failure_keeps_A_custody_and_refuses_Send(intent_turn, monkeypatch, failure): + session, events = intent_turn; value = session['agent']; lease = pick_once(session) + snapshot = lease['restore_snapshot']; calls = [] + def offline(*a, **kw): + calls.append(value.model); stop() + return {'final_response': 'offline', 'interrupted': True, 'completed': False, 'api_calls': 0} + monkeypatch.setattr(value, 'run_conversation', offline) + def refuse(*a): raise RuntimeError('offline restore/publication refusal') + monkeypatch.setattr(server, '_restore_agent_model_runtime' if failure == 'restore' + else '_persist_live_session_runtime', refuse) + session['running'] = True + server._run_prompt_submit('old', 'selected', session, 'first') + assert session['_one_turn_model_runtime'] is lease and lease['restore_snapshot'] is snapshot + assert snapshot['model'] == 'model-a' and lease['restore_failed'] and not lease['active'] + assert any(e == 'error' and p.get('error_surface', {}).get('code') == + 'one_turn_model_restore_failed' for e, _, p in events) + accept = Mock(side_effect=AssertionError('failed restoration accepted another input')) + monkeypatch.setattr(server, '_accept_tui_context_input', accept) + response = submit('refused later input') + assert response['error']['data']['error_surface']['code'] == 'one_turn_model_restore_failed' + assert response['error']['data']['durable_input_accepted'] is False + accept.assert_not_called(); assert calls == ['model-b'] + + +@pytest.mark.parametrize('once', [False, True]) +def test_actual_Stop_newer_explicit_C_remains_authoritative(intent_turn, monkeypatch, once): + session, _ = intent_turn; value = session['agent']; old_lease = pick_once(session) + seen = {}; restore = Mock(side_effect=AssertionError('old A restored over explicit C')) + monkeypatch.setattr(server, '_restore_agent_model_runtime', restore) + def offline(*a, **kw): + stop() + response = server.handle_request(select('model-c', 'endpoint-c', once=once)) + assert not response.get('error'), response + seen['lease'] = session.get('_one_turn_model_runtime') + return {'final_response': 'offline', 'interrupted': True, 'completed': False, 'api_calls': 0} + monkeypatch.setattr(value, 'run_conversation', offline) + session['running'] = True + server._run_prompt_submit('old', 'selected', session, 'first') + restore.assert_not_called(); assert value.model == 'model-c' and not session['running'] + assert session.get('_one_turn_model_runtime') is seen['lease'] + if once: + assert seen['lease'] is not old_lease and not seen['lease'].get('restore_failed') + assert not seen['lease']['active'] and session.get('one_turn_model_restore') is seen['lease']['restore_snapshot'] + else: assert seen['lease'] is None + + +@pytest.mark.parametrize('queued', [False, True]) +def test_actual_Stop_new_Send_then_Stop_does_not_borrow_cleared_successor(intent_turn, monkeypatch, queued): + session, _ = intent_turn; value = session['agent']; lease = pick_once(session) + old_driver = server._run_prompt_submit; seen = {}; sentinel = object(); markers = [] + monkeypatch.setattr(server, '_retire_turn_marker', lambda *a: markers.append('retire')) + restore = Mock(side_effect=AssertionError('old A restored after a newer stopped admission')) + monkeypatch.setattr(server, '_restore_agent_model_runtime', restore) + monkeypatch.setattr(server, '_wait_agent_for_prompt', lambda *a: None) + monkeypatch.setattr(server, '_start_agent_build', lambda *a: None) + def claim(*a, **kw): + seen['nonce'] = session['_turn_outcomes'].turns[-1]['accepted_turn']['request_id'] + value.interim_assistant_callback = sentinel + return True + def offline(*a, **kw): + stop() + monkeypatch.setattr(server, '_run_prompt_submit', claim) + response = submit('newer admission'); assert not response.get('error'), response + stop(); seen['markers'] = len(markers) + return {'final_response': 'offline', 'interrupted': True, 'completed': False, 'api_calls': 0} + monkeypatch.setattr(value, 'run_conversation', offline) + if queued: + with session['history_lock']: + server._enqueue_prompt(session, 'first', session['transport']) + assert server._drain_queued_prompt('queued', 'selected', session) + else: + session['running'] = True + old_driver('old', 'selected', session, 'first') + restore.assert_not_called(); assert value.model == 'model-b' + assert session['_one_turn_model_runtime'] is lease and lease['restore_failed'] + assert session['_turn_outcomes'].turns[-1]['accepted_turn']['request_id'] == seen['nonce'] + assert value.interim_assistant_callback is sentinel and len(markers) == seen['markers'] diff --git a/tests/tui_gateway/test_once_restore_admission.py b/tests/tui_gateway/test_once_restore_admission.py new file mode 100644 index 0000000000000..5cb9f4f0b0e3f --- /dev/null +++ b/tests/tui_gateway/test_once_restore_admission.py @@ -0,0 +1,186 @@ +"""Failed once custody at real model selection and prompt admission owners.""" +import copy +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest + +from tests.tui_gateway.test_model_intent_admission_order import ( # noqa: F401 + agent, gateway, intent_turn, live_turn, routes, select, turn_env, run_turn, +) +from tui_gateway import server + + +@pytest.fixture +def failed_once(intent_turn): + session, events = intent_turn + assert not server.handle_request(select('model-b', 'endpoint-b', once=True)).get('error') + snapshot, lease = server._consume_one_turn_model_runtime(session, session['agent']) + lease['restore_failed'] = True + lease['active'] = False + return session, events, lease, snapshot + + +def submit(text='new input', **params): + return server._methods['prompt.submit']('new-rid', {'session_id': 'selected', 'text': text, **params}) + + +def refused(response, *, accepted=False): + assert response['error']['data']['error_surface'] == { + 'layer': 'runtime', 'code': 'one_turn_model_restore_failed', 'retryable': False} + assert response['error']['data']['execution_started'] is False + assert response['error']['data']['durable_input_accepted'] is accepted + + +@pytest.mark.parametrize('truncate', [False, True]) +def test_fresh_failure_refuses_before_input_history_and_attachments(failed_once, monkeypatch, truncate): + session, _, lease, snapshot = failed_once + session['attached_images'] = ['inert.png'] + history = copy.deepcopy(session['history']) + version = session.get('history_version') + accept = Mock(side_effect=AssertionError('input acceptance forbidden')) + monkeypatch.setattr(server, '_accept_tui_context_input', accept) + response = submit(**({'truncate_before_user_ordinal': 0, 'confirm_truncate': True} if truncate else {})) + refused(response) + accept.assert_not_called() + assert session['history'] == history and session.get('history_version') == version + assert session['attached_images'] == ['inert.png'] and not session['running'] + assert session['_one_turn_model_runtime'] is lease and lease['restore_snapshot'] is snapshot + + +def test_dispatch_backstop_retains_its_own_failed_input(failed_once, monkeypatch): + session, events, lease, snapshot = failed_once + history = copy.deepcopy(session['history']) + session['running'] = True + provider = Mock(side_effect=AssertionError('provider forbidden')) + marker = Mock(side_effect=AssertionError('crash marker forbidden')) + monkeypatch.setattr(session['agent'], 'run_conversation', provider) + monkeypatch.setattr(server, 'record_turn_start', marker) + assert server._run_prompt_submit('queued', 'selected', session, 'owned queued input', + queued_prompt_generation=int(session.get('_queued_prompt_generation', 0))) is False + provider.assert_not_called();marker.assert_not_called() + assert session['history'] == history and not session['running'] + assert session['inflight_turn']['user'] == 'owned queued input' + assert session['inflight_turn']['status'] == 'error' + terminal = [p for e, _, p in events if e == 'message.complete' and p.get('status') == 'error'] + assert terminal[-1]['error_surface']['code'] == 'one_turn_model_restore_failed' + assert lease['restore_snapshot'] is snapshot and session['_one_turn_model_runtime'] is lease + + +def test_failure_racing_durable_acceptance_keeps_receipt(failed_once, monkeypatch): + session, _, lease, _ = failed_once + lease['restore_failed'] = False + receipt = SimpleNamespace(event_id='recorded-input') + accepted = [] + def accept(*args, **kwargs): + accepted.append(receipt) + lease['restore_failed'] = True + return receipt + monkeypatch.setattr(server, '_accept_tui_context_input', accept) + response = submit() + refused(response, accepted=True) + assert accepted == [receipt] + assert response['error']['data']['input_event_id'] == receipt.event_id + assert session['_one_turn_model_runtime'] is lease and not session['running'] + + +def test_successful_idle_C_retires_failed_lease(failed_once, monkeypatch): + session, _, _, _ = failed_once + assert not server.handle_request(select('model-c', 'endpoint-c')).get('error') + assert not session.get('_one_turn_model_runtime') + assert run_turn(session, monkeypatch)[0][:2] == ('model-c', 'endpoint-c') + + +def test_new_once_C_preserves_original_A(failed_once, monkeypatch): + session, _, old, snapshot = failed_once + assert not server.handle_request(select('model-c', 'endpoint-c', once=True)).get('error') + assert session['_one_turn_model_runtime'] is not old + assert session['one_turn_model_restore'] is snapshot + assert session['_one_turn_model_runtime']['restore_snapshot'] is snapshot + assert run_turn(session, monkeypatch)[0][:2] == ('model-c', 'endpoint-c') + assert session['agent'].model == 'model-a' + + +def test_failed_publication_keeps_failed_custody(failed_once, monkeypatch): + session, _, lease, snapshot = failed_once + def fail(*args): + refused(submit()) + raise RuntimeError('inert publication failure') + monkeypatch.setattr(server, '_restart_slash_worker', fail) + response = server.handle_request(select('model-c', 'endpoint-c')) + assert response.get('error') + assert session['_one_turn_model_runtime'] is lease and lease['restore_snapshot'] is snapshot + assert lease['restore_failed'] and not lease.get('superseding_intent') + + +def test_old_publication_failure_cannot_resurrect_over_newer_choice(failed_once, monkeypatch): + session, _, lease, _ = failed_once + def publish(*args): + if session['agent'].model == 'model-c': + assert not server.handle_request(select('newer-model', 'endpoint-b')).get('error') + raise RuntimeError('older publication failure') + monkeypatch.setattr(server, '_restart_slash_worker', publish) + assert server.handle_request(select('model-c', 'endpoint-c')).get('error') + assert session['agent'].model == 'newer-model' and not session.get('_one_turn_model_runtime') + assert lease['restore_failed'] + + +@pytest.mark.parametrize('kind', ['invalid', 'unconfirmed']) +def test_unsuccessful_intent_keeps_failed_lease(failed_once, monkeypatch, kind): + session, _, lease, _ = failed_once + if kind == 'invalid': + response = server.handle_request(select('model-c', 'nonexistent-provider')) + assert response.get('error') + else: + monkeypatch.setattr('hermes_cli.model_selection_guards.combined_selection_warning', + lambda *a, **k: SimpleNamespace(message='inert consent required')) + response = server.handle_request(select('model-c', 'endpoint-c', confirmed=False)) + assert response['result']['confirm_required'] + assert session['_one_turn_model_runtime'] is lease and lease['restore_failed'] + + +def test_internal_config_adoption_cannot_supersede_failed_lease(failed_once): + session, _, lease, _ = failed_once + with pytest.raises(ValueError, match='not restored'): + server._apply_model_switch('selected', session, 'model-c --provider endpoint-c', + confirm_expensive_model=True, pin_session_override=False, persist_override=False) + assert session['_one_turn_model_runtime'] is lease and session['agent'].model == 'model-b' + + +def test_eligible_deferred_C_settles_before_input_admission(failed_once, monkeypatch): + session, _, _, _ = failed_once + session['pending_model_switch'] = {'raw': 'model-c --provider endpoint-c --session', + 'confirm_expensive_model': True} + seen = [] + def accept(*args, **kwargs): + seen.append((session['agent'].model, session.get('_one_turn_model_runtime'))) + raise RuntimeError('stop after admission observation') + monkeypatch.setattr(server, '_accept_tui_context_input', accept) + assert submit().get('error') + assert seen == [('model-c', None)] and not session.get('pending_model_switch') + + +def test_Stop_does_not_clear_failed_lease(failed_once): + session, _, lease, _ = failed_once + response = server._methods['session.interrupt']('stop', {'session_id': 'selected'}) + assert not response.get('error') + assert session['_one_turn_model_runtime'] is lease and lease['restore_failed'] + refused(submit()) + + +def test_unexpected_same_session_agent_replacement_cannot_escape_failure(failed_once): + session, _, lease, _ = failed_once + session['agent'] = SimpleNamespace(model='replacement', provider='inert') + refused(submit()) + assert session['_one_turn_model_runtime'] is lease + + +def test_stale_generation_does_not_poison_replacement(failed_once): + session, _, lease, _ = failed_once + replacement = {'agent': SimpleNamespace(), 'history_lock': session['history_lock']} + server._sessions['selected'] = replacement + with session['history_lock']: + assert server._one_turn_model_restore_error('selected', session) is None + assert server._one_turn_model_restore_error('selected', replacement) is None + assert session['_one_turn_model_runtime'] is lease + assert '_one_turn_model_runtime' not in replacement diff --git a/tests/tui_gateway/test_once_restore_owner_races.py b/tests/tui_gateway/test_once_restore_owner_races.py new file mode 100644 index 0000000000000..aa30571a8a369 --- /dev/null +++ b/tests/tui_gateway/test_once_restore_owner_races.py @@ -0,0 +1,387 @@ +"""Reviewed CPO owner schedules, with actual admission and inert provider boundaries.""" +import copy +import threading +from unittest.mock import Mock + +import pytest + +from tests.tui_gateway.test_once_restore_admission import ( # noqa: F401 + failed_once, intent_turn, agent, gateway, live_turn, routes, turn_env, select, submit, +) +from hermes_cli import model_switch as ms +from tui_gateway import server + +REAL_EMIT = server._emit + + +@pytest.fixture(autouse=True) +def initialize_inert_agent_interrupt_owner(intent_turn): + # The route fixture uses AIAgent.__new__. Supply its normal pre-execution + # thread state so the real Stop implementation can run without tool I/O. + value = intent_turn[0]['agent'] + value._execution_thread_id = None + value._hard_interrupt_requested = threading.Event() + value._active_children_lock = threading.Lock() + value._active_children = [] + + +def test_unexpected_agent_replacement_retains_consumed_A(intent_turn, monkeypatch): + session, _ = intent_turn + value = session['agent'] + assert not server.handle_request(select('model-b', 'endpoint-b', once=True)).get('error') + lease = session['_one_turn_model_runtime'] + replacement = copy.copy(value) + replacement.model = 'replacement' + def offline(*a, **kw): + session['agent'] = replacement + session['running'] = False + return {'final_response': 'offline', 'completed': False, 'api_calls': 0} + monkeypatch.setattr(value, 'run_conversation', offline) + session['running'] = True + server._run_prompt_submit('old', 'selected', session, 'first') + assert session.get('_one_turn_model_runtime') is lease + assert lease['restore_snapshot']['model'] == 'model-a' and lease['restore_failed'] + response = submit() + assert response['error']['data']['durable_input_accepted'] is False + assert replacement.model == 'replacement' + + +def test_Stop_then_new_Send_during_restore_publication_keeps_successor(intent_turn, monkeypatch): + session, events = intent_turn + value = session['agent'] + assert not server.handle_request(select('model-b', 'endpoint-b', once=True)).get('error') + lease = session['_one_turn_model_runtime'] + monkeypatch.setattr(value, 'run_conversation', lambda *a, **kw: { + 'final_response': 'offline', 'completed': False, 'api_calls': 0}) + old_driver = server._run_prompt_submit + successor = {} + sentinel = object() + markers = [] + monkeypatch.setattr(server, '_retire_turn_marker', lambda *a: markers.append('retire')) + def claim_only(*a, **kw): + successor['turn'] = session['inflight_turn'] + successor['nonce'] = session['_turn_outcomes'].turns[-1]['accepted_turn']['request_id'] + return True + def restart(*a): + successor['entered'] = True + try: + successor['stop'] = server._methods['session.interrupt']('stop', {'session_id': 'selected'}) + monkeypatch.setattr(server, '_run_prompt_submit', claim_only) + monkeypatch.setattr(server, '_start_agent_build', lambda *a: None) + monkeypatch.setattr(server, '_wait_agent_for_prompt', lambda *a: None) + successor['response'] = submit('successor') + value.interim_assistant_callback = sentinel + successor['markers'] = len(markers) + except Exception as exc: + successor['exception'] = repr(exc) + raise + monkeypatch.setattr(server, '_restart_slash_worker', restart) + session['running'] = True + old_driver('old', 'selected', session, 'first') + assert successor.get('entered') and 'exception' not in successor, successor + assert not successor['stop'].get('error'), successor + assert 'turn' in successor, successor + assert not successor['response'].get('error'), successor + assert session['running'] and session['inflight_turn'] is successor['turn'] + assert value.interim_assistant_callback is sentinel + assert len(markers) == successor['markers'] + assert session['_one_turn_model_runtime'] is lease and lease['restore_failed'] + assert not any(e == 'error' and p.get('error_surface', {}).get('code') == + 'one_turn_model_restore_failed' for e, _, p in events) + + +@pytest.mark.parametrize('changed', ['session', 'agent', 'transport', 'Stop']) +def test_pending_resolution_rechecks_before_acceptance(failed_once, monkeypatch, changed): + session, events, lease, _ = failed_once + session['pending_model_switch'] = {'raw': 'model-c --provider endpoint-c --session', + 'confirm_expensive_model': True} + original = copy.deepcopy(session['history']) + session['attached_images'] = ['owned-image'] + real = ms.switch_model + replacement = copy.copy(session) + def resolve(*a, **kw): + result = real(*a, **kw) + if changed == 'session': server._sessions['selected'] = replacement + elif changed == 'agent': session['agent'] = copy.copy(session['agent']) + elif changed == 'transport': session['transport'] = object() + else: + assert not server._methods['session.interrupt']('stop', {'session_id': 'selected'}).get('error') + events.clear() + return result + monkeypatch.setattr(ms, 'switch_model', resolve) + accept = Mock(side_effect=AssertionError('stale acceptance')) + monkeypatch.setattr(server, '_accept_tui_context_input', accept) + events.clear() + response = submit('stale input') + accept.assert_not_called() + assert response.get('error') + assert response['error']['data']['durable_input_accepted'] is False + assert session['history'] == original and session['attached_images'] == ['owned-image'] + assert not events + assert session['_one_turn_model_runtime'] is lease + + +@pytest.mark.parametrize('changed', ['session', 'transport']) +def test_explicit_suffix_owner_loss_is_refusal_without_successor_delivery(failed_once, monkeypatch, changed): + session, events, lease, _ = failed_once + def restart(*a): + if changed == 'session': server._sessions['selected'] = copy.copy(session) + else: session['transport'] = object() + events.clear() + monkeypatch.setattr(server, '_restart_slash_worker', restart) + response = server.handle_request(select('model-c', 'endpoint-c')) + assert response.get('error') and 'owner changed' in response['error']['message'] + assert not events + assert session['_one_turn_model_runtime'] is lease and lease['restore_failed'] + + +def test_explicit_slash_commit_cannot_bypass_changed_owner(failed_once, monkeypatch): + session, _, lease, _ = failed_once + real = ms.switch_model + def resolve(*a, **kw): + result = real(*a, **kw) + server._sessions['selected'] = copy.copy(session) + return result + monkeypatch.setattr(ms, 'switch_model', resolve) + with pytest.raises(ValueError, match='owner changed'): + server._apply_model_switch('selected', session, 'model-c --provider endpoint-c --session', + confirm_expensive_model=True, explicit_model_intent=True, defer_if_running=False) + assert session['agent'].model == 'model-b' and session['_one_turn_model_runtime'] is lease + + +@pytest.mark.parametrize('stop_after_first', [False, True]) +@pytest.mark.parametrize('host_policy', [False, True]) +def test_each_queued_image_refusal_keeps_receipt_and_Stop_cut(failed_once, monkeypatch, tmp_path, stop_after_first, host_policy): + session, events, lease, _ = failed_once + history = copy.deepcopy(session['history']) + # Canonical governed text receipts are accepted before the inert image + # envelopes are constructed. Fresh governed image RPCs remain refused by + # the existing context owner; this does not certify that separate route. + session['agent'].context_rebase_enabled = True + paths = [] + receipts = [] + for n in range(2): + p = tmp_path / f'image{n}.png';p.write_bytes(b'inert-image');paths.append(str(p)) + receipt = server._accept_tui_context_input(session, f'owned input {n}') + receipts.append(receipt.event_id) + with session['history_lock']: + server._enqueue_prompt(session, f'owned input {n}', session['transport'], + image_paths=[str(p)], context_input_event_id=receipt.event_id) + session['attached_images'] = ['unclaimed-attachment'] + provider = Mock(side_effect=AssertionError('refused provider dispatch')) + monkeypatch.setattr(session['agent'], 'run_conversation', provider) + host = Mock(side_effect=AssertionError('refused host dispatch')) + monkeypatch.setattr(server, '_submit_prompt_to_compute_host', host) + monkeypatch.setattr(server, '_session_uses_compute_host', lambda *a, **kw: host_policy) + def capture(event, sid, payload=None): + events.append((event, sid, payload)) + REAL_EMIT(event, sid, payload) + if stop_after_first and event == 'message.complete': + assert not server._methods['session.interrupt']('stop', {'session_id': 'selected'}).get('error') + monkeypatch.setattr(server, '_emit', capture) + events.clear() + assert server._drain_queued_prompt('drain', 'selected', session) + terminal = [p for e, _, p in events if e == 'message.complete'] + assert [p.get('input_event_id') for p in terminal] == receipts[:1 if stop_after_first else 2] + assert all(p['durable_input_accepted'] and p['execution_started'] is False + and p['error_surface']['code'] == 'one_turn_model_restore_failed' for p in terminal) + provider.assert_not_called();host.assert_not_called() + db = session['agent']._session_db + assert [db.read_context_input(session['session_key'], source='tui', event_id=e).content + for e in receipts] == ['owned input 0', 'owned input 1'] + wire = [f['params']['payload'] for f in session['transport'].frames + if f.get('params', {}).get('type') == 'message.complete'] + assert [p['input_event_id'] for p in wire] == receipts[:1 if stop_after_first else 2] + assert len({p['accepted_turn']['request_id'] for p in wire}) == len(wire) + assert all(t['state'] == 'error' for t in session['_turn_outcomes'].turns) + assert all(__import__('pathlib').Path(p).read_bytes() == b'inert-image' for p in paths) + assert session['history'] == history and session['attached_images'] == ['unclaimed-attachment'] + assert session['_one_turn_model_runtime'] is lease and not session['running'] + + +def test_busy_Send_other_transport_settles_original_then_drains(intent_turn, monkeypatch): + from tests.tui_gateway.test_config_dispatch_responsiveness import Transport + session, events = intent_turn + old_sink = session['transport']; new_sink = Transport() + value = session['agent']; value.context_rebase_enabled = True + assert not server.handle_request(select('model-b', 'endpoint-b', once=True)).get('error') + calls = [] + queued = {} + monkeypatch.setattr(server, '_wait_agent_for_prompt', lambda *a: None) + monkeypatch.setattr(server, '_emit', REAL_EMIT) + def offline(*a, persist_user_event_id=None, **kw): + assert persist_user_event_id is not None + accepted = value._session_db.read_context_input(session['session_key'], + source='tui', event_id=persist_user_event_id) + assert accepted.content in ('first input', 'second input') + calls.append(value.model) + if len(calls) == 1: + queued['ack'] = server.dispatch({'id': 'busy-send', 'method': 'prompt.submit', + 'params': {'session_id': 'selected', 'text': 'second input'}}, new_sink) + return {'final_response': calls[-1], 'completed': False, 'api_calls': 0} + monkeypatch.setattr(value, 'run_conversation', offline) + first = server.dispatch({'id': 'first-send', 'method': 'prompt.submit', + 'params': {'session_id': 'selected', 'text': 'first input'}}, old_sink) + assert not first.get('error') and not queued['ack'].get('error') + assert queued['ack']['result']['status'] == 'queued' + assert calls == ['model-b', 'model-a'] + assert value.model == 'model-a' and not session.get('_one_turn_model_runtime') + assert not session['running'] and not session.get('inflight_turn') and not session.get('queued_prompt') + def finals(sink): + return [f['params']['payload'] for f in sink.frames + if f.get('params', {}).get('type') == 'message.complete'] + assert [p['text'] for p in finals(old_sink)] == ['model-b'] + assert [p['text'] for p in finals(new_sink)] == ['model-a'] + + +@pytest.mark.parametrize('successor', [False, True]) +def test_Stop_before_queued_refusal_finishes_only_captured_nonce(failed_once, monkeypatch, successor): + session, events, lease, _ = failed_once + session['agent'].context_rebase_enabled = True + receipt = server._accept_tui_context_input(session, 'accepted queued input') + with session['history_lock']: + server._enqueue_prompt(session, 'accepted queued input', session['transport'], + context_input_event_id=receipt.event_id) + real_terminal = server._emit_terminal_turn_error + captured = {}; sentinel = object() + def claim(*a, **kw): + captured['successor_turn'] = session['inflight_turn'] + session['agent'].interim_assistant_callback = sentinel + return True + def before_helper(*a, **kw): + window = session['_turn_outcomes'] + captured['nonce'] = window.turns[-1]['accepted_turn']['request_id'] + captured['stop'] = server._methods['session.interrupt']('stop', {'session_id': 'selected'}) + if successor: + captured['pick'] = server.handle_request(select('model-c', 'endpoint-c')) + monkeypatch.setattr(server, '_wait_agent_for_prompt', lambda *a: None) + monkeypatch.setattr(server, '_start_agent_build', lambda *a: None) + monkeypatch.setattr(server, '_run_prompt_submit', claim) + captured['send'] = submit('successor') + captured['events_before'] = len(events) + captured['settled'] = real_terminal(*a, **kw) + return captured['settled'] + monkeypatch.setattr(server, '_emit_terminal_turn_error', before_helper) + assert server._drain_queued_prompt('drain', 'selected', session) + assert not captured['stop'].get('error') and captured['settled'] is False + assert session['_turn_outcomes'].find(captured['nonce'])['state'] == 'interrupted' + assert len(events) == captured['events_before'] + if successor: + assert not captured['pick'].get('error') and not captured['send'].get('error'), captured + assert session['running'] and session['inflight_turn'] is captured['successor_turn'] + assert session['agent'].interim_assistant_callback is sentinel + else: + assert not session['running'] and session['_one_turn_model_runtime'] is lease + + +def test_settled_metadata_transport_callback_can_acquire_admission_lock(intent_turn, monkeypatch): + session, _ = intent_turn + observed = [] + value = session['agent'] + monkeypatch.setattr(value, 'run_conversation', lambda *a, **kw: { + 'final_response': 'offline', 'completed': False, 'api_calls': 0}) + sink = session['transport']; real_write = sink.write + def write(frame): + if frame.get('params', {}).get('type') == 'session.info' and not session['running']: + acquired = session['history_lock'].acquire(blocking=False) + observed.append(acquired) + if acquired: + session['history_lock'].release() + return real_write(frame) + monkeypatch.setattr(sink, 'write', write) + monkeypatch.setattr(server, '_emit', REAL_EMIT) + session['running'] = True + server._run_prompt_submit('old', 'selected', session, 'first') + assert observed == [True] + + +@pytest.mark.parametrize('changed', ['agent', 'transport']) +def test_pending_pop_owner_is_carried_into_switch_commit(failed_once, monkeypatch, changed): + session, events, lease, _ = failed_once + session['pending_model_switch'] = {'raw': 'model-c --provider endpoint-c --session', + 'confirm_expensive_model': True} + original_agent = session['agent']; replacement = copy.copy(original_agent) + real_apply = server._apply_model_switch + def before_entry(*a, **kw): + if changed == 'agent': session['agent'] = replacement + else: session['transport'] = object() + events.clear() + return real_apply(*a, **kw) + monkeypatch.setattr(server, '_apply_model_switch', before_entry) + accept = Mock(side_effect=AssertionError('stale acceptance')) + monkeypatch.setattr(server, '_accept_tui_context_input', accept) + response = submit('old input') + accept.assert_not_called() + assert response.get('error') and response['error']['data']['durable_input_accepted'] is False + assert replacement.model == original_agent.model == 'model-b' + assert session['_one_turn_model_runtime'] is lease and lease['restore_failed'] + assert not events + + +def test_pending_resolution_error_cannot_deliver_after_new_Send(failed_once, monkeypatch): + session, events, _, _ = failed_once + session['agent'].context_rebase_enabled = True + session['pending_model_switch'] = {'raw': 'model-c --provider endpoint-c --session', + 'confirm_expensive_model': True} + real_resolve = ms.switch_model; accepted = []; successor = {}; sentinel = object() + real_accept = server._accept_tui_context_input + def accept(*a, **kw): + accepted.append(a[1]);return real_accept(*a, **kw) + monkeypatch.setattr(server, '_accept_tui_context_input', accept) + def claim(*a, **kw): + successor['turn'] = session['inflight_turn'] + successor['nonce'] = session['_turn_outcomes'].turns[-1]['accepted_turn']['request_id'] + session['agent'].interim_assistant_callback = sentinel + return True + def resolve(*a, **kw): + monkeypatch.setattr(ms, 'switch_model', real_resolve) + successor['pick'] = server.handle_request(select('model-c', 'endpoint-c')) + monkeypatch.setattr(server, '_wait_agent_for_prompt', lambda *a: None) + monkeypatch.setattr(server, '_start_agent_build', lambda *a: None) + monkeypatch.setattr(server, '_run_prompt_submit', claim) + successor['send'] = submit('new owner') + events.clear() + raise RuntimeError('old pending resolution failed') + monkeypatch.setattr(ms, 'switch_model', resolve) + response = submit('old pending input') + assert not successor['pick'].get('error') and not successor['send'].get('error'), successor + assert response.get('error') and response['error']['data']['durable_input_accepted'] is False + assert accepted == ['new owner'] and not events + assert session['running'] and session['inflight_turn'] is successor['turn'] + assert session['_turn_outcomes'].turns[-1]['accepted_turn']['request_id'] == successor['nonce'] + assert session['agent'].interim_assistant_callback is sentinel + + +def test_Stop_new_Send_at_queue_policy_boundary_keeps_successor(failed_once, monkeypatch): + session, events, _, _ = failed_once + session['agent'].context_rebase_enabled = True + receipt = server._accept_tui_context_input(session, 'old queued input') + with session['history_lock']: + server._enqueue_prompt(session, 'old queued input', session['transport'], + context_input_event_id=receipt.event_id) + entered = []; successor = {}; sentinel = object(); markers = [] + monkeypatch.setattr(server, '_retire_turn_marker', lambda *a: markers.append('retire')) + def claim(*a, **kw): + successor['turn'] = session['inflight_turn'] + successor['nonce'] = session['_turn_outcomes'].turns[-1]['accepted_turn']['request_id'] + session['agent'].interim_assistant_callback = sentinel + return True + def policy(*a, **kw): + if not entered: + entered.append(True) + successor['stop'] = server._methods['session.interrupt']('stop', {'session_id': 'selected'}) + successor['pick'] = server.handle_request(select('model-c', 'endpoint-c')) + monkeypatch.setattr(server, '_wait_agent_for_prompt', lambda *a: None) + monkeypatch.setattr(server, '_start_agent_build', lambda *a: None) + monkeypatch.setattr(server, '_run_prompt_submit', claim) + successor['send'] = submit('new owner') + successor['markers'] = len(markers) + events.clear() + return False + monkeypatch.setattr(server, '_session_uses_compute_host', policy) + assert server._drain_queued_prompt('drain', 'selected', session) + assert all(not successor[k].get('error') for k in ('stop', 'pick', 'send')), successor + assert session['running'] and session['inflight_turn'] is successor['turn'] + assert session['_turn_outcomes'].turns[-1]['accepted_turn']['request_id'] == successor['nonce'] + assert session['agent'].interim_assistant_callback is sentinel + assert len(markers) == successor['markers'] and not events diff --git a/tests/tui_gateway/test_prompt_recovery_contract.py b/tests/tui_gateway/test_prompt_recovery_contract.py index 317940e1c9d6f..3defbaa285d38 100644 --- a/tests/tui_gateway/test_prompt_recovery_contract.py +++ b/tests/tui_gateway/test_prompt_recovery_contract.py @@ -74,11 +74,20 @@ def test_patient_first_prompt_wait_outlives_old_30_second_cliff(monkeypatch): _agent_build_thread=_BuildThread(), ) events = [] - times = iter((0.0, 31.0)) - monkeypatch.setattr(server.time, "monotonic", lambda: next(times)) + elapsed = [0.0] + original_wait = ready.wait + + def advance_slice(timeout=None): + elapsed[0] += 31.0 + return original_wait(timeout=timeout) + + monkeypatch.setattr(ready, "wait", advance_slice) + monkeypatch.setattr(server.time, "monotonic", lambda: elapsed[0]) monkeypatch.setattr(server, "_emit", lambda event, sid, payload=None: events.append((event, payload))) assert server._wait_agent_for_prompt(session, "rid", "sid") is None + assert ready.calls == 2 + assert elapsed[0] > 30.0 assert [event for event, _ in events] == [ "notification.show", "notification.clear", diff --git a/tests/tui_gateway/test_resume_history_admission.py b/tests/tui_gateway/test_resume_history_admission.py new file mode 100644 index 0000000000000..16d14af6f720c --- /dev/null +++ b/tests/tui_gateway/test_resume_history_admission.py @@ -0,0 +1,668 @@ +"""Cold-resume admission through real hydration, dispatch and temporary SQLite.""" +import io +import threading +from types import SimpleNamespace + +import pytest + +from hermes_state import SessionDB +from tests.test_tui_gateway_server import _configure_immediate_prompt_run, _session +from tests.tui_gateway.test_model_intent_admission_order import ( # noqa: F401 + REAL_THREAD, agent, assert_route, gateway, intent_turn, live_turn, routes, select, turn_env, +) +from tui_gateway import server +from tui_gateway.compute_host import ComputeHost +from tui_gateway.turn_marker import read_turn_marker, record_turn_start + +REAL_EMIT = server._emit +REAL_DRAIN = server._drain_queued_prompt +REAL_TERMINAL_ERROR = server._emit_terminal_turn_error +REAL_WAIT = server._wait_agent_for_prompt +REAL_RETIRE = server._retire_turn_marker + + +class HistoryGate(threading.Event): + def __init__(self): + super().__init__() + self.entered = threading.Event() + + def wait(self, timeout=None): + self.entered.set() + return super().wait(timeout) + + +class Transport: + def __init__(self): + self.frames = [] + self.terminal = threading.Event() + + def write(self, obj): + self.frames.append(obj) + if obj.get("params", {}).get("type") == "message.complete": + self.terminal.set() + return True + + +@pytest.fixture +def resume(monkeypatch, tmp_path): + _configure_immediate_prompt_run(monkeypatch, tmp_path, immediate_threads=False) + monkeypatch.setattr(server, "_emit", REAL_EMIT) + monkeypatch.setattr(server, "_sync_bot_capabilities", lambda *a: None) + monkeypatch.setattr(server, "_sync_agent_compression_with_config", lambda *a: None) + monkeypatch.setattr(server, "_ensure_active_session_slot", lambda *a: None) + monkeypatch.setattr(server, "_session_uses_compute_host", lambda *a: False) + monkeypatch.setattr(server, "_voice_mode_enabled", lambda: False) + monkeypatch.setattr(server, "_maybe_schedule_auto_continue", lambda *a: None) + monkeypatch.setattr(server, "_start_agent_build", lambda *a: None) + monkeypatch.setattr(server, "_notify_session_boundary", lambda *a: None) + monkeypatch.setattr(server, "_start_usage_ticker", lambda *a: (threading.Event(), SimpleNamespace(join=lambda *a, **kw: None))) + monkeypatch.setattr(server, "record_turn_start", lambda *a, **kw: None) + monkeypatch.setattr(server, "_retire_turn_marker", lambda *a: None) + monkeypatch.setattr(server, "_AGENT_BUILD_WAIT_SLICE", 0.01) + monkeypatch.setattr(server, "_agent_build_wait_cap", lambda: 5.0) + db = SessionDB(tmp_path / "history.db") + db.create_session("durable", source="tui") + db.append_message("durable", role="user", content="prior durable question") + db.append_message("durable", role="assistant", content="prior durable answer") + monkeypatch.setattr(server, "_db", db) + monkeypatch.setattr(server, "_get_db", lambda: db) + calls = [] + provider_entered = threading.Event() + provider_release = threading.Event() + provider_release.set() + + def conversation(prompt, conversation_history=None, **kwargs): + calls.append((prompt, list(conversation_history or []))) + provider_entered.set() + assert provider_release.wait(5) + return {"final_response": "offline", "messages": [], "completed": False, "api_calls": 0} + + agent = SimpleNamespace(model="inert-model", provider="inert-provider", session_id="durable", + _session_db=db, clear_interrupt=lambda: None, interrupt=lambda: None, + run_conversation=conversation) + ready = threading.Event() + ready.set() + gate = HistoryGate() + transport = Transport() + session = _session(agent=agent, session_key="durable", agent_ready=ready, + resume_history_ready=gate, resume_hydrating=True, transport=transport) + monkeypatch.setattr(server, "_sessions", {"runtime": session}) + read_entered, read_release = threading.Event(), threading.Event() + original_read = db.get_resume_conversations + error = [] + threads = [] + + def read(target): + read_entered.set() + assert read_release.wait(5) + if error: + raise RuntimeError(error[0]) + return original_read(target) + + monkeypatch.setattr(db, "get_resume_conversations", read) + state = SimpleNamespace(db=db, session=session, gate=gate, transport=transport, calls=calls, + provider_entered=provider_entered, provider_release=provider_release, read_entered=read_entered, + read_release=read_release, error=error, threads=threads) + yield state + read_release.set() + provider_release.set() + for thread in [*threads, session.get("_run_thread")]: + if thread is not None and thread is not threading.current_thread(): + thread.join(5) + assert not thread.is_alive() + assert not read_entered.is_set() or gate.wait(5) + db.close() + + +def hydrate(state): + server._schedule_resume_hydration("runtime", "durable", state.db) + assert state.read_entered.wait(5) + + +def launch(state, route): + if route == "inline": + response = server.dispatch({"id": "send", "method": "prompt.submit", "params": { + "session_id": "runtime", "text": "new question"}}, state.transport) + assert response["result"]["status"] == "streaming", response + if not response["result"].get("turn_isolation"): + state.threads.append(state.session["_run_thread"]) + return response + state.session["running"] = route != "queue" + if route == "queue": + state.session["queued_prompt"] = {"text": "new question"} + target = lambda: REAL_DRAIN("queue", "runtime", state.session) + else: + target = lambda: server._run_prompt_submit("direct", "runtime", state.session, + "new question", display_kind="auto_continue" if route == "continuation" else None) + thread = threading.Thread(target=target) + state.threads.append(thread) + thread.start() + return None + + +@pytest.mark.parametrize("route", ["inline", "direct", "queue", "continuation"]) +def test_delayed_real_sqlite_hydration_precedes_provider(resume, route): + hydrate(resume) + launch(resume, route) + assert resume.gate.entered.wait(2), f"history gate bypassed: {resume.calls!r}" + assert not resume.provider_entered.is_set() + resume.read_release.set() + assert resume.provider_entered.wait(5) + assert [row["content"] for row in resume.calls[0][1]] == [ + "prior durable question", "prior durable answer"] + + +@pytest.mark.parametrize("route", ["inline", "direct", "queue", "continuation"]) +def test_failed_hydration_delivers_correlated_terminal_to_actual_transport(resume, route): + resume.error.append("injected SQLite read failure") + hydrate(resume) + response = launch(resume, route) + assert resume.gate.entered.wait(2), f"history gate bypassed: {resume.calls!r}" + resume.read_release.set() + assert resume.transport.terminal.wait(5), resume.transport.frames + terminal = next(x["params"]["payload"] for x in resume.transport.frames + if x.get("params", {}).get("type") == "message.complete") + assert terminal["status"] == "error" + assert "injected SQLite read failure" in terminal["error"] + if response is not None: + assert terminal["accepted_turn"] == response["result"]["accepted_turn"] + else: + assert "accepted_turn" not in terminal + assert resume.calls == [] + + +@pytest.mark.parametrize("outcome", ["success", "cancel", "deadline"]) +def test_unused_once_keeps_original_restore_through_history_wait(intent_turn, monkeypatch, outcome): + session, _ = intent_turn + monkeypatch.setattr(server.threading, "Thread", REAL_THREAD) + monkeypatch.setattr(server, "_start_agent_build", lambda *a: None) + monkeypatch.setattr(server, "_maybe_schedule_auto_continue", lambda *a: None) + monkeypatch.setattr(server, "_AGENT_BUILD_WAIT_SLICE", 0.01) + monkeypatch.setattr(server, "_agent_build_wait_cap", lambda: 0.05 if outcome == "deadline" else 5.0) + value = session["agent"] + assert not server.handle_request(select("model-b", "endpoint-b", once=True)).get("error") + assert not server.handle_request(select("model-c", "endpoint-c", once=True)).get("error") + original = session["one_turn_model_restore"] + assert original["model"] == "model-a" + db = value._session_db + db.append_message(session["session_key"], role="user", content="prior durable question") + db.append_message(session["session_key"], role="assistant", content="prior durable answer") + read_entered, release_read = threading.Event(), threading.Event() + original_read = db.get_resume_conversations + def read(target): + read_entered.set() + assert release_read.wait(5) + return original_read(target) + monkeypatch.setattr(db, "get_resume_conversations", read) + gate = HistoryGate() + session.update(resume_history_ready=gate, resume_hydrating=True, running=True) + calls = [] + def conversation(prompt, **kwargs): + calls.append((value.model, value.provider)) + return {"final_response": "offline", "completed": False, "api_calls": 0} + monkeypatch.setattr(value, "run_conversation", conversation) + server._schedule_resume_hydration("selected", session["session_key"], db) + assert read_entered.wait(5) + waiter = REAL_THREAD(target=lambda: server._run_prompt_submit("once", "selected", session, "new question")) + waiter.start() + try: + assert gate.entered.wait(2) + assert calls == [] and session["one_turn_model_restore"] is original + if outcome == "cancel": + with session["history_lock"]: + session["_turn_cancel_requested"] = True + waiter.join(5) + elif outcome == "deadline": + waiter.join(5) + else: + release_read.set() + waiter.join(5) + worker = session.get("_run_thread") + if worker is not None: + worker.join(5) + assert not worker.is_alive() + assert not waiter.is_alive() + if outcome == "success": + assert calls == [("model-c", "endpoint-c")] + assert_route(value, "http://a.invalid/v1", "synthetic-key-a") + assert "one_turn_model_restore" not in session + else: + assert calls == [] + assert session["one_turn_model_restore"] is original + assert_route(value, "http://c.invalid/v1", "synthetic-key-c") + finally: + release_read.set() + waiter.join(5) + assert gate.wait(5) + + +def test_fresh_executing_host_uses_actual_owner_db_context(resume, monkeypatch): + monkeypatch.setattr(server, "_make_agent", lambda *a, **kw: resume.session["agent"]) + monkeypatch.setattr(server, "_persist_session_cwd_and_schedule_git_meta", lambda *a, **kw: None) + host = ComputeHost(stdout=io.StringIO(), heartbeat_secs=0, max_workers=1) + try: + record = host._ensure_server_session(server, {"sid": "executing-host", "session_key": "durable", + "history": [{"role": "user", "content": "stale parent mirror"}], "cwd": ".", "source": "desktop"}) + assert "resume_history_ready" not in record + record["running"] = True + assert server._run_prompt_submit("host", "executing-host", record, "host input") + record["_run_thread"].join(5) + assert not record["_run_thread"].is_alive() + assert [row["content"] for row in resume.calls[0][1]] == [ + "prior durable question", "prior durable answer"] + finally: + host.close() + + +def test_adopted_parent_history_error_does_not_gate_host_send_or_stop(resume, monkeypatch): + submitted, stopped = [], [] + class Host: + boot_id = "offline-host" + def submit_turn(self, frame, **kwargs): + frame["_admitted_host_boot_id"] = self.boot_id + submitted.append(frame) + def interrupt(self, sid, **kwargs): + stopped.append((sid, kwargs)) + return {"type": "interrupt.ack", "sid": sid, "request_id": kwargs["request_id"], + "target_request_id": kwargs["target_request_id"], "applied": True} + host = Host() + monkeypatch.setattr(server, "_get_compute_host_supervisor", lambda *a: host) + monkeypatch.setattr(server, "_session_uses_compute_host", lambda *a: True) + resume.session.update(agent=None, _compute_host_active=True, resume_history_error="display mirror unavailable") + response = launch(resume, "inline") + assert response["result"]["turn_isolation"] is True + stop = server.dispatch({"id": "stop", "method": "session.interrupt", "params": { + "session_id": "runtime"}}, resume.transport) + assert "error" not in stop, stop + assert submitted[0]["request_id"] == response["result"]["accepted_turn"]["request_id"] + assert stopped[0][1]["target_request_id"] == submitted[0]["request_id"] + assert not resume.gate.entered.is_set() + assert resume.session["agent"] is None + assert resume.calls == [] + + +@pytest.mark.parametrize("action", ["cancel", "replace", "queue_generation"]) +def test_obsolete_history_waiter_cannot_run_or_clear_successor(resume, action): + hydrate(resume) + resume.session["running"] = True + generation = 0 if action == "queue_generation" else None + thread = threading.Thread(target=lambda: server._run_prompt_submit("old", "runtime", + resume.session, "obsolete input", queued_prompt_generation=generation)) + resume.threads.append(thread) + thread.start() + assert resume.gate.entered.wait(2), f"history gate bypassed: {resume.calls!r}" + with resume.session["history_lock"]: + if action == "cancel": + resume.session["_turn_cancel_requested"] = True + elif action == "replace": + server._sessions["runtime"] = _session(running=True, inflight_turn={"user": "successor"}) + else: + resume.session["_queued_prompt_generation"] = 1 + server._start_inflight_turn(resume.session, "successor") + resume.read_release.set() + thread.join(5) + assert not thread.is_alive() + assert resume.calls == [] + if action != "cancel": + assert server._sessions["runtime"]["running"] is True + assert server._sessions["runtime"]["inflight_turn"]["user"] == "successor" + + +def test_history_only_wait_has_finite_cap(resume, monkeypatch): + hydrate(resume) + resume.session["running"] = True + monkeypatch.setattr(server, "_agent_build_wait_cap", lambda: 0.03) + result = server._wait_agent_for_prompt(resume.session, "bounded", "runtime") + assert result["error"]["code"] == 5032 + assert "was not sent" in result["error"]["message"] + assert not resume.gate.is_set() + assert resume.calls == [] + + +def test_actual_stop_then_new_send_keeps_successor_admission(resume): + hydrate(resume) + resume.session["running"] = True + old = threading.Thread(target=lambda: server._run_prompt_submit("old", "runtime", + resume.session, "obsolete input", queued_prompt_generation=0)) + resume.threads.append(old) + old.start() + assert resume.gate.entered.wait(2) + stop = server.dispatch({"id": "stop", "method": "session.interrupt", "params": { + "session_id": "runtime"}}, resume.transport) + assert "error" not in stop, stop + assert not resume.session["running"] + resume.provider_release.clear() + response = launch(resume, "inline") + successor = resume.session["inflight_turn"] + resume.read_release.set() + assert resume.provider_entered.wait(5) + old.join(5) + assert not old.is_alive() + assert [prompt for prompt, _ in resume.calls] == ["new question"] + assert resume.session["running"] + assert resume.session["inflight_turn"]["started_at"] is successor["started_at"] + assert resume.session["_turn_outcomes"].turns[-1]["accepted_turn"] == response["result"]["accepted_turn"] + + +def test_actual_inline_stop_during_hydration_does_not_replay_input(resume): + hydrate(resume) + launch(resume, "inline") + assert resume.gate.entered.wait(2) + stop = server.dispatch({"id": "stop", "method": "session.interrupt", "params": { + "session_id": "runtime"}}, resume.transport) + assert "error" not in stop, stop + resume.threads[0].join(5) + assert not resume.threads[0].is_alive() + assert not resume.session["running"] + assert any(frame["params"].get("payload", {}).get("message") == "Turn cancelled before the agent was ready" + for frame in resume.transport.frames if frame.get("method") == "event") + resume.read_release.set() + assert resume.gate.wait(5) + assert resume.calls == [] + + +def test_direct_history_deadline_delivers_terminal_without_spending_model_choice(resume, monkeypatch): + hydrate(resume) + pending = {"raw": "unused --session"} + once = {"model": "original"} + resume.session.update(pending_model_switch=pending, one_turn_model_restore=once) + monkeypatch.setattr(server, "_agent_build_wait_cap", lambda: 0.03) + launch(resume, "direct") + assert resume.transport.terminal.wait(5) + terminal = next(frame["params"]["payload"] for frame in resume.transport.frames + if frame.get("params", {}).get("type") == "message.complete") + assert terminal["status"] == "error" and "was not sent" in terminal["error"] + assert resume.session["inflight_turn"]["user"] == "new question" + assert resume.session["pending_model_switch"] is pending + assert resume.session["one_turn_model_restore"] is once + assert resume.calls == [] + + +def test_old_history_error_cannot_fail_new_send_after_stop(resume, monkeypatch, tmp_path): + resume.session["profile_home"] = tmp_path / "markers" + monkeypatch.setattr(server, "_retire_turn_marker", REAL_RETIRE) + hydrate(resume) + entered_error, release_error = threading.Event(), threading.Event() + def paused_error(*args, **kwargs): + entered_error.set() + assert release_error.wait(5) + return REAL_TERMINAL_ERROR(*args, **kwargs) + monkeypatch.setattr(server, "_emit_terminal_turn_error", paused_error) + monkeypatch.setattr(server, "_agent_build_wait_cap", lambda: 0.03) + launch(resume, "direct") + try: + assert entered_error.wait(5) + stop = server.dispatch({"id": "stop", "method": "session.interrupt", "params": { + "session_id": "runtime"}}, resume.transport) + assert "error" not in stop, stop + monkeypatch.setattr(server, "_agent_build_wait_cap", lambda: 5.0) + resume.provider_release.clear() + response = launch(resume, "inline") + successor = resume.session["inflight_turn"] + resume.read_release.set() + assert resume.provider_entered.wait(5) + marker = ("successor-provider", "successor-model") + resume.session["model_verified_for"] = marker + record_turn_start(resume.session["profile_home"], "durable", "successor") + release_error.set() + resume.threads[0].join(5) + assert not resume.threads[0].is_alive() + assert resume.session["inflight_turn"]["started_at"] is successor["started_at"] + assert resume.session["inflight_turn"].get("status") != "error" + assert resume.session["model_verified_for"] == marker + assert resume.session["running"] + assert read_turn_marker(resume.session["profile_home"], "durable")["prompt"] == "successor" + assert resume.session["_turn_outcomes"].turns[-1]["accepted_turn"] == response["result"]["accepted_turn"] + finally: + release_error.set() + + +def test_waiter_starting_after_failed_owner_removal_rejects_replacement(resume): + removed = threading.Event() + resume.session["active_session_lease"] = SimpleNamespace(release=removed.set) + resume.session["running"] = True + resume.error.append("injected SQLite read failure") + hydrate(resume) + resume.read_release.set() + assert removed.wait(5) + replacement = _session(running=True, inflight_turn={"user": "successor"}) + server._sessions["runtime"] = replacement + assert server._wait_agent_for_prompt(resume.session, "obsolete", "runtime") is None + assert server._sessions["runtime"] is replacement + assert replacement["running"] and replacement["inflight_turn"] == {"user": "successor"} + + +@pytest.mark.parametrize("route", ["inline", "direct", "queue", "continuation"]) +def test_completed_failed_owner_removal_still_delivers_admitted_error(resume, monkeypatch, route): + removed = threading.Event() + resume.session["active_session_lease"] = SimpleNamespace(release=removed.set) + def wait_after_removal(*args): + error = REAL_WAIT(*args) + assert error is not None + assert removed.wait(5) + return error + monkeypatch.setattr(server, "_wait_agent_for_prompt", wait_after_removal) + resume.error.append("injected SQLite read failure") + hydrate(resume) + response = launch(resume, route) + assert resume.gate.entered.wait(2) + resume.read_release.set() + assert resume.transport.terminal.wait(5) + assert removed.wait(5) and server._sessions.get("runtime") is None + terminal = next(frame["params"]["payload"] for frame in resume.transport.frames + if frame.get("params", {}).get("type") == "message.complete") + assert terminal["status"] == "error" and "injected SQLite read failure" in terminal["error"] + if response is not None: + assert terminal["accepted_turn"] == response["result"]["accepted_turn"] + else: + assert "accepted_turn" not in terminal + assert resume.calls == [] + + +def test_inline_cancel_does_not_emit_to_replacement_transport(resume, monkeypatch): + cancelled, release_waiter = threading.Event(), threading.Event() + def paused_wait(*args): + result = REAL_WAIT(*args) + cancelled.set() + assert release_waiter.wait(5) + return result + monkeypatch.setattr(server, "_wait_agent_for_prompt", paused_wait) + hydrate(resume) + launch(resume, "inline") + assert resume.gate.entered.wait(2) + try: + response = server.dispatch({"id": "stop", "method": "session.interrupt", "params": { + "session_id": "runtime"}}, resume.transport) + assert "error" not in response + assert cancelled.wait(5) + transport = Transport() + replacement = _session(running=True, transport=transport, inflight_turn={"user": "successor"}) + server._sessions["runtime"] = replacement + release_waiter.set() + resume.threads[0].join(5) + assert not resume.threads[0].is_alive() + assert transport.frames == [] + assert replacement["running"] and replacement["inflight_turn"] == {"user": "successor"} + assert resume.calls == [] + finally: + release_waiter.set() + + +def test_inline_error_settlement_does_not_clear_later_send(resume, monkeypatch): + settled, release_error = threading.Event(), threading.Event() + def paused_after_settlement(*args, **kwargs): + result = REAL_TERMINAL_ERROR(*args, **kwargs) + settled.set() + assert release_error.wait(5) + return result + monkeypatch.setattr(server, "_emit_terminal_turn_error", paused_after_settlement) + monkeypatch.setattr(server, "_agent_build_wait_cap", lambda: 0.03) + hydrate(resume) + response = launch(resume, "inline") + old_thread = resume.threads[0] + try: + assert settled.wait(5) + assert not resume.session["running"] + terminal = next(frame["params"]["payload"] for frame in resume.transport.frames + if frame.get("params", {}).get("type") == "message.complete") + assert terminal["accepted_turn"] == response["result"]["accepted_turn"] + monkeypatch.setattr(server, "_agent_build_wait_cap", lambda: 5.0) + resume.provider_release.clear() + successor_ack = launch(resume, "inline") + assert successor_ack["result"]["accepted_turn"] != response["result"]["accepted_turn"] + successor = resume.session["inflight_turn"] + resume.read_release.set() + assert resume.provider_entered.wait(5) + release_error.set() + old_thread.join(5) + assert not old_thread.is_alive() + assert resume.session["running"] + assert resume.session["inflight_turn"]["started_at"] is successor["started_at"] + assert resume.session["inflight_turn"].get("status") != "error" + assert resume.session["_turn_outcomes"].turns[-1]["accepted_turn"] == successor_ack["result"]["accepted_turn"] + assert [text for text, _ in resume.calls] == ["new question"] + finally: + release_error.set() + + +@pytest.mark.parametrize("original_sink", ["transport", "stdio"]) +def test_replaced_legacy_terminal_uses_original_sink_without_borrowing_correlation(resume, monkeypatch, original_sink): + ref = server._begin_turn_outcome(resume.session, "runtime", "old-admission", "inline") + original = resume.transport + if original_sink == "stdio": + resume.session["transport"] = None + monkeypatch.setattr(server, "_stdio_transport", original) + replacement_transport = Transport() + replacement = _session(running=True, transport=replacement_transport, inflight_turn={"user": "successor"}) + server._sessions["runtime"] = replacement + successor_ref = server._begin_turn_outcome(replacement, "runtime", "successor-admission", "inline") + execution_token = server._turn_outcome_execution.set((resume.session, "runtime", ref["request_id"])) + transport_token = server.bind_transport(None) + try: + server._emit("message.complete", "runtime", {"text": "old terminal", "status": "error"}) + finally: + server.reset_transport(transport_token) + server._turn_outcome_execution.reset(execution_token) + terminal = next(frame["params"]["payload"] for frame in original.frames + if frame.get("params", {}).get("type") == "message.complete") + assert terminal["text"] == "old terminal" and "accepted_turn" not in terminal + assert replacement_transport.frames == [] + assert replacement["_turn_outcomes"].turns[-1]["accepted_turn"] == successor_ref + assert replacement["_turn_outcomes"].turns[-1]["finalized"] == [] + assert replacement["running"] and replacement["inflight_turn"] == {"user": "successor"} + + +def test_replacement_during_terminal_projection_cannot_capture_old_stdio_reply(resume, monkeypatch): + entered, release = threading.Event(), threading.Event() + resume.session["transport"] = None + monkeypatch.setattr(server, "_stdio_transport", resume.transport) + ref = server._begin_turn_outcome(resume.session, "runtime", "old-admission", "inline") + original_lock = resume.session["history_lock"] + class LockGate: + def __enter__(self): + original_lock.acquire() + entered.set() + assert release.wait(5) + def __exit__(self, *args): + original_lock.release() + resume.session["history_lock"] = LockGate() + def emit(): + token = server._turn_outcome_execution.set((resume.session, "runtime", ref["request_id"])) + transport_token = server.bind_transport(None) + try: + server._emit("message.complete", "runtime", {"text": "old terminal", "status": "error"}) + finally: + server.reset_transport(transport_token) + server._turn_outcome_execution.reset(token) + thread = threading.Thread(target=emit) + resume.threads.append(thread) + thread.start() + try: + assert entered.wait(5) + transport = Transport() + replacement = _session(running=True, transport=transport, inflight_turn={"user": "successor"}) + server._sessions["runtime"] = replacement + successor_ref = server._begin_turn_outcome(replacement, "runtime", "successor-admission", "inline") + release.set() + thread.join(5) + assert not thread.is_alive() + assert transport.frames == [] + terminal = next(frame["params"]["payload"] for frame in resume.transport.frames + if frame.get("params", {}).get("type") == "message.complete") + assert terminal["text"] == "old terminal" and "accepted_turn" not in terminal + assert replacement["_turn_outcomes"].turns[-1]["accepted_turn"] == successor_ref + assert replacement["_turn_outcomes"].turns[-1]["finalized"] == [] + assert replacement["running"] and replacement["inflight_turn"] == {"user": "successor"} + finally: + release.set() + + +@pytest.mark.parametrize("route", ["inline", "direct"]) +def test_removed_history_owner_cannot_retire_new_runtime_marker(resume, monkeypatch, tmp_path, route): + from hermes_cli.active_sessions import try_acquire_active_session + + home = tmp_path / "markers" + resume.session["profile_home"] = home + lease, error = try_acquire_active_session(session_id="durable", surface="tui", + config={"max_concurrent_sessions": 1}, metadata={"live_session_id": "runtime"}) + assert lease is not None and error is None and lease.enabled + removed, entered_error, release_error = threading.Event(), threading.Event(), threading.Event() + real_release = lease.release + def release_lease(): + real_release() + removed.set() + monkeypatch.setattr(lease, "release", release_lease) + resume.session["active_session_lease"] = lease + def paused_error(*args, **kwargs): + assert removed.wait(5) and lease.released + entered_error.set() + assert release_error.wait(5) + return REAL_TERMINAL_ERROR(*args, **kwargs) + monkeypatch.setattr(server, "_emit_terminal_turn_error", paused_error) + monkeypatch.setattr(server, "_retire_turn_marker", REAL_RETIRE) + monkeypatch.setattr(server, "record_turn_start", record_turn_start) + resume.error.append("injected SQLite read failure") + hydrate(resume) + launch(resume, route) + assert resume.gate.entered.wait(2) + new_lease = None + owned_threads = [] + try: + resume.read_release.set() + assert entered_error.wait(5) and server._sessions.get("runtime") is None + new_lease, error = try_acquire_active_session(session_id="durable", surface="tui", + config={"max_concurrent_sessions": 1}, metadata={"live_session_id": "replacement-runtime"}) + assert new_lease is not None and error is None and not new_lease.released + fresh_agent = SimpleNamespace(**vars(resume.session["agent"])) + transport = Transport() + successor = _session(agent=fresh_agent, session_key="durable", transport=transport, + profile_home=home, active_session_lease=new_lease, + history=resume.db.get_messages_as_conversation("durable")) + server._sessions["replacement-runtime"] = successor + resume.provider_release.clear() + ack = server.dispatch({"id": "successor", "method": "prompt.submit", "params": { + "session_id": "replacement-runtime", "text": "successor question"}}, transport) + assert ack["result"]["status"] == "streaming", ack + assert resume.provider_entered.wait(5) + owned_threads = [t for t in threading.enumerate() + if getattr(t, "_turn_outcome_execution", (None,))[0] is successor] + resume.threads.extend(owned_threads) + turn = successor["inflight_turn"] + assert read_turn_marker(home, "durable")["prompt"] == "successor question" + release_error.set() + resume.threads[0].join(5) + assert not resume.threads[0].is_alive() + marker = read_turn_marker(home, "durable") + assert marker and marker["prompt"] == "successor question" + assert successor["running"] and successor["inflight_turn"]["started_at"] is turn["started_at"] + assert successor["_turn_outcomes"].turns[-1]["accepted_turn"] == ack["result"]["accepted_turn"] + assert not new_lease.released + finally: + release_error.set() + resume.provider_release.set() + for thread in owned_threads: + thread.join(5) + assert not thread.is_alive() + if new_lease is not None: + new_lease.release() + lease.release() diff --git a/tests/tui_gateway/test_tool_policy_selection.py b/tests/tui_gateway/test_tool_policy_selection.py new file mode 100644 index 0000000000000..b77f2596737f3 --- /dev/null +++ b/tests/tui_gateway/test_tool_policy_selection.py @@ -0,0 +1,249 @@ +"""Gateway policy survives schema construction; no provider or handler runs. + +The constructor substitute assembles schemas with the real model_tools path. +The group-member test also uses the real session.create and deferred builder, +with a temporary profile. Inference, MCP discovery, notification, and approval +side effects are inert. These tests do not qualify native-provider execution or +post-construction memory/context-engine schema injection. +""" + +from types import SimpleNamespace + +import pytest + + +@pytest.fixture +def inert_surface(monkeypatch): + import model_tools + import toolsets + from tools.registry import registry + + monkeypatch.setattr(registry, "_tools", {}) + monkeypatch.setattr(registry, "_scoped_tools", {}) + monkeypatch.setattr(registry, "_generation", registry._generation) + monkeypatch.setattr(toolsets, "_resolve_toolset_memo", {}) + names = { + "policy_canary_allowed": "policy_canary_keep", + "policy_canary_forbidden": "policy_canary_deny", + "policy_canary_mcp": "mcp_policy_canary", + } + for name, toolset in names.items(): + def handler(*args, **kwargs): + pytest.fail("an inert canary handler must never execute") + + registry.register( + name=name, toolset=toolset, + schema={"name": name, "description": "Inert policy witness", + "parameters": {"type": "object", "properties": {}}}, + handler=handler, check_fn=lambda: True, + ) + monkeypatch.setitem(toolsets.TOOLSETS, "policy_canary_composite", { + "tools": [], "includes": list(names.values()), + }) + return model_tools + + +@pytest.fixture +def construction(monkeypatch, inert_surface): + import run_agent + from tui_gateway import server + import agent.coding_context + import hermes_cli.mcp_startup + import tui_gateway.entry + + class SchemaAgent: + def __init__(self, **kwargs): + self.kwargs = kwargs + self.__dict__.update(kwargs) + self.tools = inert_surface.get_tool_definitions( + enabled_toolsets=kwargs.get("enabled_toolsets"), + disabled_toolsets=kwargs.get("disabled_toolsets"), + quiet_mode=True, skip_tool_search_assembly=True, + ) + self.valid_tool_names = { + (tool.get("function") or tool)["name"] for tool in self.tools + } + + monkeypatch.setattr(run_agent, "AIAgent", SchemaAgent) + monkeypatch.setattr(server, "_resolve_startup_runtime", lambda: ("inert-model", "inert")) + monkeypatch.setattr(server, "_resolve_runtime_with_fallback", lambda *a, **k: + SimpleNamespace(runtime={"provider": "inert", "api_mode": "chat_completions"}, + used_fallback=False)) + monkeypatch.setattr(server, "_load_provider_routing", lambda: {}) + monkeypatch.setattr(server, "_load_reasoning_config", lambda *a: None) + monkeypatch.setattr(server, "_load_service_tier", lambda: None) + monkeypatch.setattr(server, "_load_fallback_model", lambda: None) + monkeypatch.setattr(server, "_parse_tui_skills_env", lambda: []) + monkeypatch.setattr(server, "_get_db", lambda: None) + monkeypatch.setattr(server, "_agent_cbs", lambda *a: {}) + monkeypatch.setattr(agent.coding_context, "coding_selection", lambda **k: None) + monkeypatch.setattr(hermes_cli.mcp_startup, "wait_for_mcp_discovery", lambda: None) + monkeypatch.setattr(tui_gateway.entry, "wait_for_mcp_discovery", lambda: None) + monkeypatch.setattr(tui_gateway.entry, "ensure_mcp_discovery_started", lambda: None) + monkeypatch.delenv("HERMES_TUI_TOOLSETS", raising=False) + return server, SchemaAgent + + +@pytest.mark.parametrize("surface, expected", [ + ("desktop", ["desktop_ui", "project"]), ("tui", ["project"]), + ("no-canary-gui-grant", []), +]) +def test_empty_canonical_selection_never_becomes_all(monkeypatch, surface, expected): + from tui_gateway import server + import agent.coding_context + import hermes_cli.config + import hermes_cli.tools_config + + monkeypatch.delenv("HERMES_TUI_TOOLSETS", raising=False) + monkeypatch.setattr(agent.coding_context, "coding_selection", lambda **k: None) + monkeypatch.setattr(hermes_cli.config, "load_config", lambda: {}) + monkeypatch.setattr(hermes_cli.tools_config, "_get_platform_tools", lambda *a, **k: set()) + if surface == "no-canary-gui-grant": + monkeypatch.setattr(server, "_gui_surface_toolsets", lambda *a: set()) + assert server._load_enabled_toolsets(surface) == expected + + +@pytest.mark.parametrize("disabled", [ + ["policy_canary_deny", "mcp_policy_canary"], + '["policy_canary_deny", "mcp_policy_canary"]', + "['policy_canary_deny', 'mcp_policy_canary']", + [" policy_canary_deny ", "mcp_policy_canary", ""], +]) +def test_fresh_agent_subtracts_disabled_members_of_composite(construction, monkeypatch, disabled): + server, _ = construction + monkeypatch.setattr(server, "_load_cfg", lambda: {"agent": {"disabled_toolsets": disabled}}) + monkeypatch.setattr(server, "_load_enabled_toolsets", lambda *a: ["policy_canary_composite"]) + agent = server._make_agent("canary-fresh", "canary-key", platform_override="desktop") + assert agent.valid_tool_names == {"policy_canary_allowed"} + assert agent.disabled_toolsets == ["policy_canary_deny", "mcp_policy_canary"] + assert agent.platform == "desktop" + + +@pytest.mark.parametrize("agent_cfg", [None, {}, {"disabled_toolsets": []}]) +def test_unset_disabled_preserves_existing_default_surface(construction, monkeypatch, agent_cfg): + server, _ = construction + monkeypatch.setattr(server, "_load_cfg", lambda: {"agent": agent_cfg}) + monkeypatch.setattr(server, "_load_enabled_toolsets", lambda *a: None) + agent = server._make_agent("canary-default", "canary-key") + assert agent.valid_tool_names == { + "policy_canary_allowed", "policy_canary_forbidden", "policy_canary_mcp", + } + + +@pytest.mark.parametrize("enabled", [[], ["policy_canary_composite"]]) +def test_background_preserves_explicit_selection_and_disabled(construction, monkeypatch, enabled): + server, SchemaAgent = construction + monkeypatch.setattr(server, "_load_cfg", lambda: {}) + monkeypatch.setattr(server, "_load_enabled_toolsets", lambda *a: + pytest.fail("an explicit parent selection must not fall back")) + parent = SimpleNamespace(enabled_toolsets=enabled, disabled_toolsets=["policy_canary_deny", "mcp_policy_canary"], + model="inert-model", provider="inert") + kwargs = server._background_agent_kwargs(parent, "canary-child") + child = SchemaAgent(**kwargs) + assert kwargs["enabled_toolsets"] == enabled + assert kwargs["disabled_toolsets"] == parent.disabled_toolsets + assert child.valid_tool_names == ({"policy_canary_allowed"} if enabled else set()) + assert kwargs["platform"] == "tui" + kwargs["disabled_toolsets"].append("child-only") + assert "child-only" not in parent.disabled_toolsets + + +@pytest.mark.parametrize("parent", [SimpleNamespace(model="inert-model"), + SimpleNamespace(model="inert-model", enabled_toolsets=None)]) +def test_background_none_or_missing_still_uses_tui_defaults(construction, monkeypatch, parent): + server, _ = construction + calls = [] + monkeypatch.setattr(server, "_load_cfg", lambda: {}) + monkeypatch.setattr(server, "_load_enabled_toolsets", lambda platform: + calls.append(platform) or ["policy_canary_keep"]) + kwargs = server._background_agent_kwargs(parent, "canary-child") + assert kwargs["enabled_toolsets"] == ["policy_canary_keep"] + assert calls == ["tui"] + + +def test_background_explicit_empty_disabled_stays_empty(construction, monkeypatch): + server, _ = construction + monkeypatch.setattr(server, "_load_cfg", lambda: {"agent": {"disabled_toolsets": ["policy_canary_keep"]}}) + parent = SimpleNamespace(model="inert-model", enabled_toolsets=[], disabled_toolsets=[]) + kwargs = server._background_agent_kwargs(parent, "canary-child") + assert kwargs["disabled_toolsets"] == [] + + +def test_group_member_session_create_uses_its_profile_policy(construction, monkeypatch, tmp_path): + import yaml + from hermes_constants import get_hermes_home_override + import agent.credits_tracker + import tools.approval + import hermes_cli.profiles + + server, _ = construction + profile = tmp_path / "policy-canary-member" + profile.mkdir() + cfg = {"platform_toolsets": {"cli": ["policy_canary_composite"]}, + "agent": {"disabled_toolsets": ["policy_canary_deny", "mcp_policy_canary"]}, + "known_plugin_toolsets": {"cli": ["policy_canary_keep", "policy_canary_deny", "mcp_policy_canary"]}, + "mcp_servers": {}, "context": {"engine": "compressor"}} + config_file = profile / "config.yaml" + config_file.write_text(yaml.safe_dump(cfg)) + original = config_file.read_bytes() + before_override = get_hermes_home_override() + events = [] + monkeypatch.setattr(server, "_sessions", {}) + monkeypatch.setattr(hermes_cli.profiles, "get_profile_dir", lambda name: str(profile)) + monkeypatch.setattr(server, "_schedule_agent_build", lambda sid: None) + monkeypatch.setattr(server, "_schedule_session_cap_enforcement", lambda: None) + monkeypatch.setattr(server, "_enable_gateway_prompts", lambda: None) + monkeypatch.setattr(server, "_open_profile_session_db", lambda *a: None) + monkeypatch.setattr(server, "_wire_callbacks", lambda *a: None) + monkeypatch.setattr(server, "_start_notification_poller", lambda *a: None) + monkeypatch.setattr(server, "_notify_session_boundary", lambda *a: None) + monkeypatch.setattr(server, "_schedule_mcp_late_refresh", lambda *a: None) + monkeypatch.setattr(server, "_session_info", lambda agent, session: + {"tools": sorted(agent.valid_tool_names), "source": session["source"]}) + monkeypatch.setattr(server, "_emit", lambda event, sid, data: events.append((event, data))) + monkeypatch.setattr(tools.approval, "register_gateway_notify", lambda *a: None) + monkeypatch.setattr(tools.approval, "load_permanent_allowlist", lambda: None) + monkeypatch.setattr(agent.credits_tracker, "seed_credits_at_session_start", lambda *a: None) + + result = server._methods["session.create"]("policy-canary-create", { + "profile": "policy-canary-member", "source": "desktop", "hidden": True, + "title": "OFFLINE POLICY CANARY", "cwd": str(tmp_path), + "model": "inert-session-model", "provider": "inert", + }) + sid = result["result"]["session_id"] + session = server._sessions[sid] + assert session["history"] == [] + server._start_agent_build(sid, session) + assert session["agent_ready"].wait(5), "bounded inert build did not complete" + session["_agent_build_thread"].join(1) + assert session["agent_error"] is None + built = session["agent"] + assert built.valid_tool_names == {"policy_canary_allowed"} + assert built.model == "inert-session-model" + assert built.provider == "inert" + assert built.platform == "desktop" + assert session["profile_home"] == str(profile) + assert any(event == "session.info" and data["tools"] == ["policy_canary_allowed"] + for event, data in events) + assert config_file.read_bytes() == original + assert get_hermes_home_override() == before_override + + +@pytest.mark.parametrize("pin", ["all", "*"]) +def test_explicit_all_pin_keeps_existing_default_meaning(monkeypatch, pin): + from tui_gateway import server + monkeypatch.setenv("HERMES_TUI_TOOLSETS", pin) + assert server._load_enabled_toolsets("desktop") is None + + +def test_config_load_failure_keeps_existing_default_meaning(monkeypatch): + from tui_gateway import server + import agent.coding_context + import hermes_cli.config + + monkeypatch.delenv("HERMES_TUI_TOOLSETS", raising=False) + monkeypatch.setattr(agent.coding_context, "coding_selection", lambda **k: None) + def broken_config(): + raise ValueError("inert failure") + monkeypatch.setattr(hermes_cli.config, "load_config", broken_config) + assert server._load_enabled_toolsets("desktop") is None diff --git a/tests/tui_gateway/test_turn_outcome_projection.py b/tests/tui_gateway/test_turn_outcome_projection.py index 3eddb8559fd13..f0292d0b49b1f 100644 --- a/tests/tui_gateway/test_turn_outcome_projection.py +++ b/tests/tui_gateway/test_turn_outcome_projection.py @@ -26,7 +26,7 @@ def projection_env(monkeypatch): monkeypatch.setattr(server, "_load_dashboard_process_isolation_config", lambda: {}) monkeypatch.setattr(server, "_pending_clarify_request_payload", lambda sid: None) frames = [] - monkeypatch.setattr(server, "write_json", frames.append) + monkeypatch.setattr(server, "_stdio_transport", types.SimpleNamespace(write=frames.append)) yield frames # These sessions use fake host pipes; production teardown would wait for # control ACKs that this in-process fixture intentionally cannot produce. diff --git a/tools/async_delegation.py b/tools/async_delegation.py index 4965363a07eef..9b73154e1fb52 100644 --- a/tools/async_delegation.py +++ b/tools/async_delegation.py @@ -233,27 +233,38 @@ def _capture_routing_origin() -> Dict[str, Any]: return origin -def _persist_dispatch(record: Dict[str, Any]) -> None: +class _DurableAdmissionFull(RuntimeError): + """The existing pending-receipt admission policy has no free slot.""" + + +def _persist_dispatch(record: Dict[str, Any]) -> Dict[str, Any]: + """Reserve one pending slot and return an ephemeral release fence.""" now = time.time() + owner_pid = __import__("os").getpid() try: from gateway.status import get_process_start_time - owner_started_at = get_process_start_time(__import__("os").getpid()) + owner_started_at = get_process_start_time(owner_pid) except Exception: owner_started_at = None task_payload = { key: record.get(key) for key in ( "goal", "goals", "context", "toolsets", "role", "model", "is_batch", - # Routing origin (scope_id/user_id/user_name): persisted so a - # restart-recovered completion can reconstruct a full - # SessionSource — see _capture_routing_origin. "scope_id", "user_id", "user_name", ) if key in record } with _DB_LOCK, _transaction() as conn: + # Serialize the count and insert across processes using the existing + # ledger. Every pending row occupies a slot, including active/claimed. + conn.execute("BEGIN IMMEDIATE") + pending = conn.execute( + "SELECT COUNT(*) FROM async_delegations WHERE delivery_state='pending'" + ).fetchone()[0] + if pending >= _MAX_DURABLE_PENDING: + raise _DurableAdmissionFull("Pending delegation receipt capacity reached") conn.execute( - """INSERT OR REPLACE INTO async_delegations + """INSERT INTO async_delegations (delegation_id, origin_session, origin_ui_session_id, parent_session_id, state, dispatched_at, updated_at, delivery_state, delivery_attempts, owner_pid, @@ -261,55 +272,77 @@ def _persist_dispatch(record: Dict[str, Any]) -> None: VALUES (?, ?, ?, ?, 'running', ?, ?, 'pending', 0, ?, ?, ?, ?)""", (record["delegation_id"], record.get("session_key", ""), record.get("origin_ui_session_id", ""), record.get("parent_session_id"), - record["dispatched_at"], now, __import__("os").getpid(), + record["dispatched_at"], now, owner_pid, owner_started_at, json.dumps(task_payload), record.get("origin_session_id", "")), ) - _prune_durable_records() + reservation = { + "delegation_id": record["delegation_id"], "owner_pid": owner_pid, + "owner_started_at": owner_started_at, + "dispatched_at": record["dispatched_at"], "updated_at": now, + } + # The admission is committed. Maintenance failure cannot turn it into a + # refusal or cause a caller to replay work that has already been admitted. + try: + _prune_durable_records() + except Exception: + logger.exception("Delegation admitted; disposed-history maintenance failed") + return reservation + -def _delete_durable_delegation(delegation_id: str) -> None: +def _delete_durable_delegation( + delegation_id: str, *, reservation: Dict[str, Any] +) -> bool: + """Release only this process's exact reservation before submission.""" + if reservation.get("delegation_id") != delegation_id: + return False with _DB_LOCK, _transaction() as conn: - conn.execute("DELETE FROM async_delegations WHERE delegation_id=?", (delegation_id,)) + deleted = conn.execute( + """DELETE FROM async_delegations + WHERE delegation_id=? AND owner_pid=? AND owner_started_at IS ? + AND dispatched_at=? AND updated_at=? + AND state='running' AND delivery_state='pending' + AND delivery_attempts=0 AND delivery_claim IS NULL + AND delivery_claimed_at IS NULL AND completed_at IS NULL + AND event_json IS NULL AND result_json IS NULL""", + (delegation_id, reservation["owner_pid"], + reservation["owner_started_at"], reservation["dispatched_at"], + reservation["updated_at"]), + ).rowcount == 1 + return deleted + def _prune_durable_records() -> None: - """Bound terminal history, preferring delivered records for deletion.""" + """Bound disposed history without deleting pending or claimed receipts.""" now = time.time() cutoff = now - _DURABLE_RETENTION_SECONDS + eligible = ( + "state NOT IN ('running','finalizing') " + "AND delivery_state IN ('delivered','dropped') " + "AND delivery_claim IS NULL" + ) with _DB_LOCK, _transaction() as conn: conn.execute( - "DELETE FROM async_delegations WHERE delivery_state='delivered' AND updated_at < ?", + f"""DELETE FROM async_delegations WHERE {eligible} + AND delivery_state='delivered' AND updated_at < ?""", (cutoff,), ) terminal_count = conn.execute( - "SELECT COUNT(*) FROM async_delegations WHERE state NOT IN ('running','finalizing')" + f"SELECT COUNT(*) FROM async_delegations WHERE {eligible}" ).fetchone()[0] excess = max(0, terminal_count - _MAX_RETAINED_COMPLETED) if excess: conn.execute( - """DELETE FROM async_delegations WHERE delegation_id IN ( + f"""DELETE FROM async_delegations WHERE delegation_id IN ( SELECT delegation_id FROM async_delegations - WHERE state NOT IN ('running','finalizing') + WHERE {eligible} ORDER BY CASE delivery_state WHEN 'delivered' THEN 0 ELSE 1 END, updated_at ASC LIMIT ? )""", (excess,), ) - pending_count = conn.execute( - """SELECT COUNT(*) FROM async_delegations - WHERE state NOT IN ('running','finalizing') AND delivery_state='pending'""" - ).fetchone()[0] - overflow = max(0, pending_count - _MAX_DURABLE_PENDING) - if overflow: - conn.execute( - """DELETE FROM async_delegations WHERE delegation_id IN ( - SELECT delegation_id FROM async_delegations - WHERE state NOT IN ('running','finalizing') AND delivery_state='pending' - ORDER BY updated_at ASC LIMIT ? - )""", - (overflow,), - ) def _persist_completion(event: Dict[str, Any], result: Dict[str, Any]) -> None: @@ -836,6 +869,12 @@ def dispatch_async_delegation( # active_count() separately would let two concurrent dispatches (e.g. # from different gateway sessions) both pass the check and exceed the cap. with _records_lock: + if delegation_id in _records: + return { + "status": "rejected", "error_code": "duplicate_delegation_id", + "execution_started": False, + "error": "Delegation identity already exists; preserved its receipt", + } running = sum( 1 for r in _records.values() if r.get("status") in ("running", "stalling") @@ -843,6 +882,8 @@ def dispatch_async_delegation( if running >= max_async_children: return { "status": "rejected", + "error_code": "pool_capacity", + "execution_started": False, "error": ( f"Async delegation capacity reached ({max_async_children} " f"running). Wait for one to finish (its result will re-enter " @@ -853,8 +894,42 @@ def dispatch_async_delegation( } _records[delegation_id] = record - _persist_dispatch(record) - executor = _get_executor(max_async_children) + try: + reservation = _persist_dispatch(record) + except Exception as exc: + with _records_lock: + if _records.get(delegation_id) is record: + del _records[delegation_id] + return { + "status": "rejected", + "error_code": ("durable_backlog_full" if isinstance(exc, _DurableAdmissionFull) + else "durable_storage_unavailable"), + "execution_started": False, + "delegation_id": delegation_id, + "error": f"Async delegation was not submitted: {exc}", + } + try: + executor = _get_executor(max_async_children) + except Exception as exc: + # No submit call occurred. Release only the unchanged owned row; + # failed or conflicting cleanup retains custody rather than guessing. + released = False + try: + released = _delete_durable_delegation( + delegation_id, reservation=reservation + ) + except Exception: + logger.exception("Could not confirm unstarted delegation release %s", delegation_id) + if released: + with _records_lock: + if _records.get(delegation_id) is record: + del _records[delegation_id] + return { + "status": "rejected", "error_code": "executor_unavailable", + "execution_started": False, "delegation_id": delegation_id, + "durable_reservation_released": released, + "error": f"Async executor unavailable before submission: {exc}", + } def _worker() -> None: result: Dict[str, Any] = {} @@ -879,13 +954,14 @@ def _worker() -> None: # Propagate the dispatching profile so the detached child resolves # get_hermes_home() under the right profile. executor.submit(propagate_context_to_thread(_worker)) - except Exception as exc: # pragma: no cover — pool submit failure is rare - with _records_lock: - _records.pop(delegation_id, None) - _delete_durable_delegation(delegation_id) + except Exception as exc: + # submit() may enqueue before raising. Keep both ledgers and the + # handle; callers must never replay an ambiguous scheduling outcome. + logger.exception("Async submission outcome uncertain for %s", delegation_id) return { - "status": "rejected", - "error": f"Failed to schedule async delegation: {exc}", + "status": "dispatch_uncertain", "error_code": "scheduling_uncertain", + "execution_started": None, "delegation_id": delegation_id, + "error": f"Async submission outcome is unknown: {exc}", } if progress_fn is not None: _ensure_stale_monitor() @@ -897,6 +973,7 @@ def _worker() -> None: return {"status": "dispatched", "delegation_id": delegation_id} + def _finalize(delegation_id: str, result: Dict[str, Any], status: str) -> None: """Mark a record complete and push the completion event onto the queue.""" claimed = _begin_finalization(delegation_id) @@ -1080,6 +1157,12 @@ def dispatch_async_delegation_batch( "_interrupted_at": None, } with _records_lock: + if delegation_id in _records: + return { + "status": "rejected", "error_code": "duplicate_delegation_id", + "execution_started": False, + "error": "Delegation identity already exists; preserved its receipt", + } running = sum( 1 for r in _records.values() if r.get("status") in ("running", "stalling") @@ -1087,6 +1170,8 @@ def dispatch_async_delegation_batch( if running >= max_async_children: return { "status": "rejected", + "error_code": "pool_capacity", + "execution_started": False, "error": ( f"Async delegation capacity reached ({max_async_children} " f"running). Wait for one to finish (its result will re-enter " @@ -1096,8 +1181,42 @@ def dispatch_async_delegation_batch( } _records[delegation_id] = record - _persist_dispatch(record) - executor = _get_executor(max_async_children) + try: + reservation = _persist_dispatch(record) + except Exception as exc: + with _records_lock: + if _records.get(delegation_id) is record: + del _records[delegation_id] + return { + "status": "rejected", + "error_code": ("durable_backlog_full" if isinstance(exc, _DurableAdmissionFull) + else "durable_storage_unavailable"), + "execution_started": False, + "delegation_id": delegation_id, + "error": f"Async delegation was not submitted: {exc}", + } + try: + executor = _get_executor(max_async_children) + except Exception as exc: + # No submit call occurred. Release only the unchanged owned row; + # failed or conflicting cleanup retains custody rather than guessing. + released = False + try: + released = _delete_durable_delegation( + delegation_id, reservation=reservation + ) + except Exception: + logger.exception("Could not confirm unstarted delegation release %s", delegation_id) + if released: + with _records_lock: + if _records.get(delegation_id) is record: + del _records[delegation_id] + return { + "status": "rejected", "error_code": "executor_unavailable", + "execution_started": False, "delegation_id": delegation_id, + "durable_reservation_released": released, + "error": f"Async executor unavailable before submission: {exc}", + } def _worker() -> None: combined: Dict[str, Any] = {} @@ -1127,13 +1246,14 @@ def _worker() -> None: try: # Propagate the dispatching profile to the detached batch children. executor.submit(propagate_context_to_thread(_worker)) - except Exception as exc: # pragma: no cover - with _records_lock: - _records.pop(delegation_id, None) - _delete_durable_delegation(delegation_id) + except Exception as exc: + # submit() may enqueue before raising. Keep both ledgers and the + # handle; callers must never replay an ambiguous scheduling outcome. + logger.exception("Async submission outcome uncertain for %s", delegation_id) return { - "status": "rejected", - "error": f"Failed to schedule async delegation batch: {exc}", + "status": "dispatch_uncertain", "error_code": "scheduling_uncertain", + "execution_started": None, "delegation_id": delegation_id, + "error": f"Async submission outcome is unknown: {exc}", } if progress_fn is not None: _ensure_stale_monitor() @@ -1145,6 +1265,7 @@ def _worker() -> None: return {"status": "dispatched", "delegation_id": delegation_id} + def _finalize_batch( delegation_id: str, combined: Dict[str, Any], status: str ) -> None: diff --git a/tools/cronjob_tools.py b/tools/cronjob_tools.py index a938c7cca93c4..bb60dab0b1f2a 100644 --- a/tools/cronjob_tools.py +++ b/tools/cronjob_tools.py @@ -1171,9 +1171,15 @@ def _try_dispatch_background_run( except Exception: pass - # Same snapshot claim as _execute_job_now: carry the owner-bearing - # record into the run so terminal writes stay fenced by this owner. - claimed_job = claim_job_for_fire(job_id, return_job=True) + # Opt in to an ephemeral owner/field fence for releasing this claim + # only when dispatch positively proves that execution never started. + from cron.jobs import claim_job_for_fire_with_unstarted_receipt + + claim_result = claim_job_for_fire_with_unstarted_receipt(job_id) + if isinstance(claim_result, tuple) and len(claim_result) == 2: + claimed_job, unstarted_receipt = claim_result + else: + claimed_job, unstarted_receipt = None, None if not isinstance(claimed_job, dict): refreshed = get_job(job_id) if refreshed is None: @@ -1185,11 +1191,40 @@ def _try_dispatch_background_run( return {"claimed": False, "success": False, "error": reason} except Exception as e: logger.error("Failed to claim cron job %s for background run: %s", job_id, e) + # No runner started, but a failed acquisition cannot certify whether + # the owner published a claim. Never terminally mark or erase it here. + return { + "claimed": None, "dispatched": False, "success": False, + "status": "rejected", "error_code": "durable_storage_unavailable", + "execution_started": False, "error": str(e), + "claim_release_status": "write_uncertain", "claim_release_confirmed": False, + } + + def refuse_unstarted(dispatch_result: Dict[str, Any]) -> Dict[str, Any]: try: - mark_job_run(job_id, False, str(e)) - except Exception: - pass - return {"claimed": True, "dispatched": False, "success": False, "error": str(e)} + from cron.jobs import release_unstarted_fire_claim + + release = release_unstarted_fire_claim(unstarted_receipt) + release_status = release.get("status", "write_uncertain") + except Exception as exc: + logger.warning("Unstarted cron claim release failed for %s: %s", job_id, exc) + release_status = "write_uncertain" + confirmed = release_status == "released" + error = dispatch_result.get("error") or "Background dispatch refused before execution." + if not confirmed: + error += f" Claim release was not confirmed ({release_status})." + refusal = { + "claimed": True, "dispatched": False, "success": False, + "status": "rejected", + "error_code": dispatch_result.get("error_code") or "durable_storage_unavailable", + "execution_started": False, "error": error, + "claim_release_status": release_status, + "claim_release_confirmed": confirmed, + } + for field in ("delegation_id", "durable_reservation_released"): + if field in dispatch_result: + refusal[field] = dispatch_result[field] + return refusal origin_ui_session_id = "" try: @@ -1207,13 +1242,11 @@ def _try_dispatch_background_run( origin_session_id = _current_origin_session_id() except Exception as e: - logger.warning( - "cronjob run: async delegation registry unavailable (%s); " - "running job '%s' inline.", e, job_name, - ) - result = _run_claimed_job(claimed_job, extra_prompt=extra_prompt) - result["dispatched"] = False - return result + logger.warning("cronjob run: async delegation registry unavailable (%s)", e) + return refuse_unstarted({ + "status": "rejected", "error_code": "durable_storage_unavailable", + "execution_started": False, "error": str(e), + }) try: from tools.delegate_tool import _get_max_async_children @@ -1278,13 +1311,28 @@ def _runner() -> Dict[str, Any]: "delegation_id": dispatch.get("delegation_id"), } - # Pool at capacity (or submit failure): the claim is already taken and - # must not be stranded — run inline exactly as the legacy path did. + known_unstarted = ( + dispatch.get("status") == "rejected" + and dispatch.get("execution_started") is False + ) + if not known_unstarted: + # The runner may have been accepted. Preserve claim custody and do + # not retry inline or use the unstarted-only release receipt. + return { + "claimed": True, "dispatched": False, "success": False, + "status": "dispatch_uncertain", "error_code": "scheduling_uncertain", + "execution_started": None, "delegation_id": dispatch.get("delegation_id"), + "error": dispatch.get("error") or "Background dispatch outcome is uncertain.", + } + if dispatch.get("error_code") != "pool_capacity": + return refuse_unstarted(dispatch) + + # A typed pool refusal alone permits one owner-bearing inline run. logger.info( "cronjob run: background pool unavailable (%s); running job '%s' inline.", dispatch.get("error", "rejected"), job_name, ) - result = _run_claimed_job(job, extra_prompt=extra_prompt) + result = _run_claimed_job(claimed_job, extra_prompt=extra_prompt) result["dispatched"] = False return result @@ -1619,6 +1667,29 @@ def cronjob( bg = _try_dispatch_background_run( job, session_id=session_id, extra_prompt=extra_prompt ) + if bg is not None and bg.get("status") in {"rejected", "dispatch_uncertain"}: + result = _format_job(job) + result["executed"] = False if bg.get("execution_started") is False else None + response = { + "success": False, + "job": result, + "status": bg["status"], + "error_code": bg.get("error_code"), + "execution_started": bg.get("execution_started"), + "error": bg.get("error"), + } + for field in ( + "delegation_id", "durable_reservation_released", + "claim_release_status", "claim_release_confirmed", + ): + if field in bg: + response[field] = bg[field] + if bg["status"] == "dispatch_uncertain": + response["note"] = ( + "Work may already be running. Keep the delegation handle; " + "do not retry or run inline while submission remains uncertain." + ) + return json.dumps(response, indent=2) if bg is not None and bg.get("dispatched"): _notify_provider_jobs_changed_safe() result = _format_job(get_job(job_id) or {"id": job_id}) diff --git a/tools/delegate_tool.py b/tools/delegate_tool.py index d3c2a2dda8817..714bba6f6d4b1 100644 --- a/tools/delegate_tool.py +++ b/tools/delegate_tool.py @@ -4348,9 +4348,84 @@ def _batch_progress(): ) return json.dumps(payload, ensure_ascii=False) - # Pool at capacity / schedule failure — children are still attached - # (we detach above only on the parent list, but the async unit was - # never accepted, so re-attaching isn't needed: we just run inline). + known_unstarted = ( + dispatch.get("status") == "rejected" + and dispatch.get("execution_started") is False + ) + if not known_unstarted: + # Submit may already have accepted the runner. Its registry keeps + # custody; closing, draining steering or running inline could race it. + return json.dumps({ + "status": "dispatch_uncertain", + "error_code": "scheduling_uncertain", + "execution_started": None, + "mode": "background", + "delegation_id": dispatch.get("delegation_id") or live_deleg_id, + "count": len(_goals), + "goals": _goals, + "subagent_ids": [getattr(c, "_subagent_id", None) for c in _child_agents], + "error": dispatch.get("error") or "Background dispatch outcome is uncertain.", + "note": "Do not retry or run these children synchronously while dispatch is uncertain.", + }, ensure_ascii=False) + + if dispatch.get("error_code") != "pool_capacity": + # Admission positively refused execution. Close only this call's + # constructed children, preserving steering accepted before refusal. + rejected_results = [] + for index, task, child in children: + entry = { + "task_index": index, + "status": "rejected", + "error_code": dispatch.get("error_code") or "durable_storage_unavailable", + "execution_started": False, + "error": dispatch.get("error") or "Background dispatch refused before execution.", + "summary": None, + "api_calls": 0, + "duration_seconds": 0, + } + sid = getattr(child, "_subagent_id", None) + if sid: + try: + missed_steer = _close_subagent_steering(sid, child) + if missed_steer: + entry["missed_steer"] = missed_steer + except Exception: + logger.warning("Failed to close unstarted child steering: %s", sid, exc_info=True) + _unregister_subagent(sid, agent=child) + try: + if hasattr(child, "close"): + child.close() + except Exception as exc: + entry["cleanup_error"] = str(exc) + logger.warning("Failed to close unstarted child: %s", sid, exc_info=True) + writer = live_writers[index] if 0 <= index < len(live_writers) else None + if writer is not None: + try: + writer.finalize(entry) + except Exception: + logger.debug("Live transcript refusal finalize failed", exc_info=True) + if index < len(live_paths): + entry["live_transcript"] = live_paths[index] + rejected_results.append(entry) + try: + update_manifest_statuses(live_deleg_id, rejected_results) + except Exception: + logger.debug("Live transcript refusal manifest update failed", exc_info=True) + rejection_payload = { + "status": "rejected", + "error_code": dispatch.get("error_code") or "durable_storage_unavailable", + "execution_started": False, + "mode": "background", + "error": dispatch.get("error") or "Background dispatch refused before execution.", + "results": rejected_results, + "total_duration_seconds": round(time.monotonic() - overall_start, 2), + } + for field in ("delegation_id", "durable_reservation_released"): + if field in dispatch: + rejection_payload[field] = dispatch[field] + return json.dumps(rejection_payload, ensure_ascii=False) + + # Only a typed pool refusal proves inline fallback is permitted. logger.info( "delegate_task: async pool at capacity (%s); running the whole " "batch synchronously instead.", diff --git a/tools/file_tools.py b/tools/file_tools.py index 3f3730149fe92..ed8a8787e55a0 100644 --- a/tools/file_tools.py +++ b/tools/file_tools.py @@ -2158,32 +2158,58 @@ def _mark_verification_stale( session_id: str | None = None, ) -> None: """Best-effort note that successful edits made prior verification stale.""" - paths = [p for p in resolved_paths if p] + paths = list(dict.fromkeys(p for p in resolved_paths if p)) if not paths: return try: from agent.coding_context import project_facts_for from agent.verification_evidence import mark_workspace_edited - cwd = None - for path in paths: - try: - candidate = str(Path(path).parent) - except Exception: - continue - if project_facts_for(candidate): - cwd = candidate - break - if cwd is None: + except Exception: + logger.debug("verification stale marker failed", exc_info=True) + return + + grouped_paths: dict[str | None, list[str]] = {} + unresolved: list[str] = [] + for path in paths: + try: + candidate = str(Path(path).parent) + facts = project_facts_for(candidate) + except Exception: + facts = None + if facts: + root = str(facts.get("root") or candidate) + grouped_paths.setdefault(root, []).append(path) + else: + unresolved.append(path) + + if unresolved: + try: cwd = _authoritative_workspace_root(task_id) + except Exception: + cwd = None if cwd is None: try: - cwd = str(Path(paths[0]).parent) + cwd = str(Path(unresolved[0]).parent) except Exception: cwd = None - mark_workspace_edited(session_id=session_id or task_id, cwd=cwd, paths=paths) - except Exception: - logger.debug("verification stale marker failed", exc_info=True) + try: + facts = project_facts_for(cwd) + if facts: + cwd = str(facts.get("root") or cwd) + except Exception: + pass + if cwd is not None: + grouped_paths.setdefault(cwd, []).extend(unresolved) + + for cwd, changed_paths in grouped_paths.items(): + try: + mark_workspace_edited( + session_id=session_id or task_id, cwd=cwd, paths=changed_paths + ) + except Exception: + # One failed best-effort marker must not skip another workspace. + logger.debug("verification stale marker failed", exc_info=True) def _check_binary_document_write(filepath: str, task_id: str = "default") -> str | None: diff --git a/tools/mcp_tool.py b/tools/mcp_tool.py index a33710a5fe0e8..3aed173790a5b 100644 --- a/tools/mcp_tool.py +++ b/tools/mcp_tool.py @@ -114,8 +114,10 @@ from types import SimpleNamespace from typing import Callable from datetime import datetime +from pathlib import Path from typing import Any, Coroutine, Dict, List, Optional, Set, Tuple from urllib.parse import urlparse +from uuid import uuid4 from tools.registry import tool_error from tools.ansi_strip import strip_unicode_tags @@ -179,63 +181,157 @@ def _truncate_mcp_text_result(text: str, max_chars: int = _MCP_HARD_RESULT_CAP_C # the terminal while prompt_toolkit / Rich is rendering the TUI — which # corrupts the display and can hang the session. # -# Instead we redirect every stdio MCP subprocess's stderr into a shared -# per-profile log file (~/.hermes/logs/mcp-stderr.log), tagged with the -# server name so individual servers remain debuggable. +# Each stdio attempt owns a separate file under its config home's logs. +# The existing mcp-stderr.log remains an index, without interleaving child +# output from concurrent servers or borrowing another profile's handle. # # Fallback is os.devnull if opening the log file fails for any reason. -_mcp_stderr_log_fh: Optional[Any] = None +_mcp_stderr_log_files: Dict[str, Any] = {} _mcp_stderr_log_lock = threading.Lock() +_mcp_stdio_diagnostic: contextvars.ContextVar[Optional[dict]] = contextvars.ContextVar( + "mcp_stdio_diagnostic", default=None +) def _get_mcp_stderr_log() -> Any: - """Return a shared append-mode file handle for MCP subprocess stderr. + """Return this attempt's stderr file, or this config home's index. - Opened once per process and reused for every stdio server. Must have a - real OS-level file descriptor (``fileno()``) because asyncio's subprocess - machinery wires the child's stderr directly to that fd. Falls back to - ``/dev/null`` if opening the log file fails. + The no-argument seam is retained for existing stdio adapters. Attempt + handles are closed with the transport; only per-home index handles are + cached. Diagnostics must never prevent a connection attempt. """ - global _mcp_stderr_log_fh + capture = _mcp_stdio_diagnostic.get() + if capture is not None and "stream" in capture: + return capture["stream"] with _mcp_stderr_log_lock: - if _mcp_stderr_log_fh is not None: - return _mcp_stderr_log_fh + log_path = None + fh = None try: - from hermes_constants import get_hermes_home - log_dir = get_hermes_home() / "logs" - log_dir.mkdir(parents=True, exist_ok=True) - log_path = log_dir / "mcp-stderr.log" - # Line-buffered so server output lands on disk promptly; errors= - # "replace" tolerates garbled binary output from misbehaving - # servers. - fh = open(log_path, "a", encoding="utf-8", errors="replace", buffering=1) - # Sanity-check: confirm a real fd is available before we commit. + if capture is not None: + log_path = Path(capture["stderr_path"]) + else: + from hermes_constants import get_hermes_home + + log_path = get_hermes_home() / "logs" / "mcp-stderr.log" + cached = _mcp_stderr_log_files.get(str(log_path)) + if cached is not None and not cached.closed: + return cached + log_path.parent.mkdir(parents=True, exist_ok=True) + fh = open( + log_path, "x+" if capture is not None else "a", encoding="utf-8", + errors="replace", buffering=1, newline="" if capture is not None else None, + opener=lambda path, flags: os.open(path, flags, 0o600), + ) fh.fileno() - _mcp_stderr_log_fh = fh + destination = "file" except Exception as exc: # pragma: no cover — best-effort fallback - logger.debug("Failed to open MCP stderr log, using devnull: %s", exc) + if fh is not None: + try: + fh.close() + except Exception: + pass + logger.debug("Failed to open MCP stderr log: %s", type(exc).__name__) + try: + fh = open(os.devnull, "w", encoding="utf-8") + destination = "discarded" + except Exception: + fh = sys.stderr + destination = "parent_stderr" + if capture is not None: + capture["stream"] = fh + capture["destination"] = destination + elif log_path is not None: + _mcp_stderr_log_files[str(log_path)] = fh + return fh + + +def _stdio_diagnostic_fields(capture: dict) -> dict: + """Only ownership and lifecycle metadata; no argv, env or error text.""" + return {"kind": "mcp.stdio.attempt", **{key: capture.get(key) for key in ( + "attempt_id", "server", "config_home", "parent_pid", "stderr_path", + "destination", "phase", "status", "exception_type", + )}} + + +def _begin_stdio_diagnostic(server_name: str) -> dict: + capture = { + "attempt_id": uuid4().hex, "server": server_name, + "config_home": None, "parent_pid": os.getpid(), "stderr_path": None, + "phase": "transport", "status": "starting", + } + try: + from hermes_constants import get_hermes_home + + home = get_hermes_home() + capture["config_home"] = str(home) + capture["stderr_path"] = str(home / "logs" / "mcp-stderr" / f"{capture['attempt_id']}.log") + except Exception: + pass # Unavailable ownership is explicit; capture still fails softly. + capture["token"] = _mcp_stdio_diagnostic.set(capture) + capture["stream"] = _get_mcp_stderr_log() + _write_stderr_log_header(server_name) + logger.debug("MCP stdio attempt: %s", json.dumps(_stdio_diagnostic_fields(capture))) + return capture + + +def _finish_stdio_diagnostic(capture: dict) -> None: + try: + record = json.dumps(_stdio_diagnostic_fields(capture)) + fh = capture.get("stream") + try: + if fh is not None: + boundary = "\n" + if capture.get("destination") == "file": + # The transport has unwound. Inspect its last byte through + # the same owned handle; child stderr may omit a newline. + fh.flush() + fh.seek(0, os.SEEK_END) + end = fh.tell() + if end: + fh.seek(end - 1) + boundary = "" if fh.read(1) == "\n" else "\n" + else: + boundary = "" + fh.seek(0, os.SEEK_END) + fh.write(boundary + record + "\n") + fh.flush() + except Exception: + pass + if capture["status"] != "closed" or capture.get("destination") != "file": + logger.warning("MCP stdio attempt ended: %s", record) + finally: + _mcp_stdio_diagnostic.reset(capture["token"]) + fh = capture.get("stream") + if fh is not None and fh is not sys.stderr: try: - _mcp_stderr_log_fh = open(os.devnull, "w", encoding="utf-8") + fh.close() except Exception: - # Last resort: the real stderr. Not ideal for TUI users but - # it matches pre-fix behavior. - _mcp_stderr_log_fh = sys.stderr - return _mcp_stderr_log_fh + pass def _write_stderr_log_header(server_name: str) -> None: """Write a human-readable session marker before launching a server. - Gives operators a way to find each server's output in the shared - ``mcp-stderr.log`` file without needing per-line prefixes (which would - require a pipe + reader thread and complicate shutdown). + The per-home index points at an attempt-owned file, so concurrent child + output can be attributed without changing the MCP protocol streams. """ fh = _get_mcp_stderr_log() try: ts = datetime.now().strftime("%Y-%m-%d %H:%M:%S") fh.write(f"\n===== [{ts}] starting MCP server '{server_name}' =====\n") fh.flush() + capture = _mcp_stdio_diagnostic.get() + if capture is not None: + record = json.dumps(_stdio_diagnostic_fields(capture)) + fh.write(record + "\n") + token = _mcp_stdio_diagnostic.set(None) + try: + index = _get_mcp_stderr_log() + finally: + _mcp_stdio_diagnostic.reset(token) + index.write(record + "\n") + index.flush() except Exception: pass @@ -3280,12 +3376,10 @@ async def _run_stdio(self, config: dict): # Snapshot child PIDs before spawning so we can track the new one. pids_before = _snapshot_child_pids() new_pids: set = set() - # Redirect subprocess stderr into a shared log file so MCP servers - # (FastMCP banners, slack-mcp startup JSON, etc.) don't dump onto - # the user's TTY and corrupt the TUI. Preserves debuggability via - # ~/.hermes/logs/mcp-stderr.log. - _write_stderr_log_header(self.name) - _errlog = _get_mcp_stderr_log() + # Each attempt keeps its child stderr and config-home ownership. + # The per-home mcp-stderr.log is an index of those files. + _diagnostic = _begin_stdio_diagnostic(self.name) + _errlog = _diagnostic["stream"] try: async with stdio_client(server_params, errlog=_errlog) as ( read_stream, @@ -3356,13 +3450,16 @@ async def _run_stdio(self, config: dict): connect_timeout = float( config.get("connect_timeout", _DEFAULT_CONNECT_TIMEOUT) ) + _diagnostic["phase"] = "negotiate" self.initialize_result = await self._negotiate_session( session, connect_timeout ) self.session = session self._mark_lifecycle_started() + _diagnostic["phase"] = "discover_tools" await self._discover_tools() self._ready.set() + _diagnostic["phase"] = "active" self._ever_connected = True # Session is live again: clear any breaker state from a # prior outage so the first call after recovery isn't @@ -3378,7 +3475,17 @@ async def _run_stdio(self, config: dict): # _reconnect_event (e.g. future manual /mcp refresh) for # consistency with _run_http. return await self._wait_for_lifecycle_event() + except BaseException as exc: + _diagnostic["status"] = "failed" + _diagnostic["exception_type"] = type(exc).__name__ + raise finally: + if _diagnostic["status"] == "starting": + _diagnostic["status"] = "closed" + try: + _finish_stdio_diagnostic(_diagnostic) + except Exception: + logger.debug("Failed to finalize MCP stdio diagnostics", exc_info=False) # Runs on clean exit, exceptions, AND asyncio cancellation. # If any of the spawned PIDs are still alive, the SDK's # teardown failed (common when the task is cancelled mid-way @@ -4020,6 +4127,9 @@ async def run(self, config: dict): while True: try: + # Readiness belongs to this transport attempt. Exception + # retries and lazy stdio revival also rebuild the session. + self._ready.clear() if self._is_http(): lifecycle_reason = await self._run_http(config) else: @@ -8001,7 +8111,11 @@ def get_mcp_status() -> List[dict]: transport = cfg.get("transport", "http") if "url" in cfg else "stdio" enabled = _parse_boolish(cfg.get("enabled", True), default=True) server = active_servers.get(name) - if server and server.session is not None: + # A session is assigned after initialize, before tools/list completes. + # The task's ready event and error state own readiness, including on + # reconnect; a transport object alone must not advertise success. + if (server and server.session is not None and server._ready.is_set() + and server._error is None): entry = { "name": name, "transport": transport, @@ -8025,7 +8139,8 @@ def get_mcp_status() -> List[dict]: "disabled": True, "status": "disabled", }) - elif name in connecting: + elif name in connecting or (server and not server._ready.is_set() + and server._error is None): result.append({ "name": name, "transport": transport, @@ -8034,7 +8149,7 @@ def get_mcp_status() -> List[dict]: "disabled": False, "status": "connecting", }) - elif name in connect_errors: + elif name in connect_errors or (server and server._error is not None): result.append({ "name": name, "transport": transport, @@ -8042,7 +8157,7 @@ def get_mcp_status() -> List[dict]: "connected": False, "disabled": False, "status": "failed", - "error": connect_errors[name], + "error": connect_errors.get(name) or _format_connect_error(server._error), }) else: result.append({ @@ -8332,7 +8447,10 @@ def _add(schema: dict) -> bool: staged_engine_names: set = set() try: enabled = getattr(agent, "enabled_toolsets", None) - context_engine_allowed = enabled is None or "context_engine" in enabled + context_engine_allowed = ( + (enabled is None or "context_engine" in enabled) + and "context_engine" not in (getattr(agent, "disabled_toolsets", None) or []) + ) compressor = getattr(agent, "context_compressor", None) get_schemas = getattr(compressor, "get_tool_schemas", None) if compressor else None if context_engine_allowed and callable(get_schemas): diff --git a/tui_gateway/methods_prompt.py b/tui_gateway/methods_prompt.py index ba7d4cea543df..f5d587c7f752b 100644 --- a/tui_gateway/methods_prompt.py +++ b/tui_gateway/methods_prompt.py @@ -370,14 +370,49 @@ def _(rid, params: dict) -> dict: if (t := current_transport()) is not None: session["transport"] = t input_receipt = None + pending_input_admission = None + with session["history_lock"]: + input_agent = session.get("agent") + input_transport = session.get("transport") + def input_owner_error(): + # Caller holds history_lock; slow pending resolution grants no new owner. + current_window = session.get("_turn_outcomes") + current_nonce = (current_window.turns[-1]["accepted_turn"]["request_id"] + if isinstance(current_window, TurnOutcomeWindow) and current_window.owns(session, sid) and current_window.turns else None) + if (_sessions.get(sid) is not session or session.get("agent") is not input_agent + or session.get("transport") is not input_transport + or (pending_input_admission is not None and ( + int(session.get("_queued_prompt_generation", 0)) != pending_input_admission[0] + or (session.get("inflight_turn") or {}).get("started_at") is not pending_input_admission[1] + or current_window is not pending_input_admission[2] + or current_nonce != pending_input_admission[3]))): + return _err(rid, 5001, "session owner changed; input not started", + {"execution_started": False, "durable_input_accepted": input_receipt is not None, + "error_surface": {"layer": "runtime", "code": "session_owner_changed", "retryable": True}, + **({"input_event_id": input_receipt.event_id} if input_receipt is not None else {})}) + return None stop_input_deadline = time.monotonic() + 5.0 while True: if input_receipt is None: stop_error = _wait_for_stop_input_admission(rid, session, deadline=stop_input_deadline) if stop_error is not None: return stop_error + with session["history_lock"]: + failed_once = _one_turn_model_restore_error(sid, session) is not None + if failed_once: + window = session.get("_turn_outcomes") + pending_input_admission = (int(session.get("_queued_prompt_generation", 0)), + (session.get("inflight_turn") or {}).get("started_at"), window, + window.turns[-1]["accepted_turn"]["request_id"] + if isinstance(window, TurnOutcomeWindow) and window.owns(session, sid) and window.turns else None) + if failed_once: + # Only an already acknowledged, eligible explicit pick may + # settle failed custody. Resolve outside the admission lock. + _apply_pending_model_switch(sid, session, before_admission=True) busy_transport = None with session["history_lock"]: + if (owner_error := input_owner_error()) is not None: + return owner_error if input_receipt is None and (session.get("_stop_pending") or session.get("_stop_uncertain")): # Stop may have begun between the wait and lock acquisition. # Retry outside the lock, before any durable input write. @@ -385,6 +420,12 @@ def _(rid, params: dict) -> dict: return _err(rid, 5032, "Stop settlement timed out; durable input not accepted", {"durable_input_accepted": False}) continue + if (restore_error := _one_turn_model_restore_error(sid, session)) is not None: + return _err(rid, 5001, + "The saved model was not restored; choose a model explicitly before sending.", + {"error_surface": restore_error, "execution_started": False, + "durable_input_accepted": input_receipt is not None, + **({"input_event_id": input_receipt.event_id} if input_receipt is not None else {})}) # Refusals must precede durable acceptance, including busy ACKs. # A watch child's run belongs to its parent, so running alone # cannot establish that this session is available for input. @@ -433,6 +474,14 @@ def _(rid, params: dict) -> dict: else None ) with session["history_lock"]: + if (owner_error := input_owner_error()) is not None: + return owner_error + if (restore_error := _one_turn_model_restore_error(sid, session)) is not None: + return _err(rid, 5001, + "The saved model was not restored; choose a model explicitly before sending.", + {"error_surface": restore_error, "execution_started": False, + "durable_input_accepted": input_receipt is not None, + **({"input_event_id": input_receipt.event_id} if input_receipt is not None else {})}) # A watch session's run lives in the PARENT turn, so its own running # flag is False — without this, typing mid-run builds a second agent # racing the in-flight child on the same stored session (interleaved @@ -969,7 +1018,11 @@ def run_after_agent_ready() -> None: window.finish(outcome_execution[2], interrupted=bool(owns_current and session.get("_turn_cancel_requested"))) _turn_outcome_execution.reset(token) - def run_owned_after_agent_ready() -> None: + def run_owned_after_agent_ready( + admitted_started_at=(session.get("inflight_turn") or {}).get("started_at"), + admitted_generation=int(session.get("_queued_prompt_generation", 0)), + admitted_transport=session.get("transport"), + ) -> None: # Patient wait (#63078): the user's message is already the accepted # in-flight turn, so a slow deferred build must not eat it. The wait # delivers the prompt when the still-running build completes, honors a @@ -980,41 +1033,62 @@ def run_owned_after_agent_ready() -> None: # Terminal frame + retained snapshot (not a bare "error" event + # cleared inflight): if the client is disconnected right now, the # retained snapshot is the only way resume can show this failure. - _emit_terminal_turn_error( + settled = _emit_terminal_turn_error( sid, session, (err.get("error") or {}).get("message", "agent initialization failed"), # Agent construction never reached the provider: this is a # local-runtime failure (env/config/venv), not an API error. error_surface={"layer": "runtime", "code": "agent_init_failed", "retryable": True}, + expected_started_at=admitted_started_at, + expected_queue_generation=admitted_generation, + settle_running=True, + terminal_transport=admitted_transport, ) - with session["history_lock"]: - session["running"] = False - session["last_active"] = time.time() - _emit("session.info", sid, _session_info(session.get("agent"), session)) - return + if settled: + return with session["history_lock"]: + window = session.get("_turn_outcomes") + inflight = session.get("inflight_turn") + cancelled = bool(session.get("_turn_cancel_requested")) + generation = int(session.get("_queued_prompt_generation", 0)) + if (_sessions.get(sid) is not session or session.get("_closing") + or session.get("_finalized") or admitted_started_at is None + or not isinstance(inflight, dict) or inflight.get("started_at") is not admitted_started_at + or not isinstance(window, TurnOutcomeWindow) or not window.owns(session, sid) + or not window.turns + or window.turns[-1]["accepted_turn"]["request_id"] != outcome_execution[2] + or (generation != admitted_generation and not ( + cancelled and generation == session.get("_last_stop_queue_generation")))): + return if session.get("_turn_cancel_requested") or not session.get("running"): session["running"] = False _clear_inflight_turn(session) + window.finish(outcome_execution[2], interrupted=cancelled) # Surface the cancellation to the client. Without this emit the # turn vanishes silently — the Desktop sees `prompt.submit` # return `{"status": "streaming"}` but never receives a # `message.start` or `error` event, so the composer shows no # feedback (issue #63078 server-side half). Match the # `_wait_agent` error branch above: emit, then bail. - _emit( - "error", - sid, - { - "message": "Turn cancelled before the agent was ready" - if session.get("_turn_cancel_requested") - else "Session no longer running before the agent was ready" - }, - ) + transport_token = bind_transport(admitted_transport) + try: + _emit( + "error", + sid, + { + "message": "Turn cancelled before the agent was ready" + if cancelled + else "Session no longer running before the agent was ready" + }, + ) + finally: + reset_transport(transport_token) + return + if err: return _run_prompt_submit(rid, sid, session, text, display_kind=display_kind, - **({"context_input_event_id": input_receipt.event_id} if input_receipt is not None else {})) + turn_transport=admitted_transport, **({"context_input_event_id": input_receipt.event_id} if input_receipt is not None else {})) run_thread = threading.Thread(target=run_after_agent_ready, daemon=True) run_thread._turn_outcome_execution = outcome_execution diff --git a/tui_gateway/server.py b/tui_gateway/server.py index e06205106602b..7521c35cbe602 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -2611,13 +2611,45 @@ def _turn_outcomes_snapshot(sid, session): def _emit(event: str, sid: str, payload: dict | None = None): session = _sessions.get(sid) if event == "message.complete" else None - execution = _turn_outcome_execution.get() if session is not None else None + execution = _turn_outcome_execution.get() + owned_event = execution is not None and execution[1] == sid and event in { + "message.complete", "error", "session.info", + } + transport = None + if owned_event: + captured = execution[0] + current = _sessions.get(sid) + ready = captured.get("resume_history_ready") + failed_history = (ready is not None and ready.is_set() + and captured.get("resume_history_error") + and captured.get("agent_error") == captured.get("resume_history_error") + and not captured.get("_compute_host_active")) + # The real hydrator deliberately removes a failed inline record. Its + # admitted terminal still belongs to the captured client, never to a + # replacement record under the same runtime id. + transport = current_transport() or captured.get("transport") or _stdio_transport + detached_failure = ( + current is None and event == "message.complete" and failed_history + and not captured.get("_closing") and not captured.get("_finalized") + and not captured.get("_turn_cancel_requested") + ) + if current is captured or detached_failure: + session = captured + elif event == "message.complete": + # Keep legacy terminal visibility after owner loss, without a + # reply projection or metadata borrowed from the replacement. + session = None + else: + return # Execution correlation is captured before starting the worker. A late old # emitter must never borrow a newer request from the mutable session record. - if execution is not None and execution[0] is session and execution[1] == sid: + if event == "message.complete" and execution is not None and execution[0] is session and execution[1] == sid: with session["history_lock"]: window = session.get("_turn_outcomes") - if (_sessions.get(sid) is session and not session.get("_closing") + current = _sessions.get(sid) + if ((current is session or (current is None and owned_event and failed_history)) + and not session.get("_closing") + and not session.get("_finalized") and isinstance(window, TurnOutcomeWindow) and window.owns(session, sid)): try: ref = window.append(execution[2], payload or {}) @@ -2628,11 +2660,20 @@ def _emit(event: str, sid: str, payload: dict | None = None): logger.debug("finalized reply projection unavailable", exc_info=True) if ref is not None: payload = {**(payload or {}), "request_id": ref["request_id"], "accepted_turn": ref} - if session is not None and session.get("_host_turn_request_id") and _inside_compute_host_child(): + if event == "message.complete" and session is not None and session.get("_host_turn_request_id") and _inside_compute_host_child(): # The bubble is complete; only the host's matching turn terminal can # retire the accepted request (including any chained goal work). payload = {**(payload or {}), "chain_pending": True} - write_json(_event_frame(event, sid, payload)) + frame = _event_frame(event, sid, payload) + if owned_event: + # Keep the canonical stamp/replay path, but do not let write_json's + # registry precedence redirect a captured terminal to a new owner. + from tui_gateway.event_replay import _stamp_event + + _stamp_event(frame) + (transport or _stdio_transport).write(frame) + else: + write_json(frame) # Live client transports, one per connected WS peer (maintained by tui_gateway.ws). @@ -3597,24 +3638,60 @@ def _wait_agent_for_prompt(session: dict, rid: str, sid: str) -> dict | None: (``agent.build_wait_timeout``, default 600s — no infinite waits) expired on a genuinely hung build. - Returns ``None`` on success OR when the turn was cancelled mid-wait (the - caller's cancel branch owns that messaging), an ``_err`` dict otherwise. + Resumed prompts also require their real history completion event. Agent + readiness alone cannot admit a provider turn without persisted context. + Returns ``None`` on success OR cancellation/owner loss (the caller owns + that messaging), an ``_err`` dict otherwise. """ ready = session.get("agent_ready") - if ready is None: + history_ready = session.get("resume_history_ready") + gates = [gate for gate in (ready, history_ready) if gate is not None] + if not gates: return None + execution = _turn_outcome_execution.get() + if execution is not None and (execution[0] is not session or execution[1] != sid): + execution = None + registered = _sessions.get(sid) is session + with session["history_lock"]: + inflight = session.get("inflight_turn") + admission = inflight.get("started_at") if isinstance(inflight, dict) else None + generation = int(session.get("_queued_prompt_generation", 0)) start = time.monotonic() cap = _agent_build_wait_cap() notified_slow = False - while not ready.wait(timeout=_AGENT_BUILD_WAIT_SLICE): + while True: with session["history_lock"]: - cancelled = session.get("_turn_cancel_requested") or not session.get( + current = _sessions.get(sid) + window = session.get("_turn_outcomes") + inflight = session.get("inflight_turn") + if ( + (registered or history_ready is not None or execution is not None) + and current is not session + and not (current is None and history_ready is not None + and history_ready.is_set() and session.get("resume_history_error") + and session.get("agent_error") == session.get("resume_history_error") + and not session.get("_compute_host_active")) + ) or int(session.get("_queued_prompt_generation", 0)) != generation: + return None + if admission is not None and ( + not isinstance(inflight, dict) or inflight.get("started_at") is not admission + ): + return None + if execution is not None and execution[2] is not None and ( + not isinstance(window, TurnOutcomeWindow) or not window.owns(session, sid) + or not window.turns + or window.turns[-1]["accepted_turn"]["request_id"] != execution[2] + ): + return None + cancelled = session.get("_closing") or session.get("_finalized") or session.get("_turn_cancel_requested") or not session.get( "running" ) if cancelled: # The caller's cancel/not-running branch emits the user-visible # event for this — bail without an error of our own. return None + if all(gate.is_set() for gate in gates): + break waited = time.monotonic() - start if waited >= cap: return _err( @@ -3627,6 +3704,7 @@ def _wait_agent_for_prompt(session: dict, rid: str, sid: str) -> dict | None: if ( build_thread is not None and not build_thread.is_alive() + and ready is not None and not ready.is_set() ): # _build's ``finally`` guarantees ready.set(); a dead thread with @@ -3659,9 +3737,13 @@ def _wait_agent_for_prompt(session: dict, rid: str, sid: str) -> dict | None: "id": _AGENT_BUILD_SLOW_NOTICE_KEY, }, ) + wait_gate = next((gate for gate in gates if not gate.is_set()), None) + if wait_gate is None: + continue + wait_gate.wait(timeout=min(_AGENT_BUILD_WAIT_SLICE, max(0.0, cap - waited))) if notified_slow: _emit("notification.clear", sid, {"key": _AGENT_BUILD_SLOW_NOTICE_KEY}) - err = session.get("agent_error") + err = session.get("resume_history_error") or session.get("agent_error") return _err(rid, 5032, err) if err else None @@ -3672,12 +3754,10 @@ def _wait_agent_for_prompt(session: dict, rid: str, sid: str) -> dict | None: def _await_resume_history(current: dict, sid: str, key: str) -> str: """Wait for cold-resume history without holding agent init hostage. - The transcript is only display state; agent construction needs the - session's context/db/secrets/MCP, not the history rows. A slow hydration - (SQLite contention while several sessions resume at once) previously - killed agent init with TimeoutError after 300s. Now: bounded wait, a - visible slow+grace window, then degrade to a live agent with an empty - history — late hydration still fills in under history_lock. + Agent construction may proceed after the grace window, but history is + also model context. Timeout leaves resume_history_ready to the real + hydration owner; executing prompts independently await completion with + cancellation and the existing finite build-wait cap. Returns "ready" | "degraded" | "vanish" (session replaced mid-wait). """ history_ready = current.get("resume_history_ready") @@ -3701,19 +3781,15 @@ def _await_resume_history(current: dict, sid: str, key: str) -> str: if _sessions.get(sid) is not current: return "vanish" logger.warning( - "resume hydration degraded; starting agent without history runtime=%s stored=%s profile=%s stage=%s", + "resume hydration degraded; agent setup may proceed while turns await history runtime=%s stored=%s profile=%s stage=%s", sid, key, Path(profile_home).name if profile_home else "default", stage, ) _emit( "session.resume_progress", sid, {"phase": "history", "status": "degraded_timeout", "stage": stage, - "message": f"history still loading ({stage}); agent started without it"}, + "message": f"history still loading ({stage}); turns wait for complete history"}, ) - with current["history_lock"]: - current["resume_hydrating"] = False - current.setdefault("history", []) - history_ready.set() return "degraded" @@ -6202,7 +6278,9 @@ def _is_pivot_marker(entry: Any) -> bool: return isinstance(entry, dict) and entry.get("display_kind") == "personality_switch" -def _append_model_switch_marker(session: dict | None, *, model: str, provider: str) -> None: +def _append_model_switch_marker( + session: dict | None, *, model: str, provider: str, history_lock_held: bool = False, +) -> None: """Record a real system-history pivot after a live model switch. Only the most recent marker is kept: each new switch first strips any @@ -6238,7 +6316,7 @@ def _replace_markers() -> None: session["history_version"] = int(session.get("history_version", 0)) + 1 lock = session.get("history_lock") - if lock is not None: + if lock is not None and not history_lock_held: with lock: _replace_markers() else: @@ -6593,8 +6671,6 @@ def _load_enabled_toolsets(platform: str | None = None) -> list[str] | None: enabled = _get_platform_tools(cfg, "cli", include_default_mcp_servers=True) if fallback_notice is not None: print(fallback_notice, file=sys.stderr, flush=True) - if not enabled: - return None # The client-surface toolsets are off _HERMES_CORE_TOOLS (every other # platform would carry their schema for nothing), so the platform # recovery above — which keys off hermes-cli's tool universe — can't @@ -6683,13 +6759,19 @@ def _persist_model_switch(result) -> None: def _snapshot_agent_model_runtime(agent) -> dict: """Capture the current agent model runtime for a one-turn restore.""" + reasoning = copy.deepcopy(getattr(agent, "reasoning_config", None)) + primary = copy.deepcopy(getattr(agent, "_primary_runtime", None)) + if isinstance(primary, dict): + # Session reasoning can differ from an older primary snapshot. + primary["reasoning_config"] = copy.deepcopy(reasoning) return { "model": getattr(agent, "model", ""), "provider": getattr(agent, "provider", ""), "api_key": getattr(agent, "api_key", ""), "base_url": getattr(agent, "base_url", ""), "api_mode": getattr(agent, "api_mode", ""), - "primary_runtime": copy.deepcopy(getattr(agent, "_primary_runtime", None)), + "primary_runtime": primary, + "reasoning_config": reasoning, } @@ -6707,19 +6789,37 @@ def _owns_one_turn_model_runtime(session, agent, runtime=None) -> bool: ) +def _one_turn_model_restore_error(sid: str, session: dict) -> dict | None: + """Read failed once custody; callers hold the existing history lock.""" + runtime = session.get("_one_turn_model_runtime") + if (not isinstance(runtime, dict) + or (not runtime.get("restore_failed") and runtime.get("agent") is session.get("agent")) + or runtime.get("session") is not session or runtime.get("sid") != sid + or _sessions.get(sid) is not session): + return None + return {"layer": "runtime", "code": "one_turn_model_restore_failed", "retryable": False} + + def _consume_one_turn_model_runtime(session, agent): with session["history_lock"]: + runtime = session.get("_one_turn_model_runtime") + if isinstance(runtime, dict) and runtime.get("restore_failed"): + raise RuntimeError("The saved model was not restored; choose a model explicitly before sending.") snapshot = session.get("one_turn_model_restore") if not snapshot: return None, None runtime = session.get("_one_turn_model_runtime") if not _owns_one_turn_model_runtime(session, agent, runtime): - session.pop("one_turn_model_restore", None) - session.pop("_one_turn_model_runtime", None) + if (isinstance(runtime, dict) and runtime.get("session") is session + and _sessions.get(runtime.get("sid")) is session): + runtime.setdefault("restore_snapshot", snapshot) + runtime["restore_failed"] = True + runtime["active"] = False raise RuntimeError("One-turn model selection changed owner before execution") # Publish active ownership before removing the queued snapshot. Metadata # stays truthful across consumption and through the finally restore. runtime["active"] = True + runtime.setdefault("restore_snapshot", snapshot) session.pop("one_turn_model_restore", None) return snapshot, runtime @@ -6728,6 +6828,16 @@ def _restore_agent_model_runtime(agent, snapshot: dict | None) -> None: """Restore an agent model runtime captured before a one-turn override.""" if not snapshot or agent is None: return + + def restore_reasoning(): + # Missing keys retain legacy behavior; an explicit None clears the + # temporary setting only after the saved route has restored. + if "reasoning_config" in snapshot: + agent.reasoning_config = copy.deepcopy(snapshot["reasoning_config"]) + primary = getattr(agent, "_primary_runtime", None) + if isinstance(primary, dict): + primary["reasoning_config"] = copy.deepcopy(agent.reasoning_config) + primary = snapshot.get("primary_runtime") if primary and hasattr(agent, "_restore_primary_runtime"): try: @@ -6735,6 +6845,7 @@ def _restore_agent_model_runtime(agent, snapshot: dict | None) -> None: agent._fallback_activated = True agent._rate_limited_until = 0 if agent._restore_primary_runtime(): + restore_reasoning() return except Exception: logger.debug("TUI one-turn model restore via primary runtime failed", exc_info=True) @@ -6746,6 +6857,7 @@ def _restore_agent_model_runtime(agent, snapshot: dict | None) -> None: base_url=snapshot.get("base_url", ""), api_mode=snapshot.get("api_mode", ""), ) + restore_reasoning() def _apply_model_switch( @@ -6757,6 +6869,10 @@ def _apply_model_switch( pin_session_override: bool = True, parsed_flags: Any | None = None, persist_override: bool | None = None, + defer_if_running: bool = False, + supersede_pending: bool | None = None, + explicit_model_intent: bool = False, + expected_admission: tuple | None = None, ) -> dict: from hermes_cli.model_switch import ( parse_model_switch_args, @@ -6797,6 +6913,13 @@ def _apply_model_switch( raise ValueError("model value required") agent = session.get("agent") + owner_transport = session.get("transport") + # Existing restore/adoption/MoA callers explicitly suppress persistence; + # the picker marks manual intent, and the pending consumer opts out. + # A direct user /model choice otherwise keeps the pinning default. + if supersede_pending is None: + supersede_pending = pin_session_override and persist_override is None + superseded_pending = session.get("pending_model_switch") if supersede_pending else None if one_turn and not agent: raise ValueError("/model --once requires a live session") if agent: @@ -6855,8 +6978,6 @@ def _apply_model_switch( if not result.success: raise ValueError(result.error_message or "model switch failed") - restore_snapshot = _snapshot_agent_model_runtime(agent) if (one_turn and agent) else None - if agent: try: from hermes_cli.context_switch_guard import merge_preflight_compression_warning @@ -6903,76 +7024,242 @@ def _apply_model_switch( "confirm_message": confirm_msg, } - if agent: - try: - agent.switch_model( - new_model=result.new_model, - new_provider=result.target_provider, - api_key=result.api_key, - base_url=result.base_url, - api_mode=result.api_mode, + model_commit_lock = session.get("history_lock") + pending_publication = None + failed_once = None + next_once = None + def owns_expected_admission(): + if expected_admission is None: + return True + expected_agent, expected_transport, generation, started_at, expected_window, nonce = expected_admission + current = session.get("inflight_turn") + window = session.get("_turn_outcomes") + current_nonce = (window.turns[-1]["accepted_turn"]["request_id"] + if isinstance(window, TurnOutcomeWindow) and window.owns(session, sid) and window.turns else None) + return (session.get("agent") is expected_agent and session.get("transport") is expected_transport + and int(session.get("_queued_prompt_generation", 0)) == generation + and (current.get("started_at") if isinstance(current, dict) else None) is started_at + and window is expected_window and current_nonce == nonce) + with model_commit_lock if model_commit_lock is not None else contextlib.nullcontext(): + if defer_if_running or explicit_model_intent: + if (_sessions.get(sid) is not session or session.get("agent") is not agent + or session.get("transport") is not owner_transport): + raise ValueError("session owner changed; model request not applied") + if not owns_expected_admission(): + raise ValueError("session admission changed; model request not applied") + if defer_if_running: + if session.get("running"): + session["pending_model_switch"] = { + "raw": raw_input, + "confirm_expensive_model": confirm_expensive_model, + "display_model": result.new_model, + "display_provider": result.target_provider, + "after_inflight_turn": session.get("inflight_turn"), + } + return { + "value": result.new_model, "warning": result.warning_message or "", + "confirm_required": False, "confirm_message": "", "deferred": True, + "scope": "once" if one_turn else ("global" if persist_global else "session"), + } + if _one_turn_model_restore_error(sid, session) is not None: + if not explicit_model_intent: + raise ValueError("The saved model was not restored; choose a model explicitly before sending.") + failed_once = session["_one_turn_model_runtime"] + if one_turn and not failed_once.get("restore_snapshot"): + raise ValueError("The original one-turn restore target is unavailable; choose a session model explicitly.") + publication_generation = int(session.get("_queued_prompt_generation", 0)) + publication_inflight = session.get("inflight_turn") + publication_started_at = (publication_inflight.get("started_at") + if isinstance(publication_inflight, dict) else None) + restore_snapshot = None + if one_turn and agent: + queued_restore = session.get("one_turn_model_restore") + queued_runtime = session.get("_one_turn_model_runtime") + if failed_once is not None: + restore_snapshot = failed_once["restore_snapshot"] + elif (queued_restore and _owns_one_turn_model_runtime(session, agent, queued_runtime) + and not queued_runtime.get("active")): + # Replacing an unused once choice changes its one eligible + # turn, not the original durable/runtime restoration target. + restore_snapshot = queued_restore + else: + restore_snapshot = _snapshot_agent_model_runtime(agent) + + if agent: + try: + agent.switch_model( + new_model=result.new_model, + new_provider=result.target_provider, + api_key=result.api_key, + base_url=result.base_url, + api_mode=result.api_mode, + ) + except Exception as exc: + # The in-place swap rolled the agent back to the old working + # model/client and re-raised. Abort the commit: do NOT restart the + # slash worker, persist runtime, append the switch marker, set a + # session model_override, or persist to config — all of which would + # otherwise leave the session pinned to a broken model and kill the + # conversation on the next turn (#50163). A failed switch is a + # no-op; surface a clean error to the client. + logger.warning("In-place model switch failed for TUI agent: %s", exc) + raise ValueError( + f"Model switch to {result.new_model} failed ({exc}); " + f"staying on {getattr(agent, 'model', current_model)}." + ) from exc + _persist_live_session_runtime(session) + _persist_live_session_system_prompt(session) + _append_model_switch_marker( + session, model=result.new_model, provider=result.target_provider, + history_lock_held=model_commit_lock is not None, ) - except Exception as exc: - # The in-place swap rolled the agent back to the old working - # model/client and re-raised. Abort the commit: do NOT restart the - # slash worker, persist runtime, append the switch marker, set a - # session model_override, or persist to config — all of which would - # otherwise leave the session pinned to a broken model and kill the - # conversation on the next turn (#50163). A failed switch is a - # no-op; surface a clean error to the client. - logger.warning("In-place model switch failed for TUI agent: %s", exc) - raise ValueError( - f"Model switch to {result.new_model} failed ({exc}); " - f"staying on {getattr(agent, 'model', current_model)}." - ) from exc - _restart_slash_worker(sid, session) - _persist_live_session_runtime(session) - _persist_live_session_system_prompt(session) - _append_model_switch_marker( - session, model=result.new_model, provider=result.target_provider - ) - # The turn consumer uses this lock: snapshot and ownership must become - # visible (or retire) together, never as a partially published lease. - with session["history_lock"]: + # The turn consumer uses this lock: snapshot and ownership must become + # visible (or retire) together, never as a partially published lease. if one_turn: - session["one_turn_model_restore"] = restore_snapshot - session["_one_turn_model_runtime"] = { + next_once = { "session": session, "agent": agent, "sid": sid, "active": False, + "restore_snapshot": restore_snapshot, } - else: + if failed_once is None: + session["one_turn_model_restore"] = restore_snapshot + session["_one_turn_model_runtime"] = next_once + elif failed_once is None: session.pop("one_turn_model_restore", None) session.pop("_one_turn_model_runtime", None) + if failed_once is not None: + # This lease remains a refusal until this exact explicit + # choice finishes publication. A later choice owns its suffix. + failed_once["superseding_intent"] = result + + # Record the switch as a PER-SESSION override so a later rebuild of THIS + # session (e.g. /new via _reset_session_agent, or resume) re-derives the + # user's chosen model/provider instead of falling back to global config. + # + # We deliberately do NOT write process-global env vars (HERMES_MODEL / + # HERMES_INFERENCE_MODEL / HERMES_TUI_PROVIDER / HERMES_INFERENCE_PROVIDER) + # here. The desktop backend hosts every same-profile session in ONE process, + # so mutating os.environ on a /model switch leaked the new model/provider + # into every OTHER live session's next agent rebuild — switching the model + # in one session silently changed it in the others (the cross-session + # contamination bug). agent.switch_model() above already mutated the right + # agent in place; the override dict makes that choice survive a rebuild + # without touching shared process state. + if pin_session_override and isinstance(session, dict) and not one_turn: + session["model_override"] = { + "model": result.new_model, + "provider": result.target_provider, + "base_url": result.base_url, + "api_key": result.api_key, + "api_mode": result.api_mode, + } + if superseded_pending is not None and session.get("pending_model_switch") is superseded_pending: + # Keep the old intent in its canonical queue until a manual + # choice finishes publication. Overlapping choices share claims + # on this same old intent; a newer queued pick is a new dict. + pending_publication = superseded_pending.get("_model_switch_publications") + if pending_publication is None: + pending_publication = { + "owners": [], + "projection": {key: superseded_pending[key] for key in + ("display_model", "display_provider") if key in superseded_pending}, + } + superseded_pending["_model_switch_publications"] = pending_publication + pending_publication["owners"].append(result) + # Project the working pin/owned once runtime, including its + # natural restore, rather than the temporarily held old B. + superseded_pending.pop("display_model", None) + superseded_pending.pop("display_provider", None) + if isinstance(session, dict): + session.pop("model_verified_for", None) + mirror = session.get("_metadata_mirror") + if isinstance(mirror, dict): + mirror["model_ready"] = False + def owns_publication(): + # Caller holds the commit lock. Failed custody is superseded only by + # this exact explicit publication, never by a successor's projection. + if (not owns_expected_admission() or _sessions.get(sid) is not session or session.get("agent") is not agent + or session.get("transport") is not owner_transport): + return False + if failed_once is not None: + current = session.get("inflight_turn") + return (session.get("_one_turn_model_runtime") is failed_once + and failed_once.get("superseding_intent") is result + and int(session.get("_queued_prompt_generation", 0)) == publication_generation + and (current.get("started_at") if isinstance(current, dict) else None) is publication_started_at) + return True - # Record the switch as a PER-SESSION override so a later rebuild of THIS - # session (e.g. /new via _reset_session_agent, or resume) re-derives the - # user's chosen model/provider instead of falling back to global config. - # - # We deliberately do NOT write process-global env vars (HERMES_MODEL / - # HERMES_INFERENCE_MODEL / HERMES_TUI_PROVIDER / HERMES_INFERENCE_PROVIDER) - # here. The desktop backend hosts every same-profile session in ONE process, - # so mutating os.environ on a /model switch leaked the new model/provider - # into every OTHER live session's next agent rebuild — switching the model - # in one session silently changed it in the others (the cross-session - # contamination bug). agent.switch_model() above already mutated the right - # agent in place; the override dict makes that choice survive a rebuild - # without touching shared process state. - if pin_session_override and isinstance(session, dict) and not one_turn: - session["model_override"] = { - "model": result.new_model, - "provider": result.target_provider, - "base_url": result.base_url, - "api_key": result.api_key, - "api_mode": result.api_mode, - } - if isinstance(session, dict): - session.pop("model_verified_for", None) - mirror = session.get("_metadata_mirror") - if isinstance(mirror, dict): - mirror["model_ready"] = False - if agent: - _emit("session.info", sid, _session_info(agent, session)) - if persist_global: - _persist_model_switch(result) + # Worker replacement can acquire the registry lock; keep it outside + # history_lock to preserve the existing registry -> history lock order. + try: + if agent: + if explicit_model_intent: + with model_commit_lock if model_commit_lock is not None else contextlib.nullcontext(): + if not owns_publication(): + raise ValueError("session owner changed; model request not published") + _restart_slash_worker(sid, session) + with model_commit_lock if model_commit_lock is not None else contextlib.nullcontext(): + if explicit_model_intent and not owns_publication(): + raise ValueError("session owner changed; model request not published") + info = _session_info(agent, session) + # Transport callbacks may acquire history_lock; deliver outside it. + # The captured execution/transport cannot route through a successor. + token = _turn_outcome_execution.set((session, sid, None)) + transport_token = bind_transport(owner_transport) + try: + with model_commit_lock if model_commit_lock is not None else contextlib.nullcontext(): + if explicit_model_intent and not owns_publication(): + raise ValueError("session owner changed; model request not published") + _emit("session.info", sid, info) + finally: + reset_transport(transport_token) + _turn_outcome_execution.reset(token) + if persist_global: + with model_commit_lock if model_commit_lock is not None else contextlib.nullcontext(): + if explicit_model_intent and not owns_publication(): + raise ValueError("session owner changed; model request not published") + _persist_model_switch(result) + except Exception: + if failed_once is not None: + with model_commit_lock if model_commit_lock is not None else contextlib.nullcontext(): + if (_sessions.get(sid) is session and session.get("agent") is agent + and session.get("transport") is owner_transport + and session.get("_one_turn_model_runtime") is failed_once + and failed_once.get("superseding_intent") is result): + failed_once.pop("superseding_intent", None) + # A raised publication/persistence suffix means config.set cannot + # acknowledge C. Retain B's already acknowledged intent, without + # undoing the working client or resurrecting B over a newer choice. + if pending_publication is not None: + with model_commit_lock if model_commit_lock is not None else contextlib.nullcontext(): + if (_sessions.get(sid) is session + and session.get("pending_model_switch") is superseded_pending + and superseded_pending.get("_model_switch_publications") is pending_publication): + pending_publication["owners"][:] = [ + owner for owner in pending_publication["owners"] if owner is not result + ] + if not pending_publication["owners"]: + superseded_pending.pop("_model_switch_publications", None) + superseded_pending.update(pending_publication["projection"]) + if session.get("running"): + superseded_pending["after_inflight_turn"] = session.get("inflight_turn") + raise + if failed_once is not None: + with model_commit_lock if model_commit_lock is not None else contextlib.nullcontext(): + if not owns_publication(): + raise ValueError("session owner changed; model request not published") + if next_once is not None: + session["one_turn_model_restore"] = restore_snapshot + session["_one_turn_model_runtime"] = next_once + else: + session.pop("one_turn_model_restore", None) + session.pop("_one_turn_model_runtime", None) + if pending_publication is not None: + with model_commit_lock if model_commit_lock is not None else contextlib.nullcontext(): + if (_sessions.get(sid) is session + and session.get("pending_model_switch") is superseded_pending): + # Any acknowledged manual choice supersedes this captured + # old intent, even if another choice's suffix is in flight. + session.pop("pending_model_switch", None) return { "value": result.new_model, "warning": result.warning_message or "", @@ -7621,7 +7908,7 @@ def _sync_agent_compression_with_config(sid: str, session: dict) -> None: ) -def _apply_pending_model_switch(sid: str, session: dict) -> None: +def _apply_pending_model_switch(sid: str, session: dict, *, before_admission: bool = False) -> dict | None: """Apply a model switch queued while a turn was running. ``config.set model`` on a busy session doesn't mutate the live agent (the @@ -7632,31 +7919,78 @@ def _apply_pending_model_switch(sid: str, session: dict) -> None: the current model and never blocks the turn, matching ``_sync_agent_model_with_config``. """ - pending = session.pop("pending_model_switch", None) + with session["history_lock"]: + if before_admission and (session.get("running") or _sessions.get(sid) is not session + or session.get("_closing") or session.get("_finalized")): + return None + pending = session.get("pending_model_switch") + if pending and pending.get("_model_switch_publications") is not None: + return + # A pick that finished resolving AFTER Send's admission belongs to a + # later turn, even if this turn has not reached its setup yet. An + # accepted correction shallow-copies the replay dict, preserving its + # existing immutable started_at object from the same admission. + deferred_after = pending.get("after_inflight_turn") if pending else None + current_inflight = session.get("inflight_turn") + if (deferred_after is not None and ( + deferred_after is current_inflight + or (isinstance(deferred_after, dict) and isinstance(current_inflight, dict) + and deferred_after.get("started_at") is not None + and deferred_after["started_at"] is current_inflight.get("started_at")))): + return + pending = session.pop("pending_model_switch", None) + pending_agent = session.get("agent") + pending_transport = session.get("transport") + pending_generation = int(session.get("_queued_prompt_generation", 0)) + window = session.get("_turn_outcomes") + pending_admission = (pending_agent, pending_transport, pending_generation, + current_inflight.get("started_at") if isinstance(current_inflight, dict) else None, + window, + window.turns[-1]["accepted_turn"]["request_id"] + if isinstance(window, TurnOutcomeWindow) and window.owns(session, sid) and window.turns else None) if not pending or session.get("agent") is None: return + def emit_pending_error(message): + with session["history_lock"]: + current_inflight = session.get("inflight_turn") + current_window = session.get("_turn_outcomes") + current_nonce = (current_window.turns[-1]["accepted_turn"]["request_id"] + if isinstance(current_window, TurnOutcomeWindow) + and current_window.owns(session, sid) and current_window.turns else None) + if (_sessions.get(sid) is not session or session.get("agent") is not pending_agent + or session.get("transport") is not pending_transport + or int(session.get("_queued_prompt_generation", 0)) != pending_generation + or (before_admission and ( + (current_inflight.get("started_at") if isinstance(current_inflight, dict) else None) is not pending_admission[3] + or current_window is not pending_admission[4] + or current_nonce != pending_admission[5]))): + return + token = _turn_outcome_execution.set((session, sid, None)) + transport_token = bind_transport(pending_transport) + try: + _emit("error", sid, {"message": message}) + finally: + reset_transport(transport_token) + _turn_outcome_execution.reset(token) try: result = _apply_model_switch( sid, session, pending["raw"], confirm_expensive_model=bool(pending.get("confirm_expensive_model")), + supersede_pending=False, + explicit_model_intent=True, + defer_if_running=before_admission, + expected_admission=pending_admission if before_admission else None, ) # A queued pick is a deliberate user action; honour the expensive-model # confirm by NOT applying it silently — surface the warning and drop the # switch rather than spend on a pricey model the user never confirmed. if result.get("confirm_required"): - _emit( - "error", - sid, - {"message": result.get("confirm_message") or result.get("warning") or ""}, - ) + emit_pending_error(result.get("confirm_message") or result.get("warning") or "") + return result except Exception as e: - _emit( - "error", - sid, - {"message": f"Could not switch model: {e}"}, - ) + emit_pending_error(f"Could not switch model: {e}") class CompressionLockHeld(Exception): @@ -9161,6 +9495,12 @@ def _agent_fallback_model(agent): def _background_agent_kwargs(agent, task_id: str) -> dict: cfg = _load_cfg() + enabled_toolsets = getattr(agent, "enabled_toolsets", None) + if enabled_toolsets is None: + # Detached background tasks declare platform="tui" below: they have no + # UI session id, so a renderer-routed event has nowhere to land. Resolve + # defaults against that same platform, while preserving an explicit []. + enabled_toolsets = _load_enabled_toolsets("tui") return { "base_url": getattr(agent, "base_url", None) or None, @@ -9171,12 +9511,8 @@ def _background_agent_kwargs(agent, task_id: str) -> dict: "acp_args": getattr(agent, "acp_args", None) or None, "model": getattr(agent, "model", None) or _resolve_model(), "max_iterations": _cfg_max_turns(cfg, 25), - "enabled_toolsets": getattr(agent, "enabled_toolsets", None) - # Detached background tasks declare platform="tui" below: they have no - # UI session id, so a renderer-routed event has nowhere to land. Resolve - # their toolsets against that same platform rather than the gateway - # process's, so they never carry GUI schema they cannot use. - or _load_enabled_toolsets("tui"), + "enabled_toolsets": enabled_toolsets, + "disabled_toolsets": list(getattr(agent, "disabled_toolsets", None) or []), "quiet_mode": True, "verbose_logging": False, "ephemeral_system_prompt": getattr(agent, "ephemeral_system_prompt", None) @@ -9648,8 +9984,30 @@ def _make_agent( raise RuntimeError("Auth fallback resolved without a model") model = resolution.selected_model _pr = _load_provider_routing() + from agent.skill_utils import parse_config_string_list + + agent_cfg = cfg.get("agent") or {} + disabled_toolsets = [ + name.strip() + for name in parse_config_string_list( + agent_cfg.get("disabled_toolsets") if isinstance(agent_cfg, dict) else None + ) + if name.strip() + ] + # Keep global cap parsing and invalid-value handling in the initializer. + # A resolved provider cap is only its fallback when the global is unset. + model_cfg = cfg.get("model") + global_cap = model_cfg.get("max_tokens") if isinstance(model_cfg, dict) else None + provider_cap = runtime.get("max_output_tokens") + max_tokens = ( + provider_cap + if global_cap is None and isinstance(provider_cap, int) + and not isinstance(provider_cap, bool) and provider_cap > 0 + else None + ) return AIAgent( model=model, + max_tokens=max_tokens, max_iterations=_cfg_max_turns(cfg, 500), provider=runtime.get("provider"), base_url=runtime.get("base_url"), @@ -9677,6 +10035,7 @@ def _make_agent( else _load_service_tier() ), enabled_toolsets=_load_enabled_toolsets(_resolve_agent_platform(platform_override)), + disabled_toolsets=disabled_toolsets, # OpenRouter provider-routing prefs (config.yaml `provider_routing`). # Mirrors the messaging gateway + CLI so the desktop/TUI honors the same # routing instead of letting OpenRouter pick providers at random. @@ -10938,6 +11297,12 @@ def _drain_queued_prompt(rid, sid: str, session: dict) -> bool: if not queued or session.get("running"): return False queue_generation = int(session.get("_queued_prompt_generation", 0)) + claim_agent = session.get("agent") + claim_inflight = session.get("inflight_turn") + claim_started_at = claim_inflight.get("started_at") if isinstance(claim_inflight, dict) else None + claim_window = session.get("_turn_outcomes") + claim_nonce = (claim_window.turns[-1]["accepted_turn"]["request_id"] + if isinstance(claim_window, TurnOutcomeWindow) and claim_window.owns(session, sid) and claim_window.turns else None) queued_prompts = session.get("queued_prompts") or [] session["queued_prompt"] = queued_prompts.pop(0) if queued_prompts else None if not queued_prompts: @@ -10951,7 +11316,17 @@ def _drain_queued_prompt(rid, sid: str, session: dict) -> bool: if int(session.get("_last_stop_queue_generation", 0)) > queue_generation: # This exact claim predates an explicit Stop cut. Restoring it # would resurrect cancelled input ahead of post-cut arrivals. - if not session.get("_compute_host_active_request_id"): + current_inflight = session.get("inflight_turn") + current_window = session.get("_turn_outcomes") + current_nonce = (current_window.turns[-1]["accepted_turn"]["request_id"] + if isinstance(current_window, TurnOutcomeWindow) + and current_window.owns(session, sid) and current_window.turns else None) + if (not session.get("_compute_host_active_request_id") + and _sessions.get(sid) is session and session.get("agent") is claim_agent + and session.get("_turn_cancel_requested") + and (current_inflight is None or (isinstance(current_inflight, dict) + and current_inflight.get("started_at") is claim_started_at)) + and current_window is claim_window and current_nonce == claim_nonce): session["running"] = False return True # A non-Stop generation change (such as a compress re-anchor) @@ -10971,24 +11346,46 @@ def _drain_queued_prompt(rid, sid: str, session: dict) -> bool: session.pop("queued_prompts", None) session["running"] = False return True + restore_error = _one_turn_model_restore_error(sid, session) + if restore_error is not None: + # This accepted envelope owns a distinct pre-provider outcome, + # including when the normal execution policy uses a compute host. + _start_inflight_turn(session, queued["text"]) + refusal_started_at = session["inflight_turn"]["started_at"] + refusal = _begin_turn_outcome(session, sid, f"inline-turn-{uuid.uuid4().hex}", "inline") + refusal_transport = queued.get("transport") or session.get("transport") + refusal_window = session.get("_turn_outcomes") dispatch_failed = False try: - if use_compute_host: - if queued.get("image_paths"): - resp = _submit_prompt_to_compute_host( - rid, - sid, - session, - queued["text"], - image_paths=queued["image_paths"], - queued_prompt_generation=queue_generation, - ) - else: - resp = _submit_prompt_to_compute_host( - rid, sid, session, queued["text"], queued_prompt_generation=queue_generation, - **({"context_input_event_id": queued["context_input_event_id"]} - if queued.get("context_input_event_id") else {}), - ) + kwargs = {"queued_prompt_generation": queue_generation} + if queued.get("image_paths"): + kwargs["image_paths"] = queued["image_paths"] + if queued.get("context_input_event_id"): + kwargs["context_input_event_id"] = queued["context_input_event_id"] + if restore_error is not None: + token = _turn_outcome_execution.set((session, sid, refusal["request_id"])) + try: + settled = _emit_terminal_turn_error(sid, session, + "The saved model was not restored; choose a model explicitly before sending.", + error_surface=restore_error, expected_started_at=refusal_started_at, + expected_queue_generation=queue_generation, settle_running=True, + terminal_transport=refusal_transport, + context_input_event_id=queued.get("context_input_event_id")) + finally: + _turn_outcome_execution.reset(token) + if not settled: + with session["history_lock"]: + if (_sessions.get(sid) is session + and session.get("_turn_outcomes") is refusal_window + and isinstance(refusal_window, TurnOutcomeWindow) + and refusal_window.owns(session, sid) + and int(session.get("_last_stop_queue_generation", 0)) > queue_generation): + # Stop cut this exact envelope after its nonce began. + # Finish only its projection, never a successor's state. + refusal_window.finish(refusal["request_id"], interrupted=True) + dispatch_failed = True + elif use_compute_host: + resp = _submit_prompt_to_compute_host(rid, sid, session, queued["text"], **kwargs) if resp.get("error"): message = str(((resp.get("error") or {}).get("message")) or "queued prompt failed") with session["history_lock"]: @@ -10997,25 +11394,8 @@ def _drain_queued_prompt(rid, sid: str, session: dict) -> bool: _emit("error", sid, {"message": message}) dispatch_failed = True else: - if queued.get("image_paths"): - _run_prompt_submit( - rid, - sid, - session, - queued["text"], - image_paths=queued["image_paths"], - queued_prompt_generation=queue_generation, - ) - else: - _run_prompt_submit( - rid, - sid, - session, - queued["text"], - queued_prompt_generation=queue_generation, - **({"context_input_event_id": queued["context_input_event_id"]} - if queued.get("context_input_event_id") else {}), - ) + dispatch_failed = _run_prompt_submit(rid, sid, session, queued["text"], + turn_transport=queued.get("transport"), **kwargs) is False except Exception as exc: print( f"[tui_gateway] queued prompt dispatch failed: " @@ -11027,9 +11407,10 @@ def _drain_queued_prompt(rid, sid: str, session: dict) -> bool: dispatch_failed = True if dispatch_failed: with session["history_lock"]: - drain_next = bool(session.get("queued_prompt")) and not session.get( - "_turn_cancel_requested" - ) + drain_next = (bool(session.get("queued_prompt")) + and _sessions.get(sid) is session and not session.get("running") + and int(session.get("_queued_prompt_generation", 0)) == queue_generation + and not session.get("_turn_cancel_requested")) if drain_next: _drain_queued_prompt(rid, sid, session) return True @@ -11081,8 +11462,11 @@ def _inflight_snapshot(session: dict) -> dict | None: def _emit_terminal_turn_error( - sid: str, session: dict, error: Any, error_surface: Optional[dict] = None -) -> None: + sid: str, session: dict, error: Any, error_surface: Optional[dict] = None, + *, expected_started_at: Any = None, expected_queue_generation: int | None = None, + settle_running: bool = False, terminal_transport=None, + context_input_event_id: str | None = None, +) -> bool: """Close a failed turn with a terminal ``message.complete`` frame. Emits the same ``status: "error"`` frame shape the returned-error path in @@ -11094,7 +11478,13 @@ def _emit_terminal_turn_error( ``error_surface`` lets callers that already know the failing layer (e.g. agent-init failures = local runtime) pass it explicitly; exception callers leave it None and the classifier derives it here. + + Deferred pre-provider callers supply the existing admission's started_at + and queue generation and request atomic settlement. A stale/cancelled + owner returns False without mutation or delivery. Positional callers keep + their established provider-error behavior. """ + execution = _turn_outcome_execution.get() agent = session.get("agent") # Classify the failure into a {layer, code, retryable} descriptor so the # desktop can say "Provider error" / "Gateway error" with matching @@ -11111,12 +11501,42 @@ def _emit_terminal_turn_error( except Exception: error_surface = None with session["history_lock"]: + if settle_running: + current = _sessions.get(sid) + ready = session.get("resume_history_ready") + failed_history = (ready is not None and ready.is_set() + and session.get("resume_history_error") + and session.get("agent_error") == session.get("resume_history_error") + and not session.get("_compute_host_active")) + turn = session.get("inflight_turn") + window = session.get("_turn_outcomes") + if (current is not session and not (current is None and failed_history) + or session.get("_closing") or session.get("_finalized") + or session.get("_turn_cancel_requested") or not session.get("running") + or expected_started_at is None or not isinstance(turn, dict) + or turn.get("started_at") is not expected_started_at + or expected_queue_generation is None + or int(session.get("_queued_prompt_generation", 0)) != expected_queue_generation + or execution is None or execution[0] is not session or execution[1] != sid): + return False + if execution[2] is not None and ( + not isinstance(window, TurnOutcomeWindow) or not window.owns(session, sid) + or not window.turns or window.turns[-1]["accepted_turn"]["request_id"] != execution[2] + ): + return False _fail_inflight_turn(session, error, error_surface=error_surface) session.pop("model_verified_for", None) turn = session.get("inflight_turn") or {} message = str(turn.get("error") or "turn failed") partial = str(turn.get("assistant") or "") cols = int(session.get("cols", 80)) + if settle_running: + # These callers have not started a conversation or recorded its + # marker. A pre-existing marker belongs to crash recovery or a + # different runtime after failed hydration releases its lease. + # Only the code that records that marker may retire it. + session["running"] = False + session["last_active"] = time.time() text = partial or f"Error: {message}" payload = { "text": text, @@ -11127,6 +11547,10 @@ def _emit_terminal_turn_error( } if error_surface: payload["error_surface"] = error_surface + if settle_running and (error_surface or {}).get("code") == "one_turn_model_restore_failed": + payload.update(execution_started=False, durable_input_accepted=context_input_event_id is not None) + if context_input_event_id is not None: + payload["input_event_id"] = context_input_event_id if partial: payload["partial"] = True try: @@ -11135,9 +11559,44 @@ def _emit_terminal_turn_error( rendered = "" if rendered: payload["rendered"] = rendered - _retire_turn_marker(session) - _emit("message.complete", sid, payload) - _emit("session.info", sid, _session_info(agent, session)) + if not settle_running: + _retire_turn_marker(session) + _emit("message.complete", sid, payload) + _emit("session.info", sid, _session_info(agent, session)) + return True + transport_token = bind_transport(terminal_transport or _stdio_transport) + try: + _emit("message.complete", sid, payload) + with session["history_lock"]: + window = session.get("_turn_outcomes") + if (execution[2] is not None and isinstance(window, TurnOutcomeWindow) + and window.owns(session, sid)): + # A successor may already be newest. Finish only this nonce; + # the canonical emitter above appended its captured payload. + window.finish(execution[2], error=True) + turn = session.get("inflight_turn") + if (_sessions.get(sid) is session and not session.get("_closing") + and not session.get("_finalized") and isinstance(turn, dict) + and turn.get("started_at") is expected_started_at + and int(session.get("_queued_prompt_generation", 0)) == expected_queue_generation + and (execution[2] is None or (isinstance(window, TurnOutcomeWindow) + and window.owns(session, sid) and window.turns + and window.turns[-1]["accepted_turn"]["request_id"] == execution[2]))): + # No detached metadata, and no old idle snapshot after Send B. + if (error_surface or {}).get("code") == "one_turn_model_restore_failed": + settled_info = _session_info(agent, session) + else: + # Other settlement callers retain their existing owner + # contract; this correction is scoped to once refusal. + _emit("session.info", sid, _session_info(agent, session)) + settled_info = None + else: + settled_info = None + if settled_info is not None: + _emit("session.info", sid, settled_info) + finally: + reset_transport(transport_token) + return True def _restore_agent_history_after_turn_error(session: dict, agent) -> bool: @@ -13553,6 +14012,7 @@ def _run_prompt_submit( image_paths: list[str] | None = None, queued_prompt_generation: int | None = None, context_input_event_id: str | None = None, + turn_transport=None, ) -> bool: execution = _turn_outcome_execution.get() if execution is not None and (execution[0] is not session or execution[1] != sid): @@ -13567,9 +14027,127 @@ def _run_prompt_submit( window = session.get("_turn_outcomes") if isinstance(window, TurnOutcomeWindow) and window.owns(session, sid) and window.find(rid) is not None: execution = (session, sid, rid) + history_ready = session.get("resume_history_ready") + registered = _sessions.get(sid) is session with session["history_lock"]: - if session.get("_closing") or session.get("_turn_cancel_requested"): - session["running"] = False + inflight = session.get("inflight_turn") + admission = inflight.get("started_at") if isinstance(inflight, dict) else None + admission_generation = int(session.get("_queued_prompt_generation", 0)) + admission_owner_transport = session.get("transport") + # A caller can inherit another session's RPC context. The accepted + # envelope and captured session owner take precedence over that fallback. + admission_transport = turn_transport or admission_owner_transport or current_transport() + admission_window = session.get("_turn_outcomes") + admission_nonce = (admission_window.turns[-1]["accepted_turn"]["request_id"] + if isinstance(admission_window, TurnOutcomeWindow) + and admission_window.owns(session, sid) and admission_window.turns else None) + + def owns_admission(*, failed_history: bool = False) -> bool: + # Caller holds history_lock. Reuse the accepted owner; a wait cannot + # grant execution authority or retire a successor after Stop/new Send. + current = _sessions.get(sid) + if (registered or history_ready is not None or execution is not None) and current is not session: + if not (failed_history and current is None and history_ready is not None + and history_ready.is_set() and session.get("resume_history_error") + and session.get("agent_error") == session.get("resume_history_error") + and not session.get("_compute_host_active")): + return False + if int(session.get("_queued_prompt_generation", 0)) != admission_generation: + return False + if queued_prompt_generation is not None and int(session.get("_queued_prompt_generation", 0)) != queued_prompt_generation: + return False + current_inflight = session.get("inflight_turn") + if admission is not None and ( + not isinstance(current_inflight, dict) or current_inflight.get("started_at") is not admission + ): + return False + if execution is not None and execution[2] is not None: + window = session.get("_turn_outcomes") + if (not isinstance(window, TurnOutcomeWindow) or not window.owns(session, sid) + or not window.turns + or window.turns[-1]["accepted_turn"]["request_id"] != execution[2]): + return False + return True + + def owns_stopped_admission() -> bool: + # Stop revokes execution, but the same cancelled owner still settles + # its consumed once lease. A later admission (even one already stopped) + # cannot be borrowed through a cleared inflight slot or a missing nonce. + current_inflight = session.get("inflight_turn") + window = session.get("_turn_outcomes") + nonce = (window.turns[-1]["accepted_turn"]["request_id"] + if isinstance(window, TurnOutcomeWindow) and window.owns(session, sid) + and window.turns else None) + generation = int(session.get("_queued_prompt_generation", 0)) + return bool(_sessions.get(sid) is session + and session.get("_turn_cancel_requested") + and generation == session.get("_last_stop_queue_generation") + and generation > admission_generation + and ((isinstance(current_inflight, dict) + and current_inflight.get("started_at") is admission) + or (current_inflight is None and not session.get("running"))) + and window is admission_window and nonce == admission_nonce + and (execution is None or execution[2] is None or nonce == execution[2])) + + with session["history_lock"]: + restore_error = _one_turn_model_restore_error(sid, session) + if restore_error is not None: + if not owns_admission() or session.get("_closing") or session.get("_finalized") or session.get("_turn_cancel_requested"): + return False + if not isinstance(inflight, dict) or inflight.get("status") == "error": + _start_inflight_turn(session, text) + admission = session["inflight_turn"]["started_at"] + runtime = session.get("_one_turn_model_runtime") + runtime["restore_failed"] = True + runtime["active"] = False + session["running"] = True + if restore_error is not None: + token = _turn_outcome_execution.set(execution if execution is not None else (session, sid, None)) + try: + _emit_terminal_turn_error(sid, session, + "The saved model was not restored; choose a model explicitly before sending.", + error_surface=restore_error, expected_started_at=admission, + expected_queue_generation=admission_generation, settle_running=True, + terminal_transport=admission_transport, context_input_event_id=context_input_event_id) + finally: + _turn_outcome_execution.reset(token) + return False + + if history_ready is not None and (not history_ready.is_set() or session.get("resume_history_error")): + with session["history_lock"]: + if not owns_admission(failed_history=True): + return False + if (not session.get("running") or session.get("_closing") + or session.get("_finalized") or session.get("_turn_cancel_requested")): + return False + if not isinstance(inflight, dict) or inflight.get("status") == "error": + _start_inflight_turn(session, text) + admission = session["inflight_turn"]["started_at"] + terminal_execution = execution if execution is not None else (session, sid, None) + token = _turn_outcome_execution.set(terminal_execution) + try: + error = _wait_agent_for_prompt(session, rid, sid) + with session["history_lock"]: + if not owns_admission(failed_history=True): + return False + cancelled = (session.get("_closing") or session.get("_finalized") + or session.get("_turn_cancel_requested") or not session.get("running")) + if cancelled: + return False + if error is not None: + _emit_terminal_turn_error(sid, session, error["error"]["message"], + error_surface={"layer": "runtime", "code": "resume_history_unavailable", "retryable": True}, + expected_started_at=admission, expected_queue_generation=admission_generation, + settle_running=True, terminal_transport=admission_transport) + return False + finally: + _turn_outcome_execution.reset(token) + if not history_ready.is_set(): + return False + with session["history_lock"]: + if not owns_admission(): + return False + if session.get("_closing") or session.get("_finalized") or session.get("_turn_cancel_requested"): return False if ( queued_prompt_generation is not None @@ -13587,6 +14165,7 @@ def _run_prompt_submit( # by the time a new turn starts — replace it, never append onto it. if not isinstance(inflight, dict) or inflight.get("status") == "error": _start_inflight_turn(session, text) + admission = session["inflight_turn"]["started_at"] agent = session["agent"] if hasattr(agent, "clear_interrupt"): try: @@ -13619,7 +14198,7 @@ def run_with_outcome(): # dispatcher do not follow automatically. Rebind the exact transport # stored on this session generation before any tool can commission a # child; delegate_task then captures it as non-serializable authority. - transport_token = bind_transport(session.get("transport")) + transport_token = bind_transport(admission_transport) runtime_session_token = _current_runtime_session_record.set(session) # Bound eagerly so the except/finally paths below always have an agent # even if turn setup throws; re-read after _sync_bot_capabilities, @@ -13699,6 +14278,8 @@ def run_with_outcome(): # Snapshot after turn-start model sync. A deferred switch mutates # history and its version; that mutation belongs to this turn. with session["history_lock"]: + if one_turn_runtime is not None and not _owns_one_turn_model_runtime(session, agent, one_turn_runtime): + return history = list(session["history"]) history_version = int(session.get("history_version", 0)) cwd = _session_cwd(session) @@ -13946,6 +14527,9 @@ def _interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None: # message.complete. _usage_stop.set() _usage_thread.join() + with session["history_lock"]: + if not owns_admission() or session.get("agent") is not agent: + return if display_kind and isinstance(text, str): db = getattr(agent, "_session_db", None) current_session_id = getattr(agent, "session_id", None) or session.get("session_key") @@ -14187,7 +14771,12 @@ def _interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None: ) turn_error_retained = True else: - _clear_inflight_turn(session) + # Keep this admission's immutable started_at until the + # existing finally settles model restoration and lifecycle. + # A terminal response alone grants no successor ownership. + inflight = session.get("inflight_turn") + if isinstance(inflight, dict): + inflight["streaming"] = False if status == "complete": from agent.turn_finalizer import AcceptedResponseRoute @@ -14471,20 +15060,54 @@ def _interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None: pass if tts_queue is not None: tts_queue.put(None) # end-of-text sentinel — flush + finish speaking - if one_turn_restore and _owns_one_turn_model_runtime(session, agent, one_turn_runtime): + restored_once = False + def owns_restore(): + return ((owns_admission() or owns_stopped_admission()) + and _owns_one_turn_model_runtime(session, agent, one_turn_runtime) + and one_turn_runtime.get("superseding_intent") is None) + if one_turn_restore: try: - _restore_agent_model_runtime(agent, one_turn_restore) - _restart_slash_worker(sid, session) - _persist_live_session_runtime(session) - _persist_live_session_system_prompt(session) + with session["history_lock"]: + restore_owned = owns_restore() + if restore_owned: + _restore_agent_model_runtime(agent, one_turn_restore) + if restore_owned: + # Registry-taking worker replacement stays outside history_lock. + _restart_slash_worker(sid, session) + with session["history_lock"]: + if not owns_restore(): + raise RuntimeError("one-turn restore owner changed during publication") + _persist_live_session_runtime(session) + _persist_live_session_system_prompt(session) + if not owns_restore(): + raise RuntimeError("one-turn restore owner changed during publication") + restored_once = True except Exception: - one_turn_runtime["restore_failed"] = True - _emit("error", sid, {"message": "Could not restore the saved model after this one-turn selection."}) + with session["history_lock"]: + failed_owner = owns_restore() + if failed_owner: + one_turn_runtime["restore_failed"] = True + if failed_owner: + token = bind_transport(admission_transport) + try: + _emit("error", sid, {"message": "Could not restore the saved model after this one-turn selection.", + "error_surface": {"layer": "runtime", "code": "one_turn_model_restore_failed", "retryable": False}, + **({"request_id": execution[2]} if execution is not None and execution[2] is not None else {})}) + finally: + reset_transport(token) logger.debug("TUI one-turn model restore failed", exc_info=True) - if one_turn_runtime is not None and session.get("_one_turn_model_runtime") is one_turn_runtime: - one_turn_runtime["active"] = False - if not one_turn_runtime.get("restore_failed"): - session.pop("_one_turn_model_runtime", None) + with session["history_lock"]: + if (one_turn_runtime is not None and _sessions.get(sid) is session + and session.get("_one_turn_model_runtime") is one_turn_runtime + and one_turn_runtime.get("session") is session + and one_turn_runtime.get("sid") == sid + and one_turn_runtime.get("restore_snapshot") is one_turn_restore): + one_turn_runtime["active"] = False + if restored_once and owns_restore(): + session.pop("_one_turn_model_runtime", None) + else: + # Custody annotation is not restoration or execution authority. + one_turn_runtime["restore_failed"] = True try: if approval_token is not None: reset_current_session_key(approval_token) @@ -14499,12 +15122,19 @@ def _interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None: reset_transport(transport_token) # Clear the per-turn interim callback so a stale closure from # this turn can't fire during a later turn on the same agent. - agent.interim_assistant_callback = None with session["history_lock"]: + if ((not owns_admission() and not owns_stopped_admission()) + or (registered and _sessions.get(sid) is not session) + or session.get("agent") is not agent + or (one_turn_runtime is not None and one_turn_runtime.get("agent") is not agent)): + return + agent.interim_assistant_callback = None session["running"] = False session["last_active"] = time.time() if not turn_error_retained: _clear_inflight_turn(session) + _retire_turn_marker(session, marker_key) + session.pop("_auto_continue_scheduled", None) # Closing bookend of the "tui prompt accepted" record above — # fires on every path (success, returned error, exception, # interrupt), so one accepted prompt always produces exactly one @@ -14533,15 +15163,34 @@ def _interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None: ) # Backstop for turns that never reached a terminal frame (the # frame paths retire the marker as they emit). - _retire_turn_marker(session, marker_key) - session.pop("_auto_continue_scheduled", None) - _settled_owner = _sessions.get(sid) - if _settled_owner is None or _settled_owner is session: - # Publish the settled snapshot only while this session still - # owns the UI id (a post-turn rebind must not republish the - # old session's readiness onto the new owner's id). - _emit_settled_session_info(sid, session, agent) with session["history_lock"]: + if (_sessions.get(sid) not in (None, session) or session.get("agent") is not agent + or (int(session.get("_queued_prompt_generation", 0)) != admission_generation + and not owns_stopped_admission()) + or (session.get("inflight_turn") is not None and not turn_error_retained) + or session.get("running")): + return + # Preserve settled-cwd semantics while separating snapshot + # custody from transport I/O on this non-reentrant lock. + try: + _reconcile_session_cwd_from_terminal(session) + except Exception: + logger.debug("failed to reconcile settled session cwd", exc_info=True) + settled_info = _session_info(agent, session) + settled_execution_token = _turn_outcome_execution.set((session, sid, + execution[2] if execution is not None else None)) + settled_transport_token = bind_transport(admission_transport) + try: + _emit("session.info", sid, settled_info) + finally: + reset_transport(settled_transport_token) + _turn_outcome_execution.reset(settled_execution_token) + with session["history_lock"]: + if (_sessions.get(sid) not in (None, session) or session.get("agent") is not agent + or int(session.get("_queued_prompt_generation", 0)) != admission_generation + or session.get("running") + or (session.get("inflight_turn") is not None and not turn_error_retained)): + return deferred_teardown = session.pop("_run_checkpoint_teardown_deferred", None) deferred_finalize = session.pop("_run_checkpoint_finalize_deferred", None) if deferred_teardown: @@ -15114,6 +15763,8 @@ def _(rid, params: dict) -> dict: # The user gets to pick, keep typing, and send the next turn on # the new model without waiting for the swap or interrupting. if session.get("running"): + pending_agent = session.get("agent") + pending_transport = session.get("transport") parsed = parse_model_switch_args(value) try: pending_model = parsed.model_input @@ -15161,16 +15812,19 @@ def _(rid, params: dict) -> dict: "deferred": False, }, ) - session["pending_model_switch"] = { - "raw": value, - "confirm_expensive_model": confirmed, - # The resolved model/provider the next turn will run on. - # _session_info reports these while the switch is pending - # so the end-of-turn settle keeps showing the user's pick - # instead of blipping back to the still-live old model. - "display_model": pending_model, - "display_provider": pending_provider, - } + with session["history_lock"]: + if (_sessions.get(params["session_id"]) is not session + or session.get("agent") is not pending_agent + or session.get("transport") is not pending_transport): + return _err(rid, 4001, "session owner changed; request not applied") + session["pending_model_switch"] = { + "raw": value, + "confirm_expensive_model": confirmed, + # Projection names the next eligible turn's choice. + "display_model": pending_model, + "display_provider": pending_provider, + "after_inflight_turn": session.get("inflight_turn"), + } return _ok( rid, { @@ -15201,6 +15855,9 @@ def _(rid, params: dict) -> dict: params.get("confirm_expensive_model", False) ), parsed_flags=parsed_flags, + defer_if_running=True, + supersede_pending=True, + explicit_model_intent=True, ) else: result = _apply_model_switch( @@ -15220,6 +15877,7 @@ def _(rid, params: dict) -> dict: "confirm_required": result.get("confirm_required", False), "confirm_message": result.get("confirm_message", ""), "scope": result.get("scope", "session"), + **({"deferred": result["deferred"]} if "deferred" in result else {}), }, ) except Exception as e: @@ -17499,7 +18157,7 @@ def _mirror_slash_side_effects(sid: str, session: dict, command: str) -> str: try: if name == "model" and arg and agent: - result = _apply_model_switch(sid, session, arg) + result = _apply_model_switch(sid, session, arg, explicit_model_intent=True) return result.get("warning", "") elif name == "approvals" and arg: # The slash worker already persisted the new approvals.mode; the diff --git a/utils.py b/utils.py index 33d4932302f2d..196bd17a0f245 100644 --- a/utils.py +++ b/utils.py @@ -140,88 +140,53 @@ def _is_contended_windows_replace_error(exc: OSError) -> bool: ) -def _rewrite_in_place(tmp_str: str, real_path: str) -> None: - """Overwrite *real_path* with the contents of *tmp_str*, in place. - - Last-resort path for a target whose handle is still held after the retry - budget: writing through the existing file works where renaming onto it - does not. Unlike ``shutil.copyfile`` this never truncates the target to - zero first — a concurrent reader would otherwise be able to observe an - empty ``auth.json`` / ``gateway_state.json`` mid-write (measured: a - 4-thread poller sees a 0-byte read during a plain copyfile). A single - ``os.write`` of the full payload followed by ``ftruncate`` keeps the - visible content going straight from old to new. - - This is still not atomic — it is a strictly smaller window than a copy, - not the absence of one — so it runs only after the rename has genuinely - failed. Writing through the target also preserves its ACL, which - ``os.replace`` does not (the temp file's inherited ACL wins there). + + + +def _copy_fallback(tmp_str: str, real_path: str) -> None: + """Restage a cross-device source beside the target before publication. + + Copying into the existing target would destroy its previous complete + bytes on a short write or ENOSPC. Only a complete, fsynced sibling may + replace it. The incoming temp remains the caller's on staging failure. """ - with open(tmp_str, "rb") as src: - data = src.read() - flags = os.O_WRONLY | getattr(os, "O_BINARY", 0) - fd = os.open(real_path, flags) + target = Path(real_path) + fd, sibling = tempfile.mkstemp( + dir=str(target.parent), prefix=f".{target.name[:80]}.", suffix=".tmp" + ) try: - os.lseek(fd, 0, os.SEEK_SET) - written = 0 - while written < len(data): - written += os.write(fd, data[written:]) - os.ftruncate(fd, len(data)) + with os.fdopen(fd, "wb") as dst: + fd = None + with open(tmp_str, "rb") as src: + shutil.copyfileobj(src, dst) + dst.flush() + shutil.copystat(tmp_str, sibling) + os.fsync(dst.fileno()) + os.replace(sibling, real_path) + finally: + if fd is not None: + os.close(fd) try: - os.fsync(fd) + os.unlink(sibling) except OSError: pass - finally: - os.close(fd) - os.unlink(tmp_str) - - -def _copy_fallback(tmp_str: str, real_path: str) -> None: - """Copy/fsync/unlink fallback for cross-device and bind-mount renames.""" - shutil.copyfile(tmp_str, real_path) - try: - shutil.copystat(tmp_str, real_path) - except OSError: - pass - try: - with open(real_path, "rb") as f: - os.fsync(f.fileno()) - except OSError: - pass os.unlink(tmp_str) def atomic_replace(tmp_path: Union[str, Path], target: Union[str, Path]) -> str: """Atomically move *tmp_path* onto *target*, preserving symlinks. - ``os.replace(tmp, target)`` atomically swaps ``tmp`` into place at - ``target``. When ``target`` is a symlink, the symlink itself is - replaced with a regular file — silently detaching managed deployments - that symlink ``config.yaml`` / ``SOUL.md`` / ``auth.json`` etc. from - ``~/.hermes/`` to a git-tracked profile package or dotfiles repo - (GitHub #16743). - - This helper resolves the symlink first so ``os.replace`` writes to - the real file in-place while the symlink survives. For non-symlink - and non-existent paths the behavior is identical to a plain - ``os.replace`` call unless the rename fails with: - - * ``EXDEV`` / ``EBUSY`` (any platform) — cross-device, bind-mount, and - busy-file deployments fall back to copy/fsync/unlink immediately. - These never clear on retry. - * A Windows rename contended by another open handle (winerror 5/32/33). - CPython opens files without ``FILE_SHARE_DELETE``, so *any* concurrent - reader of the target blocks the rename. The rename is retried with - jittered backoff first — a retry that wins keeps the write atomic — - and only a target whose handle outlives the budget is rewritten in - place, so the update lands instead of being silently dropped. - - A genuine Windows permission failure produces the same winerror as a - contended one, so it is not classified up front: it exhausts the retry - budget, fails the in-place rewrite too, and is re-raised unchanged. - - Returns the resolved real path used for the replace, so callers that - need to re-apply permissions can target it instead of the symlink. + Resolve a symlinked target before replacing it, so the link survives. + EXDEV restages the complete source in the resolved target's directory + and publishes by rename there. EBUSY fails before writing the target. + + Windows winerror 5/32/33 may be transient handle contention or genuine + permission denial. Retry the existing bounded budget, then propagate + the last error. Never overwrite the target in place: a failed write + could destroy its previous complete contents. Successful retries keep + the ordinary rename semantics. + + Return the resolved real path so callers can reapply their metadata. """ target_str = str(target) real_path = os.path.realpath(target_str) if os.path.islink(target_str) else target_str @@ -234,8 +199,6 @@ def atomic_replace(tmp_path: Union[str, Path], target: Union[str, Path]) -> str: if exc.errno not in (errno.EXDEV, errno.EBUSY) and not contended: raise if contended: - # Lazy import: keeps ``utils`` free of a package-level dependency - # on ``agent`` for every consumer that never hits this path. from agent.retry_utils import jittered_backoff for attempt in range(1, _REPLACE_RETRY_ATTEMPTS + 1): @@ -251,28 +214,15 @@ def atomic_replace(tmp_path: Union[str, Path], target: Union[str, Path]) -> str: return real_path except OSError as retry_exc: if retry_exc.errno in (errno.EXDEV, errno.EBUSY): - # Not contention after all — stop burning the budget. exc = retry_exc - contended = False break if not _is_contended_windows_replace_error(retry_exc): raise exc = retry_exc - logger.debug( - "atomic_replace: %s -> %s failed with %s; falling back to %s", - tmp_str, - real_path, - getattr(exc, "winerror", None) - or errno.errorcode.get(exc.errno or 0, exc.errno), - "in-place rewrite" if contended else "copy", - ) - if contended: - # Re-raises the rewrite's own error (not the rename's) when the - # target is genuinely unwritable — an ACL denial stays an ACL - # denial rather than being reported as contention. - _rewrite_in_place(tmp_str, real_path) - else: - _copy_fallback(tmp_str, real_path) + if exc.errno != errno.EXDEV: + raise exc + logger.debug("atomic_replace: restaging %s beside %s after EXDEV", tmp_str, real_path) + _copy_fallback(tmp_str, real_path) return real_path