diff --git a/src/detail/copy_until_terminator.hpp b/src/detail/copy_until_terminator.hpp new file mode 100644 index 0000000..4c1d152 --- /dev/null +++ b/src/detail/copy_until_terminator.hpp @@ -0,0 +1,30 @@ +#pragma once + +#include "matches_suffix.hpp" + +namespace cpp_bindings_windows::detail +{ +struct CopyUntilTerminatorResult +{ + int bytes_copied = 0; + bool terminator_found = false; +}; + +inline auto copyUntilTerminator(unsigned char *output, int output_size, const unsigned char *input, int input_size, + const unsigned char *terminator, int terminator_size) -> CopyUntilTerminatorResult +{ + CopyUntilTerminatorResult result; + while (result.bytes_copied < input_size) + { + output[output_size + result.bytes_copied] = input[result.bytes_copied]; + ++result.bytes_copied; + + if (matchesSuffix(output, output_size + result.bytes_copied, terminator, terminator_size)) + { + result.terminator_found = true; + break; + } + } + return result; +} +} // namespace cpp_bindings_windows::detail diff --git a/src/detail/handle_types.hpp b/src/detail/handle_types.hpp index fa5ede4..61b53d9 100644 --- a/src/detail/handle_types.hpp +++ b/src/detail/handle_types.hpp @@ -1,6 +1,7 @@ #pragma once #include "common_types.hpp" +#include "read_ahead_buffer.hpp" #include "windows.hpp" #include @@ -45,6 +46,7 @@ struct HandleState std::atomic bytes_written_total{0}; std::atomic abort_read{false}; std::atomic abort_write{false}; + ReadAheadBuffer read_ahead; std::mutex pending_io_mutex; OVERLAPPED *pending_read = nullptr; OVERLAPPED *pending_write = nullptr; diff --git a/src/detail/read_ahead_buffer.hpp b/src/detail/read_ahead_buffer.hpp new file mode 100644 index 0000000..04dfdd1 --- /dev/null +++ b/src/detail/read_ahead_buffer.hpp @@ -0,0 +1,67 @@ +#pragma once + +#include "copy_until_terminator.hpp" + +#include +#include +#include +#include + +namespace cpp_bindings_windows::detail +{ +class ReadAheadBuffer +{ + public: + auto append(const unsigned char *data, int data_size) -> void + { + if (data == nullptr || data_size <= 0) + { + return; + } + + std::lock_guard lock(mutex_); + bytes_.insert(bytes_.end(), data, data + data_size); + } + + auto consume(unsigned char *output, int output_size, int max_bytes, const unsigned char *terminator, + int terminator_size) -> CopyUntilTerminatorResult + { + CopyUntilTerminatorResult result; + if (output == nullptr || max_bytes <= 0) + { + return result; + } + + std::lock_guard lock(mutex_); + while (result.bytes_copied < max_bytes && !bytes_.empty()) + { + output[output_size + result.bytes_copied] = bytes_.front(); + bytes_.pop_front(); + ++result.bytes_copied; + + if (matchesSuffix(output, output_size + result.bytes_copied, terminator, terminator_size)) + { + result.terminator_found = true; + break; + } + } + return result; + } + + auto clear() -> void + { + std::lock_guard lock(mutex_); + bytes_.clear(); + } + + [[nodiscard]] auto size() const -> int + { + std::lock_guard lock(mutex_); + return bytes_.size() > static_cast(INT_MAX) ? INT_MAX : static_cast(bytes_.size()); + } + + private: + mutable std::mutex mutex_; + std::deque bytes_; +}; +} // namespace cpp_bindings_windows::detail diff --git a/src/detail/read_impl.hpp b/src/detail/read_impl.hpp index 92ca7e9..473f629 100644 --- a/src/detail/read_impl.hpp +++ b/src/detail/read_impl.hpp @@ -2,18 +2,22 @@ #include "acquire_handle_context.hpp" #include "bytes_waiting.hpp" +#include "consume_abort.hpp" +#include "copy_until_terminator.hpp" #include "fail_win32.hpp" -#include "matches_suffix.hpp" -#include "multiplier_timeout.hpp" #include "note_bytes_transferred.hpp" #include "read_chunk.hpp" +#include "read_timeout.hpp" #include #include +#include namespace cpp_bindings_windows::detail { +inline constexpr int kTerminatedReadChunkSize = 4096; + inline auto readImpl(int64_t handle, void *buffer, int buffer_size, int timeout_ms, int multiplier, const unsigned char *terminator, int terminator_size, ErrorCallbackT error_callback) -> int { @@ -35,9 +39,24 @@ inline auto readImpl(int64_t handle, void *buffer, int buffer_size, int timeout_ { return handle_status; } + if (consumeAbort(context.state, Operation::kRead)) + { + return cpp_core::failMsg(callback, static_cast(StatusCode::Io::kAbortReadError), + "Read aborted"); + } auto *output = static_cast(buffer); - int total_read = 0; + const auto buffered = context.state->read_ahead.consume(output, 0, buffer_size, terminator, terminator_size); + int total_read = buffered.bytes_copied; + if (total_read > 0) + { + noteBytesTransferred(context.state, Operation::kRead, total_read); + } + if (buffered.terminator_found) + { + return total_read; + } + while (total_read < buffer_size) { int chunk_size = 1; @@ -48,12 +67,14 @@ inline auto readImpl(int64_t handle, void *buffer, int buffer_size, int timeout_ { return failWin32(callback, static_cast(StatusCode::Control::kGetStateError)); } + chunk_size = waiting > 0 ? std::min(waiting, buffer_size - total_read) : 1; } - const int current_timeout = - total_read == 0 ? cpp_core::clampTimeout(timeout_ms) : multiplierTimeout(timeout_ms, multiplier); - const auto result = readChunk(context, output + total_read, chunk_size, current_timeout); + const int current_timeout = readTimeout(timeout_ms, multiplier, total_read == 0, terminator_size > 0); + std::array chunk{}; + unsigned char *destination = terminator_size > 0 ? chunk.data() : output + total_read; + const auto result = readChunk(context, destination, chunk_size, current_timeout); if (result.outcome == IoOutcome::kTimedOut) { return total_read; @@ -73,9 +94,24 @@ inline auto readImpl(int64_t handle, void *buffer, int buffer_size, int timeout_ return total_read; } - noteBytesTransferred(context.state, Operation::kRead, result.bytes_transferred); - total_read += result.bytes_transferred; - if (matchesSuffix(output, total_read, terminator, terminator_size)) + if (terminator_size <= 0) + { + noteBytesTransferred(context.state, Operation::kRead, result.bytes_transferred); + total_read += result.bytes_transferred; + continue; + } + + const auto copied = copyUntilTerminator(output, total_read, chunk.data(), result.bytes_transferred, terminator, + terminator_size); + if (copied.bytes_copied < result.bytes_transferred) + { + context.state->read_ahead.append(chunk.data() + copied.bytes_copied, + result.bytes_transferred - copied.bytes_copied); + } + + noteBytesTransferred(context.state, Operation::kRead, copied.bytes_copied); + total_read += copied.bytes_copied; + if (copied.terminator_found) { return total_read; } diff --git a/src/detail/read_timeout.hpp b/src/detail/read_timeout.hpp new file mode 100644 index 0000000..3816998 --- /dev/null +++ b/src/detail/read_timeout.hpp @@ -0,0 +1,21 @@ +#pragma once + +#include "multiplier_timeout.hpp" + +#include + +namespace cpp_bindings_windows::detail +{ +inline auto readTimeout(int timeout_ms, int multiplier, bool first_read, bool terminated_read) -> int +{ + // A terminated read must keep waiting for the terminator even when callers + // use the raw-read default multiplier of zero. Otherwise it only drains the + // bytes that happened to be queued when the call started. + if (first_read || (terminated_read && multiplier <= 0)) + { + return cpp_core::clampTimeout(timeout_ms); + } + + return multiplierTimeout(timeout_ms, multiplier); +} +} // namespace cpp_bindings_windows::detail diff --git a/src/detail/wait_for_pending_io.hpp b/src/detail/wait_for_pending_io.hpp index 839f726..87dce9d 100644 --- a/src/detail/wait_for_pending_io.hpp +++ b/src/detail/wait_for_pending_io.hpp @@ -15,13 +15,22 @@ inline auto waitForPendingIo(HANDLE handle, const std::shared_ptr & if (wait_result == WAIT_TIMEOUT) { (void)CancelIoEx(handle, overlapped); - DWORD ignored = 0; - (void)GetOverlappedResult(handle, overlapped, &ignored, TRUE); + DWORD transferred = 0; + const BOOL completed = GetOverlappedResult(handle, overlapped, &transferred, TRUE); + const DWORD error = completed != FALSE ? ERROR_SUCCESS : GetLastError(); if (finishPendingIo(state, operation, overlapped)) { return {.outcome = IoOutcome::kAborted}; } - return {.outcome = IoOutcome::kTimedOut}; + if (completed != FALSE || transferred > 0) + { + return {.outcome = IoOutcome::kCompleted, .bytes_transferred = static_cast(transferred)}; + } + if (error == ERROR_OPERATION_ABORTED) + { + return {.outcome = IoOutcome::kTimedOut}; + } + return {.outcome = IoOutcome::kError, .error = error}; } if (wait_result != WAIT_OBJECT_0) diff --git a/src/read_ahead_buffer.test.cpp b/src/read_ahead_buffer.test.cpp new file mode 100644 index 0000000..6582cec --- /dev/null +++ b/src/read_ahead_buffer.test.cpp @@ -0,0 +1,86 @@ +#include "detail/read_ahead_buffer.hpp" +#include "detail/copy_until_terminator.hpp" + +#include +#include +#include +#include +#include +#include + +#include + +namespace +{ +using cpp_bindings_windows::detail::copyUntilTerminator; +using cpp_bindings_windows::detail::ReadAheadBuffer; +} // namespace + +TEST(CopyUntilTerminatorTest, StopsExactlyAfterTerminatorInLongChunk) +{ + std::string input(1000, 'A'); + input += "\ntail"; + std::vector output(input.size()); + constexpr unsigned char newline = '\n'; + + const auto result = copyUntilTerminator(output.data(), 0, reinterpret_cast(input.data()), + static_cast(input.size()), &newline, 1); + + ASSERT_TRUE(result.terminator_found); + ASSERT_EQ(result.bytes_copied, 1001); + EXPECT_EQ(std::string_view(reinterpret_cast(output.data()), result.bytes_copied), + std::string_view(input.data(), 1001)); +} + +TEST(CopyUntilTerminatorTest, FindsSequenceAcrossChunkBoundary) +{ + std::array output{}; + constexpr std::string_view prefix = "prefix-EN"; + constexpr std::string_view input = "D-tail"; + constexpr unsigned char terminator[] = {'E', 'N', 'D'}; + std::copy(prefix.begin(), prefix.end(), output.begin()); + + const auto result = copyUntilTerminator( + output.data(), static_cast(prefix.size()), reinterpret_cast(input.data()), + static_cast(input.size()), terminator, static_cast(std::size(terminator))); + + ASSERT_TRUE(result.terminator_found); + ASSERT_EQ(result.bytes_copied, 1); + EXPECT_EQ(std::string_view(reinterpret_cast(output.data()), prefix.size() + result.bytes_copied), + "prefix-END"); +} + +TEST(ReadAheadBufferTest, PreservesBytesAfterTerminatorForNextRead) +{ + ReadAheadBuffer buffer; + constexpr std::string_view input = "line\nnext"; + buffer.append(reinterpret_cast(input.data()), static_cast(input.size())); + + std::array first_output{}; + constexpr unsigned char newline = '\n'; + const auto first = buffer.consume(first_output.data(), 0, static_cast(first_output.size()), &newline, 1); + + ASSERT_TRUE(first.terminator_found); + ASSERT_EQ(first.bytes_copied, 5); + EXPECT_EQ(std::string_view(reinterpret_cast(first_output.data()), first.bytes_copied), "line\n"); + ASSERT_EQ(buffer.size(), 4); + + std::array second_output{}; + const auto second = buffer.consume(second_output.data(), 0, static_cast(second_output.size()), nullptr, 0); + + EXPECT_FALSE(second.terminator_found); + ASSERT_EQ(second.bytes_copied, 4); + EXPECT_EQ(std::string_view(reinterpret_cast(second_output.data()), second.bytes_copied), "next"); + EXPECT_EQ(buffer.size(), 0); +} + +TEST(ReadAheadBufferTest, ClearDropsBufferedBytes) +{ + ReadAheadBuffer buffer; + constexpr std::array input = {'a', 'b', 'c'}; + buffer.append(input.data(), static_cast(input.size())); + + buffer.clear(); + + EXPECT_EQ(buffer.size(), 0); +} diff --git a/src/read_timeout.test.cpp b/src/read_timeout.test.cpp new file mode 100644 index 0000000..6600486 --- /dev/null +++ b/src/read_timeout.test.cpp @@ -0,0 +1,26 @@ +#include "detail/read_timeout.hpp" + +#include + +namespace cpp_bindings_windows::detail +{ +TEST(ReadTimeoutTest, UsesBaseTimeoutForFirstRead) +{ + EXPECT_EQ(readTimeout(250, 0, true, false), 250); +} + +TEST(ReadTimeoutTest, RawReadWithZeroMultiplierDoesNotWaitForMoreData) +{ + EXPECT_EQ(readTimeout(250, 0, false, false), 0); +} + +TEST(ReadTimeoutTest, TerminatedReadWithZeroMultiplierKeepsWaitingForTerminator) +{ + EXPECT_EQ(readTimeout(250, 0, false, true), 250); +} + +TEST(ReadTimeoutTest, TerminatedReadStillAppliesPositiveMultiplier) +{ + EXPECT_EQ(readTimeout(250, 3, false, true), 750); +} +} // namespace cpp_bindings_windows::detail diff --git a/src/serial_clear_buffer_in.cpp b/src/serial_clear_buffer_in.cpp index b45afe0..226026b 100644 --- a/src/serial_clear_buffer_in.cpp +++ b/src/serial_clear_buffer_in.cpp @@ -21,6 +21,7 @@ extern "C" cpp_bindings_windows::detail::effectiveErrorCallback(error_callback), static_cast(cpp_core::StatusCode::Io::kClearBufferInError)); } + context.state->read_ahead.clear(); return static_cast(cpp_core::StatusCode::kSuccess); } diff --git a/src/serial_in_bytes_waiting.cpp b/src/serial_in_bytes_waiting.cpp index fd9f063..1496c5f 100644 --- a/src/serial_in_bytes_waiting.cpp +++ b/src/serial_in_bytes_waiting.cpp @@ -4,6 +4,9 @@ #include "detail/bytes_waiting.hpp" #include "detail/fail_win32.hpp" +#include +#include + extern "C" { @@ -23,7 +26,8 @@ extern "C" cpp_bindings_windows::detail::effectiveErrorCallback(error_callback), static_cast(cpp_core::StatusCode::Control::kGetStateError)); } - return waiting; + const auto total_waiting = static_cast(waiting) + context.state->read_ahead.size(); + return total_waiting > INT_MAX ? INT_MAX : static_cast(total_waiting); } } // extern "C" diff --git a/tests/serial_arduino.test.cpp b/tests/serial_arduino.test.cpp index 2546c92..501c74d 100644 --- a/tests/serial_arduino.test.cpp +++ b/tests/serial_arduino.test.cpp @@ -1,6 +1,9 @@ +#include #include #include #include +#include +#include #include #include #include @@ -14,6 +17,7 @@ #include #include #include +#include namespace { @@ -135,6 +139,42 @@ TEST_F(SerialArduinoTest, MultipleEchoCycles) } } +TEST_F(SerialArduinoTest, ReadUntilHandlesLongPayload) +{ + ASSERT_EQ(serialClearBufferIn(handle_, nullptr), 0); + + std::string message(1000, 'A'); + message.push_back('\n'); + ASSERT_EQ(serialWrite(handle_, message.data(), static_cast(message.size()), 3000, 1, nullptr), + static_cast(message.size())); + + std::vector buffer(message.size()); + char newline = '\n'; + const int read_bytes = + serialReadUntil(handle_, buffer.data(), static_cast(buffer.size()), 3000, 0, &newline, nullptr); + + ASSERT_EQ(read_bytes, static_cast(message.size())); + EXPECT_EQ(std::string_view(buffer.data(), static_cast(read_bytes)), message); +} + +TEST_F(SerialArduinoTest, ReadUntilSequenceHandlesLongPayload) +{ + ASSERT_EQ(serialClearBufferIn(handle_, nullptr), 0); + + std::string message(1000, 'A'); + message += "\r\n"; + ASSERT_EQ(serialWrite(handle_, message.data(), static_cast(message.size()), 3000, 1, nullptr), + static_cast(message.size())); + + std::vector buffer(message.size()); + char sequence[] = "\r\n"; + const int read_bytes = + serialReadUntilSequence(handle_, buffer.data(), static_cast(buffer.size()), 3000, 0, sequence, nullptr); + + ASSERT_EQ(read_bytes, static_cast(message.size())); + EXPECT_EQ(std::string_view(buffer.data(), static_cast(read_bytes)), message); +} + TEST_F(SerialArduinoTest, ReadTimeout) { char buffer[256];