/*
 * Copyright (c) 2024 Huawei Technologies Co., Ltd.
 * This file is a part of the CANN Open Software.
 * Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
#include "tilexr_comm.h"
#include "tilexr_internal.h"

#include <chrono>
#include <vector>
#include <mutex>
#include <map>
#include <set>
#include <thread>
#include <sstream>
#include <iomanip>

#include <hccl/hccl.h>
#include "mki/utils/log/log.h"
#include "mki/utils/env/env.h"
#include "tools/socket/tilexr_sock_exchange.h"

#include "runtime/kernel.h"
#include "runtime/mem.h"
#include "runtime/dev.h"
#include "runtime/rt_ffts.h"

#include <shmem/include/shmem.h>
#include <shmem/include/host_device/shmem_common_types.h>

enum TopologyType : int {
    TOPOLOGY_HCCS = 0,
    TOPOLOGY_PIX,
    TOPOLOGY_PIB,
    TOPOLOGY_PHB,
    TOPOLOGY_SYS,
    TOPOLOGY_SIO,
    TOPOLOGY_HCCS_SW
};

using namespace std;
using namespace chrono;
using namespace Mki;

namespace TileXR {
constexpr int HCCL_IPC_PID_ARRAY_SIZE = 1; // 固定每次只传一个PID数据
constexpr int TILEXR_INIT_TIMEOUT = 600;
constexpr size_t TILEXR_SHMEM_MIN_MEM = 1 * 1024 * 1024;  // 1 MB,仅供 shmem 内部 sync

static map<string, GM_ADDR [TILEXR_MAX_RANK_SIZE]> g_localPeerMemMap;
static map<string, int[TILEXR_MAX_RANK_SIZE]> g_devList;
static std::mutex g_mtx;


// 如果是互联的链路,返回false; 对910B2C那些不互联的链路,返回true
bool SkipUnusedChannel910B2C(int curRank, int peerRank, ChipName chipName)
{
    if (chipName == ChipName::CHIP_910B2C) {
        constexpr int rankSizePerNode = 8;
        // 双节点16P中不用的链路: 不在同一个节点 且rank在节点内序号不同; 在调用时将跳过
        if ((curRank / rankSizePerNode != peerRank / rankSizePerNode)
            && (std::abs(curRank - peerRank) != rankSizePerNode)) {
            return true;
        }
    }
    return false;
}

int TileXRComm::InitDumpAddr()
{
    constexpr uint32_t dumpCoreCnt = 75;
    constexpr uint32_t dumpSizePerCore = 1 * 1024 * 1024;
    constexpr uint32_t dumpWorkspaceSize = dumpCoreCnt * dumpSizePerCore;
    GM_ADDR dumpAddr = nullptr;
    int ret = 0;
    ret = aclrtMalloc(reinterpret_cast<void **>(&dumpAddr), dumpWorkspaceSize, ACL_MEM_MALLOC_HUGE_FIRST);
    if (ret != ACL_SUCCESS) {
        MKI_LOG(ERROR) << "aclrtMalloc err " << __LINE__ << " " << ret;
        return TILEXR_ERROR_INTERNAL;
    }
    aclrtMemset(dumpAddr, dumpWorkspaceSize, 0, dumpWorkspaceSize);
 
    GM_ADDR memory = static_cast<GM_ADDR>(std::malloc(dumpWorkspaceSize));
    if (!memory) {
        MKI_LOG(ERROR) << "std::malloc err " << __LINE__;
        return TILEXR_ERROR_INTERNAL;
    }
    // Zero out allocated memory
    for (size_t i = 0; i < dumpWorkspaceSize; ++i) {
        ((uint8_t*)memory)[i] = 0;
    }
    // 遍历每个block进行初始化
        for (uint32_t i = 0; i < dumpCoreCnt; ++i) {
        // 计算当前block的起始地址
        GM_ADDR blockStart = memory + i * dumpSizePerCore;
        GM_ADDR deviceBlockStart = dumpAddr + i * dumpSizePerCore;
        
        // 初始化BlockInfo
        LcclDumpBlockInfo* blockInfo = reinterpret_cast<LcclDumpBlockInfo*>(blockStart);
        blockInfo->len = dumpSizePerCore;
        blockInfo->core = i;
        blockInfo->blockNum = 0;
        blockInfo->dumpOffset = dumpSizePerCore - sizeof(LcclDumpBlockInfo);
        blockInfo->magic = 0; // 示例魔法值
        blockInfo->dumpAddr = reinterpret_cast<uint64_t>(deviceBlockStart + sizeof(LcclDumpBlockInfo));
    }
 
    ret = aclrtMemcpy(dumpAddr, dumpWorkspaceSize, memory, dumpWorkspaceSize, ACL_MEMCPY_HOST_TO_DEVICE);
    if (ret != ACL_SUCCESS) {
        MKI_LOG(ERROR) << "aclrtMemcpy err " << __LINE__ << " " << ret;
        return TILEXR_ERROR_INTERNAL;
    }
    std::free(memory);
 
    commArgs_.dumpAddr = dumpAddr;
    return TILEXR_SUCCESS;
}

int TileXRComm::InitUDMA()
{
    // Step 1: rank 0 生成 shmem UID
    aclshmemx_uniqueid_t shmemUid;
    if (rank_ == 0) {
        int ret = aclshmemx_get_uniqueid(&shmemUid);
        if (ret != ACLSHMEM_SUCCESS) {
            MKI_LOG(WARN) << "aclshmemx_get_uniqueid failed: " << ret << ", UDMA disabled";
            return TILEXR_SUCCESS;  // 优雅降级
        }
    }

    // Step 2: AllGather UID,所有 rank 获得相同 UID
    int ret = socketExchange_->AllGather(
        reinterpret_cast<char*>(&shmemUid),
        sizeof(aclshmemx_uniqueid_t),
        reinterpret_cast<char*>(&shmemUid)
    );
    if (ret != TILEXR_SUCCESS) {
        MKI_LOG(WARN) << "AllGather shmem UID failed: " << ret << ", UDMA disabled";
        return TILEXR_SUCCESS;  // 优雅降级
    }

    // Step 3: 设置 shmem 初始化属性,指定 UDMA 引擎
    aclshmemx_init_attr_t shmemAttr;
    ret = aclshmemx_set_attr_uniqueid_args(
        rank_,
        rankSize_,
        TILEXR_SHMEM_MIN_MEM,
        &shmemUid,
        &shmemAttr
    );
    if (ret != ACLSHMEM_SUCCESS) {
        MKI_LOG(WARN) << "aclshmemx_set_attr_uniqueid_args failed: " << ret << ", UDMA disabled";
        return TILEXR_SUCCESS;  // 优雅降级
    }

    // 设置 UDMA 引擎
    shmemAttr.option_attr.data_op_engine_type = ACLSHMEM_DATA_OP_UDMA;

    // Step 4: 初始化 shmem,失败时优雅降级
    ret = aclshmemx_init_attr(ACLSHMEMX_INIT_WITH_UNIQUEID, &shmemAttr);
    if (ret != ACLSHMEM_SUCCESS) {
        MKI_LOG(WARN) << "aclshmemx_init_attr failed: " << ret << ", UDMA disabled";
        return TILEXR_SUCCESS;  // 优雅降级
    }

    // Step 5: 从 shmem 获取 UDMA 信息(设备侧指针)
    void* udmaInfoPtr = nullptr;
    size_t udmaInfoSize = 0;
    ret = aclshmemx_get_udma_info(&udmaInfoPtr, &udmaInfoSize);
    if (ret != ACLSHMEM_SUCCESS || udmaInfoPtr == nullptr) {
        MKI_LOG(WARN) << "aclshmemx_get_udma_info failed: " << ret << ", UDMA disabled";
        aclshmem_finalize();
        return TILEXR_SUCCESS;  // 优雅降级
    }

    // Step 6: 直接使用 shmem 返回的设备侧指针
    // 注意:udmaInfoPtr 已经是设备内存地址,无需再次拷贝
    udmaInfoDev_ = reinterpret_cast<uint8_t*>(udmaInfoPtr);

    // Step 7: 设置 commArgs_ 字段
    commArgs_.udmaInfoPtr = udmaInfoDev_;
    commArgs_.extraFlag |= ExtraFlag::UDMA;

    MKI_LOG(INFO) << "InitUDMA success, rank " << rank_ << "/" << rankSize_;
    return TILEXR_SUCCESS;
}

int TileXRComm::SyncCommArgs()
{
    commArgs_.rank = rank_;
    commArgs_.localRank = localRank_;
    commArgs_.rankSize = rankSize_;
    commArgs_.localRankSize = localRankSize_;
    for (int i = 0; i < rankSize_; ++i) {
        commArgs_.peerMems[i] = peerMem_[i];    // 这里不会越界,之前有逻辑校验过越界了
    }

    if (isEnableMsprofOp_) {
        if (InitDumpAddr() != TILEXR_SUCCESS) {
            return TILEXR_ERROR_INTERNAL;
        }

        uint64_t fftsVal = 0;
        uint32_t fftsLen = 0;
        int error = rtGetC2cCtrlAddr(&fftsVal, &fftsLen);
        if (error != RT_ERROR_NONE) {
            MKI_LOG(ERROR) << "rtGetC2cCtrlAddr err:" << error;
            return TILEXR_ERROR_MKIRT;
        }
        commArgs_.fftsVal = fftsVal;
    }

    int ret = 0;
    ret = aclrtMalloc(reinterpret_cast<void **>(&commArgsPtr_), sizeof(commArgs_), ACL_MEM_MALLOC_HUGE_FIRST);
    if (ret != ACL_SUCCESS) {
        MKI_LOG(ERROR) << "aclrtMalloc err " << __LINE__ << " " << ret;
        return TILEXR_ERROR_INTERNAL;
    }
    ret = aclrtMemcpy(commArgsPtr_, sizeof(commArgs_), &commArgs_, sizeof(commArgs_), ACL_MEMCPY_HOST_TO_DEVICE);
    if (ret != ACL_SUCCESS) {
        MKI_LOG(ERROR) << "aclrtMemcpy err " << __LINE__ << " " << ret;
        return TILEXR_ERROR_INTERNAL;
    }
    return TILEXR_SUCCESS;
}

int TileXRComm::InitCommon()
{
    // enable peer device
    if (EnablePeerAccess() != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "EnablePeerAccess failed!";
        return TILEXR_ERROR_INTERNAL;
    }
    const char *lcclDeterministic = Mki::GetEnv("LCCL_DETERMINISTIC");
    if (lcclDeterministic && (string(lcclDeterministic) == "1" || string(lcclDeterministic) == "true")) {
        commArgs_.extraFlag |= ExtraFlag::DETERMINISTIC;
    }
    if (GetChipName() == ChipName::CHIP_910B2C) {
        commArgs_.extraFlag |= ExtraFlag::TOPO_910B2C;
    }
    if (GetChipName() >= ChipName::CHIP_910_9391) {
        commArgs_.extraFlag |= ExtraFlag::TOPO_910_93;
    }
    if (GetChipName() > ChipName::CHIP_910_9362) {
        commArgs_.extraFlag |= ExtraFlag::TOPO_910A5;
    }
    constexpr uint32_t AI_CORE_NUM_20 = 20;
    if (GetCoreNum(GetChipName()) > AI_CORE_NUM_20) {
        commArgs_.extraFlag |= ExtraFlag::IS_GREATER_THAN_40_AIV;
    }

    // RegistKernel(isEnableMsprofOp_);

    localRank_ = rank_ % localRankSize_;
    return TILEXR_SUCCESS;
}

void TileXRComm::CloseIpcMem()
{
    for (int i = 0; i < rankSize_; ++i) {
        if (i == rank_ || peerMem_[i] == nullptr) {
            continue;
        }
        int ret = rtIpcCloseMemory(static_cast<void *>(peerMem_[i]));
        if (ret != RT_ERROR_NONE) {
            MKI_LOG(WARN) << "Close ipc[" << i << "] memory failed! ret: " << ret;
        }
        peerMem_[i] = nullptr;
    }
}

void TileXRComm::FreePeerMem(GM_ADDR &mem) const
{
    if (mem != nullptr) {
        aclError aclRet = aclrtFree(mem);
        if (aclRet != ACL_SUCCESS) {
            MKI_LOG(ERROR) << "Free share memory failed! ret: " << aclRet;
        }
    }
    mem = nullptr;
}

int TileXRComm::Init()
{
    if (inited_) {
        return TILEXR_SUCCESS;
    }
    if (rank_ < 0 || rank_ >= rankSize_ || rankSize_ <= 0 || rankSize_ > TILEXR_MAX_RANK_SIZE) {
        MKI_LOG(ERROR) << "The rank is invalid! rank:" << rank_ << " rankSize:" << rankSize_;
        return TILEXR_ERROR_PARA_CHECK_FAIL;
    }
    if (TileXRSockExchange::CheckValid(commId_)) {
        socketExchange_ = new (nothrow) TileXRSockExchange(rank_, rankSize_, commId_);
    } else {
        socketExchange_ = new (nothrow) TileXRSockExchange(rank_, rankSize_, commDomain_);
    }
    if (socketExchange_ == nullptr) {
        MKI_LOG(ERROR) << "TileXRSockExchange create failed. rank : " << rank_ << " rankSize:" << rankSize_;
        return TILEXR_ERROR_INTERNAL;
    }
    int ret = GetDev();
    if (ret != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "init context failed! ret: " << ret;
        return ret;
    }

    MKI_LOG(INFO) << "rank " << rank_ << "/" << rankSize_ << " running devId:" << devId_;

    if (InitCommon() != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "init common failed!";
        return TILEXR_ERROR_INTERNAL;
    }

    MKI_LOG(DEBUG) << "Prepare to InitCommMem localRankSize_ -> " << localRankSize_ << ", localRank_ -> " << localRank_;
    if (InitCommMem() != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "InitCommMem failed!";
        return TILEXR_ERROR_INTERNAL;
    }
    MKI_LOG(DEBUG) << "InitCommMem " << rank_ << "/" << rankSize_ << ", localRank_ : " << localRank_ <<
            ", localRankSize_ : " << localRankSize_ << " success";

    // 新增:初始化 UDMA
    ret = InitUDMA();
    if (ret != TILEXR_SUCCESS) {
        return ret;
    }

    // set comm args in device.
    SyncCommArgs();
    MKI_LOG(INFO) << "TileXRCommInit " << rank_ << "/" << rankSize_ << " success. extraFlag:" << commArgs_.extraFlag <<
        " commArgs_.localRank : " << commArgs_.localRank << " commArgs_.localRankSize : " << commArgs_.localRankSize;
    inited_ = true;
    delete socketExchange_; // socketExchange_ 不会为空
    socketExchange_ = nullptr;
    return TILEXR_SUCCESS;
}

int TileXRComm::InitThread(const std::string &uid)
{
    if (inited_) {
        return TILEXR_SUCCESS;
    }
    if (rank_ < 0 || rank_ >= rankSize_ || rankSize_ <= 0 || rankSize_ > TILEXR_MAX_RANK_SIZE) {
        MKI_LOG(ERROR) << "The rank is invalid! rank:" << rank_ << "rankSize:" << rankSize_;
        return TILEXR_ERROR_PARA_CHECK_FAIL;
    }
    if (GetDevThread(uid) != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "get devs failed.";
        return TILEXR_ERROR_INTERNAL;
    }
    MKI_LOG(INFO) << "rank " << rank_ << "/" << rankSize_ << " running devId:" << devId_ << "uid: " << uid;

    if (InitCommon() != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "init common failed!";
        return TILEXR_ERROR_INTERNAL;
    }
    {
        lock_guard<mutex> lock(g_mtx);
        if (g_localPeerMemMap.find(uid) == g_localPeerMemMap.end()) {
            for (int i = 0; i < rankSize_; ++i) {
                g_localPeerMemMap[uid][i] = nullptr;
            }
        }
        uid_ = uid;
    }
    InitMem();
    g_localPeerMemMap[uid][rank_] = peerMem_[rank_];

    auto start = high_resolution_clock::now();
    for (int i = 0; i < rankSize_; ++i) {
        while (g_localPeerMemMap[uid][i] == nullptr) { // check other threads
            this_thread::sleep_for(1ms);
            auto elapsed = duration_cast<seconds>(high_resolution_clock::now() - start);
            if (elapsed.count() > TILEXR_INIT_TIMEOUT) {
                MKI_LOG(ERROR) << "Lccl Init timeout!";
                FreePeerMem(g_localPeerMemMap[uid][rank_]);
                return TILEXR_ERROR_TIMEOUT;
            }
        }
        peerMem_[i] = g_localPeerMemMap[uid][i];
    }
    localRank_ = rank_;
    localRankSize_ = rankSize_;

    // 注意:InitThread 为单进程多线程模式,不支持 UDMA(需要 socketExchange_ 进行跨进程协调)
    // UDMA 主要用于跨进程/跨节点通信,线程模式使用进程内共享内存即可
    MKI_LOG(DEBUG) << "Thread mode: UDMA initialization skipped (single-process multi-thread scenario)";

    SyncCommArgs();
    MKI_LOG(INFO) << "Lccl init multi thread " << rank_ << "/" << rankSize_ << " success, uid:" << uid;
    inited_ = true;
    return TILEXR_SUCCESS;
}

/**
 * @brief 函数内部会有检测,是否需要进行 aclrtDeviceEnablePeerAccess,如果芯片为310P且是HCCS链路,则不调用此函数。
 *
 *
 */
