Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 78 additions & 9 deletions src/ucode/agents/claude.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
from ucode.gateway_proxy import AI_GATEWAY_TOKEN_HEADER, AUTHORIZATION_HEADER, start_proxy
from ucode.launcher import exec_or_spawn
from ucode.managed_files import OS, current_os, write_managed_file
from ucode.smart_routing import v2 as smart_routing_v2
from ucode.smart_routing.claude_hooks import (
remove_smart_routing_hooks,
sync_smart_routing_hooks,
Expand All @@ -41,10 +42,11 @@
from ucode.tracing import tracing_env
from ucode.ui import print_err, print_note, print_success, print_warning

GATEWAY_MODEL_DISCOVERY_ENV_VAR = "ENABLE_CLAUDE_CODE_GATEWAY_MODEL_DISCOVERY"
CLAUDE_CONFIG_DIR = Path.home() / ".claude"
CLAUDE_SETTINGS_PATH = CLAUDE_CONFIG_DIR / "ucode-settings.json"
CLAUDE_USER_SETTINGS_PATH = CLAUDE_CONFIG_DIR / "settings.json"
CLAUDE_BACKUP_PATH = APP_DIR / "claude-ucode-settings.backup.json"
GATEWAY_MODEL_DISCOVERY_ENV_VAR = "ENABLE_CLAUDE_CODE_GATEWAY_MODEL_DISCOVERY"

SPEC: ToolSpec = {
"binary": "claude",
Expand Down Expand Up @@ -876,10 +878,40 @@ def _merge_claude_settings(base: dict, overlay: dict) -> dict:
return merged


def _compose_v2_settings(tool_args: list[str]) -> tuple[dict, list[str]]:
"""Compose caller settings with ucode's Claude settings for a v2 launch."""
caller_values, remaining = _extract_caller_settings(tool_args)
settings: dict = {}
for value in caller_values:
settings = _merge_claude_settings(settings, _load_caller_settings(value))
return _merge_claude_settings(settings, read_json_safe(CLAUDE_SETTINGS_PATH)), remaining


def _original_launch_model(snapshot: dict, state: dict) -> str | None:
override = state.get("_claude_launch_model")
if isinstance(override, str) and override.strip():
return override.strip()
value = snapshot.get("value") if snapshot.get("present") is True else None
if isinstance(value, str) and value.strip():
return value.strip()
return default_model(state)


def _has_explicit_model_arg(tool_args: list[str]) -> bool:
return any(arg in {"--model", "-m"} or arg.startswith("--model=") for arg in tool_args)


def _launch_model_args(tool_args: list[str], launch_model: str | None) -> list[str]:
if not launch_model or _has_explicit_model_arg(tool_args):
return []
return ["--model", launch_model]


def _build_claude_argv(
binary: str,
tool_args: list[str],
relayed: bool = False,
launch_model: str | None = None,
settings_override: dict | None = None,
) -> list[str]:
"""Build the ``claude`` argv, composing any caller ``--settings`` with
Expand All @@ -903,11 +935,19 @@ def _build_claude_argv(
filter through and shadow the subscription OAuth.
"""
source_args = ["--setting-sources", _RELAYED_SETTING_SOURCES] if relayed else []
model_args = _launch_model_args(tool_args, launch_model)
caller_values, remaining = _extract_caller_settings(tool_args)
if not caller_values and settings_override is None:
# No caller --settings: hand Claude ucode's settings file directly (the
# common path; behavior unchanged).
return [binary, *source_args, "--settings", str(CLAUDE_SETTINGS_PATH), *tool_args]
return [
binary,
*source_args,
"--settings",
str(CLAUDE_SETTINGS_PATH),
*model_args,
*tool_args,
]
caller_settings: dict = {}
for value in caller_values:
caller_settings = _merge_claude_settings(caller_settings, _load_caller_settings(value))
Expand All @@ -921,6 +961,7 @@ def _build_claude_argv(
*source_args,
"--settings",
json.dumps(merged, separators=(",", ":")),
*model_args,
*remaining,
]

Expand Down Expand Up @@ -1029,7 +1070,10 @@ def _launch_relayed(state: dict, binary: str, tool_args: list[str]) -> None:
raise SystemExit(returncode)


def _launch_gateway(state: dict, binary: str, tool_args: list[str]) -> None:
def _launch_gateway(
state: dict, binary: str, tool_args: list[str], launch_model: str | None
) -> None:
"""Launch discovery-enabled Claude through a refreshing gateway proxy."""
workspace = state["workspace"]
server, cache, client = start_proxy(
workspace,
Expand All @@ -1046,11 +1090,14 @@ def _launch_gateway(state: dict, binary: str, tool_args: list[str]) -> None:

server_thread = threading.Thread(target=server.serve_forever, daemon=True)
server_thread.start()
settings_override = {
"env": {"ANTHROPIC_BASE_URL": os.environ["ANTHROPIC_BASE_URL"]},
}
settings_override = {"env": {"ANTHROPIC_BASE_URL": os.environ["ANTHROPIC_BASE_URL"]}}
proc = subprocess.Popen(
_build_claude_argv(binary, tool_args, settings_override=settings_override)
_build_claude_argv(
binary,
tool_args,
launch_model=launch_model,
settings_override=settings_override,
)
)
try:
returncode = proc.wait()
Expand All @@ -1067,15 +1114,37 @@ def _launch_gateway(state: dict, binary: str, tool_args: list[str]) -> None:
def launch(state: dict, tool_args: list[str]) -> None:
binary = SPEC["binary"]
workspace = state.get("workspace")
# Recover a prior switch interrupted before its surgical settings restore.
smart_routing_v2.recover_claude_model_snapshots(CLAUDE_USER_SETTINGS_PATH)
model_snapshot = smart_routing_v2.snapshot_claude_model_setting(CLAUDE_USER_SETTINGS_PATH)
launch_model = _original_launch_model(model_snapshot, state)
if state.get("claude_relayed"):
_launch_relayed(state, binary, tool_args)
return
# Smart routing v2 needs Unix PTY support, which Windows does not provide.
if smart_routing_v2.enabled() and os.name == "nt":
raise RuntimeError(
"Smart routing in Claude Code is currently not supported on Windows. "
"Please use Codex or disable smart routing."
)
if smart_routing_v2.enabled() and workspace:
smart_routing_v2.launch_claude(
state,
tool_args,
binary=binary,
user_settings_path=CLAUDE_USER_SETTINGS_PATH,
model_snapshot=model_snapshot,
launch_model=launch_model,
compose_settings=_compose_v2_settings,
launch_model_args=_launch_model_args,
)
return
if workspace and os.environ.get(GATEWAY_MODEL_DISCOVERY_ENV_VAR) == "1":
_launch_gateway(state, binary, tool_args)
_launch_gateway(state, binary, tool_args, launch_model)
return
if workspace:
os.environ["OAUTH_TOKEN"] = get_databricks_token(workspace, state.get("profile"))
exec_or_spawn(_build_claude_argv(binary, tool_args))
exec_or_spawn(_build_claude_argv(binary, tool_args, launch_model=launch_model))


def validate_cmd(binary: str) -> list[str]:
Expand Down
34 changes: 33 additions & 1 deletion src/ucode/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,7 @@
download_managed_skills_on_launch,
)
from ucode.smart_routing import claude_routing, codex_routing
from ucode.smart_routing import v2 as smart_routing_v2
from ucode.state import (
STATE_PATH,
clear_state,
Expand Down Expand Up @@ -1403,6 +1404,7 @@ def claude_router_hook_cmd(
profile: Annotated[str | None, typer.Option("--profile")] = None,
use_pat: Annotated[bool, typer.Option("--use-pat")] = False,
model: Annotated[list[str] | None, typer.Option("--model")] = None,
socket_path: Annotated[str | None, typer.Option("--socket")] = None,
) -> None:
"""Run a Claude Code smart-routing lifecycle hook."""
import json
Expand All @@ -1420,6 +1422,24 @@ def claude_router_hook_cmd(
return
if not isinstance(payload, dict):
return
if event == "route-first-prompt":
if not socket_path:
from ucode.smart_routing.claude_hooks import FIRST_PROMPT_SOCKET_ENV

socket_path = os.environ.get(FIRST_PROMPT_SOCKET_ENV)
if not socket_path:
return
from pathlib import Path

from ucode.smart_routing.claude_pty import (
first_prompt_hook_output,
request_first_prompt_route,
)

output = first_prompt_hook_output(request_first_prompt_route(Path(socket_path), payload))
if output is not None:
sys.stdout.write(json.dumps(output))
return
if event == "session-start":
record_session_start(payload)
return
Expand Down Expand Up @@ -1862,7 +1882,12 @@ def _launch_tool(
managed_launch_model(managed, recommendation, tool) if managed is not None else None
)
state, resolved_model = resolve_launch_model(tool, state, managed_model)
if routing_agent is not None and routing_agent.smart_routing_enabled(state):
first_prompt_routes_claude = tool == "claude" and smart_routing_v2.enabled()
if (
routing_agent is not None
and routing_agent.smart_routing_enabled(state)
and not first_prompt_routes_claude
):
display = TOOL_SPECS[tool]["display"]
with spinner(f"Selecting a {display} model with smart routing..."):
decision, routing_error = _ROUTING_MODULES[tool].route_launch_model(
Expand Down Expand Up @@ -1957,6 +1982,13 @@ def _launch_tool(
if managed is not None and not is_dry_run():
_register_managed_mcp_servers(managed, tool, state)
_apply_managed_skills(managed, tool, state)
if tool == "claude":
# Transient launch precedence for claude.py's universal --model flag.
# An explicit choice wins, followed by a routed/managed root pick;
# neither value is persisted into workspace state.
launch_model = model or route_root_model
if launch_model:
state["_claude_launch_model"] = launch_model
print_success(f"Starting {TOOL_SPECS[tool]['display']}")
launch_agent(tool, state, ctx.args)
except RuntimeError as exc:
Expand Down
19 changes: 19 additions & 0 deletions src/ucode/smart_routing/claude_hooks.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,8 @@
from ucode.smart_routing import hooks

ROUTING_HOOK_COMMAND_MARKER = "claude-router-hook"
FIRST_PROMPT_HOOK_MARKER = "claude-router-hook route-first-prompt"
FIRST_PROMPT_SOCKET_ENV = "UCODE_CLAUDE_V2_SOCKET"


def sync_smart_routing_hooks(doc: dict, state: dict, *, enabled: bool) -> None:
Expand All @@ -28,6 +30,23 @@ def remove_smart_routing_hooks(doc: dict) -> bool:
return hooks.remove_managed_hooks(doc, ROUTING_HOOK_COMMAND_MARKER)


def sync_first_prompt_hook(doc: dict, executable: str) -> None:
"""Add the first-prompt hook to a per-launch settings document."""
groups = {
"UserPromptSubmit": [
{
"hooks": [
_routing_command_hook(
[executable, ROUTING_HOOK_COMMAND_MARKER, "route-first-prompt"],
status="Selecting a model with Smart Routing",
)
]
}
]
}
hooks.sync_managed_hooks(doc, FIRST_PROMPT_HOOK_MARKER, groups)


def _routing_hook_groups(state: dict) -> dict[str, list[dict]]:
route_argv = _routing_hook_argv(state, "route-subagent")
session_argv = _routing_hook_argv(state, "session-start")
Expand Down
Loading
Loading