diff --git a/backend/app/core/cloud/storage.py b/backend/app/core/cloud/storage.py index fcb7eb616..f1000f6f6 100644 --- a/backend/app/core/cloud/storage.py +++ b/backend/app/core/cloud/storage.py @@ -307,7 +307,7 @@ def get_cloud_storage(session: Session, project_id: int) -> CloudStorage: Method to create and configure a cloud storage instance. """ # Lazy import to avoid a top-level cycle: storage.py is imported from - # app.services.llm.providers.google_ai, which itself is wired into the + # app.services.llm.providers.google_gcp, which itself is wired into the # provider registry that app.crud transitively pulls in. from app.crud import get_project_by_id diff --git a/backend/app/core/providers.py b/backend/app/core/providers.py index 1e0c58f18..e3fe27f4d 100644 --- a/backend/app/core/providers.py +++ b/backend/app/core/providers.py @@ -14,6 +14,7 @@ class Provider(str, Enum): OPENAI = "openai" LANGFUSE = "langfuse" GOOGLE_AISTUDIO = "google-aistudio" + GOOGLE_GCP = "google-gcp" SARVAMAI = "sarvamai" ELEVENLABS = "elevenlabs" ANTHROPIC = "anthropic" @@ -68,6 +69,14 @@ class GoogleCredentials(ProviderCredentialsBase): api_key: str = Field(description="Google API key") +class GoogleGcpCredentials(ProviderCredentialsBase): + api_key: str = Field(description="Google GCP API key") + project_id: str = Field(description="GCP project ID") + location: str = Field(description="GCP region/location") + sa_key: JsonValue = Field(description="Service account key JSON") + gcs_bucket: str = Field(description="GCS bucket name") + + class WebhookSecretCredentials(ProviderCredentialsBase): webhook_secret: str = Field( description="Shared secret used to HMAC-sign outgoing webhooks" @@ -86,6 +95,7 @@ class ProxyCredentials(ProviderCredentialsBase): | ElevenLabsCredentials | AnthropicCredentials | GoogleCredentials + | GoogleGcpCredentials | WebhookSecretCredentials | ProxyCredentials, Field( @@ -141,6 +151,10 @@ def required_fields(self) -> list[str]: model=GoogleCredentials, sensitive_fields=["api_key"], ), + Provider.GOOGLE_GCP: ProviderConfig( + model=GoogleGcpCredentials, + sensitive_fields=["api_key", "sa_key"], + ), Provider.WEBHOOK_SECRET: ProviderConfig( model=WebhookSecretCredentials, sensitive_fields=["webhook_secret"] ), diff --git a/backend/app/crud/credentials.py b/backend/app/crud/credentials.py index 7736f4600..6916546dd 100644 --- a/backend/app/crud/credentials.py +++ b/backend/app/crud/credentials.py @@ -6,7 +6,7 @@ from sqlmodel import Session, select from app.core.exception_handlers import HTTPException -from app.core.providers import validate_provider +from app.core.providers import parse_provider_credentials, validate_provider from app.core.security import decrypt_credentials, encrypt_credentials from app.core.util import now from app.models import Credential, CredsCreate, CredsUpdate @@ -194,6 +194,27 @@ def update_creds_for_org( Credential.project_id == project_id, ) creds = session.exec(statement).one_or_none() + + # Merge onto the existing credentials so a partial payload (PATCH) only + # overwrites the fields it supplies instead of dropping the rest. + merged_credential_data = credential_data + if creds and creds.credential: + merged_credential_data = { + **decrypt_credentials(creds.credential), + **credential_data, + } + + try: + parse_provider_credentials(provider, merged_credential_data) + except ValueError as e: + logger.warning( + f"[update_creds_for_org] Validation error | organization_id: {org_id}, project_id: {project_id}, provider: {provider}, error: {str(e)}" + ) + raise HTTPException(status_code=400, detail=str(e)) + + # Encrypt the entire credentials object + encrypted_credentials = encrypt_credentials(merged_credential_data) + if creds is None: # Create new credential if it doesn't exist creds = Credential( diff --git a/backend/app/models/credentials.py b/backend/app/models/credentials.py index 354ebdcfe..533c2d64d 100644 --- a/backend/app/models/credentials.py +++ b/backend/app/models/credentials.py @@ -84,8 +84,12 @@ class CredsUpdate(SQLModel): provider: Provider = Field( description="Name of the provider to update/add credentials for" ) - credential: ProviderCredentials = Field( - description="Credentials for the specified provider", + credential: CredentialPayload = Field( + description=( + "Credentials for the specified provider. May be a partial payload " + "(PATCH semantics) — completeness is validated after merging with " + "any existing stored credentials, not on this raw payload." + ), ) is_active: bool | None = Field( default=None, description="Whether the credentials are active" @@ -107,15 +111,21 @@ def _parse_credential(cls, data: object) -> object: if isinstance(nested, dict): credential = nested + # An empty payload has nothing to merge with an existing stored + # credential, so it is rejected here rather than deferred to the + # crud-level merge check. + if isinstance(credential, dict) and not credential: + parse_provider_credentials(provider_key, credential) + return { **data, "provider": provider_key, - "credential": parse_provider_credentials(provider_key, credential), + "credential": credential, } def credential_payload(self) -> CredentialPayload: """Credential dict for `provider`, exactly as submitted.""" - return self.credential.model_dump(exclude_unset=True) + return self.credential class Credential(CredsBase, table=True): diff --git a/backend/app/models/llm/constants.py b/backend/app/models/llm/constants.py index 77e5ae3c0..f6e927bfc 100644 --- a/backend/app/models/llm/constants.py +++ b/backend/app/models/llm/constants.py @@ -9,6 +9,7 @@ class Provider(StrEnum): ELEVENLABS = "elevenlabs" ANTHROPIC = "anthropic" GOOGLE_AISTUDIO = "google-aistudio" + GOOGLE_GCP = "google-gcp" PROXY = "proxy" @@ -17,12 +18,14 @@ class Provider(StrEnum): # instead of leaving behind stale magic strings. STTProvider = Literal[ Provider.GOOGLE, + Provider.GOOGLE_GCP, Provider.SARVAMAI, Provider.ELEVENLABS, Provider.GOOGLE_AISTUDIO, ] TTSProvider = Literal[ Provider.GOOGLE, + Provider.GOOGLE_GCP, Provider.SARVAMAI, Provider.ELEVENLABS, Provider.GOOGLE_AISTUDIO, @@ -30,7 +33,11 @@ class Provider(StrEnum): RAGProvider = Literal[Provider.OPENAI, Provider.GOOGLE_AISTUDIO] TextProvider = Literal[ - Provider.OPENAI, Provider.GOOGLE, Provider.ANTHROPIC, Provider.GOOGLE_AISTUDIO + Provider.OPENAI, + Provider.GOOGLE, + Provider.ANTHROPIC, + Provider.GOOGLE_AISTUDIO, + Provider.GOOGLE_GCP, ] KaapiProvider = Union[TextProvider, STTProvider, TTSProvider] @@ -46,6 +53,7 @@ class Provider(StrEnum): "elevenlabs-native", "anthropic-native", "google-aistudio-native", + "google-gcp-native", ] diff --git a/backend/app/models/llm/request.py b/backend/app/models/llm/request.py index cef0c017b..c028a978a 100644 --- a/backend/app/models/llm/request.py +++ b/backend/app/models/llm/request.py @@ -342,8 +342,9 @@ class KaapiTextCompletionConfig(SQLModel): provider: TextProvider | None = Field( default=None, description=( - "LLM provider for text completions (openai, google, anthropic). " - "Omit to use the platform default for the type." + "LLM provider for text completions (openai, google, anthropic, " + "google-aistudio, google-gcp). Omit to use the platform default " + "for the type." ), ) type: Literal[CompletionType.TEXT] = Field( diff --git a/backend/app/models/model_config.py b/backend/app/models/model_config.py index a55341b9d..c9c7cf7f4 100644 --- a/backend/app/models/model_config.py +++ b/backend/app/models/model_config.py @@ -18,6 +18,7 @@ class ModelConfigBase(SQLModel): "anthropic", "proxy", "google-aistudio", + "google-gcp", ] = Field( default="openai", sa_column=sa.Column( @@ -29,12 +30,13 @@ class ModelConfigBase(SQLModel): "anthropic", "proxy", "google-aistudio", + "google-gcp", name="provider_enum", schema="global", create_type=False, ), nullable=False, - comment="provider name (e.g. openai, google, sarvamai, elevenlabs, anthropic, google-aistudio, proxy)", + comment="provider name (e.g. openai, google, sarvamai, elevenlabs, anthropic, google-aistudio, google-gcp, proxy)", ), ) diff --git a/backend/app/services/llm/mappers.py b/backend/app/services/llm/mappers.py index 3082bcf1e..f2e3d7d05 100644 --- a/backend/app/services/llm/mappers.py +++ b/backend/app/services/llm/mappers.py @@ -740,6 +740,20 @@ def transform_kaapi_config_to_native( warnings, ) + if kaapi_config.provider == Provider.GOOGLE_GCP: + # Kaapi STT/TTS param shape is identical to Google's; reuse the Google mapper. + mapped_params, warnings = map_kaapi_to_google_params( + kaapi_config.params, kaapi_config.type + ) + return ( + NativeCompletionConfig( + provider="google-gcp-native", + params=mapped_params, + type=kaapi_config.type, + ), + warnings, + ) + if kaapi_config.provider == Provider.ANTHROPIC: if kaapi_config.type != CompletionType.TEXT: raise ValueError( diff --git a/backend/app/services/llm/providers/__init__.py b/backend/app/services/llm/providers/__init__.py index b044a1c42..024cd3352 100644 --- a/backend/app/services/llm/providers/__init__.py +++ b/backend/app/services/llm/providers/__init__.py @@ -4,7 +4,7 @@ from app.services.llm.providers.eleven_ai import ElevenlabsAIProvider from app.services.llm.providers.sarvam_ai import SarvamAIProvider from app.services.llm.providers.claude import ClaudeProvider -from app.services.llm.providers.google_ai import GoogleVertexAIProvider +from app.services.llm.providers.google_gcp import GoogleGCPProvider from app.services.llm.providers.registry import ( LLMProvider, get_llm_provider, diff --git a/backend/app/services/llm/providers/google_ai.py b/backend/app/services/llm/providers/google_gcp.py similarity index 70% rename from backend/app/services/llm/providers/google_ai.py rename to backend/app/services/llm/providers/google_gcp.py index bf8bb8ed2..d486c1ab2 100644 --- a/backend/app/services/llm/providers/google_ai.py +++ b/backend/app/services/llm/providers/google_gcp.py @@ -1,9 +1,7 @@ import base64 import json import logging -import os import uuid -from pathlib import Path from typing import Any import requests @@ -17,9 +15,11 @@ from app.core.cloud.storage import upload_audio_to_gcs from app.core.config import settings from app.models.llm import ( + ImageContent, LLMCallResponse, LLMResponse, NativeCompletionConfig, + PDFContent, QueryParams, TextContent, TextOutput, @@ -27,6 +27,7 @@ ) from app.models.llm.constants import ( DEFAULT_STT_MODEL, + DEFAULT_TEXT_MODELS, DEFAULT_TTS_MODEL, DEFAULT_TTS_VOICE, CompletionType, @@ -71,8 +72,8 @@ def _load_platform_sa_info() -> dict | None: return None -class VertexClient: - """Holds Vertex AI connection details. Pure config — no SDK session. +class GoogleGCPClient: + """Holds Google GCP connection details. Pure config — no SDK session. BYOK: per-project SA JSON + GCS bucket are passed via credentials and stored directly on the client; falls back to platform-shared values @@ -105,15 +106,14 @@ def endpoint(self, model: str) -> str: ) -class GoogleVertexAIProvider(BaseProvider): - """Google Vertex AI provider using REST + API key auth. +class GoogleGCPProvider(BaseProvider): + """Google GCP provider using REST + API key auth. - Supports STT (audio → text) and TTS (text → audio) via Gemini multimodal - models on Vertex. Text-only completions are routed through the standard - `google` provider. + Supports text, STT (audio → text), and TTS (text → audio) via Gemini + multimodal models on GCP. """ - def __init__(self, client: VertexClient): + def __init__(self, client: GoogleGCPClient): super().__init__(client) self.client = client @@ -129,7 +129,7 @@ def create_client(credentials: dict[str, Any]) -> Any: source = "byok" if credentials.get("api_key") else "platform" logger.info( - f"[create_client] vertex creds | source={source}, " + f"[create_client] google-gcp creds | source={source}, " f"project_id={project_id}, location={location}" ) @@ -139,14 +139,16 @@ def create_client(credentials: dict[str, Any]) -> Any: ("api_key", api_key), ("project_id", project_id), ("location", location), + ("sa_key", sa_info), + ("gcs_bucket", gcs_bucket), ) if not value ] if missing: raise ValueError( - f"Google Vertex AI credentials missing required fields: {', '.join(missing)}" + f"Google GCP credentials missing required fields: {', '.join(missing)}" ) - return VertexClient( + return GoogleGCPClient( api_key=api_key, project_id=project_id, location=location, @@ -157,19 +159,19 @@ def create_client(credentials: dict[str, Any]) -> Any: def _post( self, model: str, payload: dict, log_context: str = "" ) -> tuple[dict | None, str | None]: - """POST to Vertex generateContent and return parsed JSON or a + """POST to GCP generateContent and return parsed JSON or a descriptive, pre-logged error message. Maps: - ``requests.Timeout`` / ``ConnectionError`` / ``RequestException`` → ``[KAAPI]`` network-side errors - - HTTP 4xx/5xx → ``[VERTEX]`` errors, branched by status code, with + - HTTP 4xx/5xx → ``[GOOGLE_GCP]`` errors, branched by status code, with Google's ``error.message`` / ``error.status`` surfaced when the response body is the standard error envelope. - - Non-JSON 200 body → ``[VERTEX]`` malformed-response error + - Non-JSON 200 body → ``[GOOGLE_GCP]`` malformed-response error """ url = self.client.endpoint(model) - logger.debug(f"[_post] vertex url={url}") + logger.debug(f"[_post] google-gcp url={url}") try: resp = requests.post( @@ -181,36 +183,36 @@ def _post( ) except requests.Timeout as e: error_message = ( - f"[KAAPI] Vertex AI request timed out after {REQUEST_TIMEOUT}s " + f"[KAAPI] Google GCP request timed out after {REQUEST_TIMEOUT}s " f"(code: {type(e).__name__}): {str(e)}. The request took too " f"long to complete — retry with a smaller payload or contact " f"Kaapi if the issue persists." ) logger.error( - f"[GoogleVertexAIProvider._post] {error_message} | model={model}, {log_context}", + f"[GoogleGCPProvider._post] {error_message} | model={model}, {log_context}", exc_info=True, ) return None, error_message except requests.ConnectionError as e: error_message = ( - f"[KAAPI] Vertex AI connection failed (code: " + f"[KAAPI] Google GCP connection failed (code: " f"{type(e).__name__}): {str(e)}. Network or DNS issue " - f"reaching Vertex — check network connectivity from the " + f"reaching Google GCP — check network connectivity from the " f"Kaapi backend. If the issue persists, contact Kaapi." ) logger.error( - f"[GoogleVertexAIProvider._post] {error_message} | model={model}, {log_context}", + f"[GoogleGCPProvider._post] {error_message} | model={model}, {log_context}", exc_info=True, ) return None, error_message except requests.RequestException as e: error_message = ( - f"[KAAPI] Vertex AI request failed (code: " + f"[KAAPI] Google GCP request failed (code: " f"{type(e).__name__}): {str(e)}. Unexpected requests-library " f"error — contact Kaapi if the issue persists." ) logger.error( - f"[GoogleVertexAIProvider._post] {error_message} | model={model}, {log_context}", + f"[GoogleGCPProvider._post] {error_message} | model={model}, {log_context}", exc_info=True, ) return None, error_message @@ -232,44 +234,44 @@ def _post( if status_code == 400: error_message = ( - f"[VERTEX] Bad request (code: 400{status_label}): " + f"[GOOGLE_GCP] Bad request (code: 400{status_label}): " f"{google_msg}. Review your config parameters and input " f"payload — the request shape, model, or content may be " - f"invalid for this Vertex endpoint." + f"invalid for this Google GCP endpoint." ) elif status_code in (401, 403): error_message = ( - f"[VERTEX] Authentication / permission denied (code: " + f"[GOOGLE_GCP] Authentication / permission denied (code: " f"{status_code}{status_label}): {google_msg}. Verify the " - f"Vertex API key is valid and not expired, the project_id " + f"Google GCP API key is valid and not expired, the project_id " f"and location are correct, and the service account has " f"access to the requested model." ) elif status_code == 404: error_message = ( - f"[VERTEX] Resource not found (code: 404{status_label}): " + f"[GOOGLE_GCP] Resource not found (code: 404{status_label}): " f"{google_msg}. Check that the model '{model}' exists and " f"is available in your project and location." ) elif status_code == 429: error_message = ( - f"[VERTEX] Rate limit / quota exceeded (code: 429" - f"{status_label}): {google_msg}. You have hit Vertex AI's " + f"[GOOGLE_GCP] Rate limit / quota exceeded (code: 429" + f"{status_label}): {google_msg}. You have hit Google GCP's " f"per-minute or per-day quota for this model. Wait at " f"least 1 minute and retry; if the issue persists, " f"request a quota increase from Google or contact Kaapi." ) elif 500 <= status_code < 600: error_message = ( - f"[VERTEX] Server error (code: {status_code}" + f"[GOOGLE_GCP] Server error (code: {status_code}" f"{status_label}): {google_msg}. This is typically " - f"transient (Vertex overloaded or internal error) — " + f"transient (Google GCP overloaded or internal error) — " f"retry in a few seconds. If the issue persists, contact " f"Kaapi." ) else: error_message = ( - f"[VERTEX] HTTP error (code: {status_code}{status_label}): " + f"[GOOGLE_GCP] HTTP error (code: {status_code}{status_label}): " f"{google_msg}. If the issue persists, contact Kaapi." ) @@ -277,7 +279,7 @@ def _post( # are caller's fault and only need a warning. log = logger.error if 500 <= status_code < 600 else logger.warning log( - f"[GoogleVertexAIProvider._post] {error_message} | " + f"[GoogleGCPProvider._post] {error_message} | " f"model={model}, {log_context}" ) return None, error_message @@ -286,12 +288,12 @@ def _post( return resp.json(), None except ValueError as e: error_message = ( - f"[VERTEX] Returned a non-JSON success response: {str(e)}. " - f"This indicates an unexpected payload shape from Vertex — " + f"[GOOGLE_GCP] Returned a non-JSON success response: {str(e)}. " + f"This indicates an unexpected payload shape from Google GCP — " f"retry the request. If the issue persists, contact Kaapi." ) logger.warning( - f"[GoogleVertexAIProvider._post] {error_message} | " + f"[GoogleGCPProvider._post] {error_message} | " f"model={model}, {log_context}" ) return None, error_message @@ -310,18 +312,118 @@ def _extract_usage(data: dict) -> Usage: reasoning_tokens=reasoning_tokens, ) + @staticmethod + def _format_content_parts(parts: list[ContentPart]) -> list[dict]: + """Render Kaapi content parts as GCP REST `parts` entries (camelCase + keys, matching the generateContent wire format used elsewhere in + this file — not the genai SDK's snake_case Python bindings).""" + items = [] + for part in parts: + if isinstance(part, TextContent): + items.append({"text": part.value}) + elif isinstance(part, (ImageContent, PDFContent)): + if part.format == "base64": + items.append( + {"inlineData": {"mimeType": part.mime_type, "data": part.value}} + ) + else: + items.append( + { + "fileData": { + "mimeType": part.mime_type, + "fileUri": part.value, + } + } + ) + return items + + def _execute_text( + self, + completion_config: NativeCompletionConfig, + resolved_input: str | list[ContentPart] | MultiModalInput, + include_provider_raw_response: bool = False, + ) -> tuple[LLMCallResponse | None, str | None]: + """Execute a text completion via Google GCP generateContent.""" + provider = completion_config.provider + params = completion_config.params + + if isinstance(resolved_input, MultiModalInput): + gemini_parts = self._format_content_parts(resolved_input.parts) + elif isinstance(resolved_input, list): + gemini_parts = self._format_content_parts(resolved_input) + else: + gemini_parts = [{"text": resolved_input}] + + model = params.get("model") or DEFAULT_TEXT_MODELS["google"] + instructions = params.get("instructions") + temperature = params.get("temperature") + + payload: dict[str, Any] = { + "contents": [{"role": "user", "parts": gemini_parts}], + } + if instructions: + payload["systemInstruction"] = {"parts": [{"text": instructions}]} + + generation_config: dict[str, Any] = {} + if temperature is not None: + generation_config["temperature"] = temperature + if generation_config: + payload["generationConfig"] = generation_config + + data, err = self._post( + model, payload, log_context=f"provider={provider}, type=text" + ) + if err: + return None, err + + try: + text = data["candidates"][0]["content"]["parts"][0]["text"] + except (KeyError, IndexError, TypeError): + error_message = ( + "[GOOGLE_GCP] Text response is missing generated content. " + "Google GCP returned a 200 response but the expected " + "candidates[0].content.parts[0].text path is absent — this " + "typically means the response was blocked by safety filters " + "or truncated. Review the prompt and safety settings, then " + "retry." + ) + logger.warning( + f"[GoogleGCPProvider._execute_text] {error_message} | " + f"provider={provider}, model={model}, response_id={data.get('responseId')}" + ) + return None, error_message + + llm_response = LLMCallResponse( + response=LLMResponse( + provider_response_id=data.get("responseId") + or f"google-gcp-{uuid.uuid4().hex}", + model=data.get("modelVersion") or model, + provider=provider, + output=TextOutput(content=TextContent(value=text)), + ), + usage=self._extract_usage(data), + ) + + if include_provider_raw_response: + llm_response.provider_raw_response = data + + logger.info( + f"[GoogleGCPProvider._execute_text] Generated text response | provider={provider}, model={model}" + ) + return llm_response, None + def _execute_stt( self, completion_config: NativeCompletionConfig, resolved_input: "AudioRef", include_provider_raw_response: bool = False, ) -> tuple[LLMCallResponse | None, str | None]: - """Execute STT via Vertex generateContent. + """Execute STT via Google GCP generateContent. Note: HTTP / network errors come back from ``_post()`` already-logged and tagged. This method only handles Kaapi-side input validation, - staging failures, and Vertex response-shape checks. + staging failures, and Google GCP response-shape checks. """ provider = completion_config.provider params = completion_config.params @@ -333,7 +435,7 @@ def _execute_stt( f"Ensure the audio is uploaded and resolved before invoking STT." ) logger.warning( - f"[GoogleVertexAIProvider._execute_stt] {error_message} | provider={provider}" + f"[GoogleGCPProvider._execute_stt] {error_message} | provider={provider}" ) return None, error_message @@ -341,11 +443,11 @@ def _execute_stt( if mime_type not in SUPPORTED_AUDIO_MIMES: error_message = ( f"[KAAPI] STT validation failed: unsupported audio mime " - f"'{mime_type}' for Vertex STT. Supported MIME types are: " + f"'{mime_type}' for Google GCP STT. Supported MIME types are: " f"{', '.join(sorted(SUPPORTED_AUDIO_MIMES))}." ) logger.warning( - f"[GoogleVertexAIProvider._execute_stt] {error_message} | provider={provider}" + f"[GoogleGCPProvider._execute_stt] {error_message} | provider={provider}" ) return None, error_message @@ -353,13 +455,13 @@ def _execute_stt( # the 20 MB inline cap. if not self.client.sa_info: error_message = ( - "[KAAPI] Vertex STT staging failed: ``google`` sa_key is " + "[KAAPI] Google GCP STT staging failed: ``google-gcp`` sa_key is " "not configured on this project's credentials, so audio " "cannot be uploaded to GCS for transcription. Add the " - "service-account key to the project's ``google`` credentials." + "service-account key to the project's ``google-gcp`` credentials." ) logger.warning( - f"[GoogleVertexAIProvider._execute_stt] {error_message} | provider={provider}" + f"[GoogleGCPProvider._execute_stt] {error_message} | provider={provider}" ) return None, error_message @@ -373,14 +475,14 @@ def _execute_stt( ) except Exception as e: error_message = ( - f"[KAAPI] Failed to stage audio for Vertex STT: GCS upload " + f"[KAAPI] Failed to stage audio for Google GCP STT: GCS upload " f"to bucket '{self.client.gcs_bucket}' failed ({str(e)}). " f"Verify the service account has write access to the bucket " f"and that the bucket exists in project " f"'{self.client.project_id}'." ) logger.error( - f"[GoogleVertexAIProvider._execute_stt] {error_message} | provider={provider}", + f"[GoogleGCPProvider._execute_stt] {error_message} | provider={provider}", exc_info=True, ) return None, error_message @@ -439,7 +541,7 @@ def _execute_stt( transcript = data["candidates"][0]["content"]["parts"][0]["text"] except (KeyError, IndexError, TypeError): error_message = ( - "[VERTEX] STT response is missing transcribed text. Vertex " + "[GOOGLE_GCP] STT response is missing transcribed text. Google GCP " "returned a 200 response but the expected " "candidates[0].content.parts[0].text path is absent — this " "typically means the response was blocked by safety filters " @@ -447,7 +549,7 @@ def _execute_stt( "retry." ) logger.warning( - f"[GoogleVertexAIProvider._execute_stt] {error_message} | " + f"[GoogleGCPProvider._execute_stt] {error_message} | " f"provider={provider}, model={model}, response_id={data.get('responseId')}" ) return None, error_message @@ -455,7 +557,7 @@ def _execute_stt( llm_response = LLMCallResponse( response=LLMResponse( provider_response_id=data.get("responseId") - or f"vertex-{uuid.uuid4().hex}", + or f"google-gcp-{uuid.uuid4().hex}", model=data.get("modelVersion") or model, provider=provider, output=TextOutput( @@ -471,7 +573,7 @@ def _execute_stt( llm_response.provider_raw_response = data logger.info( - f"[GoogleVertexAIProvider._execute_stt] Transcribed audio | provider={provider}, model={model}" + f"[GoogleGCPProvider._execute_stt] Transcribed audio | provider={provider}, model={model}" ) return llm_response, None @@ -481,12 +583,12 @@ def _execute_tts( resolved_input: str, include_provider_raw_response: bool = False, ) -> tuple[LLMCallResponse | None, str | None]: - """Execute TTS via Vertex generateContent. + """Execute TTS via Google GCP generateContent. Note: HTTP / network errors come back from ``_post()`` already-logged and tagged. This method only handles Kaapi-side input validation, - Vertex response-shape checks, and audio post-processing failures. + Google GCP response-shape checks, and audio post-processing failures. """ provider = completion_config.provider params = completion_config.params @@ -499,7 +601,7 @@ def _execute_tts( f"synthesize as a plain string." ) logger.warning( - f"[GoogleVertexAIProvider._execute_tts] {error_message} | provider={provider}" + f"[GoogleGCPProvider._execute_tts] {error_message} | provider={provider}" ) return None, error_message if not resolved_input.strip(): @@ -508,7 +610,7 @@ def _execute_tts( "whitespace-only. Provide non-empty text to synthesize." ) logger.warning( - f"[GoogleVertexAIProvider._execute_tts] {error_message} | provider={provider}" + f"[GoogleGCPProvider._execute_tts] {error_message} | provider={provider}" ) return None, error_message @@ -550,15 +652,15 @@ def _execute_tts( audio_b64 = inline["data"] except (KeyError, IndexError, TypeError): error_message = ( - "[VERTEX] TTS response is missing audio data. Vertex returned " + "[GOOGLE_GCP] TTS response is missing audio data. Google GCP returned " "a 200 response but the expected " "candidates[0].content.parts[0].inlineData path is absent — " - "this typically means Vertex was unable to generate audio " + "this typically means Google GCP was unable to generate audio " "from the input. Ensure the input text is properly formatted " "and does not contain unsupported control sequences." ) logger.warning( - f"[GoogleVertexAIProvider._execute_tts] {error_message} | " + f"[GoogleGCPProvider._execute_tts] {error_message} | " f"provider={provider}, model={model}, response_id={data.get('responseId')}" ) return None, error_message @@ -567,13 +669,13 @@ def _execute_tts( raw_pcm = base64.b64decode(audio_b64) except (ValueError, TypeError) as e: error_message = ( - f"[VERTEX] TTS returned invalid base64 audio: {str(e)}. The " + f"[GOOGLE_GCP] TTS returned invalid base64 audio: {str(e)}. The " f"audio payload could not be decoded — this indicates a " - f"corrupted response from Vertex. Retry the request; if the " + f"corrupted response from Google GCP. Retry the request; if the " f"issue persists, contact Kaapi." ) logger.warning( - f"[GoogleVertexAIProvider._execute_tts] {error_message} | " + f"[GoogleGCPProvider._execute_tts] {error_message} | " f"provider={provider}, model={model}", exc_info=True, ) @@ -581,13 +683,13 @@ def _execute_tts( if not raw_pcm: error_message = ( - "[VERTEX] TTS returned empty audio data. Vertex accepted the " + "[GOOGLE_GCP] TTS returned empty audio data. Google GCP accepted the " "request and returned a base64 payload that decoded to zero " - "bytes — this is typically a Vertex server-side issue. Wait " + "bytes — this is typically a Google GCP server-side issue. Wait " "a minute and retry; if the issue persists, contact Kaapi." ) logger.warning( - f"[GoogleVertexAIProvider._execute_tts] {error_message} | " + f"[GoogleGCPProvider._execute_tts] {error_message} | " f"provider={provider}, model={model}" ) return None, error_message @@ -601,11 +703,11 @@ def _execute_tts( if convert_err: error_message = ( f"[KAAPI] Post-processing failure: unable to convert " - f"Vertex PCM audio to MP3 ({convert_err}). Falling back " + f"Google GCP PCM audio to MP3 ({convert_err}). Falling back " f"to WAV is possible by setting response_format='wav'." ) logger.error( - f"[GoogleVertexAIProvider._execute_tts] {error_message} | " + f"[GoogleGCPProvider._execute_tts] {error_message} | " f"provider={provider}, model={model}, pcm_bytes={len(raw_pcm)}" ) return None, error_message @@ -616,11 +718,11 @@ def _execute_tts( if convert_err: error_message = ( f"[KAAPI] Post-processing failure: unable to convert " - f"Vertex PCM audio to OGG ({convert_err}). Falling back " + f"Google GCP PCM audio to OGG ({convert_err}). Falling back " f"to WAV is possible by setting response_format='wav'." ) logger.error( - f"[GoogleVertexAIProvider._execute_tts] {error_message} | " + f"[GoogleGCPProvider._execute_tts] {error_message} | " f"provider={provider}, model={model}, pcm_bytes={len(raw_pcm)}" ) return None, error_message @@ -628,14 +730,14 @@ def _execute_tts( actual_format = "ogg" elif response_format and response_format != "wav": logger.warning( - f"[GoogleVertexAIProvider._execute_tts] Unsupported response_format " + f"[GoogleGCPProvider._execute_tts] Unsupported response_format " f"'{response_format}', returning native WAV | provider={provider}" ) llm_response = LLMCallResponse( response=LLMResponse( provider_response_id=data.get("responseId") - or f"vertex-{uuid.uuid4().hex}", + or f"google-gcp-{uuid.uuid4().hex}", model=data.get("modelVersion") or model, provider=provider, output=AudioOutput( @@ -653,7 +755,7 @@ def _execute_tts( llm_response.provider_raw_response = data logger.info( - f"[GoogleVertexAIProvider._execute_tts] Synthesised audio | " + f"[GoogleGCPProvider._execute_tts] Synthesised audio | " f"provider={provider}, model={model}, format={actual_format}, " f"raw_pcm_bytes={len(raw_pcm)}" ) @@ -669,6 +771,12 @@ def execute( provider = completion_config.provider completion_type = completion_config.type try: + if completion_type == CompletionType.TEXT: + return self._execute_text( + completion_config=completion_config, + resolved_input=resolved_input, + include_provider_raw_response=include_provider_raw_response, + ) if completion_type == CompletionType.STT: return self._execute_stt( completion_config=completion_config, @@ -683,11 +791,10 @@ def execute( ) error_message = ( f"[KAAPI] Unsupported completion type '{completion_type}' for " - f"google provider. Vertex supports 'stt' and 'tts' only; " - f"use the 'google-aistudio' provider for text completions." + f"google-gcp provider. Google GCP supports 'text', 'stt' and 'tts'." ) logger.warning( - f"[GoogleVertexAIProvider.execute] {error_message} | provider={provider}" + f"[GoogleGCPProvider.execute] {error_message} | provider={provider}" ) return None, error_message @@ -695,10 +802,10 @@ def execute( error_message = ( f"[KAAPI] Invalid or unexpected parameter in Config: {str(e)}. " f"Review the completion config; one of the parameters does " - f"not match the Vertex provider's expected signature." + f"not match the Google GCP provider's expected signature." ) logger.warning( - f"[GoogleVertexAIProvider.execute] {error_message} | " + f"[GoogleGCPProvider.execute] {error_message} | " f"provider={provider}, type={completion_type}", exc_info=True, ) @@ -706,13 +813,13 @@ def execute( except Exception as e: error_message = ( - f"[KAAPI] Unexpected error while executing Vertex " + f"[KAAPI] Unexpected error while executing Google GCP " f"{completion_type or 'request'}: {str(e)}. This was not " - f"raised inside the Vertex HTTP call — likely a Kaapi-side " + f"raised inside the Google GCP HTTP call — likely a Kaapi-side " f"failure. Contact Kaapi if the issue persists." ) logger.error( - f"[GoogleVertexAIProvider.execute] {error_message} | " + f"[GoogleGCPProvider.execute] {error_message} | " f"provider={provider}, type={completion_type}", exc_info=True, ) diff --git a/backend/app/services/llm/providers/registry.py b/backend/app/services/llm/providers/registry.py index 04ae9fd7a..b6ed5b882 100644 --- a/backend/app/services/llm/providers/registry.py +++ b/backend/app/services/llm/providers/registry.py @@ -4,6 +4,7 @@ from app.services.llm.providers.base import BaseProvider from app.services.llm.providers.open_ai import OpenAIProvider from app.services.llm.providers.google_aistudio import GoogleAIProvider +from app.services.llm.providers.google_gcp import GoogleGCPProvider from app.services.llm.providers.sarvam_ai import SarvamAIProvider from app.services.llm.providers.eleven_ai import ElevenlabsAIProvider from app.services.llm.providers.claude import ClaudeProvider @@ -19,6 +20,8 @@ class LLMProvider: ANTHROPIC = "anthropic" GOOGLE_AISTUDIO = "google-aistudio" GOOGLE_AISTUDIO_NATIVE = "google-aistudio-native" + GOOGLE_GCP = "google-gcp" + GOOGLE_GCP_NATIVE = "google-gcp-native" OPENAI_NATIVE = "openai-native" GOOGLE_NATIVE = "google-native" SARVAMAI_NATIVE = "sarvamai-native" @@ -33,6 +36,8 @@ class LLMProvider: ANTHROPIC: ClaudeProvider, GOOGLE_AISTUDIO: GoogleAIProvider, GOOGLE_AISTUDIO_NATIVE: GoogleAIProvider, + GOOGLE_GCP: GoogleGCPProvider, + GOOGLE_GCP_NATIVE: GoogleGCPProvider, OPENAI_NATIVE: OpenAIProvider, GOOGLE_NATIVE: GoogleAIProvider, SARVAMAI_NATIVE: SarvamAIProvider, diff --git a/backend/app/tests/crud/test_credentials.py b/backend/app/tests/crud/test_credentials.py index df612727c..8bbfa544c 100644 --- a/backend/app/tests/crud/test_credentials.py +++ b/backend/app/tests/crud/test_credentials.py @@ -1,4 +1,5 @@ import pytest +from fastapi import HTTPException from pydantic import ValidationError from sqlmodel import Session @@ -191,6 +192,69 @@ def test_update_creds_for_org(db: Session) -> None: assert retrieved_cred["api_key"] == "updated-key" +def test_update_creds_for_org_partial_update_preserves_other_fields( + db: Session, +) -> None: + """A PATCH with only one field must not wipe the provider's other + required fields, and must still pass validation against all of them.""" + _, project = create_test_credential(db) + + creds_update = CredsUpdate( + provider="langfuse", credential={"host": "https://updated.langfuse.com"} + ) + + updated = update_creds_for_org( + session=db, + org_id=project.organization_id, + creds_in=creds_update, + project_id=project.id, + ) + + assert len(updated) == 1 + retrieved_cred = get_provider_credential( + session=db, + org_id=project.organization_id, + provider="langfuse", + project_id=project.id, + ) + assert retrieved_cred["host"] == "https://updated.langfuse.com" + assert retrieved_cred["secret_key"] + assert retrieved_cred["public_key"] + + +def test_update_creds_for_org_partial_update_with_no_existing_credential_fails( + db: Session, +) -> None: + """A partial payload has nothing to merge with when the provider has no + stored credential yet, so the crud-level completeness check must reject + it with a 400 rather than persisting an incomplete credential.""" + project = create_test_project(db) + + creds_update = CredsUpdate( + provider="langfuse", credential={"host": "https://new.langfuse.com"} + ) + + with pytest.raises(HTTPException) as exc_info: + update_creds_for_org( + session=db, + org_id=project.organization_id, + creds_in=creds_update, + project_id=project.id, + ) + + assert exc_info.value.status_code == 400 + assert "Missing required fields for langfuse" in exc_info.value.detail + assert ( + get_provider_credential( + session=db, + org_id=project.organization_id, + provider="langfuse", + project_id=project.id, + ) + is None + ) + + def test_remove_provider_credential(db: Session) -> None: """Test removing credentials for a specific provider.""" _, project = create_test_credential(db) diff --git a/backend/app/tests/services/llm/providers/test_google_ai.py b/backend/app/tests/services/llm/providers/test_google_gcp.py similarity index 73% rename from backend/app/tests/services/llm/providers/test_google_ai.py rename to backend/app/tests/services/llm/providers/test_google_gcp.py index 9ad7a7b12..98be39e2c 100644 --- a/backend/app/tests/services/llm/providers/test_google_ai.py +++ b/backend/app/tests/services/llm/providers/test_google_gcp.py @@ -1,4 +1,4 @@ -"""Tests for the Google Vertex AI provider.""" +"""Tests for the Google GCP provider.""" import base64 import json @@ -7,18 +7,18 @@ import pytest import requests -# ponytail: Vertex provider is temporarily disabled in the registry; skip only -# the routing/execute tests. Pure client tests (endpoint, sa_info, credential -# fallback) still run for coverage. -_vertex_disabled = pytest.mark.skip( - reason="Vertex provider disabled — routing swapped to GoogleAIProvider" -) - from app.core.audio_utils import AudioRef -from app.models.llm import NativeCompletionConfig, QueryParams -from app.services.llm.providers.google_ai import ( - GoogleVertexAIProvider, - VertexClient, +from app.models.llm import ( + ImageContent, + NativeCompletionConfig, + PDFContent, + QueryParams, + TextContent, +) +from app.services.llm.providers.base import MultiModalInput +from app.services.llm.providers.google_gcp import ( + GoogleGCPProvider, + GoogleGCPClient, _load_platform_sa_info, ) from app.models.llm.constants import CompletionType @@ -81,16 +81,15 @@ def _mock_http_err(status: int = 400, body: str = "bad request") -> MagicMock: def _mock_gcs(monkeypatch): """Stub out GCS upload so STT tests don't touch external services.""" monkeypatch.setattr( - "app.services.llm.providers.google_ai.upload_audio_to_gcs", + "app.services.llm.providers.google_gcp.upload_audio_to_gcs", lambda *, audio_bytes, bucket_name, sa_info, **kw: f"gs://{bucket_name}/audio/test.wav", ) -@_vertex_disabled -class TestGoogleVertexAIProvider: +class TestGoogleGCPProvider: @pytest.fixture - def client(self) -> VertexClient: - return VertexClient( + def client(self) -> GoogleGCPClient: + return GoogleGCPClient( api_key="k", project_id="p", location="us-central1", @@ -99,8 +98,8 @@ def client(self) -> VertexClient: ) @pytest.fixture - def provider(self, client) -> GoogleVertexAIProvider: - return GoogleVertexAIProvider(client=client) + def provider(self, client) -> GoogleGCPProvider: + return GoogleGCPProvider(client=client) @pytest.fixture def query(self) -> QueryParams: @@ -129,11 +128,17 @@ def tts_config(self) -> NativeCompletionConfig: # ── create_client ──────────────────────────────────────────────────────── def test_create_client_requires_all_fields(self): with pytest.raises(ValueError, match="project_id, location"): - GoogleVertexAIProvider.create_client({"api_key": "k"}) + GoogleGCPProvider.create_client({"api_key": "k"}) def test_create_client_builds_endpoint(self): - c = GoogleVertexAIProvider.create_client( - {"api_key": "k", "project_id": "p", "location": "us-central1"} + c = GoogleGCPProvider.create_client( + { + "api_key": "k", + "project_id": "p", + "location": "us-central1", + "sa_key": {"type": "service_account"}, + "gcs_bucket": "test-bucket", + } ) assert "us-central1-aiplatform.googleapis.com" in c.endpoint("m") assert "projects/p/locations/us-central1" in c.endpoint("m") @@ -142,7 +147,7 @@ def test_create_client_builds_endpoint(self): # ── STT ────────────────────────────────────────────────────────────────── def test_stt_happy_path(self, provider, stt_config, query, audio_ref): with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(_stt_response("hi there")), ) as mock_post: resp, err = provider.execute(stt_config, query, audio_ref) @@ -175,24 +180,24 @@ def test_stt_gcs_upload_failure_returns_clean_error( self, provider, stt_config, query, audio_ref, monkeypatch ): monkeypatch.setattr( - "app.services.llm.providers.google_ai.upload_audio_to_gcs", + "app.services.llm.providers.google_gcp.upload_audio_to_gcs", MagicMock(side_effect=RuntimeError("bucket denied")), ) resp, err = provider.execute(stt_config, query, audio_ref) assert resp is None - assert "Failed to stage audio for Vertex STT" in err + assert "Failed to stage audio for Google GCP STT" in err assert "bucket denied" in err def test_stt_http_error_returns_clean_message( self, provider, stt_config, query, audio_ref ): with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_err(403, "permission denied"), ): resp, err = provider.execute(stt_config, query, audio_ref) assert resp is None - assert "[VERTEX]" in err + assert "[GOOGLE_GCP]" in err assert "403" in err assert "permission denied" in err @@ -200,18 +205,18 @@ def test_stt_network_error_returns_clean_message( self, provider, stt_config, query, audio_ref ): with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", side_effect=requests.ConnectionError("dns boom"), ): resp, err = provider.execute(stt_config, query, audio_ref) assert resp is None - assert "Vertex AI connection failed" in err + assert "Google GCP connection failed" in err def test_stt_missing_transcript_returns_error( self, provider, stt_config, query, audio_ref ): with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok({"candidates": []}), ): resp, err = provider.execute(stt_config, query, audio_ref) @@ -230,7 +235,7 @@ def test_stt_input_language_overrides_prompt(self, provider, query, audio_ref): }, ) with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(_stt_response()), ) as mock_post: provider.execute(config, query, audio_ref) @@ -243,7 +248,7 @@ def test_stt_input_language_overrides_prompt(self, provider, query, audio_ref): # ── TTS ────────────────────────────────────────────────────────────────── def test_tts_happy_path_wav(self, provider, tts_config, query): with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(_tts_response()), ) as mock_post: resp, err = provider.execute(tts_config, query, "hello") @@ -275,7 +280,7 @@ def test_tts_rejects_empty_input(self, provider, tts_config, query): def test_tts_missing_audio_returns_error(self, provider, tts_config, query): with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok({"candidates": [{"content": {"parts": []}}]}), ): resp, err = provider.execute(tts_config, query, "hello") @@ -289,7 +294,7 @@ def test_tts_language_is_forwarded(self, provider, query): params={"model": "gemini-2.5-flash-preview-tts", "language": "en-US"}, ) with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(_tts_response()), ) as mock_post: provider.execute(config, query, "hi") @@ -297,22 +302,26 @@ def test_tts_language_is_forwarded(self, provider, query): assert speech["languageCode"] == "en-US" # ── execute dispatcher ─────────────────────────────────────────────────── - def test_text_completion_is_rejected(self, provider, query): + def test_text_completion_happy_path(self, provider, query): config = NativeCompletionConfig( provider="google-native", type=CompletionType.TEXT, params={"model": "gemini-2.5-flash"}, ) - resp, err = provider.execute(config, query, "hello") - assert resp is None - assert "Unsupported completion type 'text'" in err + with patch( + "app.services.llm.providers.google_gcp.requests.post", + return_value=_mock_http_ok(_stt_response("hi there")), + ): + resp, err = provider.execute(config, query, "hello") + assert err is None + assert resp.response.output.content.value == "hi there" def test_raw_response_included_when_requested( self, provider, stt_config, query, audio_ref ): raw = _stt_response() with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(raw), ): resp, _ = provider.execute( @@ -320,26 +329,130 @@ def test_raw_response_included_when_requested( ) assert resp.provider_raw_response == raw + # ── text ───────────────────────────────────────────────────────────────── + def test_text_forwards_instructions_and_temperature(self, provider, query): + config = NativeCompletionConfig( + provider="google-native", + type=CompletionType.TEXT, + params={ + "model": "gemini-2.5-flash", + "instructions": "be concise", + "temperature": 0.4, + }, + ) + with patch( + "app.services.llm.providers.google_gcp.requests.post", + return_value=_mock_http_ok(_stt_response("hi there")), + ) as mock_post: + provider.execute(config, query, "hello") + payload = mock_post.call_args.kwargs["json"] + assert payload["systemInstruction"] == {"parts": [{"text": "be concise"}]} + assert payload["generationConfig"]["temperature"] == 0.4 + + def test_text_accepts_multimodal_list_input(self, provider, query): + config = NativeCompletionConfig( + provider="google-native", + type=CompletionType.TEXT, + params={"model": "gemini-2.5-flash"}, + ) + parts = [ + TextContent(value="describe this"), + ImageContent(format="base64", value="aW1n", mime_type="image/png"), + ImageContent( + format="url", value="https://x/img.png", mime_type="image/png" + ), + PDFContent(format="base64", value="cGRm", mime_type="application/pdf"), + PDFContent( + format="url", value="https://x/doc.pdf", mime_type="application/pdf" + ), + ] + with patch( + "app.services.llm.providers.google_gcp.requests.post", + return_value=_mock_http_ok(_stt_response("described")), + ) as mock_post: + resp, err = provider.execute(config, query, parts) + assert err is None + sent_parts = mock_post.call_args.kwargs["json"]["contents"][0]["parts"] + assert sent_parts == [ + {"text": "describe this"}, + {"inlineData": {"mimeType": "image/png", "data": "aW1n"}}, + {"fileData": {"mimeType": "image/png", "fileUri": "https://x/img.png"}}, + {"inlineData": {"mimeType": "application/pdf", "data": "cGRm"}}, + { + "fileData": { + "mimeType": "application/pdf", + "fileUri": "https://x/doc.pdf", + } + }, + ] + + def test_text_accepts_multimodal_input_wrapper(self, provider, query): + config = NativeCompletionConfig( + provider="google-native", + type=CompletionType.TEXT, + params={"model": "gemini-2.5-flash"}, + ) + multimodal = MultiModalInput(parts=[TextContent(value="hi")]) + with patch( + "app.services.llm.providers.google_gcp.requests.post", + return_value=_mock_http_ok(_stt_response("hi there")), + ) as mock_post: + resp, err = provider.execute(config, query, multimodal) + assert err is None + sent_parts = mock_post.call_args.kwargs["json"]["contents"][0]["parts"] + assert sent_parts == [{"text": "hi"}] + + def test_text_missing_content_returns_error(self, provider, query): + config = NativeCompletionConfig( + provider="google-native", + type=CompletionType.TEXT, + params={"model": "gemini-2.5-flash"}, + ) + malformed = {"candidates": [], "modelVersion": "gemini-2.5-flash"} + with patch( + "app.services.llm.providers.google_gcp.requests.post", + return_value=_mock_http_ok(malformed), + ): + resp, err = provider.execute(config, query, "hello") + assert resp is None + assert "Text response is missing generated content" in err + + def test_text_raw_response_included_when_requested(self, provider, query): + raw = _stt_response("hi there") + with patch( + "app.services.llm.providers.google_gcp.requests.post", + return_value=_mock_http_ok(raw), + ): + config = NativeCompletionConfig( + provider="google-native", + type=CompletionType.TEXT, + params={"model": "gemini-2.5-flash"}, + ) + resp, _ = provider.execute( + config, query, "hello", include_provider_raw_response=True + ) + assert resp.provider_raw_response == raw + # --------------------------------------------------------------------------- # TTS payload shape — not routing-dependent, kept unskipped for coverage. # --------------------------------------------------------------------------- def test_tts_wraps_input_in_transcript_tags(): - client = VertexClient( + client = GoogleGCPClient( api_key="k", project_id="p", location="us-central1", sa_info={"type": "service_account", "project_id": "p"}, gcs_bucket="test-bucket", ) - provider = GoogleVertexAIProvider(client=client) + provider = GoogleGCPProvider(client=client) config = NativeCompletionConfig( provider="google-native", type=CompletionType.TTS, params={"model": "gemini-2.5-flash-preview-tts", "voice": "Kore"}, ) with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(_tts_response()), ) as mock_post: provider.execute(config, QueryParams(input="ignored"), "Say this text") @@ -352,15 +465,15 @@ def test_tts_wraps_input_in_transcript_tags(): # Standalone _post / _execute_tts / execute() coverage — not routing-dependent, # kept unskipped (same pattern as test_tts_wraps_input_in_transcript_tags). # --------------------------------------------------------------------------- -def _provider() -> GoogleVertexAIProvider: - client = VertexClient( +def _provider() -> GoogleGCPProvider: + client = GoogleGCPClient( api_key="k", project_id="p", location="us-central1", sa_info={"type": "service_account", "project_id": "p"}, gcs_bucket="test-bucket", ) - return GoogleVertexAIProvider(client=client) + return GoogleGCPProvider(client=client) def _tts_config(**params) -> NativeCompletionConfig: @@ -385,7 +498,7 @@ def _tts_config(**params) -> NativeCompletionConfig: def test_post_http_error_status_branches(status, expected_snippet): provider = _provider() with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_err(status, "boom"), ): resp, err = provider.execute(_tts_config(), QueryParams(input="ignored"), "hi") @@ -403,7 +516,7 @@ def test_post_http_error_uses_google_error_envelope(): "error": {"message": "invalid field", "status": "INVALID_ARGUMENT"} } with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=resp_mock, ): _, err = provider.execute(_tts_config(), QueryParams(input="ignored"), "hi") @@ -414,7 +527,7 @@ def test_post_http_error_uses_google_error_envelope(): def test_post_timeout_returns_clean_message(): provider = _provider() with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", side_effect=requests.Timeout("too slow"), ): resp, err = provider.execute(_tts_config(), QueryParams(input="ignored"), "hi") @@ -425,7 +538,7 @@ def test_post_timeout_returns_clean_message(): def test_post_request_exception_returns_clean_message(): provider = _provider() with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", side_effect=requests.RequestException("weird"), ): resp, err = provider.execute(_tts_config(), QueryParams(input="ignored"), "hi") @@ -440,7 +553,7 @@ def test_post_non_json_success_response_returns_clean_message(): resp_mock.status_code = 200 resp_mock.json.side_effect = ValueError("no json") with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=resp_mock, ): resp, err = provider.execute(_tts_config(), QueryParams(input="ignored"), "hi") @@ -465,7 +578,7 @@ def test_execute_tts_rejects_empty_input(): def test_execute_tts_missing_audio_data_returns_error(): provider = _provider() with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok({"candidates": [{"content": {"parts": []}}]}), ): resp, err = provider.execute(_tts_config(), QueryParams(input="ignored"), "hi") @@ -481,7 +594,7 @@ def test_execute_tts_invalid_base64_returns_error(): ] } with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(bad), ): resp, err = provider.execute(_tts_config(), QueryParams(input="ignored"), "hi") @@ -493,7 +606,7 @@ def test_execute_tts_empty_audio_bytes_returns_error(): provider = _provider() empty = {"candidates": [{"content": {"parts": [{"inlineData": {"data": ""}}]}}]} with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(empty), ): resp, err = provider.execute(_tts_config(), QueryParams(input="ignored"), "hi") @@ -504,10 +617,10 @@ def test_execute_tts_empty_audio_bytes_returns_error(): def test_execute_tts_mp3_conversion_success(): provider = _provider() with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(_tts_response()), ), patch( - "app.services.llm.providers.google_ai.convert_pcm_to_mp3", + "app.services.llm.providers.google_gcp.convert_pcm_to_mp3", return_value=(b"mp3bytes", None), ): resp, err = provider.execute( @@ -520,10 +633,10 @@ def test_execute_tts_mp3_conversion_success(): def test_execute_tts_mp3_conversion_failure_returns_error(): provider = _provider() with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(_tts_response()), ), patch( - "app.services.llm.providers.google_ai.convert_pcm_to_mp3", + "app.services.llm.providers.google_gcp.convert_pcm_to_mp3", return_value=(None, "ffmpeg missing"), ): resp, err = provider.execute( @@ -537,10 +650,10 @@ def test_execute_tts_mp3_conversion_failure_returns_error(): def test_execute_tts_ogg_conversion_success(): provider = _provider() with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(_tts_response()), ), patch( - "app.services.llm.providers.google_ai.convert_pcm_to_ogg", + "app.services.llm.providers.google_gcp.convert_pcm_to_ogg", return_value=(b"oggbytes", None), ): resp, err = provider.execute( @@ -553,7 +666,7 @@ def test_execute_tts_ogg_conversion_success(): def test_execute_tts_unsupported_response_format_falls_back_to_wav(): provider = _provider() with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(_tts_response()), ): resp, err = provider.execute( @@ -563,22 +676,26 @@ def test_execute_tts_unsupported_response_format_falls_back_to_wav(): assert resp.response.output.content.mime_type == "audio/wav" -def test_execute_rejects_unsupported_completion_type(): +def test_execute_text_happy_path(): provider = _provider() config = NativeCompletionConfig( provider="google-native", type=CompletionType.TEXT, params={"model": "gemini-2.5-flash"}, ) - resp, err = provider.execute(config, QueryParams(input="ignored"), "hi") - assert resp is None - assert "Unsupported completion type" in err + with patch( + "app.services.llm.providers.google_gcp.requests.post", + return_value=_mock_http_ok(_stt_response("hi there")), + ): + resp, err = provider.execute(config, QueryParams(input="ignored"), "hi") + assert err is None + assert resp.response.output.content.value == "hi there" def test_execute_tts_language_is_forwarded(): provider = _provider() with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(_tts_response()), ) as mock_post: provider.execute( @@ -592,7 +709,7 @@ def test_execute_tts_director_notes_set_system_instruction(): provider = _provider() config = _tts_config(provider_specific={"gemini": {"director_notes": "whisper"}}) with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(_tts_response()), ) as mock_post: provider.execute(config, QueryParams(input="ignored"), "hi") @@ -603,10 +720,10 @@ def test_execute_tts_director_notes_set_system_instruction(): def test_execute_tts_ogg_conversion_failure_returns_error(): provider = _provider() with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(_tts_response()), ), patch( - "app.services.llm.providers.google_ai.convert_pcm_to_ogg", + "app.services.llm.providers.google_gcp.convert_pcm_to_ogg", return_value=(None, "codec missing"), ): resp, err = provider.execute( @@ -621,7 +738,7 @@ def test_execute_tts_raw_response_included_when_requested(): provider = _provider() raw = _tts_response() with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=_mock_http_ok(raw), ): resp, _ = provider.execute( @@ -636,7 +753,7 @@ def test_execute_tts_raw_response_included_when_requested(): def test_post_connection_error_returns_clean_message(): provider = _provider() with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", side_effect=requests.ConnectionError("dns boom"), ): resp, err = provider.execute(_tts_config(), QueryParams(input="ignored"), "hi") @@ -652,7 +769,7 @@ def test_post_http_error_falls_back_to_raw_text_when_body_not_json(): resp_mock.text = "plain text error" resp_mock.json.side_effect = ValueError("not json") with patch( - "app.services.llm.providers.google_ai.requests.post", + "app.services.llm.providers.google_gcp.requests.post", return_value=resp_mock, ): _, err = provider.execute(_tts_config(), QueryParams(input="ignored"), "hi") @@ -660,9 +777,7 @@ def test_post_http_error_falls_back_to_raw_text_when_body_not_json(): def test_execute_dispatches_stt_validation(): - """Only the dispatch branch + input validation are exercised here — the - full STT network path is covered by the (currently skipped) routing - tests.""" + """Only the dispatch branch + input validation are exercised here.""" provider = _provider() config = NativeCompletionConfig( provider="google-native", @@ -695,11 +810,11 @@ def test_execute_wraps_unexpected_exception(): # --------------------------------------------------------------------------- -# VertexClient.endpoint — host changes by location +# GoogleGCPClient.endpoint — host changes by location # --------------------------------------------------------------------------- -class TestVertexEndpoint: - def _client(self, location: str) -> VertexClient: - return VertexClient( +class TestGoogleGCPEndpoint: + def _client(self, location: str) -> GoogleGCPClient: + return GoogleGCPClient( api_key="k", project_id="my-proj", location=location, @@ -742,32 +857,32 @@ def _sample_sa(self) -> dict: "private_key": "-----BEGIN PRIVATE KEY-----\nfake\n-----END PRIVATE KEY-----", } - @patch("app.services.llm.providers.google_ai.settings") + @patch("app.services.llm.providers.google_gcp.settings") def test_returns_none_when_unset(self, mock_settings): mock_settings.GCP_SA_KEY = "" assert _load_platform_sa_info() is None - @patch("app.services.llm.providers.google_ai.settings") + @patch("app.services.llm.providers.google_gcp.settings") def test_parses_raw_json_string(self, mock_settings): sa = self._sample_sa() mock_settings.GCP_SA_KEY = json.dumps(sa) assert _load_platform_sa_info() == sa - @patch("app.services.llm.providers.google_ai.settings") + @patch("app.services.llm.providers.google_gcp.settings") def test_strips_surrounding_whitespace(self, mock_settings): """env-var injection often leaves trailing newlines — must still parse.""" sa = self._sample_sa() mock_settings.GCP_SA_KEY = "\n " + json.dumps(sa) + " \n" assert _load_platform_sa_info() == sa - @patch("app.services.llm.providers.google_ai.settings") + @patch("app.services.llm.providers.google_gcp.settings") def test_returns_none_on_malformed_json(self, mock_settings): """A JSON-looking but invalid value must not raise — it returns None and lets create_client raise the missing-fields ValueError later.""" mock_settings.GCP_SA_KEY = "{not valid json" assert _load_platform_sa_info() is None - @patch("app.services.llm.providers.google_ai.settings") + @patch("app.services.llm.providers.google_gcp.settings") def test_non_json_string_returns_none(self, mock_settings): """Anything not starting with '{' is treated as non-JSON and ignored — this guards against accidentally interpreting a path or sentinel as a key.""" @@ -779,7 +894,7 @@ def test_non_json_string_returns_none(self, mock_settings): # create_client — credential precedence (BYOK overrides platform settings) # --------------------------------------------------------------------------- class TestCreateClientFallback: - @patch("app.services.llm.providers.google_ai.settings") + @patch("app.services.llm.providers.google_gcp.settings") def test_byok_overrides_platform_settings(self, mock_settings): mock_settings.GCP_VERTEX_API_KEY = "platform-key" mock_settings.GCP_PROJECT_ID = "platform-proj" @@ -787,11 +902,12 @@ def test_byok_overrides_platform_settings(self, mock_settings): mock_settings.GCP_SA_KEY = "" mock_settings.GCS_AUDIO_BUCKET = "platform-bucket" - c = GoogleVertexAIProvider.create_client( + c = GoogleGCPProvider.create_client( { "api_key": "byok-key", "project_id": "byok-proj", "location": "europe-west4", + "sa_key": {"type": "service_account"}, "gcs_bucket": "byok-bucket", } ) @@ -800,21 +916,21 @@ def test_byok_overrides_platform_settings(self, mock_settings): assert c.location == "europe-west4" assert c.gcs_bucket == "byok-bucket" - @patch("app.services.llm.providers.google_ai.settings") + @patch("app.services.llm.providers.google_gcp.settings") def test_partial_byok_fills_from_platform(self, mock_settings): """When BYOK only supplies api_key, project/location come from settings.""" mock_settings.GCP_VERTEX_API_KEY = "platform-key" mock_settings.GCP_PROJECT_ID = "platform-proj" mock_settings.GCP_VERTEX_LOCATION = "us-central1" - mock_settings.GCP_SA_KEY = "" + mock_settings.GCP_SA_KEY = json.dumps({"type": "service_account"}) mock_settings.GCS_AUDIO_BUCKET = "platform-bucket" - c = GoogleVertexAIProvider.create_client({"api_key": "byok-key"}) + c = GoogleGCPProvider.create_client({"api_key": "byok-key"}) assert c.api_key == "byok-key" assert c.project_id == "platform-proj" assert c.location == "us-central1" - @patch("app.services.llm.providers.google_ai.settings") + @patch("app.services.llm.providers.google_gcp.settings") def test_missing_everything_raises_value_error(self, mock_settings): mock_settings.GCP_VERTEX_API_KEY = "" mock_settings.GCP_PROJECT_ID = "" @@ -823,8 +939,10 @@ def test_missing_everything_raises_value_error(self, mock_settings): mock_settings.GCS_AUDIO_BUCKET = "" with pytest.raises(ValueError) as exc_info: - GoogleVertexAIProvider.create_client({}) + GoogleGCPProvider.create_client({}) msg = str(exc_info.value) assert "api_key" in msg assert "project_id" in msg assert "location" in msg + assert "sa_key" in msg + assert "gcs_bucket" in msg diff --git a/backend/app/tests/services/llm/test_mappers.py b/backend/app/tests/services/llm/test_mappers.py index aa032a23d..77a862788 100644 --- a/backend/app/tests/services/llm/test_mappers.py +++ b/backend/app/tests/services/llm/test_mappers.py @@ -939,6 +939,46 @@ def test_unsupported_language_emits_warning(self, db: Session): assert "xx-YY" in warnings[0] +class TestTransformGoogleGCPRouting: + """Routing contract for the ``google-gcp`` provider.""" + + def test_text_completion_maps_via_google_mapper(self, db: Session): + """``google-gcp`` text completions reuse the Google mapper and + produce a ``google-gcp-native`` config.""" + kaapi_config = build_kaapi_completion_config( + provider="google-gcp", + type="text", + params={"model": "gemini-2.5-pro"}, + ) + + native_config, warnings = transform_kaapi_config_to_native( + session=db, kaapi_config=kaapi_config + ) + + assert native_config.provider == "google-gcp-native" + assert native_config.type == "text" + assert native_config.params["model"] == "gemini-2.5-pro" + assert warnings == [] + + def test_stt_completion_maps_via_google_mapper(self, db: Session): + """``google-gcp`` STT completions reuse the Google mapper and + produce a ``google-gcp-native`` config.""" + kaapi_config = build_kaapi_completion_config( + provider="google-gcp", + type="stt", + params={"model": "gemini-2.5-pro", "input_language": "hi-IN"}, + ) + + native_config, warnings = transform_kaapi_config_to_native( + session=db, kaapi_config=kaapi_config + ) + + assert native_config.provider == "google-gcp-native" + assert native_config.type == "stt" + assert native_config.params["input_language"] == "hi-IN" + assert warnings == [] + + class TestBCP47ToElevenlabsLang: """Test BCP-47 language code conversion for ElevenLabs."""