Add function close_database

This commit is contained in:
2026-02-08 14:48:48 +01:00
parent 6af6176dec
commit aec5c278a1
+16 -8
View File
@@ -21,7 +21,7 @@ from peewee import (
from playhouse.pool import PooledMySQLDatabase from playhouse.pool import PooledMySQLDatabase
# Re-export PeeweeException as DatabseException, DoesNotExist and # Re-export PeeweeException as DatabaseError, DoesNotExist and
# EncodingError for convenience # EncodingError for convenience
__all__ = [ __all__ = [
'db', 'db',
@@ -36,6 +36,7 @@ __all__ = [
'get_hostname_for_user', 'get_hostname_for_user',
'get_permission', 'get_permission',
'get_user', 'get_user',
'close_database',
'DoesNotExist', 'DoesNotExist',
'EncodingError', 'EncodingError',
'DatabaseError', 'DatabaseError',
@@ -165,7 +166,7 @@ TABLE_TO_MODEL = {
} }
def init_database(config: dict): def init_database(config):
""" """
Initialize database connection based on config. Initialize database connection based on config.
@@ -210,7 +211,7 @@ def init_database(config: dict):
db.connect() db.connect()
def _migrate_table_sqlite(model_class, from_version: int, column_map: dict): def _migrate_table_sqlite(model_class, from_version, column_map):
""" """
Migrate a single SQLite table using Peewee model for schema. Migrate a single SQLite table using Peewee model for schema.
@@ -257,7 +258,7 @@ def _migrate_table_sqlite(model_class, from_version: int, column_map: dict):
db.execute_sql(f'DROP TABLE "{backup_name}"') db.execute_sql(f'DROP TABLE "{backup_name}"')
def _migrate_sqlite(from_version: int, to_version: int): def _migrate_sqlite(from_version, to_version):
"""Migrate SQLite from from_version to to_version.""" """Migrate SQLite from from_version to to_version."""
if to_version == 3: if to_version == 3:
_migrate_v3_create_permissions() _migrate_v3_create_permissions()
@@ -283,7 +284,7 @@ def _migrate_v3_create_permissions():
) )
def _migrate_mariadb(to_version: int): def _migrate_mariadb(to_version):
"""Migrate MariaDB to target version using ALTER TABLE.""" """Migrate MariaDB to target version using ALTER TABLE."""
if to_version == 2: if to_version == 2:
db.execute_sql('ALTER TABLE hostnames DROP INDEX hostnames_hostname') db.execute_sql('ALTER TABLE hostnames DROP INDEX hostnames_hostname')
@@ -332,7 +333,7 @@ def create_tables():
logging.debug("Database tables created") logging.debug("Database tables created")
def get_user(username: str) -> User: def get_user(username):
""" """
Get user by username. Get user by username.
@@ -399,7 +400,7 @@ def get_permission(user, fqdn):
raise DoesNotExist raise DoesNotExist
def get_hostname(hostname: str, zone: str) -> Hostname: def get_hostname(hostname, zone):
""" """
Get hostname by name and zone. Get hostname by name and zone.
@@ -428,7 +429,7 @@ def get_hostname(hostname: str, zone: str) -> Hostname:
def get_hostname_for_user( def get_hostname_for_user(
user: User, hostname: str, zone: str, dns_ttl: int, expiry_ttl: int): user: User, hostname: str, zone: str, dns_ttl: int, expiry_ttl: int):
""" """
Get hostname if it exists or create a new intance. Get hostname if it exists or create a new instance.
Args: Args:
user: User requesting access. user: User requesting access.
@@ -466,3 +467,10 @@ def get_hostname_for_user(
), ),
True True
) )
def close_database():
"""Close database connection."""
if not db.is_closed():
db.close()
logging.debug("Database connection closed")