test: use sqlite3 session in test_annotation_reply (#38714)

This commit is contained in:
Asuka Minato
2026-07-14 03:52:52 +00:00
committed by GitHub
parent e044292518
commit 6fa71ef7ab
2 changed files with 150 additions and 132 deletions
@@ -1,12 +1,13 @@
import logging
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.app.entities.app_invoke_entities import InvokeFrom
from core.rag.datasource.vdb.vector_factory import Vector
from core.rag.index_processor.constant.index_type import IndexTechniqueType
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.model import App, AppAnnotationSetting, Message, MessageAnnotation
from services.annotation_service import AppAnnotationService
@@ -17,24 +18,33 @@ logger = logging.getLogger(__name__)
class AnnotationReplyFeature:
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:
"""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
:param app_record: app record
:param message: message
:param query: query
:param user_id: user id
:param invoke_from: invoke from
:return:
"""
if session is None:
session = db.session()
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:
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:
return None
@@ -45,7 +55,7 @@ class AnnotationReplyFeature:
embedding_model_name = collection_binding_detail.model_name
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(
@@ -66,7 +76,7 @@ class AnnotationReplyFeature:
if documents and documents[0].metadata:
annotation_id = documents[0].metadata["annotation_id"]
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 invoke_from in {InvokeFrom.SERVICE_API, InvokeFrom.WEB_APP}:
from_source = ConversationFromSource.API
@@ -84,7 +94,7 @@ class AnnotationReplyFeature:
message.id,
from_source,
score,
session=db.session(),
session=session,
)
return annotation