From 3a82334762d647f6a79e9e8cf8c45399301631ef Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 27 Aug 2024 15:55:43 -0700 Subject: [PATCH 1/8] add basic cohere rerank --- litellm/__init__.py | 1 + litellm/rerank_api/types.py | 20 ++++++++++++++++++++ 2 files changed, 21 insertions(+) create mode 100644 litellm/rerank_api/types.py diff --git a/litellm/__init__.py b/litellm/__init__.py index a627061cfe1..581db4fcb88 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -934,6 +934,7 @@ from .proxy.proxy_cli import run_server from .router import Router from .assistants.main import * from .batches.main import * +from .rerank_api.main import * from .fine_tuning.main import * from .files.main import * from .scheduler import * diff --git a/litellm/rerank_api/types.py b/litellm/rerank_api/types.py new file mode 100644 index 00000000000..9d53cf278cb --- /dev/null +++ b/litellm/rerank_api/types.py @@ -0,0 +1,20 @@ +""" +LiteLLM Follows the cohere API format for the re rank API +https://docs.cohere.com/reference/rerank + +""" + +from pydantic import BaseModel + + +class RerankRequest(BaseModel): + model: str + query: str + top_n: int + documents: list[str] + + +class RerankResponse(BaseModel): + id: str + results: list[dict] # Contains index and relevance_score + meta: dict # Contains api_version and billed_units From b8bc185bd5ce4d3edb2692112a59faeb286b7802 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 27 Aug 2024 15:57:36 -0700 Subject: [PATCH 2/8] add main cohere ai rerank handler + test --- litellm/rerank_api/main.py | 111 +++++++++++++++++++++++++++++++++++++ 1 file changed, 111 insertions(+) create mode 100644 litellm/rerank_api/main.py diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py new file mode 100644 index 00000000000..c65dca503ec --- /dev/null +++ b/litellm/rerank_api/main.py @@ -0,0 +1,111 @@ +import asyncio +import contextvars +from functools import partial +from typing import Any, Coroutine, Dict, List, Literal, Optional, Union + +import litellm +from litellm import get_secret +from litellm._logging import verbose_logger +from litellm.llms.cohere.rerank import CohereRerank +from litellm.types.router import * +from litellm.utils import supports_httpx_timeout + +from .types import RerankRequest, RerankResponse + +####### ENVIRONMENT VARIABLES ################### +# Initialize any necessary instances or variables here +cohere_rerank = CohereRerank() +################################################# + + +async def arerank( + model: str, + query: str, + documents: List[str], + custom_llm_provider: Literal["cohere", "together_ai"] = "cohere", + top_n: int = 3, + **kwargs, +) -> Dict[str, Any]: + """ + Async: Reranks a list of documents based on their relevance to the query + """ + try: + loop = asyncio.get_event_loop() + kwargs["arerank"] = True + + func = partial( + rerank, model, query, documents, custom_llm_provider, top_n, **kwargs + ) + + ctx = contextvars.copy_context() + func_with_context = partial(ctx.run, func) + init_response = await loop.run_in_executor(None, func_with_context) + + if asyncio.iscoroutine(init_response): + response = await init_response + else: + response = init_response + return response + except Exception as e: + raise e + + +def rerank( + model: str, + query: str, + documents: List[str], + custom_llm_provider: Literal["cohere", "together_ai"] = "cohere", + top_n: int = 3, + **kwargs, +) -> Union[RerankResponse, Coroutine[Any, Any, RerankResponse]]: + """ + Reranks a list of documents based on their relevance to the query + """ + try: + _is_async = kwargs.pop("arerank", False) is True + optional_params = GenericLiteLLMParams(**kwargs) + + # Implement rerank logic here based on the custom_llm_provider + if custom_llm_provider == "cohere": + # Implement Cohere rerank logic + cohere_key = ( + optional_params.api_key + or litellm.cohere_key + or get_secret("COHERE_API_KEY") + or get_secret("CO_API_KEY") + or litellm.api_key + ) + + if cohere_key is None: + raise ValueError( + "Cohere API key is required, please set 'COHERE_API_KEY' in your environment" + ) + + api_base = ( + optional_params.api_base + or litellm.api_base + or get_secret("COHERE_API_BASE") + or "https://api.cohere.ai/v1/generate" + ) + + headers: Dict = litellm.headers or {} + + response = cohere_rerank.rerank( + model=model, + query=query, + documents=documents, + top_n=top_n, + api_key=cohere_key, + ) + pass + elif custom_llm_provider == "together_ai": + # Implement Together AI rerank logic + pass + else: + raise ValueError(f"Unsupported provider: {custom_llm_provider}") + + # Placeholder return + return response + except Exception as e: + verbose_logger.error(f"Error in rerank: {str(e)}") + raise e From dc42ad0021c377f8ced952417026525eb3b18bc4 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 27 Aug 2024 16:25:54 -0700 Subject: [PATCH 3/8] add tg ai rerank support --- litellm/llms/cohere/rerank.py | 44 ++++++++++++++++++++++++++ litellm/llms/togetherai/rerank.py | 52 +++++++++++++++++++++++++++++++ litellm/rerank_api/main.py | 44 ++++++++++++++++++++++---- 3 files changed, 134 insertions(+), 6 deletions(-) create mode 100644 litellm/llms/cohere/rerank.py create mode 100644 litellm/llms/togetherai/rerank.py diff --git a/litellm/llms/cohere/rerank.py b/litellm/llms/cohere/rerank.py new file mode 100644 index 00000000000..a547ea21890 --- /dev/null +++ b/litellm/llms/cohere/rerank.py @@ -0,0 +1,44 @@ +""" +Re rank api + +LiteLLM supports the re rank API format, no paramter transformation occurs +""" + +import httpx +from pydantic import BaseModel + +from litellm.llms.base import BaseLLM +from litellm.llms.custom_httpx.http_handler import ( + _get_async_httpx_client, + _get_httpx_client, +) +from litellm.rerank_api.types import RerankRequest, RerankResponse + + +class CohereRerank(BaseLLM): + def rerank( + self, + model: str, + api_key: str, + query: str, + documents: list[str], + top_n: int = 3, + ) -> RerankResponse: + client = _get_httpx_client() + request_data = RerankRequest( + model=model, query=query, top_n=top_n, documents=documents + ) + + response = client.post( + "https://api.cohere.com/v1/rerank", + headers={ + "accept": "application/json", + "content-type": "application/json", + "Authorization": f"bearer {api_key}", + }, + json=request_data.dict(), + ) + + return RerankResponse(**response.json()) + + pass diff --git a/litellm/llms/togetherai/rerank.py b/litellm/llms/togetherai/rerank.py new file mode 100644 index 00000000000..b4020fc6511 --- /dev/null +++ b/litellm/llms/togetherai/rerank.py @@ -0,0 +1,52 @@ +""" +Re rank api + +LiteLLM supports the re rank API format, no paramter transformation occurs +""" + +import httpx +from pydantic import BaseModel + +from litellm.llms.base import BaseLLM +from litellm.llms.custom_httpx.http_handler import ( + _get_async_httpx_client, + _get_httpx_client, +) +from litellm.rerank_api.types import RerankRequest, RerankResponse + + +class TogetherAIRerank(BaseLLM): + def rerank( + self, + model: str, + api_key: str, + query: str, + documents: list[str], + top_n: int = 3, + ) -> RerankResponse: + client = _get_httpx_client() + + request_data = RerankRequest( + model=model, query=query, top_n=top_n, documents=documents + ) + + response = client.post( + "https://api.together.xyz/v1/rerank", + headers={ + "accept": "application/json", + "content-type": "application/json", + "authorization": f"Bearer {api_key}", + }, + json=request_data.dict(), + ) + + _json_response = response.json() + response = RerankResponse( + id=_json_response.get("id"), + results=_json_response.get("results"), + meta=_json_response.get("meta") or {}, + ) + + return response + + pass diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index c65dca503ec..bb0094d001f 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -7,6 +7,7 @@ import litellm from litellm import get_secret from litellm._logging import verbose_logger from litellm.llms.cohere.rerank import CohereRerank +from litellm.llms.togetherai.rerank import TogetherAIRerank from litellm.types.router import * from litellm.utils import supports_httpx_timeout @@ -15,6 +16,7 @@ from .types import RerankRequest, RerankResponse ####### ENVIRONMENT VARIABLES ################### # Initialize any necessary instances or variables here cohere_rerank = CohereRerank() +together_rerank = TogetherAIRerank() ################################################# @@ -54,7 +56,7 @@ def rerank( model: str, query: str, documents: List[str], - custom_llm_provider: Literal["cohere", "together_ai"] = "cohere", + custom_llm_provider: Optional[Literal["cohere", "together_ai"]] = None, top_n: int = 3, **kwargs, ) -> Union[RerankResponse, Coroutine[Any, Any, RerankResponse]]: @@ -65,11 +67,21 @@ def rerank( _is_async = kwargs.pop("arerank", False) is True optional_params = GenericLiteLLMParams(**kwargs) + model, _custom_llm_provider, dynamic_api_key, api_base = ( + litellm.get_llm_provider( + model=model, + custom_llm_provider=custom_llm_provider, + api_base=optional_params.api_base, + api_key=optional_params.api_key, + ) + ) + # Implement rerank logic here based on the custom_llm_provider - if custom_llm_provider == "cohere": + if _custom_llm_provider == "cohere": # Implement Cohere rerank logic cohere_key = ( - optional_params.api_key + dynamic_api_key + or optional_params.api_key or litellm.cohere_key or get_secret("COHERE_API_KEY") or get_secret("CO_API_KEY") @@ -98,11 +110,31 @@ def rerank( api_key=cohere_key, ) pass - elif custom_llm_provider == "together_ai": + elif _custom_llm_provider == "together_ai": # Implement Together AI rerank logic - pass + together_key = ( + dynamic_api_key + or optional_params.api_key + or litellm.togetherai_api_key + or get_secret("TOGETHERAI_API_KEY") + or litellm.api_key + ) + + if together_key is None: + raise ValueError( + "TogetherAI API key is required, please set 'TOGETHERAI_API_KEY' in your environment" + ) + + response = together_rerank.rerank( + model=model, + query=query, + documents=documents, + top_n=top_n, + api_key=together_key, + ) + else: - raise ValueError(f"Unsupported provider: {custom_llm_provider}") + raise ValueError(f"Unsupported provider: {_custom_llm_provider}") # Placeholder return return response From 255ad865cd217d5878272ca4233ed2a633694ae1 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 27 Aug 2024 16:27:55 -0700 Subject: [PATCH 4/8] add rerank api tests --- litellm/tests/test_rerank.py | 93 ++++++++++++++++++++++++++++++++++++ 1 file changed, 93 insertions(+) create mode 100644 litellm/tests/test_rerank.py diff --git a/litellm/tests/test_rerank.py b/litellm/tests/test_rerank.py new file mode 100644 index 00000000000..946bfbb970f --- /dev/null +++ b/litellm/tests/test_rerank.py @@ -0,0 +1,93 @@ +import json +import os +import sys +import traceback + +from dotenv import load_dotenv + +load_dotenv() +import io +import os + +sys.path.insert( + 0, os.path.abspath("../..") +) # Adds the parent directory to the system path + +import os +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +import litellm +from litellm import RateLimitError, Timeout, completion, completion_cost, embedding +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler + + +def assert_response_shape(response, custom_llm_provider): + expected_response_shape = {"id": str, "results": list, "meta": dict} + + expected_results_shape = {"index": int, "relevance_score": float} + + expected_meta_shape = {"api_version": dict, "billed_units": dict} + + expected_api_version_shape = {"version": str} + + expected_billed_units_shape = {"search_units": int} + + assert isinstance(response.id, expected_response_shape["id"]) + assert isinstance(response.results, expected_response_shape["results"]) + for result in response.results: + assert isinstance(result["index"], expected_results_shape["index"]) + assert isinstance( + result["relevance_score"], expected_results_shape["relevance_score"] + ) + assert isinstance(response.meta, expected_response_shape["meta"]) + + if custom_llm_provider == "cohere": + + assert isinstance( + response.meta["api_version"], expected_meta_shape["api_version"] + ) + assert isinstance( + response.meta["api_version"]["version"], + expected_api_version_shape["version"], + ) + assert isinstance( + response.meta["billed_units"], expected_meta_shape["billed_units"] + ) + assert isinstance( + response.meta["billed_units"]["search_units"], + expected_billed_units_shape["search_units"], + ) + + +def test_basic_rerank(): + response = litellm.rerank( + model="cohere/rerank-english-v3.0", + query="hello", + documents=["hello", "world"], + top_n=3, + ) + + print("re rank response: ", response) + + assert response.id is not None + assert response.results is not None + + assert_response_shape(response, custom_llm_provider="cohere") + + +def test_basic_rerank_together_ai(): + response = litellm.rerank( + model="together_ai/Salesforce/Llama-Rank-V1", + query="hello", + documents=["hello", "world"], + top_n=3, + ) + + print("re rank response: ", response) + + assert response.id is not None + assert response.results is not None + + assert_response_shape(response, custom_llm_provider="together_ai") From f33dfe0b95deb62fe413cc3667e390ac7f1cd0c2 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 27 Aug 2024 16:45:39 -0700 Subject: [PATCH 5/8] add rerank params --- litellm/llms/cohere/rerank.py | 22 ++++++++++++++++++---- litellm/llms/togetherai/rerank.py | 27 +++++++++++++++++++++++---- litellm/rerank_api/main.py | 13 +++++++++++-- litellm/rerank_api/types.py | 11 ++++++++--- 4 files changed, 60 insertions(+), 13 deletions(-) diff --git a/litellm/llms/cohere/rerank.py b/litellm/llms/cohere/rerank.py index a547ea21890..0c00ea03c96 100644 --- a/litellm/llms/cohere/rerank.py +++ b/litellm/llms/cohere/rerank.py @@ -4,6 +4,8 @@ Re rank api LiteLLM supports the re rank API format, no paramter transformation occurs """ +from typing import Any, Dict, List, Optional, Union + import httpx from pydantic import BaseModel @@ -21,14 +23,26 @@ class CohereRerank(BaseLLM): model: str, api_key: str, query: str, - documents: list[str], - top_n: int = 3, + documents: list[Union[str, Dict[str, Any]]], + top_n: Optional[int] = None, + rank_fields: Optional[List[str]] = None, + return_documents: Optional[bool] = True, + max_chunks_per_doc: Optional[int] = None, ) -> RerankResponse: client = _get_httpx_client() + request_data = RerankRequest( - model=model, query=query, top_n=top_n, documents=documents + model=model, + query=query, + top_n=top_n, + documents=documents, + rank_fields=rank_fields, + return_documents=return_documents, + max_chunks_per_doc=max_chunks_per_doc, ) + request_data_dict = request_data.dict(exclude_none=True) + response = client.post( "https://api.cohere.com/v1/rerank", headers={ @@ -36,7 +50,7 @@ class CohereRerank(BaseLLM): "content-type": "application/json", "Authorization": f"bearer {api_key}", }, - json=request_data.dict(), + json=request_data_dict, ) return RerankResponse(**response.json()) diff --git a/litellm/llms/togetherai/rerank.py b/litellm/llms/togetherai/rerank.py index b4020fc6511..8a5a4668527 100644 --- a/litellm/llms/togetherai/rerank.py +++ b/litellm/llms/togetherai/rerank.py @@ -4,6 +4,8 @@ Re rank api LiteLLM supports the re rank API format, no paramter transformation occurs """ +from typing import Any, Dict, List, Optional, Union + import httpx from pydantic import BaseModel @@ -21,15 +23,28 @@ class TogetherAIRerank(BaseLLM): model: str, api_key: str, query: str, - documents: list[str], - top_n: int = 3, + documents: list[Union[str, Dict[str, Any]]], + top_n: Optional[int] = None, + rank_fields: Optional[List[str]] = None, + return_documents: Optional[bool] = True, + max_chunks_per_doc: Optional[int] = None, ) -> RerankResponse: client = _get_httpx_client() request_data = RerankRequest( - model=model, query=query, top_n=top_n, documents=documents + model=model, + query=query, + top_n=top_n, + documents=documents, + rank_fields=rank_fields, + return_documents=return_documents, ) + # exclude None values from request_data + request_data_dict = request_data.dict(exclude_none=True) + if max_chunks_per_doc is not None: + raise ValueError("TogetherAI does not support max_chunks_per_doc") + response = client.post( "https://api.together.xyz/v1/rerank", headers={ @@ -37,10 +52,14 @@ class TogetherAIRerank(BaseLLM): "content-type": "application/json", "authorization": f"Bearer {api_key}", }, - json=request_data.dict(), + json=request_data_dict, ) + if response.status_code != 200: + raise Exception(response.text) + _json_response = response.json() + response = RerankResponse( id=_json_response.get("id"), results=_json_response.get("results"), diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index bb0094d001f..6d3a27f549b 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -55,9 +55,12 @@ async def arerank( def rerank( model: str, query: str, - documents: List[str], + documents: List[Union[str, Dict[str, Any]]], custom_llm_provider: Optional[Literal["cohere", "together_ai"]] = None, - top_n: int = 3, + top_n: Optional[int] = None, + rank_fields: Optional[List[str]] = None, + return_documents: Optional[bool] = True, + max_chunks_per_doc: Optional[int] = None, **kwargs, ) -> Union[RerankResponse, Coroutine[Any, Any, RerankResponse]]: """ @@ -107,6 +110,9 @@ def rerank( query=query, documents=documents, top_n=top_n, + rank_fields=rank_fields, + return_documents=return_documents, + max_chunks_per_doc=max_chunks_per_doc, api_key=cohere_key, ) pass @@ -130,6 +136,9 @@ def rerank( query=query, documents=documents, top_n=top_n, + rank_fields=rank_fields, + return_documents=return_documents, + max_chunks_per_doc=max_chunks_per_doc, api_key=together_key, ) diff --git a/litellm/rerank_api/types.py b/litellm/rerank_api/types.py index 9d53cf278cb..605e25a2ecb 100644 --- a/litellm/rerank_api/types.py +++ b/litellm/rerank_api/types.py @@ -4,17 +4,22 @@ https://docs.cohere.com/reference/rerank """ +from typing import List, Optional, Union + from pydantic import BaseModel class RerankRequest(BaseModel): model: str query: str - top_n: int - documents: list[str] + top_n: Optional[int] = None + documents: List[Union[str, dict]] + rank_fields: Optional[List[str]] = None + return_documents: Optional[bool] = None + max_chunks_per_doc: Optional[int] = None class RerankResponse(BaseModel): id: str - results: list[dict] # Contains index and relevance_score + results: List[dict] # Contains index and relevance_score meta: dict # Contains api_version and billed_units From b3892b871ddafa8209b4c17b8c7b392a8f094d4d Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 27 Aug 2024 17:02:48 -0700 Subject: [PATCH 6/8] add async support for rerank --- litellm/llms/cohere/rerank.py | 26 +++++++++-- litellm/llms/togetherai/rerank.py | 32 +++++++++++++ litellm/rerank_api/main.py | 24 ++++++++-- litellm/tests/test_rerank.py | 78 ++++++++++++++++++++++--------- 4 files changed, 130 insertions(+), 30 deletions(-) diff --git a/litellm/llms/cohere/rerank.py b/litellm/llms/cohere/rerank.py index 0c00ea03c96..4ef523e3a14 100644 --- a/litellm/llms/cohere/rerank.py +++ b/litellm/llms/cohere/rerank.py @@ -28,9 +28,8 @@ class CohereRerank(BaseLLM): rank_fields: Optional[List[str]] = None, return_documents: Optional[bool] = True, max_chunks_per_doc: Optional[int] = None, + _is_async: Optional[bool] = False, # New parameter ) -> RerankResponse: - client = _get_httpx_client() - request_data = RerankRequest( model=model, query=query, @@ -43,6 +42,10 @@ class CohereRerank(BaseLLM): request_data_dict = request_data.dict(exclude_none=True) + if _is_async: + return self.async_rerank(request_data_dict, api_key) # type: ignore # Call async method + + client = _get_httpx_client() response = client.post( "https://api.cohere.com/v1/rerank", headers={ @@ -55,4 +58,21 @@ class CohereRerank(BaseLLM): return RerankResponse(**response.json()) - pass + async def async_rerank( + self, + request_data_dict: Dict[str, Any], + api_key: str, + ) -> RerankResponse: + client = _get_async_httpx_client() + + response = await client.post( + "https://api.cohere.com/v1/rerank", + headers={ + "accept": "application/json", + "content-type": "application/json", + "Authorization": f"bearer {api_key}", + }, + json=request_data_dict, + ) + + return RerankResponse(**response.json()) diff --git a/litellm/llms/togetherai/rerank.py b/litellm/llms/togetherai/rerank.py index 8a5a4668527..32a8cdcfdc6 100644 --- a/litellm/llms/togetherai/rerank.py +++ b/litellm/llms/togetherai/rerank.py @@ -28,6 +28,7 @@ class TogetherAIRerank(BaseLLM): rank_fields: Optional[List[str]] = None, return_documents: Optional[bool] = True, max_chunks_per_doc: Optional[int] = None, + _is_async: Optional[bool] = False, ) -> RerankResponse: client = _get_httpx_client() @@ -45,6 +46,9 @@ class TogetherAIRerank(BaseLLM): if max_chunks_per_doc is not None: raise ValueError("TogetherAI does not support max_chunks_per_doc") + if _is_async: + return self.async_rerank(request_data_dict, api_key) # type: ignore # Call async method + response = client.post( "https://api.together.xyz/v1/rerank", headers={ @@ -68,4 +72,32 @@ class TogetherAIRerank(BaseLLM): return response + async def async_rerank( # New async method + self, + request_data_dict: Dict[str, Any], + api_key: str, + ) -> RerankResponse: + client = _get_async_httpx_client() # Use async client + + response = await client.post( + "https://api.together.xyz/v1/rerank", + headers={ + "accept": "application/json", + "content-type": "application/json", + "authorization": f"Bearer {api_key}", + }, + json=request_data_dict, + ) + + if response.status_code != 200: + raise Exception(response.text) + + _json_response = response.json() + + return RerankResponse( + id=_json_response.get("id"), + results=_json_response.get("results"), + meta=_json_response.get("meta") or {}, + ) # Return response + pass diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index 6d3a27f549b..968b9b562cd 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -23,11 +23,14 @@ together_rerank = TogetherAIRerank() async def arerank( model: str, query: str, - documents: List[str], - custom_llm_provider: Literal["cohere", "together_ai"] = "cohere", - top_n: int = 3, + documents: List[Union[str, Dict[str, Any]]], + custom_llm_provider: Optional[Literal["cohere", "together_ai"]] = None, + top_n: Optional[int] = None, + rank_fields: Optional[List[str]] = None, + return_documents: Optional[bool] = True, + max_chunks_per_doc: Optional[int] = None, **kwargs, -) -> Dict[str, Any]: +) -> Union[RerankResponse, Coroutine[Any, Any, RerankResponse]]: """ Async: Reranks a list of documents based on their relevance to the query """ @@ -36,7 +39,16 @@ async def arerank( kwargs["arerank"] = True func = partial( - rerank, model, query, documents, custom_llm_provider, top_n, **kwargs + rerank, + model, + query, + documents, + custom_llm_provider, + top_n, + rank_fields, + return_documents, + max_chunks_per_doc, + **kwargs, ) ctx = contextvars.copy_context() @@ -114,6 +126,7 @@ def rerank( return_documents=return_documents, max_chunks_per_doc=max_chunks_per_doc, api_key=cohere_key, + _is_async=_is_async, ) pass elif _custom_llm_provider == "together_ai": @@ -140,6 +153,7 @@ def rerank( return_documents=return_documents, max_chunks_per_doc=max_chunks_per_doc, api_key=together_key, + _is_async=_is_async, ) else: diff --git a/litellm/tests/test_rerank.py b/litellm/tests/test_rerank.py index 946bfbb970f..a0127063f91 100644 --- a/litellm/tests/test_rerank.py +++ b/litellm/tests/test_rerank.py @@ -61,33 +61,67 @@ def assert_response_shape(response, custom_llm_provider): ) -def test_basic_rerank(): - response = litellm.rerank( - model="cohere/rerank-english-v3.0", - query="hello", - documents=["hello", "world"], - top_n=3, - ) +@pytest.mark.asyncio() +@pytest.mark.parametrize("sync_mode", [True, False]) +async def test_basic_rerank(sync_mode): + if sync_mode is True: + response = litellm.rerank( + model="cohere/rerank-english-v3.0", + query="hello", + documents=["hello", "world"], + top_n=3, + ) - print("re rank response: ", response) + print("re rank response: ", response) - assert response.id is not None - assert response.results is not None + assert response.id is not None + assert response.results is not None - assert_response_shape(response, custom_llm_provider="cohere") + assert_response_shape(response, custom_llm_provider="cohere") + else: + response = await litellm.arerank( + model="cohere/rerank-english-v3.0", + query="hello", + documents=["hello", "world"], + top_n=3, + ) + + print("async re rank response: ", response) + + assert response.id is not None + assert response.results is not None + + assert_response_shape(response, custom_llm_provider="cohere") -def test_basic_rerank_together_ai(): - response = litellm.rerank( - model="together_ai/Salesforce/Llama-Rank-V1", - query="hello", - documents=["hello", "world"], - top_n=3, - ) +@pytest.mark.asyncio() +@pytest.mark.parametrize("sync_mode", [True, False]) +async def test_basic_rerank_together_ai(sync_mode): + if sync_mode is True: + response = litellm.rerank( + model="together_ai/Salesforce/Llama-Rank-V1", + query="hello", + documents=["hello", "world"], + top_n=3, + ) - print("re rank response: ", response) + print("re rank response: ", response) - assert response.id is not None - assert response.results is not None + assert response.id is not None + assert response.results is not None - assert_response_shape(response, custom_llm_provider="together_ai") + assert_response_shape(response, custom_llm_provider="together_ai") + else: + response = await litellm.arerank( + model="together_ai/Salesforce/Llama-Rank-V1", + query="hello", + documents=["hello", "world"], + top_n=3, + ) + + print("async re rank response: ", response) + + assert response.id is not None + assert response.results is not None + + assert_response_shape(response, custom_llm_provider="together_ai") From 37ed201c5077fb9b69fef47e0ff7bd410e47c208 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 27 Aug 2024 17:09:16 -0700 Subject: [PATCH 7/8] fix install on 3.8 --- litellm/llms/cohere/rerank.py | 2 +- litellm/llms/togetherai/rerank.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/llms/cohere/rerank.py b/litellm/llms/cohere/rerank.py index 4ef523e3a14..a2a7476df87 100644 --- a/litellm/llms/cohere/rerank.py +++ b/litellm/llms/cohere/rerank.py @@ -23,7 +23,7 @@ class CohereRerank(BaseLLM): model: str, api_key: str, query: str, - documents: list[Union[str, Dict[str, Any]]], + documents: List[Union[str, Dict[str, Any]]], top_n: Optional[int] = None, rank_fields: Optional[List[str]] = None, return_documents: Optional[bool] = True, diff --git a/litellm/llms/togetherai/rerank.py b/litellm/llms/togetherai/rerank.py index 32a8cdcfdc6..5d905071c36 100644 --- a/litellm/llms/togetherai/rerank.py +++ b/litellm/llms/togetherai/rerank.py @@ -23,7 +23,7 @@ class TogetherAIRerank(BaseLLM): model: str, api_key: str, query: str, - documents: list[Union[str, Dict[str, Any]]], + documents: List[Union[str, Dict[str, Any]]], top_n: Optional[int] = None, rank_fields: Optional[List[str]] = None, return_documents: Optional[bool] = True, From fb5be57bb8d988c7de02c123d3bc2516a729d51c Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 27 Aug 2024 17:28:39 -0700 Subject: [PATCH 8/8] v0 add rerank on litellm proxy --- .../docs/proxy/guardrails/custom_guardrail.md | 1 + litellm/integrations/custom_logger.py | 1 + litellm/proxy/custom_callbacks1.py | 1 + litellm/proxy/custom_guardrail.py | 1 + .../example_config_yaml/custom_guardrail.py | 1 + .../guardrail_hooks/custom_guardrail.py | 1 + .../guardrails/guardrail_hooks/lakera_ai.py | 2 + litellm/proxy/hooks/dynamic_rate_limiter.py | 1 + litellm/proxy/proxy_server.py | 2 + litellm/proxy/rerank_endpoints/endpoints.py | 124 ++++++++++++++++++ litellm/proxy/route_llm_request.py | 2 + litellm/proxy/utils.py | 1 + 12 files changed, 138 insertions(+) create mode 100644 litellm/proxy/rerank_endpoints/endpoints.py diff --git a/docs/my-website/docs/proxy/guardrails/custom_guardrail.md b/docs/my-website/docs/proxy/guardrails/custom_guardrail.md index 469196d864e..ce4d4441499 100644 --- a/docs/my-website/docs/proxy/guardrails/custom_guardrail.md +++ b/docs/my-website/docs/proxy/guardrails/custom_guardrail.md @@ -55,6 +55,7 @@ class myCustomGuardrail(CustomGuardrail): "moderation", "audio_transcription", "pass_through_endpoint", + "rerank" ], ) -> Optional[Union[Exception, str, dict]]: """ diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 47d28ab56a7..01fd35990db 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -109,6 +109,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac "moderation", "audio_transcription", "pass_through_endpoint", + "rerank", ], ) -> Optional[ Union[Exception, str, dict] diff --git a/litellm/proxy/custom_callbacks1.py b/litellm/proxy/custom_callbacks1.py index 05028f033c0..fbfbf606099 100644 --- a/litellm/proxy/custom_callbacks1.py +++ b/litellm/proxy/custom_callbacks1.py @@ -29,6 +29,7 @@ class MyCustomHandler( "moderation", "audio_transcription", "pass_through_endpoint", + "rerank", ], ): return data diff --git a/litellm/proxy/custom_guardrail.py b/litellm/proxy/custom_guardrail.py index 2ed989cfd38..d8d63ab0a2d 100644 --- a/litellm/proxy/custom_guardrail.py +++ b/litellm/proxy/custom_guardrail.py @@ -32,6 +32,7 @@ class myCustomGuardrail(CustomGuardrail): "moderation", "audio_transcription", "pass_through_endpoint", + "rerank", ], ) -> Optional[Union[Exception, str, dict]]: """ diff --git a/litellm/proxy/example_config_yaml/custom_guardrail.py b/litellm/proxy/example_config_yaml/custom_guardrail.py index 2ed989cfd38..d8d63ab0a2d 100644 --- a/litellm/proxy/example_config_yaml/custom_guardrail.py +++ b/litellm/proxy/example_config_yaml/custom_guardrail.py @@ -32,6 +32,7 @@ class myCustomGuardrail(CustomGuardrail): "moderation", "audio_transcription", "pass_through_endpoint", + "rerank", ], ) -> Optional[Union[Exception, str, dict]]: """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py index 2ed989cfd38..d8d63ab0a2d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py @@ -32,6 +32,7 @@ class myCustomGuardrail(CustomGuardrail): "moderation", "audio_transcription", "pass_through_endpoint", + "rerank", ], ) -> Optional[Union[Exception, str, dict]]: """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py index e1ff55c82cd..364bcb22274 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py @@ -127,6 +127,7 @@ class lakeraAI_Moderation(CustomGuardrail): "moderation", "audio_transcription", "pass_through_endpoint", + "rerank", ], ): if ( @@ -288,6 +289,7 @@ class lakeraAI_Moderation(CustomGuardrail): "moderation", "audio_transcription", "pass_through_endpoint", + "rerank", ], ) -> Optional[Union[Exception, str, Dict]]: from litellm.types.guardrails import GuardrailEventHooks diff --git a/litellm/proxy/hooks/dynamic_rate_limiter.py b/litellm/proxy/hooks/dynamic_rate_limiter.py index 57985e9a690..1ef674b7e18 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter.py @@ -199,6 +199,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): "moderation", "audio_transcription", "pass_through_endpoint", + "rerank", ], ) -> Optional[ Union[Exception, str, dict] diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7e6f3c5e2fc..9d8d3fa2dfc 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -205,6 +205,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( router as pass_through_router, ) +from litellm.proxy.rerank_endpoints.endpoints import router as rerank_router from litellm.proxy.route_llm_request import route_request from litellm.proxy.secret_managers.aws_secret_manager import ( load_aws_kms, @@ -9881,6 +9882,7 @@ def cleanup_router_config_variables(): app.include_router(router) +app.include_router(rerank_router) app.include_router(fine_tuning_router) app.include_router(vertex_router) app.include_router(gemini_router) diff --git a/litellm/proxy/rerank_endpoints/endpoints.py b/litellm/proxy/rerank_endpoints/endpoints.py new file mode 100644 index 00000000000..6bc6dc94825 --- /dev/null +++ b/litellm/proxy/rerank_endpoints/endpoints.py @@ -0,0 +1,124 @@ +#### Rerank Endpoints ##### +from datetime import datetime, timedelta, timezone +from typing import List, Optional + +import fastapi +import orjson +from fastapi import APIRouter, Depends, Header, HTTPException, Request, Response, status +from fastapi.responses import ORJSONResponse + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import * +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + +router = APIRouter() +import asyncio + + +@router.post( + "/v1/rerank", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, + tags=["rerank"], +) +@router.post( + "/rerank", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, + tags=["rerank"], +) +async def rerank( + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + from litellm.proxy.proxy_server import ( + add_litellm_data_to_request, + general_settings, + get_custom_headers, + llm_router, + proxy_config, + proxy_logging_obj, + route_request, + user_model, + version, + ) + + data = {} + try: + body = await request.body() + data = orjson.loads(body) + + # Include original request and headers in the data + data = await add_litellm_data_to_request( + data=data, + request=request, + general_settings=general_settings, + user_api_key_dict=user_api_key_dict, + version=version, + proxy_config=proxy_config, + ) + + ### CALL HOOKS ### - modify incoming data / reject request before calling the model + data = await proxy_logging_obj.pre_call_hook( + user_api_key_dict=user_api_key_dict, data=data, call_type="rerank" + ) + + ## ROUTE TO CORRECT ENDPOINT ## + llm_call = await route_request( + data=data, + route_type="arerank", + llm_router=llm_router, + user_model=user_model, + ) + response = await llm_call + + ### ALERTING ### + asyncio.create_task( + proxy_logging_obj.update_request_status( + litellm_call_id=data.get("litellm_call_id", ""), status="success" + ) + ) + + ### RESPONSE HEADERS ### + hidden_params = getattr(response, "_hidden_params", {}) or {} + model_id = hidden_params.get("model_id", None) or "" + cache_key = hidden_params.get("cache_key", None) or "" + api_base = hidden_params.get("api_base", None) or "" + + fastapi_response.headers.update( + get_custom_headers( + user_api_key_dict=user_api_key_dict, + model_id=model_id, + cache_key=cache_key, + api_base=api_base, + version=version, + model_region=getattr(user_api_key_dict, "allowed_model_region", ""), + request_data=data, + ) + ) + + return response + except Exception as e: + await proxy_logging_obj.post_call_failure_hook( + user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data + ) + verbose_proxy_logger.error( + "litellm.proxy.proxy_server.rerank(): Exception occured - {}".format(str(e)) + ) + if isinstance(e, HTTPException): + raise ProxyException( + message=getattr(e, "message", str(e)), + type=getattr(e, "type", "None"), + param=getattr(e, "param", "None"), + code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), + ) + else: + error_msg = f"{str(e)}" + raise ProxyException( + message=getattr(e, "message", error_msg), + type=getattr(e, "type", "None"), + param=getattr(e, "param", "None"), + code=getattr(e, "status_code", 500), + ) diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 7a7be55b22a..361c5be0c0b 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -33,6 +33,7 @@ ROUTE_ENDPOINT_MAPPING = { "aspeech": "/audio/speech", "atranscription": "/audio/transcriptions", "amoderation": "/moderations", + "arerank": "/rerank", } @@ -48,6 +49,7 @@ async def route_request( "aspeech", "atranscription", "amoderation", + "arerank", ], ): """ diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 09fc014d58b..ffd354224c7 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -375,6 +375,7 @@ class ProxyLogging: "moderation", "audio_transcription", "pass_through_endpoint", + "rerank", ], ) -> dict: """