Move location and project_id out of InputParams in GoogleVertexLLMService, making them top-level init args instead. We do this for two reasons:
- Conceptually, these args comprise project-level setup, akin to credentials, whereas everything in `InputParams` is concerned with model configuration - Providing a `project_id` when initializing `GoogleVertexLLMService` should not be optional, but prior to the change in this commit it was (erroneously) treated as optional by dint of `InputParams` being optional This improvement was discussed [in this comment](https://github.com/pipecat-ai/pipecat/pull/2795#discussion_r2408279142).
This commit is contained in:
@@ -51,6 +51,10 @@ class GoogleVertexLLMService(OpenAILLMService):
|
|||||||
class InputParams(OpenAILLMService.InputParams):
|
class InputParams(OpenAILLMService.InputParams):
|
||||||
"""Input parameters specific to Vertex AI.
|
"""Input parameters specific to Vertex AI.
|
||||||
|
|
||||||
|
.. deprecated:: 0.0.90
|
||||||
|
Use `OpenAILLMService.InputParams` instead and provide `location` and
|
||||||
|
`project_id` as direct arguments to `GoogleVertexLLMService.__init__()`.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
location: GCP region for Vertex AI endpoint (e.g., "us-east4").
|
location: GCP region for Vertex AI endpoint (e.g., "us-east4").
|
||||||
project_id: Google Cloud project ID.
|
project_id: Google Cloud project ID.
|
||||||
@@ -60,12 +64,29 @@ class GoogleVertexLLMService(OpenAILLMService):
|
|||||||
location: str = "us-east4"
|
location: str = "us-east4"
|
||||||
project_id: str
|
project_id: str
|
||||||
|
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
"""Initializes the InputParams."""
|
||||||
|
import warnings
|
||||||
|
|
||||||
|
with warnings.catch_warnings():
|
||||||
|
warnings.simplefilter("ignore", DeprecationWarning)
|
||||||
|
warnings.warn(
|
||||||
|
"GoogleVertexLLMService.InputParams is deprecated. "
|
||||||
|
"Please use OpenAILLMService.InputParams instead and provide "
|
||||||
|
"'location' and 'project_id' as direct arguments to __init__.",
|
||||||
|
DeprecationWarning,
|
||||||
|
stacklevel=2,
|
||||||
|
)
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
credentials: Optional[str] = None,
|
credentials: Optional[str] = None,
|
||||||
credentials_path: Optional[str] = None,
|
credentials_path: Optional[str] = None,
|
||||||
model: str = "google/gemini-2.0-flash-001",
|
model: str = "google/gemini-2.0-flash-001",
|
||||||
|
location: Optional[str] = None,
|
||||||
|
project_id: Optional[str] = None,
|
||||||
params: Optional[InputParams] = None,
|
params: Optional[InputParams] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
@@ -75,11 +96,42 @@ class GoogleVertexLLMService(OpenAILLMService):
|
|||||||
credentials: JSON string of service account credentials.
|
credentials: JSON string of service account credentials.
|
||||||
credentials_path: Path to the service account JSON file.
|
credentials_path: Path to the service account JSON file.
|
||||||
model: Model identifier (e.g., "google/gemini-2.0-flash-001").
|
model: Model identifier (e.g., "google/gemini-2.0-flash-001").
|
||||||
|
location: GCP region for Vertex AI endpoint (e.g., "us-east4").
|
||||||
|
project_id: Google Cloud project ID.
|
||||||
params: Vertex AI input parameters including location and project.
|
params: Vertex AI input parameters including location and project.
|
||||||
|
|
||||||
|
.. deprecated:: 0.0.90
|
||||||
|
Use `OpenAILLMService.InputParams` instead and provide `location`
|
||||||
|
and `project_id` as direct arguments to `__init__()`.
|
||||||
|
|
||||||
**kwargs: Additional arguments passed to OpenAILLMService.
|
**kwargs: Additional arguments passed to OpenAILLMService.
|
||||||
"""
|
"""
|
||||||
params = params or OpenAILLMService.InputParams()
|
# Handle deprecated InputParams
|
||||||
base_url = self._get_base_url(params)
|
if params is not None and isinstance(params, GoogleVertexLLMService.InputParams):
|
||||||
|
# Extract location and project_id from params if not provided directly
|
||||||
|
if project_id is None:
|
||||||
|
project_id = params.project_id
|
||||||
|
if location is None:
|
||||||
|
location = params.location
|
||||||
|
# Convert to base InputParams
|
||||||
|
params = OpenAILLMService.InputParams(
|
||||||
|
**params.model_dump(exclude={"location", "project_id"}, exclude_unset=True)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
params = params or OpenAILLMService.InputParams()
|
||||||
|
|
||||||
|
# Validate parameters
|
||||||
|
# NOTE: once we remove deprecated InputParams, we can update __init__()
|
||||||
|
# signature with the following:
|
||||||
|
# - location: str = "us-east4",
|
||||||
|
# - project_id: str,
|
||||||
|
# For now, we need them as-is to maintain backward compatibility.
|
||||||
|
if project_id is None:
|
||||||
|
raise ValueError("project_id is required")
|
||||||
|
if location is None:
|
||||||
|
location = "us-east4" # Default location if not provided
|
||||||
|
|
||||||
|
base_url = self._get_base_url(location, project_id)
|
||||||
self._api_key = self._get_api_token(credentials, credentials_path)
|
self._api_key = self._get_api_token(credentials, credentials_path)
|
||||||
|
|
||||||
super().__init__(
|
super().__init__(
|
||||||
@@ -91,17 +143,14 @@ class GoogleVertexLLMService(OpenAILLMService):
|
|||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _get_base_url(params: InputParams) -> str:
|
def _get_base_url(location: str, project_id: str) -> str:
|
||||||
"""Construct the base URL for Vertex AI API."""
|
"""Construct the base URL for Vertex AI API."""
|
||||||
# Determine the correct API host based on location
|
# Determine the correct API host based on location
|
||||||
if params.location == "global":
|
if location == "global":
|
||||||
api_host = "aiplatform.googleapis.com"
|
api_host = "aiplatform.googleapis.com"
|
||||||
else:
|
else:
|
||||||
api_host = f"{params.location}-aiplatform.googleapis.com"
|
api_host = f"{location}-aiplatform.googleapis.com"
|
||||||
return (
|
return f"https://{api_host}/v1/projects/{project_id}/locations/{location}/endpoints/openapi"
|
||||||
f"https://{api_host}/v1/"
|
|
||||||
f"projects/{params.project_id}/locations/{params.location}/endpoints/openapi"
|
|
||||||
)
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _get_api_token(credentials: Optional[str], credentials_path: Optional[str]) -> str:
|
def _get_api_token(credentials: Optional[str], credentials_path: Optional[str]) -> str:
|
||||||
|
|||||||
Reference in New Issue
Block a user