Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 7 additions & 4 deletions src/app/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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(
Expand Down
30 changes: 0 additions & 30 deletions src/app/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))
28 changes: 28 additions & 0 deletions src/app/services/base.py
Original file line number Diff line number Diff line change
@@ -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):
Expand Down
21 changes: 13 additions & 8 deletions src/app/services/embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down
21 changes: 13 additions & 8 deletions src/app/services/rerank.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
Loading