diff --git a/litellm/proxy/agent_session_endpoints/session_endpoints.py b/litellm/proxy/agent_session_endpoints/session_endpoints.py new file mode 100644 index 00000000000..faf3dbe0a8e --- /dev/null +++ b/litellm/proxy/agent_session_endpoints/session_endpoints.py @@ -0,0 +1,554 @@ +""" +Session CRUD endpoints — POST/GET/DELETE /v2/sessions{,/}, 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_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 _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"), +): + 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), +): + 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. + """ + 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") + + # Else start a fresh run. + 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(), + } + ) + 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], + }