int TileXRComm::EnablePeerAccess()
{
    physicalInfo_.chipName = GetChipName();
    for (auto &dev : devList_) {
        if (devId_ == dev) {
            continue;
        }
        // 处理910B2C 16卡通信的特例
        if (SkipUnusedChannel910B2C(dev, devId_, GetChipName())) {
            continue;
        }

        int64_t value = 0;
        if (rtGetPairDevicesInfo(devId_, dev, 0, &value) != RT_ERROR_NONE) {
            MKI_LOG(WARN) << devId_ << " & " << dev << " pair devices info failed to get";
        } else {
            MKI_LOG(DEBUG) << devId_ << " <-----> " << dev << ", halGetPairDevicesInfo: *value = " << value;
        }

        // 如果310P未来通信域要支持两卡四芯的话,这里需要做更改。并且现在默认服务器上机器只有一个链路种类。
        if (value == TOPOLOGY_HCCS || value == TOPOLOGY_SIO || value == TOPOLOGY_HCCS_SW ||
            GetChipName() == ChipName::CHIP_910B2C) {
            physicalInfo_.physicalLink = PhysicalLink::HCCS;
            commArgs_.extraFlag &= ~(ExtraFlag::TOPO_PCIE);
        } else if (physicalInfo_.physicalLink == PhysicalLink::RESERVED) {
            physicalInfo_.physicalLink = PhysicalLink::PCIE;
            commArgs_.extraFlag |= ExtraFlag::TOPO_PCIE;
            if (rankSize_ > PING_PONG_SIZE) {
                MKI_LOG(ERROR) << "do not support pcie > 2 rank! rankSize_ = " << rankSize_;
                return TILEXR_ERROR_INTERNAL;
            }
        }

        physicalInfo_.coreNum = GetCoreNum(physicalInfo_.chipName);

        // value里的0实际上对应驱动枚举类的 TOPOLOGY_HCCS
        if (physicalInfo_.chipName == ChipName::CHIP_310P3 && value == 0) {
            MKI_LOG(WARN) << "warn aclrtDeviceEnablePeerAccess is skipped! peerDeviceId = " << dev;
            continue;
        }

        aclError ret = aclrtDeviceEnablePeerAccess(dev, 0);
        if (ret != ACL_SUCCESS) {
            MKI_LOG(ERROR) << "err aclrtDeviceEnablePeerAccess failed peerDeviceId = " << dev << " ,rank = " << rank_
                           << ", value = " << value << ", flags = " << 0 << "," << __LINE__ << ": " << ret;
            return TILEXR_ERROR_INTERNAL;
        }
    }
    MKI_LOG(DEBUG) << "EnablePeerAccess succeed" << rank_;
    return TILEXR_SUCCESS;
}

