Add provider option for API key

This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-02-21 21:38:55 -03:00
commit 73b23a2501

View file

@ -1,4 +1,3 @@
import os
from typing import Any, Callable, Dict, Optional, Union from typing import Any, Callable, Dict, Optional, Union
from langchain_community.chat_models.litellm import ChatLiteLLM, ChatLiteLLMException from langchain_community.chat_models.litellm import ChatLiteLLM, ChatLiteLLMException
@ -27,6 +26,18 @@ class ChatLiteLLMComponent(CustomComponent):
"required": False, "required": False,
"password": True, "password": True,
}, },
"provider": {
"display_name": "Provider",
"info": "The provider of the API key.",
"options": [
"OpenAI",
"Azure",
"Anthropic",
"Replicate",
"Cohere",
"OpenRouter",
],
},
"streaming": { "streaming": {
"display_name": "Streaming", "display_name": "Streaming",
"field_type": "bool", "field_type": "bool",
@ -96,7 +107,8 @@ class ChatLiteLLMComponent(CustomComponent):
def build( def build(
self, self,
model: str, model: str,
api_key: str, provider: str,
api_key: Optional[str] = None,
streaming: bool = True, streaming: bool = True,
temperature: Optional[float] = 0.7, temperature: Optional[float] = 0.7,
model_kwargs: Optional[Dict[str, Any]] = {}, model_kwargs: Optional[Dict[str, Any]] = {},
@ -114,13 +126,19 @@ class ChatLiteLLMComponent(CustomComponent):
litellm.set_verbose = verbose litellm.set_verbose = verbose
except ImportError: except ImportError:
raise ChatLiteLLMException( raise ChatLiteLLMException(
"Could not import litellm python package. " "Please install it with `pip install litellm`" "Could not import litellm python package. "
"Please install it with `pip install litellm`"
) )
if api_key: provider_map = {
if "perplexity" in model: "OpenAI": "openai_api_key",
os.environ["PERPLEXITYAI_API_KEY"] = api_key "Azure": "azure_api_key",
elif "replicate" in model: "Anthropic": "anthropic_api_key",
os.environ["REPLICATE_API_KEY"] = api_key "Replicate": "replicate_api_key",
"Cohere": "cohere_api_key",
"OpenRouter": "openrouter_api_key",
}
# Set the API key based on the provider
kwarg = {provider_map[provider]: api_key}
LLM = ChatLiteLLM( LLM = ChatLiteLLM(
model=model, model=model,
@ -133,5 +151,6 @@ class ChatLiteLLMComponent(CustomComponent):
n=n, n=n,
max_tokens=max_tokens, max_tokens=max_tokens,
max_retries=max_retries, max_retries=max_retries,
**kwarg,
) )
return LLM return LLM