mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
merge: integrate litellm/managed_agents/ (Epic C pivot — Python module wrapping claude-agent-sdk)
This commit is contained in:
commit
c0ecb0d9ad
20 changed files with 3195 additions and 0 deletions
69
litellm/managed_agents/__init__.py
Normal file
69
litellm/managed_agents/__init__.py
Normal 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",
|
||||
]
|
||||
331
litellm/managed_agents/agent.py
Normal file
331
litellm/managed_agents/agent.py
Normal 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 []
|
||||
120
litellm/managed_agents/agent_runtime/base.py
Normal file
120
litellm/managed_agents/agent_runtime/base.py
Normal 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
|
||||
220
litellm/managed_agents/agent_runtime/claude_sdk.py
Normal file
220
litellm/managed_agents/agent_runtime/claude_sdk.py
Normal 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)
|
||||
532
litellm/managed_agents/agent_runtime/litellm_native.py
Normal file
532
litellm/managed_agents/agent_runtime/litellm_native.py
Normal 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)
|
||||
64
litellm/managed_agents/events.py
Normal file
64
litellm/managed_agents/events.py
Normal 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)
|
||||
130
litellm/managed_agents/run.py
Normal file
130
litellm/managed_agents/run.py
Normal 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),
|
||||
)
|
||||
12
litellm/managed_agents/sandbox/__init__.py
Normal file
12
litellm/managed_agents/sandbox/__init__.py
Normal 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",
|
||||
]
|
||||
94
litellm/managed_agents/sandbox/base.py
Normal file
94
litellm/managed_agents/sandbox/base.py
Normal 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
|
||||
57
litellm/managed_agents/sandbox/ec2_ssm.py
Normal file
57
litellm/managed_agents/sandbox/ec2_ssm.py
Normal 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."
|
||||
)
|
||||
199
litellm/managed_agents/sandbox/local.py
Normal file
199
litellm/managed_agents/sandbox/local.py
Normal 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)
|
||||
323
litellm/managed_agents/session.py
Normal file
323
litellm/managed_agents/session.py
Normal 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)
|
||||
0
tests/test_litellm/managed_agents/__init__.py
Normal file
0
tests/test_litellm/managed_agents/__init__.py
Normal file
54
tests/test_litellm/managed_agents/conftest.py
Normal file
54
tests/test_litellm/managed_agents/conftest.py
Normal 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()
|
||||
230
tests/test_litellm/managed_agents/test_agent.py
Normal file
230
tests/test_litellm/managed_agents/test_agent.py
Normal 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}
|
||||
78
tests/test_litellm/managed_agents/test_runtime_claude_sdk.py
Normal file
78
tests/test_litellm/managed_agents/test_runtime_claude_sdk.py
Normal 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)
|
||||
184
tests/test_litellm/managed_agents/test_runtime_litellm_native.py
Normal file
184
tests/test_litellm/managed_agents/test_runtime_litellm_native.py
Normal 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")
|
||||
132
tests/test_litellm/managed_agents/test_sandbox_local.py
Normal file
132
tests/test_litellm/managed_agents/test_sandbox_local.py
Normal 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
|
||||
184
tests/test_litellm/managed_agents/test_session.py
Normal file
184
tests/test_litellm/managed_agents/test_session.py
Normal 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"}
|
||||
182
tests/test_litellm/managed_agents/test_subclass_override.py
Normal file
182
tests/test_litellm/managed_agents/test_subclass_override.py
Normal 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"]
|
||||
Loading…
Add table
Reference in a new issue