merge: integrate A1 (LIT-2877 /v2/sessions) into LIT-2890 base

This commit is contained in:
Ishaan Jaffer 2026-05-06 15:43:16 -07:00
commit 8131c7df0c
No known key found for this signature in database
34 changed files with 5111 additions and 0 deletions

View file

@ -0,0 +1,111 @@
-- CreateTable
CREATE TABLE "LiteLLM_Agent" (
"id" TEXT NOT NULL,
"name" TEXT NOT NULL,
"user_api_key_hash" TEXT NOT NULL,
"team_id" TEXT,
"model" TEXT NOT NULL,
"system_prompt" TEXT,
"default_repos" JSONB,
"default_env_vars" JSONB,
"tools_config" JSONB,
"metadata" JSONB,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"updated_at" TIMESTAMP(3) NOT NULL,
CONSTRAINT "LiteLLM_Agent_pkey" PRIMARY KEY ("id")
);
-- CreateTable
CREATE TABLE "LiteLLM_AgentSession" (
"id" TEXT NOT NULL,
"agent_id" TEXT NOT NULL,
"user_api_key_hash" TEXT NOT NULL,
"team_id" TEXT,
"vm_id" TEXT,
"vm_provider" TEXT,
"repos" JSONB NOT NULL,
"env_vars" JSONB,
"status" TEXT NOT NULL DEFAULT 'provisioning',
"daemon_token_hash" TEXT,
"expires_at" TIMESTAMP(3) NOT NULL,
"last_heartbeat_at" TIMESTAMP(3),
"idempotency_key" TEXT,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"updated_at" TIMESTAMP(3) NOT NULL,
"terminated_at" TIMESTAMP(3),
CONSTRAINT "LiteLLM_AgentSession_pkey" PRIMARY KEY ("id")
);
-- CreateTable
CREATE TABLE "LiteLLM_AgentRun" (
"id" TEXT NOT NULL,
"session_id" TEXT NOT NULL,
"parent_run_id" TEXT,
"status" TEXT NOT NULL DEFAULT 'queued',
"prompt" JSONB NOT NULL,
"result" TEXT,
"git_branches" JSONB,
"idempotency_key" TEXT,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"updated_at" TIMESTAMP(3) NOT NULL,
"started_at" TIMESTAMP(3),
"terminated_at" TIMESTAMP(3),
CONSTRAINT "LiteLLM_AgentRun_pkey" PRIMARY KEY ("id")
);
-- CreateTable
CREATE TABLE "LiteLLM_AgentRunEvent" (
"id" TEXT NOT NULL,
"run_id" TEXT NOT NULL,
"seq" INTEGER NOT NULL,
"event_type" TEXT NOT NULL,
"payload" JSONB NOT NULL,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
CONSTRAINT "LiteLLM_AgentRunEvent_pkey" PRIMARY KEY ("id")
);
-- CreateIndex
CREATE INDEX "LiteLLM_Agent_user_api_key_hash_idx" ON "LiteLLM_Agent"("user_api_key_hash");
-- CreateIndex
CREATE INDEX "LiteLLM_Agent_team_id_idx" ON "LiteLLM_Agent"("team_id");
-- CreateIndex
CREATE UNIQUE INDEX "LiteLLM_AgentSession_user_api_key_hash_idempotency_key_key" ON "LiteLLM_AgentSession"("user_api_key_hash", "idempotency_key");
-- CreateIndex
CREATE INDEX "LiteLLM_AgentSession_agent_id_idx" ON "LiteLLM_AgentSession"("agent_id");
-- CreateIndex
CREATE INDEX "LiteLLM_AgentSession_status_expires_at_idx" ON "LiteLLM_AgentSession"("status", "expires_at");
-- CreateIndex
CREATE INDEX "LiteLLM_AgentSession_user_api_key_hash_idx" ON "LiteLLM_AgentSession"("user_api_key_hash");
-- CreateIndex
CREATE UNIQUE INDEX "LiteLLM_AgentRun_session_id_idempotency_key_key" ON "LiteLLM_AgentRun"("session_id", "idempotency_key");
-- CreateIndex
CREATE INDEX "LiteLLM_AgentRun_session_id_status_idx" ON "LiteLLM_AgentRun"("session_id", "status");
-- CreateIndex
CREATE INDEX "LiteLLM_AgentRun_session_id_created_at_idx" ON "LiteLLM_AgentRun"("session_id", "created_at");
-- CreateIndex
CREATE UNIQUE INDEX "LiteLLM_AgentRunEvent_run_id_seq_key" ON "LiteLLM_AgentRunEvent"("run_id", "seq");
-- CreateIndex
CREATE INDEX "LiteLLM_AgentRunEvent_run_id_seq_idx" ON "LiteLLM_AgentRunEvent"("run_id", "seq");
-- AddForeignKey
ALTER TABLE "LiteLLM_AgentSession" ADD CONSTRAINT "LiteLLM_AgentSession_agent_id_fkey" FOREIGN KEY ("agent_id") REFERENCES "LiteLLM_Agent"("id") ON DELETE CASCADE ON UPDATE CASCADE;
-- AddForeignKey
ALTER TABLE "LiteLLM_AgentRun" ADD CONSTRAINT "LiteLLM_AgentRun_session_id_fkey" FOREIGN KEY ("session_id") REFERENCES "LiteLLM_AgentSession"("id") ON DELETE CASCADE ON UPDATE CASCADE;
-- AddForeignKey
ALTER TABLE "LiteLLM_AgentRunEvent" ADD CONSTRAINT "LiteLLM_AgentRunEvent_run_id_fkey" FOREIGN KEY ("run_id") REFERENCES "LiteLLM_AgentRun"("id") ON DELETE CASCADE ON UPDATE CASCADE;

View file

@ -1477,3 +1477,128 @@ model LiteLLM_AgentWorkerPairingToken {
@@index([team_id])
}
// ===========================================================================
// Agent Sessions / Runs (Cursor SDK on LiteLLM)
//
// Three-level hierarchy:
// Agent — definition (model, system prompt, default repos, tools)
// Session — VM-backed conversation, owned by an Agent
// Run — single turn within a Session
// RunEvent — append-only event log per run (for resumable SSE)
// ===========================================================================
model LiteLLM_Agent {
id String @id // "agent_<uuid>"
name String
user_api_key_hash String
team_id String?
model String
system_prompt String?
default_repos Json?
default_env_vars Json?
tools_config Json?
metadata Json?
created_at DateTime @default(now())
updated_at DateTime @updatedAt
sessions LiteLLM_AgentSession[]
@@index([user_api_key_hash])
@@index([team_id])
}
model LiteLLM_AgentSession {
id String @id // "sess_<uuid>"
agent_id String
user_api_key_hash String
team_id String?
vm_id String?
vm_provider String?
repos Json
env_vars Json?
status String @default("provisioning")
daemon_token_hash String?
expires_at DateTime
last_heartbeat_at DateTime?
idempotency_key String?
created_at DateTime @default(now())
updated_at DateTime @updatedAt
terminated_at DateTime?
agent LiteLLM_Agent @relation(fields: [agent_id], references: [id], onDelete: Cascade)
runs LiteLLM_AgentRun[]
@@unique([user_api_key_hash, idempotency_key])
@@index([agent_id])
@@index([status, expires_at])
@@index([user_api_key_hash])
}
model LiteLLM_AgentRun {
id String @id // "run_<uuid>"
session_id String
parent_run_id String?
status String @default("queued")
prompt Json
result String?
git_branches Json?
idempotency_key String?
created_at DateTime @default(now())
updated_at DateTime @updatedAt
started_at DateTime?
terminated_at DateTime?
session LiteLLM_AgentSession @relation(fields: [session_id], references: [id], onDelete: Cascade)
events LiteLLM_AgentRunEvent[]
@@unique([session_id, idempotency_key])
@@index([session_id, status])
@@index([session_id, created_at])
}
model LiteLLM_AgentRunEvent {
id String @id @default(uuid())
run_id String
seq Int
event_type String
payload Json
created_at DateTime @default(now())
run LiteLLM_AgentRun @relation(fields: [run_id], references: [id], onDelete: Cascade)
@@unique([run_id, seq])
@@index([run_id, seq])
}
// ===========================================================================
// Warm pool VM tracking (LIT-2890 / Epic B2)
//
// Tracks the lifecycle of pre-provisioned EC2 (or other provider) VMs used
// for instant session attach. Each row maps to one underlying instance.
//
// State machine:
// provisioning → warm → hydrating → attached → terminating → terminated
//
// On session end, the VM is terminated (NOT recycled) — security boundary.
// The maintenance loop refills `warm` slots; rows in `terminated` are kept
// for audit until pruned.
// ===========================================================================
model LiteLLM_AgentVM {
id String @id // EC2 instance id (e.g. "i-0abcd...")
provider String // "ec2" | "noop" | "self_hosted"
region String?
state String // provisioning|warm|hydrating|attached|terminating|terminated
team_id String // owner team — pool is per-team
pool_id String // logical pool key (currently == team_id)
attached_session_id String? // FK to LiteLLM_AgentSession.id when state=attached
created_at DateTime @default(now())
warmed_at DateTime?
last_hydrate_at DateTime?
terminated_at DateTime?
metadata Json? // public_ip, private_ip, ssm_status, etc.
@@index([state, pool_id])
@@index([team_id, state])
@@index([attached_session_id])
}

View file

@ -0,0 +1,35 @@
"""
Agent / Session / Run REST API surface — `/v2/agents`, `/v2/sessions`,
`/v2/sessions/{sid}/runs`.
The `/v1/agents` namespace is owned by the existing A2A registry
(``litellm/proxy/agent_endpoints/``). To avoid collision, all endpoints
in this module mount under ``/v2/``.
Exports the FastAPI routers consumed by ``proxy_server.py``:
* ``agent_router`` — /v2/agents CRUD
* ``session_router`` — /v2/sessions CRUD + followup + conversation
* ``run_router`` — /v2/sessions/{sid}/runs (+ events SSE, cancel)
* ``internal_router`` — /v2/sessions/{sid}/internal/* (daemon callbacks)
"""
from litellm.proxy.agent_session_endpoints.agent_endpoints import (
router as agent_router,
)
from litellm.proxy.agent_session_endpoints.internal_endpoints import (
router as internal_router,
)
from litellm.proxy.agent_session_endpoints.run_endpoints import (
router as run_router,
)
from litellm.proxy.agent_session_endpoints.session_endpoints import (
router as session_router,
)
__all__ = [
"agent_router",
"session_router",
"run_router",
"internal_router",
]

View file

