30#include <capnp/ez-rpc.h>
37#include "rgpot/rpc/Potentials.capnp.h"
50class CallbackPotImpl final :
public Potential::Server {
53 : m_callback(std::move(cb)) {}
55 kj::Promise<void> calculate(CalculateContext context)
override {
56 auto fip = context.getParams().getFip();
57 const size_t numAtoms = fip.getPos().size() / 3;
59 KJ_REQUIRE(fip.getAtmnrs().size() == numAtoms,
"AtomNumbers size mismatch");
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];
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];
75 auto capnpBox = fip.getBox();
76 for (
size_t i = 0; i < 9 && i < capnpBox.size(); ++i) {
81 std::vector<double> forces(numAtoms * 3, 0.0);
83 m_callback(
static_cast<long>(numAtoms), positions.data(), atomicNrs.data(),
84 forces.data(), &energy, box);
87 auto result = context.getResults();
88 auto pres = result.initResult();
89 pres.setEnergy(energy);
91 auto forcesList = pres.initForces(numAtoms * 3);
92 for (
size_t i = 0; i < numAtoms * 3; ++i) {
93 forcesList.set(i, forces[i]);
107 EONC_LOG_INFO(
"Starting Cap'n Proto RPC server on {}:{}", host, port);
109 capnp::EzRpcServer server(kj::heap<CallbackPotImpl>(std::move(callback)),
112 auto &waitScope = server.getWaitScope();
113 EONC_LOG_INFO(
"Server ready on port {}. Ctrl+C to stop.", port);
114 kj::NEVER_DONE.wait(waitScope);
127class PooledCallbackPotImpl final :
public Potential::Server {
129 explicit PooledCallbackPotImpl(std::vector<ForceCallback> pool)
130 : m_pool(std::move(pool)),
131 m_mutexes(m_pool.size()),
134 kj::Promise<void> calculate(CalculateContext context)
override {
135 size_t idx = m_next.fetch_add(1, std::memory_order_relaxed) % m_pool.size();
137 auto fip = context.getParams().getFip();
138 const size_t numAtoms = fip.getPos().size() / 3;
140 KJ_REQUIRE(fip.getAtmnrs().size() == numAtoms,
"AtomNumbers size mismatch");
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];
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];
155 auto capnpBox = fip.getBox();
156 for (
size_t i = 0; i < 9 && i < capnpBox.size(); ++i) {
157 box[i] = capnpBox[i];
160 std::vector<double> forces(numAtoms * 3, 0.0);
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);
167 auto result = context.getResults();
168 auto pres = result.initResult();
169 pres.setEnergy(energy);
171 auto forcesList = pres.initForces(numAtoms * 3);
172 for (
size_t i = 0; i < numAtoms * 3; ++i) {
173 forcesList.set(i, forces[i]);
176 return kj::READY_NOW;
180 std::vector<ForceCallback> m_pool;
181 std::vector<std::mutex> m_mutexes;
182 std::atomic<size_t> m_next;
188 const std::string &host, uint16_t port) {
189 EONC_LOG_INFO(
"Starting pooled RPC gateway on {}:{} with {} instances", host,
192 capnp::EzRpcServer server(kj::heap<PooledCallbackPotImpl>(std::move(pool)),
195 auto &waitScope = server.getWaitScope();
196 EONC_LOG_INFO(
"Gateway ready on port {}. Ctrl+C to stop.", port);
197 kj::NEVER_DONE.wait(waitScope);
#define EONC_LOG_INFO(...)
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.