Coverage for gws-app/gws/base/database/connection.py: 92%
74 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"""Database connection wrapper."""
3import gws
4import gws.lib.sa as sa
7class Object(gws.DatabaseConnection):
8 """Database connection."""
10 db: gws.DatabaseProvider
11 """Provider this connection belongs to."""
12 saConn: sa.Connection
13 """The underlying SQLAlchemy connection."""
15 def __init__(self, db: gws.DatabaseProvider, conn: sa.Connection):
16 """Create a connection wrapper.
18 Args:
19 db: Provider this connection belongs to.
20 conn: The SQLAlchemy connection.
21 """
22 self.db = db
23 self.saConn = conn
25 def __enter__(self):
26 return self
28 def __exit__(self, exc_type, exc_value, traceback):
29 self.close()
31 def close(self):
32 getattr(self.db, '_close_connection')()
34 def execute(self, stmt, params=None, execution_options=None):
35 return self.saConn.execute(stmt, params, execution_options=execution_options)
37 def commit(self):
38 self.saConn.commit()
40 def rollback(self):
41 self.saConn.rollback()
43 def exec(self, sql, **params):
44 if isinstance(sql, str):
45 sql = sa.text(sql)
46 return self.saConn.execute(sql, params)
48 def exec_commit(self, sql, **params):
49 if isinstance(sql, str):
50 sql = sa.text(sql)
51 try:
52 res = self.saConn.execute(sql, params)
53 self.saConn.commit()
54 return res
55 except Exception:
56 self.saConn.rollback()
57 raise
59 def exec_rollback(self, sql, **params):
60 if isinstance(sql, str):
61 sql = sa.text(sql)
62 try:
63 return self.saConn.execute(sql, params)
64 finally:
65 self.saConn.rollback()
67 def fetch_all(self, stmt, **params):
68 return [r._asdict() for r in self.exec_rollback(stmt, **params)]
70 def fetch_first(self, stmt, **params):
71 res = self.exec_rollback(stmt, **params)
72 r = res.first()
73 return r._asdict() if r else None
75 def fetch_scalars(self, stmt, **params):
76 res = self.exec_rollback(stmt, **params)
77 return list(res.scalars().all())
79 def fetch_strings(self, stmt, **params):
80 res = self.exec_rollback(stmt, **params)
81 return [_to_str(s) for s in res.scalars().all()]
83 def fetch_ints(self, stmt, **params):
84 res = self.exec_rollback(stmt, **params)
85 return [_to_int(s) for s in res.scalars().all()]
87 def fetch_scalar(self, stmt, **params):
88 res = self.exec_rollback(stmt, **params)
89 return res.scalar()
91 def fetch_string(self, stmt, **params):
92 res = self.exec_rollback(stmt, **params)
93 s = res.scalar()
94 return _to_str(s) if s is not None else None
96 def fetch_int(self, stmt, **params):
97 res = self.exec_rollback(stmt, **params)
98 s = res.scalar()
99 return _to_int(s) if s is not None else None
102##
105def _to_int(s) -> int:
106 """Return the value if it is an int, raise ``ValueError`` otherwise."""
107 if isinstance(s, int):
108 return s
109 raise ValueError(f'db: expected int, got {s=}')
112def _to_str(s) -> str:
113 """Convert a value to a string, ``None`` to an empty string."""
114 if s is None:
115 return ''
116 return str(s)