int TileXRComm::GetDev()
{
    // 这里这个nodeNum可以理解为Y轴长度,手动控制的话将这个拦截修改即可。
    int nodeNum = socketExchange_->GetNodeNum();
    if (nodeNum <= 0 || nodeNum > rankSize_) {
        MKI_LOG(ERROR) << "error! node num : " << nodeNum << " rank size: " << rankSize_;
        return TILEXR_ERROR_INTERNAL;
    }
    localRankSize_ = rankSize_ < 0 ? 0 : rankSize_ / nodeNum;
    localRank_ = rank_ % localRankSize_;
    MKI_LOG(DEBUG) << "GetDev : localRankSize_ : " << localRankSize_ << " localRank_: " << localRank_
                    << "  rank :" << rank_ << "   rankSize :" << rankSize_;
    devList_.resize(rankSize_);
    // get current id and broadcast
    aclError aclRet = aclrtGetDevice(&devId_);
    if (aclRet != ACL_SUCCESS) {
        MKI_LOG(ERROR) << "aclrtGetDevice error! ret: " << aclRet;
        return TILEXR_ERROR_INTERNAL;
    }
    // get other rank dev id, put into devList_
    int ret = socketExchange_->AllGather(&devId_, 1, devList_.data());
    if (ret != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "TileXRSockExchange AllGather error! ret: " << ret;
        return TILEXR_ERROR_INTERNAL;
    }
    std::string devIdStr = "";
    for (int i = 0; i < rankSize_; ++i) {
        devIdStr += (i == 0 ? "" : ", ");
        devIdStr += to_string(devList_[i]);
    }
    MKI_LOG(DEBUG) << "rank " << rank_ << " devId: " << devId_ << ", otherDevList : " << devIdStr;
    MKI_LOG(INFO) << "AllGather: Get other rank dev id success";
    return TILEXR_SUCCESS;
}

