services(google): make api_key argument mandatory

This commit is contained in:
Aleix Conchillo Flaqué
2024-05-24 09:05:17 -07:00
parent 4e594aa9b0
commit c45d428551
7 changed files with 25 additions and 24 deletions

View File

@@ -7,6 +7,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [Unreleased] ## [Unreleased]
### Changed
- GoogleLLMService `api_key` argument is now mandatory.
### Fixed ### Fixed
- Fixed AzureLLMService. - Fixed AzureLLMService.

View File

@@ -62,19 +62,15 @@ async def main(room_url: str, token):
) )
) )
tts = ElevenLabsTTSService(
aiohttp_session=session,
api_key=os.getenv("ELEVENLABS_API_KEY"),
voice_id=os.getenv("ELEVENLABS_VOICE_ID"),
)
user_response = UserResponseAggregator() user_response = UserResponseAggregator()
image_requester = UserImageRequester() image_requester = UserImageRequester()
vision_aggregator = VisionImageFrameAggregator() vision_aggregator = VisionImageFrameAggregator()
google = GoogleLLMService(model="gemini-1.5-flash-latest") google = GoogleLLMService(
model="gemini-1.5-flash-latest",
api_key=os.getenv("GOOGLE_API_KEY"))
tts = ElevenLabsTTSService( tts = ElevenLabsTTSService(
aiohttp_session=session, aiohttp_session=session,

View File

@@ -44,9 +44,9 @@ class AnthropicLLMService(LLMService):
def __init__( def __init__(
self, self,
api_key, api_key: str,
model="claude-3-opus-20240229", model: str = "claude-3-opus-20240229",
max_tokens=1024): max_tokens: int = 1024):
super().__init__() super().__init__()
self._client = AsyncAnthropic(api_key=api_key) self._client = AsyncAnthropic(api_key=api_key)
self._model = model self._model = model

View File

@@ -11,6 +11,7 @@ import io
from PIL import Image from PIL import Image
from typing import AsyncGenerator from typing import AsyncGenerator
from numpy import str_
from openai import AsyncAzureOpenAI from openai import AsyncAzureOpenAI
from pipecat.frames.frames import AudioRawFrame, ErrorFrame, Frame, URLImageRawFrame from pipecat.frames.frames import AudioRawFrame, ErrorFrame, Frame, URLImageRawFrame
@@ -73,10 +74,10 @@ class AzureLLMService(BaseOpenAILLMService):
def __init__( def __init__(
self, self,
*, *,
api_key, api_key: str,
endpoint, endpoint: str,
api_version="2023-12-01-preview", model: str,
model): api_version: str = "2023-12-01-preview"):
# Initialize variables before calling parent __init__() because that # Initialize variables before calling parent __init__() because that
# will call create_client() and we need those values there. # will call create_client() and we need those values there.
self._endpoint = endpoint self._endpoint = endpoint
@@ -96,12 +97,12 @@ class AzureImageGenServiceREST(ImageGenService):
def __init__( def __init__(
self, self,
*, *,
api_version="2023-06-01-preview",
image_size: str,
aiohttp_session: aiohttp.ClientSession, aiohttp_session: aiohttp.ClientSession,
api_key, image_size: str,
endpoint, api_key: str,
model, endpoint: str,
model: str,
api_version="2023-06-01-preview",
): ):
super().__init__() super().__init__()

View File

@@ -19,6 +19,6 @@ except ModuleNotFoundError as e:
class FireworksLLMService(BaseOpenAILLMService): class FireworksLLMService(BaseOpenAILLMService):
def __init__(self, def __init__(self,
model="accounts/fireworks/models/firefunction-v1", model: str = "accounts/fireworks/models/firefunction-v1",
base_url="https://api.fireworks.ai/inference/v1"): base_url: str = "https://api.fireworks.ai/inference/v1"):
super().__init__(model, base_url) super().__init__(model, base_url)

View File

@@ -40,9 +40,9 @@ class GoogleLLMService(LLMService):
franca for all LLM services, so that it is easy to switch between different LLMs. franca for all LLM services, so that it is easy to switch between different LLMs.
""" """
def __init__(self, model="gemini-1.5-flash-latest", api_key=None, **kwargs): def __init__(self, api_key: str, model: str = "gemini-1.5-flash-latest", **kwargs):
super().__init__(**kwargs) super().__init__(**kwargs)
gai.configure(api_key=api_key or os.environ["GOOGLE_API_KEY"]) gai.configure(api_key=api_key)
self._client = gai.GenerativeModel(model) self._client = gai.GenerativeModel(model)
def _get_messages_from_openai_context( def _get_messages_from_openai_context(

View File

@@ -9,5 +9,5 @@ from pipecat.services.openai import BaseOpenAILLMService
class OLLamaLLMService(BaseOpenAILLMService): class OLLamaLLMService(BaseOpenAILLMService):
def __init__(self, model="llama2", base_url="http://localhost:11434/v1"): def __init__(self, model: str = "llama2", base_url: str = "http://localhost:11434/v1"):
super().__init__(model=model, base_url=base_url, api_key="ollama") super().__init__(model=model, base_url=base_url, api_key="ollama")