Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/rag_endpoints/endpoints.py: 22%

285 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1""" 

2RAG Endpoints for LiteLLM Proxy. 

3 

4Provides: 

5- /rag/ingest: All-in-one document ingestion pipeline (Upload -> Chunk -> Embed -> Vector Store) 

6- /rag/query: RAG query pipeline (Search -> Rerank -> LLM Completion) 

7""" 

8 

9import base64 

10import json 

11from collections.abc import Mapping 

12from types import MappingProxyType 

13from typing import TYPE_CHECKING, Any, Final 

14 

15import orjson 

16from fastapi import APIRouter, Depends, HTTPException, Request, Response, status 

17from fastapi.responses import ORJSONResponse, StreamingResponse 

18from starlette.datastructures import UploadFile 

19 

20import litellm 

21from litellm._logging import verbose_proxy_logger 

22from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH 

23from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( 

24 LiteLLM_ManagedVectorStore, 

25) 

26from litellm.proxy._types import * 

27from litellm.proxy.auth.auth_utils import is_request_body_safe 

28from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth 

29from litellm.proxy.common_request_processing import ( 

30 ProxyBaseLLMRequestProcessing, 

31 open_sse_before_first_byte, 

32 ttft_keepalive_interval, 

33) 

34from litellm.proxy.common_utils.http_parsing_utils import ( 

35 _read_request_body, 

36 _safe_get_request_headers, 

37 get_form_data, 

38) 

39from litellm.proxy.rag_endpoints.upload_security import ( 

40 MAX_UPLOAD_SIZE_BYTES, 

41 EicarTestMalwareScanner, 

42 MalwareScanner, 

43 RejectedUpload, 

44 validate_upload, 

45) 

46from litellm.proxy.vector_store_endpoints.endpoints import ( 

47 build_request_data_from_managed_vector_store, 

48 reject_caller_embedding_selection_params, 

49) 

50from litellm.proxy.vector_store_endpoints.utils import ( 

51 assert_user_can_access_vector_store_id, 

52) 

53from litellm.rag.main import get_ingestion_class 

54from litellm.repositories.table_repositories import ManagedVectorStoresRepository 

55from litellm.types.utils import ModelResponse 

56 

57if TYPE_CHECKING: 57 ↛ 58line 57 didn't jump to line 58 because the condition on line 57 was never true

58 from litellm.proxy.utils import PrismaClient 

59 

60router: Final = APIRouter() 

61 

62 

63def _as_string_keyed_mapping(value: object) -> Mapping[str, object] | None: 

64 if isinstance(value, Mapping): 

65 return value 

66 return None 

67 

68 

69def _response_attr(source: object, name: str) -> object: 

70 return getattr(source, name, None) 

71 

72 

73def _upstream_status_code(error: Exception) -> int: 

74 code: Final = getattr(error, "status_code", None) 

75 return code if isinstance(code, int) else 500 

76 

77 

78def _raise_vector_store_scan_depth_exceeded() -> None: 

79 raise HTTPException( 

80 status_code=400, 

81 detail={"error": f"Max depth of {DEFAULT_MAX_RECURSE_DEPTH} exceeded while scanning vector_store_id values"}, 

82 ) 

83 

84 

85def _append_payload_to_scan_stack( 

86 payload_stack: list[tuple[object, int]], 

87 value: object, 

88 next_depth: int, 

89) -> None: 

90 if isinstance(value, dict): 

91 if next_depth > DEFAULT_MAX_RECURSE_DEPTH: 

92 _raise_vector_store_scan_depth_exceeded() 

93 payload_stack.append((value, next_depth)) 

94 elif isinstance(value, list): 

95 if next_depth > DEFAULT_MAX_RECURSE_DEPTH: 

96 if any(isinstance(item, (dict, list)) for item in value): 

97 _raise_vector_store_scan_depth_exceeded() 

98 return 

99 payload_stack.append((value, next_depth)) 

100 

101 

102def _collect_vector_store_ids_from_payload(payload: object) -> set[str]: 

103 vector_store_ids: Final[set[str]] = set() 

104 payload_stack: Final = [(payload, 0)] 

105 

106 while payload_stack: 

107 current_payload, depth = payload_stack.pop() 

108 if depth > DEFAULT_MAX_RECURSE_DEPTH: 

109 _raise_vector_store_scan_depth_exceeded() 

110 

111 if isinstance(current_payload, dict): 

