diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 27cdc483d4a..0957f394698 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -27633,7 +27633,7 @@ "type": "null" } ], - "description": "Model to use for AI chat", + "description": "Model group to use for AI chat", "title": "Model" } }, @@ -27681,7 +27681,7 @@ "paths": { "/usage/ai/chat": { "post": { - "description": "AI chat about usage data. Streams SSE events with the AI response.\nThe AI agent has access to tools that query aggregated daily activity data.", + "description": "AI chat about usage data. Streams SSE events with the AI response.\n\nThe agent queries aggregated daily activity data through a provider scoped\nto the caller: admins get a global view, non-admins are restricted to their\nown ``user_id``.", "operationId": "usage_ai_chat_usage_ai_chat_post", "requestBody": { "content": { diff --git a/litellm/proxy/management_endpoints/usage_endpoints/agent.py b/litellm/proxy/management_endpoints/usage_endpoints/agent.py new file mode 100644 index 00000000000..50a91430674 --- /dev/null +++ b/litellm/proxy/management_endpoints/usage_endpoints/agent.py @@ -0,0 +1,401 @@ +"""Ask AI agent: answers usage/spend questions by driving an LLM tool-calling +loop over the scoped usage data provider. + +The LLM call goes through the proxy's own ``llm_router`` (not the bare +``litellm`` SDK), so UI-selected model groups / aliases resolve exactly as they +do on ``/chat/completions``, and the call is credentialed, logged, budgeted, +rate-limited, and guardrailed like any other proxy request. +""" + +import json +from dataclasses import dataclass +from datetime import date +from typing import Any, AsyncIterator, Dict, List, Literal, Optional, Set, Union, cast + +from pydantic import BaseModel, TypeAdapter +from typing_extensions import TypedDict, assert_never + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.router import Router +from litellm.types.llms.openai import AllMessageValues +from litellm.types.utils import ( + ChatCompletionMessageToolCall, + Choices, + Message, + ModelResponse, +) +from litellm.proxy.management_endpoints.usage_endpoints.scoped_data import ( + ScopedUsageDataProvider, + summarise_entity_data, + summarise_usage_data, +) + +USAGE_AI_TEMPERATURE = 0.2 +MAX_CHAT_MESSAGES = 20 +MAX_TOOL_ROUNDS = 5 +USAGE_AI_MODEL_SETTING = "usage_ai_model" + + +# --------------------------------------------------------------------------- +# SSE wire events (kept identical to the v1 contract the frontend consumes) +# --------------------------------------------------------------------------- + + +class SSEStatusEvent(TypedDict): + type: Literal["status"] + message: str + + +class SSEToolCallEvent(TypedDict, total=False): + type: Literal["tool_call"] + tool_name: str + tool_label: str + arguments: Dict[str, str] + status: Literal["running", "complete", "error"] + + +class SSEChunkEvent(TypedDict): + type: Literal["chunk"] + content: str + + +class SSEDoneEvent(TypedDict): + type: Literal["done"] + + +class SSEErrorEvent(TypedDict): + type: Literal["error"] + message: str + + +SSEEvent = Union[SSEStatusEvent, SSEToolCallEvent, SSEChunkEvent, SSEDoneEvent, SSEErrorEvent] + + +def _sse(event: SSEEvent) -> str: + return f"data: {json.dumps(event)}\n\n" + + +# --------------------------------------------------------------------------- +# Terminal errors modelled as values, mapped to the SSE error contract once +# --------------------------------------------------------------------------- + + +@dataclass(frozen=True, slots=True) +class ModelNotConfigured: + pass + + +@dataclass(frozen=True, slots=True) +class RouterUnavailable: + pass + + +@dataclass(frozen=True, slots=True) +class LLMCallError: + detail: str + + +UsageAiError = Union[ModelNotConfigured, RouterUnavailable, LLMCallError] + + +def _error_event(err: UsageAiError) -> SSEErrorEvent: + match err: + case ModelNotConfigured(): + message = ( + "No model is configured for Ask AI. Pick a model in the selector, " + f"or set '{USAGE_AI_MODEL_SETTING}' under general_settings." + ) + case RouterUnavailable(): + message = "The proxy has no model router initialized yet. Try again in a moment." + case LLMCallError(): + message = "The AI request failed. Confirm the selected model is configured on this proxy and reachable." + case _: + assert_never(err) + return {"type": "error", "message": message} + + +class _RouterUnavailableError(Exception): + pass + + +def _require_router() -> Router: + from litellm.proxy.proxy_server import llm_router + + if llm_router is None: + raise _RouterUnavailableError() + return llm_router + + +def _assembled_message(chunks: List[object]) -> Optional[Message]: + """Reassemble streamed chunks into a single message (content + tool_calls).""" + built = litellm.stream_chunk_builder(chunks) + if not isinstance(built, ModelResponse) or not built.choices: + return None + choice = built.choices[0] + return choice.message if isinstance(choice, Choices) else None # pyright: ignore[reportUnnecessaryIsInstance] # choices[0] can be StreamingChoices at runtime + + +def resolve_model(requested: Optional[str]) -> Union[str, ModelNotConfigured]: + """Resolve the model group to use: explicit request wins, then the + configured ``usage_ai_model`` setting, else an actionable error value.""" + explicit = (requested or "").strip() + if explicit: + return explicit + + from litellm.proxy.proxy_server import general_settings + + configured = TypeAdapter(Optional[str]).validate_python(general_settings.get(USAGE_AI_MODEL_SETTING)) + stripped = (configured or "").strip() + return stripped or ModelNotConfigured() + + +# --------------------------------------------------------------------------- +# Tools (OpenAI function-calling schema) + typed argument validation +# --------------------------------------------------------------------------- + +_DATE_PARAMS = { + "start_date": {"type": "string", "description": "Start date in YYYY-MM-DD format"}, + "end_date": {"type": "string", "description": "End date in YYYY-MM-DD format"}, +} + +_TOOL_USAGE = { + "type": "function", + "function": { + "name": "get_usage_data", + "description": ( + "Fetch aggregated usage/spend data. Returns daily spend, token counts, " + "request counts, and breakdowns by model and provider. Use for overall " + "spend, top models, and top providers." + ), + "parameters": { + "type": "object", + "properties": { + **_DATE_PARAMS, + "user_id": { + "type": "string", + "description": "Optional user ID filter (admin only). Omit for global view.", + }, + }, + "required": ["start_date", "end_date"], + }, + }, +} + +_TOOL_TEAM = { + "type": "function", + "function": { + "name": "get_team_usage_data", + "description": ( + "Fetch usage/spend data broken down by team. Use for questions like " + "'which team spends the most' or 'show me team X usage'." + ), + "parameters": { + "type": "object", + "properties": { + **_DATE_PARAMS, + "team_ids": {"type": "string", "description": "Optional comma-separated team IDs. Omit for all teams."}, + }, + "required": ["start_date", "end_date"], + }, + }, +} + +_TOOL_TAG = { + "type": "function", + "function": { + "name": "get_tag_usage_data", + "description": ( + "Fetch usage/spend data broken down by tag. Tags are labels attached to " + "requests (features, environments, credentials)." + ), + "parameters": { + "type": "object", + "properties": { + **_DATE_PARAMS, + "tags": {"type": "string", "description": "Optional comma-separated tag names. Omit for all tags."}, + }, + "required": ["start_date", "end_date"], + }, + }, +} + + +def tools_for_role(is_admin: bool) -> List[Dict[str, Any]]: + return [_TOOL_USAGE, _TOOL_TEAM, _TOOL_TAG] if is_admin else [_TOOL_USAGE] + + +_TOOL_LABELS = { + "get_usage_data": "global usage data", + "get_team_usage_data": "team usage data", + "get_tag_usage_data": "tag usage data", +} + + +class _UsageArgs(BaseModel): + start_date: str + end_date: str + user_id: Optional[str] = None + + +class _TeamArgs(BaseModel): + start_date: str + end_date: str + team_ids: Optional[str] = None + + +class _TagArgs(BaseModel): + start_date: str + end_date: str + tags: Optional[str] = None + + +async def _dispatch_tool(name: str, raw_args: Dict[str, Any], provider: ScopedUsageDataProvider) -> str: + if name == "get_usage_data": + args = _UsageArgs.model_validate(raw_args) + data = await provider.usage(args.start_date, args.end_date, args.user_id) + return summarise_usage_data(data) + if name == "get_team_usage_data": + team_args = _TeamArgs.model_validate(raw_args) + team_data = await provider.team(team_args.start_date, team_args.end_date, team_args.team_ids) + return summarise_entity_data(team_data, "Team") + if name == "get_tag_usage_data": + tag_args = _TagArgs.model_validate(raw_args) + tag_data = await provider.tag(tag_args.start_date, tag_args.end_date, tag_args.tags) + return summarise_entity_data(tag_data, "Tag") + raise ValueError(f"Unknown tool: {name}") + + +# --------------------------------------------------------------------------- +# System prompt +# --------------------------------------------------------------------------- + +_SYSTEM_PROMPT_BASE = ( + "You are an AI assistant embedded in the LiteLLM Usage dashboard. " + "You help users understand their LLM API spend and usage data.\n\n" + "ALWAYS call the appropriate tool(s) first to fetch data before answering. " + "You may call multiple tools, across multiple turns, if the question spans " + "different dimensions or needs follow-up lookups.\n\n" + "Guidelines:\n" + "- Be concise and specific. Use exact numbers from the data.\n" + "- Format costs as dollar amounts (e.g. $12.34).\n" + "- When comparing entities, show a ranked list.\n" + "- If data is empty or no results found, say so clearly.\n" + "- Do not hallucinate data; only use what the tools return.\n" + "- Today's date is provided below; use it to interpret relative dates like " + "'this week', 'this month', or 'last 7 days'." +) + +_TOOL_DESCRIPTIONS_ADMIN = ( + "You have access to these tools:\n" + "- `get_usage_data`: Global/user-level usage (spend, models, providers)\n" + "- `get_team_usage_data`: Team-level usage breakdown\n" + "- `get_tag_usage_data`: Tag-level usage breakdown\n\n" +) + +_TOOL_DESCRIPTIONS_BASE = ( + "You have access to this tool:\n- `get_usage_data`: Your usage data (spend, models, providers)\n\n" +) + + +def _system_prompt(is_admin: bool) -> str: + tool_desc = _TOOL_DESCRIPTIONS_ADMIN if is_admin else _TOOL_DESCRIPTIONS_BASE + return f"{_SYSTEM_PROMPT_BASE}\n\n{tool_desc}Today's date: {date.today().isoformat()}" + + +# --------------------------------------------------------------------------- +# Streaming agent loop +# --------------------------------------------------------------------------- + + +async def _run_tool_call( + tc: ChatCompletionMessageToolCall, + provider: ScopedUsageDataProvider, + allowed_names: Set[str], + convo: List[Dict[str, Any]], +) -> AsyncIterator[str]: + """Execute one tool call, yield status events, and append its result to convo.""" + name = tc.function.name + try: + parsed = json.loads(tc.function.arguments or "{}") + except json.JSONDecodeError: + parsed = {} + raw_args = parsed if isinstance(parsed, dict) else {} + + if name not in allowed_names: + convo.append({"role": "tool", "tool_call_id": tc.id, "content": f"Tool not available: {name}"}) + return + + label = _TOOL_LABELS.get(name, name) + base: Dict[str, Any] = {"type": "tool_call", "tool_name": name, "tool_label": label, "arguments": raw_args} + yield _sse(cast(SSEToolCallEvent, {**base, "status": "running"})) + + try: + result = await _dispatch_tool(name, raw_args, provider) + yield _sse(cast(SSEToolCallEvent, {**base, "status": "complete"})) + except Exception as e: + verbose_proxy_logger.error("Usage AI tool %s failed: %s", name, e) + result = f"Error fetching {label}. Please try again." + yield _sse(cast(SSEToolCallEvent, {**base, "status": "error"})) + + convo.append({"role": "tool", "tool_call_id": tc.id, "content": result}) + + +async def stream_usage_ai_chat( + provider: ScopedUsageDataProvider, + messages: List[Dict[str, str]], + model: Optional[str] = None, +) -> AsyncIterator[str]: + """Stream SSE events: status -> tool_call -> chunk -> done (or a single error).""" + resolved = resolve_model(model) + if isinstance(resolved, ModelNotConfigured): + yield _sse(_error_event(resolved)) + return + + tools = tools_for_role(provider.is_admin) + allowed_names = {t["function"]["name"] for t in tools} + history = messages[-MAX_CHAT_MESSAGES:] + convo: List[Dict[str, Any]] = [{"role": "system", "content": _system_prompt(provider.is_admin)}, *history] + + try: + router = _require_router() + yield _sse({"type": "status", "message": "Thinking..."}) + + for round_index in range(MAX_TOOL_ROUNDS + 1): + use_tools = tools if round_index < MAX_TOOL_ROUNDS else None + chunks: List[object] = [] + response = await router.acompletion( + model=resolved, + messages=cast(List[AllMessageValues], convo), + tools=use_tools, + stream=True, + temperature=USAGE_AI_TEMPERATURE, + metadata={"feature": "usage_ai"}, + ) + async for chunk in response: + choices = getattr(chunk, "choices", None) + delta = choices[0].delta if choices else None + content = getattr(delta, "content", None) if delta is not None else None + if content: + yield _sse({"type": "chunk", "content": content}) + chunks.append(chunk) + + message = _assembled_message(chunks) + tool_calls = message.tool_calls if message is not None else None + + if not message or not tool_calls: + yield _sse({"type": "done"}) + return + + convo.append(message.model_dump()) + for tc in tool_calls: + async for event in _run_tool_call(tc, provider, allowed_names, convo): + yield event + + yield _sse({"type": "done"}) + + except _RouterUnavailableError: + yield _sse(_error_event(RouterUnavailable())) + except Exception as e: + verbose_proxy_logger.error("Usage AI chat failed: %s", e) + yield _sse(_error_event(LLMCallError(detail=str(e)))) diff --git a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py deleted file mode 100644 index efb714e4da0..00000000000 --- a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py +++ /dev/null @@ -1,559 +0,0 @@ -""" -AI Usage Chat - uses LLM tool calling to answer questions about -usage/spend data by querying the aggregated daily activity endpoints. -""" - -import json -from datetime import date -from typing import Any, AsyncIterator, Callable, Dict, List, Literal, Optional, cast - -from typing_extensions import TypedDict - -import litellm -from litellm._logging import verbose_proxy_logger -from litellm.constants import DEFAULT_COMPETITOR_DISCOVERY_MODEL -from litellm.types.proxy.management_endpoints.common_daily_activity import ( - SpendAnalyticsPaginatedResponse, -) - -# --------------------------------------------------------------------------- -# Constants -# --------------------------------------------------------------------------- - -USAGE_AI_TEMPERATURE = 0.2 - -TABLE_DAILY_USER_SPEND = "litellm_dailyuserspend" -TABLE_DAILY_TEAM_SPEND = "litellm_dailyteamspend" -TABLE_DAILY_TAG_SPEND = "litellm_dailytagspend" - -ENTITY_FIELD_USER = "user_id" -ENTITY_FIELD_TEAM = "team_id" -ENTITY_FIELD_TAG = "tag" - -PAGINATED_PAGE_SIZE = 200 -MAX_CHAT_MESSAGES = 20 -TOP_N_MODELS = 15 -TOP_N_PROVIDERS = 10 -TOP_N_KEYS = 10 - -# --------------------------------------------------------------------------- -# Types -# --------------------------------------------------------------------------- - - -class SSEStatusEvent(TypedDict): - type: Literal["status"] - message: str - - -class SSEToolCallEvent(TypedDict, total=False): - type: Literal["tool_call"] - tool_name: str - tool_label: str - arguments: Dict[str, str] - status: Literal["running", "complete", "error"] - error: str - - -class SSEChunkEvent(TypedDict): - type: Literal["chunk"] - content: str - - -class SSEDoneEvent(TypedDict): - type: Literal["done"] - - -class SSEErrorEvent(TypedDict): - type: Literal["error"] - message: str - - -SSEEvent = SSEStatusEvent | SSEToolCallEvent | SSEChunkEvent | SSEDoneEvent | SSEErrorEvent - - -class ToolHandler(TypedDict): - fetch: Callable[..., Any] - summarise: Callable[[Dict[str, Any]], str] - label: str - - -# --------------------------------------------------------------------------- -# Tool definitions (OpenAI function-calling schema) -# --------------------------------------------------------------------------- - -_DATE_PARAMS = { - "start_date": {"type": "string", "description": "Start date in YYYY-MM-DD format"}, - "end_date": {"type": "string", "description": "End date in YYYY-MM-DD format"}, -} - -_TOOL_USAGE = { - "type": "function", - "function": { - "name": "get_usage_data", - "description": ( - "Fetch aggregated global usage/spend data. Returns daily spend, " - "token counts, request counts, and breakdowns by model, provider, " - "and API key. Use for overall spend, top models, top providers." - ), - "parameters": { - "type": "object", - "properties": { - **_DATE_PARAMS, - "user_id": { - "type": "string", - "description": "Optional user ID filter. Omit for global view.", - }, - }, - "required": ["start_date", "end_date"], - }, - }, -} - -_TOOL_TEAM = { - "type": "function", - "function": { - "name": "get_team_usage_data", - "description": ( - "Fetch usage/spend data broken down by team. Use for questions " - "like 'which team spends the most' or 'show me team X usage'." - ), - "parameters": { - "type": "object", - "properties": { - **_DATE_PARAMS, - "team_ids": { - "type": "string", - "description": "Optional comma-separated team IDs. Omit for all teams.", - }, - }, - "required": ["start_date", "end_date"], - }, - }, -} - -_TOOL_TAG = { - "type": "function", - "function": { - "name": "get_tag_usage_data", - "description": ( - "Fetch usage/spend data broken down by tag. Tags are labels " - "attached to requests (features, environments, credentials)." - ), - "parameters": { - "type": "object", - "properties": { - **_DATE_PARAMS, - "tags": { - "type": "string", - "description": "Optional comma-separated tag names. Omit for all tags.", - }, - }, - "required": ["start_date", "end_date"], - }, - }, -} - -TOOLS_BASE = [_TOOL_USAGE] -TOOLS_ADMIN = [_TOOL_USAGE, _TOOL_TEAM, _TOOL_TAG] - - -def get_tools_for_role(is_admin: bool) -> List[Dict[str, Any]]: - """Return the tool list appropriate for the user's role.""" - return TOOLS_ADMIN if is_admin else TOOLS_BASE - - -_SYSTEM_PROMPT_BASE = ( - "You are an AI assistant embedded in the LiteLLM Usage dashboard. " - "You help users understand their LLM API spend and usage data.\n\n" - "ALWAYS call the appropriate tool(s) first to fetch data before answering. " - "You may call multiple tools if the question spans different dimensions.\n\n" - "Guidelines:\n" - "- Be concise and specific. Use exact numbers from the data.\n" - "- Format costs as dollar amounts (e.g. $12.34).\n" - "- When comparing entities, show a ranked list.\n" - "- If data is empty or no results found, say so clearly.\n" - "- Do not hallucinate data — only use what the tools return.\n" - "- Today's date will be provided below. Use it to interpret relative dates " - "like 'this week', 'this month', 'last 7 days', etc." -) - -_TOOL_DESCRIPTIONS_ADMIN = ( - "You have access to these tools:\n" - "- `get_usage_data`: Global/user-level usage (spend, models, providers, API keys)\n" - "- `get_team_usage_data`: Team-level usage breakdown\n" - "- `get_tag_usage_data`: Tag-level usage breakdown\n\n" -) - -_TOOL_DESCRIPTIONS_BASE = ( - "You have access to this tool:\n- `get_usage_data`: Your usage data (spend, models, providers, API keys)\n\n" -) - - -def _build_system_prompt(is_admin: bool) -> str: - """Build role-appropriate system prompt with today's date.""" - tool_desc = _TOOL_DESCRIPTIONS_ADMIN if is_admin else _TOOL_DESCRIPTIONS_BASE - return f"{_SYSTEM_PROMPT_BASE}\n\n{tool_desc}Today's date: {date.today().isoformat()}" - - -# keep a public reference for test assertions -SYSTEM_PROMPT = _SYSTEM_PROMPT_BASE - -# --------------------------------------------------------------------------- -# Data fetchers -# --------------------------------------------------------------------------- - - -def _parse_csv_ids(raw: Optional[str]) -> Optional[List[str]]: - if not raw: - return None - return [t.strip() for t in raw.split(",") if t.strip()] - - -async def _query_activity( - table_name: str, - entity_id_field: str, - entity_id: Optional[Any], - start_date: str, - end_date: str, - *, - use_aggregated: bool = False, -) -> SpendAnalyticsPaginatedResponse: - """Shared helper that calls the daily activity query layer.""" - from litellm.proxy.management_endpoints.common_daily_activity import ( - get_daily_activity, - get_daily_activity_aggregated, - ) - from litellm.proxy.proxy_server import prisma_client - - if use_aggregated: - return await get_daily_activity_aggregated( - prisma_client=prisma_client, - table_name=table_name, - entity_id_field=entity_id_field, - entity_id=entity_id, - entity_metadata_field=None, - start_date=start_date, - end_date=end_date, - model=None, - api_key=None, - ) - return await get_daily_activity( - prisma_client=prisma_client, - table_name=table_name, - entity_id_field=entity_id_field, - entity_id=entity_id, - entity_metadata_field=None, - start_date=start_date, - end_date=end_date, - model=None, - api_key=None, - page=1, - page_size=PAGINATED_PAGE_SIZE, - ) - - -async def _fetch_usage_data(start_date: str, end_date: str, user_id: Optional[str] = None) -> Dict[str, Any]: - resp = await _query_activity( - TABLE_DAILY_USER_SPEND, - ENTITY_FIELD_USER, - user_id, - start_date, - end_date, - use_aggregated=True, - ) - return resp.model_dump(mode="json") - - -async def _fetch_team_usage_data(start_date: str, end_date: str, team_ids: Optional[str] = None) -> Dict[str, Any]: - resp = await _query_activity( - TABLE_DAILY_TEAM_SPEND, - ENTITY_FIELD_TEAM, - _parse_csv_ids(team_ids), - start_date, - end_date, - ) - return resp.model_dump(mode="json") - - -async def _fetch_tag_usage_data(start_date: str, end_date: str, tags: Optional[str] = None) -> Dict[str, Any]: - resp = await _query_activity( - TABLE_DAILY_TAG_SPEND, - ENTITY_FIELD_TAG, - _parse_csv_ids(tags), - start_date, - end_date, - ) - return resp.model_dump(mode="json") - - -# --------------------------------------------------------------------------- -# Summarisers — convert raw JSON to concise text the LLM can reason over -# --------------------------------------------------------------------------- - - -def _accumulate_breakdown( - results: List[Dict[str, Any]], dimension: str, fields: List[str] -) -> Dict[str, Dict[str, float]]: - """Aggregate a single breakdown dimension across days.""" - totals: Dict[str, Dict[str, float]] = {} - for day in results: - for key, entry in day.get("breakdown", {}).get(dimension, {}).items(): - if key not in totals: - totals[key] = {f: 0.0 for f in fields} - m = entry.get("metrics", {}) - for f in fields: - totals[key][f] += m.get(f, 0) - return totals - - -def _ranked_lines( - totals: Dict[str, Dict[str, float]], - fmt: Callable[[str, Dict[str, float]], str], - limit: int, -) -> List[str]: - """Sort by spend descending, format each entry, and truncate.""" - return [fmt(name, vals) for name, vals in sorted(totals.items(), key=lambda x: -x[1].get("spend", 0))[:limit]] - - -def _summarise_usage_data(data: Dict[str, Any]) -> str: - meta = data.get("metadata", {}) - results = data.get("results", []) - - header = ( - f"Total Spend: ${meta.get('total_spend', 0):.4f}\n" - f"Total Requests: {meta.get('total_api_requests', 0)}\n" - f"Successful: {meta.get('total_successful_requests', 0)} | " - f"Failed: {meta.get('total_failed_requests', 0)}\n" - f"Total Tokens: {meta.get('total_tokens', 0)}" - ) - - models = _accumulate_breakdown(results, "models", ["spend", "api_requests", "total_tokens"]) - providers = _accumulate_breakdown(results, "providers", ["spend", "api_requests"]) - - model_lines = _ranked_lines( - models, - lambda n, d: f" - {n}: ${d['spend']:.4f} ({int(d['api_requests'])} reqs, {int(d['total_tokens'])} tokens)", - TOP_N_MODELS, - ) - provider_lines = _ranked_lines( - providers, - lambda n, d: f" - {n}: ${d['spend']:.4f} ({int(d['api_requests'])} reqs)", - TOP_N_PROVIDERS, - ) - - sections = [header, ""] - sections += ["Top Models by Spend:"] + (model_lines or [" (no data)"]) + [""] - sections += ["Top Providers by Spend:"] + (provider_lines or [" (no data)"]) - return "\n".join(sections) - - -def _summarise_entity_data(data: Dict[str, Any], entity_label: str) -> str: - """Summarise team/tag entity usage data.""" - results = data.get("results", []) - if not results: - return f"No {entity_label} usage data found for the given date range." - - totals: Dict[str, Dict[str, Any]] = {} - for day in results: - for eid, entry in day.get("breakdown", {}).get("entities", {}).items(): - if eid not in totals: - alias = entry.get("metadata", {}).get("alias", eid) - totals[eid] = {"alias": alias, "spend": 0.0, "requests": 0, "tokens": 0} - m = entry.get("metrics", {}) - totals[eid]["spend"] += m.get("spend", 0) - totals[eid]["requests"] += m.get("api_requests", 0) - totals[eid]["tokens"] += m.get("total_tokens", 0) - - lines = [f"{entity_label} Usage ({len(totals)} {entity_label.lower()}s):", ""] - for eid, d in sorted(totals.items(), key=lambda x: -x[1]["spend"]): - label = d["alias"] if d["alias"] != eid else eid - lines.append( - f"- {label} (ID: {eid}): ${d['spend']:.4f} | {int(d['requests'])} reqs | {int(d['tokens'])} tokens" - ) - return "\n".join(lines) - - -# --------------------------------------------------------------------------- -# Tool dispatch registry -# --------------------------------------------------------------------------- - -TOOL_HANDLERS: Dict[str, ToolHandler] = { - "get_usage_data": ToolHandler( - fetch=_fetch_usage_data, - summarise=_summarise_usage_data, - label="global usage data", - ), - "get_team_usage_data": ToolHandler( - fetch=_fetch_team_usage_data, - summarise=lambda data: _summarise_entity_data(data, "Team"), - label="team usage data", - ), - "get_tag_usage_data": ToolHandler( - fetch=_fetch_tag_usage_data, - summarise=lambda data: _summarise_entity_data(data, "Tag"), - label="tag usage data", - ), -} - - -# --------------------------------------------------------------------------- -# SSE streaming -# --------------------------------------------------------------------------- - - -def _sse(event: SSEEvent) -> str: - return f"data: {json.dumps(event)}\n\n" - - -def _resolve_fetch_kwargs( - fn_name: str, - fn_args: Dict[str, str], - user_id: Optional[str], - is_admin: bool, -) -> Dict[str, Any]: - """Build keyword arguments for a tool's fetch function.""" - start_date = fn_args.get("start_date", "") - end_date = fn_args.get("end_date", "") - if not start_date or not end_date: - raise ValueError("Missing required start_date or end_date from tool arguments") - kwargs: Dict[str, Any] = {"start_date": start_date, "end_date": end_date} - if fn_name == "get_usage_data": - if not is_admin: - if user_id is None: - # Defense-in-depth: the endpoint guard in usage_endpoints/endpoints.py - # should have already rejected this. If we ever reach here it means - # a future caller invoked the helper without scoping — fail loudly - # rather than issuing an unfiltered global query. - raise ValueError( - "Non-admin caller has user_id=None; refusing to issue an " - "unscoped query. Endpoint-level guard missing." - ) - kwargs["user_id"] = user_id - elif fn_args.get("user_id"): - kwargs["user_id"] = fn_args["user_id"] - elif fn_name == "get_team_usage_data" and fn_args.get("team_ids"): - kwargs["team_ids"] = fn_args["team_ids"] - elif fn_name == "get_tag_usage_data" and fn_args.get("tags"): - kwargs["tags"] = fn_args["tags"] - return kwargs - - -async def _execute_tool_call( - handler: ToolHandler, - fn_name: str, - fn_args: Dict[str, str], - user_id: Optional[str], - is_admin: bool, -) -> str: - """Run a single tool and return the summarised result text.""" - kwargs = _resolve_fetch_kwargs(fn_name, fn_args, user_id, is_admin) - raw_data = await handler["fetch"](**kwargs) - return handler["summarise"](raw_data) - - -async def _process_tool_call( - tc: Any, - chat_messages: List[Dict[str, Any]], - user_id: Optional[str], - is_admin: bool, -) -> AsyncIterator[str]: - """Execute a single tool call, yielding SSE events for status.""" - fn_name = tc.function.name - fn_args = json.loads(tc.function.arguments) - - allowed_names = {t["function"]["name"] for t in get_tools_for_role(is_admin)} - handler = TOOL_HANDLERS.get(fn_name) - - if fn_name not in allowed_names or not handler: - chat_messages.append( - { - "role": "tool", - "tool_call_id": tc.id, - "content": f"Tool not available: {fn_name}", - } - ) - return - - tool_event_base = { - "type": "tool_call", - "tool_name": fn_name, - "tool_label": handler["label"], - "arguments": fn_args, - } - yield _sse(cast(SSEToolCallEvent, {**tool_event_base, "status": "running"})) - - try: - tool_result = await _execute_tool_call(handler, fn_name, fn_args, user_id, is_admin) - yield _sse(cast(SSEToolCallEvent, {**tool_event_base, "status": "complete"})) - except Exception as e: - verbose_proxy_logger.error("Tool %s failed: %s", fn_name, e) - tool_result = f"Error fetching {handler['label']}. Please try again." - yield _sse(cast(SSEToolCallEvent, {**tool_event_base, "status": "error"})) - - chat_messages.append({"role": "tool", "tool_call_id": tc.id, "content": tool_result}) - - -async def _stream_final_response(model: str, chat_messages: List[Dict[str, Any]]) -> AsyncIterator[str]: - """Stream the final LLM response after tool results are appended.""" - yield _sse({"type": "status", "message": "Analyzing results..."}) - - response = await litellm.acompletion( - model=model, - messages=chat_messages, - stream=True, - temperature=USAGE_AI_TEMPERATURE, - ) - async for chunk in response: - delta = chunk.choices[0].delta.content - if delta: - yield _sse({"type": "chunk", "content": delta}) - - -async def stream_usage_ai_chat( - messages: List[Dict[str, str]], - model: Optional[str] = None, - user_id: Optional[str] = None, - is_admin: bool = False, -) -> AsyncIterator[str]: - """Stream SSE events: status → tool_call → chunk → done.""" - resolved_model = (model or "").strip() or DEFAULT_COMPETITOR_DISCOVERY_MODEL - truncated = messages[-MAX_CHAT_MESSAGES:] if len(messages) > MAX_CHAT_MESSAGES else messages - chat_messages: List[Dict[str, Any]] = [ - {"role": "system", "content": _build_system_prompt(is_admin)}, - *truncated, - ] - - try: - yield _sse({"type": "status", "message": "Thinking..."}) - tools = get_tools_for_role(is_admin) - response = await litellm.acompletion( - model=resolved_model, - messages=chat_messages, - tools=tools, - temperature=USAGE_AI_TEMPERATURE, - ) - choice = response.choices[0] # type: ignore - - if not choice.message.tool_calls: - if choice.message.content: - yield _sse({"type": "chunk", "content": choice.message.content}) - yield _sse({"type": "done"}) - return - - chat_messages.append(choice.message.model_dump()) - for tc in choice.message.tool_calls: - async for event in _process_tool_call(tc, chat_messages, user_id, is_admin): - yield event - async for event in _stream_final_response(resolved_model, chat_messages): - yield event - yield _sse({"type": "done"}) - - except Exception as e: - verbose_proxy_logger.error("AI usage chat failed: %s", e) - yield _sse( - { - "type": "error", - "message": "An internal error occurred. Please try again.", - } - ) diff --git a/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py b/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py index 26515c749fb..443936b5cc9 100644 --- a/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py +++ b/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py @@ -1,5 +1,5 @@ """ -USAGE AI CHAT ENDPOINTS +USAGE AI CHAT ENDPOINT /usage/ai/chat - Stream AI chat responses about usage data """ @@ -23,7 +23,7 @@ class ChatMessage(BaseModel): class UsageAIChatRequest(BaseModel): messages: List[ChatMessage] = Field(..., description="Chat messages (user/assistant history)") - model: Optional[str] = Field(default=None, description="Model to use for AI chat") + model: Optional[str] = Field(default=None, description="Model group to use for AI chat") @router.post( @@ -38,30 +38,35 @@ async def usage_ai_chat( ): """ AI chat about usage data. Streams SSE events with the AI response. - The AI agent has access to tools that query aggregated daily activity data. + + The agent queries aggregated daily activity data through a provider scoped + to the caller: admins get a global view, non-admins are restricted to their + own ``user_id``. """ + from litellm.proxy._types import user_api_key_has_admin_view from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_view, require_caller_user_id_for_non_admin, ) - from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import ( + from litellm.proxy.management_endpoints.usage_endpoints.agent import ( stream_usage_ai_chat, ) + from litellm.proxy.management_endpoints.usage_endpoints.scoped_data import ( + AdminScope, + ScopedUsageDataProvider, + UserScope, + ) + from litellm.proxy.proxy_server import prisma_client - is_admin = _user_has_admin_view(user_api_key_dict) - if is_admin: - user_id = user_api_key_dict.user_id + if user_api_key_has_admin_view(user_api_key_dict): + scope = AdminScope(caller_user_id=user_api_key_dict.user_id) else: - user_id = require_caller_user_id_for_non_admin(user_api_key_dict) + scope = UserScope(user_id=require_caller_user_id_for_non_admin(user_api_key_dict)) + + provider = ScopedUsageDataProvider(scope=scope, prisma_client=prisma_client) messages = [{"role": m.role, "content": m.content} for m in data.messages] return StreamingResponse( - stream_usage_ai_chat( - messages=messages, - model=data.model, - user_id=user_id, - is_admin=is_admin, - ), + stream_usage_ai_chat(provider=provider, messages=messages, model=data.model), media_type="text/event-stream", headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}, ) diff --git a/litellm/proxy/management_endpoints/usage_endpoints/scoped_data.py b/litellm/proxy/management_endpoints/usage_endpoints/scoped_data.py new file mode 100644 index 00000000000..11e8611bac6 --- /dev/null +++ b/litellm/proxy/management_endpoints/usage_endpoints/scoped_data.py @@ -0,0 +1,212 @@ +"""Authorization-scoped access to aggregated usage/spend data for Ask AI. + +The caller's authorization is baked into the provider at construction. A +non-admin caller receives a provider that can only ever read its own +``user_id``, so an out-of-scope query is unrepresentable rather than something +each tool handler has to remember to guard. Team and tag breakdowns are +admin-only and the provider refuses them for a user scope as defense in depth. + +Data is read through the same ``common_daily_activity`` helpers that back the +``/spend`` REST endpoints, so there is a single source of truth for the +aggregation logic. +""" + +from dataclasses import dataclass +from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Tuple, Union + +from litellm.types.proxy.management_endpoints.common_daily_activity import ( + DailySpendData, + MetricWithMetadata, + SpendAnalyticsPaginatedResponse, + SpendMetrics, +) + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + +TABLE_DAILY_USER_SPEND = "litellm_dailyuserspend" +TABLE_DAILY_TEAM_SPEND = "litellm_dailyteamspend" +TABLE_DAILY_TAG_SPEND = "litellm_dailytagspend" + +ENTITY_FIELD_USER = "user_id" +ENTITY_FIELD_TEAM = "team_id" +ENTITY_FIELD_TAG = "tag" + +PAGINATED_PAGE_SIZE = 200 +TOP_N_MODELS = 15 +TOP_N_PROVIDERS = 10 + + +@dataclass(frozen=True, slots=True) +class AdminScope: + """Global view. ``caller_user_id`` is the admin's own id (may be None) and + is not used to filter; admins may optionally pass an explicit user filter.""" + + caller_user_id: Optional[str] + + +@dataclass(frozen=True, slots=True) +class UserScope: + """Non-admin view. Every query is forced to this ``user_id``.""" + + user_id: str + + +AiChatScope = Union[AdminScope, UserScope] + + +def _parse_csv(raw: Optional[str]) -> Optional[List[str]]: + if not raw: + return None + return [t.strip() for t in raw.split(",") if t.strip()] + + +class ScopedUsageDataProvider: + """Reads daily activity data within the bounds of an ``AiChatScope``.""" + + def __init__(self, scope: AiChatScope, prisma_client: Optional["PrismaClient"]) -> None: + self._scope = scope + self._prisma_client = prisma_client + + @property + def is_admin(self) -> bool: + return isinstance(self._scope, AdminScope) + + async def usage( + self, start_date: str, end_date: str, user_id_filter: Optional[str] + ) -> SpendAnalyticsPaginatedResponse: + scope = self._scope + effective_user_id = scope.user_id if isinstance(scope, UserScope) else user_id_filter + from litellm.proxy.management_endpoints.common_daily_activity import ( + get_daily_activity_aggregated, + ) + + return await get_daily_activity_aggregated( + prisma_client=self._prisma_client, + table_name=TABLE_DAILY_USER_SPEND, + entity_id_field=ENTITY_FIELD_USER, + entity_id=effective_user_id, + entity_metadata_field=None, + start_date=start_date, + end_date=end_date, + model=None, + api_key=None, + ) + + async def team(self, start_date: str, end_date: str, team_ids: Optional[str]) -> SpendAnalyticsPaginatedResponse: + self._require_admin("team usage") + return await self._paginated( + TABLE_DAILY_TEAM_SPEND, ENTITY_FIELD_TEAM, _parse_csv(team_ids), start_date, end_date + ) + + async def tag(self, start_date: str, end_date: str, tags: Optional[str]) -> SpendAnalyticsPaginatedResponse: + self._require_admin("tag usage") + return await self._paginated(TABLE_DAILY_TAG_SPEND, ENTITY_FIELD_TAG, _parse_csv(tags), start_date, end_date) + + def _require_admin(self, what: str) -> None: + if not isinstance(self._scope, AdminScope): + raise PermissionError(f"{what} data is only available to admin callers") + + async def _paginated( + self, + table_name: str, + entity_id_field: str, + entity_id: Optional[List[str]], + start_date: str, + end_date: str, + ) -> SpendAnalyticsPaginatedResponse: + from litellm.proxy.management_endpoints.common_daily_activity import ( + get_daily_activity, + ) + + return await get_daily_activity( + prisma_client=self._prisma_client, + table_name=table_name, + entity_id_field=entity_id_field, + entity_id=entity_id, + entity_metadata_field=None, + start_date=start_date, + end_date=end_date, + model=None, + api_key=None, + page=1, + page_size=PAGINATED_PAGE_SIZE, + ) + + +@dataclass(frozen=True, slots=True) +class _Totals: + spend: float + api_requests: int + total_tokens: int + + def plus(self, m: "SpendMetrics") -> "_Totals": + return _Totals( + spend=self.spend + m.spend, + api_requests=self.api_requests + m.api_requests, + total_tokens=self.total_tokens + m.total_tokens, + ) + + +def _accumulate( + days: List[DailySpendData], + pick: Callable[[DailySpendData], Dict[str, MetricWithMetadata]], +) -> Dict[str, _Totals]: + """Sum one breakdown dimension across days into per-name totals.""" + totals: Dict[str, _Totals] = {} + for day in days: + for name, entry in pick(day).items(): + totals[name] = totals.get(name, _Totals(0.0, 0, 0)).plus(entry.metrics) + return totals + + +def _ranked(totals: Dict[str, _Totals], limit: int) -> List[Tuple[str, _Totals]]: + return sorted(totals.items(), key=lambda item: -item[1].spend)[:limit] + + +def summarise_usage_data(resp: SpendAnalyticsPaginatedResponse) -> str: + """Render global/user usage into concise text the LLM can reason over.""" + meta = resp.metadata + header = ( + f"Total Spend: ${meta.total_spend:.4f}\n" + f"Total Requests: {meta.total_api_requests}\n" + f"Successful: {meta.total_successful_requests} | Failed: {meta.total_failed_requests}\n" + f"Total Tokens: {meta.total_tokens}" + ) + + models = _accumulate(resp.results, lambda d: d.breakdown.models) + providers = _accumulate(resp.results, lambda d: d.breakdown.providers) + + model_lines = [ + f" - {name}: ${t.spend:.4f} ({t.api_requests} reqs, {t.total_tokens} tokens)" + for name, t in _ranked(models, TOP_N_MODELS) + ] + provider_lines = [ + f" - {name}: ${t.spend:.4f} ({t.api_requests} reqs)" for name, t in _ranked(providers, TOP_N_PROVIDERS) + ] + + sections = [header, ""] + sections += ["Top Models by Spend:"] + (model_lines or [" (no data)"]) + [""] + sections += ["Top Providers by Spend:"] + (provider_lines or [" (no data)"]) + return "\n".join(sections) + + +def _entity_alias(entry: MetricWithMetadata, entity_id: str) -> str: + raw = entry.metadata.get("alias") + return raw if isinstance(raw, str) and raw else entity_id + + +def summarise_entity_data(resp: SpendAnalyticsPaginatedResponse, entity_label: str) -> str: + """Render team/tag entity usage into concise text.""" + if not resp.results: + return f"No {entity_label} usage data found for the given date range." + + totals = _accumulate(resp.results, lambda d: d.breakdown.entities) + aliases = {eid: _entity_alias(entry, eid) for day in resp.results for eid, entry in day.breakdown.entities.items()} + + lines = [f"{entity_label} Usage ({len(totals)} {entity_label.lower()}s):", ""] + for eid, t in _ranked(totals, len(totals)): + lines.append( + f"- {aliases.get(eid, eid)} (ID: {eid}): ${t.spend:.4f} | {t.api_requests} reqs | {t.total_tokens} tokens" + ) + return "\n".join(lines) diff --git a/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_agent.py b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_agent.py new file mode 100644 index 00000000000..7c8442e6d66 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_agent.py @@ -0,0 +1,299 @@ +"""Tests for the Ask AI agent loop. + +Highest-value regression: the LLM call goes through the proxy's ``llm_router`` +(so UI-selected model groups resolve), not the bare ``litellm`` SDK, and the +scope-forced user_id survives an end-to-end tool round. +""" + +import json +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +import litellm +from litellm.types.utils import ( + ChatCompletionDeltaToolCall, + Delta, +) +from litellm.types.utils import Function as DeltaFunction +from litellm.types.utils import ( + ModelResponseStream, + StreamingChoices, +) +from litellm.proxy.management_endpoints.usage_endpoints import agent as agent_mod +from litellm.proxy.management_endpoints.usage_endpoints.agent import ( + LLMCallError, + ModelNotConfigured, + RouterUnavailable, + _error_event, + resolve_model, + stream_usage_ai_chat, + tools_for_role, +) +from litellm.proxy.management_endpoints.usage_endpoints.scoped_data import ( + AdminScope, + ScopedUsageDataProvider, + UserScope, +) +from litellm.types.proxy.management_endpoints.common_daily_activity import ( + SpendAnalyticsPaginatedResponse, +) + +DATE_ARGS = {"start_date": "2025-01-01", "end_date": "2025-01-31"} + +SAMPLE_USAGE = { + "results": [ + { + "date": "2025-01-15", + "metrics": {"spend": 12.5, "api_requests": 100}, + "breakdown": {"models": {}, "providers": {}, "entities": {}}, + } + ], + "metadata": {"total_spend": 12.5, "total_api_requests": 100}, +} + + +def _content_chunk(text: str) -> ModelResponseStream: + return ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(content=text))]) + + +def _toolcall_chunk(call_id: str, name: str, args: dict) -> ModelResponseStream: + tc = ChatCompletionDeltaToolCall( + index=0, id=call_id, type="function", function=DeltaFunction(name=name, arguments=json.dumps(args)) + ) + return ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(tool_calls=[tc]))]) + + +class FakeRouter: + """Router stub whose acompletion returns scripted streaming chunks per call.""" + + def __init__(self, scripts): + self._scripts = list(scripts) + self.calls = [] + + async def acompletion(self, **kwargs): + idx = len(self.calls) + self.calls.append(kwargs) + chunks = self._scripts[idx] + + async def _gen(): + for c in chunks: + yield c + + return _gen() + + +def _usage_response_mock(): + return SpendAnalyticsPaginatedResponse.model_validate(SAMPLE_USAGE) + + +async def _collect(provider, messages, model, router): + """Run the agent with a patched router; return parsed SSE events.""" + with ( + patch.object(agent_mod, "_require_router", return_value=router), + patch( + "litellm.proxy.management_endpoints.common_daily_activity.get_daily_activity_aggregated", + new_callable=AsyncMock, + ) as mock_agg, + ): + mock_agg.return_value = _usage_response_mock() + events = [] + async for raw in stream_usage_ai_chat(provider=provider, messages=messages, model=model): + events.append(json.loads(raw.replace("data: ", "").strip())) + return events, mock_agg + + +class TestRoutesThroughProxyRouter: + @pytest.mark.asyncio + async def test_selected_model_group_is_passed_to_llm_router_not_sdk(self): + """The core v2 fix: a UI-selected model group is sent to llm_router, + and the bare litellm.acompletion SDK is never called.""" + router = FakeRouter([[_content_chunk("Your spend is $12.50.")]]) + provider = ScopedUsageDataProvider(scope=AdminScope(caller_user_id="a1"), prisma_client=MagicMock()) + + with patch.object(litellm, "acompletion", new_callable=AsyncMock) as sdk_call: + events, _ = await _collect(provider, [{"role": "user", "content": "spend?"}], "my-model-group", router) + + assert sdk_call.call_count == 0 + assert router.calls[0]["model"] == "my-model-group" + assert router.calls[0]["metadata"] == {"feature": "usage_ai"} + chunks = [e for e in events if e["type"] == "chunk"] + assert chunks and chunks[0]["content"] == "Your spend is $12.50." + assert events[-1]["type"] == "done" + + @pytest.mark.asyncio + async def test_tool_round_then_streamed_answer(self): + router = FakeRouter( + [ + [_toolcall_chunk("c1", "get_usage_data", DATE_ARGS)], + [_content_chunk("Total spend is $12.50.")], + ] + ) + provider = ScopedUsageDataProvider(scope=AdminScope(caller_user_id="a1"), prisma_client=MagicMock()) + events, mock_agg = await _collect(provider, [{"role": "user", "content": "spend?"}], "m", router) + + assert mock_agg.call_count == 1 + tool_events = [e for e in events if e["type"] == "tool_call"] + assert {e["status"] for e in tool_events} == {"running", "complete"} + assert tool_events[0]["tool_name"] == "get_usage_data" + assert any(e["type"] == "chunk" and "$12.50" in e["content"] for e in events) + assert events[-1]["type"] == "done" + + +class TestScopeEnforcedThroughLoop: + @pytest.mark.asyncio + async def test_non_admin_tool_call_cannot_reach_other_users_data(self): + """Even if the model asks for another user's data, the scoped provider + forces the caller's own user_id at the data layer.""" + router = FakeRouter( + [ + [_toolcall_chunk("c1", "get_usage_data", {**DATE_ARGS, "user_id": "victim"})], + [_content_chunk("Here is your usage.")], + ] + ) + provider = ScopedUsageDataProvider(scope=UserScope(user_id="caller"), prisma_client=MagicMock()) + _events, mock_agg = await _collect(provider, [{"role": "user", "content": "show victim usage"}], "m", router) + + assert mock_agg.call_args.kwargs["entity_id"] == "caller" + + +class TestMultiRoundLoop: + @pytest.mark.asyncio + async def test_two_tool_rounds_then_answer(self): + router = FakeRouter( + [ + [_toolcall_chunk("c1", "get_usage_data", DATE_ARGS)], + [_toolcall_chunk("c2", "get_team_usage_data", DATE_ARGS)], + [_content_chunk("Engineering leads spend.")], + ] + ) + provider = ScopedUsageDataProvider(scope=AdminScope(caller_user_id="a1"), prisma_client=MagicMock()) + + with ( + patch.object(agent_mod, "_require_router", return_value=router), + patch( + "litellm.proxy.management_endpoints.common_daily_activity.get_daily_activity_aggregated", + new_callable=AsyncMock, + ) as mock_agg, + patch( + "litellm.proxy.management_endpoints.common_daily_activity.get_daily_activity", + new_callable=AsyncMock, + ) as mock_paginated, + ): + mock_agg.return_value = _usage_response_mock() + mock_paginated.return_value = _usage_response_mock() + events = [ + json.loads(raw.replace("data: ", "").strip()) + async for raw in stream_usage_ai_chat( + provider=provider, messages=[{"role": "user", "content": "q"}], model="m" + ) + ] + + assert len(router.calls) == 3 + assert mock_agg.call_count == 1 + assert mock_paginated.call_count == 1 + assert any(e["type"] == "chunk" and "Engineering" in e["content"] for e in events) + assert events[-1]["type"] == "done" + + @pytest.mark.asyncio + async def test_runaway_tool_calls_are_capped_and_still_answer(self): + """If the model never stops calling tools, the loop caps rounds and + forces a final (tool-less) answer instead of looping forever.""" + # More tool-call rounds scripted than MAX_TOOL_ROUNDS; the final call + # is made with tools disabled so it must return content. + scripts = [[_toolcall_chunk(f"c{i}", "get_usage_data", DATE_ARGS)] for i in range(agent_mod.MAX_TOOL_ROUNDS)] + scripts.append([_content_chunk("Final answer.")]) + router = FakeRouter(scripts) + provider = ScopedUsageDataProvider(scope=AdminScope(caller_user_id="a1"), prisma_client=MagicMock()) + + events, _ = await _collect(provider, [{"role": "user", "content": "q"}], "m", router) + + # Last router call must have had tools disabled (the safety net). + assert router.calls[-1]["tools"] is None + assert any(e["type"] == "chunk" and "Final answer." in e["content"] for e in events) + assert events[-1]["type"] == "done" + + +class TestErrorPaths: + @pytest.mark.asyncio + async def test_model_not_configured_errors_without_calling_router(self): + provider = ScopedUsageDataProvider(scope=AdminScope(caller_user_id="a1"), prisma_client=MagicMock()) + with ( + patch.object(agent_mod, "_require_router") as require_router, + patch("litellm.proxy.proxy_server.general_settings", {}), + ): + events = [ + json.loads(raw.replace("data: ", "").strip()) + async for raw in stream_usage_ai_chat( + provider=provider, messages=[{"role": "user", "content": "q"}], model=None + ) + ] + + require_router.assert_not_called() + assert len(events) == 1 + assert events[0]["type"] == "error" + assert "usage_ai_model" in events[0]["message"] + + @pytest.mark.asyncio + async def test_router_unavailable_emits_error(self): + provider = ScopedUsageDataProvider(scope=AdminScope(caller_user_id="a1"), prisma_client=MagicMock()) + with patch.object(agent_mod, "_require_router", side_effect=agent_mod._RouterUnavailableError()): + events = [ + json.loads(raw.replace("data: ", "").strip()) + async for raw in stream_usage_ai_chat( + provider=provider, messages=[{"role": "user", "content": "q"}], model="m" + ) + ] + + error_events = [e for e in events if e["type"] == "error"] + assert len(error_events) == 1 + assert "router" in error_events[0]["message"].lower() + + @pytest.mark.asyncio + async def test_llm_exception_emits_error(self): + provider = ScopedUsageDataProvider(scope=AdminScope(caller_user_id="a1"), prisma_client=MagicMock()) + broken_router = MagicMock() + broken_router.acompletion = AsyncMock(side_effect=RuntimeError("boom")) + with patch.object(agent_mod, "_require_router", return_value=broken_router): + events = [ + json.loads(raw.replace("data: ", "").strip()) + async for raw in stream_usage_ai_chat( + provider=provider, messages=[{"role": "user", "content": "q"}], model="m" + ) + ] + + error_events = [e for e in events if e["type"] == "error"] + assert len(error_events) == 1 + assert "failed" in error_events[0]["message"].lower() + + +class TestResolveModel: + def test_explicit_request_wins(self): + assert resolve_model("chosen-group") == "chosen-group" + + def test_falls_back_to_configured_setting(self): + with patch("litellm.proxy.proxy_server.general_settings", {"usage_ai_model": "configured-group"}): + assert resolve_model(None) == "configured-group" + + def test_blank_request_falls_back_to_setting(self): + with patch("litellm.proxy.proxy_server.general_settings", {"usage_ai_model": "configured-group"}): + assert resolve_model(" ") == "configured-group" + + def test_no_model_anywhere_returns_error_value(self): + with patch("litellm.proxy.proxy_server.general_settings", {}): + assert isinstance(resolve_model(None), ModelNotConfigured) + + +class TestToolsAndErrorMapping: + def test_admin_gets_all_tools_non_admin_gets_usage_only(self): + assert {t["function"]["name"] for t in tools_for_role(True)} == { + "get_usage_data", + "get_team_usage_data", + "get_tag_usage_data", + } + assert {t["function"]["name"] for t in tools_for_role(False)} == {"get_usage_data"} + + def test_error_messages_are_distinct_and_actionable(self): + assert "usage_ai_model" in _error_event(ModelNotConfigured())["message"] + assert "router" in _error_event(RouterUnavailable())["message"].lower() + assert "failed" in _error_event(LLMCallError(detail="x"))["message"].lower() diff --git a/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py deleted file mode 100644 index e8a74e41dae..00000000000 --- a/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py +++ /dev/null @@ -1,468 +0,0 @@ -""" -Tests for AI Usage Chat module. -""" - -import json -from unittest.mock import AsyncMock, MagicMock, patch - -import pytest - -from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import ( - TOOL_HANDLERS, - TOOLS_ADMIN, - TOOLS_BASE, - _build_system_prompt, - _summarise_entity_data, - _summarise_usage_data, - stream_usage_ai_chat, -) - - -SAMPLE_AGGREGATED_RESPONSE = { - "results": [ - { - "date": "2025-01-15", - "metrics": { - "spend": 50.25, - "prompt_tokens": 20000, - "completion_tokens": 10000, - "total_tokens": 30000, - "api_requests": 500, - "successful_requests": 480, - "failed_requests": 20, - "cache_read_input_tokens": 0, - "cache_creation_input_tokens": 0, - }, - "breakdown": { - "models": { - "gpt-4": { - "metrics": { - "spend": 40.0, - "api_requests": 300, - "total_tokens": 25000, - }, - "metadata": {}, - "api_key_breakdown": {}, - }, - }, - "providers": { - "openai": { - "metrics": {"spend": 50.25, "api_requests": 500}, - "metadata": {}, - "api_key_breakdown": {}, - }, - }, - "api_keys": { - "sk-test123": { - "metrics": {"spend": 50.25}, - "metadata": {"key_alias": "Production Key"}, - }, - }, - "model_groups": {}, - "mcp_servers": {}, - "entities": {}, - }, - }, - ], - "metadata": { - "total_spend": 50.25, - "total_api_requests": 500, - "total_successful_requests": 480, - "total_failed_requests": 20, - "total_tokens": 30000, - }, -} - -SAMPLE_TEAM_RESPONSE = { - "results": [ - { - "date": "2025-01-15", - "metrics": {"spend": 100.0, "api_requests": 1000, "total_tokens": 50000}, - "breakdown": { - "entities": { - "team-1": { - "metrics": { - "spend": 60.0, - "api_requests": 600, - "total_tokens": 30000, - }, - "metadata": {"alias": "Engineering"}, - "api_key_breakdown": {}, - }, - "team-2": { - "metrics": { - "spend": 40.0, - "api_requests": 400, - "total_tokens": 20000, - }, - "metadata": {"alias": "Marketing"}, - "api_key_breakdown": {}, - }, - }, - "models": {}, - "providers": {}, - "api_keys": {}, - "model_groups": {}, - "mcp_servers": {}, - }, - }, - ], - "metadata": {"total_spend": 100.0, "total_api_requests": 1000}, -} - - -class TestToolSchemas: - def test_admin_tools_include_all(self): - assert len(TOOLS_ADMIN) == 3 - names = {t["function"]["name"] for t in TOOLS_ADMIN} - assert "get_usage_data" in names - assert "get_team_usage_data" in names - assert "get_tag_usage_data" in names - - def test_base_tools_restricted_to_usage_only(self): - assert len(TOOLS_BASE) == 1 - assert TOOLS_BASE[0]["function"]["name"] == "get_usage_data" - - def test_admin_prompt_mentions_all_tools(self): - prompt = _build_system_prompt(is_admin=True) - assert "get_usage_data" in prompt - assert "get_team_usage_data" in prompt - assert "get_tag_usage_data" in prompt - - def test_non_admin_prompt_only_mentions_usage_tool(self): - prompt = _build_system_prompt(is_admin=False) - assert "get_usage_data" in prompt - assert "get_team_usage_data" not in prompt - assert "get_tag_usage_data" not in prompt - - def test_system_prompt_includes_todays_date(self): - from datetime import date - - prompt = _build_system_prompt(is_admin=True) - assert date.today().isoformat() in prompt - - -class TestSummariseUsageData: - def test_summarise_includes_totals(self): - summary = _summarise_usage_data(SAMPLE_AGGREGATED_RESPONSE) - assert "$50.25" in summary - assert "500" in summary - - def test_summarise_includes_models(self): - summary = _summarise_usage_data(SAMPLE_AGGREGATED_RESPONSE) - assert "gpt-4" in summary - - def test_summarise_includes_providers(self): - summary = _summarise_usage_data(SAMPLE_AGGREGATED_RESPONSE) - assert "openai" in summary - - def test_summarise_handles_empty_data(self): - empty = {"results": [], "metadata": {}} - summary = _summarise_usage_data(empty) - assert "no data" in summary.lower() - - -class TestSummariseEntityData: - def test_team_summary_includes_teams(self): - summary = _summarise_entity_data(SAMPLE_TEAM_RESPONSE, "Team") - assert "Engineering" in summary - assert "Marketing" in summary - assert "$60.0" in summary - assert "$40.0" in summary - - def test_team_summary_empty(self): - empty = {"results": [], "metadata": {}} - summary = _summarise_entity_data(empty, "Team") - assert "No Team usage data" in summary - - -class TestStreamUsageAiChat: - @pytest.mark.asyncio - async def test_stream_emits_status_events(self): - mock_tool_call = MagicMock() - mock_tool_call.id = "call_123" - mock_tool_call.function.name = "get_usage_data" - mock_tool_call.function.arguments = json.dumps( - { - "start_date": "2025-01-01", - "end_date": "2025-01-31", - } - ) - - mock_first_response = MagicMock() - mock_first_response.choices = [MagicMock()] - mock_first_response.choices[0].message.tool_calls = [mock_tool_call] - mock_first_response.choices[0].message.model_dump.return_value = { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_123", - "type": "function", - "function": { - "name": "get_usage_data", - "arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31"}', - }, - } - ], - } - - async def mock_stream(): - chunk = MagicMock() - chunk.choices = [MagicMock()] - chunk.choices[0].delta.content = "Total spend is $50.25" - yield chunk - - with ( - patch( - "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm" - ) as mock_litellm, - patch( - "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat._fetch_usage_data", - new_callable=AsyncMock, - ) as mock_fetch, - ): - mock_litellm.acompletion = AsyncMock( - side_effect=[ - mock_first_response, - mock_stream(), - ] - ) - mock_fetch.return_value = SAMPLE_AGGREGATED_RESPONSE - - events = [] - async for event in stream_usage_ai_chat( - messages=[{"role": "user", "content": "What is my total spend?"}], - model="gpt-4o-mini", - user_id="user-123", - is_admin=True, - ): - events.append(json.loads(event.replace("data: ", "").strip())) - - status_events = [e for e in events if e["type"] == "status"] - tool_call_events = [e for e in events if e["type"] == "tool_call"] - chunk_events = [e for e in events if e["type"] == "chunk"] - done_events = [e for e in events if e["type"] == "done"] - - assert len(status_events) >= 1 - assert "Thinking" in status_events[0]["message"] - assert len(tool_call_events) >= 1 - assert tool_call_events[0]["tool_name"] == "get_usage_data" - assert tool_call_events[0]["status"] in ("running", "complete") - assert len(chunk_events) >= 1 - assert len(done_events) == 1 - - @pytest.mark.asyncio - async def test_stream_handles_team_tool(self): - mock_tool_call = MagicMock() - mock_tool_call.id = "call_team" - mock_tool_call.function.name = "get_team_usage_data" - mock_tool_call.function.arguments = json.dumps( - { - "start_date": "2025-01-01", - "end_date": "2025-01-31", - } - ) - - mock_first_response = MagicMock() - mock_first_response.choices = [MagicMock()] - mock_first_response.choices[0].message.tool_calls = [mock_tool_call] - mock_first_response.choices[0].message.model_dump.return_value = { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_team", - "type": "function", - "function": { - "name": "get_team_usage_data", - "arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31"}', - }, - } - ], - } - - async def mock_stream(): - chunk = MagicMock() - chunk.choices = [MagicMock()] - chunk.choices[0].delta.content = "Engineering is the top team." - yield chunk - - with ( - patch( - "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm" - ) as mock_litellm, - patch( - "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat._fetch_team_usage_data", - new_callable=AsyncMock, - ) as mock_fetch, - ): - mock_litellm.acompletion = AsyncMock( - side_effect=[ - mock_first_response, - mock_stream(), - ] - ) - mock_fetch.return_value = SAMPLE_TEAM_RESPONSE - - events = [] - async for event in stream_usage_ai_chat( - messages=[{"role": "user", "content": "Which team spends the most?"}], - model="gpt-4o-mini", - is_admin=True, - ): - events.append(json.loads(event.replace("data: ", "").strip())) - - chunk_events = [e for e in events if e["type"] == "chunk"] - assert len(chunk_events) >= 1 - assert "Engineering" in chunk_events[0]["content"] - - @pytest.mark.asyncio - async def test_stream_handles_error(self): - with patch( - "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm" - ) as mock_litellm: - mock_litellm.acompletion = AsyncMock(side_effect=Exception("LLM error")) - - events = [] - async for event in stream_usage_ai_chat( - messages=[{"role": "user", "content": "test"}], - ): - events.append(json.loads(event.replace("data: ", "").strip())) - - error_events = [e for e in events if e["type"] == "error"] - assert len(error_events) == 1 - assert "internal error" in error_events[0]["message"].lower() - - @pytest.mark.asyncio - async def test_non_admin_enforces_user_id(self): - mock_tool_call = MagicMock() - mock_tool_call.id = "call_456" - mock_tool_call.function.name = "get_usage_data" - mock_tool_call.function.arguments = json.dumps( - { - "start_date": "2025-01-01", - "end_date": "2025-01-31", - "user_id": "other-user", - } - ) - - mock_first_response = MagicMock() - mock_first_response.choices = [MagicMock()] - mock_first_response.choices[0].message.tool_calls = [mock_tool_call] - mock_first_response.choices[0].message.model_dump.return_value = { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_456", - "type": "function", - "function": { - "name": "get_usage_data", - "arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31","user_id":"other-user"}', - }, - } - ], - } - - async def mock_stream(): - chunk = MagicMock() - chunk.choices = [MagicMock()] - chunk.choices[0].delta.content = "Data." - yield chunk - - mock_fetch = AsyncMock(return_value=SAMPLE_AGGREGATED_RESPONSE) - - with ( - patch( - "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm" - ) as mock_litellm, - patch.dict( - "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.TOOL_HANDLERS", - { - "get_usage_data": { - "fetch": mock_fetch, - "summarise": _summarise_usage_data, - "label": "global usage data", - } - }, - ), - ): - mock_litellm.acompletion = AsyncMock( - side_effect=[ - mock_first_response, - mock_stream(), - ] - ) - - events = [] - async for event in stream_usage_ai_chat( - messages=[{"role": "user", "content": "Show data"}], - model="gpt-4o-mini", - user_id="my-user-id", - is_admin=False, - ): - events.append(event) - - mock_fetch.assert_called_once_with( - start_date="2025-01-01", - end_date="2025-01-31", - user_id="my-user-id", - ) - - -class TestUsageAiChatServiceAccountGuard: - """ - Security regression: a non-admin caller with user_id=None (service-account - key) must be rejected at the endpoint boundary, before any tool dispatch. - """ - - @pytest.mark.asyncio - async def test_non_admin_with_user_id_none_is_rejected(self): - from fastapi import HTTPException - - from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth - from litellm.proxy.management_endpoints.usage_endpoints.endpoints import ( - ChatMessage, - UsageAIChatRequest, - usage_ai_chat, - ) - - service_account_key = UserAPIKeyAuth( - user_id=None, - user_role=LitellmUserRoles.INTERNAL_USER, - ) - request = MagicMock() - body = UsageAIChatRequest( - messages=[ChatMessage(role="user", content="hi")], - model="gpt-4o-mini", - ) - - with pytest.raises(HTTPException) as exc_info: - await usage_ai_chat( - data=body, - request=request, - user_api_key_dict=service_account_key, - ) - - assert exc_info.value.status_code == 403 - assert "Service-account keys" in str(exc_info.value.detail) - - def test_resolve_fetch_kwargs_tripwire_fires_on_none_user_id(self): - """ - Defense-in-depth: if a future endpoint forgets the entry guard and - a non-admin caller with user_id=None reaches _resolve_fetch_kwargs, - the tripwire must fire rather than issuing an unscoped query. - """ - from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import ( - _resolve_fetch_kwargs, - ) - - with pytest.raises(ValueError) as exc_info: - _resolve_fetch_kwargs( - fn_name="get_usage_data", - fn_args={"start_date": "2025-01-01", "end_date": "2025-01-31"}, - user_id=None, - is_admin=False, - ) - assert "Endpoint-level guard missing" in str(exc_info.value) diff --git a/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_endpoints.py b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_endpoints.py new file mode 100644 index 00000000000..d0fd761dda5 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_endpoints.py @@ -0,0 +1,67 @@ +"""Endpoint-boundary tests for /usage/ai/chat. + +Security regression: a non-admin caller with user_id=None (a service-account +key) must be rejected at the endpoint before any scope/provider is built, so it +can never fall through to an unscoped global query. +""" + +from unittest.mock import MagicMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.management_endpoints.usage_endpoints.endpoints import ( + ChatMessage, + UsageAIChatRequest, + usage_ai_chat, +) + + +class TestServiceAccountGuard: + @pytest.mark.asyncio + async def test_non_admin_with_user_id_none_is_rejected(self): + service_account_key = UserAPIKeyAuth(user_id=None, user_role=LitellmUserRoles.INTERNAL_USER) + body = UsageAIChatRequest(messages=[ChatMessage(role="user", content="hi")], model="m") + + with pytest.raises(HTTPException) as exc_info: + await usage_ai_chat(data=body, request=MagicMock(), user_api_key_dict=service_account_key) + + assert exc_info.value.status_code == 403 + assert "Service-account keys" in str(exc_info.value.detail) + + +class TestScopeSelection: + @pytest.mark.asyncio + async def test_admin_caller_builds_admin_scope(self): + admin_key = UserAPIKeyAuth(user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN) + body = UsageAIChatRequest(messages=[ChatMessage(role="user", content="hi")], model="m") + + captured = {} + + async def _fake_stream(*, provider, messages, model): + captured["is_admin"] = provider.is_admin + captured["messages"] = messages + if False: + yield "" # pragma: no cover + + # Inject a fake agent stream + prisma via patching the lazily imported names. + import litellm.proxy.proxy_server as proxy_server + + original_prisma = getattr(proxy_server, "prisma_client", None) + proxy_server.prisma_client = MagicMock() + try: + with pytest.MonkeyPatch.context() as mp: + mp.setattr( + "litellm.proxy.management_endpoints.usage_endpoints.agent.stream_usage_ai_chat", + _fake_stream, + ) + response = await usage_ai_chat(data=body, request=MagicMock(), user_api_key_dict=admin_key) + # Drain the streaming body so the generator runs. + async for _ in response.body_iterator: + pass + finally: + proxy_server.prisma_client = original_prisma + + assert captured["is_admin"] is True + assert captured["messages"] == [{"role": "user", "content": "hi"}] diff --git a/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_scoped_data.py b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_scoped_data.py new file mode 100644 index 00000000000..9576e0ade6a --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_scoped_data.py @@ -0,0 +1,180 @@ +"""Tests for the authorization-scoped usage data provider. + +The security property under test: a non-admin (UserScope) provider can only +ever read its own user_id, no matter what filter a tool passes, and team/tag +breakdowns are refused for non-admins. +""" + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from litellm.proxy.management_endpoints.usage_endpoints.scoped_data import ( + AdminScope, + ScopedUsageDataProvider, + UserScope, + summarise_entity_data, + summarise_usage_data, +) +from litellm.types.proxy.management_endpoints.common_daily_activity import ( + SpendAnalyticsPaginatedResponse, +) + + +def _as_response(payload) -> SpendAnalyticsPaginatedResponse: + return SpendAnalyticsPaginatedResponse.model_validate(payload) + + +SAMPLE_AGGREGATED_RESPONSE = { + "results": [ + { + "date": "2025-01-15", + "metrics": {"spend": 50.25, "total_tokens": 30000, "api_requests": 500}, + "breakdown": { + "models": { + "gpt-4": {"metrics": {"spend": 40.0, "api_requests": 300, "total_tokens": 25000}, "metadata": {}}, + }, + "providers": { + "openai": {"metrics": {"spend": 50.25, "api_requests": 500}, "metadata": {}}, + }, + "entities": {}, + }, + }, + ], + "metadata": { + "total_spend": 50.25, + "total_api_requests": 500, + "total_successful_requests": 480, + "total_failed_requests": 20, + "total_tokens": 30000, + }, +} + +SAMPLE_TEAM_RESPONSE = { + "results": [ + { + "date": "2025-01-15", + "metrics": {"spend": 100.0}, + "breakdown": { + "entities": { + "team-1": { + "metrics": {"spend": 60.0, "api_requests": 600, "total_tokens": 30000}, + "metadata": {"alias": "Engineering"}, + }, + "team-2": { + "metrics": {"spend": 40.0, "api_requests": 400, "total_tokens": 20000}, + "metadata": {"alias": "Marketing"}, + }, + }, + }, + }, + ], + "metadata": {"total_spend": 100.0}, +} + + +def _response_mock(payload): + resp = MagicMock() + resp.model_dump.return_value = payload + return resp + + +class TestUserScopeForcesOwnUserId: + @pytest.mark.asyncio + async def test_user_scope_ignores_model_supplied_user_id_filter(self): + """Cross-tenant regression: a non-admin query must be scoped to the + caller's own user_id even when the tool arguments name a different one.""" + provider = ScopedUsageDataProvider(scope=UserScope(user_id="my-user"), prisma_client=MagicMock()) + + with patch( + "litellm.proxy.management_endpoints.common_daily_activity.get_daily_activity_aggregated", + new_callable=AsyncMock, + ) as mock_agg: + mock_agg.return_value = _response_mock(SAMPLE_AGGREGATED_RESPONSE) + + await provider.usage(start_date="2025-01-01", end_date="2025-01-31", user_id_filter="other-user") + + assert mock_agg.call_args.kwargs["entity_id"] == "my-user" + + @pytest.mark.asyncio + async def test_admin_scope_honors_supplied_user_id_filter(self): + provider = ScopedUsageDataProvider(scope=AdminScope(caller_user_id="admin-1"), prisma_client=MagicMock()) + + with patch( + "litellm.proxy.management_endpoints.common_daily_activity.get_daily_activity_aggregated", + new_callable=AsyncMock, + ) as mock_agg: + mock_agg.return_value = _response_mock(SAMPLE_AGGREGATED_RESPONSE) + + await provider.usage(start_date="2025-01-01", end_date="2025-01-31", user_id_filter="target-user") + + assert mock_agg.call_args.kwargs["entity_id"] == "target-user" + + @pytest.mark.asyncio + async def test_admin_scope_global_view_when_no_filter(self): + provider = ScopedUsageDataProvider(scope=AdminScope(caller_user_id="admin-1"), prisma_client=MagicMock()) + + with patch( + "litellm.proxy.management_endpoints.common_daily_activity.get_daily_activity_aggregated", + new_callable=AsyncMock, + ) as mock_agg: + mock_agg.return_value = _response_mock(SAMPLE_AGGREGATED_RESPONSE) + + await provider.usage(start_date="2025-01-01", end_date="2025-01-31", user_id_filter=None) + + assert mock_agg.call_args.kwargs["entity_id"] is None + + +class TestAdminOnlyBreakdowns: + @pytest.mark.asyncio + async def test_user_scope_cannot_query_team_data(self): + provider = ScopedUsageDataProvider(scope=UserScope(user_id="u1"), prisma_client=MagicMock()) + with pytest.raises(PermissionError): + await provider.team(start_date="2025-01-01", end_date="2025-01-31", team_ids=None) + + @pytest.mark.asyncio + async def test_user_scope_cannot_query_tag_data(self): + provider = ScopedUsageDataProvider(scope=UserScope(user_id="u1"), prisma_client=MagicMock()) + with pytest.raises(PermissionError): + await provider.tag(start_date="2025-01-01", end_date="2025-01-31", tags=None) + + @pytest.mark.asyncio + async def test_admin_scope_can_query_team_data_with_parsed_ids(self): + provider = ScopedUsageDataProvider(scope=AdminScope(caller_user_id=None), prisma_client=MagicMock()) + with patch( + "litellm.proxy.management_endpoints.common_daily_activity.get_daily_activity", + new_callable=AsyncMock, + ) as mock_paginated: + mock_paginated.return_value = _response_mock(SAMPLE_TEAM_RESPONSE) + await provider.team(start_date="2025-01-01", end_date="2025-01-31", team_ids="team-1, team-2") + + assert mock_paginated.call_args.kwargs["entity_id"] == ["team-1", "team-2"] + + +class TestIsAdminFlag: + def test_admin_scope_is_admin(self): + assert ScopedUsageDataProvider(scope=AdminScope(caller_user_id=None), prisma_client=None).is_admin is True + + def test_user_scope_is_not_admin(self): + assert ScopedUsageDataProvider(scope=UserScope(user_id="u1"), prisma_client=None).is_admin is False + + +class TestSummarisers: + def test_usage_summary_includes_totals_models_providers(self): + summary = summarise_usage_data(_as_response(SAMPLE_AGGREGATED_RESPONSE)) + assert "$50.25" in summary + assert "gpt-4" in summary + assert "openai" in summary + + def test_usage_summary_handles_empty(self): + assert "no data" in summarise_usage_data(_as_response({"results": [], "metadata": {}})).lower() + + def test_entity_summary_ranks_by_spend(self): + summary = summarise_entity_data(_as_response(SAMPLE_TEAM_RESPONSE), "Team") + assert "Engineering" in summary + assert "Marketing" in summary + # Engineering (higher spend) must appear before Marketing + assert summary.index("Engineering") < summary.index("Marketing") + + def test_entity_summary_empty(self): + assert "No Team usage data" in summarise_entity_data(_as_response({"results": [], "metadata": {}}), "Team") diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageAIChatPanel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageAIChatPanel.test.tsx index 54bf4fc25ce..2166125ab59 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageAIChatPanel.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageAIChatPanel.test.tsx @@ -37,7 +37,7 @@ describe("UsageAIChatPanel", () => { it("should render model selector", () => { renderWithProviders(); - expect(screen.getByText("Select a model (optional, defaults to gpt-4o-mini)")).toBeInTheDocument(); + expect(screen.getByText("Select a model (uses the configured default if left empty)")).toBeInTheDocument(); }); it("should render empty state message when no conversation", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageAIChatPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageAIChatPanel.tsx index 9d6b09353a2..9285c98b780 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageAIChatPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageAIChatPanel.tsx @@ -260,7 +260,7 @@ const UsageAIChatPanel: React.FC = ({ open, onClose, acce {/* Model selector */}