Files
ai-agent-admin/backend-fastapi/ai_platform/knowledge/services/rerank_service.py
T
2026-06-08 18:14:59 +08:00

183 lines
5.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
Rerank 重排序服务
通过 Rerank 模型对检索结果进行重新排序,提升检索质量。
支持 Jina/Cohere 风格的 Rerank API(大多数提供商兼容此接口)。
"""
import logging
from dataclasses import dataclass
from typing import List, Optional
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
logger = logging.getLogger(__name__)
@dataclass
class RerankResult:
"""重排序结果"""
index: int
relevance_score: float
class RerankService:
"""
Rerank 重排序服务
通过模型 ID 获取对应的提供商,调用 Rerank API 对文档进行重排序。
支持两种 API 风格:
- Jina/Cohere 风格:POST /v1/rerank
- OpenAI 兼容风格(部分提供商)
"""
def __init__(self, db: AsyncSession):
self._db = db
self._client_cache = {}
async def _get_client_config(self, model_id: str):
"""
根据模型 ID 获取 API 配置
Returns:
(base_url, api_key, model_name)
"""
from ai_platform.models import LLMModel, LLMProvider
result = await self._db.execute(
select(LLMModel).where(
LLMModel.id == model_id,
LLMModel.is_active == True,
LLMModel.is_deleted == False
)
)
model = result.scalar_one_or_none()
if not model:
raise ValueError(f'Rerank 模型不存在或已禁用: {model_id}')
if model.model_type != 'rerank':
raise ValueError(f'模型 {model.display_name} 不是 Rerank 类型')
provider_result = await self._db.execute(
select(LLMProvider).where(
LLMProvider.id == model.provider_id,
LLMProvider.is_active == True,
LLMProvider.is_deleted == False
)
)
provider = provider_result.scalar_one_or_none()
if not provider:
raise ValueError('Rerank 模型对应的提供商不存在或已禁用')
if provider.provider_type == 'ollama':
base_url = (provider.ollama_host or 'http://localhost:11434').rstrip('/') + '/v1'
else:
base_url = provider.api_base or 'https://api.openai.com/v1'
api_key = provider.api_key or 'ollama'
return base_url, api_key, model.model_name
async def rerank(
self,
model_id: str,
query: str,
documents: List[str],
top_n: Optional[int] = None,
) -> List[RerankResult]:
"""
对文档列表进行重排序
Args:
model_id: Rerank 模型 ID
query: 查询文本
documents: 待排序的文档列表
top_n: 返回前 N 个结果(默认返回全部)
Returns:
按相关性降序排列的 RerankResult 列表
"""
if not documents:
return []
if top_n is None:
top_n = len(documents)
base_url, api_key, model_name = await self._get_client_config(model_id)
try:
return await self._call_rerank_api(
base_url=base_url,
api_key=api_key,
model_name=model_name,
query=query,
documents=documents,
top_n=top_n,
)
except Exception as e:
logger.error(f'Rerank 调用失败: {e}')
raise ValueError(f'Rerank 调用失败: {str(e)}')
async def _call_rerank_api(
self,
base_url: str,
api_key: str,
model_name: str,
query: str,
documents: List[str],
top_n: int,
) -> List[RerankResult]:
"""
调用 Rerank APIJina/Cohere 兼容风格)
POST {base_url}/rerank
{
"model": "...",
"query": "...",
"documents": ["...", "..."],
"top_n": 5
}
Response:
{
"results": [
{"index": 0, "relevance_score": 0.95},
{"index": 2, "relevance_score": 0.87},
...
]
}
"""
import httpx
url = base_url.rstrip('/') + '/rerank'
headers = {
'Content-Type': 'application/json',
'Authorization': f'Bearer {api_key}',
}
payload = {
'model': model_name,
'query': query,
'documents': documents,
'top_n': top_n,
}
async with httpx.AsyncClient(timeout=60) as client:
response = await client.post(url, json=payload, headers=headers)
response.raise_for_status()
data = response.json()
# 解析结果(兼容 Jina/Cohere/通义千问 等格式)
raw_results = data.get('results', [])
results = []
for item in raw_results:
results.append(RerankResult(
index=item.get('index', 0),
relevance_score=item.get('relevance_score', 0.0),
))
# 按相关性降序排序
results.sort(key=lambda r: r.relevance_score, reverse=True)
return results