OrcRPCTargetProcessControl.h 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415
  1. //===--- OrcRPCTargetProcessControl.h - Remote target control ---*- C++ -*-===//
  2. //
  3. // Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
  4. // See https://llvm.org/LICENSE.txt for license information.
  5. // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
  6. //
  7. //===----------------------------------------------------------------------===//
  8. //
  9. // Utilities for interacting with target processes.
  10. //
  11. //===----------------------------------------------------------------------===//
  12. #ifndef LLVM_EXECUTIONENGINE_ORC_ORCRPCTARGETPROCESSCONTROL_H
  13. #define LLVM_EXECUTIONENGINE_ORC_ORCRPCTARGETPROCESSCONTROL_H
  14. #include "llvm/ExecutionEngine/Orc/Shared/RPCUtils.h"
  15. #include "llvm/ExecutionEngine/Orc/Shared/RawByteChannel.h"
  16. #include "llvm/ExecutionEngine/Orc/TargetProcess/OrcRPCTPCServer.h"
  17. #include "llvm/ExecutionEngine/Orc/TargetProcessControl.h"
  18. #include "llvm/Support/MSVCErrorWorkarounds.h"
  19. namespace llvm {
  20. namespace orc {
  21. /// JITLinkMemoryManager implementation for a process connected via an ORC RPC
  22. /// endpoint.
  23. template <typename OrcRPCTPCImplT>
  24. class OrcRPCTPCJITLinkMemoryManager : public jitlink::JITLinkMemoryManager {
  25. private:
  26. struct HostAlloc {
  27. std::unique_ptr<char[]> Mem;
  28. uint64_t Size;
  29. };
  30. struct TargetAlloc {
  31. JITTargetAddress Address = 0;
  32. uint64_t AllocatedSize = 0;
  33. };
  34. using HostAllocMap = DenseMap<int, HostAlloc>;
  35. using TargetAllocMap = DenseMap<int, TargetAlloc>;
  36. public:
  37. class OrcRPCAllocation : public Allocation {
  38. public:
  39. OrcRPCAllocation(OrcRPCTPCJITLinkMemoryManager<OrcRPCTPCImplT> &Parent,
  40. HostAllocMap HostAllocs, TargetAllocMap TargetAllocs)
  41. : Parent(Parent), HostAllocs(std::move(HostAllocs)),
  42. TargetAllocs(std::move(TargetAllocs)) {
  43. assert(HostAllocs.size() == TargetAllocs.size() &&
  44. "HostAllocs size should match TargetAllocs");
  45. }
  46. ~OrcRPCAllocation() override {
  47. assert(TargetAllocs.empty() && "failed to deallocate");
  48. }
  49. MutableArrayRef<char> getWorkingMemory(ProtectionFlags Seg) override {
  50. auto I = HostAllocs.find(Seg);
  51. assert(I != HostAllocs.end() && "No host allocation for segment");
  52. auto &HA = I->second;
  53. return {HA.Mem.get(), static_cast<size_t>(HA.Size)};
  54. }
  55. JITTargetAddress getTargetMemory(ProtectionFlags Seg) override {
  56. auto I = TargetAllocs.find(Seg);
  57. assert(I != TargetAllocs.end() && "No target allocation for segment");
  58. return I->second.Address;
  59. }
  60. void finalizeAsync(FinalizeContinuation OnFinalize) override {
  61. std::vector<tpctypes::BufferWrite> BufferWrites;
  62. orcrpctpc::ReleaseOrFinalizeMemRequest FMR;
  63. for (auto &KV : HostAllocs) {
  64. assert(TargetAllocs.count(KV.first) &&
  65. "No target allocation for buffer");
  66. auto &HA = KV.second;
  67. auto &TA = TargetAllocs[KV.first];
  68. BufferWrites.push_back({TA.Address, StringRef(HA.Mem.get(), HA.Size)});
  69. FMR.push_back({orcrpctpc::toWireProtectionFlags(
  70. static_cast<sys::Memory::ProtectionFlags>(KV.first)),
  71. TA.Address, TA.AllocatedSize});
  72. }
  73. DEBUG_WITH_TYPE("orc", {
  74. dbgs() << "finalizeAsync " << (void *)this << ":\n";
  75. auto FMRI = FMR.begin();
  76. for (auto &B : BufferWrites) {
  77. auto Prot = FMRI->Prot;
  78. ++FMRI;
  79. dbgs() << " Writing " << formatv("{0:x16}", B.Buffer.size())
  80. << " bytes to " << ((Prot & orcrpctpc::WPF_Read) ? 'R' : '-')
  81. << ((Prot & orcrpctpc::WPF_Write) ? 'W' : '-')
  82. << ((Prot & orcrpctpc::WPF_Exec) ? 'X' : '-')
  83. << " segment: local " << (const void *)B.Buffer.data()
  84. << " -> target " << formatv("{0:x16}", B.Address) << "\n";
  85. }
  86. });
  87. if (auto Err =
  88. Parent.Parent.getMemoryAccess().writeBuffers(BufferWrites)) {
  89. OnFinalize(std::move(Err));
  90. return;
  91. }
  92. DEBUG_WITH_TYPE("orc", dbgs() << " Applying permissions...\n");
  93. if (auto Err =
  94. Parent.getEndpoint().template callAsync<orcrpctpc::FinalizeMem>(
  95. [OF = std::move(OnFinalize)](Error Err2) {
  96. // FIXME: Dispatch to work queue.
  97. std::thread([OF = std::move(OF),
  98. Err3 = std::move(Err2)]() mutable {
  99. DEBUG_WITH_TYPE(
  100. "orc", { dbgs() << " finalizeAsync complete\n"; });
  101. OF(std::move(Err3));
  102. }).detach();
  103. return Error::success();
  104. },
  105. FMR)) {
  106. DEBUG_WITH_TYPE("orc", dbgs() << " failed.\n");
  107. Parent.getEndpoint().abandonPendingResponses();
  108. Parent.reportError(std::move(Err));
  109. }
  110. DEBUG_WITH_TYPE("orc", {
  111. dbgs() << "Leaving finalizeAsync (finalization may continue in "
  112. "background)\n";
  113. });
  114. }
  115. Error deallocate() override {
  116. orcrpctpc::ReleaseOrFinalizeMemRequest RMR;
  117. for (auto &KV : TargetAllocs)
  118. RMR.push_back({orcrpctpc::toWireProtectionFlags(
  119. static_cast<sys::Memory::ProtectionFlags>(KV.first)),
  120. KV.second.Address, KV.second.AllocatedSize});
  121. TargetAllocs.clear();
  122. return Parent.getEndpoint().template callB<orcrpctpc::ReleaseMem>(RMR);
  123. }
  124. private:
  125. OrcRPCTPCJITLinkMemoryManager<OrcRPCTPCImplT> &Parent;
  126. HostAllocMap HostAllocs;
  127. TargetAllocMap TargetAllocs;
  128. };
  129. OrcRPCTPCJITLinkMemoryManager(OrcRPCTPCImplT &Parent) : Parent(Parent) {}
  130. Expected<std::unique_ptr<Allocation>>
  131. allocate(const jitlink::JITLinkDylib *JD,
  132. const SegmentsRequestMap &Request) override {
  133. orcrpctpc::ReserveMemRequest RMR;
  134. HostAllocMap HostAllocs;
  135. for (auto &KV : Request) {
  136. assert(KV.second.getContentSize() <= std::numeric_limits<size_t>::max() &&
  137. "Content size is out-of-range for host");
  138. RMR.push_back({orcrpctpc::toWireProtectionFlags(
  139. static_cast<sys::Memory::ProtectionFlags>(KV.first)),
  140. KV.second.getContentSize() + KV.second.getZeroFillSize(),
  141. KV.second.getAlignment()});
  142. HostAllocs[KV.first] = {
  143. std::make_unique<char[]>(KV.second.getContentSize()),
  144. KV.second.getContentSize()};
  145. }
  146. DEBUG_WITH_TYPE("orc", {
  147. dbgs() << "Orc remote memmgr got request:\n";
  148. for (auto &KV : Request)
  149. dbgs() << " permissions: "
  150. << ((KV.first & sys::Memory::MF_READ) ? 'R' : '-')
  151. << ((KV.first & sys::Memory::MF_WRITE) ? 'W' : '-')
  152. << ((KV.first & sys::Memory::MF_EXEC) ? 'X' : '-')
  153. << ", content size: "
  154. << formatv("{0:x16}", KV.second.getContentSize())
  155. << " + zero-fill-size: "
  156. << formatv("{0:x16}", KV.second.getZeroFillSize())
  157. << ", align: " << KV.second.getAlignment() << "\n";
  158. });
  159. // FIXME: LLVM RPC needs to be fixed to support alt
  160. // serialization/deserialization on return types. For now just
  161. // translate from std::map to DenseMap manually.
  162. auto TmpTargetAllocs =
  163. Parent.getEndpoint().template callB<orcrpctpc::ReserveMem>(RMR);
  164. if (!TmpTargetAllocs)
  165. return TmpTargetAllocs.takeError();
  166. if (TmpTargetAllocs->size() != RMR.size())
  167. return make_error<StringError>(
  168. "Number of target allocations does not match request",
  169. inconvertibleErrorCode());
  170. TargetAllocMap TargetAllocs;
  171. for (auto &E : *TmpTargetAllocs)
  172. TargetAllocs[orcrpctpc::fromWireProtectionFlags(E.Prot)] = {
  173. E.Address, E.AllocatedSize};
  174. DEBUG_WITH_TYPE("orc", {
  175. auto HAI = HostAllocs.begin();
  176. for (auto &KV : TargetAllocs)
  177. dbgs() << " permissions: "
  178. << ((KV.first & sys::Memory::MF_READ) ? 'R' : '-')
  179. << ((KV.first & sys::Memory::MF_WRITE) ? 'W' : '-')
  180. << ((KV.first & sys::Memory::MF_EXEC) ? 'X' : '-')
  181. << " assigned local " << (void *)HAI->second.Mem.get()
  182. << ", target " << formatv("{0:x16}", KV.second.Address) << "\n";
  183. });
  184. return std::make_unique<OrcRPCAllocation>(*this, std::move(HostAllocs),
  185. std::move(TargetAllocs));
  186. }
  187. private:
  188. void reportError(Error Err) { Parent.reportError(std::move(Err)); }
  189. decltype(std::declval<OrcRPCTPCImplT>().getEndpoint()) getEndpoint() {
  190. return Parent.getEndpoint();
  191. }
  192. OrcRPCTPCImplT &Parent;
  193. };
  194. /// TargetProcessControl::MemoryAccess implementation for a process connected
  195. /// via an ORC RPC endpoint.
  196. template <typename OrcRPCTPCImplT>
  197. class OrcRPCTPCMemoryAccess : public TargetProcessControl::MemoryAccess {
  198. public:
  199. OrcRPCTPCMemoryAccess(OrcRPCTPCImplT &Parent) : Parent(Parent) {}
  200. void writeUInt8s(ArrayRef<tpctypes::UInt8Write> Ws,
  201. WriteResultFn OnWriteComplete) override {
  202. writeViaRPC<orcrpctpc::WriteUInt8s>(Ws, std::move(OnWriteComplete));
  203. }
  204. void writeUInt16s(ArrayRef<tpctypes::UInt16Write> Ws,
  205. WriteResultFn OnWriteComplete) override {
  206. writeViaRPC<orcrpctpc::WriteUInt16s>(Ws, std::move(OnWriteComplete));
  207. }
  208. void writeUInt32s(ArrayRef<tpctypes::UInt32Write> Ws,
  209. WriteResultFn OnWriteComplete) override {
  210. writeViaRPC<orcrpctpc::WriteUInt32s>(Ws, std::move(OnWriteComplete));
  211. }
  212. void writeUInt64s(ArrayRef<tpctypes::UInt64Write> Ws,
  213. WriteResultFn OnWriteComplete) override {
  214. writeViaRPC<orcrpctpc::WriteUInt64s>(Ws, std::move(OnWriteComplete));
  215. }
  216. void writeBuffers(ArrayRef<tpctypes::BufferWrite> Ws,
  217. WriteResultFn OnWriteComplete) override {
  218. writeViaRPC<orcrpctpc::WriteBuffers>(Ws, std::move(OnWriteComplete));
  219. }
  220. private:
  221. template <typename WriteRPCFunction, typename WriteElementT>
  222. void writeViaRPC(ArrayRef<WriteElementT> Ws, WriteResultFn OnWriteComplete) {
  223. if (auto Err = Parent.getEndpoint().template callAsync<WriteRPCFunction>(
  224. [OWC = std::move(OnWriteComplete)](Error Err2) mutable -> Error {
  225. OWC(std::move(Err2));
  226. return Error::success();
  227. },
  228. Ws)) {
  229. Parent.reportError(std::move(Err));
  230. Parent.getEndpoint().abandonPendingResponses();
  231. }
  232. }
  233. OrcRPCTPCImplT &Parent;
  234. };
  235. // TargetProcessControl for a process connected via an ORC RPC Endpoint.
  236. template <typename RPCEndpointT>
  237. class OrcRPCTargetProcessControlBase : public TargetProcessControl {
  238. public:
  239. using ErrorReporter = unique_function<void(Error)>;
  240. using OnCloseConnectionFunction = unique_function<Error(Error)>;
  241. OrcRPCTargetProcessControlBase(std::shared_ptr<SymbolStringPool> SSP,
  242. RPCEndpointT &EP, ErrorReporter ReportError)
  243. : TargetProcessControl(std::move(SSP)),
  244. ReportError(std::move(ReportError)), EP(EP) {}
  245. void reportError(Error Err) { ReportError(std::move(Err)); }
  246. RPCEndpointT &getEndpoint() { return EP; }
  247. Expected<tpctypes::DylibHandle> loadDylib(const char *DylibPath) override {
  248. DEBUG_WITH_TYPE("orc", {
  249. dbgs() << "Loading dylib \"" << (DylibPath ? DylibPath : "") << "\" ";
  250. if (!DylibPath)
  251. dbgs() << "(process symbols)";
  252. dbgs() << "\n";
  253. });
  254. if (!DylibPath)
  255. DylibPath = "";
  256. auto H = EP.template callB<orcrpctpc::LoadDylib>(DylibPath);
  257. DEBUG_WITH_TYPE("orc", {
  258. if (H)
  259. dbgs() << " got handle " << formatv("{0:x16}", *H) << "\n";
  260. else
  261. dbgs() << " error, unable to load\n";
  262. });
  263. return H;
  264. }
  265. Expected<std::vector<tpctypes::LookupResult>>
  266. lookupSymbols(ArrayRef<LookupRequest> Request) override {
  267. std::vector<orcrpctpc::RemoteLookupRequest> RR;
  268. for (auto &E : Request) {
  269. RR.push_back({});
  270. RR.back().first = E.Handle;
  271. for (auto &KV : E.Symbols)
  272. RR.back().second.push_back(
  273. {(*KV.first).str(),
  274. KV.second == SymbolLookupFlags::WeaklyReferencedSymbol});
  275. }
  276. DEBUG_WITH_TYPE("orc", {
  277. dbgs() << "Compound lookup:\n";
  278. for (auto &R : Request) {
  279. dbgs() << " In " << formatv("{0:x16}", R.Handle) << ": {";
  280. bool First = true;
  281. for (auto &KV : R.Symbols) {
  282. dbgs() << (First ? "" : ",") << " " << *KV.first;
  283. First = false;
  284. }
  285. dbgs() << " }\n";
  286. }
  287. });
  288. return EP.template callB<orcrpctpc::LookupSymbols>(RR);
  289. }
  290. Expected<int32_t> runAsMain(JITTargetAddress MainFnAddr,
  291. ArrayRef<std::string> Args) override {
  292. DEBUG_WITH_TYPE("orc", {
  293. dbgs() << "Running as main: " << formatv("{0:x16}", MainFnAddr)
  294. << ", args = [";
  295. for (unsigned I = 0; I != Args.size(); ++I)
  296. dbgs() << (I ? "," : "") << " \"" << Args[I] << "\"";
  297. dbgs() << "]\n";
  298. });
  299. auto Result = EP.template callB<orcrpctpc::RunMain>(MainFnAddr, Args);
  300. DEBUG_WITH_TYPE("orc", {
  301. dbgs() << " call to " << formatv("{0:x16}", MainFnAddr);
  302. if (Result)
  303. dbgs() << " returned result " << *Result << "\n";
  304. else
  305. dbgs() << " failed\n";
  306. });
  307. return Result;
  308. }
  309. Expected<tpctypes::WrapperFunctionResult>
  310. runWrapper(JITTargetAddress WrapperFnAddr,
  311. ArrayRef<uint8_t> ArgBuffer) override {
  312. DEBUG_WITH_TYPE("orc", {
  313. dbgs() << "Running as wrapper function "
  314. << formatv("{0:x16}", WrapperFnAddr) << " with "
  315. << formatv("{0:x16}", ArgBuffer.size()) << " argument buffer\n";
  316. });
  317. auto Result =
  318. EP.template callB<orcrpctpc::RunWrapper>(WrapperFnAddr, ArgBuffer);
  319. // dbgs() << "Returned from runWrapper...\n";
  320. return Result;
  321. }
  322. Error closeConnection(OnCloseConnectionFunction OnCloseConnection) {
  323. DEBUG_WITH_TYPE("orc", dbgs() << "Closing connection to remote\n");
  324. return EP.template callAsync<orcrpctpc::CloseConnection>(
  325. std::move(OnCloseConnection));
  326. }
  327. Error closeConnectionAndWait() {
  328. std::promise<MSVCPError> P;
  329. auto F = P.get_future();
  330. if (auto Err = closeConnection([&](Error Err2) -> Error {
  331. P.set_value(std::move(Err2));
  332. return Error::success();
  333. })) {
  334. EP.abandonAllPendingResponses();
  335. return joinErrors(std::move(Err), F.get());
  336. }
  337. return F.get();
  338. }
  339. protected:
  340. /// Subclasses must call this during construction to initialize the
  341. /// TargetTriple and PageSize members.
  342. Error initializeORCRPCTPCBase() {
  343. if (auto TripleOrErr = EP.template callB<orcrpctpc::GetTargetTriple>())
  344. TargetTriple = Triple(*TripleOrErr);
  345. else
  346. return TripleOrErr.takeError();
  347. if (auto PageSizeOrErr = EP.template callB<orcrpctpc::GetPageSize>())
  348. PageSize = *PageSizeOrErr;
  349. else
  350. return PageSizeOrErr.takeError();
  351. return Error::success();
  352. }
  353. private:
  354. ErrorReporter ReportError;
  355. RPCEndpointT &EP;
  356. };
  357. } // end namespace orc
  358. } // end namespace llvm
  359. #endif // LLVM_EXECUTIONENGINE_ORC_ORCRPCTARGETPROCESSCONTROL_H