diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index b70c2ee..f9d9f6b 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -4,11 +4,14 @@ on: [pull_request] jobs: validate-pr: - runs-on: ubuntu-22.04 + runs-on: ubuntu-latest + container: + image: ghcr.io/jangala-dev/mini-lua-image:bookworm + options: --user 0:0 steps: - name: Check out repository - uses: actions/checkout@v2 + uses: actions/checkout@v4 - name: Run custom install script run: | @@ -25,6 +28,6 @@ jobs: cd tests luajit test.lua - - name: Run Linter - run: luacheck . + # - name: Run Linter + # run: luacheck . diff --git a/.gitignore b/.gitignore index c17485e..4f7c513 100644 --- a/.gitignore +++ b/.gitignore @@ -1 +1,2 @@ *DS_Store +/scratch* diff --git a/src/fibers/io/exec.lua b/src/fibers/io/exec.lua index 098cd03..b48c45d 100644 --- a/src/fibers/io/exec.lua +++ b/src/fibers/io/exec.lua @@ -571,7 +571,7 @@ function Command:_on_scope_exit() for _, name in ipairs { 'stdin', 'stdout', 'stderr' } do local cfg = self['_' .. name] if cfg.stream and cfg.owned then - local ok, err = cfg.stream:close() + local ok, err = op.perform_raw(cfg.stream:close_op()) if not ok then error(err or ('failed to close ' .. name .. ' stream')) end diff --git a/src/fibers/io/fd_backend/nixio.lua b/src/fibers/io/fd_backend/nixio.lua index 7aa7c70..90a1f82 100644 --- a/src/fibers/io/fd_backend/nixio.lua +++ b/src/fibers/io/fd_backend/nixio.lua @@ -66,14 +66,14 @@ local function read_fd(fd, max) -- nixio.File:read / Socket:read both follow the same style: -- data (success/EOF) -- nil, msg, errno (error) - local data, msg, eno = fd:read(max) + local data, eno, msg = fd:read(max) - if data ~= nil then + if type(data) == 'string' then -- data may be "" at EOF; that is acceptable to callers. return data, nil end - eno = eno or nixio.errno() + -- eno = eno or nixio.errno() if eno == EAGAIN or eno == EWOULDBLOCK then -- Would block, signal “not ready yet”. @@ -100,13 +100,13 @@ local function write_fd(fd, str, len) -- For files: File.write(buf, offset, length) -- For sockets: Socket.send / write(buf, offset, length) – same shape. - local n, msg, eno = fd:write(str, 0, len) + local n, eno, msg = fd:write(str, 0, len) - if n ~= nil then + if type(n) == 'number' then return n, nil end - eno = eno or nixio.errno() + -- eno = eno or nixio.errno() if eno == EAGAIN or eno == EWOULDBLOCK then -- Would block. diff --git a/src/fibers/io/file.lua b/src/fibers/io/file.lua index 10e3ca0..b981325 100644 --- a/src/fibers/io/file.lua +++ b/src/fibers/io/file.lua @@ -69,7 +69,7 @@ end ---@return Stream local function fdopen(fd, flags_or_mode, filename) -- assert(type(fd) == "number", "fdopen: fd must be a number") - assert(type(fd) ~= nil, 'fdopen: fd must be non-nil') + assert(fd ~= nil, 'fdopen: fd must be non-nil') local readable, writable @@ -182,12 +182,8 @@ local function tmpfile(perms, tmpdir) ---@param newname string ---@return boolean|nil ok, string|nil err function f:rename(newname) - -- Flush buffered data first (various stream flavours). - if self.flush_output then - self:flush_output() - elseif self.flush then - self:flush() - end + -- Flush buffered data first. + self:flush() local real_fd = io.fileno and io:fileno() or fd if real_fd then diff --git a/src/fibers/io/poller/core.lua b/src/fibers/io/poller/core.lua index d9da138..fa38e57 100644 --- a/src/fibers/io/poller/core.lua +++ b/src/fibers/io/poller/core.lua @@ -64,7 +64,7 @@ end ---@return WaitToken function Poller:wait(fd, dir, task) -- assert(type(fd) == "number", "fd must be number") - assert(type(fd) ~= nil, 'fd must be non-nil') + assert(fd ~= nil, 'fd must be non-nil') assert(dir == 'rd' or dir == 'wr', "dir must be 'rd' or 'wr'") local ws = (dir == 'rd') and self.rd or self.wr diff --git a/src/fibers/io/poller/epoll.lua b/src/fibers/io/poller/epoll.lua index 28caa02..13400da 100644 --- a/src/fibers/io/poller/epoll.lua +++ b/src/fibers/io/poller/epoll.lua @@ -24,11 +24,12 @@ local C = ffi_c.C local ffi_tonumber = ffi_c.tonumber local get_errno = ffi_c.errno +local EPERM = 1 local EINTR = 4 local ENOENT = 2 local EBADF = 9 -local jit = rawget(_G, "jit") +local jit = rawget(_G, 'jit') local ARCH = ffi.arch or ((jit and jit.arch) or 'x64') ---------------------------------------------------------------------- @@ -82,31 +83,15 @@ local EPOLL_CTL_MOD = 3 local get_event, set_event, get_data, set_data if ARCH == 'x64' or ARCH == 'x86' then - get_event = function (ev) - return ffi.cast('uint32_t*', ev.raw)[0] - end - set_event = function (ev, value) - ffi.cast('uint32_t*', ev.raw)[0] = value - end - get_data = function (ev) - return ffi.cast('uint64_t*', ev.raw + 4)[0] - end - set_data = function (ev, value) - ffi.cast('uint64_t*', ev.raw + 4)[0] = value - end + get_event = function (ev) return ffi.cast('uint32_t*', ev.raw)[0] end + set_event = function (ev, value) ffi.cast('uint32_t*', ev.raw)[0] = value end + get_data = function (ev) return ffi.cast('uint64_t*', ev.raw + 4)[0] end + set_data = function (ev, value) ffi.cast('uint64_t*', ev.raw + 4)[0] = value end else - get_event = function (ev) - return ev.events - end - set_event = function (ev, value) - ev.events = value - end - get_data = function (ev) - return ev.data - end - set_data = function (ev, value) - ev.data = value - end + get_event = function (ev) return ev.events end + set_event = function (ev, value) ev.events = value end + get_data = function (ev) return ev.data end + set_data = function (ev, value) ev.data = value end end local function wrap_error(ret) @@ -150,7 +135,6 @@ local function epoll_wait(epfd, timeout_ms, max_events) if n == -1 then local errno = get_errno() if errno == EINTR then - -- Benign interruption: report “no events”. return {}, nil, errno end local err = ffi.string(C.strerror(errno)) @@ -163,7 +147,6 @@ local function epoll_wait(epfd, timeout_ms, max_events) local event = assert(ffi_tonumber(get_event(events[i]))) res[fd] = event end - return res, nil, nil end @@ -178,6 +161,7 @@ end ---@class EpollState ---@field epfd integer ---@field active_events table +---@field unpollable table -- fds that return EPERM to epoll_ctl ---@field maxevents integer local Epoll = {} Epoll.__index = Epoll @@ -188,6 +172,7 @@ local function new_epoll() local ret = { epfd = epoll_create(), active_events = {}, + unpollable = {}, maxevents = INITIAL_MAXEVENTS, } return setmetatable(ret, Epoll) @@ -197,51 +182,92 @@ local RD = EPOLLIN + EPOLLRDHUP local WR = EPOLLOUT local ERR = EPOLLERR + EPOLLHUP +local function die_ctl(opname, fd, err, errno) + error((opname .. ' failed for fd ' .. tostring(fd) .. ' (' .. tostring(err) .. ', errno ' .. tostring(errno) .. ')')) +end + function Epoll:add(fd, events) + -- Once an fd is known to be unpollable, do not try to epoll_ctl it again. + if self.unpollable[fd] then + return + end + local active = self.active_events[fd] or 0 local eventmask = bit.bor(events, active, EPOLLONESHOT) - local ok = epoll_ctl_mod(self.epfd, fd, eventmask) - if not ok then - assert(epoll_ctl_add(self.epfd, fd, eventmask)) + + -- Try MOD first (common case). + local ok, err, eno = epoll_ctl_mod(self.epfd, fd, eventmask) + if ok then + self.active_events[fd] = eventmask + return + end + + -- EPERM: fd type not supported by epoll (e.g. regular file). Treat as unpollable. + if eno == EPERM then + self.active_events[fd] = nil + self.unpollable[fd] = true + return end - self.active_events[fd] = eventmask + + -- Not currently registered (or MOD failed): try ADD. + local ok2, err2, eno2 = epoll_ctl_add(self.epfd, fd, eventmask) + if ok2 then + self.active_events[fd] = eventmask + return + end + + if eno2 == EPERM then + self.active_events[fd] = nil + self.unpollable[fd] = true + return + end + + die_ctl('epoll_ctl(ADD)', fd, err2 or err, eno2 or eno) end function Epoll:poll(timeout_ms) - local events, err, errno = epoll_wait(self.epfd, timeout_ms or 0, self.maxevents) - if not events then + local evmap, err, errno = epoll_wait(self.epfd, timeout_ms or 0, self.maxevents) + if not evmap then error(err or ('epoll_wait failed (errno ' .. tostring(errno) .. ')')) end local count = 0 - for fd, _ in pairs(events) do + for fd, _ in pairs(evmap) do count = count + 1 self.active_events[fd] = nil end - if count == self.maxevents then self.maxevents = self.maxevents * 2 end - return events + return evmap end function Epoll:del(fd) + -- If this fd was unpollable, there is no kernel state to delete. + if self.unpollable[fd] then + self.unpollable[fd] = nil + self.active_events[fd] = nil + return + end + local ok, err, errno = epoll_ctl_del(self.epfd, fd) if not ok then - -- ENOENT/EBADF: fd already closed or never registered; just clear. if errno == ENOENT or errno == EBADF then self.active_events[fd] = nil return end - error(err or ('epoll_ctl(DEL) failed (errno ' .. tostring(errno) .. ')')) + die_ctl('epoll_ctl(DEL)', fd, err, errno) end + self.active_events[fd] = nil end function Epoll:close() epoll_close(self.epfd) self.epfd = nil + self.active_events = {} + self.unpollable = {} end ---------------------------------------------------------------------- @@ -264,18 +290,42 @@ local function on_wait_change(ep, fd, want_rd, want_wr) end end -local function poll_backend(ep, timeout_ms, _, _) - -- ep:poll already returns fd -> epoll event bits. - local evmap = ep:poll(timeout_ms) +local function poll_backend(ep, timeout_ms, rd_waitset, wr_waitset) local events = {} + -- Synthesize readiness for fds that epoll cannot watch (EPERM). + -- This keeps Poller:wait and backend registration exception-free. + local had_synthetic = false + if ep.unpollable then + for fd, _ in pairs(ep.unpollable) do + local rd = rd_waitset and (not rd_waitset:is_empty(fd)) or false + local wr = wr_waitset and (not wr_waitset:is_empty(fd)) or false + if rd or wr then + had_synthetic = true + events[fd] = { rd = rd, wr = wr, err = false } + end + end + end + + -- If we have synthetic events to deliver, do not block in epoll_wait. + local real_timeout = had_synthetic and 0 or timeout_ms + + local evmap = ep:poll(real_timeout) for fd, ev in pairs(evmap) do local flags = { rd = bit.band(ev, RD + ERR) ~= 0, wr = bit.band(ev, WR + ERR) ~= 0, err = bit.band(ev, ERR) ~= 0, } - events[fd] = flags + local cur = events[fd] + if cur then + -- Merge with any synthetic readiness. + cur.rd = cur.rd or flags.rd + cur.wr = cur.wr or flags.wr + cur.err = cur.err or flags.err + else + events[fd] = flags + end end return events diff --git a/src/fibers/io/stream.lua b/src/fibers/io/stream.lua index 64d7076..5d44f9b 100644 --- a/src/fibers/io/stream.lua +++ b/src/fibers/io/stream.lua @@ -1,65 +1,210 @@ --- Use of this source code is governed by the Apache 2.0 license; see COPYING. - ----@module 'fibers.io.stream' +-- fibers/io/stream.lua local wait = require 'fibers.wait' local bytes = require 'fibers.utils.bytes' local op = require 'fibers.op' local perform = require 'fibers.performer'.perform +local runtime = require 'fibers.runtime' local RingBuf = bytes.RingBuf local LinearBuf = bytes.LinearBuf ---- Backend interface expected by Stream. ---@class StreamBackend ----@field read_string fun(self: StreamBackend, max: integer): string|nil, string|nil ----@field write_string fun(self: StreamBackend, data: string): integer|nil, string|nil ----@field on_readable fun(self: StreamBackend, task: Task): WaitToken ----@field on_writable fun(self: StreamBackend, task: Task): WaitToken ----@field close fun(self: StreamBackend): boolean, string|nil ----@field seek fun(self: StreamBackend, whence: string, offset: integer): integer|nil, string|nil ----@field nonblock fun(self: StreamBackend)|nil ----@field block fun(self: StreamBackend)|nil ----@field filename string|nil ----@field fileno fun(self: StreamBackend): integer|nil -- optional, used by file.tmpfile - ---- Buffered IO stream over a StreamBackend. +---@field read_string fun(self: StreamBackend, max: integer): string|nil, any|nil, any|nil +---@field write_string fun(self: StreamBackend, data: string): integer|nil, any|nil, any|nil +---@field on_readable fun(self: StreamBackend, task: Task): WaitToken +---@field on_writable fun(self: StreamBackend, task: Task): WaitToken +---@field close fun(self: StreamBackend): boolean, any|nil +---@field seek fun(self: StreamBackend, whence: string, offset: integer): integer|nil, any|nil +---@field nonblock fun(self: StreamBackend)|nil +---@field block fun(self: StreamBackend)|nil +---@field filename string|nil +---@field fileno fun(self: StreamBackend): integer|nil + ---@class Stream ---@field io StreamBackend|nil ----@field rx RingBuf|nil ----@field tx RingBuf|nil ----@field line_buffering boolean # flag only; behaviour is caller-defined ----@field flush_output fun(self: Stream)|nil ----@field flush fun(self: Stream)|nil ----@field rename fun(self: Stream, newname: string): boolean|nil, string|nil +---@field rx any|nil +---@field tx any|nil +---@field line_buffering boolean +---@field _bufmode '"no"'|'"line"'|'"full"'|nil +---@field _bufsize integer|nil +---@field _ws Waitset +---@field _closed boolean +---@field _closing boolean +---@field _sticky_rerr any|nil +---@field _sticky_werr any|nil +---@field _big string|nil +---@field _big_off integer +---@field _pump_task Task +---@field _pump_token WaitToken|nil +---@field _pump_scheduled boolean +---@field _close_done boolean +---@field _close_ok boolean|nil +---@field _close_err any|nil +---@field _rd_owner any|nil +---@field _wr_owner any|nil local Stream = {} Stream.__index = Stream local DEFAULT_BUFFER_SIZE = 2 ^ 12 +local BIG_WRITE_CHUNK = 64 * 1024 + +-- Single internal wait key. Everything that could unblock someone notifies K_STATE. +local K_STATE = 'state' +local WANT_STATE = K_STATE + +---------------------------------------------------------------------- +-- Small helpers +---------------------------------------------------------------------- + +local function sched() return runtime.current_scheduler end + +-- Lifecycle predicates. +function Stream:_has_backend() return self.io ~= nil end + +-- No backend, or already fully terminated. +function Stream:_is_dead() return self._closed or (self.io == nil) end + +-- In the close handshake, but not yet torn down. +function Stream:_is_closing() return self._closing and not self._closed end + +function Stream:_is_readable() return self.rx ~= nil end + +function Stream:_is_writable() return self.tx ~= nil end + +local function token2(t1, t2) + return { + unlink = function () + if t1 and t1.unlink then t1:unlink() end + if t2 and t2.unlink then t2:unlink() end + return false + end, + } +end + +local NO_TOKEN = { unlink = function () return false end } + +function Stream:_signal_state() + self._ws:notify_all(K_STATE, sched()) +end + +local function drained_tx(self) + return (not self._big) and self.tx and (self.tx:read_avail() == 0) +end + +---------------------------------------------------------------------- +-- Lane serialisation (read/write) +---------------------------------------------------------------------- + +local function new_lane(stream, field) + local owner = {} + + local function acquire() + local cur = stream[field] + if cur == nil or cur == owner then + stream[field] = owner + return true + end + return false + end + + local function release() + if stream[field] == owner then + stream[field] = nil + stream:_signal_state() + if stream._closing and not stream._close_done then + stream:_finish_close_if_ready() + end + end + end + + -- Common wrapper; probe releases on would-block, run holds the lane. + local function wrap(step, release_on_fail) + return function () + if not acquire() then return false, WANT_STATE end + + local ok, v = step() + if ok then return true, v end + + if release_on_fail then release() end + + return false, v or WANT_STATE + end + end + + local function wrap_probe(step) return wrap(step, true) end + + local function wrap_run(step) return wrap(step, false) end + + return { release = release, wrap_probe = wrap_probe, wrap_run = wrap_run } +end + +---------------------------------------------------------------------- +-- Backend wait registration +---------------------------------------------------------------------- + +local function make_register(self, opts) + opts = opts or {} + local primed = false + + return function (task, waker, want) + -- Always register internal state, so close/pump changes wake everyone. + local t_state = self._ws:add(K_STATE, task) + + if opts.prime_once and not primed then + primed = true + waker:wakeup(task) + end + + -- Internal waits (or unspecified wants) just wait on state changes. + if want == K_STATE or want == nil or not (want == 'rd' or want == 'wr') then + return token2(t_state, NO_TOKEN) + end + + local io = self.io + if not io then + -- Ensure the task runs again and the step observes closure. + waker:wakeup(task) + return token2(t_state, NO_TOKEN) + end + + local tok + if want == 'wr' and io.on_writable then + tok = io:on_writable(task) + else + tok = io:on_readable(task) + end + + return token2(t_state, tok) + end +end + +---------------------------------------------------------------------- +-- Construction +---------------------------------------------------------------------- ---- Open a new Stream over a backend. ---@param io_backend StreamBackend ----@param readable? boolean # default true ----@param writable? boolean # default true ----@param bufsize? integer # per-direction buffer size +---@param readable? boolean +---@param writable? boolean +---@param bufsize? integer ---@return Stream local function open(io_backend, readable, writable, bufsize) + bufsize = bufsize or DEFAULT_BUFFER_SIZE + local s = setmetatable({ io = io_backend, line_buffering = false, + _ws = wait.new_waitset(), + _big_off = 0, }, Stream) - if readable ~= false then - s.rx = RingBuf.new(bufsize or DEFAULT_BUFFER_SIZE) - end - if writable ~= false then - s.tx = RingBuf.new(bufsize or DEFAULT_BUFFER_SIZE) - end + if readable ~= false then s.rx = RingBuf.new(bufsize) end + if writable ~= false then s.tx = RingBuf.new(bufsize) end + s._pump_task = { run = function () s:_pump() end } return s end ---- Check whether a value is a Stream instance. ---@param x any ---@return boolean local function is_stream(x) @@ -67,136 +212,286 @@ local function is_stream(x) end function Stream:nonblock() - if self.io and self.io.nonblock then - self.io:nonblock() - end + if self.io and self.io.nonblock then self.io:nonblock() end end function Stream:block() - if self.io and self.io.block then - self.io:block() + if self.io and self.io.block then self.io:block() end +end + +---------------------------------------------------------------------- +-- Close / terminate +---------------------------------------------------------------------- + +function Stream:_latch_close(ok, err) + if self._close_done then return end + self._close_done = true + self._close_ok = ok + self._close_err = err + self:_signal_state() +end + +function Stream:_unlink_pump_wait() + local pt = self._pump_token + self._pump_token = nil + if pt and pt.unlink then pt:unlink() end +end + +function Stream:terminate(_) + -- Idempotent. + if self._closed then + if not self._close_done then + if self._sticky_werr ~= nil then + self:_latch_close(nil, self._sticky_werr) + else + self:_latch_close(true, nil) + end + end + self:_signal_state() + return end + + self._closed = true + self._closing = true + self._rd_owner = nil + self._wr_owner = nil + + self:_unlink_pump_wait() + + local io = self.io + self.io = nil + + self.rx, self.tx = nil, nil + self._big, self._big_off = nil, 0 + + if io and io.close then + pcall(function () io:close() end) + end + + if not self._close_done then + if self._sticky_werr ~= nil then + self:_latch_close(nil, self._sticky_werr) + else + self:_latch_close(true, nil) + end + end + + self:_signal_state() +end + +function Stream:_finish_close_if_ready() + if self._close_done or self._closed then return end + if not self._closing then return end + + -- If writable, wait for drain or sticky write error. + if self.tx then + if self._sticky_werr ~= nil then + self:_latch_close(nil, self._sticky_werr) + self:terminate('closed') + return + end + if not drained_tx(self) then + return + end + end + + -- Do not tear down while a lane is owned. + if self._rd_owner ~= nil or self._wr_owner ~= nil then + return + end + + self:_latch_close(true, nil) + self:terminate('closed') +end + +function Stream:_begin_close(_) + if self._closed then + self:_signal_state() + self:_finish_close_if_ready() + return + end + self._closing = true + self:_signal_state() + self:_finish_close_if_ready() +end + +---@return Op +function Stream:close_op() + local register = make_register(self, { prime_once = true }) + + local function probe_step() + if self._close_done then + return true, function () return self._close_ok, self._close_err end + end + return false, WANT_STATE + end + + local function run_step() + if not self._closing and not self._closed then + self:_begin_close('closing') + end + + -- Ensure pending output makes progress if any. + if self.tx then self:_kick_pump() end + + self:_finish_close_if_ready() + + if self._close_done then + return true, function () return self._close_ok, self._close_err end + end + return false, WANT_STATE + end + + return wait.waitable2(register, probe_step, run_step) + :wrap(function (th) + local ok, err = th() + if ok then return true, nil end + return nil, err + end) end ---------------------------------------------------------------------- --- Internal step machines +-- Read path (choice-safe; ALWAYS returns thunks on ready) ---------------------------------------------------------------------- ---@param stream Stream ----@param buf LinearBuf +---@param buf any ---@param min integer ---@param max integer ---@param terminator string|nil ----@return fun(): boolean, ... # step() -local function make_read_step(stream, buf, min, max, terminator) +---@return fun(): boolean, any +---@return fun(): boolean, any +local function make_read_steps(stream, buf, min, max, terminator) local tally = 0 + local term_target = nil + local want_hint = WANT_STATE + + local function term_enabled() + return terminator ~= nil and terminator ~= '' + end - local function adjust_for_terminator() - if not terminator then return end + local function maybe_clamp() + if term_target or not term_enabled() or not stream.rx then return end local loc = stream.rx:find(terminator) - if loc then - local final = tally + loc + #terminator + if not loc then return end + local final = tally + loc + #terminator + if final <= max then + term_target = final min, max = final, final end end - return function () - while true do - if not stream.rx or not stream.io then - return true, buf, tally, 'stream closed' + local function drain_once() + if not stream.rx then return end + local avail = stream.rx:read_avail() + if avail <= 0 or tally >= max then return end + local need = math.min(avail, max - tally) + if need <= 0 then return end + local chunk = stream.rx:take(need) + if chunk and #chunk > 0 then + buf:append(chunk) + tally = tally + #chunk + end + end + + local function drain_all() + if not stream.rx then return end + while tally < max do + local before = tally + drain_once() + if tally == before then break end + end + end + + -- Terminal check used for both probe and run; does not perform backend I/O. + -- Always returns either nil (not terminal) or a thunk. + local function terminal_thunk() + if stream._sticky_rerr ~= nil then + maybe_clamp() + return function () + drain_all() + return buf, tally, stream._sticky_rerr end + end + if stream:_is_dead() or stream:_is_closing() then + return function () return buf, tally, 'closed' end + end + return nil + end + + local function probe_step() + local th = terminal_thunk() + if th then return true, th end - adjust_for_terminator() - - local avail = stream.rx:read_avail() - if avail > 0 and tally < max then - local need = math.min(avail, max - tally) - local chunk = stream.rx:take(need) - if #chunk > 0 then - buf:append(chunk) - tally = tally + #chunk - if tally >= min then - return true, buf, tally + maybe_clamp() + if tally >= min then + return true, function () return buf, tally, nil end + end + + -- Choice-safe: if rx has enough to satisfy min, return a drain thunk. + local rx = stream.rx + if rx and tally < max then + local avail = rx:read_avail() + if avail > 0 then + local possible = tally + math.min(avail, max - tally) + if possible >= min then + return true, function () + drain_once() + return buf, tally, nil end end end + end + + return false, want_hint + end - if not (stream.io and stream.io.read_string) then - return true, buf, tally, 'backend does not support read_string' + local function run_step() + while true do + local th = terminal_thunk() + if th then return true, th end + + maybe_clamp() + drain_once() + if tally >= min then + return true, function () return buf, tally, nil end + end + + local io = stream.io + if not (io and io.read_string) then + return true, function () return buf, tally, 'backend does not support read_string' end end local room = stream.rx:write_avail() if room <= 0 then - if tally >= min then - return true, buf, tally - end - return true, buf, tally, 'buffer capacity exhausted' + return true, function () return buf, tally, 'buffer capacity exhausted' end end - local data, err, want = stream.io:read_string(room) - if err then - return true, buf, tally, err + local data, err, want = io:read_string(room) + if err ~= nil then + stream._sticky_rerr = err + stream:_signal_state() + return true, function () return buf, tally, err end end - if not data then - if tally >= min then - return true, buf, tally - end - return false, want + + if data == nil then + want_hint = want or 'rd' + return false, want_hint end - if #data == 0 then - return true, buf, tally + + if data == '' then + -- EOF + return true, function () return buf, tally, nil end end stream.rx:put(data) end end -end - ----@param stream Stream ----@param src_str string ----@return fun(): boolean, ... # step() -local function make_write_step(stream, src_str) - local offset = 0 - local len = #src_str - - return function () - if not stream.io then - return true, offset, 'stream closed' - end - - if offset == len then - return true, len - end - - if not (stream.io and stream.io.write_string) then - return true, offset, 'backend does not support write_string' - end - - local chunk = src_str:sub(offset + 1) - local n, err, want = stream.io:write_string(chunk) - if err then - return true, offset, err - end - if n == nil then - return false, want - end - if n == 0 then - return true, offset - end - offset = offset + n - if offset >= len then - return true, offset - end - return false - end + return probe_step, run_step end ----------------------------------------------------------------------- --- Core stream primitives ----------------------------------------------------------------------- - ----@param buf LinearBuf +---@param buf any ---@param opts? { min?: integer, max?: integer, terminator?: string, eof_ok?: boolean } ---@return Op function Stream:read_into_op(buf, opts) @@ -208,44 +503,34 @@ function Stream:read_into_op(buf, opts) local terminator = opts.terminator local eof_ok = not not opts.eof_ok - local step = make_read_step(self, buf, min, max, terminator) + local lane = new_lane(self, '_rd_owner') + local probe_step, run_step = make_read_steps(self, buf, min, max, terminator) + + probe_step = lane.wrap_probe(probe_step) + run_step = lane.wrap_run(run_step) + + local register = make_register(self, { prime_once = true }) + + local function wrap(th) + local ret_buf, cnt, err = th() + lane.release() - local function wrap(ret_buf, cnt, err) if cnt == 0 and not eof_ok then - return nil, cnt, err + return nil, 0, err end return ret_buf, cnt, err end - return wait.waitable( - function (task, suspension, _, want) - local io = self.io - if not io then - -- ensure the task runs again and the step observes closure - suspension.sched:schedule(task) - return { unlink = function () end } - end - - if want == 'wr' then - return io:on_writable(task) - end - return io:on_readable(task) - end, - step, - wrap - ) + local ev = wait.waitable2(register, probe_step, run_step, wrap) + return ev:on_abort(function () lane.release() end) end ----@param opts? { min?: integer, max?: integer, terminator?: string, eof_ok?: boolean } ----@return Op -function Stream:read_string_op(opts) +function Stream:core_read_op(opts) local buf = LinearBuf.new() local ev = self:read_into_op(buf, opts) return ev:wrap(function (ret_buf, cnt, err) - if not ret_buf then - return nil, cnt, err - end + if not ret_buf then return nil, 0, err end local s = ret_buf:tostring() if cnt == 0 and s == '' then return nil, 0, err @@ -254,68 +539,46 @@ function Stream:read_string_op(opts) end) end ----@param str string ----@return Op -function Stream:write_string_op(str) - assert(self.tx, 'stream is not writable') - assert(type(str) == 'string', 'write_string_op expects a string') +function Stream:read_some_op(max) + assert(type(max) == 'number' and max >= 0, 'read_some_op: max must be non-negative') + if max == 0 then return op.always('', nil) end - local step = make_write_step(self, str) - - local function wrap(bytes_written, err) - return bytes_written, err - end - - return wait.waitable( - function (task, suspension, _, want) - local io = self.io - if not io then - -- ensure the task runs again and the step observes closure - suspension.sched:schedule(task) - return { unlink = function () end } - end - - if want == 'rd' then - return io:on_readable(task) - end - return io:on_writable(task) - end, - step, - wrap - ) + return self:core_read_op { min = 1, max = max, eof_ok = true } + :wrap(function (s, cnt, err) + if err ~= nil then return nil, err end + if not s or cnt == 0 then return nil, nil end + return s, nil + end) end ----@return Op -function Stream:flush_output_op() - -- Unbuffered write path: there is nothing to flush at the Stream level. - -- Writes only return once the backend has accepted the data (or errored). - return op.always(0, nil) -end +function Stream:read_exactly_op(n) + assert(type(n) == 'number' and n >= 0, 'read_exactly_op: n must be non-negative') + if n == 0 then return op.always('', nil) end ----------------------------------------------------------------------- --- Derived per-stream ops ----------------------------------------------------------------------- + return self:core_read_op { min = n, max = n, eof_ok = false } + :wrap(function (s, cnt, err) + if err ~= nil then return nil, err end + if not s or cnt ~= n then return nil, 'short read' end + return s, nil + end) +end ----@param opts? { terminator?: string, keep_terminator?: boolean, max?: integer } ----@return Op -- when performed: line:string|nil, err:string|nil function Stream:read_line_op(opts) assert(self.rx, 'stream is not readable') opts = opts or {} local term = opts.terminator or '\n' local keep_term = not not opts.keep_terminator - local max_bytes = opts.max or math.huge - local ev = self:read_string_op { - min = max_bytes, - max = max_bytes, + local ev = self:core_read_op { + min = math.huge, + max = math.huge, terminator = term, eof_ok = true, } return ev:wrap(function (s, cnt, err) - if err then return nil, err end - + if err ~= nil then return nil, err end if not s or cnt == 0 then return nil, nil end if not keep_term and #term > 0 and s:sub(- #term) == term then @@ -326,71 +589,312 @@ function Stream:read_line_op(opts) end) end ----@param n integer ----@return Op -- when performed: data:string|nil, err:string|nil -function Stream:read_exactly_op(n) - assert(type(n) == 'number' and n >= 0, 'read_exactly_op: n must be non-negative') +function Stream:read_all_op() + assert(self.rx, 'stream is not readable') - return self:read_string_op { - min = n, - max = n, - eof_ok = false, - }:wrap(function (s, cnt, err) - if err then return nil, err end + return self:core_read_op { min = math.huge, max = math.huge, eof_ok = true } + :wrap(function (s, _, err) + if not s then return '', err end + return s, err + end) +end - if not s or cnt ~= n then return nil, 'short read' end +---------------------------------------------------------------------- +-- Buffered write pump +---------------------------------------------------------------------- - return s, nil - end) +function Stream:_kick_pump() + if self._pump_scheduled then return end + if self:_is_dead() then return end + self._pump_scheduled = true + sched():schedule(self._pump_task) end ----@return Op -- when performed: data:string, err:string|nil -function Stream:read_all_op() - assert(self.rx, 'stream is not readable') +local function next_write_chunk(self) + if self._big then + if self._big_off >= #self._big then + self._big = nil + self._big_off = 0 + self:_signal_state() + return nil + end + local remaining = #self._big - self._big_off + local take = remaining + if take > BIG_WRITE_CHUNK then take = BIG_WRITE_CHUNK end + return self._big:sub(self._big_off + 1, self._big_off + take), 'big' + end - -- Read until EOF or error in a single op. - local ev = self:read_string_op { - min = math.huge, - max = math.huge, - eof_ok = true, - } + if self.tx and self.tx:read_avail() > 0 then + local avail = self.tx:read_avail() + if avail > BIG_WRITE_CHUNK then avail = BIG_WRITE_CHUNK end + return self.tx:peek(avail), 'ring' + end - return ev:wrap(function (s, _, err) - -- read_string_op returns: - -- s == nil, cnt == 0 : no data at all (EOF or error-before-data) - -- s ~= nil, cnt > 0 : some data read, possibly with err - if not s then return '', err end -- Normalise “no data” to empty string. + return nil +end - return s, err - end) +local function advance_after_write(self, mode, n) + if mode == 'big' then + self._big_off = self._big_off + n + if self._big_off >= #self._big then + self._big = nil + self._big_off = 0 + end + self:_signal_state() + return + end + + self.tx:advance_read(n) + self:_signal_state() +end + +function Stream:_pump() + self._pump_scheduled = false + + local io = self.io + if self:_is_dead() or not io then return false end + if self._sticky_werr then return false end + if not (self.tx or self._big) then return false end + + self:_unlink_pump_wait() + + local progressed = false + + while true do + if self._sticky_werr or self:_is_dead() then break end + + local chunk, mode = next_write_chunk(self) + if not chunk or #chunk == 0 then break end + + local n, err, want = io:write_string(chunk) + if err then + self._sticky_werr = err + self:_signal_state() + break + end + + if n == nil or n == 0 then + -- Would block: arm readiness (poller is responsible for any EPERM cases). + local w = (want == 'rd') and 'rd' or 'wr' + if w == 'rd' and io.on_readable then + self._pump_token = io:on_readable(self._pump_task) + else + self._pump_token = io:on_writable(self._pump_task) + end + break + end + + progressed = true + advance_after_write(self, mode, n) + end + + if drained_tx(self) then self:_signal_state() end + + self:_finish_close_if_ready() + return progressed end ---------------------------------------------------------------------- --- Misc and lifecycle +-- Buffered write ops ---------------------------------------------------------------------- -function Stream:flush_input() - if self.rx then - self.rx:reset() +-- Shared output-lane op builder. +-- kind: +-- * 'write' : publish bytes (buffer/big) and return (n|nil, err|nil) +-- * 'flush' : wait until outbound is drained and return (true|nil, err|nil) +local function output_lane_op(self, kind, str) + assert(kind == 'write' or kind == 'flush', 'output_lane_op: bad kind') + + local lane = new_lane(self, '_wr_owner') + local register = make_register(self) + + local function pending() + return (self._big ~= nil) or (self.tx and self.tx:read_avail() > 0) + end + + local function drained() return not pending() end + + -- Decide whether we can accept `str` into the outbound queue. + -- Terminal/error cases are checked by step() before calling this. + local function can_accept(len) + local mode = self._bufmode or 'full' + + if mode == 'no' then + -- Do not allow queuing; require fully drained output. + if pending() then return false end + return true, 'big' + end + + -- Existing buffered behaviour. + if self._big then return false end + + local cap = self.tx:capacity() + if len <= self.tx:write_avail() then return true, 'ring' end + if self.tx:read_avail() == 0 and len > cap then return true, 'big' end + + return false + end + + local function publish(mode, s) + if mode == 'ring' then + self.tx:put(s) + else + self._big = s + self._big_off = 0 + end + end + + local function rollback_published(mode) + -- Only used in the idle-fast-path failure case. + -- Safe because was_idle implies there was no prior pending output. + if mode == 'ring' then + if self.tx then self.tx:reset() end + else + self._big, self._big_off = nil, 0 + end + end + + local function step(is_probe) + -- Sticky backend write error always wins. + if self._sticky_werr ~= nil then + local e = self._sticky_werr + return true, function () return nil, e end + end + + if kind == 'write' then + -- Writes do not proceed once closing/closed or backend absent. + if self:_is_dead() or self:_is_closing() then + return true, function () return nil, 'closed' end + end + + local len = #str + local ok, mode = can_accept(len) + if not ok then + if not is_probe then self:_kick_pump() end + return false, WANT_STATE + end + + local was_idle = drained() + + return true, function () + publish(mode, str) + + -- Opportunistic progress when previously idle; surfaces peer-close promptly. + local progressed = false + if was_idle then + progressed = self:_pump() or false + end + + -- If the very first attempt discovers a terminal error before any progress, + -- fail this write (and drop the just-published bytes). + if was_idle and not progressed and self._sticky_werr ~= nil then + local e = self._sticky_werr + rollback_published(mode) + self:_signal_state() + return nil, e + end + + if pending() and not self._pump_token then + self:_kick_pump() + end + + self:_signal_state() + return len, nil + end + end + + -- kind == 'flush' + if self:_is_dead() then + if drained() then + return true, function () return true, nil end + end + return true, function () return nil, 'closed' end + end + + if drained() then + return true, function () return true, nil end + end + + if is_probe then return false, WANT_STATE end + + self:_kick_pump() + return false, WANT_STATE + end + + local function probe_step() return step(true) end + local function run_step() return step(false) end + + probe_step = lane.wrap_probe(probe_step) + run_step = lane.wrap_run(run_step) + + local function wrap(th) + local a, b = th() + lane.release() + return a, b + end + + local ev = wait.waitable2(register, probe_step, run_step, wrap) + return ev:on_abort(function () lane.release() end) +end + +function Stream:core_write_op(str) + assert(self.tx, 'stream is not writable') + assert(type(str) == 'string', 'core_write_op expects a string') + if str == '' then return op.always(0, nil) end + return output_lane_op(self, 'write', str) +end + +local function flush_required_for_write(self, str) + local mode = self._bufmode or 'full' + if mode == 'no' then + return true + end + if mode == 'line' and type(str) == 'string' then + return str:find('\n', 1, true) ~= nil end + return false end ----@return boolean ok, string|nil err -function Stream:close() - local ok, err - if self.io and self.io.close then - ok, err = self.io:close() - else - ok, err = true, nil +function Stream:write_op(...) + assert(self.tx, 'stream is not writable') + + local count = select('#', ...) + if count == 0 then return op.always(0, nil) end + + local parts = {} + for i = 1, count do + local v = select(i, ...) + parts[i] = (type(v) == 'string') and v or tostring(v) end - self.rx, self.tx, self.io = nil, nil, nil - return ok, err + local str = table.concat(parts) + return self:core_write_op(str):wrap(function (n, err) + if n == nil then return nil, err end + + if flush_required_for_write(self, str) then + local ok, ferr = perform(self:flush_op()) + if ok == nil then return nil, ferr end + end + + return n, nil + end) +end + +function Stream:flush_op() + if not self.tx then return op.always(true, nil) end + return output_lane_op(self, 'flush') +end + +---------------------------------------------------------------------- +-- Misc +---------------------------------------------------------------------- + +function Stream:flush_input() + if self.rx then self.rx:reset() end + self:_signal_state() end ----@param whence? string ----@param offset? integer ----@return integer|nil pos, string|nil err function Stream:seek(whence, offset) + self:flush() if not (self.io and self.io.seek) then return nil, 'stream is not seekable' end @@ -399,23 +903,38 @@ function Stream:seek(whence, offset) return self.io:seek(whence, offset) end ----@param mode '"no"'|'"line"'|'"full"' ----@param _ any ----@return Stream -function Stream:setvbuf(mode, _) - if mode == 'no' then - self.line_buffering = false - elseif mode == 'line' then - self.line_buffering = true - elseif mode == 'full' then - self.line_buffering = false - else +local function next_pow2(n) + if n <= 1 then return 1 end + local p = 1 + while p < n do p = p * 2 end + return p +end + +function Stream:setvbuf(mode, size) + if mode ~= 'no' and mode ~= 'line' and mode ~= 'full' then error('bad mode: ' .. tostring(mode)) end + + self._bufmode = mode + self.line_buffering = (mode == 'line') + + if size ~= nil then + assert(type(size) == 'number' and size > 0, 'setvbuf: size must be positive') + size = next_pow2(math.floor(size)) + self._bufsize = size + + if self.rx and self.rx:read_avail() == 0 then + self.rx = RingBuf.new(size) + end + if self.tx and (not self._big) and self.tx:read_avail() == 0 then + self.tx = RingBuf.new(size) + end + self:_signal_state() + end + return self end ----@return string|nil function Stream:filename() return self.io and self.io.filename end @@ -424,101 +943,50 @@ end -- Synchronous convenience wrappers ---------------------------------------------------------------------- -function Stream:read_string(opts) - return perform(self:read_string_op(opts)) -end +function Stream:read_line(opts) return perform(self:read_line_op(opts)) end -function Stream:read_all() - return perform(self:read_all_op()) -end +function Stream:read_exactly(n) return perform(self:read_exactly_op(n)) end -function Stream:read_exactly(n) - return perform(self:read_exactly_op(n)) -end +function Stream:read_some(max) return perform(self:read_some_op(max)) end -function Stream:write_string(str) - return perform(self:write_string_op(str)) -end +function Stream:read_all() return perform(self:read_all_op()) end -function Stream:flush_output() - return perform(self:flush_output_op()) -end +function Stream:write(...) return perform(self:write_op(...)) end -function Stream:flush() - return self:flush_output() -end +function Stream:flush() return perform(self:flush_op()) end + +function Stream:close() return perform(self:close_op()) end ---------------------------------------------------------------------- --- Lua compatibility surface +-- Lua io-like compatibility ---------------------------------------------------------------------- ----@param fmt? string|integer ----@return Op -- when performed: value|nil, err|string|nil function Stream:read_op(fmt) assert(self.rx, 'stream is not readable') - local t = type(fmt) - - -- Default / "*l": line without terminator if fmt == nil or fmt == '*l' then return self:read_line_op() end - - -- "*L": line with terminator if fmt == '*L' then return self:read_line_op { keep_terminator = true } end - - -- "*a": read all if fmt == '*a' then return self:read_all_op() end - -- numeric: read up to n bytes - if t == 'number' then - local n = fmt - assert(n >= 0, 'read_op: n must be non-negative') - - -- Lua: f:read(0) returns "" immediately - if n == 0 then return op.always('', nil) end + if type(fmt) == 'number' then + assert(fmt >= 0, 'read_op: n must be non-negative') + if fmt == 0 then return op.always('', nil) end - -- read up to n bytes; allow EOF - local ev = self:read_string_op { min = 1, max = n, eof_ok = true } - - return ev:wrap(function (s, cnt, err) - if err then return nil, err end - if not s or cnt == 0 then return nil, nil end -- EOF before any data - return s, nil - end) - else - error('read_op: invalid format ' .. tostring(fmt)) + return self:core_read_op { min = 1, max = fmt, eof_ok = true } + :wrap(function (s, cnt, err) + if err then return nil, err end + if not s or cnt == 0 then return nil, nil end + return s, nil + end) end + + error('read_op: invalid format ' .. tostring(fmt)) end function Stream:read(fmt) return perform(self:read_op(fmt)) end ----@param ... any ----@return Op -- when performed: bytes_written:integer, err:string|nil -function Stream:write_op(...) - assert(self.tx, 'stream is not writable') - - local n = select('#', ...) - if n == 0 then - -- Match the “no-op but succeed” flavour. - return op.always(0, nil) - end - - local parts = {} - for i = 1, n do - local v = select(i, ...) - -- Follow Lua’s io.write behaviour: tostring each argument. - parts[i] = (type(v) == 'string') and v or tostring(v) - end - - local s = table.concat(parts) - return self:write_string_op(s) -end - -function Stream:write(...) - return perform(self:write_op(...)) -end - ---------------------------------------------------------------------- -- Module-level helpers ---------------------------------------------------------------------- @@ -540,12 +1008,9 @@ local function merge_lines_op(named_streams, opts) return op.named_choice(arms) end ----------------------------------------------------------------------- --- Public API ----------------------------------------------------------------------- - return { open = open, is_stream = is_stream, merge_lines_op = merge_lines_op, + Stream = Stream, } diff --git a/src/fibers/mailbox.lua b/src/fibers/mailbox.lua index 1468a6d..aeb049b 100644 --- a/src/fibers/mailbox.lua +++ b/src/fibers/mailbox.lua @@ -39,6 +39,7 @@ local perform = require 'fibers.performer'.perform ---@field buf any|nil -- FIFO buffer when cap>0; nil for rendezvous ---@field getq any -- FIFO of waiting receivers ---@field putq any -- FIFO of waiting senders +---@field taskq any -- FIFO of task waiters for recv readiness ---@field closed boolean ---@field reason any|nil ---@field senders integer -- counted sender handles still open @@ -74,6 +75,29 @@ local function pop_active(q) end end +---@param st MailboxState +local function notify_task_waiters(st) + local q = st.taskq + if not q then return end + + while not q:empty() do + local e = q:pop() + if e and e.active then + e.active = false + e.waker:wakeup(e.task) + end + end +end + +---@param st MailboxState +---@return boolean +local function recv_may_succeed(st) + if st.closed then return true end + if st.buf and st.buf:length() > 0 then return true end + if st.putq and not st.putq:empty() then return true end + return false +end + ---@param st MailboxState ---@param reason any|nil local function record_reason(st, reason) @@ -114,6 +138,8 @@ local function close_state(st, reason) if not snd then break end snd.suspension:complete(snd.wrap, nil) end + + notify_task_waiters(st) end ---------------------------------------------------------------------- @@ -149,6 +175,7 @@ local function new(capacity, opts) buf = (capacity > 0) and fifo.new() or nil, getq = fifo.new(), putq = fifo.new(), + taskq = fifo.new(), closed = false, reason = nil, senders = 1, @@ -250,6 +277,7 @@ function Tx:send_op(v) -- (For cap==0, drop_oldest is normalised away to reject_newest.) buf:pop() buf:push(v) + notify_task_waiters(st) return true, true end @@ -267,12 +295,14 @@ function Tx:send_op(v) local recv = pop_active(getq) if recv then recv.suspension:complete(recv.wrap, v) + notify_task_waiters(st) return true, true end -- Buffered enqueue when there is space. if buf and buf:length() < cap then buf:push(v) + notify_task_waiters(st) return true, true end @@ -292,6 +322,35 @@ function Tx:send_op(v) return op.new_primitive(nil, try, block) end +--- Register a task to be woken when recv may succeed (message arrives or close). +--- This does not expose the scheduler; callers provide a waker capability. +---@param task Task +---@param waker table +---@return WaitToken +function Rx:on_message(task, waker) + local st = self._st + assert(task and type(task) == 'table' and type(task.run) == 'function', + 'on_message: task must have :run()') + assert(waker and type(waker.wakeup) == 'function', + 'on_message: waker must support :wakeup(task)') + + if recv_may_succeed(st) then + waker:wakeup(task) + return { unlink = function () return false end } + end + + local entry = { task = task, waker = waker, active = true } + st.taskq:push(entry) + + return { + unlink = function () + if not entry.active then return false end + entry.active = false + return false + end, + } +end + --- Synchronously send a message. ---@param v any ---@return boolean|nil ok diff --git a/src/fibers/op.lua b/src/fibers/op.lua index db5efa7..6733347 100644 --- a/src/fibers/op.lua +++ b/src/fibers/op.lua @@ -71,6 +71,29 @@ function Suspension:_run_cleanups() cs[i] = nil end end + +-- Waker capability (scheduler is an implementation detail) + +--- Wake a task to run “soon”. +---@param task Task +function Suspension:wakeup(task) + self.sched:schedule(task) +end + +--- Wake a task at an absolute time on the scheduler clock. +---@param t number +---@param task Task +function Suspension:at_time(t, task) + self.sched:schedule_at_time(t, task) +end + +--- Wake a task after a delay from the scheduler’s current time. +---@param dt number +---@param task Task +function Suspension:after(dt, task) + self.sched:schedule_after_sleep(dt, task) +end + --- Mark a suspension as complete and enqueue it on the scheduler. ---@param wrap WrapFn ---@param ... any diff --git a/src/fibers/sleep.lua b/src/fibers/sleep.lua index 4eb9fd4..5587b77 100644 --- a/src/fibers/sleep.lua +++ b/src/fibers/sleep.lua @@ -21,7 +21,7 @@ local function deadline_op(t) ---@param suspension Suspension ---@param wrap_fn WrapFn local function block(suspension, wrap_fn) - suspension.sched:schedule_at_time(t, suspension:complete_task(wrap_fn)) + suspension:at_time(t, suspension:complete_task(wrap_fn)) end return op.new_primitive(nil, try, block) diff --git a/src/fibers/utils/bytes/ffi.lua b/src/fibers/utils/bytes/ffi.lua index 2e09249..38a1b9e 100644 --- a/src/fibers/utils/bytes/ffi.lua +++ b/src/fibers/utils/bytes/ffi.lua @@ -125,6 +125,19 @@ function ring_mt:put(str) copy_in(self, tmp, n) end +-- Opaque mark of the current write position (for tail rollback). +-- Intended for "publish then possibly roll back" patterns in higher layers. +function ring_mt:mark_write() + return self.write_idx +end + +-- Rewind the write position to a previously obtained mark. +-- Caller must ensure no consumer progress happened since the mark. +function ring_mt:rewind_write(mark) + -- mark is expected to be the cdata returned by mark_write() + self.write_idx = mark +end + function ring_mt:take(n) assert(type(n) == 'number' and n >= 0, 'RingBuf:take expects non-negative count') local avail = self:read_avail() @@ -162,6 +175,46 @@ local function RingBuf_new(size) return ring_mt.init(self, size) end +function ring_mt:capacity() + return self.size +end + +function ring_mt:advance_read(n) + assert(type(n) == 'number' and n >= 0, 'RingBuf:advance_read expects non-negative count') + local avail = self:read_avail() + assert(n <= avail, 'RingBuf:advance_read out of range') + if n == 0 then return end + self.read_idx = self.read_idx + ffi.cast('uint32_t', n) +end + +function ring_mt:peek(n) + assert(type(n) == 'number' and n >= 0, 'RingBuf:peek expects non-negative count') + local avail = self:read_avail() + if avail == 0 or n == 0 then + return '' + end + if n > avail then + n = avail + end + + -- Like tostring() but only for n bytes, and without mutating read_idx. + local tmp = ffi.new('uint8_t[?]', n) + local size = self.size + local start = pos(self, self.read_idx) + local first = math.min(n, size - start) + + if first > 0 then + ffi.copy(tmp, self.buf + start, first) + end + + local rest = n - first + if rest > 0 then + ffi.copy(tmp + first, self.buf, rest) + end + + return ffi.string(tmp, n) +end + ---------------------------------------------------------------------- -- LinearBuf ---------------------------------------------------------------------- diff --git a/src/fibers/utils/bytes/lua.lua b/src/fibers/utils/bytes/lua.lua index af7b85e..612269f 100644 --- a/src/fibers/utils/bytes/lua.lua +++ b/src/fibers/utils/bytes/lua.lua @@ -176,6 +176,26 @@ function RingBuf_mt:put(str) self:write(str, n) end +-- Opaque mark of the current write position (for tail rollback). +-- We only need enough state to drop newly appended chunks and restore len. +function RingBuf_mt:mark_write() + return { n = #self.chunks, len = self.len } +end + +-- Rewind the write position to a previously obtained mark. +-- Caller must ensure no consumer progress happened since the mark. +function RingBuf_mt:rewind_write(mark) + assert(type(mark) == 'table', 'RingBuf:rewind_write expects mark table') + local n = mark.n or 0 + local len = mark.len or 0 + + for i = #self.chunks, n + 1, -1 do + self.chunks[i] = nil + end + + self.len = len +end + function RingBuf_mt:take(n) assert(type(n) == 'number' and n >= 0, 'RingBuf:take expects non-negative count') if self.len == 0 or n == 0 then @@ -202,6 +222,43 @@ function RingBuf_mt:find(pattern) return i and (i - 1) or nil end +function RingBuf_mt:capacity() + return self.size +end + +function RingBuf_mt:peek(n) + assert(type(n) == 'number' and n >= 0, 'RingBuf:peek expects non-negative count') + if n == 0 or self.len == 0 then + return '' + end + if n > self.len then + n = self.len + end + + -- Same as read(nil, n) but without advancing. + local out = {} + local need = n + local i = self.head_idx + local off = self.head_off + local last = #self.chunks + + while need > 0 and i <= last do + local chunk = self.chunks[i] + local rem = #chunk - off + local take = math.min(need, rem) + out[#out + 1] = chunk:sub(off + 1, off + take) + need = need - take + if take == rem then + i = i + 1 + off = 0 + else + off = off + take + end + end + + return table.concat(out) +end + ---------------------------------------------------------------------- -- LinearBuf ---------------------------------------------------------------------- diff --git a/src/fibers/wait.lua b/src/fibers/wait.lua index 08ee612..c3d1301 100644 --- a/src/fibers/wait.lua +++ b/src/fibers/wait.lua @@ -213,119 +213,142 @@ end ---------------------------------------------------------------------- -- waitable: (register, step, wrap_fn?) -> Op +-- waitable2: (register, probe_step, run_step, wrap_fn?) -> Op ---------------------------------------------------------------------- +-- Normalise "want" without restricting it to rd/wr/any. +-- * nil/false -> nil +-- * 'any' is treated specially by register_with_want +-- * everything else is passed through to register(...) local function normalise_want(want) - if want == 'rd' or want == 'wr' or want == 'any' then - return want - end - return nil + return (want == nil or want == false) and nil or want end ---- Build a waitable Op from a register function and step function. --- --- step() -> done:boolean, ... + +--- Build a waitable Op from a register function and two step functions. -- --- * done == true : the operation is ready to commit now; --- remaining values are the result. --- * done == false : not ready; the register() function must --- arrange a future call to task:run(). +-- probe_step() -> done:boolean, ... +-- * Must be non-blocking and must not yield. +-- * Should be side-effect neutral when returning done==false. +-- * May return (false, want) where want is any token understood by register(). -- --- register(task, suspension, leaf_wrap) -> token +-- run_step() -> done:boolean, ... +-- * Must be non-blocking and must not yield. +-- * May perform stateful progress (e.g. fill buffers, advance state machines). -- --- * Must arrange for task:run() to be invoked when progress may --- have been made (fd readable, space available, timer expired). --- * Returns a token table which may define token:unlink() to --- cancel any outstanding registration for this synchronisation. +-- register(task, waker, want) -> token +-- * Must arrange for task:run() when progress may be possible. +-- * want is passed through (except 'any', see below). +-- * token:unlink() (if present) is called on abort to cancel registration. -- --- wrap_fn (optional) is used as the primitive wrap for the Op. +-- waker capability: +-- * waker:wakeup(task) +-- * waker:at_time(t, task) +-- * waker:after(dt, task) -- --- Requirements on step and register: --- - Both must be non-blocking and must not yield. --- - Errors raised by step/register are not caught here; they are --- treated as bugs and surfaced by the surrounding scope/fiber. +-- Special want: +-- * want == 'any' registers both ('rd' and 'wr') and unlinks both on abort. -- --- The op participates fully in choice/with_nack/on_abort; if it loses --- a choice, any outstanding registration is cancelled via token:unlink(). ----@param register fun(task: Task, suspension: Suspension, leaf_wrap: WrapFn, want: any): WaitToken ----@param step fun(): boolean, ... +---@param register fun(task: Task, waker: table, want: any): WaitToken +---@param probe_step fun(): boolean, ... +---@param run_step fun(): boolean, ... ---@param wrap_fn? WrapFn ---@return Op -local function waitable(register, step, wrap_fn) - assert(type(register) == 'function', 'waitable: register must be a function') - assert(type(step) == 'function', 'waitable: step must be a function') +local function waitable2(register, probe_step, run_step, wrap_fn) + assert(type(register) == 'function', 'waitable2: register must be a function') + assert(type(probe_step) == 'function', 'waitable2: probe_step must be a function') + assert(type(run_step) == 'function', 'waitable2: run_step must be a function') wrap_fn = wrap_fn or id_wrap return op.guard(function () - local token - local last_want + local token, last_want, cleanup_added, waker - local function unlink_token() - if token and token.unlink then - token:unlink() - end + local function unlink() + local t = token token = nil + if t and t.unlink then t:unlink() end end - local function try() - local res = pack(step()) - if not res[1] then - last_want = normalise_want(res[2]) + + local function capture_want(step_fn) + local r = pack(step_fn()) + last_want = r[1] and nil or normalise_want(r[2]) + return r + end + + local function register_any(task, waker_) + local t1 = register(task, waker_, 'rd') + local t2 = register(task, waker_, 'wr') + return { + unlink = function () + if t1 and t1.unlink then t1:unlink() end + if t2 and t2.unlink then t2:unlink() end + return false + end, + } + end + + local function arm(task, suspension, leaf_wrap, want) + if not cleanup_added then + cleanup_added = true + suspension:add_cleanup(unlink) + end + + unlink() + + if want == 'any' then + token = register_any(task, waker) -- see note below else - last_want = nil + token = register(task, waker, want) end - return unpack(res, 1, res.n) + end + + local function try() + local r = capture_want(probe_step) + return unpack(r, 1, r.n) end local function block(suspension, leaf_wrap) - ---@class WaitTask : Task - local task + waker = { + wakeup = function (_, task_) suspension:wakeup(task_) end, + at_time = function (_, t, task_) suspension:at_time(t, task_) end, + after = function (_, dt, task_) suspension:after(dt, task_) end, + } - local function register_with_want(want) - unlink_token() - - if want == 'any' then - local t1 = register(task, suspension, leaf_wrap, 'rd') - local t2 = register(task, suspension, leaf_wrap, 'wr') - token = { - unlink = function () - if t1 and t1.unlink then t1:unlink() end - if t2 and t2.unlink then t2:unlink() end - end, - } - else - token = register(task, suspension, leaf_wrap, want) - end - end + local task task = { run = function () - if not suspension:waiting() then - return - end + if not suspension:waiting() then return end - local res = pack(step()) - local done = res[1] - if done then - unlink_token() - return suspension:complete(leaf_wrap, unpack(res, 2, res.n)) + local r = capture_want(run_step) + if r[1] then + unlink() + return suspension:complete(leaf_wrap, unpack(r, 2, r.n)) end - last_want = normalise_want(res[2]) - register_with_want(last_want) + arm(task, suspension, leaf_wrap, last_want) end, } - register_with_want(last_want) + -- Use want captured by the most recent try(). + arm(task, suspension, leaf_wrap, last_want) end - local prim = op.new_primitive(wrap_fn, try, block) - - return prim:on_abort(function () - unlink_token() - end) + return op.new_primitive(wrap_fn, try, block):on_abort(unlink) end) end + +--- Backwards-compatible wrapper: a single step is used for both probe and run. +---@param register fun(task: Task, waker: table, want: any): WaitToken +---@param step fun(): boolean, ... +---@param wrap_fn? WrapFn +---@return Op +local function waitable(register, step, wrap_fn) + return waitable2(register, step, step, wrap_fn) +end + return { new_waitset = new_waitset, waitable = waitable, + waitable2 = waitable2, } diff --git a/tests/test_io-file.lua b/tests/test_io-file.lua index 51b9804..9d01c26 100644 --- a/tests/test_io-file.lua +++ b/tests/test_io-file.lua @@ -43,7 +43,7 @@ local function test_tmpfile_roundtrip() assert(f, 'tmpfile() failed: ' .. tostring(err)) local msg = 'hello, tmpfile' - local n, werr = perform(f:write_string_op(msg)) + local n, werr = perform(f:write_op(msg)) assert(n == #msg, 'write_string_op wrote ' .. tostring(n) .. ' bytes, expected ' .. #msg) assert(werr == nil, 'write_string_op returned error: ' .. tostring(werr)) @@ -51,7 +51,7 @@ local function test_tmpfile_roundtrip() local pos, serr = f:seek('set', 0) assert(pos ~= nil, 'seek failed: ' .. tostring(serr)) - local s, cnt, rerr = perform(f:read_string_op { + local s, cnt, rerr = perform(f:core_read_op { min = #msg, max = #msg, eof_ok = true, @@ -74,11 +74,11 @@ local function test_pipe_roundtrip_and_eof() assert(r and w, 'pipe() did not return read and write streams') local msg = 'pipe-test' - local n, werr = perform(w:write_string_op(msg)) + local n, werr = perform(w:write_op(msg)) assert(n == #msg, 'pipe write_string_op wrote ' .. tostring(n) .. ' bytes, expected ' .. #msg) assert(werr == nil, 'pipe write_string_op returned error: ' .. tostring(werr)) - local s, cnt, rerr = perform(r:read_string_op { + local s, cnt, rerr = perform(r:core_read_op { min = #msg, max = #msg, eof_ok = true, @@ -92,7 +92,7 @@ local function test_pipe_roundtrip_and_eof() local okw, errw = w:close() assert(okw, 'pipe write stream close failed: ' .. tostring(errw)) - local s2, cnt2, rerr2 = perform(r:read_string_op { + local s2, cnt2, rerr2 = perform(r:core_read_op { min = 1, eof_ok = true, }) @@ -120,7 +120,7 @@ local function test_closed_stream_errors() -- Writing via an already-closed Stream should raise "stream is not writable". local ok, err = pcall(function () - return perform(w1:write_string_op('abc')) + return perform(w1:write_op('abc')) end) assert(not ok, 'expected write after close to fail with an assertion') assert(tostring(err):match('stream is not writable'), @@ -135,7 +135,7 @@ local function test_closed_stream_errors() -- Reading via an already-closed Stream should raise "stream is not readable". ok, err = pcall(function () - return perform(r2:read_string_op { + return perform(r2:core_read_op { min = 1, eof_ok = false, }) @@ -168,7 +168,7 @@ local function test_cancellation_cancels_blocked_read() end) -- Perform a read that will block (nothing is written). - local v1, v2, v3 = perform(r:read_string_op { + local v1, v2, v3 = perform(r:core_read_op { min = 1, eof_ok = true, }) diff --git a/tests/test_io-mem.lua b/tests/test_io-mem.lua index c9f5abc..abf5e79 100644 --- a/tests/test_io-mem.lua +++ b/tests/test_io-mem.lua @@ -49,14 +49,14 @@ local function test_simple_read_write() -- Writer: write once from A to B fibers.spawn(function () - local ev = a:write_string_op(payload) + local ev = a:write_op(payload) local n, err = fibers.perform(ev) assert_nil(err, 'simple write: unexpected error') assert_eq(n, #payload, 'simple write: wrong byte count') end) -- Reader: read exactly len(payload) bytes - local ev = b:read_string_op { + local ev = b:core_read_op { min = #payload, max = #payload, eof_ok = true, @@ -93,7 +93,7 @@ local function test_backpressure_and_partial() while total < #payload do -- Read at least 1 byte, at most 3 each time. - local ev = b:read_string_op { + local ev = b:core_read_op { min = 1, max = 3, eof_ok = true, @@ -117,7 +117,7 @@ local function test_backpressure_and_partial() end) -- Writer: write the full payload as one op - local ev = a:write_string_op(payload) + local ev = a:write_op(payload) local n, err = fibers.perform(ev) assert_nil(err, 'backpressure write: unexpected error') @@ -141,7 +141,7 @@ local function test_eof_behaviour() -- Write then close A's half. fibers.spawn(function () - local ev = a:write_string_op(payload) + local ev = a:write_op(payload) local n, err = fibers.perform(ev) assert_nil(err, 'EOF write: unexpected error') assert_eq(n, #payload, 'EOF write: wrong byte count') @@ -149,7 +149,7 @@ local function test_eof_behaviour() end) -- First read should get the payload. - local ev1 = b:read_string_op { + local ev1 = b:core_read_op { min = #payload, max = #payload, eof_ok = true, @@ -162,7 +162,7 @@ local function test_eof_behaviour() -- Second read should see EOF. For read_string_op: -- EOF with no data → (nil, 0, err|nil) - local ev2 = b:read_string_op { + local ev2 = b:core_read_op { min = 1, max = 16, eof_ok = true, @@ -190,7 +190,7 @@ local function test_line_terminator() local data = 'line1\nline2\n' fibers.spawn(function () - local ev = a:write_string_op(data) + local ev = a:write_op(data) local n, err = fibers.perform(ev) assert_nil(err, 'line write: unexpected error') assert_eq(n, #data, 'line write: wrong byte count') @@ -198,7 +198,7 @@ local function test_line_terminator() end) -- Read up to and including first "\n" - local ev1 = b:read_string_op { + local ev1 = b:core_read_op { min = 1, max = #data, terminator = '\n', @@ -211,7 +211,7 @@ local function test_line_terminator() assert_eq(cnt1, #s1, 'line read(1): wrong count') -- Read up to and including second "\n" - local ev2 = b:read_string_op { + local ev2 = b:core_read_op { min = 1, max = #data, terminator = '\n', @@ -224,7 +224,7 @@ local function test_line_terminator() assert_eq(cnt2, #s2, 'line read(2): wrong count') -- Third read should see EOF - local ev3 = b:read_string_op { + local ev3 = b:core_read_op { min = 1, max = 16, terminator = '\n', @@ -253,7 +253,7 @@ local function test_write_after_peer_close() b:close() -- Writing from A should report "closed" from the backend. - local ev = a:write_string_op('x') + local ev = a:core_write_op('x') local _, err = fibers.perform(ev) -- Depending on exact semantics, n may be 0 or nil; err should be "closed". diff --git a/tests/test_io-socket.lua b/tests/test_io-socket.lua index c8d2b52..672ca82 100644 --- a/tests/test_io-socket.lua +++ b/tests/test_io-socket.lua @@ -35,7 +35,7 @@ local function test_unix_socket_roundtrip(scope) local s, aerr = server:accept() assert(s, 'server accept failed: ' .. tostring(aerr)) - local msg, cnt, rerr = perform(s:read_string_op { + local msg, cnt, rerr = perform(s:core_read_op { min = 5, max = 5, eof_ok = true, @@ -45,7 +45,7 @@ local function test_unix_socket_roundtrip(scope) assert(cnt == 5, 'server read_string_op read ' .. tostring(cnt) .. ' bytes, expected 5') assert(msg == 'hello', ('server received %q, expected %q'):format(tostring(msg), 'hello')) - local n, werr = perform(s:write_string_op('world')) + local n, werr = perform(s:write_op('world')) assert(werr == nil, 'server write_string_op error: ' .. tostring(werr)) assert(n == 5, 'server write_string_op wrote ' .. tostring(n) .. ' bytes, expected 5') @@ -60,11 +60,11 @@ local function test_unix_socket_roundtrip(scope) local client, cerr = socket_mod.connect_unix(path) assert(client, 'connect_unix failed: ' .. tostring(cerr)) - local n, werr = perform(client:write_string_op('hello')) + local n, werr = perform(client:write_op('hello')) assert(werr == nil, 'client write_string_op error: ' .. tostring(werr)) assert(n == 5, 'client write_string_op wrote ' .. tostring(n) .. ' bytes, expected 5') - local resp, cnt, rerr = perform(client:read_string_op { + local resp, cnt, rerr = perform(client:core_read_op { min = 5, max = 5, eof_ok = true, diff --git a/tests/test_io-stream.lua b/tests/test_io-stream.lua index 1586abf..ef72697 100644 --- a/tests/test_io-stream.lua +++ b/tests/test_io-stream.lua @@ -1,25 +1,60 @@ -- tests/test_stream_mem.lua -- --- Synthetic tests for fibers.io.stream using an in-memory backend. +-- Synthetic tests for fibers.io.stream using in-memory backends. print('testing: fibers.io.stream') --- look one level up package.path = '../src/?.lua;' .. package.path -local fibers = require 'fibers' -local stream = require 'fibers.io.stream' -local wait = require 'fibers.wait' -local runtime = require 'fibers.runtime' -local sleep = require 'fibers.sleep' -local op = require 'fibers.op' -local perform = require 'fibers.performer'.perform +local fibers = require 'fibers' +local stream = require 'fibers.io.stream' +local wait = require 'fibers.wait' +local runtime = require 'fibers.runtime' +local sleep = require 'fibers.sleep' +local op = require 'fibers.op' +local waitgroup = require 'fibers.waitgroup' +local perform = require 'fibers.performer'.perform + +require 'fibers.scope'.set_debug(true) local function with_timeout(ev, timeout_s) -- op.boolean_choice returns: (won:boolean, ...results...) return perform(op.boolean_choice(ev, sleep.sleep_op(timeout_s))) end --- In-memory duplex backend with partial I/O and readiness notifications. +local function assert_eq(a, b, msg) + if a ~= b then + error((msg or 'assert_eq failed') .. (': got ' .. tostring(a) .. ', expected ' .. tostring(b)), 2) + end +end + +local function assert_truthy(v, msg) + if not v then error(msg or 'expected truthy') end +end + +local function assert_ok_or_zero(v, msg) + if v ~= true and v ~= 0 then + error((msg or 'expected true or 0') .. (': got ' .. tostring(v)), 2) + end +end + +local function assert_internal_ws_empty(s, msg) + local ws = s and s._ws + if not ws or not ws.buckets then return end + if next(ws.buckets) ~= nil then + local keys = {} + for k in pairs(ws.buckets) do keys[#keys + 1] = tostring(k) end + error((msg or 'internal waitset leaked entries') .. ': keys=' .. table.concat(keys, ','), 2) + end +end + +local function assert_closed_err(err) + assert_truthy(err == 'closed' or err == 'stream closed', 'expected close error, got ' .. tostring(err)) +end + +---------------------------------------------------------------------- +-- Backend 1: basic duplex, partial writes, only "rd" notifications +---------------------------------------------------------------------- + local function make_stream_pair() local shared = { buf = '', @@ -38,7 +73,6 @@ local function make_stream_pair() return nil, nil -- would block end max = max or 1 - -- Deliberately read at most 1 byte to exercise partial reads. local n = math.min(1, max, #self.shared.buf) local s = self.shared.buf:sub(1, n) self.shared.buf = self.shared.buf:sub(n + 1) @@ -81,10 +115,15 @@ local function make_stream_pair() end function rd_io:seek() return nil, 'not seekable' end + function wr_io:seek() return nil, 'not seekable' end + function rd_io:nonblock() end + function rd_io:block() end + function wr_io:nonblock() end + function wr_io:block() end local rd = stream.open(rd_io, true, false) @@ -92,15 +131,17 @@ local function make_stream_pair() return rd, wr, shared end --- Variant backend to assert 'want' propagation: --- rd_io:read_string returns want='wr' when empty; and only 'wr' waiters are notified. +---------------------------------------------------------------------- +-- Backend 2: "want='wr'" read wakeups to validate want propagation +---------------------------------------------------------------------- + local function make_stream_pair_want_wr() local shared = { - buf = '', - closed = false, - waitset = wait.new_waitset(), -- use key "wr" only for wakeups - rd_regs = 0, - wr_regs = 0, + buf = '', + closed = false, + waitset = wait.new_waitset(), -- use key "wr" only for wakeups + rd_regs = 0, + wr_regs = 0, } local rd_io = { shared = shared } @@ -111,7 +152,6 @@ local function make_stream_pair_want_wr() if self.shared.closed then return '', nil -- EOF end - -- Would block; request registration on writability. return nil, nil, 'wr' end max = max or 1 @@ -131,7 +171,6 @@ local function make_stream_pair_want_wr() local n = 1 local ch = str:sub(1, n) shared.buf = shared.buf .. ch - -- Only notify "wr". If Stream ignores want and waits on "rd", it will hang. shared.waitset:notify_all('wr', runtime.current_scheduler) return n, nil end @@ -166,10 +205,185 @@ local function make_stream_pair_want_wr() end function rd_io:seek() return nil, 'not seekable' end + + function wr_io:seek() return nil, 'not seekable' end + + function rd_io:nonblock() end + + function rd_io:block() end + + function wr_io:nonblock() end + + function wr_io:block() end + + local rd = stream.open(rd_io, true, false) + local wr = stream.open(wr_io, false, true) + return rd, wr, shared +end + +---------------------------------------------------------------------- +-- Backend 3: full duplex for buffered write/flush tests +---------------------------------------------------------------------- + +local function make_stream_pair_full() + local shared = { + wire = '', + closed = false, + waitset = wait.new_waitset(), -- keys: 'rd' + } + + local rd_io = { shared = shared } + local wr_io = { shared = shared } + + function rd_io:read_string(max) + if #self.shared.wire == 0 then + if self.shared.closed then + return '', nil -- EOF + end + return nil, nil -- would block + end + + max = max or 1 + local n = math.min(1, max, #self.shared.wire) + local s = self.shared.wire:sub(1, n) + self.shared.wire = self.shared.wire:sub(n + 1) + return s, nil + end + + function wr_io:write_string(str) + if self.shared.closed then + return nil, 'closed' + end + if #str == 0 then + return 0, nil + end + + local ch = str:sub(1, 1) + self.shared.wire = self.shared.wire .. ch + + self.shared.waitset:notify_all('rd', runtime.current_scheduler) + return 1, nil + end + + function rd_io:on_readable(task) + return self.shared.waitset:add('rd', task) + end + + function wr_io:on_writable(task) + runtime.current_scheduler:schedule(task) + return { unlink = function () end } + end + + function rd_io:close() + self.shared.closed = true + self.shared.waitset:notify_all('rd', runtime.current_scheduler) + return true + end + + function wr_io:close() + self.shared.closed = true + self.shared.waitset:notify_all('rd', runtime.current_scheduler) + return true + end + + function rd_io:seek() return nil, 'not seekable' end + + function wr_io:seek() return nil, 'not seekable' end + + function rd_io:nonblock() end + + function rd_io:block() end + + function wr_io:nonblock() end + + function wr_io:block() end + + local rd = stream.open(rd_io, true, false) + local wr = stream.open(wr_io, false, true) + return rd, wr, shared +end + +local function make_stream_pair_full_backpressure(cap) + cap = cap or 8 + + local shared = { + wire = '', + closed = false, + waitset = wait.new_waitset(), -- keys: 'rd', 'wr' + cap = cap, + } + + local rd_io = { shared = shared } + local wr_io = { shared = shared } + + function rd_io:read_string(max) + if #self.shared.wire == 0 then + if self.shared.closed then + return '', nil -- EOF + end + return nil, nil -- would block + end + + max = max or 1 + local n = math.min(1, max, #self.shared.wire) + local s = self.shared.wire:sub(1, n) + self.shared.wire = self.shared.wire:sub(n + 1) + + self.shared.waitset:notify_all('wr', runtime.current_scheduler) + return s, nil + end + + function wr_io:write_string(str) + if self.shared.closed then + return nil, 'closed' + end + if #str == 0 then + return 0, nil + end + + if #self.shared.wire >= self.shared.cap then + return nil, nil, 'wr' + end + + local ch = str:sub(1, 1) + self.shared.wire = self.shared.wire .. ch + + self.shared.waitset:notify_all('rd', runtime.current_scheduler) + return 1, nil + end + + function rd_io:on_readable(task) + return self.shared.waitset:add('rd', task) + end + + function wr_io:on_writable(task) + return self.shared.waitset:add('wr', task) + end + + function rd_io:close() + self.shared.closed = true + self.shared.waitset:notify_all('rd', runtime.current_scheduler) + self.shared.waitset:notify_all('wr', runtime.current_scheduler) + return true + end + + function wr_io:close() + self.shared.closed = true + self.shared.waitset:notify_all('rd', runtime.current_scheduler) + self.shared.waitset:notify_all('wr', runtime.current_scheduler) + return true + end + + function rd_io:seek() return nil, 'not seekable' end + function wr_io:seek() return nil, 'not seekable' end + function rd_io:nonblock() end + function rd_io:block() end + function wr_io:nonblock() end + function wr_io:block() end local rd = stream.open(rd_io, true, false) @@ -177,29 +391,129 @@ local function make_stream_pair_want_wr() return rd, wr, shared end +---------------------------------------------------------------------- +-- Backend 4: write error injection (sticky write error propagation) +---------------------------------------------------------------------- + +local function make_stream_pair_write_error(opts) + opts = opts or {} + local fail_after = opts.fail_after or 4 + + local shared = { + wire = '', + closed = false, + waitset = wait.new_waitset(), -- key 'rd' + writes = 0, + fail_after = fail_after, + } + + local rd_io = { shared = shared } + local wr_io = { shared = shared } + + function rd_io:read_string(max) + if #self.shared.wire == 0 then + if self.shared.closed then + return '', nil -- EOF + end + return nil, nil -- would block + end + max = max or 1 + local n = math.min(1, max, #self.shared.wire) + local s = self.shared.wire:sub(1, n) + self.shared.wire = self.shared.wire:sub(n + 1) + return s, nil + end + + function wr_io:write_string(str) + if self.shared.closed then + return nil, 'closed' + end + if #str == 0 then + return 0, nil + end + + self.shared.writes = self.shared.writes + 1 + if self.shared.writes >= self.shared.fail_after then + return nil, 'boom' -- injected hard error + end + + local ch = str:sub(1, 1) + self.shared.wire = self.shared.wire .. ch + self.shared.waitset:notify_all('rd', runtime.current_scheduler) + return 1, nil + end + + function rd_io:on_readable(task) + return self.shared.waitset:add('rd', task) + end + + function wr_io:on_writable(task) + runtime.current_scheduler:schedule(task) + return { unlink = function () end } + end + + function rd_io:close() + self.shared.closed = true + self.shared.waitset:notify_all('rd', runtime.current_scheduler) + return true + end + + function wr_io:close() + self.shared.closed = true + self.shared.waitset:notify_all('rd', runtime.current_scheduler) + return true + end + + function rd_io:seek() return nil, 'not seekable' end + + function wr_io:seek() return nil, 'not seekable' end + + function rd_io:nonblock() end + + function rd_io:block() end + + function wr_io:nonblock() end + + function wr_io:block() end + + local rd = stream.open(rd_io, true, false) + local wr = stream.open(wr_io, false, true) + return rd, wr, shared +end + +---------------------------------------------------------------------- +-- Tests +---------------------------------------------------------------------- + local function test_basic_line_read() local rd, wr, shared = make_stream_pair() + rd:setvbuf('full') wr:setvbuf('line') - assert(wr.line_buffering == true, "setvbuf('line') did not set line_buffering") + assert_truthy(wr.line_buffering == true, "setvbuf('line') did not set line_buffering") local message = 'hello, world\n' fibers.spawn(function () sleep.sleep(0.01) - local n, err = wr:write(message) - assert(err == nil, 'write error: ' .. tostring(err)) - assert(n == #message, 'write wrote ' .. tostring(n) .. ' bytes, expected ' .. #message) - wr:close() + local n, err = perform(wr:write_op(message)) + assert_eq(err, nil, 'write error') + assert_eq(n, #message, 'write length mismatch') + + local ok, cerr = perform(wr:close_op()) + assert_eq(ok, true, 'close ok expected') + assert_eq(cerr, nil, 'close err expected nil') end) - local line, err = rd:read('*L') - assert(err == nil, "read('*L') returned error: " .. tostring(err)) - assert(line == message, - ("read('*L') returned %q, expected %q"):format(tostring(line), tostring(message))) + local line, err = perform(rd:read_line_op { keep_terminator = true }) + assert_eq(err, nil, 'read_line_op error') + assert_eq(line, message, 'read_line_op returned wrong line') - rd:close() - assert(shared.waitset:size('rd') == 0, 'waitset still has readers after close') + local ok, cerr = perform(rd:close_op()) + assert_eq(ok, true, 'close ok expected') + assert_eq(cerr, nil, 'close err expected nil') + + assert_eq(shared.waitset:size('rd'), 0, 'waitset still has readers after close') end local function test_close_unblocks_reader_no_crash() @@ -207,31 +521,33 @@ local function test_close_unblocks_reader_no_crash() fibers.spawn(function () sleep.sleep(0.01) - rd:close() + local ok, cerr = perform(rd:close_op()) + assert_eq(ok, true) + assert_eq(cerr, nil) end) - local won, line, err = with_timeout(rd:read_op('*L'), 0.2) - assert(won == true, 'timed out waiting for blocked read to resolve on close') - assert(line == nil, 'expected nil line on close, got ' .. tostring(line)) - assert(err == 'stream closed', 'expected err "stream closed", got ' .. tostring(err)) + local won, line, err = with_timeout(rd:read_line_op { keep_terminator = true }, 0.2) + assert_eq(won, true, 'timed out waiting for blocked read to resolve on close') + assert_eq(line, nil, 'expected nil line on close') + assert_closed_err(err) - assert(shared.waitset:size('rd') == 0, 'waitset still has readers after close-unblock') + assert_eq(shared.waitset:size('rd'), 0, 'waitset still has readers after close-unblock') - wr:close() + local ok, cerr = perform(wr:close_op()) + assert_eq(ok, true) + assert_eq(cerr, nil) end local function test_abort_unlinks_waiters() local rd, wr, shared = make_stream_pair() - -- Block a read, then abort it via timeout choice. local won = with_timeout(rd:read_exactly_op(1), 0.02) - assert(won == false, 'expected timeout branch to win') + assert_eq(won, false, 'expected timeout branch to win') - -- The op lost the choice; its wait registration must be cancelled. - assert(shared.waitset:size('rd') == 0, 'waitset leaked readers after abort') + assert_eq(shared.waitset:size('rd'), 0, 'waitset leaked readers after abort') - rd:close() - wr:close() + perform(rd:close_op()) + perform(wr:close_op()) end local function test_want_wiring_wr() @@ -241,30 +557,386 @@ local function test_want_wiring_wr() fibers.spawn(function () sleep.sleep(0.01) - local n, err = wr:write(message) - assert(err == nil, 'write error: ' .. tostring(err)) - assert(n == #message, 'write wrote ' .. tostring(n) .. ' bytes, expected ' .. #message) - wr:close() + local n, err = perform(wr:write_op(message)) + assert_eq(err, nil) + assert_eq(n, #message) + perform(wr:close_op()) end) - local won, line, err = with_timeout(rd:read_op('*L'), 0.2) - assert(won == true, 'timed out: want="wr" registration did not wake') - assert(err == nil, 'read returned error: ' .. tostring(err)) - assert(line == message, ('read returned %q, expected %q'):format(tostring(line), tostring(message))) + local won, line, err = with_timeout(rd:read_line_op { keep_terminator = true }, 0.2) + assert_eq(won, true, 'timed out: want="wr" registration did not wake') + assert_eq(err, nil) + assert_eq(line, message) + + assert_truthy(shared.wr_regs > 0, 'expected on_writable registrations (want="wr")') + assert_eq(shared.rd_regs, 0, 'unexpected on_readable registrations; want wiring may be ignored') + + perform(rd:close_op()) + assert_eq(shared.waitset:size('wr'), 0, 'waitset leaked wr waiters') +end + +local function test_flush_is_noop_on_readonly() + local rd, _, _ = make_stream_pair() - -- Strong regression checks: should register on_writable (want='wr'), not on_readable. - assert(shared.wr_regs > 0, 'expected on_writable registrations (want="wr")') - assert(shared.rd_regs == 0, 'unexpected on_readable registrations; want wiring may be ignored') + local ok, err = perform(rd:flush_op()) + assert_ok_or_zero(ok, 'flush on read-only should succeed') + assert_eq(err, nil, 'flush on read-only should have nil err') - rd:close() - assert(shared.waitset:size('wr') == 0, 'waitset leaked wr waiters') + perform(rd:close_op()) end +local function test_read_some_and_exactly_and_all_eof_shapes() + local rd, wr, _ = make_stream_pair_full() + + local msg = 'abcdef' + local n, werr = perform(wr:write_op(msg)) + assert_eq(werr, nil) + assert_eq(n, #msg) + + local fok, ferr = perform(wr:flush_op()) + assert_ok_or_zero(fok); assert_eq(ferr, nil) + perform(wr:close_op()) + + local s1, e1 = perform(rd:read_some_op(2)) + assert_eq(e1, nil) + assert_truthy(type(s1) == 'string' and #s1 > 0 and #s1 <= 2, 'read_some size') + + local rest_needed = #msg - #s1 + local s2, e2 = perform(rd:read_exactly_op(rest_needed)) + assert_eq(e2, nil) + assert_eq(#s2, rest_needed) + + local s3, e3 = perform(rd:read_some_op(10)) + assert_eq(s3, nil) + assert_truthy(e3 == nil or e3 == 'eof', 'expected eof indicator') + + local all, aerr = perform(rd:read_all_op()) + assert_eq(all, '') + assert_eq(aerr, nil, 'read_all should treat eof as success') + + perform(rd:close_op()) +end + +local function test_write_buffering_write_then_flush_drains() + local rd, wr, shared = make_stream_pair_full() + + local msg = ('x'):rep(256) + + local won, n, err = with_timeout(wr:write_op(msg), 0.05) + assert_eq(won, true, 'write_op should not block on backend drain') + assert_eq(err, nil) + assert_eq(n, #msg) + + local won2, ok, ferr = with_timeout(wr:flush_op(), 0.5) + assert_eq(won2, true, 'flush_op should complete') + assert_ok_or_zero(ok) + assert_eq(ferr, nil) + + perform(wr:close_op()) + + local got, rerr = perform(rd:read_all_op()) + assert_eq(rerr, nil) + assert_eq(got, msg) + + perform(rd:close_op()) + + assert_eq(shared.waitset:size('rd'), 0, 'rd waiters leaked') + assert_eq(shared.waitset:size('wr'), 0, 'wr waiters leaked') +end + +local function test_concurrent_writers_are_serialised_no_interleave() + local rd, wr, shared = make_stream_pair_full() + local wg = require('fibers.waitgroup').new() + + local a = ('A'):rep(64) + local b = ('B'):rep(64) + + wg:add(2) + + fibers.spawn(function () + local n, err = perform(wr:write_op(a)) + assert_eq(err, nil); assert_eq(n, #a) + wg:done() + end) + + fibers.spawn(function () + local n, err = perform(wr:write_op(b)) + assert_eq(err, nil); assert_eq(n, #b) + wg:done() + end) + + wg:wait() + + local ok, ferr = perform(wr:flush_op()) + assert_ok_or_zero(ok); assert_eq(ferr, nil) + + perform(wr:close_op()) + + local all, rerr = perform(rd:read_all_op()) + assert_eq(rerr, nil) + assert_eq(#all, #a + #b) + + local ab = a .. b + local ba = b .. a + assert_truthy(all == ab or all == ba, 'write interleaving detected: got=' .. tostring(all)) + + perform(rd:close_op()) + + assert_eq(shared.waitset:size('rd'), 0) +end + +local function test_abort_unlinks_write_waiters_and_does_not_deadlock() + local rd, wr, shared = make_stream_pair_full_backpressure(8) + + local msg = ('z'):rep(512) + local n, err = perform(wr:write_op(msg)) + assert_eq(err, nil); assert_eq(n, #msg) + + local won1 = with_timeout(wr:flush_op(), 0.001) + assert_eq(won1, false, 'expected timeout branch to win (flush should block under backpressure)') + + local wr1 = shared.waitset:size('wr') + local rd1 = shared.waitset:size('rd') + assert_truthy(wr1 == 0 or wr1 == 1, ('unexpected wr waiter count after abort: %d'):format(wr1)) + assert_eq(rd1, 0, ('unexpected rd waiters after abort: %d'):format(rd1)) + + local won2 = with_timeout(wr:flush_op(), 0.001) + assert_eq(won2, false, 'expected timeout branch to win again') + + local wr2 = shared.waitset:size('wr') + local rd2 = shared.waitset:size('rd') + assert_eq(rd2, 0, ('unexpected rd waiters after second abort: %d'):format(rd2)) + assert_eq(wr2, wr1, ('wr waiter count grew across aborts: %d -> %d'):format(wr1, wr2)) + + local wg = require('fibers.waitgroup').new() + wg:add(1) + + fibers.spawn(function () + local s, rerr = perform(rd:read_exactly_op(#msg)) + assert_eq(rerr, nil) + assert_eq(#s, #msg) + wg:done() + end) + + local won3, ok3, ferr3 = with_timeout(wr:flush_op(), 0.5) + assert_eq(won3, true, 'flush did not complete after reader drained') + assert_ok_or_zero(ok3) + assert_eq(ferr3, nil) + + assert_eq(shared.waitset:size('wr'), 0, 'wr waiters not cleared after successful flush') + assert_eq(shared.waitset:size('rd'), 0, 'rd waiters not cleared after successful flush') + + perform(wr:close_op()) + wg:wait() + perform(rd:close_op()) + + assert_eq(shared.waitset:size('wr'), 0, 'wr waiters leaked at end') + assert_eq(shared.waitset:size('rd'), 0, 'rd waiters leaked at end') +end + +local function test_seek_and_setvbuf_surface() + local rd, wr, _ = make_stream_pair() + + rd:setvbuf('full') + assert_eq(rd.line_buffering, false) + + wr:setvbuf('line') + assert_eq(wr.line_buffering, true) + + wr:setvbuf('no') + assert_eq(wr.line_buffering, false) + + local pos, err = rd:seek('cur', 0) + assert_eq(pos, nil) + assert_truthy(err ~= nil) + + perform(rd:close_op()) + perform(wr:close_op()) +end + +local function test_close_is_idempotent_and_unblocks_waiters() + local rd, wr, shared = make_stream_pair_full() + + fibers.spawn(function () + sleep.sleep(0.01) + perform(rd:close_op()) + end) + + local won, line, err = with_timeout(rd:read_line_op { keep_terminator = true }, 0.2) + assert_eq(won, true) + assert_eq(line, nil) + assert_closed_err(err) + + local ok2, err2 = perform(rd:close_op()) + assert_eq(ok2, true) + assert_eq(err2, nil) + + perform(wr:close_op()) + + assert_eq(shared.waitset:size('rd'), 0) + assert_eq(shared.waitset:size('wr'), 0) +end + +---------------------------------------------------------------------- +-- New close semantics tests (for latched close + prompt begin on block) +---------------------------------------------------------------------- + +local function test_close_is_side_effect_free_when_it_loses_in_choice() + local rd, wr, shared = make_stream_pair_full() + + -- First arm is immediately ready; close_op should lose without starting close. + local won, v = perform(op.boolean_choice(op.always('win'), wr:close_op())) + assert_eq(won, true) + assert_eq(v, 'win') + + assert_truthy(not wr._closing, 'close should not begin during speculative probe') + assert_truthy(not wr._closed, 'stream should not be closed after losing close arm') + + -- Stream remains usable. + local msg = 'ok\n' + local n, werr = perform(wr:write_op(msg)) + assert_eq(werr, nil) + assert_eq(n, #msg) + local okf, ferr = perform(wr:flush_op()) + assert_ok_or_zero(okf); assert_eq(ferr, nil) + + perform(wr:close_op()) + local got, rerr = perform(rd:read_all_op()) + assert_eq(rerr, nil) + assert_eq(got, msg) + perform(rd:close_op()) + + assert_eq(shared.waitset:size('rd'), 0) + assert_eq(shared.waitset:size('wr'), 0) + assert_internal_ws_empty(wr, 'wr internal waitset leak after choice-losing close') + assert_internal_ws_empty(rd, 'rd internal waitset leak after choice-losing close') +end + +local function test_close_blocks_until_flush_completes_and_starts_promptly() + local rd, wr, shared = make_stream_pair_full_backpressure(4) + + local msg = ('m'):rep(128) + local n, err = perform(wr:write_op(msg)) + assert_eq(err, nil) + assert_eq(n, #msg) + + local box = { done = false, ok = nil, err = nil } + local wg = waitgroup.new() + wg:add(1) + + fibers.spawn(function () + local ok, cerr = perform(wr:close_op()) + box.ok = ok + box.err = cerr + box.done = true + wg:done() + end) + + -- Give the close a chance to enter blocking path and begin closing. + sleep.sleep(0.01) + assert_truthy(wr._closing or wr._closed, 'close did not begin promptly once blocked') + assert_eq(box.done, false, 'close_op returned before flush drained') + + -- Drain the wire; this should allow the writer pump to make progress. + local got, rerr = perform(rd:read_exactly_op(#msg)) + assert_eq(rerr, nil) + assert_eq(#got, #msg) + + wg:wait() + assert_eq(box.ok, true, 'close_op should succeed once drained') + assert_eq(box.err, nil) + + perform(rd:close_op()) + + assert_eq(shared.waitset:size('rd'), 0) + assert_eq(shared.waitset:size('wr'), 0) + assert_internal_ws_empty(wr, 'wr internal waitset leak after blocking close') + assert_internal_ws_empty(rd, 'rd internal waitset leak after blocking close') +end + +local function test_close_aborted_in_choice_still_completes() + local rd, wr, shared = make_stream_pair_full_backpressure(4) + + local msg = ('q'):rep(128) + local n, err = perform(wr:write_op(msg)) + assert_eq(err, nil) + assert_eq(n, #msg) + + -- Start a close, but race it against a short timeout so the close arm loses. + local won = with_timeout(wr:close_op(), 0.01) + assert_eq(won, false, 'expected timeout branch to win; close should still be pending') + + -- Allow any scheduled close/pump work to run. + sleep.sleep(0.01) + + -- Drain, which should allow close to finish in the background. + local got, rerr = perform(rd:read_exactly_op(#msg)) + assert_eq(rerr, nil) + assert_eq(#got, #msg) + + -- A subsequent close should now complete (idempotent). + local won2, ok2, err2 = with_timeout(wr:close_op(), 0.5) + assert_eq(won2, true, 'close did not complete after drain') + assert_eq(ok2, true) + assert_eq(err2, nil) + + perform(rd:close_op()) + + assert_eq(shared.waitset:size('rd'), 0) + assert_eq(shared.waitset:size('wr'), 0) + assert_internal_ws_empty(wr, 'wr internal waitset leak after aborted close') + assert_internal_ws_empty(rd, 'rd internal waitset leak after aborted close') +end + +local function test_close_reports_sticky_write_error_and_terminates() + local rd, wr, shared = make_stream_pair_write_error { fail_after = 3 } + + -- Enqueue more than fail_after bytes so the pump hits the injected error. + local msg = ('x'):rep(32) + local n, werr = perform(wr:write_op(msg)) + assert_eq(werr, nil) + assert_eq(n, #msg) + + -- Allow the pump to run and observe the backend error. + sleep.sleep(0.02) + + local ok, cerr = perform(wr:close_op()) + assert_eq(ok, nil, 'close should fail when a sticky write error is present') + assert_eq(cerr, 'boom', 'unexpected close error') + + -- Writer should be terminated best-effort. + assert_truthy(wr._closed, 'writer stream not terminated after close error') + + -- Reader should be closable and should not strand waiters. + perform(rd:close_op()) + + assert_eq(shared.waitset:size('rd'), 0) + assert_internal_ws_empty(wr, 'wr internal waitset leak after close error') + assert_internal_ws_empty(rd, 'rd internal waitset leak after close error') +end + +---------------------------------------------------------------------- +-- Main +---------------------------------------------------------------------- + local function main() test_basic_line_read() test_close_unblocks_reader_no_crash() test_abort_unlinks_waiters() test_want_wiring_wr() + + test_flush_is_noop_on_readonly() + test_read_some_and_exactly_and_all_eof_shapes() + test_write_buffering_write_then_flush_drains() + test_concurrent_writers_are_serialised_no_interleave() + test_abort_unlinks_write_waiters_and_does_not_deadlock() + test_seek_and_setvbuf_surface() + test_close_is_idempotent_and_unblocks_waiters() + + test_close_is_side_effect_free_when_it_loses_in_choice() + test_close_blocks_until_flush_completes_and_starts_promptly() + test_close_aborted_in_choice_still_completes() + test_close_reports_sticky_write_error_and_terminates() end fibers.run(main)