e-data 2.0.2.dev134__py3-none-any.whl → 2.0.2.dev135__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.dev134
3
+ Version: 2.0.2.dev135
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
@@ -697,7 +697,7 @@ Requires-Dist: pydantic<3,>=2.10
697
697
  Requires-Dist: python_dateutil<3,>=2.8
698
698
  Requires-Dist: Requests<3,>=2.31
699
699
  Requires-Dist: SQLAlchemy[asyncio]<3,>=2.0
700
- Requires-Dist: sqlmodel<0.1,>=0.0.22
700
+ Requires-Dist: sqlmodel<0.0.45,>=0.0.22
701
701
  Requires-Dist: typer<1,>=0.12
702
702
  Requires-Dist: aiosqlite<1,>=0.21
703
703
  Dynamic: license-file
@@ -1,4 +1,4 @@
1
- e_data-2.0.2.dev134.dist-info/licenses/LICENSE,sha256=OXLcl0T2SZ8Pmy2_dmlvKuetivmyPd5m1q-Gyd-zaYY,35149
1
+ e_data-2.0.2.dev135.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,9 +6,9 @@ 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=EW7YKnq1Ti_rslVAuIq5TDjf8THyWBAt-rwMxnsSvdQ,18507
10
- edata/database/models.py,sha256=4rlRpPErULqYr6QJxZar1NBpJPWTcHftDN1MB3bwGyA,5096
11
- edata/database/queries.py,sha256=yfUAwHtEv6JkcGf8aGiQ1Z5ihbgdMtKLGNiwvBikb7Y,8013
9
+ edata/database/controller.py,sha256=xKBNwv1PYfEkuaunz3Qh5zLAPI_6sDkxiohUGEP0jN4,22831
10
+ edata/database/models.py,sha256=tSkxf2OKgT1m2F1td3j3GtYcboUtSuYbhnHWQtpvJuY,5525
11
+ edata/database/queries.py,sha256=XziDSG2V-PVFmvRkj8UrEvgrc9Qc-ShhnxgsPHIsyik,8304
12
12
  edata/database/utils.py,sha256=V-FJfyAh8YJBGHB85BbWtG3Ilq09pM9Gl-L86TffB4g,1062
13
13
  edata/database/migrations/__init__.py,sha256=P8VXZzfyceOV3GknYXif2HEMB8Bo4B7f7vKdIC02ZO0,1203
14
14
  edata/database/migrations/base.py,sha256=V7jfMaTl6XPCwM6Wei4bNxfOKSKanIHF30UgH-_CJNM,444
@@ -18,19 +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=t8xrGYpTORHUH_FNZjwrvGnx6L_CBaaD0pKtvfFTITU,20588
22
- edata/providers/redata.py,sha256=uvxO2TTtSNSfPRWHZifLjbaP7zAr6O7CZ1oRQFNbIhQ,3122
23
- edata/services/bill_service.py,sha256=vukgeQAN0fMjAkyI9cqPqQ4wSRvPdOaswfSR1KheIYc,14947
24
- edata/services/data_service.py,sha256=henmYVQUynCTl1g83mZuU4Fwi2d5Qk6yhEjXOOmvhk4,18278
21
+ edata/providers/datadis.py,sha256=VZ_MhO3x1crSmOH2dn3SWJ5Eq1gWUIDCEvB9oRRMZfQ,21468
22
+ edata/providers/redata.py,sha256=i6aBjunyZMb34SGgAgNQvaNu_GFoe5P2rUIY-_vOElw,3727
23
+ edata/services/bill_service.py,sha256=KAM78u7TpoVF-KTlB7bqj2Kp-rJERiPPoGTJCc-7rM0,15264
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_datadis_connector.py,sha256=pT161wQdA_GbdutqnWDVkM7itNQlK2NIvEETCi33oDM,11711
27
+ edata/tests/test_controller.py,sha256=7sqDR3_Vrxy4S084Eh6VYkLQNuhfYAYzJFCtMSCudd8,2815
28
+ edata/tests/test_datadis_connector.py,sha256=8Xet2M2eG-kMhiRDTCSH0R5nYO2oIpQTPOuTZKMKoLo,12538
28
29
  edata/tests/test_incremental.py,sha256=dn1LRgjBi2HuNYELPKDVdZQlS01hk4lsP9L-vyvtEHo,3479
29
30
  edata/tests/test_migrations.py,sha256=3Io2nokZVdv2xU9fCZCMIvwAwPVtGc0UivClSouP3pw,6418
30
31
  edata/tests/test_redata_connector.py,sha256=SBvxsma5t-qvz2drxNisNM2P1STTurqEs1TSDgwaPJ8,508
31
- edata/tests/test_services.py,sha256=BMw-MYXs2L_c1Sg_v4TwAZGOJnx2BoWH8U50M05-ZC8,6816
32
+ edata/tests/test_services.py,sha256=Ya5cbADHj1Qsffb_IuMeDD0BViz8nqrgwazUWK0-hRc,8297
32
33
  edata/tests/test_utils.py,sha256=h6yBzT7j3eA42zbXJ-AK9lXcZkjCMf16wnob1oqGvUE,2307
33
- e_data-2.0.2.dev134.dist-info/METADATA,sha256=piYxSiQvkdZOXOfPXqGyIkoMRqrDCpChwfbqJC7J0R8,48084
34
- e_data-2.0.2.dev134.dist-info/WHEEL,sha256=K260EYznzXsJYBQGqmI8VTxEdiZYNvDZwW9cBh9-_MA,91
35
- e_data-2.0.2.dev134.dist-info/top_level.txt,sha256=Ez-fReWtUVTMcFuH0dzCnoKIgMBHgjLaBuba8QiKC7g,6
36
- e_data-2.0.2.dev134.dist-info/RECORD,,
34
+ e_data-2.0.2.dev135.dist-info/METADATA,sha256=FWIkMXs3931Au8RSJa5QzAEaFUXJzIboaKHzLt7U0x0,48087
35
+ e_data-2.0.2.dev135.dist-info/WHEEL,sha256=YVMoNqKzERt-wjUZwJ33xBGAwnFl-4cqbYkTtWa4itE,91
36
+ e_data-2.0.2.dev135.dist-info/top_level.txt,sha256=Ez-fReWtUVTMcFuH0dzCnoKIgMBHgjLaBuba8QiKC7g,6
37
+ e_data-2.0.2.dev135.dist-info/RECORD,,
@@ -1,5 +1,5 @@
1
1
  Wheel-Version: 1.0
