e-data 2.0.2.dev137__py3-none-any.whl → 2.0.2.dev138__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: e-data
3
- Version: 2.0.2.dev137
3
+ Version: 2.0.2.dev138
4
4
  Summary: Python library for managing spanish energy data from various web providers
5
5
  Author-email: VMG <vmayorg@outlook.es>
6
6
  License: GNU GENERAL PUBLIC LICENSE
@@ -696,10 +696,9 @@ Requires-Dist: Jinja2<4,>=3.1
696
696
  Requires-Dist: pydantic<3,>=2.10
697
697
  Requires-Dist: python_dateutil<3,>=2.8
698
698
  Requires-Dist: Requests<3,>=2.31
699
- Requires-Dist: SQLAlchemy[asyncio]<3,>=2.0
699
+ Requires-Dist: SQLAlchemy<3,>=2.0
700
700
  Requires-Dist: sqlmodel<0.1,>=0.0.45
701
701
  Requires-Dist: typer<1,>=0.12
702
- Requires-Dist: aiosqlite<1,>=0.21
703
702
  Dynamic: license-file
704
703
 
705
704
  [![Downloads](https://pepy.tech/badge/e-data)](https://pepy.tech/project/e-data)
@@ -1,4 +1,4 @@
1
- e_data-2.0.2.dev137.dist-info/licenses/LICENSE,sha256=OXLcl0T2SZ8Pmy2_dmlvKuetivmyPd5m1q-Gyd-zaYY,35149
1
+ e_data-2.0.2.dev138.dist-info/licenses/LICENSE,sha256=OXLcl0T2SZ8Pmy2_dmlvKuetivmyPd5m1q-Gyd-zaYY,35149
2
2
  edata/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
3
3
  edata/cli.py,sha256=CsBYDggt0grkLb4cebrckOvJyHr9PaGTnlWqT4yjuFo,5215
4
4
  edata/core/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
@@ -6,7 +6,7 @@ edata/core/completion.py,sha256=A0ekfu7eCjaLblUYO5DrvaWuDPAO9k1buzy2qdfjdNc,3914
6
6
  edata/core/const.py,sha256=X9J2f7lt9jXJnghqz0JFi-J7USa3xlYvbuiki70eHNE,43
7
7
  edata/core/utils.py,sha256=uQX8GmNjuhof56jiXXoV-Fk-d590OuIoBJbM9ymZyV8,3194
8
8
  edata/database/__init__.py,sha256=HP_pDbGLzeMY9Ufm_Dps5cxzW_3YTonT8w8dstG558c,46
9
- edata/database/controller.py,sha256=bI3cE1Q4SK9U1Eep_61EhlmyVjgYH87F5lfBLhwSpmg,23211
9
+ edata/database/controller.py,sha256=iNZcmdJYvw0sBu02VuBs9HypvK-NSpqFeOKx0ta6hMM,23967
10
10
  edata/database/models.py,sha256=WWZHYgFBgSDqHORxw5z4Vip7HvorMcQbt9oCpV61NJA,5780
11
11
  edata/database/queries.py,sha256=XziDSG2V-PVFmvRkj8UrEvgrc9Qc-ShhnxgsPHIsyik,8304
12
12
  edata/database/utils.py,sha256=V-FJfyAh8YJBGHB85BbWtG3Ilq09pM9Gl-L86TffB4g,1062
@@ -18,20 +18,20 @@ edata/models/bill.py,sha256=-qZ4l9f8BnPFlykFRoNYnmRe6-ELWmuNh6usks_k4X0,2674
18
18
  edata/models/data.py,sha256=O1aj0H_G92mnJjm6mYhC8fnl8ZOoatGiokRf8b7ziD8,2388
19
19
  edata/models/supply.py,sha256=qM0Q22-7PfBN93-Tu3ljAE58p5gfob0YWw-o5FccjrI,878
20
20
  edata/providers/__init__.py,sha256=5xaRKPttYM6EpVE1I_vhUAF25L9pEM_qK_5Ix6xnQe0,104
21
- edata/providers/datadis.py,sha256=VZ_MhO3x1crSmOH2dn3SWJ5Eq1gWUIDCEvB9oRRMZfQ,21468
21
+ edata/providers/datadis.py,sha256=3oSMWpiSHa111X1e0GUkOgnANoFvoUHeGwJN8pduW1Q,21968
22
22
  edata/providers/redata.py,sha256=i6aBjunyZMb34SGgAgNQvaNu_GFoe5P2rUIY-_vOElw,3727
23
23
  edata/services/bill_service.py,sha256=KAM78u7TpoVF-KTlB7bqj2Kp-rJERiPPoGTJCc-7rM0,15264
24
24
  edata/services/data_service.py,sha256=JJCdivjrwNt3aLhz3uMBWv06y-EL7UiK2G5EBr97l_8,19374
25
25
  edata/tests/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
26
26
  edata/tests/test_completion.py,sha256=-iyVt-juU3U1HqMuLUJXJthil2-n_4ddWiJq2GwnXTo,3551
27
- edata/tests/test_controller.py,sha256=YvilJ9cIpjhdGGtVJtlEevhRL7IHAVnU0At7AmxM5_8,3983
28
- edata/tests/test_datadis_connector.py,sha256=8Xet2M2eG-kMhiRDTCSH0R5nYO2oIpQTPOuTZKMKoLo,12538
29
- edata/tests/test_incremental.py,sha256=dn1LRgjBi2HuNYELPKDVdZQlS01hk4lsP9L-vyvtEHo,3479
30
- edata/tests/test_migrations.py,sha256=3Io2nokZVdv2xU9fCZCMIvwAwPVtGc0UivClSouP3pw,6418
27
+ edata/tests/test_controller.py,sha256=CZmT6TC1vfQsxfBNRq1sKKwD4kIWDWW6OYTYeIpLKpc,4147
28
+ edata/tests/test_datadis_connector.py,sha256=yh-cYX6zIU3HaTPimD5Rvm_QiPWJWBIskjKp2M5LpO8,12579
29
+ edata/tests/test_incremental.py,sha256=MPRc7pPSIUuBSMFzlVTpV_wafeRat4kV3wy8p6NmyFQ,3280
30
+ edata/tests/test_migrations.py,sha256=P8Qaas8u1Nwu2qNLSoqUAGL0WFqx9d59gs4Bp6L8RtY,6219
31
31
  edata/tests/test_redata_connector.py,sha256=SBvxsma5t-qvz2drxNisNM2P1STTurqEs1TSDgwaPJ8,508
32
- edata/tests/test_services.py,sha256=Ya5cbADHj1Qsffb_IuMeDD0BViz8nqrgwazUWK0-hRc,8297
32
+ edata/tests/test_services.py,sha256=DiXA8Y9zZgtAs9HGiMOfystz2shaGAGEs444GJ2AymE,8321
33
33
  edata/tests/test_utils.py,sha256=h6yBzT7j3eA42zbXJ-AK9lXcZkjCMf16wnob1oqGvUE,2307
34
- e_data-2.0.2.dev137.dist-info/METADATA,sha256=-7nOw3zSVBn4BVOgnmEXPt44pUE_qauv2ZZ2yjDjXt4,48084
35
- e_data-2.0.2.dev137.dist-info/WHEEL,sha256=YVMoNqKzERt-wjUZwJ33xBGAwnFl-4cqbYkTtWa4itE,91
36
- e_data-2.0.2.dev137.dist-info/top_level.txt,sha256=Ez-fReWtUVTMcFuH0dzCnoKIgMBHgjLaBuba8QiKC7g,6
37
- e_data-2.0.2.dev137.dist-info/RECORD,,
34
+ e_data-2.0.2.dev138.dist-info/METADATA,sha256=-DmL66kwdc79X1tIqaJYUATBtnLaFEfXPJjuCNlq7oo,48041
35
+ e_data-2.0.2.dev138.dist-info/WHEEL,sha256=YVMoNqKzERt-wjUZwJ33xBGAwnFl-4cqbYkTtWa4itE,91
36
+ e_data-2.0.2.dev138.dist-info/top_level.txt,sha256=Ez-fReWtUVTMcFuH0dzCnoKIgMBHgjLaBuba8QiKC7g,6
37
+ e_data-2.0.2.dev138.dist-info/RECORD,,
@@ -1,15 +1,15 @@
1
1
  import asyncio
2
+ import functools
2
3
  import logging
3
4
  import os
4
5
  import typing
6
+ from concurrent.futures import ThreadPoolExecutor
5
7
  from datetime import datetime
6
8
 
7
- from sqlalchemy import Select, Table, event, insert, or_
9
+ from sqlalchemy import Engine, Select, Table, create_engine, event, insert, or_
8
10
  from sqlalchemy.dialects.sqlite import insert as sqlite_insert
9
11
  from sqlalchemy.exc import IntegrityError
10
- from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
11
- from sqlmodel import SQLModel, UniqueConstraint
12
- from sqlmodel.ext.asyncio.session import AsyncSession
12
+ from sqlmodel import Session, SQLModel, UniqueConstraint
13
13
  from sqlmodel.sql.expression import SelectOfScalar
14
14
 
15
15
  import edata.database.queries as q
@@ -28,6 +28,7 @@ from edata.models.bill import Bill, EnergyPrice
28
28
  _LOGGER = logging.getLogger(__name__)
29
29
 
30
30
  T = typing.TypeVar("T", bound=SQLModel)
31
+ R = typing.TypeVar("R")
31
32
 
32
33
 
33
34
  def _set_sqlite_pragmas(dbapi_connection, connection_record) -> None:
@@ -77,53 +78,91 @@ def _conflict_columns(table: Table) -> list[str]:
77
78
  return [c.name for c in table.primary_key.columns]
78
79
 
79
80
 
81
+ def _in_db_thread(
82
+ fn: typing.Callable[..., R],
83
+ ) -> typing.Callable[..., typing.Coroutine[typing.Any, typing.Any, R]]:
84
+ """Turn a blocking ``EdataDB`` method into a coroutine run on its DB thread."""
85
+
86
+ @functools.wraps(fn)
87
+ async def wrapper(self: "EdataDB", *args: typing.Any, **kwargs: typing.Any) -> R:
88
+ loop = asyncio.get_running_loop()
89
+ return await loop.run_in_executor(
90
+ self._executor, functools.partial(fn, self, *args, **kwargs)
91
+ )
92
+
93
+ return wrapper
94
+
95
+
80
96
  class EdataDB:
97
+ """SQLite store.
98
+
99
+ The public API is async, but every database operation (SQLAlchemy, pydantic
100
+ (de)serialization and SQLite itself) runs on a single dedicated thread so a
101
+ large import never stalls the caller's event loop (e.g. Home Assistant's).
102
+ One thread also serializes access, which suits SQLite's single writer.
103
+ """
81
104
 
82
105
  _instance = None
83
- _engine: AsyncEngine | None = None
106
+ _engine: Engine | None = None
84
107
  _db_url: str | None = None
85
108
 
86
109
  def __new__(cls, sqlite_path: str):
87
- db_url = f"sqlite+aiosqlite:////{os.path.abspath(sqlite_path)}"
110
+ db_url = f"sqlite:////{os.path.abspath(sqlite_path)}"
88
111
  if cls._instance is None:
89
112
  cls._instance = super().__new__(cls)
90
113
  cls._db_url = db_url
91
114
  # Ensure parent directory exists before the first connection is opened.
92
115
  dir_path = os.path.dirname(os.path.abspath(sqlite_path))
93
116
  os.makedirs(dir_path, exist_ok=True)
94
- cls._engine = create_async_engine(db_url, future=True)
95
- event.listen(cls._engine.sync_engine, "connect", _set_sqlite_pragmas)
117
+ cls._engine = create_engine(db_url)
118
+ event.listen(cls._engine, "connect", _set_sqlite_pragmas)
96
119
  cls._instance._tables_initialized = False
97
- cls._instance._tables_lock = asyncio.Lock()
120
+ cls._instance._executor = ThreadPoolExecutor(
121
+ max_workers=1, thread_name_prefix="edata-db"
122
+ )
98
123
  elif db_url != cls._db_url:
99
124
  raise ValueError("EdataDB already initialized with a different db_url")
100
125
  return cls._instance
101
126
 
127
+ @classmethod
128
+ def reset(cls) -> None:
129
+ """Close the shared instance (its DB thread and engine) and forget it.
130
+
131
+ Waits for queued database work to finish first. A later ``EdataDB(...)``
132
+ starts afresh, possibly on another path.
133
+ """
134
+
135
+ if cls._instance is not None:
136
+ cls._instance._executor.shutdown(wait=True)
137
+ if cls._engine is not None:
138
+ cls._engine.dispose()
139
+ cls._instance = None
140
+ cls._engine = None
141
+ cls._db_url = None
142
+
102
143
  @property
103
- def engine(self) -> AsyncEngine | None:
104
- """Return the async database engine."""
144
+ def engine(self) -> Engine | None:
145
+ """Return the database engine."""
105
146
 
106
147
  return self._engine
107
148
 
108
- async def _ensure_tables(self) -> None:
109
- """Create tables if not already created (lazy init)."""
149
+ def _ensure_tables(self) -> None:
150
+ """Create tables and missing indexes if not already done (lazy init).
110
151
 
111
- if self._tables_initialized:
152
+ Only ever runs on the database thread, so concurrent first calls are
153
+ serialized and cannot race to create the same index.
154
+ """
155
+
156
+ if self._tables_initialized or not self.engine:
112
157
  return
113
- # concurrent first calls (e.g. several websocket requests right after
114
- # startup) would otherwise all see the index missing and race to create
115
- # it, failing with "index ... already exists"
116
- async with self._tables_lock:
117
- if self._tables_initialized or not self.engine:
118
- return
119
- async with self.engine.begin() as conn:
120
- await conn.run_sync(SQLModel.metadata.create_all)
121
- await conn.run_sync(_create_missing_indexes)
122
- self._tables_initialized = True
123
-
124
- async def _add_one(
158
+ with self.engine.begin() as conn:
159
+ SQLModel.metadata.create_all(conn)
160
+ _create_missing_indexes(conn)
161
+ self._tables_initialized = True
162
+
163
+ def _add_one(
125
164
  self,
126
- session: AsyncSession,
165
+ session: Session,
127
166
  record: T,
128
167
  commit: bool = True,
129
168
  ) -> T:
@@ -131,15 +170,15 @@ class EdataDB:
131
170
 
132
171
  session.add(record)
133
172
  if commit:
134
- await session.commit()
135
- await session.refresh(record)
173
+ session.commit()
174
+ session.refresh(record)
136
175
  else:
137
- await session.flush()
176
+ session.flush()
138
177
  return record
139
178
 
140
- async def _update_one(
179
+ def _update_one(
141
180
  self,
142
- session: AsyncSession,
181
+ session: Session,
143
182
  query: SelectOfScalar,
144
183
  data: typing.Any,
145
184
  commit: bool = True,
@@ -147,7 +186,7 @@ class EdataDB:
147
186
  ) -> T | None: # type: ignore
148
187
  """Update a single record in the database."""
149
188
 
150
- result = await session.exec(query)
189
+ result = session.exec(query)
151
190
  existing = result.first()
152
191
  if existing and getattr(existing, "data") == data and not overrides:
153
192
  return existing
@@ -156,15 +195,15 @@ class EdataDB:
156
195
  for key, value in overrides.items():
157
196
  setattr(existing, key, value)
158
197
  if commit:
159
- await session.commit()
160
- await session.refresh(existing)
198
+ session.commit()
199
+ session.refresh(existing)
161
200
  else:
162
- await session.flush()
201
+ session.flush()
163
202
  return existing
164
203
 
165
- async def _add_or_update_one(
204
+ def _add_or_update_one(
166
205
  self,
167
- session: AsyncSession,
206
+ session: Session,
168
207
  query: SelectOfScalar,
169
208
  record: T,
170
209
  commit: bool = True,
@@ -173,12 +212,12 @@ class EdataDB:
173
212
  """Add a single record into the database and fallback to update safely."""
174
213
 
175
214
  try:
176
- async with session.begin_nested():
215
+ with session.begin_nested():
177
216
  session.add(record)
178
- await session.flush()
217
+ session.flush()
179
218
  if commit:
180
- await session.commit()
181
- await session.refresh(record)
219
+ session.commit()
220
+ session.refresh(record)
182
221
  return record
183
222
  except IntegrityError:
184
223
  new_data = getattr(record, "data")
@@ -188,13 +227,13 @@ class EdataDB:
188
227
  override_dict = {
189
228
  x: record_json[x] for x in record.model_dump() if x in override
190
229
  }
191
- return await self._update_one(
230
+ return self._update_one(
192
231
  session, query, new_data, commit=commit, overrides=override_dict
193
232
  )
194
233
 
195
- async def _add_or_update_many(
234
+ def _add_or_update_many(
196
235
  self,
197
- session: AsyncSession,
236
+ session: Session,
198
237
  model: type[SQLModel],
199
238
  rows: list[dict[str, typing.Any]],
200
239
  override: list[str] | None = None,
@@ -232,166 +271,183 @@ class EdataDB:
232
271
  },
233
272
  where=or_(*(table.c[c].is_distinct_from(stmt.excluded[c]) for c in updated)),
234
273
  )
235
- connection = await session.connection()
236
- await connection.execute(stmt, rows)
237
- await session.commit()
274
+ connection = session.connection()
275
+ connection.execute(stmt, rows)
276
+ session.commit()
238
277
 
239
- async def get_supply(self, cups: str) -> SupplyModel | None:
278
+ @_in_db_thread
279
+ def get_supply(self, cups: str) -> SupplyModel | None:
240
280
  """Get a supply record by cups."""
241
281
 
242
- await self._ensure_tables()
243
- async with AsyncSession(self.engine) as session:
244
- result = await session.exec(q.get_supply(cups))
282
+ self._ensure_tables()
283
+ with Session(self.engine) as session:
284
+ result = session.exec(q.get_supply(cups))
245
285
  return result.first()
246
286
 
247
- async def get_contract(
287
+ @_in_db_thread
288
+ def get_contract(
248
289
  self, cups: str, date_start: datetime | None = None
249
290
  ) -> ContractModel | None:
250
291
  """Get a contract record by cups."""
251
292
 
252
- await self._ensure_tables()
253
- async with AsyncSession(self.engine) as session:
254
- result = await session.exec(q.get_contract(cups, date_start))
293
+ self._ensure_tables()
294
+ with Session(self.engine) as session:
295
+ result = session.exec(q.get_contract(cups, date_start))
255
296
  return result.first()
256
297
 
257
- async def get_last_energy(self, cups: str) -> EnergyModel | None:
298
+ @_in_db_thread
299
+ def get_last_energy(self, cups: str) -> EnergyModel | None:
258
300
  """Get the most recent Energy record by cups."""
259
301
 
260
- await self._ensure_tables()
261
- async with AsyncSession(self.engine) as session:
262
- result = await session.exec(q.get_last_energy(cups))
302
+ self._ensure_tables()
303
+ with Session(self.engine) as session:
304
+ result = session.exec(q.get_last_energy(cups))
263
305
  return result.first()
264
306
 
265
- async def get_last_power(self, cups: str) -> PowerModel | None:
307
+ @_in_db_thread
308
+ def get_last_power(self, cups: str) -> PowerModel | None:
266
309
  """Get the most recent power record by cups."""
267
310
 
268
- await self._ensure_tables()
269
- async with AsyncSession(self.engine) as session:
270
- result = await session.exec(q.get_last_power(cups))
311
+ self._ensure_tables()
312
+ with Session(self.engine) as session:
313
+ result = session.exec(q.get_last_power(cups))
271
314
  return result.first()
272
315
 
273
- async def get_last_pvpc(self) -> PVPCModel | None:
316
+ @_in_db_thread
317
+ def get_last_pvpc(self) -> PVPCModel | None:
274
318
  """Get the most recent pvpc."""
275
319
 
276
- await self._ensure_tables()
277
- async with AsyncSession(self.engine) as session:
278
- result = await session.exec(q.get_last_pvpc())
320
+ self._ensure_tables()
321
+ with Session(self.engine) as session:
322
+ result = session.exec(q.get_last_pvpc())
279
323
  return result.first()
280
324
 
281
- async def get_last_bill(self, cups: str) -> BillModel | None:
325
+ @_in_db_thread
326
+ def get_last_bill(self, cups: str) -> BillModel | None:
282
327
  """Get the most recent bill record by cups."""
283
328
 
284
- await self._ensure_tables()
285
- async with AsyncSession(self.engine) as session:
286
- result = await session.exec(q.get_last_bill(cups))
329
+ self._ensure_tables()
330
+ with Session(self.engine) as session:
331
+ result = session.exec(q.get_last_bill(cups))
287
332
  return result.first()
288
333
 
289
- async def get_last_complete_statistic(
334
+ @_in_db_thread
335
+ def get_last_complete_statistic(
290
336
  self, cups: str, type_: typing.Literal["day", "month"]
291
337
  ) -> StatisticsModel | None:
292
338
  """Get the most recent complete statistics record by cups and type."""
293
339
 
294
- await self._ensure_tables()
295
- async with AsyncSession(self.engine) as session:
296
- result = await session.exec(q.get_last_complete_statistic(cups, type_))
340
+ self._ensure_tables()
341
+ with Session(self.engine) as session:
342
+ result = session.exec(q.get_last_complete_statistic(cups, type_))
297
343
  return result.first()
298
344
 
299
- async def get_last_complete_bill(
345
+ @_in_db_thread
346
+ def get_last_complete_bill(
300
347
  self, cups: str, type_: typing.Literal["hour", "day", "month"]
301
348
  ) -> BillModel | None:
302
349
  """Get the most recent complete bill record by cups and type."""
303
350
 
304
- await self._ensure_tables()
305
- async with AsyncSession(self.engine) as session:
306
- result = await session.exec(q.get_last_complete_bill(cups, type_))
351
+ self._ensure_tables()
352
+ with Session(self.engine) as session:
353
+ result = session.exec(q.get_last_complete_bill(cups, type_))
307
354
  return result.first()
308
355
 
309
- async def add_contract(self, cups: str, contract: Contract) -> ContractModel | None:
356
+ @_in_db_thread
357
+ def add_contract(self, cups: str, contract: Contract) -> ContractModel | None:
310
358
  """Add or update a contract record."""
311
359
 
312
- await self._ensure_tables()
360
+ self._ensure_tables()
313
361
  record = ContractModel(cups=cups, date_start=contract.date_start, data=contract)
314
- async with AsyncSession(self.engine) as session:
315
- return await self._add_or_update_one(
362
+ with Session(self.engine) as session:
363
+ return self._add_or_update_one(
316
364
  session, q.get_contract(cups, contract.date_start), record
317
365
  )
318
366
 
319
- async def add_supply(self, supply: Supply) -> SupplyModel | None:
367
+ @_in_db_thread
368
+ def add_supply(self, supply: Supply) -> SupplyModel | None:
320
369
  """Add or update a supply record."""
321
370
 
322
- await self._ensure_tables()
371
+ self._ensure_tables()
323
372
  record = SupplyModel(cups=supply.cups, data=supply)
324
- async with AsyncSession(self.engine) as session:
325
- return await self._add_or_update_one(
373
+ with Session(self.engine) as session:
374
+ return self._add_or_update_one(
326
375
  session, q.get_supply(supply.cups), record
327
376
  )
328
377
 
329
- async def add_energy(self, cups: str, energy: Energy) -> EnergyModel | None:
378
+ @_in_db_thread
379
+ def add_energy(self, cups: str, energy: Energy) -> EnergyModel | None:
330
380
  """Add or update an energy record."""
331
381
 
332
- await self._ensure_tables()
382
+ self._ensure_tables()
333
383
  record = EnergyModel(
334
384
  cups=cups, delta_h=energy.delta_h, datetime=energy.datetime, data=energy
335
385
  )
336
- async with AsyncSession(self.engine) as session:
337
- return await self._add_or_update_one(
386
+ with Session(self.engine) as session:
387
+ return self._add_or_update_one(
338
388
  session, q.get_energy(cups, energy.datetime), record
339
389
  )
340
390
 
341
- async def add_energy_list(self, cups: str, energy: list[Energy]) -> None:
391
+ @_in_db_thread
392
+ def add_energy_list(self, cups: str, energy: list[Energy]) -> None:
342
393
  """Add or update a list of energy records."""
343
394
 
344
- await self._ensure_tables()
345
- async with AsyncSession(self.engine) as session:
395
+ self._ensure_tables()
396
+ with Session(self.engine) as session:
346
397
  unique_map = {item.datetime: item for item in energy}
347
398
  unique = list(unique_map.values())
348
399
  rows = [
349
400
  {"cups": cups, "delta_h": x.delta_h, "datetime": x.datetime, "data": x}
350
401
  for x in unique
351
402
  ]
352
- await self._add_or_update_many(session, EnergyModel, rows)
403
+ self._add_or_update_many(session, EnergyModel, rows)
353
404
 
354
- async def add_power(self, cups: str, power: Power) -> PowerModel | None:
405
+ @_in_db_thread
406
+ def add_power(self, cups: str, power: Power) -> PowerModel | None:
355
407
  """Add or update a power record for a given CUPS and Power instance."""
356
408
 
357
- await self._ensure_tables()
409
+ self._ensure_tables()
358
410
  record = PowerModel(cups=cups, datetime=power.datetime, data=power)
359
- async with AsyncSession(self.engine) as session:
360
- return await self._add_or_update_one(
411
+ with Session(self.engine) as session:
412
+ return self._add_or_update_one(
361
413
  session, q.get_power(cups, power.datetime), record
362
414
  )
363
415
 
364
- async def add_power_list(self, cups: str, power: list[Power]) -> None:
416
+ @_in_db_thread
417
+ def add_power_list(self, cups: str, power: list[Power]) -> None:
365
418
  """Add or update a list of power records."""
366
419
 
367
- await self._ensure_tables()
368
- async with AsyncSession(self.engine) as session:
420
+ self._ensure_tables()
421
+ with Session(self.engine) as session:
369
422
  unique_map = {item.datetime: item for item in power}
370
423
  unique = list(unique_map.values())
371
424
  rows = [{"cups": cups, "datetime": x.datetime, "data": x} for x in unique]
372
- await self._add_or_update_many(session, PowerModel, rows)
425
+ self._add_or_update_many(session, PowerModel, rows)
373
426
 
374
- async def add_pvpc(self, pvpc: EnergyPrice) -> PVPCModel | None:
427
+ @_in_db_thread
428
+ def add_pvpc(self, pvpc: EnergyPrice) -> PVPCModel | None:
375
429
  """Add or update a pvpc record."""
376
430
 
377
- await self._ensure_tables()
431
+ self._ensure_tables()
378
432
  record = PVPCModel(datetime=pvpc.datetime, data=pvpc)
379
- async with AsyncSession(self.engine) as session:
380
- return await self._add_or_update_one(
433
+ with Session(self.engine) as session:
434
+ return self._add_or_update_one(
381
435
  session, q.get_pvpc(pvpc.datetime), record
382
436
  )
383
437
 
384
- async def add_pvpc_list(self, pvpc: list[EnergyPrice]) -> None:
438
+ @_in_db_thread
439
+ def add_pvpc_list(self, pvpc: list[EnergyPrice]) -> None:
385
440
  """Add or update a list of pvpc records."""
386
441
 
387
- await self._ensure_tables()
388
- async with AsyncSession(self.engine) as session:
442
+ self._ensure_tables()
443
+ with Session(self.engine) as session:
389
444
  unique_map = {item.datetime: item for item in pvpc}
390
445
  unique = list(unique_map.values())
391
446
  rows = [{"datetime": x.datetime, "data": x} for x in unique]
392
- await self._add_or_update_many(session, PVPCModel, rows)
447
+ self._add_or_update_many(session, PVPCModel, rows)
393
448
 
394
- async def add_statistics(
449
+ @_in_db_thread
450
+ def add_statistics(
395
451
  self,
396
452
  cups: str,
397
453
  type_: typing.Literal["day", "month"],
@@ -400,19 +456,20 @@ class EdataDB:
400
456
  ) -> StatisticsModel | None:
401
457
  """Add or update a statistics record."""
402
458
 
403
- await self._ensure_tables()
459
+ self._ensure_tables()
404
460
  record = StatisticsModel(
405
461
  cups=cups, datetime=data.datetime, type=type_, data=data, complete=complete
406
462
  )
407
- async with AsyncSession(self.engine) as session:
408
- return await self._add_or_update_one(
463
+ with Session(self.engine) as session:
464
+ return self._add_or_update_one(
409
465
  session,
410
466
  q.get_statistics(cups, type_, data.datetime),
411
467
  record,
412
468
  override=["complete"],
413
469
  )
414
470
 
415
- async def add_statistics_list(
471
+ @_in_db_thread
472
+ def add_statistics_list(
416
473
  self,
417
474
  cups: str,
418
475
  type_: typing.Literal["day", "month"],
@@ -421,8 +478,8 @@ class EdataDB:
421
478
  ) -> None:
422
479
  """Add or update a list of statistics records."""
423
480
 
424
- await self._ensure_tables()
425
- async with AsyncSession(self.engine) as session:
481
+ self._ensure_tables()
482
+ with Session(self.engine) as session:
426
483
  unique_map = {item.datetime: item for item in statistics}
427
484
  rows = [
428
485
  {
@@ -434,11 +491,12 @@ class EdataDB:
434
491
  }
435
492
  for x in unique_map.values()
436
493
  ]
437
- await self._add_or_update_many(
494
+ self._add_or_update_many(
438
495
  session, StatisticsModel, rows, override=["complete"]
439
496
  )
440
497
 
441
- async def add_bill(
498
+ @_in_db_thread
499
+ def add_bill(
442
500
  self,
443
501
  cups: str,
444
502
  type_: typing.Literal["hour", "day", "month"],
@@ -448,7 +506,7 @@ class EdataDB:
448
506
  ) -> BillModel | None:
449
507
  """Add or update a bill record."""
450
508
 
451
- await self._ensure_tables()
509
+ self._ensure_tables()
452
510
  record = BillModel(
453
511
  cups=cups,
454
512
  datetime=data.datetime,
@@ -457,15 +515,16 @@ class EdataDB:
457
515
  confhash=confhash,
458
516
  data=data,
459
517
  )
460
- async with AsyncSession(self.engine) as session:
461
- return await self._add_or_update_one(
518
+ with Session(self.engine) as session:
519
+ return self._add_or_update_one(
462
520
  session,
463
521
  q.get_bill(cups, type_, data.datetime),
464
522
  record,
465
523
  override=["complete", "confhash"],
466
524
  )
467
525
 
468
- async def add_bill_list(
526
+ @_in_db_thread
527
+ def add_bill_list(
469
528
  self,
470
529
  cups: str,
471
530
  type_: typing.Literal["hour", "day", "month"],
@@ -475,8 +534,8 @@ class EdataDB:
475
534
  ) -> None:
476
535
  """Add or update a list of bill records."""
477
536
 
478
- await self._ensure_tables()
479
- async with AsyncSession(self.engine) as session:
537
+ self._ensure_tables()
538
+ with Session(self.engine) as session:
480
539
  unique_map = {item.datetime: item for item in bill}
481
540
  unique = list(unique_map.values())
482
541
  rows = [
@@ -490,51 +549,55 @@ class EdataDB:
490
549
  }
491
550
  for x in unique
492
551
  ]
493
- await self._add_or_update_many(
552
+ self._add_or_update_many(
494
553
  session, BillModel, rows, override=["complete", "confhash"]
495
554
  )
496
555
 
497
- async def clear_bills(self, cups: str, since: datetime | None = None) -> None:
556
+ @_in_db_thread
557
+ def clear_bills(self, cups: str, since: datetime | None = None) -> None:
498
558
  """Delete bill records for a cups, optionally only from a datetime onwards."""
499
559
 
500
- await self._ensure_tables()
501
- async with AsyncSession(self.engine) as session:
502
- await session.exec(q.delete_bill(cups, since)) # type: ignore[call-overload]
503
- await session.commit()
560
+ self._ensure_tables()
561
+ with Session(self.engine) as session:
562
+ session.exec(q.delete_bill(cups, since)) # type: ignore[call-overload]
563
+ session.commit()
504
564
 
505
- async def list_supplies(self) -> typing.Sequence[SupplyModel]:
565
+ @_in_db_thread
566
+ def list_supplies(self) -> typing.Sequence[SupplyModel]:
506
567
  """List all supply records."""
507
568
 
508
- await self._ensure_tables()
509
- async with AsyncSession(self.engine) as session:
510
- result = await session.exec(q.list_supply())
569
+ self._ensure_tables()
570
+ with Session(self.engine) as session:
571
+ result = session.exec(q.list_supply())
511
572
  return result.all()
512
573
 
513
- async def list_contracts(
574
+ @_in_db_thread
575
+ def list_contracts(
514
576
  self, cups: str | None = None
515
577
  ) -> typing.Sequence[ContractModel]:
516
578
  """List all contract records."""
517
579
 
518
- await self._ensure_tables()
519
- async with AsyncSession(self.engine) as session:
520
- result = await session.exec(q.list_contract(cups))
580
+ self._ensure_tables()
581
+ with Session(self.engine) as session:
582
+ result = session.exec(q.list_contract(cups))
521
583
  return result.all()
522
584
 
523
- async def _list_data(self, query: SelectOfScalar, model: type[SQLModel]) -> list:
585
+ def _list_data(self, query: SelectOfScalar, model: type[SQLModel]) -> list:
524
586
  """Return only the ``data`` payload of the rows selected by ``query``.
525
587
 
526
588
  Skips building an ORM instance per row, which nearly halves the cost of
527
589
  reading a month of hourly records.
528
590
  """
529
591
 
530
- await self._ensure_tables()
531
- async with AsyncSession(self.engine) as session:
532
- result = await session.exec(
592
+ self._ensure_tables()
593
+ with Session(self.engine) as session:
594
+ result = session.exec(
533
595
  query.with_only_columns(model.data) # type: ignore[attr-defined]
534
596
  )
535
597
  return list(result.all())
536
598
 
537
- async def list_energy_data(
599
+ @_in_db_thread
600
+ def list_energy_data(
538
601
  self,
539
602
  cups: str,
540
603
  date_from: datetime | None = None,
@@ -542,20 +605,22 @@ class EdataDB:
542
605
  ) -> list[Energy]:
543
606
  """List the energy data (without row metadata)."""
544
607
 
545
- return await self._list_data(
608
+ return self._list_data(
546
609
  q.list_energy(cups, date_from, date_to), EnergyModel
547
610
  )
548
611
 
549
- async def list_pvpc_data(
612
+ @_in_db_thread
613
+ def list_pvpc_data(
550
614
  self,
551
615
  date_from: datetime | None = None,
552
616
  date_to: datetime | None = None,
553
617
  ) -> list[EnergyPrice]:
554
618
  """List the pvpc data (without row metadata)."""
555
619
 
556
- return await self._list_data(q.list_pvpc(date_from, date_to), PVPCModel)
620
+ return self._list_data(q.list_pvpc(date_from, date_to), PVPCModel)
557
621
 
558
- async def list_bill_data(
622
+ @_in_db_thread
623
+ def list_bill_data(
559
624
  self,
560
625
  cups: str,
561
626
  type_: typing.Literal["hour", "day", "month"],
@@ -564,11 +629,12 @@ class EdataDB:
564
629
  ) -> list[Bill]:
565
630
  """List the bill data (without row metadata)."""
566
631
 
567
- return await self._list_data(
632
+ return self._list_data(
568
633
  q.list_bill(cups, type_, date_from, date_to), BillModel
569
634
  )
570
635
 
571
- async def list_energy(
636
+ @_in_db_thread
637
+ def list_energy(
572
638
  self,
573
639
  cups: str,
574
640
  date_from: datetime | None = None,
@@ -576,12 +642,13 @@ class EdataDB:
576
642
  ) -> typing.Sequence[EnergyModel]:
577
643
  """List energy records."""
578
644
 
579
- await self._ensure_tables()
580
- async with AsyncSession(self.engine) as session:
581
- result = await session.exec(q.list_energy(cups, date_from, date_to))
645
+ self._ensure_tables()
646
+ with Session(self.engine) as session:
647
+ result = session.exec(q.list_energy(cups, date_from, date_to))
582
648
  return result.all()
583
649
 
584
- async def list_power(
650
+ @_in_db_thread
651
+ def list_power(
585
652
  self,
586
653
  cups: str,
587
654
  date_from: datetime | None = None,
@@ -589,24 +656,26 @@ class EdataDB:
589
656
  ) -> typing.Sequence[PowerModel]:
590
657
  """List power records."""
591
658
 
592
- await self._ensure_tables()
593
- async with AsyncSession(self.engine) as session:
594
- result = await session.exec(q.list_power(cups, date_from, date_to))
659
+ self._ensure_tables()
660
+ with Session(self.engine) as session:
661
+ result = session.exec(q.list_power(cups, date_from, date_to))
595
662
  return result.all()
596
663
 
597
- async def list_pvpc(
664
+ @_in_db_thread
665
+ def list_pvpc(
598
666
  self,
599
667
  date_from: datetime | None = None,
600
668
  date_to: datetime | None = None,
601
669
  ) -> typing.Sequence[PVPCModel]:
602
670
  """List pvpc records."""
603
671
 
604
- await self._ensure_tables()
605
- async with AsyncSession(self.engine) as session:
606
- result = await session.exec(q.list_pvpc(date_from, date_to))
672
+ self._ensure_tables()
673
+ with Session(self.engine) as session:
674
+ result = session.exec(q.list_pvpc(date_from, date_to))
607
675
  return result.all()
608
676
 
609
- async def list_statistics(
677
+ @_in_db_thread
678
+ def list_statistics(
610
679
  self,
611
680
  cups: str,
612
681
  type_: typing.Literal["day", "month"],
@@ -616,14 +685,15 @@ class EdataDB:
616
685
  ) -> typing.Sequence[StatisticsModel]:
617
686
  """List statistics records filtered by type ('day' or 'month') and date range."""
618
687
 
619
- await self._ensure_tables()
620
- async with AsyncSession(self.engine) as session:
621
- result = await session.exec(
688
+ self._ensure_tables()
689
+ with Session(self.engine) as session:
690
+ result = session.exec(
622
691
  q.list_statistics(cups, type_, date_from, date_to, complete)
623
692
  )
624
693
  return result.all()
625
694
 
626
- async def list_bill(
695
+ @_in_db_thread
696
+ def list_bill(
627
697
  self,
628
698
  cups: str,
629
699
  type_: typing.Literal["hour", "day", "month"],
@@ -633,9 +703,9 @@ class EdataDB:
633
703
  ) -> typing.Sequence[BillModel]:
634
704
  """List bill records filtered by type ('hour', 'day' or 'month') and date range."""
635
705
 
636
- await self._ensure_tables()
637
- async with AsyncSession(self.engine) as session:
638
- result = await session.exec(
706
+ self._ensure_tables()
707
+ with Session(self.engine) as session:
708
+ result = session.exec(
639
709
  q.list_bill(cups, type_, date_from, date_to, complete)
640
710
  )
641
711
  return result.all()
@@ -3,6 +3,7 @@
3
3
  import asyncio
4
4
  import contextlib
5
5
  import hashlib
6
+ import json
6
7
  import logging
7
8
  import os
8
9
  import tempfile
@@ -69,6 +70,59 @@ def migrate_storage(storage_dir: str) -> None:
69
70
  os.remove(os.path.join(storage_dir, "edata_recent_queries_cache.json"))
70
71
 
71
72
 
73
+ def _parse_consumptions(
74
+ response: dict[str, typing.Any], start_date: datetime, end_date: datetime
75
+ ) -> list[Energy]:
76
+ """Build the energy records of a 'get_consumption_data' response."""
77
+
78
+ consumptions = []
79
+ for i in response.get("timeCurve", []):
80
+ if "consumptionKWh" in i:
81
+ if all(k in i for k in GET_CONSUMPTION_DATA_MANDATORY_FIELDS):
82
+ raw_hour = int(i["time"].split(":")[0])
83
+ # Datadis delivers hours 1..24 (end-of-interval). A sporadic
84
+ # i-DE glitch emits an extra "00:00" row on a day already
85
+ # carrying its full 24 hours, so drop it -- 01:00..24:00
86
+ # already covers the day -- instead of remapping onto 23:00
87
+ # and overlapping that day's 24:00 slot.
88
+ if raw_hour == 0:
89
+ continue
90
+ date_as_dt = datetime.strptime(
91
+ i["date"], "%Y/%m/%d"
92
+ ) + timedelta(hours=raw_hour - 1)
93
+ if not (start_date <= date_as_dt <= end_date):
94
+ continue # skip element if dt is out of range
95
+
96
+ # sanitize these values
97
+ _surplus_kwh = i.get("surplusEnergyKWh", 0)
98
+ if _surplus_kwh is None:
99
+ _surplus_kwh = 0
100
+ _generation_kwh = i.get("generationEnergyKWh", 0)
101
+ if _generation_kwh is None:
102
+ _generation_kwh = 0
103
+ _selfconsumption_kwh = i.get("selfConsumptionEnergyKWh", 0)
104
+ if _selfconsumption_kwh is None:
105
+ _selfconsumption_kwh = 0
106
+
107
+ consumptions.append(
108
+ Energy(
109
+ datetime=date_as_dt,
110
+ delta_h=1,
111
+ consumption_kwh=i["consumptionKWh"],
112
+ surplus_kwh=_surplus_kwh,
113
+ generation_kwh=_generation_kwh,
114
+ selfconsumption_kwh=_selfconsumption_kwh,
115
+ real=i["obtainMethod"] == "Real",
116
+ )
117
+ )
118
+ else:
119
+ _LOGGER.warning(
120
+ "Weird data structure while fetching consumption data, got %s",
121
+ response,
122
+ )
123
+ return consumptions
124
+
125
+
72
126
  class DatadisConnector:
73
127
  """A Datadis private API connector."""
74
128
 
@@ -237,7 +291,10 @@ class DatadisConnector:
237
291
  # it here as well would decode every (large) payload twice
238
292
  if reply.status == 200:
239
293
  try:
240
- json_data = await reply.json(content_type=None)
294
+ # decode off the event loop: a full-history
295
+ # response is several MB of JSON
296
+ body = await reply.read()
297
+ json_data = await asyncio.to_thread(json.loads, body)
241
298
  if json_data:
242
299
  response = json_data
243
300
  if not ignore_cache:
@@ -419,52 +476,11 @@ class DatadisConnector:
419
476
 
420
477
  response = await self._async_get(URL_GET_CONSUMPTION_DATA, request_data=data)
421
478
 
422
- consumptions = []
423
- for i in response.get("timeCurve", []):
424
- if "consumptionKWh" in i:
425
- if all(k in i for k in GET_CONSUMPTION_DATA_MANDATORY_FIELDS):
426
- raw_hour = int(i["time"].split(":")[0])
427
- # Datadis delivers hours 1..24 (end-of-interval). A sporadic
428
- # i-DE glitch emits an extra "00:00" row on a day already
429
- # carrying its full 24 hours, so drop it -- 01:00..24:00
430
- # already covers the day -- instead of remapping onto 23:00
431
- # and overlapping that day's 24:00 slot.
432
- if raw_hour == 0:
433
- continue
434
- date_as_dt = datetime.strptime(
435
- i["date"], "%Y/%m/%d"
436
- ) + timedelta(hours=raw_hour - 1)
437
- if not (start_date <= date_as_dt <= end_date):
438
- continue # skip element if dt is out of range
439
-
440
- # sanitize these values
441
- _surplus_kwh = i.get("surplusEnergyKWh", 0)
442
- if _surplus_kwh is None:
443
- _surplus_kwh = 0
444
- _generation_kwh = i.get("generationEnergyKWh", 0)
445
- if _generation_kwh is None:
446
- _generation_kwh = 0
447
- _selfconsumption_kwh = i.get("selfConsumptionEnergyKWh", 0)
448
- if _selfconsumption_kwh is None:
449
- _selfconsumption_kwh = 0
450
-
451
- consumptions.append(
452
- Energy(
453
- datetime=date_as_dt,
454
- delta_h=1,
455
- consumption_kwh=i["consumptionKWh"],
456
- surplus_kwh=_surplus_kwh,
457
- generation_kwh=_generation_kwh,
458
- selfconsumption_kwh=_selfconsumption_kwh,
459
- real=i["obtainMethod"] == "Real",
460
- )
461
- )
462
- else:
463
- _LOGGER.warning(
464
- "Weird data structure while fetching consumption data, got %s",
465
- response,
466
- )
467
- return consumptions
479
+ # building tens of thousands of records on a first import takes long
480
+ # enough to stall the event loop, so parse in a worker thread
481
+ return await asyncio.to_thread(
482
+ _parse_consumptions, response, start_date, end_date
483
+ )
468
484
 
469
485
  def get_consumption_data(
470
486
  self,
@@ -1,6 +1,7 @@
1
1
  """Tests for the database controller."""
2
2
 
3
3
  import asyncio
4
+ import threading
4
5
  from collections.abc import AsyncIterator
5
6
  from datetime import datetime, timedelta
6
7
 
@@ -16,16 +17,10 @@ CUPS = "ESXXXXXXXXXXXXXXXXTEST"
16
17
  START = datetime(2024, 1, 1)
17
18
 
18
19
 
19
- def _reset_singleton() -> None:
20
- EdataDB._instance = None
21
- EdataDB._engine = None
22
- EdataDB._db_url = None
23
-
24
-
25
20
  @pytest_asyncio.fixture
26
21
  async def db(tmp_path) -> AsyncIterator[EdataDB]:
27
22
  """An EdataDB on an isolated on-disk database with one supply."""
28
- _reset_singleton()
23
+ EdataDB.reset()
29
24
  database = EdataDB(str(tmp_path / "edata.db"))
30
25
  await database.add_supply(
31
26
  Supply(
@@ -42,9 +37,7 @@ async def db(tmp_path) -> AsyncIterator[EdataDB]:
42
37
  )
43
38
  )
44
39
  yield database
45
- if EdataDB._engine is not None:
46
- await EdataDB._engine.dispose()
47
- _reset_singleton()
40
+ EdataDB.reset()
48
41
 
49
42
 
50
43
  def _energy(hours: int, kwh: float) -> list[Energy]:
@@ -97,14 +90,14 @@ async def test_add_bill_list_applies_overrides(db: EdataDB) -> None:
97
90
  @pytest.mark.asyncio
98
91
  async def test_concurrent_first_calls_do_not_race_index_creation(db: EdataDB) -> None:
99
92
  # simulate a database created before the index existed, opened fresh
100
- async with db.engine.begin() as conn:
101
- await conn.exec_driver_sql("DROP INDEX ix_energy_cups_datetime")
93
+ with db.engine.begin() as conn:
94
+ conn.exec_driver_sql("DROP INDEX ix_energy_cups_datetime")
102
95
  db._tables_initialized = False
103
96
 
104
97
  await asyncio.gather(*(db.get_last_energy(CUPS) for _ in range(10)))
105
98
 
106
- async with db.engine.connect() as conn:
107
- result = await conn.exec_driver_sql(
99
+ with db.engine.connect() as conn:
100
+ result = conn.exec_driver_sql(
108
101
  "SELECT name FROM sqlite_master WHERE name='ix_energy_cups_datetime'"
109
102
  )
110
103
  assert result.first() is not None
@@ -121,3 +114,14 @@ def test_datetime_columns_are_naive() -> None:
121
114
  ]
122
115
  assert len(columns) == 20
123
116
  assert all(column.type.timezone is False for column in columns)
117
+
118
+
119
+ @pytest.mark.asyncio
120
+ async def test_reset_stops_db_thread_and_allows_new_path(db: EdataDB, tmp_path) -> None:
121
+ assert await db.list_supplies()
122
+ EdataDB.reset()
123
+
124
+ assert not any(t.name.startswith("edata-db") for t in threading.enumerate())
125
+ other = EdataDB(str(tmp_path / "other.db"))
126
+ assert other is not db
127
+ assert await other.list_supplies() == []
@@ -1,6 +1,7 @@
1
1
  """Tests for DatadisConnector (offline)."""
2
2
 
3
3
  import datetime
4
+ import json
4
5
  import os
5
6
  from unittest.mock import AsyncMock, MagicMock, patch
6
7
 
@@ -8,6 +9,12 @@ import pytest
8
9
 
9
10
  from edata.providers.datadis import DatadisConnector
10
11
 
12
+
13
+ def _json_body(payload) -> AsyncMock:
14
+ """Mock ``reply.read()`` returning ``payload`` as a JSON body."""
15
+ return AsyncMock(return_value=json.dumps(payload).encode())
16
+
17
+
11
18
  MOCK_USERNAME = "USERNAME"
12
19
  MOCK_PASSWORD = "PASSWORD"
13
20
 
@@ -126,7 +133,7 @@ def test_get_supplies(mock_token, mock_get, snapshot):
126
133
  mock_response = MagicMock()
127
134
  mock_response.status = 200
128
135
  mock_response.text = AsyncMock(return_value="text")
129
- mock_response.json = AsyncMock(return_value=SUPPLIES_RESPONSE)
136
+ mock_response.read = _json_body(SUPPLIES_RESPONSE)
130
137
  mock_get.return_value.__aenter__.return_value = mock_response
131
138
  connector = DatadisConnector(MOCK_USERNAME, MOCK_PASSWORD)
132
139
  assert connector.get_supplies() == snapshot
@@ -141,7 +148,7 @@ def test_get_contract_detail(mock_token, mock_get, snapshot):
141
148
  mock_response = MagicMock()
142
149
  mock_response.status = 200
143
150
  mock_response.text = AsyncMock(return_value="text")
144
- mock_response.json = AsyncMock(return_value=CONTRACTS_RESPONSE)
151
+ mock_response.read = _json_body(CONTRACTS_RESPONSE)
145
152
  mock_get.return_value.__aenter__.return_value = mock_response
146
153
  connector = DatadisConnector(MOCK_USERNAME, MOCK_PASSWORD)
147
154
  assert connector.get_contract_detail("ESXXXXXXXXXXXXXXXXTEST", "2") == snapshot
@@ -156,7 +163,7 @@ def test_get_consumption_data(mock_token, mock_get, snapshot):
156
163
  mock_response = MagicMock()
157
164
  mock_response.status = 200
158
165
  mock_response.text = AsyncMock(return_value="text")
159
- mock_response.json = AsyncMock(return_value=CONSUMPTIONS_RESPONSE)
166
+ mock_response.read = _json_body(CONSUMPTIONS_RESPONSE)
160
167
  mock_get.return_value.__aenter__.return_value = mock_response
161
168
  connector = DatadisConnector(MOCK_USERNAME, MOCK_PASSWORD)
162
169
  assert (
@@ -181,9 +188,7 @@ def test_get_consumption_data_skips_zero_hour(mock_token, mock_get):
181
188
  mock_response = MagicMock()
182
189
  mock_response.status = 200
183
190
  mock_response.text = AsyncMock(return_value="text")
184
- mock_response.json = AsyncMock(
185
- return_value=CONSUMPTIONS_RESPONSE_WITH_ZERO_HOUR
186
- )
191
+ mock_response.read = _json_body(CONSUMPTIONS_RESPONSE_WITH_ZERO_HOUR)
187
192
  mock_get.return_value.__aenter__.return_value = mock_response
188
193
  connector = DatadisConnector(MOCK_USERNAME, MOCK_PASSWORD)
189
194
 
@@ -215,7 +220,7 @@ def test_get_max_power(mock_token, mock_get, snapshot):
215
220
  mock_response = MagicMock()
216
221
  mock_response.status = 200
217
222
  mock_response.text = AsyncMock(return_value="text")
218
- mock_response.json = AsyncMock(return_value=MAXIMETER_RESPONSE)
223
+ mock_response.read = _json_body(MAXIMETER_RESPONSE)
219
224
  mock_get.return_value.__aenter__.return_value = mock_response
220
225
  connector = DatadisConnector(MOCK_USERNAME, MOCK_PASSWORD)
221
226
  assert (
@@ -239,7 +244,7 @@ def test_get_supplies_empty_response(mock_token, mock_get, snapshot):
239
244
  mock_response = MagicMock()
240
245
  mock_response.status = 200
241
246
  mock_response.text = AsyncMock(return_value="text")
242
- mock_response.json = AsyncMock(return_value={"supplies": []})
247
+ mock_response.read = _json_body({"supplies": []})
243
248
  mock_get.return_value.__aenter__.return_value = mock_response
244
249
  connector = DatadisConnector(MOCK_USERNAME, MOCK_PASSWORD)
245
250
  assert connector.get_supplies() == snapshot
@@ -255,7 +260,7 @@ def test_get_supplies_malformed_response(mock_token, mock_get, snapshot):
255
260
  mock_response = MagicMock()
256
261
  mock_response.status = 200
257
262
  mock_response.text = AsyncMock(return_value="text")
258
- mock_response.json = AsyncMock(return_value=malformed)
263
+ mock_response.read = _json_body(malformed)
259
264
  mock_get.return_value.__aenter__.return_value = mock_response
260
265
  connector = DatadisConnector(MOCK_USERNAME, MOCK_PASSWORD)
261
266
  assert connector.get_supplies() == snapshot
@@ -274,7 +279,7 @@ def test_get_supplies_partial_response(mock_token, mock_get, snapshot):
274
279
  mock_response = MagicMock()
275
280
  mock_response.status = 200
276
281
  mock_response.text = AsyncMock(return_value="text")
277
- mock_response.json = AsyncMock(return_value=partial)
282
+ mock_response.read = _json_body(partial)
278
283
  mock_get.return_value.__aenter__.return_value = mock_response
279
284
  connector = DatadisConnector(MOCK_USERNAME, MOCK_PASSWORD)
280
285
  assert connector.get_supplies() == snapshot
@@ -289,7 +294,7 @@ def test_get_consumption_data_cache(mock_token, mock_get, snapshot):
289
294
  mock_response = MagicMock()
290
295
  mock_response.status = 200
291
296
  mock_response.text = AsyncMock(return_value="text")
292
- mock_response.json = AsyncMock(return_value=CONSUMPTIONS_RESPONSE)
297
+ mock_response.read = _json_body(CONSUMPTIONS_RESPONSE)
293
298
  mock_get.return_value.__aenter__.return_value = mock_response
294
299
  connector = DatadisConnector(MOCK_USERNAME, MOCK_PASSWORD)
295
300
  # First call populates cache
@@ -345,7 +350,7 @@ def test_get_supplies_optional_fields_none(mock_token, mock_get, snapshot):
345
350
  mock_response = MagicMock()
346
351
  mock_response.status = 200
347
352
  mock_response.text = AsyncMock(return_value="text")
348
- mock_response.json = AsyncMock(return_value=response)
353
+ mock_response.read = _json_body(response)
349
354
  mock_get.return_value.__aenter__.return_value = mock_response
350
355
  connector = DatadisConnector(MOCK_USERNAME, MOCK_PASSWORD)
351
356
  assert connector.get_supplies() == snapshot
@@ -359,7 +364,7 @@ async def test_shared_session_is_reused(mock_token, tmp_path, snapshot):
359
364
  """A caller-provided session serves the requests and is left open."""
360
365
  mock_response = MagicMock()
361
366
  mock_response.status = 200
362
- mock_response.json = AsyncMock(return_value=SUPPLIES_RESPONSE)
367
+ mock_response.read = _json_body(SUPPLIES_RESPONSE)
363
368
  session = MagicMock()
364
369
  session.get.return_value.__aenter__.return_value = mock_response
365
370
  session.close = AsyncMock()
@@ -15,25 +15,17 @@ from edata.services.data_service import DataService
15
15
  CUPS = "ESXXXXXXXXXXXXXXXXTEST"
16
16
 
17
17
 
18
- def _reset_singleton() -> None:
19
- EdataDB._instance = None
20
- EdataDB._engine = None
21
- EdataDB._db_url = None
22
-
23
-
24
18
  @pytest_asyncio.fixture
25
19
  async def data_service(tmp_path) -> AsyncIterator[DataService]:
26
20
  """A DataService on an isolated on-disk database, singleton reset around it."""
27
- _reset_singleton()
21
+ EdataDB.reset()
28
22
  with (
29
23
  patch("edata.services.data_service.DatadisConnector"),
30
24
  patch("edata.services.data_service.REDataConnector"),
31
25
  ):
32
26
  service = DataService(CUPS, "user", "pwd", storage_path=str(tmp_path))
33
27
  yield service
34
- if EdataDB._engine is not None:
35
- await EdataDB._engine.dispose()
36
- _reset_singleton()
28
+ EdataDB.reset()
37
29
 
38
30
 
39
31
  def _hourly_energy(start: datetime, end: datetime) -> list[Energy]:
@@ -18,12 +18,6 @@ ASSETS = os.path.join(os.path.dirname(__file__), "assets")
18
18
  CUPS = "ESXXXXXXXXXXXXXXXXTEST"
19
19
 
20
20
 
21
- def _reset_singleton() -> None:
22
- EdataDB._instance = None
23
- EdataDB._engine = None
24
- EdataDB._db_url = None
25
-
26
-
27
21
  def _raw() -> dict:
28
22
  with open(os.path.join(ASSETS, "legacy_1.3.3.json"), encoding="utf-8") as f:
29
23
  return json.load(f)
@@ -40,16 +34,14 @@ def _install_legacy_file(storage_dir: str, cups: str = CUPS) -> str:
40
34
  @pytest_asyncio.fixture
41
35
  async def data_service(tmp_path) -> AsyncIterator[DataService]:
42
36
  """A DataService on an isolated on-disk database, singleton reset around it."""
43
- _reset_singleton()
37
+ EdataDB.reset()
44
38
  with (
45
39
  patch("edata.services.data_service.DatadisConnector"),
46
40
  patch("edata.services.data_service.REDataConnector"),
47
41
  ):
48
42
  service = DataService(CUPS, "user", "pwd", storage_path=str(tmp_path))
49
43
  yield service
50
- if EdataDB._engine is not None:
51
- await EdataDB._engine.dispose()
52
- _reset_singleton()
44
+ EdataDB.reset()
53
45
 
54
46
 
55
47
  # --- pure mappers ---
@@ -258,13 +258,13 @@ async def test_update_power_is_incremental(populated_data_service, power):
258
258
  async def test_missing_indexes_are_created_on_existing_db(populated_data_service):
259
259
  db = populated_data_service.db
260
260
 
261
- async with db.engine.begin() as conn:
262
- await conn.exec_driver_sql("DROP INDEX IF EXISTS ix_energy_cups_datetime")
261
+ with db.engine.begin() as conn:
262
+ conn.exec_driver_sql("DROP INDEX IF EXISTS ix_energy_cups_datetime")
263
263
  db._tables_initialized = False
264
- await db._ensure_tables()
264
+ await db.list_supplies() # any call lazily re-runs the table/index setup
265
265
 
266
- async with db.engine.connect() as conn:
267
- result = await conn.exec_driver_sql(
266
+ with db.engine.connect() as conn:
267
+ result = conn.exec_driver_sql(
268
268
  "SELECT name FROM sqlite_master WHERE type='index' AND tbl_name='energy'"
269
269
  )
270
270
  assert "ix_energy_cups_datetime" in {row[0] for row in result}