Loading...
Searching...
No Matches
ServeRpcServer.cpp
Go to the documentation of this file.
1/*
2** This file is part of eOn.
3**
4** SPDX-License-Identifier: BSD-3-Clause
5**
6** Copyright (c) 2010--present, eOn Development Team
7** All rights reserved.
8**
9** Repo:
10** https://github.com/TheochemUI/eOn
11*/
12
24
25#include "eon/ServeRpcServer.h"
26#include "eon/BaseStructures.h"
27#include "eon/EonLogger.h"
28
29#include <atomic>
30#include <capnp/ez-rpc.h>
31#include <kj/debug.h>
32#include <mutex>
33
34// Cap'n Proto generated header (from Potentials.capnp).
35// This defines `class Potential` -- which collides with eOn's Potential class,
36// hence the separate translation unit.
37#include "rgpot/rpc/Potentials.capnp.h"
38
39namespace eonc {
40
41namespace {
42
50class CallbackPotImpl final : public Potential::Server {
51public:
52 explicit CallbackPotImpl(ForceCallback cb)
53 : m_callback(std::move(cb)) {}
54
55 kj::Promise<void> calculate(CalculateContext context) override {
56 auto fip = context.getParams().getFip();
57 const size_t numAtoms = fip.getPos().size() / 3;
58
59 KJ_REQUIRE(fip.getAtmnrs().size() == numAtoms, "AtomNumbers size mismatch");
60
61 // Extract flat arrays from capnp
62 std::vector<double> positions(numAtoms * 3);
63 auto capnpPos = fip.getPos();
64 for (size_t i = 0; i < numAtoms * 3; ++i) {
65 positions[i] = capnpPos[i];
66 }
67
68 std::vector<int> atomicNrs(numAtoms);
69 auto capnpAtmnrs = fip.getAtmnrs();
70 for (size_t i = 0; i < numAtoms; ++i) {
71 atomicNrs[i] = capnpAtmnrs[i];
72 }
73
74 double box[9] = {};
75 auto capnpBox = fip.getBox();
76 for (size_t i = 0; i < 9 && i < capnpBox.size(); ++i) {
77 box[i] = capnpBox[i];
78 }
79
80 // Call the force callback
81 std::vector<double> forces(numAtoms * 3, 0.0);
82 double energy = 0.0;
83 m_callback(static_cast<long>(numAtoms), positions.data(), atomicNrs.data(),
84 forces.data(), &energy, box);
85
86 // Serialize result back to capnp
87 auto result = context.getResults();
88 auto pres = result.initResult();
89 pres.setEnergy(energy);
90
91 auto forcesList = pres.initForces(numAtoms * 3);
92 for (size_t i = 0; i < numAtoms * 3; ++i) {
93 forcesList.set(i, forces[i]);
94 }
95
96 return kj::READY_NOW;
97 }
98
99private:
100 ForceCallback m_callback;
101};
102
103} // anonymous namespace
104
105void startRpcServer(ForceCallback callback, const std::string &host,
106 uint16_t port) {
107 EONC_LOG_INFO("Starting Cap'n Proto RPC server on {}:{}", host, port);
108
109 capnp::EzRpcServer server(kj::heap<CallbackPotImpl>(std::move(callback)),
110 host, port);
111
112 auto &waitScope = server.getWaitScope();
113 EONC_LOG_INFO("Server ready on port {}. Ctrl+C to stop.", port);
114 kj::NEVER_DONE.wait(waitScope);
115}
116
117// ---------------------------------------------------------------------------
118// Pooled (round-robin gateway) server
119// ---------------------------------------------------------------------------
120
121namespace {
122
127class PooledCallbackPotImpl final : public Potential::Server {
128public:
129 explicit PooledCallbackPotImpl(std::vector<ForceCallback> pool)
130 : m_pool(std::move(pool)),
131 m_mutexes(m_pool.size()),
132 m_next(0) {}
133
134 kj::Promise<void> calculate(CalculateContext context) override {
135 size_t idx = m_next.fetch_add(1, std::memory_order_relaxed) % m_pool.size();
136
137 auto fip = context.getParams().getFip();
138 const size_t numAtoms = fip.getPos().size() / 3;
139
140 KJ_REQUIRE(fip.getAtmnrs().size() == numAtoms, "AtomNumbers size mismatch");
141
142 std::vector<double> positions(numAtoms * 3);
143 auto capnpPos = fip.getPos();
144 for (size_t i = 0; i < numAtoms * 3; ++i) {
145 positions[i] = capnpPos[i];
146 }
147
148 std::vector<int> atomicNrs(numAtoms);
149 auto capnpAtmnrs = fip.getAtmnrs();
150 for (size_t i = 0; i < numAtoms; ++i) {
151 atomicNrs[i] = capnpAtmnrs[i];
152 }
153
154 double box[9] = {};
155 auto capnpBox = fip.getBox();
156 for (size_t i = 0; i < 9 && i < capnpBox.size(); ++i) {
157 box[i] = capnpBox[i];
158 }
159
160 std::vector<double> forces(numAtoms * 3, 0.0);
161 double energy = 0.0;
162
163 std::lock_guard<std::mutex> lock(m_mutexes[idx]);
164 m_pool[idx](static_cast<long>(numAtoms), positions.data(), atomicNrs.data(),
165 forces.data(), &energy, box);
166
167 auto result = context.getResults();
168 auto pres = result.initResult();
169 pres.setEnergy(energy);
170
171 auto forcesList = pres.initForces(numAtoms * 3);
172 for (size_t i = 0; i < numAtoms * 3; ++i) {
173 forcesList.set(i, forces[i]);
174 }
175
176 return kj::READY_NOW;
177 }
178
179private:
180 std::vector<ForceCallback> m_pool;
181 std::vector<std::mutex> m_mutexes;
182 std::atomic<size_t> m_next;
183};
184
185} // anonymous namespace
186
187void startPooledRpcServer(std::vector<ForceCallback> pool,
188 const std::string &host, uint16_t port) {
189 EONC_LOG_INFO("Starting pooled RPC gateway on {}:{} with {} instances", host,
190 port, pool.size());
191
192 capnp::EzRpcServer server(kj::heap<PooledCallbackPotImpl>(std::move(pool)),
193 host, port);
194
195 auto &waitScope = server.getWaitScope();
196 EONC_LOG_INFO("Gateway ready on port {}. Ctrl+C to stop.", port);
197 kj::NEVER_DONE.wait(waitScope);
198}
199
200} // namespace eonc
#define EONC_LOG_INFO(...)
Definition EonLogger.h:249
RAII resource manager for the ARTn C library with global synchronization.
std::function< void(long nAtoms, const double *positions, const int *atomicNrs, double *forces, double *energy, const double *box)> ForceCallback
Callback type for potential energy/force evaluation.
void startPooledRpcServer(std::vector< ForceCallback > pool, const std::string &host, uint16_t port)
Start a blocking Cap'n Proto RPC server backed by a pool of force callbacks dispatched round-robin.
void startRpcServer(ForceCallback callback, const std::string &host, uint16_t port)
Start a blocking Cap'n Proto RPC server using a force callback.