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", [])