-- Copyright (c) 2026 Huawei Technologies Co., Ltd.
-- openUBMC is licensed under Mulan PSL v2.
-- You can use this software according to the terms and conditions of the Mulan PSL v2.
-- See the Mulan PSL v2 for more details.

local lu = require('luaunit')
local skynet = require('skynet')
local common_def = require('common_def')

TestControllerCollectionDumpLogs = {}

local MODULE_COLLECTION = 'controller.controller_collection'
local MODULE_CTRL_LOG_DUMP = 'controller.ctrl_log_dump'

local saved_skynet = {}

local function reload_dump_modules()
    package.loaded[MODULE_COLLECTION] = nil
    package.loaded[MODULE_CTRL_LOG_DUMP] = nil
    package.loaded['common_def'] = nil
end

local function save_skynet()
    saved_skynet.fork = skynet.fork
    saved_skynet.wait = skynet.wait
    saved_skynet.wakeup = skynet.wakeup
end

local function restore_skynet()
    if saved_skynet.fork then
        skynet.fork = saved_skynet.fork
    end
    if saved_skynet.wait then
        skynet.wait = saved_skynet.wait
    end
    if saved_skynet.wakeup then
        skynet.wakeup = saved_skynet.wakeup
    end
end

--[[
  在独立 reload 的模块上跑 dump_controller_logs,skynet/ctrl_log_dump 桩仅作用于本次调用。
  @return table { fork_count, wait_count, wakeup_count, collect_calls }
]]
local function run_dump_controller_logs(opts)
    opts = opts or {}
    reload_dump_modules()

    local controller_collection = require(MODULE_COLLECTION)
    local ctrl_log_dump = require(MODULE_CTRL_LOG_DUMP)

    local fork_count = 0
    local wait_count = 0
    local wakeup_count = 0
    local fork_queue = {}
    local collect_calls = {}

    if opts.defer_fork then
        skynet.fork = function(fn)
            fork_count = fork_count + 1
            table.insert(fork_queue, fn)
        end
        skynet.wait = function(_)
            wait_count = wait_count + 1
            for _, fn in ipairs(fork_queue) do
                fn()
            end
            fork_queue = {}
        end
        skynet.wakeup = function(_)
            wakeup_count = wakeup_count + 1
        end
    else
        skynet.fork = function(fn)
            fork_count = fork_count + 1
            fn()
        end
        skynet.wait = function(_)
            wait_count = wait_count + 1
        end
        skynet.wakeup = function(_)
            wakeup_count = wakeup_count + 1
        end
    end

    local orig_ctrl = {}
    local stubs = opts.ctrl_log_dump_stubs or {}
    for name, fn in pairs(stubs) do
        orig_ctrl[name] = ctrl_log_dump[name]
        ctrl_log_dump[name] = fn
    end

    if stubs.check_storage_ready == nil then
        orig_ctrl.check_storage_ready = ctrl_log_dump.check_storage_ready
        ctrl_log_dump.check_storage_ready = function()
            return true
        end
    end
    if stubs.is_pmc_or_huawei == nil then
        orig_ctrl.is_pmc_or_huawei = ctrl_log_dump.is_pmc_or_huawei
        ctrl_log_dump.is_pmc_or_huawei = function(type_id)
            return type_id == 1
        end
    end
    if stubs.check_precondition == nil then
        orig_ctrl.check_precondition = ctrl_log_dump.check_precondition
        ctrl_log_dump.check_precondition = function(obj)
            return obj.OOBSupport == 1
        end
    end
    if stubs.prepare_src_dir == nil then
        orig_ctrl.prepare_src_dir = ctrl_log_dump.prepare_src_dir
        ctrl_log_dump.prepare_src_dir = function(obj)
            if obj.no_dir then
                return nil
            end
            return '/data/var/log/storage/ctrllog/C' .. obj.Id
        end
    end
    if stubs.collect_raw_logs == nil then
        orig_ctrl.collect_raw_logs = ctrl_log_dump.collect_raw_logs
        ctrl_log_dump.collect_raw_logs = function(obj, src_dir)
            table.insert(collect_calls, { Id = obj.Id, src_dir = src_dir })
        end
    end

    local inst = {
        get_all_controllers = function()
            return opts.controllers or {}
        end,
    }

    local ok, err = pcall(controller_collection.dump_controller_logs, inst)
    for name, fn in pairs(orig_ctrl) do
        ctrl_log_dump[name] = fn
    end

    if not ok then
        error(err)
    end

    return {
        fork_count = fork_count,
        wait_count = wait_count,
        wakeup_count = wakeup_count,
        collect_calls = collect_calls,
    }
