Merge pull request #27 from HarnessLab/update_17_4

slash commands
This commit is contained in:
Abdelrahman Abdallah
2026-04-20 01:16:31 +02:00
committed by GitHub
40 changed files with 5717 additions and 46 deletions
+43 -36
View File
@@ -117,6 +117,7 @@ Done:
- [x] Filesystem-backed custom agent discovery from `~/.claude/agents` and `./.claude/agents`
- [x] Active agent override precedence across built-in, user, and project agent definitions
- [x] Custom agent resolution in the `Agent` tool with model, tool-filter, and initial-prompt support
- [x] Local custom-agent file creation, update, and deletion flows for project/user agent definitions
Missing:
@@ -145,6 +146,9 @@ Done:
- [x] `agent-context-raw` command
- [x] `token-budget` command
- [x] `agents` command
- [x] `agents-create` command
- [x] `agents-update` command
- [x] `agents-delete` command
- [x] Local background session mode
- [x] Local background session listing (`agent-ps`)
- [x] Local background session logs (`agent-logs`)
@@ -171,14 +175,14 @@ Missing:
- [ ] Self-hosted runner mode
- [ ] tmux fast paths
- [ ] Worktree fast paths at the CLI entrypoint level
- [ ] Node.js version check and platform setup from `setup.ts`
- [x] Python (Node.js equivalent) version check and platform detection from `setup.ts`
- [ ] Worktree creation/setup from `setup.ts`
- [ ] Terminal backup/restore from `setup.ts`
- [ ] Release notes checking from `setup.ts`
- [x] Release notes checking from `setup.ts` (local CHANGELOG.md, no network/cache layer)
- [ ] Full `entrypoints/cli.tsx` parity (version flag, feature flags, env setup, dynamic imports)
- [ ] Full `entrypoints/init.ts` parity (settings validation, OAuth, policy limits, telemetry, cleanup handlers)
- [ ] SDK entrypoint (`entrypoints/sdk/` — controlTypes, coreTypes, runtimeTypes, settingsTypes, toolTypes)
- [ ] Sandbox types/network config schema (`entrypoints/sandboxTypes.ts`)
- [ ] SDK entrypoint (`entrypoints/sdk/` — controlTypes, coreTypes, runtimeTypes, settingsTypes, toolTypes) — partial: HOOK_EVENTS, EXIT_REASONS, ModelUsage, ThinkingConfig, MCP server configs, JsonSchemaOutputFormat ported in `src/sdk_core_types.py`
- [x] Sandbox types/network config schema (`entrypoints/sandboxTypes.ts`)
## 3. Prompt Assembly
@@ -317,56 +321,56 @@ Done (53 slash command names in 37 specs):
Missing npm slash commands (from `src/commands/` — 80+ commands total):
- [x] `/add-dir` — Add a new working directory
- [x] `/agents` — Inspect local agent configurations and show active definitions
- [x] `/agents` — Inspect, create, update, and delete local agent definitions
- [x] `/branch` — Create a branch of the current conversation
- [ ] `/bridge` — Connect for remote-control sessions
- [ ] `/btw` — Quick side question without interrupting main conversation
- [ ] `/chrome` — Chrome extension settings
- [x] `/bridge` — Connect for remote-control sessions (read-only status in this runtime)
- [x] `/btw` — Quick side question without interrupting main conversation
- [x] `/chrome` — Chrome extension settings
- [x] `/color` — Set the prompt bar color for this session
- [x] `/compact` — Clear history but keep a summary in context
- [x] `/copy` — Copy Claude's last response to clipboard
- [x] `/cost` — Show total cost and duration of session
- [ ] `/desktop` — Continue session in Claude Desktop
- [x] `/desktop` — Continue session in Claude Desktop
- [x] `/diff` — View uncommitted changes and per-turn diffs
- [x] `/doctor` — Diagnose and verify installation and settings
- [x] `/effort` — Set effort level for model usage
- [x] `/exit` — Exit the REPL
- [x] `/export` — Export conversation to file or clipboard
- [ ] `/extra-usage` — Configure extra usage for rate limits
- [x] `/extra-usage` — Configure extra usage for rate limits
- [x] `/fast` — Toggle fast mode
- [ ] `/feedback` — Submit feedback
- [x] `/feedback` — Submit feedback (alias `/bug`)
- [x] `/files` — List all files currently in context
- [ ] `/ide` — Manage IDE integrations and show status
- [ ] `/install-github-app` — Set up GitHub Actions
- [ ] `/install-slack-app` — Install Slack app
- [ ] `/keybindings` — Open keybindings config file
- [ ] `/mobile`QR code for mobile app
- [ ] `/output-style`Change output style
- [ ] `/passes` — Passes management
- [ ] `/plugin` — Plugin management
- [x] `/ide` — Manage IDE integrations and show status
- [x] `/install-github-app` — Set up GitHub Actions
- [x] `/install-slack-app` — Install Slack app
- [x] `/keybindings` — Open keybindings config file
- [x] `/mobile`Mobile app store links (aliases `/ios`, `/android`)
- [x] `/output-style`Deprecation pointer to `/config`
- [x] `/passes` — Passes management
- [x] `/plugin` — Plugin management (read-only listing)
- [x] `/pr-comments`, `/pr_comments` — Get comments from a GitHub PR (prompt-type)
- [ ] `/privacy-settings` — View/update privacy settings
- [ ] `/rate-limit-options` — Show options when rate limited
- [ ] `/release-notes` — View release notes
- [ ] `/reload-plugins` — Activate pending plugin changes
- [ ] `/remote-env` — Configure default remote environment
- [ ] `/remote-setup` — Remote setup configuration
- [x] `/privacy-settings` — View/update privacy settings
- [x] `/rate-limit-options` — Show options when rate limited
- [x] `/release-notes` — View release notes
- [x] `/reload-plugins` — Activate pending plugin changes
- [x] `/remote-env` — Configure default remote environment
- [x] `/remote-setup` — Remote setup configuration (gh auth status + Claude.ai/code link)
- [x] `/rename` — Rename current conversation
- [x] `/resume`, `/continue` — Resume a previous conversation
- [x] `/rewind`, `/checkpoint` — Restore code/conversation to a previous point
- [ ] `/sandbox-toggle` — Toggle sandbox mode
- [x] `/skills` — List available skills
- [x] `/sandbox-toggle` — Toggle sandbox mode (alias `/sandbox`)
- [x] `/skills` — List available bundled skills (mirrors `commands/skills/SkillsMenu.tsx`; lists `bundled_skills.BUNDLED_SKILLS`, not slash commands)
- [x] `/stats` — Usage statistics and activity
- [ ] `/stickers` — Order stickers
- [x] `/stickers` — Order stickers
- [x] `/tag` — Toggle a searchable tag on the session
- [ ] `/theme` — Change the theme
- [ ] `/upgrade` — Upgrade to Max
- [x] `/theme` — Change the theme
- [x] `/upgrade` — Upgrade to Max
- [x] `/vim` — Toggle Vim/Normal editing modes
- [ ] `/voice` — Toggle voice mode
- [x] `/voice` — Toggle voice mode
- [ ] Feature-gated: `/buddy`, `/fork`, `/peers`, `/proactive`, `/torch`, `/workflows` (full), etc.
- [ ] Internal: `/backfill-sessions`, `/break-cache`, `/bughunter`, `/commit-push-pr`, `/init-verifiers`, `/mock-limits`, `/version`, `/ultraplan`, `/autofix-pr`, etc.
- [x] `/commit` — Create a git commit (prompt-type with injected git context)
- [ ] Full `/agents` parity for create/edit/delete flows and multi-source management UI
- [ ] Full `/agents` parity for interactive TUI/editor flows, plugin sources, and full multi-source management UX
## 6. Built-in Tools
@@ -702,21 +706,24 @@ Done:
- [x] Basic file operations in tool implementations
- [x] Basic git status snapshot
- [x] Basic shell/subprocess handling
- [x] Bundled small portable utilities — `utils/array.ts`, `utils/set.ts`, `utils/objectGroupBy.ts`, `utils/xml.ts`, `utils/uuid.ts` ported in `src/small_utils.py`
- [x] Session-scoped env-var registry (`utils/sessionEnvVars.ts`) ported in `src/session_env_vars.py`, merged into spawned subprocess env via `_build_subprocess_env` (mirrors `utils/shell/bashProvider.ts`), and dropped on `/clear` via `clear_runtime_state` (mirrors `commands/clear/caches.ts`)
- [x] Display formatters from `utils/format.ts` (`formatFileSize`, `formatSecondsShort`, `formatDuration`, `formatNumber`, `formatTokens`) ported in `src/format_utils.py`
Missing major utility categories:
- [ ] Shell utilities (`utils/bash/`, `utils/shell/`, `Shell.ts`, `ShellCommand.ts`)
- [ ] Git operations (`utils/git.ts`, `utils/gitDiff.ts`, `utils/gitSettings.ts`, `utils/commitAttribution.ts`)
- [ ] Git operations (`utils/git.ts`, `utils/gitDiff.ts`, `utils/gitSettings.ts`, `utils/commitAttribution.ts`) — partial: `findGitRoot`, `normalizeGitRemoteUrl`, `getRepoRemoteHash`, and `shouldIncludeGitInstructions` ported in `src/git_utils.py`
- [ ] File operations (`utils/file.ts`, `utils/fileRead.ts`, `utils/fileHistory.ts`, `utils/fileStateCache.ts`, `utils/fsOperations.ts`, `utils/ripgrep.ts`, `utils/glob.ts`)
- [ ] AI/Model utilities (`utils/modelCost.ts`, `utils/model/`, `utils/context.ts`, `utils/queryContext.ts`)
- [ ] AI/Model utilities (`utils/modelCost.ts`, `utils/model/`, `utils/context.ts`, `utils/queryContext.ts`) — partial: modelCost ported in `src/model_cost.py`
- [ ] Config/Settings (`utils/config.ts`, `utils/settings/`)
- [ ] Message handling (`utils/messages.ts`, `utils/messages/`, `utils/messageQueueManager.ts`)
- [ ] API/Network (`utils/api.ts`, `utils/http.ts`, `utils/proxy.ts`, `utils/auth.ts`)
- [ ] Session management (`utils/sessionStorage.ts`, `utils/sessionState.ts`, `utils/sessionStart.ts`, `utils/sessionRestore.ts`)
- [ ] Plugin/Skill utilities (`utils/plugins/`, `utils/skills/`)
- [ ] Memory/Context (`utils/memory/`, `utils/claudemd.ts`, `utils/contextAnalysis.ts`)
- [ ] IDE integration (`utils/ide.ts`, `utils/jetbrains.ts`)
- [ ] Platform/OS (`utils/platform.ts`, `utils/terminal.ts`, `utils/systemDirectories.ts`)
- [ ] IDE integration (`utils/ide.ts`, `utils/jetbrains.ts`) — partial: `utils/idePathConversion.ts` ported in `src/ide_path_conversion.py` (`WindowsToWSLConverter`, `checkWSLDistroMatch`)
- [ ] Platform/OS (`utils/platform.ts`, `utils/terminal.ts`, `utils/systemDirectories.ts`) — partial: platform detection (`getPlatform`, `getWslVersion`, `getLinuxDistroInfo`, `detectVcs`) and `getSystemDirectories` ported in `src/platform_info.py`
- [ ] Debugging (`utils/debug.ts`, `utils/diagLogs.ts`, `utils/log.ts`, `utils/profilerBase.ts`)
- [ ] Telemetry (`utils/telemetry/`)
- [ ] Deep link utilities (`utils/deepLink/`)
+31 -1
View File
@@ -152,6 +152,7 @@ Built on the public porting workspace from [instructkr/claw-code](https://github
- [x] Nested agent delegation with dependency-aware topological batching
- [x] Agent manager with lineage tracking and group membership
- [x] Filesystem-backed custom agent profiles with built-in/user/project precedence
- [x] Local custom-agent create/update/delete flows via CLI and `/agents`
- [x] Local daemon-style background command family
- [x] Local background session workflows: `agent-bg`, `agent-ps`, `agent-logs`, `agent-attach`, `agent-kill`
- [x] Local remote runtime: manifest discovery, profile listing, connect/disconnect persistence, and CLI/slash flows
@@ -451,6 +452,9 @@ python3 -m src.main agent \
| `agent-context-raw` | Show the raw context snapshot |
| `token-budget` | Show prompt-window budget, reserves, and soft/hard input limits |
| `agents [agent_type]` | List active local agent definitions or show one agent profile |
| `agents-create <agent_type>` | Create a project or user agent definition markdown file |
| `agents-update <agent_type>` | Update an existing project or user agent definition |
| `agents-delete <agent_type>` | Delete an existing project or user agent definition |
| `agent-resume <id> <prompt>` | Resume a saved session |
### Runtime Utility Commands
@@ -550,7 +554,7 @@ These are handled **locally** before the model loop:
| `/permissions` | — | Show active tool permission mode |
| `/model` | — | Show or update the active model |
| `/tools` | — | List registered tools with permission status |
| `/agents` | — | List active local agent definitions or show one profile |
| `/agents` | — | List, show, create, update, or delete local agent definitions |
| `/memory` | — | Show loaded CLAUDE.md memory bundle |
| `/status` | `/session` | Show runtime/session status summary |
| `/clear` | — | Clear ephemeral runtime state |
@@ -596,6 +600,32 @@ python3 -m src.main agent "/agents" --cwd .
python3 -m src.main agent "/agents show reviewer" --cwd .
```
Create, update, or delete agent files from the CLI:
```bash
python3 -m src.main agents-create reviewer \
--cwd . \
--description "Review implementation changes carefully." \
--prompt "Inspect code changes and summarize risks." \
--tools read_file,grep_search \
--model Qwen/Qwen3-Coder-30B-A3B-Instruct
python3 -m src.main agents-update reviewer \
--cwd . \
--description "Review implementation changes and tests carefully." \
--prompt "Focus on regressions, missing tests, and risky diffs."
python3 -m src.main agents-delete reviewer --cwd . --source project
```
Or use the local slash command management forms:
```bash
python3 -m src.main agent "/agents create reviewer :: Review implementation changes carefully. :: Inspect code changes and summarize risks." --cwd .
python3 -m src.main agent "/agents update reviewer Updated review description :: Focus on regressions and missing tests." --cwd .
python3 -m src.main agent "/agents delete reviewer" --cwd .
```
### Utility Commands
```bash
+6
View File
@@ -621,6 +621,9 @@ python3 -m src.main agent "/model demo-model" --cwd ./test_cases
python3 -m src.main agent "/tools" --cwd ./test_cases
python3 -m src.main agent "/agents" --cwd ./test_cases_agents
python3 -m src.main agent "/agents show reviewer" --cwd ./test_cases_agents
python3 -m src.main agent "/agents create reviewer-temp :: Review temporary code changes. :: Inspect code and summarize risks." --cwd ./test_cases_agents
python3 -m src.main agent "/agents update reviewer-temp Updated temp reviewer :: Focus on regressions and missing tests." --cwd ./test_cases_agents
python3 -m src.main agent "/agents delete reviewer-temp" --cwd ./test_cases_agents
python3 -m src.main agent "/memory" --cwd ./test_cases
python3 -m src.main agent "/status" --cwd ./test_cases
python3 -m src.main agent "/session" --cwd ./test_cases
@@ -700,6 +703,9 @@ python3 -m src.main agent-context-raw --cwd ./test_cases
python3 -m src.main token-budget --cwd ./test_cases
python3 -m src.main agents --cwd ./test_cases_agents
python3 -m src.main agents reviewer --cwd ./test_cases_agents
python3 -m src.main agents-create reviewer-cli --cwd ./test_cases_agents --description "CLI-created reviewer agent." --prompt "Inspect code changes and summarize risks."
python3 -m src.main agents-update reviewer-cli --cwd ./test_cases_agents --description "Updated CLI reviewer." --prompt "Focus on regressions and missing tests."
python3 -m src.main agents-delete reviewer-cli --cwd ./test_cases_agents --source project
```
### 6.2 Extra working directories and `CLAUDE.md` toggle
+14
View File
@@ -13,11 +13,18 @@ from .agent_context import (
from .agent_manager import AgentManager
from .agent_registry import (
AgentLoadError,
AgentMutationResult,
AgentRegistrySnapshot,
create_agent_definition,
delete_agent_definition,
find_agent_definition,
load_agent_registry,
normalize_mutable_source,
render_agent_detail,
render_agent_mutation,
render_agents_report,
scaffold_agent_definition,
update_agent_definition,
)
from .agent_runtime import LocalCodingAgent
from .agent_session import AgentMessage, AgentSessionState
@@ -51,6 +58,7 @@ __all__ = [
'AgentContextSnapshot',
'AgentManager',
'AgentLoadError',
'AgentMutationResult',
'AgentPermissions',
'AgentRegistrySnapshot',
'AgentRunResult',
@@ -119,8 +127,10 @@ __all__ = [
'clear_context_caches',
'clear_token_counter_cache',
'count_tokens',
'create_agent_definition',
'calculate_token_budget',
'default_tool_registry',
'delete_agent_definition',
'describe_token_counter',
'estimate_chat_overhead',
'execute_tool',
@@ -130,9 +140,13 @@ __all__ = [
'get_user_context',
'load_agent_registry',
'load_session',
'normalize_mutable_source',
'render_agent_detail',
'render_agent_mutation',
'render_agents_report',
'run_parity_audit',
'scaffold_agent_definition',
'save_session',
'set_system_prompt_injection',
'update_agent_definition',
]
+292
View File
@@ -16,6 +16,13 @@ _SOURCE_ORDER = {
'userSettings': 1,
'projectSettings': 2,
}
_MUTABLE_SOURCE_ALIASES = {
'project': 'projectSettings',
'projectSettings': 'projectSettings',
'user': 'userSettings',
'userSettings': 'userSettings',
}
_UNSET = object()
@dataclass(frozen=True)
@@ -33,6 +40,15 @@ class AgentRegistrySnapshot:
failed_files: tuple[AgentLoadError, ...]
@dataclass(frozen=True)
class AgentMutationResult:
action: str
agent_type: str
source: str
file_path: str
overwritten: bool = False
def load_agent_registry(cwd: Path) -> AgentRegistrySnapshot:
builtin_agents = tuple(get_builtin_agents())
loaded_agents: list[AgentDefinition] = list(builtin_agents)
@@ -258,6 +274,274 @@ def render_agent_detail(snapshot: AgentRegistrySnapshot, agent_type: str) -> str
return '\n'.join(lines)
def normalize_mutable_source(source: str | None, *, allow_auto: bool = False) -> str:
if source is None:
return 'auto' if allow_auto else 'projectSettings'
normalized = source.strip()
if allow_auto and normalized in {'', 'auto'}:
return 'auto'
resolved = _MUTABLE_SOURCE_ALIASES.get(normalized)
if resolved is None:
choices = ', '.join(sorted(_MUTABLE_SOURCE_ALIASES))
if allow_auto:
choices = 'auto, ' + choices
raise ValueError(f'Unsupported agent source: {source}. Expected one of: {choices}')
return resolved
def format_agent_markdown(
*,
agent_type: str,
description: str,
system_prompt: str,
tools: tuple[str, ...] | None = None,
model: str | None = None,
color: str | None = None,
permission_mode: str | None = None,
max_turns: int | None = None,
initial_prompt: str | None = None,
background: bool = False,
one_shot: bool = False,
omit_claude_md: bool = False,
) -> str:
lines = [
'---',
f'name: {agent_type}',
f'description: "{_escape_frontmatter_text(description)}"',
]
if tools is not None:
if tools:
lines.append(f'tools: {", ".join(tools)}')
else:
lines.append('tools: []')
if model:
lines.append(f'model: {model}')
if color:
lines.append(f'color: {color}')
if permission_mode:
lines.append(f'permissionMode: {permission_mode}')
if max_turns is not None:
lines.append(f'maxTurns: {max_turns}')
if initial_prompt:
lines.append(f'initialPrompt: "{_escape_frontmatter_text(initial_prompt)}"')
if background:
lines.append('background: true')
if one_shot:
lines.append('oneShot: true')
if omit_claude_md:
lines.append('omitClaudeMd: true')
lines.extend(['---', '', system_prompt.strip(), ''])
return '\n'.join(lines)
def create_agent_definition(
cwd: Path,
*,
agent_type: str,
description: str,
system_prompt: str,
source: str = 'projectSettings',
overwrite: bool = False,
tools: tuple[str, ...] | None = None,
model: str | None = None,
color: str | None = None,
permission_mode: str | None = None,
max_turns: int | None = None,
initial_prompt: str | None = None,
background: bool = False,
one_shot: bool = False,
omit_claude_md: bool = False,
) -> AgentMutationResult:
resolved_source = normalize_mutable_source(source)
file_path = get_agent_file_path(cwd, resolved_source, agent_type)
file_path.parent.mkdir(parents=True, exist_ok=True)
existed_before = file_path.exists()
if existed_before and not overwrite:
raise ValueError(f'Agent file already exists: {file_path}')
file_path.write_text(
format_agent_markdown(
agent_type=agent_type,
description=description,
system_prompt=system_prompt,
tools=tools,
model=model,
color=color,
permission_mode=permission_mode,
max_turns=max_turns,
initial_prompt=initial_prompt,
background=background,
one_shot=one_shot,
omit_claude_md=omit_claude_md,
),
encoding='utf-8',
)
return AgentMutationResult(
action='created',
agent_type=agent_type,
source=resolved_source,
file_path=str(file_path),
overwritten=existed_before and overwrite,
)
def update_agent_definition(
cwd: Path,
*,
agent_type: str,
source: str = 'auto',
description: str | object = _UNSET,
system_prompt: str | object = _UNSET,
tools: tuple[str, ...] | None | object = _UNSET,
model: str | None | object = _UNSET,
color: str | None | object = _UNSET,
permission_mode: str | None | object = _UNSET,
max_turns: int | None | object = _UNSET,
initial_prompt: str | None | object = _UNSET,
background: bool | object = _UNSET,
one_shot: bool | object = _UNSET,
omit_claude_md: bool | object = _UNSET,
) -> AgentMutationResult:
resolved_source = normalize_mutable_source(source, allow_auto=True)
snapshot = load_agent_registry(cwd)
target = find_mutable_agent(snapshot, agent_type, source=resolved_source)
if target is None:
raise ValueError(f'No editable agent definition found for: {agent_type}')
file_path = get_agent_file_path(
cwd,
target.source,
target.agent_type,
filename=target.filename,
)
file_path.parent.mkdir(parents=True, exist_ok=True)
file_path.write_text(
format_agent_markdown(
agent_type=target.agent_type,
description=target.when_to_use if description is _UNSET else str(description),
system_prompt=target.system_prompt if system_prompt is _UNSET else str(system_prompt),
tools=target.tools if tools is _UNSET else tools,
model=target.model if model is _UNSET else model,
color=target.color if color is _UNSET else color,
permission_mode=(
target.permission_mode if permission_mode is _UNSET else permission_mode
),
max_turns=target.max_turns if max_turns is _UNSET else max_turns,
initial_prompt=target.initial_prompt if initial_prompt is _UNSET else initial_prompt,
background=target.background if background is _UNSET else bool(background),
one_shot=target.one_shot if one_shot is _UNSET else bool(one_shot),
omit_claude_md=(
target.omit_claude_md if omit_claude_md is _UNSET else bool(omit_claude_md)
),
),
encoding='utf-8',
)
return AgentMutationResult(
action='updated',
agent_type=target.agent_type,
source=target.source,
file_path=str(file_path),
)
def delete_agent_definition(
cwd: Path,
*,
agent_type: str,
source: str = 'auto',
) -> AgentMutationResult:
resolved_source = normalize_mutable_source(source, allow_auto=True)
snapshot = load_agent_registry(cwd)
target = find_mutable_agent(snapshot, agent_type, source=resolved_source)
if target is None:
raise ValueError(f'No editable agent definition found for: {agent_type}')
file_path = get_agent_file_path(
cwd,
target.source,
target.agent_type,
filename=target.filename,
)
if not file_path.exists():
raise ValueError(f'Agent file does not exist: {file_path}')
file_path.unlink()
return AgentMutationResult(
action='deleted',
agent_type=target.agent_type,
source=target.source,
file_path=str(file_path),
)
def scaffold_agent_definition(
cwd: Path,
*,
agent_type: str,
source: str = 'projectSettings',
overwrite: bool = False,
description: str | None = None,
system_prompt: str | None = None,
) -> AgentMutationResult:
resolved_description = description or f'Use this agent when the task calls for {agent_type}.'
resolved_prompt = system_prompt or (
f'You are the {agent_type} agent.\n'
'Read the task carefully, use the available tools deliberately, and return a concise result.'
)
return create_agent_definition(
cwd,
agent_type=agent_type,
description=resolved_description,
system_prompt=resolved_prompt,
source=source,
overwrite=overwrite,
)
def find_mutable_agent(
snapshot: AgentRegistrySnapshot,
agent_type: str,
*,
source: str = 'auto',
) -> AgentDefinition | None:
if source == 'auto':
candidates = [
agent
for agent in snapshot.all_agents
if agent.agent_type == agent_type and agent.source in _MUTABLE_SOURCE_ALIASES.values()
]
if not candidates:
return None
return max(candidates, key=lambda agent: _source_rank(agent.source))
for agent in snapshot.all_agents:
if agent.agent_type == agent_type and agent.source == source:
return agent
return None
def get_agent_file_path(
cwd: Path,
source: str,
agent_type: str,
*,
filename: str | None = None,
) -> Path:
resolved_source = normalize_mutable_source(source)
directories = dict(iter_agent_directories(cwd))
directory = directories[resolved_source]
return directory / f'{filename or agent_type}.md'
def render_agent_mutation(result: AgentMutationResult) -> str:
return '\n'.join(
[
'# Agent',
'',
f'action={result.action}',
f'agent_type={result.agent_type}',
f'source={result.source}',
f'file_path={result.file_path}',
f'overwritten={result.overwritten}',
]
)
def _split_frontmatter(text: str) -> tuple[dict[str, Any], str]:
normalized = text.replace('\r\n', '\n')
match = _FRONTMATTER_RE.match(normalized)
@@ -303,6 +587,14 @@ def _parse_frontmatter_value(value: str) -> Any:
return value
def _escape_frontmatter_text(value: str) -> str:
return (
value.replace('\\', '\\\\')
.replace('"', '\\"')
.replace('\n', '\\\\n')
)
def _parse_tool_list(value: Any) -> tuple[str, ...] | None:
if value is None or value == '':
return None
+61
View File
@@ -15,10 +15,15 @@ from .agent_context_usage import collect_context_usage, estimate_tokens, format_
from .compact import compact_conversation
from .ask_user_runtime import AskUserRuntime
from .agent_registry import (
delete_agent_definition,
find_agent_definition,
normalize_mutable_source,
load_agent_registry,
render_agent_mutation,
render_agent_detail,
render_agents_report,
scaffold_agent_definition,
update_agent_definition,
)
from .config_runtime import ConfigRuntime
from .hook_policy import HookPolicyRuntime
@@ -62,6 +67,7 @@ from .team_runtime import TeamRuntime
from .tokenizer_runtime import describe_token_counter
from .workflow_runtime import WorkflowRuntime
from .worktree_runtime import WorktreeRuntime
from .session_env_vars import clear_session_env_vars
from .session_store import (
StoredAgentSession,
load_agent_session,
@@ -244,6 +250,8 @@ class LocalCodingAgent:
self.resume_source_session_id = None
if self.plugin_runtime is not None:
self.plugin_runtime.restore_session_state({})
# Mirror commands/clear/caches.ts: drop session-scoped env vars on /clear.
clear_session_env_vars()
def build_prompt_context(self, scratchpad_directory: Path | None = None):
return build_prompt_context(
@@ -3418,6 +3426,59 @@ class LocalCodingAgent:
snapshot = self.load_agent_registry()
return render_agent_detail(snapshot, agent_type)
def render_agent_create_report(
self,
agent_type: str,
*,
source: str = 'projectSettings',
description: str | None = None,
system_prompt: str | None = None,
overwrite: bool = False,
) -> str:
result = scaffold_agent_definition(
self.runtime_config.cwd,
agent_type=agent_type,
source=normalize_mutable_source(source),
overwrite=overwrite,
description=description,
system_prompt=system_prompt,
)
return render_agent_mutation(result)
def render_agent_update_report(
self,
agent_type: str,
*,
source: str = 'auto',
description: str | None = None,
system_prompt: str | None = None,
) -> str:
update_kwargs: dict[str, object] = {}
if description is not None:
update_kwargs['description'] = description
if system_prompt is not None:
update_kwargs['system_prompt'] = system_prompt
result = update_agent_definition(
self.runtime_config.cwd,
agent_type=agent_type,
source=normalize_mutable_source(source, allow_auto=True),
**update_kwargs,
)
return render_agent_mutation(result)
def render_agent_delete_report(
self,
agent_type: str,
*,
source: str = 'auto',
) -> str:
result = delete_agent_definition(
self.runtime_config.cwd,
agent_type=agent_type,
source=normalize_mutable_source(source, allow_auto=True),
)
return render_agent_mutation(result)
def render_memory_report(self) -> str:
prompt_context = self.build_prompt_context()
claude_md = prompt_context.user_context.get('claudeMd')
File diff suppressed because it is too large Load Diff
+6
View File
@@ -15,6 +15,7 @@ from pathlib import Path
from typing import TYPE_CHECKING, Any, Callable, Iterator, Union
from .agent_types import AgentPermissions, AgentRuntimeConfig, ToolExecutionResult
from .session_env_vars import get_session_env_vars
if TYPE_CHECKING:
from .account_runtime import AccountRuntime
@@ -3321,6 +3322,11 @@ def _build_subprocess_env(context: ToolExecutionContext) -> dict[str, str]:
for key, value in os.environ.items()
if not _is_sensitive_env_var(key)
}
# Mirror utils/shell/bashProvider.ts: session env vars (set via /env)
# apply to spawned children, layered above the parent env but below
# explicit per-call extras.
for key, value in get_session_env_vars().items():
env[key] = value
env.update(context.extra_env)
return env
+147
View File
@@ -0,0 +1,147 @@
"""Display formatters — Python port of pure helpers from ``utils/format.ts``.
Only the leaf-safe formatters are ported here (no Intl dependencies, no
Ink-specific layout). These mirror the upstream output exactly so existing
golden snapshots and tests of formatted strings stay aligned.
Ported:
- ``format_file_size`` — bytes → ``"1.5KB"`` / ``"2MB"`` / ``"3.4GB"``
- ``format_seconds_short`` — ms → ``"1.2s"``
- ``format_duration`` — ms → ``"3h 4m 5s"`` with hide/most-significant flags
- ``format_number`` — compact notation (``"1.3k"``, ``"2.5m"``)
- ``format_tokens`` — like ``format_number`` but trims trailing ``.0``
Not ported (stay in TypeScript-only paths):
``formatRelativeTime`` / ``formatRelativeTimeAgo`` / ``formatLogMetadata``
/ ``formatResetTime`` / ``formatResetText`` — they depend on
``intl.ts`` and the Ink reset-time UX.
"""
from __future__ import annotations
def _trim_trailing_zero(value: str) -> str:
return value[:-2] if value.endswith('.0') else value
def format_file_size(size_in_bytes: float) -> str:
"""Bytes to a human-readable string, mirroring the JS thresholds."""
kb = size_in_bytes / 1024
if kb < 1:
return f'{int(size_in_bytes)} bytes'
if kb < 1024:
return f'{_trim_trailing_zero(f"{kb:.1f}")}KB'
mb = kb / 1024
if mb < 1024:
return f'{_trim_trailing_zero(f"{mb:.1f}")}MB'
gb = mb / 1024
return f'{_trim_trailing_zero(f"{gb:.1f}")}GB'
def format_seconds_short(ms: float) -> str:
"""Milliseconds → ``"1.2s"`` (always one decimal)."""
return f'{ms / 1000:.1f}s'
def format_duration(
ms: float,
*,
hide_trailing_zeros: bool = False,
most_significant_only: bool = False,
) -> str:
"""Format a millisecond duration with d/h/m/s components.
Mirrors ``utils/format.ts#formatDuration`` including the rounding
carry-over (``59.5s`` rounds up to the next minute).
"""
if ms < 60_000:
if ms == 0:
return '0s'
if ms < 1:
return f'{ms / 1000:.1f}s'
return f'{int(ms // 1000)}s'
days = int(ms // 86_400_000)
hours = int((ms % 86_400_000) // 3_600_000)
minutes = int((ms % 3_600_000) // 60_000)
seconds = int(round((ms % 60_000) / 1000))
if seconds == 60:
seconds = 0
minutes += 1
if minutes == 60:
minutes = 0
hours += 1
if hours == 24:
hours = 0
days += 1
if most_significant_only:
if days > 0:
return f'{days}d'
if hours > 0:
return f'{hours}h'
if minutes > 0:
return f'{minutes}m'
return f'{seconds}s'
hide = hide_trailing_zeros
if days > 0:
if hide and hours == 0 and minutes == 0:
return f'{days}d'
if hide and minutes == 0:
return f'{days}d {hours}h'
return f'{days}d {hours}h {minutes}m'
if hours > 0:
if hide and minutes == 0 and seconds == 0:
return f'{hours}h'
if hide and seconds == 0:
return f'{hours}h {minutes}m'
return f'{hours}h {minutes}m {seconds}s'
if minutes > 0:
if hide and seconds == 0:
return f'{minutes}m'
return f'{minutes}m {seconds}s'
return f'{seconds}s'
_COMPACT_SUFFIXES = (
(1_000_000_000_000, 't'),
(1_000_000_000, 'b'),
(1_000_000, 'm'),
(1_000, 'k'),
)
def format_number(number: float) -> str:
"""Compact notation matching ``Intl.NumberFormat`` with one fraction digit.
The npm version emits e.g. ``"1.3k"`` from 1321 and ``"900"`` from 900.
For values < 1000 the integer is returned with no separator. For larger
values one fraction digit is shown when ``number >= 1000`` to mirror the
``minimumFractionDigits: 1`` consistent-decimal branch.
"""
if number < 1000:
return str(int(number))
for threshold, suffix in _COMPACT_SUFFIXES:
if number >= threshold:
scaled = number / threshold
return f'{scaled:.1f}{suffix}'
return str(int(number))
def format_tokens(count: float) -> str:
"""Like ``format_number`` but trims a trailing ``.0`` (e.g. ``"1k"``)."""
return format_number(count).replace('.0', '')
__all__ = [
'format_file_size',
'format_seconds_short',
'format_duration',
'format_number',
'format_tokens',
]
+185
View File
@@ -0,0 +1,185 @@
"""Git utility ports — subset of ``utils/git.ts`` plus ``utils/gitSettings.ts``.
This module covers the pure / filesystem-only pieces:
- ``find_git_root`` — walks up from ``start_path`` looking for ``.git``
- ``normalize_git_remote_url`` — canonicalizes SSH/HTTPS remote URLs
- ``get_repo_remote_hash`` — sha256[:16] of the normalized remote URL
- ``should_include_git_instructions`` — env-var override + settings opt-out
The shell-driven git operations (``getHead``, ``getBranch``, ``getDefaultBranch``,
``getChangedFiles``, etc.) are intentionally left for a later slice — they
need the full settings/cache plumbing the npm version uses.
"""
from __future__ import annotations
import hashlib
import os
import re
from collections import OrderedDict
from pathlib import Path
from typing import Callable
# ---------------------------------------------------------------------------
# Tiny LRU helper that mirrors lodash memoizeWithLRU semantics
# ---------------------------------------------------------------------------
def _lru_memoize(
fn: Callable[[str], str | None], max_size: int,
) -> Callable[[str], str | None]:
cache: OrderedDict[str, str | None] = OrderedDict()
def wrapper(key: str) -> str | None:
if key in cache:
cache.move_to_end(key)
return cache[key]
value = fn(key)
cache[key] = value
cache.move_to_end(key)
if len(cache) > max_size:
cache.popitem(last=False)
return value
wrapper.cache_clear = cache.clear # type: ignore[attr-defined]
return wrapper
# ---------------------------------------------------------------------------
# find_git_root
# ---------------------------------------------------------------------------
def _find_git_root_uncached(start_path: str) -> str | None:
current = Path(start_path).resolve()
while True:
git_path = current / '.git'
try:
stat = git_path.stat()
if stat.st_mode and (
git_path.is_dir() or git_path.is_file()
):
return str(current)
except OSError:
pass
parent = current.parent
if parent == current:
return None
current = parent
find_git_root = _lru_memoize(_find_git_root_uncached, max_size=50)
"""Walk up from ``start_path`` to find the first directory containing ``.git``.
Returns the absolute path of that directory, or ``None`` if not in a repo.
Memoized per ``start_path`` with an LRU cache (max 50 entries).
"""
# ---------------------------------------------------------------------------
# normalize_git_remote_url
# ---------------------------------------------------------------------------
_SSH_RE = re.compile(r'^git@([^:]+):(.+?)(?:\.git)?$')
_URL_RE = re.compile(
r'^(?:https?|ssh)://(?:[^@]+@)?([^/]+)/(.+?)(?:\.git)?$',
)
_LOCAL_HOST_IPV4 = re.compile(r'^127\.\d{1,3}\.\d{1,3}\.\d{1,3}$')
def _is_local_host(host: str) -> bool:
host_no_port = host.split(':', 1)[0]
return host_no_port == 'localhost' or bool(_LOCAL_HOST_IPV4.match(host_no_port))
def normalize_git_remote_url(url: str) -> str | None:
"""Canonicalize a git remote URL to ``host/owner/repo`` lowercased.
Returns ``None`` if the URL doesn't match a recognized SSH/HTTPS shape.
"""
trimmed = url.strip()
if not trimmed:
return None
ssh = _SSH_RE.match(trimmed)
if ssh:
return f'{ssh.group(1)}/{ssh.group(2)}'.lower()
url_match = _URL_RE.match(trimmed)
if url_match:
host = url_match.group(1)
path = url_match.group(2)
# CCR git proxy: http://...@127.0.0.1:PORT/git/[host/]owner/repo
if _is_local_host(host) and path.startswith('git/'):
proxy_path = path[len('git/'):]
segments = proxy_path.split('/')
if len(segments) >= 3 and '.' in segments[0]:
return proxy_path.lower()
return f'github.com/{proxy_path}'.lower()
return f'{host}/{path}'.lower()
return None
def get_repo_remote_hash(remote_url: str | None) -> str | None:
"""Return sha256[:16] of the normalized remote URL, or None.
Unlike the npm version this takes the URL as a parameter rather than
invoking ``git remote get-url`` itself, so it can be called from
contexts where git binary access is unavailable.
"""
if not remote_url:
return None
normalized = normalize_git_remote_url(remote_url)
if not normalized:
return None
return hashlib.sha256(normalized.encode('utf-8')).hexdigest()[:16]
# ---------------------------------------------------------------------------
# gitSettings.ts — env-var + settings opt-out for git instructions
# ---------------------------------------------------------------------------
_TRUTHY_ENV = frozenset({'1', 'true', 'yes', 'on'})
_FALSY_ENV = frozenset({'0', 'false', 'no', 'off'})
def _env_truthy(value: str | None) -> bool:
if value is None:
return False
return value.strip().lower() in _TRUTHY_ENV
def _env_defined_falsy(value: str | None) -> bool:
if value is None:
return False
return value.strip().lower() in _FALSY_ENV
def should_include_git_instructions(
*,
settings_value: bool | None = None,
env: dict[str, str] | None = None,
) -> bool:
"""Whether to surface git-aware prompt sections.
Mirrors ``utils/gitSettings.ts``: env var
``CLAUDE_CODE_DISABLE_GIT_INSTRUCTIONS`` overrides
``settings.includeGitInstructions``; default is ``True``.
"""
chosen_env = env if env is not None else os.environ
raw = chosen_env.get('CLAUDE_CODE_DISABLE_GIT_INSTRUCTIONS')
if _env_truthy(raw):
return False
if _env_defined_falsy(raw):
return True
return True if settings_value is None else settings_value
__all__ = [
'find_git_root',
'normalize_git_remote_url',
'get_repo_remote_hash',
'should_include_git_instructions',
]
+91
View File
@@ -0,0 +1,91 @@
"""IDE path conversion — Python port of ``utils/idePathConversion.ts``.
Used when Claude runs under WSL but the IDE (VS Code, JetBrains) is on the
host Windows side. Outgoing paths need to be converted to ``\\\\wsl$\\...``
form for the IDE; incoming paths from the IDE need to be converted back to
``/mnt/c/...`` form for Claude.
"""
from __future__ import annotations
import re
import subprocess
from typing import Protocol
_WSL_UNC_RE = re.compile(r'^\\\\wsl(?:\.localhost|\$)\\([^\\]+)(.*)$')
_DRIVE_RE = re.compile(r'^([A-Za-z]):')
class IDEPathConverter(Protocol):
"""Bidirectional path mapping between IDE-side and Claude-side paths."""
def to_local_path(self, ide_path: str) -> str: ...
def to_ide_path(self, local_path: str) -> str: ...
def _run_wslpath(flag: str, path: str) -> str:
"""Invoke ``wslpath`` and return the stripped stdout. Raises on failure."""
completed = subprocess.run(
['wslpath', flag, path],
capture_output=True,
text=True,
check=True,
)
return completed.stdout.strip()
def _manual_windows_to_wsl(windows_path: str) -> str:
"""Fallback when ``wslpath`` is unavailable: ``C:\\foo`` → ``/mnt/c/foo``."""
converted = windows_path.replace('\\', '/')
def _replace_drive(match: re.Match[str]) -> str:
return f'/mnt/{match.group(1).lower()}'
return _DRIVE_RE.sub(_replace_drive, converted)
class WindowsToWSLConverter:
"""Converter for the Windows IDE + WSL Claude scenario."""
def __init__(self, wsl_distro_name: str | None) -> None:
self.wsl_distro_name = wsl_distro_name
def to_local_path(self, windows_path: str) -> str:
if not windows_path:
return windows_path
if self.wsl_distro_name:
unc = _WSL_UNC_RE.match(windows_path)
if unc and unc.group(1) != self.wsl_distro_name:
# Path belongs to a different distro — wslpath would fail.
return windows_path
try:
return _run_wslpath('-u', windows_path)
except (subprocess.CalledProcessError, FileNotFoundError, OSError):
return _manual_windows_to_wsl(windows_path)
def to_ide_path(self, wsl_path: str) -> str:
if not wsl_path:
return wsl_path
try:
return _run_wslpath('-w', wsl_path)
except (subprocess.CalledProcessError, FileNotFoundError, OSError):
return wsl_path
def check_wsl_distro_match(windows_path: str, wsl_distro_name: str) -> bool:
"""True if ``windows_path`` isn't a WSL UNC path or names this distro."""
unc = _WSL_UNC_RE.match(windows_path)
if unc:
return unc.group(1) == wsl_distro_name
return True
__all__ = [
'IDEPathConverter',
'WindowsToWSLConverter',
'check_wsl_distro_match',
]
+131
View File
@@ -11,6 +11,13 @@ from typing import Callable
from .background_runtime import BackgroundSessionRuntime, build_background_worker_command
from .account_runtime import AccountRuntime
from .ask_user_runtime import AskUserRuntime
from .agent_registry import (
create_agent_definition,
delete_agent_definition,
normalize_mutable_source,
render_agent_mutation,
update_agent_definition,
)
from .agent_runtime import LocalCodingAgent
from .agent_types import (
AgentPermissions,
@@ -163,6 +170,47 @@ def _load_output_schema_config(args: argparse.Namespace) -> OutputSchemaConfig |
)
def _add_agent_management_args(parser: argparse.ArgumentParser) -> None:
parser.add_argument('--cwd', default='.')
parser.add_argument(
'--source',
default='project',
choices=('project', 'user', 'auto', 'projectSettings', 'userSettings'),
)
parser.add_argument('--description')
parser.add_argument('--prompt')
parser.add_argument('--prompt-file')
parser.add_argument('--tools')
parser.add_argument('--model')
parser.add_argument('--color')
parser.add_argument('--permission-mode')
parser.add_argument('--max-turns', type=int)
parser.add_argument('--initial-prompt')
parser.add_argument('--background', action='store_true')
parser.add_argument('--one-shot', action='store_true')
parser.add_argument('--omit-claude-md', action='store_true')
parser.add_argument('--overwrite', action='store_true')
def _resolve_agent_prompt_text(args: argparse.Namespace) -> str | None:
prompt_text = getattr(args, 'prompt', None)
prompt_file = getattr(args, 'prompt_file', None)
if prompt_text and prompt_file:
raise ValueError('Specify only one of --prompt or --prompt-file')
if prompt_file:
return Path(prompt_file).read_text(encoding='utf-8')
return prompt_text
def _parse_tools_flag(raw_tools: str | None) -> tuple[str, ...] | None:
if raw_tools is None:
return None
cleaned = raw_tools.strip()
if not cleaned or cleaned == '*':
return None
return tuple(item.strip() for item in cleaned.split(',') if item.strip())
def _build_agent(args: argparse.Namespace) -> LocalCodingAgent:
return LocalCodingAgent(
model_config=_build_model_config(args),
@@ -879,6 +927,23 @@ def build_parser() -> argparse.ArgumentParser:
agents_parser.add_argument('agent_type', nargs='?')
agents_parser.add_argument('--all', action='store_true')
_add_agent_common_args(agents_parser, include_backend=False)
agents_create_parser = subparsers.add_parser('agents-create', help='create a local agent definition markdown file')
agents_create_parser.add_argument('agent_type')
_add_agent_management_args(agents_create_parser)
agents_update_parser = subparsers.add_parser('agents-update', help='update an existing local agent definition markdown file')
agents_update_parser.add_argument('agent_type')
_add_agent_management_args(agents_update_parser)
agents_delete_parser = subparsers.add_parser('agents-delete', help='delete an existing local agent definition markdown file')
agents_delete_parser.add_argument('agent_type')
agents_delete_parser.add_argument('--cwd', default='.')
agents_delete_parser.add_argument(
'--source',
default='auto',
choices=('project', 'user', 'auto', 'projectSettings', 'userSettings'),
)
return parser
@@ -1519,6 +1584,72 @@ def main(argv: list[str] | None = None) -> int:
else:
print(agent.render_agents_report(show_all=bool(args.all)))
return 0
if args.command == 'agents-create':
prompt_text = _resolve_agent_prompt_text(args)
result = create_agent_definition(
Path(args.cwd).resolve(),
agent_type=args.agent_type,
description=args.description or f'Use this agent when the task calls for {args.agent_type}.',
system_prompt=prompt_text
or (
f'You are the {args.agent_type} agent.\n'
'Read the task carefully, use the available tools deliberately, and return a concise result.'
),
source=normalize_mutable_source(args.source),
overwrite=bool(args.overwrite),
tools=_parse_tools_flag(args.tools),
model=args.model,
color=args.color,
permission_mode=args.permission_mode,
max_turns=args.max_turns,
initial_prompt=args.initial_prompt,
background=bool(args.background),
one_shot=bool(args.one_shot),
omit_claude_md=bool(args.omit_claude_md),
)
print(render_agent_mutation(result))
return 0
if args.command == 'agents-update':
prompt_text = _resolve_agent_prompt_text(args)
update_kwargs: dict[str, object] = {}
if args.description is not None:
update_kwargs['description'] = args.description
if prompt_text is not None:
update_kwargs['system_prompt'] = prompt_text
if args.tools is not None:
update_kwargs['tools'] = _parse_tools_flag(args.tools)
if args.model is not None:
update_kwargs['model'] = args.model
if args.color is not None:
update_kwargs['color'] = args.color
if args.permission_mode is not None:
update_kwargs['permission_mode'] = args.permission_mode
if args.max_turns is not None:
update_kwargs['max_turns'] = args.max_turns
if args.initial_prompt is not None:
update_kwargs['initial_prompt'] = args.initial_prompt
if args.background:
update_kwargs['background'] = True
if args.one_shot:
update_kwargs['one_shot'] = True
if args.omit_claude_md:
update_kwargs['omit_claude_md'] = True
result = update_agent_definition(
Path(args.cwd).resolve(),
agent_type=args.agent_type,
source=normalize_mutable_source(args.source, allow_auto=True),
**update_kwargs,
)
print(render_agent_mutation(result))
return 0
if args.command == 'agents-delete':
result = delete_agent_definition(
Path(args.cwd).resolve(),
agent_type=args.agent_type,
source=normalize_mutable_source(args.source, allow_auto=True),
)
print(render_agent_mutation(result))
return 0
parser.error(f'unknown command: {args.command}')
return 2
+226
View File
@@ -0,0 +1,226 @@
"""Model pricing — Python port of utils/modelCost.ts.
Pricing values mirror the upstream tiers exactly. The npm version logs an
analytics event on unknown models; here we just fall back to the
DEFAULT_UNKNOWN_MODEL_COST tier and let callers decide what to track.
"""
from __future__ import annotations
from dataclasses import dataclass
@dataclass(frozen=True)
class ModelCosts:
"""USD per million input/output tokens (cache + web-search per request)."""
input_tokens: float
output_tokens: float
prompt_cache_write_tokens: float
prompt_cache_read_tokens: float
web_search_requests: float
# Standard pricing tier for Sonnet models: $3 input / $15 output per Mtok
COST_TIER_3_15 = ModelCosts(
input_tokens=3.0,
output_tokens=15.0,
prompt_cache_write_tokens=3.75,
prompt_cache_read_tokens=0.3,
web_search_requests=0.01,
)
# Pricing tier for Opus 4 / 4.1: $15 input / $75 output per Mtok
COST_TIER_15_75 = ModelCosts(
input_tokens=15.0,
output_tokens=75.0,
prompt_cache_write_tokens=18.75,
prompt_cache_read_tokens=1.5,
web_search_requests=0.01,
)
# Pricing tier for Opus 4.5 (also default Opus 4.6): $5 input / $25 output per Mtok
COST_TIER_5_25 = ModelCosts(
input_tokens=5.0,
output_tokens=25.0,
prompt_cache_write_tokens=6.25,
prompt_cache_read_tokens=0.5,
web_search_requests=0.01,
)
# Fast-mode pricing for Opus 4.6: $30 input / $150 output per Mtok
COST_TIER_30_150 = ModelCosts(
input_tokens=30.0,
output_tokens=150.0,
prompt_cache_write_tokens=37.5,
prompt_cache_read_tokens=3.0,
web_search_requests=0.01,
)
# Pricing for Haiku 3.5: $0.80 input / $4 output per Mtok
COST_HAIKU_35 = ModelCosts(
input_tokens=0.8,
output_tokens=4.0,
prompt_cache_write_tokens=1.0,
prompt_cache_read_tokens=0.08,
web_search_requests=0.01,
)
# Pricing for Haiku 4.5: $1 input / $5 output per Mtok
COST_HAIKU_45 = ModelCosts(
input_tokens=1.0,
output_tokens=5.0,
prompt_cache_write_tokens=1.25,
prompt_cache_read_tokens=0.1,
web_search_requests=0.01,
)
DEFAULT_UNKNOWN_MODEL_COST = COST_TIER_5_25
# Canonical short-name → cost tier. Lookup uses substring matching so that
# version-suffixed model IDs (`claude-opus-4-6-20251015`) resolve correctly.
MODEL_COSTS: dict[str, ModelCosts] = {
'claude-3-5-haiku': COST_HAIKU_35,
'claude-haiku-4-5': COST_HAIKU_45,
'claude-3-5-sonnet': COST_TIER_3_15,
'claude-3-7-sonnet': COST_TIER_3_15,
'claude-sonnet-4': COST_TIER_3_15,
'claude-sonnet-4-5': COST_TIER_3_15,
'claude-sonnet-4-6': COST_TIER_3_15,
'claude-opus-4': COST_TIER_15_75,
'claude-opus-4-1': COST_TIER_15_75,
'claude-opus-4-5': COST_TIER_5_25,
'claude-opus-4-6': COST_TIER_5_25,
}
def _resolve_model_costs(model: str) -> ModelCosts | None:
"""Return MODEL_COSTS entry for `model`, matching by longest prefix."""
canonical = model.lower()
matches = [
(key, costs)
for key, costs in MODEL_COSTS.items()
if canonical.startswith(key) or key in canonical
]
if not matches:
return None
# Prefer the most specific (longest) key match so `claude-opus-4-6` wins
# over `claude-opus-4`.
matches.sort(key=lambda item: len(item[0]), reverse=True)
return matches[0][1]
def get_opus_4_6_cost_tier(fast_mode: bool) -> ModelCosts:
"""Return the right tier for Opus 4.6 — fast mode is more expensive."""
return COST_TIER_30_150 if fast_mode else COST_TIER_5_25
def get_model_costs(model: str, *, fast_mode: bool = False) -> ModelCosts:
"""Return ModelCosts for `model`, applying the Opus 4.6 fast-mode tier."""
canonical = model.lower()
if 'claude-opus-4-6' in canonical:
return get_opus_4_6_cost_tier(fast_mode)
costs = _resolve_model_costs(model)
return costs if costs is not None else DEFAULT_UNKNOWN_MODEL_COST
def tokens_to_usd_cost(
costs: ModelCosts,
*,
input_tokens: int,
output_tokens: int,
cache_read_input_tokens: int = 0,
cache_creation_input_tokens: int = 0,
web_search_requests: int = 0,
) -> float:
"""Compute USD cost from token counts and a ModelCosts tier."""
return (
(input_tokens / 1_000_000) * costs.input_tokens
+ (output_tokens / 1_000_000) * costs.output_tokens
+ (cache_read_input_tokens / 1_000_000) * costs.prompt_cache_read_tokens
+ (cache_creation_input_tokens / 1_000_000) * costs.prompt_cache_write_tokens
+ web_search_requests * costs.web_search_requests
)
def calculate_usd_cost(
model: str,
*,
input_tokens: int,
output_tokens: int,
cache_read_input_tokens: int = 0,
cache_creation_input_tokens: int = 0,
web_search_requests: int = 0,
fast_mode: bool = False,
) -> float:
"""USD cost for a query — looks up the tier and applies tokens_to_usd_cost."""
costs = get_model_costs(model, fast_mode=fast_mode)
return tokens_to_usd_cost(
costs,
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_read_input_tokens=cache_read_input_tokens,
cache_creation_input_tokens=cache_creation_input_tokens,
web_search_requests=web_search_requests,
)
def calculate_cost_from_tokens(
model: str,
tokens: dict,
*,
fast_mode: bool = False,
) -> float:
"""Mirror calculateCostFromTokens — accepts the same camelCase token dict."""
return calculate_usd_cost(
model,
input_tokens=int(tokens.get('inputTokens', 0)),
output_tokens=int(tokens.get('outputTokens', 0)),
cache_read_input_tokens=int(tokens.get('cacheReadInputTokens', 0)),
cache_creation_input_tokens=int(tokens.get('cacheCreationInputTokens', 0)),
web_search_requests=int(tokens.get('webSearchRequests', 0)),
fast_mode=fast_mode,
)
def _format_price(price: float) -> str:
if float(price).is_integer():
return f'${int(price)}'
return f'${price:.2f}'
def format_model_pricing(costs: ModelCosts) -> str:
"""Return a human-readable pricing label like '$3/$15 per Mtok'."""
return (
f'{_format_price(costs.input_tokens)}/'
f'{_format_price(costs.output_tokens)} per Mtok'
)
def get_model_pricing_string(model: str) -> str | None:
"""Return formatted pricing string for `model`, or None if unknown."""
costs = _resolve_model_costs(model)
if costs is None:
return None
return format_model_pricing(costs)
__all__ = [
'ModelCosts',
'COST_TIER_3_15',
'COST_TIER_15_75',
'COST_TIER_5_25',
'COST_TIER_30_150',
'COST_HAIKU_35',
'COST_HAIKU_45',
'DEFAULT_UNKNOWN_MODEL_COST',
'MODEL_COSTS',
'get_opus_4_6_cost_tier',
'get_model_costs',
'tokens_to_usd_cost',
'calculate_usd_cost',
'calculate_cost_from_tokens',
'format_model_pricing',
'get_model_pricing_string',
]
+232
View File
@@ -0,0 +1,232 @@
"""Platform detection and system directories — Python ports of
``utils/platform.ts`` and ``utils/systemDirectories.ts``.
The npm functions are memoized via lodash; here a module-level cache plus
``_reset_cache`` (test-only) provides equivalent behavior. Detection is
cheap enough that callers can also bypass the cache by passing explicit
overrides to ``get_system_directories``.
"""
from __future__ import annotations
import os
import platform as _stdlib_platform
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import Literal
Platform = Literal['macos', 'windows', 'wsl', 'linux', 'unknown']
SUPPORTED_PLATFORMS: tuple[Platform, ...] = ('macos', 'wsl')
_UNSET = object()
_platform_cache: Platform | None = None
_wsl_version_cache: object = _UNSET # sentinel until first computation
def _reset_cache() -> None:
"""Clear cached platform detection — only used by tests."""
global _platform_cache, _wsl_version_cache
_platform_cache = None
_wsl_version_cache = _UNSET
def _read_proc_version() -> str:
return Path('/proc/version').read_text(encoding='utf-8')
def get_platform() -> Platform:
"""Return the current platform identifier (memoized)."""
global _platform_cache
if _platform_cache is not None:
return _platform_cache
if sys.platform == 'darwin':
_platform_cache = 'macos'
elif sys.platform.startswith('win'):
_platform_cache = 'windows'
elif sys.platform.startswith('linux'):
try:
proc_version = _read_proc_version().lower()
if 'microsoft' in proc_version or 'wsl' in proc_version:
_platform_cache = 'wsl'
else:
_platform_cache = 'linux'
except OSError:
_platform_cache = 'linux'
else:
_platform_cache = 'unknown'
return _platform_cache
def get_wsl_version() -> str | None:
"""Return the WSL major version (`'1'`/`'2'`/...), or None if not WSL."""
global _wsl_version_cache
if _wsl_version_cache is not _UNSET:
return _wsl_version_cache # type: ignore[return-value]
result: str | None = None
if sys.platform.startswith('linux'):
try:
proc_version = _read_proc_version()
except OSError:
proc_version = ''
if proc_version:
import re
match = re.search(r'WSL(\d+)', proc_version, re.IGNORECASE)
if match:
result = match.group(1)
elif 'microsoft' in proc_version.lower():
result = '1'
_wsl_version_cache = result
return result
@dataclass(frozen=True)
class LinuxDistroInfo:
linux_distro_id: str | None = None
linux_distro_version: str | None = None
linux_kernel: str | None = None
def to_dict(self) -> dict[str, str]:
out: dict[str, str] = {}
if self.linux_distro_id is not None:
out['linuxDistroId'] = self.linux_distro_id
if self.linux_distro_version is not None:
out['linuxDistroVersion'] = self.linux_distro_version
if self.linux_kernel is not None:
out['linuxKernel'] = self.linux_kernel
return out
def get_linux_distro_info() -> LinuxDistroInfo | None:
"""Return distro id/version/kernel on Linux, or None on other platforms."""
if not sys.platform.startswith('linux'):
return None
distro_id: str | None = None
distro_version: str | None = None
try:
content = Path('/etc/os-release').read_text(encoding='utf-8')
except OSError:
content = ''
for line in content.splitlines():
if '=' not in line:
continue
key, _, value = line.partition('=')
value = value.strip().strip('"')
if key == 'ID':
distro_id = value
elif key == 'VERSION_ID':
distro_version = value
return LinuxDistroInfo(
linux_distro_id=distro_id,
linux_distro_version=distro_version,
linux_kernel=_stdlib_platform.release() or None,
)
_VCS_MARKERS: tuple[tuple[str, str], ...] = (
('.git', 'git'),
('.hg', 'mercurial'),
('.svn', 'svn'),
('.p4config', 'perforce'),
('$tf', 'tfs'),
('.tfvc', 'tfs'),
('.jj', 'jujutsu'),
('.sl', 'sapling'),
)
def detect_vcs(directory: str | os.PathLike[str] | None = None) -> list[str]:
"""Detect VCS systems by marker files in ``directory`` (defaults to cwd)."""
detected: set[str] = set()
if os.environ.get('P4PORT'):
detected.add('perforce')
target = Path(directory) if directory is not None else Path.cwd()
try:
entries = {entry.name for entry in target.iterdir()}
except OSError:
entries = set()
for marker, vcs in _VCS_MARKERS:
if marker in entries:
detected.add(vcs)
return sorted(detected)
# ---------------------------------------------------------------------------
# systemDirectories.ts
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class SystemDirectories:
HOME: str
DESKTOP: str
DOCUMENTS: str
DOWNLOADS: str
def to_dict(self) -> dict[str, str]:
return {
'HOME': self.HOME,
'DESKTOP': self.DESKTOP,
'DOCUMENTS': self.DOCUMENTS,
'DOWNLOADS': self.DOWNLOADS,
}
def get_system_directories(
*,
env: dict[str, str] | None = None,
home_dir: str | None = None,
platform: Platform | None = None,
) -> SystemDirectories:
"""Cross-platform system directories matching ``getSystemDirectories``."""
chosen_platform: Platform = platform if platform is not None else get_platform()
chosen_home = home_dir if home_dir is not None else str(Path.home())
chosen_env = env if env is not None else dict(os.environ)
defaults = SystemDirectories(
HOME=chosen_home,
DESKTOP=str(Path(chosen_home) / 'Desktop'),
DOCUMENTS=str(Path(chosen_home) / 'Documents'),
DOWNLOADS=str(Path(chosen_home) / 'Downloads'),
)
if chosen_platform == 'windows':
user_profile = chosen_env.get('USERPROFILE') or chosen_home
return SystemDirectories(
HOME=chosen_home,
DESKTOP=str(Path(user_profile) / 'Desktop'),
DOCUMENTS=str(Path(user_profile) / 'Documents'),
DOWNLOADS=str(Path(user_profile) / 'Downloads'),
)
if chosen_platform in ('linux', 'wsl'):
return SystemDirectories(
HOME=chosen_home,
DESKTOP=chosen_env.get('XDG_DESKTOP_DIR') or defaults.DESKTOP,
DOCUMENTS=chosen_env.get('XDG_DOCUMENTS_DIR') or defaults.DOCUMENTS,
DOWNLOADS=chosen_env.get('XDG_DOWNLOAD_DIR') or defaults.DOWNLOADS,
)
# macOS and unknown both use the defaults.
return defaults
__all__ = [
'Platform',
'SUPPORTED_PLATFORMS',
'get_platform',
'get_wsl_version',
'LinuxDistroInfo',
'get_linux_distro_info',
'detect_vcs',
'SystemDirectories',
'get_system_directories',
]
+134
View File
@@ -0,0 +1,134 @@
"""Local release-notes parsing — Python port of utils/releaseNotes.ts.
The Python runtime has no network/cache layer, so this module only reads a
local CHANGELOG.md (typically in the project root). The npm version fetches
from GitHub and caches under ~/.claude/cache/changelog.md; here, callers
provide the changelog text or pass the project cwd so we can read it.
"""
from __future__ import annotations
import re
from pathlib import Path
from typing import Iterable
MAX_RELEASE_NOTES_SHOWN = 5
def parse_changelog(content: str) -> dict[str, list[str]]:
"""Parse a markdown CHANGELOG into {version: [bullet, ...]}.
Recognises sections starting with `## <version>` (optionally followed by
` - YYYY-MM-DD`). Bullet lines starting with `- ` become entries.
"""
if not content:
return {}
notes: dict[str, list[str]] = {}
sections = re.split(r'^## ', content, flags=re.MULTILINE)[1:]
for section in sections:
lines = section.strip().splitlines()
if not lines:
continue
version = lines[0].split(' - ')[0].strip()
if not version:
continue
bullets: list[str] = []
for line in lines[1:]:
stripped = line.strip()
if stripped.startswith('- '):
text = stripped[2:].strip()
if text:
bullets.append(text)
if bullets:
notes[version] = bullets
return notes
def _coerce_version(value: str | None) -> tuple[int, ...] | None:
if value is None:
return None
match = re.match(r'^(\d+)(?:\.(\d+))?(?:\.(\d+))?', value.strip())
if match is None:
return None
return tuple(int(part) if part else 0 for part in match.groups())
def _gt(a: str, b: str) -> bool:
parsed_a = _coerce_version(a)
parsed_b = _coerce_version(b)
if parsed_a is None or parsed_b is None:
return False
return parsed_a > parsed_b
def get_recent_release_notes(
current_version: str,
previous_version: str | None,
changelog_content: str,
) -> list[str]:
"""Return up to MAX_RELEASE_NOTES_SHOWN bullets newer than previous_version."""
notes = parse_changelog(changelog_content)
base_current = _coerce_version(current_version)
base_previous = _coerce_version(previous_version)
if base_previous is not None and base_current is not None and base_current <= base_previous:
return []
relevant = [
(version, bullets)
for version, bullets in notes.items()
if base_previous is None or _gt(version, previous_version or '0')
]
relevant.sort(key=lambda item: _coerce_version(item[0]) or (), reverse=True)
flat: list[str] = []
for _, bullets in relevant:
flat.extend(bullets)
return flat[:MAX_RELEASE_NOTES_SHOWN]
def get_all_release_notes(changelog_content: str) -> list[tuple[str, list[str]]]:
"""Return all [(version, bullets)] entries sorted oldest-first."""
notes = parse_changelog(changelog_content)
versions = sorted(notes.keys(), key=lambda v: _coerce_version(v) or ())
return [(version, notes[version]) for version in versions if notes[version]]
def read_local_changelog(cwd: Path) -> str:
"""Read CHANGELOG.md from the workspace root, returning '' if missing."""
path = cwd / 'CHANGELOG.md'
try:
return path.read_text(encoding='utf-8')
except (OSError, UnicodeDecodeError):
return ''
def check_for_release_notes(
current_version: str,
last_seen_version: str | None,
cwd: Path | None = None,
changelog_content: str | None = None,
) -> dict:
"""Return {'hasReleaseNotes': bool, 'releaseNotes': [...]}.
If `changelog_content` is supplied it is used directly; otherwise the
workspace CHANGELOG.md is read. Mirrors the npm `checkForReleaseNotes`
return shape (without the network fetch — no cache update).
"""
content = (
changelog_content
if changelog_content is not None
else read_local_changelog(cwd or Path.cwd())
)
bullets = get_recent_release_notes(current_version, last_seen_version, content)
return {
'hasReleaseNotes': bool(bullets),
'releaseNotes': bullets,
}
__all__ = [
'MAX_RELEASE_NOTES_SHOWN',
'parse_changelog',
'get_recent_release_notes',
'get_all_release_notes',
'read_local_changelog',
'check_for_release_notes',
]
+230
View File
@@ -0,0 +1,230 @@
"""Sandbox configuration types — Python port of entrypoints/sandboxTypes.ts.
The npm version uses Zod schemas with `.passthrough()` to forward unknown
fields. Here we mirror the same fields as dataclasses, with `from_dict`
classmethods that accept and round-trip arbitrary extra keys.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any
def _str_list(value: Any) -> list[str] | None:
if value is None:
return None
if not isinstance(value, list):
raise ValueError(f'expected list[str], got {type(value).__name__}')
return [str(item) for item in value]
def _opt_bool(value: Any) -> bool | None:
if value is None:
return None
if not isinstance(value, bool):
raise ValueError(f'expected bool, got {type(value).__name__}')
return value
def _opt_int(value: Any) -> int | None:
if value is None:
return None
if isinstance(value, bool):
raise ValueError('expected number, got bool')
if not isinstance(value, (int, float)):
raise ValueError(f'expected number, got {type(value).__name__}')
return int(value)
def _drop_none(d: dict[str, Any]) -> dict[str, Any]:
return {k: v for k, v in d.items() if v is not None}
@dataclass(frozen=True)
class SandboxNetworkConfig:
allowed_domains: list[str] | None = None
allow_managed_domains_only: bool | None = None
allow_unix_sockets: list[str] | None = None
allow_all_unix_sockets: bool | None = None
allow_local_binding: bool | None = None
http_proxy_port: int | None = None
socks_proxy_port: int | None = None
@classmethod
def from_dict(cls, data: dict[str, Any]) -> 'SandboxNetworkConfig':
return cls(
allowed_domains=_str_list(data.get('allowedDomains')),
allow_managed_domains_only=_opt_bool(data.get('allowManagedDomainsOnly')),
allow_unix_sockets=_str_list(data.get('allowUnixSockets')),
allow_all_unix_sockets=_opt_bool(data.get('allowAllUnixSockets')),
allow_local_binding=_opt_bool(data.get('allowLocalBinding')),
http_proxy_port=_opt_int(data.get('httpProxyPort')),
socks_proxy_port=_opt_int(data.get('socksProxyPort')),
)
def to_dict(self) -> dict[str, Any]:
return _drop_none({
'allowedDomains': self.allowed_domains,
'allowManagedDomainsOnly': self.allow_managed_domains_only,
'allowUnixSockets': self.allow_unix_sockets,
'allowAllUnixSockets': self.allow_all_unix_sockets,
'allowLocalBinding': self.allow_local_binding,
'httpProxyPort': self.http_proxy_port,
'socksProxyPort': self.socks_proxy_port,
})
@dataclass(frozen=True)
class SandboxFilesystemConfig:
allow_write: list[str] | None = None
deny_write: list[str] | None = None
deny_read: list[str] | None = None
allow_read: list[str] | None = None
allow_managed_read_paths_only: bool | None = None
@classmethod
def from_dict(cls, data: dict[str, Any]) -> 'SandboxFilesystemConfig':
return cls(
allow_write=_str_list(data.get('allowWrite')),
deny_write=_str_list(data.get('denyWrite')),
deny_read=_str_list(data.get('denyRead')),
allow_read=_str_list(data.get('allowRead')),
allow_managed_read_paths_only=_opt_bool(
data.get('allowManagedReadPathsOnly')
),
)
def to_dict(self) -> dict[str, Any]:
return _drop_none({
'allowWrite': self.allow_write,
'denyWrite': self.deny_write,
'denyRead': self.deny_read,
'allowRead': self.allow_read,
'allowManagedReadPathsOnly': self.allow_managed_read_paths_only,
})
@dataclass(frozen=True)
class SandboxRipgrepConfig:
command: str
args: list[str] | None = None
@classmethod
def from_dict(cls, data: dict[str, Any]) -> 'SandboxRipgrepConfig':
command = data.get('command')
if not isinstance(command, str) or not command:
raise ValueError('ripgrep.command must be a non-empty string')
return cls(command=command, args=_str_list(data.get('args')))
def to_dict(self) -> dict[str, Any]:
return _drop_none({'command': self.command, 'args': self.args})
@dataclass(frozen=True)
class SandboxSettings:
enabled: bool | None = None
fail_if_unavailable: bool | None = None
auto_allow_bash_if_sandboxed: bool | None = None
allow_unsandboxed_commands: bool | None = None
network: SandboxNetworkConfig | None = None
filesystem: SandboxFilesystemConfig | None = None
ignore_violations: dict[str, list[str]] | None = None
enable_weaker_nested_sandbox: bool | None = None
enable_weaker_network_isolation: bool | None = None
excluded_commands: list[str] | None = None
ripgrep: SandboxRipgrepConfig | None = None
extra: dict[str, Any] = field(default_factory=dict)
_KNOWN_KEYS = frozenset({
'enabled', 'failIfUnavailable', 'autoAllowBashIfSandboxed',
'allowUnsandboxedCommands', 'network', 'filesystem',
'ignoreViolations', 'enableWeakerNestedSandbox',
'enableWeakerNetworkIsolation', 'excludedCommands', 'ripgrep',
})
@classmethod
def from_dict(cls, data: dict[str, Any]) -> 'SandboxSettings':
if not isinstance(data, dict):
raise ValueError(f'expected mapping, got {type(data).__name__}')
network_raw = data.get('network')
filesystem_raw = data.get('filesystem')
ripgrep_raw = data.get('ripgrep')
ignore_raw = data.get('ignoreViolations')
ignore: dict[str, list[str]] | None = None
if ignore_raw is not None:
if not isinstance(ignore_raw, dict):
raise ValueError('ignoreViolations must be a mapping')
ignore = {}
for key, value in ignore_raw.items():
items = _str_list(value)
ignore[str(key)] = items if items is not None else []
extra = {k: v for k, v in data.items() if k not in cls._KNOWN_KEYS}
return cls(
enabled=_opt_bool(data.get('enabled')),
fail_if_unavailable=_opt_bool(data.get('failIfUnavailable')),
auto_allow_bash_if_sandboxed=_opt_bool(
data.get('autoAllowBashIfSandboxed')
),
allow_unsandboxed_commands=_opt_bool(
data.get('allowUnsandboxedCommands')
),
network=(
SandboxNetworkConfig.from_dict(network_raw)
if isinstance(network_raw, dict) else None
),
filesystem=(
SandboxFilesystemConfig.from_dict(filesystem_raw)
if isinstance(filesystem_raw, dict) else None
),
ignore_violations=ignore,
enable_weaker_nested_sandbox=_opt_bool(
data.get('enableWeakerNestedSandbox')
),
enable_weaker_network_isolation=_opt_bool(
data.get('enableWeakerNetworkIsolation')
),
excluded_commands=_str_list(data.get('excludedCommands')),
ripgrep=(
SandboxRipgrepConfig.from_dict(ripgrep_raw)
if isinstance(ripgrep_raw, dict) else None
),
extra=extra,
)
def to_dict(self) -> dict[str, Any]:
out: dict[str, Any] = _drop_none({
'enabled': self.enabled,
'failIfUnavailable': self.fail_if_unavailable,
'autoAllowBashIfSandboxed': self.auto_allow_bash_if_sandboxed,
'allowUnsandboxedCommands': self.allow_unsandboxed_commands,
'enableWeakerNestedSandbox': self.enable_weaker_nested_sandbox,
'enableWeakerNetworkIsolation': self.enable_weaker_network_isolation,
'excludedCommands': self.excluded_commands,
})
if self.network is not None:
out['network'] = self.network.to_dict()
if self.filesystem is not None:
out['filesystem'] = self.filesystem.to_dict()
if self.ignore_violations is not None:
out['ignoreViolations'] = {
key: list(value) for key, value in self.ignore_violations.items()
}
if self.ripgrep is not None:
out['ripgrep'] = self.ripgrep.to_dict()
# passthrough — preserve unknown keys after serialization
for key, value in self.extra.items():
out.setdefault(key, value)
return out
SandboxIgnoreViolations = dict[str, list[str]]
__all__ = [
'SandboxNetworkConfig',
'SandboxFilesystemConfig',
'SandboxRipgrepConfig',
'SandboxSettings',
'SandboxIgnoreViolations',
]
+319
View File
@@ -0,0 +1,319 @@
"""Foundational SDK types — Python port of entrypoints/sdk/coreTypes.ts.
Mirrors the public constants and a starter set of dataclasses for the
serializable types defined in `entrypoints/sdk/coreSchemas.ts`. The full
schema set (1800+ lines of Zod) is too large to port wholesale; this file
covers the foundational pieces other modules already reference and gives
later slices a place to extend.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Literal
# ---------------------------------------------------------------------------
# Const arrays (HOOK_EVENTS, EXIT_REASONS) — mirror coreTypes.ts top-level
# ---------------------------------------------------------------------------
HOOK_EVENTS: tuple[str, ...] = (
'PreToolUse',
'PostToolUse',
'PostToolUseFailure',
'Notification',
'UserPromptSubmit',
'SessionStart',
'SessionEnd',
'Stop',
'StopFailure',
'SubagentStart',
'SubagentStop',
'PreCompact',
'PostCompact',
'PermissionRequest',
'PermissionDenied',
'Setup',
'TeammateIdle',
'TaskCreated',
'TaskCompleted',
'Elicitation',
'ElicitationResult',
'ConfigChange',
'WorktreeCreate',
'WorktreeRemove',
'InstructionsLoaded',
'CwdChanged',
'FileChanged',
)
EXIT_REASONS: tuple[str, ...] = (
'clear',
'resume',
'logout',
'prompt_input_exit',
'other',
'bypass_permissions_disabled',
)
API_KEY_SOURCES: tuple[str, ...] = ('user', 'project', 'org', 'temporary', 'oauth')
CONFIG_SCOPES: tuple[str, ...] = ('local', 'user', 'project')
SDK_BETAS: tuple[str, ...] = ('context-1m-2025-08-07',)
ApiKeySource = Literal['user', 'project', 'org', 'temporary', 'oauth']
ConfigScope = Literal['local', 'user', 'project']
SdkBeta = Literal['context-1m-2025-08-07']
# ---------------------------------------------------------------------------
# Usage / model
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class ModelUsage:
input_tokens: int
output_tokens: int
cache_read_input_tokens: int
cache_creation_input_tokens: int
web_search_requests: int
cost_usd: float
context_window: int
max_output_tokens: int
@classmethod
def from_dict(cls, data: dict[str, Any]) -> 'ModelUsage':
return cls(
input_tokens=int(data['inputTokens']),
output_tokens=int(data['outputTokens']),
cache_read_input_tokens=int(data['cacheReadInputTokens']),
cache_creation_input_tokens=int(data['cacheCreationInputTokens']),
web_search_requests=int(data['webSearchRequests']),
cost_usd=float(data['costUSD']),
context_window=int(data['contextWindow']),
max_output_tokens=int(data['maxOutputTokens']),
)
def to_dict(self) -> dict[str, Any]:
return {
'inputTokens': self.input_tokens,
'outputTokens': self.output_tokens,
'cacheReadInputTokens': self.cache_read_input_tokens,
'cacheCreationInputTokens': self.cache_creation_input_tokens,
'webSearchRequests': self.web_search_requests,
'costUSD': self.cost_usd,
'contextWindow': self.context_window,
'maxOutputTokens': self.max_output_tokens,
}
# ---------------------------------------------------------------------------
# Thinking config
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class ThinkingAdaptive:
type: Literal['adaptive'] = 'adaptive'
def to_dict(self) -> dict[str, Any]:
return {'type': self.type}
@dataclass(frozen=True)
class ThinkingEnabled:
budget_tokens: int | None = None
type: Literal['enabled'] = 'enabled'
def to_dict(self) -> dict[str, Any]:
out: dict[str, Any] = {'type': self.type}
if self.budget_tokens is not None:
out['budgetTokens'] = self.budget_tokens
return out
@dataclass(frozen=True)
class ThinkingDisabled:
type: Literal['disabled'] = 'disabled'
def to_dict(self) -> dict[str, Any]:
return {'type': self.type}
ThinkingConfig = ThinkingAdaptive | ThinkingEnabled | ThinkingDisabled
def thinking_config_from_dict(data: dict[str, Any]) -> ThinkingConfig:
kind = data.get('type')
if kind == 'adaptive':
return ThinkingAdaptive()
if kind == 'enabled':
budget = data.get('budgetTokens')
return ThinkingEnabled(
budget_tokens=int(budget) if budget is not None else None,
)
if kind == 'disabled':
return ThinkingDisabled()
raise ValueError(f'unknown thinking config type: {kind!r}')
# ---------------------------------------------------------------------------
# MCP server configs
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class McpStdioServerConfig:
command: str
args: list[str] | None = None
env: dict[str, str] | None = None
type: Literal['stdio'] = 'stdio'
def to_dict(self) -> dict[str, Any]:
out: dict[str, Any] = {'type': self.type, 'command': self.command}
if self.args is not None:
out['args'] = list(self.args)
if self.env is not None:
out['env'] = dict(self.env)
return out
@dataclass(frozen=True)
class McpSSEServerConfig:
url: str
headers: dict[str, str] | None = None
type: Literal['sse'] = 'sse'
def to_dict(self) -> dict[str, Any]:
out: dict[str, Any] = {'type': self.type, 'url': self.url}
if self.headers is not None:
out['headers'] = dict(self.headers)
return out
@dataclass(frozen=True)
class McpHttpServerConfig:
url: str
headers: dict[str, str] | None = None
type: Literal['http'] = 'http'
def to_dict(self) -> dict[str, Any]:
out: dict[str, Any] = {'type': self.type, 'url': self.url}
if self.headers is not None:
out['headers'] = dict(self.headers)
return out
@dataclass(frozen=True)
class McpSdkServerConfig:
name: str
type: Literal['sdk'] = 'sdk'
def to_dict(self) -> dict[str, Any]:
return {'type': self.type, 'name': self.name}
@dataclass(frozen=True)
class McpClaudeAIProxyServerConfig:
url: str
id: str
type: Literal['claudeai-proxy'] = 'claudeai-proxy'
def to_dict(self) -> dict[str, Any]:
return {'type': self.type, 'url': self.url, 'id': self.id}
McpServerConfigForProcessTransport = (
McpStdioServerConfig
| McpSSEServerConfig
| McpHttpServerConfig
| McpSdkServerConfig
)
McpServerStatusConfig = (
McpServerConfigForProcessTransport | McpClaudeAIProxyServerConfig
)
def mcp_server_config_from_dict(data: dict[str, Any]) -> McpServerStatusConfig:
"""Discriminate by `type` and route to the matching dataclass."""
raw_type = data.get('type')
if raw_type is None or raw_type == 'stdio':
command = data.get('command')
if not isinstance(command, str):
raise ValueError('stdio server config requires a string command')
return McpStdioServerConfig(
command=command,
args=list(data['args']) if isinstance(data.get('args'), list) else None,
env=dict(data['env']) if isinstance(data.get('env'), dict) else None,
)
if raw_type == 'sse':
return McpSSEServerConfig(
url=str(data['url']),
headers=dict(data['headers']) if isinstance(data.get('headers'), dict) else None,
)
if raw_type == 'http':
return McpHttpServerConfig(
url=str(data['url']),
headers=dict(data['headers']) if isinstance(data.get('headers'), dict) else None,
)
if raw_type == 'sdk':
return McpSdkServerConfig(name=str(data['name']))
if raw_type == 'claudeai-proxy':
return McpClaudeAIProxyServerConfig(
url=str(data['url']),
id=str(data['id']),
)
raise ValueError(f'unknown MCP server config type: {raw_type!r}')
# ---------------------------------------------------------------------------
# Output format
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class JsonSchemaOutputFormat:
schema: dict[str, Any] = field(default_factory=dict)
type: Literal['json_schema'] = 'json_schema'
def to_dict(self) -> dict[str, Any]:
return {'type': self.type, 'schema': dict(self.schema)}
@classmethod
def from_dict(cls, data: dict[str, Any]) -> 'JsonSchemaOutputFormat':
if data.get('type') != 'json_schema':
raise ValueError(
f'expected output format type "json_schema", got {data.get("type")!r}'
)
schema = data.get('schema')
if not isinstance(schema, dict):
raise ValueError('json_schema output format requires a "schema" mapping')
return cls(schema=dict(schema))
OutputFormat = JsonSchemaOutputFormat
__all__ = [
'HOOK_EVENTS',
'EXIT_REASONS',
'API_KEY_SOURCES',
'CONFIG_SCOPES',
'SDK_BETAS',
'ApiKeySource',
'ConfigScope',
'SdkBeta',
'ModelUsage',
'ThinkingAdaptive',
'ThinkingEnabled',
'ThinkingDisabled',
'ThinkingConfig',
'thinking_config_from_dict',
'McpStdioServerConfig',
'McpSSEServerConfig',
'McpHttpServerConfig',
'McpSdkServerConfig',
'McpClaudeAIProxyServerConfig',
'McpServerConfigForProcessTransport',
'McpServerStatusConfig',
'mcp_server_config_from_dict',
'JsonSchemaOutputFormat',
'OutputFormat',
]
+46
View File
@@ -0,0 +1,46 @@
"""Session-scoped environment variables — Python port of
``utils/sessionEnvVars.ts``.
These are env vars set during a session (via the upstream ``/env`` slash
command in npm) and applied only to spawned child processes — not to the
host Python REPL/agent process itself. Bash and similar tool providers
read this map to merge into ``subprocess`` environments.
Mirrors the upstream module-level singleton: callers import the helpers
directly rather than passing a registry around.
"""
from __future__ import annotations
from collections.abc import Mapping
from types import MappingProxyType
_session_env_vars: dict[str, str] = {}
def get_session_env_vars() -> Mapping[str, str]:
"""Return a read-only view of the current session env vars."""
return MappingProxyType(_session_env_vars)
def set_session_env_var(name: str, value: str) -> None:
"""Set ``name=value`` for the rest of this session's child processes."""
_session_env_vars[name] = value
def delete_session_env_var(name: str) -> None:
"""Remove ``name`` from the session env (no-op if absent)."""
_session_env_vars.pop(name, None)
def clear_session_env_vars() -> None:
"""Drop every session-scoped env var."""
_session_env_vars.clear()
__all__ = [
'get_session_env_vars',
'set_session_env_var',
'delete_session_env_var',
'clear_session_env_vars',
]
+77 -2
View File
@@ -2,11 +2,61 @@ from __future__ import annotations
import platform
import sys
from dataclasses import dataclass
from dataclasses import dataclass, field
from pathlib import Path
from .deferred_init import DeferredInitResult, run_deferred_init
from .prefetch import PrefetchResult, start_keychain_prefetch, start_mdm_raw_read, start_project_scan
from .release_notes import check_for_release_notes
MIN_PYTHON_VERSION: tuple[int, int] = (3, 10)
@dataclass(frozen=True)
class RuntimeRequirementCheck:
name: str
ok: bool
detail: str
def check_runtime_requirements() -> tuple[RuntimeRequirementCheck, ...]:
checks: list[RuntimeRequirementCheck] = []
py_major, py_minor = sys.version_info[:2]
py_ok = (py_major, py_minor) >= MIN_PYTHON_VERSION
checks.append(RuntimeRequirementCheck(
name='python_version',
ok=py_ok,
detail=(
f'Python {py_major}.{py_minor} >= {MIN_PYTHON_VERSION[0]}.{MIN_PYTHON_VERSION[1]}'
if py_ok
else f'Python {py_major}.{py_minor} is below required '
f'{MIN_PYTHON_VERSION[0]}.{MIN_PYTHON_VERSION[1]}'
),
))
impl = platform.python_implementation()
checks.append(RuntimeRequirementCheck(
name='python_implementation',
ok=True,
detail=impl,
))
machine = platform.machine() or 'unknown'
system = platform.system() or 'unknown'
checks.append(RuntimeRequirementCheck(
name='platform',
ok=True,
detail=f'{system} on {machine}',
))
return tuple(checks)
def _read_package_version() -> str:
try:
from importlib.metadata import version
return version('claw-code-agent')
except Exception:
return '0.0.0'
@dataclass(frozen=True)
@@ -34,6 +84,11 @@ class SetupReport:
deferred_init: DeferredInitResult
trusted: bool
cwd: Path
runtime_checks: tuple[RuntimeRequirementCheck, ...] = field(default_factory=tuple)
release_notes: tuple[str, ...] = field(default_factory=tuple)
def has_blocking_issues(self) -> bool:
return any(not check.ok for check in self.runtime_checks)
def as_markdown(self) -> str:
lines = [
@@ -44,12 +99,21 @@ class SetupReport:
f'- Trusted mode: {self.trusted}',
f'- CWD: {self.cwd}',
'',
'Runtime checks:',
*(
f'- {check.name}: {"ok" if check.ok else "FAIL"}{check.detail}'
for check in self.runtime_checks
),
'',
'Prefetches:',
*(f'- {prefetch.name}: {prefetch.detail}' for prefetch in self.prefetches),
'',
'Deferred init:',
*self.deferred_init.as_lines(),
]
if self.release_notes:
lines.extend(['', 'Release notes (newer than last seen):'])
lines.extend(f'- {note}' for note in self.release_notes)
return '\n'.join(lines)
@@ -61,17 +125,28 @@ def build_workspace_setup() -> WorkspaceSetup:
)
def run_setup(cwd: Path | None = None, trusted: bool = True) -> SetupReport:
def run_setup(
cwd: Path | None = None,
trusted: bool = True,
last_seen_version: str | None = None,
) -> SetupReport:
root = cwd or Path(__file__).resolve().parent.parent
prefetches = [
start_mdm_raw_read(),
start_keychain_prefetch(),
start_project_scan(root),
]
release_notes_payload = check_for_release_notes(
current_version=_read_package_version(),
last_seen_version=last_seen_version,
cwd=root,
)
return SetupReport(
setup=build_workspace_setup(),
prefetches=tuple(prefetches),
deferred_init=run_deferred_init(trusted=trusted),
trusted=trusted,
cwd=root,
runtime_checks=check_runtime_requirements(),
release_notes=tuple(release_notes_payload['releaseNotes']),
)
+166
View File
@@ -0,0 +1,166 @@
"""Bundled portable utilities — Python ports of small npm `utils/` files.
This module collects narrow, dependency-free helpers from
``utils/array.ts``, ``utils/set.ts``, ``utils/objectGroupBy.ts``,
``utils/xml.ts``, and ``utils/uuid.ts``. Keeping them in one file mirrors
how a user of the npm SDK would reach for these as a small toolbox.
Design notes:
- Function names mirror upstream where idiomatic in Python; ``every`` is
renamed ``every_in`` to avoid shadowing the built-in name in callers.
- ``object_group_by`` returns a plain ``dict`` rather than a ``Mapping`` so
callers can mutate it the same way the JS object would behave.
- ``create_agent_id`` mirrors the upstream format exactly:
``a{label-}{16 hex chars}``.
"""
from __future__ import annotations
import re
import secrets
from collections.abc import Callable, Iterable
from typing import TypeVar
A = TypeVar('A')
T = TypeVar('T')
K = TypeVar('K')
# ---------------------------------------------------------------------------
# array.ts
# ---------------------------------------------------------------------------
def intersperse(items: Iterable[A], separator: Callable[[int], A]) -> list[A]:
"""Insert ``separator(i)`` between consecutive items.
Mirrors ``utils/array.ts`` ``intersperse``: the separator callable
receives the 1-based index of the item it precedes.
"""
out: list[A] = []
for i, item in enumerate(items):
if i:
out.append(separator(i))
out.append(item)
return out
def count(items: Iterable[T], predicate: Callable[[T], object]) -> int:
"""Count items where ``predicate(item)`` is truthy."""
return sum(1 for x in items if predicate(x))
def uniq(items: Iterable[T]) -> list[T]:
"""Return unique items preserving first-seen order.
Note: upstream JS uses ``[...new Set(xs)]`` which preserves insertion
order for primitive values; this matches that behavior.
"""
seen: set[T] = set()
out: list[T] = []
for item in items:
if item not in seen:
seen.add(item)
out.append(item)
return out
# ---------------------------------------------------------------------------
# objectGroupBy.ts
# ---------------------------------------------------------------------------
def object_group_by(
items: Iterable[T],
key_selector: Callable[[T, int], K],
) -> dict[K, list[T]]:
"""Group items by ``key_selector(item, index)``.
Mirrors ``Object.groupBy`` semantics from the TC39 proposal.
"""
out: dict[K, list[T]] = {}
for index, item in enumerate(items):
key = key_selector(item, index)
bucket = out.get(key)
if bucket is None:
bucket = []
out[key] = bucket
bucket.append(item)
return out
# ---------------------------------------------------------------------------
# set.ts
# ---------------------------------------------------------------------------
def difference(a: set[T], b: set[T]) -> set[T]:
"""Items in ``a`` but not ``b``."""
return {item for item in a if item not in b}
def intersects(a: set[T], b: set[T]) -> bool:
"""Whether ``a`` and ``b`` share at least one element."""
if not a or not b:
return False
return any(item in b for item in a)
def every_in(a: set[T], b: set[T]) -> bool:
"""Whether every element of ``a`` is in ``b`` (renamed from ``every``)."""
return all(item in b for item in a)
def union(a: set[T], b: set[T]) -> set[T]:
"""Set union."""
return a | b
# ---------------------------------------------------------------------------
# xml.ts
# ---------------------------------------------------------------------------
def escape_xml(value: str) -> str:
"""Escape ``& < >`` for safe interpolation between XML/HTML tags."""
return value.replace('&', '&amp;').replace('<', '&lt;').replace('>', '&gt;')
def escape_xml_attr(value: str) -> str:
"""Escape ``& < > " '`` for safe interpolation into an attribute value."""
return escape_xml(value).replace('"', '&quot;').replace("'", '&apos;')
# ---------------------------------------------------------------------------
# uuid.ts
# ---------------------------------------------------------------------------
_UUID_RE = re.compile(
r'^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$',
re.IGNORECASE,
)
def validate_uuid(maybe_uuid: object) -> str | None:
"""Return the input as a UUID string if it matches the canonical format."""
if not isinstance(maybe_uuid, str):
return None
return maybe_uuid if _UUID_RE.match(maybe_uuid) else None
def create_agent_id(label: str | None = None) -> str:
"""Generate an agent ID with the upstream ``a{label-}{hex16}`` format."""
suffix = secrets.token_hex(8)
return f'a{label}-{suffix}' if label else f'a{suffix}'
__all__ = [
'intersperse',
'count',
'uniq',
'object_group_by',
'difference',
'intersects',
'every_in',
'union',
'escape_xml',
'escape_xml_attr',
'validate_uuid',
'create_agent_id',
]
+52 -1
View File
@@ -6,7 +6,14 @@ import unittest
from pathlib import Path
from unittest.mock import patch
from src.agent_registry import load_agent_registry, render_agent_detail, render_agents_report
from src.agent_registry import (
create_agent_definition,
delete_agent_definition,
load_agent_registry,
render_agent_detail,
render_agents_report,
update_agent_definition,
)
def _write_agent(path: Path, *, name: str, description: str, body: str, extra: str = '') -> None:
@@ -22,6 +29,50 @@ def _write_agent(path: Path, *, name: str, description: str, body: str, extra: s
class AgentRegistryTests(unittest.TestCase):
def test_create_update_delete_agent_definition_roundtrip(self) -> None:
with tempfile.TemporaryDirectory() as home_dir, tempfile.TemporaryDirectory() as tmp_dir:
workspace = Path(tmp_dir)
with patch.dict(os.environ, {'HOME': home_dir}):
created = create_agent_definition(
workspace,
agent_type='reviewer',
description='Review implementation changes.',
system_prompt='Inspect diffs and summarize risks.',
source='project',
tools=('read_file', 'grep_search'),
model='demo-model',
initial_prompt='Start with changed files.',
)
self.assertEqual(created.action, 'created')
self.assertTrue(Path(created.file_path).exists())
snapshot = load_agent_registry(workspace)
detail = render_agent_detail(snapshot, 'reviewer')
self.assertIn('demo-model', detail)
self.assertIn('Start with changed files.', detail)
updated = update_agent_definition(
workspace,
agent_type='reviewer',
description='Review code and tests carefully.',
system_prompt='Focus on regressions and missing coverage.',
source='project',
)
self.assertEqual(updated.action, 'updated')
snapshot = load_agent_registry(workspace)
detail = render_agent_detail(snapshot, 'reviewer')
self.assertIn('Review code and tests carefully.', detail)
self.assertIn('Focus on regressions and missing coverage.', detail)
deleted = delete_agent_definition(
workspace,
agent_type='reviewer',
source='project',
)
self.assertEqual(deleted.action, 'deleted')
self.assertFalse(Path(deleted.file_path).exists())
def test_project_agent_overrides_built_in_agent(self) -> None:
with tempfile.TemporaryDirectory() as home_dir, tempfile.TemporaryDirectory() as tmp_dir:
workspace = Path(tmp_dir)
+40
View File
@@ -318,6 +318,46 @@ class AgentRuntimeTests(unittest.TestCase):
self.assertEqual(result.tool_calls, 0)
self.assertIn('# Permissions', result.final_output)
def test_slash_skills_lists_bundled_skills(self) -> None:
from src.bundled_skills import get_bundled_skills
with tempfile.TemporaryDirectory() as tmp_dir:
agent = LocalCodingAgent(
model_config=ModelConfig(model='Qwen/Qwen3-Coder-30B-A3B-Instruct'),
runtime_config=AgentRuntimeConfig(cwd=Path(tmp_dir)),
)
result = agent.run('/skills')
self.assertEqual(result.turns, 0)
self.assertEqual(result.tool_calls, 0)
self.assertIn('Available Skills', result.final_output)
for skill in get_bundled_skills():
if skill.user_invocable:
self.assertIn(skill.name, result.final_output)
# Slash command names (e.g. 'permissions') must NOT appear — that was
# the old buggy behavior the upstream `/skills` command never did.
self.assertNotIn('/permissions', result.final_output)
def test_clear_runtime_state_drops_session_env_vars(self) -> None:
from src.session_env_vars import (
clear_session_env_vars,
get_session_env_vars,
set_session_env_var,
)
clear_session_env_vars()
try:
with tempfile.TemporaryDirectory() as tmp_dir:
agent = LocalCodingAgent(
model_config=ModelConfig(model='Qwen/Qwen3-Coder-30B-A3B-Instruct'),
runtime_config=AgentRuntimeConfig(cwd=Path(tmp_dir)),
)
set_session_env_var('CLAW_CLEAR_TEST', 'before')
self.assertEqual(get_session_env_vars()['CLAW_CLEAR_TEST'], 'before')
agent.clear_runtime_state()
self.assertNotIn('CLAW_CLEAR_TEST', get_session_env_vars())
finally:
clear_session_env_vars()
def test_agent_persists_session_and_can_resume(self) -> None:
responses = [
{
+27
View File
@@ -149,6 +149,33 @@ class AgentSlashCommandTests(unittest.TestCase):
self.assertIn('# Agent: reviewer', detail_result.final_output)
self.assertIn('Inspect code changes and summarize risks.', detail_result.final_output)
def test_agents_command_can_create_update_and_delete_project_agent(self) -> None:
with tempfile.TemporaryDirectory() as home_dir, tempfile.TemporaryDirectory() as tmp_dir:
workspace = Path(tmp_dir)
with patch.dict(os.environ, {'HOME': home_dir}):
agent = LocalCodingAgent(
model_config=ModelConfig(model='test-model'),
runtime_config=AgentRuntimeConfig(cwd=workspace),
)
create_result = agent.run(
'/agents create reviewer :: Review code changes carefully. :: Inspect code changes and summarize risks.'
)
self.assertIn('action=created', create_result.final_output)
self.assertTrue((workspace / '.claude' / 'agents' / 'reviewer.md').exists())
update_result = agent.run(
'/agents update reviewer Updated review description :: Focus on regressions and missing tests.'
)
self.assertIn('action=updated', update_result.final_output)
detail_result = agent.run('/agents reviewer')
self.assertIn('Updated review description', detail_result.final_output)
self.assertIn('Focus on regressions and missing tests.', detail_result.final_output)
delete_result = agent.run('/agents delete reviewer')
self.assertIn('action=deleted', delete_result.final_output)
self.assertFalse((workspace / '.claude' / 'agents' / 'reviewer.md').exists())
def test_mcp_and_resource_commands_render_local_reports(self) -> None:
with tempfile.TemporaryDirectory() as tmp_dir:
workspace = Path(tmp_dir)
+35
View File
@@ -9,6 +9,7 @@ from pathlib import Path
from src.agent_tools import (
ToolExecutionError,
ToolPermissionError,
_build_subprocess_env,
_ensure_shell_allowed,
_is_sensitive_env_var,
_resolve_path,
@@ -16,6 +17,10 @@ from src.agent_tools import (
default_tool_registry,
)
from src.agent_types import AgentPermissions, AgentRuntimeConfig
from src.session_env_vars import (
clear_session_env_vars,
set_session_env_var,
)
def _make_context(
@@ -230,5 +235,35 @@ class TestIsSensitiveEnvVar(unittest.TestCase):
self.assertTrue(_is_sensitive_env_var("db_password"))
# ---------------------------------------------------------------------------
# _build_subprocess_env session env var merging
# ---------------------------------------------------------------------------
class TestBuildSubprocessEnv(unittest.TestCase):
def setUp(self):
clear_session_env_vars()
def tearDown(self):
clear_session_env_vars()
def test_session_env_var_appears_in_subprocess_env(self):
with tempfile.TemporaryDirectory() as tmp:
ctx = _make_context(tmp)
set_session_env_var("CLAW_SESSION_FOO", "from-session")
env = _build_subprocess_env(ctx)
self.assertEqual(env["CLAW_SESSION_FOO"], "from-session")
def test_extra_env_overrides_session_env(self):
with tempfile.TemporaryDirectory() as tmp:
config = AgentRuntimeConfig(cwd=Path(tmp))
ctx = build_tool_context(
config,
tool_registry=default_tool_registry(),
extra_env={"CLAW_OVERRIDE_ME": "from-extra"},
)
set_session_env_var("CLAW_OVERRIDE_ME", "from-session")
env = _build_subprocess_env(ctx)
self.assertEqual(env["CLAW_OVERRIDE_ME"], "from-extra")
if __name__ == "__main__":
unittest.main()
+145
View File
@@ -0,0 +1,145 @@
"""Tests for the informational slash commands ported from the npm source.
Covers /output-style, /release-notes, /feedback, /upgrade, /stickers, /mobile,
/desktop, /install-github-app, /install-slack-app, /privacy-settings,
/extra-usage, /passes, /rate-limit-options, /chrome, /reload-plugins.
"""
from __future__ import annotations
import os
import tempfile
import unittest
from pathlib import Path
from src.agent_runtime import LocalCodingAgent
from src.agent_types import AgentRuntimeConfig, ModelConfig
def _make_agent(tmp_dir: str) -> LocalCodingAgent:
return LocalCodingAgent(
model_config=ModelConfig(model='test-model'),
runtime_config=AgentRuntimeConfig(cwd=Path(tmp_dir)),
)
class ExternalSlashCommandsTest(unittest.TestCase):
"""Each test runs with CLAUDE_CODE_NO_BROWSER=1 so no browser opens."""
def setUp(self) -> None:
os.environ['CLAUDE_CODE_NO_BROWSER'] = '1'
def tearDown(self) -> None:
os.environ.pop('CLAUDE_CODE_NO_BROWSER', None)
def _run(self, cmd: str) -> str:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
return agent.run(cmd).final_output
def test_output_style_is_deprecated(self) -> None:
out = self._run('/output-style')
self.assertIn('deprecated', out.lower())
self.assertIn('/config', out)
def test_release_notes_falls_back_to_link(self) -> None:
out = self._run('/release-notes')
self.assertIn('CHANGELOG.md', out)
self.assertIn('https://github.com/anthropics/claude-code', out)
def test_release_notes_reads_local_changelog(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
(Path(tmp) / 'CHANGELOG.md').write_text(
'# Changelog\n\n## 1.2.3\n- did a thing\n\n## 1.2.2\n- old\n',
encoding='utf-8',
)
agent = _make_agent(tmp)
out = agent.run('/release-notes').final_output
self.assertIn('1.2.3', out)
self.assertIn('did a thing', out)
self.assertNotIn('1.2.2', out)
def test_feedback_returns_link(self) -> None:
out = self._run('/feedback')
self.assertIn('https://github.com/anthropics/claude-code/issues', out)
def test_bug_aliases_to_feedback(self) -> None:
out = self._run('/bug')
self.assertIn('https://github.com/anthropics/claude-code/issues', out)
def test_feedback_includes_user_note(self) -> None:
out = self._run('/feedback the wrap selector keeps eating my prompt')
self.assertIn('Draft note', out)
self.assertIn('wrap selector', out)
def test_upgrade_returns_link(self) -> None:
out = self._run('/upgrade')
self.assertIn('https://claude.ai/upgrade/max', out)
def test_stickers_returns_link(self) -> None:
out = self._run('/stickers')
self.assertIn('stickermule.com/claudecode', out)
def test_mobile_lists_both_stores(self) -> None:
out = self._run('/mobile')
self.assertIn('apps.apple.com', out)
self.assertIn('play.google.com', out)
def test_ios_alias(self) -> None:
out = self._run('/ios')
self.assertIn('apps.apple.com', out)
def test_android_alias(self) -> None:
out = self._run('/android')
self.assertIn('play.google.com', out)
def test_desktop_returns_link(self) -> None:
out = self._run('/desktop')
self.assertIn('claude.ai/download', out)
def test_app_aliases_to_desktop(self) -> None:
out = self._run('/app')
self.assertIn('claude.ai/download', out)
def test_install_github_app(self) -> None:
out = self._run('/install-github-app')
self.assertIn('github.com/apps/claude', out)
def test_install_slack_app(self) -> None:
out = self._run('/install-slack-app')
self.assertIn('slack.com/marketplace/A08SF47R6P4-claude', out)
def test_privacy_settings(self) -> None:
out = self._run('/privacy-settings')
self.assertIn('claude.ai/settings/data-privacy-controls', out)
def test_extra_usage_points_to_upgrade(self) -> None:
out = self._run('/extra-usage')
self.assertIn('claude.ai/upgrade/max', out)
self.assertIn('/login', out)
def test_passes_mentions_claude_ai(self) -> None:
out = self._run('/passes')
self.assertIn('claude.ai', out.lower())
self.assertIn('passes', out.lower())
def test_rate_limit_options_lists_actions(self) -> None:
out = self._run('/rate-limit-options')
self.assertIn('/upgrade', out)
self.assertIn('/extra-usage', out)
self.assertIn('/login', out)
def test_chrome_returns_link(self) -> None:
out = self._run('/chrome')
self.assertIn('claude.ai/chrome', out)
def test_reload_plugins_reports_counts(self) -> None:
out = self._run('/reload-plugins')
self.assertIn('Reloaded plugins', out)
self.assertIn('plugin(s)', out)
self.assertIn('tool(s)', out)
self.assertIn('hook(s)', out)
if __name__ == '__main__':
unittest.main()
+114
View File
@@ -0,0 +1,114 @@
"""Tests for ``src/format_utils.py``."""
from __future__ import annotations
import unittest
from src.format_utils import (
format_duration,
format_file_size,
format_number,
format_seconds_short,
format_tokens,
)
class FormatFileSizeTest(unittest.TestCase):
def test_bytes(self) -> None:
self.assertEqual(format_file_size(0), '0 bytes')
self.assertEqual(format_file_size(512), '512 bytes')
def test_kb_with_decimal(self) -> None:
self.assertEqual(format_file_size(1536), '1.5KB')
def test_kb_trims_trailing_zero(self) -> None:
self.assertEqual(format_file_size(2048), '2KB')
def test_mb(self) -> None:
self.assertEqual(format_file_size(5 * 1024 * 1024), '5MB')
def test_gb(self) -> None:
self.assertEqual(format_file_size(2 * 1024 * 1024 * 1024), '2GB')
class FormatSecondsShortTest(unittest.TestCase):
def test_basic(self) -> None:
self.assertEqual(format_seconds_short(1234), '1.2s')
def test_under_one(self) -> None:
self.assertEqual(format_seconds_short(450), '0.5s')
class FormatDurationTest(unittest.TestCase):
def test_zero(self) -> None:
self.assertEqual(format_duration(0), '0s')
def test_sub_second(self) -> None:
self.assertEqual(format_duration(0.5), '0.0s')
def test_seconds_only(self) -> None:
self.assertEqual(format_duration(5_000), '5s')
def test_minutes_seconds(self) -> None:
self.assertEqual(format_duration(125_000), '2m 5s')
def test_hours(self) -> None:
self.assertEqual(format_duration(3_725_000), '1h 2m 5s')
def test_days(self) -> None:
# 1d 2h 3m
ms = 86_400_000 + 2 * 3_600_000 + 3 * 60_000 + 0
self.assertEqual(format_duration(ms), '1d 2h 3m')
def test_hide_trailing_zeros_minutes(self) -> None:
self.assertEqual(
format_duration(120_000, hide_trailing_zeros=True), '2m',
)
def test_hide_trailing_zeros_hours(self) -> None:
self.assertEqual(
format_duration(3_600_000, hide_trailing_zeros=True), '1h',
)
def test_most_significant_only_picks_largest_unit(self) -> None:
self.assertEqual(format_duration(125_000, most_significant_only=True), '2m')
self.assertEqual(
format_duration(3_725_000, most_significant_only=True), '1h',
)
def test_rounding_carry_over(self) -> None:
# 59,500 ms rounds seconds=60 → carries to 1m 0s
self.assertEqual(format_duration(59_500 + 60_000), '2m 0s')
class FormatNumberTest(unittest.TestCase):
def test_below_thousand(self) -> None:
self.assertEqual(format_number(900), '900')
self.assertEqual(format_number(0), '0')
def test_thousands_with_decimal(self) -> None:
self.assertEqual(format_number(1321), '1.3k')
def test_thousand_keeps_decimal(self) -> None:
self.assertEqual(format_number(1000), '1.0k')
def test_millions(self) -> None:
self.assertEqual(format_number(2_500_000), '2.5m')
def test_billions(self) -> None:
self.assertEqual(format_number(3_700_000_000), '3.7b')
class FormatTokensTest(unittest.TestCase):
def test_trims_decimal_zero(self) -> None:
self.assertEqual(format_tokens(1000), '1k')
def test_keeps_meaningful_decimal(self) -> None:
self.assertEqual(format_tokens(1321), '1.3k')
def test_below_thousand(self) -> None:
self.assertEqual(format_tokens(450), '450')
if __name__ == '__main__':
unittest.main()
+164
View File
@@ -0,0 +1,164 @@
"""Tests for ``src/git_utils.py``."""
from __future__ import annotations
import hashlib
import tempfile
import unittest
from pathlib import Path
from src import git_utils
from src.git_utils import (
find_git_root,
get_repo_remote_hash,
normalize_git_remote_url,
should_include_git_instructions,
)
class FindGitRootTest(unittest.TestCase):
def setUp(self) -> None:
find_git_root.cache_clear()
def test_finds_git_directory(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
(Path(tmp) / '.git').mkdir()
self.assertEqual(find_git_root(tmp), str(Path(tmp).resolve()))
def test_walks_up_from_subdirectory(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
(Path(tmp) / '.git').mkdir()
sub = Path(tmp) / 'a' / 'b' / 'c'
sub.mkdir(parents=True)
self.assertEqual(find_git_root(str(sub)), str(Path(tmp).resolve()))
def test_finds_when_git_is_a_file(self) -> None:
# Worktrees and submodules use a .git file containing 'gitdir: ...'
with tempfile.TemporaryDirectory() as tmp:
(Path(tmp) / '.git').write_text('gitdir: /elsewhere/.git/worktrees/x')
self.assertEqual(find_git_root(tmp), str(Path(tmp).resolve()))
def test_returns_none_when_not_in_repo(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
self.assertIsNone(find_git_root(tmp))
def test_memoizes(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
(Path(tmp) / '.git').mkdir()
first = find_git_root(tmp)
# Now remove .git — second lookup should still hit cache
(Path(tmp) / '.git').rmdir()
self.assertEqual(find_git_root(tmp), first)
class NormalizeGitRemoteUrlTest(unittest.TestCase):
def test_ssh_url(self) -> None:
self.assertEqual(
normalize_git_remote_url('git@github.com:org/repo.git'),
'github.com/org/repo',
)
def test_ssh_url_no_dot_git_suffix(self) -> None:
self.assertEqual(
normalize_git_remote_url('git@github.com:org/repo'),
'github.com/org/repo',
)
def test_https_url(self) -> None:
self.assertEqual(
normalize_git_remote_url('https://github.com/org/repo.git'),
'github.com/org/repo',
)
def test_https_with_user(self) -> None:
self.assertEqual(
normalize_git_remote_url('https://user@github.com/org/repo.git'),
'github.com/org/repo',
)
def test_ssh_protocol_url(self) -> None:
self.assertEqual(
normalize_git_remote_url('ssh://git@github.com/org/repo.git'),
'github.com/org/repo',
)
def test_lowercases_result(self) -> None:
self.assertEqual(
normalize_git_remote_url('git@GitHub.com:Org/Repo.git'),
'github.com/org/repo',
)
def test_ccr_proxy_legacy_assumes_github(self) -> None:
self.assertEqual(
normalize_git_remote_url('http://x@127.0.0.1:8080/git/org/repo'),
'github.com/org/repo',
)
def test_ccr_proxy_ghe_uses_first_segment_as_host(self) -> None:
self.assertEqual(
normalize_git_remote_url(
'http://x@127.0.0.1:8080/git/ghe.example.com/org/repo',
),
'ghe.example.com/org/repo',
)
def test_returns_none_on_garbage(self) -> None:
self.assertIsNone(normalize_git_remote_url(''))
self.assertIsNone(normalize_git_remote_url(' '))
self.assertIsNone(normalize_git_remote_url('not-a-url'))
def test_localhost_alias(self) -> None:
self.assertEqual(
normalize_git_remote_url('http://localhost:8080/git/org/repo'),
'github.com/org/repo',
)
class GetRepoRemoteHashTest(unittest.TestCase):
def test_hashes_normalized_url(self) -> None:
url = 'git@github.com:Org/Repo.git'
normalized = 'github.com/org/repo'
expected = hashlib.sha256(normalized.encode('utf-8')).hexdigest()[:16]
self.assertEqual(get_repo_remote_hash(url), expected)
def test_returns_none_for_empty(self) -> None:
self.assertIsNone(get_repo_remote_hash(None))
self.assertIsNone(get_repo_remote_hash(''))
def test_returns_none_for_unparseable(self) -> None:
self.assertIsNone(get_repo_remote_hash('not-a-url'))
class ShouldIncludeGitInstructionsTest(unittest.TestCase):
def test_default_true_when_nothing_set(self) -> None:
self.assertTrue(should_include_git_instructions(env={}))
def test_env_truthy_disables(self) -> None:
for value in ('1', 'true', 'yes', 'on'):
self.assertFalse(
should_include_git_instructions(
settings_value=True,
env={'CLAUDE_CODE_DISABLE_GIT_INSTRUCTIONS': value},
),
)
def test_env_falsy_overrides_settings(self) -> None:
# Settings says off, but env explicitly says don't disable → on
self.assertTrue(
should_include_git_instructions(
settings_value=False,
env={'CLAUDE_CODE_DISABLE_GIT_INSTRUCTIONS': '0'},
),
)
def test_settings_value_used_when_env_unset(self) -> None:
self.assertFalse(
should_include_git_instructions(settings_value=False, env={}),
)
self.assertTrue(
should_include_git_instructions(settings_value=True, env={}),
)
if __name__ == '__main__':
unittest.main()
+133
View File
@@ -0,0 +1,133 @@
"""Tests for ``src/ide_path_conversion.py``."""
from __future__ import annotations
import subprocess
import unittest
from unittest import mock
from src.ide_path_conversion import (
WindowsToWSLConverter,
check_wsl_distro_match,
)
class CheckWslDistroMatchTest(unittest.TestCase):
def test_matches_named_distro(self) -> None:
self.assertTrue(
check_wsl_distro_match(r'\\wsl$\Ubuntu\home\me', 'Ubuntu'),
)
def test_matches_localhost_form(self) -> None:
self.assertTrue(
check_wsl_distro_match(
r'\\wsl.localhost\Ubuntu\home\me', 'Ubuntu',
),
)
def test_mismatch(self) -> None:
self.assertFalse(
check_wsl_distro_match(r'\\wsl$\Debian\home\me', 'Ubuntu'),
)
def test_non_unc_path_returns_true(self) -> None:
self.assertTrue(check_wsl_distro_match(r'C:\Users\me', 'Ubuntu'))
class WindowsToWSLConverterToLocalPathTest(unittest.TestCase):
def test_empty_path_passthrough(self) -> None:
conv = WindowsToWSLConverter('Ubuntu')
self.assertEqual(conv.to_local_path(''), '')
def test_uses_wslpath_when_available(self) -> None:
conv = WindowsToWSLConverter(None)
with mock.patch(
'src.ide_path_conversion.subprocess.run',
return_value=subprocess.CompletedProcess(
args=[], returncode=0, stdout='/mnt/c/Users/me\n', stderr='',
),
) as run:
self.assertEqual(conv.to_local_path(r'C:\Users\me'), '/mnt/c/Users/me')
run.assert_called_once()
args = run.call_args[0][0]
self.assertEqual(args, ['wslpath', '-u', r'C:\Users\me'])
def test_falls_back_to_manual_when_wslpath_missing(self) -> None:
conv = WindowsToWSLConverter(None)
with mock.patch(
'src.ide_path_conversion.subprocess.run',
side_effect=FileNotFoundError(),
):
self.assertEqual(
conv.to_local_path(r'C:\Users\me'), '/mnt/c/Users/me',
)
def test_falls_back_to_manual_on_called_process_error(self) -> None:
conv = WindowsToWSLConverter(None)
with mock.patch(
'src.ide_path_conversion.subprocess.run',
side_effect=subprocess.CalledProcessError(1, 'wslpath'),
):
self.assertEqual(
conv.to_local_path(r'D:\path\to\file.txt'),
'/mnt/d/path/to/file.txt',
)
def test_different_distro_path_returned_as_is(self) -> None:
conv = WindowsToWSLConverter('Ubuntu')
with mock.patch(
'src.ide_path_conversion.subprocess.run',
) as run:
self.assertEqual(
conv.to_local_path(r'\\wsl$\Debian\home\me'),
r'\\wsl$\Debian\home\me',
)
run.assert_not_called()
def test_same_distro_unc_uses_wslpath(self) -> None:
conv = WindowsToWSLConverter('Ubuntu')
with mock.patch(
'src.ide_path_conversion.subprocess.run',
return_value=subprocess.CompletedProcess(
args=[], returncode=0, stdout='/home/me\n', stderr='',
),
) as run:
self.assertEqual(
conv.to_local_path(r'\\wsl$\Ubuntu\home\me'),
'/home/me',
)
run.assert_called_once()
class WindowsToWSLConverterToIdePathTest(unittest.TestCase):
def test_empty_passthrough(self) -> None:
conv = WindowsToWSLConverter(None)
self.assertEqual(conv.to_ide_path(''), '')
def test_uses_wslpath(self) -> None:
conv = WindowsToWSLConverter(None)
with mock.patch(
'src.ide_path_conversion.subprocess.run',
return_value=subprocess.CompletedProcess(
args=[], returncode=0, stdout=r'\\wsl$\Ubuntu\home\me' + '\n',
stderr='',
),
) as run:
self.assertEqual(
conv.to_ide_path('/home/me'), r'\\wsl$\Ubuntu\home\me',
)
run.assert_called_once()
args = run.call_args[0][0]
self.assertEqual(args, ['wslpath', '-w', '/home/me'])
def test_returns_original_on_failure(self) -> None:
conv = WindowsToWSLConverter(None)
with mock.patch(
'src.ide_path_conversion.subprocess.run',
side_effect=FileNotFoundError(),
):
self.assertEqual(conv.to_ide_path('/home/me'), '/home/me')
if __name__ == '__main__':
unittest.main()
+50 -1
View File
@@ -7,7 +7,7 @@ from dataclasses import replace
from pathlib import Path
from unittest.mock import patch
from src.main import _build_runtime_config, _build_agent, _run_agent_chat_loop, build_parser
from src.main import _build_runtime_config, _build_agent, _run_agent_chat_loop, build_parser, main
class FakeHTTPResponse:
@@ -221,6 +221,55 @@ class MainCliTests(unittest.TestCase):
self.assertTrue(args.all)
self.assertEqual(args.cwd, '.')
def test_parser_accepts_agent_management_commands(self) -> None:
parser = build_parser()
create_args = parser.parse_args(
['agents-create', 'reviewer', '--cwd', '.', '--description', 'Review code', '--prompt', 'Inspect diffs']
)
self.assertEqual(create_args.command, 'agents-create')
self.assertEqual(create_args.agent_type, 'reviewer')
self.assertEqual(create_args.description, 'Review code')
update_args = parser.parse_args(['agents-update', 'reviewer', '--cwd', '.', '--source', 'auto'])
self.assertEqual(update_args.command, 'agents-update')
self.assertEqual(update_args.source, 'auto')
delete_args = parser.parse_args(['agents-delete', 'reviewer', '--cwd', '.', '--source', 'project'])
self.assertEqual(delete_args.command, 'agents-delete')
self.assertEqual(delete_args.source, 'project')
def test_main_can_create_and_delete_agent_definition(self) -> None:
with tempfile.TemporaryDirectory() as home_dir, tempfile.TemporaryDirectory() as tmp_dir:
workspace = Path(tmp_dir)
with patch.dict('os.environ', {'HOME': home_dir}):
exit_code = main(
[
'agents-create',
'reviewer',
'--cwd',
str(workspace),
'--description',
'Review code carefully',
'--prompt',
'Inspect code and summarize risks.',
]
)
self.assertEqual(exit_code, 0)
self.assertTrue((workspace / '.claude' / 'agents' / 'reviewer.md').exists())
exit_code = main(
[
'agents-delete',
'reviewer',
'--cwd',
str(workspace),
'--source',
'project',
]
)
self.assertEqual(exit_code, 0)
self.assertFalse((workspace / '.claude' / 'agents' / 'reviewer.md').exists())
def test_parser_accepts_team_runtime_commands(self) -> None:
parser = build_parser()
args = parser.parse_args(['team-create', 'reviewers', '--member', 'alice', '--cwd', '.'])
+143
View File
@@ -0,0 +1,143 @@
"""Tests for discovery slash commands ported from the npm source.
Covers /version, /init, /ide, /plugin, /remote-env.
"""
from __future__ import annotations
import json
import os
import tempfile
import unittest
from pathlib import Path
from unittest import mock
from src.agent_runtime import LocalCodingAgent
from src.agent_slash_commands import preprocess_slash_command
from src.agent_types import AgentRuntimeConfig, ModelConfig
def _make_agent(tmp_dir: str) -> LocalCodingAgent:
return LocalCodingAgent(
model_config=ModelConfig(model='test-model'),
runtime_config=AgentRuntimeConfig(cwd=Path(tmp_dir)),
)
def _local_settings(tmp_dir: str) -> dict:
path = Path(tmp_dir) / '.claude' / 'settings.local.json'
if not path.exists():
return {}
return json.loads(path.read_text(encoding='utf-8'))
class VersionCommandTest(unittest.TestCase):
def test_prints_python_runtime_version(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/version').final_output
self.assertIn('claw-code-agent', out)
self.assertIn('Python', out)
class InitCommandTest(unittest.TestCase):
def test_returns_prompt_with_claude_md_instructions(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
result = preprocess_slash_command(agent, '/init')
self.assertTrue(result.handled)
self.assertTrue(result.should_query)
self.assertIn('CLAUDE.md', result.prompt or '')
self.assertIn('analyze this codebase', (result.prompt or '').lower())
class IdeCommandTest(unittest.TestCase):
def test_no_ide_when_env_clean(self) -> None:
clean = {k: v for k, v in os.environ.items() if k not in {
'TERM_PROGRAM', 'VSCODE_INJECTION', 'VSCODE_PID',
'JETBRAINS_IDE', 'TERMINAL_EMULATOR', 'SSH_CONNECTION',
}}
with mock.patch.dict(os.environ, clean, clear=True):
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/ide').final_output
self.assertIn('No IDE detected', out)
self.assertIn('IDE auto-connect', out)
def test_detects_vscode(self) -> None:
env = {'VSCODE_PID': '1234', 'TERM_PROGRAM': 'vscode'}
with mock.patch.dict(os.environ, env, clear=True):
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/ide').final_output
self.assertIn('Visual Studio Code', out)
self.assertIn('VSCODE_PID=1234', out)
class PluginCommandTest(unittest.TestCase):
def test_lists_no_plugins_when_empty(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/plugin').final_output
self.assertIn('No installed plugins', out)
def test_help_describes_usage(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/plugin help').final_output
self.assertIn('Usage: /plugin', out)
self.assertIn('list', out)
def test_unknown_subcommand(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/plugin bogus').final_output
self.assertIn('Unknown plugin subcommand', out)
class RemoteEnvCommandTest(unittest.TestCase):
def test_lists_empty_profiles(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/remote-env').final_output
self.assertIn('Available remote environments', out)
self.assertIn('no profiles found', out)
self.assertIn('Usage:', out)
def test_clear_when_no_default_set(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/remote-env clear').final_output
self.assertIn('No default remote environment', out)
def test_unknown_profile_is_rejected(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/remote-env nope').final_output
self.assertIn('Unknown remote environment', out)
self.assertIn('nope', out)
def test_set_then_clear_persists(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
(Path(tmp) / '.remote.json').write_text(json.dumps({
'profiles': [
{'name': 'sandbox', 'mode': 'ssh', 'target': 'user@host'},
],
}), encoding='utf-8')
agent = _make_agent(tmp)
set_out = agent.run('/remote-env sandbox').final_output
self.assertIn('Default remote environment set to sandbox', set_out)
self.assertEqual(_local_settings(tmp).get('defaultRemoteEnvironment'), 'sandbox')
agent2 = _make_agent(tmp)
list_out = agent2.run('/remote-env').final_output
self.assertIn('sandbox', list_out)
self.assertIn('(default)', list_out)
clear_out = agent2.run('/remote-env clear').final_output
self.assertIn('Cleared default remote environment', clear_out)
self.assertIsNone(_local_settings(tmp).get('defaultRemoteEnvironment'))
if __name__ == '__main__':
unittest.main()
+167
View File
@@ -0,0 +1,167 @@
"""Tests for model pricing utilities ported from utils/modelCost.ts."""
from __future__ import annotations
import unittest
from src.model_cost import (
COST_HAIKU_35,
COST_HAIKU_45,
COST_TIER_3_15,
COST_TIER_5_25,
COST_TIER_15_75,
COST_TIER_30_150,
DEFAULT_UNKNOWN_MODEL_COST,
calculate_cost_from_tokens,
calculate_usd_cost,
format_model_pricing,
get_model_costs,
get_model_pricing_string,
get_opus_4_6_cost_tier,
tokens_to_usd_cost,
)
class TierConstantsTest(unittest.TestCase):
def test_sonnet_tier(self) -> None:
self.assertEqual(COST_TIER_3_15.input_tokens, 3.0)
self.assertEqual(COST_TIER_3_15.output_tokens, 15.0)
def test_opus_4_tier(self) -> None:
self.assertEqual(COST_TIER_15_75.input_tokens, 15.0)
def test_opus_4_5_tier(self) -> None:
self.assertEqual(COST_TIER_5_25.input_tokens, 5.0)
def test_fast_mode_tier(self) -> None:
self.assertEqual(COST_TIER_30_150.input_tokens, 30.0)
def test_haiku_tiers(self) -> None:
self.assertAlmostEqual(COST_HAIKU_35.input_tokens, 0.8)
self.assertEqual(COST_HAIKU_45.input_tokens, 1.0)
class GetModelCostsTest(unittest.TestCase):
def test_opus_4_6_default(self) -> None:
self.assertIs(get_model_costs('claude-opus-4-6'), COST_TIER_5_25)
def test_opus_4_6_fast_mode(self) -> None:
self.assertIs(
get_model_costs('claude-opus-4-6', fast_mode=True),
COST_TIER_30_150,
)
def test_versioned_model_name_resolves(self) -> None:
self.assertIs(
get_model_costs('claude-opus-4-6-20251015'),
COST_TIER_5_25,
)
def test_sonnet_models_use_3_15(self) -> None:
for name in ('claude-sonnet-4-6', 'claude-sonnet-4-5', 'claude-sonnet-4'):
self.assertIs(get_model_costs(name), COST_TIER_3_15)
def test_opus_4_and_4_1_use_15_75(self) -> None:
self.assertIs(get_model_costs('claude-opus-4'), COST_TIER_15_75)
self.assertIs(get_model_costs('claude-opus-4-1'), COST_TIER_15_75)
def test_haiku_4_5(self) -> None:
self.assertIs(get_model_costs('claude-haiku-4-5-20251001'), COST_HAIKU_45)
def test_haiku_3_5(self) -> None:
self.assertIs(get_model_costs('claude-3-5-haiku-20241022'), COST_HAIKU_35)
def test_unknown_falls_back_to_default(self) -> None:
self.assertIs(get_model_costs('mystery-llm-3000'), DEFAULT_UNKNOWN_MODEL_COST)
def test_get_opus_4_6_helper_matches_fast_mode(self) -> None:
self.assertIs(get_opus_4_6_cost_tier(False), COST_TIER_5_25)
self.assertIs(get_opus_4_6_cost_tier(True), COST_TIER_30_150)
class TokensToUsdCostTest(unittest.TestCase):
def test_simple_input_output(self) -> None:
cost = tokens_to_usd_cost(
COST_TIER_3_15, input_tokens=1_000_000, output_tokens=500_000,
)
self.assertAlmostEqual(cost, 3.0 + 7.5)
def test_includes_cache_tokens(self) -> None:
cost = tokens_to_usd_cost(
COST_TIER_3_15,
input_tokens=0,
output_tokens=0,
cache_read_input_tokens=1_000_000,
cache_creation_input_tokens=1_000_000,
)
self.assertAlmostEqual(
cost,
COST_TIER_3_15.prompt_cache_read_tokens
+ COST_TIER_3_15.prompt_cache_write_tokens,
)
def test_includes_web_search(self) -> None:
cost = tokens_to_usd_cost(
COST_TIER_3_15,
input_tokens=0,
output_tokens=0,
web_search_requests=10,
)
self.assertAlmostEqual(cost, 0.10)
class CalculateUsdCostTest(unittest.TestCase):
def test_resolves_model_then_costs(self) -> None:
cost = calculate_usd_cost(
'claude-sonnet-4-6',
input_tokens=1_000_000,
output_tokens=500_000,
)
self.assertAlmostEqual(cost, 10.5)
def test_fast_mode_changes_opus_46_cost(self) -> None:
normal = calculate_usd_cost(
'claude-opus-4-6', input_tokens=1_000_000, output_tokens=0,
)
fast = calculate_usd_cost(
'claude-opus-4-6', input_tokens=1_000_000, output_tokens=0,
fast_mode=True,
)
self.assertGreater(fast, normal)
self.assertAlmostEqual(normal, 5.0)
self.assertAlmostEqual(fast, 30.0)
class CalculateCostFromTokensTest(unittest.TestCase):
def test_camel_case_dict_input(self) -> None:
cost = calculate_cost_from_tokens(
'claude-opus-4-1',
{
'inputTokens': 1_000_000,
'outputTokens': 0,
'cacheReadInputTokens': 0,
'cacheCreationInputTokens': 0,
},
)
self.assertAlmostEqual(cost, 15.0)
class FormatPricingTest(unittest.TestCase):
def test_integers_no_decimals(self) -> None:
self.assertEqual(format_model_pricing(COST_TIER_3_15), '$3/$15 per Mtok')
def test_haiku_decimals(self) -> None:
self.assertEqual(format_model_pricing(COST_HAIKU_35), '$0.80/$4 per Mtok')
def test_get_pricing_string_known(self) -> None:
self.assertEqual(
get_model_pricing_string('claude-opus-4-6'),
'$5/$25 per Mtok',
)
def test_get_pricing_string_unknown(self) -> None:
self.assertIsNone(get_model_pricing_string('unknown-model'))
if __name__ == '__main__':
unittest.main()
+210
View File
@@ -0,0 +1,210 @@
"""Tests for ``src/platform_info.py`` — platform detection and system dirs."""
from __future__ import annotations
import os
import tempfile
import unittest
from pathlib import Path
from unittest import mock
from src import platform_info
from src.platform_info import (
SUPPORTED_PLATFORMS,
LinuxDistroInfo,
SystemDirectories,
detect_vcs,
get_linux_distro_info,
get_platform,
get_system_directories,
get_wsl_version,
)
class GetPlatformTest(unittest.TestCase):
def setUp(self) -> None:
platform_info._reset_cache()
def tearDown(self) -> None:
platform_info._reset_cache()
def test_macos(self) -> None:
with mock.patch('src.platform_info.sys.platform', 'darwin'):
self.assertEqual(get_platform(), 'macos')
def test_windows(self) -> None:
with mock.patch('src.platform_info.sys.platform', 'win32'):
self.assertEqual(get_platform(), 'windows')
def test_linux_no_wsl(self) -> None:
with mock.patch('src.platform_info.sys.platform', 'linux'), \
mock.patch(
'src.platform_info._read_proc_version',
return_value='Linux version 5.10 (gcc)',
):
self.assertEqual(get_platform(), 'linux')
def test_linux_wsl_microsoft_marker(self) -> None:
with mock.patch('src.platform_info.sys.platform', 'linux'), \
mock.patch(
'src.platform_info._read_proc_version',
return_value='Linux version 5.10 microsoft-standard-WSL2',
):
self.assertEqual(get_platform(), 'wsl')
def test_linux_proc_version_unreadable(self) -> None:
with mock.patch('src.platform_info.sys.platform', 'linux'), \
mock.patch(
'src.platform_info._read_proc_version',
side_effect=FileNotFoundError(),
):
self.assertEqual(get_platform(), 'linux')
def test_unknown_platform(self) -> None:
with mock.patch('src.platform_info.sys.platform', 'sunos5'):
self.assertEqual(get_platform(), 'unknown')
def test_memoized(self) -> None:
with mock.patch('src.platform_info.sys.platform', 'darwin'):
self.assertEqual(get_platform(), 'macos')
# Second call should hit cache, not re-evaluate sys.platform
with mock.patch('src.platform_info.sys.platform', 'win32'):
self.assertEqual(get_platform(), 'macos')
def test_supported_platforms_contains_expected(self) -> None:
self.assertIn('macos', SUPPORTED_PLATFORMS)
self.assertIn('wsl', SUPPORTED_PLATFORMS)
class GetWslVersionTest(unittest.TestCase):
def setUp(self) -> None:
platform_info._reset_cache()
def tearDown(self) -> None:
platform_info._reset_cache()
def test_explicit_wsl2(self) -> None:
with mock.patch('src.platform_info.sys.platform', 'linux'), \
mock.patch(
'src.platform_info._read_proc_version',
return_value='5.15.123-microsoft-standard-WSL2',
):
self.assertEqual(get_wsl_version(), '2')
def test_wsl1_fallback(self) -> None:
with mock.patch('src.platform_info.sys.platform', 'linux'), \
mock.patch(
'src.platform_info._read_proc_version',
return_value='4.4.0-19041-Microsoft (Microsoft@Microsoft.com)',
):
self.assertEqual(get_wsl_version(), '1')
def test_non_linux(self) -> None:
with mock.patch('src.platform_info.sys.platform', 'darwin'):
self.assertIsNone(get_wsl_version())
def test_linux_no_microsoft_marker(self) -> None:
with mock.patch('src.platform_info.sys.platform', 'linux'), \
mock.patch(
'src.platform_info._read_proc_version',
return_value='Linux version 6.5.0 (gcc)',
):
self.assertIsNone(get_wsl_version())
class GetLinuxDistroInfoTest(unittest.TestCase):
def test_non_linux_returns_none(self) -> None:
with mock.patch('src.platform_info.sys.platform', 'darwin'):
self.assertIsNone(get_linux_distro_info())
def test_parses_id_and_version(self) -> None:
os_release = 'NAME="Ubuntu"\nID=ubuntu\nVERSION_ID="22.04"\n'
with mock.patch('src.platform_info.sys.platform', 'linux'), \
mock.patch(
'src.platform_info.Path.read_text', return_value=os_release,
):
info = get_linux_distro_info()
assert info is not None
self.assertEqual(info.linux_distro_id, 'ubuntu')
self.assertEqual(info.linux_distro_version, '22.04')
self.assertIsNotNone(info.linux_kernel)
def test_to_dict_skips_none(self) -> None:
info = LinuxDistroInfo(linux_distro_id='fedora')
self.assertEqual(info.to_dict(), {'linuxDistroId': 'fedora'})
class DetectVcsTest(unittest.TestCase):
def test_detects_git(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
(Path(tmp) / '.git').mkdir()
self.assertEqual(detect_vcs(tmp), ['git'])
def test_detects_multiple(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
(Path(tmp) / '.git').mkdir()
(Path(tmp) / '.hg').mkdir()
self.assertEqual(detect_vcs(tmp), ['git', 'mercurial'])
def test_perforce_via_env(self) -> None:
with tempfile.TemporaryDirectory() as tmp, \
mock.patch.dict(os.environ, {'P4PORT': '1666'}):
self.assertIn('perforce', detect_vcs(tmp))
def test_unreadable_directory_returns_empty(self) -> None:
with mock.patch.dict(os.environ, {}, clear=True):
self.assertEqual(detect_vcs('/nonexistent/path/abc/xyz'), [])
class GetSystemDirectoriesTest(unittest.TestCase):
def test_macos_defaults(self) -> None:
dirs = get_system_directories(
home_dir='/Users/x', platform='macos', env={},
)
self.assertEqual(dirs.HOME, '/Users/x')
self.assertEqual(dirs.DESKTOP, '/Users/x/Desktop')
self.assertEqual(dirs.DOCUMENTS, '/Users/x/Documents')
self.assertEqual(dirs.DOWNLOADS, '/Users/x/Downloads')
def test_windows_uses_userprofile(self) -> None:
dirs = get_system_directories(
home_dir='C:/Users/old',
platform='windows',
env={'USERPROFILE': 'C:/Users/new'},
)
# Path normalizes to forward slashes on linux test runs; just check
# USERPROFILE was used as the base, not home_dir.
self.assertIn('Users/new', dirs.DESKTOP.replace('\\', '/'))
self.assertIn('Users/new', dirs.DOWNLOADS.replace('\\', '/'))
# HOME stays as the explicit home_dir
self.assertEqual(dirs.HOME, 'C:/Users/old')
def test_linux_xdg_overrides(self) -> None:
dirs = get_system_directories(
home_dir='/home/u',
platform='linux',
env={'XDG_DOWNLOAD_DIR': '/data/dl'},
)
self.assertEqual(dirs.DOWNLOADS, '/data/dl')
self.assertEqual(dirs.DESKTOP, '/home/u/Desktop')
def test_wsl_xdg_overrides(self) -> None:
dirs = get_system_directories(
home_dir='/home/u',
platform='wsl',
env={'XDG_DESKTOP_DIR': '/mnt/c/Users/x/Desktop'},
)
self.assertEqual(dirs.DESKTOP, '/mnt/c/Users/x/Desktop')
def test_to_dict_round_trip(self) -> None:
dirs = SystemDirectories(
HOME='/h', DESKTOP='/h/D', DOCUMENTS='/h/Doc', DOWNLOADS='/h/Dn',
)
self.assertEqual(
dirs.to_dict(),
{'HOME': '/h', 'DESKTOP': '/h/D', 'DOCUMENTS': '/h/Doc', 'DOWNLOADS': '/h/Dn'},
)
if __name__ == '__main__':
unittest.main()
+101
View File
@@ -0,0 +1,101 @@
"""Tests for the local release-notes parser ported from utils/releaseNotes.ts."""
from __future__ import annotations
import tempfile
import unittest
from pathlib import Path
from src.release_notes import (
MAX_RELEASE_NOTES_SHOWN,
check_for_release_notes,
get_all_release_notes,
get_recent_release_notes,
parse_changelog,
read_local_changelog,
)
SAMPLE = (
'# Changelog\n\n'
'## 1.3.0 - 2026-04-15\n'
'- new shiny\n'
'- another bullet\n\n'
'## 1.2.0\n'
'- mid bullet\n\n'
'## 1.1.0\n'
'- oldest bullet\n'
)
class ParseChangelogTest(unittest.TestCase):
def test_returns_empty_for_blank(self) -> None:
self.assertEqual(parse_changelog(''), {})
def test_extracts_versions_and_bullets(self) -> None:
parsed = parse_changelog(SAMPLE)
self.assertEqual(set(parsed.keys()), {'1.3.0', '1.2.0', '1.1.0'})
self.assertEqual(parsed['1.3.0'], ['new shiny', 'another bullet'])
self.assertEqual(parsed['1.2.0'], ['mid bullet'])
def test_skips_versions_without_bullets(self) -> None:
parsed = parse_changelog('# X\n\n## 1.0.0\nplain text\n')
self.assertEqual(parsed, {})
class RecentNotesTest(unittest.TestCase):
def test_returns_only_newer_versions(self) -> None:
notes = get_recent_release_notes('1.3.0', '1.2.0', SAMPLE)
self.assertEqual(notes, ['new shiny', 'another bullet'])
def test_first_run_returns_all(self) -> None:
notes = get_recent_release_notes('1.3.0', None, SAMPLE)
self.assertEqual(notes[0], 'new shiny')
self.assertIn('oldest bullet', notes)
def test_no_new_when_at_or_below_previous(self) -> None:
self.assertEqual(get_recent_release_notes('1.1.0', '1.3.0', SAMPLE), [])
def test_caps_at_max_shown(self) -> None:
big_changelog = '# Changelog\n\n' + ''.join(
f'## 9.9.{i}\n- bullet {i}\n\n' for i in range(20)
)
notes = get_recent_release_notes('9.9.19', '0.0.1', big_changelog)
self.assertEqual(len(notes), MAX_RELEASE_NOTES_SHOWN)
class AllNotesTest(unittest.TestCase):
def test_sorted_oldest_first(self) -> None:
all_notes = get_all_release_notes(SAMPLE)
versions = [version for version, _ in all_notes]
self.assertEqual(versions, ['1.1.0', '1.2.0', '1.3.0'])
class ReadLocalChangelogTest(unittest.TestCase):
def test_reads_when_present(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
(Path(tmp) / 'CHANGELOG.md').write_text(SAMPLE, encoding='utf-8')
self.assertIn('## 1.3.0', read_local_changelog(Path(tmp)))
def test_returns_empty_when_missing(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
self.assertEqual(read_local_changelog(Path(tmp)), '')
class CheckForReleaseNotesTest(unittest.TestCase):
def test_signals_when_changelog_present_and_newer(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
(Path(tmp) / 'CHANGELOG.md').write_text(SAMPLE, encoding='utf-8')
payload = check_for_release_notes('1.3.0', '1.2.0', cwd=Path(tmp))
self.assertTrue(payload['hasReleaseNotes'])
self.assertEqual(payload['releaseNotes'][0], 'new shiny')
def test_silent_when_missing_changelog(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
payload = check_for_release_notes('1.3.0', None, cwd=Path(tmp))
self.assertFalse(payload['hasReleaseNotes'])
self.assertEqual(payload['releaseNotes'], [])
if __name__ == '__main__':
unittest.main()
+109
View File
@@ -0,0 +1,109 @@
"""Tests for remote/bridge slash commands ported from the npm source.
Covers /bridge (aliases /remote-control, /rc) and /remote-setup
(alias /web-setup).
"""
from __future__ import annotations
import json
import tempfile
import unittest
from pathlib import Path
from unittest import mock
from src.agent_runtime import LocalCodingAgent
from src.agent_types import AgentRuntimeConfig, ModelConfig
def _make_agent(tmp_dir: str) -> LocalCodingAgent:
return LocalCodingAgent(
model_config=ModelConfig(model='test-model'),
runtime_config=AgentRuntimeConfig(cwd=Path(tmp_dir)),
)
class BridgeCommandTest(unittest.TestCase):
def test_reports_unsupported_status(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/bridge').final_output
self.assertIn('not implemented', out.lower())
self.assertIn('No active local remote connection', out)
def test_remote_control_alias(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/remote-control').final_output
self.assertIn('Remote-control bridge', out)
def test_rc_alias(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/rc').final_output
self.assertIn('Remote-control bridge', out)
def test_named_lookup_misses_unknown_profile(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/bridge nope').final_output
self.assertIn('No matching remote profile for "nope"', out)
def test_named_lookup_matches_known_profile(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
(Path(tmp) / '.remote.json').write_text(json.dumps({
'profiles': [
{'name': 'edge', 'mode': 'ssh', 'target': 'user@edge.example'},
],
}), encoding='utf-8')
agent = _make_agent(tmp)
out = agent.run('/bridge edge').final_output
self.assertIn('Matched remote profile "edge"', out)
self.assertIn('user@edge.example', out)
class RemoteSetupCommandTest(unittest.TestCase):
def test_includes_web_url(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/remote-setup').final_output
self.assertIn('https://claude.ai/code', out)
self.assertIn('GitHub CLI', out)
def test_web_setup_alias(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/web-setup').final_output
self.assertIn('https://claude.ai/code', out)
def test_handles_missing_gh(self) -> None:
with mock.patch('shutil.which', return_value=None):
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/remote-setup').final_output
self.assertIn('not_installed', out)
self.assertIn('cli.github.com', out)
def test_handles_authenticated_gh(self) -> None:
fake = mock.Mock(returncode=0, stdout='Logged in to github.com as octo', stderr='')
with mock.patch('shutil.which', return_value='/usr/bin/gh'), \
mock.patch('subprocess.run', return_value=fake):
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/remote-setup').final_output
self.assertIn('authenticated', out)
self.assertIn('gh auth token', out)
def test_handles_unauthenticated_gh(self) -> None:
fake = mock.Mock(returncode=1, stdout='', stderr='You are not logged into any GitHub hosts')
with mock.patch('shutil.which', return_value='/usr/bin/gh'), \
mock.patch('subprocess.run', return_value=fake):
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/remote-setup').final_output
self.assertIn('not_authenticated', out)
self.assertIn('gh auth login', out)
if __name__ == '__main__':
unittest.main()
+110
View File
@@ -0,0 +1,110 @@
"""Tests for sandbox configuration types ported from sandboxTypes.ts."""
from __future__ import annotations
import unittest
from src.sandbox_types import (
SandboxFilesystemConfig,
SandboxNetworkConfig,
SandboxRipgrepConfig,
SandboxSettings,
)
class SandboxNetworkConfigTest(unittest.TestCase):
def test_round_trip(self) -> None:
raw = {
'allowedDomains': ['example.com'],
'allowManagedDomainsOnly': True,
'allowUnixSockets': ['/tmp/sock'],
'allowAllUnixSockets': False,
'allowLocalBinding': True,
'httpProxyPort': 8080,
'socksProxyPort': 1080,
}
parsed = SandboxNetworkConfig.from_dict(raw)
self.assertEqual(parsed.allowed_domains, ['example.com'])
self.assertEqual(parsed.http_proxy_port, 8080)
self.assertEqual(parsed.to_dict(), raw)
def test_strips_none(self) -> None:
parsed = SandboxNetworkConfig.from_dict({'allowedDomains': ['a']})
self.assertEqual(parsed.to_dict(), {'allowedDomains': ['a']})
def test_rejects_wrong_type(self) -> None:
with self.assertRaises(ValueError):
SandboxNetworkConfig.from_dict({'allowedDomains': 'nope'})
class SandboxFilesystemConfigTest(unittest.TestCase):
def test_round_trip(self) -> None:
raw = {
'allowWrite': ['/tmp'],
'denyWrite': ['/etc'],
'denyRead': ['/etc/secrets'],
'allowRead': ['/etc/secrets/public'],
'allowManagedReadPathsOnly': False,
}
parsed = SandboxFilesystemConfig.from_dict(raw)
self.assertEqual(parsed.allow_write, ['/tmp'])
self.assertEqual(parsed.to_dict(), raw)
class SandboxRipgrepConfigTest(unittest.TestCase):
def test_requires_command(self) -> None:
with self.assertRaises(ValueError):
SandboxRipgrepConfig.from_dict({'args': ['-i']})
def test_round_trip(self) -> None:
parsed = SandboxRipgrepConfig.from_dict({'command': 'rg', 'args': ['--no-ignore']})
self.assertEqual(parsed.command, 'rg')
self.assertEqual(parsed.to_dict(), {'command': 'rg', 'args': ['--no-ignore']})
class SandboxSettingsTest(unittest.TestCase):
def test_full_round_trip_preserves_passthrough(self) -> None:
raw = {
'enabled': True,
'failIfUnavailable': False,
'autoAllowBashIfSandboxed': True,
'allowUnsandboxedCommands': True,
'network': {'allowedDomains': ['x.com']},
'filesystem': {'allowWrite': ['/tmp']},
'ignoreViolations': {'NetworkViolation': ['y.com']},
'enableWeakerNestedSandbox': False,
'enableWeakerNetworkIsolation': False,
'excludedCommands': ['rm -rf /'],
'ripgrep': {'command': 'rg'},
'enabledPlatforms': ['macos'],
'somethingFuture': 42,
}
parsed = SandboxSettings.from_dict(raw)
self.assertTrue(parsed.enabled)
self.assertEqual(parsed.network.allowed_domains, ['x.com'])
self.assertEqual(parsed.ignore_violations, {'NetworkViolation': ['y.com']})
self.assertEqual(parsed.extra['enabledPlatforms'], ['macos'])
self.assertEqual(parsed.extra['somethingFuture'], 42)
back = parsed.to_dict()
self.assertEqual(back['enabled'], True)
self.assertEqual(back['enabledPlatforms'], ['macos'])
self.assertEqual(back['somethingFuture'], 42)
def test_empty_returns_defaults(self) -> None:
parsed = SandboxSettings.from_dict({})
self.assertIsNone(parsed.enabled)
self.assertIsNone(parsed.network)
self.assertEqual(parsed.to_dict(), {})
def test_rejects_non_mapping(self) -> None:
with self.assertRaises(ValueError):
SandboxSettings.from_dict('nope') # type: ignore[arg-type]
def test_ignore_violations_must_be_mapping(self) -> None:
with self.assertRaises(ValueError):
SandboxSettings.from_dict({'ignoreViolations': ['nope']})
if __name__ == '__main__':
unittest.main()
+175
View File
@@ -0,0 +1,175 @@
"""Tests for SDK core types ported from entrypoints/sdk/coreTypes.ts."""
from __future__ import annotations
import unittest
from src.sdk_core_types import (
API_KEY_SOURCES,
CONFIG_SCOPES,
EXIT_REASONS,
HOOK_EVENTS,
SDK_BETAS,
JsonSchemaOutputFormat,
McpClaudeAIProxyServerConfig,
McpHttpServerConfig,
McpSdkServerConfig,
McpSSEServerConfig,
McpStdioServerConfig,
ModelUsage,
ThinkingAdaptive,
ThinkingDisabled,
ThinkingEnabled,
mcp_server_config_from_dict,
thinking_config_from_dict,
)
class HookEventsTest(unittest.TestCase):
def test_includes_known_events(self) -> None:
for required in (
'PreToolUse',
'PostToolUse',
'UserPromptSubmit',
'SessionStart',
'PreCompact',
'WorktreeCreate',
'CwdChanged',
'FileChanged',
):
self.assertIn(required, HOOK_EVENTS)
def test_no_duplicates(self) -> None:
self.assertEqual(len(HOOK_EVENTS), len(set(HOOK_EVENTS)))
class ExitReasonsTest(unittest.TestCase):
def test_known_exit_reasons(self) -> None:
for required in (
'clear', 'resume', 'logout', 'prompt_input_exit',
'other', 'bypass_permissions_disabled',
):
self.assertIn(required, EXIT_REASONS)
class EnumLiteralsTest(unittest.TestCase):
def test_api_key_sources(self) -> None:
self.assertEqual(set(API_KEY_SOURCES), {'user', 'project', 'org', 'temporary', 'oauth'})
def test_config_scopes(self) -> None:
self.assertEqual(set(CONFIG_SCOPES), {'local', 'user', 'project'})
def test_sdk_betas_known(self) -> None:
self.assertIn('context-1m-2025-08-07', SDK_BETAS)
class ModelUsageTest(unittest.TestCase):
def test_round_trip(self) -> None:
raw = {
'inputTokens': 100, 'outputTokens': 50,
'cacheReadInputTokens': 10, 'cacheCreationInputTokens': 5,
'webSearchRequests': 0, 'costUSD': 0.001,
'contextWindow': 200000, 'maxOutputTokens': 8192,
}
usage = ModelUsage.from_dict(raw)
self.assertEqual(usage.input_tokens, 100)
self.assertEqual(usage.cost_usd, 0.001)
self.assertEqual(usage.to_dict(), raw)
class ThinkingConfigTest(unittest.TestCase):
def test_adaptive(self) -> None:
cfg = thinking_config_from_dict({'type': 'adaptive'})
self.assertIsInstance(cfg, ThinkingAdaptive)
self.assertEqual(cfg.to_dict(), {'type': 'adaptive'})
def test_enabled_with_budget(self) -> None:
cfg = thinking_config_from_dict({'type': 'enabled', 'budgetTokens': 1024})
self.assertIsInstance(cfg, ThinkingEnabled)
self.assertEqual(cfg.to_dict(), {'type': 'enabled', 'budgetTokens': 1024})
def test_enabled_without_budget(self) -> None:
cfg = thinking_config_from_dict({'type': 'enabled'})
self.assertEqual(cfg.to_dict(), {'type': 'enabled'})
def test_disabled(self) -> None:
cfg = thinking_config_from_dict({'type': 'disabled'})
self.assertIsInstance(cfg, ThinkingDisabled)
def test_unknown_type_raises(self) -> None:
with self.assertRaises(ValueError):
thinking_config_from_dict({'type': 'something-else'})
class McpServerConfigTest(unittest.TestCase):
def test_stdio_default_type(self) -> None:
cfg = mcp_server_config_from_dict({'command': 'npx'})
self.assertIsInstance(cfg, McpStdioServerConfig)
self.assertEqual(cfg.command, 'npx')
def test_stdio_with_args_env(self) -> None:
cfg = mcp_server_config_from_dict({
'type': 'stdio',
'command': 'node',
'args': ['./mcp.js'],
'env': {'TOKEN': 'xyz'},
})
out = cfg.to_dict()
self.assertEqual(out['command'], 'node')
self.assertEqual(out['args'], ['./mcp.js'])
self.assertEqual(out['env'], {'TOKEN': 'xyz'})
def test_sse(self) -> None:
cfg = mcp_server_config_from_dict({'type': 'sse', 'url': 'https://x.example'})
self.assertIsInstance(cfg, McpSSEServerConfig)
def test_http_with_headers(self) -> None:
cfg = mcp_server_config_from_dict({
'type': 'http',
'url': 'https://x.example',
'headers': {'Auth': 'Bearer 1'},
})
self.assertIsInstance(cfg, McpHttpServerConfig)
self.assertEqual(cfg.headers, {'Auth': 'Bearer 1'})
def test_sdk(self) -> None:
cfg = mcp_server_config_from_dict({'type': 'sdk', 'name': 'my-sdk'})
self.assertIsInstance(cfg, McpSdkServerConfig)
def test_claudeai_proxy(self) -> None:
cfg = mcp_server_config_from_dict({
'type': 'claudeai-proxy',
'url': 'https://claude.ai/p',
'id': 'abc',
})
self.assertIsInstance(cfg, McpClaudeAIProxyServerConfig)
def test_unknown_type_raises(self) -> None:
with self.assertRaises(ValueError):
mcp_server_config_from_dict({'type': 'mystery'})
def test_stdio_requires_command(self) -> None:
with self.assertRaises(ValueError):
mcp_server_config_from_dict({'type': 'stdio'})
class JsonSchemaOutputFormatTest(unittest.TestCase):
def test_round_trip(self) -> None:
fmt = JsonSchemaOutputFormat.from_dict({
'type': 'json_schema',
'schema': {'type': 'object', 'properties': {}},
})
self.assertEqual(fmt.schema['type'], 'object')
self.assertEqual(fmt.to_dict()['type'], 'json_schema')
def test_rejects_wrong_type(self) -> None:
with self.assertRaises(ValueError):
JsonSchemaOutputFormat.from_dict({'type': 'text', 'schema': {}})
def test_requires_schema_mapping(self) -> None:
with self.assertRaises(ValueError):
JsonSchemaOutputFormat.from_dict({'type': 'json_schema'})
if __name__ == '__main__':
unittest.main()
+63
View File
@@ -0,0 +1,63 @@
"""Tests for ``src/session_env_vars.py``."""
from __future__ import annotations
import unittest
from src.session_env_vars import (
clear_session_env_vars,
delete_session_env_var,
get_session_env_vars,
set_session_env_var,
)
class SessionEnvVarsTest(unittest.TestCase):
def setUp(self) -> None:
clear_session_env_vars()
def tearDown(self) -> None:
clear_session_env_vars()
def test_starts_empty(self) -> None:
self.assertEqual(dict(get_session_env_vars()), {})
def test_set_and_get(self) -> None:
set_session_env_var('FOO', 'bar')
self.assertEqual(get_session_env_vars()['FOO'], 'bar')
def test_set_overwrites_existing(self) -> None:
set_session_env_var('FOO', 'one')
set_session_env_var('FOO', 'two')
self.assertEqual(get_session_env_vars()['FOO'], 'two')
def test_delete_removes(self) -> None:
set_session_env_var('FOO', 'bar')
delete_session_env_var('FOO')
self.assertNotIn('FOO', get_session_env_vars())
def test_delete_missing_is_noop(self) -> None:
delete_session_env_var('NEVER_SET')
self.assertEqual(dict(get_session_env_vars()), {})
def test_clear_drops_everything(self) -> None:
set_session_env_var('A', '1')
set_session_env_var('B', '2')
clear_session_env_vars()
self.assertEqual(dict(get_session_env_vars()), {})
def test_returned_mapping_is_read_only(self) -> None:
set_session_env_var('FOO', 'bar')
view = get_session_env_vars()
with self.assertRaises(TypeError):
view['FOO'] = 'mutated' # type: ignore[index]
def test_view_reflects_subsequent_mutations(self) -> None:
view = get_session_env_vars()
self.assertNotIn('FOO', view)
set_session_env_var('FOO', 'bar')
self.assertEqual(view['FOO'], 'bar')
if __name__ == '__main__':
unittest.main()
+179
View File
@@ -0,0 +1,179 @@
"""Tests for settings-touching slash commands ported from the npm source.
Covers /theme, /voice, /sandbox-toggle (alias /sandbox), /keybindings, /btw.
"""
from __future__ import annotations
import json
import tempfile
import unittest
from pathlib import Path
from src.agent_runtime import LocalCodingAgent
from src.agent_types import AgentRuntimeConfig, ModelConfig
def _make_agent(tmp_dir: str) -> LocalCodingAgent:
return LocalCodingAgent(
model_config=ModelConfig(model='test-model'),
runtime_config=AgentRuntimeConfig(cwd=Path(tmp_dir)),
)
def _local_settings(tmp_dir: str) -> dict:
path = Path(tmp_dir) / '.claude' / 'settings.local.json'
if not path.exists():
return {}
return json.loads(path.read_text(encoding='utf-8'))
class ThemeCommandTest(unittest.TestCase):
def test_lists_themes_when_no_arg(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/theme').final_output
self.assertIn('Available themes', out)
self.assertIn('light', out)
self.assertIn('dark', out)
self.assertIn('Usage: /theme <name>', out)
def test_rejects_unknown_theme(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/theme neon').final_output
self.assertIn('Unknown theme', out)
self.assertIn('neon', out)
def test_sets_theme_and_persists(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/theme dark').final_output
self.assertIn('Theme set to dark', out)
settings = _local_settings(tmp)
self.assertEqual(settings.get('theme'), 'dark')
def test_marks_current_theme(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
agent.run('/theme dark')
out = agent.run('/theme').final_output
self.assertIn('dark (current)', out)
class VoiceCommandTest(unittest.TestCase):
def test_toggle_enables_when_unset(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/voice').final_output
self.assertIn('Voice mode enabled', out)
self.assertEqual(_local_settings(tmp).get('voiceEnabled'), True)
def test_toggle_disables_when_enabled(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
agent.run('/voice on')
out = agent.run('/voice').final_output
self.assertIn('Voice mode disabled', out)
self.assertEqual(_local_settings(tmp).get('voiceEnabled'), False)
def test_explicit_on_off(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
self.assertIn('enabled', agent.run('/voice on').final_output)
self.assertIn('disabled', agent.run('/voice off').final_output)
def test_rejects_unknown_arg(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/voice maybe').final_output
self.assertIn('Usage', out)
class SandboxToggleCommandTest(unittest.TestCase):
def test_status_with_no_args(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/sandbox-toggle').final_output
self.assertIn('Sandbox:', out)
self.assertIn('Excluded commands', out)
self.assertIn('Usage:', out)
def test_alias_sandbox(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/sandbox').final_output
self.assertIn('Sandbox:', out)
def test_exclude_appends_pattern(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/sandbox-toggle exclude "npm run test:*"').final_output
self.assertIn('Added "npm run test:*"', out)
settings = _local_settings(tmp)
excluded = settings.get('sandbox', {}).get('excludedCommands', [])
self.assertIn('npm run test:*', excluded)
def test_exclude_dedupes(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
agent.run('/sandbox-toggle exclude "rm -rf /"')
out = agent.run('/sandbox-toggle exclude "rm -rf /"').final_output
self.assertIn('already in', out)
def test_exclude_requires_pattern(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/sandbox-toggle exclude').final_output
self.assertIn('Usage', out)
def test_unknown_subcommand(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/sandbox-toggle wat').final_output
self.assertIn('Unknown subcommand', out)
class KeybindingsCommandTest(unittest.TestCase):
def test_creates_template_when_missing(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
out = agent.run('/keybindings').final_output
path = Path(tmp) / '.claude' / 'keybindings.json'
self.assertTrue(path.exists())
self.assertIn('Created', out)
self.assertIn(str(path), out)
# Template is valid JSON-ish (has braces); strict json.loads would
# choke on the "//" comment, so just sanity-check structure.
text = path.read_text(encoding='utf-8')
self.assertIn('"bindings"', text)
def test_reports_existing_file(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
agent.run('/keybindings')
out = agent.run('/keybindings').final_output
self.assertIn('Found', out)
class BtwCommandTest(unittest.TestCase):
def test_no_question_shows_usage(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
result = agent.run('/btw').final_output
self.assertIn('Usage: /btw', result)
def test_question_returns_prompt_result(self) -> None:
from src.agent_slash_commands import preprocess_slash_command
with tempfile.TemporaryDirectory() as tmp:
agent = _make_agent(tmp)
result = preprocess_slash_command(agent, '/btw what does this codebase do?')
self.assertTrue(result.handled)
self.assertTrue(result.should_query)
self.assertIn('side question', (result.prompt or '').lower())
self.assertIn('what does this codebase do?', result.prompt or '')
if __name__ == '__main__':
unittest.main()
+64
View File
@@ -0,0 +1,64 @@
"""Tests for setup-time runtime checks ported from setup.ts."""
from __future__ import annotations
import sys
import tempfile
import unittest
from pathlib import Path
from unittest import mock
from src.setup import (
MIN_PYTHON_VERSION,
SetupReport,
check_runtime_requirements,
run_setup,
)
class RuntimeRequirementCheckTest(unittest.TestCase):
def test_python_check_passes_on_current_runtime(self) -> None:
checks = {check.name: check for check in check_runtime_requirements()}
self.assertTrue(checks['python_version'].ok)
self.assertGreaterEqual(sys.version_info[:2], MIN_PYTHON_VERSION)
def test_python_check_fails_when_below_minimum(self) -> None:
# Force a lower version_info via patching to verify the failure branch.
fake_version = mock.Mock()
fake_version.__getitem__ = lambda self, idx: (3, 8)[idx] if isinstance(idx, int) else (3, 8)[idx]
with mock.patch('src.setup.sys') as fake_sys:
fake_sys.version_info = (3, 8, 0)
checks = {check.name: check for check in check_runtime_requirements()}
self.assertFalse(checks['python_version'].ok)
self.assertIn('below required', checks['python_version'].detail)
def test_includes_platform_and_implementation(self) -> None:
names = {check.name for check in check_runtime_requirements()}
self.assertIn('python_implementation', names)
self.assertIn('platform', names)
class SetupReportTest(unittest.TestCase):
def test_run_setup_reports_runtime_and_release_notes(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
(Path(tmp) / 'CHANGELOG.md').write_text(
'# Changelog\n\n## 9.9.9\n- big release\n', encoding='utf-8',
)
report = run_setup(cwd=Path(tmp), trusted=True, last_seen_version='0.0.1')
self.assertIsInstance(report, SetupReport)
self.assertGreaterEqual(len(report.runtime_checks), 3)
self.assertIn('big release', report.release_notes)
markdown = report.as_markdown()
self.assertIn('Runtime checks', markdown)
self.assertIn('Release notes', markdown)
self.assertFalse(report.has_blocking_issues())
def test_no_release_notes_when_changelog_missing(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
report = run_setup(cwd=Path(tmp), trusted=True)
self.assertEqual(report.release_notes, ())
self.assertNotIn('Release notes', report.as_markdown())
if __name__ == '__main__':
unittest.main()
+138
View File
@@ -0,0 +1,138 @@
"""Tests for the bundled small utilities ported in ``src/small_utils.py``."""
from __future__ import annotations
import unittest
from src.small_utils import (
count,
create_agent_id,
difference,
escape_xml,
escape_xml_attr,
every_in,
intersects,
intersperse,
object_group_by,
union,
uniq,
validate_uuid,
)
class IntersperseTest(unittest.TestCase):
def test_empty(self) -> None:
self.assertEqual(intersperse([], lambda i: ','), [])
def test_single_item_no_separator(self) -> None:
self.assertEqual(intersperse(['a'], lambda i: ','), ['a'])
def test_separator_receives_index_starting_at_one(self) -> None:
seen: list[int] = []
def sep(i: int) -> str:
seen.append(i)
return f'-{i}-'
out = intersperse(['a', 'b', 'c'], sep)
self.assertEqual(out, ['a', '-1-', 'b', '-2-', 'c'])
self.assertEqual(seen, [1, 2])
class CountTest(unittest.TestCase):
def test_counts_truthy(self) -> None:
self.assertEqual(count([1, 2, 3, 4], lambda x: x % 2 == 0), 2)
def test_predicate_returning_objects_treated_as_truthy(self) -> None:
self.assertEqual(count(['', 'x', 'y'], lambda s: s), 2)
class UniqTest(unittest.TestCase):
def test_preserves_first_seen_order(self) -> None:
self.assertEqual(uniq([3, 1, 2, 1, 3, 4]), [3, 1, 2, 4])
class ObjectGroupByTest(unittest.TestCase):
def test_groups_by_key(self) -> None:
out = object_group_by(['apple', 'banana', 'avocado'], lambda s, _i: s[0])
self.assertEqual(out, {'a': ['apple', 'avocado'], 'b': ['banana']})
def test_passes_index_to_selector(self) -> None:
out = object_group_by(
['x', 'y', 'z'], lambda _s, i: 'even' if i % 2 == 0 else 'odd',
)
self.assertEqual(out, {'even': ['x', 'z'], 'odd': ['y']})
class SetOpsTest(unittest.TestCase):
def test_difference(self) -> None:
self.assertEqual(difference({1, 2, 3}, {2}), {1, 3})
def test_intersects_true(self) -> None:
self.assertTrue(intersects({1, 2}, {2, 3}))
def test_intersects_false(self) -> None:
self.assertFalse(intersects({1, 2}, {3, 4}))
def test_intersects_empty_short_circuits(self) -> None:
self.assertFalse(intersects(set(), {1}))
self.assertFalse(intersects({1}, set()))
def test_every_in(self) -> None:
self.assertTrue(every_in({1, 2}, {1, 2, 3}))
self.assertFalse(every_in({1, 4}, {1, 2, 3}))
self.assertTrue(every_in(set(), {1, 2}))
def test_union(self) -> None:
self.assertEqual(union({1, 2}, {2, 3}), {1, 2, 3})
class XmlEscapeTest(unittest.TestCase):
def test_escape_xml(self) -> None:
self.assertEqual(
escape_xml('a & b < c > d'), 'a &amp; b &lt; c &gt; d',
)
def test_escape_xml_amp_first_no_double_escape(self) -> None:
self.assertEqual(escape_xml('<&>'), '&lt;&amp;&gt;')
def test_escape_xml_attr_includes_quotes(self) -> None:
self.assertEqual(
escape_xml_attr('he said "hi" & \'bye\''),
'he said &quot;hi&quot; &amp; &apos;bye&apos;',
)
class ValidateUuidTest(unittest.TestCase):
def test_valid_lowercase(self) -> None:
u = '12345678-1234-1234-1234-123456789012'
self.assertEqual(validate_uuid(u), u)
def test_valid_uppercase(self) -> None:
u = 'ABCDEF12-1234-5678-90AB-CDEF12345678'
self.assertEqual(validate_uuid(u), u)
def test_invalid_format(self) -> None:
self.assertIsNone(validate_uuid('not-a-uuid'))
def test_non_string_returns_none(self) -> None:
self.assertIsNone(validate_uuid(123))
self.assertIsNone(validate_uuid(None))
class CreateAgentIdTest(unittest.TestCase):
def test_no_label(self) -> None:
agent_id = create_agent_id()
self.assertRegex(agent_id, r'^a[0-9a-f]{16}$')
def test_with_label(self) -> None:
agent_id = create_agent_id('compact')
self.assertRegex(agent_id, r'^acompact-[0-9a-f]{16}$')
def test_unique_across_calls(self) -> None:
ids = {create_agent_id() for _ in range(50)}
self.assertEqual(len(ids), 50)
if __name__ == '__main__':
unittest.main()