diff --git a/src/tagstudio/core/library/alchemy/library.py b/src/tagstudio/core/library/alchemy/library.py index 6316c172..28025fc7 100644 --- a/src/tagstudio/core/library/alchemy/library.py +++ b/src/tagstudio/core/library/alchemy/library.py @@ -506,13 +506,12 @@ class Library: # migrate if necessary try: - migrations = DBMigrations(library_dir, sql_filename) + with DBMigrations(library_dir, sql_filename) as migrations: + # save backup if patches will be applied + if migrations.required: + Library.save_library_backup_to_disk(library_dir) - # save backup if patches will be applied - if migrations.required: - Library.save_library_backup_to_disk(library_dir) - - migrations.run() + migrations.run() except MigrationError as e: return LibraryStatus(success=False, message=e.args[0]) diff --git a/src/tagstudio/core/library/alchemy/migrations.py b/src/tagstudio/core/library/alchemy/migrations.py index ecec5a93..c3cb0b43 100644 --- a/src/tagstudio/core/library/alchemy/migrations.py +++ b/src/tagstudio/core/library/alchemy/migrations.py @@ -78,11 +78,21 @@ class DBMigrations: f"Opening Library with DB Version {self.loaded_db_version}/{DB_VERSION}" ) + self._exited = False + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_value, exc_traceback): + self._connection.close() + self._exited = True + @property def required(self) -> bool: return self.loaded_db_version < DB_VERSION def run(self): + assert not self._exited if not self.required: return