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

1"""Base database provider.""" 

2 

3import threading 

4from typing import Optional, cast 

5 

6import gws 

7import gws.lib.sa as sa 

8 

9from . import connection 

10 

11 

12class Config(gws.Config): 

13 """Database provider""" 

14 

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.""" 

21 

22 

23_thread_local = threading.local() 

24 

25 

26class Object(gws.DatabaseProvider): 

27 """Base database provider. 

28 

29 Manages the SQLAlchemy engine and the per-thread connection, reflects table 

30 structures and describes tables and columns, and runs plain SQL text. 

31 

32 Subclasses provide ``url``, ``split_table_name``, ``join_table_name`` and 

33 ``table_bounds``, and extend ``describe_column`` for database-specific types. 

34 """ 

35 

36 saEngine: sa.Engine 

37 """SQLAlchemy engine.""" 

38 saMetaMap: dict[str, sa.MetaData] 

39 """Reflected metadata, keyed by schema name.""" 

40 

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') 

44 

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 = {} 

49 

50 def activate(self): 

51 self.saEngine = self.create_engine() 

52 self.saMetaMap = {} 

53 

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 

60 

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 

65 

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) 

70 

71 if self.cfg('withPool') is False: 

72 kwargs.setdefault('poolclass', sa.NullPool) 

73 return kwargs 

74 

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 

80 

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) 

93 

94 return kwargs 

95 

96 def inspect_schema(self, schema, options=None): 

97 if options and options.refresh: 

98 self.saMetaMap.pop(schema, None) 

99 

100 if schema in self.saMetaMap: 

101 return 

102 

103 def _load(): 

104 md = sa.MetaData(schema=schema) 

105 

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 

109 

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 

115 

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) 

123 

124 def connect(self): 

125 conn = self._open_connection() 

126 return connection.Object(self, conn) 

127 

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) 

131 

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) 

136 

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 

143 

144 setattr(_thread_local, '_connectionCount', cc + 1) 

145 # gws.log.debug(f'db.connect: open: {cc + 1}') 

146 return conn 

147 

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) 

162 

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 

168 

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) 

176 

177 def has_schema(self, schema): 

178 return schema in self.schema_names() 

179 

180 def schema_names(self): 

181 inspector = sa.inspect(self.engine()) 

182 return inspector.get_schema_names() 

183 

184 def has_table(self, table_name: str): 

185 tab = self._sa_table(table_name) 

186 return tab is not None 

187 

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) 

200 

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') 

207 

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 

211 

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 

219 

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 

229 

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.""" 

258 

259 # @TODO proper support for Z/M geoms 

260 

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.""" 

287 

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.""" 

292 

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}') 

297 

298 schema = tab.schema 

299 name = tab.name 

300 

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 ) 

311 

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 

317 

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 

324 

325 return desc 

326 

327 def describe_column(self, table, column_name): 

328 sa_col = self.column(table, column_name) 

329 

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 ) 

345 

346 col.nativeType = type(sa_col.type).__name__.upper() 

347 col.type = self.SA_TO_ATTR.get(col.nativeType, self.UNKNOWN_TYPE) 

348 

349 return col 

350 

351 

352##