Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/management_endpoints/workflow_management_endpoints.py: 56%
241 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2WORKFLOW RUN MANAGEMENT
4Generic durable state tracking for agents and automated workflows.
6POST /v1/workflows/runs - Create a workflow run
7GET /v1/workflows/runs - List runs (filter by type, status)
8GET /v1/workflows/runs/{run_id} - Get run with latest event
9PATCH /v1/workflows/runs/{run_id} - Update status, metadata, output
10POST /v1/workflows/runs/{run_id}/events - Append event (updates run status)
11GET /v1/workflows/runs/{run_id}/events - Full event log
12POST /v1/workflows/runs/{run_id}/messages - Append conversation message
13GET /v1/workflows/runs/{run_id}/messages - Fetch conversation history
14"""
16import json
17from collections.abc import Mapping, Sequence
18from typing import Final, Literal, Protocol, TypedDict
20from fastapi import APIRouter, Depends, HTTPException, Query
22try:
23 from prisma.errors import UniqueViolationError
24except ImportError:
25 UniqueViolationError = None
26from pydantic import BaseModel
28from litellm._logging import verbose_proxy_logger
29from litellm.proxy._types import (
30 CommonProxyErrors,
31 LitellmUserRoles,
32 UserAPIKeyAuth,
33 user_api_key_has_admin_view,
34)
35from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
36from litellm.repositories.table_repositories import (
37 WorkflowEventRepository,
38 WorkflowMessageRepository,
39 WorkflowRunRepository,
40)
42router: Final = APIRouter()
44_MAX_SEQUENCE_RETRIES: Final = 5
47def _json(value: object) -> str:
48 """Serialize a Python value for prisma-client-py Json fields (must be a string)."""
49 return json.dumps(value)
52def _is_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
53 return user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
56def _read_scope_caller(user_api_key_dict: UserAPIKeyAuth) -> UserAPIKeyAuth | None:
57 return None if user_api_key_has_admin_view(user_api_key_dict) else user_api_key_dict
60def _caller_key(user_api_key_dict: UserAPIKeyAuth) -> str | None:
61 """Return the hashed key token that identifies this caller, or None for master key."""
62 return user_api_key_dict.token
65# Status transitions driven by event_type
66_EVENT_STATUS_MAP: Final[Mapping[str, str]] = {
67 "step.started": "running",
68 "step.failed": "failed",
69 "hook.waiting": "paused",
70 "hook.received": "running",
71}
74# ---------------------------------------------------------------------------
75# Request / Response models
76# ---------------------------------------------------------------------------
79class WorkflowRunCreateRequest(BaseModel):
80 workflow_type: str
81 input: Mapping[str, object] | None = None
82 metadata: Mapping[str, object] | None = None
85WorkflowRunStatus = Literal["pending", "running", "paused", "completed", "failed"]
88class WorkflowRunUpdateRequest(BaseModel):
89 status: WorkflowRunStatus | None = None
90 output: Mapping[str, object] | None = None
91 metadata: Mapping[str, object] | None = None
94class WorkflowEventCreateRequest(BaseModel):
95 event_type: str
96 step_name: str
97 data: Mapping[str, object] | None = None
100class WorkflowMessageCreateRequest(BaseModel):
101 role: str
102 content: str
103 session_id: str | None = None
106class _RunRow(Protocol):
107 @property
108 def created_by(self) -> str | None: ... 108 ↛ exitline 108 didn't return from function 'created_by' because
111class _SeqRow(Protocol):
112 @property
113 def sequence_number(self) -> int: ... 113 ↛ exitline 113 didn't return from function 'sequence_number' because
116class _RunCreateData(TypedDict, total=False):
117 workflow_type: str
118 created_by: str | None
119 input: str
120 metadata: str
123class _RunWhere(TypedDict, total=False):
124 workflow_type: str
125 status: str | Mapping[str, Sequence[str]]
126 created_by: str
129class _RunUpdateData(TypedDict, total=False):
130 status: WorkflowRunStatus
131 output: str
132 metadata: str
135class _EventCreateData(TypedDict, total=False):
136 run_id: str
137 event_type: str
138 step_name: str
139 sequence_number: int
140 data: str
143class _MessageCreateData(TypedDict, total=False):
144 run_id: str
145 role: str
146 content: str
147 sequence_number: int
148 session_id: str
151# ---------------------------------------------------------------------------
152# Helpers
153# ---------------------------------------------------------------------------
156async def _get_next_sequence_number(prisma_client: object, run_id: str, table: str) -> int:
157 """Return MAX(sequence_number) + 1 for the given run, for either events or messages."""
158 if table == "events":
159 rows: Sequence[_SeqRow] = await WorkflowEventRepository(prisma_client).table.find_many(
160 where={"run_id": run_id},
161 order={"sequence_number": "desc"},
162 take=1,
163 )
164 else:
165 rows = await WorkflowMessageRepository(prisma_client).table.find_many(
166 where={"run_id": run_id},
167 order={"sequence_number": "desc"},
168 take=1,
169 )
170 return (rows[0].sequence_number + 1) if rows else 0
173async def _require_run(
174 prisma_client: object,
175 run_id: str,
176 user_api_key_dict: UserAPIKeyAuth | None = None,
177) -> _RunRow:
178 """Return the run or raise 404. For non-admin callers, also enforce key ownership."""
179 run: Final[_RunRow | None] = await WorkflowRunRepository(prisma_client).table.find_unique(where={"run_id": run_id})
180 if run is None: 180 ↛ 182line 180 didn't jump to line 182 because the condition on line 180 was always true
181 raise HTTPException(status_code=404, detail=f"Run '{run_id}' not found")
182 if user_api_key_dict is not None and not _is_admin(user_api_key_dict):
183 caller: Final = _caller_key(user_api_key_dict)
184 if not caller or run.created_by != caller:
185 raise HTTPException(status_code=404, detail=f"Run '{run_id}' not found")
186 return run
189# ---------------------------------------------------------------------------
190# Endpoints
191# ---------------------------------------------------------------------------
194@router.post(
195 "/v1/workflows/runs",
196 tags=["workflow management"],
197 dependencies=[Depends(user_api_key_auth)],
198)
199async def create_workflow_run(
200 data: WorkflowRunCreateRequest,
201 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
202):
203 """Create a new workflow run. Returns run_id and session_id.
205 The caller's API key token is stored as created_by so that non-admin keys
206 can only see and modify their own runs.
207 """
208 from litellm.proxy.proxy_server import prisma_client
210 if prisma_client is None: 210 ↛ 211line 210 didn't jump to line 211 because the condition on line 210 was never true
211 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
213 try:
214 create_data: Final[_RunCreateData] = {
215 "workflow_type": data.workflow_type,
216 "created_by": _caller_key(user_api_key_dict),
217 }
218 if data.input is not None:
219 create_data["input"] = _json(data.input)
220 if data.metadata is not None:
221 create_data["metadata"] = _json(data.metadata)
222 run: Final[_RunRow] = await WorkflowRunRepository(prisma_client).table.create(data=create_data)
223 return run
224 except Exception as e:
225 verbose_proxy_logger.exception("Error creating workflow run: %s", e)
226 raise HTTPException(status_code=500, detail=str(e))
229@router.get(
230 "/v1/workflows/runs",
231 tags=["workflow management"],
232 dependencies=[Depends(user_api_key_auth)],
233)
234async def list_workflow_runs(
235 workflow_type: str | None = Query(None),
236 status: str | None = Query(None),
237 limit: int = Query(50, ge=1, le=250),
238 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
239):
240 """List workflow runs. Filter by workflow_type and/or status.
242 Non-admin callers only see runs created by their own API key.
243 """
244 from litellm.proxy.proxy_server import prisma_client
246 if prisma_client is None: 246 ↛ 247line 246 didn't jump to line 247 because the condition on line 246 was never true
247 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
249 where: Final[_RunWhere] = {}
250 if workflow_type:
251 where["workflow_type"] = workflow_type
252 if status:
253 statuses: Final = [s.strip() for s in status.split(",")]
254 where["status"] = {"in": statuses} if len(statuses) > 1 else statuses[0]
256 # Non-admin callers are scoped to their own key.
257 if not user_api_key_has_admin_view(user_api_key_dict): 257 ↛ 258line 257 didn't jump to line 258 because the condition on line 257 was never true
258 caller: Final = _caller_key(user_api_key_dict)
259 if caller:
260 where["created_by"] = caller
262 try:
263 runs: Final[Sequence[object]] = await WorkflowRunRepository(prisma_client).table.find_many(
264 where=where,
265 order={"created_at": "desc"},
266 take=limit,
267 )
268 return {"runs": runs, "count": len(runs)}
269 except Exception as e:
270 verbose_proxy_logger.exception("Error listing workflow runs: %s", e)
271 raise HTTPException(status_code=500, detail=str(e))
274@router.get(
275 "/v1/workflows/runs/{run_id}",
276 tags=["workflow management"],
277 dependencies=[Depends(user_api_key_auth)],
278)
279async def get_workflow_run(
280 run_id: str,
281 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
282):
283 """Get a workflow run with its most recent event."""
284 from litellm.proxy.proxy_server import prisma_client
286 if prisma_client is None: 286 ↛ 287line 286 didn't jump to line 287 because the condition on line 286 was never true
287 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
289 try:
290 run: Final[_RunRow | None] = await WorkflowRunRepository(prisma_client).table.find_unique(
291 where={"run_id": run_id},
292 include={"events": {"order_by": {"sequence_number": "desc"}, "take": 1}},
293 )
294 if run is None: 294 ↛ 296line 294 didn't jump to line 296 because the condition on line 294 was always true
295 raise HTTPException(status_code=404, detail=f"Run '{run_id}' not found")
296 if not user_api_key_has_admin_view(user_api_key_dict):
297 caller: Final = _caller_key(user_api_key_dict)
298 if not caller or run.created_by != caller:
299 raise HTTPException(status_code=404, detail=f"Run '{run_id}' not found")
300 return run
301 except HTTPException:
302 raise
303 except Exception as e:
304 verbose_proxy_logger.exception("Error getting workflow run: %s", e)
305 raise HTTPException(status_code=500, detail=str(e))
308@router.patch(
309 "/v1/workflows/runs/{run_id}",
310 tags=["workflow management"],
311 dependencies=[Depends(user_api_key_auth)],
312)
313async def update_workflow_run(
314 run_id: str,
315 data: WorkflowRunUpdateRequest,
316 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
317):
318 """Update status, metadata, or output on a workflow run."""
319 from litellm.proxy.proxy_server import prisma_client
321 if prisma_client is None: 321 ↛ 322line 321 didn't jump to line 322 because the condition on line 321 was never true
322 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
324 update: Final[_RunUpdateData] = {}
325 if data.status is not None:
326 update["status"] = data.status
327 if data.output is not None:
328 update["output"] = _json(data.output)
329 if data.metadata is not None:
330 update["metadata"] = _json(data.metadata)
332 if not update:
333 raise HTTPException(status_code=400, detail="No fields to update")
335 # Enforce ownership before writing.
336 await _require_run(prisma_client, run_id, user_api_key_dict)
338 try:
339 run: Final[_RunRow | None] = await WorkflowRunRepository(prisma_client).table.update(
340 where={"run_id": run_id},
341 data=update,
342 )
343 if run is None:
344 raise HTTPException(status_code=404, detail=f"Run '{run_id}' not found")
345 return run
346 except HTTPException:
347 raise
348 except Exception as e:
349 verbose_proxy_logger.exception("Error updating workflow run: %s", e)
350 raise HTTPException(status_code=500, detail=str(e))
353@router.post(
354 "/v1/workflows/runs/{run_id}/events",
355 tags=["workflow management"],
356 dependencies=[Depends(user_api_key_auth)],
357)
358async def append_workflow_event(
359 run_id: str,
360 data: WorkflowEventCreateRequest,
361 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
362):
363 """Append an event to the run's event log. Also updates run.status if event_type maps to a status.
365 Sequence numbers use optimistic concurrency: on a unique-constraint collision
366 (concurrent append), retries up to _MAX_SEQUENCE_RETRIES times with a fresh MAX+1.
367 The event+status update is atomic in a single DB transaction.
368 """
369 from litellm.proxy.proxy_server import prisma_client
371 if prisma_client is None: 371 ↛ 372line 371 didn't jump to line 372 because the condition on line 371 was never true
372 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
374 await _require_run(prisma_client, run_id, user_api_key_dict)
376 new_status: Final = _EVENT_STATUS_MAP.get(data.event_type)
378 for attempt in range(_MAX_SEQUENCE_RETRIES):
379 try:
380 seq = await _get_next_sequence_number(prisma_client, run_id, "events")
381 event_data: _EventCreateData = {
382 "run_id": run_id,
383 "event_type": data.event_type,
384 "step_name": data.step_name,
385 "sequence_number": seq,
386 }
387 if data.data is not None:
388 event_data["data"] = _json(data.data)
390 async with prisma_client.db.tx() as tx:
391 event: object = await tx.litellm_workflowevent.create(data=event_data)
392 if new_status:
393 await tx.litellm_workflowrun.update(
394 where={"run_id": run_id},
395 data={"status": new_status},
396 )
398 return event
400 except Exception as e:
401 if UniqueViolationError is not None and isinstance(e, UniqueViolationError):
402 if attempt == _MAX_SEQUENCE_RETRIES - 1:
403 verbose_proxy_logger.exception(
404 "Sequence number collision after %d retries for run %s",
405 _MAX_SEQUENCE_RETRIES,
406 run_id,
407 )
408 raise HTTPException(
409 status_code=409,
410 detail="Concurrent write conflict — please retry",
411 )
412 continue
413 verbose_proxy_logger.exception("Error appending workflow event: %s", e)
414 raise HTTPException(status_code=500, detail=str(e))
416 raise HTTPException(status_code=500, detail="Failed to append event") # pragma: no cover
419@router.get(
420 "/v1/workflows/runs/{run_id}/events",
421 tags=["workflow management"],
422 dependencies=[Depends(user_api_key_auth)],
423)
424async def list_workflow_events(
425 run_id: str,
426 limit: int = Query(100, ge=1, le=500),
427 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
428):
429 """Fetch event log for a run, ordered by sequence_number. Default limit 100, max 500."""
430 from litellm.proxy.proxy_server import prisma_client
432 if prisma_client is None: 432 ↛ 433line 432 didn't jump to line 433 because the condition on line 432 was never true
433 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
435 await _require_run(prisma_client, run_id, _read_scope_caller(user_api_key_dict))
437 try:
438 events: Final[Sequence[object]] = await WorkflowEventRepository(prisma_client).table.find_many(
439 where={"run_id": run_id},
440 order={"sequence_number": "asc"},
441 take=limit,
442 )
443 return {"events": events, "count": len(events)}
444 except Exception as e:
445 verbose_proxy_logger.exception("Error listing workflow events: %s", e)
446 raise HTTPException(status_code=500, detail=str(e))
449@router.post(
450 "/v1/workflows/runs/{run_id}/messages",
451 tags=["workflow management"],
452 dependencies=[Depends(user_api_key_auth)],
453)
454async def append_workflow_message(
455 run_id: str,
456 data: WorkflowMessageCreateRequest,
457 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
458):
459 """Append a conversation message. Stores full content (not truncated).
461 Uses optimistic concurrency for sequence numbers.
462 """
463 from litellm.proxy.proxy_server import prisma_client
465 if prisma_client is None: 465 ↛ 466line 465 didn't jump to line 466 because the condition on line 465 was never true
466 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
468 await _require_run(prisma_client, run_id, user_api_key_dict)
470 for attempt in range(_MAX_SEQUENCE_RETRIES):
471 try:
472 seq = await _get_next_sequence_number(prisma_client, run_id, "messages")
473 msg_data: _MessageCreateData = {
474 "run_id": run_id,
475 "role": data.role,
476 "content": data.content,
477 "sequence_number": seq,
478 }
479 if data.session_id is not None:
480 msg_data["session_id"] = data.session_id
481 msg: object = await WorkflowMessageRepository(prisma_client).table.create(data=msg_data)
482 return msg
484 except Exception as e:
485 if UniqueViolationError is not None and isinstance(e, UniqueViolationError):
486 if attempt == _MAX_SEQUENCE_RETRIES - 1:
487 verbose_proxy_logger.exception(
488 "Sequence number collision after %d retries for run %s",
489 _MAX_SEQUENCE_RETRIES,
490 run_id,
491 )
492 raise HTTPException(
493 status_code=409,
494 detail="Concurrent write conflict — please retry",
495 )
496 continue
497 verbose_proxy_logger.exception("Error appending workflow message: %s", e)
498 raise HTTPException(status_code=500, detail=str(e))
500 raise HTTPException(status_code=500, detail="Failed to append message") # pragma: no cover
503@router.get(
504 "/v1/workflows/runs/{run_id}/messages",
505 tags=["workflow management"],
506 dependencies=[Depends(user_api_key_auth)],
507)
508async def list_workflow_messages(
509 run_id: str,
510 limit: int = Query(100, ge=1, le=500),
511 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
512):
513 """Fetch conversation history for a run, ordered by sequence_number. Default limit 100, max 500."""
514 from litellm.proxy.proxy_server import prisma_client
516 if prisma_client is None: 516 ↛ 517line 516 didn't jump to line 517 because the condition on line 516 was never true
517 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
519 await _require_run(prisma_client, run_id, _read_scope_caller(user_api_key_dict))
521 try:
522 messages: Final[Sequence[object]] = await WorkflowMessageRepository(prisma_client).table.find_many(
523 where={"run_id": run_id},
524 order={"sequence_number": "asc"},
525 take=limit,
526 )
527 return {"messages": messages, "count": len(messages)}
528 except Exception as e:
529 verbose_proxy_logger.exception("Error listing workflow messages: %s", e)
530 raise HTTPException(status_code=500, detail=str(e))