mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
refactor: resolve summary models through manager (#40696)
This commit is contained in:
+44
-43
@@ -16,7 +16,7 @@ from core.rag.models.document import AttachmentDocument, Document
|
||||
from extensions.storage.storage_type import StorageType
|
||||
from graphon.model_runtime.entities.llm_entities import LLMResult, LLMUsage
|
||||
from graphon.model_runtime.entities.message_entities import AssistantPromptMessage, ImagePromptMessageContent
|
||||
from graphon.model_runtime.entities.model_entities import ModelFeature
|
||||
from graphon.model_runtime.entities.model_entities import ModelFeature, ModelType
|
||||
from models.dataset import DocumentSegment, SegmentAttachmentBinding
|
||||
from models.enums import CreatorUserRole
|
||||
from models.model import UploadFile
|
||||
@@ -519,41 +519,51 @@ class TestParagraphIndexProcessor:
|
||||
with pytest.raises(ValueError, match="model_name and model_provider_name"):
|
||||
ParagraphIndexProcessor.generate_summary("tenant-1", "text", {"enable": True}, session=self.session)
|
||||
|
||||
def test_generate_summary_text_only_flow(self, caplog: pytest.LogCaptureFixture) -> None:
|
||||
def test_generate_summary_text_only_flow(self) -> None:
|
||||
model_instance = Mock()
|
||||
model_instance.credentials = {"k": "v"}
|
||||
model_instance.model_type_instance.get_model_schema.return_value = SimpleNamespace(features=[])
|
||||
model_instance.invoke_llm.return_value = self._llm_result("text summary")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"core.rag.index_processor.processor.paragraph_index_processor.create_plugin_provider_manager"
|
||||
) as mock_provider_manager,
|
||||
patch(
|
||||
"core.rag.index_processor.processor.paragraph_index_processor.ModelInstance",
|
||||
return_value=model_instance,
|
||||
),
|
||||
patch(
|
||||
"core.rag.index_processor.processor.paragraph_index_processor.deduct_llm_quota",
|
||||
side_effect=RuntimeError("quota"),
|
||||
),
|
||||
):
|
||||
mock_provider_manager.return_value.get_provider_model_bundle.return_value = Mock()
|
||||
with caplog.at_level(
|
||||
logging.WARNING, logger="core.rag.index_processor.processor.paragraph_index_processor"
|
||||
):
|
||||
summary, usage = ParagraphIndexProcessor.generate_summary(
|
||||
"tenant-1",
|
||||
"text content",
|
||||
{"enable": True, "model_name": "model-a", "model_provider_name": "provider-a"},
|
||||
document_language="English",
|
||||
session=self.session,
|
||||
)
|
||||
with patch(
|
||||
"core.rag.index_processor.processor.paragraph_index_processor.ModelManager.for_tenant"
|
||||
) as mock_model_manager:
|
||||
mock_model_manager.return_value.get_model_instance.return_value = model_instance
|
||||
summary, usage = ParagraphIndexProcessor.generate_summary(
|
||||
"tenant-1",
|
||||
"text content",
|
||||
{"enable": True, "model_name": "model-a", "model_provider_name": "provider-a"},
|
||||
document_language="English",
|
||||
session=self.session,
|
||||
)
|
||||
|
||||
assert summary == "text summary"
|
||||
assert isinstance(usage, LLMUsage)
|
||||
assert sum(1 for r in caplog.records if r.levelno == logging.WARNING) == 1
|
||||
assert any("Failed to deduct quota for summary generation" in record.message for record in caplog.records)
|
||||
mock_model_manager.assert_called_once_with(tenant_id="tenant-1")
|
||||
mock_model_manager.return_value.get_model_instance.assert_called_once_with(
|
||||
tenant_id="tenant-1",
|
||||
provider="provider-a",
|
||||
model_type=ModelType.LLM,
|
||||
model="model-a",
|
||||
)
|
||||
|
||||
def test_generate_summary_propagates_model_invocation_errors(self) -> None:
|
||||
model_instance = Mock()
|
||||
model_instance.credentials = {"k": "v"}
|
||||
model_instance.model_type_instance.get_model_schema.return_value = SimpleNamespace(features=[])
|
||||
model_instance.invoke_llm.side_effect = RuntimeError("invocation failed")
|
||||
|
||||
with patch(
|
||||
"core.rag.index_processor.processor.paragraph_index_processor.ModelManager.for_tenant"
|
||||
) as mock_model_manager:
|
||||
mock_model_manager.return_value.get_model_instance.return_value = model_instance
|
||||
with pytest.raises(RuntimeError, match="invocation failed"):
|
||||
ParagraphIndexProcessor.generate_summary(
|
||||
"tenant-1",
|
||||
"text content",
|
||||
{"enable": True, "model_name": "model-a", "model_provider_name": "provider-a"},
|
||||
session=self.session,
|
||||
)
|
||||
|
||||
def test_generate_summary_handles_vision_and_image_conversion(self) -> None:
|
||||
model_instance = Mock()
|
||||
@@ -567,12 +577,8 @@ class TestParagraphIndexProcessor:
|
||||
|
||||
with (
|
||||
patch(
|
||||
"core.rag.index_processor.processor.paragraph_index_processor.create_plugin_provider_manager"
|
||||
) as mock_provider_manager,
|
||||
patch(
|
||||
"core.rag.index_processor.processor.paragraph_index_processor.ModelInstance",
|
||||
return_value=model_instance,
|
||||
),
|
||||
"core.rag.index_processor.processor.paragraph_index_processor.ModelManager.for_tenant"
|
||||
) as mock_model_manager,
|
||||
patch.object(
|
||||
ParagraphIndexProcessor, "_extract_images_from_segment_attachments", return_value=[image_file]
|
||||
),
|
||||
@@ -581,9 +587,8 @@ class TestParagraphIndexProcessor:
|
||||
"core.rag.index_processor.processor.paragraph_index_processor.file_manager.to_prompt_message_content",
|
||||
return_value=image_content,
|
||||
),
|
||||
patch("core.rag.index_processor.processor.paragraph_index_processor.deduct_llm_quota"),
|
||||
):
|
||||
mock_provider_manager.return_value.get_provider_model_bundle.return_value = Mock()
|
||||
mock_model_manager.return_value.get_model_instance.return_value = model_instance
|
||||
summary, _ = ParagraphIndexProcessor.generate_summary(
|
||||
"tenant-1",
|
||||
"text content",
|
||||
@@ -606,12 +611,8 @@ class TestParagraphIndexProcessor:
|
||||
|
||||
with (
|
||||
patch(
|
||||
"core.rag.index_processor.processor.paragraph_index_processor.create_plugin_provider_manager"
|
||||
) as mock_provider_manager,
|
||||
patch(
|
||||
"core.rag.index_processor.processor.paragraph_index_processor.ModelInstance",
|
||||
return_value=model_instance,
|
||||
),
|
||||
"core.rag.index_processor.processor.paragraph_index_processor.ModelManager.for_tenant"
|
||||
) as mock_model_manager,
|
||||
patch(
|
||||
"core.rag.index_processor.processor.paragraph_index_processor.DEFAULT_GENERATOR_SUMMARY_PROMPT",
|
||||
"Prompt {missing}",
|
||||
@@ -623,7 +624,7 @@ class TestParagraphIndexProcessor:
|
||||
side_effect=RuntimeError("bad image"),
|
||||
),
|
||||
):
|
||||
mock_provider_manager.return_value.get_provider_model_bundle.return_value = Mock()
|
||||
mock_model_manager.return_value.get_model_instance.return_value = model_instance
|
||||
with pytest.raises(ValueError, match="Expected LLMResult"):
|
||||
with caplog.at_level(
|
||||
logging.WARNING, logger="core.rag.index_processor.processor.paragraph_index_processor"
|
||||
|
||||
Reference in New Issue
Block a user