From 35a748846dbca338d76de6c0cc65605927e15de5 Mon Sep 17 00:00:00 2001 From: Behrouz Mirabdi Date: Thu, 30 Jul 2026 15:25:24 +0200 Subject: [PATCH 1/5] feat(embedding-api): enforce MVP request envelope from env Cap sequences, AA length, FASTA size, and sync predict timeout via config instead of hardcoded route defaults, and reject oversize requests early on job and predict paths. --- services/embedding-api/config.py | 14 +++++++++++--- 1 file changed, 11 insertions(+), 3 deletions(-) diff --git a/services/embedding-api/config.py b/services/embedding-api/config.py index cc6cb25..a48ed60 100644 --- a/services/embedding-api/config.py +++ b/services/embedding-api/config.py @@ -8,9 +8,6 @@ DEFAULT_BACKEND = "esm2" DEFAULT_POOLING = "mean" DEFAULT_BATCH_SIZE = 8 -DEFAULT_MAX_LENGTH = 1280 - -MAX_FASTA_UPLOAD_BYTES = 5 * 1024 * 1024 # 5 MB GO_PREDICTION_API_URL = os.getenv("GO_PREDICTION_API_URL", "http://go-prediction-api:8000") @@ -24,3 +21,14 @@ RQ_RETRY_MAX = int(os.getenv("EMBEDDING_RQ_RETRY_MAX", "3")) RQ_RETRY_INTERVALS = [10, 60, 180] WORKER_METRICS_PORT = int(os.getenv("WORKER_METRICS_PORT", "8001")) + +# Request envelope (MVP). Tune via .env — do not hardcode in route handlers. +MAX_SEQUENCES_PER_REQUEST = int(os.getenv("MAX_SEQUENCES_PER_REQUEST", "20")) +MAX_SEQUENCE_LENGTH_AA = int(os.getenv("MAX_SEQUENCE_LENGTH_AA", "1000")) +MAX_FASTA_UPLOAD_MB = int(os.getenv("MAX_FASTA_UPLOAD_MB", "2")) +MAX_FASTA_UPLOAD_BYTES = MAX_FASTA_UPLOAD_MB * 1024 * 1024 +SYNC_PREDICT_TIMEOUT_SEC = int(os.getenv("SYNC_PREDICT_TIMEOUT_SEC", "600")) +SYNC_PREDICT_POLL_INTERVAL_SEC = float(os.getenv("SYNC_PREDICT_POLL_INTERVAL_SEC", "1.0")) + +# Tokenizer window default: match AA length cap unless overridden. +DEFAULT_MAX_LENGTH = int(os.getenv("DEFAULT_MAX_LENGTH", str(MAX_SEQUENCE_LENGTH_AA))) From 29f03be8b5dda34d1646a31349c79ba1b62d3fbe Mon Sep 17 00:00:00 2001 From: Behrouz Mirabdi Date: Thu, 30 Jul 2026 15:26:41 +0200 Subject: [PATCH 2/5] feat(streamlit-ui): mirror MVP envelope limits in client validation Read sequence/FASTA/timeout caps from env so the UI rejects bad uploads before the gateway, and surface those limits in the FASTA help text. --- services/streamlit-ui/.streamlit/config.toml | 4 +-- services/streamlit-ui/app.py | 26 +++++++++++--------- services/streamlit-ui/validation.py | 21 +++++++++++++--- tests/unit/test_ui_validation.py | 15 +++++++++++ 4 files changed, 50 insertions(+), 16 deletions(-) diff --git a/services/streamlit-ui/.streamlit/config.toml b/services/streamlit-ui/.streamlit/config.toml index 9d1654a..44679e6 100644 --- a/services/streamlit-ui/.streamlit/config.toml +++ b/services/streamlit-ui/.streamlit/config.toml @@ -1,3 +1,3 @@ [server] -# Align uploader caption + hard limit with API/nginx FASTA cap (5 MB). -maxUploadSize = 5 +# Align uploader hard limit with MAX_FASTA_UPLOAD_MB (env / API / nginx). +maxUploadSize = 2 diff --git a/services/streamlit-ui/app.py b/services/streamlit-ui/app.py index 9e12dd0..6cc02d7 100644 --- a/services/streamlit-ui/app.py +++ b/services/streamlit-ui/app.py @@ -11,7 +11,10 @@ from requests.exceptions import RequestException from validation import ( - MAX_FASTA_UPLOAD_BYTES, + MAX_FASTA_UPLOAD_MB, + MAX_SEQUENCE_LENGTH_AA, + MAX_SEQUENCES_PER_REQUEST, + SYNC_PREDICT_TIMEOUT_SEC, GatewayConfig, load_gateway_config, normalize_sequence, @@ -32,8 +35,6 @@ PREDICT_SEQUENCES_ENDPOINT = "/api/v1/predict-go-from-sequences" PREDICT_FASTA_ENDPOINT = "/api/v1/predict-go-from-fasta" MAX_TOP_K = 500 -SEQUENCE_TIMEOUT_SECONDS = 600 -FASTA_TIMEOUT_SECONDS = 1800 PREDICTION_MODE_SEQUENCE = "Prediction with sequence" PREDICTION_MODE_FASTA = "Prediction with FASTA" @@ -55,8 +56,9 @@ def build_request_payload(sequence: str, top_k: int) -> dict[str, Any]: "backend": "esm2", "pooling": "mean", "batch_size": 1, - "max_length": 1280, + "max_length": MAX_SEQUENCE_LENGTH_AA, "top_k": top_k, + "timeout_seconds": SYNC_PREDICT_TIMEOUT_SEC, "sequences": [{"id": "input_1", "sequence": sequence}], } @@ -103,7 +105,7 @@ def call_fasta_prediction_api( "backend": "esm2", "pooling": "mean", "batch_size": "8", - "max_length": "1280", + "max_length": str(MAX_SEQUENCE_LENGTH_AA), "top_k": str(top_k), "fail_fast": "true", "timeout_seconds": str(timeout_seconds), @@ -231,7 +233,7 @@ def main() -> None: gateway=gateway, sequence=cleaned, top_k=int(top_k), - timeout_seconds=SEQUENCE_TIMEOUT_SECONDS, + timeout_seconds=SYNC_PREDICT_TIMEOUT_SEC, ) else: with st.form("predict_fasta_form"): @@ -239,8 +241,9 @@ def main() -> None: "Protein FASTA file", type=["fasta", "fa", "txt"], help=( - f"Upload a UTF-8 FASTA file (max {MAX_FASTA_UPLOAD_BYTES // (1024 * 1024)} MB). " - "All records in the file are embedded and predicted. " + f"Upload a UTF-8 FASTA file (max {MAX_FASTA_UPLOAD_MB} MB, " + f"up to {MAX_SEQUENCES_PER_REQUEST} sequences, " + f"each <= {MAX_SEQUENCE_LENGTH_AA} aa). " "Non-canonical residues are normalized by the API at embedding time." ), ) @@ -265,8 +268,9 @@ def main() -> None: ) if record_count > 1: st.info( - f"This FASTA contains {record_count} sequences. " - "Large files may take several minutes to complete." + f"This FASTA contains {record_count} sequences " + f"(limit {MAX_SEQUENCES_PER_REQUEST}). " + "Larger files may take several minutes to complete." ) with st.spinner(f"Submitting request to {PREDICT_FASTA_ENDPOINT} ..."): @@ -275,7 +279,7 @@ def main() -> None: file_bytes=file_bytes, filename=fasta_file.name, top_k=int(top_k), - timeout_seconds=FASTA_TIMEOUT_SECONDS, + timeout_seconds=SYNC_PREDICT_TIMEOUT_SEC, ) if not ok: diff --git a/services/streamlit-ui/validation.py b/services/streamlit-ui/validation.py index 6d1582b..e3c0d72 100644 --- a/services/streamlit-ui/validation.py +++ b/services/streamlit-ui/validation.py @@ -2,11 +2,17 @@ from __future__ import annotations +import os import re from dataclasses import dataclass AA_PATTERN = re.compile(r"^[ACDEFGHIKLMNPQRSTVWY]+$") -MAX_FASTA_UPLOAD_BYTES = 5 * 1024 * 1024 # must match embedding-api config + nginx route + +MAX_SEQUENCES_PER_REQUEST = int(os.getenv("MAX_SEQUENCES_PER_REQUEST", "20")) +MAX_SEQUENCE_LENGTH_AA = int(os.getenv("MAX_SEQUENCE_LENGTH_AA", "1000")) +MAX_FASTA_UPLOAD_MB = int(os.getenv("MAX_FASTA_UPLOAD_MB", "2")) +MAX_FASTA_UPLOAD_BYTES = MAX_FASTA_UPLOAD_MB * 1024 * 1024 +SYNC_PREDICT_TIMEOUT_SEC = int(os.getenv("SYNC_PREDICT_TIMEOUT_SEC", "600")) @dataclass(frozen=True) @@ -25,6 +31,11 @@ def normalize_sequence(raw_sequence: str) -> str: def validate_sequence(sequence: str) -> tuple[bool, str]: if not sequence: return False, "Sequence is empty after whitespace cleanup." + if len(sequence) > MAX_SEQUENCE_LENGTH_AA: + return ( + False, + f"Sequence length {len(sequence)} exceeds the max of {MAX_SEQUENCE_LENGTH_AA} amino acids.", + ) if not AA_PATTERN.fullmatch(sequence): return ( False, @@ -37,8 +48,7 @@ def validate_fasta_upload(file_bytes: bytes, filename: str) -> tuple[bool, str]: if not file_bytes: return False, "Uploaded FASTA file is empty." if len(file_bytes) > MAX_FASTA_UPLOAD_BYTES: - max_mb = MAX_FASTA_UPLOAD_BYTES // (1024 * 1024) - return False, f"FASTA file exceeds the {max_mb} MB upload limit." + return False, f"FASTA file exceeds the {MAX_FASTA_UPLOAD_MB} MB upload limit." try: fasta_text = file_bytes.decode("utf-8") except UnicodeDecodeError: @@ -48,6 +58,11 @@ def validate_fasta_upload(file_bytes: bytes, filename: str) -> tuple[bool, str]: record_count = sum(1 for line in fasta_text.splitlines() if line.startswith(">")) if record_count == 0: return False, "FASTA file has no records (lines starting with '>')." + if record_count > MAX_SEQUENCES_PER_REQUEST: + return ( + False, + f"FASTA has {record_count} sequences; max allowed is {MAX_SEQUENCES_PER_REQUEST}.", + ) return True, "" diff --git a/tests/unit/test_ui_validation.py b/tests/unit/test_ui_validation.py index fc6ca83..8fa1ce5 100644 --- a/tests/unit/test_ui_validation.py +++ b/tests/unit/test_ui_validation.py @@ -52,6 +52,21 @@ def test_validate_fasta_upload_too_large(ui_validation) -> None: assert "exceeds" in msg.lower() +def test_validate_sequence_too_long(ui_validation) -> None: + too_long = "A" * (ui_validation.MAX_SEQUENCE_LENGTH_AA + 1) + ok, msg = ui_validation.validate_sequence(too_long) + assert ok is False + assert "exceeds" in msg.lower() + + +def test_validate_fasta_too_many_sequences(ui_validation) -> None: + limit = ui_validation.MAX_SEQUENCES_PER_REQUEST + records = "\n".join(f">p{i}\nACDE" for i in range(limit + 1)).encode() + ok, msg = ui_validation.validate_fasta_upload(records, "many.fasta") + assert ok is False + assert "max allowed" in msg.lower() + + def test_load_gateway_config_ok(ui_validation) -> None: cfg, err = ui_validation.load_gateway_config( base_url="http://nginx/", From 0abbfe7acb026222ef42f1fe7b7d2a3a996dffd6 Mon Sep 17 00:00:00 2001 From: Behrouz Mirabdi Date: Thu, 30 Jul 2026 15:27:14 +0200 Subject: [PATCH 3/5] chore(infra): wire MVP envelope env and align nginx predict timeouts Expose envelope vars in Compose for embedding-api and streamlit-ui, document them in .env.example, and set predict routes to 2m body / 660s read timeout for cold starts. --- .env.example | 7 +++++++ docker-compose.yml | 9 +++++++++ nginx/nginx.conf | 8 +++++--- 3 files changed, 21 insertions(+), 3 deletions(-) diff --git a/.env.example b/.env.example index a9c2e59..c49759b 100644 --- a/.env.example +++ b/.env.example @@ -28,6 +28,13 @@ REDIS_URL=redis://redis:6379/0 EMBEDDING_JOB_TIMEOUT_SEC=3600 TRAINING_JOB_TIMEOUT_SEC=86400 +# Request envelope (MVP production limits) +MAX_SEQUENCES_PER_REQUEST=20 # 20 sequences per request +MAX_SEQUENCE_LENGTH_AA=1000 # 1000 amino acids per sequence +MAX_FASTA_UPLOAD_MB=2 # 2MB +SYNC_PREDICT_TIMEOUT_SEC=600 # 10 minutes +SYNC_PREDICT_POLL_INTERVAL_SEC=1.0 # 1 second + # MinIO (S3-compatible artifact store) MINIO_ROOT_USER=mlflow-minio MINIO_ROOT_PASSWORD=change-me-minio diff --git a/docker-compose.yml b/docker-compose.yml index 6f87101..06ba0f2 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -245,6 +245,11 @@ services: REDIS_URL: redis://redis:6379/0 EMBEDDING_JOB_TIMEOUT_SEC: ${EMBEDDING_JOB_TIMEOUT_SEC:-3600} EMBEDDING_ARTIFACT_ROOT: /app/outputs/service_artifacts + MAX_SEQUENCES_PER_REQUEST: ${MAX_SEQUENCES_PER_REQUEST:-20} + MAX_SEQUENCE_LENGTH_AA: ${MAX_SEQUENCE_LENGTH_AA:-1000} + MAX_FASTA_UPLOAD_MB: ${MAX_FASTA_UPLOAD_MB:-2} + SYNC_PREDICT_TIMEOUT_SEC: ${SYNC_PREDICT_TIMEOUT_SEC:-600} + SYNC_PREDICT_POLL_INTERVAL_SEC: ${SYNC_PREDICT_POLL_INTERVAL_SEC:-1.0} depends_on: go-prediction-api: condition: service_started @@ -323,6 +328,10 @@ services: GATEWAY_USER_PASSWORD: ${GATEWAY_USER_PASSWORD:-change-me-gateway-user} # Compose uses plain HTTP to nginx; disable TLS verification explicitly. GATEWAY_VERIFY_TLS: "false" + MAX_SEQUENCES_PER_REQUEST: ${MAX_SEQUENCES_PER_REQUEST:-20} + MAX_SEQUENCE_LENGTH_AA: ${MAX_SEQUENCE_LENGTH_AA:-1000} + MAX_FASTA_UPLOAD_MB: ${MAX_FASTA_UPLOAD_MB:-2} + SYNC_PREDICT_TIMEOUT_SEC: ${SYNC_PREDICT_TIMEOUT_SEC:-600} depends_on: - embedding-api - go-prediction-api diff --git a/nginx/nginx.conf b/nginx/nginx.conf index c24641c..72c5b7a 100644 --- a/nginx/nginx.conf +++ b/nginx/nginx.conf @@ -128,12 +128,12 @@ http { auth_basic_user_file /etc/nginx/.htpasswd-user; # auth basic user file for GO prediction API } - # Training API — regular user + admin (.htpasswd-user lists both) location /api/v1/predict-go-from-sequences { - client_max_body_size 512m; # client max body size for NGINX + client_max_body_size 2m; # aligned with MAX_FASTA_UPLOAD_MB MVP envelope limit_req zone=rl_predict burst=80 nodelay; # limit request zone for predict API set $embedding_api_upstream embedding-api:8000; proxy_pass http://$embedding_api_upstream; # proxy pass for GO prediction API + proxy_read_timeout 660s; auth_basic "Prediction API"; # auth basic for GO prediction API auth_basic_user_file /etc/nginx/.htpasswd-user; # auth basic user file for GO prediction API @@ -141,10 +141,12 @@ http { # FASTA upload → embed → predict GO — regular user + admin (.htpasswd-user lists both) location /api/v1/predict-go-from-fasta { - client_max_body_size 5m; + client_max_body_size 2m; # aligned with MAX_FASTA_UPLOAD_MB limit_req zone=rl_predict burst=80 nodelay; set $embedding_api_upstream embedding-api:8000; proxy_pass http://$embedding_api_upstream; + # SYNC_PREDICT_TIMEOUT_SEC default 600; keep headroom for cold starts + proxy_read_timeout 660s; auth_basic "Prediction API"; auth_basic_user_file /etc/nginx/.htpasswd-user; From 494c3f46ac8b900e47bd403d04152f3b83b91a2c Mon Sep 17 00:00:00 2001 From: Behrouz Mirabdi Date: Thu, 30 Jul 2026 15:28:45 +0200 Subject: [PATCH 4/5] chore(examples): add human PLSCR1 sequence to small_sequences.fasta Include a longer UniProt record for local MVP envelope and predict smoke checks. --- examples/small_sequences.fasta | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/examples/small_sequences.fasta b/examples/small_sequences.fasta index 307d643..353500b 100644 --- a/examples/small_sequences.fasta +++ b/examples/small_sequences.fasta @@ -10,3 +10,10 @@ RIISSIEQKEENKGGEDKLKMIREYRQMVETELKLICCDILDVLDKHLIPAANTGESKVF YYKMKGDYHRYLAEFATGNDRKEAAENSLVAYKAASDIAMTELPPTHPIRLGLALNFSVF YYEILNSPDRACRLAKAAFDDAIAELDTLSEESYKDSTLIMQLLRDNLTLWTSDMQGDGE EQNKEALQDVEDENQ +>sp|O15162|PLS1_HUMAN Phospholipid scramblase 1 OS=Homo sapiens OX=9606 GN=PLSCR1 PE=1 SV=1 +MDKQNSQMNASHPETNLPVGYPPQYPPTAFQGPPGYSGYPGPQVSYPPPPAGHSGPGPAG +FPVPNQPVYNQPVYNQPVGAAGVPWMPAPQPPLNCPPGLEYLSQIDQILIHQQIELLEVL +TGFETNNKYEIKNSFGQRVYFAAEDTDCCTRNCCGPSRPFTLRIIDNMGQEVITLERPLR +CSSCCCPCCLQEIEIQAPPGVPIGYVIQTWHPCLPKFTIQNEKREDVLKISGPCVVCSCC +GDVDFEIKSLDEQCVVGKISKHWTGILREAFTDADNFGIQFPLDLDVKMKAVMIGACFLI +DFMFFESTGSQEQKSGVW From 8020a65e2c10a31b0622634b9b63e0f656827583 Mon Sep 17 00:00:00 2001 From: Behrouz Mirabdi Date: Thu, 30 Jul 2026 15:29:29 +0200 Subject: [PATCH 5/5] feat: add env-driven MVP request envelope across predict stack Limit sequences, AA length, FASTA upload size, and sync predict timeouts end-to-end (API, nginx, Streamlit, Compose) so cold-start GPU deploys stay within production MVP bounds. --- services/embedding-api/main.py | 78 ++++++++++++++++++++++--------- services/embedding-api/schemas.py | 32 +++++++++---- 2 files changed, 81 insertions(+), 29 deletions(-) diff --git a/services/embedding-api/main.py b/services/embedding-api/main.py index 7f949db..781738e 100644 --- a/services/embedding-api/main.py +++ b/services/embedding-api/main.py @@ -16,9 +16,16 @@ from config import ( API_PREFIX, ARTIFACT_ROOT, + DEFAULT_BATCH_SIZE, + DEFAULT_MAX_LENGTH, GO_PREDICTION_API_URL, JOBS_DATABASE_URL, MAX_FASTA_UPLOAD_BYTES, + MAX_FASTA_UPLOAD_MB, + MAX_SEQUENCE_LENGTH_AA, + MAX_SEQUENCES_PER_REQUEST, + SYNC_PREDICT_POLL_INTERVAL_SEC, + SYNC_PREDICT_TIMEOUT_SEC, ) from job_store import JobStore from queueing import enqueue_embedding_job, get_queue @@ -161,12 +168,33 @@ def metrics() -> Response: return Response(content=generate_latest(registry), media_type=CONTENT_TYPE_LATEST) +def _enforce_sequence_envelope(sequences: list[str]) -> None: + """Reject requests that exceed MVP count / AA-length caps from env.""" + if len(sequences) > MAX_SEQUENCES_PER_REQUEST: + raise HTTPException( + status_code=400, + detail=( + f"TOO_MANY_SEQUENCES: max {MAX_SEQUENCES_PER_REQUEST}, " + f"got {len(sequences)}" + ), + ) + for index, sequence in enumerate(sequences): + aa_len = len("".join(sequence.split())) + if aa_len > MAX_SEQUENCE_LENGTH_AA: + raise HTTPException( + status_code=400, + detail=( + f"SEQUENCE_TOO_LONG: index={index} length={aa_len} " + f"max={MAX_SEQUENCE_LENGTH_AA}" + ), + ) + + @app.post(API_PREFIX + "/jobs", response_model=CreateJobResponse, status_code=202) def create_job(request: CreateJobRequest) -> CreateJobResponse: - _observe_sequence_lengths( - backend=request.backend, - sequences=[seq.sequence for seq in request.sequences], - ) + sequences = [seq.sequence for seq in request.sequences] + _enforce_sequence_envelope(sequences) + _observe_sequence_lengths(backend=request.backend, sequences=sequences) job_id = str(uuid.uuid4()) _create_and_enqueue(job_id, request.model_dump()) return CreateJobResponse( @@ -181,16 +209,15 @@ async def create_fasta_job( fasta_file: UploadFile = File(...), backend: Literal["esm2", "protbert", "t5"] = Form(default="esm2"), pooling: Literal["mean", "cls"] = Form(default="mean"), - batch_size: int = Form(default=8), - max_length: int = Form(default=1280), + batch_size: int = Form(default=DEFAULT_BATCH_SIZE), + max_length: int = Form(default=DEFAULT_MAX_LENGTH), ) -> CreateJobResponse: - fasta_text = (await fasta_file.read()).decode("utf-8", errors="replace") - if not fasta_text.strip(): - raise HTTPException(status_code=400, detail="Uploaded FASTA is empty.") + fasta_text = await _read_fasta_upload(fasta_file) try: _, sequences = parse_fasta_text(fasta_text) except ValueError as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc + _enforce_sequence_envelope(sequences) _observe_sequence_lengths(backend=backend, sequences=sequences) job_id = str(uuid.uuid4()) @@ -362,6 +389,7 @@ def _parse_and_validate_fasta(fasta_text: str, backend: str) -> None: status_code=400, detail=f"FASTA records with empty sequences: {preview}{suffix}", ) + _enforce_sequence_envelope(sequences) _observe_sequence_lengths(backend=backend, sequences=sequences) @@ -379,8 +407,13 @@ def _validate_predict_form_params( raise HTTPException(status_code=400, detail="max_length must be between 8 and 8192") if not 1 <= top_k <= 500: raise HTTPException(status_code=400, detail="top_k must be between 1 and 500") - if not 5 <= timeout_seconds <= 7200: - raise HTTPException(status_code=400, detail="timeout_seconds must be between 5 and 7200") + if not 5 <= timeout_seconds <= SYNC_PREDICT_TIMEOUT_SEC: + raise HTTPException( + status_code=400, + detail=( + f"timeout_seconds must be between 5 and {SYNC_PREDICT_TIMEOUT_SEC}" + ), + ) if not 0.1 < poll_interval_seconds <= 5.0: raise HTTPException( status_code=400, @@ -391,10 +424,9 @@ def _validate_predict_form_params( async def _read_fasta_upload(fasta_file: UploadFile) -> str: raw = await fasta_file.read(MAX_FASTA_UPLOAD_BYTES + 1) if len(raw) > MAX_FASTA_UPLOAD_BYTES: - max_mb = MAX_FASTA_UPLOAD_BYTES // (1024 * 1024) raise HTTPException( status_code=413, - detail=f"FASTA_FILE_TOO_LARGE: max {max_mb} MB", + detail=f"FASTA_FILE_TOO_LARGE: max {MAX_FASTA_UPLOAD_MB} MB", ) fasta_text = raw.decode("utf-8", errors="replace") if not fasta_text.strip(): @@ -441,10 +473,14 @@ def _wait_for_job_completion(job_id: str, timeout_seconds: int, poll_interval_se @app.post(API_PREFIX + "/predict-go-from-sequences", response_model=PredictGoResponse) def predict_go_from_sequences(request: PredictGoFromSequencesRequest) -> PredictGoResponse: - _observe_sequence_lengths( - backend=request.backend, - sequences=[seq.sequence for seq in request.sequences], - ) + sequences = [seq.sequence for seq in request.sequences] + _enforce_sequence_envelope(sequences) + if request.timeout_seconds > SYNC_PREDICT_TIMEOUT_SEC: + raise HTTPException( + status_code=400, + detail=f"timeout_seconds must be <= {SYNC_PREDICT_TIMEOUT_SEC}", + ) + _observe_sequence_lengths(backend=request.backend, sequences=sequences) job_payload = { "stage": "test", "backend": request.backend, @@ -468,12 +504,12 @@ async def predict_go_from_fasta( fasta_file: UploadFile = File(...), backend: Literal["esm2", "protbert", "t5"] = Form(default="esm2"), pooling: Literal["mean", "cls"] = Form(default="mean"), - batch_size: int = Form(default=8), - max_length: int = Form(default=1280), + batch_size: int = Form(default=DEFAULT_BATCH_SIZE), + max_length: int = Form(default=DEFAULT_MAX_LENGTH), top_k: int = Form(default=10), fail_fast: bool = Form(default=True), - timeout_seconds: int = Form(default=1800), - poll_interval_seconds: float = Form(default=1.0), + timeout_seconds: int = Form(default=SYNC_PREDICT_TIMEOUT_SEC), + poll_interval_seconds: float = Form(default=SYNC_PREDICT_POLL_INTERVAL_SEC), ) -> PredictGoResponse: _validate_predict_form_params( batch_size=batch_size, diff --git a/services/embedding-api/schemas.py b/services/embedding-api/schemas.py index 976cb22..498393c 100644 --- a/services/embedding-api/schemas.py +++ b/services/embedding-api/schemas.py @@ -4,6 +4,14 @@ from pydantic import BaseModel, Field +from config import ( + DEFAULT_BATCH_SIZE, + DEFAULT_MAX_LENGTH, + MAX_SEQUENCES_PER_REQUEST, + SYNC_PREDICT_POLL_INTERVAL_SEC, + SYNC_PREDICT_TIMEOUT_SEC, +) + class SequenceItem(BaseModel): id: str = Field(min_length=1) @@ -14,9 +22,9 @@ class CreateJobRequest(BaseModel): stage: Literal["test"] = "test" backend: Literal["esm2", "protbert", "t5"] = "esm2" pooling: Literal["mean", "cls"] = "mean" - batch_size: int = Field(default=8, ge=1, le=128) - max_length: int = Field(default=1280, ge=8, le=8192) - sequences: list[SequenceItem] = Field(min_length=1) + batch_size: int = Field(default=DEFAULT_BATCH_SIZE, ge=1, le=128) + max_length: int = Field(default=DEFAULT_MAX_LENGTH, ge=8, le=8192) + sequences: list[SequenceItem] = Field(min_length=1, max_length=MAX_SEQUENCES_PER_REQUEST) class Progress(BaseModel): @@ -78,11 +86,19 @@ class PredictGoResponse(BaseModel): class PredictGoFromSequencesRequest(BaseModel): backend: Literal["esm2", "protbert", "t5"] = "esm2" pooling: Literal["mean", "cls"] = "mean" - batch_size: int = Field(default=8, ge=1, le=128) - max_length: int = Field(default=1280, ge=8, le=8192) - sequences: list[SequenceItem] = Field(min_length=1) + batch_size: int = Field(default=DEFAULT_BATCH_SIZE, ge=1, le=128) + max_length: int = Field(default=DEFAULT_MAX_LENGTH, ge=8, le=8192) + sequences: list[SequenceItem] = Field(min_length=1, max_length=MAX_SEQUENCES_PER_REQUEST) top_k: int = Field(default=10, ge=1, le=500) indices: list[int] | None = None fail_fast: bool = True - timeout_seconds: int = Field(default=1800, ge=5, le=7200) - poll_interval_seconds: float = Field(default=1.0, gt=0.1, le=5.0) \ No newline at end of file + timeout_seconds: int = Field( + default=SYNC_PREDICT_TIMEOUT_SEC, + ge=5, + le=SYNC_PREDICT_TIMEOUT_SEC, + ) + poll_interval_seconds: float = Field( + default=SYNC_PREDICT_POLL_INTERVAL_SEC, + gt=0.1, + le=5.0, + )