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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 23 additions & 3 deletions client/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -891,17 +891,18 @@ class _RunApprovalPolicy:
output: OutputFacade
watch: _ApprovalWatchState
auto_approve_tools: set[str]
auto_approve_all: bool = False
seen_disallowed: set[str] = dataclasses.field(default_factory=set)
attempted: set[str] = dataclasses.field(default_factory=set)

async def on_pending(self, request_id: str, tool_name: str, capability_id: str) -> None:
normalized_request = request_id.strip() if request_id.strip() else "?"
normalized_tool = tool_name.strip() if tool_name.strip() else "?"
self.watch.mark_pending(normalized_request)
if not self.auto_approve_tools:
if not self.auto_approve_all and not self.auto_approve_tools:
return

if normalized_tool not in self.auto_approve_tools:
if not self.auto_approve_all and normalized_tool not in self.auto_approve_tools:
if normalized_request in self.seen_disallowed:
return
self.seen_disallowed.add(normalized_request)
Expand Down Expand Up @@ -1606,6 +1607,12 @@ def _build_parser() -> argparse.ArgumentParser:
default=None,
help="extra tool name eligible for auto-approve (repeatable)",
)
run.add_argument(
"--full-auto",
action="store_true",
help="fully autonomous mode: auto-approve ALL tool invocations and "
"auto-respond to ask_user questions (no human interaction required)",
)
run.add_argument(
"--headless",
action="store_true",
Expand Down Expand Up @@ -1815,6 +1822,16 @@ async def main(argv: list[str] | None = None) -> int:
output.emit_data(_serialize(_list_session_payload(session_store=session_store)))
return 0

# --full-auto: inject AutoUserInputHandler so ask_user never blocks.
full_auto = getattr(args, "full_auto", False)
if full_auto:
from dataclasses import replace as _replace

from dare_framework.tool._internal.tools.ask_user import AutoUserInputHandler

options = _replace(options, user_input_handler=AutoUserInputHandler())
output.info("full-auto mode: ask_user will auto-respond without human input")

try:
runtime = await bootstrap_runtime(options)
except Exception as exc: # noqa: BLE001
Expand Down Expand Up @@ -1905,14 +1922,17 @@ async def main(argv: list[str] | None = None) -> int:
for tool in args.auto_approve_tool
if isinstance(tool, str) and tool.strip()
)
if auto_tools:
if full_auto:
output.info("full-auto mode: all tool approvals will be auto-granted")
elif auto_tools:
output.info(f"run auto-approve enabled for tools={','.join(sorted(auto_tools))}")
approval_watch = _ApprovalWatchState()
approval_policy = _RunApprovalPolicy(
action_client=action_client,
output=output,
watch=approval_watch,
auto_approve_tools=auto_tools,
auto_approve_all=full_auto,
)

