[Feat] Add support for using Bedrock Knowledge Bases with LiteLLM /chat/completions requests (#10413)

* add make_bedrock_kb_retrieve_request

* working bedrock KB hook

* working bedrock KB hook

* test_openai_with_knowledge_base_mock_openai

* fix linting

* fix BedrockKnowledgeBaseHook

* docs using bedrock kb with litellm

* docs kb with litellm

* fix bedrock kb test

* DynamicPromptManagementParamLiteral

* fix _should_run_prompt_management_hooks_without_prompt_id

* test_init_custom_logger_compatible_class_as_callback
This commit is contained in:
Ishaan Jaff 2025-04-29 17:29:02 -07:00 • committed by GitHub
parent 36264d4764
commit f30871ef13
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 829 additions and 3 deletions

View file

@ -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';
<Tabs>
<TabItem value="curl" label="Curl">
```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"]
}'
```
</TabItem>
<TabItem value="openai-sdk" label="OpenAI Python SDK">
```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)
```
</TabItem>
</Tabs>
## 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 |

View file

@ -271,6 +271,7 @@ const sidebars = {
"reasoning_content",
"completion/prompt_caching",
"completion/predict_outputs",
"completion/knowledgebase",
"completion/prefix",
"completion/drop_params",
"completion/prompt_formatting",

View file

@ -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(

View file

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

View file

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

View file

@ -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):

View file

@ -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,

View file

@ -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
#########################################################################
#########################################################################
#########################################################################

View file

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

View file

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

View file

@ -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 = {