2
- Generator: setuptools (83.0.0)
2
+ Generator: setuptools (84.0.0)
3
3
  Root-Is-Purelib: true
4
4
  Tag: py3-none-any
5
5
 
@@ -3,10 +3,11 @@ import os
3
3
  import typing
4
4
  from datetime import datetime
5
5
 
6
- from sqlalchemy import Select, event, insert
6
+ from sqlalchemy import Select, Table, event, insert, or_
7
+ from sqlalchemy.dialects.sqlite import insert as sqlite_insert
7
8
  from sqlalchemy.exc import IntegrityError
8
9
  from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
9
- from sqlmodel import SQLModel
10
+ from sqlmodel import SQLModel, UniqueConstraint
10
11
  from sqlmodel.ext.asyncio.session import AsyncSession
11
12
  from sqlmodel.sql.expression import SelectOfScalar
12
13
 
@@ -47,6 +48,34 @@ def _set_sqlite_pragmas(dbapi_connection, connection_record) -> None:
47
48
  cursor.close()
48
49
 
49
50
 
51
+ def _create_missing_indexes(connection) -> None:
52
+ """Create indexes added after a table was first created.
53
+
54
+ ``create_all`` skips tables that already exist, indexes included, so
55
+ databases created by an older version would never get new indexes.
56
+ """
57
+
58
+ for table in SQLModel.metadata.sorted_tables:
59
+ for index in table.indexes:
60
+ index.create(connection, checkfirst=True)
61
+
62
+
63
+ def _conflict_columns(table: Table) -> list[str]:
64
+ """Return the columns that identify a row for upserts on ``table``.
65
+
66
+ That is the table's unique constraint, else its unique index, else its
67
+ primary key.
68
+ """
69
+
70
+ for constraint in table.constraints:
71
+ if isinstance(constraint, UniqueConstraint):
72
+ return [c.name for c in constraint.columns]
73
+ for index in table.indexes:
74
+ if index.unique:
75
+ return [c.name for c in index.columns]
76
+ return [c.name for c in table.primary_key.columns]
77
+
78
+
50
79
  class EdataDB:
51
80
 
52
81
  _instance = None
@@ -82,6 +111,7 @@ class EdataDB:
82
111
  if self.engine:
83
112
  async with self.engine.begin() as conn:
84
113
  await conn.run_sync(SQLModel.metadata.create_all)
114
+ await conn.run_sync(_create_missing_indexes)
85
115
  self._tables_initialized = True
86
116
 
87
117
  async def _add_one(
@@ -158,33 +188,46 @@ class EdataDB:
158
188
  async def _add_or_update_many(
159
189
  self,
160
190
  session: AsyncSession,
161
- queries: list[SelectOfScalar],
162
- records: list[T],
163
- batch_size: int = 100,
191
+ model: type[SQLModel],
192
+ rows: list[dict[str, typing.Any]],
164
193
  override: list[str] | None = None,
165
- ) -> list[T]:
166
- """Updates many records in the database"""
167
- if not records:
168
- return []
169
-
170
- for i in range(0, len(records), batch_size):
171
- chunk_records = records[i : i + batch_size]
172
- chunk_queries = queries[i : i + batch_size]
173
- try:
174
- async with session.begin_nested():
175
- session.add_all(chunk_records)
176
- await session.flush()
177
- except IntegrityError:
178
- for j, record in enumerate(chunk_records):
179
- await self._add_or_update_one(
180
- session,
181
- chunk_queries[j],
182
- record,
183
- commit=False,
184
- override=override,
185
- )
194
+ ) -> None:
195
+ """Insert many rows, updating the existing ones, in a single statement.
196
+
197
+ Uses SQLite's ``INSERT ... ON CONFLICT DO UPDATE`` so re-syncing rows that
198
+ are already stored costs one bulk statement instead of a savepoint, a
199
+ failed insert and a lookup per row. Existing rows are only rewritten (and
200
+ their ``updated_at`` bumped) when ``data`` or an ``override`` column
201
+ actually changed.
202
+
203
+ Rows are plain column dicts: building an ORM instance per row only to
204
+ read it back cost about half of a full-history import. Columns missing
205
+ from a row take the model defaults, resolved once per batch.
206
+ """
207
+ if not rows:
208
+ return
209
+
210
+ table = model.__table__ # type: ignore[attr-defined]
211
+ defaults = {
212
+ c.name: model.model_fields[c.name].get_default(call_default_factory=True)
213
+ for c in table.columns
214
+ if c.name != "id" and c.name not in rows[0]
215
+ }
216
+ rows = [{**defaults, **row} for row in rows]
217
+
218
+ updated = ["data", *(override or [])]
219
+ stmt = sqlite_insert(table)
220
+ stmt = stmt.on_conflict_do_update(
221
+ index_elements=_conflict_columns(table),
222
+ set_={
223
+ **{c: stmt.excluded[c] for c in updated},
224
+ "updated_at": stmt.excluded.updated_at,
225
+ },
226
+ where=or_(*(table.c[c].is_distinct_from(stmt.excluded[c]) for c in updated)),
227
+ )
228
+ connection = await session.connection()
229
+ await connection.execute(stmt, rows)
186
230
  await session.commit()
187
- return records
188
231
 
189
232
  async def get_supply(self, cups: str) -> SupplyModel | None:
190
233
  """Get a supply record by cups."""
@@ -212,6 +255,14 @@ class EdataDB:
212
255
  result = await session.exec(q.get_last_energy(cups))
213
256
  return result.first()
214
257
 