112 for key, value in current_payload.items(): 

113 if key == "vector_store_id": 

114 if not isinstance(value, str) or not value: 

115 raise HTTPException( 

116 status_code=400, 

117 detail={"error": "vector_store_id must be a non-empty string"}, 

118 ) 

119 vector_store_ids.add(value) 

120 continue 

121 if isinstance(value, (dict, list)): 

122 _append_payload_to_scan_stack( 

123 payload_stack=payload_stack, 

124 value=value, 

125 next_depth=depth + 1, 

126 ) 

127 elif isinstance(current_payload, list): 

128 for item in current_payload: 

129 _append_payload_to_scan_stack( 

130 payload_stack=payload_stack, 

131 value=item, 

132 next_depth=depth + 1, 

133 ) 

134 

135 return vector_store_ids 

136 

137 

138async def _authorize_nested_vector_store_ids( 

139 payload: object, 

140 user_api_key_dict: UserAPIKeyAuth, 

141) -> Mapping[str, LiteLLM_ManagedVectorStore]: 

142 """Authorize every nested vector store id and return the managed stores it resolved.""" 

143 return MappingProxyType( 

144 { 

145 vector_store_id: store 

146 for vector_store_id in sorted(_collect_vector_store_ids_from_payload(payload)) 

147 if ( 

148 store := await assert_user_can_access_vector_store_id( 

149 vector_store_id=vector_store_id, 

150 user_api_key_dict=user_api_key_dict, 

151 ) 

152 ) 

153 is not None 

154 } 

155 ) 

156 

157 

158def _ingest_provider_error(vector_store_config: Mapping[str, object]) -> str | None: 

159 provider: Final = vector_store_config.get("custom_llm_provider", "openai") 

160 if not isinstance(provider, str): 

161 return "custom_llm_provider must be a string" 

162 try: 

163 get_ingestion_class(provider) 

164 except ValueError as error: 

165 return str(error) 

166 return None 

167 

168 

169_MANAGED_STORE_CALLER_OPTIONS: Final = frozenset( 

170 { 

171 "vector_store_id", 

172 "data_source_id", 

173 "wait_for_ingestion", 

174 "ingestion_timeout", 

175 "custom_metadata", 

176 "file_description", 

177 "max_embedding_requests_per_min", 

178 } 

179) 

180 

181 

182def _caller_vector_store_options( 

183 request_vector_store_config: Mapping[str, object], 

184 managed_store: LiteLLM_ManagedVectorStore | None, 

185) -> Mapping[str, object]: 

186 if managed_store is None: 

187 return request_vector_store_config 

188 return MappingProxyType( 

189 {key: value for key, value in request_vector_store_config.items() if key in _MANAGED_STORE_CALLER_OPTIONS} 

190 ) 

191 

192 

193def _managed_store_overrides(managed_store: LiteLLM_ManagedVectorStore | None) -> Mapping[str, object]: 

194 if managed_store is None: 

195 return MappingProxyType({}) 

196 return MappingProxyType( 

197 { 

198 key: value 

199 for key, value in build_request_data_from_managed_vector_store(managed_store).items() 

200 if value is not None 

201 } 

202 ) 

203 

204 

205def _build_file_metadata_entry( 

206 response: object, 

207 file_data: tuple[str, bytes, str] | None = None, 

208 file_url: str | None = None, 

209) -> Mapping[str, str | int | None]: 

210 """ 

211 Build a file metadata entry for storing in vector_store_metadata. 

212 

213 Args: 

214 response: The response from litellm.aingest containing file_id 

215 file_data: Optional tuple of (filename, content, content_type) 

216 file_url: Optional URL if file was ingested from URL 

217 

218 Returns: 

219 Dictionary with file metadata (file_id, filename, file_url, ingested_at, etc.) 

220 """ 

221 from datetime import datetime, timezone 

222 

223 # Extract file_id from response 

224 mapping_response: Final = _as_string_keyed_mapping(response) 

225 raw_file_id: Final = ( 

226 mapping_response.get("file_id") if mapping_response is not None else _response_attr(response, "file_id") 

227 ) 

228 file_id: Final = raw_file_id if isinstance(raw_file_id, str) else None 

229 

230 # Extract file information from file_data tuple 

231 filename = None 

232 file_size = None 

233 content_type = None 

234 

235 if file_data: 

236 filename = file_data[0] 

