Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py: 41%
214 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"""
2AI Usage Chat - uses LLM tool calling to answer questions about
3usage/spend data by querying the aggregated daily activity endpoints.
4"""
6import json
7from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Mapping, Sequence
8from datetime import date
9from typing import Final, Literal, NamedTuple, Protocol, cast, overload
11from typing_extensions import ReadOnly, TypedDict
13import litellm
14from litellm._logging import verbose_proxy_logger
15from litellm.constants import DEFAULT_COMPETITOR_DISCOVERY_MODEL
16from litellm.types.proxy.management_endpoints.common_daily_activity import (
17 SpendAnalyticsPaginatedResponse,
18)
19from litellm.types.utils import ChatCompletionMessageToolCall
21# ---------------------------------------------------------------------------
22# Constants
23# ---------------------------------------------------------------------------
25USAGE_AI_TEMPERATURE: Final = 0.2
27TABLE_DAILY_USER_SPEND: Final = "litellm_dailyuserspend"
28TABLE_DAILY_TEAM_SPEND: Final = "litellm_dailyteamspend"
29TABLE_DAILY_TAG_SPEND: Final = "litellm_dailytagspend"
31ENTITY_FIELD_USER: Final = "user_id"
32ENTITY_FIELD_TEAM: Final = "team_id"
33ENTITY_FIELD_TAG: Final = "tag"
35PAGINATED_PAGE_SIZE: Final = 200
36MAX_CHAT_MESSAGES: Final = 20
37TOP_N_MODELS: Final = 15
38TOP_N_PROVIDERS: Final = 10
39TOP_N_KEYS: Final = 10
41# ---------------------------------------------------------------------------
42# Types
43# ---------------------------------------------------------------------------
46class SSEStatusEvent(TypedDict):
47 type: Literal["status"]
48 message: str
51class SSEToolCallEvent(TypedDict, total=False):
52 type: Literal["tool_call"]
53 tool_name: str
54 tool_label: str
55 arguments: dict[str, str]
56 status: Literal["running", "complete", "error"]
57 error: str
60class SSEChunkEvent(TypedDict):
61 type: Literal["chunk"]
62 content: str
65class SSEDoneEvent(TypedDict):
66 type: Literal["done"]
69class SSEErrorEvent(TypedDict):
70 type: Literal["error"]
71 message: str
74SSEEvent = SSEStatusEvent | SSEToolCallEvent | SSEChunkEvent | SSEDoneEvent | SSEErrorEvent
77class _EntityEntry(TypedDict, total=False):
78 metrics: ReadOnly[Mapping[str, float]]
79 metadata: ReadOnly[Mapping[str, str]]
82class _DayDump(TypedDict, total=False):
83 breakdown: ReadOnly[Mapping[str, Mapping[str, _EntityEntry]]]
86class _EntityTotal(NamedTuple):
87 """Running per-entity totals accumulated while summarising a usage dump."""
89 alias: str
90 spend: float
91 requests: float
92 tokens: float
95class _UsageDump(Protocol):
96 @overload
97 def get(self, key: Literal["metadata"], default: Mapping[str, float], /) -> Mapping[str, float]: ... 97 ↛ exitline 97 didn't return from function 'get' because
98 @overload
99 def get(self, key: Literal["results"], default: Sequence[_DayDump], /) -> Sequence[_DayDump]: ... 99 ↛ exitline 99 didn't return from function 'get' because
102class _ToolFunctionDef(TypedDict):
103 name: ReadOnly[str]
104 description: ReadOnly[str]
105 parameters: ReadOnly[Mapping[str, object]]
108class _ToolDef(TypedDict):
109 type: ReadOnly[str]
110 function: ReadOnly[_ToolFunctionDef]
113class ToolHandler(TypedDict):
114 fetch: Callable[..., Awaitable[_UsageDump]]
115 summarise: Callable[[_UsageDump], str]
116 label: str
119# ---------------------------------------------------------------------------
120# Tool definitions (OpenAI function-calling schema)
121# ---------------------------------------------------------------------------
123_DATE_PARAMS: Final = {
124 "start_date": {"type": "string", "description": "Start date in YYYY-MM-DD format"},
125 "end_date": {"type": "string", "description": "End date in YYYY-MM-DD format"},
126}
128_TOOL_USAGE: Final[_ToolDef] = {
129 "type": "function",
130 "function": {
131 "name": "get_usage_data",
132 "description": (
133 "Fetch aggregated global usage/spend data. Returns daily spend, "
134 "token counts, request counts, and breakdowns by model, provider, "
135 "and API key. Use for overall spend, top models, top providers."
136 ),
137 "parameters": {
138 "type": "object",
139 "properties": {
140 **_DATE_PARAMS,
141 "user_id": {
142 "type": "string",
143 "description": "Optional user ID filter. Omit for global view.",
144 },
145 },
146 "required": ["start_date", "end_date"],
147 },
148 },
149}
151_TOOL_TEAM: Final[_ToolDef] = {
152 "type": "function",
153 "function": {
154 "name": "get_team_usage_data",
155 "description": (
156 "Fetch usage/spend data broken down by team. Use for questions "
157 "like 'which team spends the most' or 'show me team X usage'."
158 ),
159 "parameters": {
160 "type": "object",
161 "properties": {
162 **_DATE_PARAMS,
163 "team_ids": {
164 "type": "string",
165 "description": "Optional comma-separated team IDs. Omit for all teams.",
166 },
167 },
168 "required": ["start_date", "end_date"],
169 },
170 },
171}
173_TOOL_TAG: Final[_ToolDef] = {
174 "type": "function",
175 "function": {
176 "name": "get_tag_usage_data",
177 "description": (
178 "Fetch usage/spend data broken down by tag. Tags are labels "
179 "attached to requests (features, environments, credentials)."
180 ),
181 "parameters": {
182 "type": "object",
183 "properties": {
184 **_DATE_PARAMS,
185 "tags": {
186 "type": "string",
187 "description": "Optional comma-separated tag names. Omit for all tags.",
188 },
189 },
190 "required": ["start_date", "end_date"],
191 },
192 },
193}
195TOOLS_BASE: Final = [_TOOL_USAGE]
196TOOLS_ADMIN: Final = [_TOOL_USAGE, _TOOL_TEAM, _TOOL_TAG]
199def get_tools_for_role(is_admin: bool) -> list[_ToolDef]:
200 """Return the tool list appropriate for the user's role."""
201 return TOOLS_ADMIN if is_admin else TOOLS_BASE
204_SYSTEM_PROMPT_BASE: Final = (
205 "You are an AI assistant embedded in the LiteLLM Usage dashboard. "
206 "You help users understand their LLM API spend and usage data.\n\n"
207 "ALWAYS call the appropriate tool(s) first to fetch data before answering. "
208 "You may call multiple tools if the question spans different dimensions.\n\n"
209 "Guidelines:\n"
210 "- Be concise and specific. Use exact numbers from the data.\n"
211 "- Format costs as dollar amounts (e.g. $12.34).\n"
212 "- When comparing entities, show a ranked list.\n"
213 "- If data is empty or no results found, say so clearly.\n"
214 "- Do not hallucinate data — only use what the tools return.\n"
215 "- Today's date will be provided below. Use it to interpret relative dates "
216 "like 'this week', 'this month', 'last 7 days', etc."
217)
219_TOOL_DESCRIPTIONS_ADMIN: Final = (
220 "You have access to these tools:\n"
221 "- `get_usage_data`: Global/user-level usage (spend, models, providers, API keys)\n"
222 "- `get_team_usage_data`: Team-level usage breakdown\n"
223 "- `get_tag_usage_data`: Tag-level usage breakdown\n\n"
224)
226_TOOL_DESCRIPTIONS_BASE: Final = (
227 "You have access to this tool:\n- `get_usage_data`: Your usage data (spend, models, providers, API keys)\n\n"
228)
231def _build_system_prompt(is_admin: bool) -> str:
232 """Build role-appropriate system prompt with today's date."""
233 tool_desc: Final = _TOOL_DESCRIPTIONS_ADMIN if is_admin else _TOOL_DESCRIPTIONS_BASE
234 return f"{_SYSTEM_PROMPT_BASE}\n\n{tool_desc}Today's date: {date.today().isoformat()}"
237# keep a public reference for test assertions
238SYSTEM_PROMPT: Final = _SYSTEM_PROMPT_BASE
240# ---------------------------------------------------------------------------
241# Data fetchers
242# ---------------------------------------------------------------------------
245def _parse_csv_ids(raw: str | None) -> list[str] | None:
246 if not raw:
247 return None
248 return [t.strip() for t in raw.split(",") if t.strip()]
251async def _query_activity(
252 table_name: str,
253 entity_id_field: str,
254 entity_id: str | list[str] | None,
255 start_date: str,
256 end_date: str,
257 *,
258 use_aggregated: bool = False,
259) -> SpendAnalyticsPaginatedResponse:
260 """Shared helper that calls the daily activity query layer."""
261 from litellm.proxy.management_endpoints.common_daily_activity import (
262 get_daily_activity,
263 get_daily_activity_aggregated,
264 )
265 from litellm.proxy.proxy_server import prisma_client
267 if use_aggregated:
268 return await get_daily_activity_aggregated(
269 prisma_client=prisma_client,
270 table_name=table_name,
271 entity_id_field=entity_id_field,
272 entity_id=entity_id,
273 entity_metadata_field=None,
274 start_date=start_date,
275 end_date=end_date,
276 model=None,
277 api_key=None,
278 )
279 return await get_daily_activity(
280 prisma_client=prisma_client,
281 table_name=table_name,
282 entity_id_field=entity_id_field,
283 entity_id=entity_id,
284 entity_metadata_field=None,
285 start_date=start_date,
286 end_date=end_date,
287 model=None,
288 api_key=None,
289 page=1,
290 page_size=PAGINATED_PAGE_SIZE,
291 )
294async def _fetch_usage_data(start_date: str, end_date: str, user_id: str | None = None) -> _UsageDump:
295 resp: Final = await _query_activity(
296 TABLE_DAILY_USER_SPEND,
297 ENTITY_FIELD_USER,
298 user_id,
299 start_date,
300 end_date,
301 use_aggregated=True,
302 )
303 return resp.model_dump(mode="json")
306async def _fetch_team_usage_data(start_date: str, end_date: str, team_ids: str | None = None) -> _UsageDump:
307 resp: Final = await _query_activity(
308 TABLE_DAILY_TEAM_SPEND,
309 ENTITY_FIELD_TEAM,
310 _parse_csv_ids(team_ids),
311 start_date,
312 end_date,
313 )
314 return resp.model_dump(mode="json")
317async def _fetch_tag_usage_data(start_date: str, end_date: str, tags: str | None = None) -> _UsageDump:
318 resp: Final = await _query_activity(
319 TABLE_DAILY_TAG_SPEND,
320 ENTITY_FIELD_TAG,
321 _parse_csv_ids(tags),
322 start_date,
323 end_date,
324 )
325 return resp.model_dump(mode="json")
328# ---------------------------------------------------------------------------
329# Summarisers — convert raw JSON to concise text the LLM can reason over
330# ---------------------------------------------------------------------------
333def _accumulate_breakdown(
334 results: Sequence[_DayDump], dimension: str, fields: Sequence[str]
335) -> dict[str, dict[str, float]]:
336 """Aggregate a single breakdown dimension across days."""
337 totals: Final[dict[str, dict[str, float]]] = {}
338 for day in results:
339 for key, entry in day.get("breakdown", {}).get(dimension, {}).items():
340 if key not in totals:
341 totals[key] = {f: 0.0 for f in fields}
342 m = entry.get("metrics", {})
343 for f in fields:
344 totals[key][f] += m.get(f, 0)
345 return totals
348def _ranked_lines(
349 totals: dict[str, dict[str, float]],
350 fmt: Callable[[str, dict[str, float]], str],
351 limit: int,
352) -> list[str]:
353 """Sort by spend descending, format each entry, and truncate."""
354 return [fmt(name, vals) for name, vals in sorted(totals.items(), key=lambda x: -x[1].get("spend", 0))[:limit]]
357def _summarise_usage_data(data: _UsageDump) -> str:
358 meta: Final = data.get("metadata", {})
359 results: Final = data.get("results", [])
361 header: Final = (
362 f"Total Spend: ${meta.get('total_spend', 0):.4f}\n"
363 f"Total Requests: {meta.get('total_api_requests', 0)}\n"
364 f"Successful: {meta.get('total_successful_requests', 0)} | "
365 f"Failed: {meta.get('total_failed_requests', 0)}\n"
366 f"Total Tokens: {meta.get('total_tokens', 0)}"
367 )
369 models: Final = _accumulate_breakdown(results, "models", ["spend", "api_requests", "total_tokens"])
370 providers: Final = _accumulate_breakdown(results, "providers", ["spend", "api_requests"])
372 model_lines: Final = _ranked_lines(
373 models,
374 lambda n, d: f" - {n}: ${d['spend']:.4f} ({int(d['api_requests'])} reqs, {int(d['total_tokens'])} tokens)",
375 TOP_N_MODELS,
376 )
377 provider_lines: Final = _ranked_lines(
378 providers,
379 lambda n, d: f" - {n}: ${d['spend']:.4f} ({int(d['api_requests'])} reqs)",
380 TOP_N_PROVIDERS,
381 )
383 sections = [header, ""]
384 sections += ["Top Models by Spend:"] + (model_lines or [" (no data)"]) + [""]
385 sections += ["Top Providers by Spend:"] + (provider_lines or [" (no data)"])
386 return "\n".join(sections)
389def _summarise_entity_data(data: _UsageDump, entity_label: str) -> str:
390 """Summarise team/tag entity usage data."""
391 results: Final = data.get("results", [])
392 if not results:
393 return f"No {entity_label} usage data found for the given date range."
395 totals: Final[dict[str, _EntityTotal]] = {}
396 for day in results:
397 for eid, entry in day.get("breakdown", {}).get("entities", {}).items():
398 previous = totals.get(eid)
399 m = entry.get("metrics", {})
400 totals[eid] = _EntityTotal(
401 alias=previous.alias if previous is not None else entry.get("metadata", {}).get("alias", eid),
402 spend=(previous.spend if previous is not None else 0.0) + m.get("spend", 0),
403 requests=(previous.requests if previous is not None else 0) + m.get("api_requests", 0),
404 tokens=(previous.tokens if previous is not None else 0) + m.get("total_tokens", 0),
405 )
407 lines: Final = [f"{entity_label} Usage ({len(totals)} {entity_label.lower()}s):", ""]
408 for eid, d in sorted(totals.items(), key=lambda x: -x[1].spend):
409 label = d.alias if d.alias != eid else eid
410 lines.append(f"- {label} (ID: {eid}): ${d.spend:.4f} | {int(d.requests)} reqs | {int(d.tokens)} tokens")
411 return "\n".join(lines)
414# ---------------------------------------------------------------------------
415# Tool dispatch registry
416# ---------------------------------------------------------------------------
418TOOL_HANDLERS: Final[dict[str, ToolHandler]] = {
419 "get_usage_data": ToolHandler(
420 fetch=_fetch_usage_data,
421 summarise=_summarise_usage_data,
422 label="global usage data",
423 ),
424 "get_team_usage_data": ToolHandler(
425 fetch=_fetch_team_usage_data,
426 summarise=lambda data: _summarise_entity_data(data, "Team"),
427 label="team usage data",
428 ),
429 "get_tag_usage_data": ToolHandler(
430 fetch=_fetch_tag_usage_data,
431 summarise=lambda data: _summarise_entity_data(data, "Tag"),
432 label="tag usage data",
433 ),
434}
437# ---------------------------------------------------------------------------
438# SSE streaming
439# ---------------------------------------------------------------------------
442def _sse(event: SSEEvent) -> str:
443 return f"data: {json.dumps(event)}\n\n"
446def _resolve_fetch_kwargs(
447 fn_name: str,
448 fn_args: Mapping[str, str],
449 user_id: str | None,
450 is_admin: bool,
451) -> dict[str, str]:
452 """Build keyword arguments for a tool's fetch function."""
453 start_date: Final = fn_args.get("start_date", "")
454 end_date: Final = fn_args.get("end_date", "")
455 if not start_date or not end_date:
456 raise ValueError("Missing required start_date or end_date from tool arguments")
457 kwargs: Final[dict[str, str]] = {"start_date": start_date, "end_date": end_date}
458 if fn_name == "get_usage_data":
459 if not is_admin:
460 if user_id is None:
461 # Defense-in-depth: the endpoint guard in usage_endpoints/endpoints.py
462 # should have already rejected this. If we ever reach here it means
463 # a future caller invoked the helper without scoping — fail loudly
464 # rather than issuing an unfiltered global query.
465 raise ValueError(
466 "Non-admin caller has user_id=None; refusing to issue an "
467 "unscoped query. Endpoint-level guard missing."
468 )
469 kwargs["user_id"] = user_id
470 elif fn_args.get("user_id"):
471 kwargs["user_id"] = fn_args["user_id"]
472 elif fn_name == "get_team_usage_data" and fn_args.get("team_ids"):
473 kwargs["team_ids"] = fn_args["team_ids"]
474 elif fn_name == "get_tag_usage_data" and fn_args.get("tags"):
475 kwargs["tags"] = fn_args["tags"]
476 return kwargs
479async def _execute_tool_call(
480 handler: ToolHandler,
481 fn_name: str,
482 fn_args: Mapping[str, str],
483 user_id: str | None,
484 is_admin: bool,
485) -> str:
486 """Run a single tool and return the summarised result text."""
487 kwargs: Final = _resolve_fetch_kwargs(fn_name, fn_args, user_id, is_admin)
488 raw_data: Final = await handler["fetch"](**kwargs)
489 return handler["summarise"](raw_data)
492async def _process_tool_call(
493 tc: ChatCompletionMessageToolCall,
494 chat_messages: list[Mapping[str, object]],
495 user_id: str | None,
496 is_admin: bool,
497) -> AsyncIterator[str]:
498 """Execute a single tool call, yielding SSE events for status."""
499 fn_name: Final = tc.function.name
500 fn_args: Final[Mapping[str, str]] = json.loads(tc.function.arguments)
502 allowed_names: Final = {t["function"]["name"] for t in get_tools_for_role(is_admin)}
503 handler: Final = TOOL_HANDLERS.get(fn_name) if fn_name is not None else None
505 if fn_name is None or fn_name not in allowed_names or not handler:
506 chat_messages.append(
507 {
508 "role": "tool",
509 "tool_call_id": tc.id,
510 "content": f"Tool not available: {fn_name}",
511 }
512 )
513 return
515 tool_event_base: Final = {
516 "type": "tool_call",
517 "tool_name": fn_name,
518 "tool_label": handler["label"],
519 "arguments": fn_args,
520 }
521 yield _sse(cast(SSEToolCallEvent, {**tool_event_base, "status": "running"}))
523 try:
524 tool_result = await _execute_tool_call(handler, fn_name, fn_args, user_id, is_admin)
525 yield _sse(cast(SSEToolCallEvent, {**tool_event_base, "status": "complete"}))
526 except Exception as e:
527 verbose_proxy_logger.error("Tool %s failed: %s", fn_name, e)
528 tool_result = f"Error fetching {handler['label']}. Please try again."
529 yield _sse(cast(SSEToolCallEvent, {**tool_event_base, "status": "error"}))
531 chat_messages.append({"role": "tool", "tool_call_id": tc.id, "content": tool_result})
534async def _stream_final_response(model: str, chat_messages: list[Mapping[str, object]]) -> AsyncIterator[str]:
535 """Stream the final LLM response after tool results are appended."""
536 yield _sse({"type": "status", "message": "Analyzing results..."})
538 response: Final = await litellm.acompletion(
539 model=model,
540 messages=chat_messages,
541 stream=True,
542 temperature=USAGE_AI_TEMPERATURE,
543 )
544 async for chunk in response:
545 delta = chunk.choices[0].delta.content
546 if delta:
547 yield _sse({"type": "chunk", "content": delta})
550async def stream_usage_ai_chat(
551 messages: list[dict[str, str]],
552 model: str | None = None,
553 user_id: str | None = None,
554 is_admin: bool = False,
555) -> AsyncGenerator[str, None]:
556 """Stream SSE events: status → tool_call → chunk → done."""
557 resolved_model: Final = (model or "").strip() or DEFAULT_COMPETITOR_DISCOVERY_MODEL
558 truncated: Final = messages[-MAX_CHAT_MESSAGES:] if len(messages) > MAX_CHAT_MESSAGES else messages
559 chat_messages: Final[list[Mapping[str, object]]] = [
560 {"role": "system", "content": _build_system_prompt(is_admin)},
561 *truncated,
562 ]
564 try:
565 yield _sse({"type": "status", "message": "Thinking..."})
566 tools: Final = get_tools_for_role(is_admin)
567 response: Final = await litellm.acompletion(
568 model=resolved_model,
569 messages=chat_messages,
570 tools=tools,
571 temperature=USAGE_AI_TEMPERATURE,
572 )
573 choice: Final = response.choices[0]
575 if not choice.message.tool_calls:
576 if choice.message.content:
577 yield _sse({"type": "chunk", "content": choice.message.content})
578 yield _sse({"type": "done"})
579 return
581 chat_messages.append(choice.message.model_dump())
582 for tc in choice.message.tool_calls:
583 async for event in _process_tool_call(tc, chat_messages, user_id, is_admin):
584 yield event
585 async for event in _stream_final_response(resolved_model, chat_messages):
586 yield event
587 yield _sse({"type": "done"})
589 except Exception as e:
590 verbose_proxy_logger.error("AI usage chat failed: %s", e)
591 yield _sse(
592 {
593 "type": "error",
594 "message": "An internal error occurred. Please try again.",
595 }
596 )