[Haik]: ckpt, pulled out shared mode resolution logic into its own helper class

This commit is contained in:
haikdc
2026-04-05 12:03:25 -07:00
parent 76a0b7c3ca
commit f9b1c8277a
3 changed files with 67 additions and 38 deletions
@@ -0,0 +1,48 @@
from typing import Optional, Tuple, List
from backend.apps.modes.modes import get_mode_by_id
from backend.apps.modes.Mode import Mode
from backend.apps.settings.settings import load_settings
from backend.apps.agents.ResolvedModeConfig.compose_system_prompt import compose_system_prompt
from backend.core.tools.shared_structs.Toolkit import Toolkit
from typeguard import typechecked
from pydantic import BaseModel, Field
class ResolvedModeConfig(BaseModel):
system_prompt: Optional[str] = None
allowed_tools: List[str] = Field(default_factory=list)
disallowed_tools: List[str] = Field(default_factory=list)
cwd: Optional[str] = None
@typechecked
@classmethod
async def create(
cls,
mode_id: str,
session_prompt: Optional[str],
toolkit: Toolkit,
) -> "ResolvedModeConfig":
settings = load_settings()
mode_def: Optional[Mode] = await get_mode_by_id(mode_id)
system_prompt: Optional[str] = compose_system_prompt(
global_default=settings.default_system_prompt,
mode_prompt=mode_def.system_prompt if mode_def else None,
session_prompt=session_prompt,
)
allowed_tools, disallowed_tools = toolkit.collect_tool_permissions()
if mode_def and mode_def.tools is not None:
mode_tool_set: set[str] = set[str](mode_def.tools)
disallowed_tools += [t for t in allowed_tools if t not in mode_tool_set]
allowed_tools = [t for t in allowed_tools if t in mode_tool_set]
cwd: Optional[str] = mode_def.default_folder if mode_def and mode_def.default_folder else None
return cls(
system_prompt=system_prompt,
allowed_tools=allowed_tools,
disallowed_tools=disallowed_tools,
cwd=cwd,
)
+19 -38
View File
@@ -23,14 +23,12 @@ from backend.core.Agent.Agent import Agent
from backend.core.db.PydanticStore import PydanticStore
from backend.core.shared_structs.agent.Message.Message import UserMessage
from backend.core.events.events import AgentStatusEvent, AgentClosedEvent, BranchSwitchedEvent
from backend.apps.agents.agent_utils.compose_system_prompt import compose_system_prompt
from backend.core.tools.make_builtin_toolkit.make_builtin_toolkit import make_builtin_toolkit
from backend.apps.agents.agent_utils.create_sdk_hooks import create_sdk_hooks
from backend.apps.agents.agent_utils.build_search_text import build_search_text
from backend.apps.agents.ResolvedModeConfig.ResolvedModeConfig import ResolvedModeConfig
from backend.apps.agents.COMMS_MANAGER.COMMS_MANAGER import COMMS_MANAGER
from backend.apps.settings.settings import load_settings
from backend.apps.modes.modes import get_mode_by_id
from backend.apps.modes.Mode import Mode
from backend.core.llm.resolve_sdk_env import resolve_sdk_env
from backend.ports import NINE_ROUTER_PORT
from claude_agent_sdk import ClaudeAgentOptions
@@ -104,14 +102,6 @@ class LaunchBody(BaseModel):
@agents.router.post("/launch")
async def launch(body: LaunchBody) -> dict:
settings = load_settings()
mode_def: Optional[Mode] = await get_mode_by_id(body.mode)
system_prompt = compose_system_prompt(
global_default=settings.default_system_prompt,
mode_prompt=mode_def.system_prompt if mode_def else None,
session_prompt=body.system_prompt or None,
)
agent: Agent = Agent(
model=body.model,
mode=body.mode,
@@ -124,15 +114,16 @@ async def launch(body: LaunchBody) -> dict:
toolkit: Toolkit = make_builtin_toolkit(agent, SESSIONS, COMMS_MANAGER.send_browser_command)
agent.toolkit = toolkit
mcp_servers: Dict[str, McpServerConfig] = toolkit.collect_mcp_servers()
allowed_tools, disallowed_tools = toolkit.collect_tool_permissions()
if mode_def and mode_def.tools is not None:
mode_tool_set: set[str] = set[str](mode_def.tools)
disallowed_tools += [t for t in allowed_tools if t not in mode_tool_set]
allowed_tools = [t for t in allowed_tools if t in mode_tool_set]
resolved_mode_config: ResolvedModeConfig = await ResolvedModeConfig.create(
mode_id=body.mode,
session_prompt=body.system_prompt or None,
toolkit=toolkit,
)
can_use_tool, pre_tool_hook, post_tool_hook = create_sdk_hooks(agent)
settings = load_settings()
env = resolve_sdk_env(
api_key=settings.anthropic_api_key,
nine_router_port=NINE_ROUTER_PORT if not settings.anthropic_api_key else None,
@@ -141,12 +132,12 @@ async def launch(body: LaunchBody) -> dict:
agent.config = ClaudeAgentOptions(
env=env,
model=body.model,
system_prompt=system_prompt,
system_prompt=resolved_mode_config.system_prompt,
max_turns=body.max_turns,
cwd=mode_def.default_folder if mode_def and mode_def.default_folder else None,
cwd=resolved_mode_config.cwd,
mcp_servers=mcp_servers if mcp_servers else None,
allowed_tools=allowed_tools,
disallowed_tools=disallowed_tools,
allowed_tools=resolved_mode_config.allowed_tools,
disallowed_tools=resolved_mode_config.disallowed_tools,
permission_mode="default",
can_use_tool=can_use_tool,
hooks={
@@ -213,27 +204,17 @@ async def send_message(session_id: str, body: MessageBody) -> dict:
agent.model = body.model # type: ignore[assignment]
if mode_changed or model_changed:
settings = load_settings()
mode_def: Optional[Mode] = await get_mode_by_id(agent.mode)
system_prompt = compose_system_prompt(
global_default=settings.default_system_prompt,
mode_prompt=mode_def.system_prompt if mode_def else None,
resolved_mode_config: ResolvedModeConfig = await ResolvedModeConfig.create(
mode_id=agent.mode,
session_prompt=agent.config.system_prompt,
toolkit=agent.toolkit,
)
allowed_tools, disallowed_tools = agent.toolkit.collect_tool_permissions()
if mode_def and mode_def.tools is not None:
mode_tool_set = set(mode_def.tools)
disallowed_tools += [t for t in allowed_tools if t not in mode_tool_set]
allowed_tools = [t for t in allowed_tools if t in mode_tool_set]
agent.config.system_prompt = system_prompt
agent.config.system_prompt = resolved_mode_config.system_prompt
agent.config.model = agent.model
agent.config.allowed_tools = allowed_tools
agent.config.disallowed_tools = disallowed_tools
if mode_def and mode_def.default_folder:
agent.config.cwd = mode_def.default_folder
agent.config.allowed_tools = resolved_mode_config.allowed_tools
agent.config.disallowed_tools = resolved_mode_config.disallowed_tools
if resolved_mode_config.cwd:
agent.config.cwd = resolved_mode_config.cwd
msg: UserMessage = UserMessage(
content=body.prompt,