237 file_size = len(file_data[1]) if len(file_data) > 1 else None 

238 content_type = file_data[2] if len(file_data) > 2 else None 

239 

240 # Build file metadata entry 

241 file_entry: Final[dict[str, str | int | None]] = { 

242 "file_id": file_id, 

243 "filename": filename, 

244 "file_url": file_url, 

245 "ingested_at": datetime.now(timezone.utc).isoformat(), 

246 } 

247 

248 # Add optional fields if available 

249 if file_size is not None: 

250 file_entry["file_size"] = file_size 

251 if content_type is not None: 

252 file_entry["content_type"] = content_type 

253 

254 return file_entry 

255 

256 

257async def _save_vector_store_to_db_from_rag_ingest( 

258 response: object, 

259 ingest_options: Mapping[str, dict[str, str | None]], 

260 prisma_client: "PrismaClient", 

261 user_api_key_dict: UserAPIKeyAuth, 

262 file_data: tuple[str, bytes, str] | None = None, 

263 file_url: str | None = None, 

264 *, 

265 store_is_managed: bool = False, 

266) -> None: 

267 """ 

268 Helper function to save a newly created vector store from RAG ingest to the database. 

269 

270 This function: 

271 - Extracts vector store ID and config from the ingest response 

272 - Checks if the vector store already exists in the database 

273 - Creates a new database entry if it doesn't exist and the store is not registry-managed 

274 - Adds the vector store to the registry 

275 - Tracks team_id and user_id for access control 

276 

277 Args: 

278 response: The response from litellm.aingest() 

279 ingest_options: The ingest options containing vector store config 

280 prisma_client: The Prisma database client 

281 user_api_key_dict: User API key authentication info 

282 store_is_managed: True when the requested id resolved to a managed store, so a missing row means 

283 the store is config-registered and must not get a database row 

284 """ 

285 from litellm.proxy.vector_store_endpoints.management_endpoints import ( 

286 create_vector_store_in_db, 

287 ) 

288 

289 # Handle both dict and object responses 

290 mapping_response: Final = _as_string_keyed_mapping(response) 

291 if mapping_response is not None: 

292 vector_store_id = mapping_response.get("vector_store_id") 

293 elif hasattr(response, "vector_store_id"): 

294 vector_store_id = _response_attr(response, "vector_store_id") 

295 else: 

296 verbose_proxy_logger.warning("Unable to extract vector_store_id from response type: %s", type(response)) 

297 return 

298 

299 if vector_store_id is None or not isinstance(vector_store_id, str): 

300 verbose_proxy_logger.warning("Vector store ID is None or not a string, skipping database save") 

301 return 

302 

303 vector_store_config: Final = ingest_options.get("vector_store", {}) 

304 custom_llm_provider: Final = vector_store_config.get("custom_llm_provider") 

305 

306 # Extract litellm_vector_store_params for custom name and description 

307 litellm_vector_store_params: Final = ingest_options.get("litellm_vector_store_params", {}) 

308 custom_vector_store_name: Final = litellm_vector_store_params.get("vector_store_name") 

309 custom_vector_store_description: Final = litellm_vector_store_params.get("vector_store_description") 

310 

311 # Extract provider-specific params from vector_store_config to save as litellm_params 

312 # This ensures params like aws_region_name, embedding_model, etc. are available for search 

313 provider_specific_params: Final = {} 

314 excluded_keys: Final = {"custom_llm_provider", "vector_store_id"} 

315 for key, value in vector_store_config.items(): 

316 if key not in excluded_keys and value is not None: 

317 provider_specific_params[key] = value 

318 

319 # Build file metadata entry using helper 

320 file_entry: Final = _build_file_metadata_entry( 

321 response=response, 

322 file_data=file_data, 

323 file_url=file_url, 

324 ) 

325 

326 try: 

327 # Check if vector store already exists in database 

328 existing_vector_store: Final = await ManagedVectorStoresRepository(prisma_client).table.find_unique( 

329 where={"vector_store_id": vector_store_id} 

330 ) 

331 

332 if existing_vector_store is None and store_is_managed: 

333 verbose_proxy_logger.info("Vector store %s is config-registered, skipping database save", vector_store_id) 

334 return 

335 

336 # Only create if it doesn't exist 

337 if existing_vector_store is None: 

338 verbose_proxy_logger.info("Saving newly created vector store %s to database", vector_store_id) 

