From e067a4a60f58ff6bc68a47607ee95dc387aa41e8 Mon Sep 17 00:00:00 2001 From: Rapeter Date: Sun, 20 Sep 2026 08:41:43 +0800 Subject: [PATCH] fix(xinference): support string reranker documents --- .../model/reranker.py | 9 ++++++- apps/models_provider/tests.py | 26 +++++++++++++++++++ 2 files changed, 34 insertions(+), 1 deletion(-) diff --git a/apps/models_provider/impl/xinference_model_provider/model/reranker.py b/apps/models_provider/impl/xinference_model_provider/model/reranker.py index 405e8322474..86c21b7dfb8 100644 --- a/apps/models_provider/impl/xinference_model_provider/model/reranker.py +++ b/apps/models_provider/impl/xinference_model_provider/model/reranker.py @@ -27,6 +27,13 @@ def new_instance(model_type, model_name, model_credential: Dict[str, object], ** return XInferenceReranker(server_url=model_credential.get('server_url'), model_uid=model_name, api_key=model_credential.get('api_key'), top_n=model_kwargs.get('top_n', 3)) + @staticmethod + def _get_document_text(result: Dict[str, Any]) -> str: + document = result.get('document') + if isinstance(document, dict): + return document.get('text', '') + return document or '' + top_n: Optional[int] = 3 def compress_documents(self, documents: Sequence[Document], query: str, callbacks: Optional[Callbacks] = None) -> \ @@ -50,5 +57,5 @@ def compress_documents(self, documents: Sequence[Document], query: str, callback client = RESTfulClient(self.server_url, self.api_key) model: RESTfulRerankModelHandle = client.get_model(self.model_uid) res = model.rerank([document.page_content for document in documents], query, self.top_n, return_documents=True) - return [Document(page_content=d.get('document', {}).get('text'), + return [Document(page_content=self._get_document_text(d), metadata={'relevance_score': d.get('relevance_score')}) for d in res.get('results', [])] diff --git a/apps/models_provider/tests.py b/apps/models_provider/tests.py index 44e7258f159..8f4297b2c9d 100644 --- a/apps/models_provider/tests.py +++ b/apps/models_provider/tests.py @@ -2,8 +2,10 @@ from unittest.mock import patch from django.test import SimpleTestCase +from langchain_core.documents import Document from models_provider.impl.vllm_model_provider.model.whisper_sst import VllmWhisperSpeechToText +from models_provider.impl.xinference_model_provider.model.reranker import XInferenceReranker class VllmWhisperSpeechToTextTest(SimpleTestCase): @@ -24,3 +26,27 @@ def test_normalizes_trailing_slash_in_v1_base_url(self, openai_mock): base_url='https://vllm.example/v1', ) self.assertEqual(result, 'transcript') + + +class XInferenceRerankerTest(SimpleTestCase): + @patch('xinference_client.RESTfulClient') + def test_compress_documents_accepts_string_documents(self, client_mock): + client_mock.return_value.get_model.return_value.rerank.return_value = { + 'results': [{'document': 'current result', 'relevance_score': 0.9}] + } + model = XInferenceReranker(server_url='http://localhost', model_uid='reranker', api_key=None) + + result = model.compress_documents([Document(page_content='query text')], 'query') + + self.assertEqual(result[0].page_content, 'current result') + self.assertEqual(result[0].metadata['relevance_score'], 0.9) + + def test_extracts_text_from_legacy_document_object(self): + result = {'document': {'text': 'legacy result'}} + + self.assertEqual(XInferenceReranker._get_document_text(result), 'legacy result') + + def test_extracts_text_from_current_document_string(self): + result = {'document': 'current result'} + + self.assertEqual(XInferenceReranker._get_document_text(result), 'current result')