mirror of
https://github.com/langgenius/dify.git
synced 2026-08-28 19:38:30 +08:00
refactor(models): dep-inject Session on @property-dead accessors
Follow-up to #40370. Two @property methods on TypeBase-backed models still reached for ``db.session`` internally, which forces callers to remain coupled to the Flask app-context session even in code paths that already thread a Session explicitly: - ``SavedMessage.message`` - ``WorkflowTriggerLog.created_by_account`` / ``created_by_end_user`` None of these are called from any runtime code path today (grep confirms zero external callers outside their own definition files); each drops the ``@property`` decorator, accepts ``session: Session`` as an explicit argument, and forwards the query against it. Matches the pattern established by ``Dataset.get_created_by_account`` in the referenced PR. Also drops the now-unused ``from .engine import db`` in both files. Tests: adds ``tests/unit_tests/models/test_web_models.py`` and ``tests/unit_tests/models/test_trigger_models.py`` verifying the accessors forward to the supplied session and respect role-based dispatch. Refs: #40372
This commit is contained in:
parent
dfac3e524e
commit
e6509680e7
@ -8,7 +8,7 @@ from uuid import uuid4
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import DateTime, Index, Integer, String, UniqueConstraint, func
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
from sqlalchemy.orm import Mapped, Session, mapped_column
|
||||
|
||||
from core.plugin.entities.plugin_daemon import CredentialType
|
||||
from core.trigger.entities.api_entities import TriggerProviderSubscriptionApiEntity
|
||||
@ -18,7 +18,6 @@ from libs.datetime_utils import naive_utc_now
|
||||
from libs.uuid_utils import uuidv7
|
||||
|
||||
from .base import TypeBase
|
||||
from .engine import db
|
||||
from .enums import AppTriggerStatus, AppTriggerType, CreatorUserRole, PermissionEnum, WorkflowTriggerStatus
|
||||
from .model import Account
|
||||
from .types import EnumText, LongText, StringUUID
|
||||
@ -291,17 +290,15 @@ class WorkflowTriggerLog(TypeBase):
|
||||
triggered_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
finished_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
|
||||
@property
|
||||
def created_by_account(self):
|
||||
def created_by_account(self, session: Session) -> Account | None:
|
||||
created_by_role = CreatorUserRole(self.created_by_role)
|
||||
return db.session.get(Account, self.created_by) if created_by_role == CreatorUserRole.ACCOUNT else None
|
||||
return session.get(Account, self.created_by) if created_by_role == CreatorUserRole.ACCOUNT else None
|
||||
|
||||
@property
|
||||
def created_by_end_user(self):
|
||||
def created_by_end_user(self, session: Session):
|
||||
from .model import EndUser
|
||||
|
||||
created_by_role = CreatorUserRole(self.created_by_role)
|
||||
return db.session.get(EndUser, self.created_by) if created_by_role == CreatorUserRole.END_USER else None
|
||||
return session.get(EndUser, self.created_by) if created_by_role == CreatorUserRole.END_USER else None
|
||||
|
||||
def to_dict(self) -> WorkflowTriggerLogDict:
|
||||
"""Convert to dictionary for API responses"""
|
||||
|
||||
@ -3,10 +3,9 @@ from uuid import uuid4
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import DateTime, func, select
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
from sqlalchemy.orm import Mapped, Session, mapped_column
|
||||
|
||||
from .base import TypeBase
|
||||
from .engine import db
|
||||
from .enums import CreatorUserRole
|
||||
from .model import Message
|
||||
from .types import EnumText, StringUUID
|
||||
@ -36,9 +35,8 @@ class SavedMessage(TypeBase):
|
||||
init=False,
|
||||
)
|
||||
|
||||
@property
|
||||
def message(self):
|
||||
return db.session.scalar(select(Message).where(Message.id == self.message_id))
|
||||
def message(self, session: Session) -> Message | None:
|
||||
return session.scalar(select(Message).where(Message.id == self.message_id))
|
||||
|
||||
|
||||
class PinnedConversation(TypeBase):
|
||||
|
||||
70
api/tests/unit_tests/models/test_trigger_models.py
Normal file
70
api/tests/unit_tests/models/test_trigger_models.py
Normal file
@ -0,0 +1,70 @@
|
||||
"""Regression coverage for ``models.trigger.WorkflowTriggerLog`` account accessors.
|
||||
|
||||
Ensures the ``@property``→session-parameter refactor preserves the role-based dispatch:
|
||||
``created_by_account`` looks up an Account only when role is ACCOUNT; ``created_by_end_user``
|
||||
looks up an EndUser only when role is END_USER.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from models.enums import CreatorUserRole, WorkflowTriggerStatus
|
||||
from models.trigger import WorkflowTriggerLog
|
||||
|
||||
|
||||
def _log(role: CreatorUserRole, created_by: str = "00000000-0000-0000-0000-000000000abc") -> WorkflowTriggerLog:
|
||||
"""Construct a WorkflowTriggerLog without touching the database."""
|
||||
return WorkflowTriggerLog(
|
||||
tenant_id="00000000-0000-0000-0000-000000000001",
|
||||
app_id="00000000-0000-0000-0000-000000000002",
|
||||
workflow_id="00000000-0000-0000-0000-000000000003",
|
||||
workflow_run_id=None,
|
||||
root_node_id=None,
|
||||
trigger_metadata="{}",
|
||||
trigger_type="manual",
|
||||
trigger_data="{}",
|
||||
inputs="{}",
|
||||
status=WorkflowTriggerStatus.SUCCEEDED,
|
||||
queue_name="default",
|
||||
created_by_role=role,
|
||||
created_by=created_by,
|
||||
)
|
||||
|
||||
|
||||
class TestCreatedByAccount:
|
||||
def test_returns_account_lookup_when_role_is_account(self) -> None:
|
||||
log = _log(CreatorUserRole.ACCOUNT)
|
||||
session = MagicMock()
|
||||
sentinel = object()
|
||||
session.get.return_value = sentinel
|
||||
|
||||
result = log.created_by_account(session=session)
|
||||
|
||||
assert result is sentinel
|
||||
session.get.assert_called_once()
|
||||
|
||||
def test_returns_none_when_role_is_end_user(self) -> None:
|
||||
log = _log(CreatorUserRole.END_USER)
|
||||
session = MagicMock()
|
||||
|
||||
assert log.created_by_account(session=session) is None
|
||||
session.get.assert_not_called()
|
||||
|
||||
|
||||
class TestCreatedByEndUser:
|
||||
def test_returns_end_user_lookup_when_role_is_end_user(self) -> None:
|
||||
log = _log(CreatorUserRole.END_USER)
|
||||
session = MagicMock()
|
||||
sentinel = object()
|
||||
session.get.return_value = sentinel
|
||||
|
||||
result = log.created_by_end_user(session=session)
|
||||
|
||||
assert result is sentinel
|
||||
session.get.assert_called_once()
|
||||
|
||||
def test_returns_none_when_role_is_account(self) -> None:
|
||||
log = _log(CreatorUserRole.ACCOUNT)
|
||||
session = MagicMock()
|
||||
|
||||
assert log.created_by_end_user(session=session) is None
|
||||
session.get.assert_not_called()
|
||||
45
api/tests/unit_tests/models/test_web_models.py
Normal file
45
api/tests/unit_tests/models/test_web_models.py
Normal file
@ -0,0 +1,45 @@
|
||||
"""Regression coverage for ``models.web.SavedMessage.message`` accessor.
|
||||
|
||||
Ensures the property→method refactor (drop of ``db.session`` in favor of a caller-provided
|
||||
``Session``) preserves query intent: the accessor forwards ``self.message_id`` to the
|
||||
supplied session and returns whatever the session returns.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from models.enums import CreatorUserRole
|
||||
from models.web import SavedMessage
|
||||
|
||||
|
||||
def _saved_message() -> SavedMessage:
|
||||
"""Construct a SavedMessage without touching the database."""
|
||||
return SavedMessage(
|
||||
app_id="00000000-0000-0000-0000-000000000001",
|
||||
message_id="00000000-0000-0000-0000-000000000002",
|
||||
created_by_role=CreatorUserRole.END_USER,
|
||||
created_by="00000000-0000-0000-0000-000000000003",
|
||||
)
|
||||
|
||||
|
||||
def test_message_forwards_query_to_supplied_session() -> None:
|
||||
saved = _saved_message()
|
||||
session = MagicMock()
|
||||
sentinel = object()
|
||||
session.scalar.return_value = sentinel
|
||||
|
||||
result = saved.message(session=session)
|
||||
|
||||
assert result is sentinel
|
||||
# scalar() is called with a Select; asserting the argument is a Select-shaped
|
||||
# object is enough to prove the accessor still issues a lookup by message_id.
|
||||
session.scalar.assert_called_once()
|
||||
args, _ = session.scalar.call_args
|
||||
assert args, "expected the Select statement to be passed as a positional argument"
|
||||
|
||||
|
||||
def test_message_returns_none_when_session_yields_none() -> None:
|
||||
saved = _saved_message()
|
||||
session = MagicMock()
|
||||
session.scalar.return_value = None
|
||||
|
||||
assert saved.message(session=session) is None
|
||||
Loading…
Reference in New Issue
Block a user