mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +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])
|
||||
}
|
||||
|
||||
// ===========================================================================
|
||||
// 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
|
||||
shared_aiohttp_session = await _initialize_shared_aiohttp_session()
|
||||
|
||||
## /v2/agents+sessions cleanup sweeper (Epic A — Cursor SDK).
|
||||
## Idempotent; safe even if the agent_session_endpoints module never
|
||||
## sees traffic.
|
||||
try:
|
||||
from litellm.proxy.agent_session_endpoints.cleanup import (
|
||||
start_cleanup_sweeper,
|
||||
)
|
||||
|
||||
start_cleanup_sweeper()
|
||||
except Exception as exc:
|
||||
verbose_proxy_logger.warning(
|
||||
"agent_session_endpoints cleanup sweeper failed to start: %s", exc
|
||||
)
|
||||
|
||||
# End of startup event
|
||||
yield
|
||||
|
||||
## Stop the agent session cleanup sweeper.
|
||||
try:
|
||||
from litellm.proxy.agent_session_endpoints.cleanup import (
|
||||
stop_cleanup_sweeper,
|
||||
)
|
||||
|
||||
stop_cleanup_sweeper()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Shutdown event - close shared aiohttp session
|
||||
if shared_aiohttp_session is not None:
|
||||
try:
|
||||
|
|
@ -14900,6 +14924,45 @@ app.include_router(agent_pool_status_router)
|
|||
# Eager: /models/{name}:method overlaps with the OpenAI /models endpoint.
|
||||
app.include_router(google_router)
|
||||
|
||||
# /v2/agents, /v2/sessions — Cursor SDK agent runtime (Epic A).
|
||||
# Mounted under /v2/ to avoid collision with the existing /v1/agents
|
||||
# (A2A registry in litellm/proxy/agent_endpoints/).
|
||||
#
|
||||
# SECURITY: the daemon JWT secret is a separate credential from the proxy
|
||||
# master key. If ``LITELLM_AGENT_JWT_SECRET`` is not set, refuse to mount
|
||||
# these routers — silently signing daemon tokens with the master key (or
|
||||
# any default) would conflate two distinct auth surfaces and let a
|
||||
# captured daemon JWT mint master-key-authority API keys.
|
||||
from litellm.proxy.agent_session_endpoints.auth import (
|
||||
is_agent_jwt_secret_configured,
|
||||
)
|
||||
|
||||
if is_agent_jwt_secret_configured():
|
||||
from litellm.proxy.agent_session_endpoints import (
|
||||
agent_router as agent_session_agent_router,
|
||||
)
|
||||
from litellm.proxy.agent_session_endpoints import (
|
||||
internal_router as agent_session_internal_router,
|
||||
)
|
||||
from litellm.proxy.agent_session_endpoints import (
|
||||
run_router as agent_session_run_router,
|
||||
)
|
||||
from litellm.proxy.agent_session_endpoints import (
|
||||
session_router as agent_session_session_router,
|
||||
)
|
||||
|
||||
app.include_router(agent_session_agent_router)
|
||||
app.include_router(agent_session_session_router)
|
||||
app.include_router(agent_session_run_router)
|
||||
app.include_router(agent_session_internal_router)
|
||||
else:
|
||||
verbose_proxy_logger.error(
|
||||
"agent_session_endpoints (/v2/agents, /v2/sessions) NOT mounted: "
|
||||
"LITELLM_AGENT_JWT_SECRET is not set. Set this env var to a "
|
||||
"dedicated random secret (distinct from LITELLM_MASTER_KEY) to "
|
||||
"enable the Cursor SDK agent runtime."
|
||||
)
|
||||
|
||||
attach_lazy_features(app)
|
||||
app.add_middleware(
|
||||
RequestSizeLimitMiddleware,
|
||||
|
|
|
|||
|
|
@ -1477,3 +1477,128 @@ model LiteLLM_AgentWorkerPairingToken {
|
|||
|
||||
@@index([team_id])
|
||||
}
|
||||
|
||||
// ===========================================================================
|
||||
// Agent Sessions / Runs (Cursor SDK on LiteLLM)
|
||||
//
|
||||
// Three-level hierarchy:
|
||||
// Agent — definition (model, system prompt, default repos, tools)
|
||||
// Session — VM-backed conversation, owned by an Agent
|
||||
// Run — single turn within a Session
|
||||
// RunEvent — append-only event log per run (for resumable SSE)
|
||||
// ===========================================================================
|
||||
|
||||
model LiteLLM_Agent {
|
||||
id String @id // "agent_<uuid>"
|
||||
name String
|
||||
user_api_key_hash String
|
||||
team_id String?
|
||||
model String
|
||||
system_prompt String?
|
||||
default_repos Json?
|
||||
default_env_vars Json?
|
||||
tools_config Json?
|
||||
metadata Json?
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
|
||||
sessions LiteLLM_AgentSession[]
|
||||
|
||||
@@index([user_api_key_hash])
|
||||
@@index([team_id])
|
||||
}
|
||||
|
||||
model LiteLLM_AgentSession {
|
||||
id String @id // "sess_<uuid>"
|
||||
agent_id String
|
||||
user_api_key_hash String
|
||||
team_id String?
|
||||
vm_id String?
|
||||
vm_provider String?
|
||||
repos Json
|
||||
env_vars Json?
|
||||
status String @default("provisioning")
|
||||
daemon_token_hash String?
|
||||
expires_at DateTime
|
||||
last_heartbeat_at DateTime?
|
||||
idempotency_key String?
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
terminated_at DateTime?
|
||||
|
||||
agent LiteLLM_Agent @relation(fields: [agent_id], references: [id], onDelete: Cascade)
|
||||
runs LiteLLM_AgentRun[]
|
||||
|
||||
@@unique([user_api_key_hash, idempotency_key])
|
||||
@@index([agent_id])
|
||||
@@index([status, expires_at])
|
||||
@@index([user_api_key_hash])
|
||||
}
|
||||
|
||||
model LiteLLM_AgentRun {
|
||||
id String @id // "run_<uuid>"
|
||||
session_id String
|
||||
parent_run_id String?
|
||||
status String @default("queued")
|
||||
prompt Json
|
||||
result String?
|
||||
git_branches Json?
|
||||
idempotency_key String?
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
started_at DateTime?
|
||||
terminated_at DateTime?
|
||||
|
||||
session LiteLLM_AgentSession @relation(fields: [session_id], references: [id], onDelete: Cascade)
|
||||
events LiteLLM_AgentRunEvent[]
|
||||
|
||||
@@unique([session_id, idempotency_key])
|
||||
@@index([session_id, status])
|
||||
@@index([session_id, created_at])
|
||||
}
|
||||
|
||||
model LiteLLM_AgentRunEvent {
|
||||
id String @id @default(uuid())
|
||||
run_id String
|
||||
seq Int
|
||||
event_type String
|
||||
payload Json
|
||||
created_at DateTime @default(now())
|
||||
|
||||
run LiteLLM_AgentRun @relation(fields: [run_id], references: [id], onDelete: Cascade)
|
||||
|
||||
@@unique([run_id, seq])
|
||||
@@index([run_id, seq])
|
||||
}
|
||||
|
||||
// ===========================================================================
|
||||
// Warm pool VM tracking (LIT-2890 / Epic B2)
|
||||
//
|
||||
// Tracks the lifecycle of pre-provisioned EC2 (or other provider) VMs used
|
||||
// for instant session attach. Each row maps to one underlying instance.
|
||||
//
|
||||
// State machine:
|
||||
// provisioning → warm → hydrating → attached → terminating → terminated
|
||||
//
|
||||
// On session end, the VM is terminated (NOT recycled) — security boundary.
|
||||
// The maintenance loop refills `warm` slots; rows in `terminated` are kept
|
||||
// for audit until pruned.
|
||||
// ===========================================================================
|
||||
model LiteLLM_AgentVM {
|
||||
id String @id // EC2 instance id (e.g. "i-0abcd...")
|
||||
provider String // "ec2" | "noop" | "self_hosted"
|
||||
region String?
|
||||
state String // provisioning|warm|hydrating|attached|terminating|terminated
|
||||
team_id String // owner team — pool is per-team
|
||||
pool_id String // logical pool key (currently == team_id)
|
||||
attached_session_id String? // FK to LiteLLM_AgentSession.id when state=attached
|
||||
created_at DateTime @default(now())
|
||||
warmed_at DateTime?
|
||||
last_hydrate_at DateTime?
|
||||
terminated_at DateTime?
|
||||
metadata Json? // public_ip, private_ip, ssm_status, etc.
|
||||
|
||||
@@index([state, pool_id])
|
||||
@@index([team_id, state])
|
||||
@@index([attached_session_id])
|
||||
}
|
||||
|
|
|
|||
125
schema.prisma
125
schema.prisma
|
|
@ -1477,3 +1477,128 @@ model LiteLLM_AgentWorkerPairingToken {
|
|||
|
||||
@@index([team_id])
|
||||
}
|
||||
|
||||
// ===========================================================================
|
||||
// Agent Sessions / Runs (Cursor SDK on LiteLLM)
|
||||
//
|
||||
// Three-level hierarchy:
|
||||
// Agent — definition (model, system prompt, default repos, tools)
|
||||
// Session — VM-backed conversation, owned by an Agent
|
||||
// Run — single turn within a Session
|
||||
// RunEvent — append-only event log per run (for resumable SSE)
|
||||
// ===========================================================================
|
||||
|
||||
model LiteLLM_Agent {
|
||||
id String @id // "agent_<uuid>"
|
||||
name String
|
||||
user_api_key_hash String
|
||||
team_id String?
|
||||
model String
|
||||
system_prompt String?
|
||||
default_repos Json?
|
||||
default_env_vars Json?
|
||||
tools_config Json?
|
||||
metadata Json?
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
|
||||
sessions LiteLLM_AgentSession[]
|
||||
|
||||
@@index([user_api_key_hash])
|
||||
@@index([team_id])
|
||||
}
|
||||
|
||||
model LiteLLM_AgentSession {
|
||||
id String @id // "sess_<uuid>"
|
||||
agent_id String
|
||||
user_api_key_hash String
|
||||
team_id String?
|
||||
vm_id String?
|
||||
vm_provider String?
|
||||
repos Json
|
||||
env_vars Json?
|
||||
status String @default("provisioning")
|
||||
daemon_token_hash String?
|
||||
expires_at DateTime
|
||||
last_heartbeat_at DateTime?
|
||||
idempotency_key String?
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
terminated_at DateTime?
|
||||
|
||||
agent LiteLLM_Agent @relation(fields: [agent_id], references: [id], onDelete: Cascade)
|
||||
runs LiteLLM_AgentRun[]
|
||||
|
||||
@@unique([user_api_key_hash, idempotency_key])
|
||||
@@index([agent_id])
|
||||
@@index([status, expires_at])
|
||||
@@index([user_api_key_hash])
|
||||
}
|
||||
|
||||
model LiteLLM_AgentRun {
|
||||
id String @id // "run_<uuid>"
|
||||
session_id String
|
||||
parent_run_id String?
|
||||
status String @default("queued")
|
||||
prompt Json
|
||||
result String?
|
||||
git_branches Json?
|
||||
idempotency_key String?
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
started_at DateTime?
|
||||
terminated_at DateTime?
|
||||
|
||||
session LiteLLM_AgentSession @relation(fields: [session_id], references: [id], onDelete: Cascade)
|
||||
events LiteLLM_AgentRunEvent[]
|
||||
|
||||
@@unique([session_id, idempotency_key])
|
||||
@@index([session_id, status])
|
||||
@@index([session_id, created_at])
|
||||
}
|
||||
|
||||
model LiteLLM_AgentRunEvent {
|
||||
id String @id @default(uuid())
|
||||
run_id String
|
||||
seq Int
|
||||
event_type String
|
||||
payload Json
|
||||
created_at DateTime @default(now())
|
||||
|
||||
run LiteLLM_AgentRun @relation(fields: [run_id], references: [id], onDelete: Cascade)
|
||||
|
||||
@@unique([run_id, seq])
|
||||
@@index([run_id, seq])
|
||||
}
|
||||
|
||||
// ===========================================================================
|
||||
// Warm pool VM tracking (LIT-2890 / Epic B2)
|
||||
//
|
||||
// Tracks the lifecycle of pre-provisioned EC2 (or other provider) VMs used
|
||||
// for instant session attach. Each row maps to one underlying instance.
|
||||
//
|
||||
// State machine:
|
||||
// provisioning → warm → hydrating → attached → terminating → terminated
|
||||
//
|
||||
// On session end, the VM is terminated (NOT recycled) — security boundary.
|
||||
// The maintenance loop refills `warm` slots; rows in `terminated` are kept
|
||||
// for audit until pruned.
|
||||
// ===========================================================================
|
||||
model LiteLLM_AgentVM {
|
||||
id String @id // EC2 instance id (e.g. "i-0abcd...")
|
||||
provider String // "ec2" | "noop" | "self_hosted"
|
||||
region String?
|
||||
state String // provisioning|warm|hydrating|attached|terminating|terminated
|
||||
team_id String // owner team — pool is per-team
|
||||
pool_id String // logical pool key (currently == team_id)
|
||||
attached_session_id String? // FK to LiteLLM_AgentSession.id when state=attached
|
||||
created_at DateTime @default(now())
|
||||
warmed_at DateTime?
|
||||
last_hydrate_at DateTime?
|
||||
terminated_at DateTime?
|
||||
metadata Json? // public_ip, private_ip, ssm_status, etc.
|
||||
|
||||
@@index([state, pool_id])
|
||||
@@index([team_id, state])
|
||||
@@index([attached_session_id])
|
||||
}
|
||||
|
|
|
|||
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