diff --git a/src/app/main.py b/src/app/main.py index 4c5c3d8..a3d2bb1 100644 --- a/src/app/main.py +++ b/src/app/main.py @@ -23,7 +23,7 @@ RerankRequest, RerankResponse, ) -from .models import get_model as get_model, get_model_or_400 +from .models import get_model as get_model from .config import ( EMBEDDING_MODELS, RERANK_MODELS, @@ -37,6 +37,7 @@ EmbeddingService, RerankService, ) +from .services.base import get_validated_model from .services.embedding import ( _determine_ruri_prefix as _determine_ruri_prefix, _apply_prefix as _apply_prefix, @@ -173,16 +174,18 @@ def _get_model_or_400(model_name: str, model_type: str) -> Any: Helper for backwards compatibility with legacy tests calling _get_model_or_400. """ supported_models = EMBEDDING_MODELS if model_type == "embedding" else RERANK_MODELS - return get_model_or_400(model_name, supported_models, model_type) + return get_validated_model( + model_name, supported_models, model_type, loader=get_model + ) # Dependency Injection Providers def get_embedding_service() -> BaseEmbeddingService: - return EmbeddingService(proxy_to_tei_func=_proxy_to_tei) + return EmbeddingService(proxy_to_tei_func=_proxy_to_tei, model_loader=get_model) def get_rerank_service() -> BaseRerankService: - return RerankService(proxy_to_tei_func=_proxy_to_tei) + return RerankService(proxy_to_tei_func=_proxy_to_tei, model_loader=get_model) @app.post( diff --git a/src/app/models.py b/src/app/models.py index e30b007..492495c 100644 --- a/src/app/models.py +++ b/src/app/models.py @@ -120,33 +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, valid_models: Any, model_type: str = "model" -) -> Any: - """ - Validates model_name against valid_models and fetches model instance via get_model. - Raises HTTPException(400) if model is invalid or if get_model raises a ValueError. - """ - from fastapi import HTTPException - - if model_name not in valid_models: - suffix = "s" if not model_type.endswith("s") else "" - raise HTTPException( - status_code=400, - detail=f"Model '{model_name}' not found for {model_type}{suffix}.", - ) - - try: - # Dynamically import main to support patches on app.main.get_model in tests - try: - import app.main as main_mod - - fetch_func = getattr(main_mod, "get_model", get_model) - except Exception: - fetch_func = get_model - - return fetch_func(model_name) - except ValueError as e: - raise HTTPException(status_code=400, detail=str(e)) diff --git a/src/app/services/base.py b/src/app/services/base.py index 3a72d2d..1934b1e 100644 --- a/src/app/services/base.py +++ b/src/app/services/base.py @@ -1,5 +1,33 @@ from abc import ABC, abstractmethod +from typing import Any, Callable, Collection, Optional +from fastapi import HTTPException from ..schemas import EmbeddingRequest, EmbeddingResponse, RerankRequest, RerankResponse +from ..models import get_model + + +def get_validated_model( + model_name: str, + allowed_models: Collection[str], + service_name: str, + loader: Optional[Callable[[str], Any]] = None, +) -> Any: + """ + Validates model_name against allowed_models and loads the model via loader/get_model. + Raises HTTPException(400) if model is invalid or if loading raises a ValueError. + """ + if model_name not in allowed_models: + suffix = "s" if not service_name.endswith("s") else "" + raise HTTPException( + status_code=400, + detail=f"Model '{model_name}' not found for {service_name}{suffix}.", + ) + + fetch_func = loader or get_model + + try: + return fetch_func(model_name) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) class BaseEmbeddingService(ABC): diff --git a/src/app/services/embedding.py b/src/app/services/embedding.py index c74f253..cea8cf9 100644 --- a/src/app/services/embedding.py +++ b/src/app/services/embedding.py @@ -5,9 +5,8 @@ from PIL import Image from fastapi import HTTPException -from .base import BaseEmbeddingService +from .base import BaseEmbeddingService, get_validated_model from ..image_utils import load_image_from_source -from ..models import get_model_or_400 from ..schemas import ( EmbeddingRequest, EmbeddingResponse, @@ -143,17 +142,18 @@ def _tokenize_and_truncate_embeddings( return processed_inputs, usage -def _get_model_or_400(model_name: str) -> Any: - return get_model_or_400(model_name, EMBEDDING_MODELS, "embedding") - - class EmbeddingService(BaseEmbeddingService): """ Default production implementation of BaseEmbeddingService. """ - def __init__(self, proxy_to_tei_func: Optional[Any] = None): + def __init__( + self, + proxy_to_tei_func: Optional[Any] = None, + model_loader: Optional[Any] = None, + ): self.proxy_to_tei_func = proxy_to_tei_func + self.model_loader = model_loader async def create_embeddings(self, request: EmbeddingRequest) -> EmbeddingResponse: import app.main as main_mod @@ -192,7 +192,12 @@ async def create_embeddings(self, request: EmbeddingRequest) -> EmbeddingRespons ) return EmbeddingResponse(**data) - model = _get_model_or_400(request.model) + model = get_validated_model( + request.model, + EMBEDDING_MODELS, + "embedding", + loader=self.model_loader, + ) is_multimodal = getattr(model, "supports_multimodal", False) is True if has_image and not is_multimodal: diff --git a/src/app/services/rerank.py b/src/app/services/rerank.py index 6aab873..00fc5c2 100644 --- a/src/app/services/rerank.py +++ b/src/app/services/rerank.py @@ -2,8 +2,7 @@ from typing import Any, List, Optional from fastapi import HTTPException -from .base import BaseRerankService -from ..models import get_model_or_400 +from .base import BaseRerankService, get_validated_model from ..schemas import RerankRequest, RerankResponse, RerankData, Usage from ..config import RERANK_MODELS @@ -47,17 +46,18 @@ def _sort_and_format_rerank_results( return [RerankData(**result) for result in sorted_results] -def _get_model_or_400(model_name: str) -> Any: - return get_model_or_400(model_name, RERANK_MODELS, "rerank") - - class RerankService(BaseRerankService): """ Default production implementation of BaseRerankService. """ - def __init__(self, proxy_to_tei_func: Optional[Any] = None): + def __init__( + self, + proxy_to_tei_func: Optional[Any] = None, + model_loader: Optional[Any] = None, + ): self.proxy_to_tei_func = proxy_to_tei_func + self.model_loader = model_loader async def create_rerank(self, request: RerankRequest) -> RerankResponse: import app.main as main_mod @@ -95,7 +95,12 @@ async def create_rerank(self, request: RerankRequest) -> RerankResponse: usage=usage, ) - model = _get_model_or_400(request.model) + model = get_validated_model( + request.model, + RERANK_MODELS, + "rerank", + loader=self.model_loader, + ) pairs = [[request.query, doc] for doc in request.documents] usage = _calculate_rerank_tokens(model, request.query, request.documents)