refactor(api): migrate core RAG layer to SQLAlchemy 2.0 select() API (#34965)

This commit is contained in:
wdeveloper16
2026-04-11 16:32:20 +00:00
committed by GitHub
parent 50206ae8a7
commit 12814b55d2
8 changed files with 66 additions and 81 deletions
@@ -258,10 +258,10 @@ class TestParentChildIndexProcessor:
session.commit.assert_called_once()
def test_clean_deletes_summaries_when_requested(self, processor: ParentChildIndexProcessor, dataset: Mock) -> None:
segment_query = Mock()
segment_query.filter.return_value.all.return_value = [SimpleNamespace(id="seg-1")]
scalars_result = Mock()
scalars_result.all.return_value = [SimpleNamespace(id="seg-1")]
session = Mock()
session.query.return_value = segment_query
session.scalars.return_value = scalars_result
session_ctx = MagicMock()
session_ctx.__enter__.return_value = session
session_ctx.__exit__.return_value = False
@@ -220,10 +220,10 @@ class TestQAIndexProcessor:
self, processor: QAIndexProcessor, dataset: Mock
) -> None:
mock_segment = SimpleNamespace(id="seg-1")
mock_query = Mock()
mock_query.filter.return_value.all.return_value = [mock_segment]
scalars_result = Mock()
scalars_result.all.return_value = [mock_segment]
mock_session = Mock()
mock_session.query.return_value = mock_query
mock_session.scalars.return_value = scalars_result
session_context = MagicMock()
session_context.__enter__.return_value = mock_session
session_context.__exit__.return_value = False
@@ -8,7 +8,6 @@ import pytest
from flask import Flask, current_app
from graphon.model_runtime.entities.llm_entities import LLMUsage
from graphon.model_runtime.entities.model_entities import ModelFeature
from sqlalchemy import column
from core.app.app_config.entities import (
DatasetEntity,
@@ -4039,21 +4038,9 @@ class TestDatasetRetrievalAdditionalHelpers:
def test_get_available_datasets(self, retrieval: DatasetRetrieval) -> None:
session = Mock()
subquery_query = Mock()
subquery_query.where.return_value = subquery_query
subquery_query.group_by.return_value = subquery_query
subquery_query.having.return_value = subquery_query
subquery_query.subquery.return_value = SimpleNamespace(
c=SimpleNamespace(
dataset_id=column("dataset_id"), available_document_count=column("available_document_count")
)
)
dataset_query = Mock()
dataset_query.outerjoin.return_value = dataset_query
dataset_query.where.return_value = dataset_query
dataset_query.all.return_value = [SimpleNamespace(id="d1"), None, SimpleNamespace(id="d2")]
session.query.side_effect = [subquery_query, dataset_query]
scalars_result = Mock()
scalars_result.all.return_value = [SimpleNamespace(id="d1"), None, SimpleNamespace(id="d2")]
session.scalars.return_value = scalars_result
session_ctx = MagicMock()
session_ctx.__enter__.return_value = session
@@ -4902,9 +4889,6 @@ class TestInternalHooksCoverage:
_scalars(segments),
_scalars(bindings),
]
query = Mock()
query.where.return_value = query
session.query.return_value = query
session_ctx = MagicMock()
session_ctx.__enter__.return_value = session
session_ctx.__exit__.return_value = False
@@ -4919,7 +4903,7 @@ class TestInternalHooksCoverage:
):
retrieval._on_retrieval_end(flask_app=app, documents=docs, message_id="m1", timer={"cost": 1})
query.update.assert_called_once()
session.execute.assert_called_once()
mock_trace.assert_called_once()
def test_retriever_variants(self, retrieval: DatasetRetrieval) -> None: