diff --git a/src/langbot/pkg/api/http/service/provider.py b/src/langbot/pkg/api/http/service/provider.py index 0a3a1cc75..74647d7ab 100644 --- a/src/langbot/pkg/api/http/service/provider.py +++ b/src/langbot/pkg/api/http/service/provider.py @@ -331,8 +331,15 @@ class ModelProviderService: embedding_models = await self.ap.embedding_models_service.get_embedding_models_by_provider( context, provider_uuid ) + rerank_service = getattr(self.ap, 'rerank_models_service', None) + rerank_models = ( + await rerank_service.get_rerank_models_by_provider(context, provider_uuid) + if rerank_service is not None + else [] + ) existing_llm_names = {model['name'] for model in llm_models} existing_embedding_names = {model['name'] for model in embedding_models} + existing_rerank_names = {model['name'] for model in rerank_models} filtered_models = [] for model in scanned_models: @@ -359,6 +366,8 @@ class ModelProviderService: 'already_added': ( model_name in existing_embedding_names if scanned_type == 'embedding' + else model_name in existing_rerank_names + if scanned_type == 'rerank' else model_name in existing_llm_names ), } diff --git a/src/langbot/pkg/pipeline/process/handlers/command.py b/src/langbot/pkg/pipeline/process/handlers/command.py index 09fa5379b..03a2ba9e6 100644 --- a/src/langbot/pkg/pipeline/process/handlers/command.py +++ b/src/langbot/pkg/pipeline/process/handlers/command.py @@ -1,10 +1,12 @@ from __future__ import annotations import typing +import sqlalchemy from .. import handler from ... import entities from ... import plugin_diagnostics +from ....entity.persistence.bot import BotAdmin import langbot_plugin.api.entities.builtin.provider.message as provider_message import langbot_plugin.api.entities.builtin.provider.session as provider_session import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query @@ -24,7 +26,14 @@ class CommandHandler(handler.MessageHandler): privilege = 1 - if f'{query.launcher_type.value}_{query.launcher_id}' in self.ap.instance_config.data['admins']: + admins = await self.ap.persistence_mgr.execute_async( + sqlalchemy.select(BotAdmin).where( + BotAdmin.bot_uuid == (query.bot_uuid or ''), + BotAdmin.launcher_type == query.launcher_type.value, + BotAdmin.launcher_id == str(query.launcher_id), + ) + ) + if admins.first() is not None: privilege = 2 spt = command_text.split(' ') diff --git a/src/langbot/pkg/platform/sources/aiocqhttp.py b/src/langbot/pkg/platform/sources/aiocqhttp.py index 6fc0ca918..9bee40d6e 100644 --- a/src/langbot/pkg/platform/sources/aiocqhttp.py +++ b/src/langbot/pkg/platform/sources/aiocqhttp.py @@ -241,7 +241,11 @@ class AiocqhttpMessageConverter(abstract_platform_adapter.AbstractMessageConvert async def process_message_data(msg_data, reply_list): if msg_data['type'] == 'image': image_base64, image_format = await image.qq_image_url_to_base64(msg_data['data']['url']) - reply_list.append(platform_message.Image(base64=f'data:image/{image_format};base64,{image_base64}')) + reply_list.append( + platform_message.Image( + url=msg_data['data']['url'], base64=f'data:image/{image_format};base64,{image_base64}' + ) + ) elif msg_data['type'] == 'text': reply_list.append(platform_message.Plain(text=msg_data['data']['text'])) @@ -286,7 +290,9 @@ class AiocqhttpMessageConverter(abstract_platform_adapter.AbstractMessageConvert image_msg = platform_message.Face(face_id=face_id, face_name=face_name) else: image_base64, image_format = await image.qq_image_url_to_base64(msg.data['url']) - image_msg = platform_message.Image(base64=f'data:image/{image_format};base64,{image_base64}') + image_msg = platform_message.Image( + url=msg.data['url'], base64=f'data:image/{image_format};base64,{image_base64}' + ) yiri_msg_list.append(image_msg) elif msg.type == 'forward': # 暂时不太合理 diff --git a/src/langbot/pkg/platform/sources/discord.py b/src/langbot/pkg/platform/sources/discord.py index 1be977d61..69f349056 100644 --- a/src/langbot/pkg/platform/sources/discord.py +++ b/src/langbot/pkg/platform/sources/discord.py @@ -764,7 +764,9 @@ class DiscordMessageConverter(abstract_platform_adapter.AbstractMessageConverter ) image_base64 = (await asyncio.to_thread(base64.b64encode, image_data)).decode('utf-8') image_format = response.headers['Content-Type'] - element_list.append(platform_message.Image(base64=f'data:{image_format};base64,{image_base64}')) + element_list.append( + platform_message.Image(url=attachment.url, base64=f'data:{image_format};base64,{image_base64}') + ) return platform_message.MessageChain(element_list) diff --git a/src/langbot/pkg/platform/sources/qqofficial.py b/src/langbot/pkg/platform/sources/qqofficial.py index bd4f97feb..598061f6c 100644 --- a/src/langbot/pkg/platform/sources/qqofficial.py +++ b/src/langbot/pkg/platform/sources/qqofficial.py @@ -102,7 +102,7 @@ class QQOfficialMessageConverter(abstract_platform_adapter.AbstractMessageConver yiri_msg_list.append(platform_message.Source(id=message_id, time=datetime.datetime.now())) if pic_url is not None: base64_url = await image.get_qq_official_image_base64(pic_url=pic_url, content_type=content_type) - yiri_msg_list.append(platform_message.Image(base64=base64_url)) + yiri_msg_list.append(platform_message.Image(url=pic_url, base64=base64_url)) yiri_msg_list.append(platform_message.Plain(text=message)) chain = platform_message.MessageChain(yiri_msg_list) diff --git a/src/langbot/pkg/platform/sources/slack.py b/src/langbot/pkg/platform/sources/slack.py index e577e3bdd..df18a0765 100644 --- a/src/langbot/pkg/platform/sources/slack.py +++ b/src/langbot/pkg/platform/sources/slack.py @@ -47,7 +47,7 @@ class SlackMessageConverter(abstract_platform_adapter.AbstractMessageConverter): yiri_msg_list.append(platform_message.Source(id=message_id, time=datetime.datetime.now())) if pic_url is not None: base64_url = await image.get_slack_image_to_base64(pic_url=pic_url, bot_token=bot.bot_token) - yiri_msg_list.append(platform_message.Image(base64=base64_url)) + yiri_msg_list.append(platform_message.Image(url=pic_url, base64=base64_url)) yiri_msg_list.append(platform_message.Plain(text=message)) chain = platform_message.MessageChain(yiri_msg_list) diff --git a/src/langbot/pkg/platform/sources/telegram.py b/src/langbot/pkg/platform/sources/telegram.py index 107894094..d24b84eb6 100644 --- a/src/langbot/pkg/platform/sources/telegram.py +++ b/src/langbot/pkg/platform/sources/telegram.py @@ -181,7 +181,10 @@ class TelegramMessageConverter(abstract_platform_adapter.AbstractMessageConverte encoded = await asyncio.to_thread(base64.b64encode, file_bytes) message_components.append( - platform_message.Image(base64=f'data:{file_format};base64,{encoded.decode("utf-8")}') + platform_message.Image( + url=file.file_path, + base64=f'data:{file_format};base64,{encoded.decode("utf-8")}', + ) ) if message.voice: @@ -889,4 +892,4 @@ class TelegramAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter): await self.logger.info('Telegram adapter stopped') self.msg_stream_id.clear() self._form_action_titles.clear() - return True + return True \ No newline at end of file diff --git a/src/langbot/pkg/platform/sources/wecom.py b/src/langbot/pkg/platform/sources/wecom.py index bb8ba2b17..93aaf1f92 100644 --- a/src/langbot/pkg/platform/sources/wecom.py +++ b/src/langbot/pkg/platform/sources/wecom.py @@ -133,7 +133,9 @@ class WecomMessageConverter(abstract_platform_adapter.AbstractMessageConverter): yiri_msg_list = [] yiri_msg_list.append(platform_message.Source(id=message_id, time=datetime.datetime.now())) image_base64, image_format = await image.get_wecom_image_base64(pic_url=picurl) - yiri_msg_list.append(platform_message.Image(base64=f'data:image/{image_format};base64,{image_base64}')) + yiri_msg_list.append( + platform_message.Image(url=picurl, base64=f'data:image/{image_format};base64,{image_base64}') + ) chain = platform_message.MessageChain(yiri_msg_list) return chain diff --git a/src/langbot/pkg/provider/modelmgr/modelmgr.py b/src/langbot/pkg/provider/modelmgr/modelmgr.py index 925b7de95..73eb65535 100644 --- a/src/langbot/pkg/provider/modelmgr/modelmgr.py +++ b/src/langbot/pkg/provider/modelmgr/modelmgr.py @@ -529,6 +529,7 @@ class ModelManager: model['uuid']: model for model in await self.ap.embedding_models_service.get_embedding_models(context, include_secret=True) } + existing_rerank_models = {m['uuid']: m for m in await self.ap.rerank_models_service.get_rerank_models()} created = 0 updated = 0 @@ -597,6 +598,33 @@ class ModelManager: ) 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: self.ap.logger.info(f'Synced models from LangBot Space: {created} added, {updated} updated.') diff --git a/src/langbot/pkg/provider/modelmgr/requesters/litellmchat.py b/src/langbot/pkg/provider/modelmgr/requesters/litellmchat.py index 982f5609b..f4daaf1df 100644 --- a/src/langbot/pkg/provider/modelmgr/requesters/litellmchat.py +++ b/src/langbot/pkg/provider/modelmgr/requesters/litellmchat.py @@ -944,16 +944,21 @@ class LiteLLMRequester(requester.ProviderAPIRequester): if 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] = { 'model': model_name, 'query': query, 'documents': documents, 'top_n': top_n, } - if extra_args: - payload.update(extra_args) + if request_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: async with httpx.AsyncClient( diff --git a/tests/unit_tests/api/service/test_provider_service.py b/tests/unit_tests/api/service/test_provider_service.py index 81f5ee6c1..5667ec9fe 100644 --- a/tests/unit_tests/api/service/test_provider_service.py +++ b/tests/unit_tests/api/service/test_provider_service.py @@ -833,6 +833,44 @@ class TestModelProviderServiceScanProviderModels: assert len(result['models']) == 1 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): """Raises ValueError when scan not implemented.""" # Setup diff --git a/tests/unit_tests/api/service/test_space_service.py b/tests/unit_tests/api/service/test_space_service.py index 6fb74d20b..d92265743 100644 --- a/tests/unit_tests/api/service/test_space_service.py +++ b/tests/unit_tests/api/service/test_space_service.py @@ -756,7 +756,7 @@ class TestSpaceServiceGetModels: 'uuid': 'uuid-2', 'model_id': 'model-2', 'provider': 'provider-2', - 'category': 'chat', + 'category': 'rerank', 'status': 'active', }, ] @@ -778,6 +778,7 @@ class TestSpaceServiceGetModels: # Verify assert len(result) == 2 + assert result[1].category == 'rerank' async def test_get_models_api_error(self): """Raises ValueError on API error.""" diff --git a/tests/unit_tests/pipeline/test_command_handler.py b/tests/unit_tests/pipeline/test_command_handler.py index 00bd5b681..bc5d39481 100644 --- a/tests/unit_tests/pipeline/test_command_handler.py +++ b/tests/unit_tests/pipeline/test_command_handler.py @@ -158,17 +158,21 @@ class TestCommandHandlerReal: @pytest.mark.asyncio async def test_admin_privilege_check(self, fake_app, mock_event_ctx, mock_execute_factory): - """Admin users get privilege level 2.""" + """A per-bot admin from the database is marked as admin in command events.""" from langbot_plugin.api.entities.builtin.provider.session import LauncherTypes command = get_command_handler() - fake_app.instance_config.data = {'admins': ['person_12345']} + admin_result = Mock() + admin_result.first.return_value = Mock() + fake_app.persistence_mgr.execute_async = AsyncMock(return_value=admin_result) + fake_app.instance_config.data = {} fake_app.plugin_connector.emit_event = AsyncMock(return_value=mock_event_ctx) fake_app.cmd_mgr.execute = mock_execute_factory() handler = command.CommandHandler(fake_app) query = command_query('status') + query.bot_uuid = 'bot-1' query.launcher_type = LauncherTypes.PERSON query.launcher_id = 12345 @@ -176,23 +180,28 @@ class TestCommandHandlerReal: async for result in handler.handle(query): results.append(result) + fake_app.persistence_mgr.execute_async.assert_awaited_once() call_args = fake_app.plugin_connector.emit_event.call_args event = call_args[0][0] assert event.is_admin is True @pytest.mark.asyncio async def test_non_admin_privilege_check(self, fake_app, mock_event_ctx, mock_execute_factory): - """Non-admin users get privilege level 1.""" + """A launcher absent from the per-bot admin table is not an admin.""" from langbot_plugin.api.entities.builtin.provider.session import LauncherTypes command = get_command_handler() - fake_app.instance_config.data = {'admins': ['person_12345']} + admin_result = Mock() + admin_result.first.return_value = None + fake_app.persistence_mgr.execute_async = AsyncMock(return_value=admin_result) + fake_app.instance_config.data = {} fake_app.plugin_connector.emit_event = AsyncMock(return_value=mock_event_ctx) fake_app.cmd_mgr.execute = mock_execute_factory() handler = command.CommandHandler(fake_app) query = command_query('status') + query.bot_uuid = 'bot-1' query.launcher_type = LauncherTypes.PERSON query.launcher_id = 67890 @@ -200,6 +209,7 @@ class TestCommandHandlerReal: async for result in handler.handle(query): results.append(result) + fake_app.persistence_mgr.execute_async.assert_awaited_once() call_args = fake_app.plugin_connector.emit_event.call_args event = call_args[0][0] assert event.is_admin is False diff --git a/tests/unit_tests/provider/test_litellmchat.py b/tests/unit_tests/provider/test_litellmchat.py index a22a589a7..d0332f726 100644 --- a/tests/unit_tests/provider/test_litellmchat.py +++ b/tests/unit_tests/provider/test_litellmchat.py @@ -1147,6 +1147,46 @@ class TestInvokeRerank: assert results[0]['relevance_score'] == 1.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: """Test _convert_messages method""" diff --git a/tests/unit_tests/provider/test_model_manager.py b/tests/unit_tests/provider/test_model_manager.py index de343140b..6f288adee 100644 --- a/tests/unit_tests/provider/test_model_manager.py +++ b/tests/unit_tests/provider/test_model_manager.py @@ -86,6 +86,48 @@ async def test_model_manager_skips_legacy_space_sync_in_cloud_runtime(mock_app_f app.workspace_service.get_local_execution_binding.assert_not_awaited() +@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 # ============================================================================ @@ -1067,4 +1109,4 @@ def test_provider_not_found_error_str(): error = provider_errors.ProviderNotFoundError('test-provider') assert str(error) == 'Provider test-provider not found' - assert error.provider_name == 'test-provider' + assert error.provider_name == 'test-provider' \ No newline at end of file