feat: support Space rerank models (#2358)

Co-authored-by: chan <dadachann@users.noreply.github.com>
This commit is contained in:
Hyu
2026-07-26 16:37:34 +08:00
committed by GitHub
parent 38e35d328a
commit dd8d1007a1
7 changed files with 171 additions and 5 deletions
@@ -232,8 +232,15 @@ class ModelProviderService:
llm_models = await self.ap.llm_model_service.get_llm_models_by_provider(provider_uuid) llm_models = await self.ap.llm_model_service.get_llm_models_by_provider(provider_uuid)
embedding_models = await self.ap.embedding_models_service.get_embedding_models_by_provider(provider_uuid) embedding_models = await self.ap.embedding_models_service.get_embedding_models_by_provider(provider_uuid)
rerank_service = getattr(self.ap, 'rerank_models_service', None)
rerank_models = (
await rerank_service.get_rerank_models_by_provider(provider_uuid)
if rerank_service is not None
else []
)
existing_llm_names = {model['name'] for model in llm_models} existing_llm_names = {model['name'] for model in llm_models}
existing_embedding_names = {model['name'] for model in embedding_models} existing_embedding_names = {model['name'] for model in embedding_models}
existing_rerank_names = {model['name'] for model in rerank_models}
filtered_models = [] filtered_models = []
for model in scanned_models: for model in scanned_models:
@@ -260,6 +267,8 @@ class ModelProviderService:
'already_added': ( 'already_added': (
model_name in existing_embedding_names model_name in existing_embedding_names
if scanned_type == 'embedding' if scanned_type == 'embedding'
else model_name in existing_rerank_names
if scanned_type == 'rerank'
else model_name in existing_llm_names else model_name in existing_llm_names
), ),
} }
@@ -193,6 +193,9 @@ class ModelManager:
existing_embedding_models = { existing_embedding_models = {
m['uuid']: m for m in await self.ap.embedding_models_service.get_embedding_models() m['uuid']: m for m in await self.ap.embedding_models_service.get_embedding_models()
} }
existing_rerank_models = {
m['uuid']: m for m in await self.ap.rerank_models_service.get_rerank_models()
}
created = 0 created = 0
updated = 0 updated = 0
@@ -258,6 +261,33 @@ class ModelManager:
await self.ap.embedding_models_service.update_embedding_model(space_model.uuid, dict(desired)) await self.ap.embedding_models_service.update_embedding_model(space_model.uuid, dict(desired))
updated += 1 updated += 1
elif space_model.category == 'rerank':
existing = existing_rerank_models.get(space_model.uuid)
if existing is None:
await self.ap.rerank_models_service.create_rerank_model(
{
'uuid': space_model.uuid,
'name': space_model.model_id,
'provider_uuid': space_model_provider.uuid,
'extra_args': {},
'prefered_ranking': space_model.featured_order,
},
preserve_uuid=True,
)
created += 1
elif existing.get('provider_uuid') == space_model_provider.uuid:
desired = {
'name': space_model.model_id,
'provider_uuid': space_model_provider.uuid,
'prefered_ranking': space_model.featured_order,
}
if (
existing.get('name') != desired['name']
or existing.get('prefered_ranking') != desired['prefered_ranking']
):
await self.ap.rerank_models_service.update_rerank_model(space_model.uuid, dict(desired))
updated += 1
if created or updated: if created or updated:
self.ap.logger.info(f'Synced models from LangBot Space: {created} added, {updated} updated.') self.ap.logger.info(f'Synced models from LangBot Space: {created} added, {updated} updated.')
@@ -943,16 +943,21 @@ class LiteLLMRequester(requester.ProviderAPIRequester):
if api_key: if api_key:
headers['Authorization'] = f'Bearer {api_key}' headers['Authorization'] = f'Bearer {api_key}'
request_args = dict(extra_args)
rerank_url = request_args.pop('rerank_url', None)
rerank_path = request_args.pop('rerank_path', 'rerank')
payload: dict[str, typing.Any] = { payload: dict[str, typing.Any] = {
'model': model_name, 'model': model_name,
'query': query, 'query': query,
'documents': documents, 'documents': documents,
'top_n': top_n, 'top_n': top_n,
} }
if extra_args: if request_args:
payload.update(extra_args) payload.update(request_args)
rerank_url = f'{base_url}/rerank' if not rerank_url:
rerank_url = f'{base_url}/{str(rerank_path).strip("/")}'
try: try:
async with httpx.AsyncClient(timeout=timeout) as client: async with httpx.AsyncClient(timeout=timeout) as client:
@@ -777,6 +777,44 @@ class TestModelProviderServiceScanProviderModels:
assert len(result['models']) == 1 assert len(result['models']) == 1
assert result['models'][0]['type'] == 'llm' assert result['models'][0]['type'] == 'llm'
async def test_scan_provider_marks_existing_rerank_model(self):
"""Rerank scan results use the rerank service when computing already_added."""
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.llm_model_service = SimpleNamespace()
ap.embedding_models_service = SimpleNamespace()
ap.rerank_models_service = SimpleNamespace()
provider = _create_mock_provider(provider_uuid='rerank-scan-uuid')
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_mock_result([], first_item=provider))
ap.persistence_mgr.serialize_model = Mock(
return_value={
'uuid': 'rerank-scan-uuid',
'name': 'New API',
'requester': 'new-api-chat-completions',
'base_url': 'https://new-api.example.com/v1',
'api_keys': ['key'],
}
)
runtime_provider = Mock()
runtime_provider.token_mgr.get_token.return_value = 'token'
runtime_provider.requester.scan_models = AsyncMock(
return_value={'models': [{'id': 'Qwen3-Reranker-8B', 'type': 'rerank'}]}
)
ap.model_mgr.load_provider = AsyncMock(return_value=runtime_provider)
ap.llm_model_service.get_llm_models_by_provider = AsyncMock(return_value=[])
ap.embedding_models_service.get_embedding_models_by_provider = AsyncMock(return_value=[])
ap.rerank_models_service.get_rerank_models_by_provider = AsyncMock(
return_value=[{'name': 'Qwen3-Reranker-8B'}]
)
result = await ModelProviderService(ap).scan_provider_models('rerank-scan-uuid', model_type='rerank')
assert result['models'][0]['type'] == 'rerank'
assert result['models'][0]['already_added'] is True
async def test_scan_provider_not_implemented_raises_error(self): async def test_scan_provider_not_implemented_raises_error(self):
"""Raises ValueError when scan not implemented.""" """Raises ValueError when scan not implemented."""
# Setup # Setup
@@ -693,7 +693,7 @@ class TestSpaceServiceGetModels:
'uuid': 'uuid-2', 'uuid': 'uuid-2',
'model_id': 'model-2', 'model_id': 'model-2',
'provider': 'provider-2', 'provider': 'provider-2',
'category': 'chat', 'category': 'rerank',
'status': 'active', 'status': 'active',
}, },
] ]
@@ -714,6 +714,7 @@ class TestSpaceServiceGetModels:
# Verify # Verify
assert len(result) == 2 assert len(result) == 2
assert result[1].category == 'rerank'
async def test_get_models_api_error(self): async def test_get_models_api_error(self):
"""Raises ValueError on API error.""" """Raises ValueError on API error."""
@@ -1147,6 +1147,46 @@ class TestInvokeRerank:
assert results[0]['relevance_score'] == 1.0 assert results[0]['relevance_score'] == 1.0
assert results[1]['relevance_score'] == 0.0 assert results[1]['relevance_score'] == 0.0
@pytest.mark.asyncio
@pytest.mark.parametrize(
('model_extra_args', 'expected_url'),
[
({'rerank_path': 'reranks'}, 'https://gateway.example.com/v1/reranks'),
({'rerank_url': 'https://rerank.example.com/api/rerank'}, 'https://rerank.example.com/api/rerank'),
],
)
async def test_invoke_rerank_openai_compatible_endpoint_override(self, model_extra_args, expected_url):
"""Endpoint configuration controls routing and is not sent in the Cohere body."""
requester = litellmchat.LiteLLMRequester(
ap=Mock(),
config={
'base_url': 'https://gateway.example.com/v1/',
'custom_llm_provider': 'openai',
},
)
model = MockRuntimeRerankModel('Qwen3-Reranker-8B', 'test-api-key')
model.model_entity.extra_args = model_extra_args
mock_resp = Mock()
mock_resp.raise_for_status = Mock()
mock_resp.json = Mock(return_value={'results': [{'index': 0, 'relevance_score': 0.8}]})
mock_client = AsyncMock()
mock_client.post = AsyncMock(return_value=mock_resp)
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
mock_client.__aexit__ = AsyncMock(return_value=False)
with patch('httpx.AsyncClient', return_value=mock_client):
await requester.invoke_rerank(model=model, query='query', documents=['document'])
assert mock_client.post.call_args.args[0] == expected_url
payload = mock_client.post.call_args.kwargs['json']
assert payload == {
'model': 'Qwen3-Reranker-8B',
'query': 'query',
'documents': ['document'],
'top_n': 1,
}
class TestConvertMessages: class TestConvertMessages:
"""Test _convert_messages method""" """Test _convert_messages method"""
@@ -8,7 +8,8 @@ and error handling without calling real LLM APIs.
from __future__ import annotations from __future__ import annotations
import pytest import pytest
from unittest.mock import Mock from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
from langbot.pkg.provider.modelmgr.modelmgr import ModelManager from langbot.pkg.provider.modelmgr.modelmgr import ModelManager
from langbot.pkg.provider.modelmgr import requester from langbot.pkg.provider.modelmgr import requester
@@ -62,6 +63,48 @@ async def test_model_manager_skips_space_sync_when_disabled(mock_app_for_modelmg
app.space_service.get_models.assert_not_called() app.space_service.get_models.assert_not_called()
@pytest.mark.asyncio
async def test_sync_new_models_from_space_creates_rerank_models(mock_app_for_modelmgr):
"""Space rerank entries are discovered and persisted under the shared provider."""
app = mock_app_for_modelmgr
provider = persistence_model.ModelProvider(
uuid='space-provider',
name='LangBot Space',
requester='space-chat-completions',
base_url='https://api.langbot.cloud/v1',
api_keys=['space-key'],
)
app.persistence_mgr.execute_async = AsyncMock(return_value=_make_mock_result([provider], first_item=provider))
app.space_service.get_models = AsyncMock(
return_value=[
SimpleNamespace(
uuid='rerank-model-uuid',
model_id='Qwen3-Reranker-8B',
category='rerank',
featured_order=10,
)
]
)
app.llm_model_service.get_llm_models = AsyncMock(return_value=[])
app.embedding_models_service.get_embedding_models = AsyncMock(return_value=[])
app.rerank_models_service = AsyncMock()
app.rerank_models_service.get_rerank_models = AsyncMock(return_value=[])
model_mgr = ModelManager(app)
await model_mgr.sync_new_models_from_space()
app.rerank_models_service.create_rerank_model.assert_awaited_once_with(
{
'uuid': 'rerank-model-uuid',
'name': 'Qwen3-Reranker-8B',
'provider_uuid': 'space-provider',
'extra_args': {},
'prefered_ranking': 10,
},
preserve_uuid=True,
)
# ============================================================================ # ============================================================================
# Model Loading Tests # Model Loading Tests
# ============================================================================ # ============================================================================