339 

340 # Initialize metadata with first file 

341 initial_metadata: Final = {"ingested_files": [file_entry]} 

342 

343 # Use custom name if provided, otherwise default 

344 vector_store_name: Final = custom_vector_store_name or f"RAG Vector Store - {vector_store_id[:8]}" 

345 vector_store_description: Final = custom_vector_store_description or "Created via RAG ingest endpoint" 

346 

347 await create_vector_store_in_db( 

348 vector_store_id=vector_store_id, 

349 custom_llm_provider=custom_llm_provider or "openai", 

350 prisma_client=prisma_client, 

351 vector_store_name=vector_store_name, 

352 vector_store_description=vector_store_description, 

353 vector_store_metadata=initial_metadata, 

354 litellm_params=(provider_specific_params if provider_specific_params else None), 

355 team_id=user_api_key_dict.team_id, 

356 user_id=user_api_key_dict.user_id, 

357 ) 

358 

359 verbose_proxy_logger.info("Vector store %s saved to database successfully", vector_store_id) 

360 else: 

361 verbose_proxy_logger.info("Vector store %s already exists, appending file to metadata", vector_store_id) 

362 

363 # Update existing vector store with new file 

364 stored_metadata: Final = existing_vector_store.vector_store_metadata or {} 

365 existing_metadata: dict[str, object] = ( 

366 json.loads(stored_metadata) if isinstance(stored_metadata, str) else stored_metadata 

367 ) 

368 

369 previous_files: Final = existing_metadata.get("ingested_files", []) 

370 ingested_files: Final = [*previous_files, file_entry] if isinstance(previous_files, list) else [file_entry] 

371 existing_metadata["ingested_files"] = ingested_files 

372 

373 # Update the vector store 

374 from litellm.proxy.utils import safe_dumps 

375 

376 await ManagedVectorStoresRepository(prisma_client).table.update( 

377 where={"vector_store_id": vector_store_id}, 

378 data={"vector_store_metadata": safe_dumps(existing_metadata)}, 

379 ) 

380 

381 verbose_proxy_logger.info( 

382 "Added file %s to vector store %s metadata", 

383 file_entry.get("filename") or file_entry.get("file_url", "Unknown"), 

384 vector_store_id, 

385 ) 

386 except Exception as db_error: 

387 # Log the error but don't fail the request since ingestion succeeded 

388 verbose_proxy_logger.exception("Failed to save vector store %s to database: %s", vector_store_id, db_error) 

389 

390 

391def _secure_uploaded_file( 

392 file_data: tuple[str, bytes, str], 

393 scanner: MalwareScanner, 

394) -> tuple[str, bytes, str]: 

395 validation: Final = validate_upload(content=file_data[1], scanner=scanner) 

396 if isinstance(validation, RejectedUpload): 

397 raise HTTPException( 

398 status_code=400, 

399 detail={"error": validation.message, "reason": validation.reason.value}, 

400 ) 

401 return validation.safe_filename, file_data[1], validation.content_type 

402 

403 

404async def parse_rag_ingest_request( 

405 request: Request, 

406 scanner: MalwareScanner, 

407) -> tuple[dict[str, Any], tuple[str, bytes, str] | None, str | None, str | None]: 

408 """ 

409 Parse RAG ingest request. 

410 

411 Supports: 

412 - Form: file + request JSON in form field 

413 - JSON body for URL-based ingestion 

414 

415 Uploaded file bytes are validated against the vector-store upload controls 

416 (size limit, format allowlist with content inspection, archive rejection, 

417 and the injected malware scanner) and given a server-generated filename 

418 before they are returned. 

419 

420 Returns: 

421 Tuple of (ingest_options, file_data, file_url, file_id) 

422 """ 

423 headers: Final = _safe_get_request_headers(request) 

424 content_type = headers.get("content-type", "") 

425 

426 file_data: tuple[str, bytes, str] | None = None 

427 file_url: str | None = None 

428 file_id: str | None = None 

429 ingest_options: dict[str, Any] = {} 

430 

431 if "multipart/form-data" in content_type: 431 ↛ 433line 431 didn't jump to line 433 because the condition on line 431 was never true

432 # Form upload 

433 form_data: Final = await get_form_data(request) 

434 

435 # Get file 

436 file_obj = form_data.get("file") 

437 if isinstance(file_obj, UploadFile): 

