diff --git a/litellm/managed_agents/__init__.py b/litellm/managed_agents/__init__.py new file mode 100644 index 00000000000..32fbc4ade73 --- /dev/null +++ b/litellm/managed_agents/__init__.py @@ -0,0 +1,69 @@ +""" +``litellm.managed_agents`` — Python SDK for spawning and driving managed agents. + +Public surface: + + * ``Agent`` — handle for one ``LiteLLM_Agent`` row; + spawns sessions. + * ``Session`` — one running session; ``send(prompt)`` + drives the runtime and yields events. + * ``Run`` — read-only view of one ``LiteLLM_AgentRun`` + row; ``stream(starting_seq=N)`` replays + persisted events. + * ``Event`` — the dataclass yielded by both runtimes + and ``Run.stream``. + * ``AgentRuntime`` (+ subclasses ``ClaudeSDKAgentRuntime``, + ``LiteLLMAgentRuntime``) — the LLM tool-loop driver. Subclass to + add ``before_tool_call`` / ``after_tool_call`` hooks. + * ``Sandbox`` (+ subclasses ``LocalSandbox``, ``EC2SandboxViaSSM``) + — where tool calls actually execute. + +Typical use:: + + from litellm.managed_agents import ( + Agent, LiteLLMAgentRuntime, LocalSandbox, + ) + + agent = await Agent.from_db_row( + row, db=prisma_client.db, + runtime=LiteLLMAgentRuntime(), + sandbox=LocalSandbox(), + ) + session = await agent.create_session() + async for event in session.send("hello"): + print(event.type, event.data) +""" + +from litellm.managed_agents.agent import Agent +from litellm.managed_agents.agent_runtime import ( + AgentConfig, + AgentRuntime, + ClaudeSDKAgentRuntime, + LiteLLMAgentRuntime, + SessionState, +) +from litellm.managed_agents.events import Event +from litellm.managed_agents.run import Run +from litellm.managed_agents.sandbox import ( + EC2SandboxViaSSM, + LocalSandbox, + Sandbox, + ToolResult, +) +from litellm.managed_agents.session import Session + +__all__ = [ + "Agent", + "Session", + "Run", + "Event", + "AgentRuntime", + "AgentConfig", + "SessionState", + "ClaudeSDKAgentRuntime", + "LiteLLMAgentRuntime", + "Sandbox", + "ToolResult", + "LocalSandbox", + "EC2SandboxViaSSM", +] diff --git a/litellm/managed_agents/agent.py b/litellm/managed_agents/agent.py new file mode 100644 index 00000000000..ca7542d4cd7 --- /dev/null +++ b/litellm/managed_agents/agent.py @@ -0,0 +1,331 @@ +""" +Agent — Python-side handle for one managed-agent definition. + +This is the SDK-side counterpart to the public Agent CRUD HTTP endpoints +(``/v2/agents``). The proxy persists every agent to ``LiteLLM_Agent``; +``Agent`` is the typed, behaviour-rich accessor in-process code uses to +build sessions, list past sessions, and tear an agent down. + +Boundary with the proxy +----------------------- +The proxy owns the HTTP surface and ownership/auth checks. The +``managed_agents`` layer (this file + ``Session``, ``Run``) is a pure +Python API on top of the same DB tables. They co-exist deliberately: +the HTTP endpoints are how external SDKs talk to the proxy, but +in-process callers (sub-agents, tests, internal tools) shouldn't have +to round-trip through HTTP just to spawn a session. + +Lifecycle: + * ``Agent.from_db_row(row, db, runtime, sandbox)`` — build from a + Prisma row (the proxy wiring layer is responsible for fetching the + row, applying ownership checks, then handing both off to us). + * ``await agent.create_session(repos=..., env_vars=...)`` — INSERT a + session row in ``ready`` status, mint a daemon JWT, and return a + ``Session`` ready to ``send()`` prompts to. + * ``await agent.get_session(session_id)`` — fetch one of the agent's + existing sessions as a ``Session`` instance. + * ``await agent.list_sessions()`` — list all of the agent's sessions. + * ``await agent.delete()`` — hard-delete the agent row. Cascade rules + on ``LiteLLM_Agent`` -> ``LiteLLM_AgentSession`` mean the DB will + drop the sessions and runs too. + +Why ``status=ready`` and not ``provisioning``? +----------------------------------------------- +The proxy's HTTP path uses ``provisioning`` because it kicks off a real +VM provider call in the background (EC2 cold boot can take ~60s). The +``managed_agents`` Python layer doesn't talk to a VM provider — the +sandbox is constructed in-process and is "ready" the moment it exists. +We mint the same daemon JWT for the same ``LITELLM_AGENT_JWT_SECRET`` +auth surface so tests / internal callers can hit the existing daemon +endpoints if they want, but they don't have to. +""" + +from dataclasses import dataclass, field +from datetime import datetime, timedelta, timezone +from typing import Any, Dict, List, Optional + +import prisma + +from litellm.managed_agents.agent_runtime.base import ( + AgentConfig, + AgentRuntime, +) +from litellm.managed_agents.sandbox.base import Sandbox +from litellm.managed_agents.session import Session +from litellm.proxy.agent_session_endpoints.auth import ( + hash_daemon_token, + mint_daemon_token, +) +from litellm.proxy.agent_session_endpoints.constants import ( + DEFAULT_MAX_SESSION_MINUTES, + SESSION_STATUS_READY, +) +from litellm.proxy.agent_session_endpoints.ids import new_session_id + + +def _now() -> datetime: + return datetime.now(timezone.utc) + + +@dataclass +class Agent: + """A live, in-process handle on one ``LiteLLM_Agent`` row. + + Holds the static agent definition (name, model, system_prompt, + tools_config) plus the runtime + sandbox that any session it spawns + will use. Tests can construct an ``Agent`` directly with a mock + runtime and sandbox; production code goes through ``from_db_row``. + + The same ``runtime`` and ``sandbox`` instance get reused across all + sessions this Agent spawns. That's a deliberate choice: most runtimes + are stateless, and most sandboxes either are too (``LocalSandbox`` + creates its own per-instance tmpdir) or document explicitly when + they're not. Callers that need per-session isolation should + construct a fresh ``Agent`` per session. + """ + + id: str + name: str + model: str + system_prompt: Optional[str] = None + tools_config: Optional[Dict[str, Any]] = None + metadata: Dict[str, Any] = field(default_factory=dict) + default_repos: List[Dict[str, Any]] = field(default_factory=list) + default_env_vars: Dict[str, str] = field(default_factory=dict) + user_api_key_hash: str = "" + team_id: Optional[str] = None + runtime: Optional[AgentRuntime] = None + sandbox: Optional[Sandbox] = None + db: Any = None + + # ------------------------------------------------------------------ + # Derived helpers + # ------------------------------------------------------------------ + + def to_runtime_config(self) -> AgentConfig: + """Project this row into the smaller ``AgentConfig`` the runtime sees.""" + return AgentConfig( + name=self.name, + model=self.model, + system_prompt=self.system_prompt, + tools_config=self.tools_config, + metadata=dict(self.metadata), + ) + + @classmethod + async def from_db_row( + cls, + row: Any, + db: Any, + runtime: Optional[AgentRuntime] = None, + sandbox: Optional[Sandbox] = None, + ) -> "Agent": + """Build an ``Agent`` from a Prisma ``LiteLLM_Agent`` row. + + ``runtime`` and ``sandbox`` are required for any session-spawning + operation; pass ``None`` for read-only flows (e.g. ``list_sessions`` + / ``get_session`` from a CLI that just wants to inspect history). + """ + return cls( + id=getattr(row, "id"), + name=getattr(row, "name"), + model=getattr(row, "model"), + system_prompt=getattr(row, "system_prompt", None), + tools_config=_coerce_dict_or_none(getattr(row, "tools_config", None)), + metadata=_coerce_dict(getattr(row, "metadata", None)), + default_repos=_coerce_list_of_dict(getattr(row, "default_repos", None)), + default_env_vars=_coerce_dict(getattr(row, "default_env_vars", None)), + user_api_key_hash=getattr(row, "user_api_key_hash", "") or "", + team_id=getattr(row, "team_id", None), + runtime=runtime, + sandbox=sandbox, + db=db, + ) + + # ------------------------------------------------------------------ + # Session lifecycle + # ------------------------------------------------------------------ + + async def create_session( + self, + repos: Optional[List[Dict[str, Any]]] = None, + env_vars: Optional[Dict[str, str]] = None, + max_session_minutes: int = DEFAULT_MAX_SESSION_MINUTES, + ) -> Session: + """Insert a session row and return a ``Session`` ready to send prompts. + + ``repos`` / ``env_vars`` overlay the agent defaults exactly the + same way the HTTP endpoint does: + * ``repos``: caller-provided wholly replaces the default list. + * ``env_vars``: merged key by key, caller wins on collisions. + + Status starts at ``ready`` (not ``provisioning``) because the + managed-agents layer owns the sandbox lifecycle directly — there + is no remote VM to wait for. See module docstring. + """ + self._require_db() + self._require_runtime_and_sandbox(operation="create_session") + + resolved_repos = self._resolve_repos(repos) + resolved_env_vars = self._resolve_env_vars(env_vars) + + session_id = new_session_id() + expires_at = _now() + timedelta(minutes=max_session_minutes) + daemon_token = mint_daemon_token( + session_id=session_id, + agent_id=self.id, + expires_at_epoch=int(expires_at.timestamp()), + ) + + # Mirrors the HTTP path's payload shape so the same row format + # works whether the session was created via /v2/sessions or via + # this Python API. Json columns must be wrapped via prisma.Json, + # relations must use {"connect": {"id": ...}}. + payload: Dict[str, Any] = { + "id": session_id, + "agent": {"connect": {"id": self.id}}, + "user_api_key_hash": self.user_api_key_hash, + "team_id": self.team_id, + "repos": prisma.Json(resolved_repos or []), + "status": SESSION_STATUS_READY, + "daemon_token_hash": hash_daemon_token(daemon_token), + "expires_at": expires_at, + "updated_at": _now(), + } + if resolved_env_vars is not None: + payload["env_vars"] = prisma.Json(resolved_env_vars) + + row = await self.db.litellm_agentsession.create(data=payload) + + return await Session.from_db_row( + row=row, + db=self.db, + runtime=self.runtime, + sandbox=self.sandbox, + agent_config=self.to_runtime_config(), + daemon_token=daemon_token, + ) + + async def get_session(self, session_id: str) -> Session: + """Fetch one of this agent's sessions by id. + + Raises ``LookupError`` if the session doesn't exist or doesn't + belong to this agent — callers that need richer auth handling + should wrap and remap to HTTP errors. + """ + self._require_db() + row = await self.db.litellm_agentsession.find_unique(where={"id": session_id}) + if row is None or getattr(row, "agent_id", None) != self.id: + raise LookupError(f"Session {session_id!r} not found for agent {self.id!r}") + return await Session.from_db_row( + row=row, + db=self.db, + runtime=self.runtime, + sandbox=self.sandbox, + agent_config=self.to_runtime_config(), + daemon_token=None, + ) + + async def list_sessions(self) -> List[Session]: + """Return every session this agent has spawned, newest first. + + No pagination here — call sites that expect lots of sessions + should query ``self.db.litellm_agentsession`` directly with + ``take``/``skip``. This helper is a convenience for tests and + small-scale callers. + """ + self._require_db() + rows = await self.db.litellm_agentsession.find_many( + where={"agent_id": self.id}, + order={"created_at": "desc"}, + ) + return [ + await Session.from_db_row( + row=row, + db=self.db, + runtime=self.runtime, + sandbox=self.sandbox, + agent_config=self.to_runtime_config(), + daemon_token=None, + ) + for row in rows + ] + + async def delete(self) -> None: + """Hard-delete the agent row. + + Cascade rules in ``schema.prisma`` (``LiteLLM_AgentSession`` + ``onDelete: Cascade``, then ``LiteLLM_AgentRun`` and + ``LiteLLM_AgentRunEvent`` cascading from there) drop everything + downstream — no need to enumerate sessions here. + """ + self._require_db() + await self.db.litellm_agent.delete(where={"id": self.id}) + + # ------------------------------------------------------------------ + # Internals + # ------------------------------------------------------------------ + + def _resolve_repos( + self, + body_repos: Optional[List[Dict[str, Any]]], + ) -> List[Dict[str, Any]]: + """Caller-provided repos override agent defaults entirely (whole-list + replace). Mirrors ``_resolve_repos`` in ``session_endpoints.py``. + """ + if body_repos is not None: + return [r for r in body_repos if isinstance(r, dict)] + return list(self.default_repos) + + def _resolve_env_vars( + self, + body_env_vars: Optional[Dict[str, str]], + ) -> Optional[Dict[str, str]]: + """Merge: agent defaults first, caller overrides on top.""" + if body_env_vars is None and not self.default_env_vars: + return None + merged: Dict[str, str] = {} + if self.default_env_vars: + merged.update({str(k): str(v) for k, v in self.default_env_vars.items()}) + if body_env_vars: + merged.update({str(k): str(v) for k, v in body_env_vars.items()}) + return merged or None + + def _require_db(self) -> None: + if self.db is None: + raise RuntimeError( + "Agent operation requires a Prisma client; pass db= " + "via Agent.from_db_row(...) or set self.db before calling." + ) + + def _require_runtime_and_sandbox(self, operation: str) -> None: + if self.runtime is None or self.sandbox is None: + raise RuntimeError( + f"Agent.{operation} requires both runtime and sandbox; " + "construct the Agent via Agent.from_db_row(row, db, runtime, sandbox)." + ) + + +# --------------------------------------------------------------------------- +# Prisma JSON-column coercion helpers — same shape as ``serialization.py`` +# but kept local so this module doesn't depend on the proxy serializer +# (which exists for the HTTP wire shape, not the Python API). +# --------------------------------------------------------------------------- + + +def _coerce_dict(value: Any) -> Dict[str, Any]: + if isinstance(value, dict): + return value + return {} + + +def _coerce_dict_or_none(value: Any) -> Optional[Dict[str, Any]]: + if isinstance(value, dict): + return value + return None + + +def _coerce_list_of_dict(value: Any) -> List[Dict[str, Any]]: + if isinstance(value, list): + return [v for v in value if isinstance(v, dict)] + return [] diff --git a/litellm/managed_agents/agent_runtime/base.py b/litellm/managed_agents/agent_runtime/base.py new file mode 100644 index 00000000000..201e2f40e8d --- /dev/null +++ b/litellm/managed_agents/agent_runtime/base.py @@ -0,0 +1,120 @@ +""" +AgentRuntime ABC — drives the LLM tool loop. + +Implementations: + + * ``ClaudeSDKAgentRuntime`` — wraps Anthropic's ``claude-agent-sdk`` + * ``LiteLLMAgentRuntime`` — uses ``litellm.acompletion`` for multi-provider + +Customers can subclass either to hook into ``before_tool_call`` / +``after_tool_call`` for audit, telemetry, or tool-result rewriting. + +The runtime sees ``AgentConfig`` (the static agent definition: model, +system_prompt, tools_config) and ``SessionState`` (the live per-session +state: cwd hint, env_vars, repos), but never the DB row directly. That +keeps the runtime layer pure and testable without a Prisma client. +""" + +from abc import ABC, abstractmethod +from dataclasses import dataclass, field +from typing import Any, AsyncIterator, Dict, List, Optional + +from litellm.managed_agents.events import Event +from litellm.managed_agents.sandbox.base import Sandbox + + +@dataclass +class AgentConfig: + """Static agent config the runtime needs to drive an LLM tool loop. + + Mirrors the columns of ``LiteLLM_Agent`` that matter at runtime; the + DB row stays in ``Agent`` (the public class), the runtime sees only + this snapshot. Decoupling the two means a runtime can be unit tested + with a literal ``AgentConfig(...)`` and no DB at all. + """ + + name: str + model: str + system_prompt: Optional[str] = None + tools_config: Optional[Dict[str, Any]] = None + metadata: Dict[str, Any] = field(default_factory=dict) + + +@dataclass +class SessionState: + """Live per-session state. + + ``cwd`` is the working directory for tool execution (provided by the + sandbox). ``env_vars`` and ``repos`` are exposed so the runtime can + seed its system prompt with context (e.g. "you have access to repos + X, Y at /workspace/foo, /workspace/bar"). + """ + + session_id: str + cwd: Optional[str] = None + env_vars: Dict[str, str] = field(default_factory=dict) + repos: List[Dict[str, Any]] = field(default_factory=list) + + +class AgentRuntime(ABC): + """Drives the LLM tool loop. Yields ``Event`` instances. + + The contract: + + * ``run(prompt, sandbox, session_state, agent_config)`` is an + async iterator that yields events as the LLM produces them. + * The runtime decides when to stop (typically on a terminal + ``run_finished`` event from the LLM). + * The runtime is responsible for routing tool calls through the + provided ``sandbox`` (with the documented ClaudeSDKAgentRuntime + exception — see its docstring). + * ``before_tool_call`` and ``after_tool_call`` are optional hooks + subclasses can override to inject behaviour without forking the + whole runtime. Default implementations are no-ops. + """ + + @abstractmethod + async def run( + self, + prompt: str, + sandbox: Sandbox, + session_state: SessionState, + agent_config: AgentConfig, + ) -> AsyncIterator[Event]: + """Drive the LLM tool loop and yield events.""" + # The ``yield`` here makes this an async generator (so the ABC + # signature matches subclasses). Subclasses MUST implement. + if False: + yield Event(type="placeholder", data={}) + raise NotImplementedError + + async def before_tool_call( + self, + tool_name: str, + tool_input: Dict[str, Any], + ) -> Dict[str, Any]: + """Hook: called right before each tool execution. + + Return the (possibly rewritten) ``tool_input``. Subclasses can + override to log, validate, or rewrite tool calls before they hit + the sandbox. Default returns ``tool_input`` unchanged. + + Raise to abort the tool call entirely — the runtime should catch + and surface the failure as a ``tool_result`` with ``is_error=True``. + """ + return tool_input + + async def after_tool_call( + self, + tool_name: str, + tool_input: Dict[str, Any], + result: Any, + ) -> Any: + """Hook: called right after each tool execution. + + Return the (possibly rewritten) ``result``. Subclasses can + override to log, redact, or rewrite tool outputs before they get + shown to the LLM on its next turn. Default returns ``result`` + unchanged. + """ + return result diff --git a/litellm/managed_agents/agent_runtime/claude_sdk.py b/litellm/managed_agents/agent_runtime/claude_sdk.py new file mode 100644 index 00000000000..b9490019641 --- /dev/null +++ b/litellm/managed_agents/agent_runtime/claude_sdk.py @@ -0,0 +1,220 @@ +""" +ClaudeSDKAgentRuntime — wraps Anthropic's ``claude-agent-sdk``. + +This is the default runtime. It delegates the full LLM tool loop (model +choice, tool calling, hooks, MCP, sub-agents) to ``claude-agent-sdk`` and +just translates the SDK's message stream into our ``Event`` shape. + +Sandbox interop note (read this before assuming things) +------------------------------------------------------- + +``claude-agent-sdk`` runs its built-in tools (``Read``, ``Write``, +``Edit``, ``Bash``, etc.) IN-PROCESS, in its own subprocess. Routing +those through our ``Sandbox`` abstraction would require overriding each +built-in with a custom MCP-backed tool that calls into the sandbox. + +For v1 we explicitly do NOT do that: ``ClaudeSDKAgentRuntime`` works +correctly only with ``LocalSandbox`` (the SDK runs locally, sandbox-side +state == process-side state). For remote-execution agents that need to +run on EC2 / Docker, callers should use ``LiteLLMAgentRuntime`` instead, +which fully honors the ``Sandbox`` abstraction. The MCP-backed override +is tracked as a follow-up. + +We do still consume ``sandbox.cwd`` though — ``ClaudeAgentOptions(cwd=...)`` +controls where the SDK runs its built-in tools, so a ``LocalSandbox`` +configured with ``working_dir="/tmp/foo"`` will have the SDK do its +work there. That's enough to support "create file foo.txt" tests. +""" + +from typing import Any, AsyncIterator, Dict, Optional + +from litellm.managed_agents.agent_runtime.base import ( + AgentConfig, + AgentRuntime, + SessionState, +) +from litellm.managed_agents.events import ( + EVENT_TYPE_ASSISTANT_MESSAGE, + EVENT_TYPE_RUN_FINISHED, + EVENT_TYPE_SYSTEM, + EVENT_TYPE_THINKING, + EVENT_TYPE_TOOL_RESULT, + EVENT_TYPE_TOOL_USE, + Event, +) +from litellm.managed_agents.sandbox.base import Sandbox + + +class ClaudeSDKAgentRuntime(AgentRuntime): + """Default runtime — wraps ``claude-agent-sdk``. + + Construct with optional overrides; per-run values fall back to the + ``AgentConfig`` passed to ``run()``. This split lets callers either + (a) bake everything into the runtime instance and reuse it across + agents, or (b) construct a fresh runtime per agent that defers + everything to the agent config. + + ``permission_mode`` defaults to ``"bypassPermissions"`` so the SDK + doesn't prompt for tool approval — agents are typically invoked + headlessly. Override to ``"default"`` for interactive flows. + """ + + def __init__( + self, + model: Optional[str] = None, + system_prompt: Optional[str] = None, + permission_mode: str = "bypassPermissions", + max_turns: Optional[int] = None, + extra_options: Optional[Dict[str, Any]] = None, + ) -> None: + self.model = model + self.system_prompt = system_prompt + self.permission_mode = permission_mode + self.max_turns = max_turns + self.extra_options = dict(extra_options or {}) + + def _build_options( + self, + sandbox: Sandbox, + session_state: SessionState, + agent_config: AgentConfig, + ): + """Translate (config + state + sandbox) into a ClaudeAgentOptions.""" + from claude_agent_sdk import ClaudeAgentOptions + + opts: Dict[str, Any] = { + "model": self.model or agent_config.model, + "system_prompt": self.system_prompt or agent_config.system_prompt, + "permission_mode": self.permission_mode, + } + if self.max_turns is not None: + opts["max_turns"] = self.max_turns + + cwd = sandbox.cwd or session_state.cwd + if cwd is not None: + opts["cwd"] = cwd + + if session_state.env_vars: + opts["env"] = dict(session_state.env_vars) + + # Drop any unset / unsupported keys before constructing. + opts = {k: v for k, v in opts.items() if v is not None} + + # Caller can layer on raw SDK options too (e.g. ``allowed_tools``, + # ``hooks``, ``mcp_servers``). These win over our derived defaults. + opts.update(self.extra_options) + + return ClaudeAgentOptions(**opts) + + async def run( + self, + prompt: str, + sandbox: Sandbox, + session_state: SessionState, + agent_config: AgentConfig, + ) -> AsyncIterator[Event]: + from claude_agent_sdk import ( + AssistantMessage, + ResultMessage, + SystemMessage, + TextBlock, + ThinkingBlock, + ToolResultBlock, + ToolUseBlock, + UserMessage, + query, + ) + + # Make sure the sandbox is ready (creates tmpdir for LocalSandbox). + await sandbox.setup() + + options = self._build_options(sandbox, session_state, agent_config) + + async for message in query(prompt=prompt, options=options): + # AssistantMessage carries an array of content blocks; each + # block becomes its own event (assistant_message / tool_use / + # thinking) so consumers can render them individually. + if isinstance(message, AssistantMessage): + for block in message.content or []: + event = self._block_to_event(block) + if event is not None: + yield event + elif isinstance(message, UserMessage): + # UserMessage from the SDK is the SDK echoing the tool + # results it just got back (so the LLM saw them on the + # next turn). We surface those as ``tool_result`` events. + for block in message.content or []: + if isinstance(block, ToolResultBlock): + yield Event( + type=EVENT_TYPE_TOOL_RESULT, + data={ + "tool_use_id": block.tool_use_id, + "output": _stringify(block.content), + "is_error": bool(block.is_error), + }, + ) + elif isinstance(message, SystemMessage): + yield Event( + type=EVENT_TYPE_SYSTEM, + data={"subtype": message.subtype, "data": message.data}, + ) + elif isinstance(message, ResultMessage): + yield Event( + type=EVENT_TYPE_RUN_FINISHED, + data={ + "result": message.result, + "is_error": bool(message.is_error), + "stop_reason": message.stop_reason, + "num_turns": message.num_turns, + "duration_ms": message.duration_ms, + "total_cost_usd": message.total_cost_usd, + }, + ) + return + # Other message types (StreamEvent, RateLimitEvent) ignored + # for now — they're noise for our use case. + + @staticmethod + def _block_to_event(block: Any) -> Optional[Event]: + from claude_agent_sdk import TextBlock, ThinkingBlock, ToolUseBlock + + if isinstance(block, TextBlock): + return Event( + type=EVENT_TYPE_ASSISTANT_MESSAGE, + data={"content": block.text}, + ) + if isinstance(block, ToolUseBlock): + return Event( + type=EVENT_TYPE_TOOL_USE, + data={ + "tool_use_id": block.id, + "tool": block.name, + "input": block.input, + }, + ) + if isinstance(block, ThinkingBlock): + return Event( + type=EVENT_TYPE_THINKING, + data={"content": block.thinking}, + ) + return None + + +def _stringify(content: Any) -> str: + """Tool results in ``ToolResultBlock`` come as either a string or a + list of content dicts (the Anthropic SDK shape). Squash to a string + for the wire payload.""" + if isinstance(content, str): + return content + if isinstance(content, list): + parts = [] + for item in content: + if isinstance(item, dict): + if item.get("type") == "text": + parts.append(item.get("text", "")) + else: + parts.append(str(item)) + else: + parts.append(str(item)) + return "\n".join(parts) + return str(content) diff --git a/litellm/managed_agents/agent_runtime/litellm_native.py b/litellm/managed_agents/agent_runtime/litellm_native.py new file mode 100644 index 00000000000..126a36497d5 --- /dev/null +++ b/litellm/managed_agents/agent_runtime/litellm_native.py @@ -0,0 +1,532 @@ +""" +LiteLLMAgentRuntime — drives a manual tool loop via ``litellm.acompletion``. + +Why a second runtime when ClaudeSDKAgentRuntime already exists? + + * Multi-provider. ``litellm.acompletion`` speaks Anthropic, OpenAI, Gemini, + Bedrock, etc. through one unified surface, so the same agent definition + can run against any provider that supports tool calling. + * Sandbox-honest. Unlike ``ClaudeSDKAgentRuntime`` (which runs its built-in + tools in-process), this runtime routes EVERY tool call through + ``sandbox.execute_tool(...)``. That makes ``EC2SandboxViaSSM`` and any + future remote-execution sandbox actually work. + * Hookable. The ``before_tool_call`` / ``after_tool_call`` ABC hooks fire + around every tool call so subclasses can audit, redact, or rewrite + without forking the whole loop. + +Wire shape: events emitted match the snake_case wire format the integration +branch already serves (``assistant_message`` / ``tool_use`` / ``tool_result`` +/ ``run_finished``). + +Tool source of truth: tools are configured on the ``AgentConfig`` via +``tools_config``. Two shapes are accepted: + + * ``{"tools": [, ...]}`` — pre-formatted OpenAI/LiteLLM + tool definitions. Passed straight through to ``acompletion``. + * ``[, ...]`` — bare list, same handling. + +If ``tools_config`` is empty / missing, the runtime falls back to the +default LocalSandbox tool surface (Bash/Read/Write/Edit/ls). That keeps +the "create file foo.txt" smoke test working out of the box without +forcing every config to spell out tool schemas. +""" + +import json +from typing import Any, AsyncIterator, Dict, List, Optional + +import litellm +from litellm.managed_agents.agent_runtime.base import ( + AgentConfig, + AgentRuntime, + SessionState, +) +from litellm.managed_agents.events import ( + EVENT_TYPE_ASSISTANT_MESSAGE, + EVENT_TYPE_RUN_FINISHED, + EVENT_TYPE_TOOL_RESULT, + EVENT_TYPE_TOOL_USE, + Event, +) +from litellm.managed_agents.sandbox.base import Sandbox, ToolResult + + +# Default upper bound on tool-loop iterations. The LLM gets this many +# turns to call tools before we forcibly stop and yield run_finished. +# Picked to match claude-agent-sdk's default; configurable per-instance. +DEFAULT_MAX_TURNS = 25 + + +# Default tool surface — these mirror what LocalSandbox knows how to +# execute. We hand them to providers that need OpenAI-shaped tool defs +# (which is most of them) when the AgentConfig doesn't supply its own. +_DEFAULT_TOOLS: List[Dict[str, Any]] = [ + { + "type": "function", + "function": { + "name": "Bash", + "description": "Run a shell command in the sandbox working directory.", + "parameters": { + "type": "object", + "properties": { + "command": { + "type": "string", + "description": "Shell command to execute.", + }, + }, + "required": ["command"], + }, + }, + }, + { + "type": "function", + "function": { + "name": "Read", + "description": "Read the contents of a file inside the sandbox.", + "parameters": { + "type": "object", + "properties": { + "path": { + "type": "string", + "description": "File path (relative or absolute).", + }, + }, + "required": ["path"], + }, + }, + }, + { + "type": "function", + "function": { + "name": "Write", + "description": "Write text to a file inside the sandbox.", + "parameters": { + "type": "object", + "properties": { + "path": {"type": "string", "description": "File path."}, + "content": { + "type": "string", + "description": "Text content to write.", + }, + }, + "required": ["path", "content"], + }, + }, + }, + { + "type": "function", + "function": { + "name": "Edit", + "description": "Replace one occurrence of old_string with new_string in a file.", + "parameters": { + "type": "object", + "properties": { + "path": {"type": "string"}, + "old_string": {"type": "string"}, + "new_string": {"type": "string"}, + }, + "required": ["path", "old_string", "new_string"], + }, + }, + }, + { + "type": "function", + "function": { + "name": "ls", + "description": "List entries in a sandbox directory.", + "parameters": { + "type": "object", + "properties": { + "path": {"type": "string"}, + }, + }, + }, + }, +] + + +class LiteLLMAgentRuntime(AgentRuntime): + """Manual tool-loop runtime built on ``litellm.acompletion``. + + Construct with optional overrides; per-run values fall back to the + ``AgentConfig`` passed to ``run()``. Same split as + ``ClaudeSDKAgentRuntime`` so callers can either bake everything into + the runtime instance or defer to the agent config. + """ + + def __init__( + self, + model: Optional[str] = None, + system_prompt: Optional[str] = None, + max_turns: int = DEFAULT_MAX_TURNS, + extra_completion_kwargs: Optional[Dict[str, Any]] = None, + ) -> None: + self.model = model + self.system_prompt = system_prompt + self.max_turns = max_turns + self.extra_completion_kwargs = dict(extra_completion_kwargs or {}) + + def _resolve_tools(self, agent_config: AgentConfig) -> List[Dict[str, Any]]: + """Pick the tool list to send to the LLM. + + Accepts either a ``{"tools": [...]}`` wrapper or a bare list, since + both shapes show up in real configs. Falls back to the default + LocalSandbox surface when nothing is configured. + """ + cfg = agent_config.tools_config + if isinstance(cfg, dict) and isinstance(cfg.get("tools"), list): + return list(cfg["tools"]) + if isinstance(cfg, list): + return list(cfg) + return list(_DEFAULT_TOOLS) + + def _build_initial_messages( + self, + prompt: str, + agent_config: AgentConfig, + session_state: SessionState, + ) -> List[Dict[str, Any]]: + """Compose the message list passed to the first ``acompletion`` call. + + We seed an extra ``system`` line describing the live session so the + LLM has cwd/repos context without the agent author having to thread + it into every prompt. + """ + messages: List[Dict[str, Any]] = [] + + sys_text = self.system_prompt or agent_config.system_prompt + if sys_text: + messages.append({"role": "system", "content": sys_text}) + + # Session context line — only added when there's something to say. + ctx_lines: List[str] = [] + if session_state.cwd: + ctx_lines.append(f"Working directory: {session_state.cwd}") + if session_state.repos: + repo_summaries = ", ".join( + str(r.get("url") or r.get("path") or "?") for r in session_state.repos + ) + ctx_lines.append(f"Repositories available: {repo_summaries}") + if ctx_lines: + messages.append({"role": "system", "content": "\n".join(ctx_lines)}) + + messages.append({"role": "user", "content": prompt}) + return messages + + async def _execute_tool_call( + self, + sandbox: Sandbox, + tool_name: str, + tool_input: Dict[str, Any], + ) -> ToolResult: + """Run a single tool call through the hooks + sandbox. + + Hook failures bubble out as ``ToolResult(is_error=True)`` rather + than raising, so the LLM sees the failure on its next turn. + """ + try: + tool_input = await self.before_tool_call(tool_name, tool_input) + except Exception as exc: # noqa: BLE001 — surface to LLM + return ToolResult( + output=f"before_tool_call hook raised: {exc}", + is_error=True, + metadata={"hook": "before_tool_call"}, + ) + + result = await sandbox.execute_tool(tool_name, tool_input) + + try: + rewritten = await self.after_tool_call(tool_name, tool_input, result) + except Exception as exc: # noqa: BLE001 + return ToolResult( + output=f"after_tool_call hook raised: {exc}", + is_error=True, + metadata={"hook": "after_tool_call"}, + ) + # Hook may return a fully replaced ToolResult or a bare value; only + # treat the former as authoritative. + if isinstance(rewritten, ToolResult): + return rewritten + return result + + async def run( + self, + prompt: str, + sandbox: Sandbox, + session_state: SessionState, + agent_config: AgentConfig, + ) -> AsyncIterator[Event]: + await sandbox.setup() + + model = self.model or agent_config.model + tools = self._resolve_tools(agent_config) + messages = self._build_initial_messages(prompt, agent_config, session_state) + + completion_kwargs: Dict[str, Any] = { + "model": model, + "messages": messages, + } + if tools: + completion_kwargs["tools"] = tools + completion_kwargs.update(self.extra_completion_kwargs) + + last_text: Optional[str] = None + + for turn in range(self.max_turns): + response = await litellm.acompletion(**completion_kwargs) + + choice = _first_choice(response) + if choice is None: + # Defensive — provider returned no choices, treat as done. + yield Event( + type=EVENT_TYPE_RUN_FINISHED, + data={ + "result": last_text, + "is_error": False, + "stop_reason": "no_choice", + "num_turns": turn + 1, + }, + ) + return + + assistant_msg = _choice_message(choice) + content = _message_content(assistant_msg) + tool_calls = _message_tool_calls(assistant_msg) + + if content: + last_text = content + yield Event( + type=EVENT_TYPE_ASSISTANT_MESSAGE, + data={"content": content}, + ) + + # Append the assistant turn to the conversation BEFORE handling + # tool calls — the next request needs the assistant's tool_calls + # array to interleave correctly with tool messages. + messages.append(_assistant_dict(assistant_msg)) + + # No tool call -> the LLM is done. Emit run_finished and stop. + if not tool_calls: + yield Event( + type=EVENT_TYPE_RUN_FINISHED, + data={ + "result": last_text, + "is_error": False, + "stop_reason": _choice_finish_reason(choice) or "stop", + "num_turns": turn + 1, + }, + ) + return + + for call in tool_calls: + tool_use_id = _tool_call_id(call) + tool_name = _tool_call_name(call) + tool_input = _tool_call_arguments(call) + + yield Event( + type=EVENT_TYPE_TOOL_USE, + data={ + "tool_use_id": tool_use_id, + "tool": tool_name, + "input": tool_input, + }, + ) + + result = await self._execute_tool_call(sandbox, tool_name, tool_input) + output_str = _stringify(result.output) + + yield Event( + type=EVENT_TYPE_TOOL_RESULT, + data={ + "tool_use_id": tool_use_id, + "output": output_str, + "is_error": result.is_error, + }, + ) + + # Feed the result back into the conversation so the LLM + # sees it on the next turn. OpenAI shape: role=tool with a + # tool_call_id linking back to the assistant's call. + messages.append( + { + "role": "tool", + "tool_call_id": tool_use_id, + "name": tool_name, + "content": output_str, + } + ) + + # Hit max_turns without the LLM signalling done. Yield a terminal + # event so the caller's run row gets closed — surface as a finished + # run with stop_reason=max_turns rather than an error (matches the + # claude-agent-sdk behaviour). + yield Event( + type=EVENT_TYPE_RUN_FINISHED, + data={ + "result": last_text, + "is_error": False, + "stop_reason": "max_turns", + "num_turns": self.max_turns, + }, + ) + + +# --------------------------------------------------------------------------- +# Response shape adapters. +# +# litellm.acompletion can return either ModelResponse (pydantic-ish) or a +# plain dict depending on caller config and provider. The helpers below +# normalise both shapes so the main loop above doesn't have to care. +# --------------------------------------------------------------------------- + + +def _first_choice(response: Any) -> Optional[Any]: + choices = getattr(response, "choices", None) + if choices is None and isinstance(response, dict): + choices = response.get("choices") + if not choices: + return None + return choices[0] + + +def _choice_message(choice: Any) -> Any: + msg = getattr(choice, "message", None) + if msg is None and isinstance(choice, dict): + msg = choice.get("message") + return msg + + +def _choice_finish_reason(choice: Any) -> Optional[str]: + reason = getattr(choice, "finish_reason", None) + if reason is None and isinstance(choice, dict): + reason = choice.get("finish_reason") + return reason + + +def _message_content(message: Any) -> Optional[str]: + if message is None: + return None + content = getattr(message, "content", None) + if content is None and isinstance(message, dict): + content = message.get("content") + if isinstance(content, list): + # Anthropic-via-litellm sometimes returns content as a list of + # blocks; concat any text blocks for the assistant_message event. + parts = [] + for block in content: + if isinstance(block, dict) and block.get("type") == "text": + parts.append(block.get("text", "")) + return "\n".join(p for p in parts if p) or None + if isinstance(content, str) and content: + return content + return None + + +def _message_tool_calls(message: Any) -> List[Any]: + if message is None: + return [] + calls = getattr(message, "tool_calls", None) + if calls is None and isinstance(message, dict): + calls = message.get("tool_calls") + return list(calls or []) + + +def _assistant_dict(message: Any) -> Dict[str, Any]: + """Serialize the assistant message back into the shape acompletion expects. + + We can't always re-send the raw ModelResponse object — the next call + needs a plain dict with ``role``, ``content``, and optional + ``tool_calls`` (each with ``id``, ``type``, ``function``). + """ + out: Dict[str, Any] = {"role": "assistant"} + content = _message_content(message) or "" + out["content"] = content + calls = _message_tool_calls(message) + if calls: + out["tool_calls"] = [_tool_call_to_dict(c) for c in calls] + return out + + +def _tool_call_to_dict(call: Any) -> Dict[str, Any]: + if isinstance(call, dict): + fn = call.get("function") or {} + if not isinstance(fn, dict): + fn = { + "name": getattr(fn, "name", None), + "arguments": getattr(fn, "arguments", None), + } + return { + "id": call.get("id"), + "type": call.get("type", "function"), + "function": { + "name": fn.get("name"), + "arguments": fn.get("arguments", "{}"), + }, + } + fn = getattr(call, "function", None) + return { + "id": getattr(call, "id", None), + "type": getattr(call, "type", "function"), + "function": { + "name": getattr(fn, "name", None) if fn is not None else None, + "arguments": getattr(fn, "arguments", "{}") if fn is not None else "{}", + }, + } + + +def _tool_call_id(call: Any) -> str: + if isinstance(call, dict): + return str(call.get("id") or "") + return str(getattr(call, "id", "") or "") + + +def _tool_call_name(call: Any) -> str: + if isinstance(call, dict): + fn = call.get("function") or {} + if isinstance(fn, dict): + return str(fn.get("name") or "") + return str(getattr(fn, "name", "") or "") + fn = getattr(call, "function", None) + return str(getattr(fn, "name", "") or "") if fn is not None else "" + + +def _tool_call_arguments(call: Any) -> Dict[str, Any]: + """Tool arguments come back as a JSON string — parse to dict. + + Defensive: providers occasionally emit malformed JSON; surface that as + an empty dict so the loop can still call the sandbox (which will then + return an error the LLM can see). + """ + if isinstance(call, dict): + fn = call.get("function") or {} + raw = ( + fn.get("arguments") + if isinstance(fn, dict) + else getattr(fn, "arguments", None) + ) + else: + fn = getattr(call, "function", None) + raw = getattr(fn, "arguments", None) if fn is not None else None + + if raw is None: + return {} + if isinstance(raw, dict): + return raw + if isinstance(raw, str): + try: + parsed = json.loads(raw) + except json.JSONDecodeError: + return {} + if isinstance(parsed, dict): + return parsed + return {} + + +def _stringify(output: Any) -> str: + if isinstance(output, str): + return output + if output is None: + return "" + try: + return json.dumps(output) + except (TypeError, ValueError): + return str(output) diff --git a/litellm/managed_agents/events.py b/litellm/managed_agents/events.py new file mode 100644 index 00000000000..1ed3a74609c --- /dev/null +++ b/litellm/managed_agents/events.py @@ -0,0 +1,64 @@ +""" +Event types emitted by an ``AgentRuntime`` while driving an LLM tool loop. + +Events are the canonical wire shape between a runtime and the calling +``Session`` — they get persisted to ``LiteLLM_AgentRunEvent`` (one row per +event, monotonically increasing ``seq``) and streamed back to clients via +the existing /v2/sessions/{id}/runs/{rid}/events SSE endpoint. + +Wire shape (snake_case, matches what the integration branch already serves): + + {"seq": 1, "event_type": "assistant_message", "payload": {"content": "..."}} + +The Python ``Event`` dataclass below is a small in-memory wrapper. The +runtime yields ``Event`` instances; the ``Session`` is responsible for +turning them into rows. +""" + +from dataclasses import dataclass, field +from typing import Any, Dict + + +# --------------------------------------------------------------------------- +# Event type constants — keep in sync with +# litellm/proxy/agent_session_endpoints/constants.py for the lifecycle ones +# (run_started, run_finished, run_cancelled, run_error, user_message). The +# runtime-specific event types below are new and only flow through Epic C. +# --------------------------------------------------------------------------- + +# Lifecycle (also defined in agent_session_endpoints.constants) +EVENT_TYPE_RUN_STARTED = "run_started" +EVENT_TYPE_RUN_FINISHED = "run_finished" +EVENT_TYPE_RUN_CANCELLED = "run_cancelled" +EVENT_TYPE_RUN_ERROR = "run_error" +EVENT_TYPE_USER_MESSAGE = "user_message" + +# Runtime-emitted (the actual LLM tool-loop events) +EVENT_TYPE_ASSISTANT_MESSAGE = "assistant_message" +EVENT_TYPE_TOOL_USE = "tool_use" +EVENT_TYPE_TOOL_RESULT = "tool_result" +EVENT_TYPE_THINKING = "thinking" +EVENT_TYPE_SYSTEM = "system" + + +@dataclass +class Event: + """A single event from a runtime. + + ``type`` is the canonical event type string (snake_case). ``data`` is a + free-form dict that gets persisted as the JSON ``payload`` column on + ``LiteLLM_AgentRunEvent``. + + Example:: + + Event(type="tool_use", data={"tool": "Bash", "input": {"command": "ls"}}) + """ + + type: str + data: Dict[str, Any] = field(default_factory=dict) + + def to_payload(self) -> Dict[str, Any]: + """Return the wire payload — currently identity, but reserved for + future field renames so call sites have one place to change. + """ + return dict(self.data) diff --git a/litellm/managed_agents/run.py b/litellm/managed_agents/run.py new file mode 100644 index 00000000000..e58fcf2ca7c --- /dev/null +++ b/litellm/managed_agents/run.py @@ -0,0 +1,130 @@ +""" +Run — read-only Python accessor for one ``LiteLLM_AgentRun`` row. + +This is the SDK-side counterpart to the ``GET /v2/sessions/{sid}/runs/{rid}`` +HTTP endpoint. The proxy persists every run to ``LiteLLM_AgentRun`` and every +event to ``LiteLLM_AgentRunEvent``; ``Run`` is the typed accessor in-process +code uses to read those rows back. + +Lifecycle notes: + * Construction does not touch the DB. Use ``Run.from_db_row(...)`` to + build an instance from a Prisma row. + * ``stream(starting_seq=N)`` is an async iterator over events ordered by + ``seq``; pass ``N>0`` to resume from where you stopped. Iteration is + snapshot-based — it yields events that are persisted at the time of + the call and stops. Live tailing (waiting for new events as the run + progresses) is the SSE endpoint's job, not this class. + * Read-only by design. ``Run`` never writes to the DB. The owning + ``Session`` writes; ``Run`` only reads. +""" + +from dataclasses import dataclass, field +from typing import Any, AsyncIterator, Dict, Optional + +from litellm.managed_agents.events import Event + + +# Default page size when pulling events from the DB. Picked to keep memory +# footprint modest on chatty runs (a busy run can have thousands of events). +_DEFAULT_EVENT_PAGE_SIZE = 200 + + +@dataclass +class Run: + """Read-only view of one ``LiteLLM_AgentRun`` row. + + The ``db`` reference is the Prisma client — kept here so ``stream()`` + can fetch events without callers having to wire the client through. + Tests and offline tools that don't need ``stream()`` can pass + ``db=None`` and just consume the static fields. + """ + + id: str + session_id: str + status: str + prompt: Dict[str, Any] = field(default_factory=dict) + result: Optional[str] = None + parent_run_id: Optional[str] = None + db: Any = None + + @classmethod + async def from_db_row(cls, row: Any, db: Any) -> "Run": + """Build a ``Run`` from a Prisma ``LiteLLM_AgentRun`` row. + + Defensive on field access so partial selects (e.g. via Prisma's + ``select=...`` parameter) don't blow up — missing fields just + default. The caller is responsible for fetching all the columns + it needs for whatever it plans to do with the ``Run``. + """ + prompt = _coerce_dict(getattr(row, "prompt", None)) + return cls( + id=getattr(row, "id"), + session_id=getattr(row, "session_id"), + status=getattr(row, "status", "unknown"), + prompt=prompt, + result=getattr(row, "result", None), + parent_run_id=getattr(row, "parent_run_id", None), + db=db, + ) + + async def stream( + self, + starting_seq: int = 0, + ) -> AsyncIterator[Event]: + """Yield persisted events in ``seq`` order, starting from ``starting_seq``. + + ``starting_seq`` is exclusive of itself when passed (matches the + SSE endpoint contract: client passes the last seq it saw, server + replays everything strictly greater). Pass ``0`` for "from the + beginning". + + Iterates in batches to keep peak memory bounded — ``LiteLLM_AgentRunEvent`` + is wide (JSON payload) and a busy run can have thousands of rows. + Stops as soon as the DB returns fewer rows than the page size, + which is the canonical "no more pages" signal in cursor-style + pagination. + """ + if self.db is None: + raise RuntimeError( + "Run.stream() requires a Prisma client; construct with db= " + "or via Run.from_db_row(row, db=)." + ) + + cursor_seq = int(starting_seq) + while True: + rows = await self.db.litellm_agentrunevent.find_many( + where={ + "run_id": self.id, + "seq": {"gt": cursor_seq}, + }, + order={"seq": "asc"}, + take=_DEFAULT_EVENT_PAGE_SIZE, + ) + if not rows: + return + + for row in rows: + yield _event_from_row(row) + cursor_seq = max(cursor_seq, int(getattr(row, "seq", cursor_seq))) + + # Short-circuit when the page wasn't full — no more rows exist. + if len(rows) < _DEFAULT_EVENT_PAGE_SIZE: + return + + +def _coerce_dict(value: Any) -> Dict[str, Any]: + """Prisma JSON columns come back as dicts; pass strings/None through.""" + if isinstance(value, dict): + return value + return {} + + +def _event_from_row(row: Any) -> Event: + """Translate a ``LiteLLM_AgentRunEvent`` row into our ``Event`` dataclass.""" + payload = getattr(row, "payload", None) + if not isinstance(payload, dict): + payload = {} + return Event( + type=getattr(row, "event_type", "unknown"), + data=dict(payload), + ) diff --git a/litellm/managed_agents/sandbox/__init__.py b/litellm/managed_agents/sandbox/__init__.py new file mode 100644 index 00000000000..b35c11b6b08 --- /dev/null +++ b/litellm/managed_agents/sandbox/__init__.py @@ -0,0 +1,12 @@ +"""Pluggable sandboxes: where tool calls actually execute.""" + +from litellm.managed_agents.sandbox.base import Sandbox, ToolResult +from litellm.managed_agents.sandbox.ec2_ssm import EC2SandboxViaSSM +from litellm.managed_agents.sandbox.local import LocalSandbox + +__all__ = [ + "Sandbox", + "ToolResult", + "LocalSandbox", + "EC2SandboxViaSSM", +] diff --git a/litellm/managed_agents/sandbox/base.py b/litellm/managed_agents/sandbox/base.py new file mode 100644 index 00000000000..81939adf5b0 --- /dev/null +++ b/litellm/managed_agents/sandbox/base.py @@ -0,0 +1,94 @@ +""" +Sandbox abstraction — the venue where tool calls actually execute. + +The runtime separates *deciding what tool to call* (LLM tool loop) from +*actually running the tool* (filesystem, shell, network). This lets the +same runtime drive a tool loop against: + + * ``LocalSandbox`` — execute in the proxy process (dev only) + * ``EC2SandboxViaSSM`` — execute on a remote VM via SSM RunCommand + * ``DockerSandbox`` — execute in a container (future) + +Each implementation only needs to honor ``execute_tool(name, input)``. +""" + +from abc import ABC, abstractmethod +from dataclasses import dataclass, field +from typing import Any, Dict, Optional + + +@dataclass +class ToolResult: + """Result of executing a tool inside a sandbox. + + ``output`` is whatever the tool returned (string, dict, bytes — caller + decides). ``is_error`` is True when the tool failed; the runtime maps + this onto the ``tool_result`` event ``is_error`` field which the LLM + then sees on its next turn. ``metadata`` is open-ended for sandbox + implementations that want to surface execution-venue details (exit + code, vm_id, region, duration_ms, etc.). + """ + + output: Any + is_error: bool = False + metadata: Dict[str, Any] = field(default_factory=dict) + + +class Sandbox(ABC): + """Where tool calls execute. + + Subclasses implement ``execute_tool(tool_name, tool_input) -> ToolResult``. + + The runtime calls into this for any tool the LLM requests. For example, + when the LLM emits a ``tool_use`` block ``{"name": "Bash", "input": + {"command": "ls"}}``, the runtime invokes + ``await sandbox.execute_tool("Bash", {"command": "ls"})`` and feeds the + result back into the next LLM turn as a ``tool_result``. + + Implementations should be safe to share across concurrent runs only if + they document so explicitly. The default contract is "one Sandbox per + Session" — each Session owns its sandbox for the duration of its life. + """ + + @abstractmethod + async def execute_tool( + self, + tool_name: str, + tool_input: Dict[str, Any], + ) -> ToolResult: + """Execute a tool and return its result. + + Implementations must NOT raise on tool-level errors (e.g. the + command exited non-zero, the file did not exist). Instead, return + ``ToolResult(output=, is_error=True)`` so the LLM can + see the failure and react. Raise only on infrastructure failures + (sandbox unreachable, OOM) so the runtime can decide whether to + abort the run or retry. + """ + + async def setup(self) -> None: + """Optional: prepare the sandbox before the first tool call. + + Override for sandboxes that need to provision a VM, clone repos, + install deps, etc. Default is no-op so simple sandboxes + (``LocalSandbox``) can ignore it. + """ + return None + + async def teardown(self) -> None: + """Optional: clean up after the last tool call. + + Override for sandboxes that need to release a VM, delete a + container, etc. Default is no-op. + """ + return None + + @property + def cwd(self) -> Optional[str]: + """Optional working directory hint for runtimes that need one + (e.g. claude-agent-sdk's ``ClaudeAgentOptions.cwd``). + + Returning ``None`` means "no hint" — the runtime falls back to + whatever default it normally uses. + """ + return None diff --git a/litellm/managed_agents/sandbox/ec2_ssm.py b/litellm/managed_agents/sandbox/ec2_ssm.py new file mode 100644 index 00000000000..7d8e74a83e0 --- /dev/null +++ b/litellm/managed_agents/sandbox/ec2_ssm.py @@ -0,0 +1,57 @@ +""" +EC2SandboxViaSSM — placeholder. + +Production sandbox: the proxy issues an SSM ``RunCommand`` to a VM +provisioned via Epic B's vm_providers (default: EC2 in customer AWS) and +shells out the tool call there. The VM no longer needs a custom binary — +stock Ubuntu + the SSM agent is enough. + +This placeholder exists so: + + * The public API (``from litellm.managed_agents.sandbox import + EC2SandboxViaSSM``) is stable from day one. + * Validation criterion #1 (clean import) passes. + * Future PR can swap in the real implementation without touching call + sites. + +The real implementation belongs in a follow-up that wires +``litellm/proxy/agent_session_endpoints/vm_providers/`` (already on the +integration branch) into ``execute_tool``. Tracked as a deferred item on +LIT-2879. +""" + +from typing import Any, Dict, Optional + +from litellm.managed_agents.sandbox.base import Sandbox, ToolResult + + +class EC2SandboxViaSSM(Sandbox): + """Placeholder for SSM-backed remote execution. Not yet implemented. + + Construct it freely (the public API is stable). Calling + ``execute_tool`` raises ``NotImplementedError`` until the SSM wiring + lands. The constructor takes ``team_id`` because the real implementation + will look up team-scoped AWS creds + the VM provisioning provider via + ``vm_providers.get_vm_provider("ec2")``. + """ + + def __init__( + self, + team_id: Optional[str] = None, + vm_id: Optional[str] = None, + region: Optional[str] = None, + ) -> None: + self.team_id = team_id + self.vm_id = vm_id + self.region = region + + async def execute_tool( + self, + tool_name: str, + tool_input: Dict[str, Any], + ) -> ToolResult: + raise NotImplementedError( + "EC2SandboxViaSSM is a placeholder. Wire it to " + "litellm/proxy/agent_session_endpoints/vm_providers/ec2 in a " + "follow-up PR." + ) diff --git a/litellm/managed_agents/sandbox/local.py b/litellm/managed_agents/sandbox/local.py new file mode 100644 index 00000000000..5cfe46442a4 --- /dev/null +++ b/litellm/managed_agents/sandbox/local.py @@ -0,0 +1,199 @@ +""" +LocalSandbox — executes tool calls in the proxy process. + +For dev mode only. Multi-tenant unsafe (no isolation between sessions). +Production deploys should use ``EC2SandboxViaSSM`` instead. + +The set of tools recognised here mirrors the names ``claude-agent-sdk`` +uses for its built-ins (``Bash``, ``Read``, ``Write``, ``Edit``) so the +``LiteLLMAgentRuntime`` and any custom runtime can reuse them. + +For ``ClaudeSDKAgentRuntime``, ``LocalSandbox.execute_tool`` is intentionally +unused: claude-agent-sdk runs its built-in tools in-process directly, and the +sandbox-routing override is deferred to a follow-up PR. The ``cwd`` property +is still consumed though so the SDK runs in the right directory. +""" + +import asyncio +import os +import shutil +from pathlib import Path +from typing import Any, Dict, Optional + +from litellm.managed_agents.sandbox.base import Sandbox, ToolResult + + +class LocalSandbox(Sandbox): + """Executes tools in the proxy process. + + ``working_dir`` defaults to a fresh temp dir per-instance. Pass an + explicit dir to share state across runs (e.g. a cloned repo). + + ``shell_timeout_seconds`` bounds individual ``Bash`` calls — the LLM + can request long-running commands and we don't want a runaway process + holding up a whole run. + """ + + def __init__( + self, + working_dir: Optional[str] = None, + shell_timeout_seconds: float = 60.0, + ) -> None: + self._working_dir: Optional[str] = working_dir + self._owned_tmpdir: Optional[str] = None + self._shell_timeout = shell_timeout_seconds + + @property + def cwd(self) -> Optional[str]: + return self._working_dir + + async def setup(self) -> None: + if self._working_dir is None: + import tempfile + + self._owned_tmpdir = tempfile.mkdtemp(prefix="litellm_managed_agent_") + self._working_dir = self._owned_tmpdir + + async def teardown(self) -> None: + if self._owned_tmpdir is not None and os.path.isdir(self._owned_tmpdir): + shutil.rmtree(self._owned_tmpdir, ignore_errors=True) + self._owned_tmpdir = None + self._working_dir = None + + async def execute_tool( + self, + tool_name: str, + tool_input: Dict[str, Any], + ) -> ToolResult: + if self._working_dir is None: + await self.setup() + # mypy: setup populates _working_dir + cwd = self._working_dir or "." + + name = tool_name.lower() + if name in {"bash", "shell", "exec"}: + return await self._run_bash(tool_input, cwd) + if name in {"read", "read_file"}: + return self._run_read(tool_input, cwd) + if name in {"write", "write_file"}: + return self._run_write(tool_input, cwd) + if name in {"edit", "edit_file"}: + return self._run_edit(tool_input, cwd) + if name == "ls": + return self._run_ls(tool_input, cwd) + return ToolResult( + output=f"unknown tool: {tool_name}", + is_error=True, + metadata={"sandbox": "local"}, + ) + + # ------------------------------------------------------------------ + # Tool implementations — kept tiny on purpose. These exist so + # LiteLLMAgentRuntime + LocalSandbox can do useful work end-to-end + # without depending on claude-agent-sdk's built-in tool surface. + # ------------------------------------------------------------------ + + async def _run_bash(self, tool_input: Dict[str, Any], cwd: str) -> ToolResult: + cmd = tool_input.get("command") + if not isinstance(cmd, str) or not cmd: + return ToolResult(output="missing or empty 'command'", is_error=True) + try: + proc = await asyncio.create_subprocess_shell( + cmd, + cwd=cwd, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + ) + try: + stdout, stderr = await asyncio.wait_for( + proc.communicate(), timeout=self._shell_timeout + ) + except asyncio.TimeoutError: + proc.kill() + await proc.wait() + return ToolResult( + output=f"command timed out after {self._shell_timeout}s", + is_error=True, + metadata={"timeout": self._shell_timeout}, + ) + output = (stdout or b"").decode(errors="replace") + err = (stderr or b"").decode(errors="replace") + if proc.returncode != 0: + return ToolResult( + output=f"{output}\n{err}".strip(), + is_error=True, + metadata={"exit_code": proc.returncode}, + ) + return ToolResult( + output=output, + metadata={"exit_code": proc.returncode}, + ) + except Exception as exc: + return ToolResult(output=f"bash failed: {exc}", is_error=True) + + def _resolve_path(self, tool_input: Dict[str, Any], cwd: str) -> Optional[Path]: + path = tool_input.get("path") or tool_input.get("file_path") + if not isinstance(path, str) or not path: + return None + p = Path(path) + if not p.is_absolute(): + p = Path(cwd) / p + return p + + def _run_read(self, tool_input: Dict[str, Any], cwd: str) -> ToolResult: + p = self._resolve_path(tool_input, cwd) + if p is None: + return ToolResult(output="missing 'path'", is_error=True) + try: + return ToolResult(output=p.read_text()) + except FileNotFoundError: + return ToolResult(output=f"no such file: {p}", is_error=True) + except Exception as exc: + return ToolResult(output=f"read failed: {exc}", is_error=True) + + def _run_write(self, tool_input: Dict[str, Any], cwd: str) -> ToolResult: + p = self._resolve_path(tool_input, cwd) + if p is None: + return ToolResult(output="missing 'path'", is_error=True) + content = tool_input.get("content", "") + if not isinstance(content, str): + return ToolResult(output="'content' must be a string", is_error=True) + try: + p.parent.mkdir(parents=True, exist_ok=True) + p.write_text(content) + return ToolResult(output=f"wrote {len(content)} bytes to {p}") + except Exception as exc: + return ToolResult(output=f"write failed: {exc}", is_error=True) + + def _run_edit(self, tool_input: Dict[str, Any], cwd: str) -> ToolResult: + p = self._resolve_path(tool_input, cwd) + if p is None: + return ToolResult(output="missing 'path'", is_error=True) + old = tool_input.get("old_string") + new = tool_input.get("new_string", "") + if not isinstance(old, str) or not isinstance(new, str): + return ToolResult( + output="'old_string' and 'new_string' must be strings", is_error=True + ) + try: + text = p.read_text() + if old not in text: + return ToolResult( + output="'old_string' not found in file", is_error=True + ) + p.write_text(text.replace(old, new, 1)) + return ToolResult(output=f"edited {p}") + except FileNotFoundError: + return ToolResult(output=f"no such file: {p}", is_error=True) + except Exception as exc: + return ToolResult(output=f"edit failed: {exc}", is_error=True) + + def _run_ls(self, tool_input: Dict[str, Any], cwd: str) -> ToolResult: + path = tool_input.get("path") or cwd + try: + entries = sorted(os.listdir(path)) + return ToolResult(output="\n".join(entries)) + except FileNotFoundError: + return ToolResult(output=f"no such dir: {path}", is_error=True) + except Exception as exc: + return ToolResult(output=f"ls failed: {exc}", is_error=True) diff --git a/litellm/managed_agents/session.py b/litellm/managed_agents/session.py new file mode 100644 index 00000000000..46920bee625 --- /dev/null +++ b/litellm/managed_agents/session.py @@ -0,0 +1,323 @@ +""" +Session — Python-side handle for one managed-agent session. + +A ``Session`` is the unit of work in the managed-agents API. It holds the +sandbox + runtime that any prompts sent through it will use, and it owns +the persistence of runs and events to the DB. + +Lifecycle: + * Construction is via ``Session.from_db_row(row, db, runtime, sandbox, + agent_config)``. The owning ``Agent`` calls this internally so SDK + users normally don't. + * ``await session.send(prompt)`` is the primary entry point: it + INSERTs a ``LiteLLM_AgentRun`` row, drives the runtime, persists + every event to ``LiteLLM_AgentRunEvent``, and yields each event back + to the caller as the runtime emits it. The session and run statuses + flip in lock-step (``ready`` -> ``busy`` -> ``ready`` for the + session; ``queued`` -> ``running`` -> ``finished``/``error`` for the + run) so polling clients see consistent state. + * Other helpers (``get_run``, ``list_runs``, ``conversation``) wrap + read-only queries — same shape the HTTP endpoints serve, no + surprises. + +Concurrency contract: one in-flight ``send()`` per session. The session +status flag enforces this in the DB (busy -> reject), but Python callers +should not call ``send()`` twice concurrently on the same instance — the +runtime and sandbox are not assumed to be concurrent-safe. +""" + +from dataclasses import dataclass, field +from datetime import datetime, timezone +from typing import Any, AsyncIterator, Dict, List, Optional + +import prisma + +from litellm.managed_agents.agent_runtime.base import ( + AgentConfig, + AgentRuntime, + SessionState, +) +from litellm.managed_agents.events import Event +from litellm.managed_agents.run import Run +from litellm.managed_agents.sandbox.base import Sandbox +from litellm.proxy.agent_session_endpoints.constants import ( + RUN_STATUS_ERROR, + RUN_STATUS_FINISHED, + RUN_STATUS_QUEUED, + RUN_STATUS_RUNNING, + SESSION_STATUS_BUSY, + SESSION_STATUS_READY, +) +from litellm.proxy.agent_session_endpoints.ids import new_run_id + + +def _now() -> datetime: + return datetime.now(timezone.utc) + + +@dataclass +class Session: + """A live, in-process handle on one ``LiteLLM_AgentSession`` row. + + Holds runtime + sandbox + agent_config (a snapshot of the parent + agent's runtime-relevant fields). The ``daemon_token`` field is set + only at create-time and is the same JWT the proxy stores; it's + surfaced here so callers that still want to talk to the proxy's + HTTP endpoints (e.g. SSE event tail) have what they need. + """ + + id: str + agent_id: str + status: str + runtime: AgentRuntime + sandbox: Sandbox + agent_config: AgentConfig + db: Any + repos: List[Dict[str, Any]] = field(default_factory=list) + env_vars: Dict[str, str] = field(default_factory=dict) + daemon_token: Optional[str] = None + + # ------------------------------------------------------------------ + # Construction + # ------------------------------------------------------------------ + + @classmethod + async def from_db_row( + cls, + row: Any, + db: Any, + runtime: Optional[AgentRuntime], + sandbox: Optional[Sandbox], + agent_config: AgentConfig, + daemon_token: Optional[str] = None, + ) -> "Session": + if runtime is None or sandbox is None: + raise RuntimeError( + "Session requires both runtime and sandbox; the parent Agent " + "must be constructed with both before spawning sessions." + ) + return cls( + id=getattr(row, "id"), + agent_id=getattr(row, "agent_id"), + status=getattr(row, "status", "unknown"), + runtime=runtime, + sandbox=sandbox, + agent_config=agent_config, + db=db, + repos=_coerce_list_of_dict(getattr(row, "repos", None)), + env_vars=_coerce_dict(getattr(row, "env_vars", None)), + daemon_token=daemon_token, + ) + + # ------------------------------------------------------------------ + # The interesting one — send a prompt and yield events as they happen + # ------------------------------------------------------------------ + + async def send(self, prompt: str) -> AsyncIterator[Event]: + """Drive one run end-to-end and yield events as the runtime emits them. + + Steps (mirrors what the HTTP create_run + event_stream pair does + on the proxy): + 1. INSERT ``LiteLLM_AgentRun`` (status=``queued``, prompt={text}). + 2. UPDATE ``self.status = busy`` and the run row to ``running``. + 3. ``async for event in self.runtime.run(...)``: + - INSERT ``LiteLLM_AgentRunEvent(run_id, seq, event_type, payload)``. + - yield the event. + 4. On clean exit: UPDATE run -> ``finished``, session -> ``ready``. + 5. On exception: UPDATE run -> ``error``, session -> ``ready``, + then re-raise. We always restore the session to ``ready`` + so a single failed send doesn't lock the session forever. + + The ``finally`` block restores session status even on early + cancellation by the consumer (their generator close raises + ``GeneratorExit``). That's important because async generators + can be closed without ever reaching their tail. + """ + run_id = await self._insert_run(prompt) + + await self._set_session_status(SESSION_STATUS_BUSY) + await self._set_run_status(run_id, RUN_STATUS_RUNNING, started_at=_now()) + + seq = 0 + terminal_status = RUN_STATUS_FINISHED + terminal_result: Optional[str] = None + + session_state = SessionState( + session_id=self.id, + cwd=self.sandbox.cwd, + env_vars=dict(self.env_vars), + repos=list(self.repos), + ) + + try: + async for event in self.runtime.run( + prompt=prompt, + sandbox=self.sandbox, + session_state=session_state, + agent_config=self.agent_config, + ): + seq += 1 + await self._insert_event(run_id, seq, event) + + # Track the LLM's last "result" so we can store it on the + # run row — matches what the HTTP endpoint does when the + # daemon reports a run_finished event. + if event.type == "run_finished": + result = event.data.get("result") + if isinstance(result, str): + terminal_result = result + + yield event + except Exception: + terminal_status = RUN_STATUS_ERROR + raise + finally: + # Always close out the run + restore the session, even if the + # consumer cancelled early. The runtime is responsible for + # any sandbox-side cleanup it owns. + await self._set_run_status( + run_id, + terminal_status, + terminated_at=_now(), + result=terminal_result, + ) + await self._set_session_status(SESSION_STATUS_READY) + + # ------------------------------------------------------------------ + # Read-only queries + # ------------------------------------------------------------------ + + async def get_run(self, run_id: str) -> Run: + """Fetch a single run by id; raises ``LookupError`` if not found / not ours.""" + row = await self.db.litellm_agentrun.find_unique(where={"id": run_id}) + if row is None or getattr(row, "session_id", None) != self.id: + raise LookupError(f"Run {run_id!r} not found for session {self.id!r}") + return await Run.from_db_row(row, db=self.db) + + async def list_runs(self) -> List[Run]: + """Return all runs in this session, newest first.""" + rows = await self.db.litellm_agentrun.find_many( + where={"session_id": self.id}, + order={"created_at": "desc"}, + ) + return [await Run.from_db_row(row, db=self.db) for row in rows] + + async def conversation(self) -> List[Dict[str, Any]]: + """Return every event across every run in this session, in order. + + Wire shape matches ``GET /v2/sessions/{sid}/conversation``: each + item is ``{run_id, seq, event_type, payload, created_at}``. The + order key is ``(created_at ASC, seq ASC)`` so events from + concurrent runs interleave by wall-clock — the HTTP serializer + sorts the same way. + """ + runs = await self.db.litellm_agentrun.find_many( + where={"session_id": self.id}, + ) + run_ids = [r.id for r in runs] + if not run_ids: + return [] + + events = await self.db.litellm_agentrunevent.find_many( + where={"run_id": {"in": run_ids}}, + order={"created_at": "asc"}, + ) + + out: List[Dict[str, Any]] = [] + for ev in events: + payload = getattr(ev, "payload", None) + out.append( + { + "run_id": getattr(ev, "run_id"), + "seq": getattr(ev, "seq"), + "event_type": getattr(ev, "event_type"), + "payload": payload if isinstance(payload, dict) else {}, + "created_at": _iso(getattr(ev, "created_at", None)), + } + ) + return out + + # ------------------------------------------------------------------ + # DB write helpers — kept small + named so the send() flow reads + # like a state machine instead of a wall of prisma calls. + # ------------------------------------------------------------------ + + async def _insert_run(self, prompt: str) -> str: + """Create the LiteLLM_AgentRun row in queued state. Returns the run id. + + We wrap the prompt in ``{text: prompt}`` so the JSON column has + a stable shape — the HTTP RunCreate schema does the same and + downstream serializers expect a dict. + """ + run_id = new_run_id() + payload: Dict[str, Any] = { + "id": run_id, + "session": {"connect": {"id": self.id}}, + "status": RUN_STATUS_QUEUED, + "prompt": prisma.Json({"text": prompt}), + "updated_at": _now(), + } + await self.db.litellm_agentrun.create(data=payload) + return run_id + + async def _set_run_status( + self, + run_id: str, + status: str, + started_at: Optional[datetime] = None, + terminated_at: Optional[datetime] = None, + result: Optional[str] = None, + ) -> None: + data: Dict[str, Any] = {"status": status, "updated_at": _now()} + if started_at is not None: + data["started_at"] = started_at + if terminated_at is not None: + data["terminated_at"] = terminated_at + if result is not None: + data["result"] = result + await self.db.litellm_agentrun.update(where={"id": run_id}, data=data) + + async def _set_session_status(self, status: str) -> None: + """Flip the session row's status. Also keeps a local mirror in sync.""" + await self.db.litellm_agentsession.update( + where={"id": self.id}, + data={"status": status, "updated_at": _now()}, + ) + self.status = status + + async def _insert_event(self, run_id: str, seq: int, event: Event) -> None: + await self.db.litellm_agentrunevent.create( + data={ + "run_id": run_id, + "seq": seq, + "event_type": event.type, + "payload": prisma.Json(event.to_payload()), + } + ) + + +# --------------------------------------------------------------------------- +# Coercion helpers — same minimal shape as in ``agent.py``. Inlined here +# instead of pulling from a shared module to keep ``Session`` and ``Agent`` +# independently importable. +# --------------------------------------------------------------------------- + + +def _coerce_dict(value: Any) -> Dict[str, Any]: + if isinstance(value, dict): + return value + return {} + + +def _coerce_list_of_dict(value: Any) -> List[Dict[str, Any]]: + if isinstance(value, list): + return [v for v in value if isinstance(v, dict)] + return [] + + +def _iso(value: Any) -> Optional[str]: + if value is None: + return None + if isinstance(value, datetime): + return value.isoformat() + return str(value) diff --git a/tests/test_litellm/managed_agents/__init__.py b/tests/test_litellm/managed_agents/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/managed_agents/conftest.py b/tests/test_litellm/managed_agents/conftest.py new file mode 100644 index 00000000000..233bfe6db3e --- /dev/null +++ b/tests/test_litellm/managed_agents/conftest.py @@ -0,0 +1,54 @@ +""" +Shared fixtures for ``litellm.managed_agents`` tests. + +Re-uses the in-memory Prisma stand-in from +``tests/test_litellm/proxy/agent_session_endpoints/conftest.py`` so the +managed_agents Python API and the proxy HTTP endpoints exercise the +same fake DB behaviour. Tests don't need to reach into proxy_server's +global state because the managed_agents API talks to the Prisma client +directly via ``self.db``. +""" + +import os +import sys +from pathlib import Path + +import pytest + +# Set a JWT secret BEFORE any module under test (Agent / Session import the +# proxy auth helpers which read this env var at call time). +os.environ.setdefault("LITELLM_AGENT_JWT_SECRET", "test-agent-jwt-secret") +os.environ.setdefault("LITELLM_MASTER_KEY", "sk-1234") + +# Provide ``prisma.Json`` as an identity function for tests. The real symbol +# is generated by ``prisma generate`` from the schema and isn't available +# unless the full prisma codegen has run. The fake DB stores whatever payload +# we hand it, so an identity wrapper is sufficient — the call sites only +# need ``prisma.Json(value)`` to evaluate to ``value`` in tests. +import prisma # noqa: E402 + +if not hasattr(prisma, "Json"): + prisma.Json = lambda value: value # type: ignore[attr-defined] + +# Make the agent_session_endpoints test conftest importable from here. We +# reach across test packages instead of duplicating the FakePrismaClient +# so any future evolution of the fake DB is shared with proxy tests. +_ENDPOINTS_TEST_DIR = ( + Path(__file__).resolve().parent.parent / "proxy" / "agent_session_endpoints" +) +if str(_ENDPOINTS_TEST_DIR) not in sys.path: + sys.path.insert(0, str(_ENDPOINTS_TEST_DIR)) + + +@pytest.fixture +def fake_db(): + """Bare in-memory Prisma client; use this directly via ``client.db``. + + Most managed_agents tests don't need monkeypatching of + ``proxy_server.prisma_client`` because Agent/Session take ``db`` as + a parameter — pass ``fake_db.db`` straight through. + """ + # Local import: see _ENDPOINTS_TEST_DIR sys.path manipulation above. + from conftest import FakePrismaClient # type: ignore # noqa: E402 + + return FakePrismaClient() diff --git a/tests/test_litellm/managed_agents/test_agent.py b/tests/test_litellm/managed_agents/test_agent.py new file mode 100644 index 00000000000..fc68c3d5c32 --- /dev/null +++ b/tests/test_litellm/managed_agents/test_agent.py @@ -0,0 +1,230 @@ +""" +Unit tests for ``litellm.managed_agents.Agent``. + +Covers the lifecycle helpers (``from_db_row``, ``create_session``, +``get_session``, ``list_sessions``, ``delete``) end-to-end against the +in-memory Prisma stand-in. The runtime + sandbox here are minimal stubs +since these tests focus on the DB orchestration, not LLM behaviour. +""" + +import asyncio +from typing import Any, AsyncIterator, Dict + +import pytest + +from litellm.managed_agents.agent import Agent +from litellm.managed_agents.agent_runtime.base import ( + AgentConfig, + AgentRuntime, + SessionState, +) +from litellm.managed_agents.events import EVENT_TYPE_RUN_FINISHED, Event +from litellm.managed_agents.sandbox.local import LocalSandbox + + +class _StubRuntime(AgentRuntime): + """Minimal AgentRuntime that immediately yields run_finished.""" + + async def run( + self, + prompt: str, + sandbox, + session_state: SessionState, + agent_config: AgentConfig, + ) -> AsyncIterator[Event]: + yield Event( + type=EVENT_TYPE_RUN_FINISHED, + data={"result": f"echo: {prompt}"}, + ) + + +def _seed_agent_row(db, *, agent_id: str = "agent_1", **overrides: Any): + """Insert a LiteLLM_Agent row directly into the fake DB.""" + return asyncio.get_event_loop().run_until_complete( + db.litellm_agent.create( + data={ + "id": agent_id, + "name": overrides.get("name", "test-agent"), + "model": overrides.get("model", "gpt-4o-mini"), + "user_api_key_hash": overrides.get("user_api_key_hash", "hash-A"), + "team_id": overrides.get("team_id", None), + "system_prompt": overrides.get("system_prompt", "you are helpful"), + "tools_config": overrides.get("tools_config", None), + "metadata": overrides.get("metadata", {}), + "default_repos": overrides.get("default_repos", []), + "default_env_vars": overrides.get("default_env_vars", {}), + } + ) + ) + + +@pytest.mark.asyncio +async def test_from_db_row_populates_fields(fake_db): + row = await fake_db.db.litellm_agent.create( + data={ + "id": "agent_x", + "name": "alpha", + "model": "claude-sonnet-4", + "system_prompt": "hi", + "user_api_key_hash": "h", + "tools_config": {"tools": []}, + "metadata": {"k": "v"}, + "default_repos": [{"url": "a", "ref": "main"}], + "default_env_vars": {"NPM": "1"}, + "team_id": "team-1", + } + ) + agent = await Agent.from_db_row( + row, db=fake_db.db, runtime=_StubRuntime(), sandbox=LocalSandbox() + ) + assert agent.id == "agent_x" + assert agent.name == "alpha" + assert agent.model == "claude-sonnet-4" + assert agent.system_prompt == "hi" + assert agent.tools_config == {"tools": []} + assert agent.metadata == {"k": "v"} + assert agent.default_repos == [{"url": "a", "ref": "main"}] + assert agent.default_env_vars == {"NPM": "1"} + assert agent.team_id == "team-1" + assert agent.user_api_key_hash == "h" + + +@pytest.mark.asyncio +async def test_create_session_inserts_row_in_ready(fake_db): + row = await fake_db.db.litellm_agent.create( + data={ + "id": "agent_1", + "name": "test", + "model": "gpt-4o-mini", + "user_api_key_hash": "h", + "default_repos": [{"url": "default-repo"}], + "default_env_vars": {"DEFAULT": "x"}, + } + ) + agent = await Agent.from_db_row( + row, db=fake_db.db, runtime=_StubRuntime(), sandbox=LocalSandbox() + ) + + session = await agent.create_session(env_vars={"OVERRIDE": "y"}) + assert session.id.startswith("sess_") + assert session.agent_id == "agent_1" + assert session.status == "ready" + # Defaults flow through; overrides win on collisions. + assert session.env_vars == {"DEFAULT": "x", "OVERRIDE": "y"} + # repos default carried through (no override). + assert session.repos == [{"url": "default-repo"}] + # daemon token returned exactly once at create. + assert isinstance(session.daemon_token, str) and session.daemon_token + + # Same row should be persisted in the fake DB. + persisted = await fake_db.db.litellm_agentsession.find_unique( + where={"id": session.id} + ) + assert persisted is not None + assert persisted.status == "ready" + + +@pytest.mark.asyncio +async def test_create_session_with_repos_replaces_defaults(fake_db): + row = await fake_db.db.litellm_agent.create( + data={ + "id": "agent_1", + "name": "x", + "model": "m", + "user_api_key_hash": "h", + "default_repos": [{"url": "default"}], + } + ) + agent = await Agent.from_db_row( + row, db=fake_db.db, runtime=_StubRuntime(), sandbox=LocalSandbox() + ) + session = await agent.create_session(repos=[{"url": "override"}]) + assert session.repos == [{"url": "override"}] + + +@pytest.mark.asyncio +async def test_get_session_returns_existing_row(fake_db): + row = await fake_db.db.litellm_agent.create( + data={"id": "agent_1", "name": "x", "model": "m", "user_api_key_hash": "h"} + ) + agent = await Agent.from_db_row( + row, db=fake_db.db, runtime=_StubRuntime(), sandbox=LocalSandbox() + ) + s1 = await agent.create_session() + s2 = await agent.get_session(s1.id) + assert s2.id == s1.id + assert s2.agent_id == "agent_1" + + +@pytest.mark.asyncio +async def test_get_session_raises_for_unknown_id(fake_db): + row = await fake_db.db.litellm_agent.create( + data={"id": "agent_1", "name": "x", "model": "m", "user_api_key_hash": "h"} + ) + agent = await Agent.from_db_row( + row, db=fake_db.db, runtime=_StubRuntime(), sandbox=LocalSandbox() + ) + with pytest.raises(LookupError): + await agent.get_session("sess_does_not_exist") + + +@pytest.mark.asyncio +async def test_list_sessions_returns_all_for_agent(fake_db): + row = await fake_db.db.litellm_agent.create( + data={"id": "agent_1", "name": "x", "model": "m", "user_api_key_hash": "h"} + ) + agent = await Agent.from_db_row( + row, db=fake_db.db, runtime=_StubRuntime(), sandbox=LocalSandbox() + ) + s_a = await agent.create_session() + s_b = await agent.create_session() + + sessions = await agent.list_sessions() + ids = {s.id for s in sessions} + assert ids == {s_a.id, s_b.id} + + +@pytest.mark.asyncio +async def test_create_session_requires_runtime_and_sandbox(fake_db): + row = await fake_db.db.litellm_agent.create( + data={"id": "agent_1", "name": "x", "model": "m", "user_api_key_hash": "h"} + ) + agent = await Agent.from_db_row(row, db=fake_db.db) # no runtime/sandbox + with pytest.raises(RuntimeError, match="runtime and sandbox"): + await agent.create_session() + + +@pytest.mark.asyncio +async def test_delete_removes_agent_row(fake_db): + row = await fake_db.db.litellm_agent.create( + data={"id": "agent_1", "name": "x", "model": "m", "user_api_key_hash": "h"} + ) + agent = await Agent.from_db_row( + row, db=fake_db.db, runtime=_StubRuntime(), sandbox=LocalSandbox() + ) + await agent.delete() + assert await fake_db.db.litellm_agent.find_unique(where={"id": "agent_1"}) is None + + +@pytest.mark.asyncio +async def test_to_runtime_config_projects_subset(fake_db): + row = await fake_db.db.litellm_agent.create( + data={ + "id": "agent_1", + "name": "x", + "model": "m", + "user_api_key_hash": "h", + "system_prompt": "sp", + "tools_config": {"tools": []}, + "metadata": {"a": 1}, + } + ) + agent = await Agent.from_db_row( + row, db=fake_db.db, runtime=_StubRuntime(), sandbox=LocalSandbox() + ) + cfg = agent.to_runtime_config() + assert cfg.name == "x" + assert cfg.model == "m" + assert cfg.system_prompt == "sp" + assert cfg.tools_config == {"tools": []} + assert cfg.metadata == {"a": 1} diff --git a/tests/test_litellm/managed_agents/test_runtime_claude_sdk.py b/tests/test_litellm/managed_agents/test_runtime_claude_sdk.py new file mode 100644 index 00000000000..c2c395cb99b --- /dev/null +++ b/tests/test_litellm/managed_agents/test_runtime_claude_sdk.py @@ -0,0 +1,78 @@ +""" +Slow-but-real integration test for ``ClaudeSDKAgentRuntime``. + +Hits the real Anthropic API via ``claude-agent-sdk``. Skipped unless +``ANTHROPIC_API_KEY`` is set in the env (typically loaded from .env in +local dev). Marked ``@pytest.mark.slow`` so the default unit-test pass +(``-k 'not slow'``) skips it. + +What this proves end-to-end: + * The runtime constructs a ``ClaudeAgentOptions`` with the right cwd. + * The SDK actually picks up the cwd and writes its tool output there. + * Our event translation captures at least one assistant_message and + a terminal run_finished event. + +This is an integration smoke — it intentionally does NOT make detailed +assertions about LLM output. The model can phrase things differently +between runs. +""" + +import os +import shutil +import tempfile +from pathlib import Path + +import pytest + +from litellm.managed_agents.agent_runtime.base import AgentConfig, SessionState +from litellm.managed_agents.agent_runtime.claude_sdk import ClaudeSDKAgentRuntime +from litellm.managed_agents.events import EVENT_TYPE_RUN_FINISHED +from litellm.managed_agents.sandbox.local import LocalSandbox + + +pytestmark = [ + pytest.mark.slow, + pytest.mark.skipif( + not os.environ.get("ANTHROPIC_API_KEY"), + reason="ANTHROPIC_API_KEY not set; skipping live Claude SDK test", + ), +] + + +@pytest.mark.asyncio +async def test_create_file_via_claude_sdk(): + pytest.importorskip("claude_agent_sdk") + + workdir = tempfile.mkdtemp(prefix="litellm_managed_agents_test_") + try: + sandbox = LocalSandbox(working_dir=workdir) + runtime = ClaudeSDKAgentRuntime() + session_state = SessionState(session_id="sess_test", cwd=workdir) + agent_config = AgentConfig( + name="filemaker", + model=None, # let the SDK pick its default + system_prompt=( + "You are a filesystem agent. Use the Write tool to do exactly " + "what the user asks." + ), + ) + + terminal_seen = False + async for event in runtime.run( + prompt='Use the Write tool to create a file named "foo.txt" in the ' + 'current working directory containing exactly the text "bar". ' + "Then stop.", + sandbox=sandbox, + session_state=session_state, + agent_config=agent_config, + ): + if event.type == EVENT_TYPE_RUN_FINISHED: + terminal_seen = True + break + + assert terminal_seen, "expected a run_finished event from claude-agent-sdk" + created = Path(workdir) / "foo.txt" + assert created.exists(), f"expected {created} to exist" + assert "bar" in created.read_text() + finally: + shutil.rmtree(workdir, ignore_errors=True) diff --git a/tests/test_litellm/managed_agents/test_runtime_litellm_native.py b/tests/test_litellm/managed_agents/test_runtime_litellm_native.py new file mode 100644 index 00000000000..9681cd893a8 --- /dev/null +++ b/tests/test_litellm/managed_agents/test_runtime_litellm_native.py @@ -0,0 +1,184 @@ +""" +Unit tests for ``LiteLLMAgentRuntime``. + +These tests exercise the manual tool loop without hitting a real LLM. +``litellm.acompletion`` is monkeypatched to a scripted async function +that returns a pre-baked sequence of completions, letting us assert: + * tool calls flow into ``sandbox.execute_tool`` + * the loop terminates on a no-tool-call assistant message + * tool results are appended to the message history correctly +""" + +import json +from typing import Any, Dict, List + +import pytest + +from litellm.managed_agents.agent_runtime.base import AgentConfig, SessionState +from litellm.managed_agents.agent_runtime.litellm_native import LiteLLMAgentRuntime +from litellm.managed_agents.events import ( + EVENT_TYPE_ASSISTANT_MESSAGE, + EVENT_TYPE_RUN_FINISHED, + EVENT_TYPE_TOOL_RESULT, + EVENT_TYPE_TOOL_USE, +) +from litellm.managed_agents.sandbox.base import Sandbox, ToolResult + + +class _RecordingSandbox(Sandbox): + """Sandbox that records every tool call and returns a configurable result.""" + + def __init__(self, response: str = "ok"): + self.calls: List[Dict[str, Any]] = [] + self.response = response + + async def execute_tool(self, tool_name: str, tool_input: Dict[str, Any]): + self.calls.append({"name": tool_name, "input": tool_input}) + return ToolResult(output=self.response) + + +def _completion(content=None, tool_calls=None, finish_reason="stop"): + """Build an OpenAI-shaped completion dict the runtime knows how to parse.""" + msg: Dict[str, Any] = {"role": "assistant", "content": content or ""} + if tool_calls is not None: + msg["tool_calls"] = tool_calls + return {"choices": [{"index": 0, "finish_reason": finish_reason, "message": msg}]} + + +def _make_acompletion_stub(scripted: List[Dict[str, Any]]): + """Return an async stub that yields the next scripted response per call. + + Accepts arbitrary kwargs (matching litellm.acompletion's signature) so + a stale stub doesn't fail with 'unexpected keyword argument' if the + runtime starts passing more kwargs in the future. + """ + calls = {"n": 0, "kwargs": []} + + async def fake_acompletion(*args, **kwargs): + calls["kwargs"].append(kwargs) + idx = calls["n"] + calls["n"] += 1 + if idx >= len(scripted): + raise AssertionError( + f"acompletion called more times than scripted: {idx + 1} > {len(scripted)}" + ) + return scripted[idx] + + return fake_acompletion, calls + + +@pytest.mark.asyncio +async def test_terminates_when_assistant_emits_text_only(monkeypatch): + fake, _ = _make_acompletion_stub( + [_completion(content="all done")], + ) + monkeypatch.setattr("litellm.acompletion", fake) + + runtime = LiteLLMAgentRuntime() + sandbox = _RecordingSandbox() + events = [] + async for ev in runtime.run( + prompt="say done", + sandbox=sandbox, + session_state=SessionState(session_id="sess"), + agent_config=AgentConfig(name="x", model="gpt-4o-mini"), + ): + events.append(ev) + + types = [e.type for e in events] + assert types == [EVENT_TYPE_ASSISTANT_MESSAGE, EVENT_TYPE_RUN_FINISHED] + assert events[0].data["content"] == "all done" + # No tool calls should have been routed to the sandbox. + assert sandbox.calls == [] + + +@pytest.mark.asyncio +async def test_routes_tool_call_to_sandbox(monkeypatch): + tool_call = { + "id": "call_123", + "type": "function", + "function": { + "name": "Bash", + "arguments": json.dumps({"command": "echo hi"}), + }, + } + fake, calls = _make_acompletion_stub( + [ + _completion(tool_calls=[tool_call], finish_reason="tool_calls"), + _completion(content="all done"), + ] + ) + monkeypatch.setattr("litellm.acompletion", fake) + + runtime = LiteLLMAgentRuntime() + sandbox = _RecordingSandbox(response="hi\n") + events = [] + async for ev in runtime.run( + prompt="run echo hi", + sandbox=sandbox, + session_state=SessionState(session_id="sess"), + agent_config=AgentConfig(name="x", model="gpt-4o-mini"), + ): + events.append(ev) + + types = [e.type for e in events] + assert types == [ + EVENT_TYPE_TOOL_USE, + EVENT_TYPE_TOOL_RESULT, + EVENT_TYPE_ASSISTANT_MESSAGE, + EVENT_TYPE_RUN_FINISHED, + ] + assert sandbox.calls == [{"name": "Bash", "input": {"command": "echo hi"}}] + + # Second acompletion call should have included the tool result message. + second_kwargs = calls["kwargs"][1] + msgs = second_kwargs["messages"] + assert any(m.get("role") == "tool" and m.get("content") == "hi\n" for m in msgs) + + +@pytest.mark.asyncio +async def test_default_tools_passed_when_config_empty(monkeypatch): + fake, calls = _make_acompletion_stub([_completion(content="bye")]) + monkeypatch.setattr("litellm.acompletion", fake) + + runtime = LiteLLMAgentRuntime() + async for _ in runtime.run( + prompt="hi", + sandbox=_RecordingSandbox(), + session_state=SessionState(session_id="sess"), + agent_config=AgentConfig(name="x", model="gpt-4o-mini"), + ): + pass + + first_kwargs = calls["kwargs"][0] + tool_names = {t["function"]["name"] for t in first_kwargs["tools"]} + assert {"Bash", "Read", "Write", "Edit", "ls"} <= tool_names + + +@pytest.mark.asyncio +async def test_max_turns_terminates_with_max_turns_reason(monkeypatch): + """When the LLM keeps asking for tool calls forever, we still stop cleanly.""" + forever_tool_call = { + "id": "call_forever", + "type": "function", + "function": {"name": "Bash", "arguments": "{}"}, + } + # Always return the same tool call so the loop runs to max_turns. + scripted = [ + _completion(tool_calls=[forever_tool_call], finish_reason="tool_calls") + for _ in range(3) + ] + fake, _ = _make_acompletion_stub(scripted) + monkeypatch.setattr("litellm.acompletion", fake) + + runtime = LiteLLMAgentRuntime(max_turns=3) + types = [] + async for ev in runtime.run( + prompt="loop", + sandbox=_RecordingSandbox(), + session_state=SessionState(session_id="sess"), + agent_config=AgentConfig(name="x", model="gpt-4o-mini"), + ): + types.append((ev.type, ev.data.get("stop_reason"))) + + assert types[-1] == (EVENT_TYPE_RUN_FINISHED, "max_turns") diff --git a/tests/test_litellm/managed_agents/test_sandbox_local.py b/tests/test_litellm/managed_agents/test_sandbox_local.py new file mode 100644 index 00000000000..275095d64bc --- /dev/null +++ b/tests/test_litellm/managed_agents/test_sandbox_local.py @@ -0,0 +1,132 @@ +""" +Unit tests for ``LocalSandbox``. + +Covers the tool surface (Bash / Read / Write / Edit / ls) plus the +sandbox lifecycle (setup creates a tmpdir, teardown removes it). +""" + +import os +import tempfile +from pathlib import Path + +import pytest + +from litellm.managed_agents.sandbox.base import ToolResult +from litellm.managed_agents.sandbox.local import LocalSandbox + + +@pytest.mark.asyncio +async def test_setup_creates_tmpdir_and_teardown_removes_it(): + sb = LocalSandbox() + await sb.setup() + assert sb.cwd is not None and os.path.isdir(sb.cwd) + cwd = sb.cwd + await sb.teardown() + assert not os.path.exists(cwd) + assert sb.cwd is None + + +@pytest.mark.asyncio +async def test_setup_respects_explicit_working_dir(): + with tempfile.TemporaryDirectory() as tmp: + sb = LocalSandbox(working_dir=tmp) + await sb.setup() + assert sb.cwd == tmp + await sb.teardown() + # Explicit dir is left intact (we don't own it). + assert os.path.isdir(tmp) + + +@pytest.mark.asyncio +async def test_bash_runs_command_and_returns_stdout(): + with tempfile.TemporaryDirectory() as tmp: + sb = LocalSandbox(working_dir=tmp) + result = await sb.execute_tool("Bash", {"command": "echo hello"}) + assert isinstance(result, ToolResult) + assert not result.is_error + assert result.output.strip() == "hello" + + +@pytest.mark.asyncio +async def test_bash_nonzero_exit_marks_is_error(): + with tempfile.TemporaryDirectory() as tmp: + sb = LocalSandbox(working_dir=tmp) + result = await sb.execute_tool("Bash", {"command": "exit 7"}) + assert result.is_error + assert result.metadata.get("exit_code") == 7 + + +@pytest.mark.asyncio +async def test_bash_empty_command_returns_error(): + sb = LocalSandbox() + result = await sb.execute_tool("Bash", {}) + assert result.is_error + assert "command" in result.output + + +@pytest.mark.asyncio +async def test_write_then_read_roundtrip(): + with tempfile.TemporaryDirectory() as tmp: + sb = LocalSandbox(working_dir=tmp) + write = await sb.execute_tool("Write", {"path": "a.txt", "content": "hi"}) + assert not write.is_error + assert (Path(tmp) / "a.txt").read_text() == "hi" + + read = await sb.execute_tool("Read", {"path": "a.txt"}) + assert not read.is_error + assert read.output == "hi" + + +@pytest.mark.asyncio +async def test_edit_replaces_substring(): + with tempfile.TemporaryDirectory() as tmp: + sb = LocalSandbox(working_dir=tmp) + await sb.execute_tool("Write", {"path": "a.txt", "content": "hello world"}) + edit = await sb.execute_tool( + "Edit", + {"path": "a.txt", "old_string": "world", "new_string": "moon"}, + ) + assert not edit.is_error + assert (Path(tmp) / "a.txt").read_text() == "hello moon" + + +@pytest.mark.asyncio +async def test_edit_missing_old_string_is_error(): + with tempfile.TemporaryDirectory() as tmp: + sb = LocalSandbox(working_dir=tmp) + await sb.execute_tool("Write", {"path": "a.txt", "content": "hello"}) + edit = await sb.execute_tool( + "Edit", + {"path": "a.txt", "old_string": "nope", "new_string": "_"}, + ) + assert edit.is_error + assert "not found" in edit.output + + +@pytest.mark.asyncio +async def test_ls_lists_dir_entries(): + with tempfile.TemporaryDirectory() as tmp: + (Path(tmp) / "a.txt").write_text("") + (Path(tmp) / "b.txt").write_text("") + sb = LocalSandbox(working_dir=tmp) + result = await sb.execute_tool("ls", {"path": tmp}) + assert not result.is_error + names = sorted(result.output.split("\n")) + assert names == ["a.txt", "b.txt"] + + +@pytest.mark.asyncio +async def test_unknown_tool_returns_error(): + sb = LocalSandbox() + result = await sb.execute_tool("Magic", {}) + assert result.is_error + assert "unknown tool" in result.output + + +@pytest.mark.asyncio +async def test_bash_timeout_kills_command(): + with tempfile.TemporaryDirectory() as tmp: + sb = LocalSandbox(working_dir=tmp, shell_timeout_seconds=0.5) + result = await sb.execute_tool("Bash", {"command": "sleep 5"}) + assert result.is_error + assert "timed out" in result.output diff --git a/tests/test_litellm/managed_agents/test_session.py b/tests/test_litellm/managed_agents/test_session.py new file mode 100644 index 00000000000..525e0c9465a --- /dev/null +++ b/tests/test_litellm/managed_agents/test_session.py @@ -0,0 +1,184 @@ +""" +Unit tests for ``litellm.managed_agents.Session``. + +Covers the ``send()`` happy path: events are persisted to +``LiteLLM_AgentRunEvent`` in seq order, and session/run statuses flip +in lock-step. Also covers the error path where the runtime raises and +we still restore the session to ``ready``. +""" + +from typing import AsyncIterator + +import pytest + +from litellm.managed_agents.agent import Agent +from litellm.managed_agents.agent_runtime.base import ( + AgentConfig, + AgentRuntime, + SessionState, +) +from litellm.managed_agents.events import ( + EVENT_TYPE_ASSISTANT_MESSAGE, + EVENT_TYPE_RUN_FINISHED, + EVENT_TYPE_TOOL_RESULT, + EVENT_TYPE_TOOL_USE, + Event, +) +from litellm.managed_agents.sandbox.local import LocalSandbox + + +class _ScriptedRuntime(AgentRuntime): + """Yield a fixed list of events (or raise) so tests can assert wire shape.""" + + def __init__(self, events=None, raises=None): + self.events = events or [] + self.raises = raises + self.runs_called = 0 + + async def run( + self, prompt, sandbox, session_state: SessionState, agent_config: AgentConfig + ) -> AsyncIterator[Event]: + self.runs_called += 1 + for ev in self.events: + yield ev + if self.raises is not None: + raise self.raises + + +async def _make_session(fake_db, runtime): + row = await fake_db.db.litellm_agent.create( + data={ + "id": "agent_1", + "name": "x", + "model": "m", + "user_api_key_hash": "h", + "system_prompt": "sp", + } + ) + agent = await Agent.from_db_row( + row, db=fake_db.db, runtime=runtime, sandbox=LocalSandbox() + ) + return await agent.create_session() + + +@pytest.mark.asyncio +async def test_send_persists_events_and_yields_them(fake_db): + events = [ + Event(type=EVENT_TYPE_ASSISTANT_MESSAGE, data={"content": "hello"}), + Event( + type=EVENT_TYPE_TOOL_USE, + data={"tool_use_id": "t1", "tool": "Bash", "input": {"command": "ls"}}, + ), + Event( + type=EVENT_TYPE_TOOL_RESULT, + data={"tool_use_id": "t1", "output": "ok", "is_error": False}, + ), + Event( + type=EVENT_TYPE_RUN_FINISHED, + data={"result": "done", "is_error": False, "stop_reason": "stop"}, + ), + ] + session = await _make_session(fake_db, _ScriptedRuntime(events=events)) + + yielded = [] + async for event in session.send("do the thing"): + yielded.append(event) + assert [e.type for e in yielded] == [e.type for e in events] + + # Persisted row count == number of events; seq is monotonic. + persisted = await fake_db.db.litellm_agentrunevent.find_many(where={}) + persisted.sort(key=lambda r: r.seq) + assert [r.event_type for r in persisted] == [e.type for e in events] + assert [r.seq for r in persisted] == [1, 2, 3, 4] + + +@pytest.mark.asyncio +async def test_send_flips_session_and_run_status(fake_db): + session = await _make_session( + fake_db, + _ScriptedRuntime( + events=[ + Event( + type=EVENT_TYPE_RUN_FINISHED, + data={"result": "ok"}, + ) + ] + ), + ) + + # ready before + s_pre = await fake_db.db.litellm_agentsession.find_unique(where={"id": session.id}) + assert s_pre.status == "ready" + + # drain + async for _ in session.send("hi"): + pass + + # ready after (busy was visible mid-flight; we settle it back). + s_post = await fake_db.db.litellm_agentsession.find_unique(where={"id": session.id}) + assert s_post.status == "ready" + assert session.status == "ready" + + # Run row terminal. + runs = await fake_db.db.litellm_agentrun.find_many(where={"session_id": session.id}) + assert len(runs) == 1 + assert runs[0].status == "finished" + assert runs[0].result == "ok" + assert runs[0].terminated_at is not None + assert runs[0].started_at is not None + + +@pytest.mark.asyncio +async def test_send_restores_session_on_runtime_error(fake_db): + session = await _make_session( + fake_db, + _ScriptedRuntime(events=[], raises=RuntimeError("boom")), + ) + with pytest.raises(RuntimeError, match="boom"): + async for _ in session.send("hi"): + pass + + # session bounced back to ready even though we raised + s_post = await fake_db.db.litellm_agentsession.find_unique(where={"id": session.id}) + assert s_post.status == "ready" + runs = await fake_db.db.litellm_agentrun.find_many(where={"session_id": session.id}) + assert len(runs) == 1 + assert runs[0].status == "error" + + +@pytest.mark.asyncio +async def test_get_run_returns_run(fake_db): + session = await _make_session( + fake_db, + _ScriptedRuntime( + events=[Event(type=EVENT_TYPE_RUN_FINISHED, data={"result": "ok"})] + ), + ) + async for _ in session.send("hi"): + pass + runs = await session.list_runs() + assert len(runs) == 1 + fetched = await session.get_run(runs[0].id) + assert fetched.id == runs[0].id + assert fetched.status == "finished" + + +@pytest.mark.asyncio +async def test_conversation_includes_persisted_events(fake_db): + session = await _make_session( + fake_db, + _ScriptedRuntime( + events=[ + Event(type=EVENT_TYPE_ASSISTANT_MESSAGE, data={"content": "hi"}), + Event(type=EVENT_TYPE_RUN_FINISHED, data={"result": "ok"}), + ] + ), + ) + async for _ in session.send("hello"): + pass + convo = await session.conversation() + assert [m["event_type"] for m in convo] == [ + EVENT_TYPE_ASSISTANT_MESSAGE, + EVENT_TYPE_RUN_FINISHED, + ] + assert convo[0]["payload"] == {"content": "hi"} diff --git a/tests/test_litellm/managed_agents/test_subclass_override.py b/tests/test_litellm/managed_agents/test_subclass_override.py new file mode 100644 index 00000000000..e6de3b32919 --- /dev/null +++ b/tests/test_litellm/managed_agents/test_subclass_override.py @@ -0,0 +1,182 @@ +""" +Tests for subclassing ``AgentRuntime`` to add ``before_tool_call`` / +``after_tool_call`` hooks. + +This is the primary extension point for callers who need to audit, +redact, or rewrite tool calls without forking the entire runtime. +""" + +import json +from typing import Any, Dict, List + +import pytest + +from litellm.managed_agents.agent_runtime.base import AgentConfig, SessionState +from litellm.managed_agents.agent_runtime.litellm_native import LiteLLMAgentRuntime +from litellm.managed_agents.events import ( + EVENT_TYPE_RUN_FINISHED, + EVENT_TYPE_TOOL_RESULT, +) +from litellm.managed_agents.sandbox.base import Sandbox, ToolResult + + +class _PassthroughSandbox(Sandbox): + """Sandbox that echoes the tool name + input back so we can see what hook saw.""" + + def __init__(self): + self.invocations: List[Dict[str, Any]] = [] + + async def execute_tool(self, tool_name: str, tool_input: Dict[str, Any]): + self.invocations.append({"name": tool_name, "input": tool_input}) + return ToolResult(output=f"executed {tool_name} with {tool_input!r}") + + +class _AuditingRuntime(LiteLLMAgentRuntime): + """Subclass that records every hook fire and rewrites tool input.""" + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.before_calls: List[Dict[str, Any]] = [] + self.after_calls: List[Dict[str, Any]] = [] + + async def before_tool_call(self, tool_name, tool_input): + self.before_calls.append({"name": tool_name, "input": dict(tool_input)}) + # Demonstrate rewriting: append " (audited)" to any command. + if tool_name == "Bash" and "command" in tool_input: + return {**tool_input, "command": tool_input["command"] + " (audited)"} + return tool_input + + async def after_tool_call(self, tool_name, tool_input, result): + self.after_calls.append( + { + "name": tool_name, + "input": dict(tool_input), + "is_error": result.is_error, + } + ) + return result + + +def _completion_with_call(call_id, name, args): + return { + "choices": [ + { + "index": 0, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": call_id, + "type": "function", + "function": { + "name": name, + "arguments": json.dumps(args), + }, + } + ], + }, + } + ] + } + + +def _completion_text(text): + return { + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": text}, + } + ] + } + + +@pytest.mark.asyncio +async def test_before_and_after_hooks_fire_with_rewrite(monkeypatch): + scripted = [ + _completion_with_call("call_1", "Bash", {"command": "ls"}), + _completion_text("done"), + ] + call_state = {"i": 0} + + async def fake_acompletion(**kwargs): + i = call_state["i"] + call_state["i"] += 1 + return scripted[i] + + monkeypatch.setattr("litellm.acompletion", fake_acompletion) + + runtime = _AuditingRuntime() + sandbox = _PassthroughSandbox() + + events = [] + async for ev in runtime.run( + prompt="run ls", + sandbox=sandbox, + session_state=SessionState(session_id="sess"), + agent_config=AgentConfig(name="x", model="gpt-4o-mini"), + ): + events.append(ev) + + # both hooks fired exactly once (one tool call) + assert len(runtime.before_calls) == 1 + assert runtime.before_calls[0]["name"] == "Bash" + assert runtime.before_calls[0]["input"] == {"command": "ls"} + + assert len(runtime.after_calls) == 1 + assert runtime.after_calls[0]["name"] == "Bash" + + # rewrite is honoured: sandbox saw the augmented command, not the original + assert sandbox.invocations == [ + {"name": "Bash", "input": {"command": "ls (audited)"}} + ] + + # tool_result event in the stream reflects what the sandbox returned + tool_results = [e for e in events if e.type == EVENT_TYPE_TOOL_RESULT] + assert len(tool_results) == 1 + assert "ls (audited)" in tool_results[0].data["output"] + + # And we hit the terminal event at the end. + assert events[-1].type == EVENT_TYPE_RUN_FINISHED + + +class _BeforeRaisesRuntime(LiteLLMAgentRuntime): + async def before_tool_call(self, tool_name, tool_input): + raise ValueError("nope") + + +@pytest.mark.asyncio +async def test_before_hook_raises_surfaces_as_tool_result_error(monkeypatch): + scripted = [ + _completion_with_call("call_1", "Bash", {"command": "ls"}), + _completion_text("ok"), + ] + call_state = {"i": 0} + + async def fake_acompletion(**kwargs): + i = call_state["i"] + call_state["i"] += 1 + return scripted[i] + + monkeypatch.setattr("litellm.acompletion", fake_acompletion) + + runtime = _BeforeRaisesRuntime() + sandbox = _PassthroughSandbox() + events = [] + async for ev in runtime.run( + prompt="run ls", + sandbox=sandbox, + session_state=SessionState(session_id="sess"), + agent_config=AgentConfig(name="x", model="gpt-4o-mini"), + ): + events.append(ev) + + # Sandbox was never reached — the hook short-circuited. + assert sandbox.invocations == [] + tool_results = [e for e in events if e.type == EVENT_TYPE_TOOL_RESULT] + assert len(tool_results) == 1 + assert tool_results[0].data["is_error"] is True + assert "before_tool_call" in tool_results[0].data["output"]