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

1""" 

2AI Usage Chat - uses LLM tool calling to answer questions about 

3usage/spend data by querying the aggregated daily activity endpoints. 

4""" 

5 

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 

10 

11from typing_extensions import ReadOnly, TypedDict 

12 

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 

20 

21# --------------------------------------------------------------------------- 

22# Constants 

23# --------------------------------------------------------------------------- 

24 

25USAGE_AI_TEMPERATURE: Final = 0.2 

26 

27TABLE_DAILY_USER_SPEND: Final = "litellm_dailyuserspend" 

28TABLE_DAILY_TEAM_SPEND: Final = "litellm_dailyteamspend" 

29TABLE_DAILY_TAG_SPEND: Final = "litellm_dailytagspend" 

30 

31ENTITY_FIELD_USER: Final = "user_id" 

32ENTITY_FIELD_TEAM: Final = "team_id" 

33ENTITY_FIELD_TAG: Final = "tag" 

34 

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 

40 

41# --------------------------------------------------------------------------- 

42# Types 

43# --------------------------------------------------------------------------- 

44 

45 

46class SSEStatusEvent(TypedDict): 

47 type: Literal["status"] 

48 message: str 

49 

50 

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 

58 

59 

60class SSEChunkEvent(TypedDict): 

61 type: Literal["chunk"] 

62 content: str 

63 

64 

65class SSEDoneEvent(TypedDict): 

66 type: Literal["done"] 

67 

68 

69class SSEErrorEvent(TypedDict): 

70 type: Literal["error"] 

71 message: str 

72 

73 

74SSEEvent = SSEStatusEvent | SSEToolCallEvent | SSEChunkEvent | SSEDoneEvent | SSEErrorEvent 

75 

76 

77class _EntityEntry(TypedDict, total=False): 

78 metrics: ReadOnly[Mapping[str, float]] 

79 metadata: ReadOnly[Mapping[str, str]] 

80 

81 

82class _DayDump(TypedDict, total=False): 

83 breakdown: ReadOnly[Mapping[str, Mapping[str, _EntityEntry]]] 

84 

85 

86class _EntityTotal(NamedTuple): 

87 """Running per-entity totals accumulated while summarising a usage dump.""" 

88 

89 alias: str 

90 spend: float 

91 requests: float 

92 tokens: float 

93 

94 

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

100 

101 

102class _ToolFunctionDef(TypedDict): 

103 name: ReadOnly[str] 

104 description: ReadOnly[str] 

105 parameters: ReadOnly[Mapping[str, object]] 

106 

107 

108class _ToolDef(TypedDict): 

109 type: ReadOnly[str] 

110 function: ReadOnly[_ToolFunctionDef] 

111 

112 

113class ToolHandler(TypedDict): 

114 fetch: Callable[..., Awaitable[_UsageDump]] 

115 summarise: Callable[[_UsageDump], str] 

116 label: str 

117 

118 

119# --------------------------------------------------------------------------- 

120# Tool definitions (OpenAI function-calling schema) 

121# --------------------------------------------------------------------------- 

122 

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} 

127 

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} 

150 

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} 

172 

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} 

194 

195TOOLS_BASE: Final = [_TOOL_USAGE] 

196TOOLS_ADMIN: Final = [_TOOL_USAGE, _TOOL_TEAM, _TOOL_TAG] 

197 

198 

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 

202 

203 

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) 

218 

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) 

225 

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) 

229 

230 

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()}" 

235 

236 

237# keep a public reference for test assertions 

238SYSTEM_PROMPT: Final = _SYSTEM_PROMPT_BASE 

239 

240# --------------------------------------------------------------------------- 

241# Data fetchers 

242# --------------------------------------------------------------------------- 

243 

244 

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()] 

249 

250 

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 

266 

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 ) 

292 

293 

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") 

304 

305 

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") 

315 

316 

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") 

326 

327 

328# --------------------------------------------------------------------------- 

329# Summarisers — convert raw JSON to concise text the LLM can reason over 

330# --------------------------------------------------------------------------- 

331 

332 

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 

346 

347 

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]] 

355 

356 

357def _summarise_usage_data(data: _UsageDump) -> str: 

358 meta: Final = data.get("metadata", {}) 

359 results: Final = data.get("results", []) 

360 

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 ) 

368 

369 models: Final = _accumulate_breakdown(results, "models", ["spend", "api_requests", "total_tokens"]) 

370 providers: Final = _accumulate_breakdown(results, "providers", ["spend", "api_requests"]) 

371 

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 ) 

382 

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) 

387 

388 

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." 

394 

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 ) 

406 

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) 

412 

413 

414# --------------------------------------------------------------------------- 

415# Tool dispatch registry 

416# --------------------------------------------------------------------------- 

417 

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} 

435 

436 

437# --------------------------------------------------------------------------- 

438# SSE streaming 

439# --------------------------------------------------------------------------- 

440 

441 

442def _sse(event: SSEEvent) -> str: 

443 return f"data: {json.dumps(event)}\n\n" 

444 

445 

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 

477 

478 

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) 

490 

491 

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) 

501 

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 

504 

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 

514 

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"})) 

522 

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"})) 

530 

531 chat_messages.append({"role": "tool", "tool_call_id": tc.id, "content": tool_result}) 

532 

533 

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..."}) 

537 

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}) 

548 

549 

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 ] 

563 

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] 

574 

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 

580 

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"}) 

588 

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 )