diff --git a/dare_framework/agent/base_agent.py b/dare_framework/agent/base_agent.py index 636548a8..f4cc12fd 100644 --- a/dare_framework/agent/base_agent.py +++ b/dare_framework/agent/base_agent.py @@ -15,7 +15,12 @@ from dare_framework.agent.status import AgentStatus from dare_framework.plan.types import RunResult, Task from dare_framework.transport.interaction.payloads import build_error_payload, build_success_payload -from dare_framework.transport.types import EnvelopeKind, TransportEnvelope, new_envelope_id +from dare_framework.transport.types import ( + EnvelopeKind, + TransportEnvelope, + TransportEventType, + new_envelope_id, +) if TYPE_CHECKING: from dare_framework.agent.builder import DareAgentBuilder, ReactAgentBuilder, SimpleChatAgentBuilder @@ -274,6 +279,7 @@ async def _send_transport_result( envelope = TransportEnvelope( id=new_envelope_id(), reply_to=reply_to, + event_type=TransportEventType.RESULT.value, payload={ **build_success_payload( kind="message", @@ -378,6 +384,7 @@ async def _send_transport_error( id=new_envelope_id(), kind=EnvelopeKind.MESSAGE, reply_to=envelope_id, + event_type=TransportEventType.ERROR.value, payload=build_error_payload( kind="message", target=target, diff --git a/dare_framework/agent/dare_agent.py b/dare_framework/agent/dare_agent.py index a28f6ba7..6a3cf9e3 100644 --- a/dare_framework/agent/dare_agent.py +++ b/dare_framework/agent/dare_agent.py @@ -47,16 +47,11 @@ ValidatedPlan, VerifyResult, ) -from dare_framework.tool._internal.control.approval_manager import ( - ApprovalDecision, - ApprovalEvaluationStatus, +from dare_framework.tool._internal.governed_tool_gateway import ( + ApprovalInvokeContext, + GovernedToolGateway, ) from dare_framework.tool.types import CapabilityKind -from dare_framework.transport.interaction.payloads import ( - build_approval_pending_payload, - build_approval_resolved_payload, -) -from dare_framework.transport.types import TransportEnvelope, new_envelope_id @dataclass @@ -169,6 +164,11 @@ def __init__( # Tool components self._tool_gateway = tool_gateway + self._governed_tool_gateway = GovernedToolGateway( + tool_gateway, + approval_manager=approval_manager, + logger=self._logger, + ) self._mcp_manager = mcp_manager self._exec_ctl = execution_control self._approval_manager = approval_manager @@ -207,7 +207,6 @@ def __init__( # Runtime state (set during execution) self._session_state: SessionState | None = None - self._active_transport: AgentChannel | None = None self._conversation_id: str | None = None self._token_usage: dict[str, int] = {"input_tokens": 0, "output_tokens": 0, "cached_tokens": 0} @@ -285,8 +284,6 @@ async def execute( transport: AgentChannel | None = None, ) -> RunResult: """Execute a task with automatic mode selection.""" - previous_transport = self._active_transport - self._active_transport = transport previous_conversation_id = self._conversation_id if isinstance(task, Task): task_obj = task @@ -314,13 +311,12 @@ async def execute( result: RunResult | None = None error: Exception | None = None try: - result = await self._run_session_loop(task_obj) + result = await self._run_session_loop(task_obj, transport=transport) return self._with_normalized_output_text(result) except Exception as exc: error = exc raise finally: - self._active_transport = previous_transport self._conversation_id = previous_conversation_id duration_ms = (time.perf_counter() - start_time) * 1000.0 errors: list[str] = [] @@ -348,7 +344,12 @@ async def execute( # Session Loop (Layer 1) # ========================================================================= - async def _run_session_loop(self, task: Task) -> RunResult: + async def _run_session_loop( + self, + task: Task, + *, + transport: AgentChannel | None = None, + ) -> RunResult: """Run the session loop - top-level task lifecycle.""" # Initialize session state if self._session_state is None: @@ -417,7 +418,7 @@ async def _run_session_loop(self, task: Task) -> RunResult: if self._exec_ctl is not None: self._poll_or_raise() - result = await self._run_milestone_loop(milestone) + result = await self._run_milestone_loop(milestone, transport=transport) milestone_results.append(result) self._log(f"Milestone {idx + 1} result: success={result.success}") @@ -455,7 +456,12 @@ async def _run_session_loop(self, task: Task) -> RunResult: # Milestone Loop (Layer 2) # ========================================================================= - async def _run_milestone_loop(self, milestone: Milestone) -> MilestoneResult: + async def _run_milestone_loop( + self, + milestone: Milestone, + *, + transport: AgentChannel | None = None, + ) -> MilestoneResult: """Run the milestone loop - sub-goal tracking.""" milestone_start = time.perf_counter() await self._emit_hook(HookPhase.BEFORE_MILESTONE, { @@ -488,7 +494,7 @@ async def _run_milestone_loop(self, milestone: Milestone) -> MilestoneResult: # Run execute loop self._log("Running execute loop...") - execute_result = await self._run_execute_loop(validated_plan) + execute_result = await self._run_execute_loop(validated_plan, transport=transport) self._log(f"Execute loop done, result keys={list(execute_result.keys())}") # Handle plan tool encountered @@ -670,7 +676,12 @@ async def _run_plan_loop(self, milestone: Milestone) -> ValidatedPlan | None: # Execute Loop (Layer 4) # ========================================================================= - async def _run_execute_loop(self, plan: ValidatedPlan | None) -> dict[str, Any]: + async def _run_execute_loop( + self, + plan: ValidatedPlan | None, + *, + transport: AgentChannel | None = None, + ) -> dict[str, Any]: """Run the execute loop - model-driven execution.""" self._log("Starting execute loop") execute_start = time.perf_counter() @@ -827,6 +838,7 @@ async def _run_execute_loop(self, plan: ValidatedPlan | None) -> dict[str, Any]: capability_id=capability_id, params=tool_call.get("arguments", {}), ), + transport=transport, tool_name=tool_name, tool_call_id=tool_call_id, descriptor=descriptor, @@ -836,6 +848,8 @@ async def _run_execute_loop(self, plan: ValidatedPlan | None) -> dict[str, Any]: result_success = tool_result.get("success", False) result_output = tool_result.get("output", {}) result_error = tool_result.get("error", "") + # Keep explicit result categories for downstream policy-aware handling. + result_status = tool_result.get("status", "success" if result_success else "fail") if result_success and self._is_skill_tool_call(descriptor): self._mount_skill_from_result(result_output) @@ -843,17 +857,18 @@ async def _run_execute_loop(self, plan: ValidatedPlan | None) -> dict[str, Any]: if result_success: self._log(f" ✅ Success: {result_output}") else: - self._log(f" ❌ Failed: {result_error}") + self._log(f" ❌ Failed({result_status}): {result_error}") # Add tool result as message to STM (CRITICAL: LLM needs to see result!) - tool_result_content = json.dumps({ - "success": result_success, - "output": result_output, - "error": result_error, - }) if not result_success else json.dumps({ - "success": True, - "output": result_output, - }) + tool_result_content = json.dumps( + { + "success": result_success, + "status": result_status, + "output": result_output, + "error": None if result_success else result_error, + }, + default=str, + ) tool_msg = Message( role="tool", name=tool_call_id or capability_id, # Use tool_call_id for OpenAI API format @@ -864,7 +879,7 @@ async def _run_execute_loop(self, plan: ValidatedPlan | None) -> dict[str, Any]: outputs.append(tool_result) if not result_success: - errors.append(result_error or "tool failed") + errors.append(result_error or ("tool not allowed" if result_status == "not_allow" else "tool failed")) # Reassemble context with new messages for next iteration before_context_dispatch = await self._emit_hook(HookPhase.BEFORE_CONTEXT_ASSEMBLE, {}) @@ -902,6 +917,7 @@ async def _run_tool_loop( self, request: ToolLoopRequest, *, + transport: AgentChannel | None = None, tool_name: str, tool_call_id: str, descriptor: Any | None = None, @@ -922,46 +938,6 @@ async def _run_tool_loop( self._context.budget_check() self._context.budget_use("tool_calls", 1) - if requires_approval: - allowed, approval_error = await self._resolve_tool_approval( - capability_id=request.capability_id, - params=request.params, - session_id=session_id, - tool_name=tool_name, - tool_call_id=tool_call_id, - ) - if not allowed: - await self._log_event( - "tool.error", - { - "tool_name": tool_name, - "tool_call_id": tool_call_id, - "capability_id": request.capability_id, - "error": approval_error, - "attempt": attempts, - }, - ) - await self._emit_hook( - HookPhase.AFTER_TOOL, - { - "tool_call_id": tool_call_id, - "tool_name": tool_name, - "capability_id": request.capability_id, - "attempt": attempts, - "success": False, - "error": approval_error, - "approved": False, - "evidence_collected": False, - "duration_ms": 0.0, - "budget_stats": self._budget_stats(), - }, - ) - return { - "success": False, - "error": approval_error, - "output": {}, - } - tool_start = time.perf_counter() before_tool_dispatch = await self._emit_hook( HookPhase.BEFORE_TOOL, @@ -1008,8 +984,17 @@ async def _run_tool_loop( }) try: - result = await self._tool_gateway.invoke( + approval_ctx = ApprovalInvokeContext( + session_id=session_id, + transport=transport, + tool_name=tool_name, + tool_call_id=tool_call_id, + event_logger=self._log_event, + runtime_context=self._context, + ) + result = await self._governed_tool_gateway.invoke( request.capability_id, + approval_ctx, envelope=request.envelope, **request.params, ) @@ -1026,6 +1011,15 @@ async def _run_tool_loop( if hasattr(result, "success") and not result.success: tool_success = False evidence_collected = bool(getattr(result, "evidence", [])) + denied_status = "fail" + denied_output = getattr(result, "output", {}) + approved = True + if not tool_success: + if isinstance(denied_output, dict): + candidate_status = denied_output.get("status") + if isinstance(candidate_status, str) and candidate_status: + denied_status = candidate_status + approved = denied_status != "not_allow" await self._emit_hook( HookPhase.AFTER_TOOL, { @@ -1035,7 +1029,7 @@ async def _run_tool_loop( "attempt": attempts, "success": tool_success, "error": result.error if hasattr(result, "error") else None, - "approved": True, + "approved": approved, "evidence_collected": evidence_collected, "duration_ms": (time.perf_counter() - tool_start) * 1000.0, "budget_stats": self._budget_stats(), @@ -1045,8 +1039,9 @@ async def _run_tool_loop( if not tool_success: return { "success": False, + "status": denied_status, "error": result.error or "tool failed", - "output": getattr(result, "output", {}), + "output": denied_output, "result": result, } @@ -1059,6 +1054,7 @@ async def _run_tool_loop( if done_predicate is None or _done_predicate_satisfied(done_predicate, result): return { "success": True, + "status": "success", "output": getattr(result, "output", {}), "error": getattr(result, "error", None), "result": result, @@ -1067,6 +1063,7 @@ async def _run_tool_loop( if max_calls is not None and attempts >= max_calls: return { "success": False, + "status": "fail", "error": "done predicate not satisfied before budget exhausted", "output": getattr(result, "output", {}), "result": result, @@ -1105,104 +1102,11 @@ async def _run_tool_loop( ) return { "success": False, + "status": "fail", "error": str(e), "output": {}, } - async def _resolve_tool_approval( - self, - *, - capability_id: str, - params: dict[str, Any], - session_id: str | None, - tool_name: str, - tool_call_id: str, - ) -> tuple[bool, str | None]: - if self._approval_manager is None: - return False, "tool requires approval but no approval manager is configured" - - evaluation = await self._approval_manager.evaluate( - capability_id=capability_id, - params=params, - session_id=session_id, - reason=f"Tool {capability_id} requires approval", - ) - if evaluation.status == ApprovalEvaluationStatus.ALLOW: - await self._log_event( - "tool.approval", - { - "tool_name": tool_name, - "tool_call_id": tool_call_id, - "capability_id": capability_id, - "status": "allow", - "source": "rule", - "rule_id": evaluation.rule.rule_id if evaluation.rule is not None else None, - }, - ) - return True, None - - if evaluation.status == ApprovalEvaluationStatus.DENY: - await self._log_event( - "tool.approval", - { - "tool_name": tool_name, - "tool_call_id": tool_call_id, - "capability_id": capability_id, - "status": "deny", - "source": "rule", - "rule_id": evaluation.rule.rule_id if evaluation.rule is not None else None, - }, - ) - return False, "tool invocation denied by approval rule" - - if evaluation.request is None: - return False, "tool invocation requires approval" - - request_id = evaluation.request.request_id - await self._emit_approval_pending_message( - request=evaluation.request.to_dict(), - capability_id=capability_id, - tool_name=tool_name, - tool_call_id=tool_call_id, - ) - await self._log_event( - "exec.waiting_human", - { - "checkpoint_id": request_id, - "reason": evaluation.request.reason, - "mode": "approval_memory_wait", - }, - ) - decision = await self._approval_manager.wait_for_resolution(request_id) - await self._log_event( - "exec.resume", - { - "checkpoint_id": request_id, - "decision": decision.value, - }, - ) - await self._log_event( - "tool.approval", - { - "tool_name": tool_name, - "tool_call_id": tool_call_id, - "capability_id": capability_id, - "status": decision.value, - "source": "pending_request", - "request_id": request_id, - }, - ) - await self._emit_approval_resolved_message( - request_id=request_id, - decision=decision.value, - capability_id=capability_id, - tool_name=tool_name, - tool_call_id=tool_call_id, - ) - if decision == ApprovalDecision.ALLOW: - return True, None - return False, "tool invocation denied by human approval" - # ========================================================================= # Verify # ========================================================================= @@ -1487,53 +1391,6 @@ async def _emit_hook(self, phase: HookPhase, payload: dict[str, Any]) -> HookRes return HookResult(decision=HookDecision.ALLOW) return HookResult(decision=HookDecision.ALLOW) - async def _emit_approval_pending_message( - self, - *, - request: dict[str, Any], - capability_id: str, - tool_name: str, - tool_call_id: str, - ) -> None: - payload = build_approval_pending_payload( - request=request, - capability_id=capability_id, - tool_name=tool_name, - tool_call_id=tool_call_id, - ) - await self._send_transport_payload(payload) - - async def _emit_approval_resolved_message( - self, - *, - request_id: str, - decision: str, - capability_id: str, - tool_name: str, - tool_call_id: str, - ) -> None: - payload = build_approval_resolved_payload( - request_id=request_id, - decision=decision, - capability_id=capability_id, - tool_name=tool_name, - tool_call_id=tool_call_id, - ) - await self._send_transport_payload(payload) - - async def _send_transport_payload(self, payload: dict[str, Any]) -> None: - channel = self._active_transport - if channel is None: - return - envelope = TransportEnvelope( - id=new_envelope_id(), - payload=payload, - ) - try: - await channel.send(envelope) - except Exception: - self._logger.exception("agent approval transport send failed") - def _record_token_usage(self, usage: dict[str, Any] | None) -> None: if not usage: return diff --git a/dare_framework/hook/_internal/agent_event_transport_hook.py b/dare_framework/hook/_internal/agent_event_transport_hook.py index 1d58b85c..0f5aa4d6 100644 --- a/dare_framework/hook/_internal/agent_event_transport_hook.py +++ b/dare_framework/hook/_internal/agent_event_transport_hook.py @@ -9,7 +9,7 @@ from dare_framework.hook.types import HookPhase from dare_framework.infra.component import ComponentType from dare_framework.transport.kernel import AgentChannel -from dare_framework.transport.types import TransportEnvelope, new_envelope_id +from dare_framework.transport.types import TransportEnvelope, TransportEventType, new_envelope_id _logger = logging.getLogger("dare.hook") @@ -35,8 +35,8 @@ async def invoke(self, phase: HookPhase, *args: Any, **kwargs: Any) -> Any: payload = {} envelope = TransportEnvelope( id=new_envelope_id(), + event_type=TransportEventType.HOOK.value, payload={ - "type": "hook", "phase": phase.value, "payload": payload, }, diff --git a/dare_framework/tool/_internal/control/approval_manager.py b/dare_framework/tool/_internal/control/approval_manager.py index 49db7b72..5038d421 100644 --- a/dare_framework/tool/_internal/control/approval_manager.py +++ b/dare_framework/tool/_internal/control/approval_manager.py @@ -147,6 +147,10 @@ class ApprovalEvaluation: class _PendingApproval: request: PendingApprovalRequest fingerprint: str + # Track all sessions that are currently blocked on this deduplicated request. + # The first requester is also recorded so session-filtered polling can match it + # even after subsequent evaluate() calls deduplicate to the same request id. + session_ids: set[str] = field(default_factory=set) event: asyncio.Event = field(default_factory=asyncio.Event) resolution: ApprovalDecision | None = None @@ -217,8 +221,10 @@ def __init__( self._pending_by_id: dict[str, _PendingApproval] = {} self._pending_by_fingerprint: dict[str, _PendingApproval] = {} self._resolved_by_id: dict[str, ApprovalDecision] = {} - self._pending_available = asyncio.Event() self._lock = asyncio.Lock() + # Condition-based wakeups avoid tight loops when polling with a session filter + # while unrelated pending requests exist. + self._pending_state_changed = asyncio.Condition(self._lock) @classmethod def from_paths(cls, *, workspace_dir: str | Path, user_dir: str | Path) -> ToolApprovalManager: @@ -231,31 +237,36 @@ def list_pending(self) -> list[PendingApprovalRequest]: pending.sort(key=lambda item: (item.created_at, item.request_id)) return pending - async def poll_pending(self, *, timeout_seconds: float | None = None) -> PendingApprovalRequest | None: - """Return the oldest pending approval request, optionally waiting for one.""" + async def poll_pending( + self, + *, + timeout_seconds: float | None = None, + session_id: str | None = None, + ) -> PendingApprovalRequest | None: + """Return the oldest pending approval request, optionally filtered by session.""" if timeout_seconds is not None and timeout_seconds < 0: raise ValueError("timeout_seconds must be >= 0") loop = asyncio.get_running_loop() deadline = None if timeout_seconds is None else loop.time() + timeout_seconds - while True: - async with self._lock: - request = self._oldest_pending_locked() + + async with self._pending_state_changed: + while True: + request = self._oldest_pending_locked(session_id=session_id) if request is not None: return request - wait_event = self._pending_available - if deadline is None: - await wait_event.wait() - continue + if deadline is None: + await self._pending_state_changed.wait() + continue - remaining = deadline - loop.time() - if remaining <= 0: - return None - try: - await asyncio.wait_for(wait_event.wait(), timeout=remaining) - except asyncio.TimeoutError: - return None + remaining = deadline - loop.time() + if remaining <= 0: + return None + try: + await asyncio.wait_for(self._pending_state_changed.wait(), timeout=remaining) + except asyncio.TimeoutError: + return None def list_rules(self) -> list[ApprovalRule]: combined = [ @@ -278,7 +289,8 @@ async def evaluate( command = _extract_command(params) fingerprint = _request_fingerprint(capability_id, params_hash) - async with self._lock: + # Use the condition lock consistently for pending-state mutations. + async with self._pending_state_changed: matched_rule = self._find_matching_rule( capability_id=capability_id, params_hash=params_hash, @@ -313,9 +325,16 @@ async def evaluate( created_at=self._time_fn(), ) existing = _PendingApproval(request=request, fingerprint=fingerprint) + self._track_pending_session_locked(existing, session_id) self._pending_by_fingerprint[fingerprint] = existing self._pending_by_id[request.request_id] = existing - self._pending_available.set() + self._pending_state_changed.notify_all() + else: + # A deduplicated request can gain new interested sessions later. + # Wake session-filtered poll waiters when that subscriber set expands. + added_session = self._track_pending_session_locked(existing, session_id) + if added_session: + self._pending_state_changed.notify_all() return ApprovalEvaluation( status=ApprovalEvaluationStatus.PENDING, @@ -329,7 +348,9 @@ async def wait_for_resolution( *, timeout_seconds: float | None = None, ) -> ApprovalDecision: - async with self._lock: + # Keep all pending/resolution map access on the same condition-backed lock + # so concurrency audits only need to reason about one synchronization surface. + async with self._pending_state_changed: resolved = self._resolved_by_id.pop(request_id, None) if resolved is not None: return resolved @@ -345,7 +366,7 @@ async def wait_for_resolution( if pending.resolution is None: raise RuntimeError(f"Approval request resolved without decision: {request_id}") - async with self._lock: + async with self._pending_state_changed: self._resolved_by_id.pop(request_id, None) return pending.resolution @@ -382,7 +403,7 @@ async def deny( ) async def revoke(self, rule_id: str) -> bool: - async with self._lock: + async with self._pending_state_changed: removed = self._remove_rule(rule_id) if removed: self._persist_rules_for_scope(removed.scope) @@ -397,7 +418,8 @@ async def _resolve_request( matcher: ApprovalMatcherKind, matcher_value: str | None, ) -> ApprovalRule | None: - async with self._lock: + # Keep pending-state transitions on the condition lock to avoid mixed styles. + async with self._pending_state_changed: pending = self._pending_by_id.get(request_id) if pending is None: raise KeyError(f"Unknown approval request: {request_id}") @@ -419,8 +441,7 @@ async def _resolve_request( self._resolved_by_id[request_id] = decision self._pending_by_id.pop(request_id, None) self._pending_by_fingerprint.pop(pending.fingerprint, None) - if not self._pending_by_id: - self._pending_available.clear() + self._pending_state_changed.notify_all() return rule def _append_rule(self, rule: ApprovalRule) -> None: @@ -509,15 +530,33 @@ def _find_matching_rule( return rule return None - def _oldest_pending_locked(self) -> PendingApprovalRequest | None: + def _oldest_pending_locked(self, *, session_id: str | None = None) -> PendingApprovalRequest | None: if not self._pending_by_id: return None + candidates = list(self._pending_by_id.values()) + if session_id is not None: + candidates = [ + item + for item in candidates + if session_id in item.session_ids or item.request.session_id == session_id + ] + if not candidates: + return None oldest = min( - self._pending_by_id.values(), + candidates, key=lambda item: (item.request.created_at, item.request.request_id), ) return oldest.request + @staticmethod + def _track_pending_session_locked(pending: _PendingApproval, session_id: str | None) -> bool: + if isinstance(session_id, str) and session_id: + if session_id in pending.session_ids: + return False + pending.session_ids.add(session_id) + return True + return False + def _rule_matches( rule: ApprovalRule, diff --git a/dare_framework/tool/_internal/governed_tool_gateway.py b/dare_framework/tool/_internal/governed_tool_gateway.py new file mode 100644 index 00000000..a465100c --- /dev/null +++ b/dare_framework/tool/_internal/governed_tool_gateway.py @@ -0,0 +1,327 @@ +"""Gateway-level tool invocation governance (approval + execution). + +This module keeps policy/approval decisions at the tool invocation boundary so +agent orchestration can focus on the loop itself. +""" + +from __future__ import annotations + +import logging +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any, Awaitable, Callable, Literal + +from dare_framework.tool._internal.control.approval_manager import ( + ApprovalDecision, + ApprovalEvaluationStatus, + ToolApprovalManager, +) +from dare_framework.tool._internal.runtime_context_override import ( + RUNTIME_CONTEXT_PARAM, + RuntimeContextOverride, +) +from dare_framework.tool.kernel import IToolGateway +from dare_framework.tool.types import CapabilityDescriptor, ToolResult +from dare_framework.transport.interaction.payloads import build_approval_pending_payload +from dare_framework.transport.types import ( + EnvelopeKind, + TransportEnvelope, + TransportEventType, + new_envelope_id, +) + +if TYPE_CHECKING: + from dare_framework.context import Context + from dare_framework.plan.types import Envelope + from dare_framework.transport.kernel import AgentChannel + +ApprovalEventLogger = Callable[[str, dict[str, Any]], Awaitable[None]] + + +@dataclass(frozen=True) +class ApprovalInvokeContext: + """Gateway-local approval governance context carried outside tool params.""" + + session_id: str | None = None + transport: AgentChannel | None = None + tool_name: str | None = None + tool_call_id: str | None = None + event_logger: ApprovalEventLogger | None = None + runtime_context: Context | None = None + + +@dataclass(frozen=True) +class ApprovalResolution: + """Approval decision normalized for invoke-layer status mapping.""" + + verdict: Literal["allow", "deny", "error"] + error: str | None = None + + +class GovernedToolGateway(IToolGateway): + """IToolGateway wrapper that applies approval memory before tool execution.""" + + def __init__( + self, + delegate: IToolGateway, + *, + approval_manager: ToolApprovalManager | None = None, + logger: logging.Logger | None = None, + ) -> None: + self._delegate = delegate + self._approval_manager = approval_manager + self._logger = logger or logging.getLogger("dare.tool.governed_gateway") + self._runtime_context_param = RUNTIME_CONTEXT_PARAM + + def list_capabilities(self) -> list[CapabilityDescriptor]: + return self._delegate.list_capabilities() + + async def invoke( + self, + capability_id: str, + approval: ApprovalInvokeContext | None = None, + *, + envelope: Envelope, + context: Context | None = None, + **params: Any, + ) -> ToolResult: + session_id = approval.session_id if approval is not None else None + transport = approval.transport if approval is not None else None + tool_name = approval.tool_name if approval is not None else None + tool_call_id = approval.tool_call_id if approval is not None else None + approval_event_logger = approval.event_logger if approval is not None else None + runtime_context = approval.runtime_context if approval is not None else None + if runtime_context is None: + runtime_context = context + + delegate_params = params + if approval is not None and approval.runtime_context is not None and context is not None: + # When callers pass tool argument key `context`, Python binds it to this + # named parameter. Re-inject it into tool params so schema-defined + # arguments are preserved while runtime context stays out-of-band. + delegate_params = dict(params) + delegate_params.setdefault("context", context) + + requires_approval = self._requires_approval(capability_id) + if requires_approval: + approval_resolution = await self._resolve_approval( + capability_id=capability_id, + params=dict(delegate_params), + session_id=session_id, + transport=transport, + tool_name=tool_name or capability_id, + tool_call_id=tool_call_id or "unknown", + event_logger=approval_event_logger, + ) + if approval_resolution.verdict != "allow": + status = "not_allow" if approval_resolution.verdict == "deny" else "fail" + error = approval_resolution.error or "approval check failed" + return ToolResult( + success=False, + output={"status": status}, + error=error, + ) + + result = await self._delegate.invoke( + capability_id, + envelope=envelope, + **self._build_delegate_invoke_kwargs( + runtime_context=runtime_context, + params=delegate_params, + ), + ) + return result + + def _requires_approval(self, capability_id: str) -> bool: + descriptor = self._find_capability(capability_id) + if descriptor is None: + return False + metadata = descriptor.metadata + return bool(metadata and metadata.get("requires_approval", False)) + + def _find_capability(self, capability_id: str) -> CapabilityDescriptor | None: + for descriptor in self._delegate.list_capabilities(): + if descriptor.id == capability_id: + return descriptor + return None + + async def _resolve_approval( + self, + *, + capability_id: str, + params: dict[str, Any], + session_id: str | None, + transport: AgentChannel | None, + tool_name: str, + tool_call_id: str, + event_logger: ApprovalEventLogger | None, + ) -> ApprovalResolution: + if self._approval_manager is None: + return ApprovalResolution( + verdict="error", + error="tool requires approval but no approval manager is configured", + ) + + evaluation = await self._approval_manager.evaluate( + capability_id=capability_id, + params=params, + session_id=session_id, + reason=f"Tool {capability_id} requires approval", + ) + if evaluation.status == ApprovalEvaluationStatus.ALLOW: + await self._emit_approval_event( + event_logger, + "tool.approval", + { + "tool_name": tool_name, + "tool_call_id": tool_call_id, + "capability_id": capability_id, + "status": "allow", + "source": "rule", + "rule_id": evaluation.rule.rule_id if evaluation.rule is not None else None, + }, + ) + return ApprovalResolution(verdict="allow") + if evaluation.status == ApprovalEvaluationStatus.DENY: + await self._emit_approval_event( + event_logger, + "tool.approval", + { + "tool_name": tool_name, + "tool_call_id": tool_call_id, + "capability_id": capability_id, + "status": "deny", + "source": "rule", + "rule_id": evaluation.rule.rule_id if evaluation.rule is not None else None, + }, + ) + return ApprovalResolution( + verdict="deny", + error="tool invocation denied by approval rule", + ) + if evaluation.request is None: + return ApprovalResolution( + verdict="error", + error="tool invocation requires approval", + ) + + request_id = evaluation.request.request_id + await self._emit_approval_pending_message( + request=evaluation.request.to_dict(), + transport=transport, + capability_id=capability_id, + tool_name=tool_name, + tool_call_id=tool_call_id, + ) + await self._emit_approval_event( + event_logger, + "exec.waiting_human", + { + "checkpoint_id": request_id, + "reason": evaluation.request.reason, + "mode": "approval_memory_wait", + }, + ) + decision = await self._approval_manager.wait_for_resolution(request_id) + await self._emit_approval_event( + event_logger, + "exec.resume", + { + "checkpoint_id": request_id, + "decision": decision.value, + }, + ) + await self._emit_approval_event( + event_logger, + "tool.approval", + { + "tool_name": tool_name, + "tool_call_id": tool_call_id, + "capability_id": capability_id, + "status": decision.value, + "source": "pending_request", + "request_id": request_id, + }, + ) + if decision == ApprovalDecision.ALLOW: + return ApprovalResolution(verdict="allow") + return ApprovalResolution( + verdict="deny", + error="tool invocation denied by human approval", + ) + + async def _emit_approval_pending_message( + self, + *, + request: dict[str, Any], + transport: AgentChannel | None, + capability_id: str, + tool_name: str, + tool_call_id: str, + ) -> None: + if transport is None: + return + payload = build_approval_pending_payload( + request=request, + capability_id=capability_id, + tool_name=tool_name, + tool_call_id=tool_call_id, + ) + # Approval pending is an explicit user-choice interaction shape. + resp = payload.get("resp") + if isinstance(resp, dict): + resp.setdefault( + "options", + [ + {"label": "allow", "description": "Approve this tool invocation."}, + {"label": "deny", "description": "Deny this tool invocation."}, + ], + ) + + envelope = TransportEnvelope( + id=new_envelope_id(), + kind=EnvelopeKind.SELECT, + event_type=TransportEventType.APPROVAL_PENDING.value, + payload=payload, + ) + try: + await transport.send(envelope) + except Exception: + self._logger.exception("approval pending transport send failed") + + async def _emit_approval_event( + self, + event_logger: ApprovalEventLogger | None, + event_type: str, + payload: dict[str, Any], + ) -> None: + if event_logger is None: + return + try: + await event_logger(event_type, payload) + except Exception: + self._logger.exception("approval event emission failed: %s", event_type) + + def _build_delegate_invoke_kwargs( + self, + *, + runtime_context: Context | None, + params: dict[str, Any], + ) -> dict[str, Any]: + """Assemble delegate invoke kwargs without colliding with tool arg keys.""" + kwargs: dict[str, Any] = { + "params": dict(params), + } + delegate_params: dict[str, Any] = kwargs["params"] + + if runtime_context is None: + return delegate_params + + if "context" in delegate_params: + delegate_params[self._runtime_context_param] = RuntimeContextOverride(runtime_context) + return delegate_params + + delegate_params["context"] = runtime_context + return delegate_params + + +__all__ = ["ApprovalInvokeContext", "ApprovalResolution", "GovernedToolGateway"] diff --git a/dare_framework/tool/_internal/runtime_context_override.py b/dare_framework/tool/_internal/runtime_context_override.py new file mode 100644 index 00000000..5dc0d8db --- /dev/null +++ b/dare_framework/tool/_internal/runtime_context_override.py @@ -0,0 +1,22 @@ +"""Internal marker for runtime-context override between gateway layers.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from dare_framework.context import Context + + +RUNTIME_CONTEXT_PARAM = "__dare_runtime_context__" + + +@dataclass(frozen=True) +class RuntimeContextOverride: + """Opaque wrapper so user payload keys cannot spoof runtime context.""" + + context: Context | None + + +__all__ = ["RUNTIME_CONTEXT_PARAM", "RuntimeContextOverride"] diff --git a/dare_framework/tool/action_handler.py b/dare_framework/tool/action_handler.py index aacabc58..a19492ba 100644 --- a/dare_framework/tool/action_handler.py +++ b/dare_framework/tool/action_handler.py @@ -77,7 +77,11 @@ async def invoke( if action == ResourceAction.APPROVALS_POLL: timeout_seconds = _parse_timeout_seconds(params) - request = await self._approval_manager.poll_pending(timeout_seconds=timeout_seconds) + session_id = _optional_session_id(params.get("session_id")) + request = await self._approval_manager.poll_pending( + timeout_seconds=timeout_seconds, + session_id=session_id, + ) return {"request": _pending_to_dict(request) if request is not None else None} if action == ResourceAction.APPROVALS_GRANT: @@ -180,6 +184,13 @@ def _optional_matcher_value(raw: Any) -> str | None: return text or None +def _optional_session_id(raw: Any) -> str | None: + if raw is None: + return None + text = str(raw).strip() + return text or None + + def _parse_timeout_seconds(params: dict[str, Any]) -> float | None: raw_seconds = params.get("timeout_seconds") raw_millis = params.get("timeout_ms") diff --git a/dare_framework/tool/tool_gateway.py b/dare_framework/tool/tool_gateway.py index 75a1e707..633ae435 100644 --- a/dare_framework/tool/tool_gateway.py +++ b/dare_framework/tool/tool_gateway.py @@ -2,10 +2,15 @@ from dare_framework.context import Context from dare_framework.plan import Envelope +from dare_framework.tool._internal.runtime_context_override import ( + RUNTIME_CONTEXT_PARAM, + RuntimeContextOverride, +) from dare_framework.tool import IToolGateway, IToolManager, ToolResult, CapabilityDescriptor, RunContext class ToolGateway(IToolGateway): + _RUNTIME_CONTEXT_PARAM = RUNTIME_CONTEXT_PARAM def __init__(self, tool_manager: IToolManager): self._tool_manager = tool_manager @@ -24,6 +29,17 @@ async def invoke( ) -> ToolResult: if envelope.allowed_capability_ids and capability_id not in envelope.allowed_capability_ids: raise PermissionError(f"Capability '{capability_id}' not allowed by envelope") + tool_params = dict(params) + runtime_context = context + runtime_context_override = tool_params.pop(self._RUNTIME_CONTEXT_PARAM, None) + # Ignore caller-provided values on this reserved key unless they carry + # the internal wrapper type injected by GovernedToolGateway. + if isinstance(runtime_context_override, RuntimeContextOverride): + runtime_context = runtime_context_override.context + if context is not None: + # `context` was consumed by this gateway's reserved kwarg slot; + # recover it as an explicit tool argument when collision occurs. + tool_params.setdefault("context", context) tool = self._tool_manager.get_tool(capability_id) - tool_context = RunContext(context) - return await tool.execute(run_context=tool_context, **params) + tool_context = RunContext(runtime_context) + return await tool.execute(run_context=tool_context, **tool_params) diff --git a/dare_framework/transport/__init__.py b/dare_framework/transport/__init__.py index a4180658..616d0568 100644 --- a/dare_framework/transport/__init__.py +++ b/dare_framework/transport/__init__.py @@ -1,9 +1,11 @@ """transport domain facade.""" -from dare_framework.transport.interfaces import AgentChannel, ClientChannel +from dare_framework.transport.interfaces import AgentChannel, ClientChannel, PollableClientChannel from dare_framework.transport.types import ( EnvelopeKind, + TransportEventType, TransportEnvelope, + normalize_transport_event_type, new_envelope_id, Receiver, Sender, @@ -23,8 +25,11 @@ __all__ = [ "AgentChannel", "ClientChannel", + "PollableClientChannel", "EnvelopeKind", + "TransportEventType", "TransportEnvelope", + "normalize_transport_event_type", "new_envelope_id", "Receiver", "Sender", diff --git a/dare_framework/transport/_internal/adapters.py b/dare_framework/transport/_internal/adapters.py index fd21f245..5d2e79ee 100644 --- a/dare_framework/transport/_internal/adapters.py +++ b/dare_framework/transport/_internal/adapters.py @@ -7,8 +7,16 @@ from typing import Any, Callable from dare_framework.transport.interaction.controls import AgentControl -from dare_framework.transport.kernel import ClientChannel -from dare_framework.transport.types import EnvelopeKind, Receiver, Sender, TransportEnvelope, new_envelope_id +from dare_framework.transport.kernel import ClientChannel, PollableClientChannel +from dare_framework.transport.types import ( + EnvelopeKind, + normalize_transport_event_type, + Receiver, + Sender, + TransportEnvelope, + TransportEventType, + new_envelope_id, +) class StdioClientChannel(ClientChannel): @@ -30,10 +38,10 @@ def attach_agent_envelope_sender(self, sender: Sender) -> None: def agent_envelope_receiver(self) -> Receiver: async def recv(msg: TransportEnvelope) -> None: + event_type = _resolve_transport_event_type(msg) payload = msg.payload if isinstance(payload, dict): - payload_type = payload.get("type") - if payload_type == "result": + if event_type == TransportEventType.RESULT.value: kind = payload.get("kind") resp = payload.get("resp") if kind == "message": @@ -43,9 +51,9 @@ async def recv(msg: TransportEnvelope) -> None: output = payload.get("output") else: output = resp if resp is not None else payload - elif payload_type == "error": + elif event_type == TransportEventType.ERROR.value: output = payload.get("reason") or payload.get("error") - elif payload_type == "approval_pending": + elif event_type == TransportEventType.APPROVAL_PENDING.value: resp = payload.get("resp") request_id = None if isinstance(resp, dict): @@ -53,7 +61,7 @@ async def recv(msg: TransportEnvelope) -> None: if isinstance(request, dict): request_id = request.get("request_id") output = f"approval pending: request_id={request_id or '?'}" - elif payload_type == "approval_resolved": + elif event_type == TransportEventType.APPROVAL_RESOLVED.value: resp = payload.get("resp") request_id = None decision = None @@ -61,7 +69,7 @@ async def recv(msg: TransportEnvelope) -> None: request_id = resp.get("request_id") decision = resp.get("decision") output = f"approval resolved: request_id={request_id or '?'} decision={decision or '?'}" - elif payload_type == "hook": + elif event_type == TransportEventType.HOOK.value: output = payload.get("event") else: output = payload @@ -142,7 +150,7 @@ async def handle_ws_message(self, raw: Any) -> None: await self._sender(envelope) -class DirectClientChannel(ClientChannel): +class DirectClientChannel(PollableClientChannel): """Direct in-process adapter for request/response patterns.""" def __init__(self) -> None: @@ -172,6 +180,7 @@ async def ask(self, req: TransportEnvelope, timeout: float = 30.0) -> TransportE id=new_envelope_id(), reply_to=req.reply_to, kind=req.kind, + event_type=req.event_type, payload=req.payload, meta=req.meta, stream_id=req.stream_id, @@ -203,6 +212,7 @@ def _default_serialize(msg: TransportEnvelope) -> str: "id": msg.id, "reply_to": msg.reply_to, "kind": msg.kind, + "event_type": msg.event_type, "payload": msg.payload, "meta": msg.meta, "stream_id": msg.stream_id, @@ -226,6 +236,7 @@ def _default_deserialize(raw: Any) -> TransportEnvelope: id=str(data.get("id") or new_envelope_id()), reply_to=data.get("reply_to"), kind=data.get("kind"), + event_type=data.get("event_type"), payload=data.get("payload"), meta=data.get("meta") or {}, stream_id=data.get("stream_id"), @@ -233,4 +244,11 @@ def _default_deserialize(raw: Any) -> TransportEnvelope: ) +def _resolve_transport_event_type(msg: TransportEnvelope) -> str | None: + """Resolve event_type for receiver routing.""" + if isinstance(msg.event_type, str): + return normalize_transport_event_type(msg.event_type) + return None + + __all__ = ["StdioClientChannel", "WebSocketClientChannel", "DirectClientChannel"] diff --git a/dare_framework/transport/_internal/default_channel.py b/dare_framework/transport/_internal/default_channel.py index 624271f7..55c16b44 100644 --- a/dare_framework/transport/_internal/default_channel.py +++ b/dare_framework/transport/_internal/default_channel.py @@ -16,6 +16,7 @@ Receiver, Sender, TransportEnvelope, + TransportEventType, new_envelope_id, ) @@ -258,6 +259,7 @@ async def _send_result( id=new_envelope_id(), reply_to=reply_to, kind=EnvelopeKind.MESSAGE, + event_type=TransportEventType.RESULT.value, payload=build_success_payload( kind=kind, target=target, @@ -280,6 +282,7 @@ async def _send_error( id=new_envelope_id(), reply_to=reply_to, kind=EnvelopeKind.MESSAGE, + event_type=TransportEventType.ERROR.value, payload=build_error_payload( kind=kind, target=target, diff --git a/dare_framework/transport/interaction/payloads.py b/dare_framework/transport/interaction/payloads.py index d41b18cc..6fa4eda2 100644 --- a/dare_framework/transport/interaction/payloads.py +++ b/dare_framework/transport/interaction/payloads.py @@ -8,7 +8,6 @@ def build_success_payload(*, kind: str, target: str, resp: Any) -> dict[str, Any]: """Build a unified success payload for action/control/message paths.""" return { - "type": "result", "kind": kind, "target": target, "ok": True, @@ -20,7 +19,6 @@ def build_error_payload(*, kind: str, target: str, code: str, reason: str) -> di """Build a unified error payload with deterministic error fields.""" detail = {"code": code, "reason": reason} return { - "type": "error", "kind": kind, "target": target, "ok": False, @@ -40,7 +38,6 @@ def build_approval_pending_payload( ) -> dict[str, Any]: """Build a transport payload for a pending tool approval request.""" return { - "type": "approval_pending", "kind": "approval", "target": capability_id, "ok": True, @@ -63,7 +60,6 @@ def build_approval_resolved_payload( ) -> dict[str, Any]: """Build a transport payload for a resolved tool approval request.""" return { - "type": "approval_resolved", "kind": "approval", "target": capability_id, "ok": True, diff --git a/dare_framework/transport/interfaces.py b/dare_framework/transport/interfaces.py index e1e30c5a..24dfffb2 100644 --- a/dare_framework/transport/interfaces.py +++ b/dare_framework/transport/interfaces.py @@ -2,6 +2,6 @@ from __future__ import annotations -from dare_framework.transport.kernel import AgentChannel, ClientChannel +from dare_framework.transport.kernel import AgentChannel, ClientChannel, PollableClientChannel -__all__ = ["AgentChannel", "ClientChannel"] +__all__ = ["AgentChannel", "ClientChannel", "PollableClientChannel"] diff --git a/dare_framework/transport/kernel.py b/dare_framework/transport/kernel.py index ff6f6750..4f8d2476 100644 --- a/dare_framework/transport/kernel.py +++ b/dare_framework/transport/kernel.py @@ -2,7 +2,7 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Protocol +from typing import TYPE_CHECKING, Protocol, runtime_checkable from dare_framework.transport.types import ( Receiver, @@ -25,6 +25,14 @@ def agent_envelope_receiver(self) -> Receiver: """Return the receiver used to deliver envelopes from the agent outbox.""" +@runtime_checkable +class PollableClientChannel(ClientChannel, Protocol): + """Extension for client channels that support polling unsolicited events.""" + + async def poll(self, timeout: float | None = None) -> TransportEnvelope | None: + """Poll unsolicited envelopes emitted by the agent channel.""" + + class AgentChannel(Protocol): """Agent-facing channel contract for transport.""" @@ -75,4 +83,4 @@ def build( ) -__all__ = ["AgentChannel", "ClientChannel"] +__all__ = ["AgentChannel", "ClientChannel", "PollableClientChannel"] diff --git a/dare_framework/transport/types.py b/dare_framework/transport/types.py index 93445e2c..297cede1 100644 --- a/dare_framework/transport/types.py +++ b/dare_framework/transport/types.py @@ -12,10 +12,38 @@ class EnvelopeKind(StrEnum): """Strong envelope categories for transport dispatch.""" MESSAGE = "message" + SELECT = "select" ACTION = "action" CONTROL = "control" +class TransportEventType(StrEnum): + """Canonical event categories carried by message envelopes.""" + + RESULT = "result" + ERROR = "error" + HOOK = "hook" + APPROVAL_PENDING = "approval.pending" + APPROVAL_RESOLVED = "approval.resolved" + + +_LEGACY_PAYLOAD_EVENT_TYPE_MAP: dict[str, str] = { + # Only legacy aliases that differ from canonical event_type values. + "approval_pending": TransportEventType.APPROVAL_PENDING.value, + "approval_resolved": TransportEventType.APPROVAL_RESOLVED.value, +} + + +def normalize_transport_event_type(raw: str | None) -> str | None: + """Normalize legacy/new event_type strings into canonical values.""" + if raw is None: + return None + normalized = raw.strip() + if not normalized: + return None + return _LEGACY_PAYLOAD_EVENT_TYPE_MAP.get(normalized, normalized) + + @dataclass(frozen=True) class TransportEnvelope: """Transport envelope for agent/client messages.""" @@ -23,6 +51,7 @@ class TransportEnvelope: id: str reply_to: str | None = None kind: EnvelopeKind = EnvelopeKind.MESSAGE + event_type: str | None = None payload: Any = None meta: dict[str, Any] = field(default_factory=dict) stream_id: str | None = None @@ -35,10 +64,23 @@ def __post_init__(self) -> None: object.__setattr__(self, "kind", EnvelopeKind(kind)) except ValueError as exc: raise ValueError(f"invalid envelope kind: {kind!r}") from exc - return + kind = self.kind if not isinstance(kind, EnvelopeKind): raise TypeError(f"invalid envelope kind type: {type(kind).__name__}") + event_type = self.event_type + if isinstance(event_type, TransportEventType): + object.__setattr__(self, "event_type", event_type.value) + return + if event_type is None: + return + if not isinstance(event_type, str): + raise TypeError(f"invalid event_type type: {type(event_type).__name__}") + normalized = normalize_transport_event_type(event_type) + if normalized is None: + raise ValueError("event_type must not be empty") + object.__setattr__(self, "event_type", normalized) + def new_envelope_id() -> str: """Generate a new envelope id.""" @@ -51,6 +93,8 @@ def new_envelope_id() -> str: __all__ = [ "EnvelopeKind", + "TransportEventType", + "normalize_transport_event_type", "TransportEnvelope", "new_envelope_id", "Sender", diff --git a/docs/guides/Tool_Approval_Memory.md b/docs/guides/Tool_Approval_Memory.md index a3dd2558..42548b14 100644 --- a/docs/guides/Tool_Approval_Memory.md +++ b/docs/guides/Tool_Approval_Memory.md @@ -115,6 +115,7 @@ 可选参数: - `timeout_seconds` 或 `timeout_ms` +- `session_id`(仅返回该 session 的最早 pending) 输出: @@ -134,8 +135,10 @@ 当 `requires_approval=true` 的工具调用进入 pending 时,runtime 会通过 transport 主动发出: -- `type="approval_pending"`:包含 pending request 详情 -- `type="approval_resolved"`:包含 request_id 与最终 decision(allow/deny) +- `event_type="approval.pending"`:包含 pending request 详情 +- `event_type="approval.resolved"`:包含 request_id 与最终 decision(allow/deny) + +协议说明:客户端应仅使用 `event_type` 做分流;`payload.type` 已移除。 这允许客户端实现 Codex/Claude Code 风格的“消息流里出现审批卡片”,同时再用 `approvals:grant|deny` 完成决策。 @@ -155,7 +158,7 @@ `examples/05-dare-coding-agent-enhanced/cli.py` 与 `examples/06-dare-coding-agent-mcp/cli.py` 都支持: - `/approvals list` -- `/approvals poll [timeout_ms=30000]` +- `/approvals poll [timeout_ms=30000] [session_id=...]` - `/approvals grant [scope=workspace] [matcher=exact_params] [matcher_value=...]` - `/approvals deny [scope=once] [matcher=exact_params] [matcher_value=...]` - `/approvals revoke ` diff --git a/examples/05-dare-coding-agent-enhanced/README.md b/examples/05-dare-coding-agent-enhanced/README.md index c7bdbbc8..ce149822 100644 --- a/examples/05-dare-coding-agent-enhanced/README.md +++ b/examples/05-dare-coding-agent-enhanced/README.md @@ -29,7 +29,7 @@ python main.py - `/reject`:取消当前计划 - `/status`:查看状态 - `/approvals list`:查看待审批请求与当前审批规则 -- `/approvals poll [timeout_ms=30000]`:阻塞等待下一个待审批请求(无请求则超时返回) +- `/approvals poll [timeout_ms=30000] [session_id=...]`:阻塞等待下一个待审批请求(可按 session 过滤) - `/approvals grant [scope=workspace] [matcher=exact_params] [matcher_value=...]`:批准请求并可写入规则 - `/approvals deny [scope=once] [matcher=exact_params] [matcher_value=...]`:拒绝请求并可写入规则 - `/approvals revoke `:撤销审批规则 diff --git a/examples/05-dare-coding-agent-enhanced/cli.py b/examples/05-dare-coding-agent-enhanced/cli.py index 9cc2a1e4..b68435f9 100644 --- a/examples/05-dare-coding-agent-enhanced/cli.py +++ b/examples/05-dare-coding-agent-enhanced/cli.py @@ -28,8 +28,8 @@ from dare_framework.knowledge import create_knowledge from dare_framework.model import OpenRouterModelAdapter from dare_framework.plan import DefaultPlanner, DefaultRemediator, Task -from dare_framework.tool.action_handler import ApprovalsActionHandler from dare_framework.tool._internal.tools import ReadFileTool, RunCommandTool, SearchCodeTool, WriteFileTool +from dare_framework.transport import AgentChannel, DirectClientChannel, EnvelopeKind, TransportEnvelope, new_envelope_id from dare_framework.transport.interaction.resource_action import ResourceAction from validators.file_validator import FileExistsValidator @@ -269,6 +269,7 @@ def _create_builder( display: CLIDisplay, *, config: Config | None = None, + agent_channel: AgentChannel | None = None, ) -> DareAgentBuilder: """创建 DareAgentBuilder;MCP 与 initial_skill_path 由 builder.build() 内部从 config 读取。""" model = OpenRouterModelAdapter( @@ -303,6 +304,8 @@ def _create_builder( .with_remediator(DefaultRemediator(model, verbose=False)) .with_event_log(event_log) ) + if agent_channel is not None: + builder = builder.with_agent_channel(agent_channel) return builder @@ -433,17 +436,75 @@ async def _handle_mcp_command( def _approvals_usage(display: CLIDisplay) -> None: display.info("/approvals list") - display.info("/approvals poll [timeout_ms=30000]") + display.info("/approvals poll [timeout_ms=30000] [session_id=...]") display.info("/approvals grant [scope=workspace] [matcher=exact_params] [matcher_value=...]") display.info("/approvals deny [scope=once] [matcher=exact_params] [matcher_value=...]") display.info("/approvals revoke ") -def _resolve_approvals_handler(agent: Any) -> ApprovalsActionHandler | None: - approval_manager = getattr(agent, "_approval_manager", None) - if approval_manager is None: - return None - return ApprovalsActionHandler(approval_manager) +async def _invoke_approval_action( + approval_client: DirectClientChannel, + action: ResourceAction, + *, + params: dict[str, Any] | None = None, +) -> dict[str, Any]: + meta = dict(params or {}) + request = TransportEnvelope( + id=new_envelope_id(), + kind=EnvelopeKind.ACTION, + payload=action.value, + meta=meta, + ) + response = await approval_client.ask( + request, + timeout=_approval_action_timeout_seconds(action, meta), + ) + payload = response.payload + if not isinstance(payload, dict): + raise RuntimeError(f"unexpected action response payload: {payload!r}") + + event_type = response.event_type + if not isinstance(event_type, str) or not event_type: + raise RuntimeError("invalid action response: missing event_type") + if event_type == "error": + raise RuntimeError(str(payload.get("reason") or payload.get("error") or "action failed")) + + resp = payload.get("resp") + if not isinstance(resp, dict): + raise RuntimeError(f"unexpected action response shape: {payload!r}") + + result = resp.get("result") + if not isinstance(result, dict): + raise RuntimeError(f"unexpected action result shape: {payload!r}") + return result + + +def _approval_action_timeout_seconds(action: ResourceAction, params: dict[str, Any]) -> float: + default_timeout = 30.0 + if action != ResourceAction.APPROVALS_POLL: + return default_timeout + + poll_timeout_seconds = _parse_poll_timeout_seconds(params) + if poll_timeout_seconds is None: + return default_timeout + # Leave a small transport cushion so ask() does not time out first. + return max(default_timeout, poll_timeout_seconds + 5.0) + + +def _parse_poll_timeout_seconds(params: dict[str, Any]) -> float | None: + raw_seconds = params.get("timeout_seconds") + raw_millis = params.get("timeout_ms") + if raw_seconds is not None: + seconds = float(raw_seconds) + if seconds < 0: + raise ValueError("timeout_seconds must be >= 0") + return seconds + if raw_millis is not None: + millis = float(raw_millis) + if millis < 0: + raise ValueError("timeout_ms must be >= 0") + return millis / 1000.0 + return None def _parse_key_value_args(tokens: list[str]) -> tuple[list[str], dict[str, str]]: @@ -473,15 +534,23 @@ def _build_approval_action_params( return params +def _build_approval_poll_params(trailing_args: list[str]) -> dict[str, Any]: + _positional, options = _parse_key_value_args(trailing_args) + params: dict[str, Any] = {} + for key in ("timeout_ms", "timeout_seconds", "session_id"): + if key in options and options[key]: + params[key] = options[key] + return params + + async def _handle_approvals_command( args: list[str], *, - agent: Any, + approval_client: DirectClientChannel | None, display: CLIDisplay, ) -> None: - handler = _resolve_approvals_handler(agent) - if handler is None: - display.warn("approval manager unavailable on current agent") + if approval_client is None: + display.warn("approval transport unavailable") return if not args: _approvals_usage(display) @@ -490,7 +559,7 @@ async def _handle_approvals_command( subcommand = args[0].lower() try: if subcommand == "list": - result = await handler.invoke(ResourceAction.APPROVALS_LIST) + result = await _invoke_approval_action(approval_client, ResourceAction.APPROVALS_LIST) pending = result.get("pending", []) rules = result.get("rules", []) display.info(f"pending={len(pending)} rules={len(rules)}") @@ -498,12 +567,12 @@ async def _handle_approvals_command( return if subcommand == "poll": - _positional, options = _parse_key_value_args(args[1:]) - params: dict[str, Any] = {} - for key in ("timeout_ms", "timeout_seconds"): - if key in options and options[key]: - params[key] = options[key] - result = await handler.invoke(ResourceAction.APPROVALS_POLL, **params) + params = _build_approval_poll_params(args[1:]) + result = await _invoke_approval_action( + approval_client, + ResourceAction.APPROVALS_POLL, + params=params, + ) request = result.get("request") if isinstance(request, dict): display.info(f"pending request: {request.get('request_id', '?')}") @@ -524,7 +593,7 @@ async def _handle_approvals_command( else ResourceAction.APPROVALS_DENY ) params = _build_approval_action_params(request_id=request_id, trailing_args=args[2:]) - result = await handler.invoke(action, **params) + result = await _invoke_approval_action(approval_client, action, params=params) display.ok(f"{subcommand} applied: {request_id}") print(json.dumps(result, ensure_ascii=False, indent=2), flush=True) return @@ -535,9 +604,10 @@ async def _handle_approvals_command( _approvals_usage(display) return rule_id = args[1] - result = await handler.invoke( + result = await _invoke_approval_action( + approval_client, ResourceAction.APPROVALS_REVOKE, - rule_id=rule_id, + params={"rule_id": rule_id}, ) if result.get("removed"): display.ok(f"revoked rule: {rule_id}") @@ -604,6 +674,7 @@ async def run_cli_loop( model: OpenRouterModelAdapter, display: CLIDisplay, mcp_state: MCPRuntimeState | None = None, + approval_client: DirectClientChannel | None = None, state: CLISessionState | None = None, background_execute: bool = False, ) -> tuple[CLISessionState, bool]: @@ -642,7 +713,7 @@ async def run_cli_loop( if cmd.type == CommandType.APPROVALS: await _handle_approvals_command( cmd.args, - agent=agent, + approval_client=approval_client, display=display, ) continue @@ -763,6 +834,8 @@ async def main(argv: list[str] | None = None) -> None: if config.mcp_paths: display.info("MCP config found. Start local_mcp_server.py in another terminal if using local_math.") + approval_client = DirectClientChannel() + approval_channel = AgentChannel.build(approval_client) builder = _create_builder( workspace, model_name, @@ -771,8 +844,10 @@ async def main(argv: list[str] | None = None) -> None: timeout_seconds, display, config=config, + agent_channel=approval_channel, ) agent = await builder.build() + await approval_channel.start() mcp_state = MCPRuntimeState( config_provider=config_provider, config=config, @@ -786,50 +861,56 @@ async def main(argv: list[str] | None = None) -> None: http_client_options={"timeout": timeout_seconds}, ) - if args.demo: - script_path = Path(args.demo) - lines = load_script_lines(script_path) - await run_cli_loop( - lines, - agent=agent, - model=model, - display=display, - mcp_state=mcp_state, - ) - return + try: + if args.demo: + script_path = Path(args.demo) + lines = load_script_lines(script_path) + await run_cli_loop( + lines, + agent=agent, + model=model, + display=display, + mcp_state=mcp_state, + approval_client=approval_client, + ) + return - if args.script: - script_path = Path(args.script) - lines = load_script_lines(script_path) - await run_cli_loop( - lines, - agent=agent, - model=model, - display=display, - mcp_state=mcp_state, - ) - return + if args.script: + script_path = Path(args.script) + lines = load_script_lines(script_path) + await run_cli_loop( + lines, + agent=agent, + model=model, + display=display, + mcp_state=mcp_state, + approval_client=approval_client, + ) + return - display.info("type /help for commands. /quit to exit.") - cli_state = CLISessionState() - while True: - try: - raw = input("dare> ").strip() - except (EOFError, KeyboardInterrupt): - break - if not raw: - continue - cli_state, quit_requested = await run_cli_loop( - [raw], - agent=agent, - model=model, - display=display, - mcp_state=mcp_state, - state=cli_state, - background_execute=True, - ) - if quit_requested: - break + display.info("type /help for commands. /quit to exit.") + cli_state = CLISessionState() + while True: + try: + raw = input("dare> ").strip() + except (EOFError, KeyboardInterrupt): + break + if not raw: + continue + cli_state, quit_requested = await run_cli_loop( + [raw], + agent=agent, + model=model, + display=display, + mcp_state=mcp_state, + approval_client=approval_client, + state=cli_state, + background_execute=True, + ) + if quit_requested: + break + finally: + await approval_channel.stop() if __name__ == "__main__": diff --git a/examples/06-dare-coding-agent-mcp/README.md b/examples/06-dare-coding-agent-mcp/README.md index 3f5721c1..e42e3624 100644 --- a/examples/06-dare-coding-agent-mcp/README.md +++ b/examples/06-dare-coding-agent-mcp/README.md @@ -69,7 +69,7 @@ MCP server 定义在 `.dare/mcp/local_math.json`: - `/reject`:取消待审批计划 - `/status`:查看当前状态 - `/approvals list`:查看待审批请求与当前审批规则 -- `/approvals poll [timeout_ms=30000]`:阻塞等待下一个待审批请求(无请求则超时返回) +- `/approvals poll [timeout_ms=30000] [session_id=...]`:阻塞等待下一个待审批请求(可按 session 过滤) - `/approvals grant [scope=workspace] [matcher=exact_params] [matcher_value=...]`:批准请求并可写入规则 - `/approvals deny [scope=once] [matcher=exact_params] [matcher_value=...]`:拒绝请求并可写入规则 - `/approvals revoke `:撤销审批规则 diff --git a/examples/06-dare-coding-agent-mcp/cli.py b/examples/06-dare-coding-agent-mcp/cli.py index d6e21953..566bf99f 100644 --- a/examples/06-dare-coding-agent-mcp/cli.py +++ b/examples/06-dare-coding-agent-mcp/cli.py @@ -28,8 +28,8 @@ from dare_framework.event.types import Event, RuntimeSnapshot from dare_framework.model import OpenRouterModelAdapter from dare_framework.plan import DefaultPlanner, DefaultRemediator, Task -from dare_framework.tool.action_handler import ApprovalsActionHandler from dare_framework.tool._internal.tools import ReadFileTool, RunCommandTool, SearchCodeTool, WriteFileTool +from dare_framework.transport import AgentChannel, DirectClientChannel, EnvelopeKind, TransportEnvelope, new_envelope_id from dare_framework.transport.interaction.resource_action import ResourceAction from validators.file_validator import FileExistsValidator @@ -290,6 +290,8 @@ async def build_agent( timeout_seconds: float, display: CLIDisplay, config: Config, + *, + agent_channel: AgentChannel | None = None, ) -> Any: model = OpenRouterModelAdapter( model=model_name, @@ -316,6 +318,8 @@ async def build_agent( .with_remediator(DefaultRemediator(model, verbose=False)) .with_event_log(event_log) ) + if agent_channel is not None: + builder = builder.with_agent_channel(agent_channel) agent = await builder.build() return agent @@ -521,17 +525,75 @@ async def _handle_mcp_command( def _approvals_usage(display: CLIDisplay) -> None: display.info("/approvals list") - display.info("/approvals poll [timeout_ms=30000]") + display.info("/approvals poll [timeout_ms=30000] [session_id=...]") display.info("/approvals grant [scope=workspace] [matcher=exact_params] [matcher_value=...]") display.info("/approvals deny [scope=once] [matcher=exact_params] [matcher_value=...]") display.info("/approvals revoke ") -def _resolve_approvals_handler(agent: Any) -> ApprovalsActionHandler | None: - approval_manager = getattr(agent, "_approval_manager", None) - if approval_manager is None: - return None - return ApprovalsActionHandler(approval_manager) +async def _invoke_approval_action( + approval_client: DirectClientChannel, + action: ResourceAction, + *, + params: dict[str, Any] | None = None, +) -> dict[str, Any]: + meta = dict(params or {}) + request = TransportEnvelope( + id=new_envelope_id(), + kind=EnvelopeKind.ACTION, + payload=action.value, + meta=meta, + ) + response = await approval_client.ask( + request, + timeout=_approval_action_timeout_seconds(action, meta), + ) + payload = response.payload + if not isinstance(payload, dict): + raise RuntimeError(f"unexpected action response payload: {payload!r}") + + event_type = response.event_type + if not isinstance(event_type, str) or not event_type: + raise RuntimeError("invalid action response: missing event_type") + if event_type == "error": + raise RuntimeError(str(payload.get("reason") or payload.get("error") or "action failed")) + + resp = payload.get("resp") + if not isinstance(resp, dict): + raise RuntimeError(f"unexpected action response shape: {payload!r}") + + result = resp.get("result") + if not isinstance(result, dict): + raise RuntimeError(f"unexpected action result shape: {payload!r}") + return result + + +def _approval_action_timeout_seconds(action: ResourceAction, params: dict[str, Any]) -> float: + default_timeout = 30.0 + if action != ResourceAction.APPROVALS_POLL: + return default_timeout + + poll_timeout_seconds = _parse_poll_timeout_seconds(params) + if poll_timeout_seconds is None: + return default_timeout + # Leave a small transport cushion so ask() does not time out first. + return max(default_timeout, poll_timeout_seconds + 5.0) + + +def _parse_poll_timeout_seconds(params: dict[str, Any]) -> float | None: + raw_seconds = params.get("timeout_seconds") + raw_millis = params.get("timeout_ms") + if raw_seconds is not None: + seconds = float(raw_seconds) + if seconds < 0: + raise ValueError("timeout_seconds must be >= 0") + return seconds + if raw_millis is not None: + millis = float(raw_millis) + if millis < 0: + raise ValueError("timeout_ms must be >= 0") + return millis / 1000.0 + return None def _parse_key_value_args(tokens: list[str]) -> tuple[list[str], dict[str, str]]: @@ -561,15 +623,23 @@ def _build_approval_action_params( return params +def _build_approval_poll_params(trailing_args: list[str]) -> dict[str, Any]: + _positional, options = _parse_key_value_args(trailing_args) + params: dict[str, Any] = {} + for key in ("timeout_ms", "timeout_seconds", "session_id"): + if key in options and options[key]: + params[key] = options[key] + return params + + async def _handle_approvals_command( args: list[str], *, - agent: Any, + approval_client: DirectClientChannel | None, display: CLIDisplay, ) -> None: - handler = _resolve_approvals_handler(agent) - if handler is None: - display.warn("approval manager unavailable on current agent") + if approval_client is None: + display.warn("approval transport unavailable") return if not args: _approvals_usage(display) @@ -578,7 +648,7 @@ async def _handle_approvals_command( subcommand = args[0].lower() try: if subcommand == "list": - result = await handler.invoke(ResourceAction.APPROVALS_LIST) + result = await _invoke_approval_action(approval_client, ResourceAction.APPROVALS_LIST) pending = result.get("pending", []) rules = result.get("rules", []) display.info(f"pending={len(pending)} rules={len(rules)}") @@ -586,12 +656,12 @@ async def _handle_approvals_command( return if subcommand == "poll": - _positional, options = _parse_key_value_args(args[1:]) - params: dict[str, Any] = {} - for key in ("timeout_ms", "timeout_seconds"): - if key in options and options[key]: - params[key] = options[key] - result = await handler.invoke(ResourceAction.APPROVALS_POLL, **params) + params = _build_approval_poll_params(args[1:]) + result = await _invoke_approval_action( + approval_client, + ResourceAction.APPROVALS_POLL, + params=params, + ) request = result.get("request") if isinstance(request, dict): display.info(f"pending request: {request.get('request_id', '?')}") @@ -612,7 +682,7 @@ async def _handle_approvals_command( else ResourceAction.APPROVALS_DENY ) params = _build_approval_action_params(request_id=request_id, trailing_args=args[2:]) - result = await handler.invoke(action, **params) + result = await _invoke_approval_action(approval_client, action, params=params) display.ok(f"{subcommand} applied: {request_id}") print(json.dumps(result, ensure_ascii=False, indent=2), flush=True) return @@ -623,9 +693,10 @@ async def _handle_approvals_command( _approvals_usage(display) return rule_id = args[1] - result = await handler.invoke( + result = await _invoke_approval_action( + approval_client, ResourceAction.APPROVALS_REVOKE, - rule_id=rule_id, + params={"rule_id": rule_id}, ) if result.get("removed"): display.ok(f"revoked rule: {rule_id}") @@ -692,6 +763,7 @@ async def run_cli_loop( model: OpenRouterModelAdapter, display: CLIDisplay, mcp_state: MCPRuntimeState | None = None, + approval_client: DirectClientChannel | None = None, state: CLISessionState | None = None, background_execute: bool = False, ) -> tuple[CLISessionState, bool]: @@ -730,7 +802,7 @@ async def run_cli_loop( if cmd.type == CommandType.APPROVALS: await _handle_approvals_command( cmd.args, - agent=agent, + approval_client=approval_client, display=display, ) continue @@ -853,6 +925,8 @@ async def main(argv: list[str] | None = None) -> None: else: display.warn("No mcp_paths configured. MCP tools will not be loaded by default.") + approval_client = DirectClientChannel() + approval_channel = AgentChannel.build(approval_client) agent = await build_agent( workspace, model_name, @@ -861,7 +935,9 @@ async def main(argv: list[str] | None = None) -> None: timeout_seconds, display, config, + agent_channel=approval_channel, ) + await approval_channel.start() mcp_state = MCPRuntimeState( config_provider=config_provider, config=config, @@ -876,50 +952,56 @@ async def main(argv: list[str] | None = None) -> None: http_client_options={"timeout": timeout_seconds}, ) - if args.demo: - script_path = Path(args.demo) - lines = load_script_lines(script_path) - await run_cli_loop( - lines, - agent=agent, - model=model, - display=display, - mcp_state=mcp_state, - ) - return + try: + if args.demo: + script_path = Path(args.demo) + lines = load_script_lines(script_path) + await run_cli_loop( + lines, + agent=agent, + model=model, + display=display, + mcp_state=mcp_state, + approval_client=approval_client, + ) + return - if args.script: - script_path = Path(args.script) - lines = load_script_lines(script_path) - await run_cli_loop( - lines, - agent=agent, - model=model, - display=display, - mcp_state=mcp_state, - ) - return + if args.script: + script_path = Path(args.script) + lines = load_script_lines(script_path) + await run_cli_loop( + lines, + agent=agent, + model=model, + display=display, + mcp_state=mcp_state, + approval_client=approval_client, + ) + return - display.info("type /help for commands. /quit to exit.") - cli_state = CLISessionState() - while True: - try: - raw = input("dare> ").strip() - except (EOFError, KeyboardInterrupt): - break - if not raw: - continue - cli_state, quit_requested = await run_cli_loop( - [raw], - agent=agent, - model=model, - display=display, - mcp_state=mcp_state, - state=cli_state, - background_execute=True, - ) - if quit_requested: - break + display.info("type /help for commands. /quit to exit.") + cli_state = CLISessionState() + while True: + try: + raw = input("dare> ").strip() + except (EOFError, KeyboardInterrupt): + break + if not raw: + continue + cli_state, quit_requested = await run_cli_loop( + [raw], + agent=agent, + model=model, + display=display, + mcp_state=mcp_state, + approval_client=approval_client, + state=cli_state, + background_execute=True, + ) + if quit_requested: + break + finally: + await approval_channel.stop() if __name__ == "__main__": diff --git a/examples/07-tool-approval-memory/README.md b/examples/07-tool-approval-memory/README.md index 3cc8896a..da8a660e 100644 --- a/examples/07-tool-approval-memory/README.md +++ b/examples/07-tool-approval-memory/README.md @@ -18,7 +18,7 @@ python main.py ## 你会看到什么 -- 第一次 run:先收到 `approval_pending` 通知并打印 `pending request: ...`,随后调用 `approvals:grant`。 +- 第一次 run:先收到 `event_type=approval.pending` 通知并打印 `pending request: ...`,随后调用 `approvals:grant`。 - 第一次 run 同时演示 `approvals:poll`(控制面阻塞拉取待审批请求)。 - 第二次 run:`pending_count=0`(自动放行)。 - 撤销规则后第三次 run:再次出现 `pending request after revoke: ...`。 @@ -30,7 +30,7 @@ python main.py - 通道侧:`DirectClientChannel + AgentChannel` - 审批 action: - `approvals:list` - - `approvals:poll` + - `approvals:poll`(可选 `session_id` 过滤) - `approvals:grant` - `approvals:revoke` - 规则持久化路径: diff --git a/examples/07-tool-approval-memory/main.py b/examples/07-tool-approval-memory/main.py index 4ed3aeaa..d0c89f04 100644 --- a/examples/07-tool-approval-memory/main.py +++ b/examples/07-tool-approval-memory/main.py @@ -117,7 +117,10 @@ async def _invoke_action( payload = response.payload if not isinstance(payload, dict): raise RuntimeError(f"unexpected action response payload: {payload!r}") - if payload.get("type") == "error": + event_type = response.event_type + if not isinstance(event_type, str) or not event_type: + raise RuntimeError("invalid action response: missing event_type") + if event_type == "error": raise RuntimeError(f"action failed ({action_id}): {payload.get('reason')}") resp = payload.get("resp") @@ -143,7 +146,8 @@ async def _wait_for_pending_request_id( payload = envelope.payload if not isinstance(payload, dict): continue - if payload.get("type") != "approval_pending": + event_type = envelope.event_type + if event_type != "approval.pending": continue resp = payload.get("resp") if not isinstance(resp, dict): @@ -151,7 +155,7 @@ async def _wait_for_pending_request_id( request = resp.get("request") if isinstance(request, dict) and isinstance(request.get("request_id"), str): return request["request_id"] - raise TimeoutError("approval_pending event was not received in time") + raise TimeoutError("approval.pending event was not received in time") def _new_prompt_envelope(prompt: str) -> TransportEnvelope: @@ -166,7 +170,8 @@ def _extract_run_success(response: TransportEnvelope) -> bool: payload = response.payload if not isinstance(payload, dict): return False - if payload.get("type") == "error": + event_type = response.event_type + if event_type != "result": return False raw_success = payload.get("success") if isinstance(raw_success, bool): diff --git a/examples/10-agentscope-compat-single-agent/cli.py b/examples/10-agentscope-compat-single-agent/cli.py index 61962ca0..01d881bd 100644 --- a/examples/10-agentscope-compat-single-agent/cli.py +++ b/examples/10-agentscope-compat-single-agent/cli.py @@ -100,7 +100,14 @@ def _parse_response(response: TransportEnvelope) -> tuple[bool, str]: payload = response.payload if not isinstance(payload, dict): return (False, _render_output_text(payload)) - if payload.get("type") == "error": + + # Prefer envelope-level event typing, fallback to legacy payload.type for + # compatibility with older transports. + event_type = response.event_type + if not isinstance(event_type, str) or not event_type: + legacy = payload.get("type") + event_type = legacy if isinstance(legacy, str) else None + if event_type == "error": reason = payload.get("reason") or payload.get("error") or "unknown transport error" return (False, str(reason)) @@ -108,7 +115,7 @@ def _parse_response(response: TransportEnvelope) -> tuple[bool, str]: # callers may still rely on top-level fallbacks. resp = payload.get("resp") if isinstance(resp, dict): - success = bool(resp.get("success", payload.get("success", True))) + success = bool(resp.get("success", payload.get("success", payload.get("ok", True)))) output = resp.get("output", payload.get("output")) errors = resp.get("errors", payload.get("errors", [])) text = _render_output_text(output) @@ -116,7 +123,7 @@ def _parse_response(response: TransportEnvelope) -> tuple[bool, str]: text = f"{text}\nerrors={errors}" if text else f"errors={errors}" return (success, text) - success = bool(payload.get("success", payload.get("ok", True))) + success = bool(payload.get("success", payload.get("ok", event_type != "error"))) return (success, _render_output_text(payload.get("output"))) diff --git a/tests/unit/test_agent_event_transport_hook.py b/tests/unit/test_agent_event_transport_hook.py new file mode 100644 index 00000000..ad2fecbf --- /dev/null +++ b/tests/unit/test_agent_event_transport_hook.py @@ -0,0 +1,31 @@ +from __future__ import annotations + +from typing import Any + +import pytest + +from dare_framework.hook._internal.agent_event_transport_hook import AgentEventTransportHook +from dare_framework.hook.types import HookPhase +from dare_framework.transport import TransportEventType + + +class _RecordingTransport: + def __init__(self) -> None: + self.sent: list[Any] = [] + + async def send(self, msg: Any) -> None: + self.sent.append(msg) + + +@pytest.mark.asyncio +async def test_agent_event_transport_hook_sets_event_type() -> None: + transport = _RecordingTransport() + hook = AgentEventTransportHook(transport) + + await hook.invoke(HookPhase.BEFORE_PLAN, payload={"task_id": "task-1"}) + + assert len(transport.sent) == 1 + envelope = transport.sent[0] + assert envelope.event_type == TransportEventType.HOOK.value + assert isinstance(envelope.payload, dict) + assert envelope.payload.get("phase") == HookPhase.BEFORE_PLAN.value diff --git a/tests/unit/test_base_agent_transport_contract.py b/tests/unit/test_base_agent_transport_contract.py index 3ca4fec5..bf6c2375 100644 --- a/tests/unit/test_base_agent_transport_contract.py +++ b/tests/unit/test_base_agent_transport_contract.py @@ -8,7 +8,7 @@ from dare_framework.agent.base_agent import BaseAgent from dare_framework.agent.status import AgentStatus from dare_framework.plan.types import RunResult -from dare_framework.transport import EnvelopeKind, TransportEnvelope +from dare_framework.transport import EnvelopeKind, TransportEnvelope, TransportEventType class _CaptureAgent(BaseAgent): @@ -224,6 +224,7 @@ async def test_transport_loop_accepts_batched_messages_from_poll() -> None: assert agent.seen_tasks == ["first", "second"] assert len(channel._sent) == 2 assert [envelope.reply_to for envelope in channel._sent] == ["m1", "m2"] + assert all(envelope.event_type == TransportEventType.RESULT.value for envelope in channel._sent) @pytest.mark.asyncio @@ -238,12 +239,18 @@ async def test_transport_loop_returns_structured_error_for_invalid_message_paylo error_payloads = [ envelope.payload for envelope in channel._sent - if isinstance(envelope.payload, dict) and envelope.payload.get("type") == "error" + if getattr(envelope, "event_type", None) == TransportEventType.ERROR.value and isinstance(envelope.payload, dict) ] assert len(error_payloads) == 1 payload = error_payloads[0] assert payload.get("code") == "INVALID_MESSAGE_PAYLOAD" assert payload.get("kind") == "message" + error_envelope = next( + envelope + for envelope in channel._sent + if getattr(envelope, "event_type", None) == TransportEventType.ERROR.value and isinstance(envelope.payload, dict) + ) + assert error_envelope.event_type == TransportEventType.ERROR.value @pytest.mark.asyncio @@ -254,12 +261,15 @@ async def test_transport_loop_returns_structured_error_when_execute_raises() -> await agent._run_transport_loop() - error_payloads = [ - envelope.payload + error_envelopes = [ + envelope for envelope in channel._sent - if isinstance(envelope.payload, dict) and envelope.payload.get("type") == "error" + if getattr(envelope, "event_type", None) == TransportEventType.ERROR.value and isinstance(envelope.payload, dict) ] - assert len(error_payloads) == 1 - payload = error_payloads[0] + assert len(error_envelopes) == 1 + error_envelope = error_envelopes[0] + assert error_envelope.event_type == TransportEventType.ERROR.value + payload = error_envelope.payload + assert isinstance(payload, dict) assert payload.get("code") == "AGENT_EXECUTION_FAILED" assert "simulated model timeout" in str(payload.get("reason")) diff --git a/tests/unit/test_dare_agent_hook_transport_boundary.py b/tests/unit/test_dare_agent_hook_transport_boundary.py index 606fb760..37a9d0e2 100644 --- a/tests/unit/test_dare_agent_hook_transport_boundary.py +++ b/tests/unit/test_dare_agent_hook_transport_boundary.py @@ -10,7 +10,7 @@ from dare_framework.context import Context from dare_framework.model.types import ModelInput, ModelResponse from dare_framework.tool.types import ToolResult -from dare_framework.transport import AgentChannel, TransportEnvelope +from dare_framework.transport import AgentChannel, TransportEnvelope, TransportEventType class _Model: @@ -86,11 +86,24 @@ async def test_dare_agent_does_not_emit_hook_messages_without_transport_hook() - hook_payloads = [ envelope.payload for envelope in transport.sent - if isinstance(envelope.payload, dict) and envelope.payload.get("type") == "hook" + if getattr(envelope, "event_type", None) == TransportEventType.HOOK.value ] assert hook_payloads == [] +def test_dare_agent_does_not_expose_transport_payload_helper() -> None: + agent = DareAgent( + name="transport-explicitness", + model=_Model(), + context=Context(config=Config()), + tool_gateway=_ToolGateway(), + ) + + # Transport payload emission for approval flow now lives at tool-gateway + # boundary; agent should not expose the legacy helper. + assert not hasattr(agent, "_send_transport_payload") + + @pytest.mark.asyncio async def test_dare_builder_registers_agent_event_transport_hook_when_channel_present() -> None: channel = _RecordingChannel() diff --git a/tests/unit/test_example_10_agentscope_compat.py b/tests/unit/test_example_10_agentscope_compat.py index c2940e22..52da48a1 100644 --- a/tests/unit/test_example_10_agentscope_compat.py +++ b/tests/unit/test_example_10_agentscope_compat.py @@ -154,6 +154,27 @@ def _new_scripted_model() -> _DeterministicCompatTestModel: return _DeterministicCompatTestModel() +def test_example_10_cli_parse_response_uses_event_type_for_errors() -> None: + cli_module = _load_example_cli_module() + success, text = cli_module._parse_response( # type: ignore[attr-defined] + TransportEnvelope( + id="resp-error", + kind=EnvelopeKind.MESSAGE, + event_type="error", + payload={ + "kind": "message", + "target": "agent", + "ok": False, + "reason": "boom", + "error": "boom", + "resp": {"code": "runtime_error", "reason": "boom"}, + }, + ) + ) + assert success is False + assert "boom" in text + + class _DeterministicSimpleLoopModel(IModelAdapter): def __init__(self) -> None: self.called_tools: list[str] = [] @@ -365,8 +386,8 @@ async def test_single_agent_demo_transport_message_loop(tmp_path: Path) -> None: await bundle.agent.stop() assert isinstance(response.payload, dict) + assert response.event_type == "result" payload = response.payload - assert payload.get("type") == "result" resp = payload.get("resp") assert isinstance(resp, dict) assert resp.get("success") is True diff --git a/tests/unit/test_examples_cli.py b/tests/unit/test_examples_cli.py index ddfde28b..0395dc8c 100644 --- a/tests/unit/test_examples_cli.py +++ b/tests/unit/test_examples_cli.py @@ -4,15 +4,19 @@ import contextlib from pathlib import Path import sys +from typing import Any import pytest +from dare_framework.tool.action_handler import ApprovalsActionHandler from dare_framework.tool._internal.control.approval_manager import ( ApprovalDecision, ApprovalEvaluationStatus, JsonApprovalRuleStore, ToolApprovalManager, ) +from dare_framework.transport import EnvelopeKind, TransportEnvelope +from dare_framework.transport.interaction.resource_action import ResourceAction def _resolve_example_dir() -> Path: @@ -99,6 +103,67 @@ def show_mode(self, _mode) -> None: return +class _HandlerBackedApprovalClient: + """Minimal ask-capable client shim backed by the real approvals action handler.""" + + def __init__(self, manager: ToolApprovalManager) -> None: + self._handler = ApprovalsActionHandler(manager) + + async def ask(self, req: TransportEnvelope, timeout: float = 30.0) -> TransportEnvelope: + _ = timeout + action = ResourceAction(str(req.payload)) + result = await self._handler.invoke(action, **dict(req.meta)) + return TransportEnvelope( + id=f"resp-{req.id}", + reply_to=req.id, + kind=EnvelopeKind.MESSAGE, + event_type="result", + payload={"resp": {"result": result}}, + ) + + +class _CaptureApprovalClient: + def __init__(self) -> None: + self.last_meta: dict[str, Any] | None = None + + async def ask(self, req: TransportEnvelope, timeout: float = 30.0) -> TransportEnvelope: + _ = timeout + self.last_meta = dict(req.meta) + return TransportEnvelope( + id=f"resp-{req.id}", + reply_to=req.id, + kind=EnvelopeKind.MESSAGE, + event_type="result", + payload={"resp": {"result": {"request": None}}}, + ) + + +class _CaptureTimeoutApprovalClient: + def __init__(self) -> None: + self.last_timeout: float | None = None + + async def ask(self, req: TransportEnvelope, timeout: float = 30.0) -> TransportEnvelope: + _ = req + self.last_timeout = timeout + return TransportEnvelope( + id="resp-timeout", + kind=EnvelopeKind.MESSAGE, + event_type="result", + payload={"resp": {"result": {"request": None}}}, + ) + + +class _MissingEventTypeApprovalClient: + async def ask(self, req: TransportEnvelope, timeout: float = 30.0) -> TransportEnvelope: + _ = timeout + return TransportEnvelope( + id=f"resp-{req.id}", + reply_to=req.id, + kind=EnvelopeKind.MESSAGE, + payload={"resp": {"result": {"request": None}}}, + ) + + @pytest.mark.asyncio async def test_handle_approvals_command_list_and_grant(tmp_path: Path) -> None: manager = ToolApprovalManager( @@ -115,21 +180,19 @@ async def test_handle_approvals_command_list_and_grant(tmp_path: Path) -> None: assert evaluation.request is not None request_id = evaluation.request.request_id - class _FakeAgent: - _approval_manager = manager - display = _CaptureDisplay() + approval_client = _HandlerBackedApprovalClient(manager) await cli._handle_approvals_command( # type: ignore[attr-defined] ["list"], - agent=_FakeAgent(), + approval_client=approval_client, display=display, ) assert any("pending" in msg for level, msg in display.messages if level == "info") await cli._handle_approvals_command( # type: ignore[attr-defined] ["poll", "timeout_ms=10"], - agent=_FakeAgent(), + approval_client=approval_client, display=display, ) assert any("pending request:" in msg for level, msg in display.messages if level == "info") @@ -137,12 +200,50 @@ class _FakeAgent: wait_task = asyncio.create_task(manager.wait_for_resolution(request_id)) await cli._handle_approvals_command( # type: ignore[attr-defined] ["grant", request_id, "scope=workspace", "matcher=exact_params"], - agent=_FakeAgent(), + approval_client=approval_client, display=display, ) assert await wait_task == ApprovalDecision.ALLOW +@pytest.mark.asyncio +async def test_handle_approvals_poll_forwards_session_filter() -> None: + display = _CaptureDisplay() + approval_client = _CaptureApprovalClient() + await cli._handle_approvals_command( # type: ignore[attr-defined] + ["poll", "timeout_ms=10", "session_id=session-42"], + approval_client=approval_client, + display=display, + ) + assert approval_client.last_meta is not None + assert approval_client.last_meta.get("session_id") == "session-42" + + +@pytest.mark.asyncio +async def test_handle_approvals_poll_uses_user_timeout_for_transport_wait() -> None: + display = _CaptureDisplay() + approval_client = _CaptureTimeoutApprovalClient() + await cli._handle_approvals_command( # type: ignore[attr-defined] + ["poll", "timeout_seconds=60"], + approval_client=approval_client, + display=display, + ) + assert approval_client.last_timeout is not None + assert approval_client.last_timeout >= 60.0 + + +@pytest.mark.asyncio +async def test_handle_approvals_command_requires_event_type_in_response() -> None: + display = _CaptureDisplay() + approval_client = _MissingEventTypeApprovalClient() + await cli._handle_approvals_command( # type: ignore[attr-defined] + ["list"], + approval_client=approval_client, + display=display, + ) + assert any("missing event_type" in msg for level, msg in display.messages if level == "error") + + @pytest.mark.asyncio async def test_run_cli_loop_background_execute_allows_followup_commands(monkeypatch) -> None: started = asyncio.Event() diff --git a/tests/unit/test_examples_cli_mcp.py b/tests/unit/test_examples_cli_mcp.py index a98d4b54..553101a9 100644 --- a/tests/unit/test_examples_cli_mcp.py +++ b/tests/unit/test_examples_cli_mcp.py @@ -4,15 +4,19 @@ import importlib.util from pathlib import Path import sys +from typing import Any import pytest +from dare_framework.tool.action_handler import ApprovalsActionHandler from dare_framework.tool._internal.control.approval_manager import ( ApprovalDecision, ApprovalEvaluationStatus, JsonApprovalRuleStore, ToolApprovalManager, ) +from dare_framework.transport import EnvelopeKind, TransportEnvelope +from dare_framework.transport.interaction.resource_action import ResourceAction def _load_cli_module(module_name: str, relative_cli_path: str): @@ -64,6 +68,67 @@ def show_mode(self, _mode) -> None: return +class _HandlerBackedApprovalClient: + """Minimal ask-capable client shim backed by the real approvals action handler.""" + + def __init__(self, manager: ToolApprovalManager) -> None: + self._handler = ApprovalsActionHandler(manager) + + async def ask(self, req: TransportEnvelope, timeout: float = 30.0) -> TransportEnvelope: + _ = timeout + action = ResourceAction(str(req.payload)) + result = await self._handler.invoke(action, **dict(req.meta)) + return TransportEnvelope( + id=f"resp-{req.id}", + reply_to=req.id, + kind=EnvelopeKind.MESSAGE, + event_type="result", + payload={"resp": {"result": result}}, + ) + + +class _CaptureApprovalClient: + def __init__(self) -> None: + self.last_meta: dict[str, Any] | None = None + + async def ask(self, req: TransportEnvelope, timeout: float = 30.0) -> TransportEnvelope: + _ = timeout + self.last_meta = dict(req.meta) + return TransportEnvelope( + id=f"resp-{req.id}", + reply_to=req.id, + kind=EnvelopeKind.MESSAGE, + event_type="result", + payload={"resp": {"result": {"request": None}}}, + ) + + +class _CaptureTimeoutApprovalClient: + def __init__(self) -> None: + self.last_timeout: float | None = None + + async def ask(self, req: TransportEnvelope, timeout: float = 30.0) -> TransportEnvelope: + _ = req + self.last_timeout = timeout + return TransportEnvelope( + id="resp-timeout", + kind=EnvelopeKind.MESSAGE, + event_type="result", + payload={"resp": {"result": {"request": None}}}, + ) + + +class _MissingEventTypeApprovalClient: + async def ask(self, req: TransportEnvelope, timeout: float = 30.0) -> TransportEnvelope: + _ = timeout + return TransportEnvelope( + id=f"resp-{req.id}", + reply_to=req.id, + kind=EnvelopeKind.MESSAGE, + payload={"resp": {"result": {"request": None}}}, + ) + + @pytest.mark.asyncio async def test_handle_approvals_command_list_and_grant_mcp_cli(tmp_path: Path) -> None: cli_mcp = _load_cli_module( @@ -84,20 +149,18 @@ async def test_handle_approvals_command_list_and_grant_mcp_cli(tmp_path: Path) - assert evaluation.request is not None request_id = evaluation.request.request_id - class _FakeAgent: - _approval_manager = manager - display = _CaptureDisplay() + approval_client = _HandlerBackedApprovalClient(manager) await cli_mcp._handle_approvals_command( # type: ignore[attr-defined] ["list"], - agent=_FakeAgent(), + approval_client=approval_client, display=display, ) assert any("pending" in msg for level, msg in display.messages if level == "info") await cli_mcp._handle_approvals_command( # type: ignore[attr-defined] ["poll", "timeout_ms=10"], - agent=_FakeAgent(), + approval_client=approval_client, display=display, ) assert any("pending request:" in msg for level, msg in display.messages if level == "info") @@ -105,12 +168,62 @@ class _FakeAgent: wait_task = asyncio.create_task(manager.wait_for_resolution(request_id)) await cli_mcp._handle_approvals_command( # type: ignore[attr-defined] ["grant", request_id, "scope=workspace", "matcher=exact_params"], - agent=_FakeAgent(), + approval_client=approval_client, display=display, ) assert await wait_task == ApprovalDecision.ALLOW +@pytest.mark.asyncio +async def test_handle_approvals_poll_forwards_session_filter_mcp_cli() -> None: + cli_mcp = _load_cli_module( + "examples_06_cli_poll_filter", + "examples/06-dare-coding-agent-mcp/cli.py", + ) + display = _CaptureDisplay() + approval_client = _CaptureApprovalClient() + await cli_mcp._handle_approvals_command( # type: ignore[attr-defined] + ["poll", "timeout_ms=10", "session_id=session-42"], + approval_client=approval_client, + display=display, + ) + assert approval_client.last_meta is not None + assert approval_client.last_meta.get("session_id") == "session-42" + + +@pytest.mark.asyncio +async def test_handle_approvals_poll_uses_user_timeout_for_transport_wait_mcp_cli() -> None: + cli_mcp = _load_cli_module( + "examples_06_cli_poll_timeout", + "examples/06-dare-coding-agent-mcp/cli.py", + ) + display = _CaptureDisplay() + approval_client = _CaptureTimeoutApprovalClient() + await cli_mcp._handle_approvals_command( # type: ignore[attr-defined] + ["poll", "timeout_seconds=60"], + approval_client=approval_client, + display=display, + ) + assert approval_client.last_timeout is not None + assert approval_client.last_timeout >= 60.0 + + +@pytest.mark.asyncio +async def test_handle_approvals_command_requires_event_type_in_response_mcp_cli() -> None: + cli_mcp = _load_cli_module( + "examples_06_cli_missing_event_type", + "examples/06-dare-coding-agent-mcp/cli.py", + ) + display = _CaptureDisplay() + approval_client = _MissingEventTypeApprovalClient() + await cli_mcp._handle_approvals_command( # type: ignore[attr-defined] + ["list"], + approval_client=approval_client, + display=display, + ) + assert any("missing event_type" in msg for level, msg in display.messages if level == "error") + + @pytest.mark.asyncio async def test_run_cli_loop_background_execute_allows_followup_commands_mcp_cli(monkeypatch) -> None: cli_mcp = _load_cli_module( diff --git a/tests/unit/test_five_layer_agent.py b/tests/unit/test_five_layer_agent.py index c121d0d7..24abd366 100644 --- a/tests/unit/test_five_layer_agent.py +++ b/tests/unit/test_five_layer_agent.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio +import json from dataclasses import dataclass from typing import Any from unittest.mock import AsyncMock, MagicMock @@ -21,6 +22,7 @@ ToolApprovalManager, ) from dare_framework.tool.types import CapabilityDescriptor, CapabilityKind, CapabilityType +from dare_framework.transport.types import EnvelopeKind # ============================================================================= @@ -517,11 +519,11 @@ async def test_no_planner_emits_transport_approval_pending_message(self, tmp_pat request_id: str | None = None for _ in range(100): for envelope in transport.sent: + if getattr(envelope, "event_type", None) != "approval.pending": + continue payload = getattr(envelope, "payload", None) if not isinstance(payload, dict): continue - if payload.get("type") != "approval_pending": - continue resp = payload.get("resp") if not isinstance(resp, dict): continue @@ -533,16 +535,305 @@ async def test_no_planner_emits_transport_approval_pending_message(self, tmp_pat break await asyncio.sleep(0.01) assert request_id is not None + assert any( + getattr(envelope, "event_type", None) == "approval.pending" + and getattr(envelope, "kind", None) == EnvelopeKind.SELECT + for envelope in transport.sent + ) + + await approval_manager.grant( + request_id, + scope=ApprovalScope.ONCE, + matcher=ApprovalMatcherKind.EXACT_PARAMS, + ) + + result = await run_task + assert result.success is True + + @pytest.mark.asyncio + async def test_no_planner_denied_approval_reports_not_allow_without_resolved_event(self, tmp_path) -> None: + capability = CapabilityDescriptor( + id="run_command", + type=CapabilityType.TOOL, + name="run_command", + description="Run a shell command.", + input_schema={"type": "object", "properties": {"command": {"type": "string"}}}, + metadata={"requires_approval": True}, + ) + tool_gateway = MockToolGateway([capability]) + approval_manager = ToolApprovalManager( + workspace_store=JsonApprovalRuleStore(tmp_path / "workspace" / "approvals.json"), + user_store=JsonApprovalRuleStore(tmp_path / "user" / "approvals.json"), + ) + model = MockModelAdapter( + [ + ModelResponse( + content="Running command...", + tool_calls=[{"name": "run_command", "arguments": {"command": "git status --short"}}], + ), + ModelResponse(content="Denied.", tool_calls=[]), + ] + ) + agent = _make_agent( + name="react-agent-approval-denied", + model=model, + tool_gateway=tool_gateway, + approval_manager=approval_manager, + ) + transport = RecordingTransportChannel() + + run_task = asyncio.create_task(agent("Run git status", transport=transport)) + request_id: str | None = None + for _ in range(100): + pending = approval_manager.list_pending() + if pending: + request_id = pending[0].request_id + break + await asyncio.sleep(0.01) + assert request_id is not None + + await approval_manager.deny( + request_id, + scope=ApprovalScope.ONCE, + matcher=ApprovalMatcherKind.EXACT_PARAMS, + ) + + await run_task + + tool_messages = [msg for msg in agent._context.stm_get() if msg.role == "tool"] + assert tool_messages + tool_payload = json.loads(tool_messages[-1].content) + assert tool_payload.get("status") == "not_allow" + assert tool_payload.get("success") is False + + event_types = [getattr(envelope, "event_type", None) for envelope in transport.sent] + assert "approval.pending" in event_types + assert "approval.resolved" not in event_types + + @pytest.mark.asyncio + async def test_no_planner_emits_approval_lifecycle_events_for_event_log_auto_resolution(self, tmp_path) -> None: + capability = CapabilityDescriptor( + id="run_command", + type=CapabilityType.TOOL, + name="run_command", + description="Run shell command", + input_schema={"type": "object", "properties": {"command": {"type": "string"}}}, + metadata={"requires_approval": True}, + ) + tool_gateway = MockToolGateway([capability]) + approval_manager = ToolApprovalManager( + workspace_store=JsonApprovalRuleStore(tmp_path / "workspace" / "approvals.json"), + user_store=JsonApprovalRuleStore(tmp_path / "user" / "approvals.json"), + ) + + class AutoApproveEventLog(MockEventLog): + async def append(self, event_type: str, payload: dict[str, Any]) -> None: + await super().append(event_type, payload) + if event_type != "exec.waiting_human": + return + checkpoint_id = payload.get("checkpoint_id") + if isinstance(checkpoint_id, str) and checkpoint_id: + await approval_manager.grant( + checkpoint_id, + scope=ApprovalScope.ONCE, + matcher=ApprovalMatcherKind.EXACT_PARAMS, + ) + + event_log = AutoApproveEventLog() + model = MockModelAdapter( + [ + ModelResponse( + content="Running command...", + tool_calls=[{"name": "run_command", "arguments": {"command": "git status --short"}}], + ), + ModelResponse(content="Approved and finished.", tool_calls=[]), + ] + ) + agent = _make_agent( + name="react-agent-approval-lifecycle-events", + model=model, + tool_gateway=tool_gateway, + approval_manager=approval_manager, + event_log=event_log, + ) + + result = await asyncio.wait_for(agent("Run git status"), timeout=1.0) + assert result.success is True + + event_types = [event_type for event_type, _ in event_log.events] + assert "exec.waiting_human" in event_types + assert "exec.resume" in event_types + assert "tool.approval" in event_types + + @pytest.mark.asyncio + async def test_no_planner_tool_params_session_id_does_not_collide_with_governance(self, tmp_path) -> None: + capability = CapabilityDescriptor( + id="run_command", + type=CapabilityType.TOOL, + name="run_command", + description="Run shell command", + input_schema={ + "type": "object", + "properties": { + "command": {"type": "string"}, + "session_id": {"type": "string"}, + }, + }, + metadata={"requires_approval": True}, + ) + tool_gateway = MockToolGateway([capability]) + approval_manager = ToolApprovalManager( + workspace_store=JsonApprovalRuleStore(tmp_path / "workspace" / "approvals.json"), + user_store=JsonApprovalRuleStore(tmp_path / "user" / "approvals.json"), + ) + model = MockModelAdapter( + [ + ModelResponse( + content="Run command with tool-level session_id argument.", + tool_calls=[ + { + "name": "run_command", + "arguments": { + "command": "git status --short", + "session_id": "tool-arg-session", + }, + } + ], + ), + ModelResponse(content="Done.", tool_calls=[]), + ] + ) + agent = _make_agent( + name="react-agent-tool-param-session-id", + model=model, + tool_gateway=tool_gateway, + approval_manager=approval_manager, + ) + + run_task = asyncio.create_task(agent("Run command with tool arg session_id")) + request_id: str | None = None + for _ in range(100): + pending = approval_manager.list_pending() + if pending: + request_id = pending[0].request_id + break + await asyncio.sleep(0.01) + assert request_id is not None await approval_manager.grant( request_id, scope=ApprovalScope.ONCE, matcher=ApprovalMatcherKind.EXACT_PARAMS, ) + result = await run_task + assert result.success is True + + assert tool_gateway.invoke_calls + _capability_id, params, _envelope = tool_gateway.invoke_calls[-1] + assert params.get("session_id") == "tool-arg-session" + + @pytest.mark.asyncio + async def test_no_planner_tool_params_context_does_not_collide_with_runtime_context(self, tmp_path) -> None: + capability = CapabilityDescriptor( + id="run_command", + type=CapabilityType.TOOL, + name="run_command", + description="Run shell command", + input_schema={ + "type": "object", + "properties": { + "command": {"type": "string"}, + "context": {"type": "string"}, + }, + }, + metadata={"requires_approval": True}, + ) + tool_gateway = MockToolGateway([capability]) + approval_manager = ToolApprovalManager( + workspace_store=JsonApprovalRuleStore(tmp_path / "workspace" / "approvals.json"), + user_store=JsonApprovalRuleStore(tmp_path / "user" / "approvals.json"), + ) + model = MockModelAdapter( + [ + ModelResponse( + content="Run command with tool-level context argument.", + tool_calls=[ + { + "name": "run_command", + "arguments": { + "command": "git status --short", + "context": "tool-arg-context", + }, + } + ], + ), + ModelResponse(content="Done.", tool_calls=[]), + ] + ) + agent = _make_agent( + name="react-agent-tool-param-context", + model=model, + tool_gateway=tool_gateway, + approval_manager=approval_manager, + ) + + run_task = asyncio.create_task(agent("Run command with tool arg context")) + request_id: str | None = None + for _ in range(100): + pending = approval_manager.list_pending() + if pending: + request_id = pending[0].request_id + break + await asyncio.sleep(0.01) + assert request_id is not None + await approval_manager.grant( + request_id, + scope=ApprovalScope.ONCE, + matcher=ApprovalMatcherKind.EXACT_PARAMS, + ) result = await run_task assert result.success is True + assert tool_gateway.invoke_calls + _capability_id, params, _envelope = tool_gateway.invoke_calls[-1] + assert params.get("context") == "tool-arg-context" + + @pytest.mark.asyncio + async def test_no_planner_missing_approval_manager_reports_fail_status(self) -> None: + capability = CapabilityDescriptor( + id="run_command", + type=CapabilityType.TOOL, + name="run_command", + description="Run shell command", + input_schema={"type": "object", "properties": {"command": {"type": "string"}}}, + metadata={"requires_approval": True}, + ) + tool_gateway = MockToolGateway([capability]) + model = MockModelAdapter( + [ + ModelResponse( + content="Try running command without approval manager.", + tool_calls=[{"name": "run_command", "arguments": {"command": "git status --short"}}], + ), + ModelResponse(content="Done.", tool_calls=[]), + ] + ) + agent = _make_agent( + name="react-agent-missing-approval-manager", + model=model, + tool_gateway=tool_gateway, + approval_manager=None, + ) + + await agent("Run command") + tool_messages = [msg for msg in agent._context.stm_get() if msg.role == "tool"] + assert tool_messages + tool_payload = json.loads(tool_messages[-1].content) + assert tool_payload.get("success") is False + assert tool_payload.get("status") == "fail" + assert "no approval manager" in str(tool_payload.get("error", "")) + # ============================================================================= # Tests: Full Five-Layer Mode diff --git a/tests/unit/test_governed_tool_gateway.py b/tests/unit/test_governed_tool_gateway.py new file mode 100644 index 00000000..3fdf6770 --- /dev/null +++ b/tests/unit/test_governed_tool_gateway.py @@ -0,0 +1,95 @@ +from __future__ import annotations + +from typing import Any + +import pytest + +from dare_framework.plan.types import Envelope +from dare_framework.tool._internal.control.approval_manager import ( + ApprovalEvaluation, + ApprovalEvaluationStatus, +) +from dare_framework.tool._internal.governed_tool_gateway import ( + ApprovalInvokeContext, + GovernedToolGateway, +) +from dare_framework.tool.types import CapabilityDescriptor, CapabilityType, ToolResult + + +class _RecordingDelegateGateway: + def __init__(self, capability: CapabilityDescriptor) -> None: + self._capability = capability + self.invoke_calls: list[dict[str, Any]] = [] + + def list_capabilities(self) -> list[CapabilityDescriptor]: + return [self._capability] + + async def invoke(self, capability_id: str, *, envelope: Envelope, **params: Any) -> ToolResult: + self.invoke_calls.append( + { + "capability_id": capability_id, + "envelope": envelope, + "params": dict(params), + } + ) + return ToolResult(success=True, output={"ok": True}) + + +class _RecordingApprovalManager: + def __init__(self) -> None: + self.evaluate_calls: list[dict[str, Any]] = [] + + async def evaluate( + self, + *, + capability_id: str, + params: dict[str, Any], + session_id: str | None, + reason: str, + ) -> ApprovalEvaluation: + self.evaluate_calls.append( + { + "capability_id": capability_id, + "params": dict(params), + "session_id": session_id, + "reason": reason, + } + ) + return ApprovalEvaluation(status=ApprovalEvaluationStatus.ALLOW) + + +@pytest.mark.asyncio +async def test_governed_gateway_approval_uses_effective_params_with_context_collision() -> None: + capability = CapabilityDescriptor( + id="run_command", + type=CapabilityType.TOOL, + name="run_command", + description="run command", + input_schema={"type": "object", "properties": {}}, + metadata={"requires_approval": True}, + ) + delegate = _RecordingDelegateGateway(capability) + approval_manager = _RecordingApprovalManager() + gateway = GovernedToolGateway(delegate, approval_manager=approval_manager) + + runtime_context = object() + envelope = Envelope() + result = await gateway.invoke( + capability.id, + approval=ApprovalInvokeContext(runtime_context=runtime_context), + envelope=envelope, + command="echo hello", + context="tool-arg-context", + ) + + assert result.success is True + assert approval_manager.evaluate_calls + assert approval_manager.evaluate_calls[0]["params"] == { + "command": "echo hello", + "context": "tool-arg-context", + } + + assert delegate.invoke_calls + delegate_params = delegate.invoke_calls[0]["params"] + assert delegate_params["command"] == "echo hello" + assert delegate_params["context"] == "tool-arg-context" diff --git a/tests/unit/test_tool_approval_action_handler.py b/tests/unit/test_tool_approval_action_handler.py index 36526ee2..f9dd5f21 100644 --- a/tests/unit/test_tool_approval_action_handler.py +++ b/tests/unit/test_tool_approval_action_handler.py @@ -102,3 +102,52 @@ async def test_approvals_action_handler_poll_timeout_returns_null_request(manage polled = await handler.invoke(ResourceAction.APPROVALS_POLL, timeout_seconds=0.05) assert polled["request"] is None + + +@pytest.mark.asyncio +async def test_approvals_action_handler_poll_filters_by_session_id(manager) -> None: + handler = ApprovalsActionHandler(manager) + first = await manager.evaluate( + capability_id="run_command", + params={"command": "echo session-a"}, + session_id="session-a", + reason="Tool run_command requires approval", + ) + second = await manager.evaluate( + capability_id="run_command", + params={"command": "echo session-b"}, + session_id="session-b", + reason="Tool run_command requires approval", + ) + assert first.request is not None + assert second.request is not None + + polled = await handler.invoke(ResourceAction.APPROVALS_POLL, session_id="session-b") + request = polled["request"] + assert isinstance(request, dict) + assert request["request_id"] == second.request.request_id + + +@pytest.mark.asyncio +async def test_approvals_action_handler_poll_matches_deduplicated_pending_across_sessions(manager) -> None: + handler = ApprovalsActionHandler(manager) + first = await manager.evaluate( + capability_id="run_command", + params={"command": "echo same-command"}, + session_id="session-a", + reason="Tool run_command requires approval", + ) + second = await manager.evaluate( + capability_id="run_command", + params={"command": "echo same-command"}, + session_id="session-b", + reason="Tool run_command requires approval", + ) + assert first.request is not None + assert second.request is not None + assert first.request.request_id == second.request.request_id + + polled = await handler.invoke(ResourceAction.APPROVALS_POLL, session_id="session-b", timeout_seconds=0.05) + request = polled["request"] + assert isinstance(request, dict) + assert request["request_id"] == first.request.request_id diff --git a/tests/unit/test_tool_approval_manager.py b/tests/unit/test_tool_approval_manager.py index d6af1b0c..839829ad 100644 --- a/tests/unit/test_tool_approval_manager.py +++ b/tests/unit/test_tool_approval_manager.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +import contextlib import json import pytest @@ -166,6 +167,95 @@ async def test_poll_pending_returns_next_request_when_available(manager) -> None assert polled.request_id == first.request.request_id +@pytest.mark.asyncio +async def test_poll_pending_filters_by_session_id(manager) -> None: + first = await manager.evaluate( + capability_id="run_command", + params={"command": "echo session-a"}, + session_id="session-a", + reason="Tool run_command requires approval", + ) + second = await manager.evaluate( + capability_id="run_command", + params={"command": "echo session-b"}, + session_id="session-b", + reason="Tool run_command requires approval", + ) + assert first.request is not None + assert second.request is not None + + polled = await manager.poll_pending(session_id="session-b") + assert polled is not None + assert polled.request_id == second.request.request_id + + +@pytest.mark.asyncio +async def test_poll_pending_session_filter_matches_deduplicated_pending_request(manager) -> None: + first = await manager.evaluate( + capability_id="run_command", + params={"command": "echo same-command"}, + session_id="session-a", + reason="Tool run_command requires approval", + ) + second = await manager.evaluate( + capability_id="run_command", + params={"command": "echo same-command"}, + session_id="session-b", + reason="Tool run_command requires approval", + ) + assert first.request is not None + assert second.request is not None + # Dedup keeps one pending request id for both sessions. + assert second.request.request_id == first.request.request_id + + polled = await manager.poll_pending(session_id="session-b", timeout_seconds=0.05) + assert polled is not None + assert polled.request_id == first.request.request_id + + +@pytest.mark.asyncio +async def test_poll_pending_session_filter_timeout_returns_none(manager) -> None: + await manager.evaluate( + capability_id="run_command", + params={"command": "echo session-a"}, + session_id="session-a", + reason="Tool run_command requires approval", + ) + + polled = await manager.poll_pending(session_id="session-miss", timeout_seconds=0.05) + assert polled is None + + +@pytest.mark.asyncio +async def test_poll_pending_session_filter_waiter_does_not_busy_loop_on_unrelated_pending(manager, monkeypatch) -> None: + await manager.evaluate( + capability_id="run_command", + params={"command": "echo session-a"}, + session_id="session-a", + reason="Tool run_command requires approval", + ) + + call_count = 0 + original = manager._oldest_pending_locked + + def tracked_oldest_pending(*, session_id: str | None = None): + nonlocal call_count + call_count += 1 + return original(session_id=session_id) + + monkeypatch.setattr(manager, "_oldest_pending_locked", tracked_oldest_pending) + waiter = asyncio.create_task(manager.poll_pending(session_id="session-miss")) + await asyncio.sleep(0.05) + + assert waiter.done() is False + # A healthy waiter should sleep on condition changes instead of tight-loop polling. + assert call_count < 30 + + waiter.cancel() + with contextlib.suppress(asyncio.CancelledError): + await waiter + + @pytest.mark.asyncio async def test_poll_pending_waits_until_request_arrives(manager) -> None: waiter = asyncio.create_task(manager.poll_pending(timeout_seconds=1.0)) @@ -184,6 +274,33 @@ async def test_poll_pending_waits_until_request_arrives(manager) -> None: assert polled.request_id == first.request.request_id +@pytest.mark.asyncio +async def test_poll_pending_session_waiter_wakes_when_dedup_adds_matching_session(manager) -> None: + first = await manager.evaluate( + capability_id="run_command", + params={"command": "echo dedup-shared"}, + session_id="session-a", + reason="Tool run_command requires approval", + ) + assert first.request is not None + + waiter = asyncio.create_task(manager.poll_pending(session_id="session-b", timeout_seconds=1.0)) + await asyncio.sleep(0.05) + + second = await manager.evaluate( + capability_id="run_command", + params={"command": "echo dedup-shared"}, + session_id="session-b", + reason="Tool run_command requires approval", + ) + assert second.request is not None + assert second.request.request_id == first.request.request_id + + polled = await asyncio.wait_for(waiter, timeout=0.2) + assert polled is not None + assert polled.request_id == first.request.request_id + + @pytest.mark.asyncio async def test_poll_pending_timeout_returns_none(manager) -> None: polled = await manager.poll_pending(timeout_seconds=0.05) diff --git a/tests/unit/test_tool_signature_contract.py b/tests/unit/test_tool_signature_contract.py index 6c5f0d9a..f7095016 100644 --- a/tests/unit/test_tool_signature_contract.py +++ b/tests/unit/test_tool_signature_contract.py @@ -162,3 +162,23 @@ async def test_gateway_invokes_tool_with_keyword_arguments() -> None: assert result.success is True assert tool.captured["message"] == "hello" assert isinstance(tool.captured["run_context"], RunContext) + + +@pytest.mark.asyncio +async def test_gateway_ignores_untrusted_runtime_context_override_param() -> None: + manager = ToolManager(load_entrypoints=False) + tool = _KeywordProbeTool() + descriptor = manager.register_tool(tool) + gateway = ToolGateway(manager) + + result = await gateway.invoke( + descriptor.id, + envelope=Envelope(allowed_capability_ids=[descriptor.id]), + message="hello", + __dare_runtime_context__={"spoofed": True}, + ) + + assert result.success is True + assert tool.captured["message"] == "hello" + assert isinstance(tool.captured["run_context"], RunContext) + assert tool.captured["run_context"].deps is None diff --git a/tests/unit/test_transport_adapters.py b/tests/unit/test_transport_adapters.py index cf49dc1c..cbc15fb3 100644 --- a/tests/unit/test_transport_adapters.py +++ b/tests/unit/test_transport_adapters.py @@ -1,9 +1,11 @@ import asyncio +import json import pytest +from dare_framework.transport import PollableClientChannel from dare_framework.transport._internal.adapters import DirectClientChannel, StdioClientChannel, WebSocketClientChannel -from dare_framework.transport.types import EnvelopeKind, TransportEnvelope +from dare_framework.transport.types import EnvelopeKind, TransportEnvelope, TransportEventType @pytest.mark.asyncio @@ -34,6 +36,14 @@ async def send(self, _msg) -> None: # pragma: no cover - not used in this test return None +class _CaptureWS: + def __init__(self) -> None: + self.sent: list[str] = [] + + async def send(self, msg: str) -> None: + self.sent.append(msg) + + @pytest.mark.asyncio async def test_websocket_requires_explicit_kind() -> None: ws = WebSocketClientChannel(_DummyWS()) @@ -47,6 +57,24 @@ async def sender(_msg: TransportEnvelope) -> None: await ws.handle_ws_message({"id": "req-1", "payload": "hello"}) +@pytest.mark.asyncio +async def test_websocket_serializer_includes_event_type() -> None: + ws = _CaptureWS() + channel = WebSocketClientChannel(ws) + receiver = channel.agent_envelope_receiver() + await receiver( + TransportEnvelope( + id="evt-1", + kind=EnvelopeKind.MESSAGE, + event_type=TransportEventType.RESULT.value, + payload={"ok": True}, + ) + ) + assert ws.sent + data = json.loads(ws.sent[0]) + assert data["event_type"] == TransportEventType.RESULT.value + + @pytest.mark.asyncio async def test_direct_client_channel_poll_receives_unmatched_agent_messages() -> None: channel = DirectClientChannel() @@ -61,15 +89,15 @@ async def sender(msg: TransportEnvelope) -> None: TransportEnvelope( id="event-1", kind=EnvelopeKind.MESSAGE, - payload={"type": "approval_pending", "request_id": "req-1"}, + event_type=TransportEventType.APPROVAL_PENDING.value, + payload={"request_id": "req-1"}, ) ) polled = await channel.poll(timeout=0.2) assert polled is not None assert polled.id == "event-1" - assert isinstance(polled.payload, dict) - assert polled.payload.get("type") == "approval_pending" + assert polled.event_type == TransportEventType.APPROVAL_PENDING.value @pytest.mark.asyncio @@ -82,3 +110,46 @@ async def sender(msg: TransportEnvelope) -> None: channel.attach_agent_envelope_sender(sender) polled = await channel.poll(timeout=0.05) assert polled is None + + +def test_direct_client_channel_matches_pollable_protocol() -> None: + channel = DirectClientChannel() + assert isinstance(channel, PollableClientChannel) + + +@pytest.mark.asyncio +async def test_stdio_receiver_uses_event_type_without_legacy_payload_type(capsys) -> None: + channel = StdioClientChannel() + receiver = channel.agent_envelope_receiver() + await receiver( + TransportEnvelope( + id="evt-2", + kind=EnvelopeKind.MESSAGE, + event_type=TransportEventType.APPROVAL_PENDING.value, + payload={"resp": {"request": {"request_id": "req-42"}}}, + ) + ) + captured = capsys.readouterr() + assert "approval pending: request_id=req-42" in captured.out + + +@pytest.mark.asyncio +async def test_stdio_receiver_does_not_route_by_payload_type_without_event_type(capsys) -> None: + channel = StdioClientChannel() + receiver = channel.agent_envelope_receiver() + + await receiver( + TransportEnvelope( + id="evt-legacy-result", + kind=EnvelopeKind.MESSAGE, + payload={ + "type": "result", + "kind": "message", + "resp": {"output": "hello"}, + }, + ) + ) + + captured = capsys.readouterr() + assert "Assistant: {'type': 'result'" in captured.out + assert "Assistant: hello" not in captured.out diff --git a/tests/unit/test_transport_channel.py b/tests/unit/test_transport_channel.py index e42ed153..66515f8c 100644 --- a/tests/unit/test_transport_channel.py +++ b/tests/unit/test_transport_channel.py @@ -3,7 +3,13 @@ import pytest -from dare_framework.transport import AgentChannel, EnvelopeKind, TransportEnvelope, new_envelope_id +from dare_framework.transport import ( + AgentChannel, + EnvelopeKind, + TransportEnvelope, + TransportEventType, + new_envelope_id, +) from dare_framework.transport.interaction.resource_action import ResourceAction from dare_framework.transport.interaction.control_handler import AgentControlHandler from dare_framework.transport.interaction.dispatcher import ActionHandlerDispatcher @@ -290,9 +296,9 @@ async def receiver(msg: TransportEnvelope) -> None: await channel.stop() assert len(seen) == 1 + assert seen[0].event_type == TransportEventType.ERROR.value payload = seen[0].payload assert isinstance(payload, dict) - assert payload.get("type") == "error" assert payload.get("kind") == "control" assert payload.get("ok") is False assert payload.get("code") == "INVALID_CONTROL_PAYLOAD" @@ -326,9 +332,9 @@ async def receiver(msg: TransportEnvelope) -> None: await channel.stop() assert len(seen) == 1 + assert seen[0].event_type == TransportEventType.ERROR.value payload = seen[0].payload assert isinstance(payload, dict) - assert payload.get("type") == "error" assert payload.get("kind") == "message" assert payload.get("ok") is False assert payload.get("code") == "INVALID_MESSAGE_PAYLOAD" @@ -377,9 +383,9 @@ async def receiver(msg: TransportEnvelope) -> None: await channel.stop() assert msg.payload == "hello-after-timeout" + assert seen[0].event_type == TransportEventType.ERROR.value timeout_payload = seen[0].payload assert isinstance(timeout_payload, dict) - assert timeout_payload.get("type") == "error" assert timeout_payload.get("kind") == "action" assert timeout_payload.get("ok") is False assert timeout_payload.get("code") == "ACTION_TIMEOUT" diff --git a/tests/unit/test_transport_types.py b/tests/unit/test_transport_types.py new file mode 100644 index 00000000..7e0cb349 --- /dev/null +++ b/tests/unit/test_transport_types.py @@ -0,0 +1,50 @@ +from __future__ import annotations + +import pytest + +from dare_framework.transport import normalize_transport_event_type +from dare_framework.transport.types import EnvelopeKind, TransportEnvelope, TransportEventType + + +def test_transport_envelope_does_not_derive_event_type_from_payload_type() -> None: + envelope = TransportEnvelope( + id="evt-no-derive", + kind=EnvelopeKind.MESSAGE, + payload={"type": "result"}, + ) + assert envelope.event_type is None + + +def test_transport_envelope_normalizes_legacy_event_type() -> None: + envelope = TransportEnvelope( + id="evt-legacy", + kind="message", + event_type="approval_pending", + payload={"type": "approval_pending"}, + ) + assert envelope.kind == EnvelopeKind.MESSAGE + assert envelope.event_type == TransportEventType.APPROVAL_PENDING.value + + +def test_transport_envelope_accepts_select_kind() -> None: + envelope = TransportEnvelope( + id="evt-select", + kind="select", + event_type=TransportEventType.APPROVAL_PENDING.value, + payload={"kind": "approval"}, + ) + assert envelope.kind == EnvelopeKind.SELECT + + +def test_transport_envelope_rejects_empty_event_type() -> None: + with pytest.raises(ValueError, match="event_type must not be empty"): + TransportEnvelope( + id="evt-empty", + kind=EnvelopeKind.MESSAGE, + event_type=" ", + payload={"type": "result"}, + ) + + +def test_transport_facade_re_exports_event_type_normalizer() -> None: + assert normalize_transport_event_type("approval_pending") == TransportEventType.APPROVAL_PENDING.value