From 739349139fff15129cbb504cb709c799a18b51d0 Mon Sep 17 00:00:00 2001 From: chottokun <29515187+chottokun@users.noreply.github.com> Date: Sat, 12 Sep 2026 07:58:32 +0000 Subject: [PATCH 1/2] =?UTF-8?q?=F0=9F=A7=B9=20refactor(architecture):=20de?= =?UTF-8?q?couple=20models=20layer=20from=20web=20framework=20and=20move?= =?UTF-8?q?=20validation=20to=20service=20layer?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: google-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com> --- src/app/main.py | 5 +++-- src/app/models.py | 30 ---------------------------- src/app/services/base.py | 37 +++++++++++++++++++++++++++++++++++ src/app/services/embedding.py | 21 ++++++++++++-------- src/app/services/rerank.py | 21 ++++++++++++-------- 5 files changed, 66 insertions(+), 48 deletions(-) diff --git a/src/app/main.py b/src/app/main.py index 4c5c3d8..abb261b 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,7 +174,7 @@ 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) # Dependency Injection Providers 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..f8bc225 100644 --- a/src/app/services/base.py +++ b/src/app/services/base.py @@ -1,5 +1,42 @@ 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}.", + ) + + # If no explicit loader passed, try importing get_model from main_mod to support test patches on app.main.get_model + if loader is None: + try: + import app.main as main_mod + + fetch_func = getattr(main_mod, "get_model", get_model) + except Exception: + fetch_func = get_model + else: + fetch_func = loader + + 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) From 5d190bec0ee5d3b7fd1285bcca25cbc49f76e6c2 Mon Sep 17 00:00:00 2001 From: chottokun <29515187+chottokun@users.noreply.github.com> Date: Sat, 12 Sep 2026 08:17:20 +0000 Subject: [PATCH 2/2] =?UTF-8?q?=F0=9F=A7=B9=20refactor(architecture):=20el?= =?UTF-8?q?iminate=20dynamic=20reflection=20in=20service=20layer=20via=20d?= =?UTF-8?q?ependency=20injection?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: google-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com> --- src/app/main.py | 8 +++++--- src/app/services/base.py | 11 +---------- 2 files changed, 6 insertions(+), 13 deletions(-) diff --git a/src/app/main.py b/src/app/main.py index abb261b..a3d2bb1 100644 --- a/src/app/main.py +++ b/src/app/main.py @@ -174,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_validated_model(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/services/base.py b/src/app/services/base.py index f8bc225..1934b1e 100644 --- a/src/app/services/base.py +++ b/src/app/services/base.py @@ -22,16 +22,7 @@ def get_validated_model( detail=f"Model '{model_name}' not found for {service_name}{suffix}.", ) - # If no explicit loader passed, try importing get_model from main_mod to support test patches on app.main.get_model - if loader is None: - try: - import app.main as main_mod - - fetch_func = getattr(main_mod, "get_model", get_model) - except Exception: - fetch_func = get_model - else: - fetch_func = loader + fetch_func = loader or get_model try: return fetch_func(model_name)