Compare commits
173
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b472f91f19 | ||
|
|
c9ce05d95a | ||
|
|
33ca3da593 | ||
|
|
d5623b9fcb | ||
|
|
b1b6a309cb | ||
|
|
8dc3f71aaf | ||
|
|
2be4f78f8a | ||
|
|
1a958cb102 | ||
|
|
c851bad829 | ||
|
|
09abe05f66 | ||
|
|
0924e8df52 | ||
|
|
72ac7e7300 | ||
|
|
18083822ea | ||
|
|
3beeb5a461 | ||
|
|
631b14e2ec | ||
|
|
231c340ea6 | ||
|
|
c24b6684cc | ||
|
|
d9e82038dc | ||
|
|
e59c192097 | ||
|
|
2a47dab272 | ||
|
|
9db7f63274 | ||
|
|
ebf96de4f1 | ||
|
|
5b9d60e80a | ||
|
|
963e43e3e4 | ||
|
|
39c11ae606 | ||
|
|
3eec701448 | ||
|
|
489e5b596a | ||
|
|
567ddb9779 | ||
|
|
0fb51d338c | ||
|
|
b812d5ac46 | ||
|
|
a9458394f2 | ||
|
|
4884a19e3d | ||
|
|
afe653bafa | ||
|
|
25a08c47a4 | ||
|
|
ffc693e995 | ||
|
|
5855fff451 | ||
|
|
1ea2f659e0 | ||
|
|
2e125dbd53 | ||
|
|
28486a83c7 | ||
|
|
a78401c2ef | ||
|
|
009856a50d | ||
|
|
43c2ddce1a | ||
|
|
d7f04aee8c | ||
|
|
8dbb59a3b7 | ||
|
|
2992bbb0ef | ||
|
|
22faeb45bb | ||
|
|
c37dfacfc7 | ||
|
|
16b050d821 | ||
|
|
d3d914b117 | ||
|
|
fc4af857e3 | ||
|
|
0958d9a2e9 | ||
|
|
b405aea88b | ||
|
|
bfc7aa3031 | ||
|
|
2f63b9630c | ||
|
|
d07a083e03 | ||
|
|
b65f700d56 | ||
|
|
f4cea3874b | ||
|
|
134f0abb5f | ||
|
|
3e04b15656 | ||
|
|
f2e8f6a8e7 | ||
|
|
90a03e7fd6 | ||
|
|
d3fc90b320 | ||
|
|
d4acfc438a | ||
|
|
c26160b10b | ||
|
|
efbe36d7c0 | ||
|
|
f663981cdb | ||
|
|
da05fd2f09 | ||
|
|
188168b16a | ||
|
|
f1e7ce3133 | ||
|
|
e3d0d17cac | ||
|
|
33f5de0591 | ||
|
|
839faaf95b | ||
|
|
82bb329500 | ||
|
|
17057ed41e | ||
|
|
227b146c2c | ||
|
|
32a20785ee | ||
|
|
2544514f52 | ||
|
|
96e85fb48e | ||
|
|
b71009620a | ||
|
|
8b91146c29 | ||
|
|
682e06d256 | ||
|
|
d7f0731d1c | ||
|
|
9bc46f8d2d | ||
|
|
6f8cb05eab | ||
|
|
effe434637 | ||
|
|
9d3e28a529 | ||
|
|
7e61b0a938 | ||
|
|
e8521351d7 | ||
|
|
8d4f496ff4 | ||
|
|
78f3cc4776 | ||
|
|
96e88861d4 | ||
|
|
57f2459f60 | ||
|
|
c61fc2d4ba | ||
|
|
d35ac7d752 | ||
|
|
ef48032adb | ||
|
|
7e266a21eb | ||
|
|
4284bd40a7 | ||
|
|
bedea196c3 | ||
|
|
df54f5518b | ||
|
|
18fe18281b | ||
|
|
355c8eb8b8 | ||
|
|
1ffdaebbbd | ||
|
|
c64693b8e5 | ||
|
|
c89d083942 | ||
|
|
6090b98ab0 | ||
|
|
489cd7df3d | ||
|
|
35cdc7f22a | ||
|
|
75d685f04d | ||
|
|
935e68846c | ||
|
|
4767a7af60 | ||
|
|
45d87f36f7 | ||
|
|
35b95f5c88 | ||
|
|
9d05f5ea62 | ||
|
|
360cddc05b | ||
|
|
14b846549e | ||
|
|
68165f0e01 | ||
|
|
2dc469274a | ||
|
|
702a1cec00 | ||
|
|
49d15d8ffe | ||
|
|
c89cea4953 | ||
|
|
d9c82e39fb | ||
|
|
eac7860b51 | ||
|
|
da7d0f9d62 | ||
|
|
e49a6ad963 | ||
|
|
d70b985eb0 | ||
|
|
e3c47b2762 | ||
|
|
dbf17f3592 | ||
|
|
ae75e4582d | ||
|
|
ee1264b66b | ||
|
|
51da8c5716 | ||
|
|
81bd32f613 | ||
|
|
ee3031013e | ||
|
|
5026967cee | ||
|
|
5481fbc99c | ||
|
|
e480a84a44 | ||
|
|
64d3882c16 | ||
|
|
6b67b8e0f8 | ||
|
|
0cb94d85ec | ||
|
|
9da88db221 | ||
|
|
1a3aaea933 | ||
|
|
bf7fd71a21 | ||
|
|
962ba26c7c | ||
|
|
da236643f2 | ||
|
|
bd09523e94 | ||
|
|
53f1245d83 | ||
|
|
51f712f602 | ||
|
|
f8b1e5fc71 | ||
|
|
a9830c42d8 | ||
|
|
8aa7316b26 | ||
|
|
32d93bba2a | ||
|
|
0d988a9b28 | ||
|
|
ef7ea6b971 | ||
|
|
6cc6382515 | ||
|
|
ef2bd3c9c5 | ||
|
|
cc2c02a2e2 | ||
|
|
b2e26f0b17 | ||
|
|
8975acc48b | ||
|
|
6cfeb2b865 | ||
|
|
dba9e28540 | ||
|
|
2bc5d6ea9a | ||
|
|
3ec663e138 | ||
|
|
048414c5cb | ||
|
|
9ce3f2a0b8 | ||
|
|
0fba7cfe11 | ||
|
|
d8303eaa3d | ||
|
|
8da1f13e60 | ||
|
|
de77019ce3 | ||
|
|
c2b1b7b751 | ||
|
|
3628ac51e5 | ||
|
|
1756192270 | ||
|
|
66ec9979cc | ||
|
|
c1a5d7a425 | ||
|
|
1e0b235cef |
@@ -8,3 +8,6 @@ data
|
||||
openapi
|
||||
src
|
||||
|
||||
# Frontend host build artifacts — built inside the node stage, not needed from context
|
||||
frontend/node_modules
|
||||
frontend/dist
|
||||
|
||||
@@ -7,8 +7,14 @@ APP_DATABASE_URL=sqlite:////app/data/app.db
|
||||
AUTH_BOOTSTRAP_USERNAME=admin
|
||||
AUTH_BOOTSTRAP_PASSWORD=change-me
|
||||
|
||||
# Required by Docker Compose for the WarmteLink serial device. Set these only in
|
||||
# your local .env; use a stable /dev/serial/by-id path and its numeric host GID.
|
||||
# WARMTELINK_DEVICE_PATH=/dev/serial/by-id/<stable-by-id-name>
|
||||
# WARMTELINK_SERIAL_GID=<host-serial-gid>
|
||||
|
||||
# Optional: runtime overrides.
|
||||
# Leave these commented out to use the application's built-in defaults.
|
||||
# TZ=Europe/Amsterdam
|
||||
# APP_DEBUG=
|
||||
# AUTH_SESSION_COOKIE_NAME=
|
||||
# AUTH_SESSION_TTL_HOURS=
|
||||
@@ -28,3 +34,40 @@ TICKTICK_CLIENT_ID=
|
||||
TICKTICK_CLIENT_SECRET=
|
||||
TICKTICK_TOKEN=
|
||||
HOME_ASSISTANT_ACTION_TASK_PROJECT_ID=
|
||||
|
||||
# Optional: Modbus polling (global kill-switch; default false — opt-in).
|
||||
# MODBUS_POLLING_ENABLED=false
|
||||
|
||||
# Optional: MQTT broker connection.
|
||||
# Leave MQTT_ENABLED=false (or unset) when MQTT is not needed.
|
||||
MQTT_ENABLED=false
|
||||
MQTT_BROKER_HOST=
|
||||
MQTT_BROKER_PORT=1883
|
||||
MQTT_USERNAME=
|
||||
MQTT_PASSWORD=
|
||||
MQTT_TLS_ENABLED=false
|
||||
# MQTT_CLIENT_ID must be a non-empty ASCII slug; use a distinct value per deployment.
|
||||
MQTT_CLIENT_ID=home-automation
|
||||
|
||||
# Optional: Home Assistant MQTT Discovery.
|
||||
# Requires MQTT_ENABLED=true and a running MQTT broker.
|
||||
# HA_DISCOVERY_PREFIX is the config topic prefix (must be "homeassistant" for HA auto-discovery).
|
||||
# HA_STATE_TOPIC_PREFIX is the prefix for state/availability topics (separate from discovery).
|
||||
HA_DISCOVERY_ENABLED=false
|
||||
HA_DISCOVERY_PREFIX=homeassistant
|
||||
HA_STATE_TOPIC_PREFIX=home_automation
|
||||
|
||||
# Optional: DSMR smart-meter ingest via MQTT (M6; default off — opt-in).
|
||||
# Set DSMR_INGEST_ENABLED=true to subscribe to the DSMR MQTT topic and
|
||||
# store 10-second downsampled telegram frames in dsmr_reading.
|
||||
# DSMR_INGEST_ENABLED=false
|
||||
# DSMR_MQTT_TOPIC=dsmr/json
|
||||
# DSMR_SAMPLE_INTERVAL_S=10
|
||||
# DSMR tariff topic: publishes "1" (dal/off-peak) or "2" (normal/peak).
|
||||
# Empty string disables tariff-aware pricing (buy/sell_price_now always show normal rate).
|
||||
# DSMR_TARIFF_TOPIC=dsmr/meter-stats/electricity_tariff
|
||||
|
||||
# Optional: Tibber dynamic pricing credentials (M6).
|
||||
# Only used when an active energy contract with kind=tibber is configured.
|
||||
# TIBBER_API_TOKEN=
|
||||
# TIBBER_HOME_ID=
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
name: frontend
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- "**"
|
||||
pull_request:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
frontend:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Node
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: "22"
|
||||
cache: npm
|
||||
cache-dependency-path: frontend/package-lock.json
|
||||
|
||||
- name: Install dependencies
|
||||
working-directory: frontend
|
||||
run: npm ci
|
||||
|
||||
- name: Check codegen is in sync
|
||||
working-directory: frontend
|
||||
run: |
|
||||
npm run codegen
|
||||
git diff --exit-code src/api/schema.d.ts
|
||||
|
||||
- name: Lint
|
||||
working-directory: frontend
|
||||
run: npm run lint
|
||||
|
||||
- name: Type-check
|
||||
working-directory: frontend
|
||||
run: npm run typecheck
|
||||
|
||||
- name: Test
|
||||
working-directory: frontend
|
||||
run: npm run test
|
||||
|
||||
- name: Build
|
||||
working-directory: frontend
|
||||
run: npm run build
|
||||
@@ -0,0 +1,208 @@
|
||||
# AGENTS.md — Home Automation Backend
|
||||
|
||||
本文件是本仓库 coding agent 指引的 **single source of truth**;`CLAUDE.md` 通过符号链接指向本文件。它定义本项目的**工作流程、文档位置、commit 规范**。支持对应项目指引的 agent 在动手前应完整读取本文件。
|
||||
|
||||
## 项目速览
|
||||
|
||||
- 个人用 home-automation 应用:**FastAPI + React SPA + SQLite + SQLAlchemy + Alembic**,前后端同源托管。
|
||||
- 单 admin 鉴权(Argon2 + server-side session cookie),runtime config 落 `app_config` 表。
|
||||
- 模块:public IPv4 monitor、SMTP 通知、location / poo recorder、Home Assistant in/out、TickTick OAuth、Modbus / DSMR 能耗采集、MQTT / HA Discovery、动态电价与电费计算。
|
||||
- 已发布 `v1.5.1`。M1、M2、M4-M7 已完成;M3 token / 移动端仍为远期方向。
|
||||
- **当前现实**:已收敛为单一 `app.db`、一套 DeclarativeBase 和一条 Alembic 链;只有历史数据迁移 runbook 会读取旧 location / poo 数据库。
|
||||
- 明确不做:Notion 模块。
|
||||
|
||||
## 文档地图与「开工前必读」
|
||||
|
||||
文档都在 `docs/`:
|
||||
|
||||
| 路径 | 作用 |
|
||||
| --- | --- |
|
||||
| `docs/roadmap.md` | 全局规划与里程碑总览 |
|
||||
| `docs/design/README.md` | **协作契约**:任务卡格式、原子任务定义、校验闸门、数据安全红线 |
|
||||
| `docs/design/m1-db-consolidation.md` | M1 原子任务(含真实代码现状盘点 + 人工 runbook) |
|
||||
| `docs/design/m2-frontend-v2.md` | M2 原子任务 + API 契约 + 前端校验闸门 |
|
||||
| `docs/design/m3-token-mobile.md` | M3(远期,暂缓) |
|
||||
| `docs/design/m4-login-hardening.md` | M4 登录加固(已完成) |
|
||||
| `docs/design/m5-iot-energy.md` | M5 IoT / 能耗采集(已完成) |
|
||||
| `docs/design/m6-tibber-dynamic-energy.md` | M6 动态电价、DSMR 与电费计算(已完成) |
|
||||
| `docs/design/m7-meter-epochs-archival.md` | M7 电表生命周期 / 换表归档(已完成) |
|
||||
| `docs/*.md`(auth / public-ip-monitor / location-recorder …) | 各模块说明,按需读 |
|
||||
|
||||
**开工时读取顺序**:
|
||||
1. `docs/design/README.md`(每轮都读,它是流程与验收的共同契约)。
|
||||
2. 本轮对应的 milestone 文档(如 `docs/design/m1-db-consolidation.md`),定位要做的任务卡。
|
||||
3. 任务卡 `Files` 列出的源文件 + 该模块的 `docs/*.md`(按需)。
|
||||
4. `docs/roadmap.md` 仅在需要全局视角时读。
|
||||
|
||||
## 工作流程
|
||||
|
||||
### 实现模式(由用户的提示词决定)
|
||||
|
||||
- **默认逐步**:给一个 milestone 文档,按其中原子任务**一步一步**实现。
|
||||
- **(a) 只实现一步**:用户说"只实现一步 / 这一个任务"时,**只做那一个任务卡**,跑完校验闸门后停下,等用户确认,不要顺手往下做。
|
||||
- **(b) 完成整个 milestone**:仅当用户在提示词里**显式要求启用 sub-agent**时,才起 implementer / reviewer / fixer sub-agent(按下方**『默认能力档位』**选择模型,用户人工指定则覆盖),按任务依赖顺序跑完整条链。
|
||||
- **Sub-agent 纪律**:只在用户显式要求时才 spawn sub-agent;单步/小改动在主线内联完成。当前 harness 支持独立 sub-agent 时,使用其原生机制按下方**『默认能力档位』**派发;不支持时不得假装已创建 sub-agent,应明确说明限制,并仅在用户允许 fallback 时由主 agent 继续。
|
||||
|
||||
### 默认能力档位(实现模式 sub-agent;可被人工指定覆盖)
|
||||
|
||||
起 implementer / reviewer / fixer sub-agent 时,**默认**按下列能力档位选择当前 harness 支持的模型,无需用户每次人工指定:
|
||||
|
||||
| 角色 | 通用模型要求 | 推理档位 | Harness 示例(非强制) |
|
||||
| --- | --- | --- | --- |
|
||||
| **Implementer** | 平衡型代码实现模型 | `medium` 或等效档位 | Claude Code:Sonnet;Codex/OpenAI:GPT-5.6 Terra (`gpt-5.6-terra`) |
|
||||
| **Fixer**(返工) | 平衡型代码实现模型 | `medium` 或等效档位 | Claude Code:Sonnet;Codex/OpenAI:GPT-5.6 Terra (`gpt-5.6-terra`) |
|
||||
| **Reviewer** | 当前 harness 支持的最强通用推理 / 代码模型 | `extra-high` / `xhigh` 或等效档位 | Claude Code:Opus;Codex/OpenAI:GPT-5.6 Sol (`gpt-5.6-sol`) |
|
||||
|
||||
- **示例非强制**:示例模型只表示当前推荐映射,不构成跨 harness 的硬性模型 ID;当前 harness 不支持时,选择最符合「通用模型要求」的可用模型。
|
||||
- **选择优先级**:用户显式指定 > 当前 harness 的原生角色配置 > 上表的 harness 示例 > 按通用模型要求自动选择。
|
||||
- **推理档位说明**:若 harness 提供独立的 reasoning-effort 设置,按上表设置;若不提供,在 spawn prompt 中明确 implementer/fixer 按平衡深度思考,reviewer 按对抗性外部审计强度复核。
|
||||
|
||||
### 角色(Orchestrator → Implementer → Reviewer → Fixer)
|
||||
|
||||
- 我(主线)= **Orchestrator**:挑依赖已满足的下一个任务、派发、转述结果、维护任务 `Status`。
|
||||
- **Implementer**(平衡型代码实现模型,medium 或等效档位):一次一个任务,严格按任务卡,不扩范围。
|
||||
- **Reviewer**(最强通用推理 / 代码模型,extra-high / xhigh 或等效档位):实现完成后起 Reviewer sub-agent,按任务卡 `Acceptance criteria` + `Reviewer checklist` 复核、**独立重跑校验闸门**,驱动返工直到本轮 PASS。
|
||||
- **Fixer**(平衡型代码实现模型,medium 或等效档位):按 reviewer 的编号返工清单返工;**每轮返工起一个干净的 Fixer**(与首次实现的 Implementer 分开冷启动),先读对应 `review-notes/<task>-review-<n>.md` 再改。
|
||||
|
||||
#### Reviewer 盲审纪律(M1 教训)
|
||||
|
||||
M1 里 review **从未触发过一次 rework**,根因是 orchestrator 把自己的结论 / 辩护喂给了 reviewer,造成 context bleed、review 沦为橡皮图章。所以:
|
||||
|
||||
- reviewer 必须**使用全新、独立的 sub-agent / thread 冷启动,并最小化喂料**——spawn prompt 只给:① 任务卡(`Acceptance criteria` + `Reviewer checklist`)、② 对应的 `review-notes/<task>-impl|rework-<n>.md` 路径、③ 要审的 diff / commit 范围。
|
||||
- **不要**在 prompt 里塞 orchestrator 自己的判断、"我觉得没问题"、对实现选择的辩护,或上一轮 reviewer 的倾向性结论。让它**独立得出结论、独立重跑校验闸门**。
|
||||
- 事后另起的整库**独立盲审**(如对抗复审)同理:使用全新独立的 agent / thread、最小上下文,把它当"**外部审计**"而非"确认自己没错"。
|
||||
|
||||
### 校验闸门(每个任务结束都要全绿)
|
||||
|
||||
根目录、激活 `.venv` 后:
|
||||
```bash
|
||||
pytest # 权威闸门(CI 跑的就是它)
|
||||
ruff check . # line-length=100
|
||||
python scripts/export_openapi.py && git diff --exit-code openapi/ # 改了路由/schema 才需要,且产物须入库
|
||||
```
|
||||
前端任务(M2)在 `frontend/` 下另跑 `npm run lint && npm run typecheck && npm run test && npm run build`(详见 m2 文档 §8)。
|
||||
**不过闸门就不算完成**,不得跳过、不得留红给下一轮。
|
||||
|
||||
**Repo-meta 例外**:纯文档、agent 指引、符号链接等不影响可执行代码、构建与 API 契约的变更,可在用户明确同意时跳过代码闸门。仍须完成针对性校验(如链接目标、文件类型、diff 与 Git 状态),并在结果中明确记录未运行哪些闸门。
|
||||
|
||||
#### API 契约同步:`openapi/` 与 `schema.d.ts` 是**两步**(v1.4.0 后教训)
|
||||
|
||||
**只跑 `export_openapi.py` 不够。** 前端的 `frontend/src/api/schema.d.ts` 是由 `openapi/openapi.json` 二次生成的,CI(`.github/workflows/frontend.yml` 的 *Check codegen is in sync*)会重跑 codegen 并 `git diff --exit-code src/api/schema.d.ts`。漏了第二步 → 本地闸门全绿、远端 CI 红。真出过:`d07a083` 改了 `/api/energy/prices` 的 docstring 并同步了 `openapi.json`,但没重跑 codegen,只差一行注释就把 CI 挂了。
|
||||
|
||||
所以**只要动了路由 / Pydantic schema / 路由 docstring**(docstring 也会进 OpenAPI description!),两步都要跑,两个产物都要入库:
|
||||
|
||||
```bash
|
||||
# 1) 后端契约
|
||||
python scripts/export_openapi.py && git diff --exit-code openapi/
|
||||
# 2) 前端类型(在 frontend/ 下)
|
||||
npm run codegen && git diff --exit-code src/api/schema.d.ts
|
||||
```
|
||||
|
||||
- 判据:`git diff --exit-code openapi/` 有输出 → **必然**还要跑一次 `npm run codegen`。
|
||||
- 反过来也成立:`schema.d.ts` 不要手改,它是生成物。
|
||||
- Reviewer 审"动了路由 / schema / 路由 docstring"类任务时,把**这两个产物是否都已重新生成并入库**当作 acceptance 的一部分。
|
||||
|
||||
### 构建上下文完整性(M1 Dockerfile 教训)
|
||||
|
||||
`docker build` **不在 pytest/ruff 闸门里**——M1 删了 `alembic_location/poo` 后忘了同步 `Dockerfile` 的 `COPY`,单元闸门全绿却把坏掉的镜像构建一路漏到 release tag。所以:
|
||||
|
||||
- 任务**删除 / 移动 / 重命名文件或目录**时,必须 grep 构建清单是否还在引用它们:`Dockerfile`(尤其 `COPY` 源)、`docker/`、`*.ini`、CI workflow、`requirements*.txt` 等。
|
||||
- 已有回归测试 `tests/test_deployment.py::test_dockerfile_copy_sources_exist` 守"Dockerfile `COPY` 源必须存在于构建上下文";新增 / 改动 `COPY` 时确保它仍覆盖得到。
|
||||
- Reviewer 审"删 / 移文件"类任务时,**必须顺带核对构建清单引用**,把它当 acceptance 的一部分。
|
||||
|
||||
## 每轮简报(`review-notes/`)
|
||||
|
||||
由 milestone 任务卡驱动的每轮实现、返工或 review,都要在 `review-notes/` 下产出**中文简报**。该目录**已在 `.gitignore` 忽略**,纯本地、不入库——它是 agent 之间和与人之间的交接载体,不是仓库产物。纯讨论、只读分析与不进入正式任务链的 repo-meta 变更无需产出简报,除非用户明确要求。
|
||||
|
||||
- **实现 / 返工简报**:每轮实现完成后(无论首次实现还是返工),写一份。文件名建议 `<task-id>-impl-<n>.md` / `<task-id>-rework-<n>.md`(如 `M1-T03-impl-1.md`、`M1-T03-rework-1.md`)。至少包含:
|
||||
1. **本轮修改的具体内容**(改了哪些文件、做了什么、为什么)。
|
||||
2. **自动化测试结果**(`pytest` / `ruff` / 前端闸门的实际输出或结论,通过/失败逐项写清)。
|
||||
3. **若需人工 walkthrough**:写明具体步骤(怎么启动、点哪里、预期看到什么);若无需人工验证,明确写"无需人工 walkthrough"。
|
||||
- **review 简报**:每轮 review 后写一份,文件名建议 `<task-id>-review-<n>.md`(如 `M1-T03-review-1.md`)。至少包含:评审结论(`PASS` 或带编号的返工清单)、对照任务卡 `Acceptance criteria` + `Reviewer checklist` 的逐条核对、reviewer 独立重跑校验闸门的结果。
|
||||
|
||||
**用途**:① reviewer 审核时参考对应的实现简报;② implementer 返工时参考对应的 review 简报;③ 人类(用户)通读这些简报确认有无问题。简报之间用文件名里的 `<task-id>` 与轮次 `<n>` 对应起来。
|
||||
|
||||
### Orchestrator 派发契约(让简报真正被读到)
|
||||
|
||||
**关键**:sub-agent 冷启动、不继承主线上下文,**不会因为本文件提到简报就自动去读**对应文件。简报能流转,靠的是 orchestrator(主线)在**每次 spawn 时把路径显式写进 prompt**,而不是被动约定。所以派发时必须做到:
|
||||
|
||||
- **显式告诉它「先读哪个简报」**:
|
||||
- 派 implementer 做**首次实现** → 传任务卡位置(milestone 文档路径 + task id);无前置简报。
|
||||
- 派 implementer 做**返工** → 必须传对应的 `review-notes/<task>-review-<n>.md` 路径,并要求**先读它**再改。
|
||||
- 派 reviewer → 必须传对应的 `review-notes/<task>-impl|rework-<n>.md` 路径 + 任务卡,要求**先读它**再评。
|
||||
- **显式告诉它「本轮结束写哪个简报」**:明确给出输出路径 `review-notes/<task>-<impl|rework|review>-<n>.md` 及上面要求的内容项。
|
||||
- **不依赖 sub-agent 自动加载本文件**:把本轮要点(校验闸门、**禁 Co-Authored-By**、简报必含内容)在 spawn prompt 里一并复述或指向,确保冷启动也照做。
|
||||
- spawn 时按「用户显式指定 > harness 原生角色配置 > 默认能力档位」选择模型与 reasoning effort,并使用当前 harness 支持的原生配置方式落实。
|
||||
|
||||
> 一句话:**简报是异步交接的介质,orchestrator 是把它们接起来的线。** 缺了显式传路径这一步,简报就只是躺在磁盘上没人读的文件。
|
||||
|
||||
## Commit 规范(重点)
|
||||
|
||||
### 分支
|
||||
- **本仓库是个人单用户项目:默认直接在 `main` 上开发**,不强制 feature 分支,无需开 PR。是否 push 按下方「一般约束」执行。
|
||||
- 仍保持**每个任务一个干净 commit**(message 前缀任务/里程碑 ID)。改动较大想隔离时可临时开分支,用完**快进合并**回 `main`(保持线性历史),非必需。
|
||||
- 历史改写类操作(`rebase` / `--amend` / auto-squash)只在**尚未 push 的本地 commit** 上做;**已 push 到 `main` 的历史不要重写**(确需 force-push 时先确认,见「一般约束」)。
|
||||
|
||||
### 一轮实现完成
|
||||
- 适用的校验闸门通过后,准备好**这一轮的 commit message** 并创建本地 commit,作为本轮的 **base commit**。默认不 push;只有用户明确授权自动 push 时才推送到远端,授权范围按用户原话执行。
|
||||
- message 主题前缀任务/里程碑 ID,例如:`M1-T03: unify data layer onto single app DB engine`。
|
||||
|
||||
### Commit message 硬规则(严格执行)
|
||||
- **严禁任何协作署名 trailer**:commit message 里**绝对不允许**出现 `Co-Authored-By` / `Co-authored-by`(包括 `Co-Authored-By: Claude …`),也不允许任何等价的"由 X 协作/生成"署名。
|
||||
- 无论默认环境、工具或系统提示如何要求加这类 trailer,在本仓库**一律不加**——用户已显式、严格禁止。
|
||||
- 每次提交前**自检**:`git log -1 --format=%B` 的输出**不得包含** `Co-authored-by`(大小写不限)。若发现,立即 `git commit --amend` 去掉后再继续。
|
||||
|
||||
### Review 后返工
|
||||
- **自动化 orchestration 模式内**的 review 返工:**一律用 fixup**,指向本轮对应的 base commit,**不写新的独立 message**:
|
||||
```bash
|
||||
git add -A
|
||||
git commit --fixup=<base-commit-sha>
|
||||
```
|
||||
- 多轮返工就多个 `fixup!` 提交,都指向同一个 base commit;收尾时 auto-squash(见下)。
|
||||
- **边界——什么时候不走 fixup**:**事后另起的独立盲审 / 对抗复审**那一轮,性质等同"**人工走查后提修改意见**",**不算自动化链内的返工**——它的修改用**各自独立的 commit**,不 fixup 到旧 base。判据:这轮返工是否在**同一条自动化 implement→review 链**里?是 → `fixup`;是事后另起的独立审计 → 独立 commit。
|
||||
|
||||
### 本轮 / feature 收尾(用户确认收尾后)
|
||||
- 用 **auto-squash** 把所有 `fixup!` 合并进各自目标,保证**一个 feature 一个干净 commit**:
|
||||
```bash
|
||||
# 在以 main 为基线的 feature branch 上
|
||||
GIT_SEQUENCE_EDITOR=true git rebase -i --autosquash main
|
||||
|
||||
# 直接在 main 上整理尚未 push 的本地提交
|
||||
GIT_SEQUENCE_EDITOR=true git rebase -i --autosquash origin/main
|
||||
```
|
||||
- 执行前确认选定的基线位于 base commit 之前,以便 base commit 与对应 `fixup!` 都进入 rebase 范围。用 `GIT_SEQUENCE_EDITOR=true` 让它**非交互**执行(不弹编辑器,自动接受 autosquash 排好的 todo)。
|
||||
- autosquash **改写历史**:仅在 push / 开 PR **之前**做。若该分支已 push,需要 force-push——属对外操作,**先取得用户确认再做**。
|
||||
|
||||
### 一般约束
|
||||
- **个人单用户仓库:默认直接在 `main` 上开发并创建本地 commit**。默认不 push;只有用户明确授权自动 push 时才推送到远端。
|
||||
- 始终需要**单独、明确授权**的操作:**force-push / 改写已推送历史**,以及**打 tag**(会触发镜像 CI / 对外发布;且打 tag 前须按下方「发版前置走查」真跑一次 `docker build`)。
|
||||
|
||||
## 发版前置走查(打 tag 前必做)
|
||||
|
||||
单元闸门绿 ≠ 真的能跑、能构建、能用。M1 出过"绿了但 docker 构建坏了"的事故,所以**打版本 tag(触发镜像 CI)之前**,除了 `pytest` / `ruff` 全绿,还要:
|
||||
|
||||
- **真起 app**:迁移(`python -m scripts.run_migrations`)→ `uvicorn app.main:app ...`,确认能正常启动、关键路由不 500。
|
||||
- **真跑镜像构建**:本地 `docker build`(多阶段就跑完整条),确认构建通过、`COPY` 源都在。
|
||||
- **关键功能人工瞄一眼**:尤其前端 / 可视化类(M2 的热力图、首页地图)——自动闸门判断不了"渲染对不对、UX 顺不顺",这部分**靠看跑起来的 app,不靠读代码**。
|
||||
- 上述任一不过 → **不打 tag**。tag 一旦 push 会触发 docker 镜像 CI / 对外发布,属对外操作,**先确认**。
|
||||
|
||||
## 数据安全红线(不可违反)
|
||||
|
||||
- 任何脚本 / migration **都不得删除或覆盖用户数据文件**(旧 `.db`、备份、volume)。删除只能是人工、事后、保留归档的独立步骤(见 `docs/design/m1-db-consolidation.md` §6 runbook)。
|
||||
- 涉及历史数据的迁移**先在备份副本上演练**;迁移脚本必须幂等且搬完对账行数。
|
||||
- Review 时只要发现"删文件 / drop 有数据的表 / truncate"出现在自动化任务里,直接判返工。
|
||||
|
||||
## 常用命令
|
||||
|
||||
```bash
|
||||
# 环境
|
||||
python -m venv .venv && source .venv/bin/activate && pip install -r dev-requirements.txt
|
||||
# 迁移(初始化/适配 DB)
|
||||
python -m scripts.run_migrations
|
||||
# 起服务
|
||||
uvicorn app.main:app --reload --host 0.0.0.0 --port 8000
|
||||
# 测试 / lint / OpenAPI 导出
|
||||
pytest
|
||||
ruff check .
|
||||
python scripts/export_openapi.py
|
||||
```
|
||||
@@ -1,137 +0,0 @@
|
||||
# CLAUDE.md — Home Automation Backend
|
||||
|
||||
本文件每次会话自动加载。它定义本项目的**工作流程、文档位置、commit 规范**。请在动手前先读完。
|
||||
|
||||
## 项目速览
|
||||
|
||||
- 个人用 home-automation 后端:**FastAPI + SQLite + SQLAlchemy + Alembic**,服务端模板(Jinja,M2 将换成 React SPA)。
|
||||
- 单 admin 鉴权(Argon2 + server-side session cookie),runtime config 落 `app_config` 表。
|
||||
- 模块:public IPv4 monitor、SMTP 通知、location recorder、poo recorder、Home Assistant in/out、TickTick OAuth。
|
||||
- 已发布 `v1.0.3`。下一阶段方向:**M1 单库化 → M2 React 前端 → M3 token/移动端(远期,M2 后再说)**。
|
||||
- **当前现实**:在 M1 完成前仍是**三个独立 SQLite 库**(app / location / poo),三套 DeclarativeBase、三条 Alembic 链。不要假设已经单库——以代码现状为准。
|
||||
- 明确不做:Notion 模块。
|
||||
|
||||
## 文档地图与「开工前必读」
|
||||
|
||||
文档都在 `docs/`:
|
||||
|
||||
| 路径 | 作用 |
|
||||
| --- | --- |
|
||||
| `docs/roadmap.md` | 全局规划与里程碑总览 |
|
||||
| `docs/design/README.md` | **协作契约**:任务卡格式、原子任务定义、校验闸门、数据安全红线 |
|
||||
| `docs/design/m1-db-consolidation.md` | M1 原子任务(含真实代码现状盘点 + 人工 runbook) |
|
||||
| `docs/design/m2-frontend-v2.md` | M2 原子任务 + API 契约 + 前端校验闸门 |
|
||||
| `docs/design/m3-token-mobile.md` | M3(远期,暂缓) |
|
||||
| `docs/*.md`(auth / public-ip-monitor / location-recorder …) | 各模块说明,按需读 |
|
||||
|
||||
**开工时读取顺序**:
|
||||
1. `docs/design/README.md`(每轮都读,它是流程与验收的共同契约)。
|
||||
2. 本轮对应的 milestone 文档(如 `docs/design/m1-db-consolidation.md`),定位要做的任务卡。
|
||||
3. 任务卡 `Files` 列出的源文件 + 该模块的 `docs/*.md`(按需)。
|
||||
4. `docs/roadmap.md` 仅在需要全局视角时读。
|
||||
|
||||
## 工作流程
|
||||
|
||||
### 实现模式(由用户的提示词决定)
|
||||
|
||||
- **默认逐步**:给一个 milestone 文档,按其中原子任务**一步一步**实现。
|
||||
- **(a) 只实现一步**:用户说"只实现一步 / 这一个任务"时,**只做那一个任务卡**,跑完校验闸门后停下,等用户确认,不要顺手往下做。
|
||||
- **(b) 完成整个 milestone**:仅当用户在提示词里**显式要求启用 sub-agent 并指定模型**时,才用指定模型起 implementer sub-agent,按任务依赖顺序跑完整条链。
|
||||
- **Sub-agent 纪律**:只在用户显式要求时才 spawn sub-agent;单步/小改动在主线内联完成。起 sub-agent 时用用户**指定的模型**(Agent 工具的 `model` 覆盖)。
|
||||
|
||||
### 角色(Orchestrator → Implementer → Reviewer)
|
||||
|
||||
- 我(主线)= **Orchestrator**:挑依赖已满足的下一个任务、派发、转述结果、维护任务 `Status`。
|
||||
- **Implementer**(便宜模型,用户指定):一次一个任务,严格按任务卡,不扩范围。
|
||||
- **Reviewer**(强模型,用户指定):实现完成后起 Reviewer sub-agent,按任务卡 `Acceptance criteria` + `Reviewer checklist` 复核、**独立重跑校验闸门**,驱动 implementer 返工直到本轮 PASS。
|
||||
|
||||
### 校验闸门(每个任务结束都要全绿)
|
||||
|
||||
根目录、激活 `.venv` 后:
|
||||
```bash
|
||||
pytest # 权威闸门(CI 跑的就是它)
|
||||
ruff check . # line-length=100
|
||||
python scripts/export_openapi.py && git diff --exit-code openapi/ # 改了路由/schema 才需要,且产物须入库
|
||||
```
|
||||
前端任务(M2)在 `frontend/` 下另跑 `npm run lint && npm run typecheck && npm run test && npm run build`(详见 m2 文档 §8)。
|
||||
**不过闸门就不算完成**,不得跳过、不得留红给下一轮。
|
||||
|
||||
## 每轮简报(`review-notes/`)
|
||||
|
||||
每轮工作都要在 `review-notes/` 下产出**中文简报**。该目录**已在 `.gitignore` 忽略**,纯本地、不入库——它是 agent 之间和与人之间的交接载体,不是仓库产物。
|
||||
|
||||
- **实现 / 返工简报**:每轮实现完成后(无论首次实现还是返工),写一份。文件名建议 `<task-id>-impl-<n>.md` / `<task-id>-rework-<n>.md`(如 `M1-T03-impl-1.md`、`M1-T03-rework-1.md`)。至少包含:
|
||||
1. **本轮修改的具体内容**(改了哪些文件、做了什么、为什么)。
|
||||
2. **自动化测试结果**(`pytest` / `ruff` / 前端闸门的实际输出或结论,通过/失败逐项写清)。
|
||||
3. **若需人工 walkthrough**:写明具体步骤(怎么启动、点哪里、预期看到什么);若无需人工验证,明确写"无需人工 walkthrough"。
|
||||
- **review 简报**:每轮 review 后写一份,文件名建议 `<task-id>-review-<n>.md`(如 `M1-T03-review-1.md`)。至少包含:评审结论(`PASS` 或带编号的返工清单)、对照任务卡 `Acceptance criteria` + `Reviewer checklist` 的逐条核对、reviewer 独立重跑校验闸门的结果。
|
||||
|
||||
**用途**:① reviewer 审核时参考对应的实现简报;② implementer 返工时参考对应的 review 简报;③ 人类(用户)通读这些简报确认有无问题。简报之间用文件名里的 `<task-id>` 与轮次 `<n>` 对应起来。
|
||||
|
||||
### Orchestrator 派发契约(让简报真正被读到)
|
||||
|
||||
**关键**:sub-agent 冷启动、不继承主线上下文,**不会因为本文件提到简报就自动去读**对应文件。简报能流转,靠的是 orchestrator(主线)在**每次 spawn 时把路径显式写进 prompt**,而不是被动约定。所以派发时必须做到:
|
||||
|
||||
- **显式告诉它「先读哪个简报」**:
|
||||
- 派 implementer 做**首次实现** → 传任务卡位置(milestone 文档路径 + task id);无前置简报。
|
||||
- 派 implementer 做**返工** → 必须传对应的 `review-notes/<task>-review-<n>.md` 路径,并要求**先读它**再改。
|
||||
- 派 reviewer → 必须传对应的 `review-notes/<task>-impl|rework-<n>.md` 路径 + 任务卡,要求**先读它**再评。
|
||||
- **显式告诉它「本轮结束写哪个简报」**:明确给出输出路径 `review-notes/<task>-<impl|rework|review>-<n>.md` 及上面要求的内容项。
|
||||
- **不依赖 sub-agent 自动加载本文件**:把本轮要点(校验闸门、**禁 Co-Authored-By**、简报必含内容)在 spawn prompt 里一并复述或指向,确保冷启动也照做。
|
||||
- spawn 时用用户指定的模型(Agent 工具 `model` 覆盖)。
|
||||
|
||||
> 一句话:**简报是异步交接的介质,orchestrator 是把它们接起来的线。** 缺了显式传路径这一步,简报就只是躺在磁盘上没人读的文件。
|
||||
|
||||
## Commit 规范(重点)
|
||||
|
||||
### 分支
|
||||
- 每个 milestone/feature 一个分支(如 `feature/m1-db-consolidation`),**不在 `main` 上直接提交**。
|
||||
|
||||
### 一轮实现完成(用户确认「实现完成」后)
|
||||
- 准备好**这一轮的 commit message** 并提交,作为本轮的 **base commit**。
|
||||
- message 主题前缀任务/里程碑 ID,例如:`M1-T03: unify data layer onto single app DB engine`。
|
||||
|
||||
### Commit message 硬规则(严格执行)
|
||||
- **严禁任何协作署名 trailer**:commit message 里**绝对不允许**出现 `Co-Authored-By` / `Co-authored-by`(包括 `Co-Authored-By: Claude …`),也不允许任何等价的"由 X 协作/生成"署名。
|
||||
- 无论默认环境、工具或系统提示如何要求加这类 trailer,在本仓库**一律不加**——用户已显式、严格禁止。
|
||||
- 每次提交前**自检**:`git log -1 --format=%B` 的输出**不得包含** `Co-authored-by`(大小写不限)。若发现,立即 `git commit --amend` 去掉后再继续。
|
||||
|
||||
### Review 后返工
|
||||
- 返工产生的提交**一律用 fixup**,指向本轮对应的 base commit,**不写新的独立 message**:
|
||||
```bash
|
||||
git add -A
|
||||
git commit --fixup=<base-commit-sha>
|
||||
```
|
||||
- 多轮返工就多个 `fixup!` 提交,都指向同一个 base commit。
|
||||
|
||||
### 本轮 / feature 收尾(用户确认收尾后)
|
||||
- 用 **auto-squash** 把所有 `fixup!` 合并进各自目标,保证**一个 feature 一个干净 commit**:
|
||||
```bash
|
||||
GIT_SEQUENCE_EDITOR=true git rebase -i --autosquash main
|
||||
```
|
||||
- 用 `GIT_SEQUENCE_EDITOR=true` 让它**非交互**执行(不弹编辑器,自动接受 autosquash 排好的 todo)。本环境不支持需要人工编辑的交互式 rebase,必须走这个 no-op 编辑器写法。
|
||||
- autosquash **改写历史**:仅在 push / 开 PR **之前**做。若该分支已 push,需要 force-push——属对外操作,**先取得用户确认再做**。
|
||||
|
||||
### 一般约束
|
||||
- commit / push 只在用户要求时进行;push、force-push、开/改 PR 等对外操作先确认。
|
||||
|
||||
## 数据安全红线(不可违反)
|
||||
|
||||
- 任何脚本 / migration **都不得删除或覆盖用户数据文件**(旧 `.db`、备份、volume)。删除只能是人工、事后、保留归档的独立步骤(见 `docs/design/m1-db-consolidation.md` §6 runbook)。
|
||||
- 涉及历史数据的迁移**先在备份副本上演练**;迁移脚本必须幂等且搬完对账行数。
|
||||
- Review 时只要发现"删文件 / drop 有数据的表 / truncate"出现在自动化任务里,直接判返工。
|
||||
|
||||
## 常用命令
|
||||
|
||||
```bash
|
||||
# 环境
|
||||
python -m venv .venv && source .venv/bin/activate && pip install -r dev-requirements.txt
|
||||
# 迁移(初始化/适配 DB)
|
||||
python -m scripts.run_migrations
|
||||
# 起服务
|
||||
uvicorn app.main:app --reload --host 0.0.0.0 --port 8000
|
||||
# 测试 / lint / OpenAPI 导出
|
||||
pytest
|
||||
ruff check .
|
||||
python scripts/export_openapi.py
|
||||
```
|
||||
+20
-4
@@ -1,3 +1,20 @@
|
||||
# Stage 1: build the React SPA.
|
||||
# Pin to the native build host ($BUILDPLATFORM) so the Node/V8 build never runs
|
||||
# under QEMU during multi-arch builds — emulated Node crashes V8's baseline JIT
|
||||
# (SIGTRAP / exit 133 on `npm ci`). The dist/ output is static JS/CSS, i.e.
|
||||
# architecture-independent, so building it once and COPYing it into each
|
||||
# target-arch runtime stage is both correct and avoids the emulator entirely.
|
||||
FROM --platform=$BUILDPLATFORM node:22-slim AS frontend-build
|
||||
|
||||
WORKDIR /frontend
|
||||
|
||||
COPY frontend/package.json frontend/package-lock.json ./
|
||||
RUN npm ci
|
||||
|
||||
COPY frontend/ ./
|
||||
RUN npm run build
|
||||
|
||||
# Stage 2: python runtime (no node)
|
||||
FROM python:3.12-slim
|
||||
|
||||
ENV PYTHONDONTWRITEBYTECODE=1 \
|
||||
@@ -11,15 +28,14 @@ RUN pip install --no-cache-dir -r requirements.txt
|
||||
COPY app ./app
|
||||
COPY alembic_app ./alembic_app
|
||||
COPY alembic_app.ini ./
|
||||
COPY alembic_location ./alembic_location
|
||||
COPY alembic_location.ini ./
|
||||
COPY alembic_poo ./alembic_poo
|
||||
COPY alembic_poo.ini ./
|
||||
COPY scripts ./scripts
|
||||
COPY docker ./docker
|
||||
COPY README.md ./
|
||||
RUN mkdir -p /app/data
|
||||
|
||||
# Copy the built SPA dist from the frontend-build stage
|
||||
COPY --from=frontend-build /frontend/dist ./frontend/dist
|
||||
|
||||
EXPOSE 8000
|
||||
|
||||
ENTRYPOINT ["/app/docker/entrypoint.sh"]
|
||||
|
||||
@@ -4,16 +4,25 @@
|
||||
|
||||
当前系统已经包含:
|
||||
|
||||
- FastAPI Web 应用与服务端模板页面
|
||||
- FastAPI Web 应用(React SPA 前端 + JSON API)
|
||||
- SQLite + SQLAlchemy + Alembic 的单库结构
|
||||
- username/password + server-side session 鉴权
|
||||
- username/password + server-side session 鉴权(含登录加固,见下文)
|
||||
- runtime config 页面与 app DB 持久化
|
||||
- public IPv4 monitor、历史持久化与定时检查
|
||||
- SMTP 配置、测试发信与 public IPv4 changed 邮件通知
|
||||
- location recorder
|
||||
- poo recorder
|
||||
- Home Assistant inbound / outbound integration
|
||||
- Home Assistant inbound / outbound integration(REST 通道)
|
||||
- TickTick OAuth 与 action task 集成
|
||||
- **Modbus 设备采集**:通过 YAML profile(首个:SDM120 电表)按设备周期轮询 Modbus-TCP 网关,解码工程量并落通用读数表(`modbus_device` + `modbus_reading`)
|
||||
- **MQTT + Home Assistant Discovery**:以可勾选方式把 Modbus 设备/工程量注册为 HA device/entity(含 binary_sensor online),state 周期发布;配置变更可重连重发
|
||||
- **前端侧边栏 + Energy 视图**:侧边导航替换顶栏;Energy 页管理 Modbus 设备、展示最新读数与 Recharts 走势图;Config 页 Accordion 分区展开;Expose 设置勾选 HA 可暴露实体
|
||||
- **DSMR 实时电表接入**:订阅 DSMR Reader 的 `dsmr/json`(每秒一帧)、整帧 JSON blob 按 10 秒降采样落库(`dsmr_reading`)
|
||||
- **通用电价合同层**:YAML profile 定合同结构(manual 固定/双费率 / tibber 动态电价);`EnergyContract`+`EnergyContractVersion` 存 UI 可填的数值,改价加新版本旧版本保留;price strategy 按 kind 出价
|
||||
- **实时买卖电费计算**:每 15 分钟按寄存器差值(`_1`=dal/低、`_2`=normal/高)× 买/卖价算计量电费,快照不可变;日/月/年汇总加固定费减 heffingskorting
|
||||
- **反哺 Home Assistant Energy**:当前买/卖价 + 累计买电支出/卖电收入(`total_increasing`)发成 HA 实体,可直接挂 HA Energy 仪表盘
|
||||
- **多数据源 Meter 与 WarmteLink**:DSMR MQTT 与只读 WarmteLink P1 serial source 统一为 Source → Channel → Binding → Meter;WarmteLink 提供 heating `GJ` 与 hot-water `m³` 的 Decimal history、质量与重连
|
||||
- **热力合同与成本**:electricity / thermal scope 可各有一个 active 合同;热力按 15 分钟账本计算 variable、fixed 与 all-in 成本,并可按需暴露给 HA
|
||||
- pytest 测试与 OpenAPI 导出脚本
|
||||
- Docker / Compose 部署入口
|
||||
|
||||
@@ -30,6 +39,15 @@
|
||||
- public IPv4 当前状态与变化历史
|
||||
- location 记录(`location` 表)
|
||||
- poo 记录(`poo_records` 表)
|
||||
- Modbus 设备定义(`modbus_device` 表)
|
||||
- Modbus 通用读数(`modbus_reading` 表,JSON payload)
|
||||
- HA 实体暴露开关(`exposed_entity_toggle` 表)
|
||||
- DSMR 电表实时读数(`dsmr_reading` 表,整帧 JSON blob,10s 降采样)
|
||||
- 电价合同(`energy_contract` 表)与版本(`energy_contract_version` 表,values JSON)
|
||||
- Tibber 15 分钟电价缓存(`tibber_price` 表,不可变)
|
||||
- 每 15 分钟计量电费(`energy_cost_period` 表,快照价,不可变)
|
||||
- meter source、channel 与 binding(`meter_source`、`meter_source_channel`、`meter_source_binding`)
|
||||
- WarmteLink scalar 历史(`warmtelink_reading`)与热力 15 分钟成本账本(`meter_cost_period`)
|
||||
|
||||
配置层只保留一个数据库环境变量:
|
||||
|
||||
@@ -41,17 +59,19 @@
|
||||
python -m scripts.run_migrations
|
||||
```
|
||||
|
||||
该命令会通过 Alembic 将 `app.db` 初始化或升级到最新 head(含 `location` / `poo_records` 表)。
|
||||
该命令会通过 Alembic 将 `app.db` 初始化或升级到最新 head(包括 Modbus、DSMR、Source/Channel/Binding、WarmteLink、electricity/thermal 合同与成本账本)。
|
||||
|
||||
## 当前目录
|
||||
|
||||
主要目录如下:
|
||||
|
||||
- `app/`: FastAPI 应用代码
|
||||
- `alembic_app/`: App DB 的 Alembic migration 环境(同时管理 `location` / `poo_records` 表)
|
||||
- `app/`: FastAPI 应用代码(包含 JSON API、业务服务、数据模型)
|
||||
- `frontend/`: React SPA 前端(Vite + React + TypeScript + Mantine)
|
||||
- `alembic_app/`: App DB 的唯一 Alembic migration 环境(管理所有 app 表,包括 Modbus、DSMR、Meter source、WarmteLink、合同与成本账本)
|
||||
- `tests/`: pytest 测试
|
||||
- `docs/`: 当前系统说明文档
|
||||
- `scripts/`: 辅助脚本,例如 OpenAPI 导出
|
||||
- `openapi/`: OpenAPI schema 静态产物(`openapi.json` / `openapi.yaml`),纳入版本控制
|
||||
|
||||
## 依赖管理
|
||||
|
||||
@@ -112,11 +132,63 @@ uvicorn app.main:app --reload --host 0.0.0.0 --port 8000
|
||||
|
||||
启动后可访问:
|
||||
|
||||
- 应用首页:`http://localhost:8000/`
|
||||
- 应用首页(React SPA):`http://localhost:8000/`
|
||||
- 健康检查:`http://localhost:8000/status`
|
||||
- Swagger UI:`http://localhost:8000/docs`
|
||||
- ReDoc:`http://localhost:8000/redoc`
|
||||
|
||||
## 前端 v2(React SPA)
|
||||
|
||||
M2 用 React SPA 取代了原有 Jinja 服务端模板,由 FastAPI 同源托管(同一容器、同一 origin)。
|
||||
|
||||
### 技术栈
|
||||
|
||||
- **Vite + React + TypeScript + Mantine**(组件库)
|
||||
- **TanStack Query**(数据请求/缓存)
|
||||
- **Leaflet / react-leaflet**(地图与热力图)
|
||||
- **Recharts**(Energy 视图走势图,M5 引入)
|
||||
- **openapi-typescript + openapi-fetch**(类型化 API client,由 `openapi/openapi.json` 生成)
|
||||
|
||||
### 本地开发(前端)
|
||||
|
||||
前端开发服务器会把 `/api`、`/location`、`/poo`、`/public-ip`、`/homeassistant`、`/ticktick`、`/status` 等路径代理到后端 FastAPI(`:8000`)。
|
||||
|
||||
```bash
|
||||
cd frontend
|
||||
npm install
|
||||
npm run dev # 启动 Vite dev server(默认 :5173),代理后端
|
||||
```
|
||||
|
||||
### 构建
|
||||
|
||||
```bash
|
||||
cd frontend
|
||||
npm run build # 产出 frontend/dist
|
||||
```
|
||||
|
||||
FastAPI 启动时若 `frontend/dist/index.html` 存在,则自动挂载该目录,并对非 `/api` 路径做 SPA fallback(返回 `index.html`)。该路径可通过环境变量 `SPA_DIST_DIR` 覆盖(默认值为 `frontend/dist`,与多阶段 Dockerfile 中 `COPY` 到 `/app/frontend/dist` 一致)。
|
||||
|
||||
### 类型化 API Client
|
||||
|
||||
前端 API client 由后端 OpenAPI schema 自动生成:
|
||||
|
||||
```bash
|
||||
cd frontend
|
||||
npm run codegen # 从 ../openapi/openapi.json 生成 src/api/schema.d.ts
|
||||
```
|
||||
|
||||
生成物(`src/api/schema.d.ts`)已提交入库,CI 会校验它与 `openapi/openapi.json` 保持同步。
|
||||
|
||||
### 前端校验闸门
|
||||
|
||||
```bash
|
||||
cd frontend
|
||||
npm run lint # ESLint
|
||||
npm run typecheck # TypeScript 类型检查
|
||||
npm run test # Vitest 单元测试
|
||||
npm run build # 构建,确认产出 dist
|
||||
```
|
||||
|
||||
## 数据库与 Alembic
|
||||
|
||||
当前使用单一 SQLite 数据库文件:
|
||||
@@ -124,7 +196,7 @@ uvicorn app.main:app --reload --host 0.0.0.0 --port 8000
|
||||
- App DB:`sqlite:///./data/app.db`
|
||||
- 数据目录:`./data/`
|
||||
|
||||
所有模型(auth / config / public_ip / location / poo)共用同一个 `Base`,均通过单一 Alembic 链管理:
|
||||
所有模型(auth / config / public_ip / location / poo / modbus / expose)共用同一个 `Base`,均通过单一 Alembic 链管理:
|
||||
|
||||
- Alembic 环境:`alembic_app.ini` + `alembic_app/`
|
||||
- 统一 migration job:`python -m scripts.run_migrations`
|
||||
@@ -142,9 +214,9 @@ python -m scripts.migrate_legacy_data
|
||||
|
||||
- 认证模型:`username/password`
|
||||
- 会话模型:server-side session + cookie
|
||||
- 当前主要受保护页面:`/config`
|
||||
- 当前公开页面:`/login`
|
||||
- 当前公开 API:现有业务 API 暂未在这一轮统一收口到 auth 下
|
||||
- 当前受保护入口:React SPA(`/` 等客户端路由)调用 `/api/*` JSON 端点
|
||||
- 当前公开页面:`/login`(SPA 登录页)
|
||||
- 当前公开 API:裸 ingestion 端点(`/location/record`、`/poo/record` 等设备调用端点)暂未收口到 session 保护(M3 再做)
|
||||
|
||||
安全实现的当前边界:
|
||||
|
||||
@@ -152,7 +224,7 @@ python -m scripts.migrate_legacy_data
|
||||
- session cookie 使用 `HttpOnly`
|
||||
- `Secure` 默认随 `APP_ENV` 切换:非 development 时默认开启
|
||||
- `SameSite=Lax`
|
||||
- 登录表单和登出表单都有基础 CSRF 防护
|
||||
- 写请求(POST/PUT/PATCH/DELETE)需携带 `X-CSRF-Token` header(SameSite=Lax + 自定义 header 纵深防御,无需 per-session token 值比对)
|
||||
|
||||
首次启动时,如果 `APP_DATABASE_URL` 对应的 auth DB 里还没有用户,应用会使用:
|
||||
|
||||
@@ -166,12 +238,203 @@ python -m scripts.migrate_legacy_data
|
||||
|
||||
首次登录后会被要求立即修改密码。这个 bootstrap 只用于首个用户落库,不是后续的完整配置管理方案。
|
||||
|
||||
当前前端主要有两条页面路径:
|
||||
React SPA 主要页面路由(客户端路由,均由 FastAPI fallback 到 `index.html`):
|
||||
|
||||
- `/login`
|
||||
- `/config`
|
||||
- `/login`:登录页
|
||||
- `/`:首页(地图热力图主视图)
|
||||
- `/config`:配置页(取代原 Jinja `/config`)
|
||||
- `/records`:记录管理列表页
|
||||
|
||||
无论是本地 `host:port` 还是反向代理后的域名访问,登录成功后都使用相对路径跳转到 `/config`。
|
||||
无论是本地 `host:port` 还是反向代理后的域名访问,登录成功后进入 SPA 首页(`/`)。
|
||||
|
||||
## M4 登录加固
|
||||
|
||||
M4 在基础鉴权之上叠加了三层防御,详细说明见 [`docs/auth.md`](./docs/auth.md)。
|
||||
|
||||
### 防爆破 / 指数退避
|
||||
|
||||
登录失败超过 3 次后进入指数退避(`wait = min(900s, 1s × 2^(failures-3))`),期间请求返回 `429 Too Many Requests`(含 `Retry-After` 响应头);成功登录后自动清零。退避按 **client IP** 与 **username** 双键取较大值,不会因此永久锁定账号(只是延迟,不是封号)。
|
||||
|
||||
- 全局开关:`AUTH_LOGIN_THROTTLE_ENABLED`(CONFIG_FIELDS,默认 `true`)
|
||||
- 反代后需要 `AUTH_TRUST_FORWARDED_FOR=true` 才会读 `X-Forwarded-For`(`.env` 部署级配置,默认 `false`)
|
||||
|
||||
### CLI 逃生通道
|
||||
|
||||
拿到服务器 CLI 权限时,可以在**不依赖任何已存凭据**(无需密码、恢复码)的情况下重置密码、解锁退避、关停 TOTP:
|
||||
|
||||
```bash
|
||||
# 重置密码(不加 --password 则交互式输入,不回显)
|
||||
python -m scripts.admin_cli reset-password admin
|
||||
|
||||
# 解锁退避(被 429 挡住时使用)
|
||||
python -m scripts.admin_cli unlock --all # 清所有退避行
|
||||
python -m scripts.admin_cli unlock --ip 1.2.3.4 # 按 IP 清
|
||||
python -m scripts.admin_cli unlock --username admin # 按 username 清
|
||||
|
||||
# TOTP 相关(需要先启用,见下文)
|
||||
python -m scripts.admin_cli disable-totp admin # 关停 TOTP(零凭据,逃生口)
|
||||
python -m scripts.admin_cli reissue-totp admin # 重新发放 TOTP secret(打印新 URI)
|
||||
|
||||
# 查看用户列表
|
||||
python -m scripts.admin_cli list-admin
|
||||
```
|
||||
|
||||
在 Docker 容器内执行时:
|
||||
|
||||
```bash
|
||||
docker compose exec app python -m scripts.admin_cli <command>
|
||||
```
|
||||
|
||||
### 可选 TOTP 二次验证
|
||||
|
||||
admin 可在 React SPA 设置页(`/config`)自选启用 RFC 6238 TOTP:
|
||||
|
||||
1. 设置页点「启用 TOTP」→ 后端生成 `otpauth://` URI,前端渲染二维码(`qrcode.react`)
|
||||
2. 用 Authenticator App(如 Google Authenticator、Authy)扫码
|
||||
3. 输入当前 6 位动态码确认 → TOTP 启用
|
||||
4. 妥善保存一次性展示的 10 个恢复码(格式 `xxxx-xxxx`)
|
||||
|
||||
启用后,登录需要两步:密码 → 6 位动态码(或恢复码,一次性)。不启用则维持纯密码登录,行为不变。
|
||||
|
||||
恢复码丢失时,可用 CLI 逃生:`python -m scripts.admin_cli disable-totp admin`,随后即可纯密码登录。
|
||||
|
||||
TOTP issuer 标签(显示在 Authenticator 里)通过 `AUTH_TOTP_ISSUER` 环境变量配置(`.env` 部署级),默认回退 `app_name`。
|
||||
|
||||
## M5 Modbus 设备采集 / Energy / MQTT + HA Discovery
|
||||
|
||||
M5 给后端接入家庭 IoT 生态,新增通用 Modbus 采集链路(首个领域:能耗)、MQTT + Home Assistant Discovery 发布,以及前端侧边栏与 Energy 视图。
|
||||
|
||||
### 依赖
|
||||
|
||||
后端新增:
|
||||
- `pymodbus`:Modbus-TCP 客户端(轮询电表等 slave 设备)
|
||||
- `paho-mqtt`:MQTT 客户端(HA Discovery 与 state 发布)
|
||||
- `pyyaml`:YAML profile 加载(设备协议声明式描述)
|
||||
|
||||
前端新增:
|
||||
- `recharts`:Energy 视图走势图
|
||||
|
||||
### Modbus 设备采集
|
||||
|
||||
采集链路采用两层分离:**YAML profile**(协议知识,随代码走)+ **`modbus_device` 数据库行**(部署/可配置信息)+ **`modbus_reading` 通用读数表**(JSON payload 遥测)。
|
||||
|
||||
- **profile**(如 `sdm120.yaml`)描述:读哪些寄存器(FC04 块读)、每个量的 key/unit/device_class/ha_component。纯协议知识,不含 unit_id / friendly_name 等部署项。
|
||||
- **`modbus_device` 行**:friendly_name、网关 host/port、Modbus slave `unit_id`(电表 Meter ID,设备面板可改故落 DB)、选用哪个 profile、采样周期、是否启用。
|
||||
- **`modbus_reading` 行**:device_id FK、recorded_at、payload(JSON,如 `{"voltage": 230.2, "current": 1.3, ...}`)。
|
||||
- 多设备可共享同一 profile(如两块 SDM120 共用 `sdm120` profile,各自独立 unit_id 和 friendly_name)。
|
||||
- APScheduler 后台 job 周期轮询所有 `enabled` 设备,更新 `last_poll_at` / `last_poll_ok`。全局开关 `MODBUS_POLLING_ENABLED`(CONFIG_FIELDS)。
|
||||
|
||||
**手工命令行试读**(不依赖 DB,最快验证网关连通性):
|
||||
|
||||
```bash
|
||||
# 按 profile 解码读一次(验证整套解码链路)
|
||||
python -m scripts.modbus_cli read --host <网关IP> --port 502 --unit 1 --profile sdm120
|
||||
|
||||
# 手工指定请求内容(first-contact 验证,不依赖 profile)
|
||||
python -m scripts.modbus_cli probe --host <网关IP> --port 502 --unit 1 --fc 4 --address 0x0000 --count 2 --decode float32
|
||||
```
|
||||
|
||||
CLI 工具为受控手工验证而设(设备需接市电),仅暴露读功能码(FC03/04),无写寄存器子命令。
|
||||
|
||||
### MQTT + Home Assistant Discovery
|
||||
|
||||
后端作为 MQTT 发布方,把 Modbus 设备/工程量以 HA Discovery 协议自动注册为 device/entity:
|
||||
|
||||
- 每个 Modbus 设备 = 一个 HA device(`unique_id` 锚定于设备 `uuid`,不随改名变)
|
||||
- 各工程量 = sensor entity(device_class/unit 取自 YAML profile);另有 binary_sensor `online`(取 `last_poll_ok`)
|
||||
- 每次轮询成功后推 state topic;连接成功或勾选变更时重发 retained discovery config
|
||||
- `POST /api/config/mqtt/test`:试连 broker 并发布一条测试消息(可用 MQTT Explorer 验证链路)
|
||||
|
||||
**Config 页配置流程**:
|
||||
1. 在 `/config` 的「MQTT」section 填写 broker host/port/username/password,启用 `MQTT_ENABLED`
|
||||
2. 点「发送测试消息」确认 broker 链路通
|
||||
3. 启用 `HA_DISCOVERY_ENABLED`
|
||||
4. 在「Home Assistant Expose」面板勾选要暴露的实体,点「重新发布 discovery」
|
||||
5. 在 Home Assistant 确认对应 device/entity 出现
|
||||
|
||||
### API 端点(M5 新增)
|
||||
|
||||
| 端点 | 用途 |
|
||||
| --- | --- |
|
||||
| `GET /api/modbus/devices` | 列出 Modbus 设备 |
|
||||
| `POST /api/modbus/devices` | 新建设备 |
|
||||
| `GET /api/modbus/devices/{uuid}` | 单个设备 |
|
||||
| `PATCH /api/modbus/devices/{uuid}` | 修改设备(含 enable/disable)|
|
||||
| `DELETE /api/modbus/devices/{uuid}` | 删除设备;有读数时 409 |
|
||||
| `GET /api/modbus/devices/{uuid}/metrics` | 该设备 profile 的量目录(key/unit/device_class)|
|
||||
| `GET /api/modbus/devices/{uuid}/latest` | 最新一条读数 payload |
|
||||
| `GET /api/modbus/devices/{uuid}/readings` | 时间范围读数(start/end/limit),供走势图 |
|
||||
| `POST /api/modbus/devices/{uuid}/test` | 即时试读(不落库) |
|
||||
| `GET /api/modbus/profiles` | 列出可用 profile 名 + 描述 |
|
||||
| `GET /api/expose` | 可暴露实体目录 + 勾选状态 + MQTT/Discovery 状态 |
|
||||
| `PUT /api/expose` | 设置逐 key 暴露开关 |
|
||||
| `POST /api/expose/republish` | 手动重发 discovery |
|
||||
| `POST /api/config/mqtt/test` | 试连 broker 并发布测试消息 |
|
||||
|
||||
### 前端视图(M5 新增)
|
||||
|
||||
- **侧边栏**:把顶栏改为侧边导航(Home / Records / Energy / Config + 主题切换 + 注销),当前路由高亮,移动端可折叠。
|
||||
- **`/energy`(Energy 视图)**:设备 CRUD(新建/编辑/删除,删除有二次确认;有读数时引导改用禁用);最新读数卡片(字段标签/单位取自 profile metrics);时间序列走势图(Recharts,支持电压/电流/功率/电能,带时间范围选择)。
|
||||
- **Config 页 Accordion**:各大 config section 可独立折叠/展开;「Home Assistant Expose」面板按设备分组勾选可暴露实体、显示 MQTT/Discovery 连接状态、「重新发布 discovery」按钮。
|
||||
- SPA 路由新增 `/energy`。
|
||||
|
||||
## M6 DSMR 接入 / 电价合同 / 实时电费计算 / HA Energy 反哺
|
||||
|
||||
M6 在 M5 IoT 基建之上接入 DSMR 实时智能电表数据,建立通用电价合同层,按每 15 分钟算出实际买卖电费并反哺 HA Energy。
|
||||
|
||||
### 依赖
|
||||
|
||||
M6 **不新增任何 Python 依赖**,复用 M5 已有的 `httpx`(Tibber GraphQL)、`paho-mqtt`(DSMR 订阅)、`pyyaml`(pricing profile 加载)、`apscheduler`(抓价 job、计费 job)。
|
||||
|
||||
### DSMR 实时电表接入
|
||||
|
||||
订阅 DSMR Reader 的 `dsmr/json` topic(每秒一帧完整 telegram),整帧存为 JSON blob、按 `dsmr_sample_interval_s`(默认 10 秒)降采样落 `dsmr_reading`(`source_id` 幂等去重)。`dsmr_ingest_enabled`(默认 false,opt-in)。
|
||||
|
||||
### 电价合同层
|
||||
|
||||
- **YAML profile 定结构**(仓库内,不放数值):`manual.yaml`(固定/双费率:buy_normal/dal、sell_normal/dal、energy_tax、ode、固定费、heffingskorting);`tibber.yaml`(动态:source=tibber_api,energy_tax、sell_adjust)
|
||||
- **`EnergyContract` + `EnergyContractVersion`**(UI 填数值):改价 = 加新版本行(带 `effective_from`),旧版本保留(审计链);一次只有一个 active 合同
|
||||
- **price strategy**:`manual` 用双费率常数(`buy = energy_buy_档 + energy_tax`,`sell = sell_档`);`tibber` 用 `tibber_price.total` 作买价(已含税,demo 确认 `total=energy+tax`),`total − energy_tax − sell_adjust` 作卖价(卖价残差 `sell_adjust` 默认 0,待真实账单核定)
|
||||
|
||||
### 每 15 分钟计量电费(不可变)
|
||||
|
||||
APScheduler 1 分钟 tick,取每个闭合 15 分钟窗口的 DSMR 寄存器差值(`delivered_1/2`,`returned_1/2`;`_1`=dal/低,`_2`=normal/高,NL 惯例)× 当时合同版本的 strategy 出价,upsert `energy_cost_period`(快照当时价 + `contract_version_id`)。缺价/缺数据时标 `degraded`。`POST /api/energy/costs/recompute` 显式重算。
|
||||
|
||||
日/月/年汇总 = Σnet + 固定费(network_fee + management_fee 按月→天 × 天数)- heffingskorting(按年→天 × 天数),读时计算、不落表。能源税 `energy_tax` 参考值约 0.1108 EUR/kWh(2026 第一档含 VAT,待真实账单核定;该值由 UI 填入合同版本,YAML profile 仅声明字段 unit,代码无写死默认数值)。
|
||||
|
||||
### Tibber 动态电价
|
||||
|
||||
`app/integrations/tibber/client.py` httpx POST GraphQL(`priceInfoRange(QUARTER_HOURLY, first=96)`),解析 `startsAt`/`total`/`energy`/`tax`/`level`,按 `starts_at` upsert `tibber_price`(幂等)。启动 + 每小时抓取今明两天 15 分钟价、幂等 upsert(hourly trigger,确保每日刷新且可补重试);仅当 active 合同 kind=tibber 且 `tibber_api_token` 存在时运行。`POST /api/energy/tibber/test` 试连三态(success 带当前价 / config-error / failed)。
|
||||
|
||||
### 反哺 Home Assistant Energy
|
||||
|
||||
`_energy_cost_provider` 向 expose 框架注册 4 个实体:`buy_price_now`、`sell_price_now`(€/kWh sensor)、`import_cost_total`、`export_revenue_total`(`total_increasing` monetary,可直接挂 HA Energy 仪表盘)。默认未勾选,在 Expose 面板启用。
|
||||
|
||||
### API 端点(M6 新增)
|
||||
|
||||
| 端点 | 用途 |
|
||||
| --- | --- |
|
||||
| `GET /api/energy/contracts` | 列出合同 + active 标记 |
|
||||
| `POST /api/energy/contracts` | 新建合同(kind + 首版本值,按 profile 校验)|
|
||||
| `GET /api/energy/contracts/{id}` | 单个合同 + 版本历史 |
|
||||
| `PATCH /api/energy/contracts/{id}` | 改名 / 激活 |
|
||||
| `POST /api/energy/contracts/{id}/versions` | 加新版本(改价,带生效日期)|
|
||||
| `GET /api/energy/profiles` | 列出 pricing profile 结构(前端按它渲染表单)|
|
||||
| `GET /api/energy/prices` | 区间价格点(曲线)|
|
||||
| `GET /api/energy/costs` | 区间 `energy_cost_period`(走势/明细)|
|
||||
| `GET /api/energy/costs/summary` | 区间汇总(计量电费 + 固定费 − 抵扣)|
|
||||
| `POST /api/energy/costs/recompute` | 幂等重算 |
|
||||
| `GET /api/energy/dsmr/latest` | 最新 `dsmr_reading` |
|
||||
| `POST /api/energy/tibber/test` | 试连 Tibber + 拉当前价,三态 |
|
||||
|
||||
DSMR/Tibber 标量配置复用现有 `GET/PUT /api/config`(新增 `dsmr_ingest_enabled`、`dsmr_mqtt_topic`、`dsmr_sample_interval_s`、`tibber_api_token`(secret)、`tibber_home_id`)。
|
||||
|
||||
### 前端视图(M6 新增,并入 Energy 视图)
|
||||
|
||||
- **Contracts Tab**:合同列表 + 新建/编辑(表单按 `/api/energy/profiles` 结构渲染,不 hardcode 字段)+ 激活 + 改价加版本 + 版本历史只读。
|
||||
- **Prices Tab**:15 分钟价格曲线(tibber 动态或 manual 档位),复用 Recharts。
|
||||
- **Costs Tab**:费用走势/明细 + 汇总卡片(含固定费/抵扣)。
|
||||
- **Config 页 Tibber 测试**:三态(success/config-error/failed)。
|
||||
|
||||
## Config 持久化
|
||||
|
||||
@@ -198,6 +461,11 @@ python -m scripts.migrate_legacy_data
|
||||
- SMTP 基础配置
|
||||
- TickTick OAuth 配置
|
||||
- Home Assistant 配置
|
||||
- MQTT broker 配置(`MQTT_ENABLED`、`MQTT_BROKER_HOST/PORT/USERNAME/PASSWORD`、`MQTT_TLS_ENABLED`)
|
||||
- Home Assistant Discovery 配置(`HA_DISCOVERY_ENABLED`、`HA_DISCOVERY_PREFIX`)
|
||||
- Modbus 采集配置(`MODBUS_POLLING_ENABLED`)
|
||||
- DSMR 接入配置(`DSMR_INGEST_ENABLED`、`DSMR_MQTT_TOPIC`、`DSMR_SAMPLE_INTERVAL_S`)
|
||||
- Tibber 凭据(`TIBBER_API_TOKEN`(secret)、`TIBBER_HOME_ID`)
|
||||
|
||||
其中 SMTP password 与其他 secret 字段一致:
|
||||
|
||||
@@ -230,8 +498,8 @@ python -m scripts.migrate_legacy_data
|
||||
|
||||
当前系统已经提供最小可用的 SMTP 能力:
|
||||
|
||||
- SMTP 配置可在 `/config` 页面填写并保存到 `app_config`
|
||||
- 可通过 config 页面发送测试邮件
|
||||
- SMTP 配置可在 React SPA `/config` 页面填写并保存到 `app_config`(通过 `PUT /api/config`)
|
||||
- 可通过 config 页面发送测试邮件(`POST /api/config/smtp/test`)
|
||||
- 邮件 `From` 头支持显示名,例如 `Home Automation <sender@example.com>`
|
||||
|
||||
当前 SMTP 配置项包括:
|
||||
@@ -283,18 +551,34 @@ python scripts/export_openapi.py
|
||||
|
||||
当前 Compose 分成两层:
|
||||
|
||||
- `docker-compose.yml`:默认使用 registry image,适合部署 / 生产拉取
|
||||
- `docker-compose.override.yml`:仅为本地开发追加 `build: .`
|
||||
- `docker-compose.yml`:默认使用 registry image,适合部署 / 生产拉取(暴露 8881)
|
||||
- `docker-compose.dev.yml`:本地开发显式叠加层——追加 `build: .`、独立 project /
|
||||
容器名(`-dev` 后缀)、暴露 8001,并把 DB 指向挂载的 `./data` 副本,可与生产栈在同一台机器上并存
|
||||
|
||||
本地开发启动方式:
|
||||
WarmteLink serial access is configured directly by both Compose combinations. Before starting either
|
||||
one, set these host-specific values in your uncommitted local `.env` (use a stable `/dev/serial/by-id/...`
|
||||
path, never a transient `/dev/ttyUSB*` name):
|
||||
|
||||
```bash
|
||||
docker compose up -d --build
|
||||
```dotenv
|
||||
WARMTELINK_DEVICE_PATH=/dev/serial/by-id/<stable-by-id-name>
|
||||
WARMTELINK_SERIAL_GID=<host-serial-gid>
|
||||
```
|
||||
|
||||
上面的命令会自动叠加 `docker-compose.override.yml`,因此本地仍然会按当前工作目录重新 build。
|
||||
Only `app` receives the device as `/dev/warmtelink:rw` and the serial group; `migration` does not.
|
||||
The app remains non-root, non-privileged, and has no added capabilities. One physical serial port may
|
||||
have only one owner: stop the app before running the Pre-M8 P1 probe. Never remove `./data`, databases,
|
||||
or volumes while changing this configuration.
|
||||
|
||||
如果要按生产方式直接从 registry 拉取并启动,显式只使用基础 compose 文件:
|
||||
本地开发启动方式(显式叠加 dev 层):
|
||||
|
||||
```bash
|
||||
docker compose -f docker-compose.yml -f docker-compose.dev.yml up -d --build
|
||||
```
|
||||
|
||||
dev 层刻意不沿用 `docker-compose.override.yml` 这种会被 `docker compose up` 自动叠加的文件名,
|
||||
因此默认的 `docker compose up` 只用生产基础文件,不会把开发端口 / 配置误带到生产。
|
||||
|
||||
如果要按生产方式直接从 registry 拉取并启动,使用基础 compose 文件:
|
||||
|
||||
```bash
|
||||
docker compose -f docker-compose.yml pull
|
||||
|
||||
+12
-1
@@ -6,10 +6,21 @@ from sqlalchemy import engine_from_config, pool
|
||||
from app.config import get_settings
|
||||
from app.db import Base
|
||||
from app.models.config import AppConfigEntry # noqa: F401
|
||||
from app.models.auth import AuthSession, AuthUser # noqa: F401
|
||||
from app.models.auth import AuthSession, AuthUser, RecoveryCode # noqa: F401
|
||||
from app.models.auth_throttle import LoginThrottle # noqa: F401
|
||||
from app.models.public_ip import PublicIPHistory, PublicIPState # noqa: F401
|
||||
from app.models.location import Location # noqa: F401
|
||||
from app.models.poo import PooRecord # noqa: F401
|
||||
from app.models.modbus import ModbusDevice, ModbusReading # noqa: F401
|
||||
from app.models.expose import ExposedEntityToggle # noqa: F401
|
||||
from app.models.energy import ( # noqa: F401
|
||||
DsmrReading,
|
||||
EnergyContract,
|
||||
EnergyContractVersion,
|
||||
TibberPrice,
|
||||
EnergyCostPeriod,
|
||||
)
|
||||
from app.models.meter_source import MeterSource, MeterSourceBinding, MeterSourceChannel # noqa: F401
|
||||
|
||||
config = context.config
|
||||
|
||||
|
||||
@@ -0,0 +1,43 @@
|
||||
"""add auth_login_throttle table for exponential back-off throttling
|
||||
|
||||
Revision ID: 20260621_07_auth_login_throttle
|
||||
Revises: 20260611_06_merge_location_poo_tables
|
||||
Create Date: 2026-06-21 00:00:00.000000
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
revision: str = "20260621_07_auth_login_throttle"
|
||||
down_revision: Union[str, None] = "20260611_06_merge_location_poo_tables"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"auth_login_throttle",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("key", sa.String(length=255), nullable=False),
|
||||
sa.Column("scope", sa.String(length=16), nullable=False),
|
||||
sa.Column("failures", sa.Integer(), nullable=False),
|
||||
sa.Column("first_failed_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("last_failed_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("next_allowed_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("scope", "key", name="uq_auth_login_throttle_scope_key"),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_auth_login_throttle_scope_key",
|
||||
"auth_login_throttle",
|
||||
["scope", "key"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_auth_login_throttle_scope_key", table_name="auth_login_throttle")
|
||||
op.drop_table("auth_login_throttle")
|
||||
@@ -0,0 +1,61 @@
|
||||
"""add TOTP fields to auth_users and create auth_recovery_code table
|
||||
|
||||
Revision ID: 20260621_08_totp
|
||||
Revises: 20260621_07_auth_login_throttle
|
||||
Create Date: 2026-06-21 00:00:00.000000
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "20260621_08_totp"
|
||||
down_revision: Union[str, None] = "20260621_07_auth_login_throttle"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Add totp_secret (nullable) to auth_users — existing rows get NULL, which is correct.
|
||||
op.add_column("auth_users", sa.Column("totp_secret", sa.String(length=64), nullable=True))
|
||||
# Add totp_enabled (NOT NULL) with server_default="0" so existing rows default to false.
|
||||
op.add_column(
|
||||
"auth_users",
|
||||
sa.Column(
|
||||
"totp_enabled",
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default="0",
|
||||
),
|
||||
)
|
||||
|
||||
# Create auth_recovery_code table for one-time TOTP recovery codes.
|
||||
op.create_table(
|
||||
"auth_recovery_code",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("user_id", sa.Integer(), nullable=False),
|
||||
sa.Column("code_hash", sa.String(length=255), nullable=False),
|
||||
sa.Column("used_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.ForeignKeyConstraint(["user_id"], ["auth_users.id"]),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_auth_recovery_code_user_id"),
|
||||
"auth_recovery_code",
|
||||
["user_id"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Drop auth_recovery_code table first (has FK to auth_users).
|
||||
op.drop_index(op.f("ix_auth_recovery_code_user_id"), table_name="auth_recovery_code")
|
||||
op.drop_table("auth_recovery_code")
|
||||
|
||||
# Drop the two TOTP columns from auth_users.
|
||||
# Use batch_alter_table for SQLite compatibility (alembic's portable column-drop path).
|
||||
with op.batch_alter_table("auth_users") as batch_op:
|
||||
batch_op.drop_column("totp_enabled")
|
||||
batch_op.drop_column("totp_secret")
|
||||
@@ -0,0 +1,81 @@
|
||||
"""add modbus_device and modbus_reading tables
|
||||
|
||||
Revision ID: 20260622_09_modbus_tables
|
||||
Revises: 20260621_08_totp
|
||||
Create Date: 2026-06-22 00:00:00.000000
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "20260622_09_modbus_tables"
|
||||
down_revision: Union[str, None] = "20260621_08_totp"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# modbus_device — deployment/configurable metadata for each polled device.
|
||||
op.create_table(
|
||||
"modbus_device",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("uuid", sa.String(length=36), nullable=False),
|
||||
sa.Column("friendly_name", sa.String(length=255), nullable=False),
|
||||
sa.Column("transport", sa.String(length=16), nullable=False),
|
||||
sa.Column("host", sa.String(length=255), nullable=False),
|
||||
sa.Column("port", sa.Integer(), nullable=False),
|
||||
sa.Column("unit_id", sa.Integer(), nullable=False),
|
||||
sa.Column("profile", sa.String(length=64), nullable=False),
|
||||
sa.Column("poll_interval_s", sa.Integer(), nullable=False),
|
||||
sa.Column("enabled", sa.Boolean(), nullable=False),
|
||||
sa.Column("last_poll_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("last_poll_ok", sa.Boolean(), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("uuid", name="uq_modbus_device_uuid"),
|
||||
)
|
||||
|
||||
# modbus_reading — generic telemetry, one row per device per poll cycle.
|
||||
op.create_table(
|
||||
"modbus_reading",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("device_id", sa.Integer(), nullable=False),
|
||||
sa.Column("recorded_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("payload", sa.JSON(), nullable=False),
|
||||
sa.ForeignKeyConstraint(
|
||||
["device_id"],
|
||||
["modbus_device.id"],
|
||||
ondelete="RESTRICT",
|
||||
),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
|
||||
# Individual index on recorded_at (from the ORM-level index=True).
|
||||
op.create_index(
|
||||
"ix_modbus_reading_recorded_at",
|
||||
"modbus_reading",
|
||||
["recorded_at"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
# Composite index for efficient time-range queries per device.
|
||||
op.create_index(
|
||||
"ix_modbus_reading_device_recorded",
|
||||
"modbus_reading",
|
||||
["device_id", "recorded_at"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Drop the reading table first (it has a FK referencing modbus_device).
|
||||
op.drop_index("ix_modbus_reading_device_recorded", table_name="modbus_reading")
|
||||
op.drop_index("ix_modbus_reading_recorded_at", table_name="modbus_reading")
|
||||
op.drop_table("modbus_reading")
|
||||
|
||||
# Drop the device table.
|
||||
op.drop_table("modbus_device")
|
||||
@@ -0,0 +1,44 @@
|
||||
"""add exposed_entity_toggle table
|
||||
|
||||
Revision ID: 20260622_10_exposed_entities
|
||||
Revises: 20260622_09_modbus_tables
|
||||
Create Date: 2026-06-22 00:00:00.000000
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "20260622_10_exposed_entities"
|
||||
down_revision: Union[str, None] = "20260622_09_modbus_tables"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# exposed_entity_toggle — per-entity on/off switch for MQTT / HA Discovery.
|
||||
op.create_table(
|
||||
"exposed_entity_toggle",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("key", sa.String(length=255), nullable=False),
|
||||
sa.Column("enabled", sa.Boolean(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("key", name="uq_exposed_entity_toggle_key"),
|
||||
)
|
||||
|
||||
# Index on key for fast single-key lookups (toggle by key).
|
||||
op.create_index(
|
||||
"ix_exposed_entity_toggle_key",
|
||||
"exposed_entity_toggle",
|
||||
["key"],
|
||||
unique=True,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Drop index then table — only removes what this revision created.
|
||||
op.drop_index("ix_exposed_entity_toggle_key", table_name="exposed_entity_toggle")
|
||||
op.drop_table("exposed_entity_toggle")
|
||||
@@ -0,0 +1,132 @@
|
||||
"""add energy pricing and DSMR metering tables
|
||||
|
||||
Revision ID: 20260623_11_energy_tables
|
||||
Revises: 20260622_10_exposed_entities
|
||||
Create Date: 2026-06-23 00:00:00.000000
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "20260623_11_energy_tables"
|
||||
down_revision: Union[str, None] = "20260622_10_exposed_entities"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# dsmr_reading — raw DSMR telegram blobs (10-second down-sampled).
|
||||
# No device table: single P1 smart meter; a ``source`` column can be added later
|
||||
# if a second meter is introduced.
|
||||
op.create_table(
|
||||
"dsmr_reading",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("recorded_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("source_id", sa.Integer(), nullable=True),
|
||||
sa.Column("payload", sa.JSON(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("source_id", name="uq_dsmr_reading_source_id"),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_dsmr_reading_recorded_at",
|
||||
"dsmr_reading",
|
||||
["recorded_at"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
# energy_contract — contract head (manual or tibber, one active at a time).
|
||||
op.create_table(
|
||||
"energy_contract",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("name", sa.String(length=255), nullable=False),
|
||||
sa.Column("kind", sa.String(length=32), nullable=False),
|
||||
sa.Column("active", sa.Boolean(), nullable=False),
|
||||
sa.Column("currency", sa.String(length=8), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
|
||||
# energy_contract_version — versioned pricing values; append-only for auditability.
|
||||
# Must be created after energy_contract because of the FK dependency.
|
||||
op.create_table(
|
||||
"energy_contract_version",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("contract_id", sa.Integer(), nullable=False),
|
||||
sa.Column("effective_from", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("effective_to", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("values", sa.JSON(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(
|
||||
["contract_id"],
|
||||
["energy_contract.id"],
|
||||
ondelete="RESTRICT",
|
||||
),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
|
||||
# tibber_price — cached Tibber 15-minute spot prices (immutable once fetched).
|
||||
op.create_table(
|
||||
"tibber_price",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("starts_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("resolution", sa.String(length=32), nullable=False),
|
||||
sa.Column("energy", sa.Float(), nullable=False),
|
||||
sa.Column("tax", sa.Float(), nullable=False),
|
||||
sa.Column("total", sa.Float(), nullable=False),
|
||||
sa.Column("level", sa.String(length=32), nullable=True),
|
||||
sa.Column("currency", sa.String(length=8), nullable=False),
|
||||
sa.Column("fetched_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("starts_at", name="uq_tibber_price_starts_at"),
|
||||
)
|
||||
|
||||
# energy_cost_period — computed 15-minute billing periods (immutable snapshot).
|
||||
# Must be created after energy_contract_version because of the FK dependency.
|
||||
op.create_table(
|
||||
"energy_cost_period",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("period_start", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("d1_kwh", sa.Float(), nullable=False),
|
||||
sa.Column("d2_kwh", sa.Float(), nullable=False),
|
||||
sa.Column("r1_kwh", sa.Float(), nullable=False),
|
||||
sa.Column("r2_kwh", sa.Float(), nullable=False),
|
||||
sa.Column("import_cost", sa.Float(), nullable=False),
|
||||
sa.Column("export_revenue", sa.Float(), nullable=False),
|
||||
sa.Column("net_cost", sa.Float(), nullable=False),
|
||||
sa.Column("currency", sa.String(length=8), nullable=False),
|
||||
sa.Column("pricing", sa.JSON(), nullable=False),
|
||||
sa.Column("contract_version_id", sa.Integer(), nullable=True),
|
||||
sa.Column("degraded", sa.Boolean(), nullable=False),
|
||||
sa.Column("computed_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(
|
||||
["contract_version_id"],
|
||||
["energy_contract_version.id"],
|
||||
ondelete="RESTRICT",
|
||||
),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("period_start", name="uq_energy_cost_period_period_start"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Drop tables in reverse dependency order: tables with FKs first.
|
||||
|
||||
# energy_cost_period has a FK to energy_contract_version — drop it first.
|
||||
op.drop_table("energy_cost_period")
|
||||
|
||||
# tibber_price has no FK dependencies on our new tables.
|
||||
op.drop_table("tibber_price")
|
||||
|
||||
# energy_contract_version has a FK to energy_contract — drop it before energy_contract.
|
||||
op.drop_table("energy_contract_version")
|
||||
|
||||
# energy_contract has no FK dependencies on our new tables.
|
||||
op.drop_table("energy_contract")
|
||||
|
||||
# dsmr_reading has no FK dependencies.
|
||||
op.drop_index("ix_dsmr_reading_recorded_at", table_name="dsmr_reading")
|
||||
op.drop_table("dsmr_reading")
|
||||
@@ -0,0 +1,58 @@
|
||||
"""decouple dsmr_reading from the telegram id
|
||||
|
||||
The DSMR Reader's own ``id`` field is auto-incrementing but is known to overflow
|
||||
and require a manual reset to zero (a long-standing DSMR firmware quirk). If we
|
||||
keep a UNIQUE constraint on ``source_id`` (the telegram id) and use it for
|
||||
idempotency, an overflow/reset would make legitimately-new telegrams collide with
|
||||
old ids and be silently dropped as "duplicates" — data loss.
|
||||
|
||||
This migration removes that coupling:
|
||||
|
||||
- Drops the UNIQUE constraint on ``source_id`` (the column is kept as a plain,
|
||||
nullable reference value; it is no longer relied upon for uniqueness/dedup).
|
||||
- Makes ``recorded_at`` (the telegram timestamp) UNIQUE instead — a single P1
|
||||
meter emits one telegram per timestamp, so this is a robust, telegram-id-
|
||||
independent idempotency key. The table's own autoincrement ``id`` PK remains
|
||||
the stable internal identity.
|
||||
|
||||
Revision ID: 20260624_12_dsmr_decouple_telegram_id
|
||||
Revises: 20260623_11_energy_tables
|
||||
Create Date: 2026-06-24 00:00:00.000000
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "20260624_12_dsmr_decouple_telegram_id"
|
||||
down_revision: Union[str, None] = "20260623_11_energy_tables"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# SQLite cannot ALTER away a constraint in place, so recreate the table via
|
||||
# Alembic's batch mode. The table is recreated and rows are copied; existing
|
||||
# data (if any) is preserved. recorded_at must be unique for this to succeed
|
||||
# — a single meter never emits two telegrams at the exact same timestamp.
|
||||
with op.batch_alter_table("dsmr_reading", schema=None) as batch_op:
|
||||
# Old non-unique index on recorded_at is replaced by the unique constraint.
|
||||
batch_op.drop_index("ix_dsmr_reading_recorded_at")
|
||||
# Telegram id is no longer a uniqueness/idempotency key.
|
||||
batch_op.drop_constraint("uq_dsmr_reading_source_id", type_="unique")
|
||||
# Timestamp becomes the telegram-id-independent dedup key.
|
||||
batch_op.create_unique_constraint(
|
||||
"uq_dsmr_reading_recorded_at", ["recorded_at"]
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
with op.batch_alter_table("dsmr_reading", schema=None) as batch_op:
|
||||
batch_op.drop_constraint("uq_dsmr_reading_recorded_at", type_="unique")
|
||||
batch_op.create_unique_constraint(
|
||||
"uq_dsmr_reading_source_id", ["source_id"]
|
||||
)
|
||||
batch_op.create_index(
|
||||
"ix_dsmr_reading_recorded_at", ["recorded_at"], unique=False
|
||||
)
|
||||
@@ -0,0 +1,238 @@
|
||||
"""add meter table and energy_cost_period.meter_id
|
||||
|
||||
Introduces the ``meter`` table (one row per physical meter installation epoch)
|
||||
and a nullable FK column ``energy_cost_period.meter_id`` that attributes each
|
||||
billing period to a specific physical meter.
|
||||
|
||||
**Backfill logic (§3.7 of the M7 design doc)**:
|
||||
|
||||
If the database already contains any ``dsmr_reading`` or
|
||||
``energy_cost_period`` rows, one initial ``meter`` row is created:
|
||||
|
||||
label = "Initial meter"
|
||||
commodity = "electricity"
|
||||
started_at = earliest dsmr_reading.recorded_at
|
||||
(or, if none, earliest energy_cost_period.period_start,
|
||||
or, if still none, the migration timestamp)
|
||||
ended_at = NULL (still active)
|
||||
reason = "initial"
|
||||
|
||||
All existing ``energy_cost_period`` rows are then back-filled with that
|
||||
initial meter's id.
|
||||
|
||||
**Idempotency**: the backfill is guarded with a check for any existing
|
||||
``meter`` row whose ``reason = 'initial'`` and ``commodity = 'electricity'``
|
||||
and ``ended_at IS NULL``, so repeating the upgrade does not create duplicate
|
||||
meters or overwrite already-filled meter_id values.
|
||||
|
||||
**Audit (on-non-degraded periods only)**: after backfilling, the number of
|
||||
non-degraded ``energy_cost_period`` rows with ``meter_id IS NULL`` must be
|
||||
zero; if it is not, the migration raises a ``RuntimeError`` and rolls back.
|
||||
|
||||
**Data safety**: this migration is additive only — no existing rows are
|
||||
deleted or overwritten; it only creates a new table, adds a nullable column,
|
||||
and back-fills that column.
|
||||
|
||||
Revision ID: 20260625_13_meter_table
|
||||
Revises: 20260624_12_dsmr_decouple_telegram_id
|
||||
Create Date: 2026-06-25 00:00:00.000000
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "20260625_13_meter_table"
|
||||
down_revision: Union[str, None] = "20260624_12_dsmr_decouple_telegram_id"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ------------------------------------------------------------------ #
|
||||
# 1. Create the meter table. #
|
||||
# ------------------------------------------------------------------ #
|
||||
op.create_table(
|
||||
"meter",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("label", sa.String(length=255), nullable=False),
|
||||
sa.Column("commodity", sa.String(length=32), nullable=False),
|
||||
sa.Column("started_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("ended_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("reason", sa.String(length=64), nullable=False),
|
||||
sa.Column("note", sa.String(length=1024), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# 2. Add meter_id column to energy_cost_period (nullable FK). #
|
||||
# ------------------------------------------------------------------ #
|
||||
with op.batch_alter_table("energy_cost_period", schema=None) as batch_op:
|
||||
batch_op.add_column(
|
||||
sa.Column("meter_id", sa.Integer(), nullable=True)
|
||||
)
|
||||
batch_op.create_foreign_key(
|
||||
"fk_energy_cost_period_meter_id",
|
||||
"meter",
|
||||
["meter_id"],
|
||||
["id"],
|
||||
ondelete="RESTRICT",
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# 3. Backfill initial meter (idempotent). #
|
||||
# ------------------------------------------------------------------ #
|
||||
conn = op.get_bind()
|
||||
|
||||
# Check if there is already an initial meter (idempotency guard).
|
||||
existing_initial = conn.execute(
|
||||
sa.text(
|
||||
"SELECT id FROM meter "
|
||||
"WHERE reason = 'initial' AND commodity = 'electricity' AND ended_at IS NULL "
|
||||
"LIMIT 1"
|
||||
)
|
||||
).fetchone()
|
||||
|
||||
if existing_initial is not None:
|
||||
# Already backfilled — nothing to do.
|
||||
return
|
||||
|
||||
# Determine whether there is any historical data to create a meter for.
|
||||
has_readings = conn.execute(
|
||||
sa.text("SELECT 1 FROM dsmr_reading LIMIT 1")
|
||||
).fetchone()
|
||||
has_periods = conn.execute(
|
||||
sa.text("SELECT 1 FROM energy_cost_period LIMIT 1")
|
||||
).fetchone()
|
||||
|
||||
if not has_readings and not has_periods:
|
||||
# Empty database: no historical data, so no initial meter is needed.
|
||||
# meter_id will remain NULL on any future rows until T02 service layer
|
||||
# starts populating it.
|
||||
return
|
||||
|
||||
# Determine started_at: earliest dsmr_reading.recorded_at, falling back to
|
||||
# earliest energy_cost_period.period_start, and finally to now().
|
||||
earliest_reading_row = conn.execute(
|
||||
sa.text("SELECT MIN(recorded_at) AS ts FROM dsmr_reading")
|
||||
).fetchone()
|
||||
earliest_period_row = conn.execute(
|
||||
sa.text("SELECT MIN(period_start) AS ts FROM energy_cost_period")
|
||||
).fetchone()
|
||||
|
||||
started_at_value: datetime | None = None
|
||||
if earliest_reading_row and earliest_reading_row[0] is not None:
|
||||
# SQLite returns ISO strings for datetime columns; parse to datetime.
|
||||
raw = earliest_reading_row[0]
|
||||
started_at_value = _parse_sqlite_datetime(raw)
|
||||
if started_at_value is None and earliest_period_row and earliest_period_row[0] is not None:
|
||||
raw = earliest_period_row[0]
|
||||
started_at_value = _parse_sqlite_datetime(raw)
|
||||
if started_at_value is None:
|
||||
started_at_value = datetime.now(tz=timezone.utc)
|
||||
|
||||
now_utc = datetime.now(tz=timezone.utc)
|
||||
|
||||
# Insert the initial meter row.
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"INSERT INTO meter (label, commodity, started_at, ended_at, reason, note, created_at) "
|
||||
"VALUES (:label, :commodity, :started_at, NULL, :reason, NULL, :created_at)"
|
||||
),
|
||||
{
|
||||
"label": "Initial meter",
|
||||
"commodity": "electricity",
|
||||
"started_at": _iso(started_at_value),
|
||||
"reason": "initial",
|
||||
"created_at": _iso(now_utc),
|
||||
},
|
||||
)
|
||||
|
||||
# Retrieve the newly created meter id.
|
||||
meter_row = conn.execute(
|
||||
sa.text(
|
||||
"SELECT id FROM meter "
|
||||
"WHERE reason = 'initial' AND commodity = 'electricity' AND ended_at IS NULL "
|
||||
"LIMIT 1"
|
||||
)
|
||||
).fetchone()
|
||||
assert meter_row is not None, "Initial meter row not found after insert"
|
||||
meter_id: int = meter_row[0]
|
||||
|
||||
# Back-fill all existing energy_cost_period rows that have meter_id IS NULL.
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE energy_cost_period SET meter_id = :mid WHERE meter_id IS NULL"
|
||||
),
|
||||
{"mid": meter_id},
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# 4. Audit: verify all non-degraded periods have a meter_id. #
|
||||
# ------------------------------------------------------------------ #
|
||||
unmatched_row = conn.execute(
|
||||
sa.text(
|
||||
"SELECT COUNT(*) FROM energy_cost_period "
|
||||
"WHERE meter_id IS NULL AND degraded = 0"
|
||||
)
|
||||
).fetchone()
|
||||
unmatched_count: int = unmatched_row[0] if unmatched_row else 0
|
||||
|
||||
if unmatched_count != 0:
|
||||
raise RuntimeError(
|
||||
f"Meter backfill audit failed: {unmatched_count} non-degraded "
|
||||
"energy_cost_period row(s) still have meter_id IS NULL after backfill. "
|
||||
"Migration aborted to protect data integrity."
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Remove the FK column from energy_cost_period first (references meter).
|
||||
with op.batch_alter_table("energy_cost_period", schema=None) as batch_op:
|
||||
batch_op.drop_constraint("fk_energy_cost_period_meter_id", type_="foreignkey")
|
||||
batch_op.drop_column("meter_id")
|
||||
|
||||
# Drop the meter table.
|
||||
op.drop_table("meter")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _parse_sqlite_datetime(value: str | datetime) -> datetime:
|
||||
"""Parse a SQLite datetime value into an aware UTC datetime.
|
||||
|
||||
SQLite stores datetimes as ISO 8601 strings. SQLAlchemy may return them
|
||||
as plain strings or as naive datetimes (no tzinfo) depending on the driver
|
||||
and column declaration. This helper normalises both forms to an aware UTC
|
||||
``datetime``.
|
||||
"""
|
||||
if isinstance(value, datetime):
|
||||
if value.tzinfo is None:
|
||||
return value.replace(tzinfo=timezone.utc)
|
||||
return value
|
||||
# String form — strip trailing Z or +00:00 variants, then attach UTC.
|
||||
s = str(value).strip()
|
||||
for suffix in ("+00:00", "Z", " UTC"):
|
||||
if s.endswith(suffix):
|
||||
s = s[: -len(suffix)]
|
||||
# SQLite uses space as the T separator in some formats.
|
||||
s = s.replace(" ", "T")
|
||||
try:
|
||||
dt = datetime.fromisoformat(s)
|
||||
except ValueError:
|
||||
# Fallback: strip subseconds if present to handle unusual formats.
|
||||
dt = datetime.strptime(s[:19], "%Y-%m-%dT%H:%M:%S")
|
||||
return dt.replace(tzinfo=timezone.utc)
|
||||
|
||||
|
||||
def _iso(dt: datetime) -> str:
|
||||
"""Serialise a datetime to an ISO 8601 string for SQLite storage."""
|
||||
if dt.tzinfo is not None:
|
||||
dt = dt.astimezone(timezone.utc).replace(tzinfo=None)
|
||||
return dt.strftime("%Y-%m-%dT%H:%M:%S")
|
||||
@@ -0,0 +1,104 @@
|
||||
"""add uuid column to meter table
|
||||
|
||||
Adds a stable ``uuid`` (UUID v4 string) column to the ``meter`` table so that
|
||||
each meter epoch has a durable identity anchor suitable for use as an HA
|
||||
Discovery ``unique_id``.
|
||||
|
||||
**Migration strategy (SQLite-safe)**:
|
||||
|
||||
SQLite does not support adding a NOT NULL + UNIQUE column to a non-empty table
|
||||
in a single ``ALTER TABLE ADD COLUMN`` statement (adding a NOT NULL column
|
||||
without a default value is rejected if the table already has rows). The
|
||||
safe approach used here is:
|
||||
|
||||
1. Add ``uuid`` as a **nullable** column (SQLite allows this).
|
||||
2. **Back-fill** every existing ``meter`` row with a distinct ``str(uuid4())``
|
||||
value. Each row gets its *own* random UUID — not a shared value — so the
|
||||
subsequent UNIQUE constraint is satisfied.
|
||||
3. Use ``batch_alter_table`` (which re-creates the table under the hood in
|
||||
SQLite) to alter the column to ``NOT NULL`` and add a UNIQUE constraint.
|
||||
|
||||
**Idempotency**: only rows where ``uuid IS NULL`` are back-filled; rows that
|
||||
already have a uuid (e.g. from a repeated upgrade after a partial failure) are
|
||||
left untouched.
|
||||
|
||||
**Audit**: after back-fill, the count of rows with ``uuid IS NULL`` must be
|
||||
exactly zero; if not, the migration raises ``RuntimeError`` and rolls back.
|
||||
|
||||
**Data safety**: this migration is additive only — no existing rows are deleted
|
||||
or overwritten; it only adds a new column and fills it in.
|
||||
|
||||
Revision ID: 20260625_14_meter_uuid
|
||||
Revises: 20260625_13_meter_table
|
||||
Create Date: 2026-06-25 00:00:00.000000
|
||||
"""
|
||||
|
||||
import uuid as _uuid
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "20260625_14_meter_uuid"
|
||||
down_revision: Union[str, None] = "20260625_13_meter_table"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# 1. Add uuid as a nullable column. #
|
||||
# ------------------------------------------------------------------ #
|
||||
with op.batch_alter_table("meter", schema=None) as batch_op:
|
||||
batch_op.add_column(
|
||||
sa.Column("uuid", sa.String(length=36), nullable=True)
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# 2. Back-fill: assign a distinct UUID to every row that has #
|
||||
# uuid IS NULL. Each row gets its own random value so that the #
|
||||
# subsequent UNIQUE constraint is satisfied. #
|
||||
# ------------------------------------------------------------------ #
|
||||
rows = conn.execute(sa.text("SELECT id FROM meter WHERE uuid IS NULL")).fetchall()
|
||||
for (meter_id,) in rows:
|
||||
new_uuid = str(_uuid.uuid4())
|
||||
conn.execute(
|
||||
sa.text("UPDATE meter SET uuid = :uuid WHERE id = :mid"),
|
||||
{"uuid": new_uuid, "mid": meter_id},
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# 3. Audit: verify no rows remain with uuid IS NULL. #
|
||||
# ------------------------------------------------------------------ #
|
||||
null_count_row = conn.execute(
|
||||
sa.text("SELECT COUNT(*) FROM meter WHERE uuid IS NULL")
|
||||
).fetchone()
|
||||
null_count: int = null_count_row[0] if null_count_row else 0
|
||||
|
||||
if null_count != 0:
|
||||
raise RuntimeError(
|
||||
f"meter.uuid back-fill audit failed: {null_count} meter row(s) still have "
|
||||
"uuid IS NULL after back-fill. Migration aborted to protect data integrity."
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# 4. Alter column to NOT NULL + UNIQUE (requires batch on SQLite). #
|
||||
# batch_alter_table re-creates the table, so the UNIQUE constraint #
|
||||
# and NOT NULL are applied atomically. #
|
||||
# ------------------------------------------------------------------ #
|
||||
with op.batch_alter_table("meter", schema=None) as batch_op:
|
||||
batch_op.alter_column(
|
||||
"uuid",
|
||||
existing_type=sa.String(length=36),
|
||||
nullable=False,
|
||||
)
|
||||
batch_op.create_unique_constraint("uq_meter_uuid", ["uuid"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Drop the UNIQUE constraint and the uuid column (batch on SQLite).
|
||||
with op.batch_alter_table("meter", schema=None) as batch_op:
|
||||
batch_op.drop_constraint("uq_meter_uuid", type_="unique")
|
||||
batch_op.drop_column("uuid")
|
||||
@@ -0,0 +1,107 @@
|
||||
"""add protocol-agnostic meter source, channel, and binding tables
|
||||
|
||||
Revision ID: 20260822_15_meter_sources
|
||||
Revises: 20260625_14_meter_uuid
|
||||
Create Date: 2026-08-22 00:00:00.000000
|
||||
|
||||
This revision is additive on upgrade. It deliberately does not backfill
|
||||
existing DSMR data; that adoption is a later, separately audited migration.
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "20260822_15_meter_sources"
|
||||
down_revision: Union[str, None] = "20260625_14_meter_uuid"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"meter_source",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("uuid", sa.String(length=36), nullable=False),
|
||||
sa.Column("name", sa.String(length=255), nullable=False),
|
||||
sa.Column("kind", sa.String(length=64), nullable=False),
|
||||
sa.Column("enabled", sa.Boolean(), nullable=False),
|
||||
sa.Column("config", sa.JSON(), nullable=False),
|
||||
sa.Column("status", sa.String(length=32), nullable=False),
|
||||
sa.Column("last_seen_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("last_error", sa.String(length=1024), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("uuid", name="uq_meter_source_uuid"),
|
||||
)
|
||||
op.create_index("ix_meter_source_kind_enabled", "meter_source", ["kind", "enabled"])
|
||||
|
||||
op.create_table(
|
||||
"meter_source_channel",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("uuid", sa.String(length=36), nullable=False),
|
||||
sa.Column("source_id", sa.Integer(), nullable=False),
|
||||
sa.Column("channel_key", sa.String(length=128), nullable=False),
|
||||
sa.Column("label", sa.String(length=255), nullable=False),
|
||||
sa.Column("suggested_commodity", sa.String(length=32), nullable=True),
|
||||
sa.Column("unit", sa.String(length=32), nullable=False),
|
||||
sa.Column("device_type", sa.String(length=64), nullable=True),
|
||||
sa.Column("fingerprint", sa.String(length=64), nullable=True),
|
||||
sa.Column("latest_value", sa.Numeric(precision=20, scale=6), nullable=True),
|
||||
sa.Column("latest_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("latest_quality", sa.String(length=32), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(["source_id"], ["meter_source.id"], ondelete="RESTRICT"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("uuid", name="uq_meter_source_channel_uuid"),
|
||||
sa.UniqueConstraint("source_id", "channel_key", name="uq_meter_source_channel_source_key"),
|
||||
)
|
||||
op.create_index("ix_meter_source_channel_source_id", "meter_source_channel", ["source_id"])
|
||||
|
||||
op.create_table(
|
||||
"meter_source_binding",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("uuid", sa.String(length=36), nullable=False),
|
||||
sa.Column("meter_id", sa.Integer(), nullable=False),
|
||||
sa.Column("channel_id", sa.Integer(), nullable=False),
|
||||
sa.Column("started_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("ended_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(["meter_id"], ["meter.id"], ondelete="RESTRICT"),
|
||||
sa.ForeignKeyConstraint(
|
||||
["channel_id"], ["meter_source_channel.id"], ondelete="RESTRICT"
|
||||
),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("uuid", name="uq_meter_source_binding_uuid"),
|
||||
)
|
||||
op.create_index("ix_meter_source_binding_meter_id", "meter_source_binding", ["meter_id"])
|
||||
op.create_index("ix_meter_source_binding_channel_id", "meter_source_binding", ["channel_id"])
|
||||
|
||||
with op.batch_alter_table("energy_cost_period", schema=None) as batch_op:
|
||||
batch_op.add_column(sa.Column("source_binding_id", sa.Integer(), nullable=True))
|
||||
batch_op.create_foreign_key(
|
||||
"fk_energy_cost_period_source_binding_id",
|
||||
"meter_source_binding",
|
||||
["source_binding_id"],
|
||||
["id"],
|
||||
ondelete="RESTRICT",
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
with op.batch_alter_table("energy_cost_period", schema=None) as batch_op:
|
||||
batch_op.drop_constraint("fk_energy_cost_period_source_binding_id", type_="foreignkey")
|
||||
batch_op.drop_column("source_binding_id")
|
||||
|
||||
op.drop_index("ix_meter_source_binding_channel_id", table_name="meter_source_binding")
|
||||
op.drop_index("ix_meter_source_binding_meter_id", table_name="meter_source_binding")
|
||||
op.drop_table("meter_source_binding")
|
||||
op.drop_index("ix_meter_source_channel_source_id", table_name="meter_source_channel")
|
||||
op.drop_table("meter_source_channel")
|
||||
op.drop_index("ix_meter_source_kind_enabled", table_name="meter_source")
|
||||
op.drop_table("meter_source")
|
||||
@@ -0,0 +1,269 @@
|
||||
"""adopt historical DSMR rows into the source and binding model
|
||||
|
||||
Revision ID: 20260822_16_dsmr_source_adoption
|
||||
Revises: 20260822_15_meter_sources
|
||||
Create Date: 2026-08-22 00:00:00.000000
|
||||
|
||||
The upgrade is deliberately data-preserving: it creates one migration-owned
|
||||
DSMR source/channel, moves the telegram identifier to ``telegram_id``, and
|
||||
audits every reading and cost row before committing. Old ``app_config`` rows,
|
||||
payload JSON, and cost snapshots are never deleted or rewritten.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "20260822_16_dsmr_source_adoption"
|
||||
down_revision: Union[str, None] = "20260822_15_meter_sources"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def _as_bool(value: str | None) -> bool:
|
||||
return value is not None and value.strip().lower() in {"1", "true", "yes", "on"}
|
||||
|
||||
|
||||
def _iso_now() -> str:
|
||||
return datetime.now(tz=timezone.utc).replace(tzinfo=None).isoformat(sep=" ")
|
||||
|
||||
|
||||
def _config(connection: sa.Connection) -> dict[str, object]:
|
||||
rows = connection.execute(sa.text("SELECT key, value FROM app_config")).all()
|
||||
values = {str(key): str(value) for key, value in rows}
|
||||
# An unconfigured historical DSMR installation needs a disabled identity,
|
||||
# not guessed connection details. Preserve every legacy value we model
|
||||
# when any legacy DSMR/MQTT configuration was explicitly present.
|
||||
legacy_keys = {
|
||||
"MQTT_BROKER_HOST", "MQTT_BROKER_PORT", "MQTT_USERNAME", "MQTT_PASSWORD",
|
||||
"MQTT_TLS_ENABLED", "DSMR_MQTT_TOPIC", "DSMR_TARIFF_TOPIC", "DSMR_SAMPLE_INTERVAL_S",
|
||||
}
|
||||
if not legacy_keys & values.keys():
|
||||
return {}
|
||||
return {
|
||||
"broker_host": values.get("MQTT_BROKER_HOST", ""),
|
||||
"broker_port": int(values.get("MQTT_BROKER_PORT", "1883")),
|
||||
"username": values.get("MQTT_USERNAME", ""),
|
||||
"password": values.get("MQTT_PASSWORD", ""),
|
||||
"tls_enabled": _as_bool(values.get("MQTT_TLS_ENABLED")),
|
||||
"topic": values.get("DSMR_MQTT_TOPIC", "dsmr/json"),
|
||||
"tariff_topic": values.get("DSMR_TARIFF_TOPIC", "dsmr/meter-stats/electricity_tariff"),
|
||||
"sample_interval_s": int(values.get("DSMR_SAMPLE_INTERVAL_S", "10")),
|
||||
}
|
||||
|
||||
|
||||
def _count(connection: sa.Connection, table: str) -> int:
|
||||
return int(connection.execute(sa.text(f"SELECT COUNT(*) FROM {table}")).scalar_one())
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
connection = op.get_bind()
|
||||
readings_before = _count(connection, "dsmr_reading")
|
||||
costs_before = _count(connection, "energy_cost_period")
|
||||
sources_before = _count(connection, "meter_source")
|
||||
channels_before = _count(connection, "meter_source_channel")
|
||||
bindings_before = _count(connection, "meter_source_binding")
|
||||
now = _iso_now()
|
||||
|
||||
# A source exists even without historical configuration/readings. It stays
|
||||
# disabled unless the old explicit DSMR switch was enabled, so no broker or
|
||||
# topic is guessed at runtime.
|
||||
config = _config(connection)
|
||||
source_result = connection.execute(
|
||||
sa.text(
|
||||
"INSERT INTO meter_source "
|
||||
"(uuid, name, kind, enabled, config, status, last_seen_at, last_error, created_at, updated_at) "
|
||||
"VALUES (:uuid, :name, 'dsmr_mqtt', :enabled, :config, 'unknown', NULL, NULL, :now, :now)"
|
||||
),
|
||||
{
|
||||
"uuid": str(uuid.uuid4()),
|
||||
"name": "Migrated DSMR source",
|
||||
"enabled": _as_bool(
|
||||
connection.execute(
|
||||
sa.text("SELECT value FROM app_config WHERE key = 'DSMR_INGEST_ENABLED'")
|
||||
).scalar_one_or_none()
|
||||
),
|
||||
"config": __import__("json").dumps(config),
|
||||
"now": now,
|
||||
},
|
||||
)
|
||||
source_id = source_result.lastrowid
|
||||
if source_id is None:
|
||||
raise RuntimeError("DSMR source adoption failed to create a source")
|
||||
channel_result = connection.execute(
|
||||
sa.text(
|
||||
"INSERT INTO meter_source_channel "
|
||||
"(uuid, source_id, channel_key, label, suggested_commodity, unit, device_type, fingerprint, "
|
||||
"latest_value, latest_at, latest_quality, created_at, updated_at) "
|
||||
"VALUES (:uuid, :source_id, 'electricity-total', 'DSMR electricity total', 'electricity', "
|
||||
"'kWh', NULL, NULL, NULL, NULL, NULL, :now, :now)"
|
||||
),
|
||||
{"uuid": str(uuid.uuid4()), "source_id": source_id, "now": now},
|
||||
)
|
||||
channel_id = channel_result.lastrowid
|
||||
if channel_id is None:
|
||||
raise RuntimeError("DSMR source adoption failed to create an electricity channel")
|
||||
if _count(connection, "meter_source") != sources_before + 1:
|
||||
raise RuntimeError("DSMR source adoption source row-count audit failed")
|
||||
if _count(connection, "meter_source_channel") != channels_before + 1:
|
||||
raise RuntimeError("DSMR source adoption channel row-count audit failed")
|
||||
|
||||
# Rename/add while nullable, back-fill all rows, then make the FK non-null
|
||||
# and replace the legacy timestamp-only uniqueness in a SQLite batch rebuild.
|
||||
with op.batch_alter_table("dsmr_reading", schema=None) as batch_op:
|
||||
batch_op.alter_column("source_id", new_column_name="telegram_id")
|
||||
batch_op.add_column(sa.Column("meter_source_id", sa.Integer(), nullable=True))
|
||||
connection.execute(
|
||||
sa.text("UPDATE dsmr_reading SET meter_source_id = :source_id WHERE meter_source_id IS NULL"),
|
||||
{"source_id": source_id},
|
||||
)
|
||||
with op.batch_alter_table("dsmr_reading", schema=None) as batch_op:
|
||||
batch_op.drop_constraint("uq_dsmr_reading_recorded_at", type_="unique")
|
||||
batch_op.alter_column("meter_source_id", existing_type=sa.Integer(), nullable=False)
|
||||
batch_op.create_foreign_key(
|
||||
"fk_dsmr_reading_meter_source_id", "meter_source", ["meter_source_id"], ["id"],
|
||||
ondelete="RESTRICT",
|
||||
)
|
||||
batch_op.create_unique_constraint(
|
||||
"uq_dsmr_reading_source_recorded_at", ["meter_source_id", "recorded_at"]
|
||||
)
|
||||
batch_op.create_index("ix_dsmr_reading_meter_source_id", ["meter_source_id"])
|
||||
adopted_readings = int(
|
||||
connection.execute(
|
||||
sa.text("SELECT COUNT(*) FROM dsmr_reading WHERE meter_source_id = :source_id"),
|
||||
{"source_id": source_id},
|
||||
).scalar_one()
|
||||
)
|
||||
if adopted_readings != readings_before:
|
||||
raise RuntimeError("DSMR source adoption reading source audit failed")
|
||||
|
||||
# Bind each electricity meter only where it overlaps the actual DSMR data.
|
||||
data_window = connection.execute(
|
||||
sa.text("SELECT MIN(recorded_at), MAX(recorded_at) FROM dsmr_reading")
|
||||
).one()
|
||||
expected_binding_count = 0
|
||||
if data_window[0] is not None:
|
||||
meters = connection.execute(
|
||||
sa.text(
|
||||
"SELECT id, started_at, ended_at FROM meter WHERE commodity = 'electricity' "
|
||||
"ORDER BY started_at, id"
|
||||
)
|
||||
).all()
|
||||
for meter_id, started_at, ended_at in meters:
|
||||
# Intersect [meter start, meter end) with the inclusive historical
|
||||
# samples. A closed boundary at the final sample remains valid for
|
||||
# the preceding interval; an empty intersection gets no fake binding.
|
||||
if started_at > data_window[1] or (ended_at is not None and ended_at <= data_window[0]):
|
||||
continue
|
||||
expected_binding_count += 1
|
||||
binding_start = max(started_at, data_window[0])
|
||||
binding_end = ended_at
|
||||
connection.execute(
|
||||
sa.text(
|
||||
"INSERT INTO meter_source_binding "
|
||||
"(uuid, meter_id, channel_id, started_at, ended_at, created_at, updated_at) "
|
||||
"VALUES (:uuid, :meter_id, :channel_id, :started_at, :ended_at, :now, :now)"
|
||||
),
|
||||
{
|
||||
"uuid": str(uuid.uuid4()), "meter_id": meter_id, "channel_id": channel_id,
|
||||
"started_at": binding_start, "ended_at": binding_end, "now": now,
|
||||
},
|
||||
)
|
||||
|
||||
# A cost period may be linked only if exactly one binding covers both its
|
||||
# start and end. Historical boundary/unknown rows remain auditable but are
|
||||
# explicitly degraded instead of being silently attributed to a current meter.
|
||||
periods = connection.execute(
|
||||
sa.text("SELECT id, meter_id, period_start, degraded FROM energy_cost_period")
|
||||
).all()
|
||||
resolvable_normal_periods: dict[int, int] = {}
|
||||
unresolved_period_ids: set[int] = set()
|
||||
for period_id, meter_id, period_start, degraded_before in periods:
|
||||
candidates = []
|
||||
if meter_id is not None:
|
||||
candidates = connection.execute(
|
||||
sa.text(
|
||||
"SELECT id FROM meter_source_binding "
|
||||
"WHERE meter_id = :meter_id AND started_at <= :start "
|
||||
"AND (ended_at IS NULL OR julianday(ended_at) > julianday(:start, '+15 minutes'))"
|
||||
),
|
||||
{"meter_id": meter_id, "start": period_start},
|
||||
).all()
|
||||
if len(candidates) == 1:
|
||||
if not degraded_before:
|
||||
resolvable_normal_periods[period_id] = candidates[0][0]
|
||||
connection.execute(
|
||||
sa.text("UPDATE energy_cost_period SET source_binding_id = :binding_id WHERE id = :id"),
|
||||
{"binding_id": candidates[0][0], "id": period_id},
|
||||
)
|
||||
else:
|
||||
unresolved_period_ids.add(period_id)
|
||||
connection.execute(
|
||||
sa.text(
|
||||
"UPDATE energy_cost_period SET degraded = 1, source_binding_id = NULL WHERE id = :id"
|
||||
),
|
||||
{"id": period_id},
|
||||
)
|
||||
|
||||
readings_after = _count(connection, "dsmr_reading")
|
||||
costs_after = _count(connection, "energy_cost_period")
|
||||
if readings_after != readings_before or costs_after != costs_before:
|
||||
raise RuntimeError("DSMR source adoption row-count audit failed")
|
||||
if _count(connection, "meter_source_binding") != bindings_before + expected_binding_count:
|
||||
raise RuntimeError("DSMR source adoption binding row-count audit failed")
|
||||
if int(
|
||||
connection.execute(
|
||||
sa.text("SELECT COUNT(*) FROM meter_source_binding WHERE channel_id = :channel_id"),
|
||||
{"channel_id": channel_id},
|
||||
).scalar_one()
|
||||
) != expected_binding_count:
|
||||
raise RuntimeError("DSMR source adoption binding channel audit failed")
|
||||
for period_id, binding_id in resolvable_normal_periods.items():
|
||||
bound, degraded = connection.execute(
|
||||
sa.text("SELECT source_binding_id, degraded FROM energy_cost_period WHERE id = :id"),
|
||||
{"id": period_id},
|
||||
).one()
|
||||
if bound != binding_id or degraded:
|
||||
raise RuntimeError("DSMR source adoption resolvable cost audit failed")
|
||||
if unresolved_period_ids:
|
||||
unresolved_count = int(
|
||||
connection.execute(
|
||||
sa.text(
|
||||
"SELECT COUNT(*) FROM energy_cost_period "
|
||||
"WHERE id IN :period_ids AND (degraded != 1 OR source_binding_id IS NOT NULL)"
|
||||
).bindparams(sa.bindparam("period_ids", expanding=True)),
|
||||
{"period_ids": list(unresolved_period_ids)},
|
||||
).scalar_one()
|
||||
)
|
||||
if unresolved_count:
|
||||
raise RuntimeError("DSMR source adoption unresolved cost audit failed")
|
||||
orphan_rows = connection.execute(sa.text("PRAGMA foreign_key_check")).all()
|
||||
if orphan_rows:
|
||||
raise RuntimeError("DSMR source adoption foreign-key audit failed")
|
||||
normal_unbound = int(
|
||||
connection.execute(
|
||||
sa.text("SELECT COUNT(*) FROM energy_cost_period WHERE degraded = 0 AND source_binding_id IS NULL")
|
||||
).scalar_one()
|
||||
)
|
||||
if normal_unbound:
|
||||
raise RuntimeError(f"DSMR source adoption left {normal_unbound} normal cost period(s) unbound")
|
||||
if _count(connection, "meter_source") < 1:
|
||||
raise RuntimeError("DSMR source adoption source audit failed")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Schema-only downgrade for isolated test databases. It intentionally does
|
||||
# not delete migration-created source/channel/binding rows.
|
||||
with op.batch_alter_table("dsmr_reading", schema=None) as batch_op:
|
||||
batch_op.drop_index("ix_dsmr_reading_meter_source_id")
|
||||
batch_op.drop_constraint("uq_dsmr_reading_source_recorded_at", type_="unique")
|
||||
batch_op.drop_constraint("fk_dsmr_reading_meter_source_id", type_="foreignkey")
|
||||
batch_op.drop_column("meter_source_id")
|
||||
batch_op.alter_column("telegram_id", new_column_name="source_id")
|
||||
batch_op.create_unique_constraint("uq_dsmr_reading_recorded_at", ["recorded_at"])
|
||||
@@ -0,0 +1,52 @@
|
||||
"""add normalized WarmteLink scalar reading history
|
||||
|
||||
Revision ID: 20260822_17_warmtelink_readings
|
||||
Revises: 20260822_16_dsmr_source_adoption
|
||||
Create Date: 2026-08-22 00:00:00.000000
|
||||
|
||||
The upgrade is additive: existing business rows are neither changed nor
|
||||
removed. The downgrade is schema-only and is exercised only on isolated test
|
||||
databases.
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "20260822_17_warmtelink_readings"
|
||||
down_revision: Union[str, None] = "20260822_16_dsmr_source_adoption"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"warmtelink_reading",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("channel_id", sa.Integer(), nullable=False),
|
||||
sa.Column("recorded_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("received_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("value", sa.Numeric(precision=15, scale=3), nullable=False),
|
||||
sa.Column("unit", sa.String(length=32), nullable=False),
|
||||
sa.Column("quality", sa.String(length=32), nullable=False),
|
||||
sa.Column("equipment_fingerprint", sa.String(length=64), nullable=False),
|
||||
sa.CheckConstraint(
|
||||
"quality IN ('valid', 'invalid', 'unverifiable')",
|
||||
name="ck_warmtelink_reading_quality",
|
||||
),
|
||||
sa.ForeignKeyConstraint(
|
||||
["channel_id"], ["meter_source_channel.id"], ondelete="RESTRICT"
|
||||
),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint(
|
||||
"channel_id", "recorded_at", name="uq_warmtelink_reading_channel_recorded_at"
|
||||
),
|
||||
)
|
||||
op.create_index("ix_warmtelink_reading_recorded_at", "warmtelink_reading", ["recorded_at"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_warmtelink_reading_recorded_at", table_name="warmtelink_reading")
|
||||
op.drop_table("warmtelink_reading")
|
||||
@@ -0,0 +1,107 @@
|
||||
"""add a billing scope to energy contracts
|
||||
|
||||
Revision ID: 20260822_18_contract_scopes
|
||||
Revises: 20260822_17_warmtelink_readings
|
||||
Create Date: 2026-08-22 00:00:00.000000
|
||||
|
||||
The upgrade preserves every existing contract, version and cost row. Existing
|
||||
contracts predate scopes and therefore deterministically belong to electricity.
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "20260822_18_contract_scopes"
|
||||
down_revision: Union[str, None] = "20260822_17_warmtelink_readings"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def _count(connection: sa.Connection, table: str) -> int:
|
||||
return int(connection.execute(sa.text(f"SELECT COUNT(*) FROM {table}")).scalar_one())
|
||||
|
||||
|
||||
def _orphan_count(connection: sa.Connection) -> int:
|
||||
version_orphans = connection.execute(
|
||||
sa.text(
|
||||
"SELECT COUNT(*) FROM energy_contract_version v "
|
||||
"LEFT JOIN energy_contract c ON c.id = v.contract_id WHERE c.id IS NULL"
|
||||
)
|
||||
).scalar_one()
|
||||
cost_orphans = connection.execute(
|
||||
sa.text(
|
||||
"SELECT COUNT(*) FROM energy_cost_period p "
|
||||
"LEFT JOIN energy_contract_version v ON v.id = p.contract_version_id "
|
||||
"WHERE p.contract_version_id IS NOT NULL AND v.id IS NULL"
|
||||
)
|
||||
).scalar_one()
|
||||
return int(version_orphans) + int(cost_orphans)
|
||||
|
||||
|
||||
def _audit_scope_upgrade(connection: sa.Connection, before: dict[str, int], orphan_before: int) -> None:
|
||||
after = {table: _count(connection, table) for table in before}
|
||||
if after != before:
|
||||
raise RuntimeError("contract scope migration row-count audit failed")
|
||||
if _orphan_count(connection) != orphan_before:
|
||||
raise RuntimeError("contract scope migration FK audit failed")
|
||||
invalid_scope_count = connection.execute(
|
||||
sa.text("SELECT COUNT(*) FROM energy_contract WHERE scope IS NULL OR scope != 'electricity'")
|
||||
).scalar_one()
|
||||
if invalid_scope_count:
|
||||
raise RuntimeError("contract scope migration backfill audit failed")
|
||||
|
||||
# Kept on Alembic's Config attributes rather than an environment switch so
|
||||
# isolated migration tests can deterministically exercise the rollback
|
||||
# boundary without changing production behavior.
|
||||
failure_injector = op.get_context().config.attributes.get("m8_t12_post_ddl_audit_failure")
|
||||
if callable(failure_injector):
|
||||
failure_injector()
|
||||
|
||||
|
||||
def _apply_scope_schema() -> None:
|
||||
# SQLite batch mode reconstructs the table. The server default gives every
|
||||
# historical row its deterministic value during reconstruction.
|
||||
with op.batch_alter_table("energy_contract", schema=None) as batch_op:
|
||||
batch_op.add_column(
|
||||
sa.Column("scope", sa.String(length=32), nullable=False, server_default="electricity")
|
||||
)
|
||||
batch_op.create_index("ix_energy_contract_scope", ["scope"])
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
connection = op.get_bind()
|
||||
before = {
|
||||
table: _count(connection, table)
|
||||
for table in ("energy_contract", "energy_contract_version", "energy_cost_period")
|
||||
}
|
||||
orphan_before = _orphan_count(connection)
|
||||
|
||||
if connection.dialect.name != "sqlite":
|
||||
_apply_scope_schema()
|
||||
_audit_scope_upgrade(connection, before, orphan_before)
|
||||
return
|
||||
|
||||
# Alembic marks SQLite batch DDL as non-transactional. SQLite itself can
|
||||
# nevertheless atomically roll back CREATE/COPY/DROP/RENAME when an
|
||||
# explicit transaction owns the complete batch operation. Keep the audit
|
||||
# inside that boundary so a failed audit cannot strand a revision-17 DB
|
||||
# with a revision-18 table shape.
|
||||
connection.exec_driver_sql("BEGIN IMMEDIATE")
|
||||
try:
|
||||
_apply_scope_schema()
|
||||
_audit_scope_upgrade(connection, before, orphan_before)
|
||||
except BaseException:
|
||||
connection.exec_driver_sql("ROLLBACK")
|
||||
raise
|
||||
else:
|
||||
connection.exec_driver_sql("COMMIT")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Schema reversibility is only exercised against isolated temporary test DBs.
|
||||
with op.batch_alter_table("energy_contract", schema=None) as batch_op:
|
||||
batch_op.drop_index("ix_energy_contract_scope")
|
||||
batch_op.drop_column("scope")
|
||||
@@ -0,0 +1,95 @@
|
||||
"""add generic commodity-scoped meter cost periods
|
||||
|
||||
Revision ID: 20260822_19_meter_cost_periods
|
||||
Revises: 20260822_18_contract_scopes
|
||||
Create Date: 2026-08-22 00:00:00.000000
|
||||
|
||||
This additive migration creates a separate audit ledger for non-electricity
|
||||
meter costs. It deliberately does not alter, migrate, or delete rows from the
|
||||
existing electricity-only energy_cost_period table.
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "20260822_19_meter_cost_periods"
|
||||
down_revision: Union[str, None] = "20260822_18_contract_scopes"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
class ExactDecimal(sa.TypeDecorator):
|
||||
"""Use SQLite text storage while retaining Numeric semantics elsewhere."""
|
||||
|
||||
impl = sa.Numeric
|
||||
cache_ok = True
|
||||
|
||||
def __init__(self, precision: int, scale: int) -> None:
|
||||
self.precision = precision
|
||||
self.scale = scale
|
||||
super().__init__(precision=precision, scale=scale)
|
||||
|
||||
def load_dialect_impl(self, dialect):
|
||||
if dialect.name == "sqlite":
|
||||
return dialect.type_descriptor(sa.String(self.precision + 2))
|
||||
return dialect.type_descriptor(sa.Numeric(self.precision, self.scale, asdecimal=True))
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"meter_cost_period",
|
||||
sa.Column("id", sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column("commodity", sa.String(length=32), nullable=False),
|
||||
sa.Column("period_start", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("period_end", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("meter_id", sa.Integer(), nullable=True),
|
||||
sa.Column("source_binding_id", sa.Integer(), nullable=True),
|
||||
sa.Column("contract_version_id", sa.Integer(), nullable=True),
|
||||
# SQLite NUMERIC coercion binds Decimal values as binary floats. Store
|
||||
# fixed-width decimal text there, while retaining Numeric semantics on
|
||||
# other supported dialects.
|
||||
sa.Column("quantity", ExactDecimal(15, 6), nullable=False),
|
||||
sa.Column("cost", ExactDecimal(15, 9), nullable=False),
|
||||
sa.Column("currency", sa.String(length=8), nullable=False),
|
||||
sa.Column("cost_breakdown", sa.JSON(), nullable=False),
|
||||
sa.Column("pricing_snapshot", sa.JSON(), nullable=False),
|
||||
sa.Column("quality", sa.String(length=32), nullable=False),
|
||||
sa.Column("degraded", sa.Boolean(), nullable=False, server_default=sa.false()),
|
||||
sa.Column("degraded_reason", sa.String(length=255), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(["meter_id"], ["meter.id"], ondelete="RESTRICT"),
|
||||
sa.ForeignKeyConstraint(
|
||||
["source_binding_id"], ["meter_source_binding.id"], ondelete="RESTRICT"
|
||||
),
|
||||
sa.ForeignKeyConstraint(
|
||||
["contract_version_id"], ["energy_contract_version.id"], ondelete="RESTRICT"
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"degraded OR (meter_id IS NOT NULL AND source_binding_id IS NOT NULL "
|
||||
"AND contract_version_id IS NOT NULL)",
|
||||
name="ck_meter_cost_period_normal_audit_links",
|
||||
),
|
||||
sa.CheckConstraint("period_end > period_start", name="ck_meter_cost_period_positive_interval"),
|
||||
sa.CheckConstraint(
|
||||
"NOT degraded OR (degraded_reason IS NOT NULL AND length(trim(degraded_reason)) > 0)",
|
||||
name="ck_meter_cost_period_degraded_reason",
|
||||
),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("commodity", "period_start", name="uq_meter_cost_period_commodity_start"),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_meter_cost_period_commodity_start", "meter_cost_period", ["commodity", "period_start"]
|
||||
)
|
||||
op.create_index(
|
||||
"ix_meter_cost_period_source_binding_id", "meter_cost_period", ["source_binding_id"]
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_meter_cost_period_source_binding_id", table_name="meter_cost_period")
|
||||
op.drop_index("ix_meter_cost_period_commodity_start", table_name="meter_cost_period")
|
||||
op.drop_table("meter_cost_period")
|
||||
@@ -0,0 +1,324 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from fastapi.responses import JSONResponse
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.routes.api.deps import require_csrf, require_session
|
||||
from app.config import Settings, get_settings
|
||||
from app.dependencies import get_app_settings, get_db
|
||||
from app.integrations.mqtt import MQTT_SETTINGS_KEYS, mqtt_manager, mqtt_test_client_id
|
||||
from app.schemas.config import (
|
||||
ConfigField,
|
||||
ConfigResponse,
|
||||
ConfigSection,
|
||||
ConfigUpdateRequest,
|
||||
ConfigUpdateResponse,
|
||||
MqttTestResponse,
|
||||
SmtpTestResponse,
|
||||
)
|
||||
from app.services.auth import AuthenticatedSession
|
||||
from app.services.config_page import ConfigSaveError, build_config_sections, save_config_updates
|
||||
from app.services.email import EmailConfigurationError, EmailDeliveryError, send_smtp_test_email
|
||||
from app.services.tibber_prices import active_tibber_contract_exists, trigger_tibber_refresh
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/api", tags=["api-config"])
|
||||
|
||||
|
||||
def _sections_from_raw(sections_raw: list[dict]) -> list[ConfigSection]:
|
||||
result = []
|
||||
for section in sections_raw:
|
||||
fields = [ConfigField(**f) for f in section["fields"]]
|
||||
result.append(ConfigSection(name=section["name"], fields=fields))
|
||||
return result
|
||||
|
||||
|
||||
@router.get("/config", response_model=ConfigResponse)
|
||||
def get_config(
|
||||
db: Session = Depends(get_db),
|
||||
settings: Settings = Depends(get_app_settings),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> ConfigResponse:
|
||||
"""Return all configuration sections. Secret field values are masked (empty string)."""
|
||||
sections_raw = build_config_sections(db, settings)
|
||||
return ConfigResponse(sections=_sections_from_raw(sections_raw))
|
||||
|
||||
|
||||
@router.put("/config", response_model=ConfigUpdateResponse)
|
||||
def put_config(
|
||||
body: ConfigUpdateRequest,
|
||||
db: Session = Depends(get_db),
|
||||
settings: Settings = Depends(get_app_settings),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> ConfigUpdateResponse:
|
||||
"""
|
||||
Save configuration updates.
|
||||
|
||||
- Blank secret value keeps the existing stored value (no change).
|
||||
- Invalid values return 422 and nothing is written to the database.
|
||||
- If MQTT-related settings changed, the MQTT client reconnects automatically.
|
||||
"""
|
||||
# Detect whether any MQTT-related key is being submitted (non-secret change
|
||||
# or non-blank secret change) so we know to reconnect after saving.
|
||||
mqtt_keys_submitted = any(k.lower() in MQTT_SETTINGS_KEYS for k in body.updates)
|
||||
tibber_values_before = (
|
||||
settings.tibber_api_token,
|
||||
settings.tibber_home_id,
|
||||
)
|
||||
|
||||
try:
|
||||
save_config_updates(db, body.updates, settings)
|
||||
except ConfigSaveError as exc:
|
||||
logger.warning("Rejected config update via API: %s", exc)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="invalid config submission",
|
||||
) from exc
|
||||
|
||||
# Re-read settings after save (save_config_updates clears the settings cache).
|
||||
# Use build_runtime_settings so the reconnect picks up DB-stored values, not
|
||||
# just the bootstrap env (otherwise a broker configured via the UI is ignored).
|
||||
from app.services.config_page import build_runtime_settings
|
||||
refreshed_settings = build_runtime_settings(db, get_settings())
|
||||
|
||||
# Reconnect MQTT client if any MQTT setting was updated.
|
||||
if mqtt_keys_submitted:
|
||||
logger.info("MQTT settings changed — triggering reconnect.")
|
||||
mqtt_manager.reconnect(refreshed_settings)
|
||||
|
||||
# Re-apply the DSMR subscription so enabling/disabling DSMR ingest (or changing
|
||||
# its topic / sample interval) takes effect immediately, without an app restart.
|
||||
# Done after any MQTT reconnect so it operates on the current client.
|
||||
from app.services.dsmr_ingest import apply_dsmr_subscription
|
||||
apply_dsmr_subscription(refreshed_settings)
|
||||
|
||||
tibber_values_changed = tibber_values_before != (
|
||||
refreshed_settings.tibber_api_token,
|
||||
refreshed_settings.tibber_home_id,
|
||||
)
|
||||
if tibber_values_changed and active_tibber_contract_exists(db):
|
||||
trigger_tibber_refresh()
|
||||
|
||||
sections_raw = build_config_sections(db, refreshed_settings)
|
||||
return ConfigUpdateResponse(sections=_sections_from_raw(sections_raw))
|
||||
|
||||
|
||||
@router.post(
|
||||
"/config/smtp/test",
|
||||
responses={
|
||||
200: {"model": SmtpTestResponse},
|
||||
400: {"model": SmtpTestResponse},
|
||||
502: {"model": SmtpTestResponse},
|
||||
},
|
||||
)
|
||||
def post_smtp_test(
|
||||
settings: Settings = Depends(get_app_settings),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> JSONResponse:
|
||||
"""
|
||||
Send a test SMTP email using the current runtime settings.
|
||||
|
||||
Returns a structured result indicating success or the category of failure.
|
||||
Three possible outcomes:
|
||||
- 200 { "result": "success", "message": ... }
|
||||
- 400 { "result": "config-error", "message": ... } (EmailConfigurationError)
|
||||
- 502 { "result": "failed", "message": ... } (EmailDeliveryError)
|
||||
|
||||
SMTP credentials are never echoed in the response.
|
||||
"""
|
||||
try:
|
||||
send_smtp_test_email(settings)
|
||||
except EmailConfigurationError as exc:
|
||||
logger.warning("SMTP test rejected due to configuration: %s", exc)
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
content={"result": "config-error", "message": str(exc)},
|
||||
)
|
||||
except EmailDeliveryError as exc:
|
||||
logger.warning("SMTP test delivery failed: %s", exc)
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
content={"result": "failed", "message": str(exc)},
|
||||
)
|
||||
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_200_OK,
|
||||
content={"result": "success", "message": "Test email sent successfully."},
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# POST /api/config/mqtt/test — M5-T10
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class _MqttConfigurationError(ValueError):
|
||||
"""Raised when MQTT settings are incomplete or disabled."""
|
||||
|
||||
|
||||
class _MqttConnectionError(RuntimeError):
|
||||
"""Raised when MQTT broker connection or publish fails."""
|
||||
|
||||
|
||||
@router.post(
|
||||
"/config/mqtt/test",
|
||||
responses={
|
||||
200: {"model": MqttTestResponse},
|
||||
400: {"model": MqttTestResponse},
|
||||
502: {"model": MqttTestResponse},
|
||||
},
|
||||
)
|
||||
def post_mqtt_test(
|
||||
settings: Settings = Depends(get_app_settings),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> JSONResponse:
|
||||
"""
|
||||
Test MQTT broker connectivity by attempting to connect and publishing a
|
||||
test message to ``<ha_discovery_prefix>/home-automation/test``.
|
||||
|
||||
The message is visible in MQTT Explorer (or any subscriber) so users can
|
||||
confirm the full broker publish path is working.
|
||||
|
||||
Three possible outcomes:
|
||||
- 200 { "result": "success", "message": ... }
|
||||
- 400 { "result": "config-error", "message": ... } (not configured)
|
||||
- 502 { "result": "failed", "message": ... } (connection/publish error)
|
||||
|
||||
MQTT credentials are never echoed in the response.
|
||||
"""
|
||||
try:
|
||||
_run_mqtt_test(settings)
|
||||
except _MqttConfigurationError as exc:
|
||||
logger.warning("MQTT test rejected due to configuration: %s", exc)
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
content={"result": "config-error", "message": str(exc)},
|
||||
)
|
||||
except _MqttConnectionError as exc:
|
||||
logger.warning("MQTT test connection/publish failed: %s", exc)
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
content={"result": "failed", "message": str(exc)},
|
||||
)
|
||||
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_200_OK,
|
||||
content={
|
||||
"result": "success",
|
||||
"message": (
|
||||
f"Test message published to "
|
||||
f"{settings.ha_discovery_prefix}/home-automation/test."
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _run_mqtt_test(settings: Settings) -> None:
|
||||
"""Attempt a transient MQTT connection and publish a test message.
|
||||
|
||||
Raises
|
||||
------
|
||||
_MqttConfigurationError
|
||||
When MQTT is not enabled or broker host is not configured.
|
||||
_MqttConnectionError
|
||||
When the broker is unreachable or the publish fails.
|
||||
"""
|
||||
import socket
|
||||
|
||||
import paho.mqtt.client as mqtt
|
||||
|
||||
if not settings.mqtt_broker_host:
|
||||
raise _MqttConfigurationError("MQTT broker host is not configured.")
|
||||
|
||||
test_topic = f"{settings.ha_discovery_prefix}/home-automation/test"
|
||||
test_payload = '{"source": "home-automation", "event": "mqtt_test"}'
|
||||
|
||||
connected_event = __import__("threading").Event()
|
||||
published_event = __import__("threading").Event()
|
||||
connect_error: list[str] = []
|
||||
|
||||
client = mqtt.Client(
|
||||
callback_api_version=mqtt.CallbackAPIVersion.VERSION2,
|
||||
client_id=mqtt_test_client_id(settings.mqtt_client_id),
|
||||
)
|
||||
|
||||
def _on_connect(
|
||||
_client: mqtt.Client,
|
||||
_userdata: object,
|
||||
_flags: mqtt.ConnectFlags,
|
||||
reason_code: mqtt.ReasonCode,
|
||||
_properties: mqtt.Properties | None,
|
||||
) -> None:
|
||||
if reason_code.is_failure:
|
||||
connect_error.append(f"Broker refused connection: {reason_code}")
|
||||
connected_event.set()
|
||||
|
||||
def _on_publish(
|
||||
_client: mqtt.Client,
|
||||
_userdata: object,
|
||||
_mid: int,
|
||||
_reason_code: mqtt.ReasonCode,
|
||||
_properties: mqtt.Properties | None,
|
||||
) -> None:
|
||||
published_event.set()
|
||||
|
||||
client.on_connect = _on_connect
|
||||
client.on_publish = _on_publish
|
||||
|
||||
if settings.mqtt_tls_enabled:
|
||||
try:
|
||||
client.tls_set()
|
||||
except Exception as exc:
|
||||
raise _MqttConnectionError(f"TLS setup failed: {exc}") from exc
|
||||
|
||||
if settings.mqtt_username:
|
||||
client.username_pw_set(
|
||||
username=settings.mqtt_username,
|
||||
password=settings.mqtt_password or None,
|
||||
)
|
||||
|
||||
client.loop_start()
|
||||
try:
|
||||
try:
|
||||
client.connect(
|
||||
host=settings.mqtt_broker_host,
|
||||
port=settings.mqtt_broker_port,
|
||||
keepalive=10,
|
||||
)
|
||||
except (OSError, socket.error) as exc:
|
||||
# Sanitise: ensure password never leaks into the error message.
|
||||
msg = _sanitize_mqtt_error(str(exc), settings.mqtt_password)
|
||||
raise _MqttConnectionError(f"Cannot reach broker: {msg}") from exc
|
||||
|
||||
# Wait up to 5 s for connection acknowledgement.
|
||||
if not connected_event.wait(timeout=5):
|
||||
raise _MqttConnectionError("Broker connection timed out (5 s).")
|
||||
|
||||
if connect_error:
|
||||
raise _MqttConnectionError(connect_error[0])
|
||||
|
||||
# Publish test message.
|
||||
client.publish(test_topic, payload=test_payload, qos=1, retain=False)
|
||||
# Wait up to 5 s for the publish ACK (QoS 1).
|
||||
if not published_event.wait(timeout=5):
|
||||
raise _MqttConnectionError("Publish timed out (5 s) — broker reachable but no ACK.")
|
||||
|
||||
finally:
|
||||
try:
|
||||
client.disconnect()
|
||||
except Exception:
|
||||
pass
|
||||
client.loop_stop()
|
||||
|
||||
|
||||
def _sanitize_mqtt_error(message: str, password: str | None) -> str:
|
||||
"""Replace *password* in *message* with ``[redacted]``."""
|
||||
if password:
|
||||
return message.replace(password, "[redacted]")
|
||||
return message
|
||||
@@ -0,0 +1,275 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException, Query, status
|
||||
from sqlalchemy import desc, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.routes.api.deps import require_csrf, require_session
|
||||
from app.dependencies import get_db
|
||||
from app.models.location import Location
|
||||
from app.models.poo import PooRecord
|
||||
from app.models.public_ip import PublicIPHistory, PublicIPState
|
||||
from app.schemas.data import (
|
||||
LocationRecord,
|
||||
LocationUpdateRequest,
|
||||
LocationsResponse,
|
||||
PooRecord as PooRecordSchema,
|
||||
PooResponse,
|
||||
PooUpdateRequest,
|
||||
PublicIPHistorySchema,
|
||||
PublicIPResponse,
|
||||
PublicIPStateSchema,
|
||||
)
|
||||
from app.services.auth import AuthenticatedSession
|
||||
from app.services.location import delete_location, update_location
|
||||
from app.services.poo import delete_poo_record, update_poo_record
|
||||
|
||||
router = APIRouter(prefix="/api", tags=["api-data"])
|
||||
|
||||
|
||||
@router.get("/locations", response_model=LocationsResponse)
|
||||
def get_locations(
|
||||
limit: int = Query(default=1000, ge=1, le=5000),
|
||||
offset: int = Query(default=0, ge=0),
|
||||
start: str | None = Query(default=None),
|
||||
end: str | None = Query(default=None),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> LocationsResponse:
|
||||
"""
|
||||
Return location records with optional time-window filtering and pagination.
|
||||
|
||||
- ``start`` / ``end`` are ISO8601 strings; filtering is **inclusive** on both bounds.
|
||||
- Results are ordered by ``datetime`` ascending.
|
||||
- ``limit`` is capped at 5000 to prevent full-table exports.
|
||||
"""
|
||||
stmt = select(Location)
|
||||
|
||||
if start is not None:
|
||||
stmt = stmt.where(Location.datetime >= start)
|
||||
if end is not None:
|
||||
stmt = stmt.where(Location.datetime <= end)
|
||||
|
||||
stmt = stmt.order_by(Location.datetime).offset(offset).limit(limit)
|
||||
|
||||
rows = db.execute(stmt).scalars().all()
|
||||
|
||||
items = [
|
||||
LocationRecord(
|
||||
person=row.person,
|
||||
datetime=row.datetime,
|
||||
latitude=row.latitude,
|
||||
longitude=row.longitude,
|
||||
altitude=row.altitude,
|
||||
)
|
||||
for row in rows
|
||||
]
|
||||
|
||||
return LocationsResponse(items=items, limit=limit, offset=offset)
|
||||
|
||||
|
||||
@router.get("/poo", response_model=PooResponse)
|
||||
def get_poo(
|
||||
limit: int = Query(default=100, ge=1, le=1000),
|
||||
offset: int = Query(default=0, ge=0),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> PooResponse:
|
||||
"""
|
||||
Return poo records ordered by timestamp descending (most recent first).
|
||||
|
||||
``limit`` is capped at 1000 to prevent full-table exports.
|
||||
"""
|
||||
stmt = (
|
||||
select(PooRecord)
|
||||
.order_by(desc(PooRecord.timestamp))
|
||||
.offset(offset)
|
||||
.limit(limit)
|
||||
)
|
||||
|
||||
rows = db.execute(stmt).scalars().all()
|
||||
|
||||
items = [
|
||||
PooRecordSchema(
|
||||
timestamp=row.timestamp,
|
||||
status=row.status,
|
||||
latitude=row.latitude,
|
||||
longitude=row.longitude,
|
||||
)
|
||||
for row in rows
|
||||
]
|
||||
|
||||
return PooResponse(items=items, limit=limit, offset=offset)
|
||||
|
||||
|
||||
@router.get("/public-ip", response_model=PublicIPResponse)
|
||||
def get_public_ip(
|
||||
limit: int = Query(default=100, ge=1, le=1000),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> PublicIPResponse:
|
||||
"""
|
||||
Return the current public IP state and recent history.
|
||||
|
||||
- ``state`` is ``null`` if no IP check has been performed yet.
|
||||
- ``history`` is ordered by ``observed_at`` descending (most recent first).
|
||||
- ``limit`` applies to the history list and is capped at 1000.
|
||||
"""
|
||||
state_row = db.execute(
|
||||
select(PublicIPState).where(PublicIPState.id == 1).limit(1)
|
||||
).scalar_one_or_none()
|
||||
|
||||
history_rows = db.execute(
|
||||
select(PublicIPHistory).order_by(desc(PublicIPHistory.observed_at)).limit(limit)
|
||||
).scalars().all()
|
||||
|
||||
state = PublicIPStateSchema.model_validate(state_row) if state_row is not None else None
|
||||
history = [PublicIPHistorySchema.model_validate(row) for row in history_rows]
|
||||
|
||||
return PublicIPResponse(state=state, history=history)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# PATCH /api/locations/{person}/{datetime}
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.patch("/locations/{person}/{datetime}", response_model=LocationRecord)
|
||||
def patch_location(
|
||||
person: str,
|
||||
datetime: str,
|
||||
body: LocationUpdateRequest = Body(default=LocationUpdateRequest()),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> LocationRecord:
|
||||
"""
|
||||
Update the non-PK fields of a single location record.
|
||||
|
||||
- ``person`` and ``datetime`` identify the row (composite PK) and are immutable.
|
||||
- Only ``latitude``, ``longitude``, and ``altitude`` may be updated.
|
||||
- Omitted body fields are left unchanged.
|
||||
- Returns **404** if the PK does not exist.
|
||||
"""
|
||||
row = update_location(
|
||||
db,
|
||||
person,
|
||||
datetime,
|
||||
latitude=body.latitude,
|
||||
longitude=body.longitude,
|
||||
altitude=body.altitude,
|
||||
)
|
||||
if row is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="location record not found",
|
||||
)
|
||||
return LocationRecord(
|
||||
person=row.person,
|
||||
datetime=row.datetime,
|
||||
latitude=row.latitude,
|
||||
longitude=row.longitude,
|
||||
altitude=row.altitude,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DELETE /api/locations/{person}/{datetime}
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/locations/{person}/{datetime}",
|
||||
status_code=status.HTTP_204_NO_CONTENT,
|
||||
response_model=None,
|
||||
)
|
||||
def delete_location_record(
|
||||
person: str,
|
||||
datetime: str,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> None:
|
||||
"""
|
||||
Delete the single location record identified by its composite PK.
|
||||
|
||||
- Exactly one row is deleted; **404** if the PK does not exist.
|
||||
- No batch delete / truncate path is available.
|
||||
"""
|
||||
deleted = delete_location(db, person, datetime)
|
||||
if not deleted:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="location record not found",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# PATCH /api/poo/{timestamp}
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.patch("/poo/{timestamp}", response_model=PooRecordSchema)
|
||||
def patch_poo(
|
||||
timestamp: str,
|
||||
body: PooUpdateRequest = Body(default=PooUpdateRequest()),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> PooRecordSchema:
|
||||
"""
|
||||
Update the non-PK fields of a single poo record.
|
||||
|
||||
- ``timestamp`` is the PK and is immutable.
|
||||
- Only ``status``, ``latitude``, and ``longitude`` may be updated.
|
||||
- Omitted body fields are left unchanged.
|
||||
- Returns **404** if the PK does not exist.
|
||||
"""
|
||||
row = update_poo_record(
|
||||
db,
|
||||
timestamp,
|
||||
status=body.status,
|
||||
latitude=body.latitude,
|
||||
longitude=body.longitude,
|
||||
)
|
||||
if row is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="poo record not found",
|
||||
)
|
||||
return PooRecordSchema(
|
||||
timestamp=row.timestamp,
|
||||
status=row.status,
|
||||
latitude=row.latitude,
|
||||
longitude=row.longitude,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DELETE /api/poo/{timestamp}
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/poo/{timestamp}",
|
||||
status_code=status.HTTP_204_NO_CONTENT,
|
||||
response_model=None,
|
||||
)
|
||||
def delete_poo(
|
||||
timestamp: str,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> None:
|
||||
"""
|
||||
Delete the single poo record identified by its PK.
|
||||
|
||||
- Exactly one row is deleted; **404** if the PK does not exist.
|
||||
- No batch delete / truncate path is available.
|
||||
"""
|
||||
deleted = delete_poo_record(db, timestamp)
|
||||
if not deleted:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="poo record not found",
|
||||
)
|
||||
@@ -0,0 +1,28 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import Depends, Header, HTTPException, status
|
||||
|
||||
from app.dependencies import get_current_auth_session
|
||||
from app.services.auth import AuthenticatedSession
|
||||
|
||||
|
||||
def require_session(
|
||||
auth: AuthenticatedSession | None = Depends(get_current_auth_session),
|
||||
) -> AuthenticatedSession:
|
||||
if auth is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="authentication required",
|
||||
)
|
||||
return auth
|
||||
|
||||
|
||||
def require_csrf(
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
x_csrf_token: str | None = Header(default=None, alias="X-CSRF-Token"),
|
||||
) -> None:
|
||||
if not x_csrf_token:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="missing CSRF token",
|
||||
)
|
||||
@@ -0,0 +1,607 @@
|
||||
"""Energy data API: prices, costs, summary, DSMR, recompute, Tibber test (M6-T09).
|
||||
|
||||
All endpoints are under /api/energy, require an authenticated session, and
|
||||
write endpoints (POST) additionally require a non-empty X-CSRF-Token header.
|
||||
|
||||
Route prefix note
|
||||
-----------------
|
||||
This router shares the ``/api/energy`` prefix with ``energy_contracts.py``
|
||||
(which handles contract CRUD at /contracts/* and /profiles). The sub-paths
|
||||
used here (/prices, /costs, /costs/summary, /costs/recompute, /dsmr/latest,
|
||||
/tibber/test) are disjoint from the contract router's paths, so there is no
|
||||
conflict.
|
||||
|
||||
Tibber token security
|
||||
---------------------
|
||||
``POST /api/energy/tibber/test`` calls the Tibber API but **never** echoes the
|
||||
token in the response body or in log messages. Three-state logic mirrors the
|
||||
MQTT test endpoint (M5-T10, app/api/routes/api/config.py::post_mqtt_test):
|
||||
|
||||
200 { result: "success", message: ..., price: {...} }
|
||||
400 { result: "config-error", message: ... }
|
||||
502 { result: "failed", message: ... }
|
||||
|
||||
Recompute safety
|
||||
----------------
|
||||
``POST /api/energy/costs/recompute`` is idempotent: it calls
|
||||
``energy_cost.recompute_range`` which upserts existing rows without deleting
|
||||
anything. The endpoint enforces a maximum time-window of 366 days to avoid
|
||||
unbounded recomputation triggered by erroneous client requests.
|
||||
|
||||
Prices endpoint behaviour
|
||||
-------------------------
|
||||
``GET /api/energy/prices`` queries the ``tibber_price`` table for tibber
|
||||
contracts, or derives the effective fixed-tariff prices for manual contracts
|
||||
using the same formula as the billing engine (_manual_strategy in strategies.py):
|
||||
|
||||
buy_dal = energy.buy.dal + energy.energy_tax + energy.ode
|
||||
buy_normal = energy.buy.normal + energy.energy_tax + energy.ode
|
||||
sell_dal = energy.sell.dal (no tax added to sell price)
|
||||
sell_normal = energy.sell.normal
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import Any, Literal
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, status
|
||||
from fastapi.responses import JSONResponse
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.routes.api.deps import require_csrf, require_session
|
||||
from app.dependencies import get_app_settings, get_db
|
||||
from app.config import Settings
|
||||
from app.integrations.tibber.client import (
|
||||
TibberAuthError,
|
||||
TibberError,
|
||||
fetch_current_price,
|
||||
)
|
||||
from app.models.energy import DsmrReading, EnergyCostPeriod, TibberPrice
|
||||
from app.schemas.energy import (
|
||||
CostPeriodSchema,
|
||||
CostsResponse,
|
||||
DsmrLatestResponse,
|
||||
ManualTariffSchema,
|
||||
PricePointSchema,
|
||||
PricesResponse,
|
||||
RecomputeResponse,
|
||||
SummaryResponse,
|
||||
TibberTestPriceSchema,
|
||||
TibberTestResponse,
|
||||
)
|
||||
from app.services.auth import AuthenticatedSession
|
||||
from app.services.contracts import active_contract_version_at
|
||||
from app.services.energy_cost import recompute_range, summarize
|
||||
from app.services.timezone import local_midnight_utc, local_now
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/api/energy", tags=["api-energy"])
|
||||
|
||||
# Maximum number of cost periods returned per request (mirrors modbus readings cap).
|
||||
_COSTS_LIMIT_MAX = 5000
|
||||
|
||||
# Maximum allowed time-window for recompute to prevent unbounded computation.
|
||||
_RECOMPUTE_MAX_DAYS = 366
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _as_utc(dt: datetime) -> datetime:
|
||||
"""Attach UTC tzinfo to a naive datetime (SQLite read-back workaround)."""
|
||||
if dt.tzinfo is None:
|
||||
return dt.replace(tzinfo=UTC)
|
||||
return dt
|
||||
|
||||
|
||||
def _manual_tariff_from_values(values: dict[str, Any]) -> ManualTariffSchema:
|
||||
"""Derive the effective fixed tariff from a manual contract version's values dict.
|
||||
|
||||
Mirrors the _manual_strategy formula (strategies.py):
|
||||
buy_dal = energy.buy.dal + energy.energy_tax + energy.ode
|
||||
buy_normal = energy.buy.normal + energy.energy_tax + energy.ode
|
||||
sell_dal = energy.sell.dal (no tax added)
|
||||
sell_normal = energy.sell.normal
|
||||
"""
|
||||
from decimal import Decimal
|
||||
|
||||
def _d(v: Any) -> Decimal:
|
||||
return Decimal(str(v or 0))
|
||||
|
||||
energy = values.get("energy", {})
|
||||
buy = energy.get("buy", {})
|
||||
sell = energy.get("sell", {})
|
||||
energy_tax = _d(energy.get("energy_tax", 0))
|
||||
ode = _d(energy.get("ode", 0))
|
||||
|
||||
buy_dal = _d(buy.get("dal", 0)) + energy_tax + ode
|
||||
buy_normal = _d(buy.get("normal", 0)) + energy_tax + ode
|
||||
sell_dal = _d(sell.get("dal", 0))
|
||||
sell_normal = _d(sell.get("normal", 0))
|
||||
|
||||
return ManualTariffSchema(
|
||||
buy_dal=float(buy_dal),
|
||||
buy_normal=float(buy_normal),
|
||||
sell_dal=float(sell_dal),
|
||||
sell_normal=float(sell_normal),
|
||||
)
|
||||
|
||||
|
||||
def _electricity_prices_response(response: PricesResponse) -> JSONResponse:
|
||||
"""Preserve the exact pre-scope electricity response body."""
|
||||
return JSONResponse(
|
||||
content=response.model_dump(
|
||||
mode="json", include={"kind", "currency", "points", "tariff"}
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /api/energy/prices
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/prices", response_model=PricesResponse, response_model_exclude_none=True)
|
||||
def get_prices(
|
||||
scope: Literal["electricity", "thermal"] = Query("electricity"),
|
||||
start: datetime | None = Query(
|
||||
default=None,
|
||||
description="Inclusive start of the time window (ISO 8601). "
|
||||
"Defaults to the start of today UTC when omitted.",
|
||||
),
|
||||
end: datetime | None = Query(
|
||||
default=None,
|
||||
description="Inclusive end of the time window (ISO 8601). "
|
||||
"Defaults to the end of tomorrow UTC when omitted.",
|
||||
),
|
||||
limit: int = Query(
|
||||
default=500,
|
||||
ge=1,
|
||||
le=_COSTS_LIMIT_MAX,
|
||||
description="Maximum number of Tibber price points to return.",
|
||||
),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> PricesResponse:
|
||||
"""Return the price curve for the active contract.
|
||||
|
||||
**Tibber contracts** (kind="tibber"):
|
||||
Fetches ``tibber_price`` rows within ``[start, end]``, ordered ascending
|
||||
by ``starts_at``. At most ``limit`` rows are returned (most recent first
|
||||
within the window, then reversed to ascending order — identical to the
|
||||
modbus readings pattern).
|
||||
|
||||
Response ``points`` carries per-slot:
|
||||
- ``buy = total`` (Tibber all-inclusive price)
|
||||
- ``sell = total − energy_tax − sell_fee − sell_adjust`` (from active version values)
|
||||
- ``level`` (Tibber price level, may be null)
|
||||
|
||||
``tariff`` is null.
|
||||
|
||||
**Manual contracts** (kind="manual"):
|
||||
``points`` is empty. ``tariff`` carries the four effective prices
|
||||
derived using the billing engine formula:
|
||||
- ``buy_dal = energy.buy.dal + energy_tax + ode``
|
||||
- ``buy_normal = energy.buy.normal + energy_tax + ode``
|
||||
- ``sell_dal = energy.sell.dal``
|
||||
- ``sell_normal = energy.sell.normal``
|
||||
|
||||
**No active contract**: returns kind=null, currency="EUR", points=[], tariff=null (200).
|
||||
"""
|
||||
now = datetime.now(UTC)
|
||||
|
||||
if scope == "thermal":
|
||||
version = active_contract_version_at(db, now, scope="thermal")
|
||||
if version is None:
|
||||
return PricesResponse(kind=None, currency="EUR", points=[], tariff=None)
|
||||
return PricesResponse(
|
||||
kind="district_heating", currency=version.contract.currency,
|
||||
contract_version_id=version.id, effective_from=_as_utc(version.effective_from),
|
||||
effective_to=_as_utc(version.effective_to) if version.effective_to else None,
|
||||
values=version.values, points=[], tariff=None,
|
||||
)
|
||||
|
||||
# Default window: today + tomorrow.
|
||||
if start is None:
|
||||
start = now.replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
if end is None:
|
||||
end = (start + timedelta(days=2)).replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
|
||||
start_utc = _as_utc(start)
|
||||
end_utc = _as_utc(end)
|
||||
|
||||
# Resolve the active contract version at the start of the window.
|
||||
version = active_contract_version_at(db, start_utc)
|
||||
|
||||
if version is None:
|
||||
return _electricity_prices_response(PricesResponse(
|
||||
kind=None,
|
||||
currency="EUR",
|
||||
points=[],
|
||||
tariff=None,
|
||||
))
|
||||
|
||||
contract = version.contract
|
||||
currency = contract.currency
|
||||
|
||||
if contract.kind == "tibber":
|
||||
# Fetch tibber_price rows in the window.
|
||||
stmt = (
|
||||
select(TibberPrice)
|
||||
.where(
|
||||
TibberPrice.starts_at >= start_utc,
|
||||
TibberPrice.starts_at <= end_utc,
|
||||
)
|
||||
.order_by(TibberPrice.starts_at.desc())
|
||||
.limit(limit)
|
||||
)
|
||||
rows = list(reversed(db.execute(stmt).scalars().all()))
|
||||
|
||||
# Derive sell price per-point using version values
|
||||
# (energy_tax + sell_fee + sell_adjust).
|
||||
from decimal import Decimal
|
||||
|
||||
def _d(v: Any) -> Decimal:
|
||||
return Decimal(str(v or 0))
|
||||
|
||||
energy = version.values.get("energy", {}) if version.values else {}
|
||||
energy_tax = _d(energy.get("energy_tax", 0))
|
||||
sell_fee = _d(energy.get("sell_fee", 0))
|
||||
sell_adjust = _d(energy.get("sell_adjust", 0))
|
||||
|
||||
points = []
|
||||
for row in rows:
|
||||
total = _d(row.total)
|
||||
sell = float(total - energy_tax - sell_fee - sell_adjust)
|
||||
points.append(
|
||||
PricePointSchema(
|
||||
starts_at=_as_utc(row.starts_at),
|
||||
buy=row.total,
|
||||
sell=sell,
|
||||
level=row.level,
|
||||
)
|
||||
)
|
||||
|
||||
return _electricity_prices_response(PricesResponse(
|
||||
kind="tibber",
|
||||
currency=currency,
|
||||
points=points,
|
||||
tariff=None,
|
||||
))
|
||||
|
||||
elif contract.kind == "manual":
|
||||
tariff = _manual_tariff_from_values(version.values or {})
|
||||
return _electricity_prices_response(PricesResponse(
|
||||
kind="manual",
|
||||
currency=currency,
|
||||
points=[],
|
||||
tariff=tariff,
|
||||
))
|
||||
|
||||
else:
|
||||
# Unknown kind — return empty response gracefully.
|
||||
return _electricity_prices_response(PricesResponse(
|
||||
kind=contract.kind,
|
||||
currency=currency,
|
||||
points=[],
|
||||
tariff=None,
|
||||
))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /api/energy/costs
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/costs", response_model=CostsResponse)
|
||||
def get_costs(
|
||||
start: datetime | None = Query(
|
||||
default=None,
|
||||
description="Inclusive lower bound for period_start (ISO 8601).",
|
||||
),
|
||||
end: datetime | None = Query(
|
||||
default=None,
|
||||
description="Inclusive upper bound for period_start (ISO 8601).",
|
||||
),
|
||||
limit: int = Query(
|
||||
default=500,
|
||||
ge=1,
|
||||
le=_COSTS_LIMIT_MAX,
|
||||
description=f"Maximum number of cost periods to return (default 500, max {_COSTS_LIMIT_MAX}).",
|
||||
),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> CostsResponse:
|
||||
"""Return energy_cost_period rows within a time window.
|
||||
|
||||
Rows are ordered by ``period_start`` ascending. When the window contains
|
||||
more rows than ``limit``, the **most recent** N rows are returned (DESC LIMIT),
|
||||
then reversed to ascending order — identical to the modbus readings pattern.
|
||||
|
||||
Query parameters:
|
||||
- ``start``: inclusive lower bound on ``period_start`` (ISO 8601 datetime).
|
||||
- ``end``: inclusive upper bound on ``period_start`` (ISO 8601 datetime).
|
||||
- ``limit``: max rows to return (default 500, max {_COSTS_LIMIT_MAX}).
|
||||
"""
|
||||
stmt = (
|
||||
select(EnergyCostPeriod)
|
||||
.order_by(EnergyCostPeriod.period_start.desc())
|
||||
.limit(limit)
|
||||
)
|
||||
|
||||
if start is not None:
|
||||
stmt = stmt.where(EnergyCostPeriod.period_start >= _as_utc(start))
|
||||
if end is not None:
|
||||
stmt = stmt.where(EnergyCostPeriod.period_start <= _as_utc(end))
|
||||
|
||||
rows = list(reversed(db.execute(stmt).scalars().all()))
|
||||
items = [CostPeriodSchema.model_validate(r) for r in rows]
|
||||
return CostsResponse(items=items, total=len(items))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /api/energy/costs/summary
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/costs/summary", response_model=SummaryResponse)
|
||||
def get_costs_summary(
|
||||
start: datetime | None = Query(
|
||||
default=None,
|
||||
description=(
|
||||
"Inclusive start of the summary interval (ISO 8601). "
|
||||
"Defaults to the start of the current UTC day."
|
||||
),
|
||||
),
|
||||
end: datetime | None = Query(
|
||||
default=None,
|
||||
description=(
|
||||
"Exclusive end of the summary interval (ISO 8601). "
|
||||
"Defaults to the start of the next UTC day (i.e. today's full data)."
|
||||
),
|
||||
),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> SummaryResponse:
|
||||
"""Aggregate billing for a time interval.
|
||||
|
||||
Calls ``energy_cost.summarize(session, start, end)`` which computes:
|
||||
|
||||
total_payable = Σ(net_cost) + fixed_costs − credits
|
||||
|
||||
where ``fixed_costs`` is (network_fee + management_fee) apportioned to the
|
||||
interval length in days (÷30 per month), and ``credits`` is heffingskorting
|
||||
apportioned similarly (÷365 per year).
|
||||
|
||||
Both ``fixed_costs`` and ``credits`` are derived from the **currently active
|
||||
contract version at ``end``**. When no active contract exists they are 0.
|
||||
"""
|
||||
if start is None or end is None:
|
||||
# Default to the server's local today: [local_midnight, next_local_midnight).
|
||||
# This ensures "today" aligns with the local calendar day (NL time) rather
|
||||
# than UTC midnight.
|
||||
local_today = local_now().date()
|
||||
if start is None:
|
||||
start = local_midnight_utc(local_today)
|
||||
if end is None:
|
||||
end = local_midnight_utc(local_today + timedelta(days=1))
|
||||
|
||||
result = summarize(db, _as_utc(start), _as_utc(end))
|
||||
return SummaryResponse(**result)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /api/energy/dsmr/latest
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/dsmr/latest", response_model=DsmrLatestResponse)
|
||||
def get_dsmr_latest(
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> DsmrLatestResponse:
|
||||
"""Return the most recent dsmr_reading row.
|
||||
|
||||
Returns ``{"found": false, "recorded_at": null, "payload": null}`` (200, not
|
||||
404) when no rows exist yet, so the front-end can distinguish "no data" from
|
||||
a server error.
|
||||
"""
|
||||
row = db.execute(
|
||||
select(DsmrReading)
|
||||
.order_by(DsmrReading.recorded_at.desc())
|
||||
.limit(1)
|
||||
).scalar_one_or_none()
|
||||
|
||||
if row is None:
|
||||
return DsmrLatestResponse(found=False)
|
||||
|
||||
return DsmrLatestResponse(
|
||||
found=True,
|
||||
recorded_at=_as_utc(row.recorded_at),
|
||||
payload=row.payload,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# POST /api/energy/costs/recompute
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post(
|
||||
"/costs/recompute",
|
||||
responses={
|
||||
200: {"model": RecomputeResponse},
|
||||
422: {"description": "Validation error (missing window or range too large)"},
|
||||
},
|
||||
)
|
||||
def post_recompute(
|
||||
start: datetime = Query(
|
||||
...,
|
||||
description="Inclusive start of the recompute window (ISO 8601). Required.",
|
||||
),
|
||||
end: datetime = Query(
|
||||
...,
|
||||
description="Exclusive end of the recompute window (ISO 8601). Required.",
|
||||
),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> RecomputeResponse:
|
||||
"""Idempotently recompute billing records in a time window.
|
||||
|
||||
Calls ``energy_cost.recompute_range(session, start, end)`` which overwrites
|
||||
existing rows (including successful ones) for every UTC quarter-hour boundary
|
||||
in ``[start, end)``.
|
||||
|
||||
**Idempotency**: repeated calls with the same window produce the same
|
||||
outcome. No rows are deleted; only upserted.
|
||||
|
||||
**Window constraint**: the maximum allowed range is {_RECOMPUTE_MAX_DAYS} days.
|
||||
Requests exceeding this return 422.
|
||||
|
||||
Returns the number of periods for which a billing record was written.
|
||||
Periods skipped due to missing contract or missing Tibber price are not counted.
|
||||
"""
|
||||
start_utc = _as_utc(start)
|
||||
end_utc = _as_utc(end)
|
||||
|
||||
if end_utc <= start_utc:
|
||||
from fastapi import HTTPException
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="'end' must be strictly after 'start'.",
|
||||
)
|
||||
|
||||
span_days = (end_utc - start_utc).total_seconds() / 86400
|
||||
if span_days > _RECOMPUTE_MAX_DAYS:
|
||||
from fastapi import HTTPException
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=(
|
||||
f"Time window is {span_days:.1f} days which exceeds the maximum of "
|
||||
f"{_RECOMPUTE_MAX_DAYS} days. Use a smaller window."
|
||||
),
|
||||
)
|
||||
|
||||
n = recompute_range(db, start_utc, end_utc)
|
||||
logger.info(
|
||||
"POST /api/energy/costs/recompute [%s, %s): wrote %d period(s).",
|
||||
start_utc.isoformat(),
|
||||
end_utc.isoformat(),
|
||||
n,
|
||||
)
|
||||
return RecomputeResponse(recomputed=n)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# POST /api/energy/tibber/test
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post(
|
||||
"/tibber/test",
|
||||
responses={
|
||||
200: {"model": TibberTestResponse},
|
||||
400: {"model": TibberTestResponse},
|
||||
502: {"model": TibberTestResponse},
|
||||
},
|
||||
)
|
||||
def post_tibber_test(
|
||||
settings: Settings = Depends(get_app_settings),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> JSONResponse:
|
||||
"""Test Tibber API connectivity by fetching the current price point.
|
||||
|
||||
Three possible outcomes:
|
||||
|
||||
- **200** ``{ result: "success", message: ..., price: {...} }``
|
||||
The Tibber API responded with a valid current price. ``price`` contains
|
||||
starts_at, total, energy, tax, currency, and level.
|
||||
|
||||
- **400** ``{ result: "config-error", message: ... }``
|
||||
The Tibber API token is empty or not configured.
|
||||
|
||||
- **502** ``{ result: "failed", message: ... }``
|
||||
The API call failed (authentication rejected, network error, timeout,
|
||||
unexpected response, etc.).
|
||||
|
||||
The API token is **never** included in the response body or logged.
|
||||
"""
|
||||
token = settings.tibber_api_token
|
||||
home_id = settings.tibber_home_id or None # treat empty string as None
|
||||
|
||||
if not token:
|
||||
logger.info("POST /api/energy/tibber/test: no token configured.")
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
content=TibberTestResponse(
|
||||
result="config-error",
|
||||
message=(
|
||||
"Tibber API token is not configured. "
|
||||
"Set TIBBER_API_TOKEN in the Config page."
|
||||
),
|
||||
price=None,
|
||||
).model_dump(mode="json"),
|
||||
)
|
||||
|
||||
try:
|
||||
price_point = fetch_current_price(token, home_id)
|
||||
except TibberAuthError:
|
||||
logger.warning("POST /api/energy/tibber/test: authentication failed.")
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
content=TibberTestResponse(
|
||||
result="failed",
|
||||
message=(
|
||||
"Tibber API authentication failed. "
|
||||
"Check that your API token is correct."
|
||||
),
|
||||
price=None,
|
||||
).model_dump(mode="json"),
|
||||
)
|
||||
except TibberError as exc:
|
||||
logger.warning("POST /api/energy/tibber/test: API call failed — %s", exc)
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
content=TibberTestResponse(
|
||||
result="failed",
|
||||
message=f"Tibber API call failed: {exc}",
|
||||
price=None,
|
||||
).model_dump(mode="json"),
|
||||
)
|
||||
|
||||
price_schema = TibberTestPriceSchema(
|
||||
starts_at=price_point.starts_at,
|
||||
total=price_point.total,
|
||||
energy=price_point.energy,
|
||||
tax=price_point.tax,
|
||||
currency=price_point.currency,
|
||||
level=price_point.level,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"POST /api/energy/tibber/test: success (starts_at=%s, total=%s %s).",
|
||||
price_point.starts_at.isoformat(),
|
||||
price_point.total,
|
||||
price_point.currency,
|
||||
)
|
||||
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_200_OK,
|
||||
content=TibberTestResponse(
|
||||
result="success",
|
||||
message=(
|
||||
f"Tibber API connected. Current price: "
|
||||
f"{price_point.total} {price_point.currency}/kWh "
|
||||
f"(starts {price_point.starts_at.isoformat()})."
|
||||
),
|
||||
price=price_schema,
|
||||
).model_dump(mode="json"),
|
||||
)
|
||||
@@ -0,0 +1,345 @@
|
||||
"""EnergyContract CRUD, versioning, and pricing-profile API (M6-T04).
|
||||
|
||||
All endpoints are under /api/energy, require an authenticated session, and
|
||||
write endpoints (POST/PATCH) additionally require a non-empty X-CSRF-Token
|
||||
header.
|
||||
|
||||
There is deliberately no DELETE endpoint for contracts: the FK RESTRICT
|
||||
constraint on energy_contract_version.contract_id and
|
||||
energy_cost_period.contract_version_id prevents accidental deletion of
|
||||
contracts that have billing history. Users should instead deactivate a
|
||||
contract (PATCH active=false) to stop it from being used.
|
||||
|
||||
Route ordering note
|
||||
-------------------
|
||||
GET /api/energy/profiles and GET /api/energy/contracts/{id} are on separate
|
||||
path prefixes (/energy/profiles vs /energy/contracts/{id}) so there is no
|
||||
ambiguity even without extra ordering gymnastics.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.routes.api.deps import require_csrf, require_session
|
||||
from app.dependencies import get_db
|
||||
from app.integrations.pricing.profiles import (
|
||||
ProfileNotFoundError,
|
||||
ProfileValidationError,
|
||||
list_profiles,
|
||||
)
|
||||
from app.schemas.energy_contract import (
|
||||
ContractCreate,
|
||||
ContractDetailResponse,
|
||||
ContractListResponse,
|
||||
ContractPatch,
|
||||
ContractResponse,
|
||||
ContractVersionResponse,
|
||||
ProfilesResponse,
|
||||
VersionCreate,
|
||||
)
|
||||
from app.services.auth import AuthenticatedSession
|
||||
from app.services import timezone as _tz_mod
|
||||
from app.services.contracts import (
|
||||
CONTRACT_KIND_SCOPES,
|
||||
ContractVersionError,
|
||||
ContractScopeError,
|
||||
activate_contract,
|
||||
add_version,
|
||||
create_contract,
|
||||
deactivate_contract,
|
||||
get_contract_or_none,
|
||||
list_contracts,
|
||||
)
|
||||
from app.services.tibber_prices import trigger_tibber_refresh
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/api/energy", tags=["api-energy-contracts"])
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_contract_or_404(db: Session, contract_id: int):
|
||||
"""Return the contract with the given id or raise 404."""
|
||||
contract = get_contract_or_none(db, contract_id)
|
||||
if contract is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"Energy contract {contract_id!r} not found.",
|
||||
)
|
||||
return contract
|
||||
|
||||
|
||||
def _contract_detail(db: Session, contract) -> ContractDetailResponse:
|
||||
"""Build a ContractDetailResponse for *contract*, loading versions."""
|
||||
# Eager-load versions ordered by effective_from for consistent display.
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.models.energy import EnergyContractVersion
|
||||
|
||||
versions = list(
|
||||
db.execute(
|
||||
select(EnergyContractVersion)
|
||||
.where(EnergyContractVersion.contract_id == contract.id)
|
||||
.order_by(EnergyContractVersion.effective_from)
|
||||
)
|
||||
.scalars()
|
||||
.all()
|
||||
)
|
||||
|
||||
version_schemas = [ContractVersionResponse.model_validate(v) for v in versions]
|
||||
return ContractDetailResponse(
|
||||
id=contract.id,
|
||||
name=contract.name,
|
||||
kind=contract.kind,
|
||||
scope=contract.scope,
|
||||
active=contract.active,
|
||||
currency=contract.currency,
|
||||
created_at=contract.created_at,
|
||||
updated_at=contract.updated_at,
|
||||
versions=version_schemas,
|
||||
)
|
||||
|
||||
|
||||
def _raise_422_for_profile_error(exc: Exception) -> Any:
|
||||
"""Convert a ProfileValidationError or ProfileNotFoundError into HTTP 422."""
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=str(exc),
|
||||
)
|
||||
|
||||
|
||||
def _localize_effective_from(dt: datetime | None) -> datetime:
|
||||
"""Resolve *dt* to an aware UTC datetime for storage.
|
||||
|
||||
Rules (Principle A):
|
||||
- If *dt* is None → use ``datetime.now(UTC)`` (unchanged from before).
|
||||
- If *dt* is timezone-aware → convert to UTC as-is.
|
||||
- If *dt* is timezone-naive → interpret as server local wall-clock time,
|
||||
localize with ``local_tz()``, then convert to UTC.
|
||||
|
||||
This means a front-end that sends ``"2026-06-25T00:00:00"`` (no Z) has it
|
||||
interpreted as local midnight (e.g. CEST = UTC+2 → stored as 2026-06-24T22:00:00Z),
|
||||
not as UTC midnight.
|
||||
"""
|
||||
if dt is None:
|
||||
return datetime.now(UTC)
|
||||
if dt.tzinfo is not None:
|
||||
# Already aware: convert to UTC.
|
||||
return dt.astimezone(UTC)
|
||||
# Naive: assume server local wall-clock.
|
||||
tz = _tz_mod.local_tz()
|
||||
local_dt = dt.replace(tzinfo=tz)
|
||||
return local_dt.astimezone(UTC)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /api/energy/profiles
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/profiles", response_model=ProfilesResponse)
|
||||
def get_profiles(
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> ProfilesResponse:
|
||||
"""List all available pricing profile structures.
|
||||
|
||||
Returns the full profile structure for each supported contract kind
|
||||
(``manual`` and ``tibber``). The front-end uses this to dynamically
|
||||
render the correct fields and labels for the contract creation/editing form.
|
||||
"""
|
||||
return ProfilesResponse(profiles=list_profiles())
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /api/energy/contracts
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/contracts", response_model=ContractListResponse)
|
||||
def list_energy_contracts(
|
||||
scope: str = "electricity",
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> ContractListResponse:
|
||||
"""List all energy contracts with their active status.
|
||||
|
||||
Scope defaults to ``electricity`` for old clients. Returns a flat list (no embedded version history); use
|
||||
GET /api/energy/contracts/{id} to fetch the full version history for a
|
||||
specific contract.
|
||||
"""
|
||||
if scope not in set(CONTRACT_KIND_SCOPES.values()):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"Unknown energy contract scope: {scope!r}",
|
||||
)
|
||||
contracts = list_contracts(db, scope=scope)
|
||||
items = [ContractResponse.model_validate(c) for c in contracts]
|
||||
return ContractListResponse(items=items, total=len(items))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# POST /api/energy/contracts
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post(
|
||||
"/contracts",
|
||||
response_model=ContractDetailResponse,
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
)
|
||||
def create_energy_contract(
|
||||
body: ContractCreate,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> ContractDetailResponse:
|
||||
"""Create a new energy contract with an initial pricing version.
|
||||
|
||||
The ``values`` dict is validated against the YAML profile for the given
|
||||
``kind``; non-conforming values result in 422 Unprocessable Entity.
|
||||
The new contract is created with ``active=False``; use
|
||||
PATCH /api/energy/contracts/{id} with ``active=true`` to activate it.
|
||||
"""
|
||||
effective_from = _localize_effective_from(body.effective_from)
|
||||
|
||||
try:
|
||||
contract = create_contract(
|
||||
db,
|
||||
name=body.name,
|
||||
kind=body.kind,
|
||||
currency=body.currency,
|
||||
scope=body.scope,
|
||||
values=body.values,
|
||||
effective_from=effective_from,
|
||||
)
|
||||
except (ProfileNotFoundError, ProfileValidationError, ContractScopeError) as exc:
|
||||
_raise_422_for_profile_error(exc)
|
||||
|
||||
db.commit()
|
||||
db.refresh(contract)
|
||||
logger.info("Created energy contract id=%d name=%r kind=%s", contract.id, contract.name, contract.kind)
|
||||
return _contract_detail(db, contract)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /api/energy/contracts/{id}
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/contracts/{contract_id}", response_model=ContractDetailResponse)
|
||||
def get_energy_contract(
|
||||
contract_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> ContractDetailResponse:
|
||||
"""Return a single energy contract with its full version history.
|
||||
|
||||
Versions are ordered by ``effective_from`` ascending so the caller can
|
||||
easily inspect the pricing timeline.
|
||||
"""
|
||||
contract = _get_contract_or_404(db, contract_id)
|
||||
return _contract_detail(db, contract)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# PATCH /api/energy/contracts/{id}
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.patch("/contracts/{contract_id}", response_model=ContractDetailResponse)
|
||||
def patch_energy_contract(
|
||||
contract_id: int,
|
||||
body: ContractPatch,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> ContractDetailResponse:
|
||||
"""Partially update a contract: rename or change activation status.
|
||||
|
||||
- ``name``: updates the human-readable label.
|
||||
- ``active=true``: activates this contract (same-scope contracts are deactivated).
|
||||
- ``active=false``: deactivates this contract (no effect on others).
|
||||
|
||||
At most one contract may be active per scope; the service layer enforces
|
||||
scope-local mutual exclusion.
|
||||
"""
|
||||
contract = _get_contract_or_404(db, contract_id)
|
||||
was_active = contract.active
|
||||
|
||||
if body.name is not None:
|
||||
contract.name = body.name
|
||||
contract.updated_at = datetime.now(UTC)
|
||||
|
||||
if body.active is True:
|
||||
activate_contract(db, contract)
|
||||
elif body.active is False:
|
||||
deactivate_contract(db, contract)
|
||||
|
||||
db.commit()
|
||||
db.refresh(contract)
|
||||
if body.active is True and not was_active and contract.kind == "tibber":
|
||||
trigger_tibber_refresh()
|
||||
return _contract_detail(db, contract)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# POST /api/energy/contracts/{id}/versions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post(
|
||||
"/contracts/{contract_id}/versions",
|
||||
response_model=ContractDetailResponse,
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
)
|
||||
def add_contract_version(
|
||||
contract_id: int,
|
||||
body: VersionCreate,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> ContractDetailResponse:
|
||||
"""Add a new pricing version to an existing contract.
|
||||
|
||||
This is how price changes are recorded: the current open version is
|
||||
automatically closed (its ``effective_to`` is set to ``body.effective_from``)
|
||||
and a new version is created starting at ``body.effective_from``.
|
||||
|
||||
The ``values`` dict must conform to the contract's pricing profile.
|
||||
Non-conforming values return 422. If ``effective_from`` is not strictly
|
||||
after the previous version's ``effective_from``, 422 is returned without
|
||||
writing any rows.
|
||||
|
||||
Historical versions are never modified; this endpoint is append-only.
|
||||
"""
|
||||
contract = _get_contract_or_404(db, contract_id)
|
||||
effective_from = _localize_effective_from(body.effective_from)
|
||||
|
||||
try:
|
||||
add_version(
|
||||
db,
|
||||
contract,
|
||||
effective_from=effective_from,
|
||||
values=body.values,
|
||||
)
|
||||
except ContractVersionError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=str(exc),
|
||||
)
|
||||
except (ProfileNotFoundError, ProfileValidationError) as exc:
|
||||
_raise_422_for_profile_error(exc)
|
||||
|
||||
db.commit()
|
||||
db.refresh(contract)
|
||||
return _contract_detail(db, contract)
|
||||
@@ -0,0 +1,221 @@
|
||||
"""Expose API routes (M5-T12).
|
||||
|
||||
Three endpoints:
|
||||
GET /api/expose — Return catalog of exposable entities, per-key
|
||||
toggle state, and MQTT/Discovery connection status.
|
||||
PUT /api/expose — Set per-key toggle state (map key → bool);
|
||||
triggers a HA Discovery re-publish on success.
|
||||
POST /api/expose/republish — Manually trigger a full HA Discovery re-publish.
|
||||
|
||||
Auth model:
|
||||
- GET: session required (no CSRF — read-only).
|
||||
- PUT, POST: session + CSRF required.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.routes.api.deps import require_csrf, require_session
|
||||
from app.config import get_settings
|
||||
from app.dependencies import get_db
|
||||
from app.services.config_page import build_runtime_settings
|
||||
from app.integrations.expose import CatalogEntry, build_catalog
|
||||
from app.integrations.mqtt import mqtt_manager
|
||||
from app.models.expose import ExposedEntityToggle
|
||||
from app.schemas.expose import (
|
||||
CatalogEntrySchema,
|
||||
DeviceInfoSchema,
|
||||
ExposeResponse,
|
||||
ExposeUpdateRequest,
|
||||
ExposeUpdateResponse,
|
||||
ExposableEntitySchema,
|
||||
MqttStatusSchema,
|
||||
RepublishResponse,
|
||||
)
|
||||
from app.services.auth import AuthenticatedSession
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/api", tags=["api-expose"])
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_mqtt_status(db: Session) -> MqttStatusSchema:
|
||||
"""Build MQTT/Discovery status from live runtime state (DB-merged settings)."""
|
||||
settings = build_runtime_settings(db, get_settings())
|
||||
return MqttStatusSchema(
|
||||
mqtt_configured=mqtt_manager.is_configured(settings),
|
||||
mqtt_connected=mqtt_manager.is_connected,
|
||||
discovery_enabled=bool(settings.ha_discovery_enabled),
|
||||
)
|
||||
|
||||
|
||||
def _entry_to_schema(entry: CatalogEntry) -> CatalogEntrySchema:
|
||||
"""Convert a CatalogEntry to its Pydantic schema form.
|
||||
|
||||
Excludes ``value_getter`` because it is a non-serialisable callable.
|
||||
"""
|
||||
entity = entry.entity
|
||||
return CatalogEntrySchema(
|
||||
entity=ExposableEntitySchema(
|
||||
key=entity.key,
|
||||
component=entity.component,
|
||||
device=DeviceInfoSchema(
|
||||
identifiers=list(entity.device.identifiers),
|
||||
name=entity.device.name,
|
||||
),
|
||||
device_class=entity.device_class,
|
||||
unit=entity.unit,
|
||||
name=entity.name,
|
||||
state_class=entity.state_class,
|
||||
),
|
||||
enabled=entry.enabled,
|
||||
)
|
||||
|
||||
|
||||
def _build_response_data(
|
||||
session: Session,
|
||||
) -> tuple[list[CatalogEntrySchema], MqttStatusSchema]:
|
||||
"""Return (catalog_entries, mqtt_status) for building GET / PUT responses."""
|
||||
catalog = build_catalog(session)
|
||||
catalog_schema = [_entry_to_schema(e) for e in catalog]
|
||||
mqtt_status = _get_mqtt_status(session)
|
||||
return catalog_schema, mqtt_status
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /api/expose
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/expose", response_model=ExposeResponse)
|
||||
def get_expose(
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> ExposeResponse:
|
||||
"""Return the full exposable-entity catalog with toggle states and MQTT status.
|
||||
|
||||
The catalog is computed dynamically from registered providers (e.g. the
|
||||
Modbus provider enumerates all enabled devices and their metric entities).
|
||||
Toggle states come from the ``exposed_entity_toggle`` table; entities with
|
||||
no row default to ``enabled=False``.
|
||||
"""
|
||||
catalog_schema, mqtt_status = _build_response_data(db)
|
||||
return ExposeResponse(catalog=catalog_schema, mqtt_status=mqtt_status)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# PUT /api/expose
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.put("/expose", response_model=ExposeUpdateResponse)
|
||||
def put_expose(
|
||||
body: ExposeUpdateRequest,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> ExposeUpdateResponse:
|
||||
"""Set per-entity toggle state.
|
||||
|
||||
Accepts a map of ``{key: bool}`` and upserts rows in the
|
||||
``exposed_entity_toggle`` table. Only keys present in ``body.toggles``
|
||||
are touched; other entities' toggles are left unchanged.
|
||||
|
||||
After writing the toggles, triggers a HA Discovery re-publish so any
|
||||
changes (enabled ↔ disabled) are reflected in Home Assistant immediately.
|
||||
"""
|
||||
now = datetime.now(UTC)
|
||||
|
||||
for key, enabled in body.toggles.items():
|
||||
existing = (
|
||||
db.query(ExposedEntityToggle)
|
||||
.filter(ExposedEntityToggle.key == key)
|
||||
.first()
|
||||
)
|
||||
if existing is not None:
|
||||
existing.enabled = enabled
|
||||
existing.updated_at = now
|
||||
else:
|
||||
db.add(
|
||||
ExposedEntityToggle(
|
||||
key=key,
|
||||
enabled=enabled,
|
||||
updated_at=now,
|
||||
)
|
||||
)
|
||||
|
||||
db.commit()
|
||||
|
||||
# Trigger discovery re-publish after toggle change.
|
||||
_trigger_republish(db)
|
||||
|
||||
catalog_schema, mqtt_status = _build_response_data(db)
|
||||
return ExposeUpdateResponse(catalog=catalog_schema, mqtt_status=mqtt_status)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# POST /api/expose/republish
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/expose/republish", response_model=RepublishResponse)
|
||||
def post_expose_republish(
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> RepublishResponse:
|
||||
"""Manually trigger a full HA Discovery re-publish.
|
||||
|
||||
Calls ``publish_discovery(session)`` from the HA discovery service (M5-T11).
|
||||
Returns a status indicating whether the publish was attempted (or skipped
|
||||
because MQTT / discovery is not enabled / connected).
|
||||
"""
|
||||
settings = build_runtime_settings(db, get_settings())
|
||||
if not (settings.mqtt_enabled and settings.ha_discovery_enabled and mqtt_manager.is_connected):
|
||||
return RepublishResponse(
|
||||
ok=False,
|
||||
message="MQTT / HA Discovery not enabled or broker not connected.",
|
||||
)
|
||||
|
||||
try:
|
||||
_trigger_republish(db)
|
||||
return RepublishResponse(
|
||||
ok=True,
|
||||
message="HA Discovery re-published successfully.",
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.exception("post_expose_republish: unexpected error during publish")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Discovery re-publish failed.",
|
||||
) from exc
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Internal: trigger discovery re-publish
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _trigger_republish(session: Session) -> None:
|
||||
"""Call publish_discovery from the HA discovery service.
|
||||
|
||||
No-op if MQTT / discovery is not enabled or broker is not connected
|
||||
(publish_discovery guards internally). All errors are swallowed to avoid
|
||||
breaking the API response.
|
||||
"""
|
||||
try:
|
||||
from app.services.ha_discovery import publish_discovery
|
||||
|
||||
publish_discovery(session)
|
||||
except Exception:
|
||||
logger.exception("_trigger_republish: publish_discovery raised an error")
|
||||
@@ -0,0 +1,150 @@
|
||||
"""Authenticated API for the thermal 15-minute cost ledger."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from decimal import Decimal
|
||||
from typing import Literal
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.routes.api.deps import require_csrf, require_session
|
||||
from app.dependencies import get_db
|
||||
from app.models.energy import MeterCostPeriod
|
||||
from app.schemas.meter_cost import (
|
||||
MeterCostPeriodSchema,
|
||||
MeterCostRecomputeResponse,
|
||||
MeterCostsResponse,
|
||||
ThermalCostSummaryResponse,
|
||||
)
|
||||
from app.services.auth import AuthenticatedSession
|
||||
from app.services.meter_cost import recompute_range, summarize
|
||||
from app.services.timezone import local_midnight_utc, local_now
|
||||
|
||||
router = APIRouter(prefix="/api/energy/meter-costs", tags=["api-energy"])
|
||||
|
||||
_LIMIT_MAX = 5000
|
||||
_RECOMPUTE_MAX_DAYS = 31
|
||||
_QUARTER = timedelta(minutes=15)
|
||||
|
||||
|
||||
def _utc(value: datetime) -> datetime:
|
||||
return value.replace(tzinfo=UTC) if value.tzinfo is None else value.astimezone(UTC)
|
||||
|
||||
|
||||
def _decimal_strings(value: object) -> object:
|
||||
if isinstance(value, dict):
|
||||
return {str(key): _decimal_strings(item) for key, item in value.items()}
|
||||
if isinstance(value, Decimal):
|
||||
return format(value, "f")
|
||||
return str(value) if isinstance(value, (int, float)) else value
|
||||
|
||||
|
||||
def _row_schema(row: MeterCostPeriod) -> MeterCostPeriodSchema:
|
||||
return MeterCostPeriodSchema(
|
||||
commodity=row.commodity,
|
||||
period_start=_utc(row.period_start), period_end=_utc(row.period_end),
|
||||
meter_id=row.meter_id, source_binding_id=row.source_binding_id,
|
||||
contract_version_id=row.contract_version_id, quantity=format(row.quantity, "f"),
|
||||
cost=format(row.cost, "f"), currency=row.currency,
|
||||
cost_breakdown=_decimal_strings(row.cost_breakdown),
|
||||
pricing_snapshot=_decimal_strings(row.pricing_snapshot), quality=row.quality,
|
||||
degraded=row.degraded, degraded_reason=row.degraded_reason,
|
||||
)
|
||||
|
||||
|
||||
@router.get("", response_model=MeterCostsResponse)
|
||||
def get_meter_costs(
|
||||
scope: Literal["thermal"] = Query("thermal"),
|
||||
commodity: Literal["heating", "hot_water"] | None = None,
|
||||
start: datetime | None = None,
|
||||
end: datetime | None = None,
|
||||
limit: int = Query(500, ge=1, le=_LIMIT_MAX),
|
||||
offset: int = Query(0, ge=0),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> MeterCostsResponse:
|
||||
"""List thermal rows in a half-open time window with stable pagination."""
|
||||
del scope
|
||||
if start is not None and end is not None and _utc(end) <= _utc(start):
|
||||
raise HTTPException(status.HTTP_422_UNPROCESSABLE_ENTITY, "'end' must be after 'start'.")
|
||||
clauses = []
|
||||
if commodity is not None:
|
||||
clauses.append(MeterCostPeriod.commodity == commodity)
|
||||
if start is not None:
|
||||
clauses.append(MeterCostPeriod.period_start >= _utc(start))
|
||||
if end is not None:
|
||||
clauses.append(MeterCostPeriod.period_start < _utc(end))
|
||||
total = db.scalar(select(func.count()).select_from(MeterCostPeriod).where(*clauses)) or 0
|
||||
rows = db.execute(
|
||||
select(MeterCostPeriod).where(*clauses).order_by(MeterCostPeriod.period_start, MeterCostPeriod.id)
|
||||
.offset(offset).limit(limit)
|
||||
).scalars().all()
|
||||
return MeterCostsResponse(items=[_row_schema(row) for row in rows], total=total)
|
||||
|
||||
|
||||
@router.get("/summary", response_model=ThermalCostSummaryResponse)
|
||||
def get_meter_cost_summary(
|
||||
scope: Literal["thermal"] = Query("thermal"),
|
||||
start: datetime | None = None,
|
||||
end: datetime | None = None,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> ThermalCostSummaryResponse:
|
||||
"""Summarize thermal variable and once-per-contract daily fixed costs."""
|
||||
del scope
|
||||
if start is None or end is None:
|
||||
today = local_now().date()
|
||||
start = start or local_midnight_utc(today)
|
||||
end = end or local_midnight_utc(today + timedelta(days=1))
|
||||
start, end = _utc(start), _utc(end)
|
||||
if end <= start:
|
||||
raise HTTPException(status.HTTP_422_UNPROCESSABLE_ENTITY, "'end' must be after 'start'.")
|
||||
result = summarize(db, start, end)
|
||||
breakdown = result["breakdown"]
|
||||
fixed_breakdown = result["fixed_breakdown"]
|
||||
fixed = result["fixed_cost"]
|
||||
return ThermalCostSummaryResponse(
|
||||
currency=result["currency"], heating=format(breakdown["heating"], "f"),
|
||||
hot_water_heating=format(breakdown["hot_water_heating"], "f"),
|
||||
hot_water=format(breakdown["hot_water"], "f"), hot_water_tax=format(breakdown["hot_water_tax"], "f"),
|
||||
variable_subtotal=format(result["variable_cost"], "f"),
|
||||
fixed_breakdown={key: format(value, "f") for key, value in fixed_breakdown.items()},
|
||||
fixed_subtotal=format(fixed, "f"),
|
||||
all_in=format(result["total_cost"], "f"), period_count=result["period_count"],
|
||||
degraded_count=result["degraded_count"],
|
||||
)
|
||||
|
||||
|
||||
@router.post("/recompute", response_model=MeterCostRecomputeResponse)
|
||||
def post_meter_cost_recompute(
|
||||
scope: Literal["thermal"] = Query("thermal"),
|
||||
start: datetime = Query(...), end: datetime = Query(...),
|
||||
db: Session = Depends(get_db), _auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> MeterCostRecomputeResponse:
|
||||
"""Atomically overwrite closed, UTC-quarter thermal rows in a bounded window."""
|
||||
del scope
|
||||
start, end = _utc(start), _utc(end)
|
||||
if end <= start or end - start > timedelta(days=_RECOMPUTE_MAX_DAYS):
|
||||
raise HTTPException(status.HTTP_422_UNPROCESSABLE_ENTITY, "invalid or overlarge recompute window")
|
||||
if start.minute % 15 or start.second or start.microsecond or end.minute % 15 or end.second or end.microsecond:
|
||||
raise HTTPException(status.HTTP_422_UNPROCESSABLE_ENTITY, "start and end must align to UTC quarters")
|
||||
if end > datetime.now(UTC):
|
||||
raise HTTPException(status.HTTP_422_UNPROCESSABLE_ENTITY, "recompute window must be closed")
|
||||
try:
|
||||
processed = recompute_range(db, start, end, commit=False)
|
||||
# Ensure pending upserts participate in this transaction before the
|
||||
# counts are read; a flush/query failure must still roll everything back.
|
||||
db.flush()
|
||||
rows = db.execute(select(MeterCostPeriod.degraded).where(
|
||||
MeterCostPeriod.period_start >= start, MeterCostPeriod.period_start < end
|
||||
)).scalars().all()
|
||||
degraded = sum(bool(value) for value in rows)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
return MeterCostRecomputeResponse(processed=processed, normal=len(rows) - degraded, degraded=degraded)
|
||||
@@ -0,0 +1,414 @@
|
||||
"""Authenticated HTTP contract for meter sources, channels, and bindings."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.routes.api.deps import require_csrf, require_session
|
||||
from app.config import get_settings
|
||||
from app.dependencies import get_db
|
||||
from app.integrations.meter_sources import SourceProfileError, list_source_profiles, sanitize_source_config
|
||||
from app.models.energy import DsmrReading, Meter
|
||||
from app.models.meter_source import MeterSource, MeterSourceBinding, MeterSourceChannel, WarmteLinkReading
|
||||
from app.schemas.meter_source import (
|
||||
BindingCreate, BindingListResponse, BindingPatch, BindingResponse, ChannelBindingSummaryResponse,
|
||||
BindingTransferRequest, BindingTransferResponse,
|
||||
ChannelReadingResponse,
|
||||
ChannelReadingsResponse, CommoditiesResponse, CommodityResponse, DiscoverResponse,
|
||||
DiscoverChannelResponse,
|
||||
MeterSourceChannelListResponse, MeterSourceChannelResponse, MeterSourceCreate,
|
||||
MeterSourceListResponse, MeterSourcePatch, MeterSourceResponse, SourceConfigFieldResponse,
|
||||
SourceProfileResponse, SourceProfilesResponse,
|
||||
)
|
||||
from app.services.auth import AuthenticatedSession
|
||||
from app.services.config_page import build_runtime_settings
|
||||
from app.services.dsmr_ingest import apply_dsmr_subscription
|
||||
from app.services.meter_sources import (
|
||||
BindingNotFoundError, ChannelNotFoundError, MeterNotFoundError,
|
||||
MeterSourceError, SourceDeleteRestrictedError, SourceNotFoundError, create_binding,
|
||||
create_source, delete_source, list_bindings, list_sources, transfer_binding, update_binding, update_source,
|
||||
)
|
||||
from app.services.energy_cost import recompute_range as electricity_recompute_range
|
||||
from app.services import timezone as _tz_mod
|
||||
from app.services.warmtelink_worker import warmtelink_worker_manager
|
||||
|
||||
router = APIRouter(prefix="/api/energy", tags=["api-energy-meter-sources"])
|
||||
|
||||
|
||||
def _reconcile_runtimes_after_commit(db: Session) -> None:
|
||||
"""Best-effort runtime convergence after a durable source CRUD commit."""
|
||||
try:
|
||||
warmtelink_worker_manager.reconcile()
|
||||
except Exception:
|
||||
# The manager records individual source failures itself. Do not turn a
|
||||
# successful durable create/update/delete into a misleading HTTP 500.
|
||||
pass
|
||||
try:
|
||||
apply_dsmr_subscription(build_runtime_settings(db, get_settings()))
|
||||
except Exception:
|
||||
# DSMR owns independent source clients. Its failure must neither undo
|
||||
# durable CRUD nor prevent the WarmteLink manager from converging.
|
||||
pass
|
||||
finally:
|
||||
# DSMR health callbacks use short independent sessions. Make a CRUD
|
||||
# response observe any durable status change they just committed.
|
||||
db.expire_all()
|
||||
|
||||
|
||||
def _as_utc(value: datetime) -> datetime:
|
||||
if value.tzinfo is None:
|
||||
return value.replace(tzinfo=_tz_mod.local_tz()).astimezone(UTC)
|
||||
return value.astimezone(UTC)
|
||||
|
||||
|
||||
def _source_or_404(db: Session, uuid: str) -> MeterSource:
|
||||
source = db.execute(select(MeterSource).where(MeterSource.uuid == uuid)).scalar_one_or_none()
|
||||
if source is None:
|
||||
raise HTTPException(status_code=404, detail="Meter source not found.")
|
||||
return source
|
||||
|
||||
|
||||
def _channel_or_404(db: Session, source: MeterSource, uuid: str) -> MeterSourceChannel:
|
||||
channel = db.execute(
|
||||
select(MeterSourceChannel).where(
|
||||
MeterSourceChannel.uuid == uuid, MeterSourceChannel.source_id == source.id
|
||||
)
|
||||
).scalar_one_or_none()
|
||||
if channel is None:
|
||||
raise HTTPException(status_code=404, detail="Meter source channel not found.")
|
||||
return channel
|
||||
|
||||
|
||||
def _source_response(source: MeterSource) -> MeterSourceResponse:
|
||||
return MeterSourceResponse(
|
||||
uuid=source.uuid, name=source.name, kind=source.kind, enabled=source.enabled,
|
||||
config=sanitize_source_config(source.kind, source.config), status=source.status,
|
||||
last_seen_at=source.last_seen_at, last_error=source.last_error,
|
||||
created_at=source.created_at, updated_at=source.updated_at,
|
||||
)
|
||||
|
||||
|
||||
def binding_response(binding: MeterSourceBinding) -> BindingResponse:
|
||||
return BindingResponse(
|
||||
uuid=binding.uuid, meter_id=binding.meter_id, source_channel_uuid=binding.channel.uuid,
|
||||
source_uuid=binding.channel.source.uuid, started_at=binding.started_at, ended_at=binding.ended_at,
|
||||
created_at=binding.created_at, updated_at=binding.updated_at,
|
||||
)
|
||||
|
||||
|
||||
def _binding_error(exc: MeterSourceError) -> HTTPException:
|
||||
if isinstance(exc, (SourceNotFoundError, ChannelNotFoundError, MeterNotFoundError, BindingNotFoundError)):
|
||||
return HTTPException(status_code=404, detail=str(exc))
|
||||
return HTTPException(status_code=422, detail=str(exc))
|
||||
|
||||
|
||||
def _recompute_binding_commodity(db: Session, commodity: str, start: datetime) -> None:
|
||||
end = datetime.now(UTC)
|
||||
if start >= end:
|
||||
return
|
||||
if commodity == "electricity":
|
||||
electricity_recompute_range(db, start, end, commit=False, strict=True)
|
||||
else:
|
||||
from app.services.meter_cost import recompute_range
|
||||
recompute_range(db, start, end, commit=False)
|
||||
|
||||
|
||||
def _republish_after_commit(db: Session) -> None:
|
||||
try:
|
||||
from app.services.ha_discovery import publish_discovery
|
||||
publish_discovery(db)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@router.get("/source-profiles", response_model=SourceProfilesResponse)
|
||||
def source_profiles(_auth: AuthenticatedSession = Depends(require_session)) -> SourceProfilesResponse:
|
||||
"""Return profile metadata; default secrets are never populated with stored values."""
|
||||
profiles = []
|
||||
for profile in list_source_profiles():
|
||||
fields = [
|
||||
SourceConfigFieldResponse(name=f.name, value_type=f.value_type.__name__, default=f.default,
|
||||
required=f.required, secret=f.secret)
|
||||
for f in profile.fields
|
||||
]
|
||||
profiles.append(SourceProfileResponse(
|
||||
kind=profile.kind, fields=fields,
|
||||
defaults={f.name: f.default for f in profile.fields if not f.required},
|
||||
capabilities=sorted(profile.capabilities), allowed_units=sorted(profile.allowed_units),
|
||||
))
|
||||
return SourceProfilesResponse(items=profiles)
|
||||
|
||||
|
||||
@router.get("/commodities", response_model=CommoditiesResponse)
|
||||
def commodities(_auth: AuthenticatedSession = Depends(require_session)) -> CommoditiesResponse:
|
||||
return CommoditiesResponse(items=[
|
||||
CommodityResponse(key="electricity", unit="kWh", capabilities=["meter", "binding", "cost"]),
|
||||
CommodityResponse(key="heating", unit="GJ", capabilities=["meter", "binding"]),
|
||||
CommodityResponse(key="hot_water", unit="m³", capabilities=["meter", "binding"]),
|
||||
])
|
||||
|
||||
|
||||
@router.get("/sources", response_model=MeterSourceListResponse)
|
||||
def get_sources(db: Session = Depends(get_db), _auth: AuthenticatedSession = Depends(require_session)) -> MeterSourceListResponse:
|
||||
items = [_source_response(source) for source in list_sources(db)]
|
||||
return MeterSourceListResponse(items=items, total=len(items))
|
||||
|
||||
|
||||
@router.post("/sources", response_model=MeterSourceResponse, status_code=status.HTTP_201_CREATED)
|
||||
def post_source(body: MeterSourceCreate, db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session), _csrf: None = Depends(require_csrf)) -> MeterSourceResponse:
|
||||
try:
|
||||
source = create_source(db, name=body.name, kind=body.kind, config=body.config, enabled=body.enabled)
|
||||
db.commit()
|
||||
_reconcile_runtimes_after_commit(db)
|
||||
return _source_response(_source_or_404(db, source.uuid))
|
||||
except (SourceProfileError, MeterSourceError) as exc:
|
||||
db.rollback()
|
||||
raise HTTPException(status_code=422, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.get("/sources/{source_uuid}", response_model=MeterSourceResponse)
|
||||
def get_source_detail(source_uuid: str, db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session)) -> MeterSourceResponse:
|
||||
return _source_response(_source_or_404(db, source_uuid))
|
||||
|
||||
|
||||
@router.patch("/sources/{source_uuid}", response_model=MeterSourceResponse)
|
||||
def patch_source(source_uuid: str, body: MeterSourcePatch, db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session), _csrf: None = Depends(require_csrf)) -> MeterSourceResponse:
|
||||
source = _source_or_404(db, source_uuid)
|
||||
try:
|
||||
updated = update_source(db, source.id, name=body.name, enabled=body.enabled, config_patch=body.config)
|
||||
db.commit()
|
||||
_reconcile_runtimes_after_commit(db)
|
||||
return _source_response(_source_or_404(db, updated.uuid))
|
||||
except (SourceProfileError, MeterSourceError) as exc:
|
||||
db.rollback()
|
||||
raise _binding_error(exc) from exc
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/sources/{source_uuid}", status_code=status.HTTP_204_NO_CONTENT, response_model=None
|
||||
)
|
||||
def remove_source(source_uuid: str, db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session), _csrf: None = Depends(require_csrf)) -> None:
|
||||
source = _source_or_404(db, source_uuid)
|
||||
# DSMR readings are not a relationship on MeterSource to avoid loading a large history.
|
||||
if db.execute(select(DsmrReading.id).where(DsmrReading.meter_source_id == source.id).limit(1)).scalar() is not None:
|
||||
raise HTTPException(status_code=409, detail="Meter source has dependent readings.")
|
||||
try:
|
||||
delete_source(db, source.id)
|
||||
db.commit()
|
||||
_reconcile_runtimes_after_commit(db)
|
||||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||
except SourceDeleteRestrictedError as exc:
|
||||
db.rollback()
|
||||
raise HTTPException(status_code=409, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/sources/{source_uuid}/discover", response_model=DiscoverResponse)
|
||||
def discover_source(source_uuid: str, db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session), _csrf: None = Depends(require_csrf)) -> DiscoverResponse:
|
||||
source = _source_or_404(db, source_uuid)
|
||||
if source.kind == "warmtelink_serial":
|
||||
if not source.enabled:
|
||||
return DiscoverResponse(
|
||||
requested=False, supported=True, status="error",
|
||||
detail="The WarmteLink source is disabled.", channels=_discover_channels(db, source),
|
||||
)
|
||||
# This merely schedules lifecycle convergence. It never opens a serial
|
||||
# descriptor or waits for a frame in the request thread; the one managed
|
||||
# worker remains the sole owner of serial I/O and can keep reconnecting.
|
||||
request = warmtelink_worker_manager.request_discovery(source.id)
|
||||
if request.completed.is_set():
|
||||
# A worker may have accepted a frame during the bounded wait.
|
||||
# Refresh only durable accepted metadata, never candidates/raw data.
|
||||
db.expire_all()
|
||||
source = _source_or_404(db, source_uuid)
|
||||
return DiscoverResponse(
|
||||
requested=request.status != "error", supported=True, status=request.status,
|
||||
request_id=request.request_id or None, detail=request.detail,
|
||||
channels=_discover_channels(db, source),
|
||||
)
|
||||
return DiscoverResponse(requested=False, supported=True, status="managed_by_runtime",
|
||||
detail="This source is discovered by its runtime subscription; no connection was opened.")
|
||||
|
||||
|
||||
@router.get("/sources/{source_uuid}/channels", response_model=MeterSourceChannelListResponse)
|
||||
def source_channels(source_uuid: str, db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session)) -> MeterSourceChannelListResponse:
|
||||
source = _source_or_404(db, source_uuid)
|
||||
channels = db.execute(select(MeterSourceChannel).where(MeterSourceChannel.source_id == source.id)).scalars().all()
|
||||
items = []
|
||||
for channel in channels:
|
||||
bindings = list_bindings(db, channel_id=channel.id)
|
||||
meter_ids = [binding.meter_id for binding in bindings]
|
||||
items.append(MeterSourceChannelResponse(
|
||||
uuid=channel.uuid, label=channel.label, suggested_commodity=channel.suggested_commodity,
|
||||
unit=channel.unit, device_type=channel.device_type, latest_value=channel.latest_value,
|
||||
latest_at=channel.latest_at, latest_quality=channel.latest_quality, binding_count=len(bindings),
|
||||
bound_meter_ids=meter_ids,
|
||||
binding_summary=ChannelBindingSummaryResponse(count=len(bindings), meter_ids=meter_ids),
|
||||
))
|
||||
return MeterSourceChannelListResponse(items=items, total=len(items), source_status=source.status)
|
||||
|
||||
|
||||
@router.get("/sources/{source_uuid}/channels/{channel_uuid}/readings", response_model=ChannelReadingsResponse)
|
||||
def channel_readings(source_uuid: str, channel_uuid: str, limit: int = Query(default=100, ge=1, le=1000),
|
||||
from_: datetime | None = Query(default=None, alias="from"),
|
||||
to: datetime | None = Query(default=None), db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session)) -> ChannelReadingsResponse:
|
||||
source = _source_or_404(db, source_uuid)
|
||||
channel = _channel_or_404(db, source, channel_uuid)
|
||||
if from_ is not None and to is not None and _as_utc(from_) >= _as_utc(to):
|
||||
raise HTTPException(status_code=422, detail="'from' must be earlier than 'to'.")
|
||||
if source.kind == "warmtelink_serial":
|
||||
statement = select(WarmteLinkReading).where(WarmteLinkReading.channel_id == channel.id)
|
||||
model = WarmteLinkReading
|
||||
else:
|
||||
# DSMR remains a source-level protocol history. Its channel is the
|
||||
# public electricity identity, while payload/telegram diagnostics stay
|
||||
# private to ingestion and the legacy latest endpoint.
|
||||
statement = select(DsmrReading).where(DsmrReading.meter_source_id == source.id)
|
||||
model = DsmrReading
|
||||
if from_ is not None:
|
||||
statement = statement.where(model.recorded_at >= _as_utc(from_))
|
||||
if to is not None:
|
||||
statement = statement.where(model.recorded_at < _as_utc(to))
|
||||
rows = list(db.execute(statement.order_by(model.recorded_at.asc()).limit(limit)).scalars())
|
||||
return ChannelReadingsResponse(
|
||||
items=[ChannelReadingResponse(
|
||||
recorded_at=row.recorded_at,
|
||||
value=getattr(row, "value", None), quality=getattr(row, "quality", None),
|
||||
) for row in rows],
|
||||
total=len(rows),
|
||||
)
|
||||
|
||||
|
||||
def _discover_channels(db: Session, source: MeterSource) -> list[DiscoverChannelResponse]:
|
||||
"""Return only public, accepted channel metadata for discover responses."""
|
||||
return [
|
||||
DiscoverChannelResponse(
|
||||
uuid=channel.uuid, label=channel.label, unit=channel.unit,
|
||||
latest_value=channel.latest_value, latest_at=channel.latest_at,
|
||||
latest_quality=channel.latest_quality,
|
||||
)
|
||||
for channel in db.execute(
|
||||
select(MeterSourceChannel).where(MeterSourceChannel.source_id == source.id)
|
||||
).scalars()
|
||||
]
|
||||
|
||||
|
||||
@router.get("/meters/{meter_id}/bindings", response_model=BindingListResponse)
|
||||
def meter_bindings(meter_id: int, db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session)) -> BindingListResponse:
|
||||
if db.get(Meter, meter_id) is None:
|
||||
raise HTTPException(status_code=404, detail="Meter not found.")
|
||||
items = [binding_response(binding) for binding in list_bindings(db, meter_id=meter_id)]
|
||||
return BindingListResponse(items=items, total=len(items))
|
||||
|
||||
|
||||
@router.post("/meters/{meter_id}/bindings", response_model=BindingResponse, status_code=status.HTTP_201_CREATED)
|
||||
def post_meter_binding(meter_id: int, body: BindingCreate, db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session), _csrf: None = Depends(require_csrf)) -> BindingResponse:
|
||||
channel = db.execute(select(MeterSourceChannel).where(MeterSourceChannel.uuid == body.source_channel_uuid)).scalar_one_or_none()
|
||||
if channel is None:
|
||||
raise HTTPException(status_code=404, detail="Meter source channel not found.")
|
||||
try:
|
||||
binding = create_binding(db, meter_id=meter_id, channel_id=channel.id, started_at=_as_utc(body.started_at),
|
||||
ended_at=_as_utc(body.ended_at) if body.ended_at else None)
|
||||
meter = db.get(Meter, meter_id)
|
||||
db.flush()
|
||||
_recompute_binding_commodity(db, meter.commodity, _as_utc(body.started_at))
|
||||
db.commit()
|
||||
db.refresh(binding)
|
||||
_republish_after_commit(db)
|
||||
return binding_response(binding)
|
||||
except MeterSourceError as exc:
|
||||
db.rollback()
|
||||
raise _binding_error(exc) from exc
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
|
||||
@router.patch("/bindings/{binding_uuid}", response_model=BindingResponse)
|
||||
def patch_binding(binding_uuid: str, body: BindingPatch, db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session), _csrf: None = Depends(require_csrf)) -> BindingResponse:
|
||||
binding = db.execute(select(MeterSourceBinding).where(MeterSourceBinding.uuid == binding_uuid)).scalar_one_or_none()
|
||||
if binding is None:
|
||||
raise HTTPException(status_code=404, detail="Meter source binding not found.")
|
||||
try:
|
||||
# ``ended_at`` has three meaningful states in the service layer: omitted
|
||||
# keeps the existing boundary, null reopens the interval, and a datetime
|
||||
# changes the exclusive end. Do not collapse omitted into null here.
|
||||
changes: dict[str, datetime | None] = {}
|
||||
if "started_at" in body.model_fields_set:
|
||||
changes["started_at"] = _as_utc(body.started_at) if body.started_at is not None else None
|
||||
if "ended_at" in body.model_fields_set:
|
||||
changes["ended_at"] = _as_utc(body.ended_at) if body.ended_at is not None else None
|
||||
old_started_at = _as_utc(binding.started_at)
|
||||
old_ended_at = _as_utc(binding.ended_at) if binding.ended_at is not None else None
|
||||
updated = update_binding(db, binding.id, **changes)
|
||||
meter = db.get(Meter, updated.meter_id)
|
||||
if "started_at" in changes:
|
||||
earliest = min(old_started_at, _as_utc(updated.started_at))
|
||||
elif "ended_at" in changes:
|
||||
new_ended_at = _as_utc(updated.ended_at) if updated.ended_at is not None else None
|
||||
changed_ends = [value for value in (old_ended_at, new_ended_at) if value is not None]
|
||||
earliest = min(changed_ends) if changed_ends else old_started_at
|
||||
else:
|
||||
earliest = old_started_at
|
||||
db.flush()
|
||||
_recompute_binding_commodity(db, meter.commodity, earliest)
|
||||
db.commit()
|
||||
db.refresh(updated)
|
||||
_republish_after_commit(db)
|
||||
return binding_response(updated)
|
||||
except MeterSourceError as exc:
|
||||
db.rollback()
|
||||
raise _binding_error(exc) from exc
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
|
||||
@router.post("/meters/{meter_id}/bindings/transfer", response_model=BindingTransferResponse)
|
||||
def post_binding_transfer(
|
||||
meter_id: int, body: BindingTransferRequest, db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session), _csrf: None = Depends(require_csrf),
|
||||
) -> BindingTransferResponse:
|
||||
source = db.execute(select(MeterSourceBinding).where(
|
||||
MeterSourceBinding.uuid == body.from_binding_uuid
|
||||
)).scalar_one_or_none()
|
||||
channel = db.execute(select(MeterSourceChannel).where(
|
||||
MeterSourceChannel.uuid == body.to_source_channel_uuid
|
||||
)).scalar_one_or_none()
|
||||
if source is None or channel is None:
|
||||
raise HTTPException(status_code=404, detail="Meter source binding or channel not found.")
|
||||
effective_at = _as_utc(body.effective_at)
|
||||
try:
|
||||
closed, created = transfer_binding(db, target_meter_id=meter_id, from_binding_id=source.id,
|
||||
to_channel_id=channel.id, effective_at=effective_at)
|
||||
meter = db.get(Meter, meter_id)
|
||||
if closed.meter_id == meter.id:
|
||||
earliest = effective_at
|
||||
else:
|
||||
earliest = min(_as_utc(closed.ended_at), effective_at)
|
||||
db.flush()
|
||||
_recompute_binding_commodity(db, meter.commodity, earliest)
|
||||
db.commit()
|
||||
db.refresh(closed)
|
||||
db.refresh(created)
|
||||
except MeterSourceError as exc:
|
||||
db.rollback()
|
||||
raise _binding_error(exc) from exc
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
_republish_after_commit(db)
|
||||
return BindingTransferResponse(closed_binding=binding_response(closed), created_binding=binding_response(created))
|
||||
@@ -0,0 +1,449 @@
|
||||
"""Meter CRUD, swap declaration, and retroactive recompute API (M7-T05).
|
||||
|
||||
All endpoints are under /api/energy/meters, require an authenticated session,
|
||||
and write endpoints (POST/PATCH) additionally require a non-empty X-CSRF-Token
|
||||
header.
|
||||
|
||||
Route semantics
|
||||
---------------
|
||||
GET /api/energy/meters — list all meter epochs (ascending started_at)
|
||||
POST /api/energy/meters — declare a meter swap / initial meter epoch
|
||||
PATCH /api/energy/meters/{id} — edit label / note, or correct started_at (retroactive)
|
||||
|
||||
Retroactive recompute
|
||||
---------------------
|
||||
Whenever a write operation changes a meter's ``started_at`` (new declaration
|
||||
or PATCH correction), the affected billing window is re-judged via
|
||||
``recompute_range``:
|
||||
|
||||
- **POST** (new meter, possibly retroactive):
|
||||
window = [new_meter.started_at, now)
|
||||
Rationale: the new meter's ``started_at`` closes the previous meter at that
|
||||
point; all periods from that boundary forward may have a different meter
|
||||
attribution. Using ``now`` as the upper bound is safe because
|
||||
``recompute_range`` only processes closed quarters and the operation is
|
||||
idempotent.
|
||||
|
||||
- **PATCH started_at** (retroactive correction):
|
||||
window = [min(old_started_at, new_started_at), now)
|
||||
Rationale: shifting the boundary in either direction affects all periods
|
||||
between the old and new boundary (and potentially beyond if re-attribution
|
||||
cascades). Using the minimum of the two timestamps guarantees the entire
|
||||
affected range is covered; using ``now`` as the upper bound is safe and
|
||||
idempotent.
|
||||
|
||||
``started_at`` localisation (Principle A, FU10 convention)
|
||||
----------------------------------------------------------
|
||||
If the client sends a timezone-naive ``started_at`` value, it is interpreted as
|
||||
the **server's local wall-clock time** and converted to UTC before storage.
|
||||
Timezone-aware values are converted to UTC as-is. This is identical to the
|
||||
``_localize_effective_from`` convention used in ``energy_contracts.py``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from sqlalchemy.orm import Session
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.api.routes.api.deps import require_csrf, require_session
|
||||
from app.dependencies import get_db
|
||||
from app.models.energy import Meter
|
||||
from app.models.meter_source import MeterSourceChannel
|
||||
from app.schemas.meter import (
|
||||
MeterCloseRequest,
|
||||
MeterDeclareRequest,
|
||||
MeterBindingSummary,
|
||||
MeterListResponse,
|
||||
MeterPatchRequest,
|
||||
MeterResponse,
|
||||
)
|
||||
from app.services.meter_sources import (
|
||||
ChannelNotFoundError,
|
||||
MeterSourceError,
|
||||
create_binding,
|
||||
create_binding_for_meter_swap,
|
||||
close_open_bindings_for_meter,
|
||||
)
|
||||
from app.services import timezone as _tz_mod
|
||||
from app.services.auth import AuthenticatedSession
|
||||
from app.services.energy_cost import recompute_range
|
||||
from app.services.meters import (
|
||||
MeterIntervalError,
|
||||
MeterOverlapError,
|
||||
declare_meter,
|
||||
close_meter,
|
||||
list_meters,
|
||||
update_meter,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/api/energy", tags=["api-energy-meters"])
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _trigger_discovery_republish(session: Session) -> None:
|
||||
"""Call publish_discovery after a meter write operation (best-effort).
|
||||
|
||||
No-op if MQTT / discovery is not enabled or the broker is not connected
|
||||
(publish_discovery guards internally). All errors are swallowed so that a
|
||||
discovery failure never breaks the API response.
|
||||
|
||||
Must be called **after** db.commit() so that publish_discovery sees the
|
||||
final committed state of the meter table when it rebuilds the catalog.
|
||||
"""
|
||||
try:
|
||||
from app.services.ha_discovery import publish_discovery
|
||||
|
||||
publish_discovery(session)
|
||||
except Exception:
|
||||
logger.exception("_trigger_discovery_republish: publish_discovery raised an error")
|
||||
|
||||
|
||||
def _get_meter_or_404(db: Session, meter_id: int) -> Meter:
|
||||
"""Return the meter with the given id or raise 404."""
|
||||
meter: Optional[Meter] = db.get(Meter, meter_id)
|
||||
if meter is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"Meter {meter_id!r} not found.",
|
||||
)
|
||||
return meter
|
||||
|
||||
|
||||
def _localize_started_at(dt: datetime) -> datetime:
|
||||
"""Resolve *dt* to an aware UTC datetime for storage.
|
||||
|
||||
Follows the same Principle-A convention as ``_localize_effective_from``
|
||||
in ``energy_contracts.py`` (FU10):
|
||||
|
||||
- Timezone-aware → convert to UTC as-is.
|
||||
- Timezone-naive → interpret as server local wall-clock time, localize
|
||||
with ``local_tz()``, then convert to UTC.
|
||||
|
||||
A front-end sending ``"2026-06-25T00:00:00"`` (no Z) has it interpreted
|
||||
as local midnight (e.g. CEST = UTC+2 → stored as 2026-06-24T22:00:00Z),
|
||||
not as UTC midnight.
|
||||
"""
|
||||
if dt.tzinfo is not None:
|
||||
return dt.astimezone(UTC)
|
||||
tz = _tz_mod.local_tz()
|
||||
local_dt = dt.replace(tzinfo=tz)
|
||||
return local_dt.astimezone(UTC)
|
||||
|
||||
|
||||
def _trigger_recompute(db: Session, start: datetime, label: str) -> int:
|
||||
"""Trigger recompute_range from *start* to now (UTC).
|
||||
|
||||
This is the standard "retroactive window" call: everything from the
|
||||
affected boundary up to the current moment needs re-attribution.
|
||||
Using ``now`` as the upper bound is safe because ``recompute_range``
|
||||
only touches closed quarter-hour periods and the operation is idempotent.
|
||||
"""
|
||||
end = datetime.now(UTC)
|
||||
if start >= end:
|
||||
# started_at is in the future — nothing to recompute.
|
||||
logger.info("%s: started_at (%s) is in the future, skipping recompute.", label, start)
|
||||
return 0
|
||||
n = recompute_range(db, start, end, commit=False, strict=True)
|
||||
logger.info(
|
||||
"%s: recomputed %d period(s) in window [%s, %s).",
|
||||
label,
|
||||
n,
|
||||
start.isoformat(),
|
||||
end.isoformat(),
|
||||
)
|
||||
return n
|
||||
|
||||
|
||||
def _recompute_commodity(db: Session, commodity: str, start: datetime, label: str) -> int:
|
||||
if commodity == "electricity":
|
||||
return _trigger_recompute(db, start, label)
|
||||
from app.services.meter_cost import recompute_range as thermal_recompute_range
|
||||
|
||||
end = datetime.now(UTC)
|
||||
if start >= end:
|
||||
return 0
|
||||
return thermal_recompute_range(db, start, end, commit=False)
|
||||
|
||||
|
||||
def _meter_response(meter: Meter) -> MeterResponse:
|
||||
"""Serialize meter plus binding summaries without exposing source config."""
|
||||
response = MeterResponse.model_validate(meter)
|
||||
response.bindings = [
|
||||
MeterBindingSummary(
|
||||
uuid=binding.uuid,
|
||||
source_channel_uuid=binding.channel.uuid,
|
||||
source_uuid=binding.channel.source.uuid,
|
||||
started_at=binding.started_at,
|
||||
ended_at=binding.ended_at,
|
||||
)
|
||||
for binding in meter.source_bindings
|
||||
]
|
||||
return response
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /api/energy/meters
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/meters", response_model=MeterListResponse)
|
||||
def list_energy_meters(
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> MeterListResponse:
|
||||
"""List all meter epochs in ascending ``started_at`` order.
|
||||
|
||||
Returns the full historical sequence of meter installations across all
|
||||
commodities. The active meter (``ended_at=null``) appears last because it
|
||||
has the latest ``started_at``.
|
||||
"""
|
||||
meters = list_meters(db)
|
||||
items = [_meter_response(m) for m in meters]
|
||||
return MeterListResponse(items=items, total=len(items))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# POST /api/energy/meters
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post(
|
||||
"/meters",
|
||||
response_model=MeterResponse,
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
)
|
||||
def declare_energy_meter(
|
||||
body: MeterDeclareRequest,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> MeterResponse:
|
||||
"""Declare a new meter epoch (swap, home move, or initial declaration).
|
||||
|
||||
Closes the current active meter for the given commodity at ``started_at``
|
||||
and opens a new active meter. If no active meter exists, the new meter is
|
||||
simply created without closing anything.
|
||||
|
||||
**Validation**: ``started_at`` must be **≥** the current active meter's
|
||||
own ``started_at`` (no chronological backdate below the active epoch's
|
||||
start). Equal timestamps are allowed (replaces the current meter at the
|
||||
same logical moment). Violation → 422.
|
||||
|
||||
**Retroactive recompute**: if ``started_at`` is in the past, billing
|
||||
records from that point forward are re-judged via ``recompute_range`` to
|
||||
reflect the new meter attribution. The recompute is transparent — the
|
||||
response body is the created meter (``MeterResponse``) only and does **not**
|
||||
include a recompute count; callers should re-fetch costs if they need the
|
||||
updated totals.
|
||||
"""
|
||||
started_at_utc = _localize_started_at(body.started_at)
|
||||
|
||||
try:
|
||||
old_meter = db.execute(
|
||||
select(Meter).where(Meter.commodity == body.commodity, Meter.ended_at.is_(None))
|
||||
).scalar_one_or_none()
|
||||
new_meter = declare_meter(
|
||||
db,
|
||||
label=body.label,
|
||||
started_at=started_at_utc,
|
||||
reason=body.reason.value,
|
||||
commodity=body.commodity,
|
||||
note=body.note,
|
||||
)
|
||||
db.flush() # assign PK before an optional binding and recompute
|
||||
# A closed predecessor must never retain an open interval. For a
|
||||
# meter swap with no selected channel we can safely hand off exactly
|
||||
# one compatible open channel; ambiguity is fail-closed.
|
||||
auto_channel = None
|
||||
if body.source_channel_uuid is None and old_meter is not None and body.reason.value == "meter_swap":
|
||||
candidates = [b for b in old_meter.source_bindings if b.ended_at is None and b.channel.unit == {"electricity": "kWh", "heating": "GJ", "hot_water": "m³"}.get(body.commodity)]
|
||||
if len(candidates) > 1:
|
||||
raise MeterSourceError("Meter swap has ambiguous open bindings; select a channel explicitly.")
|
||||
if len(candidates) == 1:
|
||||
auto_channel = candidates[0].channel
|
||||
if body.source_channel_uuid is not None:
|
||||
channel = db.execute(
|
||||
select(MeterSourceChannel).where(MeterSourceChannel.uuid == body.source_channel_uuid)
|
||||
).scalar_one_or_none()
|
||||
if channel is None:
|
||||
raise ChannelNotFoundError("Meter source channel was not found.")
|
||||
if body.reason.value == "meter_swap":
|
||||
create_binding_for_meter_swap(
|
||||
db,
|
||||
old_meter_id=old_meter.id if old_meter is not None else None,
|
||||
new_meter_id=new_meter.id,
|
||||
channel_id=channel.id,
|
||||
started_at=started_at_utc,
|
||||
)
|
||||
else:
|
||||
create_binding(
|
||||
db,
|
||||
meter_id=new_meter.id,
|
||||
channel_id=channel.id,
|
||||
started_at=started_at_utc,
|
||||
)
|
||||
elif auto_channel is not None:
|
||||
create_binding_for_meter_swap(db, old_meter_id=old_meter.id, new_meter_id=new_meter.id,
|
||||
channel_id=auto_channel.id, started_at=started_at_utc)
|
||||
if old_meter is not None:
|
||||
close_open_bindings_for_meter(db, old_meter.id, ended_at=started_at_utc)
|
||||
|
||||
# Keep recompute in this transaction: a failure must not leave a new
|
||||
# meter, its predecessor, or either binding at a half-applied boundary.
|
||||
now = datetime.now(UTC)
|
||||
if started_at_utc < now:
|
||||
db.flush()
|
||||
_recompute_commodity(db, body.commodity, started_at_utc, "POST /api/energy/meters")
|
||||
db.commit()
|
||||
except (MeterIntervalError, MeterOverlapError, MeterSourceError) as exc:
|
||||
db.rollback()
|
||||
raise HTTPException(
|
||||
status_code=(status.HTTP_404_NOT_FOUND if isinstance(exc, ChannelNotFoundError)
|
||||
else status.HTTP_422_UNPROCESSABLE_ENTITY),
|
||||
detail=str(exc),
|
||||
)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
db.refresh(new_meter)
|
||||
|
||||
# Trigger HA discovery re-publish so the new active meter's energy-cost
|
||||
# device/sensor configuration is pushed to Home Assistant. Best-effort:
|
||||
# failures are logged and swallowed; the API response is not affected.
|
||||
_trigger_discovery_republish(db)
|
||||
|
||||
logger.info(
|
||||
"POST /api/energy/meters: declared %r meter id=%d label=%r started_at=%s",
|
||||
body.commodity,
|
||||
new_meter.id,
|
||||
new_meter.label,
|
||||
started_at_utc.isoformat(),
|
||||
)
|
||||
return _meter_response(new_meter)
|
||||
|
||||
|
||||
@router.post("/meters/{meter_id}/close", response_model=MeterResponse)
|
||||
def close_energy_meter(
|
||||
meter_id: int, body: MeterCloseRequest, db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session), _csrf: None = Depends(require_csrf),
|
||||
) -> MeterResponse:
|
||||
meter = _get_meter_or_404(db, meter_id)
|
||||
boundary = _localize_started_at(body.ended_at)
|
||||
try:
|
||||
close_meter(db, meter, ended_at=boundary)
|
||||
close_open_bindings_for_meter(db, meter.id, ended_at=boundary)
|
||||
db.flush()
|
||||
_recompute_commodity(db, meter.commodity, boundary, f"POST /api/energy/meters/{meter_id}/close")
|
||||
db.commit()
|
||||
except (MeterIntervalError, MeterSourceError) as exc:
|
||||
db.rollback()
|
||||
raise HTTPException(status_code=422, detail=str(exc)) from exc
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
db.refresh(meter)
|
||||
_trigger_discovery_republish(db)
|
||||
return _meter_response(meter)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# PATCH /api/energy/meters/{id}
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.patch("/meters/{meter_id}", response_model=MeterResponse)
|
||||
def patch_energy_meter(
|
||||
meter_id: int,
|
||||
body: MeterPatchRequest,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> MeterResponse:
|
||||
"""Partially update a meter epoch: rename, edit note, or correct started_at.
|
||||
|
||||
- ``label``: updates the human-readable label.
|
||||
- ``note``: updates the free-form note.
|
||||
- ``started_at``: **retroactive correction** — shifts this meter's start
|
||||
boundary. The service layer maintains timeline continuity by also
|
||||
updating the preceding meter's ``ended_at``. Validation:
|
||||
* Must be strictly after the previous meter's own ``started_at``.
|
||||
* Must be strictly before this meter's ``ended_at`` (if closed).
|
||||
Violation → 422.
|
||||
|
||||
**Retroactive recompute when ``started_at`` changes**: billing records in
|
||||
the window ``[min(old, new), now)`` are re-judged to reflect the corrected
|
||||
meter attribution.
|
||||
|
||||
Not found → 404.
|
||||
"""
|
||||
meter = _get_meter_or_404(db, meter_id)
|
||||
|
||||
# Capture old started_at before mutation (needed for recompute window).
|
||||
old_started_at: Optional[datetime] = meter.started_at
|
||||
|
||||
# Localise started_at if provided.
|
||||
new_started_at_utc: Optional[datetime] = None
|
||||
if body.started_at is not None:
|
||||
new_started_at_utc = _localize_started_at(body.started_at)
|
||||
|
||||
try:
|
||||
update_meter(
|
||||
db,
|
||||
meter,
|
||||
label=body.label,
|
||||
note=body.note,
|
||||
started_at=new_started_at_utc,
|
||||
)
|
||||
|
||||
# Retroactive recompute if started_at was changed.
|
||||
if new_started_at_utc is not None and old_started_at is not None:
|
||||
# Normalise old_started_at to UTC-aware for comparison.
|
||||
if old_started_at.tzinfo is None:
|
||||
old_started_at = old_started_at.replace(tzinfo=UTC)
|
||||
# Window = [min(old, new), now) — covers all periods whose attribution
|
||||
# may have changed due to the boundary shift in either direction.
|
||||
window_start = min(old_started_at, new_started_at_utc)
|
||||
db.flush()
|
||||
_recompute_commodity(
|
||||
db,
|
||||
meter.commodity,
|
||||
window_start,
|
||||
f"PATCH /api/energy/meters/{meter_id}",
|
||||
)
|
||||
|
||||
db.commit()
|
||||
except MeterIntervalError as exc:
|
||||
db.rollback()
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=str(exc),
|
||||
)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
db.refresh(meter)
|
||||
|
||||
# Trigger HA discovery re-publish so label renames on the active meter
|
||||
# propagate to the HA device name. Best-effort: failures are logged and
|
||||
# swallowed; the API response is not affected.
|
||||
_trigger_discovery_republish(db)
|
||||
|
||||
logger.info(
|
||||
"PATCH /api/energy/meters/%d: updated meter label=%r started_at=%s",
|
||||
meter_id,
|
||||
meter.label,
|
||||
meter.started_at,
|
||||
)
|
||||
return _meter_response(meter)
|
||||
@@ -0,0 +1,511 @@
|
||||
"""Modbus device CRUD, readings, metrics, and test-read API (M5-T05).
|
||||
|
||||
All endpoints are under /api/modbus, require an authenticated session, and
|
||||
write endpoints (POST/PATCH/DELETE + test-read) additionally require a
|
||||
non-empty X-CSRF-Token header.
|
||||
|
||||
Deletion safety:
|
||||
DELETE /api/modbus/devices/{uuid} queries the application layer for
|
||||
existing modbus_reading rows before deleting. Without ``cascade`` it
|
||||
returns 409 and suggests disabling the device instead — a friendly guard
|
||||
rather than letting the DB raise. With ``cascade=true`` it removes the
|
||||
readings (and expose toggles) first, then the device. SQLite FK RESTRICT
|
||||
*is* enforced at runtime here (the app sets PRAGMA foreign_keys=ON; see
|
||||
app/db.py), so the cascade delete order matters.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
from sqlalchemy import delete as sa_delete, func, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.routes.api.deps import require_csrf, require_session
|
||||
from app.dependencies import get_db
|
||||
from app.integrations.modbus import driver as modbus_driver
|
||||
from app.integrations.modbus.driver import ModbusDriverError
|
||||
from app.integrations.modbus.profiles import (
|
||||
ProfileNotFoundError,
|
||||
decode as decode_profile,
|
||||
list_profiles,
|
||||
load_profile,
|
||||
)
|
||||
from app.models.expose import ExposedEntityToggle
|
||||
from app.models.modbus import ModbusDevice, ModbusReading
|
||||
from app.schemas.modbus import (
|
||||
MetricInfo,
|
||||
ModbusDeleteResponse,
|
||||
ModbusDeviceCreate,
|
||||
ModbusDeviceListResponse,
|
||||
ModbusDeviceResponse,
|
||||
ModbusDeviceUpdate,
|
||||
ModbusLatestResponse,
|
||||
ModbusMetricsResponse,
|
||||
ModbusProfilesResponse,
|
||||
ModbusReadingResponse,
|
||||
ModbusReadingsResponse,
|
||||
ModbusTestReadResponse,
|
||||
ProfileSummary,
|
||||
)
|
||||
from app.services.auth import AuthenticatedSession
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/api/modbus", tags=["api-modbus"])
|
||||
|
||||
# Maximum number of readings that can be returned per request.
|
||||
_READINGS_LIMIT_MAX = 5000
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_device_or_404(db: Session, uuid: str) -> ModbusDevice:
|
||||
"""Return the device with the given UUID or raise 404."""
|
||||
device = db.execute(
|
||||
select(ModbusDevice).where(ModbusDevice.uuid == uuid)
|
||||
).scalar_one_or_none()
|
||||
if device is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"Device '{uuid}' not found.",
|
||||
)
|
||||
return device
|
||||
|
||||
|
||||
def _label_from_key(key: str) -> str:
|
||||
"""Derive a human-readable label from a metric key.
|
||||
|
||||
Example: ``"active_power"`` → ``"Active Power"``
|
||||
"""
|
||||
return key.replace("_", " ").title()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /api/modbus/profiles
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/profiles", response_model=ModbusProfilesResponse)
|
||||
def get_profiles(
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> ModbusProfilesResponse:
|
||||
"""List all available Modbus YAML profiles (name + description).
|
||||
|
||||
Intended for the front-end's device-creation profile drop-down.
|
||||
"""
|
||||
pairs = list_profiles()
|
||||
return ModbusProfilesResponse(
|
||||
profiles=[ProfileSummary(name=name, description=desc) for name, desc in pairs]
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Device CRUD
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/devices", response_model=ModbusDeviceListResponse)
|
||||
def list_devices(
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> ModbusDeviceListResponse:
|
||||
"""Return all Modbus devices (no pagination — device counts stay small)."""
|
||||
devices = db.execute(select(ModbusDevice)).scalars().all()
|
||||
items = [ModbusDeviceResponse.model_validate(d) for d in devices]
|
||||
return ModbusDeviceListResponse(items=items, total=len(items))
|
||||
|
||||
|
||||
@router.post(
|
||||
"/devices",
|
||||
response_model=ModbusDeviceResponse,
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
)
|
||||
def create_device(
|
||||
body: ModbusDeviceCreate,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> ModbusDeviceResponse:
|
||||
"""Create a new Modbus device.
|
||||
|
||||
- Validates that the referenced ``profile`` exists; returns 422 if not.
|
||||
- Returns 201 with the created device on success.
|
||||
"""
|
||||
# Validate profile exists (422 if not)
|
||||
try:
|
||||
load_profile(body.profile)
|
||||
except ProfileNotFoundError:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"Profile '{body.profile}' does not exist.",
|
||||
)
|
||||
|
||||
now = datetime.now(UTC)
|
||||
device = ModbusDevice(
|
||||
friendly_name=body.friendly_name,
|
||||
transport=body.transport,
|
||||
host=body.host,
|
||||
port=body.port,
|
||||
unit_id=body.unit_id,
|
||||
profile=body.profile,
|
||||
poll_interval_s=body.poll_interval_s,
|
||||
enabled=body.enabled,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
db.add(device)
|
||||
db.commit()
|
||||
db.refresh(device)
|
||||
logger.info("Created Modbus device %r (uuid=%s)", device.friendly_name, device.uuid)
|
||||
return ModbusDeviceResponse.model_validate(device)
|
||||
|
||||
|
||||
@router.get("/devices/{uuid}", response_model=ModbusDeviceResponse)
|
||||
def get_device(
|
||||
uuid: str,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> ModbusDeviceResponse:
|
||||
"""Return a single Modbus device by UUID."""
|
||||
device = _get_device_or_404(db, uuid)
|
||||
return ModbusDeviceResponse.model_validate(device)
|
||||
|
||||
|
||||
@router.patch("/devices/{uuid}", response_model=ModbusDeviceResponse)
|
||||
def patch_device(
|
||||
uuid: str,
|
||||
body: ModbusDeviceUpdate,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> ModbusDeviceResponse:
|
||||
"""Partially update a Modbus device (including enable/disable).
|
||||
|
||||
Only fields explicitly provided in the request body are updated.
|
||||
Providing ``profile`` triggers a profile-existence check (422 if unknown).
|
||||
"""
|
||||
device = _get_device_or_404(db, uuid)
|
||||
|
||||
update_data = body.model_dump(exclude_none=True)
|
||||
|
||||
# If profile is being changed, validate it exists.
|
||||
if "profile" in update_data:
|
||||
try:
|
||||
load_profile(update_data["profile"])
|
||||
except ProfileNotFoundError:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"Profile '{update_data['profile']}' does not exist.",
|
||||
)
|
||||
|
||||
for field, value in update_data.items():
|
||||
setattr(device, field, value)
|
||||
|
||||
device.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
db.refresh(device)
|
||||
logger.info("Updated Modbus device %r (uuid=%s)", device.friendly_name, device.uuid)
|
||||
return ModbusDeviceResponse.model_validate(device)
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/devices/{uuid}",
|
||||
status_code=status.HTTP_204_NO_CONTENT,
|
||||
response_model=None,
|
||||
)
|
||||
def delete_device(
|
||||
uuid: str,
|
||||
cascade: bool = Query(default=False),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> None | ModbusDeleteResponse:
|
||||
"""Delete a Modbus device.
|
||||
|
||||
**Default behaviour (cascade=false)**:
|
||||
Returns 409 Conflict if the device has any associated readings; use
|
||||
``enabled=false`` to disable it instead.
|
||||
|
||||
**Cascade delete (cascade=true)**:
|
||||
Permanently deletes the device together with all its readings and any
|
||||
``ExposedEntityToggle`` rows whose key matches ``modbus.<uuid>.*``.
|
||||
Also makes a best-effort attempt to clear the device's HA Discovery
|
||||
config topics from MQTT (empty retained payload) before the DB rows
|
||||
are removed. MQTT failures are swallowed — the DB deletion proceeds
|
||||
regardless.
|
||||
Returns HTTP 200 with a ``ModbusDeleteResponse`` JSON body on success.
|
||||
|
||||
**Application-layer safety**: the 409 guard uses an explicit SELECT COUNT
|
||||
query to return a friendly message. FK RESTRICT is enforced at runtime
|
||||
(the app sets ``PRAGMA foreign_keys=ON``), so the cascade path deletes
|
||||
readings before the device.
|
||||
"""
|
||||
from app.services.ha_discovery import clear_device_discovery
|
||||
|
||||
device = _get_device_or_404(db, uuid)
|
||||
|
||||
# Application-layer guard: refuse deletion if any readings exist.
|
||||
reading_count: int = db.execute(
|
||||
select(func.count(ModbusReading.id)).where(ModbusReading.device_id == device.id)
|
||||
).scalar_one()
|
||||
|
||||
if reading_count > 0 and not cascade:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=(
|
||||
f"Device '{uuid}' has {reading_count} associated reading(s) and cannot be "
|
||||
"deleted. Disable the device (set enabled=false) instead to stop polling "
|
||||
"without losing historical data."
|
||||
),
|
||||
)
|
||||
|
||||
if cascade:
|
||||
# --- Cascade deletion path ---
|
||||
logger.info(
|
||||
"Cascade-deleting Modbus device %r (uuid=%s): %d reading(s)",
|
||||
device.friendly_name,
|
||||
uuid,
|
||||
reading_count,
|
||||
)
|
||||
|
||||
# 1. Best-effort HA Discovery cleanup (before DB rows are gone).
|
||||
try:
|
||||
clear_device_discovery(db, uuid)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception(
|
||||
"delete_device(cascade): HA discovery cleanup raised for uuid=%s; continuing",
|
||||
uuid,
|
||||
)
|
||||
|
||||
# 2. Delete all readings for this device (must happen before device row).
|
||||
deleted_readings_result = db.execute(
|
||||
sa_delete(ModbusReading).where(ModbusReading.device_id == device.id)
|
||||
)
|
||||
readings_deleted: int = deleted_readings_result.rowcount
|
||||
|
||||
# 3. Delete all ExposedEntityToggle rows for this device.
|
||||
# Keys follow the pattern "modbus.<uuid>.<metric>".
|
||||
toggle_prefix = f"modbus.{uuid}.%"
|
||||
deleted_toggles_result = db.execute(
|
||||
sa_delete(ExposedEntityToggle).where(ExposedEntityToggle.key.like(toggle_prefix))
|
||||
)
|
||||
toggles_deleted: int = deleted_toggles_result.rowcount
|
||||
|
||||
# 4. Delete the device itself.
|
||||
db.delete(device)
|
||||
db.commit()
|
||||
|
||||
logger.info(
|
||||
"Cascade-delete complete for device uuid=%s: %d reading(s), %d toggle(s) removed",
|
||||
uuid,
|
||||
readings_deleted,
|
||||
toggles_deleted,
|
||||
)
|
||||
|
||||
from fastapi.responses import JSONResponse
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_200_OK,
|
||||
content={
|
||||
"deleted": True,
|
||||
"readings_deleted": readings_deleted,
|
||||
"toggles_deleted": toggles_deleted,
|
||||
},
|
||||
)
|
||||
|
||||
# --- Non-cascade path (no readings exist at this point) ---
|
||||
logger.info("Deleting Modbus device %r (uuid=%s)", device.friendly_name, uuid)
|
||||
db.delete(device)
|
||||
db.commit()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Readings endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/devices/{uuid}/latest", response_model=ModbusLatestResponse)
|
||||
def get_latest_reading(
|
||||
uuid: str,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> ModbusLatestResponse:
|
||||
"""Return the most recent reading for a device.
|
||||
|
||||
If no readings exist yet, returns ``{"found": false, "recorded_at": null,
|
||||
"payload": null}`` (200, not 404) so the front-end can distinguish
|
||||
"device exists but has no data" from "device not found".
|
||||
"""
|
||||
device = _get_device_or_404(db, uuid)
|
||||
|
||||
row = db.execute(
|
||||
select(ModbusReading)
|
||||
.where(ModbusReading.device_id == device.id)
|
||||
.order_by(ModbusReading.recorded_at.desc())
|
||||
.limit(1)
|
||||
).scalar_one_or_none()
|
||||
|
||||
if row is None:
|
||||
return ModbusLatestResponse(found=False, recorded_at=None, payload=None)
|
||||
|
||||
return ModbusLatestResponse(
|
||||
found=True,
|
||||
recorded_at=row.recorded_at,
|
||||
payload=row.payload,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/devices/{uuid}/readings", response_model=ModbusReadingsResponse)
|
||||
def get_readings(
|
||||
uuid: str,
|
||||
start: datetime | None = Query(default=None),
|
||||
end: datetime | None = Query(default=None),
|
||||
limit: int = Query(default=500, ge=1, le=_READINGS_LIMIT_MAX),
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> ModbusReadingsResponse:
|
||||
"""Return time-range readings for a device.
|
||||
|
||||
When the window contains more rows than ``limit``, the **most recent** N rows
|
||||
are returned (``ORDER BY recorded_at DESC LIMIT n``), then reversed to
|
||||
ascending order before being sent to the client. This ensures that for long
|
||||
time-range requests (e.g. 6 h / 24 h) the caller always sees the latest data
|
||||
rather than the oldest segment of the window.
|
||||
|
||||
When the window has fewer rows than ``limit`` the full window is returned in
|
||||
ascending order — behaviour is identical to a plain ascending query.
|
||||
|
||||
The response schema is unchanged: items are always ``recorded_at`` ascending.
|
||||
|
||||
The query uses the ``(device_id, recorded_at)`` composite index for
|
||||
efficient time-window scans.
|
||||
|
||||
Query parameters:
|
||||
- ``start``: inclusive lower bound (ISO8601 datetime)
|
||||
- ``end``: inclusive upper bound (ISO8601 datetime)
|
||||
- ``limit``: max rows to return (default 500, max {_READINGS_LIMIT_MAX})
|
||||
"""
|
||||
device = _get_device_or_404(db, uuid)
|
||||
|
||||
stmt = (
|
||||
select(ModbusReading)
|
||||
.where(ModbusReading.device_id == device.id)
|
||||
.order_by(ModbusReading.recorded_at.desc())
|
||||
.limit(limit)
|
||||
)
|
||||
|
||||
if start is not None:
|
||||
stmt = stmt.where(ModbusReading.recorded_at >= start)
|
||||
if end is not None:
|
||||
stmt = stmt.where(ModbusReading.recorded_at <= end)
|
||||
|
||||
# Fetch most-recent N rows, then reverse to restore ascending order for the response.
|
||||
rows = list(reversed(db.execute(stmt).scalars().all()))
|
||||
items = [ModbusReadingResponse(recorded_at=r.recorded_at, payload=r.payload) for r in rows]
|
||||
return ModbusReadingsResponse(items=items)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Metrics endpoint
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/devices/{uuid}/metrics", response_model=ModbusMetricsResponse)
|
||||
def get_metrics(
|
||||
uuid: str,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> ModbusMetricsResponse:
|
||||
"""Return the metric catalogue for a device's profile.
|
||||
|
||||
Each entry has ``key``, ``label`` (derived from key if the profile has
|
||||
none — underscores → spaces, title-cased), ``unit``, and ``device_class``.
|
||||
This is the authoritative metadata source for front-end card labels and
|
||||
chart axis labels.
|
||||
"""
|
||||
device = _get_device_or_404(db, uuid)
|
||||
|
||||
try:
|
||||
profile = load_profile(device.profile)
|
||||
except ProfileNotFoundError:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"Device profile '{device.profile}' could not be loaded.",
|
||||
)
|
||||
|
||||
metrics = [
|
||||
MetricInfo(
|
||||
key=m.key,
|
||||
label=_label_from_key(m.key),
|
||||
unit=m.unit,
|
||||
device_class=m.device_class,
|
||||
)
|
||||
for m in profile.metrics
|
||||
]
|
||||
|
||||
return ModbusMetricsResponse(profile=profile.name, metrics=metrics)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test-read endpoint
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/devices/{uuid}/test", response_model=ModbusTestReadResponse)
|
||||
def test_read(
|
||||
uuid: str,
|
||||
db: Session = Depends(get_db),
|
||||
_auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> ModbusTestReadResponse:
|
||||
"""Immediately read and decode the device once without persisting any data.
|
||||
|
||||
Useful for validating gateway connectivity and Modbus address settings
|
||||
before relying on the background polling job.
|
||||
|
||||
This is a write-class endpoint (it initiates network I/O on demand) and
|
||||
therefore requires a CSRF token. The response payload is never stored.
|
||||
"""
|
||||
device = _get_device_or_404(db, uuid)
|
||||
|
||||
try:
|
||||
profile = load_profile(device.profile)
|
||||
except ProfileNotFoundError:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"Device profile '{device.profile}' could not be loaded.",
|
||||
)
|
||||
|
||||
try:
|
||||
registers = modbus_driver.read_blocks(
|
||||
device.host,
|
||||
device.port,
|
||||
device.unit_id,
|
||||
[{"start": b.start, "count": b.count} for b in profile.blocks],
|
||||
function_code=profile.function_code,
|
||||
)
|
||||
payload: dict[str, Any] = decode_profile(profile, registers)
|
||||
return ModbusTestReadResponse(ok=True, payload=payload)
|
||||
|
||||
except ModbusDriverError as exc:
|
||||
logger.warning(
|
||||
"Test read failed for device %r (uuid=%s): %s",
|
||||
device.friendly_name,
|
||||
uuid,
|
||||
exc,
|
||||
)
|
||||
return ModbusTestReadResponse(ok=False, error=str(exc))
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning(
|
||||
"Unexpected error during test read for device %r (uuid=%s): %s",
|
||||
device.friendly_name,
|
||||
uuid,
|
||||
exc,
|
||||
)
|
||||
return ModbusTestReadResponse(ok=False, error=f"Unexpected error: {exc}")
|
||||
@@ -0,0 +1,324 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.routes.api.deps import require_csrf, require_session
|
||||
from app.config import Settings
|
||||
from app.dependencies import get_app_settings, get_db
|
||||
from app.schemas.session import (
|
||||
LoginRequest,
|
||||
PasswordChangeRequest,
|
||||
SessionResponse,
|
||||
SessionUser,
|
||||
)
|
||||
from app.schemas.totp import (
|
||||
TotpDisableRequest,
|
||||
TotpEnableRequest,
|
||||
TotpSetupResponse,
|
||||
TotpStatusResponse,
|
||||
)
|
||||
from app.services.auth import (
|
||||
AuthPasswordChangeError,
|
||||
AuthenticatedSession,
|
||||
authenticate_user,
|
||||
change_password,
|
||||
create_session,
|
||||
revoke_session,
|
||||
)
|
||||
from app.services import login_throttle
|
||||
from app.services import totp as totp_service
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/api", tags=["api-session"])
|
||||
|
||||
|
||||
def _build_session_response(auth: AuthenticatedSession) -> SessionResponse:
|
||||
return SessionResponse(
|
||||
user=SessionUser(
|
||||
username=auth.user.username,
|
||||
force_password_change=auth.user.force_password_change,
|
||||
),
|
||||
csrf_token=auth.session.csrf_token,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/session", response_model=SessionResponse)
|
||||
def get_session(
|
||||
auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> SessionResponse:
|
||||
"""Return the current session user and CSRF token. Returns 401 if not authenticated."""
|
||||
return _build_session_response(auth)
|
||||
|
||||
|
||||
def _get_client_ip(request: Request, *, trust_forwarded_for: bool) -> str:
|
||||
"""Extract the client IP address from the request.
|
||||
|
||||
When ``trust_forwarded_for`` is True (reverse-proxy deployments) the
|
||||
left-most value from the ``X-Forwarded-For`` header is used. Otherwise the
|
||||
direct socket IP (``request.client.host``) is used.
|
||||
"""
|
||||
if trust_forwarded_for:
|
||||
xff = request.headers.get("X-Forwarded-For", "")
|
||||
if xff:
|
||||
return xff.split(",")[0].strip()
|
||||
if request.client is not None:
|
||||
return request.client.host
|
||||
return "unknown"
|
||||
|
||||
|
||||
@router.post("/auth/login", response_model=SessionResponse)
|
||||
def post_login(
|
||||
body: LoginRequest,
|
||||
request: Request,
|
||||
response: Response,
|
||||
db: Session = Depends(get_db),
|
||||
settings: Settings = Depends(get_app_settings),
|
||||
) -> SessionResponse:
|
||||
"""
|
||||
Authenticate with username and password.
|
||||
|
||||
On success, sets an HttpOnly session cookie and returns the session user + CSRF token.
|
||||
On failure, returns 401 with no cookie set.
|
||||
Repeated failures trigger exponential back-off (429 + Retry-After).
|
||||
No X-CSRF-Token required (unauthenticated endpoint).
|
||||
"""
|
||||
client_ip = _get_client_ip(request, trust_forwarded_for=settings.auth_trust_forwarded_for)
|
||||
|
||||
# --- Throttle check (before any password verification) ---
|
||||
if settings.auth_login_throttle_enabled:
|
||||
wait_seconds = login_throttle.check_and_get_wait(
|
||||
db, ip=client_ip, username=body.username
|
||||
)
|
||||
if wait_seconds > 0:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
|
||||
detail="too many failed login attempts; please try again later",
|
||||
headers={"Retry-After": str(wait_seconds)},
|
||||
)
|
||||
|
||||
# --- Password verification ---
|
||||
user = authenticate_user(db, username=body.username, password=body.password)
|
||||
if user is None:
|
||||
if settings.auth_login_throttle_enabled:
|
||||
login_throttle.register_failure(db, ip=client_ip, username=body.username)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="invalid username or password",
|
||||
)
|
||||
|
||||
# --- TOTP second-factor check (only when TOTP is enabled for the user) ---
|
||||
if user.totp_enabled:
|
||||
if not body.totp_code:
|
||||
# Password correct but no TOTP code supplied: signal the front-end to
|
||||
# prompt for the second factor. Do NOT issue a session.
|
||||
# Deliberate: this is a normal two-step protocol step from a legitimate
|
||||
# user — we do NOT register a throttle failure here. The attacker must
|
||||
# already have the correct password to reach this branch, and counting
|
||||
# each normal first-step as a failure would risk locking out the real
|
||||
# admin on every login attempt.
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail={"totp_required": True},
|
||||
)
|
||||
|
||||
# TOTP code supplied — verify against time-based code or a recovery code.
|
||||
totp_ok = totp_service.verify_totp_code(user, body.totp_code)
|
||||
if not totp_ok:
|
||||
totp_ok = totp_service.verify_recovery_code(db, user=user, code=body.totp_code)
|
||||
|
||||
if not totp_ok:
|
||||
# Second-factor failure is an active attack signal — register failure
|
||||
# so that repeated wrong TOTP/recovery-code guesses trigger back-off.
|
||||
if settings.auth_login_throttle_enabled:
|
||||
login_throttle.register_failure(db, ip=client_ip, username=body.username)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="invalid username or password",
|
||||
)
|
||||
|
||||
# --- Success: clear back-off state and issue session ---
|
||||
if settings.auth_login_throttle_enabled:
|
||||
login_throttle.clear(db, ip=client_ip, username=body.username)
|
||||
|
||||
auth_session, raw_token = create_session(db, user=user, settings=settings)
|
||||
logger.info("Created API authenticated session for user '%s'", user.username)
|
||||
|
||||
response.set_cookie(
|
||||
key=settings.auth_session_cookie_name,
|
||||
value=raw_token,
|
||||
max_age=settings.auth_session_ttl_hours * 3600,
|
||||
httponly=True,
|
||||
secure=settings.auth_cookie_secure,
|
||||
samesite="lax",
|
||||
path="/",
|
||||
)
|
||||
|
||||
auth = AuthenticatedSession(user=user, session=auth_session)
|
||||
return _build_session_response(auth)
|
||||
|
||||
|
||||
@router.post("/auth/logout")
|
||||
def post_logout(
|
||||
response: Response,
|
||||
db: Session = Depends(get_db),
|
||||
settings: Settings = Depends(get_app_settings),
|
||||
auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> Response:
|
||||
"""
|
||||
Revoke the current session and clear the session cookie.
|
||||
Requires authentication and X-CSRF-Token header.
|
||||
Returns 204 No Content.
|
||||
"""
|
||||
revoke_session(db, auth_session=auth.session)
|
||||
logger.info("Revoked API authenticated session for user '%s'", auth.user.username)
|
||||
no_content = Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||
no_content.delete_cookie(settings.auth_session_cookie_name, path="/")
|
||||
return no_content
|
||||
|
||||
|
||||
@router.post("/auth/password")
|
||||
def post_change_password(
|
||||
body: PasswordChangeRequest,
|
||||
db: Session = Depends(get_db),
|
||||
auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> Response:
|
||||
"""
|
||||
Change the current user's password.
|
||||
Requires authentication and X-CSRF-Token header.
|
||||
On AuthPasswordChangeError returns 400 with a generic message.
|
||||
On success, force_password_change becomes False (handled by the service).
|
||||
Returns 204 No Content.
|
||||
"""
|
||||
try:
|
||||
change_password(
|
||||
db,
|
||||
user=auth.user,
|
||||
current_password=body.current_password,
|
||||
new_password=body.new_password,
|
||||
confirm_password=body.confirm_password,
|
||||
)
|
||||
except AuthPasswordChangeError as exc:
|
||||
logger.info(
|
||||
"Rejected password change for user '%s': %s",
|
||||
auth.user.username,
|
||||
exc,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="password change failed",
|
||||
) from exc
|
||||
|
||||
logger.info("Password updated for user '%s'", auth.user.username)
|
||||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# TOTP endpoints (M4-T05)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/auth/totp/setup", response_model=TotpSetupResponse)
|
||||
def post_totp_setup(
|
||||
db: Session = Depends(get_db),
|
||||
settings: Settings = Depends(get_app_settings),
|
||||
auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> TotpSetupResponse:
|
||||
"""
|
||||
Generate a new pending TOTP secret, otpauth URI, and one-time recovery codes.
|
||||
|
||||
The secret is stored in the DB but TOTP is NOT yet enabled (totp_enabled stays
|
||||
False until the user confirms with POST /api/auth/totp/enable).
|
||||
|
||||
Recovery codes are returned here as plaintext exactly once; their Argon2 hashes
|
||||
are persisted immediately so enable only needs to flip the enabled flag.
|
||||
|
||||
Repeating this call replaces any prior pending secret and regenerates codes.
|
||||
|
||||
Requires: session cookie + X-CSRF-Token.
|
||||
"""
|
||||
secret, otpauth_uri, recovery_codes = totp_service.setup(
|
||||
db,
|
||||
user=auth.user,
|
||||
issuer=settings.effective_totp_issuer,
|
||||
)
|
||||
return TotpSetupResponse(
|
||||
secret=secret,
|
||||
otpauth_uri=otpauth_uri,
|
||||
recovery_codes=recovery_codes,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/auth/totp/enable", status_code=status.HTTP_204_NO_CONTENT)
|
||||
def post_totp_enable(
|
||||
body: TotpEnableRequest,
|
||||
db: Session = Depends(get_db),
|
||||
auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> Response:
|
||||
"""
|
||||
Enable TOTP by confirming with the current 6-digit code from the authenticator app.
|
||||
|
||||
Requires a prior call to POST /api/auth/totp/setup (so that a pending secret
|
||||
exists). On success, totp_enabled becomes True.
|
||||
|
||||
Returns 400 if the code is wrong or there is no pending secret.
|
||||
Requires: session cookie + X-CSRF-Token.
|
||||
"""
|
||||
ok = totp_service.enable(db, user=auth.user, code=body.code)
|
||||
if not ok:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="invalid TOTP code or no pending setup",
|
||||
)
|
||||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||
|
||||
|
||||
@router.post("/auth/totp/disable", status_code=status.HTTP_204_NO_CONTENT)
|
||||
def post_totp_disable(
|
||||
body: TotpDisableRequest,
|
||||
db: Session = Depends(get_db),
|
||||
auth: AuthenticatedSession = Depends(require_session),
|
||||
_csrf: None = Depends(require_csrf),
|
||||
) -> Response:
|
||||
"""
|
||||
Disable TOTP. The caller must provide exactly one of:
|
||||
- ``password``: the user's current login password, OR
|
||||
- ``code``: the current 6-digit TOTP code.
|
||||
|
||||
On success: totp_enabled=False, totp_secret cleared, all recovery codes deleted.
|
||||
Returns 400 if neither credential matches or neither is provided.
|
||||
Requires: session cookie + X-CSRF-Token.
|
||||
"""
|
||||
ok = totp_service.disable(
|
||||
db,
|
||||
user=auth.user,
|
||||
password=body.password,
|
||||
code=body.code,
|
||||
)
|
||||
if not ok:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="invalid credential; provide a valid password or TOTP code",
|
||||
)
|
||||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||
|
||||
|
||||
@router.get("/auth/totp", response_model=TotpStatusResponse)
|
||||
def get_totp_status(
|
||||
auth: AuthenticatedSession = Depends(require_session),
|
||||
) -> TotpStatusResponse:
|
||||
"""
|
||||
Return the current TOTP status for the authenticated user.
|
||||
|
||||
Response contains only ``{"enabled": bool}``.
|
||||
Secret and recovery codes are NEVER returned here.
|
||||
Requires: session cookie only (no CSRF — read-only).
|
||||
"""
|
||||
return TotpStatusResponse(enabled=auth.user.totp_enabled)
|
||||
@@ -1,234 +0,0 @@
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
from fastapi import APIRouter, Depends, Form, Request, status
|
||||
from fastapi.responses import HTMLResponse, RedirectResponse, Response
|
||||
from fastapi.templating import Jinja2Templates
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.config import Settings
|
||||
from app.dependencies import get_app_settings, get_db, get_current_auth_session
|
||||
from app.services.auth import (
|
||||
AuthenticatedSession,
|
||||
authenticate_user,
|
||||
change_password,
|
||||
create_session,
|
||||
AuthPasswordChangeError,
|
||||
issue_login_csrf_token,
|
||||
revoke_session,
|
||||
validate_csrf_token,
|
||||
)
|
||||
from app.services.config_page import build_config_sections, is_ticktick_oauth_ready
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
templates = Jinja2Templates(directory=str(Path(__file__).resolve().parents[2] / "templates"))
|
||||
router = APIRouter(tags=["auth"])
|
||||
|
||||
LOGIN_CSRF_COOKIE_NAME = "login_csrf"
|
||||
|
||||
|
||||
@router.get("/login", response_class=HTMLResponse)
|
||||
def login_page(
|
||||
request: Request,
|
||||
settings: Settings = Depends(get_app_settings),
|
||||
current_auth: AuthenticatedSession | None = Depends(get_current_auth_session),
|
||||
) -> Response:
|
||||
if current_auth is not None:
|
||||
return RedirectResponse(url="/config", status_code=status.HTTP_303_SEE_OTHER)
|
||||
|
||||
csrf_token = issue_login_csrf_token()
|
||||
response = templates.TemplateResponse(
|
||||
request,
|
||||
"login.html",
|
||||
{
|
||||
"app_name": settings.app_name,
|
||||
"app_env": settings.app_env,
|
||||
"csrf_token": csrf_token,
|
||||
"error_message": None,
|
||||
},
|
||||
)
|
||||
_set_login_csrf_cookie(response, settings=settings, token=csrf_token)
|
||||
return response
|
||||
|
||||
|
||||
@router.post("/login", response_class=HTMLResponse)
|
||||
def login_submit(
|
||||
request: Request,
|
||||
username: str = Form(),
|
||||
password: str = Form(),
|
||||
csrf_token: str = Form(),
|
||||
session: Session = Depends(get_db),
|
||||
settings: Settings = Depends(get_app_settings),
|
||||
) -> Response:
|
||||
cookie_csrf_token = request.cookies.get(LOGIN_CSRF_COOKIE_NAME)
|
||||
if not validate_csrf_token(expected=cookie_csrf_token, actual=csrf_token):
|
||||
logger.warning("Rejected login attempt due to CSRF validation failure")
|
||||
return _render_login_error(
|
||||
request,
|
||||
settings=settings,
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
error_message="invalid login request",
|
||||
)
|
||||
|
||||
user = authenticate_user(session, username=username, password=password)
|
||||
if user is None:
|
||||
return _render_login_error(
|
||||
request,
|
||||
settings=settings,
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
error_message="invalid username or password",
|
||||
)
|
||||
|
||||
auth_session, raw_token = create_session(session, user=user, settings=settings)
|
||||
response = RedirectResponse(url="/config", status_code=status.HTTP_303_SEE_OTHER)
|
||||
response.delete_cookie(LOGIN_CSRF_COOKIE_NAME, path="/login")
|
||||
response.set_cookie(
|
||||
key=settings.auth_session_cookie_name,
|
||||
value=raw_token,
|
||||
max_age=settings.auth_session_ttl_hours * 3600,
|
||||
httponly=True,
|
||||
secure=settings.auth_cookie_secure,
|
||||
samesite="lax",
|
||||
path="/",
|
||||
)
|
||||
logger.info("Created authenticated session for user '%s'", user.username)
|
||||
return response
|
||||
|
||||
|
||||
@router.post("/config/change-password", response_class=HTMLResponse)
|
||||
def change_password_submit(
|
||||
request: Request,
|
||||
current_password: str = Form(),
|
||||
new_password: str = Form(),
|
||||
confirm_password: str = Form(),
|
||||
csrf_token: str = Form(),
|
||||
session: Session = Depends(get_db),
|
||||
settings: Settings = Depends(get_app_settings),
|
||||
current_auth: AuthenticatedSession | None = Depends(get_current_auth_session),
|
||||
) -> Response:
|
||||
if current_auth is None:
|
||||
return RedirectResponse(url="/login", status_code=status.HTTP_303_SEE_OTHER)
|
||||
|
||||
if not validate_csrf_token(expected=current_auth.session.csrf_token, actual=csrf_token):
|
||||
logger.warning("Rejected password change attempt due to CSRF validation failure")
|
||||
return _render_config_page(
|
||||
request,
|
||||
settings=settings,
|
||||
auth_db_session=session,
|
||||
current_auth=current_auth,
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
password_change_error="invalid password change request",
|
||||
)
|
||||
|
||||
try:
|
||||
change_password(
|
||||
session,
|
||||
user=current_auth.user,
|
||||
current_password=current_password,
|
||||
new_password=new_password,
|
||||
confirm_password=confirm_password,
|
||||
)
|
||||
except AuthPasswordChangeError as exc:
|
||||
logger.info(
|
||||
"Rejected password change for user '%s': %s",
|
||||
current_auth.user.username,
|
||||
exc,
|
||||
)
|
||||
return _render_config_page(
|
||||
request,
|
||||
settings=settings,
|
||||
auth_db_session=session,
|
||||
current_auth=current_auth,
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
password_change_error="password change failed",
|
||||
)
|
||||
|
||||
logger.info("Password updated for user '%s'", current_auth.user.username)
|
||||
return RedirectResponse(url="/config", status_code=status.HTTP_303_SEE_OTHER)
|
||||
|
||||
|
||||
@router.post("/logout")
|
||||
def logout(
|
||||
request: Request,
|
||||
csrf_token: str = Form(),
|
||||
session: Session = Depends(get_db),
|
||||
settings: Settings = Depends(get_app_settings),
|
||||
current_auth: AuthenticatedSession | None = Depends(get_current_auth_session),
|
||||
) -> RedirectResponse:
|
||||
if current_auth is not None and validate_csrf_token(
|
||||
expected=current_auth.session.csrf_token, actual=csrf_token
|
||||
):
|
||||
revoke_session(session, auth_session=current_auth.session)
|
||||
logger.info("Revoked authenticated session for user '%s'", current_auth.user.username)
|
||||
else:
|
||||
logger.warning("Rejected logout request due to missing session or invalid CSRF token")
|
||||
|
||||
response = RedirectResponse(url="/login", status_code=status.HTTP_303_SEE_OTHER)
|
||||
response.delete_cookie(settings.auth_session_cookie_name, path="/")
|
||||
return response
|
||||
|
||||
|
||||
def _render_login_error(
|
||||
request: Request,
|
||||
*,
|
||||
settings: Settings,
|
||||
status_code: int,
|
||||
error_message: str,
|
||||
) -> HTMLResponse:
|
||||
csrf_token = issue_login_csrf_token()
|
||||
response = templates.TemplateResponse(
|
||||
request,
|
||||
"login.html",
|
||||
{
|
||||
"app_name": settings.app_name,
|
||||
"app_env": settings.app_env,
|
||||
"csrf_token": csrf_token,
|
||||
"error_message": error_message,
|
||||
},
|
||||
status_code=status_code,
|
||||
)
|
||||
_set_login_csrf_cookie(response, settings=settings, token=csrf_token)
|
||||
return response
|
||||
|
||||
|
||||
def _set_login_csrf_cookie(response: HTMLResponse, *, settings: Settings, token: str) -> None:
|
||||
response.set_cookie(
|
||||
key=LOGIN_CSRF_COOKIE_NAME,
|
||||
value=token,
|
||||
max_age=1800,
|
||||
httponly=True,
|
||||
secure=settings.auth_cookie_secure,
|
||||
samesite="lax",
|
||||
path="/login",
|
||||
)
|
||||
|
||||
|
||||
def _render_config_page(
|
||||
request: Request,
|
||||
*,
|
||||
settings: Settings,
|
||||
auth_db_session: Session,
|
||||
current_auth: AuthenticatedSession,
|
||||
status_code: int,
|
||||
password_change_error: str | None,
|
||||
) -> HTMLResponse:
|
||||
return templates.TemplateResponse(
|
||||
request,
|
||||
"config.html",
|
||||
{
|
||||
"app_name": settings.app_name,
|
||||
"app_env": settings.app_env,
|
||||
"current_username": current_auth.user.username,
|
||||
"csrf_token": current_auth.session.csrf_token,
|
||||
"force_password_change": current_auth.user.force_password_change,
|
||||
"password_change_error": password_change_error,
|
||||
"config_error": None,
|
||||
"config_saved": False,
|
||||
"config_sections": build_config_sections(auth_db_session, settings),
|
||||
"ticktick_oauth_ready": is_ticktick_oauth_ready(settings),
|
||||
"ticktick_redirect_uri": settings.ticktick_redirect_uri,
|
||||
"ticktick_oauth_notice": None,
|
||||
"ticktick_oauth_error": None,
|
||||
},
|
||||
status_code=status_code,
|
||||
)
|
||||
@@ -1,240 +0,0 @@
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
from fastapi import APIRouter, Depends, Request, status
|
||||
from fastapi.responses import HTMLResponse, RedirectResponse, Response
|
||||
from fastapi.templating import Jinja2Templates
|
||||
|
||||
from app.config import Settings, get_settings
|
||||
from app.dependencies import get_app_settings, get_db, get_current_auth_session
|
||||
from app.services.auth import AuthenticatedSession
|
||||
from app.services.config_page import (
|
||||
ConfigSaveError,
|
||||
build_config_sections,
|
||||
is_ticktick_oauth_ready,
|
||||
save_config_updates,
|
||||
)
|
||||
from app.services.email import EmailConfigurationError, EmailDeliveryError, is_smtp_ready, send_smtp_test_email
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
templates = Jinja2Templates(directory=str(Path(__file__).resolve().parents[2] / "templates"))
|
||||
router = APIRouter(tags=["pages"])
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _ticktick_oauth_notice(status_value: str | None) -> tuple[str | None, str | None]:
|
||||
if status_value == "success":
|
||||
return "TickTick authorization completed successfully.", None
|
||||
if status_value == "invalid-state":
|
||||
return None, "TickTick authorization failed due to invalid OAuth state. Start the flow again."
|
||||
if status_value == "invalid-callback":
|
||||
return None, "TickTick authorization callback was missing required parameters."
|
||||
if status_value == "failed":
|
||||
return None, "TickTick authorization failed. Check server logs for the provider response and verify TickTick app credentials and redirect URI."
|
||||
return None, None
|
||||
|
||||
|
||||
def _smtp_test_notice(status_value: str | None) -> tuple[str | None, str | None]:
|
||||
if status_value == "success":
|
||||
return "SMTP test email sent successfully.", None
|
||||
if status_value == "config-error":
|
||||
return None, "SMTP test failed. Check required SMTP settings before sending a test email."
|
||||
if status_value == "failed":
|
||||
return None, "SMTP test failed. Check saved SMTP settings and server reachability."
|
||||
return None, None
|
||||
|
||||
|
||||
def _build_config_context(
|
||||
*,
|
||||
auth_db_session: Session,
|
||||
settings: Settings,
|
||||
current_auth: AuthenticatedSession,
|
||||
config_saved: bool,
|
||||
config_error: str | None,
|
||||
password_change_error: str | None,
|
||||
ticktick_oauth_notice: str | None,
|
||||
ticktick_oauth_error: str | None,
|
||||
smtp_test_notice: str | None,
|
||||
smtp_test_error: str | None,
|
||||
) -> dict[str, object]:
|
||||
return {
|
||||
"app_name": settings.app_name,
|
||||
"app_env": settings.app_env,
|
||||
"current_username": current_auth.user.username,
|
||||
"csrf_token": current_auth.session.csrf_token,
|
||||
"force_password_change": current_auth.user.force_password_change,
|
||||
"password_change_error": password_change_error,
|
||||
"config_error": config_error,
|
||||
"config_saved": config_saved,
|
||||
"config_sections": build_config_sections(auth_db_session, settings),
|
||||
"ticktick_oauth_ready": is_ticktick_oauth_ready(settings),
|
||||
"ticktick_redirect_uri": settings.ticktick_redirect_uri,
|
||||
"ticktick_oauth_notice": ticktick_oauth_notice,
|
||||
"ticktick_oauth_error": ticktick_oauth_error,
|
||||
"smtp_test_ready": is_smtp_ready(settings),
|
||||
"smtp_test_notice": smtp_test_notice,
|
||||
"smtp_test_error": smtp_test_error,
|
||||
}
|
||||
|
||||
|
||||
@router.get("/", response_class=HTMLResponse)
|
||||
def home(
|
||||
request: Request,
|
||||
current_auth: AuthenticatedSession | None = Depends(get_current_auth_session),
|
||||
) -> RedirectResponse:
|
||||
if current_auth is None:
|
||||
return RedirectResponse(url="/login", status_code=status.HTTP_303_SEE_OTHER)
|
||||
return RedirectResponse(url="/config", status_code=status.HTTP_303_SEE_OTHER)
|
||||
|
||||
|
||||
@router.get("/admin", response_class=HTMLResponse)
|
||||
def admin_redirect(
|
||||
request: Request,
|
||||
current_auth: AuthenticatedSession | None = Depends(get_current_auth_session),
|
||||
) -> RedirectResponse:
|
||||
if current_auth is None:
|
||||
return RedirectResponse(url="/login", status_code=status.HTTP_303_SEE_OTHER)
|
||||
return RedirectResponse(url="/config", status_code=status.HTTP_303_SEE_OTHER)
|
||||
|
||||
|
||||
@router.get("/config", response_class=HTMLResponse)
|
||||
def config_page(
|
||||
request: Request,
|
||||
auth_db_session: Session = Depends(get_db),
|
||||
settings: Settings = Depends(get_app_settings),
|
||||
current_auth: AuthenticatedSession | None = Depends(get_current_auth_session),
|
||||
) -> Response:
|
||||
if current_auth is None:
|
||||
return RedirectResponse(url="/login", status_code=status.HTTP_303_SEE_OTHER)
|
||||
|
||||
ticktick_oauth_notice, ticktick_oauth_error = _ticktick_oauth_notice(
|
||||
request.query_params.get("ticktick_oauth")
|
||||
)
|
||||
smtp_test_notice, smtp_test_error = _smtp_test_notice(request.query_params.get("smtp_test"))
|
||||
context = _build_config_context(
|
||||
auth_db_session=auth_db_session,
|
||||
settings=settings,
|
||||
current_auth=current_auth,
|
||||
config_saved=request.query_params.get("saved") == "1",
|
||||
config_error=None,
|
||||
password_change_error=None,
|
||||
ticktick_oauth_notice=ticktick_oauth_notice,
|
||||
ticktick_oauth_error=ticktick_oauth_error,
|
||||
smtp_test_notice=smtp_test_notice,
|
||||
smtp_test_error=smtp_test_error,
|
||||
)
|
||||
return templates.TemplateResponse(request, "config.html", context)
|
||||
|
||||
|
||||
@router.post("/config", response_class=HTMLResponse)
|
||||
async def config_submit(
|
||||
request: Request,
|
||||
auth_db_session: Session = Depends(get_db),
|
||||
settings: Settings = Depends(get_app_settings),
|
||||
current_auth: AuthenticatedSession | None = Depends(get_current_auth_session),
|
||||
) -> Response:
|
||||
if current_auth is None:
|
||||
return RedirectResponse(url="/login", status_code=status.HTTP_303_SEE_OTHER)
|
||||
|
||||
form = await request.form()
|
||||
csrf_token = form.get("csrf_token")
|
||||
if csrf_token != current_auth.session.csrf_token:
|
||||
logger.warning("Rejected config update due to CSRF validation failure")
|
||||
context = _build_config_context(
|
||||
auth_db_session=auth_db_session,
|
||||
settings=settings,
|
||||
current_auth=current_auth,
|
||||
config_saved=False,
|
||||
config_error="invalid config update request",
|
||||
password_change_error=None,
|
||||
ticktick_oauth_notice=None,
|
||||
ticktick_oauth_error=None,
|
||||
smtp_test_notice=None,
|
||||
smtp_test_error=None,
|
||||
)
|
||||
return templates.TemplateResponse(
|
||||
request,
|
||||
"config.html",
|
||||
context,
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
)
|
||||
|
||||
try:
|
||||
save_config_updates(auth_db_session, dict(form), settings)
|
||||
except ConfigSaveError:
|
||||
logger.warning("Rejected config update due to invalid submitted values")
|
||||
refreshed_settings = get_settings()
|
||||
context = _build_config_context(
|
||||
auth_db_session=auth_db_session,
|
||||
settings=refreshed_settings,
|
||||
current_auth=current_auth,
|
||||
config_saved=False,
|
||||
config_error="invalid config submission",
|
||||
password_change_error=None,
|
||||
ticktick_oauth_notice=None,
|
||||
ticktick_oauth_error=None,
|
||||
smtp_test_notice=None,
|
||||
smtp_test_error=None,
|
||||
)
|
||||
return templates.TemplateResponse(
|
||||
request,
|
||||
"config.html",
|
||||
context,
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
)
|
||||
|
||||
return RedirectResponse(url="/config?saved=1", status_code=status.HTTP_303_SEE_OTHER)
|
||||
|
||||
|
||||
@router.post("/config/smtp/test", response_class=HTMLResponse)
|
||||
async def smtp_test_submit(
|
||||
request: Request,
|
||||
auth_db_session: Session = Depends(get_db),
|
||||
settings: Settings = Depends(get_app_settings),
|
||||
current_auth: AuthenticatedSession | None = Depends(get_current_auth_session),
|
||||
) -> Response:
|
||||
if current_auth is None:
|
||||
return RedirectResponse(url="/login", status_code=status.HTTP_303_SEE_OTHER)
|
||||
|
||||
form = await request.form()
|
||||
csrf_token = form.get("csrf_token")
|
||||
if csrf_token != current_auth.session.csrf_token:
|
||||
logger.warning("Rejected SMTP test due to CSRF validation failure")
|
||||
context = _build_config_context(
|
||||
auth_db_session=auth_db_session,
|
||||
settings=settings,
|
||||
current_auth=current_auth,
|
||||
config_saved=False,
|
||||
config_error=None,
|
||||
password_change_error=None,
|
||||
ticktick_oauth_notice=None,
|
||||
ticktick_oauth_error=None,
|
||||
smtp_test_notice=None,
|
||||
smtp_test_error="invalid SMTP test request",
|
||||
)
|
||||
return templates.TemplateResponse(
|
||||
request,
|
||||
"config.html",
|
||||
context,
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
)
|
||||
|
||||
try:
|
||||
send_smtp_test_email(settings)
|
||||
except EmailConfigurationError as exc:
|
||||
logger.warning("SMTP test email rejected due to configuration: %s", exc)
|
||||
return RedirectResponse(
|
||||
url="/config?smtp_test=config-error",
|
||||
status_code=status.HTTP_303_SEE_OTHER,
|
||||
)
|
||||
except EmailDeliveryError as exc:
|
||||
logger.warning("SMTP test email failed: %s", exc)
|
||||
return RedirectResponse(
|
||||
url="/config?smtp_test=failed",
|
||||
status_code=status.HTTP_303_SEE_OTHER,
|
||||
)
|
||||
|
||||
return RedirectResponse(
|
||||
url="/config?smtp_test=success",
|
||||
status_code=status.HTTP_303_SEE_OTHER,
|
||||
)
|
||||
+56
-1
@@ -1,7 +1,8 @@
|
||||
from functools import lru_cache
|
||||
from pathlib import Path
|
||||
import re
|
||||
|
||||
from pydantic import computed_field
|
||||
from pydantic import computed_field, field_validator
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
@@ -37,6 +38,43 @@ class Settings(BaseSettings):
|
||||
auth_session_cookie_name: str = "home_automation_session"
|
||||
auth_session_ttl_hours: int = 12
|
||||
auth_cookie_secure_override: bool | None = True
|
||||
auth_login_throttle_enabled: bool = True
|
||||
auth_trust_forwarded_for: bool = False
|
||||
auth_totp_issuer: str = "" # defaults to app_name when empty
|
||||
|
||||
# Modbus polling — global kill-switch (CONFIG_FIELDS registered in T08).
|
||||
# Off by default: polling is opt-in so a fresh deploy writes no modbus_reading
|
||||
# rows until the admin explicitly enables it (matches mqtt/discovery defaults).
|
||||
modbus_polling_enabled: bool = False
|
||||
|
||||
# MQTT broker connection (T08 wires into CONFIG_FIELDS/UI; T10 builds the client).
|
||||
mqtt_enabled: bool = False
|
||||
mqtt_broker_host: str = ""
|
||||
mqtt_broker_port: int = 1883
|
||||
mqtt_username: str = ""
|
||||
mqtt_password: str = ""
|
||||
mqtt_tls_enabled: bool = False
|
||||
mqtt_client_id: str = "home-automation"
|
||||
|
||||
# Home Assistant MQTT Discovery (T08 wires into CONFIG_FIELDS/UI; T11 does publishing).
|
||||
ha_discovery_enabled: bool = False
|
||||
ha_discovery_prefix: str = "homeassistant"
|
||||
# State/availability topics use a separate prefix so they live outside the HA discovery
|
||||
# namespace. HA still receives state via the state_topic declared in the discovery config.
|
||||
ha_state_topic_prefix: str = "home_automation"
|
||||
|
||||
# DSMR smart-meter ingest via MQTT (M6; default off — opt-in).
|
||||
# Subscribe to dsmr_mqtt_topic when dsmr_ingest_enabled=True.
|
||||
dsmr_ingest_enabled: bool = False
|
||||
dsmr_mqtt_topic: str = "dsmr/json"
|
||||
dsmr_sample_interval_s: int = 10
|
||||
# DSMR dual-tariff topic: publishes "1" (dal/off-peak) or "2" (normal/peak).
|
||||
# Empty string = do not subscribe (tariff-aware pricing disabled).
|
||||
dsmr_tariff_topic: str = "dsmr/meter-stats/electricity_tariff"
|
||||
|
||||
# Tibber dynamic pricing credentials (M6; only used when active contract kind=tibber).
|
||||
tibber_api_token: str = ""
|
||||
tibber_home_id: str = ""
|
||||
|
||||
model_config = SettingsConfigDict(
|
||||
env_file=".env",
|
||||
@@ -45,6 +83,17 @@ class Settings(BaseSettings):
|
||||
extra="ignore",
|
||||
)
|
||||
|
||||
@field_validator("mqtt_client_id", mode="before")
|
||||
@classmethod
|
||||
def validate_mqtt_client_id(cls, value: object) -> str:
|
||||
"""Normalize a broker-safe base client identity used by every MQTT client."""
|
||||
if not isinstance(value, str):
|
||||
raise ValueError("MQTT client ID must be a string")
|
||||
normalized = value.strip()
|
||||
if not re.fullmatch(r"[A-Za-z0-9][A-Za-z0-9_-]{0,63}", normalized):
|
||||
raise ValueError("MQTT client ID must be a non-empty ASCII slug")
|
||||
return normalized
|
||||
|
||||
@computed_field
|
||||
@property
|
||||
def is_development(self) -> bool:
|
||||
@@ -86,6 +135,12 @@ class Settings(BaseSettings):
|
||||
return self.auth_cookie_secure_override
|
||||
return not self.is_development
|
||||
|
||||
@computed_field
|
||||
@property
|
||||
def effective_totp_issuer(self) -> str:
|
||||
"""The issuer label shown in Authenticator apps. Falls back to app_name."""
|
||||
return self.auth_totp_issuer.strip() or self.app_name
|
||||
|
||||
|
||||
@lru_cache
|
||||
def get_settings() -> Settings:
|
||||
|
||||
@@ -28,6 +28,7 @@ def _get_engine(database_url: str) -> Engine:
|
||||
def _enable_sqlite_wal(dbapi_connection, _connection_record):
|
||||
cursor = dbapi_connection.cursor()
|
||||
cursor.execute("PRAGMA journal_mode=WAL")
|
||||
cursor.execute("PRAGMA foreign_keys=ON")
|
||||
cursor.close()
|
||||
|
||||
return engine
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,171 @@
|
||||
"""Registry and configuration helpers for meter-source integrations.
|
||||
|
||||
The registry is deliberately I/O-free. Workers and HTTP handlers use these
|
||||
helpers to share one config contract without opening a broker or serial port.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
|
||||
SECRET_MASK = ""
|
||||
|
||||
|
||||
class SourceProfileError(ValueError):
|
||||
"""Raised when a source kind or its configuration is invalid."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SourceConfigField:
|
||||
"""One source configuration field and its public metadata."""
|
||||
|
||||
name: str
|
||||
value_type: type
|
||||
default: Any = None
|
||||
required: bool = False
|
||||
secret: bool = False
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MeterSourceProfile:
|
||||
"""A supported source kind's config, capabilities, and channel units."""
|
||||
|
||||
kind: str
|
||||
fields: tuple[SourceConfigField, ...]
|
||||
capabilities: frozenset[str]
|
||||
allowed_units: frozenset[str]
|
||||
|
||||
|
||||
DSMR_MQTT_PROFILE = MeterSourceProfile(
|
||||
kind="dsmr_mqtt",
|
||||
fields=(
|
||||
SourceConfigField("broker_host", str, default=""),
|
||||
SourceConfigField("broker_port", int, default=1883),
|
||||
SourceConfigField("username", str, default="", secret=True),
|
||||
SourceConfigField("password", str, default="", secret=True),
|
||||
SourceConfigField("tls_enabled", bool, default=False),
|
||||
SourceConfigField("topic", str, default="dsmr/json"),
|
||||
SourceConfigField("tariff_topic", str, default="dsmr/meter-stats/electricity_tariff"),
|
||||
SourceConfigField("sample_interval_s", int, default=10),
|
||||
),
|
||||
capabilities=frozenset({"discover", "mqtt_subscribe", "tariff"}),
|
||||
allowed_units=frozenset({"kWh"}),
|
||||
)
|
||||
|
||||
WARMTELINK_SERIAL_PROFILE = MeterSourceProfile(
|
||||
kind="warmtelink_serial",
|
||||
fields=(
|
||||
SourceConfigField("path", str, required=True),
|
||||
SourceConfigField("baudrate", int, default=115200),
|
||||
SourceConfigField("data_bits", int, default=7),
|
||||
SourceConfigField("parity", str, default="N"),
|
||||
SourceConfigField("stop_bits", int, default=1),
|
||||
),
|
||||
capabilities=frozenset({"discover", "read_only_serial"}),
|
||||
allowed_units=frozenset({"GJ", "m³"}),
|
||||
)
|
||||
|
||||
SOURCE_PROFILES: dict[str, MeterSourceProfile] = {
|
||||
DSMR_MQTT_PROFILE.kind: DSMR_MQTT_PROFILE,
|
||||
WARMTELINK_SERIAL_PROFILE.kind: WARMTELINK_SERIAL_PROFILE,
|
||||
}
|
||||
|
||||
|
||||
def get_source_profile(kind: str) -> MeterSourceProfile:
|
||||
"""Return the profile for *kind*, or raise a stable validation error."""
|
||||
try:
|
||||
return SOURCE_PROFILES[kind]
|
||||
except KeyError as exc:
|
||||
raise SourceProfileError(f"Unsupported meter source kind: {kind!r}") from exc
|
||||
|
||||
|
||||
def list_source_profiles() -> list[MeterSourceProfile]:
|
||||
"""Return profiles in deterministic kind order for a future API/UI."""
|
||||
return [SOURCE_PROFILES[kind] for kind in sorted(SOURCE_PROFILES)]
|
||||
|
||||
|
||||
def _check_type(field: SourceConfigField, value: Any) -> None:
|
||||
# bool is a subclass of int; accept it only for explicitly boolean fields.
|
||||
if type(value) is not field.value_type:
|
||||
raise SourceProfileError(
|
||||
f"Config field {field.name!r} must be a {field.value_type.__name__}."
|
||||
)
|
||||
|
||||
|
||||
def _validate_field_value(kind: str, field: SourceConfigField, value: Any) -> None:
|
||||
_check_type(field, value)
|
||||
if field.name == "path" and not value.startswith("/dev/"):
|
||||
raise SourceProfileError("warmtelink_serial config path must start with '/dev/'.")
|
||||
if field.name in {"broker_port", "sample_interval_s", "baudrate"} and value <= 0:
|
||||
raise SourceProfileError(f"Config field {field.name!r} must be greater than zero.")
|
||||
if field.name == "data_bits" and value != 7:
|
||||
raise SourceProfileError("warmtelink_serial data_bits must be 7.")
|
||||
if field.name == "parity" and value != "N":
|
||||
raise SourceProfileError("warmtelink_serial parity must be 'N'.")
|
||||
if field.name == "stop_bits" and value != 1:
|
||||
raise SourceProfileError("warmtelink_serial stop_bits must be 1.")
|
||||
|
||||
|
||||
def validate_source_config(kind: str, config: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Validate a complete config and return it with profile defaults filled.
|
||||
|
||||
Unknown keys are rejected to make configuration additions explicit. Secret
|
||||
masking is intentionally not interpreted here: callers must merge a PATCH
|
||||
with its stored config first.
|
||||
"""
|
||||
profile = get_source_profile(kind)
|
||||
if not isinstance(config, dict):
|
||||
raise SourceProfileError("Source config must be an object.")
|
||||
fields = {field.name: field for field in profile.fields}
|
||||
unknown = set(config) - set(fields)
|
||||
if unknown:
|
||||
raise SourceProfileError(f"Unknown {kind} config field(s): {sorted(unknown)!r}")
|
||||
|
||||
validated: dict[str, Any] = {}
|
||||
for field in profile.fields:
|
||||
if field.name in config:
|
||||
value = config[field.name]
|
||||
elif field.required:
|
||||
raise SourceProfileError(f"Missing required {kind} config field: {field.name!r}")
|
||||
else:
|
||||
value = field.default
|
||||
_validate_field_value(kind, field, value)
|
||||
validated[field.name] = value
|
||||
return validated
|
||||
|
||||
|
||||
def sanitize_source_config(kind: str, config: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Validate and return a response-safe config with secrets masked."""
|
||||
profile = get_source_profile(kind)
|
||||
sanitized = validate_source_config(kind, config)
|
||||
for field in profile.fields:
|
||||
if field.secret:
|
||||
sanitized[field.name] = SECRET_MASK
|
||||
return sanitized
|
||||
|
||||
|
||||
def merge_source_config(kind: str, current: dict[str, Any], patch: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Merge a partial PATCH into stored config, retaining masked secrets.
|
||||
|
||||
An empty secret value is the public response mask and therefore means
|
||||
"keep the old value". New sources use :func:`validate_source_config`
|
||||
instead, so an explicitly empty secret can still be initially configured.
|
||||
"""
|
||||
profile = get_source_profile(kind)
|
||||
current_validated = validate_source_config(kind, current)
|
||||
if not isinstance(patch, dict):
|
||||
raise SourceProfileError("Source config patch must be an object.")
|
||||
fields = {field.name: field for field in profile.fields}
|
||||
unknown = set(patch) - set(fields)
|
||||
if unknown:
|
||||
raise SourceProfileError(f"Unknown {kind} config field(s): {sorted(unknown)!r}")
|
||||
|
||||
merged = dict(current_validated)
|
||||
for name, value in patch.items():
|
||||
field = fields[name]
|
||||
if field.secret and value == SECRET_MASK:
|
||||
continue
|
||||
merged[name] = value
|
||||
return validate_source_config(kind, merged)
|
||||
@@ -0,0 +1 @@
|
||||
"""Modbus integration package — driver, profile loader, and decoder."""
|
||||
@@ -0,0 +1,230 @@
|
||||
"""Modbus TCP driver — thin wrapper around pymodbus.
|
||||
|
||||
This module provides a single public function ``read_blocks`` that performs
|
||||
one or more block reads against a Modbus TCP gateway — using either FC04
|
||||
(Read Input Registers) or FC03 (Read Holding Registers), selected per call
|
||||
via the ``function_code`` argument — and returns a flat
|
||||
``dict[register_address -> 16-bit_value]`` map. The function code comes from
|
||||
the device profile (e.g. SDM120 uses FC04, DDSU666 uses FC03).
|
||||
|
||||
Design decisions
|
||||
----------------
|
||||
- **Only read function codes** (FC03/FC04) are exposed. There is no write path.
|
||||
- float32 decoding uses ``struct.unpack('>f', ...)`` directly rather than the
|
||||
pymodbus ``BinaryPayloadDecoder`` helper, which has had API churn across 3.x
|
||||
releases. Big-endian word order + big-endian byte order means the 4 raw bytes
|
||||
are already in standard network (big-endian) order.
|
||||
- ``ModbusTcpClient`` is used in synchronous mode (``connect()`` / ``close()``).
|
||||
The client is created fresh per call; this keeps the driver stateless and
|
||||
avoids threading concerns at the cost of one TCP handshake per read cycle.
|
||||
APScheduler jobs call this from a thread pool, so stateless is safer.
|
||||
- pymodbus 3.13.x uses ``device_id=`` as the slave-address keyword argument
|
||||
(renamed from ``slave=`` in earlier 3.x releases).
|
||||
|
||||
Exceptions
|
||||
----------
|
||||
``ModbusDriverError``
|
||||
Base exception for all driver-level errors.
|
||||
``ModbusConnectionError``
|
||||
Raised when the TCP connection to the gateway cannot be established or is
|
||||
lost during the read.
|
||||
``ModbusResponseError``
|
||||
Raised when the gateway responds with a Modbus exception frame or when the
|
||||
returned register count does not match the requested count.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import struct
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from pymodbus.client import ModbusTcpClient
|
||||
from pymodbus.exceptions import ConnectionException, ModbusException
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Sequence
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Custom exceptions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ModbusDriverError(RuntimeError):
|
||||
"""Base class for all Modbus driver errors."""
|
||||
|
||||
|
||||
class ModbusConnectionError(ModbusDriverError):
|
||||
"""Cannot connect to (or lost connection with) the Modbus TCP gateway."""
|
||||
|
||||
|
||||
class ModbusResponseError(ModbusDriverError):
|
||||
"""Modbus gateway returned an exception frame or an unexpected response."""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Block definition type alias
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
#: A single read block: ``{"start": int, "count": int}``.
|
||||
Block = dict[str, int]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public API
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def read_blocks(
|
||||
host: str,
|
||||
port: int,
|
||||
unit_id: int,
|
||||
blocks: Sequence[Block],
|
||||
*,
|
||||
function_code: int = 4,
|
||||
timeout: float = 3.0,
|
||||
) -> dict[int, int]:
|
||||
"""Read one or more contiguous register blocks via FC03 or FC04.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
host:
|
||||
Hostname or IP address of the Modbus TCP gateway.
|
||||
port:
|
||||
TCP port (typically 502).
|
||||
unit_id:
|
||||
Modbus slave / unit address (the device's Meter ID for SDM120).
|
||||
blocks:
|
||||
Sequence of ``{"start": int, "count": int}`` dicts describing the
|
||||
contiguous register ranges to read. ``count`` is the number of
|
||||
16-bit registers (not bytes).
|
||||
function_code:
|
||||
Modbus read function code: ``4`` for FC04 (Read Input Registers,
|
||||
default — SDM120) or ``3`` for FC03 (Read Holding Registers —
|
||||
DDSU666 and other devices that expose measurements as holding
|
||||
registers). Comes from the device profile's ``function_code`` field.
|
||||
timeout:
|
||||
TCP connect/read timeout in seconds (default 3 s).
|
||||
|
||||
Returns
|
||||
-------
|
||||
dict[int, int]
|
||||
Flat mapping of ``{register_address: 16-bit_register_value}`` for
|
||||
every register returned across all blocks.
|
||||
|
||||
Raises
|
||||
------
|
||||
ModbusDriverError
|
||||
If ``function_code`` is neither 3 nor 4 (validated before any
|
||||
connection is attempted).
|
||||
ModbusConnectionError
|
||||
If the TCP connection to the gateway fails.
|
||||
ModbusResponseError
|
||||
If the gateway returns a Modbus exception frame or an unexpected
|
||||
number of registers.
|
||||
"""
|
||||
if function_code not in (3, 4):
|
||||
raise ModbusDriverError(
|
||||
f"Unsupported read function code FC{function_code:02d} "
|
||||
f"(only FC03 holding-register and FC04 input-register reads are supported)"
|
||||
)
|
||||
|
||||
client = ModbusTcpClient(host, port=port, timeout=timeout)
|
||||
try:
|
||||
connected = client.connect()
|
||||
if not connected:
|
||||
raise ModbusConnectionError(
|
||||
f"Could not connect to Modbus gateway at {host}:{port}"
|
||||
)
|
||||
|
||||
registers: dict[int, int] = {}
|
||||
for block in blocks:
|
||||
start: int = block["start"]
|
||||
count: int = block["count"]
|
||||
_read_block(client, unit_id, start, count, registers, function_code=function_code)
|
||||
|
||||
return registers
|
||||
|
||||
except ConnectionException as exc:
|
||||
raise ModbusConnectionError(
|
||||
f"Connection to Modbus gateway {host}:{port} failed: {exc}"
|
||||
) from exc
|
||||
except ModbusException as exc:
|
||||
raise ModbusResponseError(f"Modbus protocol error: {exc}") from exc
|
||||
finally:
|
||||
client.close()
|
||||
|
||||
|
||||
def _read_block(
|
||||
client: ModbusTcpClient,
|
||||
unit_id: int,
|
||||
start: int,
|
||||
count: int,
|
||||
result: dict[int, int],
|
||||
*,
|
||||
function_code: int,
|
||||
) -> None:
|
||||
"""Read one block and merge into *result*. Raises on any error.
|
||||
|
||||
``function_code`` is assumed already validated to be 3 or 4 by the caller
|
||||
(``read_blocks``); 3 dispatches FC03 (holding) and 4 dispatches FC04 (input).
|
||||
"""
|
||||
try:
|
||||
if function_code == 3:
|
||||
response = client.read_holding_registers(start, count=count, device_id=unit_id)
|
||||
else: # function_code == 4 (input registers)
|
||||
response = client.read_input_registers(start, count=count, device_id=unit_id)
|
||||
except ConnectionException as exc:
|
||||
raise ModbusConnectionError(
|
||||
f"Lost connection while reading registers 0x{start:04X}+{count}: {exc}"
|
||||
) from exc
|
||||
|
||||
if response.isError():
|
||||
raise ModbusResponseError(
|
||||
f"Modbus exception reading registers 0x{start:04X}+{count}: {response}"
|
||||
)
|
||||
|
||||
regs = response.registers
|
||||
if len(regs) != count:
|
||||
raise ModbusResponseError(
|
||||
f"Expected {count} registers from 0x{start:04X}, got {len(regs)}"
|
||||
)
|
||||
|
||||
for offset, value in enumerate(regs):
|
||||
result[start + offset] = value
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Decoding helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def registers_to_float(hi_reg: int, lo_reg: int) -> float:
|
||||
"""Decode two 16-bit registers to a big-endian IEEE-754 float32.
|
||||
|
||||
The SDM120 (and most Modbus float devices) use big-endian word order:
|
||||
the high 16-bit register comes first, then the low register. Combined
|
||||
with big-endian byte order within each register, the raw 4-byte sequence
|
||||
is standard network-byte-order (big-endian) float32.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
hi_reg:
|
||||
The first (higher-address-range, most-significant) 16-bit register.
|
||||
lo_reg:
|
||||
The second (lower-address-range, least-significant) 16-bit register.
|
||||
|
||||
Returns
|
||||
-------
|
||||
float
|
||||
The decoded IEEE-754 single-precision floating-point value.
|
||||
|
||||
Example
|
||||
-------
|
||||
>>> registers_to_float(0x4366, 0x3334)
|
||||
230.20001220703125
|
||||
"""
|
||||
raw = struct.pack(">HH", hi_reg, lo_reg)
|
||||
return struct.unpack(">f", raw)[0]
|
||||
@@ -0,0 +1,231 @@
|
||||
"""Modbus profile loader, validator, and decoder.
|
||||
|
||||
A *profile* is a YAML file that describes the protocol-level knowledge for
|
||||
a particular Modbus device model:
|
||||
|
||||
- Which registers to read (``blocks`` for efficient bulk reads).
|
||||
- How to decode each quantity (``metrics``: register address, type, unit, …).
|
||||
- Home Assistant metadata per metric (``device_class``, ``ha_component``, …).
|
||||
|
||||
Profiles live in ``app/integrations/modbus/profiles/<name>.yaml`` and are
|
||||
located at runtime relative to *this file* (not CWD), so they work whether
|
||||
the package is installed as a wheel, run in-process, or invoked via
|
||||
``python -m scripts.modbus_cli``.
|
||||
|
||||
Design notes
|
||||
------------
|
||||
- **Purely data + functions** — no abstract base classes or inheritance.
|
||||
- ``ModbusProfile`` and ``MetricSpec`` are Pydantic models: they validate on
|
||||
construction and raise ``pydantic.ValidationError`` for malformed YAML.
|
||||
- ``decode`` uses ``driver.registers_to_float`` for float32 quantities.
|
||||
Unsupported types are silently skipped with a warning (future-proofing for
|
||||
int16/uint16 etc.).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import yaml
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from app.integrations.modbus.driver import registers_to_float
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Directory containing the YAML profiles, located relative to *this* file.
|
||||
_PROFILES_DIR = Path(__file__).parent / "profiles"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic models for profile validation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MetricSpec(BaseModel):
|
||||
"""Specification for a single measurable quantity within a profile."""
|
||||
|
||||
key: str
|
||||
"""Stable identifier used as the key in decoded payloads (e.g. ``"voltage"``)."""
|
||||
|
||||
address: int
|
||||
"""Starting Modbus register address (0-based, e.g. ``0x0000`` for voltage)."""
|
||||
|
||||
type: str
|
||||
"""Encoding type — currently only ``"float32"`` is supported."""
|
||||
|
||||
unit: str
|
||||
"""Physical unit string (e.g. ``"V"``, ``"A"``, ``"kWh"``). Empty string for dimensionless."""
|
||||
|
||||
device_class: str
|
||||
"""Home Assistant device class (e.g. ``"voltage"``, ``"power"``, ``"energy"``)."""
|
||||
|
||||
ha_component: str
|
||||
"""Home Assistant component type (e.g. ``"sensor"``)."""
|
||||
|
||||
state_class: Optional[str] = None
|
||||
"""Home Assistant state class, if applicable (e.g. ``"total_increasing"``)."""
|
||||
|
||||
|
||||
class BlockSpec(BaseModel):
|
||||
"""A contiguous register block for bulk reading."""
|
||||
|
||||
start: int
|
||||
"""First register address in the block."""
|
||||
|
||||
count: int
|
||||
"""Number of 16-bit registers to read."""
|
||||
|
||||
|
||||
class ModbusProfile(BaseModel):
|
||||
"""Complete description of a Modbus device's protocol-level knowledge."""
|
||||
|
||||
name: str
|
||||
"""Short identifier matching the YAML file name (e.g. ``"sdm120"``)."""
|
||||
|
||||
description: str
|
||||
"""Human-readable description of the device."""
|
||||
|
||||
function_code: int
|
||||
"""Modbus function code for reading measurements (4 = FC04 input registers)."""
|
||||
|
||||
word_order: str
|
||||
"""Word (register) order: ``"big"`` means high register first."""
|
||||
|
||||
byte_order: str
|
||||
"""Byte order within each register: ``"big"`` means standard network order."""
|
||||
|
||||
blocks: list[BlockSpec]
|
||||
"""Register blocks to read in bulk, reducing Modbus transaction count."""
|
||||
|
||||
metrics: list[MetricSpec]
|
||||
"""List of individual measurable quantities and their decoding rules."""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public API
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ProfileNotFoundError(FileNotFoundError):
|
||||
"""Raised when the requested profile YAML file does not exist."""
|
||||
|
||||
|
||||
class ProfileValidationError(ValueError):
|
||||
"""Raised when a profile YAML file fails Pydantic validation."""
|
||||
|
||||
|
||||
def load_profile(name: str) -> ModbusProfile:
|
||||
"""Load and validate a Modbus device profile by name.
|
||||
|
||||
The profile file is expected at
|
||||
``app/integrations/modbus/profiles/<name>.yaml``.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name:
|
||||
Profile name without extension (e.g. ``"sdm120"``).
|
||||
|
||||
Returns
|
||||
-------
|
||||
ModbusProfile
|
||||
Validated profile model.
|
||||
|
||||
Raises
|
||||
------
|
||||
ProfileNotFoundError
|
||||
If ``profiles/<name>.yaml`` does not exist.
|
||||
ProfileValidationError
|
||||
If the YAML is syntactically valid but the content fails schema
|
||||
validation (missing required fields, wrong types, etc.).
|
||||
"""
|
||||
path = _PROFILES_DIR / f"{name}.yaml"
|
||||
if not path.exists():
|
||||
raise ProfileNotFoundError(
|
||||
f"Modbus profile '{name}' not found (looked for {path})"
|
||||
)
|
||||
|
||||
with path.open("r", encoding="utf-8") as fh:
|
||||
raw = yaml.safe_load(fh)
|
||||
|
||||
try:
|
||||
profile = ModbusProfile.model_validate(raw)
|
||||
except ValidationError as exc:
|
||||
raise ProfileValidationError(
|
||||
f"Profile '{name}' failed validation: {exc}"
|
||||
) from exc
|
||||
|
||||
return profile
|
||||
|
||||
|
||||
def list_profiles() -> list[tuple[str, str]]:
|
||||
"""Return ``(name, description)`` pairs for all available profiles.
|
||||
|
||||
Scans ``app/integrations/modbus/profiles/*.yaml``, loads each, and
|
||||
returns a sorted list of ``(name, description)`` tuples. Profiles that
|
||||
fail to load are skipped with a warning.
|
||||
"""
|
||||
results: list[tuple[str, str]] = []
|
||||
if not _PROFILES_DIR.exists():
|
||||
return results
|
||||
|
||||
for path in sorted(_PROFILES_DIR.glob("*.yaml")):
|
||||
name = path.stem
|
||||
try:
|
||||
profile = load_profile(name)
|
||||
results.append((profile.name, profile.description))
|
||||
except (ProfileNotFoundError, ProfileValidationError, Exception) as exc:
|
||||
logger.warning("Skipping malformed profile '%s': %s", name, exc)
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def decode(profile: ModbusProfile, registers: dict[int, int]) -> dict[str, float]:
|
||||
"""Decode raw register values into a mapping of ``{key: float_value}``.
|
||||
|
||||
For each metric in *profile*, the function looks up the two registers at
|
||||
``address`` and ``address + 1`` in *registers* and decodes them as a
|
||||
big-endian float32 (high register first).
|
||||
|
||||
Metrics whose registers are absent from *registers* are skipped (not
|
||||
an error — this can happen if a block was not requested or a device does
|
||||
not populate that register). Metrics with unsupported ``type`` values
|
||||
are also skipped.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
profile:
|
||||
Validated ``ModbusProfile`` describing the device.
|
||||
registers:
|
||||
Mapping of ``{register_address: 16-bit_value}`` as returned by
|
||||
``driver.read_blocks``.
|
||||
|
||||
Returns
|
||||
-------
|
||||
dict[str, float]
|
||||
Decoded engineering values keyed by each metric's ``key`` field.
|
||||
"""
|
||||
result: dict[str, float] = {}
|
||||
for metric in profile.metrics:
|
||||
if metric.type == "float32":
|
||||
hi_addr = metric.address
|
||||
lo_addr = metric.address + 1
|
||||
if hi_addr not in registers or lo_addr not in registers:
|
||||
logger.debug(
|
||||
"Skipping metric '%s': registers 0x%04X/0x%04X not available",
|
||||
metric.key,
|
||||
hi_addr,
|
||||
lo_addr,
|
||||
)
|
||||
continue
|
||||
value = registers_to_float(registers[hi_addr], registers[lo_addr])
|
||||
result[metric.key] = value
|
||||
else:
|
||||
logger.warning(
|
||||
"Metric '%s' has unsupported type '%s', skipping",
|
||||
metric.key,
|
||||
metric.type,
|
||||
)
|
||||
return result
|
||||
@@ -0,0 +1,82 @@
|
||||
name: ddsu666
|
||||
description: CHINT DDSU666 single-phase smart meter
|
||||
function_code: 3 # holding registers (FC03) — DDSU666 has NO input registers (no FC04)
|
||||
word_order: big # high register first — ASSUMED; verify with a known voltage reading
|
||||
byte_order: big # high byte first within each register (confirmed by manual CRC example)
|
||||
|
||||
# NOTE: the manual gives no worked float-decode example, so word_order is a best-guess
|
||||
# (standard big-endian, high register first, matching sdm120). After wiring the meter,
|
||||
# read 0x2000 (voltage) — it should decode to ~230 V. If it decodes to garbage, the
|
||||
# device uses the opposite word order and this profile (and the decoder) need adjusting.
|
||||
# Byte order is confirmed big-endian from the manual (Appendix A, Table A.4: 0x1388 -> 13 88).
|
||||
|
||||
blocks:
|
||||
# Instantaneous quantities 0x2000–0x200F: voltage, current, P, Q, (rsv), PF, (rsv), Freq.
|
||||
# 16 contiguous registers — single bulk read. (DDSU666 manual Table 9.)
|
||||
- { start: 0x2000, count: 0x0010 }
|
||||
# Active energy — import (0x4000) and export (0x400A) read as two small blocks rather
|
||||
# than one span, to avoid touching the undocumented/reserved 0x4002–0x4009 gap.
|
||||
- { start: 0x4000, count: 0x0002 }
|
||||
- { start: 0x400A, count: 0x0002 }
|
||||
|
||||
metrics:
|
||||
# Addresses are the raw Modbus protocol addresses (hex) from DDSU666 manual Table 9,
|
||||
# read via FC03. Each float32 occupies two consecutive 16-bit registers.
|
||||
|
||||
- key: voltage
|
||||
address: 0x2000 # U — Voltage (V)
|
||||
type: float32
|
||||
unit: "V"
|
||||
device_class: voltage
|
||||
ha_component: sensor
|
||||
|
||||
- key: current
|
||||
address: 0x2002 # I — Current (A)
|
||||
type: float32
|
||||
unit: "A"
|
||||
device_class: current
|
||||
ha_component: sensor
|
||||
|
||||
- key: active_power
|
||||
address: 0x2004 # P — Active power. Manual unit is kW (NOT W like sdm120).
|
||||
type: float32
|
||||
unit: "kW"
|
||||
device_class: power
|
||||
ha_component: sensor
|
||||
|
||||
- key: reactive_power
|
||||
address: 0x2006 # Q — Reactive power (kvar)
|
||||
type: float32
|
||||
unit: "kvar"
|
||||
device_class: reactive_power
|
||||
ha_component: sensor
|
||||
|
||||
- key: power_factor
|
||||
address: 0x200A # PF — Power factor (dimensionless)
|
||||
type: float32
|
||||
unit: ""
|
||||
device_class: power_factor
|
||||
ha_component: sensor
|
||||
|
||||
- key: frequency
|
||||
address: 0x200E # Freq — Frequency (Hz)
|
||||
type: float32
|
||||
unit: "Hz"
|
||||
device_class: frequency
|
||||
ha_component: sensor
|
||||
|
||||
- key: import_energy
|
||||
address: 0x4000 # Ep — positive/forward active energy (kWh)
|
||||
type: float32
|
||||
unit: "kWh"
|
||||
device_class: energy
|
||||
state_class: total_increasing
|
||||
ha_component: sensor
|
||||
|
||||
- key: export_energy
|
||||
address: 0x400A # -Ep — reverse active energy (kWh)
|
||||
type: float32
|
||||
unit: "kWh"
|
||||
device_class: energy
|
||||
state_class: total_increasing
|
||||
ha_component: sensor
|
||||
@@ -0,0 +1,74 @@
|
||||
name: sdm120
|
||||
description: Eastron SDM120 single-phase energy meter
|
||||
function_code: 4 # input registers (FC04)
|
||||
word_order: big # high register first (big-endian word order)
|
||||
byte_order: big # high byte first within each register
|
||||
|
||||
blocks:
|
||||
# Registers 30001–30095 (0x0000–0x005F): voltage through max export demand
|
||||
- { start: 0x0000, count: 0x0060 }
|
||||
# Registers 30343–30346 (0x0156–0x0159): total active/reactive energy
|
||||
- { start: 0x0156, count: 0x0004 }
|
||||
|
||||
metrics:
|
||||
# Core measurements — addresses from SDM120 Modbus Protocol §4 (FC04 input registers)
|
||||
# Each float32 occupies two consecutive 16-bit registers; address = Hex start in §4.
|
||||
|
||||
- key: voltage
|
||||
address: 0x0000 # register 30001 — Voltage (V)
|
||||
type: float32
|
||||
unit: "V"
|
||||
device_class: voltage
|
||||
ha_component: sensor
|
||||
|
||||
- key: current
|
||||
address: 0x0006 # register 30007 — Current (A)
|
||||
type: float32
|
||||
unit: "A"
|
||||
device_class: current
|
||||
ha_component: sensor
|
||||
|
||||
- key: active_power
|
||||
address: 0x000C # register 30013 — Active power (W)
|
||||
type: float32
|
||||
unit: "W"
|
||||
device_class: power
|
||||
ha_component: sensor
|
||||
|
||||
- key: power_factor
|
||||
address: 0x001E # register 30031 — Power factor (dimensionless)
|
||||
type: float32
|
||||
unit: ""
|
||||
device_class: power_factor
|
||||
ha_component: sensor
|
||||
|
||||
- key: frequency
|
||||
address: 0x0046 # register 30071 — Frequency (Hz)
|
||||
type: float32
|
||||
unit: "Hz"
|
||||
device_class: frequency
|
||||
ha_component: sensor
|
||||
|
||||
- key: import_energy
|
||||
address: 0x0048 # register 30073 — Import active energy (kWh)
|
||||
type: float32
|
||||
unit: "kWh"
|
||||
device_class: energy
|
||||
state_class: total_increasing
|
||||
ha_component: sensor
|
||||
|
||||
- key: export_energy
|
||||
address: 0x004A # register 30075 — Export active energy (kWh)
|
||||
type: float32
|
||||
unit: "kWh"
|
||||
device_class: energy
|
||||
state_class: total_increasing
|
||||
ha_component: sensor
|
||||
|
||||
- key: total_energy
|
||||
address: 0x0156 # register 30343 — Total active energy (kWh)
|
||||
type: float32
|
||||
unit: "kWh"
|
||||
device_class: energy
|
||||
state_class: total_increasing
|
||||
ha_component: sensor
|
||||
@@ -0,0 +1,571 @@
|
||||
"""MQTT client integration (paho-mqtt 2.x).
|
||||
|
||||
Provides :class:`MqttManager`: a long-lived MQTT client that runs paho's
|
||||
loop in a background thread. Designed to be started in FastAPI's lifespan
|
||||
and shared across the application as a module-level singleton (``mqtt_manager``).
|
||||
|
||||
Key design decisions
|
||||
--------------------
|
||||
- **paho 2.x API**: ``Client`` requires ``CallbackAPIVersion.VERSION2`` as the
|
||||
first positional argument. VERSION2 on_connect callback receives
|
||||
``(client, userdata, connect_flags, reason_code, properties)`` — not the
|
||||
older RC int.
|
||||
- **Background thread**: ``loop_start()`` spawns a daemon thread; ``loop_stop()``
|
||||
joins it cleanly on shutdown.
|
||||
- **Never crash the process**: all paho operations are wrapped in try/except;
|
||||
connection failure is logged but does not propagate to the caller.
|
||||
- **Password safety**: the MQTT password is never passed to ``logger`` calls.
|
||||
- **Reconnect**: ``reconnect(settings)`` tears down the old client and
|
||||
establishes a fresh connection with the new settings. Callers (e.g. PUT
|
||||
/api/config) call this after saving MQTT-related config values.
|
||||
- **Subscribe**: ``subscribe(topic, handler)`` registers a topic → handler
|
||||
mapping. Subscriptions are re-established automatically on reconnect.
|
||||
The ``on_message`` callback dispatches incoming payloads to the registered
|
||||
handler; handler exceptions are caught and logged so they never crash the
|
||||
paho loop thread or the network connection.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import paho.mqtt.client as mqtt
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from app.config import Settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# MQTT settings keys that, when changed, should trigger a reconnect.
|
||||
MQTT_SETTINGS_KEYS = {
|
||||
"mqtt_enabled",
|
||||
"mqtt_broker_host",
|
||||
"mqtt_broker_port",
|
||||
"mqtt_username",
|
||||
"mqtt_password",
|
||||
"mqtt_tls_enabled",
|
||||
"mqtt_client_id",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class _SourceClientState:
|
||||
"""One installed source generation and its in-flight callback count."""
|
||||
|
||||
client: mqtt.Client
|
||||
generation: int
|
||||
in_flight: int = 0
|
||||
|
||||
|
||||
def _is_configured(settings: Settings) -> bool:
|
||||
"""Return True if MQTT is enabled *and* the broker host is set."""
|
||||
return bool(settings.mqtt_enabled and settings.mqtt_broker_host)
|
||||
|
||||
|
||||
def mqtt_source_client_id(base_client_id: str, source_id: int) -> str:
|
||||
"""Return a stable, deployment-scoped identity for one DSMR source."""
|
||||
return f"{base_client_id}-dsmr-source-{source_id}"
|
||||
|
||||
|
||||
def mqtt_test_client_id(base_client_id: str) -> str:
|
||||
"""Return a transient test identity that cannot evict a long-lived client."""
|
||||
return f"{base_client_id}-test"
|
||||
|
||||
|
||||
class MqttManager:
|
||||
"""Long-lived MQTT client wrapper.
|
||||
|
||||
Lifecycle
|
||||
---------
|
||||
1. ``connect(settings)`` — start background loop + connect to broker.
|
||||
2. ``publish(...)`` — publish messages while connected.
|
||||
3. ``disconnect()`` — graceful teardown (called on app shutdown).
|
||||
4. ``reconnect(settings)`` — disconnect then re-connect with new settings
|
||||
(called after saving updated MQTT configuration).
|
||||
|
||||
When MQTT is not enabled or not configured, all methods are no-ops.
|
||||
Connection failures are caught and logged; they do not raise.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._client: mqtt.Client | None = None
|
||||
self._lock = threading.RLock()
|
||||
# Source replacement/removal is serialized independently from callback
|
||||
# bookkeeping. In particular, paho's loop_stop() joins its network
|
||||
# thread, whose callback completion also needs ``_lock``.
|
||||
self._source_lifecycle_lock = threading.Lock()
|
||||
self._source_idle = threading.Condition(self._lock)
|
||||
self._connected = False
|
||||
# topic → handler registry; persists across reconnects so subscriptions
|
||||
# are automatically re-established when the client reconnects.
|
||||
self._subscriptions: dict[str, Callable[[bytes], None]] = {}
|
||||
# DSMR sources are independent connections: their credentials and TLS
|
||||
# configuration belong to MeterSource.config, not app_config.
|
||||
self._source_clients: dict[int, mqtt.Client] = {}
|
||||
self._source_subscriptions: dict[int, dict[str, Callable[[bytes], None]]] = {}
|
||||
self._source_connected: set[int] = set()
|
||||
# Each replacement gets a distinct identity. A paho callback can run
|
||||
# after its client was stopped, so source id alone is not sufficient.
|
||||
self._source_generations: dict[int, int] = {}
|
||||
self._source_states: dict[int, _SourceClientState] = {}
|
||||
self._next_source_generation = 0
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Public properties
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@property
|
||||
def is_connected(self) -> bool:
|
||||
"""True if the underlying paho client is currently connected."""
|
||||
return self._connected
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Public interface
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def is_configured(self, settings: Settings) -> bool:
|
||||
"""Return True if MQTT is enabled and broker host is configured."""
|
||||
return _is_configured(settings)
|
||||
|
||||
def connect(self, settings: Settings) -> None:
|
||||
"""Connect to the MQTT broker and start the background loop thread.
|
||||
|
||||
No-op if MQTT is not enabled or broker host is not set.
|
||||
Connection errors are logged but do not raise.
|
||||
"""
|
||||
if not _is_configured(settings):
|
||||
logger.debug("MQTT not configured or not enabled — skipping connect.")
|
||||
return
|
||||
|
||||
with self._lock:
|
||||
self._start_client(settings)
|
||||
|
||||
def disconnect(self) -> None:
|
||||
"""Disconnect from the broker and stop the background loop thread.
|
||||
|
||||
No-op if no client is active.
|
||||
"""
|
||||
with self._lock:
|
||||
self._stop_client()
|
||||
with self._source_lifecycle_lock:
|
||||
with self._lock:
|
||||
source_ids = list(self._source_clients)
|
||||
for source_id in source_ids:
|
||||
self._stop_source_client(source_id)
|
||||
|
||||
def reconnect(self, settings: Settings) -> None:
|
||||
"""Disconnect the current client (if any) and reconnect with *settings*.
|
||||
|
||||
Call this after saving updated MQTT configuration values.
|
||||
No-op if MQTT is not enabled or broker host is not set.
|
||||
"""
|
||||
with self._lock:
|
||||
self._stop_client()
|
||||
self.connect(settings)
|
||||
|
||||
def publish(
|
||||
self,
|
||||
topic: str,
|
||||
payload: str | bytes | None,
|
||||
*,
|
||||
retain: bool = False,
|
||||
qos: int = 0,
|
||||
) -> None:
|
||||
"""Publish a message to *topic*.
|
||||
|
||||
If the client is not connected the call is silently skipped.
|
||||
Errors are logged but do not raise.
|
||||
"""
|
||||
with self._lock:
|
||||
client = self._client
|
||||
|
||||
if client is None or not self._connected:
|
||||
logger.debug("MQTT publish skipped — not connected (topic=%s).", topic)
|
||||
return
|
||||
|
||||
try:
|
||||
client.publish(topic, payload=payload, qos=qos, retain=retain)
|
||||
except Exception:
|
||||
logger.exception("MQTT publish error (topic=%s).", topic)
|
||||
|
||||
def subscribe(self, topic: str, handler: Callable[[bytes], None]) -> None:
|
||||
"""Register *handler* to be called when a message arrives on *topic*.
|
||||
|
||||
If the client is already connected, the subscription is sent to the
|
||||
broker immediately. Otherwise it is queued and will be established
|
||||
the next time ``_on_connect`` fires (including after a reconnect).
|
||||
|
||||
Handler exceptions are swallowed by the ``on_message`` dispatcher so
|
||||
that a buggy handler can never crash the paho loop thread.
|
||||
"""
|
||||
self._subscriptions[topic] = handler
|
||||
# Grab the client reference outside the lock to avoid holding the lock
|
||||
# while calling back into paho (which may itself acquire internal locks).
|
||||
with self._lock:
|
||||
client = self._client
|
||||
connected = self._connected
|
||||
|
||||
if client is not None and connected:
|
||||
try:
|
||||
client.subscribe(topic)
|
||||
logger.debug("MQTT subscribed to topic=%s (immediate).", topic)
|
||||
except Exception:
|
||||
logger.exception("MQTT subscribe error (topic=%s).", topic)
|
||||
|
||||
def unsubscribe(self, topic: str) -> None:
|
||||
"""Remove the handler for *topic* and unsubscribe from the broker.
|
||||
|
||||
Idempotent: unknown topics are ignored. Used to apply config changes
|
||||
without a restart (e.g. when DSMR ingest is turned off or its topic
|
||||
changes). If the client is connected, the broker unsubscribe is sent
|
||||
immediately; either way the handler is removed from the registry so it
|
||||
will not be re-subscribed on the next ``_on_connect``.
|
||||
"""
|
||||
self._subscriptions.pop(topic, None)
|
||||
with self._lock:
|
||||
client = self._client
|
||||
connected = self._connected
|
||||
|
||||
if client is not None and connected:
|
||||
try:
|
||||
client.unsubscribe(topic)
|
||||
logger.debug("MQTT unsubscribed from topic=%s.", topic)
|
||||
except Exception:
|
||||
logger.exception("MQTT unsubscribe error (topic=%s).", topic)
|
||||
|
||||
def replace_source(
|
||||
self,
|
||||
source_id: int,
|
||||
*,
|
||||
host: str,
|
||||
port: int,
|
||||
username: str,
|
||||
password: str,
|
||||
tls_enabled: bool,
|
||||
subscriptions: dict[str, Callable[[bytes], None]],
|
||||
base_client_id: str = "home-automation",
|
||||
state_handler: Callable[[str], None] | None = None,
|
||||
) -> bool:
|
||||
"""Replace one source-owned client and its handlers.
|
||||
|
||||
This intentionally does not touch the legacy app-wide client or any
|
||||
other source client. It is also safe for a source to be temporarily
|
||||
unconfigured: handlers are retained in the source registry but no
|
||||
connection is attempted until a host is supplied.
|
||||
"""
|
||||
with self._source_lifecycle_lock:
|
||||
self._stop_source_client(source_id)
|
||||
if not host:
|
||||
self._report_source_state(state_handler, "error", source_id)
|
||||
return False
|
||||
self._next_source_generation += 1
|
||||
generation = self._next_source_generation
|
||||
captured_subscriptions = dict(subscriptions)
|
||||
client = mqtt.Client(
|
||||
callback_api_version=mqtt.CallbackAPIVersion.VERSION2,
|
||||
client_id=mqtt_source_client_id(base_client_id, source_id),
|
||||
)
|
||||
|
||||
def _on_connect(
|
||||
connected_client: mqtt.Client,
|
||||
_userdata: object,
|
||||
_flags: mqtt.ConnectFlags,
|
||||
reason_code: mqtt.ReasonCode,
|
||||
_properties: mqtt.Properties | None,
|
||||
) -> None:
|
||||
with self._lock:
|
||||
if not self._is_current_source_client(source_id, generation, connected_client):
|
||||
return
|
||||
if reason_code.is_failure:
|
||||
self._source_connected.discard(source_id)
|
||||
logger.warning("DSMR MQTT connection refused for source_id=%s", source_id)
|
||||
state = "error"
|
||||
else:
|
||||
self._source_connected.add(source_id)
|
||||
state = "online"
|
||||
for topic in captured_subscriptions:
|
||||
try:
|
||||
connected_client.subscribe(topic)
|
||||
except Exception:
|
||||
logger.exception("DSMR MQTT re-subscribe failed for source_id=%s", source_id)
|
||||
self._report_source_state(state_handler, state, source_id)
|
||||
|
||||
def _on_disconnect(
|
||||
disconnected_client: mqtt.Client,
|
||||
_userdata: object,
|
||||
_flags: mqtt.DisconnectFlags,
|
||||
_reason_code: mqtt.ReasonCode,
|
||||
_properties: mqtt.Properties | None,
|
||||
) -> None:
|
||||
with self._lock:
|
||||
if not self._is_current_source_client(source_id, generation, disconnected_client):
|
||||
return
|
||||
self._source_connected.discard(source_id)
|
||||
self._report_source_state(state_handler, "error", source_id)
|
||||
|
||||
def _on_message(
|
||||
message_client: mqtt.Client,
|
||||
_userdata: object,
|
||||
message: mqtt.MQTTMessage,
|
||||
) -> None:
|
||||
with self._lock:
|
||||
if not self._is_current_source_client(source_id, generation, message_client):
|
||||
return
|
||||
handler = captured_subscriptions.get(message.topic)
|
||||
state = self._source_states.get(source_id)
|
||||
if handler is None or state is None:
|
||||
return
|
||||
# This permit covers the entire handler call. Teardown first
|
||||
# invalidates the state and then waits for all permits, so an
|
||||
# old callback cannot run after teardown returns.
|
||||
state.in_flight += 1
|
||||
try:
|
||||
handler(message.payload)
|
||||
except Exception:
|
||||
logger.exception("DSMR source handler raised (source_id=%s)", source_id)
|
||||
finally:
|
||||
with self._lock:
|
||||
state.in_flight -= 1
|
||||
if state.in_flight == 0:
|
||||
self._source_idle.notify_all()
|
||||
|
||||
client.on_connect = _on_connect
|
||||
client.on_disconnect = _on_disconnect
|
||||
client.on_message = _on_message
|
||||
if tls_enabled:
|
||||
try:
|
||||
client.tls_set()
|
||||
except Exception:
|
||||
logger.exception("DSMR MQTT TLS setup failed for source_id=%s", source_id)
|
||||
self._report_source_state(state_handler, "error", source_id)
|
||||
return False
|
||||
if username:
|
||||
client.username_pw_set(username=username, password=password or None)
|
||||
# Register ownership before network processing begins. A broker
|
||||
# may deliver CONNACK synchronously from connect(), or on the loop
|
||||
# thread before connect() returns.
|
||||
with self._lock:
|
||||
self._source_clients[source_id] = client
|
||||
self._source_subscriptions[source_id] = captured_subscriptions
|
||||
self._source_generations[source_id] = generation
|
||||
self._source_states[source_id] = _SourceClientState(client, generation)
|
||||
self._report_source_state(state_handler, "connecting", source_id)
|
||||
client.loop_start()
|
||||
try:
|
||||
client.connect(host=host, port=port, keepalive=60)
|
||||
except Exception:
|
||||
logger.exception("DSMR MQTT connect failed (source_id=%s, host=%s)", source_id, host)
|
||||
self._report_source_state(state_handler, "error", source_id)
|
||||
self._stop_source_client(source_id)
|
||||
return False
|
||||
return True
|
||||
|
||||
def remove_source(self, source_id: int) -> None:
|
||||
"""Drop one source client and its handlers, including queued callbacks."""
|
||||
with self._source_lifecycle_lock:
|
||||
self._stop_source_client(source_id)
|
||||
|
||||
def source_is_active(self, source_id: int) -> bool:
|
||||
"""Whether a source-owned client is currently installed for callbacks."""
|
||||
with self._lock:
|
||||
return source_id in self._source_clients
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _start_client(self, settings: Settings) -> None:
|
||||
"""Build a fresh paho Client, configure it, and call loop_start + connect."""
|
||||
client = mqtt.Client(
|
||||
callback_api_version=mqtt.CallbackAPIVersion.VERSION2,
|
||||
client_id=settings.mqtt_client_id,
|
||||
)
|
||||
|
||||
# Callbacks — VERSION2 on_connect signature:
|
||||
# (client, userdata, connect_flags, reason_code, properties)
|
||||
def _on_connect(
|
||||
_client: mqtt.Client,
|
||||
_userdata: object,
|
||||
_flags: mqtt.ConnectFlags,
|
||||
reason_code: mqtt.ReasonCode,
|
||||
_properties: mqtt.Properties | None,
|
||||
) -> None:
|
||||
if reason_code.is_failure:
|
||||
logger.warning("MQTT connection refused: %s", reason_code)
|
||||
self._connected = False
|
||||
else:
|
||||
logger.info("MQTT connected (reason_code=%s).", reason_code)
|
||||
self._connected = True
|
||||
# Re-establish all registered subscriptions after (re-)connect.
|
||||
for sub_topic in self._subscriptions:
|
||||
try:
|
||||
_client.subscribe(sub_topic)
|
||||
logger.debug("MQTT re-subscribed to topic=%s.", sub_topic)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"MQTT re-subscribe failed (topic=%s).", sub_topic
|
||||
)
|
||||
|
||||
# VERSION2 on_disconnect signature:
|
||||
# (client, userdata, disconnect_flags, reason_code, properties)
|
||||
def _on_disconnect(
|
||||
_client: mqtt.Client,
|
||||
_userdata: object,
|
||||
_disconnect_flags: mqtt.DisconnectFlags,
|
||||
reason_code: mqtt.ReasonCode,
|
||||
_properties: mqtt.Properties | None,
|
||||
) -> None:
|
||||
self._connected = False
|
||||
if reason_code.is_failure:
|
||||
logger.warning("MQTT disconnected unexpectedly (reason_code=%s).", reason_code)
|
||||
else:
|
||||
logger.info("MQTT disconnected cleanly.")
|
||||
|
||||
# VERSION2 on_message signature: (client, userdata, message)
|
||||
def _on_message(
|
||||
_client: mqtt.Client,
|
||||
_userdata: object,
|
||||
message: mqtt.MQTTMessage,
|
||||
) -> None:
|
||||
handler = self._subscriptions.get(message.topic)
|
||||
if handler is None:
|
||||
logger.debug(
|
||||
"MQTT on_message: no handler for topic=%s.", message.topic
|
||||
)
|
||||
return
|
||||
try:
|
||||
handler(message.payload)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"MQTT message handler raised for topic=%s (swallowed).",
|
||||
message.topic,
|
||||
)
|
||||
|
||||
client.on_connect = _on_connect
|
||||
client.on_disconnect = _on_disconnect
|
||||
client.on_message = _on_message
|
||||
|
||||
# TLS
|
||||
if settings.mqtt_tls_enabled:
|
||||
try:
|
||||
client.tls_set()
|
||||
except Exception:
|
||||
logger.exception("MQTT TLS setup failed.")
|
||||
return
|
||||
|
||||
# Credentials — never log the password
|
||||
if settings.mqtt_username:
|
||||
client.username_pw_set(
|
||||
username=settings.mqtt_username,
|
||||
password=settings.mqtt_password or None,
|
||||
)
|
||||
|
||||
# Start background loop thread *before* connect so paho can handle
|
||||
# the connection handshake asynchronously.
|
||||
client.loop_start()
|
||||
|
||||
try:
|
||||
client.connect(
|
||||
host=settings.mqtt_broker_host,
|
||||
port=settings.mqtt_broker_port,
|
||||
keepalive=60,
|
||||
)
|
||||
except Exception:
|
||||
# Log without leaking the password
|
||||
logger.exception(
|
||||
"MQTT connect failed (host=%s, port=%s). Stopping loop.",
|
||||
settings.mqtt_broker_host,
|
||||
settings.mqtt_broker_port,
|
||||
)
|
||||
try:
|
||||
client.loop_stop()
|
||||
except Exception:
|
||||
pass
|
||||
return
|
||||
|
||||
self._client = client
|
||||
logger.info(
|
||||
"MQTT client started (host=%s, port=%s).",
|
||||
settings.mqtt_broker_host,
|
||||
settings.mqtt_broker_port,
|
||||
)
|
||||
|
||||
def _stop_client(self) -> None:
|
||||
"""Disconnect and stop the background loop. Idempotent."""
|
||||
client = self._client
|
||||
if client is None:
|
||||
return
|
||||
self._client = None
|
||||
self._connected = False
|
||||
try:
|
||||
client.disconnect()
|
||||
except Exception:
|
||||
logger.debug("MQTT disconnect raised (ignoring).", exc_info=True)
|
||||
try:
|
||||
client.loop_stop()
|
||||
except Exception:
|
||||
logger.debug("MQTT loop_stop raised (ignoring).", exc_info=True)
|
||||
logger.info("MQTT client stopped.")
|
||||
|
||||
def _stop_source_client(self, source_id: int) -> None:
|
||||
"""Detach then stop a source client without blocking callback bookkeeping.
|
||||
|
||||
Callers hold ``_source_lifecycle_lock``. The first phase makes the
|
||||
generation unreachable while holding ``_lock``. Paho operations and
|
||||
the in-flight wait are deliberately outside that lock: loop_stop()
|
||||
joins paho's network thread, and an active callback needs ``_lock`` to
|
||||
release its permit in ``_on_message``'s finally block.
|
||||
"""
|
||||
with self._lock:
|
||||
state = self._source_states.pop(source_id, None)
|
||||
client = self._source_clients.pop(source_id, None)
|
||||
self._source_subscriptions.pop(source_id, None)
|
||||
self._source_connected.discard(source_id)
|
||||
# Invalidate callbacks even when there was no successfully
|
||||
# installed client (for example after a failed replacement).
|
||||
self._source_generations.pop(source_id, None)
|
||||
if client is not None:
|
||||
try:
|
||||
client.disconnect()
|
||||
except Exception:
|
||||
logger.debug("DSMR MQTT disconnect raised (source_id=%s)", source_id, exc_info=True)
|
||||
try:
|
||||
client.loop_stop()
|
||||
except Exception:
|
||||
logger.debug("DSMR MQTT loop_stop raised (source_id=%s)", source_id, exc_info=True)
|
||||
if state is not None:
|
||||
with self._lock:
|
||||
while state.in_flight:
|
||||
self._source_idle.wait()
|
||||
|
||||
def _is_current_source_client(
|
||||
self, source_id: int, generation: int, client: mqtt.Client
|
||||
) -> bool:
|
||||
"""Check callback ownership while ``_lock`` is held."""
|
||||
return (
|
||||
self._source_generations.get(source_id) == generation
|
||||
and self._source_clients.get(source_id) is client
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _report_source_state(
|
||||
state_handler: Callable[[str], None] | None, state: str, source_id: int
|
||||
) -> None:
|
||||
"""Invoke an optional health callback without exposing connection credentials."""
|
||||
if state_handler is None:
|
||||
return
|
||||
try:
|
||||
state_handler(state)
|
||||
except Exception:
|
||||
logger.exception("DSMR MQTT source state update failed for source_id=%s", source_id)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Module-level singleton — shared across lifespan and route handlers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
mqtt_manager = MqttManager()
|
||||
@@ -0,0 +1,192 @@
|
||||
"""Pure, privacy-preserving parser for DSMR and WarmteLink P1 telegrams."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field as dataclass_field
|
||||
from decimal import Decimal, InvalidOperation
|
||||
from enum import StrEnum
|
||||
import hashlib
|
||||
import re
|
||||
|
||||
|
||||
_OBIS_LINE = re.compile(r"^(?P<code>\d+-\d+:\d+\.\d+\.\d+)(?P<values>(?:\([^)]*\))*)$")
|
||||
_NUMBER_WITH_UNIT = re.compile(r"^(?P<number>[+-]?\d+(?:\.\d+)?)(?:\*(?P<unit>.+))?$")
|
||||
_CHANNEL_OBIS = re.compile(r"^0-(?P<channel>[1-9]\d*):(24|96)\.")
|
||||
_EQUIPMENT_ID_CODES = re.compile(r"^0-(?:0|[1-9]\d*):96\.1\.[01]$")
|
||||
|
||||
|
||||
class IntegrityStatus(StrEnum):
|
||||
"""Whether a frame has a verifiable standard DSMR checksum."""
|
||||
|
||||
VALID = "valid"
|
||||
INVALID = "invalid"
|
||||
UNVERIFIABLE = "unverifiable"
|
||||
|
||||
|
||||
class P1ParseError(ValueError):
|
||||
"""A parse error whose message never includes telegram contents."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ObisField:
|
||||
"""One sanitized OBIS line, including values not understood by the parser."""
|
||||
|
||||
code: str
|
||||
raw_values: tuple[str, ...]
|
||||
value: Decimal | None = None
|
||||
unit: str | None = None
|
||||
comparison_token: str | None = dataclass_field(default=None, repr=False)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class P1Channel:
|
||||
"""Fields associated with one M-Bus channel, without its raw identifier."""
|
||||
|
||||
number: int
|
||||
device_type: str | None
|
||||
equipment_fingerprint: str | None = dataclass_field(repr=False)
|
||||
readings: tuple[ObisField, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class P1Telegram:
|
||||
"""A parsed telegram with only sanitized, persistence-safe data."""
|
||||
|
||||
frame_length: int
|
||||
integrity: IntegrityStatus
|
||||
integrity_reason: str
|
||||
timestamp: str | None
|
||||
equipment_fingerprint: str | None = dataclass_field(repr=False)
|
||||
fields: tuple[ObisField, ...]
|
||||
channels: tuple[P1Channel, ...]
|
||||
|
||||
|
||||
class TelegramFramer:
|
||||
"""Incrementally extract newline-terminated variable-length telegrams."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._buffer = bytearray()
|
||||
|
||||
def feed(self, chunk: bytes) -> list[bytes]:
|
||||
"""Append *chunk* and return every complete frame now available."""
|
||||
|
||||
self._buffer.extend(chunk)
|
||||
frames: list[bytes] = []
|
||||
while (bang := self._buffer.find(b"!")) >= 0:
|
||||
newline = self._buffer.find(b"\n", bang)
|
||||
if newline < 0:
|
||||
break
|
||||
standard_start = self._buffer.find(b"/")
|
||||
start = standard_start if 0 <= standard_start < bang else 0
|
||||
frames.append(bytes(self._buffer[start : newline + 1]))
|
||||
del self._buffer[: newline + 1]
|
||||
return frames
|
||||
|
||||
|
||||
def dsmr_crc16(data: bytes) -> int:
|
||||
"""Return the DSMR CRC-16 over *data* (normally from ``/`` through ``!``)."""
|
||||
|
||||
crc = 0
|
||||
for byte in data:
|
||||
crc ^= byte
|
||||
for _ in range(8):
|
||||
crc = (crc >> 1) ^ 0xA001 if crc & 1 else crc >> 1
|
||||
return crc & 0xFFFF
|
||||
|
||||
|
||||
def parse_telegram(frame: bytes) -> P1Telegram:
|
||||
"""Parse one complete frame without retaining raw telegram bytes or IDs."""
|
||||
|
||||
bang = frame.find(b"!")
|
||||
if bang < 0:
|
||||
raise P1ParseError("telegram has no footer marker")
|
||||
body = frame[:bang]
|
||||
footer = frame[bang + 1 :].rstrip(b"\r\n")
|
||||
integrity, reason = _integrity(frame, bang, footer)
|
||||
fields, identifiers = _parse_obis_fields(body)
|
||||
timestamp = _field_value(fields, "0-0:1.0.0")
|
||||
return P1Telegram(
|
||||
frame_length=len(frame),
|
||||
integrity=integrity,
|
||||
integrity_reason=reason,
|
||||
timestamp=timestamp,
|
||||
equipment_fingerprint=identifiers.get("0-0:96.1.1"),
|
||||
fields=tuple(fields),
|
||||
channels=_parse_channels(fields, identifiers),
|
||||
)
|
||||
|
||||
|
||||
def _integrity(frame: bytes, bang: int, footer: bytes) -> tuple[IntegrityStatus, str]:
|
||||
if not frame.startswith(b"/"):
|
||||
return IntegrityStatus.UNVERIFIABLE, "missing standard DSMR '/' header"
|
||||
if len(footer) != 4 or not all(chr(byte) in "0123456789abcdefABCDEF" for byte in footer):
|
||||
return IntegrityStatus.UNVERIFIABLE, "footer is not a four-digit hexadecimal CRC"
|
||||
expected = int(footer, 16)
|
||||
actual = dsmr_crc16(frame[: bang + 1])
|
||||
if actual == expected:
|
||||
return IntegrityStatus.VALID, "CRC16 verified from '/' through '!'"
|
||||
return IntegrityStatus.INVALID, f"CRC16 mismatch: expected {expected:04X}, calculated {actual:04X}"
|
||||
|
||||
|
||||
def _parse_obis_fields(body: bytes) -> tuple[list[ObisField], dict[str, str]]:
|
||||
fields: list[ObisField] = []
|
||||
identifiers: dict[str, str] = {}
|
||||
for line in body.decode("ascii", errors="replace").splitlines()[1:]:
|
||||
match = _OBIS_LINE.fullmatch(line)
|
||||
if not match:
|
||||
continue
|
||||
code = match.group("code")
|
||||
raw_values = tuple(re.findall(r"\(([^)]*)\)", match.group("values")))
|
||||
if _EQUIPMENT_ID_CODES.fullmatch(code):
|
||||
comparison_token = _fingerprint(raw_values[-1]) if raw_values else None
|
||||
if comparison_token is not None:
|
||||
identifiers[code] = comparison_token
|
||||
fields.append(ObisField(code, ("<redacted>",), comparison_token=comparison_token))
|
||||
continue
|
||||
value, unit = _numeric_value(raw_values)
|
||||
fields.append(ObisField(code, raw_values, value, unit))
|
||||
return fields, identifiers
|
||||
|
||||
|
||||
def _fingerprint(identifier: str) -> str:
|
||||
"""Hash an identifier locally; callers never receive its original value."""
|
||||
|
||||
return hashlib.sha256(identifier.encode("ascii", errors="replace")).hexdigest()
|
||||
|
||||
|
||||
def _numeric_value(raw_values: tuple[str, ...]) -> tuple[Decimal | None, str | None]:
|
||||
if not raw_values:
|
||||
return None, None
|
||||
match = _NUMBER_WITH_UNIT.fullmatch(raw_values[-1])
|
||||
if not match:
|
||||
return None, None
|
||||
try:
|
||||
return Decimal(match.group("number")), match.group("unit")
|
||||
except InvalidOperation:
|
||||
return None, None
|
||||
|
||||
|
||||
def _field_value(fields: list[ObisField], code: str) -> str | None:
|
||||
field = next((item for item in fields if item.code == code), None)
|
||||
return field.raw_values[-1] if field and field.raw_values else None
|
||||
|
||||
|
||||
def _parse_channels(fields: list[ObisField], identifiers: dict[str, str]) -> tuple[P1Channel, ...]:
|
||||
by_channel: dict[int, list[ObisField]] = {}
|
||||
for field in fields:
|
||||
match = _CHANNEL_OBIS.match(field.code)
|
||||
if match:
|
||||
by_channel.setdefault(int(match.group("channel")), []).append(field)
|
||||
return tuple(
|
||||
P1Channel(
|
||||
number=number,
|
||||
device_type=_field_value(channel_fields, f"0-{number}:24.1.0"),
|
||||
equipment_fingerprint=identifiers.get(f"0-{number}:96.1.0"),
|
||||
readings=tuple(
|
||||
field
|
||||
for field in channel_fields
|
||||
if field.code == f"0-{number}:24.2.1" and field.value is not None
|
||||
),
|
||||
)
|
||||
for number, channel_fields in sorted(by_channel.items())
|
||||
)
|
||||
@@ -0,0 +1,11 @@
|
||||
"""Pricing profile framework and strategy registry.
|
||||
|
||||
This package provides:
|
||||
|
||||
- ``profiles``: Pydantic models describing the structure of pricing profiles
|
||||
(``ManualProfile`` / ``TibberProfile``), YAML loaders, and contract-values
|
||||
validation (``load_profile`` / ``list_profiles`` / ``validate_values``).
|
||||
- ``strategies``: A lightweight registry of price-calculation strategies
|
||||
(``register_strategy`` / ``get_strategy``), with built-in implementations for
|
||||
the ``manual`` (fixed dual-tariff) and ``tibber`` (dynamic API) kinds.
|
||||
"""
|
||||
@@ -0,0 +1,556 @@
|
||||
"""Pricing profile loader, validator, and contract-values checker.
|
||||
|
||||
A *pricing profile* is a YAML file that describes the **structure** of an energy
|
||||
contract — which fields exist, their units, and which have defaults. Actual
|
||||
pricing values (the numbers the user fills in via the UI) are stored in the DB
|
||||
as ``EnergyContractVersion.values`` (a JSON blob) and must conform to the
|
||||
corresponding profile's structure.
|
||||
|
||||
Profiles live in ``app/integrations/pricing/profiles/<kind>.yaml`` and are
|
||||
located at runtime relative to *this file* (not CWD), matching the pattern
|
||||
established by the Modbus profile loader.
|
||||
|
||||
Design notes
|
||||
------------
|
||||
- **Purely data + functions** — no abstract base classes or inheritance.
|
||||
- Two Pydantic models — ``ManualProfile`` and ``TibberProfile`` — capture the
|
||||
different structures of the two supported kinds. A thin union dispatcher in
|
||||
``load_profile`` picks the right one.
|
||||
- ``validate_values(kind, values)`` fills in fields that carry a ``default`` and
|
||||
raises ``ProfileValidationError`` for missing required fields or wrong types.
|
||||
- All errors are subclasses of built-ins so callers need not import this module
|
||||
just to catch them.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from decimal import Decimal, InvalidOperation
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
import yaml
|
||||
from pydantic import BaseModel, ValidationError, model_validator
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Directory containing the YAML profiles, located relative to *this* file.
|
||||
_PROFILES_DIR = Path(__file__).parent / "profiles"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Custom exceptions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ProfileNotFoundError(FileNotFoundError):
|
||||
"""Raised when the requested profile YAML file does not exist."""
|
||||
|
||||
|
||||
class ProfileValidationError(ValueError):
|
||||
"""Raised when a profile YAML fails Pydantic validation or when contract
|
||||
values do not conform to the profile structure."""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Leaf-node models (a single pricing field: unit + optional default)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class FieldSpec(BaseModel):
|
||||
"""Specification for a single numeric pricing field."""
|
||||
|
||||
unit: str
|
||||
"""Physical / monetary unit string (e.g. ``"EUR/kWh"``, ``"EUR/month"``)."""
|
||||
|
||||
default: Optional[float] = None
|
||||
"""Default value used when the field is absent from contract values.
|
||||
``None`` means the field is required (no default)."""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ManualProfile — fixed / variable dual-tariff contract structure
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ManualBuySpec(BaseModel):
|
||||
normal: FieldSpec # high-tariff buy price (delivered_2 registers)
|
||||
dal: FieldSpec # low-tariff buy price (delivered_1 registers)
|
||||
|
||||
|
||||
class ManualSellSpec(BaseModel):
|
||||
normal: FieldSpec # high-tariff sell / return price
|
||||
dal: FieldSpec # low-tariff sell / return price
|
||||
|
||||
|
||||
class ManualEnergySpec(BaseModel):
|
||||
dual_tariff: bool # always True for manual profiles
|
||||
buy: ManualBuySpec
|
||||
sell: ManualSellSpec
|
||||
energy_tax: FieldSpec # added to buy price; includes VAT
|
||||
ode: FieldSpec # currently merged into energy_tax; default 0
|
||||
|
||||
|
||||
class ManualStandingSpec(BaseModel):
|
||||
network_fee: FieldSpec # EUR/month
|
||||
management_fee: FieldSpec # EUR/month
|
||||
|
||||
|
||||
class ManualCreditsSpec(BaseModel):
|
||||
heffingskorting: FieldSpec # EUR/year — energy-tax credit deducted at summary
|
||||
|
||||
|
||||
class ManualProfile(BaseModel):
|
||||
"""Complete structure description for a ``manual`` pricing contract."""
|
||||
|
||||
kind: str
|
||||
label: str
|
||||
energy: ManualEnergySpec
|
||||
standing: ManualStandingSpec
|
||||
credits: ManualCreditsSpec
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_kind(self) -> "ManualProfile":
|
||||
if self.kind != "manual":
|
||||
raise ValueError(f"ManualProfile requires kind='manual', got {self.kind!r}")
|
||||
return self
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# TibberProfile — dynamic Tibber API contract structure
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TibberEnergySpec(BaseModel):
|
||||
source: str # must be "tibber_api"
|
||||
energy_tax: FieldSpec # subtracted from total to derive sell price
|
||||
sell_fee: FieldSpec # verkoopvergoeding (feed-in fee); always subtracted from sell; default 0.0248
|
||||
sell_adjust: FieldSpec # additional sell-price adjustment; default 0
|
||||
|
||||
|
||||
class TibberStandingSpec(BaseModel):
|
||||
management_fee: FieldSpec # EUR/month; has a default
|
||||
network_fee: FieldSpec # EUR/month
|
||||
|
||||
|
||||
class TibberCreditsSpec(BaseModel):
|
||||
heffingskorting: FieldSpec # EUR/year
|
||||
|
||||
|
||||
class TibberProfile(BaseModel):
|
||||
"""Complete structure description for a ``tibber`` pricing contract."""
|
||||
|
||||
kind: str
|
||||
label: str
|
||||
energy: TibberEnergySpec
|
||||
standing: TibberStandingSpec
|
||||
credits: TibberCreditsSpec
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_kind(self) -> "TibberProfile":
|
||||
if self.kind != "tibber":
|
||||
raise ValueError(f"TibberProfile requires kind='tibber', got {self.kind!r}")
|
||||
return self
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DistrictHeatingProfile — user-entered thermal contract structure
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class DistrictHeatingFieldSpec(BaseModel):
|
||||
"""A Decimal-safe thermal tariff field displayed to the user."""
|
||||
|
||||
model_config = {"extra": "forbid"}
|
||||
|
||||
unit: str
|
||||
label: str
|
||||
help: str
|
||||
minimum: Decimal = Decimal("0")
|
||||
default: Decimal | None = None
|
||||
|
||||
|
||||
class DistrictHeatingVariableSpec(BaseModel):
|
||||
model_config = {"extra": "forbid"}
|
||||
|
||||
heating: DistrictHeatingFieldSpec
|
||||
hot_water_heating: DistrictHeatingFieldSpec
|
||||
hot_water: DistrictHeatingFieldSpec
|
||||
hot_water_tax: DistrictHeatingFieldSpec
|
||||
|
||||
|
||||
class DistrictHeatingStandingSpec(BaseModel):
|
||||
model_config = {"extra": "forbid"}
|
||||
|
||||
heating_network: DistrictHeatingFieldSpec
|
||||
metering: DistrictHeatingFieldSpec
|
||||
delivery_set: DistrictHeatingFieldSpec
|
||||
hot_water_network: DistrictHeatingFieldSpec
|
||||
other: DistrictHeatingFieldSpec
|
||||
|
||||
|
||||
class DistrictHeatingProfile(BaseModel):
|
||||
"""Complete structure description for a ``district_heating`` contract."""
|
||||
|
||||
model_config = {"extra": "forbid"}
|
||||
|
||||
kind: str
|
||||
label: str
|
||||
variable: DistrictHeatingVariableSpec
|
||||
standing: DistrictHeatingStandingSpec
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_kind(self) -> "DistrictHeatingProfile":
|
||||
if self.kind != "district_heating":
|
||||
raise ValueError(
|
||||
"DistrictHeatingProfile requires kind='district_heating', "
|
||||
f"got {self.kind!r}"
|
||||
)
|
||||
units = {
|
||||
"heating": "EUR/GJ",
|
||||
"hot_water_heating": "EUR/m³",
|
||||
"hot_water": "EUR/m³",
|
||||
"hot_water_tax": "EUR/m³",
|
||||
}
|
||||
for key, unit in units.items():
|
||||
field = getattr(self.variable, key)
|
||||
if field.unit != unit or field.minimum != 0 or field.default is not None:
|
||||
raise ValueError(f"district_heating.variable.{key} must be required {unit} with minimum 0")
|
||||
for key in DistrictHeatingStandingSpec.model_fields:
|
||||
field = getattr(self.standing, key)
|
||||
if field.unit != "EUR/year" or field.minimum != 0 or field.default != 0:
|
||||
raise ValueError(
|
||||
f"district_heating.standing.{key} must be EUR/year with default and minimum 0"
|
||||
)
|
||||
return self
|
||||
|
||||
|
||||
# A union type for type hints where either profile is acceptable.
|
||||
AnyProfile = ManualProfile | TibberProfile | DistrictHeatingProfile
|
||||
|
||||
# Map kind → Pydantic model class used for validation.
|
||||
_PROFILE_MODELS: dict[str, type[ManualProfile] | type[TibberProfile] | type[DistrictHeatingProfile]] = {
|
||||
"manual": ManualProfile,
|
||||
"tibber": TibberProfile,
|
||||
"district_heating": DistrictHeatingProfile,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public API
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def load_profile(kind: str) -> AnyProfile:
|
||||
"""Load and validate a pricing profile by kind name.
|
||||
|
||||
The profile file is expected at
|
||||
``app/integrations/pricing/profiles/<kind>.yaml``.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
kind:
|
||||
Profile kind without extension (``"manual"`` or ``"tibber"``).
|
||||
|
||||
Returns
|
||||
-------
|
||||
ManualProfile | TibberProfile
|
||||
Validated profile model.
|
||||
|
||||
Raises
|
||||
------
|
||||
ProfileNotFoundError
|
||||
If ``profiles/<kind>.yaml`` does not exist.
|
||||
ProfileValidationError
|
||||
If the YAML is syntactically valid but fails schema validation.
|
||||
"""
|
||||
path = _PROFILES_DIR / f"{kind}.yaml"
|
||||
if not path.exists():
|
||||
raise ProfileNotFoundError(
|
||||
f"Pricing profile '{kind}' not found (looked for {path})"
|
||||
)
|
||||
|
||||
with path.open("r", encoding="utf-8") as fh:
|
||||
raw = yaml.safe_load(fh)
|
||||
|
||||
if not isinstance(raw, dict):
|
||||
raise ProfileValidationError(
|
||||
f"Profile '{kind}': expected a YAML mapping, got {type(raw).__name__}"
|
||||
)
|
||||
|
||||
# YAML's implicit float conversion must never contaminate the thermal profile.
|
||||
# Existing electricity profiles intentionally retain their established defaults.
|
||||
if raw.get("kind", kind) == "district_heating":
|
||||
_reject_yaml_floats(raw, path)
|
||||
|
||||
# Choose the right Pydantic model based on the ``kind`` field in the YAML.
|
||||
yaml_kind = raw.get("kind", kind)
|
||||
model_cls = _PROFILE_MODELS.get(yaml_kind)
|
||||
if model_cls is None:
|
||||
raise ProfileValidationError(
|
||||
f"Profile '{kind}' has unknown kind={yaml_kind!r}; "
|
||||
f"supported: {list(_PROFILE_MODELS)}"
|
||||
)
|
||||
|
||||
try:
|
||||
return model_cls.model_validate(raw)
|
||||
except ValidationError as exc:
|
||||
raise ProfileValidationError(
|
||||
f"Profile '{kind}' failed validation: {exc}"
|
||||
) from exc
|
||||
|
||||
|
||||
def list_profiles() -> list[dict[str, Any]]:
|
||||
"""Return structural data for all available pricing profiles.
|
||||
|
||||
Scans ``app/integrations/pricing/profiles/*.yaml``, loads each, and returns
|
||||
a list of dicts (Pydantic model dumps). Profiles that fail to load are
|
||||
skipped with a warning.
|
||||
|
||||
Returns
|
||||
-------
|
||||
list[dict[str, Any]]
|
||||
One entry per valid profile, suitable for JSON serialisation and
|
||||
front-end form rendering.
|
||||
"""
|
||||
results: list[dict[str, Any]] = []
|
||||
if not _PROFILES_DIR.exists():
|
||||
return results
|
||||
|
||||
for path in sorted(_PROFILES_DIR.glob("*.yaml")):
|
||||
kind = path.stem
|
||||
try:
|
||||
profile = load_profile(kind)
|
||||
results.append(profile.model_dump())
|
||||
except (ProfileNotFoundError, ProfileValidationError, Exception) as exc:
|
||||
logger.warning("Skipping malformed pricing profile '%s': %s", kind, exc)
|
||||
|
||||
return results
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Contract-values validation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _fill_defaults_manual(values: dict[str, Any], profile: ManualProfile) -> dict[str, Any]:
|
||||
"""Return a copy of *values* with any ``ode`` default applied if missing."""
|
||||
filled = dict(values)
|
||||
|
||||
energy = dict(filled.get("energy", {}))
|
||||
|
||||
# Apply default for ode (default=0) if absent.
|
||||
if "ode" not in energy and profile.energy.ode.default is not None:
|
||||
energy["ode"] = profile.energy.ode.default
|
||||
|
||||
filled["energy"] = energy
|
||||
return filled
|
||||
|
||||
|
||||
def _fill_defaults_tibber(values: dict[str, Any], profile: TibberProfile) -> dict[str, Any]:
|
||||
"""Return a copy of *values* with sell_fee, sell_adjust and management_fee defaults applied."""
|
||||
filled = dict(values)
|
||||
|
||||
energy = dict(filled.get("energy", {}))
|
||||
# Apply default for sell_fee (default=0.0248) if absent.
|
||||
if "sell_fee" not in energy and profile.energy.sell_fee.default is not None:
|
||||
energy["sell_fee"] = profile.energy.sell_fee.default
|
||||
# Apply default for sell_adjust (default=0) if absent.
|
||||
if "sell_adjust" not in energy and profile.energy.sell_adjust.default is not None:
|
||||
energy["sell_adjust"] = profile.energy.sell_adjust.default
|
||||
filled["energy"] = energy
|
||||
|
||||
standing = dict(filled.get("standing", {}))
|
||||
# Apply default for management_fee if absent.
|
||||
if (
|
||||
"management_fee" not in standing
|
||||
and profile.standing.management_fee.default is not None
|
||||
):
|
||||
standing["management_fee"] = profile.standing.management_fee.default
|
||||
filled["standing"] = standing
|
||||
|
||||
return filled
|
||||
|
||||
|
||||
def _require_numeric(section: str, key: str, container: dict[str, Any]) -> None:
|
||||
"""Assert that *container[key]* exists and is a number; raise ProfileValidationError."""
|
||||
if key not in container:
|
||||
raise ProfileValidationError(
|
||||
f"Contract values missing required field '{section}.{key}'"
|
||||
)
|
||||
val = container[key]
|
||||
if not isinstance(val, (int, float)):
|
||||
raise ProfileValidationError(
|
||||
f"Contract values field '{section}.{key}' must be a number, "
|
||||
f"got {type(val).__name__!r}"
|
||||
)
|
||||
|
||||
|
||||
def _validate_manual_values(values: dict[str, Any], profile: ManualProfile) -> dict[str, Any]:
|
||||
"""Validate and fill-defaults for manual contract values.
|
||||
|
||||
Returns the filled values dict on success. Raises ProfileValidationError
|
||||
on missing required fields or wrong types.
|
||||
"""
|
||||
filled = _fill_defaults_manual(values, profile)
|
||||
energy = filled.get("energy", {})
|
||||
buy = energy.get("buy", {})
|
||||
sell = energy.get("sell", {})
|
||||
standing = filled.get("standing", {})
|
||||
credits = filled.get("credits", {})
|
||||
|
||||
# Required energy.buy fields.
|
||||
_require_numeric("energy.buy", "normal", buy)
|
||||
_require_numeric("energy.buy", "dal", buy)
|
||||
# Required energy.sell fields.
|
||||
_require_numeric("energy.sell", "normal", sell)
|
||||
_require_numeric("energy.sell", "dal", sell)
|
||||
# Required energy fields.
|
||||
_require_numeric("energy", "energy_tax", energy)
|
||||
_require_numeric("energy", "ode", energy)
|
||||
# Required standing fields.
|
||||
_require_numeric("standing", "network_fee", standing)
|
||||
_require_numeric("standing", "management_fee", standing)
|
||||
# Required credits fields.
|
||||
_require_numeric("credits", "heffingskorting", credits)
|
||||
|
||||
return filled
|
||||
|
||||
|
||||
def _validate_tibber_values(values: dict[str, Any], profile: TibberProfile) -> dict[str, Any]:
|
||||
"""Validate and fill-defaults for tibber contract values.
|
||||
|
||||
Returns the filled values dict on success. Raises ProfileValidationError
|
||||
on missing required fields or wrong types.
|
||||
"""
|
||||
filled = _fill_defaults_tibber(values, profile)
|
||||
energy = filled.get("energy", {})
|
||||
standing = filled.get("standing", {})
|
||||
credits = filled.get("credits", {})
|
||||
|
||||
# Required energy fields.
|
||||
_require_numeric("energy", "energy_tax", energy)
|
||||
_require_numeric("energy", "sell_fee", energy)
|
||||
_require_numeric("energy", "sell_adjust", energy)
|
||||
# Required standing fields.
|
||||
_require_numeric("standing", "management_fee", standing)
|
||||
_require_numeric("standing", "network_fee", standing)
|
||||
# Required credits fields.
|
||||
_require_numeric("credits", "heffingskorting", credits)
|
||||
|
||||
return filled
|
||||
|
||||
|
||||
def _reject_yaml_floats(value: Any, path: Path) -> None:
|
||||
"""Reject implicit YAML floats for district-heating profile metadata."""
|
||||
if isinstance(value, float):
|
||||
raise ProfileValidationError(
|
||||
f"Profile '{path.stem}' must not contain YAML float values; use integer 0 or strings."
|
||||
)
|
||||
if isinstance(value, dict):
|
||||
for child in value.values():
|
||||
_reject_yaml_floats(child, path)
|
||||
elif isinstance(value, list):
|
||||
for child in value:
|
||||
_reject_yaml_floats(child, path)
|
||||
|
||||
|
||||
def _decimal_value(section: str, key: str, value: Any) -> str:
|
||||
"""Validate and normalise one thermal amount without passing through float."""
|
||||
if isinstance(value, bool) or isinstance(value, float) or not isinstance(value, (str, int, Decimal)):
|
||||
raise ProfileValidationError(
|
||||
f"Contract values field '{section}.{key}' must be a Decimal-compatible string or integer, "
|
||||
f"got {type(value).__name__!r}"
|
||||
)
|
||||
try:
|
||||
amount = Decimal(str(value))
|
||||
except (InvalidOperation, ValueError) as exc:
|
||||
raise ProfileValidationError(
|
||||
f"Contract values field '{section}.{key}' must be a Decimal-compatible value"
|
||||
) from exc
|
||||
if not amount.is_finite() or amount < 0:
|
||||
raise ProfileValidationError(
|
||||
f"Contract values field '{section}.{key}' must be a non-negative finite Decimal"
|
||||
)
|
||||
return format(amount, "f")
|
||||
|
||||
|
||||
def _validate_district_heating_values(
|
||||
values: dict[str, Any], profile: DistrictHeatingProfile
|
||||
) -> dict[str, Any]:
|
||||
"""Validate thermal values and return a complete JSON-safe Decimal snapshot."""
|
||||
if not isinstance(values, dict):
|
||||
raise ProfileValidationError("District-heating contract values must be a mapping")
|
||||
expected_sections = {"variable", "standing"}
|
||||
unknown_sections = set(values) - expected_sections
|
||||
if unknown_sections:
|
||||
raise ProfileValidationError(
|
||||
f"District-heating contract values contain unknown section(s): {sorted(unknown_sections)}"
|
||||
)
|
||||
|
||||
def normalise_section(
|
||||
section: str, specs: Any, *, defaults_allowed: bool
|
||||
) -> dict[str, str]:
|
||||
supplied = values.get(section, {})
|
||||
if not isinstance(supplied, dict):
|
||||
raise ProfileValidationError(f"Contract values section '{section}' must be a mapping")
|
||||
expected = set(type(specs).model_fields)
|
||||
unknown = set(supplied) - expected
|
||||
if unknown:
|
||||
raise ProfileValidationError(
|
||||
f"Contract values section '{section}' contains unknown field(s): {sorted(unknown)}"
|
||||
)
|
||||
normalised: dict[str, str] = {}
|
||||
for key in type(specs).model_fields:
|
||||
if key not in supplied:
|
||||
field = getattr(specs, key)
|
||||
if not defaults_allowed or field.default is None:
|
||||
raise ProfileValidationError(f"Contract values missing required field '{section}.{key}'")
|
||||
normalised[key] = format(field.default, "f")
|
||||
else:
|
||||
normalised[key] = _decimal_value(section, key, supplied[key])
|
||||
return normalised
|
||||
|
||||
return {
|
||||
"variable": normalise_section("variable", profile.variable, defaults_allowed=False),
|
||||
"standing": normalise_section("standing", profile.standing, defaults_allowed=True),
|
||||
}
|
||||
|
||||
|
||||
def validate_values(kind: str, values: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Validate a contract-values dict against the named profile structure.
|
||||
|
||||
Fields that carry a ``default`` in the profile (e.g. ``ode``,
|
||||
``sell_fee``, ``sell_adjust``, tibber ``management_fee``) are silently
|
||||
filled in when absent from *values*. Fields with no default that are
|
||||
absent, or fields whose value is not a number, cause a
|
||||
``ProfileValidationError``.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
kind:
|
||||
Profile kind (``"manual"`` or ``"tibber"``).
|
||||
values:
|
||||
Contract values dict as stored in ``EnergyContractVersion.values``.
|
||||
|
||||
Returns
|
||||
-------
|
||||
dict[str, Any]
|
||||
A (possibly mutated) copy of *values* with defaults applied.
|
||||
|
||||
Raises
|
||||
------
|
||||
ProfileNotFoundError
|
||||
If the profile YAML for *kind* does not exist.
|
||||
ProfileValidationError
|
||||
If required fields are missing or have wrong types.
|
||||
"""
|
||||
profile = load_profile(kind)
|
||||
if isinstance(profile, ManualProfile):
|
||||
return _validate_manual_values(values, profile)
|
||||
if isinstance(profile, TibberProfile):
|
||||
return _validate_tibber_values(values, profile)
|
||||
if isinstance(profile, DistrictHeatingProfile):
|
||||
return _validate_district_heating_values(values, profile)
|
||||
# Unreachable with current kinds, but guard for future extensions.
|
||||
raise ProfileValidationError(f"No validator implemented for kind={kind!r}")
|
||||
@@ -0,0 +1,56 @@
|
||||
kind: district_heating
|
||||
label: 区域供热
|
||||
|
||||
variable:
|
||||
heating:
|
||||
unit: EUR/GJ
|
||||
label: 供暖热量
|
||||
help: 按供暖用热量计收;请录入合同中的实际金额。
|
||||
minimum: 0
|
||||
hot_water_heating:
|
||||
unit: EUR/m³
|
||||
label: 热水加热
|
||||
help: 按热水体积计收的加热部分;请录入合同中的实际金额。
|
||||
minimum: 0
|
||||
hot_water:
|
||||
unit: EUR/m³
|
||||
label: 热水用量
|
||||
help: 按热水体积计收的用量部分;请录入合同中的实际金额。
|
||||
minimum: 0
|
||||
hot_water_tax:
|
||||
unit: EUR/m³
|
||||
label: 热水税费
|
||||
help: 按热水体积计收的税费部分;请录入合同中的实际金额。
|
||||
minimum: 0
|
||||
|
||||
standing:
|
||||
heating_network:
|
||||
unit: EUR/year
|
||||
label: 供暖网络费
|
||||
help: 年度固定费用;默认零,按合同实际金额录入。
|
||||
minimum: 0
|
||||
default: 0
|
||||
metering:
|
||||
unit: EUR/year
|
||||
label: 计量费
|
||||
help: 年度固定费用;默认零,按合同实际金额录入。
|
||||
minimum: 0
|
||||
default: 0
|
||||
delivery_set:
|
||||
unit: EUR/year
|
||||
label: 交付装置费
|
||||
help: 年度固定费用;默认零,按合同实际金额录入。
|
||||
minimum: 0
|
||||
default: 0
|
||||
hot_water_network:
|
||||
unit: EUR/year
|
||||
label: 热水网络费
|
||||
help: 年度固定费用;默认零,按合同实际金额录入。
|
||||
minimum: 0
|
||||
default: 0
|
||||
other:
|
||||
unit: EUR/year
|
||||
label: 其他固定费
|
||||
help: 年度固定费用;默认零,按合同实际金额录入。
|
||||
minimum: 0
|
||||
default: 0
|
||||
@@ -0,0 +1,20 @@
|
||||
kind: manual
|
||||
label: 固定 / 可变费率(NL,双费率)
|
||||
|
||||
energy:
|
||||
dual_tariff: true # use delivered_1/2, returned_1/2 for low/high tariffs
|
||||
buy:
|
||||
normal: { unit: EUR/kWh } # high-tariff buy price (delivered_2 register)
|
||||
dal: { unit: EUR/kWh } # low-tariff buy price (delivered_1 register)
|
||||
sell:
|
||||
normal: { unit: EUR/kWh } # high-tariff sell / return price
|
||||
dal: { unit: EUR/kWh } # low-tariff sell / return price (currently equal to normal)
|
||||
energy_tax: { unit: EUR/kWh } # energy tax added to buy price (incl. VAT)
|
||||
ode: { unit: EUR/kWh, default: 0 } # currently merged into energy_tax
|
||||
|
||||
standing: # fixed charges; UI fills per month, engine prorates to days
|
||||
network_fee: { unit: EUR/month }
|
||||
management_fee: { unit: EUR/month }
|
||||
|
||||
credits:
|
||||
heffingskorting: { unit: EUR/year } # energy-tax credit; deducted at summary layer
|
||||
@@ -0,0 +1,15 @@
|
||||
kind: tibber
|
||||
label: Tibber 动态电价(15 分钟)
|
||||
|
||||
energy:
|
||||
source: tibber_api # buy = total (from API); sell = total − energy_tax − sell_fee − sell_adjust
|
||||
energy_tax: { unit: EUR/kWh } # subtracted from total to derive sell price (incl. VAT)
|
||||
sell_fee: { unit: EUR/kWh, default: 0.0248 } # verkoopvergoeding (feed-in fee, incl. VAT); always subtracted from sell
|
||||
sell_adjust: { unit: EUR/kWh, default: 0 } # manual sell-price adjustment; net-metering: set = −energy_tax to refund the tax
|
||||
|
||||
standing: # fixed charges; UI fills per month, engine prorates to days
|
||||
management_fee: { unit: EUR/month, default: 5.99 }
|
||||
network_fee: { unit: EUR/month }
|
||||
|
||||
credits:
|
||||
heffingskorting: { unit: EUR/year } # energy-tax credit; deducted at summary layer
|
||||
@@ -0,0 +1,301 @@
|
||||
"""Price-calculation strategy registry and built-in implementations.
|
||||
|
||||
A *strategy* is a plain function registered under a ``kind`` string. Given
|
||||
per-register kWh deltas, the period start time, the active contract-version
|
||||
values, and a DB session, it returns a result dict with:
|
||||
|
||||
{
|
||||
"import_cost": Decimal, # total cost of electricity drawn from grid
|
||||
"export_revenue": Decimal, # total revenue from electricity fed to grid
|
||||
"net_cost": Decimal, # import_cost − export_revenue
|
||||
"pricing": dict, # snapshot of the price inputs used (for auditing)
|
||||
}
|
||||
|
||||
**Import and export are kept separate throughout** — net_cost is only derived
|
||||
at the end. This supports the Dutch no-netting rule (saldering afgebouwd).
|
||||
|
||||
**All monetary arithmetic uses Decimal** converted via ``Decimal(str(x))`` to
|
||||
avoid float binary rounding errors.
|
||||
|
||||
Design notes
|
||||
------------
|
||||
- **Purely data + functions**: no ABC, no class hierarchy. A ``dict``-based
|
||||
registry is the simplest structure that supports the two current kinds and
|
||||
leaves the door open for future additions (e.g. ``octopus``, ``frank``).
|
||||
- ``deltas`` is a plain dataclass ``PeriodDeltas`` — typed, but no ORM.
|
||||
- Tibber strategy queries ``TibberPrice`` directly from the session; it does not
|
||||
call the Tibber API (that is T05's job).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
from typing import Any, Callable
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Shared data types
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class PeriodDeltas:
|
||||
"""Per-register kWh deltas for one 15-minute billing period.
|
||||
|
||||
All values are ``Decimal`` and represent ``end_reading − start_reading``
|
||||
for the corresponding cumulative energy register.
|
||||
|
||||
Attributes
|
||||
----------
|
||||
d1: Decimal
|
||||
Delivered (consumed) kWh on the low/dal tariff register (delivered_1).
|
||||
d2: Decimal
|
||||
Delivered (consumed) kWh on the normal/high tariff register (delivered_2).
|
||||
r1: Decimal
|
||||
Returned (fed-to-grid) kWh on the low/dal tariff register (returned_1).
|
||||
r2: Decimal
|
||||
Returned (fed-to-grid) kWh on the normal/high tariff register (returned_2).
|
||||
"""
|
||||
|
||||
d1: Decimal # delivered low-tariff
|
||||
d2: Decimal # delivered high-tariff
|
||||
r1: Decimal # returned low-tariff
|
||||
r2: Decimal # returned high-tariff
|
||||
|
||||
|
||||
# Strategy callable signature.
|
||||
StrategyFn = Callable[
|
||||
[PeriodDeltas, datetime, dict[str, Any], Session],
|
||||
dict[str, Any],
|
||||
]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Registry
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_REGISTRY: dict[str, StrategyFn] = {}
|
||||
|
||||
|
||||
def register_strategy(kind: str, fn: StrategyFn) -> None:
|
||||
"""Register *fn* as the price-calculation strategy for *kind*.
|
||||
|
||||
Overwrites any previously registered strategy for the same kind (allows
|
||||
monkey-patching in tests).
|
||||
"""
|
||||
_REGISTRY[kind] = fn
|
||||
|
||||
|
||||
def get_strategy(kind: str) -> StrategyFn:
|
||||
"""Return the registered strategy for *kind*.
|
||||
|
||||
Raises
|
||||
------
|
||||
KeyError
|
||||
If no strategy is registered for *kind*.
|
||||
"""
|
||||
if kind not in _REGISTRY:
|
||||
raise KeyError(
|
||||
f"No price strategy registered for kind={kind!r}. "
|
||||
f"Available: {sorted(_REGISTRY)}"
|
||||
)
|
||||
return _REGISTRY[kind]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _to_decimal(value: Any) -> Decimal:
|
||||
"""Convert *value* to Decimal via str() to avoid float binary rounding."""
|
||||
return Decimal(str(value))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Manual strategy — fixed dual-tariff
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _manual_strategy(
|
||||
deltas: PeriodDeltas,
|
||||
t0: datetime,
|
||||
values: dict[str, Any],
|
||||
session: Session,
|
||||
) -> dict[str, Any]:
|
||||
"""Calculate billing costs for one period using fixed dual-tariff rates.
|
||||
|
||||
Formula (§3.4):
|
||||
- ``buy_dal = energy.buy.dal + energy.energy_tax + energy.ode``
|
||||
- ``buy_normal = energy.buy.normal + energy.energy_tax + energy.ode``
|
||||
- ``import_cost = Δd1 × buy_dal + Δd2 × buy_normal``
|
||||
- ``sell_dal = energy.sell.dal`` (no tax on return price)
|
||||
- ``sell_normal = energy.sell.normal``
|
||||
- ``export_revenue = Δr1 × sell_dal + Δr2 × sell_normal``
|
||||
- ``net_cost = import_cost − export_revenue``
|
||||
|
||||
The ``pricing`` snapshot contains all per-unit prices used so that the
|
||||
result is fully auditable without re-querying the contract version.
|
||||
"""
|
||||
energy = values.get("energy", {})
|
||||
buy = energy.get("buy", {})
|
||||
sell = energy.get("sell", {})
|
||||
|
||||
energy_tax = _to_decimal(energy.get("energy_tax", 0))
|
||||
ode = _to_decimal(energy.get("ode", 0))
|
||||
|
||||
buy_dal_base = _to_decimal(buy.get("dal", 0))
|
||||
buy_normal_base = _to_decimal(buy.get("normal", 0))
|
||||
sell_dal = _to_decimal(sell.get("dal", 0))
|
||||
sell_normal = _to_decimal(sell.get("normal", 0))
|
||||
|
||||
# Effective buy prices including all taxes.
|
||||
buy_dal = buy_dal_base + energy_tax + ode
|
||||
buy_normal = buy_normal_base + energy_tax + ode
|
||||
|
||||
import_cost = deltas.d1 * buy_dal + deltas.d2 * buy_normal
|
||||
export_revenue = deltas.r1 * sell_dal + deltas.r2 * sell_normal
|
||||
net_cost = import_cost - export_revenue
|
||||
|
||||
pricing_snapshot = {
|
||||
"kind": "manual",
|
||||
"buy_dal": str(buy_dal),
|
||||
"buy_normal": str(buy_normal),
|
||||
"sell_dal": str(sell_dal),
|
||||
"sell_normal": str(sell_normal),
|
||||
"energy_tax": str(energy_tax),
|
||||
"ode": str(ode),
|
||||
}
|
||||
|
||||
return {
|
||||
"import_cost": import_cost,
|
||||
"export_revenue": export_revenue,
|
||||
"net_cost": net_cost,
|
||||
"pricing": pricing_snapshot,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tibber strategy — dynamic 15-minute spot price
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TibberPriceNotFoundError(LookupError):
|
||||
"""Raised when no TibberPrice row covers the requested period start.
|
||||
|
||||
The billing engine (T07) catches this to mark the period as ``degraded``
|
||||
or skip it until a price becomes available.
|
||||
"""
|
||||
|
||||
|
||||
def _tibber_strategy(
|
||||
deltas: PeriodDeltas,
|
||||
t0: datetime,
|
||||
values: dict[str, Any],
|
||||
session: Session,
|
||||
) -> dict[str, Any]:
|
||||
"""Calculate billing costs for one period using Tibber dynamic pricing.
|
||||
|
||||
Price lookup:
|
||||
Takes the most recent ``TibberPrice`` row where ``starts_at ≤ t0``
|
||||
(i.e. the slot that was in effect at *t0*). This is a DESC LIMIT 1
|
||||
query on ``starts_at``.
|
||||
|
||||
Formula (§3.4):
|
||||
- ``buy = total`` (Tibber's all-inclusive price; already includes energy
|
||||
tax, VAT and the buy-side ``inkoopvergoeding``)
|
||||
- ``sell = total − energy_tax − sell_fee − sell_adjust``
|
||||
- ``import_cost = (Δd1 + Δd2) × buy``
|
||||
- ``export_revenue = (Δr1 + Δr2) × sell``
|
||||
- ``net_cost = import_cost − export_revenue``
|
||||
|
||||
``sell_fee`` models Tibber's per-kWh **verkoopvergoeding** (feed-in fee,
|
||||
€0.0248/kWh incl. VAT since 2026-01-01). It is always deducted from the
|
||||
feed-in payout: even under the net-metering (saldering) scheme, Tibber pays
|
||||
``total − verkoopvergoeding`` per returned kWh (Tibber NL: "€0,28 − €0,0248
|
||||
= €0,2552"). ``total`` already contains the equal buy-side
|
||||
``inkoopvergoeding``, so the two fees do **not** cancel — the feed-in price
|
||||
sits ``sell_fee`` below the buy price. ``sell_adjust`` is a separate manual
|
||||
correction: under net metering it carries back the refunded energy tax
|
||||
(``sell_adjust = −energy_tax``), leaving ``sell = total − sell_fee``.
|
||||
|
||||
Tibber does not differentiate tariff slots (dal vs normal) — the 15-minute
|
||||
API price applies to the full delivered/returned volume.
|
||||
|
||||
Negative ``total`` (extreme negative spot prices):
|
||||
When ``total`` is negative the ``sell`` price will also be negative,
|
||||
meaning ``export_revenue`` becomes negative (feeding to grid *costs*
|
||||
money). This is the mathematically correct outcome and is left as-is.
|
||||
|
||||
Raises
|
||||
------
|
||||
TibberPriceNotFoundError
|
||||
If no ``TibberPrice`` row exists with ``starts_at ≤ t0``. The caller
|
||||
(T07 billing engine) should catch this and mark the period as
|
||||
``degraded`` or skip it for later recomputation.
|
||||
"""
|
||||
from app.models.energy import TibberPrice # local import to avoid circular
|
||||
from sqlalchemy import desc
|
||||
|
||||
price_row: TibberPrice | None = (
|
||||
session.query(TibberPrice)
|
||||
.filter(TibberPrice.starts_at <= t0)
|
||||
.order_by(desc(TibberPrice.starts_at))
|
||||
.first()
|
||||
)
|
||||
|
||||
if price_row is None:
|
||||
raise TibberPriceNotFoundError(
|
||||
f"No TibberPrice found covering t0={t0.isoformat()!r}; "
|
||||
"period cannot be billed until price data is available."
|
||||
)
|
||||
|
||||
energy = values.get("energy", {})
|
||||
energy_tax = _to_decimal(energy.get("energy_tax", 0))
|
||||
sell_fee = _to_decimal(energy.get("sell_fee", 0))
|
||||
sell_adjust = _to_decimal(energy.get("sell_adjust", 0))
|
||||
|
||||
total = _to_decimal(price_row.total)
|
||||
buy = total
|
||||
sell = total - energy_tax - sell_fee - sell_adjust
|
||||
|
||||
total_delivered = deltas.d1 + deltas.d2
|
||||
total_returned = deltas.r1 + deltas.r2
|
||||
|
||||
import_cost = total_delivered * buy
|
||||
export_revenue = total_returned * sell
|
||||
net_cost = import_cost - export_revenue
|
||||
|
||||
pricing_snapshot = {
|
||||
"kind": "tibber",
|
||||
"tibber_price_starts_at": price_row.starts_at.isoformat(),
|
||||
"tibber_price_id": price_row.id,
|
||||
"total": str(total),
|
||||
"buy": str(buy),
|
||||
"sell": str(sell),
|
||||
"energy_tax": str(energy_tax),
|
||||
"sell_fee": str(sell_fee),
|
||||
"sell_adjust": str(sell_adjust),
|
||||
}
|
||||
|
||||
return {
|
||||
"import_cost": import_cost,
|
||||
"export_revenue": export_revenue,
|
||||
"net_cost": net_cost,
|
||||
"pricing": pricing_snapshot,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Register built-in strategies at module import time
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
register_strategy("manual", _manual_strategy)
|
||||
register_strategy("tibber", _tibber_strategy)
|
||||
@@ -0,0 +1,10 @@
|
||||
"""Tibber dynamic electricity pricing integration.
|
||||
|
||||
This package provides:
|
||||
- ``client``: HTTP client for the Tibber GraphQL API (price fetching).
|
||||
- Custom exceptions: ``TibberError`` (base) and ``TibberAuthError`` (auth failure).
|
||||
|
||||
Usage::
|
||||
|
||||
from app.integrations.tibber.client import fetch_price_range, TibberAuthError
|
||||
"""
|
||||
@@ -0,0 +1,357 @@
|
||||
"""Tibber GraphQL API client for fetching electricity price data.
|
||||
|
||||
Design decisions
|
||||
----------------
|
||||
- Uses ``httpx`` (already a project dependency) for all HTTP requests.
|
||||
- Token is **never** written to log messages or exception strings to prevent
|
||||
credential leakage into log aggregators.
|
||||
- ``fetch_price_range`` returns a list of ``PricePoint`` objects with UTC-aware
|
||||
``starts_at`` datetimes; the number of nodes is not assumed — all returned
|
||||
nodes are parsed regardless of count.
|
||||
- ``fetch_current_price`` uses a separate, simpler query that asks for the
|
||||
*current* price point only; used by the connection-test endpoint (T09) to
|
||||
produce a fast three-state result (success / auth-error / network-error).
|
||||
- Authentication failures (HTTP 401/403) are raised as ``TibberAuthError`` so
|
||||
that callers (the test endpoint) can distinguish them from generic network or
|
||||
API errors.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_TIBBER_API_URL = "https://api.tibber.com/v1-beta/gql"
|
||||
_DEFAULT_TIMEOUT = 15.0
|
||||
|
||||
# GraphQL query to fetch the forward-looking today + tomorrow price curve at
|
||||
# 15-minute resolution.
|
||||
#
|
||||
# ``priceInfo(resolution: QUARTER_HOURLY)`` returns two node lists:
|
||||
# * ``today`` — always the full current local day (96 quarter-hourly slots,
|
||||
# 00:00 → 23:45 local), regardless of the current time.
|
||||
# * ``tomorrow`` — the full next local day (96 slots) once Tibber publishes the
|
||||
# day-ahead prices (around 13:00–15:00 local); empty before that.
|
||||
#
|
||||
# This is deliberately NOT ``priceInfoRange``: that field is a historical cursor
|
||||
# connection whose range ends at "now" (it never returns future slots), so it
|
||||
# cannot supply upcoming prices. ``priceInfo`` is forward-looking, so every
|
||||
# 15-minute slot's price is present in the DB *before* the slot closes — which is
|
||||
# what makes per-slot billing accurate (each period finds its own exact slot
|
||||
# instead of falling back to a stale earlier price) and keeps the live current-
|
||||
# price entity fresh. The hourly refresh job re-runs this query, so tomorrow's
|
||||
# prices are picked up within an hour of publication without a restart.
|
||||
_PRICE_RANGE_QUERY = """
|
||||
{
|
||||
viewer {
|
||||
homes {
|
||||
id
|
||||
currentSubscription {
|
||||
priceInfo(resolution: QUARTER_HOURLY) {
|
||||
today {
|
||||
startsAt
|
||||
total
|
||||
energy
|
||||
tax
|
||||
currency
|
||||
level
|
||||
}
|
||||
tomorrow {
|
||||
startsAt
|
||||
total
|
||||
energy
|
||||
tax
|
||||
currency
|
||||
level
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
"""
|
||||
|
||||
# GraphQL query to fetch the single *current* price point.
|
||||
_CURRENT_PRICE_QUERY = """
|
||||
{
|
||||
viewer {
|
||||
homes {
|
||||
id
|
||||
currentSubscription {
|
||||
priceInfo {
|
||||
current {
|
||||
startsAt
|
||||
total
|
||||
energy
|
||||
tax
|
||||
currency
|
||||
level
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
"""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Custom exceptions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TibberError(Exception):
|
||||
"""Base class for all Tibber client errors.
|
||||
|
||||
Raised for network failures, timeouts, unexpected HTTP status codes, and
|
||||
malformed API responses. The exception message will **never** contain the
|
||||
API token.
|
||||
"""
|
||||
|
||||
|
||||
class TibberAuthError(TibberError):
|
||||
"""Raised when the Tibber API rejects the provided token (HTTP 401/403).
|
||||
|
||||
Callers that implement a three-state connection test should catch this
|
||||
exception separately from ``TibberError`` to distinguish auth problems
|
||||
from network / API problems.
|
||||
"""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Data model
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class PricePoint:
|
||||
"""One 15-minute price slot returned by the Tibber API.
|
||||
|
||||
All monetary values are in ``currency`` and include VAT (user-facing).
|
||||
``starts_at`` is always a timezone-aware UTC datetime regardless of the
|
||||
timezone offset that Tibber returns in ``startsAt``.
|
||||
"""
|
||||
|
||||
starts_at: datetime # UTC-aware
|
||||
total: float # full all-in price (energy + tax)
|
||||
energy: float # spot energy component
|
||||
tax: float # tax component
|
||||
currency: str # ISO 4217 (typically "EUR")
|
||||
level: str | None # e.g. "CHEAP", "NORMAL", "EXPENSIVE"; None if absent
|
||||
resolution: str # label for the slot resolution (e.g. "QUARTER_HOURLY")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _parse_starts_at(raw: str) -> datetime:
|
||||
"""Parse an ISO 8601 timestamp with timezone offset and convert to UTC.
|
||||
|
||||
Tibber returns timestamps like ``2026-06-23T00:00:00.000+02:00``.
|
||||
``datetime.fromisoformat`` handles this format in Python 3.11+; the result
|
||||
is then converted to UTC via ``.astimezone(UTC)``.
|
||||
"""
|
||||
return datetime.fromisoformat(raw).astimezone(UTC)
|
||||
|
||||
|
||||
def _parse_node(node: dict[str, Any], resolution: str) -> PricePoint:
|
||||
"""Parse a single price node dict into a ``PricePoint``."""
|
||||
return PricePoint(
|
||||
starts_at=_parse_starts_at(node["startsAt"]),
|
||||
total=float(node["total"]),
|
||||
energy=float(node["energy"]),
|
||||
tax=float(node["tax"]),
|
||||
currency=str(node["currency"]),
|
||||
level=node.get("level") or None,
|
||||
resolution=resolution,
|
||||
)
|
||||
|
||||
|
||||
def _pick_home(homes: list[dict[str, Any]], home_id: str | None) -> dict[str, Any]:
|
||||
"""Return the home dict matching *home_id*, or the first home if None."""
|
||||
if not homes:
|
||||
raise TibberError("Tibber API returned no homes")
|
||||
if home_id is not None:
|
||||
for h in homes:
|
||||
if h.get("id") == home_id:
|
||||
return h
|
||||
raise TibberError("Tibber home id not found in API response")
|
||||
return homes[0]
|
||||
|
||||
|
||||
def _post_graphql(token: str, query: str, timeout: float) -> dict[str, Any]:
|
||||
"""POST the GraphQL *query* and return the parsed ``data`` dict.
|
||||
|
||||
Raises
|
||||
------
|
||||
TibberAuthError
|
||||
On HTTP 401 or 403.
|
||||
TibberError
|
||||
On timeouts, network errors, non-2xx responses, or unexpected body shape.
|
||||
"""
|
||||
headers = {
|
||||
"Authorization": f"Bearer {token}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
try:
|
||||
response = httpx.post(
|
||||
_TIBBER_API_URL,
|
||||
json={"query": query},
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
except httpx.TimeoutException as exc:
|
||||
raise TibberError("Tibber API request timed out") from exc
|
||||
except httpx.HTTPError as exc:
|
||||
raise TibberError("Tibber API request failed") from exc
|
||||
|
||||
if response.status_code in (401, 403):
|
||||
raise TibberAuthError("Tibber API authentication failed")
|
||||
|
||||
try:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise TibberError(f"Tibber API returned unexpected status {response.status_code}") from exc
|
||||
|
||||
try:
|
||||
body = response.json()
|
||||
except Exception as exc:
|
||||
raise TibberError("Tibber API returned non-JSON response") from exc
|
||||
|
||||
if "errors" in body:
|
||||
# GraphQL errors are not HTTP errors; surface them as TibberError.
|
||||
# Do not include token in the message.
|
||||
raise TibberError("Tibber GraphQL returned errors")
|
||||
|
||||
data = body.get("data")
|
||||
if data is None:
|
||||
raise TibberError("Tibber API response missing 'data' field")
|
||||
|
||||
return data
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public API
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def fetch_price_range(
|
||||
token: str,
|
||||
home_id: str | None = None,
|
||||
*,
|
||||
timeout: float = _DEFAULT_TIMEOUT,
|
||||
) -> list[PricePoint]:
|
||||
"""Fetch the forward-looking today + tomorrow 15-minute price curve from Tibber.
|
||||
|
||||
Sends the ``priceInfo(resolution: QUARTER_HOURLY) { today tomorrow }`` query
|
||||
and parses every node from both lists (today first, then tomorrow) into a
|
||||
``PricePoint``. ``priceInfo`` is forward-looking — ``today`` is always the
|
||||
full current local day and ``tomorrow`` is populated once Tibber publishes the
|
||||
day-ahead prices — so upcoming slots are returned, unlike ``priceInfoRange``
|
||||
which only reaches "now". ``tomorrow`` may be empty (before publication); the
|
||||
number of nodes is not assumed and all returned nodes are parsed.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
token:
|
||||
Tibber API token. **Never** logged or included in exception messages.
|
||||
home_id:
|
||||
If given, the home with this Tibber home ID is selected. If ``None``,
|
||||
the first home in the account is used.
|
||||
timeout:
|
||||
HTTP request timeout in seconds (default 15 s).
|
||||
|
||||
Returns
|
||||
-------
|
||||
list[PricePoint]
|
||||
List of price points with UTC-aware ``starts_at`` values.
|
||||
|
||||
Raises
|
||||
------
|
||||
TibberAuthError
|
||||
If the token is rejected (HTTP 401/403).
|
||||
TibberError
|
||||
For network failures, timeouts, or unexpected API responses.
|
||||
"""
|
||||
data = _post_graphql(token, _PRICE_RANGE_QUERY, timeout)
|
||||
|
||||
try:
|
||||
homes = data["viewer"]["homes"]
|
||||
except (KeyError, TypeError) as exc:
|
||||
raise TibberError("Tibber API response has unexpected shape") from exc
|
||||
|
||||
home = _pick_home(homes, home_id)
|
||||
|
||||
try:
|
||||
price_info = home["currentSubscription"]["priceInfo"]
|
||||
today = price_info["today"]
|
||||
tomorrow = price_info["tomorrow"]
|
||||
except (KeyError, TypeError) as exc:
|
||||
raise TibberError("Tibber API response missing priceInfo today/tomorrow") from exc
|
||||
|
||||
# tomorrow is null/empty until Tibber publishes the day-ahead prices; treat
|
||||
# a missing list as empty so we still return today's slots.
|
||||
nodes = list(today or []) + list(tomorrow or [])
|
||||
return [_parse_node(node, "QUARTER_HOURLY") for node in nodes]
|
||||
|
||||
|
||||
def fetch_current_price(
|
||||
token: str,
|
||||
home_id: str | None = None,
|
||||
*,
|
||||
timeout: float = _DEFAULT_TIMEOUT,
|
||||
) -> PricePoint:
|
||||
"""Fetch the *current* price point from the Tibber API.
|
||||
|
||||
Uses the lighter ``priceInfo { current { ... } }`` query rather than the
|
||||
full range query, making it suitable for a fast connection test.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
token:
|
||||
Tibber API token. **Never** logged or included in exception messages.
|
||||
home_id:
|
||||
If given, the home with this Tibber home ID is selected. If ``None``,
|
||||
the first home in the account is used.
|
||||
timeout:
|
||||
HTTP request timeout in seconds (default 15 s).
|
||||
|
||||
Returns
|
||||
-------
|
||||
PricePoint
|
||||
The current price point with a UTC-aware ``starts_at``.
|
||||
|
||||
Raises
|
||||
------
|
||||
TibberAuthError
|
||||
If the token is rejected (HTTP 401/403).
|
||||
TibberError
|
||||
For network failures, timeouts, missing current price, or unexpected
|
||||
API responses.
|
||||
"""
|
||||
data = _post_graphql(token, _CURRENT_PRICE_QUERY, timeout)
|
||||
|
||||
try:
|
||||
homes = data["viewer"]["homes"]
|
||||
except (KeyError, TypeError) as exc:
|
||||
raise TibberError("Tibber API response has unexpected shape") from exc
|
||||
|
||||
home = _pick_home(homes, home_id)
|
||||
|
||||
try:
|
||||
current = home["currentSubscription"]["priceInfo"]["current"]
|
||||
except (KeyError, TypeError) as exc:
|
||||
raise TibberError("Tibber API response missing current price") from exc
|
||||
|
||||
if current is None:
|
||||
raise TibberError("Tibber API returned null for current price")
|
||||
|
||||
return _parse_node(current, "QUARTER_HOURLY")
|
||||
+297
-8
@@ -1,15 +1,30 @@
|
||||
import logging
|
||||
import os
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import Callable
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi import FastAPI, HTTPException, Request
|
||||
from fastapi.responses import FileResponse
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from apscheduler.schedulers.background import BackgroundScheduler
|
||||
from apscheduler.triggers.cron import CronTrigger
|
||||
from apscheduler.triggers.interval import IntervalTrigger
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app import models # noqa: F401
|
||||
from app.api.routes.auth import router as auth_router
|
||||
from app.api.routes import pages, status
|
||||
from app.api.routes.api.config import router as api_config_router
|
||||
from app.api.routes.api.data import router as api_data_router
|
||||
from app.api.routes.api.energy import router as api_energy_router
|
||||
from app.api.routes.api.energy_contracts import router as api_energy_contracts_router
|
||||
from app.api.routes.api.meter_costs import router as api_meter_costs_router
|
||||
from app.api.routes.api.expose import router as api_expose_router
|
||||
from app.api.routes.api.meters import router as api_meters_router
|
||||
from app.api.routes.api.meter_sources import router as api_meter_sources_router
|
||||
from app.api.routes.api.modbus import router as api_modbus_router
|
||||
from app.api.routes.api.session import router as api_session_router
|
||||
from app.api.routes import status
|
||||
from app.db import get_session_local
|
||||
from app.api.routes.homeassistant import router as homeassistant_router
|
||||
from app.api.routes.location import router as location_router
|
||||
@@ -17,11 +32,31 @@ from app.api.routes.poo import router as poo_router
|
||||
from app.api.routes.public_ip import router as public_ip_router
|
||||
from app.api.routes.ticktick import router as ticktick_router
|
||||
from app.config import get_settings
|
||||
from app.integrations.mqtt import mqtt_manager
|
||||
from app.services.auth import AuthBootstrapError, initialize_auth_schema
|
||||
from app.services.config_page import seed_missing_config_from_bootstrap, sync_app_hostname_from_bootstrap
|
||||
from app.services.config_page import build_runtime_settings, seed_missing_config_from_bootstrap, sync_app_hostname_from_bootstrap
|
||||
from app.services.dsmr_ingest import apply_dsmr_subscription
|
||||
from app.services.public_ip import check_public_ipv4_and_notify
|
||||
from app.services.modbus_poll import poll_all_enabled_devices, BASE_POLL_TICK_SECONDS
|
||||
from app.services.ha_discovery import publish_discovery, publish_states
|
||||
from app.services.tibber_prices import run_tibber_refresh_best_effort
|
||||
from app.services.energy_cost import compute_closed_periods
|
||||
from app.services.meter_cost import compute_closed_periods as compute_closed_meter_cost_periods
|
||||
from app.services.warmtelink_worker import warmtelink_worker_manager
|
||||
from app.services.timezone import local_tz
|
||||
from scripts.app_db_adopt import AppDatabaseAdoptionError, validate_app_runtime_db
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_REPO_ROOT = Path(__file__).resolve().parents[1]
|
||||
|
||||
|
||||
def _get_spa_dist_dir() -> Path:
|
||||
env_val = os.environ.get("SPA_DIST_DIR")
|
||||
if env_val:
|
||||
return Path(env_val)
|
||||
return _REPO_ROOT / "frontend" / "dist"
|
||||
|
||||
|
||||
def _run_scheduled_public_ip_check() -> None:
|
||||
session_local = get_session_local()
|
||||
@@ -32,6 +67,134 @@ def _run_scheduled_public_ip_check() -> None:
|
||||
session.close()
|
||||
|
||||
|
||||
def _run_scheduled_modbus_poll() -> None:
|
||||
session_local = get_session_local()
|
||||
session: Session = session_local()
|
||||
try:
|
||||
poll_all_enabled_devices(session, bootstrap_settings=get_settings())
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def _run_scheduled_tibber_refresh() -> None:
|
||||
"""Scheduled job: fetch Tibber 15-minute prices and upsert into tibber_price.
|
||||
|
||||
Runs every hour so that:
|
||||
- Today's prices are available from startup.
|
||||
- Tomorrow's prices (published by Tibber around 13:00 CET / 11:00 UTC) are
|
||||
picked up within an hour of publication without requiring a server restart.
|
||||
|
||||
The job is a no-op when:
|
||||
- No active energy contract with kind="tibber" exists.
|
||||
- The Tibber API token is empty in the runtime settings.
|
||||
|
||||
Any client exceptions (auth failures, network errors) are caught and logged
|
||||
so that a single failed fetch does not crash the scheduler or affect the
|
||||
other background jobs.
|
||||
"""
|
||||
run_tibber_refresh_best_effort()
|
||||
|
||||
|
||||
def _run_scheduled_energy_cost() -> None:
|
||||
"""Scheduled job: compute billing records for all uncalculated closed 15-minute periods.
|
||||
|
||||
Runs every minute so that a new period is picked up within 1 minute of
|
||||
closing. The job is a no-op when:
|
||||
- No active energy contract with a version covering the period exists.
|
||||
- DSMR data has not yet arrived for the period boundaries.
|
||||
- The period's billing record already exists and is not degraded.
|
||||
|
||||
Any unexpected exceptions are caught and logged so that a single failure
|
||||
does not crash the scheduler or affect the other background jobs.
|
||||
"""
|
||||
session_local = get_session_local()
|
||||
|
||||
def run_scope(label: str, operation: Callable[[Session], None]) -> None:
|
||||
"""Run one best-effort scope in an isolated transaction/session."""
|
||||
session: Session | None = None
|
||||
try:
|
||||
session = session_local()
|
||||
operation(session)
|
||||
except Exception:
|
||||
logger.exception("_run_scheduled_energy_cost: %s failed", label)
|
||||
if session is not None:
|
||||
try:
|
||||
session.rollback()
|
||||
except Exception:
|
||||
# A failed cleanup must not replace the operation/factory
|
||||
# error or prevent the following independent scope.
|
||||
logger.exception("_run_scheduled_energy_cost: %s rollback failed", label)
|
||||
finally:
|
||||
if session is not None:
|
||||
try:
|
||||
session.close()
|
||||
except Exception:
|
||||
# Sessions are intentionally isolated; close failures are
|
||||
# diagnostic only and must remain best-effort too.
|
||||
logger.exception("_run_scheduled_energy_cost: %s close failed", label)
|
||||
|
||||
# Electricity, thermal and HA publishing must not share failed transaction
|
||||
# state or accidentally commit each other's partially-flushed changes.
|
||||
run_scope("electricity computation", compute_closed_periods)
|
||||
run_scope("thermal computation", compute_closed_meter_cost_periods)
|
||||
|
||||
def publish(session: Session) -> None:
|
||||
# publish_states is internally guarded by _should_publish (MQTT
|
||||
# disabled / disconnected -> no-op), but gets a clean Session anyway.
|
||||
from app.services.ha_discovery import publish_states
|
||||
|
||||
publish_states(session)
|
||||
|
||||
run_scope("publish_states (non-fatal)", publish)
|
||||
|
||||
|
||||
def _run_scheduled_ha_state_publish() -> None:
|
||||
"""Periodic job: publish discovery configs + state + availability for all enabled exposed entities.
|
||||
|
||||
Runs every 60 seconds. When the MQTT broker is connected:
|
||||
- Publishes (or re-publishes) all HA Discovery configs (retained, idempotent).
|
||||
- Publishes the current state / availability for all enabled entities.
|
||||
|
||||
This also serves as the reliable "publish discovery after connect" mechanism:
|
||||
because paho connects asynchronously, a synchronous call immediately after
|
||||
``mqtt_manager.connect()`` would fire before the TCP handshake completes and
|
||||
be a no-op. Instead, this periodic job picks it up within 60 seconds of the
|
||||
broker becoming available — retained payloads make repeated publishes harmless.
|
||||
|
||||
Additionally, if MQTT is configured in DB but the manager is not yet connected
|
||||
(e.g. MQTT was enabled via UI after startup), this job attempts to connect so
|
||||
the user does not need to restart the server.
|
||||
"""
|
||||
session_local = get_session_local()
|
||||
session: Session = session_local()
|
||||
try:
|
||||
runtime_settings = build_runtime_settings(session, get_settings())
|
||||
# Reconnect if MQTT is configured (in DB) but not yet connected.
|
||||
if mqtt_manager.is_configured(runtime_settings) and not mqtt_manager.is_connected:
|
||||
logger.info("_run_scheduled_ha_state_publish: MQTT configured but not connected — attempting connect.")
|
||||
mqtt_manager.connect(runtime_settings)
|
||||
publish_discovery(session)
|
||||
publish_states(session)
|
||||
except Exception:
|
||||
logger.exception("_run_scheduled_ha_state_publish: unexpected error")
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def _run_midnight_state_publish() -> None:
|
||||
"""本地午夜后不久专门发布一次状态,让 *_today 的每日归零稳稳落在午夜之后
|
||||
(对 HA 钟慢几秒鲁棒)。best-effort:失败仅记日志,不影响调度器。"""
|
||||
session_local = get_session_local()
|
||||
session = session_local()
|
||||
try:
|
||||
from app.services.ha_discovery import publish_states
|
||||
publish_states(session)
|
||||
except Exception:
|
||||
logger.exception("_run_midnight_state_publish: failed (non-fatal)")
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def ensure_auth_db_ready() -> None:
|
||||
session_local = get_session_local()
|
||||
session: Session = session_local()
|
||||
@@ -67,9 +230,94 @@ async def lifespan(_: FastAPI):
|
||||
max_instances=1,
|
||||
coalesce=True,
|
||||
)
|
||||
scheduler.add_job(
|
||||
_run_scheduled_modbus_poll,
|
||||
trigger=IntervalTrigger(seconds=BASE_POLL_TICK_SECONDS),
|
||||
id="modbus-poll",
|
||||
replace_existing=True,
|
||||
max_instances=1,
|
||||
coalesce=True,
|
||||
)
|
||||
# Periodic HA state / availability publish (60-second fallback sweep).
|
||||
scheduler.add_job(
|
||||
_run_scheduled_ha_state_publish,
|
||||
trigger=IntervalTrigger(seconds=60),
|
||||
id="ha-state-publish",
|
||||
replace_existing=True,
|
||||
max_instances=1,
|
||||
coalesce=True,
|
||||
)
|
||||
# Tibber price refresh: fetch today + tomorrow every hour.
|
||||
# The job is a no-op when no active tibber contract or token is configured,
|
||||
# so it is safe to register unconditionally.
|
||||
scheduler.add_job(
|
||||
_run_scheduled_tibber_refresh,
|
||||
trigger=IntervalTrigger(hours=1),
|
||||
id="tibber-refresh",
|
||||
replace_existing=True,
|
||||
max_instances=1,
|
||||
coalesce=True,
|
||||
# APScheduler otherwise waits one full interval before its first run.
|
||||
# This preserves the hourly cadence while requesting a non-blocking
|
||||
# startup fetch as soon as the scheduler starts.
|
||||
next_run_time=datetime.now(UTC),
|
||||
)
|
||||
# Energy cost billing: compute uncalculated closed 15-minute periods every minute.
|
||||
# The job is a no-op when no active contract or DSMR data is present, so it is
|
||||
# safe to register unconditionally.
|
||||
scheduler.add_job(
|
||||
_run_scheduled_energy_cost,
|
||||
trigger=IntervalTrigger(minutes=1),
|
||||
id="energy-cost",
|
||||
replace_existing=True,
|
||||
max_instances=1,
|
||||
coalesce=True,
|
||||
)
|
||||
# Dedicated midnight publish: fire at local 00:00:10 so *_today grace (5 s) has
|
||||
# already elapsed and the day-rolled value is pushed to HA immediately, rather
|
||||
# than waiting for the next 60-second ha-state-publish sweep.
|
||||
scheduler.add_job(
|
||||
_run_midnight_state_publish,
|
||||
trigger=CronTrigger(hour=0, minute=0, second=10, timezone=local_tz()),
|
||||
id="midnight-today-publish",
|
||||
replace_existing=True,
|
||||
max_instances=1,
|
||||
coalesce=True,
|
||||
)
|
||||
scheduler.start()
|
||||
yield
|
||||
scheduler.shutdown(wait=False)
|
||||
|
||||
# MQTT: connect using DB-merged runtime settings so broker configured via UI
|
||||
# is picked up on restart (not just from env/bootstrap settings).
|
||||
# Discovery will be published by the first run of _run_scheduled_ha_state_publish
|
||||
# (within 60 s of startup), after the async paho handshake completes.
|
||||
_startup_session_local = get_session_local()
|
||||
_startup_session: Session = _startup_session_local()
|
||||
try:
|
||||
_startup_runtime_settings = build_runtime_settings(_startup_session, get_settings())
|
||||
finally:
|
||||
_startup_session.close()
|
||||
serial_started = False
|
||||
try:
|
||||
mqtt_manager.connect(_startup_runtime_settings)
|
||||
|
||||
# DSMR sources carry their own runtime configuration and are reconciled
|
||||
# after the MQTT manager is connected.
|
||||
apply_dsmr_subscription(_startup_runtime_settings)
|
||||
# Mark it before reconcile: a partial reconcile can already own a fd or
|
||||
# a non-daemon thread and must receive the same orderly shutdown.
|
||||
serial_started = True
|
||||
warmtelink_worker_manager.start()
|
||||
|
||||
yield
|
||||
finally:
|
||||
# Serial descriptors/workers must be handled first on every exit path.
|
||||
if serial_started:
|
||||
try:
|
||||
warmtelink_worker_manager.shutdown()
|
||||
except Exception:
|
||||
logger.exception("WarmteLink shutdown failed")
|
||||
mqtt_manager.disconnect()
|
||||
scheduler.shutdown(wait=False)
|
||||
|
||||
|
||||
def create_app() -> FastAPI:
|
||||
@@ -89,13 +337,54 @@ def create_app() -> FastAPI:
|
||||
app.mount("/static", StaticFiles(directory=static_dir), name="static")
|
||||
|
||||
app.include_router(status.router)
|
||||
app.include_router(auth_router)
|
||||
app.include_router(pages.router)
|
||||
app.include_router(api_config_router)
|
||||
app.include_router(api_data_router)
|
||||
app.include_router(api_energy_router)
|
||||
app.include_router(api_energy_contracts_router)
|
||||
app.include_router(api_meter_costs_router)
|
||||
app.include_router(api_meters_router)
|
||||
app.include_router(api_meter_sources_router)
|
||||
app.include_router(api_expose_router)
|
||||
app.include_router(api_modbus_router)
|
||||
app.include_router(api_session_router)
|
||||
app.include_router(homeassistant_router)
|
||||
app.include_router(location_router)
|
||||
app.include_router(poo_router)
|
||||
app.include_router(public_ip_router)
|
||||
app.include_router(ticktick_router)
|
||||
|
||||
# SPA hosting: mount frontend/dist if it exists and has index.html.
|
||||
# If the SPA dist is absent (e.g. backend-only CI), skip SPA serving entirely
|
||||
# so that pytest stays green with only the API routes registered.
|
||||
spa_dist = _get_spa_dist_dir()
|
||||
spa_index = spa_dist / "index.html"
|
||||
if spa_dist.is_dir() and spa_index.is_file():
|
||||
spa_assets = spa_dist / "assets"
|
||||
if spa_assets.is_dir():
|
||||
app.mount("/assets", StaticFiles(directory=spa_assets), name="spa-assets")
|
||||
|
||||
# Resolve the dist root once so the containment check is fast and consistent.
|
||||
_spa_root = spa_dist.resolve()
|
||||
|
||||
@app.get("/{full_path:path}", include_in_schema=False)
|
||||
async def spa_fallback(full_path: str, request: Request) -> FileResponse: # noqa: RUF029
|
||||
# Explicit 404 for unmatched /api/* — never return index.html for API paths.
|
||||
if full_path.startswith("api/"):
|
||||
raise HTTPException(status_code=404, detail="not found")
|
||||
# Resolve candidate to an absolute path and verify it stays within the SPA
|
||||
# dist root. Without this check, URL-encoded ".." sequences (e.g. "..%2f")
|
||||
# bypass Starlette's path parameter handling and allow arbitrary file reads.
|
||||
candidate = (spa_dist / full_path).resolve()
|
||||
if candidate.is_file() and candidate.is_relative_to(_spa_root):
|
||||
return FileResponse(candidate)
|
||||
# For any path outside the dist root, or for SPA client routes that don't
|
||||
# correspond to a real file, return index.html so the SPA router handles it.
|
||||
return FileResponse(spa_index)
|
||||
else:
|
||||
logger.warning(
|
||||
"SPA dist not found at %s — SPA hosting disabled (API-only mode).", spa_dist
|
||||
)
|
||||
|
||||
return app
|
||||
|
||||
|
||||
|
||||
@@ -5,12 +5,16 @@ from app.models.config import AppConfigEntry
|
||||
from app.models.location import Location
|
||||
from app.models.poo import PooRecord
|
||||
from app.models.public_ip import PublicIPHistory, PublicIPState
|
||||
from app.models.meter_source import MeterSource, MeterSourceBinding, MeterSourceChannel
|
||||
|
||||
__all__ = [
|
||||
"AppConfigEntry",
|
||||
"AuthSession",
|
||||
"AuthUser",
|
||||
"Location",
|
||||
"MeterSource",
|
||||
"MeterSourceBinding",
|
||||
"MeterSourceChannel",
|
||||
"PooRecord",
|
||||
"PublicIPHistory",
|
||||
"PublicIPState",
|
||||
|
||||
@@ -15,8 +15,12 @@ class AuthUser(Base):
|
||||
is_active: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True)
|
||||
force_password_change: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
# TOTP fields (Phase B, M4-T04) — nullable/false by default so existing users are unaffected
|
||||
totp_secret: Mapped[str | None] = mapped_column(String(64), nullable=True)
|
||||
totp_enabled: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
|
||||
|
||||
sessions: Mapped[list["AuthSession"]] = relationship(back_populates="user")
|
||||
recovery_codes: Mapped[list["RecoveryCode"]] = relationship(back_populates="user")
|
||||
|
||||
|
||||
class AuthSession(Base):
|
||||
@@ -31,3 +35,21 @@ class AuthSession(Base):
|
||||
revoked_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
user: Mapped[AuthUser] = relationship(back_populates="sessions")
|
||||
|
||||
|
||||
class RecoveryCode(Base):
|
||||
"""One-time TOTP recovery codes for AuthUser.
|
||||
|
||||
``code_hash`` stores the Argon2 hash of the plaintext code (plaintext is
|
||||
returned only at setup time and never stored). ``used_at`` is NULL while
|
||||
the code is still valid; set to the consumption timestamp when consumed.
|
||||
"""
|
||||
|
||||
__tablename__ = "auth_recovery_code"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
user_id: Mapped[int] = mapped_column(ForeignKey("auth_users.id"), nullable=False, index=True)
|
||||
code_hash: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
used_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
user: Mapped[AuthUser] = relationship(back_populates="recovery_codes")
|
||||
|
||||
@@ -0,0 +1,31 @@
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import DateTime, Index, Integer, String, UniqueConstraint
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.db import Base
|
||||
|
||||
|
||||
class LoginThrottle(Base):
|
||||
"""Tracks per-key login failure state for exponential back-off throttling.
|
||||
|
||||
``scope`` is either ``'ip'`` or ``'user'``; ``key`` is the IP address or
|
||||
username string respectively. The pair ``(scope, key)`` is unique — one
|
||||
row per tracked entity. A row is deleted on successful login (clear).
|
||||
"""
|
||||
|
||||
__tablename__ = "auth_login_throttle"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("scope", "key", name="uq_auth_login_throttle_scope_key"),
|
||||
Index("ix_auth_login_throttle_scope_key", "scope", "key"),
|
||||
)
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
key: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
scope: Mapped[str] = mapped_column(String(16), nullable=False) # 'ip' | 'user'
|
||||
failures: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
|
||||
first_failed_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
last_failed_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
next_allowed_at: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True), nullable=True
|
||||
)
|
||||
@@ -0,0 +1,558 @@
|
||||
"""SQLAlchemy models for the energy pricing and DSMR metering subsystem.
|
||||
|
||||
Six tables:
|
||||
- meter: physical electricity meter lifecycle epoch.
|
||||
- dsmr_reading: raw DSMR telegram blobs (10-second down-sampled).
|
||||
- energy_contract: contract head (manual or tibber, one active at a time).
|
||||
- energy_contract_version: versioned pricing values; append-only for auditability.
|
||||
- tibber_price: cached Tibber 15-minute spot prices (immutable).
|
||||
- energy_cost_period: computed 15-minute billing periods (immutable snapshot).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid as _uuid
|
||||
from datetime import datetime, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
from sqlalchemy import (
|
||||
Boolean,
|
||||
CheckConstraint,
|
||||
DateTime,
|
||||
Float,
|
||||
ForeignKey,
|
||||
Index,
|
||||
Integer,
|
||||
Numeric,
|
||||
String,
|
||||
UniqueConstraint,
|
||||
event,
|
||||
text,
|
||||
)
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship, synonym, validates
|
||||
from sqlalchemy.types import JSON, TypeDecorator
|
||||
|
||||
from app.db import Base
|
||||
from app.models.meter_source import MeterSourceBinding
|
||||
|
||||
|
||||
def _uuid4_str() -> str:
|
||||
return str(_uuid.uuid4())
|
||||
|
||||
|
||||
def _decimal_json(value: Any) -> Any:
|
||||
"""Make auditable JSON portable without admitting binary numeric values."""
|
||||
if isinstance(value, Decimal):
|
||||
return format(value, "f")
|
||||
if isinstance(value, float) or (isinstance(value, int) and not isinstance(value, bool)):
|
||||
raise ValueError("JSON amounts and quantities must be decimal strings, not numeric JSON values")
|
||||
if isinstance(value, dict):
|
||||
return {key: _decimal_json(child) for key, child in value.items()}
|
||||
if isinstance(value, list):
|
||||
return [_decimal_json(child) for child in value]
|
||||
return value
|
||||
|
||||
|
||||
def _validate_fixed_decimal(value: Decimal, precision: int, scale: int, field: str) -> Decimal:
|
||||
if not isinstance(value, Decimal):
|
||||
raise ValueError(f"{field} must be a Decimal, not a binary float or other numeric type")
|
||||
if not value.is_finite():
|
||||
raise ValueError(f"{field} must be finite")
|
||||
if -value.as_tuple().exponent > scale:
|
||||
raise ValueError(f"{field} exceeds scale {scale}")
|
||||
integer_digits = 0 if value.is_zero() else max(value.copy_abs().adjusted() + 1, 0)
|
||||
if integer_digits > precision - scale:
|
||||
raise ValueError(f"{field} exceeds precision {precision},{scale}")
|
||||
return value
|
||||
|
||||
|
||||
class ExactDecimal(TypeDecorator[Decimal]):
|
||||
"""Fixed-point Decimal which uses SQLite TEXT, never a binary float."""
|
||||
|
||||
impl = Numeric
|
||||
cache_ok = True
|
||||
|
||||
def __init__(self, precision: int, scale: int) -> None:
|
||||
self.precision = precision
|
||||
self.scale = scale
|
||||
super().__init__(precision=precision, scale=scale)
|
||||
|
||||
def load_dialect_impl(self, dialect):
|
||||
if dialect.name == "sqlite":
|
||||
return dialect.type_descriptor(String(self.precision + 2))
|
||||
return dialect.type_descriptor(Numeric(self.precision, self.scale, asdecimal=True))
|
||||
|
||||
def process_bind_param(self, value: Decimal | None, dialect) -> Decimal | str | None:
|
||||
if value is None:
|
||||
return None
|
||||
value = _validate_fixed_decimal(value, self.precision, self.scale, "decimal value")
|
||||
if dialect.name == "sqlite":
|
||||
return format(value, f".{self.scale}f")
|
||||
return value
|
||||
|
||||
def process_result_value(self, value: Decimal | str | None, _dialect) -> Decimal | None:
|
||||
return None if value is None else Decimal(value)
|
||||
|
||||
|
||||
class DecimalJSON(TypeDecorator[dict]):
|
||||
"""JSON which serializes Decimal values as strings on every write path."""
|
||||
|
||||
impl = JSON
|
||||
cache_ok = True
|
||||
|
||||
def process_bind_param(self, value: Any, _dialect) -> Any:
|
||||
return None if value is None else _decimal_json(value)
|
||||
|
||||
|
||||
class UTCDateTime(TypeDecorator[datetime]):
|
||||
"""UTC timestamps that preserve instant identity on SQLite and other dialects."""
|
||||
|
||||
impl = DateTime(timezone=True)
|
||||
cache_ok = True
|
||||
|
||||
def __init__(self, field: str) -> None:
|
||||
self.field = field
|
||||
super().__init__()
|
||||
|
||||
def process_bind_param(self, value: datetime | None, _dialect) -> datetime | None:
|
||||
return None if value is None else _normalise_utc_period(value, self.field)
|
||||
|
||||
def process_result_value(self, value: datetime | None, _dialect) -> datetime | None:
|
||||
if value is None:
|
||||
return None
|
||||
if value.tzinfo is None or value.utcoffset() is None:
|
||||
return value.replace(tzinfo=timezone.utc)
|
||||
return value.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def _require_aware_period(start: datetime, end: datetime) -> None:
|
||||
if start.tzinfo is None or start.utcoffset() is None:
|
||||
raise ValueError("period_start must be timezone-aware")
|
||||
if end.tzinfo is None or end.utcoffset() is None:
|
||||
raise ValueError("period_end must be timezone-aware")
|
||||
if end <= start:
|
||||
raise ValueError("period_end must be after period_start")
|
||||
|
||||
|
||||
def _normalise_utc_period(value: datetime, field: str) -> datetime:
|
||||
if value.tzinfo is None or value.utcoffset() is None:
|
||||
raise ValueError(f"{field} must be timezone-aware")
|
||||
return value.astimezone(timezone.utc)
|
||||
|
||||
|
||||
class Meter(Base):
|
||||
"""One physical electricity meter's installation epoch.
|
||||
|
||||
A ``meter`` record represents the period ``[started_at, ended_at)`` during
|
||||
which a particular physical meter was installed and active. Replacing a meter
|
||||
(swap, home move, etc.) is modelled by closing the current record
|
||||
(``ended_at = swap_timestamp``) and opening a new one
|
||||
(``started_at = swap_timestamp``).
|
||||
|
||||
**Invariant**: for each ``commodity`` there is at most one active meter
|
||||
(``ended_at IS NULL``) at any point in time. The service layer enforces
|
||||
this — no DB-level constraint is added to keep the migration simple and to
|
||||
allow the application to return a meaningful error message.
|
||||
|
||||
``commodity`` defaults to ``"electricity"``; the field is a free-form string
|
||||
(no CHECK constraint) so future commodities (``gas``, ``heating``) can be
|
||||
added without a schema change.
|
||||
|
||||
``reason`` captures why this epoch started — one of ``initial``,
|
||||
``meter_swap``, ``home_move``, or ``other`` — stored as a plain string so
|
||||
the application layer controls the allowed set.
|
||||
"""
|
||||
|
||||
__tablename__ = "meter"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
|
||||
# Stable internal identity — used as HA Discovery unique_id anchor.
|
||||
uuid: Mapped[str] = mapped_column(String(36), unique=True, nullable=False, default=_uuid4_str)
|
||||
|
||||
# Human-readable label for this physical meter (e.g. address, serial, tariff zone).
|
||||
label: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
|
||||
# Energy commodity this meter measures. Defaults to "electricity".
|
||||
commodity: Mapped[str] = mapped_column(String(32), nullable=False, default="electricity")
|
||||
|
||||
# UTC timestamp when this meter epoch starts (inclusive). May be in the past
|
||||
# (retroactive declaration); effective billing start = max(started_at, data start).
|
||||
started_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
|
||||
# UTC timestamp when this meter epoch ends (exclusive). NULL = currently active.
|
||||
ended_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
# Why this epoch was created. Application-layer validation enforces the
|
||||
# allowed set; no CHECK constraint to keep migrations simple.
|
||||
reason: Mapped[str] = mapped_column(String(64), nullable=False)
|
||||
|
||||
# Free-form note (e.g. location, physical meter id, reason details).
|
||||
note: Mapped[str | None] = mapped_column(String(1024), nullable=True)
|
||||
|
||||
# UTC timestamp of when this row was created.
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
|
||||
# Relationship to cost periods attributed to this meter epoch (not loaded eagerly).
|
||||
cost_periods: Mapped[list["EnergyCostPeriod"]] = relationship(
|
||||
back_populates="meter", cascade="save-update, merge"
|
||||
)
|
||||
|
||||
source_bindings: Mapped[list["MeterSourceBinding"]] = relationship(
|
||||
back_populates="meter", cascade="save-update, merge"
|
||||
)
|
||||
|
||||
|
||||
class DsmrReading(Base):
|
||||
"""One down-sampled DSMR telegram stored as a full JSON blob.
|
||||
|
||||
Identity & idempotency are **independent of the DSMR Reader's telegram id**
|
||||
(that field overflows and must be manually reset to zero — a known DSMR
|
||||
quirk — so relying on it for uniqueness risks silently dropping new data).
|
||||
The table's own autoincrement ``id`` PK is the stable internal identity, and
|
||||
``(meter_source_id, recorded_at)`` is the UNIQUE de-duplication key: each
|
||||
configured P1 source emits at most one telegram per timestamp, while
|
||||
different sources may legitimately emit at the same instant.
|
||||
|
||||
``recorded_at`` is a real column (not inside the payload) so time-range
|
||||
queries are efficient. The entire telegram frame is stored verbatim in
|
||||
``payload``; no field allow-list is applied so future commodities (gas,
|
||||
heating, three-phase) are accommodated without a schema change.
|
||||
"""
|
||||
|
||||
__tablename__ = "dsmr_reading"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
|
||||
# UTC timestamp of the sample. Idempotency is per configured source, so
|
||||
# distinct P1 sources may legitimately emit at the same instant.
|
||||
recorded_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
|
||||
# Telegram's own id (DSMR Reader assigns it). Stored only as a reference /
|
||||
# debugging aid — NOT used for uniqueness or idempotency (it overflows and
|
||||
# gets reset to zero). Nullable because some DSMR sources may not emit one.
|
||||
telegram_id: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
|
||||
# Compatibility for the pre-M8 ingest implementation. This is an ORM
|
||||
# alias only; the physical database column is ``telegram_id``.
|
||||
source_id = synonym("telegram_id")
|
||||
|
||||
# The configured source is the durable identity of the cumulative reading
|
||||
# stream. It is non-null after the revision-16 historical adoption.
|
||||
meter_source_id: Mapped[int] = mapped_column(
|
||||
ForeignKey("meter_source.id", ondelete="RESTRICT"), nullable=False, index=True
|
||||
)
|
||||
|
||||
# Full telegram frame as a JSON object; values are typically JSON strings
|
||||
# (e.g. "20915.154") — callers must cast to Decimal before arithmetic.
|
||||
payload: Mapped[dict] = mapped_column(JSON, nullable=False)
|
||||
|
||||
__table_args__ = (
|
||||
UniqueConstraint(
|
||||
"meter_source_id", "recorded_at", name="uq_dsmr_reading_source_recorded_at"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@event.listens_for(DsmrReading, "before_insert")
|
||||
def _supply_legacy_dsmr_source(_mapper, connection, target: DsmrReading) -> None:
|
||||
"""Keep the pre-T04 single-source writer working during the schema handoff."""
|
||||
if target.meter_source_id is None:
|
||||
target.meter_source_id = connection.execute(
|
||||
text("SELECT id FROM meter_source WHERE kind = 'dsmr_mqtt' ORDER BY id LIMIT 1")
|
||||
).scalar_one()
|
||||
|
||||
|
||||
class EnergyContract(Base):
|
||||
"""Contract head: a named energy contract with a chosen pricing strategy.
|
||||
|
||||
``kind`` determines which price strategy is used (``manual`` for fixed
|
||||
dual-tariff rates entered by the user, ``tibber`` for dynamic API prices).
|
||||
A contract belongs to an energy ``scope`` (currently electricity; thermal
|
||||
profiles are reserved for the next milestone). Only one contract may be
|
||||
``active`` per scope; the service layer enforces mutual exclusion. Specific
|
||||
pricing values live in ``EnergyContractVersion``
|
||||
so that price changes can be tracked without modifying historical records.
|
||||
"""
|
||||
|
||||
__tablename__ = "energy_contract"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
|
||||
# Human-readable label; freely editable by the user.
|
||||
name: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
|
||||
# Strategy selector: "manual" or "tibber". Application-layer validation
|
||||
# enforces the allowed set; no DB CHECK constraint is added to keep the
|
||||
# migration simple and the strategy registry extensible.
|
||||
kind: Mapped[str] = mapped_column(String(32), nullable=False)
|
||||
|
||||
# Billing domain. The service registry derives this from ``kind`` so API
|
||||
# callers cannot move a pricing strategy into an incompatible domain.
|
||||
scope: Mapped[str] = mapped_column(
|
||||
String(32), nullable=False, default="electricity", index=True
|
||||
)
|
||||
|
||||
# Whether this is the currently active contract (at most one should be True).
|
||||
active: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
|
||||
|
||||
# ISO 4217 currency code for all monetary values in this contract.
|
||||
currency: Mapped[str] = mapped_column(String(8), nullable=False, default="EUR")
|
||||
|
||||
# Audit timestamps.
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
|
||||
# Relationship to versions (back-reference; not loaded eagerly).
|
||||
versions: Mapped[list["EnergyContractVersion"]] = relationship(
|
||||
back_populates="contract", cascade="save-update, merge"
|
||||
)
|
||||
|
||||
|
||||
class EnergyContractVersion(Base):
|
||||
"""One time-bounded version of an energy contract's pricing values.
|
||||
|
||||
Pricing changes are modelled as new versions (append-only); existing versions
|
||||
are never modified so that historical ``EnergyCostPeriod`` records remain
|
||||
fully auditable. ``effective_to`` is ``NULL`` for the currently open version.
|
||||
"""
|
||||
|
||||
__tablename__ = "energy_contract_version"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
|
||||
# FK to the parent contract. RESTRICT prevents deletion of a contract that
|
||||
# still has versioned pricing rows attached to it.
|
||||
contract_id: Mapped[int] = mapped_column(
|
||||
ForeignKey("energy_contract.id", ondelete="RESTRICT"), nullable=False
|
||||
)
|
||||
|
||||
# Start of this version's validity window (inclusive, UTC).
|
||||
effective_from: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
|
||||
# End of this version's validity window (exclusive, UTC). NULL means open-ended
|
||||
# (i.e. this is the most recent / current version).
|
||||
effective_to: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
# Pricing values as a JSON object conforming to the profile structure for
|
||||
# ``contract.kind`` (validated by the application layer against the YAML profile).
|
||||
values: Mapped[dict] = mapped_column(JSON, nullable=False)
|
||||
|
||||
# Creation timestamp.
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
|
||||
# Relationship back to the parent contract.
|
||||
contract: Mapped["EnergyContract"] = relationship(back_populates="versions")
|
||||
|
||||
# Relationship to cost periods that reference this version.
|
||||
cost_periods: Mapped[list["EnergyCostPeriod"]] = relationship(
|
||||
back_populates="contract_version", cascade="save-update, merge"
|
||||
)
|
||||
|
||||
|
||||
class TibberPrice(Base):
|
||||
"""Cached Tibber 15-minute spot price point (immutable once fetched).
|
||||
|
||||
``starts_at`` is unique so that upserts are idempotent. Past prices are
|
||||
never overwritten; the fetch job only adds rows for future time slots.
|
||||
"""
|
||||
|
||||
__tablename__ = "tibber_price"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
|
||||
# UTC start of the 15-minute slot; unique so upsert is idempotent.
|
||||
starts_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, unique=True
|
||||
)
|
||||
|
||||
# Resolution label as returned by the Tibber API (e.g. "QUARTER_HOURLY").
|
||||
resolution: Mapped[str] = mapped_column(String(32), nullable=False)
|
||||
|
||||
# Price components in the contract currency (all include VAT, user-facing).
|
||||
energy: Mapped[float] = mapped_column(Float, nullable=False)
|
||||
tax: Mapped[float] = mapped_column(Float, nullable=False)
|
||||
total: Mapped[float] = mapped_column(Float, nullable=False)
|
||||
|
||||
# Tibber price level (e.g. "NORMAL", "CHEAP", "EXPENSIVE"); may be absent.
|
||||
level: Mapped[str | None] = mapped_column(String(32), nullable=True)
|
||||
|
||||
# ISO 4217 currency code as returned by the API.
|
||||
currency: Mapped[str] = mapped_column(String(8), nullable=False)
|
||||
|
||||
# UTC timestamp of when this row was fetched/inserted.
|
||||
fetched_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
|
||||
|
||||
class EnergyCostPeriod(Base):
|
||||
"""Computed billing record for one 15-minute metering period (immutable snapshot).
|
||||
|
||||
Each row captures the per-register kWh deltas, the resulting import cost and
|
||||
export revenue, and a full snapshot of the pricing values used so that the
|
||||
calculation is fully auditable and reproducible without re-querying the
|
||||
contract version. Rows are written once and never modified; explicit
|
||||
recomputation via the API is the only way to overwrite a period.
|
||||
"""
|
||||
|
||||
__tablename__ = "energy_cost_period"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
|
||||
# UTC start of the 15-minute period; unique so upsert is idempotent.
|
||||
period_start: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, unique=True
|
||||
)
|
||||
|
||||
# Per-register kWh deltas for the period (end minus start of cumulative registers).
|
||||
# _1 = dal/low-tariff, _2 = normal/high-tariff (NL convention).
|
||||
d1_kwh: Mapped[float] = mapped_column(Float, nullable=False) # delivered low
|
||||
d2_kwh: Mapped[float] = mapped_column(Float, nullable=False) # delivered high
|
||||
r1_kwh: Mapped[float] = mapped_column(Float, nullable=False) # returned low
|
||||
r2_kwh: Mapped[float] = mapped_column(Float, nullable=False) # returned high
|
||||
|
||||
# Computed monetary amounts for the period (in ``currency``).
|
||||
import_cost: Mapped[float] = mapped_column(Float, nullable=False)
|
||||
export_revenue: Mapped[float] = mapped_column(Float, nullable=False)
|
||||
net_cost: Mapped[float] = mapped_column(Float, nullable=False)
|
||||
|
||||
# ISO 4217 currency code matching the contract.
|
||||
currency: Mapped[str] = mapped_column(String(8), nullable=False)
|
||||
|
||||
# Full snapshot of the pricing inputs used during computation. This makes
|
||||
# each row self-contained and auditable even if the contract is later changed.
|
||||
pricing: Mapped[dict] = mapped_column(JSON, nullable=False)
|
||||
|
||||
# FK to the exact contract version whose values were used. RESTRICT prevents
|
||||
# deletion of a version that has cost records attached. Nullable to support
|
||||
# periods computed in ``degraded`` mode (missing price data).
|
||||
contract_version_id: Mapped[int | None] = mapped_column(
|
||||
ForeignKey("energy_contract_version.id", ondelete="RESTRICT"), nullable=True
|
||||
)
|
||||
|
||||
# FK to the meter epoch this period belongs to. RESTRICT prevents deletion of
|
||||
# a meter that still has attributed cost periods. Nullable for backwards
|
||||
# compatibility (pre-M7 rows) and degraded periods where the meter was not
|
||||
# determinable.
|
||||
meter_id: Mapped[int | None] = mapped_column(
|
||||
ForeignKey("meter.id", ondelete="RESTRICT"), nullable=True
|
||||
)
|
||||
|
||||
# Nullable for historical and degraded rows. Every new normal period
|
||||
# points at the one binding that supplied both cumulative endpoints.
|
||||
source_binding_id: Mapped[int | None] = mapped_column(
|
||||
ForeignKey("meter_source_binding.id", ondelete="RESTRICT"), nullable=True
|
||||
)
|
||||
|
||||
# True when the period was computed with incomplete data (missing readings or
|
||||
# missing price); serves as a flag for later recomputation.
|
||||
degraded: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
|
||||
|
||||
# UTC timestamp of when this row was computed/inserted.
|
||||
computed_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
|
||||
# Relationship back to the contract version.
|
||||
contract_version: Mapped["EnergyContractVersion | None"] = relationship(
|
||||
back_populates="cost_periods"
|
||||
)
|
||||
|
||||
# Relationship back to the meter epoch.
|
||||
meter: Mapped["Meter | None"] = relationship(back_populates="cost_periods")
|
||||
|
||||
source_binding: Mapped["MeterSourceBinding | None"] = relationship(
|
||||
back_populates="cost_periods"
|
||||
)
|
||||
|
||||
|
||||
class MeterCostPeriod(Base):
|
||||
"""Auditable commodity-scoped ledger row for one half-open metering period.
|
||||
|
||||
A degraded row intentionally permits missing audit links; services must
|
||||
still require them before writing a normal row.
|
||||
"""
|
||||
|
||||
__tablename__ = "meter_cost_period"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
commodity: Mapped[str] = mapped_column(String(32), nullable=False)
|
||||
period_start: Mapped[datetime] = mapped_column(UTCDateTime("period_start"), nullable=False)
|
||||
period_end: Mapped[datetime] = mapped_column(UTCDateTime("period_end"), nullable=False)
|
||||
meter_id: Mapped[int | None] = mapped_column(
|
||||
ForeignKey("meter.id", ondelete="RESTRICT"), nullable=True
|
||||
)
|
||||
source_binding_id: Mapped[int | None] = mapped_column(
|
||||
ForeignKey("meter_source_binding.id", ondelete="RESTRICT"), nullable=True
|
||||
)
|
||||
contract_version_id: Mapped[int | None] = mapped_column(
|
||||
ForeignKey("energy_contract_version.id", ondelete="RESTRICT"), nullable=True
|
||||
)
|
||||
|
||||
# SQLite reliably round-trips at most fifteen significant decimal digits.
|
||||
# WarmteLink itself reports 0.001 units, so nine cost fractional digits
|
||||
# retain a six-place tariff times that source precision without float loss.
|
||||
quantity: Mapped[Decimal] = mapped_column(ExactDecimal(15, 6), nullable=False)
|
||||
cost: Mapped[Decimal] = mapped_column(ExactDecimal(15, 9), nullable=False)
|
||||
currency: Mapped[str] = mapped_column(String(8), nullable=False)
|
||||
cost_breakdown: Mapped[dict] = mapped_column(DecimalJSON(), nullable=False, default=dict)
|
||||
pricing_snapshot: Mapped[dict] = mapped_column(DecimalJSON(), nullable=False, default=dict)
|
||||
quality: Mapped[str] = mapped_column(String(32), nullable=False, default="valid")
|
||||
degraded: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
|
||||
degraded_reason: Mapped[str | None] = mapped_column(String(255), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
|
||||
meter: Mapped["Meter | None"] = relationship()
|
||||
source_binding: Mapped["MeterSourceBinding | None"] = relationship()
|
||||
contract_version: Mapped["EnergyContractVersion | None"] = relationship()
|
||||
|
||||
__table_args__ = (
|
||||
CheckConstraint(
|
||||
"degraded OR (meter_id IS NOT NULL AND source_binding_id IS NOT NULL "
|
||||
"AND contract_version_id IS NOT NULL)",
|
||||
name="ck_meter_cost_period_normal_audit_links",
|
||||
),
|
||||
CheckConstraint("period_end > period_start", name="ck_meter_cost_period_positive_interval"),
|
||||
CheckConstraint(
|
||||
"NOT degraded OR (degraded_reason IS NOT NULL AND length(trim(degraded_reason)) > 0)",
|
||||
name="ck_meter_cost_period_degraded_reason",
|
||||
),
|
||||
UniqueConstraint("commodity", "period_start", name="uq_meter_cost_period_commodity_start"),
|
||||
Index("ix_meter_cost_period_commodity_start", "commodity", "period_start"),
|
||||
Index("ix_meter_cost_period_source_binding_id", "source_binding_id"),
|
||||
)
|
||||
|
||||
@validates("cost_breakdown", "pricing_snapshot")
|
||||
def _validate_decimal_json(self, _key: str, value: dict) -> dict:
|
||||
return _decimal_json(value)
|
||||
|
||||
@validates("quantity", "cost")
|
||||
def _validate_fixed_decimal(self, key: str, value: Decimal) -> Decimal:
|
||||
precision, scale = (15, 6) if key == "quantity" else (15, 9)
|
||||
return _validate_fixed_decimal(value, precision, scale, key)
|
||||
|
||||
@validates("period_start", "period_end")
|
||||
def _normalise_period(self, key: str, value: datetime) -> datetime:
|
||||
return _normalise_utc_period(value, key)
|
||||
|
||||
|
||||
@event.listens_for(MeterCostPeriod, "before_insert")
|
||||
@event.listens_for(MeterCostPeriod, "before_update")
|
||||
def _validate_meter_cost_period(_mapper, _connection, target: MeterCostPeriod) -> None:
|
||||
_require_aware_period(target.period_start, target.period_end)
|
||||
if not target.degraded and (
|
||||
target.meter_id is None
|
||||
or target.source_binding_id is None
|
||||
or target.contract_version_id is None
|
||||
):
|
||||
raise ValueError("normal meter cost periods require meter, binding, and contract version")
|
||||
if target.degraded and not target.degraded_reason:
|
||||
raise ValueError("degraded meter cost periods require a degraded_reason")
|
||||
|
||||
|
||||
# Index on recorded_at for efficient time-range queries on DSMR readings.
|
||||
# (The ORM-level index=True on recorded_at already creates ix_dsmr_reading_recorded_at;
|
||||
# no composite index is needed for single-meter deployments.)
|
||||
|
||||
# Index on period_start is covered by the unique constraint (SQLite creates an
|
||||
# implicit index for UNIQUE columns), so no additional index is required.
|
||||
|
||||
# Index on starts_at for TibberPrice is covered by the unique constraint similarly.
|
||||
@@ -0,0 +1,49 @@
|
||||
"""SQLAlchemy model for the exposed-entity toggle table.
|
||||
|
||||
``ExposedEntityToggle`` stores the per-entity enable/disable state for the
|
||||
MQTT / HA Discovery expose framework. The *catalog* of what *can* be
|
||||
exposed is computed dynamically by provider functions (see
|
||||
``app/integrations/expose.py``); this table only records which entries the
|
||||
user has explicitly enabled.
|
||||
|
||||
Key design decisions
|
||||
--------------------
|
||||
- ``key`` is a stable string identifier derived from device uuid + metric key
|
||||
(e.g. ``"modbus.<uuid>.voltage"``), so it does not drift if rows are
|
||||
deleted and re-inserted.
|
||||
- Default state for an entity not yet in this table is **disabled** (false).
|
||||
``build_catalog`` in expose.py treats a missing row as ``enabled=False``.
|
||||
- Only one row per entity key (``key`` is UNIQUE).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, DateTime, Integer, String
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.db import Base
|
||||
|
||||
|
||||
class ExposedEntityToggle(Base):
|
||||
"""Per-entity on/off switch for MQTT / HA Discovery publishing.
|
||||
|
||||
Rows are created on demand (first toggle); entities with no row are
|
||||
treated as disabled.
|
||||
"""
|
||||
|
||||
__tablename__ = "exposed_entity_toggle"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
|
||||
# Stable entity key, e.g. "modbus.<uuid>.voltage" or "modbus.<uuid>.online".
|
||||
key: Mapped[str] = mapped_column(String(255), unique=True, nullable=False)
|
||||
|
||||
# Whether this entity should be published via MQTT / HA Discovery.
|
||||
enabled: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
|
||||
|
||||
# Last time this row was created or modified.
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False
|
||||
)
|
||||
@@ -0,0 +1,165 @@
|
||||
"""Protocol-agnostic source, channel, and meter-binding identity models."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid as _uuid
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import (
|
||||
Boolean,
|
||||
CheckConstraint,
|
||||
DateTime,
|
||||
ForeignKey,
|
||||
Index,
|
||||
Integer,
|
||||
Numeric,
|
||||
String,
|
||||
UniqueConstraint,
|
||||
)
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
from sqlalchemy.types import JSON
|
||||
|
||||
from app.db import Base
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from app.models.energy import EnergyCostPeriod, Meter
|
||||
|
||||
|
||||
def _uuid4_str() -> str:
|
||||
return str(_uuid.uuid4())
|
||||
|
||||
|
||||
def half_open_intervals_overlap(
|
||||
started_at: datetime,
|
||||
ended_at: datetime | None,
|
||||
other_started_at: datetime,
|
||||
other_ended_at: datetime | None,
|
||||
) -> bool:
|
||||
"""Return whether two ``[started_at, ended_at)`` intervals overlap.
|
||||
|
||||
``None`` denotes an open-ended interval. Equal boundaries do not overlap,
|
||||
which lets a source binding hand off at one exact timestamp.
|
||||
"""
|
||||
return (other_ended_at is None or started_at < other_ended_at) and (
|
||||
ended_at is None or other_started_at < ended_at
|
||||
)
|
||||
|
||||
|
||||
class MeterSource(Base):
|
||||
"""A configured protocol connection that discovers one or more channels."""
|
||||
|
||||
__tablename__ = "meter_source"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
uuid: Mapped[str] = mapped_column(String(36), unique=True, nullable=False, default=_uuid4_str)
|
||||
name: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
kind: Mapped[str] = mapped_column(String(64), nullable=False)
|
||||
enabled: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True)
|
||||
config: Mapped[dict] = mapped_column(JSON, nullable=False, default=dict)
|
||||
status: Mapped[str] = mapped_column(String(32), nullable=False, default="unknown")
|
||||
last_seen_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
last_error: Mapped[str | None] = mapped_column(String(1024), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
|
||||
channels: Mapped[list["MeterSourceChannel"]] = relationship(
|
||||
back_populates="source", cascade="save-update, merge"
|
||||
)
|
||||
|
||||
|
||||
class MeterSourceChannel(Base):
|
||||
"""A stable cumulative measurement identity discovered from a source."""
|
||||
|
||||
__tablename__ = "meter_source_channel"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
uuid: Mapped[str] = mapped_column(String(36), unique=True, nullable=False, default=_uuid4_str)
|
||||
source_id: Mapped[int] = mapped_column(
|
||||
ForeignKey("meter_source.id", ondelete="RESTRICT"), nullable=False, index=True
|
||||
)
|
||||
channel_key: Mapped[str] = mapped_column(String(128), nullable=False)
|
||||
label: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
suggested_commodity: Mapped[str | None] = mapped_column(String(32), nullable=True)
|
||||
unit: Mapped[str] = mapped_column(String(32), nullable=False)
|
||||
device_type: Mapped[str | None] = mapped_column(String(64), nullable=True)
|
||||
fingerprint: Mapped[str | None] = mapped_column(String(64), nullable=True)
|
||||
latest_value: Mapped[float | None] = mapped_column(Numeric(20, 6), nullable=True)
|
||||
latest_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
latest_quality: Mapped[str | None] = mapped_column(String(32), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
|
||||
source: Mapped["MeterSource"] = relationship(back_populates="channels")
|
||||
bindings: Mapped[list["MeterSourceBinding"]] = relationship(
|
||||
back_populates="channel", cascade="save-update, merge"
|
||||
)
|
||||
warmtelink_readings: Mapped[list["WarmteLinkReading"]] = relationship(
|
||||
back_populates="channel", cascade="save-update, merge"
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
UniqueConstraint("source_id", "channel_key", name="uq_meter_source_channel_source_key"),
|
||||
)
|
||||
|
||||
|
||||
class WarmteLinkReading(Base):
|
||||
"""One accepted scalar cumulative reading from a WarmteLink channel."""
|
||||
|
||||
__tablename__ = "warmtelink_reading"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
channel_id: Mapped[int] = mapped_column(
|
||||
ForeignKey("meter_source_channel.id", ondelete="RESTRICT"), nullable=False
|
||||
)
|
||||
# These timestamps retain their UTC-aware application semantics. SQLite
|
||||
# stores them without an offset, so callers must always supply aware UTC.
|
||||
recorded_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
received_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
# SQLite's ORM Numeric path can exactly round-trip this 12-integer-digit
|
||||
# range at scale 3. That is ample for a long-lived cumulative meter while
|
||||
# retaining the protocol's 0.001 resolution without float conversion.
|
||||
value: Mapped[Decimal] = mapped_column(Numeric(15, 3), nullable=False)
|
||||
unit: Mapped[str] = mapped_column(String(32), nullable=False)
|
||||
quality: Mapped[str] = mapped_column(String(32), nullable=False)
|
||||
equipment_fingerprint: Mapped[str] = mapped_column(String(64), nullable=False)
|
||||
|
||||
channel: Mapped["MeterSourceChannel"] = relationship(back_populates="warmtelink_readings")
|
||||
|
||||
__table_args__ = (
|
||||
CheckConstraint(
|
||||
"quality IN ('valid', 'invalid', 'unverifiable')",
|
||||
name="ck_warmtelink_reading_quality",
|
||||
),
|
||||
UniqueConstraint("channel_id", "recorded_at", name="uq_warmtelink_reading_channel_recorded_at"),
|
||||
Index("ix_warmtelink_reading_recorded_at", "recorded_at"),
|
||||
)
|
||||
|
||||
|
||||
class MeterSourceBinding(Base):
|
||||
"""Connect one source channel to one physical meter for a half-open window."""
|
||||
|
||||
__tablename__ = "meter_source_binding"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
uuid: Mapped[str] = mapped_column(String(36), unique=True, nullable=False, default=_uuid4_str)
|
||||
meter_id: Mapped[int] = mapped_column(
|
||||
ForeignKey("meter.id", ondelete="RESTRICT"), nullable=False, index=True
|
||||
)
|
||||
channel_id: Mapped[int] = mapped_column(
|
||||
ForeignKey("meter_source_channel.id", ondelete="RESTRICT"), nullable=False, index=True
|
||||
)
|
||||
started_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
ended_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
|
||||
meter: Mapped["Meter"] = relationship(back_populates="source_bindings")
|
||||
channel: Mapped["MeterSourceChannel"] = relationship(back_populates="bindings")
|
||||
cost_periods: Mapped[list["EnergyCostPeriod"]] = relationship(
|
||||
back_populates="source_binding", cascade="save-update, merge", passive_deletes="all"
|
||||
)
|
||||
|
||||
|
||||
Index("ix_meter_source_kind_enabled", MeterSource.kind, MeterSource.enabled)
|
||||
@@ -0,0 +1,114 @@
|
||||
"""SQLAlchemy models for Modbus device management and telemetry.
|
||||
|
||||
Two tables:
|
||||
- modbus_device: deployment/configurable metadata for each polled device.
|
||||
- modbus_reading: generic telemetry rows (one per device per poll cycle).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid as _uuid
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, DateTime, ForeignKey, Index, Integer, String
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
from sqlalchemy.types import JSON
|
||||
|
||||
from app.db import Base
|
||||
|
||||
|
||||
def _uuid4_str() -> str:
|
||||
return str(_uuid.uuid4())
|
||||
|
||||
|
||||
class ModbusDevice(Base):
|
||||
"""Deployment-layer record for a Modbus slave device reachable via a TCP gateway.
|
||||
|
||||
Protocol knowledge (register map, decoding rules) lives in the YAML profile
|
||||
referenced by ``profile``. Per-device variables (network address, slave ID,
|
||||
display name, poll rate) live here.
|
||||
"""
|
||||
|
||||
__tablename__ = "modbus_device"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
|
||||
# Stable internal identity — used as the API path key and HA Discovery unique_id anchor.
|
||||
uuid: Mapped[str] = mapped_column(
|
||||
String(36), unique=True, nullable=False, default=_uuid4_str
|
||||
)
|
||||
|
||||
# Human-readable label (may be changed; changing it triggers re-publish of HA Discovery).
|
||||
friendly_name: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
|
||||
# Transport protocol — only "tcp" is supported for now.
|
||||
transport: Mapped[str] = mapped_column(String(16), nullable=False, default="tcp")
|
||||
|
||||
# TCP gateway address.
|
||||
host: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
port: Mapped[int] = mapped_column(Integer, nullable=False, default=502)
|
||||
|
||||
# Modbus slave address (set on the device panel; different meters on the same gateway
|
||||
# must have distinct unit_ids).
|
||||
unit_id: Mapped[int] = mapped_column(Integer, nullable=False, default=1)
|
||||
|
||||
# Which YAML profile to use for register mapping and decoding (e.g. "sdm120").
|
||||
profile: Mapped[str] = mapped_column(String(64), nullable=False)
|
||||
|
||||
# Polling interval in seconds.
|
||||
poll_interval_s: Mapped[int] = mapped_column(Integer, nullable=False, default=5)
|
||||
|
||||
# Whether this device is included in the periodic poll sweep.
|
||||
enabled: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True)
|
||||
|
||||
# Poll-status fields (updated by the poll service after each attempt).
|
||||
last_poll_at: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True), nullable=True
|
||||
)
|
||||
last_poll_ok: Mapped[bool | None] = mapped_column(Boolean, nullable=True)
|
||||
|
||||
# Audit timestamps.
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False
|
||||
)
|
||||
|
||||
# Relationship to readings (back-reference; not loaded eagerly).
|
||||
readings: Mapped[list["ModbusReading"]] = relationship(
|
||||
back_populates="device", cascade="save-update, merge"
|
||||
)
|
||||
|
||||
|
||||
class ModbusReading(Base):
|
||||
"""Generic telemetry row: one device, one poll instant, all decoded metrics as JSON.
|
||||
|
||||
``payload`` contains the full dict of engineering values produced by the YAML profile
|
||||
decoder, e.g. ``{"voltage": 230.2, "current": 1.3, ...}``. The profile is the key
|
||||
to interpreting which keys exist and what units they carry.
|
||||
"""
|
||||
|
||||
__tablename__ = "modbus_reading"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
|
||||
# FK to the device that produced this reading.
|
||||
# ON DELETE RESTRICT: prevents accidental deletion of a device that has historical data.
|
||||
device_id: Mapped[int] = mapped_column(
|
||||
ForeignKey("modbus_device.id", ondelete="RESTRICT"), nullable=False
|
||||
)
|
||||
|
||||
# The UTC timestamp of when the sample was taken — real indexed column, NOT embedded in payload.
|
||||
recorded_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, index=True
|
||||
)
|
||||
|
||||
# Profile-decoded engineering values as a JSON object.
|
||||
payload: Mapped[dict] = mapped_column(JSON, nullable=False)
|
||||
|
||||
device: Mapped["ModbusDevice"] = relationship(back_populates="readings")
|
||||
|
||||
|
||||
# Composite index for efficient time-range queries scoped to a single device.
|
||||
Index("ix_modbus_reading_device_recorded", ModbusReading.device_id, ModbusReading.recorded_at)
|
||||
@@ -0,0 +1,47 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class ConfigField(BaseModel):
|
||||
env_name: str
|
||||
label: str
|
||||
value: str
|
||||
secret: bool
|
||||
input_type: str
|
||||
configured: bool
|
||||
|
||||
|
||||
class ConfigSection(BaseModel):
|
||||
name: str
|
||||
fields: list[ConfigField]
|
||||
|
||||
|
||||
class ConfigResponse(BaseModel):
|
||||
sections: list[ConfigSection]
|
||||
|
||||
|
||||
class ConfigUpdateRequest(BaseModel):
|
||||
"""Flat mapping of env_name → value, mirroring the existing form semantics."""
|
||||
|
||||
updates: dict[str, str]
|
||||
|
||||
|
||||
class ConfigUpdateResponse(BaseModel):
|
||||
sections: list[ConfigSection]
|
||||
|
||||
|
||||
class SmtpTestResponse(BaseModel):
|
||||
"""Response from POST /api/config/smtp/test."""
|
||||
|
||||
result: Literal["success", "config-error", "failed"]
|
||||
message: str
|
||||
|
||||
|
||||
class MqttTestResponse(BaseModel):
|
||||
"""Response from POST /api/config/mqtt/test."""
|
||||
|
||||
result: Literal["success", "config-error", "failed"]
|
||||
message: str
|
||||
@@ -0,0 +1,92 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Location
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class LocationRecord(BaseModel):
|
||||
person: str
|
||||
datetime: str
|
||||
latitude: float
|
||||
longitude: float
|
||||
altitude: float | None
|
||||
|
||||
|
||||
class LocationsResponse(BaseModel):
|
||||
items: list[LocationRecord]
|
||||
limit: int
|
||||
offset: int
|
||||
|
||||
|
||||
class LocationUpdateRequest(BaseModel):
|
||||
"""PATCH body for a location record — all fields optional; PK fields excluded."""
|
||||
|
||||
latitude: float | None = None
|
||||
longitude: float | None = None
|
||||
altitude: float | None = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Poo
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class PooRecord(BaseModel):
|
||||
timestamp: str
|
||||
status: str
|
||||
latitude: float
|
||||
longitude: float
|
||||
|
||||
|
||||
class PooResponse(BaseModel):
|
||||
items: list[PooRecord]
|
||||
limit: int
|
||||
offset: int
|
||||
|
||||
|
||||
class PooUpdateRequest(BaseModel):
|
||||
"""PATCH body for a poo record — all fields optional; PK field excluded."""
|
||||
|
||||
status: str | None = None
|
||||
latitude: float | None = None
|
||||
longitude: float | None = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public IP
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class PublicIPStateSchema(BaseModel):
|
||||
id: int
|
||||
current_ipv4: str
|
||||
previous_ipv4: str | None
|
||||
first_seen_at: datetime
|
||||
last_checked_at: datetime
|
||||
last_changed_at: datetime | None
|
||||
last_check_status: str
|
||||
last_check_error: str | None
|
||||
last_provider: str | None
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class PublicIPHistorySchema(BaseModel):
|
||||
id: int
|
||||
ipv4: str
|
||||
observed_at: datetime
|
||||
change_type: str
|
||||
provider: str | None
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class PublicIPResponse(BaseModel):
|
||||
state: PublicIPStateSchema | None
|
||||
history: list[PublicIPHistorySchema]
|
||||
@@ -0,0 +1,253 @@
|
||||
"""Pydantic schemas for the Energy data API (M6-T09).
|
||||
|
||||
Covers six endpoint groups under /api/energy:
|
||||
- GET /prices — price curve (tibber 15min points or manual tariff)
|
||||
- GET /costs — energy_cost_period rows (time-range, paginated)
|
||||
- GET /costs/summary — aggregated metered + standing charges − credits
|
||||
- GET /dsmr/latest — most recent dsmr_reading blob
|
||||
- POST /costs/recompute — explicit idempotent recompute (returns count)
|
||||
- POST /tibber/test — three-state Tibber connection test
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /api/energy/prices
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class PricePointSchema(BaseModel):
|
||||
"""A single 15-minute price point (tibber) or placeholder entry."""
|
||||
|
||||
starts_at: datetime
|
||||
buy: float = Field(description="All-in buy price in EUR/kWh (including taxes).")
|
||||
sell: float = Field(description="Net sell price in EUR/kWh.")
|
||||
level: str | None = Field(
|
||||
default=None,
|
||||
description="Tibber price level (CHEAP / NORMAL / EXPENSIVE); null for manual.",
|
||||
)
|
||||
|
||||
|
||||
class ManualTariffSchema(BaseModel):
|
||||
"""Fixed-tariff breakdown for manual contracts.
|
||||
|
||||
Prices are the *effective* buy prices as used by the billing engine
|
||||
(energy_buy_x + energy_tax + ode) and the raw sell prices.
|
||||
"""
|
||||
|
||||
buy_dal: float = Field(description="Effective buy price, low-tariff / dal (EUR/kWh).")
|
||||
buy_normal: float = Field(description="Effective buy price, normal / high-tariff (EUR/kWh).")
|
||||
sell_dal: float = Field(description="Sell price, low-tariff / dal (EUR/kWh).")
|
||||
sell_normal: float = Field(description="Sell price, normal / high-tariff (EUR/kWh).")
|
||||
|
||||
|
||||
class PricesResponse(BaseModel):
|
||||
"""Response for GET /api/energy/prices.
|
||||
|
||||
``kind`` mirrors the active contract kind:
|
||||
- ``"tibber"`` → ``points`` has actual 15-min price entries; ``tariff`` is null.
|
||||
- ``"manual"`` → ``points`` is empty; ``tariff`` carries the fixed-rate table.
|
||||
- ``None`` → no active contract; both ``points`` and ``tariff`` are empty/null.
|
||||
|
||||
``currency`` comes from the active contract (or "EUR" fallback).
|
||||
``points`` is always ascending by ``starts_at``.
|
||||
"""
|
||||
|
||||
kind: str | None = Field(
|
||||
description="Active contract kind ('tibber' or 'manual'), or null if no active contract."
|
||||
)
|
||||
currency: str = Field(description="ISO 4217 currency code.")
|
||||
points: list[PricePointSchema] = Field(
|
||||
description=(
|
||||
"15-minute price points for tibber contracts (ascending by starts_at). "
|
||||
"Empty for manual contracts or when no active contract exists."
|
||||
)
|
||||
)
|
||||
tariff: ManualTariffSchema | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Fixed tariff table for manual contracts. "
|
||||
"Null for tibber contracts and when no active contract exists."
|
||||
),
|
||||
)
|
||||
contract_version_id: int | None = Field(
|
||||
default=None,
|
||||
description="Thermal active contract version identifier; omitted for electricity.",
|
||||
)
|
||||
effective_from: datetime | None = Field(
|
||||
default=None,
|
||||
description="Thermal contract version start; omitted for electricity.",
|
||||
)
|
||||
effective_to: datetime | None = Field(
|
||||
default=None,
|
||||
description="Thermal contract version end; omitted for electricity.",
|
||||
)
|
||||
values: dict[str, dict[str, str]] | None = Field(
|
||||
default=None,
|
||||
description="Thermal normalized Decimal-string contract values; omitted for electricity.",
|
||||
)
|
||||
|
||||
model_config = {"ser_json_exclude_none": True}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /api/energy/costs
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class CostPeriodSchema(BaseModel):
|
||||
"""One 15-minute billing record from the energy_cost_period table."""
|
||||
|
||||
period_start: datetime
|
||||
d1_kwh: float = Field(description="Delivered low-tariff kWh for this period.")
|
||||
d2_kwh: float = Field(description="Delivered normal-tariff kWh for this period.")
|
||||
r1_kwh: float = Field(description="Returned low-tariff kWh for this period.")
|
||||
r2_kwh: float = Field(description="Returned normal-tariff kWh for this period.")
|
||||
import_cost: float = Field(description="Cost of electricity drawn from grid (EUR).")
|
||||
export_revenue: float = Field(description="Revenue from electricity fed to grid (EUR).")
|
||||
net_cost: float = Field(description="import_cost − export_revenue (EUR).")
|
||||
currency: str = Field(description="ISO 4217 currency code.")
|
||||
degraded: bool = Field(description="True when the period was computed with incomplete data.")
|
||||
contract_version_id: int | None = Field(
|
||||
default=None,
|
||||
description="FK to the contract version used for this billing period (null when degraded).",
|
||||
)
|
||||
source_binding_id: int | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"FK to the source binding that supplied both cumulative endpoints "
|
||||
"(null for legacy or degraded periods)."
|
||||
),
|
||||
)
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class CostsResponse(BaseModel):
|
||||
"""Response for GET /api/energy/costs."""
|
||||
|
||||
items: list[CostPeriodSchema]
|
||||
total: int = Field(description="Number of items returned.")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /api/energy/costs/summary
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class SummaryResponse(BaseModel):
|
||||
"""Response for GET /api/energy/costs/summary.
|
||||
|
||||
Monetary values are in ``currency``; the ``*_kwh`` fields are energy totals
|
||||
in kWh. ``metered_import``/``metered_export`` are **money**, not energy —
|
||||
only the ``_kwh``-suffixed fields carry kWh.
|
||||
|
||||
``total_payable = metered_net + fixed_costs − credits``
|
||||
"""
|
||||
|
||||
currency: str
|
||||
metered_import: float = Field(
|
||||
description="Σ import_cost for non-degraded periods (money, in `currency`)."
|
||||
)
|
||||
metered_export: float = Field(
|
||||
description="Σ export_revenue for non-degraded periods (money, in `currency`)."
|
||||
)
|
||||
metered_net: float = Field(
|
||||
description="Σ net_cost for non-degraded periods (money, in `currency`)."
|
||||
)
|
||||
metered_import_kwh: float = Field(
|
||||
description="Σ (d1_kwh + d2_kwh) for non-degraded periods (energy imported, kWh)."
|
||||
)
|
||||
metered_export_kwh: float = Field(
|
||||
description="Σ (r1_kwh + r2_kwh) for non-degraded periods (energy exported, kWh)."
|
||||
)
|
||||
fixed_costs: float = Field(
|
||||
description="Standing charges (network_fee + management_fee) apportioned over the interval."
|
||||
)
|
||||
credits: float = Field(
|
||||
description="Energy-tax credit (heffingskorting) apportioned over the interval."
|
||||
)
|
||||
total_payable: float = Field(
|
||||
description="metered_net + fixed_costs − credits (actual amount owed)."
|
||||
)
|
||||
period_count: int = Field(description="Number of non-degraded billing periods in range.")
|
||||
degraded_count: int = Field(description="Number of degraded billing periods in range.")
|
||||
days: float = Field(description="Interval length in days.")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /api/energy/dsmr/latest
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class DsmrLatestResponse(BaseModel):
|
||||
"""Response for GET /api/energy/dsmr/latest.
|
||||
|
||||
``found`` is False when no dsmr_reading rows exist yet. The front-end
|
||||
should check ``found`` before reading ``recorded_at`` or ``payload``.
|
||||
"""
|
||||
|
||||
found: bool
|
||||
recorded_at: datetime | None = None
|
||||
payload: dict[str, Any] | None = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# POST /api/energy/costs/recompute
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class RecomputeResponse(BaseModel):
|
||||
"""Response for POST /api/energy/costs/recompute."""
|
||||
|
||||
recomputed: int = Field(
|
||||
description=(
|
||||
"Number of 15-minute periods for which a billing record was written "
|
||||
"(inserted or updated). Periods skipped due to missing contract or "
|
||||
"missing Tibber price are not counted."
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# POST /api/energy/tibber/test
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TibberTestPriceSchema(BaseModel):
|
||||
"""Current Tibber price point returned on a successful test.
|
||||
|
||||
Carries enough fields for the front-end to confirm the API is working and
|
||||
display the live price. The API token is **never** included.
|
||||
"""
|
||||
|
||||
starts_at: datetime
|
||||
total: float
|
||||
energy: float
|
||||
tax: float
|
||||
currency: str
|
||||
level: str | None = None
|
||||
|
||||
|
||||
class TibberTestResponse(BaseModel):
|
||||
"""Three-state response for POST /api/energy/tibber/test.
|
||||
|
||||
Possible ``result`` values:
|
||||
|
||||
- ``"success"`` — Tibber API responded with a valid price.
|
||||
- ``"config-error"`` — Token is missing or not configured.
|
||||
- ``"failed"`` — API call failed (auth rejected, network error, timeout, etc.).
|
||||
|
||||
``price`` is populated only on ``"success"``; it is null otherwise.
|
||||
``message`` always contains a human-readable explanation.
|
||||
"""
|
||||
|
||||
result: Literal["success", "config-error", "failed"]
|
||||
message: str
|
||||
price: TibberTestPriceSchema | None = None
|
||||
@@ -0,0 +1,151 @@
|
||||
"""Pydantic schemas for the EnergyContract CRUD + versioning API (M6-T04).
|
||||
|
||||
Schema hierarchy
|
||||
----------------
|
||||
ContractVersionResponse — single version row (id, dates, values, created_at)
|
||||
ContractResponse — contract head (id, name, kind, active, currency, timestamps)
|
||||
ContractDetailResponse — contract head + embedded versions list (for GET{id})
|
||||
ContractListResponse — paginated list of ContractResponse items
|
||||
ContractCreate — POST /api/energy/contracts body
|
||||
ContractPatch — PATCH /api/energy/contracts/{id} body (all fields optional)
|
||||
VersionCreate — POST /api/energy/contracts/{id}/versions body
|
||||
ProfilesResponse — GET /api/energy/profiles response
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Version schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ContractVersionResponse(BaseModel):
|
||||
"""Response schema for a single EnergyContractVersion row."""
|
||||
|
||||
id: int
|
||||
effective_from: datetime
|
||||
effective_to: datetime | None
|
||||
values: dict[str, Any]
|
||||
created_at: datetime
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Contract schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ContractResponse(BaseModel):
|
||||
"""Response schema for a single EnergyContract (without embedded versions).
|
||||
|
||||
Used for list responses where embedding all versions would be expensive.
|
||||
"""
|
||||
|
||||
id: int
|
||||
name: str
|
||||
kind: str
|
||||
scope: str
|
||||
active: bool
|
||||
currency: str
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class ContractDetailResponse(BaseModel):
|
||||
"""Response schema for a single EnergyContract with full version history.
|
||||
|
||||
Returned by GET /api/energy/contracts/{id} and by successful POST / PATCH
|
||||
operations where the caller needs to see all version data.
|
||||
"""
|
||||
|
||||
id: int
|
||||
name: str
|
||||
kind: str
|
||||
scope: str
|
||||
active: bool
|
||||
currency: str
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
versions: list[ContractVersionResponse]
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class ContractListResponse(BaseModel):
|
||||
"""Response schema for GET /api/energy/contracts."""
|
||||
|
||||
items: list[ContractResponse]
|
||||
total: int
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Request schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ContractCreate(BaseModel):
|
||||
"""Request body for POST /api/energy/contracts.
|
||||
|
||||
``effective_from`` defaults to the current UTC time if not provided,
|
||||
giving the first version an open-ended start from "now".
|
||||
``kind`` is validated at the application layer against the profile registry;
|
||||
clients should send ``"manual"`` or ``"tibber"``.
|
||||
"""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=255)
|
||||
kind: str = Field(..., min_length=1, max_length=32)
|
||||
scope: str | None = Field(default=None, min_length=1, max_length=32)
|
||||
currency: str = Field(default="EUR", min_length=1, max_length=8)
|
||||
values: dict[str, Any]
|
||||
effective_from: datetime | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"UTC datetime from which the first pricing version is effective. "
|
||||
"Defaults to the current UTC time when omitted."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class ContractPatch(BaseModel):
|
||||
"""Request body for PATCH /api/energy/contracts/{id}.
|
||||
|
||||
All fields are optional. Sending ``active=true`` activates this contract
|
||||
(deactivating all others); ``active=false`` deactivates it without affecting
|
||||
other contracts.
|
||||
"""
|
||||
|
||||
name: str | None = Field(default=None, min_length=1, max_length=255)
|
||||
active: bool | None = None
|
||||
|
||||
|
||||
class VersionCreate(BaseModel):
|
||||
"""Request body for POST /api/energy/contracts/{id}/versions."""
|
||||
|
||||
effective_from: datetime
|
||||
values: dict[str, Any]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Profile response schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ProfilesResponse(BaseModel):
|
||||
"""Response schema for GET /api/energy/profiles.
|
||||
|
||||
``profiles`` is a list of raw profile dicts as produced by
|
||||
``list_profiles()`` (Pydantic model dumps). Each entry contains at minimum
|
||||
``kind`` and ``label``; the full nested structure allows the front-end to
|
||||
render a type-appropriate form for each pricing profile.
|
||||
"""
|
||||
|
||||
profiles: list[dict[str, Any]]
|
||||
@@ -0,0 +1,95 @@
|
||||
"""Pydantic schemas for the Expose API (M5-T12).
|
||||
|
||||
Three endpoints:
|
||||
GET /api/expose — catalog + toggle state + MQTT/Discovery status
|
||||
PUT /api/expose — set toggles (map key → bool)
|
||||
POST /api/expose/republish — trigger discovery re-publish
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Nested schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class DeviceInfoSchema(BaseModel):
|
||||
"""HA device grouping info for an exposable entity."""
|
||||
|
||||
identifiers: list[str]
|
||||
name: str
|
||||
|
||||
|
||||
class ExposableEntitySchema(BaseModel):
|
||||
"""One exposable entity in the catalog.
|
||||
|
||||
``value_getter`` is intentionally excluded — it is a non-serialisable
|
||||
callable and is only used internally by the HA Discovery service.
|
||||
"""
|
||||
|
||||
key: str
|
||||
component: str
|
||||
device: DeviceInfoSchema
|
||||
device_class: str | None
|
||||
unit: str
|
||||
name: str
|
||||
state_class: str | None = None
|
||||
|
||||
|
||||
class CatalogEntrySchema(BaseModel):
|
||||
"""An entity from the catalog with its current toggle state."""
|
||||
|
||||
entity: ExposableEntitySchema
|
||||
enabled: bool
|
||||
|
||||
|
||||
class MqttStatusSchema(BaseModel):
|
||||
"""Connection status for MQTT and HA Discovery."""
|
||||
|
||||
mqtt_configured: bool
|
||||
mqtt_connected: bool
|
||||
discovery_enabled: bool
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Response schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ExposeResponse(BaseModel):
|
||||
"""Response for GET /api/expose."""
|
||||
|
||||
catalog: list[CatalogEntrySchema]
|
||||
mqtt_status: MqttStatusSchema
|
||||
|
||||
|
||||
class ExposeUpdateResponse(BaseModel):
|
||||
"""Response for PUT /api/expose (returns updated catalog + status)."""
|
||||
|
||||
catalog: list[CatalogEntrySchema]
|
||||
mqtt_status: MqttStatusSchema
|
||||
|
||||
|
||||
class RepublishResponse(BaseModel):
|
||||
"""Response for POST /api/expose/republish."""
|
||||
|
||||
ok: bool
|
||||
message: str
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Request schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ExposeUpdateRequest(BaseModel):
|
||||
"""Request body for PUT /api/expose.
|
||||
|
||||
``toggles`` is a map from entity key to desired enabled state (bool).
|
||||
Only keys present in the map are updated; absent keys are untouched.
|
||||
"""
|
||||
|
||||
toggles: dict[str, bool]
|
||||
@@ -0,0 +1,154 @@
|
||||
"""Pydantic schemas for the Meter CRUD + swap declaration API (M7-T05).
|
||||
|
||||
Schema hierarchy
|
||||
----------------
|
||||
MeterResponse — single meter row (id/label/commodity/started_at/ended_at/reason/note/created_at)
|
||||
MeterListResponse — ordered list of MeterResponse items
|
||||
MeterDeclareRequest — POST /api/energy/meters body (declare a swap or initial meter)
|
||||
MeterPatchRequest — PATCH /api/energy/meters/{id} body (all fields optional)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from enum import Enum
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Enums
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MeterReason(str, Enum):
|
||||
"""Allowed values for the meter epoch creation reason."""
|
||||
|
||||
initial = "initial"
|
||||
meter_swap = "meter_swap"
|
||||
home_move = "home_move"
|
||||
other = "other"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Response schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MeterResponse(BaseModel):
|
||||
"""Response schema for a single Meter epoch row.
|
||||
|
||||
``ended_at`` is ``null`` for the currently active meter.
|
||||
"""
|
||||
|
||||
id: int
|
||||
label: str
|
||||
commodity: str
|
||||
started_at: datetime
|
||||
ended_at: datetime | None
|
||||
reason: str
|
||||
note: str | None
|
||||
created_at: datetime
|
||||
bindings: list["MeterBindingSummary"] = Field(default_factory=list)
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class MeterBindingSummary(BaseModel):
|
||||
"""Stable, non-sensitive binding identity embedded in meter responses."""
|
||||
|
||||
uuid: str
|
||||
source_channel_uuid: str
|
||||
source_uuid: str
|
||||
started_at: datetime
|
||||
ended_at: datetime | None
|
||||
|
||||
|
||||
class MeterListResponse(BaseModel):
|
||||
"""Response schema for GET /api/energy/meters.
|
||||
|
||||
Meters are returned in ascending ``started_at`` order so the caller sees
|
||||
the historical installation sequence.
|
||||
"""
|
||||
|
||||
items: list[MeterResponse]
|
||||
total: int
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Request schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_VALID_REASONS = ", ".join(r.value for r in MeterReason)
|
||||
|
||||
|
||||
class MeterDeclareRequest(BaseModel):
|
||||
"""Request body for POST /api/energy/meters.
|
||||
|
||||
Declares a new meter epoch (swap, home move, or initial declaration). The
|
||||
service layer closes the current active meter for the given commodity at
|
||||
``started_at`` and opens a new one.
|
||||
|
||||
``started_at`` follows the Principle-A localisation convention: a
|
||||
timezone-naive value is interpreted as the **server's local wall-clock time**
|
||||
(e.g. CEST midnight → stored as UTC the night before); a timezone-aware
|
||||
value is converted to UTC as-is. Omitting ``started_at`` is not allowed —
|
||||
every meter declaration must carry an explicit start timestamp.
|
||||
|
||||
``commodity`` defaults to ``"electricity"``; the field is available for
|
||||
future use with ``gas`` or ``heating``.
|
||||
"""
|
||||
|
||||
label: str = Field(..., min_length=1, max_length=255)
|
||||
started_at: datetime = Field(
|
||||
...,
|
||||
description=(
|
||||
"UTC (or server-local naive) datetime from which this meter epoch starts. "
|
||||
"May be in the past (retroactive declaration)."
|
||||
),
|
||||
)
|
||||
reason: MeterReason = Field(
|
||||
...,
|
||||
description=f"Why this epoch was created. One of: {_VALID_REASONS}.",
|
||||
)
|
||||
note: str | None = Field(default=None, max_length=1024)
|
||||
commodity: str = Field(
|
||||
default="electricity",
|
||||
min_length=1,
|
||||
max_length=32,
|
||||
description="Energy commodity this meter measures. Defaults to 'electricity'.",
|
||||
)
|
||||
source_channel_uuid: str | None = Field(
|
||||
default=None,
|
||||
min_length=1,
|
||||
max_length=36,
|
||||
description="Optional compatible source channel to bind atomically to this meter.",
|
||||
)
|
||||
|
||||
|
||||
class MeterPatchRequest(BaseModel):
|
||||
"""Request body for PATCH /api/energy/meters/{id}.
|
||||
|
||||
All fields are optional. Only non-``None`` values are applied.
|
||||
|
||||
Updating ``started_at`` is a **retroactive correction**: the service layer
|
||||
maintains timeline continuity (adjusting the preceding meter's ``ended_at``)
|
||||
and the API layer triggers ``recompute_range`` over the affected window so
|
||||
that billing attribution is re-judged.
|
||||
"""
|
||||
|
||||
label: str | None = Field(default=None, min_length=1, max_length=255)
|
||||
note: str | None = Field(default=None, max_length=1024)
|
||||
started_at: datetime | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Retroactive correction of the meter epoch start timestamp. "
|
||||
"Triggers billing recompute over the affected window."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class MeterCloseRequest(BaseModel):
|
||||
"""Close the active meter epoch at an exclusive end boundary."""
|
||||
|
||||
ended_at: datetime
|
||||
@@ -0,0 +1,62 @@
|
||||
"""Schemas for the commodity-scoped thermal cost ledger."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class MeterCostPeriodSchema(BaseModel):
|
||||
"""One auditable thermal ledger row; all Decimal values are JSON strings."""
|
||||
|
||||
commodity: Literal["heating", "hot_water"]
|
||||
period_start: datetime
|
||||
period_end: datetime
|
||||
meter_id: int | None
|
||||
source_binding_id: int | None
|
||||
contract_version_id: int | None
|
||||
quantity: str
|
||||
cost: str
|
||||
currency: str
|
||||
cost_breakdown: dict[str, str]
|
||||
pricing_snapshot: dict[str, dict[str, str]]
|
||||
quality: str
|
||||
degraded: bool
|
||||
degraded_reason: str | None
|
||||
|
||||
|
||||
class MeterCostsResponse(BaseModel):
|
||||
items: list[MeterCostPeriodSchema]
|
||||
total: int = Field(description="Total matching rows before pagination.")
|
||||
|
||||
|
||||
class ThermalFixedBreakdown(BaseModel):
|
||||
"""D11 annual-standing charges accrued per settled local day, as Decimal strings."""
|
||||
|
||||
heating_network: str
|
||||
metering: str
|
||||
delivery_set: str
|
||||
hot_water_network: str
|
||||
other: str
|
||||
|
||||
|
||||
class ThermalCostSummaryResponse(BaseModel):
|
||||
currency: str
|
||||
heating: str
|
||||
hot_water_heating: str
|
||||
hot_water: str
|
||||
hot_water_tax: str
|
||||
variable_subtotal: str
|
||||
fixed_breakdown: ThermalFixedBreakdown
|
||||
fixed_subtotal: str
|
||||
all_in: str
|
||||
period_count: int
|
||||
degraded_count: int
|
||||
|
||||
|
||||
class MeterCostRecomputeResponse(BaseModel):
|
||||
processed: int
|
||||
normal: int
|
||||
degraded: int
|
||||
@@ -0,0 +1,162 @@
|
||||
"""Public schemas for protocol-agnostic meter sources and bindings."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class SourceConfigFieldResponse(BaseModel):
|
||||
name: str
|
||||
value_type: str
|
||||
default: Any = None
|
||||
required: bool
|
||||
secret: bool
|
||||
|
||||
|
||||
class SourceProfileResponse(BaseModel):
|
||||
kind: str
|
||||
fields: list[SourceConfigFieldResponse]
|
||||
defaults: dict[str, Any]
|
||||
capabilities: list[str]
|
||||
allowed_units: list[str]
|
||||
|
||||
|
||||
class SourceProfilesResponse(BaseModel):
|
||||
items: list[SourceProfileResponse]
|
||||
|
||||
|
||||
class MeterSourceCreate(BaseModel):
|
||||
name: str = Field(..., min_length=1, max_length=255)
|
||||
kind: str = Field(..., min_length=1, max_length=64)
|
||||
config: dict[str, Any] = Field(default_factory=dict)
|
||||
enabled: bool = True
|
||||
|
||||
|
||||
class MeterSourcePatch(BaseModel):
|
||||
name: str | None = Field(default=None, min_length=1, max_length=255)
|
||||
config: dict[str, Any] | None = None
|
||||
enabled: bool | None = None
|
||||
|
||||
|
||||
class MeterSourceResponse(BaseModel):
|
||||
uuid: str
|
||||
name: str
|
||||
kind: str
|
||||
enabled: bool
|
||||
config: dict[str, Any]
|
||||
status: str
|
||||
last_seen_at: datetime | None
|
||||
last_error: str | None
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class MeterSourceListResponse(BaseModel):
|
||||
items: list[MeterSourceResponse]
|
||||
total: int
|
||||
|
||||
|
||||
class DiscoverResponse(BaseModel):
|
||||
requested: bool
|
||||
supported: bool
|
||||
status: str
|
||||
request_id: int | None = None
|
||||
detail: str | None = None
|
||||
channels: list["DiscoverChannelResponse"] = Field(default_factory=list)
|
||||
|
||||
|
||||
class DiscoverChannelResponse(BaseModel):
|
||||
uuid: str
|
||||
label: str
|
||||
unit: str
|
||||
latest_value: Decimal | None
|
||||
latest_at: datetime | None
|
||||
latest_quality: str | None
|
||||
|
||||
|
||||
class ChannelBindingSummaryResponse(BaseModel):
|
||||
count: int
|
||||
meter_ids: list[int]
|
||||
|
||||
|
||||
class CommodityResponse(BaseModel):
|
||||
key: str
|
||||
unit: str
|
||||
capabilities: list[str]
|
||||
|
||||
|
||||
class CommoditiesResponse(BaseModel):
|
||||
items: list[CommodityResponse]
|
||||
|
||||
|
||||
class MeterSourceChannelResponse(BaseModel):
|
||||
uuid: str
|
||||
label: str
|
||||
suggested_commodity: str | None
|
||||
unit: str
|
||||
device_type: str | None
|
||||
latest_value: Decimal | None
|
||||
latest_at: datetime | None
|
||||
latest_quality: str | None
|
||||
binding_count: int
|
||||
bound_meter_ids: list[int]
|
||||
binding_summary: ChannelBindingSummaryResponse
|
||||
|
||||
|
||||
class MeterSourceChannelListResponse(BaseModel):
|
||||
items: list[MeterSourceChannelResponse]
|
||||
total: int
|
||||
source_status: str
|
||||
|
||||
|
||||
class ChannelReadingResponse(BaseModel):
|
||||
recorded_at: datetime
|
||||
value: Decimal | None = None
|
||||
quality: str | None = None
|
||||
|
||||
|
||||
class ChannelReadingsResponse(BaseModel):
|
||||
items: list[ChannelReadingResponse]
|
||||
total: int
|
||||
|
||||
|
||||
class BindingCreate(BaseModel):
|
||||
source_channel_uuid: str = Field(..., min_length=1, max_length=36)
|
||||
started_at: datetime
|
||||
ended_at: datetime | None = None
|
||||
|
||||
|
||||
class BindingPatch(BaseModel):
|
||||
started_at: datetime | None = None
|
||||
ended_at: datetime | None = None
|
||||
|
||||
|
||||
class BindingTransferRequest(BaseModel):
|
||||
from_binding_uuid: str = Field(..., min_length=1, max_length=36)
|
||||
to_source_channel_uuid: str = Field(..., min_length=1, max_length=36)
|
||||
effective_at: datetime
|
||||
|
||||
|
||||
class BindingResponse(BaseModel):
|
||||
uuid: str
|
||||
meter_id: int
|
||||
source_channel_uuid: str
|
||||
source_uuid: str
|
||||
started_at: datetime
|
||||
ended_at: datetime | None
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class BindingListResponse(BaseModel):
|
||||
items: list[BindingResponse]
|
||||
total: int
|
||||
|
||||
|
||||
class BindingTransferResponse(BaseModel):
|
||||
closed_binding: BindingResponse
|
||||
created_binding: BindingResponse
|
||||
@@ -0,0 +1,171 @@
|
||||
"""Pydantic schemas for the Modbus device CRUD + readings + metrics API (M5-T05)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Device schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ModbusDeviceCreate(BaseModel):
|
||||
"""Request body for POST /api/modbus/devices."""
|
||||
|
||||
friendly_name: str = Field(..., min_length=1, max_length=255)
|
||||
transport: str = Field(default="tcp", max_length=16)
|
||||
host: str = Field(..., min_length=1, max_length=255)
|
||||
port: int = Field(default=502, ge=1, le=65535)
|
||||
unit_id: int = Field(default=1, ge=0, le=247)
|
||||
profile: str = Field(..., min_length=1, max_length=64)
|
||||
poll_interval_s: int = Field(default=5, ge=1, le=3600)
|
||||
enabled: bool = Field(default=True)
|
||||
|
||||
|
||||
class ModbusDeviceUpdate(BaseModel):
|
||||
"""Request body for PATCH /api/modbus/devices/{uuid} — all fields optional."""
|
||||
|
||||
friendly_name: str | None = Field(default=None, min_length=1, max_length=255)
|
||||
transport: str | None = Field(default=None, max_length=16)
|
||||
host: str | None = Field(default=None, min_length=1, max_length=255)
|
||||
port: int | None = Field(default=None, ge=1, le=65535)
|
||||
unit_id: int | None = Field(default=None, ge=0, le=247)
|
||||
profile: str | None = Field(default=None, min_length=1, max_length=64)
|
||||
poll_interval_s: int | None = Field(default=None, ge=1, le=3600)
|
||||
enabled: bool | None = None
|
||||
|
||||
|
||||
class ModbusDeviceResponse(BaseModel):
|
||||
"""Response schema for a single Modbus device."""
|
||||
|
||||
uuid: str
|
||||
friendly_name: str
|
||||
transport: str
|
||||
host: str
|
||||
port: int
|
||||
unit_id: int
|
||||
profile: str
|
||||
poll_interval_s: int
|
||||
enabled: bool
|
||||
last_poll_at: datetime | None
|
||||
last_poll_ok: bool | None
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class ModbusDeviceListResponse(BaseModel):
|
||||
"""Response schema for listing Modbus devices."""
|
||||
|
||||
items: list[ModbusDeviceResponse]
|
||||
total: int
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Reading schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ModbusReadingResponse(BaseModel):
|
||||
"""A single reading row: timestamp + decoded payload."""
|
||||
|
||||
recorded_at: datetime
|
||||
payload: dict[str, Any]
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class ModbusReadingsResponse(BaseModel):
|
||||
"""Response schema for the readings time-range endpoint."""
|
||||
|
||||
items: list[ModbusReadingResponse]
|
||||
|
||||
|
||||
class ModbusLatestResponse(BaseModel):
|
||||
"""Response for the /latest endpoint.
|
||||
|
||||
``found`` is False and ``recorded_at``/``payload`` are None when the device
|
||||
has no readings yet. Callers should check ``found`` before using the values.
|
||||
"""
|
||||
|
||||
found: bool
|
||||
recorded_at: datetime | None
|
||||
payload: dict[str, Any] | None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Metrics / profile schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MetricInfo(BaseModel):
|
||||
"""Metadata for a single measurable quantity in a device's profile."""
|
||||
|
||||
key: str
|
||||
label: str
|
||||
unit: str
|
||||
device_class: str
|
||||
|
||||
|
||||
class ModbusMetricsResponse(BaseModel):
|
||||
"""Response schema for GET /api/modbus/devices/{uuid}/metrics."""
|
||||
|
||||
profile: str
|
||||
metrics: list[MetricInfo]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Profile list schema
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ProfileSummary(BaseModel):
|
||||
"""One entry in the GET /api/modbus/profiles response."""
|
||||
|
||||
name: str
|
||||
description: str
|
||||
|
||||
|
||||
class ModbusProfilesResponse(BaseModel):
|
||||
"""Response schema for GET /api/modbus/profiles."""
|
||||
|
||||
profiles: list[ProfileSummary]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test-read schema
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ModbusTestReadResponse(BaseModel):
|
||||
"""Response for POST /api/modbus/devices/{uuid}/test.
|
||||
|
||||
On success ``ok=True`` and ``payload`` contains the decoded values.
|
||||
On failure ``ok=False`` and ``error`` describes the problem.
|
||||
"""
|
||||
|
||||
ok: bool
|
||||
payload: dict[str, Any] | None = None
|
||||
error: str | None = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cascade-delete schema
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ModbusDeleteResponse(BaseModel):
|
||||
"""Response for DELETE /api/modbus/devices/{uuid}?cascade=true.
|
||||
|
||||
Returned only when cascade deletion succeeds (HTTP 200). Non-cascade
|
||||
successful deletes continue to return HTTP 204 (no body).
|
||||
"""
|
||||
|
||||
deleted: bool
|
||||
readings_deleted: int
|
||||
toggles_deleted: int
|
||||
@@ -0,0 +1,25 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class SessionUser(BaseModel):
|
||||
username: str
|
||||
force_password_change: bool
|
||||
|
||||
|
||||
class SessionResponse(BaseModel):
|
||||
user: SessionUser
|
||||
csrf_token: str
|
||||
|
||||
|
||||
class LoginRequest(BaseModel):
|
||||
username: str
|
||||
password: str
|
||||
totp_code: str | None = None
|
||||
|
||||
|
||||
class PasswordChangeRequest(BaseModel):
|
||||
current_password: str
|
||||
new_password: str
|
||||
confirm_password: str
|
||||
@@ -0,0 +1,59 @@
|
||||
"""Pydantic schemas for TOTP setup / enable / disable / status endpoints (M4-T05)."""
|
||||
from __future__ import annotations
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Setup (POST /api/auth/totp/setup) — response
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TotpSetupResponse(BaseModel):
|
||||
"""Returned once after a setup call.
|
||||
|
||||
``secret`` and ``recovery_codes`` are **one-time plaintext values**.
|
||||
They are never returned again by any subsequent API call.
|
||||
The frontend must display and instruct the user to save them before confirming.
|
||||
"""
|
||||
|
||||
secret: str
|
||||
otpauth_uri: str
|
||||
recovery_codes: list[str]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Enable (POST /api/auth/totp/enable) — request
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TotpEnableRequest(BaseModel):
|
||||
"""The user confirms setup by providing the 6-digit TOTP code."""
|
||||
|
||||
code: str
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Disable (POST /api/auth/totp/disable) — request
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TotpDisableRequest(BaseModel):
|
||||
"""Disable TOTP by proving identity.
|
||||
|
||||
Exactly one of ``password`` or ``code`` must be provided.
|
||||
"""
|
||||
|
||||
password: str | None = None
|
||||
code: str | None = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Status (GET /api/auth/totp) — response
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TotpStatusResponse(BaseModel):
|
||||
"""Minimal status response — never exposes secret or recovery codes."""
|
||||
|
||||
enabled: bool
|
||||
@@ -25,9 +25,9 @@ class ConfigField:
|
||||
CONFIG_FIELDS: tuple[ConfigField, ...] = (
|
||||
ConfigField("System", "APP_NAME", "app_name", "App Name"),
|
||||
ConfigField("System", "APP_ENV", "app_env", "App Env"),
|
||||
ConfigField("System", "APP_DEBUG", "app_debug", "App Debug"),
|
||||
ConfigField("System", "APP_DEBUG", "app_debug", "App Debug", input_type="checkbox"),
|
||||
ConfigField("System", "APP_HOSTNAME", "app_hostname", "App Hostname"),
|
||||
ConfigField("SMTP", "SMTP_ENABLED", "smtp_enabled", "SMTP Enabled"),
|
||||
ConfigField("SMTP", "SMTP_ENABLED", "smtp_enabled", "SMTP Enabled", input_type="checkbox"),
|
||||
ConfigField("SMTP", "SMTP_HOST", "smtp_host", "SMTP Host"),
|
||||
ConfigField("SMTP", "SMTP_PORT", "smtp_port", "SMTP Port"),
|
||||
ConfigField("SMTP", "SMTP_USERNAME", "smtp_username", "SMTP Username"),
|
||||
@@ -35,7 +35,7 @@ CONFIG_FIELDS: tuple[ConfigField, ...] = (
|
||||
ConfigField("SMTP", "SMTP_FROM_NAME", "smtp_from_name", "SMTP From Name"),
|
||||
ConfigField("SMTP", "SMTP_FROM_ADDRESS", "smtp_from_address", "SMTP From Address"),
|
||||
ConfigField("SMTP", "SMTP_TO_ADDRESS", "smtp_to_address", "SMTP To Address"),
|
||||
ConfigField("SMTP", "SMTP_USE_STARTTLS", "smtp_use_starttls", "SMTP Use STARTTLS"),
|
||||
ConfigField("SMTP", "SMTP_USE_STARTTLS", "smtp_use_starttls", "SMTP Use STARTTLS", input_type="checkbox"),
|
||||
ConfigField(
|
||||
"Authentication",
|
||||
"AUTH_SESSION_COOKIE_NAME",
|
||||
@@ -49,6 +49,13 @@ CONFIG_FIELDS: tuple[ConfigField, ...] = (
|
||||
"auth_cookie_secure_override",
|
||||
"Cookie Secure Override",
|
||||
),
|
||||
ConfigField(
|
||||
"Authentication",
|
||||
"AUTH_LOGIN_THROTTLE_ENABLED",
|
||||
"auth_login_throttle_enabled",
|
||||
"Login Throttle Enabled",
|
||||
input_type="checkbox",
|
||||
),
|
||||
ConfigField("Poo", "POO_WEBHOOK_ID", "poo_webhook_id", "Poo Webhook ID", secret=True),
|
||||
ConfigField(
|
||||
"Poo",
|
||||
@@ -96,6 +103,41 @@ CONFIG_FIELDS: tuple[ConfigField, ...] = (
|
||||
"home_assistant_action_task_project_id",
|
||||
"Home Assistant Action Task Project ID",
|
||||
),
|
||||
ConfigField("MQTT", "MQTT_ENABLED", "mqtt_enabled", "MQTT Enabled", input_type="checkbox"),
|
||||
ConfigField("MQTT", "MQTT_BROKER_HOST", "mqtt_broker_host", "MQTT Broker Host"),
|
||||
ConfigField("MQTT", "MQTT_BROKER_PORT", "mqtt_broker_port", "MQTT Broker Port", input_type="number"),
|
||||
ConfigField("MQTT", "MQTT_USERNAME", "mqtt_username", "MQTT Username"),
|
||||
ConfigField("MQTT", "MQTT_PASSWORD", "mqtt_password", "MQTT Password", secret=True),
|
||||
ConfigField("MQTT", "MQTT_TLS_ENABLED", "mqtt_tls_enabled", "MQTT TLS Enabled", input_type="checkbox"),
|
||||
ConfigField("MQTT", "MQTT_CLIENT_ID", "mqtt_client_id", "MQTT Client ID"),
|
||||
ConfigField(
|
||||
"Home Assistant Discovery",
|
||||
"HA_DISCOVERY_ENABLED",
|
||||
"ha_discovery_enabled",
|
||||
"HA Discovery Enabled",
|
||||
input_type="checkbox",
|
||||
),
|
||||
ConfigField(
|
||||
"Home Assistant Discovery",
|
||||
"HA_DISCOVERY_PREFIX",
|
||||
"ha_discovery_prefix",
|
||||
"HA Discovery Prefix",
|
||||
),
|
||||
ConfigField(
|
||||
"Home Assistant Discovery",
|
||||
"HA_STATE_TOPIC_PREFIX",
|
||||
"ha_state_topic_prefix",
|
||||
"HA State Topic Prefix",
|
||||
),
|
||||
ConfigField("Modbus", "MODBUS_POLLING_ENABLED", "modbus_polling_enabled", "Modbus Polling Enabled", input_type="checkbox"),
|
||||
ConfigField(
|
||||
"Tibber",
|
||||
"TIBBER_API_TOKEN",
|
||||
"tibber_api_token",
|
||||
"Tibber API Token",
|
||||
secret=True,
|
||||
),
|
||||
ConfigField("Tibber", "TIBBER_HOME_ID", "tibber_home_id", "Tibber Home ID"),
|
||||
)
|
||||
|
||||
|
||||
@@ -181,7 +223,12 @@ def save_config_updates(session: Session, form_data: dict[str, str], bootstrap_s
|
||||
else:
|
||||
merged_values[field.env_name] = submitted_value
|
||||
|
||||
_validate_config_values(merged_values, bootstrap_settings)
|
||||
validated_settings = _validate_config_values(merged_values, bootstrap_settings)
|
||||
# Persist the canonical client identity as well as using it at runtime. A
|
||||
# whitespace-padded value must not survive in app_config and unexpectedly
|
||||
# reappear in another consumer of the stored settings.
|
||||
if "MQTT_CLIENT_ID" in merged_values:
|
||||
merged_values["MQTT_CLIENT_ID"] = validated_settings.mqtt_client_id
|
||||
_persist_config_values(session, merged_values)
|
||||
get_settings.cache_clear()
|
||||
reset_db_caches()
|
||||
@@ -196,7 +243,9 @@ def save_config_value(
|
||||
) -> None:
|
||||
current_values = _read_config_values(session)
|
||||
current_values[env_name] = value
|
||||
_validate_config_values(current_values, bootstrap_settings)
|
||||
validated_settings = _validate_config_values(current_values, bootstrap_settings)
|
||||
if env_name == "MQTT_CLIENT_ID":
|
||||
current_values[env_name] = validated_settings.mqtt_client_id
|
||||
_persist_config_values(session, current_values)
|
||||
get_settings.cache_clear()
|
||||
reset_db_caches()
|
||||
@@ -215,14 +264,14 @@ def _read_config_values(session: Session) -> dict[str, str]:
|
||||
return {row.key: row.value for row in rows}
|
||||
|
||||
|
||||
def _validate_config_values(config_values: dict[str, str], bootstrap_settings: Settings) -> None:
|
||||
def _validate_config_values(config_values: dict[str, str], bootstrap_settings: Settings) -> Settings:
|
||||
payload = _settings_payload(bootstrap_settings)
|
||||
for field in CONFIG_FIELDS:
|
||||
if field.env_name in config_values:
|
||||
payload[field.setting_attr] = config_values[field.env_name]
|
||||
|
||||
try:
|
||||
Settings(_env_file=None, **payload)
|
||||
return Settings(_env_file=None, **payload)
|
||||
except Exception as exc:
|
||||
raise ConfigSaveError("invalid config submission") from exc
|
||||
|
||||
@@ -284,4 +333,19 @@ def _settings_payload(settings: Settings) -> dict[str, Any]:
|
||||
"auth_session_cookie_name": settings.auth_session_cookie_name,
|
||||
"auth_session_ttl_hours": settings.auth_session_ttl_hours,
|
||||
"auth_cookie_secure_override": settings.auth_cookie_secure_override,
|
||||
"auth_login_throttle_enabled": settings.auth_login_throttle_enabled,
|
||||
"auth_trust_forwarded_for": settings.auth_trust_forwarded_for,
|
||||
"modbus_polling_enabled": settings.modbus_polling_enabled,
|
||||
"mqtt_enabled": settings.mqtt_enabled,
|
||||
"mqtt_broker_host": settings.mqtt_broker_host,
|
||||
"mqtt_broker_port": settings.mqtt_broker_port,
|
||||
"mqtt_username": settings.mqtt_username,
|
||||
"mqtt_password": settings.mqtt_password,
|
||||
"mqtt_tls_enabled": settings.mqtt_tls_enabled,
|
||||
"mqtt_client_id": settings.mqtt_client_id,
|
||||
"ha_discovery_enabled": settings.ha_discovery_enabled,
|
||||
"ha_discovery_prefix": settings.ha_discovery_prefix,
|
||||
"ha_state_topic_prefix": settings.ha_state_topic_prefix,
|
||||
"tibber_api_token": settings.tibber_api_token,
|
||||
"tibber_home_id": settings.tibber_home_id,
|
||||
}
|
||||
|
||||
@@ -0,0 +1,409 @@
|
||||
"""Service layer for EnergyContract CRUD, versioning, and activation.
|
||||
|
||||
All functions accept an explicit SQLAlchemy Session; callers are responsible
|
||||
for committing or rolling back the transaction.
|
||||
|
||||
Design decisions
|
||||
----------------
|
||||
- ``create_contract``: validates values against the pricing profile *before*
|
||||
writing any rows; raises ``ProfileValidationError`` on non-compliance.
|
||||
- ``add_version``: append-only; closes the previous open version by
|
||||
setting its ``effective_to`` to the new version's ``effective_from``; raises
|
||||
``ContractVersionError`` if the new date is strictly earlier than the previous
|
||||
version's ``effective_from``.
|
||||
- ``activate_contract``: scope-local mutual exclusion; sets other contracts in
|
||||
the target scope inactive, then sets the given contract active.
|
||||
- ``active_contract_version_at``: returns the single version of the currently
|
||||
active contract that covers *ts* (``effective_from ≤ ts < effective_to``,
|
||||
or open-ended when ``effective_to`` is None).
|
||||
|
||||
SQLite timezone note
|
||||
--------------------
|
||||
SQLite stores ``DateTime(timezone=True)`` columns as naive UTC strings; on
|
||||
read-back they come out as **timezone-naive** datetimes. Wherever this code
|
||||
compares timestamps from the DB against timezone-aware values (e.g. from
|
||||
Pydantic or ``datetime.now(UTC)``), it calls ``_as_utc()`` to make both sides
|
||||
comparable without tripping on "offset-naive vs offset-aware" TypeErrors.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import select, update
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.integrations.pricing.profiles import validate_values
|
||||
from app.models.energy import EnergyContract, EnergyContractVersion
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# This is deliberately separate from the pricing-profile loader. T12 needs to
|
||||
# reserve the thermal domain before T13 supplies its actual profile.
|
||||
CONTRACT_KIND_SCOPES: dict[str, str] = {
|
||||
"manual": "electricity",
|
||||
"tibber": "electricity",
|
||||
"district_heating": "thermal",
|
||||
}
|
||||
|
||||
|
||||
class ContractScopeError(ValueError):
|
||||
"""Raised when a contract kind is unknown or its supplied scope disagrees."""
|
||||
|
||||
|
||||
def contract_scope_for_kind(kind: str, requested_scope: str | None = None) -> str:
|
||||
"""Return the registry-owned scope for *kind*, rejecting client mismatches."""
|
||||
try:
|
||||
scope = CONTRACT_KIND_SCOPES[kind]
|
||||
except KeyError as exc:
|
||||
raise ContractScopeError(f"Unknown energy contract kind: {kind!r}") from exc
|
||||
if requested_scope is not None and requested_scope != scope:
|
||||
raise ContractScopeError(
|
||||
f"Contract kind {kind!r} belongs to scope {scope!r}, not {requested_scope!r}."
|
||||
)
|
||||
return scope
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _as_utc(dt: datetime) -> datetime:
|
||||
"""Return *dt* as a timezone-aware UTC datetime.
|
||||
|
||||
SQLite's DateTime(timezone=True) column type stores datetimes as naive UTC
|
||||
strings and gives them back as naive datetimes on read. This helper
|
||||
re-attaches the UTC timezone info when it is missing, making cross-origin
|
||||
comparisons safe.
|
||||
"""
|
||||
if dt.tzinfo is None:
|
||||
return dt.replace(tzinfo=UTC)
|
||||
return dt
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Custom exception
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ContractVersionError(ValueError):
|
||||
"""Raised when a new contract version has an invalid effective date.
|
||||
|
||||
Specifically: the new version's ``effective_from`` must be greater than or
|
||||
equal to the previous open version's ``effective_from``. Allowing equal
|
||||
timestamps would cause ambiguous overlap; the service therefore also rejects
|
||||
strictly-equal values (same second) to avoid silent data loss.
|
||||
"""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def get_contract_or_none(session: Session, contract_id: int) -> EnergyContract | None:
|
||||
"""Return the contract with the given id, or None if not found."""
|
||||
return session.execute(
|
||||
select(EnergyContract).where(EnergyContract.id == contract_id)
|
||||
).scalar_one_or_none()
|
||||
|
||||
|
||||
def list_contracts(session: Session, *, scope: str = "electricity") -> list[EnergyContract]:
|
||||
"""Return contracts in one scope, ordered by id (ascending)."""
|
||||
return list(
|
||||
session.execute(
|
||||
select(EnergyContract)
|
||||
.where(EnergyContract.scope == scope)
|
||||
.order_by(EnergyContract.id)
|
||||
).scalars().all()
|
||||
)
|
||||
|
||||
|
||||
def _open_version(session: Session, contract: EnergyContract) -> EnergyContractVersion | None:
|
||||
"""Return the current open version (effective_to IS NULL) for *contract*, or None."""
|
||||
return session.execute(
|
||||
select(EnergyContractVersion)
|
||||
.where(
|
||||
EnergyContractVersion.contract_id == contract.id,
|
||||
EnergyContractVersion.effective_to.is_(None),
|
||||
)
|
||||
.order_by(EnergyContractVersion.effective_from.desc())
|
||||
.limit(1)
|
||||
).scalar_one_or_none()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Core service functions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def create_contract(
|
||||
session: Session,
|
||||
*,
|
||||
name: str,
|
||||
kind: str,
|
||||
currency: str = "EUR",
|
||||
scope: str | None = None,
|
||||
values: dict[str, Any],
|
||||
effective_from: datetime,
|
||||
) -> EnergyContract:
|
||||
"""Create a new energy contract with its first pricing version.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
session:
|
||||
Active SQLAlchemy session. Caller must commit after this returns.
|
||||
name:
|
||||
Human-readable label for the contract.
|
||||
kind:
|
||||
Pricing strategy identifier (``"manual"``, ``"tibber"``, or
|
||||
``"district_heating"``).
|
||||
currency:
|
||||
ISO 4217 currency code (default ``"EUR"``).
|
||||
values:
|
||||
Pricing values dict conforming to the named profile's structure. The
|
||||
district-heating profile normalises its Decimal-safe values to strings
|
||||
before the JSON snapshot is stored.
|
||||
Validated via ``validate_values(kind, values)`` before any writes.
|
||||
effective_from:
|
||||
UTC datetime at which the first pricing version takes effect.
|
||||
|
||||
Returns
|
||||
-------
|
||||
EnergyContract
|
||||
The newly created contract (not yet committed).
|
||||
|
||||
Raises
|
||||
------
|
||||
ProfileNotFoundError
|
||||
If no YAML profile exists for *kind*.
|
||||
ProfileValidationError
|
||||
If *values* does not conform to the profile structure.
|
||||
"""
|
||||
resolved_scope = contract_scope_for_kind(kind, scope)
|
||||
# Validate (and fill defaults) before any DB write.
|
||||
filled_values = validate_values(kind, values)
|
||||
|
||||
now = datetime.now(UTC)
|
||||
contract = EnergyContract(
|
||||
name=name,
|
||||
kind=kind,
|
||||
scope=resolved_scope,
|
||||
currency=currency,
|
||||
active=False, # New contracts are inactive; caller must explicitly activate.
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
session.add(contract)
|
||||
session.flush() # Assign contract.id so we can reference it in the version FK.
|
||||
|
||||
version = EnergyContractVersion(
|
||||
contract_id=contract.id,
|
||||
effective_from=effective_from,
|
||||
effective_to=None,
|
||||
values=filled_values,
|
||||
created_at=now,
|
||||
)
|
||||
session.add(version)
|
||||
|
||||
logger.info("Created energy contract %r (kind=%s, id=%d)", name, kind, contract.id)
|
||||
return contract
|
||||
|
||||
|
||||
def add_version(
|
||||
session: Session,
|
||||
contract: EnergyContract,
|
||||
*,
|
||||
effective_from: datetime,
|
||||
values: dict[str, Any],
|
||||
) -> EnergyContractVersion:
|
||||
"""Add a new pricing version to an existing contract (append-only).
|
||||
|
||||
The previous open version's ``effective_to`` is automatically set to the
|
||||
new version's ``effective_from`` (version closure), ensuring there is never
|
||||
a gap or overlap between consecutive versions.
|
||||
|
||||
The new ``effective_from`` **must be strictly greater than** the previous
|
||||
open version's ``effective_from``. Equal timestamps are rejected because
|
||||
they would produce two versions starting at the same instant, making it
|
||||
impossible to determine which is current. If this constraint is not met,
|
||||
``ContractVersionError`` is raised and **no rows are written**.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
session:
|
||||
Active SQLAlchemy session. Caller must commit after this returns.
|
||||
contract:
|
||||
The parent ``EnergyContract`` to append a version to.
|
||||
effective_from:
|
||||
UTC datetime at which this version's pricing takes effect.
|
||||
values:
|
||||
New pricing values dict conforming to the contract's profile.
|
||||
|
||||
Returns
|
||||
-------
|
||||
EnergyContractVersion
|
||||
The newly created version (not yet committed).
|
||||
|
||||
Raises
|
||||
------
|
||||
ContractVersionError
|
||||
If *effective_from* is not strictly after the previous open version's
|
||||
``effective_from``.
|
||||
ProfileNotFoundError
|
||||
If no YAML profile exists for the contract's kind.
|
||||
ProfileValidationError
|
||||
If *values* does not conform to the profile structure.
|
||||
"""
|
||||
# Validate (and fill defaults) before any DB write.
|
||||
filled_values = validate_values(contract.kind, values)
|
||||
|
||||
prev = _open_version(session, contract)
|
||||
if prev is not None:
|
||||
# Enforce strictly-after constraint to prevent ambiguous overlap.
|
||||
# Use _as_utc() on both sides: the incoming value may be tz-aware
|
||||
# while the DB-read value is tz-naive (SQLite limitation).
|
||||
if _as_utc(effective_from) <= _as_utc(prev.effective_from):
|
||||
raise ContractVersionError(
|
||||
f"New version effective_from ({effective_from.isoformat()}) must be strictly "
|
||||
f"after the previous open version's effective_from "
|
||||
f"({prev.effective_from.isoformat()})."
|
||||
)
|
||||
# Close the previous open version.
|
||||
prev.effective_to = effective_from
|
||||
|
||||
now = datetime.now(UTC)
|
||||
new_version = EnergyContractVersion(
|
||||
contract_id=contract.id,
|
||||
effective_from=effective_from,
|
||||
effective_to=None,
|
||||
values=filled_values,
|
||||
created_at=now,
|
||||
)
|
||||
session.add(new_version)
|
||||
|
||||
logger.info(
|
||||
"Added version to contract id=%d (kind=%s, effective_from=%s)",
|
||||
contract.id,
|
||||
contract.kind,
|
||||
effective_from.isoformat(),
|
||||
)
|
||||
return new_version
|
||||
|
||||
|
||||
def activate_contract(session: Session, contract: EnergyContract) -> None:
|
||||
"""Activate a contract with mutual exclusion.
|
||||
|
||||
Sets every other contract in the same scope inactive, then sets the given
|
||||
contract active. This guarantees at most one active contract per scope.
|
||||
|
||||
Caller must commit after this returns.
|
||||
"""
|
||||
# This bulk update is a single write statement inside the caller's
|
||||
# transaction. SQLite serializes writers, and another scope is never touched.
|
||||
session.execute(
|
||||
update(EnergyContract)
|
||||
.where(EnergyContract.scope == contract.scope, EnergyContract.id != contract.id)
|
||||
.values(active=False)
|
||||
)
|
||||
contract.active = True
|
||||
contract.updated_at = datetime.now(UTC)
|
||||
logger.info("Activated contract %r (id=%d)", contract.name, contract.id)
|
||||
|
||||
|
||||
def deactivate_contract(session: Session, contract: EnergyContract) -> None:
|
||||
"""Deactivate a contract without touching other contracts.
|
||||
|
||||
Caller must commit after this returns.
|
||||
"""
|
||||
contract.active = False
|
||||
contract.updated_at = datetime.now(UTC)
|
||||
logger.info("Deactivated contract %r (id=%d)", contract.name, contract.id)
|
||||
|
||||
|
||||
def active_contract_versions(
|
||||
session: Session, *, scope: str = "electricity"
|
||||
) -> list[EnergyContractVersion]:
|
||||
"""Return all versions of the currently active contract, ordered by effective_from ascending.
|
||||
|
||||
Returns an empty list when there is no active contract. The list spans the
|
||||
full history of the active contract (all closed + the current open version)
|
||||
and is used to iterate over pricing-rate segments for cross-version fixed-
|
||||
cost / credit accumulation (Principle C).
|
||||
"""
|
||||
active = session.execute(
|
||||
select(EnergyContract)
|
||||
.where(EnergyContract.active.is_(True), EnergyContract.scope == scope)
|
||||
.limit(1)
|
||||
).scalar_one_or_none()
|
||||
|
||||
if active is None:
|
||||
return []
|
||||
|
||||
return list(
|
||||
session.execute(
|
||||
select(EnergyContractVersion)
|
||||
.where(EnergyContractVersion.contract_id == active.id)
|
||||
.order_by(EnergyContractVersion.effective_from)
|
||||
)
|
||||
.scalars()
|
||||
.all()
|
||||
)
|
||||
|
||||
|
||||
def active_contract_version_at(
|
||||
session: Session, ts: datetime, *, scope: str = "electricity"
|
||||
) -> EnergyContractVersion | None:
|
||||
"""Return the active contract's version that covers *ts*.
|
||||
|
||||
A version covers *ts* when:
|
||||
``effective_from ≤ ts`` AND (``effective_to IS NULL`` OR ``ts < effective_to``)
|
||||
|
||||
Returns None when:
|
||||
- There is no active contract.
|
||||
- The active contract has no version covering *ts*.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
session:
|
||||
Active SQLAlchemy session (read-only usage).
|
||||
ts:
|
||||
UTC datetime to look up.
|
||||
|
||||
Returns
|
||||
-------
|
||||
EnergyContractVersion | None
|
||||
"""
|
||||
active = session.execute(
|
||||
select(EnergyContract)
|
||||
.where(EnergyContract.active.is_(True), EnergyContract.scope == scope)
|
||||
.limit(1)
|
||||
).scalar_one_or_none()
|
||||
|
||||
if active is None:
|
||||
return None
|
||||
|
||||
# Build a query for all versions of the active contract covering ts.
|
||||
stmt = (
|
||||
select(EnergyContractVersion)
|
||||
.where(
|
||||
EnergyContractVersion.contract_id == active.id,
|
||||
EnergyContractVersion.effective_from <= ts,
|
||||
)
|
||||
.order_by(EnergyContractVersion.effective_from.desc())
|
||||
.limit(1)
|
||||
)
|
||||
version = session.execute(stmt).scalar_one_or_none()
|
||||
|
||||
if version is None:
|
||||
return None
|
||||
|
||||
# Exclude versions whose effective window has already closed before ts.
|
||||
if version.effective_to is not None and _as_utc(ts) >= _as_utc(version.effective_to):
|
||||
return None
|
||||
|
||||
return version
|
||||
@@ -0,0 +1,431 @@
|
||||
"""DSMR MQTT ingest, keyed by durable meter-source identity."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import sqlalchemy.exc
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.db import get_session_local
|
||||
from app.models.energy import DsmrReading, Meter
|
||||
from app.models.meter_source import MeterSource, MeterSourceBinding, MeterSourceChannel
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from app.config import Settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DsmrSourceSnapshot:
|
||||
"""The only DSMR runtime configuration a network handler may use."""
|
||||
|
||||
source_id: int
|
||||
topic: str
|
||||
tariff_topic: str
|
||||
sample_interval_s: int
|
||||
broker_host: str = ""
|
||||
broker_port: int = 1883
|
||||
username: str = ""
|
||||
password: str = ""
|
||||
tls_enabled: bool = False
|
||||
|
||||
|
||||
_subscriptions: dict[int, DsmrSourceSnapshot] = {}
|
||||
_subscription_client_ids: dict[int, str] = {}
|
||||
_subscription_lock = threading.RLock()
|
||||
# A configuration value is not an ownership identity: disable and re-enable
|
||||
# can produce an equal snapshot. Each installed handler therefore captures a
|
||||
# fresh token and verifies object identity before it can write or update tariff.
|
||||
_subscription_tokens: dict[int, object] = {}
|
||||
_reconcile_lock = threading.RLock()
|
||||
_tariffs: dict[int, int] = {}
|
||||
_tariff_lock = threading.Lock()
|
||||
# Kept only for legacy direct test callers of set_current_tariff(value). Runtime
|
||||
# handlers never write this value; production callers resolve a binding first.
|
||||
_current_tariff: int | None = None
|
||||
_UNSET = object()
|
||||
|
||||
|
||||
def get_current_tariff(meter_source_id: int | None = None) -> int | None:
|
||||
"""Return a source tariff, or resolve the active electricity binding.
|
||||
|
||||
The no-argument form is retained for the pre-M8 expose integration. It
|
||||
opens a short session to select the current electricity binding, so a
|
||||
source's MQTT callback can never make another source's tariff current.
|
||||
``_current_tariff`` is solely a test-era fallback when no binding database
|
||||
is available; runtime MQTT handlers do not update it.
|
||||
"""
|
||||
if meter_source_id is None:
|
||||
session_local = get_session_local()
|
||||
session = session_local()
|
||||
try:
|
||||
source_id = _current_electricity_source_id(session, datetime.now(timezone.utc))
|
||||
if source_id is not None:
|
||||
return get_current_tariff(source_id)
|
||||
except Exception:
|
||||
logger.debug("DSMR legacy tariff lookup could not resolve a binding", exc_info=True)
|
||||
finally:
|
||||
session.close()
|
||||
return _current_tariff
|
||||
with _tariff_lock:
|
||||
return _tariffs.get(meter_source_id)
|
||||
|
||||
|
||||
def set_current_tariff(meter_source_id: int, value: int | None | object = _UNSET) -> None:
|
||||
"""Set or clear an individual source's tariff state."""
|
||||
global _current_tariff
|
||||
if value is _UNSET:
|
||||
# Compatibility with older direct callers. Do not route runtime source
|
||||
# updates through this global fallback.
|
||||
_current_tariff = meter_source_id if meter_source_id in (1, 2) else None
|
||||
return
|
||||
with _tariff_lock:
|
||||
if value is None:
|
||||
_tariffs.pop(meter_source_id, None)
|
||||
else:
|
||||
_tariffs[meter_source_id] = value
|
||||
|
||||
|
||||
def _current_electricity_source_id(session: Session, at: datetime) -> int | None:
|
||||
"""Return the DSMR source bound to electricity at ``at``, if any."""
|
||||
return session.scalar(
|
||||
select(MeterSource.id)
|
||||
.join(MeterSourceChannel, MeterSourceChannel.source_id == MeterSource.id)
|
||||
.join(MeterSourceBinding, MeterSourceBinding.channel_id == MeterSourceChannel.id)
|
||||
.join(Meter, Meter.id == MeterSourceBinding.meter_id)
|
||||
.where(
|
||||
Meter.commodity == "electricity",
|
||||
MeterSourceBinding.started_at <= at,
|
||||
(MeterSourceBinding.ended_at.is_(None)) | (MeterSourceBinding.ended_at > at),
|
||||
)
|
||||
.order_by(MeterSourceBinding.started_at.desc())
|
||||
.limit(1)
|
||||
)
|
||||
|
||||
|
||||
def get_current_electricity_tariff(session: Session, at: datetime | None = None) -> int | None:
|
||||
"""Resolve tariff through the current electricity binding, never globally."""
|
||||
source_id = _current_electricity_source_id(session, at or datetime.now(timezone.utc))
|
||||
if source_id is None:
|
||||
return None
|
||||
return get_current_tariff(source_id)
|
||||
|
||||
|
||||
def handle_tariff_message(payload_bytes: bytes, meter_source_id: int) -> None:
|
||||
"""Parse one source's tariff payload without raising in paho's thread."""
|
||||
try:
|
||||
raw = (
|
||||
payload_bytes.decode("utf-8", errors="replace").strip()
|
||||
if isinstance(payload_bytes, (bytes, bytearray))
|
||||
else str(payload_bytes).strip()
|
||||
)
|
||||
value = int(raw)
|
||||
if value in (1, 2):
|
||||
set_current_tariff(meter_source_id, value)
|
||||
except Exception:
|
||||
logger.debug("DSMR tariff payload ignored for source_id=%s", meter_source_id)
|
||||
|
||||
|
||||
def _snapshot(source: MeterSource) -> DsmrSourceSnapshot:
|
||||
config = source.config
|
||||
return DsmrSourceSnapshot(
|
||||
source_id=source.id,
|
||||
topic=str(config.get("topic", "dsmr/json")),
|
||||
tariff_topic=str(config.get("tariff_topic", "")),
|
||||
sample_interval_s=int(config.get("sample_interval_s", 10)),
|
||||
broker_host=str(config.get("broker_host", "")),
|
||||
broker_port=int(config.get("broker_port", 1883)),
|
||||
username=str(config.get("username", "")),
|
||||
password=str(config.get("password", "")),
|
||||
tls_enabled=bool(config.get("tls_enabled", False)),
|
||||
)
|
||||
|
||||
|
||||
def _enabled_snapshots() -> list[DsmrSourceSnapshot]:
|
||||
session_local = get_session_local()
|
||||
session = session_local()
|
||||
try:
|
||||
sources = session.scalars(
|
||||
select(MeterSource).where(MeterSource.kind == "dsmr_mqtt", MeterSource.enabled.is_(True))
|
||||
).all()
|
||||
return [_snapshot(source) for source in sources]
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def apply_dsmr_subscription(settings: "Settings | None" = None) -> None:
|
||||
"""Reconcile enabled DSMR source subscriptions from the database.
|
||||
|
||||
``settings`` provides the DB-merged app-wide MQTT identity. Individual
|
||||
DSMR broker settings continue to come solely from MeterSource records.
|
||||
"""
|
||||
from app.integrations.mqtt import mqtt_manager
|
||||
from app.config import get_settings
|
||||
|
||||
base_client_id = (settings or get_settings()).mqtt_client_id
|
||||
|
||||
try:
|
||||
desired = {snapshot.source_id: snapshot for snapshot in _enabled_snapshots()}
|
||||
except Exception:
|
||||
logger.exception("DSMR subscription reconcile failed while reading sources")
|
||||
return
|
||||
|
||||
# Do not retain _subscription_lock while stopping MQTT clients: a callback
|
||||
# may currently hold it through its complete dispatch, and MqttManager
|
||||
# waits for that callback before teardown returns.
|
||||
with _reconcile_lock:
|
||||
# Every source owns a distinct MQTT client, so equal topics from
|
||||
# different sources/brokers are dispatchable. A telegram and tariff
|
||||
# topic on the *same* client would overwrite one handler, however.
|
||||
rejected_source_ids: set[int] = set()
|
||||
for source_id, snapshot in list(desired.items()):
|
||||
if snapshot.tariff_topic and snapshot.topic == snapshot.tariff_topic:
|
||||
logger.error("DSMR source_id=%s rejected: telegram/tariff topic collision", source_id)
|
||||
rejected_source_ids.add(source_id)
|
||||
desired.pop(source_id)
|
||||
|
||||
with _subscription_lock:
|
||||
stale_source_ids = [
|
||||
source_id
|
||||
for source_id, current in _subscriptions.items()
|
||||
if desired.get(source_id) != current
|
||||
or _subscription_client_ids.get(source_id) != base_client_id
|
||||
]
|
||||
for source_id in stale_source_ids:
|
||||
_subscriptions.pop(source_id, None)
|
||||
_subscription_client_ids.pop(source_id, None)
|
||||
_subscription_tokens.pop(source_id, None)
|
||||
set_current_tariff(source_id, None)
|
||||
|
||||
# A disabled source has no installed MQTT owner. Persist that fact
|
||||
# after invalidating its callback token, so a retained callback cannot
|
||||
# revive an earlier online state while teardown is in progress.
|
||||
for source_id in stale_source_ids:
|
||||
_mark_disabled_source_inactive(source_id)
|
||||
|
||||
# A topic collision is a configuration error for an enabled source,
|
||||
# not a disabled-state transition. It must therefore replace any
|
||||
# earlier online state even when this process started without a
|
||||
# matching runtime subscription to tear down.
|
||||
for source_id in rejected_source_ids:
|
||||
_mark_rejected_source_error(source_id)
|
||||
|
||||
for source_id in stale_source_ids:
|
||||
mqtt_manager.remove_source(source_id)
|
||||
|
||||
for source_id, snapshot in desired.items():
|
||||
with _subscription_lock:
|
||||
current = _subscriptions.get(source_id)
|
||||
current_client_id = _subscription_client_ids.get(source_id)
|
||||
if (
|
||||
current == snapshot
|
||||
and current_client_id == base_client_id
|
||||
and mqtt_manager.source_is_active(source_id)
|
||||
):
|
||||
continue
|
||||
if current is not None:
|
||||
# The client went inactive outside reconcile. Invalidate its
|
||||
# old token before rebuilding the same snapshot.
|
||||
with _subscription_lock:
|
||||
if _subscriptions.get(source_id) == current:
|
||||
_subscriptions.pop(source_id, None)
|
||||
_subscription_client_ids.pop(source_id, None)
|
||||
_subscription_tokens.pop(source_id, None)
|
||||
mqtt_manager.remove_source(source_id)
|
||||
token = object()
|
||||
handlers = {
|
||||
snapshot.topic: lambda payload, captured=snapshot, captured_token=token: (
|
||||
handle_captured_message(payload, captured, captured_token)
|
||||
)
|
||||
}
|
||||
if snapshot.tariff_topic:
|
||||
handlers[snapshot.tariff_topic] = (
|
||||
lambda payload, captured=snapshot, captured_token=token: (
|
||||
handle_captured_tariff_message(payload, captured, captured_token)
|
||||
)
|
||||
)
|
||||
with _subscription_lock:
|
||||
_subscriptions[source_id] = snapshot
|
||||
_subscription_client_ids[source_id] = base_client_id
|
||||
_subscription_tokens[source_id] = token
|
||||
applied = mqtt_manager.replace_source(
|
||||
source_id,
|
||||
host=snapshot.broker_host,
|
||||
port=snapshot.broker_port,
|
||||
username=snapshot.username,
|
||||
password=snapshot.password,
|
||||
tls_enabled=snapshot.tls_enabled,
|
||||
subscriptions=handlers,
|
||||
base_client_id=base_client_id,
|
||||
state_handler=lambda state, captured=snapshot, captured_token=token: (
|
||||
handle_captured_source_state(captured, captured_token, state)
|
||||
),
|
||||
)
|
||||
if not applied:
|
||||
with _subscription_lock:
|
||||
if _subscription_tokens.get(source_id) is token:
|
||||
_subscriptions.pop(source_id, None)
|
||||
_subscription_client_ids.pop(source_id, None)
|
||||
_subscription_tokens.pop(source_id, None)
|
||||
|
||||
|
||||
def handle_message(payload_bytes: bytes, snapshot: DsmrSourceSnapshot) -> None:
|
||||
"""Persist one down-sampled frame under its captured source identity."""
|
||||
try:
|
||||
_handle_message_inner(payload_bytes, snapshot)
|
||||
except Exception:
|
||||
logger.exception("DSMR ingest handler failed for source_id=%s (swallowed)", snapshot.source_id)
|
||||
|
||||
|
||||
def handle_captured_message(
|
||||
payload_bytes: bytes, snapshot: DsmrSourceSnapshot, token: object | None = None
|
||||
) -> None:
|
||||
"""Run a broker callback only while its exact source generation is active."""
|
||||
with _subscription_lock:
|
||||
if token is not None:
|
||||
if _subscription_tokens.get(snapshot.source_id) is not token:
|
||||
return
|
||||
elif _subscriptions.get(snapshot.source_id) != snapshot:
|
||||
return
|
||||
handle_message(payload_bytes, snapshot)
|
||||
|
||||
|
||||
def handle_captured_tariff_message(
|
||||
payload_bytes: bytes, snapshot: DsmrSourceSnapshot, token: object | None = None
|
||||
) -> None:
|
||||
"""Ignore tariff callbacks retained from a removed/replaced source."""
|
||||
with _subscription_lock:
|
||||
if token is not None:
|
||||
if _subscription_tokens.get(snapshot.source_id) is not token:
|
||||
return
|
||||
elif _subscriptions.get(snapshot.source_id) != snapshot:
|
||||
return
|
||||
handle_tariff_message(payload_bytes, snapshot.source_id)
|
||||
|
||||
|
||||
def handle_captured_source_state(snapshot: DsmrSourceSnapshot, token: object, state: str) -> None:
|
||||
"""Persist one active generation's connection health in a short DB session."""
|
||||
with _subscription_lock:
|
||||
if _subscription_tokens.get(snapshot.source_id) is not token:
|
||||
return
|
||||
session_local = get_session_local()
|
||||
session = session_local()
|
||||
try:
|
||||
source = session.get(MeterSource, snapshot.source_id)
|
||||
if source is None or not source.enabled or source.kind != "dsmr_mqtt":
|
||||
return
|
||||
source.status = state
|
||||
source.last_error = "MQTT connection failed." if state == "error" else None
|
||||
source.updated_at = datetime.now(timezone.utc)
|
||||
session.commit()
|
||||
except Exception:
|
||||
session.rollback()
|
||||
logger.exception("DSMR source health update failed for source_id=%s", snapshot.source_id)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def _mark_disabled_source_inactive(source_id: int) -> None:
|
||||
"""Clear an obsolete online health state for a disabled DSMR source.
|
||||
|
||||
The caller has already invalidated the source's generation token. This
|
||||
helper deliberately opens its own short session so reconcile never shares
|
||||
a callback-thread transaction. Deleted sources simply have no row left
|
||||
to update.
|
||||
"""
|
||||
session_local = get_session_local()
|
||||
session = session_local()
|
||||
try:
|
||||
source = session.get(MeterSource, source_id)
|
||||
if source is None or source.enabled or source.kind != "dsmr_mqtt":
|
||||
return
|
||||
source.status = "unknown"
|
||||
source.last_error = None
|
||||
source.updated_at = datetime.now(timezone.utc)
|
||||
session.commit()
|
||||
except Exception:
|
||||
session.rollback()
|
||||
logger.exception("DSMR disabled source health update failed for source_id=%s", source_id)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def _mark_rejected_source_error(source_id: int) -> None:
|
||||
"""Persist a non-sensitive error for an enabled source rejected by reconcile."""
|
||||
session_local = get_session_local()
|
||||
session = session_local()
|
||||
try:
|
||||
source = session.get(MeterSource, source_id)
|
||||
if source is None or not source.enabled or source.kind != "dsmr_mqtt":
|
||||
return
|
||||
source.status = "error"
|
||||
source.last_error = "DSMR source configuration invalid."
|
||||
source.updated_at = datetime.now(timezone.utc)
|
||||
session.commit()
|
||||
except Exception:
|
||||
session.rollback()
|
||||
logger.exception("DSMR rejected source health update failed for source_id=%s", source_id)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def _handle_message_inner(payload_bytes: bytes, snapshot: DsmrSourceSnapshot) -> None:
|
||||
try:
|
||||
data = json.loads(payload_bytes)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return
|
||||
if not isinstance(data, dict):
|
||||
return
|
||||
try:
|
||||
raw_ts = data["timestamp"]
|
||||
ts_utc = datetime.fromisoformat(raw_ts.replace("Z", "+00:00"))
|
||||
if ts_utc.tzinfo is None:
|
||||
ts_utc = ts_utc.replace(tzinfo=timezone.utc)
|
||||
except (KeyError, ValueError, TypeError, AttributeError):
|
||||
return
|
||||
if snapshot.sample_interval_s > 0 and ts_utc.second % snapshot.sample_interval_s:
|
||||
return
|
||||
telegram_id = data.get("id")
|
||||
if telegram_id is not None and not isinstance(telegram_id, int):
|
||||
telegram_id = None
|
||||
|
||||
session_local = get_session_local()
|
||||
session = session_local()
|
||||
try:
|
||||
exists = session.scalar(
|
||||
select(DsmrReading.id).where(
|
||||
DsmrReading.meter_source_id == snapshot.source_id,
|
||||
DsmrReading.recorded_at == ts_utc,
|
||||
)
|
||||
)
|
||||
if exists is None:
|
||||
session.add(
|
||||
DsmrReading(
|
||||
meter_source_id=snapshot.source_id,
|
||||
recorded_at=ts_utc,
|
||||
telegram_id=telegram_id,
|
||||
payload=data,
|
||||
)
|
||||
)
|
||||
source = session.get(MeterSource, snapshot.source_id)
|
||||
if source is not None and source.enabled and source.kind == "dsmr_mqtt":
|
||||
source.status = "online"
|
||||
source.last_seen_at = datetime.now(timezone.utc)
|
||||
source.last_error = None
|
||||
source.updated_at = datetime.now(timezone.utc)
|
||||
session.commit()
|
||||
except sqlalchemy.exc.IntegrityError:
|
||||
session.rollback()
|
||||
except Exception:
|
||||
session.rollback()
|
||||
logger.exception("DSMR database write failed for source_id=%s", snapshot.source_id)
|
||||
finally:
|
||||
session.close()
|
||||
@@ -0,0 +1,992 @@
|
||||
"""Billing engine for DSMR 15-minute energy metering periods.
|
||||
|
||||
This module implements the two-layer billing model described in §3.4 of the
|
||||
M6 design document, extended in M7-T03 to be meter-aware:
|
||||
|
||||
**Layer 1 — per-period metering cost (immutable, price-snapshot)**
|
||||
``compute_period(session, t0)`` computes the import cost, export revenue, and
|
||||
net cost for the 15-minute period ``[t0, t0+15min)``. The result is written
|
||||
to ``energy_cost_period`` with a full pricing snapshot so each row is
|
||||
self-contained and auditable. Existing *successful* rows are never overwritten
|
||||
by the normal tick path; only an explicit ``recompute_range`` call passes
|
||||
``overwrite=True``.
|
||||
|
||||
**Layer 2 — summary (computed at read time, not stored)**
|
||||
``summarize(session, start, end)`` aggregates all non-degraded
|
||||
``energy_cost_period`` rows in ``[start, end)``, then adds the daily
|
||||
standing charges (network_fee + management_fee, apportioned at EUR/month
|
||||
÷ 30 per day) and subtracts the energy-tax credit (heffingskorting,
|
||||
apportioned at EUR/year ÷ 365 per day).
|
||||
|
||||
The summary reports **both** money and energy: ``metered_import`` /
|
||||
``metered_export`` are monetary totals (Σ import_cost / Σ export_revenue),
|
||||
while ``metered_import_kwh`` / ``metered_export_kwh`` are the corresponding
|
||||
metered energy totals in kWh. The ``_kwh`` suffix is the only thing that
|
||||
distinguishes them — always check it before labelling a value in a UI.
|
||||
|
||||
Design notes
|
||||
------------
|
||||
- **Decimal arithmetic throughout**: all monetary computations use
|
||||
``decimal.Decimal`` to avoid float binary rounding errors. Only when
|
||||
writing to ``EnergyCostPeriod`` columns (Float) are values converted to
|
||||
float. ``summarize`` converts back to Decimal for summation.
|
||||
- **UTC quarter-hour grid**: period boundaries are aligned to UTC 00/15/30/45
|
||||
minutes (``floor_to_quarter``). NL local time (CET/CEST) is always a whole
|
||||
number of hours from UTC, so the quarter-hour grid is the same in both
|
||||
timezone representations.
|
||||
- **Register keys**: DSMR payload uses JSON strings like ``"20915.154"``
|
||||
for cumulative kWh registers. ``register_at`` converts them to Decimal.
|
||||
- **Degraded vs skip semantics**:
|
||||
- *No unique meter coverage* (no sole electricity meter at t0): write a
|
||||
``degraded=True`` row with ``meter_id=None``.
|
||||
- *Cross-meter boundary* (m0.id != m1.id for t0/t1): write a ``degraded=True``
|
||||
row with ``meter_id=m0.id``; losing this one period at the swap boundary is
|
||||
acceptable (D5 decision).
|
||||
- *Missing readings* (``register_at`` returns None for start or end
|
||||
boundary within the meter window): write a ``degraded=True`` row with
|
||||
``meter_id=m0.id`` so the period is tracked and can be retried by
|
||||
``compute_closed_periods``.
|
||||
- *Negative or excessively large delta* (delta sanity guard D6): write a
|
||||
``degraded=True`` row with ``meter_id=m0.id``; prevents negative costs and
|
||||
grossly inflated costs from meter resets, DSMR rollover, or data spikes.
|
||||
- *Missing Tibber price* (``TibberPriceNotFoundError``): skip entirely (do
|
||||
not write a row); the period will be retried once prices arrive.
|
||||
- *Missing active contract version*: skip (no contract to compute against).
|
||||
- **Meter-aware register lookup**: ``register_at`` now accepts a ``meter``
|
||||
parameter and restricts the DSMR reading query to readings within
|
||||
``[meter.started_at, meter.ended_at)`` (half-open), preventing old-meter
|
||||
readings from leaking into a new-meter epoch.
|
||||
- **Lookback window in ``compute_closed_periods``**: to avoid scanning all
|
||||
historical DSMR data on every tick, the function looks back at most 7 days
|
||||
from the current time. This covers typical short outages (no data / no
|
||||
contract) while staying bounded. Periods older than 7 days must be
|
||||
recovered via an explicit ``recompute_range`` call.
|
||||
|
||||
Meter-aware compute_period ordering rationale (M7-T03)
|
||||
-------------------------------------------------------
|
||||
The order of checks inside ``compute_period`` is:
|
||||
|
||||
1. **Immutability guard** (existing non-degraded row, overwrite=False) → return False.
|
||||
2. **Meter determination** (m0/m1 each resolve to one electricity Meter):
|
||||
- No unique meter (m0 is None) → write degraded, meter_id=None.
|
||||
- Cross-meter boundary (m0.id != m1.id) → write degraded, meter_id=m0.id.
|
||||
3. **Active contract version check** → skip (no write) if absent.
|
||||
4. **Boundary register readings** within m0's window → write degraded if missing.
|
||||
5. **Delta sanity guard** → write degraded if any delta < 0 or > _MAX_DELTA_KWH.
|
||||
6. **Price strategy** → skip (no write) if Tibber price missing.
|
||||
7. **Upsert billing record** with meter_id=m0.id.
|
||||
|
||||
Why meter before contract? The meter is a *structural* prerequisite: without a
|
||||
known meter epoch we cannot trust the delta at all, so we commit a degraded row
|
||||
immediately. The contract skip, by contrast, is transient (the period can be
|
||||
re-computed once a contract is configured), so it produces no row.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.integrations.pricing.strategies import (
|
||||
PeriodDeltas,
|
||||
TibberPriceNotFoundError,
|
||||
get_strategy,
|
||||
)
|
||||
from app.models.energy import DsmrReading, EnergyCostPeriod, Meter
|
||||
from app.models.meter_source import MeterSource, MeterSourceBinding, MeterSourceChannel
|
||||
from app.services.contracts import active_contract_version_at, active_contract_versions
|
||||
from app.services.timezone import local_date, local_now
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Constants
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_PERIOD_MINUTES = 15
|
||||
_LOOKBACK_DAYS = 7 # maximum lookback window for compute_closed_periods
|
||||
|
||||
# Maximum age a DSMR reading may have relative to the boundary being queried.
|
||||
# Under normal operation DSMR readings arrive every ~10 seconds, so a reading
|
||||
# more than one period (15 minutes) old at the boundary indicates either a data
|
||||
# gap or — critically — a *future* boundary being resolved against the last
|
||||
# historical reading. In both cases the reading is considered stale and
|
||||
# ``register_at`` returns None, letting the period be marked degraded instead of
|
||||
# producing a spurious zero-delta "successful" row.
|
||||
_READING_MAX_STALENESS = timedelta(minutes=_PERIOD_MINUTES)
|
||||
|
||||
# Maximum plausible kWh delta for a single 15-minute period (D6 sanity guard).
|
||||
# A typical Dutch household uses well under 5 kWh per quarter hour even under
|
||||
# heavy load. 100 kWh per 15 minutes corresponds to ~400 kW — far beyond any
|
||||
# residential consumption — but is lenient enough to never fire on legitimate
|
||||
# data. Any delta at or above this threshold indicates a meter reset, DSMR
|
||||
# rollover, sign error, or other data anomaly, and the period is marked
|
||||
# degraded to prevent negative costs or grossly inflated charges.
|
||||
_MAX_DELTA_KWH = Decimal("100")
|
||||
|
||||
# 每日固定费/税补在"本地午夜后多久"才结算入账。延后到 01:05 是为了让累计成本的
|
||||
# 整天阶跃落在新一天、且避开 01:00 整点(HA 长期统计的小时桶边界)。
|
||||
_SETTLEMENT_OFFSET = timedelta(hours=1, minutes=5)
|
||||
|
||||
# DSMR payload register keys (cumulative kWh, JSON string values).
|
||||
_KEY_D1 = "electricity_delivered_1" # delivered low-tariff (dal / _1)
|
||||
_KEY_D2 = "electricity_delivered_2" # delivered high-tariff (normal / _2)
|
||||
_KEY_R1 = "electricity_returned_1" # returned low-tariff
|
||||
_KEY_R2 = "electricity_returned_2" # returned high-tariff
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def floor_to_quarter(dt: datetime) -> datetime:
|
||||
"""Return *dt* floored to the nearest UTC quarter-hour boundary.
|
||||
|
||||
The result always has seconds=0 and microseconds=0, and minutes in
|
||||
{0, 15, 30, 45}. Timezone info is preserved if present.
|
||||
"""
|
||||
floored_minute = (dt.minute // _PERIOD_MINUTES) * _PERIOD_MINUTES
|
||||
return dt.replace(minute=floored_minute, second=0, microsecond=0)
|
||||
|
||||
|
||||
def _to_decimal(value: Any) -> Decimal:
|
||||
"""Convert *value* to Decimal via str() to avoid float binary rounding."""
|
||||
return Decimal(str(value))
|
||||
|
||||
|
||||
def _as_utc(dt: datetime) -> datetime:
|
||||
"""Attach UTC tzinfo to a naive datetime (SQLite read-back workaround)."""
|
||||
if dt.tzinfo is None:
|
||||
return dt.replace(tzinfo=UTC)
|
||||
return dt
|
||||
|
||||
|
||||
def _existing_period(session: Session, t0: datetime) -> EnergyCostPeriod | None:
|
||||
"""Return the EnergyCostPeriod row for period_start=t0, or None."""
|
||||
return session.execute(
|
||||
select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == t0)
|
||||
).scalar_one_or_none()
|
||||
|
||||
|
||||
def _unique_electricity_meter_at(session: Session, boundary: datetime) -> Meter | None:
|
||||
"""Return the sole electricity meter covering *boundary*, if one exists.
|
||||
|
||||
Billing must treat overlapping meter epochs as a structural ambiguity rather
|
||||
than relying on ``meter_at``'s newest-started tie breaker. A cumulative
|
||||
delta is safe only when exactly one electricity meter covers each endpoint.
|
||||
"""
|
||||
candidates = session.execute(
|
||||
select(Meter).where(
|
||||
Meter.commodity == "electricity",
|
||||
Meter.started_at <= boundary,
|
||||
(Meter.ended_at.is_(None)) | (Meter.ended_at > boundary),
|
||||
)
|
||||
).scalars().all()
|
||||
return candidates[0] if len(candidates) == 1 else None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# register_at — boundary reading lookup (meter-aware)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def register_at(
|
||||
session: Session,
|
||||
boundary: datetime,
|
||||
meter: Meter,
|
||||
*,
|
||||
meter_source_id: int | None = None,
|
||||
) -> dict[str, Decimal] | None:
|
||||
"""Return the four cumulative kWh register values at *boundary*, within *meter*'s window.
|
||||
|
||||
Queries the most recent ``DsmrReading`` with:
|
||||
``recorded_at ≤ boundary``
|
||||
AND ``recorded_at ≥ meter.started_at``
|
||||
AND (``meter.ended_at IS NULL`` OR ``recorded_at < meter.ended_at``)
|
||||
|
||||
The meter window constraint (half-open ``[started_at, ended_at)``) ensures
|
||||
that readings from a previous meter epoch are never used to anchor a new
|
||||
meter's computation. Without this guard, the final reading of the old meter
|
||||
would be visible at the start of the new meter's epoch and produce a
|
||||
cross-meter delta, defeating the isolation guarantee.
|
||||
|
||||
Extracts the four energy registers from ``payload``:
|
||||
|
||||
d1 — electricity_delivered_1 (delivered low-tariff / dal)
|
||||
d2 — electricity_delivered_2 (delivered high-tariff / normal)
|
||||
r1 — electricity_returned_1 (returned low-tariff)
|
||||
r2 — electricity_returned_2 (returned high-tariff)
|
||||
|
||||
Returns
|
||||
-------
|
||||
dict[str, Decimal] with keys ``d1``, ``d2``, ``r1``, ``r2``, or ``None``
|
||||
when:
|
||||
- No ``DsmrReading`` row exists with ``recorded_at ≤ boundary`` within
|
||||
*meter*'s epoch window.
|
||||
- The most recent such reading is older than ``_READING_MAX_STALENESS``
|
||||
relative to *boundary* (freshness guard).
|
||||
- Any of the four register keys is absent from the payload.
|
||||
- Any of the four register values is ``None`` (null in JSON).
|
||||
|
||||
SQLite naive datetime note
|
||||
--------------------------
|
||||
``recorded_at`` is stored as a naive UTC datetime in SQLite. Comparisons
|
||||
against *boundary* (always tz-aware UTC) use ``_as_utc()`` for the
|
||||
freshness check. The SQL ``WHERE`` clause comparisons work correctly
|
||||
because SQLAlchemy's SQLite dialect strips tzinfo when binding parameters
|
||||
(leaving the wall-clock UTC value unchanged), consistent with the storage
|
||||
format.
|
||||
"""
|
||||
# Build the meter-window constraints: [started_at, ended_at).
|
||||
meter_lower = meter.started_at # DsmrReading.recorded_at >= meter.started_at
|
||||
meter_upper = meter.ended_at # DsmrReading.recorded_at < meter.ended_at (if set)
|
||||
|
||||
stmt = (
|
||||
select(DsmrReading)
|
||||
.where(
|
||||
DsmrReading.recorded_at <= boundary,
|
||||
DsmrReading.recorded_at >= meter_lower,
|
||||
)
|
||||
.order_by(DsmrReading.recorded_at.desc())
|
||||
.limit(1)
|
||||
)
|
||||
# Apply the upper bound only when the meter is closed (ended_at is not None).
|
||||
if meter_upper is not None:
|
||||
stmt = stmt.where(DsmrReading.recorded_at < meter_upper)
|
||||
if meter_source_id is not None:
|
||||
stmt = stmt.where(DsmrReading.meter_source_id == meter_source_id)
|
||||
|
||||
row: DsmrReading | None = session.execute(stmt).scalar_one_or_none()
|
||||
|
||||
if row is None:
|
||||
return None
|
||||
|
||||
# Freshness guard: reject readings that are too old relative to *boundary*.
|
||||
# ``recorded_at`` is stored as a naive UTC datetime in SQLite; attach UTC
|
||||
# tzinfo before comparing with *boundary* (which is always tz-aware UTC) to
|
||||
# avoid an "offset-naive vs offset-aware" TypeError.
|
||||
if _as_utc(row.recorded_at) < _as_utc(boundary) - _READING_MAX_STALENESS:
|
||||
return None
|
||||
|
||||
payload = row.payload or {}
|
||||
try:
|
||||
d1_raw = payload[_KEY_D1]
|
||||
d2_raw = payload[_KEY_D2]
|
||||
r1_raw = payload[_KEY_R1]
|
||||
r2_raw = payload[_KEY_R2]
|
||||
except KeyError:
|
||||
return None
|
||||
|
||||
if any(v is None for v in (d1_raw, d2_raw, r1_raw, r2_raw)):
|
||||
return None
|
||||
|
||||
return {
|
||||
"d1": _to_decimal(d1_raw),
|
||||
"d2": _to_decimal(d2_raw),
|
||||
"r1": _to_decimal(r1_raw),
|
||||
"r2": _to_decimal(r2_raw),
|
||||
}
|
||||
|
||||
|
||||
def _binding_at(
|
||||
session: Session, boundary: datetime, meter: Meter
|
||||
) -> tuple[MeterSourceBinding, int] | None:
|
||||
"""Resolve the sole DSMR binding for *meter* at one period boundary.
|
||||
|
||||
Costing must not infer a cumulative domain from whichever reading happens
|
||||
to be latest. A binding anchors both the physical meter epoch and its
|
||||
source stream. Any missing or overlapping binding is therefore
|
||||
deliberately unresolvable.
|
||||
"""
|
||||
candidates = session.execute(
|
||||
select(MeterSourceBinding, MeterSourceChannel.source_id)
|
||||
.join(MeterSourceChannel, MeterSourceChannel.id == MeterSourceBinding.channel_id)
|
||||
.join(MeterSource, MeterSource.id == MeterSourceChannel.source_id)
|
||||
.where(
|
||||
MeterSourceBinding.meter_id == meter.id,
|
||||
MeterSourceBinding.started_at <= boundary,
|
||||
(MeterSourceBinding.ended_at.is_(None)) | (MeterSourceBinding.ended_at > boundary),
|
||||
MeterSource.kind == "dsmr_mqtt",
|
||||
)
|
||||
).all()
|
||||
if len(candidates) != 1:
|
||||
return None
|
||||
binding, source_id = candidates[0]
|
||||
return binding, source_id
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# compute_period — single 15-minute period
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def compute_period(session: Session, t0: datetime, *, overwrite: bool = False) -> bool:
|
||||
"""Compute and upsert the billing record for the period ``[t0, t0+15min)``.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
session:
|
||||
Active SQLAlchemy session. Caller is responsible for committing.
|
||||
t0:
|
||||
UTC start of the 15-minute period. **Must** lie on a quarter-hour
|
||||
grid boundary (minutes ∈ {0, 15, 30, 45}, seconds=0, microseconds=0).
|
||||
overwrite:
|
||||
If ``True``, overwrite an existing *successful* row (i.e. re-compute
|
||||
even when a non-degraded record already exists). The normal tick path
|
||||
always passes ``False``; only ``recompute_range`` passes ``True``.
|
||||
|
||||
Returns
|
||||
-------
|
||||
bool
|
||||
``True`` if a record was written (inserted or updated), ``False`` if
|
||||
the period was skipped (missing contract or missing Tibber price).
|
||||
|
||||
Side-effects
|
||||
------------
|
||||
- Inserts or updates an ``EnergyCostPeriod`` row keyed on ``period_start=t0``.
|
||||
- If no unique meter covers t0: inserts/updates a degraded row with
|
||||
``meter_id=None``.
|
||||
- If the period spans a meter boundary (m0.id != m1.id):
|
||||
inserts/updates a degraded row with ``meter_id=m0.id`` (D5 decision).
|
||||
- If readings are missing at either boundary within the meter window:
|
||||
inserts/updates a degraded row with ``meter_id=m0.id``.
|
||||
- If any delta is negative or exceeds ``_MAX_DELTA_KWH`` (D6 sanity guard):
|
||||
inserts/updates a degraded row with ``meter_id=m0.id``.
|
||||
- If the active contract version is missing: **skips** (returns False, no write).
|
||||
- If the Tibber price is missing (TibberPriceNotFoundError): **skips**
|
||||
(returns False, no write).
|
||||
"""
|
||||
t1 = t0 + timedelta(minutes=_PERIOD_MINUTES)
|
||||
now = datetime.now(UTC)
|
||||
|
||||
# Immutability guard: skip if a successful record already exists and we
|
||||
# are not in overwrite mode.
|
||||
existing = _existing_period(session, t0)
|
||||
if existing is not None and not existing.degraded and not overwrite:
|
||||
return False
|
||||
|
||||
# --- Meter determination (structural prerequisite, checked before contract) ---
|
||||
#
|
||||
# A missing or cross-boundary meter is a structural problem: we cannot trust
|
||||
# the delta at all, so we write a degraded row immediately. This is different
|
||||
# from the contract skip (transient, no write): the degraded row ensures the
|
||||
# period appears in the history and can be revisited once the meter timeline
|
||||
# is corrected and a recompute_range is triggered.
|
||||
#
|
||||
# Ordering rationale:
|
||||
# 1. No unique meter (m0 is None) → degraded(meter_id=None): no unambiguous epoch for t0.
|
||||
# 2. Cross-meter boundary (m0.id != m1.id) → degraded(meter_id=m0.id): D5.
|
||||
# 3. (Single meter, proceed) → contract check → readings → delta guard → price.
|
||||
#
|
||||
# We place meter before contract so that "cross-table period" is always
|
||||
# marked degraded regardless of contract state. If we checked contract
|
||||
# first, a missing-contract skip would silently discard the cross-table
|
||||
# evidence; once a contract is added and recompute runs, the engine would
|
||||
# incorrectly use cross-table reads.
|
||||
m0 = _unique_electricity_meter_at(session, t0)
|
||||
m1 = _unique_electricity_meter_at(session, t1)
|
||||
|
||||
if m0 is None:
|
||||
# No unambiguous meter epoch covers t0 — degraded with no attribution.
|
||||
logger.debug(
|
||||
"compute_period(%s): no unique active meter at t0 — writing degraded (meter_id=None).",
|
||||
t0.isoformat(),
|
||||
)
|
||||
_upsert_degraded(session, t0, now, existing, meter_id=None)
|
||||
return True
|
||||
|
||||
if m1 is None or m0.id != m1.id:
|
||||
# Period spans a meter boundary or t1 has no meter. Degrade with m0's id
|
||||
# (t0's meter attribution): the period's start belongs to m0's epoch.
|
||||
logger.debug(
|
||||
"compute_period(%s): period crosses meter boundary "
|
||||
"(m0.id=%s, m1.id=%s) — writing degraded.",
|
||||
t0.isoformat(),
|
||||
m0.id,
|
||||
m1.id if m1 is not None else None,
|
||||
)
|
||||
_upsert_degraded(session, t0, now, existing, meter_id=m0.id)
|
||||
return True
|
||||
|
||||
# Both endpoints must resolve to the same binding and source before a
|
||||
# cumulative subtraction is permitted. This is checked before contract
|
||||
# lookup so structural inconsistencies remain visible as degraded rows.
|
||||
bound0 = _binding_at(session, t0, m0)
|
||||
bound1 = _binding_at(session, t1, m1)
|
||||
if bound0 is None or bound1 is None or bound0[0].id != bound1[0].id or bound0[1] != bound1[1]:
|
||||
_upsert_degraded(session, t0, now, existing, meter_id=m0.id)
|
||||
return True
|
||||
binding, meter_source_id = bound0
|
||||
|
||||
# --- Active contract version at t0 ---
|
||||
# If there is no active contract covering t0, skip the period entirely.
|
||||
# We do not write a degraded row — there is no meaningful state to recover
|
||||
# without a contract (we would not know which strategy to apply once data
|
||||
# arrives). The period can be recovered via an explicit recompute_range once
|
||||
# a contract is configured and activated.
|
||||
version = active_contract_version_at(session, t0)
|
||||
if version is None:
|
||||
logger.debug("compute_period(%s): no active contract version — skipping.", t0.isoformat())
|
||||
return False
|
||||
|
||||
# --- Boundary readings within m0's meter window ---
|
||||
start_regs = register_at(session, t0, m0, meter_source_id=meter_source_id)
|
||||
end_regs = register_at(session, t1, m0, meter_source_id=meter_source_id)
|
||||
|
||||
if start_regs is None or end_regs is None:
|
||||
# Missing readings within the meter window → degraded with m0 attribution.
|
||||
_upsert_degraded(session, t0, now, existing, meter_id=m0.id)
|
||||
return True # a record was written (degraded)
|
||||
|
||||
# --- Compute deltas (end − start) ---
|
||||
deltas = PeriodDeltas(
|
||||
d1=end_regs["d1"] - start_regs["d1"],
|
||||
d2=end_regs["d2"] - start_regs["d2"],
|
||||
r1=end_regs["r1"] - start_regs["r1"],
|
||||
r2=end_regs["r2"] - start_regs["r2"],
|
||||
)
|
||||
|
||||
# --- Delta sanity guard (D6) ---
|
||||
# Any negative delta indicates a meter reset, DSMR rollover, or data error.
|
||||
# Any delta exceeding _MAX_DELTA_KWH (100 kWh per 15 min = 400 kW average)
|
||||
# is implausible for residential use and indicates an anomaly.
|
||||
# Both cases produce a degraded row so no negative or grossly inflated cost
|
||||
# is ever written to the billing record.
|
||||
all_deltas = (deltas.d1, deltas.d2, deltas.r1, deltas.r2)
|
||||
if any(d < Decimal("0") for d in all_deltas) or any(d > _MAX_DELTA_KWH for d in all_deltas):
|
||||
logger.debug(
|
||||
"compute_period(%s): delta sanity guard triggered "
|
||||
"(d1=%s, d2=%s, r1=%s, r2=%s) — writing degraded.",
|
||||
t0.isoformat(),
|
||||
deltas.d1,
|
||||
deltas.d2,
|
||||
deltas.r1,
|
||||
deltas.r2,
|
||||
)
|
||||
_upsert_degraded(session, t0, now, existing, meter_id=m0.id)
|
||||
return True
|
||||
|
||||
# --- Price strategy ---
|
||||
strategy = get_strategy(version.contract.kind)
|
||||
try:
|
||||
result = strategy(deltas, t0, version.values, session)
|
||||
except TibberPriceNotFoundError:
|
||||
# Missing Tibber price → skip the period; it will be retried once the
|
||||
# price arrives (e.g. after the next Tibber refresh job runs).
|
||||
logger.debug("compute_period(%s): no Tibber price found — skipping.", t0.isoformat())
|
||||
return False
|
||||
|
||||
# --- Upsert the billing record ---
|
||||
import_cost: Decimal = result["import_cost"]
|
||||
export_revenue: Decimal = result["export_revenue"]
|
||||
net_cost: Decimal = result["net_cost"]
|
||||
pricing: dict = result["pricing"]
|
||||
|
||||
if existing is not None:
|
||||
# Update in-place (overwrite=True or previous record was degraded).
|
||||
existing.d1_kwh = float(deltas.d1)
|
||||
existing.d2_kwh = float(deltas.d2)
|
||||
existing.r1_kwh = float(deltas.r1)
|
||||
existing.r2_kwh = float(deltas.r2)
|
||||
existing.import_cost = float(import_cost)
|
||||
existing.export_revenue = float(export_revenue)
|
||||
existing.net_cost = float(net_cost)
|
||||
existing.currency = version.contract.currency
|
||||
existing.pricing = pricing
|
||||
existing.contract_version_id = version.id
|
||||
existing.meter_id = m0.id
|
||||
existing.source_binding_id = binding.id
|
||||
existing.degraded = False
|
||||
existing.computed_at = now
|
||||
else:
|
||||
period = EnergyCostPeriod(
|
||||
period_start=t0,
|
||||
d1_kwh=float(deltas.d1),
|
||||
d2_kwh=float(deltas.d2),
|
||||
r1_kwh=float(deltas.r1),
|
||||
r2_kwh=float(deltas.r2),
|
||||
import_cost=float(import_cost),
|
||||
export_revenue=float(export_revenue),
|
||||
net_cost=float(net_cost),
|
||||
currency=version.contract.currency,
|
||||
pricing=pricing,
|
||||
contract_version_id=version.id,
|
||||
meter_id=m0.id,
|
||||
source_binding_id=binding.id,
|
||||
degraded=False,
|
||||
computed_at=now,
|
||||
)
|
||||
session.add(period)
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def _upsert_degraded(
|
||||
session: Session,
|
||||
t0: datetime,
|
||||
now: datetime,
|
||||
existing: EnergyCostPeriod | None,
|
||||
*,
|
||||
meter_id: int | None,
|
||||
) -> None:
|
||||
"""Insert or update a degraded placeholder for period *t0*.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
session:
|
||||
Active SQLAlchemy session.
|
||||
t0:
|
||||
UTC start of the 15-minute period.
|
||||
now:
|
||||
Current UTC timestamp for the ``computed_at`` field.
|
||||
existing:
|
||||
The existing ``EnergyCostPeriod`` row for this period, or ``None``.
|
||||
meter_id:
|
||||
The meter ID to attribute this degraded period to, or ``None`` when
|
||||
no meter epoch covers the period (no-meter degraded case).
|
||||
|
||||
When *existing* is not None (row was previously written — either degraded
|
||||
or successful), the row is explicitly reset to the standard degraded state.
|
||||
This is required for the ``recompute_range`` (overwrite=True) path: if the
|
||||
row was previously a *successful* computation and the boundary readings have
|
||||
since disappeared, the stale non-zero costs must be cleared so the row
|
||||
accurately reflects the current "missing readings" state rather than
|
||||
masquerading as a valid result.
|
||||
|
||||
The ``meter_id`` is always updated to reflect the current meter attribution
|
||||
judgment (the result of ``meter_at`` at the time of recompute). This
|
||||
ensures that a retroactive ``started_at`` change + ``recompute_range`` will
|
||||
re-attribute historical degraded periods to the correct meter epoch.
|
||||
"""
|
||||
if existing is not None:
|
||||
# Explicitly reset to degraded state — identical field values to the
|
||||
# new-row path below. This covers the recompute-over-successful-row
|
||||
# case where old non-zero costs must not survive the downgrade.
|
||||
existing.d1_kwh = 0.0
|
||||
existing.d2_kwh = 0.0
|
||||
existing.r1_kwh = 0.0
|
||||
existing.r2_kwh = 0.0
|
||||
existing.import_cost = 0.0
|
||||
existing.export_revenue = 0.0
|
||||
existing.net_cost = 0.0
|
||||
existing.pricing = {}
|
||||
existing.contract_version_id = None
|
||||
existing.meter_id = meter_id
|
||||
existing.source_binding_id = None
|
||||
existing.degraded = True
|
||||
existing.computed_at = now
|
||||
else:
|
||||
period = EnergyCostPeriod(
|
||||
period_start=t0,
|
||||
d1_kwh=0.0,
|
||||
d2_kwh=0.0,
|
||||
r1_kwh=0.0,
|
||||
r2_kwh=0.0,
|
||||
import_cost=0.0,
|
||||
export_revenue=0.0,
|
||||
net_cost=0.0,
|
||||
currency="EUR", # placeholder; real currency known after contract lookup
|
||||
pricing={},
|
||||
contract_version_id=None,
|
||||
meter_id=meter_id,
|
||||
source_binding_id=None,
|
||||
degraded=True,
|
||||
computed_at=now,
|
||||
)
|
||||
session.add(period)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# compute_closed_periods — periodic tick
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def compute_closed_periods(session: Session) -> int:
|
||||
"""Find and compute all uncalculated closed 15-minute periods.
|
||||
|
||||
A period ``[t0, t1)`` is *closed* when ``t1 ≤ now``. This function:
|
||||
|
||||
1. Determines the lookback window: from ``now − LOOKBACK_DAYS`` to ``now``,
|
||||
floored to the nearest quarter-hour. This avoids an unbounded full
|
||||
historical scan on every tick while still covering the typical recovery
|
||||
window (short outages, missing contract, etc.). Periods older than
|
||||
``LOOKBACK_DAYS`` must be recovered via an explicit ``recompute_range``.
|
||||
2. Iterates over all quarter-hour boundaries in that window where
|
||||
``t1 ≤ now`` (i.e. the period has already closed).
|
||||
3. For each boundary, calls ``compute_period(overwrite=False)``, which:
|
||||
- Skips periods that already have a *successful* (non-degraded) record.
|
||||
- Retries periods that have a *degraded* record.
|
||||
- Writes degraded rows for periods with no meter or cross-meter boundaries.
|
||||
- Skips periods for which no active contract version exists or the
|
||||
Tibber price is unavailable (without writing a degraded row).
|
||||
|
||||
Returns
|
||||
-------
|
||||
int
|
||||
Number of periods for which a record was written (inserted or updated).
|
||||
Does not count skipped periods.
|
||||
"""
|
||||
now = datetime.now(UTC)
|
||||
# Current period boundary (the one whose t1 has not yet passed).
|
||||
current_t0 = floor_to_quarter(now)
|
||||
# Earliest boundary to consider.
|
||||
earliest_t0 = floor_to_quarter(now - timedelta(days=_LOOKBACK_DAYS))
|
||||
|
||||
written = 0
|
||||
t0 = earliest_t0
|
||||
while t0 < current_t0:
|
||||
t1 = t0 + timedelta(minutes=_PERIOD_MINUTES)
|
||||
if t1 <= now:
|
||||
try:
|
||||
did_write = compute_period(session, t0, overwrite=False)
|
||||
if did_write:
|
||||
written += 1
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"compute_closed_periods: unexpected error for t0=%s — continuing.",
|
||||
t0.isoformat(),
|
||||
)
|
||||
t0 += timedelta(minutes=_PERIOD_MINUTES)
|
||||
|
||||
if written:
|
||||
session.commit()
|
||||
logger.info("compute_closed_periods: wrote %d period(s).", written)
|
||||
return written
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# recompute_range — explicit full recompute
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def recompute_range(
|
||||
session: Session, start: datetime, end: datetime, *, commit: bool = True, strict: bool = False
|
||||
) -> int:
|
||||
"""Recompute (overwrite) all 15-minute periods in ``[start, end)``.
|
||||
|
||||
This is the *explicit opt-in* path for recovering from:
|
||||
- Periods where readings or prices arrived late.
|
||||
- Price corrections (new contract version retroactively applied).
|
||||
- Retroactive meter changes (``update_meter`` with new ``started_at``) —
|
||||
re-running this function will re-judge meter attribution and re-compute
|
||||
costs using the corrected epoch boundaries.
|
||||
- Any other reason to override the immutability guard.
|
||||
|
||||
The function iterates over every UTC quarter-hour boundary in
|
||||
``[floor(start), end)`` and calls ``compute_period(overwrite=True)``.
|
||||
Existing rows (including successful ones) are overwritten; their
|
||||
``meter_id`` fields will reflect the *current* ``meter_at`` judgment for
|
||||
each period's start timestamp, naturally re-attributing periods when meter
|
||||
``started_at`` values have been retroactively corrected.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
session:
|
||||
Active SQLAlchemy session.
|
||||
commit:
|
||||
When true (the default), commit after all periods have been processed.
|
||||
Callers composing this recompute with other writes may pass false and
|
||||
own the surrounding transaction themselves.
|
||||
strict:
|
||||
When true, propagate a failed period computation to the caller. This
|
||||
is for lifecycle transactions which must roll back their meter/binding
|
||||
mutation together with the cost recompute. The default remains
|
||||
best-effort for existing background and standalone callers.
|
||||
start:
|
||||
Inclusive start datetime (floored to the nearest quarter-hour internally).
|
||||
end:
|
||||
Exclusive end datetime.
|
||||
|
||||
Returns
|
||||
-------
|
||||
int
|
||||
Number of periods for which a record was written (inserted or updated).
|
||||
Periods skipped due to missing contract or missing Tibber price are
|
||||
*not* counted.
|
||||
"""
|
||||
t0 = floor_to_quarter(_as_utc(start))
|
||||
end_utc = _as_utc(end)
|
||||
now = datetime.now(UTC)
|
||||
|
||||
written = 0
|
||||
while t0 < end_utc:
|
||||
t1 = t0 + timedelta(minutes=_PERIOD_MINUTES)
|
||||
# Only recompute periods that have already closed (t1 ≤ now). Even
|
||||
# when the caller passes a future *end*, we must not write speculative
|
||||
# "future" rows: register_at would resolve to the latest historical
|
||||
# reading for both boundaries, producing a zero-delta fake-success row
|
||||
# that blocks the real computation once the period actually closes.
|
||||
if t1 <= now:
|
||||
try:
|
||||
did_write = compute_period(session, t0, overwrite=True)
|
||||
if did_write:
|
||||
written += 1
|
||||
except Exception:
|
||||
if strict:
|
||||
raise
|
||||
logger.exception(
|
||||
"recompute_range: unexpected error for t0=%s — continuing.",
|
||||
t0.isoformat(),
|
||||
)
|
||||
t0 += timedelta(minutes=_PERIOD_MINUTES)
|
||||
|
||||
if commit:
|
||||
session.commit()
|
||||
logger.info(
|
||||
"recompute_range(%s, %s): wrote %d period(s).",
|
||||
start.isoformat(),
|
||||
end.isoformat(),
|
||||
written,
|
||||
)
|
||||
return written
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# summarize — layer-2 aggregation (read-time, not stored)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def summarize(session: Session, start: datetime, end: datetime) -> dict[str, Any]:
|
||||
"""Aggregate billing for the interval ``[start, end)``.
|
||||
|
||||
Computes the total payable as:
|
||||
|
||||
total_payable = Σ(net_cost) -- metered electricity
|
||||
+ fixed_costs -- per-day standing charges, cross-version
|
||||
- credits -- per-day heffingskorting, cross-version
|
||||
|
||||
**Fixed-cost / credit counting — Principle C (symmetric begin/end)**:
|
||||
Both the local calendar day in which *start* falls and the local calendar
|
||||
day in which *end* falls are counted as full days. Fixed charges
|
||||
(network_fee, management_fee) and the energy-tax credit (heffingskorting)
|
||||
are assessed on a "service-is-active" basis — if the meter was online on a
|
||||
given calendar day, the full day's charge/credit applies, regardless of
|
||||
whether the window starts at midnight or mid-morning.
|
||||
|
||||
For each local calendar date D in the range
|
||||
``[local_date(start), min(local_date(end), today_local)]``:
|
||||
- D ≤ today_local (only elapsed / today days count as "whole days").
|
||||
- The contract version whose effective_from local-date ≤ D is used.
|
||||
- This is cross-version: if V1 is from June 1 and V2 from June 25,
|
||||
querying June 1–30 uses V1 for days 1-24 and V2 for day 25.
|
||||
|
||||
The end-day (last_counted) is included when the end's local midnight falls
|
||||
strictly before end_utc; combined with the always-counted start day this
|
||||
makes the begin/end handling symmetric. A short same-day window therefore
|
||||
counts its single local day. A window contributes 0 days only when the
|
||||
counted range is empty (first_counted > last_counted) — e.g. a window lying
|
||||
entirely in the future, since last_counted is capped at today_local.
|
||||
|
||||
Daily getters (``*_today``) use windows exactly aligned to local midnight,
|
||||
so their ``first_counted`` is always today — unaffected by this fix.
|
||||
|
||||
Days that have not yet started in local time (D > today_local) are
|
||||
never counted. Switching versions never resets the counter.
|
||||
|
||||
**Timezone note**: the ``days`` field in the returned dict still represents
|
||||
the window length in calendar days (total_seconds / 86400), for backward
|
||||
compatibility with existing API consumers. The fixed/credit calculation
|
||||
independently counts whole local days as described above.
|
||||
|
||||
All arithmetic uses Decimal; the returned dict contains Python floats for
|
||||
JSON-serialisation convenience.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
session:
|
||||
Active read-only SQLAlchemy session.
|
||||
start:
|
||||
Inclusive start of the summary interval (UTC or naive-UTC).
|
||||
end:
|
||||
Exclusive end of the summary interval (UTC or naive-UTC).
|
||||
|
||||
Returns
|
||||
-------
|
||||
dict with keys:
|
||||
|
||||
currency str ISO 4217 currency (from contract, or "EUR" fallback)
|
||||
metered_import float Σ import_cost from non-degraded periods (money)
|
||||
metered_export float Σ export_revenue from non-degraded periods (money)
|
||||
metered_net float Σ net_cost from non-degraded periods (money)
|
||||
metered_import_kwh float Σ (d1_kwh + d2_kwh) from non-degraded periods (energy)
|
||||
metered_export_kwh float Σ (r1_kwh + r2_kwh) from non-degraded periods (energy)
|
||||
fixed_costs float standing charges for elapsed whole local days
|
||||
credits float energy-tax credit for elapsed whole local days
|
||||
total_payable float metered_net + fixed_costs − credits
|
||||
period_count int number of non-degraded periods in range
|
||||
degraded_count int number of degraded periods in range
|
||||
days float interval length in days (total_seconds / 86400)
|
||||
"""
|
||||
from datetime import timedelta as _td, date as _date
|
||||
|
||||
start_utc = _as_utc(start)
|
||||
end_utc = _as_utc(end)
|
||||
|
||||
# --- Fetch all EnergyCostPeriod rows in [start, end) ---
|
||||
rows = (
|
||||
session.execute(
|
||||
select(EnergyCostPeriod).where(
|
||||
EnergyCostPeriod.period_start >= start_utc,
|
||||
EnergyCostPeriod.period_start < end_utc,
|
||||
)
|
||||
)
|
||||
.scalars()
|
||||
.all()
|
||||
)
|
||||
|
||||
good_rows = [r for r in rows if not r.degraded]
|
||||
degraded_rows = [r for r in rows if r.degraded]
|
||||
|
||||
# Σ monetary amounts (Decimal arithmetic).
|
||||
sum_import = sum((_to_decimal(r.import_cost) for r in good_rows), Decimal("0"))
|
||||
sum_export = sum((_to_decimal(r.export_revenue) for r in good_rows), Decimal("0"))
|
||||
sum_net = sum((_to_decimal(r.net_cost) for r in good_rows), Decimal("0"))
|
||||
|
||||
# Σ metered energy (kWh), summed across both tariff registers. Reuses the
|
||||
# already-fetched ``good_rows`` so no extra query is issued.
|
||||
sum_import_kwh = sum(
|
||||
(_to_decimal(r.d1_kwh) + _to_decimal(r.d2_kwh) for r in good_rows), Decimal("0")
|
||||
)
|
||||
sum_export_kwh = sum(
|
||||
(_to_decimal(r.r1_kwh) + _to_decimal(r.r2_kwh) for r in good_rows), Decimal("0")
|
||||
)
|
||||
|
||||
# --- Interval length in days (window, not elapsed — kept for API compat) ---
|
||||
total_seconds = (end_utc - start_utc).total_seconds()
|
||||
days = _to_decimal(str(total_seconds)) / _to_decimal("86400")
|
||||
|
||||
# --- Fixed costs and credits: Principle C cross-version whole-day counting ---
|
||||
_now_local = local_now()
|
||||
today_local: _date = _now_local.date()
|
||||
|
||||
# --- Compute [first_counted, last_counted] local date range (inclusive) ---
|
||||
#
|
||||
# Both the window-start day and the window-end day are counted as complete
|
||||
# local calendar days, regardless of whether the window starts/ends at midnight.
|
||||
#
|
||||
# Principle C (symmetric begin/end):
|
||||
# • first_counted = local calendar date of start_utc (start day always counted)
|
||||
# • last_counted = local calendar date of end_utc (end day counted if its
|
||||
# local midnight is strictly before end_utc)
|
||||
# • Both are then capped at today_local (only elapsed / today days count).
|
||||
#
|
||||
# Why symmetric? Fixed charges (network_fee, management_fee) and energy-tax credits
|
||||
# (heffingskorting) are assessed on a "service-is-active" basis, not on how many
|
||||
# hours the service was actually running within that calendar day. If the meter
|
||||
# anchor (started_at) falls at 09:18 on June 24, the full June 24 standing charge
|
||||
# and credit still apply because the service was online for that day.
|
||||
#
|
||||
# Previous asymmetric behaviour: a start_utc that was *later* than the local
|
||||
# midnight of local_start_date caused first_counted to be bumped to the *next*
|
||||
# day, silently dropping the anchor day's charges/credits. This was incorrect
|
||||
# for cumulative entities (import_cost_total / export_revenue_total) whose anchor
|
||||
# is often an above-midnight started_at. The end-side had always been symmetric
|
||||
# (counted if local midnight < end_utc), creating an inconsistency.
|
||||
#
|
||||
# Daily getters (window = [local today 00:00, local tomorrow 00:00)) are
|
||||
# unaffected: local_date(local_midnight_utc(today)) == today, so first_counted
|
||||
# is still today regardless of the fix.
|
||||
#
|
||||
# A short same-day window now counts its single local day (the start day is
|
||||
# always counted, symmetric with the end side). A window contributes 0 days
|
||||
# only when the counted range is empty (first_counted > last_counted) — e.g. a
|
||||
# window lying entirely in the future, where last_counted is capped at
|
||||
# today_local while first_counted is later.
|
||||
from app.services.timezone import local_midnight_utc as _lmu
|
||||
|
||||
# Start side: always count the local calendar day in which start_utc falls.
|
||||
first_counted: _date = local_date(start_utc)
|
||||
|
||||
local_end_date: _date = local_date(end_utc)
|
||||
# Is the midnight of local_end_date < end_utc? If yes, that day's midnight is in range.
|
||||
if _lmu(local_end_date) < end_utc:
|
||||
last_counted: _date = local_end_date
|
||||
else:
|
||||
last_counted = local_end_date - _td(days=1)
|
||||
|
||||
# Settlement cap: "today" is only counted once the local clock has passed the
|
||||
# settlement offset since midnight (01:05). This defers the daily standing-charge
|
||||
# step from 00:00 to 01:05, ensuring the cumulative-cost adiabatic jump lands
|
||||
# inside the new calendar day and avoids the HA 01:00 hourly-bucket boundary.
|
||||
_today_midnight_utc = _lmu(today_local)
|
||||
if _now_local >= _today_midnight_utc + _SETTLEMENT_OFFSET:
|
||||
settled_cap = today_local
|
||||
else:
|
||||
settled_cap = today_local - _td(days=1)
|
||||
last_counted = min(last_counted, settled_cap)
|
||||
|
||||
# If the range is empty (first_counted > last_counted), no days are counted.
|
||||
|
||||
# Fetch all versions of the active contract once (single DB call).
|
||||
versions = active_contract_versions(session)
|
||||
|
||||
fixed_dec = Decimal("0")
|
||||
credits_dec = Decimal("0")
|
||||
currency = "EUR"
|
||||
|
||||
if versions:
|
||||
# Currency comes from the contract regardless of day count.
|
||||
currency = versions[0].contract.currency
|
||||
|
||||
if versions and first_counted <= last_counted:
|
||||
# For each version segment, find its overlap with [first_counted, last_counted]
|
||||
# and accumulate whole days.
|
||||
version_segments: list[tuple[_date, _date | None, dict]] = []
|
||||
for v in versions:
|
||||
v_start_local = local_date(_as_utc(v.effective_from))
|
||||
v_end_local = (
|
||||
local_date(_as_utc(v.effective_to)) if v.effective_to is not None else None
|
||||
)
|
||||
version_segments.append((v_start_local, v_end_local, v.values or {}))
|
||||
|
||||
for v_start, v_end_excl, v_values in version_segments:
|
||||
# The version covers [v_start, v_end_excl) in local dates,
|
||||
# where v_end_excl=None means open-ended (no upper bound).
|
||||
# Intersection with [first_counted, last_counted]:
|
||||
seg_start = max(first_counted, v_start)
|
||||
if v_end_excl is not None:
|
||||
seg_end_excl = min(last_counted + _td(days=1), v_end_excl)
|
||||
else:
|
||||
seg_end_excl = last_counted + _td(days=1)
|
||||
|
||||
if seg_start >= seg_end_excl:
|
||||
continue # no overlap
|
||||
|
||||
n_days = (seg_end_excl - seg_start).days
|
||||
if n_days <= 0:
|
||||
continue
|
||||
|
||||
standing: dict = v_values.get("standing", {})
|
||||
creds: dict = v_values.get("credits", {})
|
||||
|
||||
network_fee = _to_decimal(standing.get("network_fee", 0))
|
||||
management_fee = _to_decimal(standing.get("management_fee", 0))
|
||||
heffingskorting = _to_decimal(creds.get("heffingskorting", 0))
|
||||
|
||||
# Standing charges: EUR/month → EUR/day (÷ 30) × n_days.
|
||||
fixed_dec += (network_fee + management_fee) / Decimal("30") * Decimal(str(n_days))
|
||||
# Energy-tax credit: EUR/year → EUR/day (÷ 365) × n_days.
|
||||
credits_dec += heffingskorting / Decimal("365") * Decimal(str(n_days))
|
||||
|
||||
total_payable = sum_net + fixed_dec - credits_dec
|
||||
|
||||
return {
|
||||
"currency": currency,
|
||||
"metered_import": float(sum_import),
|
||||
"metered_export": float(sum_export),
|
||||
"metered_net": float(sum_net),
|
||||
"metered_import_kwh": float(sum_import_kwh),
|
||||
"metered_export_kwh": float(sum_export_kwh),
|
||||
"fixed_costs": float(fixed_dec),
|
||||
"credits": float(credits_dec),
|
||||
"total_payable": float(total_payable),
|
||||
"period_count": len(good_rows),
|
||||
"degraded_count": len(degraded_rows),
|
||||
"days": float(days),
|
||||
}
|
||||
@@ -0,0 +1,590 @@
|
||||
"""HA MQTT Discovery service — build and publish discovery configs + state.
|
||||
|
||||
This module handles:
|
||||
|
||||
1. Building HA MQTT Discovery payloads for each ``ExposableEntity``.
|
||||
2. Publishing discovery configs (retained) for enabled entities.
|
||||
3. Clearing (empty-payload) discovery configs for disabled entities.
|
||||
4. Publishing current entity state values.
|
||||
5. A per-device helper for use in the Modbus poll loop.
|
||||
|
||||
Design
|
||||
------
|
||||
- **No-op when MQTT / discovery is not configured**: every public function
|
||||
checks ``ha_discovery_enabled`` and ``mqtt_manager.is_connected`` before
|
||||
doing any work. Callers never need to guard these conditions.
|
||||
- **Stable unique_id**: derived from device ``uuid`` + metric ``key``, never
|
||||
from mutable fields like ``friendly_name`` or auto-increment DB ids.
|
||||
- **Discovery topic format**: ``<discovery_prefix>/<component>/<node_id>/<object_id>/config``
|
||||
where ``node_id`` is the device uuid (slugified to be safe) and
|
||||
``object_id`` is a uuid-prefixed stable string. Uses ``ha_discovery_prefix``
|
||||
(default ``"homeassistant"``), which must stay under the HA discovery namespace.
|
||||
- **State topic format**: ``<state_prefix>/<component>/<node_id>/<object_id>/state``
|
||||
Uses ``ha_state_topic_prefix`` (default ``"home_automation"``), separate from
|
||||
the discovery namespace so HA discovery and state/availability topics live under
|
||||
different prefixes.
|
||||
- **Availability topic**: ``<state_prefix>/modbus/<node_id>/availability`` (shared
|
||||
across all entities of the same device; publishes "online"/"offline"). Also uses
|
||||
the state prefix, not the discovery prefix.
|
||||
- **best-effort**: functions catch all exceptions internally so that callers
|
||||
(e.g. the Modbus poll loop) never crash due to MQTT publish errors.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.integrations.expose import ExposableEntity, build_catalog
|
||||
from app.integrations.mqtt import mqtt_manager
|
||||
from app.services.config_page import build_runtime_settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _should_publish(settings: Any) -> bool:
|
||||
"""Return True only if MQTT and HA Discovery are both enabled and connected."""
|
||||
return bool(
|
||||
settings.mqtt_enabled
|
||||
and settings.ha_discovery_enabled
|
||||
and mqtt_manager.is_connected
|
||||
)
|
||||
|
||||
|
||||
def _node_id(device_uuid: str) -> str:
|
||||
"""Stable MQTT node_id derived from device uuid (replace hyphens for safety)."""
|
||||
return device_uuid.replace("-", "_")
|
||||
|
||||
|
||||
def _object_id(entity: ExposableEntity) -> str:
|
||||
"""Stable MQTT object_id derived from entity key (dots + hyphens → underscores)."""
|
||||
return entity.key.replace(".", "_").replace("-", "_")
|
||||
|
||||
|
||||
def _discovery_topic(entity: ExposableEntity, prefix: str) -> str:
|
||||
"""Build the HA Discovery config topic for *entity*.
|
||||
|
||||
Format: ``<prefix>/<component>/<node_id>/<object_id>/config``
|
||||
"""
|
||||
node = _node_id(entity.device.identifiers[1]) # device uuid
|
||||
obj = _object_id(entity)
|
||||
return f"{prefix}/{entity.component}/{node}/{obj}/config"
|
||||
|
||||
|
||||
def _state_topic(entity: ExposableEntity, prefix: str) -> str:
|
||||
"""Build the state publish topic for *entity*.
|
||||
|
||||
Format: ``<prefix>/<component>/<node_id>/<object_id>/state``
|
||||
"""
|
||||
node = _node_id(entity.device.identifiers[1])
|
||||
obj = _object_id(entity)
|
||||
return f"{prefix}/{entity.component}/{node}/{obj}/state"
|
||||
|
||||
|
||||
def _availability_topic(device_uuid: str, prefix: str) -> str:
|
||||
"""Shared availability topic for all entities of a device.
|
||||
|
||||
Format: ``<prefix>/modbus/<node_id>/availability``
|
||||
"""
|
||||
node = _node_id(device_uuid)
|
||||
return f"{prefix}/modbus/{node}/availability"
|
||||
|
||||
|
||||
def _availability_id(entity: ExposableEntity) -> str:
|
||||
"""Return the identity which owns this entity's liveness topic.
|
||||
|
||||
M8 meters deliberately retain their own UUID as HA node/unique identity,
|
||||
while their availability is supplied by a MeterSource UUID.
|
||||
"""
|
||||
return entity.device.availability_id or entity.device.identifiers[1]
|
||||
|
||||
|
||||
def _unique_id(entity: ExposableEntity) -> str:
|
||||
"""Stable unique_id — device uuid + metric key (never from mutable fields)."""
|
||||
device_uuid = entity.device.identifiers[1]
|
||||
# entity.key is already "modbus.<uuid>.<metric_key>" — use it as the seed
|
||||
return f"{device_uuid}_{entity.key.replace('.', '_')}"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public: build discovery payload
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def build_discovery_payload(
|
||||
entity: ExposableEntity,
|
||||
discovery_prefix: str,
|
||||
state_prefix: str | None = None,
|
||||
) -> tuple[str, dict[str, Any]]:
|
||||
"""Build the HA MQTT Discovery config topic and payload dict for *entity*.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
entity:
|
||||
An ``ExposableEntity`` (from ``build_catalog``).
|
||||
discovery_prefix:
|
||||
The HA Discovery config topic prefix (e.g. ``"homeassistant"``).
|
||||
This determines where HA looks for the discovery config topic.
|
||||
state_prefix:
|
||||
The prefix used for state and availability topics (e.g.
|
||||
``"home_automation"``). When omitted or ``None``, falls back to
|
||||
``discovery_prefix`` for backward compatibility (e.g. in tests that
|
||||
pass only one prefix).
|
||||
|
||||
Returns
|
||||
-------
|
||||
tuple[str, dict]
|
||||
``(topic, config_dict)`` — the config topic and the HA Discovery
|
||||
payload (not yet JSON-serialised).
|
||||
"""
|
||||
if state_prefix is None:
|
||||
state_prefix = discovery_prefix
|
||||
|
||||
avail_topic = _availability_topic(_availability_id(entity), state_prefix)
|
||||
state_t = _state_topic(entity, state_prefix)
|
||||
topic = _discovery_topic(entity, discovery_prefix)
|
||||
|
||||
config: dict[str, Any] = {
|
||||
"unique_id": _unique_id(entity),
|
||||
"name": entity.name,
|
||||
"state_topic": state_t,
|
||||
"device": {
|
||||
"identifiers": list(entity.device.identifiers),
|
||||
"name": entity.device.name,
|
||||
},
|
||||
}
|
||||
|
||||
# Only declare an availability topic for devices that actually publish an
|
||||
# online/offline heartbeat (e.g. Modbus, via its "online" binary_sensor).
|
||||
# Devices without a heartbeat (e.g. energy-cost) omit availability so HA
|
||||
# treats their entities as always-available — otherwise HA would mark them
|
||||
# ``unavailable`` forever even though their state is being published.
|
||||
if entity.device.provides_availability:
|
||||
config["availability"] = [{"topic": avail_topic}]
|
||||
config["availability_mode"] = "all"
|
||||
|
||||
# Component-specific fields
|
||||
if entity.component == "binary_sensor":
|
||||
# "online" binary sensor: payload ON/OFF
|
||||
config["payload_on"] = "ON"
|
||||
config["payload_off"] = "OFF"
|
||||
if entity.device_class:
|
||||
config["device_class"] = entity.device_class
|
||||
else:
|
||||
# sensor (and others)
|
||||
if entity.device_class:
|
||||
config["device_class"] = entity.device_class
|
||||
if entity.unit:
|
||||
config["unit_of_measurement"] = entity.unit
|
||||
if entity.state_class:
|
||||
config["state_class"] = entity.state_class
|
||||
|
||||
return topic, config
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public: publish discovery
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def publish_discovery(session: Session) -> None:
|
||||
"""Publish HA Discovery configs for all entities in the catalog.
|
||||
|
||||
- Enabled entities: publish retained JSON config payload.
|
||||
- Disabled entities: publish empty (b"") payload to clear any previously
|
||||
retained config from HA.
|
||||
|
||||
No-op if MQTT / discovery is not enabled or the client is not connected.
|
||||
All MQTT errors are caught internally — this function never raises.
|
||||
"""
|
||||
from app.config import get_settings
|
||||
settings = build_runtime_settings(session, get_settings())
|
||||
if not _should_publish(settings):
|
||||
logger.debug("publish_discovery: skipped (MQTT/Discovery not enabled or not connected)")
|
||||
return
|
||||
|
||||
discovery_prefix = settings.ha_discovery_prefix
|
||||
state_prefix = settings.ha_state_topic_prefix
|
||||
|
||||
try:
|
||||
catalog = build_catalog(session)
|
||||
except Exception:
|
||||
logger.exception("publish_discovery: failed to build catalog; aborting")
|
||||
return
|
||||
|
||||
# Meter UUIDs are intentionally identity-changing epochs. Discovery config
|
||||
# is retained, so clear only the precisely enumerable old M8 identities;
|
||||
# never wildcard a provider/topic and risk removing another source's card.
|
||||
try:
|
||||
stale_entities = _stale_m8_entities(session)
|
||||
except Exception:
|
||||
logger.exception("publish_discovery: unable to enumerate old M8 identities")
|
||||
stale_entities = []
|
||||
for old_entity in stale_entities:
|
||||
try:
|
||||
old_topic, _ = build_discovery_payload(old_entity, discovery_prefix, state_prefix)
|
||||
mqtt_manager.publish(old_topic, b"", retain=True)
|
||||
except Exception:
|
||||
logger.exception("publish_discovery: unable to clear old identity %r", old_entity.key)
|
||||
|
||||
for entry in catalog:
|
||||
entity = entry.entity
|
||||
try:
|
||||
topic, config = build_discovery_payload(entity, discovery_prefix, state_prefix)
|
||||
if entry.enabled:
|
||||
payload = json.dumps(config)
|
||||
mqtt_manager.publish(topic, payload, retain=True)
|
||||
logger.debug(
|
||||
"publish_discovery: published config for %r → %s", entity.key, topic
|
||||
)
|
||||
else:
|
||||
# Clear retained config for disabled entities.
|
||||
mqtt_manager.publish(topic, b"", retain=True)
|
||||
logger.debug(
|
||||
"publish_discovery: cleared config for disabled entity %r → %s",
|
||||
entity.key,
|
||||
topic,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"publish_discovery: error processing entity %r; continuing", entity.key
|
||||
)
|
||||
|
||||
|
||||
def _stale_m8_entities(session: Session) -> list[ExposableEntity]:
|
||||
"""Return synthetic discovery entries for superseded thermal identities.
|
||||
|
||||
This is deliberately a narrow, best-effort cleanup: ended Meter UUIDs and
|
||||
historically possible thermal combinations only; current identities are excluded.
|
||||
"""
|
||||
from app.integrations.expose import DeviceInfo
|
||||
from app.models.energy import Meter
|
||||
from sqlalchemy import select
|
||||
|
||||
meters = session.execute(select(Meter).where(
|
||||
Meter.commodity.in_(("electricity", "heating", "hot_water"))
|
||||
)).scalars().all()
|
||||
current = {meter.commodity: meter for meter in meters if meter.ended_at is None}
|
||||
old = [meter for meter in meters if meter.ended_at is not None]
|
||||
entities: list[ExposableEntity] = []
|
||||
for meter in old:
|
||||
info = DeviceInfo(identifiers=("meter", meter.uuid), name=meter.label)
|
||||
for suffix in ("total", "today"):
|
||||
entities.append(ExposableEntity(
|
||||
key=f"meter.{meter.uuid}.{suffix}", component="sensor", device=info,
|
||||
device_class=None, unit="", name="obsolete",
|
||||
))
|
||||
heatings = [meter for meter in meters if meter.commodity == "heating"]
|
||||
waters = [meter for meter in meters if meter.commodity == "hot_water"]
|
||||
current_identity = (
|
||||
".".join(sorted((current["heating"].uuid, current["hot_water"].uuid)))
|
||||
if current.get("heating") is not None and current.get("hot_water") is not None else None
|
||||
)
|
||||
for heating in heatings:
|
||||
for water in waters:
|
||||
if heating.ended_at is None and water.ended_at is None:
|
||||
continue
|
||||
# A thermal identity can only have been published when both Meter
|
||||
# epochs were current at the same instant. Do not form a Cartesian
|
||||
# product of historical records: that would tombstone identities
|
||||
# which have never existed in HA.
|
||||
heating_start, water_start = heating.started_at, water.started_at
|
||||
heating_end, water_end = heating.ended_at, water.ended_at
|
||||
if (heating_end is not None and water_start >= heating_end) or (
|
||||
water_end is not None and heating_start >= water_end
|
||||
):
|
||||
continue
|
||||
identity = ".".join(sorted((heating.uuid, water.uuid)))
|
||||
if identity == current_identity:
|
||||
continue
|
||||
info = DeviceInfo(identifiers=("thermal-cost", identity), name="obsolete")
|
||||
for metric in ("heating", "hot_water_heating", "water", "water_tax", "fixed", "all_in"):
|
||||
for suffix in ("total", "today"):
|
||||
entities.append(ExposableEntity(
|
||||
key=f"thermal_cost.{identity}.{metric}_{suffix}", component="sensor", device=info,
|
||||
device_class=None, unit="", name="obsolete",
|
||||
))
|
||||
return entities
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public: publish states
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def publish_states(session: Session) -> None:
|
||||
"""Publish current state values for all enabled entities.
|
||||
|
||||
Calls each entity's ``value_getter`` to obtain the current value and
|
||||
publishes it to the entity's state topic. Also publishes availability
|
||||
(online/offline) for each device.
|
||||
|
||||
No-op if MQTT / discovery is not enabled or the client is not connected.
|
||||
All errors are caught internally — this function never raises.
|
||||
"""
|
||||
from app.config import get_settings
|
||||
settings = build_runtime_settings(session, get_settings())
|
||||
if not _should_publish(settings):
|
||||
logger.debug("publish_states: skipped (MQTT/Discovery not enabled or not connected)")
|
||||
return
|
||||
|
||||
state_prefix = settings.ha_state_topic_prefix
|
||||
|
||||
try:
|
||||
catalog = build_catalog(session)
|
||||
except Exception:
|
||||
logger.exception("publish_states: failed to build catalog; aborting")
|
||||
return
|
||||
|
||||
# Collect enabled entries and publish
|
||||
for entry in catalog:
|
||||
if not entry.enabled:
|
||||
continue
|
||||
entity = entry.entity
|
||||
try:
|
||||
_publish_entity_state(entity, state_prefix, session)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"publish_states: error publishing state for %r; continuing", entity.key
|
||||
)
|
||||
|
||||
|
||||
def _publish_entity_state(
|
||||
entity: ExposableEntity,
|
||||
prefix: str,
|
||||
session: Session,
|
||||
) -> None:
|
||||
"""Publish one entity's current state value to its state topic.
|
||||
|
||||
Also publishes the availability topic for ``binary_sensor`` "online" entities.
|
||||
"""
|
||||
state_t = _state_topic(entity, prefix)
|
||||
# Source-backed entities can have a different liveness identity from their
|
||||
# HA device identity. Publish it before the state; a None value below is
|
||||
# intentionally not converted to a synthetic zero.
|
||||
if entity.device.provides_availability and entity.device.availability_getter is not None:
|
||||
try:
|
||||
available = bool(entity.device.availability_getter(session))
|
||||
mqtt_manager.publish(
|
||||
_availability_topic(_availability_id(entity), prefix),
|
||||
"online" if available else "offline",
|
||||
retain=False,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("availability_getter raised for entity %r", entity.key)
|
||||
|
||||
if entity.component == "binary_sensor" and "online" in entity.key:
|
||||
# The online sensor represents device availability.
|
||||
# Ask the value_getter for the current ON/OFF string, then also publish
|
||||
# the shared availability topic (online/offline).
|
||||
raw_value = None
|
||||
if entity.value_getter is not None:
|
||||
try:
|
||||
raw_value = entity.value_getter(session)
|
||||
except Exception:
|
||||
logger.exception("value_getter raised for online entity %r", entity.key)
|
||||
# Default to offline when no reading is available.
|
||||
online = (raw_value == "ON")
|
||||
avail_payload = "online" if online else "offline"
|
||||
avail_topic = _availability_topic(_availability_id(entity), prefix)
|
||||
mqtt_manager.publish(avail_topic, avail_payload, retain=False)
|
||||
# The state of the binary_sensor itself
|
||||
state_payload = "ON" if online else "OFF"
|
||||
mqtt_manager.publish(state_t, state_payload, retain=False)
|
||||
else:
|
||||
# Regular sensor: call value_getter(session) if available.
|
||||
value = None
|
||||
if entity.value_getter is not None:
|
||||
try:
|
||||
value = entity.value_getter(session)
|
||||
except Exception:
|
||||
logger.exception("value_getter raised for entity %r", entity.key)
|
||||
|
||||
if value is None:
|
||||
logger.debug("publish_states: no value for %r — skipping state publish", entity.key)
|
||||
return
|
||||
|
||||
mqtt_manager.publish(state_t, str(value), retain=False)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public: per-device state push (called by modbus_poll after each poll)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def publish_device_state(session: Session, device: Any) -> None:
|
||||
"""Publish state + availability for all enabled entities of *device*.
|
||||
|
||||
Called by ``modbus_poll.poll_device`` after a successful (or failed) poll.
|
||||
Publishes:
|
||||
- The current availability topic ("online" / "offline") for the device.
|
||||
- The state topic for each **enabled** entity of the device.
|
||||
|
||||
No-op if MQTT / discovery is not enabled or the client is not connected.
|
||||
All errors are caught internally — never raises.
|
||||
"""
|
||||
from app.config import get_settings
|
||||
settings = build_runtime_settings(session, get_settings())
|
||||
if not _should_publish(settings):
|
||||
return
|
||||
|
||||
state_prefix = settings.ha_state_topic_prefix
|
||||
device_uuid = device.uuid
|
||||
online = bool(device.last_poll_ok)
|
||||
|
||||
# Publish availability topic first.
|
||||
try:
|
||||
avail_topic = _availability_topic(device_uuid, state_prefix)
|
||||
avail_payload = "online" if online else "offline"
|
||||
mqtt_manager.publish(avail_topic, avail_payload, retain=False)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"publish_device_state: failed to publish availability for device uuid=%s", device_uuid
|
||||
)
|
||||
|
||||
if not online:
|
||||
# Device is offline — no point pushing stale state values.
|
||||
return
|
||||
|
||||
# Publish state for each enabled entity of this device.
|
||||
try:
|
||||
catalog = build_catalog(session)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"publish_device_state: failed to build catalog for device uuid=%s", device_uuid
|
||||
)
|
||||
return
|
||||
|
||||
for entry in catalog:
|
||||
if not entry.enabled:
|
||||
continue
|
||||
entity = entry.entity
|
||||
# Only process entities belonging to this device.
|
||||
if entity.device.identifiers[1] != device_uuid:
|
||||
continue
|
||||
# Skip the online binary_sensor itself (availability handled above).
|
||||
if entity.component == "binary_sensor" and "online" in entity.key:
|
||||
# Publish the binary_sensor state too.
|
||||
try:
|
||||
state_t = _state_topic(entity, state_prefix)
|
||||
mqtt_manager.publish(state_t, "ON", retain=False)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"publish_device_state: error publishing online sensor state for %r",
|
||||
entity.key,
|
||||
)
|
||||
continue
|
||||
|
||||
# Regular sensor — call value_getter(session).
|
||||
try:
|
||||
value = None
|
||||
if entity.value_getter is not None:
|
||||
value = entity.value_getter(session)
|
||||
if value is None:
|
||||
continue
|
||||
state_t = _state_topic(entity, state_prefix)
|
||||
mqtt_manager.publish(state_t, str(value), retain=False)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"publish_device_state: error publishing state for entity %r", entity.key
|
||||
)
|
||||
|
||||
|
||||
def clear_device_discovery(session: Session, device_uuid: str) -> None:
|
||||
"""Best-effort: send empty retained payloads to all HA Discovery config topics
|
||||
for the given device, effectively removing the device's entities from HA.
|
||||
|
||||
This must be called **before** the device rows are deleted from the DB, so
|
||||
that ``build_catalog`` can still enumerate the device's entities.
|
||||
|
||||
Behaviour
|
||||
---------
|
||||
- No-op if MQTT / HA Discovery is not enabled or the MQTT client is not
|
||||
connected.
|
||||
- All exceptions are caught internally; this function never raises.
|
||||
The caller proceeds with the DB deletion regardless of MQTT outcome.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
session:
|
||||
Active SQLAlchemy session (device must still exist in DB at call time).
|
||||
device_uuid:
|
||||
UUID string of the device being deleted.
|
||||
"""
|
||||
from app.config import get_settings
|
||||
settings = build_runtime_settings(session, get_settings())
|
||||
if not _should_publish(settings):
|
||||
logger.debug(
|
||||
"clear_device_discovery: skipped (MQTT/Discovery not enabled or not connected)"
|
||||
)
|
||||
return
|
||||
|
||||
discovery_prefix = settings.ha_discovery_prefix
|
||||
state_prefix = settings.ha_state_topic_prefix
|
||||
|
||||
try:
|
||||
catalog = build_catalog(session)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"clear_device_discovery: failed to build catalog for device uuid=%s; skipping HA cleanup",
|
||||
device_uuid,
|
||||
)
|
||||
return
|
||||
|
||||
cleared = 0
|
||||
for entry in catalog:
|
||||
entity = entry.entity
|
||||
# Only clear entities belonging to this device.
|
||||
if entity.device.identifiers[1] != device_uuid:
|
||||
continue
|
||||
try:
|
||||
topic, _config = build_discovery_payload(entity, discovery_prefix, state_prefix)
|
||||
mqtt_manager.publish(topic, b"", retain=True)
|
||||
cleared += 1
|
||||
logger.debug(
|
||||
"clear_device_discovery: cleared config topic for entity %r → %s",
|
||||
entity.key,
|
||||
topic,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"clear_device_discovery: error clearing entity %r; continuing", entity.key
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"clear_device_discovery: cleared %d HA discovery topic(s) for device uuid=%s",
|
||||
cleared,
|
||||
device_uuid,
|
||||
)
|
||||
|
||||
|
||||
def publish_device_offline(session: Session, device_uuid: str) -> None:
|
||||
"""Publish an "offline" availability payload for *device_uuid*.
|
||||
|
||||
Called by ``modbus_poll.poll_device`` when a poll fails, to immediately
|
||||
reflect the device going offline in HA (before the next full state sweep).
|
||||
|
||||
No-op if MQTT / discovery is not enabled or the client is not connected.
|
||||
Never raises.
|
||||
"""
|
||||
from app.config import get_settings
|
||||
settings = build_runtime_settings(session, get_settings())
|
||||
if not _should_publish(settings):
|
||||
return
|
||||
|
||||
try:
|
||||
state_prefix = settings.ha_state_topic_prefix
|
||||
avail_topic = _availability_topic(device_uuid, state_prefix)
|
||||
mqtt_manager.publish(avail_topic, "offline", retain=False)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"publish_device_offline: failed for device uuid=%s", device_uuid
|
||||
)
|
||||
@@ -1,6 +1,6 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import insert
|
||||
from sqlalchemy import delete, insert, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models.location import Location
|
||||
@@ -40,3 +40,58 @@ def record_location(session: Session, payload: LocationRecordRequest) -> None:
|
||||
)
|
||||
session.execute(stmt)
|
||||
session.commit()
|
||||
|
||||
|
||||
def update_location(
|
||||
session: Session,
|
||||
person: str,
|
||||
datetime_pk: str,
|
||||
*,
|
||||
latitude: float | None,
|
||||
longitude: float | None,
|
||||
altitude: float | None,
|
||||
) -> Location | None:
|
||||
"""Update non-PK fields of a single location row.
|
||||
|
||||
Returns the updated ORM object, or ``None`` if the PK does not exist.
|
||||
The caller must not pass PK fields — they are immutable.
|
||||
Only fields with a non-``None`` value are written; ``altitude`` being
|
||||
``None`` in the request means "leave unchanged", not "clear to NULL".
|
||||
"""
|
||||
row = session.execute(
|
||||
select(Location).where(
|
||||
Location.person == person,
|
||||
Location.datetime == datetime_pk,
|
||||
)
|
||||
).scalar_one_or_none()
|
||||
|
||||
if row is None:
|
||||
return None
|
||||
|
||||
if latitude is not None:
|
||||
row.latitude = latitude
|
||||
if longitude is not None:
|
||||
row.longitude = longitude
|
||||
if altitude is not None:
|
||||
row.altitude = altitude
|
||||
|
||||
session.commit()
|
||||
session.refresh(row)
|
||||
return row
|
||||
|
||||
|
||||
def delete_location(session: Session, person: str, datetime_pk: str) -> bool:
|
||||
"""Delete the single location row identified by its full composite PK.
|
||||
|
||||
Returns ``True`` if exactly one row was deleted, ``False`` if the PK did
|
||||
not exist (caller should raise 404). The DELETE is scoped to the exact PK
|
||||
— no batch/truncate path exists.
|
||||
"""
|
||||
result = session.execute(
|
||||
delete(Location).where(
|
||||
Location.person == person,
|
||||
Location.datetime == datetime_pk,
|
||||
)
|
||||
)
|
||||
session.commit()
|
||||
return result.rowcount == 1
|
||||
|
||||
@@ -0,0 +1,179 @@
|
||||
"""Exponential back-off throttle service for the login endpoint.
|
||||
|
||||
Design
|
||||
------
|
||||
Failures are tracked per (scope, key) in the ``auth_login_throttle`` table:
|
||||
|
||||
* scope ``'ip'`` → key is the client IP address
|
||||
* scope ``'user'`` → key is the username
|
||||
|
||||
On each login attempt the endpoint checks **both** keys and takes the **maximum**
|
||||
remaining wait (in seconds). A wait > 0 means the caller is still inside a
|
||||
back-off window and should receive 429.
|
||||
|
||||
Exponential formula
|
||||
-------------------
|
||||
The first ``N_FREE`` failures are free (no delay). After that::
|
||||
|
||||
wait = min(BACKOFF_CAP_S, BACKOFF_BASE_S * 2 ** (failures - N_FREE))
|
||||
|
||||
Constants (module-level so they can be tuned without touching the algorithm):
|
||||
|
||||
* ``N_FREE`` = 3 — failures before any delay starts
|
||||
* ``BACKOFF_BASE_S`` = 1 — base wait in seconds (1st delayed attempt → 2 s)
|
||||
* ``BACKOFF_CAP_S`` = 900 — hard cap (15 min)
|
||||
|
||||
Example delay schedule:
|
||||
failures 1, 2, 3 → 0 s (free)
|
||||
failures 4 → 2 s
|
||||
failures 5 → 4 s
|
||||
failures 6 → 8 s
|
||||
failures 7 → 16 s
|
||||
…
|
||||
failures 13+ → 900 s (capped)
|
||||
|
||||
Timezone notes
|
||||
--------------
|
||||
The model uses ``DateTime(timezone=True)``. SQLite stores timestamps without
|
||||
timezone info, so when rows are read back the column value may be timezone-naive.
|
||||
We normalise every stored value to UTC-aware via ``_as_utc`` before comparing
|
||||
against ``datetime.now(UTC)``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
from sqlalchemy import delete, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models.auth_throttle import LoginThrottle
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tunable constants
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
N_FREE: int = 3 # failures before any backoff starts
|
||||
BACKOFF_BASE_S: int = 1 # seconds — wait after the first delayed failure
|
||||
BACKOFF_CAP_S: int = 900 # seconds — maximum wait (15 min)
|
||||
|
||||
# Internal scope labels (keep in sync with model constraint)
|
||||
_SCOPE_IP = "ip"
|
||||
_SCOPE_USER = "user"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public API
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def check_and_get_wait(session: Session, *, ip: str, username: str) -> int:
|
||||
"""Return the number of seconds the caller must still wait, or 0 if allowed.
|
||||
|
||||
Checks both the IP row and the username row; returns the *larger* remaining
|
||||
wait so that both keys must have their back-off window expire before the
|
||||
caller is permitted to try again.
|
||||
"""
|
||||
now = datetime.now(UTC)
|
||||
wait_ip = _remaining_wait(session, scope=_SCOPE_IP, key=ip, now=now)
|
||||
wait_user = _remaining_wait(session, scope=_SCOPE_USER, key=username, now=now)
|
||||
return max(wait_ip, wait_user)
|
||||
|
||||
|
||||
def register_failure(session: Session, *, ip: str, username: str) -> None:
|
||||
"""Record one login failure for both the IP and the username keys.
|
||||
|
||||
Updates (or creates) the throttle row for each key, increments the failure
|
||||
counter, and recomputes ``next_allowed_at`` using the exponential formula.
|
||||
"""
|
||||
now = datetime.now(UTC)
|
||||
_upsert_failure(session, scope=_SCOPE_IP, key=ip, now=now)
|
||||
_upsert_failure(session, scope=_SCOPE_USER, key=username, now=now)
|
||||
session.commit()
|
||||
|
||||
|
||||
def clear(session: Session, *, ip: str, username: str) -> None:
|
||||
"""Delete the throttle rows for the given IP and username (successful login).
|
||||
|
||||
It is safe to call this even if no rows exist for the given keys.
|
||||
"""
|
||||
session.execute(
|
||||
delete(LoginThrottle).where(
|
||||
(LoginThrottle.scope == _SCOPE_IP) & (LoginThrottle.key == ip)
|
||||
| (LoginThrottle.scope == _SCOPE_USER) & (LoginThrottle.key == username)
|
||||
)
|
||||
)
|
||||
session.commit()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _remaining_wait(session: Session, *, scope: str, key: str, now: datetime) -> int:
|
||||
"""Return the number of seconds that ``key`` in ``scope`` must still wait.
|
||||
|
||||
Returns 0 if the key has no throttle row or if the back-off window has
|
||||
already expired.
|
||||
"""
|
||||
row = session.scalar(
|
||||
select(LoginThrottle)
|
||||
.where(LoginThrottle.scope == scope, LoginThrottle.key == key)
|
||||
.limit(1)
|
||||
)
|
||||
if row is None or row.next_allowed_at is None:
|
||||
return 0
|
||||
|
||||
next_allowed = _as_utc(row.next_allowed_at)
|
||||
if next_allowed <= now:
|
||||
return 0
|
||||
|
||||
remaining = (next_allowed - now).total_seconds()
|
||||
# Round up so callers are never told "0 seconds" while still blocked.
|
||||
return math.ceil(remaining)
|
||||
|
||||
|
||||
def _compute_next_allowed_at(failures: int, now: datetime) -> datetime | None:
|
||||
"""Compute the earliest time the key should be allowed to try again.
|
||||
|
||||
Returns ``None`` for the first ``N_FREE`` failures (no delay).
|
||||
"""
|
||||
if failures <= N_FREE:
|
||||
return None
|
||||
wait_s = min(BACKOFF_CAP_S, BACKOFF_BASE_S * (2 ** (failures - N_FREE)))
|
||||
return now + timedelta(seconds=wait_s)
|
||||
|
||||
|
||||
def _upsert_failure(session: Session, *, scope: str, key: str, now: datetime) -> None:
|
||||
"""Create or update the throttle row for (scope, key) with one more failure."""
|
||||
row = session.scalar(
|
||||
select(LoginThrottle)
|
||||
.where(LoginThrottle.scope == scope, LoginThrottle.key == key)
|
||||
.limit(1)
|
||||
)
|
||||
if row is None:
|
||||
new_failures = 1
|
||||
row = LoginThrottle(
|
||||
scope=scope,
|
||||
key=key,
|
||||
failures=new_failures,
|
||||
first_failed_at=now,
|
||||
last_failed_at=now,
|
||||
next_allowed_at=_compute_next_allowed_at(new_failures, now),
|
||||
)
|
||||
session.add(row)
|
||||
else:
|
||||
row.failures += 1
|
||||
row.last_failed_at = now
|
||||
row.next_allowed_at = _compute_next_allowed_at(row.failures, now)
|
||||
|
||||
|
||||
def _as_utc(value: datetime | None) -> datetime | None:
|
||||
"""Normalise a possibly-naive datetime to UTC-aware."""
|
||||
if value is None:
|
||||
return None
|
||||
if value.tzinfo is None:
|
||||
return value.replace(tzinfo=UTC)
|
||||
return value.astimezone(UTC)
|
||||
@@ -0,0 +1,298 @@
|
||||
"""Thermal (WarmteLink) 15-minute cost ledger.
|
||||
|
||||
This module deliberately does not share the electricity ledger: thermal has
|
||||
two independently-bound cumulative domains and Decimal database columns.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import UTC, date, datetime, time, timedelta
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models.energy import Meter, MeterCostPeriod
|
||||
from app.models.meter_source import MeterSource, MeterSourceBinding, MeterSourceChannel, WarmteLinkReading
|
||||
from app.services.contracts import active_contract_version_at, active_contract_versions
|
||||
from app.services.energy_cost import floor_to_quarter
|
||||
from app.services import timezone as timezone_service
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_PERIOD = timedelta(minutes=15)
|
||||
_FRESHNESS = timedelta(seconds=120)
|
||||
_LIMITS = {"heating": Decimal("0.1"), "hot_water": Decimal("1")}
|
||||
_ACCEPTED_QUALITIES = {"valid", "unverifiable"}
|
||||
_SETTLEMENT_TIME = time(1, 5)
|
||||
|
||||
|
||||
def _utc(value: datetime) -> datetime:
|
||||
return value.replace(tzinfo=UTC) if value.tzinfo is None else value.astimezone(UTC)
|
||||
|
||||
|
||||
def _decimal(value: Any) -> Decimal:
|
||||
return value if isinstance(value, Decimal) else Decimal(str(value))
|
||||
|
||||
|
||||
def _existing(session: Session, commodity: str, start: datetime) -> MeterCostPeriod | None:
|
||||
return session.execute(
|
||||
select(MeterCostPeriod).where(
|
||||
MeterCostPeriod.commodity == commodity, MeterCostPeriod.period_start == start
|
||||
)
|
||||
).scalar_one_or_none()
|
||||
|
||||
|
||||
def _meter_at(session: Session, commodity: str, instant: datetime) -> Meter | None:
|
||||
meters = session.execute(
|
||||
select(Meter).where(
|
||||
Meter.commodity == commodity,
|
||||
Meter.started_at <= instant,
|
||||
(Meter.ended_at.is_(None)) | (Meter.ended_at > instant),
|
||||
)
|
||||
).scalars().all()
|
||||
return meters[0] if len(meters) == 1 else None
|
||||
|
||||
|
||||
def _binding_at(
|
||||
session: Session, meter: Meter, instant: datetime
|
||||
) -> tuple[MeterSourceBinding, MeterSourceChannel] | None:
|
||||
rows = session.execute(
|
||||
select(MeterSourceBinding, MeterSourceChannel)
|
||||
.join(MeterSourceChannel, MeterSourceChannel.id == MeterSourceBinding.channel_id)
|
||||
.join(MeterSource, MeterSource.id == MeterSourceChannel.source_id)
|
||||
.where(
|
||||
MeterSourceBinding.meter_id == meter.id,
|
||||
MeterSourceBinding.started_at <= instant,
|
||||
(MeterSourceBinding.ended_at.is_(None)) | (MeterSourceBinding.ended_at > instant),
|
||||
MeterSource.kind == "warmtelink_serial",
|
||||
)
|
||||
).all()
|
||||
return rows[0] if len(rows) == 1 else None
|
||||
|
||||
|
||||
def _reading_at(
|
||||
session: Session,
|
||||
channel: MeterSourceChannel,
|
||||
meter: Meter,
|
||||
binding: MeterSourceBinding,
|
||||
target: datetime,
|
||||
) -> WarmteLinkReading | None:
|
||||
"""Choose the nearest accepted reading inside this cumulative domain.
|
||||
|
||||
Freshness alone is insufficient: a frame immediately before a meter or
|
||||
source hand-off belongs to a different cumulative register and must never
|
||||
be used as the other side of a delta.
|
||||
"""
|
||||
window_start, window_end = target - _FRESHNESS, target + _FRESHNESS
|
||||
rows = session.execute(
|
||||
select(WarmteLinkReading)
|
||||
.where(
|
||||
WarmteLinkReading.channel_id == channel.id,
|
||||
WarmteLinkReading.recorded_at >= window_start,
|
||||
WarmteLinkReading.recorded_at <= window_end,
|
||||
)
|
||||
).scalars().all()
|
||||
domain_start = max(_utc(meter.started_at), _utc(binding.started_at))
|
||||
domain_ends = (meter.ended_at, binding.ended_at)
|
||||
domain_end = min((_utc(value) for value in domain_ends if value is not None), default=None)
|
||||
accepted = [
|
||||
row for row in rows
|
||||
if row.quality in _ACCEPTED_QUALITIES
|
||||
and _utc(row.recorded_at) >= domain_start
|
||||
and (domain_end is None or _utc(row.recorded_at) < domain_end)
|
||||
]
|
||||
if not accepted:
|
||||
return None
|
||||
return min(
|
||||
accepted,
|
||||
key=lambda row: (abs((_utc(row.recorded_at) - target).total_seconds()), _utc(row.recorded_at)),
|
||||
)
|
||||
|
||||
|
||||
def _degrade(
|
||||
session: Session,
|
||||
commodity: str,
|
||||
start: datetime,
|
||||
end: datetime,
|
||||
existing: MeterCostPeriod | None,
|
||||
reason: str,
|
||||
meter_id: int | None = None,
|
||||
binding_id: int | None = None,
|
||||
) -> None:
|
||||
now = datetime.now(UTC)
|
||||
fields = dict(
|
||||
period_end=end, meter_id=meter_id, source_binding_id=binding_id,
|
||||
contract_version_id=None, quantity=Decimal("0"), cost=Decimal("0"), currency="EUR",
|
||||
cost_breakdown={}, pricing_snapshot={}, quality="invalid", degraded=True,
|
||||
degraded_reason=reason, updated_at=now,
|
||||
)
|
||||
if existing is None:
|
||||
session.add(MeterCostPeriod(commodity=commodity, period_start=start, created_at=now, **fields))
|
||||
else:
|
||||
for key, value in fields.items():
|
||||
setattr(existing, key, value)
|
||||
|
||||
|
||||
def compute_period(
|
||||
session: Session, commodity: str, period_start: datetime, *, overwrite: bool = False
|
||||
) -> bool:
|
||||
"""Compute one closed thermal period, recording every unsafe input as degraded."""
|
||||
if commodity not in _LIMITS:
|
||||
raise ValueError("commodity must be heating or hot_water")
|
||||
start = floor_to_quarter(_utc(period_start))
|
||||
end = start + _PERIOD
|
||||
existing = _existing(session, commodity, start)
|
||||
if existing is not None and not existing.degraded and not overwrite:
|
||||
return False
|
||||
|
||||
meter0, meter1 = _meter_at(session, commodity, start), _meter_at(session, commodity, end)
|
||||
if meter0 is None:
|
||||
_degrade(session, commodity, start, end, existing, "missing_or_ambiguous_meter")
|
||||
return True
|
||||
if meter1 is None or meter1.id != meter0.id:
|
||||
_degrade(session, commodity, start, end, existing, "cross_meter_epoch", meter0.id)
|
||||
return True
|
||||
bound0, bound1 = _binding_at(session, meter0, start), _binding_at(session, meter1, end)
|
||||
if bound0 is None or bound1 is None:
|
||||
_degrade(session, commodity, start, end, existing, "missing_or_ambiguous_binding", meter0.id)
|
||||
return True
|
||||
binding, channel = bound0
|
||||
if bound1[0].id != binding.id or bound1[1].id != channel.id:
|
||||
_degrade(session, commodity, start, end, existing, "cross_source_binding", meter0.id, binding.id)
|
||||
return True
|
||||
first = _reading_at(session, channel, meter0, binding, start)
|
||||
last = _reading_at(session, channel, meter0, binding, end)
|
||||
if first is None or last is None:
|
||||
_degrade(session, commodity, start, end, existing, "missing_stale_or_invalid_reading", meter0.id, binding.id)
|
||||
return True
|
||||
delta = _decimal(last.value) - _decimal(first.value)
|
||||
if delta < 0:
|
||||
_degrade(session, commodity, start, end, existing, "negative_delta", meter0.id, binding.id)
|
||||
return True
|
||||
if delta > _LIMITS[commodity]:
|
||||
_degrade(session, commodity, start, end, existing, "delta_limit_exceeded", meter0.id, binding.id)
|
||||
return True
|
||||
version = active_contract_version_at(session, start, scope="thermal")
|
||||
if version is None:
|
||||
_degrade(session, commodity, start, end, existing, "missing_contract", meter0.id, binding.id)
|
||||
return True
|
||||
|
||||
values = {key: _decimal(value) for key, value in version.values["variable"].items()}
|
||||
if commodity == "heating":
|
||||
breakdown = {"heating": delta * values["heating"]}
|
||||
else:
|
||||
breakdown = {
|
||||
key: delta * values[key]
|
||||
for key in ("hot_water_heating", "hot_water", "hot_water_tax")
|
||||
}
|
||||
cost = sum(breakdown.values(), Decimal("0"))
|
||||
now = datetime.now(UTC)
|
||||
fields = dict(
|
||||
period_end=end, meter_id=meter0.id, source_binding_id=binding.id,
|
||||
contract_version_id=version.id, quantity=delta, cost=cost, currency=version.contract.currency,
|
||||
cost_breakdown=breakdown, pricing_snapshot=dict(version.values),
|
||||
quality="valid" if first.quality == last.quality == "valid" else "unverifiable",
|
||||
degraded=False, degraded_reason=None, updated_at=now,
|
||||
)
|
||||
if existing is None:
|
||||
session.add(MeterCostPeriod(commodity=commodity, period_start=start, created_at=now, **fields))
|
||||
else:
|
||||
for key, value in fields.items():
|
||||
setattr(existing, key, value)
|
||||
return True
|
||||
|
||||
|
||||
def compute_closed_periods(session: Session, *, now: datetime | None = None) -> int:
|
||||
"""Retry incomplete thermal rows and fill recent closed periods without touching good rows."""
|
||||
now = _utc(now or datetime.now(UTC))
|
||||
first = floor_to_quarter(now - timedelta(days=7))
|
||||
written = 0
|
||||
cursor = first
|
||||
while cursor + _PERIOD <= now:
|
||||
for commodity in ("heating", "hot_water"):
|
||||
if compute_period(session, commodity, cursor):
|
||||
written += 1
|
||||
cursor += _PERIOD
|
||||
session.commit()
|
||||
return written
|
||||
|
||||
|
||||
def recompute_range(session: Session, start: datetime, end: datetime, *, commit: bool = True) -> int:
|
||||
"""Recompute a thermal range.
|
||||
|
||||
The historical service entry point remains self-committing for the scheduler
|
||||
and direct callers. HTTP callers pass ``commit=False`` so validation,
|
||||
recomputation, response statistics, and the single commit share one
|
||||
transaction owned by the route.
|
||||
"""
|
||||
cursor, end = floor_to_quarter(_utc(start)), _utc(end)
|
||||
now, written = datetime.now(UTC), 0
|
||||
while cursor < end:
|
||||
if cursor + _PERIOD <= now:
|
||||
for commodity in ("heating", "hot_water"):
|
||||
if compute_period(session, commodity, cursor, overwrite=True):
|
||||
written += 1
|
||||
cursor += _PERIOD
|
||||
if commit:
|
||||
session.commit()
|
||||
return written
|
||||
|
||||
|
||||
def _settled_end_date(now: datetime) -> date:
|
||||
local_now = timezone_service.to_local(now)
|
||||
return local_now.date() if local_now.timetz().replace(tzinfo=None) >= _SETTLEMENT_TIME else local_now.date() - timedelta(days=1)
|
||||
|
||||
|
||||
def summarize(session: Session, start: datetime, end: datetime, *, now: datetime | None = None) -> dict[str, Any]:
|
||||
"""Return thermal variable/fixed totals; standing is charged once per contract/day."""
|
||||
start, end = _utc(start), _utc(end)
|
||||
rows = session.execute(select(MeterCostPeriod).where(
|
||||
MeterCostPeriod.period_start >= start, MeterCostPeriod.period_start < end
|
||||
)).scalars().all()
|
||||
good = [row for row in rows if not row.degraded]
|
||||
variable = sum((_decimal(row.cost) for row in good), Decimal("0"))
|
||||
breakdown: dict[str, Decimal] = {key: Decimal("0") for key in (
|
||||
"heating", "hot_water_heating", "hot_water", "hot_water_tax")}
|
||||
for row in good:
|
||||
for key, value in row.cost_breakdown.items():
|
||||
breakdown[key] = breakdown.get(key, Decimal("0")) + _decimal(value)
|
||||
|
||||
# A summary is half-open. ``end`` at local midnight has no overlap with
|
||||
# that next local date, and an empty/reversed range owns no standing day.
|
||||
final_day = timezone_service.local_date(end - timedelta(microseconds=1))
|
||||
final_day = min(final_day, _settled_end_date(now or datetime.now(UTC)))
|
||||
day = timezone_service.local_date(start)
|
||||
fixed_breakdown: dict[str, Decimal] = {key: Decimal("0") for key in (
|
||||
"heating_network", "metering", "delivery_set", "hot_water_network", "other"
|
||||
)}
|
||||
versions = active_contract_versions(session, scope="thermal")
|
||||
while start < end and day <= final_day:
|
||||
day_start = datetime.combine(day, time.min, tzinfo=timezone_service.local_tz()).astimezone(UTC)
|
||||
next_day_start = datetime.combine(
|
||||
day + timedelta(days=1), time.min, tzinfo=timezone_service.local_tz()
|
||||
).astimezone(UTC)
|
||||
local_day_seconds = Decimal(str((next_day_start - day_start).total_seconds()))
|
||||
# A rate revision part-way through a local date is attributable only
|
||||
# to its effective interval. This preserves one contract-level daily
|
||||
# charge while correctly handling first-version and intra-day changes.
|
||||
for version in versions:
|
||||
segment_start = max(day_start, _utc(version.effective_from))
|
||||
version_end = _utc(version.effective_to) if version.effective_to is not None else next_day_start
|
||||
segment_end = min(next_day_start, version_end)
|
||||
if segment_start >= segment_end:
|
||||
continue
|
||||
values = version.values["standing"]
|
||||
fraction = Decimal(str((segment_end - segment_start).total_seconds())) / local_day_seconds
|
||||
for key in fixed_breakdown:
|
||||
fixed_breakdown[key] += _decimal(values.get(key, "0")) / Decimal("365") * fraction
|
||||
day += timedelta(days=1)
|
||||
fixed = sum(fixed_breakdown.values(), Decimal("0"))
|
||||
return {
|
||||
"currency": good[0].currency if good else "EUR", "variable_cost": variable,
|
||||
"fixed_cost": fixed, "fixed_breakdown": fixed_breakdown,
|
||||
"total_cost": variable + fixed, "breakdown": breakdown,
|
||||
"period_count": len(good), "degraded_count": len(rows) - len(good),
|
||||
}
|
||||
@@ -0,0 +1,529 @@
|
||||
"""Service layer for source/channel discovery and meter-source bindings.
|
||||
|
||||
All mutating functions receive a caller-owned :class:`~sqlalchemy.orm.Session`
|
||||
and never commit. This lets HTTP handlers compose source and meter changes in
|
||||
one transaction later without exposing any connection I/O here.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import or_, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.integrations.meter_sources import (
|
||||
SourceProfileError,
|
||||
get_source_profile,
|
||||
merge_source_config,
|
||||
validate_source_config,
|
||||
)
|
||||
from app.models.energy import Meter
|
||||
from app.models.meter_source import (
|
||||
MeterSource,
|
||||
MeterSourceBinding,
|
||||
MeterSourceChannel,
|
||||
half_open_intervals_overlap,
|
||||
)
|
||||
|
||||
|
||||
class MeterSourceError(ValueError):
|
||||
"""Base class for source-domain validation errors."""
|
||||
|
||||
|
||||
class SourceNotFoundError(MeterSourceError):
|
||||
"""Raised when the requested source does not exist."""
|
||||
|
||||
|
||||
class ChannelNotFoundError(MeterSourceError):
|
||||
"""Raised when the requested source channel does not exist."""
|
||||
|
||||
|
||||
class MeterNotFoundError(MeterSourceError):
|
||||
"""Raised when the requested meter does not exist."""
|
||||
|
||||
|
||||
class BindingNotFoundError(MeterSourceError):
|
||||
"""Raised when the requested binding does not exist."""
|
||||
|
||||
|
||||
class BindingValidationError(MeterSourceError):
|
||||
"""Raised for an incompatible unit, commodity, or invalid interval."""
|
||||
|
||||
|
||||
class BindingOverlapError(BindingValidationError):
|
||||
"""Raised when a meter or channel already has an overlapping binding."""
|
||||
|
||||
|
||||
class SourceDeleteRestrictedError(MeterSourceError):
|
||||
"""Raised when a source has retained channel, binding, or reading history."""
|
||||
|
||||
|
||||
COMMODITY_UNITS = {"electricity": "kWh", "heating": "GJ", "hot_water": "m³"}
|
||||
_UNSET = object()
|
||||
|
||||
|
||||
def _utc_now() -> datetime:
|
||||
return datetime.now(UTC)
|
||||
|
||||
|
||||
def _as_utc(value: datetime) -> datetime:
|
||||
return value.replace(tzinfo=UTC) if value.tzinfo is None else value.astimezone(UTC)
|
||||
|
||||
|
||||
def get_source(session: Session, source_id: int) -> MeterSource:
|
||||
source = session.get(MeterSource, source_id)
|
||||
if source is None:
|
||||
raise SourceNotFoundError(f"Meter source {source_id} was not found.")
|
||||
return source
|
||||
|
||||
|
||||
def list_sources(session: Session, *, kind: str | None = None) -> list[MeterSource]:
|
||||
statement = select(MeterSource).order_by(MeterSource.id)
|
||||
if kind is not None:
|
||||
get_source_profile(kind)
|
||||
statement = statement.where(MeterSource.kind == kind)
|
||||
return list(session.execute(statement).scalars())
|
||||
|
||||
|
||||
def create_source(
|
||||
session: Session,
|
||||
*,
|
||||
name: str,
|
||||
kind: str,
|
||||
config: dict[str, Any],
|
||||
enabled: bool = True,
|
||||
) -> MeterSource:
|
||||
"""Add a source after validating its complete kind-specific config."""
|
||||
now = _utc_now()
|
||||
source = MeterSource(
|
||||
name=name,
|
||||
kind=kind,
|
||||
enabled=enabled,
|
||||
config=validate_source_config(kind, config),
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
session.add(source)
|
||||
return source
|
||||
|
||||
|
||||
def update_source(
|
||||
session: Session,
|
||||
source_id: int,
|
||||
*,
|
||||
name: str | None = None,
|
||||
enabled: bool | None = None,
|
||||
config_patch: dict[str, Any] | None = None,
|
||||
) -> MeterSource:
|
||||
"""Update source metadata and/or merge a partial source config without commit."""
|
||||
source = get_source(session, source_id)
|
||||
if name is not None:
|
||||
source.name = name
|
||||
if enabled is not None:
|
||||
source.enabled = enabled
|
||||
if config_patch is not None:
|
||||
source.config = merge_source_config(source.kind, source.config, config_patch)
|
||||
source.updated_at = _utc_now()
|
||||
return source
|
||||
|
||||
|
||||
def delete_source(session: Session, source_id: int) -> None:
|
||||
"""Delete an entirely unused source; history is always retained instead."""
|
||||
source = get_source(session, source_id)
|
||||
has_channel = session.execute(
|
||||
select(MeterSourceChannel.id).where(MeterSourceChannel.source_id == source.id).limit(1)
|
||||
).scalar_one_or_none()
|
||||
has_binding = session.execute(
|
||||
select(MeterSourceBinding.id)
|
||||
.join(MeterSourceChannel)
|
||||
.where(MeterSourceChannel.source_id == source.id)
|
||||
.limit(1)
|
||||
).scalar_one_or_none()
|
||||
if has_channel is not None or has_binding is not None:
|
||||
raise SourceDeleteRestrictedError(
|
||||
f"Meter source {source_id} has dependent channels, bindings, or readings."
|
||||
)
|
||||
session.delete(source)
|
||||
|
||||
|
||||
def get_channel(session: Session, channel_id: int) -> MeterSourceChannel:
|
||||
channel = session.get(MeterSourceChannel, channel_id)
|
||||
if channel is None:
|
||||
raise ChannelNotFoundError(f"Meter source channel {channel_id} was not found.")
|
||||
return channel
|
||||
|
||||
|
||||
def upsert_discovered_channel(
|
||||
session: Session,
|
||||
*,
|
||||
source_id: int,
|
||||
channel_key: str,
|
||||
label: str,
|
||||
unit: str,
|
||||
suggested_commodity: str | None = None,
|
||||
device_type: str | None = None,
|
||||
fingerprint: str | None = None,
|
||||
latest_value: Any = None,
|
||||
latest_at: datetime | None = None,
|
||||
latest_quality: str | None = None,
|
||||
) -> MeterSourceChannel:
|
||||
"""Idempotently create or refresh a discovered channel's metadata.
|
||||
|
||||
``suggested_commodity`` remains metadata only; this function never creates
|
||||
a meter or a binding.
|
||||
"""
|
||||
source = get_source(session, source_id)
|
||||
if unit not in get_source_profile(source.kind).allowed_units:
|
||||
raise SourceProfileError(f"Unit {unit!r} is not allowed for source kind {source.kind!r}.")
|
||||
channel = session.execute(
|
||||
select(MeterSourceChannel).where(
|
||||
MeterSourceChannel.source_id == source.id,
|
||||
MeterSourceChannel.channel_key == channel_key,
|
||||
)
|
||||
).scalar_one_or_none()
|
||||
now = _utc_now()
|
||||
if channel is None:
|
||||
channel = MeterSourceChannel(
|
||||
source_id=source.id,
|
||||
channel_key=channel_key,
|
||||
label=label,
|
||||
unit=unit,
|
||||
suggested_commodity=suggested_commodity,
|
||||
device_type=device_type,
|
||||
fingerprint=fingerprint,
|
||||
latest_value=latest_value,
|
||||
latest_at=latest_at,
|
||||
latest_quality=latest_quality,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
session.add(channel)
|
||||
return channel
|
||||
|
||||
if channel.unit != unit:
|
||||
binding_id = session.execute(
|
||||
select(MeterSourceBinding.id)
|
||||
.where(MeterSourceBinding.channel_id == channel.id)
|
||||
.limit(1)
|
||||
).scalar_one_or_none()
|
||||
if binding_id is not None:
|
||||
raise BindingValidationError(
|
||||
f"Cannot change unit of bound channel {channel.id} from {channel.unit!r} to {unit!r}."
|
||||
)
|
||||
|
||||
channel.label = label
|
||||
channel.unit = unit
|
||||
channel.suggested_commodity = suggested_commodity
|
||||
channel.device_type = device_type
|
||||
channel.fingerprint = fingerprint
|
||||
channel.latest_value = latest_value
|
||||
channel.latest_at = latest_at
|
||||
channel.latest_quality = latest_quality
|
||||
channel.updated_at = now
|
||||
return channel
|
||||
|
||||
|
||||
def list_bindings(
|
||||
session: Session, *, meter_id: int | None = None, channel_id: int | None = None
|
||||
) -> list[MeterSourceBinding]:
|
||||
statement = select(MeterSourceBinding).order_by(MeterSourceBinding.started_at, MeterSourceBinding.id)
|
||||
if meter_id is not None:
|
||||
statement = statement.where(MeterSourceBinding.meter_id == meter_id)
|
||||
if channel_id is not None:
|
||||
statement = statement.where(MeterSourceBinding.channel_id == channel_id)
|
||||
return list(session.execute(statement).scalars())
|
||||
|
||||
|
||||
def _get_meter(session: Session, meter_id: int) -> Meter:
|
||||
meter = session.get(Meter, meter_id)
|
||||
if meter is None:
|
||||
raise MeterNotFoundError(f"Meter {meter_id} was not found.")
|
||||
return meter
|
||||
|
||||
|
||||
def _validate_binding(
|
||||
session: Session,
|
||||
*,
|
||||
meter_id: int,
|
||||
channel_id: int,
|
||||
started_at: datetime,
|
||||
ended_at: datetime | None,
|
||||
excluding_ids: set[int] | None = None,
|
||||
) -> None:
|
||||
meter = _get_meter(session, meter_id)
|
||||
channel = get_channel(session, channel_id)
|
||||
expected_unit = COMMODITY_UNITS.get(meter.commodity)
|
||||
if expected_unit is None:
|
||||
raise BindingValidationError(f"Commodity {meter.commodity!r} cannot be bound to a source channel.")
|
||||
if channel.unit != expected_unit:
|
||||
raise BindingValidationError(
|
||||
f"Meter commodity {meter.commodity!r} requires unit {expected_unit!r}, "
|
||||
f"but channel has {channel.unit!r}."
|
||||
)
|
||||
if ended_at is not None and _as_utc(ended_at) <= _as_utc(started_at):
|
||||
raise BindingValidationError("Binding ended_at must be strictly after started_at.")
|
||||
if _as_utc(started_at) < _as_utc(meter.started_at):
|
||||
raise BindingValidationError("Binding must not start before its meter epoch.")
|
||||
if meter.ended_at is None:
|
||||
if ended_at is not None:
|
||||
# A historical binding on an active epoch is valid, but it must be
|
||||
# wholly within that epoch (whose upper bound is open).
|
||||
pass
|
||||
else:
|
||||
meter_end = _as_utc(meter.ended_at)
|
||||
if ended_at is None or _as_utc(ended_at) > meter_end:
|
||||
raise BindingValidationError("Closed meter bindings must end within the meter epoch.")
|
||||
|
||||
excluded = excluding_ids or set()
|
||||
candidates = session.execute(
|
||||
select(MeterSourceBinding).where(
|
||||
or_(
|
||||
MeterSourceBinding.meter_id == meter_id,
|
||||
MeterSourceBinding.channel_id == channel_id,
|
||||
)
|
||||
)
|
||||
).scalars()
|
||||
for existing in candidates:
|
||||
if existing.id in excluded:
|
||||
continue
|
||||
if half_open_intervals_overlap(
|
||||
_as_utc(started_at),
|
||||
_as_utc(ended_at) if ended_at is not None else None,
|
||||
_as_utc(existing.started_at),
|
||||
_as_utc(existing.ended_at) if existing.ended_at is not None else None,
|
||||
):
|
||||
side = "meter" if existing.meter_id == meter_id else "channel"
|
||||
raise BindingOverlapError(f"Binding overlaps existing {side} binding {existing.id}.")
|
||||
|
||||
|
||||
def create_binding(
|
||||
session: Session,
|
||||
*,
|
||||
meter_id: int,
|
||||
channel_id: int,
|
||||
started_at: datetime,
|
||||
ended_at: datetime | None = None,
|
||||
) -> MeterSourceBinding:
|
||||
"""Create a compatible non-overlapping half-open source binding."""
|
||||
now = _utc_now()
|
||||
if _as_utc(started_at) > now or (ended_at is not None and _as_utc(ended_at) > now):
|
||||
raise BindingValidationError("Binding boundaries must not be in the future.")
|
||||
_validate_binding(
|
||||
session,
|
||||
meter_id=meter_id,
|
||||
channel_id=channel_id,
|
||||
started_at=started_at,
|
||||
ended_at=ended_at,
|
||||
)
|
||||
binding = MeterSourceBinding(
|
||||
meter_id=meter_id,
|
||||
channel_id=channel_id,
|
||||
started_at=started_at,
|
||||
ended_at=ended_at,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
session.add(binding)
|
||||
return binding
|
||||
|
||||
|
||||
def create_binding_for_meter_swap(
|
||||
session: Session,
|
||||
*,
|
||||
old_meter_id: int | None,
|
||||
new_meter_id: int,
|
||||
channel_id: int,
|
||||
started_at: datetime,
|
||||
) -> MeterSourceBinding:
|
||||
"""Create a binding during a physical meter swap, handing off one channel if safe.
|
||||
|
||||
A channel is transferable only when exactly one of its bindings covered the
|
||||
instant immediately before ``started_at`` and that binding belongs to the
|
||||
meter which this declaration just closed. All other occupied or ambiguous
|
||||
cases retain the normal fail-closed overlap behaviour.
|
||||
|
||||
This function deliberately does not commit. The caller must keep the meter
|
||||
declaration, binding handoff, and any billing recompute in one transaction.
|
||||
"""
|
||||
new_meter = _get_meter(session, new_meter_id)
|
||||
channel = get_channel(session, channel_id)
|
||||
expected_unit = COMMODITY_UNITS.get(new_meter.commodity)
|
||||
if expected_unit is None or channel.unit != expected_unit:
|
||||
raise BindingValidationError(
|
||||
f"Meter commodity {new_meter.commodity!r} requires unit {expected_unit!r}, "
|
||||
f"but channel has {channel.unit!r}."
|
||||
)
|
||||
|
||||
boundary = _as_utc(started_at)
|
||||
if _as_utc(new_meter.started_at) != boundary:
|
||||
raise BindingValidationError(
|
||||
"Meter-swap binding must start at the new meter's started_at boundary."
|
||||
)
|
||||
covering_bindings = [
|
||||
binding
|
||||
for binding in session.execute(
|
||||
select(MeterSourceBinding).where(MeterSourceBinding.channel_id == channel_id)
|
||||
).scalars()
|
||||
if _as_utc(binding.started_at) < boundary
|
||||
and (binding.ended_at is None or _as_utc(binding.ended_at) >= boundary)
|
||||
]
|
||||
|
||||
if not covering_bindings:
|
||||
return create_binding(
|
||||
session,
|
||||
meter_id=new_meter_id,
|
||||
channel_id=channel_id,
|
||||
started_at=started_at,
|
||||
)
|
||||
|
||||
if old_meter_id is None or len(covering_bindings) != 1:
|
||||
raise BindingOverlapError("Channel is occupied or has an ambiguous binding at meter swap.")
|
||||
|
||||
old_meter = _get_meter(session, old_meter_id)
|
||||
old_binding = covering_bindings[0]
|
||||
if (
|
||||
old_meter.commodity != new_meter.commodity
|
||||
or old_meter.ended_at is None
|
||||
or _as_utc(old_meter.ended_at) != boundary
|
||||
or old_binding.meter_id != old_meter.id
|
||||
):
|
||||
raise BindingOverlapError("Channel is occupied by a binding that cannot be handed off.")
|
||||
|
||||
update_binding(session, old_binding.id, ended_at=started_at)
|
||||
return create_binding(
|
||||
session,
|
||||
meter_id=new_meter_id,
|
||||
channel_id=channel_id,
|
||||
started_at=started_at,
|
||||
)
|
||||
|
||||
|
||||
def update_binding(
|
||||
session: Session,
|
||||
binding_id: int,
|
||||
*,
|
||||
meter_id: int | None = None,
|
||||
channel_id: int | None = None,
|
||||
started_at: datetime | None = None,
|
||||
ended_at: datetime | None | object = _UNSET,
|
||||
) -> MeterSourceBinding:
|
||||
"""Correct a binding while preserving half-open timeline constraints."""
|
||||
binding = session.get(MeterSourceBinding, binding_id)
|
||||
if binding is None:
|
||||
raise BindingNotFoundError(f"Meter source binding {binding_id} was not found.")
|
||||
new_meter_id = binding.meter_id if meter_id is None else meter_id
|
||||
new_channel_id = binding.channel_id if channel_id is None else channel_id
|
||||
new_started_at = binding.started_at if started_at is None else started_at
|
||||
new_ended_at = binding.ended_at if ended_at is _UNSET else ended_at
|
||||
now = _utc_now()
|
||||
if _as_utc(new_started_at) > now or (new_ended_at is not None and _as_utc(new_ended_at) > now):
|
||||
raise BindingValidationError("Binding boundaries must not be in the future.")
|
||||
_validate_binding(
|
||||
session,
|
||||
meter_id=new_meter_id,
|
||||
channel_id=new_channel_id,
|
||||
started_at=new_started_at,
|
||||
ended_at=new_ended_at,
|
||||
excluding_ids={binding.id},
|
||||
)
|
||||
binding.meter_id = new_meter_id
|
||||
binding.channel_id = new_channel_id
|
||||
binding.started_at = new_started_at
|
||||
binding.ended_at = new_ended_at
|
||||
binding.updated_at = _utc_now()
|
||||
return binding
|
||||
|
||||
|
||||
def close_binding(session: Session, binding_id: int, *, ended_at: datetime) -> MeterSourceBinding:
|
||||
"""Close an existing binding at its exclusive end boundary."""
|
||||
return update_binding(session, binding_id, ended_at=ended_at)
|
||||
|
||||
|
||||
def close_open_bindings_for_meter(session: Session, meter_id: int, *, ended_at: datetime) -> list[MeterSourceBinding]:
|
||||
"""Close every open binding on a meter at one shared boundary."""
|
||||
bindings = list(session.execute(
|
||||
select(MeterSourceBinding).where(
|
||||
MeterSourceBinding.meter_id == meter_id, MeterSourceBinding.ended_at.is_(None)
|
||||
)
|
||||
).scalars())
|
||||
for binding in bindings:
|
||||
update_binding(session, binding.id, ended_at=ended_at)
|
||||
return bindings
|
||||
|
||||
|
||||
def transfer_binding(
|
||||
session: Session, *, target_meter_id: int, from_binding_id: int, to_channel_id: int,
|
||||
effective_at: datetime,
|
||||
) -> tuple[MeterSourceBinding, MeterSourceBinding]:
|
||||
"""Atomically close a binding and open its replacement on the target meter."""
|
||||
source = session.get(MeterSourceBinding, from_binding_id)
|
||||
if source is None:
|
||||
raise BindingNotFoundError(f"Meter source binding {from_binding_id} was not found.")
|
||||
target = _get_meter(session, target_meter_id)
|
||||
old_meter = _get_meter(session, source.meter_id)
|
||||
effective_at = _as_utc(effective_at)
|
||||
now = _utc_now()
|
||||
if effective_at > now:
|
||||
raise BindingValidationError("Binding transfer effective_at must not be in the future.")
|
||||
if old_meter.commodity != target.commodity:
|
||||
raise BindingValidationError("Binding transfer meters must have the same commodity.")
|
||||
if source.ended_at is not None:
|
||||
raise BindingValidationError("Only an open binding can be transferred.")
|
||||
if old_meter.id == target.id:
|
||||
close_at = effective_at
|
||||
else:
|
||||
# Recovery is deliberately narrow: the source meter must be the one
|
||||
# and only most-recent closed predecessor in this commodity's timeline.
|
||||
# A manually closed meter may leave an intentional epoch gap before the
|
||||
# target is declared, so adjacency is not required.
|
||||
if old_meter.ended_at is None:
|
||||
raise BindingValidationError("Source binding must belong to a closed predecessor meter.")
|
||||
timeline = list(session.execute(
|
||||
select(Meter).where(Meter.commodity == target.commodity)
|
||||
).scalars())
|
||||
predecessors = [
|
||||
meter for meter in timeline
|
||||
if meter.id != target.id
|
||||
and meter.ended_at is not None
|
||||
and _as_utc(meter.ended_at) <= _as_utc(target.started_at)
|
||||
]
|
||||
if not predecessors:
|
||||
raise BindingValidationError("Source meter is not the unique immediately preceding meter.")
|
||||
latest_end = max(_as_utc(meter.ended_at) for meter in predecessors)
|
||||
latest = [meter for meter in predecessors if _as_utc(meter.ended_at) == latest_end]
|
||||
if len(latest) != 1 or latest[0].id != old_meter.id:
|
||||
raise BindingValidationError("Source meter is not the unique immediately preceding meter.")
|
||||
# Reject any overlapping epoch around either endpoint. A separate
|
||||
# meter inside the gap is already excluded by the predecessor check;
|
||||
# one extending into either endpoint is an ambiguous timeline too.
|
||||
for meter in timeline:
|
||||
if meter.id in {old_meter.id, target.id}:
|
||||
continue
|
||||
meter_end = _as_utc(meter.ended_at) if meter.ended_at is not None else None
|
||||
if (
|
||||
half_open_intervals_overlap(
|
||||
_as_utc(old_meter.started_at), _as_utc(old_meter.ended_at),
|
||||
_as_utc(meter.started_at), meter_end,
|
||||
)
|
||||
or half_open_intervals_overlap(
|
||||
_as_utc(target.started_at),
|
||||
_as_utc(target.ended_at) if target.ended_at is not None else None,
|
||||
_as_utc(meter.started_at), meter_end,
|
||||
)
|
||||
):
|
||||
raise BindingValidationError("Source meter has an ambiguous commodity timeline.")
|
||||
close_at = _as_utc(old_meter.ended_at)
|
||||
if effective_at < _as_utc(target.started_at):
|
||||
raise BindingValidationError("Transfer effective_at must be within the target meter epoch.")
|
||||
if effective_at < _as_utc(source.started_at):
|
||||
raise BindingValidationError("Transfer effective_at precedes the source binding.")
|
||||
# Validate the target before mutating the old row, then close/create in one session.
|
||||
_validate_binding(session, meter_id=target.id, channel_id=to_channel_id,
|
||||
started_at=effective_at, ended_at=None,
|
||||
excluding_ids={source.id})
|
||||
update_binding(session, source.id, ended_at=close_at)
|
||||
created = create_binding(session, meter_id=target.id, channel_id=to_channel_id,
|
||||
started_at=effective_at)
|
||||
return source, created
|
||||
@@ -0,0 +1,554 @@
|
||||
"""Service layer for Meter epoch CRUD, swap/close, and time-range lookup.
|
||||
|
||||
All functions accept an explicit SQLAlchemy Session; callers are responsible
|
||||
for committing or rolling back the transaction.
|
||||
|
||||
Design decisions
|
||||
----------------
|
||||
- ``meter_at``: half-open interval ``[started_at, ended_at)`` lookup;
|
||||
returns the meter whose epoch covers *ts* for the given commodity.
|
||||
SQLite naive datetime is normalised via ``_as_utc()`` before comparison.
|
||||
|
||||
- ``declare_meter``: validates that the new ``started_at`` is **not earlier
|
||||
than** the current active meter's ``started_at`` (rejects back-dating below
|
||||
the active meter's own start). Equal timestamps are allowed because the
|
||||
typical "swap now" use-case sets ``started_at`` to the current moment, which
|
||||
coincides with the active meter's ``started_at`` only in degenerate test
|
||||
scenarios — but blocking equal values would make that workflow impossible.
|
||||
The old active meter is closed (``ended_at = started_at``) and a new active
|
||||
meter is opened in the same operation, guaranteeing continuity: the old
|
||||
meter's ``ended_at`` equals the new meter's ``started_at`` (contiguous,
|
||||
no gap, no overlap).
|
||||
|
||||
- ``update_meter``: when ``started_at`` is modified, the service keeps the
|
||||
timeline contiguous by also updating the **previous** meter's ``ended_at``
|
||||
(the one whose ``ended_at`` matched the old ``started_at``) to the new
|
||||
``started_at``. Validation ensures the new ``started_at``:
|
||||
* is strictly after the previous meter's own ``started_at``
|
||||
(cannot push the boundary before the previous meter even started);
|
||||
* is strictly before the current meter's ``ended_at``, if set
|
||||
(cannot push the boundary past where the current meter was already
|
||||
closed).
|
||||
|
||||
- Mutual exclusion (at most one active meter per commodity) is enforced by the
|
||||
service layer; no DB-level unique partial index is added to keep migrations
|
||||
simple and to allow the application to return a meaningful error message.
|
||||
|
||||
SQLite timezone note
|
||||
--------------------
|
||||
SQLite stores ``DateTime(timezone=True)`` columns as naive UTC strings; on
|
||||
read-back they come out as **timezone-naive** datetimes. Wherever this code
|
||||
compares timestamps from the DB against timezone-aware values, it calls
|
||||
``_as_utc()`` to make both sides comparable without tripping on
|
||||
"offset-naive vs offset-aware" TypeErrors.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models.energy import Meter
|
||||
from app.models.meter_source import MeterSourceBinding
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _as_utc(dt: datetime) -> datetime:
|
||||
"""Return *dt* as a timezone-aware UTC datetime.
|
||||
|
||||
SQLite's ``DateTime(timezone=True)`` column type stores datetimes as naive
|
||||
UTC strings and gives them back as naive datetimes on read. This helper
|
||||
re-attaches the UTC timezone info when it is missing, making cross-origin
|
||||
comparisons safe.
|
||||
"""
|
||||
if dt.tzinfo is None:
|
||||
return dt.replace(tzinfo=UTC)
|
||||
return dt
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Custom exceptions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MeterError(ValueError):
|
||||
"""Base class for meter service validation errors."""
|
||||
|
||||
|
||||
class MeterOverlapError(MeterError):
|
||||
"""Raised when a new meter's ``started_at`` would create an overlap or backdate.
|
||||
|
||||
Specifically: the new meter's ``started_at`` must be greater than or equal
|
||||
to the current active meter's ``started_at``. Allowing a value strictly
|
||||
earlier than the active meter's start would mean the new epoch begins before
|
||||
the current one, which is chronologically inconsistent.
|
||||
"""
|
||||
|
||||
|
||||
class MeterIntervalError(MeterError):
|
||||
"""Raised when an ``update_meter`` call would produce an inconsistent interval.
|
||||
|
||||
Examples of inconsistent intervals:
|
||||
- New ``started_at`` ≥ this meter's ``ended_at`` (epoch would be empty/inverted).
|
||||
- New ``started_at`` ≤ the previous meter's own ``started_at`` (previous meter
|
||||
would become empty/inverted after its ``ended_at`` is updated).
|
||||
"""
|
||||
|
||||
|
||||
def close_meter(session: Session, meter: Meter, *, ended_at: datetime) -> Meter:
|
||||
"""Close an active meter at a valid, non-future exclusive boundary."""
|
||||
boundary = _as_utc(ended_at)
|
||||
if meter.ended_at is not None:
|
||||
raise MeterIntervalError("Only an active meter can be closed.")
|
||||
if boundary <= _as_utc(meter.started_at):
|
||||
raise MeterIntervalError("Meter ended_at must be strictly after started_at.")
|
||||
if boundary > datetime.now(UTC):
|
||||
raise MeterIntervalError("Meter ended_at must not be in the future.")
|
||||
_validate_bindings_fit_meter_end(session, meter, boundary)
|
||||
meter.ended_at = boundary
|
||||
return meter
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Internal query helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _active_meter(session: Session, commodity: str) -> Optional[Meter]:
|
||||
"""Return the current active (``ended_at IS NULL``) meter for *commodity*, or None."""
|
||||
return session.execute(
|
||||
select(Meter)
|
||||
.where(
|
||||
Meter.commodity == commodity,
|
||||
Meter.ended_at.is_(None),
|
||||
)
|
||||
.limit(1)
|
||||
).scalar_one_or_none()
|
||||
|
||||
|
||||
def _validate_bindings_fit_meter_end(session: Session, meter: Meter, boundary: datetime) -> None:
|
||||
"""Reject an epoch close that would put any retained binding out of bounds.
|
||||
|
||||
Closed binding history is immutable here. Open bindings may subsequently
|
||||
be closed by the caller at the shared meter boundary, but only when that
|
||||
produces a non-empty interval.
|
||||
"""
|
||||
for binding in session.execute(
|
||||
select(MeterSourceBinding).where(MeterSourceBinding.meter_id == meter.id)
|
||||
).scalars():
|
||||
if binding.ended_at is None:
|
||||
if _as_utc(binding.started_at) >= boundary:
|
||||
raise MeterIntervalError("Open binding cannot be closed within the proposed meter epoch.")
|
||||
elif _as_utc(binding.ended_at) > boundary:
|
||||
raise MeterIntervalError("Closed binding extends beyond the proposed meter epoch.")
|
||||
|
||||
|
||||
def _meter_before(session: Session, meter: Meter) -> Optional[Meter]:
|
||||
"""Return the meter whose ``ended_at`` equals *meter*'s ``started_at``.
|
||||
|
||||
This is the meter that was closed when *meter* was opened; its ``ended_at``
|
||||
needs to stay equal to *meter*'s ``started_at`` to maintain timeline
|
||||
continuity. Returns None if *meter* is the first epoch for its commodity.
|
||||
|
||||
The lookup compares naive/aware datetimes via string to avoid SQLite timezone
|
||||
quirks: both are formatted as ISO 8601 UTC strings for the WHERE clause.
|
||||
We rely on the fact that the service layer always stores the same timestamp
|
||||
object as both ``prev.ended_at`` and ``new.started_at``, so their string
|
||||
representations are identical.
|
||||
"""
|
||||
# Normalise the target to a tz-aware UTC datetime for comparison.
|
||||
# We scan in Python (rather than SQL) to avoid SQLite naive-vs-aware
|
||||
# mismatch issues when comparing DateTime columns against tz-aware values.
|
||||
target = _as_utc(meter.started_at)
|
||||
|
||||
# Fetch all closed meters of the same commodity and find the one whose
|
||||
# ended_at equals this meter's started_at (the standard contiguous handoff).
|
||||
candidates = session.execute(
|
||||
select(Meter)
|
||||
.where(
|
||||
Meter.commodity == meter.commodity,
|
||||
Meter.ended_at.is_not(None),
|
||||
)
|
||||
).scalars().all()
|
||||
|
||||
for candidate in candidates:
|
||||
if candidate.id == meter.id:
|
||||
continue
|
||||
candidate_ended = _as_utc(candidate.ended_at)
|
||||
if candidate_ended == target:
|
||||
return candidate
|
||||
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Core service functions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def meter_at(
|
||||
session: Session,
|
||||
ts: datetime,
|
||||
commodity: str = "electricity",
|
||||
) -> Optional[Meter]:
|
||||
"""Return the meter epoch that covers *ts* for *commodity*.
|
||||
|
||||
A meter covers *ts* when:
|
||||
``started_at ≤ ts`` AND (``ended_at IS NULL`` OR ``ts < ended_at``)
|
||||
|
||||
This is the standard half-open interval ``[started_at, ended_at)`` lookup,
|
||||
consistent with ``active_contract_version_at`` in ``contracts.py``.
|
||||
|
||||
Returns ``None`` when no meter covers *ts* (e.g. before the first epoch
|
||||
was declared, or after a gap).
|
||||
|
||||
Parameters
|
||||
----------
|
||||
session:
|
||||
Active SQLAlchemy session (read-only usage).
|
||||
ts:
|
||||
UTC datetime to look up.
|
||||
commodity:
|
||||
Energy commodity (default ``"electricity"``).
|
||||
|
||||
Returns
|
||||
-------
|
||||
Meter | None
|
||||
|
||||
Implementation note — why the upper bound is pushed into SQL
|
||||
------------------------------------------------------------
|
||||
When ``declare_meter`` is called with ``started_at == active.started_at``
|
||||
(the "equal-timestamp swap" allowed by §3.5), the old meter is closed as a
|
||||
zero-width epoch ``[T0, T0)`` and the new active meter also has
|
||||
``started_at == T0``. Two rows share the same ``started_at``; SQLite's
|
||||
row-ordering for ``ORDER BY started_at DESC LIMIT 1`` is then determined by
|
||||
rowid (i.e. insertion order), which returns the older zero-width row first.
|
||||
|
||||
If the upper bound were checked in Python *after* fetching that one row, the
|
||||
condition ``ts < ended_at`` would be False for **any** ``ts >= T0`` (because
|
||||
``ended_at == T0``), causing the function to return ``None`` and making the
|
||||
new active meter permanently invisible.
|
||||
|
||||
Pushing both bounds into the SQL ``WHERE`` clause eliminates the ambiguity:
|
||||
the zero-width row is excluded by ``ts < ended_at`` before ``LIMIT 1`` is
|
||||
applied, so only the genuinely covering row survives.
|
||||
|
||||
SQLite naive-vs-aware datetime note: SQLAlchemy's SQLite dialect strips the
|
||||
``tzinfo`` from aware datetimes when binding parameters (it does *not*
|
||||
convert to UTC first). Since all datetimes in this codebase are UTC (either
|
||||
naive-UTC from the DB or aware-UTC from ``datetime.now(UTC)``), stripping
|
||||
the tzinfo leaves the wall-clock value unchanged and comparisons remain
|
||||
correct. This is consistent with how ``contracts.py`` handles the same
|
||||
situation (see OBS 2 in the M7-T02 review notes).
|
||||
"""
|
||||
# Push both the lower and upper bounds into the SQL WHERE clause so that
|
||||
# zero-width epochs (ended_at == started_at) are excluded *before* LIMIT 1
|
||||
# is applied. This prevents an equal-timestamp swap from making the new
|
||||
# active meter invisible (see implementation note above).
|
||||
stmt = (
|
||||
select(Meter)
|
||||
.where(
|
||||
Meter.commodity == commodity,
|
||||
Meter.started_at <= ts,
|
||||
(Meter.ended_at.is_(None)) | (Meter.ended_at > ts),
|
||||
)
|
||||
.order_by(Meter.started_at.desc())
|
||||
.limit(1)
|
||||
)
|
||||
return session.execute(stmt).scalar_one_or_none()
|
||||
|
||||
|
||||
def declare_meter(
|
||||
session: Session,
|
||||
*,
|
||||
label: str,
|
||||
started_at: datetime,
|
||||
reason: str,
|
||||
commodity: str = "electricity",
|
||||
note: Optional[str] = None,
|
||||
) -> Meter:
|
||||
"""Declare a meter swap or initial meter epoch.
|
||||
|
||||
Closes the current active meter for *commodity* (if one exists) by setting
|
||||
its ``ended_at`` to *started_at*, then opens a new active meter. The
|
||||
resulting timeline is **contiguous**: old meter's ``ended_at`` equals new
|
||||
meter's ``started_at``.
|
||||
|
||||
If there is no current active meter (first-ever declaration for this
|
||||
commodity), the new meter is simply opened without closing anything.
|
||||
|
||||
Validation
|
||||
----------
|
||||
- If a current active meter exists, *started_at* must be **≥** that meter's
|
||||
own ``started_at``. A value strictly earlier would place the new epoch
|
||||
entirely before the current active meter, which is a chronological
|
||||
contradiction ("back-dating before the active epoch's start"). Raises
|
||||
``MeterOverlapError`` when this constraint is violated.
|
||||
|
||||
Note: equal timestamps (``started_at == active.started_at``) are allowed
|
||||
because that scenario effectively replaces the current meter at the same
|
||||
logical moment (e.g. correcting a mis-entry), which is a valid use-case.
|
||||
The old meter is then closed with ``ended_at == started_at`` (a zero-width
|
||||
epoch), which is intentional and auditable.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
session:
|
||||
Active SQLAlchemy session. Caller must commit after this returns.
|
||||
label:
|
||||
Human-readable label for the new meter.
|
||||
started_at:
|
||||
UTC datetime at which this meter epoch starts. May be in the past.
|
||||
reason:
|
||||
Why this epoch was created (``"initial"``, ``"meter_swap"``,
|
||||
``"home_move"``, or ``"other"``).
|
||||
commodity:
|
||||
Energy commodity (default ``"electricity"``).
|
||||
note:
|
||||
Optional free-form note.
|
||||
|
||||
Returns
|
||||
-------
|
||||
Meter
|
||||
The newly created, not-yet-committed active meter.
|
||||
|
||||
Raises
|
||||
------
|
||||
MeterOverlapError
|
||||
If *started_at* is strictly earlier than the current active meter's
|
||||
``started_at`` (chronological backdate below the active epoch's start).
|
||||
"""
|
||||
if _as_utc(started_at) > datetime.now(UTC):
|
||||
raise MeterIntervalError("Meter started_at must not be in the future.")
|
||||
active = _active_meter(session, commodity)
|
||||
|
||||
if active is not None:
|
||||
# Reject back-dating: new started_at must be ≥ current active's started_at.
|
||||
if _as_utc(started_at) < _as_utc(active.started_at):
|
||||
raise MeterOverlapError(
|
||||
f"New meter started_at ({started_at.isoformat()}) must be ≥ the current "
|
||||
f"active {commodity!r} meter's started_at "
|
||||
f"({active.started_at.isoformat()}). "
|
||||
"Declare a started_at on or after the active meter's start to avoid "
|
||||
"a chronologically inconsistent epoch ordering."
|
||||
)
|
||||
# Validate before changing the epoch: retained closed binding history
|
||||
# must never be silently truncated by a later declaration.
|
||||
_validate_bindings_fit_meter_end(session, active, _as_utc(started_at))
|
||||
# Close the current active meter at the swap point (contiguous handoff).
|
||||
active.ended_at = started_at
|
||||
logger.info(
|
||||
"Closed active %r meter id=%d (ended_at=%s)",
|
||||
commodity,
|
||||
active.id,
|
||||
started_at.isoformat(),
|
||||
)
|
||||
|
||||
now = datetime.now(UTC)
|
||||
new_meter = Meter(
|
||||
label=label,
|
||||
commodity=commodity,
|
||||
started_at=started_at,
|
||||
ended_at=None,
|
||||
reason=reason,
|
||||
note=note,
|
||||
created_at=now,
|
||||
)
|
||||
session.add(new_meter)
|
||||
|
||||
logger.info(
|
||||
"Declared new %r meter %r (started_at=%s, reason=%s)",
|
||||
commodity,
|
||||
label,
|
||||
started_at.isoformat(),
|
||||
reason,
|
||||
)
|
||||
return new_meter
|
||||
|
||||
|
||||
def list_meters(
|
||||
session: Session,
|
||||
commodity: Optional[str] = None,
|
||||
) -> list[Meter]:
|
||||
"""List all meter epochs, ordered by started_at ascending.
|
||||
|
||||
Active meters (``ended_at IS NULL``) sort naturally to the end of the
|
||||
timeline since they have the latest ``started_at``. Within a single
|
||||
commodity, the ascending ``started_at`` order reflects the historical
|
||||
sequence of installed meters.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
session:
|
||||
Active SQLAlchemy session (read-only usage).
|
||||
commodity:
|
||||
If provided, filter to this commodity only. If ``None``, return all
|
||||
meters across all commodities.
|
||||
|
||||
Returns
|
||||
-------
|
||||
list[Meter]
|
||||
Meters ordered by (started_at ASC).
|
||||
"""
|
||||
stmt = select(Meter).order_by(Meter.started_at.asc())
|
||||
if commodity is not None:
|
||||
stmt = stmt.where(Meter.commodity == commodity)
|
||||
|
||||
return list(session.execute(stmt).scalars().all())
|
||||
|
||||
|
||||
def update_meter(
|
||||
session: Session,
|
||||
meter: Meter,
|
||||
*,
|
||||
label: Optional[str] = None,
|
||||
note: Optional[str] = None,
|
||||
started_at: Optional[datetime] = None,
|
||||
) -> Meter:
|
||||
"""Update a meter's mutable fields (label, note, started_at).
|
||||
|
||||
Passing ``None`` for a field leaves it unchanged. At least one keyword
|
||||
argument must be non-``None``; calling with all-``None`` is a no-op but
|
||||
is not an error.
|
||||
|
||||
Updating ``started_at`` (retroactive correction)
|
||||
------------------------------------------------
|
||||
When ``started_at`` is provided the service maintains **timeline
|
||||
continuity** across the adjacent meter boundaries:
|
||||
|
||||
1. **Previous meter's ``ended_at``** — if the meter immediately before
|
||||
this one has ``ended_at == meter.started_at`` (the standard contiguous
|
||||
handoff), its ``ended_at`` is updated to the new ``started_at`` so the
|
||||
boundary between the two epochs stays seamless.
|
||||
|
||||
2. **Validation** — the new ``started_at`` is checked for consistency:
|
||||
a. It must be **strictly after** the previous meter's own ``started_at``
|
||||
(otherwise the previous meter's epoch would collapse to zero or
|
||||
invert).
|
||||
b. It must be **strictly before** this meter's ``ended_at`` (if set),
|
||||
so this meter's epoch remains non-empty.
|
||||
c. Every binding on this meter and its affected predecessor must remain
|
||||
wholly inside its proposed epoch. The service rejects the correction
|
||||
rather than rewriting binding history.
|
||||
|
||||
Note: triggering a billing recompute (``recompute_range``) after a
|
||||
retroactive ``started_at`` change is **out of scope** for this service
|
||||
layer; that is the API layer's responsibility (M7-T05).
|
||||
|
||||
Parameters
|
||||
----------
|
||||
session:
|
||||
Active SQLAlchemy session. Caller must commit after this returns.
|
||||
meter:
|
||||
The ``Meter`` ORM instance to update (already loaded from the session).
|
||||
label:
|
||||
New human-readable label; ``None`` = keep existing.
|
||||
note:
|
||||
New free-form note; ``None`` = keep existing.
|
||||
started_at:
|
||||
New start timestamp for this meter epoch; ``None`` = keep existing.
|
||||
|
||||
Returns
|
||||
-------
|
||||
Meter
|
||||
The updated ``Meter`` instance (not yet committed).
|
||||
|
||||
Raises
|
||||
------
|
||||
MeterIntervalError
|
||||
If the new ``started_at`` would produce an invalid (empty or inverted)
|
||||
epoch for this meter or the immediately preceding one.
|
||||
"""
|
||||
# Validate the proposed epoch boundary before touching *any* mutable
|
||||
# field. PATCH accepts label/note together with started_at, so doing this
|
||||
# first keeps an invalid future timestamp from leaking a partial in-session
|
||||
# update before the API's rollback boundary is reached.
|
||||
proposed_started_at = _as_utc(started_at) if started_at is not None else None
|
||||
if proposed_started_at is not None and proposed_started_at > datetime.now(UTC):
|
||||
raise MeterIntervalError("Meter started_at must not be in the future.")
|
||||
|
||||
if label is not None:
|
||||
meter.label = label
|
||||
logger.info("Updated meter id=%d label=%r", meter.id, label)
|
||||
|
||||
if note is not None:
|
||||
meter.note = note
|
||||
logger.info("Updated meter id=%d note=%r", meter.id, note)
|
||||
|
||||
if started_at is not None:
|
||||
old_started_at = meter.started_at
|
||||
assert proposed_started_at is not None
|
||||
|
||||
# --- Validate upper bound: new started_at must be < this meter's ended_at (if set).
|
||||
if meter.ended_at is not None:
|
||||
if proposed_started_at >= _as_utc(meter.ended_at):
|
||||
raise MeterIntervalError(
|
||||
f"New started_at ({started_at.isoformat()}) must be strictly before "
|
||||
f"this meter's ended_at ({meter.ended_at.isoformat()}). "
|
||||
"The meter epoch would become empty or inverted."
|
||||
)
|
||||
|
||||
# --- Find the immediately preceding meter (its ended_at == meter's old started_at).
|
||||
prev = _meter_before(session, meter)
|
||||
|
||||
# --- Validate lower bound: new started_at must be strictly after prev's started_at.
|
||||
if prev is not None:
|
||||
if proposed_started_at <= _as_utc(prev.started_at):
|
||||
raise MeterIntervalError(
|
||||
f"New started_at ({started_at.isoformat()}) must be strictly after "
|
||||
f"the previous meter's started_at ({prev.started_at.isoformat()}). "
|
||||
"Moving the boundary that far back would collapse the previous meter's epoch."
|
||||
)
|
||||
# A boundary correction changes both adjacent meter epochs. Fail closed
|
||||
# rather than silently rewriting binding history: every existing binding
|
||||
# must still fit in its proposed epoch before either Meter is mutated.
|
||||
affected_meters = [
|
||||
(meter, proposed_started_at, _as_utc(meter.ended_at) if meter.ended_at is not None else None)
|
||||
]
|
||||
if prev is not None:
|
||||
affected_meters.append((prev, _as_utc(prev.started_at), proposed_started_at))
|
||||
for affected_meter, proposed_start, proposed_end in affected_meters:
|
||||
bindings = session.execute(
|
||||
select(MeterSourceBinding).where(MeterSourceBinding.meter_id == affected_meter.id)
|
||||
).scalars()
|
||||
for binding in bindings:
|
||||
if _as_utc(binding.started_at) < proposed_start:
|
||||
raise MeterIntervalError(
|
||||
f"Binding {binding.id} starts before meter {affected_meter.id}'s epoch."
|
||||
)
|
||||
if proposed_end is not None and (
|
||||
binding.ended_at is None or _as_utc(binding.ended_at) > proposed_end
|
||||
):
|
||||
raise MeterIntervalError(
|
||||
f"Binding {binding.id} would fall outside meter {affected_meter.id}'s epoch."
|
||||
)
|
||||
|
||||
if prev is not None:
|
||||
# Maintain continuity: update the previous meter's ended_at to match the new start.
|
||||
prev.ended_at = started_at
|
||||
logger.info(
|
||||
"Updated previous meter id=%d ended_at=%s (boundary shift from %s)",
|
||||
prev.id,
|
||||
started_at.isoformat(),
|
||||
old_started_at.isoformat(),
|
||||
)
|
||||
|
||||
meter.started_at = started_at
|
||||
logger.info(
|
||||
"Updated meter id=%d started_at=%s (was %s)",
|
||||
meter.id,
|
||||
started_at.isoformat(),
|
||||
old_started_at.isoformat(),
|
||||
)
|
||||
|
||||
return meter
|
||||
@@ -0,0 +1,210 @@
|
||||
"""Modbus polling service — periodic read-and-store for enabled Modbus devices.
|
||||
|
||||
Design notes
|
||||
------------
|
||||
- ``poll_device`` loads the device's YAML profile, calls the driver to bulk-read
|
||||
all register blocks in a single TCP connection, decodes the raw registers into
|
||||
an engineering-value dict via the profile, and persists one ``ModbusReading``
|
||||
row. It also stamps ``last_poll_at`` / ``last_poll_ok`` on the device.
|
||||
- All exceptions are caught and logged inside ``poll_device``; the function
|
||||
never re-raises. This prevents a single broken device from aborting the
|
||||
sweep for all other devices.
|
||||
- ``poll_all_enabled_devices`` respects per-device ``poll_interval_s``: it skips
|
||||
a device whose ``last_poll_at`` is recent enough. This lets the APScheduler
|
||||
job run on a short base tick (e.g. 5 s) while each device self-regulates its
|
||||
own cadence.
|
||||
- The global ``modbus_polling_enabled`` flag from ``Settings`` (merged with any
|
||||
DB overrides via ``build_runtime_settings``) gates the entire sweep as a
|
||||
no-op. T08 will expose this flag in CONFIG_FIELDS/UI; for now it is
|
||||
controlled by the ``MODBUS_POLLING_ENABLED`` env var (or the default True).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.config import Settings
|
||||
from app.integrations.modbus import driver, profiles
|
||||
from app.models.modbus import ModbusDevice, ModbusReading
|
||||
from app.services.config_page import build_runtime_settings
|
||||
from app.services.ha_discovery import publish_device_offline, publish_device_state
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Base polling tick in seconds (used by APScheduler; each device checks its own interval).
|
||||
BASE_POLL_TICK_SECONDS = 5
|
||||
|
||||
|
||||
def poll_device(session: Session, device: ModbusDevice) -> ModbusReading | None:
|
||||
"""Read one enabled Modbus device, decode, and persist a reading row.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
session:
|
||||
Open SQLAlchemy session. The caller is responsible for session
|
||||
lifecycle (open / commit / close).
|
||||
device:
|
||||
``ModbusDevice`` ORM instance to poll.
|
||||
|
||||
Returns
|
||||
-------
|
||||
ModbusReading | None
|
||||
The newly-persisted reading on success, or ``None`` on any failure.
|
||||
Failure is logged and the device's ``last_poll_ok`` is set to ``False``.
|
||||
"""
|
||||
now = datetime.now(UTC)
|
||||
try:
|
||||
profile = profiles.load_profile(device.profile)
|
||||
registers = driver.read_blocks(
|
||||
device.host,
|
||||
device.port,
|
||||
device.unit_id,
|
||||
[{"start": b.start, "count": b.count} for b in profile.blocks],
|
||||
function_code=profile.function_code,
|
||||
)
|
||||
payload = profiles.decode(profile, registers)
|
||||
|
||||
reading = ModbusReading(
|
||||
device_id=device.id,
|
||||
recorded_at=now,
|
||||
payload=payload,
|
||||
)
|
||||
session.add(reading)
|
||||
|
||||
device.last_poll_at = now
|
||||
device.last_poll_ok = True
|
||||
session.commit()
|
||||
|
||||
logger.debug(
|
||||
"Polled device %r (id=%d): %d metrics recorded at %s",
|
||||
device.friendly_name,
|
||||
device.id,
|
||||
len(payload),
|
||||
now.isoformat(),
|
||||
)
|
||||
|
||||
# Best-effort: push MQTT state after successful poll.
|
||||
# Must never let MQTT errors crash the poll loop.
|
||||
try:
|
||||
publish_device_state(session, device)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception(
|
||||
"poll_device: MQTT state publish failed for device %r (id=%d); "
|
||||
"continuing (poll itself succeeded)",
|
||||
device.friendly_name,
|
||||
device.id,
|
||||
)
|
||||
|
||||
return reading
|
||||
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception(
|
||||
"Failed to poll device %r (id=%d, host=%s:%d, unit_id=%d)",
|
||||
device.friendly_name,
|
||||
device.id,
|
||||
device.host,
|
||||
device.port,
|
||||
device.unit_id,
|
||||
)
|
||||
# Rollback first — if session.add(reading) already ran in the success
|
||||
# branch before commit() failed, that pending ModbusReading must be
|
||||
# discarded. Without rollback it would remain staged and could be
|
||||
# flushed/committed by a subsequent operation (even for a different
|
||||
# device), producing a spurious reading row for a poll that failed.
|
||||
# Rollback also expires all ORM objects; device fields are re-set below.
|
||||
try:
|
||||
session.rollback()
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception(
|
||||
"session.rollback() failed for device id=%d; session may be unusable", device.id
|
||||
)
|
||||
try:
|
||||
device.last_poll_at = now
|
||||
device.last_poll_ok = False
|
||||
session.commit()
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception(
|
||||
"Also failed to persist last_poll_ok=False for device id=%d", device.id
|
||||
)
|
||||
|
||||
# Best-effort: publish "offline" availability after failed poll.
|
||||
try:
|
||||
publish_device_offline(session, device.uuid)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception(
|
||||
"poll_device: MQTT offline publish failed for device %r (id=%d)",
|
||||
device.friendly_name,
|
||||
device.id,
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def poll_all_enabled_devices(
|
||||
session: Session,
|
||||
*,
|
||||
bootstrap_settings: Settings,
|
||||
) -> None:
|
||||
"""Sweep all enabled Modbus devices and poll each one whose interval has elapsed.
|
||||
|
||||
Global kill-switch: if ``build_runtime_settings(...).modbus_polling_enabled``
|
||||
is ``False``, this function returns immediately without touching any device.
|
||||
|
||||
Per-device interval: a device is skipped if its ``last_poll_at`` is not yet
|
||||
``poll_interval_s`` seconds in the past. This allows the APScheduler job to
|
||||
run on a short fixed tick while each device self-regulates its own cadence.
|
||||
|
||||
Exceptions from individual devices are swallowed inside ``poll_device``;
|
||||
this function itself will not raise.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
session:
|
||||
Open SQLAlchemy session.
|
||||
bootstrap_settings:
|
||||
The bootstrap ``Settings`` instance (from ``get_settings()``), used as
|
||||
the base for ``build_runtime_settings`` to pick up any DB overrides.
|
||||
"""
|
||||
try:
|
||||
runtime_settings = build_runtime_settings(session, bootstrap_settings)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception("Failed to build runtime settings for Modbus poll; aborting sweep")
|
||||
return
|
||||
|
||||
if not runtime_settings.modbus_polling_enabled:
|
||||
logger.debug("Modbus polling is disabled (modbus_polling_enabled=False); skipping sweep")
|
||||
return
|
||||
|
||||
try:
|
||||
devices = session.execute(
|
||||
select(ModbusDevice).where(ModbusDevice.enabled == True) # noqa: E712
|
||||
).scalars().all()
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception("Failed to query enabled Modbus devices; aborting sweep")
|
||||
return
|
||||
|
||||
now = datetime.now(UTC)
|
||||
for device in devices:
|
||||
# Respect per-device poll interval: skip if last poll was too recent.
|
||||
if device.last_poll_at is not None:
|
||||
# SQLite may return a naive datetime despite DateTime(timezone=True) —
|
||||
# treat naive values as UTC so the comparison is safe.
|
||||
last_poll = device.last_poll_at
|
||||
if last_poll.tzinfo is None:
|
||||
last_poll = last_poll.replace(tzinfo=UTC)
|
||||
elapsed = (now - last_poll).total_seconds()
|
||||
if elapsed < device.poll_interval_s:
|
||||
logger.debug(
|
||||
"Skipping device %r (id=%d): %.1fs elapsed < %ds interval",
|
||||
device.friendly_name,
|
||||
device.id,
|
||||
elapsed,
|
||||
device.poll_interval_s,
|
||||
)
|
||||
continue
|
||||
|
||||
poll_device(session, device)
|
||||
+48
-1
@@ -4,7 +4,7 @@ from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
import logging
|
||||
|
||||
from sqlalchemy import desc, insert, select
|
||||
from sqlalchemy import delete, desc, insert, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.config import Settings
|
||||
@@ -74,6 +74,53 @@ def record_poo(
|
||||
logger.warning("Failed to trigger poo webhook on Home Assistant: %s", exc)
|
||||
|
||||
|
||||
def update_poo_record(
|
||||
session: Session,
|
||||
timestamp_pk: str,
|
||||
*,
|
||||
status: str | None,
|
||||
latitude: float | None,
|
||||
longitude: float | None,
|
||||
) -> PooRecord | None:
|
||||
"""Update non-PK fields of a single poo record row.
|
||||
|
||||
Returns the updated ORM object, or ``None`` if the PK does not exist.
|
||||
The ``timestamp`` PK is immutable and must not be passed as an update field.
|
||||
Only fields with a non-``None`` value are written.
|
||||
"""
|
||||
row = session.execute(
|
||||
select(PooRecord).where(PooRecord.timestamp == timestamp_pk)
|
||||
).scalar_one_or_none()
|
||||
|
||||
if row is None:
|
||||
return None
|
||||
|
||||
if status is not None:
|
||||
row.status = status
|
||||
if latitude is not None:
|
||||
row.latitude = latitude
|
||||
if longitude is not None:
|
||||
row.longitude = longitude
|
||||
|
||||
session.commit()
|
||||
session.refresh(row)
|
||||
return row
|
||||
|
||||
|
||||
def delete_poo_record(session: Session, timestamp_pk: str) -> bool:
|
||||
"""Delete the single poo record row identified by its PK.
|
||||
|
||||
Returns ``True`` if exactly one row was deleted, ``False`` if the PK did
|
||||
not exist (caller should raise 404). The DELETE is scoped to the exact PK
|
||||
— no batch/truncate path exists.
|
||||
"""
|
||||
result = session.execute(
|
||||
delete(PooRecord).where(PooRecord.timestamp == timestamp_pk)
|
||||
)
|
||||
session.commit()
|
||||
return result.rowcount == 1
|
||||
|
||||
|
||||
def get_latest_poo_record(session: Session) -> LatestPooRecord | None:
|
||||
stmt = select(PooRecord).order_by(desc(PooRecord.timestamp)).limit(1)
|
||||
record = session.execute(stmt).scalar_one_or_none()
|
||||
|
||||
@@ -0,0 +1,202 @@
|
||||
"""Service layer for fetching and persisting Tibber 15-minute electricity prices.
|
||||
|
||||
``refresh_prices`` is the main entry point. It is designed to be called from a
|
||||
scheduled background job (see ``app/main.py``) and from tests.
|
||||
|
||||
Design decisions
|
||||
----------------
|
||||
- **Guard clause**: if no active contract with ``kind="tibber"`` exists, or if
|
||||
the Tibber API token is empty in the runtime settings, the function is a
|
||||
complete no-op (returns 0) and does not raise. This means the scheduler can
|
||||
call ``refresh_prices`` unconditionally; the service itself decides whether to
|
||||
do anything based on current configuration.
|
||||
- **Upsert idempotency**: rows are matched by ``starts_at`` (the unique
|
||||
constraint on ``tibber_price``). If a row already exists, its price fields
|
||||
and ``fetched_at`` are updated in-place; if it does not exist, a new row is
|
||||
inserted. Running ``refresh_prices`` twice in a row must not double-insert.
|
||||
- **No destructive operations**: this service never deletes ``tibber_price``
|
||||
rows. Only INSERT or UPDATE.
|
||||
- **Exception propagation**: exceptions from the Tibber client are *not*
|
||||
swallowed here. The scheduled job wrapper in ``main.py`` is responsible for
|
||||
catching and logging errors so that a single fetch failure does not crash the
|
||||
scheduler.
|
||||
- **SQLite timezone note**: ``DateTime(timezone=True)`` columns come back as
|
||||
timezone-naive UTC datetimes on read. We store UTC-aware datetimes on write
|
||||
(``starts_at`` from ``PricePoint`` is already UTC-aware; ``fetched_at`` is
|
||||
``datetime.now(UTC)``). Reads elsewhere that compare against these values
|
||||
must normalise accordingly.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from threading import Lock, Thread
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.integrations.tibber.client import fetch_price_range
|
||||
from app.models.energy import EnergyContract, TibberPrice
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
_background_refresh_lock = Lock()
|
||||
|
||||
|
||||
def active_tibber_contract_exists(session: Session) -> bool:
|
||||
"""Return True if there is an active contract with kind='tibber'."""
|
||||
row = session.execute(
|
||||
select(EnergyContract).where(
|
||||
EnergyContract.active.is_(True),
|
||||
EnergyContract.kind == "tibber",
|
||||
).limit(1)
|
||||
).scalar_one_or_none()
|
||||
return row is not None
|
||||
|
||||
|
||||
def refresh_prices(session: Session, settings: object) -> int:
|
||||
"""Fetch today-and-tomorrow Tibber prices and upsert them into ``tibber_price``.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
session:
|
||||
An active SQLAlchemy session. The caller is responsible for closing it;
|
||||
this function commits each upserted row individually to keep transactions
|
||||
short.
|
||||
settings:
|
||||
A runtime settings object that exposes ``tibber_api_token`` and
|
||||
``tibber_home_id`` attributes (typically a ``Settings`` or merged
|
||||
runtime-settings instance from ``build_runtime_settings``).
|
||||
|
||||
Returns
|
||||
-------
|
||||
int
|
||||
Number of rows upserted (inserted or updated). Returns 0 for a no-op.
|
||||
|
||||
Notes
|
||||
-----
|
||||
- No-op (returns 0) when:
|
||||
* No active contract exists with ``kind="tibber"``, OR
|
||||
* ``settings.tibber_api_token`` is empty or whitespace.
|
||||
- Exceptions from the Tibber client are *not* caught here; they propagate to
|
||||
the caller (the scheduled job wrapper swallows them and logs).
|
||||
"""
|
||||
token: str = getattr(settings, "tibber_api_token", "") or ""
|
||||
if not token.strip():
|
||||
logger.debug("refresh_prices: tibber_api_token is empty — no-op")
|
||||
return 0
|
||||
|
||||
if not active_tibber_contract_exists(session):
|
||||
logger.debug("refresh_prices: no active tibber contract — no-op")
|
||||
return 0
|
||||
|
||||
home_id: str = getattr(settings, "tibber_home_id", "") or ""
|
||||
home_id_or_none: str | None = home_id.strip() or None
|
||||
|
||||
# Neither the API token nor the selected home identifier is safe to emit in
|
||||
# diagnostics. The fetch client receives them, but logs only describe the
|
||||
# operation itself.
|
||||
logger.info("refresh_prices: fetching Tibber price range")
|
||||
|
||||
# May raise TibberError or TibberAuthError — let them propagate.
|
||||
price_points = fetch_price_range(token, home_id_or_none)
|
||||
|
||||
if not price_points:
|
||||
logger.info("refresh_prices: API returned zero price points — nothing to upsert")
|
||||
return 0
|
||||
|
||||
fetched_at = datetime.now(UTC)
|
||||
upserted = 0
|
||||
|
||||
for point in price_points:
|
||||
existing = session.execute(
|
||||
select(TibberPrice).where(TibberPrice.starts_at == point.starts_at)
|
||||
).scalar_one_or_none()
|
||||
|
||||
if existing is not None:
|
||||
# Update in-place — price fields may change (e.g. Tibber corrects a
|
||||
# forecast), but starts_at and resolution stay the same.
|
||||
existing.total = point.total
|
||||
existing.energy = point.energy
|
||||
existing.tax = point.tax
|
||||
existing.level = point.level
|
||||
existing.currency = point.currency
|
||||
existing.fetched_at = fetched_at
|
||||
else:
|
||||
session.add(
|
||||
TibberPrice(
|
||||
starts_at=point.starts_at,
|
||||
resolution=point.resolution,
|
||||
total=point.total,
|
||||
energy=point.energy,
|
||||
tax=point.tax,
|
||||
level=point.level,
|
||||
currency=point.currency,
|
||||
fetched_at=fetched_at,
|
||||
)
|
||||
)
|
||||
|
||||
upserted += 1
|
||||
|
||||
session.commit()
|
||||
|
||||
logger.info("refresh_prices: upserted %d price points", upserted)
|
||||
return upserted
|
||||
|
||||
|
||||
def run_tibber_refresh_best_effort() -> bool:
|
||||
"""Run one refresh with an isolated session, skipping concurrent requests.
|
||||
|
||||
This is shared by the hourly scheduler and immediate post-commit triggers.
|
||||
It intentionally catches all failures: refresh is advisory and must never
|
||||
make app startup or a successfully committed configuration/contract update
|
||||
appear to have failed. The boolean reports whether this invocation owned
|
||||
the work; it is primarily useful for tests and diagnostics.
|
||||
"""
|
||||
if not _background_refresh_lock.acquire(blocking=False):
|
||||
logger.debug("Tibber refresh already running; skipping duplicate request")
|
||||
return False
|
||||
|
||||
session: Session | None = None
|
||||
try:
|
||||
# Local imports keep the pure refresh service free of app startup import
|
||||
# cycles, while every background invocation gets a fresh DB session.
|
||||
from app.config import get_settings
|
||||
from app.db import get_session_local
|
||||
from app.services.config_page import build_runtime_settings
|
||||
|
||||
session = get_session_local()()
|
||||
refresh_prices(session, build_runtime_settings(session, get_settings()))
|
||||
except Exception as exc:
|
||||
# Exception text can contain remote request details. Keep diagnostics
|
||||
# useful without allowing a token or home id to escape through logging.
|
||||
logger.warning("Tibber price refresh failed (%s)", type(exc).__name__)
|
||||
if session is not None:
|
||||
try:
|
||||
session.rollback()
|
||||
except Exception:
|
||||
logger.warning("Tibber price refresh rollback failed")
|
||||
finally:
|
||||
if session is not None:
|
||||
try:
|
||||
session.close()
|
||||
except Exception:
|
||||
logger.warning("Tibber price refresh session close failed")
|
||||
_background_refresh_lock.release()
|
||||
return True
|
||||
|
||||
|
||||
def trigger_tibber_refresh() -> None:
|
||||
"""Request a non-blocking, best-effort Tibber refresh after a DB commit."""
|
||||
try:
|
||||
Thread(
|
||||
target=run_tibber_refresh_best_effort,
|
||||
name="tibber-price-refresh",
|
||||
daemon=True,
|
||||
).start()
|
||||
except Exception as exc:
|
||||
# Starting the optional worker must not turn an already committed API
|
||||
# operation into a failure; avoid logging exception text for secrecy.
|
||||
logger.warning("Unable to start Tibber price refresh (%s)", type(exc).__name__)
|
||||
@@ -0,0 +1,98 @@
|
||||
"""Server local-timezone helpers.
|
||||
|
||||
All time is stored as UTC (no change). This module provides a single,
|
||||
monkeypatch-friendly entry point for converting UTC moments to the server's
|
||||
local timezone — used for business-day boundary calculations (standing charges,
|
||||
effective-from date interpretation, etc.).
|
||||
|
||||
Design goals
|
||||
------------
|
||||
- Zero third-party dependencies (no tzlocal, no pytz).
|
||||
- Testable without depending on CI host timezone:
|
||||
monkeypatch ``local_tz`` to a fixed ``ZoneInfo`` in unit tests.
|
||||
- DST-correct: conversions use ``astimezone(local_tz())`` per instant, so each
|
||||
moment is independently correct even during DST transitions.
|
||||
- Consistent with the existing ``.astimezone()`` pattern already used in
|
||||
``homeassistant_inbound.py`` and ``poo.py``.
|
||||
|
||||
Priority for resolving the local timezone
|
||||
-----------------------------------------
|
||||
1. ``TZ`` environment variable — ``ZoneInfo(os.environ["TZ"])``.
|
||||
Set ``TZ=Europe/Amsterdam`` in the deployment env for correct NL handling.
|
||||
2. The DST-aware ``Europe/Amsterdam`` business timezone. This keeps local-day
|
||||
calculations deterministic when a deployment does not set ``TZ``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from datetime import date, datetime, timezone
|
||||
from typing import TYPE_CHECKING
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from datetime import tzinfo
|
||||
|
||||
|
||||
def local_tz() -> "tzinfo":
|
||||
"""Return the server's local timezone.
|
||||
|
||||
Resolution order:
|
||||
1. ``TZ`` environment variable (``ZoneInfo(TZ)``). Set
|
||||
``TZ=Europe/Amsterdam`` in production for correct NL/DST handling.
|
||||
2. The DST-aware ``Europe/Amsterdam`` business timezone.
|
||||
|
||||
**Monkeypatch this function in tests** to get deterministic timezone
|
||||
behaviour regardless of CI host configuration::
|
||||
|
||||
monkeypatch.setattr("app.services.timezone.local_tz",
|
||||
lambda: ZoneInfo("Europe/Amsterdam"))
|
||||
"""
|
||||
tz_env = os.environ.get("TZ", "").strip()
|
||||
if tz_env:
|
||||
return ZoneInfo(tz_env)
|
||||
return ZoneInfo("Europe/Amsterdam")
|
||||
|
||||
|
||||
def to_local(dt: datetime) -> datetime:
|
||||
"""Convert *dt* to the server local timezone.
|
||||
|
||||
Handles three input cases:
|
||||
- timezone-aware (any zone): converted via ``astimezone``.
|
||||
- timezone-naive: assumed UTC (SQLite read-back convention), then converted.
|
||||
|
||||
Returns a timezone-aware datetime in ``local_tz()``.
|
||||
"""
|
||||
if dt.tzinfo is None:
|
||||
dt = dt.replace(tzinfo=timezone.utc)
|
||||
return dt.astimezone(local_tz())
|
||||
|
||||
|
||||
def local_now() -> datetime:
|
||||
"""Return the current moment as a timezone-aware local datetime."""
|
||||
return datetime.now(local_tz())
|
||||
|
||||
|
||||
def local_date(dt: datetime) -> date:
|
||||
"""Return the local calendar date for *dt*.
|
||||
|
||||
Equivalent to ``to_local(dt).date()``. A distinct helper to make
|
||||
call sites readable.
|
||||
"""
|
||||
return to_local(dt).date()
|
||||
|
||||
|
||||
def local_midnight_utc(d: date) -> datetime:
|
||||
"""Return the UTC instant that corresponds to local midnight on *d*.
|
||||
|
||||
Example (Europe/Amsterdam, CEST=UTC+2):
|
||||
local_midnight_utc(date(2026, 6, 25))
|
||||
→ datetime(2026, 6, 24, 22, 0, 0, tzinfo=UTC)
|
||||
|
||||
This is used to convert a local calendar date back to a UTC boundary, e.g.
|
||||
to compute the start/end of a local day for DB queries.
|
||||
"""
|
||||
tz = local_tz()
|
||||
# Build a naive local datetime at midnight, then localize and convert to UTC.
|
||||
local_midnight = datetime(d.year, d.month, d.day, 0, 0, 0, tzinfo=tz)
|
||||
return local_midnight.astimezone(timezone.utc)
|
||||
@@ -0,0 +1,234 @@
|
||||
"""TOTP service: secret management, recovery-code lifecycle, setup/enable/disable (M4-T05).
|
||||
|
||||
State-machine summary
|
||||
---------------------
|
||||
DISABLED (default)
|
||||
totp_secret=None, totp_enabled=False, no recovery codes
|
||||
|
||||
PENDING (after setup, before enable)
|
||||
totp_secret=<base32>, totp_enabled=False, recovery codes stored as hashes
|
||||
- Re-calling setup replaces the pending secret and regenerates recovery codes.
|
||||
|
||||
ENABLED (after enable)
|
||||
totp_secret=<base32>, totp_enabled=True, recovery codes stored as hashes
|
||||
|
||||
After disable:
|
||||
totp_secret=None, totp_enabled=False, recovery codes deleted → back to DISABLED
|
||||
|
||||
|
||||
Recovery-code timing
|
||||
--------------------
|
||||
1. setup → generate 10 plaintext codes, persist their Argon2 hashes immediately,
|
||||
return plaintext to caller (ONLY time).
|
||||
2. enable → verify 6-digit TOTP code; if ok, flip totp_enabled=True.
|
||||
Recovery codes are already in the DB from step 1.
|
||||
3. disable → clear secret, clear enabled flag, delete all recovery codes.
|
||||
|
||||
Security note: ``totp_secret`` is stored as plaintext (same posture as other
|
||||
secrets in this project: protected by file-system permissions on the SQLite
|
||||
database). Recovery codes are stored as Argon2 hashes only.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import secrets
|
||||
import string
|
||||
|
||||
import pyotp
|
||||
from sqlalchemy import delete, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models.auth import AuthUser, RecoveryCode
|
||||
from app.services.auth import hash_password, verify_password
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Number of recovery codes to generate per setup
|
||||
_RECOVERY_CODE_COUNT = 10
|
||||
# Alphabet for each 4-character segment: lowercase letters + digits, no ambiguous chars
|
||||
_CODE_ALPHABET = string.ascii_lowercase + string.digits
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _generate_recovery_code() -> str:
|
||||
"""Return a single recovery code in ``xxxx-xxxx`` format."""
|
||||
segment = lambda: "".join(secrets.choice(_CODE_ALPHABET) for _ in range(4)) # noqa: E731
|
||||
return f"{segment()}-{segment()}"
|
||||
|
||||
|
||||
def _verify_totp_code(secret: str, code: str) -> bool:
|
||||
"""Verify a 6-digit TOTP code against the given base-32 secret (±1 window)."""
|
||||
return pyotp.TOTP(secret).verify(code, valid_window=1)
|
||||
|
||||
|
||||
def _delete_recovery_codes(db: Session, *, user_id: int) -> None:
|
||||
"""Delete ALL recovery codes for a user (called on setup re-run and disable)."""
|
||||
db.execute(delete(RecoveryCode).where(RecoveryCode.user_id == user_id))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public API
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def setup(
|
||||
db: Session,
|
||||
*,
|
||||
user: AuthUser,
|
||||
issuer: str,
|
||||
) -> tuple[str, str, list[str]]:
|
||||
"""Generate a new pending TOTP secret and recovery codes.
|
||||
|
||||
The secret is stored in ``user.totp_secret`` (``totp_enabled`` stays
|
||||
``False``). Any pre-existing pending secret and recovery codes are
|
||||
replaced atomically.
|
||||
|
||||
Returns
|
||||
-------
|
||||
(secret, otpauth_uri, plaintext_recovery_codes)
|
||||
|
||||
``plaintext_recovery_codes`` are returned **once** here and MUST NOT be
|
||||
stored or returned elsewhere.
|
||||
"""
|
||||
# Generate new TOTP secret
|
||||
new_secret = pyotp.random_base32()
|
||||
otpauth_uri = pyotp.TOTP(new_secret).provisioning_uri(
|
||||
name=user.username,
|
||||
issuer_name=issuer,
|
||||
)
|
||||
|
||||
# Generate plaintext recovery codes
|
||||
plaintext_codes = [_generate_recovery_code() for _ in range(_RECOVERY_CODE_COUNT)]
|
||||
|
||||
# --- Persist atomically ---
|
||||
# 1. Delete any old pending recovery codes (idempotent on re-setup)
|
||||
assert user.id is not None # mypy guard; always set after DB insertion
|
||||
_delete_recovery_codes(db, user_id=user.id)
|
||||
|
||||
# 2. Update the secret (pending — totp_enabled stays False)
|
||||
user.totp_secret = new_secret
|
||||
|
||||
# 3. Persist new recovery codes as Argon2 hashes
|
||||
for plaintext in plaintext_codes:
|
||||
db.add(RecoveryCode(user_id=user.id, code_hash=hash_password(plaintext)))
|
||||
|
||||
db.commit()
|
||||
db.refresh(user)
|
||||
|
||||
logger.info("TOTP setup (pending) for user '%s'; %d recovery codes generated.", user.username, _RECOVERY_CODE_COUNT)
|
||||
return new_secret, otpauth_uri, plaintext_codes
|
||||
|
||||
|
||||
def enable(db: Session, *, user: AuthUser, code: str) -> bool:
|
||||
"""Enable TOTP after the user proves they can generate the correct code.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
code:
|
||||
The current 6-digit TOTP code from the user's authenticator app.
|
||||
|
||||
Returns
|
||||
-------
|
||||
True on success; False if the code is invalid or no secret is pending.
|
||||
"""
|
||||
if not user.totp_secret:
|
||||
logger.info("TOTP enable rejected for '%s': no pending secret.", user.username)
|
||||
return False
|
||||
|
||||
if not _verify_totp_code(user.totp_secret, code):
|
||||
logger.info("TOTP enable rejected for '%s': wrong code.", user.username)
|
||||
return False
|
||||
|
||||
user.totp_enabled = True
|
||||
db.commit()
|
||||
db.refresh(user)
|
||||
logger.info("TOTP enabled for user '%s'.", user.username)
|
||||
return True
|
||||
|
||||
|
||||
def disable(
|
||||
db: Session,
|
||||
*,
|
||||
user: AuthUser,
|
||||
password: str | None = None,
|
||||
code: str | None = None,
|
||||
) -> bool:
|
||||
"""Disable TOTP. Caller must supply exactly one of ``password`` or ``code``.
|
||||
|
||||
On success: ``totp_enabled=False``, ``totp_secret=None``, all recovery
|
||||
codes deleted.
|
||||
|
||||
Returns
|
||||
-------
|
||||
True on success; False if neither credential is valid.
|
||||
"""
|
||||
if not password and not code:
|
||||
logger.info("TOTP disable rejected for '%s': no credential provided.", user.username)
|
||||
return False
|
||||
|
||||
if password:
|
||||
if not verify_password(password, user.password_hash):
|
||||
logger.info("TOTP disable rejected for '%s': wrong password.", user.username)
|
||||
return False
|
||||
elif code:
|
||||
if not user.totp_secret or not _verify_totp_code(user.totp_secret, code):
|
||||
logger.info("TOTP disable rejected for '%s': wrong TOTP code.", user.username)
|
||||
return False
|
||||
|
||||
# Clear TOTP state
|
||||
assert user.id is not None
|
||||
_delete_recovery_codes(db, user_id=user.id)
|
||||
user.totp_enabled = False
|
||||
user.totp_secret = None
|
||||
db.commit()
|
||||
db.refresh(user)
|
||||
logger.info("TOTP disabled for user '%s'.", user.username)
|
||||
return True
|
||||
|
||||
|
||||
def verify_totp_code(user: AuthUser, code: str) -> bool:
|
||||
"""Verify a 6-digit TOTP code for a user.
|
||||
|
||||
Returns True if the code is valid (±1 time window), False otherwise.
|
||||
The user must have a ``totp_secret`` set; returns False if not.
|
||||
|
||||
This is the public entry-point for T06 login two-factor verification.
|
||||
"""
|
||||
if not user.totp_secret:
|
||||
return False
|
||||
return _verify_totp_code(user.totp_secret, code)
|
||||
|
||||
|
||||
def verify_recovery_code(db: Session, *, user: AuthUser, code: str) -> bool:
|
||||
"""Verify and consume a one-time recovery code.
|
||||
|
||||
Finds the first unused recovery code whose hash matches ``code``, marks it
|
||||
as consumed (sets ``used_at``), and returns ``True``. Returns ``False`` if
|
||||
no matching unused code exists.
|
||||
|
||||
This function is provided for T06 (login two-factor) but lives here so the
|
||||
TOTP service owns all recovery-code logic.
|
||||
"""
|
||||
from datetime import UTC, datetime
|
||||
|
||||
unused = db.execute(
|
||||
select(RecoveryCode).where(
|
||||
RecoveryCode.user_id == user.id,
|
||||
RecoveryCode.used_at.is_(None),
|
||||
)
|
||||
).scalars().all()
|
||||
|
||||
for rc in unused:
|
||||
if verify_password(code, rc.code_hash):
|
||||
rc.used_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
logger.info(
|
||||
"Recovery code consumed for user '%s' (id=%d).", user.username, rc.id
|
||||
)
|
||||
return True
|
||||
|
||||
return False
|
||||
@@ -0,0 +1,401 @@
|
||||
"""Privacy-preserving WarmteLink frame admission and minute sampling.
|
||||
|
||||
The serial worker added later owns I/O. This module deliberately only accepts
|
||||
already parsed :class:`P1Telegram` instances (or a parser callable at its
|
||||
small convenience entry point), so rejected telegram bytes never enter the
|
||||
database or an exception message.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from decimal import Decimal
|
||||
import re
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.db import get_session_local
|
||||
from app.integrations.p1 import IntegrityStatus, P1Channel, P1Telegram, parse_telegram
|
||||
from app.models.meter_source import MeterSource, WarmteLinkReading
|
||||
from app.services.meter_sources import upsert_discovered_channel
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ChannelSample:
|
||||
key: str
|
||||
label: str
|
||||
value: Decimal
|
||||
unit: str
|
||||
device_type: str | None
|
||||
fingerprint: str | None
|
||||
identity: tuple[object, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _FrameSnapshot:
|
||||
recorded_at: datetime
|
||||
fingerprint: str | None
|
||||
samples: tuple[_ChannelSample, ...]
|
||||
|
||||
|
||||
class WarmteLinkIngestor:
|
||||
"""Keep unverifiable candidates isolated by source for one worker lifetime."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
clock: Callable[[], datetime] | None = None,
|
||||
) -> None:
|
||||
self._clock = clock or (lambda: datetime.now(UTC))
|
||||
self._previous: dict[int, _FrameSnapshot] = {}
|
||||
|
||||
def ingest(
|
||||
self,
|
||||
session: Session,
|
||||
*,
|
||||
source_id: int,
|
||||
telegram: P1Telegram,
|
||||
received_at: datetime | None = None,
|
||||
) -> bool:
|
||||
"""Apply one parsed telegram in the caller's transaction.
|
||||
|
||||
Returns whether the frame was admitted. Callers that own a session
|
||||
must commit on success; :func:`handle_frame` is the failure-contained
|
||||
worker-facing entry point.
|
||||
"""
|
||||
source = session.get(MeterSource, source_id)
|
||||
if source is None:
|
||||
raise ValueError("WarmteLink source was not found")
|
||||
if source.kind != "warmtelink_serial":
|
||||
raise ValueError("Source is not a WarmteLink serial source")
|
||||
|
||||
received_at = _utc(self._clock()) if received_at is None else _utc(received_at)
|
||||
try:
|
||||
integrity = _integrity_status(telegram.integrity)
|
||||
except Exception:
|
||||
self._previous.pop(source_id, None)
|
||||
self._diagnose(source, "WarmteLink frame could not be normalized", now=received_at)
|
||||
return False
|
||||
|
||||
if integrity is IntegrityStatus.INVALID:
|
||||
self._previous.pop(source_id, None)
|
||||
self._diagnose(source, "WarmteLink frame checksum is invalid", now=received_at)
|
||||
return False
|
||||
|
||||
try:
|
||||
snapshot = _snapshot(telegram, received_at)
|
||||
fingerprints = _final_fingerprints(snapshot)
|
||||
except Exception:
|
||||
# A parser DTO remains an untrusted boundary. Do not let an
|
||||
# unnormalizable DTO bridge two otherwise matching candidates.
|
||||
self._previous.pop(source_id, None)
|
||||
self._diagnose(source, "WarmteLink frame could not be normalized", now=received_at)
|
||||
return False
|
||||
if fingerprints is None:
|
||||
self._previous.pop(source_id, None)
|
||||
self._diagnose(source, "WarmteLink frame fingerprint is invalid", now=received_at)
|
||||
return False
|
||||
if integrity is IntegrityStatus.UNVERIFIABLE:
|
||||
previous = self._previous.get(source_id)
|
||||
if previous is None:
|
||||
self._previous[source_id] = snapshot
|
||||
self._diagnose(
|
||||
source, "Awaiting a second matching unverifiable WarmteLink frame", now=received_at
|
||||
)
|
||||
return False
|
||||
reason = _continuity_problem(previous, snapshot)
|
||||
if reason is not None:
|
||||
self._previous[source_id] = snapshot
|
||||
self._diagnose(source, reason, now=received_at)
|
||||
return False
|
||||
try:
|
||||
self._admit(session, source, snapshot, integrity.value, fingerprints, received_at)
|
||||
except Exception:
|
||||
# ``ingest`` owns flushes and may be used directly by tests or
|
||||
# future callers. A failed write must never become a speculative
|
||||
# predecessor for this source.
|
||||
self._previous.pop(source_id, None)
|
||||
raise
|
||||
if integrity is IntegrityStatus.UNVERIFIABLE:
|
||||
# Advance the sliding predecessor only once every database write
|
||||
# for this frame has succeeded. ``handle_frame`` also clears it
|
||||
# if the caller's later commit fails.
|
||||
self._previous[source_id] = snapshot
|
||||
else:
|
||||
# A verified frame has no need for a speculative predecessor.
|
||||
self._previous.pop(source_id, None)
|
||||
return True
|
||||
|
||||
def handle_frame(
|
||||
self,
|
||||
source_id: int,
|
||||
frame: bytes,
|
||||
*,
|
||||
session_factory: Callable[[], Session] = get_session_local,
|
||||
parser: Callable[[bytes], P1Telegram] = parse_telegram,
|
||||
) -> bool:
|
||||
"""Parse and persist one frame, containing both parse and DB failures.
|
||||
|
||||
A failed write is rolled back before a fresh transaction records only
|
||||
a generic source error. Consequently no partial latest/history update
|
||||
survives and the following frame may recover normally.
|
||||
"""
|
||||
received_at = _utc(self._clock())
|
||||
try:
|
||||
telegram = parser(frame)
|
||||
except Exception:
|
||||
self._previous.pop(source_id, None)
|
||||
self._record_error(
|
||||
session_factory, source_id, "WarmteLink frame could not be parsed", now=received_at
|
||||
)
|
||||
return False
|
||||
try:
|
||||
with session_factory() as session:
|
||||
admitted = self.ingest(
|
||||
session, source_id=source_id, telegram=telegram, received_at=received_at
|
||||
)
|
||||
session.commit()
|
||||
return admitted
|
||||
except Exception:
|
||||
self._previous.pop(source_id, None)
|
||||
self._record_error(session_factory, source_id, "WarmteLink ingest failed", now=received_at)
|
||||
return False
|
||||
|
||||
def _admit(
|
||||
self,
|
||||
session: Session,
|
||||
source: MeterSource,
|
||||
snapshot: _FrameSnapshot,
|
||||
quality: str,
|
||||
fingerprints: tuple[str, ...],
|
||||
received_at: datetime,
|
||||
) -> None:
|
||||
for sample, fingerprint in zip(snapshot.samples, fingerprints, strict=True):
|
||||
channel = upsert_discovered_channel(
|
||||
session,
|
||||
source_id=source.id,
|
||||
channel_key=sample.key,
|
||||
label=sample.label,
|
||||
unit=sample.unit,
|
||||
suggested_commodity=_suggestion(sample.unit),
|
||||
device_type=sample.device_type,
|
||||
fingerprint=fingerprint,
|
||||
latest_value=sample.value,
|
||||
latest_at=snapshot.recorded_at,
|
||||
latest_quality=quality,
|
||||
)
|
||||
session.flush()
|
||||
bucket = snapshot.recorded_at.replace(second=0, microsecond=0)
|
||||
exists = session.scalar(
|
||||
select(WarmteLinkReading.id)
|
||||
.where(
|
||||
WarmteLinkReading.channel_id == channel.id,
|
||||
WarmteLinkReading.recorded_at >= bucket,
|
||||
WarmteLinkReading.recorded_at < bucket + timedelta(minutes=1),
|
||||
)
|
||||
.limit(1)
|
||||
)
|
||||
if exists is None:
|
||||
session.add(
|
||||
WarmteLinkReading(
|
||||
channel_id=channel.id,
|
||||
recorded_at=snapshot.recorded_at,
|
||||
received_at=received_at,
|
||||
value=sample.value,
|
||||
unit=sample.unit,
|
||||
quality=quality,
|
||||
equipment_fingerprint=fingerprint,
|
||||
)
|
||||
)
|
||||
source.status = "online"
|
||||
source.last_seen_at = received_at
|
||||
source.last_error = None
|
||||
source.updated_at = received_at
|
||||
|
||||
def _diagnose(self, source: MeterSource, reason: str, *, now: datetime | None = None) -> None:
|
||||
now = _utc(self._clock()) if now is None else now
|
||||
source.status = "error"
|
||||
source.last_error = reason
|
||||
source.updated_at = now
|
||||
|
||||
def _record_error(
|
||||
self,
|
||||
session_factory: Callable[[], Session],
|
||||
source_id: int,
|
||||
message: str,
|
||||
*,
|
||||
now: datetime,
|
||||
) -> None:
|
||||
try:
|
||||
with session_factory() as session:
|
||||
source = session.get(MeterSource, source_id)
|
||||
if source is not None:
|
||||
self._diagnose(source, message, now=now)
|
||||
session.commit()
|
||||
except Exception:
|
||||
# Error reporting itself must not kill another source worker.
|
||||
return
|
||||
|
||||
|
||||
def _snapshot(telegram: P1Telegram, received_at: datetime) -> _FrameSnapshot:
|
||||
recorded_at = _parse_timestamp(telegram.timestamp, received_at)
|
||||
samples = tuple(_sample(channel) for channel in telegram.channels)
|
||||
if not samples:
|
||||
raise ValueError("WarmteLink telegram contains no cumulative channels")
|
||||
return _FrameSnapshot(recorded_at, telegram.equipment_fingerprint, samples)
|
||||
|
||||
|
||||
def _sample(channel: P1Channel) -> _ChannelSample:
|
||||
profile = _canonical_channel_profile(channel.number)
|
||||
if channel.device_type != profile.device_type:
|
||||
raise ValueError("WarmteLink channel device type is not canonical")
|
||||
if len(channel.readings) != 1:
|
||||
raise ValueError("WarmteLink channel has no unambiguous cumulative reading")
|
||||
reading = channel.readings[0]
|
||||
if reading.code != profile.reading_code or reading.value is None or reading.unit != profile.raw_unit:
|
||||
raise ValueError("WarmteLink cumulative reading is incomplete")
|
||||
# Parser annotations are not a trust boundary: fake or future parser DTOs
|
||||
# must not put a float (including NaN) or a non-finite Decimal into the
|
||||
# per-source unverifiable candidate state. Do not coerce here: accepting
|
||||
# another numeric type would make the persistence and continuity paths
|
||||
# disagree about the cumulative-value contract.
|
||||
if not isinstance(reading.value, Decimal) or not reading.value.is_finite():
|
||||
raise ValueError("WarmteLink cumulative reading is not a finite Decimal")
|
||||
fingerprint = channel.equipment_fingerprint
|
||||
return _ChannelSample(
|
||||
key=profile.key,
|
||||
label=profile.label,
|
||||
value=reading.value,
|
||||
unit=profile.unit,
|
||||
device_type=profile.device_type,
|
||||
fingerprint=fingerprint,
|
||||
identity=(channel.number, profile.device_type, fingerprint, profile.reading_code, profile.unit),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _CanonicalChannelProfile:
|
||||
key: str
|
||||
label: str
|
||||
device_type: str
|
||||
reading_code: str
|
||||
raw_unit: str
|
||||
unit: str
|
||||
|
||||
|
||||
_CANONICAL_CHANNELS = {
|
||||
1: _CanonicalChannelProfile(
|
||||
key="channel-1",
|
||||
label="WarmteLink channel 1",
|
||||
device_type="006",
|
||||
reading_code="0-1:24.2.1",
|
||||
raw_unit="m3",
|
||||
unit="m³",
|
||||
),
|
||||
2: _CanonicalChannelProfile(
|
||||
key="channel-2",
|
||||
label="WarmteLink channel 2",
|
||||
device_type="012",
|
||||
reading_code="0-2:24.2.1",
|
||||
raw_unit="GJ",
|
||||
unit="GJ",
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def _canonical_channel_profile(number: object) -> _CanonicalChannelProfile:
|
||||
# Do not format or coerce parser supplied channel numbers: that could turn
|
||||
# an arbitrary object into an identity key before it is rejected.
|
||||
if type(number) is not int:
|
||||
raise ValueError("WarmteLink channel number is not canonical")
|
||||
try:
|
||||
return _CANONICAL_CHANNELS[number]
|
||||
except KeyError as exc:
|
||||
raise ValueError("WarmteLink channel is not supported by the profile") from exc
|
||||
|
||||
|
||||
def _integrity_status(value: object) -> IntegrityStatus:
|
||||
"""Require a real parser integrity enum, never a look-alike value object."""
|
||||
if not isinstance(value, IntegrityStatus):
|
||||
raise ValueError("WarmteLink integrity status is not canonical")
|
||||
return value
|
||||
|
||||
|
||||
def _continuity_problem(previous: _FrameSnapshot, current: _FrameSnapshot) -> str | None:
|
||||
if current.recorded_at <= previous.recorded_at:
|
||||
return "Unverifiable WarmteLink timestamp is not strictly increasing"
|
||||
if current.recorded_at - previous.recorded_at != timedelta(seconds=10):
|
||||
return "Unverifiable WarmteLink frame cadence is not 10 seconds"
|
||||
if current.fingerprint != previous.fingerprint:
|
||||
return "WarmteLink equipment metadata changed"
|
||||
if tuple(sample.identity for sample in current.samples) != tuple(sample.identity for sample in previous.samples):
|
||||
return "WarmteLink channel metadata changed or channel set changed"
|
||||
old_values = {sample.key: sample.value for sample in previous.samples}
|
||||
if any(sample.value < old_values[sample.key] for sample in current.samples):
|
||||
return "WarmteLink cumulative value decreased"
|
||||
return None
|
||||
|
||||
|
||||
_FINGERPRINT_PATTERN = re.compile(r"[0-9a-f]{64}")
|
||||
|
||||
|
||||
def _final_fingerprints(snapshot: _FrameSnapshot) -> tuple[str, ...] | None:
|
||||
"""Return only canonical SHA-256 hexdigests safe to persist.
|
||||
|
||||
A parser DTO is an untrusted boundary: even a field named ``fingerprint``
|
||||
can contain a raw equipment identifier. Validate the complete DTO before
|
||||
choosing persisted values: the top-level value and every channel value
|
||||
must independently be canonical hashes. This deliberately does not use
|
||||
a top-level fallback for an absent or malformed channel fingerprint.
|
||||
"""
|
||||
values = (snapshot.fingerprint, *(sample.fingerprint for sample in snapshot.samples))
|
||||
if any(not _is_canonical_fingerprint(value) for value in values):
|
||||
return None
|
||||
return tuple(sample.fingerprint for sample in snapshot.samples if sample.fingerprint is not None)
|
||||
|
||||
|
||||
def _is_canonical_fingerprint(value: str | None) -> bool:
|
||||
return value is not None and _FINGERPRINT_PATTERN.fullmatch(value) is not None
|
||||
|
||||
|
||||
_AMSTERDAM = ZoneInfo("Europe/Amsterdam")
|
||||
_MAX_CLOCK_SKEW = timedelta(minutes=5)
|
||||
|
||||
|
||||
def _parse_timestamp(value: str | None, received_at: datetime) -> datetime:
|
||||
if value is None or len(value) != 13 or value[-1] not in {"S", "W"} or not value[:-1].isdigit():
|
||||
raise ValueError("WarmteLink timestamp is unavailable")
|
||||
naive = datetime.strptime(value[:-1], "%y%m%d%H%M%S")
|
||||
# The S/W marker is only advisory: deployed devices have emitted W while
|
||||
# on CEST. Validate both folds by a UTC round trip, which rejects spring
|
||||
# gaps and leaves one (ordinary) or two (fall-back) real instants.
|
||||
candidates: list[datetime] = []
|
||||
for fold in (0, 1):
|
||||
candidate = naive.replace(tzinfo=_AMSTERDAM, fold=fold).astimezone(UTC)
|
||||
local = candidate.astimezone(_AMSTERDAM)
|
||||
if local.replace(tzinfo=None) == naive and local.fold == fold and candidate not in candidates:
|
||||
candidates.append(candidate)
|
||||
if not candidates:
|
||||
raise ValueError("WarmteLink timestamp is unavailable")
|
||||
received_at = _utc(received_at)
|
||||
recorded_at = min(candidates, key=lambda candidate: abs(candidate - received_at))
|
||||
if abs(recorded_at - received_at) > _MAX_CLOCK_SKEW:
|
||||
raise ValueError("WarmteLink timestamp is unavailable")
|
||||
return recorded_at
|
||||
|
||||
|
||||
def _suggestion(unit: str) -> str | None:
|
||||
return {"GJ": "heating", "m³": "hot_water"}.get(unit)
|
||||
|
||||
|
||||
def _canonical_unit(unit: str) -> str:
|
||||
"""Map the P1 spelling of cubic metres to the source-profile unit."""
|
||||
return "m³" if unit == "m3" else unit
|
||||
|
||||
|
||||
def _utc(value: datetime) -> datetime:
|
||||
return value.replace(tzinfo=UTC) if value.tzinfo is None else value.astimezone(UTC)
|
||||
@@ -0,0 +1,436 @@
|
||||
"""Read-only WarmteLink serial workers and their lifecycle manager.
|
||||
|
||||
The worker deliberately owns no long-lived SQLAlchemy session and never
|
||||
retains telegram bytes after handing a complete frame to the ingestor.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime, timedelta
|
||||
import logging
|
||||
import threading
|
||||
from typing import Protocol
|
||||
|
||||
import serial
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.db import get_session_local
|
||||
from app.integrations.p1 import TelegramFramer
|
||||
from app.models.meter_source import MeterSource
|
||||
from app.services.warmtelink_ingest import WarmteLinkIngestor
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_BACKOFF_SECONDS = (1, 2, 4, 8, 16, 32, 60)
|
||||
_JOIN_TIMEOUT_SECONDS = 5
|
||||
_DISCOVERY_LOCK_TIMEOUT_SECONDS = 0.05
|
||||
_DISCOVERY_WAIT_SECONDS = 0.1
|
||||
_DISCOVERY_TIMEOUT_SECONDS = 5
|
||||
|
||||
|
||||
class ReadOnlySerial(Protocol):
|
||||
def read(self, size: int = 1) -> bytes: ...
|
||||
|
||||
def close(self) -> None: ...
|
||||
|
||||
|
||||
SerialFactory = Callable[[dict], ReadOnlySerial]
|
||||
SessionFactory = Callable[[], Session]
|
||||
|
||||
|
||||
def _default_session_factory() -> Session:
|
||||
"""Resolve the cached sessionmaker at call time, then open one session."""
|
||||
return get_session_local()()
|
||||
|
||||
|
||||
class WorkerClock(Protocol):
|
||||
"""Injectable interruptible clock, keeping retry tests deterministic."""
|
||||
|
||||
def wait(self, stop_event: threading.Event, seconds: float) -> bool: ...
|
||||
|
||||
|
||||
class _EventClock:
|
||||
def wait(self, stop_event: threading.Event, seconds: float) -> bool:
|
||||
return stop_event.wait(seconds)
|
||||
|
||||
|
||||
def open_warmtelink_serial(config: dict) -> ReadOnlySerial:
|
||||
"""Open the fixed WarmteLink P1 profile; no write-capable API is exposed."""
|
||||
return serial.Serial(
|
||||
port=config["path"], baudrate=115200, bytesize=serial.SEVENBITS,
|
||||
parity=serial.PARITY_NONE, stopbits=serial.STOPBITS_ONE, timeout=1,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _WorkerConfig:
|
||||
source_id: int
|
||||
config: dict
|
||||
|
||||
|
||||
@dataclass
|
||||
class DiscoveryRequest:
|
||||
"""One source-scoped request, completed only by its serial owner."""
|
||||
|
||||
request_id: int
|
||||
source_id: int
|
||||
deadline: datetime
|
||||
status: str = "pending"
|
||||
detail: str | None = None
|
||||
completed: threading.Event = field(default_factory=threading.Event)
|
||||
|
||||
|
||||
class WarmteLinkWorker:
|
||||
"""One interruptible, read-only serial loop for one meter source."""
|
||||
|
||||
def __init__(
|
||||
self, source_id: int, config: dict, *, session_factory: SessionFactory = _default_session_factory,
|
||||
serial_factory: SerialFactory = open_warmtelink_serial,
|
||||
stop_event: threading.Event | None = None,
|
||||
ingestor: WarmteLinkIngestor | None = None,
|
||||
clock: WorkerClock | None = None,
|
||||
) -> None:
|
||||
self.source_id = source_id
|
||||
self.config = dict(config)
|
||||
self._session_factory = session_factory
|
||||
self._serial_factory = serial_factory
|
||||
self._stop_event = stop_event or threading.Event()
|
||||
self._ingestor = ingestor or WarmteLinkIngestor()
|
||||
self._clock = clock or _EventClock()
|
||||
self._serial: ReadOnlySerial | None = None
|
||||
self._serial_lock = threading.Lock()
|
||||
self._discovery_lock = threading.Lock()
|
||||
self._discoveries: list[DiscoveryRequest] = []
|
||||
# Never inherit a daemon flag from a caller's background thread: a serial
|
||||
# descriptor and its orderly shutdown must remain visible to the process.
|
||||
self._thread = threading.Thread(
|
||||
target=self._run, name=f"warmtelink-{source_id}", daemon=False
|
||||
)
|
||||
|
||||
@property
|
||||
def thread(self) -> threading.Thread:
|
||||
return self._thread
|
||||
|
||||
def start(self) -> None:
|
||||
self._thread.start()
|
||||
|
||||
def stop(self) -> None:
|
||||
self._stop_event.set()
|
||||
self._close_serial()
|
||||
|
||||
def request_discovery(self, request: DiscoveryRequest) -> None:
|
||||
"""Queue a read request; this worker remains the sole serial owner."""
|
||||
with self._discovery_lock:
|
||||
self._discoveries.append(request)
|
||||
|
||||
def _finish_discoveries(self, status: str, detail: str | None = None) -> None:
|
||||
now = datetime.now(UTC)
|
||||
with self._discovery_lock:
|
||||
pending, self._discoveries = self._discoveries, []
|
||||
for request in pending:
|
||||
if request.completed.is_set():
|
||||
continue
|
||||
if request.deadline <= now and status == "completed":
|
||||
request.status, request.detail = "error", "Discovery timed out."
|
||||
else:
|
||||
request.status, request.detail = status, detail
|
||||
request.completed.set()
|
||||
|
||||
def _expire_discoveries(self) -> None:
|
||||
now = datetime.now(UTC)
|
||||
with self._discovery_lock:
|
||||
expired = [request for request in self._discoveries if request.deadline <= now]
|
||||
self._discoveries = [request for request in self._discoveries if request.deadline > now]
|
||||
for request in expired:
|
||||
if request.completed.is_set():
|
||||
continue
|
||||
request.status, request.detail = "error", "Discovery timed out."
|
||||
request.completed.set()
|
||||
|
||||
def join(self, timeout: float = _JOIN_TIMEOUT_SECONDS) -> bool:
|
||||
self._thread.join(timeout)
|
||||
return not self._thread.is_alive()
|
||||
|
||||
def _close_serial(self) -> None:
|
||||
with self._serial_lock:
|
||||
device, self._serial = self._serial, None
|
||||
if device is not None:
|
||||
with suppress(Exception):
|
||||
device.close()
|
||||
|
||||
def _record_error(self, message: str) -> None:
|
||||
try:
|
||||
with self._session_factory() as session:
|
||||
source = session.get(MeterSource, self.source_id)
|
||||
if source is not None:
|
||||
source.status = "error"
|
||||
source.last_error = message
|
||||
source.updated_at = datetime.now(UTC)
|
||||
session.commit()
|
||||
except Exception:
|
||||
# A source-status failure must not end another source's worker.
|
||||
return
|
||||
|
||||
def _run(self) -> None:
|
||||
backoff_index = 0
|
||||
while not self._stop_event.is_set():
|
||||
try:
|
||||
device = self._serial_factory(self.config)
|
||||
with self._serial_lock:
|
||||
if self._stop_event.is_set():
|
||||
with suppress(Exception):
|
||||
device.close()
|
||||
return
|
||||
self._serial = device
|
||||
# A disconnect makes any bytes buffered from the previous
|
||||
# descriptor untrustworthy. In particular, never let a
|
||||
# trailing partial telegram be completed by a newly opened
|
||||
# device.
|
||||
framer = TelegramFramer()
|
||||
while not self._stop_event.is_set():
|
||||
self._expire_discoveries()
|
||||
chunk = device.read(1024)
|
||||
if not chunk:
|
||||
# ``timeout`` reads are normal, but still yield so a bad
|
||||
# fake/device cannot turn an empty read into a busy spin.
|
||||
self._clock.wait(self._stop_event, 0.05)
|
||||
continue
|
||||
frames = framer.feed(chunk)
|
||||
for frame in frames:
|
||||
if self._stop_event.is_set():
|
||||
break
|
||||
admitted = self._ingestor.handle_frame(
|
||||
self.source_id, frame, session_factory=self._session_factory
|
||||
)
|
||||
if admitted:
|
||||
self._finish_discoveries("completed")
|
||||
# A complete frame proves transport recovery even if its
|
||||
# contents are rejected by the privacy/admission layer.
|
||||
backoff_index = 0
|
||||
except Exception:
|
||||
self._finish_discoveries("error", "WarmteLink discovery failed.")
|
||||
self._record_error("WarmteLink serial connection failed")
|
||||
delay = _BACKOFF_SECONDS[min(backoff_index, len(_BACKOFF_SECONDS) - 1)]
|
||||
backoff_index += 1
|
||||
self._clock.wait(self._stop_event, delay)
|
||||
finally:
|
||||
self._close_serial()
|
||||
|
||||
|
||||
class WarmteLinkWorkerManager:
|
||||
"""Reconcile enabled serial sources into exactly one worker each."""
|
||||
|
||||
def __init__(
|
||||
self, *, session_factory: SessionFactory = _default_session_factory,
|
||||
serial_factory: SerialFactory = open_warmtelink_serial,
|
||||
worker_factory: Callable[..., WarmteLinkWorker] = WarmteLinkWorker,
|
||||
) -> None:
|
||||
self._session_factory = session_factory
|
||||
self._serial_factory = serial_factory
|
||||
self._worker_factory = worker_factory
|
||||
self._workers: dict[int, tuple[_WorkerConfig, WarmteLinkWorker]] = {}
|
||||
self._lock = threading.Lock()
|
||||
self._reapers: set[int] = set()
|
||||
self._next_discovery_id = 0
|
||||
self._shutting_down = False
|
||||
|
||||
@property
|
||||
def worker_count(self) -> int:
|
||||
with self._lock:
|
||||
return len(self._workers)
|
||||
|
||||
def reconcile(self) -> None:
|
||||
# Reading desired state under the same lock which applies it prevents a
|
||||
# delayed pre-commit snapshot from rolling a newer commit backwards.
|
||||
with self._lock:
|
||||
if self._shutting_down:
|
||||
return
|
||||
desired = self._read_desired()
|
||||
self._reconcile_locked(desired)
|
||||
|
||||
def start(self) -> None:
|
||||
"""Enable reconciliation for a newly entered application lifespan."""
|
||||
with self._lock:
|
||||
self._shutting_down = False
|
||||
self.reconcile()
|
||||
|
||||
def request_discovery(self, source_id: int) -> DiscoveryRequest:
|
||||
"""Ask the current source worker for one bounded read/discovery attempt.
|
||||
|
||||
This intentionally does not reconcile or open a descriptor. Lifecycle
|
||||
convergence remains separate; a request can neither replace nor stop a
|
||||
worker when an HTTP client times out or disconnects.
|
||||
"""
|
||||
now = datetime.now(UTC)
|
||||
request = DiscoveryRequest(0, source_id, now)
|
||||
if not self._lock.acquire(timeout=_DISCOVERY_LOCK_TIMEOUT_SECONDS):
|
||||
request.status, request.detail = "error", "Discovery queue is busy."
|
||||
request.completed.set()
|
||||
return request
|
||||
try:
|
||||
self._next_discovery_id += 1
|
||||
request.request_id = self._next_discovery_id
|
||||
request.deadline = now + timedelta(seconds=_DISCOVERY_TIMEOUT_SECONDS)
|
||||
if self._shutting_down:
|
||||
request.status, request.detail = "error", "WarmteLink manager is stopped."
|
||||
request.completed.set()
|
||||
elif (entry := self._workers.get(source_id)) is None:
|
||||
request.status, request.detail = "error", "WarmteLink worker is not running."
|
||||
request.completed.set()
|
||||
else:
|
||||
entry[1].request_discovery(request)
|
||||
timer = threading.Timer(_DISCOVERY_TIMEOUT_SECONDS, self._timeout_discovery, args=(request,))
|
||||
timer.daemon = True
|
||||
timer.start()
|
||||
finally:
|
||||
self._lock.release()
|
||||
# A tiny bounded wait makes an immediately available frame observable,
|
||||
# without turning an HTTP call into serial I/O or an unbounded wait.
|
||||
request.completed.wait(_DISCOVERY_WAIT_SECONDS)
|
||||
return request
|
||||
|
||||
@staticmethod
|
||||
def _timeout_discovery(request: DiscoveryRequest) -> None:
|
||||
"""Resolve a stale HTTP request without touching its healthy worker."""
|
||||
if not request.completed.is_set():
|
||||
request.status, request.detail = "error", "Discovery timed out."
|
||||
request.completed.set()
|
||||
|
||||
def _read_desired(self) -> dict[int, _WorkerConfig]:
|
||||
with self._session_factory() as session:
|
||||
return {
|
||||
source.id: _WorkerConfig(source.id, dict(source.config))
|
||||
for source in session.execute(
|
||||
select(MeterSource).where(
|
||||
MeterSource.kind == "warmtelink_serial", MeterSource.enabled.is_(True)
|
||||
)
|
||||
).scalars()
|
||||
}
|
||||
|
||||
def _record_manager_error(self, source_id: int) -> None:
|
||||
"""Best-effort, deliberately non-sensitive lifecycle failure status."""
|
||||
try:
|
||||
with self._session_factory() as session:
|
||||
source = session.get(MeterSource, source_id)
|
||||
if source is not None:
|
||||
source.status = "error"
|
||||
source.last_error = "WarmteLink worker failed"
|
||||
source.updated_at = datetime.now(UTC)
|
||||
session.commit()
|
||||
except Exception:
|
||||
return
|
||||
|
||||
def _reconcile_locked(self, desired: dict[int, _WorkerConfig]) -> None:
|
||||
stale = [
|
||||
source_id for source_id, (config, _) in self._workers.items()
|
||||
if source_id not in desired or desired[source_id] != config
|
||||
]
|
||||
blocked: set[int] = set()
|
||||
for source_id in stale:
|
||||
_, worker = self._workers[source_id]
|
||||
try:
|
||||
worker.stop()
|
||||
stopped = worker.join()
|
||||
except Exception:
|
||||
self._record_manager_error(source_id)
|
||||
blocked.add(source_id)
|
||||
continue
|
||||
if stopped:
|
||||
self._workers.pop(source_id, None)
|
||||
else:
|
||||
logger.error("WarmteLink worker did not stop for source %s", source_id)
|
||||
blocked.add(source_id)
|
||||
self._schedule_reaper_locked(source_id, worker)
|
||||
for source_id, config in desired.items():
|
||||
if source_id in self._workers or source_id in blocked:
|
||||
continue
|
||||
worker: WarmteLinkWorker | None = None
|
||||
try:
|
||||
worker = self._worker_factory(
|
||||
source_id, config.config, session_factory=self._session_factory,
|
||||
serial_factory=self._serial_factory,
|
||||
)
|
||||
self._workers[source_id] = (config, worker)
|
||||
worker.start()
|
||||
except Exception:
|
||||
self._record_manager_error(source_id)
|
||||
# A failed start normally has no thread. If an unusual worker
|
||||
# did start before raising, keep it tracked until it is reaped.
|
||||
if not self._worker_is_alive(worker):
|
||||
self._workers.pop(source_id, None)
|
||||
else:
|
||||
self._schedule_reaper_locked(source_id, worker)
|
||||
|
||||
@staticmethod
|
||||
def _worker_is_alive(worker: object | None) -> bool:
|
||||
thread = getattr(worker, "thread", None)
|
||||
return bool(thread is not None and thread.is_alive())
|
||||
|
||||
def _schedule_reaper_locked(self, source_id: int, worker: WarmteLinkWorker) -> None:
|
||||
if source_id in self._reapers:
|
||||
return
|
||||
self._reapers.add(source_id)
|
||||
threading.Thread(
|
||||
target=self._reap_worker, args=(source_id, worker),
|
||||
# This bookkeeping watcher must not turn a deliberately bounded
|
||||
# application shutdown into an unbounded process wait. The actual
|
||||
# serial worker itself is explicitly non-daemon.
|
||||
name=f"warmtelink-reaper-{source_id}", daemon=True,
|
||||
).start()
|
||||
|
||||
def _reap_worker(self, source_id: int, worker: WarmteLinkWorker) -> None:
|
||||
"""Wait for one timed-out worker, then converge without another API call."""
|
||||
try:
|
||||
while True:
|
||||
try:
|
||||
if worker.join():
|
||||
break
|
||||
except Exception:
|
||||
self._record_manager_error(source_id)
|
||||
return
|
||||
# A custom worker can report a bounded join timeout immediately;
|
||||
# yield before asking again so its reaper cannot busy-spin.
|
||||
threading.Event().wait(0.05)
|
||||
with self._lock:
|
||||
current = self._workers.get(source_id)
|
||||
if current is not None and current[1] is worker:
|
||||
self._workers.pop(source_id)
|
||||
self._reapers.discard(source_id)
|
||||
should_reconcile = not self._shutting_down
|
||||
if should_reconcile:
|
||||
self.reconcile()
|
||||
finally:
|
||||
with self._lock:
|
||||
self._reapers.discard(source_id)
|
||||
|
||||
def shutdown(self) -> None:
|
||||
with self._lock:
|
||||
self._shutting_down = True
|
||||
workers = list(self._workers.items())
|
||||
for source_id, (_, worker) in workers:
|
||||
try:
|
||||
worker.stop()
|
||||
except Exception:
|
||||
self._record_manager_error(source_id)
|
||||
for source_id, (_, worker) in workers:
|
||||
try:
|
||||
stopped = worker.join()
|
||||
except Exception:
|
||||
self._record_manager_error(source_id)
|
||||
continue
|
||||
if not stopped:
|
||||
logger.error("WarmteLink worker did not stop during shutdown for source %s", source_id)
|
||||
with self._lock:
|
||||
self._schedule_reaper_locked(source_id, worker)
|
||||
else:
|
||||
with self._lock:
|
||||
current = self._workers.get(source_id)
|
||||
if current is not None and current[1] is worker:
|
||||
self._workers.pop(source_id)
|
||||
|
||||
|
||||
warmtelink_worker_manager = WarmteLinkWorkerManager()
|
||||
@@ -1,16 +0,0 @@
|
||||
<!DOCTYPE html>
|
||||
<html lang="zh-CN">
|
||||
<head>
|
||||
<meta charset="utf-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1">
|
||||
<title>{% block title %}{{ app_name }}{% endblock %}</title>
|
||||
<link rel="icon" href="data:,">
|
||||
<link rel="stylesheet" href="/static/styles.css">
|
||||
</head>
|
||||
<body>
|
||||
<main class="shell">
|
||||
{% block content %}{% endblock %}
|
||||
</main>
|
||||
</body>
|
||||
</html>
|
||||
|
||||
@@ -1,139 +0,0 @@
|
||||
{% extends "base.html" %}
|
||||
|
||||
{% block title %}Config · {{ app_name }}{% endblock %}
|
||||
|
||||
{% block content %}
|
||||
<section class="panel">
|
||||
<p class="eyebrow">Configuration</p>
|
||||
<h1>Config</h1>
|
||||
|
||||
{% if force_password_change %}
|
||||
<div class="alert">
|
||||
首次登录后需要先修改密码。完成后再继续长期使用当前配置页面。
|
||||
</div>
|
||||
{% endif %}
|
||||
|
||||
{% if password_change_error %}
|
||||
<div class="alert">{{ password_change_error }}</div>
|
||||
{% endif %}
|
||||
|
||||
{% if config_error %}
|
||||
<div class="alert">{{ config_error }}</div>
|
||||
{% endif %}
|
||||
|
||||
{% if config_saved %}
|
||||
<div class="notice">config saved to the app database. Some changes may require an app restart.</div>
|
||||
{% endif %}
|
||||
|
||||
{% if ticktick_oauth_error %}
|
||||
<div class="alert">{{ ticktick_oauth_error }}</div>
|
||||
{% endif %}
|
||||
|
||||
{% if ticktick_oauth_notice %}
|
||||
<div class="notice">{{ ticktick_oauth_notice }}</div>
|
||||
{% endif %}
|
||||
|
||||
{% if smtp_test_error %}
|
||||
<div class="alert">{{ smtp_test_error }}</div>
|
||||
{% endif %}
|
||||
|
||||
{% if smtp_test_notice %}
|
||||
<div class="notice">{{ smtp_test_notice }}</div>
|
||||
{% endif %}
|
||||
|
||||
<div class="meta single-column">
|
||||
<div>
|
||||
<dt>当前用户</dt>
|
||||
<dd>admin</dd>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<section class="config-block">
|
||||
<h2>Change Password</h2>
|
||||
<form class="auth-form" method="post" action="/config/change-password">
|
||||
<input type="hidden" name="csrf_token" value="{{ csrf_token }}">
|
||||
|
||||
<label>
|
||||
<span>Current Password</span>
|
||||
<input type="password" name="current_password" autocomplete="current-password" required>
|
||||
</label>
|
||||
|
||||
<label>
|
||||
<span>New Password</span>
|
||||
<input type="password" name="new_password" autocomplete="new-password" required>
|
||||
</label>
|
||||
|
||||
<label>
|
||||
<span>Confirm New Password</span>
|
||||
<input type="password" name="confirm_password" autocomplete="new-password" required>
|
||||
</label>
|
||||
|
||||
<button type="submit">修改密码</button>
|
||||
</form>
|
||||
</section>
|
||||
|
||||
<section class="config-block">
|
||||
<h2>Config</h2>
|
||||
<form class="config-form" method="post" action="/config">
|
||||
<input type="hidden" name="csrf_token" value="{{ csrf_token }}">
|
||||
|
||||
{% for section in config_sections %}
|
||||
<fieldset class="config-section">
|
||||
<legend>{{ section.name }}</legend>
|
||||
{% for field in section.fields %}
|
||||
<label>
|
||||
<span>{{ field.label }}</span>
|
||||
{% if field.secret %}
|
||||
<input type="{{ field.input_type }}" name="{{ field.env_name }}" value="" placeholder="leave blank to keep current value">
|
||||
<small>{% if field.configured %}configured{% else %}not configured{% endif %}</small>
|
||||
{% else %}
|
||||
<input type="{{ field.input_type }}" name="{{ field.env_name }}" value="{{ field.value }}">
|
||||
{% endif %}
|
||||
</label>
|
||||
{% endfor %}
|
||||
|
||||
{% if section.name == "TickTick" %}
|
||||
<div class="integration-action-row">
|
||||
<div>
|
||||
<p class="integration-action-title">TickTick OAuth</p>
|
||||
<p class="integration-action-copy">Redirect URI: {{ ticktick_redirect_uri or "configure APP_HOSTNAME to generate the callback URI" }}</p>
|
||||
{% if ticktick_oauth_ready %}
|
||||
<p class="integration-action-copy">Use the saved TickTick client settings to start the authorization flow.</p>
|
||||
{% else %}
|
||||
<p class="integration-action-copy">Fill in App Hostname, TickTick Client ID, and TickTick Client Secret before starting OAuth.</p>
|
||||
{% endif %}
|
||||
</div>
|
||||
{% if ticktick_oauth_ready %}
|
||||
<a class="button-link" href="/ticktick/auth/start">Authorize TickTick</a>
|
||||
{% else %}
|
||||
<span class="button-link disabled" aria-disabled="true">Authorize TickTick</span>
|
||||
{% endif %}
|
||||
</div>
|
||||
{% endif %}
|
||||
|
||||
{% if section.name == "SMTP" %}
|
||||
<div class="integration-action-row">
|
||||
<div>
|
||||
<p class="integration-action-title">SMTP Test Email</p>
|
||||
<p class="integration-action-copy">Save the SMTP settings first, then send a simple plaintext test email to the configured recipient.</p>
|
||||
</div>
|
||||
{% if smtp_test_ready %}
|
||||
<button type="submit" formaction="/config/smtp/test" formmethod="post">Send SMTP Test</button>
|
||||
{% else %}
|
||||
<span class="button-link disabled" aria-disabled="true">Send SMTP Test</span>
|
||||
{% endif %}
|
||||
</div>
|
||||
{% endif %}
|
||||
</fieldset>
|
||||
{% endfor %}
|
||||
|
||||
<button type="submit">Save Config</button>
|
||||
</form>
|
||||
</section>
|
||||
|
||||
<form class="logout-form" method="post" action="/logout">
|
||||
<input type="hidden" name="csrf_token" value="{{ csrf_token }}">
|
||||
<button type="submit">登出</button>
|
||||
</form>
|
||||
</section>
|
||||
{% endblock %}
|
||||
@@ -1,36 +0,0 @@
|
||||
{% extends "base.html" %}
|
||||
|
||||
{% block title %}{{ app_name }}{% endblock %}
|
||||
|
||||
{% block content %}
|
||||
<section class="panel">
|
||||
<p class="eyebrow">Python Rewrite Skeleton</p>
|
||||
<h1>{{ app_name }}</h1>
|
||||
<p class="lead">
|
||||
这是当前 Go 后端的 Python 重构基础骨架。此阶段仅提供应用入口、配置、数据库、
|
||||
测试、模板和容器化基础,不包含业务逻辑迁移。
|
||||
</p>
|
||||
<dl class="meta">
|
||||
<div>
|
||||
<dt>运行环境</dt>
|
||||
<dd>{{ app_env }}</dd>
|
||||
</div>
|
||||
<div>
|
||||
<dt>健康检查</dt>
|
||||
<dd><a href="/status">/status</a></dd>
|
||||
</div>
|
||||
<div>
|
||||
<dt>OpenAPI</dt>
|
||||
<dd><a href="/docs">/docs</a></dd>
|
||||
</div>
|
||||
<div>
|
||||
<dt>登录</dt>
|
||||
<dd><a href="/login">/login</a></dd>
|
||||
</div>
|
||||
<div>
|
||||
<dt>Notion</dt>
|
||||
<dd>{{ notion_status }}</dd>
|
||||
</div>
|
||||
</dl>
|
||||
</section>
|
||||
{% endblock %}
|
||||
@@ -1,33 +0,0 @@
|
||||
{% extends "base.html" %}
|
||||
|
||||
{% block title %}登录 · {{ app_name }}{% endblock %}
|
||||
|
||||
{% block content %}
|
||||
<section class="panel auth-panel">
|
||||
<p class="eyebrow">Authentication</p>
|
||||
<h1>登录</h1>
|
||||
<p class="lead">
|
||||
登录成功后会进入受保护的 config 页面。
|
||||
</p>
|
||||
|
||||
{% if error_message %}
|
||||
<div class="alert">{{ error_message }}</div>
|
||||
{% endif %}
|
||||
|
||||
<form class="auth-form" method="post" action="/login">
|
||||
<input type="hidden" name="csrf_token" value="{{ csrf_token }}">
|
||||
|
||||
<label>
|
||||
<span>Username</span>
|
||||
<input type="text" name="username" autocomplete="username" required>
|
||||
</label>
|
||||
|
||||
<label>
|
||||
<span>Password</span>
|
||||
<input type="password" name="password" autocomplete="current-password" required>
|
||||
</label>
|
||||
|
||||
<button type="submit">登录</button>
|
||||
</form>
|
||||
</section>
|
||||
{% endblock %}
|
||||
@@ -53,19 +53,17 @@ idna==3.11
|
||||
# httpx
|
||||
iniconfig==2.3.0
|
||||
# via pytest
|
||||
jinja2==3.1.6
|
||||
# via -r requirements.in
|
||||
mako==1.3.11
|
||||
# via alembic
|
||||
markupsafe==3.0.3
|
||||
# via
|
||||
# jinja2
|
||||
# mako
|
||||
# via mako
|
||||
packaging==26.1
|
||||
# via
|
||||
# build
|
||||
# pytest
|
||||
# wheel
|
||||
paho-mqtt==2.1.0
|
||||
# via -r requirements.in
|
||||
pip-tools==7.5.3
|
||||
# via -r dev-requirements.in
|
||||
pluggy==1.6.0
|
||||
@@ -82,10 +80,16 @@ pydantic-settings==2.13.1
|
||||
# via -r requirements.in
|
||||
pygments==2.20.0
|
||||
# via pytest
|
||||
pymodbus==3.13.1
|
||||
# via -r requirements.in
|
||||
pyotp==2.10.0
|
||||
# via -r requirements.in
|
||||
pyproject-hooks==1.2.0
|
||||
# via
|
||||
# build
|
||||
# pip-tools
|
||||
pyserial==3.5
|
||||
# via -r requirements.in
|
||||
pytest==8.4.2
|
||||
# via -r dev-requirements.in
|
||||
python-dotenv==1.2.2
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
# Local dev override — use explicitly:
|
||||
# docker compose -f docker-compose.yml -f docker-compose.dev.yml up --build
|
||||
# Isolated from the production stack so both can run on this host at once:
|
||||
# - distinct compose project name (separate network/grouping)
|
||||
# - distinct container names (-dev suffix; Docker rejects duplicate names)
|
||||
# - distinct image tag (local build doesn't clobber the prod :latest tag)
|
||||
name: home-automation-dev
|
||||
|
||||
services:
|
||||
migration:
|
||||
build: .
|
||||
image: home-automation:dev
|
||||
container_name: home-automation-migration-dev
|
||||
environment:
|
||||
# In-container path for the mounted ./data volume (./data -> /app/data).
|
||||
# Overrides the host-absolute APP_DATABASE_URL in .env for local compose runs.
|
||||
APP_DATABASE_URL: "sqlite:////app/data/app.db"
|
||||
|
||||
app:
|
||||
build: .
|
||||
image: home-automation:dev
|
||||
container_name: home-automation-app-dev
|
||||
# Publish on 8002 for dev. `!override` REPLACES the base ports list instead of
|
||||
# appending to it, so the dev stack does NOT also bind the production 8881.
|
||||
ports: !override
|
||||
- "127.0.0.1:8002:8000"
|
||||
environment:
|
||||
APP_DATABASE_URL: "sqlite:////app/data/app.db"
|
||||
devices: !override
|
||||
- "${WARMTELINK_DEVICE_PATH:?Set a stable /dev/serial/by-id path}:/dev/warmtelink:rw"
|
||||
group_add: !override
|
||||
- "${WARMTELINK_SERIAL_GID:?Set the host serial device GID}"
|
||||
@@ -1,6 +0,0 @@
|
||||
services:
|
||||
migration:
|
||||
build: .
|
||||
|
||||
app:
|
||||
build: .
|
||||
+10
-1
@@ -6,9 +6,12 @@ services:
|
||||
restart: "no"
|
||||
init: true
|
||||
command: ["python", "-m", "scripts.run_migrations"]
|
||||
environment:
|
||||
TZ: "${TZ:-Europe/Amsterdam}"
|
||||
volumes:
|
||||
- ./data:/app/data
|
||||
- ./.env:/app/.env:ro
|
||||
- /etc/localtime:/etc/localtime:ro
|
||||
|
||||
app:
|
||||
container_name: home-automation-app
|
||||
@@ -16,12 +19,18 @@ services:
|
||||
user: "1000:1000"
|
||||
restart: unless-stopped
|
||||
init: true
|
||||
environment:
|
||||
TZ: "${TZ:-Europe/Amsterdam}"
|
||||
depends_on:
|
||||
migration:
|
||||
condition: service_completed_successfully
|
||||
ports:
|
||||
- "127.0.0.1:8881:8000"
|
||||
devices:
|
||||
- "${WARMTELINK_DEVICE_PATH:?Set a stable /dev/serial/by-id path}:/dev/warmtelink:rw"
|
||||
group_add:
|
||||
- "${WARMTELINK_SERIAL_GID:?Set the host serial device GID}"
|
||||
volumes:
|
||||
- ./data:/app/data
|
||||
- ./.env:/app/.env:ro
|
||||
|
||||
- /etc/localtime:/etc/localtime:ro
|
||||
|
||||
+108
-14
@@ -19,35 +19,42 @@
|
||||
|
||||
- `main.py`
|
||||
- FastAPI app factory
|
||||
- lifespan
|
||||
- lifespan(APScheduler 启停、MQTT 客户端起停、连接后触发 HA Discovery 发布;注册 DSMR MQTT source,并启动/关闭每个 enabled WarmteLink 业务只读 serial worker;`tibber-refresh`、electricity 与 thermal cost tick 均由 scheduler 驱动)
|
||||
- 基础路由注册
|
||||
- `config.py`
|
||||
- 环境变量驱动的 settings
|
||||
- 环境变量驱动的 settings(含 M5 新增的 MQTT/HA Discovery/Modbus 配置项;M6 新增 `dsmr_ingest_enabled`、`dsmr_mqtt_topic`、`dsmr_sample_interval_s`、`tibber_api_token`(secret)、`tibber_home_id`)
|
||||
- `db.py`
|
||||
- 统一数据层:一个 `Base`、一个绑定 `app_database_url` 的 cached engine(SQLite WAL)、`get_engine` / `get_session_local` / `reset_db_caches` / `get_db_session`
|
||||
- `dependencies.py`
|
||||
- 通用依赖注入
|
||||
- `api/`
|
||||
- HTTP routes
|
||||
- 当前已迁入 `/login`、`/logout`、`/admin`
|
||||
- 当前已迁入 `GET /public-ip/check`
|
||||
- 当前已迁入 `POST /homeassistant/publish` 第一版入口
|
||||
- 当前已迁入 `POST /poo/record` 与 `GET /poo/latest`
|
||||
- `api/routes/api/`:JSON API(`/api/*` 前缀),供 React SPA 调用:会话/鉴权、配置读写、记录 CRUD、Modbus、Expose 与 MQTT 测试;Energy 包含 source profile/source/channel/history/discover/binding/Meter API、scope-aware contracts/prices/costs,以及兼容的 DSMR latest API
|
||||
- 裸 ingestion 端点:`GET /public-ip/check`、`POST /homeassistant/publish`、`POST /poo/record`、`GET /poo/latest`、TickTick OAuth 等
|
||||
- `models/`
|
||||
- SQLAlchemy models
|
||||
- 所有模型(auth / config / public_ip / location / poo)共用同一个 `Base`,均落在单一 `app.db` 中
|
||||
- 所有模型(auth / config / public_ip / location / poo / modbus / expose / energy / meter_source)共用同一个 `Base`,均落在单一 `app.db` 中
|
||||
- M5 新增:`ModbusDevice`(设备部署层)、`ModbusReading`(通用遥测,JSON payload)、`ExposedEntityToggle`(HA 实体暴露开关)
|
||||
- Energy:`MeterSource` / `MeterSourceChannel` / `MeterSourceBinding` 将协议连接、稳定测量 channel 与 Meter epoch 分离;`DsmrReading`、`WarmteLinkReading` 分别保存 JSON 与 Decimal scalar 历史;electricity `EnergyCostPeriod` 绑定 source binding,thermal `MeterCostPeriod` 保存审计账本;合同按 electricity / thermal scope 共存
|
||||
- `schemas/`
|
||||
- Pydantic schemas
|
||||
- Pydantic schemas(包括 `modbus.py`、`expose.py`、`energy_contract.py`、`energy.py`、`meter_source.py`)
|
||||
- `services/`
|
||||
- 业务服务层
|
||||
- 当前已迁入 config page 的 DB 持久化逻辑
|
||||
- 当前已迁入 public IPv4 检查、状态持久化与变化通知逻辑
|
||||
- 当前已迁入 SMTP 发信与测试发信逻辑
|
||||
- M5 新增:`modbus_poll.py`(采集 service,逐设备 poll + 落库 + 推 MQTT state)、`ha_discovery.py`(构建 HA Discovery payload、发布 retained config、发布 state)
|
||||
- Energy:`dsmr_ingest.py` 按 source 入库;`warmtelink_ingest.py` 接纳连续确认的业务只读 P1 scalar,`warmtelink_worker.py` 管理 interruptible serial reconnect(pyserial 的 POSIX `O_RDWR` 打开由非 root、非 privileged、无 `m` 的 Docker `rw` device rule 支持;worker 只 read/close,绝不 write);`energy_cost.py` 与 `meter_cost.py` 分别计算 electricity/thermal 账本,均拒绝跨 Meter/binding 相减
|
||||
- `integrations/`
|
||||
- 外部系统适配层
|
||||
- 当前已迁入 Home Assistant outbound adapter
|
||||
- `templates/`
|
||||
- Jinja2 模板
|
||||
- Home Assistant outbound adapter(REST 通道,原有)
|
||||
- M5 新增:`modbus/`(pymodbus 薄封装:`driver.py` 块读 + float32 解码;`profiles.py` YAML profile 加载/校验/解码;`profiles/sdm120.yaml` SDM120 协议声明)
|
||||
- M5 新增:`mqtt.py`(paho-mqtt 长连接 `MqttManager`:lifespan 起/停、配置变更重连、`publish(topic, payload, retain)`)
|
||||
- M5 新增:`expose.py`(通用 expose 框架:`ExposableEntity`、provider 注册表、`build_catalog`;Modbus provider 从 YAML profile 派生 sensor/binary_sensor 实体目录)
|
||||
- M6 新增:`pricing/`(通用电价层:`profiles.py` pydantic 加载/校验 YAML profile + `validate_values`;`strategies.py` manual/tibber 出价策略注册表;`profiles/manual.yaml` 固定/双费率结构;`profiles/tibber.yaml` 动态电价结构)
|
||||
- M6 新增:`tibber/`(`client.py`:httpx GraphQL 客户端,`priceInfoRange(QUARTER_HOURLY)` 查询,按 `starts_at` 解析,不假设固定节点数)
|
||||
- M6 扩展:`mqtt.py` 新增订阅端(`subscribe(topic, handler)`,`on_connect` 里 subscribe,`on_message` 按 topic 分发,handler 异常吞掉不崩连接)
|
||||
- M6 扩展:`expose.py` 新增 `_energy_cost_provider`(`buy_price_now`、`sell_price_now`、`import_cost_total`/`export_revenue_total` 均 `total_increasing`,反哺 HA Energy)
|
||||
- `static/`
|
||||
- 极简静态资源
|
||||
|
||||
@@ -63,17 +70,104 @@ pytest 测试目录。后续可以在这里自然扩展:
|
||||
- mock tests
|
||||
- integration tests
|
||||
|
||||
### `frontend/`
|
||||
|
||||
React SPA 前端(M2 引入)。Vite + React + TypeScript + Mantine,由 FastAPI 同源托管。
|
||||
|
||||
- `src/`:React 源码
|
||||
- `src/api/`:由 `openapi/openapi.json` 生成的类型化 client(`schema.d.ts`)+ fetch 封装
|
||||
- `dist/`:`npm run build` 产物,由 FastAPI 的 `SPA_DIST_DIR` 挂载并对非 `/api` 路径做 fallback
|
||||
|
||||
### `scripts/`
|
||||
|
||||
辅助脚本目录。当前包含 OpenAPI 导出脚本。
|
||||
辅助脚本目录。当前包含:
|
||||
|
||||
- `export_openapi.py`:导出 OpenAPI schema 静态产物
|
||||
- `run_migrations.py`:运行 Alembic migration
|
||||
- `app_db_adopt.py`:App DB 接管 / 初始化
|
||||
- `migrate_legacy_data.py`:一次性历史数据搬迁脚本
|
||||
- `admin_cli.py`:Admin CLI 逃生通道(M4),见下方"登录加固"说明
|
||||
- `modbus_cli.py`(M5):Modbus 手工试读 CLI(`read`/`probe` 两个只读子命令),不依赖 DB,供受控手工验证链路连通性
|
||||
|
||||
### `openapi/`
|
||||
|
||||
OpenAPI schema 静态产物(`openapi.json` / `openapi.yaml`),由 `python scripts/export_openapi.py` 生成,纳入版本控制。前端 codegen 以此为契约源。
|
||||
|
||||
## 登录加固(M4)
|
||||
|
||||
M4 在基础 Argon2 + server-side session 鉴权之上叠加了三层防御:
|
||||
|
||||
**防爆破 / 指数退避**:`app/services/login_throttle.py` 按 client IP 与 username 双键记失败计数,失败超过 3 次后指数增长等待时间(最长 15 分钟),`POST /api/auth/login` 在退避窗口内直接返回 `429 + Retry-After`,不执行 Argon2 验证。成功登录后清零。退避是延迟而非永久封号;全局开关 `AUTH_LOGIN_THROTTLE_ENABLED`(CONFIG_FIELDS,默认开);反代后需设 `AUTH_TRUST_FORWARDED_FOR=true`(`.env` 部署级,默认 false)。
|
||||
|
||||
**CLI 逃生通道**:`scripts/admin_cli.py`(入口 `python -m scripts.admin_cli`)直连本地 DB,**无需 HTTP 服务运行、无需任何已存凭据**,支持:重置密码(`reset-password`)、解锁退避(`unlock`)、关停 TOTP(`disable-totp`,零凭据最终逃生)、重新发放 TOTP secret(`reissue-totp`)、查看用户列表(`list-admin`)。CLI 只动 auth 行,不触碰用户数据表。
|
||||
|
||||
**可选 TOTP 二次验证**:admin 可在设置页自选启用 RFC 6238 TOTP。启用后登录为两步(密码 → 6 位动态码或一次性恢复码);不启用维持纯密码。TOTP secret 明文存库(与其他 secret 一致,靠文件权限保护);恢复码以 Argon2 哈希存储,使用后消费(一次性)。后端使用 `pyotp`,二维码由前端 `qrcode.react` 渲染。issuer 标签由 `AUTH_TOTP_ISSUER`(`.env` 部署级)配置,默认回退 `app_name`。
|
||||
|
||||
详细说明:[`docs/auth.md`](./auth.md)
|
||||
|
||||
## M5 — 通用 Modbus 采集链路与 MQTT 通道
|
||||
|
||||
### Modbus 采集链路
|
||||
|
||||
```
|
||||
YAML profile(协议知识) modbus_device 行(部署信息,DB)
|
||||
- 寄存器块/地址/类型 - host / port(网关 IP)
|
||||
- 每量:key/unit/device_class - unit_id(Modbus slave 地址)
|
||||
- ha_component - friendly_name / profile / poll_interval_s / enabled
|
||||
│ │
|
||||
└────────────┬───────────────────────┘
|
||||
│ APScheduler job(轮询所有 enabled 设备)
|
||||
│ driver.py: ModbusTcpClient → FC04 块读 → 大端 float32 解码
|
||||
│ profiles.py: decode(profile, registers) → dict[key → value]
|
||||
▼
|
||||
modbus_reading 行
|
||||
device_id · recorded_at · payload(JSON)
|
||||
{"voltage": 230.2, "current": 1.3, "active_power": 295.0, ...}
|
||||
```
|
||||
|
||||
**关键设计决策**:
|
||||
- 命名分层:存储/采集/API 全部通用 `modbus_*`(`/api/modbus/devices`);面向用户的领域视图叫 **Energy**(第一个)。接入新设备型号只需新增 YAML profile,不改表/不改 API。
|
||||
- 读数为 JSON `payload`(无固定列),SQLite `json_extract` 在 DB 端做 AVG/GROUP BY(被 `(device_id, recorded_at)` 索引圈住)。
|
||||
- FK `ON DELETE RESTRICT`:有读数的设备拒删,引导改为 `enabled=false`。
|
||||
- 设备 `uuid`(uuid4 内部生成)是**稳定身份锚点**:API 路径键、HA Discovery `unique_id` 来源,不随 friendly_name 改变。
|
||||
|
||||
### 第二条 MQTT 通道(HA Discovery 发布)
|
||||
|
||||
与已有 REST 通道(`app/integrations/homeassistant.py`、`POST /homeassistant/publish`)**并行、不冲突**:
|
||||
|
||||
```
|
||||
MQTT 通道(M5 新增):
|
||||
paho-mqtt MqttManager(lifespan 长连接,配置变更可重连)
|
||||
│
|
||||
├─► Discovery config(retained)
|
||||
│ topic: <prefix>/<component>/<node>/<object>/config
|
||||
│ 内容:device 块(identifiers=uuid)、state_topic、unique_id(uuid+key)、
|
||||
│ name(friendly_name)、device_class、unit_of_measurement、availability
|
||||
│ 时机:连接成功时 / 目录或勾选变更时(全量重发);
|
||||
│ 取消勾选时发空 payload(清除 entity)
|
||||
│
|
||||
└─► State / Availability(非 retained)
|
||||
时机:每次轮询成功后推该设备所有 enabled entity 的最新值;
|
||||
周期兜底 job 重推所有 enabled entity + online topic
|
||||
|
||||
expose 框架:
|
||||
provider 动态产出 ExposableEntity 目录(元数据从 YAML profile 派生)
|
||||
ExposedEntityToggle 表:只存逐 key 开关(default=disabled)
|
||||
build_catalog(session) → 目录 + 勾选状态(合并所有 provider)
|
||||
```
|
||||
|
||||
**HA entity 身份模型(Z2M 语义)**:
|
||||
- `unique_id` = `f"{device.uuid}_{metric.key}"`(稳定,改名不变)
|
||||
- `name` = `friendly_name`(改名重发 discovery,HA 显示名跟着变、历史不丢)
|
||||
- 每设备除各 sensor entity 外,另有 `binary_sensor` `online`(取 `last_poll_ok`)——这是"不止 sensor"的体现
|
||||
|
||||
## 当前约束
|
||||
|
||||
- 当前只搭骨架,不迁业务逻辑
|
||||
- 当前数据库继续使用 SQLite
|
||||
- 当前不引入前后端分离
|
||||
- ~~当前不引入前后端分离~~ **已退役(M2)**:现为 React SPA + JSON `/api` 层,由 FastAPI 同源托管
|
||||
- 当前不设计 Notion 模块
|
||||
- 当前通知能力仍保持极小范围,不引入独立通知中心或多渠道抽象
|
||||
- Modbus 当前**仅 TCP**(Waveshare RTU↔TCP 网关),**只读**(FC03/04),不写设备寄存器
|
||||
|
||||
## 关于 Notion
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user