258
+ async def get_last_power(self, cups: str) -> PowerModel | None:
259
+ """Get the most recent power record by cups."""
260
+
261
+ await self._ensure_tables()
262
+ async with AsyncSession(self.engine) as session:
263
+ result = await session.exec(q.get_last_power(cups))
264
+ return result.first()
265
+
215
266
  async def get_last_pvpc(self) -> PVPCModel | None:
216
267
  """Get the most recent pvpc."""
217
268
 
@@ -287,12 +338,11 @@ class EdataDB:
287
338
  async with AsyncSession(self.engine) as session:
288
339
  unique_map = {item.datetime: item for item in energy}
289
340
  unique = list(unique_map.values())
290
- queries = [q.get_energy(cups, x.datetime) for x in unique]
291
- items = [
292
- EnergyModel(cups=cups, delta_h=x.delta_h, datetime=x.datetime, data=x)
341
+ rows = [
342
+ {"cups": cups, "delta_h": x.delta_h, "datetime": x.datetime, "data": x}
293
343
  for x in unique
294
344
  ]
295
- await self._add_or_update_many(session, queries, items)
345
+ await self._add_or_update_many(session, EnergyModel, rows)
296
346
 
297
347
  async def add_power(self, cups: str, power: Power) -> PowerModel | None:
298
348
  """Add or update a power record for a given CUPS and Power instance."""
@@ -311,9 +361,8 @@ class EdataDB:
311
361
  async with AsyncSession(self.engine) as session:
312
362
  unique_map = {item.datetime: item for item in power}
313
363
  unique = list(unique_map.values())
314
- queries = [q.get_power(cups, x.datetime) for x in unique]
315
- items = [PowerModel(cups=cups, datetime=x.datetime, data=x) for x in unique]
316
- await self._add_or_update_many(session, queries, items)
364
+ rows = [{"cups": cups, "datetime": x.datetime, "data": x} for x in unique]
365
+ await self._add_or_update_many(session, PowerModel, rows)
317
366
 
318
367
  async def add_pvpc(self, pvpc: EnergyPrice) -> PVPCModel | None:
319
368
  """Add or update a pvpc record."""
@@ -332,9 +381,8 @@ class EdataDB:
332
381
  async with AsyncSession(self.engine) as session:
333
382
  unique_map = {item.datetime: item for item in pvpc}
334
383
  unique = list(unique_map.values())
335
- queries = [q.get_pvpc(x.datetime) for x in unique]
336
- items = [PVPCModel(datetime=x.datetime, data=x) for x in unique]
337
- await self._add_or_update_many(session, queries, items)
384
+ rows = [{"datetime": x.datetime, "data": x} for x in unique]
385
+ await self._add_or_update_many(session, PVPCModel, rows)
338
386
 
339
387
  async def add_statistics(
340
388
  self,
@@ -357,6 +405,32 @@ class EdataDB:
357
405
  override=["complete"],
358
406
  )
359
407
 
408
+ async def add_statistics_list(
409
+ self,
410
+ cups: str,
411
+ type_: typing.Literal["day", "month"],
412
+ complete: bool,
413
+ statistics: list[Statistics],
414
+ ) -> None:
415
+ """Add or update a list of statistics records."""
416
+
417
+ await self._ensure_tables()
418
+ async with AsyncSession(self.engine) as session:
419
+ unique_map = {item.datetime: item for item in statistics}
420
+ rows = [
421
+ {
422
+ "cups": cups,
423
+ "datetime": x.datetime,
424
+ "type": type_,
425
+ "complete": complete,
426
+ "data": x,
427
+ }
428
+ for x in unique_map.values()
429
+ ]
430
+ await self._add_or_update_many(
431
+ session, StatisticsModel, rows, override=["complete"]
432
+ )
433
+
360
434
  async def add_bill(
361
435
  self,
362
436
  cups: str,
@@ -398,20 +472,19 @@ class EdataDB:
398
472
  async with AsyncSession(self.engine) as session:
399
473
  unique_map = {item.datetime: item for item in bill}
400
474
  unique = list(unique_map.values())
401
- queries = [q.get_bill(cups, type_, x.datetime) for x in unique]
402
- items = [
403
- BillModel(
404
- cups=cups,
405
- datetime=x.datetime,
406
- type=type_,
407
- confhash=confhash,
408
- complete=complete,
409
- data=x,
410
- )
475
+ rows = [
476
+ {
477
+ "cups": cups,
478
+ "datetime": x.datetime,
479
+ "type": type_,
480
+ "confhash": confhash,
481
+ "complete": complete,
482
+ "data": x,
483
+ }
411
484
  for x in unique
412
485
  ]
413
486
  await self._add_or_update_many(
414
- session, queries, items, override=["complete", "confhash"]
487
+ session, BillModel, rows, override=["complete", "confhash"]
415
488
  )
416
489
 
417
490
  async def clear_bills(self, cups: str, since: datetime | None = None) -> None:
@@ -440,6 +513,54 @@ class EdataDB:
440
513
  result = await session.exec(q.list_contract(cups))
441
514
  return result.all()
442
515
 
516
+ async def _list_data(self, query: SelectOfScalar, model: type[SQLModel]) -> list:
517
+ """Return only the ``data`` payload of the rows selected by ``query``.
518
+
519
+ Skips building an ORM instance per row, which nearly halves the cost of
520
+ reading a month of hourly records.
521
+ """
522
+
523
+ await self._ensure_tables()
524
+ async with AsyncSession(self.engine) as session:
525
+ result = await session.exec(
526
+ query.with_only_columns(model.data) # type: ignore[attr-defined]
527
+ )
528
+ return list(result.all())
529
+
530
+ async def list_energy_data(
531
+ self,
532
+ cups: str,
533
+ date_from: datetime | None = None,
534
+ date_to: datetime | None = None,
535
+ ) -> list[Energy]:
536
+ """List the energy data (without row metadata)."""
537
+
538
+ return await self._list_data(
539
+ q.list_energy(cups, date_from, date_to), EnergyModel
540
+ )
541
+
542
+ async def list_pvpc_data(
543
+ self,
544
+ date_from: datetime | None = None,
545
+ date_to: datetime | None = None,
546
+ ) -> list[EnergyPrice]:
547
+ """List the pvpc data (without row metadata)."""
548
+
549
+ return await self._list_data(q.list_pvpc(date_from, date_to), PVPCModel)
550
+
551
+ async def list_bill_data(
552
+ self,
553
+ cups: str,
554
+ type_: typing.Literal["hour", "day", "month"],
555
+ date_from: datetime | None = None,
556
+ date_to: datetime | None = None,
557
+ ) -> list[Bill]:
558
+ """List the bill data (without row metadata)."""
559
+
560
+ return await self._list_data(
561
+ q.list_bill(cups, type_, date_from, date_to), BillModel
562
+ )
563
+
443
564
  async def list_energy(
444
565
  self,
445
566
  cups: str,
edata/database/models.py CHANGED
@@ -1,12 +1,23 @@
1
1
  import typing
2
2
  from datetime import datetime as dt
3
3
 
4
- from sqlmodel import AutoString, Column, Field, SQLModel, UniqueConstraint
4
+ from sqlmodel import AutoString, Column, Field, Index, SQLModel, UniqueConstraint
5
5
 
6
6
  from edata.database.utils import PydanticJSON
7
7
  from edata.models import Bill, Contract, Energy, EnergyPrice, Power, Statistics, Supply
8
8
 
9
9
 
10
+ def _now() -> dt:
11
+ """Return the current time.
12
+
13
+ A plain Python function on purpose: with the ``dt.now`` builtin as
14
+ ``default_factory`` pydantic re-inspects its signature on every instance,
15
+ which made building each row ~6x slower.
16
+ """
17
+
18
+ return dt.now()
19
+
20
+
10
21
  class SupplyModel(SQLModel, table=True):
11
22
 
12
23
  __tablename__ = "supply" # type: ignore
@@ -16,9 +27,9 @@ class SupplyModel(SQLModel, table=True):
16
27
  cups: str = Field(default=None, primary_key=True)
17
28
  data: Supply = Field(sa_column=Column(PydanticJSON(Supply)))
18
29
  version: int = Field(default=1)
19
- created_at: dt = Field(default_factory=dt.now, nullable=False)
30
+ created_at: dt = Field(default_factory=_now, nullable=False)
20
31
  updated_at: dt = Field(
21
- default_factory=dt.now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
32
+ default_factory=_now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
22
33
  )
23
34
 
24
35
 
@@ -37,9 +48,9 @@ class ContractModel(SQLModel, table=True):
37
48
  data: Contract = Field(sa_column=Column(PydanticJSON(Contract)))
38
49
 
39
50
  version: int = Field(default=1)
40
- created_at: dt = Field(default_factory=dt.now, nullable=False)
51
+ created_at: dt = Field(default_factory=_now, nullable=False)
41
52
  updated_at: dt = Field(
42
- default_factory=dt.now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
53
+ default_factory=_now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
43
54
  )
44
55
 
45
56
 
@@ -51,6 +62,9 @@ class EnergyModel(SQLModel, table=True):
51
62
  UniqueConstraint(
52
63
  "cups", "delta_h", "datetime", name="uq_energy_cups_delta_datetime"
53
64
  ),
65
+ # serves the per-cups range scans and "latest record" lookups without
66
+ # sorting the whole table
67
+ Index("ix_energy_cups_datetime", "cups", "datetime"),
54
68
  {"extend_existing": True},
55
69
  )
56
70
 
@@ -62,9 +76,9 @@ class EnergyModel(SQLModel, table=True):
62
76
  data: Energy = Field(sa_column=Column(PydanticJSON(Energy)))
63
77
 
64
78
  version: int = Field(default=1)
65
- created_at: dt = Field(default_factory=dt.now, nullable=False)
79
+ created_at: dt = Field(default_factory=_now, nullable=False)
66
80
  updated_at: dt = Field(
67
- default_factory=dt.now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
81
+ default_factory=_now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
68
82
  )
69
83
 
70
84
 
@@ -84,9 +98,9 @@ class PowerModel(SQLModel, table=True):
84
98
  data: Power = Field(sa_column=Column(PydanticJSON(Power)))
85
99
 
86
100
  version: int = Field(default=1)
87
- created_at: dt = Field(default_factory=dt.now, nullable=False)
101
+ created_at: dt = Field(default_factory=_now, nullable=False)
88
102
  updated_at: dt = Field(
89
- default_factory=dt.now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
103
+ default_factory=_now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
90
104
  )
91
105
 
92
106
 
@@ -112,9 +126,9 @@ class StatisticsModel(SQLModel, table=True):
112
126
  data: Statistics = Field(sa_column=Column(PydanticJSON(Statistics)))
113
127
 
114
128
  version: int = Field(default=1)
115
- created_at: dt = Field(default_factory=dt.now, nullable=False)
129
+ created_at: dt = Field(default_factory=_now, nullable=False)
116
130
  updated_at: dt = Field(
117
- default_factory=dt.now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
131
+ default_factory=_now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
118
132
  )
119
133
 
120
134
 
@@ -127,9 +141,9 @@ class PVPCModel(SQLModel, table=True):
127
141
  data: EnergyPrice = Field(sa_column=Column(PydanticJSON(EnergyPrice)))
128
142
 
129
143
  version: int = Field(default=1)
130
- created_at: dt = Field(default_factory=dt.now, nullable=False)
144
+ created_at: dt = Field(default_factory=_now, nullable=False)
131
145
  updated_at: dt = Field(
132
- default_factory=dt.now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
146
+ default_factory=_now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
133
147
  )
134
148
 
135
149
 
@@ -156,7 +170,7 @@ class BillModel(SQLModel, table=True):
156
170
  data: Bill = Field(sa_column=Column(PydanticJSON(Bill)))
157
171
 
158
172
  version: int = Field(default=1)
159
- created_at: dt = Field(default_factory=dt.now, nullable=False)
173
+ created_at: dt = Field(default_factory=_now, nullable=False)
160
174
  updated_at: dt = Field(
161
- default_factory=dt.now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
175
+ default_factory=_now, nullable=False, sa_column_kwargs={"onupdate": dt.now}
162
176
  )
edata/database/queries.py CHANGED
@@ -118,6 +118,18 @@ def list_power(
118
118
  return query
119
119
 
120
120
 
121
+ def get_last_power(
122
+ cups: str,
123
+ ) -> SelectOfScalar[PowerModel]:
124
+ """Query that selects the most recent power record."""
125
+
126
+ query = select(PowerModel).where(PowerModel.cups == cups)
127
+ query = query.order_by(desc(PowerModel.datetime))
128
+ query = query.limit(1)
129
+
130
+ return query
131
+
132
+
121
133
  # Queries for "statistics" table
122
134
  def get_statistics(
123
135
  cups: str, type_: typing.Literal["day", "month"], datetime_: datetime | None = None
@@ -78,7 +78,15 @@ class DatadisConnector:
78
78
  password: str,
79
79
  enable_smart_fetch: bool = True,
80
80
  storage_path: str | None = None,
81
+ session: aiohttp.ClientSession | None = None,
81
82
  ) -> None:
83
+ """Init the connector.
84
+
85
+ ``session`` lets the caller share a long-lived aiohttp session (e.g. Home
86
+ Assistant's) so requests reuse pooled TLS connections; without it a
87
+ short-lived session is opened per request.
88
+ """
89
+ self._session = session
82
90
  self._usr = username
83
91
  self._pwd = password
84
92
  self._token = {}
@@ -96,6 +104,15 @@ class DatadisConnector:
96
104
  os.makedirs(self._recent_cache_dir, exist_ok=True)
97
105
  self._cache = diskcache.Cache(self._recent_cache_dir)
98
106
 
107
+ @contextlib.asynccontextmanager
108
+ async def _client(self) -> typing.AsyncIterator[aiohttp.ClientSession]:
109
+ """Yield the shared session, or a short-lived one when none was given."""
110
+ if self._session is not None:
111
+ yield self._session
112
+ else:
113
+ async with aiohttp.ClientSession() as session:
114
+ yield session
115
+
99
116
  def _get_hash(self, item: str) -> str:
100
117
  """Return a hash."""
101
118
  return hashlib.md5(item.encode()).hexdigest()
@@ -127,7 +144,7 @@ class DatadisConnector:
127
144
  _LOGGER.debug("No token found, fetching a new one")
128
145
  is_valid_token = False
129
146
  timeout = aiohttp.ClientTimeout(total=TIMEOUT)
130
- async with aiohttp.ClientSession(timeout=timeout) as session:
147
+ async with self._client() as session:
131
148
  try:
132
149
  async with session.post(
133
150
  URL_TOKEN,
@@ -135,6 +152,7 @@ class DatadisConnector:
135
152
  TOKEN_USERNAME: self._usr,
136
153
  TOKEN_PASSWD: self._pwd,
137
154
  },
155
+ timeout=timeout,
138
156
  ) as response:
139
157
  text = await response.text()
140
158
  if response.status == 200:
@@ -209,12 +227,14 @@ class DatadisConnector:
209
227
  if self._token.get("headers"):
210
228
  headers.update(self._token["headers"])
211
229
  timeout = aiohttp.ClientTimeout(total=TIMEOUT)
212
- async with aiohttp.ClientSession(timeout=timeout) as session:
230
+ async with self._client() as session:
213
231
  async with session.get(
214
232
  url + params,
215
233
  headers=headers,
234
+ timeout=timeout,
216
235
  ) as reply:
217
- text = await reply.text()
236
+ # the body is only needed as text to log errors; decoding
237
+ # it here as well would decode every (large) payload twice
218
238
  if reply.status == 200:
219
239
  try:
220
240
  json_data = await reply.json(content_type=None)
@@ -247,7 +267,7 @@ class DatadisConnector:
247
267
  _LOGGER.warning(
248
268
  "%s with message '%s'",
249
269
  reply.status,
250
- text,
270
+ await reply.text(),
251
271
  )
252
272
  if not ignore_cache:
253
273
  await asyncio.to_thread(self._set_cache, url + params)
@@ -256,7 +276,7 @@ class DatadisConnector:
256
276
  _LOGGER.warning(
257
277
  "%s with message '%s'. %s. %s",
258
278
  reply.status,
259
- text,
279
+ await reply.text(),
260
280
  "Query temporary disabled",
261
281
  "Future 500 code errors for this query will be silenced until restart",
262
282
  )
edata/providers/redata.py CHANGED
@@ -1,8 +1,10 @@
1
1
  """A REData API connector"""
2
2
 
3
3
  import asyncio
4
+ import contextlib
4
5
  import datetime as dt
5
6
  import logging
7
+ import typing
6
8
 
7
9
  import aiohttp
8
10
  from dateutil import parser
@@ -26,8 +28,23 @@ class REDataConnector:
26
28
 
27
29
  def __init__(
28
30
  self,
31
+ session: aiohttp.ClientSession | None = None,
29
32
  ) -> None:
30
- """Init method for REDataConnector"""
33
+ """Init method for REDataConnector
34
+
35
+ ``session`` lets the caller share a long-lived aiohttp session; without
36
+ it a short-lived session is opened per request.
37
+ """
38
+ self._session = session
39
+
40
+ @contextlib.asynccontextmanager
41
+ async def _client(self) -> typing.AsyncIterator[aiohttp.ClientSession]:
42
+ """Yield the shared session, or a short-lived one when none was given."""
43
+ if self._session is not None:
44
+ yield self._session
45
+ else:
46
+ async with aiohttp.ClientSession() as session:
47
+ yield session
31
48
 
32
49
  async def async_get_realtime_prices(
33
50
  self, dt_from: dt.datetime, dt_to: dt.datetime, is_ceuta_melilla: bool = False
@@ -41,10 +58,9 @@ class REDataConnector:
41
58
  data = []
42
59
  _LOGGER.info("GET %s", url)
43
60
  timeout = aiohttp.ClientTimeout(total=REQUESTS_TIMEOUT)
44
- async with aiohttp.ClientSession(timeout=timeout) as session:
61
+ async with self._client() as session:
45
62
  try:
46
- async with session.get(url) as res:
47
- text = await res.text()
63
+ async with session.get(url, timeout=timeout) as res:
48
64
  if res.status == 200:
49
65
  try:
50
66
  res_json = await res.json()
@@ -53,7 +69,7 @@ class REDataConnector:
53
69
  _LOGGER.error(
54
70
  "%s returned a malformed response: %s ",
55
71
  url,
56
- text,
72
+ await res.text(),
57
73
  )
58
74
  return data
59
75
  for element in res_list:
@@ -70,7 +86,7 @@ class REDataConnector:
70
86
  _LOGGER.error(
71
87
  "%s returned %s with code %s",
72
88
  url,
73
- text,
89
+ await res.text(),
74
90
  res.status,
75
91
  )
76
92
  except Exception as e:
@@ -44,8 +44,7 @@ class BillService:
44
44
  ) -> list[Bill]:
45
45
  """Return the list of bills."""
46
46
 
47
- res = await self.db.list_bill(self._cups, type_, start, end)
48
- return [x.data for x in res]
47
+ return await self.db.list_bill_data(self._cups, type_, start, end)
49
48
 
50
49
  async def update(
51
50
  self,
@@ -234,27 +233,39 @@ class BillService:
234
233
  )
235
234
  data = await self.get_bills(month, month_end, "hour")
236
235
 
236
+ # aggregate the whole month in a single pass (and a single thread hop)
237
+ day_bills, month_bills = await asyncio.to_thread(
238
+ self._compile_day_and_month, data
239
+ )
240
+ by_day = {x.datetime: x for x in day_bills}
241
+
242
+ done: dict[bool, list[Bill]] = {True: [], False: []}
237
243
  for day in ledger.days_in(month):
238
- day_end = day + timedelta(days=1) - timedelta(microseconds=1)
239
- day_data = [x for x in data if day <= x.datetime <= day_end]
240
- stat = await asyncio.to_thread(
241
- self._compile_statistics, day_data, get_day
242
- )
243
- delta_h = stat[0].delta_h if stat else 0.0
244
- final = is_day_final(day, delta_h, now)
245
- if stat:
246
- await self.db.add_bill(self._cups, "day", stat[0], "mix", final)
244
+ bill = by_day.get(day)
245
+ final = is_day_final(day, bill.delta_h if bill else 0.0, now)
246
+ if bill:
247
+ done[final].append(bill)
247
248
  ledger.resolve_day(day, final=final)
249
+ for complete, bills in done.items():
250
+ if bills:
251
+ await self.db.add_bill_list(
252
+ self._cups, "day", "mix", complete, bills
253
+ )
248
254
 
249
- month_stat = await asyncio.to_thread(
250
- self._compile_statistics, data, get_month
251
- )
252
- delta_h = month_stat[0].delta_h if month_stat else 0.0
255
+ delta_h = month_bills[0].delta_h if month_bills else 0.0
253
256
  final = is_month_final(month, delta_h, now)
254
- if month_stat:
255
- await self.db.add_bill(self._cups, "month", month_stat[0], "mix", final)
257
+ if month_bills:
258
+ await self.db.add_bill(self._cups, "month", month_bills[0], "mix", final)
256
259
  ledger.resolve_month(month, final=final)
257
260
 
261
+ def _compile_day_and_month(self, data: list[Bill]) -> tuple[list[Bill], list[Bill]]:
262
+ """Return the daily and monthly aggregates of a month of hourly bills."""
263
+
264
+ return (
265
+ self._compile_statistics(data, get_day),
266
+ self._compile_statistics(data, get_month),
267
+ )
268
+
258
269
  def _compile_statistics(
259
270
  self,
260
271
  data: list[Bill],
@@ -422,15 +433,13 @@ class BillService:
422
433
  self, start: datetime | None = None, end: datetime | None = None
423
434
  ) -> list[Energy]:
424
435
  """Get energy."""
425
- res = await self.db.list_energy(self._cups, start, end)
426
- return [x.data for x in res]
436
+ return await self.db.list_energy_data(self._cups, start, end)
427
437
 
428
438
  async def _get_pvpc(
429
439
  self, start: datetime | None = None, end: datetime | None = None
430
440
  ) -> list[EnergyPrice]:
431
441
  """Get PVPC."""
432
- res = await self.db.list_pvpc(start, end)
433
- return [x.data for x in res]
442
+ return await self.db.list_pvpc_data(start, end)
434
443
 
435
444
  async def _get_last_bill_dt(self) -> datetime | None:
436
445
  """Return the timestamp of the latest bill record."""
@@ -8,6 +8,7 @@ from datetime import datetime, timedelta
8
8
  from pathlib import Path
9
9
  from tempfile import gettempdir
10
10
 
11
+ import aiohttp
11
12
  from dateutil import relativedelta
12
13
 
13
14
  from edata.core.completion import PendingLedger, is_day_final, is_month_final
@@ -37,12 +38,13 @@ class DataService:
37
38
  datadis_pwd: str,
38
39
  storage_path: str,
39
40
  datadis_authorized_nif: str | None = None,
41
+ session: aiohttp.ClientSession | None = None,
40
42
  ) -> None:
41
43
 
42
44
  self.datadis = DatadisConnector(
43
- datadis_user, datadis_pwd, storage_path=storage_path
45
+ datadis_user, datadis_pwd, storage_path=storage_path, session=session
44
46
  )
45
- self.redata = REDataConnector()
47
+ self.redata = REDataConnector(session=session)
46
48
 
47
49
  # params
48
50
  self._cups = cups
@@ -86,8 +88,7 @@ class DataService:
86
88
  ) -> list[Energy]:
87
89
  """Return a list of energy records for the selected cups."""
88
90
 
89
- res = await self.db.list_energy(self._cups, start, end)
90
- return [x.data for x in res]
91
+ return await self.db.list_energy_data(self._cups, start, end)
91
92
 
92
93
  async def get_power(
93
94
  self, start: datetime | None = None, end: datetime | None = None
@@ -102,8 +103,7 @@ class DataService:
102
103
  ) -> list[EnergyPrice]:
103
104
  """Return a list of pvpc records (energy prices) for the selected cups."""
104
105
 
105
- res = await self.db.list_pvpc(start, end)
106
- return [x.data for x in res]
106
+ return await self.db.list_pvpc_data(start, end)
107
107
 
108
108
  async def get_statistics(
109
109
  self,
@@ -164,6 +164,7 @@ class DataService:
164
164
  supply.date_end,
165
165
  )
166
166
 
167
+ explicit_start = start_date is not None
167
168
  if not start_date:
168
169
  start_date = supply.date_start
169
170
  _LOGGER.debug(
@@ -218,8 +219,12 @@ class DataService:
218
219
  # we have no data yet, fetch from start
219
220
  await self.update_energy(start_date, end_date)
220
221
 
221
- # update power records
222
- await self.update_power(start_date, end_date)
222
+ # update power records; unless a start is forced, only refetch from the
223
+ # month of the latest stored peak instead of the whole supply history
224
+ power_start = start_date
225
+ if not explicit_start and (last_power_dt := await self._get_last_power_dt()):
226
+ power_start = max(start_date, get_month(last_power_dt))
227
+ await self.update_power(power_start, end_date)
223
228
 
224
229
  # fetch pvpc data
225
230
  await self.update_pvpc(start_date, end_date)
@@ -394,27 +399,41 @@ class DataService:
394
399
  )
395
400
  data = await self.get_energy(month, month_end)
396
401
 
402
+ # aggregate the whole month in a single pass (and a single thread hop)
403
+ day_stats, month_stats = await asyncio.to_thread(
404
+ self._compile_day_and_month, data
405
+ )
406
+ by_day = {x.datetime: x for x in day_stats}
407
+
408
+ done: dict[bool, list[Statistics]] = {True: [], False: []}
397
409
  for day in ledger.days_in(month):
398
- day_end = day + timedelta(days=1) - timedelta(microseconds=1)
399
- day_data = [x for x in data if day <= x.datetime <= day_end]
400
- stat = await asyncio.to_thread(
401
- self._compile_statistics, day_data, get_day
402
- )
403
- delta_h = stat[0].delta_h if stat else 0.0
404
- final = is_day_final(day, delta_h, now)
410
+ stat = by_day.get(day)
411
+ final = is_day_final(day, stat.delta_h if stat else 0.0, now)
405
412
  if stat:
406
- await self.db.add_statistics(self._cups, "day", stat[0], final)
413
+ done[final].append(stat)
407
414
  ledger.resolve_day(day, final=final)
415
+ for complete, stats in done.items():
416
+ if stats:
417
+ await self.db.add_statistics_list(
418
+ self._cups, "day", complete, stats
419
+ )
408
420
 
409
- month_stat = await asyncio.to_thread(
410
- self._compile_statistics, data, get_month
411
- )
412
- delta_h = month_stat[0].delta_h if month_stat else 0.0
421
+ delta_h = month_stats[0].delta_h if month_stats else 0.0
413
422
  final = is_month_final(month, delta_h, now)
414
- if month_stat:
415
- await self.db.add_statistics(self._cups, "month", month_stat[0], final)
423
+ if month_stats:
424
+ await self.db.add_statistics(self._cups, "month", month_stats[0], final)
416
425
  ledger.resolve_month(month, final=final)
417
426
 
427
+ def _compile_day_and_month(
428
+ self, data: list[Energy]
429
+ ) -> tuple[list[Statistics], list[Statistics]]:
430
+ """Return the daily and monthly aggregates of a month of energy data."""
431
+
432
+ return (
433
+ self._compile_statistics(data, get_day),
434
+ self._compile_statistics(data, get_month),
435
+ )
436
+
418
437
  def _compile_statistics(
419
438
  self,
420
439
  data: list[Energy],
@@ -491,6 +510,13 @@ class DataService:
491
510
  if last_record:
492
511
  return last_record.datetime
493
512
 
513
+ async def _get_last_power_dt(self) -> datetime | None:
514
+ """Return the timestamp of the latest power record."""
515
+
516
+ last_record = await self.db.get_last_power(self._cups)
517
+ if last_record:
518
+ return last_record.datetime
519
+
494
520
  async def _get_last_pvpc_dt(self) -> datetime | None:
495
521
  """Return the timestamp of the latest pvpc record."""
496
522
 
@@ -0,0 +1,91 @@
1
+ """Bulk upsert tests for the database controller."""
2
+
3
+ from collections.abc import AsyncIterator
4
+ from datetime import datetime, timedelta
5
+
6
+ import pytest
7
+ import pytest_asyncio
8
+
9
+ from edata.database.controller import EdataDB
10
+ from edata.models import Bill, Energy, Supply
11
+
12
+ CUPS = "ESXXXXXXXXXXXXXXXXTEST"
13
+ START = datetime(2024, 1, 1)
14
+
15
+
16
+ def _reset_singleton() -> None:
17
+ EdataDB._instance = None
18
+ EdataDB._engine = None
19
+ EdataDB._db_url = None
20
+
21
+
22
+ @pytest_asyncio.fixture
23
+ async def db(tmp_path) -> AsyncIterator[EdataDB]:
24
+ """An EdataDB on an isolated on-disk database with one supply."""
25
+ _reset_singleton()
26
+ database = EdataDB(str(tmp_path / "edata.db"))
27
+ await database.add_supply(
28
+ Supply(
29
+ cups=CUPS,
30
+ date_start=START,
31
+ date_end=START + timedelta(days=30),
32
+ address=None,
33
+ postal_code=None,
34
+ province=None,
35
+ municipality=None,
36
+ distributor=None,
37
+ point_type=5,
38
+ distributor_code="2",
39
+ )
40
+ )
41
+ yield database
42
+ if EdataDB._engine is not None:
43
+ await EdataDB._engine.dispose()
44
+ _reset_singleton()
45
+
46
+
47
+ def _energy(hours: int, kwh: float) -> list[Energy]:
48
+ return [
49
+ Energy(
50
+ datetime=START + timedelta(hours=h),
51
+ delta_h=1.0,
52
+ consumption_kwh=kwh,
53
+ real=True,
54
+ )
55
+ for h in range(hours)
56
+ ]
57
+
58
+
59
+ @pytest.mark.asyncio
60
+ async def test_add_energy_list_upserts(db: EdataDB) -> None:
61
+ await db.add_energy_list(CUPS, _energy(24, 0.1))
62
+ before = {x.datetime: x for x in await db.list_energy(CUPS)}
63
+
64
+ # overlapping batch: first 12 hours unchanged, last 12 changed, 12 new
65
+ await db.add_energy_list(CUPS, _energy(12, 0.1) + _energy(36, 0.5)[12:])
66
+ after = {x.datetime: x for x in await db.list_energy(CUPS)}
67
+
68
+ assert len(after) == 36
69
+ for h in range(36):
70
+ dt = START + timedelta(hours=h)
71
+ assert after[dt].data.consumption_kwh == (0.1 if h < 12 else 0.5)
72
+ if h < 12:
73
+ # unchanged rows are not rewritten
74
+ assert after[dt].id == before[dt].id
75
+ assert after[dt].updated_at == before[dt].updated_at
76
+ elif h < 24:
77
+ assert after[dt].id == before[dt].id
78
+ assert after[dt].updated_at > before[dt].updated_at
79
+
80
+
81
+ @pytest.mark.asyncio
82
+ async def test_add_bill_list_applies_overrides(db: EdataDB) -> None:
83
+ bills = [Bill(datetime=START + timedelta(hours=h), delta_h=1) for h in range(3)]
84
+ await db.add_bill_list(CUPS, "hour", "hash-a", False, bills)
85
+
86
+ # same data, only the override columns change
87
+ await db.add_bill_list(CUPS, "hour", "hash-b", True, bills)
88
+
89
+ stored = await db.list_bill(CUPS, "hour")
90
+ assert len(stored) == 3
91
+ assert all(x.complete and x.confhash == "hash-b" for x in stored)
@@ -4,6 +4,8 @@ import datetime
4
4
  import os
5
5
  from unittest.mock import AsyncMock, MagicMock, patch
6
6
 
7
+ import pytest
8
+
7
9
  from edata.providers.datadis import DatadisConnector
8
10
 
9
11
  MOCK_USERNAME = "USERNAME"
@@ -347,3 +349,26 @@ def test_get_supplies_optional_fields_none(mock_token, mock_get, snapshot):
347
349
  mock_get.return_value.__aenter__.return_value = mock_response
348
350
  connector = DatadisConnector(MOCK_USERNAME, MOCK_PASSWORD)
349
351
  assert connector.get_supplies() == snapshot
352
+
353
+
354
+ @pytest.mark.asyncio
355
+ @patch.object(
356
+ DatadisConnector, "_async_get_token", new_callable=AsyncMock, return_value=True
357
+ )
358
+ async def test_shared_session_is_reused(mock_token, tmp_path, snapshot):
359
+ """A caller-provided session serves the requests and is left open."""
360
+ mock_response = MagicMock()
361
+ mock_response.status = 200
362
+ mock_response.json = AsyncMock(return_value=SUPPLIES_RESPONSE)
363
+ session = MagicMock()
364
+ session.get.return_value.__aenter__.return_value = mock_response
365
+ session.close = AsyncMock()
366
+
367
+ connector = DatadisConnector(
368
+ MOCK_USERNAME, MOCK_PASSWORD, storage_path=str(tmp_path), session=session
369
+ )
370
+ supplies = await connector.async_get_supplies()
371
+
372
+ assert supplies == snapshot
373
+ session.get.assert_called_once()
374
+ session.close.assert_not_awaited()
@@ -8,7 +8,7 @@ import pytest
8
8
  import pytest_asyncio
9
9
  from syrupy.assertion import SnapshotAssertion
10
10
 
11
- from edata.core.utils import get_day
11
+ from edata.core.utils import get_day, get_month
12
12
  from edata.models.bill import BillingRules
13
13
  from edata.models.data import Energy, Power
14
14
  from edata.models.supply import Contract, Supply
@@ -231,3 +231,40 @@ async def test_update_pvpc_clamps_range_to_min_date(populated_data_service):
231
231
  mock_fetch.assert_awaited_once()
232
232
  called_start = mock_fetch.await_args.args[0]
233
233
  assert called_start >= get_day(now) - timedelta(days=28)
234
+
235
+
236
+ @pytest.mark.asyncio
237
+ async def test_update_power_is_incremental(populated_data_service, power):
238
+ ds = populated_data_service
239
+ last_power_dt = max(x.datetime for x in power)
240
+
241
+ with (
242
+ patch.object(ds, "update_energy", AsyncMock(return_value=True)),
243
+ patch.object(ds, "update_pvpc", AsyncMock(return_value=True)),
244
+ patch.object(ds, "update_statistics_incremental", AsyncMock()),
245
+ patch.object(ds, "update_power", AsyncMock(return_value=True)) as mock_power,
246
+ ):
247
+ await ds.update()
248
+ # only refetch from the month of the latest stored peak
249
+ assert mock_power.await_args.args[0] == get_month(last_power_dt)
250
+
251
+ # an explicit start date still forces the full range
252
+ forced_start = datetime(2020, 1, 1)
253
+ await ds.update(start_date=forced_start)
254
+ assert mock_power.await_args.args[0] == forced_start
255
+
256
+
257
+ @pytest.mark.asyncio
258
+ async def test_missing_indexes_are_created_on_existing_db(populated_data_service):
259
+ db = populated_data_service.db
260
+
261
+ async with db.engine.begin() as conn:
262
+ await conn.exec_driver_sql("DROP INDEX IF EXISTS ix_energy_cups_datetime")
263
+ db._tables_initialized = False
264
+ await db._ensure_tables()
265
+
266
+ async with db.engine.connect() as conn:
267
+ result = await conn.exec_driver_sql(
268
+ "SELECT name FROM sqlite_master WHERE type='index' AND tbl_name='energy'"
269
+ )
270
+ assert "ix_energy_cups_datetime" in {row[0] for row in result}