diff --git a/src/utils_flask_sqla/generic.py b/src/utils_flask_sqla/generic.py index 85b3be1..eee8759 100644 --- a/src/utils_flask_sqla/generic.py +++ b/src/utils_flask_sqla/generic.py @@ -4,6 +4,7 @@ import sqlalchemy as sa from dateutil import parser from flask_sqlalchemy import SQLAlchemy +from sqlalchemy import MetaData, inspect from sqlalchemy.dialects.postgresql import UUID from sqlalchemy.types import Boolean, Date, DateTime, Integer, Numeric from werkzeug.exceptions import BadRequest @@ -104,17 +105,46 @@ def __init__(self, tableName, schemaName, engine): - engine : sqlalchemy instance engine for exemple : DB.engine if DB = Sqlalchemy() """ + # En sqlalchemy 2.0, il faut utiliser MetaData() + meta = MetaData() + # Try to reflect just the specific table instead of all tables try: - with db.engine.connect() as conn: - self.tableDef = sa.Table( - tableName, - db.metadata, - schema=schemaName, - autoload_with=conn, - ) - except KeyError: - raise KeyError("Table {}.{} doesn't exists".format(schemaName, tableName)) + + meta.reflect(only=[tableName], schema=schemaName, views=True, bind=engine) + table_key = f"{schemaName}.{tableName}" + + if table_key in meta.tables: + self.tableDef = meta.tables[table_key] + # If not found with schema, try without schema + elif tableName in meta.tables: + self.tableDef = meta.tables[tableName] + else: + # Si on ne trouve pas la table, en essaye de la trouver dans le schema et les vues + + inspector = inspect(engine) + available_views = inspector.get_view_names(schema=schemaName) + available_tables = inspector.get_table_names(schema=schemaName) + if tableName in available_views or tableName in available_tables: + # Force reflection with explicit view flag + meta = MetaData() + meta.reflect( + only=[tableName], + schema=schemaName, + views=True, + bind=engine, + extend_existing=True, + ) + table_key = f"{schemaName}.{tableName}" + if table_key in meta.tables: + self.tableDef = meta.tables[table_key] + else: + raise KeyError(f"table {schemaName}.{tableName} doesn't exist") + else: + raise KeyError(f"table {schemaName}.{tableName} doesn't exist") + except Exception as e: + # If any error occurs, provide a detailed error message + raise KeyError(f"Error accessing table {schemaName}.{tableName}: {str(e)}") # Mise en place d'un mapping des colonnes en vue d'une sérialisation self.serialize_columns, self.db_cols = self.get_serialized_columns() @@ -256,8 +286,8 @@ def raw_query(self, process_filter=True): Renvoie la requete 'brute' (sans .all) - process_filter: application des filtres (et du sort) """ - - q = self.DB.session.query(self.view.tableDef) + # Use select() instead of query() + q = self.DB.select(self.view.tableDef) if not process_filter: return q @@ -274,13 +304,24 @@ def query(self): """ Lance la requete et retourne l'objet sqlalchemy """ - q = self.DB.session.query(self.view.tableDef) - nb_result_without_filter = q.count() + # Use select() instead of query() + nb_result_without_filter = self.DB.session.scalar( + self.DB.select(self.DB.func.count()).select_from(self.view.tableDef) + ) + # Get filtered query using raw_query q = self.raw_query(process_filter=True) total_filtered = q.limit(None).count() if self.filters else nb_result_without_filter - data = q.all() + # Calculate total filtered rows + if self.filters: + count_stmt = self.DB.select(self.DB.func.count()).select_from(q.subquery()) + total_filtered = self.DB.session.scalar(count_stmt) + else: + total_filtered = nb_result_without_filter + + # Execute query + data = self.DB.session.execute(q).all() return data, nb_result_without_filter, total_filtered diff --git a/src/utils_flask_sqla/serializers.py b/src/utils_flask_sqla/serializers.py index 2011e39..58e84a7 100644 --- a/src/utils_flask_sqla/serializers.py +++ b/src/utils_flask_sqla/serializers.py @@ -8,6 +8,7 @@ from itertools import chain from functools import lru_cache from uuid import UUID +from flask import current_app from sqlalchemy.orm import ColumnProperty from sqlalchemy import inspect @@ -358,7 +359,7 @@ def populatefn(self, dict_in, recursif=False): recursif: si on renseigne les relationships """ - + db = current_app.extensions["sqlalchemy"] cls_db_columns_key = list(map(lambda x: x[0], get_cls_db_columns())) # populate cls_db_columns @@ -403,10 +404,11 @@ def populatefn(self, dict_in, recursif=False): # preload with id # pour faire une seule requête - ids = filter(lambda x: x, map(lambda x: x.get(id_field_name), values)) - preload_res_with_ids = Model.query.where( - getattr(Model, id_field_name).in_(ids) - ).all() + ids = set(filter(lambda x: x, map(lambda x: x.get(id_field_name), values))) + + stmt = select(Model).where(getattr(Model, id_field_name).in_(ids)) + + preload_res_with_ids = db.session.execute(stmt).scalars().all() # resul v_obj = [] @@ -414,16 +416,18 @@ def populatefn(self, dict_in, recursif=False): for data in values: id_value = data.pop(id_field_name, None) + # On filtre la liste des objets préchargés + filtered_results = list( + filter( + lambda x: getattr(x, id_field_name) == id_value, + preload_res_with_ids, + ) + ) + res = ( - # si on a une id -> on recupère dans la liste preload_res_with_ids - # TODO trouver un find plus propre ? - list( - filter( - lambda x: getattr(x, id_field_name) == id_value, - preload_res_with_ids, - ) - )[0] - if id_value and len(preload_res_with_ids) + # si on a une id et qu'on a trouvé au moins un résultat + filtered_results[0] + if id_value and filtered_results # sinon on cree une nouvelle instance else Model() )