LCOV - code coverage report
Current view: top level - base_comm/resources/ccu/ccu_transport - ccu_transport_.h (source / functions) Coverage Total Hit
Test: coverage.info Lines: 100.0 % 46 46
Test Date: 2026-08-18 17:47:01 Functions: 100.0 % 17 17

            Line data    Source code
       1              : /**
       2              :  * Copyright (c) 2025 Huawei Technologies Co., Ltd.
       3              :  * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
       4              :  * CANN Open Software License Agreement Version 2.0 (the "License").
       5              :  * Please refer to the License for details. You may not use this file except in compliance with the License.
       6              :  * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
       7              :  * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
       8              :  * See LICENSE in the root of the software repository for the full text of the License.
       9              :  */
      10              : 
      11              : #ifndef HCOMM_CCU_TRANSPORT_H
      12              : #define HCOMM_CCU_TRANSPORT_H
      13              : 
      14              : #include <memory>
      15              : #include <vector>
      16              : #include <shared_mutex>
      17              : 
      18              : #include "../../../../../legacy/ascend950/unified_platform/resource/socket/socket.h"
      19              : #include "env_config.h"
      20              : #include "op_mode.h"
      21              : #include "binary_stream.h"
      22              : #include "ccu_conn.h"
      23              : #include "remote_rma_buffer.h"
      24              : 
      25              : namespace hcomm {
      26              : 
      27              : class CcuTransport {
      28              : public:
      29              :     // 缩减channel预留的cke、xn量,避免与ccu instance需求量冲突
      30              :     static constexpr uint32_t INIT_CKE_NUM = 4;
      31              :     static constexpr uint32_t INIT_XN_NUM = 4;
      32          571 :     MAKE_ENUM(
      33              :         TransStatus, INIT, SEND_DATA_SIZE, RECV_DATA_SIZE, SEND_ALL_INFO, RECV_ALL_INFO, SEND_TRANS_RES, RECV_TRANS_RES,
      34              :         SEND_FIN, RECV_FIN, RECVING_FIN, RECVING_TRANS_RES, READY, CONNECT_FAILED, SOCKET_TIMEOUT)
      35              : 
      36              :     struct CclBufferInfo {
      37              :         uint64_t addr{0};
      38              :         uint32_t size{0};
      39              :         uint32_t tokenId{0};
      40              :         uint32_t tokenValue{0};
      41              :         CommMemType type{COMM_MEM_TYPE_INVALID};
      42              :         std::array<char, HCCL_RES_TAG_MAX_LEN> memInfo{};
      43              : 
      44           41 :         explicit CclBufferInfo() = default;
      45           30 :         CclBufferInfo(const uint64_t addr, const uint32_t size, const uint32_t tokenId, const uint32_t tokenValue)
      46           30 :             : addr(addr),
      47           30 :               size(size),
      48           30 :               tokenId(tokenId),
      49           30 :               tokenValue(tokenValue)
      50           30 :         {}
      51              : 
      52           16 :         CclBufferInfo(
      53              :             const uint64_t addr, const uint32_t size, const uint32_t tokenId, const uint32_t tokenValue,
      54              :             const CommMemType type, const std::array<char, HCCL_RES_TAG_MAX_LEN>& memInfo)
      55           16 :             : addr(addr),
      56           16 :               size(size),
      57           16 :               tokenId(tokenId),
      58           16 :               tokenValue(tokenValue),
      59           16 :               type(type),
      60           16 :               memInfo(memInfo)
      61           16 :         {}
      62              : 
      63            2 :         void Pack(Hccl::BinaryStream& binaryStream) const
      64              :         {
      65            2 :             binaryStream << addr << size << tokenId << tokenValue << type;
      66              :             // 逐个字节传输
      67          512 :             for (uint32_t i = 0; i < HCCL_RES_TAG_MAX_LEN; ++i) {
      68          510 :                 binaryStream << static_cast<u8>(memInfo[i]);
      69              :             }
      70            2 :             HCCL_INFO("Pack Ccl Buffer Info: addr[%llu] size[%u] memInfo[%s]", addr, size, memInfo.data());
      71            2 :         }
      72              : 
      73            2 :         void Unpack(Hccl::BinaryStream& binaryStream)
      74              :         {
      75            2 :             binaryStream >> addr >> size >> tokenId >> tokenValue >> type;
      76          512 :             for (uint32_t i = 0; i < HCCL_RES_TAG_MAX_LEN; ++i) {
      77              :                 u8 byte;
      78          510 :                 binaryStream >> byte;
      79          510 :                 memInfo[i] = static_cast<char>(byte);
      80              :             }
      81            4 :             HCCL_INFO("Unpack Ccl Buffer Info: addr[%llu] size[%u] memInfo[%s]", addr, size, memInfo.data());
      82            2 :         }
      83              :     };
      84              : 
      85           30 :     MAKE_ENUM(CcuConnectionType, UBC_TP, UB_CTP, UBC_CTP = UB_CTP);
      86              :     struct CcuConnectionInfo {
      87              :         CcuConnectionType type{CcuConnectionType::UBC_TP};
      88              :         CommAddr locAddr{};
      89              :         CommAddr rmtAddr{};
      90              :         CcuChannelInfo channelInfo{};
      91              :         std::vector<CcuJetty*> ccuJettys{};
      92              :         uint32_t qos{EnvConfig::UB_QOS_DEFAULT};
      93              : 
      94              :         explicit CcuConnectionInfo() = default;
      95           15 :         CcuConnectionInfo(
      96              :             const CcuConnectionType type, const CommAddr& locAddr, const CommAddr& rmtAddr,
      97              :             const CcuChannelInfo& channelInfo, const std::vector<CcuJetty*>& ccuJettys,
      98              :             uint32_t qos = EnvConfig::UB_QOS_DEFAULT)
      99           15 :             : type(type),
     100           15 :               locAddr(locAddr),
     101           15 :               rmtAddr(rmtAddr),
     102           15 :               channelInfo(channelInfo),
     103           15 :               ccuJettys(ccuJettys),
     104           15 :               qos(qos)
     105           15 :         {}
     106              :     };
     107              : 
     108              :     CcuTransport(Hccl::Socket* socket, std::unique_ptr<CcuConnection>&& connection, const CclBufferInfo& locCclBufInfo);
     109              :     CcuTransport(
     110              :         Hccl::Socket* socket, std::unique_ptr<CcuConnection>&& connection,
     111              :         const std::vector<CclBufferInfo>& bufferInfos);
     112              :     CcuTransport(const CcuTransport& that) = delete;
     113              :     CcuTransport& operator=(const CcuTransport& other) = delete;
     114              :     ~CcuTransport();
     115              :     HcclResult Init();
     116              :     TransStatus GetStatus();
     117              :     void Clean();
     118              :     HcclResult GetRemoteMems(uint32_t* memNum, CommMem** remoteMem, char*** memInfos);
     119              :     HcclResult CheckSocketStatus();
     120              :     HcclResult UpdateMemInfo(std::vector<CcuTransport::CclBufferInfo>& bufferVecTemp);
     121              :     HcclResult ResUpdate(std::vector<std::string>& resGroupTags);
     122              : 
     123              :     // 下面接口为平台层接口,不能在框架层使用
     124              :     uint32_t GetDieId() const;
     125              :     uint32_t GetChannelId() const;
     126              :     HcclResult GetLocCkeByIndex(const uint32_t index, uint32_t& locCkeId) const;
     127              :     HcclResult GetLocXnByIndex(const uint32_t index, uint32_t& locXnId) const;
     128              :     HcclResult GetRmtCkeByIndex(const uint32_t index, uint32_t& rmtCkeId) const;
     129              :     HcclResult GetRmtXnByIndex(const uint32_t index, uint32_t& rmtXnId) const;
     130              :     HcclResult GetRmtWishCntXnAddr(const std::string& resGroupTag, uint64_t& wishCntXnAddr) const;
     131              :     HcclResult GetLocBuffer(CclBufferInfo& bufferInfo, const uint32_t& bufNum) const;
     132              :     HcclResult GetRmtBuffer(CclBufferInfo& bufferInfo, const uint32_t& bufNum) const;
     133              :     HcclResult GetCkeNum(uint32_t& ckeNum) const;
     134              :     HcclResult GetRmtVarAddrByIndex(uint32_t index, uint64_t& rmtXnAddr) const;
     135              :     HcclResult GetRmtSignalAddrByIndex(uint32_t index, uint64_t& rmtCkeAddr) const;
     136              :     HcclResult GetRmtCcuBufferTokenInfo(uint32_t& rmtTokenId, uint32_t& rmtTokenValue) const;
     137              :     std::string Describe() const;
     138              :     HcclResult Describe(std::string& dfxMsg);
     139              : 
     140              : public:
     141              :     struct Attribution {
     142              :         Hccl::OpMode opMode{Hccl::OpMode::OPBASE};
     143              :         u32 devicePhyId{0};
     144              :         std::vector<char> handshakeMsg{};
     145            1 :         std::string Describe() const
     146              :         {
     147              :             return Hccl::StringFormat(
     148            1 :                 "CcuTransportAttribution[opMode=%s, devicePhyId=%u, handshakeMsg=%s]", opMode.Describe().c_str(),
     149            2 :                 devicePhyId, Hccl::Bytes2hex(handshakeMsg.data(), handshakeMsg.size()).c_str());
     150              :         }
     151              :     };
     152              : 
     153            1 :     std::vector<char>& GetRmtHandshakeMsg() // 返回握手消息
     154              :     {
     155            1 :         return rmtHandshakeMsg_;
     156              :     }
     157              : 
     158            1 :     std::vector<char>& GetLocalHandshakeMsg() // 返回握手消息
     159              :     {
     160            1 :         return attr_.handshakeMsg;
     161              :     }
     162              : 
     163            1 :     void SetHandshakeMsg(const std::vector<char>& handshakeMsg) { attr_.handshakeMsg = handshakeMsg; }
     164              : 
     165              : private:
     166              :     HcclResult StatusMachine();
     167              :     HcclResult AppendCkes(uint32_t ckesNum);
     168              :     HcclResult AppendXns(uint32_t xnsNum);
     169              :     HcclResult AppendCntXns();
     170              :     HcclResult SendFinish();
     171              :     HcclResult RecvFinish();
     172              :     HcclResult CheckFinish();
     173              :     HcclResult RecvDataProcess();
     174              :     HcclResult RecvTransInfoProcess();
     175              :     HcclResult ReleaseTransRes();
     176              :     HcclResult SendConnAndTransInfo();
     177              :     HcclResult RecvConnAndTransInfo();
     178              :     HcclResult SendDataSize();
     179              :     HcclResult RecvDataSize();
     180              :     HcclResult SendTransInfo();
     181              :     HcclResult RecvTransInfo();
     182              :     HcclResult HandshakeMsgPack(Hccl::BinaryStream& binaryStream);
     183              :     HcclResult ConnInfoPack(Hccl::BinaryStream& binaryStream) const;
     184              :     HcclResult TransResPack(Hccl::BinaryStream& binaryStream);
     185              :     HcclResult TransCntXnResPack(Hccl::BinaryStream& binaryStream);
     186              :     HcclResult BufferInfoPack(Hccl::BinaryStream& binaryStream, std::vector<CclBufferInfo>& bufferVec) const;
     187              :     HcclResult HandshakeMsgUnpack(Hccl::BinaryStream& binaryStream);
     188              :     HcclResult ConnInfoUnpackProc(Hccl::BinaryStream& binaryStream) const;
     189              :     HcclResult TransResUnpackProc(Hccl::BinaryStream& binaryStream);
     190              :     HcclResult TransCntXnResUnpackProc(Hccl::BinaryStream& binaryStream);
     191              :     HcclResult BufferInfoUnpack(Hccl::BinaryStream& binaryStream);
     192              :     HcclResult GetRmtVarAddrByXnId(const uint32_t rmtXnId, uint64_t& rmtXnAddr) const;
     193              : 
     194              :     HcclResult ReturnErrorStatus(const std::string& funcName);
     195              : 
     196              : private:
     197              :     // 保存transport中需要使用的cke,xn等ccu资源
     198              :     struct TransRes {
     199              :         std::vector<uint32_t> ckes{};
     200              :         std::vector<uint32_t> xns{};
     201              :         std::map<std::string, uint32_t> cntXns{}; // {groupTag, wishCntXn}
     202              :     };
     203              : 
     204              :     uint32_t dieId_{0};
     205              :     int32_t devLogicId_{0};
     206              :     Attribution attr_{};
     207              :     std::vector<char> rmtHandshakeMsg_{0}; // 远端握手消息
     208              :     Hccl::Socket* socket_{nullptr};
     209              :     std::unique_ptr<CcuConnection> ccuConnection_;
     210              :     TransRes locRes_{};
     211              :     TransRes rmtRes_{};
     212              :     TransStatus transStatus_{TransStatus::INVALID};
     213              :     std::vector<std::vector<ResInfo>> ckesRes_{};
     214              :     std::vector<std::vector<ResInfo>> xnsRes_{};
     215              :     std::vector<CclBufferInfo> locBufferInfos_{};
     216              :     CclBufferInfo rmtHcclBufferInfo_{};
     217              :     std::vector<std::unique_ptr<Hccl::RemoteUbRmaBuffer>> rmtBufferVec_{};
     218              :     uint32_t exchangeDataSize_{0};
     219              :     std::vector<char> recvData_{};
     220              :     std::vector<char> recvTrans_{};
     221              :     std::vector<char> sendData_{};
     222              :     std::vector<char> sendTrans_{};
     223              :     std::vector<char> recvFinishMsg_{};
     224              :     std::vector<char> sendFinishMsg_{};
     225              :     bool cacheValid_ = false;                // GetUserRemoteMem 的缓存标识
     226              :     std::mutex remoteMemsMutex_;             // 远端内存列表互斥锁
     227              :     std::vector<CommMem> remoteUserMems_;    // 内存基本信息缓存
     228              :     std::vector<std::string> memInfoCopies_; // 储存 Tag 字符串副本
     229              :     std::vector<char*> memInfoPointers_;     // Tag 缓存
     230              : };
     231              : 
     232              : HcclResult BuildCcuConnection(
     233              :     const CcuTransport::CcuConnectionInfo& ccuConnectionInfo, std::unique_ptr<CcuConnection>& ccuConnection);
     234              : 
     235              : HcclResult CcuCreateTransport(
     236              :     Hccl::Socket* socket, const CcuTransport::CcuConnectionInfo& ccuConnectionInfo,
     237              :     const CcuTransport::CclBufferInfo& cclBufferInfo, std::unique_ptr<CcuTransport>& ccuTransport);
     238              : 
     239              : HcclResult CcuCreateTransport(
     240              :     Hccl::Socket* socket, const CcuTransport::CcuConnectionInfo& ccuConnectionInfo,
     241              :     const std::vector<CcuTransport::CclBufferInfo>& bufferInfos, std::unique_ptr<CcuTransport>& ccuTransport);
     242              : } // namespace hcomm
     243              : #endif // HCOMM_CCU_TRANSPORT_H
        

Generated by: LCOV version 2.0-1