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
55 changes: 29 additions & 26 deletions astrbot/core/provider/sources/bailian_rerank_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,12 @@ def __init__(self, provider_config: dict, provider_settings: dict) -> None:
logger.info(f"AstrBot 百炼 Rerank 初始化完成。模型: {self.model}")

def _uses_compatible_api(self) -> bool:
"""判断当前 base_url 是否为 compatible-api/compatible-mode 端点。

通过解析 URL path 而非简单的字符串包含判断,
避免 query string 中出现 "compatible-api" 导致误判,
同时支持中国站(compatible-api)和新加坡站(compatible-mode)两种路径。
"""
base_url_path = urlsplit(self.base_url).path.rstrip("/")
return base_url_path.endswith(self.COMPATIBLE_API_PATH_SUFFIXES)

Expand All @@ -97,40 +103,39 @@ def _build_payload(
"""
normalized_model = self.model.strip().lower()
normalized_top_n = top_n if top_n is not None and top_n > 0 else None
is_qwen3_rerank = normalized_model == self.QWEN3_RERANK_MODEL
is_compatible_api = self._uses_compatible_api()

if normalized_model == self.QWEN3_RERANK_MODEL and is_compatible_api:
payload = {
if is_qwen3_rerank and self.return_documents:
logger.warning(
"qwen3-rerank does not support return_documents; "
"this option will be ignored."
)

if is_compatible_api:
payload: dict[str, Any] = {
"model": self.model,
"query": query,
"documents": documents,
}
if normalized_top_n is not None:
payload["top_n"] = normalized_top_n
if self.instruct:
# instruct 仅 qwen3-rerank 支持
if is_qwen3_rerank and self.instruct:
payload["instruct"] = self.instruct
if self.return_documents:
logger.warning(
"qwen3-rerank does not support return_documents; "
"this option will be ignored."
)
if self.return_documents and not is_qwen3_rerank:
payload["return_documents"] = True
return payload

payload_input = {"query": query, "documents": documents}
params = {
k: v
for k, v in [
("top_n", normalized_top_n),
("return_documents", True if self.return_documents else None),
(
"instruct",
self.instruct
if self.instruct and normalized_model == self.QWEN3_RERANK_MODEL
else None,
),
]
if v is not None
}
# 原生端点:input 包装格式
payload_input: dict[str, Any] = {"query": query, "documents": documents}
params: dict[str, Any] = {}
if normalized_top_n is not None:
params["top_n"] = normalized_top_n
if self.return_documents and not is_qwen3_rerank:
params["return_documents"] = True
if is_qwen3_rerank and self.instruct:
params["instruct"] = self.instruct

base: dict[str, Any] = {"model": self.model, "input": payload_input}
if params:
Expand All @@ -151,9 +156,7 @@ def _parse_results(self, data: dict) -> list[RerankResult]:
BailianAPIError: API返回错误
KeyError: 结果缺少必要字段
"""
is_compatible_api = self._uses_compatible_api()

if is_compatible_api:
if self._uses_compatible_api():
code = data.get("code")
if code:
raise BailianAPIError(
Expand Down
127 changes: 124 additions & 3 deletions tests/test_bailian_rerank_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,7 @@

import astrbot.core.provider.sources.bailian_rerank_source as bailian_rerank_module
from astrbot.core.config.default import CONFIG_METADATA_2
from astrbot.core.provider.sources.bailian_rerank_source import (
BailianRerankProvider,
)
from astrbot.core.provider.sources.bailian_rerank_source import BailianRerankProvider

CHINA_COMPATIBLE_URL = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks"
SINGAPORE_COMPATIBLE_URL = (
Expand Down Expand Up @@ -142,3 +140,126 @@ def test_native_endpoint_parses_nested_results(provider):
assert len(results) == 1
assert results[0].index == 0
assert results[0].relevance_score == 0.75


# ---------- qwen3-rerank additional cases ----------


def test_qwen3_native_endpoint_ignores_return_documents(provider):
provider.base_url = NATIVE_URL
provider.return_documents = True

payload = provider._build_payload("q", ["d1"], None)

assert "parameters" not in payload


def test_qwen3_native_endpoint_no_params_when_top_n_zero(provider):
provider.base_url = NATIVE_URL

payload = provider._build_payload("q", ["d1"], 0)

assert "parameters" not in payload


def test_qwen3_compatible_endpoint_ignores_return_documents(provider):
provider.base_url = CHINA_COMPATIBLE_URL
provider.return_documents = True

payload = provider._build_payload("q", ["d1"], None)

assert "return_documents" not in payload


def test_qwen3_compatible_endpoint_no_optional_fields(provider):
provider.base_url = CHINA_COMPATIBLE_URL

payload = provider._build_payload("q", ["d1"], None)

assert payload == {
"model": "qwen3-rerank",
"query": "q",
"documents": ["d1"],
}


def test_qwen3_compatible_endpoint_includes_instruct(provider):
provider.base_url = CHINA_COMPATIBLE_URL
provider.instruct = "focus on facts"

payload = provider._build_payload("q", ["d1", "d2"], 3)

assert payload == {
"model": "qwen3-rerank",
"query": "q",
"documents": ["d1", "d2"],
"top_n": 3,
"instruct": "focus on facts",
}


# ---------- non-qwen3 model (gte-rerank) ----------


def test_gte_rerank_native_endpoint_uses_wrapped_format_with_return_documents(provider):
provider.model = "gte-rerank-v2"
provider.base_url = NATIVE_URL
provider.return_documents = True

payload = provider._build_payload("q", ["d1"], 2)

assert payload == {
"model": "gte-rerank-v2",
"input": {"query": "q", "documents": ["d1"]},
"parameters": {"top_n": 2, "return_documents": True},
}


def test_gte_rerank_native_endpoint_no_instruct(provider):
provider.model = "gte-rerank-v2"
provider.base_url = NATIVE_URL
provider.instruct = "ignored"

payload = provider._build_payload("q", ["d1"], None)

assert payload["input"] == {"query": "q", "documents": ["d1"]}
assert "parameters" not in payload


def test_gte_rerank_compatible_endpoint_uses_flat_format(provider):
provider.model = "gte-rerank-v2"
provider.base_url = CHINA_COMPATIBLE_URL
provider.return_documents = True

payload = provider._build_payload("q", ["d1"], 2)

assert payload == {
"model": "gte-rerank-v2",
"query": "q",
"documents": ["d1"],
"top_n": 2,
"return_documents": True,
}


def test_gte_rerank_compatible_endpoint_no_instruct(provider):
provider.model = "gte-rerank-v2"
provider.base_url = CHINA_COMPATIBLE_URL
provider.instruct = "ignored"

payload = provider._build_payload("q", ["d1"], None)

assert "instruct" not in payload


def test_gte_rerank_compatible_endpoint_no_optional_fields(provider):
provider.model = "gte-rerank-v2"
provider.base_url = SINGAPORE_COMPATIBLE_URL

payload = provider._build_payload("q", ["d1"], None)

assert payload == {
"model": "gte-rerank-v2",
"query": "q",
"documents": ["d1"],
}