dify/api/tests/unit_tests/services/test_summary_index_service.py
2026-07-15 06:48:28 +00:00

1207 lines
47 KiB
Python

"""Unit tests for services.summary_index_service."""
from __future__ import annotations
import logging
import sys
from dataclasses import dataclass
from datetime import UTC, datetime
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from sqlalchemy import create_engine, select
from sqlalchemy.exc import SAWarning
from sqlalchemy.orm import sessionmaker
import services.summary_index_service as summary_module
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
from models.dataset import DocumentSegmentSummary
from models.enums import SegmentStatus, SummaryStatus
from services.summary_index_service import SummaryIndexService
@dataclass(frozen=True)
class _SessionContext:
session: MagicMock
def __enter__(self) -> MagicMock:
return self.session
def __exit__(self, exc_type, exc, tb) -> None:
return None
def _dataset(*, indexing_technique: str = IndexTechniqueType.HIGH_QUALITY) -> MagicMock:
dataset = MagicMock(name="dataset")
dataset.id = "dataset-1"
dataset.tenant_id = "tenant-1"
dataset.indexing_technique = indexing_technique
dataset.embedding_model_provider = "openai"
dataset.embedding_model = "text-embedding"
return dataset
def _segment(*, has_document: bool = True) -> MagicMock:
segment = MagicMock(name="segment")
segment.id = "seg-1"
segment.document_id = "doc-1"
segment.dataset_id = "dataset-1"
segment.content = "hello world"
segment.enabled = True
segment.status = SegmentStatus.COMPLETED
segment.position = 1
if has_document:
doc = MagicMock(name="document")
doc.doc_language = "en"
doc.doc_form = IndexStructureType.PARAGRAPH_INDEX
segment.document = doc
segment.get_document.return_value = doc
else:
segment.document = None
segment.get_document.return_value = None
return segment
def _summary_record(*, summary_content: str = "summary", node_id: str | None = None) -> MagicMock:
record = MagicMock(spec=summary_module.DocumentSegmentSummary, name="summary_record")
record.id = "sum-1"
record.dataset_id = "dataset-1"
record.document_id = "doc-1"
record.chunk_id = "seg-1"
record.summary_content = summary_content
record.summary_index_node_id = node_id
record.summary_index_node_hash = None
record.tokens = None
record.status = SummaryStatus.GENERATING
record.error = None
record.enabled = True
record.created_at = datetime(2024, 1, 1, tzinfo=UTC)
record.updated_at = datetime(2024, 1, 1, tzinfo=UTC)
record.disabled_at = None
record.disabled_by = None
return record
def test_generate_summary_for_segment_passes_document_language(monkeypatch: pytest.MonkeyPatch) -> None:
usage = MagicMock()
usage.total_tokens = 10
usage.prompt_tokens = 3
usage.completion_tokens = 7
paragraph_module = SimpleNamespace(
ParagraphIndexProcessor=SimpleNamespace(generate_summary=MagicMock(return_value=("sum", usage)))
)
monkeypatch.setitem(
sys.modules,
"core.rag.index_processor.processor.paragraph_index_processor",
paragraph_module,
)
segment = _segment(has_document=True)
dataset = _dataset()
session = MagicMock()
content, got_usage = SummaryIndexService.generate_summary_for_segment(segment, dataset, {"a": 1}, session=session)
assert content == "sum"
assert got_usage is usage
paragraph_module.ParagraphIndexProcessor.generate_summary.assert_called_once()
_, kwargs = paragraph_module.ParagraphIndexProcessor.generate_summary.call_args
assert kwargs["document_language"] == "en"
assert kwargs["session"] is session
segment.get_document.assert_called_once_with(session=session)
def test_generate_summary_for_segment_raises_when_empty(monkeypatch: pytest.MonkeyPatch) -> None:
paragraph_module = SimpleNamespace(
ParagraphIndexProcessor=SimpleNamespace(generate_summary=MagicMock(return_value=("", MagicMock())))
)
monkeypatch.setitem(
sys.modules,
"core.rag.index_processor.processor.paragraph_index_processor",
paragraph_module,
)
with pytest.raises(ValueError, match="Generated summary is empty"):
SummaryIndexService.generate_summary_for_segment(_segment(), _dataset(), {"a": 1}, session=MagicMock())
def test_create_summary_record_updates_existing_and_reenables() -> None:
existing = _summary_record(summary_content="old", node_id="n1")
existing.enabled = False
existing.disabled_at = datetime(2024, 1, 1)
existing.disabled_by = "u"
session = MagicMock(name="session")
session.scalar.return_value = existing
segment = _segment()
dataset = _dataset()
result = SummaryIndexService.create_summary_record(
segment, dataset, "new", status=SummaryStatus.GENERATING, session=session
)
assert result is existing
assert existing.summary_content == "new"
assert existing.status == SummaryStatus.GENERATING
assert existing.enabled is True
assert existing.disabled_at is None
assert existing.disabled_by is None
assert existing.error is None
session.add.assert_called_once_with(existing)
session.flush.assert_called_once()
def test_create_summary_record_creates_new() -> None:
session = MagicMock(name="session")
session.scalar.return_value = None
record = SummaryIndexService.create_summary_record(
_segment(), _dataset(), "new", status=SummaryStatus.GENERATING, session=session
)
assert record.dataset_id == "dataset-1"
assert record.chunk_id == "seg-1"
assert record.summary_content == "new"
assert record.enabled is True
session.add.assert_called_once()
session.flush.assert_called_once()
def test_vectorize_summary_skips_non_high_quality(monkeypatch: pytest.MonkeyPatch) -> None:
vector_cls = MagicMock()
monkeypatch.setattr(summary_module, "Vector", vector_cls)
dataset = _dataset(indexing_technique=IndexTechniqueType.ECONOMY)
SummaryIndexService.vectorize_summary(_summary_record(), _segment(), dataset)
vector_cls.assert_not_called()
def test_vectorize_summary_raises_for_blank_content() -> None:
with pytest.raises(ValueError, match="Summary content is empty"):
SummaryIndexService.vectorize_summary(_summary_record(summary_content=" "), _segment(), _dataset())
def test_vectorize_summary_retries_connection_errors_then_succeeds(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
segment = _segment()
summary = _summary_record(summary_content="sum", node_id=None)
monkeypatch.setattr(summary_module.uuid, "uuid4", MagicMock(return_value="uuid-1"))
monkeypatch.setattr(summary_module.helper, "generate_text_hash", MagicMock(return_value="hash-1"))
embedding_model = MagicMock()
embedding_model.get_text_embedding_num_tokens.return_value = [5]
model_manager = MagicMock()
model_manager.get_model_instance.return_value = embedding_model
monkeypatch.setattr(summary_module.ModelManager, "for_tenant", MagicMock(return_value=model_manager))
vector_instance = MagicMock()
vector_instance.add_texts.side_effect = [RuntimeError("connection timeout"), None]
monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance))
session = MagicMock(name="provided_session")
merged = _summary_record(summary_content="sum")
session.merge.return_value = merged
monkeypatch.setattr(summary_module.time, "sleep", MagicMock())
SummaryIndexService.vectorize_summary(summary, segment, dataset, session=session)
assert vector_instance.add_texts.call_count == 2
summary_module.time.sleep.assert_called_once() # type: ignore[attr-defined]
session.flush.assert_called_once()
assert summary.status == SummaryStatus.COMPLETED
assert summary.summary_index_node_id == "uuid-1"
assert summary.summary_index_node_hash == "hash-1"
assert summary.tokens == 5
def test_vectorize_summary_without_session_creates_record_when_missing(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
segment = _segment()
summary = _summary_record(summary_content="sum", node_id="old-node")
monkeypatch.setattr(summary_module.helper, "generate_text_hash", MagicMock(return_value="hash-1"))
# Force deletion branch to run and swallow delete failures.
vector_for_delete = MagicMock()
vector_for_delete.delete_by_ids.side_effect = RuntimeError("delete failed")
vector_for_add = MagicMock()
vector_for_add.add_texts.return_value = None
vector_cls = MagicMock(side_effect=[vector_for_delete, vector_for_add])
monkeypatch.setattr(summary_module, "Vector", vector_cls)
model_manager = MagicMock()
model_manager.get_model_instance.side_effect = RuntimeError("no model")
monkeypatch.setattr(summary_module.ModelManager, "for_tenant", MagicMock(return_value=model_manager))
# New session used after vectorization succeeds (record not found by id nor chunk_id).
session = MagicMock(name="session")
session.scalar.side_effect = [None, None]
create_session_mock = MagicMock(return_value=_SessionContext(session))
monkeypatch.setattr(summary_module, "session_factory", SimpleNamespace(create_session=create_session_mock))
SummaryIndexService.vectorize_summary(summary, segment, dataset, session=None)
# Vector initialization and the record update both obtain local sessions.
create_session_mock.assert_called()
assert all(call.kwargs["session"] is session for call in vector_cls.call_args_list)
session.add.assert_called()
session.commit.assert_called_once()
assert summary.status == SummaryStatus.COMPLETED
assert summary.summary_index_node_id == "old-node" # reused
def test_vectorize_summary_final_failure_updates_error_status(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
segment = _segment()
summary = _summary_record(summary_content="sum", node_id=None)
monkeypatch.setattr(summary_module.uuid, "uuid4", MagicMock(return_value="uuid-1"))
monkeypatch.setattr(summary_module.helper, "generate_text_hash", MagicMock(return_value="hash-1"))
monkeypatch.setattr(summary_module.time, "sleep", MagicMock())
vector_instance = MagicMock()
vector_instance.add_texts.side_effect = RuntimeError("boom")
monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance))
# error_session should find record and commit status update
error_session = MagicMock(name="error_session")
error_session.scalar.return_value = summary
create_session_mock = MagicMock(return_value=_SessionContext(error_session))
monkeypatch.setattr(summary_module, "session_factory", SimpleNamespace(create_session=create_session_mock))
with pytest.raises(RuntimeError, match="boom"):
SummaryIndexService.vectorize_summary(summary, segment, dataset, session=None)
assert summary.status == SummaryStatus.ERROR
assert "Vectorization failed" in (summary.error or "")
error_session.commit.assert_called_once()
def test_batch_create_summary_records_no_segments_noop(monkeypatch: pytest.MonkeyPatch) -> None:
create_session_mock = MagicMock()
monkeypatch.setattr(summary_module, "session_factory", SimpleNamespace(create_session=create_session_mock))
SummaryIndexService.batch_create_summary_records([], _dataset())
create_session_mock.assert_not_called()
def test_batch_create_summary_records_creates_and_updates(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
s1 = _segment()
s2 = _segment()
s2.id = "seg-2"
s2.document_id = "doc-2"
existing = _summary_record()
existing.chunk_id = "seg-2"
existing.enabled = False
session = MagicMock()
session.scalars.return_value.all.return_value = [existing]
monkeypatch.setattr(
summary_module,
"session_factory",
SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))),
)
SummaryIndexService.batch_create_summary_records([s1, s2], dataset, status=SummaryStatus.NOT_STARTED)
session.commit.assert_called_once()
assert existing.enabled is True
def test_update_summary_record_error_updates_when_exists(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
segment = _segment()
record = _summary_record()
session = MagicMock()
session.scalar.return_value = record
monkeypatch.setattr(
summary_module,
"session_factory",
SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))),
)
SummaryIndexService.update_summary_record_error(segment, dataset, "err")
assert record.status == SummaryStatus.ERROR
assert record.error == "err"
session.commit.assert_called_once()
def test_generate_and_vectorize_summary_success(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
segment = _segment()
record = _summary_record(summary_content="")
session = MagicMock()
session.scalar.return_value = record
phase_events: list[str] = []
session.commit.side_effect = lambda: phase_events.append("commit")
generate_summary = MagicMock(
side_effect=lambda *_args, **_kwargs: phase_events.append("generate") or ("sum", MagicMock(total_tokens=0))
)
monkeypatch.setattr(SummaryIndexService, "generate_summary_for_segment", generate_summary)
vectorize_summary = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("vectorize"))
monkeypatch.setattr(SummaryIndexService, "vectorize_summary", vectorize_summary)
out = SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}, session=session)
assert out is record
session.refresh.assert_called_once_with(record)
assert phase_events == ["commit", "generate", "vectorize", "commit"]
def test_generate_and_vectorize_summary_vectorize_failure_sets_error(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
segment = _segment()
record = _summary_record(summary_content="")
session = MagicMock()
session.scalar.return_value = record
monkeypatch.setattr(
SummaryIndexService, "generate_summary_for_segment", MagicMock(return_value=("sum", MagicMock(total_tokens=0)))
)
monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(side_effect=RuntimeError("boom")))
with pytest.raises(RuntimeError, match="boom"):
SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}, session=session)
assert record.status == SummaryStatus.ERROR
# Outer exception handler overwrites the error with the raw exception message.
assert record.error == "boom"
session.rollback.assert_called_once()
def test_generate_and_vectorize_summary_rolls_back_failed_transaction_before_recording_error(
monkeypatch: pytest.MonkeyPatch,
) -> None:
dataset = _dataset()
segment = _segment()
record = _summary_record(summary_content="")
session = MagicMock()
rolled_back = False
scalar_calls = 0
def scalar(*_args, **_kwargs):
nonlocal scalar_calls
scalar_calls += 1
if scalar_calls > 1 and not rolled_back:
raise PendingRollbackError("rollback required")
return record
def rollback() -> None:
nonlocal rolled_back
rolled_back = True
session.scalar.side_effect = scalar
session.rollback.side_effect = rollback
session.flush.side_effect = RuntimeError("flush failed")
monkeypatch.setattr(
SummaryIndexService,
"generate_summary_for_segment",
MagicMock(return_value=("sum", MagicMock(total_tokens=0))),
)
vectorize_summary = MagicMock()
monkeypatch.setattr(SummaryIndexService, "vectorize_summary", vectorize_summary)
with pytest.raises(RuntimeError, match="flush failed"):
SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}, session=session)
assert rolled_back is True
assert record.status == SummaryStatus.ERROR
assert record.error == "flush failed"
vectorize_summary.assert_not_called()
assert session.commit.call_count == 2
def test_vectorize_summary_updates_existing_record_found_by_chunk_id(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
segment = _segment()
summary = _summary_record(summary_content="sum", node_id=None)
monkeypatch.setattr(summary_module.uuid, "uuid4", MagicMock(return_value="uuid-1"))
monkeypatch.setattr(summary_module.helper, "generate_text_hash", MagicMock(return_value="hash-1"))
vector_instance = MagicMock()
vector_instance.add_texts.return_value = None
monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance))
monkeypatch.setattr(
summary_module.ModelManager,
"for_tenant",
MagicMock(return_value=MagicMock(get_model_instance=MagicMock(return_value=None))),
)
existing = _summary_record(summary_content="old", node_id="old-node")
existing.id = "other-id"
session = MagicMock(name="session")
session.scalar.side_effect = [None, existing] # miss by id, hit by chunk_id
monkeypatch.setattr(
summary_module,
"session_factory",
SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))),
)
SummaryIndexService.vectorize_summary(summary, segment, dataset, session=None)
session.commit.assert_called_once()
assert existing.summary_index_node_id == "uuid-1"
def test_vectorize_summary_updates_existing_record_found_by_id(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
segment = _segment()
summary = _summary_record(summary_content="sum", node_id=None)
monkeypatch.setattr(summary_module.uuid, "uuid4", MagicMock(return_value="uuid-1"))
monkeypatch.setattr(summary_module.helper, "generate_text_hash", MagicMock(return_value="hash-1"))
monkeypatch.setattr(
summary_module, "Vector", MagicMock(return_value=MagicMock(add_texts=MagicMock(return_value=None)))
)
monkeypatch.setattr(
summary_module.ModelManager,
"for_tenant",
MagicMock(return_value=MagicMock(get_model_instance=MagicMock(return_value=None))),
)
existing = _summary_record(summary_content="old", node_id="old-node")
session = MagicMock(name="session")
session.scalar.return_value = existing # hit by id
monkeypatch.setattr(
summary_module,
"session_factory",
SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))),
)
SummaryIndexService.vectorize_summary(summary, segment, dataset, session=None)
session.commit.assert_called_once()
assert existing.summary_index_node_hash == "hash-1"
def test_vectorize_summary_session_enter_returns_none_triggers_runtime_error(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
segment = _segment()
summary = _summary_record(summary_content="sum", node_id=None)
monkeypatch.setattr(summary_module.uuid, "uuid4", MagicMock(return_value="uuid-1"))
monkeypatch.setattr(summary_module.helper, "generate_text_hash", MagicMock(return_value="hash-1"))
monkeypatch.setattr(
summary_module, "Vector", MagicMock(return_value=MagicMock(add_texts=MagicMock(return_value=None)))
)
monkeypatch.setattr(
summary_module.ModelManager,
"for_tenant",
MagicMock(return_value=MagicMock(get_model_instance=MagicMock(return_value=None))),
)
class _BadContext:
def __enter__(self):
return None
def __exit__(self, exc_type, exc, tb) -> None:
return None
error_session = MagicMock()
error_session.scalar.return_value = summary
vector_session = MagicMock()
create_session_mock = MagicMock(
side_effect=[_SessionContext(vector_session), _BadContext(), _SessionContext(error_session)]
)
monkeypatch.setattr(summary_module, "session_factory", SimpleNamespace(create_session=create_session_mock))
with pytest.raises(RuntimeError, match="Session should not be None"):
SummaryIndexService.vectorize_summary(summary, segment, dataset, session=None)
def test_vectorize_summary_created_record_becomes_none_triggers_guard(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
segment = _segment()
summary = _summary_record(summary_content="sum", node_id=None)
monkeypatch.setattr(summary_module.uuid, "uuid4", MagicMock(return_value="uuid-1"))
monkeypatch.setattr(summary_module.helper, "generate_text_hash", MagicMock(return_value="hash-1"))
monkeypatch.setattr(
summary_module, "Vector", MagicMock(return_value=MagicMock(add_texts=MagicMock(return_value=None)))
)
monkeypatch.setattr(
summary_module.ModelManager,
"for_tenant",
MagicMock(return_value=MagicMock(get_model_instance=MagicMock(return_value=None))),
)
session = MagicMock()
session.scalar.side_effect = [None, None] # miss by id and chunk_id
error_session = MagicMock()
error_session.scalar.return_value = summary
vector_session = MagicMock()
create_session_mock = MagicMock(
side_effect=[_SessionContext(vector_session), _SessionContext(session), _SessionContext(error_session)]
)
monkeypatch.setattr(summary_module, "session_factory", SimpleNamespace(create_session=create_session_mock))
# Force the created record to be None so the "should not be None" guard triggers.
# Also mock select() so SQLAlchemy doesn't validate the mocked DocumentSegmentSummary as a real column clause.
monkeypatch.setattr(summary_module, "select", MagicMock(return_value=MagicMock()))
monkeypatch.setattr(summary_module, "DocumentSegmentSummary", MagicMock(return_value=None))
with pytest.raises(RuntimeError, match="summary_record_in_session should not be None"):
SummaryIndexService.vectorize_summary(summary, segment, dataset, session=None)
def test_vectorize_summary_error_handler_tries_chunk_id_lookup_and_can_warn_not_found(
monkeypatch: pytest.MonkeyPatch,
) -> None:
dataset = _dataset()
segment = _segment()
summary = _summary_record(summary_content="sum", node_id=None)
monkeypatch.setattr(summary_module.uuid, "uuid4", MagicMock(return_value="uuid-1"))
monkeypatch.setattr(summary_module.helper, "generate_text_hash", MagicMock(return_value="hash-1"))
monkeypatch.setattr(summary_module.time, "sleep", MagicMock())
monkeypatch.setattr(
summary_module,
"Vector",
MagicMock(return_value=MagicMock(add_texts=MagicMock(side_effect=RuntimeError("boom")))),
)
error_session = MagicMock(name="error_session")
error_session.scalar.side_effect = [None, None] # not found by id, not found by chunk_id
monkeypatch.setattr(
summary_module,
"session_factory",
SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(error_session))),
)
with pytest.raises(RuntimeError, match="boom"):
SummaryIndexService.vectorize_summary(summary, segment, dataset, session=None)
# No record -> no commit in error session.
error_session.commit.assert_not_called()
def test_update_summary_record_error_warns_when_missing(
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
) -> None:
dataset = _dataset()
segment = _segment()
session = MagicMock()
session.scalar.return_value = None
monkeypatch.setattr(
summary_module,
"session_factory",
SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))),
)
with caplog.at_level(logging.WARNING, logger="services.summary_index_service"):
SummaryIndexService.update_summary_record_error(segment, dataset, "err")
assert any(r.levelno >= logging.WARNING for r in caplog.records)
def test_generate_and_vectorize_summary_creates_missing_record_and_logs_usage(
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
) -> None:
dataset = _dataset()
segment = _segment()
session = MagicMock()
session.scalar.return_value = None
usage = MagicMock(total_tokens=4, prompt_tokens=1, completion_tokens=3)
monkeypatch.setattr(SummaryIndexService, "generate_summary_for_segment", MagicMock(return_value=("sum", usage)))
monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(return_value=None))
with caplog.at_level(logging.INFO, logger="services.summary_index_service"):
result = SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}, session=session)
assert result.status in {SummaryStatus.GENERATING, SummaryStatus.COMPLETED}
assert any(r.levelno >= logging.INFO for r in caplog.records)
def test_generate_summaries_for_document_skip_conditions(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset(indexing_technique=IndexTechniqueType.ECONOMY)
document = MagicMock(spec=summary_module.DatasetDocument)
document.id = "doc-1"
document.doc_form = IndexStructureType.PARAGRAPH_INDEX
assert SummaryIndexService.generate_summaries_for_document(dataset, document, {"enable": True}) == []
dataset = _dataset()
assert SummaryIndexService.generate_summaries_for_document(dataset, document, {"enable": False}) == []
document.doc_form = IndexStructureType.QA_INDEX
assert SummaryIndexService.generate_summaries_for_document(dataset, document, {"enable": True}) == []
def test_generate_summaries_for_document_runs_and_handles_errors(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
document = MagicMock(spec=summary_module.DatasetDocument)
document.id = "doc-1"
document.doc_form = IndexStructureType.PARAGRAPH_INDEX
seg1 = _segment()
seg2 = _segment()
seg2.id = "seg-2"
session = MagicMock()
session.scalars.return_value.all.return_value = [seg1, seg2]
monkeypatch.setattr(
summary_module,
"session_factory",
SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))),
)
monkeypatch.setattr(SummaryIndexService, "batch_create_summary_records", MagicMock())
monkeypatch.setattr(
SummaryIndexService,
"generate_and_vectorize_summary",
MagicMock(side_effect=[MagicMock(), RuntimeError("boom")]),
)
update_err_mock = MagicMock()
monkeypatch.setattr(SummaryIndexService, "update_summary_record_error", update_err_mock)
records = SummaryIndexService.generate_summaries_for_document(dataset, document, {"enable": True})
assert len(records) == 1
update_err_mock.assert_called_once()
def test_generate_summaries_for_document_no_segments_returns_empty(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
document = MagicMock(spec=summary_module.DatasetDocument)
document.id = "doc-1"
document.doc_form = IndexStructureType.PARAGRAPH_INDEX
session = MagicMock()
session.scalars.return_value.all.return_value = []
monkeypatch.setattr(
summary_module,
"session_factory",
SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))),
)
assert SummaryIndexService.generate_summaries_for_document(dataset, document, {"enable": True}) == []
def test_generate_summaries_for_document_applies_segment_ids_and_only_parent_chunks(
monkeypatch: pytest.MonkeyPatch,
) -> None:
dataset = _dataset()
document = MagicMock(spec=summary_module.DatasetDocument)
document.id = "doc-1"
document.doc_form = IndexStructureType.PARAGRAPH_INDEX
seg = _segment()
session = MagicMock()
session.scalars.return_value.all.return_value = [seg]
monkeypatch.setattr(
summary_module,
"session_factory",
SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))),
)
monkeypatch.setattr(SummaryIndexService, "batch_create_summary_records", MagicMock())
monkeypatch.setattr(SummaryIndexService, "generate_and_vectorize_summary", MagicMock(return_value=MagicMock()))
SummaryIndexService.generate_summaries_for_document(
dataset,
document,
{"enable": True},
segment_ids=[seg.id],
only_parent_chunks=True,
)
session.scalars.assert_called()
def test_disable_summaries_for_segments_updates_sqlite_records() -> None:
dataset = SimpleNamespace(id="dataset-1", indexing_technique=IndexTechniqueType.ECONOMY)
engine = create_engine("sqlite+pysqlite:///:memory:")
DocumentSegmentSummary.__table__.create(engine)
summary_rows = [
{
"id": "sum-1",
"dataset_id": dataset.id,
"document_id": "doc-1",
"chunk_id": "seg-1",
"summary_content": "s",
"summary_index_node_id": "n1",
"status": SummaryStatus.COMPLETED,
"enabled": True,
},
{
"id": "sum-2",
"dataset_id": dataset.id,
"document_id": "doc-1",
"chunk_id": "seg-1",
"summary_content": "s",
"summary_index_node_id": None,
"status": SummaryStatus.COMPLETED,
"enabled": True,
},
]
with engine.begin() as connection:
connection.execute(DocumentSegmentSummary.__table__.insert(), summary_rows)
session_maker = sessionmaker(bind=engine, expire_on_commit=False)
summary_module.session_factory.configure(engine, expire_on_commit=False)
SummaryIndexService.disable_summaries_for_segments(dataset, segment_ids=["seg-1"], disabled_by="u")
with session_maker() as session:
summaries = session.scalars(select(DocumentSegmentSummary).order_by(DocumentSegmentSummary.id)).all()
assert [(summary.id, summary.enabled, summary.disabled_by) for summary in summaries] == [
("sum-1", False, "u"),
("sum-2", False, "u"),
]
assert all(summary.disabled_at is not None for summary in summaries)
def test_disable_summaries_for_segments_no_summaries_noop(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
session = MagicMock()
session.scalars.return_value.all.return_value = []
monkeypatch.setattr(
summary_module,
"session_factory",
SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))),
)
monkeypatch.setitem(
sys.modules, "libs.datetime_utils", SimpleNamespace(naive_utc_now=MagicMock(return_value=datetime(2024, 1, 1)))
)
SummaryIndexService.disable_summaries_for_segments(dataset)
session.commit.assert_not_called()
def test_enable_summaries_for_segments_skips_non_high_quality() -> None:
SummaryIndexService.enable_summaries_for_segments(_dataset(indexing_technique=IndexTechniqueType.ECONOMY))
def test_enable_summaries_for_segments_revectorizes_and_enables(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
summary = _summary_record(summary_content="sum", node_id="n1")
summary.enabled = False
segment = _segment()
segment.id = summary.chunk_id
segment.enabled = True
segment.status = SegmentStatus.COMPLETED
session = MagicMock()
session.scalars.return_value.all.return_value = [summary]
session.scalar.return_value = segment
monkeypatch.setattr(
summary_module,
"session_factory",
SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))),
)
vec_mock = MagicMock()
monkeypatch.setattr(SummaryIndexService, "vectorize_summary", vec_mock)
SummaryIndexService.enable_summaries_for_segments(dataset, segment_ids=[summary.chunk_id])
vec_mock.assert_called_once()
assert summary.enabled is True
session.commit.assert_called_once()
def test_enable_summaries_for_segments_no_summaries_noop(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
session = MagicMock()
session.scalars.return_value.all.return_value = []
monkeypatch.setattr(
summary_module,
"session_factory",
SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))),
)
SummaryIndexService.enable_summaries_for_segments(dataset)
session.commit.assert_not_called()
def test_enable_summaries_for_segments_skips_segment_or_content_and_handles_vectorize_error(
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
) -> None:
dataset = _dataset()
summary1 = _summary_record(summary_content="sum", node_id="n1")
summary1.enabled = False
summary2 = _summary_record(summary_content="", node_id="n2")
summary2.enabled = False
summary3 = _summary_record(summary_content="sum3", node_id="n3")
summary3.enabled = False
bad_segment = _segment()
bad_segment.enabled = False
bad_segment.status = SegmentStatus.COMPLETED
good_segment = _segment()
good_segment.enabled = True
good_segment.status = SegmentStatus.COMPLETED
session = MagicMock()
session.scalars.return_value.all.return_value = [summary1, summary2, summary3]
session.scalar.side_effect = [bad_segment, good_segment, good_segment]
monkeypatch.setattr(
summary_module,
"session_factory",
SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))),
)
monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(side_effect=RuntimeError("boom")))
with caplog.at_level(logging.ERROR, logger="services.summary_index_service"):
SummaryIndexService.enable_summaries_for_segments(dataset)
assert any(r.levelno >= logging.ERROR for r in caplog.records)
session.commit.assert_called_once()
def test_delete_summaries_for_segments_deletes_vectors_and_records(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
summary = _summary_record(summary_content="sum", node_id="n1")
session = MagicMock()
session.scalars.return_value.all.return_value = [summary]
vector_instance = MagicMock()
monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance))
SummaryIndexService.delete_summaries_for_segments(dataset, segment_ids=[summary.chunk_id], session=session)
vector_instance.delete_by_ids.assert_called_once_with(["n1"])
session.delete.assert_called_once_with(summary)
session.flush.assert_called_once()
def test_delete_summaries_for_segments_no_summaries_noop(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
session = MagicMock()
session.scalars.return_value.all.return_value = []
SummaryIndexService.delete_summaries_for_segments(dataset, session=session)
session.flush.assert_not_called()
def test_update_summary_for_segment_skip_conditions() -> None:
session = MagicMock()
economy_dataset = _dataset(indexing_technique=IndexTechniqueType.ECONOMY)
assert SummaryIndexService.update_summary_for_segment(_segment(), economy_dataset, "x", session=session) is None
seg = _segment(has_document=True)
seg.get_document.return_value.doc_form = IndexStructureType.QA_INDEX
assert SummaryIndexService.update_summary_for_segment(seg, _dataset(), "x", session=session) is None
def test_update_summary_for_segment_empty_content_deletes_existing(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
segment = _segment()
record = _summary_record(summary_content="old", node_id="n1")
session = MagicMock()
session.scalar.return_value = record
vector_instance = MagicMock()
monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance))
assert SummaryIndexService.update_summary_for_segment(segment, dataset, " ", session=session) is None
vector_instance.delete_by_ids.assert_called_once_with(["n1"])
session.delete.assert_called_once_with(record)
session.commit.assert_called_once()
def test_update_summary_for_segment_empty_content_delete_vector_warns(
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
) -> None:
dataset = _dataset()
segment = _segment()
record = _summary_record(summary_content="old", node_id="n1")
session = MagicMock()
session.scalar.return_value = record
vector_instance = MagicMock()
vector_instance.delete_by_ids.side_effect = RuntimeError("boom")
monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance))
with caplog.at_level(logging.WARNING, logger="services.summary_index_service"):
assert SummaryIndexService.update_summary_for_segment(segment, dataset, "", session=session) is None
assert any(r.levelno >= logging.WARNING for r in caplog.records)
def test_update_summary_for_segment_empty_content_no_record_noop(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
segment = _segment()
session = MagicMock()
session.scalar.return_value = None
assert SummaryIndexService.update_summary_for_segment(segment, dataset, " ", session=session) is None
def test_update_summary_for_segment_updates_existing_and_vectorizes(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
segment = _segment()
record = _summary_record(summary_content="old", node_id="n1")
session = MagicMock()
session.scalar.return_value = record
vector_instance = MagicMock()
monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance))
vectorize_mock = MagicMock()
monkeypatch.setattr(SummaryIndexService, "vectorize_summary", vectorize_mock)
out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new summary", session=session)
assert out is record
vectorize_mock.assert_called_once()
session.refresh.assert_called_once_with(record)
session.commit.assert_called()
def test_update_summary_for_segment_existing_vector_delete_warns(
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
) -> None:
dataset = _dataset()
segment = _segment()
record = _summary_record(summary_content="old", node_id="n1")
session = MagicMock()
session.scalar.return_value = record
vector_instance = MagicMock()
vector_instance.delete_by_ids.side_effect = RuntimeError("boom")
monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance))
monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(return_value=None))
with caplog.at_level(logging.WARNING, logger="services.summary_index_service"):
SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session)
assert any(r.levelno >= logging.WARNING for r in caplog.records)
def test_update_summary_for_segment_existing_vectorize_failure_returns_error_record(
monkeypatch: pytest.MonkeyPatch,
) -> None:
dataset = _dataset()
segment = _segment()
record = _summary_record(summary_content="old", node_id="n1")
session = MagicMock(is_active=True)
session.scalar.return_value = record
monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(side_effect=RuntimeError("boom")))
out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session)
assert out is record
assert out.summary_content == "new"
assert out.status == SummaryStatus.ERROR
assert "Vectorization failed" in (out.error or "")
session.rollback.assert_not_called()
session.add.assert_called_with(record)
session.commit.assert_called_once()
def test_update_summary_for_segment_failed_flush_persists_new_content(monkeypatch: pytest.MonkeyPatch) -> None:
engine = create_engine("sqlite+pysqlite:///:memory:")
DocumentSegmentSummary.__table__.create(engine)
session_maker = sessionmaker(bind=engine, expire_on_commit=False)
with session_maker() as session:
record = DocumentSegmentSummary(
dataset_id="dataset-1",
document_id="doc-1",
chunk_id="seg-1",
summary_content="old",
status=SummaryStatus.COMPLETED,
)
record.id = "sum-1"
session.add(record)
session.commit()
def fail_flush(*_args, **_kwargs) -> None:
duplicate = DocumentSegmentSummary(
dataset_id="dataset-1",
document_id="doc-1",
chunk_id="seg-2",
summary_content="duplicate",
)
duplicate.id = record.id
session.add(duplicate)
session.flush()
segment = _segment()
segment.get_document.return_value = SimpleNamespace(doc_form=IndexStructureType.PARAGRAPH_INDEX)
monkeypatch.setattr(SummaryIndexService, "vectorize_summary", fail_flush)
with pytest.warns(SAWarning, match="conflicts with persistent instance"):
out = SummaryIndexService.update_summary_for_segment(segment, _dataset(), "new", session=session)
session.expire_all()
persisted = session.get(DocumentSegmentSummary, record.id)
assert out is record
assert persisted is not None
assert persisted.summary_content == "new"
assert persisted.status == SummaryStatus.ERROR
def test_update_summary_for_segment_new_record_success(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
segment = _segment()
session = MagicMock()
session.scalar.return_value = None
created = _summary_record(summary_content="new", node_id=None)
monkeypatch.setattr(SummaryIndexService, "create_summary_record", MagicMock(return_value=created))
monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(return_value=None))
out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session)
assert out is created
session.refresh.assert_called()
session.commit.assert_called()
def test_update_summary_for_segment_outer_exception_sets_error_and_reraises(monkeypatch: pytest.MonkeyPatch) -> None:
dataset = _dataset()
segment = _segment()
record = _summary_record(summary_content="old", node_id="n1")
session = MagicMock()
session.scalar.return_value = record
session.flush.side_effect = RuntimeError("flush boom")
with pytest.raises(RuntimeError, match="flush boom"):
SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session)
assert record.status == SummaryStatus.ERROR
assert record.error == "flush boom"
session.rollback.assert_called_once()
session.commit.assert_called()
def test_get_segment_summary_and_document_summaries() -> None:
record = _summary_record(summary_content="sum", node_id="n1")
session = MagicMock()
session.scalar.return_value = record
session.scalars.return_value.all.return_value = [record]
assert SummaryIndexService.get_segment_summary("seg-1", "dataset-1", session=session) is record
assert SummaryIndexService.get_document_summaries("doc-1", "dataset-1", segment_ids=["seg-1"], session=session) == [
record
]
def test_get_segments_summaries_non_empty() -> None:
record1 = _summary_record()
record1.chunk_id = "seg-1"
record2 = _summary_record()
record2.chunk_id = "seg-2"
session = MagicMock()
session.scalars.return_value.all.return_value = [record1, record2]
out = SummaryIndexService.get_segments_summaries(["seg-1", "seg-2"], "dataset-1", session=session)
assert set(out.keys()) == {"seg-1", "seg-2"}
def test_get_document_summary_index_status_no_segments_returns_none() -> None:
session = MagicMock()
session.scalars.return_value.all.return_value = []
assert (
SummaryIndexService.get_document_summary_index_status("doc-1", "dataset-1", "tenant-1", session=session) is None
)
def test_get_documents_summary_index_status_empty_input() -> None:
assert (
SummaryIndexService.get_documents_summary_index_status([], "dataset-1", "tenant-1", session=MagicMock()) == {}
)
def test_get_documents_summary_index_status_no_pending_sets_none(monkeypatch: pytest.MonkeyPatch) -> None:
session = MagicMock()
session.execute.return_value.all.return_value = [SimpleNamespace(id="seg-1", document_id="doc-1")]
monkeypatch.setattr(
SummaryIndexService,
"get_segments_summaries",
MagicMock(return_value={"seg-1": SimpleNamespace(status=SummaryStatus.COMPLETED)}),
)
result = SummaryIndexService.get_documents_summary_index_status(["doc-1"], "dataset-1", "tenant-1", session=session)
assert result["doc-1"] is None
def test_update_summary_for_segment_creates_new_and_vectorize_fails_returns_error_record(
monkeypatch: pytest.MonkeyPatch,
) -> None:
dataset = _dataset()
segment = _segment()
session = MagicMock(is_active=False)
session.scalar.return_value = None
created = _summary_record(summary_content="new", node_id=None)
monkeypatch.setattr(SummaryIndexService, "create_summary_record", MagicMock(return_value=created))
monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(side_effect=RuntimeError("boom")))
out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session)
assert out.status == SummaryStatus.ERROR
assert "Vectorization failed" in (out.error or "")
session.rollback.assert_called_once()
session.add.assert_called_with(created)
session.commit.assert_called_once()
def test_get_segments_summaries_empty_list() -> None:
assert SummaryIndexService.get_segments_summaries([], "dataset-1", session=MagicMock()) == {}
def test_get_document_summary_index_status_and_documents_status(monkeypatch: pytest.MonkeyPatch) -> None:
seg_row = SimpleNamespace(id="seg-1", document_id="doc-1")
session = MagicMock()
session.scalars.return_value.all.return_value = ["seg-1"] # get_document_summary_index_status returns IDs
monkeypatch.setattr(
SummaryIndexService,
"get_segments_summaries",
MagicMock(return_value={"seg-1": SimpleNamespace(status=SummaryStatus.GENERATING)}),
)
assert (
SummaryIndexService.get_document_summary_index_status("doc-1", "dataset-1", "tenant-1", session=session)
== "SUMMARIZING"
)
# Multiple docs
session2 = MagicMock()
session2.execute.return_value.all.return_value = [seg_row] # get_documents_summary_index_status uses execute
monkeypatch.setattr(
SummaryIndexService,
"get_segments_summaries",
MagicMock(return_value={"seg-1": SimpleNamespace(status=SummaryStatus.NOT_STARTED)}),
)
result = SummaryIndexService.get_documents_summary_index_status(
["doc-1", "doc-2"], "dataset-1", "tenant-1", session=session2
)
assert result["doc-1"] == "SUMMARIZING"
assert result["doc-2"] is None
def test_get_document_summary_status_detail_counts_and_previews(monkeypatch: pytest.MonkeyPatch) -> None:
segment1 = _segment()
segment1.id = "seg-1"
segment1.position = 1
segment2 = _segment()
segment2.id = "seg-2"
segment2.position = 2
summary1 = _summary_record(summary_content="x" * 150, node_id="n1")
summary1.chunk_id = "seg-1"
summary1.status = SummaryStatus.COMPLETED
summary1.error = None
summary1.created_at = datetime(2024, 1, 1, tzinfo=UTC)
summary1.updated_at = datetime(2024, 1, 2, tzinfo=UTC)
segment_service = SimpleNamespace(get_segments_by_document_and_dataset=MagicMock(return_value=[segment1, segment2]))
monkeypatch.setitem(sys.modules, "services.dataset_service", SimpleNamespace(SegmentService=segment_service))
monkeypatch.setattr(SummaryIndexService, "get_document_summaries", MagicMock(return_value=[summary1]))
detail = SummaryIndexService.get_document_summary_status_detail("doc-1", "dataset-1", MagicMock())
assert detail["total_segments"] == 2
assert detail["summary_status"]["completed"] == 1
assert detail["summary_status"]["not_started"] == 1
assert detail["summaries"][0]["summary_preview"].endswith("...")
assert detail["summaries"][1]["status"] == "not_started"