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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2RAG Endpoints for LiteLLM Proxy.
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"""
9import base64
10import json
11from collections.abc import Mapping
12from types import MappingProxyType
13from typing import TYPE_CHECKING, Any, Final
15import orjson
16from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
17from fastapi.responses import ORJSONResponse, StreamingResponse
18from starlette.datastructures import UploadFile
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
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
60router: Final = APIRouter()
63def _as_string_keyed_mapping(value: object) -> Mapping[str, object] | None:
64 if isinstance(value, Mapping):
65 return value
66 return None
69def _response_attr(source: object, name: str) -> object:
70 return getattr(source, name, None)
73def _upstream_status_code(error: Exception) -> int:
74 code: Final = getattr(error, "status_code", None)
75 return code if isinstance(code, int) else 500
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 )
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))
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)]
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()
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 )
135 return vector_store_ids
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 )
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
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)
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 )
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 )
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.
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
218 Returns:
219 Dictionary with file metadata (file_id, filename, file_url, ingested_at, etc.)
220 """
221 from datetime import datetime, timezone
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
230 # Extract file information from file_data tuple
231 filename = None
232 file_size = None
233 content_type = None
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
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 }
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
254 return file_entry
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.
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
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 )
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
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
303 vector_store_config: Final = ingest_options.get("vector_store", {})
304 custom_llm_provider: Final = vector_store_config.get("custom_llm_provider")
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")
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
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 )
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 )
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
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)
340 # Initialize metadata with first file
341 initial_metadata: Final = {"ingested_files": [file_entry]}
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"
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 )
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)
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 )
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
373 # Update the vector store
374 from litellm.proxy.utils import safe_dumps
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 )
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)
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
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.
411 Supports:
412 - Form: file + request JSON in form field
413 - JSON body for URL-based ingestion
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.
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", "")
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] = {}
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)
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 "")
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")
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")
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")
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 )
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 )
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 )
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 )
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 )
525 return ingest_options, secured_file_data, file_url, file_id
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.
548 Supports form upload (for files) or JSON body (for URLs).
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 ```
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 ```
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 )
586 try:
587 # Parse request
588 ingest_options, file_data, file_url, file_id = await parse_rag_ingest_request(
589 request, scanner=EicarTestMalwareScanner()
590 )
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 )
604 resolved_stores: Final = await _authorize_nested_vector_store_ids(
605 payload=ingest_options,
606 user_api_key_dict=user_api_key_dict,
607 )
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)})
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 }
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 )
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 )
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 )
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 )
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 )
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 )
687 return response
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 )
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.
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
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 ```
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 )
771 try:
772 # Parse request body
773 data: Final = await _read_request_body(request)
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)
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 )
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 }
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 )
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 )
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 )
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 )
872 if stream:
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 )
887 return await open_sse_before_first_byte(
888 produce_stream(),
889 ping_interval_seconds=ttft_keepalive_interval(data, llm_router),
890 )
892 response: Final = await query()
893 fastapi_response.headers.update(custom_headers_for(response))
894 return response
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 )