Coverage for gws-app/gws/base/database/provider.py: 87%
189 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-10-05 13:35 +0200
« prev ^ index » next coverage.py v7.16.2, created at 2026-10-05 13:35 +0200
1"""Base database provider."""
3import threading
4from typing import Optional, cast
6import gws
7import gws.lib.sa as sa
9from . import connection
12class Config(gws.Config):
13 """Database provider"""
15 schemaCacheLifeTime: gws.Duration = '3600'
16 """How long table structures read from the database are cached."""
17 withPool: Optional[bool] = False
18 """Keep and reuse database connections in a pool."""
19 pool: Optional[dict]
20 """Connection pool options."""
23_thread_local = threading.local()
26class Object(gws.DatabaseProvider):
27 """Base database provider.
29 Manages the SQLAlchemy engine and the per-thread connection, reflects table
30 structures and describes tables and columns, and runs plain SQL text.
32 Subclasses provide ``url``, ``split_table_name``, ``join_table_name`` and
33 ``table_bounds``, and extend ``describe_column`` for database-specific types.
34 """
36 saEngine: sa.Engine
37 """SQLAlchemy engine."""
38 saMetaMap: dict[str, sa.MetaData]
39 """Reflected metadata, keyed by schema name."""
41 def __getstate__(self):
42 """Return the state for pickling, without the engine and the metadata."""
43 return gws.u.omit(vars(self), 'saMetaMap', 'saEngine')
45 def configure(self):
46 # init a dummy engine just to check things
47 self.saEngine = self.create_engine(poolclass=sa.NullPool)
48 self.saMetaMap = {}
50 def activate(self):
51 self.saEngine = self.create_engine()
52 self.saMetaMap = {}
54 def engine(self):
55 eng = getattr(self, 'saEngine', None)
56 if eng is not None:
57 return eng
58 self.saEngine = self.create_engine()
59 return self.saEngine
61 def create_engine(self, **kwargs):
62 eng = sa.create_engine(self.url(), **self.engine_options(**kwargs))
63 # setattr(eng, '_connection_cls', connection.Object)
64 return eng
66 def engine_options(self, **kwargs):
67 if self.root.app.developer_option('db.engine_echo'):
68 kwargs.setdefault('echo', True)
69 kwargs.setdefault('echo_pool', True)
71 if self.cfg('withPool') is False:
72 kwargs.setdefault('poolclass', sa.NullPool)
73 return kwargs
75 pool = self.cfg('pool') or {}
76 p = pool.get('disabled')
77 if p is True:
78 kwargs.setdefault('poolclass', sa.NullPool)
79 return kwargs
81 p = pool.get('pre_ping')
82 if p is True:
83 kwargs.setdefault('pool_pre_ping', True)
84 p = pool.get('size')
85 if isinstance(p, int):
86 kwargs.setdefault('pool_size', p)
87 p = pool.get('recycle')
88 if isinstance(p, int):
89 kwargs.setdefault('pool_recycle', p)
90 p = pool.get('timeout')
91 if isinstance(p, int):
92 kwargs.setdefault('pool_timeout', p)
94 return kwargs
96 def inspect_schema(self, schema, options=None):
97 if options and options.refresh:
98 self.saMetaMap.pop(schema, None)
100 if schema in self.saMetaMap:
101 return
103 def _load():
104 md = sa.MetaData(schema=schema)
106 # introspecting the whole schema is generally faster
107 # but what if we only need a single table from a big schema?
108 # @TODO add options for reflection
110 gws.debug.time_start(f'AUTOLOAD {self.uid=} {schema=}')
111 with self.connect() as conn:
112 md.reflect(conn.saConn, schema, resolve_fks=False, views=True)
113 gws.debug.time_end()
114 return md
116 life_time = self.cfg('schemaCacheLifeTime', 0)
117 if options and options.cacheLifeTime is not None:
118 life_time = options.cacheLifeTime
119 if not life_time:
120 self.saMetaMap[schema] = _load()
121 else:
122 self.saMetaMap[schema] = gws.u.get_cached_object(f'database_metadata_schema_{schema}', life_time, _load)
124 def connect(self):
125 conn = self._open_connection()
126 return connection.Object(self, conn)
128 def _sa_connection(self) -> sa.Connection | None:
129 """Return the open connection of the current thread, if any."""
130 return getattr(_thread_local, '_connection', None)
132 def _open_connection(self) -> sa.Connection:
133 """Return the connection of the current thread, opening it if needed, and increment the counter."""
134 conn = getattr(_thread_local, '_connection', None)
135 cc = getattr(_thread_local, '_connectionCount', 0)
137 if conn is None:
138 assert cc == 0
139 conn = self.engine().connect()
140 setattr(_thread_local, '_connection', conn)
141 else:
142 assert cc > 0
144 setattr(_thread_local, '_connectionCount', cc + 1)
145 # gws.log.debug(f'db.connect: open: {cc + 1}')
146 return conn
148 def _close_connection(self):
149 """Decrement the connection counter of the current thread, closing the connection at zero."""
150 conn = getattr(_thread_local, '_connection', None)
151 cc = getattr(_thread_local, '_connectionCount', 0)
152 assert conn is not None
153 assert cc > 0
154 # gws.log.debug(f'db.connect: close: {cc}')
155 if cc == 1:
156 if conn:
157 conn.close()
158 setattr(_thread_local, '_connection', None)
159 setattr(_thread_local, '_connectionCount', 0)
160 else:
161 setattr(_thread_local, '_connectionCount', cc - 1)
163 def table(self, table, **kwargs):
164 tab = self._sa_table(table)
165 if tab is None:
166 raise sa.Error(f'table not found: {table!r}')
167 return tab
169 def count(self, table):
170 tab = self._sa_table(table)
171 if tab is None:
172 return 0
173 sql = sa.select(sa.func.count()).select_from(tab)
174 with self.connect() as conn:
175 return conn.fetch_int(sql)
177 def has_schema(self, schema):
178 return schema in self.schema_names()
180 def schema_names(self):
181 inspector = sa.inspect(self.engine())
182 return inspector.get_schema_names()
184 def has_table(self, table_name: str):
185 tab = self._sa_table(table_name)
186 return tab is not None
188 def _sa_table(self, tab_or_name) -> sa.Table | None:
189 """Return a reflected table by name, or ``None`` if the table does not exist."""
190 if isinstance(tab_or_name, sa.Table):
191 return tab_or_name
192 schema, name = self.split_table_name(tab_or_name)
193 self.inspect_schema(schema)
194 # see _get_table_key in sqlalchemy/sql/schema.py
195 table_key = schema + '.' + name
196 sm = self.saMetaMap.get(schema)
197 if sm is None:
198 raise sa.Error(f'schema {schema!r} not found')
199 return sm.tables.get(table_key)
201 def column(self, table, column_name):
202 tab = self.table(table)
203 try:
204 return tab.columns[column_name]
205 except KeyError:
206 raise sa.Error(f'column {str(table)}.{column_name!r} not found')
208 def has_column(self, table, column_name):
209 tab = self._sa_table(table)
210 return tab is not None and column_name in tab.columns
212 def select_text(self, sql, **kwargs):
213 with self.connect() as conn:
214 try:
215 return [gws.u.to_dict(r) for r in conn.execute(sa.text(sql), kwargs)]
216 except sa.Error:
217 conn.rollback()
218 raise
220 def execute_text(self, sql, **kwargs):
221 with self.connect() as conn:
222 try:
223 res = conn.execute(sa.text(sql), kwargs)
224 conn.commit()
225 return res
226 except sa.Error:
227 conn.rollback()
228 raise
230 SA_TO_ATTR = {
231 # common: sqlalchemy.sql.sqltypes
232 'BIGINT': gws.AttributeType.int,
233 'BOOLEAN': gws.AttributeType.bool,
234 'CHAR': gws.AttributeType.str,
235 'DATE': gws.AttributeType.date,
236 'DOUBLE_PRECISION': gws.AttributeType.float,
237 'INTEGER': gws.AttributeType.int,
238 'NUMERIC': gws.AttributeType.float,
239 'REAL': gws.AttributeType.float,
240 'SMALLINT': gws.AttributeType.int,
241 'TEXT': gws.AttributeType.str,
242 # 'UUID': ...,
243 'VARCHAR': gws.AttributeType.str,
244 # postgres specific: sqlalchemy.dialects.postgresql.types
245 # 'JSON': ...,
246 # 'JSONB': ...,
247 # 'BIT': ...,
248 'BYTEA': gws.AttributeType.bytes,
249 # 'CIDR': ...,
250 # 'INET': ...,
251 # 'MACADDR': ...,
252 # 'MACADDR8': ...,
253 # 'MONEY': ...,
254 'TIME': gws.AttributeType.time,
255 'TIMESTAMP': gws.AttributeType.datetime,
256 }
257 """Attribute types for SQLAlchemy type names."""
259 # @TODO proper support for Z/M geoms
261 SA_TO_GEOM = {
262 'POINT': gws.GeometryType.point,
263 'POINTM': gws.GeometryType.point,
264 'POINTZ': gws.GeometryType.point,
265 'POINTZM': gws.GeometryType.point,
266 'LINESTRING': gws.GeometryType.linestring,
267 'LINESTRINGM': gws.GeometryType.linestring,
268 'LINESTRINGZ': gws.GeometryType.linestring,
269 'LINESTRINGZM': gws.GeometryType.linestring,
270 'POLYGON': gws.GeometryType.polygon,
271 'POLYGONM': gws.GeometryType.polygon,
272 'POLYGONZ': gws.GeometryType.polygon,
273 'POLYGONZM': gws.GeometryType.polygon,
274 'MULTIPOINT': gws.GeometryType.multipoint,
275 'MULTIPOINTM': gws.GeometryType.multipoint,
276 'MULTIPOINTZ': gws.GeometryType.multipoint,
277 'MULTIPOINTZM': gws.GeometryType.multipoint,
278 'MULTILINESTRING': gws.GeometryType.multilinestring,
279 'MULTILINESTRINGM': gws.GeometryType.multilinestring,
280 'MULTILINESTRINGZ': gws.GeometryType.multilinestring,
281 'MULTILINESTRINGZM': gws.GeometryType.multilinestring,
282 'MULTIPOLYGON': gws.GeometryType.multipolygon,
283 # 'GEOMETRYCOLLECTION': gws.GeometryType.geometrycollection,
284 # 'CURVE': gws.GeometryType.curve,
285 }
286 """Geometry types for database geometry type names."""
288 UNKNOWN_TYPE = gws.AttributeType.str
289 """Attribute type for columns of unknown types."""
290 UNKNOWN_ARRAY_TYPE = gws.AttributeType.strlist
291 """Attribute type for array columns of unknown item types."""
293 def describe(self, table):
294 tab = self._sa_table(table)
295 if tab is None:
296 raise sa.Error(f'table not found: {table!r}')
298 schema = tab.schema
299 name = tab.name
301 desc = gws.DataSetDescription(
302 columns=[],
303 columnMap={},
304 fullName=self.join_table_name(schema or '', name),
305 geometryName='',
306 geometrySrid=0,
307 geometryType='',
308 name=name,
309 schema=schema,
310 )
312 for n, sa_col in enumerate(cast(list[sa.Column], tab.columns)):
313 col = self.describe_column(table, sa_col.name)
314 col.columnIndex = n
315 desc.columns.append(col)
316 desc.columnMap[col.name] = col
318 for col in desc.columns:
319 if col.geometryType:
320 desc.geometryName = col.name
321 desc.geometryType = col.geometryType
322 desc.geometrySrid = col.geometrySrid
323 break
325 return desc
327 def describe_column(self, table, column_name):
328 sa_col = self.column(table, column_name)
330 col = gws.ColumnDescription(
331 columnIndex=0,
332 comment=str(sa_col.comment or ''),
333 default=sa_col.default,
334 geometrySrid=0,
335 geometryType='',
336 isAutoincrement=bool(sa_col.autoincrement),
337 isNullable=bool(sa_col.nullable),
338 isPrimaryKey=bool(sa_col.primary_key),
339 isUnique=bool(sa_col.unique),
340 hasDefault=sa_col.server_default is not None,
341 name=str(sa_col.name),
342 nativeType='',
343 type='',
344 )
346 col.nativeType = type(sa_col.type).__name__.upper()
347 col.type = self.SA_TO_ATTR.get(col.nativeType, self.UNKNOWN_TYPE)
349 return col
352##