Loading...
Searching...
No Matches
SocketNWChemPot.cpp
Go to the documentation of this file.
2#include "eon/EonLogger.h"
3#include "eon/Parameters.h"
4
5#include <algorithm>
6#include <array>
7#include <cctype>
8#include <cerrno>
9#include <cstring>
10#include <format>
11#include <fstream>
12#include <stdexcept>
13
14#include <arpa/inet.h>
15#include <sys/socket.h>
16#include <sys/un.h>
17#include <unistd.h>
18
19#include <chrono>
20#include <cstdint>
21#include <readcon-core.hpp>
22#include <thread>
23
25 : eonc::Potential(eonc::PotType::SocketNWChem, p) {
26
31
32 if (unix_socket_mode) {
34 // NWChem's Fortran i-PI driver truncates the socket name to ~30 chars.
35 // The full path is /tmp/ipi_<basename>, so basename must be short.
37 if (server_address.size() > 30) {
39 "UNIX socket path '{}' is {} characters, past NWChem's ~30 character "
40 "limit. Shorten unix_socket_path (currently '{}') to at most {} "
41 "characters.",
43 throw std::runtime_error(
44 "unix_socket_path too long for NWChem (max ~21 chars, got " +
45 std::to_string(unix_socket_basename.size()) + ")");
46 }
47 port = -1;
48 EONC_LOG_INFO("SocketNWChemPot UNIX socket {}", server_address);
49 } else {
52 EONC_LOG_INFO("SocketNWChemPot TCP {}:{}", server_address, port);
53 }
54
56}
57
59 if (is_connected) {
60 EONC_LOG_INFO("Closing connection to NWChem client");
61 try {
62 send_header("EXIT");
63 } catch (...) {
64 // Ignore errors during shutdown
65 }
66 }
67 if (conn_fd >= 0)
68 ::close(conn_fd);
69 if (listen_fd >= 0)
70 ::close(listen_fd);
71 if (unix_socket_mode) {
72 ::unlink(server_address.c_str());
73 }
74}
75
76// =============================================
77// Public Methods
78// =============================================
79
81 const std::string &filename, long N,
82 const std::vector<std::string> &atom_symbols) {
83 std::ofstream outfile(filename);
84 if (!outfile.is_open()) {
85 throw std::runtime_error("Could not open file to write NWChem template: " +
86 filename);
87 }
88
89 outfile << "start nwchem_socket_job\n";
90 outfile << "title \"NWChem Server for eOn\"\n\n";
91 outfile << "memory " << mem_in_gb << " gb\n\n";
92 outfile << "geometry units bohr noautosym nocenter noautoz\n";
93 // This geometry block is only a template for memory allocation.
94 // The atom types and count are what matter.
95 for (long i = 0; i < N; ++i) {
96 outfile << " " << atom_symbols[i] << " 0.0 0.0 " << static_cast<double>(i)
97 << "\n";
98 }
99 outfile << "end\n\n";
100 outfile << "include " << nwchem_settings << "\n\n";
101 outfile << "driver\n";
102 if (unix_socket_mode) {
103 // For the NWChem input, we provide only the basename. NWChem adds the
104 // prefix.
105 outfile << " socket unix " << unix_socket_basename << "\n";
106 } else {
107 outfile << " socket ipi_client " << server_address << ":" << port << "\n";
108 }
109 outfile << "end\n\n";
110 outfile << "task scf optimize\n";
111
112 outfile.close();
113}
114
115void SocketNWChemPot::force(long N, const double *R, const int *atomicNrs,
116 double *F, double *U, double *variance,
117 const double *box) {
118 try {
119 forceOnce(N, R, atomicNrs, F, U, variance, box);
120 return;
121 } catch (const std::runtime_error &) {
123 }
124 forceOnce(N, R, atomicNrs, F, U, variance, box);
125}
126
127void SocketNWChemPot::forceOnce(long N, const double *R, const int *atomicNrs,
128 double *F, double *U, double *variance,
129 const double *box) {
130 if (!is_connected) {
131 std::vector<std::string> symbols;
132 symbols.reserve(N);
133 for (long i = 0; i < N; ++i) {
134 const int z = atomicNrs[i];
135 symbols.emplace_back(
136 z > 0 ? readcon::z_to_symbol(static_cast<uint64_t>(z)) : "X");
137 }
139 write_nwchem_template("nwchem_socket.nwi", N, symbols);
140 }
141
142 EONC_LOG_INFO("Waiting for NWChem client connection");
144 EONC_LOG_INFO("NWChem client connected");
145
146 // 1. eOn acts as the server: after accepting the NWChem client connection,
147 // eOn sends "STATUS" to the client to query its status.
148 std::array<char, MSG_LEN + 1> status_buffer{};
149 send_header("STATUS");
150
151 // 2. eOn (acting as server) then waits for NWChem (the client) to respond
152 // with "READY".
153 recv_header(status_buffer.data());
154 if (std::string(status_buffer.data()) != "READY") {
155 throw std::runtime_error(
156 "Handshake failed: NWChem client not READY. It sent: " +
157 std::string(status_buffer.data()));
158 }
159 EONC_LOG_INFO("NWChem server is connected and READY");
160 }
161
162 // Check status for this specific force call
163 std::array<char, MSG_LEN + 1> status_buffer{};
164 send_header("STATUS");
165 recv_header(status_buffer.data());
166
167 if (std::string(status_buffer.data()) == "NEEDINIT") {
168 send_header("INIT");
169 // Send dummy INIT payload (bead index, number of bytes in extra string)
170 std::array<int32_t, 2> init_payload{{0, 1}}; // bead_index=0, nbytes=1
171 char dummy_byte = 0;
172 send_exact(init_payload.data(), init_payload.size() * sizeof(int32_t));
173 send_exact(&dummy_byte,
174 sizeof(dummy_byte)); // No extra string (just a null terminator)
175 send_header("STATUS");
176 recv_header(status_buffer.data());
177 }
178
179 if (std::string(status_buffer.data()) != "READY") {
180 throw std::runtime_error("NWChem server not ready for new positions!");
181 }
182
183 // Convert positions to Bohr
184 std::vector<double> pos_bohr(N * 3);
185 for (size_t i = 0; i < pos_bohr.size(); ++i) {
186 pos_bohr[i] = R[i] / BOHR_IN_ANGSTROM;
187 }
188
189 // Per i-PI spec, cell and inverse cell must be sent. NWChem does not use
190 // this information for non-periodic calculations, so the frame carries an
191 // identity cell and its inverse.
192 std::array<double, 9> invcell_T{{1, 0, 0, 0, 1, 0, 0, 0, 1}};
193 std::array<double, 9> cell_T{{1, 0, 0, 0, 1, 0, 0, 0, 1}};
194
195 send_header("POSDATA");
196 int32_t nat = static_cast<int32_t>(N);
197 send_exact(cell_T.data(), cell_T.size() * sizeof(double));
198 send_exact(invcell_T.data(), invcell_T.size() * sizeof(double));
199 send_exact(&nat, sizeof(nat));
200 send_exact(pos_bohr.data(), pos_bohr.size() * sizeof(double));
201
202 // Poll for results
203 while (true) {
204 send_header("STATUS");
205 recv_header(status_buffer.data());
206 if (std::string(status_buffer.data()) == "HAVEDATA") {
207 break;
208 }
209 // A small sleep to prevent busy-waiting that consumes 100% CPU.
210 std::this_thread::sleep_for(std::chrono::milliseconds(10));
211 }
212
213 // Request and receive results ---
214 send_header("GETFORCE");
215 recv_header(status_buffer.data());
216 if (std::string(status_buffer.data()) != "FORCEREADY") {
217 throw std::runtime_error("Expected FORCEREADY, got " +
218 std::string(status_buffer.data()));
219 }
220
221 // Unpack the results payload.
222 double energy_ha;
223 int32_t nat_back;
224 std::vector<double> forces_ha_bohr(N * 3);
225 std::array<double, 9> virial_ha{};
226 int32_t extra_len;
227
228 recv_exact(&energy_ha, sizeof(energy_ha));
229 recv_exact(&nat_back, sizeof(nat_back));
230 if (nat_back != N)
231 throw std::runtime_error("Atom count mismatch from NWChem");
232 recv_exact(forces_ha_bohr.data(), forces_ha_bohr.size() * sizeof(double));
233 recv_exact(virial_ha.data(), virial_ha.size() * sizeof(double));
234 // i-PI virial is in Hartree and is positive when the system pushes
235 // outward, so sigma = (1/V) dE/dε = -virial / V. The buffer is
236 // symmetric, so row-major and column-major agree after averaging.
237 const double volume =
238 box == nullptr ? 0.0 : std::abs(Matrix3d::Map(box).determinant());
239 stress_.setZero();
240 haveStress_ = volume > 0.0;
241 if (haveStress_) {
242 const double scale = -HARTREE_IN_EV / volume;
243 for (int row = 0; row < 3; ++row) {
244 for (int col = 0; col < 3; ++col) {
245 const double vij =
246 0.5 * (virial_ha[static_cast<size_t>(row * 3 + col)] +
247 virial_ha[static_cast<size_t>(col * 3 + row)]);
248 stress_(row, col) = vij * scale;
249 }
250 }
251 }
252 recv_exact(&extra_len, sizeof(extra_len));
253 if (extra_len > 0) {
254 std::vector<char> extra_buf(extra_len);
255 recv_exact(extra_buf.data(), extra_len);
256 }
257
258 // Convert results back to eOn units (eV and Angstrom)
259 *U = energy_ha * HARTREE_IN_EV;
260 for (int i = 0; i < N * 3; ++i) {
261 F[i] = forces_ha_bohr[i] * (HARTREE_IN_EV / BOHR_IN_ANGSTROM);
262 }
263 if (variance != nullptr) {
264 *variance = 0.0;
265 }
266}
267
268// =================================================
269// Private Helper Methods for Socket Communication
270// =================================================
271
272namespace {
273
274[[noreturn]] void socket_fail(int fd, const char *what) {
275 const int err = errno;
276 if (fd >= 0) {
277 ::close(fd);
278 }
279 throw std::runtime_error(std::format("{}: {}", what, std::strerror(err)));
280}
281
282} // namespace
283
285 int domain = unix_socket_mode ? AF_UNIX : AF_INET;
286 listen_fd = socket(domain, SOCK_STREAM, 0);
287 if (listen_fd < 0) {
288 throw std::runtime_error(
289 std::format("Failed to create socket: {}", std::strerror(errno)));
290 }
291
292 if (unix_socket_mode) {
293 ::unlink(server_address.c_str()); // Remove stale socket file if it exists
294 sockaddr_un sock_addr{};
295 sock_addr.sun_family = AF_UNIX;
296 std::strncpy(sock_addr.sun_path, server_address.c_str(),
297 sizeof(sock_addr.sun_path) - 1);
298
299 socklen_t addr_len = static_cast<socklen_t>(
300 sizeof(sock_addr.sun_family) + std::strlen(sock_addr.sun_path));
301 if (::bind(listen_fd, reinterpret_cast<sockaddr *>(&sock_addr), addr_len) <
302 0) {
303 const int fd = listen_fd;
304 listen_fd = -1;
305 socket_fail(fd, "Failed to bind UNIX socket");
306 }
307 } else {
308 int opt = 1;
309 if (setsockopt(listen_fd, SOL_SOCKET, SO_REUSEADDR, &opt, sizeof(opt)) <
310 0) {
311 const int fd = listen_fd;
312 listen_fd = -1;
313 socket_fail(fd, "Failed to set SO_REUSEADDR");
314 }
315 sockaddr_in sock_addr{};
316 sock_addr.sin_family = AF_INET;
317 sock_addr.sin_addr.s_addr = inet_addr(server_address.c_str());
318 // inet_addr accepts only a dotted IPv4 address. A name fails here instead
319 // of being bound as INADDR_NONE.
320 if (sock_addr.sin_addr.s_addr == INADDR_NONE) {
321 ::close(listen_fd);
322 listen_fd = -1;
323 throw std::runtime_error(
324 "Failed to parse the TCP host as an IPv4 address");
325 }
326 sock_addr.sin_port = htons(static_cast<uint16_t>(port));
327
328 if (::bind(listen_fd, reinterpret_cast<sockaddr *>(&sock_addr),
329 sizeof(sock_addr)) < 0) {
330 const int fd = listen_fd;
331 listen_fd = -1;
332 socket_fail(fd, "Failed to bind TCP socket");
333 }
334 }
335
336 if (::listen(listen_fd, 1) < 0) {
337 const int fd = listen_fd;
338 listen_fd = -1;
339 socket_fail(fd, "Socket listen() failed");
340 }
341}
342
344 conn_fd = ::accept(listen_fd, nullptr, nullptr);
345 if (conn_fd < 0) {
346 throw std::runtime_error(std::format(
347 "Failed to accept client connection: {}", std::strerror(errno)));
348 }
349 is_connected = true;
350}
351
353 if (conn_fd >= 0) {
354 ::close(conn_fd);
355 conn_fd = -1;
356 }
357 is_connected = false;
358}
359
360void SocketNWChemPot::send_header(const char *msg) {
361 std::array<char, MSG_LEN> buffer{};
362 const std::size_t n = std::min(std::strlen(msg), buffer.size());
363 std::copy_n(msg, n, buffer.begin());
364 send_exact(buffer.data(), buffer.size());
365}
366
368 recv_exact(buffer, MSG_LEN);
369 buffer[MSG_LEN] = '\0'; // Null-terminate
370 // Trim trailing whitespace
371 for (int i = MSG_LEN - 1;
372 i >= 0 && std::isspace(static_cast<unsigned char>(buffer[i])); --i) {
373 buffer[i] = '\0';
374 }
375}
376
377void SocketNWChemPot::send_exact(const void *buffer, size_t n_bytes) {
378 size_t sent = 0;
379 while (sent < n_bytes) {
380 ssize_t n = ::send(conn_fd, static_cast<const char *>(buffer) + sent,
381 n_bytes - sent, 0);
382 if (n <= 0) {
383 throw std::runtime_error(
384 "send_exact failed: connection closed or error.");
385 }
386 sent += n;
387 }
388}
389
390void SocketNWChemPot::recv_exact(void *buffer, size_t n_bytes) {
391 size_t recvd = 0;
392 while (recvd < n_bytes) {
393 ssize_t n = ::recv(conn_fd, static_cast<char *>(buffer) + recvd,
394 n_bytes - recvd, 0);
395 if (n <= 0) {
396 throw std::runtime_error(
397 "recv_exact failed: connection closed or error.");
398 }
399 recvd += n;
400 }
401}
#define EONC_LOG_ERROR(...)
Definition EonLogger.h:261
#define EONC_LOG_INFO(...)
Definition EonLogger.h:249
std::string unix_socket_basename
~SocketNWChemPot() override
static constexpr int MSG_LEN
void send_header(const char *msg)
void recv_exact(void *buffer, size_t n_bytes)
void send_exact(const void *buffer, size_t n_bytes)
std::string server_address
static constexpr double BOHR_IN_ANGSTROM
void recv_header(char *buffer)
static constexpr double HARTREE_IN_EV
void write_nwchem_template(const std::string &filename, long N, const std::vector< std::string > &atom_symbols)
Generates an NWChem input file (.nwi) configured to connect to this server.
std::string nwchem_settings
void force(long N, const double *R, const int *atomicNrs, double *F, double *U, double *variance, const double *box) override
The method called to compute forces and energy.
SocketNWChemPot(const eonc::Parameters &p)
void forceOnce(long N, const double *R, const int *atomicNrs, double *F, double *U, double *variance, const double *box)
const socket_nwchem_options_t & socket_nwchem_options() const
Potential(PotType a_ptype)
Production default: construction-scope registry, else PotRegistry::get().
RAII resource manager for the ARTn C library with global synchronization.