// test/cpp/cpu/test_memory_manager.cpp
#include <gtest/gtest.h>
#include "cpu/memory_manager.h"
#include <cstring>

namespace cpu_inference {

class MemoryManagerTest : public ::testing::Test {
protected:
    void SetUp() override {
        MemoryManager::get();
    }

    void TearDown() override {
        MemoryManager::get();
    }
};

// 辅助函数:从 Tensor 的指定 NUMA 节点拷出数据
template<typename T>
std::vector<T> read_from_tensor(const Tensor& tensor, int numa_node, size_t count) {
    const void* ptr = tensor.at(numa_node, {0});
    if (!ptr) return {};
    std::vector<T> out(count);
    std::memcpy(out.data(), ptr, count * sizeof(T));
    return out;
}

// 辅助函数:向 Tensor 的指定 NUMA 节点写入数据
template<typename T>
bool write_to_tensor(Tensor& tensor, int numa_node, const std::vector<T>& data) {
    void* ptr = tensor.at(numa_node, {0});
    if (!ptr) return false;
    std::memcpy(ptr, data.data(), data.size() * sizeof(T));
    return true;
}

// 测试 allocSingleNuma
TEST_F(MemoryManagerTest, AllocSingleNuma_Basic) {
    auto& manager = MemoryManager::get();
    std::vector<float> input = {1.0f, 2.0f, 3.0f, 4.0f};

    Tensor& tensor = manager.alloc_single_numa("single", {4}, 0, "float32");

    // 写入
    ASSERT_TRUE(write_to_tensor(tensor, 0, input));

    // 读出
    auto output = read_from_tensor<float>(tensor, 0, input.size());
    EXPECT_EQ(input, output);
}

// 测试 allocMultiNuma(需 >=2 NUMA)
TEST_F(MemoryManagerTest, AllocMultiNuma_TwoNodes) {
    auto& manager = MemoryManager::get();
    int tot_numa_cnt = manager.get_tot_numa_cnt();
    printf("Total NUMA nodes available: %d\n", tot_numa_cnt);

    if (tot_numa_cnt < 2) {
        GTEST_SKIP() << "Only " << tot_numa_cnt << " NUMA node(s) available.";
    }

    Tensor& tensor = manager.alloc_multi_numa("multi", {5}, {0, 1}, "float32");
    std::vector<float> data0 = {10.0f, 20.0f, 30.0f, 40.0f, 50.0f};
    std::vector<float> data1 = {40.0f, 50.0f, 10.0f, 20.0f, 30.0f};

    ASSERT_TRUE(write_to_tensor(tensor, 0, data0));
    ASSERT_TRUE(write_to_tensor(tensor, 1, data1));

    auto out0 = read_from_tensor<float>(tensor, 0, data0.size());
    auto out1 = read_from_tensor<float>(tensor, 1, data1.size());

    EXPECT_EQ(data0, out0);
    EXPECT_EQ(data1, out1);
}

// 测试 allocAllNuma
TEST_F(MemoryManagerTest, AllocAllNuma_Basic) {
    auto& manager = MemoryManager::get();

    Tensor& tensor = manager.alloc_all_numa("all", {8}, "float32");

    // 给每个 NUMA 节点写入不同值
    size_t total = 8;

    for (int node = 0; node < 4; ++node) {
        std::vector<float> block(total, static_cast<float>(100 + node));
        ASSERT_TRUE(write_to_tensor(tensor, node, block));
    }

    // 验证第一个节点
    auto first = read_from_tensor<float>(tensor, 0, 1);
    EXPECT_FLOAT_EQ(first[0], 100.0f);
}

// 测试 get_memory
TEST_F(MemoryManagerTest, GetMemory_Exists) {
    auto& manager = MemoryManager::get();
    manager.alloc_single_numa("lookup", {3}, 0, "float32");

    Tensor& tensor = manager.get_memory("lookup");

    std::vector<float> in = {1.5f, 2.5f, 3.5f};
    ASSERT_TRUE(write_to_tensor(tensor, 0, in));

    auto out = read_from_tensor<float>(tensor, 0, in.size());
    EXPECT_EQ(in, out);
}

// 测试异常:重复名称
TEST_F(MemoryManagerTest, AllocSingleNuma_NameConflict) {
    auto& manager = MemoryManager::get();
    manager.alloc_single_numa("dup", {1}, 0, "float32");
    EXPECT_THROW(
        manager.alloc_single_numa("dup", {1}, 0, "float32"),
        std::runtime_error
    );
}

// 测试异常:无效 NUMA
TEST_F(MemoryManagerTest, AllocSingleNuma_InvalidNuma) {
    auto& manager = MemoryManager::get();
    EXPECT_THROW(
        manager.alloc_single_numa("bad", {1}, -1, "float32"),
        std::out_of_range
    );
    EXPECT_THROW(
        manager.alloc_single_numa("bad", {1}, manager.get_tot_numa_cnt(), "float32"),
        std::out_of_range
    );
}

} // namespace cpu_inference