From 9f55b91277caf90c72a1a4bc4f82d6e26b1d7998 Mon Sep 17 00:00:00 2001 From: Prajna1999 Date: Mon, 17 Aug 2026 09:35:58 +0530 Subject: [PATCH 01/15] 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 5cc91d94840a21143e4e60d15269ed66bea81464 Mon Sep 17 00:00:00 2001 From: Prashant Vasudevan <71649489+vprashrex@users.noreply.github.com> Date: Wed, 19 Aug 2026 10:27:29 +0530 Subject: [PATCH 02/15] feat(bucket): Implement GCS bucket provider with signed URL generation - Added GCSBucketProvider for handling Google Cloud Storage (GCS) operations, including signed URL generation. - Introduced GCSClient to manage GCS client instances with default bucket support. - Created a global bucket-provider registry to resolve and manage different bucket providers. - Implemented tests for GCS bucket provider functionality, including signed URL generation and credential handling. - Enhanced attachment resolution to support GCS URIs, allowing for both native and signed URL handling based on provider type. - Updated documentation to reflect changes in bucket provider architecture and functionality. --- backend/app/core/batch/__init__.py | 11 +- backend/app/core/batch/client.py | 50 ++++ backend/app/core/batch/vertex.py | 266 ++++++++++++++++++ backend/app/crud/assessment/batch.py | 11 + backend/app/models/llm/constants.py | 7 +- backend/app/services/assessment/api/batch.py | 46 ++- .../app/services/assessment/api/submission.py | 3 +- backend/app/services/assessment/tasks.py | 15 + .../services/assessment/utils/attachments.py | 54 +++- backend/app/services/buckets/__init__.py | 0 backend/app/services/buckets/attachments.py | 83 ++++++ .../services/buckets/providers/__init__.py | 0 .../app/services/buckets/providers/base.py | 48 ++++ backend/app/services/buckets/providers/gcs.py | 112 ++++++++ .../services/buckets/providers/registry.py | 92 ++++++ .../app/tests/assessment/test_api_batch.py | 63 ++++- .../tests/assessment/test_api_submission.py | 7 + backend/app/tests/assessment/test_batch.py | 43 +++ backend/app/tests/core/batch/test_client.py | 76 +++++ backend/app/tests/core/batch/test_vertex.py | 185 ++++++++++++ .../app/tests/services/buckets/__init__.py | 0 .../services/buckets/test_attachments.py | 126 +++++++++ .../app/tests/services/buckets/test_gcs.py | 156 ++++++++++ .../tests/services/buckets/test_registry.py | 102 +++++++ docs/wiki/domain-map.md | 4 +- docs/wiki/modules/assessment.md | 4 +- docs/wiki/modules/platform.md | 1 + 27 files changed, 1538 insertions(+), 27 deletions(-) create mode 100644 backend/app/core/batch/vertex.py create mode 100644 backend/app/services/buckets/__init__.py create mode 100644 backend/app/services/buckets/attachments.py create mode 100644 backend/app/services/buckets/providers/__init__.py create mode 100644 backend/app/services/buckets/providers/base.py create mode 100644 backend/app/services/buckets/providers/gcs.py create mode 100644 backend/app/services/buckets/providers/registry.py create mode 100644 backend/app/tests/core/batch/test_client.py create mode 100644 backend/app/tests/core/batch/test_vertex.py create mode 100644 backend/app/tests/services/buckets/__init__.py create mode 100644 backend/app/tests/services/buckets/test_attachments.py create mode 100644 backend/app/tests/services/buckets/test_gcs.py create mode 100644 backend/app/tests/services/buckets/test_registry.py diff --git a/backend/app/core/batch/__init__.py b/backend/app/core/batch/__init__.py index 1e8202f96..291dc63db 100644 --- a/backend/app/core/batch/__init__.py +++ b/backend/app/core/batch/__init__.py @@ -2,7 +2,12 @@ from .anthropic import AnthropicBatchProvider, MessageBatchStatus from .base import BATCH_KEY, BatchProvider -from .client import GeminiClient, GeminiClientError +from .client import ( + GeminiClient, + GeminiClientError, + get_gemini_batch_provider, + is_vertex_batch_provider, +) from .gemini import ( BatchJobState, GeminiBatchProvider, @@ -11,6 +16,7 @@ extract_text_from_response_dict, ) from .openai import OpenAIBatchProvider +from .vertex import VertexBatchProvider from .operations import ( download_batch_results, process_completed_batch, @@ -28,6 +34,9 @@ "GeminiClient", "GeminiClientError", "GeminiBatchProvider", + "VertexBatchProvider", + "get_gemini_batch_provider", + "is_vertex_batch_provider", "OpenAIBatchProvider", "create_stt_batch_requests", "create_tts_batch_requests", diff --git a/backend/app/core/batch/client.py b/backend/app/core/batch/client.py index 804793a69..66bc06c0d 100644 --- a/backend/app/core/batch/client.py +++ b/backend/app/core/batch/client.py @@ -9,6 +9,8 @@ from fastapi import HTTPException from app.crud.credentials import get_provider_credential +from .base import BatchProvider + logger = logging.getLogger(__name__) @@ -91,3 +93,51 @@ def from_credentials( f"org_id: {org_id}, project_id: {project_id}" ) return cls(api_key=api_key) + + +# Providers whose batch jobs run on Vertex AI (GCS-backed) rather than AI-Studio. +VERTEX_BATCH_PROVIDERS = {"google-gcp", "google-gcp-native"} + + +def is_vertex_batch_provider(provider_name: str) -> bool: + """Whether a provider name routes batch jobs through Vertex AI (vs AI-Studio).""" + return provider_name in VERTEX_BATCH_PROVIDERS + + +def get_gemini_batch_provider( + *, + session: Session, + organization_id: int, + project_id: int, + provider_name: str, + model: str | None = None, +) -> BatchProvider: + """Resolve a Gemini-family batch provider from the provider name. + + ``google-gcp`` routes to Vertex (GCS-backed); ``google``/``google-aistudio`` + to AI-Studio (File API). Both honor the same BatchProvider contract, so any + service (assessment, evaluations, STT/TTS) can call this single resolver. + """ + from .gemini import GeminiBatchProvider + from .vertex import VertexBatchProvider + + if is_vertex_batch_provider(provider_name): + cred = get_provider_credential( + session=session, + provider="google-gcp", + project_id=project_id, + org_id=organization_id, + ) + if not cred: + raise HTTPException( + status_code=404, + detail="google-gcp credentials not configured for this project", + ) + if not isinstance(cred, dict): + raise GeminiClientError("Expected decrypted google-gcp credentials dict") + return VertexBatchProvider.from_credentials(cred, model=model) + + gemini = GeminiClient.from_credentials( + session=session, org_id=organization_id, project_id=project_id + ) + return GeminiBatchProvider(client=gemini.client, model=model) diff --git a/backend/app/core/batch/vertex.py b/backend/app/core/batch/vertex.py new file mode 100644 index 000000000..01870a954 --- /dev/null +++ b/backend/app/core/batch/vertex.py @@ -0,0 +1,266 @@ +"""Google GCP Vertex AI batch provider implementation. + +Vertex batch prediction reads its input JSONL from GCS and writes results back to +GCS (no File API, unlike the AI-Studio ``GeminiBatchProvider``). Input/output +therefore ride the project's ``google-gcp`` credential (SA key + ``gcs_bucket``). +""" + +import json +import logging +import time +from typing import Any +from uuid import uuid4 + +from google import genai +from google.genai import types +from google.cloud import storage as gcs +from google.oauth2 import service_account + +from app.core.cloud.storage import GCS_SCOPES, CloudStorageError + +from .base import BATCH_KEY, BatchProvider +from .gemini import BatchJobState + +logger = logging.getLogger(__name__) + +# Terminal Vertex job states (superset of AI-Studio: Vertex adds PAUSED). +_TERMINAL_STATES = { + BatchJobState.SUCCEEDED.value, + BatchJobState.FAILED.value, + BatchJobState.CANCELLED.value, + BatchJobState.EXPIRED.value, + "JOB_STATE_PAUSED", +} +_FAILED_STATES = { + BatchJobState.FAILED.value, + BatchJobState.CANCELLED.value, + BatchJobState.EXPIRED.value, +} + +_DEFAULT_INPUT_PREFIX = "batch-input" +_DEFAULT_OUTPUT_PREFIX = "batch-output" + + +def _parse_gs_uri(uri: str) -> tuple[str, str]: + """Split ``gs://bucket/key`` into ``(bucket, key)``.""" + if not uri.startswith("gs://"): + raise ValueError(f"Expected a gs:// URI, got '{uri}'.") + bucket, _, key = uri[len("gs://") :].partition("/") + return bucket, key + + +class VertexBatchProvider(BatchProvider): + """Vertex AI implementation of the BatchProvider interface (GCS in/out). + + Each JSONL line is the Vertex request schema, e.g. + {"request": {"contents": [{"parts": [...], "role": "user"}]}} + """ + + DEFAULT_MODEL = "gemini-2.5-pro" + + def __init__( + self, + client: genai.Client, + storage_client: gcs.Client, + gcs_bucket: str, + model: str | None = None, + input_prefix: str = _DEFAULT_INPUT_PREFIX, + output_prefix: str = _DEFAULT_OUTPUT_PREFIX, + ) -> None: + self._client = client + self._storage = storage_client + self._bucket = gcs_bucket + self._model = model or self.DEFAULT_MODEL + self._input_prefix = input_prefix + self._output_prefix = output_prefix + + @classmethod + def from_credentials( + cls, credentials: dict[str, Any], model: str | None = None + ) -> "VertexBatchProvider": + """Build a Vertex batch provider from a ``google-gcp`` credential dict.""" + project_id = credentials.get("project_id") + location = credentials.get("location") + gcs_bucket = credentials.get("gcs_bucket") + sa_info = credentials.get("sa_key") + missing = [ + name + for name, value in ( + ("project_id", project_id), + ("location", location), + ("gcs_bucket", gcs_bucket), + ("sa_key", sa_info), + ) + if not value + ] + if missing: + raise ValueError( + f"Vertex batch provider missing required fields: {', '.join(missing)}" + ) + + creds = service_account.Credentials.from_service_account_info( + sa_info, scopes=list(GCS_SCOPES) + ) + client = genai.Client( + vertexai=True, project=project_id, location=location, credentials=creds + ) + storage_client = gcs.Client(project=project_id, credentials=creds) + return cls( + client=client, + storage_client=storage_client, + gcs_bucket=gcs_bucket, + model=model, + ) + + def create_batch( + self, jsonl_data: list[dict[str, Any]], config: dict[str, Any] + ) -> dict[str, Any]: + """Upload input JSONL to GCS and start a Vertex batch prediction job.""" + model = config.get("model", self._model) + display_name = config.get("display_name", f"batch-{int(time.time())}") + + jsonl_content = "\n".join( + json.dumps(item, ensure_ascii=False) for item in jsonl_data + ) + src_uri = self.upload_file(jsonl_content, purpose="batch") + dest_uri = f"gs://{self._bucket}/{self._output_prefix}/{uuid4().hex}/" + + logger.info( + f"[create_batch] Creating Vertex batch | items={len(jsonl_data)} | " + f"model={model} | src={src_uri} | dest={dest_uri}" + ) + + try: + batch_job = self._client.batches.create( + model=model, + src=src_uri, + config=types.CreateBatchJobConfig( + dest=dest_uri, display_name=display_name + ), + ) + initial_state = batch_job.state.name if batch_job.state else "UNKNOWN" + result = { + "provider_batch_id": batch_job.name, + "provider_file_id": src_uri, + "provider_output_prefix": dest_uri, + "provider_status": initial_state, + "total_items": len(jsonl_data), + } + logger.info( + f"[create_batch] Created Vertex batch | batch_id={batch_job.name} | " + f"status={initial_state} | items={len(jsonl_data)}" + ) + return result + except Exception as e: + logger.error(f"[create_batch] Failed to create Vertex batch | {e}") + raise + + def get_batch_status(self, batch_id: str) -> dict[str, Any]: + """Poll Vertex for batch job status.""" + logger.info(f"[get_batch_status] Polling Vertex batch | batch_id={batch_id}") + try: + batch_job = self._client.batches.get(name=batch_id) + state = batch_job.state.name if batch_job.state else "UNKNOWN" + # Results live in the job's GCS dest; get_batch_status re-fetches it. + output_uri = ( + batch_job.dest.gcs_uri + if batch_job.dest and batch_job.dest.gcs_uri + else None + ) + result: dict[str, Any] = { + "provider_status": state, + "provider_output_file_id": output_uri or batch_id, + } + if state in _FAILED_STATES: + message = batch_job.error.message if batch_job.error else state + result["error_message"] = message + logger.info( + f"[get_batch_status] Vertex batch status | batch_id={batch_id} | " + f"status={state}" + ) + return result + except Exception as e: + logger.error( + f"[get_batch_status] Failed to poll Vertex batch | " + f"batch_id={batch_id} | {e}" + ) + raise + + def download_batch_results(self, output_file_id: str) -> list[dict[str, Any]]: + """Read prediction JSONL files from the batch job's GCS output prefix. + + Vertex echoes the input ``key`` per line; line order is only the fallback. + """ + logger.info( + f"[download_batch_results] Reading Vertex results | src={output_file_id}" + ) + output_uri = output_file_id + if not output_uri.startswith("gs://"): + batch_job = self._client.batches.get(name=output_file_id) + state = batch_job.state.name if batch_job.state else "UNKNOWN" + if state != BatchJobState.SUCCEEDED.value: + raise ValueError(f"Batch job not complete. Current state: {state}") + if not (batch_job.dest and batch_job.dest.gcs_uri): + raise ValueError(f"Batch job has no GCS output | id={output_file_id}") + output_uri = batch_job.dest.gcs_uri + + try: + bucket_name, prefix = _parse_gs_uri(output_uri) + bucket = self._storage.bucket(bucket_name) + results: list[dict[str, Any]] = [] + index = 0 + for blob in self._storage.list_blobs(bucket, prefix=prefix): + if not blob.name.endswith(".jsonl"): + continue + content = blob.download_as_text() + for line in content.strip().split("\n"): + if not line: + continue + parsed = json.loads(line) + custom_id = parsed.get("key") or str(index) + response_obj = parsed.get("response") + error_obj = parsed.get("error") or parsed.get("status") + results.append( + { + BATCH_KEY: custom_id, + "response": response_obj, + "error": str(error_obj) if error_obj else None, + } + ) + index += 1 + logger.info( + f"[download_batch_results] Read Vertex results | src={output_uri} | " + f"results={len(results)}" + ) + return results + except Exception as e: + logger.error( + f"[download_batch_results] Failed to read Vertex results | " + f"src={output_uri} | {e}" + ) + raise + + def upload_file(self, content: str, purpose: str = "batch") -> str: + """Upload a JSONL string to GCS and return its ``gs://`` URI.""" + key = f"{self._input_prefix}/{int(time.time())}-{uuid4().hex}.jsonl" + logger.info(f"[upload_file] Uploading batch input to GCS | key={key}") + try: + blob = self._storage.bucket(self._bucket).blob(key) + blob.upload_from_string(content, content_type="application/jsonl") + return f"gs://{self._bucket}/{key}" + except Exception as e: + logger.error(f"[upload_file] Failed to upload batch input to GCS | {e}") + raise CloudStorageError(f"GCS upload failed: {e}") from e + + def download_file(self, file_id: str) -> str: + """Download a ``gs://`` object's content as text.""" + logger.info(f"[download_file] Downloading from GCS | uri={file_id}") + try: + bucket_name, key = _parse_gs_uri(file_id) + blob = self._storage.bucket(bucket_name).blob(key) + return blob.download_as_text() + except Exception as e: + logger.error( + f"[download_file] Failed to download from GCS | uri={file_id} | {e}" + ) + raise CloudStorageError(f"GCS download failed: {e}") from e diff --git a/backend/app/crud/assessment/batch.py b/backend/app/crud/assessment/batch.py index d77b239e5..a977e29d0 100644 --- a/backend/app/crud/assessment/batch.py +++ b/backend/app/crud/assessment/batch.py @@ -42,6 +42,7 @@ build_anthropic_attachment_parts, build_gemini_attachment_parts, resolve_attachment_values, + rewrite_gcs_attachment_urls, ) from app.services.llm.providers.registry import LLMProvider from app.utils import get_anthropic_client, get_openai_client @@ -419,6 +420,16 @@ def submit_assessment_batch( # Determine the base provider (openai or google) base_provider = provider_name.replace("-native", "") + # Resolve attachments url to provider-reachable URLs before building JSONL. + rows = rewrite_gcs_attachment_urls( + session=session, + rows=rows, + attachments=attachments, + llm_provider=provider_name, + project_id=project_id, + organization_id=organization_id, + ) + if base_provider == LLMProvider.OPENAI: mapped_params, warnings = map_kaapi_to_openai_params( session=session, diff --git a/backend/app/models/llm/constants.py b/backend/app/models/llm/constants.py index 7ac980b48..9a4cddd86 100644 --- a/backend/app/models/llm/constants.py +++ b/backend/app/models/llm/constants.py @@ -42,7 +42,12 @@ class Provider(StrEnum): Provider.GOOGLE_AISTUDIO, ] -TextProvider = Literal[Provider.OPENAI, Provider.GOOGLE, Provider.ANTHROPIC] +TextProvider = Literal[ + Provider.OPENAI, + Provider.GOOGLE, + Provider.GOOGLE_GCP, + Provider.ANTHROPIC, +] # Native provider names are the Kaapi providers with a "-native" suffix. # Kept as explicit strings since there's no corresponding enum member. diff --git a/backend/app/services/assessment/api/batch.py b/backend/app/services/assessment/api/batch.py index 993c84f92..afc8724c7 100644 --- a/backend/app/services/assessment/api/batch.py +++ b/backend/app/services/assessment/api/batch.py @@ -25,16 +25,16 @@ BATCH_KEY, AnthropicBatchProvider, BatchJobState, - GeminiBatchProvider, MessageBatchStatus, OpenAIBatchProvider, extract_text_from_response_dict, + get_gemini_batch_provider, + is_vertex_batch_provider, poll_batch_status, process_completed_batch, start_batch_job, ) from app.core.batch.base import BatchProvider -from app.core.batch.client import GeminiClient from app.core.config import settings from app.core.db import engine from app.crud.assessment import api @@ -43,6 +43,7 @@ build_google_jsonl, build_openai_jsonl, ) +from app.services.assessment.utils.attachments import rewrite_gcs_attachment_urls from app.crud.job import get_batch_job from app.models.assessment import ( Assessment, @@ -126,6 +127,7 @@ class StageKind(StrEnum): _SUPPORTED_PROVIDERS = { LLMProvider.OPENAI, LLMProvider.GOOGLE, + LLMProvider.GOOGLE_GCP, LLMProvider.ANTHROPIC, } @@ -291,6 +293,16 @@ def _submit_provider_batch( description: str, ) -> BatchJob: """Build provider JSONL and submit it via the shared batch infra.""" + # Resolve gs:// attachments to provider-reachable URLs before building JSONL. + rows = rewrite_gcs_attachment_urls( + session=session, + rows=rows, + attachments=attachments, + llm_provider=provider_name, + project_id=project_id, + organization_id=organization_id, + ) + if provider_name == LLMProvider.OPENAI: mapped, _ = map_kaapi_to_openai_params(session=session, kaapi_params=params) jsonl = build_openai_jsonl( @@ -306,16 +318,24 @@ def _submit_provider_batch( "description": description, "completion_window": "24h", } - elif provider_name == LLMProvider.GOOGLE: + elif provider_name in (LLMProvider.GOOGLE, LLMProvider.GOOGLE_GCP): mapped, _ = map_kaapi_to_google_params(params) jsonl = build_google_jsonl( rows, text_columns, attachments, prompt, mapped, row_indices ) - gemini = GeminiClient.from_credentials( - session=session, org_id=organization_id, project_id=project_id + provider = get_gemini_batch_provider( + session=session, + organization_id=organization_id, + project_id=project_id, + provider_name=provider_name, + model=model, + ) + # Vertex takes a bare model id; AI-Studio uses the "models/" prefix. + config = ( + {"display_name": description} + if is_vertex_batch_provider(provider_name) + else {"display_name": description, "model": f"models/{model}"} ) - provider = GeminiBatchProvider(client=gemini.client, model=f"models/{model}") - config = {"display_name": description, "model": f"models/{model}"} elif provider_name == LLMProvider.ANTHROPIC: mapped, _ = map_kaapi_to_anthropic_params(params) jsonl = build_anthropic_jsonl( @@ -362,11 +382,13 @@ def _build_batch_provider( session=session, org_id=organization_id, project_id=project_id ) ) - if provider_name == LLMProvider.GOOGLE: - gemini = GeminiClient.from_credentials( - session=session, org_id=organization_id, project_id=project_id + if provider_name in (LLMProvider.GOOGLE, LLMProvider.GOOGLE_GCP): + return get_gemini_batch_provider( + session=session, + organization_id=organization_id, + project_id=project_id, + provider_name=provider_name, ) - return GeminiBatchProvider(client=gemini.client) if provider_name == LLMProvider.ANTHROPIC: return AnthropicBatchProvider( client=get_anthropic_client( @@ -453,7 +475,7 @@ def _parse_one(result: dict[str, Any], provider_name: str) -> ParsedResult: "response_id": response.get("id"), } - if provider_name == LLMProvider.GOOGLE: + if provider_name in (LLMProvider.GOOGLE, LLMProvider.GOOGLE_GCP): response = result.get("response") text = extract_text_from_response_dict(response) if response else None return { diff --git a/backend/app/services/assessment/api/submission.py b/backend/app/services/assessment/api/submission.py index 1034a8681..3edc8774d 100644 --- a/backend/app/services/assessment/api/submission.py +++ b/backend/app/services/assessment/api/submission.py @@ -31,7 +31,8 @@ logger = logging.getLogger(__name__) # Attachment cell values are provided as URLs (base64 is unsupported for batch). -_URL_PREFIXES = ("http://", "https://") +# gs:// is allowed: it is resolved to a provider-reachable URL before batch build. +_URL_PREFIXES = ("http://", "https://", "gs://") _ATTACHMENT_TYPES = ("image", "pdf") diff --git a/backend/app/services/assessment/tasks.py b/backend/app/services/assessment/tasks.py index c4d6b45a5..852725e7b 100644 --- a/backend/app/services/assessment/tasks.py +++ b/backend/app/services/assessment/tasks.py @@ -28,6 +28,8 @@ ) from app.models.config.config import ConfigTag from app.services.assessment.prefilter import resolve_prefilter_settings +from app.services.assessment.prefilter.constants import ASSESSMENT_PREFILTER_PROVIDER +from app.services.assessment.utils.attachments import rewrite_gcs_attachment_urls from app.services.assessment.stages import ( GATE_STAGES, STAGE_PARSERS, @@ -242,6 +244,19 @@ def _submit_stage( selected = cfg.get("tr_attachment_columns") if selected is not None: attachments = [a for a in attachments if a.column in set(selected)] + if attachments: + resolved_rows = rewrite_gcs_attachment_urls( + session=session, + rows=[r for _, r in rows_with_idx], + attachments=attachments, + llm_provider=ASSESSMENT_PREFILTER_PROVIDER, + project_id=project_id, + organization_id=organization_id, + ) + rows_with_idx = [ + (idx, resolved_rows[pos]) + for pos, (idx, _) in enumerate(rows_with_idx) + ] jsonl = build_prefilter_requests(stage, rows_with_idx, cfg, attachments) batch_job = submit_prefilter_batch( session=session, diff --git a/backend/app/services/assessment/utils/attachments.py b/backend/app/services/assessment/utils/attachments.py index b130fa736..7e3f28dce 100644 --- a/backend/app/services/assessment/utils/attachments.py +++ b/backend/app/services/assessment/utils/attachments.py @@ -8,10 +8,14 @@ import logging import re -from typing import Any +from typing import Any, cast from urllib.parse import urlparse +from sqlmodel import Session + from app.models.assessment import AssessmentAttachment +from app.models.llm.constants import KaapiProvider +from app.services.buckets.attachments import is_gcs_uri, resolve_attachments logger = logging.getLogger(__name__) @@ -34,6 +38,54 @@ def split_attachment_urls(value: str) -> list[str]: return [part.strip() for part in re.split(r"[\n,]+", value) if part.strip()] +def rewrite_gcs_attachment_urls( + *, + session: Session, + rows: list[dict[str, str]], + attachments: list[AssessmentAttachment], + llm_provider: str, + project_id: int, + organization_id: int, +) -> list[dict[str, str]]: + """Replace gs:// tokens in attachment cells with LLM-reachable URLs. + + Bulk-resolves every gs:// URI across the batch once, then rewrites each cell; + non-gs values are left as-is. + """ + gcs_uris = { + attachment_url + for row in rows + for att in attachments + for attachment_url in split_attachment_urls(row.get(att.column, "")) + if is_gcs_uri(attachment_url) + } + if not gcs_uris: + return rows + + resolved = resolve_attachments( + session=session, + source=list(gcs_uris), + llm_provider=cast(KaapiProvider, llm_provider), + project_id=project_id, + organization_id=organization_id, + ) + assert isinstance(resolved, dict) # list input always yields a dict + + rewritten: list[dict[str, str]] = [] + for row in rows: + new_row = dict(row) + for att in attachments: + value = row.get(att.column) + if not value: + continue + new_row[att.column] = ", ".join( + resolved.get(attachment_url, attachment_url) + for attachment_url in split_attachment_urls(value) + ) + rewritten.append(new_row) + return rewritten + + def to_direct_attachment_url(url: str, attachment_type: str) -> str: """Normalize share-page attachment URLs into provider-fetchable direct URLs. diff --git a/backend/app/services/buckets/__init__.py b/backend/app/services/buckets/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/backend/app/services/buckets/attachments.py b/backend/app/services/buckets/attachments.py new file mode 100644 index 000000000..af28ce903 --- /dev/null +++ b/backend/app/services/buckets/attachments.py @@ -0,0 +1,83 @@ +"""Attachment utilities for LLM providers: gs:// path selection, URL resolution, etc.""" + +from enum import Enum + +from sqlmodel import Session + +from app.models.llm.constants import KaapiProvider, Provider +from app.services.buckets.providers.registry import get_bucket_provider + +GCS_URI_SCHEME = "gs" +DEFAULT_BUCKET_PROVIDER = "gcs" + +# LLM providers that read gs:// URIs natively (Vertex/google-gcp, incl. native key). +NATIVE_PROVIDERS: frozenset[str] = frozenset( + {Provider.GOOGLE_GCP, f"{Provider.GOOGLE_GCP}-native"} +) + + +class BucketPathStrategyEnum(str, Enum): + NATIVE = "native" # Path A: pass the gs:// URI straight to the provider. + SIGNED_URL = "signed_url" # Path B: convert to a signed HTTPS URL. + + +def is_gcs_uri(uri: str) -> bool: + return uri.startswith(f"{GCS_URI_SCHEME}://") + + +def resolve_bucket_path_strategy( + *, + llm_provider: KaapiProvider, + source_uri: str, + credential: dict | None = None, +) -> BucketPathStrategyEnum: + """Native only for a gs:// source the provider reads directly; else signed.""" + if is_gcs_uri(source_uri) and llm_provider in NATIVE_PROVIDERS: + return BucketPathStrategyEnum.NATIVE + return BucketPathStrategyEnum.SIGNED_URL + + +def resolve_attachments( + *, + session: Session, + source: str | list[str], + llm_provider: KaapiProvider, + project_id: int, + organization_id: int, + expires_in: int = 3600, + bucket_provider_type: str = DEFAULT_BUCKET_PROVIDER, +) -> str | dict[str, str]: + """Make attachment(s) LLM-reachable. Accepts a single URL or a list. + + Per URI: non-``gs://`` (http(s)/direct) is returned by reference; ``gs://`` on + a native provider passes through (Path A); otherwise it is signed (Path B). + Path-B URIs are bulk-signed with a single bucket-provider client per call. + Returns a ``str`` for a single input, a ``uri -> url`` dict for a list. + """ + is_single = isinstance(source, str) + uris = [source] if is_single else list(source) + + resolved: dict[str, str] = {} + to_sign: list[str] = [] + for uri in uris: + if not is_gcs_uri(uri): + resolved[uri] = uri + continue + strategy = resolve_bucket_path_strategy( + llm_provider=llm_provider, source_uri=uri + ) + if strategy is BucketPathStrategyEnum.NATIVE: + resolved[uri] = uri + else: + to_sign.append(uri) + + if to_sign: + provider = get_bucket_provider( + session=session, + provider_type=bucket_provider_type, + project_id=project_id, + organization_id=organization_id, + ) + resolved.update(provider.get_bulk_signed_urls(to_sign, expires_in=expires_in)) + + return resolved[uris[0]] if is_single else resolved diff --git a/backend/app/services/buckets/providers/__init__.py b/backend/app/services/buckets/providers/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/backend/app/services/buckets/providers/base.py b/backend/app/services/buckets/providers/base.py new file mode 100644 index 000000000..1951de7ae --- /dev/null +++ b/backend/app/services/buckets/providers/base.py @@ -0,0 +1,48 @@ +"""Base provider interface for object-storage bucket providers.""" + +import logging +from abc import ABC, abstractmethod +from typing import Any + +logger = logging.getLogger(__name__) + + +class BaseBucketProvider(ABC): + """Abstract base class for bucket providers.""" + + # URI scheme this provider handles (e.g. "gs", "s3"). + SCHEME: str = "" + + # Cap on signed-URL lifetime (24h), matching AmazonCloudStorage. + MAX_SIGNED_URL_EXPIRY: int = 86400 + + def __init__(self, client: Any): + self.client = client + + @staticmethod + @abstractmethod + def create_client(credentials: dict[str, Any]) -> Any: + """Instantiate a storage client from decrypted credentials.""" + raise NotImplementedError("Bucket providers must implement create_client") + + @abstractmethod + def get_signed_url(self, uri: str, expires_in: int = 3600) -> str: + """Generate a time-limited signed URL for a single object.""" + raise NotImplementedError("Bucket providers must implement get_signed_url") + + def get_bulk_signed_urls( + self, uris: list[str], expires_in: int = 3600 + ) -> dict[str, str]: + """Sign many object URIs, reusing this provider's single client.""" + signed: dict[str, str] = {} + for uri in uris: + signed[uri] = self.get_signed_url(uri, expires_in=expires_in) + return signed + + def to_public_url(self, uri: str, expires_in: int = 3600) -> str: + """Resolve a private object URI to a fetchable signed URL.""" + return self.get_signed_url(uri, expires_in=expires_in) + + def get_provider_name(self) -> str: + """Return the provider name derived from the class name.""" + return self.__class__.__name__.replace("BucketProvider", "").lower() diff --git a/backend/app/services/buckets/providers/gcs.py b/backend/app/services/buckets/providers/gcs.py new file mode 100644 index 000000000..62967953d --- /dev/null +++ b/backend/app/services/buckets/providers/gcs.py @@ -0,0 +1,112 @@ +"""GCS bucket provider: signed-URL generation over Google Cloud Storage.""" + +import logging +from typing import Any +from datetime import timedelta +from urllib.parse import urlparse + +from google.cloud import storage as gcs +from google.oauth2 import service_account + +from app.core.config import settings +from app.core.cloud.storage import GCS_SCOPES, CloudStorageError +from app.services.buckets.providers.base import BaseBucketProvider +from app.services.llm.providers.google_gcp import _load_platform_sa_info + +logger = logging.getLogger(__name__) + +GCS_URI_SCHEME = "gs" + + +class GCSClient: + # default_bucket carried for parity with the credential row; signing derives + # the bucket from each URI. + def __init__(self, storage_client: gcs.Client, default_bucket: str | None): + self.storage_client = storage_client + self.default_bucket = default_bucket + + +class GCSBucketProvider(BaseBucketProvider): + """Bucket provider for Google Cloud Storage (``gs://`` scheme).""" + + SCHEME = GCS_URI_SCHEME + + def __init__(self, client: GCSClient): + super().__init__(client) + self.client = client + + @staticmethod + def create_client(credentials: dict[str, Any]) -> GCSClient: + """Build a signing-capable GCS client with BYOK-over-settings precedence.""" + credentials = credentials or {} + gcs_bucket = credentials.get("gcs_bucket") or settings.GCS_AUDIO_BUCKET + sa_info = credentials.get("sa_key") or _load_platform_sa_info() + + source = "byok" if credentials.get("sa_key") else "platform" + logger.info( + f"[GCSBucketProvider.create_client] gcs creds | source={source}, " + f"bucket={gcs_bucket}" + ) + + # Signing needs the SA private key, so a missing sa_info is fatal. + if not sa_info: + raise ValueError( + "GCS bucket provider requires a service-account key (sa_key) to " + "sign URLs; none configured for this project or platform default." + ) + + creds = service_account.Credentials.from_service_account_info( + sa_info, scopes=list(GCS_SCOPES) + ) + storage_client = gcs.Client( + project=sa_info.get("project_id"), credentials=creds + ) + return GCSClient(storage_client=storage_client, default_bucket=gcs_bucket) + + @staticmethod + def _parse_gcs_uri(uri: str) -> tuple[str, str]: + """Split ``gs://bucket/key`` into ``(bucket, key)``.""" + parsed = urlparse(uri) + if parsed.scheme != GCS_URI_SCHEME or not parsed.netloc: + raise ValueError( + f"Invalid GCS URI '{uri}'; expected 'gs://bucket/key'." + ) + key = parsed.path.lstrip("/") + if not key: + raise ValueError(f"GCS URI '{uri}' is missing an object key.") + return parsed.netloc, key + + def get_signed_url(self, uri: str, expires_in: int = 3600) -> str: + expires_in = min(expires_in, self.MAX_SIGNED_URL_EXPIRY) + bucket_name, key = self._parse_gcs_uri(uri) + + try: + blob = self.client.storage_client.bucket(bucket_name).blob(key) + signed_url = blob.generate_signed_url( + version="v4", + expiration=timedelta(seconds=expires_in), + method="GET", + ) + except Exception as e: + logger.error( + f"[GCSBucketProvider.get_signed_url] GCS signing failed | " + f"bucket={bucket_name}, key={key}, error={e}", + exc_info=True, + ) + raise CloudStorageError(f"GCS signing failed: {e} ({uri})") from e + + logger.info( + f"[GCSBucketProvider.get_signed_url] Signed URL generated | " + f"bucket={bucket_name}, key={key}, expires_in={expires_in}" + ) + return signed_url + + def get_bulk_signed_urls( + self, uris: list[str], expires_in: int = 3600 + ) -> dict[str, str]: + """Sign each URI reusing this provider's single client.""" + logger.info( + f"[GCSBucketProvider.get_bulk_signed_urls] Signing batch | " + f"count={len(uris)}, expires_in={expires_in}" + ) + return {uri: self.get_signed_url(uri, expires_in=expires_in) for uri in uris} diff --git a/backend/app/services/buckets/providers/registry.py b/backend/app/services/buckets/providers/registry.py new file mode 100644 index 000000000..b9c2e7c09 --- /dev/null +++ b/backend/app/services/buckets/providers/registry.py @@ -0,0 +1,92 @@ +"""Global bucket-provider registry and resolver.""" + +import logging + +from sqlmodel import Session + +from app.services.buckets.providers.base import BaseBucketProvider +from app.services.buckets.providers.gcs import GCSBucketProvider + +logger = logging.getLogger(__name__) + + +class BucketProvider: + GCS = "gcs" + + _registry: dict[str, type[BaseBucketProvider]] = { + GCS: GCSBucketProvider, + } + + # Bucket-provider key -> credential-provider key. GCS reuses google-gcp. + _credential_provider: dict[str, str] = { + GCS: "google-gcp", + } + + @classmethod + def get_provider_class(cls, provider_type: str) -> type[BaseBucketProvider]: + """Return the bucket-provider class for a given name.""" + provider = cls._registry.get(provider_type) + if not provider: + raise ValueError( + f"Bucket provider '{provider_type}' is not supported. " + f"Supported providers: {', '.join(cls._registry.keys())}" + ) + return provider + + @classmethod + def supported_providers(cls) -> list[str]: + """Return a list of supported bucket-provider names.""" + return list(cls._registry.keys()) + + @classmethod + def get_credential_provider(cls, provider_type: str) -> str: + """Return the credential-provider key backing a bucket provider.""" + credential_provider = cls._credential_provider.get(provider_type) + if not credential_provider: + raise ValueError( + f"Bucket provider '{provider_type}' has no credential mapping." + ) + return credential_provider + + +def get_bucket_provider( + session: Session, provider_type: str, project_id: int, organization_id: int +) -> BaseBucketProvider: + # Lazy import to avoid the crud <-> provider-registry cycle. + from app.crud.credentials import get_provider_credential + + provider_class = BucketProvider.get_provider_class(provider_type) + credential_provider = BucketProvider.get_credential_provider(provider_type) + + credentials = get_provider_credential( + session=session, + provider=credential_provider, + project_id=project_id, + org_id=organization_id, + ) + + if not credentials: + raise ValueError( + f"Credentials for provider '{credential_provider}' not configured " + f"for this project." + ) + + # Default fetch yields the decrypted dict create_client needs. + if not isinstance(credentials, dict): + raise ValueError( + f"Expected decrypted credentials dict for provider " + f"'{credential_provider}', got {type(credentials).__name__}." + ) + + try: + client = provider_class.create_client(credentials=credentials) + return provider_class(client=client) + except ValueError: + # Credential/config errors are the caller's to fix; surface as-is. + raise + except Exception as e: + logger.error( + f"[get_bucket_provider] Failed to initialize {provider_type} client: {e}", + exc_info=True, + ) + raise RuntimeError(f"Could not connect to {provider_type} bucket services.") diff --git a/backend/app/tests/assessment/test_api_batch.py b/backend/app/tests/assessment/test_api_batch.py index 592da00ab..8fcbc38d9 100644 --- a/backend/app/tests/assessment/test_api_batch.py +++ b/backend/app/tests/assessment/test_api_batch.py @@ -388,6 +388,14 @@ def test_google_empty_response(self) -> None: assert result["output"] is None assert result["error"] == "Empty response" + def test_google_gcp_parses_like_google(self) -> None: + with patch( + "app.services.assessment.api.batch.extract_text_from_response_dict", + return_value="vertex text", + ): + result = _parse_one({"response": {"candidates": []}}, "google-gcp") + assert result["output"] == "vertex text" + def test_unknown_provider(self) -> None: result = _parse_one({"response": {}}, "cohere") assert result["output"] is None @@ -860,10 +868,10 @@ def test_openai(self, db) -> None: def test_google(self, db) -> None: auth = get_user_test_auth_context(db) - gemini = MagicMock() - gemini.client = MagicMock() - with patch("app.services.assessment.api.batch.GeminiClient") as gemini_cls: - gemini_cls.from_credentials.return_value = gemini + with patch( + "app.services.assessment.api.batch.get_gemini_batch_provider", + return_value=MagicMock(), + ) as get_provider: provider = _build_batch_provider( session=db, provider_name="google", @@ -871,6 +879,21 @@ def test_google(self, db) -> None: project_id=auth.project_id, ) assert provider is not None + assert get_provider.call_args.kwargs["provider_name"] in ("google", "google-native") + + def test_google_gcp_routes_to_vertex(self, db) -> None: + auth = get_user_test_auth_context(db) + with patch( + "app.services.assessment.api.batch.get_gemini_batch_provider", + return_value=MagicMock(), + ) as get_provider: + _build_batch_provider( + session=db, + provider_name="google-gcp", + organization_id=auth.organization_id, + project_id=auth.project_id, + ) + assert get_provider.call_args.kwargs["provider_name"] == "google-gcp" def test_anthropic(self, db) -> None: auth = get_user_test_auth_context(db) @@ -915,24 +938,48 @@ def _kwargs(self, db, auth, provider_name, params): def test_google_branch(self, db) -> None: auth = get_user_test_auth_context(db) - gemini = MagicMock() - gemini.client = MagicMock() job = _make_batch_job( db, org_id=auth.organization_id, project_id=auth.project_id ) with ( - patch("app.services.assessment.api.batch.GeminiClient") as gemini_cls, + patch( + "app.services.assessment.api.batch.get_gemini_batch_provider", + return_value=MagicMock(), + ) as get_provider, patch( "app.services.assessment.api.batch.start_batch_job", return_value=job, ) as start, ): - gemini_cls.from_credentials.return_value = gemini result = _submit_provider_batch( **self._kwargs(db, auth, "google", {"model": "gemini-2.5-pro"}) ) assert result.id == job.id assert start.call_args.kwargs["provider_name"] == "google" + assert get_provider.call_args.kwargs["provider_name"] in ("google", "google-native") + + def test_google_gcp_branch_routes_to_vertex(self, db) -> None: + auth = get_user_test_auth_context(db) + job = _make_batch_job( + db, org_id=auth.organization_id, project_id=auth.project_id + ) + with ( + patch( + "app.services.assessment.api.batch.get_gemini_batch_provider", + return_value=MagicMock(), + ) as get_provider, + patch( + "app.services.assessment.api.batch.start_batch_job", + return_value=job, + ) as start, + ): + _submit_provider_batch( + **self._kwargs(db, auth, "google-gcp", {"model": "gemini-2.5-pro"}) + ) + assert start.call_args.kwargs["provider_name"] == "google-gcp" + assert get_provider.call_args.kwargs["provider_name"] == "google-gcp" + # Vertex config omits the "models/" prefixed model (uses bare id). + assert "model" not in start.call_args.kwargs["config"] def test_anthropic_branch_sets_max_tokens(self, db) -> None: auth = get_user_test_auth_context(db) diff --git a/backend/app/tests/assessment/test_api_submission.py b/backend/app/tests/assessment/test_api_submission.py index 0f2198f6f..d82f6fc17 100644 --- a/backend/app/tests/assessment/test_api_submission.py +++ b/backend/app/tests/assessment/test_api_submission.py @@ -159,6 +159,13 @@ def test_row_attachment_value_not_url_is_422(self, db) -> None: assert "input.data[0]" in exc.value.detail assert "must be a URL" in exc.value.detail + def test_gs_uri_attachment_value_passes_validation(self) -> None: + # gs:// is allowed at submit-time; it is resolved before batch build. + submission._validate_rows_against_schema( + [{"img": "gs://bucket/key.png"}], + {"img": {"type": "image", "format": "url"}}, + ) + def test_unsupported_provider_is_422(self, db) -> None: auth = get_user_test_auth_context(db) # Build a config whose stored blob names an unsupported batch provider. diff --git a/backend/app/tests/assessment/test_batch.py b/backend/app/tests/assessment/test_batch.py index 16a9e0f1a..9e66c7627 100644 --- a/backend/app/tests/assessment/test_batch.py +++ b/backend/app/tests/assessment/test_batch.py @@ -26,10 +26,53 @@ build_gemini_attachment_parts, resolve_attachment_values, resolve_item_type, + rewrite_gcs_attachment_urls, split_attachment_urls, to_direct_attachment_url, ) +_REWRITE_RESOLVER = "app.services.assessment.utils.attachments.resolve_attachments" + + +class TestRewriteGcsAttachmentUrls: + def test_rewrites_gcs_leaves_https_untouched(self): + att = AssessmentAttachment(column="img", type="image", format="url") + rows = [{"img": "gs://b/1.png, https://x/2.png"}, {"img": "gs://b/3.png"}] + with patch( + _REWRITE_RESOLVER, + return_value={"gs://b/1.png": "https://s1", "gs://b/3.png": "https://s3"}, + ) as mock_resolve: + out = rewrite_gcs_attachment_urls( + session=MagicMock(), + rows=rows, + attachments=[att], + llm_provider="openai", + project_id=1, + organization_id=2, + ) + assert out[0]["img"] == "https://s1, https://x/2.png" + assert out[1]["img"] == "https://s3" + # single bulk resolve for both gs:// URIs across all rows + mock_resolve.assert_called_once() + _, kwargs = mock_resolve.call_args + assert set(kwargs["source"]) == {"gs://b/1.png", "gs://b/3.png"} + assert kwargs["llm_provider"] == "openai" + + def test_no_gcs_returns_rows_unchanged_without_resolving(self): + att = AssessmentAttachment(column="img", type="image", format="url") + rows = [{"img": "https://x/1.png"}] + with patch(_REWRITE_RESOLVER) as mock_resolve: + out = rewrite_gcs_attachment_urls( + session=MagicMock(), + rows=rows, + attachments=[att], + llm_provider="anthropic", + project_id=1, + organization_id=2, + ) + assert out is rows + mock_resolve.assert_not_called() + def _make_run() -> MagicMock: run = MagicMock() diff --git a/backend/app/tests/core/batch/test_client.py b/backend/app/tests/core/batch/test_client.py new file mode 100644 index 000000000..dd44453b8 --- /dev/null +++ b/backend/app/tests/core/batch/test_client.py @@ -0,0 +1,76 @@ +"""Test cases for the Gemini-family batch provider resolver.""" + +from unittest.mock import MagicMock, patch + +import pytest +from fastapi import HTTPException + +from app.core.batch.client import ( + get_gemini_batch_provider, + is_vertex_batch_provider, +) +from app.core.batch.gemini import GeminiBatchProvider + + +class TestIsVertexBatchProvider: + @pytest.mark.parametrize("name", ["google-gcp", "google-gcp-native"]) + def test_vertex_providers(self, name): + assert is_vertex_batch_provider(name) is True + + @pytest.mark.parametrize("name", ["google", "google-aistudio", "openai"]) + def test_non_vertex_providers(self, name): + assert is_vertex_batch_provider(name) is False + + +class TestGetGeminiBatchProvider: + def test_google_routes_to_aistudio(self): + gemini = MagicMock() + gemini.client = MagicMock() + with patch( + "app.core.batch.client.GeminiClient.from_credentials", + return_value=gemini, + ) as from_cred: + provider = get_gemini_batch_provider( + session=MagicMock(), + organization_id=1, + project_id=2, + provider_name="google", + ) + assert isinstance(provider, GeminiBatchProvider) + from_cred.assert_called_once() + + def test_google_gcp_routes_to_vertex(self): + sentinel = MagicMock() + with ( + patch( + "app.core.batch.client.get_provider_credential", + return_value={"gcs_bucket": "b", "sa_key": {}}, + ) as get_cred, + patch( + "app.core.batch.vertex.VertexBatchProvider.from_credentials", + return_value=sentinel, + ) as from_cred, + ): + provider = get_gemini_batch_provider( + session=MagicMock(), + organization_id=1, + project_id=2, + provider_name="google-gcp", + model="gemini-2.5-pro", + ) + assert provider is sentinel + assert get_cred.call_args.kwargs["provider"] == "google-gcp" + assert from_cred.call_args.kwargs["model"] == "gemini-2.5-pro" + + def test_missing_gcp_credential_raises_404(self): + with patch( + "app.core.batch.client.get_provider_credential", return_value=None + ): + with pytest.raises(HTTPException) as exc: + get_gemini_batch_provider( + session=MagicMock(), + organization_id=1, + project_id=2, + provider_name="google-gcp", + ) + assert exc.value.status_code == 404 diff --git a/backend/app/tests/core/batch/test_vertex.py b/backend/app/tests/core/batch/test_vertex.py new file mode 100644 index 000000000..5c14af33d --- /dev/null +++ b/backend/app/tests/core/batch/test_vertex.py @@ -0,0 +1,185 @@ +"""Test cases for VertexBatchProvider (Vertex AI batch, GCS-backed).""" + +import json +from unittest.mock import MagicMock, patch + +import pytest + +from app.core.batch.vertex import VertexBatchProvider, _parse_gs_uri + +_BUCKET = "test-bucket" + + +@pytest.fixture +def mock_genai(): + return MagicMock() + + +@pytest.fixture +def mock_storage(): + return MagicMock() + + +@pytest.fixture +def provider(mock_genai, mock_storage): + return VertexBatchProvider( + client=mock_genai, + storage_client=mock_storage, + gcs_bucket=_BUCKET, + model="gemini-2.5-pro", + ) + + +class TestParseGsUri: + def test_parses_bucket_and_key(self): + assert _parse_gs_uri("gs://b/dir/f.jsonl") == ("b", "dir/f.jsonl") + + def test_rejects_non_gs(self): + with pytest.raises(ValueError, match="gs://"): + _parse_gs_uri("https://b/f.jsonl") + + +class TestCreateBatch: + def test_uploads_to_gcs_and_starts_job(self, provider, mock_genai, mock_storage): + job = MagicMock() + job.name = "batches/123" + job.state.name = "JOB_STATE_PENDING" + mock_genai.batches.create.return_value = job + + result = provider.create_batch( + [{"request": {"contents": []}}], {"display_name": "run-1"} + ) + + # Input JSONL uploaded to gs://bucket/batch-input/... + blob = mock_storage.bucket.return_value.blob.return_value + blob.upload_from_string.assert_called_once() + # batches.create called with a gs:// src and a gs:// dest in config. + kwargs = mock_genai.batches.create.call_args.kwargs + assert kwargs["src"].startswith(f"gs://{_BUCKET}/batch-input/") + assert kwargs["config"].dest.startswith(f"gs://{_BUCKET}/batch-output/") + assert result["provider_batch_id"] == "batches/123" + assert result["provider_status"] == "JOB_STATE_PENDING" + assert result["total_items"] == 1 + + +class TestGetBatchStatus: + def test_succeeded_returns_output_uri(self, provider, mock_genai): + job = MagicMock() + job.state.name = "JOB_STATE_SUCCEEDED" + job.dest.gcs_uri = "gs://b/out/" + mock_genai.batches.get.return_value = job + + result = provider.get_batch_status("batches/123") + assert result["provider_status"] == "JOB_STATE_SUCCEEDED" + assert result["provider_output_file_id"] == "gs://b/out/" + assert "error_message" not in result + + def test_failed_sets_error_message(self, provider, mock_genai): + job = MagicMock() + job.state.name = "JOB_STATE_FAILED" + job.dest = None + job.error.message = "boom" + mock_genai.batches.get.return_value = job + + result = provider.get_batch_status("batches/123") + assert result["provider_status"] == "JOB_STATE_FAILED" + assert result["error_message"] == "boom" + assert result["provider_output_file_id"] == "batches/123" + + +class TestDownloadBatchResults: + def _blob(self, name, text): + blob = MagicMock() + blob.name = name + blob.download_as_text.return_value = text + return blob + + def test_reads_jsonl_and_preserves_echoed_key(self, provider, mock_storage): + lines = "\n".join( + json.dumps({"key": k, "response": {"text": t}}) + for k, t in (("row_3", "a"), ("row_7", "b")) + ) + mock_storage.list_blobs.return_value = [ + self._blob("out/predictions.jsonl", lines), + self._blob("out/_SUCCESS", "ignore-me"), # non-jsonl skipped + ] + + results = provider.download_batch_results("gs://b/out/") + assert [r["custom_id"] for r in results] == ["row_3", "row_7"] + assert results[0]["response"] == {"text": "a"} + assert results[0]["error"] is None + + def test_falls_back_to_line_order_without_key(self, provider, mock_storage): + lines = "\n".join(json.dumps({"response": {"text": t}}) for t in ("a", "b")) + mock_storage.list_blobs.return_value = [ + self._blob("out/predictions.jsonl", lines) + ] + results = provider.download_batch_results("gs://b/out/") + assert [r["custom_id"] for r in results] == ["0", "1"] + + def test_batch_name_resolves_dest_then_reads(self, provider, mock_genai, mock_storage): + job = MagicMock() + job.state.name = "JOB_STATE_SUCCEEDED" + job.dest.gcs_uri = "gs://b/out/" + mock_genai.batches.get.return_value = job + mock_storage.list_blobs.return_value = [ + self._blob("out/p.jsonl", json.dumps({"key": "row_5", "error": "quota"})) + ] + + results = provider.download_batch_results("batches/123") + assert results[0]["custom_id"] == "row_5" + assert results[0]["error"] == "quota" + assert results[0]["response"] is None + + def test_incomplete_job_raises(self, provider, mock_genai): + job = MagicMock() + job.state.name = "JOB_STATE_RUNNING" + mock_genai.batches.get.return_value = job + with pytest.raises(ValueError, match="not complete"): + provider.download_batch_results("batches/123") + + +class TestFileIO: + def test_upload_file_returns_gs_uri(self, provider, mock_storage): + uri = provider.upload_file('{"x":1}') + assert uri.startswith(f"gs://{_BUCKET}/batch-input/") + blob = mock_storage.bucket.return_value.blob.return_value + blob.upload_from_string.assert_called_once() + assert blob.upload_from_string.call_args.kwargs["content_type"] == ( + "application/jsonl" + ) + + def test_download_file_reads_text(self, provider, mock_storage): + mock_storage.bucket.return_value.blob.return_value.download_as_text.return_value = ( + "hello" + ) + assert provider.download_file("gs://b/k.jsonl") == "hello" + + +class TestFromCredentials: + _CRED = { + "project_id": "proj", + "location": "us-central1", + "gcs_bucket": _BUCKET, + "sa_key": {"type": "service_account", "project_id": "proj"}, + } + + def test_builds_provider(self): + with ( + patch("app.core.batch.vertex.service_account") as sa, + patch("app.core.batch.vertex.genai.Client") as genai_client, + patch("app.core.batch.vertex.gcs.Client") as gcs_client, + ): + provider = VertexBatchProvider.from_credentials(self._CRED) + assert isinstance(provider, VertexBatchProvider) + sa.Credentials.from_service_account_info.assert_called_once() + assert genai_client.call_args.kwargs["vertexai"] is True + gcs_client.assert_called_once() + + @pytest.mark.parametrize( + "drop", ["project_id", "location", "gcs_bucket", "sa_key"] + ) + def test_missing_field_raises(self, drop): + cred = {k: v for k, v in self._CRED.items() if k != drop} + with pytest.raises(ValueError, match=drop): + VertexBatchProvider.from_credentials(cred) diff --git a/backend/app/tests/services/buckets/__init__.py b/backend/app/tests/services/buckets/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/backend/app/tests/services/buckets/test_attachments.py b/backend/app/tests/services/buckets/test_attachments.py new file mode 100644 index 000000000..aa640df4f --- /dev/null +++ b/backend/app/tests/services/buckets/test_attachments.py @@ -0,0 +1,126 @@ +"""Tests for the bucket attachment utilities (path strategy + URL resolution).""" + +from unittest.mock import MagicMock, patch + +import pytest + +from app.services.buckets.attachments import ( + BucketPathStrategyEnum, + is_gcs_uri, + resolve_attachments, + resolve_bucket_path_strategy, +) + +_RESOLVER_PATH = "app.services.buckets.attachments.get_bucket_provider" + + +class TestResolveBucketPathStrategy: + @pytest.mark.parametrize("llm_provider", ["google-gcp", "google-gcp-native"]) + def test_gcs_uri_native_provider_is_native(self, llm_provider): + strategy = resolve_bucket_path_strategy( + llm_provider=llm_provider, + source_uri="gs://bucket/key.wav", + ) + assert strategy is BucketPathStrategyEnum.NATIVE + + @pytest.mark.parametrize("llm_provider", ["openai", "anthropic", "google-aistudio"]) + def test_gcs_uri_non_native_provider_is_signed_url(self, llm_provider): + strategy = resolve_bucket_path_strategy( + llm_provider=llm_provider, + source_uri="gs://bucket/key.wav", + ) + assert strategy is BucketPathStrategyEnum.SIGNED_URL + + def test_non_gcs_uri_is_signed_url_even_for_native_provider(self): + strategy = resolve_bucket_path_strategy( + llm_provider="google-gcp", + source_uri="https://example.com/key.wav", + ) + assert strategy is BucketPathStrategyEnum.SIGNED_URL + + +class TestIsGcsUri: + def test_gcs_uri(self): + assert is_gcs_uri("gs://bucket/key.wav") is True + + def test_non_gcs_uri(self): + assert is_gcs_uri("https://example.com/key.wav") is False + + +class TestResolveAttachmentsSingle: + def test_https_returned_as_is_without_provider(self): + with patch(_RESOLVER_PATH) as mock_get: + url = resolve_attachments( + session=MagicMock(), + source="https://example.com/img.png", + llm_provider="openai", + project_id=1, + organization_id=2, + ) + assert url == "https://example.com/img.png" + mock_get.assert_not_called() + + def test_gcs_native_passthrough(self): + with patch(_RESOLVER_PATH) as mock_get: + url = resolve_attachments( + session=MagicMock(), + source="gs://bucket/key.png", + llm_provider="google-gcp", + project_id=1, + organization_id=2, + ) + assert url == "gs://bucket/key.png" + mock_get.assert_not_called() + + def test_gcs_signed_for_non_native_provider(self): + provider = MagicMock() + provider.get_bulk_signed_urls.return_value = {"gs://bucket/key.png": "https://signed"} + with patch(_RESOLVER_PATH, return_value=provider): + url = resolve_attachments( + session=MagicMock(), + source="gs://bucket/key.png", + llm_provider="openai", + project_id=1, + organization_id=2, + expires_in=1200, + ) + assert url == "https://signed" + provider.get_bulk_signed_urls.assert_called_once_with( + ["gs://bucket/key.png"], expires_in=1200 + ) + + +class TestResolveAttachmentsList: + def test_mixed_schemes_partitioned(self): + provider = MagicMock() + provider.get_bulk_signed_urls.return_value = {"gs://b/2.png": "https://signed2"} + with patch(_RESOLVER_PATH, return_value=provider): + result = resolve_attachments( + session=MagicMock(), + source=["https://x/1.png", "gs://b/2.png"], + llm_provider="anthropic", + project_id=1, + organization_id=2, + ) + assert result == { + "https://x/1.png": "https://x/1.png", + "gs://b/2.png": "https://signed2", + } + provider.get_bulk_signed_urls.assert_called_once_with( + ["gs://b/2.png"], expires_in=3600 + ) + + def test_all_native_skips_provider(self): + with patch(_RESOLVER_PATH) as mock_get: + result = resolve_attachments( + session=MagicMock(), + source=["gs://b/1.wav", "gs://b/2.wav"], + llm_provider="google-gcp", + project_id=1, + organization_id=2, + ) + assert result == { + "gs://b/1.wav": "gs://b/1.wav", + "gs://b/2.wav": "gs://b/2.wav", + } + mock_get.assert_not_called() diff --git a/backend/app/tests/services/buckets/test_gcs.py b/backend/app/tests/services/buckets/test_gcs.py new file mode 100644 index 000000000..a4ca51d0c --- /dev/null +++ b/backend/app/tests/services/buckets/test_gcs.py @@ -0,0 +1,156 @@ +"""Tests for the GCS bucket provider.""" + +from datetime import timedelta +from unittest.mock import MagicMock, patch + +import pytest + +from app.core.cloud.storage import GCS_SCOPES +from app.services.buckets.providers.gcs import ( + GCSBucketProvider, + GCSClient, +) + + +def _make_provider() -> tuple[GCSBucketProvider, MagicMock]: + """Build a provider over a mock storage client; return the blob mock too.""" + blob = MagicMock() + storage_client = MagicMock() + storage_client.bucket.return_value.blob.return_value = blob + provider = GCSBucketProvider( + client=GCSClient(storage_client=storage_client, default_bucket="b") + ) + return provider, blob + + +class TestCreateClient: + def test_byok_credentials_win_over_settings(self, monkeypatch): + monkeypatch.setattr( + "app.services.buckets.providers.gcs.settings.GCS_AUDIO_BUCKET", + "platform-bucket", + ) + byok_sa = {"project_id": "byok-project"} + + with ( + patch( + "app.services.buckets.providers.gcs._load_platform_sa_info", + return_value={"project_id": "platform-project"}, + ), + patch( + "app.services.buckets.providers.gcs.service_account." + "Credentials.from_service_account_info" + ) as mock_from_info, + patch( + "app.services.buckets.providers.gcs.gcs.Client" + ) as mock_client, + ): + client = GCSBucketProvider.create_client( + {"gcs_bucket": "byok-bucket", "sa_key": byok_sa} + ) + + assert client.default_bucket == "byok-bucket" + mock_from_info.assert_called_once_with(byok_sa, scopes=list(GCS_SCOPES)) + mock_client.assert_called_once_with( + project="byok-project", credentials=mock_from_info.return_value + ) + + def test_falls_back_to_settings_and_platform_sa(self, monkeypatch): + monkeypatch.setattr( + "app.services.buckets.providers.gcs.settings.GCS_AUDIO_BUCKET", + "platform-bucket", + ) + + with ( + patch( + "app.services.buckets.providers.gcs._load_platform_sa_info", + return_value={"project_id": "platform-project"}, + ), + patch( + "app.services.buckets.providers.gcs.service_account." + "Credentials.from_service_account_info" + ) as mock_from_info, + patch("app.services.buckets.providers.gcs.gcs.Client"), + ): + client = GCSBucketProvider.create_client({}) + + assert client.default_bucket == "platform-bucket" + mock_from_info.assert_called_once_with( + {"project_id": "platform-project"}, scopes=list(GCS_SCOPES) + ) + + def test_missing_sa_info_raises(self): + with patch( + "app.services.buckets.providers.gcs._load_platform_sa_info", + return_value=None, + ): + with pytest.raises(ValueError) as exc_info: + GCSBucketProvider.create_client({}) + + assert "sa_key" in str(exc_info.value) + + +class TestParseGcsUri: + def test_parses_bucket_and_key(self): + assert GCSBucketProvider._parse_gcs_uri( + "gs://my-bucket/path/to/object.wav" + ) == ("my-bucket", "path/to/object.wav") + + def test_non_gs_scheme_raises(self): + with pytest.raises(ValueError): + GCSBucketProvider._parse_gcs_uri("s3://my-bucket/key") + + def test_missing_key_raises(self): + with pytest.raises(ValueError): + GCSBucketProvider._parse_gcs_uri("gs://my-bucket") + + +class TestGetSignedUrl: + def test_signs_with_v4_and_get(self): + provider, blob = _make_provider() + blob.generate_signed_url.return_value = "https://signed.example/obj" + + url = provider.get_signed_url("gs://my-bucket/key.wav", expires_in=1800) + + assert url == "https://signed.example/obj" + provider.client.storage_client.bucket.assert_called_once_with("my-bucket") + provider.client.storage_client.bucket.return_value.blob.assert_called_once_with( + "key.wav" + ) + blob.generate_signed_url.assert_called_once_with( + version="v4", + expiration=timedelta(seconds=1800), + method="GET", + ) + + def test_expiry_capped_at_max(self): + provider, blob = _make_provider() + blob.generate_signed_url.return_value = "https://signed.example/obj" + + provider.get_signed_url( + "gs://my-bucket/key.wav", + expires_in=provider.MAX_SIGNED_URL_EXPIRY + 10_000, + ) + + _, kwargs = blob.generate_signed_url.call_args + assert kwargs["expiration"] == timedelta( + seconds=provider.MAX_SIGNED_URL_EXPIRY + ) + + +class TestGetBulkSignedUrls: + def test_returns_uri_to_url_map_reusing_one_client(self): + provider, blob = _make_provider() + blob.generate_signed_url.side_effect = [ + "https://signed.example/a", + "https://signed.example/b", + ] + + result = provider.get_bulk_signed_urls( + ["gs://bucket/a.wav", "gs://bucket/b.wav"] + ) + + assert result == { + "gs://bucket/a.wav": "https://signed.example/a", + "gs://bucket/b.wav": "https://signed.example/b", + } + assert blob.generate_signed_url.call_count == 2 diff --git a/backend/app/tests/services/buckets/test_registry.py b/backend/app/tests/services/buckets/test_registry.py new file mode 100644 index 000000000..228780075 --- /dev/null +++ b/backend/app/tests/services/buckets/test_registry.py @@ -0,0 +1,102 @@ +"""Tests for the bucket provider registry.""" + +import pytest +from unittest.mock import MagicMock, patch + +from sqlmodel import Session + +from app.services.buckets.providers.base import BaseBucketProvider +from app.services.buckets.providers.gcs import GCSBucketProvider +from app.services.buckets.providers.registry import ( + BucketProvider, + get_bucket_provider, +) +from app.tests.utils.utils import get_project + + +class TestBucketProviderRegistry: + def test_get_provider_class_returns_gcs(self): + assert BucketProvider.get_provider_class("gcs") is GCSBucketProvider + + def test_registry_values_are_provider_classes(self): + for provider_type, provider_class in BucketProvider._registry.items(): + assert issubclass( + provider_class, BaseBucketProvider + ), f"Provider '{provider_type}' must inherit from BaseBucketProvider" + + def test_get_provider_class_unknown_raises(self): + with pytest.raises(ValueError) as exc_info: + BucketProvider.get_provider_class("s3") + message = str(exc_info.value) + assert "s3" in message + assert "is not supported" in message + + +class TestGetBucketProvider: + def test_get_bucket_provider_with_gcs(self, db: Session): + project = get_project(db) + + credential = { + "gcs_bucket": "byok-bucket", + "sa_key": {"project_id": "byok-project"}, + } + + with ( + patch( + "app.crud.credentials.get_provider_credential" + ) as mock_get_creds, + patch( + "app.services.buckets.providers.gcs.service_account." + "Credentials.from_service_account_info" + ), + patch("app.services.buckets.providers.gcs.gcs.Client") as mock_client, + ): + mock_get_creds.return_value = credential + mock_client.return_value = MagicMock() + + provider = get_bucket_provider( + session=db, + provider_type="gcs", + project_id=project.id, + organization_id=project.organization_id, + ) + + assert isinstance(provider, GCSBucketProvider) + mock_get_creds.assert_called_once_with( + session=db, + provider="google-gcp", + project_id=project.id, + org_id=project.organization_id, + ) + + def test_get_bucket_provider_missing_credential_raises(self, db: Session): + project = get_project(db) + + with patch( + "app.crud.credentials.get_provider_credential" + ) as mock_get_creds: + mock_get_creds.return_value = None + + with pytest.raises(ValueError) as exc_info: + get_bucket_provider( + session=db, + provider_type="gcs", + project_id=project.id, + organization_id=project.organization_id, + ) + + assert "google-gcp" in str(exc_info.value) + assert "not configured" in str(exc_info.value) + + def test_get_bucket_provider_unknown_type_raises(self, db: Session): + project = get_project(db) + + with pytest.raises(ValueError) as exc_info: + get_bucket_provider( + session=db, + provider_type="s3", + project_id=project.id, + organization_id=project.organization_id, + ) + + assert "s3" in str(exc_info.value) diff --git a/docs/wiki/domain-map.md b/docs/wiki/domain-map.md index 4d6d95425..b490401cd 100644 --- a/docs/wiki/domain-map.md +++ b/docs/wiki/domain-map.md @@ -22,7 +22,7 @@ APIKey → Organization, Project, User # programmatic access | Project | project.py | Organization | nearly all tables; unit of permissioning | | User | user.py | — | UserProject, APIKey, Notification | | APIKey | api_key.py | Organization, Project, User | auth dependency on every API route (logical) | -| Credential | credentials.py | Organization, Project | provider clients: OpenAI/Gemini/Anthropic calls (logical) | +| Credential | credentials.py | Organization, Project | provider clients: OpenAI/Gemini/Anthropic calls (logical); bucket providers (`google-gcp` → GCS signing, logical) | | Config | config/config.py | Project | ConfigVersion; LLM call path; Assessment; EvaluationRun (`config_id`) | | ConfigVersion | config/version.py | Config | resolved by `LLMCallConfig` saved references (logical) | | LlmCall | llm/request.py | Job, LlmChain, Org, Project | Langfuse traces (logical); analytics | @@ -54,7 +54,7 @@ APIKey → Organization, Project, User # programmatic access - **Langfuse** — every LLM call and evaluation run writes traces/scores. A change to run scoring or trace shape ripples here. - **kaapi-frontend console** — reads run results, annotation queues, config CRUD. A response-shape change ripples here. - **Provider Batch APIs** (OpenAI, Gemini, Anthropic in `core/batch/`) — eval/assessment payload shape changes ripple here. -- **Object storage** (`core/cloud/storage.py`) — files, dataset artifacts. +- **Object storage** (`core/cloud/storage.py` S3; `services/buckets/` GCS signed URLs) — files, dataset artifacts, gs:// attachments. ## Blast-radius procedure diff --git a/docs/wiki/modules/assessment.md b/docs/wiki/modules/assessment.md index e8e0b4186..097f665e4 100644 --- a/docs/wiki/modules/assessment.md +++ b/docs/wiki/modules/assessment.md @@ -21,6 +21,7 @@ Config version (tag=ASSESSMENT, `models/config/assessment_blob.py`) owns system `assessment.id` is a **UUID** (like config/job/llm_call). Per-item result = `AssessmentResult {output: {assessment, pre_filter}, error}` (no `metadata` — the provider/model/usage block was removed from the API-client output) where `output.assessment` = the LLM output parsed to an object when the config has a `json_output_schema`, else string (null for gated/failed rows), and `output.pre_filter` groups the two verdicts — `{topic_relevance, duplicate_detection}`, each `{verdict, reasoning}` or null — and is itself null when no pre-filter ran. Delivery is **webhook-only**: the `POST /assessments` ack is the flat `AssessmentSubmitResponse {assessment_id, status, message, inserted_at, updated_at}`, and the result is delivered solely by POSTing the `AssessmentCallback {assessment_id, status, data, request_metadata}` to the request's required `callback_url` on completion — where `data` is a single `AssessmentResult` (RESPONSE) or an `AssessmentBatchResult {total_items, counts, items}` (BATCH); `status` lives on the envelope only. Pre-filter `stop_on_fail` flag (`config/assessment_blob.py`) drives which filters hard-stop the chain on a failing verdict vs pass-through (record only). ## Services / CRUD +- `services/assessment/utils/attachments.py` — cell→provider attachment conversion (Drive URL normalization; OpenAI/Anthropic/Gemini part builders). `rewrite_gcs_attachment_urls` bulk-resolves `gs://` cells to provider-reachable URLs via `services/buckets/` (Path A native passthrough for google-gcp / Path B signed HTTPS otherwise) **before** JSONL build. Called in: `services/assessment/api/batch.py::_submit_provider_batch` (API pipeline, all stages), `crud/assessment/batch.py::submit_assessment_batch` (legacy L2), `services/assessment/tasks.py` (legacy prefilter). Submit-time validation (`services/assessment/api/submission.py`) allows `gs://` alongside `http(s)://`. - `services/assessment/` — legacy RUN pipeline (service, stages, processing, batch, cron, tasks) - `services/assessment/api/` — API-client pipeline: `submission.py` (submit), `batch.py` (staged provider batches — gate pre-filters → pass-through → assessment, over `core/batch`; `PREFILTER_VERDICT_SCHEMA`), `results.py` (builds `AssessmentBatchResult`), `callbacks.py` (webhook) - `crud/assessment/api.py` — new API-client crud (method-based Assessment/AssessmentRun writes): `create_assessment`, `set_assessment_job`, `create_execution`, `set_execution_batch_job`, `update_status`, `list_executions` (no `get_assessment` — delivery is webhook-only, so there is no request-time fetch). Namespaced under `api` (`from app.crud.assessment import api`) to avoid colliding with the legacy `create_assessment`. @@ -31,4 +32,5 @@ Config version (tag=ASSESSMENT, `models/config/assessment_blob.py`) owns system - API-client BATCH: `celery/tasks/job_execution.py::run_assessment_api_batch` drives the staged pipeline (poll → parse verdict/gate → advance/finalize → callback), self-re-enqueuing between stages. Staged state + `callback_url`/`request_metadata` live in the `assessment_run.execution` bag. ## External -- Provider Batch APIs, object storage for attachments. +- Provider Batch APIs, object storage for attachments (incl. `gs://` attachments resolved via `services/buckets/`). +- Gemini-family batch provider selection is centralized in `core/batch/client.py::get_gemini_batch_provider(provider_name=...)`: `google-gcp` -> `VertexBatchProvider` (Vertex, GCS in/out, `core/batch/vertex.py`), `google` -> `GeminiBatchProvider` (AI-Studio, File API). Shared by any service (assessment/evaluations/STT/TTS); assessment's `api/batch.py` routes through it. diff --git a/docs/wiki/modules/platform.md b/docs/wiki/modules/platform.md index 07b5fa1aa..755c4769b 100644 --- a/docs/wiki/modules/platform.md +++ b/docs/wiki/modules/platform.md @@ -12,5 +12,6 @@ All paths relative to `backend/app/`. | Languages | `api/routes/languages.py` | `global.languages` (`models/language.py`) | `crud/language.py` | | Credentials | `api/routes/credentials.py` | `credential` (`models/credentials.py`) | `crud/credentials.py`; provider keys per org/project; envelope encryption (KMS-wrapped data key + AES-GCM), prefix-versioned ciphertexts | | Model config | `api/routes/model_config.py` | `model_config` (`models/model_config.py`) | `crud/model_config.py` | +| Bucket providers | — | reuses `credential` (`google-gcp`) | `services/buckets/` — global registry + resolver (`providers/`), GCS signed/bulk-signed/private-to-public URLs (`providers/gcs.py`), attachment path selection + URL resolution (`attachments.py`) | | Cron | `api/routes/cron.py` | — | triggers batch polling (`crud/evaluations/cron.py`) | | Jobs | — | `job` (`models/job.py`), `batch_job` (`models/batch_job.py`) | `crud/jobs.py`, `crud/job/`, `services/job_monitoring.py` | From 80adb817b627046f80684782914c71b6469cd62d Mon Sep 17 00:00:00 2001 From: Prashant Vasudevan <71649489+vprashrex@users.noreply.github.com> Date: Wed, 19 Aug 2026 10:28:03 +0530 Subject: [PATCH 03/15] refactor: Clean up code formatting and improve readability in various files --- backend/app/services/assessment/tasks.py | 3 +-- backend/app/services/buckets/providers/gcs.py | 4 +--- backend/app/tests/assessment/test_api_batch.py | 10 ++++++++-- backend/app/tests/core/batch/test_client.py | 4 +--- backend/app/tests/core/batch/test_vertex.py | 8 ++++---- backend/app/tests/services/buckets/test_attachments.py | 4 +++- backend/app/tests/services/buckets/test_gcs.py | 8 ++------ backend/app/tests/services/buckets/test_registry.py | 8 ++------ 8 files changed, 22 insertions(+), 27 deletions(-) diff --git a/backend/app/services/assessment/tasks.py b/backend/app/services/assessment/tasks.py index 852725e7b..33e066385 100644 --- a/backend/app/services/assessment/tasks.py +++ b/backend/app/services/assessment/tasks.py @@ -254,8 +254,7 @@ def _submit_stage( organization_id=organization_id, ) rows_with_idx = [ - (idx, resolved_rows[pos]) - for pos, (idx, _) in enumerate(rows_with_idx) + (idx, resolved_rows[pos]) for pos, (idx, _) in enumerate(rows_with_idx) ] jsonl = build_prefilter_requests(stage, rows_with_idx, cfg, attachments) batch_job = submit_prefilter_batch( diff --git a/backend/app/services/buckets/providers/gcs.py b/backend/app/services/buckets/providers/gcs.py index 62967953d..16fa01cde 100644 --- a/backend/app/services/buckets/providers/gcs.py +++ b/backend/app/services/buckets/providers/gcs.py @@ -68,9 +68,7 @@ def _parse_gcs_uri(uri: str) -> tuple[str, str]: """Split ``gs://bucket/key`` into ``(bucket, key)``.""" parsed = urlparse(uri) if parsed.scheme != GCS_URI_SCHEME or not parsed.netloc: - raise ValueError( - f"Invalid GCS URI '{uri}'; expected 'gs://bucket/key'." - ) + raise ValueError(f"Invalid GCS URI '{uri}'; expected 'gs://bucket/key'.") key = parsed.path.lstrip("/") if not key: raise ValueError(f"GCS URI '{uri}' is missing an object key.") diff --git a/backend/app/tests/assessment/test_api_batch.py b/backend/app/tests/assessment/test_api_batch.py index 8fcbc38d9..864303b8c 100644 --- a/backend/app/tests/assessment/test_api_batch.py +++ b/backend/app/tests/assessment/test_api_batch.py @@ -879,7 +879,10 @@ def test_google(self, db) -> None: project_id=auth.project_id, ) assert provider is not None - assert get_provider.call_args.kwargs["provider_name"] in ("google", "google-native") + assert get_provider.call_args.kwargs["provider_name"] in ( + "google", + "google-native", + ) def test_google_gcp_routes_to_vertex(self, db) -> None: auth = get_user_test_auth_context(db) @@ -956,7 +959,10 @@ def test_google_branch(self, db) -> None: ) assert result.id == job.id assert start.call_args.kwargs["provider_name"] == "google" - assert get_provider.call_args.kwargs["provider_name"] in ("google", "google-native") + assert get_provider.call_args.kwargs["provider_name"] in ( + "google", + "google-native", + ) def test_google_gcp_branch_routes_to_vertex(self, db) -> None: auth = get_user_test_auth_context(db) diff --git a/backend/app/tests/core/batch/test_client.py b/backend/app/tests/core/batch/test_client.py index dd44453b8..6383ef6a9 100644 --- a/backend/app/tests/core/batch/test_client.py +++ b/backend/app/tests/core/batch/test_client.py @@ -63,9 +63,7 @@ def test_google_gcp_routes_to_vertex(self): assert from_cred.call_args.kwargs["model"] == "gemini-2.5-pro" def test_missing_gcp_credential_raises_404(self): - with patch( - "app.core.batch.client.get_provider_credential", return_value=None - ): + with patch("app.core.batch.client.get_provider_credential", return_value=None): with pytest.raises(HTTPException) as exc: get_gemini_batch_provider( session=MagicMock(), diff --git a/backend/app/tests/core/batch/test_vertex.py b/backend/app/tests/core/batch/test_vertex.py index 5c14af33d..36ac54373 100644 --- a/backend/app/tests/core/batch/test_vertex.py +++ b/backend/app/tests/core/batch/test_vertex.py @@ -117,7 +117,9 @@ def test_falls_back_to_line_order_without_key(self, provider, mock_storage): results = provider.download_batch_results("gs://b/out/") assert [r["custom_id"] for r in results] == ["0", "1"] - def test_batch_name_resolves_dest_then_reads(self, provider, mock_genai, mock_storage): + def test_batch_name_resolves_dest_then_reads( + self, provider, mock_genai, mock_storage + ): job = MagicMock() job.state.name = "JOB_STATE_SUCCEEDED" job.dest.gcs_uri = "gs://b/out/" @@ -176,9 +178,7 @@ def test_builds_provider(self): assert genai_client.call_args.kwargs["vertexai"] is True gcs_client.assert_called_once() - @pytest.mark.parametrize( - "drop", ["project_id", "location", "gcs_bucket", "sa_key"] - ) + @pytest.mark.parametrize("drop", ["project_id", "location", "gcs_bucket", "sa_key"]) def test_missing_field_raises(self, drop): cred = {k: v for k, v in self._CRED.items() if k != drop} with pytest.raises(ValueError, match=drop): diff --git a/backend/app/tests/services/buckets/test_attachments.py b/backend/app/tests/services/buckets/test_attachments.py index aa640df4f..6b9f36590 100644 --- a/backend/app/tests/services/buckets/test_attachments.py +++ b/backend/app/tests/services/buckets/test_attachments.py @@ -74,7 +74,9 @@ def test_gcs_native_passthrough(self): def test_gcs_signed_for_non_native_provider(self): provider = MagicMock() - provider.get_bulk_signed_urls.return_value = {"gs://bucket/key.png": "https://signed"} + provider.get_bulk_signed_urls.return_value = { + "gs://bucket/key.png": "https://signed" + } with patch(_RESOLVER_PATH, return_value=provider): url = resolve_attachments( session=MagicMock(), diff --git a/backend/app/tests/services/buckets/test_gcs.py b/backend/app/tests/services/buckets/test_gcs.py index a4ca51d0c..078b432da 100644 --- a/backend/app/tests/services/buckets/test_gcs.py +++ b/backend/app/tests/services/buckets/test_gcs.py @@ -40,9 +40,7 @@ def test_byok_credentials_win_over_settings(self, monkeypatch): "app.services.buckets.providers.gcs.service_account." "Credentials.from_service_account_info" ) as mock_from_info, - patch( - "app.services.buckets.providers.gcs.gcs.Client" - ) as mock_client, + patch("app.services.buckets.providers.gcs.gcs.Client") as mock_client, ): client = GCSBucketProvider.create_client( {"gcs_bucket": "byok-bucket", "sa_key": byok_sa} @@ -132,9 +130,7 @@ def test_expiry_capped_at_max(self): ) _, kwargs = blob.generate_signed_url.call_args - assert kwargs["expiration"] == timedelta( - seconds=provider.MAX_SIGNED_URL_EXPIRY - ) + assert kwargs["expiration"] == timedelta(seconds=provider.MAX_SIGNED_URL_EXPIRY) class TestGetBulkSignedUrls: diff --git a/backend/app/tests/services/buckets/test_registry.py b/backend/app/tests/services/buckets/test_registry.py index 228780075..858e122e4 100644 --- a/backend/app/tests/services/buckets/test_registry.py +++ b/backend/app/tests/services/buckets/test_registry.py @@ -42,9 +42,7 @@ def test_get_bucket_provider_with_gcs(self, db: Session): } with ( - patch( - "app.crud.credentials.get_provider_credential" - ) as mock_get_creds, + patch("app.crud.credentials.get_provider_credential") as mock_get_creds, patch( "app.services.buckets.providers.gcs.service_account." "Credentials.from_service_account_info" @@ -72,9 +70,7 @@ def test_get_bucket_provider_with_gcs(self, db: Session): def test_get_bucket_provider_missing_credential_raises(self, db: Session): project = get_project(db) - with patch( - "app.crud.credentials.get_provider_credential" - ) as mock_get_creds: + with patch("app.crud.credentials.get_provider_credential") as mock_get_creds: mock_get_creds.return_value = None with pytest.raises(ValueError) as exc_info: From e4ee220c268568dbd3a8d49308d4a788bd4fd6b9 Mon Sep 17 00:00:00 2001 From: Prashant Vasudevan <71649489+vprashrex@users.noreply.github.com> Date: Wed, 19 Aug 2026 12:13:21 +0530 Subject: [PATCH 04/15] feat(batch): Implement VertexBatchProvider for Google GCP integration and refactor batch provider selection --- backend/app/core/batch/__init__.py | 11 +-- backend/app/core/batch/client.py | 50 ------------- .../core/batch/{vertex.py => google_gcp.py} | 0 backend/app/services/assessment/api/batch.py | 64 +++++++++++----- .../app/services/assessment/api/submission.py | 1 - .../services/assessment/utils/attachments.py | 8 +- .../app/tests/assessment/test_api_batch.py | 69 ++++++++++------- backend/app/tests/core/batch/test_client.py | 74 ------------------- .../{test_vertex.py => test_google_gcp.py} | 8 +- docs/wiki/modules/assessment.md | 2 +- 10 files changed, 96 insertions(+), 191 deletions(-) rename backend/app/core/batch/{vertex.py => google_gcp.py} (100%) delete mode 100644 backend/app/tests/core/batch/test_client.py rename backend/app/tests/core/batch/{test_vertex.py => test_google_gcp.py} (95%) diff --git a/backend/app/core/batch/__init__.py b/backend/app/core/batch/__init__.py index 291dc63db..c35948171 100644 --- a/backend/app/core/batch/__init__.py +++ b/backend/app/core/batch/__init__.py @@ -2,12 +2,7 @@ from .anthropic import AnthropicBatchProvider, MessageBatchStatus from .base import BATCH_KEY, BatchProvider -from .client import ( - GeminiClient, - GeminiClientError, - get_gemini_batch_provider, - is_vertex_batch_provider, -) +from .client import GeminiClient, GeminiClientError from .gemini import ( BatchJobState, GeminiBatchProvider, @@ -16,7 +11,7 @@ extract_text_from_response_dict, ) from .openai import OpenAIBatchProvider -from .vertex import VertexBatchProvider +from .google_gcp import VertexBatchProvider from .operations import ( download_batch_results, process_completed_batch, @@ -35,8 +30,6 @@ "GeminiClientError", "GeminiBatchProvider", "VertexBatchProvider", - "get_gemini_batch_provider", - "is_vertex_batch_provider", "OpenAIBatchProvider", "create_stt_batch_requests", "create_tts_batch_requests", diff --git a/backend/app/core/batch/client.py b/backend/app/core/batch/client.py index 66bc06c0d..804793a69 100644 --- a/backend/app/core/batch/client.py +++ b/backend/app/core/batch/client.py @@ -9,8 +9,6 @@ from fastapi import HTTPException from app.crud.credentials import get_provider_credential -from .base import BatchProvider - logger = logging.getLogger(__name__) @@ -93,51 +91,3 @@ def from_credentials( f"org_id: {org_id}, project_id: {project_id}" ) return cls(api_key=api_key) - - -# Providers whose batch jobs run on Vertex AI (GCS-backed) rather than AI-Studio. -VERTEX_BATCH_PROVIDERS = {"google-gcp", "google-gcp-native"} - - -def is_vertex_batch_provider(provider_name: str) -> bool: - """Whether a provider name routes batch jobs through Vertex AI (vs AI-Studio).""" - return provider_name in VERTEX_BATCH_PROVIDERS - - -def get_gemini_batch_provider( - *, - session: Session, - organization_id: int, - project_id: int, - provider_name: str, - model: str | None = None, -) -> BatchProvider: - """Resolve a Gemini-family batch provider from the provider name. - - ``google-gcp`` routes to Vertex (GCS-backed); ``google``/``google-aistudio`` - to AI-Studio (File API). Both honor the same BatchProvider contract, so any - service (assessment, evaluations, STT/TTS) can call this single resolver. - """ - from .gemini import GeminiBatchProvider - from .vertex import VertexBatchProvider - - if is_vertex_batch_provider(provider_name): - cred = get_provider_credential( - session=session, - provider="google-gcp", - project_id=project_id, - org_id=organization_id, - ) - if not cred: - raise HTTPException( - status_code=404, - detail="google-gcp credentials not configured for this project", - ) - if not isinstance(cred, dict): - raise GeminiClientError("Expected decrypted google-gcp credentials dict") - return VertexBatchProvider.from_credentials(cred, model=model) - - gemini = GeminiClient.from_credentials( - session=session, org_id=organization_id, project_id=project_id - ) - return GeminiBatchProvider(client=gemini.client, model=model) diff --git a/backend/app/core/batch/vertex.py b/backend/app/core/batch/google_gcp.py similarity index 100% rename from backend/app/core/batch/vertex.py rename to backend/app/core/batch/google_gcp.py diff --git a/backend/app/services/assessment/api/batch.py b/backend/app/services/assessment/api/batch.py index afc8724c7..1f92c5208 100644 --- a/backend/app/services/assessment/api/batch.py +++ b/backend/app/services/assessment/api/batch.py @@ -27,14 +27,17 @@ BatchJobState, MessageBatchStatus, OpenAIBatchProvider, + GeminiBatchProvider, + VertexBatchProvider, extract_text_from_response_dict, - get_gemini_batch_provider, - is_vertex_batch_provider, poll_batch_status, process_completed_batch, start_batch_job, ) from app.core.batch.base import BatchProvider +from app.core.batch.client import GeminiClient +from app.crud.credentials import get_provider_credential +from fastapi import HTTPException from app.core.config import settings from app.core.db import engine from app.crud.assessment import api @@ -323,19 +326,29 @@ def _submit_provider_batch( jsonl = build_google_jsonl( rows, text_columns, attachments, prompt, mapped, row_indices ) - provider = get_gemini_batch_provider( - session=session, - organization_id=organization_id, - project_id=project_id, - provider_name=provider_name, - model=model, - ) - # Vertex takes a bare model id; AI-Studio uses the "models/" prefix. - config = ( - {"display_name": description} - if is_vertex_batch_provider(provider_name) - else {"display_name": description, "model": f"models/{model}"} - ) + if provider_name == LLMProvider.GOOGLE_GCP: + cred = get_provider_credential( + session=session, + provider="google-gcp", + project_id=project_id, + org_id=organization_id, + ) + if not isinstance(cred, dict): + raise HTTPException( + status_code=404, + detail="google-gcp credentials not configured for this project", + ) + provider = VertexBatchProvider.from_credentials(cred, model=model) + config = {"display_name": description} # Vertex uses a bare model id + else: + gemini = GeminiClient.from_credentials( + session=session, org_id=organization_id, project_id=project_id + ) + provider = GeminiBatchProvider( + client=gemini.client, model=f"models/{model}" + ) + config = {"display_name": description, "model": f"models/{model}"} + elif provider_name == LLMProvider.ANTHROPIC: mapped, _ = map_kaapi_to_anthropic_params(params) jsonl = build_anthropic_jsonl( @@ -383,12 +396,23 @@ def _build_batch_provider( ) ) if provider_name in (LLMProvider.GOOGLE, LLMProvider.GOOGLE_GCP): - return get_gemini_batch_provider( - session=session, - organization_id=organization_id, - project_id=project_id, - provider_name=provider_name, + if provider_name == LLMProvider.GOOGLE_GCP: + cred = get_provider_credential( + session=session, + provider="google-gcp", + project_id=project_id, + org_id=organization_id, + ) + if not isinstance(cred, dict): + raise HTTPException( + status_code=404, + detail="google-gcp credentials not configured for this project", + ) + return VertexBatchProvider.from_credentials(cred) + gemini = GeminiClient.from_credentials( + session=session, org_id=organization_id, project_id=project_id ) + return GeminiBatchProvider(client=gemini.client) if provider_name == LLMProvider.ANTHROPIC: return AnthropicBatchProvider( client=get_anthropic_client( diff --git a/backend/app/services/assessment/api/submission.py b/backend/app/services/assessment/api/submission.py index 3edc8774d..1778d8695 100644 --- a/backend/app/services/assessment/api/submission.py +++ b/backend/app/services/assessment/api/submission.py @@ -31,7 +31,6 @@ logger = logging.getLogger(__name__) # Attachment cell values are provided as URLs (base64 is unsupported for batch). -# gs:// is allowed: it is resolved to a provider-reachable URL before batch build. _URL_PREFIXES = ("http://", "https://", "gs://") _ATTACHMENT_TYPES = ("image", "pdf") diff --git a/backend/app/services/assessment/utils/attachments.py b/backend/app/services/assessment/utils/attachments.py index 7e3f28dce..c5002065e 100644 --- a/backend/app/services/assessment/utils/attachments.py +++ b/backend/app/services/assessment/utils/attachments.py @@ -1,10 +1,4 @@ -"""Attachment resolution utilities for assessment batch builds. - -URL-only: dataset cells hold attachment URLs. Handles Google Drive URL -normalization and conversion of cell values into provider input objects. -Attachments are passed to providers by reference (URL), never inlined as base64, -to keep the batch build memory-light. -""" +"""Attachment utilities for assessment batch builds: URL normalization, gs:// resolution, provider parts.""" import logging import re diff --git a/backend/app/tests/assessment/test_api_batch.py b/backend/app/tests/assessment/test_api_batch.py index 864303b8c..e6db26017 100644 --- a/backend/app/tests/assessment/test_api_batch.py +++ b/backend/app/tests/assessment/test_api_batch.py @@ -868,10 +868,10 @@ def test_openai(self, db) -> None: def test_google(self, db) -> None: auth = get_user_test_auth_context(db) - with patch( - "app.services.assessment.api.batch.get_gemini_batch_provider", - return_value=MagicMock(), - ) as get_provider: + gemini = MagicMock() + gemini.client = MagicMock() + with patch("app.services.assessment.api.batch.GeminiClient") as gemini_cls: + gemini_cls.from_credentials.return_value = gemini provider = _build_batch_provider( session=db, provider_name="google", @@ -879,24 +879,43 @@ def test_google(self, db) -> None: project_id=auth.project_id, ) assert provider is not None - assert get_provider.call_args.kwargs["provider_name"] in ( - "google", - "google-native", - ) def test_google_gcp_routes_to_vertex(self, db) -> None: auth = get_user_test_auth_context(db) - with patch( - "app.services.assessment.api.batch.get_gemini_batch_provider", - return_value=MagicMock(), - ) as get_provider: - _build_batch_provider( + sentinel = MagicMock() + with ( + patch( + "app.services.assessment.api.batch.get_provider_credential", + return_value={"gcs_bucket": "b", "sa_key": {}}, + ), + patch( + "app.services.assessment.api.batch.VertexBatchProvider.from_credentials", + return_value=sentinel, + ) as vertex_from_cred, + ): + provider = _build_batch_provider( session=db, provider_name="google-gcp", organization_id=auth.organization_id, project_id=auth.project_id, ) - assert get_provider.call_args.kwargs["provider_name"] == "google-gcp" + assert provider is sentinel + vertex_from_cred.assert_called_once() + + def test_google_gcp_missing_credential_raises_404(self, db) -> None: + auth = get_user_test_auth_context(db) + with patch( + "app.services.assessment.api.batch.get_provider_credential", + return_value=None, + ): + with pytest.raises(HTTPException) as exc: + _build_batch_provider( + session=db, + provider_name="google-gcp", + organization_id=auth.organization_id, + project_id=auth.project_id, + ) + assert exc.value.status_code == 404 def test_anthropic(self, db) -> None: auth = get_user_test_auth_context(db) @@ -941,28 +960,24 @@ def _kwargs(self, db, auth, provider_name, params): def test_google_branch(self, db) -> None: auth = get_user_test_auth_context(db) + gemini = MagicMock() + gemini.client = MagicMock() job = _make_batch_job( db, org_id=auth.organization_id, project_id=auth.project_id ) with ( - patch( - "app.services.assessment.api.batch.get_gemini_batch_provider", - return_value=MagicMock(), - ) as get_provider, + patch("app.services.assessment.api.batch.GeminiClient") as gemini_cls, patch( "app.services.assessment.api.batch.start_batch_job", return_value=job, ) as start, ): + gemini_cls.from_credentials.return_value = gemini result = _submit_provider_batch( **self._kwargs(db, auth, "google", {"model": "gemini-2.5-pro"}) ) assert result.id == job.id assert start.call_args.kwargs["provider_name"] == "google" - assert get_provider.call_args.kwargs["provider_name"] in ( - "google", - "google-native", - ) def test_google_gcp_branch_routes_to_vertex(self, db) -> None: auth = get_user_test_auth_context(db) @@ -971,9 +986,13 @@ def test_google_gcp_branch_routes_to_vertex(self, db) -> None: ) with ( patch( - "app.services.assessment.api.batch.get_gemini_batch_provider", + "app.services.assessment.api.batch.get_provider_credential", + return_value={"gcs_bucket": "b", "sa_key": {}}, + ), + patch( + "app.services.assessment.api.batch.VertexBatchProvider.from_credentials", return_value=MagicMock(), - ) as get_provider, + ) as vertex_from_cred, patch( "app.services.assessment.api.batch.start_batch_job", return_value=job, @@ -983,7 +1002,7 @@ def test_google_gcp_branch_routes_to_vertex(self, db) -> None: **self._kwargs(db, auth, "google-gcp", {"model": "gemini-2.5-pro"}) ) assert start.call_args.kwargs["provider_name"] == "google-gcp" - assert get_provider.call_args.kwargs["provider_name"] == "google-gcp" + vertex_from_cred.assert_called_once() # Vertex config omits the "models/" prefixed model (uses bare id). assert "model" not in start.call_args.kwargs["config"] diff --git a/backend/app/tests/core/batch/test_client.py b/backend/app/tests/core/batch/test_client.py deleted file mode 100644 index 6383ef6a9..000000000 --- a/backend/app/tests/core/batch/test_client.py +++ /dev/null @@ -1,74 +0,0 @@ -"""Test cases for the Gemini-family batch provider resolver.""" - -from unittest.mock import MagicMock, patch - -import pytest -from fastapi import HTTPException - -from app.core.batch.client import ( - get_gemini_batch_provider, - is_vertex_batch_provider, -) -from app.core.batch.gemini import GeminiBatchProvider - - -class TestIsVertexBatchProvider: - @pytest.mark.parametrize("name", ["google-gcp", "google-gcp-native"]) - def test_vertex_providers(self, name): - assert is_vertex_batch_provider(name) is True - - @pytest.mark.parametrize("name", ["google", "google-aistudio", "openai"]) - def test_non_vertex_providers(self, name): - assert is_vertex_batch_provider(name) is False - - -class TestGetGeminiBatchProvider: - def test_google_routes_to_aistudio(self): - gemini = MagicMock() - gemini.client = MagicMock() - with patch( - "app.core.batch.client.GeminiClient.from_credentials", - return_value=gemini, - ) as from_cred: - provider = get_gemini_batch_provider( - session=MagicMock(), - organization_id=1, - project_id=2, - provider_name="google", - ) - assert isinstance(provider, GeminiBatchProvider) - from_cred.assert_called_once() - - def test_google_gcp_routes_to_vertex(self): - sentinel = MagicMock() - with ( - patch( - "app.core.batch.client.get_provider_credential", - return_value={"gcs_bucket": "b", "sa_key": {}}, - ) as get_cred, - patch( - "app.core.batch.vertex.VertexBatchProvider.from_credentials", - return_value=sentinel, - ) as from_cred, - ): - provider = get_gemini_batch_provider( - session=MagicMock(), - organization_id=1, - project_id=2, - provider_name="google-gcp", - model="gemini-2.5-pro", - ) - assert provider is sentinel - assert get_cred.call_args.kwargs["provider"] == "google-gcp" - assert from_cred.call_args.kwargs["model"] == "gemini-2.5-pro" - - def test_missing_gcp_credential_raises_404(self): - with patch("app.core.batch.client.get_provider_credential", return_value=None): - with pytest.raises(HTTPException) as exc: - get_gemini_batch_provider( - session=MagicMock(), - organization_id=1, - project_id=2, - provider_name="google-gcp", - ) - assert exc.value.status_code == 404 diff --git a/backend/app/tests/core/batch/test_vertex.py b/backend/app/tests/core/batch/test_google_gcp.py similarity index 95% rename from backend/app/tests/core/batch/test_vertex.py rename to backend/app/tests/core/batch/test_google_gcp.py index 36ac54373..183bdf6dd 100644 --- a/backend/app/tests/core/batch/test_vertex.py +++ b/backend/app/tests/core/batch/test_google_gcp.py @@ -5,7 +5,7 @@ import pytest -from app.core.batch.vertex import VertexBatchProvider, _parse_gs_uri +from app.core.batch.google_gcp import VertexBatchProvider, _parse_gs_uri _BUCKET = "test-bucket" @@ -168,9 +168,9 @@ class TestFromCredentials: def test_builds_provider(self): with ( - patch("app.core.batch.vertex.service_account") as sa, - patch("app.core.batch.vertex.genai.Client") as genai_client, - patch("app.core.batch.vertex.gcs.Client") as gcs_client, + patch("app.core.batch.google_gcp.service_account") as sa, + patch("app.core.batch.google_gcp.genai.Client") as genai_client, + patch("app.core.batch.google_gcp.gcs.Client") as gcs_client, ): provider = VertexBatchProvider.from_credentials(self._CRED) assert isinstance(provider, VertexBatchProvider) diff --git a/docs/wiki/modules/assessment.md b/docs/wiki/modules/assessment.md index 097f665e4..97047be65 100644 --- a/docs/wiki/modules/assessment.md +++ b/docs/wiki/modules/assessment.md @@ -33,4 +33,4 @@ Config version (tag=ASSESSMENT, `models/config/assessment_blob.py`) owns system ## External - Provider Batch APIs, object storage for attachments (incl. `gs://` attachments resolved via `services/buckets/`). -- Gemini-family batch provider selection is centralized in `core/batch/client.py::get_gemini_batch_provider(provider_name=...)`: `google-gcp` -> `VertexBatchProvider` (Vertex, GCS in/out, `core/batch/vertex.py`), `google` -> `GeminiBatchProvider` (AI-Studio, File API). Shared by any service (assessment/evaluations/STT/TTS); assessment's `api/batch.py` routes through it. +- Gemini-family batch provider is chosen inline in `api/batch.py` (`_submit_provider_batch` / `_build_batch_provider`): `google-gcp` -> `VertexBatchProvider` (Vertex, GCS in/out, `core/batch/vertex.py`, built from the `google-gcp` credential), `google` -> `GeminiBatchProvider` (AI-Studio, File API, via `GeminiClient`). From 8942549337060f638c002a9b7fbb8648b71826d2 Mon Sep 17 00:00:00 2001 From: Prajna1999 Date: Thu, 20 Aug 2026 22:04:55 +0530 Subject: [PATCH 05/15] 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 From b2082558b5c4ad37c61dcb25e987b0ee1a61d6b0 Mon Sep 17 00:00:00 2001 From: Prashant Vasudevan <71649489+vprashrex@users.noreply.github.com> Date: Fri, 21 Aug 2026 10:40:56 +0530 Subject: [PATCH 06/15] feat(bucket): Add configurable signed URL expiry settings for GCS bucket provider --- backend/app/core/config.py | 4 +++ .../services/assessment/utils/attachments.py | 2 ++ backend/app/services/buckets/attachments.py | 6 +++- .../app/services/buckets/providers/base.py | 18 ++++------ backend/app/services/buckets/providers/gcs.py | 33 +++++++++++++++++-- 5 files changed, 48 insertions(+), 15 deletions(-) diff --git a/backend/app/core/config.py b/backend/app/core/config.py index 4d1544da3..baccbb43a 100644 --- a/backend/app/core/config.py +++ b/backend/app/core/config.py @@ -117,6 +117,10 @@ def SQLALCHEMY_DATABASE_URI(self) -> PostgresDsn: # Used by the registry fallback when a project has no ``google`` row. GCP_SA_KEY: str = "" GCS_AUDIO_BUCKET: str = "" + # Signed-URL lifetime for private bucket attachments (Path B); default 1h. + SIGNED_URL_EXPIRY_SECONDS: int = 3600 + # Hard cap on signed-URL lifetime (24h) enforced by bucket providers. + MAX_SIGNED_URL_EXPIRY_SECONDS: int = 86400 # RabbitMQ configuration for Celery broker RABBITMQ_HOST: str = "localhost" diff --git a/backend/app/services/assessment/utils/attachments.py b/backend/app/services/assessment/utils/attachments.py index c5002065e..96caed7b2 100644 --- a/backend/app/services/assessment/utils/attachments.py +++ b/backend/app/services/assessment/utils/attachments.py @@ -7,6 +7,7 @@ from sqlmodel import Session +from app.core.config import settings from app.models.assessment import AssessmentAttachment from app.models.llm.constants import KaapiProvider from app.services.buckets.attachments import is_gcs_uri, resolve_attachments @@ -62,6 +63,7 @@ def rewrite_gcs_attachment_urls( llm_provider=cast(KaapiProvider, llm_provider), project_id=project_id, organization_id=organization_id, + expires_in=settings.MAX_SIGNED_URL_EXPIRY_SECONDS, ) assert isinstance(resolved, dict) # list input always yields a dict diff --git a/backend/app/services/buckets/attachments.py b/backend/app/services/buckets/attachments.py index af28ce903..6ced39128 100644 --- a/backend/app/services/buckets/attachments.py +++ b/backend/app/services/buckets/attachments.py @@ -4,6 +4,7 @@ from sqlmodel import Session +from app.core.config import settings from app.models.llm.constants import KaapiProvider, Provider from app.services.buckets.providers.registry import get_bucket_provider @@ -44,7 +45,7 @@ def resolve_attachments( llm_provider: KaapiProvider, project_id: int, organization_id: int, - expires_in: int = 3600, + expires_in: int | None = None, bucket_provider_type: str = DEFAULT_BUCKET_PROVIDER, ) -> str | dict[str, str]: """Make attachment(s) LLM-reachable. Accepts a single URL or a list. @@ -54,6 +55,9 @@ def resolve_attachments( Path-B URIs are bulk-signed with a single bucket-provider client per call. Returns a ``str`` for a single input, a ``uri -> url`` dict for a list. """ + if expires_in is None: + expires_in = settings.SIGNED_URL_EXPIRY_SECONDS + is_single = isinstance(source, str) uris = [source] if is_single else list(source) diff --git a/backend/app/services/buckets/providers/base.py b/backend/app/services/buckets/providers/base.py index 1951de7ae..d5894fecf 100644 --- a/backend/app/services/buckets/providers/base.py +++ b/backend/app/services/buckets/providers/base.py @@ -4,6 +4,8 @@ from abc import ABC, abstractmethod from typing import Any +from app.core.config import settings + logger = logging.getLogger(__name__) @@ -13,8 +15,8 @@ class BaseBucketProvider(ABC): # URI scheme this provider handles (e.g. "gs", "s3"). SCHEME: str = "" - # Cap on signed-URL lifetime (24h), matching AmazonCloudStorage. - MAX_SIGNED_URL_EXPIRY: int = 86400 + # Cap on signed-URL lifetime, enforced by concrete providers. + MAX_SIGNED_URL_EXPIRY: int = settings.MAX_SIGNED_URL_EXPIRY_SECONDS def __init__(self, client: Any): self.client = client @@ -26,23 +28,15 @@ def create_client(credentials: dict[str, Any]) -> Any: raise NotImplementedError("Bucket providers must implement create_client") @abstractmethod - def get_signed_url(self, uri: str, expires_in: int = 3600) -> str: + def get_signed_url(self, uri: str, expires_in: int | None = None) -> str: """Generate a time-limited signed URL for a single object.""" raise NotImplementedError("Bucket providers must implement get_signed_url") def get_bulk_signed_urls( - self, uris: list[str], expires_in: int = 3600 + self, uris: list[str], expires_in: int | None = None ) -> dict[str, str]: """Sign many object URIs, reusing this provider's single client.""" signed: dict[str, str] = {} for uri in uris: signed[uri] = self.get_signed_url(uri, expires_in=expires_in) return signed - - def to_public_url(self, uri: str, expires_in: int = 3600) -> str: - """Resolve a private object URI to a fetchable signed URL.""" - return self.get_signed_url(uri, expires_in=expires_in) - - def get_provider_name(self) -> str: - """Return the provider name derived from the class name.""" - return self.__class__.__name__.replace("BucketProvider", "").lower() diff --git a/backend/app/services/buckets/providers/gcs.py b/backend/app/services/buckets/providers/gcs.py index 16fa01cde..772f9933d 100644 --- a/backend/app/services/buckets/providers/gcs.py +++ b/backend/app/services/buckets/providers/gcs.py @@ -74,9 +74,33 @@ def _parse_gcs_uri(uri: str) -> tuple[str, str]: raise ValueError(f"GCS URI '{uri}' is missing an object key.") return parsed.netloc, key - def get_signed_url(self, uri: str, expires_in: int = 3600) -> str: + def _assert_in_credential_bucket(self, bucket_name: str, uri: str) -> None: + """Path A: SA can only sign its own bucket; a foreign URI would 403 mid-batch.""" + expected = self.client.default_bucket + if expected and bucket_name != expected: + raise ValueError( + f"Attachment bucket '{bucket_name}' does not match the credential " + f"bucket '{expected}' ({uri}); the service account cannot access it." + ) + + def _assert_all_in_credential_bucket(self, uris: list[str]) -> None: + """Validate the whole batch upfront so every mismatch surfaces before signing.""" + expected = self.client.default_bucket + if not expected: + return + mismatched = [uri for uri in uris if self._parse_gcs_uri(uri)[0] != expected] + if mismatched: + raise ValueError( + f"{len(mismatched)} attachment(s) not in credential bucket " + f"'{expected}'; the service account cannot access them: {mismatched}" + ) + + def get_signed_url(self, uri: str, expires_in: int | None = None) -> str: + if expires_in is None: + expires_in = settings.SIGNED_URL_EXPIRY_SECONDS expires_in = min(expires_in, self.MAX_SIGNED_URL_EXPIRY) bucket_name, key = self._parse_gcs_uri(uri) + self._assert_in_credential_bucket(bucket_name, uri) try: blob = self.client.storage_client.bucket(bucket_name).blob(key) @@ -100,9 +124,14 @@ def get_signed_url(self, uri: str, expires_in: int = 3600) -> str: return signed_url def get_bulk_signed_urls( - self, uris: list[str], expires_in: int = 3600 + self, uris: list[str], expires_in: int | None = None ) -> dict[str, str]: """Sign each URI reusing this provider's single client.""" + if expires_in is None: + expires_in = settings.SIGNED_URL_EXPIRY_SECONDS + + self._assert_all_in_credential_bucket(uris) + logger.info( f"[GCSBucketProvider.get_bulk_signed_urls] Signing batch | " f"count={len(uris)}, expires_in={expires_in}" From 9cac6c5bb57c0fb48cf6661017c0484fda3868b3 Mon Sep 17 00:00:00 2001 From: Prashant Vasudevan <71649489+vprashrex@users.noreply.github.com> Date: Mon, 24 Aug 2026 14:44:11 +0530 Subject: [PATCH 07/15] feat(gcp): Update Google GCP provider to support new model and enhance credential handling --- backend/app/core/batch/google_gcp.py | 2 +- backend/app/models/llm/constants.py | 1 + backend/app/services/assessment/api/batch.py | 38 +++--- .../services/assessment/utils/attachments.py | 14 +- backend/app/services/buckets/attachments.py | 4 - backend/app/services/buckets/providers/gcs.py | 28 +--- .../app/services/llm/providers/google_gcp.py | 121 +++++++++++++++++- .../services/buckets/test_attachments.py | 2 +- .../services/llm/providers/test_google_gcp.py | 38 ++++-- docs/wiki/modules/assessment.md | 2 +- docs/wiki/modules/platform.md | 2 +- 11 files changed, 183 insertions(+), 69 deletions(-) diff --git a/backend/app/core/batch/google_gcp.py b/backend/app/core/batch/google_gcp.py index 01870a954..6be7ce900 100644 --- a/backend/app/core/batch/google_gcp.py +++ b/backend/app/core/batch/google_gcp.py @@ -56,7 +56,7 @@ class VertexBatchProvider(BatchProvider): {"request": {"contents": [{"parts": [...], "role": "user"}]}} """ - DEFAULT_MODEL = "gemini-2.5-pro" + DEFAULT_MODEL = "gemini-3.1-pro-preview" def __init__( self, diff --git a/backend/app/models/llm/constants.py b/backend/app/models/llm/constants.py index 9a4cddd86..a85d70e5f 100644 --- a/backend/app/models/llm/constants.py +++ b/backend/app/models/llm/constants.py @@ -86,6 +86,7 @@ class Modality(StrEnum): "anthropic": "claude-sonnet-4-6", "openai": "gpt-4.1-mini", "google": "gemini-2.5-pro", + "google-gcp": "gemini-3.1-pro-preview", } DEFAULT_ANTHROPIC_MAX_TOKENS = 4096 diff --git a/backend/app/services/assessment/api/batch.py b/backend/app/services/assessment/api/batch.py index 1f92c5208..d31f7659e 100644 --- a/backend/app/services/assessment/api/batch.py +++ b/backend/app/services/assessment/api/batch.py @@ -280,6 +280,24 @@ def _stage_provider_model(blob: AssessmentConfigBlob, stage: str) -> tuple[str, return blob.assessment.provider, blob.assessment.params["model"] +def _google_gcp_credential( + *, session: Session, organization_id: int, project_id: int +) -> dict[str, Any]: + """Vertex needs the SA key + bucket; a missing credential is a client-fixable 404.""" + cred = get_provider_credential( + session=session, + provider=LLMProvider.GOOGLE_GCP, + project_id=project_id, + org_id=organization_id, + ) + if not isinstance(cred, dict): + raise HTTPException( + status_code=404, + detail="google-gcp credentials not configured for this project", + ) + return cred + + def _submit_provider_batch( *, session: Session, @@ -327,17 +345,11 @@ def _submit_provider_batch( rows, text_columns, attachments, prompt, mapped, row_indices ) if provider_name == LLMProvider.GOOGLE_GCP: - cred = get_provider_credential( + cred = _google_gcp_credential( session=session, - provider="google-gcp", + organization_id=organization_id, project_id=project_id, - org_id=organization_id, ) - if not isinstance(cred, dict): - raise HTTPException( - status_code=404, - detail="google-gcp credentials not configured for this project", - ) provider = VertexBatchProvider.from_credentials(cred, model=model) config = {"display_name": description} # Vertex uses a bare model id else: @@ -397,17 +409,11 @@ def _build_batch_provider( ) if provider_name in (LLMProvider.GOOGLE, LLMProvider.GOOGLE_GCP): if provider_name == LLMProvider.GOOGLE_GCP: - cred = get_provider_credential( + cred = _google_gcp_credential( session=session, - provider="google-gcp", + organization_id=organization_id, project_id=project_id, - org_id=organization_id, ) - if not isinstance(cred, dict): - raise HTTPException( - status_code=404, - detail="google-gcp credentials not configured for this project", - ) return VertexBatchProvider.from_credentials(cred) gemini = GeminiClient.from_credentials( session=session, org_id=organization_id, project_id=project_id diff --git a/backend/app/services/assessment/utils/attachments.py b/backend/app/services/assessment/utils/attachments.py index 96caed7b2..d4c41af14 100644 --- a/backend/app/services/assessment/utils/attachments.py +++ b/backend/app/services/assessment/utils/attachments.py @@ -74,10 +74,16 @@ def rewrite_gcs_attachment_urls( value = row.get(att.column) if not value: continue - new_row[att.column] = ", ".join( - resolved.get(attachment_url, attachment_url) - for attachment_url in split_attachment_urls(value) - ) + rewritten_urls: list[str] = [] + for attachment_url in split_attachment_urls(value): + target = resolved.get(attachment_url) + if target is None and is_gcs_uri(attachment_url): + logger.warning( + "[rewrite_gcs_attachment_urls] Unresolved gs:// attachment | column=%s", + att.column, + ) + rewritten_urls.append(target or attachment_url) + new_row[att.column] = ", ".join(rewritten_urls) rewritten.append(new_row) return rewritten diff --git a/backend/app/services/buckets/attachments.py b/backend/app/services/buckets/attachments.py index 6ced39128..2e5b60bb1 100644 --- a/backend/app/services/buckets/attachments.py +++ b/backend/app/services/buckets/attachments.py @@ -4,7 +4,6 @@ from sqlmodel import Session -from app.core.config import settings from app.models.llm.constants import KaapiProvider, Provider from app.services.buckets.providers.registry import get_bucket_provider @@ -55,9 +54,6 @@ def resolve_attachments( Path-B URIs are bulk-signed with a single bucket-provider client per call. Returns a ``str`` for a single input, a ``uri -> url`` dict for a list. """ - if expires_in is None: - expires_in = settings.SIGNED_URL_EXPIRY_SECONDS - is_single = isinstance(source, str) uris = [source] if is_single else list(source) diff --git a/backend/app/services/buckets/providers/gcs.py b/backend/app/services/buckets/providers/gcs.py index 772f9933d..d4e08f2fd 100644 --- a/backend/app/services/buckets/providers/gcs.py +++ b/backend/app/services/buckets/providers/gcs.py @@ -74,33 +74,15 @@ def _parse_gcs_uri(uri: str) -> tuple[str, str]: raise ValueError(f"GCS URI '{uri}' is missing an object key.") return parsed.netloc, key - def _assert_in_credential_bucket(self, bucket_name: str, uri: str) -> None: - """Path A: SA can only sign its own bucket; a foreign URI would 403 mid-batch.""" - expected = self.client.default_bucket - if expected and bucket_name != expected: - raise ValueError( - f"Attachment bucket '{bucket_name}' does not match the credential " - f"bucket '{expected}' ({uri}); the service account cannot access it." - ) - - def _assert_all_in_credential_bucket(self, uris: list[str]) -> None: - """Validate the whole batch upfront so every mismatch surfaces before signing.""" - expected = self.client.default_bucket - if not expected: - return - mismatched = [uri for uri in uris if self._parse_gcs_uri(uri)[0] != expected] - if mismatched: - raise ValueError( - f"{len(mismatched)} attachment(s) not in credential bucket " - f"'{expected}'; the service account cannot access them: {mismatched}" - ) - def get_signed_url(self, uri: str, expires_in: int | None = None) -> str: if expires_in is None: expires_in = settings.SIGNED_URL_EXPIRY_SECONDS + if expires_in < 1: + raise ValueError( + f"Signed-URL expiry must be at least 1 second, got {expires_in}." + ) expires_in = min(expires_in, self.MAX_SIGNED_URL_EXPIRY) bucket_name, key = self._parse_gcs_uri(uri) - self._assert_in_credential_bucket(bucket_name, uri) try: blob = self.client.storage_client.bucket(bucket_name).blob(key) @@ -130,8 +112,6 @@ def get_bulk_signed_urls( if expires_in is None: expires_in = settings.SIGNED_URL_EXPIRY_SECONDS - self._assert_all_in_credential_bucket(uris) - logger.info( f"[GCSBucketProvider.get_bulk_signed_urls] Signing batch | " f"count={len(uris)}, expires_in={expires_in}" diff --git a/backend/app/services/llm/providers/google_gcp.py b/backend/app/services/llm/providers/google_gcp.py index a6ca98661..5d49a7384 100644 --- a/backend/app/services/llm/providers/google_gcp.py +++ b/backend/app/services/llm/providers/google_gcp.py @@ -25,10 +25,12 @@ ) from app.models.llm.constants import ( DEFAULT_STT_MODEL, + DEFAULT_TEXT_MODELS, DEFAULT_TTS_MODEL, DEFAULT_TTS_VOICE, CompletionType, ) +from app.models.llm.request import ImageContent, PDFContent from app.models.llm.response import AudioContent, AudioOutput from app.services.llm.providers.base import BaseProvider, ContentPart, MultiModalInput @@ -106,9 +108,8 @@ def endpoint(self, model: str) -> str: class GoogleGCPProvider(BaseProvider): """Google GCP provider using REST + API key auth. - Supports STT (audio → text) and TTS (text → audio) via Gemini multimodal - models on GCP. Text-only completions are routed through the standard - `google` provider. + Supports TEXT (prompt → text), STT (audio → text) and TTS (text → audio) + via Gemini multimodal models on GCP. """ def __init__(self, client: GoogleGCPClient): @@ -655,6 +656,110 @@ def _execute_tts( ) return llm_response, None + @staticmethod + def _format_parts_rest(parts: list[ContentPart]) -> list[dict[str, Any]]: + """Map resolved content parts to REST generateContent ``parts`` (camelCase).""" + items: list[dict[str, Any]] = [] + 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": {"data": part.value, "mimeType": part.mime_type}} + ) + else: + items.append( + { + "fileData": { + "fileUri": part.value, + "mimeType": part.mime_type, + } + } + ) + 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. + + HTTP / network errors return pre-logged from ``_post()``; this method + only handles payload building and response-shape validation. + """ + provider = completion_config.provider + params = completion_config.params + model = params.get("model") or DEFAULT_TEXT_MODELS["google-gcp"] + + if isinstance(resolved_input, MultiModalInput): + parts = self._format_parts_rest(resolved_input.parts) + elif isinstance(resolved_input, list): + parts = self._format_parts_rest(resolved_input) + else: + parts = [{"text": resolved_input}] + + instructions = params.get("instructions") + temperature = params.get("temperature") + max_output_tokens = params.get("max_output_tokens") + + generation_config: dict[str, Any] = {} + if temperature is not None: + generation_config["temperature"] = temperature + if max_output_tokens is not None: + generation_config["maxOutputTokens"] = max_output_tokens + + payload: dict[str, Any] = {"contents": [{"role": "user", "parts": parts}]} + if generation_config: + payload["generationConfig"] = generation_config + if instructions: + payload["systemInstruction"] = {"parts": [{"text": instructions}]} + + 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 by token limits. 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 | " + f"provider={provider}, model={model}" + ) + return llm_response, None + def execute( self, completion_config: NativeCompletionConfig, @@ -665,6 +770,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, @@ -679,8 +790,8 @@ def execute( ) error_message = ( f"[KAAPI] Unsupported completion type '{completion_type}' for " - f"google-gcp provider. Google GCP supports 'stt' and 'tts' only; " - f"use the 'google-aistudio' provider for text completions." + f"google-gcp provider. Google GCP supports 'text', 'stt', and " + f"'tts' only." ) logger.warning( f"[GoogleGCPProvider.execute] {error_message} | provider={provider}" diff --git a/backend/app/tests/services/buckets/test_attachments.py b/backend/app/tests/services/buckets/test_attachments.py index 6b9f36590..913e5ddde 100644 --- a/backend/app/tests/services/buckets/test_attachments.py +++ b/backend/app/tests/services/buckets/test_attachments.py @@ -109,7 +109,7 @@ def test_mixed_schemes_partitioned(self): "gs://b/2.png": "https://signed2", } provider.get_bulk_signed_urls.assert_called_once_with( - ["gs://b/2.png"], expires_in=3600 + ["gs://b/2.png"], expires_in=None ) def test_all_native_skips_provider(self): 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 f69073578..edeb03568 100644 --- a/backend/app/tests/services/llm/providers/test_google_gcp.py +++ b/backend/app/tests/services/llm/providers/test_google_gcp.py @@ -295,15 +295,23 @@ 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", + provider="google-gcp-native", type=CompletionType.TEXT, - params={"model": "gemini-2.5-flash"}, + params={"model": "gemini-2.5-flash", "temperature": 0.2}, ) - 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("generated answer")), + ) as mock_post: + resp, err = provider.execute(config, query, "hello") + + assert err is None + assert resp.response.output.content.value == "generated answer" + payload = mock_post.call_args.kwargs["json"] + assert payload["contents"][0]["parts"] == [{"text": "hello"}] + assert payload["generationConfig"]["temperature"] == 0.2 def test_raw_response_included_when_requested( self, provider, stt_config, query, audio_ref @@ -561,16 +569,22 @@ 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_forwards_system_instruction(): provider = _provider() config = NativeCompletionConfig( - provider="google-native", + provider="google-gcp-native", type=CompletionType.TEXT, - params={"model": "gemini-2.5-flash"}, + params={"model": "gemini-2.5-flash", "instructions": "be terse"}, ) - 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("ok")), + ) as mock_post: + resp, err = provider.execute(config, QueryParams(input="ignored"), "hi") + + assert err is None + payload = mock_post.call_args.kwargs["json"] + assert payload["systemInstruction"] == {"parts": [{"text": "be terse"}]} def test_execute_tts_language_is_forwarded(): diff --git a/docs/wiki/modules/assessment.md b/docs/wiki/modules/assessment.md index 97047be65..092d60a3a 100644 --- a/docs/wiki/modules/assessment.md +++ b/docs/wiki/modules/assessment.md @@ -33,4 +33,4 @@ Config version (tag=ASSESSMENT, `models/config/assessment_blob.py`) owns system ## External - Provider Batch APIs, object storage for attachments (incl. `gs://` attachments resolved via `services/buckets/`). -- Gemini-family batch provider is chosen inline in `api/batch.py` (`_submit_provider_batch` / `_build_batch_provider`): `google-gcp` -> `VertexBatchProvider` (Vertex, GCS in/out, `core/batch/vertex.py`, built from the `google-gcp` credential), `google` -> `GeminiBatchProvider` (AI-Studio, File API, via `GeminiClient`). +- Gemini-family batch provider is chosen inline in `api/batch.py` (`_submit_provider_batch` / `_build_batch_provider`): `google-gcp` -> `VertexBatchProvider` (Vertex, GCS in/out, `core/batch/google_gcp.py`, built from the `google-gcp` credential), `google` -> `GeminiBatchProvider` (AI-Studio, File API, via `GeminiClient`). diff --git a/docs/wiki/modules/platform.md b/docs/wiki/modules/platform.md index 755c4769b..70b17c84b 100644 --- a/docs/wiki/modules/platform.md +++ b/docs/wiki/modules/platform.md @@ -12,6 +12,6 @@ All paths relative to `backend/app/`. | Languages | `api/routes/languages.py` | `global.languages` (`models/language.py`) | `crud/language.py` | | Credentials | `api/routes/credentials.py` | `credential` (`models/credentials.py`) | `crud/credentials.py`; provider keys per org/project; envelope encryption (KMS-wrapped data key + AES-GCM), prefix-versioned ciphertexts | | Model config | `api/routes/model_config.py` | `model_config` (`models/model_config.py`) | `crud/model_config.py` | -| Bucket providers | — | reuses `credential` (`google-gcp`) | `services/buckets/` — global registry + resolver (`providers/`), GCS signed/bulk-signed/private-to-public URLs (`providers/gcs.py`), attachment path selection + URL resolution (`attachments.py`) | +| Bucket providers | — | reuses `credential` (`google-gcp`) | `services/buckets/` — global registry + resolver (`providers/`), GCS V4 signed + bulk-signed URLs with a 24h expiry cap (`providers/gcs.py`, cap in `providers/base.py`), attachment path selection + URL resolution (`attachments.py`) | | Cron | `api/routes/cron.py` | — | triggers batch polling (`crud/evaluations/cron.py`) | | Jobs | — | `job` (`models/job.py`), `batch_job` (`models/batch_job.py`) | `crud/jobs.py`, `crud/job/`, `services/job_monitoring.py` | From 01b3443209e518c7838027538c031248e0b9eef9 Mon Sep 17 00:00:00 2001 From: Prashant Vasudevan <71649489+vprashrex@users.noreply.github.com> Date: Mon, 24 Aug 2026 18:36:09 +0530 Subject: [PATCH 08/15] feat(gcp): Rename VertexBatchProvider to GoogleGCPBatchProvider for clarity and consistency --- backend/app/core/batch/__init__.py | 4 ++-- backend/app/core/batch/google_gcp.py | 4 ++-- backend/app/services/assessment/api/batch.py | 6 +++--- backend/app/tests/assessment/test_api_batch.py | 4 ++-- backend/app/tests/core/batch/test_google_gcp.py | 12 ++++++------ docs/wiki/modules/assessment.md | 2 +- 6 files changed, 16 insertions(+), 16 deletions(-) diff --git a/backend/app/core/batch/__init__.py b/backend/app/core/batch/__init__.py index c35948171..666bf5af0 100644 --- a/backend/app/core/batch/__init__.py +++ b/backend/app/core/batch/__init__.py @@ -11,7 +11,7 @@ extract_text_from_response_dict, ) from .openai import OpenAIBatchProvider -from .google_gcp import VertexBatchProvider +from .google_gcp import GoogleGCPBatchProvider from .operations import ( download_batch_results, process_completed_batch, @@ -29,7 +29,7 @@ "GeminiClient", "GeminiClientError", "GeminiBatchProvider", - "VertexBatchProvider", + "GoogleGCPBatchProvider", "OpenAIBatchProvider", "create_stt_batch_requests", "create_tts_batch_requests", diff --git a/backend/app/core/batch/google_gcp.py b/backend/app/core/batch/google_gcp.py index 6be7ce900..04184b9c1 100644 --- a/backend/app/core/batch/google_gcp.py +++ b/backend/app/core/batch/google_gcp.py @@ -49,7 +49,7 @@ def _parse_gs_uri(uri: str) -> tuple[str, str]: return bucket, key -class VertexBatchProvider(BatchProvider): +class GoogleGCPBatchProvider(BatchProvider): """Vertex AI implementation of the BatchProvider interface (GCS in/out). Each JSONL line is the Vertex request schema, e.g. @@ -77,7 +77,7 @@ def __init__( @classmethod def from_credentials( cls, credentials: dict[str, Any], model: str | None = None - ) -> "VertexBatchProvider": + ) -> "GoogleGCPBatchProvider": """Build a Vertex batch provider from a ``google-gcp`` credential dict.""" project_id = credentials.get("project_id") location = credentials.get("location") diff --git a/backend/app/services/assessment/api/batch.py b/backend/app/services/assessment/api/batch.py index d31f7659e..5fdb90be1 100644 --- a/backend/app/services/assessment/api/batch.py +++ b/backend/app/services/assessment/api/batch.py @@ -28,7 +28,7 @@ MessageBatchStatus, OpenAIBatchProvider, GeminiBatchProvider, - VertexBatchProvider, + GoogleGCPBatchProvider, extract_text_from_response_dict, poll_batch_status, process_completed_batch, @@ -350,7 +350,7 @@ def _submit_provider_batch( organization_id=organization_id, project_id=project_id, ) - provider = VertexBatchProvider.from_credentials(cred, model=model) + provider = GoogleGCPBatchProvider.from_credentials(cred, model=model) config = {"display_name": description} # Vertex uses a bare model id else: gemini = GeminiClient.from_credentials( @@ -414,7 +414,7 @@ def _build_batch_provider( organization_id=organization_id, project_id=project_id, ) - return VertexBatchProvider.from_credentials(cred) + return GoogleGCPBatchProvider.from_credentials(cred) gemini = GeminiClient.from_credentials( session=session, org_id=organization_id, project_id=project_id ) diff --git a/backend/app/tests/assessment/test_api_batch.py b/backend/app/tests/assessment/test_api_batch.py index e6db26017..4aeeb8c1a 100644 --- a/backend/app/tests/assessment/test_api_batch.py +++ b/backend/app/tests/assessment/test_api_batch.py @@ -889,7 +889,7 @@ def test_google_gcp_routes_to_vertex(self, db) -> None: return_value={"gcs_bucket": "b", "sa_key": {}}, ), patch( - "app.services.assessment.api.batch.VertexBatchProvider.from_credentials", + "app.services.assessment.api.batch.GoogleGCPBatchProvider.from_credentials", return_value=sentinel, ) as vertex_from_cred, ): @@ -990,7 +990,7 @@ def test_google_gcp_branch_routes_to_vertex(self, db) -> None: return_value={"gcs_bucket": "b", "sa_key": {}}, ), patch( - "app.services.assessment.api.batch.VertexBatchProvider.from_credentials", + "app.services.assessment.api.batch.GoogleGCPBatchProvider.from_credentials", return_value=MagicMock(), ) as vertex_from_cred, patch( diff --git a/backend/app/tests/core/batch/test_google_gcp.py b/backend/app/tests/core/batch/test_google_gcp.py index 183bdf6dd..36585778c 100644 --- a/backend/app/tests/core/batch/test_google_gcp.py +++ b/backend/app/tests/core/batch/test_google_gcp.py @@ -1,11 +1,11 @@ -"""Test cases for VertexBatchProvider (Vertex AI batch, GCS-backed).""" +"""Test cases for GoogleGCPBatchProvider (Vertex AI batch, GCS-backed).""" import json from unittest.mock import MagicMock, patch import pytest -from app.core.batch.google_gcp import VertexBatchProvider, _parse_gs_uri +from app.core.batch.google_gcp import GoogleGCPBatchProvider, _parse_gs_uri _BUCKET = "test-bucket" @@ -22,7 +22,7 @@ def mock_storage(): @pytest.fixture def provider(mock_genai, mock_storage): - return VertexBatchProvider( + return GoogleGCPBatchProvider( client=mock_genai, storage_client=mock_storage, gcs_bucket=_BUCKET, @@ -172,8 +172,8 @@ def test_builds_provider(self): patch("app.core.batch.google_gcp.genai.Client") as genai_client, patch("app.core.batch.google_gcp.gcs.Client") as gcs_client, ): - provider = VertexBatchProvider.from_credentials(self._CRED) - assert isinstance(provider, VertexBatchProvider) + provider = GoogleGCPBatchProvider.from_credentials(self._CRED) + assert isinstance(provider, GoogleGCPBatchProvider) sa.Credentials.from_service_account_info.assert_called_once() assert genai_client.call_args.kwargs["vertexai"] is True gcs_client.assert_called_once() @@ -182,4 +182,4 @@ def test_builds_provider(self): def test_missing_field_raises(self, drop): cred = {k: v for k, v in self._CRED.items() if k != drop} with pytest.raises(ValueError, match=drop): - VertexBatchProvider.from_credentials(cred) + GoogleGCPBatchProvider.from_credentials(cred) diff --git a/docs/wiki/modules/assessment.md b/docs/wiki/modules/assessment.md index 092d60a3a..78122b862 100644 --- a/docs/wiki/modules/assessment.md +++ b/docs/wiki/modules/assessment.md @@ -33,4 +33,4 @@ Config version (tag=ASSESSMENT, `models/config/assessment_blob.py`) owns system ## External - Provider Batch APIs, object storage for attachments (incl. `gs://` attachments resolved via `services/buckets/`). -- Gemini-family batch provider is chosen inline in `api/batch.py` (`_submit_provider_batch` / `_build_batch_provider`): `google-gcp` -> `VertexBatchProvider` (Vertex, GCS in/out, `core/batch/google_gcp.py`, built from the `google-gcp` credential), `google` -> `GeminiBatchProvider` (AI-Studio, File API, via `GeminiClient`). +- Gemini-family batch provider is chosen inline in `api/batch.py` (`_submit_provider_batch` / `_build_batch_provider`): `google-gcp` -> `GoogleGCPBatchProvider` (Vertex, GCS in/out, `core/batch/google_gcp.py`, built from the `google-gcp` credential), `google` -> `GeminiBatchProvider` (AI-Studio, File API, via `GeminiClient`). From 8a1b6673a3738dee456e53f393a38a51db9c7a88 Mon Sep 17 00:00:00 2001 From: Prashant Vasudevan <71649489+vprashrex@users.noreply.github.com> Date: Tue, 25 Aug 2026 11:43:03 +0530 Subject: [PATCH 09/15] feat(buckets): configurable signed-URL TTL and BYOK-only GCS provider - Single MAX_SIGNED_URL_EXPIRY_SECONDS (24h) setting so attachment URLs outlive a long batch run; drop the 1h default - expires_in is now a required int end to end (no None), validated >= 1 - GCS bucket provider is BYOK-only: require sa_key + gcs_bucket from the supplied credentials, no platform SA / GCS_AUDIO_BUCKET fallback Co-Authored-By: Claude Opus 4.8 --- backend/app/core/config.py | 4 +- backend/app/services/buckets/attachments.py | 2 +- .../app/services/buckets/providers/base.py | 6 +-- backend/app/services/buckets/providers/gcs.py | 38 ++++++-------- .../services/buckets/test_attachments.py | 6 ++- .../app/tests/services/buckets/test_gcs.py | 50 ++++--------------- 6 files changed, 34 insertions(+), 72 deletions(-) diff --git a/backend/app/core/config.py b/backend/app/core/config.py index baccbb43a..367851da2 100644 --- a/backend/app/core/config.py +++ b/backend/app/core/config.py @@ -117,9 +117,7 @@ def SQLALCHEMY_DATABASE_URI(self) -> PostgresDsn: # Used by the registry fallback when a project has no ``google`` row. GCP_SA_KEY: str = "" GCS_AUDIO_BUCKET: str = "" - # Signed-URL lifetime for private bucket attachments (Path B); default 1h. - SIGNED_URL_EXPIRY_SECONDS: int = 3600 - # Hard cap on signed-URL lifetime (24h) enforced by bucket providers. + # A batch can run for hours; sign attachment URLs for 24h so they don't expire mid-run. MAX_SIGNED_URL_EXPIRY_SECONDS: int = 86400 # RabbitMQ configuration for Celery broker diff --git a/backend/app/services/buckets/attachments.py b/backend/app/services/buckets/attachments.py index 2e5b60bb1..7746ca1c6 100644 --- a/backend/app/services/buckets/attachments.py +++ b/backend/app/services/buckets/attachments.py @@ -44,7 +44,7 @@ def resolve_attachments( llm_provider: KaapiProvider, project_id: int, organization_id: int, - expires_in: int | None = None, + expires_in: int, bucket_provider_type: str = DEFAULT_BUCKET_PROVIDER, ) -> str | dict[str, str]: """Make attachment(s) LLM-reachable. Accepts a single URL or a list. diff --git a/backend/app/services/buckets/providers/base.py b/backend/app/services/buckets/providers/base.py index d5894fecf..094e8dcb9 100644 --- a/backend/app/services/buckets/providers/base.py +++ b/backend/app/services/buckets/providers/base.py @@ -28,13 +28,11 @@ def create_client(credentials: dict[str, Any]) -> Any: raise NotImplementedError("Bucket providers must implement create_client") @abstractmethod - def get_signed_url(self, uri: str, expires_in: int | None = None) -> str: + def get_signed_url(self, uri: str, expires_in: int) -> str: """Generate a time-limited signed URL for a single object.""" raise NotImplementedError("Bucket providers must implement get_signed_url") - def get_bulk_signed_urls( - self, uris: list[str], expires_in: int | None = None - ) -> dict[str, str]: + def get_bulk_signed_urls(self, uris: list[str], expires_in: int) -> dict[str, str]: """Sign many object URIs, reusing this provider's single client.""" signed: dict[str, str] = {} for uri in uris: diff --git a/backend/app/services/buckets/providers/gcs.py b/backend/app/services/buckets/providers/gcs.py index d4e08f2fd..fa853e519 100644 --- a/backend/app/services/buckets/providers/gcs.py +++ b/backend/app/services/buckets/providers/gcs.py @@ -8,10 +8,8 @@ from google.cloud import storage as gcs from google.oauth2 import service_account -from app.core.config import settings from app.core.cloud.storage import GCS_SCOPES, CloudStorageError from app.services.buckets.providers.base import BaseBucketProvider -from app.services.llm.providers.google_gcp import _load_platform_sa_info logger = logging.getLogger(__name__) @@ -37,24 +35,25 @@ def __init__(self, client: GCSClient): @staticmethod def create_client(credentials: dict[str, Any]) -> GCSClient: - """Build a signing-capable GCS client with BYOK-over-settings precedence.""" + """Build a signing-capable GCS client from the given credentials.""" credentials = credentials or {} - gcs_bucket = credentials.get("gcs_bucket") or settings.GCS_AUDIO_BUCKET - sa_info = credentials.get("sa_key") or _load_platform_sa_info() + sa_info = credentials.get("sa_key") + gcs_bucket = credentials.get("gcs_bucket") - source = "byok" if credentials.get("sa_key") else "platform" - logger.info( - f"[GCSBucketProvider.create_client] gcs creds | source={source}, " - f"bucket={gcs_bucket}" - ) - - # Signing needs the SA private key, so a missing sa_info is fatal. if not sa_info: raise ValueError( - "GCS bucket provider requires a service-account key (sa_key) to " - "sign URLs; none configured for this project or platform default." + "GCS bucket provider requires a service-account key (sa_key) in " + "the credentials to sign URLs." + ) + if not gcs_bucket: + raise ValueError( + "GCS bucket provider requires a gcs_bucket in the credentials." ) + logger.info( + f"[GCSBucketProvider.create_client] gcs creds | bucket={gcs_bucket}" + ) + creds = service_account.Credentials.from_service_account_info( sa_info, scopes=list(GCS_SCOPES) ) @@ -74,9 +73,7 @@ def _parse_gcs_uri(uri: str) -> tuple[str, str]: raise ValueError(f"GCS URI '{uri}' is missing an object key.") return parsed.netloc, key - def get_signed_url(self, uri: str, expires_in: int | None = None) -> str: - if expires_in is None: - expires_in = settings.SIGNED_URL_EXPIRY_SECONDS + def get_signed_url(self, uri: str, expires_in: int) -> str: if expires_in < 1: raise ValueError( f"Signed-URL expiry must be at least 1 second, got {expires_in}." @@ -105,13 +102,8 @@ def get_signed_url(self, uri: str, expires_in: int | None = None) -> str: ) return signed_url - def get_bulk_signed_urls( - self, uris: list[str], expires_in: int | None = None - ) -> dict[str, str]: + def get_bulk_signed_urls(self, uris: list[str], expires_in: int) -> dict[str, str]: """Sign each URI reusing this provider's single client.""" - if expires_in is None: - expires_in = settings.SIGNED_URL_EXPIRY_SECONDS - logger.info( f"[GCSBucketProvider.get_bulk_signed_urls] Signing batch | " f"count={len(uris)}, expires_in={expires_in}" diff --git a/backend/app/tests/services/buckets/test_attachments.py b/backend/app/tests/services/buckets/test_attachments.py index 913e5ddde..7a9aa2ab6 100644 --- a/backend/app/tests/services/buckets/test_attachments.py +++ b/backend/app/tests/services/buckets/test_attachments.py @@ -56,6 +56,7 @@ def test_https_returned_as_is_without_provider(self): llm_provider="openai", project_id=1, organization_id=2, + expires_in=86400, ) assert url == "https://example.com/img.png" mock_get.assert_not_called() @@ -68,6 +69,7 @@ def test_gcs_native_passthrough(self): llm_provider="google-gcp", project_id=1, organization_id=2, + expires_in=86400, ) assert url == "gs://bucket/key.png" mock_get.assert_not_called() @@ -103,13 +105,14 @@ def test_mixed_schemes_partitioned(self): llm_provider="anthropic", project_id=1, organization_id=2, + expires_in=86400, ) assert result == { "https://x/1.png": "https://x/1.png", "gs://b/2.png": "https://signed2", } provider.get_bulk_signed_urls.assert_called_once_with( - ["gs://b/2.png"], expires_in=None + ["gs://b/2.png"], expires_in=86400 ) def test_all_native_skips_provider(self): @@ -120,6 +123,7 @@ def test_all_native_skips_provider(self): llm_provider="google-gcp", project_id=1, organization_id=2, + expires_in=86400, ) assert result == { "gs://b/1.wav": "gs://b/1.wav", diff --git a/backend/app/tests/services/buckets/test_gcs.py b/backend/app/tests/services/buckets/test_gcs.py index 078b432da..01b4d472a 100644 --- a/backend/app/tests/services/buckets/test_gcs.py +++ b/backend/app/tests/services/buckets/test_gcs.py @@ -24,18 +24,10 @@ def _make_provider() -> tuple[GCSBucketProvider, MagicMock]: class TestCreateClient: - def test_byok_credentials_win_over_settings(self, monkeypatch): - monkeypatch.setattr( - "app.services.buckets.providers.gcs.settings.GCS_AUDIO_BUCKET", - "platform-bucket", - ) + def test_uses_byok_credentials(self): byok_sa = {"project_id": "byok-project"} with ( - patch( - "app.services.buckets.providers.gcs._load_platform_sa_info", - return_value={"project_id": "platform-project"}, - ), patch( "app.services.buckets.providers.gcs.service_account." "Credentials.from_service_account_info" @@ -52,40 +44,18 @@ def test_byok_credentials_win_over_settings(self, monkeypatch): project="byok-project", credentials=mock_from_info.return_value ) - def test_falls_back_to_settings_and_platform_sa(self, monkeypatch): - monkeypatch.setattr( - "app.services.buckets.providers.gcs.settings.GCS_AUDIO_BUCKET", - "platform-bucket", - ) - - with ( - patch( - "app.services.buckets.providers.gcs._load_platform_sa_info", - return_value={"project_id": "platform-project"}, - ), - patch( - "app.services.buckets.providers.gcs.service_account." - "Credentials.from_service_account_info" - ) as mock_from_info, - patch("app.services.buckets.providers.gcs.gcs.Client"), - ): - client = GCSBucketProvider.create_client({}) - - assert client.default_bucket == "platform-bucket" - mock_from_info.assert_called_once_with( - {"project_id": "platform-project"}, scopes=list(GCS_SCOPES) - ) - def test_missing_sa_info_raises(self): - with patch( - "app.services.buckets.providers.gcs._load_platform_sa_info", - return_value=None, - ): - with pytest.raises(ValueError) as exc_info: - GCSBucketProvider.create_client({}) + with pytest.raises(ValueError) as exc_info: + GCSBucketProvider.create_client({"gcs_bucket": "b"}) assert "sa_key" in str(exc_info.value) + def test_missing_bucket_raises(self): + with pytest.raises(ValueError) as exc_info: + GCSBucketProvider.create_client({"sa_key": {"project_id": "p"}}) + + assert "gcs_bucket" in str(exc_info.value) + class TestParseGcsUri: def test_parses_bucket_and_key(self): @@ -142,7 +112,7 @@ def test_returns_uri_to_url_map_reusing_one_client(self): ] result = provider.get_bulk_signed_urls( - ["gs://bucket/a.wav", "gs://bucket/b.wav"] + ["gs://bucket/a.wav", "gs://bucket/b.wav"], expires_in=86400 ) assert result == { From edfa349c836216e2e5d6aa76fa8e9114cc71a3b1 Mon Sep 17 00:00:00 2001 From: Prajna1999 Date: Wed, 26 Aug 2026 16:44:58 +0530 Subject: [PATCH 10/15] fix: add models for GCPProvider, partial updates --- backend/app/core/providers.py | 17 ++++++++++------- backend/app/crud/credentials.py | 13 ++----------- backend/app/models/credentials.py | 18 ++++++++++++++---- 3 files changed, 26 insertions(+), 22 deletions(-) diff --git a/backend/app/core/providers.py b/backend/app/core/providers.py index 06e0f69ff..e3fe27f4d 100644 --- a/backend/app/core/providers.py +++ b/backend/app/core/providers.py @@ -69,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" @@ -87,6 +95,7 @@ class ProxyCredentials(ProviderCredentialsBase): | ElevenLabsCredentials | AnthropicCredentials | GoogleCredentials + | GoogleGcpCredentials | WebhookSecretCredentials | ProxyCredentials, Field( @@ -143,13 +152,7 @@ def required_fields(self) -> list[str]: sensitive_fields=["api_key"], ), Provider.GOOGLE_GCP: ProviderConfig( - required_fields=[ - "api_key", - "project_id", - "location", - "sa_key", - "gcs_bucket", - ], + model=GoogleGcpCredentials, sensitive_fields=["api_key", "sa_key"], ), Provider.WEBHOOK_SECRET: ProviderConfig( diff --git a/backend/app/crud/credentials.py b/backend/app/crud/credentials.py index cc5f5b8db..3f0b4747c 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 @@ -181,15 +181,6 @@ def update_creds_for_org( if not creds_in.provider or not creds_in.credential: raise ValueError("Provider and credential must be provided") - # Auto-unwrap nested format: {"google": {"api_key": "..."}} -> {"api_key": "..."} - # so the same payload shape works for both create and update. - credential_data = creds_in.credential - if ( - isinstance(credential_data, dict) - and creds_in.provider in credential_data - and isinstance(credential_data[creds_in.provider], dict) - ): - credential_data = credential_data[creds_in.provider] provider = creds_in.provider.value credential_data = creds_in.credential_payload() @@ -214,7 +205,7 @@ def update_creds_for_org( } try: - validate_provider_credentials(creds_in.provider, merged_credential_data) + parse_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)}" 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): From 600379e1da247e2d8d2d48777b5dab96ad881a64 Mon Sep 17 00:00:00 2001 From: Prajna1999 Date: Wed, 26 Aug 2026 17:47:01 +0530 Subject: [PATCH 11/15] codecov and remove redundant constant variable --- backend/app/models/llm/constants.py | 10 +- backend/app/models/llm/request.py | 8 +- .../app/services/llm/providers/google_gcp.py | 117 +++++++++++++++++- .../services/llm/providers/test_google_gcp.py | 24 ++-- .../app/tests/services/llm/test_mappers.py | 40 ++++++ 5 files changed, 173 insertions(+), 26 deletions(-) diff --git a/backend/app/models/llm/constants.py b/backend/app/models/llm/constants.py index c903d5ab2..f6e927bfc 100644 --- a/backend/app/models/llm/constants.py +++ b/backend/app/models/llm/constants.py @@ -32,18 +32,12 @@ class Provider(StrEnum): ] RAGProvider = Literal[Provider.OPENAI, Provider.GOOGLE_AISTUDIO] -KaapiProvider = Literal[ +TextProvider = Literal[ Provider.OPENAI, Provider.GOOGLE, - Provider.GOOGLE_GCP, - Provider.SARVAMAI, - Provider.ELEVENLABS, Provider.ANTHROPIC, Provider.GOOGLE_AISTUDIO, -] - -TextProvider = Literal[ - Provider.OPENAI, Provider.GOOGLE, Provider.ANTHROPIC, Provider.GOOGLE_AISTUDIO + Provider.GOOGLE_GCP, ] KaapiProvider = Union[TextProvider, STTProvider, TTSProvider] diff --git a/backend/app/models/llm/request.py b/backend/app/models/llm/request.py index 89b09d5bd..c028a978a 100644 --- a/backend/app/models/llm/request.py +++ b/backend/app/models/llm/request.py @@ -342,11 +342,9 @@ class KaapiTextCompletionConfig(SQLModel): provider: TextProvider | None = Field( default=None, description=( - "LLM provider (openai, google, sarvamai, elevenlabs, anthropic, " - "google-aistudio, google-gcp). 'google-aistudio' uses Google AI " - "Studio; 'google-gcp' uses Google GCP directly for STT/TTS." - "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/services/llm/providers/google_gcp.py b/backend/app/services/llm/providers/google_gcp.py index 1221e3ebb..d486c1ab2 100644 --- a/backend/app/services/llm/providers/google_gcp.py +++ b/backend/app/services/llm/providers/google_gcp.py @@ -15,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, @@ -25,6 +27,7 @@ ) from app.models.llm.constants import ( DEFAULT_STT_MODEL, + DEFAULT_TEXT_MODELS, DEFAULT_TTS_MODEL, DEFAULT_TTS_VOICE, CompletionType, @@ -106,9 +109,8 @@ def endpoint(self, model: str) -> str: class GoogleGCPProvider(BaseProvider): """Google GCP provider using REST + API key auth. - Supports STT (audio → text) and TTS (text → audio) via Gemini multimodal - models on GCP. 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: GoogleGCPClient): @@ -310,6 +312,106 @@ 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, @@ -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,8 +791,7 @@ def execute( ) error_message = ( f"[KAAPI] Unsupported completion type '{completion_type}' for " - f"google-gcp provider. Google GCP 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"[GoogleGCPProvider.execute] {error_message} | provider={provider}" 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 f69073578..bc7ccb4c0 100644 --- a/backend/app/tests/services/llm/providers/test_google_gcp.py +++ b/backend/app/tests/services/llm/providers/test_google_gcp.py @@ -295,15 +295,19 @@ 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 @@ -561,16 +565,20 @@ 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(): 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.""" From 7d4e13580ba90c5e29ba5c140355a149c76e2268 Mon Sep 17 00:00:00 2001 From: Prashant Vasudevan <71649489+vprashrex@users.noreply.github.com> Date: Wed, 26 Aug 2026 19:31:58 +0530 Subject: [PATCH 12/15] feat: refactor Google GCP provider integration and credential handling --- backend/app/core/batch/google_gcp.py | 45 ++++++-------- backend/app/core/cloud/storage.py | 13 +++- backend/app/services/assessment/api/batch.py | 59 +++++++++++-------- backend/app/services/assessment/service.py | 4 ++ backend/app/services/assessment/stages.py | 21 ++++++- backend/app/services/buckets/providers/gcs.py | 7 +-- .../app/tests/core/batch/test_google_gcp.py | 5 +- .../app/tests/services/buckets/test_gcs.py | 10 ++-- .../tests/services/buckets/test_registry.py | 5 +- 9 files changed, 98 insertions(+), 71 deletions(-) diff --git a/backend/app/core/batch/google_gcp.py b/backend/app/core/batch/google_gcp.py index 04184b9c1..51310a15c 100644 --- a/backend/app/core/batch/google_gcp.py +++ b/backend/app/core/batch/google_gcp.py @@ -8,15 +8,19 @@ import json import logging import time -from typing import Any +from typing import Any, cast from uuid import uuid4 from google import genai from google.genai import types from google.cloud import storage as gcs -from google.oauth2 import service_account -from app.core.cloud.storage import GCS_SCOPES, CloudStorageError +from app.core.cloud.storage import CloudStorageError, build_gcp_sa_credentials +from app.core.providers import ( + GoogleGcpCredentials, + Provider, + parse_provider_credentials, +) from .base import BATCH_KEY, BatchProvider from .gemini import BatchJobState @@ -79,36 +83,23 @@ def from_credentials( cls, credentials: dict[str, Any], model: str | None = None ) -> "GoogleGCPBatchProvider": """Build a Vertex batch provider from a ``google-gcp`` credential dict.""" - project_id = credentials.get("project_id") - location = credentials.get("location") - gcs_bucket = credentials.get("gcs_bucket") - sa_info = credentials.get("sa_key") - missing = [ - name - for name, value in ( - ("project_id", project_id), - ("location", location), - ("gcs_bucket", gcs_bucket), - ("sa_key", sa_info), - ) - if not value - ] - if missing: - raise ValueError( - f"Vertex batch provider missing required fields: {', '.join(missing)}" - ) - - creds = service_account.Credentials.from_service_account_info( - sa_info, scopes=list(GCS_SCOPES) + creds_model = cast( + GoogleGcpCredentials, + parse_provider_credentials(Provider.GOOGLE_GCP, credentials), ) + + creds = build_gcp_sa_credentials(cast(dict[str, Any], creds_model.sa_key)) client = genai.Client( - vertexai=True, project=project_id, location=location, credentials=creds + vertexai=True, + project=creds_model.project_id, + location=creds_model.location, + credentials=creds, ) - storage_client = gcs.Client(project=project_id, credentials=creds) + storage_client = gcs.Client(project=creds_model.project_id, credentials=creds) return cls( client=client, storage_client=storage_client, - gcs_bucket=gcs_bucket, + gcs_bucket=creds_model.gcs_bucket, model=model, ) diff --git a/backend/app/core/cloud/storage.py b/backend/app/core/cloud/storage.py index f1000f6f6..a0b451edd 100644 --- a/backend/app/core/cloud/storage.py +++ b/backend/app/core/cloud/storage.py @@ -10,6 +10,7 @@ from urllib.parse import ParseResult, urlparse, urlunparse from abc import ABC, abstractmethod +from typing import Any import boto3 from fastapi import UploadFile from botocore.exceptions import ClientError @@ -329,6 +330,14 @@ def get_cloud_storage(session: Session, project_id: int) -> CloudStorage: GCS_SCOPES = ("https://www.googleapis.com/auth/cloud-platform",) + +def build_gcp_sa_credentials(sa_key: dict[str, Any]) -> service_account.Credentials: + """Build signing-capable SA credentials from a service-account key dict.""" + return service_account.Credentials.from_service_account_info( + sa_key, scopes=list(GCS_SCOPES) + ) + + MAX_AUDIO_UPLOAD_BYTES = 50 * 1024 * 1024 # 50 MB _MIME_TO_EXT = { @@ -400,9 +409,7 @@ def upload_audio_to_gcs( key = f"{key_prefix}/{uuid4().hex}{ext}" try: - creds = service_account.Credentials.from_service_account_info( - sa_info, scopes=list(GCS_SCOPES) - ) + creds = build_gcp_sa_credentials(sa_info) client = gcs.Client( project=project_id or sa_info.get("project_id"), credentials=creds ) diff --git a/backend/app/services/assessment/api/batch.py b/backend/app/services/assessment/api/batch.py index 1159a99a4..2942516bc 100644 --- a/backend/app/services/assessment/api/batch.py +++ b/backend/app/services/assessment/api/batch.py @@ -25,10 +25,10 @@ BATCH_KEY, AnthropicBatchProvider, BatchJobState, + GoogleGCPBatchProvider, MessageBatchStatus, OpenAIBatchProvider, GeminiBatchProvider, - GoogleGCPBatchProvider, extract_text_from_response_dict, poll_batch_status, process_completed_batch, @@ -36,11 +36,11 @@ ) from app.core.batch.base import BatchProvider from app.core.batch.client import GeminiClient -from app.crud.credentials import get_provider_credential from fastapi import HTTPException from app.core.config import settings from app.core.db import engine from app.crud.assessment import api +from app.crud.credentials import get_provider_credential from app.crud.assessment.batch import ( build_anthropic_jsonl, build_google_jsonl, @@ -72,10 +72,32 @@ map_kaapi_to_openai_params, ) from app.services.llm.providers.registry import LLMProvider -from app.utils import get_anthropic_client, get_openai_client +from app.utils import ( + get_anthropic_client, + get_openai_client, +) logger = logging.getLogger(__name__) + +def _google_gcp_credential( + *, session: Session, organization_id: int, project_id: int +) -> dict[str, Any]: + """Vertex needs the SA key + bucket; a missing credential is client-fixable.""" + cred = get_provider_credential( + session=session, + provider=LLMProvider.GOOGLE_GCP, + project_id=project_id, + org_id=organization_id, + ) + if not isinstance(cred, dict): + raise HTTPException( + status_code=404, + detail="google-gcp credentials not configured for this project", + ) + return cred + + # Re-poll cadence for a stage's provider batch, mirroring the assessment cron tick. POLL_COUNTDOWN_SECONDS = settings.CRON_INTERVAL_MINUTES * 60 @@ -128,6 +150,7 @@ class StageKind(StrEnum): _SUPPORTED_PROVIDERS = { LLMProvider.OPENAI, LLMProvider.GOOGLE, + LLMProvider.GOOGLE_AISTUDIO, LLMProvider.GOOGLE_GCP, LLMProvider.ANTHROPIC, } @@ -272,24 +295,6 @@ def _stage_provider_model(blob: AssessmentConfigBlob, stage: str) -> tuple[str, return blob.assessment.provider, blob.assessment.params["model"] -def _google_gcp_credential( - *, session: Session, organization_id: int, project_id: int -) -> dict[str, Any]: - """Vertex needs the SA key + bucket; a missing credential is a client-fixable 404.""" - cred = get_provider_credential( - session=session, - provider=LLMProvider.GOOGLE_GCP, - project_id=project_id, - org_id=organization_id, - ) - if not isinstance(cred, dict): - raise HTTPException( - status_code=404, - detail="google-gcp credentials not configured for this project", - ) - return cred - - def _submit_provider_batch( *, session: Session, @@ -331,7 +336,11 @@ def _submit_provider_batch( "description": description, "completion_window": "24h", } - elif provider_name in (LLMProvider.GOOGLE, LLMProvider.GOOGLE_GCP): + elif provider_name in ( + LLMProvider.GOOGLE, + LLMProvider.GOOGLE_AISTUDIO, + LLMProvider.GOOGLE_GCP, + ): mapped, _ = map_kaapi_to_google_params(params) jsonl = build_google_jsonl( rows, text_columns, attachments, prompt, mapped, row_indices @@ -399,7 +408,11 @@ def _build_batch_provider( session=session, org_id=organization_id, project_id=project_id ) ) - if provider_name in (LLMProvider.GOOGLE, LLMProvider.GOOGLE_GCP): + if provider_name in ( + LLMProvider.GOOGLE, + LLMProvider.GOOGLE_AISTUDIO, + LLMProvider.GOOGLE_GCP, + ): if provider_name == LLMProvider.GOOGLE_GCP: cred = _google_gcp_credential( session=session, diff --git a/backend/app/services/assessment/service.py b/backend/app/services/assessment/service.py index 6c2548f34..888a949bb 100644 --- a/backend/app/services/assessment/service.py +++ b/backend/app/services/assessment/service.py @@ -39,6 +39,10 @@ LLMProvider.OPENAI_NATIVE, LLMProvider.GOOGLE, LLMProvider.GOOGLE_NATIVE, + LLMProvider.GOOGLE_AISTUDIO, + LLMProvider.GOOGLE_AISTUDIO_NATIVE, + LLMProvider.GOOGLE_GCP, + LLMProvider.GOOGLE_GCP_NATIVE, LLMProvider.ANTHROPIC, LLMProvider.ANTHROPIC_NATIVE, } diff --git a/backend/app/services/assessment/stages.py b/backend/app/services/assessment/stages.py index 0ac03fa2f..0676ce05e 100644 --- a/backend/app/services/assessment/stages.py +++ b/backend/app/services/assessment/stages.py @@ -29,7 +29,10 @@ parse_topic_relevance_results, ) from app.services.llm.providers.registry import LLMProvider -from app.utils import get_anthropic_client, get_openai_client +from app.utils import ( + get_anthropic_client, + get_openai_client, +) logger = logging.getLogger(__name__) @@ -147,11 +150,27 @@ def _get_batch_provider( if provider_name in ( LLMProvider.GOOGLE, LLMProvider.GOOGLE_NATIVE, + LLMProvider.GOOGLE_AISTUDIO, + LLMProvider.GOOGLE_AISTUDIO_NATIVE, ): gemini_client = GeminiClient.from_credentials( session=session, org_id=organization_id, project_id=project_id ) return GeminiBatchProvider(client=gemini_client.client) + if provider_name in (LLMProvider.GOOGLE_GCP, LLMProvider.GOOGLE_GCP_NATIVE): + # Lazy to avoid a crud<->stages cycle. + from app.core.batch import GoogleGCPBatchProvider + from app.crud.credentials import get_provider_credential + + cred = get_provider_credential( + session=session, + provider=LLMProvider.GOOGLE_GCP, + project_id=project_id, + org_id=organization_id, + ) + if not isinstance(cred, dict): + raise ValueError("google-gcp credentials not configured for this project") + return GoogleGCPBatchProvider.from_credentials(cred) if provider_name in (LLMProvider.ANTHROPIC, LLMProvider.ANTHROPIC_NATIVE): return AnthropicBatchProvider( client=get_anthropic_client( diff --git a/backend/app/services/buckets/providers/gcs.py b/backend/app/services/buckets/providers/gcs.py index fa853e519..e59151e80 100644 --- a/backend/app/services/buckets/providers/gcs.py +++ b/backend/app/services/buckets/providers/gcs.py @@ -6,9 +6,8 @@ from urllib.parse import urlparse from google.cloud import storage as gcs -from google.oauth2 import service_account -from app.core.cloud.storage import GCS_SCOPES, CloudStorageError +from app.core.cloud.storage import CloudStorageError, build_gcp_sa_credentials from app.services.buckets.providers.base import BaseBucketProvider logger = logging.getLogger(__name__) @@ -54,9 +53,7 @@ def create_client(credentials: dict[str, Any]) -> GCSClient: f"[GCSBucketProvider.create_client] gcs creds | bucket={gcs_bucket}" ) - creds = service_account.Credentials.from_service_account_info( - sa_info, scopes=list(GCS_SCOPES) - ) + creds = build_gcp_sa_credentials(sa_info) storage_client = gcs.Client( project=sa_info.get("project_id"), credentials=creds ) diff --git a/backend/app/tests/core/batch/test_google_gcp.py b/backend/app/tests/core/batch/test_google_gcp.py index 36585778c..434d68a91 100644 --- a/backend/app/tests/core/batch/test_google_gcp.py +++ b/backend/app/tests/core/batch/test_google_gcp.py @@ -160,6 +160,7 @@ def test_download_file_reads_text(self, provider, mock_storage): class TestFromCredentials: _CRED = { + "api_key": "test-key", "project_id": "proj", "location": "us-central1", "gcs_bucket": _BUCKET, @@ -168,13 +169,13 @@ class TestFromCredentials: def test_builds_provider(self): with ( - patch("app.core.batch.google_gcp.service_account") as sa, + patch("app.core.batch.google_gcp.build_gcp_sa_credentials") as build_creds, patch("app.core.batch.google_gcp.genai.Client") as genai_client, patch("app.core.batch.google_gcp.gcs.Client") as gcs_client, ): provider = GoogleGCPBatchProvider.from_credentials(self._CRED) assert isinstance(provider, GoogleGCPBatchProvider) - sa.Credentials.from_service_account_info.assert_called_once() + build_creds.assert_called_once() assert genai_client.call_args.kwargs["vertexai"] is True gcs_client.assert_called_once() diff --git a/backend/app/tests/services/buckets/test_gcs.py b/backend/app/tests/services/buckets/test_gcs.py index 01b4d472a..67d933629 100644 --- a/backend/app/tests/services/buckets/test_gcs.py +++ b/backend/app/tests/services/buckets/test_gcs.py @@ -5,7 +5,6 @@ import pytest -from app.core.cloud.storage import GCS_SCOPES from app.services.buckets.providers.gcs import ( GCSBucketProvider, GCSClient, @@ -29,9 +28,8 @@ def test_uses_byok_credentials(self): with ( patch( - "app.services.buckets.providers.gcs.service_account." - "Credentials.from_service_account_info" - ) as mock_from_info, + "app.services.buckets.providers.gcs.build_gcp_sa_credentials" + ) as mock_build, patch("app.services.buckets.providers.gcs.gcs.Client") as mock_client, ): client = GCSBucketProvider.create_client( @@ -39,9 +37,9 @@ def test_uses_byok_credentials(self): ) assert client.default_bucket == "byok-bucket" - mock_from_info.assert_called_once_with(byok_sa, scopes=list(GCS_SCOPES)) + mock_build.assert_called_once_with(byok_sa) mock_client.assert_called_once_with( - project="byok-project", credentials=mock_from_info.return_value + project="byok-project", credentials=mock_build.return_value ) def test_missing_sa_info_raises(self): diff --git a/backend/app/tests/services/buckets/test_registry.py b/backend/app/tests/services/buckets/test_registry.py index 858e122e4..631b697e1 100644 --- a/backend/app/tests/services/buckets/test_registry.py +++ b/backend/app/tests/services/buckets/test_registry.py @@ -43,10 +43,7 @@ def test_get_bucket_provider_with_gcs(self, db: Session): with ( patch("app.crud.credentials.get_provider_credential") as mock_get_creds, - patch( - "app.services.buckets.providers.gcs.service_account." - "Credentials.from_service_account_info" - ), + patch("app.services.buckets.providers.gcs.build_gcp_sa_credentials"), patch("app.services.buckets.providers.gcs.gcs.Client") as mock_client, ): mock_get_creds.return_value = credential From 00307db471258381103f809425da95438d19f394 Mon Sep 17 00:00:00 2001 From: Prashant Vasudevan <71649489+vprashrex@users.noreply.github.com> Date: Wed, 26 Aug 2026 21:01:46 +0530 Subject: [PATCH 13/15] feat: add support for GOOGLE_AISTUDIO and GOOGLE_AISTUDIO_NATIVE providers --- backend/app/crud/assessment/batch.py | 2 +- backend/app/crud/assessment/processing.py | 2 ++ backend/app/services/assessment/api/batch.py | 6 +++++- 3 files changed, 8 insertions(+), 2 deletions(-) diff --git a/backend/app/crud/assessment/batch.py b/backend/app/crud/assessment/batch.py index 0b4fe4e12..1c4464525 100644 --- a/backend/app/crud/assessment/batch.py +++ b/backend/app/crud/assessment/batch.py @@ -475,7 +475,7 @@ def submit_assessment_batch( config=batch_config, ) - elif base_provider == LLMProvider.GOOGLE: + elif base_provider in (LLMProvider.GOOGLE, LLMProvider.GOOGLE_AISTUDIO): mapped_params, warnings = map_kaapi_to_google_params(params) if warnings: logger.info("[submit_assessment_batch] Mapper warnings: %s", warnings) diff --git a/backend/app/crud/assessment/processing.py b/backend/app/crud/assessment/processing.py index 4b38762ac..edf15f12a 100644 --- a/backend/app/crud/assessment/processing.py +++ b/backend/app/crud/assessment/processing.py @@ -236,6 +236,8 @@ def parse_assessment_output( elif provider_name in ( LLMProvider.GOOGLE, LLMProvider.GOOGLE_NATIVE, + LLMProvider.GOOGLE_AISTUDIO, + LLMProvider.GOOGLE_AISTUDIO_NATIVE, ): response = result.get("response") error = result.get("error") diff --git a/backend/app/services/assessment/api/batch.py b/backend/app/services/assessment/api/batch.py index 2942516bc..79e464edd 100644 --- a/backend/app/services/assessment/api/batch.py +++ b/backend/app/services/assessment/api/batch.py @@ -510,7 +510,11 @@ def _parse_one(result: dict[str, Any], provider_name: str) -> ParsedResult: "response_id": response.get("id"), } - if provider_name in (LLMProvider.GOOGLE, LLMProvider.GOOGLE_GCP): + if provider_name in ( + LLMProvider.GOOGLE, + LLMProvider.GOOGLE_AISTUDIO, + LLMProvider.GOOGLE_GCP, + ): response = result.get("response") text = extract_text_from_response_dict(response) if response else None return { From 4f654ab3a2601567dc53f24bf50fe7e3a6cd579e Mon Sep 17 00:00:00 2001 From: Prashant Vasudevan <71649489+vprashrex@users.noreply.github.com> Date: Wed, 26 Aug 2026 21:38:25 +0530 Subject: [PATCH 14/15] feat(tests): add comprehensive tests for GCS bucket provider and Google GCP integration --- .../app/tests/assessment/test_api_batch.py | 9 +++ backend/app/tests/assessment/test_batch.py | 16 +++++ .../assessment/test_prefilter_batching.py | 44 ++++++++++++ .../app/tests/assessment/test_processing.py | 36 ++++++++++ .../app/tests/core/batch/test_google_gcp.py | 45 ++++++++++++ backend/app/tests/core/cloud/__init__.py | 0 backend/app/tests/core/cloud/test_storage.py | 16 +++++ .../app/tests/services/buckets/test_base.py | 39 +++++++++++ .../app/tests/services/buckets/test_gcs.py | 12 ++++ .../tests/services/buckets/test_registry.py | 70 +++++++++++++++++++ .../services/llm/providers/test_google_gcp.py | 48 +++++++++++++ 11 files changed, 335 insertions(+) create mode 100644 backend/app/tests/core/cloud/__init__.py create mode 100644 backend/app/tests/core/cloud/test_storage.py create mode 100644 backend/app/tests/services/buckets/test_base.py diff --git a/backend/app/tests/assessment/test_api_batch.py b/backend/app/tests/assessment/test_api_batch.py index a72c66a2a..208077b92 100644 --- a/backend/app/tests/assessment/test_api_batch.py +++ b/backend/app/tests/assessment/test_api_batch.py @@ -400,6 +400,15 @@ def test_google_gcp_parses_like_google(self) -> None: result = _parse_one({"response": {"candidates": []}}, "google-gcp") assert result["output"] == "vertex text" + def test_google_aistudio_parses_like_google(self) -> None: + with patch( + "app.services.assessment.api.batch.extract_text_from_response_dict", + return_value="aistudio text", + ): + result = _parse_one({"response": {"candidates": []}}, "google-aistudio") + assert result["output"] == "aistudio text" + assert result["error"] is None + def test_unknown_provider(self) -> None: result = _parse_one({"response": {}}, "cohere") assert result["output"] is None diff --git a/backend/app/tests/assessment/test_batch.py b/backend/app/tests/assessment/test_batch.py index 9e66c7627..d00ba8af5 100644 --- a/backend/app/tests/assessment/test_batch.py +++ b/backend/app/tests/assessment/test_batch.py @@ -73,6 +73,22 @@ def test_no_gcs_returns_rows_unchanged_without_resolving(self): assert out is rows mock_resolve.assert_not_called() + def test_empty_cell_skipped_and_unresolved_gcs_left_as_is(self): + att = AssessmentAttachment(column="img", type="image", format="url") + rows = [{"img": ""}, {"img": "gs://b/missing.png"}] + with patch(_REWRITE_RESOLVER, return_value={}): + out = rewrite_gcs_attachment_urls( + session=MagicMock(), + rows=rows, + attachments=[att], + llm_provider="openai", + project_id=1, + organization_id=2, + ) + assert out[0]["img"] == "" + # unresolved gs:// URI stays put rather than becoming a broken/empty value + assert out[1]["img"] == "gs://b/missing.png" + def _make_run() -> MagicMock: run = MagicMock() diff --git a/backend/app/tests/assessment/test_prefilter_batching.py b/backend/app/tests/assessment/test_prefilter_batching.py index 9c86d8267..a8081491c 100644 --- a/backend/app/tests/assessment/test_prefilter_batching.py +++ b/backend/app/tests/assessment/test_prefilter_batching.py @@ -140,6 +140,50 @@ def test_zero_accepted_advances(self) -> None: tasks._submit_stage(session, run, 1, 1) advance.assert_called_once() + def test_prefilter_rewrites_gcs_attachments_before_submit(self) -> None: + run = _run( + stage=Stage.PRE_FILTER_TOPIC_RELEVANCE, + stage_status=StageStatus.PENDING, + stage_batches={}, + ) + session = MagicMock() + assessment = SimpleNamespace( + input={ + **_ASSESSMENT_INPUT, + "attachments": [{"column": "a", "type": "image", "format": "url"}], + } + ) + batch_job = SimpleNamespace(id=7, total_items=3) + with ( + patch.object( + tasks, + "_resolve_run_context", + return_value=(assessment, MagicMock(), SimpleNamespace(), None), + ), + patch.object(tasks, "_load_dataset_rows", return_value=[{"a": "1"}] * 3), + patch.object(tasks, "_accepted_indices", return_value=[0, 1, 2]), + patch.object(tasks, "recompute_assessment_status"), + patch.object(assessment_core, "flag_modified"), + patch.object( + tasks, + "rewrite_gcs_attachment_urls", + return_value=[{"a": "https://signed/1"}] * 3, + ) as rewrite, + patch.object( + tasks, "build_prefilter_requests", return_value=[{"key": "tr_0"}] + ) as build_reqs, + patch.object(tasks, "submit_prefilter_batch", return_value=batch_job), + ): + tasks._submit_stage(session, run, 1, 1) + + rewrite.assert_called_once() + # the rewritten rows (not the raw gs:// rows) are what get built into JSONL + _, build_kwargs = build_reqs.call_args + rows_with_idx = ( + build_kwargs.get("rows_with_idx") or build_reqs.call_args.args[1] + ) + assert [r for _, r in rows_with_idx] == [{"a": "https://signed/1"}] * 3 + class TestAcceptedIndices: def test_uses_persisted_indices_without_downloading(self) -> None: diff --git a/backend/app/tests/assessment/test_processing.py b/backend/app/tests/assessment/test_processing.py index c7afb2750..e96ddfb92 100644 --- a/backend/app/tests/assessment/test_processing.py +++ b/backend/app/tests/assessment/test_processing.py @@ -385,6 +385,42 @@ def test_google_provider_returned(self) -> None: ) mock_batch_cls.assert_called_once_with(client=mock_gemini.client) + @pytest.mark.parametrize("provider_name", ["google-gcp", "google-gcp-native"]) + def test_google_gcp_builds_from_credentials(self, provider_name: str) -> None: + session = MagicMock() + cred = {"gcs_bucket": "b", "sa_key": {"project_id": "p"}} + with ( + patch( + "app.crud.credentials.get_provider_credential", return_value=cred + ) as mock_get_cred, + patch("app.core.batch.GoogleGCPBatchProvider") as mock_cls, + ): + result = _get_batch_provider( + session=session, + provider_name=provider_name, + organization_id=1, + project_id=1, + ) + mock_get_cred.assert_called_once() + mock_cls.from_credentials.assert_called_once_with(cred) + assert result is mock_cls.from_credentials.return_value + + def test_google_gcp_non_dict_credential_raises(self) -> None: + session = MagicMock() + with ( + patch("app.crud.credentials.get_provider_credential", return_value=None), + patch("app.core.batch.GoogleGCPBatchProvider"), + ): + with pytest.raises( + ValueError, match="google-gcp credentials not configured" + ): + _get_batch_provider( + session=session, + provider_name="google-gcp", + organization_id=1, + project_id=1, + ) + class TestProcessRunBatches: def _parent(self): diff --git a/backend/app/tests/core/batch/test_google_gcp.py b/backend/app/tests/core/batch/test_google_gcp.py index 434d68a91..07401747c 100644 --- a/backend/app/tests/core/batch/test_google_gcp.py +++ b/backend/app/tests/core/batch/test_google_gcp.py @@ -6,6 +6,7 @@ import pytest from app.core.batch.google_gcp import GoogleGCPBatchProvider, _parse_gs_uri +from app.core.cloud.storage import CloudStorageError _BUCKET = "test-bucket" @@ -61,6 +62,11 @@ def test_uploads_to_gcs_and_starts_job(self, provider, mock_genai, mock_storage) assert result["provider_status"] == "JOB_STATE_PENDING" assert result["total_items"] == 1 + def test_create_failure_propagates(self, provider, mock_genai): + mock_genai.batches.create.side_effect = RuntimeError("vertex down") + with pytest.raises(RuntimeError, match="vertex down"): + provider.create_batch([{"request": {"contents": []}}], {}) + class TestGetBatchStatus: def test_succeeded_returns_output_uri(self, provider, mock_genai): @@ -86,6 +92,11 @@ def test_failed_sets_error_message(self, provider, mock_genai): assert result["error_message"] == "boom" assert result["provider_output_file_id"] == "batches/123" + def test_poll_failure_propagates(self, provider, mock_genai): + mock_genai.batches.get.side_effect = RuntimeError("poll boom") + with pytest.raises(RuntimeError, match="poll boom"): + provider.get_batch_status("batches/123") + class TestDownloadBatchResults: def _blob(self, name, text): @@ -140,6 +151,28 @@ def test_incomplete_job_raises(self, provider, mock_genai): with pytest.raises(ValueError, match="not complete"): provider.download_batch_results("batches/123") + def test_succeeded_job_without_dest_raises(self, provider, mock_genai): + job = MagicMock() + job.state.name = "JOB_STATE_SUCCEEDED" + job.dest = None + mock_genai.batches.get.return_value = job + with pytest.raises(ValueError, match="no GCS output"): + provider.download_batch_results("batches/123") + + def test_blank_lines_are_skipped(self, provider, mock_storage): + content = ( + '{"key": "row_1", "response": {"text": "a"}}\n\n' + '{"key": "row_2", "response": {"text": "b"}}' + ) + mock_storage.list_blobs.return_value = [self._blob("out/p.jsonl", content)] + results = provider.download_batch_results("gs://b/out/") + assert [r["custom_id"] for r in results] == ["row_1", "row_2"] + + def test_read_failure_propagates(self, provider, mock_storage): + mock_storage.list_blobs.side_effect = RuntimeError("list boom") + with pytest.raises(RuntimeError, match="list boom"): + provider.download_batch_results("gs://b/out/") + class TestFileIO: def test_upload_file_returns_gs_uri(self, provider, mock_storage): @@ -157,6 +190,18 @@ def test_download_file_reads_text(self, provider, mock_storage): ) assert provider.download_file("gs://b/k.jsonl") == "hello" + def test_upload_failure_wraps_cloud_storage_error(self, provider, mock_storage): + blob = mock_storage.bucket.return_value.blob.return_value + blob.upload_from_string.side_effect = RuntimeError("denied") + with pytest.raises(CloudStorageError, match="GCS upload failed"): + provider.upload_file('{"x":1}') + + def test_download_failure_wraps_cloud_storage_error(self, provider, mock_storage): + blob = mock_storage.bucket.return_value.blob.return_value + blob.download_as_text.side_effect = RuntimeError("gone") + with pytest.raises(CloudStorageError, match="GCS download failed"): + provider.download_file("gs://b/k.jsonl") + class TestFromCredentials: _CRED = { diff --git a/backend/app/tests/core/cloud/__init__.py b/backend/app/tests/core/cloud/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/backend/app/tests/core/cloud/test_storage.py b/backend/app/tests/core/cloud/test_storage.py new file mode 100644 index 000000000..95e9ff7ff --- /dev/null +++ b/backend/app/tests/core/cloud/test_storage.py @@ -0,0 +1,16 @@ +"""Tests for app.core.cloud.storage helpers.""" + +from unittest.mock import patch + +from app.core.cloud.storage import GCS_SCOPES, build_gcp_sa_credentials + + +def test_build_gcp_sa_credentials_passes_key_and_scopes(): + sa_key = {"type": "service_account", "project_id": "p"} + with patch( + "app.core.cloud.storage.service_account.Credentials.from_service_account_info" + ) as mock_from_info: + creds = build_gcp_sa_credentials(sa_key) + + mock_from_info.assert_called_once_with(sa_key, scopes=list(GCS_SCOPES)) + assert creds is mock_from_info.return_value diff --git a/backend/app/tests/services/buckets/test_base.py b/backend/app/tests/services/buckets/test_base.py new file mode 100644 index 000000000..ff1e5555d --- /dev/null +++ b/backend/app/tests/services/buckets/test_base.py @@ -0,0 +1,39 @@ +"""Tests for the abstract bucket-provider base class.""" + +from typing import Any + +import pytest + +from app.services.buckets.providers.base import BaseBucketProvider + + +class _StubProvider(BaseBucketProvider): + """Concrete provider that signs by echoing the URI, for the bulk loop.""" + + @staticmethod + def create_client(credentials: dict[str, Any]) -> Any: + return None + + def get_signed_url(self, uri: str, expires_in: int) -> str: + return f"{uri}?exp={expires_in}" + + +class TestBaseBucketProviderAbstractMethods: + def test_create_client_body_raises_not_implemented(self): + with pytest.raises(NotImplementedError, match="create_client"): + BaseBucketProvider.create_client({}) + + def test_get_signed_url_body_raises_not_implemented(self): + provider = _StubProvider(client=None) + with pytest.raises(NotImplementedError, match="get_signed_url"): + BaseBucketProvider.get_signed_url(provider, "gs://b/k", 60) + + +class TestGetBulkSignedUrls: + def test_maps_each_uri_via_get_signed_url(self): + provider = _StubProvider(client=None) + result = provider.get_bulk_signed_urls(["gs://b/a", "gs://b/c"], expires_in=120) + assert result == { + "gs://b/a": "gs://b/a?exp=120", + "gs://b/c": "gs://b/c?exp=120", + } diff --git a/backend/app/tests/services/buckets/test_gcs.py b/backend/app/tests/services/buckets/test_gcs.py index 67d933629..913161402 100644 --- a/backend/app/tests/services/buckets/test_gcs.py +++ b/backend/app/tests/services/buckets/test_gcs.py @@ -5,6 +5,7 @@ import pytest +from app.core.cloud.storage import CloudStorageError from app.services.buckets.providers.gcs import ( GCSBucketProvider, GCSClient, @@ -100,6 +101,17 @@ def test_expiry_capped_at_max(self): _, kwargs = blob.generate_signed_url.call_args assert kwargs["expiration"] == timedelta(seconds=provider.MAX_SIGNED_URL_EXPIRY) + def test_expiry_below_one_raises(self): + provider, _ = _make_provider() + with pytest.raises(ValueError, match="at least 1 second"): + provider.get_signed_url("gs://my-bucket/key.wav", expires_in=0) + + def test_signing_failure_raises_cloud_storage_error(self): + provider, blob = _make_provider() + blob.generate_signed_url.side_effect = RuntimeError("no signer creds") + with pytest.raises(CloudStorageError, match="GCS signing failed"): + provider.get_signed_url("gs://my-bucket/key.wav", expires_in=1800) + class TestGetBulkSignedUrls: def test_returns_uri_to_url_map_reusing_one_client(self): diff --git a/backend/app/tests/services/buckets/test_registry.py b/backend/app/tests/services/buckets/test_registry.py index 631b697e1..3849fb9a6 100644 --- a/backend/app/tests/services/buckets/test_registry.py +++ b/backend/app/tests/services/buckets/test_registry.py @@ -31,6 +31,16 @@ def test_get_provider_class_unknown_raises(self): assert "s3" in message assert "is not supported" in message + def test_supported_providers_lists_gcs(self): + assert BucketProvider.supported_providers() == ["gcs"] + + def test_get_credential_provider_maps_gcs_to_google_gcp(self): + assert BucketProvider.get_credential_provider("gcs") == "google-gcp" + + def test_get_credential_provider_unknown_raises(self): + with pytest.raises(ValueError, match="no credential mapping"): + BucketProvider.get_credential_provider("s3") + class TestGetBucketProvider: def test_get_bucket_provider_with_gcs(self, db: Session): @@ -93,3 +103,63 @@ def test_get_bucket_provider_unknown_type_raises(self, db: Session): ) assert "s3" in str(exc_info.value) + + def test_get_bucket_provider_non_dict_credential_raises(self, db: Session): + project = get_project(db) + + with patch("app.crud.credentials.get_provider_credential") as mock_get_creds: + mock_get_creds.return_value = "not-a-dict" + + with pytest.raises(ValueError) as exc_info: + get_bucket_provider( + session=db, + provider_type="gcs", + project_id=project.id, + organization_id=project.organization_id, + ) + + assert "decrypted credentials dict" in str(exc_info.value) + + def test_get_bucket_provider_client_value_error_surfaces_as_is(self, db: Session): + project = get_project(db) + credential = {"gcs_bucket": "b", "sa_key": {"project_id": "p"}} + + with ( + patch("app.crud.credentials.get_provider_credential") as mock_get_creds, + patch.object( + GCSBucketProvider, + "create_client", + side_effect=ValueError("bad sa_key"), + ), + ): + mock_get_creds.return_value = credential + + with pytest.raises(ValueError, match="bad sa_key"): + get_bucket_provider( + session=db, + provider_type="gcs", + project_id=project.id, + organization_id=project.organization_id, + ) + + def test_get_bucket_provider_client_error_wraps_in_runtime_error(self, db: Session): + project = get_project(db) + credential = {"gcs_bucket": "b", "sa_key": {"project_id": "p"}} + + with ( + patch("app.crud.credentials.get_provider_credential") as mock_get_creds, + patch.object( + GCSBucketProvider, + "create_client", + side_effect=RuntimeError("boom"), + ), + ): + mock_get_creds.return_value = credential + + with pytest.raises(RuntimeError, match="Could not connect to gcs"): + get_bucket_provider( + session=db, + provider_type="gcs", + project_id=project.id, + organization_id=project.organization_id, + ) 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 7b355c6bf..5263a588d 100644 --- a/backend/app/tests/services/llm/providers/test_google_gcp.py +++ b/backend/app/tests/services/llm/providers/test_google_gcp.py @@ -680,6 +680,41 @@ def test_execute_tts_unsupported_response_format_falls_back_to_wav(): assert resp.response.output.content.mime_type == "audio/wav" +def test_execute_text_forwards_max_output_tokens(): + provider = _provider() + config = NativeCompletionConfig( + provider="google-gcp-native", + type=CompletionType.TEXT, + params={"model": "gemini-2.5-flash", "max_output_tokens": 256}, + ) + with patch( + "app.services.llm.providers.google_gcp.requests.post", + return_value=_mock_http_ok(_stt_response("ok")), + ) as mock_post: + resp, err = provider.execute(config, QueryParams(input="ignored"), "hi") + + assert err is None + payload = mock_post.call_args.kwargs["json"] + assert payload["generationConfig"]["maxOutputTokens"] == 256 + + +def test_execute_text_http_error_returns_message(): + provider = _provider() + config = NativeCompletionConfig( + provider="google-gcp-native", + type=CompletionType.TEXT, + params={"model": "gemini-2.5-flash"}, + ) + with patch( + "app.services.llm.providers.google_gcp.requests.post", + return_value=_mock_http_err(403, "permission denied"), + ): + resp, err = provider.execute(config, QueryParams(input="ignored"), "hi") + assert resp is None + assert "403" in err + assert "permission denied" in err + + def test_execute_text_forwards_system_instruction(): provider = _provider() config = NativeCompletionConfig( @@ -815,6 +850,19 @@ def test_execute_wraps_unexpected_exception(): assert "kaboom" in err +def test_execute_unsupported_completion_type_returns_error(): + provider = _provider() + # SimpleNamespace bypasses NativeCompletionConfig's enum validation to reach + # the dispatcher's unsupported-type fallthrough. + from types import SimpleNamespace + + config = SimpleNamespace(provider="google-gcp-native", type="embedding") + resp, err = provider.execute(config, QueryParams(input="ignored"), "hi") + assert resp is None + assert "Unsupported completion type" in err + assert "embedding" in err + + # --------------------------------------------------------------------------- # GoogleGCPClient.endpoint — host changes by location # --------------------------------------------------------------------------- From bc64178f066444603ebbd1cb99f243b3e7320335 Mon Sep 17 00:00:00 2001 From: Prashant Vasudevan <71649489+vprashrex@users.noreply.github.com> Date: Thu, 27 Aug 2026 00:28:30 +0530 Subject: [PATCH 15/15] feat(tests): replace KaapiCompletionConfig with build_kaapi_completion_config in evaluation tests --- backend/app/tests/api/routes/test_evaluation_iteration_v2.py | 4 ++-- backend/app/tests/services/evaluations/test_iteration.py | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/backend/app/tests/api/routes/test_evaluation_iteration_v2.py b/backend/app/tests/api/routes/test_evaluation_iteration_v2.py index a56131233..95d178279 100644 --- a/backend/app/tests/api/routes/test_evaluation_iteration_v2.py +++ b/backend/app/tests/api/routes/test_evaluation_iteration_v2.py @@ -13,7 +13,7 @@ from app.core.config import settings from app.models import Config, EvaluationDataset from app.models.evaluation_iteration import EvaluationIterationRun -from app.models.llm.request import ConfigBlob, KaapiCompletionConfig +from app.models.llm.request import ConfigBlob, build_kaapi_completion_config from app.tests.utils.auth import TestAuthContext from app.tests.utils.test_data import ( create_test_config, @@ -34,7 +34,7 @@ def _make_dataset(*, db: Session, user_api_key: TestAuthContext) -> EvaluationDa def _make_text_config(db: Session, project_id: int) -> Config: blob = ConfigBlob( - completion=KaapiCompletionConfig( + completion=build_kaapi_completion_config( provider="openai", type="text", params={"model": "gpt-4o-iter-route-test", "temperature": 0.7}, diff --git a/backend/app/tests/services/evaluations/test_iteration.py b/backend/app/tests/services/evaluations/test_iteration.py index 8f283daf7..92fdc3fd4 100644 --- a/backend/app/tests/services/evaluations/test_iteration.py +++ b/backend/app/tests/services/evaluations/test_iteration.py @@ -24,7 +24,7 @@ EvaluationIterationRun, EvaluationIterationStatusEnum, ) -from app.models.llm.request import ConfigBlob, KaapiCompletionConfig +from app.models.llm.request import ConfigBlob, build_kaapi_completion_config from app.services.evaluations.iteration import ( compute_round_scores, validate_and_start_evaluation_iteration, @@ -100,7 +100,7 @@ def _make_dataset(*, db: Session, user_api_key: TestAuthContext) -> EvaluationDa def _make_text_config(db: Session, project_id: int) -> Config: blob = ConfigBlob( - completion=KaapiCompletionConfig( + completion=build_kaapi_completion_config( provider="openai", type="text", params={"model": "gpt-4o-iteration-test", "temperature": 0.7},