438 file_content = await file_obj.read(MAX_UPLOAD_SIZE_BYTES + 1) 

439 file_data = (file_obj.filename or "", file_content, file_obj.content_type or "") 

440 

441 # Parse JSON from 'request' form field (contains full request body as JSON) 

442 request_json_str: Final[str | bytes | None] = form_data.get("request") 

443 if request_json_str: 

444 request_data: Final = orjson.loads(request_json_str) 

445 ingest_options = request_data.get("ingest_options", {}) 

446 file_url = request_data.get("file_url") 

447 file_id = request_data.get("file_id") 

448 

449 else: 

450 # JSON body 

451 data: Final = await _read_request_body(request) 

452 ingest_options = data.get("ingest_options", {}) 

453 file_url = data.get("file_url") 

454 file_id = data.get("file_id") 

455 

456 # Handle base64-encoded file in JSON body 

457 file_obj = data.get("file") 

458 if file_obj and isinstance(file_obj, dict): 458 ↛ 459line 458 didn't jump to line 459 because the condition on line 458 was never true

459 filename: Final = file_obj.get("filename") 

460 content_b64: Final = file_obj.get("content") 

461 content_type = file_obj.get("content_type", "application/octet-stream") 

462 

463 if filename and content_b64: 

464 try: 

465 file_content = base64.b64decode(content_b64) 

466 file_data = (filename, file_content, content_type) 

467 except Exception as e: 

468 raise HTTPException( 

469 status_code=400, 

470 detail={"error": f"Invalid base64 content: {e}"}, 

471 ) 

472 

473 # Validate 

474 if file_data is None and file_url is None and file_id is None: 474 ↛ 480line 474 didn't jump to line 480 because the condition on line 474 was always true

475 raise HTTPException( 

476 status_code=400, 

477 detail={"error": "Must provide file, file_url, or file_id"}, 

478 ) 

479 

480 secured_file_data: Final[tuple[str, bytes, str] | None] = ( 

481 _secure_uploaded_file(file_data, scanner) if file_data is not None else None 

482 ) 

483 

484 if "vector_store" not in ingest_options: 

485 raise HTTPException( 

486 status_code=400, 

487 detail={"error": "ingest_options must contain 'vector_store' configuration"}, 

488 ) 

489 

490 # Credential fields must come from server configuration, not user requests. 

491 # Accepting user-supplied credentials (e.g. vertex_credentials with 

492 # type=external_account + credential_source.file=/proc/1/environ) allows 

493 # any authenticated user to exfiltrate host secrets via SSRF through 

494 # google-auth's identity_pool credential refresh. 

495 # api_base is also blocked: a user-controlled base URL causes the server 

496 # to send its configured provider credentials to an attacker endpoint. 

497 _BLOCKED_VECTOR_STORE_CREDENTIAL_PARAMS: Final = { 

498 "vertex_credentials", 

499 "vertex_ai_credentials", 

500 "aws_access_key_id", 

501 "aws_secret_access_key", 

502 "aws_session_token", 

503 "aws_web_identity_token", 

504 "aws_role_name", 

505 "aws_session_name", 

506 "aws_profile_name", 

507 "aws_sts_endpoint", 

508 "aws_external_id", 

509 "azure_ad_token", 

510 "api_key", 

511 "api_base", 

512 } 

513 vector_store_opts: Final[object] = ingest_options.get("vector_store", {}) 

514 if isinstance(vector_store_opts, dict): 

515 for field in _BLOCKED_VECTOR_STORE_CREDENTIAL_PARAMS: 

516 if field in vector_store_opts: 

517 raise HTTPException( 

518 status_code=400, 

519 detail={ 

520 "error": f"'{field}' cannot be set in ingest_options.vector_store. " 

521 "Credentials must be configured server-side." 

522 }, 

523 ) 

524 

525 return ingest_options, secured_file_data, file_url, file_id 

526 

527 

528@router.post( 

529 "/v1/rag/ingest", 

530 dependencies=[Depends(user_api_key_auth)], 

531 response_class=ORJSONResponse, 

532 tags=["rag"], 

533) 

534@router.post( 

535 "/rag/ingest", 

536 dependencies=[Depends(user_api_key_auth)], 

537 response_class=ORJSONResponse, 

538 tags=["rag"], 

539) 

540async def rag_ingest( 

541 request: Request, 

542 fastapi_response: Response, 

543 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

544): 

