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 "Potentials.capnp.h"
38
39namespace {
40
48class CallbackPotImpl final : public Potential::Server {
49public:
50 explicit CallbackPotImpl(ForceCallback cb)
51 : m_callback(std::move(cb)) {}
52
53 kj::Promise<void> calculate(CalculateContext context) override {
54 auto fip = context.getParams().getFip();
55 const size_t numAtoms = fip.getPos().size() / 3;
56
57 KJ_REQUIRE(fip.getAtmnrs().size() == numAtoms, "AtomNumbers size mismatch");
58
59 // Extract flat arrays from capnp
60 std::vector<double> positions(numAtoms * 3);
61 auto capnpPos = fip.getPos();
62 for (size_t i = 0; i < numAtoms * 3; ++i) {
63 positions[i] = capnpPos[i];
64 }
65
66 std::vector<int> atomicNrs(numAtoms);
67 auto capnpAtmnrs = fip.getAtmnrs();
68 for (size_t i = 0; i < numAtoms; ++i) {
69 atomicNrs[i] = capnpAtmnrs[i];
70 }
71
72 double box[9] = {};
73 auto capnpBox = fip.getBox();
74 for (size_t i = 0; i < 9 && i < capnpBox.size(); ++i) {
75 box[i] = capnpBox[i];
76 }
77
78 // Call the force callback
79 std::vector<double> forces(numAtoms * 3, 0.0);
80 double energy = 0.0;
81 m_callback(static_cast<long>(numAtoms), positions.data(), atomicNrs.data(),
82 forces.data(), &energy, box);
83
84 // Serialize result back to capnp
85 auto result = context.getResults();
86 auto pres = result.initResult();
87 pres.setEnergy(energy);
88
89 auto forcesList = pres.initForces(numAtoms * 3);
90 for (size_t i = 0; i < numAtoms * 3; ++i) {
91 forcesList.set(i, forces[i]);
92 }
93
94 return kj::READY_NOW;
95 }
96
97private:
98 ForceCallback m_callback;
99};
100
101} // anonymous namespace
102
103void eonc::startRpcServer(ForceCallback callback, const std::string &host,
104 uint16_t port) {
105 EONC_LOG_INFO("Starting Cap'n Proto RPC server on {}:{}", host, port);
106
107 capnp::EzRpcServer server(kj::heap<CallbackPotImpl>(std::move(callback)),
108 host, port);
109
110 auto &waitScope = server.getWaitScope();
111 EONC_LOG_INFO("Server ready on port {}. Ctrl+C to stop.", port);
112 kj::NEVER_DONE.wait(waitScope);
113}
114
115// ---------------------------------------------------------------------------
116// Pooled (round-robin gateway) server
117// ---------------------------------------------------------------------------
118
119namespace {
120
125class PooledCallbackPotImpl final : public Potential::Server {
126public:
127 explicit PooledCallbackPotImpl(std::vector<ForceCallback> pool)
128 : m_pool(std::move(pool)),
129 m_mutexes(m_pool.size()),
130 m_next(0) {}
131
132 kj::Promise<void> calculate(CalculateContext context) override {
133 size_t idx = m_next.fetch_add(1, std::memory_order_relaxed) % m_pool.size();
134
135 auto fip = context.getParams().getFip();
136 const size_t numAtoms = fip.getPos().size() / 3;
137
138 KJ_REQUIRE(fip.getAtmnrs().size() == numAtoms, "AtomNumbers size mismatch");
139
140 std::vector<double> positions(numAtoms * 3);
141 auto capnpPos = fip.getPos();
142 for (size_t i = 0; i < numAtoms * 3; ++i) {
143 positions[i] = capnpPos[i];
144 }
145
146 std::vector<int> atomicNrs(numAtoms);
147 auto capnpAtmnrs = fip.getAtmnrs();
148 for (size_t i = 0; i < numAtoms; ++i) {
149 atomicNrs[i] = capnpAtmnrs[i];
150 }
151
152 double box[9] = {};
153 auto capnpBox = fip.getBox();
154 for (size_t i = 0; i < 9 && i < capnpBox.size(); ++i) {
155 box[i] = capnpBox[i];
156 }
157
158 std::vector<double> forces(numAtoms * 3, 0.0);
159 double energy = 0.0;
160
161 std::lock_guard<std::mutex> lock(m_mutexes[idx]);
162 m_pool[idx](static_cast<long>(numAtoms), positions.data(), atomicNrs.data(),
163 forces.data(), &energy, box);
164
165 auto result = context.getResults();
166 auto pres = result.initResult();
167 pres.setEnergy(energy);
168
169 auto forcesList = pres.initForces(numAtoms * 3);
170 for (size_t i = 0; i < numAtoms * 3; ++i) {
171 forcesList.set(i, forces[i]);
172 }
173
174 return kj::READY_NOW;
175 }
176
177private:
178 std::vector<ForceCallback> m_pool;
179 std::vector<std::mutex> m_mutexes;
180 std::atomic<size_t> m_next;
181};
182
183} // anonymous namespace
184
185void eonc::startPooledRpcServer(std::vector<ForceCallback> pool,
186 const std::string &host, uint16_t port) {
187 EONC_LOG_INFO("Starting pooled RPC gateway on {}:{} with {} instances", host,
188 port, pool.size());
189
190 capnp::EzRpcServer server(kj::heap<PooledCallbackPotImpl>(std::move(pool)),
191 host, port);
192
193 auto &waitScope = server.getWaitScope();
194 EONC_LOG_INFO("Gateway ready on port {}. Ctrl+C to stop.", port);
195 kj::NEVER_DONE.wait(waitScope);
196}
#define EONC_LOG_INFO(...)
Definition EonLogger.h:250
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.