fix(mods): harden streamed import cleanup

This commit is contained in:
HighDrexler
2026-08-20 23:46:02 -04:00
parent 911e11a372
commit c465f58006
3 changed files with 171 additions and 51 deletions
+89 -49
View File
@@ -388,8 +388,16 @@ local function openImportSource(path)
-- read into one Lua string before validation. -- read into one Lua string before validation.
local native = io.open(path, "rb") local native = io.open(path, "rb")
if native then if native then
local size = native:seek("end") local size, sizeErr = native:seek("end")
if size then native:seek("set", 0) end if size == nil or size == false then
native:close()
return nil, sizeErr or "could not determine source file size"
end
local reset, resetErr = native:seek("set", 0)
if reset == nil or reset == false then
native:close()
return nil, resetErr or "could not rewind source file"
end
return { return {
size = size, size = size,
read = function(_, n) return native:read(n) end, read = function(_, n) return native:read(n) end,
@@ -410,7 +418,6 @@ local function openImportSource(path)
end end
return nil, "streaming source access is unavailable" return nil, "streaming source access is unavailable"
end end
local function streamRequiredImport(manifest, importId, source) local function streamRequiredImport(manifest, importId, source)
local RequiredImports = require("src.mods.RequiredImports") local RequiredImports = require("src.mods.RequiredImports")
local spec = RequiredImports.spec(manifest, importId) local spec = RequiredImports.spec(manifest, importId)
@@ -426,58 +433,91 @@ local function streamRequiredImport(manifest, importId, source)
local CacheFs = require("src.import.CacheFs") local CacheFs = require("src.import.CacheFs")
local destination = RequiredImports.path(manifest, spec) local destination = RequiredImports.path(manifest, spec)
local savedPrefix = CacheFs.prefix local savedPrefix = CacheFs.prefix
local output
local inputClosed, outputClosed = false, false
local function closeInput()
if inputClosed then return end
inputClosed = true
pcall(function() input:close() end)
end
local function closeOutput()
if outputClosed or not output then return end
outputClosed = true
pcall(function() output:close() end)
end
local resultOk, resultDetail
CacheFs.prefix = "" CacheFs.prefix = ""
CacheFs.remove(RequiredImports.receiptPath(manifest, spec)) local ran, thrown = xpcall(function()
CacheFs.remove(destination) CacheFs.remove(RequiredImports.receiptPath(manifest, spec))
local output, makeErr = CacheFs.openWrite(destination)
if not output then
CacheFs.prefix = savedPrefix
input:close()
return nil, makeErr or "could not create imported file"
end
local function cleanupDestination()
CacheFs.remove(destination) CacheFs.remove(destination)
CacheFs.prefix = savedPrefix local makeErr
end output, makeErr = CacheFs.openWrite(destination)
if not output then
local MD5 = require("src.mods.StreamMD5") resultDetail = makeErr or "could not create imported file"
local md5 = MD5.new() return
local total, chunkBytes = 0, 4 * 1024 * 1024
while true do
local chunk = input:read(chunkBytes)
if not chunk or #chunk == 0 then break end
md5:update(chunk)
local wrote, writeErr = output:write(chunk)
if wrote == false or wrote == nil then
input:close(); output:close(); cleanupDestination()
return nil, "could not copy import: " .. tostring(writeErr or "write failed")
end end
total = total + #chunk
if #chunk < chunkBytes then break end
end
input:close()
output:close()
if input.size and total ~= input.size then local MD5 = require("src.mods.StreamMD5")
cleanupDestination() local md5 = MD5.new()
return nil, ("source read ended early (expected %d bytes, copied %d)"):format(input.size, total) local total, chunkBytes = 0, 4 * 1024 * 1024
end while true do
local storedSizeErr = RequiredImports.sizeError(spec, total, true) local chunk = input:read(chunkBytes)
if storedSizeErr then cleanupDestination(); return nil, storedSizeErr end if not chunk or #chunk == 0 then break end
local digest = md5:final() md5:update(chunk)
CacheFs.prefix = savedPrefix local wrote, writeErr = output:write(chunk)
local accepted, detail = RequiredImports.acceptStoredDigest( if wrote == false or wrote == nil then
manifest, importId, digest, love.filesystem) resultDetail = "could not copy import: "
if not accepted then .. tostring(writeErr or "write failed")
CacheFs.prefix = "" return
CacheFs.remove(destination) end
total = total + #chunk
if #chunk < chunkBytes then break end
end
closeInput()
closeOutput()
if input.size and total ~= input.size then
resultDetail = ("source read ended early (expected %d bytes, copied %d)")
:format(input.size, total)
return
end
local storedSizeErr = RequiredImports.sizeError(spec, total, true)
if storedSizeErr then resultDetail = storedSizeErr; return end
local digest = md5:final()
-- acceptStoredDigest uses the normal engine path/receipt rules, so
-- restore the caller's prefix before handing control to it.
CacheFs.prefix = savedPrefix CacheFs.prefix = savedPrefix
return nil, detail local accepted, detail = RequiredImports.acceptStoredDigest(
end manifest, importId, digest, love.filesystem)
return true, detail if not accepted then resultDetail = detail; return end
end resultOk, resultDetail = true, detail
end, function(err)
if debug and debug.traceback then
return debug.traceback(tostring(err), 2)
end
return tostring(err)
end)
-- finally: resource handles and the process-global CacheFs prefix must
-- be restored even when a read/hash/write helper raises a Lua error.
closeInput()
closeOutput()
CacheFs.prefix = ""
if not ran or not resultOk then
pcall(function() CacheFs.remove(destination) end)
end
CacheFs.prefix = savedPrefix
if not ran then
return nil, "could not copy import: " .. tostring(thrown)
end
if not resultOk then return nil, resultDetail end
return true, resultDetail
end
local function readDroppedFile(file) local function readDroppedFile(file)
local ok, openError = file:open("r") local ok, openError = file:open("r")
if not ok then return nil, openError end if not ok then return nil, openError end
+14 -2
View File
@@ -228,9 +228,21 @@ function RequiredImports.acceptStoredDigest(manifest, importId, digest, fs)
if sizeErr then return nil, sizeErr end if sizeErr then return nil, sizeErr end
if love and fs == love.filesystem then if love and fs == love.filesystem then
local savedPrefix = CacheFs.prefix local savedPrefix = CacheFs.prefix
CacheFs.prefix = "" local ok, prefixErr = xpcall(function()
CacheFs.remove(removedMarker(manifest, spec)) CacheFs.prefix = ""
CacheFs.remove(removedMarker(manifest, spec))
CacheFs.prefix = savedPrefix
-- writeReceipt has its own temporary CacheFs prefix switch. Keep it
-- inside this guard too so a write error cannot leak global state.
writeReceipt(manifest, spec, digest, info, fs)
end, function(err)
return tostring(err)
end)
CacheFs.prefix = savedPrefix CacheFs.prefix = savedPrefix
if not ok then
return nil, "could not finalize import receipt: " .. tostring(prefixErr)
end
return true, digest
elseif fs and fs.remove then elseif fs and fs.remove then
fs.remove(removedMarker(manifest, spec)) fs.remove(removedMarker(manifest, spec))
end end
@@ -88,6 +88,74 @@ T.eq(love.filesystem.read("mods/stream_probe/baseroms/source.bin"), "abc",
T.eq(importer.requiredImportNotice, nil, "successful stream leaves no import error") T.eq(importer.requiredImportNotice, nil, "successful stream leaves no import error")
T.eq(importer._refreshed, true, "successful stream refreshes the mod list") T.eq(importer._refreshed, true, "successful stream refreshes the mod list")
local CacheFs = require("src.import.CacheFs")
-- A thrown writer error must not leak the temporary empty CacheFs prefix.
local workingNewFile = love.filesystem.newFile
local workingPrefix = CacheFs.prefix
CacheFs.prefix = "sentinel/"
love.filesystem.newFile = function(path)
local file = workingNewFile(path)
if path == "mods/stream_probe/baseroms/source.bin" then
function file:write()
error("forced writer failure")
end
end
return file
end
importer.requiredImportNotice = nil
RequiredImports.LARGE_WARN_BYTES = 2
local failedWrite = importer:_importRequiredSource(
"stream_probe", "source", source, true)
RequiredImports.LARGE_WARN_BYTES = oldWarn
T.eq(failedWrite, nil, "thrown streaming writer error is contained")
T.eq(CacheFs.prefix, "sentinel/",
"streaming writer error restores CacheFs.prefix")
love.filesystem.newFile = workingNewFile
CacheFs.prefix = workingPrefix
-- acceptStoredDigest has its own prefix switch for marker/receipt I/O.
-- Even an unexpected CacheFs failure must restore the caller's prefix.
love.filesystem.write("mods/stream_probe/baseroms/source.bin", "abc")
local workingRemove = CacheFs.remove
CacheFs.prefix = "sentinel/"
CacheFs.remove = function()
error("forced marker removal failure")
end
local accepted = RequiredImports.acceptStoredDigest(
manifest, "source", "900150983cd24fb0d6963f7d28e17f72", love.filesystem)
T.eq(accepted, nil, "acceptStoredDigest contains CacheFs failure")
T.eq(CacheFs.prefix, "sentinel/",
"acceptStoredDigest failure restores CacheFs.prefix")
CacheFs.remove = workingRemove
CacheFs.prefix = workingPrefix
-- Native sources are rejected cleanly if seek-to-start fails after the
-- size probe; never return a handle left sitting at EOF.
local realIoOpen = io.open
local fakeCloses = 0
io.open = function(path, mode)
if path ~= source then return realIoOpen(path, mode) end
local calls = 0
return {
seek = function(_, whence)
calls = calls + 1
if whence == "end" then return 3 end
if whence == "set" then return nil, "forced rewind failure" end
return nil, "unexpected seek"
end,
close = function() fakeCloses = fakeCloses + 1 end,
}
end
importer.requiredImportNotice = nil
RequiredImports.LARGE_WARN_BYTES = 2
local failedSeek = importer:_importRequiredSource(
"stream_probe", "source", source, true)
RequiredImports.LARGE_WARN_BYTES = oldWarn
io.open = realIoOpen
T.eq(failedSeek, nil, "failed native rewind rejects import source")
T.ok(fakeCloses >= 2, "failed native rewind closes probed source handles")
love.filesystem.remove("mods/stream_probe/baseroms/source.bin") love.filesystem.remove("mods/stream_probe/baseroms/source.bin")
love.filesystem.remove("mods/stream_probe/baseroms/source.iso") love.filesystem.remove("mods/stream_probe/baseroms/source.iso")
os.remove(source) os.remove(source)