@ -0,0 +1,218 @@
"""
Agent CRUD endpoints — POST/GET/PATCH/DELETE /v2/agents{,/<id>}.
Note: /v1/agents is the existing A2A registry (litellm/proxy/agent_endpoints/);
this module mounts under /v2/ to avoid collision.
Agents are pure definitions: model, system prompt, default repos, tools.
Creating an agent does NOT spin up a VM — that happens at session create.
DELETE cascades to sessions: every non-terminal session under the agent
is torn down (run cancellation + provider.terminate) before the agent
row is removed.
"""
import asyncio
from datetime import datetime, timezone
from typing import List, Optional
from fastapi import APIRouter, Depends, HTTPException, Query
from fastapi.responses import ORJSONResponse
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_session_endpoints.ids import new_agent_id
from litellm.proxy.agent_session_endpoints.ownership import (
assert_caller_can_mutate,
assert_caller_owns_agent,
caller_api_key_hash,
owner_filter_for_caller,
)
from litellm.proxy.agent_session_endpoints.schemas import (
AgentCreate,
AgentResponse,
AgentUpdate,
)
from litellm.proxy.agent_session_endpoints.serialization import (
agent_row_to_response,
)
from litellm.proxy.agent_session_endpoints.session_endpoints import (
_terminate_session_internal,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
router = APIRouter()
def _now() -> datetime:
return datetime.now(timezone.utc)
def _agent_create_payload(body: AgentCreate, user_api_key_dict: UserAPIKeyAuth) -> dict:
"""Build the Prisma create payload for an agent.
Pydantic models are converted to plain dicts so Prisma can serialize them
into Json columns. ``model_dump(exclude_none=True)`` keeps the row tight —
null Json columns stay null instead of becoming the string "null".
"""
return {
"id": new_agent_id(),
"name": body.name,
"user_api_key_hash": caller_api_key_hash(user_api_key_dict),
"team_id": user_api_key_dict.team_id,
"model": body.model,
"system_prompt": body.system_prompt,
"default_repos": (
[r.model_dump(exclude_none=True) for r in body.default_repos]
if body.default_repos
else None
),
"default_env_vars": body.default_env_vars,
"tools_config": body.tools_config,
"metadata": body.metadata,
"updated_at": _now(),
}
def _agent_update_payload(body: AgentUpdate) -> dict:
"""Build the Prisma update payload — only fields the caller actually
set. Pydantic ``exclude_unset=True`` is the canonical way to do this."""
raw = body.model_dump(exclude_unset=True)
if "default_repos" in raw and raw["default_repos"] is not None:
# Re-dump nested RepoSpec models as plain dicts.
raw["default_repos"] = [
r.model_dump(exclude_none=True) if hasattr(r, "model_dump") else r
for r in body.default_repos or []
]
raw["updated_at"] = _now()
return raw
async def _get_prisma_client_or_503():
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=503, detail="Database unavailable")
return prisma_client
@router.post(
"/v2/agents",
response_class=ORJSONResponse,
response_model=AgentResponse,
tags=["agents"],
dependencies=[Depends(user_api_key_auth)],
)
async def create_agent(
body: AgentCreate,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
assert_caller_can_mutate(user_api_key_dict)
prisma_client = await _get_prisma_client_or_503()
payload = _agent_create_payload(body, user_api_key_dict)
row = await prisma_client.db.litellm_agent.create(data=payload)
verbose_proxy_logger.info(
"agent.create id=%s name=%s model=%s", row.id, row.name, row.model
)
return agent_row_to_response(row)
@router.get(
"/v2/agents/{agent_id}",
response_class=ORJSONResponse,
response_model=AgentResponse,
tags=["agents"],
dependencies=[Depends(user_api_key_auth)],
)
async def get_agent(
agent_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
prisma_client = await _get_prisma_client_or_503()
row = await prisma_client.db.litellm_agent.find_unique(where={"id": agent_id})
assert_caller_owns_agent(user_api_key_dict, row)
return agent_row_to_response(row)
@router.patch(
"/v2/agents/{agent_id}",
response_class=ORJSONResponse,
response_model=AgentResponse,
tags=["agents"],
dependencies=[Depends(user_api_key_auth)],
)
async def update_agent(
agent_id: str,
body: AgentUpdate,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
assert_caller_can_mutate(user_api_key_dict)
prisma_client = await _get_prisma_client_or_503()
existing = await prisma_client.db.litellm_agent.find_unique(where={"id": agent_id})
assert_caller_owns_agent(user_api_key_dict, existing)
payload = _agent_update_payload(body)
if not any(k for k in payload if k != "updated_at"):
# No-op patch — return existing row unchanged (still 200, idempotent).
return agent_row_to_response(existing)
updated = await prisma_client.db.litellm_agent.update(
where={"id": agent_id}, data=payload
)
return agent_row_to_response(updated)
@router.get(
"/v2/agents",
response_class=ORJSONResponse,
tags=["agents"],
dependencies=[Depends(user_api_key_auth)],
)
async def list_agents(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
limit: int = Query(default=100, ge=1, le=500),
offset: int = Query(default=0, ge=0),
):
prisma_client = await _get_prisma_client_or_503()
where_filter = owner_filter_for_caller(user_api_key_dict)
rows = await prisma_client.db.litellm_agent.find_many(
where=where_filter,
order={"created_at": "desc"},
take=limit,
skip=offset,
)
return {"data": [agent_row_to_response(r).model_dump() for r in rows]}
@router.delete(
"/v2/agents/{agent_id}",
response_class=ORJSONResponse,
tags=["agents"],
dependencies=[Depends(user_api_key_auth)],
)
async def delete_agent(
agent_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
assert_caller_can_mutate(user_api_key_dict)
prisma_client = await _get_prisma_client_or_503()
existing = await prisma_client.db.litellm_agent.find_unique(where={"id": agent_id})
assert_caller_owns_agent(user_api_key_dict, existing)
# Cascade: terminate every active session under this agent first. We
# gather them in parallel because each call hits the VM provider.
sessions = await prisma_client.db.litellm_agentsession.find_many(
where={"agent_id": agent_id}
)
active = [s for s in sessions if s.terminated_at is None]
if active:
await asyncio.gather(
*[
_terminate_session_internal(s.id, reason="agent_deleted")
for s in active
],
return_exceptions=True,
)
# Now drop the agent row. Prisma cascade handles FK chains
# (sessions -> runs -> events) for any session still on disk.
await prisma_client.db.litellm_agent.delete(where={"id": agent_id})
return {"id": agent_id, "deleted": True}

View file

@ -0,0 +1,178 @@
"""
JWT minting + validation for the daemon-side internal endpoints.
Daemon tokens are minted at session creation and scoped to a single session.
The token's SHA-256 hash is stored on the session row so revoking the
session also revokes the token without rotating signing keys.
"""
import hashlib
import os
import time
from typing import Any, Dict, Optional
import jwt
from fastapi import Header, HTTPException
from litellm.proxy.agent_session_endpoints.constants import (
AGENT_JWT_ALGORITHM,
AGENT_JWT_SECRET_ENV,
AGENT_RUNTIME_SCOPE,
SESSION_TERMINAL_STATUSES,
)
class AgentDaemonTokenError(Exception):
"""Raised when a daemon token is invalid (used internally by validators)."""
class AgentJWTSecretNotConfiguredError(RuntimeError):
"""Raised when ``LITELLM_AGENT_JWT_SECRET`` is unset.
The daemon JWT secret MUST be a separate credential from the proxy
master key. Falling back to ``LITELLM_MASTER_KEY`` would conflate two
distinct auth surfaces — a captured daemon JWT could then be used to
mint regular API keys with master-key authority.
"""
def is_agent_jwt_secret_configured() -> bool:
"""Return True iff a non-empty ``LITELLM_AGENT_JWT_SECRET`` is present.
Used at startup by ``proxy_server.py`` to decide whether to mount the
agent_session_endpoints routers. Mounting the routers when this is
False would silently expose an unsigned auth surface — refuse instead.
"""
return bool(os.environ.get(AGENT_JWT_SECRET_ENV))
def _get_signing_secret() -> str:
"""Resolve the JWT signing secret.
The daemon JWT secret is a SEPARATE credential from the proxy master
key. There is no fallback — if ``LITELLM_AGENT_JWT_SECRET`` is unset,
we refuse to mint or validate any token. Mounting the agent_session
routers without this env var is itself a startup error (see
``proxy_server.py``).
"""
secret = os.environ.get(AGENT_JWT_SECRET_ENV)
if not secret:
raise AgentJWTSecretNotConfiguredError(
f"{AGENT_JWT_SECRET_ENV} is not set; cannot mint or validate "
"daemon JWTs. This env var must be a dedicated random secret, "
"distinct from LITELLM_MASTER_KEY."
)
return secret
def hash_daemon_token(token: str) -> str:
"""Return ``sha256(token)`` as hex — what we store on the session row."""
return hashlib.sha256(token.encode("utf-8")).hexdigest()
def mint_daemon_token(
session_id: str,
agent_id: str,
expires_at_epoch: int,
) -> str:
"""Mint a session-scoped JWT for the daemon.
Claims:
sub: session id (the only session this token can act on)
agent_id: parent agent id (informational)
iat: issued-at
exp: epoch seconds
scope: "agent_runtime_internal" — gated by a dedicated dependency
so this token CANNOT be used as a normal user virtual key.
"""
now = int(time.time())
payload: Dict[str, Any] = {
"sub": session_id,
"agent_id": agent_id,
"iat": now,
"exp": expires_at_epoch,
"scope": AGENT_RUNTIME_SCOPE,
}
return jwt.encode(payload, _get_signing_secret(), algorithm=AGENT_JWT_ALGORITHM)
def decode_daemon_token(token: str) -> Dict[str, Any]:
"""Decode + verify signature + verify ``exp`` and ``scope``.
Raises HTTPException(401) for any failure so callers can re-raise
directly.
"""
try:
payload = jwt.decode(
token,
_get_signing_secret(),
algorithms=[AGENT_JWT_ALGORITHM],
)
except jwt.ExpiredSignatureError as exc:
raise HTTPException(status_code=401, detail="Daemon token expired") from exc
except jwt.InvalidTokenError as exc:
raise HTTPException(status_code=401, detail="Invalid daemon token") from exc
if payload.get("scope") != AGENT_RUNTIME_SCOPE:
raise HTTPException(
status_code=401,
detail="Daemon token has wrong scope",
)
return payload
def _extract_bearer(authorization: Optional[str]) -> str:
if not authorization:
raise HTTPException(status_code=401, detail="Missing Authorization header")
parts = authorization.split(" ", 1)
if len(parts) != 2 or parts[0].lower() != "bearer":
raise HTTPException(
status_code=401, detail="Authorization header must be 'Bearer <token>'"
)
return parts[1].strip()
async def daemon_token_auth(
session_id: str,
authorization: Optional[str] = Header(default=None),
) -> Dict[str, Any]:
"""FastAPI dependency: validate the daemon JWT for a path-scoped session.
Validation steps (in order):
1. Decode + verify signature, ``exp``, ``scope``.
2. Confirm ``sub == session_id`` (cross-session abuse).
3. Look up session, confirm it's not terminated.
4. Confirm ``sha256(token) == session.daemon_token_hash`` (revoked tokens).
Returns the validated payload dict (with ``_session_row`` injected) so
downstream code can avoid re-fetching the session.
"""
token = _extract_bearer(authorization)
payload = decode_daemon_token(token)
if payload.get("sub") != session_id:
raise HTTPException(status_code=403, detail="Token not valid for this session")
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(
status_code=503, detail="Database unavailable for daemon token check"
)
session = await prisma_client.db.litellm_agentsession.find_unique(
where={"id": session_id}
)
if session is None:
raise HTTPException(status_code=404, detail="Session not found")
if session.status in SESSION_TERMINAL_STATUSES:
# Stored hash check would also fail after termination, but we want
# a 410 Gone so the daemon's systemd unit knows to shutdown.
raise HTTPException(status_code=410, detail="Session terminated")
expected_hash = session.daemon_token_hash
if not expected_hash or hash_daemon_token(token) != expected_hash:
raise HTTPException(status_code=401, detail="Daemon token revoked")
return {**payload, "_session_row": session}

View file

@ -0,0 +1,249 @@
"""
Background sweeper that handles three failure modes:
1. Expired sessions — past ``expires_at``, still non-terminal:
run the standard terminate flow.
2. Dead daemons — non-terminal session whose ``last_heartbeat_at`` is
older than ``DAEMON_HEARTBEAT_DEAD_AFTER_SECONDS``: mark error.
3. Stuck runs — running runs that haven't been touched within
``RUN_IDLE_TIMEOUT_SECONDS``: mark error.
The sweeper is started from ``proxy_server.py``'s startup hooks. It loops
until cancelled, sleeping `CLEANUP_SWEEPER_INTERVAL_SECONDS` between
passes.
"""
import asyncio
from datetime import datetime, timedelta, timezone
from typing import Optional
from litellm._logging import verbose_proxy_logger
from litellm.proxy.agent_session_endpoints.constants import (
CLEANUP_SWEEPER_INTERVAL_SECONDS,
DAEMON_HEARTBEAT_DEAD_AFTER_SECONDS,
EVENT_TYPE_RUN_ERROR,
RUN_ACTIVE_STATUSES,
RUN_IDLE_TIMEOUT_SECONDS,
RUN_STATUS_ERROR,
SESSION_STATUS_ERROR,
SESSION_STATUS_PROVISIONING,
SESSION_STATUS_TERMINATED,
SESSION_TERMINAL_STATUSES,
)
def _now() -> datetime:
return datetime.now(timezone.utc)
async def _sweep_expired_sessions(prisma_client) -> int:
"""Terminate sessions whose ``expires_at`` has passed and are not
already terminal. Returns count terminated."""
from litellm.proxy.agent_session_endpoints.session_endpoints import (
_terminate_session_internal,
)
rows = await prisma_client.db.litellm_agentsession.find_many(
where={
"expires_at": {"lt": _now()},
"status": {
"notIn": list(SESSION_TERMINAL_STATUSES),
},
},
take=200,
)
for row in rows:
try:
await _terminate_session_internal(row.id, reason="expired")
except Exception as exc:
verbose_proxy_logger.exception(
"sweeper: failed to terminate expired session=%s: %s", row.id, exc
)
return len(rows)
async def _sweep_dead_daemons(prisma_client) -> int:
"""Mark sessions ``error`` whose daemon hasn't heartbeat in
``DAEMON_HEARTBEAT_DEAD_AFTER_SECONDS``.
Skips ``provisioning`` sessions (they haven't registered yet — their
own timeout is governed by `expires_at`).
"""
threshold = _now() - timedelta(seconds=DAEMON_HEARTBEAT_DEAD_AFTER_SECONDS)
rows = await prisma_client.db.litellm_agentsession.find_many(
where={
"last_heartbeat_at": {"lt": threshold},
"status": {
"notIn": list(SESSION_TERMINAL_STATUSES)
+ [SESSION_STATUS_PROVISIONING],
},
},
take=200,
)
if not rows:
return 0
now = _now()
ids = [r.id for r in rows]
await prisma_client.db.litellm_agentsession.update_many(
where={"id": {"in": ids}},
data={
"status": SESSION_STATUS_ERROR,
"terminated_at": now,
"updated_at": now,
},
)
# Also flip any active runs in those sessions to error.
runs = await prisma_client.db.litellm_agentrun.find_many(
where={
"session_id": {"in": ids},
"status": {"in": list(RUN_ACTIVE_STATUSES)},
},
take=500,
)
for run in runs:
await prisma_client.db.litellm_agentrun.update(
where={"id": run.id},
data={
"status": RUN_STATUS_ERROR,
"terminated_at": now,
"updated_at": now,
},
)
last_evt = await prisma_client.db.litellm_agentrunevent.find_first(
where={"run_id": run.id}, order={"seq": "desc"}
)
next_seq = (last_evt.seq + 1) if last_evt else 1
try:
await prisma_client.db.litellm_agentrunevent.create(
data={
"run_id": run.id,
"seq": next_seq,
"event_type": EVENT_TYPE_RUN_ERROR,
"payload": {"reason": "daemon_dead"},
}
)
except Exception as exc:
verbose_proxy_logger.warning(
"sweeper: skipped run_error event run=%s seq=%s: %s",
run.id,
next_seq,
exc,
)
verbose_proxy_logger.info(
"sweeper: marked %d sessions error (dead daemon)", len(ids)
)
return len(ids)
async def _sweep_stuck_runs(prisma_client) -> int:
"""Mark runs ``error`` if they've been ``running`` past the idle timeout.
Sweeper-driven only — clients may legitimately have long-running runs
so the threshold is generous (``RUN_IDLE_TIMEOUT_SECONDS``, 30 min by
default).
"""
threshold = _now() - timedelta(seconds=RUN_IDLE_TIMEOUT_SECONDS)
rows = await prisma_client.db.litellm_agentrun.find_many(
where={
"status": "running",
"updated_at": {"lt": threshold},
},
take=200,
)
if not rows:
return 0
now = _now()
for run in rows:
await prisma_client.db.litellm_agentrun.update(
where={"id": run.id},
data={
"status": RUN_STATUS_ERROR,
"terminated_at": now,
"updated_at": now,
},
)
last_evt = await prisma_client.db.litellm_agentrunevent.find_first(
where={"run_id": run.id}, order={"seq": "desc"}
)
next_seq = (last_evt.seq + 1) if last_evt else 1
try:
await prisma_client.db.litellm_agentrunevent.create(
data={
"run_id": run.id,
"seq": next_seq,
"event_type": EVENT_TYPE_RUN_ERROR,
"payload": {"reason": "run_idle_timeout"},
}
)
except Exception as exc:
verbose_proxy_logger.warning(
"sweeper: skipped run_error event run=%s seq=%s: %s",
run.id,
next_seq,
exc,
)
return len(rows)
async def run_cleanup_pass(prisma_client) -> dict:
"""Single pass — exposed for tests and on-demand triggers."""
expired = await _sweep_expired_sessions(prisma_client)
dead = await _sweep_dead_daemons(prisma_client)
stuck = await _sweep_stuck_runs(prisma_client)
return {
"expired_sessions": expired,
"dead_daemon_sessions": dead,
"stuck_runs": stuck,
}
_sweeper_task: Optional[asyncio.Task] = None
async def _cleanup_loop() -> None:
"""Long-running loop. Cancellation-safe — bails out cleanly on
``asyncio.CancelledError``."""
while True:
try:
from litellm.proxy.proxy_server import prisma_client
if prisma_client is not None:
summary = await run_cleanup_pass(prisma_client)
if any(summary.values()):
verbose_proxy_logger.info(
"agent_session_cleanup_sweeper: %s", summary
)
except asyncio.CancelledError:
raise
except Exception as exc:
verbose_proxy_logger.exception(
"agent_session_cleanup_sweeper iteration failed: %s", exc
)
try:
await asyncio.sleep(CLEANUP_SWEEPER_INTERVAL_SECONDS)
except asyncio.CancelledError:
raise
def start_cleanup_sweeper() -> None:
"""Idempotent: spawn the sweeper task once. Subsequent calls are no-ops.
Called from ``proxy_server.startup_event`` so it ships with the proxy.
"""
global _sweeper_task
if _sweeper_task is not None and not _sweeper_task.done():
return
_sweeper_task = asyncio.create_task(_cleanup_loop())
verbose_proxy_logger.info("agent_session_cleanup_sweeper started")
def stop_cleanup_sweeper() -> None:
"""Cancel the sweeper task. Safe to call multiple times."""
global _sweeper_task
if _sweeper_task is not None and not _sweeper_task.done():
_sweeper_task.cancel()
_sweeper_task = None

View file

@ -0,0 +1,61 @@
"""Shared constants for the agent_session_endpoints module."""
# Session lifecycle
DEFAULT_MAX_SESSION_MINUTES = 60 * 4 # 4 hours
DAEMON_HEARTBEAT_DEAD_AFTER_SECONDS = 90
RUN_IDLE_TIMEOUT_SECONDS = 60 * 30 # 30 minutes
# Long-poll for next-run (internal daemon endpoint)
NEXT_RUN_LONG_POLL_TIMEOUT_SECONDS = 30
NEXT_RUN_POLL_INTERVAL_SECONDS = 0.5
# SSE event stream
SSE_KEEPALIVE_INTERVAL_SECONDS = 15
SSE_TERMINAL_QUIESCE_SECONDS = 1.0
SSE_POLL_INTERVAL_SECONDS = 0.25
# Cleanup sweeper
CLEANUP_SWEEPER_INTERVAL_SECONDS = 60
# JWT / auth scope
AGENT_RUNTIME_SCOPE = "agent_runtime_internal"
AGENT_JWT_SECRET_ENV = "LITELLM_AGENT_JWT_SECRET"
AGENT_JWT_ALGORITHM = "HS256"
# Session statuses
SESSION_STATUS_PROVISIONING = "provisioning"
SESSION_STATUS_READY = "ready"
SESSION_STATUS_BUSY = "busy"
SESSION_STATUS_ERROR = "error"
SESSION_STATUS_TERMINATED = "terminated"
SESSION_TERMINAL_STATUSES = {SESSION_STATUS_TERMINATED, SESSION_STATUS_ERROR}
SESSION_ACCEPTING_RUN_STATUSES = {SESSION_STATUS_READY, SESSION_STATUS_BUSY}
# Run statuses
RUN_STATUS_QUEUED = "queued"
RUN_STATUS_RUNNING = "running"
RUN_STATUS_FINISHED = "finished"
RUN_STATUS_CANCELLED = "cancelled"
RUN_STATUS_ERROR = "error"
RUN_ACTIVE_STATUSES = {RUN_STATUS_QUEUED, RUN_STATUS_RUNNING}
RUN_TERMINAL_STATUSES = {RUN_STATUS_FINISHED, RUN_STATUS_CANCELLED, RUN_STATUS_ERROR}
# Event types
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"
RUN_TERMINAL_EVENT_TYPES = {
EVENT_TYPE_RUN_FINISHED: RUN_STATUS_FINISHED,
EVENT_TYPE_RUN_CANCELLED: RUN_STATUS_CANCELLED,
EVENT_TYPE_RUN_ERROR: RUN_STATUS_ERROR,
}
# ID prefixes
AGENT_ID_PREFIX = "agent_"
SESSION_ID_PREFIX = "sess_"
RUN_ID_PREFIX = "run_"

View file

@ -0,0 +1,38 @@
"""
ID generation helpers for the agent_session_endpoints module.
Each ID is ``<prefix>_<32 hex chars>`` (16 random bytes). The prefixes line
up with the model namespaces so logs and DB rows are self-describing:
agent_<...> -> LiteLLM_Agent
sess_<...> -> LiteLLM_AgentSession
run_<...> -> LiteLLM_AgentRun
Implemented as a single ``new_id`` helper rather than three near-duplicate
functions to keep the convention enforced in one place.
"""
import secrets
from litellm.proxy.agent_session_endpoints.constants import (
AGENT_ID_PREFIX,
RUN_ID_PREFIX,
SESSION_ID_PREFIX,
)
def _new_id(prefix: str) -> str:
"""Return ``f"{prefix}{16 random bytes hex}"``."""
return f"{prefix}{secrets.token_hex(16)}"
def new_agent_id() -> str:
return _new_id(AGENT_ID_PREFIX)
def new_session_id() -> str:
return _new_id(SESSION_ID_PREFIX)
def new_run_id() -> str:
return _new_id(RUN_ID_PREFIX)

View file

@ -0,0 +1,291 @@
"""
Daemon-side internal endpoints under ``/v2/sessions/{sid}/internal/...``.
These four routes are how the on-VM daemon talks back to the proxy. They
authenticate with the session-scoped JWT minted on session create, NOT a
regular user virtual key:
* POST /v2/sessions/{sid}/internal/register
Daemon "I'm alive" — flips session provisioning -> ready.
* POST /v2/sessions/{sid}/internal/heartbeat
Periodic liveness ping; bumps last_heartbeat_at.
* GET /v2/sessions/{sid}/runs/next/internal/poll
Long-poll for the next queued run; sets it to running.
* POST /v2/sessions/{sid}/runs/{rid}/events:append
Append an event; flips run status if terminal.
"""
import asyncio
from datetime import datetime, timezone
from typing import Any, Dict, Optional
from fastapi import APIRouter, Depends, HTTPException, Path
from fastapi.responses import ORJSONResponse, Response
from litellm._logging import verbose_proxy_logger
from litellm.proxy.agent_session_endpoints.auth import daemon_token_auth
from litellm.proxy.agent_session_endpoints.constants import (
NEXT_RUN_LONG_POLL_TIMEOUT_SECONDS,
NEXT_RUN_POLL_INTERVAL_SECONDS,
RUN_STATUS_QUEUED,
RUN_STATUS_RUNNING,
RUN_TERMINAL_EVENT_TYPES,
SESSION_STATUS_PROVISIONING,
SESSION_STATUS_READY,
SESSION_TERMINAL_STATUSES,
)
from litellm.proxy.agent_session_endpoints.schemas import (
DaemonHeartbeatRequest,
DaemonRegisterRequest,
EventAppend,
NextRunResponse,
)
router = APIRouter()
def _now() -> datetime:
return datetime.now(timezone.utc)
async def _get_prisma_client_or_503():
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=503, detail="Database unavailable")
return prisma_client
@router.post(
"/v2/sessions/{session_id}/internal/register",
response_class=ORJSONResponse,
tags=["agent-internal"],
)
async def daemon_register(
body: DaemonRegisterRequest,
session_id: str = Path(...),
daemon: Dict[str, Any] = Depends(daemon_token_auth),
):
"""Daemon announces itself: flip session provisioning -> ready.
Idempotent: re-registering an already-ready session is a no-op (200).
Vending the registration is what moves the session out of provisioning,
so the cleanup sweeper's "stuck provisioning" rule doesn't fire after.
"""
prisma_client = await _get_prisma_client_or_503()
now = _now()
session_row = daemon["_session_row"]
update_data: Dict[str, Any] = {
"last_heartbeat_at": now,
"updated_at": now,
}
if body.vm_id and not session_row.vm_id:
update_data["vm_id"] = body.vm_id
if session_row.status == SESSION_STATUS_PROVISIONING:
update_data["status"] = SESSION_STATUS_READY
updated = await prisma_client.db.litellm_agentsession.update(
where={"id": session_id},
data=update_data,
)
verbose_proxy_logger.info(
"session.register session_id=%s status=%s vm_id=%s",
session_id,
updated.status,
updated.vm_id,
)
return {
"session_id": session_id,
"status": updated.status,
"registered_at": now.isoformat(),
}
@router.post(
"/v2/sessions/{session_id}/internal/heartbeat",
response_class=ORJSONResponse,
tags=["agent-internal"],
)
async def daemon_heartbeat(
body: DaemonHeartbeatRequest,
session_id: str = Path(...),
daemon: Dict[str, Any] = Depends(daemon_token_auth),
):
"""Bump ``last_heartbeat_at``. The cleanup sweeper uses this to detect
dead daemons (no heartbeat for ``DAEMON_HEARTBEAT_DEAD_AFTER_SECONDS``)."""
prisma_client = await _get_prisma_client_or_503()
now = _now()
await prisma_client.db.litellm_agentsession.update(
where={"id": session_id},
data={"last_heartbeat_at": now, "updated_at": now},
)
return {"session_id": session_id, "last_heartbeat_at": now.isoformat()}
async def _claim_next_queued_run(prisma_client, session_id: str) -> Optional[Any]:
"""Atomically claim the oldest queued run for ``session_id``.
Best-effort optimistic claim: read oldest queued run, then UPDATE with
``where status = 'queued'`` so two concurrent daemons can't both move
the same row to ``running`` (Prisma's ``update_many`` returns a count;
we re-read to confirm we won).
A future version could lift this to a single ``RETURNING`` statement
via raw SQL or a Prisma ``transaction`` — for the noop provider this
is sufficient because the daemon side runs single-tenant anyway.
"""
candidate = await prisma_client.db.litellm_agentrun.find_first(
where={"session_id": session_id, "status": RUN_STATUS_QUEUED},
order={"created_at": "asc"},
)
if candidate is None:
return None
now = _now()
result = await prisma_client.db.litellm_agentrun.update_many(
where={"id": candidate.id, "status": RUN_STATUS_QUEUED},
data={
"status": RUN_STATUS_RUNNING,
"started_at": now,
"updated_at": now,
},
)
# Prisma returns either a count (int) or an object with `.count`.
count = getattr(result, "count", result)
if not count:
return None # someone else won — caller can retry
return await prisma_client.db.litellm_agentrun.find_unique(
where={"id": candidate.id}
)
@router.get(
"/v2/sessions/{session_id}/runs/next/internal/poll",
response_class=ORJSONResponse,
response_model=NextRunResponse,
tags=["agent-internal"],
)
async def daemon_next_run(
session_id: str = Path(...),
daemon: Dict[str, Any] = Depends(daemon_token_auth),
):
"""Long-poll: return the next queued run, claiming it as ``running``.
Loops up to ``NEXT_RUN_LONG_POLL_TIMEOUT_SECONDS`` polling every
``NEXT_RUN_POLL_INTERVAL_SECONDS``. Returns 204 if nothing showed up
(daemon will re-poll). The claim step uses an optimistic
``update_many WHERE status='queued'`` so two daemons can't both grab
the same row.
"""
prisma_client = await _get_prisma_client_or_503()
# ``get_running_loop`` (not ``get_event_loop``) is the correct API
# inside a running coroutine — see Python 3.10 deprecation.
deadline = asyncio.get_running_loop().time() + NEXT_RUN_LONG_POLL_TIMEOUT_SECONDS
while True:
claimed = await _claim_next_queued_run(prisma_client, session_id)
if claimed is not None:
return NextRunResponse(
run_id=claimed.id,
prompt=claimed.prompt or {},
)
if asyncio.get_running_loop().time() >= deadline:
return Response(status_code=204)
await asyncio.sleep(NEXT_RUN_POLL_INTERVAL_SECONDS)
@router.post(
"/v2/sessions/{session_id}/runs/{run_id}/events:append",
response_class=ORJSONResponse,
tags=["agent-internal"],
)
async def daemon_append_event(
body: EventAppend,
session_id: str = Path(...),
run_id: str = Path(...),
daemon: Dict[str, Any] = Depends(daemon_token_auth),
):
"""Append an event to a run.
Seq is computed as ``MAX(seq) + 1`` on the server side. The
``(run_id, seq)`` unique constraint is the safety net — if two
daemon emits race, the loser gets an IntegrityError, retries, and
inserts at the next seq.
If ``event_type`` is one of ``run_finished | run_cancelled |
run_error``, also flip the run's ``status`` and ``terminated_at``
inside the same transaction-ish flow.
"""
prisma_client = await _get_prisma_client_or_503()
# Cross-tenant defense: token is good for ``session_id`` but the run
# itself must belong to that session.
run = await prisma_client.db.litellm_agentrun.find_unique(where={"id": run_id})
if run is None or run.session_id != session_id:
raise HTTPException(status_code=404, detail="Run not found")
# Compute next seq. Retry once on collision; the unique constraint
# makes this safe.
last_evt = await prisma_client.db.litellm_agentrunevent.find_first(
where={"run_id": run_id},
order={"seq": "desc"},
)
next_seq = (last_evt.seq + 1) if last_evt else 1
try:
evt = await prisma_client.db.litellm_agentrunevent.create(
data={
"run_id": run_id,
"seq": next_seq,
"event_type": body.event_type,
"payload": body.payload,
}
)
except Exception as exc:
# Most likely a (run_id, seq) unique-constraint collision. Reread
# MAX(seq) and try once more — that's the defensive retry the
# ticket calls out.
last_evt = await prisma_client.db.litellm_agentrunevent.find_first(
where={"run_id": run_id},
order={"seq": "desc"},
)
next_seq = (last_evt.seq + 1) if last_evt else 1
try:
evt = await prisma_client.db.litellm_agentrunevent.create(
data={
"run_id": run_id,
"seq": next_seq,
"event_type": body.event_type,
"payload": body.payload,
}
)
except Exception as exc2:
verbose_proxy_logger.exception(
"events:append double-collision run_id=%s: %s", run_id, exc2
)
raise HTTPException(status_code=409, detail="event_seq_collision") from exc
# Terminal-event handling: roll the run forward.
new_status = RUN_TERMINAL_EVENT_TYPES.get(body.event_type)
if new_status is not None:
now = _now()
await prisma_client.db.litellm_agentrun.update(
where={"id": run_id},
data={
"status": new_status,
"terminated_at": now,
"updated_at": now,
"result": (
body.payload.get("result")
if isinstance(body.payload, dict)
else None
),
},
)
return {
"run_id": run_id,
"seq": evt.seq,
"event_type": evt.event_type,
}

View file

@ -0,0 +1,144 @@
"""
Ownership / access-control helpers for the three-level agent hierarchy
(Agent -> Session -> Run).
Owner identity for an agent is the SHA-256 hash of the API key that created
it (``user_api_key_hash``). This matches the existing pattern used by
``litellm_managedobjecttable`` for containers — but we store the hash on
the row directly because there are no nested resource lookups here.
Sessions and runs inherit ownership from their parent agent — so if a
caller can read the parent, they can read the child. Cross-tenant isolation
is enforced uniformly by ``check_agent_ownership`` which is called from
every public endpoint before any read or write.
"""
from typing import Any, Optional
from fastapi import HTTPException
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy._types import hash_token as _hash_litellm_api_key
def caller_api_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str:
"""Return the SHA-256 hash of the caller's API key.
``user_api_key_dict.api_key`` may already be hashed (DB-issued virtual
keys store the hash), or it may be a master key in plain form. Either
way we hash whatever we get — hashing an already-hashed value gives a
deterministic value that can never collide with a real raw key.
"""
api_key = user_api_key_dict.api_key
if not api_key:
raise HTTPException(status_code=401, detail="Missing API key on auth context")
return _hash_litellm_api_key(api_key)
def is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
"""True iff the caller is a full proxy admin (read AND write).
``PROXY_ADMIN_VIEW_ONLY`` is intentionally NOT included here. For
write paths, this function being False means the view-only admin
falls back to the per-tenant ownership check, which they fail (no
matching ``user_api_key_hash``) and get a 404 — closing the
privilege escalation that previously let view-only admins create /
update / delete other tenants' agents, sessions, and runs.
For read paths, callers should use :func:`is_proxy_admin_read` so
view-only admins still get cross-tenant visibility.
"""
role = user_api_key_dict.user_role
if role is None:
return False
return role == LitellmUserRoles.PROXY_ADMIN
def is_proxy_admin_read(user_api_key_dict: UserAPIKeyAuth) -> bool:
"""True iff the caller is any flavor of proxy admin (incl. view-only).
Used only on read paths where the view-only admin is allowed to see
other tenants' resources.
"""
role = user_api_key_dict.user_role
if role is None:
return False
return role in {
LitellmUserRoles.PROXY_ADMIN,
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
}
def assert_caller_can_mutate(user_api_key_dict: UserAPIKeyAuth) -> None:
"""Reject write access for view-only admins.
Every state-mutating endpoint (POST/PUT/PATCH/DELETE) must call this
before any DB write. View-only admins are granted read access to all
tenants via :func:`is_proxy_admin_read` but MUST NOT be allowed to
mutate state — they otherwise inherit master-admin write authority
over every tenant's agents, sessions, and runs.
"""
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY:
raise HTTPException(
status_code=403,
detail="View-only admins cannot perform write operations",
)
def assert_caller_owns_agent(
user_api_key_dict: UserAPIKeyAuth,
agent_row: Any,
) -> None:
"""Raise 404 if the caller is not the agent's owner (and not an admin).
404 (not 403) on purpose — leaking existence of another tenant's
resource is a fingerprinting risk.
Both full and view-only admins pass this read-side check; write
endpoints must additionally call :func:`assert_caller_can_mutate` to
block view-only admins from mutating other tenants' rows.
"""
if agent_row is None:
raise HTTPException(status_code=404, detail="Agent not found")
if is_proxy_admin_read(user_api_key_dict):
return
expected_hash = caller_api_key_hash(user_api_key_dict)
if agent_row.user_api_key_hash != expected_hash:
raise HTTPException(status_code=404, detail="Agent not found")
def assert_caller_owns_session(
user_api_key_dict: UserAPIKeyAuth,
session_row: Any,
) -> None:
"""Raise 404 if the caller is not the session's owner (and not admin).
Sessions carry their own ``user_api_key_hash`` so we don't need to load
the parent agent to check ownership.
Both full and view-only admins pass this read-side check; write
endpoints must additionally call :func:`assert_caller_can_mutate` to
block view-only admins from mutating other tenants' rows.
"""
if session_row is None:
raise HTTPException(status_code=404, detail="Session not found")
if is_proxy_admin_read(user_api_key_dict):
return
expected_hash = caller_api_key_hash(user_api_key_dict)
if session_row.user_api_key_hash != expected_hash:
raise HTTPException(status_code=404, detail="Session not found")
def owner_filter_for_caller(
user_api_key_dict: UserAPIKeyAuth,
) -> Optional[dict]:
"""Build a Prisma ``where`` clause restricting list queries to the
caller's own rows.
Returns ``None`` for proxy admins (no filter) so callers can spread it
into an existing where dict only when needed. Both full and view-only
admins get the unfiltered read.
"""
if is_proxy_admin_read(user_api_key_dict):
return None
return {"user_api_key_hash": caller_api_key_hash(user_api_key_dict)}

View file

@ -0,0 +1,415 @@
"""
Run endpoints — POST/GET runs under a session, plus the SSE events stream
and explicit cancel.
Every endpoint is owner-scoped (caller must own the parent session). The
SSE stream is resumable — clients pass ``starting_seq=N`` to skip events
they've already seen, and the stream emits keep-alive comments so proxies
don't reap the connection.
"""
import asyncio
import json
from datetime import datetime, timezone
from typing import Any, Dict, List, Optional
from fastapi import APIRouter, Depends, HTTPException, Header, Query
from fastapi.responses import ORJSONResponse, StreamingResponse
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_session_endpoints.constants import (
EVENT_TYPE_RUN_CANCELLED,
RUN_ACTIVE_STATUSES,
RUN_STATUS_CANCELLED,
RUN_STATUS_QUEUED,
RUN_TERMINAL_STATUSES,
SESSION_TERMINAL_STATUSES,
SSE_KEEPALIVE_INTERVAL_SECONDS,
SSE_POLL_INTERVAL_SECONDS,
SSE_TERMINAL_QUIESCE_SECONDS,
)
from litellm.proxy.agent_session_endpoints.ids import new_run_id
from litellm.proxy.agent_session_endpoints.ownership import (
assert_caller_can_mutate,
assert_caller_owns_session,
)
from litellm.proxy.agent_session_endpoints.schemas import RunCreate, RunResponse
from litellm.proxy.agent_session_endpoints.serialization import run_row_to_response
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
router = APIRouter()
def _now() -> datetime:
return datetime.now(timezone.utc)
async def _get_prisma_client_or_503():
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=503, detail="Database unavailable")
return prisma_client
async def _load_session_or_404(prisma_client, session_id: str):
session = await prisma_client.db.litellm_agentsession.find_unique(
where={"id": session_id}
)
if session is None:
raise HTTPException(status_code=404, detail="Session not found")
return session
async def _find_idempotent_run(prisma_client, session_id: str, idempotency_key: str):
"""Return the existing run for ``(session_id, idempotency_key)`` if any."""
return await prisma_client.db.litellm_agentrun.find_first(
where={
"session_id": session_id,
"idempotency_key": idempotency_key,
}
)
async def _has_active_run(prisma_client, session_id: str) -> bool:
"""True iff session has any run in queued/running."""
existing = await prisma_client.db.litellm_agentrun.find_first(
where={
"session_id": session_id,
"status": {"in": list(RUN_ACTIVE_STATUSES)},
}
)
return existing is not None
async def _next_event_seq(prisma_client, run_id: str) -> int:
last = await prisma_client.db.litellm_agentrunevent.find_first(
where={"run_id": run_id},
order={"seq": "desc"},
)
if last is None:
return 1
return last.seq + 1
@router.post(
"/v2/sessions/{session_id}/runs",
response_class=ORJSONResponse,
response_model=RunResponse,
tags=["runs"],
dependencies=[Depends(user_api_key_auth)],
)
async def create_run(
session_id: str,
body: RunCreate,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
idempotency_key: Optional[str] = Header(default=None, alias="Idempotency-Key"),
):
"""Start a new run within a session.
Concurrency rules:
* 409 ``run_busy`` if another run is in queued/running.
* 409 ``session_not_accepting_runs`` if session is not ready/busy.
* Idempotency-Key returns the existing run on retry (same row).
"""
assert_caller_can_mutate(user_api_key_dict)
prisma_client = await _get_prisma_client_or_503()
session = await _load_session_or_404(prisma_client, session_id)
assert_caller_owns_session(user_api_key_dict, session)
if session.status in SESSION_TERMINAL_STATUSES:
raise HTTPException(
status_code=409,
detail=f"Session is {session.status}; cannot create runs",
)
# Idempotent retry shortcuts the busy check — if the same key already
# produced a run, return that run regardless of current state.
if idempotency_key:
existing = await _find_idempotent_run(
prisma_client, session_id, idempotency_key
)
if existing is not None:
return run_row_to_response(existing)
if await _has_active_run(prisma_client, session_id):
raise HTTPException(
status_code=409,
detail="run_busy: another run is queued/running for this session",
)
# Insert. The unique ``(session_id, idempotency_key)`` constraint is
# the actual safety net for racing duplicates — if two requests with
# the same key collide on insert, one gets the IntegrityError, retries
# the find_first, and returns the winner's row.
payload: Dict[str, Any] = {
"id": new_run_id(),
"session_id": session_id,
"status": RUN_STATUS_QUEUED,
"prompt": body.prompt,
"idempotency_key": idempotency_key,
"updated_at": _now(),
}
try:
row = await prisma_client.db.litellm_agentrun.create(data=payload)
except Exception as exc:
# Idempotent-collision recovery.
if idempotency_key:
existing = await _find_idempotent_run(
prisma_client, session_id, idempotency_key
)
if existing is not None:
return run_row_to_response(existing)
# Active-run race: another caller won the busy check between our
# check and our insert. Return the canonical 409.
active_other = await prisma_client.db.litellm_agentrun.find_first(
where={
"session_id": session_id,
"status": {"in": list(RUN_ACTIVE_STATUSES)},
}
)
if active_other is not None:
raise HTTPException(
status_code=409,
detail="run_busy: another run is queued/running for this session",
) from exc
raise
verbose_proxy_logger.info("run.create id=%s session_id=%s", row.id, session_id)
return run_row_to_response(row)
@router.get(
"/v2/sessions/{session_id}/runs/{run_id}",
response_class=ORJSONResponse,
response_model=RunResponse,
tags=["runs"],
dependencies=[Depends(user_api_key_auth)],
)
async def get_run(
session_id: str,
run_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
prisma_client = await _get_prisma_client_or_503()
session = await _load_session_or_404(prisma_client, session_id)
assert_caller_owns_session(user_api_key_dict, session)
run = await prisma_client.db.litellm_agentrun.find_unique(where={"id": run_id})
if run is None or run.session_id != session_id:
raise HTTPException(status_code=404, detail="Run not found")
return run_row_to_response(run)
@router.get(
"/v2/sessions/{session_id}/runs",
response_class=ORJSONResponse,
tags=["runs"],
dependencies=[Depends(user_api_key_auth)],
)
async def list_runs(
session_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
limit: int = Query(default=100, ge=1, le=500),
offset: int = Query(default=0, ge=0),
):
prisma_client = await _get_prisma_client_or_503()
session = await _load_session_or_404(prisma_client, session_id)
assert_caller_owns_session(user_api_key_dict, session)
rows = await prisma_client.db.litellm_agentrun.find_many(
where={"session_id": session_id},
order={"created_at": "desc"},
take=limit,
skip=offset,
)
return {"data": [run_row_to_response(r).model_dump() for r in rows]}
@router.post(
"/v2/sessions/{session_id}/runs/{run_id}/cancel",
response_class=ORJSONResponse,
response_model=RunResponse,
tags=["runs"],
dependencies=[Depends(user_api_key_auth)],
)
async def cancel_run(
session_id: str,
run_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""Mark a run cancelled and emit ``run_cancelled`` event.
Idempotent: cancelling an already-terminal run returns the row
unchanged (200, no extra event emitted).
"""
assert_caller_can_mutate(user_api_key_dict)
prisma_client = await _get_prisma_client_or_503()
session = await _load_session_or_404(prisma_client, session_id)
assert_caller_owns_session(user_api_key_dict, session)
run = await prisma_client.db.litellm_agentrun.find_unique(where={"id": run_id})
if run is None or run.session_id != session_id:
raise HTTPException(status_code=404, detail="Run not found")
if run.status in RUN_TERMINAL_STATUSES:
return run_row_to_response(run)
now = _now()
updated = await prisma_client.db.litellm_agentrun.update(
where={"id": run_id},
data={
"status": RUN_STATUS_CANCELLED,
"terminated_at": now,
"updated_at": now,
},
)
next_seq = await _next_event_seq(prisma_client, run_id)
try:
await prisma_client.db.litellm_agentrunevent.create(
data={
"run_id": run_id,
"seq": next_seq,
"event_type": EVENT_TYPE_RUN_CANCELLED,
"payload": {"reason": "user_cancel"},
}
)
except Exception as exc:
verbose_proxy_logger.warning(
"run.cancel: skipped run_cancelled emit run=%s seq=%s: %s",
run_id,
next_seq,
exc,
)
return run_row_to_response(updated)
# ---------------------------------------------------------------------------
# SSE events stream — resumable via ``?starting_seq=N``
# ---------------------------------------------------------------------------
def _sse_event_data(seq: int, event_type: str, payload: Dict[str, Any]) -> str:
"""Format a single SSE event frame.
The ``id:`` field carries the seq so clients can resume by using the
HTTP ``Last-Event-ID`` header (or ``starting_seq`` query param).
"""
body = json.dumps(
{"seq": seq, "event_type": event_type, "payload": payload},
default=str,
)
return f"id: {seq}\ndata: {body}\n\n"
async def _events_after(prisma_client, run_id: str, last_seen_seq: int) -> List[Any]:
return await prisma_client.db.litellm_agentrunevent.find_many(
where={"run_id": run_id, "seq": {"gt": last_seen_seq}},
order={"seq": "asc"},
)
async def _stream_run_events(
run_id: str,
starting_seq: int,
):
"""Async generator yielding SSE-formatted bytes.
Loop:
1. Fetch all events with seq > last_seen.
2. Yield each (and bump last_seen).
3. If the run is terminal AND no new events arrived this tick AND
the quiesce window has elapsed, close cleanly.
4. Sleep `SSE_POLL_INTERVAL_SECONDS`. Emit `:keepalive` every
`SSE_KEEPALIVE_INTERVAL_SECONDS`.
"""
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return
last_seen = max(starting_seq - 1, 0)
# ``get_running_loop`` (not ``get_event_loop``) is the correct API
# inside a running coroutine — see PEP 654 / Python 3.10 deprecation.
last_keepalive = asyncio.get_running_loop().time()
terminal_first_seen_at: Optional[float] = None
while True:
events = await _events_after(prisma_client, run_id, last_seen)
for evt in events:
yield _sse_event_data(evt.seq, evt.event_type, _safe_payload(evt.payload))
last_seen = evt.seq
run = await prisma_client.db.litellm_agentrun.find_unique(where={"id": run_id})
if run is None:
# Run vanished (cascade delete). Close.
return
now = asyncio.get_running_loop().time()
if run.status in RUN_TERMINAL_STATUSES:
if terminal_first_seen_at is None:
terminal_first_seen_at = now
elif (
now - terminal_first_seen_at >= SSE_TERMINAL_QUIESCE_SECONDS
and not events
):
# Run is finished AND we've drained events for at least
# the quiesce window AND nothing new came this tick.
return
if now - last_keepalive >= SSE_KEEPALIVE_INTERVAL_SECONDS:
yield ": keepalive\n\n"
last_keepalive = now
await asyncio.sleep(SSE_POLL_INTERVAL_SECONDS)
def _safe_payload(value: Any) -> Dict[str, Any]:
if isinstance(value, dict):
return value
return {"value": value}
@router.get(
"/v2/sessions/{session_id}/runs/{run_id}/events",
tags=["runs"],
dependencies=[Depends(user_api_key_auth)],
)
async def stream_run_events(
session_id: str,
run_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
starting_seq: int = Query(default=0, ge=0),
last_event_id: Optional[str] = Header(default=None, alias="Last-Event-ID"),
):
"""SSE stream of every event emitted on a run.
Resumable: pass ``?starting_seq=N`` (or the standard SSE
``Last-Event-ID`` header) to skip already-seen events.
"""
prisma_client = await _get_prisma_client_or_503()
session = await _load_session_or_404(prisma_client, session_id)
assert_caller_owns_session(user_api_key_dict, session)
run = await prisma_client.db.litellm_agentrun.find_unique(where={"id": run_id})
if run is None or run.session_id != session_id:
raise HTTPException(status_code=404, detail="Run not found")
# `Last-Event-ID` header beats explicit query param, matching the
# SSE-reconnect convention browsers use.
effective_starting_seq = starting_seq
if last_event_id:
try:
effective_starting_seq = max(int(last_event_id) + 1, starting_seq)
except ValueError:
pass # ignore malformed header, fall back to query param
return StreamingResponse(
_stream_run_events(run_id, effective_starting_seq),
media_type="text/event-stream",
headers={
"Cache-Control": "no-cache",
"X-Accel-Buffering": "no",
},
)

View file

@ -0,0 +1,181 @@
"""
Pydantic request/response models for the public agent_session_endpoints API.
These models are the canonical wire shape — Epic D (TS SDK) generates types
from this file via ``model_json_schema``. Keep field names and types stable
once shipped.
"""
from typing import Any, Dict, List, Literal, Optional
from pydantic import BaseModel, Field
# ---------------------------------------------------------------------------
# Repo / env_vars (shared building blocks)
# ---------------------------------------------------------------------------
class RepoSpec(BaseModel):
"""A git repository to clone into the session VM."""
url: str
startingRef: Optional[str] = Field(
default=None,
description="Branch, tag, or SHA to check out. Defaults to default branch.",
)
model_config = {"extra": "allow"}
# ---------------------------------------------------------------------------
# Agent
# ---------------------------------------------------------------------------
class AgentCreate(BaseModel):
name: str
model: str
system_prompt: Optional[str] = None
default_repos: Optional[List[RepoSpec]] = None
default_env_vars: Optional[Dict[str, str]] = None
tools_config: Optional[Dict[str, Any]] = None
metadata: Optional[Dict[str, Any]] = None
class AgentUpdate(BaseModel):
"""All fields optional — PATCH semantics. ``None`` means "leave alone"."""
name: Optional[str] = None
model: Optional[str] = None
system_prompt: Optional[str] = None
default_repos: Optional[List[RepoSpec]] = None
default_env_vars: Optional[Dict[str, str]] = None
tools_config: Optional[Dict[str, Any]] = None
metadata: Optional[Dict[str, Any]] = None
class AgentResponse(BaseModel):
id: str
name: str
model: str
system_prompt: Optional[str] = None
default_repos: Optional[List[Dict[str, Any]]] = None
default_env_vars: Optional[Dict[str, str]] = None
tools_config: Optional[Dict[str, Any]] = None
metadata: Optional[Dict[str, Any]] = None
team_id: Optional[str] = None
created_at: Optional[str] = None
updated_at: Optional[str] = None
# ---------------------------------------------------------------------------
# Session
# ---------------------------------------------------------------------------
class SessionCreate(BaseModel):
agent_id: str
repos: Optional[List[RepoSpec]] = None
env_vars: Optional[Dict[str, str]] = None
max_session_minutes: Optional[int] = Field(
default=None,
description="Override default 4h max. Capped at 24h.",
ge=1,
le=24 * 60,
)
class SessionResponse(BaseModel):
id: str
agent_id: str
status: str
vm_id: Optional[str] = None
vm_provider: Optional[str] = None
repos: List[Dict[str, Any]] = Field(default_factory=list)
expires_at: Optional[str] = None
last_heartbeat_at: Optional[str] = None
created_at: Optional[str] = None
updated_at: Optional[str] = None
terminated_at: Optional[str] = None
daemon_token: Optional[str] = Field(
default=None,
description="Returned ONLY on initial create. Subsequent reads return None.",
)
# ---------------------------------------------------------------------------
# Run
# ---------------------------------------------------------------------------
class RunCreate(BaseModel):
prompt: Dict[str, Any] = Field(
...,
description='Free-form prompt object. ``{"text": "hi"}`` is the minimum.',
)
class RunResponse(BaseModel):
id: str
session_id: str
status: str
prompt: Dict[str, Any]
parent_run_id: Optional[str] = None
result: Optional[str] = None
git_branches: Optional[List[Dict[str, Any]]] = None
created_at: Optional[str] = None
updated_at: Optional[str] = None
started_at: Optional[str] = None
terminated_at: Optional[str] = None
# ---------------------------------------------------------------------------
# Followup
# ---------------------------------------------------------------------------
class FollowupCreate(BaseModel):
prompt: Dict[str, Any]
class FollowupResponse(BaseModel):
run_id: str
action: Literal["queued", "new_run"]
# ---------------------------------------------------------------------------
# Internal endpoints (daemon callbacks)
# ---------------------------------------------------------------------------
class DaemonRegisterRequest(BaseModel):
vm_id: Optional[str] = None
daemon_version: Optional[str] = None
class DaemonHeartbeatRequest(BaseModel):
vm_id: Optional[str] = None
class EventAppend(BaseModel):
event_type: str
payload: Dict[str, Any]
class NextRunResponse(BaseModel):
run_id: str
prompt: Dict[str, Any]
class ConversationMessage(BaseModel):
run_id: str
seq: int
event_type: str
payload: Dict[str, Any]
created_at: Optional[str] = None
class ConversationResponse(BaseModel):
session_id: str
messages: List[ConversationMessage]

View file

@ -0,0 +1,107 @@
"""
Serialization helpers — Prisma row -> response Pydantic model.
Prisma returns ``datetime`` objects and Json columns as Python native types
already, so we just convert datetimes to ISO-8601 strings and pass JSON
columns through. Centralized here so every endpoint serializes the same
way (consistent timestamps in the SDK).
"""
from datetime import datetime
from typing import Any, Dict, List, Optional
from litellm.proxy.agent_session_endpoints.schemas import (
AgentResponse,
ConversationMessage,
RunResponse,
SessionResponse,
)
def _iso(value: Any) -> Optional[str]:
"""Format a Prisma datetime as ISO-8601 (UTC); pass through strings."""
if value is None:
return None
if isinstance(value, datetime):
return value.isoformat()
return str(value)
def _as_dict(value: Any) -> Optional[Dict[str, Any]]:
"""Coerce Prisma JSON column into a dict (or ``None`` if missing)."""
if value is None:
return None
if isinstance(value, dict):
return value
return None
def _as_list_of_dict(value: Any) -> List[Dict[str, Any]]:
if not value:
return []
if isinstance(value, list):
return [v for v in value if isinstance(v, dict)]
return []
def agent_row_to_response(row: Any) -> AgentResponse:
return AgentResponse(
id=row.id,
name=row.name,
model=row.model,
system_prompt=row.system_prompt,
default_repos=_as_list_of_dict(row.default_repos) or None,
default_env_vars=_as_dict(row.default_env_vars),
tools_config=_as_dict(row.tools_config),
metadata=_as_dict(row.metadata),
team_id=row.team_id,
created_at=_iso(row.created_at),
updated_at=_iso(row.updated_at),
)
def session_row_to_response(
row: Any,
daemon_token: Optional[str] = None,
) -> SessionResponse:
return SessionResponse(
id=row.id,
agent_id=row.agent_id,
status=row.status,
vm_id=row.vm_id,
vm_provider=row.vm_provider,
repos=_as_list_of_dict(row.repos),
expires_at=_iso(row.expires_at),
last_heartbeat_at=_iso(row.last_heartbeat_at),
created_at=_iso(row.created_at),
updated_at=_iso(row.updated_at),
terminated_at=_iso(row.terminated_at),
daemon_token=daemon_token,
)
def run_row_to_response(row: Any) -> RunResponse:
prompt = _as_dict(row.prompt) or {}
return RunResponse(
id=row.id,
session_id=row.session_id,
status=row.status,
prompt=prompt,
parent_run_id=row.parent_run_id,
result=row.result,
git_branches=_as_list_of_dict(row.git_branches) or None,
created_at=_iso(row.created_at),
updated_at=_iso(row.updated_at),
started_at=_iso(row.started_at),
terminated_at=_iso(row.terminated_at),
)
def event_row_to_message(row: Any) -> ConversationMessage:
return ConversationMessage(
run_id=row.run_id,
seq=row.seq,
event_type=row.event_type,
payload=_as_dict(row.payload) or {},
created_at=_iso(row.created_at),
)

View file

@ -0,0 +1,603 @@
"""
Session CRUD endpoints — POST/GET/DELETE /v2/sessions{,/<id>}, plus
``/followup`` (smart inject vs. new-run) and ``/conversation`` (stateless
snapshot of all events across runs).
Sessions are VM-backed. ``POST /v2/sessions``:
1. validates the parent agent exists + caller owns it
2. resolves repos/env_vars (overlay caller-provided over agent defaults)
3. inserts the session row in ``provisioning`` status
4. mints a daemon JWT, stores its hash for revocation
5. spawns the VM provider call as a background task (no client wait)
6. returns the session JSON immediately so the client can subscribe
The daemon JWT is returned exactly once on create — subsequent reads
return ``daemon_token=null``. Callers that lose it must DELETE and recreate.
"""
import asyncio
from datetime import datetime, timedelta, timezone
from typing import Any, Dict, List, Optional
from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request
from fastapi.responses import ORJSONResponse
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
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,
EVENT_TYPE_RUN_CANCELLED,
EVENT_TYPE_USER_MESSAGE,
RUN_ACTIVE_STATUSES,
RUN_STATUS_CANCELLED,
RUN_STATUS_QUEUED,
SESSION_STATUS_ERROR,
SESSION_STATUS_PROVISIONING,
SESSION_STATUS_READY,
SESSION_STATUS_TERMINATED,
SESSION_TERMINAL_STATUSES,
)
from litellm.proxy.agent_session_endpoints.ids import new_run_id, new_session_id
from litellm.proxy.agent_session_endpoints.ownership import (
assert_caller_can_mutate,
assert_caller_owns_agent,
assert_caller_owns_session,
caller_api_key_hash,
owner_filter_for_caller,
)
from litellm.proxy.agent_session_endpoints.schemas import (
FollowupCreate,
FollowupResponse,
SessionCreate,
SessionResponse,
)
from litellm.proxy.agent_session_endpoints.serialization import (
event_row_to_message,
session_row_to_response,
)
from litellm.proxy.agent_session_endpoints.vm_providers.registry import (
get_vm_provider,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
router = APIRouter()
DEFAULT_VM_PROVIDER_NAME = "noop"
def _now() -> datetime:
return datetime.now(timezone.utc)
async def _has_active_run(prisma_client, session_id: str) -> bool:
"""True iff ``session_id`` has any run in queued/running.
Mirrors the ``_has_active_run`` helper in ``run_endpoints.py``. We
duplicate it here (instead of importing) to keep ``session_endpoints``
free of any dependency on ``run_endpoints``.
"""
existing = await prisma_client.db.litellm_agentrun.find_first(
where={
"session_id": session_id,
"status": {"in": list(RUN_ACTIVE_STATUSES)},
}
)
return existing is not None
async def _get_prisma_client_or_503():
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=503, detail="Database unavailable")
return prisma_client
def _resolve_repos(
body_repos: Optional[List[Any]],
agent_default_repos: Any,
) -> List[Dict[str, Any]]:
"""Caller-provided repos override agent defaults entirely (whole-list
replace, not merge). If caller passes nothing, fall back to defaults.
"""
if body_repos is not None:
return [r.model_dump(exclude_none=True) for r in body_repos]
if isinstance(agent_default_repos, list):
return [r for r in agent_default_repos if isinstance(r, dict)]
return []
def _resolve_env_vars(
body_env_vars: Optional[Dict[str, str]],
agent_default_env_vars: Any,
) -> Optional[Dict[str, str]]:
"""Merge: agent defaults first, caller overrides on top, key by key.
This matches CRT-style env_var resolution and lets agent-level secrets
(e.g. NPM_TOKEN) sit alongside per-session overrides without re-typing.
"""
if body_env_vars is None and not isinstance(agent_default_env_vars, dict):
return None
merged: Dict[str, str] = {}
if isinstance(agent_default_env_vars, dict):
merged.update({str(k): str(v) for k, v in agent_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 _proxy_base_url() -> str:
"""Best-effort proxy base URL for the daemon to call back into.
Checks ``LITELLM_PROXY_BASE_URL`` env var first; falls back to localhost.
Production deploys MUST set the env var.
"""
import os
return os.environ.get("LITELLM_PROXY_BASE_URL", "http://localhost:4000")
async def _provision_in_background(
session_id: str,
agent_id: str,
repos: List[Dict[str, Any]],
env_vars: Optional[Dict[str, str]],
daemon_token: str,
provider_name: str,
) -> None:
"""Background task: call provider.provision and update the session row.
Failure paths flip status to ``error`` so the cleanup sweeper can chase
the row. We never raise — this runs detached and a raise would crash
the event loop's exception handler.
"""
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
verbose_proxy_logger.error(
"session.provision failed: prisma_client is None (session_id=%s)",
session_id,
)
return
try:
provider = get_vm_provider(provider_name)
result = await provider.provision(
session_id=session_id,
agent_id=agent_id,
repos=repos,
env_vars=env_vars,
daemon_token=daemon_token,
proxy_base_url=_proxy_base_url(),
)
await prisma_client.db.litellm_agentsession.update(
where={"id": session_id},
data={
"vm_id": result.vm_id,
"vm_provider": provider_name,
"updated_at": _now(),
},
)
verbose_proxy_logger.info(
"session.provision ok session_id=%s vm_id=%s", session_id, result.vm_id
)
except Exception as exc:
verbose_proxy_logger.exception(
"session.provision failed session_id=%s: %s", session_id, exc
)
try:
await prisma_client.db.litellm_agentsession.update(
where={"id": session_id},
data={
"status": SESSION_STATUS_ERROR,
"updated_at": _now(),
"terminated_at": _now(),
},
)
except Exception as inner:
verbose_proxy_logger.exception(
"session.provision: failed to mark session=%s as error: %s",
session_id,
inner,
)
async def _find_idempotent_session(user_api_key_hash: str, idempotency_key: str):
"""Return the existing session row for ``(user, idempotency_key)`` if any."""
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return None
return await prisma_client.db.litellm_agentsession.find_first(
where={
"user_api_key_hash": user_api_key_hash,
"idempotency_key": idempotency_key,
}
)
@router.post(
"/v2/sessions",
response_class=ORJSONResponse,
response_model=SessionResponse,
tags=["sessions"],
dependencies=[Depends(user_api_key_auth)],
)
async def create_session(
body: SessionCreate,
request: Request,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
idempotency_key: Optional[str] = Header(default=None, alias="Idempotency-Key"),
):
assert_caller_can_mutate(user_api_key_dict)
prisma_client = await _get_prisma_client_or_503()
# Idempotency: same (caller, key) returns same session — no daemon
# token re-mint, no second provision call.
user_hash = caller_api_key_hash(user_api_key_dict)
if idempotency_key:
existing = await _find_idempotent_session(user_hash, idempotency_key)
if existing is not None:
return session_row_to_response(existing, daemon_token=None)
# Validate parent agent + ownership.
agent_row = await prisma_client.db.litellm_agent.find_unique(
where={"id": body.agent_id}
)
assert_caller_owns_agent(user_api_key_dict, agent_row)
# Resolve repos/env_vars (overlay caller over agent defaults).
resolved_repos = _resolve_repos(body.repos, agent_row.default_repos)
resolved_env_vars = _resolve_env_vars(body.env_vars, agent_row.default_env_vars)
# Compute expiry — default 4h, capped 24h via Pydantic validator.
max_minutes = body.max_session_minutes or DEFAULT_MAX_SESSION_MINUTES
expires_at = _now() + timedelta(minutes=max_minutes)
session_id = new_session_id()
daemon_token = mint_daemon_token(
session_id=session_id,
agent_id=body.agent_id,
expires_at_epoch=int(expires_at.timestamp()),
)
payload = {
"id": session_id,
"agent_id": body.agent_id,
"user_api_key_hash": user_hash,
"team_id": user_api_key_dict.team_id,
"vm_provider": DEFAULT_VM_PROVIDER_NAME,
"repos": resolved_repos,
"env_vars": resolved_env_vars,
"status": SESSION_STATUS_PROVISIONING,
"daemon_token_hash": hash_daemon_token(daemon_token),
"expires_at": expires_at,
"idempotency_key": idempotency_key,
"updated_at": _now(),
}
row = await prisma_client.db.litellm_agentsession.create(data=payload)
# Fire-and-forget VM provisioning. The client polls / subscribes
# for status flips.
asyncio.create_task(
_provision_in_background(
session_id=session_id,
agent_id=body.agent_id,
repos=resolved_repos,
env_vars=resolved_env_vars,
daemon_token=daemon_token,
provider_name=DEFAULT_VM_PROVIDER_NAME,
)
)
verbose_proxy_logger.info(
"session.create id=%s agent_id=%s expires_at=%s",
session_id,
body.agent_id,
expires_at.isoformat(),
)
return session_row_to_response(row, daemon_token=daemon_token)
@router.get(
"/v2/sessions/{session_id}",
response_class=ORJSONResponse,
response_model=SessionResponse,
tags=["sessions"],
dependencies=[Depends(user_api_key_auth)],
)
async def get_session(
session_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
prisma_client = await _get_prisma_client_or_503()
row = await prisma_client.db.litellm_agentsession.find_unique(
where={"id": session_id}
)
assert_caller_owns_session(user_api_key_dict, row)
return session_row_to_response(row, daemon_token=None)
@router.get(
"/v2/sessions",
response_class=ORJSONResponse,
tags=["sessions"],
dependencies=[Depends(user_api_key_auth)],
)
async def list_sessions(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
agent_id: Optional[str] = Query(default=None),
limit: int = Query(default=100, ge=1, le=500),
offset: int = Query(default=0, ge=0),
):
prisma_client = await _get_prisma_client_or_503()
where: Dict[str, Any] = {}
owner = owner_filter_for_caller(user_api_key_dict)
if owner:
where.update(owner)
if agent_id:
where["agent_id"] = agent_id
rows = await prisma_client.db.litellm_agentsession.find_many(
where=where or None,
order={"created_at": "desc"},
take=limit,
skip=offset,
)
return {
"data": [
session_row_to_response(r, daemon_token=None).model_dump() for r in rows
]
}
@router.delete(
"/v2/sessions/{session_id}",
response_class=ORJSONResponse,
tags=["sessions"],
dependencies=[Depends(user_api_key_auth)],
)
async def delete_session(
session_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
assert_caller_can_mutate(user_api_key_dict)
prisma_client = await _get_prisma_client_or_503()
row = await prisma_client.db.litellm_agentsession.find_unique(
where={"id": session_id}
)
assert_caller_owns_session(user_api_key_dict, row)
if row.status not in SESSION_TERMINAL_STATUSES:
await _terminate_session_internal(session_id, reason="user_delete")
return {"id": session_id, "deleted": True}
async def _next_event_seq(prisma_client, run_id: str) -> int:
"""Return ``MAX(seq) + 1`` for a run, or 1 if no events yet.
The endpoint that calls this still relies on the DB unique constraint
``(run_id, seq)`` for correctness — this lookup is just a best-effort
starting point so retries collide and increment quickly.
"""
last = await prisma_client.db.litellm_agentrunevent.find_first(
where={"run_id": run_id},
order={"seq": "desc"},
)
if last is None:
return 1
return last.seq + 1
async def _terminate_session_internal(session_id: str, reason: str) -> None:
"""Internal helper: cancel runs, mark session terminated, fire provider.terminate.
Used by:
- DELETE /v2/sessions/{id}
- DELETE /v2/agents/{id} (cascade)
- cleanup sweeper
- daemon-dead detector
"""
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return
session = await prisma_client.db.litellm_agentsession.find_unique(
where={"id": session_id}
)
if session is None:
return
if session.status in SESSION_TERMINAL_STATUSES:
return
# 1. Cancel any non-terminal runs and emit run_cancelled events.
active_runs = await prisma_client.db.litellm_agentrun.find_many(
where={
"session_id": session_id,
"status": {"in": list(RUN_ACTIVE_STATUSES)},
}
)
now = _now()
for run in active_runs:
await prisma_client.db.litellm_agentrun.update(
where={"id": run.id},
data={
"status": RUN_STATUS_CANCELLED,
"terminated_at": now,
"updated_at": now,
},
)
next_seq = await _next_event_seq(prisma_client, run.id)
try:
await prisma_client.db.litellm_agentrunevent.create(
data={
"run_id": run.id,
"seq": next_seq,
"event_type": EVENT_TYPE_RUN_CANCELLED,
"payload": {"reason": reason},
}
)
except Exception as exc:
# Swallow seq-race; the daemon may have just emitted run_finished.
verbose_proxy_logger.warning(
"session.terminate: skipped run_cancelled emit run=%s seq=%s: %s",
run.id,
next_seq,
exc,
)
# 2. Mark session terminated.
await prisma_client.db.litellm_agentsession.update(
where={"id": session_id},
data={
"status": SESSION_STATUS_TERMINATED,
"terminated_at": now,
"updated_at": now,
},
)
# 3. Fire provider.terminate (best-effort; never blocks API caller).
try:
provider = get_vm_provider(session.vm_provider or DEFAULT_VM_PROVIDER_NAME)
await provider.terminate(
session_id=session_id, vm_id=session.vm_id, metadata=None
)
except Exception as exc:
verbose_proxy_logger.exception(
"session.terminate: provider.terminate failed session=%s: %s",
session_id,
exc,
)
@router.post(
"/v2/sessions/{session_id}/followup",
response_class=ORJSONResponse,
response_model=FollowupResponse,
tags=["sessions"],
dependencies=[Depends(user_api_key_auth)],
)
async def followup(
session_id: str,
body: FollowupCreate,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""Smart followup: if the latest run is active, append a user_message
event to it; if terminal or no runs, start a new run.
Matches Cursor's ``/followup`` semantics. Daemon picks up the
``user_message`` event via the events stream and weaves it into the
in-flight LLM turn.
"""
assert_caller_can_mutate(user_api_key_dict)
prisma_client = await _get_prisma_client_or_503()
session = await prisma_client.db.litellm_agentsession.find_unique(
where={"id": session_id}
)
assert_caller_owns_session(user_api_key_dict, session)
if session.status in SESSION_TERMINAL_STATUSES:
raise HTTPException(
status_code=409,
detail=f"Session is {session.status}; cannot followup",
)
latest_run = await prisma_client.db.litellm_agentrun.find_first(
where={"session_id": session_id},
order={"created_at": "desc"},
)
if latest_run is not None and latest_run.status in RUN_ACTIVE_STATUSES:
# Inject as a user_message event on the active run.
next_seq = await _next_event_seq(prisma_client, latest_run.id)
await prisma_client.db.litellm_agentrunevent.create(
data={
"run_id": latest_run.id,
"seq": next_seq,
"event_type": EVENT_TYPE_USER_MESSAGE,
"payload": body.prompt,
}
)
return FollowupResponse(run_id=latest_run.id, action="queued")
# Concurrency guard: matches POST /runs. Without this, two concurrent
# /followup requests on an idle session both pass the
# ``latest_run.status in RUN_ACTIVE_STATUSES`` check above and both
# fall through to ``create``, breaking the "one active run at a time"
# invariant. Re-check for any active run RIGHT before insert and
# return 409 ``run_busy`` instead, just like POST /runs does.
if await _has_active_run(prisma_client, session_id):
raise HTTPException(
status_code=409,
detail="run_busy: another run is queued/running for this session",
)
# Else start a fresh run.
try:
new_run = await prisma_client.db.litellm_agentrun.create(
data={
"id": new_run_id(),
"session_id": session_id,
"status": RUN_STATUS_QUEUED,
"prompt": body.prompt,
"updated_at": _now(),
}
)
except Exception as exc:
# Last line of defense: another /followup or POST /runs may have
# raced past the busy check and won the insert. Surface the same
# 409 so the client can retry deterministically.
active_other = await prisma_client.db.litellm_agentrun.find_first(
where={
"session_id": session_id,
"status": {"in": list(RUN_ACTIVE_STATUSES)},
}
)
if active_other is not None:
raise HTTPException(
status_code=409,
detail="run_busy: another run is queued/running for this session",
) from exc
raise
return FollowupResponse(run_id=new_run.id, action="new_run")
@router.get(
"/v2/sessions/{session_id}/conversation",
response_class=ORJSONResponse,
tags=["sessions"],
dependencies=[Depends(user_api_key_auth)],
)
async def get_conversation(
session_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""Stateless snapshot: every event across every run for this session,
in the order the daemon emitted them. Used by SDK consumers that don't
need the SSE stream — read-once-and-render."""
prisma_client = await _get_prisma_client_or_503()
session = await prisma_client.db.litellm_agentsession.find_unique(
where={"id": session_id}
)
assert_caller_owns_session(user_api_key_dict, session)
runs = await prisma_client.db.litellm_agentrun.find_many(
where={"session_id": session_id},
order={"created_at": "asc"},
)
if not runs:
return {"session_id": session_id, "messages": []}
run_ids = [r.id for r in runs]
events = await prisma_client.db.litellm_agentrunevent.find_many(
where={"run_id": {"in": run_ids}},
order=[{"created_at": "asc"}, {"seq": "asc"}],
)
return {
"session_id": session_id,
"messages": [event_row_to_message(e).model_dump() for e in events],
}

View file

@ -0,0 +1,130 @@
"""
Session and Run state machines.
This module is pure — no I/O. Endpoints call into these helpers to validate
state transitions before persisting them.
"""
from typing import Optional
from litellm.proxy.agent_session_endpoints.constants import (
RUN_ACTIVE_STATUSES,
RUN_STATUS_CANCELLED,
RUN_STATUS_ERROR,
RUN_STATUS_FINISHED,
RUN_STATUS_QUEUED,
RUN_STATUS_RUNNING,
RUN_TERMINAL_STATUSES,
SESSION_ACCEPTING_RUN_STATUSES,
SESSION_STATUS_BUSY,
SESSION_STATUS_ERROR,
SESSION_STATUS_PROVISIONING,
SESSION_STATUS_READY,
SESSION_STATUS_TERMINATED,
SESSION_TERMINAL_STATUSES,
)
# ---------------------------------------------------------------------------
# Session state machine
# ---------------------------------------------------------------------------
# Valid forward transitions. ``terminated`` is a sink — nothing transitions out.
_SESSION_TRANSITIONS = {
SESSION_STATUS_PROVISIONING: {
SESSION_STATUS_READY,
SESSION_STATUS_ERROR,
SESSION_STATUS_TERMINATED,
},
SESSION_STATUS_READY: {
SESSION_STATUS_BUSY,
SESSION_STATUS_ERROR,
SESSION_STATUS_TERMINATED,
},
SESSION_STATUS_BUSY: {
SESSION_STATUS_READY,
SESSION_STATUS_ERROR,
SESSION_STATUS_TERMINATED,
},
SESSION_STATUS_ERROR: {SESSION_STATUS_TERMINATED},
SESSION_STATUS_TERMINATED: set(),
}
def is_valid_session_transition(current: str, target: str) -> bool:
"""Return True if a session can move from ``current`` to ``target``."""
if current == target:
return True
return target in _SESSION_TRANSITIONS.get(current, set())
def session_can_accept_runs(status: str) -> bool:
"""Sessions in `ready` or `busy` accept new run inserts.
`busy` returns True here because the lock-then-check flow in
`POST /runs` lets `busy` through to a more specific 409 reason
(`run_busy`) inside the transaction. Sessions that are
`provisioning`, `error`, or `terminated` are rejected up-front.
"""
return status in SESSION_ACCEPTING_RUN_STATUSES
def session_is_terminal(status: str) -> bool:
return status in SESSION_TERMINAL_STATUSES
# ---------------------------------------------------------------------------
# Run state machine
# ---------------------------------------------------------------------------
_RUN_TRANSITIONS = {
RUN_STATUS_QUEUED: {
RUN_STATUS_RUNNING,
RUN_STATUS_CANCELLED,
RUN_STATUS_ERROR,
},
RUN_STATUS_RUNNING: {
RUN_STATUS_FINISHED,
RUN_STATUS_CANCELLED,
RUN_STATUS_ERROR,
},
RUN_STATUS_FINISHED: set(),
RUN_STATUS_CANCELLED: set(),
RUN_STATUS_ERROR: set(),
}
def is_valid_run_transition(current: str, target: str) -> bool:
if current == target:
return True
return target in _RUN_TRANSITIONS.get(current, set())
def run_is_active(status: str) -> bool:
return status in RUN_ACTIVE_STATUSES
def run_is_terminal(status: str) -> bool:
return status in RUN_TERMINAL_STATUSES
def derive_session_status_from_runs(
current_session_status: str,
has_active_run: bool,
) -> Optional[str]:
"""Return the new session status or None if no transition is needed.
This is the single source of truth for the busy<->ready oscillation
triggered by run start/finish events. It is intentionally a pure
function — callers persist the result inside their own transaction.
"""
if current_session_status in SESSION_TERMINAL_STATUSES:
return None
if current_session_status == SESSION_STATUS_PROVISIONING:
# Session must transition to ready (via daemon registration)
# before run-driven busy/ready toggling kicks in.
return None
target = SESSION_STATUS_BUSY if has_active_run else SESSION_STATUS_READY
if target == current_session_status:
return None
return target

View file

@ -0,0 +1,33 @@
"""Process-wide registry of VM providers, keyed by provider name."""
from typing import Dict
from litellm.proxy.agent_session_endpoints.vm_providers.base import AgentVMProvider
from litellm.proxy.agent_session_endpoints.vm_providers.noop import NoopVMProvider
_REGISTRY: Dict[str, AgentVMProvider] = {}
def register_vm_provider(provider: AgentVMProvider) -> None:
"""Register a provider by ``provider.name``. Last-write-wins; tests
use this to swap in a fresh ``NoopVMProvider`` between cases."""
_REGISTRY[provider.name] = provider
def get_vm_provider(name: str) -> AgentVMProvider:
"""Return the registered provider for ``name``.
Lazily instantiates a default ``NoopVMProvider`` on first access so
tests don't need a setup hook just to use the noop.
"""
if name not in _REGISTRY:
if name == "noop":
_REGISTRY[name] = NoopVMProvider()
else:
raise KeyError(f"No VM provider registered for '{name}'")
return _REGISTRY[name]
def reset_vm_provider_registry() -> None:
"""Test helper: drop all registered providers."""
_REGISTRY.clear()

View file

@ -939,9 +939,33 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915
## Initialize shared aiohttp session for connection reuse
shared_aiohttp_session = await _initialize_shared_aiohttp_session()
## /v2/agents+sessions cleanup sweeper (Epic A — Cursor SDK).
## Idempotent; safe even if the agent_session_endpoints module never
## sees traffic.
try:
from litellm.proxy.agent_session_endpoints.cleanup import (
start_cleanup_sweeper,
)
start_cleanup_sweeper()
except Exception as exc:
verbose_proxy_logger.warning(
"agent_session_endpoints cleanup sweeper failed to start: %s", exc
)
# End of startup event
yield
## Stop the agent session cleanup sweeper.
try:
from litellm.proxy.agent_session_endpoints.cleanup import (
stop_cleanup_sweeper,
)
stop_cleanup_sweeper()
except Exception:
pass
# Shutdown event - close shared aiohttp session
if shared_aiohttp_session is not None:
try:
@ -14900,6 +14924,45 @@ app.include_router(agent_pool_status_router)
# Eager: /models/{name}:method overlaps with the OpenAI /models endpoint.
app.include_router(google_router)
# /v2/agents, /v2/sessions — Cursor SDK agent runtime (Epic A).
# Mounted under /v2/ to avoid collision with the existing /v1/agents
# (A2A registry in litellm/proxy/agent_endpoints/).
#
# SECURITY: the daemon JWT secret is a separate credential from the proxy
# master key. If ``LITELLM_AGENT_JWT_SECRET`` is not set, refuse to mount
# these routers — silently signing daemon tokens with the master key (or
# any default) would conflate two distinct auth surfaces and let a
# captured daemon JWT mint master-key-authority API keys.
from litellm.proxy.agent_session_endpoints.auth import (
is_agent_jwt_secret_configured,
)
if is_agent_jwt_secret_configured():
from litellm.proxy.agent_session_endpoints import (
agent_router as agent_session_agent_router,
)
from litellm.proxy.agent_session_endpoints import (
internal_router as agent_session_internal_router,
)
from litellm.proxy.agent_session_endpoints import (
run_router as agent_session_run_router,
)
from litellm.proxy.agent_session_endpoints import (
session_router as agent_session_session_router,
)
app.include_router(agent_session_agent_router)
app.include_router(agent_session_session_router)
app.include_router(agent_session_run_router)
app.include_router(agent_session_internal_router)
else:
verbose_proxy_logger.error(
"agent_session_endpoints (/v2/agents, /v2/sessions) NOT mounted: "
"LITELLM_AGENT_JWT_SECRET is not set. Set this env var to a "
"dedicated random secret (distinct from LITELLM_MASTER_KEY) to "
"enable the Cursor SDK agent runtime."
)
attach_lazy_features(app)
app.add_middleware(
RequestSizeLimitMiddleware,

View file

@ -1477,3 +1477,128 @@ model LiteLLM_AgentWorkerPairingToken {
@@index([team_id])
}
// ===========================================================================
// Agent Sessions / Runs (Cursor SDK on LiteLLM)
//
// Three-level hierarchy:
// Agent — definition (model, system prompt, default repos, tools)
// Session — VM-backed conversation, owned by an Agent
// Run — single turn within a Session
// RunEvent — append-only event log per run (for resumable SSE)
// ===========================================================================
model LiteLLM_Agent {
id String @id // "agent_<uuid>"
name String
user_api_key_hash String
team_id String?
model String
system_prompt String?
default_repos Json?
default_env_vars Json?
tools_config Json?
metadata Json?
created_at DateTime @default(now())
updated_at DateTime @updatedAt
sessions LiteLLM_AgentSession[]
@@index([user_api_key_hash])
@@index([team_id])
}
model LiteLLM_AgentSession {
id String @id // "sess_<uuid>"
agent_id String
user_api_key_hash String
team_id String?
vm_id String?
vm_provider String?
repos Json
env_vars Json?
status String @default("provisioning")
daemon_token_hash String?
expires_at DateTime
last_heartbeat_at DateTime?
idempotency_key String?
created_at DateTime @default(now())
updated_at DateTime @updatedAt
terminated_at DateTime?
agent LiteLLM_Agent @relation(fields: [agent_id], references: [id], onDelete: Cascade)
runs LiteLLM_AgentRun[]
@@unique([user_api_key_hash, idempotency_key])
@@index([agent_id])
@@index([status, expires_at])
@@index([user_api_key_hash])
}
model LiteLLM_AgentRun {
id String @id // "run_<uuid>"
session_id String
parent_run_id String?
status String @default("queued")
prompt Json
result String?
git_branches Json?
idempotency_key String?
created_at DateTime @default(now())
updated_at DateTime @updatedAt
started_at DateTime?
terminated_at DateTime?
session LiteLLM_AgentSession @relation(fields: [session_id], references: [id], onDelete: Cascade)
events LiteLLM_AgentRunEvent[]
@@unique([session_id, idempotency_key])
@@index([session_id, status])
@@index([session_id, created_at])
}
model LiteLLM_AgentRunEvent {
id String @id @default(uuid())
run_id String
seq Int
event_type String
payload Json
created_at DateTime @default(now())
run LiteLLM_AgentRun @relation(fields: [run_id], references: [id], onDelete: Cascade)
@@unique([run_id, seq])
@@index([run_id, seq])
}
// ===========================================================================
// Warm pool VM tracking (LIT-2890 / Epic B2)
//
// Tracks the lifecycle of pre-provisioned EC2 (or other provider) VMs used
// for instant session attach. Each row maps to one underlying instance.
//
// State machine:
// provisioning → warm → hydrating → attached → terminating → terminated
//
// On session end, the VM is terminated (NOT recycled) — security boundary.
// The maintenance loop refills `warm` slots; rows in `terminated` are kept
// for audit until pruned.
// ===========================================================================
model LiteLLM_AgentVM {
id String @id // EC2 instance id (e.g. "i-0abcd...")
provider String // "ec2" | "noop" | "self_hosted"
region String?
state String // provisioning|warm|hydrating|attached|terminating|terminated
team_id String // owner team — pool is per-team
pool_id String // logical pool key (currently == team_id)
attached_session_id String? // FK to LiteLLM_AgentSession.id when state=attached
created_at DateTime @default(now())
warmed_at DateTime?
last_hydrate_at DateTime?
terminated_at DateTime?
metadata Json? // public_ip, private_ip, ssm_status, etc.
@@index([state, pool_id])
@@index([team_id, state])
@@index([attached_session_id])
}

View file

@ -1477,3 +1477,128 @@ model LiteLLM_AgentWorkerPairingToken {
@@index([team_id])
}
// ===========================================================================
// Agent Sessions / Runs (Cursor SDK on LiteLLM)
//
// Three-level hierarchy:
// Agent — definition (model, system prompt, default repos, tools)
// Session — VM-backed conversation, owned by an Agent
// Run — single turn within a Session
// RunEvent — append-only event log per run (for resumable SSE)
// ===========================================================================
model LiteLLM_Agent {
id String @id // "agent_<uuid>"
name String
user_api_key_hash String
team_id String?
model String
system_prompt String?
default_repos Json?
default_env_vars Json?
tools_config Json?
metadata Json?
created_at DateTime @default(now())
updated_at DateTime @updatedAt
sessions LiteLLM_AgentSession[]
@@index([user_api_key_hash])
@@index([team_id])
}
model LiteLLM_AgentSession {
id String @id // "sess_<uuid>"
agent_id String
user_api_key_hash String
team_id String?
vm_id String?
vm_provider String?
repos Json
env_vars Json?
status String @default("provisioning")
daemon_token_hash String?
expires_at DateTime
last_heartbeat_at DateTime?
idempotency_key String?
created_at DateTime @default(now())
updated_at DateTime @updatedAt
terminated_at DateTime?
agent LiteLLM_Agent @relation(fields: [agent_id], references: [id], onDelete: Cascade)
runs LiteLLM_AgentRun[]
@@unique([user_api_key_hash, idempotency_key])
@@index([agent_id])
@@index([status, expires_at])
@@index([user_api_key_hash])
}
model LiteLLM_AgentRun {
id String @id // "run_<uuid>"
session_id String
parent_run_id String?
status String @default("queued")
prompt Json
result String?
git_branches Json?
idempotency_key String?
created_at DateTime @default(now())
updated_at DateTime @updatedAt
started_at DateTime?
terminated_at DateTime?
session LiteLLM_AgentSession @relation(fields: [session_id], references: [id], onDelete: Cascade)
events LiteLLM_AgentRunEvent[]
@@unique([session_id, idempotency_key])
@@index([session_id, status])
@@index([session_id, created_at])
}
model LiteLLM_AgentRunEvent {
id String @id @default(uuid())
run_id String
seq Int
event_type String
payload Json
created_at DateTime @default(now())
run LiteLLM_AgentRun @relation(fields: [run_id], references: [id], onDelete: Cascade)
@@unique([run_id, seq])
@@index([run_id, seq])
}
// ===========================================================================
// Warm pool VM tracking (LIT-2890 / Epic B2)
//
// Tracks the lifecycle of pre-provisioned EC2 (or other provider) VMs used
// for instant session attach. Each row maps to one underlying instance.
//
// State machine:
// provisioning → warm → hydrating → attached → terminating → terminated
//
// On session end, the VM is terminated (NOT recycled) — security boundary.
// The maintenance loop refills `warm` slots; rows in `terminated` are kept
// for audit until pruned.
// ===========================================================================
model LiteLLM_AgentVM {
id String @id // EC2 instance id (e.g. "i-0abcd...")
provider String // "ec2" | "noop" | "self_hosted"
region String?
state String // provisioning|warm|hydrating|attached|terminating|terminated
team_id String // owner team — pool is per-team
pool_id String // logical pool key (currently == team_id)
attached_session_id String? // FK to LiteLLM_AgentSession.id when state=attached
created_at DateTime @default(now())
warmed_at DateTime?
last_hydrate_at DateTime?
terminated_at DateTime?
metadata Json? // public_ip, private_ip, ssm_status, etc.
@@index([state, pool_id])
@@index([team_id, state])
@@index([attached_session_id])
}

View file

@ -0,0 +1,308 @@
"""
Shared fixtures for `litellm/proxy/agent_session_endpoints/` tests.
Provides:
* ``fake_prisma_client`` — an in-memory stand-in for the proxy's Prisma
client. It implements only the methods our endpoints actually call —
no network, no schema. Tests assert against the data structures
directly.
* ``client`` — FastAPI TestClient with all four routers
mounted and ``user_api_key_auth`` overridden to a fixed proxy admin.
* ``other_tenant_client`` — TestClient where the auth dep returns a
different (non-admin) caller, used for cross-tenant isolation tests.
"""
import os
import pytest
# Set a JWT secret BEFORE any module under test is imported.
os.environ.setdefault("LITELLM_AGENT_JWT_SECRET", "test-agent-jwt-secret")
os.environ.setdefault("LITELLM_MASTER_KEY", "sk-1234")
from collections import defaultdict
from datetime import datetime, timezone
from typing import Any, Dict, List, Optional
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
# ---------------------------------------------------------------------------
# In-memory Prisma stand-in
# ---------------------------------------------------------------------------
class _Row:
"""Plain object mimicking Prisma row attribute access.
Missing attributes resolve to ``None`` to mirror Prisma's behavior
of returning a row with optional columns left null.
"""
def __init__(self, **fields: Any) -> None:
for k, v in fields.items():
setattr(self, k, v)
def __getattr__(self, name: str) -> Any:
# Only fires for attribute access that misses the instance dict;
# underscore-prefixed names (e.g. dunders) should error normally.
if name.startswith("_"):
raise AttributeError(name)
return None
def to_dict(self) -> Dict[str, Any]:
return self.__dict__.copy()
def _matches(row: _Row, where: Optional[Dict[str, Any]]) -> bool:
if not where:
return True
for key, expected in where.items():
actual = getattr(row, key, None)
if isinstance(expected, dict):
if "in" in expected:
if actual not in expected["in"]:
return False
elif "notIn" in expected:
if actual in expected["notIn"]:
return False
elif "lt" in expected:
if actual is None or not (actual < expected["lt"]):
return False
elif "gt" in expected:
if actual is None or not (actual > expected["gt"]):
return False
else:
# Unknown operator dict — fall back to equality on raw dict.
if actual != expected:
return False
else:
if actual != expected:
return False
return True
def _order_rows(rows: List[_Row], order: Any) -> List[_Row]:
if not order:
return rows
if isinstance(order, dict):
order = [order]
for o in reversed(order):
for k, direction in o.items():
rows.sort(
key=lambda r: getattr(r, k) or 0,
reverse=(direction == "desc"),
)
return rows
class _Table:
"""In-memory table with the few async methods our endpoints use."""
def __init__(self) -> None:
self.rows: List[_Row] = []
async def create(self, data: Dict[str, Any]) -> _Row:
# Defaults that real Prisma would apply
now = datetime.now(timezone.utc)
defaults = {"created_at": now, "updated_at": now}
merged = {**defaults, **data}
row = _Row(**merged)
self.rows.append(row)
return row
async def find_unique(self, where: Dict[str, Any]) -> Optional[_Row]:
for row in self.rows:
if _matches(row, where):
return row
return None
async def find_first(
self,
where: Optional[Dict[str, Any]] = None,
order: Any = None,
) -> Optional[_Row]:
results = [r for r in self.rows if _matches(r, where)]
results = _order_rows(results, order)
return results[0] if results else None
async def find_many(
self,
where: Optional[Dict[str, Any]] = None,
order: Any = None,
take: Optional[int] = None,
skip: Optional[int] = None,
) -> List[_Row]:
results = [r for r in self.rows if _matches(r, where)]
results = _order_rows(results, order)
if skip:
results = results[skip:]
if take:
results = results[:take]
return results
async def update(self, where: Dict[str, Any], data: Dict[str, Any]) -> _Row:
for row in self.rows:
if _matches(row, where):
for k, v in data.items():
setattr(row, k, v)
return row
raise RuntimeError(f"No row to update for {where}")
async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]):
count = 0
for row in self.rows:
if _matches(row, where):
for k, v in data.items():
setattr(row, k, v)
count += 1
# Mimic Prisma's BatchPayload-ish object.
return _Row(count=count)
async def delete(self, where: Dict[str, Any]) -> _Row:
for i, row in enumerate(self.rows):
if _matches(row, where):
return self.rows.pop(i)
raise RuntimeError(f"No row to delete for {where}")
class FakePrismaClient:
"""Drop-in for ``prisma_client`` in ``litellm.proxy.proxy_server``."""
def __init__(self) -> None:
self.db = _DB()
class _DB:
def __init__(self) -> None:
self.litellm_agent = _Table()
self.litellm_agentsession = _Table()
self.litellm_agentrun = _AgentRunTable()
self.litellm_agentrunevent = _AgentRunEventTable()
class _AgentRunTable(_Table):
"""Subclass that enforces the (session_id, idempotency_key) unique constraint."""
async def create(self, data: Dict[str, Any]) -> _Row:
sid = data.get("session_id")
idem = data.get("idempotency_key")
if idem is not None:
for row in self.rows:
if (
getattr(row, "session_id", None) == sid
and getattr(row, "idempotency_key", None) == idem
):
raise RuntimeError("idempotency_collision")
return await super().create(data)
class _AgentRunEventTable(_Table):
"""Enforces the (run_id, seq) unique constraint."""
async def create(self, data: Dict[str, Any]) -> _Row:
rid = data.get("run_id")
seq = data.get("seq")
for row in self.rows:
if getattr(row, "run_id", None) == rid and getattr(row, "seq", None) == seq:
raise RuntimeError("event_seq_collision")
return await super().create(data)
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture
def fake_prisma_client(monkeypatch):
"""Patch ``litellm.proxy.proxy_server.prisma_client`` for the duration of
the test with our in-memory stand-in.
"""
from litellm.proxy import proxy_server
fake = FakePrismaClient()
monkeypatch.setattr(proxy_server, "prisma_client", fake)
return fake
def _build_test_app(
role: LitellmUserRoles, api_key: str = "sk-test-caller"
) -> TestClient:
from litellm.proxy.agent_session_endpoints import (
agent_router,
internal_router,
run_router,
session_router,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
app = FastAPI()
app.include_router(agent_router)
app.include_router(session_router)
app.include_router(run_router)
app.include_router(internal_router)
def _fake_auth() -> UserAPIKeyAuth:
return UserAPIKeyAuth(
user_id="test-user",
user_role=role,
api_key=api_key,
team_id=None,
)
app.dependency_overrides[user_api_key_auth] = _fake_auth
return TestClient(app)
@pytest.fixture
def client(fake_prisma_client):
"""TestClient where caller is a non-admin (regular tenant)."""
return _build_test_app(LitellmUserRoles.INTERNAL_USER, api_key="sk-tenant-A")
@pytest.fixture
def admin_client(fake_prisma_client):
"""TestClient where caller is a proxy admin (sees everything)."""
return _build_test_app(LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin-key")
@pytest.fixture
def view_only_admin_client(fake_prisma_client):
"""TestClient where caller is a view-only proxy admin.
View-only admins are allowed to READ across tenants (so the support
UI can render any tenant's resources) but MUST NOT be allowed to
mutate state on any tenant's resources — see ``assert_caller_can_mutate``.
"""
return _build_test_app(
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, api_key="sk-view-only-admin"
)
@pytest.fixture
def other_tenant_client(fake_prisma_client):
"""TestClient where caller is a different tenant. Used for
cross-tenant isolation tests."""
return _build_test_app(LitellmUserRoles.INTERNAL_USER, api_key="sk-tenant-B")
@pytest.fixture
def noop_provider(monkeypatch):
"""Reset the VM provider registry to a fresh ``NoopVMProvider``."""
from litellm.proxy.agent_session_endpoints.vm_providers import (
NoopVMProvider,
register_vm_provider,
)
from litellm.proxy.agent_session_endpoints.vm_providers.registry import (
reset_vm_provider_registry,
)
reset_vm_provider_registry()
provider = NoopVMProvider()
register_vm_provider(provider)
yield provider
reset_vm_provider_registry()

View file

@ -0,0 +1,57 @@
"""
Validation #3 — one agent, multiple sessions; agent_id is stable across
sessions.
"""
def _create_agent(client) -> str:
res = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "test", "model": "gpt-4"},
)
assert res.status_code == 200, res.text
return res.json()["id"]
def _create_session(client, agent_id: str) -> str:
res = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": agent_id, "repos": []},
)
assert res.status_code == 200, res.text
return res.json()["id"]
def test_agent_reused_across_three_sessions(client, noop_provider):
agent_id = _create_agent(client)
session_ids = [_create_session(client, agent_id) for _ in range(3)]
# Sessions are distinct IDs but all reference the same agent_id.
assert len(set(session_ids)) == 3
for sid in session_ids:
get_res = client.get(
f"/v2/sessions/{sid}", headers={"Authorization": "Bearer k"}
)
assert get_res.status_code == 200
assert get_res.json()["agent_id"] == agent_id
def test_agent_id_format(client, noop_provider):
agent_id = _create_agent(client)
assert agent_id.startswith("agent_")
def test_sessions_under_agent_listed_by_filter(client, noop_provider):
agent_id = _create_agent(client)
sids = [_create_session(client, agent_id) for _ in range(2)]
res = client.get(
f"/v2/sessions?agent_id={agent_id}",
headers={"Authorization": "Bearer k"},
)
assert res.status_code == 200
listed_ids = {s["id"] for s in res.json()["data"]}
assert listed_ids == set(sids)

View file

@ -0,0 +1,87 @@
"""
Validation #11 — cascade delete + VM termination.
DELETE /v2/agents/{id} cancels active runs, terminates sessions, calls
provider.terminate.
"""
from litellm.proxy.agent_session_endpoints.constants import (
SESSION_STATUS_TERMINATED,
)
def test_delete_agent_terminates_sessions_and_calls_provider(
client, noop_provider, fake_prisma_client
):
agent = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "t", "model": "gpt-4"},
).json()
aid = agent["id"]
sess_a = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": aid, "repos": []},
).json()
sess_b = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": aid, "repos": []},
).json()
# Create a run on session A so cascade exercises run-cancel path too.
run = client.post(
f"/v2/sessions/{sess_a['id']}/runs",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "hi"}},
).json()
# DELETE the agent. Cascade: terminate both sessions, cancel the run.
res = client.delete(f"/v2/agents/{aid}", headers={"Authorization": "Bearer k"})
assert res.status_code == 200
assert res.json()["deleted"] is True
# Both sessions should be marked terminated in DB.
sessions_in_db = fake_prisma_client.db.litellm_agentsession.rows
for s in sessions_in_db:
assert s.status == SESSION_STATUS_TERMINATED
# provider.terminate was called once per session.
terminate_session_ids = {c["session_id"] for c in noop_provider.terminate_calls}
assert sess_a["id"] in terminate_session_ids
assert sess_b["id"] in terminate_session_ids
def test_delete_session_terminates_active_runs(
client, noop_provider, fake_prisma_client
):
agent = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "t", "model": "gpt-4"},
).json()
sess = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": agent["id"], "repos": []},
).json()
sid = sess["id"]
run = client.post(
f"/v2/sessions/{sid}/runs",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "x"}},
).json()
rid = run["id"]
res = client.delete(f"/v2/sessions/{sid}", headers={"Authorization": "Bearer k"})
assert res.status_code == 200
final_run = client.get(
f"/v2/sessions/{sid}/runs/{rid}", headers={"Authorization": "Bearer k"}
).json()
assert final_run["status"] == "cancelled"
# provider.terminate hit.
assert any(c["session_id"] == sid for c in noop_provider.terminate_calls)

View file

@ -0,0 +1,121 @@
"""
Validation #13 — cleanup sweeper.
Drives:
* Force ``expires_at`` past + status=ready → sweeper marks terminated, calls provider.terminate
* Force ``last_heartbeat_at`` > 90s ago → sweeper marks error
* Force ``status=running, updated_at`` > idle timeout → sweeper marks run error
"""
from datetime import datetime, timedelta, timezone
import pytest
from litellm.proxy.agent_session_endpoints.cleanup import run_cleanup_pass
from litellm.proxy.agent_session_endpoints.constants import (
DAEMON_HEARTBEAT_DEAD_AFTER_SECONDS,
RUN_IDLE_TIMEOUT_SECONDS,
RUN_STATUS_ERROR,
SESSION_STATUS_ERROR,
SESSION_STATUS_READY,
SESSION_STATUS_TERMINATED,
)
@pytest.mark.asyncio
async def test_sweeper_terminates_expired_sessions(
client, noop_provider, fake_prisma_client
):
a = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "t", "model": "gpt-4"},
).json()
sess = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": a["id"], "repos": []},
).json()
sid = sess["id"]
# Force expiry into the past.
row = fake_prisma_client.db.litellm_agentsession.rows[0]
row.expires_at = datetime.now(timezone.utc) - timedelta(minutes=1)
row.status = SESSION_STATUS_READY
summary = await run_cleanup_pass(fake_prisma_client)
assert summary["expired_sessions"] == 1
assert row.status == SESSION_STATUS_TERMINATED
# provider.terminate was called.
assert any(c["session_id"] == sid for c in noop_provider.terminate_calls)
@pytest.mark.asyncio
async def test_sweeper_marks_dead_daemon_sessions_error(
client, noop_provider, fake_prisma_client
):
a = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "t", "model": "gpt-4"},
).json()
sess = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": a["id"], "repos": []},
).json()
row = fake_prisma_client.db.litellm_agentsession.rows[0]
row.status = SESSION_STATUS_READY
row.last_heartbeat_at = datetime.now(timezone.utc) - timedelta(
seconds=DAEMON_HEARTBEAT_DEAD_AFTER_SECONDS + 30
)
summary = await run_cleanup_pass(fake_prisma_client)
assert summary["dead_daemon_sessions"] == 1
assert row.status == SESSION_STATUS_ERROR
@pytest.mark.asyncio
async def test_sweeper_marks_stuck_runs_error(
client, noop_provider, fake_prisma_client
):
a = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "t", "model": "gpt-4"},
).json()
sess = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": a["id"], "repos": []},
).json()
run = client.post(
f"/v2/sessions/{sess['id']}/runs",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "x"}},
).json()
run_row = fake_prisma_client.db.litellm_agentrun.rows[0]
run_row.status = "running"
run_row.updated_at = datetime.now(timezone.utc) - timedelta(
seconds=RUN_IDLE_TIMEOUT_SECONDS + 60
)
summary = await run_cleanup_pass(fake_prisma_client)
assert summary["stuck_runs"] == 1
assert run_row.status == RUN_STATUS_ERROR
@pytest.mark.asyncio
async def test_sweeper_pass_with_nothing_to_do(
client, noop_provider, fake_prisma_client
):
"""Empty-DB pass returns all-zero summary, no exception."""
summary = await run_cleanup_pass(fake_prisma_client)
assert summary == {
"expired_sessions": 0,
"dead_daemon_sessions": 0,
"stuck_runs": 0,
}

View file

@ -0,0 +1,92 @@
"""
Validation #7 — POST /runs concurrency.
Fire 5 runs at the same session simultaneously. Expect: exactly 1
succeeds (200), 4 get 409 run_busy.
The fake Prisma client is sequential (no real concurrency on a single
event loop), but we can still assert the busy-check semantics by
firing the requests as ``asyncio.gather`` of httpx tasks.
"""
import asyncio
import pytest
def _bootstrap_ready(client, noop_provider):
agent = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "t", "model": "gpt-4"},
).json()
sess = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": agent["id"], "repos": []},
).json()
daemon_token = sess["daemon_token"]
sid = sess["id"]
client.post(
f"/v2/sessions/{sid}/internal/register",
headers={"Authorization": f"Bearer {daemon_token}"},
json={"vm_id": "i-noop"},
)
return sid
def test_only_one_run_per_session(client, noop_provider):
sid = _bootstrap_ready(client, noop_provider)
statuses = []
for _ in range(5):
res = client.post(
f"/v2/sessions/{sid}/runs",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "go"}},
)
statuses.append(res.status_code)
# First succeeds, rest are 409 run_busy.
assert statuses.count(200) == 1
assert statuses.count(409) == 4
def test_idempotent_post_returns_same_run_id(client, noop_provider):
sid = _bootstrap_ready(client, noop_provider)
headers = {
"Authorization": "Bearer k",
"Idempotency-Key": "client-uuid-1",
}
a = client.post(
f"/v2/sessions/{sid}/runs",
headers=headers,
json={"prompt": {"text": "x"}},
)
b = client.post(
f"/v2/sessions/{sid}/runs",
headers=headers,
json={"prompt": {"text": "x"}},
)
assert a.status_code == 200
assert b.status_code == 200
assert a.json()["id"] == b.json()["id"]
def test_distinct_idempotency_keys_get_busy(client, noop_provider):
sid = _bootstrap_ready(client, noop_provider)
a = client.post(
f"/v2/sessions/{sid}/runs",
headers={"Authorization": "Bearer k", "Idempotency-Key": "k1"},
json={"prompt": {"text": "x"}},
)
b = client.post(
f"/v2/sessions/{sid}/runs",
headers={"Authorization": "Bearer k", "Idempotency-Key": "k2"},
json={"prompt": {"text": "y"}},
)
assert a.status_code == 200
assert b.status_code == 409

View file

@ -0,0 +1,130 @@
"""
Validation #16 — `/followup` does not bypass the run-busy guard.
A prior version of ``followup`` skipped the ``_has_active_run`` check
when ``latest_run`` was terminal-or-absent and went straight to
``litellm_agentrun.create``. Two concurrent ``POST /followup`` requests
on an idle session both passed the ``latest_run.status in
RUN_ACTIVE_STATUSES`` check and both inserted runs, breaking the
"one active run per session" invariant that ``POST /runs`` enforces
via 409 ``run_busy``.
This test reproduces the race deterministically by setting up a session
that already has an active run, then sending /followup with a fresh
request — followup must return 409 ``run_busy`` rather than enqueue a
duplicate run.
"""
def _bootstrap_session(client):
agent = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "concurrent", "model": "gpt-4"},
).json()
sess = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": agent["id"], "repos": []},
).json()
return sess["id"]
def test_followup_returns_409_when_active_run_is_not_the_latest(
client, fake_prisma_client, noop_provider
):
"""Reproducer for the original bug.
Scenario: a session has both
* an OLDER run still in ``running`` status, and
* a NEWER run already in a terminal status (``finished``).
The buggy code path used ``latest_run`` (newest by ``created_at``) to
decide whether to fall through to the "create new run" branch. Since
the newest run is terminal, the buggy version skipped ``_has_active_run``
and inserted a duplicate run — breaking the one-active-run invariant.
Two concurrent /followup calls trigger this same race in production
even when there's a single run: both observe the same terminal-or-absent
``latest_run`` and both reach the create branch. We can't easily race
asyncio in a unit test, so we model the equivalent state directly: an
older active run + newer terminal run. With the fix in place, the
busy check catches the older active run and returns 409.
"""
from datetime import datetime, timedelta, timezone
sid = _bootstrap_session(client)
earlier = datetime.now(timezone.utc) - timedelta(minutes=10)
later = datetime.now(timezone.utc)
# Inject directly into the fake prisma rows so we control created_at
# exactly. Bypassing ``.create()`` keeps the ``_now()``-based default
# from collapsing the timestamps.
from tests.test_litellm.proxy.agent_session_endpoints.conftest import _Row
fake_prisma_client.db.litellm_agentrun.rows.append(
_Row(
id="run_old_active",
session_id=sid,
status="running",
prompt={"text": "old"},
created_at=earlier,
updated_at=earlier,
)
)
fake_prisma_client.db.litellm_agentrun.rows.append(
_Row(
id="run_newer_terminal",
session_id=sid,
status="finished",
prompt={"text": "newer"},
created_at=later,
updated_at=later,
)
)
res = client.post(
f"/v2/sessions/{sid}/followup",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "should be blocked"}},
)
assert res.status_code == 409
assert "run_busy" in res.json()["detail"]
def test_followup_normal_path_creates_run_when_no_active_runs(client, noop_provider):
"""The new busy check must not break the happy path: when there is
no active run, /followup still creates a fresh queued run."""
sid = _bootstrap_session(client)
res = client.post(
f"/v2/sessions/{sid}/followup",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "first"}},
)
assert res.status_code == 200
body = res.json()
assert body["action"] == "new_run"
assert body["run_id"]
def test_followup_appends_event_when_run_is_active(client, noop_provider):
"""Other happy path: latest run IS active — /followup injects a
user_message event rather than a new run. Unchanged by the new
busy guard (we only fall through to the create branch when the
latest run is terminal/absent)."""
sid = _bootstrap_session(client)
# First followup creates a queued run.
first = client.post(
f"/v2/sessions/{sid}/followup",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "first"}},
).json()
# Second followup should append an event onto the now-active run.
second = client.post(
f"/v2/sessions/{sid}/followup",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "appended"}},
).json()
assert second["action"] == "queued"
assert second["run_id"] == first["run_id"]

View file

@ -0,0 +1,109 @@
"""
Validation #6 — followup smart behavior.
Case 1: latest run is active -> followup adds user_message event to it.
Case 2: latest run is terminal (or no runs exist) -> followup creates NEW run.
"""
from litellm.proxy.agent_session_endpoints.constants import (
EVENT_TYPE_RUN_FINISHED,
EVENT_TYPE_USER_MESSAGE,
)
def _bootstrap_ready(client, noop_provider):
agent = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "t", "model": "gpt-4"},
).json()
sess = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": agent["id"], "repos": []},
).json()
daemon_token = sess["daemon_token"]
sid = sess["id"]
client.post(
f"/v2/sessions/{sid}/internal/register",
headers={"Authorization": f"Bearer {daemon_token}"},
json={"vm_id": "i-noop"},
)
return sid, daemon_token
def test_followup_no_runs_creates_new_run(client, noop_provider):
sid, _ = _bootstrap_ready(client, noop_provider)
res = client.post(
f"/v2/sessions/{sid}/followup",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "first message"}},
)
assert res.status_code == 200
body = res.json()
assert body["action"] == "new_run"
assert body["run_id"]
def test_followup_with_active_run_appends_user_message(
client, noop_provider, fake_prisma_client
):
sid, daemon_token = _bootstrap_ready(client, noop_provider)
run = client.post(
f"/v2/sessions/{sid}/runs",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "hi"}},
).json()
rid = run["id"]
res = client.post(
f"/v2/sessions/{sid}/followup",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "by the way..."}},
)
assert res.status_code == 200, res.text
body = res.json()
assert body["action"] == "queued"
assert body["run_id"] == rid
# The user_message event should be appended to the active run.
events = fake_prisma_client.db.litellm_agentrunevent.rows
user_messages = [
e for e in events if e.event_type == EVENT_TYPE_USER_MESSAGE and e.run_id == rid
]
assert len(user_messages) == 1
assert user_messages[0].payload == {"text": "by the way..."}
def test_followup_after_run_finishes_creates_new_run(
client, noop_provider, fake_prisma_client
):
sid, daemon_token = _bootstrap_ready(client, noop_provider)
run = client.post(
f"/v2/sessions/{sid}/runs",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "first"}},
).json()
rid = run["id"]
# Finish the run so it goes terminal.
client.get(
f"/v2/sessions/{sid}/runs/next/internal/poll",
headers={"Authorization": f"Bearer {daemon_token}"},
)
client.post(
f"/v2/sessions/{sid}/runs/{rid}/events:append",
headers={"Authorization": f"Bearer {daemon_token}"},
json={"event_type": EVENT_TYPE_RUN_FINISHED, "payload": {"result": "ok"}},
)
res = client.post(
f"/v2/sessions/{sid}/followup",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "second turn"}},
)
assert res.status_code == 200
body = res.json()
assert body["action"] == "new_run"
assert body["run_id"] != rid

View file

@ -0,0 +1,65 @@
"""
Validation #9 — idempotency.
POST /sessions twice with same Idempotency-Key -> same session_id.
POST /runs twice with same Idempotency-Key -> same run_id, count == 1.
"""
def _create_agent(client) -> str:
return client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "t", "model": "gpt-4"},
).json()["id"]
def test_session_idempotency(client, noop_provider, fake_prisma_client):
agent_id = _create_agent(client)
headers = {"Authorization": "Bearer k", "Idempotency-Key": "uuid-A"}
a = client.post(
"/v2/sessions",
headers=headers,
json={"agent_id": agent_id, "repos": []},
)
b = client.post(
"/v2/sessions",
headers=headers,
json={"agent_id": agent_id, "repos": []},
)
assert a.status_code == 200 and b.status_code == 200
assert a.json()["id"] == b.json()["id"]
# Only one row in the table.
assert len(fake_prisma_client.db.litellm_agentsession.rows) == 1
def test_run_idempotency(client, noop_provider, fake_prisma_client):
agent_id = _create_agent(client)
sess = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": agent_id, "repos": []},
).json()
sid = sess["id"]
daemon_token = sess["daemon_token"]
client.post(
f"/v2/sessions/{sid}/internal/register",
headers={"Authorization": f"Bearer {daemon_token}"},
json={"vm_id": "i-noop"},
)
headers = {"Authorization": "Bearer k", "Idempotency-Key": "run-key-1"}
a = client.post(
f"/v2/sessions/{sid}/runs",
headers=headers,
json={"prompt": {"text": "x"}},
)
b = client.post(
f"/v2/sessions/{sid}/runs",
headers=headers,
json={"prompt": {"text": "x"}},
)
assert a.status_code == 200 and b.status_code == 200
assert a.json()["id"] == b.json()["id"]
assert len(fake_prisma_client.db.litellm_agentrun.rows) == 1

View file

@ -0,0 +1,97 @@
"""
Validation #15 — daemon JWT secret is REQUIRED, no master-key fallback.
The daemon JWT secret (``LITELLM_AGENT_JWT_SECRET``) MUST be a separate
credential from the proxy master key. A prior version of
``auth._get_signing_secret`` fell back to ``LITELLM_MASTER_KEY`` when
the dedicated env var was unset — which conflated two distinct auth
surfaces and meant a captured daemon JWT could be used to mint regular
API keys with master-key authority.
This test file enforces:
1. ``_get_signing_secret`` raises when ``LITELLM_AGENT_JWT_SECRET`` is
unset, even when ``LITELLM_MASTER_KEY`` IS set (no silent fallback).
2. ``is_agent_jwt_secret_configured()`` returns False in that case so
``proxy_server.py`` knows to skip mounting the routers.
3. ``mint_daemon_token`` and ``decode_daemon_token`` both surface the
same error — every caller is gated by the env var.
"""
import pytest
from litellm.proxy.agent_session_endpoints.auth import (
AgentJWTSecretNotConfiguredError,
_get_signing_secret,
decode_daemon_token,
is_agent_jwt_secret_configured,
mint_daemon_token,
)
def test_signing_secret_raises_without_dedicated_env_var(monkeypatch):
"""The dedicated env var is the ONLY source of the secret.
Setting LITELLM_MASTER_KEY (and clearing LITELLM_AGENT_JWT_SECRET)
must NOT silently provide a fallback.
"""
monkeypatch.delenv("LITELLM_AGENT_JWT_SECRET", raising=False)
monkeypatch.setenv("LITELLM_MASTER_KEY", "sk-master-key-which-must-not-be-used")
with pytest.raises(AgentJWTSecretNotConfiguredError):
_get_signing_secret()
def test_is_agent_jwt_secret_configured_reports_unset(monkeypatch):
"""Used at startup to decide whether to mount the routers.
Returns False when the env var is unset, even if the master key is
set — preventing the proxy from silently exposing an
auth-conflated /v2/agents surface.
"""
monkeypatch.delenv("LITELLM_AGENT_JWT_SECRET", raising=False)
monkeypatch.setenv("LITELLM_MASTER_KEY", "sk-1234")
assert is_agent_jwt_secret_configured() is False
def test_is_agent_jwt_secret_configured_reports_set(monkeypatch):
monkeypatch.setenv("LITELLM_AGENT_JWT_SECRET", "dedicated-jwt-secret-value")
assert is_agent_jwt_secret_configured() is True
def test_is_agent_jwt_secret_configured_rejects_empty_string(monkeypatch):
"""Empty string should be treated as unset — same as ``os.environ.get`` semantics."""
monkeypatch.setenv("LITELLM_AGENT_JWT_SECRET", "")
assert is_agent_jwt_secret_configured() is False
def test_mint_token_raises_without_dedicated_env_var(monkeypatch):
monkeypatch.delenv("LITELLM_AGENT_JWT_SECRET", raising=False)
monkeypatch.setenv("LITELLM_MASTER_KEY", "sk-1234")
with pytest.raises(AgentJWTSecretNotConfiguredError):
mint_daemon_token(
session_id="sess_abc",
agent_id="agt_abc",
expires_at_epoch=10**10,
)
def test_decode_token_raises_without_dedicated_env_var(monkeypatch):
"""Even decoding (which the daemon-auth dependency calls on every
request) must refuse to operate without the dedicated secret."""
monkeypatch.delenv("LITELLM_AGENT_JWT_SECRET", raising=False)
monkeypatch.setenv("LITELLM_MASTER_KEY", "sk-1234")
with pytest.raises(AgentJWTSecretNotConfiguredError):
decode_daemon_token("any-token-bytes")
def test_proxy_server_skips_mounting_when_secret_missing(monkeypatch):
"""End-to-end: simulate the startup check ``proxy_server.py`` runs.
When the env var is unset, ``is_agent_jwt_secret_configured`` returns
False — and that is the function ``proxy_server.py`` inspects to
decide whether to mount the routers.
"""
monkeypatch.delenv("LITELLM_AGENT_JWT_SECRET", raising=False)
# No master key fallback even when LITELLM_MASTER_KEY is set.
monkeypatch.setenv("LITELLM_MASTER_KEY", "sk-1234")
assert is_agent_jwt_secret_configured() is False

View file

@ -0,0 +1,93 @@
"""
Validation #10 — cross-tenant isolation at all 3 levels.
Tenant A creates an agent + session + run. Tenant B (different api_key)
must get 404 on every read/write of those resources.
"""
def test_cross_tenant_isolation(client, other_tenant_client, noop_provider):
# Tenant A creates an agent + session + run.
agent = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "secret", "model": "gpt-4"},
).json()
aid = agent["id"]
sess = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": aid, "repos": []},
).json()
sid = sess["id"]
run = client.post(
f"/v2/sessions/{sid}/runs",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "x"}},
).json()
rid = run["id"]
# Tenant B sees nothing.
assert (
other_tenant_client.get(
f"/v2/agents/{aid}", headers={"Authorization": "Bearer other"}
).status_code
== 404
)
assert (
other_tenant_client.get(
f"/v2/sessions/{sid}", headers={"Authorization": "Bearer other"}
).status_code
== 404
)
assert (
other_tenant_client.get(
f"/v2/sessions/{sid}/runs/{rid}",
headers={"Authorization": "Bearer other"},
).status_code
== 404
)
# Tenant B's list endpoints return only their own (none).
assert (
other_tenant_client.get(
"/v2/agents", headers={"Authorization": "Bearer other"}
).json()["data"]
== []
)
assert (
other_tenant_client.get(
"/v2/sessions", headers={"Authorization": "Bearer other"}
).json()["data"]
== []
)
# Tenant B can't delete tenant A's agent/session.
assert (
other_tenant_client.delete(
f"/v2/agents/{aid}", headers={"Authorization": "Bearer other"}
).status_code
== 404
)
assert (
other_tenant_client.delete(
f"/v2/sessions/{sid}", headers={"Authorization": "Bearer other"}
).status_code
== 404
)
def test_admin_sees_all_tenants(client, admin_client, noop_provider):
"""Proxy admin can read across tenants — admins exist for support
and ops and bypassing the owner filter is intentional."""
agent = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "tenant-A", "model": "gpt-4"},
).json()
# Admin can fetch tenant A's agent.
res = admin_client.get(
f"/v2/agents/{agent['id']}", headers={"Authorization": "Bearer admin"}
)
assert res.status_code == 200

View file

@ -0,0 +1,180 @@
"""
Validation #5 — run state machine.
queued -> running -> finished/cancelled/error.
Each transition emits a status event.
"""
from litellm.proxy.agent_session_endpoints.constants import (
EVENT_TYPE_RUN_CANCELLED,
EVENT_TYPE_RUN_ERROR,
EVENT_TYPE_RUN_FINISHED,
RUN_STATUS_CANCELLED,
RUN_STATUS_ERROR,
RUN_STATUS_FINISHED,
RUN_STATUS_QUEUED,
RUN_STATUS_RUNNING,
)
from litellm.proxy.agent_session_endpoints.state_machine import (
is_valid_run_transition,
run_is_active,
run_is_terminal,
)
def test_run_state_machine_pure():
assert is_valid_run_transition("queued", "running")
assert is_valid_run_transition("queued", "cancelled")
assert is_valid_run_transition("queued", "error")
assert is_valid_run_transition("running", "finished")
assert is_valid_run_transition("running", "cancelled")
assert is_valid_run_transition("running", "error")
# No transitions out of terminal
assert not is_valid_run_transition("finished", "running")
assert not is_valid_run_transition("cancelled", "running")
assert not is_valid_run_transition("error", "finished")
def test_run_helpers():
assert run_is_active("queued")
assert run_is_active("running")
assert not run_is_active("finished")
assert run_is_terminal("finished")
assert run_is_terminal("cancelled")
assert run_is_terminal("error")
assert not run_is_terminal("queued")
def _bootstrap(client, noop_provider):
agent = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "t", "model": "gpt-4"},
).json()
sess = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": agent["id"], "repos": []},
).json()
return agent, sess
def test_run_starts_queued(client, noop_provider):
_, sess = _bootstrap(client, noop_provider)
daemon_token = sess["daemon_token"]
sid = sess["id"]
client.post(
f"/v2/sessions/{sid}/internal/register",
headers={"Authorization": f"Bearer {daemon_token}"},
json={"vm_id": "i-noop"},
)
run = client.post(
f"/v2/sessions/{sid}/runs",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "hi"}},
).json()
assert run["status"] == RUN_STATUS_QUEUED
def test_run_finishes_via_events_append(client, noop_provider):
_, sess = _bootstrap(client, noop_provider)
daemon_token = sess["daemon_token"]
sid = sess["id"]
client.post(
f"/v2/sessions/{sid}/internal/register",
headers={"Authorization": f"Bearer {daemon_token}"},
json={"vm_id": "i-noop"},
)
run = client.post(
f"/v2/sessions/{sid}/runs",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "hi"}},
).json()
rid = run["id"]
# Daemon claims the run via long-poll (turns it running).
poll = client.get(
f"/v2/sessions/{sid}/runs/next/internal/poll",
headers={"Authorization": f"Bearer {daemon_token}"},
)
assert poll.status_code == 200
assert poll.json()["run_id"] == rid
after_poll = client.get(
f"/v2/sessions/{sid}/runs/{rid}", headers={"Authorization": "Bearer k"}
).json()
assert after_poll["status"] == RUN_STATUS_RUNNING
# Daemon emits run_finished — flips to finished.
append = client.post(
f"/v2/sessions/{sid}/runs/{rid}/events:append",
headers={"Authorization": f"Bearer {daemon_token}"},
json={"event_type": EVENT_TYPE_RUN_FINISHED, "payload": {"result": "done"}},
)
assert append.status_code == 200
final = client.get(
f"/v2/sessions/{sid}/runs/{rid}", headers={"Authorization": "Bearer k"}
).json()
assert final["status"] == RUN_STATUS_FINISHED
assert final["result"] == "done"
def test_run_cancel_endpoint(client, noop_provider):
_, sess = _bootstrap(client, noop_provider)
sid = sess["id"]
run = client.post(
f"/v2/sessions/{sid}/runs",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "hi"}},
).json()
rid = run["id"]
cancel = client.post(
f"/v2/sessions/{sid}/runs/{rid}/cancel",
headers={"Authorization": "Bearer k"},
)
assert cancel.status_code == 200
assert cancel.json()["status"] == RUN_STATUS_CANCELLED
# Idempotent — cancelling again returns same status.
cancel2 = client.post(
f"/v2/sessions/{sid}/runs/{rid}/cancel",
headers={"Authorization": "Bearer k"},
)
assert cancel2.status_code == 200
assert cancel2.json()["status"] == RUN_STATUS_CANCELLED
def test_terminal_event_via_append_emits_status_change(client, noop_provider):
"""run_error event flips run status."""
_, sess = _bootstrap(client, noop_provider)
daemon_token = sess["daemon_token"]
sid = sess["id"]
client.post(
f"/v2/sessions/{sid}/internal/register",
headers={"Authorization": f"Bearer {daemon_token}"},
json={"vm_id": "i-noop"},
)
run = client.post(
f"/v2/sessions/{sid}/runs",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "boom"}},
).json()
rid = run["id"]
client.get(
f"/v2/sessions/{sid}/runs/next/internal/poll",
headers={"Authorization": f"Bearer {daemon_token}"},
)
client.post(
f"/v2/sessions/{sid}/runs/{rid}/events:append",
headers={"Authorization": f"Bearer {daemon_token}"},
json={"event_type": EVENT_TYPE_RUN_ERROR, "payload": {"reason": "boom"}},
)
final = client.get(
f"/v2/sessions/{sid}/runs/{rid}", headers={"Authorization": "Bearer k"}
).json()
assert final["status"] == RUN_STATUS_ERROR

View file

@ -0,0 +1,111 @@
"""
Validation #4 — session state machine.
Drives:
* provisioning -> ready (via daemon register)
* ready -> busy (via run create) -> ready (via run finish event)
* provisioning -> error (provisioning failure)
* ready -> error (daemon goes silent for 90s — covered by sweeper test)
"""
import asyncio
import pytest
from litellm.proxy.agent_session_endpoints.auth import mint_daemon_token
from litellm.proxy.agent_session_endpoints.constants import (
SESSION_STATUS_ERROR,
SESSION_STATUS_PROVISIONING,
SESSION_STATUS_READY,
)
from litellm.proxy.agent_session_endpoints.state_machine import (
derive_session_status_from_runs,
is_valid_session_transition,
session_can_accept_runs,
)
def test_state_machine_transitions_pure():
# provisioning -> ready, error, terminated
assert is_valid_session_transition("provisioning", "ready")
assert is_valid_session_transition("provisioning", "error")
assert is_valid_session_transition("provisioning", "terminated")
# ready -> busy
assert is_valid_session_transition("ready", "busy")
# busy -> ready
assert is_valid_session_transition("busy", "ready")
# terminated is a sink
assert not is_valid_session_transition("terminated", "ready")
assert not is_valid_session_transition("terminated", "busy")
def test_session_can_accept_runs():
assert session_can_accept_runs("ready")
assert session_can_accept_runs("busy")
assert not session_can_accept_runs("provisioning")
assert not session_can_accept_runs("error")
assert not session_can_accept_runs("terminated")
def test_derive_session_status_from_runs():
# provisioning is gated until daemon registers
assert derive_session_status_from_runs("provisioning", True) is None
# ready -> busy when there's an active run
assert derive_session_status_from_runs("ready", True) == "busy"
# busy -> ready when no active runs
assert derive_session_status_from_runs("busy", False) == "ready"
# terminal sessions never transition
assert derive_session_status_from_runs("terminated", False) is None
assert derive_session_status_from_runs("error", True) is None
# No-op when target == current
assert derive_session_status_from_runs("ready", False) is None
assert derive_session_status_from_runs("busy", True) is None
def test_session_starts_provisioning(client, noop_provider):
res = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "t", "model": "gpt-4"},
)
agent_id = res.json()["id"]
res = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": agent_id, "repos": []},
)
assert res.status_code == 200
body = res.json()
assert body["status"] == SESSION_STATUS_PROVISIONING
assert body["daemon_token"]
def test_daemon_register_flips_to_ready(client, noop_provider, fake_prisma_client):
res = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "t", "model": "gpt-4"},
)
agent_id = res.json()["id"]
res = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": agent_id, "repos": []},
)
body = res.json()
sid = body["id"]
daemon_token = body["daemon_token"]
# Daemon registers — endpoint requires its own JWT.
reg = client.post(
f"/v2/sessions/{sid}/internal/register",
headers={"Authorization": f"Bearer {daemon_token}"},
json={"vm_id": "i-noop"},
)
assert reg.status_code == 200, reg.text
assert reg.json()["status"] == SESSION_STATUS_READY
after = client.get(f"/v2/sessions/{sid}", headers={"Authorization": "Bearer k"})
assert after.json()["status"] == SESSION_STATUS_READY

View file

@ -0,0 +1,115 @@
"""
Validation #8 — SSE resume.
Read 3 events, kill, reconnect with starting_seq=3, get rest. No gaps,
no dupes.
Implementation note: we don't open a real long-running SSE stream — we
shape the test as "given a run with N events, the events stream should
emit exactly the unseen ones and close once the run is terminal." We
check via the StreamingResponse body iterator with a finished run so
the loop terminates quickly.
"""
import json
import re
from litellm.proxy.agent_session_endpoints.constants import (
EVENT_TYPE_RUN_FINISHED,
)
def _bootstrap_ready(client, noop_provider):
agent = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "t", "model": "gpt-4"},
).json()
sess = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": agent["id"], "repos": []},
).json()
daemon_token = sess["daemon_token"]
sid = sess["id"]
client.post(
f"/v2/sessions/{sid}/internal/register",
headers={"Authorization": f"Bearer {daemon_token}"},
json={"vm_id": "i-noop"},
)
return sid, daemon_token
def _seed_events(client, sid, daemon_token, count: int) -> str:
run = client.post(
f"/v2/sessions/{sid}/runs",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "hi"}},
).json()
rid = run["id"]
client.get(
f"/v2/sessions/{sid}/runs/next/internal/poll",
headers={"Authorization": f"Bearer {daemon_token}"},
)
for i in range(count):
client.post(
f"/v2/sessions/{sid}/runs/{rid}/events:append",
headers={"Authorization": f"Bearer {daemon_token}"},
json={"event_type": "log", "payload": {"i": i}},
)
# End the run so the SSE stream can quiesce + close.
client.post(
f"/v2/sessions/{sid}/runs/{rid}/events:append",
headers={"Authorization": f"Bearer {daemon_token}"},
json={"event_type": EVENT_TYPE_RUN_FINISHED, "payload": {"result": "done"}},
)
return rid
def _parse_seqs(body_text: str):
return [int(m.group(1)) for m in re.finditer(r"^id:\s*(\d+)$", body_text, re.M)]
def test_sse_emits_all_events_from_zero(client, noop_provider):
sid, daemon_token = _bootstrap_ready(client, noop_provider)
rid = _seed_events(client, sid, daemon_token, count=5)
res = client.get(
f"/v2/sessions/{sid}/runs/{rid}/events",
headers={"Authorization": "Bearer k"},
)
assert res.status_code == 200
seqs = _parse_seqs(res.text)
# Seqs 1..6 (5 logs + 1 run_finished)
assert seqs == [1, 2, 3, 4, 5, 6]
def test_sse_resumes_with_starting_seq(client, noop_provider):
sid, daemon_token = _bootstrap_ready(client, noop_provider)
rid = _seed_events(client, sid, daemon_token, count=5)
res = client.get(
f"/v2/sessions/{sid}/runs/{rid}/events?starting_seq=4",
headers={"Authorization": "Bearer k"},
)
assert res.status_code == 200
seqs = _parse_seqs(res.text)
# starting_seq=4 means "last seen was 3, give me >= 4".
assert seqs == [4, 5, 6]
def test_sse_last_event_id_header_takes_precedence(client, noop_provider):
sid, daemon_token = _bootstrap_ready(client, noop_provider)
rid = _seed_events(client, sid, daemon_token, count=3)
res = client.get(
f"/v2/sessions/{sid}/runs/{rid}/events",
headers={
"Authorization": "Bearer k",
"Last-Event-ID": "2",
},
)
assert res.status_code == 200
seqs = _parse_seqs(res.text)
# Last-Event-ID=2 means "give me > 2", so we expect 3 + 4 (run_finished).
assert seqs == [3, 4]

View file

@ -0,0 +1,152 @@
"""
Validation #12 — JWT scoping.
* Daemon JWT for session A cannot append events to session B (sub mismatch)
* Expired JWT rejected with 401
* Terminated session rejects its still-valid JWT (status check)
* Daemon JWT cannot call /v1/chat/completions (wrong scope) — covered by
the dedicated daemon_token_auth dependency, since it's only mounted on
internal endpoints.
"""
import time
import jwt as pyjwt
import pytest
from litellm.proxy.agent_session_endpoints.auth import (
AGENT_RUNTIME_SCOPE,
decode_daemon_token,
hash_daemon_token,
mint_daemon_token,
)
from litellm.proxy.agent_session_endpoints.constants import (
AGENT_JWT_ALGORITHM,
)
def test_mint_and_decode_roundtrip():
token = mint_daemon_token(
session_id="sess_a",
agent_id="agent_a",
expires_at_epoch=int(time.time()) + 3600,
)
payload = decode_daemon_token(token)
assert payload["sub"] == "sess_a"
assert payload["agent_id"] == "agent_a"
assert payload["scope"] == AGENT_RUNTIME_SCOPE
def test_token_for_session_a_rejected_by_session_b(client, noop_provider):
# Create two sessions, get two daemon tokens.
a = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "x", "model": "gpt-4"},
).json()
sess_a = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": a["id"], "repos": []},
).json()
sess_b = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": a["id"], "repos": []},
).json()
# Use sess_a's daemon token to register sess_b.
res = client.post(
f"/v2/sessions/{sess_b['id']}/internal/register",
headers={"Authorization": f"Bearer {sess_a['daemon_token']}"},
json={"vm_id": "i-x"},
)
assert res.status_code == 403, res.text
def test_expired_jwt_rejected(client, noop_provider):
a = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "t", "model": "gpt-4"},
).json()
sess = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": a["id"], "repos": []},
).json()
# Mint a token that expired 1 minute ago.
expired = mint_daemon_token(
session_id=sess["id"],
agent_id=a["id"],
expires_at_epoch=int(time.time()) - 60,
)
res = client.post(
f"/v2/sessions/{sess['id']}/internal/register",
headers={"Authorization": f"Bearer {expired}"},
json={"vm_id": "i-x"},
)
assert res.status_code == 401
assert "expired" in res.text.lower()
def test_terminated_session_rejects_its_token(
client, noop_provider, fake_prisma_client
):
a = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "t", "model": "gpt-4"},
).json()
sess = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": a["id"], "repos": []},
).json()
sid = sess["id"]
daemon_token = sess["daemon_token"]
# Delete the session.
client.delete(f"/v2/sessions/{sid}", headers={"Authorization": "Bearer k"})
# Token should now be rejected with 410 Gone.
res = client.post(
f"/v2/sessions/{sid}/internal/heartbeat",
headers={"Authorization": f"Bearer {daemon_token}"},
json={"vm_id": "i-x"},
)
assert res.status_code == 410
def test_wrong_scope_rejected(client, noop_provider):
"""A token signed with the right secret but wrong scope is rejected."""
import os
secret = os.environ["LITELLM_AGENT_JWT_SECRET"]
bad = pyjwt.encode(
{
"sub": "sess_x",
"agent_id": "agent_x",
"iat": int(time.time()),
"exp": int(time.time()) + 3600,
"scope": "user_api_key", # NOT agent_runtime_internal
},
secret,
algorithm=AGENT_JWT_ALGORITHM,
)
res = client.post(
"/v2/sessions/sess_x/internal/register",
headers={"Authorization": f"Bearer {bad}"},
json={"vm_id": "i-x"},
)
assert res.status_code == 401
def test_token_hash_helper():
h1 = hash_daemon_token("abc")
h2 = hash_daemon_token("abc")
h3 = hash_daemon_token("abd")
assert h1 == h2
assert h1 != h3
assert len(h1) == 64 # sha256 hex

View file

@ -0,0 +1,162 @@
"""
Validation #14 — view-only admin cannot mutate any /v2/agents or
/v2/sessions resource.
The ``PROXY_ADMIN_VIEW_ONLY`` role is intended to grant cross-tenant
READ access (e.g. for a support UI) without granting write access. A
prior version of ``ownership.is_proxy_admin`` returned True for both
``PROXY_ADMIN`` and ``PROXY_ADMIN_VIEW_ONLY``, which let view-only
admins bypass the per-tenant ownership assertion on every write
endpoint and create / update / delete other tenants' rows.
Every state-mutating endpoint on the four routers (POST/PUT/PATCH/DELETE)
must call :func:`assert_caller_can_mutate` and return 403 for the
view-only admin role. Read endpoints (GET) must continue to work — that's
the whole point of the role.
"""
def _create_tenant_a_resources(client):
"""Bootstrap an agent + session + run owned by tenant A so we have
something for the view-only admin to try (and fail) to mutate."""
agent = client.post(
"/v2/agents",
headers={"Authorization": "Bearer k"},
json={"name": "tenant-a-agent", "model": "gpt-4"},
).json()
session = client.post(
"/v2/sessions",
headers={"Authorization": "Bearer k"},
json={"agent_id": agent["id"], "repos": []},
).json()
run = client.post(
f"/v2/sessions/{session['id']}/runs",
headers={"Authorization": "Bearer k"},
json={"prompt": {"text": "hi"}},
).json()
return agent["id"], session["id"], run["id"]
def test_view_only_admin_cannot_create_agent(view_only_admin_client, noop_provider):
res = view_only_admin_client.post(
"/v2/agents",
headers={"Authorization": "Bearer view-only"},
json={"name": "evil", "model": "gpt-4"},
)
assert res.status_code == 403
assert "view-only" in res.json()["detail"].lower()
def test_view_only_admin_cannot_update_other_tenant_agent(
client, view_only_admin_client, noop_provider
):
agent_id, _, _ = _create_tenant_a_resources(client)
res = view_only_admin_client.patch(
f"/v2/agents/{agent_id}",
headers={"Authorization": "Bearer view-only"},
json={"name": "hacked"},
)
assert res.status_code == 403
def test_view_only_admin_cannot_delete_other_tenant_agent(
client, view_only_admin_client, noop_provider
):
agent_id, _, _ = _create_tenant_a_resources(client)
res = view_only_admin_client.delete(
f"/v2/agents/{agent_id}",
headers={"Authorization": "Bearer view-only"},
)
assert res.status_code == 403
def test_view_only_admin_cannot_create_session(view_only_admin_client, noop_provider):
# Even minting a session with a non-existent agent must short-circuit
# to 403 BEFORE any DB activity — the role check is the first guard.
res = view_only_admin_client.post(
"/v2/sessions",
headers={"Authorization": "Bearer view-only"},
json={"agent_id": "agt_doesnotexist", "repos": []},
)
assert res.status_code == 403
def test_view_only_admin_cannot_delete_other_tenant_session(
client, view_only_admin_client, noop_provider
):
_, session_id, _ = _create_tenant_a_resources(client)
res = view_only_admin_client.delete(
f"/v2/sessions/{session_id}",
headers={"Authorization": "Bearer view-only"},
)
assert res.status_code == 403
def test_view_only_admin_cannot_create_run(
client, view_only_admin_client, noop_provider
):
_, session_id, _ = _create_tenant_a_resources(client)
res = view_only_admin_client.post(
f"/v2/sessions/{session_id}/runs",
headers={"Authorization": "Bearer view-only"},
json={"prompt": {"text": "evil"}},
)
assert res.status_code == 403
def test_view_only_admin_cannot_cancel_run(
client, view_only_admin_client, noop_provider
):
_, session_id, run_id = _create_tenant_a_resources(client)
res = view_only_admin_client.post(
f"/v2/sessions/{session_id}/runs/{run_id}/cancel",
headers={"Authorization": "Bearer view-only"},
)
assert res.status_code == 403
def test_view_only_admin_cannot_followup(client, view_only_admin_client, noop_provider):
_, session_id, _ = _create_tenant_a_resources(client)
res = view_only_admin_client.post(
f"/v2/sessions/{session_id}/followup",
headers={"Authorization": "Bearer view-only"},
json={"prompt": {"text": "evil followup"}},
)
assert res.status_code == 403
def test_view_only_admin_can_still_read_other_tenant_resources(
client, view_only_admin_client, noop_provider
):
"""The whole point of view-only is cross-tenant READ access — make
sure we didn't accidentally lock that down too."""
agent_id, session_id, run_id = _create_tenant_a_resources(client)
# Reads must succeed.
assert (
view_only_admin_client.get(
f"/v2/agents/{agent_id}",
headers={"Authorization": "Bearer view-only"},
).status_code
== 200
)
assert (
view_only_admin_client.get(
f"/v2/sessions/{session_id}",
headers={"Authorization": "Bearer view-only"},
).status_code
== 200
)
assert (
view_only_admin_client.get(
f"/v2/sessions/{session_id}/runs/{run_id}",
headers={"Authorization": "Bearer view-only"},
).status_code
== 200
)
# And list endpoints return tenant A's data (not filtered out).
res = view_only_admin_client.get(
"/v2/agents", headers={"Authorization": "Bearer view-only"}
)
assert res.status_code == 200
assert any(a["id"] == agent_id for a in res.json()["data"])