diff --git a/backend/app/core/batch/__init__.py b/backend/app/core/batch/__init__.py index 1e8202f96..666bf5af0 100644 --- a/backend/app/core/batch/__init__.py +++ b/backend/app/core/batch/__init__.py @@ -11,6 +11,7 @@ extract_text_from_response_dict, ) from .openai import OpenAIBatchProvider +from .google_gcp import GoogleGCPBatchProvider from .operations import ( download_batch_results, process_completed_batch, @@ -28,6 +29,7 @@ "GeminiClient", "GeminiClientError", "GeminiBatchProvider", + "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 new file mode 100644 index 000000000..51310a15c --- /dev/null +++ b/backend/app/core/batch/google_gcp.py @@ -0,0 +1,257 @@ +"""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, cast +from uuid import uuid4 + +from google import genai +from google.genai import types +from google.cloud import storage as gcs + +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 + +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 GoogleGCPBatchProvider(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-3.1-pro-preview" + + 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 + ) -> "GoogleGCPBatchProvider": + """Build a Vertex batch provider from a ``google-gcp`` credential dict.""" + 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=creds_model.project_id, + location=creds_model.location, + credentials=creds, + ) + storage_client = gcs.Client(project=creds_model.project_id, credentials=creds) + return cls( + client=client, + storage_client=storage_client, + gcs_bucket=creds_model.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/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/core/config.py b/backend/app/core/config.py index 82fd72fff..e6d8673e7 100644 --- a/backend/app/core/config.py +++ b/backend/app/core/config.py @@ -118,6 +118,8 @@ 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 = "" + # 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 RABBITMQ_HOST: str = "localhost" diff --git a/backend/app/crud/assessment/batch.py b/backend/app/crud/assessment/batch.py index 9250aa5b3..1c4464525 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.mappers import kaapi_params_as_dict from app.services.llm.providers.registry import LLMProvider @@ -423,6 +424,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, @@ -464,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/models/llm/constants.py b/backend/app/models/llm/constants.py index f6e927bfc..a731b00a7 100644 --- a/backend/app/models/llm/constants.py +++ b/backend/app/models/llm/constants.py @@ -186,6 +186,7 @@ def normalize_bcp47_language(value: str) -> str: "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 ff4ac2a1e..79e464edd 100644 --- a/backend/app/services/assessment/api/batch.py +++ b/backend/app/services/assessment/api/batch.py @@ -25,9 +25,10 @@ BATCH_KEY, AnthropicBatchProvider, BatchJobState, - GeminiBatchProvider, + GoogleGCPBatchProvider, MessageBatchStatus, OpenAIBatchProvider, + GeminiBatchProvider, extract_text_from_response_dict, poll_batch_status, process_completed_batch, @@ -35,14 +36,17 @@ ) from app.core.batch.base import BatchProvider from app.core.batch.client import GeminiClient +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, 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, @@ -68,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 @@ -124,6 +150,8 @@ class StageKind(StrEnum): _SUPPORTED_PROVIDERS = { LLMProvider.OPENAI, LLMProvider.GOOGLE, + LLMProvider.GOOGLE_AISTUDIO, + LLMProvider.GOOGLE_GCP, LLMProvider.ANTHROPIC, } @@ -283,6 +311,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( @@ -298,16 +336,32 @@ def _submit_provider_batch( "description": description, "completion_window": "24h", } - elif provider_name == LLMProvider.GOOGLE: + 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 ) - 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}"} + if provider_name == LLMProvider.GOOGLE_GCP: + cred = _google_gcp_credential( + session=session, + organization_id=organization_id, + project_id=project_id, + ) + provider = GoogleGCPBatchProvider.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( @@ -354,7 +408,18 @@ def _build_batch_provider( session=session, org_id=organization_id, project_id=project_id ) ) - if provider_name == LLMProvider.GOOGLE: + if provider_name in ( + LLMProvider.GOOGLE, + LLMProvider.GOOGLE_AISTUDIO, + LLMProvider.GOOGLE_GCP, + ): + if provider_name == LLMProvider.GOOGLE_GCP: + cred = _google_gcp_credential( + session=session, + organization_id=organization_id, + project_id=project_id, + ) + return GoogleGCPBatchProvider.from_credentials(cred) gemini = GeminiClient.from_credentials( session=session, org_id=organization_id, project_id=project_id ) @@ -445,7 +510,11 @@ 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_AISTUDIO, + 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 f98d8cd8c..6ea711feb 100644 --- a/backend/app/services/assessment/api/submission.py +++ b/backend/app/services/assessment/api/submission.py @@ -31,7 +31,7 @@ logger = logging.getLogger(__name__) # Attachment cell values are provided as URLs (base64 is unsupported for batch). -_URL_PREFIXES = ("http://", "https://") +_URL_PREFIXES = ("http://", "https://", "gs://") _ATTACHMENT_TYPES = ("image", "pdf") 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/assessment/tasks.py b/backend/app/services/assessment/tasks.py index c4d6b45a5..33e066385 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,18 @@ 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..d4c41af14 100644 --- a/backend/app/services/assessment/utils/attachments.py +++ b/backend/app/services/assessment/utils/attachments.py @@ -1,17 +1,16 @@ -"""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 -from typing import Any +from typing import Any, cast from urllib.parse import urlparse +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 logger = logging.getLogger(__name__) @@ -34,6 +33,61 @@ 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, + expires_in=settings.MAX_SIGNED_URL_EXPIRY_SECONDS, + ) + 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 + 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 + + 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..7746ca1c6 --- /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, + 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..094e8dcb9 --- /dev/null +++ b/backend/app/services/buckets/providers/base.py @@ -0,0 +1,40 @@ +"""Base provider interface for object-storage bucket providers.""" + +import logging +from abc import ABC, abstractmethod +from typing import Any + +from app.core.config import settings + +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, enforced by concrete providers. + MAX_SIGNED_URL_EXPIRY: int = settings.MAX_SIGNED_URL_EXPIRY_SECONDS + + 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) -> 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) -> 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 diff --git a/backend/app/services/buckets/providers/gcs.py b/backend/app/services/buckets/providers/gcs.py new file mode 100644 index 000000000..e59151e80 --- /dev/null +++ b/backend/app/services/buckets/providers/gcs.py @@ -0,0 +1,108 @@ +"""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 app.core.cloud.storage import CloudStorageError, build_gcp_sa_credentials +from app.services.buckets.providers.base import BaseBucketProvider + +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 from the given credentials.""" + credentials = credentials or {} + sa_info = credentials.get("sa_key") + gcs_bucket = credentials.get("gcs_bucket") + + if not sa_info: + raise ValueError( + "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 = build_gcp_sa_credentials(sa_info) + 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) -> str: + 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) + + 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) -> 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/services/llm/providers/google_gcp.py b/backend/app/services/llm/providers/google_gcp.py index d486c1ab2..172d24953 100644 --- a/backend/app/services/llm/providers/google_gcp.py +++ b/backend/app/services/llm/providers/google_gcp.py @@ -32,6 +32,7 @@ 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 @@ -761,6 +762,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, 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/assessment/test_api_batch.py b/backend/app/tests/assessment/test_api_batch.py index 45092502c..208077b92 100644 --- a/backend/app/tests/assessment/test_api_batch.py +++ b/backend/app/tests/assessment/test_api_batch.py @@ -392,6 +392,23 @@ 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_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 @@ -876,6 +893,43 @@ def test_google(self, db) -> None: ) assert provider is not None + def test_google_gcp_routes_to_vertex(self, db) -> None: + auth = get_user_test_auth_context(db) + 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.GoogleGCPBatchProvider.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 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) with patch( @@ -938,6 +992,33 @@ def test_google_branch(self, db) -> None: assert result.id == job.id assert start.call_args.kwargs["provider_name"] == "google" + 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_provider_credential", + return_value={"gcs_bucket": "b", "sa_key": {}}, + ), + patch( + "app.services.assessment.api.batch.GoogleGCPBatchProvider.from_credentials", + return_value=MagicMock(), + ) as vertex_from_cred, + 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" + 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"] + def test_anthropic_branch_sets_max_tokens(self, db) -> None: auth = get_user_test_auth_context(db) job = _make_batch_job( diff --git a/backend/app/tests/assessment/test_api_submission.py b/backend/app/tests/assessment/test_api_submission.py index bf540a32a..2ddbb5854 100644 --- a/backend/app/tests/assessment/test_api_submission.py +++ b/backend/app/tests/assessment/test_api_submission.py @@ -162,6 +162,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..d00ba8af5 100644 --- a/backend/app/tests/assessment/test_batch.py +++ b/backend/app/tests/assessment/test_batch.py @@ -26,10 +26,69 @@ 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 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 new file mode 100644 index 000000000..07401747c --- /dev/null +++ b/backend/app/tests/core/batch/test_google_gcp.py @@ -0,0 +1,231 @@ +"""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 GoogleGCPBatchProvider, _parse_gs_uri +from app.core.cloud.storage import CloudStorageError + +_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 GoogleGCPBatchProvider( + 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 + + 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): + 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" + + 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): + 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") + + 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): + 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" + + 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 = { + "api_key": "test-key", + "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.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) + build_creds.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): + GoogleGCPBatchProvider.from_credentials(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/__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..7a9aa2ab6 --- /dev/null +++ b/backend/app/tests/services/buckets/test_attachments.py @@ -0,0 +1,132 @@ +"""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, + expires_in=86400, + ) + 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, + expires_in=86400, + ) + 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, + 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=86400 + ) + + 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, + expires_in=86400, + ) + 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_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 new file mode 100644 index 000000000..913161402 --- /dev/null +++ b/backend/app/tests/services/buckets/test_gcs.py @@ -0,0 +1,132 @@ +"""Tests for the GCS bucket provider.""" + +from datetime import timedelta +from unittest.mock import MagicMock, patch + +import pytest + +from app.core.cloud.storage import CloudStorageError +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_uses_byok_credentials(self): + byok_sa = {"project_id": "byok-project"} + + with ( + patch( + "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( + {"gcs_bucket": "byok-bucket", "sa_key": byok_sa} + ) + + assert client.default_bucket == "byok-bucket" + mock_build.assert_called_once_with(byok_sa) + mock_client.assert_called_once_with( + project="byok-project", credentials=mock_build.return_value + ) + + def test_missing_sa_info_raises(self): + 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): + 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) + + 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): + 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"], expires_in=86400 + ) + + 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..3849fb9a6 --- /dev/null +++ b/backend/app/tests/services/buckets/test_registry.py @@ -0,0 +1,165 @@ +"""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 + + 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): + 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.build_gcp_sa_credentials"), + 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) + + 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/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}, 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 98be39e2c..5263a588d 100644 --- a/backend/app/tests/services/llm/providers/test_google_gcp.py +++ b/backend/app/tests/services/llm/providers/test_google_gcp.py @@ -304,17 +304,21 @@ def test_tts_language_is_forwarded(self, provider, query): # ── execute dispatcher ─────────────────────────────────────────────────── 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}, ) with patch( "app.services.llm.providers.google_gcp.requests.post", - return_value=_mock_http_ok(_stt_response("hi there")), - ): + 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 == "hi there" + 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 @@ -676,20 +680,57 @@ def test_execute_tts_unsupported_response_format_falls_back_to_wav(): assert resp.response.output.content.mime_type == "audio/wav" -def test_execute_text_happy_path(): +def test_execute_text_forwards_max_output_tokens(): provider = _provider() config = NativeCompletionConfig( - provider="google-native", + 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_ok(_stt_response("hi there")), + 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( + provider="google-gcp-native", + type=CompletionType.TEXT, + params={"model": "gemini-2.5-flash", "instructions": "be terse"}, + ) + 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 - assert resp.response.output.content.value == "hi there" + payload = mock_post.call_args.kwargs["json"] + assert payload["systemInstruction"] == {"parts": [{"text": "be terse"}]} def test_execute_tts_language_is_forwarded(): @@ -809,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 # --------------------------------------------------------------------------- diff --git a/docs/wiki/domain-map.md b/docs/wiki/domain-map.md index 714e0f540..7187f0129 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 | @@ -55,7 +55,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 43d221f3d..4314faade 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` holds the `{topic_relevance}` verdict (`{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 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`). diff --git a/docs/wiki/modules/platform.md b/docs/wiki/modules/platform.md index 0bde6b828..a41197c24 100644 --- a/docs/wiki/modules/platform.md +++ b/docs/wiki/modules/platform.md @@ -12,6 +12,7 @@ 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 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` |