diff --git a/docs/my-website/docs/completion/knowledgebase.md b/docs/my-website/docs/completion/knowledgebase.md new file mode 100644 index 00000000000..d527066e16b --- /dev/null +++ b/docs/my-website/docs/completion/knowledgebase.md @@ -0,0 +1,140 @@ +# Using Knowledge Bases with LiteLLM + +LiteLLM integrates with AWS Bedrock Knowledge Bases, allowing your models to access your organization's data for more accurate and contextually relevant responses. + +## Quick Start + +In order to use a Bedrock Knowledge Base with LiteLLM, you need to pass `knowledge_bases` as a parameter to the completion request. Where `knowledge_bases` is a list of Bedrock Knowledge Base IDs. + +### LiteLLM Python SDK + +```python showLineNumbers title="Basic Bedrock Knowledge Base Usage" +import os +import litellm + + +# Make a completion request with knowledge_bases parameter +response = await litellm.acompletion( + model="anthropic/claude-3-5-sonnet", + messages=[{"role": "user", "content": "What is litellm?"}], + knowledge_bases=["YOUR_KNOWLEDGE_BASE_ID"] # e.g., "T37J8R4WTM" +) + +print(response.choices[0].message.content) +``` + +### LiteLLM Proxy + +#### 1. Configure your proxy + +```yaml showLineNumbers title="config.yaml" +model_list: + - model_name: claude-3-5-sonnet + litellm_params: + model: anthropic/claude-3-5-sonnet + api_key: os.environ/ANTHROPIC_API_KEY + +``` + +#### 2. Make a request with knowledge_bases parameter + +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + + + + +```bash showLineNumbers title="Curl Request to LiteLLM Proxy" +curl http://localhost:4000/v1/chat/completions \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer $LITELLM_API_KEY" \ + -d '{ + "model": "claude-3-5-sonnet", + "messages": [{"role": "user", "content": "What is litellm?"}], + "knowledge_bases": ["YOUR_KNOWLEDGE_BASE_ID"] + }' +``` + + + + + +```python showLineNumbers title="OpenAI Python SDK Request" +from openai import OpenAI + +# Initialize client with your LiteLLM proxy URL +client = OpenAI( + base_url="http://localhost:4000", + api_key="your-litellm-api-key" +) + +# Make a completion request with knowledge_bases parameter +response = client.chat.completions.create( + model="claude-3-5-sonnet", + messages=[{"role": "user", "content": "What is litellm?"}], + extra_body={"knowledge_bases": ["YOUR_KNOWLEDGE_BASE_ID"]} +) + +print(response.choices[0].message.content) +``` + + + + +## How It Works + +LiteLLM implements a `BedrockKnowledgeBaseHook` that intercepts your completion requests for handling the integration with Bedrock Knowledge Bases. + +1. You make a completion request with the `knowledge_bases` parameter +2. LiteLLM automatically: + - Uses your last message as the query to retrieve relevant information from the Knowledge Base + - Adds the retrieved context to your conversation + - Sends the augmented messages to the model + +### Example Transformation + +When you pass `knowledge_bases=["YOUR_KNOWLEDGE_BASE_ID"]`, your request flows through these steps: + +**1. Original Request to LiteLLM:** +```json +{ + "model": "anthropic/claude-3-5-sonnet", + "messages": [ + {"role": "user", "content": "What is litellm?"} + ], + "knowledge_bases": ["YOUR_KNOWLEDGE_BASE_ID"] +} +``` + +**2. Request to AWS Bedrock Knowledge Base:** +```json +{ + "retrievalQuery": { + "text": "What is litellm?" + } +} +``` +This is sent to: `https://bedrock-agent-runtime.{aws_region}.amazonaws.com/knowledgebases/YOUR_KNOWLEDGE_BASE_ID/retrieve` + +**3. Final Request to LiteLLM:** +```json +{ + "model": "anthropic/claude-3-5-sonnet", + "messages": [ + {"role": "user", "content": "What is litellm?"}, + {"role": "user", "content": "Context: \n\nLiteLLM is an open-source SDK to simplify LLM API calls across providers (OpenAI, Claude, etc). It provides a standardized interface with robust error handling, streaming, and observability tools."} + ] +} +``` + +This process happens automatically whenever you include the `knowledge_bases` parameter in your request. + +## API Reference + +### LiteLLM Completion Knowledge Base Parameters + +When using the Knowledge Base integration with LiteLLM, you can include the following parameters: + +| Parameter | Type | Description | +|-----------|------|-------------| +| `knowledge_bases` | List[str] | List of Bedrock Knowledge Base IDs to query | diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index 7a648954920..c1573db3504 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -271,6 +271,7 @@ const sidebars = { "reasoning_content", "completion/prompt_caching", "completion/predict_outputs", + "completion/knowledgebase", "completion/prefix", "completion/drop_params", "completion/prompt_formatting", diff --git a/litellm/__init__.py b/litellm/__init__.py index 59c8c78eb9e..fcaf71c60c0 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -115,6 +115,7 @@ _custom_logger_compatible_callbacks_literal = Literal[ "gcs_pubsub", "agentops", "anthropic_cache_control_hook", + "bedrock_knowledgebase_hook", ] logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None _known_custom_logger_compatible_callbacks: List = list( diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 18cb8e8d7f6..c5295ada195 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -77,7 +77,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac model: str, messages: List[AllMessageValues], non_default_params: dict, - prompt_id: str, + prompt_id: Optional[str], prompt_variables: Optional[dict], dynamic_callback_params: StandardCallbackDynamicParams, ) -> Tuple[str, List[AllMessageValues], dict]: diff --git a/litellm/integrations/rag_hooks/bedrock_knowledgebase.py b/litellm/integrations/rag_hooks/bedrock_knowledgebase.py new file mode 100644 index 00000000000..7dddf594e7a --- /dev/null +++ b/litellm/integrations/rag_hooks/bedrock_knowledgebase.py @@ -0,0 +1,287 @@ +# +-------------------------------------------------------------+ +# +# Add Bedrock Knowledge Base Context to your LLM calls +# +# +-------------------------------------------------------------+ +# Thank you users! We ❤️ you! - Krrish & Ishaan + +import json +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple + +from litellm._logging import verbose_logger, verbose_proxy_logger +from litellm.integrations.custom_logger import CustomLogger +from litellm.integrations.custom_prompt_management import CustomPromptManagement +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) +from litellm.types.integrations.rag.bedrock_knowledgebase import ( + BedrockKBContent, + BedrockKBGuardrailConfiguration, + BedrockKBRequest, + BedrockKBResponse, + BedrockKBRetrievalConfiguration, + BedrockKBRetrievalQuery, + BedrockKBRetrievalResult, +) +from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import StandardCallbackDynamicParams +else: + StandardCallbackDynamicParams = Any + + +class BedrockKnowledgeBaseHook(CustomPromptManagement, BaseAWSLLM): + CONTENT_PREFIX_STRING = "Context: \n\n" + + def __init__( + self, + **kwargs, + ): + self.async_handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.LoggingCallback + ) + + # store kwargs as optional_params + self.optional_params = kwargs + + super().__init__(**kwargs) + BaseAWSLLM.__init__(self) + + async def async_get_chat_completion_prompt( + self, + model: str, + messages: List[AllMessageValues], + non_default_params: dict, + prompt_id: Optional[str], + prompt_variables: Optional[dict], + dynamic_callback_params: StandardCallbackDynamicParams, + ) -> Tuple[str, List[AllMessageValues], dict]: + """ + Retrieves the context from the Bedrock Knowledge Base and appends it to the messages. + """ + knowledge_bases = non_default_params.pop("knowledge_bases", None) + if knowledge_bases: + for knowledge_base in knowledge_bases: + response = await self.make_bedrock_kb_retrieve_request( + knowledge_base_id=knowledge_base, + query=self._get_kb_query_from_messages(messages), + ) + verbose_logger.debug(f"Bedrock Knowledge Base Response: {response}") + + context_message = ( + self.get_chat_completion_message_from_bedrock_kb_response(response) + ) + if context_message is not None: + messages.append(context_message) + return model, messages, non_default_params + + def _get_kb_query_from_messages(self, messages: List[AllMessageValues]) -> str: + """ + Uses the text `content` field of the last message in the list of messages + """ + if len(messages) == 0: + return "" + last_message = messages[-1] + last_message_content = last_message.get("content", None) + if last_message_content is None: + return "" + if isinstance(last_message_content, str): + return last_message_content + elif isinstance(last_message_content, list): + return "\n".join([item.get("text", "") for item in last_message_content]) + return "" + + def _prepare_request( + self, + credentials: Any, + data: BedrockKBRequest, + optional_params: dict, + aws_region_name: str, + api_base: str, + extra_headers: Optional[dict] = None, + ) -> Any: + """ + Prepare a signed AWS request. + + Args: + credentials: AWS credentials + data: Request data + optional_params: Additional parameters + aws_region_name: AWS region name + api_base: Base API URL + extra_headers: Additional headers + + Returns: + AWSRequest: A signed AWS request + """ + try: + from botocore.auth import SigV4Auth + from botocore.awsrequest import AWSRequest + except ImportError: + raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") + + sigv4 = SigV4Auth(credentials, "bedrock", aws_region_name) + + encoded_data = json.dumps(data).encode("utf-8") + headers = {"Content-Type": "application/json"} + if extra_headers is not None: + headers = {"Content-Type": "application/json", **extra_headers} + + request = AWSRequest( + method="POST", url=api_base, data=encoded_data, headers=headers + ) + sigv4.add_auth(request) + if extra_headers is not None and "Authorization" in extra_headers: + # prevent sigv4 from overwriting the auth header + request.headers["Authorization"] = extra_headers["Authorization"] + + return request.prepare() + + async def make_bedrock_kb_retrieve_request( + self, + knowledge_base_id: str, + query: str, + guardrail_id: Optional[str] = None, + guardrail_version: Optional[str] = None, + next_token: Optional[str] = None, + retrieval_configuration: Optional[BedrockKBRetrievalConfiguration] = None, + ) -> BedrockKBResponse: + """ + Make a Bedrock Knowledge Base retrieve request. + + Args: + knowledge_base_id (str): The unique identifier of the knowledge base to query + query (str): The query text to search for + guardrail_id (Optional[str]): The guardrail ID to apply + guardrail_version (Optional[str]): The version of the guardrail to apply + next_token (Optional[str]): Token for pagination + retrieval_configuration (Optional[BedrockKBRetrievalConfiguration]): Configuration for the retrieval process + + Returns: + BedrockKBRetrievalResponse: A typed response object containing the retrieval results + """ + from fastapi import HTTPException + + credentials = self.get_credentials() + aws_region_name = self._get_aws_region_name( + optional_params=self.optional_params + ) + + # Prepare request data + request_data: BedrockKBRequest = BedrockKBRequest( + retrievalQuery=BedrockKBRetrievalQuery(text=query), + ) + if next_token: + request_data["nextToken"] = next_token + if retrieval_configuration: + request_data["retrievalConfiguration"] = retrieval_configuration + if guardrail_id and guardrail_version: + request_data["guardrailConfiguration"] = BedrockKBGuardrailConfiguration( + guardrailId=guardrail_id, guardrailVersion=guardrail_version + ) + verbose_logger.debug( + f"Request Data: {json.dumps(request_data, indent=4, default=str)}" + ) + + # Prepare the request + api_base = f"https://bedrock-agent-runtime.{aws_region_name}.amazonaws.com/knowledgebases/{knowledge_base_id}/retrieve" + + prepared_request = self._prepare_request( + credentials=credentials, + data=request_data, + optional_params=self.optional_params, + aws_region_name=aws_region_name, + api_base=api_base, + ) + + verbose_proxy_logger.debug( + "Bedrock Knowledge Base request body: %s, url %s, headers: %s", + request_data, + prepared_request.url, + prepared_request.headers, + ) + + response = await self.async_handler.post( + url=prepared_request.url, + data=prepared_request.body, # type: ignore + headers=prepared_request.headers, # type: ignore + ) + + verbose_proxy_logger.debug("Bedrock Knowledge Base response: %s", response.text) + + if response.status_code == 200: + response_data = response.json() + return BedrockKBResponse(**response_data) + else: + verbose_proxy_logger.error( + "Bedrock Knowledge Base: error in response. Status code: %s, response: %s", + response.status_code, + response.text, + ) + raise HTTPException( + status_code=response.status_code, + detail={ + "error": "Error calling Bedrock Knowledge Base", + "response": response.text, + }, + ) + + @staticmethod + def should_use_prompt_management_hook(non_default_params: Dict) -> bool: + if non_default_params.get("knowledge_bases", None): + return True + return False + + @staticmethod + def get_initialized_custom_logger( + non_default_params: Dict, + ) -> Optional[CustomLogger]: + from litellm.litellm_core_utils.litellm_logging import ( + _init_custom_logger_compatible_class, + ) + + if BedrockKnowledgeBaseHook.should_use_prompt_management_hook( + non_default_params + ): + return _init_custom_logger_compatible_class( + logging_integration="bedrock_knowledgebase_hook", + internal_usage_cache=None, + llm_router=None, + ) + return None + + @staticmethod + def get_chat_completion_message_from_bedrock_kb_response( + response: BedrockKBResponse, + ) -> Optional[ChatCompletionUserMessage]: + """ + Retrieves the context from the Bedrock Knowledge Base response and returns a ChatCompletionUserMessage object. + """ + retrieval_results: Optional[List[BedrockKBRetrievalResult]] = response.get( + "retrievalResults", None + ) + if retrieval_results is None: + return None + + # string to combine the context from the knowledge base + context_string: str = BedrockKnowledgeBaseHook.CONTENT_PREFIX_STRING + for retrieval_result in retrieval_results: + retrieval_result_content: Optional[BedrockKBContent] = ( + retrieval_result.get("content", None) or {} + ) + if retrieval_result_content is None: + continue + retrieval_result_text: Optional[str] = retrieval_result_content.get( + "text", None + ) + if retrieval_result_text is None: + continue + context_string += retrieval_result_text + message = ChatCompletionUserMessage( + role="user", + content=context_string, + ) + return message diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 3f8c0c65e77..61d4f65c879 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -43,6 +43,9 @@ from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.mlflow import MlflowLogger from litellm.integrations.pagerduty.pagerduty import PagerDutyAlerting +from litellm.integrations.rag_hooks.bedrock_knowledgebase import ( + BedrockKnowledgeBaseHook, +) from litellm.litellm_core_utils.get_litellm_params import get_litellm_params from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import ( StandardBuiltInToolCostTracking, @@ -67,6 +70,7 @@ from litellm.types.rerank import RerankResponse from litellm.types.router import CustomPricingLiteLLMParams from litellm.types.utils import ( CallTypes, + DynamicPromptManagementParamLiteral, EmbeddingResponse, ImageResponse, LiteLLMBatch, @@ -465,10 +469,26 @@ class Logging(LiteLLMLoggingBaseClass): """ if prompt_id: return True - if AnthropicCacheControlHook.should_use_anthropic_cache_control_hook( - non_default_params + + if self._should_run_prompt_management_hooks_without_prompt_id( + non_default_params=non_default_params ): return True + + return False + + def _should_run_prompt_management_hooks_without_prompt_id( + self, + non_default_params: Dict, + ) -> bool: + """ + Certain prompt management hooks don't need a `prompt_id` to be passed in, they are triggered by dynamic params + + eg. AnthropicCacheControlHook and BedrockKnowledgeBaseHook both don't require a `prompt_id` to be passed in, they are triggered by dynamic params + """ + for param in non_default_params: + if param in DynamicPromptManagementParamLiteral.list_all_params(): + return True return False def get_chat_completion_prompt( @@ -503,6 +523,38 @@ class Logging(LiteLLMLoggingBaseClass): self.messages = messages return model, messages, non_default_params + async def async_get_chat_completion_prompt( + self, + model: str, + messages: List[AllMessageValues], + non_default_params: Dict, + prompt_id: Optional[str], + prompt_variables: Optional[dict], + prompt_management_logger: Optional[CustomLogger] = None, + ) -> Tuple[str, List[AllMessageValues], dict]: + custom_logger = ( + prompt_management_logger + or self.get_custom_logger_for_prompt_management( + model=model, non_default_params=non_default_params + ) + ) + + if custom_logger: + ( + model, + messages, + non_default_params, + ) = await custom_logger.async_get_chat_completion_prompt( + model=model, + messages=messages, + non_default_params=non_default_params or {}, + prompt_id=prompt_id, + prompt_variables=prompt_variables, + dynamic_callback_params=self.standard_callback_dynamic_params, + ) + self.messages = messages + return model, messages, non_default_params + def get_custom_logger_for_prompt_management( self, model: str, non_default_params: Dict ) -> Optional[CustomLogger]: @@ -547,6 +599,13 @@ class Logging(LiteLLMLoggingBaseClass): ) return anthropic_cache_control_logger + if bedrock_knowledgebase_logger := BedrockKnowledgeBaseHook.get_initialized_custom_logger( + non_default_params + ): + self.model_call_details["prompt_integration"] = ( + bedrock_knowledgebase_logger.__class__.__name__ + ) + return bedrock_knowledgebase_logger return None def get_custom_logger_for_anthropic_cache_control_hook( @@ -2964,6 +3023,13 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 anthropic_cache_control_hook = AnthropicCacheControlHook() _in_memory_loggers.append(anthropic_cache_control_hook) return anthropic_cache_control_hook # type: ignore + elif logging_integration == "bedrock_knowledgebase_hook": + for callback in _in_memory_loggers: + if isinstance(callback, BedrockKnowledgeBaseHook): + return callback + bedrock_knowledgebase_hook = BedrockKnowledgeBaseHook() + _in_memory_loggers.append(bedrock_knowledgebase_hook) + return bedrock_knowledgebase_hook # type: ignore elif logging_integration == "gcs_pubsub": for callback in _in_memory_loggers: if isinstance(callback, GcsPubSubLogger): @@ -3106,6 +3172,10 @@ def get_custom_logger_compatible_class( # noqa: PLR0915 for callback in _in_memory_loggers: if isinstance(callback, AnthropicCacheControlHook): return callback + elif logging_integration == "bedrock_knowledgebase_hook": + for callback in _in_memory_loggers: + if isinstance(callback, BedrockKnowledgeBaseHook): + return callback elif logging_integration == "gcs_pubsub": for callback in _in_memory_loggers: if isinstance(callback, GcsPubSubLogger): diff --git a/litellm/main.py b/litellm/main.py index ec00e1a3491..489ed258803 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -405,6 +405,31 @@ async def acompletion( loop = asyncio.get_event_loop() custom_llm_provider = kwargs.get("custom_llm_provider", None) + + ## PROMPT MANAGEMENT HOOKS ## + ######################################################### + ######################################################### + litellm_logging_obj = kwargs.get("litellm_logging_obj", None) + if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and ( + litellm_logging_obj.should_run_prompt_management_hooks( + prompt_id=kwargs.get("prompt_id", None), non_default_params=kwargs + ) + ): + ( + model, + messages, + _, + ) = await litellm_logging_obj.async_get_chat_completion_prompt( + model=model, + messages=messages, + non_default_params=kwargs, + prompt_id=kwargs.get("prompt_id", None), + prompt_variables=kwargs.get("prompt_variables", None), + ) + + ######################################################### + ######################################################### + # Adjusted to use explicit arguments instead of *args and **kwargs completion_kwargs = { "model": model, diff --git a/litellm/types/integrations/rag/bedrock_knowledgebase.py b/litellm/types/integrations/rag/bedrock_knowledgebase.py new file mode 100644 index 00000000000..dd60ccc198d --- /dev/null +++ b/litellm/types/integrations/rag/bedrock_knowledgebase.py @@ -0,0 +1,141 @@ +from typing import Any, Dict, List, Literal, Optional, TypedDict, Union + + +class BedrockKBLocation(TypedDict, total=False): + """Location information for a retrieved document.""" + + type: str + s3Location: Optional[dict] + webLocation: Optional[dict] + kendraDocumentLocation: Optional[dict] + salesforceLocation: Optional[dict] + sharePointLocation: Optional[dict] + confluenceLocation: Optional[dict] + customDocumentLocation: Optional[dict] + sqlLocation: Optional[dict] + + +class BedrockKBRowValue(TypedDict): + """Row value in a retrieved document.""" + + columnName: str + columnValue: str + type: str + + +class BedrockKBContent(TypedDict, total=False): + """Content of a retrieved document.""" + + type: str + text: Optional[str] + byteContent: Optional[str] + row: Optional[List[BedrockKBRowValue]] + + +class BedrockKBRetrievalResult(TypedDict, total=False): + """Individual result from a knowledge base retrieval.""" + + content: Optional[BedrockKBContent] + location: Optional[BedrockKBLocation] + score: Optional[float] + metadata: Optional[Dict[str, Any]] + + +class BedrockKBResponse(TypedDict, total=False): + """Response from a Bedrock Knowledge Base retrieval request.""" + + guardrailAction: Optional[Literal["INTERVENED", "NONE"]] + nextToken: Optional[str] + retrievalResults: Optional[List[BedrockKBRetrievalResult]] + + +################ Bedrock Knowledge Base Request Types ################# +######################################################################### +######################################################################### + + +class BedrockKBMetadataAttribute(TypedDict, total=False): + """Metadata attribute configuration for implicit filtering.""" + + description: Optional[str] + key: Optional[str] + type: Optional[str] + + +class BedrockKBImplicitFilterConfiguration(TypedDict, total=False): + """Configuration for implicit filtering.""" + + metadataAttributes: Optional[List[BedrockKBMetadataAttribute]] + modelArn: Optional[str] + + +class BedrockKBSelectiveModeConfiguration(TypedDict, total=False): + """Configuration for selective mode in reranking.""" + + pass # This can be expanded based on actual requirements + + +class BedrockKBMetadataConfiguration(TypedDict, total=False): + """Metadata configuration for reranking.""" + + selectionMode: Optional[str] + selectiveModeConfiguration: Optional[BedrockKBSelectiveModeConfiguration] + + +class BedrockKBModelConfiguration(TypedDict, total=False): + """Model configuration for reranking.""" + + additionalModelRequestFields: Optional[Dict[str, Any]] + modelArn: Optional[str] + + +class BedrockKBRerankingConfiguration(TypedDict, total=False): + """Configuration for reranking in vector search.""" + + bedrockRerankingConfiguration: Optional[ + Dict[str, Any] + ] # This could be further typed if needed + type: Optional[str] + + +class BedrockKBVectorSearchConfiguration(TypedDict, total=False): + """Configuration for vector search.""" + + filter: Optional[Dict[str, Any]] + implicitFilterConfiguration: Optional[BedrockKBImplicitFilterConfiguration] + numberOfResults: Optional[int] + overrideSearchType: Optional[str] + rerankingConfiguration: Optional[BedrockKBRerankingConfiguration] + + +class BedrockKBRetrievalConfiguration(TypedDict, total=False): + """Configuration for retrieval.""" + + vectorSearchConfiguration: Optional[BedrockKBVectorSearchConfiguration] + + +class BedrockKBRetrievalQuery(TypedDict, total=False): + """Query structure for retrieval.""" + + text: Optional[str] + + +class BedrockKBGuardrailConfiguration(TypedDict, total=False): + """Configuration for guardrails.""" + + guardrailId: Optional[str] + guardrailVersion: Optional[str] + + +class BedrockKBRequest(TypedDict, total=False): + """Complete request structure for Bedrock Knowledge Base retrieval.""" + + guardrailConfiguration: Optional[BedrockKBGuardrailConfiguration] + nextToken: Optional[str] + retrievalConfiguration: Optional[BedrockKBRetrievalConfiguration] + retrievalQuery: BedrockKBRetrievalQuery + + +######################################################################### +######################################################################### +######################################################################### diff --git a/litellm/types/utils.py b/litellm/types/utils.py index c126f143e8c..18e5ebd43a3 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2263,3 +2263,16 @@ class SpecialEnums(Enum): LLMResponseTypes = Union[ ModelResponse, EmbeddingResponse, ImageResponse, OpenAIFileObject ] + + +class DynamicPromptManagementParamLiteral(str, Enum): + """ + If any of these params are passed, the user is trying to use dynamic prompt management + """ + + CACHE_CONTROL_INJECTION_POINTS = "cache_control_injection_points" + KNOWLEDGE_BASES = "knowledge_bases" + + @classmethod + def list_all_params(cls): + return [param.value for param in cls] diff --git a/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py b/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py new file mode 100644 index 00000000000..4575f602162 --- /dev/null +++ b/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py @@ -0,0 +1,146 @@ +import io +import os +import sys + + +sys.path.insert(0, os.path.abspath("../..")) + +import asyncio +import litellm +import gzip +import json +import logging +import time +from unittest.mock import AsyncMock, patch, Mock + +import pytest + +import litellm +from litellm import completion +from litellm._logging import verbose_logger +from litellm.integrations.rag_hooks.bedrock_knowledgebase import BedrockKnowledgeBaseHook +from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler + + +@pytest.mark.asyncio +async def test_basic_bedrock_knowledgebase_retrieval(): + + bedrock_knowledgebase_hook = BedrockKnowledgeBaseHook() + response = await bedrock_knowledgebase_hook.make_bedrock_kb_retrieve_request( + knowledge_base_id="T37J8R4WTM", + query="what is litellm?", + ) + assert response is not None + + +@pytest.mark.asyncio +async def test_e2e_bedrock_knowledgebase_retrieval_with_completion(): + litellm._turn_on_debug() + client = AsyncHTTPHandler() + + with patch.object(client, "post") as mock_post: + # Mock the response for the LLM call + mock_response = Mock() + mock_response.status_code = 200 + mock_response.headers = {"Content-Type": "application/json"} + mock_response.json = lambda: json.loads(mock_response.text) + mock_post.return_value = mock_response + try: + response = await litellm.acompletion( + model="anthropic/claude-3.5-sonnet", + messages=[{"role": "user", "content": "what is litellm?"}], + knowledge_bases = [ + "T37J8R4WTM" + ], + client=client + ) + except Exception as e: + print(f"Error: {e}") + + # Verify the LLM request was made + mock_post.assert_called_once() + + # Verify the request body + print("call args:", mock_post.call_args) + request_body = mock_post.call_args.kwargs["json"] + print("Request body:", json.dumps(request_body, indent=4, default=str)) + + # Assert content from the knowedge base was applied to the request + + # 1. we should have 2 content blocks, the first is the user message, the second is the context from the knowledge base + content = request_body["messages"][0]["content"] + assert len(content) == 2 + assert content[0]["type"] == "text" + assert content[1]["type"] == "text" + + # 2. the message with the context should have the bedrock knowledge base prefix string + # this helps confirm that the context from the knowledge base was applied to the request + assert BedrockKnowledgeBaseHook.CONTENT_PREFIX_STRING in content[1]["text"] + + + +@pytest.mark.asyncio +async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call(): + """ + Test that the Bedrock Knowledge Base Hook works when making a real llm api call + """ + litellm._turn_on_debug() + async_client = AsyncHTTPHandler() + litellm.callbacks = [BedrockKnowledgeBaseHook()] + response = await litellm.acompletion( + model="anthropic/claude-3-5-haiku-latest", + messages=[{"role": "user", "content": "what is litellm?"}], + knowledge_bases = [ + "T37J8R4WTM" + ], + client=async_client + ) + assert response is not None + + +@pytest.mark.asyncio +async def test_openai_with_knowledge_base_mock_openai(): + """ + Tests that knowledge base content is correctly passed to the OpenAI API call + """ + litellm.callbacks = [BedrockKnowledgeBaseHook()] + litellm.set_verbose = True + from openai import AsyncOpenAI + + client = AsyncOpenAI(api_key="fake-api-key") + + with patch.object( + client.chat.completions.with_raw_response, "create" + ) as mock_client: + try: + await litellm.acompletion( + model="gpt-4", + messages=[{"role": "user", "content": "what is litellm?"}], + knowledge_bases = [ + "T37J8R4WTM" + ], + client=client, + ) + except Exception as e: + print(f"Error: {e}") + + # Verify the API was called + mock_client.assert_called_once() + request_body = mock_client.call_args.kwargs + + # Verify the request contains messages with knowledge base context + assert "messages" in request_body + messages = request_body["messages"] + + # We expect at least 2 messages: + # 1. User message with the question + # 2. User message with the knowledge base context + assert len(messages) >= 2 + + print("request messages:", json.dumps(messages, indent=4, default=str)) + + # assert message[1] is the user message with the knowledge base context + assert messages[1]["role"] == "user" + assert BedrockKnowledgeBaseHook.CONTENT_PREFIX_STRING in messages[1]["content"] + + diff --git a/tests/logging_callback_tests/test_unit_tests_init_callbacks.py b/tests/logging_callback_tests/test_unit_tests_init_callbacks.py index 445c773d999..025f0008596 100644 --- a/tests/logging_callback_tests/test_unit_tests_init_callbacks.py +++ b/tests/logging_callback_tests/test_unit_tests_init_callbacks.py @@ -34,6 +34,7 @@ from litellm.integrations.opentelemetry import OpenTelemetry from litellm.integrations.mlflow import MlflowLogger from litellm.integrations.argilla import ArgillaLogger from litellm.integrations.anthropic_cache_control_hook import AnthropicCacheControlHook +from litellm.integrations.rag_hooks.bedrock_knowledgebase import BedrockKnowledgeBaseHook from litellm.integrations.langfuse.langfuse_prompt_management import ( LangfusePromptManagement, ) @@ -77,6 +78,7 @@ callback_class_str_to_classType = { "gcs_pubsub": GcsPubSubLogger, "anthropic_cache_control_hook": AnthropicCacheControlHook, "agentops": AgentOps, + "bedrock_knowledgebase_hook": BedrockKnowledgeBaseHook, } expected_env_vars = {