From 39dd85e7f032065d582c4d8d00fe5c7f989fe40f Mon Sep 17 00:00:00 2001 From: Ahmex000 Date: Fri, 20 Mar 2026 14:17:31 +0100 Subject: [PATCH] Fix model config override and code review issues - config.py: cli-config.json LLM vars now always override shell env, preventing stale shell values from reverting the configured model - checkpoint_restore.py: extract shared restore logic from cli/tui to eliminate code duplication - cli.py / tui.py: use shared checkpoint_restore module, add double-save guard via threading.Event - agents_graph_actions.py: add _agents_lock for thread-safe access to _running_agents and _agent_instances, fix mutable default arg in restore_sub_agents Co-Authored-By: Claude Sonnet 4.6 --- strix/config/config.py | 14 +- strix/interface/checkpoint_restore.py | 156 ++++++++++++++++ strix/interface/cli.py | 166 +----------------- strix/interface/tui.py | 137 +-------------- .../agents_graph/agents_graph_actions.py | 30 ++-- 5 files changed, 196 insertions(+), 307 deletions(-) create mode 100644 strix/interface/checkpoint_restore.py diff --git a/strix/config/config.py b/strix/config/config.py index bad994ad..4cf6e600 100644 --- a/strix/config/config.py +++ b/strix/config/config.py @@ -139,17 +139,15 @@ class Config: env_vars.pop(var_name, None) if cls._config_file_override is None: cls.save({"env": env_vars}) - if cls._llm_env_changed(env_vars): - for var_name in cls._llm_env_vars(): - env_vars.pop(var_name, None) - if cls._config_file_override is None: - cls.save({"env": env_vars}) applied = {} + llm_vars = cls._llm_env_vars() for var_name, var_value in env_vars.items(): - if var_name in cls.tracked_vars() and (force or var_name not in os.environ): - os.environ[var_name] = var_value - applied[var_name] = var_value + if var_name in cls.tracked_vars(): + # LLM vars in cli-config.json always win over shell env + if var_name in llm_vars or force or var_name not in os.environ: + os.environ[var_name] = var_value + applied[var_name] = var_value return applied diff --git a/strix/interface/checkpoint_restore.py b/strix/interface/checkpoint_restore.py new file mode 100644 index 00000000..40930bbc --- /dev/null +++ b/strix/interface/checkpoint_restore.py @@ -0,0 +1,156 @@ +"""Shared helpers for restoring scan state from a checkpoint. + +Used by both cli.py and tui.py to avoid code duplication. +""" + +from __future__ import annotations + +import asyncio +import threading +from typing import Any + + +def restore_sub_agents(checkpoint_data: Any, llm_config: Any) -> list[str]: + """Spawn previously-running sub-agents from a checkpoint. + + Each sub-agent is restored with its saved AgentState, blocking flags + reset, sandbox cleared, and a resume-context message injected. Agents + are started in topological order (parents before children) in daemon + threads. + + Returns the list of agent_ids that were successfully restored. + """ + from strix.agents.StrixAgent import StrixAgent + from strix.agents.state import AgentState + from strix.tools.agents_graph import agents_graph_actions + + sub_agent_states: dict[str, Any] = checkpoint_data.sub_agent_states or {} + if not sub_agent_states: + return [] + + _memo: dict[str, int] = {} + + def _depth(aid: str) -> int: + if aid in _memo: + return _memo[aid] + parent = sub_agent_states.get(aid, {}).get("parent_id") + _memo[aid] = ( + 0 + if (parent is None or parent not in sub_agent_states) + else 1 + _depth(parent) + ) + return _memo[aid] + + restored_ids: list[str] = [] + for agent_id in sorted(sub_agent_states.keys(), key=_depth): + state_dict = sub_agent_states[agent_id] + try: + state = AgentState.model_validate(state_dict) + state.waiting_for_input = False + state.waiting_start_time = None + state.stop_requested = False + state.completed = False + state.llm_failed = False + state.sandbox_id = None + state.sandbox_token = None + state.sandbox_info = None + state.max_iterations = state.iteration + checkpoint_data.original_max_iterations + state.max_iterations_warning_sent = False + state.add_message( + "user", + f"[SYSTEM - SUB-AGENT RESUMED]\n" + f"You were interrupted at iteration {state.iteration}. " + f"Your sandbox has been reset and a fresh one will be created. " + f"Review your conversation history above and continue your task " + f"from where you left off. " + f"Call agent_finish only when your task is genuinely complete.", + ) + + agent_cfg: dict[str, Any] = { + "llm_config": llm_config, + "state": state, + "is_resumed": True, + } + agent = StrixAgent(agent_cfg) + + def _run(a: Any = agent, s: Any = state) -> None: + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + try: + loop.run_until_complete(a.agent_loop(s.task)) + except Exception: # noqa: BLE001 + pass + finally: + loop.close() + + t = threading.Thread( + target=_run, + daemon=True, + name=f"ResumedAgent-{state.agent_name}-{agent_id[:8]}", + ) + t.start() + with agents_graph_actions._agents_lock: + agents_graph_actions._running_agents[agent_id] = t + restored_ids.append(agent_id) + + except Exception: # noqa: BLE001 + pass + + return restored_ids + + +def build_root_resume_message( + state: Any, + checkpoint_data: Any, + restored_ids: list[str] | None = None, +) -> None: + """Inject a user message telling the root agent what happened and what is live.""" + iteration = checkpoint_data.iteration + sub_agent_states: dict[str, Any] = checkpoint_data.sub_agent_states or {} + + if restored_ids: + lines = [ + "\n\nThe following sub-agents have been AUTOMATICALLY RESTORED and are " + "already running. Communicate with them using their original IDs:" + ] + for aid in restored_ids: + sd = sub_agent_states.get(aid, {}) + name = sd.get("agent_name", "sub-agent") + task = (sd.get("task") or "")[:200] + lines.append(f" • {name} (ID: {aid}) — task: {task}") + lines.append("\nDo NOT re-spawn these agents — they are already active.") + sub_agent_section = "\n".join(lines) + else: + dead: list[dict[str, Any]] = [] + for aid, node in (checkpoint_data.tracer_agents or {}).items(): + if node.get("parent_id") is None: + continue + status = node.get("status", "unknown") + if status not in ("completed", "finished", "stopped", "error", "failed"): + dead.append( + {"name": node.get("name", "sub-agent"), "task": (node.get("task") or "")[:300]} + ) + if dead: + lines = [ + "\n\nThe following sub-agents were active at interruption but could " + "not be automatically restored. Re-spawn them if their work is incomplete:" + ] + for sa in dead: + lines.append(f" • {sa['name']} (task: {sa['task']})") + sub_agent_section = "\n".join(lines) + else: + sub_agent_section = "" + + msg = ( + f"[SYSTEM - SCAN RESUMED]\n" + f"This penetration test was interrupted at iteration {iteration}. " + f"A fresh sandbox will be created automatically.\n\n" + f"CRITICAL: Any agent IDs that appear in the conversation history ABOVE " + f"this message are from the old session and are DEAD — do not interact " + f"with them. Only the agents listed below are currently alive." + f"{sub_agent_section}\n\n" + f"Review the conversation history to understand what has already been done, " + f"then CONTINUE the penetration test. " + f"Do NOT call finish_scan unless all testing is genuinely complete." + ) + state.add_message("user", msg) diff --git a/strix/interface/cli.py b/strix/interface/cli.py index 3c2982cb..8356713c 100644 --- a/strix/interface/cli.py +++ b/strix/interface/cli.py @@ -97,165 +97,8 @@ def _replay_previous_output( console.print() -def _restore_sub_agents(checkpoint_data: Any, llm_config: Any) -> list[str]: - """Spawn previously-running sub-agents from checkpoint with their full history. - - Each sub-agent is restored exactly like the root agent: its saved AgentState - (messages, iteration, task) is loaded, blocking flags reset, sandbox cleared, - and a resume-context message injected. The agent is then started in a daemon - thread with is_resumed=True so it continues from where it left off. - - Returns a list of the agent_ids that were successfully restored. - """ - import asyncio - import threading - - from strix.agents.StrixAgent import StrixAgent - from strix.agents.state import AgentState - from strix.tools.agents_graph import agents_graph_actions - - sub_agent_states: dict[str, Any] = checkpoint_data.sub_agent_states or {} - if not sub_agent_states: - return [] - - # Topological sort so parents are created (and registered in _agent_graph) - # before their children. - def _depth(aid: str, _memo: dict = {}) -> int: # noqa: B006 - if aid in _memo: - return _memo[aid] - parent = sub_agent_states.get(aid, {}).get("parent_id") - _memo[aid] = 0 if (parent is None or parent not in sub_agent_states) else 1 + _depth(parent, _memo) - return _memo[aid] - - ordered_ids = sorted(sub_agent_states.keys(), key=_depth) - - restored_ids: list[str] = [] - for agent_id in ordered_ids: - state_dict = sub_agent_states[agent_id] - try: - state = AgentState.model_validate(state_dict) - - # Reset any blocking flags captured at the moment of interruption - state.waiting_for_input = False - state.waiting_start_time = None - state.stop_requested = False - state.completed = False - state.llm_failed = False - - # Clear old sandbox — it no longer exists - state.sandbox_id = None - state.sandbox_token = None - state.sandbox_info = None - - # Give a fresh iteration budget from the resume point - state.max_iterations = state.iteration + checkpoint_data.original_max_iterations - state.max_iterations_warning_sent = False - - # Inject a sub-agent resume message so the LLM knows to continue - state.add_message( - "user", - f"[SYSTEM - SUB-AGENT RESUMED]\n" - f"You were interrupted at iteration {state.iteration}. " - f"Your sandbox has been reset and a fresh one will be created. " - f"Review your conversation history above and continue your task " - f"from where you left off. " - f"Call agent_finish only when your task is genuinely complete.", - ) - - agent_cfg: dict[str, Any] = { - "llm_config": llm_config, - "state": state, - "is_resumed": True, - } - agent = StrixAgent(agent_cfg) - - def _run(a: Any = agent, s: Any = state) -> None: - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - try: - loop.run_until_complete(a.agent_loop(s.task)) - except Exception: # noqa: BLE001 - pass - finally: - loop.close() - - t = threading.Thread( - target=_run, - daemon=True, - name=f"ResumedAgent-{state.agent_name}-{agent_id[:8]}", - ) - t.start() - agents_graph_actions._running_agents[agent_id] = t - restored_ids.append(agent_id) - - except Exception: # noqa: BLE001 - pass - - return restored_ids - - -def _build_resume_context_message( - state: Any, - checkpoint_data: Any, - restored_ids: "list[str] | None" = None, -) -> None: - """Inject a user message telling the root agent what happened and what is live. - - If sub-agents were already restored automatically, the message tells the root - agent their IDs so it can communicate with them directly. Otherwise it lists - what was running so the root agent knows what to re-spawn. - """ - iteration = checkpoint_data.iteration - sub_agent_states: dict[str, Any] = checkpoint_data.sub_agent_states or {} - - if restored_ids: - lines = [ - "\n\nThe following sub-agents have been AUTOMATICALLY RESTORED and are " - "already running. Communicate with them using their original IDs:" - ] - for aid in restored_ids: - sd = sub_agent_states.get(aid, {}) - name = sd.get("agent_name", "sub-agent") - task = (sd.get("task") or "")[:200] - lines.append(f" • {name} (ID: {aid}) — task: {task}") - lines.append("\nDo NOT re-spawn these agents — they are already active.") - sub_agent_section = "\n".join(lines) - else: - # No saved sub-agent states — fall back to listing from tracer_agents - dead: list[dict[str, Any]] = [] - for aid, node in (checkpoint_data.tracer_agents or {}).items(): - if node.get("parent_id") is None: - continue - status = node.get("status", "unknown") - if status not in ("completed", "finished", "stopped", "error", "failed"): - dead.append({ - "name": node.get("name", "sub-agent"), - "task": (node.get("task") or "")[:300], - }) - if dead: - lines = [ - "\n\nThe following sub-agents were active at interruption but could " - "not be automatically restored. Re-spawn them if their work is incomplete:" - ] - for sa in dead: - lines.append(f" • {sa['name']} (task: {sa['task']})") - sub_agent_section = "\n".join(lines) - else: - sub_agent_section = "" - - msg = ( - f"[SYSTEM - SCAN RESUMED]\n" - f"This penetration test was interrupted at iteration {iteration}. " - f"A fresh sandbox will be created automatically.\n\n" - f"CRITICAL: Any agent IDs that appear in the conversation history ABOVE " - f"this message are from the old session and are DEAD — do not interact " - f"with them. Only the agents listed below are currently alive." - f"{sub_agent_section}\n\n" - f"Review the conversation history to understand what has already been done, " - f"then CONTINUE the penetration test. " - f"Do NOT call finish_scan unless all testing is genuinely complete." - ) - state.add_message("user", msg) +from strix.interface.checkpoint_restore import build_root_resume_message as _build_resume_context_message # noqa: E402 +from strix.interface.checkpoint_restore import restore_sub_agents as _restore_sub_agents # noqa: E402 async def run_cli(args: Any) -> None: # noqa: PLR0915 @@ -415,10 +258,15 @@ async def run_cli(args: Any) -> None: # noqa: PLR0915 # update the reference once the agent is created. _agent_ref: list[Any] = [] + _checkpoint_saved = threading.Event() + def _save_checkpoint_on_interrupt() -> None: """Persist current agent state before exit so the scan can be resumed.""" if not checkpoint_manager or not _agent_ref: return + if _checkpoint_saved.is_set(): + return + _checkpoint_saved.set() try: agent_instance = _agent_ref[0] checkpoint_manager.save( diff --git a/strix/interface/tui.py b/strix/interface/tui.py index 2d8b8454..11cde76d 100644 --- a/strix/interface/tui.py +++ b/strix/interface/tui.py @@ -41,136 +41,8 @@ from strix.telemetry.tracer import Tracer, set_global_tracer logger = logging.getLogger(__name__) -def _restore_sub_agents_tui(checkpoint_data: Any, llm_config: Any) -> list[str]: - """Restore previously-running sub-agents from checkpoint (TUI variant).""" - import asyncio - import threading - - from strix.agents.StrixAgent import StrixAgent - from strix.agents.state import AgentState - from strix.tools.agents_graph import agents_graph_actions - - sub_agent_states: dict[str, Any] = checkpoint_data.sub_agent_states or {} - if not sub_agent_states: - return [] - - def _depth(aid: str, _memo: dict = {}) -> int: # noqa: B006 - if aid in _memo: - return _memo[aid] - parent = sub_agent_states.get(aid, {}).get("parent_id") - _memo[aid] = 0 if (parent is None or parent not in sub_agent_states) else 1 + _depth(parent, _memo) - return _memo[aid] - - restored_ids: list[str] = [] - for agent_id in sorted(sub_agent_states.keys(), key=_depth): - state_dict = sub_agent_states[agent_id] - try: - state = AgentState.model_validate(state_dict) - state.waiting_for_input = False - state.waiting_start_time = None - state.stop_requested = False - state.completed = False - state.llm_failed = False - state.sandbox_id = None - state.sandbox_token = None - state.sandbox_info = None - state.max_iterations = state.iteration + checkpoint_data.original_max_iterations - state.max_iterations_warning_sent = False - state.add_message( - "user", - f"[SYSTEM - SUB-AGENT RESUMED]\n" - f"You were interrupted at iteration {state.iteration}. " - f"Your sandbox has been reset and a fresh one will be created. " - f"Review your conversation history above and continue your task " - f"from where you left off. " - f"Call agent_finish only when your task is genuinely complete.", - ) - agent_cfg: dict[str, Any] = { - "llm_config": llm_config, - "state": state, - "is_resumed": True, - } - agent = StrixAgent(agent_cfg) - - def _run(a: Any = agent, s: Any = state) -> None: - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - try: - loop.run_until_complete(a.agent_loop(s.task)) - except Exception: # noqa: BLE001 - pass - finally: - loop.close() - - t = threading.Thread( - target=_run, - daemon=True, - name=f"ResumedAgent-{state.agent_name}-{agent_id[:8]}", - ) - t.start() - agents_graph_actions._running_agents[agent_id] = t - restored_ids.append(agent_id) - except Exception: # noqa: BLE001 - pass - - return restored_ids - - -def _build_root_resume_message( - state: Any, - checkpoint_data: Any, - restored_ids: "list[str] | None" = None, -) -> None: - """Inject the root agent's resume context message after sub-agents are known.""" - iteration = checkpoint_data.iteration - sub_agent_states: dict[str, Any] = checkpoint_data.sub_agent_states or {} - - if restored_ids: - lines = [ - "\n\nThe following sub-agents have been AUTOMATICALLY RESTORED and are " - "already running. Communicate with them using their original IDs:" - ] - for aid in restored_ids: - sd = sub_agent_states.get(aid, {}) - name = sd.get("agent_name", "sub-agent") - task = (sd.get("task") or "")[:200] - lines.append(f" \u2022 {name} (ID: {aid}) \u2014 task: {task}") - lines.append("\nDo NOT re-spawn these agents \u2014 they are already active.") - sub_agent_section = "\n".join(lines) - else: - dead: list[dict[str, Any]] = [] - for aid, node in (checkpoint_data.tracer_agents or {}).items(): - if node.get("parent_id") is None: - continue - status = node.get("status", "unknown") - if status not in ("completed", "finished", "stopped", "error", "failed"): - dead.append({"name": node.get("name", "sub-agent"), - "task": (node.get("task") or "")[:300]}) - if dead: - lines = [ - "\n\nThe following sub-agents were active at interruption but could " - "not be automatically restored. Re-spawn them if their work is incomplete:" - ] - for sa in dead: - lines.append(f" \u2022 {sa['name']} (task: {sa['task']})") - sub_agent_section = "\n".join(lines) - else: - sub_agent_section = "" - - msg = ( - f"[SYSTEM - SCAN RESUMED]\n" - f"This penetration test was interrupted at iteration {iteration}. " - f"A fresh sandbox will be created automatically.\n\n" - f"CRITICAL: Any agent IDs that appear in the conversation history ABOVE " - f"this message are from the old session and are DEAD. " - f"Only the agents listed below are currently alive." - f"{sub_agent_section}\n\n" - f"Review the conversation history to understand what has already been done, " - f"then CONTINUE the penetration test. " - f"Do NOT call finish_scan unless all testing is genuinely complete." - ) - state.add_message("user", msg) - +from strix.interface.checkpoint_restore import build_root_resume_message as _build_root_resume_message # noqa: E402 +from strix.interface.checkpoint_restore import restore_sub_agents as _restore_sub_agents_tui # noqa: E402 def get_package_version() -> str: try: return pkg_version("strix-agent") @@ -935,11 +807,16 @@ class StrixTUIApp(App): # type: ignore[misc] def _setup_cleanup_handlers(self) -> None: # Added for Resume Feature — save checkpoint on interrupt + _checkpoint_saved = threading.Event() + def _save_checkpoint_on_interrupt() -> None: mgr = self.agent_config.get("checkpoint_manager") agent = getattr(self, "_current_agent", None) if not mgr or not agent: return + if _checkpoint_saved.is_set(): + return + _checkpoint_saved.set() try: mgr.save( agent.state, diff --git a/strix/tools/agents_graph/agents_graph_actions.py b/strix/tools/agents_graph/agents_graph_actions.py index a5f5d51f..cdba4984 100644 --- a/strix/tools/agents_graph/agents_graph_actions.py +++ b/strix/tools/agents_graph/agents_graph_actions.py @@ -20,6 +20,9 @@ _agent_instances: dict[str, Any] = {} _agent_states: dict[str, Any] = {} +# Lock guarding concurrent reads/writes to _running_agents and _agent_instances +_agents_lock: threading.Lock = threading.Lock() + def _run_agent_in_thread( agent: Any, state: Any, inherited_messages: list[dict[str, Any]] @@ -86,8 +89,9 @@ def _run_agent_in_thread( _agent_graph["nodes"][state.agent_id]["status"] = "error" _agent_graph["nodes"][state.agent_id]["finished_at"] = datetime.now(UTC).isoformat() _agent_graph["nodes"][state.agent_id]["result"] = {"error": str(e)} - _running_agents.pop(state.agent_id, None) - _agent_instances.pop(state.agent_id, None) + with _agents_lock: + _running_agents.pop(state.agent_id, None) + _agent_instances.pop(state.agent_id, None) raise else: if state.stop_requested: @@ -96,8 +100,9 @@ def _run_agent_in_thread( _agent_graph["nodes"][state.agent_id]["status"] = "completed" _agent_graph["nodes"][state.agent_id]["finished_at"] = datetime.now(UTC).isoformat() _agent_graph["nodes"][state.agent_id]["result"] = result - _running_agents.pop(state.agent_id, None) - _agent_instances.pop(state.agent_id, None) + with _agents_lock: + _running_agents.pop(state.agent_id, None) + _agent_instances.pop(state.agent_id, None) return {"result": result} @@ -227,7 +232,8 @@ def create_agent( from strix.agents.state import AgentState from strix.llm.config import LLMConfig - parent_agent = _agent_instances.get(parent_id) + with _agents_lock: + parent_agent = _agent_instances.get(parent_id) timeout = None scan_mode = "deep" @@ -278,7 +284,8 @@ def create_agent( ) ] - _agent_instances[state.agent_id] = agent + with _agents_lock: + _agent_instances[state.agent_id] = agent thread = threading.Thread( target=_run_agent_in_thread, @@ -287,7 +294,8 @@ def create_agent( name=f"Agent-{name}-{state.agent_id}", ) thread.start() - _running_agents[state.agent_id] = thread + with _agents_lock: + _running_agents[state.agent_id] = thread except Exception as e: # noqa: BLE001 return {"success": False, "error": f"Failed to create agent: {e}", "agent_id": None} @@ -471,7 +479,8 @@ def agent_finish( parent_notified = True - _running_agents.pop(agent_id, None) + with _agents_lock: + _running_agents.pop(agent_id, None) return { "agent_completed": True, @@ -518,8 +527,9 @@ def stop_agent(agent_id: str) -> dict[str, Any]: agent_state = _agent_states[agent_id] agent_state.request_stop() - if agent_id in _agent_instances: - agent_instance = _agent_instances[agent_id] + with _agents_lock: + agent_instance = _agent_instances.get(agent_id) + if agent_instance is not None: if hasattr(agent_instance, "state"): agent_instance.state.request_stop() if hasattr(agent_instance, "cancel_current_execution"):