e-data 2.0.2.dev136__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.
edata/database/models.py CHANGED
@@ -1,6 +1,7 @@
1
1
  import typing
2
2
  from datetime import datetime as dt
3
3
 
4
+ from pydantic import NaiveDatetime
4
5
  from sqlmodel import AutoString, Column, Field, Index, SQLModel, UniqueConstraint
5
6
 
6
7
  from edata.database.utils import PydanticJSON
@@ -27,8 +28,8 @@ class SupplyModel(SQLModel, table=True):
27
28
  cups: str = Field(default=None, primary_key=True)
28
29
  data: Supply = Field(sa_column=Column(PydanticJSON(Supply)))
29
30
  version: int = Field(default=1)
30
- created_at: dt = Field(default_factory=_now, nullable=False)
31
- updated_at: dt = Field(
31
+ created_at: NaiveDatetime = Field(default_factory=_now, nullable=False)
32
+ updated_at: NaiveDatetime = Field(
32
33
  default_factory=_now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
33
34
  )
34
35
 
@@ -44,12 +45,12 @@ class ContractModel(SQLModel, table=True):
44
45
 
45
46
  id: int | None = Field(default=None, primary_key=True)
46
47
  cups: str = Field(foreign_key="supply.cups", index=True)
47
- date_start: dt = Field(index=True)
48
+ date_start: NaiveDatetime = Field(index=True)
48
49
  data: Contract = Field(sa_column=Column(PydanticJSON(Contract)))
49
50
 
50
51
  version: int = Field(default=1)
51
- created_at: dt = Field(default_factory=_now, nullable=False)
52
- updated_at: dt = Field(
52
+ created_at: NaiveDatetime = Field(default_factory=_now, nullable=False)
53
+ updated_at: NaiveDatetime = Field(
53
54
  default_factory=_now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
54
55
  )
55
56
 
@@ -71,13 +72,13 @@ class EnergyModel(SQLModel, table=True):
71
72
  id: int | None = Field(default=None, primary_key=True)
72
73
  cups: str = Field(foreign_key="supply.cups", index=True)
73
74
  delta_h: float
74
- datetime: dt = Field(index=True)
75
+ datetime: NaiveDatetime = Field(index=True)
75
76
 
76
77
  data: Energy = Field(sa_column=Column(PydanticJSON(Energy)))
77
78
 
78
79
  version: int = Field(default=1)
79
- created_at: dt = Field(default_factory=_now, nullable=False)
80
- updated_at: dt = Field(
80
+ created_at: NaiveDatetime = Field(default_factory=_now, nullable=False)
81
+ updated_at: NaiveDatetime = Field(
81
82
  default_factory=_now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
82
83
  )
83
84
 
@@ -93,13 +94,13 @@ class PowerModel(SQLModel, table=True):
93
94
 
94
95
  id: int | None = Field(default=None, primary_key=True)
95
96
  cups: str = Field(foreign_key="supply.cups", index=True)
96
- datetime: dt = Field(index=True)
97
+ datetime: NaiveDatetime = Field(index=True)
97
98
 
98
99
  data: Power = Field(sa_column=Column(PydanticJSON(Power)))
99
100
 
100
101
  version: int = Field(default=1)
101
- created_at: dt = Field(default_factory=_now, nullable=False)
102
- updated_at: dt = Field(
102
+ created_at: NaiveDatetime = Field(default_factory=_now, nullable=False)
103
+ updated_at: NaiveDatetime = Field(
103
104
  default_factory=_now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
104
105
  )
105
106
 
@@ -120,14 +121,14 @@ class StatisticsModel(SQLModel, table=True):
120
121
 
121
122
  id: int | None = Field(default=None, primary_key=True)
122
123
  cups: str = Field(foreign_key="supply.cups", index=True)
123
- datetime: dt = Field(index=True)
124
+ datetime: NaiveDatetime = Field(index=True)
124
125
  type: typing.Literal["day", "month"] = Field(index=True, sa_type=AutoString)
125
126
  complete: bool = Field(False)
126
127
  data: Statistics = Field(sa_column=Column(PydanticJSON(Statistics)))
127
128
 
128
129
  version: int = Field(default=1)
129
- created_at: dt = Field(default_factory=_now, nullable=False)
130
- updated_at: dt = Field(
130
+ created_at: NaiveDatetime = Field(default_factory=_now, nullable=False)
131
+ updated_at: NaiveDatetime = Field(
131
132
  default_factory=_now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
132
133
  )
133
134
 
@@ -137,12 +138,12 @@ class PVPCModel(SQLModel, table=True):
137
138
  __tablename__ = "pvpc" # type: ignore
138
139
 
139
140
  id: int | None = Field(default=None, primary_key=True)
140
- datetime: dt = Field(index=True, unique=True)
141
+ datetime: NaiveDatetime = Field(index=True, unique=True)
141
142
  data: EnergyPrice = Field(sa_column=Column(PydanticJSON(EnergyPrice)))
142
143
 
143
144
  version: int = Field(default=1)
144
- created_at: dt = Field(default_factory=_now, nullable=False)
145
- updated_at: dt = Field(
145
+ created_at: NaiveDatetime = Field(default_factory=_now, nullable=False)
146
+ updated_at: NaiveDatetime = Field(
146
147
  default_factory=_now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
147
148
  )
148
149
 
@@ -163,14 +164,14 @@ class BillModel(SQLModel, table=True):
163
164
 
164
165
  id: int | None = Field(default=None, primary_key=True)
165
166
  cups: str = Field(foreign_key="supply.cups", index=True)
166
- datetime: dt = Field(index=True)
167
+ datetime: NaiveDatetime = Field(index=True)
167
168
  type: typing.Literal["hour", "day", "month"] = Field(index=True, sa_type=AutoString)
168
169
  complete: bool = Field(False)
169
170
  confhash: str
170
171
  data: Bill = Field(sa_column=Column(PydanticJSON(Bill)))
171
172
 
172
173
  version: int = Field(default=1)
173
- created_at: dt = Field(default_factory=_now, nullable=False)
174
- updated_at: dt = Field(
174
+ created_at: NaiveDatetime = Field(default_factory=_now, nullable=False)
175
+ updated_at: NaiveDatetime = Field(
175
176
  default_factory=_now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
176
177
  )
@@ -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,11 +1,14 @@
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
 
7
8
  import pytest
8
9
  import pytest_asyncio
10
+ from sqlalchemy import DateTime
11
+ from sqlmodel import SQLModel
9
12
 
10
13
  from edata.database.controller import EdataDB
11
14
  from edata.models import Bill, Energy, Supply
@@ -14,16 +17,10 @@ CUPS = "ESXXXXXXXXXXXXXXXXTEST"
14
17
  START = datetime(2024, 1, 1)
15
18
 
16
19
 
17
- def _reset_singleton() -> None:
18
- EdataDB._instance = None
19
- EdataDB._engine = None
20
- EdataDB._db_url = None
21
-
22
-
23
20
  @pytest_asyncio.fixture
24
21
  async def db(tmp_path) -> AsyncIterator[EdataDB]:
25
22
  """An EdataDB on an isolated on-disk database with one supply."""
26
- _reset_singleton()
23
+ EdataDB.reset()
27
24
  database = EdataDB(str(tmp_path / "edata.db"))
28
25
  await database.add_supply(
29
26
  Supply(
@@ -40,9 +37,7 @@ async def db(tmp_path) -> AsyncIterator[EdataDB]:
40
37
  )
41
38
  )
42
39
  yield database
43
- if EdataDB._engine is not None:
44
- await EdataDB._engine.dispose()
45
- _reset_singleton()
40
+ EdataDB.reset()
46
41
 
47
42
 
48
43
  def _energy(hours: int, kwh: float) -> list[Energy]:
@@ -95,14 +90,38 @@ async def test_add_bill_list_applies_overrides(db: EdataDB) -> None:
95
90
  @pytest.mark.asyncio
96
91
  async def test_concurrent_first_calls_do_not_race_index_creation(db: EdataDB) -> None:
97
92
  # simulate a database created before the index existed, opened fresh
98
- async with db.engine.begin() as conn:
99
- 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")
100
95
  db._tables_initialized = False
101
96
 
102
97
  await asyncio.gather(*(db.get_last_energy(CUPS) for _ in range(10)))
103
98
 
104
- async with db.engine.connect() as conn:
105
- result = await conn.exec_driver_sql(
99
+ with db.engine.connect() as conn:
100
+ result = conn.exec_driver_sql(
106
101
  "SELECT name FROM sqlite_master WHERE name='ix_energy_cups_datetime'"
107
102
  )
108
103
  assert result.first() is not None
104
+
105
+
106
+ def test_datetime_columns_are_naive() -> None:
107
+ # the supply timezone is unknown, so datetimes are stored as-is (naive);
108
+ # plain ``datetime`` fields map to UTC-aware columns on sqlmodel>=0.0.45
109
+ columns = [
110
+ column
111
+ for table in SQLModel.metadata.sorted_tables
112
+ for column in table.columns
113
+ if isinstance(column.type, DateTime)
114
+ ]
115
+ assert len(columns) == 20
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}