int TileXRComm::GetDevThread(const std::string &uid)
{
    devList_.resize(rankSize_);
    // get current id and broadcast
    aclError aclRet = aclrtGetDevice(&devId_);
    if (aclRet != ACL_SUCCESS) {
        MKI_LOG(ERROR) << "aclrtGetDevice error! ret: " << aclRet;
        return TILEXR_ERROR_INTERNAL;
    }
    {
        std::lock_guard<std::mutex> lock(g_mtx);
        if (g_devList.find(uid) == g_devList.end()) {
            for (int i = 0; i < rankSize_; ++i) {
                g_devList[uid][i] = 0;
            }
        }
    }
    g_devList[uid][rank_] = devId_ + 1; // 0 is invalid
    auto start = high_resolution_clock::now();
    for (int i = 0; i < rankSize_; ++i) {
        while (g_devList[uid][i] == 0) { // check other threads
            this_thread::sleep_for(1ms);
            auto elapsed = duration_cast<seconds>(high_resolution_clock::now() - start);
            if (elapsed.count() > TILEXR_INIT_TIMEOUT) {
                MKI_LOG(ERROR) << "Lccl Init timeout!";
                return TILEXR_ERROR_TIMEOUT;
            }
        }
        devList_.at(i) = g_devList[uid][i] - 1;
    }
    return TILEXR_SUCCESS;
}

