bge-reranker/init.py

56 lines
1.8 KiB
Python

# -*- coding:utf-8 -*-
from traceback import format_exc
from ahserver.serverenv import ServerEnv
from appPublic.registerfunction import RegisterFunction
from appPublic.log import debug, exception
import json
async def status_handler(request, params_kw, *args, **kwargs):
"""Status endpoint"""
return json.dumps({
"service": "bge-reranker",
"model": "BAAI/bge-reranker-v2-m3",
"model_path": "/data/ymq/models/BAAI/bge-reranker-v2-m3",
"endpoints": ["/api/status", "/api/rerank"]
}, indent=2, ensure_ascii=False)
async def rerank_handler(request, params_kw, *args, **kwargs):
"""Rerank endpoint"""
try:
query = params_kw.get("query")
documents = params_kw.get("documents", [])
top_k = params_kw.get("top_k", 5)
if not query or not documents:
return json.dumps({"error": "query and documents required"})
if not isinstance(documents, list):
return json.dumps({"error": "documents must be a list"})
from workers.bge_reranker import rerank
import time
start = time.time()
result = rerank(query, documents, top_k)
elapsed = round(time.time() - start, 4)
return json.dumps({
"status": "SUCCEEDED",
"query": query,
"ranked_docs": result["ranked_docs"],
"total": result["total"],
"returned": result["returned"],
"elapsed": elapsed
}, ensure_ascii=False)
except Exception as e:
exception(f"{e}, {format_exc()}")
return json.dumps({"error": str(e)})
def load_bge_reranker():
"""Register API handlers"""
env = ServerEnv()
rf = RegisterFunction()
rf.register("status", status_handler)
rf.register("rerank", rerank_handler)