#include "navigation/navmesh_bitset_fusion.hpp"
#include <queue>
#include <unordered_map>
#include <algorithm>

namespace navigation {

bool NavMeshBitsetFusion::InitializeFromNativeTerrain(const memory::NativeTerrainData& terrain) {
    if (!terrain.IsValid()) return false;

    m_cols = terrain.cols;
    m_rows = terrain.rows;
    m_gridToWorldScale = (terrain.gridToWorldScale > 0.001f) ? terrain.gridToWorldScale : 10.875f;
    m_worldOriginX = terrain.worldOriginX;
    m_worldOriginY = terrain.worldOriginY;

    m_wordsPerRow = (m_cols + 63) / 64;
    const size_t totalWords = static_cast<size_t>(m_rows) * m_wordsPerRow;

    m_walkableBits.assign(totalWords, 0ULL);
    m_exploredBits.assign(totalWords, 0ULL);
    m_totalWalkableCells = 0;
    m_exploredWalkableCells = 0;

    for (uint32_t y = 0; y < m_rows; ++y) {
        for (uint32_t x = 0; x < m_cols; ++x) {
            if (terrain.IsWalkableGrid(x, y)) {
                const size_t wIdx = WordIndex(x, y);
                m_walkableBits[wIdx] |= BitMask(x);
                ++m_totalWalkableCells;
            }
        }
    }

    return m_totalWalkableCells > 0;
}

bool NavMeshBitsetFusion::IsWalkable(uint32_t gx, uint32_t gy) const {
    if (gx >= m_cols || gy >= m_rows || m_walkableBits.empty()) return false;
    return (m_walkableBits[WordIndex(gx, gy)] & BitMask(gx)) != 0ULL;
}

bool NavMeshBitsetFusion::IsExplored(uint32_t gx, uint32_t gy) const {
    if (gx >= m_cols || gy >= m_rows || m_exploredBits.empty()) return false;
    return (m_exploredBits[WordIndex(gx, gy)] & BitMask(gx)) != 0ULL;
}

bool NavMeshBitsetFusion::WorldToGrid(float wx, float wy, uint32_t& outGx, uint32_t& outGy) const {
    if (m_gridToWorldScale <= 0.001f || m_cols == 0 || m_rows == 0) return false;
    const float localX = (wx - m_worldOriginX) / m_gridToWorldScale;
    const float localY = (wy - m_worldOriginY) / m_gridToWorldScale;
    if (localX < 0.0f || localY < 0.0f) return false;
    outGx = static_cast<uint32_t>(localX);
    outGy = static_cast<uint32_t>(localY);
    return outGx < m_cols && outGy < m_rows;
}

bool NavMeshBitsetFusion::GridToWorld(uint32_t gx, uint32_t gy, float& outWx, float& outWy) const {
    if (gx >= m_cols || gy >= m_rows) return false;
    outWx = m_worldOriginX + (static_cast<float>(gx) + 0.5f) * m_gridToWorldScale;
    outWy = m_worldOriginY + (static_cast<float>(gy) + 0.5f) * m_gridToWorldScale;
    return true;
}

bool NavMeshBitsetFusion::IsWalkableWorld(float wx, float wy) const {
    uint32_t gx = 0, gy = 0;
    if (!WorldToGrid(wx, wy, gx, gy)) return false;
    return IsWalkable(gx, gy);
}

bool NavMeshBitsetFusion::IsExploredWorld(float wx, float wy) const {
    uint32_t gx = 0, gy = 0;
    if (!WorldToGrid(wx, wy, gx, gy)) return false;
    return IsExplored(gx, gy);
}

void NavMeshBitsetFusion::SetExploredGrid(uint32_t gx, uint32_t gy, bool explored) {
    if (gx >= m_cols || gy >= m_rows) return;
    const size_t wIdx = WordIndex(gx, gy);
    const uint64_t mask = BitMask(gx);
    const bool wasExplored = (m_exploredBits[wIdx] & mask) != 0ULL;
    const bool isWalk = (m_walkableBits[wIdx] & mask) != 0ULL;

    if (explored && !wasExplored) {
        m_exploredBits[wIdx] |= mask;
        if (isWalk) ++m_exploredWalkableCells;
    } else if (!explored && wasExplored) {
        m_exploredBits[wIdx] &= ~mask;
        if (isWalk && m_exploredWalkableCells > 0) --m_exploredWalkableCells;
    }
}

void NavMeshBitsetFusion::RevealFogAroundWorldPos(float worldX, float worldY, float revealRadius) {
    if (m_cols == 0 || m_rows == 0 || revealRadius <= 0.0f) return;

    const float gridRadius = revealRadius / m_gridToWorldScale;
    const int32_t rCells = static_cast<int32_t>(std::ceil(gridRadius));
    const float rSq = gridRadius * gridRadius;

    uint32_t centerGx = 0, centerGy = 0;
    if (!WorldToGrid(worldX, worldY, centerGx, centerGy)) return;

    const int32_t minX = std::max<int32_t>(0, static_cast<int32_t>(centerGx) - rCells);
    const int32_t maxX = std::min<int32_t>(static_cast<int32_t>(m_cols) - 1, static_cast<int32_t>(centerGx) + rCells);
    const int32_t minY = std::max<int32_t>(0, static_cast<int32_t>(centerGy) - rCells);
    const int32_t maxY = std::min<int32_t>(static_cast<int32_t>(m_rows) - 1, static_cast<int32_t>(centerGy) + rCells);

    for (int32_t gy = minY; gy <= maxY; ++gy) {
        const float dy = static_cast<float>(gy - static_cast<int32_t>(centerGy));
        const float dySq = dy * dy;
        for (int32_t gx = minX; gx <= maxX; ++gx) {
            const float dx = static_cast<float>(gx - static_cast<int32_t>(centerGx));
            if (dx * dx + dySq <= rSq) {
                SetExploredGrid(static_cast<uint32_t>(gx), static_cast<uint32_t>(gy), true);
            }
        }
    }
}

std::vector<GridCoord> NavMeshBitsetFusion::ExtractFrontierCells(uint32_t maxFrontiers) const {
    std::vector<GridCoord> frontiers;
    if (m_cols == 0 || m_rows == 0 || maxFrontiers == 0) return frontiers;
    frontiers.reserve(maxFrontiers);

    // Frontier: Walkable == 1 && Explored == 0, và tiếp giáp ít nhất 1 ô Explored == 1
    static const int32_t dx[4] = { 1, -1, 0, 0 };
    static const int32_t dy[4] = { 0, 0, 1, -1 };

    for (uint32_t y = 1; y + 1 < m_rows; ++y) {
        for (uint32_t x = 1; x + 1 < m_cols; ++x) {
            const size_t wIdx = WordIndex(x, y);
            const uint64_t mask = BitMask(x);

            // Kiểm tra Walkable && !Explored
            if ((m_walkableBits[wIdx] & mask) != 0ULL && (m_exploredBits[wIdx] & mask) == 0ULL) {
                // Kiểm tra xem có ô lân cận nào đã khám phá không
                bool hasExploredNeighbor = false;
                for (int i = 0; i < 4; ++i) {
                    const uint32_t nx = x + dx[i];
                    const uint32_t ny = y + dy[i];
                    if (IsExplored(nx, ny)) {
                        hasExploredNeighbor = true;
                        break;
                    }
                }
                if (hasExploredNeighbor) {
                    frontiers.push_back({ static_cast<int32_t>(x), static_cast<int32_t>(y) });
                    if (frontiers.size() >= maxFrontiers) return frontiers;
                }
            }
        }
    }

    return frontiers;
}

std::vector<Vec2> NavMeshBitsetFusion::FindGlobalPath(
    const Vec2& startWorld,
    const Vec2& goalWorld,
    float* outComputeTimeMs
) const {
    const auto t0 = std::chrono::high_resolution_clock::now();
    std::vector<Vec2> path;

    uint32_t sgx = 0, sgy = 0, ggx = 0, ggy = 0;
    if (!WorldToGrid(startWorld.x, startWorld.y, sgx, sgy) ||
        !WorldToGrid(goalWorld.x, goalWorld.y, ggx, ggy)) {
        if (outComputeTimeMs) *outComputeTimeMs = 0.0f;
        return path;
    }

    if (sgx == ggx && sgy == ggy) {
        path.push_back(startWorld);
        path.push_back(goalWorld);
        if (outComputeTimeMs) *outComputeTimeMs = 0.01f;
        return path;
    }

    // A* trên lưới Bitset tối ưu O(1) walkability lookup và mảng phẳng cache-friendly
    struct Node {
        uint32_t x, y;
        float g, f;
        bool operator>(const Node& o) const { return f > o.f; }
    };

    auto Heuristic = [](uint32_t x1, uint32_t y1, uint32_t x2, uint32_t y2) -> float {
        const float ddx = static_cast<float>(x1) - static_cast<float>(x2);
        const float ddy = static_cast<float>(y1) - static_cast<float>(y2);
        return std::sqrt(ddx * ddx + ddy * ddy);
    };

    const size_t totalCells = static_cast<size_t>(m_cols) * m_rows;
    std::vector<float> gScore(totalCells, 1e9f);
    std::vector<uint32_t> cameFrom(totalCells, UINT32_MAX);

    auto CellIdx = [this](uint32_t x, uint32_t y) -> size_t {
        return static_cast<size_t>(y) * m_cols + x;
    };

    std::priority_queue<Node, std::vector<Node>, std::greater<Node>> openSet;

    const size_t startIdx = CellIdx(sgx, sgy);
    const size_t goalIdx = CellIdx(ggx, ggy);
    gScore[startIdx] = 0.0f;
    openSet.push({ sgx, sgy, 0.0f, Heuristic(sgx, sgy, ggx, ggy) });

    static const int32_t dirs[8][2] = {
        { 1, 0 }, { -1, 0 }, { 0, 1 }, { 0, -1 },
        { 1, 1 }, { 1, -1 }, { -1, 1 }, { -1, -1 }
    };
    static const float stepCosts[8] = {
        1.0f, 1.0f, 1.0f, 1.0f,
        1.4142f, 1.4142f, 1.4142f, 1.4142f
    };

    uint32_t nodesExplored = 0;
    const uint32_t maxNodes = 6000;
    bool found = false;

    while (!openSet.empty() && nodesExplored++ < maxNodes) {
        const Node current = openSet.top();
        openSet.pop();

        if (current.x == ggx && current.y == ggy) {
            found = true;
            break;
        }

        const size_t curIdx = CellIdx(current.x, current.y);
        if (current.g > gScore[curIdx]) continue;

        for (int i = 0; i < 8; ++i) {
            const int32_t nx = static_cast<int32_t>(current.x) + dirs[i][0];
            const int32_t ny = static_cast<int32_t>(current.y) + dirs[i][1];
            if (nx < 0 || nx >= static_cast<int32_t>(m_cols) ||
                ny < 0 || ny >= static_cast<int32_t>(m_rows)) continue;

            const uint32_t unx = static_cast<uint32_t>(nx);
            const uint32_t uny = static_cast<uint32_t>(ny);

            // Kiểm tra vật cản trực tiếp trên RAM bitset O(1)
            if (!IsWalkable(unx, uny)) continue;

            // Kiểm tra đi chéo không xuyên góc tường (Corner Cutting Protection)
            if (i >= 4) {
                if (!IsWalkable(current.x, uny) || !IsWalkable(unx, current.y)) {
                    continue;
                }
            }

            const float tentativeG = current.g + stepCosts[i];
            const size_t neighborIdx = CellIdx(unx, uny);

            if (tentativeG < gScore[neighborIdx]) {
                cameFrom[neighborIdx] = static_cast<uint32_t>(curIdx);
                gScore[neighborIdx] = tentativeG;
                const float f = tentativeG + Heuristic(unx, uny, ggx, ggy);
                openSet.push({ unx, uny, tentativeG, f });
            }
        }
    }

    if (found) {
        size_t curr = goalIdx;
        std::vector<Vec2> rawGridPath;
        while (curr != startIdx && curr != UINT32_MAX) {
            const uint32_t cx = static_cast<uint32_t>(curr % m_cols);
            const uint32_t cy = static_cast<uint32_t>(curr / m_cols);
            float wx = 0.0f, wy = 0.0f;
            GridToWorld(cx, cy, wx, wy);
            rawGridPath.push_back({ wx, wy });
            curr = cameFrom[curr];
        }
        rawGridPath.push_back(startWorld);
        std::reverse(rawGridPath.begin(), rawGridPath.end());
        path = std::move(rawGridPath);
    }

    const auto t1 = std::chrono::high_resolution_clock::now();
    const float durationMs = std::chrono::duration<float, std::milli>(t1 - t0).count();
    if (outComputeTimeMs) *outComputeTimeMs = durationMs;

    return path;
}

bool NavMeshBitsetFusion::FindBestFrontierTarget(
    const Vec2& playerWorld,
    Vec2& outTargetWorld,
    float* outScore
) const {
    const auto frontiers = ExtractFrontierCells(128);
    if (frontiers.empty()) return false;

    float bestUtility = -1e9f;
    GridCoord bestCell{ 0, 0 };

    for (const auto& f : frontiers) {
        float fx = 0.0f, fy = 0.0f;
        GridToWorld(static_cast<uint32_t>(f.x), static_cast<uint32_t>(f.y), fx, fy);
        const float dist = playerWorld.Distance({ fx, fy });

        // Information gain ước tính: khoảng cách càng gần và chưa đi qua nhiều càng ưu tiên
        // Utility = 1000.0 / (dist + 50.0)
        const float utility = 1000.0f / (dist + 40.0f);
        if (utility > bestUtility) {
            bestUtility = utility;
            bestCell = f;
        }
    }

    GridToWorld(static_cast<uint32_t>(bestCell.x), static_cast<uint32_t>(bestCell.y), outTargetWorld.x, outTargetWorld.y);
    if (outScore) *outScore = bestUtility;
    return true;
}

float NavMeshBitsetFusion::ExplorationProgress() const {
    if (m_totalWalkableCells == 0) return 0.0f;
    return static_cast<float>(m_exploredWalkableCells) / static_cast<float>(m_totalWalkableCells);
}

} // namespace navigation