int TileXRComm::InitMem()
{
    // 申请并初始化IpcBuff
    constexpr int32_t bufferSizeUint = 1024 * 1024;
    int tilexrBuffSize = bufferSize_ * bufferSizeUint + TILEXR_FLAG_BUFF_BYTES;

    MKI_LOG(DEBUG) << "tilexr buffer size " << tilexrBuffSize;
    aclError ret = aclrtMalloc(
        reinterpret_cast<void **>(&peerMem_[rank_]), tilexrBuffSize,
        (GetChipName() == ChipName::CHIP_310P3) ? ACL_MEM_MALLOC_HUGE_FIRST_P2P : ACL_MEM_MALLOC_HUGE_FIRST);
    if (ret != ACL_SUCCESS) {
        MKI_LOG(ERROR) << "allocate device mem error " << __FILE__ << ":" << __LINE__ << " " << ret;
        return TILEXR_ERROR_INTERNAL;
    }
    MKI_LOG(DEBUG) << "peerMem[rank" << rank_ << "], allocate finished.";
    aclrtMemset(peerMem_[rank_], tilexrBuffSize, 0, tilexrBuffSize);
    return TILEXR_SUCCESS;
}

int TileXRComm::GetPid(uint32_t *pids)
{
    if (rtDeviceGetBareTgid(&pids[rank_]) != RT_ERROR_NONE) {  // 获取docker外的进程id,bare指docker�?        MKI_LOG(ERROR) << "DeviceGetBareTgid err " << __LINE__;
        return TILEXR_ERROR_INTERNAL;
    }
    int ret = socketExchange_->AllGather(&pids[rank_], 1, pids);
    if (ret != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "TileXRSockExchange AllGather error! ret: " << ret;
        return ret;
    }
    for (int i = 0; i < rankSize_; ++i) {
        MKI_LOG(DEBUG) << "rank : " << rank_ << ", otherRank : " << i << " pid[" << i << "]: " << pids[i];
    }
    MKI_LOG(DEBUG) << "AllGather: Get other rank pid";
    return TILEXR_SUCCESS;
}

