merge: integrate litellm/managed_agents/ (Epic C pivot — Python module wrapping claude-agent-sdk)

This commit is contained in:
Ishaan Jaffer 2026-05-06 18:03:16 -07:00
commit c0ecb0d9ad
No known key found for this signature in database
20 changed files with 3195 additions and 0 deletions

View file

@ -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",
]

View file

@ -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=<client> "
"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 []

View file

@ -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

View file

@ -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)

View file

@ -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": [<openai-tool-dict>, ...]}`` — pre-formatted OpenAI/LiteLLM
tool definitions. Passed straight through to ``acompletion``.
* ``[<openai-tool-dict>, ...]`` — 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)

View file

@ -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)

View file

@ -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=<client> "
"or via Run.from_db_row(row, db=<client>)."
)
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),
)

View file

@ -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",
]

View file

@ -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=<error msg>, 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

View file

@ -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."
)

View file

@ -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)

View file

@ -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)

View file

@ -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()

View file

@ -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}

View file

@ -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)

View file

@ -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")

View file

@ -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

View file

@ -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"}

View file

@ -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"]