diff --git a/res/modules/internal/stream_providers/file.lua b/res/modules/internal/stream_providers/file.lua index 484e869e4..840ff93f2 100644 --- a/res/modules/internal/stream_providers/file.lua +++ b/res/modules/internal/stream_providers/file.lua @@ -4,6 +4,7 @@ local lib = { read = file.__read_descriptor, write = file.__write_descriptor, seek = file.__seek_descriptor, + tell = file.__tell_descriptor, flush = file.__flush_descriptor, is_alive = file.__has_descriptor, close = file.__close_descriptor @@ -14,6 +15,7 @@ file.__open_descriptor = nil file.__read_descriptor = nil file.__write_descriptor = nil file.__seek_descriptor = nil +file.__tell_descriptor = nil file.__flush_descriptor = nil file.__has_descriptor = nil file.__close_descriptor = nil diff --git a/res/modules/io_stream.lua b/res/modules/io_stream.lua index b608a22a9..8a4507a3b 100644 --- a/res/modules/io_stream.lua +++ b/res/modules/io_stream.lua @@ -373,6 +373,10 @@ function io_stream:seek(mode, offset) self.ioLib.seek(self.descriptor, mode, offset) end +function io_stream:tell() + return self.ioLib.tell(self.descriptor) +end + function io_stream:is_alive() return self.ioLib.is_alive(self.descriptor) end @@ -399,4 +403,4 @@ function io_stream:flush() if self.flushMode ~= FLUSH_MODE_ONLY_BUFFER then self.ioLib.flush(self.descriptor) end end -return io_stream \ No newline at end of file +return io_stream diff --git a/src/logic/scripting/io_descriptors.cpp b/src/logic/scripting/io_descriptors.cpp index e6fda04dc..132eea0eb 100644 --- a/src/logic/scripting/io_descriptors.cpp +++ b/src/logic/scripting/io_descriptors.cpp @@ -14,6 +14,7 @@ using namespace scripting; namespace { struct StreamDescriptor { + // TODO: std::iostream? std::unique_ptr in; std::unique_ptr out; }; @@ -34,6 +35,31 @@ std::ostream* io_descriptors::get_output(int id) { return ::descriptors[id]->out.get(); } +static StreamDescriptor& require_descriptor(int id) { + if (!io_descriptors::has_descriptor(id)) { + throw std::runtime_error( + "io-descriptor with id " + std::to_string(id) + " does not exists" + ); + } + return *::descriptors[id]; +} + +std::istream& io_descriptors::require_input(int id) { + const auto& descriptor = require_descriptor(id); + if (descriptor.in) { + return *descriptor.in; + } + throw std::runtime_error("io-descriptor is not readable"); +} + +std::ostream& io_descriptors::require_output(int id) { + const auto& descriptor = require_descriptor(id); + if (descriptor.out) { + return *descriptor.out; + } + throw std::runtime_error("io-descriptor is not writeable"); +} + void io_descriptors::flush(int id) { if (is_writeable(id)) { ::descriptors[id]->out->flush(); @@ -42,15 +68,11 @@ void io_descriptors::flush(int id) { bool io_descriptors::has_descriptor(int id) { return id >= 0 && id < static_cast(::descriptors.size()) && - ::descriptors[id].has_value() && - (::descriptors[id]->in != nullptr || - ::descriptors[id]->out != nullptr); + ::descriptors[id].has_value(); } bool io_descriptors::is_readable(int id) { - return id >= 0 && id < static_cast(::descriptors.size()) - && ::descriptors[id].has_value() - && ::descriptors[id]->in != nullptr; + return has_descriptor(id) && ::descriptors[id]->in != nullptr; } bool io_descriptors::is_writeable(int id) { diff --git a/src/logic/scripting/io_descriptors.hpp b/src/logic/scripting/io_descriptors.hpp index 0400a920a..7e6207791 100644 --- a/src/logic/scripting/io_descriptors.hpp +++ b/src/logic/scripting/io_descriptors.hpp @@ -9,6 +9,9 @@ namespace scripting::io_descriptors { std::istream* get_input(int id); std::ostream* get_output(int id); + std::istream& require_input(int id); + std::ostream& require_output(int id); + void flush(int id); bool has_descriptor(int id); diff --git a/src/logic/scripting/lua/libs/libfile.cpp b/src/logic/scripting/lua/libs/libfile.cpp index 964c6d487..1a09e9962 100644 --- a/src/logic/scripting/lua/libs/libfile.cpp +++ b/src/logic/scripting/lua/libs/libfile.cpp @@ -302,46 +302,26 @@ static int l_has_descriptor(lua::State* L) { static int l_read_descriptor(lua::State* L) { int descriptor = lua::tointeger(L, 1); - - if (!io_descriptors::has_descriptor(descriptor)) { - throw std::runtime_error("unknown descriptor"); - } - - if (!io_descriptors::is_readable(descriptor)) { - throw std::runtime_error("descriptor is not readable"); - } - int maxlen = lua::tointeger(L, 2); - auto* stream = io_descriptors::get_input(descriptor); + auto& stream = io_descriptors::require_input(descriptor); + if (stream.eof()) { + stream.clear(); + } util::Buffer buffer(maxlen); - - stream->read(buffer.data(), maxlen); - - std::streamsize read_len = stream->gcount(); - + stream.read(buffer.data(), maxlen); + std::streamsize read_len = stream.gcount(); return lua::create_bytearray(L, buffer.data(), read_len); } static int l_write_descriptor(lua::State* L) { int descriptor = lua::tointeger(L, 1); - - if (!io_descriptors::has_descriptor(descriptor)) { - throw std::runtime_error("unknown descriptor"); - } - - if (!io_descriptors::is_writeable(descriptor)) { - throw std::runtime_error("descriptor is not writeable"); - } - auto data = lua::bytearray_as_string(L, 2); - auto* stream = io_descriptors::get_output(descriptor); - - stream->write(data.data(), static_cast(data.size())); - - if (!stream->good()) { + auto& stream = io_descriptors::require_output(descriptor); + stream.write(data.data(), static_cast(data.size())); + if (!stream.good()) { throw std::runtime_error("failed to write to stream"); } return 0; @@ -354,8 +334,9 @@ static int l_seek_descriptor(lua::State* L) { throw std::runtime_error("unknown descriptor"); } - std::string mode = lua::require_string(L, 2); + auto mode = lua::require_string(L, 2); std::ios_base::seekdir dir; + auto position = lua::tointeger(L, 3); switch (mode[0]) { case 'b': @@ -371,15 +352,37 @@ static int l_seek_descriptor(lua::State* L) { throw std::runtime_error("invalid seek mode"); } - auto* stream = io_descriptors::get_output(descriptor); + if (io_descriptors::is_writeable(descriptor)) { + auto& stream = io_descriptors::require_output(descriptor); + stream.seekp(position, dir); + if (!stream.good()) { + throw std::runtime_error("failed to seek stream"); + } + } + if (io_descriptors::is_readable(descriptor)) { + auto& stream = io_descriptors::require_input(descriptor); + stream.seekg(position, dir); + if (!stream.good()) { + throw std::runtime_error("failed to seek stream"); + } + } + return 0; +} - stream->seekp(lua::tointeger(L, 3), dir); +static int l_tell_descriptor(lua::State* L) { + int descriptor = lua::tointeger(L, 1); - if (!stream->good()) { - throw std::runtime_error("failed to seek stream"); + if (!io_descriptors::has_descriptor(descriptor)) { + throw std::runtime_error("unknown descriptor"); } - return 0; + if (io_descriptors::is_writeable(descriptor)) { + auto& stream = io_descriptors::require_output(descriptor); + return lua::pushinteger(L, stream.tellp()); + } else { + auto& stream = io_descriptors::require_input(descriptor); + return lua::pushinteger(L, stream.tellg()); + } } static int l_flush_descriptor(lua::State* L) { @@ -442,6 +445,7 @@ const luaL_Reg filelib[] = { {"__read_descriptor", lua::wrap}, {"__write_descriptor", lua::wrap}, {"__seek_descriptor", lua::wrap}, + {"__tell_descriptor", lua::wrap}, {"__flush_descriptor", lua::wrap}, {"__close_descriptor", lua::wrap}, {"__close_all_descriptors", lua::wrap},