mirror of
https://github.com/langgenius/dify.git
synced 2026-07-21 03:55:25 -04:00
refactor(api): pass session explicitly through recommended app service chain
This commit is contained in:
@@ -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")
|
||||
|
||||
|
||||
@@ -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)})
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
+8
-7
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user