int TileXRComm::GetSidId(int64_t sdids[TILEXR_MAX_RANK_SIZE], int rankSize)
{
    if (rank_ >= rankSize) {
        MKI_LOG(ERROR) << "TileXRComm::GetSidId err rank_ >= rankSize " << rank_ << ">=" << rankSize;
        return TILEXR_ERROR_INTERNAL;
    }
    if ((physicalInfo_.chipName >= ChipName::CHIP_910_9391) && (physicalInfo_.chipName < ChipName::RESERVED)) {
        const int rtModuleTypeSystem = 0;
        const int infoTypeSdid = 26;
        if (rtGetDeviceInfo(devList_[rank_], rtModuleTypeSystem, infoTypeSdid, &sdids[rank_]) != RT_ERROR_NONE) {
            MKI_LOG(ERROR) << "DeviceGetDeviceInfo err " << __LINE__;
            return TILEXR_ERROR_INTERNAL;
        }
        MKI_LOG(DEBUG) << "rank " << rank_ << " dev id: " << devList_[rank_]
                       << " rtGetDeviceInfo sdid: " << sdids[rank_];

        int ret = socketExchange_->AllGather(&sdids[rank_], 1, sdids);
        if (ret != TILEXR_SUCCESS) {
            MKI_LOG(ERROR) << "TileXRSockExchange AllGather error! ret: " << ret;
            return ret;
        }
        for (int i = 0; i < rankSize_; ++i) {
            MKI_LOG(DEBUG) << "rank " << i << " sdid: " << sdids[i];
        }
        MKI_LOG(DEBUG) << "AllGather: Get other rank sdid";
    }
    return TILEXR_SUCCESS;
}

