diff --git a/lib/ota/src/update_parser.h b/lib/ota/src/update_parser.h index 0c6114c..128a03c 100644 --- a/lib/ota/src/update_parser.h +++ b/lib/ota/src/update_parser.h @@ -50,6 +50,8 @@ class UpdateParser { void feed(const uint8_t* data, size_t len); bool end(); // no more data: true if the update was installed + // The whole image the header announced has arrived (the sender need not close the connection). + bool complete() const { return state_ == State::Image && received_ == imageSize_; } State state() const { return state_; } const std::string& error() const { return error_; } diff --git a/scripts/ota_push.py b/scripts/ota_push.py index aaf06a3..b035e8b 100755 --- a/scripts/ota_push.py +++ b/scripts/ota_push.py @@ -17,16 +17,29 @@ def main(): data = open(path, "rb").read() with socket.create_connection((host, PORT), timeout=15) as s: s.settimeout(60) - sent = 0 + sent, last = 0, -1 while sent < len(data): chunk = data[sent:sent + 4096] s.sendall(chunk) sent += len(chunk) - print(f"\rsending {sent * 100 // len(data):3d}%", end="", flush=True) - s.shutdown(socket.SHUT_WR) # end of file: the device checks it and answers + pct = sent * 100 // len(data) + if pct != last: + print(f"\rsending {pct:3d}%", end="", flush=True) + last = pct + # The header announces the image size, so current firmware answers once it has it all. + # Firmware from before that change only knows the file ended when we half-close. reply = b"" - while not reply.endswith(b"\n"): - part = s.recv(256) + s.settimeout(3) + try: + reply = s.recv(256) + except socket.timeout: + s.shutdown(socket.SHUT_WR) + s.settimeout(60) + while reply == b"" or not reply.endswith(b"\n"): + try: + part = s.recv(256) + except socket.timeout: + break if not part: break reply += part diff --git a/src/services/update_service.cpp b/src/services/update_service.cpp index 11fb45e..44c498a 100644 --- a/src/services/update_service.cpp +++ b/src/services/update_service.cpp @@ -132,6 +132,10 @@ void UpdateService::install(UpdateSource& source, const char* via) { if (incoming_.empty() && parser.state() == UpdateParser::State::Image) incoming_ = parser.version(); percent_ = parser.percent(); if (parser.state() == UpdateParser::State::Failed) break; + if (parser.complete()) { // all announced bytes are here: answer while the connection is open + ok = parser.end(); + break; + } } if (!ok && parser.error().empty()) parser.end(); // e.g. a stall: abort the slot diff --git a/test/test_ota/test_ota.cpp b/test/test_ota/test_ota.cpp index 27f7589..c8af9a8 100644 --- a/test/test_ota/test_ota.cpp +++ b/test/test_ota/test_ota.cpp @@ -138,6 +138,18 @@ void test_progress_counts_image_bytes() { TEST_ASSERT_EQUAL(25, p.percent()); } +void test_complete_once_the_declared_image_size_has_arrived() { + FakeVerifier v; + MemorySink sink; + auto file = makeFile("v0.3.0", image(300)); + UpdateParser p(v, sink, 100000, "v0.2.1"); + p.feed(file.data(), file.size() - 1); + TEST_ASSERT_FALSE(p.complete()); + p.feed(file.data() + file.size() - 1, 1); + TEST_ASSERT_TRUE(p.complete()); + TEST_ASSERT_TRUE(p.end()); +} + void test_not_an_update_file_is_refused_before_writing() { FakeVerifier v; MemorySink sink; @@ -273,6 +285,7 @@ int main() { RUN_TEST(test_valid_file_in_one_chunk_is_installed); RUN_TEST(test_valid_file_byte_by_byte); RUN_TEST(test_progress_counts_image_bytes); + RUN_TEST(test_complete_once_the_declared_image_size_has_arrived); RUN_TEST(test_not_an_update_file_is_refused_before_writing); RUN_TEST(test_bad_signature_is_refused_before_writing); RUN_TEST(test_tampered_version_breaks_the_signature);