refactor(api): standardize session factory wiring (#41085)

This commit is contained in:
非法操作 2026-08-24 01:57:17 +00:00 committed by GitHub
parent ed45016475
commit 5018f8aa8a
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
7 changed files with 24 additions and 24 deletions

View File

@ -168,7 +168,7 @@ def build_application_services(
initialization_password: str,
redis: RedisClientWrapper,
) -> ApplicationServices:
installation_state = InstallationStateRepository(client=database_client)
installation_state = InstallationStateRepository(session_factory=database_client)
data_source_api_key_auth_bindings = SQLAlchemyDataSourceApiKeyAuthBindingRepository(session_factory=database_client)
app_definition_repository = AppDefinitionQueryRepository(session_factory=database_client)
feature_gateway = FeatureServiceGateway()
@ -183,7 +183,7 @@ def build_application_services(
database=database_catalog,
builtin=builtin_catalog,
)
workspace_query_repository = WorkspaceQueryRepository(client=database_client)
workspace_query_repository = WorkspaceQueryRepository(session_factory=database_client)
return ApplicationServices(
accounts=AccountServices(
avatar=AccountAvatarService(
@ -248,7 +248,7 @@ def build_application_services(
),
account_activation=AccountActivationService(
tokens=RegisterServiceInvitationTokenStore(),
accounts=SQLAlchemyAccountActivationRepository(database_client),
accounts=SQLAlchemyAccountActivationRepository(session_factory=database_client),
workspace_policy=DeploymentWorkspaceInvitePolicy(),
eligibility=BillingAccountActivationEligibility(
enabled=deployment_edition == DeploymentEdition.CLOUD,
@ -281,18 +281,18 @@ def build_application_services(
),
web_app_runtime=WebAppRuntimeQueryService(
runtime=app_definition_repository,
file_service=FileService(database_client),
file_service=FileService(session_factory=database_client),
workspace_features=feature_gateway.get_workspace_features,
files_url=dify_config.FILES_URL,
),
explore_banner_queries=ExploreBannerQueryService(
banners=ExploreBannerQueryRepository(client=database_client),
banners=ExploreBannerQueryRepository(session_factory=database_client),
enabled=FeatureService.is_explore_banner_enabled(),
),
schema_definitions=SchemaDefinitionService(source_factory=SchemaManager),
setup=SetupService(
state=installation_state,
accounts=RegisterServiceAccountProvisioner(client=database_client),
accounts=RegisterServiceAccountProvisioner(session_factory=database_client),
lock=RedisSetupLock(client=redis),
setup_required=deployment_edition != DeploymentEdition.CLOUD,
),

View File

@ -11,8 +11,8 @@ from services.explore_banner_query_service import ExploreBannerQuery, ExploreBan
class ExploreBannerQueryRepository(ExploreBannerQuery):
def __init__(self, client: sessionmaker[Session]) -> None:
self._client = client
def __init__(self, session_factory: sessionmaker[Session]) -> None:
self._session_factory = session_factory
@override
def list_enabled(self, language: str) -> tuple[ExploreBannerRecord, ...]:
@ -32,7 +32,7 @@ class ExploreBannerQueryRepository(ExploreBannerQuery):
.order_by(ExporleBanner.sort)
)
with self._client() as session:
with self._session_factory() as session:
rows = session.execute(stmt).all()
return tuple(
ExploreBannerRecord(

View File

@ -12,16 +12,16 @@ from models.model import DifySetup
class InstallationStateRepository:
"""Read persistent state shared by installation bootstrap use cases."""
def __init__(self, client: sessionmaker[Session]) -> None:
self._client = client
def __init__(self, session_factory: sessionmaker[Session]) -> None:
self._session_factory = session_factory
def get_setup_at(self) -> datetime | None:
with self._client() as session:
with self._session_factory() as session:
return session.scalar(select(DifySetup.setup_at).limit(1))
def is_setup(self) -> bool:
return self.get_setup_at() is not None
def has_tenants(self) -> bool:
with self._client() as session:
with self._session_factory() as session:
return session.scalar(select(exists().select_from(Tenant))) is True

View File

@ -11,8 +11,8 @@ from services.workspace_query_service import WorkspaceQuery, WorkspaceRecord
class WorkspaceQueryRepository(WorkspaceQuery, AccountWorkspaceMembershipQuery):
def __init__(self, client: sessionmaker[Session]) -> None:
self._client = client
def __init__(self, session_factory: sessionmaker[Session]) -> None:
self._session_factory = session_factory
@override
def list_for_account(self, account_id: str) -> tuple[WorkspaceRecord, ...]:
@ -32,7 +32,7 @@ class WorkspaceQueryRepository(WorkspaceQuery, AccountWorkspaceMembershipQuery):
.order_by(Tenant.created_at.asc())
)
with self._client() as session:
with self._session_factory() as session:
rows = session.execute(stmt).all()
return tuple(
WorkspaceRecord(
@ -48,5 +48,5 @@ class WorkspaceQueryRepository(WorkspaceQuery, AccountWorkspaceMembershipQuery):
@override
def list_ids_for_account(self, account_id: str) -> tuple[str, ...]:
stmt = select(TenantAccountJoin.tenant_id).where(TenantAccountJoin.account_id == account_id)
with self._client() as session:
with self._session_factory() as session:
return tuple(session.scalars(stmt).all())

View File

@ -14,12 +14,12 @@ _SETUP_LOCK_TIMEOUT_SECONDS = 300
class RegisterServiceAccountProvisioner(SetupAccountProvisioner):
def __init__(self, client: sessionmaker[Session]) -> None:
self._client = client
def __init__(self, session_factory: sessionmaker[Session]) -> None:
self._session_factory = session_factory
@override
def provision(self, setup: SetupInput) -> None:
with self._client() as session:
with self._session_factory() as session:
RegisterService.setup(
email=setup.email,
name=setup.name,

View File

@ -8,7 +8,7 @@ from repositories.installation_state_repository import InstallationStateReposito
def test_empty_database_has_no_installation_state(sqlite_session_factory: sessionmaker[Session]) -> None:
repository = InstallationStateRepository(client=sqlite_session_factory)
repository = InstallationStateRepository(session_factory=sqlite_session_factory)
assert repository.get_setup_at() is None
assert repository.is_setup() is False
@ -23,7 +23,7 @@ def test_get_setup_at_returns_persisted_timestamp(
sqlite_session.add(setup)
sqlite_session.commit()
sqlite_session.refresh(setup)
repository = InstallationStateRepository(client=sqlite_session_factory)
repository = InstallationStateRepository(session_factory=sqlite_session_factory)
assert repository.get_setup_at() == setup.setup_at
assert repository.is_setup() is True
@ -35,6 +35,6 @@ def test_has_tenants_detects_existing_tenant(
) -> None:
sqlite_session.add(Tenant(name="Existing workspace"))
sqlite_session.commit()
repository = InstallationStateRepository(client=sqlite_session_factory)
repository = InstallationStateRepository(session_factory=sqlite_session_factory)
assert repository.has_tenants() is True

View File

@ -14,7 +14,7 @@ from services.setup_service import SetupInput
def test_provision_delegates_to_register_service_with_managed_session(
sqlite_session_factory: sessionmaker[Session],
) -> None:
provisioner = RegisterServiceAccountProvisioner(client=sqlite_session_factory)
provisioner = RegisterServiceAccountProvisioner(session_factory=sqlite_session_factory)
setup = SetupInput(
email="admin@example.com",
name="Admin",