end

function TestControllerCollectionDumpLogs:setUp()
    save_skynet()
    reload_dump_modules()
end

function TestControllerCollectionDumpLogs:tearDown()
    restore_skynet()
    reload_dump_modules()
end

function TestControllerCollectionDumpLogs:test_dump_controller_logs_storage_not_ready()
    local ret = run_dump_controller_logs({
        ctrl_log_dump_stubs = {
            check_storage_ready = function()
                return false
            end,
        },
        controllers = {
            { Id = 1, TypeId = 1, OOBSupport = 1 },
        },
    })
    lu.assertEquals(ret.fork_count, 0)
    lu.assertEquals(ret.wait_count, 0)
    lu.assertEquals(#ret.collect_calls, 0)
end

function TestControllerCollectionDumpLogs:test_dump_controller_logs_no_controllers()
    local ret = run_dump_controller_logs({
        controllers = {},
    })
    lu.assertEquals(ret.fork_count, 0)
    lu.assertEquals(ret.wait_count, 0)
    lu.assertEquals(#ret.collect_calls, 0)
end

function TestControllerCollectionDumpLogs:test_dump_controller_logs_skip_filters()
    local ret = run_dump_controller_logs({
        controllers = {
            { Id = 0, TypeId = 0, OOBSupport = 1 },
            { Id = common_def.SML_MAX_RAID_CONTROLLER, TypeId = 1, OOBSupport = 1 },
            { Id = 1, TypeId = 1, OOBSupport = 0 },
            { Id = 2, TypeId = 1, OOBSupport = 1, no_dir = true },
            { Id = 3, TypeId = 1, OOBSupport = 1 },
        },
    })
    lu.assertEquals(ret.fork_count, 1)
    lu.assertEquals(#ret.collect_calls, 1)
    lu.assertEquals(ret.collect_calls[1].Id, 3)
    lu.assertStrContains(ret.collect_calls[1].src_dir, 'C3')
end

function TestControllerCollectionDumpLogs:test_dump_controller_logs_collect_success_sync()
    local ret = run_dump_controller_logs({
        controllers = {
            { Id = 1, TypeId = 1, OOBSupport = 1 },
        },
    })
    lu.assertEquals(ret.fork_count, 1)
    lu.assertEquals(ret.wait_count, 0)
    lu.assertEquals(ret.wakeup_count, 1)
    lu.assertEquals(#ret.collect_calls, 1)
    lu.assertEquals(ret.collect_calls[1].Id, 1)
end

function TestControllerCollectionDumpLogs:test_dump_controller_logs_collect_failed_sync()
    local ret = run_dump_controller_logs({
        ctrl_log_dump_stubs = {
            collect_raw_logs = function()
                error('stub collect error')
            end,
        },
        controllers = {
            { Id = 2, TypeId = 1, OOBSupport = 1 },
        },
    })
    lu.assertEquals(ret.fork_count, 1)
    lu.assertEquals(#ret.collect_calls, 0)
end

function TestControllerCollectionDumpLogs:test_dump_controller_logs_deferred_fork_wait()
    local ret = run_dump_controller_logs({
        defer_fork = true,
        controllers = {
            { Id = 4, TypeId = 1, OOBSupport = 1 },
            { Id = 5, TypeId = 1, OOBSupport = 1 },
        },
    })
    lu.assertEquals(ret.fork_count, 2)
    lu.assertEquals(ret.wait_count, 1)
    lu.assertEquals(ret.wakeup_count, 1)
    lu.assertEquals(#ret.collect_calls, 2)
    lu.assertEquals(ret.collect_calls[1].Id, 4)
    lu.assertEquals(ret.collect_calls[2].Id, 5)
end

function TestControllerCollectionDumpLogs:test_dump_controller_logs_multiple_sync_parallel()
    local ret = run_dump_controller_logs({
        controllers = {
            { Id = 6, TypeId = 1, OOBSupport = 1 },
            { Id = 7, TypeId = 1, OOBSupport = 1 },
        },
    })
    lu.assertEquals(ret.fork_count, 2)
    lu.assertEquals(ret.wait_count, 0)
    lu.assertTrue(ret.wakeup_count >= 1)
    lu.assertEquals(#ret.collect_calls, 2)
end