Something went wrong. Try again.
A fork of https://github.com/crosspoint-reader/crosspoint-reader
Something went wrong. Try again.
13 kB · 357 lines
C++
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358#include "FirmwareFlasher.h"
#include <Arduino.h>#include <HalStorage.h>#include <Logging.h>#include <esp_ota_ops.h>#include <esp_partition.h>#include <mbedtls/sha256.h>#include <spi_flash_mmap.h>
#include <algorithm>#include <cstring>#include <memory>
#include "FirmwareBoardTag.h"#include "OtaBootSwitch.h"
namespace firmware_flash {
namespace {constexpr uint8_t ESP_IMAGE_MAGIC = 0xE9;constexpr size_t MIN_FIRMWARE_SIZE = 64 * 1024;constexpr size_t SEC = SPI_FLASH_SEC_SIZE; // 4 KiBconstexpr size_t BLK = 64 * 1024; // 64 KiB block-erase granularityconstexpr size_t CHUNK = 4096;constexpr size_t SHA_TRAILER = 32;constexpr uint8_t CHECKSUM_SEED = 0xEF;constexpr size_t HEADER_SIZE = 24;constexpr size_t SEG_HEADER_SIZE = 8;} // namespace
const char* resultName(Result r) { switch (r) { case Result::OK: return "OK"; case Result::OPEN_FAIL: return "OPEN_FAIL"; case Result::TOO_SMALL: return "TOO_SMALL"; case Result::TOO_LARGE: return "TOO_LARGE"; case Result::BAD_MAGIC: return "BAD_MAGIC"; case Result::BAD_SEGMENTS: return "BAD_SEGMENTS"; case Result::BAD_CHECKSUM: return "BAD_CHECKSUM"; case Result::BAD_SHA: return "BAD_SHA"; case Result::BAD_CHIP: return "BAD_CHIP"; case Result::WRONG_BOARD: return "WRONG_BOARD"; case Result::BAD_SIZE: return "BAD_SIZE"; case Result::NO_PARTITION: return "NO_PARTITION"; case Result::OOM: return "OOM"; case Result::READ_FAIL: return "READ_FAIL"; case Result::ERASE_FAIL: return "ERASE_FAIL"; case Result::WRITE_FAIL: return "WRITE_FAIL"; case Result::OTADATA_FAIL: return "OTADATA_FAIL"; } return "?";}
uint16_t runningPartitionChipId() { // esp_partition_read hits SPI flash; cache the running slot's chip_id so we // only pay that cost once per boot. The running image is immutable at // runtime, so a function-local static is safe here. static uint16_t cached = [] { const esp_partition_t* run = esp_ota_get_running_partition(); if (!run) return static_cast<uint16_t>(0xFFFF); uint16_t id = 0xFFFF; // chip_id sits at offset 12 of esp_image_header_t. memcpy target is a // uint16_t local, so RISC-V alignment is guaranteed. if (esp_partition_read(run, 12, &id, sizeof(id)) != ESP_OK) return static_cast<uint16_t>(0xFFFF); return id; }(); return cached;}
namespace {// Stream `length` bytes from `file` starting at the current read offset, feeding them through// both the XOR-checksum and SHA256 accumulators. Used by validateImageFile so the whole image// is verified end-to-end without holding it in RAM (ESP32-C3 only has ~380 KB).Result feedHashAndChecksum(HalFile& file, size_t length, uint8_t* xorAccum, mbedtls_sha256_context* sha, uint8_t* buf, board_tag::Scanner* tagScanner) { size_t remaining = length; while (remaining > 0) { const size_t want = std::min<size_t>(CHUNK, remaining); const int got = file.read(buf, want); if (got <= 0 || static_cast<size_t>(got) != want) return Result::READ_FAIL; if (sha) mbedtls_sha256_update(sha, buf, want); if (tagScanner) tagScanner->feed(buf, want); if (xorAccum) { uint8_t acc = *xorAccum; for (size_t i = 0; i < want; i++) acc ^= buf[i]; *xorAccum = acc; } remaining -= want; } return Result::OK;}} // namespace
Result validateImageFile(const char* sdPath, size_t partitionSize) { HalFile file; if (!Storage.openFileForRead("FLASH", sdPath, file) || !file) { LOG_ERR("FLASH", "validate: open failed: %s", sdPath); return Result::OPEN_FAIL; }
const size_t fileSize = file.fileSize(); if (fileSize < MIN_FIRMWARE_SIZE) { LOG_ERR("FLASH", "validate: too small: %u", static_cast<unsigned>(fileSize)); file.close(); return Result::TOO_SMALL; } if (partitionSize > 0 && fileSize > partitionSize) { LOG_ERR("FLASH", "validate: too large: %u > %u", static_cast<unsigned>(fileSize), static_cast<unsigned>(partitionSize)); file.close(); return Result::TOO_LARGE; }
uint8_t header[HEADER_SIZE]; if (file.read(header, HEADER_SIZE) != static_cast<int>(HEADER_SIZE)) { LOG_ERR("FLASH", "validate: header read failed"); file.close(); return Result::READ_FAIL; } if (header[0] != ESP_IMAGE_MAGIC) { LOG_ERR("FLASH", "validate: bad magic 0x%02X", header[0]); file.close(); return Result::BAD_MAGIC; } // Reject an image built for a different MCU family before it can brick the // device. chip_id lives at esp_image_header_t offset 12; compare it against // the running slot's own chip_id (self-describing, no chip enumeration). uint16_t imageChip; std::memcpy(&imageChip, header + 12, sizeof(imageChip)); const uint16_t deviceChip = runningPartitionChipId(); if (deviceChip != 0xFFFF && imageChip != deviceChip) { LOG_ERR("FLASH", "validate: wrong chip: image=0x%04X device=0x%04X", imageChip, deviceChip); file.close(); return Result::BAD_CHIP; } const uint8_t segCount = header[1]; const bool hashAppended = header[23] != 0;
auto buf = std::unique_ptr<uint8_t[]>(new (std::nothrow) uint8_t[CHUNK]); if (!buf) { file.close(); return Result::OOM; }
mbedtls_sha256_context shaCtx; mbedtls_sha256_init(&shaCtx); mbedtls_sha256_starts(&shaCtx, /*is224=*/0); mbedtls_sha256_update(&shaCtx, header, HEADER_SIZE);
uint8_t xorAccum = CHECKSUM_SEED; size_t pos = HEADER_SIZE; // Board tag: scanned from the same segment stream the hash pass already // reads, so the check is free of extra I/O. Only a present-and-mismatched // tag rejects; untagged images (forks, other projects) pass. board_tag::Scanner tagScanner;
for (uint8_t i = 0; i < segCount; i++) { if (pos + SEG_HEADER_SIZE > fileSize) { LOG_ERR("FLASH", "validate: seg %u header overruns EOF at %u", i, static_cast<unsigned>(pos)); mbedtls_sha256_free(&shaCtx); file.close(); return Result::BAD_SEGMENTS; } uint8_t segHdr[SEG_HEADER_SIZE]; if (file.read(segHdr, SEG_HEADER_SIZE) != static_cast<int>(SEG_HEADER_SIZE)) { mbedtls_sha256_free(&shaCtx); file.close(); return Result::READ_FAIL; } mbedtls_sha256_update(&shaCtx, segHdr, SEG_HEADER_SIZE); pos += SEG_HEADER_SIZE;
uint32_t dataLen; std::memcpy(&dataLen, segHdr + 4, sizeof(dataLen)); if (pos + dataLen > fileSize) { LOG_ERR("FLASH", "validate: seg %u data overruns EOF (%u + %u > %u)", i, static_cast<unsigned>(pos), static_cast<unsigned>(dataLen), static_cast<unsigned>(fileSize)); mbedtls_sha256_free(&shaCtx); file.close(); return Result::BAD_SEGMENTS; }
const Result feedRes = feedHashAndChecksum(file, dataLen, &xorAccum, &shaCtx, buf.get(), &tagScanner); if (feedRes != Result::OK) { mbedtls_sha256_free(&shaCtx); file.close(); return feedRes; } pos += dataLen; }
if (tagScanner.mismatch()) { LOG_ERR("FLASH", "validate: wrong board: image=%s device=%.*s", tagScanner.foundName(), static_cast<int>(board_tag::boardNameLen()), board_tag::boardName()); mbedtls_sha256_free(&shaCtx); file.close(); return Result::WRONG_BOARD; }
// pad_end is the 16-byte aligned offset at which the checksum byte sits at pad_end - 1. const size_t padEnd = (pos + 16) & ~static_cast<size_t>(15); const size_t expectedTotal = padEnd + (hashAppended ? SHA_TRAILER : 0); if (expectedTotal != fileSize) { LOG_ERR("FLASH", "validate: size mismatch body+pad=%u sha=%u expected=%u actual=%u", static_cast<unsigned>(padEnd), static_cast<unsigned>(hashAppended ? SHA_TRAILER : 0), static_cast<unsigned>(expectedTotal), static_cast<unsigned>(fileSize)); mbedtls_sha256_free(&shaCtx); file.close(); return Result::BAD_SIZE; }
// Read the padding bytes (which include the stored checksum at the last byte) into the SHA stream. const size_t padLen = padEnd - pos; uint8_t padBuf[16]; if (padLen > sizeof(padBuf)) { mbedtls_sha256_free(&shaCtx); file.close(); return Result::BAD_SIZE; } if (padLen > 0 && file.read(padBuf, padLen) != static_cast<int>(padLen)) { mbedtls_sha256_free(&shaCtx); file.close(); return Result::READ_FAIL; } mbedtls_sha256_update(&shaCtx, padBuf, padLen);
const uint8_t storedChecksum = padBuf[padLen - 1]; if ((xorAccum & 0xFF) != storedChecksum) { LOG_ERR("FLASH", "validate: checksum mismatch computed=0x%02X stored=0x%02X", xorAccum, storedChecksum); mbedtls_sha256_free(&shaCtx); file.close(); return Result::BAD_CHECKSUM; }
if (hashAppended) { uint8_t computed[SHA_TRAILER]; mbedtls_sha256_finish(&shaCtx, computed); uint8_t stored[SHA_TRAILER]; if (file.read(stored, SHA_TRAILER) != static_cast<int>(SHA_TRAILER)) { mbedtls_sha256_free(&shaCtx); file.close(); return Result::READ_FAIL; } if (std::memcmp(computed, stored, SHA_TRAILER) != 0) { LOG_ERR("FLASH", "validate: SHA256 mismatch"); mbedtls_sha256_free(&shaCtx); file.close(); return Result::BAD_SHA; } }
mbedtls_sha256_free(&shaCtx); file.close(); return Result::OK;}
Result flashFromSdPath(const char* sdPath, ProgressCb onProgress, void* ctx, bool alreadyValidated) { // Resolve destination first so we can size-check during validation. The full image-integrity // pass below verifies header, segment table, XOR checksum and SHA256 trailer end-to-end before // we touch otadata, so a truncated/corrupted .bin can never become the next boot target. const esp_partition_t* dest = esp_ota_get_next_update_partition(nullptr); if (!dest) { LOG_ERR("FLASH", "no next-update partition"); return Result::NO_PARTITION; }
// When the caller already ran validateImageFile() against this same partition // size (e.g. SdFirmwareUpdateActivity validates before the confirmation // prompt), skip the redundant integrity scan. We still keep the partition // lookup so the rest of the flashing path stays unchanged. if (!alreadyValidated) { const Result validateRes = validateImageFile(sdPath, dest->size); if (validateRes != Result::OK) { LOG_ERR("FLASH", "image validation failed: %s", resultName(validateRes)); return validateRes; } }
HalFile file; if (!Storage.openFileForRead("FLASH", sdPath, file) || !file) { LOG_ERR("FLASH", "open failed: %s", sdPath); return Result::OPEN_FAIL; }
const size_t firmwareSize = file.fileSize(); LOG_INF("FLASH", "src=%s size=%u dest=%s @0x%x partsize=%u", sdPath, static_cast<unsigned>(firmwareSize), dest->label, static_cast<unsigned>(dest->address), static_cast<unsigned>(dest->size));
auto buffer = std::unique_ptr<uint8_t[]>(new (std::nothrow) uint8_t[CHUNK]); if (!buffer) { LOG_ERR("FLASH", "OOM"); file.close(); return Result::OOM; }
// Interleave erase + write so the progress bar advances 0→100% smoothly // rather than stalling for several seconds during a single up-front erase. size_t streamPos = 0; size_t erasedUpto = 0; while (streamPos < firmwareSize) { if (streamPos >= erasedUpto) { size_t eraseLen = std::min<size_t>(BLK, dest->size - streamPos); eraseLen = (eraseLen + SEC - 1) & ~(SEC - 1); eraseLen = std::min<size_t>(eraseLen, dest->size - streamPos); if (esp_partition_erase_range(dest, streamPos, eraseLen) != ESP_OK) { LOG_ERR("FLASH", "erase @%u (len=%u) failed", static_cast<unsigned>(streamPos), static_cast<unsigned>(eraseLen)); file.close(); return Result::ERASE_FAIL; } erasedUpto = streamPos + eraseLen; }
const size_t want = std::min<size_t>(CHUNK, firmwareSize - streamPos); const int read = file.read(buffer.get(), want); if (read <= 0 || static_cast<size_t>(read) != want) { LOG_ERR("FLASH", "read @%u: got=%d want=%u", static_cast<unsigned>(streamPos), read, static_cast<unsigned>(want)); file.close(); return Result::READ_FAIL; } if (esp_partition_write(dest, streamPos, buffer.get(), want) != ESP_OK) { LOG_ERR("FLASH", "write @%u failed", static_cast<unsigned>(streamPos)); file.close(); return Result::WRITE_FAIL; } streamPos += want; if (onProgress) onProgress(streamPos, firmwareSize, ctx); delay(1); } file.close();
if (!ota_boot::switchTo(dest)) { LOG_ERR("FLASH", "otadata switch failed"); return Result::OTADATA_FAIL; } return Result::OK;}
} // namespace firmware_flash