Skip to content
Closed
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
23 changes: 4 additions & 19 deletions src/app/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,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-.]+")

Expand Down Expand Up @@ -182,25 +186,6 @@ def _get_model_or_400(model_name: str, model_type: str) -> Any:
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)
Expand Down
Loading