mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
[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:
parent
36264d4764
commit
f30871ef13
11 changed files with 829 additions and 3 deletions
140
docs/my-website/docs/completion/knowledgebase.md
Normal file
140
docs/my-website/docs/completion/knowledgebase.md
Normal 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 |
|
||||
|
|
@ -271,6 +271,7 @@ const sidebars = {
|
|||
"reasoning_content",
|
||||
"completion/prompt_caching",
|
||||
"completion/predict_outputs",
|
||||
"completion/knowledgebase",
|
||||
"completion/prefix",
|
||||
"completion/drop_params",
|
||||
"completion/prompt_formatting",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
287
litellm/integrations/rag_hooks/bedrock_knowledgebase.py
Normal file
287
litellm/integrations/rag_hooks/bedrock_knowledgebase.py
Normal 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
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
141
litellm/types/integrations/rag/bedrock_knowledgebase.py
Normal file
141
litellm/types/integrations/rag/bedrock_knowledgebase.py
Normal 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
|
||||
|
||||
|
||||
#########################################################################
|
||||
#########################################################################
|
||||
#########################################################################
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
146
tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py
Normal file
146
tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py
Normal 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"]
|
||||
|
||||
|
||||
|
|
@ -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 = {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue