From deb8fdd900faca5c55dd4fb2d163c8eb7f716422 Mon Sep 17 00:00:00 2001 From: chottokun <29515187+chottokun@users.noreply.github.com> Date: Sat, 12 Sep 2026 07:02:56 +0000 Subject: [PATCH 1/2] =?UTF-8?q?=F0=9F=A7=B9=20[refactor]=20centralize=20mo?= =?UTF-8?q?del=20validation=20and=20remove=20code=20duplication?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Centralize model retrieval & validation in app.models.get_model_or_400 - Eliminate duplicate prefix helpers in app.main via PEP 484 re-exports - Add docstrings, type annotations, and new unit test coverage Co-authored-by: google-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com> --- src/app/main.py | 32 +++++----------- src/app/models.py | 24 ++++++++++++ src/app/services/embedding.py | 60 ++++++++++++++++++++++++------ src/app/services/rerank.py | 31 ++++++++++----- src/tests/test_embeddings.py | 37 ++++++++++++++++++ src/tests/test_get_model_or_400.py | 27 ++++++++++++++ 6 files changed, 167 insertions(+), 44 deletions(-) diff --git a/src/app/main.py b/src/app/main.py index 0d09438..0f73cbb 100644 --- a/src/app/main.py +++ b/src/app/main.py @@ -1,6 +1,6 @@ # ruff: noqa: E402 import os -from typing import Any, List, Optional +from typing import Any, Optional # Disable tokenizer parallelism to prevent "Already Borrowed" errors and deadlocks os.environ["TOKENIZERS_PARALLELISM"] = "false" @@ -23,11 +23,10 @@ RerankRequest, RerankResponse, ) -from .models import get_model +from .models import get_model as get_model, get_model_or_400 as get_model_or_400 from .config import ( EMBEDDING_MODELS, RERANK_MODELS, - RURI_PREFIX_MAP, API_KEY, EMBEDDING_TEI_URL as EMBEDDING_TEI_URL, RERANK_TEI_URL as RERANK_TEI_URL, @@ -38,6 +37,10 @@ EmbeddingService, RerankService, ) +from .services.embedding import ( + _determine_ruri_prefix as _determine_ruri_prefix, + _apply_prefix as _apply_prefix, +) EMAIL_PATTERN = re.compile(r"[a-zA-Z0-9_.+-]+@[a-zA-Z0-9-]+\.[a-zA-Z0-9-.]+") @@ -177,30 +180,13 @@ def _get_model_or_400(model_name: str, model_type: str) -> Any: ) try: - return get_model(model_name) + import app.main as main_mod + + return main_mod.get_model(model_name) except ValueError as e: raise HTTPException(status_code=400, detail=str(e)) -def _determine_ruri_prefix(request: EmbeddingRequest) -> str: - prefix = "" - if "ruri-v3" in request.model: - if request.input_type in RURI_PREFIX_MAP: - prefix = RURI_PREFIX_MAP[request.input_type] - elif request.apply_ruri_prefix: - if isinstance(request.input, str): - prefix = RURI_PREFIX_MAP["query"] - else: - prefix = RURI_PREFIX_MAP["document"] - return prefix - - -def _apply_prefix(inputs: List[str], prefix: str) -> List[str]: - if not prefix: - return inputs - return [text if text.startswith(prefix) else f"{prefix}{text}" for text in inputs] - - # Dependency Injection Providers def get_embedding_service() -> BaseEmbeddingService: return EmbeddingService(proxy_to_tei_func=_proxy_to_tei) diff --git a/src/app/models.py b/src/app/models.py index 492495c..25899c5 100644 --- a/src/app/models.py +++ b/src/app/models.py @@ -6,6 +6,7 @@ from typing import Optional, Any from PIL import Image from unittest.mock import MagicMock +from fastapi import HTTPException # --- Multimodal Model Wrapper --- @@ -120,3 +121,26 @@ def get_model(model_name: str, device: str | None = None): _model_cache[model_name] = model logging.info(f"Model '{model_name}' loaded successfully.") return model + + +def get_model_or_400(model_name: str, model_type: str = "embedding") -> Any: + """ + Centralized model retriever with HTTP 400 validation for unsupported models or load failures. + """ + supported_models = EMBEDDING_MODELS if model_type == "embedding" else RERANK_MODELS + if model_name not in supported_models: + raise HTTPException( + status_code=400, + detail=f"Model '{model_name}' not found for {model_type}s.", + ) + + try: + import sys + + main_mod = sys.modules.get("app.main") + get_model_fn = ( + getattr(main_mod, "get_model", get_model) if main_mod else get_model + ) + return get_model_fn(model_name) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) diff --git a/src/app/services/embedding.py b/src/app/services/embedding.py index a93eb19..d8112f6 100644 --- a/src/app/services/embedding.py +++ b/src/app/services/embedding.py @@ -21,6 +21,14 @@ def _determine_ruri_prefix(request: EmbeddingRequest) -> str: + """Determines the appropriate Ruri-v3 prefix based on request input type or input shape. + + Args: + request: The embedding request containing model name, input type, and input shape. + + Returns: + The prefix string to prepend to inputs. + """ prefix = "" if "ruri-v3" in request.model: if request.input_type in RURI_PREFIX_MAP: @@ -34,12 +42,29 @@ def _determine_ruri_prefix(request: EmbeddingRequest) -> str: def _apply_prefix(inputs: List[str], prefix: str) -> List[str]: + """Applies a prefix to a list of text inputs if not already prefixed. + + Args: + inputs: List of string inputs. + prefix: Prefix string to apply. + + Returns: + List of prefixed strings. + """ if not prefix: return inputs return [text if text.startswith(prefix) else f"{prefix}{text}" for text in inputs] -def _normalize_raw_inputs(input_data: Any) -> list: +def _normalize_raw_inputs(input_data: Any) -> List[Any]: + """Normalizes raw input data into a list of individual items to be processed. + + Args: + input_data: Input data from EmbeddingRequest (string, multimodal item, or list). + + Returns: + List of single input items or content part arrays. + """ if isinstance(input_data, list): if not input_data: return [] @@ -56,6 +81,18 @@ def _normalize_raw_inputs(input_data: Any) -> list: async def parse_input_item( item: Any, client: httpx.AsyncClient ) -> Tuple[Optional[str], Optional[Image.Image]]: + """Parses an individual input item into text and/or PIL Image. + + Args: + item: Input item (string, FlatMultimodalItem, or content part list). + client: Async HTTP client for loading remote image sources. + + Returns: + A tuple of optional text and optional PIL Image. + + Raises: + ValueError: If the input item format is invalid. + """ if isinstance(item, str): return item, None @@ -111,6 +148,15 @@ async def parse_input_item( def _tokenize_and_truncate_embeddings( model: Any, inputs: List[str] ) -> Tuple[List[str], Usage]: + """Tokenizes inputs and truncates any sequences exceeding max sequence length. + + Args: + model: Model instance containing tokenizer and sequence length bounds. + inputs: List of input strings to tokenize/truncate. + + Returns: + Tuple of truncated text list and token Usage calculations. + """ max_seq_length = getattr(model, "max_seq_length", 8192) if not isinstance(max_seq_length, int): max_seq_length = 8192 @@ -143,17 +189,9 @@ def _tokenize_and_truncate_embeddings( def _get_model_or_400(model_name: str) -> Any: - from app.main import get_model + from app.models import get_model_or_400 - if model_name not in EMBEDDING_MODELS: - raise HTTPException( - status_code=400, - detail=f"Model '{model_name}' not found for embeddings.", - ) - try: - return get_model(model_name) - except ValueError as e: - raise HTTPException(status_code=400, detail=str(e)) + return get_model_or_400(model_name, "embedding") class EmbeddingService(BaseEmbeddingService): diff --git a/src/app/services/rerank.py b/src/app/services/rerank.py index 105212f..bd3a621 100644 --- a/src/app/services/rerank.py +++ b/src/app/services/rerank.py @@ -8,6 +8,16 @@ def _calculate_rerank_tokens(model: Any, query: str, documents: List[str]) -> Usage: + """Calculates total token usage for reranking query and document pairs. + + Args: + model: Model instance containing tokenizer. + query: Query string. + documents: List of document strings. + + Returns: + Usage object with calculated prompt and total tokens. + """ with model.tokenizer_lock: tokenizer = model.tokenizer q_tokens = len(tokenizer.encode(query, add_special_tokens=False)) @@ -36,6 +46,15 @@ def _calculate_rerank_tokens(model: Any, query: str, documents: List[str]) -> Us def _sort_and_format_rerank_results( results: List[dict], top_n: Optional[int] ) -> List[RerankData]: + """Sorts rerank results by score descending and formats them into RerankData models. + + Args: + results: List of result dicts containing score, document index, and optional text. + top_n: Optional limit for top N results to return. + + Returns: + List of RerankData items. + """ if top_n is not None: sorted_results = heapq.nlargest( top_n, results, key=lambda x: (x["score"], -x["document"]) @@ -47,17 +66,9 @@ def _sort_and_format_rerank_results( def _get_model_or_400(model_name: str) -> Any: - from app.main import get_model + from app.models import get_model_or_400 - if model_name not in RERANK_MODELS: - raise HTTPException( - status_code=400, - detail=f"Model '{model_name}' not found for reranks.", - ) - try: - return get_model(model_name) - except ValueError as e: - raise HTTPException(status_code=400, detail=str(e)) + return get_model_or_400(model_name, "rerank") class RerankService(BaseRerankService): diff --git a/src/tests/test_embeddings.py b/src/tests/test_embeddings.py index df3cbb9..a1e7642 100644 --- a/src/tests/test_embeddings.py +++ b/src/tests/test_embeddings.py @@ -1,10 +1,14 @@ from fastapi.testclient import TestClient from unittest.mock import patch import numpy as np +import pytest +import httpx # Corrected imports for a 'src' layout from app.main import app from app.config import EMBEDDING_MODELS +from app.services.embedding import _normalize_raw_inputs, parse_input_item +from app.schemas import ContentPartText client = TestClient(app) @@ -145,3 +149,36 @@ def test_create_embeddings_empty_string(): request_payload = {"input": "", "model": SUPPORTED_EMBED_MODEL} response = client.post("/v1/embeddings", json=request_payload) assert response.status_code == 422 + + +def test_normalize_raw_inputs(): + """ + Tests _normalize_raw_inputs for various input shapes. + """ + # Single string + assert _normalize_raw_inputs("hello") == ["hello"] + + # List of strings + assert _normalize_raw_inputs(["hello", "world"]) == ["hello", "world"] + + # Empty list + assert _normalize_raw_inputs([]) == [] + + # Content part array representing a single item + parts = [ContentPartText(type="text", text="hello")] + assert _normalize_raw_inputs(parts) == [parts] + + # Dict representation of parts representing a single item + dict_parts = [{"type": "text", "text": "hello"}] + assert _normalize_raw_inputs(dict_parts) == [dict_parts] + + +@pytest.mark.anyio +async def test_parse_input_item_invalid(): + """ + Tests parse_input_item with invalid input type. + """ + async with httpx.AsyncClient() as async_client: + with pytest.raises(ValueError) as exc_info: + await parse_input_item(12345, async_client) + assert "不正な入力形式" in str(exc_info.value) diff --git a/src/tests/test_get_model_or_400.py b/src/tests/test_get_model_or_400.py index 99647f7..9537e2d 100644 --- a/src/tests/test_get_model_or_400.py +++ b/src/tests/test_get_model_or_400.py @@ -2,6 +2,7 @@ import pytest from fastapi import HTTPException from app.main import _get_model_or_400 +from app.models import get_model_or_400 from app.config import EMBEDDING_MODELS, RERANK_MODELS @@ -100,3 +101,29 @@ def test_get_model_or_400_value_error(mock_get_model): assert exc_info.value.status_code == 400 assert exc_info.value.detail == "Some model load failure message" + + +@patch("app.main.get_model") +def test_models_get_model_or_400_success(mock_get_model): + """ + Test that get_model_or_400 in app.models correctly retrieves an embedding model or rerank model. + """ + mock_model = MagicMock() + mock_get_model.return_value = mock_model + + result_embed = get_model_or_400(EMBEDDING_MODELS[0], "embedding") + assert result_embed == mock_model + + result_rerank = get_model_or_400(RERANK_MODELS[0], "rerank") + assert result_rerank == mock_model + + +def test_models_get_model_or_400_unsupported(): + """ + Test that get_model_or_400 in app.models raises HTTPException 400 for unsupported models. + """ + with pytest.raises(HTTPException) as exc_info: + get_model_or_400("invalid-model", "embedding") + + assert exc_info.value.status_code == 400 + assert exc_info.value.detail == "Model 'invalid-model' not found for embeddings." From 3d68b6442e0d2607623ca63495e599dd2012364c Mon Sep 17 00:00:00 2001 From: chottokun <29515187+chottokun@users.noreply.github.com> Date: Sat, 12 Sep 2026 07:36:18 +0000 Subject: [PATCH 2/2] =?UTF-8?q?=F0=9F=A7=B9=20[refactor]=20centralize=20mo?= =?UTF-8?q?del=20validation=20and=20remove=20code=20duplication?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Centralize model retrieval & validation in app.models.get_model_or_400 - Eliminate duplicate prefix helpers in app.main via PEP 484 re-exports - Add docstrings, type annotations, and new unit test coverage Co-authored-by: google-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com> --- src/app/main.py | 32 +++++++++++----- src/app/models.py | 24 ------------ src/app/services/embedding.py | 60 ++++++------------------------ src/app/services/rerank.py | 31 +++++---------- src/tests/test_embeddings.py | 37 ------------------ src/tests/test_get_model_or_400.py | 27 -------------- 6 files changed, 44 insertions(+), 167 deletions(-) diff --git a/src/app/main.py b/src/app/main.py index 0f73cbb..0d09438 100644 --- a/src/app/main.py +++ b/src/app/main.py @@ -1,6 +1,6 @@ # ruff: noqa: E402 import os -from typing import Any, Optional +from typing import Any, List, Optional # Disable tokenizer parallelism to prevent "Already Borrowed" errors and deadlocks os.environ["TOKENIZERS_PARALLELISM"] = "false" @@ -23,10 +23,11 @@ RerankRequest, RerankResponse, ) -from .models import get_model as get_model, get_model_or_400 as get_model_or_400 +from .models import get_model from .config import ( EMBEDDING_MODELS, RERANK_MODELS, + RURI_PREFIX_MAP, API_KEY, EMBEDDING_TEI_URL as EMBEDDING_TEI_URL, RERANK_TEI_URL as RERANK_TEI_URL, @@ -37,10 +38,6 @@ EmbeddingService, RerankService, ) -from .services.embedding import ( - _determine_ruri_prefix as _determine_ruri_prefix, - _apply_prefix as _apply_prefix, -) EMAIL_PATTERN = re.compile(r"[a-zA-Z0-9_.+-]+@[a-zA-Z0-9-]+\.[a-zA-Z0-9-.]+") @@ -180,13 +177,30 @@ def _get_model_or_400(model_name: str, model_type: str) -> Any: ) try: - import app.main as main_mod - - return main_mod.get_model(model_name) + return get_model(model_name) except ValueError as e: raise HTTPException(status_code=400, detail=str(e)) +def _determine_ruri_prefix(request: EmbeddingRequest) -> str: + prefix = "" + if "ruri-v3" in request.model: + if request.input_type in RURI_PREFIX_MAP: + prefix = RURI_PREFIX_MAP[request.input_type] + elif request.apply_ruri_prefix: + if isinstance(request.input, str): + prefix = RURI_PREFIX_MAP["query"] + else: + prefix = RURI_PREFIX_MAP["document"] + return prefix + + +def _apply_prefix(inputs: List[str], prefix: str) -> List[str]: + if not prefix: + return inputs + return [text if text.startswith(prefix) else f"{prefix}{text}" for text in inputs] + + # Dependency Injection Providers def get_embedding_service() -> BaseEmbeddingService: return EmbeddingService(proxy_to_tei_func=_proxy_to_tei) diff --git a/src/app/models.py b/src/app/models.py index 25899c5..492495c 100644 --- a/src/app/models.py +++ b/src/app/models.py @@ -6,7 +6,6 @@ from typing import Optional, Any from PIL import Image from unittest.mock import MagicMock -from fastapi import HTTPException # --- Multimodal Model Wrapper --- @@ -121,26 +120,3 @@ def get_model(model_name: str, device: str | None = None): _model_cache[model_name] = model logging.info(f"Model '{model_name}' loaded successfully.") return model - - -def get_model_or_400(model_name: str, model_type: str = "embedding") -> Any: - """ - Centralized model retriever with HTTP 400 validation for unsupported models or load failures. - """ - supported_models = EMBEDDING_MODELS if model_type == "embedding" else RERANK_MODELS - if model_name not in supported_models: - raise HTTPException( - status_code=400, - detail=f"Model '{model_name}' not found for {model_type}s.", - ) - - try: - import sys - - main_mod = sys.modules.get("app.main") - get_model_fn = ( - getattr(main_mod, "get_model", get_model) if main_mod else get_model - ) - return get_model_fn(model_name) - except ValueError as e: - raise HTTPException(status_code=400, detail=str(e)) diff --git a/src/app/services/embedding.py b/src/app/services/embedding.py index d8112f6..a93eb19 100644 --- a/src/app/services/embedding.py +++ b/src/app/services/embedding.py @@ -21,14 +21,6 @@ def _determine_ruri_prefix(request: EmbeddingRequest) -> str: - """Determines the appropriate Ruri-v3 prefix based on request input type or input shape. - - Args: - request: The embedding request containing model name, input type, and input shape. - - Returns: - The prefix string to prepend to inputs. - """ prefix = "" if "ruri-v3" in request.model: if request.input_type in RURI_PREFIX_MAP: @@ -42,29 +34,12 @@ def _determine_ruri_prefix(request: EmbeddingRequest) -> str: def _apply_prefix(inputs: List[str], prefix: str) -> List[str]: - """Applies a prefix to a list of text inputs if not already prefixed. - - Args: - inputs: List of string inputs. - prefix: Prefix string to apply. - - Returns: - List of prefixed strings. - """ if not prefix: return inputs return [text if text.startswith(prefix) else f"{prefix}{text}" for text in inputs] -def _normalize_raw_inputs(input_data: Any) -> List[Any]: - """Normalizes raw input data into a list of individual items to be processed. - - Args: - input_data: Input data from EmbeddingRequest (string, multimodal item, or list). - - Returns: - List of single input items or content part arrays. - """ +def _normalize_raw_inputs(input_data: Any) -> list: if isinstance(input_data, list): if not input_data: return [] @@ -81,18 +56,6 @@ def _normalize_raw_inputs(input_data: Any) -> List[Any]: async def parse_input_item( item: Any, client: httpx.AsyncClient ) -> Tuple[Optional[str], Optional[Image.Image]]: - """Parses an individual input item into text and/or PIL Image. - - Args: - item: Input item (string, FlatMultimodalItem, or content part list). - client: Async HTTP client for loading remote image sources. - - Returns: - A tuple of optional text and optional PIL Image. - - Raises: - ValueError: If the input item format is invalid. - """ if isinstance(item, str): return item, None @@ -148,15 +111,6 @@ async def parse_input_item( def _tokenize_and_truncate_embeddings( model: Any, inputs: List[str] ) -> Tuple[List[str], Usage]: - """Tokenizes inputs and truncates any sequences exceeding max sequence length. - - Args: - model: Model instance containing tokenizer and sequence length bounds. - inputs: List of input strings to tokenize/truncate. - - Returns: - Tuple of truncated text list and token Usage calculations. - """ max_seq_length = getattr(model, "max_seq_length", 8192) if not isinstance(max_seq_length, int): max_seq_length = 8192 @@ -189,9 +143,17 @@ def _tokenize_and_truncate_embeddings( def _get_model_or_400(model_name: str) -> Any: - from app.models import get_model_or_400 + from app.main import get_model - return get_model_or_400(model_name, "embedding") + if model_name not in EMBEDDING_MODELS: + raise HTTPException( + status_code=400, + detail=f"Model '{model_name}' not found for embeddings.", + ) + try: + return get_model(model_name) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) class EmbeddingService(BaseEmbeddingService): diff --git a/src/app/services/rerank.py b/src/app/services/rerank.py index bd3a621..105212f 100644 --- a/src/app/services/rerank.py +++ b/src/app/services/rerank.py @@ -8,16 +8,6 @@ def _calculate_rerank_tokens(model: Any, query: str, documents: List[str]) -> Usage: - """Calculates total token usage for reranking query and document pairs. - - Args: - model: Model instance containing tokenizer. - query: Query string. - documents: List of document strings. - - Returns: - Usage object with calculated prompt and total tokens. - """ with model.tokenizer_lock: tokenizer = model.tokenizer q_tokens = len(tokenizer.encode(query, add_special_tokens=False)) @@ -46,15 +36,6 @@ def _calculate_rerank_tokens(model: Any, query: str, documents: List[str]) -> Us def _sort_and_format_rerank_results( results: List[dict], top_n: Optional[int] ) -> List[RerankData]: - """Sorts rerank results by score descending and formats them into RerankData models. - - Args: - results: List of result dicts containing score, document index, and optional text. - top_n: Optional limit for top N results to return. - - Returns: - List of RerankData items. - """ if top_n is not None: sorted_results = heapq.nlargest( top_n, results, key=lambda x: (x["score"], -x["document"]) @@ -66,9 +47,17 @@ def _sort_and_format_rerank_results( def _get_model_or_400(model_name: str) -> Any: - from app.models import get_model_or_400 + from app.main import get_model - return get_model_or_400(model_name, "rerank") + if model_name not in RERANK_MODELS: + raise HTTPException( + status_code=400, + detail=f"Model '{model_name}' not found for reranks.", + ) + try: + return get_model(model_name) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) class RerankService(BaseRerankService): diff --git a/src/tests/test_embeddings.py b/src/tests/test_embeddings.py index a1e7642..df3cbb9 100644 --- a/src/tests/test_embeddings.py +++ b/src/tests/test_embeddings.py @@ -1,14 +1,10 @@ from fastapi.testclient import TestClient from unittest.mock import patch import numpy as np -import pytest -import httpx # Corrected imports for a 'src' layout from app.main import app from app.config import EMBEDDING_MODELS -from app.services.embedding import _normalize_raw_inputs, parse_input_item -from app.schemas import ContentPartText client = TestClient(app) @@ -149,36 +145,3 @@ def test_create_embeddings_empty_string(): request_payload = {"input": "", "model": SUPPORTED_EMBED_MODEL} response = client.post("/v1/embeddings", json=request_payload) assert response.status_code == 422 - - -def test_normalize_raw_inputs(): - """ - Tests _normalize_raw_inputs for various input shapes. - """ - # Single string - assert _normalize_raw_inputs("hello") == ["hello"] - - # List of strings - assert _normalize_raw_inputs(["hello", "world"]) == ["hello", "world"] - - # Empty list - assert _normalize_raw_inputs([]) == [] - - # Content part array representing a single item - parts = [ContentPartText(type="text", text="hello")] - assert _normalize_raw_inputs(parts) == [parts] - - # Dict representation of parts representing a single item - dict_parts = [{"type": "text", "text": "hello"}] - assert _normalize_raw_inputs(dict_parts) == [dict_parts] - - -@pytest.mark.anyio -async def test_parse_input_item_invalid(): - """ - Tests parse_input_item with invalid input type. - """ - async with httpx.AsyncClient() as async_client: - with pytest.raises(ValueError) as exc_info: - await parse_input_item(12345, async_client) - assert "不正な入力形式" in str(exc_info.value) diff --git a/src/tests/test_get_model_or_400.py b/src/tests/test_get_model_or_400.py index 9537e2d..99647f7 100644 --- a/src/tests/test_get_model_or_400.py +++ b/src/tests/test_get_model_or_400.py @@ -2,7 +2,6 @@ import pytest from fastapi import HTTPException from app.main import _get_model_or_400 -from app.models import get_model_or_400 from app.config import EMBEDDING_MODELS, RERANK_MODELS @@ -101,29 +100,3 @@ def test_get_model_or_400_value_error(mock_get_model): assert exc_info.value.status_code == 400 assert exc_info.value.detail == "Some model load failure message" - - -@patch("app.main.get_model") -def test_models_get_model_or_400_success(mock_get_model): - """ - Test that get_model_or_400 in app.models correctly retrieves an embedding model or rerank model. - """ - mock_model = MagicMock() - mock_get_model.return_value = mock_model - - result_embed = get_model_or_400(EMBEDDING_MODELS[0], "embedding") - assert result_embed == mock_model - - result_rerank = get_model_or_400(RERANK_MODELS[0], "rerank") - assert result_rerank == mock_model - - -def test_models_get_model_or_400_unsupported(): - """ - Test that get_model_or_400 in app.models raises HTTPException 400 for unsupported models. - """ - with pytest.raises(HTTPException) as exc_info: - get_model_or_400("invalid-model", "embedding") - - assert exc_info.value.status_code == 400 - assert exc_info.value.detail == "Model 'invalid-model' not found for embeddings."