refactor(api): pass session explicitly through recommended app service chain

This commit is contained in:
林玮 (Jade Lin)
2026-07-15 15:50:05 +08:00
committed by GareArc
parent 5c0a0759f0
commit ab2ff4eddf
12 changed files with 41 additions and 87 deletions
@@ -106,7 +106,7 @@ class RecommendedAppListApi(Resource):
language_prefix = _resolve_language(args.language, current_user)
return RecommendedAppListResponse.model_validate(
RecommendedAppService.get_recommended_apps_and_categories(db.session, language_prefix),
RecommendedAppService.get_recommended_apps_and_categories(language_prefix, session=db.session()),
from_attributes=True,
).model_dump(mode="json")
+5 -5
View File
@@ -293,11 +293,11 @@ class MemberInviteEmailApi(Resource):
)
except SeatsLimitExceededError:
invitation_results.append(
MemberInviteFailedResponse(
status="failed",
email=invitee_email,
message="Licensed seats limit exceeded.",
)
{
"status": "failed",
"email": invitee_email,
"message": "Licensed seats limit exceeded.",
}
)
except Exception as e:
invitation_results.append({"status": "failed", "email": invitee_email, "message": str(e)})
+2 -2
View File
@@ -19,7 +19,7 @@ import time
from collections.abc import Iterator, Mapping, Sequence
from contextlib import contextmanager
from mimetypes import guess_type
from typing import Literal, Protocol
from typing import Literal, Protocol, cast
import zstandard
from pydantic import BaseModel, TypeAdapter, ValidationError
@@ -461,7 +461,7 @@ class PluginService:
plugin_id=plugin_id,
version=manifest.latest_version,
unique_identifier=manifest.latest_package_identifier,
status=manifest.status,
status=cast(Literal["active", "deleted"], manifest.status),
deprecated_reason=manifest.deprecated_reason,
alternative_plugin_id=manifest.alternative_plugin_id,
)
@@ -4,6 +4,7 @@ from pathlib import Path
from typing import Any, override
from flask import current_app
from sqlalchemy.orm import Session
from services.recommend_app.database.database_retrieval import DatabaseRecommendAppRetrieval
from services.recommend_app.recommend_app_base import RecommendAppRetrievalBase
@@ -32,7 +33,8 @@ class BuildInRecommendAppRetrieval(RecommendAppRetrievalBase):
return result
@override
def get_recommend_app_detail(self, app_id: str):
def get_recommend_app_detail(self, app_id: str, *, session: Session | None = None):
del session
result = self.fetch_recommended_app_detail_from_builtin(app_id)
return result
@@ -1,6 +1,7 @@
from typing import Any, NotRequired, TypedDict, override
from sqlalchemy import select
from sqlalchemy.orm import Session
from constants.languages import languages
from extensions.ext_database import db
@@ -55,8 +56,10 @@ class DatabaseRecommendAppRetrieval(RecommendAppRetrievalBase):
return result
@override
def get_recommend_app_detail(self, app_id: str) -> RecommendedAppDetailDict | None:
result = self.fetch_recommended_app_detail_from_db(app_id)
def get_recommend_app_detail(
self, app_id: str, *, session: Session | None = None
) -> RecommendedAppDetailDict | None:
result = self.fetch_recommended_app_detail_from_db(app_id, session=session)
return result
@override
@@ -146,14 +149,17 @@ class DatabaseRecommendAppRetrieval(RecommendAppRetrievalBase):
)
@classmethod
def fetch_recommended_app_detail_from_db(cls, app_id: str) -> RecommendedAppDetailDict | None:
def fetch_recommended_app_detail_from_db(
cls, app_id: str, *, session: Session | None = None
) -> RecommendedAppDetailDict | None:
"""
Fetch recommended app detail from db.
:param app_id: App ID
:return:
"""
# is in public recommended list
recommended_app = db.session.scalar(
query_session = session if session is not None else db.session
recommended_app = query_session.scalar(
select(RecommendedApp).where(RecommendedApp.is_listed == True, RecommendedApp.app_id == app_id).limit(1)
)
@@ -161,7 +167,7 @@ class DatabaseRecommendAppRetrieval(RecommendAppRetrievalBase):
return None
# get app detail
app_model = db.session.get(App, app_id)
app_model = query_session.get(App, app_id)
if not app_model or not app_model.is_public:
return None
@@ -1,5 +1,7 @@
from typing import Any, Protocol
from sqlalchemy.orm import Session
class RecommendAppRetrievalBase(Protocol):
"""Interface for recommend app retrieval."""
@@ -8,6 +10,6 @@ class RecommendAppRetrievalBase(Protocol):
def get_learn_dify_apps(self, language: str) -> Any: ...
def get_recommend_app_detail(self, app_id: str) -> Any: ...
def get_recommend_app_detail(self, app_id: str, *, session: Session | None = None) -> Any: ...
def get_type(self) -> str: ...
@@ -3,6 +3,7 @@ from typing import Any, override
import httpx
from flask import has_request_context, request
from sqlalchemy.orm import Session
from configs import dify_config
from services.recommend_app.buildin.buildin_retrieval import BuildInRecommendAppRetrieval
@@ -33,7 +34,8 @@ class RemoteRecommendAppRetrieval(RecommendAppRetrievalBase):
"""
@override
def get_recommend_app_detail(self, app_id: str):
def get_recommend_app_detail(self, app_id: str, *, session: Session | None = None):
del session
try:
result = self.fetch_recommended_app_detail_from_dify_official(app_id)
except Exception as e:
+1 -1
View File
@@ -99,6 +99,6 @@ class RecommendedAppService:
session.commit()
@staticmethod
def _can_trial_app(session: scoped_session, app_id: str) -> bool:
def _can_trial_app(session: Session | scoped_session, app_id: str) -> bool:
trial_app_model = session.scalar(select(TrialApp).where(TrialApp.app_id == app_id).limit(1))
return trial_app_model is not None
@@ -9,6 +9,7 @@ import pytest
from sqlalchemy import select
from sqlalchemy.orm import Session
from extensions.ext_database import db
from models.model import AccountTrialAppRecord, App, AppMode, TrialApp
from services import recommended_app_service as service_module
from services.recommended_app_service import RecommendedAppService
@@ -154,7 +155,7 @@ class TestRecommendedAppServiceGetApps:
mock_factory = MagicMock(return_value=mock_instance)
mock_factory_class.get_recommend_app_factory.return_value = mock_factory
result = RecommendedAppService.get_recommended_apps_and_categories(db.session, "en-US")
result = RecommendedAppService.get_recommended_apps_and_categories("en-US", session=db.session())
assert result == expected
assert len(result["recommended_apps"]) == 2
@@ -179,7 +180,7 @@ class TestRecommendedAppServiceGetApps:
mock_builtin_instance.fetch_recommended_apps_from_builtin.return_value = builtin_response
mock_factory_class.get_buildin_recommend_app_retrieval.return_value = mock_builtin_instance
result = RecommendedAppService.get_recommended_apps_and_categories(db.session, "zh-CN")
result = RecommendedAppService.get_recommended_apps_and_categories("zh-CN", session=db.session())
assert result == builtin_response
assert result["recommended_apps"][0]["id"] == "builtin-1"
@@ -200,7 +201,7 @@ class TestRecommendedAppServiceGetApps:
mock_builtin_instance.fetch_recommended_apps_from_builtin.return_value = builtin_response
mock_factory_class.get_buildin_recommend_app_retrieval.return_value = mock_builtin_instance
result = RecommendedAppService.get_recommended_apps_and_categories(db.session, "en-US")
result = RecommendedAppService.get_recommended_apps_and_categories("en-US", session=db.session())
assert result == builtin_response
mock_builtin_instance.fetch_recommended_apps_from_builtin.assert_called_once()
@@ -218,7 +219,7 @@ class TestRecommendedAppServiceGetApps:
mock_instance.get_recommended_apps_and_categories.return_value = lang_response
mock_factory_class.get_recommend_app_factory.return_value = MagicMock(return_value=mock_instance)
result = RecommendedAppService.get_recommended_apps_and_categories(db.session, language)
result = RecommendedAppService.get_recommended_apps_and_categories(language, session=db.session())
assert result["recommended_apps"][0]["id"] == f"app-{language}"
mock_instance.get_recommended_apps_and_categories.assert_called_with(language)
@@ -233,7 +234,7 @@ class TestRecommendedAppServiceGetApps:
mock_instance.get_recommended_apps_and_categories.return_value = response
mock_factory_class.get_recommend_app_factory.return_value = MagicMock(return_value=mock_instance)
RecommendedAppService.get_recommended_apps_and_categories(db.session, "en-US")
RecommendedAppService.get_recommended_apps_and_categories("en-US", session=db.session())
mock_factory_class.get_recommend_app_factory.assert_called_with(mode)
@@ -402,7 +403,7 @@ class TestRecommendedAppServiceTrialFeatures:
MagicMock(return_value=SimpleNamespace(enable_trial_app=False)),
)
result = RecommendedAppService.get_recommended_apps_and_categories(db.session, "en-US")
result = RecommendedAppService.get_recommended_apps_and_categories("en-US", session=db.session())
assert result == expected
retrieval_instance.get_recommended_apps_and_categories.assert_called_once_with("en-US")
@@ -433,7 +434,7 @@ class TestRecommendedAppServiceTrialFeatures:
MagicMock(return_value=SimpleNamespace(enable_trial_app=True)),
)
result = RecommendedAppService.get_recommended_apps_and_categories(db.session, "ja-JP")
result = RecommendedAppService.get_recommended_apps_and_categories("ja-JP", session=db.session())
builtin_instance.fetch_recommended_apps_from_builtin.assert_called_once_with("en-US")
assert result["recommended_apps"][0]["can_trial"] is True
@@ -32,7 +32,7 @@ class TestRecommendedAppListApi:
):
result = method(api, make_account("fr-FR"))
service_mock.assert_called_once_with(ANY, "en-US")
service_mock.assert_called_once_with("en-US", session=ANY)
assert result == result_data
def test_get_fallback_to_user_language(self, app: Flask):
@@ -51,7 +51,7 @@ class TestRecommendedAppListApi:
):
result = method(api, make_account("fr-FR"))
service_mock.assert_called_once_with(ANY, "fr-FR")
service_mock.assert_called_once_with("fr-FR", session=ANY)
assert result == result_data
def test_get_fallback_to_default_language(self, app: Flask):
@@ -70,7 +70,7 @@ class TestRecommendedAppListApi:
):
result = method(api, make_account(None))
service_mock.assert_called_once_with(ANY, module.languages[0])
service_mock.assert_called_once_with(module.languages[0], session=ANY)
assert result == result_data
@@ -14,7 +14,7 @@ from core.app.layers.pause_state_persist_layer import (
_WorkflowGenerateEntityWrapper,
)
from core.workflow.system_variables import SystemVariableKey
from graphon.entities.pause_reason import HitlRequired, SchedulingPause
from graphon.entities.pause_reason import SchedulingPause
from graphon.filters import GraphEventFilterContext, ResponseStreamFilter
from graphon.graph_engine.entities.commands import GraphEngineCommand
from graphon.graph_engine.layers.base import GraphEngineLayerNotInitializedError
@@ -283,64 +283,6 @@ class TestPauseStatePersistenceLayer:
assert isinstance(pause_reasons, list)
def test_on_event_enriches_hitl_pause_reasons_before_persisting(self, monkeypatch: pytest.MonkeyPatch):
session_factory = Mock(name="session_factory")
generate_entity = self._create_generate_entity(workflow_execution_id="run-123")
layer = PauseStatePersistenceLayer(
session_factory=session_factory,
state_owner_user_id="owner-123",
generate_entity=generate_entity,
response_stream_filter=_create_initialized_response_stream_filter(),
)
mock_repo = Mock()
mock_factory = Mock(return_value=mock_repo)
mock_form_repository = Mock(name="form_repository")
enriched_reason = HumanInputRequired(
form_id="form-123",
form_content="Rendered content",
inputs=[],
actions=[],
node_id="node-123",
node_title="Ask for approval",
)
enrich_mock = Mock(return_value=[enriched_reason])
monkeypatch.setattr(DifyAPIRepositoryFactory, "create_api_workflow_run_repository", mock_factory)
monkeypatch.setattr(
pause_layer_module,
"HumanInputFormSubmissionRepository",
Mock(return_value=mock_form_repository),
raising=False,
)
monkeypatch.setattr(
pause_layer_module,
"enrich_graph_pause_reasons",
enrich_mock,
raising=False,
)
graph_runtime_state = MockReadOnlyGraphRuntimeState(
workflow_execution_id="run-123",
)
command_channel = MockCommandChannel()
layer.initialize(graph_runtime_state, command_channel)
raw_reason = HitlRequired(
session_id="session-123",
node_id="node-123",
node_title="Ask for approval",
)
event = GraphRunPausedEvent(reasons=[raw_reason], outputs={})
layer.on_event(event)
enrich_mock.assert_called_once_with(
reasons=[raw_reason],
form_repository=mock_form_repository,
variable_pool=graph_runtime_state.variable_pool,
)
assert mock_repo.create_workflow_pause.call_args.kwargs["pause_reasons"] == [enriched_reason]
def test_on_event_ignores_non_paused_events(self, monkeypatch: pytest.MonkeyPatch):
session_factory = Mock(name="session_factory")
layer = PauseStatePersistenceLayer(
@@ -399,7 +399,6 @@ class TestPluginAppBackwardsInvocation:
route = mocker.patch.object(PluginAppBackwardsInvocation, "invoke_workflow_app", return_value={"ok": True})
result = PluginAppBackwardsInvocation.invoke_app(
MagicMock(),
app_id="app",
user_id="wecom-sender-1",
tenant_id="tenant",