58 lines
1.6 KiB
Python
58 lines
1.6 KiB
Python
import os
|
|
from typing import Any
|
|
|
|
import httpx
|
|
from fastapi import FastAPI, HTTPException
|
|
from pydantic import BaseModel, Field
|
|
|
|
|
|
TEI_URL = os.getenv("TEI_URL", "http://tei-rerank:80").rstrip("/")
|
|
app = FastAPI(title="WeKnora TEI Rerank Adapter")
|
|
|
|
|
|
class RerankRequest(BaseModel):
|
|
model: str = "bge-reranker-base"
|
|
query: str
|
|
documents: list[str] = Field(min_length=1)
|
|
additional_data: dict[str, Any] = Field(default_factory=dict)
|
|
truncate_prompt_tokens: int = 511
|
|
|
|
|
|
@app.get("/health")
|
|
async def health() -> dict[str, str]:
|
|
return {"status": "ok"}
|
|
|
|
|
|
@app.post("/rerank")
|
|
async def rerank(request: RerankRequest) -> dict[str, Any]:
|
|
payload = {
|
|
"query": request.query,
|
|
"texts": request.documents,
|
|
"raw_scores": False,
|
|
"return_text": False,
|
|
}
|
|
try:
|
|
async with httpx.AsyncClient(timeout=60) as client:
|
|
response = await client.post(f"{TEI_URL}/rerank", json=payload)
|
|
response.raise_for_status()
|
|
scores = response.json()
|
|
except httpx.HTTPError as exc:
|
|
raise HTTPException(status_code=502, detail=f"TEI rerank request failed: {exc}") from exc
|
|
|
|
results = []
|
|
for item in scores:
|
|
index = int(item["index"])
|
|
results.append(
|
|
{
|
|
"index": index,
|
|
"relevance_score": float(item["score"]),
|
|
"document": {"text": request.documents[index]},
|
|
}
|
|
)
|
|
return {
|
|
"id": "local-tei-rerank",
|
|
"model": request.model,
|
|
"usage": {"total_tokens": 0},
|
|
"results": results,
|
|
}
|