diff --git a/packages/pi_agent/src/pi_agent/__init__.py b/packages/pi_agent/src/pi_agent/__init__.py index 1ff3895..a6b6cfa 100644 --- a/packages/pi_agent/src/pi_agent/__init__.py +++ b/packages/pi_agent/src/pi_agent/__init__.py @@ -1,3 +1,88 @@ """Agent runtime: pure loop + stateful Agent SDK.""" -__all__: list[str] = [] +from pi_agent.agent import Agent +from pi_agent.loop import agent_loop, agent_loop_continue, run_agent_loop, run_agent_loop_continue +from pi_agent.types import ( + AfterToolCallContext, + AfterToolCallResult, + AgentContext, + AgentEndEvent, + AgentEvent, + AgentLoopConfig, + AgentMessage, + AgentStartEvent, + AgentTool, + AgentToolResult, + AssistantMessage, + AssistantStreamDone, + AssistantStreamError, + AssistantStreamEvent, + AssistantStreamStart, + AssistantStreamTextDelta, + AssistantStreamToolCallDelta, + BeforeToolCallContext, + BeforeToolCallResult, + CustomMessage, + LlmMessage, + MessageEndEvent, + MessageStartEvent, + MessageUpdateEvent, + QueueMode, + StreamFn, + StreamRequest, + ToolCall, + ToolExecutionEndEvent, + ToolExecutionMode, + ToolExecutionStartEvent, + ToolExecutionUpdateEvent, + ToolResultMessage, + TurnEndEvent, + TurnStartEvent, + UserMessage, + default_convert_to_llm, +) + +__all__ = [ + "AfterToolCallContext", + "AfterToolCallResult", + "Agent", + "AgentContext", + "AgentEndEvent", + "AgentEvent", + "AgentLoopConfig", + "AgentMessage", + "AgentStartEvent", + "AgentTool", + "AgentToolResult", + "AssistantMessage", + "AssistantStreamDone", + "AssistantStreamError", + "AssistantStreamEvent", + "AssistantStreamStart", + "AssistantStreamTextDelta", + "AssistantStreamToolCallDelta", + "BeforeToolCallContext", + "BeforeToolCallResult", + "CustomMessage", + "LlmMessage", + "MessageEndEvent", + "MessageStartEvent", + "MessageUpdateEvent", + "QueueMode", + "StreamFn", + "StreamRequest", + "ToolCall", + "ToolExecutionEndEvent", + "ToolExecutionMode", + "ToolExecutionStartEvent", + "ToolExecutionUpdateEvent", + "ToolResultMessage", + "TurnEndEvent", + "TurnStartEvent", + "UserMessage", + "agent_loop", + "agent_loop_continue", + "default_convert_to_llm", + "run_agent_loop", + "run_agent_loop_continue", +] diff --git a/packages/pi_agent/src/pi_agent/agent.py b/packages/pi_agent/src/pi_agent/agent.py new file mode 100644 index 0000000..555848d --- /dev/null +++ b/packages/pi_agent/src/pi_agent/agent.py @@ -0,0 +1,288 @@ +"""Stateful Agent facade over the pure agent loop.""" + +from __future__ import annotations + +import asyncio +import inspect +from collections.abc import Awaitable, Callable, Sequence +from typing import cast + +from pi_agent.loop import run_agent_loop, run_agent_loop_continue +from pi_agent.types import ( + AfterToolCallFn, + AgentContext, + AgentEndEvent, + AgentEvent, + AgentLoopConfig, + AgentMessage, + AgentTool, + BeforeToolCallFn, + MessageEndEvent, + MessageStartEvent, + MessageUpdateEvent, + QueueMode, + StreamFn, + ToolExecutionMode, + UserMessage, + default_convert_to_llm, +) + +Listener = Callable[[AgentEvent], Awaitable[None] | None] + + +class _PendingQueue: + def __init__(self, mode: QueueMode = "one-at-a-time") -> None: + self.mode = mode + self._messages: list[AgentMessage] = [] + + def enqueue(self, message: AgentMessage) -> None: + self._messages.append(message) + + def has_items(self) -> bool: + return bool(self._messages) + + def drain(self) -> list[AgentMessage]: + if self.mode == "all": + drained = list(self._messages) + self._messages.clear() + return drained + if not self._messages: + return [] + first = self._messages.pop(0) + return [first] + + def clear(self) -> None: + self._messages.clear() + + +class Agent: + """Owns transcript/tools/queues; awaits subscribe listeners as settlement barrier.""" + + def __init__( + self, + *, + stream_fn: StreamFn, + system_prompt: str = "", + tools: Sequence[AgentTool] | None = None, + messages: Sequence[AgentMessage] | None = None, + convert_to_llm: Callable | None = None, + transform_context: Callable | None = None, + before_tool_call: BeforeToolCallFn | None = None, + after_tool_call: AfterToolCallFn | None = None, + tool_execution: ToolExecutionMode = "parallel", + steering_mode: QueueMode = "one-at-a-time", + follow_up_mode: QueueMode = "one-at-a-time", + ) -> None: + self.stream_fn = stream_fn + self.system_prompt = system_prompt + self._tools = list(tools or []) + self._messages = list(messages or []) + self.convert_to_llm = convert_to_llm or default_convert_to_llm + self.transform_context = transform_context + self.before_tool_call = before_tool_call + self.after_tool_call = after_tool_call + self.tool_execution = tool_execution + self._steering = _PendingQueue(steering_mode) + self._follow_up = _PendingQueue(follow_up_mode) + self._listeners: list[Listener] = [] + self._active = False + self._streaming_message: AgentMessage | None = None + self._idle = asyncio.Event() + self._idle.set() + + @property + def messages(self) -> list[AgentMessage]: + return self._messages + + @messages.setter + def messages(self, value: Sequence[AgentMessage]) -> None: + self._messages = list(value) + + @property + def tools(self) -> list[AgentTool]: + return self._tools + + @tools.setter + def tools(self, value: Sequence[AgentTool]) -> None: + self._tools = list(value) + + @property + def is_streaming(self) -> bool: + return self._active + + @property + def streaming_message(self) -> AgentMessage | None: + return self._streaming_message + + @property + def steering_mode(self) -> QueueMode: + return self._steering.mode + + @steering_mode.setter + def steering_mode(self, mode: QueueMode) -> None: + self._steering.mode = mode + + @property + def follow_up_mode(self) -> QueueMode: + return self._follow_up.mode + + @follow_up_mode.setter + def follow_up_mode(self, mode: QueueMode) -> None: + self._follow_up.mode = mode + + def subscribe(self, listener: Listener) -> Callable[[], None]: + self._listeners.append(listener) + + def unsubscribe() -> None: + if listener in self._listeners: + self._listeners.remove(listener) + + return unsubscribe + + def steer(self, message: AgentMessage) -> None: + self._steering.enqueue(message) + + def follow_up(self, message: AgentMessage) -> None: + self._follow_up.enqueue(message) + + def clear_steering_queue(self) -> None: + self._steering.clear() + + def clear_follow_up_queue(self) -> None: + self._follow_up.clear() + + def clear_all_queues(self) -> None: + self.clear_steering_queue() + self.clear_follow_up_queue() + + async def wait_for_idle(self) -> None: + await self._idle.wait() + + async def prompt(self, input: str | AgentMessage | Sequence[AgentMessage]) -> None: + if self._active: + raise RuntimeError( + "Agent is already processing a prompt. Use steer() or follow_up() " + "to queue messages, or wait for completion." + ) + messages = self._normalize_prompt(input) + await self._run_prompt_messages(messages) + + async def continue_(self) -> None: + """Continue from current transcript without appending a new user message.""" + if self._active: + raise RuntimeError( + "Agent is already processing. Wait for completion before continuing." + ) + + if not self._messages: + raise ValueError("No messages to continue from") + + last = self._messages[-1] + if last.role == "assistant": + steered = self._steering.drain() + if steered: + await self._run_prompt_messages(steered, skip_initial_steering_poll=True) + return + follow = self._follow_up.drain() + if follow: + await self._run_prompt_messages(follow) + return + raise ValueError("Cannot continue from message role: assistant") + + await self._run_continuation() + + def _normalize_prompt( + self, input: str | AgentMessage | Sequence[AgentMessage] + ) -> list[AgentMessage]: + if isinstance(input, str): + return [UserMessage(content=input)] + if isinstance(input, list): + return cast(list[AgentMessage], list(input)) + if isinstance(input, tuple): + return cast(list[AgentMessage], list(input)) + return [cast(AgentMessage, input)] + + async def _run_prompt_messages( + self, + messages: list[AgentMessage], + *, + skip_initial_steering_poll: bool = False, + ) -> None: + await self._run_with_lifecycle( + lambda: run_agent_loop( + messages, + self._context_snapshot(), + self._loop_config(skip_initial_steering_poll=skip_initial_steering_poll), + self._process_events, + self.stream_fn, + ) + ) + + async def _run_continuation(self) -> None: + await self._run_with_lifecycle( + lambda: run_agent_loop_continue( + self._context_snapshot(), + self._loop_config(), + self._process_events, + self.stream_fn, + ) + ) + + def _context_snapshot(self) -> AgentContext: + return AgentContext( + system_prompt=self.system_prompt, + messages=list(self._messages), + tools=list(self._tools), + ) + + def _loop_config(self, *, skip_initial_steering_poll: bool = False) -> AgentLoopConfig: + skip = skip_initial_steering_poll + + async def get_steering() -> list[AgentMessage]: + nonlocal skip + if skip: + skip = False + return [] + return self._steering.drain() + + async def get_follow_up() -> list[AgentMessage]: + return self._follow_up.drain() + + return AgentLoopConfig( + convert_to_llm=self.convert_to_llm, + transform_context=self.transform_context, + tool_execution=self.tool_execution, + before_tool_call=self.before_tool_call, + after_tool_call=self.after_tool_call, + get_steering_messages=get_steering, + get_follow_up_messages=get_follow_up, + ) + + async def _run_with_lifecycle(self, executor: Callable[[], Awaitable[object]]) -> None: + if self._active: + raise RuntimeError("Agent is already processing.") + self._active = True + self._idle.clear() + try: + await executor() + finally: + self._active = False + self._idle.set() + + async def _process_events(self, event: AgentEvent) -> None: + if isinstance(event, (MessageStartEvent, MessageUpdateEvent)): + self._streaming_message = event.message + elif isinstance(event, MessageEndEvent): + # Barrier: transcript includes the message before tool preflight / later phases. + self._streaming_message = None + self._messages.append(event.message) + elif isinstance(event, AgentEndEvent): + self._streaming_message = None + for listener in list(self._listeners): + result = listener(event) + if inspect.isawaitable(result): + await result + + +# `continue` is a Python keyword; expose the upstream method name via setattr. +setattr(Agent, "continue", Agent.continue_) diff --git a/packages/pi_agent/src/pi_agent/loop.py b/packages/pi_agent/src/pi_agent/loop.py new file mode 100644 index 0000000..5985d1e --- /dev/null +++ b/packages/pi_agent/src/pi_agent/loop.py @@ -0,0 +1,700 @@ +"""Pure agent loop: events, LLM stream boundary, tool execution.""" + +from __future__ import annotations + +import asyncio +import inspect +from collections.abc import AsyncIterator, Awaitable, Callable, Sequence +from dataclasses import dataclass +from typing import Any, TypeVar, cast + +from pi_agent.types import ( + AfterToolCallContext, + AgentContext, + AgentEndEvent, + AgentEvent, + AgentLoopConfig, + AgentMessage, + AgentStartEvent, + AgentTool, + AgentToolResult, + AssistantMessage, + AssistantStreamDone, + AssistantStreamError, + AssistantStreamEvent, + AssistantStreamStart, + AssistantStreamTextDelta, + AssistantStreamToolCallDelta, + BeforeToolCallContext, + MessageEndEvent, + MessageStartEvent, + MessageUpdateEvent, + StreamFn, + StreamRequest, + ToolCall, + ToolExecutionEndEvent, + ToolExecutionStartEvent, + ToolExecutionUpdateEvent, + ToolResultMessage, + TurnEndEvent, + TurnStartEvent, + UserMessage, +) + +T = TypeVar("T") +Emit = Callable[[AgentEvent], Awaitable[None] | None] + + +@dataclass(slots=True) +class _ToolBatch: + messages: list[ToolResultMessage] + terminate: bool + + +@dataclass(slots=True) +class _FinalizedToolCall: + tool_call: ToolCall + result: AgentToolResult + is_error: bool + + +@dataclass(slots=True) +class _PreparedToolCall: + tool_call: ToolCall + tool: AgentTool + args: dict[str, Any] + + +@dataclass(slots=True) +class _ImmediateToolCall: + result: AgentToolResult + is_error: bool + + +@dataclass(slots=True) +class _ExecutedToolCall: + result: AgentToolResult + is_error: bool + + +async def _maybe_await(value: T | Awaitable[T]) -> T: + if inspect.isawaitable(value): + return cast(T, await value) + return cast(T, value) + + +async def agent_loop( + prompts: Sequence[AgentMessage], + context: AgentContext, + config: AgentLoopConfig, + stream_fn: StreamFn, +) -> AsyncIterator[AgentEvent]: + """Observational async iterator over one prompt run (does not await external subscribers).""" + queue: asyncio.Queue[AgentEvent | None] = asyncio.Queue() + + async def emit(event: AgentEvent) -> None: + await queue.put(event) + + async def runner() -> None: + try: + await run_agent_loop(list(prompts), context, config, emit, stream_fn) + finally: + await queue.put(None) + + task = asyncio.create_task(runner()) + try: + while True: + item = await queue.get() + if item is None: + break + yield item + finally: + await task + + +async def agent_loop_continue( + context: AgentContext, + config: AgentLoopConfig, + stream_fn: StreamFn, +) -> AsyncIterator[AgentEvent]: + """Continue from existing context without appending a new prompt message.""" + queue: asyncio.Queue[AgentEvent | None] = asyncio.Queue() + + async def emit(event: AgentEvent) -> None: + await queue.put(event) + + async def runner() -> None: + try: + await run_agent_loop_continue(context, config, emit, stream_fn) + finally: + await queue.put(None) + + task = asyncio.create_task(runner()) + try: + while True: + item = await queue.get() + if item is None: + break + yield item + finally: + await task + + +async def run_agent_loop( + prompts: list[AgentMessage], + context: AgentContext, + config: AgentLoopConfig, + emit: Emit, + stream_fn: StreamFn, +) -> list[AgentMessage]: + new_messages: list[AgentMessage] = list(prompts) + current = AgentContext( + system_prompt=context.system_prompt, + messages=[*context.messages, *prompts], + tools=list(context.tools), + ) + + await _emit(emit, AgentStartEvent()) + await _emit(emit, TurnStartEvent()) + for prompt in prompts: + await _emit(emit, MessageStartEvent(message=prompt)) + await _emit(emit, MessageEndEvent(message=prompt)) + + await _run_loop(current, new_messages, config, emit, stream_fn) + return new_messages + + +async def run_agent_loop_continue( + context: AgentContext, + config: AgentLoopConfig, + emit: Emit, + stream_fn: StreamFn, +) -> list[AgentMessage]: + if not context.messages: + raise ValueError("Cannot continue: no messages in context") + if context.messages[-1].role == "assistant": + raise ValueError("Cannot continue from message role: assistant") + + new_messages: list[AgentMessage] = [] + current = AgentContext( + system_prompt=context.system_prompt, + messages=list(context.messages), + tools=list(context.tools), + ) + + await _emit(emit, AgentStartEvent()) + await _emit(emit, TurnStartEvent()) + await _run_loop(current, new_messages, config, emit, stream_fn) + return new_messages + + +async def _emit(emit: Emit, event: AgentEvent) -> None: + result = emit(event) + if inspect.isawaitable(result): + await result + + +async def _run_loop( + current_context: AgentContext, + new_messages: list[AgentMessage], + config: AgentLoopConfig, + emit: Emit, + stream_fn: StreamFn, +) -> None: + first_turn = True + pending = await _poll_messages(config.get_steering_messages) + + while True: + has_more_tool_calls = True + while has_more_tool_calls or pending: + if not first_turn: + await _emit(emit, TurnStartEvent()) + else: + first_turn = False + + if pending: + for message in pending: + await _emit(emit, MessageStartEvent(message=message)) + await _emit(emit, MessageEndEvent(message=message)) + current_context.messages.append(message) + new_messages.append(message) + pending = [] + + message = await _stream_assistant_response(current_context, config, emit, stream_fn) + new_messages.append(message) + + if message.stop_reason in ("error", "aborted"): + await _emit(emit, TurnEndEvent(message=message, tool_results=[])) + await _emit(emit, AgentEndEvent(messages=list(new_messages))) + return + + tool_results: list[ToolResultMessage] = [] + has_more_tool_calls = False + if message.tool_calls: + if message.stop_reason == "length": + batch = await _fail_truncated_tool_calls(message.tool_calls, emit) + else: + batch = await _execute_tool_calls(current_context, message, config, emit) + tool_results = batch.messages + has_more_tool_calls = not batch.terminate + for result in tool_results: + current_context.messages.append(result) + new_messages.append(result) + + await _emit(emit, TurnEndEvent(message=message, tool_results=tool_results)) + + if config.should_stop_after_turn and await _maybe_await( + config.should_stop_after_turn( + message=message, + tool_results=tool_results, + context=current_context, + new_messages=new_messages, + ) + ): + await _emit(emit, AgentEndEvent(messages=list(new_messages))) + return + + pending = await _poll_messages(config.get_steering_messages) + + follow_ups = await _poll_messages(config.get_follow_up_messages) + if follow_ups: + pending = follow_ups + continue + break + + await _emit(emit, AgentEndEvent(messages=list(new_messages))) + + +async def _poll_messages( + getter: Callable[[], Awaitable[list[AgentMessage]] | list[AgentMessage]] | None, +) -> list[AgentMessage]: + if getter is None: + return [] + return list(await _maybe_await(getter())) + + +async def _stream_assistant_response( + context: AgentContext, + config: AgentLoopConfig, + emit: Emit, + stream_fn: StreamFn, +) -> AssistantMessage: + messages: Sequence[AgentMessage] = context.messages + if config.transform_context: + try: + messages = await _maybe_await(config.transform_context(messages)) + except Exception: + messages = context.messages + + try: + llm_messages = list(await _maybe_await(config.convert_to_llm(messages))) + except Exception: + llm_messages = [ + m for m in messages if isinstance(m, (UserMessage, AssistantMessage, ToolResultMessage)) + ] + request = StreamRequest( + system_prompt=context.system_prompt, + messages=llm_messages, + tools=list(context.tools), + ) + + stream_result = stream_fn(request) + if inspect.isawaitable(stream_result): + stream = await cast(Awaitable[AsyncIterator[AssistantStreamEvent]], stream_result) + else: + stream = cast(AsyncIterator[AssistantStreamEvent], stream_result) + + partial: AssistantMessage | None = None + added_partial = False + + async for event in stream: + if isinstance(event, AssistantStreamStart): + partial = event.partial + context.messages.append(partial) + added_partial = True + await _emit(emit, MessageStartEvent(message=partial)) + elif isinstance(event, (AssistantStreamTextDelta, AssistantStreamToolCallDelta)): + if partial is not None: + partial = event.partial + context.messages[-1] = partial + await _emit( + emit, + MessageUpdateEvent(message=partial, assistant_message_event=event), + ) + elif isinstance(event, (AssistantStreamDone, AssistantStreamError)): + final = event.message + if added_partial: + context.messages[-1] = final + else: + context.messages.append(final) + await _emit(emit, MessageStartEvent(message=final)) + await _emit(emit, MessageEndEvent(message=final)) + return final + + raise RuntimeError("StreamFn ended without AssistantStreamDone/Error") + + +async def _fail_truncated_tool_calls( + tool_calls: list[ToolCall], + emit: Emit, +) -> _ToolBatch: + messages: list[ToolResultMessage] = [] + for tool_call in tool_calls: + await _emit( + emit, + ToolExecutionStartEvent( + tool_call_id=tool_call.id, + tool_name=tool_call.name, + args=tool_call.arguments, + ), + ) + result = AgentToolResult( + content=( + f'Tool call "{tool_call.name}" was not executed: the response hit the output ' + "token limit, so its arguments may be truncated. Re-issue the tool call with " + "complete arguments." + ), + ) + await _emit( + emit, + ToolExecutionEndEvent( + tool_call_id=tool_call.id, + tool_name=tool_call.name, + result=result, + is_error=True, + ), + ) + msg = ToolResultMessage( + tool_call_id=tool_call.id, + tool_name=tool_call.name, + content=result.content, + is_error=True, + details=result.details, + ) + await _emit(emit, MessageStartEvent(message=msg)) + await _emit(emit, MessageEndEvent(message=msg)) + messages.append(msg) + return _ToolBatch(messages=messages, terminate=False) + + +async def _execute_tool_calls( + current_context: AgentContext, + assistant_message: AssistantMessage, + config: AgentLoopConfig, + emit: Emit, +) -> _ToolBatch: + tool_calls = assistant_message.tool_calls + has_sequential = False + for tc in tool_calls: + tool = _find_tool(current_context.tools, tc.name) + if tool is not None and tool.execution_mode == "sequential": + has_sequential = True + break + if config.tool_execution == "sequential" or has_sequential: + return await _execute_tool_calls_sequential( + current_context, assistant_message, tool_calls, config, emit + ) + return await _execute_tool_calls_parallel( + current_context, assistant_message, tool_calls, config, emit + ) + + +def _find_tool(tools: list[AgentTool], name: str) -> AgentTool | None: + for tool in tools: + if tool.name == name: + return tool + return None + + +async def _execute_tool_calls_sequential( + current_context: AgentContext, + assistant_message: AssistantMessage, + tool_calls: list[ToolCall], + config: AgentLoopConfig, + emit: Emit, +) -> _ToolBatch: + finalized: list[_FinalizedToolCall] = [] + messages: list[ToolResultMessage] = [] + + for tool_call in tool_calls: + await _emit( + emit, + ToolExecutionStartEvent( + tool_call_id=tool_call.id, + tool_name=tool_call.name, + args=tool_call.arguments, + ), + ) + preparation = await _prepare_tool_call( + current_context, assistant_message, tool_call, config + ) + if isinstance(preparation, _ImmediateToolCall): + outcome = _FinalizedToolCall( + tool_call=tool_call, + result=preparation.result, + is_error=preparation.is_error, + ) + else: + executed = await _execute_prepared_tool_call(preparation, emit) + outcome = await _finalize_executed_tool_call( + current_context, assistant_message, preparation, executed, config + ) + + await _emit( + emit, + ToolExecutionEndEvent( + tool_call_id=outcome.tool_call.id, + tool_name=outcome.tool_call.name, + result=outcome.result, + is_error=outcome.is_error, + ), + ) + msg = _tool_result_message(outcome) + await _emit(emit, MessageStartEvent(message=msg)) + await _emit(emit, MessageEndEvent(message=msg)) + finalized.append(outcome) + messages.append(msg) + + return _ToolBatch(messages=messages, terminate=_should_terminate(finalized)) + + +async def _execute_tool_calls_parallel( + current_context: AgentContext, + assistant_message: AssistantMessage, + tool_calls: list[ToolCall], + config: AgentLoopConfig, + emit: Emit, +) -> _ToolBatch: + # Serialize emit so Agent.subscribe handlers settle without interleaving. + emit_lock = asyncio.Lock() + + async def locked_emit(event: AgentEvent) -> None: + async with emit_lock: + await _emit(emit, event) + + coros: list[Awaitable[_FinalizedToolCall]] = [] + + for tool_call in tool_calls: + await locked_emit( + ToolExecutionStartEvent( + tool_call_id=tool_call.id, + tool_name=tool_call.name, + args=tool_call.arguments, + ), + ) + preparation = await _prepare_tool_call( + current_context, assistant_message, tool_call, config + ) + if isinstance(preparation, _ImmediateToolCall): + outcome = _FinalizedToolCall( + tool_call=tool_call, + result=preparation.result, + is_error=preparation.is_error, + ) + await locked_emit( + ToolExecutionEndEvent( + tool_call_id=tool_call.id, + tool_name=tool_call.name, + result=outcome.result, + is_error=outcome.is_error, + ), + ) + coros.append(_resolve(outcome)) + continue + + async def run_one(prep: _PreparedToolCall = preparation) -> _FinalizedToolCall: + executed = await _execute_prepared_tool_call(prep, locked_emit) + finalized = await _finalize_executed_tool_call( + current_context, assistant_message, prep, executed, config + ) + await locked_emit( + ToolExecutionEndEvent( + tool_call_id=finalized.tool_call.id, + tool_name=finalized.tool_call.name, + result=finalized.result, + is_error=finalized.is_error, + ), + ) + return finalized + + coros.append(run_one()) + + ordered = list(await asyncio.gather(*coros)) + messages: list[ToolResultMessage] = [] + for outcome in ordered: + msg = _tool_result_message(outcome) + await locked_emit(MessageStartEvent(message=msg)) + await locked_emit(MessageEndEvent(message=msg)) + messages.append(msg) + + return _ToolBatch(messages=messages, terminate=_should_terminate(ordered)) + + +async def _resolve(value: _FinalizedToolCall) -> _FinalizedToolCall: + return value + + +def _should_terminate(finalized: list[_FinalizedToolCall]) -> bool: + return bool(finalized) and all(item.result.terminate is True for item in finalized) + + +def _tool_result_message(outcome: _FinalizedToolCall) -> ToolResultMessage: + return ToolResultMessage( + tool_call_id=outcome.tool_call.id, + tool_name=outcome.tool_call.name, + content=outcome.result.content, + is_error=outcome.is_error, + details=outcome.result.details, + ) + + +def _validate_tool_arguments(tool: AgentTool, args: dict[str, Any]) -> dict[str, Any]: + """Lightweight JSON-Schema-ish check against AgentTool.parameters when present.""" + schema = tool.parameters + if not schema: + return args + if not isinstance(args, dict): + raise TypeError(f"Tool {tool.name} arguments must be an object") + required = schema.get("required") + if isinstance(required, list): + missing = [key for key in required if key not in args] + if missing: + raise ValueError( + f"Tool {tool.name} missing required arguments: {', '.join(map(str, missing))}" + ) + return args + + +async def _prepare_tool_call( + current_context: AgentContext, + assistant_message: AssistantMessage, + tool_call: ToolCall, + config: AgentLoopConfig, +) -> _PreparedToolCall | _ImmediateToolCall: + tool = _find_tool(current_context.tools, tool_call.name) + if tool is None: + return _ImmediateToolCall( + result=AgentToolResult(content=f"Tool {tool_call.name} not found"), + is_error=True, + ) + if tool.execute is None: + return _ImmediateToolCall( + result=AgentToolResult(content=f"Tool {tool_call.name} has no execute handler"), + is_error=True, + ) + + try: + args = dict(tool_call.arguments) + if tool.prepare_arguments is not None: + args = tool.prepare_arguments(args) + args = _validate_tool_arguments(tool, args) + if config.before_tool_call is not None: + before = await _maybe_await( + config.before_tool_call( + BeforeToolCallContext( + assistant_message=assistant_message, + tool_call=tool_call, + args=args, + context=current_context, + ) + ) + ) + if before is not None and before.block: + return _ImmediateToolCall( + result=AgentToolResult( + content=before.reason or "Tool execution was blocked", + ), + is_error=True, + ) + return _PreparedToolCall(tool_call=tool_call, tool=tool, args=args) + except Exception as exc: + return _ImmediateToolCall( + result=AgentToolResult(content=str(exc)), + is_error=True, + ) + + +async def _execute_prepared_tool_call(prepared: _PreparedToolCall, emit: Emit) -> _ExecutedToolCall: + tool = prepared.tool + tool_call = prepared.tool_call + accepting = True + update_tasks: list[asyncio.Task[None]] = [] + + def on_update(partial: AgentToolResult) -> None: + nonlocal accepting + if not accepting: + return + + async def _emit_update() -> None: + await _emit( + emit, + ToolExecutionUpdateEvent( + tool_call_id=tool_call.id, + tool_name=tool_call.name, + args=tool_call.arguments, + partial_result=partial, + ), + ) + + update_tasks.append(asyncio.create_task(_emit_update())) + + assert tool.execute is not None + try: + result = await tool.execute(tool_call.id, prepared.args, on_update=on_update) + accepting = False + if update_tasks: + await asyncio.gather(*update_tasks) + return _ExecutedToolCall(result=result, is_error=False) + except Exception as exc: + accepting = False + if update_tasks: + await asyncio.gather(*update_tasks) + return _ExecutedToolCall( + result=AgentToolResult(content=str(exc)), + is_error=True, + ) + + +async def _finalize_executed_tool_call( + current_context: AgentContext, + assistant_message: AssistantMessage, + prepared: _PreparedToolCall, + executed: _ExecutedToolCall, + config: AgentLoopConfig, +) -> _FinalizedToolCall: + result = executed.result + is_error = executed.is_error + + if config.after_tool_call is not None: + try: + after = await _maybe_await( + config.after_tool_call( + AfterToolCallContext( + assistant_message=assistant_message, + tool_call=prepared.tool_call, + args=prepared.args, + result=result, + is_error=is_error, + context=current_context, + ) + ) + ) + if after is not None: + result = AgentToolResult( + content=after.content if after.content is not None else result.content, + details=after.details if after.details is not None else result.details, + terminate=( + after.terminate if after.terminate is not None else result.terminate + ), + ) + if after.is_error is not None: + is_error = after.is_error + except Exception as exc: + result = AgentToolResult(content=str(exc)) + is_error = True + + return _FinalizedToolCall( + tool_call=prepared.tool_call, + result=result, + is_error=is_error, + ) diff --git a/packages/pi_agent/src/pi_agent/types.py b/packages/pi_agent/src/pi_agent/types.py new file mode 100644 index 0000000..45ca1e6 --- /dev/null +++ b/packages/pi_agent/src/pi_agent/types.py @@ -0,0 +1,285 @@ +"""Public types for the agent loop and stateful Agent SDK.""" + +from __future__ import annotations + +from collections.abc import AsyncIterator, Awaitable, Callable, Sequence +from dataclasses import dataclass, field +from typing import Any, Literal, Protocol + +StopReason = Literal["stop", "toolUse", "length", "error", "aborted"] +ToolExecutionMode = Literal["sequential", "parallel"] +QueueMode = Literal["all", "one-at-a-time"] + + +@dataclass(frozen=True, slots=True) +class ToolCall: + id: str + name: str + arguments: dict[str, Any] = field(default_factory=dict) + + +@dataclass(slots=True) +class UserMessage: + content: str + role: Literal["user"] = "user" + + +@dataclass(slots=True) +class AssistantMessage: + content: str | None = None + tool_calls: list[ToolCall] = field(default_factory=list) + stop_reason: StopReason | None = None + error_message: str | None = None + role: Literal["assistant"] = "assistant" + + +@dataclass(slots=True) +class ToolResultMessage: + tool_call_id: str + tool_name: str + content: str + is_error: bool = False + details: Any = None + role: Literal["toolResult"] = "toolResult" + + +@dataclass(slots=True) +class CustomMessage: + """UI/app-only transcript message; must be filtered or mapped in convert_to_llm.""" + + role: str + content: Any = None + data: dict[str, Any] = field(default_factory=dict) + + +LlmMessage = UserMessage | AssistantMessage | ToolResultMessage +AgentMessage = LlmMessage | CustomMessage + + +@dataclass(frozen=True, slots=True) +class AgentToolResult: + content: str + details: Any = None + terminate: bool = False + + +@dataclass(slots=True) +class AgentTool: + name: str + description: str = "" + label: str = "" + parameters: dict[str, Any] = field(default_factory=dict) + execution_mode: ToolExecutionMode | None = None + prepare_arguments: Callable[[dict[str, Any]], dict[str, Any]] | None = None + execute: Callable[..., Awaitable[AgentToolResult]] | None = None + + +@dataclass(slots=True) +class AgentContext: + system_prompt: str = "" + messages: list[AgentMessage] = field(default_factory=list) + tools: list[AgentTool] = field(default_factory=list) + + +@dataclass(frozen=True, slots=True) +class StreamRequest: + """LLM-facing request built at the stream boundary after convert hooks.""" + + system_prompt: str + messages: list[LlmMessage] + tools: list[AgentTool] = field(default_factory=list) + + +@dataclass(frozen=True, slots=True) +class AssistantStreamStart: + partial: AssistantMessage + + +@dataclass(frozen=True, slots=True) +class AssistantStreamTextDelta: + partial: AssistantMessage + delta: str + + +@dataclass(frozen=True, slots=True) +class AssistantStreamToolCallDelta: + partial: AssistantMessage + tool_call_id: str + name: str | None = None + arguments_delta: str | None = None + + +@dataclass(frozen=True, slots=True) +class AssistantStreamDone: + message: AssistantMessage + + +@dataclass(frozen=True, slots=True) +class AssistantStreamError: + message: AssistantMessage + + +AssistantStreamEvent = ( + AssistantStreamStart + | AssistantStreamTextDelta + | AssistantStreamToolCallDelta + | AssistantStreamDone + | AssistantStreamError +) + + +class StreamFn(Protocol): + def __call__( + self, request: StreamRequest + ) -> AsyncIterator[AssistantStreamEvent] | Awaitable[AsyncIterator[AssistantStreamEvent]]: ... + + +@dataclass(frozen=True, slots=True) +class BeforeToolCallResult: + block: bool = False + reason: str | None = None + + +@dataclass(frozen=True, slots=True) +class AfterToolCallResult: + content: str | None = None + details: Any = None + is_error: bool | None = None + terminate: bool | None = None + + +@dataclass(slots=True) +class BeforeToolCallContext: + assistant_message: AssistantMessage + tool_call: ToolCall + args: dict[str, Any] + context: AgentContext + + +@dataclass(slots=True) +class AfterToolCallContext: + assistant_message: AssistantMessage + tool_call: ToolCall + args: dict[str, Any] + result: AgentToolResult + is_error: bool + context: AgentContext + + +ConvertToLlm = Callable[ + [Sequence[AgentMessage]], + Sequence[LlmMessage] | Awaitable[Sequence[LlmMessage]], +] +TransformContext = Callable[ + [Sequence[AgentMessage]], + Sequence[AgentMessage] | Awaitable[Sequence[AgentMessage]], +] +BeforeToolCallFn = Callable[ + [BeforeToolCallContext], + Awaitable[BeforeToolCallResult | None] | BeforeToolCallResult | None, +] +AfterToolCallFn = Callable[ + [AfterToolCallContext], + Awaitable[AfterToolCallResult | None] | AfterToolCallResult | None, +] +MessageQueueFn = Callable[[], Awaitable[list[AgentMessage]] | list[AgentMessage]] + + +@dataclass(slots=True) +class AgentLoopConfig: + convert_to_llm: ConvertToLlm + transform_context: TransformContext | None = None + tool_execution: ToolExecutionMode = "parallel" + before_tool_call: BeforeToolCallFn | None = None + after_tool_call: AfterToolCallFn | None = None + get_steering_messages: MessageQueueFn | None = None + get_follow_up_messages: MessageQueueFn | None = None + should_stop_after_turn: Callable[..., Awaitable[bool] | bool] | None = None + + +@dataclass(frozen=True, slots=True) +class AgentStartEvent: + type: Literal["agent_start"] = "agent_start" + + +@dataclass(frozen=True, slots=True) +class AgentEndEvent: + messages: list[AgentMessage] + type: Literal["agent_end"] = "agent_end" + + +@dataclass(frozen=True, slots=True) +class TurnStartEvent: + type: Literal["turn_start"] = "turn_start" + + +@dataclass(frozen=True, slots=True) +class TurnEndEvent: + message: AgentMessage + tool_results: list[ToolResultMessage] + type: Literal["turn_end"] = "turn_end" + + +@dataclass(frozen=True, slots=True) +class MessageStartEvent: + message: AgentMessage + type: Literal["message_start"] = "message_start" + + +@dataclass(frozen=True, slots=True) +class MessageUpdateEvent: + message: AgentMessage + assistant_message_event: AssistantStreamEvent + type: Literal["message_update"] = "message_update" + + +@dataclass(frozen=True, slots=True) +class MessageEndEvent: + message: AgentMessage + type: Literal["message_end"] = "message_end" + + +@dataclass(frozen=True, slots=True) +class ToolExecutionStartEvent: + tool_call_id: str + tool_name: str + args: dict[str, Any] + type: Literal["tool_execution_start"] = "tool_execution_start" + + +@dataclass(frozen=True, slots=True) +class ToolExecutionUpdateEvent: + tool_call_id: str + tool_name: str + args: dict[str, Any] + partial_result: AgentToolResult + type: Literal["tool_execution_update"] = "tool_execution_update" + + +@dataclass(frozen=True, slots=True) +class ToolExecutionEndEvent: + tool_call_id: str + tool_name: str + result: AgentToolResult + is_error: bool + type: Literal["tool_execution_end"] = "tool_execution_end" + + +AgentEvent = ( + AgentStartEvent + | AgentEndEvent + | TurnStartEvent + | TurnEndEvent + | MessageStartEvent + | MessageUpdateEvent + | MessageEndEvent + | ToolExecutionStartEvent + | ToolExecutionUpdateEvent + | ToolExecutionEndEvent +) + + +def default_convert_to_llm(messages: Sequence[AgentMessage]) -> list[LlmMessage]: + return [ + m for m in messages if isinstance(m, (UserMessage, AssistantMessage, ToolResultMessage)) + ] diff --git a/packages/pi_agent/tests/test_agent.py b/packages/pi_agent/tests/test_agent.py new file mode 100644 index 0000000..9b6fd72 --- /dev/null +++ b/packages/pi_agent/tests/test_agent.py @@ -0,0 +1,211 @@ +"""Seam: stateful Agent — subscribe settlement, prompt/continue/steer/follow_up.""" + +from __future__ import annotations + +import asyncio +from collections.abc import AsyncIterator + +import pytest +from pi_agent.types import ( + AssistantStreamDone, + AssistantStreamEvent, + AssistantStreamStart, +) + +from pi_agent import ( + Agent, + AgentEndEvent, + AgentEvent, + AgentTool, + AgentToolResult, + AssistantMessage, + MessageEndEvent, + StreamRequest, + ToolCall, + ToolExecutionStartEvent, + UserMessage, +) + + +def _scripted_stream( + responses: list[AssistantMessage], +): + remaining = list(responses) + + async def stream_fn( + request: StreamRequest, + ) -> AsyncIterator[AssistantStreamEvent]: + msg = remaining.pop(0) + yield AssistantStreamStart(partial=AssistantMessage(content="")) + yield AssistantStreamDone(message=msg) + + return stream_fn + + +async def test_subscribe_awaits_handlers_before_runtime_advances() -> None: + order: list[str] = [] + gate = asyncio.Event() + + async def execute(*_a, **_k) -> AgentToolResult: + return await _async_result(order, "exec") + + agent = Agent( + stream_fn=_scripted_stream( + [ + AssistantMessage( + content=None, + tool_calls=[ToolCall(id="1", name="t", arguments={})], + stop_reason="toolUse", + ), + AssistantMessage(content="done", stop_reason="stop"), + ] + ), + tools=[AgentTool(name="t", execute=execute)], + tool_execution="sequential", + ) + + async def on_event(event: AgentEvent) -> None: + if isinstance(event, MessageEndEvent) and event.message.role == "assistant": + order.append("assistant_end_handler_start") + await gate.wait() + order.append("assistant_end_handler_done") + assert agent.messages[-1].role == "assistant" + + agent.subscribe(on_event) + + async def release() -> None: + await asyncio.sleep(0.01) + assert "exec" not in order + gate.set() + + releaser = asyncio.create_task(release()) + await agent.prompt("go") + await releaser + + assert order[:3] == [ + "assistant_end_handler_start", + "assistant_end_handler_done", + "exec", + ] + assert agent.messages[-1].content == "done" + + +async def _async_result(order: list[str], label: str) -> AgentToolResult: + order.append(label) + return AgentToolResult(content=label) + + +async def test_prompt_rejects_concurrent_prompt() -> None: + started = asyncio.Event() + release = asyncio.Event() + + async def slow_stream( + request: StreamRequest, + ) -> AsyncIterator[AssistantStreamEvent]: + started.set() + await release.wait() + yield AssistantStreamStart(partial=AssistantMessage(content="")) + yield AssistantStreamDone(message=AssistantMessage(content="ok", stop_reason="stop")) + + agent = Agent(stream_fn=slow_stream) + task = asyncio.create_task(agent.prompt("one")) + await started.wait() + with pytest.raises(RuntimeError, match="already processing"): + await agent.prompt("two") + release.set() + await task + + +async def test_steer_injects_after_turn_tools() -> None: + responses = [ + AssistantMessage( + content=None, + tool_calls=[ToolCall(id="1", name="t", arguments={})], + stop_reason="toolUse", + ), + AssistantMessage(content="after-steer", stop_reason="stop"), + ] + + async def execute(_id: str, _args: dict, **_k) -> AgentToolResult: + return AgentToolResult(content="tool-ok") + + agent = Agent( + stream_fn=_scripted_stream(responses), + tools=[AgentTool(name="t", execute=execute)], + tool_execution="sequential", + ) + + async def on_event(event: AgentEvent) -> None: + if isinstance(event, ToolExecutionStartEvent): + agent.steer(UserMessage(content="steer now")) + + agent.subscribe(on_event) + await agent.prompt("go") + + contents = [getattr(m, "content", None) for m in agent.messages] + assert "steer now" in contents + assert contents[-1] == "after-steer" + + +async def test_follow_up_runs_after_agent_would_stop() -> None: + responses = [ + AssistantMessage(content="first", stop_reason="stop"), + AssistantMessage(content="follow", stop_reason="stop"), + ] + agent = Agent(stream_fn=_scripted_stream(responses)) + + async def on_event(event: AgentEvent) -> None: + if ( + isinstance(event, MessageEndEvent) + and getattr(event.message, "content", None) == "first" + ): + agent.follow_up(UserMessage(content="more")) + + agent.subscribe(on_event) + await agent.prompt("go") + assert [getattr(m, "content", None) for m in agent.messages] == [ + "go", + "first", + "more", + "follow", + ] + + +async def test_continue_resumes_without_new_prompt() -> None: + agent = Agent( + stream_fn=_scripted_stream([AssistantMessage(content="continued", stop_reason="stop")]), + messages=[UserMessage(content="already there")], + ) + await agent.continue_() + assert [getattr(m, "content", None) for m in agent.messages] == [ + "already there", + "continued", + ] + + +async def test_await_prompt_waits_for_agent_end_listeners() -> None: + order: list[str] = [] + gate = asyncio.Event() + + agent = Agent( + stream_fn=_scripted_stream([AssistantMessage(content="x", stop_reason="stop")]), + ) + + async def on_event(event: AgentEvent) -> None: + if isinstance(event, AgentEndEvent): + order.append("end_start") + await gate.wait() + order.append("end_done") + + agent.subscribe(on_event) + + async def release() -> None: + await asyncio.sleep(0.01) + assert "end_done" not in order + gate.set() + + releaser = asyncio.create_task(release()) + await agent.prompt("hi") + await releaser + assert order == ["end_start", "end_done"] + assert agent.is_streaming is False diff --git a/packages/pi_agent/tests/test_convert.py b/packages/pi_agent/tests/test_convert.py new file mode 100644 index 0000000..92aef76 --- /dev/null +++ b/packages/pi_agent/tests/test_convert.py @@ -0,0 +1,62 @@ +"""Seam: AgentMessage stays in transcript; LLM Messages only at stream time.""" + +from __future__ import annotations + +from collections.abc import AsyncIterator + +from pi_agent.types import ( + AssistantStreamDone, + AssistantStreamEvent, + AssistantStreamStart, +) + +from pi_agent import ( + AgentContext, + AgentEndEvent, + AgentLoopConfig, + AssistantMessage, + CustomMessage, + StreamRequest, + UserMessage, + agent_loop, + default_convert_to_llm, +) + + +async def test_custom_messages_filtered_before_stream() -> None: + seen: list[StreamRequest] = [] + + async def stream_fn( + request: StreamRequest, + ) -> AsyncIterator[AssistantStreamEvent]: + seen.append(request) + final = AssistantMessage(content="ok", stop_reason="stop") + yield AssistantStreamStart(partial=AssistantMessage(content="")) + yield AssistantStreamDone(message=final) + + context = AgentContext( + system_prompt="sys", + messages=[CustomMessage(role="notification", content="ignore me")], + ) + prompt = UserMessage(content="hi") + config = AgentLoopConfig( + convert_to_llm=default_convert_to_llm, + transform_context=lambda msgs: [m for m in msgs if getattr(m, "role", None) != "noise"], + ) + + events = [ + e + async for e in agent_loop( + [prompt, CustomMessage(role="noise", content="drop")], + context, + config, + stream_fn=stream_fn, + ) + ] + + assert seen[0].messages == [UserMessage(content="hi")] + assert all(m.role != "notification" for m in seen[0].messages) + assert isinstance(events[-1], AgentEndEvent) + roles = [m.role for m in events[-1].messages] + assert "noise" in roles + assert "user" in roles diff --git a/packages/pi_agent/tests/test_loop_prompt.py b/packages/pi_agent/tests/test_loop_prompt.py new file mode 100644 index 0000000..1679bec --- /dev/null +++ b/packages/pi_agent/tests/test_loop_prompt.py @@ -0,0 +1,78 @@ +"""Seam: agent_loop public surface with injected fake StreamFn. + +Prompt without tools emits the settled agent/turn/message event sequence. +""" + +from __future__ import annotations + +from collections.abc import AsyncIterator + +from pi_agent.types import ( + AssistantStreamDone, + AssistantStreamEvent, + AssistantStreamStart, + AssistantStreamTextDelta, +) + +from pi_agent import ( + AgentContext, + AgentEndEvent, + AgentLoopConfig, + AssistantMessage, + MessageEndEvent, + MessageStartEvent, + StreamRequest, + TurnEndEvent, + UserMessage, + agent_loop, + default_convert_to_llm, +) + + +async def _fake_text_stream( + request: StreamRequest, +) -> AsyncIterator[AssistantStreamEvent]: + assert request.messages[-1].role == "user" + partial = AssistantMessage(content="", stop_reason=None) + yield AssistantStreamStart(partial=partial) + partial = AssistantMessage(content="Hi", stop_reason=None) + yield AssistantStreamTextDelta(partial=partial, delta="Hi") + final = AssistantMessage(content="Hi", stop_reason="stop") + yield AssistantStreamDone(message=final) + + +async def test_agent_loop_prompt_without_tools_emits_fixed_event_sequence() -> None: + prompt = UserMessage(content="Hello") + context = AgentContext(system_prompt="You are helpful.", messages=[]) + config = AgentLoopConfig(convert_to_llm=default_convert_to_llm) + + events = [ + event + async for event in agent_loop( + [prompt], + context, + config, + stream_fn=_fake_text_stream, + ) + ] + + types = [e.type for e in events] + assert types == [ + "agent_start", + "turn_start", + "message_start", + "message_end", + "message_start", + "message_update", + "message_end", + "turn_end", + "agent_end", + ] + assert isinstance(events[2], MessageStartEvent) + assert events[2].message.content == "Hello" + assert isinstance(events[6], MessageEndEvent) + assert events[6].message.content == "Hi" + assert isinstance(events[7], TurnEndEvent) + assert events[7].tool_results == [] + assert isinstance(events[8], AgentEndEvent) + assert events[8].messages[-1].content == "Hi" diff --git a/packages/pi_agent/tests/test_loop_tools.py b/packages/pi_agent/tests/test_loop_tools.py new file mode 100644 index 0000000..de9254d --- /dev/null +++ b/packages/pi_agent/tests/test_loop_tools.py @@ -0,0 +1,346 @@ +"""Seam: tool batching with before/after hooks under sequential|parallel modes.""" + +from __future__ import annotations + +import asyncio +from collections.abc import AsyncIterator + +from pi_agent.types import ( + AssistantStreamDone, + AssistantStreamEvent, + AssistantStreamStart, +) + +from pi_agent import ( + AfterToolCallContext, + AfterToolCallResult, + AgentContext, + AgentEndEvent, + AgentLoopConfig, + AgentTool, + AgentToolResult, + AssistantMessage, + BeforeToolCallContext, + BeforeToolCallResult, + MessageStartEvent, + StreamRequest, + ToolCall, + ToolExecutionEndEvent, + ToolResultMessage, + TurnEndEvent, + UserMessage, + agent_loop, + default_convert_to_llm, +) + + +async def test_sequential_tools_run_with_before_after_hooks() -> None: + order: list[str] = [] + + async def execute_echo(tool_call_id: str, args: dict, **_kwargs) -> AgentToolResult: + order.append(f"exec:{args['name']}") + return AgentToolResult(content=f"ok:{args['name']}") + + async def before(ctx: BeforeToolCallContext) -> BeforeToolCallResult | None: + order.append(f"before:{ctx.tool_call.name}") + return None + + async def after(ctx: AfterToolCallContext) -> AfterToolCallResult | None: + order.append(f"after:{ctx.tool_call.name}") + return AfterToolCallResult(content=f"after:{ctx.result.content}") + + tools = [ + AgentTool(name="echo", description="echo", execute=execute_echo), + ] + calls = [ + ToolCall(id="1", name="echo", arguments={"name": "a"}), + ToolCall(id="2", name="echo", arguments={"name": "b"}), + ] + responses = [ + AssistantMessage(content=None, tool_calls=calls, stop_reason="toolUse"), + AssistantMessage(content="done", stop_reason="stop"), + ] + + async def stream_fn( + request: StreamRequest, + ) -> AsyncIterator[AssistantStreamEvent]: + msg = responses.pop(0) + yield AssistantStreamStart(partial=AssistantMessage(content=msg.content, tool_calls=[])) + yield AssistantStreamDone(message=msg) + + events = [ + e + async for e in agent_loop( + [UserMessage(content="go")], + AgentContext(tools=tools), + AgentLoopConfig( + convert_to_llm=default_convert_to_llm, + tool_execution="sequential", + before_tool_call=before, + after_tool_call=after, + ), + stream_fn=stream_fn, + ) + ] + + assert order == [ + "before:echo", + "exec:a", + "after:echo", + "before:echo", + "exec:b", + "after:echo", + ] + types = [e.type for e in events] + tool_end_idxs = [i for i, t in enumerate(types) if t == "tool_execution_end"] + assert len(tool_end_idxs) == 2 + assert types[tool_end_idxs[0] + 1 : tool_end_idxs[0] + 3] == [ + "message_start", + "message_end", + ] + first_result = events[tool_end_idxs[0] + 1] + assert isinstance(first_result, MessageStartEvent) + assert isinstance(first_result.message, ToolResultMessage) + assert first_result.message.content == "after:ok:a" + assert events[-1].type == "agent_end" + + +async def test_parallel_emits_end_in_completion_order_results_in_source_order() -> None: + started = asyncio.Event() + release_slow = asyncio.Event() + + async def execute_slow(tool_call_id: str, args: dict, **_kwargs) -> AgentToolResult: + started.set() + await release_slow.wait() + return AgentToolResult(content="slow") + + async def execute_fast(tool_call_id: str, args: dict, **_kwargs) -> AgentToolResult: + await started.wait() + return AgentToolResult(content="fast") + + tools = [ + AgentTool(name="slow", execute=execute_slow), + AgentTool(name="fast", execute=execute_fast), + ] + calls = [ + ToolCall(id="s", name="slow", arguments={}), + ToolCall(id="f", name="fast", arguments={}), + ] + responses = [ + AssistantMessage(content=None, tool_calls=calls, stop_reason="toolUse"), + AssistantMessage(content="done", stop_reason="stop"), + ] + + async def stream_fn( + request: StreamRequest, + ) -> AsyncIterator[AssistantStreamEvent]: + msg = responses.pop(0) + yield AssistantStreamStart(partial=AssistantMessage()) + yield AssistantStreamDone(message=msg) + + async def collect() -> list: + return [ + e + async for e in agent_loop( + [UserMessage(content="go")], + AgentContext(tools=tools), + AgentLoopConfig( + convert_to_llm=default_convert_to_llm, + tool_execution="parallel", + ), + stream_fn=stream_fn, + ) + ] + + task = asyncio.create_task(collect()) + await started.wait() + await asyncio.sleep(0.01) + release_slow.set() + events = await task + + end_names = [e.tool_name for e in events if isinstance(e, ToolExecutionEndEvent)] + assert end_names == ["fast", "slow"] + + result_msgs = [ + e.message + for e in events + if isinstance(e, MessageStartEvent) and isinstance(e.message, ToolResultMessage) + ] + assert [m.tool_name for m in result_msgs] == ["slow", "fast"] + assert [m.content for m in result_msgs] == ["slow", "fast"] + + +async def test_missing_required_argument_becomes_error_result() -> None: + async def execute_boom(_id: str, _args: dict, **_kwargs) -> AgentToolResult: + raise AssertionError("should not run") + + responses = [ + AssistantMessage( + content=None, + tool_calls=[ToolCall(id="1", name="need", arguments={})], + stop_reason="toolUse", + ), + AssistantMessage(content="done", stop_reason="stop"), + ] + + async def stream_fn( + request: StreamRequest, + ) -> AsyncIterator[AssistantStreamEvent]: + msg = responses.pop(0) + yield AssistantStreamStart(partial=AssistantMessage()) + yield AssistantStreamDone(message=msg) + + events = [ + e + async for e in agent_loop( + [UserMessage(content="go")], + AgentContext( + tools=[ + AgentTool( + name="need", + parameters={"required": ["path"]}, + execute=execute_boom, + ) + ] + ), + AgentLoopConfig( + convert_to_llm=default_convert_to_llm, + tool_execution="sequential", + ), + stream_fn=stream_fn, + ) + ] + ends = [e for e in events if isinstance(e, ToolExecutionEndEvent)] + assert ends[0].is_error is True + assert "path" in ends[0].result.content + + +async def test_before_hook_can_block_tool() -> None: + async def execute_boom(_id: str, _args: dict, **_kwargs) -> AgentToolResult: + raise AssertionError("should not run") + + async def before(_ctx: BeforeToolCallContext) -> BeforeToolCallResult: + return BeforeToolCallResult(block=True, reason="nope") + + responses = [ + AssistantMessage( + content=None, + tool_calls=[ToolCall(id="1", name="boom", arguments={})], + stop_reason="toolUse", + ), + AssistantMessage(content="done", stop_reason="stop"), + ] + + async def stream_fn( + request: StreamRequest, + ) -> AsyncIterator[AssistantStreamEvent]: + msg = responses.pop(0) + yield AssistantStreamStart(partial=AssistantMessage()) + yield AssistantStreamDone(message=msg) + + events = [ + e + async for e in agent_loop( + [UserMessage(content="go")], + AgentContext(tools=[AgentTool(name="boom", execute=execute_boom)]), + AgentLoopConfig( + convert_to_llm=default_convert_to_llm, + tool_execution="sequential", + before_tool_call=before, + ), + stream_fn=stream_fn, + ) + ] + ends = [e for e in events if isinstance(e, ToolExecutionEndEvent)] + assert ends[0].is_error is True + assert ends[0].result.content == "nope" + + +async def test_prompt_with_tools_emits_full_settled_event_sequence() -> None: + async def execute(_id: str, _args: dict, **_k) -> AgentToolResult: + return AgentToolResult(content="tool-out") + + responses = [ + AssistantMessage( + content=None, + tool_calls=[ToolCall(id="1", name="echo", arguments={"x": 1})], + stop_reason="toolUse", + ), + AssistantMessage(content="final", stop_reason="stop"), + ] + + async def stream_fn( + request: StreamRequest, + ) -> AsyncIterator[AssistantStreamEvent]: + msg = responses.pop(0) + yield AssistantStreamStart(partial=AssistantMessage()) + yield AssistantStreamDone(message=msg) + + events = [ + e + async for e in agent_loop( + [UserMessage(content="go")], + AgentContext(tools=[AgentTool(name="echo", execute=execute)]), + AgentLoopConfig( + convert_to_llm=default_convert_to_llm, + tool_execution="sequential", + ), + stream_fn=stream_fn, + ) + ] + + assert [e.type for e in events] == [ + "agent_start", + "turn_start", + "message_start", + "message_end", + "message_start", + "message_end", + "tool_execution_start", + "tool_execution_end", + "message_start", + "message_end", + "turn_end", + "turn_start", + "message_start", + "message_end", + "turn_end", + "agent_end", + ] + + +async def test_batch_terminate_skips_follow_up_llm_call() -> None: + stream_calls = 0 + + async def execute(_id: str, _args: dict, **_k) -> AgentToolResult: + return AgentToolResult(content="bye", terminate=True) + + async def stream_fn( + request: StreamRequest, + ) -> AsyncIterator[AssistantStreamEvent]: + nonlocal stream_calls + stream_calls += 1 + msg = AssistantMessage( + content=None, + tool_calls=[ToolCall(id="1", name="done", arguments={})], + stop_reason="toolUse", + ) + yield AssistantStreamStart(partial=AssistantMessage()) + yield AssistantStreamDone(message=msg) + + events = [ + e + async for e in agent_loop( + [UserMessage(content="go")], + AgentContext(tools=[AgentTool(name="done", execute=execute)]), + AgentLoopConfig( + convert_to_llm=default_convert_to_llm, + tool_execution="sequential", + ), + stream_fn=stream_fn, + ) + ] + + assert stream_calls == 1 + assert isinstance(events[-1], AgentEndEvent) + assert isinstance(events[-2], TurnEndEvent)