mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
merge: integrate A1 (LIT-2877 /v2/sessions) into LIT-2890 base
This commit is contained in:
commit
8131c7df0c
34 changed files with 5111 additions and 0 deletions
|
|
@ -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;
|
||||||
|
|
@ -1477,3 +1477,128 @@ model LiteLLM_AgentWorkerPairingToken {
|
||||||
|
|
||||||
@@index([team_id])
|
@@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])
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
|
]
|
||||||
218
litellm/proxy/agent_session_endpoints/agent_endpoints.py
Normal file
218
litellm/proxy/agent_session_endpoints/agent_endpoints.py
Normal 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}
|
||||||
178
litellm/proxy/agent_session_endpoints/auth.py
Normal file
178
litellm/proxy/agent_session_endpoints/auth.py
Normal 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}
|
||||||
249
litellm/proxy/agent_session_endpoints/cleanup.py
Normal file
249
litellm/proxy/agent_session_endpoints/cleanup.py
Normal 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
|
||||||
61
litellm/proxy/agent_session_endpoints/constants.py
Normal file
61
litellm/proxy/agent_session_endpoints/constants.py
Normal 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_"
|
||||||
38
litellm/proxy/agent_session_endpoints/ids.py
Normal file
38
litellm/proxy/agent_session_endpoints/ids.py
Normal 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)
|
||||||
291
litellm/proxy/agent_session_endpoints/internal_endpoints.py
Normal file
291
litellm/proxy/agent_session_endpoints/internal_endpoints.py
Normal 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,
|
||||||
|
}
|
||||||
144
litellm/proxy/agent_session_endpoints/ownership.py
Normal file
144
litellm/proxy/agent_session_endpoints/ownership.py
Normal 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)}
|
||||||
415
litellm/proxy/agent_session_endpoints/run_endpoints.py
Normal file
415
litellm/proxy/agent_session_endpoints/run_endpoints.py
Normal 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",
|
||||||
|
},
|
||||||
|
)
|
||||||
181
litellm/proxy/agent_session_endpoints/schemas.py
Normal file
181
litellm/proxy/agent_session_endpoints/schemas.py
Normal 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]
|
||||||
107
litellm/proxy/agent_session_endpoints/serialization.py
Normal file
107
litellm/proxy/agent_session_endpoints/serialization.py
Normal 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),
|
||||||
|
)
|
||||||
603
litellm/proxy/agent_session_endpoints/session_endpoints.py
Normal file
603
litellm/proxy/agent_session_endpoints/session_endpoints.py
Normal 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],
|
||||||
|
}
|
||||||
130
litellm/proxy/agent_session_endpoints/state_machine.py
Normal file
130
litellm/proxy/agent_session_endpoints/state_machine.py
Normal 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
|
||||||
|
|
@ -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()
|
||||||
|
|
@ -939,9 +939,33 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915
|
||||||
## Initialize shared aiohttp session for connection reuse
|
## Initialize shared aiohttp session for connection reuse
|
||||||
shared_aiohttp_session = await _initialize_shared_aiohttp_session()
|
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
|
# End of startup event
|
||||||
yield
|
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
|
# Shutdown event - close shared aiohttp session
|
||||||
if shared_aiohttp_session is not None:
|
if shared_aiohttp_session is not None:
|
||||||
try:
|
try:
|
||||||
|
|
@ -14900,6 +14924,45 @@ app.include_router(agent_pool_status_router)
|
||||||
# Eager: /models/{name}:method overlaps with the OpenAI /models endpoint.
|
# Eager: /models/{name}:method overlaps with the OpenAI /models endpoint.
|
||||||
app.include_router(google_router)
|
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)
|
attach_lazy_features(app)
|
||||||
app.add_middleware(
|
app.add_middleware(
|
||||||
RequestSizeLimitMiddleware,
|
RequestSizeLimitMiddleware,
|
||||||
|
|
|
||||||
|
|
@ -1477,3 +1477,128 @@ model LiteLLM_AgentWorkerPairingToken {
|
||||||
|
|
||||||
@@index([team_id])
|
@@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])
|
||||||
|
}
|
||||||
|
|
|
||||||
125
schema.prisma
125
schema.prisma
|
|
@ -1477,3 +1477,128 @@ model LiteLLM_AgentWorkerPairingToken {
|
||||||
|
|
||||||
@@index([team_id])
|
@@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])
|
||||||
|
}
|
||||||
|
|
|
||||||
308
tests/test_litellm/proxy/agent_session_endpoints/conftest.py
Normal file
308
tests/test_litellm/proxy/agent_session_endpoints/conftest.py
Normal 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()
|
||||||
|
|
@ -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)
|
||||||
|
|
@ -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)
|
||||||
|
|
@ -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,
|
||||||
|
}
|
||||||
|
|
@ -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
|
||||||
|
|
@ -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"]
|
||||||
|
|
@ -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
|
||||||
|
|
@ -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
|
||||||
|
|
@ -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
|
||||||
|
|
@ -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
|
||||||
|
|
@ -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
|
||||||
|
|
@ -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
|
||||||
|
|
@ -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]
|
||||||
|
|
@ -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
|
||||||
|
|
@ -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"])
|
||||||
Loading…
Add table
Reference in a new issue