int TileXRComm::GetName(string &name, char names[TILEXR_MAX_RANK_SIZE][IPC_NAME_SIZE]) const
{
    int ret = socketExchange_->AllGather<char>(name.c_str(), IPC_NAME_SIZE, names[0]);
    if (ret != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "TileXRSockExchange AllGather error! ret: " << ret;
        return TILEXR_ERROR_INTERNAL;
    }
    for (int i = 0; i < rankSize_; ++i) {
        MKI_LOG(DEBUG) << "rank " << i << " mem name: " << names[i];
    }
    MKI_LOG(DEBUG) << "AllGather: Get other rank mem name";
    return TILEXR_SUCCESS;
}

int TileXRComm::InitCommMem()
{
    int ret = InitMem();
    if (ret != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "InitMem error! ret: " << ret;
        return ret;
    }

    // 获取所有进程的pid
    uint32_t pids[TILEXR_MAX_RANK_SIZE] = {0};
    ret = GetPid(pids);
    if (ret != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "GetPid error! ret: " << ret;
        return ret;
    }

    // 获取所有进程的sdid
    int64_t sdids[TILEXR_MAX_RANK_SIZE] = {0};
    ret = GetSidId(sdids, rankSize_);
    if (ret != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "GetSidId error! ret: " << ret;
        return ret;
    }

    // 获取所有进程的mem name
    string name;
    if (SetMemoryName(name) != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "SetMemoryName err ";
        return TILEXR_ERROR_INTERNAL;
    }

    if (SetIpcPidSdid(name, pids, sdids) != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "SetIpcPidSdid failed!";
        return TILEXR_ERROR_INTERNAL;
    }

    MKI_LOG(DEBUG) << "rank " << rank_ << " mem name: " << name << " name len: " << name.size();
    char names[TILEXR_MAX_RANK_SIZE][IPC_NAME_SIZE];
    name.resize(IPC_NAME_SIZE);
    ret = GetName(name, names);
    if (ret != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "GetName error! ret: " << ret;
        return ret;
    }

    if (OpenIpcMem(names) != TILEXR_SUCCESS) {
        MKI_LOG(ERROR) << "rank: " << rank_ << " OpenIpcMem failed!";
        return TILEXR_ERROR_INTERNAL;
    }
    return TILEXR_SUCCESS;
}

int TileXRComm::OpenIpcMem(const char names[TILEXR_MAX_RANK_SIZE][IPC_NAME_SIZE])
{
    static mutex mut;
    lock_guard<mutex> lock(mut);
    for (int i = 0; i < rankSize_; ++i) {
        if (i == rank_) {
            continue;
        }
        // 处理910B2C 16卡通信的特例
        if (SkipUnusedChannel910B2C(rank_, i, GetChipName())) {
            continue;
        }
        int ret = rtIpcOpenMemory(reinterpret_cast<void **>(&peerMem_[i]), names[i]);
        if (ret != RT_ERROR_NONE) {
            CloseIpcMem();
            MKI_LOG(ERROR) << "rank : " << rank_ << " localRank : " << localRank_ << " peerMem: " << i <<
                " IpcOpenMemory err " << ret;
            return TILEXR_ERROR_INTERNAL;
        }
    }
    ipcMemInited_ = true;
    return TILEXR_SUCCESS;
}

int TileXRComm::SetMemoryName(string &name)
{
    char nameModified[IPC_NAME_SIZE] = {};
    int memRank = rank_;
    constexpr int32_t bufferSizeUint = 1024 * 1024;
    int tilexrBuffSize = bufferSize_ * bufferSizeUint + TILEXR_FLAG_BUFF_BYTES;
    if (rtIpcSetMemoryName(peerMem_[memRank], tilexrBuffSize, nameModified, IPC_NAME_SIZE) != RT_ERROR_NONE) {
        return TILEXR_ERROR_INTERNAL;
    }
    name = nameModified;
    return TILEXR_SUCCESS;
}

int TileXRComm::SetIpcPidSdid(string &name, const uint32_t *pids, const int64_t *sdids) const
{
    for (int i = 0; i < rankSize_; ++i) {
        if (i == rank_) {
            continue;
        }

        if (physicalInfo_.chipName < ChipName::CHIP_910_9391) {
            // 910B
            int32_t pidInt32 = pids[i];
            int rtRet = rtSetIpcMemPid(name.c_str(), &pidInt32, HCCL_IPC_PID_ARRAY_SIZE);
            if (rtRet != RT_ERROR_NONE) {
                MKI_LOG(ERROR) << "err " << rtRet;
                return TILEXR_ERROR_INTERNAL;
            }
        } else {
            // 910A3
            int32_t pidInt32 = pids[i];
            int rtRet = rtSetIpcMemorySuperPodPid(name.c_str(), sdids[i], &pidInt32, HCCL_IPC_PID_ARRAY_SIZE);
            if (rtRet != RT_ERROR_NONE) {
                MKI_LOG(ERROR) << "err " << rtRet;
                return TILEXR_ERROR_INTERNAL;
            }
        }
    }
    return TILEXR_SUCCESS;
}

