diff --git a/backend/app/core/batch/__init__.py b/backend/app/core/batch/__init__.py index 1e8202f96..c35948171 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 VertexBatchProvider from .operations import ( download_batch_results, process_completed_batch, @@ -28,6 +29,7 @@ "GeminiClient", "GeminiClientError", "GeminiBatchProvider", + "VertexBatchProvider", "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..01870a954 --- /dev/null +++ b/backend/app/core/batch/google_gcp.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..1f92c5208 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, MessageBatchStatus, OpenAIBatchProvider, + GeminiBatchProvider, + VertexBatchProvider, extract_text_from_response_dict, poll_batch_status, process_completed_batch, @@ -35,6 +36,8 @@ ) 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 @@ -43,6 +46,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 +130,7 @@ class StageKind(StrEnum): _SUPPORTED_PROVIDERS = { LLMProvider.OPENAI, LLMProvider.GOOGLE, + LLMProvider.GOOGLE_GCP, LLMProvider.ANTHROPIC, } @@ -291,6 +296,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 +321,34 @@ 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 = GeminiBatchProvider(client=gemini.client, model=f"models/{model}") - config = {"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( @@ -362,7 +395,20 @@ 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_GCP): + 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 ) @@ -453,7 +499,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..1778d8695 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/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..c5002065e 100644 --- a/backend/app/services/assessment/utils/attachments.py +++ b/backend/app/services/assessment/utils/attachments.py @@ -1,17 +1,15 @@ -"""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.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 +32,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..16fa01cde --- /dev/null +++ b/backend/app/services/buckets/providers/gcs.py @@ -0,0 +1,110 @@ +"""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..e6db26017 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 @@ -872,6 +880,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.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 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( @@ -934,6 +979,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.VertexBatchProvider.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 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_google_gcp.py b/backend/app/tests/core/batch/test_google_gcp.py new file mode 100644 index 000000000..183bdf6dd --- /dev/null +++ b/backend/app/tests/core/batch/test_google_gcp.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.google_gcp 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.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) + 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..6b9f36590 --- /dev/null +++ b/backend/app/tests/services/buckets/test_attachments.py @@ -0,0 +1,128 @@ +"""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..078b432da --- /dev/null +++ b/backend/app/tests/services/buckets/test_gcs.py @@ -0,0 +1,152 @@ +"""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..858e122e4 --- /dev/null +++ b/backend/app/tests/services/buckets/test_registry.py @@ -0,0 +1,98 @@ +"""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..97047be65 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 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`). 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` |