545 """ 

546 RAG Ingest endpoint - all-in-one document ingestion pipeline. 

547 

548 Supports form upload (for files) or JSON body (for URLs). 

549 

550 ## Form upload (for files): 

551 ```bash 

552 curl -X POST "http://localhost:4000/v1/rag/ingest" \\ 

553 -H "Authorization: Bearer sk-1234" \\ 

554 -F file="@document.pdf" \\ 

555 -F 'ingest_options={"vector_store": {"custom_llm_provider": "openai"}}' 

556 ``` 

557 

558 ## JSON body (for URLs): 

559 ```bash 

560 curl -X POST "http://localhost:4000/v1/rag/ingest" \\ 

561 -H "Authorization: Bearer sk-1234" \\ 

562 -H "Content-Type: application/json" \\ 

563 -d '{ 

564 "file_url": "https://example.com/document.pdf", 

565 "ingest_options": {"vector_store": {"custom_llm_provider": "openai"}} 

566 }' 

567 ``` 

568 

569 ## Bedrock: 

570 ```bash 

571 curl -X POST "http://localhost:4000/v1/rag/ingest" \\ 

572 -H "Authorization: Bearer sk-1234" \\ 

573 -F file="@document.pdf" \\ 

574 -F 'ingest_options={"vector_store": {"custom_llm_provider": "bedrock"}}' 

575 ``` 

576 """ 

577 from litellm.proxy.proxy_server import ( 

578 add_litellm_data_to_request, 

579 general_settings, 

580 llm_router, 

581 prisma_client, 

582 proxy_config, 

583 version, 

584 ) 

585 

586 try: 

587 # Parse request 

588 ingest_options, file_data, file_url, file_id = await parse_rag_ingest_request( 

589 request, scanner=EicarTestMalwareScanner() 

590 ) 

591 

592 # INTERNAL_USER_VIEW_ONLY can ingest to existing vector stores only 

593 if user_api_key_dict.user_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value and not ingest_options.get( 

594 "vector_store", {} 

595 ).get("vector_store_id"): 

596 raise HTTPException( 

597 status_code=status.HTTP_403_FORBIDDEN, 

598 detail={ 

599 "error": "internal_user_viewer role can only ingest files to an existing vector store. " 

600 "Provide 'vector_store_id' in ingest_options.vector_store." 

601 }, 

602 ) 

603 

604 resolved_stores: Final = await _authorize_nested_vector_store_ids( 

605 payload=ingest_options, 

606 user_api_key_dict=user_api_key_dict, 

607 ) 

608 

609 request_vector_store_config: Final = ingest_options.get("vector_store", {}) 

610 try: 

611 is_request_body_safe( 

612 request_body=request_vector_store_config, 

613 general_settings=general_settings, 

614 llm_router=llm_router, 

615 model="", 

616 ) 

617 except ValueError as e: 

618 raise HTTPException(status_code=400, detail={"error": str(e)}) 

619 

620 managed_store: Final = resolved_stores.get(request_vector_store_config.get("vector_store_id")) 

621 merged_vector_store_config: Final = { # mutable-ok: ingestion classes mutate it when loading credentials 

622 **_caller_vector_store_options(request_vector_store_config, managed_store), 

623 **_managed_store_overrides(managed_store), 

624 } 

625 merged_ingest_options: Final = { # mutable-ok: litellm.aingest takes a plain dict payload 

626 **ingest_options, 

627 "vector_store": merged_vector_store_config, 

628 } 

629 

630 provider_error: Final = _ingest_provider_error(merged_vector_store_config) 

631 if provider_error is not None: 

632 raise HTTPException( 

633 status_code=400, 

634 detail={"error": provider_error}, # mutable-ok: FastAPI serializes the detail as JSON 

635 ) 

636 

637 # Add litellm data 

638 request_data: dict[str, Any] = {} 

639 request_data = await add_litellm_data_to_request( 

640 data=request_data, 

641 request=request, 

642 general_settings=general_settings, 

643 user_api_key_dict=user_api_key_dict, 

644 version=version, 

645 proxy_config=proxy_config, 

646 ) 

647 

648 verbose_proxy_logger.debug( 

649 "RAG Ingest - options: %s, custom_llm_provider: %s", 

650 ingest_options, 

651 merged_vector_store_config.get("custom_llm_provider", "openai"), 

652 ) 

653 

654 # Call ingest 

