From ceeeef6b048cc19950b8238337b4a57a0fb2751a Mon Sep 17 00:00:00 2001 From: GuJi08233 <3439651960@qq.com> Date: Sun, 26 Jul 2026 01:17:15 +0800 Subject: [PATCH] fix(provider): add proxy support for rerank providers Rerank providers use aiohttp.ClientSession without proxy configuration, causing requests to bypass the configured proxy. Add proxy field to vLLM, Bailian, NVIDIA, and TEI rerank provider configs and pass it to each aiohttp request call. Fixes #9383 --- astrbot/core/config/default.py | 4 ++++ astrbot/core/provider/sources/bailian_rerank_source.py | 5 ++++- astrbot/core/provider/sources/nvidia_rerank_source.py | 5 ++++- astrbot/core/provider/sources/tei_rerank_source.py | 7 +++++-- astrbot/core/provider/sources/vllm_rerank_source.py | 2 ++ 5 files changed, 19 insertions(+), 4 deletions(-) diff --git a/astrbot/core/config/default.py b/astrbot/core/config/default.py index 52e7036320..c2f7f4dd27 100644 --- a/astrbot/core/config/default.py +++ b/astrbot/core/config/default.py @@ -1896,6 +1896,7 @@ "rerank_api_suffix": "/v1/rerank", "rerank_model": "BAAI/bge-reranker-base", "timeout": 20, + "proxy": "", }, "Xinference Rerank": { "id": "xinference_rerank", @@ -1921,6 +1922,7 @@ "timeout": 30, "return_documents": False, "instruct": "", + "proxy": "", }, "NVIDIA Rerank": { "id": "nvidia_rerank", @@ -1934,6 +1936,7 @@ "nvidia_rerank_model_endpoint": "/reranking", "timeout": 20, "nvidia_rerank_truncate": "", + "proxy": "", }, "TEI Rerank": { "id": "tei_rerank", @@ -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", diff --git a/astrbot/core/provider/sources/bailian_rerank_source.py b/astrbot/core/provider/sources/bailian_rerank_source.py index 65356e100b..a38c7bd903 100644 --- a/astrbot/core/provider/sources/bailian_rerank_source.py +++ b/astrbot/core/provider/sources/bailian_rerank_source.py @@ -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", @@ -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() diff --git a/astrbot/core/provider/sources/nvidia_rerank_source.py b/astrbot/core/provider/sources/nvidia_rerank_source.py index c168da4a6e..0ac7fdbd99 100644 --- a/astrbot/core/provider/sources/nvidia_rerank_source.py +++ b/astrbot/core/provider/sources/nvidia_rerank_source.py @@ -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) @@ -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() diff --git a/astrbot/core/provider/sources/tei_rerank_source.py b/astrbot/core/provider/sources/tei_rerank_source.py index 9a5d58aaf7..4ab337337b 100644 --- a/astrbot/core/provider/sources/tei_rerank_source.py +++ b/astrbot/core/provider/sources/tei_rerank_source.py @@ -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: @@ -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() @@ -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}: " diff --git a/astrbot/core/provider/sources/vllm_rerank_source.py b/astrbot/core/provider/sources/vllm_rerank_source.py index e5ed791160..0773d98d0a 100644 --- a/astrbot/core/provider/sources/vllm_rerank_source.py +++ b/astrbot/core/provider/sources/vllm_rerank_source.py @@ -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: @@ -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", [])