mirror of
https://github.com/langgenius/dify.git
synced 2026-09-05 16:55:14 +08:00
test: use sqlite3 session in test_annotation_reply (#38714)
This commit is contained in:
parent
e044292518
commit
6fa71ef7ab
@ -1,12 +1,13 @@
|
|||||||
import logging
|
import logging
|
||||||
|
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||||
from core.rag.datasource.vdb.vector_factory import Vector
|
from core.rag.datasource.vdb.vector_factory import Vector
|
||||||
from core.rag.index_processor.constant.index_type import IndexTechniqueType
|
from core.rag.index_processor.constant.index_type import IndexTechniqueType
|
||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
from models.dataset import Dataset
|
from models.dataset import Dataset, DatasetCollectionBinding
|
||||||
from models.enums import CollectionBindingType, ConversationFromSource
|
from models.enums import CollectionBindingType, ConversationFromSource
|
||||||
from models.model import App, AppAnnotationSetting, Message, MessageAnnotation
|
from models.model import App, AppAnnotationSetting, Message, MessageAnnotation
|
||||||
from services.annotation_service import AppAnnotationService
|
from services.annotation_service import AppAnnotationService
|
||||||
@ -17,24 +18,33 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
class AnnotationReplyFeature:
|
class AnnotationReplyFeature:
|
||||||
def query(
|
def query(
|
||||||
self, app_record: App, message: Message, query: str, user_id: str, invoke_from: InvokeFrom
|
self,
|
||||||
|
app_record: App,
|
||||||
|
message: Message,
|
||||||
|
query: str,
|
||||||
|
user_id: str,
|
||||||
|
invoke_from: InvokeFrom,
|
||||||
|
*,
|
||||||
|
session: Session | None = None,
|
||||||
) -> MessageAnnotation | None:
|
) -> MessageAnnotation | None:
|
||||||
|
"""Return the closest annotation reply and record a hit in ``session``.
|
||||||
|
|
||||||
|
The caller may provide its transaction so the setting lookup, annotation
|
||||||
|
lookup, and hit-history write share one session. Runtime callers that do
|
||||||
|
not provide one continue to use Flask-SQLAlchemy's scoped session.
|
||||||
|
Vector-search failures are logged and return ``None``; transaction
|
||||||
|
cleanup remains the caller's responsibility.
|
||||||
"""
|
"""
|
||||||
Query app annotations to reply
|
if session is None:
|
||||||
:param app_record: app record
|
session = db.session()
|
||||||
:param message: message
|
|
||||||
:param query: query
|
|
||||||
:param user_id: user id
|
|
||||||
:param invoke_from: invoke from
|
|
||||||
:return:
|
|
||||||
"""
|
|
||||||
stmt = select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_record.id)
|
stmt = select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_record.id)
|
||||||
annotation_setting = db.session.scalar(stmt)
|
annotation_setting = session.scalar(stmt)
|
||||||
|
|
||||||
if not annotation_setting:
|
if not annotation_setting:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
collection_binding_detail = annotation_setting.collection_binding_detail
|
collection_binding_detail = session.get(DatasetCollectionBinding, annotation_setting.collection_binding_id)
|
||||||
|
|
||||||
if not collection_binding_detail:
|
if not collection_binding_detail:
|
||||||
return None
|
return None
|
||||||
@ -45,7 +55,7 @@ class AnnotationReplyFeature:
|
|||||||
embedding_model_name = collection_binding_detail.model_name
|
embedding_model_name = collection_binding_detail.model_name
|
||||||
|
|
||||||
dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding(
|
dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding(
|
||||||
embedding_provider_name, embedding_model_name, db.session(), CollectionBindingType.ANNOTATION
|
embedding_provider_name, embedding_model_name, session, CollectionBindingType.ANNOTATION
|
||||||
)
|
)
|
||||||
|
|
||||||
dataset = Dataset(
|
dataset = Dataset(
|
||||||
@ -66,7 +76,7 @@ class AnnotationReplyFeature:
|
|||||||
if documents and documents[0].metadata:
|
if documents and documents[0].metadata:
|
||||||
annotation_id = documents[0].metadata["annotation_id"]
|
annotation_id = documents[0].metadata["annotation_id"]
|
||||||
score = documents[0].metadata["score"]
|
score = documents[0].metadata["score"]
|
||||||
annotation = AppAnnotationService.get_annotation_by_id(annotation_id, session=db.session())
|
annotation = AppAnnotationService.get_annotation_by_id(annotation_id, session=session)
|
||||||
if annotation:
|
if annotation:
|
||||||
if invoke_from in {InvokeFrom.SERVICE_API, InvokeFrom.WEB_APP}:
|
if invoke_from in {InvokeFrom.SERVICE_API, InvokeFrom.WEB_APP}:
|
||||||
from_source = ConversationFromSource.API
|
from_source = ConversationFromSource.API
|
||||||
@ -84,7 +94,7 @@ class AnnotationReplyFeature:
|
|||||||
message.id,
|
message.id,
|
||||||
from_source,
|
from_source,
|
||||||
score,
|
score,
|
||||||
session=db.session(),
|
session=session,
|
||||||
)
|
)
|
||||||
|
|
||||||
return annotation
|
return annotation
|
||||||
|
|||||||
@ -3,163 +3,171 @@ from types import SimpleNamespace
|
|||||||
from unittest.mock import Mock, patch
|
from unittest.mock import Mock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||||
from core.app.features.annotation_reply.annotation_reply import AnnotationReplyFeature
|
from core.app.features.annotation_reply.annotation_reply import AnnotationReplyFeature
|
||||||
|
from models.dataset import DatasetCollectionBinding
|
||||||
|
from models.enums import CollectionBindingType, ConversationFromSource
|
||||||
|
from models.model import AppAnnotationHitHistory, AppAnnotationSetting, MessageAnnotation
|
||||||
|
|
||||||
|
TABLES = (AppAnnotationSetting, DatasetCollectionBinding, MessageAnnotation, AppAnnotationHitHistory)
|
||||||
|
|
||||||
|
|
||||||
|
def _persist_binding(session: Session) -> DatasetCollectionBinding:
|
||||||
|
binding = DatasetCollectionBinding(
|
||||||
|
provider_name="prov",
|
||||||
|
model_name="model",
|
||||||
|
type=CollectionBindingType.ANNOTATION,
|
||||||
|
collection_name="annotation-collection",
|
||||||
|
)
|
||||||
|
session.add(binding)
|
||||||
|
session.flush()
|
||||||
|
return binding
|
||||||
|
|
||||||
|
|
||||||
|
def _persist_setting(
|
||||||
|
session: Session,
|
||||||
|
*,
|
||||||
|
app_id: str = "app-1",
|
||||||
|
collection_binding_id: str,
|
||||||
|
score_threshold: float = 0.5,
|
||||||
|
) -> AppAnnotationSetting:
|
||||||
|
setting = AppAnnotationSetting(
|
||||||
|
app_id=app_id,
|
||||||
|
score_threshold=score_threshold,
|
||||||
|
collection_binding_id=collection_binding_id,
|
||||||
|
created_user_id="user-1",
|
||||||
|
updated_user_id="user-1",
|
||||||
|
)
|
||||||
|
session.add(setting)
|
||||||
|
session.flush()
|
||||||
|
return setting
|
||||||
|
|
||||||
|
|
||||||
|
def _persist_annotation(session: Session) -> MessageAnnotation:
|
||||||
|
annotation = MessageAnnotation(
|
||||||
|
app_id="app-1",
|
||||||
|
question="question",
|
||||||
|
content="content",
|
||||||
|
account_id="acct-1",
|
||||||
|
)
|
||||||
|
session.add(annotation)
|
||||||
|
session.flush()
|
||||||
|
return annotation
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
|
||||||
class TestAnnotationReplyFeature:
|
class TestAnnotationReplyFeature:
|
||||||
def test_query_returns_none_when_setting_missing(self):
|
def test_query_returns_none_when_setting_missing(self, sqlite_session: Session):
|
||||||
feature = AnnotationReplyFeature()
|
binding = _persist_binding(sqlite_session)
|
||||||
|
_persist_setting(sqlite_session, app_id="other-app", collection_binding_id=binding.id)
|
||||||
|
|
||||||
with patch("core.app.features.annotation_reply.annotation_reply.db") as mock_db:
|
result = AnnotationReplyFeature().query(
|
||||||
mock_db.session.scalar.return_value = None
|
app_record=SimpleNamespace(id="app-1", tenant_id="tenant-1"),
|
||||||
|
message=SimpleNamespace(id="msg-1"),
|
||||||
result = feature.query(
|
query="hi",
|
||||||
app_record=SimpleNamespace(id="app-1", tenant_id="tenant-1"),
|
user_id="user-1",
|
||||||
message=SimpleNamespace(id="msg-1"),
|
invoke_from=InvokeFrom.SERVICE_API,
|
||||||
query="hi",
|
session=sqlite_session,
|
||||||
user_id="user-1",
|
)
|
||||||
invoke_from=InvokeFrom.SERVICE_API,
|
|
||||||
)
|
|
||||||
|
|
||||||
assert result is None
|
assert result is None
|
||||||
|
|
||||||
def test_query_returns_none_when_binding_missing(self):
|
def test_query_returns_none_when_binding_missing(self, sqlite_session: Session):
|
||||||
feature = AnnotationReplyFeature()
|
_persist_setting(sqlite_session, collection_binding_id="missing-binding")
|
||||||
annotation_setting = SimpleNamespace(collection_binding_detail=None)
|
|
||||||
|
|
||||||
with patch("core.app.features.annotation_reply.annotation_reply.db") as mock_db:
|
result = AnnotationReplyFeature().query(
|
||||||
mock_db.session.scalar.return_value = annotation_setting
|
app_record=SimpleNamespace(id="app-1", tenant_id="tenant-1"),
|
||||||
|
message=SimpleNamespace(id="msg-1"),
|
||||||
result = feature.query(
|
query="hi",
|
||||||
app_record=SimpleNamespace(id="app-1", tenant_id="tenant-1"),
|
user_id="user-1",
|
||||||
message=SimpleNamespace(id="msg-1"),
|
invoke_from=InvokeFrom.SERVICE_API,
|
||||||
query="hi",
|
session=sqlite_session,
|
||||||
user_id="user-1",
|
)
|
||||||
invoke_from=InvokeFrom.SERVICE_API,
|
|
||||||
)
|
|
||||||
|
|
||||||
assert result is None
|
assert result is None
|
||||||
|
|
||||||
def test_query_returns_annotation_and_records_history_for_api(self):
|
def test_query_returns_annotation_and_persists_history_for_api(self, sqlite_session: Session):
|
||||||
feature = AnnotationReplyFeature()
|
binding = _persist_binding(sqlite_session)
|
||||||
annotation_setting = SimpleNamespace(
|
_persist_setting(sqlite_session, collection_binding_id=binding.id, score_threshold=0)
|
||||||
score_threshold=None,
|
annotation = _persist_annotation(sqlite_session)
|
||||||
collection_binding_detail=SimpleNamespace(provider_name="prov", model_name="model"),
|
document = SimpleNamespace(metadata={"annotation_id": annotation.id, "score": 0.8})
|
||||||
)
|
|
||||||
dataset_binding = SimpleNamespace(id="binding-1")
|
|
||||||
annotation = SimpleNamespace(
|
|
||||||
id="ann-1",
|
|
||||||
question_text="question",
|
|
||||||
content="content",
|
|
||||||
account_id="acct-1",
|
|
||||||
account=SimpleNamespace(name="Alice"),
|
|
||||||
)
|
|
||||||
document = SimpleNamespace(metadata={"annotation_id": "ann-1", "score": 0.8})
|
|
||||||
vector_instance = Mock()
|
vector_instance = Mock()
|
||||||
vector_instance.search_by_vector.return_value = [document]
|
vector_instance.search_by_vector.return_value = [document]
|
||||||
|
|
||||||
with (
|
with patch("core.app.features.annotation_reply.annotation_reply.Vector", return_value=vector_instance):
|
||||||
patch("core.app.features.annotation_reply.annotation_reply.db") as mock_db,
|
result = AnnotationReplyFeature().query(
|
||||||
patch(
|
|
||||||
"core.app.features.annotation_reply.annotation_reply.DatasetCollectionBindingService"
|
|
||||||
) as mock_binding_service,
|
|
||||||
patch("core.app.features.annotation_reply.annotation_reply.Vector") as mock_vector,
|
|
||||||
patch(
|
|
||||||
"core.app.features.annotation_reply.annotation_reply.AppAnnotationService"
|
|
||||||
) as mock_annotation_service,
|
|
||||||
):
|
|
||||||
mock_db.session.scalar.return_value = annotation_setting
|
|
||||||
mock_binding_service.get_dataset_collection_binding.return_value = dataset_binding
|
|
||||||
mock_vector.return_value = vector_instance
|
|
||||||
mock_annotation_service.get_annotation_by_id.return_value = annotation
|
|
||||||
|
|
||||||
result = feature.query(
|
|
||||||
app_record=SimpleNamespace(id="app-1", tenant_id="tenant-1"),
|
app_record=SimpleNamespace(id="app-1", tenant_id="tenant-1"),
|
||||||
message=SimpleNamespace(id="msg-1"),
|
message=SimpleNamespace(id="msg-1"),
|
||||||
query="hi",
|
query="hi",
|
||||||
user_id="user-1",
|
user_id="user-1",
|
||||||
invoke_from=InvokeFrom.SERVICE_API,
|
invoke_from=InvokeFrom.SERVICE_API,
|
||||||
|
session=sqlite_session,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result == annotation
|
assert result is annotation
|
||||||
mock_annotation_service.add_annotation_history.assert_called_once()
|
vector_instance.search_by_vector.assert_called_once_with(
|
||||||
_, _, _, _, _, _, _, from_source, score = mock_annotation_service.add_annotation_history.call_args[0]
|
query="hi", top_k=1, score_threshold=1, filter={"group_id": ["app-1"]}
|
||||||
assert from_source == "api"
|
)
|
||||||
assert score == 0.8
|
sqlite_session.refresh(annotation)
|
||||||
|
assert annotation.hit_count == 1
|
||||||
|
history = sqlite_session.scalar(select(AppAnnotationHitHistory))
|
||||||
|
assert history is not None
|
||||||
|
assert history.annotation_id == annotation.id
|
||||||
|
assert history.app_id == "app-1"
|
||||||
|
assert history.message_id == "msg-1"
|
||||||
|
assert history.account_id == "user-1"
|
||||||
|
assert history.source == ConversationFromSource.API
|
||||||
|
assert history.score == 0.8
|
||||||
|
|
||||||
def test_query_returns_annotation_and_records_history_for_console(self):
|
def test_query_returns_annotation_and_persists_history_for_console(self, sqlite_session: Session):
|
||||||
feature = AnnotationReplyFeature()
|
binding = _persist_binding(sqlite_session)
|
||||||
annotation_setting = SimpleNamespace(
|
_persist_setting(sqlite_session, collection_binding_id=binding.id)
|
||||||
score_threshold=0.5,
|
annotation = _persist_annotation(sqlite_session)
|
||||||
collection_binding_detail=SimpleNamespace(provider_name="prov", model_name="model"),
|
document = SimpleNamespace(metadata={"annotation_id": annotation.id, "score": 0.6})
|
||||||
)
|
|
||||||
dataset_binding = SimpleNamespace(id="binding-1")
|
|
||||||
annotation = SimpleNamespace(
|
|
||||||
id="ann-1",
|
|
||||||
question_text="question",
|
|
||||||
content="content",
|
|
||||||
account_id="acct-1",
|
|
||||||
account=None,
|
|
||||||
)
|
|
||||||
document = SimpleNamespace(metadata={"annotation_id": "ann-1", "score": 0.6})
|
|
||||||
vector_instance = Mock()
|
vector_instance = Mock()
|
||||||
vector_instance.search_by_vector.return_value = [document]
|
vector_instance.search_by_vector.return_value = [document]
|
||||||
|
|
||||||
with (
|
with patch("core.app.features.annotation_reply.annotation_reply.Vector", return_value=vector_instance):
|
||||||
patch("core.app.features.annotation_reply.annotation_reply.db") as mock_db,
|
result = AnnotationReplyFeature().query(
|
||||||
patch(
|
|
||||||
"core.app.features.annotation_reply.annotation_reply.DatasetCollectionBindingService"
|
|
||||||
) as mock_binding_service,
|
|
||||||
patch("core.app.features.annotation_reply.annotation_reply.Vector") as mock_vector,
|
|
||||||
patch(
|
|
||||||
"core.app.features.annotation_reply.annotation_reply.AppAnnotationService"
|
|
||||||
) as mock_annotation_service,
|
|
||||||
):
|
|
||||||
mock_db.session.scalar.return_value = annotation_setting
|
|
||||||
mock_binding_service.get_dataset_collection_binding.return_value = dataset_binding
|
|
||||||
mock_vector.return_value = vector_instance
|
|
||||||
mock_annotation_service.get_annotation_by_id.return_value = annotation
|
|
||||||
|
|
||||||
result = feature.query(
|
|
||||||
app_record=SimpleNamespace(id="app-1", tenant_id="tenant-1"),
|
app_record=SimpleNamespace(id="app-1", tenant_id="tenant-1"),
|
||||||
message=SimpleNamespace(id="msg-1"),
|
message=SimpleNamespace(id="msg-1"),
|
||||||
query="hi",
|
query="hi",
|
||||||
user_id="user-1",
|
user_id="user-1",
|
||||||
invoke_from=InvokeFrom.EXPLORE,
|
invoke_from=InvokeFrom.EXPLORE,
|
||||||
|
session=sqlite_session,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert result == annotation
|
assert result is annotation
|
||||||
_, _, _, _, _, _, _, from_source, _ = mock_annotation_service.add_annotation_history.call_args[0]
|
history = sqlite_session.scalar(select(AppAnnotationHitHistory))
|
||||||
assert from_source == "console"
|
assert history is not None
|
||||||
|
assert history.source == ConversationFromSource.CONSOLE
|
||||||
|
|
||||||
def test_query_logs_and_returns_none_on_exception(self, caplog: pytest.LogCaptureFixture):
|
def test_query_logs_and_returns_none_on_exception(self, sqlite_session: Session, caplog: pytest.LogCaptureFixture):
|
||||||
feature = AnnotationReplyFeature()
|
binding = _persist_binding(sqlite_session)
|
||||||
annotation_setting = SimpleNamespace(
|
_persist_setting(sqlite_session, collection_binding_id=binding.id)
|
||||||
score_threshold=None,
|
vector_instance = Mock()
|
||||||
collection_binding_detail=SimpleNamespace(provider_name="prov", model_name="model"),
|
vector_instance.search_by_vector.side_effect = RuntimeError("boom")
|
||||||
)
|
|
||||||
|
|
||||||
with (
|
with (
|
||||||
patch("core.app.features.annotation_reply.annotation_reply.db") as mock_db,
|
|
||||||
patch(
|
patch(
|
||||||
"core.app.features.annotation_reply.annotation_reply.DatasetCollectionBindingService"
|
"core.app.features.annotation_reply.annotation_reply.Vector",
|
||||||
) as mock_binding_service,
|
return_value=vector_instance,
|
||||||
patch("core.app.features.annotation_reply.annotation_reply.Vector") as mock_vector,
|
),
|
||||||
|
caplog.at_level(logging.WARNING),
|
||||||
):
|
):
|
||||||
mock_db.session.scalar.return_value = annotation_setting
|
result = AnnotationReplyFeature().query(
|
||||||
mock_binding_service.get_dataset_collection_binding.return_value = SimpleNamespace(id="binding-1")
|
app_record=SimpleNamespace(id="app-1", tenant_id="tenant-1"),
|
||||||
mock_vector.return_value.search_by_vector.side_effect = RuntimeError("boom")
|
message=SimpleNamespace(id="msg-1"),
|
||||||
|
query="hi",
|
||||||
with caplog.at_level(logging.WARNING):
|
user_id="user-1",
|
||||||
result = feature.query(
|
invoke_from=InvokeFrom.SERVICE_API,
|
||||||
app_record=SimpleNamespace(id="app-1", tenant_id="tenant-1"),
|
session=sqlite_session,
|
||||||
message=SimpleNamespace(id="msg-1"),
|
)
|
||||||
query="hi",
|
|
||||||
user_id="user-1",
|
|
||||||
invoke_from=InvokeFrom.SERVICE_API,
|
|
||||||
)
|
|
||||||
|
|
||||||
assert result is None
|
assert result is None
|
||||||
assert "Query annotation failed" in caplog.text
|
assert "Query annotation failed" in caplog.text
|
||||||
|
assert sqlite_session.scalar(select(AppAnnotationHitHistory)) is None
|
||||||
|
assert sqlite_session.is_active
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user