async def _handle_run_approval_pending(
Expand Down
4 changes: 4 additions & 0 deletions client/runtime/bootstrap.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
from dare_framework.model.types import Prompt
from dare_framework.plan import DefaultPlanner, DefaultRemediator
from dare_framework.tool._internal.tools import ReadFileTool, RunCommandTool, SearchCodeTool, WriteFileTool
from dare_framework.tool._internal.tools.ask_user import IUserInputHandler
from dare_framework.transport import AgentChannel, DirectClientChannel


Expand All @@ -33,6 +34,7 @@ class RuntimeOptions:
system_prompt_mode: str | None = None
system_prompt_text: str | None = None
system_prompt_file: str | None = None
user_input_handler: IUserInputHandler | None = None


@dataclass
Expand Down Expand Up @@ -240,6 +242,8 @@ async def bootstrap_runtime(options: RuntimeOptions) -> ClientRuntime:
.with_planner(DefaultPlanner(model, verbose=False))
.with_remediator(DefaultRemediator(model, verbose=False))
)
if options.user_input_handler is not None:
builder = builder.with_user_input_handler(options.user_input_handler)
prompt_override = _resolve_system_prompt_override(config=config, model=model)
if prompt_override is not None:
builder = builder.with_prompt(prompt_override)
Expand Down
1 change: 1 addition & 0 deletions dare_framework/tool/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,7 @@
"ToolManager",
# Built-in ask_user
"AskUserTool",
"AutoUserInputHandler",
"CLIUserInputHandler",
"IUserInputHandler",
]
4 changes: 4 additions & 0 deletions dare_framework/tool/_exports.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,10 @@
"dare_framework.tool._internal.tools.ask_user",
"CLIUserInputHandler",
),
"AutoUserInputHandler": (
"dare_framework.tool._internal.tools.ask_user",
"AutoUserInputHandler",
),
"IUserInputHandler": (
"dare_framework.tool._internal.tools.ask_user",
"IUserInputHandler",
Expand Down
5 changes: 5 additions & 0 deletions dare_framework/tool/_internal/_exports.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@

__all__ = [
"AskUserTool",
"AutoUserInputHandler",
"CLIUserInputHandler",
"IUserInputHandler",
"Checkpoint",
Expand All @@ -28,6 +29,10 @@
"dare_framework.tool._internal.tools.ask_user",
"CLIUserInputHandler",
),
"AutoUserInputHandler": (
"dare_framework.tool._internal.tools.ask_user",
"AutoUserInputHandler",
),
"IUserInputHandler": (
"dare_framework.tool._internal.tools.ask_user",
"IUserInputHandler",
Expand Down
2 changes: 2 additions & 0 deletions dare_framework/tool/_internal/tools/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

from dare_framework.tool._internal.tools.ask_user import (
AskUserTool,
AutoUserInputHandler,
CLIUserInputHandler,
IUserInputHandler,
)
Expand All @@ -19,6 +20,7 @@

__all__ = [
"AskUserTool",
"AutoUserInputHandler",
"CLIUserInputHandler",
"IUserInputHandler",
"EchoTool",
Expand Down
27 changes: 27 additions & 0 deletions dare_framework/tool/_internal/tools/ask_user.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,32 @@ async def handle(self, questions: list[dict[str, Any]]) -> dict[str, str]:
# ---------------------------------------------------------------------------


class AutoUserInputHandler(IUserInputHandler):
"""Handler that responds automatically without blocking on user input.

Used in fully autonomous execution modes (e.g. ``dare run --full-auto``)
where the agent should never block waiting for human interaction.
When options are available, the first option is selected; otherwise a
configurable default response is returned.
"""

DEFAULT_RESPONSE = "Proceed with your best judgment."

def __init__(self, default_response: str | None = None) -> None:
self._default_response = default_response or self.DEFAULT_RESPONSE

async def handle(self, questions: list[dict[str, Any]]) -> dict[str, str]:
answers: dict[str, str] = {}
for q in questions:
question_text = q.get("question", "")
options = q.get("options", [])
if options:
answers[question_text] = options[0].get("label", self._default_response)
else:
answers[question_text] = self._default_response
return answers


class CLIUserInputHandler(IUserInputHandler):
"""Simple stdin/stdout handler for command-line applications."""

Expand Down Expand Up @@ -308,6 +334,7 @@ async def execute(

__all__ = [
"AskUserTool",
"AutoUserInputHandler",
"CLIUserInputHandler",
"IUserInputHandler",
]
95 changes: 95 additions & 0 deletions tests/unit/test_client_cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -1593,6 +1593,101 @@ def test_default_auto_approve_tools_exclude_write_file() -> None:
assert runtime_bootstrap.WriteFileTool().name not in client_main.DEFAULT_AUTO_APPROVE_TOOLS


def test_full_auto_flag_accepted_by_run_parser() -> None:
client_main = importlib.import_module("client.main")
parser = client_main._build_parser()
args = parser.parse_args(["run", "--task", "hello", "--full-auto"])
assert args.full_auto is True


def test_full_auto_flag_defaults_to_false() -> None:
client_main = importlib.import_module("client.main")
parser = client_main._build_parser()
args = parser.parse_args(["run", "--task", "hello"])
assert args.full_auto is False


def test_auto_user_input_handler_picks_first_option() -> None:
from dare_framework.tool._internal.tools.ask_user import AutoUserInputHandler

handler = AutoUserInputHandler()
questions = [
{
"question": "Which approach?",
"header": "Approach",
"options": [
{"label": "Option A", "description": "First"},
{"label": "Option B", "description": "Second"},
],
}
]
answers = asyncio.run(handler.handle(questions))
assert answers["Which approach?"] == "Option A"


def test_auto_user_input_handler_uses_default_when_no_options() -> None:
from dare_framework.tool._internal.tools.ask_user import AutoUserInputHandler

handler = AutoUserInputHandler()
questions = [{"question": "What now?", "header": "Q", "options": []}]
answers = asyncio.run(handler.handle(questions))
assert answers["What now?"] == AutoUserInputHandler.DEFAULT_RESPONSE


def test_auto_user_input_handler_custom_default() -> None:
from dare_framework.tool._internal.tools.ask_user import AutoUserInputHandler

handler = AutoUserInputHandler(default_response="YOLO")
questions = [{"question": "What now?", "header": "Q", "options": []}]
answers = asyncio.run(handler.handle(questions))
assert answers["What now?"] == "YOLO"


def test_run_approval_policy_auto_approve_all() -> None:
"""When auto_approve_all is True, _RunApprovalPolicy approves any tool."""
client_main = importlib.import_module("client.main")

# Build a minimal mock for action_client, output, and watch
class _FakeActionClient:
def __init__(self):
self.invocations = []

async def invoke_action(self, action, **kwargs):
self.invocations.append((action, kwargs))

class _FakeOutput:
def __init__(self):
self.messages = []

def info(self, msg):
self.messages.append(msg)

def ok(self, msg):
self.messages.append(msg)

def display(self, msg, level="info"):
self.messages.append(msg)

fake_client = _FakeActionClient()
fake_output = _FakeOutput()
watch = client_main._ApprovalWatchState()

policy = client_main._RunApprovalPolicy(
action_client=fake_client,
output=fake_output,
watch=watch,
auto_approve_tools=set(),
auto_approve_all=True,
)

# Even an unknown tool should be auto-approved
asyncio.run(
policy.on_pending("req-1", "dangerous_tool", "dangerous_tool")
)
assert len(fake_client.invocations) == 1
assert fake_client.invocations[0][1]["request_id"] == "req-1"


def test_cli_raises_system_exit(monkeypatch: pytest.MonkeyPatch) -> None:
client_main = importlib.import_module("client.main")
monkeypatch.setattr(client_main, "sync_main", lambda argv=None: 5)
Expand Down
Loading