Skip to content
Open
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
4 changes: 4 additions & 0 deletions astrbot/core/config/default.py
Original file line number Diff line number Diff line change
Expand Up @@ -1896,6 +1896,7 @@
"rerank_api_suffix": "/v1/rerank",
"rerank_model": "BAAI/bge-reranker-base",
"timeout": 20,
"proxy": "",
},
"Xinference Rerank": {
"id": "xinference_rerank",
Expand All @@ -1921,6 +1922,7 @@
"timeout": 30,
"return_documents": False,
"instruct": "",
"proxy": "",
},
"NVIDIA Rerank": {
"id": "nvidia_rerank",
Expand All @@ -1934,6 +1936,7 @@
"nvidia_rerank_model_endpoint": "/reranking",
"timeout": 20,
"nvidia_rerank_truncate": "",
"proxy": "",
},
"TEI Rerank": {
"id": "tei_rerank",
Expand All @@ -1948,6 +1951,7 @@
"tei_rerank_truncation_direction": "Right",
"tei_rerank_raw_scores": False,
"tei_rerank_return_text": False,
"proxy": "",
},
"Xinference STT": {
"id": "xinference_stt",
Expand Down
5 changes: 4 additions & 1 deletion astrbot/core/provider/sources/bailian_rerank_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,7 @@ def __init__(self, provider_config: dict, provider_settings: dict) -> None:
self.timeout = provider_config.get("timeout", 30)
self.return_documents = provider_config.get("return_documents", False)
self.instruct = provider_config.get("instruct", "")
self.proxy = provider_config.get("proxy", "") or None

self.base_url = provider_config.get(
"rerank_api_base",
Expand Down Expand Up @@ -232,7 +233,9 @@ async def rerank(
)

# 发送请求
async with self.client.post(self.base_url, json=payload) as response:
async with self.client.post(
self.base_url, json=payload, proxy=self.proxy
) as response:
response.raise_for_status()
response_data = await response.json()

Expand Down
5 changes: 4 additions & 1 deletion astrbot/core/provider/sources/nvidia_rerank_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@ def __init__(self, provider_config: dict, provider_settings: dict) -> None:
"nvidia_rerank_model_endpoint", "/reranking"
)
self.truncate = provider_config.get("nvidia_rerank_truncate", "")
self.proxy = provider_config.get("proxy", "") or None

self.client = None
self.set_model(self.model)
Expand Down Expand Up @@ -130,7 +131,9 @@ async def rerank(
payload = self._build_payload(query, documents)
request_url = self._get_endpoint()

async with client.post(request_url, json=payload) as response:
async with client.post(
request_url, json=payload, proxy=self.proxy
) as response:
if response.status != 200:
try:
response_data = await response.json()
Expand Down
7 changes: 5 additions & 2 deletions astrbot/core/provider/sources/tei_rerank_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@ def __init__(self, provider_config: dict, provider_settings: dict) -> None:
).lower()
self.raw_scores = provider_config.get("tei_rerank_raw_scores", False)
self.return_text = provider_config.get("tei_rerank_return_text", False)
self.proxy = provider_config.get("proxy", "") or None

h = {}
if self.api_key:
Expand Down Expand Up @@ -75,7 +76,9 @@ async def rerank(
f"[TEI Rerank] Request: query='{query[:50]}...', "
f"doc_count={len(documents)}"
)
async with self.client.post(rerank_url, json=payload) as response:
async with self.client.post(
rerank_url, json=payload, proxy=self.proxy
) as response:
if response.status != 200:
try:
error_data = await response.json()
Expand Down Expand Up @@ -128,7 +131,7 @@ async def test(self) -> None:

health_url = f"{self.base_url}/health"
try:
async with self.client.get(health_url) as response:
async with self.client.get(health_url, proxy=self.proxy) as response:
if response.status != 200:
raise Exception(
f"TEI service health check failed at {self.base_url}: "
Expand Down
2 changes: 2 additions & 0 deletions astrbot/core/provider/sources/vllm_rerank_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@ def __init__(self, provider_config: dict, provider_settings: dict) -> None:
self.api_suffix = "/" + self.api_suffix
self.timeout = provider_config.get("timeout", 20)
self.model = provider_config.get("rerank_model", "BAAI/bge-reranker-base")
self.proxy = provider_config.get("proxy", "") or None

h = {}
if self.auth_key:
Expand Down Expand Up @@ -54,6 +55,7 @@ async def rerank(
async with self.client.post(
rerank_url,
json=payload,
proxy=self.proxy,
) as response:
response_data = await response.json()
results = response_data.get("results", [])
Expand Down