test: use sqlite3 session in test_annotation_reply (#38714)

This commit is contained in:
Asuka Minato 2026-07-14 12:52:52 +09:00 committed by GitHub
parent e044292518
commit 6fa71ef7ab
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 150 additions and 132 deletions

View File

@ -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

View File

@ -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