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
46 changes: 26 additions & 20 deletions astrbot/core/provider/sources/bailian_rerank_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,9 @@ def __init__(self, provider_config: dict, provider_settings: dict) -> None:
"rerank_api_base",
"https://dashscope.aliyuncs.com/api/v1/services/rerank/text-rerank/text-rerank",
)
# 请求/响应格式由端点决定:compatible-api 用扁平/OpenAI 格式,
# 原生端点用 input-wrapped 格式
self.is_compatible_api = "compatible-api" in self.base_url

# 设置HTTP客户端
headers = {
Expand Down Expand Up @@ -88,33 +91,38 @@ 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

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

if self.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),
]
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 @@ -135,9 +143,7 @@ def _parse_results(self, data: dict) -> list[RerankResult]:
BailianAPIError: API返回错误
KeyError: 结果缺少必要字段
"""
is_compatible_api = "compatible-api" in self.base_url

if is_compatible_api:
if self.is_compatible_api:
code = data.get("code")
if code:
raise BailianAPIError(
Expand Down
191 changes: 191 additions & 0 deletions tests/test_bailian_rerank_source.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,191 @@
from unittest.mock import MagicMock, patch

from astrbot.core.provider.sources.bailian_rerank_source import BailianRerankProvider

NATIVE_URL = "https://dashscope.aliyuncs.com/api/v1/services/rerank/text-rerank/text-rerank"
COMPATIBLE_URL = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks"


def _make_provider(
model: str = "qwen3-rerank",
base_url: str = NATIVE_URL,
instruct: str = "",
return_documents: bool = False,
) -> BailianRerankProvider:
config = {
"rerank_api_key": "test-key",
"rerank_model": model,
"rerank_api_base": base_url,
"instruct": instruct,
"return_documents": return_documents,
"timeout": 30,
}
with patch(
"astrbot.core.provider.sources.bailian_rerank_source.aiohttp.ClientSession"
):
return BailianRerankProvider(config, {})


# ---------- qwen3-rerank, native endpoint ----------


def test_qwen3_rerank_native_endpoint_uses_input_wrapped_format():
provider = _make_provider("qwen3-rerank", NATIVE_URL)

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

assert payload == {
"model": "qwen3-rerank",
"input": {"query": "q", "documents": ["d1", "d2"]},
"parameters": {"top_n": 5},
}


def test_qwen3_rerank_native_endpoint_puts_instruct_in_parameters():
provider = _make_provider("qwen3-rerank", NATIVE_URL, instruct="rank by relevance")

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

assert payload["input"] == {"query": "q", "documents": ["d1"]}
assert payload["parameters"] == {"instruct": "rank by relevance"}


def test_qwen3_rerank_native_endpoint_ignores_return_documents():
provider = _make_provider("qwen3-rerank", NATIVE_URL, return_documents=True)

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

assert "parameters" not in payload


def test_qwen3_rerank_native_endpoint_no_params_when_top_n_zero():
provider = _make_provider("qwen3-rerank", NATIVE_URL)

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

assert "parameters" not in payload


# ---------- qwen3-rerank, compatible-api endpoint ----------


def test_qwen3_rerank_compatible_endpoint_uses_flat_format():
provider = _make_provider(
"qwen3-rerank", COMPATIBLE_URL, 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",
}


def test_qwen3_rerank_compatible_endpoint_ignores_return_documents():
provider = _make_provider("qwen3-rerank", COMPATIBLE_URL, return_documents=True)

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

assert "return_documents" not in payload


def test_qwen3_rerank_compatible_endpoint_no_optional_fields():
provider = _make_provider("qwen3-rerank", COMPATIBLE_URL)

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

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


# ---------- non-qwen3 model (gte-rerank), native endpoint ----------


def test_gte_rerank_native_endpoint_keeps_input_wrapped_format():
provider = _make_provider("gte-rerank-v2", NATIVE_URL, 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 = _make_provider("gte-rerank-v2", NATIVE_URL, 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_native_endpoint_default_url():
"""is_compatible_api defaults to False for the default native URL."""
config = {"rerank_api_key": "test-key"}
with patch(
"astrbot.core.provider.sources.bailian_rerank_source.aiohttp.ClientSession"
):
provider = BailianRerankProvider(config, {})

assert provider.is_compatible_api is False


# ---------- non-qwen3 model (gte-rerank), compatible-api endpoint ----------


def test_gte_rerank_compatible_endpoint_uses_flat_format():
provider = _make_provider("gte-rerank-v2", COMPATIBLE_URL, 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 = _make_provider("gte-rerank-v2", COMPATIBLE_URL, instruct="ignored")

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

assert "instruct" not in payload


def test_gte_rerank_compatible_endpoint_no_optional_fields():
provider = _make_provider("gte-rerank-v2", COMPATIBLE_URL)

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

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


def test_compatible_api_flag_set_from_base_url():
"""is_compatible_api is computed from base_url during __init__."""
config = {
"rerank_api_key": "test-key",
"rerank_api_base": COMPATIBLE_URL,
}
with patch(
"astrbot.core.provider.sources.bailian_rerank_source.aiohttp.ClientSession"
):
provider = BailianRerankProvider(config, {})

assert provider.is_compatible_api is True