From 56d74df55a51326adcae18e35cdebd463ddf42d4 Mon Sep 17 00:00:00 2001 From: Alex Hope-O'Connor Date: Wed, 2 Sep 2026 07:14:20 +1000 Subject: [PATCH] Release ArduinoHA v3.1.0 --- .github/workflows/ci.yml | 10 + .github/workflows/ha-contract.yml | 36 + .github/workflows/release.yml | 28 + .gitignore | 4 + CHANGELOG.md | 20 +- README.md | 4 +- docs/README.md | 1 + docs/compatibility.md | 41 ++ docs/device-and-discovery.md | 65 +- docs/entities.md | 2 +- docs/getting-started.md | 8 +- docs/mqtt-usage.md | 5 + library.json | 4 +- library.properties | 2 +- platformio.ini | 31 + scripts/bump-version.sh | 7 + scripts/prepare-release.sh | 12 + src/ArduinoHADefines.h | 2 +- src/HADevice.cpp | 166 ++++- src/HADevice.h | 15 +- src/HAMqtt.cpp | 676 +++++++++++++++--- src/HAMqtt.h | 136 +++- src/device-types/HABaseDeviceType.cpp | 120 +++- src/device-types/HABaseDeviceType.h | 42 +- src/device-types/HAText.cpp | 72 +- src/device-types/HAText.h | 20 +- src/mocks/PubSubClientMock.cpp | 61 +- src/mocks/PubSubClientMock.h | 8 +- src/utils/HAAvailabilityConfig.cpp | 69 +- src/utils/HAAvailabilityConfig.h | 5 + src/utils/HAJson.cpp | 174 +++++ src/utils/HAJson.h | 46 ++ src/utils/HASerializer.cpp | 415 ++++++++--- src/utils/HASerializerArray.cpp | 56 +- test/native/include/Arduino.h | 235 ++++++ test/native/include/Client.h | 10 + test/native/include/IPAddress.h | 33 + test/test_device_metadata/test_main.cpp | 109 +++ test/test_entities_basic/test_main.cpp | 4 + test/test_entities_basic/test_main.h | 4 + .../tests/test_binary_sensor.cpp | 1 - .../test_entities_basic/tests/test_button.cpp | 1 - .../test_entities_basic/tests/test_sensor.cpp | 1 - .../test_entities_basic/tests/test_switch.cpp | 3 +- test/test_entities_basic/tests/test_text.cpp | 61 +- .../tests/test_camera.cpp | 3 +- .../tests/test_cover.cpp | 3 +- .../test_entities_extended/tests/test_fan.cpp | 3 +- .../tests/test_hvac.cpp | 3 +- .../tests/test_light.cpp | 1 - .../tests/test_lock.cpp | 3 +- .../tests/test_device_tracker.cpp | 3 +- .../tests/test_device_trigger.cpp | 2 +- test/test_entities_misc/tests/test_scene.cpp | 3 +- .../tests/test_tag_scanner.cpp | 2 +- .../tests/test_number.cpp | 3 +- .../tests/test_select.cpp | 1 - test/test_native_core/test_main.cpp | 355 +++++++++ test/test_utils_json/test_main.cpp | 86 +++ tests/ha-contract/Dockerfile | 8 + tests/ha-contract/README.md | 45 ++ tests/ha-contract/compose.yaml | 43 ++ tests/ha-contract/configuration.yaml | 3 + tests/ha-contract/mosquitto.conf | 6 + tests/ha-contract/requirements.txt | 4 + tests/ha-contract/test_contract.py | 441 ++++++++++++ 66 files changed, 3456 insertions(+), 390 deletions(-) create mode 100644 .github/workflows/ha-contract.yml create mode 100644 docs/compatibility.md create mode 100644 src/utils/HAJson.cpp create mode 100644 src/utils/HAJson.h create mode 100644 test/native/include/Arduino.h create mode 100644 test/native/include/Client.h create mode 100644 test/native/include/IPAddress.h create mode 100644 test/test_device_metadata/test_main.cpp create mode 100644 test/test_native_core/test_main.cpp create mode 100644 test/test_utils_json/test_main.cpp create mode 100644 tests/ha-contract/Dockerfile create mode 100644 tests/ha-contract/README.md create mode 100644 tests/ha-contract/compose.yaml create mode 100644 tests/ha-contract/configuration.yaml create mode 100644 tests/ha-contract/mosquitto.conf create mode 100644 tests/ha-contract/requirements.txt create mode 100644 tests/ha-contract/test_contract.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 32eba91..ff2e373 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -20,6 +20,16 @@ jobs: - uses: actions/checkout@v4 - run: bash -n scripts/*.sh - run: ./scripts/check-docs.sh + native-core-tests: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-python@v5 + with: + python-version: '3.11' + - run: python -m pip install --upgrade platformio==6.1.19 + - run: pio test -e native --filter test_native_core + compile-tests: runs-on: ubuntu-latest strategy: diff --git a/.github/workflows/ha-contract.yml b/.github/workflows/ha-contract.yml new file mode 100644 index 0000000..06ae001 --- /dev/null +++ b/.github/workflows/ha-contract.yml @@ -0,0 +1,36 @@ +name: Home Assistant MQTT contract + +on: + workflow_dispatch: + schedule: + - cron: "17 3 * * 1" + +permissions: + contents: read + +jobs: + contract: + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + home_assistant: ["2024.11.3", stable, dev] + env: + HA_VERSION: ${{ matrix.home_assistant }} + COMPOSE_FILE: tests/ha-contract/compose.yaml + steps: + - uses: actions/checkout@v4 + - name: Start broker and Home Assistant + run: docker compose up -d mqtt homeassistant + - name: Check single-to-device migration + run: docker compose run --rm tests + - name: Restart Home Assistant from retained data + run: docker compose restart homeassistant + - name: Check retained discovery after restart + run: docker compose run --rm -e CONTRACT_MODE=retained-restart tests + - name: Collect Home Assistant logs on failure + if: failure() + run: docker compose logs --no-color homeassistant mqtt + - name: Remove contract volumes + if: always() + run: docker compose down -v diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 97a149d..d33c8e5 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -9,7 +9,34 @@ permissions: contents: write jobs: + ha-contract: + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + home_assistant: ["2024.11.3", stable, dev] + env: + HA_VERSION: ${{ matrix.home_assistant }} + COMPOSE_FILE: tests/ha-contract/compose.yaml + steps: + - uses: actions/checkout@v4 + - name: Start broker and Home Assistant + run: docker compose up -d mqtt homeassistant + - name: Check single-to-device migration + run: docker compose run --rm tests + - name: Restart Home Assistant from retained data + run: docker compose restart homeassistant + - name: Check retained discovery after restart + run: docker compose run --rm -e CONTRACT_MODE=retained-restart tests + - name: Collect Home Assistant logs on failure + if: failure() + run: docker compose logs --no-color homeassistant mqtt + - name: Remove contract volumes + if: always() + run: docker compose down -v + publish: + needs: ha-contract runs-on: ubuntu-latest steps: - uses: actions/checkout@v4 @@ -18,6 +45,7 @@ jobs: python-version: '3.11' - run: python -m pip install --upgrade platformio==6.1.19 - run: ./scripts/test.sh compile --platform esp8266 + - run: pio test -e native --filter test_native_core - run: ./scripts/test.sh compile --platform esp32 - run: ./scripts/check-docs.sh - run: ./scripts/prepare-release.sh "$GITHUB_REF_NAME" diff --git a/.gitignore b/.gitignore index 329c13d..a48bc94 100644 --- a/.gitignore +++ b/.gitignore @@ -1,6 +1,10 @@ .DS_Store tmp/ +# Python test tooling +__pycache__/ +*.py[cod] + # PlatformIO build output .pio/ diff --git a/CHANGELOG.md b/CHANGELOG.md index 2d353f3..7a34f29 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,9 +1,25 @@ # Changelog +## 3.1.0 + +**Home Assistant MQTT compatibility:** + +- Add an explicit, retry-safe single-component to device-discovery migration sequence: retained `migrate_discovery` markers, device config, then legacy-topic cleanup. +- Match Home Assistant device-component removal semantics with a platform-only tombstone followed by a compacted payload; `republishDiscovery()` re-adds a removed component. +- Stop serializing removed `obj_id`; keep `setObjectId()` source-compatible and use `setDefaultEntityId()` for new entities. + +**Reliability and maintenance:** + +- Reject invalid discovery topic tokens, escape dynamic JSON values, and fail discovery publication on partial writes or missing device IDs. +- Register entities safely regardless of whether they are constructed before or after `HAMqtt`, diagnose entity-limit drops, and make explicit disconnect callbacks observable. +- Make `HAText` own its current state and bounded command copy; align device-discovery origin version with the package version. +- Add native AddressSanitizer/UndefinedBehaviorSanitizer core regression coverage and release-script version consistency checks. + ## 3.0.2 - Align the Arduino IDE `library.properties` version with `library.json` so both package formats identify the same release. + ## 3.0.1 - Build and test against the Arduino 3-compatible pioarduino ESP32 platform. @@ -31,8 +47,8 @@ **Migration notes:** * Unit tests now live under `test/` as PlatformIO Unity suites (`pio test`). The legacy `tests/` tree (AUnit, EpoxyDuino, Make, AUniter) has been removed. * Home Assistant deprecated MQTT `object_id` in favor of `default_entity_id`, and newer Home Assistant versions may warn on or remove `object_id` handling in discovery payloads. -* Existing code using `setObjectId()` remains supported as a legacy fallback, but new projects should migrate to `setDefaultEntityId()`. -* If you enable device discovery mode, avoid publishing per-entity discovery topics manually. Use `republishDiscovery()` when a runtime config change needs to refresh discovery state. +* `setObjectId()` remains source-compatible but no longer serializes the removed `obj_id` discovery property; new projects should use `setDefaultEntityId()`. +* Existing single-component devices must use the staged migration API before device discovery; use `republishDiscovery()` when a runtime config change needs to refresh discovery state. ## 2.1.0 diff --git a/README.md b/README.md index b55d9cb..f2b4b39 100644 --- a/README.md +++ b/README.md @@ -3,7 +3,7 @@ [![](https://img.shields.io/github/v/release/alexhopeoconnor/arduino-home-assistant?label=Version)](https://github.com/alexhopeoconnor/arduino-home-assistant/releases) [![](https://img.shields.io/badge/Documentation-40BC13)](docs/README.md) -ArduinoHA is the maintained MQTT-discovery library behind compact Arduino, ESP8266, and ESP32 integrations with Home Assistant. It uses Arduino's standard network `Client` API and is continuously compile-tested on ESP8266 and ESP32. +ArduinoHA is the maintained MQTT-discovery library behind compact Arduino, ESP8266, and ESP32 integrations with Home Assistant. It uses Arduino's standard network `Client` API and is continuously compile-tested on ESP8266 and ESP32. Device discovery requires Home Assistant 2024.11.0 or newer; single-component discovery remains supported. ## Why use it @@ -49,7 +49,7 @@ Read [Getting started](docs/getting-started.md) before copying this into product ```ini lib_deps = - home-assistant-integration=https://github.com/alexhopeoconnor/arduino-home-assistant.git#v3.0.2 + home-assistant-integration=https://github.com/alexhopeoconnor/arduino-home-assistant.git#v3.1.0 ``` PlatformIO clones the Git repository and checks out the ref after `#`; that ref is a release tag, not a GitHub Release asset. Arduino IDE is supported through the included [`library.properties`](library.properties); see [Getting started](docs/getting-started.md#install-the-library). diff --git a/docs/README.md b/docs/README.md index 74a5e57..17ed728 100644 --- a/docs/README.md +++ b/docs/README.md @@ -11,3 +11,4 @@ User-facing notes for this library. API details live in the headers under [`src/ | [Examples](../examples/README.md) | Curated paths through the standalone sketches | Class-level API details live in headers under [`src/`](../src/). Return to the [project overview](../README.md). +| [Compatibility baseline](compatibility.md) | Audited upstream/HA targets and deliberate exclusions | diff --git a/docs/compatibility.md b/docs/compatibility.md new file mode 100644 index 0000000..f7a8c5c --- /dev/null +++ b/docs/compatibility.md @@ -0,0 +1,41 @@ +# Compatibility baseline + +This maintenance line begins from fork commit `84cc0037b1c0` (release `v3.0.2`). +It intentionally tracks Home Assistant's MQTT discovery contract without merging +upstream development wholesale. + +| Reference | Audited revision / target | +| --- | --- | +| Fork baseline | `84cc0037b1c0` (`v3.0.2`) | +| Upstream main | `1d333ab229b2` (`v2.1.0`) | +| Upstream develop | `a7039fad810b` (unreleased WIP 2.2.0) | +| Device discovery minimum | Home Assistant `2024.11.0` | +| Contract matrix | `2024.11.3`, current `stable`, current `dev` | + +The only imported upstream code fix is the four missing-device-ID guards from +upstream commit `9c9d074`. Device discovery migration, JSON validation, +lifecycle handling, and tests are fork-native changes. + +## Deliberate exclusions + +- No merge or wholesale cherry-pick of upstream `develop` / WIP 2.2.0. +- No `obj_id` serialization: Home Assistant removed that discovery field. Use + `setDefaultEntityId()` for new code. +- No binary-sensor `state_class`: it is not valid in the current HA MQTT binary + sensor schema. +- No upstream IMqttClient abstraction or unreviewed entity-type feature PRs. + +## Ongoing audit routine + +The `upstream` remote has no usable push URL. Fetch and compare explicitly: + +```bash +git fetch --prune upstream +git log --oneline main..upstream/develop +git diff --stat main...upstream/develop +``` + +Port only independently reviewed changes with regression tests; treat open +upstream pull requests as proposals, not release inputs. Run the native, +board-compile, and HA contract gates before publishing a new release. + diff --git a/docs/device-and-discovery.md b/docs/device-and-discovery.md index d7e649a..d3449d1 100644 --- a/docs/device-and-discovery.md +++ b/docs/device-and-discovery.md @@ -22,7 +22,13 @@ String setters take **pointers whose contents are not copied** — use literals When MQTT connects, the library publishes Home Assistant **MQTT discovery** payloads so entities appear automatically. -**Rule:** construct entity objects **after** `HAMqtt` so they can register with it. +Entities can be constructed either before or after `HAMqtt`. Entities which already +exist when `HAMqtt` is constructed are registered then; entities constructed later +register immediately. Both the device, MQTT object, and entities must have a lifetime +that outlasts `mqtt.loop()`. + +`HAMqtt` has a configurable entity limit (24 by default). Check +`getRegisteredDeviceTypeCount()`, `getDeviceTypeLimit()`, and `getDeviceTypeRegistrationFailures()` in firmware diagnostics: registrations above the limit are rejected and logged rather than silently disappearing from discovery. ### Topic prefixes @@ -41,7 +47,51 @@ mqtt.setDataPrefix("myDataPrefix"); ### Single-component vs device discovery - **Default:** one retained discovery topic per entity (single-component discovery). -- **Optional:** call `HAMqtt::enableDeviceDiscovery()` to publish a single **device** discovery payload with components under `cmps` (see project README for migration notes). +- **Device discovery:** `HAMqtt::enableDeviceDiscovery()` publishes one retained device payload with components under `cmps`. + +Both formats remain supported by Home Assistant. Device discovery requires Home Assistant **2024.11.0 or newer**. Use `enableDeviceDiscovery()` only for a new device which has never published this library's single-component discovery topics. + +### Migrating an existing device to device discovery + +Do not switch an existing device by calling `enableDeviceDiscovery()` alone. Home Assistant derives an entity's discovery identity from the topic, and a direct switch causes retained-topic conflicts. The migration deliberately publishes, in order: + +1. `{"migrate_discovery":true}` to every old retained component config topic; +2. the retained `/device//config` payload; and only then +3. an empty retained payload to every old component config topic. + +Keep the device ID and each entity ID unchanged. They preserve the mapping from the old `///config` topic to `cmps[entityId]` in the new device payload. + +Start migration once during setup, then advance *one stage per loop iteration* after the MQTT connection is available. Each method is retry-safe: do not advance if it returns `false`. + +```cpp +bool migrationStarted = false; + +void setup() { + // Create device, mqtt, and entities; keep their IDs unchanged. + migrationStarted = mqtt.beginDeviceDiscoveryMigration(); + mqtt.begin("192.168.1.50", "user", "password"); +} + +void loop() { + mqtt.loop(); + + switch (mqtt.getDeviceDiscoveryMigrationState()) { + case HAMqtt::DeviceDiscoveryMigrationMarkersPending: + mqtt.publishDeviceDiscoveryMigrationMarkers(); + break; + case HAMqtt::DeviceDiscoveryMigrationMarkersPublished: + mqtt.publishDeviceDiscoveryMigrationConfig(); + break; + case HAMqtt::DeviceDiscoveryMigrationDevicePublished: + mqtt.completeDeviceDiscoveryMigration(); + break; + default: + break; + } +} +``` + +The migration stage is held in RAM. If the board reboots before completion, start the same migration again; the retained marker/config/cleanup publishes are idempotent. If you must abandon a started migration while connected, call `rollbackDeviceDiscoveryMigration()`. Once a device config has been published, it first writes `{"migrate_discovery":true}` to the device discovery topic, restores the legacy configs, then clears the device config and returns to single-component mode. A failed rollback remains pending and suppresses automatic device-bundle publication until retried successfully. Run all publishing stages from `loop()`, not inside an inbound MQTT callback. Device discovery can also publish richer origin/device metadata, for example: @@ -56,6 +106,7 @@ mqtt.setOriginSupportUrl("https://example.com/device-help"); ``` For entity identifiers in Home Assistant, prefer **`setDefaultEntityId()`** over legacy **`setObjectId()`**. +`setObjectId()` remains source-compatible but no longer serializes the removed MQTT `obj_id` property. `setDefaultEntityId()` influences Home Assistant's entity ID only when it first creates the entity; existing users can retain a customized entity ID in the entity registry. Common entity discovery metadata can be configured on most entity types via: @@ -70,9 +121,15 @@ Common entity discovery metadata can be configured on most entity types via: After changing discovery-related settings at runtime: - `HABaseDeviceType::republishDiscovery()` to refresh discovery. -- `HABaseDeviceType::removeFromDiscovery()` to clear retained discovery for one entity. +- `HABaseDeviceType::removeFromDiscovery()` to remove one entity from discovery. -If device discovery mode is enabled, the library clears stale per-entity configs when refreshing. +In single-component mode, removal clears that entity's retained config topic. In device discovery mode, Home Assistant requires a platform-only component marker (`{"p":"sensor"}`, for example), followed by a compacted device bundle that omits the component. ArduinoHA performs both publishes. Call `republishDiscovery()` on that entity to add it back. + +## Identifier and lifetime checklist + +- Give `HADevice` a stable, non-empty unique ID and each entity a stable, topic-safe ID (`A-Z`, `a-z`, `0-9`, `_`, `-`). +- Use `enableExtendedUniqueIds()` when multiple devices could otherwise reuse the same entity IDs. +- Metadata setters store pointers; retain their backing strings. `HAText` copies the current state and command text it receives, so its state does not borrow a transient MQTT buffer. ## Trimming flash (optional) diff --git a/docs/entities.md b/docs/entities.md index f1078e5..5e91799 100644 --- a/docs/entities.md +++ b/docs/entities.md @@ -26,7 +26,7 @@ The library does not currently implement alarm control panels, events, humidifie ## Common lifecycle -Create `HADevice`, `HAMqtt`, then entities in that order; create them before `mqtt.begin(...)`. Call `mqtt.loop()` regularly. Read [Getting started](getting-started.md) for the connection flow and [device discovery](device-and-discovery.md) for discovery settings. +Create long-lived `HADevice`, `HAMqtt`, and entity objects, then call `mqtt.begin(...)` once and `mqtt.loop()` regularly. Entities may be constructed before or after `HAMqtt`; both orders register safely. Read [Getting started](getting-started.md) for the connection flow and [device discovery](device-and-discovery.md) for discovery settings. For shared online/offline state, use `device.enableSharedAvailability()` and `device.enableLastWill()` before connecting. The [availability examples](../examples/availability/) show the simplest setup, while [advanced availability](../examples/advanced-availability/) covers custom payloads and Last Will behaviour. diff --git a/docs/getting-started.md b/docs/getting-started.md index 51f6e5b..93e3fe6 100644 --- a/docs/getting-started.md +++ b/docs/getting-started.md @@ -11,7 +11,7 @@ ArduinoHA talks to Home Assistant over **MQTT** (TCP). You need an MQTT broker r ```ini lib_deps = - home-assistant-integration=https://github.com/alexhopeoconnor/arduino-home-assistant.git#v3.0.2 + home-assistant-integration=https://github.com/alexhopeoconnor/arduino-home-assistant.git#v3.1.0 ``` **Arduino IDE:** this fork is not indexed by Library Manager. Download the @@ -26,7 +26,7 @@ Arduino networking API. 1. Create **`HADevice`** and **`HAMqtt`** once (global or inside a long-lived object). 2. Call **`HAMqtt::begin(...)`** once, at the **end** of `setup()` — it only stores broker settings; the actual connection runs during **`HAMqtt::loop()`**. 3. Call **`mqtt.loop()`** regularly in `loop()` (not necessarily every iteration). -4. Construct **entity classes** (sensors, switches, …) **after** `HAMqtt`, and register them before `begin()` where the API requires it. +4. Construct **entity classes** (sensors, switches, …) before or after `HAMqtt`; ArduinoHA registers both orders. Keep all objects alive for the whole MQTT lifetime and configure discovery before the first connection. ### Ethernet (example) @@ -87,6 +87,10 @@ All are valid; pick one. Hostnames work instead of IP addresses. - `mqtt.begin("192.168.1.50", "user", "pass")` — credentials, port 1883 - `mqtt.begin("192.168.1.50", 8888, "user", "pass")` — credentials + custom port +`begin()` is intentionally a one-time configuration call. Reconnects are owned by `mqtt.loop()`, which uses the configured reconnect interval. Do not call `begin()` in a reconnect timer. `mqtt.disconnect()` now also produces the registered disconnected and state-change callbacks, making an explicit shutdown observable in the same way as a transport loss. + +Use Home Assistant **2024.11.0 or newer** when opting into device discovery. See [Device & discovery](device-and-discovery.md) for the required staged migration from an existing single-component device. + ## Security note Credentials go over **plain TCP** unless you use a TLS-capable stack and broker setup. On a trusted LAN this is often acceptable; treat untrusted networks accordingly. diff --git a/docs/mqtt-usage.md b/docs/mqtt-usage.md index faffdbd..20f18dd 100644 --- a/docs/mqtt-usage.md +++ b/docs/mqtt-usage.md @@ -114,6 +114,11 @@ myButton.setPayloadPress("PRESS"); Defined in `ArduinoHADefines.h` or via build flags. +- **`ARDUINOHA_DISABLE_STDFUNCTION`** — disables the default capturing-lambda and + `std::bind` overloads on entity command APIs. Use it on targets without a complete + `std::function` implementation or where code size/RAM matters (constrained AVR + builds are the usual case). ESP8266/ESP32 support the overloads. `HAMqtt` + lifecycle callbacks remain ordinary function pointers. - **`ARDUINOHA_DEBUG`** — enables ArduinoHA logging by default and sets the initial maximum verbosity to `Debug`. Without this flag, structured logs are compiled in but remain disabled until you call `arduinoHASetLogEnabled(true)`. Structured logging is available through: diff --git a/library.json b/library.json index 4b50681..be8090f 100644 --- a/library.json +++ b/library.json @@ -1,6 +1,6 @@ { "name": "home-assistant-integration", - "version": "3.0.2", + "version": "3.1.0", "description": "Maintained Home Assistant MQTT discovery and entity integration for Arduino and ESP devices.", "keywords": [ "mqtt", @@ -23,7 +23,7 @@ "maintainer": true } ], - "license": "MIT", + "license": "AGPL-3.0", "homepage": "https://github.com/alexhopeoconnor/arduino-home-assistant", "repository": { "type": "git", diff --git a/library.properties b/library.properties index 00efaad..22a4adc 100644 --- a/library.properties +++ b/library.properties @@ -1,5 +1,5 @@ name=home-assistant-integration -version=3.0.2 +version=3.1.0 author=Dawid Chyrzynski , Alex Hope-O'Connor maintainer=Alex Hope-O'Connor sentence=Home Assistant MQTT integration for Arduino diff --git a/platformio.ini b/platformio.ini index fd35c73..a705672 100644 --- a/platformio.ini +++ b/platformio.ini @@ -27,12 +27,43 @@ lib_deps = knolleary/PubSubClient@2.8.0 [env:esp8266] extends = common +test_ignore = test_native_core platform = espressif8266 board = d1_mini [env:esp32] extends = common +test_ignore = test_native_core platform = https://github.com/pioarduino/platform-espressif32/releases/download/51.03.05/platform-espressif32.zip board = esp32dev build_unflags = -std=gnu++11 build_flags = ${common.build_flags} -std=gnu++14 + +; Host-only regression checks for serialization, topic validation, and the +; PubSubClient mock path. The Arduino compatibility files live exclusively +; under test/native and are never included in firmware environments. +[env:native] +platform = native +test_framework = unity +test_build_src = yes +lib_deps = throwtheswitch/Unity@^2.6.1 +build_flags = + -DARDUINOHA_TEST + -Itest/native/include + -fsanitize=address,undefined + -fno-omit-frame-pointer +build_src_filter = + -<*> + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/scripts/bump-version.sh b/scripts/bump-version.sh index 603d196..9d87549 100755 --- a/scripts/bump-version.sh +++ b/scripts/bump-version.sh @@ -13,12 +13,18 @@ root="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" version="${tag#v}" repo_url="https://github.com/alexhopeoconnor/arduino-home-assistant.git" reference_files=(README.md docs/getting-started.md) +defines_file="$root/src/ArduinoHADefines.h" current_version="$(sed -n 's/.*"version": "\([^"]*\)".*/\1/p' "$root/library.json" | head -n 1)" [[ "$current_version" != "$version" ]] || { echo "library.json already declares $version; choose a new version." >&2 exit 1 } +version_macro_count="$(grep -Ec '^#define ARDUINOHA_LIBRARY_VERSION "[0-9]+\.[0-9]+\.[0-9]+"$' "$defines_file" || true)" +[[ "$version_macro_count" -eq 1 ]] || { + echo "src/ArduinoHADefines.h must contain exactly one ARDUINOHA_LIBRARY_VERSION macro." >&2 + exit 1 +} grep -q "^## $version$" "$root/CHANGELOG.md" && { echo "CHANGELOG.md already has a $version section; choose a new version." >&2 exit 1 @@ -26,6 +32,7 @@ grep -q "^## $version$" "$root/CHANGELOG.md" && { sed -i -E '0,/"version": "[0-9]+\.[0-9]+\.[0-9]+"/s//"version": "'"$version"'"/' "$root/library.json" sed -i -E "s/^version=[0-9]+\.[0-9]+\.[0-9]+$/version=$version/" "$root/library.properties" +sed -i -E "s/^#define ARDUINOHA_LIBRARY_VERSION \"[0-9]+\.[0-9]+\.[0-9]+\"$/#define ARDUINOHA_LIBRARY_VERSION \"$version\"/" "$defines_file" for file in "${reference_files[@]}"; do sed -i -E "s|${repo_url}#v[0-9]+\.[0-9]+\.[0-9]+|${repo_url}#v${version}|g" "$root/$file" done diff --git a/scripts/prepare-release.sh b/scripts/prepare-release.sh index 6170e9a..076ac74 100755 --- a/scripts/prepare-release.sh +++ b/scripts/prepare-release.sh @@ -18,6 +18,18 @@ if [[ "$manifest_version" != "$version" ]]; then echo "library.json is $manifest_version; expected $version for $tag" >&2 exit 1 fi +defines_file="$root/src/ArduinoHADefines.h" +version_macro_count="$(grep -Ec '^#define ARDUINOHA_LIBRARY_VERSION "[0-9]+\.[0-9]+\.[0-9]+"$' "$defines_file" || true)" +if [[ "$version_macro_count" -ne 1 ]]; then + echo "src/ArduinoHADefines.h must contain exactly one ARDUINOHA_LIBRARY_VERSION macro." >&2 + exit 1 +fi + +defines_version="$(sed -n 's/^#define ARDUINOHA_LIBRARY_VERSION "\([^"]*\)"$/\1/p' "$defines_file")" +if [[ "$defines_version" != "$version" ]]; then + echo "src/ArduinoHADefines.h is $defines_version; expected $version for $tag" >&2 + exit 1 +fi if [[ -f "$root/library.properties" ]]; then properties_version="$(sed -n 's/^version=//p' "$root/library.properties" | head -n 1)" diff --git a/src/ArduinoHADefines.h b/src/ArduinoHADefines.h index 9e8cf13..511789d 100644 --- a/src/ArduinoHADefines.h +++ b/src/ArduinoHADefines.h @@ -32,7 +32,7 @@ #endif // Current library version used in discovery origin metadata. -#define ARDUINOHA_LIBRARY_VERSION "2.1.0" +#define ARDUINOHA_LIBRARY_VERSION "3.1.0" #if defined(ARDUINOHA_DEBUG) #include diff --git a/src/HADevice.cpp b/src/HADevice.cpp index f34f0f6..59e794b 100644 --- a/src/HADevice.cpp +++ b/src/HADevice.cpp @@ -3,45 +3,136 @@ #include "HAMqtt.h" #include "utils/HAUtils.h" #include "utils/HADictionary.h" +#include "utils/HAJson.h" #include "utils/HASerializer.h" #include static bool appendEscapedJsonString(char*& cursor, char* end, const char* value) { - if (!cursor || !value || cursor >= end) { - return false; - } - - if (cursor + 1 >= end) { - return false; - } - - *cursor++ = '"'; - for (const char* p = value; *p != '\0'; p++) { - if ((*p == '"' || *p == '\\') && cursor + 2 >= end) { - return false; - } - - if (*p == '"' || *p == '\\') { - *cursor++ = '\\'; - } - - if (cursor + 1 >= end) { - return false; - } - - *cursor++ = *p; - } - - if (cursor + 1 >= end) { - return false; - } - - *cursor++ = '"'; - *cursor = 0; - return true; + return HAJson::appendEscapedString(cursor, end, value); } +namespace { +void skipJsonWhitespace(const char*& cursor) +{ + while (*cursor == ' ' || *cursor == '\t' || *cursor == '\n' || *cursor == '\r') { + cursor++; + } +} + +bool isHexDigit(const char value) +{ + return (value >= '0' && value <= '9') || + (value >= 'a' && value <= 'f') || + (value >= 'A' && value <= 'F'); +} + +bool parseJsonString(const char*& cursor) +{ + if (*cursor != '"') { + return false; + } + + cursor++; + while (*cursor != '\0') { + const unsigned char value = static_cast(*cursor++); + if (value == '"') { + return true; + } + + if (value < 0x20) { + return false; + } + + if (value != '\\') { + continue; + } + + const char escape = *cursor++; + if (escape == '\0') { + return false; + } + + if (escape == '"' || escape == '\\' || escape == '/' || + escape == 'b' || escape == 'f' || escape == 'n' || + escape == 'r' || escape == 't') { + continue; + } + + if (escape != 'u') { + return false; + } + + for (uint8_t i = 0; i < 4; i++) { + if (!isHexDigit(*cursor++)) { + return false; + } + } + } + + return false; +} + +bool isValidConnectionsJson(const char* value, const size_t maxLength) +{ + if (!value || value[0] == '\0' || strlen(value) >= maxLength) { + return false; + } + + const char* cursor = value; + skipJsonWhitespace(cursor); + if (*cursor++ != '[') { + return false; + } + + skipJsonWhitespace(cursor); + if (*cursor == ']') { + cursor++; + skipJsonWhitespace(cursor); + return *cursor == '\0'; + } + + while (true) { + if (*cursor++ != '[') { + return false; + } + + skipJsonWhitespace(cursor); + if (!parseJsonString(cursor)) { + return false; + } + + skipJsonWhitespace(cursor); + if (*cursor++ != ',') { + return false; + } + + skipJsonWhitespace(cursor); + if (!parseJsonString(cursor)) { + return false; + } + + skipJsonWhitespace(cursor); + if (*cursor++ != ']') { + return false; + } + + skipJsonWhitespace(cursor); + if (*cursor == ']') { + cursor++; + skipJsonWhitespace(cursor); + return *cursor == '\0'; + } + + if (*cursor++ != ',') { + return false; + } + + skipJsonWhitespace(cursor); + } +} +} // namespace + #define HADEVICE_INIT \ _ownsUniqueId(false), \ _serializer(new HASerializer(nullptr, 16)), \ @@ -250,14 +341,13 @@ bool HADevice::addConnection(const char* type, const char* value) return true; } -void HADevice::setConnectionsJson(const char* connectionsJson) +bool HADevice::setConnectionsJson(const char* connectionsJson) { - if (!connectionsJson || connectionsJson[0] == '\0') { - return; + if (!isValidConnectionsJson(connectionsJson, MaxConnectionsJsonLength)) { + return false; } - strncpy(_connectionsJson, connectionsJson, MaxConnectionsJsonLength - 1); - _connectionsJson[MaxConnectionsJsonLength - 1] = 0; + strcpy(_connectionsJson, connectionsJson); _hasConnections = true; if (!_connectionsPropertyRegistered) { @@ -268,6 +358,8 @@ void HADevice::setConnectionsJson(const char* connectionsJson) ); _connectionsPropertyRegistered = true; } + + return true; } void HADevice::setPayloadAvailable(const char* payload) diff --git a/src/HADevice.h b/src/HADevice.h index 78295b6..b944a13 100644 --- a/src/HADevice.h +++ b/src/HADevice.h @@ -41,6 +41,11 @@ public: */ ~HADevice(); + HADevice(const HADevice&) = delete; + HADevice& operator=(const HADevice&) = delete; + HADevice(HADevice&&) = delete; + HADevice& operator=(HADevice&&) = delete; + /** * Returns pointer to the unique ID. It can be nullptr if the device has no ID assigned. */ @@ -147,9 +152,15 @@ public: /** * Sets the `connections` array as raw JSON (e.g. [[\"mac\",\"aa:bb:cc:dd:ee:ff\"]]). - * The payload is copied into an internal buffer. + * + * This legacy escape hatch accepts only a complete JSON array of two-string + * connection tuples and rejects an oversized or malformed value rather than + * truncating a retained discovery document. Prefer addConnection(). + * + * @returns false when the value is malformed or does not fit in the + * internal discovery buffer. */ - void setConnectionsJson(const char* connectionsJson); + bool setConnectionsJson(const char* connectionsJson); void setPayloadAvailable(const char* payload); void setPayloadNotAvailable(const char* payload); diff --git a/src/HAMqtt.cpp b/src/HAMqtt.cpp index 7df1ef2..8ee51a2 100644 --- a/src/HAMqtt.cpp +++ b/src/HAMqtt.cpp @@ -1,6 +1,7 @@ #include "HAMqtt.h" #include +#include #include #ifndef ARDUINOHA_TEST @@ -13,6 +14,7 @@ #include "device-types/HABaseDeviceType.h" #include "mocks/PubSubClientMock.h" #include "utils/HADictionary.h" +#include "utils/HAJson.h" #include "utils/HASerializer.h" namespace { @@ -30,6 +32,7 @@ constexpr char kDiscovery[] = "discovery"; _discoveryPrefix(DefaultDiscoveryPrefix), \ _dataPrefix(DefaultDataPrefix), \ _deviceDiscoveryEnabled(false), \ + _deviceDiscoveryMigrationState(DeviceDiscoveryMigrationIdle), \ _originSupportUrl(nullptr), \ _username(nullptr), \ _password(nullptr), \ @@ -39,6 +42,7 @@ constexpr char kDiscovery[] = "discovery"; _maxDevicesTypesNb(maxDevicesTypesNb), \ _devicesTypes(new HABaseDeviceType*[maxDevicesTypesNb]), \ _lastWillTopic(nullptr), \ + _deviceTypeRegistrationFailures(0), \ _lastWillMessage(nullptr), \ _lastWillRetain(false), \ _currentState(StateDisconnected), \ @@ -56,6 +60,7 @@ constexpr char kDiscovery[] = "discovery"; static const char* DefaultDiscoveryPrefix = "homeassistant"; static const char* DefaultDataPrefix = "aha"; static const char* DeviceDiscoveryOriginName = "ArduinoHA"; +static const char* DeviceDiscoveryMigrationPayload = "{\"migrate_discovery\":true}"; HAMqtt* HAMqtt::_instance = nullptr; @@ -78,6 +83,7 @@ HAMqtt::HAMqtt( HAMQTT_INIT { _instance = this; + HABaseDeviceType::registerAllWith(*this); } #else HAMqtt::HAMqtt( @@ -90,6 +96,7 @@ HAMqtt::HAMqtt( HAMQTT_INIT { _instance = this; + HABaseDeviceType::registerAllWith(*this); } #endif @@ -243,6 +250,9 @@ bool HAMqtt::disconnect() _initialized = false; _lastConnectionAttemptAt = 0; _mqtt->disconnect(); + if (_currentState != StateDisconnected) { + setState(StateDisconnected); + } return true; } @@ -323,13 +333,52 @@ void HAMqtt::setReconnectInterval(uint16_t interval) } } -void HAMqtt::addDeviceType(HABaseDeviceType* deviceType) +bool HAMqtt::addDeviceType(HABaseDeviceType* deviceType) { - if (_devicesTypesNb + 1 > _maxDevicesTypesNb) { - return; + if (!deviceType) { + return false; + } + + for (uint8_t i = 0; i < _devicesTypesNb; i++) { + if (_devicesTypes[i] == deviceType) { + return true; + } + } + + if (_devicesTypesNb >= _maxDevicesTypesNb) { + _deviceTypeRegistrationFailures++; + arduinoHALog( + ArduinoHALogLevel::Error, + kMqtt, + String(F("entity registration dropped registered=")) + String(_devicesTypesNb) + + F(" limit=") + String(_maxDevicesTypesNb) + ); + return false; } _devicesTypes[_devicesTypesNb++] = deviceType; + return true; +} + +bool HAMqtt::removeDeviceType(HABaseDeviceType* deviceType) +{ + if (!deviceType) { + return false; + } + + for (uint8_t i = 0; i < _devicesTypesNb; i++) { + if (_devicesTypes[i] != deviceType) { + continue; + } + + for (uint8_t j = i + 1; j < _devicesTypesNb; j++) { + _devicesTypes[j - 1] = _devicesTypes[j]; + } + _devicesTypes[--_devicesTypesNb] = nullptr; + return true; + } + + return false; } bool HAMqtt::publish(const char* topic, const char* payload, bool retained) @@ -377,15 +426,19 @@ bool HAMqtt::publish(const char* topic, const char* payload, bool retained) ); return false; } - _mqtt->write(reinterpret_cast(payload), payloadLength); + const bool written = writePayload( + reinterpret_cast(payload), + payloadLength + ); const bool connBeforeEnd = isConnected(); const int psBeforeEnd = getPubSubState(); - const bool ok = _mqtt->endPublish(); + const bool ended = _mqtt->endPublish(); + const bool ok = written && ended; if (!ok) { arduinoHALog( ArduinoHALogLevel::Warn, kMqtt, - String(F("endPublish failed topic=")) + topic + + String(written ? F("endPublish failed topic=") : F("payload write failed topic=")) + topic + formatDirectPublishFailureDiagnostics(connBeforeEnd, psBeforeEnd) ); } else { @@ -443,18 +496,26 @@ bool HAMqtt::beginPublish( return true; } -void HAMqtt::writePayload(const char* data, const uint16_t length) +bool HAMqtt::writePayload(const char* data, const uint16_t length) { - writePayload(reinterpret_cast(data), length); + if (!data && length > 0) { + return false; + } + + return writePayload(reinterpret_cast(data), length); } -void HAMqtt::writePayload(const uint8_t* data, const uint16_t length) +bool HAMqtt::writePayload(const uint8_t* data, const uint16_t length) { + if (!data && length > 0) { + return false; + } + if (isProcessingMessage() && _deferredBuilder.active) { if (!_deferredBuilder.valid || (static_cast(_deferredBuilder.writtenLength) + length) > _deferredBuilder.expectedLength) { _deferredBuilder.valid = false; - return; + return false; } if (length > 0) { @@ -466,21 +527,25 @@ void HAMqtt::writePayload(const uint8_t* data, const uint16_t length) } _deferredBuilder.writtenLength = static_cast(_deferredBuilder.writtenLength + length); - return; + return true; } - _mqtt->write(data, length); + return _mqtt->write(data, length) == length; } -void HAMqtt::writePayload(const __FlashStringHelper* src) +bool HAMqtt::writePayload(const __FlashStringHelper* src) { + if (!src) { + return false; + } + if (isProcessingMessage() && _deferredBuilder.active) { PGM_P p = reinterpret_cast(src); const uint16_t chunkLen = static_cast(strlen_P(p)); if (!_deferredBuilder.valid || (static_cast(_deferredBuilder.writtenLength) + chunkLen) > _deferredBuilder.expectedLength) { _deferredBuilder.valid = false; - return; + return false; } if (chunkLen > 0) { @@ -488,10 +553,11 @@ void HAMqtt::writePayload(const __FlashStringHelper* src) _deferredBuilder.writtenLength = static_cast(_deferredBuilder.writtenLength + chunkLen); } - return; + return true; } - _mqtt->print(src); + const uint16_t length = static_cast(strlen_P(reinterpret_cast(src))); + return _mqtt->print(src) == length; } bool HAMqtt::endPublish() @@ -627,9 +693,255 @@ void HAMqtt::connectToServer() } } +bool HAMqtt::beginDeviceDiscoveryMigration() +{ + const char* deviceUniqueId = _device.getUniqueId(); + if ( + _deviceDiscoveryMigrationState != DeviceDiscoveryMigrationIdle || + !deviceUniqueId || + !HAJson::isValidDiscoveryTopicToken(deviceUniqueId) || + deviceUniqueId[0] == '\0' + ) { + return false; + } + + for (uint8_t i = 0; i < _devicesTypesNb; i++) { + HABaseDeviceType* deviceType = _devicesTypes[i]; + const char* uniqueId = deviceType ? deviceType->uniqueId() : nullptr; + if ( + deviceType && + deviceType->supportsDeviceDiscovery() && + uniqueId && + HAJson::isValidDiscoveryTopicToken(uniqueId) && + uniqueId[0] != '\0' + ) { + _deviceDiscoveryEnabled = true; + _deviceDiscoveryMigrationState = DeviceDiscoveryMigrationMarkersPending; + return true; + } + } + + return false; +} + +bool HAMqtt::publishDeviceDiscoveryMigrationMarker(HABaseDeviceType* deviceType) +{ + if ( + !deviceType || + !deviceType->supportsDeviceDiscovery() || + !deviceType->uniqueId() || + !HAJson::isValidDiscoveryTopicToken(deviceType->uniqueId()) + ) { + return false; + } + + const uint16_t topicLength = HASerializer::calculateConfigTopicLength( + deviceType->componentName(), + deviceType->uniqueId() + ); + if (topicLength == 0) { + return false; + } + + char topic[topicLength]; + if (!HASerializer::generateConfigTopic( + topic, + deviceType->componentName(), + deviceType->uniqueId() + )) { + return false; + } + + return publish(topic, DeviceDiscoveryMigrationPayload, true); +} + +bool HAMqtt::publishDeviceDiscoveryMigrationMarker() +{ + const char* deviceUniqueId = _device.getUniqueId(); + if ( + !_discoveryPrefix || + !deviceUniqueId || + !HAJson::isValidDiscoveryTopicToken(deviceUniqueId) + ) { + return false; + } + + const size_t topicLength = + strlen(_discoveryPrefix) + 1 + + strlen_P(HAComponentDevice) + 1 + + strlen(deviceUniqueId) + 1 + + strlen_P(HAConfigTopic) + 1; + if (topicLength > UINT16_MAX) { + return false; + } + + char topic[topicLength]; + strcpy(topic, _discoveryPrefix); + strcat_P(topic, HASerializerSlash); + strcat_P(topic, HAComponentDevice); + strcat_P(topic, HASerializerSlash); + strcat(topic, deviceUniqueId); + strcat_P(topic, HASerializerSlash); + strcat_P(topic, HAConfigTopic); + + return publish(topic, DeviceDiscoveryMigrationPayload, true); +} + +bool HAMqtt::publishDeviceDiscoveryMigrationMarkers() +{ + if ( + _deviceDiscoveryMigrationState != DeviceDiscoveryMigrationMarkersPending || + isProcessingMessage() + ) { + return false; + } + + bool hasComponents = false; + for (uint8_t i = 0; i < _devicesTypesNb; i++) { + HABaseDeviceType* deviceType = _devicesTypes[i]; + const char* uniqueId = deviceType ? deviceType->uniqueId() : nullptr; + if ( + !deviceType || + !deviceType->supportsDeviceDiscovery() || + !uniqueId || + !HAJson::isValidDiscoveryTopicToken(uniqueId) || + uniqueId[0] == '\0' + ) { + continue; + } + + hasComponents = true; + if (!publishDeviceDiscoveryMigrationMarker(deviceType)) { + return false; + } + } + + if (!hasComponents) { + return false; + } + + _deviceDiscoveryMigrationState = DeviceDiscoveryMigrationMarkersPublished; + return true; +} + +bool HAMqtt::publishDeviceDiscoveryMigrationConfig() +{ + if ( + _deviceDiscoveryMigrationState != DeviceDiscoveryMigrationMarkersPublished || + isProcessingMessage() + ) { + return false; + } + + if (!publishDeviceDiscoveryPayload()) { + return false; + } + + _deviceDiscoveryMigrationState = DeviceDiscoveryMigrationDevicePublished; + return true; +} + +bool HAMqtt::completeDeviceDiscoveryMigration() +{ + if ( + _deviceDiscoveryMigrationState != DeviceDiscoveryMigrationDevicePublished || + isProcessingMessage() + ) { + return false; + } + + bool hasComponents = false; + for (uint8_t i = 0; i < _devicesTypesNb; i++) { + HABaseDeviceType* deviceType = _devicesTypes[i]; + const char* uniqueId = deviceType ? deviceType->uniqueId() : nullptr; + if ( + !deviceType || + !deviceType->supportsDeviceDiscovery() || + !uniqueId || + !HAJson::isValidDiscoveryTopicToken(uniqueId) || + uniqueId[0] == '\0' + ) { + continue; + } + + hasComponents = true; + if (!deviceType->removeSingleComponentDiscovery()) { + return false; + } + } + + if (!hasComponents) { + return false; + } + + _deviceDiscoveryMigrationState = DeviceDiscoveryMigrationCompleted; + return true; +} + +bool HAMqtt::rollbackDeviceDiscoveryMigration() +{ + if ( + _deviceDiscoveryMigrationState == DeviceDiscoveryMigrationIdle || + isProcessingMessage() + ) { + return false; + } + + if (_deviceDiscoveryMigrationState == DeviceDiscoveryMigrationMarkersPending) { + _deviceDiscoveryEnabled = false; + _deviceDiscoveryMigrationState = DeviceDiscoveryMigrationIdle; + return true; + } + + const bool clearDeviceConfig = + _deviceDiscoveryMigrationState == DeviceDiscoveryMigrationDevicePublished || + _deviceDiscoveryMigrationState == DeviceDiscoveryMigrationCompleted || + _deviceDiscoveryMigrationState == DeviceDiscoveryMigrationRollbackPending; + if ( + _deviceDiscoveryMigrationState != DeviceDiscoveryMigrationMarkersPublished && + !clearDeviceConfig + ) { + return false; + } + + if (clearDeviceConfig) { + _deviceDiscoveryMigrationState = DeviceDiscoveryMigrationRollbackPending; + } + + if (clearDeviceConfig && !publishDeviceDiscoveryMigrationMarker()) { + return false; + } + + for (uint8_t i = 0; i < _devicesTypesNb; i++) { + HABaseDeviceType* deviceType = _devicesTypes[i]; + const char* uniqueId = deviceType ? deviceType->uniqueId() : nullptr; + if ( + !deviceType || + !deviceType->supportsDeviceDiscovery() || + !uniqueId || + !HAJson::isValidDiscoveryTopicToken(uniqueId) || + deviceType->_deviceDiscoveryRemoved + ) { + continue; + } + + if (!deviceType->publishConfig()) { + return false; + } + } + + if (clearDeviceConfig && !clearDeviceDiscoveryConfig()) { + return false; + } + + _deviceDiscoveryEnabled = false; + _deviceDiscoveryMigrationState = DeviceDiscoveryMigrationIdle; + return true; +} + void HAMqtt::onConnectedLogic() { - if (_deviceDiscoveryEnabled) { + if (_deviceDiscoveryEnabled && !isDeviceDiscoveryMigrationInProgress()) { publishDeviceDiscovery(); } @@ -646,19 +958,37 @@ void HAMqtt::onConnectedLogic() bool HAMqtt::publishDeviceDiscovery() { - if (!_device.getUniqueId()) { + if (isDeviceDiscoveryMigrationInProgress()) { + return false; + } + + return publishDeviceDiscoveryPayload(); +} + +bool HAMqtt::publishDeviceDiscoveryPayload(HABaseDeviceType* removalType) +{ + const char* deviceUniqueId = _device.getUniqueId(); + if ( + !_discoveryPrefix || + _discoveryPrefix[0] == '\0' || + !deviceUniqueId || + !HAJson::isValidDiscoveryTopicToken(deviceUniqueId) + ) { return false; } const HASerializer* deviceSerializer = _device.getSerializer(); - if (!deviceSerializer) { + const uint16_t deviceSerializerSize = deviceSerializer + ? deviceSerializer->calculateSize() + : 0; + if (deviceSerializerSize == 0 || _devicesTypesNb == 0) { return false; } HABaseDeviceType* componentTypes[_devicesTypesNb]; HASerializer* componentSerializers[_devicesTypesNb]; uint8_t componentSerializerCount = 0; - uint16_t componentsPayloadLength = 2; // {} + uint32_t componentsPayloadLength = 2; // {} for (uint8_t i = 0; i < _devicesTypesNb; i++) { HABaseDeviceType* deviceType = _devicesTypes[i]; @@ -666,20 +996,66 @@ bool HAMqtt::publishDeviceDiscovery() continue; } - HASerializer* serializer = deviceType->buildDeviceDiscoverySerializer(); - if (!serializer) { + const char* uniqueId = deviceType->uniqueId(); + if (!uniqueId || !HAJson::isValidDiscoveryTopicToken(uniqueId)) { + arduinoHALog(ArduinoHALogLevel::Error, kDiscovery, F("device discovery rejected invalid component ID")); + for (uint8_t j = 0; j < componentSerializerCount; j++) { + delete componentSerializers[j]; + } + return false; + } + + if (deviceType->_deviceDiscoveryRemoved && deviceType != removalType) { continue; } + HASerializer* serializer = nullptr; + if (deviceType == removalType) { + serializer = new (std::nothrow) HASerializer(deviceType, 1); + if (serializer) { + serializer->set( + AHATOFSTR(HAPlatformProperty), + deviceType->componentName(), + HASerializer::ProgmemPropertyValue + ); + } + } else { + serializer = deviceType->buildDeviceDiscoverySerializer(); + } + if (!serializer) { + arduinoHALog(ArduinoHALogLevel::Error, kDiscovery, F("device discovery component serializer unavailable")); + for (uint8_t j = 0; j < componentSerializerCount; j++) { + delete componentSerializers[j]; + } + return false; + } + + const uint16_t serializerSize = serializer->calculateSize(); + const uint16_t componentKeySize = HAJson::calculateEscapedStringSize(uniqueId); + if ( + serializerSize == 0 || + componentKeySize == 0 || + componentKeySize != static_cast(strlen(uniqueId)) + 2 + ) { + delete serializer; + arduinoHALog(ArduinoHALogLevel::Error, kDiscovery, F("device discovery component serializer invalid or too large")); + for (uint8_t j = 0; j < componentSerializerCount; j++) { + delete componentSerializers[j]; + } + return false; + } + if (componentSerializerCount > 0) { componentsPayloadLength += strlen_P(HASerializerJsonPropertiesSeparator); } - - componentsPayloadLength += - strlen_P(HASerializerJsonPropertyPrefix) + - strlen(deviceType->uniqueId()) + - strlen_P(HASerializerJsonPropertySuffix) + - serializer->calculateSize(); + componentsPayloadLength += componentKeySize + 1 + serializerSize; + if (componentsPayloadLength > UINT16_MAX) { + delete serializer; + for (uint8_t j = 0; j < componentSerializerCount; j++) { + delete componentSerializers[j]; + } + return false; + } componentTypes[componentSerializerCount] = deviceType; componentSerializers[componentSerializerCount++] = serializer; @@ -689,71 +1065,71 @@ bool HAMqtt::publishDeviceDiscovery() return false; } - char originPayload[192]; - originPayload[0] = 0; + HASerializer originSerializer(nullptr, 3); + originSerializer.set(AHATOFSTR(HANameProperty), DeviceDiscoveryOriginName); + originSerializer.set( + AHATOFSTR(HADeviceSoftwareVersionProperty), + ARDUINOHA_LIBRARY_VERSION + ); if (_originSupportUrl && _originSupportUrl[0] != '\0') { - snprintf( - originPayload, - sizeof(originPayload), - "{\"name\":\"%s\",\"sw\":\"%s\",\"url\":\"%s\"}", - DeviceDiscoveryOriginName, - ARDUINOHA_LIBRARY_VERSION, - _originSupportUrl - ); - } else { - snprintf( - originPayload, - sizeof(originPayload), - "{\"name\":\"%s\",\"sw\":\"%s\"}", - DeviceDiscoveryOriginName, - ARDUINOHA_LIBRARY_VERSION - ); + originSerializer.set(AHATOFSTR(HAOriginSupportUrlProperty), _originSupportUrl); } - const uint16_t originPayloadLength = strlen(originPayload); - - const uint16_t topicLength = - strlen(_discoveryPrefix) + 1 + - strlen_P(HAComponentDevice) + 1 + - strlen(_device.getUniqueId()) + 1 + - strlen_P(HAConfigTopic) + 1; - if (topicLength == 0) { + const uint16_t originSerializerSize = originSerializer.calculateSize(); + if (originSerializerSize == 0) { for (uint8_t i = 0; i < componentSerializerCount; i++) { delete componentSerializers[i]; } - return false; } - const uint16_t payloadLength = + const uint32_t topicLength = + strlen(_discoveryPrefix) + 1 + + strlen_P(HAComponentDevice) + 1 + + strlen(deviceUniqueId) + 1 + + strlen_P(HAConfigTopic) + 1; + if (topicLength == 0 || topicLength > UINT16_MAX) { + for (uint8_t i = 0; i < componentSerializerCount; i++) { + delete componentSerializers[i]; + } + return false; + } + + const uint32_t payloadLength = strlen_P(HASerializerJsonDataPrefix) + strlen_P(HASerializerJsonPropertyPrefix) + strlen_P(HADeviceProperty) + strlen_P(HASerializerJsonPropertySuffix) + - deviceSerializer->calculateSize() + + deviceSerializerSize + strlen_P(HASerializerJsonPropertiesSeparator) + strlen_P(HASerializerJsonPropertyPrefix) + strlen_P(HAOriginProperty) + strlen_P(HASerializerJsonPropertySuffix) + - originPayloadLength + + originSerializerSize + strlen_P(HASerializerJsonPropertiesSeparator) + strlen_P(HASerializerJsonPropertyPrefix) + strlen_P(HAComponentsProperty) + strlen_P(HASerializerJsonPropertySuffix) + componentsPayloadLength + strlen_P(HASerializerJsonDataSuffix); + if (payloadLength > UINT16_MAX) { + for (uint8_t i = 0; i < componentSerializerCount; i++) { + delete componentSerializers[i]; + } + return false; + } char topic[topicLength]; strcpy(topic, _discoveryPrefix); strcat_P(topic, HASerializerSlash); strcat_P(topic, HAComponentDevice); strcat_P(topic, HASerializerSlash); - strcat(topic, _device.getUniqueId()); + strcat(topic, deviceUniqueId); strcat_P(topic, HASerializerSlash); strcat_P(topic, HAConfigTopic); const bool discConnBefore = isConnected(); const int discPsBefore = getPubSubState(); - if (!_mqtt->beginPublish(topic, payloadLength, true)) { + if (!beginPublish(topic, static_cast(payloadLength), true)) { arduinoHALog( ArduinoHALogLevel::Warn, kDiscovery, @@ -764,50 +1140,172 @@ bool HAMqtt::publishDeviceDiscovery() for (uint8_t i = 0; i < componentSerializerCount; i++) { delete componentSerializers[i]; } - return false; } - writePayload(AHATOFSTR(HASerializerJsonDataPrefix)); - - writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)); - writePayload(AHATOFSTR(HADeviceProperty)); - writePayload(AHATOFSTR(HASerializerJsonPropertySuffix)); - deviceSerializer->flush(); - - writePayload(AHATOFSTR(HASerializerJsonPropertiesSeparator)); - writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)); - writePayload(AHATOFSTR(HAOriginProperty)); - writePayload(AHATOFSTR(HASerializerJsonPropertySuffix)); - writePayload(originPayload, originPayloadLength); - - writePayload(AHATOFSTR(HASerializerJsonPropertiesSeparator)); - writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)); - writePayload(AHATOFSTR(HAComponentsProperty)); - writePayload(AHATOFSTR(HASerializerJsonPropertySuffix)); - writePayload(AHATOFSTR(HASerializerJsonDataPrefix)); + bool written = + writePayload(AHATOFSTR(HASerializerJsonDataPrefix)) && + writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)) && + writePayload(AHATOFSTR(HADeviceProperty)) && + writePayload(AHATOFSTR(HASerializerJsonPropertySuffix)) && + deviceSerializer->flush() && + writePayload(AHATOFSTR(HASerializerJsonPropertiesSeparator)) && + writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)) && + writePayload(AHATOFSTR(HAOriginProperty)) && + writePayload(AHATOFSTR(HASerializerJsonPropertySuffix)) && + originSerializer.flush() && + writePayload(AHATOFSTR(HASerializerJsonPropertiesSeparator)) && + writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)) && + writePayload(AHATOFSTR(HAComponentsProperty)) && + writePayload(AHATOFSTR(HASerializerJsonPropertySuffix)) && + writePayload(AHATOFSTR(HASerializerJsonDataPrefix)); for (uint8_t i = 0; i < componentSerializerCount; i++) { - if (i > 0) { - writePayload(AHATOFSTR(HASerializerJsonPropertiesSeparator)); + if (written && i > 0) { + written = writePayload(AHATOFSTR(HASerializerJsonPropertiesSeparator)); } - writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)); - writePayload(componentTypes[i]->uniqueId(), strlen(componentTypes[i]->uniqueId())); - writePayload(AHATOFSTR(HASerializerJsonPropertySuffix)); - componentSerializers[i]->flush(); + if (written) { + const char* componentId = componentTypes[i]->uniqueId(); + const char quote = '"'; + const char colon = ':'; + written = writePayload("e, 1) && + writePayload(componentId, strlen(componentId)) && + writePayload("e, 1) && + writePayload(&colon, 1) && + componentSerializers[i]->flush(); + } delete componentSerializers[i]; } - writePayload(AHATOFSTR(HASerializerJsonDataSuffix)); - writePayload(AHATOFSTR(HASerializerJsonDataSuffix)); - const bool published = endPublish(); + if (written) { + written = writePayload(AHATOFSTR(HASerializerJsonDataSuffix)) && + writePayload(AHATOFSTR(HASerializerJsonDataSuffix)); + } + const bool published = endPublish() && written; if (!published) { - arduinoHALog(ArduinoHALogLevel::Warn, kDiscovery, F("device discovery endPublish failed")); + arduinoHALog(ArduinoHALogLevel::Warn, kDiscovery, F("device discovery publish failed")); } return published; } +bool HAMqtt::clearDeviceDiscoveryConfig() +{ + const char* deviceUniqueId = _device.getUniqueId(); + if ( + !_discoveryPrefix || + !deviceUniqueId || + !HAJson::isValidDiscoveryTopicToken(deviceUniqueId) + ) { + return false; + } + + const size_t topicLength = + strlen(_discoveryPrefix) + 1 + + strlen_P(HAComponentDevice) + 1 + + strlen(deviceUniqueId) + 1 + + strlen_P(HAConfigTopic) + 1; + if (topicLength > UINT16_MAX) { + return false; + } + + char topic[topicLength]; + strcpy(topic, _discoveryPrefix); + strcat_P(topic, HASerializerSlash); + strcat_P(topic, HAComponentDevice); + strcat_P(topic, HASerializerSlash); + strcat(topic, deviceUniqueId); + strcat_P(topic, HASerializerSlash); + strcat_P(topic, HAConfigTopic); + + return publish(topic, "", true); +} + +bool HAMqtt::removeDeviceDiscoveryComponent(HABaseDeviceType* deviceType) +{ + const char* uniqueId = deviceType ? deviceType->uniqueId() : nullptr; + if ( + !_deviceDiscoveryEnabled || + isDeviceDiscoveryMigrationInProgress() || + !deviceType || + !deviceType->supportsDeviceDiscovery() || + !uniqueId || + !HAJson::isValidDiscoveryTopicToken(uniqueId) || + deviceType->_deviceDiscoveryRemoved + ) { + return false; + } + + bool isRegistered = false; + bool hasOtherComponents = false; + for (uint8_t i = 0; i < _devicesTypesNb; i++) { + HABaseDeviceType* candidate = _devicesTypes[i]; + if (candidate == deviceType) { + isRegistered = true; + } + + const char* candidateUniqueId = candidate ? candidate->uniqueId() : nullptr; + if ( + candidate == deviceType || + !candidate || + !candidate->supportsDeviceDiscovery() || + !candidateUniqueId || + !HAJson::isValidDiscoveryTopicToken(candidateUniqueId) || + candidate->_deviceDiscoveryRemoved + ) { + continue; + } + + hasOtherComponents = true; + } + + if (!isRegistered || !publishDeviceDiscoveryPayload(deviceType)) { + return false; + } + + deviceType->_deviceDiscoveryRemoved = true; + if (!hasOtherComponents) { + return true; + } + + return publishDeviceDiscoveryPayload(); +} + +bool HAMqtt::republishDeviceDiscoveryComponent(HABaseDeviceType* deviceType) +{ + const char* uniqueId = deviceType ? deviceType->uniqueId() : nullptr; + if ( + !_deviceDiscoveryEnabled || + isDeviceDiscoveryMigrationInProgress() || + !deviceType || + !deviceType->supportsDeviceDiscovery() || + !uniqueId || + !HAJson::isValidDiscoveryTopicToken(uniqueId) + ) { + return false; + } + + bool isRegistered = false; + for (uint8_t i = 0; i < _devicesTypesNb; i++) { + if (_devicesTypes[i] == deviceType) { + isRegistered = true; + break; + } + } + if (!isRegistered) { + return false; + } + + const bool wasRemoved = deviceType->_deviceDiscoveryRemoved; + deviceType->_deviceDiscoveryRemoved = false; + if (publishDeviceDiscoveryPayload()) { + return true; + } + + deviceType->_deviceDiscoveryRemoved = wasRemoved; + return false; +} + void HAMqtt::setState(ConnectionState state) { ConnectionState previousState = _currentState; diff --git a/src/HAMqtt.h b/src/HAMqtt.h index a317512..eece6e6 100644 --- a/src/HAMqtt.h +++ b/src/HAMqtt.h @@ -67,6 +67,19 @@ public: StateUnauthorized = 5 }; + /** + * Explicit stages for safely migrating retained single-component discovery + * to Home Assistant device discovery. + */ + enum DeviceDiscoveryMigrationState : uint8_t { + DeviceDiscoveryMigrationIdle = 0, + DeviceDiscoveryMigrationMarkersPending, + DeviceDiscoveryMigrationMarkersPublished, + DeviceDiscoveryMigrationDevicePublished, + DeviceDiscoveryMigrationCompleted, + DeviceDiscoveryMigrationRollbackPending + }; + /** * Returns existing instance (singleton) of the HAMqtt class. * It may be a null pointer if the HAMqtt object was never constructed or it was destroyed. @@ -100,6 +113,11 @@ public: * Removes singleton of the HAMqtt class. */ ~HAMqtt(); + HAMqtt(const HAMqtt&) = delete; + HAMqtt& operator=(const HAMqtt&) = delete; + HAMqtt(HAMqtt&&) = delete; + HAMqtt& operator=(HAMqtt&&) = delete; + /** * Sets the prefix of the Home Assistant discovery topics. @@ -146,6 +164,59 @@ public: inline bool isDeviceDiscoveryEnabled() const { return _deviceDiscoveryEnabled; } + /** + * Starts an explicit Home Assistant-safe migration from retained + * single-component discovery to device discovery. + * + * This method enables device discovery but does not publish anything. Call + * the subsequent migration methods after MQTT is connected, in order. + */ + bool beginDeviceDiscoveryMigration(); + + /** + * Publishes retained migration markers to each legacy component config + * topic. The state advances only when every marker is published. + */ + bool publishDeviceDiscoveryMigrationMarkers(); + + /** + * Publishes the retained device discovery config after all legacy markers + * have been published. + */ + bool publishDeviceDiscoveryMigrationConfig(); + + /** + * Clears the retained legacy component config topics after the device + * config was published. The state advances only when every cleanup + * publish succeeds. + */ + bool completeDeviceDiscoveryMigration(); + + /** + * Reverses a staged device discovery migration. A migration marker is + * published to the device discovery topic before legacy retained configs + * are restored and the device config is cleared. + * The method never runs automatically and returns to single-component + * discovery only after every required retained publish succeeds. + */ + bool rollbackDeviceDiscoveryMigration(); + + /** + * Returns the current in-memory device discovery migration stage. + */ + inline DeviceDiscoveryMigrationState getDeviceDiscoveryMigrationState() const + { return _deviceDiscoveryMigrationState; } + + /** + * Returns true while an explicit device discovery migration awaits a + * marker, device config, or legacy-topic cleanup step. + */ + inline bool isDeviceDiscoveryMigrationInProgress() const + { + return _deviceDiscoveryMigrationState != DeviceDiscoveryMigrationIdle && + _deviceDiscoveryMigrationState != DeviceDiscoveryMigrationCompleted; + } + /** * Republishes the current MQTT device discovery payload when device * discovery is enabled. @@ -348,7 +419,31 @@ public: * @note The HAMqtt class doesn't take ownership of the given pointer. * @param deviceType Instance of the device's type (HASwitch, HABinarySensor, etc.). */ - void addDeviceType(HABaseDeviceType* deviceType); + bool addDeviceType(HABaseDeviceType* deviceType); + + /** + * Removes a destroyed entity from the connection registry. + * The MQTT instance does not own entity lifetimes. + */ + bool removeDeviceType(HABaseDeviceType* deviceType); + + /** + * Number of entities currently registered for connection callbacks. + */ + inline uint8_t getRegisteredDeviceTypeCount() const + { return _devicesTypesNb; } + + /** + * Maximum number of entities that can be registered in this HAMqtt instance. + */ + inline uint8_t getDeviceTypeLimit() const + { return _maxDevicesTypesNb; } + + /** + * Number of registrations dropped because the configured entity limit was reached. + */ + inline uint16_t getDeviceTypeRegistrationFailures() const + { return _deviceTypeRegistrationFailures; } /** * Publishes the MQTT message with given topic and payload. @@ -383,7 +478,7 @@ public: * @param data The string to publish. * @param length Length of the data (bytes). */ - void writePayload(const char* data, const uint16_t length); + bool writePayload(const char* data, const uint16_t length); /** * Writes given data to the TCP stream. @@ -393,7 +488,7 @@ public: * @param data The data to publish. * @param length Length of the data (bytes). */ - void writePayload(const uint8_t* data, const uint16_t length); + bool writePayload(const uint8_t* data, const uint16_t length); /** * Writes given progmem data to the TCP stream. @@ -402,7 +497,7 @@ public: * * @param data Progmem data to publish. */ - void writePayload(const __FlashStringHelper* data); + bool writePayload(const __FlashStringHelper* data); /** * Finishes publishing of a message. @@ -512,6 +607,31 @@ private: */ void onConnectedLogic(); + bool publishDeviceDiscoveryPayload(HABaseDeviceType* removalType = nullptr); + + bool clearDeviceDiscoveryConfig(); + + bool publishDeviceDiscoveryMigrationMarker(HABaseDeviceType* deviceType); + + /** + * Marks the retained device discovery config for Home Assistant's reverse + * migration protocol before restoring legacy component configs. + */ + bool publishDeviceDiscoveryMigrationMarker(); + + /** + * Removes a component from a device discovery payload using Home + * Assistant's required platform-only marker, followed by a compacted + * bundle when another component remains. + */ + bool removeDeviceDiscoveryComponent(HABaseDeviceType* deviceType); + + /** + * Re-adds (or refreshes) a component in the current device discovery + * payload. + */ + bool republishDeviceDiscoveryComponent(HABaseDeviceType* deviceType); + /** * Sets the state of the MQTT connection. */ @@ -604,6 +724,9 @@ private: /// Enables MQTT device discovery mode when set to true. bool _deviceDiscoveryEnabled; + /// Current in-memory stage of an explicit device discovery migration. + DeviceDiscoveryMigrationState _deviceDiscoveryMigrationState; + const char* _originSupportUrl; /// The username used for the authentication. It's set in the HAMqtt::begin method. @@ -630,6 +753,9 @@ private: /// The last will topic set by HAMqtt::setLastWill const char* _lastWillTopic; + /// Count of entity registrations rejected because the configured cap was full. + uint16_t _deviceTypeRegistrationFailures; + /// The last will message set by HAMqtt::setLastWill const char* _lastWillMessage; @@ -659,6 +785,8 @@ private: bool _deferredFlushFailedForTest = false; uint8_t _lastDeferredFlushErrorForTest = 0; #endif + + friend class HABaseDeviceType; }; #endif diff --git a/src/device-types/HABaseDeviceType.cpp b/src/device-types/HABaseDeviceType.cpp index 2de558b..da44487 100644 --- a/src/device-types/HABaseDeviceType.cpp +++ b/src/device-types/HABaseDeviceType.cpp @@ -6,6 +6,15 @@ #include "../utils/HASerializer.h" #include +HABaseDeviceType* HABaseDeviceType::_firstInstance = nullptr; + +void HABaseDeviceType::registerAllWith(HAMqtt& mqttInstance) +{ + for (HABaseDeviceType* entity = _firstInstance; entity; entity = entity->_nextInstance) { + mqttInstance.addDeviceType(entity); + } +} + HABaseDeviceType::HABaseDeviceType( const __FlashStringHelper* componentName, const char* uniqueId @@ -27,15 +36,34 @@ HABaseDeviceType::HABaseDeviceType( _payloadNotAvailable(nullptr), _availabilityMode(nullptr), _availabilityList(), - _availability(AvailabilityDefault) + _deviceDiscoveryRemoved(false), + _availability(AvailabilityDefault), + _nextInstance(_firstInstance) { - if (mqtt()) { - mqtt()->addDeviceType(this); + _firstInstance = this; + if (HAMqtt* mqttInstance = mqtt()) { + mqttInstance->addDeviceType(this); } } HABaseDeviceType::~HABaseDeviceType() { + if (_firstInstance == this) { + _firstInstance = _nextInstance; + } else { + HABaseDeviceType* previous = _firstInstance; + while (previous && previous->_nextInstance != this) { + previous = previous->_nextInstance; + } + if (previous) { + previous->_nextInstance = _nextInstance; + } + } + + if (HAMqtt* mqttInstance = mqtt()) { + mqttInstance->removeDeviceType(this); + } + destroySerializer(); } @@ -47,6 +75,25 @@ void HABaseDeviceType::setAvailability(bool online) bool HABaseDeviceType::removeFromDiscovery() { + HAMqtt* mqttInstance = mqtt(); + if (!mqttInstance) { + return false; + } + + if (mqttInstance->isDeviceDiscoveryEnabled() && supportsDeviceDiscovery()) { + return mqttInstance->removeDeviceDiscoveryComponent(this); + } + + return removeSingleComponentDiscovery(); +} + +bool HABaseDeviceType::removeSingleComponentDiscovery() +{ + HAMqtt* mqttInstance = mqtt(); + if (!mqttInstance) { + return false; + } + const uint16_t topicLength = HASerializer::calculateConfigTopicLength( componentName(), uniqueId() @@ -61,11 +108,11 @@ bool HABaseDeviceType::removeFromDiscovery() } destroySerializer(); - if (!mqtt()->beginPublish(topic, 0, true)) { + if (!mqttInstance->beginPublish(topic, 0, true)) { return false; } - return mqtt()->endPublish(); + return mqttInstance->endPublish(); } bool HABaseDeviceType::republishDiscovery() @@ -79,10 +126,7 @@ bool HABaseDeviceType::republishDiscovery() return publishConfig(); } - // Clear any stale per-entity retained config so device discovery remains - // the single source of truth for supported entities. - removeFromDiscovery(); - return mqttInstance->publishDeviceDiscovery(); + return mqttInstance->republishDeviceDiscoveryComponent(this); } HAMqtt* HABaseDeviceType::mqtt() @@ -136,8 +180,12 @@ void HABaseDeviceType::destroySerializer() bool HABaseDeviceType::publishConfig() { - buildSerializer(); + HAMqtt* mqttInstance = mqtt(); + if (!mqttInstance) { + return false; + } + buildSerializer(); if (_serializer == nullptr) { return false; } @@ -151,15 +199,19 @@ bool HABaseDeviceType::publishConfig() bool published = false; if (topicLength > 0 && dataLength > 0) { char topic[topicLength]; - HASerializer::generateConfigTopic( + if (!HASerializer::generateConfigTopic( topic, componentName(), uniqueId() - ); + )) { + destroySerializer(); + return false; + } - if (mqtt()->beginPublish(topic, dataLength, true)) { - _serializer->flush(); - published = mqtt()->endPublish(); + if (mqttInstance->beginPublish(topic, dataLength, true)) { + const bool flushed = _serializer->flush(); + const bool ended = mqttInstance->endPublish(); + published = flushed && ended; } } @@ -169,7 +221,12 @@ bool HABaseDeviceType::publishConfig() void HABaseDeviceType::publishAvailability() { - const HADevice* device = mqtt()->getDevice(); + HAMqtt* mqttInstance = mqtt(); + if (!mqttInstance) { + return; + } + + const HADevice* device = mqttInstance->getDevice(); if ( !device || device->isSharedAvailabilityEnabled() || @@ -251,13 +308,15 @@ bool HABaseDeviceType::publishAbsolute( return false; } + HAMqtt* mqttInstance = mqtt(); const uint16_t len = strlen(payload); - if (!mqtt()->beginPublish(fullTopic, len, retained)) { + if (!mqttInstance->beginPublish(fullTopic, len, retained)) { return false; } - mqtt()->writePayload(payload, len); - return mqtt()->endPublish(); + const bool written = mqttInstance->writePayload(payload, len); + const bool ended = mqttInstance->endPublish(); + return written && ended; } bool HABaseDeviceType::publishOnDataTopic( @@ -305,7 +364,12 @@ bool HABaseDeviceType::publishOnDataTopic( bool isProgmemData ) { - if (!payload) { + HAMqtt* mqttInstance = mqtt(); + if (!payload || !mqttInstance) { + return false; + } + + if (!topic) { return false; } @@ -326,14 +390,16 @@ bool HABaseDeviceType::publishOnDataTopic( return false; } - if (mqtt()->beginPublish(fullTopic, length, retained)) { + if (mqttInstance->beginPublish(fullTopic, length, retained)) { + bool written = false; if (isProgmemData) { - mqtt()->writePayload(AHATOFSTR(payload)); + written = mqttInstance->writePayload(AHATOFSTR(payload)); } else { - mqtt()->writePayload(payload, length); + written = mqttInstance->writePayload(payload, length); } - return mqtt()->endPublish(); + const bool ended = mqttInstance->endPublish(); + return written && ended; } return false; @@ -353,12 +419,6 @@ void HABaseDeviceType::setEntityIdProperty(HASerializer* serializer) const const char* defaultEntityId = nonEmptyString(_defaultEntityId); if (defaultEntityId) { serializer->set(AHATOFSTR(HADefaultEntityIdProperty), defaultEntityId); - return; - } - - const char* objectId = nonEmptyString(_objectId); - if (objectId) { - serializer->set(AHATOFSTR(HAObjectIdProperty), objectId); } } diff --git a/src/device-types/HABaseDeviceType.h b/src/device-types/HABaseDeviceType.h index e039dac..7078873 100644 --- a/src/device-types/HABaseDeviceType.h +++ b/src/device-types/HABaseDeviceType.h @@ -40,6 +40,11 @@ public: ); virtual ~HABaseDeviceType(); + HABaseDeviceType(const HABaseDeviceType&) = delete; + HABaseDeviceType& operator=(const HABaseDeviceType&) = delete; + HABaseDeviceType(HABaseDeviceType&&) = delete; + HABaseDeviceType& operator=(HABaseDeviceType&&) = delete; + /** * Returns unique ID of the device type. @@ -100,8 +105,12 @@ public: { return _defaultEntityId; } /** - * Legacy alias for the MQTT `object_id` discovery property. - * Prefer setDefaultEntityId() for new code. + * Legacy compatibility setter. Home Assistant no longer supports the MQTT + * discovery `obj_id` payload property, so this value is retained only for + * source compatibility and is not published. + * + * Use setDefaultEntityId("domain.entity_id") to suggest an entity ID on + * first discovery. It does not rename existing entities. * * @param objectId The object ID. */ @@ -171,12 +180,23 @@ public: /** * Removes this entity from MQTT discovery by publishing an empty retained - * payload on its config topic. + * payload on its config topic in single-component mode. In device + * discovery mode, publishes Home Assistant's required component-removal + * marker and updates the retained device config. */ bool removeFromDiscovery(); + /** + * Returns true after this entity was removed from the retained device + * discovery payload and before it is republished. + */ + inline bool isRemovedFromDeviceDiscovery() const + { return _deviceDiscoveryRemoved; } + /** * Republishes MQTT discovery config for this entity. + * In device discovery mode, this also re-adds an entity previously removed + * from the retained device payload. */ bool republishDiscovery(); @@ -318,8 +338,8 @@ protected: /** * Adds the preferred entity ID property to the serializer. - * `default_entity_id` takes precedence and the legacy `object_id` is only - * emitted when no default entity ID was configured. + * Only `default_entity_id` is emitted. The legacy `obj_id` field is not + * supported by current Home Assistant MQTT discovery schemas. */ void setEntityIdProperty(HASerializer* serializer) const; @@ -371,6 +391,13 @@ protected: const char* _availabilityMode; HAAvailabilityConfig _availabilityList; + /// Tracks an entity removed from the retained device discovery payload. + bool _deviceDiscoveryRemoved; + + static void registerAllWith(HAMqtt& mqtt); + + static HABaseDeviceType* _firstInstance; + private: enum Availability { AvailabilityDefault = 0, @@ -382,8 +409,13 @@ private: Availability _availability; const char* effectivePayloadAvailable() const; + /// Intrusive list entry used to register entities created before HAMqtt. + HABaseDeviceType* _nextInstance; + const char* effectivePayloadNotAvailable() const; + bool removeSingleComponentDiscovery(); + friend class HAMqtt; friend class HASerializer; }; diff --git a/src/device-types/HAText.cpp b/src/device-types/HAText.cpp index fc91148..2bc3627 100644 --- a/src/device-types/HAText.cpp +++ b/src/device-types/HAText.cpp @@ -4,6 +4,8 @@ #include "../HAMqtt.h" #include "../utils/HADictionary.h" #include "../utils/HASerializer.h" +#include +#include HAText::HAText(const char* uniqueId) : HABaseDeviceType(AHATOFSTR(HAComponentText), uniqueId), @@ -25,6 +27,40 @@ HAText::HAText(const char* uniqueId) : } +HAText::~HAText() +{ + delete[] _currentState; +} + +void HAText::setCurrentState(const char* state) +{ + setCurrentStateInternal(state); +} + +bool HAText::setCurrentStateInternal(const char* state) +{ + if (!state) { + delete[] _currentState; + _currentState = nullptr; + return true; + } + + const size_t stateLength = strlen(state); + if (stateLength > MaxCommandLength) { + return false; + } + + char* stateCopy = new (std::nothrow) char[stateLength + 1]; + if (!stateCopy) { + return false; + } + + memcpy(stateCopy, state, stateLength + 1); + delete[] _currentState; + _currentState = stateCopy; + return true; +} + void HAText::setValueTemplate(const char* valueTemplate) { _valueTemplate = valueTemplate; @@ -49,9 +85,11 @@ bool HAText::setState(const char* state, const bool force) return true; } - const bool published = publishState(state); - _currentState = state; - return published; + if (!setCurrentStateInternal(state)) { + return false; + } + + return publishState(_currentState); } void HAText::buildSerializer() @@ -222,14 +260,26 @@ void HAText::onMqttMessage( #endif ; - if (hasCommandCallback && HASerializer::compareDataTopics( - topic, - uniqueId(), - AHATOFSTR(HACommandTopic) - )) { - char value[length + 1]; + if ( + hasCommandCallback && + length <= MaxCommandLength && + (length == 0 || payload) && + HASerializer::compareDataTopics( + topic, + uniqueId(), + AHATOFSTR(HACommandTopic) + ) + ) { + char* value = new (std::nothrow) char[static_cast(length) + 1]; + if (!value) { + return; + } + value[length] = 0; - memcpy(value, payload, length); + if (length > 0) { + memcpy(value, payload, length); + } + if (_commandCallback) { _commandCallback(value, this); } @@ -238,6 +288,8 @@ void HAText::onMqttMessage( _commandStdCallback(value, this); } #endif + + delete[] value; } } diff --git a/src/device-types/HAText.h b/src/device-types/HAText.h index 567ae16..65f7468 100644 --- a/src/device-types/HAText.h +++ b/src/device-types/HAText.h @@ -28,11 +28,16 @@ public: ModePassword }; + /// Maximum number of bytes accepted in a text command payload. + static const uint16_t MaxCommandLength = 255; + /** * @param uniqueId The unique ID of the text entity. It needs to be unique in a scope of your device. */ HAText(const char* uniqueId); + ~HAText() override; + /** * Changes state of the text and publishes MQTT message. * Please note that if a new value is the same as previous one, @@ -51,8 +56,7 @@ public: * * @param state New state of the text. */ - inline void setCurrentState(const char* state) - { _currentState = state; } + void setCurrentState(const char* state); /** * Returns last known state of the text. @@ -174,6 +178,14 @@ private: */ bool publishState(const char* state); + /** + * Copies state into the owned current-state buffer, or clears it when state + * is nullptr. + * + * @returns Returns false when allocating the copy fails. + */ + bool setCurrentStateInternal(const char* state); + /** * Returns progmem string representing mode of the text. */ @@ -203,8 +215,8 @@ private: const char* _valueTemplate; const char* _commandTemplate; - /// The current state of the text. It can be nullptr if state wasn't set. - const char* _currentState; + /// Owned current state of the text. It can be nullptr if state wasn't set. + char* _currentState; /// The callback that will be called when command is received from the HA. HATEXT_CALLBACK(_commandCallback); diff --git a/src/mocks/PubSubClientMock.cpp b/src/mocks/PubSubClientMock.cpp index bc17e9e..8c209cb 100644 --- a/src/mocks/PubSubClientMock.cpp +++ b/src/mocks/PubSubClientMock.cpp @@ -2,6 +2,7 @@ #ifdef ARDUINOHA_TEST #include "../ArduinoHADefines.h" +#include PubSubClientMock::PubSubClientMock() : _pendingMessage(nullptr), @@ -156,12 +157,25 @@ bool PubSubClientMock::beginPublish( size_t PubSubClientMock::write(const uint8_t *buffer, size_t size) { - if (!_pendingMessage || !_pendingMessage->buffer) { + if (!_pendingMessage || !_pendingMessage->buffer || !buffer) { return 0; } - strncat(_pendingMessage->buffer, (const char*)buffer, size); - return size; + const size_t capacity = _pendingMessage->bufferSize - 1; + if (_pendingMessage->writtenSize >= capacity) { + return 0; + } + + const size_t available = capacity - _pendingMessage->writtenSize; + const size_t written = size < available ? size : available; + if (written == 0) { + return 0; + } + + memcpy(_pendingMessage->buffer + _pendingMessage->writtenSize, buffer, written); + _pendingMessage->writtenSize += written; + _pendingMessage->buffer[_pendingMessage->writtenSize] = 0; + return written; } size_t PubSubClientMock::print(const __FlashStringHelper* buffer) @@ -184,35 +198,46 @@ int PubSubClientMock::endPublish() return 0; } - size_t messageSize = _pendingMessage->bufferSize; - uint8_t index = _flushedMessagesNb; + if (_pendingMessage->writtenSize != _pendingMessage->bufferSize - 1 || + _flushedMessagesNb == UINT8_MAX) { + return 0; + } - _flushedMessagesNb++; - _flushedMessages = static_cast( - realloc(_flushedMessages, _flushedMessagesNb * sizeof(MqttMessage*)) + MqttMessage** expanded = static_cast( + realloc(_flushedMessages, (_flushedMessagesNb + 1) * sizeof(MqttMessage*)) ); + if (!expanded) { + return 0; + } + + _flushedMessages = expanded; + _flushedMessages[_flushedMessagesNb++] = _pendingMessage; - _flushedMessages[index] = _pendingMessage; // handover memory responsibility _pendingMessage = nullptr; // do not call destructor - return messageSize; + return _flushedMessages[_flushedMessagesNb - 1]->bufferSize; } bool PubSubClientMock::subscribe(const char* topic) { - uint8_t index = _subscriptionsNb; + if (!topic || _subscriptionsNb == UINT8_MAX) { + return false; + } - _subscriptionsNb++; - _subscriptions = static_cast( - realloc(_subscriptions, _subscriptionsNb * sizeof(MqttSubscription*)) + MqttSubscription** expanded = static_cast( + realloc(_subscriptions, (_subscriptionsNb + 1) * sizeof(MqttSubscription*)) ); + if (!expanded) { + return false; + } size_t topicSize = strlen(topic) + 1; MqttSubscription* subscription = new MqttSubscription(); subscription->topic = new char[topicSize]; memcpy(subscription->topic, topic, topicSize); - _subscriptions[index] = subscription; + _subscriptions = expanded; + _subscriptions[_subscriptionsNb++] = subscription; return true; } @@ -223,7 +248,8 @@ void PubSubClientMock::clearFlushedMessages() delete _flushedMessages[i]; } - delete _flushedMessages; + free(_flushedMessages); + _flushedMessages = nullptr; } _flushedMessagesNb = 0; @@ -236,7 +262,8 @@ void PubSubClientMock::clearSubscriptions() delete _subscriptions[i]; } - delete _subscriptions; + free(_subscriptions); + _subscriptions = nullptr; } _subscriptionsNb = 0; diff --git a/src/mocks/PubSubClientMock.h b/src/mocks/PubSubClientMock.h index 58daf30..3c47122 100644 --- a/src/mocks/PubSubClientMock.h +++ b/src/mocks/PubSubClientMock.h @@ -19,6 +19,7 @@ struct MqttMessage size_t topicSize; char* buffer; size_t bufferSize; + size_t writtenSize; bool retained; MqttMessage() : @@ -26,6 +27,7 @@ struct MqttMessage topicSize(0), buffer(nullptr), bufferSize(0), + writtenSize(0), retained(false) { @@ -34,11 +36,11 @@ struct MqttMessage ~MqttMessage() { if (topic) { - delete topic; + delete[] topic; } if (buffer) { - delete buffer; + delete[] buffer; } } }; @@ -55,7 +57,7 @@ struct MqttSubscription { ~MqttSubscription() { if (topic) { - delete topic; + delete[] topic; } } }; diff --git a/src/utils/HAAvailabilityConfig.cpp b/src/utils/HAAvailabilityConfig.cpp index e046c4b..706a71d 100644 --- a/src/utils/HAAvailabilityConfig.cpp +++ b/src/utils/HAAvailabilityConfig.cpp @@ -2,20 +2,17 @@ #include #include "HAAvailabilityConfig.h" #include "HADictionary.h" +#include "HAJson.h" static uint16_t jsonEscapedStringSize(const char* s) { - if (!s) { - return 0; - } - return 2 * strlen_P(HASerializerJsonEscapeChar) + strlen(s); + return HAJson::calculateEscapedStringSize(s); } -static void appendEscapedString(char* buf, const char* s) +static bool appendEscapedString(char* buf, char* end, const char* s) { - strcat_P(buf, HASerializerJsonEscapeChar); - strcat(buf, s); - strcat_P(buf, HASerializerJsonEscapeChar); + char* cursor = buf + strlen(buf); + return HAJson::appendEscapedString(cursor, end, s); } HAAvailabilityConfig::HAAvailabilityConfig() : @@ -90,12 +87,12 @@ void HAAvailabilityConfig::clear() uint16_t HAAvailabilityConfig::calculateJsonSize() const { - uint16_t size = + uint32_t size = strlen_P(HASerializerJsonArrayPrefix) + strlen_P(HASerializerJsonArraySuffix); if (_count == 0) { - return size; + return static_cast(size); } size += (_count - 1) * strlen_P(HASerializerJsonPropertiesSeparator); @@ -107,39 +104,59 @@ uint16_t HAAvailabilityConfig::calculateJsonSize() const size += strlen_P(HASerializerJsonDataSuffix); // "t":"..." + const uint16_t topicSize = jsonEscapedStringSize(e.topic); + if (topicSize == 0) { + return 0; + } size += strlen_P(HASerializerJsonPropertyPrefix) + strlen_P(HATopic) + strlen_P(HASerializerJsonPropertySuffix) + - jsonEscapedStringSize(e.topic); + topicSize; if (e.valueTemplate && e.valueTemplate[0] != '\0') { size += strlen_P(HASerializerJsonPropertiesSeparator); + const uint16_t valueTemplateSize = jsonEscapedStringSize(e.valueTemplate); + if (valueTemplateSize == 0) { + return 0; + } size += strlen_P(HASerializerJsonPropertyPrefix) + strlen_P(HAValueTemplateProperty) + strlen_P(HASerializerJsonPropertySuffix) + - jsonEscapedStringSize(e.valueTemplate); + valueTemplateSize; } if (e.payloadAvailable && e.payloadAvailable[0] != '\0') { size += strlen_P(HASerializerJsonPropertiesSeparator); + const uint16_t payloadAvailableSize = jsonEscapedStringSize(e.payloadAvailable); + if (payloadAvailableSize == 0) { + return 0; + } size += strlen_P(HASerializerJsonPropertyPrefix) + strlen_P(HAPayloadAvailableProperty) + strlen_P(HASerializerJsonPropertySuffix) + - jsonEscapedStringSize(e.payloadAvailable); + payloadAvailableSize; } if (e.payloadNotAvailable && e.payloadNotAvailable[0] != '\0') { size += strlen_P(HASerializerJsonPropertiesSeparator); + const uint16_t payloadNotAvailableSize = jsonEscapedStringSize(e.payloadNotAvailable); + if (payloadNotAvailableSize == 0) { + return 0; + } size += strlen_P(HASerializerJsonPropertyPrefix) + strlen_P(HAPayloadNotAvailableProperty) + strlen_P(HASerializerJsonPropertySuffix) + - jsonEscapedStringSize(e.payloadNotAvailable); + payloadNotAvailableSize; } } - return size; + if (size > UINT16_MAX) { + return 0; + } + + return static_cast(size); } bool HAAvailabilityConfig::serialize(char* output) const @@ -148,6 +165,12 @@ bool HAAvailabilityConfig::serialize(char* output) const return false; } + const uint16_t jsonSize = calculateJsonSize(); + if (jsonSize == 0) { + return false; + } + + char* const end = output + jsonSize; output[0] = 0; strcat_P(output, HASerializerJsonArrayPrefix); @@ -161,7 +184,9 @@ bool HAAvailabilityConfig::serialize(char* output) const strcat_P(output, HASerializerJsonPropertyPrefix); strcat_P(output, HATopic); strcat_P(output, HASerializerJsonPropertySuffix); - appendEscapedString(output, _entries[i].topic); + if (!appendEscapedString(output, end, _entries[i].topic)) { + return false; + } const Entry& e = _entries[i]; if (e.valueTemplate && e.valueTemplate[0] != '\0') { @@ -169,21 +194,27 @@ bool HAAvailabilityConfig::serialize(char* output) const strcat_P(output, HASerializerJsonPropertyPrefix); strcat_P(output, HAValueTemplateProperty); strcat_P(output, HASerializerJsonPropertySuffix); - appendEscapedString(output, e.valueTemplate); + if (!appendEscapedString(output, end, e.valueTemplate)) { + return false; + } } if (e.payloadAvailable && e.payloadAvailable[0] != '\0') { strcat_P(output, HASerializerJsonPropertiesSeparator); strcat_P(output, HASerializerJsonPropertyPrefix); strcat_P(output, HAPayloadAvailableProperty); strcat_P(output, HASerializerJsonPropertySuffix); - appendEscapedString(output, e.payloadAvailable); + if (!appendEscapedString(output, end, e.payloadAvailable)) { + return false; + } } if (e.payloadNotAvailable && e.payloadNotAvailable[0] != '\0') { strcat_P(output, HASerializerJsonPropertiesSeparator); strcat_P(output, HASerializerJsonPropertyPrefix); strcat_P(output, HAPayloadNotAvailableProperty); strcat_P(output, HASerializerJsonPropertySuffix); - appendEscapedString(output, e.payloadNotAvailable); + if (!appendEscapedString(output, end, e.payloadNotAvailable)) { + return false; + } } strcat_P(output, HASerializerJsonDataSuffix); diff --git a/src/utils/HAAvailabilityConfig.h b/src/utils/HAAvailabilityConfig.h index 5362e8e..00b3ea3 100644 --- a/src/utils/HAAvailabilityConfig.h +++ b/src/utils/HAAvailabilityConfig.h @@ -22,6 +22,11 @@ public: HAAvailabilityConfig(); ~HAAvailabilityConfig(); + HAAvailabilityConfig(const HAAvailabilityConfig&) = delete; + HAAvailabilityConfig& operator=(const HAAvailabilityConfig&) = delete; + HAAvailabilityConfig(HAAvailabilityConfig&&) = delete; + HAAvailabilityConfig& operator=(HAAvailabilityConfig&&) = delete; + /** * Adds an availability entry. `topic` must be a full MQTT topic string. * @return false when full or topic is null/empty. diff --git a/src/utils/HAJson.cpp b/src/utils/HAJson.cpp new file mode 100644 index 0000000..8aa16c1 --- /dev/null +++ b/src/utils/HAJson.cpp @@ -0,0 +1,174 @@ +#include "HAJson.h" + +#include +#include +#include + +namespace { +uint8_t escapedByteSize(const uint8_t value) +{ + switch (value) { + case '"': + case '\\': + case '\b': + case '\f': + case '\n': + case '\r': + case '\t': + return 2; + default: + return value < 0x20 ? 6 : 1; + } +} + +char hexDigit(const uint8_t value) +{ + return value < 10 ? static_cast('0' + value) : static_cast('A' + (value - 10)); +} + +void appendEscapedByte(char*& cursor, const uint8_t value) +{ + switch (value) { + case '"': + *cursor++ = '\\'; + *cursor++ = '"'; + return; + case '\\': + *cursor++ = '\\'; + *cursor++ = '\\'; + return; + case '\b': + *cursor++ = '\\'; + *cursor++ = 'b'; + return; + case '\f': + *cursor++ = '\\'; + *cursor++ = 'f'; + return; + case '\n': + *cursor++ = '\\'; + *cursor++ = 'n'; + return; + case '\r': + *cursor++ = '\\'; + *cursor++ = 'r'; + return; + case '\t': + *cursor++ = '\\'; + *cursor++ = 't'; + return; + default: + if (value < 0x20) { + *cursor++ = '\\'; + *cursor++ = 'u'; + *cursor++ = '0'; + *cursor++ = '0'; + *cursor++ = hexDigit(static_cast(value >> 4)); + *cursor++ = hexDigit(static_cast(value & 0x0F)); + } else { + *cursor++ = static_cast(value); + } + } +} +} // namespace + +uint16_t HAJson::calculateEscapedStringSize(const char* value) +{ + if (!value) { + return 0; + } + + uint32_t size = 2; // surrounding quotes + for (const uint8_t* p = reinterpret_cast(value); *p != 0; p++) { + size += escapedByteSize(*p); + if (size > UINT16_MAX) { + return 0; + } + } + + return static_cast(size); +} + +uint16_t HAJson::calculateEscapedProgmemStringSize(const char* value) +{ + if (!value) { + return 0; + } + + uint32_t size = 2; // surrounding quotes + for (uint16_t i = 0; ; i++) { + const uint8_t byte = pgm_read_byte(value + i); + if (byte == 0) { + break; + } + + size += escapedByteSize(byte); + if (size > UINT16_MAX) { + return 0; + } + } + + return static_cast(size); +} + +bool HAJson::appendEscapedString(char*& cursor, char* end, const char* value) +{ + if (!cursor || !end || !value || cursor > end) { + return false; + } + + const uint16_t size = calculateEscapedStringSize(value); + if (size == 0 || static_cast(end - cursor) < size) { + return false; + } + + *cursor++ = '"'; + for (const uint8_t* p = reinterpret_cast(value); *p != 0; p++) { + appendEscapedByte(cursor, *p); + } + *cursor++ = '"'; + *cursor = 0; + return true; +} + +bool HAJson::appendEscapedProgmemString(char*& cursor, char* end, const char* value) +{ + if (!cursor || !end || !value || cursor > end) { + return false; + } + + const uint16_t size = calculateEscapedProgmemStringSize(value); + if (size == 0 || static_cast(end - cursor) < size) { + return false; + } + + *cursor++ = '"'; + for (uint16_t i = 0; ; i++) { + const uint8_t byte = pgm_read_byte(value + i); + if (byte == 0) { + break; + } + appendEscapedByte(cursor, byte); + } + *cursor++ = '"'; + *cursor = 0; + return true; +} + +bool HAJson::isValidDiscoveryTopicToken(const char* value) +{ + if (!value || value[0] == '\0') { + return false; + } + + for (const unsigned char* p = reinterpret_cast(value); *p != 0; p++) { + const bool isLower = *p >= 'a' && *p <= 'z'; + const bool isUpper = *p >= 'A' && *p <= 'Z'; + const bool isDigit = *p >= '0' && *p <= '9'; + if (!isLower && !isUpper && !isDigit && *p != '_' && *p != '-') { + return false; + } + } + + return true; +} diff --git a/src/utils/HAJson.h b/src/utils/HAJson.h new file mode 100644 index 0000000..efa0c5d --- /dev/null +++ b/src/utils/HAJson.h @@ -0,0 +1,46 @@ +#ifndef AHA_JSON_H +#define AHA_JSON_H + +#include + +/** + * Small JSON helpers used by discovery serializers. + * + * The library streams discovery documents directly to PubSubClient, so the + * calculated size must exactly match the escaped JSON representation before a + * retained payload is started. + */ +namespace HAJson +{ + /** + * Returns the number of bytes needed to encode value as a JSON string, + * including its surrounding quotes. Returns zero for a null value or when + * the result cannot fit in a uint16_t. + */ + uint16_t calculateEscapedStringSize(const char* value); + + /** + * As calculateEscapedStringSize(), but reads value from program memory. + */ + uint16_t calculateEscapedProgmemStringSize(const char* value); + + /** + * Appends a program-memory value as a JSON string to [cursor, end). + */ + bool appendEscapedProgmemString(char*& cursor, char* end, const char* value); + + /** + * Appends value as a JSON string to [cursor, end). end points at the last + * usable byte for the terminating null character. The output is left + * unchanged when the complete escaped value cannot fit. + */ + bool appendEscapedString(char*& cursor, char* end, const char* value); + + /** + * Home Assistant discovery node/object IDs may only contain these topic + * token characters: [A-Za-z0-9_-]. + */ + bool isValidDiscoveryTopicToken(const char* value); +} + +#endif diff --git a/src/utils/HASerializer.cpp b/src/utils/HASerializer.cpp index e023ae6..ffb6f90 100644 --- a/src/utils/HASerializer.cpp +++ b/src/utils/HASerializer.cpp @@ -12,9 +12,112 @@ #include "../HAMqtt.h" #include "../utils/HAUtils.h" #include "../utils/HANumeric.h" +#include "../utils/HAJson.h" #include "../utils/HAAvailabilityConfig.h" #include "../device-types/HABaseDeviceType.h" +namespace { +bool writeJsonEscapedByte(HAMqtt* mqtt, const uint8_t value) +{ + char output[6]; + uint8_t length = 0; + + switch (value) { + case '"': + output[0] = '\\'; + output[1] = '"'; + length = 2; + break; + case '\\': + output[0] = '\\'; + output[1] = '\\'; + length = 2; + break; + case '\b': + output[0] = '\\'; + output[1] = 'b'; + length = 2; + break; + case '\f': + output[0] = '\\'; + output[1] = 'f'; + length = 2; + break; + case '\n': + output[0] = '\\'; + output[1] = 'n'; + length = 2; + break; + case '\r': + output[0] = '\\'; + output[1] = 'r'; + length = 2; + break; + case '\t': + output[0] = '\\'; + output[1] = 't'; + length = 2; + break; + default: + if (value < 0x20) { + static const char hex[] = "0123456789ABCDEF"; + output[0] = '\\'; + output[1] = 'u'; + output[2] = '0'; + output[3] = '0'; + output[4] = hex[value >> 4]; + output[5] = hex[value & 0x0F]; + length = 6; + } else { + output[0] = static_cast(value); + length = 1; + } + break; + } + + return mqtt && mqtt->writePayload(output, length); +} + +bool writeJsonString(HAMqtt* mqtt, const char* value, const bool progmem) +{ + if (!mqtt || !value) { + return false; + } + + const char quote = '"'; + if (!mqtt->writePayload("e, 1)) { + return false; + } + for (size_t i = 0; ; i++) { + const uint8_t byte = progmem + ? pgm_read_byte(value + i) + : static_cast(value[i]); + if (byte == 0) { + break; + } + if (!writeJsonEscapedByte(mqtt, byte)) { + return false; + } + } + return mqtt->writePayload("e, 1); +} + +bool writeJsonStringContents(HAMqtt* mqtt, const char* value) +{ + if (!mqtt || !value) { + return false; + } + + for (size_t i = 0; value[i] != '\0'; i++) { + if (!writeJsonEscapedByte(mqtt, static_cast(value[i]))) { + return false; + } + } + + return true; +} +} // namespace + uint16_t HASerializer::calculateConfigTopicLength( const __FlashStringHelper* componentName, const char* objectId @@ -26,7 +129,10 @@ uint16_t HASerializer::calculateConfigTopicLength( !objectId || !mqtt || !mqtt->getDiscoveryPrefix() || - !mqtt->getDevice() + !mqtt->getDevice() || + !mqtt->getDevice()->getUniqueId() || + !HAJson::isValidDiscoveryTopicToken(mqtt->getDevice()->getUniqueId()) || + !HAJson::isValidDiscoveryTopicToken(objectId) ) { return 0; } @@ -52,7 +158,10 @@ bool HASerializer::generateConfigTopic( !objectId || !mqtt || !mqtt->getDiscoveryPrefix() || - !mqtt->getDevice() + !mqtt->getDevice() || + !mqtt->getDevice()->getUniqueId() || + !HAJson::isValidDiscoveryTopicToken(mqtt->getDevice()->getUniqueId()) || + !HAJson::isValidDiscoveryTopicToken(objectId) ) { return false; } @@ -83,7 +192,8 @@ uint16_t HASerializer::calculateDataTopicLength( !topic || !mqtt || !mqtt->getDataPrefix() || - !mqtt->getDevice() + !mqtt->getDevice() || + !mqtt->getDevice()->getUniqueId() ) { return 0; } @@ -112,7 +222,8 @@ bool HASerializer::generateDataTopic( !topic || !mqtt || !mqtt->getDataPrefix() || - !mqtt->getDevice() + !mqtt->getDevice() || + !mqtt->getDevice()->getUniqueId() ) { return false; } @@ -259,14 +370,14 @@ HASerializer::SerializerEntry* HASerializer::addEntry() uint16_t HASerializer::calculateSize() const { - uint16_t size = + uint32_t size = strlen_P(HASerializerJsonDataPrefix) + strlen_P(HASerializerJsonDataSuffix); for (uint8_t i = 0; i < _entriesNb; i++) { const uint16_t entrySize = calculateEntrySize(&_entries[i]); if (entrySize == 0) { - continue; + return 0; } size += entrySize; @@ -275,9 +386,13 @@ uint16_t HASerializer::calculateSize() const if (i > 0) { size += strlen_P(HASerializerJsonPropertiesSeparator); } + + if (size > UINT16_MAX) { + return 0; + } } - return size; + return static_cast(size); } bool HASerializer::flush() const @@ -287,11 +402,19 @@ bool HASerializer::flush() const return false; } - mqtt->writePayload(AHATOFSTR(HASerializerJsonDataPrefix)); + if (calculateSize() == 0) { + return false; + } + + if (!mqtt->writePayload(AHATOFSTR(HASerializerJsonDataPrefix))) { + return false; + } for (uint8_t i = 0; i < _entriesNb; i++) { if (i > 0) { - mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertiesSeparator)); + if (!mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertiesSeparator))) { + return false; + } } if (!flushEntry(&_entries[i])) { @@ -299,21 +422,27 @@ bool HASerializer::flush() const } } - mqtt->writePayload(AHATOFSTR(HASerializerJsonDataSuffix)); - return true; + return mqtt->writePayload(AHATOFSTR(HASerializerJsonDataSuffix)); } uint16_t HASerializer::calculateEntrySize(const SerializerEntry* entry) const { switch (entry->type) { - case PropertyEntryType: - return - // property name + case PropertyEntryType: { + if (!entry->property) { + return 0; + } + const uint16_t valueSize = calculatePropertyValueSize(entry); + if (valueSize == 0) { + return 0; + } + const uint32_t size = strlen_P(HASerializerJsonPropertyPrefix) + strlen_P(AHAFROMFSTR(entry->property)) + strlen_P(HASerializerJsonPropertySuffix) + - // property value - calculatePropertyValueSize(entry); + valueSize; + return size > UINT16_MAX ? 0 : static_cast(size); + } case TopicEntryType: return calculateTopicEntrySize(entry); @@ -335,56 +464,74 @@ uint16_t HASerializer::calculateTopicEntrySize( const SerializerEntry* entry ) const { - uint16_t size = 0; - - // property name - size += + uint32_t size = strlen_P(HASerializerJsonPropertyPrefix) + strlen_P(AHAFROMFSTR(entry->property)) + strlen_P(HASerializerJsonPropertySuffix); - // topic escape - size += 2 * strlen_P(HASerializerJsonEscapeChar); - - // topic + uint16_t topicSize = 0; if (entry->value) { - size += strlen(static_cast(entry->value)); + topicSize = HAJson::calculateEscapedStringSize( + static_cast(entry->value) + ); } else { - if (!_deviceType) { + if (!_deviceType || !_deviceType->uniqueId()) { return 0; } - size += calculateDataTopicLength( + const uint16_t length = calculateDataTopicLength( _deviceType->uniqueId(), entry->property - ) - 1; // exclude null terminator + ); + if (length == 0) { + return 0; + } + + char topic[length]; + if (!generateDataTopic(topic, _deviceType->uniqueId(), entry->property)) { + return 0; + } + topicSize = HAJson::calculateEscapedStringSize(topic); } - return size; + if (topicSize == 0 || (size + topicSize) > UINT16_MAX) { + return 0; + } + + return static_cast(size + topicSize); } uint16_t HASerializer::calculateAvailabilityArrayEntrySize( const SerializerEntry* entry ) const { - if (!entry->value) { + if (!entry->value || !entry->property) { return 0; } const HAAvailabilityConfig* cfg = static_cast( entry->value ); + const uint16_t jsonSize = cfg->calculateJsonSize(); + if (jsonSize == 0) { + return 0; + } - return + const uint32_t size = strlen_P(HASerializerJsonPropertyPrefix) + strlen_P(AHAFROMFSTR(entry->property)) + strlen_P(HASerializerJsonPropertySuffix) + - cfg->calculateJsonSize(); + jsonSize; + + return size > UINT16_MAX ? 0 : static_cast(size); } uint16_t HASerializer::calculateFlagSize(const FlagType flag) const { const HAMqtt* mqtt = HAMqtt::instance(); + if (!mqtt || !mqtt->getDevice()) { + return 0; + } const HADevice* device = mqtt->getDevice(); if (flag == WithDevice && device->getSerializer()) { @@ -393,26 +540,44 @@ uint16_t HASerializer::calculateFlagSize(const FlagType flag) const return 0; } - return + const uint32_t size = strlen_P(HASerializerJsonPropertyPrefix) + strlen_P(HADeviceProperty) + strlen_P(HASerializerJsonPropertySuffix) + deviceLength; - } else if (flag == WithUniqueId && _deviceType) { - uint16_t uniqueIdLength = strlen(_deviceType->uniqueId()); - - if (device->isExtendedUniqueIdsEnabled()) { - uniqueIdLength += strlen(device->getUniqueId()) + 1; // with separator + return size > UINT16_MAX ? 0 : static_cast(size); + } else if (flag == WithUniqueId && _deviceType && _deviceType->uniqueId()) { + const uint16_t uniqueIdSize = HAJson::calculateEscapedStringSize( + _deviceType->uniqueId() + ); + if (uniqueIdSize == 0) { + return 0; } - return - // property name + uint32_t valueSize = uniqueIdSize; + + if (device->isExtendedUniqueIdsEnabled()) { + if (!device->getUniqueId()) { + return 0; + } + + const uint16_t deviceIdSize = HAJson::calculateEscapedStringSize( + device->getUniqueId() + ); + if (deviceIdSize == 0) { + return 0; + } + + // Both helper sizes include quotes; the combined value has one pair. + valueSize = deviceIdSize + uniqueIdSize - 1; + } + + const uint32_t size = strlen_P(HASerializerJsonPropertyPrefix) + strlen_P(HAUniqueIdProperty) + strlen_P(HASerializerJsonPropertySuffix) + - // property value - 2 * strlen_P(HASerializerJsonEscapeChar) + - uniqueIdLength; + valueSize; + return size > UINT16_MAX ? 0 : static_cast(size); } return 0; @@ -426,9 +591,9 @@ uint16_t HASerializer::calculatePropertyValueSize( case ConstCharPropertyValue: case ProgmemPropertyValue: { const char* value = static_cast(entry->value); - const uint16_t len = - entry->subtype == ConstCharPropertyValue ? strlen(value) : strlen_P(value); - return 2 * strlen_P(HASerializerJsonEscapeChar) + len; + return entry->subtype == ConstCharPropertyValue + ? HAJson::calculateEscapedStringSize(value) + : HAJson::calculateEscapedProgmemStringSize(value); } case BoolPropertyType: { @@ -447,7 +612,7 @@ uint16_t HASerializer::calculatePropertyValueSize( const HASerializerArray* array = static_cast( entry->value ); - return array->calculateSize(); + return array ? array->calculateSize() : 0; } case JsonLiteralPropertyValue: { @@ -466,9 +631,11 @@ bool HASerializer::flushEntry(const SerializerEntry* entry) const switch (entry->type) { case PropertyEntryType: { - mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)); - mqtt->writePayload(entry->property); - mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertySuffix)); + if (!mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)) || + !mqtt->writePayload(entry->property) || + !mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertySuffix))) { + return false; + } return flushEntryValue(entry); } @@ -495,22 +662,12 @@ bool HASerializer::flushEntryValue(const SerializerEntry* entry) const case ConstCharPropertyValue: case ProgmemPropertyValue: { const char* value = static_cast(entry->value); - mqtt->writePayload(AHATOFSTR(HASerializerJsonEscapeChar)); - - if (entry->subtype == ConstCharPropertyValue) { - mqtt->writePayload(value, strlen(value)); - } else { - mqtt->writePayload(AHATOFSTR(value)); - } - - mqtt->writePayload(AHATOFSTR(HASerializerJsonEscapeChar)); - return true; + return writeJsonString(mqtt, value, entry->subtype == ProgmemPropertyValue); } case BoolPropertyType: { const bool value = *static_cast(entry->value); - mqtt->writePayload(AHATOFSTR(value ? HATrue : HAFalse)); - return true; + return mqtt->writePayload(AHATOFSTR(value ? HATrue : HAFalse)); } case NumberPropertyType: { @@ -521,21 +678,34 @@ bool HASerializer::flushEntryValue(const SerializerEntry* entry) const char tmp[HANumeric::MaxDigitsNb + 1]; const uint16_t length = value->toStr(tmp); - mqtt->writePayload(tmp, length); - return true; + return mqtt->writePayload(tmp, length); } case ArrayPropertyType: { const HASerializerArray* array = static_cast( entry->value ); - const uint16_t size = array->calculateSize(); - char tmp[size + 1]; // including null terminator - tmp[0] = 0; - array->serialize(tmp); - mqtt->writePayload(tmp, size); + if (!array) { + return false; + } - return true; + const uint16_t size = array->calculateSize(); + if (size == 0) { + return false; + } + + char* tmp = new (std::nothrow) char[size + 1]; + if (!tmp) { + return false; + } + tmp[0] = 0; + bool serialized = array->serialize(tmp); + if (serialized) { + serialized = mqtt->writePayload(tmp, size); + } + delete[] tmp; + + return serialized; } case JsonLiteralPropertyValue: { @@ -544,8 +714,7 @@ bool HASerializer::flushEntryValue(const SerializerEntry* entry) const return false; } - mqtt->writePayload(value, strlen(value)); - return true; + return mqtt->writePayload(value, strlen(value)); } default: @@ -558,17 +727,20 @@ bool HASerializer::flushTopic(const SerializerEntry* entry) const HAMqtt* mqtt = HAMqtt::instance(); // property name - mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)); - mqtt->writePayload(entry->property); - mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertySuffix)); - - // value (escaped) - mqtt->writePayload(AHATOFSTR(HASerializerJsonEscapeChar)); + if (!mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)) || + !mqtt->writePayload(entry->property) || + !mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertySuffix))) { + return false; + } if (entry->value) { const char* topic = static_cast(entry->value); - mqtt->writePayload(topic, strlen(topic)); + return writeJsonString(mqtt, topic, false); } else { + if (!_deviceType || !_deviceType->uniqueId()) { + return false; + } + const uint16_t length = calculateDataTopicLength( _deviceType->uniqueId(), entry->property @@ -578,82 +750,103 @@ bool HASerializer::flushTopic(const SerializerEntry* entry) const } char topic[length]; - generateDataTopic( + if (!generateDataTopic( topic, _deviceType->uniqueId(), entry->property - ); + )) { + return false; + } - mqtt->writePayload(topic, length - 1); + return writeJsonString(mqtt, topic, false); } - - mqtt->writePayload(AHATOFSTR(HASerializerJsonEscapeChar)); - return true; } bool HASerializer::flushAvailabilityArray(const SerializerEntry* entry) const { HAMqtt* mqtt = HAMqtt::instance(); - if (!entry->value) { + if (!mqtt || !entry->value || !entry->property) { return false; } const HAAvailabilityConfig* cfg = static_cast( entry->value ); - - mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)); - mqtt->writePayload(entry->property); - mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertySuffix)); - const uint16_t jsonSize = cfg->calculateJsonSize(); - if (jsonSize >= 512) { + if (jsonSize == 0) { return false; } - char buf[512]; - if (!cfg->serialize(buf)) { + char* buf = new (std::nothrow) char[jsonSize + 1]; + if (!buf) { return false; } - mqtt->writePayload(buf, jsonSize); - return true; + const bool serialized = cfg->serialize(buf); + if (!serialized) { + delete[] buf; + return false; + } + + const bool written = mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)) && + mqtt->writePayload(entry->property) && + mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertySuffix)) && + mqtt->writePayload(buf, jsonSize); + delete[] buf; + return written; } bool HASerializer::flushFlag(const SerializerEntry* entry) const { HAMqtt* mqtt = HAMqtt::instance(); + if (!mqtt || !mqtt->getDevice()) { + return false; + } const HADevice* device = mqtt->getDevice(); const FlagType flag = static_cast(entry->subtype); - if (flag == WithDevice && device) { + if (flag == WithDevice && device->getSerializer()) { // property name - mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)); - mqtt->writePayload(AHATOFSTR(HADeviceProperty)); - mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertySuffix)); + if (!mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)) || + !mqtt->writePayload(AHATOFSTR(HADeviceProperty)) || + !mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertySuffix))) { + return false; + } // property value return device->getSerializer()->flush(); - } else if (flag == WithUniqueId && _deviceType) { + } else if (flag == WithUniqueId && _deviceType && _deviceType->uniqueId()) { + if (device->isExtendedUniqueIdsEnabled() && !device->getUniqueId()) { + return false; + } + // property name - mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)); - mqtt->writePayload(AHATOFSTR(HAUniqueIdProperty)); - mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertySuffix)); + if (!mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertyPrefix)) || + !mqtt->writePayload(AHATOFSTR(HAUniqueIdProperty)) || + !mqtt->writePayload(AHATOFSTR(HASerializerJsonPropertySuffix))) { + return false; + } // value const char* uniqueId = _deviceType->uniqueId(); - mqtt->writePayload(AHATOFSTR(HASerializerJsonEscapeChar)); + const char quote = '"'; + if (!mqtt->writePayload("e, 1)) { + return false; + } if (device->isExtendedUniqueIdsEnabled()) { const char* deviceUniqueId = device->getUniqueId(); - mqtt->writePayload(deviceUniqueId, strlen(deviceUniqueId)); - mqtt->writePayload(AHATOFSTR(HASerializerUnderscore)); + if (!writeJsonStringContents(mqtt, deviceUniqueId)) { + return false; + } + const char separator = '_'; + if (!mqtt->writePayload(&separator, 1)) { + return false; + } } - mqtt->writePayload(uniqueId, strlen(uniqueId)); - mqtt->writePayload(AHATOFSTR(HASerializerJsonEscapeChar)); - - return true; + return writeJsonStringContents(mqtt, uniqueId) && + mqtt->writePayload("e, 1); } return false; diff --git a/src/utils/HASerializerArray.cpp b/src/utils/HASerializerArray.cpp index 20446f0..7db6038 100644 --- a/src/utils/HASerializerArray.cpp +++ b/src/utils/HASerializerArray.cpp @@ -3,6 +3,7 @@ #include "HASerializerArray.h" #include "HADictionary.h" +#include "HAJson.h" HASerializerArray::HASerializerArray(const uint8_t size, const bool progmemItems) : _progmemItems(progmemItems), @@ -39,24 +40,36 @@ const char* HASerializerArray::getItem(const uint8_t index) const uint16_t HASerializerArray::calculateSize() const { - uint16_t size = + uint32_t size = strlen_P(HASerializerJsonArrayPrefix) + strlen_P(HASerializerJsonArraySuffix); if (_itemsNb == 0) { - return size; + return static_cast(size); } // separators between elements size += (_itemsNb - 1) * strlen_P(HASerializerJsonPropertiesSeparator); for (uint8_t i = 0; i < _itemsNb; i++) { - size += - 2 * strlen_P(HASerializerJsonEscapeChar) - + (_progmemItems ? strlen_P(_items[i]) : strlen(_items[i])); + if (!_items[i]) { + return 0; + } + + const uint16_t itemSize = _progmemItems + ? HAJson::calculateEscapedProgmemStringSize(_items[i]) + : HAJson::calculateEscapedStringSize(_items[i]); + if (itemSize == 0) { + return 0; + } + + size += itemSize; + if (size > UINT16_MAX) { + return 0; + } } - return size; + return static_cast(size); } bool HASerializerArray::serialize(char* output) const @@ -65,26 +78,33 @@ bool HASerializerArray::serialize(char* output) const return false; } - strcat_P(output, HASerializerJsonArrayPrefix); + const uint16_t size = calculateSize(); + if (size == 0) { + return false; + } + + char* cursor = output; + char* const end = output + size; + *cursor++ = '['; + *cursor = 0; for (uint8_t i = 0; i < _itemsNb; i++) { if (i > 0) { - strcat_P(output, HASerializerJsonPropertiesSeparator); + *cursor++ = ','; + *cursor = 0; } - strcat_P(output, HASerializerJsonEscapeChar); - - if (_progmemItems) { - strcat_P(output, _items[i]); - } else { - strcat(output, _items[i]); + const bool serialized = _progmemItems + ? HAJson::appendEscapedProgmemString(cursor, end, _items[i]) + : HAJson::appendEscapedString(cursor, end, _items[i]); + if (!serialized) { + return false; } - - strcat_P(output, HASerializerJsonEscapeChar); } - strcat_P(output, HASerializerJsonArraySuffix); - return true; + *cursor++ = ']'; + *cursor = 0; + return static_cast(cursor - output) == size; } void HASerializerArray::clear() diff --git a/test/native/include/Arduino.h b/test/native/include/Arduino.h new file mode 100644 index 0000000..f04ae97 --- /dev/null +++ b/test/native/include/Arduino.h @@ -0,0 +1,235 @@ +#ifndef AHA_NATIVE_ARDUINO_H +#define AHA_NATIVE_ARDUINO_H + +// Minimal Arduino API shim for the host-only PlatformIO native test target. +// It is deliberately test-only: production builds continue to use each +// platform's Arduino core and PROGMEM implementation. + +#include +#include +#include +#include + +#include +#include + +typedef uint8_t byte; + +class __FlashStringHelper; + +#ifndef PROGMEM +#define PROGMEM +#endif + +#ifndef PGM_P +typedef const char* PGM_P; +#endif + +#define F(value) reinterpret_cast(value) +#ifndef pgm_read_byte +#define pgm_read_byte(address) (*reinterpret_cast(address)) +#endif + +#ifndef strlen_P +inline size_t strlen_P(PGM_P value) +{ + return value ? strlen(value) : 0; +} +#endif + +#ifndef strcpy_P +inline char* strcpy_P(char* destination, PGM_P source) +{ + return strcpy(destination, source); +} +#endif + +#ifndef strncpy_P +inline char* strncpy_P(char* destination, PGM_P source, size_t count) +{ + return strncpy(destination, source, count); +} +#endif + +#ifndef strcat_P +inline char* strcat_P(char* destination, PGM_P source) +{ + return strcat(destination, source); +} +#endif + +#ifndef strcmp_P +inline int strcmp_P(const char* left, PGM_P right) +{ + return strcmp(left, right); +} +#endif + +#ifndef memcpy_P +inline void* memcpy_P(void* destination, PGM_P source, size_t count) +{ + return memcpy(destination, source, count); +} +#endif + +class String +{ +public: + String() = default; + + String(const char* value) : + _value(value ? value : "") + { + + } + + String(const __FlashStringHelper* value) : + _value(value ? reinterpret_cast(value) : "") + { + + } + + String(const std::string& value) : + _value(value) + { + + } + + String(char value) : + _value(1, value) + { + + } + + String(bool value) : + _value(value ? "1" : "0") + { + + } + + String(int value) : _value(std::to_string(value)) { } + String(unsigned int value) : _value(std::to_string(value)) { } + String(long value) : _value(std::to_string(value)) { } + String(unsigned long value) : _value(std::to_string(value)) { } + String(long long value) : _value(std::to_string(value)) { } + String(unsigned long long value) : _value(std::to_string(value)) { } + String(float value) : _value(std::to_string(value)) { } + String(double value) : _value(std::to_string(value)) { } + + const char* c_str() const + { + return _value.c_str(); + } + + size_t length() const + { + return _value.length(); + } + + String& operator+=(const String& value) + { + _value += value._value; + return *this; + } + + String& operator+=(const char* value) + { + _value += value ? value : ""; + return *this; + } + + String& operator+=(const __FlashStringHelper* value) + { + _value += value ? reinterpret_cast(value) : ""; + return *this; + } + + String operator+(const String& value) const + { + return String(_value + value._value); + } + + String operator+(const char* value) const + { + return String(_value + (value ? value : "")); + } + + String operator+(const __FlashStringHelper* value) const + { + return String(_value + (value ? reinterpret_cast(value) : "")); + } + +private: + std::string _value; +}; + +inline String operator+(const char* left, const String& right) +{ + return String(left) + right; +} + +inline String operator+(const __FlashStringHelper* left, const String& right) +{ + return String(left) + right; +} + +class NativeSerial +{ +public: + void begin(unsigned long) { } + + template + size_t print(const T& value) + { + std::cout << value; + return 1; + } + + size_t print(const __FlashStringHelper* value) + { + std::cout << reinterpret_cast(value); + return 1; + } + + size_t print(const String& value) + { + std::cout << value.c_str(); + return value.length(); + } + + template + size_t println(const T& value) + { + print(value); + std::cout << '\n'; + return 1; + } + + size_t println() + { + std::cout << '\n'; + return 1; + } +}; + +static NativeSerial Serial; + +inline uint32_t& nativeArduinoMillisStorage() +{ + static uint32_t value = 0; + return value; +} + +inline unsigned long millis() +{ + return nativeArduinoMillisStorage(); +} + +inline void delay(unsigned long duration) +{ + nativeArduinoMillisStorage() += static_cast(duration); +} + +inline void yield() { } + +#endif diff --git a/test/native/include/Client.h b/test/native/include/Client.h new file mode 100644 index 0000000..b6436af --- /dev/null +++ b/test/native/include/Client.h @@ -0,0 +1,10 @@ +#ifndef AHA_NATIVE_CLIENT_H +#define AHA_NATIVE_CLIENT_H + +class Client +{ +public: + virtual ~Client() = default; +}; + +#endif diff --git a/test/native/include/IPAddress.h b/test/native/include/IPAddress.h new file mode 100644 index 0000000..eb18faf --- /dev/null +++ b/test/native/include/IPAddress.h @@ -0,0 +1,33 @@ +#ifndef AHA_NATIVE_IPADDRESS_H +#define AHA_NATIVE_IPADDRESS_H + +#include + +class IPAddress +{ +public: + IPAddress() : + _octets{0, 0, 0, 0} + { + + } + + IPAddress(uint8_t a, uint8_t b, uint8_t c, uint8_t d) : + _octets{a, b, c, d} + { + + } + + String toString() const + { + return String(static_cast(_octets[0])) + "." + + String(static_cast(_octets[1])) + "." + + String(static_cast(_octets[2])) + "." + + String(static_cast(_octets[3])); + } + +private: + uint8_t _octets[4]; +}; + +#endif diff --git a/test/test_device_metadata/test_main.cpp b/test/test_device_metadata/test_main.cpp new file mode 100644 index 0000000..fb0fae0 --- /dev/null +++ b/test/test_device_metadata/test_main.cpp @@ -0,0 +1,109 @@ +#include +#include + +#include +#include "mocks/PubSubClientMock.h" + +using TestFn = void (*)(void); + +struct TestCase { + const char* name; + TestFn fn; + uint16_t line; +}; + +#define TEST_ENTRY(fn) { #fn, fn, __LINE__ } + +static void flushDeviceSerializer( + PubSubClientMock* mock, + const HADevice& device +) +{ + mock->connectDummy(); + const HASerializer* serializer = device.getSerializer(); + TEST_ASSERT_NOT_NULL(serializer); + TEST_ASSERT_TRUE(mock->beginPublish("test/device/config", serializer->calculateSize(), true)); + TEST_ASSERT_TRUE(serializer->flush()); + TEST_ASSERT_TRUE(mock->endPublish()); +} + +void test_DeviceMetadata_add_connection_escapes_json() +{ + PubSubClientMock* mock = new PubSubClientMock(); + HADevice device("testDevice"); + HAMqtt mqtt(mock, device); + + TEST_ASSERT_TRUE(device.addConnection("mac", "aa\"\\bb\n")); + flushDeviceSerializer(mock, device); + + TEST_ASSERT_EQUAL_UINT8(1, mock->getFlushedMessagesNb()); + TEST_ASSERT_EQUAL_STRING( + "{\"ids\":\"testDevice\",\"cns\":[[\"mac\",\"aa\\\"\\\\bb\\n\"]]}", + mock->getFlushedMessages()[0]->buffer + ); +} + +void test_DeviceMetadata_raw_connections_reject_malformed_input() +{ + PubSubClientMock* mock = new PubSubClientMock(); + HADevice device("testDevice"); + HAMqtt mqtt(mock, device); + + TEST_ASSERT_FALSE(device.setConnectionsJson("[[\"mac\",invalid]]")); + flushDeviceSerializer(mock, device); + + TEST_ASSERT_EQUAL_STRING( + "{\"ids\":\"testDevice\"}", + mock->getFlushedMessages()[0]->buffer + ); +} + +void test_DeviceMetadata_raw_connections_accept_complete_tuple_array() +{ + PubSubClientMock* mock = new PubSubClientMock(); + HADevice device("testDevice"); + HAMqtt mqtt(mock, device); + + TEST_ASSERT_TRUE(device.setConnectionsJson("[[\"mac\",\"aa:bb\"],[\"serial\",\"42\"]]")); + flushDeviceSerializer(mock, device); + + TEST_ASSERT_EQUAL_STRING( + "{\"ids\":\"testDevice\",\"cns\":[[\"mac\",\"aa:bb\"],[\"serial\",\"42\"]]}", + mock->getFlushedMessages()[0]->buffer + ); +} + +static TestCase tests[] = { + TEST_ENTRY(test_DeviceMetadata_add_connection_escapes_json), + TEST_ENTRY(test_DeviceMetadata_raw_connections_reject_malformed_input), + TEST_ENTRY(test_DeviceMetadata_raw_connections_accept_complete_tuple_array), +}; + +static const size_t testCount = sizeof(tests) / sizeof(tests[0]); +static size_t nextTest = 0; +static bool begun = false; + +void setUp(void) { } +void tearDown(void) { } + +void setup() +{ + Serial.begin(115200); + delay(500); + UNITY_BEGIN(); + begun = true; +} + +void loop() +{ + if (begun && nextTest < testCount) { + TestCase& test = tests[nextTest++]; + UnityDefaultTestRun(test.fn, test.name, test.line); + return; + } + + if (begun) { + UNITY_END(); + begun = false; + } +} diff --git a/test/test_entities_basic/test_main.cpp b/test/test_entities_basic/test_main.cpp index 218150e..7b4996c 100644 --- a/test/test_entities_basic/test_main.cpp +++ b/test/test_entities_basic/test_main.cpp @@ -125,6 +125,10 @@ static TestCase tests[] = { TEST_ENTRY(test_TextTest_publish_state_debounce), TEST_ENTRY(test_TextTest_callback_publish_is_deferred_until_after_dispatch), TEST_ENTRY(test_TextTest_retain_setter), + TEST_ENTRY(test_TextTest_current_state_is_owned), + TEST_ENTRY(test_TextTest_oversized_state_is_rejected), + TEST_ENTRY(test_TextTest_callback_state_is_owned_after_dispatch), + TEST_ENTRY(test_TextTest_oversized_command_is_ignored), }; static const size_t TEST_COUNT = sizeof(tests) / sizeof(tests[0]); diff --git a/test/test_entities_basic/test_main.h b/test/test_entities_basic/test_main.h index 4937ff2..f239599 100644 --- a/test/test_entities_basic/test_main.h +++ b/test/test_entities_basic/test_main.h @@ -135,5 +135,9 @@ extern void test_TextTest_publish_state(void); extern void test_TextTest_publish_state_debounce(void); extern void test_TextTest_retain_setter(void); extern void test_TextTest_callback_publish_is_deferred_until_after_dispatch(void); +extern void test_TextTest_current_state_is_owned(void); +extern void test_TextTest_oversized_state_is_rejected(void); +extern void test_TextTest_callback_state_is_owned_after_dispatch(void); +extern void test_TextTest_oversized_command_is_ignored(void); #endif diff --git a/test/test_entities_basic/tests/test_binary_sensor.cpp b/test/test_entities_basic/tests/test_binary_sensor.cpp index e100846..42213e1 100644 --- a/test/test_entities_basic/tests/test_binary_sensor.cpp +++ b/test/test_entities_basic/tests/test_binary_sensor.cpp @@ -110,7 +110,6 @@ void test_BinarySensorTest_object_id_setter(void) { sensor, ( "{" - "\"obj_id\":\"testId\"," "\"uniq_id\":\"uniqueSensor\"," "\"dev\":{\"ids\":\"testDevice\"}," "\"stat_t\":\"testData/testDevice/uniqueSensor/stat_t\"" diff --git a/test/test_entities_basic/tests/test_button.cpp b/test/test_entities_basic/tests/test_button.cpp index 3ac6581..471f115 100644 --- a/test/test_entities_basic/tests/test_button.cpp +++ b/test/test_entities_basic/tests/test_button.cpp @@ -139,7 +139,6 @@ void test_ButtonTest_object_id_setter(void) { button, ( "{" - "\"obj_id\":\"testId\"," "\"uniq_id\":\"uniqueButton\"," "\"dev\":{\"ids\":\"testDevice\"}," "\"cmd_t\":\"testData/testDevice/uniqueButton/cmd_t\"" diff --git a/test/test_entities_basic/tests/test_sensor.cpp b/test/test_entities_basic/tests/test_sensor.cpp index 068b87d..39f3d1c 100644 --- a/test/test_entities_basic/tests/test_sensor.cpp +++ b/test/test_entities_basic/tests/test_sensor.cpp @@ -100,7 +100,6 @@ void test_SensorTest_object_id_setter(void) { sensor, ( "{" - "\"obj_id\":\"testId\"," "\"uniq_id\":\"uniqueSensor\"," "\"dev\":{\"ids\":\"testDevice\"}," "\"stat_t\":\"testData/testDevice/uniqueSensor/stat_t\"" diff --git a/test/test_entities_basic/tests/test_switch.cpp b/test/test_entities_basic/tests/test_switch.cpp index 0c9d185..32ec104 100644 --- a/test/test_entities_basic/tests/test_switch.cpp +++ b/test/test_entities_basic/tests/test_switch.cpp @@ -110,7 +110,7 @@ void test_SwitchTest_device_discovery_payload(void) { ( "{" "\"dev\":{\"ids\":\"testDevice\"}," - "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"2.1.0\"}," + "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"3.0.2\"}," "\"cmps\":{" "\"uniqueSwitch\":{" "\"p\":\"switch\"," @@ -206,7 +206,6 @@ void test_SwitchTest_object_id_setter(void) { testSwitch, ( "{" - "\"obj_id\":\"testId\"," "\"uniq_id\":\"uniqueSwitch\"," "\"dev\":{\"ids\":\"testDevice\"}," "\"stat_t\":\"testData/testDevice/uniqueSwitch/stat_t\"," diff --git a/test/test_entities_basic/tests/test_text.cpp b/test/test_entities_basic/tests/test_text.cpp index 013bca6..8918800 100644 --- a/test/test_entities_basic/tests/test_text.cpp +++ b/test/test_entities_basic/tests/test_text.cpp @@ -51,6 +51,14 @@ void onCommandDeferredPublish(const char* value, HAText* caller) TEST_ASSERT_TRUE(caller->setState(value)); } +void onCommandSetStateAndOverwrite(const char* value, HAText* caller) +{ + TEST_ASSERT_TRUE(caller->setState(value)); + char* mutableValue = const_cast(value); + mutableValue[0] = 'x'; + TEST_ASSERT_EQUAL_STRING("hello", caller->getCurrentState()); +} + void test_TextTest_invalid_unique_id(void) { prepareTest @@ -178,7 +186,6 @@ void test_TextTest_object_id_setter(void) { text, ( "{" - "\"obj_id\":\"testId\"," "\"uniq_id\":\"uniqueText\"," "\"dev\":{\"ids\":\"testDevice\"}," "\"stat_t\":\"testData/testDevice/uniqueText/stat_t\"," @@ -318,6 +325,32 @@ void test_TextTest_publish_state_debounce(void) { TEST_ASSERT_EQUAL(0, mock->getFlushedMessagesNb()); } +void test_TextTest_current_state_is_owned(void) { + prepareTest + + HAText text(testUniqueId); + char state[] = "initial"; + text.setCurrentState(state); + state[0] = 'x'; + + TEST_ASSERT_EQUAL_STRING("initial", text.getCurrentState()); +} + +void test_TextTest_oversized_state_is_rejected(void) { + prepareTest + + HAText text(testUniqueId); + text.setCurrentState("initial"); + char state[HAText::MaxCommandLength + 2]; + memset(state, 'x', sizeof(state) - 1); + state[sizeof(state) - 1] = 0; + + text.setCurrentState(state); + TEST_ASSERT_EQUAL_STRING("initial", text.getCurrentState()); + TEST_ASSERT_FALSE(text.setState(state)); + TEST_ASSERT_EQUAL_STRING("initial", text.getCurrentState()); +} + void test_TextTest_command_callback(void) { prepareTest @@ -343,6 +376,32 @@ void test_TextTest_callback_publish_is_deferred_until_after_dispatch(void) { AHA_ASSERT_MQTT_MESSAGE(mock, 0, AHATOFSTR(StateTopic), "hello", true); } +void test_TextTest_callback_state_is_owned_after_dispatch(void) { + prepareTest + + mock->connectDummy(); + HAText text(testUniqueId); + text.onCommand(onCommandSetStateAndOverwrite); + + mock->fakeMessage(AHATOFSTR(CommandTopic), F("hello")); + + TEST_ASSERT_EQUAL_STRING("hello", text.getCurrentState()); +} + +void test_TextTest_oversized_command_is_ignored(void) { + prepareTest + + HAText text(testUniqueId); + text.onCommand(onCommandReceived); + char command[HAText::MaxCommandLength + 2]; + memset(command, 'x', sizeof(command) - 1); + command[sizeof(command) - 1] = 0; + + mock->fakeMessage(AHATOFSTR(CommandTopic), command); + + assertCommandCallbackNotCalled() +} + void test_TextTest_different_text_command(void) { prepareTest diff --git a/test/test_entities_extended/tests/test_camera.cpp b/test/test_entities_extended/tests/test_camera.cpp index ea0d7dc..275a620 100644 --- a/test/test_entities_extended/tests/test_camera.cpp +++ b/test/test_entities_extended/tests/test_camera.cpp @@ -66,7 +66,7 @@ void test_CameraTest_device_discovery_payload(void) { ( "{" "\"dev\":{\"ids\":\"testDevice\"}," - "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"2.1.0\"}," + "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"3.0.2\"}," "\"cmps\":{" "\"uniqueCamera\":{" "\"p\":\"camera\"," @@ -127,7 +127,6 @@ void test_CameraTest_object_id_setter(void) { camera, ( "{" - "\"obj_id\":\"testId\"," "\"uniq_id\":\"uniqueCamera\"," "\"dev\":{\"ids\":\"testDevice\"}," "\"t\":\"testData/testDevice/uniqueCamera/t\"" diff --git a/test/test_entities_extended/tests/test_cover.cpp b/test/test_entities_extended/tests/test_cover.cpp index 52752db..73c3aaa 100644 --- a/test/test_entities_extended/tests/test_cover.cpp +++ b/test/test_entities_extended/tests/test_cover.cpp @@ -104,7 +104,7 @@ void test_CoverTest_device_discovery_payload(void) { ( "{" "\"dev\":{\"ids\":\"testDevice\"}," - "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"2.1.0\"}," + "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"3.0.2\"}," "\"cmps\":{" "\"uniqueCover\":{" "\"p\":\"cover\"," @@ -234,7 +234,6 @@ void test_CoverTest_object_id_setter(void) { cover, ( "{" - "\"obj_id\":\"testId\"," "\"uniq_id\":\"uniqueCover\"," "\"dev\":{\"ids\":\"testDevice\"}," "\"stat_t\":\"testData/testDevice/uniqueCover/stat_t\"," diff --git a/test/test_entities_extended/tests/test_fan.cpp b/test/test_entities_extended/tests/test_fan.cpp index 6ed8e12..1f3bd88 100644 --- a/test/test_entities_extended/tests/test_fan.cpp +++ b/test/test_entities_extended/tests/test_fan.cpp @@ -136,7 +136,7 @@ void test_FanTest_device_discovery_payload(void) { ( "{" "\"dev\":{\"ids\":\"testDevice\"}," - "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"2.1.0\"}," + "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"3.0.2\"}," "\"cmps\":{" "\"uniqueFan\":{" "\"p\":\"fan\"," @@ -279,7 +279,6 @@ void test_FanTest_object_id_setter(void) { fan, ( "{" - "\"obj_id\":\"testId\"," "\"uniq_id\":\"uniqueFan\"," "\"dev\":{\"ids\":\"testDevice\"}," "\"stat_t\":\"testData/testDevice/uniqueFan/stat_t\"," diff --git a/test/test_entities_extended/tests/test_hvac.cpp b/test/test_entities_extended/tests/test_hvac.cpp index 1bde807..055fbd4 100644 --- a/test/test_entities_extended/tests/test_hvac.cpp +++ b/test/test_entities_extended/tests/test_hvac.cpp @@ -257,7 +257,7 @@ void test_HVACTest_device_discovery_payload(void) { ( "{" "\"dev\":{\"ids\":\"testDevice\"}," - "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"2.1.0\"}," + "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"3.0.2\"}," "\"cmps\":{" "\"uniqueHVAC\":{" "\"p\":\"climate\"," @@ -579,7 +579,6 @@ void test_HVACTest_object_id_setter(void) { hvac, ( "{" - "\"obj_id\":\"testId\"," "\"uniq_id\":\"uniqueHVAC\"," "\"curr_temp_t\":\"testData/testDevice/uniqueHVAC/curr_temp_t\"," "\"dev\":{\"ids\":\"testDevice\"}" diff --git a/test/test_entities_extended/tests/test_light.cpp b/test/test_entities_extended/tests/test_light.cpp index ec13dca..59920fb 100644 --- a/test/test_entities_extended/tests/test_light.cpp +++ b/test/test_entities_extended/tests/test_light.cpp @@ -429,7 +429,6 @@ void test_LightTest_object_id_setter(void) { light, ( "{" - "\"obj_id\":\"testId\"," "\"uniq_id\":\"uniqueLight\"," "\"dev\":{\"ids\":\"testDevice\"}," "\"stat_t\":\"testData/testDevice/uniqueLight/stat_t\"," diff --git a/test/test_entities_extended/tests/test_lock.cpp b/test/test_entities_extended/tests/test_lock.cpp index b3380ee..a13f671 100644 --- a/test/test_entities_extended/tests/test_lock.cpp +++ b/test/test_entities_extended/tests/test_lock.cpp @@ -103,7 +103,7 @@ void test_LockTest_device_discovery_payload(void) { ( "{" "\"dev\":{\"ids\":\"testDevice\"}," - "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"2.1.0\"}," + "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"3.0.2\"}," "\"cmps\":{" "\"uniqueLock\":{" "\"p\":\"lock\"," @@ -198,7 +198,6 @@ void test_LockTest_object_id_setter(void) { lock, ( "{" - "\"obj_id\":\"testId\"," "\"uniq_id\":\"uniqueLock\"," "\"dev\":{\"ids\":\"testDevice\"}," "\"stat_t\":\"testData/testDevice/uniqueLock/stat_t\"," diff --git a/test/test_entities_misc/tests/test_device_tracker.cpp b/test/test_entities_misc/tests/test_device_tracker.cpp index 483fd7e..a419f22 100644 --- a/test/test_entities_misc/tests/test_device_tracker.cpp +++ b/test/test_entities_misc/tests/test_device_tracker.cpp @@ -68,7 +68,7 @@ void test_DeviceTrackerTest_device_discovery_payload(void) { ( "{" "\"dev\":{\"ids\":\"testDevice\"}," - "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"2.1.0\"}," + "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"3.0.2\"}," "\"cmps\":{" "\"uniqueTracker\":{" "\"p\":\"device_tracker\"," @@ -215,7 +215,6 @@ void test_DeviceTrackerTest_object_id_setter(void) { tracker, ( "{" - "\"obj_id\":\"testId\"," "\"uniq_id\":\"uniqueTracker\"," "\"dev\":{\"ids\":\"testDevice\"}," "\"stat_t\":\"testData/testDevice/uniqueTracker/stat_t\"" diff --git a/test/test_entities_misc/tests/test_device_trigger.cpp b/test/test_entities_misc/tests/test_device_trigger.cpp index 52fb211..2c50f98 100644 --- a/test/test_entities_misc/tests/test_device_trigger.cpp +++ b/test/test_entities_misc/tests/test_device_trigger.cpp @@ -99,7 +99,7 @@ void test_DeviceTriggerTest_device_discovery_payload(void) { ( "{" "\"dev\":{\"ids\":\"testDevice\"}," - "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"2.1.0\"}," + "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"3.0.2\"}," "\"cmps\":{" "\"myType_mySubtype\":{" "\"p\":\"device_automation\"," diff --git a/test/test_entities_misc/tests/test_scene.cpp b/test/test_entities_misc/tests/test_scene.cpp index 23c31e0..83c16e2 100644 --- a/test/test_entities_misc/tests/test_scene.cpp +++ b/test/test_entities_misc/tests/test_scene.cpp @@ -95,7 +95,7 @@ void test_SceneTest_device_discovery_payload(void) { ( "{" "\"dev\":{\"ids\":\"testDevice\"}," - "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"2.1.0\"}," + "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"3.0.2\"}," "\"cmps\":{" "\"uniqueScene\":{" "\"p\":\"scene\"," @@ -167,7 +167,6 @@ void test_SceneTest_object_id_setter(void) { scene, ( "{" - "\"obj_id\":\"testId\"," "\"uniq_id\":\"uniqueScene\"," "\"pl_on\":\"ON\"," "\"cmd_t\":\"testData/testDevice/uniqueScene/cmd_t\"" diff --git a/test/test_entities_misc/tests/test_tag_scanner.cpp b/test/test_entities_misc/tests/test_tag_scanner.cpp index 5183450..98c9191 100644 --- a/test/test_entities_misc/tests/test_tag_scanner.cpp +++ b/test/test_entities_misc/tests/test_tag_scanner.cpp @@ -46,7 +46,7 @@ void test_TagScannerTest_device_discovery_payload(void) { ( "{" "\"dev\":{\"ids\":\"testDevice\"}," - "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"2.1.0\"}," + "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"3.0.2\"}," "\"cmps\":{" "\"uniqueScanner\":{" "\"p\":\"tag\"," diff --git a/test/test_entities_numeric/tests/test_number.cpp b/test/test_entities_numeric/tests/test_number.cpp index 947794d..bf8f77d 100644 --- a/test/test_entities_numeric/tests/test_number.cpp +++ b/test/test_entities_numeric/tests/test_number.cpp @@ -351,7 +351,6 @@ void test_NumberTest_object_id_setter(void) { number, ( "{" - "\"obj_id\":\"testId\"," "\"uniq_id\":\"uniqueNumber\"," "\"dev\":{\"ids\":\"testDevice\"}," "\"stat_t\":\"testData/testDevice/uniqueNumber/stat_t\"," @@ -926,7 +925,7 @@ void test_NumberTest_update_min_max_step_republishes_device_discovery_when_enabl ( "{" "\"dev\":{\"ids\":\"testDevice\"}," - "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"2.1.0\"}," + "\"o\":{\"name\":\"ArduinoHA\",\"sw\":\"3.0.2\"}," "\"cmps\":{" "\"uniqueNumber\":{" "\"p\":\"number\"," diff --git a/test/test_entities_numeric/tests/test_select.cpp b/test/test_entities_numeric/tests/test_select.cpp index 0bdb350..c573e90 100644 --- a/test/test_entities_numeric/tests/test_select.cpp +++ b/test/test_entities_numeric/tests/test_select.cpp @@ -280,7 +280,6 @@ void test_SelectTest_object_id_setter(void) { select, ( "{" - "\"obj_id\":\"testId\"," "\"uniq_id\":\"uniqueSelect\"," "\"options\":[\"Option A\",\"B\",\"C\"]," "\"dev\":{\"ids\":\"testDevice\"}," diff --git a/test/test_native_core/test_main.cpp b/test/test_native_core/test_main.cpp new file mode 100644 index 0000000..acb455b --- /dev/null +++ b/test/test_native_core/test_main.cpp @@ -0,0 +1,355 @@ +#include +#include +#include + +#include "HADevice.h" +#include "HAMqtt.h" +#include "device-types/HABaseDeviceType.h" +#include "device-types/HAText.h" +#include "mocks/PubSubClientMock.h" +#include "utils/HAAvailabilityConfig.h" +#include "utils/HAJson.h" +#include "utils/HASerializer.h" +#include "utils/HASerializerArray.h" + +namespace { +class NativeSerializerEntity : public HABaseDeviceType +{ +public: + explicit NativeSerializerEntity(const char* uniqueId) : + HABaseDeviceType(F("sensor"), uniqueId) + { + + } + +protected: + void onMqttConnected() override { } +}; +class NativeDiscoveryEntity : public HABaseDeviceType +{ +public: + explicit NativeDiscoveryEntity(const char* uniqueId) : + HABaseDeviceType(F("sensor"), uniqueId) + { + + } + +protected: + void buildSerializer() override + { + if (_serializer) { + return; + } + + _serializer = new HASerializer(this, 2); + _serializer->set(F("name"), uniqueId()); + _serializer->set(HASerializer::WithUniqueId); + } + + HASerializer* buildDeviceDiscoverySerializer() override + { + HASerializer* serializer = new HASerializer(this, 3); + serializer->set( + F("p"), + componentName(), + HASerializer::ProgmemPropertyValue + ); + serializer->set(F("name"), uniqueId()); + serializer->set(HASerializer::WithUniqueId); + return serializer; + } + + bool supportsDeviceDiscovery() const override + { + return true; + } + + void onMqttConnected() override { } +}; + +uint8_t disconnectedCallbackCalls = 0; +uint8_t stateCallbackCalls = 0; + +void onNativeDisconnected() +{ + disconnectedCallbackCalls++; +} + +void onNativeStateChanged(HAMqtt::ConnectionState) +{ + stateCallbackCalls++; +} + + +void test_json_helpers_escape_control_bytes_and_preserve_cursor_contract() +{ + const char value[] = "quote\" slash\\ newline\n tab\t control\x01"; + char output[96] = {}; + char* cursor = output; + + TEST_ASSERT_EQUAL_UINT16( + strlen("\"quote\\\" slash\\\\ newline\\n tab\\t control\\u0001\""), + HAJson::calculateEscapedStringSize(value) + ); + TEST_ASSERT_TRUE(HAJson::appendEscapedString(cursor, output + sizeof(output) - 1, value)); + TEST_ASSERT_EQUAL_STRING( + "\"quote\\\" slash\\\\ newline\\n tab\\t control\\u0001\"", + output + ); + TEST_ASSERT_EQUAL_PTR(output + strlen(output), cursor); +} + +void test_availability_and_serializer_array_escape_json_values() +{ + HAAvailabilityConfig availability; + TEST_ASSERT_TRUE(availability.add("availability/\"main\"", "{{ value_json.\\n }}", "on\nline", "off\\line")); + + const uint16_t availabilitySize = availability.calculateJsonSize(); + char availabilityJson[availabilitySize + 1]; + TEST_ASSERT_TRUE(availability.serialize(availabilityJson)); + TEST_ASSERT_EQUAL_STRING( + "[{\"t\":\"availability/\\\"main\\\"\",\"val_tpl\":\"{{ value_json.\\\\n }}\",\"pl_avail\":\"on\\nline\",\"pl_not_avail\":\"off\\\\line\"}]", + availabilityJson + ); + TEST_ASSERT_EQUAL_UINT16(strlen(availabilityJson), availabilitySize); + + HASerializerArray values(2, false); + TEST_ASSERT_TRUE(values.add("first\"item")); + TEST_ASSERT_TRUE(values.add("second\nitem")); + const uint16_t valuesSize = values.calculateSize(); + char valuesJson[valuesSize + 1]; + TEST_ASSERT_TRUE(values.serialize(valuesJson)); + TEST_ASSERT_EQUAL_STRING("[\"first\\\"item\",\"second\\nitem\"]", valuesJson); + TEST_ASSERT_EQUAL_UINT16(strlen(valuesJson), valuesSize); +} + +void test_discovery_topic_tokens_are_rejected_before_topic_generation() +{ + HADevice device("native_device"); + PubSubClientMock* mock = new PubSubClientMock(); + HAMqtt mqtt(mock, device); + + const __FlashStringHelper* component = F("sensor"); + const uint16_t topicSize = HASerializer::calculateConfigTopicLength(component, "valid_entity-1"); + char topic[topicSize]; + TEST_ASSERT_TRUE(HASerializer::generateConfigTopic(topic, component, "valid_entity-1")); + TEST_ASSERT_EQUAL_STRING("homeassistant/sensor/native_device/valid_entity-1/config", topic); + + TEST_ASSERT_EQUAL_UINT16(0, HASerializer::calculateConfigTopicLength(component, "invalid/entity")); + TEST_ASSERT_FALSE(HASerializer::generateConfigTopic(topic, component, "invalid entity")); +} + +void test_streaming_serializer_writes_exact_escaped_payload_to_mqtt_mock() +{ + HADevice device("native_device"); + PubSubClientMock* mock = new PubSubClientMock(); + HAMqtt mqtt(mock, device); + NativeSerializerEntity entity("native_entity"); + + TEST_ASSERT_TRUE(mqtt.begin("native-host")); + TEST_ASSERT_TRUE(mock->connectDummy()); + + HASerializer serializer(&entity, 2); + serializer.set(F("name"), "Name \"quoted\"\\line\nnext"); + serializer.set(F("unit_of_meas"), "C\tunit"); + + const uint16_t payloadSize = serializer.calculateSize(); + TEST_ASSERT_NOT_EQUAL(0, payloadSize); + TEST_ASSERT_TRUE(mqtt.beginPublish("native/serializer", payloadSize, true)); + TEST_ASSERT_TRUE(serializer.flush()); + TEST_ASSERT_TRUE(mqtt.endPublish()); + + TEST_ASSERT_EQUAL_UINT8(1, mock->getFlushedMessagesNb()); + MqttMessage* message = mock->getFlushedMessages()[0]; + TEST_ASSERT_EQUAL_STRING("native/serializer", message->topic); + TEST_ASSERT_EQUAL_STRING( + "{\"name\":\"Name \\\"quoted\\\"\\\\line\\nnext\",\"unit_of_meas\":\"C\\tunit\"}", + message->buffer + ); + TEST_ASSERT_EQUAL_UINT16(strlen(message->buffer), message->writtenSize); + TEST_ASSERT_EQUAL_UINT16(payloadSize, message->writtenSize); +} +void test_staged_migration_orders_markers_device_payload_and_legacy_cleanup() +{ + HADevice device("native_device"); + PubSubClientMock* mock = new PubSubClientMock(); + HAMqtt mqtt(mock, device); + NativeDiscoveryEntity first("first"); + NativeDiscoveryEntity second("second"); + TEST_ASSERT_TRUE(mock->connectDummy()); + + TEST_ASSERT_TRUE(mqtt.beginDeviceDiscoveryMigration()); + TEST_ASSERT_EQUAL_INT(HAMqtt::DeviceDiscoveryMigrationMarkersPending, + mqtt.getDeviceDiscoveryMigrationState()); + TEST_ASSERT_FALSE(mqtt.publishDeviceDiscovery()); + TEST_ASSERT_EQUAL_UINT8(0, mock->getFlushedMessagesNb()); + + TEST_ASSERT_TRUE(mqtt.publishDeviceDiscoveryMigrationMarkers()); + TEST_ASSERT_EQUAL_UINT8(2, mock->getFlushedMessagesNb()); + TEST_ASSERT_EQUAL_STRING("homeassistant/sensor/native_device/first/config", + mock->getFlushedMessages()[0]->topic); + TEST_ASSERT_EQUAL_STRING("{\"migrate_discovery\":true}", + mock->getFlushedMessages()[0]->buffer); + TEST_ASSERT_EQUAL_STRING("homeassistant/sensor/native_device/second/config", + mock->getFlushedMessages()[1]->topic); + + TEST_ASSERT_TRUE(mqtt.publishDeviceDiscoveryMigrationConfig()); + TEST_ASSERT_EQUAL_UINT8(3, mock->getFlushedMessagesNb()); + TEST_ASSERT_EQUAL_STRING("homeassistant/device/native_device/config", + mock->getFlushedMessages()[2]->topic); + TEST_ASSERT_NOT_NULL(strstr(mock->getFlushedMessages()[2]->buffer, "\"cmps\"")); + TEST_ASSERT_NOT_NULL(strstr(mock->getFlushedMessages()[2]->buffer, "\"first\"")); + TEST_ASSERT_NOT_NULL(strstr(mock->getFlushedMessages()[2]->buffer, "\"second\"")); + + TEST_ASSERT_TRUE(mqtt.completeDeviceDiscoveryMigration()); + TEST_ASSERT_EQUAL_UINT8(5, mock->getFlushedMessagesNb()); + TEST_ASSERT_EQUAL_STRING("homeassistant/sensor/native_device/first/config", + mock->getFlushedMessages()[3]->topic); + TEST_ASSERT_EQUAL_STRING("", mock->getFlushedMessages()[3]->buffer); + TEST_ASSERT_EQUAL_STRING("homeassistant/sensor/native_device/second/config", + mock->getFlushedMessages()[4]->topic); + TEST_ASSERT_EQUAL_STRING("", mock->getFlushedMessages()[4]->buffer); + TEST_ASSERT_EQUAL_INT(HAMqtt::DeviceDiscoveryMigrationCompleted, + mqtt.getDeviceDiscoveryMigrationState()); + + TEST_ASSERT_TRUE(mqtt.rollbackDeviceDiscoveryMigration()); + TEST_ASSERT_EQUAL_UINT8(9, mock->getFlushedMessagesNb()); + TEST_ASSERT_EQUAL_STRING("homeassistant/device/native_device/config", + mock->getFlushedMessages()[5]->topic); + TEST_ASSERT_EQUAL_STRING("{\"migrate_discovery\":true}", + mock->getFlushedMessages()[5]->buffer); + TEST_ASSERT_EQUAL_STRING("homeassistant/sensor/native_device/first/config", + mock->getFlushedMessages()[6]->topic); + TEST_ASSERT_EQUAL_STRING("homeassistant/sensor/native_device/second/config", + mock->getFlushedMessages()[7]->topic); + TEST_ASSERT_EQUAL_STRING("homeassistant/device/native_device/config", + mock->getFlushedMessages()[8]->topic); + TEST_ASSERT_EQUAL_STRING("", mock->getFlushedMessages()[8]->buffer); + TEST_ASSERT_EQUAL_INT(HAMqtt::DeviceDiscoveryMigrationIdle, + mqtt.getDeviceDiscoveryMigrationState()); + TEST_ASSERT_FALSE(mqtt.isDeviceDiscoveryEnabled()); +} + +void test_failed_rollback_stays_staged_until_it_can_complete() +{ + HADevice device("native_device"); + PubSubClientMock* mock = new PubSubClientMock(); + HAMqtt mqtt(mock, device); + NativeDiscoveryEntity entity("entity"); + TEST_ASSERT_TRUE(mock->connectDummy()); + + TEST_ASSERT_TRUE(mqtt.beginDeviceDiscoveryMigration()); + TEST_ASSERT_TRUE(mqtt.publishDeviceDiscoveryMigrationMarkers()); + TEST_ASSERT_TRUE(mqtt.publishDeviceDiscoveryMigrationConfig()); + + mock->failNextBeginPublish(); + TEST_ASSERT_FALSE(mqtt.rollbackDeviceDiscoveryMigration()); + TEST_ASSERT_EQUAL_INT(HAMqtt::DeviceDiscoveryMigrationRollbackPending, + mqtt.getDeviceDiscoveryMigrationState()); + TEST_ASSERT_TRUE(mqtt.isDeviceDiscoveryMigrationInProgress()); + TEST_ASSERT_FALSE(mqtt.publishDeviceDiscovery()); + + TEST_ASSERT_TRUE(mqtt.rollbackDeviceDiscoveryMigration()); + TEST_ASSERT_EQUAL_INT(HAMqtt::DeviceDiscoveryMigrationIdle, + mqtt.getDeviceDiscoveryMigrationState()); + TEST_ASSERT_FALSE(mqtt.isDeviceDiscoveryEnabled()); +} + +void test_device_component_removal_uses_marker_then_omission_and_readd() +{ + HADevice device("native_device"); + PubSubClientMock* mock = new PubSubClientMock(); + HAMqtt mqtt(mock, device); + NativeDiscoveryEntity first("first"); + NativeDiscoveryEntity second("second"); + TEST_ASSERT_TRUE(mock->connectDummy()); + mqtt.enableDeviceDiscovery(); + + TEST_ASSERT_TRUE(first.removeFromDiscovery()); + TEST_ASSERT_TRUE(first.isRemovedFromDeviceDiscovery()); + TEST_ASSERT_EQUAL_UINT8(2, mock->getFlushedMessagesNb()); + TEST_ASSERT_NOT_NULL(strstr(mock->getFlushedMessages()[0]->buffer, + "\"first\":{\"p\":\"sensor\"}")); + TEST_ASSERT_NULL(strstr(mock->getFlushedMessages()[1]->buffer, "\"first\"")); + TEST_ASSERT_NOT_NULL(strstr(mock->getFlushedMessages()[1]->buffer, "\"second\"")); + + TEST_ASSERT_TRUE(first.republishDiscovery()); + TEST_ASSERT_FALSE(first.isRemovedFromDeviceDiscovery()); + TEST_ASSERT_EQUAL_UINT8(3, mock->getFlushedMessagesNb()); + TEST_ASSERT_NOT_NULL(strstr(mock->getFlushedMessages()[2]->buffer, "\"first\"")); +} + +void test_entities_before_or_after_mqtt_have_safe_registration_lifetimes() +{ + NativeDiscoveryEntity before("before"); + HADevice device("native_device"); + PubSubClientMock* mock = new PubSubClientMock(); + HAMqtt mqtt(mock, device, 2); + TEST_ASSERT_EQUAL_UINT8(1, mqtt.getRegisteredDeviceTypeCount()); + + NativeDiscoveryEntity* after = new NativeDiscoveryEntity("after"); + TEST_ASSERT_EQUAL_UINT8(2, mqtt.getRegisteredDeviceTypeCount()); + delete after; + TEST_ASSERT_EQUAL_UINT8(1, mqtt.getRegisteredDeviceTypeCount()); +} + +void test_text_state_is_bounded_and_retains_the_last_valid_value() +{ + HADevice device("native_device"); + PubSubClientMock* mock = new PubSubClientMock(); + HAMqtt mqtt(mock, device); + HAText text("text"); + char oversizedState[HAText::MaxCommandLength + 2]; + memset(oversizedState, 'x', sizeof(oversizedState) - 1); + oversizedState[sizeof(oversizedState) - 1] = 0; + + text.setCurrentState("initial"); + text.setCurrentState(oversizedState); + TEST_ASSERT_EQUAL_STRING("initial", text.getCurrentState()); + TEST_ASSERT_FALSE(text.setState(oversizedState)); + TEST_ASSERT_EQUAL_STRING("initial", text.getCurrentState()); +} + +void test_registration_cap_and_explicit_disconnect_are_reported() +{ + HADevice device("native_device"); + PubSubClientMock* mock = new PubSubClientMock(); + HAMqtt mqtt(mock, device, 1); + NativeDiscoveryEntity first("first"); + NativeDiscoveryEntity second("second"); + TEST_ASSERT_EQUAL_UINT8(1, mqtt.getRegisteredDeviceTypeCount()); + TEST_ASSERT_EQUAL_UINT16(1, mqtt.getDeviceTypeRegistrationFailures()); + + disconnectedCallbackCalls = 0; + stateCallbackCalls = 0; + mqtt.onDisconnected(onNativeDisconnected); + mqtt.onStateChanged(onNativeStateChanged); + TEST_ASSERT_TRUE(mqtt.begin("native-host")); + TEST_ASSERT_TRUE(mock->connectDummy()); + mock->setState(HAMqtt::StateConnected); + mqtt.loop(); + TEST_ASSERT_TRUE(mqtt.disconnect()); + TEST_ASSERT_EQUAL_UINT8(1, disconnectedCallbackCalls); + TEST_ASSERT_TRUE(stateCallbackCalls >= 2); +} + +} // namespace + +void setUp(void) { } +void tearDown(void) { } + +int main(int, char**) +{ + UNITY_BEGIN(); + RUN_TEST(test_json_helpers_escape_control_bytes_and_preserve_cursor_contract); + RUN_TEST(test_availability_and_serializer_array_escape_json_values); + RUN_TEST(test_discovery_topic_tokens_are_rejected_before_topic_generation); + RUN_TEST(test_streaming_serializer_writes_exact_escaped_payload_to_mqtt_mock); + RUN_TEST(test_staged_migration_orders_markers_device_payload_and_legacy_cleanup); + RUN_TEST(test_failed_rollback_stays_staged_until_it_can_complete); + RUN_TEST(test_device_component_removal_uses_marker_then_omission_and_readd); + RUN_TEST(test_entities_before_or_after_mqtt_have_safe_registration_lifetimes); + RUN_TEST(test_text_state_is_bounded_and_retains_the_last_valid_value); + RUN_TEST(test_registration_cap_and_explicit_disconnect_are_reported); + return UNITY_END(); +} diff --git a/test/test_utils_json/test_main.cpp b/test/test_utils_json/test_main.cpp new file mode 100644 index 0000000..9f92b55 --- /dev/null +++ b/test/test_utils_json/test_main.cpp @@ -0,0 +1,86 @@ +#include +#include +#include + +#include "utils/HAJson.h" + +using TestFn = void (*)(void); + +struct TestCase { + const char* name; + TestFn fn; + uint16_t line; +}; + +#define TEST_ENTRY(fn) { #fn, fn, __LINE__ } + +void test_HAJson_escapes_all_required_json_characters() +{ + const char value[] = "quote\" slash\\ newline\n tab\t control\x01"; + char output[96] = {}; + char* cursor = output; + + TEST_ASSERT_EQUAL_UINT16( + strlen("\"quote\\\" slash\\\\ newline\\n tab\\t control\\u0001\""), + HAJson::calculateEscapedStringSize(value) + ); + TEST_ASSERT_TRUE(HAJson::appendEscapedString(cursor, output + sizeof(output) - 1, value)); + TEST_ASSERT_EQUAL_STRING( + "\"quote\\\" slash\\\\ newline\\n tab\\t control\\u0001\"", + output + ); +} + +void test_HAJson_rejects_too_small_output_buffer_without_partial_output() +{ + char output[5] = "ok"; + char* cursor = output; + + TEST_ASSERT_FALSE(HAJson::appendEscapedString(cursor, output + sizeof(output) - 1, "toolong")); + TEST_ASSERT_EQUAL_STRING("ok", output); + TEST_ASSERT_EQUAL_PTR(output, cursor); +} + +void test_HAJson_validates_home_assistant_discovery_topic_tokens() +{ + TEST_ASSERT_TRUE(HAJson::isValidDiscoveryTopicToken("device_01-A")); + TEST_ASSERT_FALSE(HAJson::isValidDiscoveryTopicToken("")); + TEST_ASSERT_FALSE(HAJson::isValidDiscoveryTopicToken("device/id")); + TEST_ASSERT_FALSE(HAJson::isValidDiscoveryTopicToken("device id")); + TEST_ASSERT_FALSE(HAJson::isValidDiscoveryTopicToken("device\n")); +} + +static TestCase tests[] = { + TEST_ENTRY(test_HAJson_escapes_all_required_json_characters), + TEST_ENTRY(test_HAJson_rejects_too_small_output_buffer_without_partial_output), + TEST_ENTRY(test_HAJson_validates_home_assistant_discovery_topic_tokens), +}; + +static const size_t testCount = sizeof(tests) / sizeof(tests[0]); +static size_t nextTest = 0; +static bool begun = false; + +void setUp(void) { } +void tearDown(void) { } + +void setup() +{ + Serial.begin(115200); + delay(500); + UNITY_BEGIN(); + begun = true; +} + +void loop() +{ + if (begun && nextTest < testCount) { + TestCase& test = tests[nextTest++]; + UnityDefaultTestRun(test.fn, test.name, test.line); + return; + } + + if (begun) { + UNITY_END(); + begun = false; + } +} diff --git a/tests/ha-contract/Dockerfile b/tests/ha-contract/Dockerfile new file mode 100644 index 0000000..52c99a1 --- /dev/null +++ b/tests/ha-contract/Dockerfile @@ -0,0 +1,8 @@ +FROM python:3.12-slim + +WORKDIR /tests +COPY requirements.txt . +RUN pip install --no-cache-dir -r requirements.txt +COPY test_contract.py . + +CMD ["python", "test_contract.py"] diff --git a/tests/ha-contract/README.md b/tests/ha-contract/README.md new file mode 100644 index 0000000..3906785 --- /dev/null +++ b/tests/ha-contract/README.md @@ -0,0 +1,45 @@ +# Home Assistant MQTT contract tests + +This harness runs retained MQTT discovery messages through Mosquitto and a real +Home Assistant container. It verifies the behavior that firmware unit tests +cannot: Home Assistant's entity-registry identity, user-owned registry changes, +and retained discovery after a Home Assistant restart. + +It has two ordered modes: + +1. `migration` publishes single-component discovery, applies a user rename and + disable, then sends `migrate_discovery` markers, a device bundle, and legacy + retained-topic cleanup. +2. `retained-restart` runs after Home Assistant restarts and verifies the + retained device bundle plus preserved registry customization. + +The migration case also rejects malformed retained JSON, checks that a direct +single-to-device publication does not duplicate registry identity, exercises the +device-component tombstone/omission sequence, and makes the current +`def_ent_id`/no-`obj_id` schema expectation explicit. + +## Run locally + +From the repository root: + +```bash +HA_VERSION=2024.11.3 docker compose -f tests/ha-contract/compose.yaml up -d mqtt homeassistant +HA_VERSION=2024.11.3 docker compose -f tests/ha-contract/compose.yaml run --rm tests +HA_VERSION=2024.11.3 docker compose -f tests/ha-contract/compose.yaml restart homeassistant +CONTRACT_MODE=retained-restart HA_VERSION=2024.11.3 docker compose -f tests/ha-contract/compose.yaml run --rm tests +HA_VERSION=2024.11.3 docker compose -f tests/ha-contract/compose.yaml down -v +``` + +Use `HA_VERSION=stable` and `HA_VERSION=dev` for the current supported and +development Home Assistant images. The test uses only ephemeral named volumes; +`down -v` removes its broker data, Home Assistant config, owner token, and +registry state. + +## Scope + +The firmware's native Unity suite covers JSON escaping, invalid topic tokens, +serializer preflight, migration ordering, component removal, and lifecycle +behavior. This container suite covers Home Assistant's persistence contract. It +does not need a physical board: fixture discovery documents mirror retained +payloads emitted by ArduinoHA and isolate Home Assistant/MQTT compatibility. + diff --git a/tests/ha-contract/compose.yaml b/tests/ha-contract/compose.yaml new file mode 100644 index 0000000..dd6a68e --- /dev/null +++ b/tests/ha-contract/compose.yaml @@ -0,0 +1,43 @@ +services: + mqtt: + image: eclipse-mosquitto:2 + command: ["mosquitto", "-c", "/mosquitto/config/mosquitto.conf"] + volumes: + - ./mosquitto.conf:/mosquitto/config/mosquitto.conf:ro + healthcheck: + test: ["CMD-SHELL", "mosquitto_sub -h localhost -t '$$SYS/broker/version' -C 1 -W 5 >/dev/null 2>&1"] + interval: 3s + timeout: 5s + retries: 20 + + homeassistant: + image: ghcr.io/home-assistant/home-assistant:${HA_VERSION:-stable} + environment: + TZ: UTC + volumes: + - ha_config:/config + - ./configuration.yaml:/config/configuration.yaml:ro + depends_on: + mqtt: + condition: service_healthy + + tests: + build: + context: . + environment: + HA_URL: http://homeassistant:8123 + MQTT_HOST: mqtt + MQTT_PORT: "1883" + CONTRACT_MODE: ${CONTRACT_MODE:-migration} + volumes: + - contract_state:/state + depends_on: + mqtt: + condition: service_healthy + homeassistant: + condition: service_started + +volumes: + ha_config: + contract_state: + diff --git a/tests/ha-contract/configuration.yaml b/tests/ha-contract/configuration.yaml new file mode 100644 index 0000000..9965618 --- /dev/null +++ b/tests/ha-contract/configuration.yaml @@ -0,0 +1,3 @@ +default_config: + + diff --git a/tests/ha-contract/mosquitto.conf b/tests/ha-contract/mosquitto.conf new file mode 100644 index 0000000..a6849fc --- /dev/null +++ b/tests/ha-contract/mosquitto.conf @@ -0,0 +1,6 @@ +persistence true +persistence_location /mosquitto/data/ +allow_anonymous true +listener 1883 +log_type all + diff --git a/tests/ha-contract/requirements.txt b/tests/ha-contract/requirements.txt new file mode 100644 index 0000000..40a499f --- /dev/null +++ b/tests/ha-contract/requirements.txt @@ -0,0 +1,4 @@ +paho-mqtt==2.1.0 +requests==2.32.3 +websocket-client==1.8.0 + diff --git a/tests/ha-contract/test_contract.py b/tests/ha-contract/test_contract.py new file mode 100644 index 0000000..2e81cc6 --- /dev/null +++ b/tests/ha-contract/test_contract.py @@ -0,0 +1,441 @@ +"""Home Assistant MQTT discovery contract checks. + +The test intentionally uses retained MQTT messages, like a deployed firmware +node. It inspects HA's entity registry over the authenticated WebSocket API so +the migration assertion is about HA's persistent identity, not just payload +shape. +""" + +import json +import os +import pathlib +import time +import uuid + +import paho.mqtt.client as mqtt +import requests +import websocket + + +HA_URL = os.environ.get("HA_URL", "http://homeassistant:8123").rstrip("/") +MQTT_HOST = os.environ.get("MQTT_HOST", "mqtt") +MQTT_PORT = int(os.environ.get("MQTT_PORT", "1883")) +MODE = os.environ.get("CONTRACT_MODE", "migration") +STATE = pathlib.Path("/state") +TOKEN_FILE = STATE / "ha-token" +DEVICE_ID = "contract_device" + + +class ContractFailure(RuntimeError): + pass + + +def fail(message): + raise ContractFailure(message) + + +def wait_until(description, predicate, timeout=90, interval=1): + deadline = time.monotonic() + timeout + last_error = None + while time.monotonic() < deadline: + try: + result = predicate() + if result: + return result + except Exception as error: # HA is expected to be starting initially. + last_error = error + time.sleep(interval) + suffix = f" (last error: {last_error})" if last_error else "" + fail(f"timed out waiting for {description}{suffix}") + + +def wait_for_home_assistant(): + def ready(): + response = requests.get(f"{HA_URL}/api/", timeout=5) + return response.status_code in (200, 401) + + wait_until("Home Assistant HTTP API", ready, timeout=180) + + +def response_json(response, context): + if not response.ok: + fail(f"{context} failed ({response.status_code}): {response.text}") + try: + return response.json() + except ValueError as error: + fail(f"{context} returned invalid JSON: {error}") + + +def onboarding_token(): + """Create an ephemeral owner token once and share it across restart checks.""" + if TOKEN_FILE.exists(): + return TOKEN_FILE.read_text(encoding="utf-8").strip() + + client_id = "http://contract-tests.local/" + user = { + "client_id": client_id, + "name": "Contract Test Owner", + "username": "contract-owner", + "password": "contract-test-password", + "language": "en", + } + response = requests.post(f"{HA_URL}/api/onboarding/users", json=user, timeout=10) + created = response_json(response, "Home Assistant onboarding user creation") + auth_code = created.get("auth_code") + if not auth_code: + fail("Home Assistant onboarding did not return an auth_code") + + token_response = requests.post( + f"{HA_URL}/auth/token", + data={ + "client_id": client_id, + "grant_type": "authorization_code", + "code": auth_code, + }, + timeout=10, + ) + token = response_json(token_response, "Home Assistant token exchange").get("access_token") + if not token: + fail("Home Assistant token exchange did not return an access_token") + + headers = {"Authorization": f"Bearer {token}"} + # These operations are idempotent across HA versions that still expose them. + for path, payload in ( + ("/api/onboarding/core_config", {}), + ("/api/onboarding/analytics", {"preferences": {}}), + ): + response = requests.post(f"{HA_URL}{path}", headers=headers, json=payload, timeout=10) + if response.status_code not in (200, 201, 400, 404): + fail(f"Home Assistant onboarding step {path} failed: {response.status_code} {response.text}") + + STATE.mkdir(parents=True, exist_ok=True) + TOKEN_FILE.write_text(token, encoding="utf-8") + return token + + +class HAWebSocket: + def __init__(self, token): + scheme = "wss" if HA_URL.startswith("https://") else "ws" + address = HA_URL.split("://", 1)[1] + self._socket = websocket.create_connection(f"{scheme}://{address}/api/websocket", timeout=15) + required = json.loads(self._socket.recv()) + if required.get("type") != "auth_required": + fail(f"unexpected Home Assistant WebSocket greeting: {required}") + self._socket.send(json.dumps({"type": "auth", "access_token": token})) + authenticated = json.loads(self._socket.recv()) + if authenticated.get("type") != "auth_ok": + fail(f"Home Assistant WebSocket authentication failed: {authenticated}") + self._next_id = 1 + + def close(self): + self._socket.close() + + def call(self, message_type, **kwargs): + message_id = self._next_id + self._next_id += 1 + self._socket.send(json.dumps({"id": message_id, "type": message_type, **kwargs})) + while True: + result = json.loads(self._socket.recv()) + if result.get("id") != message_id: + continue + if not result.get("success"): + fail(f"WebSocket {message_type} failed: {result}") + return result.get("result") + + def registry_entries(self): + return self.call("config/entity_registry/list") + + +def configure_mqtt_integration(token): + """Create HA's MQTT config entry; broker YAML options are no longer accepted.""" + headers = {"Authorization": f"Bearer {token}"} + response = requests.post( + f"{HA_URL}/api/config/config_entries/flow", + headers=headers, + json={"handler": "mqtt"}, + timeout=15, + ) + flow = response_json(response, "MQTT config-entry flow creation") + if flow.get("type") == "create_entry": + return + if flow.get("type") != "form" or not flow.get("flow_id"): + fail(f"unexpected MQTT config-entry flow result: {flow}") + + user_input = {"broker": MQTT_HOST, "port": MQTT_PORT} + # Newer HA versions group TLS/transport values in a required section. + # Older supported versions do not expose the section and must not receive + # unknown fields, so negotiate from the returned data schema. + schema_names = { + field.get("name") + for field in flow.get("data_schema", []) + if isinstance(field, dict) + } + if "other_settings" in schema_names: + user_input["other_settings"] = { + "set_client_cert": False, + "set_ca_cert": "off", + "transport": "tcp", + } + + response = requests.post( + f"{HA_URL}/api/config/config_entries/flow/{flow['flow_id']}", + headers=headers, + json=user_input, + timeout=15, + ) + configured = response_json(response, "MQTT config-entry flow configuration") + if configured.get("type") != "create_entry": + fail(f"MQTT config-entry flow did not create an entry: {configured}") + + # The successful flow result is returned before the integration has had a + # chance to subscribe to retained discovery topics. + time.sleep(3) + + +class RetainedPublisher: + def __init__(self): + self._client = mqtt.Client(mqtt.CallbackAPIVersion.VERSION2, client_id=f"contract-{uuid.uuid4()}") + self._client.connect(MQTT_HOST, MQTT_PORT, keepalive=15) + self._client.loop_start() + + def close(self): + self._client.loop_stop() + self._client.disconnect() + + def publish(self, topic, payload): + info = self._client.publish(topic, payload, qos=1, retain=True) + info.wait_for_publish(timeout=10) + if not info.is_published(): + fail(f"MQTT publish timed out for {topic}") + + def retained_payload(self, topic): + received = [] + + def on_message(_client, _userdata, message): + if message.topic == topic: + received.append(message.payload.decode("utf-8")) + + self._client.on_message = on_message + self._client.subscribe(topic, qos=1) + value = wait_until(f"retained MQTT message {topic}", lambda: received[0] if received else None) + self._client.unsubscribe(topic) + self._client.on_message = None + return value + + +def legacy_topic(object_id): + return f"homeassistant/sensor/{DEVICE_ID}/{object_id}/config" + + +def device_topic(): + return f"homeassistant/device/{DEVICE_ID}/config" + + +def component(object_id, unique_id=None, **extra): + payload = { + "p": "sensor", + "name": object_id.replace("_", " ").title(), + "uniq_id": unique_id or f"{DEVICE_ID}_{object_id}", + "stat_t": f"contract/{object_id}/state", + } + payload.update(extra) + return payload + + +def device_payload(components): + return { + "dev": {"ids": [DEVICE_ID], "name": "ArduinoHA contract device"}, + "o": {"name": "ArduinoHA", "sw": "3.0.2"}, + "cmps": components, + } + + +def legacy_payload(object_id, unique_id=None, **extra): + payload = component(object_id, unique_id, **extra) + payload["dev"] = {"ids": [DEVICE_ID], "name": "ArduinoHA contract device"} + payload.pop("p") + return payload + + +def find_unique(entries, unique_id): + matches = [entry for entry in entries if entry.get("unique_id") == unique_id] + if len(matches) > 1: + fail(f"duplicate entity-registry entries for {unique_id}: {matches}") + return matches[0] if matches else None + + +def wait_for_entry(ws, unique_id): + return wait_until( + f"entity registry entry {unique_id}", + lambda: find_unique(ws.registry_entries(), unique_id), + timeout=60, + ) + + +def migration_contract(): + publisher = RetainedPublisher() + token = onboarding_token() + configure_mqtt_integration(token) + ws = HAWebSocket(token) + try: + # Existing single-component entity and a user-owned registry customization. + unique = f"{DEVICE_ID}_temperature" + publisher.publish(legacy_topic("temperature"), json.dumps(legacy_payload("temperature", unique))) + publisher.publish("contract/temperature/state", "21.5") + original = wait_for_entry(ws, unique) + original_id = original["id"] + renamed_entity_id = "sensor.contract_temperature_user_name" + ws.call( + "config/entity_registry/update", + entity_id=original["entity_id"], + new_entity_id=renamed_entity_id, + disabled_by="user", + ) + + # HA's required sequence: marker, device payload, then retained cleanup. + publisher.publish(legacy_topic("temperature"), '{"migrate_discovery":true}') + publisher.publish(device_topic(), json.dumps(device_payload({"temperature": component("temperature", unique)}))) + publisher.publish(legacy_topic("temperature"), "") + + migrated = wait_for_entry(ws, unique) + if migrated["id"] != original_id: + fail("single-to-device migration changed the entity registry ID") + if migrated["entity_id"] != renamed_entity_id or migrated.get("disabled_by") != "user": + fail(f"migration did not preserve user registry settings: {migrated}") + if find_unique(ws.registry_entries(), unique) is None: + fail("migrated entity disappeared from the entity registry") + + # Reverse migration is also ordered: marker the device topic, restore + # legacy discovery, then clear the device topic. This is the protocol + # HA documents for preserving the existing registry entry. + publisher.publish(device_topic(), '{"migrate_discovery":true}') + publisher.publish(legacy_topic("temperature"), json.dumps(legacy_payload("temperature", unique))) + publisher.publish(device_topic(), "") + rolled_back = wait_for_entry(ws, unique) + if rolled_back["id"] != original_id: + fail("device-to-single rollback changed the entity registry ID") + if rolled_back["entity_id"] != renamed_entity_id or rolled_back.get("disabled_by") != "user": + fail(f"rollback did not preserve user registry settings: {rolled_back}") + + # Restore device discovery so the retained-restart check continues to + # exercise the forward migration form used by deployed firmware. + publisher.publish(legacy_topic("temperature"), '{"migrate_discovery":true}') + publisher.publish(device_topic(), json.dumps(device_payload({"temperature": component("temperature", unique)}))) + publisher.publish(legacy_topic("temperature"), "") + remigrated = wait_for_entry(ws, unique) + if remigrated["id"] != original_id: + fail("repeat single-to-device migration changed the entity registry ID") + + # Direct publication is deliberately not a migration protocol. It must not + # create a second registry entry for the same stable unique ID. + direct_unique = f"{DEVICE_ID}_direct" + publisher.publish(legacy_topic("direct"), json.dumps(legacy_payload("direct", direct_unique))) + direct_entry = wait_for_entry(ws, direct_unique) + publisher.publish(device_topic(), json.dumps(device_payload({"direct": component("direct", direct_unique)}))) + time.sleep(2) + entries = [entry for entry in ws.registry_entries() if entry.get("unique_id") == direct_unique] + if len(entries) != 1 or entries[0]["id"] != direct_entry["id"]: + fail("direct device discovery publish created a duplicate registry entity") + + # Device-mode removal is two root updates: platform tombstone then omission. + removable_unique = f"{DEVICE_ID}_removable" + publisher.publish( + device_topic(), + json.dumps(device_payload({ + "anchor": component("anchor"), + "removable": component("removable", removable_unique), + })), + ) + wait_for_entry(ws, removable_unique) + publisher.publish( + device_topic(), + json.dumps(device_payload({ + "anchor": component("anchor"), + "removable": {"p": "sensor"}, + })), + ) + publisher.publish(device_topic(), json.dumps(device_payload({"anchor": component("anchor")}))) + + # HA 2026.5 fixed cleanup for discovered entities that start disabled. + # Exercise the same device-component tombstone/omission sequence so a + # future regression is caught by the stable/dev contract matrix. + disabled_unique = f"{DEVICE_ID}_disabled" + publisher.publish( + device_topic(), + json.dumps(device_payload({ + "anchor": component("anchor"), + "disabled": component( + "disabled", disabled_unique, enabled_by_default=False + ), + })), + ) + disabled_entry = wait_for_entry(ws, disabled_unique) + if disabled_entry.get("disabled_by") != "integration": + fail(f"expected initially disabled entity to be integration-disabled: {disabled_entry}") + publisher.publish( + device_topic(), + json.dumps(device_payload({ + "anchor": component("anchor"), + "disabled": {"p": "sensor"}, + })), + ) + publisher.publish(device_topic(), json.dumps(device_payload({"anchor": component("anchor")}))) + wait_until( + "disabled device component cleanup", + lambda: find_unique(ws.registry_entries(), disabled_unique) is None, + timeout=60, + ) + + # Invalid retained discovery JSON never creates a registry entry. Escaped + # strings are covered by the firmware-native serializer tests. + publisher.publish(legacy_topic("malformed"), "{not-json") + time.sleep(2) + if find_unique(ws.registry_entries(), f"{DEVICE_ID}_malformed"): + fail("malformed discovery payload created an entity") + + # This fixture documents the current HA field contract: def_ent_id is + # allowed on first creation; obsolete obj_id is intentionally absent. + default_payload = legacy_payload("default_name", def_ent_id="contract_default_name") + if "obj_id" in default_payload: + fail("contract fixture accidentally contains obsolete obj_id") + publisher.publish(legacy_topic("default_name"), json.dumps(default_payload)) + wait_for_entry(ws, f"{DEVICE_ID}_default_name") + + STATE.mkdir(parents=True, exist_ok=True) + (STATE / "migration-complete").write_text("ok", encoding="utf-8") + finally: + ws.close() + publisher.close() + + +def retained_restart_contract(): + if not (STATE / "migration-complete").exists(): + fail("retained-restart mode requires the migration contract to run first") + publisher = RetainedPublisher() + token = onboarding_token() + ws = HAWebSocket(token) + try: + migrated = wait_for_entry(ws, f"{DEVICE_ID}_temperature") + if migrated["entity_id"] != "sensor.contract_temperature_user_name": + fail("HA restart lost the user-owned entity rename") + retained = json.loads(publisher.retained_payload(device_topic())) + if "cmps" not in retained: + fail("broker restart check did not receive retained device discovery") + finally: + ws.close() + publisher.close() + + +def main(): + wait_for_home_assistant() + if MODE == "migration": + migration_contract() + elif MODE == "retained-restart": + retained_restart_contract() + else: + fail(f"unknown CONTRACT_MODE: {MODE}") + print(f"Home Assistant MQTT contract mode {MODE} passed") + + +if __name__ == "__main__": + main()