mirror of
https://github.com/langgenius/dify.git
synced 2026-07-21 02:28:30 +08:00
fix(api): harden trial apps access for recommended apps (#38970)
This commit is contained in:
parent
b3386192cc
commit
edca982b1e
@ -18,6 +18,7 @@ from controllers.console.app.error import AppNotFoundError
|
||||
from extensions.ext_database import db
|
||||
from libs.login import current_account_with_tenant
|
||||
from models import App, AppMode, TrialApp
|
||||
from services.recommended_app_service import RecommendedAppService
|
||||
|
||||
__all__ = ["get_app_model", "get_app_model_with_trial", "with_session"]
|
||||
|
||||
@ -156,6 +157,8 @@ def get_app_model_with_trial[**P, R](
|
||||
*,
|
||||
mode: AppMode | list[AppMode] | None = None,
|
||||
) -> Callable[P, R] | Callable[[Callable[P, R]], Callable[P, R]]:
|
||||
"""Inject an app registered for trial or available from the recommended catalog."""
|
||||
|
||||
def decorator(view_func: Callable[P, R]) -> Callable[P, R]:
|
||||
@wraps(view_func)
|
||||
def decorated_view(*args: P.args, **kwargs: P.kwargs) -> R:
|
||||
@ -168,6 +171,8 @@ def get_app_model_with_trial[**P, R](
|
||||
del kwargs["app_id"]
|
||||
|
||||
app_model = _load_app_model_with_trial(app_id)
|
||||
if app_model is None:
|
||||
app_model = RecommendedAppService.get_app(app_id, session=db.session())
|
||||
|
||||
if not app_model:
|
||||
raise AppNotFoundError()
|
||||
|
||||
@ -4,12 +4,23 @@ from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from configs import dify_config
|
||||
from models.model import AccountTrialAppRecord, TrialApp
|
||||
from models.model import AccountTrialAppRecord, App, TrialApp
|
||||
from services.feature_service import FeatureService
|
||||
from services.recommend_app.recommend_app_factory import RecommendAppRetrievalFactory
|
||||
|
||||
|
||||
class RecommendedAppService:
|
||||
@classmethod
|
||||
def get_app(cls, app_id: str, *, session: Session) -> App | None:
|
||||
"""Return a normal app only when it belongs to the recommended catalog."""
|
||||
mode = dify_config.HOSTED_FETCH_APP_TEMPLATES_MODE
|
||||
retrieval_instance = RecommendAppRetrievalFactory.get_recommend_app_factory(mode)()
|
||||
recommended_app_detail = retrieval_instance.get_recommend_app_detail(app_id, session=session)
|
||||
if recommended_app_detail is None:
|
||||
return None
|
||||
|
||||
return session.scalar(select(App).where(App.id == app_id, App.status == "normal").limit(1))
|
||||
|
||||
@classmethod
|
||||
def get_recommended_apps_and_categories(cls, language: str, *, session: Session):
|
||||
"""
|
||||
|
||||
@ -1,9 +1,11 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import Select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from controllers.common.session import with_session
|
||||
from controllers.console.app import wraps as wraps_module
|
||||
@ -64,7 +66,11 @@ def test_get_app_model_with_trial_requires_trial_app_registration(monkeypatch: p
|
||||
)
|
||||
return None if has_trial_app_join else app_model
|
||||
|
||||
monkeypatch.setattr(wraps_module.db, "session", SimpleNamespace(scalar=scalar))
|
||||
scoped_session = MagicMock()
|
||||
scoped_session.scalar.side_effect = scalar
|
||||
recommended_get_app = MagicMock(return_value=None)
|
||||
monkeypatch.setattr(wraps_module.db, "session", scoped_session)
|
||||
monkeypatch.setattr(wraps_module.RecommendedAppService, "get_app", recommended_get_app)
|
||||
|
||||
@wraps_module.get_app_model_with_trial
|
||||
def handler(app_model):
|
||||
@ -73,6 +79,43 @@ def test_get_app_model_with_trial_requires_trial_app_registration(monkeypatch: p
|
||||
with pytest.raises(AppNotFoundError):
|
||||
handler(app_id="app-1")
|
||||
|
||||
recommended_get_app.assert_called_once_with("app-1", session=scoped_session.return_value)
|
||||
|
||||
|
||||
def test_get_app_model_with_trial_falls_back_to_recommended_app(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
app_model = SimpleNamespace(id="app-1", mode=AppMode.CHAT.value, status="normal", tenant_id="t1")
|
||||
session = MagicMock(spec=Session)
|
||||
scoped_session = MagicMock(return_value=session)
|
||||
trial_app_loader = MagicMock(return_value=None)
|
||||
recommended_get_app = MagicMock(return_value=app_model)
|
||||
monkeypatch.setattr(wraps_module.db, "session", scoped_session)
|
||||
monkeypatch.setattr(wraps_module, "_load_app_model_with_trial", trial_app_loader)
|
||||
monkeypatch.setattr(wraps_module.RecommendedAppService, "get_app", recommended_get_app)
|
||||
|
||||
@wraps_module.get_app_model_with_trial
|
||||
def handler(app_model):
|
||||
return app_model.id
|
||||
|
||||
assert handler(app_id="app-1") == "app-1"
|
||||
trial_app_loader.assert_called_once_with("app-1")
|
||||
recommended_get_app.assert_called_once_with("app-1", session=session)
|
||||
|
||||
|
||||
def test_get_app_model_with_trial_prefers_trial_registration(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
app_model = SimpleNamespace(id="app-1", mode=AppMode.CHAT.value, status="normal", tenant_id="t1")
|
||||
trial_app_loader = MagicMock(return_value=app_model)
|
||||
recommended_get_app = MagicMock()
|
||||
monkeypatch.setattr(wraps_module, "_load_app_model_with_trial", trial_app_loader)
|
||||
monkeypatch.setattr(wraps_module.RecommendedAppService, "get_app", recommended_get_app)
|
||||
|
||||
@wraps_module.get_app_model_with_trial
|
||||
def handler(app_model):
|
||||
return app_model.id
|
||||
|
||||
assert handler(app_id="app-1") == "app-1"
|
||||
trial_app_loader.assert_called_once_with("app-1")
|
||||
recommended_get_app.assert_not_called()
|
||||
|
||||
|
||||
def test_get_app_model_requires_app_id() -> None:
|
||||
@wraps_module.get_app_model
|
||||
|
||||
@ -10,14 +10,14 @@ import pytest
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from models.model import AccountTrialAppRecord, TrialApp
|
||||
from models.model import AccountTrialAppRecord, App, AppMode, TrialApp
|
||||
from services import recommended_app_service as service_module
|
||||
from services.feature_service import SystemFeatureModel
|
||||
from services.recommended_app_service import RecommendedAppService
|
||||
|
||||
pytestmark = pytest.mark.parametrize(
|
||||
"sqlite_session",
|
||||
[(TrialApp, AccountTrialAppRecord)],
|
||||
[(TrialApp, AccountTrialAppRecord, App)],
|
||||
indirect=True,
|
||||
)
|
||||
|
||||
@ -110,6 +110,37 @@ def _mock_factory_for_apps(
|
||||
return retrieval_instance, builtin_instance
|
||||
|
||||
|
||||
def _mock_factory_for_app_detail(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
*,
|
||||
result: RecommendedAppPayload | None,
|
||||
) -> MagicMock:
|
||||
retrieval_instance = MagicMock()
|
||||
retrieval_instance.get_recommend_app_detail.return_value = result
|
||||
retrieval_factory = MagicMock(return_value=retrieval_instance)
|
||||
monkeypatch.setattr(service_module.dify_config, "HOSTED_FETCH_APP_TEMPLATES_MODE", "remote", raising=False)
|
||||
monkeypatch.setattr(
|
||||
service_module.RecommendAppRetrievalFactory,
|
||||
"get_recommend_app_factory",
|
||||
MagicMock(return_value=retrieval_factory),
|
||||
)
|
||||
return retrieval_instance
|
||||
|
||||
|
||||
def _persist_app(session: Session, *, name: str) -> App:
|
||||
app = App(
|
||||
tenant_id=str(uuid.uuid4()),
|
||||
name=name,
|
||||
mode=AppMode.CHAT,
|
||||
enable_site=True,
|
||||
enable_api=True,
|
||||
)
|
||||
app.id = str(uuid.uuid4())
|
||||
session.add(app)
|
||||
session.commit()
|
||||
return app
|
||||
|
||||
|
||||
# ── Pure logic tests: get_recommended_apps_and_categories ──────────────
|
||||
|
||||
|
||||
@ -219,6 +250,39 @@ class TestRecommendedAppServiceGetApps:
|
||||
mock_factory_class.get_recommend_app_factory.assert_called_with(mode)
|
||||
|
||||
|
||||
# ── Database-backed tests: get_app ─────────────────────────────────────
|
||||
|
||||
|
||||
class TestRecommendedAppServiceGetApp:
|
||||
def test_returns_normal_recommended_app(self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
|
||||
app = _persist_app(sqlite_session, name="Recommended App")
|
||||
|
||||
retrieval_instance = _mock_factory_for_app_detail(
|
||||
monkeypatch,
|
||||
result=RecommendedAppPayload(id=app.id),
|
||||
)
|
||||
feature_lookup = MagicMock(side_effect=AssertionError("get_app must not inspect trial features"))
|
||||
monkeypatch.setattr(service_module.FeatureService, "get_system_features", feature_lookup)
|
||||
|
||||
result = RecommendedAppService.get_app(app.id, session=sqlite_session)
|
||||
|
||||
assert result is app
|
||||
retrieval_instance.get_recommend_app_detail.assert_called_once_with(app.id, session=sqlite_session)
|
||||
feature_lookup.assert_not_called()
|
||||
|
||||
def test_returns_none_when_app_is_not_recommended(
|
||||
self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
app = _persist_app(sqlite_session, name="Private App")
|
||||
|
||||
retrieval_instance = _mock_factory_for_app_detail(monkeypatch, result=None)
|
||||
|
||||
result = RecommendedAppService.get_app(app.id, session=sqlite_session)
|
||||
|
||||
assert result is None
|
||||
retrieval_instance.get_recommend_app_detail.assert_called_once_with(app.id, session=sqlite_session)
|
||||
|
||||
|
||||
# ── Pure logic tests: get_recommend_app_detail ─────────────────────────
|
||||
|
||||
|
||||
|
||||
Loading…
Reference in New Issue
Block a user