From 9f55b91277caf90c72a1a4bc4f82d6e26b1d7998 Mon Sep 17 00:00:00 2001 From: Prajna1999 Date: Mon, 17 Aug 2026 09:35:58 +0530 Subject: [PATCH 1/2] feat: add google-gcp as provider --- backend/app/core/cloud/storage.py | 2 +- backend/app/core/providers.py | 11 ++ backend/app/models/llm/constants.py | 5 + backend/app/models/llm/request.py | 4 +- backend/app/models/model_config.py | 4 +- backend/app/services/llm/mappers.py | 14 ++ .../app/services/llm/providers/__init__.py | 2 +- .../providers/{google_ai.py => google_gcp.py} | 150 +++++++++--------- .../app/services/llm/providers/registry.py | 5 + .../{test_google_ai.py => test_google_gcp.py} | 148 ++++++++--------- 10 files changed, 186 insertions(+), 159 deletions(-) rename backend/app/services/llm/providers/{google_ai.py => google_gcp.py} (81%) rename backend/app/tests/services/llm/providers/{test_google_ai.py => test_google_gcp.py} (86%) 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 743ee1883..5e5c4c8ad 100644 --- a/backend/app/core/providers.py +++ b/backend/app/core/providers.py @@ -12,6 +12,7 @@ class Provider(str, Enum): OPENAI = "openai" LANGFUSE = "langfuse" GOOGLE_AISTUDIO = "google-aistudio" + GOOGLE_GCP = "google-gcp" SARVAMAI = "sarvamai" ELEVENLABS = "elevenlabs" ANTHROPIC = "anthropic" @@ -55,6 +56,16 @@ class ProviderConfig: ], sensitive_fields=["api_key"], ), + Provider.GOOGLE_GCP: ProviderConfig( + required_fields=[ + "api_key", + "project_id", + "location", + "sa_key", + "gcs_bucket", + ], + sensitive_fields=["api_key", "sa_key"], + ), Provider.WEBHOOK_SECRET: ProviderConfig( required_fields=["webhook_secret"], sensitive_fields=["webhook_secret"] ), diff --git a/backend/app/models/llm/constants.py b/backend/app/models/llm/constants.py index b936535f1..7ac980b48 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, @@ -32,6 +35,7 @@ class Provider(StrEnum): KaapiProvider = Literal[ Provider.OPENAI, Provider.GOOGLE, + Provider.GOOGLE_GCP, Provider.SARVAMAI, Provider.ELEVENLABS, Provider.ANTHROPIC, @@ -49,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 59242006e..8b15cb047 100644 --- a/backend/app/models/llm/request.py +++ b/backend/app/models/llm/request.py @@ -295,8 +295,8 @@ class KaapiCompletionConfig(SQLModel): None, description=( "LLM provider (openai, google, sarvamai, elevenlabs, anthropic, " - "google-aistudio). 'google' routes via Google Vertex AI; " - "'google-aistudio' uses Google AI Studio." + "google-aistudio, google-gcp). 'google-aistudio' uses Google AI " + "Studio; 'google-gcp' uses Google GCP directly for STT/TTS." ), ) 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 75cf626f9..261be3d21 100644 --- a/backend/app/services/llm/mappers.py +++ b/backend/app/services/llm/mappers.py @@ -696,6 +696,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 81% rename from backend/app/services/llm/providers/google_ai.py rename to backend/app/services/llm/providers/google_gcp.py index bbbd96792..017b82903 100644 --- a/backend/app/services/llm/providers/google_ai.py +++ b/backend/app/services/llm/providers/google_gcp.py @@ -71,8 +71,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 +105,15 @@ 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 + models on GCP. Text-only completions are routed through the standard `google` provider. """ - 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}" ) @@ -144,9 +144,9 @@ def create_client(credentials: dict[str, Any]) -> Any: ] 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 +157,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 +181,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 +232,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 +277,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 +286,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 @@ -316,12 +316,12 @@ def _execute_stt( 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 +333,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 +341,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 +353,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 +373,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 +439,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 +447,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 +455,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(content=TextContent(value=transcript.strip())), @@ -467,7 +467,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 @@ -477,12 +477,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 @@ -495,7 +495,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(): @@ -504,7 +504,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 @@ -546,15 +546,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 @@ -563,13 +563,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, ) @@ -577,13 +577,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 @@ -597,11 +597,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 @@ -612,11 +612,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 @@ -624,14 +624,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( @@ -649,7 +649,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)}" ) @@ -679,11 +679,11 @@ def execute( ) error_message = ( f"[KAAPI] Unsupported completion type '{completion_type}' for " - f"google provider. Vertex supports 'stt' and 'tts' only; " + f"google-gcp provider. Google GCP supports 'stt' and 'tts' only; " f"use the 'google-aistudio' provider for text completions." ) logger.warning( - f"[GoogleVertexAIProvider.execute] {error_message} | provider={provider}" + f"[GoogleGCPProvider.execute] {error_message} | provider={provider}" ) return None, error_message @@ -691,10 +691,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, ) @@ -702,13 +702,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/services/llm/providers/test_google_ai.py b/backend/app/tests/services/llm/providers/test_google_gcp.py similarity index 86% 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..350f6c03c 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,11 @@ 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.services.llm.providers.google_gcp import ( + GoogleGCPProvider, + GoogleGCPClient, _load_platform_sa_info, ) from app.models.llm.constants import CompletionType @@ -81,16 +74,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 +91,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,10 +121,10 @@ 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( + c = GoogleGCPProvider.create_client( {"api_key": "k", "project_id": "p", "location": "us-central1"} ) assert "us-central1-aiplatform.googleapis.com" in c.endpoint("m") @@ -142,7 +134,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 +167,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 +192,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 +222,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 +235,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 +267,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 +281,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") @@ -312,7 +304,7 @@ def test_raw_response_included_when_requested( ): 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( @@ -325,21 +317,21 @@ def test_raw_response_included_when_requested( # 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 +344,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 +377,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 +395,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 +406,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 +417,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 +432,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 +457,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 +473,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 +485,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 +496,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 +512,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 +529,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 +545,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( @@ -578,7 +570,7 @@ def test_execute_rejects_unsupported_completion_type(): 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 +584,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 +595,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 +613,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 +628,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 +644,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 +652,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 +685,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 +732,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 +769,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,7 +777,7 @@ 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", @@ -800,7 +790,7 @@ 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" @@ -809,12 +799,12 @@ def test_partial_byok_fills_from_platform(self, mock_settings): mock_settings.GCP_SA_KEY = "" 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,7 +813,7 @@ 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 From 8942549337060f638c002a9b7fbb8648b71826d2 Mon Sep 17 00:00:00 2001 From: Prajna1999 Date: Thu, 20 Aug 2026 22:04:55 +0530 Subject: [PATCH 2/2] partial update for creds --- backend/app/crud/credentials.py | 28 +++++++++++------ .../app/services/llm/providers/google_gcp.py | 4 +-- backend/app/tests/crud/test_credentials.py | 30 +++++++++++++++++++ .../services/llm/providers/test_google_gcp.py | 13 ++++++-- 4 files changed, 62 insertions(+), 13 deletions(-) diff --git a/backend/app/crud/credentials.py b/backend/app/crud/credentials.py index cc070852b..860e9371b 100644 --- a/backend/app/crud/credentials.py +++ b/backend/app/crud/credentials.py @@ -199,8 +199,25 @@ def update_creds_for_org( ): credential_data = credential_data[creds_in.provider] + statement = select(Credential).where( + Credential.organization_id == org_id, + Credential.provider == creds_in.provider, + Credential.is_active.is_(True), + 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: - validate_provider_credentials(creds_in.provider, credential_data) + validate_provider_credentials(creds_in.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: {creds_in.provider}, error: {str(e)}" @@ -208,15 +225,8 @@ def update_creds_for_org( raise HTTPException(status_code=400, detail=str(e)) # Encrypt the entire credentials object - encrypted_credentials = encrypt_credentials(credential_data) + encrypted_credentials = encrypt_credentials(merged_credential_data) - statement = select(Credential).where( - Credential.organization_id == org_id, - Credential.provider == creds_in.provider, - Credential.is_active.is_(True), - Credential.project_id == project_id, - ) - creds = session.exec(statement).one_or_none() if creds is None: # Create new credential if it doesn't exist creds = Credential( diff --git a/backend/app/services/llm/providers/google_gcp.py b/backend/app/services/llm/providers/google_gcp.py index 017b82903..a6ca98661 100644 --- a/backend/app/services/llm/providers/google_gcp.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 @@ -139,6 +137,8 @@ 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 ] diff --git a/backend/app/tests/crud/test_credentials.py b/backend/app/tests/crud/test_credentials.py index 8b3269ca6..c0c603c97 100644 --- a/backend/app/tests/crud/test_credentials.py +++ b/backend/app/tests/crud/test_credentials.py @@ -188,6 +188,36 @@ 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_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_gcp.py b/backend/app/tests/services/llm/providers/test_google_gcp.py index 350f6c03c..f69073578 100644 --- a/backend/app/tests/services/llm/providers/test_google_gcp.py +++ b/backend/app/tests/services/llm/providers/test_google_gcp.py @@ -125,7 +125,13 @@ def test_create_client_requires_all_fields(self): def test_create_client_builds_endpoint(self): c = GoogleGCPProvider.create_client( - {"api_key": "k", "project_id": "p", "location": "us-central1"} + { + "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") @@ -782,6 +788,7 @@ def test_byok_overrides_platform_settings(self, mock_settings): "api_key": "byok-key", "project_id": "byok-proj", "location": "europe-west4", + "sa_key": {"type": "service_account"}, "gcs_bucket": "byok-bucket", } ) @@ -796,7 +803,7 @@ def test_partial_byok_fills_from_platform(self, mock_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 = GoogleGCPProvider.create_client({"api_key": "byok-key"}) @@ -818,3 +825,5 @@ def test_missing_everything_raises_value_error(self, mock_settings): assert "api_key" in msg assert "project_id" in msg assert "location" in msg + assert "sa_key" in msg + assert "gcs_bucket" in msg