diff --git a/api/extensions/ext_application_services.py b/api/extensions/ext_application_services.py index 287344cdd81..dc4015767ca 100644 --- a/api/extensions/ext_application_services.py +++ b/api/extensions/ext_application_services.py @@ -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, ), diff --git a/api/repositories/explore_banner_query_repository.py b/api/repositories/explore_banner_query_repository.py index 011828b7977..254929af958 100644 --- a/api/repositories/explore_banner_query_repository.py +++ b/api/repositories/explore_banner_query_repository.py @@ -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( diff --git a/api/repositories/installation_state_repository.py b/api/repositories/installation_state_repository.py index 755a3134933..f34b3c54fa9 100644 --- a/api/repositories/installation_state_repository.py +++ b/api/repositories/installation_state_repository.py @@ -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 diff --git a/api/repositories/workspace_query_repository.py b/api/repositories/workspace_query_repository.py index 3e1a7eaa95a..ed4b41293e5 100644 --- a/api/repositories/workspace_query_repository.py +++ b/api/repositories/workspace_query_repository.py @@ -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()) diff --git a/api/services/setup_adapters.py b/api/services/setup_adapters.py index e8fc6098905..01f0465ac8e 100644 --- a/api/services/setup_adapters.py +++ b/api/services/setup_adapters.py @@ -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, diff --git a/api/tests/unit_tests/repositories/test_installation_state_repository.py b/api/tests/unit_tests/repositories/test_installation_state_repository.py index 9f82ef8cd25..8f32539c617 100644 --- a/api/tests/unit_tests/repositories/test_installation_state_repository.py +++ b/api/tests/unit_tests/repositories/test_installation_state_repository.py @@ -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 diff --git a/api/tests/unit_tests/services/test_setup_adapters.py b/api/tests/unit_tests/services/test_setup_adapters.py index c99c97386c2..691a78cf557 100644 --- a/api/tests/unit_tests/services/test_setup_adapters.py +++ b/api/tests/unit_tests/services/test_setup_adapters.py @@ -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",