diff --git a/io.c b/io.c index 522c8da0484f66..53aa9cda113ccb 100644 --- a/io.c +++ b/io.c @@ -1360,6 +1360,55 @@ rb_io_write_memory(rb_io_t *fptr, const void *buf, size_t count) return (ssize_t)rb_io_blocking_region_wait(fptr, internal_write_func, &iis, RUBY_IO_WRITABLE); } +struct io_write_buffer_arguments { + VALUE scheduler; + rb_io_t *fptr; + const void *buffer; + size_t size; + ssize_t result; + int error; +}; + +static VALUE +io_write_buffer_fiber_scheduler_body(VALUE argument) +{ + struct io_write_buffer_arguments *args = (struct io_write_buffer_arguments *)argument; + VALUE result = rb_fiber_scheduler_io_write_memory(args->scheduler, args->fptr->self, + args->buffer, args->size); + + if (!UNDEF_P(result)) { + args->result = rb_fiber_scheduler_io_result_apply(result); + args->error = errno; + } + return result; +} + +static bool +io_write_buffer_fiber_scheduler(VALUE scheduler, rb_io_t *fptr, const void *buffer, + size_t size, ssize_t *result) +{ + struct io_write_buffer_arguments args = { + .scheduler = scheduler, + .fptr = fptr, + .buffer = buffer, + .size = size, + }; + int state = 0; + VALUE ret = rb_protect(io_write_buffer_fiber_scheduler_body, (VALUE)&args, &state); + + if (state) { + /* Retrying after unwinding without a byte count may replay written bytes. */ + fptr->wbuf.off = 0; + fptr->wbuf.len = 0; + rb_jump_tag(state); + } + if (UNDEF_P(ret)) return false; + + *result = args.result; + errno = args.error; + return true; +} + #ifdef HAVE_WRITEV static ssize_t rb_writev_internal(rb_io_t *fptr, const struct iovec *iov, int iovcnt) @@ -1371,10 +1420,17 @@ rb_writev_internal(rb_io_t *fptr, const struct iovec *iov, int iovcnt) VALUE scheduler = rb_fiber_scheduler_current_for_threadptr(th); if (scheduler != Qnil) { // This path assumes at least one `iov`: - VALUE result = rb_fiber_scheduler_io_write_memory(scheduler, fptr->self, iov[0].iov_base, iov[0].iov_len); - - if (!UNDEF_P(result)) { - return rb_fiber_scheduler_io_result_apply(result); + if (fptr->wbuf.len && iov[0].iov_base == fptr->wbuf.ptr + fptr->wbuf.off) { + ssize_t result; + if (io_write_buffer_fiber_scheduler(scheduler, fptr, iov[0].iov_base, iov[0].iov_len, &result)) { + return result; + } + } + else { + VALUE result = rb_fiber_scheduler_io_write_memory(scheduler, fptr->self, iov[0].iov_base, iov[0].iov_len); + if (!UNDEF_P(result)) { + return rb_fiber_scheduler_io_result_apply(result); + } } } @@ -1425,16 +1481,15 @@ io_flush_buffer_sync(void *arg) static inline VALUE io_flush_buffer_fiber_scheduler(VALUE scheduler, rb_io_t *fptr) { - VALUE ret = rb_fiber_scheduler_io_write_memory(scheduler, fptr->self, fptr->wbuf.ptr+fptr->wbuf.off, fptr->wbuf.len); - if (!UNDEF_P(ret)) { - ssize_t result = rb_fiber_scheduler_io_result_apply(ret); + ssize_t result; + if (io_write_buffer_fiber_scheduler(scheduler, fptr, fptr->wbuf.ptr + fptr->wbuf.off, fptr->wbuf.len, &result)) { if (result > 0) { fptr->wbuf.off += result; fptr->wbuf.len -= result; } return result >= 0 ? (VALUE)0 : (VALUE)-1; } - return ret; + return RUBY_Qundef; } static VALUE diff --git a/test/fiber/test_io.rb b/test/fiber/test_io.rb index eea06f97c829e7..16aa72040c708d 100644 --- a/test/fiber/test_io.rb +++ b/test/fiber/test_io.rb @@ -5,6 +5,22 @@ class TestFiberIO < Test::Unit::TestCase MESSAGE = "Hello World" + class YieldAfterWriteScheduler < IOBufferScheduler + def initialize + super + @yielded = false + end + + def io_write(...) + result = super + if result > 0 && !@yielded + @yielded = true + transfer + end + result + end + end + def test_read omit unless defined?(UNIXSocket) @@ -275,4 +291,50 @@ def test_close_while_reading_on_thread reading_thread&.join rescue nil end end + + def test_interrupted_flush_does_not_replay_buffered_bytes + assert_no_replayed_buffer { |io| io.flush } + end + + def test_interrupted_write_does_not_replay_buffered_bytes + assert_no_replayed_buffer { |io| io.write("x" * 16_384) } + end + + def test_interrupted_writev_does_not_replay_buffered_bytes + assert_no_replayed_buffer { |io| io.write("x" * 16_384, "y") } + end + + private + + def assert_no_replayed_buffer + omit("UNIXSocket is not defined!") unless defined?(UNIXSocket) + + UNIXSocket.pair do |reader, writer| + writer.sync = false + thread = Thread.new do + Fiber.set_scheduler(YieldAfterWriteScheduler.new) + fiber = Fiber.schedule do + begin + writer.write(MESSAGE) + yield(writer) + ensure + writer.close + end + end + assert_equal(MESSAGE, reader.read_nonblock(MESSAGE.bytesize)) + assert_raise(Interrupt) { fiber.raise(Interrupt) } + ensure + Fiber.set_scheduler(nil) + end + + thread.value + assert_predicate(writer, :closed?) + assert_equal("", reader.read) + ensure + if thread&.alive? + thread.kill + thread.join + end + end + end end