655 response: Final = await litellm.aingest( 

656 ingest_options=merged_ingest_options, 

657 file_data=file_data, 

658 file_url=file_url, 

659 file_id=file_id, 

660 router=llm_router, 

661 **request_data, 

662 ) 

663 

664 # Save vector store to database if it was newly created and prisma_client is available 

665 verbose_proxy_logger.debug( 

666 "RAG Ingest - Checking database save conditions: prisma_client=%s, response=%s, response_type=%s", 

667 prisma_client is not None, 

668 response is not None, 

669 type(response), 

670 ) 

671 

672 if prisma_client is not None and response is not None: 

673 await _save_vector_store_to_db_from_rag_ingest( 

674 response=response, 

675 ingest_options=ingest_options, 

676 prisma_client=prisma_client, 

677 user_api_key_dict=user_api_key_dict, 

678 file_data=file_data, 

679 file_url=file_url, 

680 store_is_managed=managed_store is not None, 

681 ) 

682 else: 

683 verbose_proxy_logger.warning( 

684 "Skipping database save: prisma_client=%s, response=%s", prisma_client is not None, response is not None 

685 ) 

686 

687 return response 

688 

689 except HTTPException: 

690 raise 

691 except Exception as e: 

692 verbose_proxy_logger.exception("RAG Ingest failed: %s", e) 

693 raise HTTPException( 

694 status_code=500, 

695 detail={"error": str(e)}, 

696 ) 

697 

698 

699@router.post( 

700 "/v1/rag/query", 

701 dependencies=[Depends(user_api_key_auth)], 

702 response_class=ORJSONResponse, 

703 tags=["rag"], 

704) 

705@router.post( 

706 "/rag/query", 

707 dependencies=[Depends(user_api_key_auth)], 

708 response_class=ORJSONResponse, 

709 tags=["rag"], 

710) 

711async def rag_query( 

712 request: Request, 

713 fastapi_response: Response, 

714 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

715): 

716 """ 

717 RAG Query endpoint - search vector store, optionally rerank, and generate LLM response. 

718 

719 This endpoint: 

720 1. Extracts the query from the last user message 

721 2. Searches the vector store for relevant context 

722 3. Optionally reranks the results 

723 4. Generates an LLM response with the retrieved context 

724 

725 ## Example Request: 

726 ```bash 

727 curl -X POST "http://localhost:4000/v1/rag/query" \\ 

728 -H "Authorization: Bearer sk-1234" \\ 

729 -H "Content-Type: application/json" \\ 

730 -d '{ 

731 "model": "gpt-4o-mini", 

732 "messages": [{"role": "user", "content": "What is LiteLLM?"}], 

733 "retrieval_config": { 

734 "vector_store_id": "vs_abc123", 

735 "custom_llm_provider": "openai", 

736 "top_k": 5 

737 } 

738 }' 

739 ``` 

740 

741 ## With Reranking: 

742 ```bash 

743 curl -X POST "http://localhost:4000/v1/rag/query" \\ 

744 -H "Authorization: Bearer sk-1234" \\ 

745 -H "Content-Type: application/json" \\ 

746 -d '{ 

747 "model": "gpt-4o-mini", 

748 "messages": [{"role": "user", "content": "What is LiteLLM?"}], 

749 "retrieval_config": { 

750 "vector_store_id": "vs_abc123", 

751 "custom_llm_provider": "openai", 

752 "top_k": 10 

753 }, 

754 "rerank": { 

755 "enabled": true, 

756 "model": "cohere/rerank-english-v3.0", 

757 "top_n": 3 

758 } 

759 }' 

760 ``` 

761 """ 

762 from litellm.proxy.proxy_server import ( 

763 add_litellm_data_to_request, 

764 general_settings, 

765 llm_router, 

766 proxy_config, 

767 select_data_generator, 

768 version, 

769 ) 

770 

771 try: 

772 # Parse request body 

773 data: Final = await _read_request_body(request) 

774 

775 # Extract required fields 

776 model: Final = data.get("model") 

777 messages: Final = data.get("messages") 

778 retrieval_config: Final = data.get("retrieval_config") 

779 rerank: Final = data.get("rerank") 

780 stream: Final = data.get("stream", False) 

781 

782 # Validate required fields 

783 if not model: 783 ↛ 788line 783 didn't jump to line 788 because the condition on line 783 was always true

