mirror of
https://github.com/langgenius/dify.git
synced 2026-08-29 03:45:08 +08:00
refactor(models): dep-inject Session on remaining @property accessors (#40797)
Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
@@ -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
-5
@@ -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):
|
||||
|
||||
@@ -0,0 +1,89 @@
|
||||
"""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.
|
||||
|
||||
Both accessors are exercised against the real ``sqlite_session`` fixture (a genuine
|
||||
SQLAlchemy ``Session`` bound to a pristine full-schema SQLite database) so the assertions
|
||||
cover actual query behaviour rather than a mock's recorded call.
|
||||
"""
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from models.account import Account
|
||||
from models.enums import CreatorUserRole, EndUserType, WorkflowTriggerStatus
|
||||
from models.model import EndUser
|
||||
from models.trigger import WorkflowTriggerLog
|
||||
|
||||
|
||||
def _log(role: CreatorUserRole, created_by: str) -> 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="{}",
|
||||
outputs=None,
|
||||
status=WorkflowTriggerStatus.SUCCEEDED,
|
||||
error=None,
|
||||
queue_name="default",
|
||||
celery_task_id=None,
|
||||
created_by_role=role,
|
||||
created_by=created_by,
|
||||
)
|
||||
|
||||
|
||||
class TestCreatedByAccount:
|
||||
def test_returns_account_lookup_when_role_is_account(self, sqlite_session: Session) -> None:
|
||||
account = Account(name="Test Account", email="test@example.com")
|
||||
sqlite_session.add(account)
|
||||
sqlite_session.flush()
|
||||
log = _log(CreatorUserRole.ACCOUNT, created_by=account.id)
|
||||
|
||||
result = log.created_by_account(session=sqlite_session)
|
||||
|
||||
assert result is not None
|
||||
assert result.id == account.id
|
||||
|
||||
def test_returns_none_when_role_is_end_user(self, sqlite_session: Session) -> None:
|
||||
account = Account(name="Test Account", email="test@example.com")
|
||||
sqlite_session.add(account)
|
||||
sqlite_session.flush()
|
||||
log = _log(CreatorUserRole.END_USER, created_by=account.id)
|
||||
|
||||
assert log.created_by_account(session=sqlite_session) is None
|
||||
|
||||
|
||||
class TestCreatedByEndUser:
|
||||
def test_returns_end_user_lookup_when_role_is_end_user(self, sqlite_session: Session) -> None:
|
||||
end_user = EndUser(
|
||||
tenant_id="00000000-0000-0000-0000-000000000001",
|
||||
type=EndUserType.BROWSER,
|
||||
session_id="session-1",
|
||||
)
|
||||
sqlite_session.add(end_user)
|
||||
sqlite_session.flush()
|
||||
log = _log(CreatorUserRole.END_USER, created_by=end_user.id)
|
||||
|
||||
result = log.created_by_end_user(session=sqlite_session)
|
||||
|
||||
assert result is not None
|
||||
assert result.id == end_user.id
|
||||
|
||||
def test_returns_none_when_role_is_account(self, sqlite_session: Session) -> None:
|
||||
end_user = EndUser(
|
||||
tenant_id="00000000-0000-0000-0000-000000000001",
|
||||
type=EndUserType.BROWSER,
|
||||
session_id="session-1",
|
||||
)
|
||||
sqlite_session.add(end_user)
|
||||
sqlite_session.flush()
|
||||
log = _log(CreatorUserRole.ACCOUNT, created_by=end_user.id)
|
||||
|
||||
assert log.created_by_end_user(session=sqlite_session) is None
|
||||
@@ -0,0 +1,67 @@
|
||||
"""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 the matching :class:`Message` (or ``None`` when absent).
|
||||
|
||||
The accessor is exercised against the real ``sqlite_session`` fixture (a genuine SQLAlchemy
|
||||
``Session`` bound to a pristine full-schema SQLite database) so the assertions cover actual
|
||||
query behaviour rather than a mock's recorded call.
|
||||
"""
|
||||
|
||||
from decimal import Decimal
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from models.enums import CreatorUserRole
|
||||
from models.model import ConversationFromSource, Message
|
||||
from models.web import SavedMessage
|
||||
|
||||
|
||||
def _persist_message(session: Session, *, app_id: str) -> Message:
|
||||
"""Persist a minimal valid Message row and return it."""
|
||||
message = Message(
|
||||
app_id=app_id,
|
||||
conversation_id=str(uuid4()),
|
||||
inputs={},
|
||||
query="hello",
|
||||
message=[{"role": "user", "text": "hello"}],
|
||||
message_unit_price=Decimal(0),
|
||||
message_price_unit=Decimal(0),
|
||||
answer="hi",
|
||||
answer_unit_price=Decimal(0),
|
||||
answer_price_unit=Decimal(0),
|
||||
currency="USD",
|
||||
from_source=ConversationFromSource.API,
|
||||
)
|
||||
session.add(message)
|
||||
session.flush()
|
||||
return message
|
||||
|
||||
|
||||
def _saved_message(*, app_id: str, message_id: str) -> SavedMessage:
|
||||
"""Construct a SavedMessage without touching the database."""
|
||||
return SavedMessage(
|
||||
app_id=app_id,
|
||||
message_id=message_id,
|
||||
created_by_role=CreatorUserRole.END_USER,
|
||||
created_by=str(uuid4()),
|
||||
)
|
||||
|
||||
|
||||
def test_message_returns_persisted_message(sqlite_session: Session) -> None:
|
||||
app_id = str(uuid4())
|
||||
message = _persist_message(sqlite_session, app_id=app_id)
|
||||
saved = _saved_message(app_id=app_id, message_id=message.id)
|
||||
|
||||
result = saved.message(session=sqlite_session)
|
||||
|
||||
assert result is not None
|
||||
assert result.id == message.id
|
||||
|
||||
|
||||
def test_message_returns_none_when_message_missing(sqlite_session: Session) -> None:
|
||||
saved = _saved_message(app_id=str(uuid4()), message_id=str(uuid4()))
|
||||
|
||||
assert saved.message(session=sqlite_session) is None
|
||||
Reference in New Issue
Block a user