TileXRComm::~TileXRComm()
{
    {
        lock_guard<mutex> lock(g_mtx);
        if (g_localPeerMemMap.find(uid_) != g_localPeerMemMap.end()) {
            g_localPeerMemMap.erase(uid_);
        }
    }
    if (ipcMemInited_) {
        CloseIpcMem();
        ipcMemInited_ = false;
    }
    if (socketExchange_) {
        delete socketExchange_;
        socketExchange_ = nullptr;
    }
    FreePeerMem(commArgs_.dumpAddr);
    FreePeerMem(peerMem_[rank_]);
    FreePeerMem(commArgsPtr_);

    // 清理 UDMA 资源(如果已初始化)
    if (udmaInfoDev_ != nullptr) {
        // udmaInfoDev_ 指向 shmem 管理的设备内存,由 shmem finalize 释放
        aclshmem_finalize();
        udmaInfoDev_ = nullptr;
    }
}

TileXRComm::TileXRComm(int rank, int rankSize) : rank_(rank), rankSize_(rankSize)
{
}

TileXRComm::TileXRComm(int rank, int rankSize, int commDomain, int bufferSize)
    : rank_(rank), rankSize_(rankSize), commDomain_(commDomain), bufferSize_(bufferSize)
{
}

TileXRComm::TileXRComm(int rank, int rankSize, TileXRUniqueId commId)
    : rank_(rank), rankSize_(rankSize), commId_(commId)
{
}

int TileXRComm::GetRank() const
{
    return rank_;
}

int TileXRComm::GetRankSize() const
{
    return rankSize_;
}

int TileXRComm::GetCommSize() const
{
    return commSize_;
}

const PhysicalInfo &TileXRComm::GetPhysicalInfo() const
{
    return physicalInfo_;
}

GM_ADDR TileXRComm::GetCommArgsPtr()
{
    return commArgsPtr_;
}

CommArgs* TileXRComm::GetCommArgs()
{
    return &commArgs_;
}

std::string TileXRComm::PrintDFX()
{
    if (commArgsPtr_ == nullptr) {
        return "no comm args";
    }

    int ret = aclrtMemcpy(&commArgs_, sizeof(commArgs_), commArgsPtr_, sizeof(commArgs_),
                          ACL_MEMCPY_DEVICE_TO_HOST);
    if (ret != ACL_SUCCESS) {
        MKI_LOG(ERROR) << "aclrtMemcpy err " << __LINE__ << " " << ret;
        return "acl mem copy error";
    }
    stringstream ss;
    // 输出CommArgs基本属性
    ss << "CommArgs {"
       << "\n  rank: " << commArgs_.rank
       << "\n  localRank: " << commArgs_.localRank
       << "\n  rankSize: " << commArgs_.rankSize
       << "\n  localRankSize: " << commArgs_.localRankSize
       << "\n  extraFlag:  0x" << std::hex << std::setfill('0') << commArgs_.extraFlag;

    // 输出peerMems数组内容
    ss << "\n  peerMems: [";
    for (int i = 0; i < TILEXR_MAX_RANK_SIZE; ++i) {
        if (commArgs_.peerMems[i] == nullptr) {
            continue;
        }
        if (i > 0) {
            ss << ", ";
        }
        ss << "{id: " << static_cast<void *>(commArgs_.peerMems[i]) << "}";
    }
    ss << "]";

    // magic数组内容
    ss << "\n  magics: [";
    for (int i = 0; i < rankSize_; ++i) {
        ss << std::dec << commArgs_.magics[i] << ",";
    }
    ss << "] \n";

    // 输出dfx数组内容
    ss << "\n  dfx: [";
    const int dfxGroupCount = 5;
    for (int i = 0; i < DFX_COUNT; ++i) {
        if (i % dfxGroupCount == 0) {
            ss << "\n    " << std::dec << setw(dfxGroupCount) << i << ": ";
        }
        ss << "0x"<< std::hex << commArgs_.dfx[i] << ", ";
    }
    ss << "\n    ]";

    ss << "\n}";
    return ss.str();
}

}  // TileXR