784 raise HTTPException( 

785 status_code=400, 

786 detail={"error": "model is required"}, 

787 ) 

788 if not messages: 

789 raise HTTPException( 

790 status_code=400, 

791 detail={"error": "messages is required"}, 

792 ) 

793 if not retrieval_config: 

794 raise HTTPException( 

795 status_code=400, 

796 detail={"error": "retrieval_config is required"}, 

797 ) 

798 if not isinstance(retrieval_config, dict): 

799 raise HTTPException( 

800 status_code=400, 

801 detail={"error": "retrieval_config must be an object"}, 

802 ) 

803 if "vector_store_id" not in retrieval_config: 

804 raise HTTPException( 

805 status_code=400, 

806 detail={"error": "retrieval_config must contain 'vector_store_id'"}, 

807 ) 

808 reject_caller_embedding_selection_params(payload=retrieval_config, source="retrieval_config") 

809 resolved_stores: Final = await _authorize_nested_vector_store_ids( 

810 payload=retrieval_config, 

811 user_api_key_dict=user_api_key_dict, 

812 ) 

813 

814 # Merge litellm-managed vector store params (provider, region, embedding 

815 # model, credentials, ...) from the registry: the same source the direct 

816 # /vector_stores/{id}/search endpoint uses. Store-managed keys win on 

817 # conflict so callers cannot override the store's provider or credentials. 

818 managed_store: Final = resolved_stores.get(retrieval_config["vector_store_id"]) 

819 store_data: Final = ( 

820 build_request_data_from_managed_vector_store(managed_store) 

821 if managed_store is not None 

822 else MappingProxyType({}) 

823 ) 

824 merged_retrieval_config: Final = { 

825 **retrieval_config, 

826 **store_data, 

827 } 

828 

829 # Add litellm data 

830 request_data: dict[str, object] = {} 

831 request_data = await add_litellm_data_to_request( 

832 data=request_data, 

833 request=request, 

834 general_settings=general_settings, 

835 user_api_key_dict=user_api_key_dict, 

836 version=version, 

837 proxy_config=proxy_config, 

838 ) 

839 

840 verbose_proxy_logger.debug( 

841 "RAG Query - model: %s, vector_store_id: %s, custom_llm_provider: %s", 

842 model, 

843 retrieval_config["vector_store_id"], 

844 merged_retrieval_config.get("custom_llm_provider"), 

845 ) 

846 

847 async def query() -> ModelResponse: 

848 return await litellm.aquery( 

849 model=model, 

850 messages=messages, 

851 retrieval_config=merged_retrieval_config, 

852 vector_store_params=store_data, 

853 rerank=rerank, 

854 stream=stream, 

855 router=llm_router, 

856 **request_data, 

857 ) 

858 

859 def custom_headers_for(response: ModelResponse) -> Mapping[str, str]: 

860 hidden_params: Final = getattr(response, "_hidden_params", {}) or {} 

861 return ProxyBaseLLMRequestProcessing.get_custom_headers( 

862 user_api_key_dict=user_api_key_dict, 

863 call_id=hidden_params.get("litellm_call_id", None) or "", 

864 model_id=hidden_params.get("model_id", None) or "", 

865 cache_key=hidden_params.get("cache_key", None) or "", 

866 api_base=hidden_params.get("api_base", None) or "", 

867 version=version, 

868 response_cost=hidden_params.get("response_cost", None), 

869 request_data=request_data, 

870 ) 

871 

872 if stream: 

873 

874 async def produce_stream() -> StreamingResponse: 

875 response: Final = await query() 

876 return StreamingResponse( 

877 select_data_generator( 

878 response=response, 

879 user_api_key_dict=user_api_key_dict, 

880 request_data=request_data, 

881 request=request, 

882 ), 

883 media_type="text/event-stream", 

884 headers=custom_headers_for(response), 

885 ) 

886 

887 return await open_sse_before_first_byte( 

888 produce_stream(), 

889 ping_interval_seconds=ttft_keepalive_interval(data, llm_router), 

890 ) 

891 

892 response: Final = await query() 

893 fastapi_response.headers.update(custom_headers_for(response)) 

894 return response 

895 

896 except HTTPException: 

897 raise 

898 except Exception as e: 

899 verbose_proxy_logger.exception("RAG Query failed: %s", e) 

900 raise HTTPException( 

901 status_code=_upstream_status_code(e), 

902 detail={"error": str(e)}, 

903 )