diff --git a/docs/advanced/input_files/input-main.md b/docs/advanced/input_files/input-main.md index 567f0f87640..71eb9eaefd7 100644 --- a/docs/advanced/input_files/input-main.md +++ b/docs/advanced/input_files/input-main.md @@ -10,6 +10,7 @@ - [ntype](#ntype) - [cell\_replica](#cell_replica) - [calculation](#calculation) + - [socket\_driver](#socket_driver) - [esolver\_type](#esolver_type) - [symmetry](#symmetry) - [symmetry\_prec](#symmetry_prec) @@ -637,6 +638,20 @@ - test_neighbour: obtain information of neighboring atoms (for LCAO basis only), please specify a positive search_radius manually - **Default**: scf +### socket_driver + +- **Type**: Boolean +- **Description**: If set to True, ABACUS keeps the calculation type as scf and receives atomic positions from an external driver through the i-PI socket protocol. + + > Note: Use calculation = scf with socket_driver = True. ABACUS connects to the external i-PI server selected by ABACUS_SOCKET_ADDRESS. If ABACUS_SOCKET_ADDRESS is unset, ABACUS uses localhost:31415. The value can use one of two forms: + + - host:port, for example localhost:31415 or 127.0.0.1:31415, opens a TCP connection to that host and port. Use this when the i-PI server listens on a TCP port. + - path:UNIX, for example /tmp/ipi_abacus_si:UNIX, opens a Unix-domain socket at the given filesystem path. The :UNIX suffix tells ABACUS that the preceding value is a local socket path rather than a TCP host name. This form only works on the same machine. + When using the ASE AbacusSocketIO interface, this environment variable is set automatically from the port or unixsocket calculator argument. + + Socket mode always computes energy. Force and stress extraction follows cal_force and cal_stress independently; disabled properties are sent as protocol padding and marked absent in the ABACUS i-PI extras metadata, not reported as physical zero values. This metadata extension is required for safe optional-property handling: a legacy response with empty extras is accepted only for energy-only use, while a generic client that ignores extras cannot distinguish padding from a computed zero. A non-converged SCF step is returned with scf_converged=false metadata so an external driver can choose its policy. +- **Default**: False + ### esolver_type - **Type**: String @@ -686,6 +701,7 @@ - **Type**: Boolean - **Description**: If set to True, calculate the force at the end of the electronic iteration. + In socket_driver mode, this flag controls whether the returned frame advertises forces; it is not forced on by the socket protocol. - **Default**: False ### kpar @@ -801,6 +817,7 @@ - **Type**: Boolean - **Description**: If set to True, calculate the stress at the end of the electronic iteration. + In socket_driver mode, this flag independently controls whether the returned frame advertises stress/virial. - **Default**: False ### diago_proc @@ -899,7 +916,12 @@ ### chg_extrap - **Type**: String -- **Description**: Charge extrapolation method for MD and relaxation calculations. +- **Description**: Charge extrapolation method for MD, relaxation, and socket-driven calculations. + + When set to default, ABACUS chooses second-order for md, first-order for + relax/cell-relax and socket_driver calculations, and atomic for other calculations. Socket-driven + molecular dynamics can explicitly set second-order if the external driver + updates structures smoothly enough for second-order extrapolation. - **Default**: default ### nb2d diff --git a/docs/advanced/interface/ase.md b/docs/advanced/interface/ase.md index e9b7b062809..b7b02d3a279 100644 --- a/docs/advanced/interface/ase.md +++ b/docs/advanced/interface/ase.md @@ -103,6 +103,104 @@ In the new implementation, we limit the range of functionalties supported to mai Please read the examples in `interfaces/ASE_interface/examples/` for more details. +### Socket I/O with ASE + +#### When to use socket mode + +`AbacusSocketIO` is designed for a sequence of electronic-structure +evaluations in which the atomic positions change while the simulation context +remains fixed. Reuse one socket calculator only when the cell and periodic +boundary conditions, atom count and species, pseudopotentials and orbitals, +k-point sampling, spin settings, and other electronic-structure parameters do +not change. The socket session can then keep one ABACUS process alive and +receive successive position updates. + +This pattern is suitable for fixed-cell ASE optimization and molecular +dynamics, fixed-cell NEB (use an independent calculator/session for each image), +finite-displacement phonon or ASE finite-difference frequency calculations, +position-only P-RFO or transition-state searches, and repeated fixed-cell +force evaluations in larger workflows such as thermal-property or active- +learning data generation. These workflows can use the socket calculator only +when their driver calls the ASE calculator interface; the existing Phonopy, +ShengBTE, DP-GEN, or transition-state tools are not automatically converted +to socket workflows by installing abacuslite. + +Use the regular `Abacus` FileIO calculator when the cell, composition, or +electronic-structure settings must change. Direct DFPT or dynamical-matrix +calculations, and external workflows that require properties beyond energy, +forces, and stress, also remain outside the current socket property interface. + +For socket-driven ASE workflows, use the `AbacusSocketIO` calculator. ASE runs the i-PI socket server, while ABACUS keeps `calculation=scf` and is launched with `socket_driver=1` as the client. Energy, forces, and stress are independent properties controlled by `cal_force` and `cal_stress`; the fixed i-PI wire layout still contains padding fields, while extras metadata identifies which values were actually computed. See the [ASE socket I/O documentation](https://ase-lib.org/ase/calculators/socketio/socketio.html) and the i-PI reference paper, [Ceriotti et al., Comput. Phys. Commun. 185, 1019-1026 (2014)](https://doi.org/10.1016/j.cpc.2013.10.027), for the protocol background. + +Build ABACUS as usual before using this interface. PW-only builds work with `basis_type=pw`; LCAO socket calculations require an LCAO-enabled executable. No extra socket library is required. + +With CMake, choose the executable according to the basis: + +```bash +cmake -S . -B build-pw -DENABLE_MPI=ON -DENABLE_LCAO=OFF +cmake --build build-pw --target abacus_pw_para -j + +cmake -S . -B build-lcao -DENABLE_MPI=ON -DENABLE_LCAO=ON +cmake --build build-lcao --target abacus_basic_para -j +``` + +With the ABACUS toolchain workflow, build the normal ABACUS executable with LCAO support when `basis_type=lcao` is needed, then pass that executable to `AbacusProfile(command=...)`. The command can include an MPI launcher, for example `mpirun -np 4 /path/to/abacus`; ABACUS rank 0 opens the socket connection and broadcasts the i-PI data to the other ranks internally. On managed clusters, keep scheduler-specific launch options outside the calculator when possible and test the exact launcher command on a compute node. + +For PW calculations on CUDA/ROCm with multiple MPI ranks, use a k-point layout compatible with ABACUS' GPU parallelization. In practice, make sure each k-point pool contains one MPI rank; for example, a 4-rank PW GPU socket calculation should use at least four k-points so the default GPU `kpar` adjustment can assign one rank per pool. A one-k-point PW GPU job with several MPI ranks can fail in the PW GPU transform path; reduce the rank count or use a denser k-point mesh such as a smaller `kspacing`. + +The ASE interface can be installed from this repository with: + +```bash +cd interfaces/ASE_interface +pip install . +``` + +A minimal socket calculator setup is: + +```python +from ase.optimize import BFGS +from abacuslite import AbacusProfile, AbacusSocketIO + +aprof = AbacusProfile( + command="mpirun -np 4 /path/to/abacus", + pseudo_dir="/path/to/pseudopotentials", + orbital_dir="/path/to/orbitals", + omp_num_threads=1, +) + +abacus = AbacusSocketIO( + profile=aprof, + directory="socketio", + unixsocket="abacus_si", + pseudopotentials={"Si": "Si_ONCV_PBE-1.0.upf"}, + basissets={"Si": "Si_gga_8au_100Ry_2s2p1d.orb"}, + inp={"calculation": "scf", "basis_type": "lcao", "kspacing": 0.1}, +) + +with abacus as calc: + atoms.calc = calc + BFGS(atoms).run(fmax=0.05) +``` + +`AbacusSocketIO` sets `socket_driver=1` automatically. The adapter enables properties requested through ASE, restarting the client if a later request expands the active property set. Set `inp={'cal_force': 1}` and/or `inp={'cal_stress': 1}` when a fixed-cell optimizer, MD integrator, or stress evaluation client needs those properties. Energy is always available. The interface selects the socket endpoint and passes it to ABACUS through `ABACUS_SOCKET_ADDRESS`, so users normally do not set this environment variable by hand when using abacuslite. + +There are two endpoint styles: + +- `unixsocket="abacus_si"` uses a local Unix-domain socket. ASE creates and listens on `/tmp/ipi_abacus_si`; abacuslite launches ABACUS with `ABACUS_SOCKET_ADDRESS=/tmp/ipi_abacus_si:UNIX`. The `:UNIX` suffix is part of ABACUS' address syntax and means that `/tmp/ipi_abacus_si` is a filesystem socket path, not a TCP host. This is usually the best choice when ASE and ABACUS run on the same node because it avoids TCP port conflicts. +- `port=31415` uses a TCP socket. abacuslite launches ABACUS with `ABACUS_SOCKET_ADDRESS=localhost:31415`, meaning host `localhost` and TCP port `31415`. Use this style when the socket server should listen on a TCP port. If ABACUS is launched manually instead of through `AbacusSocketIO`, set `ABACUS_SOCKET_ADDRESS` yourself to the same `host:port` or `path:UNIX` endpoint. + +Calling `atoms.get_potential_energy()` does not force a force or stress calculation. If a requested property was disabled, ASE raises `PropertyNotImplementedError`; zero-filled i-PI padding is never treated as a physical result. When SCF does not converge, `AbacusSocketIO.last_scf_converged` is set to `False` and the caller decides whether to continue or stop. + +The ABACUS metadata extension is required to expose force/stress presence safely. If a legacy client returns an empty extras field, the adapter accepts only an energy-only response and refuses to infer forces or stress from the fixed-wire padding. Generic i-PI/ASE clients that ignore ABACUS extras cannot distinguish mandatory padding from a computed zero; use `AbacusSocketIO` or another metadata-aware client when requesting optional properties. When launching ABACUS with a generic client, explicitly set `cal_force=1` for force-driven workflows and `cal_stress=1` for stress evaluation; an omitted switch defaults to disabled. Such clients also need their own policy for unconverged SCF results. + +A socket calculator owns one ABACUS process initialized from one fixed `INPUT`/`STRU` setup. Reuse the same `AbacusSocketIO` instance only for position updates under the same electronic-structure settings and the same cell. Do not change `kpts`, `kspacing`, `nspin`, `basis_type`, `basissets`, pseudopotentials, species, atom count, cell, or other core `INPUT`/`STRU` parameters through an existing socket calculator; create a new `AbacusSocketIO` instance and a new ABACUS client process for those changes. `AbacusSocketIO` rejects cell changes before sending them to ABACUS, and the ABACUS socket driver also checks incoming POSDATA cells against the initial `STRU` cell and exits if they differ. + +In socket mode, ABACUS keeps one client process alive. All SCF evaluations produced by the same `AbacusSocketIO` instance are appended to the same `OUT.ABACUS/running_scf.log`, because the ABACUS calculation type remains `scf`. The authoritative per-step energy and force results are returned through the i-PI socket to ASE. Use ASE trajectory and optimizer log files, such as `BFGS(atoms, trajectory="opt.traj", logfile="opt.log")`, when each optimizer or MD step should be saved separately. Treat `running_scf.log` mainly as the ABACUS diagnostic log for the socket client, not as one independent FileIO result per structure. + +The i-PI protocol does not transmit element symbols. `AbacusSocketIO` therefore sorts the internal socket atoms with the same first-occurrence species grouping used when writing `STRU`, and maps returned forces back to the original ASE `Atoms` order. This avoids silent force/atom mismatches when structures are read from CIF, extxyz, POSCAR, or other formats whose atom order is not already grouped for ABACUS. Users should not manually reorder atoms for socket I/O; pass the physical ASE `Atoms` object directly to the calculator. + +A complete fixed-cell validation and benchmark example is available in `interfaces/ASE_interface/examples/socketio.py`. + ## SPAP Analysis [SPAP](https://github.com/chuanxun/StructurePrototypeAnalysisPackage) (Structure Prototype Analysis Package) is written by Dr. Chuanxun Su to analyze symmetry and compare similarity of large amount of atomic structures. The coordination characterization function (CCF) is used to @@ -114,3 +212,10 @@ If you use this program and method in your research, please read and cite the pu `Su C, Lv J, Li Q, Wang H, Zhang L, Wang Y, Ma Y. Construction of crystal structure prototype database: methods and applications. J Phys Condens Matter. 2017 Apr 26;29(16):165901.` and you should install it first with command `pip install spap`. + +Socket results are read directly from the completed in-memory solver frame, not +parsed from output files. The client clears cached results and convergence +metadata before a new request and publishes them only after validating the full +response. A failed request therefore leaves no previous-frame result available +in the calculator cache. This protocol guarantee does not establish SCF +convergence or numerical agreement with independent single-point calculations. diff --git a/docs/parameters.yaml b/docs/parameters.yaml index 3cc85903fc3..f82f1c89a77 100644 --- a/docs/parameters.yaml +++ b/docs/parameters.yaml @@ -46,6 +46,21 @@ parameters: default_value: scf unit: "" availability: "" + - name: socket_driver + category: System variables + type: Boolean + description: | + If set to True, ABACUS keeps the calculation type as scf and receives atomic positions from an external driver through the i-PI socket protocol. + + [NOTE] Use calculation = scf with socket_driver = True. ABACUS connects to the external i-PI server selected by ABACUS_SOCKET_ADDRESS. If ABACUS_SOCKET_ADDRESS is unset, ABACUS uses localhost:31415. The value can use one of two forms: + * host:port, for example localhost:31415 or 127.0.0.1:31415, opens a TCP connection to that host and port. Use this when the i-PI server listens on a TCP port. + * path:UNIX, for example /tmp/ipi_abacus_si:UNIX, opens a Unix-domain socket at the given filesystem path. The :UNIX suffix tells ABACUS that the preceding value is a local socket path rather than a TCP host name. This form only works on the same machine. + When using the ASE AbacusSocketIO interface, this environment variable is set automatically from the port or unixsocket calculator argument. + + Socket mode always computes energy. Force and stress extraction follows cal_force and cal_stress independently; disabled properties are sent as protocol padding and marked absent in the ABACUS i-PI extras metadata, not reported as physical zero values. This metadata extension is required for safe optional-property handling: a legacy response with empty extras is accepted only for energy-only use, while a generic client that ignores extras cannot distinguish padding from a computed zero. A non-converged SCF step is returned with scf_converged=false metadata so an external driver can choose its policy. + default_value: "False" + unit: "" + availability: "" - name: esolver_type category: System variables type: String @@ -102,6 +117,7 @@ parameters: type: Boolean description: | If set to True, calculate the force at the end of the electronic iteration. + In socket_driver mode, this flag controls whether the returned frame advertises forces; it is not forced on by the socket protocol. default_value: "False" unit: "" availability: "" @@ -230,6 +246,7 @@ parameters: type: Boolean description: | If set to True, calculate the stress at the end of the electronic iteration. + In socket_driver mode, this flag independently controls whether the returned frame advertises stress/virial. default_value: "False" unit: "" availability: "" @@ -350,7 +367,12 @@ parameters: category: System variables type: String description: | - Charge extrapolation method for MD and relaxation calculations. + Charge extrapolation method for MD, relaxation, and socket-driven calculations. + + When set to default, ABACUS chooses second-order for md, first-order for + relax/cell-relax and socket_driver calculations, and atomic for other calculations. Socket-driven + molecular dynamics can explicitly set second-order if the external driver + updates structures smoothly enough for second-order extrapolation. default_value: default unit: "" availability: "" diff --git a/interfaces/ASE_interface/README.md b/interfaces/ASE_interface/README.md index 824e69900de..543d4774836 100644 --- a/interfaces/ASE_interface/README.md +++ b/interfaces/ASE_interface/README.md @@ -7,15 +7,17 @@ abacuslite is a lightweight plugin for ABACUS (Atomic-orbital Based Ab-initio Co ### Key Features - **Lightweight Design**: Implemented as a plugin, no need to modify ASE core code -- **Version Compatibility**: No longer restricted to specific ASE versions, works with most ASE versions +- **Version Compatibility**: Supports ASE versions satisfying the package requirement `ase>=3.22` - **ASE Integration**: Uses ASE as the running platform, making ABACUS a callable calculator within it -- **Function Support**: Currently only supports SCF (Self-Consistent Field) functionality, returning energy, forces, stress, etc. +- **Function Support**: Provides SCF-based energy, force, and stress evaluations through ASE. ASE can use these evaluations for relaxation, molecular dynamics, NEB, band-structure, and density-of-states workflows. +- **Socket Support**: `AbacusSocketIO` provides fixed-cell i-PI socket calculations, with energy always available and forces/stress enabled independently when requested. ## Installation -Installation is very simple, just execute the following command in the project root directory: +Install the plugin from the ASE interface directory: ```bash +cd interfaces/ASE_interface pip install . ``` @@ -32,8 +34,10 @@ Please refer to the example scripts in the `examples` folder. Recommended learni 7. **constraintmd.py** - Constrained molecular dynamics simulation 8. **metadynamics.py** - Metadynamics simulation 9. **neb.py** - Nudged Elastic Band (NEB) calculation +10. **soc.py** - Noncollinear spin-orbit coupling calculation +11. **socketio.py** - Fixed-cell ASE optimization with `AbacusSocketIO`, running ABACUS as an i-PI socket client -More usage examples will be provided in future versions. +The regular `Abacus` calculator runs one ABACUS calculation for each ASE property evaluation. ASE controls the relaxation, molecular-dynamics, and other workflow steps. The socket calculator reuses one ABACUS process for position updates, while the cell and electronic-structure settings remain fixed for that calculator instance. ## Authors @@ -48,10 +52,10 @@ Thanks to the ABACUS development team for their support and contributions. ## License -[Fill in according to the actual project license] +The applicable license terms are provided in the repository [LICENSE](../../LICENSE). ## Contact If you have any questions or suggestions, please contact us through: -- GitHub: [deepmodeling/abacus-develop](https://github.com/deepmodeling/abacus-develop) \ No newline at end of file +- GitHub: [deepmodeling/abacus-develop](https://github.com/deepmodeling/abacus-develop) diff --git a/interfaces/ASE_interface/abacuslite/core.py b/interfaces/ASE_interface/abacuslite/core.py index e285db03385..fffa7b63dab 100644 --- a/interfaces/ASE_interface/abacuslite/core.py +++ b/interfaces/ASE_interface/abacuslite/core.py @@ -30,6 +30,7 @@ @author: Huang Yi-ke ''' +import json import os import re import shutil @@ -45,6 +46,7 @@ GenericFileIOCalculator, read_stdout ) +from ase.calculators.socketio import SocketIOCalculator from ase.atoms import Atoms from ase.dft.kpoints import BandPath from ase.io import read @@ -124,17 +126,35 @@ def __init__(self, @staticmethod def parse_version(stdout) -> str: - # up to the ABACUS version v3.9.0.17, the run of command - # `abacus --version` would returns the information organized - # in the following way: - # ABACUS version v3.9.0.17 - return re.match(r'ABACUS version (\S+)', stdout).group(1) + # MPI launchers may add informational lines before ABACUS output. + match = re.search(r'ABACUS version (\S+)', stdout or '') + if match is None: + raise RuntimeError( + 'Could not parse ABACUS version from command output. ' + 'Expected a line like "ABACUS version vX.Y.Z".' + ) + return match.group(1) def get_calculator_command(self, inputfile) -> List[str]: # because ABACUS run in the folder where there are INPUT files, so the # additional inputfile argument is not used. return [] + def socketio_argv_inet(self, port: Optional[int] = None) -> List[str]: + port = 31415 if port is None else port + return [ + 'env', + f'ABACUS_SOCKET_ADDRESS=localhost:{port}', + *self._split_command, + ] + + def socketio_argv_unix(self, socket: str) -> List[str]: + return [ + 'env', + f'ABACUS_SOCKET_ADDRESS=/tmp/ipi_{socket}:UNIX', + *self._split_command, + ] + def version(self) -> str: '''get the abacus version information''' cmd_ = [*self._split_command, '--version'] @@ -443,6 +463,17 @@ def __init__(self, directory=directory, ) + def write_input(self, atoms, properties=None, system_changes=None): + if properties is None: + properties = self.template.implemented_properties + self.template.write_input( + profile=self.profile, + directory=Path(self.directory), + atoms=atoms, + parameters=self.parameters, + properties=properties, + ) + @classmethod def restart(cls, profile=None, directory='.', **kwargs): '''instantiate one ABACUS calculator from an existing job directory, @@ -558,11 +589,475 @@ def band_structure(self, efermi=None): from ase.spectrum.band_structure import get_band_structure return get_band_structure(calc=self, reference=efermi) +class AbacusSocketIO(SocketIOCalculator): + """ASE socket I/O calculator that launches ABACUS as an i-PI client. + + A socket calculator owns one ABACUS process with one fixed INPUT/STRU + setup. The i-PI protocol can update positions, but electronic-structure + parameters such as k-points, spin, basis, pseudopotentials, and species + require a new calculator instance. Energy, forces, and stress are + independently controlled by ABACUS INPUT. The fixed-layout i-PI response + uses zero padding for absent fields and an extras metadata record so + padding is never exposed as a computed property. + """ + + def __init__(self, + profile=None, + directory='.', + port=None, + unixsocket=None, + timeout=None, + log=None, + **kwargs): + inp = dict(kwargs.pop('inp', {})) + self._property_constraints = {} + for keyword, property_name in (('cal_force', 'forces'), + ('cal_stress', 'stress')): + if keyword in inp: + self._property_constraints[property_name] = self._input_bool( + inp[keyword], keyword) + self.implemented_properties = [ + 'energy', 'free_energy', 'forces', 'stress'] + self._active_properties = None + self._last_socket_metadata = None + self.last_scf_converged = None + inp = self._socket_inp(inp) + self.abacus = Abacus( + profile=profile, + directory=directory, + inp=inp, + **kwargs, + ) + self._reference_cell = None + super().__init__( + port=port, + unixsocket=unixsocket, + timeout=timeout, + log=log, + launch_client=self._launch_client, + ) + + def calculate(self, atoms=None, properties=None, system_changes=None): + from ase.calculators.calculator import ( + PropertyNotImplementedError, + all_changes, + ) + from ase.stress import full_3x3_to_voigt_6_stress + + # A failed new request must not expose results or convergence from the + # previous geometry. Publish only a completely validated response. + self.results = {} + self.last_scf_converged = None + self._last_socket_metadata = None + if system_changes is None: + system_changes = all_changes + if atoms is None: + atoms = self.atoms + if atoms is None: + raise ValueError('AbacusSocketIO.calculate requires atoms') + + requested = self._normalize_socket_properties(properties) + self._check_requested_properties(requested) + + bad = [change for change in system_changes + if change not in self.supported_changes] + if self.atoms is not None and any(bad): + raise PropertyNotImplementedError( + 'Cannot change {} through IPI protocol. ' + 'Please create new socket calculator.' + .format(bad if len(bad) > 1 else bad[0])) + + desired = set(requested) + desired.discard('free_energy') + desired.add('energy') + for property_name, enabled in self._property_constraints.items(): + if enabled: + desired.add(property_name) + active = set(self._active_properties or ()) + if not active: + active.update(desired) + elif not desired.issubset(active): + active.update(desired) + if self.server is not None: + self._close_socket_session() + self._active_properties = tuple( + name for name in ('energy', 'forces', 'stress') if name in active) + + self._check_fixed_cell(atoms) + order = self._socket_sort_indices(atoms) + socket_atoms = atoms[order] + self.atoms = atoms.copy() + + if self.server is None: + self.server = self.launch_server() + proc = self.launch_client(socket_atoms, list(self._active_properties), + port=self._port, + unixsocket=self._unixsocket) + self.server.proc = proc + + raw_results = self.server.calculate(socket_atoms) + if not isinstance(raw_results, dict): + raise ValueError('ABACUS socket server returned a non-mapping result') + results = dict(raw_results) + metadata = self._decode_socket_metadata(results.pop('morebytes', None)) + if metadata is None: + if set(self._active_properties) != {'energy'}: + raise ValueError( + 'ABACUS socket response omitted property metadata; refusing ' + 'to infer forces or stress from fixed-wire padding') + present = {'energy'} + converged = None + else: + present = set(metadata['present']) + converged = metadata['scf_converged'] + + if 'energy' not in present or 'energy' not in results: + raise ValueError('ABACUS socket response did not provide energy') + energy = float(results['energy']) + if not np.isfinite(energy): + raise ValueError('ABACUS socket energy is not finite') + free_energy = float(results.get('free_energy', energy)) + if not np.isfinite(free_energy): + raise ValueError('ABACUS socket free energy is not finite') + current = {'energy': energy, 'free_energy': free_energy} + + if 'forces' in present: + if 'forces' not in results: + raise ValueError( + 'ABACUS socket metadata advertises forces, but wire response omitted them') + forces = np.asarray(results['forces'], dtype=np.float64) + expected_shape = (len(socket_atoms), 3) + if forces.shape != expected_shape or not np.all(np.isfinite(forces)): + raise ValueError('ABACUS socket forces have invalid shape or values') + current['forces'] = self._forces_to_input_order(forces, order) + + if 'stress' in present: + virial = results.get('virial') + if virial is None: + raise ValueError( + 'ABACUS socket metadata advertises stress, but wire response omitted virial') + if self.atoms.cell.rank != 3 or not any(self.atoms.pbc): + raise PropertyNotImplementedError( + 'ABACUS socket stress requires a periodic rank-3 cell') + virial = np.asarray(virial, dtype=np.float64) + if virial.shape != (3, 3) or not np.all(np.isfinite(virial)): + raise ValueError('ABACUS socket virial is not a finite 3x3 matrix') + vol = float(atoms.get_volume()) + if not np.isfinite(vol) or vol <= 0.0: + raise ValueError('ABACUS socket stress requires a positive cell volume') + current['stress'] = -full_3x3_to_voigt_6_stress(virial) / vol + + missing = [name for name in requested if name not in current] + if missing: + raise PropertyNotImplementedError( + 'ABACUS socket response did not provide requested {}'.format( + ', '.join(missing))) + self.results = current + self._last_socket_metadata = metadata + self.last_scf_converged = converged + + def _check_fixed_cell(self, atoms): + from ase.calculators.calculator import PropertyNotImplementedError + + cell = atoms.cell.array.copy() + if self._reference_cell is None: + self._reference_cell = cell + return + max_delta = np.max(np.abs(cell - self._reference_cell)) + if max_delta > 1.0e-10: + raise PropertyNotImplementedError( + 'AbacusSocketIO is fixed-cell only; create a new socket ' + 'calculator for a changed cell, or use the normal Abacus ' + 'FileIO calculator for variable-cell workflows.' + ) + + def set(self, **kwargs): + if kwargs: + raise ValueError( + 'AbacusSocketIO input parameters are fixed after construction; ' + 'create a new AbacusSocketIO calculator to change k-points, ' + 'spin, basis, pseudopotentials, species, or other INPUT/STRU ' + 'settings.' + ) + return super().set(**kwargs) + + def _check_requested_properties(self, requested): + from ase.calculators.calculator import PropertyNotImplementedError + + constraints = self._property_constraints + keywords = {'forces': 'cal_force', 'stress': 'cal_stress'} + for property_name, keyword in keywords.items(): + if property_name in requested and constraints.get(property_name) is False: + raise PropertyNotImplementedError( + '{}=0 disables requested {}'.format(keyword, property_name)) + + @staticmethod + def _normalize_socket_properties(properties): + from ase.calculators.calculator import PropertyNotImplementedError + + if properties is None: + names = ['energy'] + elif isinstance(properties, str): + names = [properties] + else: + names = list(properties) + if not names: + names = ['energy'] + allowed = {'energy', 'free_energy', 'forces', 'stress'} + unknown = [name for name in names if name not in allowed] + if unknown: + raise PropertyNotImplementedError( + 'ABACUS socket does not implement {}'.format(', '.join(unknown))) + return tuple(dict.fromkeys(names)) + + def _close_socket_session(self): + server = getattr(self, 'server', None) + if server is not None: + close = getattr(server, 'close', None) + if callable(close): + close() + self.server = None + self.results = {} + + @staticmethod + def _decode_socket_metadata(raw): + if raw is None: + return None + if isinstance(raw, str): + payload = raw.encode('utf-8') + elif isinstance(raw, (bytes, bytearray, memoryview)): + payload = bytes(raw) + else: + payload = np.asarray(raw, dtype=np.uint8).tobytes() + if not payload: + return None + try: + metadata = json.loads(payload.decode('utf-8')) + except (UnicodeDecodeError, ValueError) as error: + raise ValueError('ABACUS socket extras are not valid UTF-8 JSON') from error + if not isinstance(metadata, dict): + raise ValueError('ABACUS socket extras must be a JSON object') + if metadata.get('schema') != 'abacus.socket.properties.v1': + raise ValueError('unsupported ABACUS socket extras schema') + present = metadata.get('present') + if not isinstance(present, list): + raise ValueError('ABACUS socket extras present must be a list') + allowed = {'energy', 'forces', 'stress'} + if any(not isinstance(name, str) or name not in allowed for name in present): + raise ValueError('ABACUS socket extras contain an unknown property') + if 'energy' not in present: + raise ValueError('ABACUS socket extras must include energy') + scf_converged = metadata.get('scf_converged') + if not isinstance(scf_converged, bool): + raise ValueError('ABACUS socket extras scf_converged must be Boolean') + return { + 'present': tuple(dict.fromkeys(present)), + 'scf_converged': scf_converged, + } + + def _launch_client(self, atoms, properties=None, port=None, unixsocket=None): + from subprocess import Popen + + if properties is None: + properties = list(self._active_properties or ('energy',)) + properties = set(properties) + properties.discard('free_energy') + properties.add('energy') + # The i-PI response has fixed force/virial fields, but ABACUS must be + # told explicitly which expensive quantities to evaluate. Keep these + # switches synchronized with the session mask before writing INPUT. + self.abacus.parameters['cal_force'] = int('forces' in properties) + self.abacus.parameters['cal_stress'] = int('stress' in properties) + properties = [name for name in ('energy', 'forces', 'stress') + if name in properties] + + directory = Path(self.abacus.directory) + directory.mkdir(exist_ok=True, parents=True) + + if hasattr(self.abacus, 'write_inputfiles'): + self.abacus.write_inputfiles(atoms, properties) + else: + self.abacus.write_input(atoms, properties=properties) + + if unixsocket is not None: + argv = self.abacus.profile.socketio_argv_unix(socket=unixsocket) + else: + argv = self.abacus.profile.socketio_argv_inet(port=port) + + stdout = open(directory / self.abacus.template.outputname, 'w') + stderr = open(directory / self.abacus.template.errorname, 'w') + try: + return Popen(argv, cwd=directory, env=os.environ, + stdout=stdout, stderr=stderr) + finally: + stdout.close() + stderr.close() + + @staticmethod + def _socket_inp(inp): + inp = dict(inp) + calculation = inp.get('calculation', 'scf') + if calculation != 'scf': + raise ValueError('ABACUS socket I/O requires calculation="scf"') + for keyword in ('cal_force', 'cal_stress'): + if keyword in inp: + inp[keyword] = int(AbacusSocketIO._input_bool(inp[keyword], keyword)) + inp.update({ + 'calculation': 'scf', + 'socket_driver': 1, + }) + return inp + + @staticmethod + def _input_bool(value, name): + if isinstance(value, bool): + return value + if isinstance(value, int) and value in (0, 1): + return bool(value) + if isinstance(value, str): + normalized = value.strip().lower() + if normalized in ('true', '1'): + return True + if normalized in ('false', '0'): + return False + raise ValueError('{} must be one of true, false, 1, or 0'.format(name)) + + @staticmethod + def _socket_sort_indices(atoms): + return species_group_indices(atoms.get_chemical_symbols()) + + @staticmethod + def _forces_to_input_order(forces, order): + reordered = np.empty_like(forces) + for sorted_index, original_index in enumerate(order): + reordered[original_index] = forces[sorted_index] + return reordered + + class TestAbacusCalculator(unittest.TestCase): here = Path(__file__).parent pporb = here.parent.parent.parent / 'tests' / 'PP_ORB' + def test_socketio_species_order_mapping(self): + atoms = Atoms(symbols=['Si', 'O', 'C', 'Si', 'O', 'C']) + order = AbacusSocketIO._socket_sort_indices(atoms) + self.assertEqual(order, [0, 3, 1, 4, 2, 5]) + + socket_forces = np.arange(18).reshape(6, 3) + input_forces = AbacusSocketIO._forces_to_input_order( + socket_forces, order) + + expected = np.empty_like(socket_forces) + for sorted_index, original_index in enumerate(order): + expected[original_index] = socket_forces[sorted_index] + np.testing.assert_array_equal(input_forces, expected) + + def test_socketio_rejects_parameter_changes(self): + calc = object.__new__(AbacusSocketIO) + with self.assertRaisesRegex(ValueError, 'fixed after construction'): + calc.set(kpts={'mode': 'mp-sampling', 'nk': [2, 2, 2]}) + + def test_socketio_rejects_cell_changes(self): + from ase.calculators.calculator import PropertyNotImplementedError + + calc = object.__new__(AbacusSocketIO) + calc.atoms = Atoms('Si', cell=[5.0, 5.0, 5.0], pbc=True) + calc._reference_cell = calc.atoms.cell.array.copy() + + changed = calc.atoms.copy() + changed.cell[0, 0] = 5.1 + with self.assertRaisesRegex(PropertyNotImplementedError, 'fixed-cell'): + calc._check_fixed_cell(changed) + + def test_socketio_input_keeps_independent_property_switches(self): + self.assertEqual( + AbacusSocketIO._socket_inp({'cal_force': 0, 'cal_stress': 1}), + {'calculation': 'scf', 'socket_driver': 1, + 'cal_force': 0, 'cal_stress': 1}) + + def test_socketio_boolean_parser_rejects_ambiguous_values(self): + for value in ('yes', 'no', 'on', 'off', ''): + with self.assertRaisesRegex(ValueError, 'cal_force'): + AbacusSocketIO._input_bool(value, 'cal_force') + + def test_socketio_metadata_rejects_unknown_property(self): + metadata = json.dumps({ + 'schema': 'abacus.socket.properties.v1', + 'present': ['energy', 'charges'], + 'scf_converged': True, + }).encode('utf-8') + with self.assertRaisesRegex(ValueError, 'unknown property'): + AbacusSocketIO._decode_socket_metadata( + np.frombuffer(metadata, dtype=np.int8)) + + def test_socketio_metadata_marks_padding_absent(self): + metadata = json.dumps({ + 'schema': 'abacus.socket.properties.v1', + 'present': ['energy'], + 'scf_converged': False, + }).encode('utf-8') + decoded = AbacusSocketIO._decode_socket_metadata( + np.frombuffer(metadata, dtype=np.int8)) + self.assertEqual(decoded['present'], ('energy',)) + self.assertFalse(decoded['scf_converged']) + + def test_socketio_legacy_response_cannot_infer_force_from_padding(self): + class LegacyServer: + def calculate(self, atoms): + return { + 'energy': 1.0, + 'forces': np.zeros((len(atoms), 3)), + 'virial': np.zeros((3, 3)), + 'morebytes': b'', + } + + calc = object.__new__(AbacusSocketIO) + calc.variable_cell = False + calc._property_constraints = {} + calc._active_properties = ('energy', 'forces') + calc._reference_cell = None + calc.atoms = None + calc.server = LegacyServer() + atoms = Atoms('Si') + with self.assertRaisesRegex(ValueError, 'refusing to infer forces'): + calc.calculate(atoms=atoms, properties=('forces',), system_changes=()) + + def test_socketio_failed_response_clears_previous_frame(self): + class BrokenServer: + def calculate(self, atoms): + raise EOFError('incomplete frame') + + calc = object.__new__(AbacusSocketIO) + calc._property_constraints = {} + calc._active_properties = ('energy',) + calc._reference_cell = None + calc.atoms = None + calc.results = {'energy': 123.0} + calc.last_scf_converged = True + calc._last_socket_metadata = {'scf_converged': True} + calc.server = BrokenServer() + with self.assertRaises(EOFError): + calc.calculate(Atoms('Si'), properties=('energy',), system_changes=()) + self.assertEqual(calc.results, {}) + self.assertIsNone(calc.last_scf_converged) + self.assertIsNone(calc._last_socket_metadata) + + def test_socketio_requested_disabled_property_is_rejected(self): + calc = object.__new__(AbacusSocketIO) + calc._property_constraints = {'forces': False, 'stress': False} + from ase.calculators.calculator import PropertyNotImplementedError + with self.assertRaises(PropertyNotImplementedError): + calc._check_requested_properties(('forces',)) + + def test_parse_version_allows_launcher_noise(self): + stdout = 'launcher info\nABACUS version v3.11.0-beta6\n' + self.assertEqual(AbacusProfile.parse_version(stdout), 'v3.11.0-beta6') + + def test_parse_version_rejects_missing_version(self): + with self.assertRaisesRegex(RuntimeError, 'ABACUS version'): + AbacusProfile.parse_version('launcher failed before abacus started') + def test_calculator_results(self): from ase.build.bulk import bulk silicon = bulk('Si', crystalstructure='diamond', a=5.43) diff --git a/interfaces/ASE_interface/examples/socketio.py b/interfaces/ASE_interface/examples/socketio.py new file mode 100644 index 00000000000..e3087453ffc --- /dev/null +++ b/interfaces/ASE_interface/examples/socketio.py @@ -0,0 +1,157 @@ +""" +This example validates and benchmarks ABACUS socket I/O from ASE. + +ASE runs as the i-PI socket server and ABACUS runs as the socket client. +ABACUS keeps calculation=scf and enables socket_driver internally. + +The script checks two PR-review relevant points: +1. socket SCF gives the same energy and forces as a normal non-socket SCF; +2. repeated socket calculations avoid relaunching ABACUS and are faster than + the normal FileIO calculator for a sequence of SCF force evaluations. + +The i-PI protocol does not carry element symbols. AbacusSocketIO handles +the required STRU/socket atom-order alignment internally and returns forces +in the original ASE Atoms order. +""" +import os +import shutil +import time +from pathlib import Path + +import numpy as np +from ase import Atoms +from abacuslite import Abacus, AbacusProfile, AbacusSocketIO + +here = Path(__file__).parent +pporb = here.parent.parent.parent / 'tests' / 'PP_ORB' + +aprof = AbacusProfile( + command=os.environ.get('ABACUS_COMMAND', 'mpirun -np 4 abacus'), + pseudo_dir=pporb, + orbital_dir=pporb, + omp_num_threads=1, +) + +common_kwargs = { + 'pseudopotentials': {'Si': 'Si_ONCV_PBE-1.0.upf'}, + 'basissets': {'Si': 'Si_gga_8au_100Ry_2s2p1d.orb'}, + 'inp': { + 'calculation': 'scf', + 'nspin': 1, + 'basis_type': 'lcao', + 'ks_solver': 'scalapack_gvx', + 'ecutwfc': 30, + 'symmetry': 0, + 'kspacing': 0.5, + 'scf_thr': 1e-8, + 'scf_nmax': 40, + 'chg_extrap': 'atomic', + 'cal_force': 1, + }, +} + +base_atoms = Atoms( + 'Si2', + positions=[[0.0, 0.0, 0.0], [1.25, 1.25, 1.25]], + cell=[5.43, 5.43, 5.43], + pbc=True, +) + + +def clean(directory): + shutil.rmtree(directory, ignore_errors=True) + + +def run_fileio(atoms, directory): + clean(directory) + calc = Abacus(profile=aprof, directory=str(directory), **common_kwargs) + atoms = atoms.copy() + atoms.calc = calc + forces = atoms.get_forces() + energy = atoms.get_potential_energy() + return energy, forces + + +def run_socketio(atoms, directory, socket_name): + clean(directory) + calc = AbacusSocketIO( + profile=aprof, + directory=str(directory), + unixsocket=socket_name, + timeout=120, + **common_kwargs, + ) + atoms = atoms.copy() + with calc: + atoms.calc = calc + energy = atoms.get_potential_energy() + forces = atoms.get_forces() + return energy, forces + + +def displaced_structures(): + structures = [] + for scale in (0.00, 0.03, -0.02, 0.05): + atoms = base_atoms.copy() + atoms.positions[1] += scale + structures.append(atoms) + return structures + + +fileio_dir = here / 'socketio_fileio_scf' +socket_dir = here / 'socketio_socket_scf' +bench_fileio_dir = here / 'socketio_bench_fileio' +bench_socket_dir = here / 'socketio_bench_socket' + +try: + reference_energy, reference_forces = run_fileio(base_atoms, fileio_dir) + socket_energy, socket_forces = run_socketio( + base_atoms, socket_dir, 'abacus_si_check') + + energy_diff = abs(socket_energy - reference_energy) + force_diff = np.max(np.abs(socket_forces - reference_forces)) + print(f'FileIO SCF energy: {reference_energy:.12f} eV') + print(f'Socket SCF energy: {socket_energy:.12f} eV') + print(f'|dE|: {energy_diff:.3e} eV') + print(f'max |dF|: {force_diff:.3e} eV/Angstrom') + assert energy_diff < 1e-4 + assert force_diff < 1e-5 + + structures = displaced_structures() + + clean(bench_fileio_dir) + fileio_calc = Abacus( + profile=aprof, + directory=str(bench_fileio_dir), + **common_kwargs, + ) + t0 = time.perf_counter() + for atoms in structures: + atoms = atoms.copy() + atoms.calc = fileio_calc + atoms.get_forces() + fileio_seconds = time.perf_counter() - t0 + + clean(bench_socket_dir) + socket_calc = AbacusSocketIO( + profile=aprof, + directory=str(bench_socket_dir), + unixsocket='abacus_si_bench', + timeout=120, + **common_kwargs, + ) + t0 = time.perf_counter() + with socket_calc as calc: + for atoms in structures: + atoms = atoms.copy() + atoms.calc = calc + atoms.get_forces() + socket_seconds = time.perf_counter() - t0 + + speedup = fileio_seconds / socket_seconds + print(f'FileIO repeated SCF force time: {fileio_seconds:.2f} s') + print(f'Socket repeated SCF force time: {socket_seconds:.2f} s') + print(f'Socket speedup vs FileIO: {speedup:.2f}x') +finally: + for directory in (fileio_dir, socket_dir, bench_fileio_dir, bench_socket_dir): + clean(directory) diff --git a/source/Makefile.Objects b/source/Makefile.Objects index dd2e861167b..c3ccb41ca0f 100644 --- a/source/Makefile.Objects +++ b/source/Makefile.Objects @@ -535,6 +535,9 @@ OBJS_PW=fft_bundle.o\ pw_op.o\ OBJS_RELAXATION=relax_data.o\ + socket_ipi.o\ + socket_frame.o\ + socket_driver.o\ cg_base.o\ bfgs_basic.o\ relax_driver.o\ diff --git a/source/source_esolver/esolver_ks.cpp b/source/source_esolver/esolver_ks.cpp index dec3fdd5770..55216e2ee7e 100644 --- a/source/source_esolver/esolver_ks.cpp +++ b/source/source_esolver/esolver_ks.cpp @@ -170,6 +170,7 @@ void ESolver_KS::runner(BaseCell& basecell, const int istep) // 7) after scf this->after_scf(ucell, istep, conv_esolver); + this->conv_esolver = conv_esolver; ModuleBase::timer::end(this->classname, "runner"); return; diff --git a/source/source_esolver/esolver_ks_lcao.cpp b/source/source_esolver/esolver_ks_lcao.cpp index 880509021f5..f65eb1f5376 100644 --- a/source/source_esolver/esolver_ks_lcao.cpp +++ b/source/source_esolver/esolver_ks_lcao.cpp @@ -583,7 +583,7 @@ void ESolver_KS_LCAO::after_scf(UnitCell& ucell, const int istep, const this->orb_, this->pw_wfc, this->pw_rho, this->pw_big, this->sf, this->pw_rhod, this->locpp.vloc, this->solvent, this->rdmft_solver, this->deepks, this->exx_nao, this->exx_info_, - this->conv_esolver, this->scf_nmax_flag, istep); + conv_esolver, this->scf_nmax_flag, istep); //! 3) Clean up RA, which is used to serach for adjacent atoms if (!this->inp_->cal_force && !this->inp_->cal_stress) diff --git a/source/source_hsolver/kernels/cuda/diag_cusolvermp.cu b/source/source_hsolver/kernels/cuda/diag_cusolvermp.cu index 6d184fd59ef..34c97b20508 100644 --- a/source/source_hsolver/kernels/cuda/diag_cusolvermp.cu +++ b/source/source_hsolver/kernels/cuda/diag_cusolvermp.cu @@ -1,6 +1,7 @@ #ifdef __CUSOLVERMP #include "diag_cusolvermp.cuh" #include "source_base/module_device/device_check.h" +#include "source_base/global_function.h" #include diff --git a/source/source_io/module_parameter/input_parameter.h b/source/source_io/module_parameter/input_parameter.h index 1637871f4d1..7dae21a97f3 100644 --- a/source/source_io/module_parameter/input_parameter.h +++ b/source/source_io/module_parameter/input_parameter.h @@ -19,6 +19,7 @@ struct Input_para std::string calculation = "scf"; ///< "scf" : self consistent calculation. ///< "nscf" : non-self consistent calculation. ///< "relax" : cell relaxations + bool socket_driver = false; ///< run ABACUS as an i-PI socket client std::string esolver_type = "ksdft"; ///< the energy solver: ksdft, sdft, ofdft, tddft, lj, dp /* symmetry level: -1, no symmetry at all; diff --git a/source/source_io/module_parameter/read_inp_sys.cpp b/source/source_io/module_parameter/read_inp_sys.cpp index d3ef1756ee4..795a6208c99 100644 --- a/source/source_io/module_parameter/read_inp_sys.cpp +++ b/source/source_io/module_parameter/read_inp_sys.cpp @@ -168,6 +168,30 @@ void ReadInput::item_system() sync_string(input.calculation); this->add_item(item); } + { + Input_Item item("socket_driver"); + item.annotation = "run as a socket client for external drivers using the i-PI protocol"; + item.category = "System variables"; + item.type = "Boolean"; + item.description = R"(If set to True, ABACUS keeps the calculation type as scf and receives atomic positions from an external driver through the i-PI socket protocol. + +[NOTE] Use calculation = scf with socket_driver = True. ABACUS connects to the external i-PI server selected by ABACUS_SOCKET_ADDRESS. If ABACUS_SOCKET_ADDRESS is unset, ABACUS uses localhost:31415. The value can use one of two forms: +* host:port, for example localhost:31415 or 127.0.0.1:31415, opens a TCP connection to that host and port. Use this when the i-PI server listens on a TCP port. +* path:UNIX, for example /tmp/ipi_abacus_si:UNIX, opens a Unix-domain socket at the given filesystem path. The :UNIX suffix tells ABACUS that the preceding value is a local socket path rather than a TCP host name. This form only works on the same machine. +When using the ASE AbacusSocketIO interface, this environment variable is set automatically from the port or unixsocket calculator argument.)"; + item.description += R"( + +Socket mode always computes energy. Force and stress extraction follows cal_force and cal_stress independently; disabled properties are sent as protocol padding and marked absent in the ABACUS i-PI extras metadata, not reported as physical zero values. This metadata extension is required for safe optional-property handling: a legacy response with empty extras is accepted only for energy-only use, while a generic client that ignores extras cannot distinguish padding from a computed zero. A non-converged SCF step is returned with scf_converged=false metadata so an external driver can choose its policy.)"; + item.default_value = "False"; + read_sync_bool(input.socket_driver); + item.check_value = [](const Input_Item& item, const Parameter& para) { + if (para.input.socket_driver && para.input.calculation != "scf") + { + ModuleBase::WARNING_QUIT("ReadInput", "socket_driver is only supported with calculation = scf."); + } + }; + this->add_item(item); + } { Input_Item item("esolver_type"); item.annotation = "the energy solver: ksdft, sdft, ofdft, tdofdft, tddft, lj, dp, ks-lr, lr, dfpt"; @@ -302,7 +326,8 @@ void ReadInput::item_system() item.annotation = "if calculate the force at the end of the electronic iteration"; item.category = "System variables"; item.type = "Boolean"; - item.description = "If set to True, calculate the force at the end of the electronic iteration."; + item.description = R"(If set to True, calculate the force at the end of the electronic iteration. +In socket_driver mode, this flag controls whether the returned frame advertises forces; it is not forced on by the socket protocol.)"; item.default_value = "False"; item.reset_value = [](const Input_Item& item, Parameter& para) { std::vector use_force = {"cell-relax", "relax", "md"}; @@ -606,7 +631,8 @@ For `basis_type=lcao_in_pw`, `init_wfc` is automatically set to `nao`. item.annotation = "calculate the stress or not"; item.category = "System variables"; item.type = "Boolean"; - item.description = "If set to True, calculate the stress at the end of the electronic iteration."; + item.description = R"(If set to True, calculate the stress at the end of the electronic iteration. +In socket_driver mode, this flag independently controls whether the returned frame advertises stress/virial.)"; item.default_value = "False"; item.reset_value = [](const Input_Item& item, Parameter& para) { if (para.input.calculation == "md") @@ -938,7 +964,12 @@ Available options are: item.annotation = "atomic; first-order; second-order; dm:coefficients of SIA"; item.category = "System variables"; item.type = "String"; - item.description = "Charge extrapolation method for MD and relaxation calculations."; + item.description = R"(Charge extrapolation method for MD, relaxation, and socket-driven calculations. + +When set to default, ABACUS chooses second-order for md, first-order for +relax/cell-relax and socket_driver calculations, and atomic for other calculations. Socket-driven +molecular dynamics can explicitly set second-order if the external driver +updates structures smoothly enough for second-order extrapolation.)"; item.default_value = "default"; read_sync_string(input.chg_extrap); item.reset_value = [](const Input_Item& item, Parameter& para) { @@ -947,7 +978,7 @@ Available options are: para.input.chg_extrap = "second-order"; } else if (para.input.chg_extrap == "default" - && (para.input.calculation == "relax" || para.input.calculation == "cell-relax")) + && (para.input.calculation == "relax" || para.input.calculation == "cell-relax" || para.input.socket_driver)) { para.input.chg_extrap = "first-order"; } diff --git a/source/source_io/test_serial/read_input_item_test.cpp b/source/source_io/test_serial/read_input_item_test.cpp index d581263ae68..399b21f810e 100644 --- a/source/source_io/test_serial/read_input_item_test.cpp +++ b/source/source_io/test_serial/read_input_item_test.cpp @@ -113,6 +113,27 @@ TEST_F(InputTest, Item_test) EXPECT_EXIT(it->second.check_value(it->second, param), ::testing::ExitedWithCode(1), ""); output = testing::internal::GetCapturedStdout(); EXPECT_THAT(output, testing::HasSubstr("NOTICE")); + + param.input.calculation = "socket"; + testing::internal::CaptureStdout(); + EXPECT_EXIT(it->second.check_value(it->second, param), ::testing::ExitedWithCode(1), ""); + output = testing::internal::GetCapturedStdout(); + EXPECT_THAT(output, testing::HasSubstr("NOTICE")); + } + + { // socket_driver + auto it = find_label("socket_driver", readinput.input_lists); + param.input.socket_driver = true; + param.input.calculation = "nscf"; + testing::internal::CaptureStdout(); + EXPECT_EXIT(it->second.check_value(it->second, param), ::testing::ExitedWithCode(1), ""); + output = testing::internal::GetCapturedStdout(); + EXPECT_THAT(output, testing::HasSubstr("NOTICE")); + + param.input.socket_driver = true; + param.input.calculation = "scf"; + EXPECT_NO_THROW(it->second.check_value(it->second, param)); + param.input.socket_driver = false; } { // esolver_type @@ -287,6 +308,7 @@ TEST_F(InputTest, Item_test) auto it = find_label("cal_force", readinput.input_lists); param.input.calculation = "cell-relax"; param.input.cal_force = false; + param.input.socket_driver = false; it->second.reset_value(it->second, param); EXPECT_EQ(param.input.cal_force, true); @@ -294,6 +316,13 @@ TEST_F(InputTest, Item_test) param.input.cal_force = true; it->second.reset_value(it->second, param); EXPECT_EQ(param.input.cal_force, false); + + param.input.calculation = "scf"; + param.input.socket_driver = true; + param.input.cal_force = false; + it->second.reset_value(it->second, param); + EXPECT_EQ(param.input.cal_force, false); + param.input.socket_driver = false; } { // ecutrho auto it = find_label("ecutrho", readinput.input_lists); @@ -403,8 +432,15 @@ TEST_F(InputTest, Item_test) it->second.reset_value(it->second, param); EXPECT_EQ(param.input.chg_extrap, "first-order"); + param.input.chg_extrap = "default"; + param.input.calculation = "scf"; + param.input.socket_driver = true; + it->second.reset_value(it->second, param); + EXPECT_EQ(param.input.chg_extrap, "first-order"); + param.input.chg_extrap = "default"; param.input.calculation = "none"; + param.input.socket_driver = false; it->second.reset_value(it->second, param); EXPECT_EQ(param.input.chg_extrap, "atomic"); diff --git a/source/source_relax/CMakeLists.txt b/source/source_relax/CMakeLists.txt index fde71a0f37e..b32dda901cc 100644 --- a/source/source_relax/CMakeLists.txt +++ b/source/source_relax/CMakeLists.txt @@ -2,6 +2,9 @@ add_library( relax OBJECT relax_data.cpp + socket_ipi.cpp + socket_frame.cpp + socket_driver.cpp cg_base.cpp relax_driver.cpp relax_sync.cpp diff --git a/source/source_relax/relax_driver.cpp b/source/source_relax/relax_driver.cpp index 1d2fd132eed..f7a91032355 100644 --- a/source/source_relax/relax_driver.cpp +++ b/source/source_relax/relax_driver.cpp @@ -1,4 +1,5 @@ #include "relax_driver.h" +#include "socket_driver.h" #include "source_base/formatter.h" #include "source_base/global_file.h" #include "source_base/version.h" @@ -21,6 +22,14 @@ void Relax_Driver::relax_driver( ModuleBase::TITLE("Relax_Driver", "relax_driver"); ModuleBase::timer::start("Relax_Driver", "relax_driver"); + if (inp.socket_driver) + { + Socket_Driver socket_driver; + socket_driver.socket_driver(p_esolver, ucell, inp, ofs_running); + ModuleBase::timer::end("Relax_Driver", "relax_driver"); + return; + } + this->init_relax(ucell.nat, inp); // steps[0]: istep (main iteration step) diff --git a/source/source_relax/socket_driver.cpp b/source/source_relax/socket_driver.cpp new file mode 100644 index 00000000000..a5bf94cde21 --- /dev/null +++ b/source/source_relax/socket_driver.cpp @@ -0,0 +1,875 @@ +#include "socket_driver.h" + +#include "source_relax/socket_ipi.h" +#include "source_relax/socket_frame.h" +#include "source_base/global_function.h" +#include "source_base/mathzone.h" +#include "source_base/parallel_common.h" +#include "source_base/timer.h" +#include "source_cell/unitcell.h" +#include "source_cell/update_cell.h" +#include "source_esolver/esolver.h" +#include "source_io/module_parameter/input_parameter.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace +{ +constexpr double RY_TO_HARTREE = 0.5; +constexpr int IPI_RANK_ROOT = 0; +constexpr double MAX_CELL_CONDITION = 1.0e12; +constexpr double INVERSE_ABSOLUTE_TOLERANCE + = 64.0 * std::numeric_limits::epsilon(); +constexpr double INVERSE_RELATIVE_TOLERANCE = 64.0; +constexpr double STRESS_ABSOLUTE_TOLERANCE = 1.0e-10; +constexpr double STRESS_RELATIVE_TOLERANCE = 1.0e-8; +constexpr std::int32_t MAX_INIT_BYTES = INT32_C(1048576); + +enum class DriverState +{ + NeedInit, + Ready, + HasData +}; + +struct ComputedFrame +{ + bool valid = false; + bool forces_present = false; + bool stress_present = false; + bool scf_converged = true; + double energy_hartree = 0.0; + std::vector forces_hartree_per_bohr; + SocketFrame::Matrix9 virial_wire_hartree = {{0.0}}; +}; + +bool all_ranks_converged(const bool local_converged) +{ + int converged = local_converged ? 1 : 0; +#ifdef __MPI + MPI_Allreduce(MPI_IN_PLACE, &converged, 1, MPI_INT, MPI_MIN, MPI_COMM_WORLD); +#endif + return converged != 0; +} + +void throw_if_any_rank_failed(int local_failed, std::string local_message) +{ + int any_failed = local_failed; +#ifdef __MPI + MPI_Allreduce(MPI_IN_PLACE, &any_failed, 1, MPI_INT, MPI_MAX, MPI_COMM_WORLD); +#endif + if (any_failed != 0) + { + if (local_message.empty()) + { + local_message = "socket frame validation failed on another MPI rank"; + } + throw std::runtime_error(local_message); + } +} + +[[noreturn]] void fail_during_collective_stage(const char* stage, + const std::string& message) +{ +#ifdef __MPI + int rank = -1; + MPI_Comm_rank(MPI_COMM_WORLD, &rank); + std::fprintf(stderr, + "ABACUS_SOCKET_MPI_FATAL stage=%s rank=%d message=%s\n", + stage, + rank, + message.c_str()); + std::fflush(stderr); + MPI_Abort(MPI_COMM_WORLD, EXIT_FAILURE); + std::abort(); +#else + (void)stage; + throw std::runtime_error(message); +#endif +} + +std::string properties_extra(const ComputedFrame& frame) +{ + std::ostringstream extra; + extra << "{\"schema\":\"abacus.socket.properties.v1\",\"present\":[\"energy\""; + if (frame.forces_present) + { + extra << ",\"forces\""; + } + if (frame.stress_present) + { + extra << ",\"stress\""; + } + extra << "],\"scf_converged\":" + << (frame.scf_converged ? "true" : "false") << "}"; + return extra.str(); +} + +bool is_root() +{ +#ifdef __MPI + int rank = IPI_RANK_ROOT; + MPI_Comm_rank(MPI_COMM_WORLD, &rank); + return rank == IPI_RANK_ROOT; +#else + return true; +#endif +} + +void bcast_double_vector(std::vector& values) +{ +#ifdef __MPI + if (!values.empty()) + { + Parallel_Common::bcast_double(values.data(), static_cast(values.size())); + } +#else + (void)values; +#endif +} + +void bcast_socket_int(int& value) +{ +#ifdef __MPI + Parallel_Common::bcast_int(value); +#else + (void)value; +#endif +} + +void bcast_socket_int32(std::int32_t& value) +{ +#ifdef __MPI + MPI_Bcast(&value, 1, MPI_INT32_T, IPI_RANK_ROOT, MPI_COMM_WORLD); +#else + (void)value; +#endif +} + +void bcast_socket_chars(char* value, const int size) +{ +#ifdef __MPI + Parallel_Common::bcast_char(value, size); +#else + (void)value; + (void)size; +#endif +} + +void bcast_socket_string(std::string& value) +{ + int size = static_cast(value.size()); + bcast_socket_int(size); + if (!is_root()) + { + value.resize(static_cast(size)); + } + if (size > 0) + { + bcast_socket_chars(&value[0], size); + } +} + +void quit_if_root_io_failed(int root_failed, std::string root_message) +{ + bcast_socket_int(root_failed); + bcast_socket_string(root_message); + if (root_failed != 0) + { + ModuleBase::WARNING_QUIT("ABACUS socket", root_message.empty() ? "i-PI socket I/O failed" : root_message); + } +} + +std::string bcast_header(std::string header) +{ + bcast_socket_string(header); + return header; +} + +std::string socket_address() +{ + const char* env = std::getenv("ABACUS_SOCKET_ADDRESS"); + if (env == nullptr || std::string(env).empty()) + { + return "localhost:31415"; + } + return std::string(env); +} + +std::vector ipi_cell_bohr_from_unitcell(const UnitCell& ucell) +{ + const double lat0 = ucell.lat0; + // ASE/i-PI sends POSDATA cell as cell.T in C order. ABACUS stores + // lattice vectors as rows in latvec, so use the transposed order here. + return { + ucell.latvec.e11 * lat0, ucell.latvec.e21 * lat0, ucell.latvec.e31 * lat0, + ucell.latvec.e12 * lat0, ucell.latvec.e22 * lat0, ucell.latvec.e32 * lat0, + ucell.latvec.e13 * lat0, ucell.latvec.e23 * lat0, ucell.latvec.e33 * lat0, + }; +} + +double max_wrapped_direct_delta_from_unitcell(const UnitCell& ucell, const std::vector& positions_bohr) +{ + if (positions_bohr.size() != static_cast(3 * ucell.nat)) + { + return 1.0e99; + } + + double out = 0.0; + int iat = 0; + for (int it = 0; it < ucell.ntype; ++it) + { + const Atom* atom = &ucell.atoms[it]; + for (int ia = 0; ia < atom->na; ++ia) + { + const double tau_x = positions_bohr[3 * iat + 0] / ucell.lat0; + const double tau_y = positions_bohr[3 * iat + 1] / ucell.lat0; + const double tau_z = positions_bohr[3 * iat + 2] / ucell.lat0; + + double dx = 0.0; + double dy = 0.0; + double dz = 0.0; + ModuleBase::Mathzone::Cartesian_to_Direct(tau_x, + tau_y, + tau_z, + ucell.latvec.e11, + ucell.latvec.e12, + ucell.latvec.e13, + ucell.latvec.e21, + ucell.latvec.e22, + ucell.latvec.e23, + ucell.latvec.e31, + ucell.latvec.e32, + ucell.latvec.e33, + dx, + dy, + dz); + + double ddx = dx - atom->taud[ia].x; + double ddy = dy - atom->taud[ia].y; + double ddz = dz - atom->taud[ia].z; + ddx -= std::round(ddx); + ddy -= std::round(ddy); + ddz -= std::round(ddz); + out = std::max(out, std::abs(ddx)); + out = std::max(out, std::abs(ddy)); + out = std::max(out, std::abs(ddz)); + ++iat; + } + } + return out; +} + +double max_abs_delta(const std::vector& a, const std::vector& b) +{ + if (a.size() != b.size()) + { + return 1.0e99; + } + double out = 0.0; + for (std::size_t i = 0; i < a.size(); ++i) + { + out = std::max(out, std::abs(a[i] - b[i])); + } + return out; +} + +double unchanged_cell_tolerance(const SocketFrame::Matrix9& cell) +{ + double maximum = 0.0; + for (std::size_t index = 0; index < cell.size(); ++index) + { + maximum = std::max(maximum, std::fabs(cell[index])); + } + return 32.0 * std::numeric_limits::epsilon() * std::max(1.0, maximum); +} + +void set_positions_from_ipi_bohr(UnitCell& ucell, const std::vector& positions_bohr) +{ + if (positions_bohr.size() != static_cast(3 * ucell.nat)) + { + ModuleBase::WARNING_QUIT("ABACUS socket", "POSDATA atom count does not match STRU."); + } + + int iat = 0; + for (int it = 0; it < ucell.ntype; ++it) + { + Atom* atom = &ucell.atoms[it]; + for (int ia = 0; ia < atom->na; ++ia) + { + const double tau_x = positions_bohr[3 * iat + 0] / ucell.lat0; + const double tau_y = positions_bohr[3 * iat + 1] / ucell.lat0; + const double tau_z = positions_bohr[3 * iat + 2] / ucell.lat0; + + double dx = 0.0; + double dy = 0.0; + double dz = 0.0; + ModuleBase::Mathzone::Cartesian_to_Direct(tau_x, + tau_y, + tau_z, + ucell.latvec.e11, + ucell.latvec.e12, + ucell.latvec.e13, + ucell.latvec.e21, + ucell.latvec.e22, + ucell.latvec.e23, + ucell.latvec.e31, + ucell.latvec.e32, + ucell.latvec.e33, + dx, + dy, + dz); + + atom->dis[ia].x = dx - atom->taud[ia].x; + atom->dis[ia].y = dy - atom->taud[ia].y; + atom->dis[ia].z = dz - atom->taud[ia].z; + atom->taud[ia].x = dx; + atom->taud[ia].y = dy; + atom->taud[ia].z = dz; + atom->tau[ia].x = tau_x; + atom->tau[ia].y = tau_y; + atom->tau[ia].z = tau_z; + ++iat; + } + } + unitcell::periodic_boundary_adjustment(ucell.atoms, ucell.latvec, ucell.ntype); + ucell.ionic_position_updated = true; + ucell.cell_parameter_updated = false; +} + +std::vector flatten_forces_hartree_per_bohr(const ModuleBase::matrix& force, const int nat) +{ + if (nat < 0 || force.nr != nat || force.nc != 3) + { + throw std::runtime_error("force matrix must have nat rows and three columns"); + } + std::vector out(static_cast(force.nr * force.nc)); + for (int iat = 0; iat < force.nr; ++iat) + { + for (int idir = 0; idir < force.nc; ++idir) + { + const double value = force(iat, idir); + if (!std::isfinite(value)) + { + throw std::runtime_error("force entries must be finite"); + } + out[static_cast(3 * iat + idir)] = value * RY_TO_HARTREE; + } + } + return out; +} + +SocketFrame::Matrix9 matrix9_from_stress(const ModuleBase::matrix& stress) +{ + if (stress.nr != 3 || stress.nc != 3) + { + throw std::runtime_error("stress matrix must have three rows and three columns"); + } + SocketFrame::Matrix9 values; + for (int row = 0; row < 3; ++row) + { + for (int column = 0; column < 3; ++column) + { + values[3 * row + column] = stress(row, column); + } + } + return values; +} + +std::vector vector_from_matrix9(const SocketFrame::Matrix9& values) +{ + return std::vector(values.begin(), values.end()); +} +} // namespace + +void Socket_Driver::socket_driver(ModuleESolver::ESolver* p_esolver, + UnitCell& ucell, + const Input_para& inp, + std::ofstream& ofs_running) +{ + ModuleBase::TITLE("Socket_Driver", "socket_driver"); + ModuleBase::timer::start("Socket_Driver", "socket_driver"); + + if (p_esolver == nullptr) + { + ModuleBase::WARNING_QUIT("ABACUS socket", "socket driver requires a valid ESolver."); + } + IpiSocket socket; + + try + { + int io_failed = 0; + std::string io_message; + if (is_root()) + { + try + { + const std::string address = socket_address(); + ofs_running << " ABACUS socket driver connecting to i-PI endpoint " << address << std::endl; + socket.connect(address); + } + catch (const std::exception& exc) + { + io_failed = 1; + io_message = exc.what(); + } + } + quit_if_root_io_failed(io_failed, io_message); + + DriverState state = DriverState::NeedInit; + int istep = 0; + const int nat_return = ucell.nat; + ComputedFrame published; + + const std::vector reference_cell = ipi_cell_bohr_from_unitcell(ucell); + bool checked_initial_positions = false; + + while (true) + { + std::string header; + io_failed = 0; + io_message.clear(); + if (is_root()) + { + try + { + header = socket.read_header(); + } + catch (const IpiSocketClosed&) + { + if (state == DriverState::HasData) + { + io_failed = 1; + io_message = "i-PI peer closed while a computed frame was pending"; + } + else + { + header.clear(); + } + } + catch (const std::exception& exc) + { + io_failed = 1; + io_message = exc.what(); + } + } + quit_if_root_io_failed(io_failed, io_message); + header = bcast_header(header); + + if (header.empty()) + { + if (is_root()) + { + ofs_running << " ABACUS socket driver exiting after peer closed connection" << std::endl; + } + break; + } + else if (header == "STATUS") + { + io_failed = 0; + io_message.clear(); + if (is_root()) + { + try + { + if (state == DriverState::HasData) + { + socket.write_header("HAVEDATA"); + } + else if (state == DriverState::Ready) + { + socket.write_header("READY"); + } + else + { + socket.write_header("NEEDINIT"); + } + } + catch (const std::exception& exc) + { + io_failed = 1; + io_message = exc.what(); + } + } + quit_if_root_io_failed(io_failed, io_message); + } + else if (header == "INIT") + { + std::int32_t rid = 0; + std::int32_t nbytes = 0; + std::string params; + io_failed = 0; + io_message.clear(); + if (is_root()) + { + if (state != DriverState::NeedInit) + { + io_failed = 1; + io_message = "INIT requires NEEDINIT state"; + } + else + { + try + { + rid = socket.read_int32(); + nbytes = socket.read_int32(); + if (nbytes < 0) + { + io_failed = 1; + io_message = "negative INIT payload length from i-PI socket"; + } + else if (nbytes > MAX_INIT_BYTES) + { + io_failed = 1; + io_message = "INIT payload exceeds the 1 MiB socket limit"; + } + else if (nbytes > 0) + { + params = socket.read_string(static_cast(nbytes)); + } + } + catch (const std::exception& exc) + { + io_failed = 1; + io_message = exc.what(); + } + } + } + quit_if_root_io_failed(io_failed, io_message); + bcast_socket_int32(rid); + bcast_socket_int32(nbytes); + if (nbytes > 0 && is_root()) + { + ofs_running << " ABACUS socket INIT params bytes " << nbytes << std::endl; + } + state = DriverState::Ready; + if (is_root()) + { + ofs_running << " ABACUS socket INIT replica " << rid << std::endl; + } + } + else if (header == "POSDATA") + { + SocketFrame::Matrix9 cell = {{0.0}}; + SocketFrame::Matrix9 inv_cell = {{0.0}}; + std::int32_t nat_socket = 0; + std::vector positions; + io_failed = 0; + io_message.clear(); + if (is_root()) + { + if (state != DriverState::Ready) + { + io_failed = 1; + io_message = "POSDATA requires READY state"; + } + else + { + try + { + const std::vector cell_values = socket.read_doubles(9); + const std::vector inverse_values = socket.read_doubles(9); + std::copy(cell_values.begin(), cell_values.end(), cell.begin()); + std::copy(inverse_values.begin(), inverse_values.end(), inv_cell.begin()); + nat_socket = socket.read_int32(); + SocketFrame::CellValidation validation + = SocketFrame::validate_ipi_cell(cell, + inv_cell, + MAX_CELL_CONDITION, + INVERSE_ABSOLUTE_TOLERANCE, + INVERSE_RELATIVE_TOLERANCE); + if (!validation.ok) + { + io_failed = 1; + io_message = "invalid POSDATA cell: " + validation.message; + } + std::size_t coordinate_count = 0; + if (io_failed == 0 + && !SocketFrame::checked_position_count(nat_socket, + ucell.nat, + coordinate_count, + io_message)) + { + io_failed = 1; + } + if (io_failed == 0) + { + positions = socket.read_doubles(coordinate_count); + if (!SocketFrame::validate_positions(positions, + coordinate_count, + io_message)) + { + io_failed = 1; + } + } + } + catch (const std::exception& exc) + { + io_failed = 1; + io_message = exc.what(); + } + } + } + quit_if_root_io_failed(io_failed, io_message); + bcast_socket_int32(nat_socket); + std::vector cell_values(cell.begin(), cell.end()); + std::vector inverse_values(inv_cell.begin(), inv_cell.end()); + bcast_double_vector(cell_values); + bcast_double_vector(inverse_values); + if (!is_root()) + { + cell = {{0.0}}; + inv_cell = {{0.0}}; + std::copy(cell_values.begin(), cell_values.end(), cell.begin()); + std::copy(inverse_values.begin(), inverse_values.end(), inv_cell.begin()); + if (nat_socket >= 0) + { + positions.assign(static_cast(3 * nat_socket), 0.0); + } + } + bcast_double_vector(positions); + + const double max_cell_delta_bohr = max_abs_delta(std::vector(cell.begin(), cell.end()), reference_cell); + if (max_cell_delta_bohr > unchanged_cell_tolerance(cell)) + { + ModuleBase::WARNING_QUIT("ABACUS socket", "variable-cell socket updates are not supported yet."); + } + if (!checked_initial_positions) + { + checked_initial_positions = true; + if (max_wrapped_direct_delta_from_unitcell(ucell, positions) > 1.0e-5 && is_root()) + { + ModuleBase::WARNING( + "ABACUS socket", + "first POSDATA positions are not PBC-equivalent to STRU atom order; " + "i-PI POSDATA carries no species, so the client atoms should use the same atom order as STRU."); + } + } + + try + { + set_positions_from_ipi_bohr(ucell, positions); + } + catch (const std::exception& exc) + { + fail_during_collective_stage("set_positions", exc.what()); + } + catch (...) + { + fail_during_collective_stage("set_positions", + "unknown socket position update failure"); + } + try + { + p_esolver->runner(ucell, istep); + } + catch (const std::exception& exc) + { + fail_during_collective_stage("runner", exc.what()); + } + catch (...) + { + fail_during_collective_stage("runner", + "unknown socket runner failure"); + } + ComputedFrame computed; + computed.scf_converged = all_ranks_converged(p_esolver->conv_esolver); + if (!computed.scf_converged && is_root()) + { + ModuleBase::WARNING( + "ABACUS socket", + "SCF did not converge; returning the available frame and marking it in i-PI extras."); + } + double energy_ry = 0.0; + try + { + energy_ry = p_esolver->cal_energy(); + } + catch (const std::exception& exc) + { + fail_during_collective_stage("cal_energy", exc.what()); + } + catch (...) + { + fail_during_collective_stage("cal_energy", + "unknown socket energy failure"); + } + int local_failed = std::isfinite(energy_ry) ? 0 : 1; + throw_if_any_rank_failed(local_failed, + local_failed == 0 ? "" : "socket energy is not finite"); + if (!std::isfinite(energy_ry)) + { + ModuleBase::WARNING_QUIT("ABACUS socket", "socket energy is not finite."); + } + computed.energy_hartree = energy_ry * RY_TO_HARTREE; + if (is_root()) + { + ofs_running << " ABACUS socket return energy " + << energy_ry << " Ry, " + << energy_ry * ModuleBase::Ry_to_eV << " eV, " + << computed.energy_hartree << " Ha" << std::endl; + } + ModuleBase::matrix force; + if (inp.cal_force) + { + try + { + p_esolver->cal_force(ucell, force); + } + catch (const std::exception& exc) + { + fail_during_collective_stage("cal_force", exc.what()); + } + catch (...) + { + fail_during_collective_stage("cal_force", + "unknown socket force failure"); + } + local_failed = 0; + std::string local_message; + try + { + computed.forces_hartree_per_bohr = flatten_forces_hartree_per_bohr(force, ucell.nat); + } + catch (const std::exception& exc) + { + local_failed = 1; + local_message = exc.what(); + } + catch (...) + { + local_failed = 1; + local_message = "unknown socket force validation failure"; + } + throw_if_any_rank_failed(local_failed, local_message); + computed.forces_present = true; + } + if (inp.cal_stress) + { + ModuleBase::matrix stress; + try + { + p_esolver->cal_stress(ucell, stress); + } + catch (const std::exception& exc) + { + fail_during_collective_stage("cal_stress", exc.what()); + } + catch (...) + { + fail_during_collective_stage("cal_stress", + "unknown socket stress failure"); + } + local_failed = 0; + std::string local_message; + try + { + const SocketFrame::VirialConversion virial + = SocketFrame::make_ipi_virial(matrix9_from_stress(stress), + ucell.omega, + STRESS_ABSOLUTE_TOLERANCE, + STRESS_RELATIVE_TOLERANCE); + if (!virial.ok) + { + throw std::runtime_error(virial.message); + } + computed.virial_wire_hartree = virial.wire_virial_hartree; + } + catch (const std::exception& exc) + { + local_failed = 1; + local_message = exc.what(); + } + catch (...) + { + local_failed = 1; + local_message = "unknown socket stress validation failure"; + } + throw_if_any_rank_failed(local_failed, local_message); + computed.stress_present = true; + } + computed.valid = true; + published = computed; + ++istep; + state = DriverState::HasData; + } + else if (header == "GETFORCE") + { + io_failed = 0; + io_message.clear(); + if (is_root()) + { + try + { + if (state != DriverState::HasData || !published.valid) + { + throw std::runtime_error("GETFORCE requires HAVEDATA state and a valid frame"); + } + socket.write_header("FORCEREADY"); + socket.write_double(published.energy_hartree); + socket.write_int32(static_cast(nat_return)); + const std::vector forces + = published.forces_present + ? published.forces_hartree_per_bohr + : std::vector(static_cast(3 * nat_return), 0.0); + socket.write_doubles(forces); + socket.write_doubles(vector_from_matrix9(published.virial_wire_hartree)); + const std::string extra = properties_extra(published); + if (extra.size() > static_cast(std::numeric_limits::max())) + { + throw std::overflow_error("i-PI extras payload is larger than int32"); + } + socket.write_int32(static_cast(extra.size())); + socket.write_string(extra); + } + catch (const std::exception& exc) + { + io_failed = 1; + io_message = exc.what(); + } + } + quit_if_root_io_failed(io_failed, io_message); + published = ComputedFrame(); + state = DriverState::Ready; + } + else if (header == "EXIT") + { + if (is_root()) + { + ofs_running << " ABACUS socket driver received i-PI EXIT" << std::endl; + } + break; + } + else + { + if (is_root()) + { + io_failed = 1; + io_message = "unknown i-PI header: " + header; + } + quit_if_root_io_failed(io_failed, io_message); + } + } + } + catch (const std::exception& exc) + { + ModuleBase::WARNING_QUIT("ABACUS socket", exc.what()); + } + + if (is_root()) + { + socket.close(); + } + + ModuleBase::timer::end("Socket_Driver", "socket_driver"); +} diff --git a/source/source_relax/socket_driver.h b/source/source_relax/socket_driver.h new file mode 100644 index 00000000000..86c180fc42e --- /dev/null +++ b/source/source_relax/socket_driver.h @@ -0,0 +1,26 @@ +#ifndef ABACUS_SOURCE_RELAX_SOCKET_DRIVER_H +#define ABACUS_SOURCE_RELAX_SOCKET_DRIVER_H + +#include + +class UnitCell; +struct Input_para; + +namespace ModuleESolver +{ +class ESolver; +} + +class Socket_Driver +{ + public: + Socket_Driver() = default; + ~Socket_Driver() = default; + + void socket_driver(ModuleESolver::ESolver* p_esolver, + UnitCell& ucell, + const Input_para& inp, + std::ofstream& ofs_running); +}; + +#endif diff --git a/source/source_relax/socket_frame.cpp b/source/source_relax/socket_frame.cpp new file mode 100644 index 00000000000..9f125310761 --- /dev/null +++ b/source/source_relax/socket_frame.cpp @@ -0,0 +1,426 @@ +#include "socket_frame.h" + +#include +#include +#include + +namespace +{ +const int MATRIX_DIMENSION = 3; +const int MAX_JACOBI_SWEEPS = 32; + +bool is_finite_matrix(const SocketFrame::Matrix9& values) +{ + for (std::size_t index = 0; index < values.size(); ++index) + { + if (!std::isfinite(values[index])) + { + return false; + } + } + return true; +} + +double column_norm_squared(const SocketFrame::Matrix9& values, int column) +{ + double norm_squared = 0.0; + for (int row = 0; row < MATRIX_DIMENSION; ++row) + { + const double value = values[row * MATRIX_DIMENSION + column]; + norm_squared += value * value; + } + return norm_squared; +} + +double column_dot(const SocketFrame::Matrix9& values, int first, int second) +{ + double dot = 0.0; + for (int row = 0; row < MATRIX_DIMENSION; ++row) + { + dot += values[row * MATRIX_DIMENSION + first] * values[row * MATRIX_DIMENSION + second]; + } + return dot; +} + +bool columns_are_orthogonal(const SocketFrame::Matrix9& values) +{ + const double multiplier = 32.0 * std::numeric_limits::epsilon(); + const int pairs[3][2] = {{0, 1}, {0, 2}, {1, 2}}; + for (int pair = 0; pair < 3; ++pair) + { + const int first = pairs[pair][0]; + const int second = pairs[pair][1]; + const double first_norm = column_norm_squared(values, first); + const double second_norm = column_norm_squared(values, second); + const double tolerance = multiplier * std::sqrt(first_norm * second_norm); + if (std::fabs(column_dot(values, first, second)) > tolerance) + { + return false; + } + } + return true; +} + +void rotate_columns(SocketFrame::Matrix9& values, int first, int second, double cosine, double sine) +{ + for (int row = 0; row < MATRIX_DIMENSION; ++row) + { + const int first_index = row * MATRIX_DIMENSION + first; + const int second_index = row * MATRIX_DIMENSION + second; + const double first_value = values[first_index]; + const double second_value = values[second_index]; + values[first_index] = cosine * first_value - sine * second_value; + values[second_index] = sine * first_value + cosine * second_value; + } +} + +bool one_sided_jacobi(SocketFrame::Matrix9& columns, SocketFrame::Matrix9& right_vectors) +{ + right_vectors = {{1.0, 0.0, 0.0, + 0.0, 1.0, 0.0, + 0.0, 0.0, 1.0}}; + const double multiplier = 32.0 * std::numeric_limits::epsilon(); + const int pairs[3][2] = {{0, 1}, {0, 2}, {1, 2}}; + + for (int sweep = 0; sweep < MAX_JACOBI_SWEEPS; ++sweep) + { + for (int pair = 0; pair < 3; ++pair) + { + const int first = pairs[pair][0]; + const int second = pairs[pair][1]; + const double first_norm = column_norm_squared(columns, first); + const double second_norm = column_norm_squared(columns, second); + const double dot = column_dot(columns, first, second); + const double tolerance = multiplier * std::sqrt(first_norm * second_norm); + if (std::fabs(dot) <= tolerance) + { + continue; + } + + const double tau = (second_norm - first_norm) / (2.0 * dot); + const double tangent + = std::copysign(1.0 / (std::fabs(tau) + std::hypot(1.0, tau)), tau); + const double cosine = 1.0 / std::sqrt(1.0 + tangent * tangent); + const double sine = tangent * cosine; + rotate_columns(columns, first, second, cosine, sine); + rotate_columns(right_vectors, first, second, cosine, sine); + } + + if (columns_are_orthogonal(columns)) + { + return true; + } + } + return false; +} + +long double scaled_determinant(const SocketFrame::Matrix9& values) +{ + const long double a00 = values[0]; + const long double a01 = values[1]; + const long double a02 = values[2]; + const long double a10 = values[3]; + const long double a11 = values[4]; + const long double a12 = values[5]; + const long double a20 = values[6]; + const long double a21 = values[7]; + const long double a22 = values[8]; + return a00 * (a11 * a22 - a12 * a21) + - a01 * (a10 * a22 - a12 * a20) + + a02 * (a10 * a21 - a11 * a20); +} + +double received_inverse_residual(const SocketFrame::Matrix9& cell, + const SocketFrame::Matrix9& inverse, + bool transpose_inverse) +{ + long double maximum = 0.0L; + for (int row = 0; row < MATRIX_DIMENSION; ++row) + { + for (int column = 0; column < MATRIX_DIMENSION; ++column) + { + long double product = 0.0L; + for (int inner = 0; inner < MATRIX_DIMENSION; ++inner) + { + const int inverse_index = transpose_inverse + ? column * MATRIX_DIMENSION + inner + : inner * MATRIX_DIMENSION + column; + product += static_cast(cell[row * MATRIX_DIMENSION + inner]) + * inverse[inverse_index]; + } + const long double expected = row == column ? 1.0L : 0.0L; + maximum = std::max(maximum, std::fabs(product - expected)); + } + } + return static_cast(maximum); +} +} // namespace + +namespace SocketFrame +{ +Matrix9 transpose_matrix9(const Matrix9& values) +{ + return {{values[0], values[3], values[6], + values[1], values[4], values[7], + values[2], values[5], values[8]}}; +} + +CellValidation validate_ipi_cell(const Matrix9& cell_wire, + const Matrix9& inverse_wire, + double max_condition_number, + double inverse_absolute_tolerance, + double inverse_relative_tolerance) +{ + CellValidation result; + result.ok = false; + result.message.clear(); + result.determinant_bohr3 = 0.0; + result.condition_number_2 = std::numeric_limits::infinity(); + result.inverse_residual = std::numeric_limits::infinity(); + result.computed_inverse_wire_bohr_inv.fill(0.0); + + if (!is_finite_matrix(cell_wire) || !is_finite_matrix(inverse_wire)) + { + result.message = "cell and received inverse entries must be finite"; + return result; + } + if (!std::isfinite(max_condition_number) || max_condition_number <= 0.0 + || !std::isfinite(inverse_absolute_tolerance) || inverse_absolute_tolerance < 0.0 + || !std::isfinite(inverse_relative_tolerance) || inverse_relative_tolerance < 0.0) + { + result.message = "cell validation tolerances must be finite and nonnegative"; + return result; + } + + double scale = 0.0; + for (std::size_t index = 0; index < cell_wire.size(); ++index) + { + scale = std::max(scale, std::fabs(cell_wire[index])); + } + if (scale == 0.0) + { + result.message = "cell determinant must be positive"; + return result; + } + + Matrix9 scaled_cell; + for (std::size_t index = 0; index < cell_wire.size(); ++index) + { + scaled_cell[index] = cell_wire[index] / scale; + } + const long double determinant_scaled = scaled_determinant(scaled_cell); + if (determinant_scaled <= 0.0L) + { + result.message = "cell determinant must be positive"; + return result; + } + const long double scale_long = scale; + const long double determinant + = determinant_scaled * scale_long * scale_long * scale_long; + if (!std::isfinite(determinant) + || determinant > static_cast(std::numeric_limits::max())) + { + result.message = "cell determinant is not representable as a finite double"; + return result; + } + result.determinant_bohr3 = static_cast(determinant); + if (!std::isfinite(result.determinant_bohr3) || result.determinant_bohr3 <= 0.0) + { + result.message = "cell determinant is not representable as a positive finite double"; + return result; + } + + Matrix9 orthogonal_columns = scaled_cell; + Matrix9 right_vectors; + if (!one_sided_jacobi(orthogonal_columns, right_vectors)) + { + result.message = "cell singular-value iteration did not converge"; + return result; + } + + double singular_values[MATRIX_DIMENSION]; + double largest_singular = 0.0; + double smallest_singular = std::numeric_limits::infinity(); + for (int column = 0; column < MATRIX_DIMENSION; ++column) + { + singular_values[column] = std::sqrt(column_norm_squared(orthogonal_columns, column)); + largest_singular = std::max(largest_singular, singular_values[column]); + smallest_singular = std::min(smallest_singular, singular_values[column]); + } + if (smallest_singular == 0.0 || !std::isfinite(smallest_singular)) + { + result.message = "cell is singular"; + return result; + } + result.condition_number_2 = largest_singular / smallest_singular; + if (!std::isfinite(result.condition_number_2) + || result.condition_number_2 >= max_condition_number) + { + result.message = "cell condition number is not below the configured maximum"; + return result; + } + + for (int row = 0; row < MATRIX_DIMENSION; ++row) + { + for (int column = 0; column < MATRIX_DIMENSION; ++column) + { + long double inverse_value = 0.0L; + for (int singular = 0; singular < MATRIX_DIMENSION; ++singular) + { + const long double sigma = singular_values[singular]; + inverse_value + += static_cast(right_vectors[row * MATRIX_DIMENSION + singular]) + * orthogonal_columns[column * MATRIX_DIMENSION + singular] + / (static_cast(scale) * sigma * sigma); + } + result.computed_inverse_wire_bohr_inv[row * MATRIX_DIMENSION + column] + = static_cast(inverse_value); + } + } + + const double direct_inverse_residual + = received_inverse_residual(cell_wire, inverse_wire, false); + const double transposed_inverse_residual + = received_inverse_residual(cell_wire, inverse_wire, true); + result.inverse_residual = std::min(direct_inverse_residual, transposed_inverse_residual); + const double residual_limit + = inverse_absolute_tolerance + + inverse_relative_tolerance * result.condition_number_2 + * std::numeric_limits::epsilon(); + if (!std::isfinite(result.inverse_residual) || result.inverse_residual > residual_limit) + { + result.message = "received cell inverse is inconsistent with the cell"; + return result; + } + + result.ok = true; + return result; +} + +bool validate_positions(const std::vector& positions_bohr, + std::size_t coordinate_count, + std::string& message) +{ + if (positions_bohr.size() != coordinate_count) + { + message = "position coordinate count does not match the validated atom count"; + return false; + } + for (std::size_t index = 0; index < positions_bohr.size(); ++index) + { + if (!std::isfinite(positions_bohr[index])) + { + message = "position coordinates must be finite"; + return false; + } + } + message.clear(); + return true; +} + +bool checked_position_count(std::int32_t nat_socket, + int nat_expected, + std::size_t& coordinate_count, + std::string& message) +{ + if (nat_socket != nat_expected) + { + message = "socket atom count does not match the expected atom count"; + return false; + } + if (nat_socket < 0) + { + message = "socket atom count must not be negative"; + return false; + } + const std::size_t atom_count = static_cast(nat_socket); + if (atom_count > std::numeric_limits::max() / 3) + { + message = "socket position coordinate count is not representable"; + return false; + } + coordinate_count = 3 * atom_count; + message.clear(); + return true; +} + +VirialConversion make_ipi_virial(const Matrix9& stress_ry_per_bohr3, + double volume_bohr3, + double antisymmetric_absolute_tolerance, + double antisymmetric_relative_tolerance) +{ + VirialConversion result; + result.ok = false; + result.message.clear(); + result.wire_virial_hartree.fill(0.0); + result.max_antisymmetric_component = 0.0; + + if (!is_finite_matrix(stress_ry_per_bohr3)) + { + result.message = "stress entries must be finite"; + return result; + } + if (!std::isfinite(volume_bohr3) || volume_bohr3 <= 0.0) + { + result.message = "cell volume must be finite and positive"; + return result; + } + if (!std::isfinite(antisymmetric_absolute_tolerance) + || antisymmetric_absolute_tolerance < 0.0 + || !std::isfinite(antisymmetric_relative_tolerance) + || antisymmetric_relative_tolerance < 0.0) + { + result.message = "stress symmetry tolerances must be finite and nonnegative"; + return result; + } + + double maximum_stress = 0.0; + for (std::size_t index = 0; index < stress_ry_per_bohr3.size(); ++index) + { + maximum_stress = std::max(maximum_stress, std::fabs(stress_ry_per_bohr3[index])); + } + for (int row = 0; row < MATRIX_DIMENSION; ++row) + { + for (int column = row + 1; column < MATRIX_DIMENSION; ++column) + { + const double difference + = std::fabs(stress_ry_per_bohr3[row * MATRIX_DIMENSION + column] + - stress_ry_per_bohr3[column * MATRIX_DIMENSION + row]); + result.max_antisymmetric_component + = std::max(result.max_antisymmetric_component, difference); + } + } + const double symmetry_limit + = antisymmetric_absolute_tolerance + antisymmetric_relative_tolerance * maximum_stress; + if (!std::isfinite(result.max_antisymmetric_component) + || result.max_antisymmetric_component > symmetry_limit) + { + result.message = "stress tensor is not symmetric within tolerance"; + return result; + } + + Matrix9 virial; + for (int row = 0; row < MATRIX_DIMENSION; ++row) + { + for (int column = 0; column < MATRIX_DIMENSION; ++column) + { + const long double symmetric_stress + = 0.5L + * (static_cast(stress_ry_per_bohr3[row * MATRIX_DIMENSION + column]) + + stress_ry_per_bohr3[column * MATRIX_DIMENSION + row]); + const long double converted = 0.5L * volume_bohr3 * symmetric_stress; + if (!std::isfinite(converted) + || std::fabs(converted) + > static_cast(std::numeric_limits::max())) + { + result.message = "converted virial is not representable as finite doubles"; + return result; + } + virial[row * MATRIX_DIMENSION + column] = static_cast(converted); + } + } + result.wire_virial_hartree = transpose_matrix9(virial); + result.ok = true; + return result; +} +} // namespace SocketFrame diff --git a/source/source_relax/socket_frame.h b/source/source_relax/socket_frame.h new file mode 100644 index 00000000000..759a4ee13a2 --- /dev/null +++ b/source/source_relax/socket_frame.h @@ -0,0 +1,51 @@ +#ifndef SOURCE_RELAX_SOCKET_FRAME_H +#define SOURCE_RELAX_SOCKET_FRAME_H + +#include +#include +#include +#include +#include + +namespace SocketFrame +{ +using Matrix9 = std::array; + +struct CellValidation +{ + bool ok; + std::string message; + double determinant_bohr3; + double condition_number_2; + double inverse_residual; + Matrix9 computed_inverse_wire_bohr_inv; +}; + +struct VirialConversion +{ + bool ok; + std::string message; + Matrix9 wire_virial_hartree; + double max_antisymmetric_component; +}; + +Matrix9 transpose_matrix9(const Matrix9& values); +CellValidation validate_ipi_cell(const Matrix9& cell_wire, + const Matrix9& inverse_wire, + double max_condition_number, + double inverse_absolute_tolerance, + double inverse_relative_tolerance); +bool validate_positions(const std::vector& positions_bohr, + std::size_t coordinate_count, + std::string& message); +bool checked_position_count(std::int32_t nat_socket, + int nat_expected, + std::size_t& coordinate_count, + std::string& message); +VirialConversion make_ipi_virial(const Matrix9& stress_ry_per_bohr3, + double volume_bohr3, + double antisymmetric_absolute_tolerance, + double antisymmetric_relative_tolerance); +} // namespace SocketFrame + +#endif diff --git a/source/source_relax/socket_ipi.cpp b/source/source_relax/socket_ipi.cpp new file mode 100644 index 00000000000..d0ba19869aa --- /dev/null +++ b/source/source_relax/socket_ipi.cpp @@ -0,0 +1,295 @@ +#include "source_relax/socket_ipi.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +static_assert(sizeof(std::int32_t) == 4, "i-PI requires a 4-byte integer"); +static_assert(sizeof(double) == 8, "i-PI requires an 8-byte float"); +static_assert(std::numeric_limits::is_iec559, + "i-PI requires IEEE-754 double precision"); + +namespace +{ +constexpr std::size_t IPI_HEADER_LEN = 12; + +std::string errno_message(const std::string& prefix) +{ + return prefix + ": " + std::strerror(errno); +} + +std::string trim_header(const char* data) +{ + std::string value(data, IPI_HEADER_LEN); + while (!value.empty() && value.back() == ' ') + { + value.pop_back(); + } + return value; +} + +std::string padded_header(const std::string& header) +{ + if (header.size() > IPI_HEADER_LEN) + { + throw std::runtime_error("i-PI header is longer than 12 bytes: " + header); + } + std::string out = header; + out.resize(IPI_HEADER_LEN, ' '); + return out; +} + +std::size_t checked_double_bytes(std::size_t n) +{ + if (n > SIZE_MAX / sizeof(double)) + { + throw std::overflow_error("i-PI double payload byte count overflows for " + std::to_string(n) + " elements"); + } + return n * sizeof(double); +} +} // namespace + +IpiSocketClosed::IpiSocketClosed(const std::string& message) : std::runtime_error(message) +{ +} + +IpiSocket::~IpiSocket() +{ + this->close(); +} + +void IpiSocket::connect(const std::string& address) +{ + this->close(); + const std::size_t colon = address.rfind(':'); + if (colon == std::string::npos) + { + throw std::runtime_error("i-PI address must be host:port or path:UNIX, got " + address); + } + const std::string host = address.substr(0, colon); + const std::string service = address.substr(colon + 1); + + if (service == "UNIX") + { + fd_ = ::socket(AF_UNIX, SOCK_STREAM, 0); + if (fd_ < 0) + { + throw std::runtime_error(errno_message("failed to create UNIX socket")); + } + sockaddr_un addr; + std::memset(&addr, 0, sizeof(addr)); + addr.sun_family = AF_UNIX; + if (host.size() >= sizeof(addr.sun_path)) + { + this->close(); + throw std::runtime_error("UNIX socket path too long: " + host); + } + std::strncpy(addr.sun_path, host.c_str(), sizeof(addr.sun_path) - 1); + if (::connect(fd_, reinterpret_cast(&addr), sizeof(addr)) != 0) + { + const std::string msg = errno_message("failed to connect UNIX i-PI socket " + host); + this->close(); + throw std::runtime_error(msg); + } + return; + } + + addrinfo hints; + std::memset(&hints, 0, sizeof(hints)); + hints.ai_family = AF_UNSPEC; + hints.ai_socktype = SOCK_STREAM; + + addrinfo* result = nullptr; + const int gai = ::getaddrinfo(host.c_str(), service.c_str(), &hints, &result); + if (gai != 0) + { + throw std::runtime_error("failed to resolve i-PI socket " + address + ": " + ::gai_strerror(gai)); + } + + std::string last_error; + for (addrinfo* rp = result; rp != nullptr; rp = rp->ai_next) + { + fd_ = ::socket(rp->ai_family, rp->ai_socktype, rp->ai_protocol); + if (fd_ < 0) + { + last_error = errno_message("failed to create INET socket"); + continue; + } + if (::connect(fd_, rp->ai_addr, rp->ai_addrlen) == 0) + { + ::freeaddrinfo(result); + return; + } + last_error = errno_message("failed to connect INET i-PI socket " + address); + this->close(); + } + ::freeaddrinfo(result); + throw std::runtime_error(last_error.empty() ? "failed to connect i-PI socket " + address : last_error); +} + +void IpiSocket::close() +{ + if (fd_ >= 0) + { + ::close(fd_); + fd_ = -1; + } +} + +std::string IpiSocket::read_header() +{ + char header[IPI_HEADER_LEN]; + std::size_t done = 0; + while (done < sizeof(header)) + { + const ssize_t nread = ::recv(fd_, header + done, sizeof(header) - done, 0); + if (nread == 0) + { + if (done == 0) + { + throw IpiSocketClosed("i-PI socket closed before next header"); + } + throw std::runtime_error("i-PI socket closed while reading header"); + } + if (nread < 0) + { + if (errno == EINTR) + { + continue; + } + if (errno == ECONNRESET && done == 0) + { + throw IpiSocketClosed("i-PI socket peer reset before next header"); + } + throw std::runtime_error(errno_message("i-PI socket header read failed")); + } + done += static_cast(nread); + } + return trim_header(header); +} + +void IpiSocket::write_header(const std::string& header) +{ + const std::string padded = padded_header(header); + this->write_exact(padded.data(), padded.size()); +} + +std::int32_t IpiSocket::read_int32() +{ + std::int32_t value = 0; + this->read_exact(&value, sizeof(value)); + return value; +} + +void IpiSocket::write_int32(std::int32_t value) +{ + this->write_exact(&value, sizeof(value)); +} + +double IpiSocket::read_double() +{ + double value = 0.0; + this->read_exact(&value, sizeof(value)); + return value; +} + +void IpiSocket::write_double(double value) +{ + this->write_exact(&value, sizeof(value)); +} + +std::vector IpiSocket::read_doubles(std::size_t n) +{ + const std::size_t nbytes = checked_double_bytes(n); + std::vector values(n); + if (!values.empty()) + { + this->read_exact(values.data(), nbytes); + } + return values; +} + +void IpiSocket::write_doubles(const std::vector& values) +{ + const std::size_t nbytes = checked_double_bytes(values.size()); + if (!values.empty()) + { + this->write_exact(values.data(), nbytes); + } +} + +std::string IpiSocket::read_string(std::size_t nbytes) +{ + std::string value(nbytes, '\0'); + if (nbytes > 0) + { + this->read_exact(&value[0], nbytes); + } + return value; +} + +void IpiSocket::write_string(const std::string& value) +{ + if (!value.empty()) + { + this->write_exact(value.data(), value.size()); + } +} + +void IpiSocket::read_exact(void* data, std::size_t nbytes) +{ + char* cursor = static_cast(data); + std::size_t done = 0; + while (done < nbytes) + { + const ssize_t nread = ::recv(fd_, cursor + done, nbytes - done, 0); + if (nread == 0) + { + throw IpiSocketClosed("i-PI socket closed while reading"); + } + if (nread < 0) + { + if (errno == EINTR) + { + continue; + } + throw std::runtime_error(errno_message("i-PI socket read failed")); + } + done += static_cast(nread); + } +} + +void IpiSocket::write_exact(const void* data, std::size_t nbytes) +{ + const char* cursor = static_cast(data); + std::size_t done = 0; + while (done < nbytes) + { +#ifdef MSG_NOSIGNAL + const int flags = MSG_NOSIGNAL; +#else + const int flags = 0; +#endif + const ssize_t nwritten = ::send(fd_, cursor + done, nbytes - done, flags); + if (nwritten == 0) + { + throw std::runtime_error("i-PI socket closed while writing"); + } + if (nwritten < 0) + { + if (errno == EINTR) + { + continue; + } + throw std::runtime_error(errno_message("i-PI socket write failed")); + } + done += static_cast(nwritten); + } +} diff --git a/source/source_relax/socket_ipi.h b/source/source_relax/socket_ipi.h new file mode 100644 index 00000000000..ea8183fd466 --- /dev/null +++ b/source/source_relax/socket_ipi.h @@ -0,0 +1,49 @@ +#ifndef ABACUS_SOCKET_IPI_H +#define ABACUS_SOCKET_IPI_H + +#include +#include +#include +#include +#include + +class IpiSocketClosed : public std::runtime_error +{ + public: + explicit IpiSocketClosed(const std::string& message); +}; + +class IpiSocket +{ + public: + IpiSocket() = default; + ~IpiSocket(); + + IpiSocket(const IpiSocket&) = delete; + IpiSocket& operator=(const IpiSocket&) = delete; + + void connect(const std::string& address); + void close(); + + std::string read_header(); + void write_header(const std::string& header); + + std::int32_t read_int32(); + void write_int32(std::int32_t value); + + double read_double(); + void write_double(double value); + + std::vector read_doubles(std::size_t n); + void write_doubles(const std::vector& values); + std::string read_string(std::size_t nbytes); + void write_string(const std::string& value); + + private: + int fd_ = -1; + + void read_exact(void* data, std::size_t nbytes); + void write_exact(const void* data, std::size_t nbytes); +}; + +#endif diff --git a/source/source_relax/test/CMakeLists.txt b/source/source_relax/test/CMakeLists.txt index 4493ced0a06..a5fdadb78c1 100644 --- a/source/source_relax/test/CMakeLists.txt +++ b/source/source_relax/test/CMakeLists.txt @@ -6,6 +6,29 @@ abacus_disable_feature_definitions(__ROCM) install(DIRECTORY support DESTINATION ${CMAKE_CURRENT_BINARY_DIR}) + +AddTest( + TARGET MODULE_RELAX_socket_ipi_test + SOURCES socket_ipi_test.cpp ../socket_ipi.cpp +) + +AddTest( + TARGET MODULE_RELAX_socket_frame_test + SOURCES socket_frame_test.cpp ../socket_frame.cpp +) + +AddTest( + TARGET MODULE_RELAX_socket_driver_test + LIBS base device + SOURCES socket_driver_test.cpp + ../socket_driver.cpp + ../socket_frame.cpp + ../socket_ipi.cpp + ../../source_cell/update_cell.cpp + ../../source_cell/bcast_cell.cpp +) +set_tests_properties(MODULE_RELAX_socket_driver_test PROPERTIES TIMEOUT 15) + AddTest( TARGET MODULE_RELAX_relax_new_line_search LIBS parameter diff --git a/source/source_relax/test/socket_driver_test.cpp b/source/source_relax/test/socket_driver_test.cpp new file mode 100644 index 00000000000..fe361e55d65 --- /dev/null +++ b/source/source_relax/test/socket_driver_test.cpp @@ -0,0 +1,615 @@ +#include "source_relax/socket_driver.h" + +#include "gmock/gmock.h" +#include "gtest/gtest.h" +#include "source_cell/unitcell.h" +#include "source_esolver/esolver.h" +#include "source_io/module_parameter/input_parameter.h" +#include "for_test.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace +{ +constexpr std::size_t IPI_HEADER_LEN = 12; + +std::string errno_message(const std::string& prefix) +{ + return prefix + ": " + std::strerror(errno); +} + +void send_all(const int fd, const void* data, const std::size_t nbytes) +{ + const char* cursor = static_cast(data); + std::size_t done = 0; + while (done < nbytes) + { +#ifdef MSG_NOSIGNAL + const int flags = MSG_NOSIGNAL; +#else + const int flags = 0; +#endif + const ssize_t sent = ::send(fd, cursor + done, nbytes - done, flags); + if (sent < 0) + { + if (errno == EINTR) + { + continue; + } + throw std::runtime_error(errno_message("send failed")); + } + if (sent == 0) + { + throw std::runtime_error("send returned zero"); + } + done += static_cast(sent); + } +} + +template +void send_value(const int fd, const T& value) +{ + send_all(fd, &value, sizeof(value)); +} + +void send_header(const int fd, const std::string& header) +{ + std::string padded = header; + padded.resize(IPI_HEADER_LEN, ' '); + send_all(fd, padded.data(), padded.size()); +} + +bool try_send_status(const int fd) +{ + try + { + send_header(fd, "STATUS"); + return true; + } + catch (const std::runtime_error&) + { + if (errno == EPIPE || errno == ECONNRESET) + { + return false; + } + throw; + } +} + +std::string read_header_or_close(const int fd) +{ + char header[IPI_HEADER_LEN]; + std::size_t done = 0; + while (done < sizeof(header)) + { + const ssize_t received = ::recv(fd, header + done, sizeof(header) - done, 0); + if (received == 0 || (received < 0 && errno == ECONNRESET)) + { + if (done == 0) + { + return ""; + } + throw std::runtime_error("socket closed during response header"); + } + if (received < 0) + { + if (errno == EINTR) + { + continue; + } + throw std::runtime_error(errno_message("receive failed")); + } + done += static_cast(received); + } + + std::string value(header, sizeof(header)); + while (!value.empty() && value.back() == ' ') + { + value.pop_back(); + } + return value; +} + +class UnixSocketServer +{ + public: + UnixSocketServer() + { + char dir_template[] = "/tmp/abacus_socket_driver_test_XXXXXX"; + char* made_dir = ::mkdtemp(dir_template); + if (made_dir == nullptr) + { + throw std::runtime_error(errno_message("mkdtemp failed")); + } + dir_ = made_dir; + path_ = dir_ + "/ipi.sock"; + + listen_fd_ = ::socket(AF_UNIX, SOCK_STREAM, 0); + if (listen_fd_ < 0) + { + throw std::runtime_error(errno_message("socket failed")); + } + + sockaddr_un address; + std::memset(&address, 0, sizeof(address)); + address.sun_family = AF_UNIX; + std::strncpy(address.sun_path, path_.c_str(), sizeof(address.sun_path) - 1); + if (::bind(listen_fd_, reinterpret_cast(&address), sizeof(address)) != 0) + { + throw std::runtime_error(errno_message("bind failed")); + } + if (::listen(listen_fd_, 1) != 0) + { + throw std::runtime_error(errno_message("listen failed")); + } + } + + ~UnixSocketServer() + { + if (listen_fd_ >= 0) + { + ::close(listen_fd_); + } + if (!path_.empty()) + { + ::unlink(path_.c_str()); + } + if (!dir_.empty()) + { + ::rmdir(dir_.c_str()); + } + } + + UnixSocketServer(const UnixSocketServer&) = delete; + UnixSocketServer& operator=(const UnixSocketServer&) = delete; + + std::string address() const + { + return path_ + ":UNIX"; + } + + int accept_once() const + { + const int fd = ::accept(listen_fd_, nullptr, nullptr); + if (fd < 0) + { + throw std::runtime_error(errno_message("accept failed")); + } + return fd; + } + + private: + int listen_fd_ = -1; + std::string dir_; + std::string path_; +}; + +class FakeESolver : public ModuleESolver::ESolver +{ + public: + explicit FakeESolver(const bool converged) : converged_(converged) + { + } + + void before_all_runners(BaseCell&, const Input_para&) override + { + } + + void runner(BaseCell& cell, const int step) override + { + position_ = dynamic_cast(cell).atoms[0].tau[0].x; + this->conv_esolver = converged_ && step == 0; + } + + void after_all_runners(BaseCell&) override + { + } + + double cal_energy() override + { + return 4.0 + position_; + } + + void cal_force(BaseCell& cell, ModuleBase::matrix& force) override + { + force.create(cell.nat(), 3); + force(0, 0) = 4.0 + 2.0 * position_; + } + + void cal_stress(BaseCell&, ModuleBase::matrix& stress) override + { + stress.create(3, 3); + stress(0, 0) = 2.0 + 3.0 * position_; + stress(1, 1) = 2.0; + stress(2, 2) = 2.0; + } + + private: + double position_ = 0.0; + bool converged_; +}; + +struct DriverResult +{ + int exit_code = -1; + std::string response_header; + std::string diagnostic; +}; + +struct ForceResponse +{ + std::string header; + double energy_hartree = 0.0; + std::int32_t nat = 0; + std::vector forces_hartree_per_bohr; + std::vector virial_wire_hartree; + std::string extra; +}; + +void initialize_one_atom_cell(UnitCell& ucell) +{ + ucell.lat0 = 1.0; + ucell.latvec.Identity(); + ucell.omega = 1.0; + ucell.ntype = 1; + ucell.nat = 1; + ucell.atoms[0].na = 1; + ucell.atoms[0].tau.resize(1); + ucell.atoms[0].taud.resize(1); + ucell.atoms[0].dis.resize(1); +} + +void send_positions(const int fd, const double x) +{ + const double identity[9] = {1.0, 0.0, 0.0, + 0.0, 1.0, 0.0, + 0.0, 0.0, 1.0}; + const std::int32_t nat = 1; + const double position[3] = {x, 0.0, 0.0}; + send_header(fd, "POSDATA"); + send_all(fd, identity, sizeof(identity)); + send_all(fd, identity, sizeof(identity)); + send_value(fd, nat); + send_all(fd, position, sizeof(position)); +} + +void send_fixed_cell_frame(const int fd) +{ + const std::int32_t replica = 0; + const std::int32_t parameter_bytes = 0; + send_header(fd, "INIT"); + send_value(fd, replica); + send_value(fd, parameter_bytes); + + send_positions(fd, 0.0); +} + +void read_all(const int fd, void* data, const std::size_t nbytes) +{ + char* cursor = static_cast(data); + std::size_t done = 0; + while (done < nbytes) + { + const ssize_t received = ::recv(fd, cursor + done, nbytes - done, 0); + if (received <= 0) + { + throw std::runtime_error("socket closed while reading response"); + } + done += static_cast(received); + } +} + +template +T read_value(const int fd) +{ + T value; + read_all(fd, &value, sizeof(value)); + return value; +} + +std::vector read_doubles(const int fd, const std::size_t count) +{ + std::vector values(count); + if (!values.empty()) + { + read_all(fd, values.data(), values.size() * sizeof(double)); + } + return values; +} + +ForceResponse read_force_response(const int fd) +{ + ForceResponse response; + response.header = read_header_or_close(fd); + if (response.header.empty()) + { + return response; + } + response.energy_hartree = read_value(fd); + response.nat = read_value(fd); + response.forces_hartree_per_bohr + = read_doubles(fd, static_cast(3 * response.nat)); + response.virial_wire_hartree = read_doubles(fd, 9); + const std::int32_t extra_bytes = read_value(fd); + if (extra_bytes < 0) + { + throw std::runtime_error("negative extras length"); + } + response.extra.resize(static_cast(extra_bytes)); + if (!response.extra.empty()) + { + read_all(fd, &response.extra[0], response.extra.size()); + } + return response; +} + +std::string read_pipe(const int fd) +{ + std::string output; + char buffer[512]; + while (true) + { + const ssize_t nread = ::read(fd, buffer, sizeof(buffer)); + if (nread == 0) + { + break; + } + if (nread < 0) + { + if (errno == EINTR) + { + continue; + } + throw std::runtime_error(errno_message("pipe read failed")); + } + output.append(buffer, static_cast(nread)); + } + return output; +} + +DriverResult run_driver_frame(const bool converged, + const bool cal_force, + const bool cal_stress, + const std::function& peer_action) +{ + UnixSocketServer server; + int output_pipe[2]; + if (::pipe(output_pipe) != 0) + { + throw std::runtime_error(errno_message("pipe failed")); + } + + const pid_t child = ::fork(); + if (child < 0) + { + ::close(output_pipe[0]); + ::close(output_pipe[1]); + throw std::runtime_error(errno_message("fork failed")); + } + if (child == 0) + { + ::close(output_pipe[0]); + ::dup2(output_pipe[1], STDOUT_FILENO); + ::dup2(output_pipe[1], STDERR_FILENO); + ::close(output_pipe[1]); + ::setenv("ABACUS_SOCKET_ADDRESS", server.address().c_str(), 1); + + UnitCell ucell; + initialize_one_atom_cell(ucell); + Input_para input; + input.cal_force = cal_force; + input.cal_stress = cal_stress; + FakeESolver solver(converged); + std::ofstream running("/dev/null"); + Socket_Driver driver; + driver.socket_driver(&solver, ucell, input, running); + std::cout.flush(); + std::cerr.flush(); + ::_exit(0); + } + + ::close(output_pipe[1]); + DriverResult result; + std::exception_ptr peer_error; + int peer_fd = -1; + try + { + peer_fd = server.accept_once(); + timeval timeout; + timeout.tv_sec = 5; + timeout.tv_usec = 0; + if (::setsockopt(peer_fd, SOL_SOCKET, SO_RCVTIMEO, &timeout, sizeof(timeout)) != 0) + { + throw std::runtime_error(errno_message("setsockopt failed")); + } + send_fixed_cell_frame(peer_fd); + peer_action(peer_fd); + } + catch (...) + { + peer_error = std::current_exception(); + } + if (peer_fd >= 0) + { + ::close(peer_fd); + } + + result.diagnostic = read_pipe(output_pipe[0]); + ::close(output_pipe[0]); + int status = 0; + while (::waitpid(child, &status, 0) < 0) + { + if (errno != EINTR) + { + throw std::runtime_error(errno_message("waitpid failed")); + } + } + if (WIFEXITED(status)) + { + result.exit_code = WEXITSTATUS(status); + } + + if (peer_error) + { + std::rethrow_exception(peer_error); + } + return result; +} +} // namespace + +TEST(SocketDriverTest, NonconvergedFrameIsPublishedWithMetadata) +{ + ForceResponse response; + const DriverResult result = run_driver_frame( + false, true, false, + [&](const int fd) { + send_header(fd, "GETFORCE"); + response = read_force_response(fd); + }); + + EXPECT_EQ("FORCEREADY", response.header); + EXPECT_EQ(0, result.exit_code); + EXPECT_THAT(response.extra, testing::HasSubstr("\"scf_converged\":false")); +} + +TEST(SocketDriverTest, EnergyOnlyFrameMarksForceAndStressAbsent) +{ + ForceResponse response; + const DriverResult result = run_driver_frame( + true, false, false, + [&](const int fd) { + send_header(fd, "GETFORCE"); + response = read_force_response(fd); + }); + + EXPECT_EQ("FORCEREADY", response.header); + EXPECT_EQ(0, result.exit_code); + EXPECT_THAT(response.extra, testing::HasSubstr("\"present\":[\"energy\"]")); + EXPECT_THAT(response.extra, testing::Not(testing::HasSubstr("\"forces\""))); + EXPECT_THAT(response.extra, testing::Not(testing::HasSubstr("\"stress\""))); + EXPECT_THAT(response.forces_hartree_per_bohr, + testing::ElementsAre(0.0, 0.0, 0.0)); + EXPECT_THAT(response.virial_wire_hartree, + testing::ElementsAre(0.0, 0.0, 0.0, + 0.0, 0.0, 0.0, + 0.0, 0.0, 0.0)); +} + +TEST(SocketDriverTest, EnergyAndStressFrameDoesNotAdvertiseForce) +{ + ForceResponse response; + const DriverResult result = run_driver_frame( + true, false, true, + [&](const int fd) { + send_header(fd, "GETFORCE"); + response = read_force_response(fd); + }); + + EXPECT_EQ("FORCEREADY", response.header); + EXPECT_EQ(0, result.exit_code); + EXPECT_THAT(response.extra, testing::HasSubstr("\"present\":[\"energy\",\"stress\"]")); + EXPECT_THAT(response.extra, testing::Not(testing::HasSubstr("\"forces\""))); + EXPECT_NE(0.0, response.virial_wire_hartree[0]); +} + +TEST(SocketDriverTest, EnergyAndForceFrameAdvertisesOnlyForce) +{ + ForceResponse response; + const DriverResult result = run_driver_frame( + true, true, false, + [&](const int fd) { + send_header(fd, "GETFORCE"); + response = read_force_response(fd); + }); + + EXPECT_EQ("FORCEREADY", response.header); + EXPECT_EQ(0, result.exit_code); + EXPECT_THAT(response.extra, testing::HasSubstr("\"present\":[\"energy\",\"forces\"]")); + EXPECT_THAT(response.extra, testing::Not(testing::HasSubstr("\"stress\""))); + EXPECT_THAT(response.forces_hartree_per_bohr, + testing::ElementsAre(2.0, 0.0, 0.0)); + EXPECT_THAT(response.virial_wire_hartree, + testing::ElementsAre(0.0, 0.0, 0.0, + 0.0, 0.0, 0.0, + 0.0, 0.0, 0.0)); +} + +TEST(SocketDriverTest, EnergyForceAndStressFrameAdvertisesBothDerivatives) +{ + ForceResponse response; + const DriverResult result = run_driver_frame( + true, true, true, + [&](const int fd) { + send_header(fd, "GETFORCE"); + response = read_force_response(fd); + }); + + EXPECT_EQ("FORCEREADY", response.header); + EXPECT_EQ(0, result.exit_code); + EXPECT_THAT(response.extra, + testing::HasSubstr("\"present\":[\"energy\",\"forces\",\"stress\"]")); + EXPECT_EQ(3u, response.forces_hartree_per_bohr.size()); + EXPECT_THAT(response.forces_hartree_per_bohr, + testing::ElementsAre(2.0, 0.0, 0.0)); + EXPECT_NE(0.0, response.virial_wire_hartree[0]); +} + +TEST(SocketDriverTest, ConsecutiveFramesKeepGeometryResultsAndConvergenceTogether) +{ + ForceResponse first, second; + const DriverResult result = run_driver_frame(true, true, true, [&](const int fd) { + send_header(fd, "GETFORCE"); + first = read_force_response(fd); + send_positions(fd, 0.25); + send_header(fd, "GETFORCE"); + second = read_force_response(fd); + }); + EXPECT_EQ(0, result.exit_code); + EXPECT_DOUBLE_EQ(2.0, first.energy_hartree); + EXPECT_DOUBLE_EQ(2.125, second.energy_hartree); + EXPECT_DOUBLE_EQ(2.0, first.forces_hartree_per_bohr.at(0)); + EXPECT_DOUBLE_EQ(2.25, second.forces_hartree_per_bohr.at(0)); + EXPECT_DOUBLE_EQ(1.0, first.virial_wire_hartree.at(0)); + EXPECT_DOUBLE_EQ(1.375, second.virial_wire_hartree.at(0)); + EXPECT_THAT(first.extra, testing::HasSubstr("\"scf_converged\":true")); + EXPECT_THAT(second.extra, testing::HasSubstr("\"scf_converged\":false")); +} + +TEST(SocketDriverTest, ConsumedFrameCannotBeReturnedTwice) +{ + const DriverResult result = run_driver_frame(true, true, false, [&](const int fd) { + send_header(fd, "GETFORCE"); + read_force_response(fd); + send_header(fd, "GETFORCE"); + EXPECT_EQ("", read_header_or_close(fd)); + }); + EXPECT_NE(0, result.exit_code); + EXPECT_THAT(result.diagnostic, testing::HasSubstr("GETFORCE requires HAVEDATA")); +} + +TEST(SocketDriverTest, InvalidNextFrameCannotReturnPreviousResults) +{ + const DriverResult result = run_driver_frame(true, true, false, [&](const int fd) { + send_header(fd, "GETFORCE"); + read_force_response(fd); + send_positions(fd, std::numeric_limits::quiet_NaN()); + EXPECT_EQ("", read_header_or_close(fd)); + }); + EXPECT_NE(0, result.exit_code); + EXPECT_THAT(result.diagnostic, testing::HasSubstr("finite")); +} diff --git a/source/source_relax/test/socket_frame_test.cpp b/source/source_relax/test/socket_frame_test.cpp new file mode 100644 index 00000000000..e420cab824f --- /dev/null +++ b/source/source_relax/test/socket_frame_test.cpp @@ -0,0 +1,377 @@ +#include "../socket_frame.h" + +#include "gtest/gtest.h" + +#include +#include +#include +#include +#include + +namespace +{ +using SocketFrame::CellValidation; +using SocketFrame::Matrix9; +using SocketFrame::VirialConversion; +using SocketFrame::checked_position_count; +using SocketFrame::make_ipi_virial; +using SocketFrame::transpose_matrix9; +using SocketFrame::validate_ipi_cell; +using SocketFrame::validate_positions; + +const double EPSILON = std::numeric_limits::epsilon(); + +CellValidation validate_with_driver_thresholds(const Matrix9& cell, const Matrix9& inverse) +{ + return validate_ipi_cell(cell, inverse, 1.0e12, 64.0 * EPSILON, 64.0); +} + +void expect_matrix_near(const Matrix9& expected, const Matrix9& actual, double tolerance) +{ + for (std::size_t index = 0; index < expected.size(); ++index) + { + EXPECT_NEAR(expected[index], actual[index], tolerance) << "matrix index " << index; + } +} +} // namespace + +TEST(SocketFrameTest, TransposeKeepsAllNineUniqueEntries) +{ + Matrix9 in = {{1, 2, 3, 4, 5, 6, 7, 8, 9}}; + Matrix9 expected = {{1, 4, 7, 2, 5, 8, 3, 6, 9}}; + EXPECT_EQ(expected, transpose_matrix9(in)); +} + +TEST(SocketFrameTest, VirialUsesPositiveHalfVolumeAndWireTranspose) +{ + Matrix9 stress = {{1, 2, 3, 2, 5, 6, 3, 6, 9}}; + VirialConversion out = make_ipi_virial(stress, 4.0, 1e-12, 1e-12); + Matrix9 expected = {{2, 4, 6, 4, 10, 12, 6, 12, 18}}; + ASSERT_TRUE(out.ok) << out.message; + EXPECT_EQ(expected, out.wire_virial_hartree); +} + +TEST(SocketFrameTest, RightHandedTriclinicCellReturnsKnownInverse) +{ + const Matrix9 cell = {{2.0, 1.0, 0.0, 0.0, 3.0, 1.0, 0.0, 0.0, 4.0}}; + const Matrix9 inverse = {{0.5, -1.0 / 6.0, 1.0 / 24.0, + 0.0, 1.0 / 3.0, -1.0 / 12.0, + 0.0, 0.0, 0.25}}; + + const CellValidation out = validate_with_driver_thresholds(cell, inverse); + + ASSERT_TRUE(out.ok) << out.message; + EXPECT_DOUBLE_EQ(24.0, out.determinant_bohr3); + EXPECT_NEAR(0.0, out.inverse_residual, 16.0 * EPSILON); + expect_matrix_near(inverse, out.computed_inverse_wire_bohr_inv, 16.0 * EPSILON); +} + +TEST(SocketFrameTest, AseTriclinicInverseWireLayoutIsAcceptedAndRecomputedFromCell) +{ + // ASE stores row lattice vectors A in Angstrom, sends H = A^T / Bohr, + // and sends pinv(A) * Bohr as the inverse field. For this nonsingular + // cell that received field is inv(H)^T, not inv(H). + const double bohr_angstrom = 0.5291772105638411; + const Matrix9 cell_wire = {{5.0 / bohr_angstrom, 0.5 / bohr_angstrom, 0.25 / bohr_angstrom, + 0.0, 4.0 / bohr_angstrom, 0.75 / bohr_angstrom, + 0.0, 0.0, 3.0 / bohr_angstrom}}; + const Matrix9 ase_inverse_wire = {{bohr_angstrom / 5.0, 0.0, 0.0, + -bohr_angstrom / 40.0, bohr_angstrom / 4.0, 0.0, + -bohr_angstrom / 96.0, -bohr_angstrom / 16.0, + bohr_angstrom / 3.0}}; + const Matrix9 inverse_computed_from_cell = {{bohr_angstrom / 5.0, + -bohr_angstrom / 40.0, + -bohr_angstrom / 96.0, + 0.0, + bohr_angstrom / 4.0, + -bohr_angstrom / 16.0, + 0.0, + 0.0, + bohr_angstrom / 3.0}}; + + const CellValidation out = validate_with_driver_thresholds(cell_wire, ase_inverse_wire); + + ASSERT_TRUE(out.ok) << out.message; + EXPECT_NEAR(0.0, out.inverse_residual, 16.0 * EPSILON); + expect_matrix_near(inverse_computed_from_cell, + out.computed_inverse_wire_bohr_inv, + 16.0 * EPSILON); +} + +TEST(SocketFrameTest, RotatedDiagonalTracksRightSingularVectorsAndInverseOrder) +{ + // Hand-multiplied U diag(5, 2, 0.5) V^T, with rational plane rotations. + const Matrix9 cell = {{4.0, -0.72, -0.96, + 3.0, 0.96, 1.28, + 0.0, -0.4, 0.3}}; + const Matrix9 inverse = {{0.16, 0.12, 0.0, + -0.18, 0.24, -1.6, + -0.24, 0.32, 1.2}}; + + const CellValidation out = validate_with_driver_thresholds(cell, inverse); + + ASSERT_TRUE(out.ok) << out.message; + EXPECT_NEAR(5.0, out.determinant_bohr3, 64.0 * EPSILON); + EXPECT_NEAR(10.0, out.condition_number_2, 256.0 * EPSILON); + expect_matrix_near(inverse, out.computed_inverse_wire_bohr_inv, 64.0 * EPSILON); +} + +TEST(SocketFrameTest, InconsistentReceivedInverseIsRejected) +{ + const Matrix9 cell = {{2.0, 1.0, 0.0, 0.0, 3.0, 1.0, 0.0, 0.0, 4.0}}; + // This is neither inv(cell) nor inv(cell)^T, so both supported wire + // layouts must reject it. + const Matrix9 wrong_inverse = {{0.6, -1.0 / 6.0, 1.0 / 24.0, + 0.0, 1.0 / 3.0, -1.0 / 12.0, + 0.0, 0.0, 0.25}}; + + const CellValidation out = validate_with_driver_thresholds(cell, wrong_inverse); + + EXPECT_FALSE(out.ok); + EXPECT_NE(std::string::npos, out.message.find("inverse")); + EXPECT_GT(out.inverse_residual, 0.1); +} + +TEST(SocketFrameTest, ReceivedInverseResidualUsesConditionScaledRelativeTolerance) +{ + const Matrix9 cell = {{1.0, 0.0, 0.0, + 0.0, 1.0, 0.0, + 0.0, 0.0, 1.0e-6}}; + Matrix9 accepted_inverse = {{1.0 + 1.0e-8, 0.0, 0.0, + 0.0, 1.0, 0.0, + 0.0, 0.0, 1.0e6}}; + Matrix9 rejected_inverse = accepted_inverse; + rejected_inverse[0] = 1.0 + 2.0e-8; + + const CellValidation accepted + = validate_ipi_cell(cell, accepted_inverse, 1.0e12, 0.0, 64.0); + const CellValidation rejected + = validate_ipi_cell(cell, rejected_inverse, 1.0e12, 0.0, 64.0); + + ASSERT_TRUE(accepted.ok) << accepted.message; + EXPECT_DOUBLE_EQ(1.0e6, accepted.condition_number_2); + EXPECT_NEAR(1.0e-8, accepted.inverse_residual, EPSILON); + EXPECT_FALSE(rejected.ok); + EXPECT_NE(std::string::npos, rejected.message.find("inverse")); + EXPECT_NEAR(2.0e-8, rejected.inverse_residual, EPSILON); +} + +TEST(SocketFrameTest, NegativeAndZeroDeterminantsAreRejected) +{ + const Matrix9 identity = {{1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0}}; + const Matrix9 left_handed = {{-1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0}}; + const Matrix9 singular = {{1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0}}; + + EXPECT_FALSE(validate_with_driver_thresholds(left_handed, identity).ok); + EXPECT_FALSE(validate_with_driver_thresholds(singular, identity).ok); +} + +TEST(SocketFrameTest, NonrepresentablePositiveCellVolumeIsRejected) +{ + const Matrix9 huge_cell = {{1.0e200, 0.0, 0.0, + 0.0, 1.0e200, 0.0, + 0.0, 0.0, 1.0e200}}; + const Matrix9 tiny_inverse = {{1.0e-200, 0.0, 0.0, + 0.0, 1.0e-200, 0.0, + 0.0, 0.0, 1.0e-200}}; + + const CellValidation out = validate_with_driver_thresholds(huge_cell, tiny_inverse); + + EXPECT_FALSE(out.ok); + EXPECT_NE(std::string::npos, out.message.find("determinant")); +} + +TEST(SocketFrameTest, UnderflowedCellVolumeIsRejectedAsZeroDeterminant) +{ + const Matrix9 tiny_cell = {{1.0e-200, 0.0, 0.0, + 0.0, 1.0e-200, 0.0, + 0.0, 0.0, 1.0e-200}}; + const Matrix9 huge_inverse = {{1.0e200, 0.0, 0.0, + 0.0, 1.0e200, 0.0, + 0.0, 0.0, 1.0e200}}; + + const CellValidation out = validate_with_driver_thresholds(tiny_cell, huge_inverse); + + EXPECT_FALSE(out.ok); + EXPECT_NE(std::string::npos, out.message.find("determinant")); +} + +TEST(SocketFrameTest, ConditionNumberMustBeStrictlyBelowMaximum) +{ + const Matrix9 below = {{1.0, 0.0, 0.0, 0.0, 1.0e-6, 0.0, 0.0, 0.0, 2.0e-12}}; + const Matrix9 below_inverse = {{1.0, 0.0, 0.0, 0.0, 1.0e6, 0.0, 0.0, 0.0, 5.0e11}}; + const Matrix9 at = {{1.0, 0.0, 0.0, 0.0, 1.0e-6, 0.0, 0.0, 0.0, 1.0e-12}}; + const Matrix9 at_inverse = {{1.0, 0.0, 0.0, 0.0, 1.0e6, 0.0, 0.0, 0.0, 1.0e12}}; + const Matrix9 above = {{1.0, 0.0, 0.0, 0.0, 1.0e-6, 0.0, 0.0, 0.0, 5.0e-13}}; + const Matrix9 above_inverse = {{1.0, 0.0, 0.0, 0.0, 1.0e6, 0.0, 0.0, 0.0, 2.0e12}}; + + EXPECT_TRUE(validate_with_driver_thresholds(below, below_inverse).ok); + const CellValidation boundary = validate_with_driver_thresholds(at, at_inverse); + EXPECT_FALSE(boundary.ok); + EXPECT_NE(std::string::npos, boundary.message.find("condition")); + EXPECT_DOUBLE_EQ(1.0e12, boundary.condition_number_2); + EXPECT_FALSE(validate_with_driver_thresholds(above, above_inverse).ok); +} + +TEST(SocketFrameTest, NonfiniteCellOrReceivedInverseIsRejected) +{ + const double nan = std::numeric_limits::quiet_NaN(); + const double infinity = std::numeric_limits::infinity(); + const Matrix9 identity = {{1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0}}; + Matrix9 bad_cell = identity; + Matrix9 bad_inverse = identity; + bad_cell[4] = nan; + EXPECT_FALSE(validate_with_driver_thresholds(bad_cell, identity).ok); + bad_cell = identity; + bad_cell[7] = infinity; + EXPECT_FALSE(validate_with_driver_thresholds(bad_cell, identity).ok); + bad_inverse[1] = nan; + EXPECT_FALSE(validate_with_driver_thresholds(identity, bad_inverse).ok); + bad_inverse = identity; + bad_inverse[8] = -infinity; + EXPECT_FALSE(validate_with_driver_thresholds(identity, bad_inverse).ok); +} + +TEST(SocketFrameTest, PositionCountRequiresMatchingNonnegativeAtomCount) +{ + std::size_t coordinate_count = 77; + std::string message; + + EXPECT_FALSE(checked_position_count(-1, 2, coordinate_count, message)); + EXPECT_EQ(77u, coordinate_count); + EXPECT_NE(std::string::npos, message.find("match")); + + message.clear(); + EXPECT_FALSE(checked_position_count(-1, -1, coordinate_count, message)); + EXPECT_EQ(77u, coordinate_count); + EXPECT_NE(std::string::npos, message.find("negative")); + + message.clear(); + EXPECT_TRUE(checked_position_count(3, 3, coordinate_count, message)) << message; + EXPECT_EQ(9u, coordinate_count); + EXPECT_TRUE(message.empty()); +} + +TEST(SocketFrameTest, PositionCountRejectsMismatchBeforeDerivingAllocationSize) +{ + std::size_t coordinate_count = 123; + std::string message; + + EXPECT_FALSE(checked_position_count(std::numeric_limits::max(), + 1, + coordinate_count, + message)); + EXPECT_EQ(123u, coordinate_count); + EXPECT_NE(std::string::npos, message.find("match")); +} + +TEST(SocketFrameTest, PositionsRequireExactSizeAndFiniteCoordinates) +{ + std::string message; + const std::vector valid = {1.0, -2.0, 3.0}; + EXPECT_TRUE(validate_positions(valid, 3, message)) << message; + + message.clear(); + EXPECT_FALSE(validate_positions(valid, 6, message)); + EXPECT_NE(std::string::npos, message.find("count")); + + std::vector nonfinite = valid; + nonfinite[1] = std::numeric_limits::quiet_NaN(); + message.clear(); + EXPECT_FALSE(validate_positions(nonfinite, 3, message)); + EXPECT_NE(std::string::npos, message.find("finite")); + + nonfinite[1] = std::numeric_limits::infinity(); + message.clear(); + EXPECT_FALSE(validate_positions(nonfinite, 3, message)); + EXPECT_NE(std::string::npos, message.find("finite")); +} + +TEST(SocketFrameTest, SmallStressAsymmetryIsAveragedBeforeConversion) +{ + const Matrix9 stress = {{1.0, 2.1, 3.2, + 1.9, 5.0, 6.3, + 2.8, 5.7, 9.0}}; + const Matrix9 expected = {{1.0, 2.0, 3.0, + 2.0, 5.0, 6.0, + 3.0, 6.0, 9.0}}; + + const VirialConversion out = make_ipi_virial(stress, 2.0, 0.61, 0.0); + + ASSERT_TRUE(out.ok) << out.message; + expect_matrix_near(expected, out.wire_virial_hartree, 4.0 * EPSILON); + EXPECT_NEAR(0.6, out.max_antisymmetric_component, 4.0 * EPSILON); +} + +TEST(SocketFrameTest, ExcessiveStressAsymmetryIsRejected) +{ + const Matrix9 stress = {{1.0, 2.1, 3.2, + 1.9, 5.0, 6.3, + 2.8, 5.7, 9.0}}; + + const VirialConversion out = make_ipi_virial(stress, 2.0, 0.59, 0.0); + + EXPECT_FALSE(out.ok); + EXPECT_NE(std::string::npos, out.message.find("symmetric")); + EXPECT_NEAR(0.6, out.max_antisymmetric_component, 4.0 * EPSILON); +} + +TEST(SocketFrameTest, StressAsymmetryUsesAbsolutePlusRelativeTolerance) +{ + const Matrix9 accepted_stress = {{10.0, 2.0 + 4.0e-8, 3.0, + 2.0 - 4.0e-8, 5.0, 6.0, + 3.0, 6.0, 9.0}}; + Matrix9 rejected_stress = accepted_stress; + rejected_stress[1] = 2.0 + 6.0e-8; + rejected_stress[3] = 2.0 - 6.0e-8; + const Matrix9 expected = {{10.0, 2.0, 3.0, + 2.0, 5.0, 6.0, + 3.0, 6.0, 9.0}}; + + const VirialConversion accepted + = make_ipi_virial(accepted_stress, 2.0, 1.0e-10, 1.0e-8); + const VirialConversion rejected + = make_ipi_virial(rejected_stress, 2.0, 1.0e-10, 1.0e-8); + + ASSERT_TRUE(accepted.ok) << accepted.message; + expect_matrix_near(expected, accepted.wire_virial_hartree, 4.0 * EPSILON); + EXPECT_NEAR(8.0e-8, accepted.max_antisymmetric_component, EPSILON); + EXPECT_FALSE(rejected.ok); + EXPECT_NE(std::string::npos, rejected.message.find("symmetric")); + EXPECT_NEAR(1.2e-7, rejected.max_antisymmetric_component, EPSILON); +} + +TEST(SocketFrameTest, NonpositiveOrNonfiniteVolumeIsRejected) +{ + const Matrix9 zero_stress = {{0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0}}; + + EXPECT_FALSE(make_ipi_virial(zero_stress, 0.0, 1.0e-10, 1.0e-8).ok); + EXPECT_FALSE(make_ipi_virial(zero_stress, -1.0, 1.0e-10, 1.0e-8).ok); + EXPECT_FALSE(make_ipi_virial(zero_stress, + std::numeric_limits::infinity(), + 1.0e-10, + 1.0e-8) + .ok); +} + +TEST(SocketFrameTest, FiniteStressAndVolumeRejectConvertedVirialOverflow) +{ + const double largest_finite = std::numeric_limits::max(); + const Matrix9 stress = {{largest_finite, 0.0, 0.0, + 0.0, 1.0, 0.0, + 0.0, 0.0, 1.0}}; + + const VirialConversion out = make_ipi_virial(stress, 4.0, 1.0e-10, 1.0e-8); + + EXPECT_FALSE(out.ok); + EXPECT_NE(std::string::npos, out.message.find("representable")); +} + +TEST(SocketFrameTest, NonfiniteStressIsRejected) +{ + Matrix9 stress = {{1.0, 2.0, 3.0, 2.0, 5.0, 6.0, 3.0, 6.0, 9.0}}; + stress[2] = std::numeric_limits::quiet_NaN(); + EXPECT_FALSE(make_ipi_virial(stress, 4.0, 1.0e-10, 1.0e-8).ok); + stress[2] = std::numeric_limits::infinity(); + EXPECT_FALSE(make_ipi_virial(stress, 4.0, 1.0e-10, 1.0e-8).ok); +} diff --git a/source/source_relax/test/socket_ipi_test.cpp b/source/source_relax/test/socket_ipi_test.cpp new file mode 100644 index 00000000000..c2b7f415334 --- /dev/null +++ b/source/source_relax/test/socket_ipi_test.cpp @@ -0,0 +1,474 @@ +#include "../socket_ipi.h" + +#include "gtest/gtest.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace +{ +constexpr std::size_t IPI_HEADER_LEN = 12; + +std::string errno_message(const std::string& prefix) +{ + return prefix + ": " + std::strerror(errno); +} + +void send_all(int fd, const void* data, std::size_t nbytes) +{ + const char* cursor = static_cast(data); + std::size_t done = 0; + while (done < nbytes) + { + const ssize_t sent = ::send(fd, cursor + done, nbytes - done, 0); + if (sent < 0) + { + if (errno == EINTR) + { + continue; + } + throw std::runtime_error(errno_message("send failed")); + } + if (sent == 0) + { + throw std::runtime_error("send returned zero"); + } + done += static_cast(sent); + } +} + +void recv_all(int fd, void* data, std::size_t nbytes) +{ + char* cursor = static_cast(data); + std::size_t done = 0; + while (done < nbytes) + { + const ssize_t received = ::recv(fd, cursor + done, nbytes - done, 0); + if (received < 0) + { + if (errno == EINTR) + { + continue; + } + throw std::runtime_error(errno_message("recv failed")); + } + if (received == 0) + { + throw std::runtime_error("socket closed while receiving test data"); + } + done += static_cast(received); + } +} + +std::string padded_header(const std::string& header) +{ + std::string padded = header; + padded.resize(IPI_HEADER_LEN, ' '); + return padded; +} + +class UnixSocketServer +{ + public: + UnixSocketServer() + { + char dir_template[] = "/tmp/abacus_socket_ipi_test_XXXXXX"; + char* made_dir = ::mkdtemp(dir_template); + if (made_dir == nullptr) + { + throw std::runtime_error(errno_message("mkdtemp failed")); + } + dir_ = made_dir; + path_ = dir_ + "/ipi.sock"; + + listen_fd_ = ::socket(AF_UNIX, SOCK_STREAM, 0); + if (listen_fd_ < 0) + { + throw std::runtime_error(errno_message("socket failed")); + } + + sockaddr_un addr; + std::memset(&addr, 0, sizeof(addr)); + addr.sun_family = AF_UNIX; + std::strncpy(addr.sun_path, path_.c_str(), sizeof(addr.sun_path) - 1); + if (::bind(listen_fd_, reinterpret_cast(&addr), sizeof(addr)) != 0) + { + throw std::runtime_error(errno_message("bind failed")); + } + if (::listen(listen_fd_, 1) != 0) + { + throw std::runtime_error(errno_message("listen failed")); + } + } + + ~UnixSocketServer() + { + if (listen_fd_ >= 0) + { + ::close(listen_fd_); + } + if (!path_.empty()) + { + ::unlink(path_.c_str()); + } + if (!dir_.empty()) + { + ::rmdir(dir_.c_str()); + } + } + + UnixSocketServer(const UnixSocketServer&) = delete; + UnixSocketServer& operator=(const UnixSocketServer&) = delete; + + std::string address() const + { + return path_ + ":UNIX"; + } + + int accept_once() + { + const int fd = ::accept(listen_fd_, nullptr, nullptr); + if (fd < 0) + { + throw std::runtime_error(errno_message("accept failed")); + } + return fd; + } + + private: + int listen_fd_ = -1; + std::string dir_; + std::string path_; +}; + +void rethrow_thread_error(const std::exception_ptr& thread_error) +{ + if (thread_error) + { + std::rethrow_exception(thread_error); + } +} +} // namespace + +TEST(IpiSocketTest, WriteHeaderPadsToTwelveBytes) +{ + UnixSocketServer server; + std::string received; + std::exception_ptr thread_error; + std::thread peer([&]() { + try + { + const int fd = server.accept_once(); + char buffer[IPI_HEADER_LEN]; + recv_all(fd, buffer, sizeof(buffer)); + received.assign(buffer, sizeof(buffer)); + ::close(fd); + } + catch (...) + { + thread_error = std::current_exception(); + } + }); + + IpiSocket socket; + socket.connect(server.address()); + socket.write_header("READY"); + socket.close(); + + peer.join(); + rethrow_thread_error(thread_error); + EXPECT_EQ(padded_header("READY"), received); +} + +TEST(IpiSocketTest, CleanPeerCloseBeforeNextHeaderThrowsDedicatedSignal) +{ + UnixSocketServer server; + std::exception_ptr thread_error; + std::thread peer([&]() { + try + { + const int fd = server.accept_once(); + const std::string header = padded_header("STATUS"); + send_all(fd, header.data(), header.size()); + ::close(fd); + } + catch (...) + { + thread_error = std::current_exception(); + } + }); + + IpiSocket socket; + socket.connect(server.address()); + EXPECT_EQ("STATUS", socket.read_header()); + EXPECT_THROW(socket.read_header(), IpiSocketClosed); + socket.close(); + + peer.join(); + rethrow_thread_error(thread_error); +} + +TEST(IpiSocketTest, PartialHeaderCloseStaysRuntimeError) +{ + UnixSocketServer server; + std::exception_ptr thread_error; + std::thread peer([&]() { + try + { + const int fd = server.accept_once(); + const std::string partial = "STAT"; + send_all(fd, partial.data(), partial.size()); + ::close(fd); + } + catch (...) + { + thread_error = std::current_exception(); + } + }); + + IpiSocket socket; + socket.connect(server.address()); + try + { + static_cast(socket.read_header()); + FAIL() << "partial header EOF should throw"; + } + catch (const IpiSocketClosed&) + { + FAIL() << "partial header EOF must not be treated as clean peer close"; + } + catch (const std::runtime_error& exc) + { + EXPECT_NE(std::string::npos, std::string(exc.what()).find("closed while reading header")); + } + socket.close(); + + peer.join(); + rethrow_thread_error(thread_error); +} + +TEST(IpiSocketTest, Int32UsesExactlyFourNativeEndianBytes) +{ + UnixSocketServer server; + const std::int32_t expected = INT32_C(0x12345678); + std::vector received(4); + std::exception_ptr thread_error; + std::thread peer([&]() { + try + { + const int fd = server.accept_once(); + recv_all(fd, received.data(), received.size()); + ::close(fd); + } + catch (...) + { + thread_error = std::current_exception(); + } + }); + + IpiSocket socket; + socket.connect(server.address()); + socket.write_int32(expected); + socket.close(); + + peer.join(); + rethrow_thread_error(thread_error); + EXPECT_EQ(0, std::memcmp(received.data(), &expected, 4)); +} + +TEST(IpiSocketTest, DoubleUsesExactlyEightNativeEndianBytes) +{ + UnixSocketServer server; + const double expected = -1234.5; + std::vector received(8); + std::exception_ptr thread_error; + std::thread peer([&]() { + try + { + const int fd = server.accept_once(); + recv_all(fd, received.data(), received.size()); + ::close(fd); + } + catch (...) + { + thread_error = std::current_exception(); + } + }); + + IpiSocket socket; + socket.connect(server.address()); + socket.write_double(expected); + socket.close(); + + peer.join(); + rethrow_thread_error(thread_error); + EXPECT_EQ(0, std::memcmp(received.data(), &expected, 8)); +} + +TEST(IpiSocketTest, ReadInt32HandlesSplitPayload) +{ + UnixSocketServer server; + const std::int32_t expected = INT32_C(0x12345678); + std::exception_ptr thread_error; + std::thread peer([&]() { + try + { + const int fd = server.accept_once(); + const char* bytes = reinterpret_cast(&expected); + send_all(fd, bytes, 2); + send_all(fd, bytes + 2, 2); + ::close(fd); + } + catch (...) + { + thread_error = std::current_exception(); + } + }); + + IpiSocket socket; + socket.connect(server.address()); + EXPECT_EQ(expected, socket.read_int32()); + socket.close(); + + peer.join(); + rethrow_thread_error(thread_error); +} + +TEST(IpiSocketTest, ReadInt32RejectsMidPayloadClose) +{ + UnixSocketServer server; + const std::int32_t value = INT32_C(0x12345678); + std::exception_ptr thread_error; + std::thread peer([&]() { + try + { + const int fd = server.accept_once(); + send_all(fd, &value, 2); + ::close(fd); + } + catch (...) + { + thread_error = std::current_exception(); + } + }); + + IpiSocket socket; + socket.connect(server.address()); + EXPECT_THROW(socket.read_int32(), IpiSocketClosed); + socket.close(); + + peer.join(); + rethrow_thread_error(thread_error); +} + +TEST(IpiSocketTest, WriteDoublesCompletesLargePayloadWithSmallPeerReads) +{ + UnixSocketServer server; + std::vector expected(1 << 18); + for (std::size_t i = 0; i < expected.size(); ++i) + { + expected[i] = -1234.5 + static_cast(i) * 0.25; + } + std::vector received(expected.size() * sizeof(double)); + std::exception_ptr thread_error; + std::thread peer([&]() { + try + { + const int fd = server.accept_once(); + std::size_t done = 0; + while (done < received.size()) + { + const std::size_t remaining = received.size() - done; + const std::size_t chunk = remaining < 37 ? remaining : 37; + const ssize_t nread = ::recv(fd, received.data() + done, chunk, 0); + if (nread < 0) + { + if (errno == EINTR) + { + continue; + } + throw std::runtime_error(errno_message("recv failed")); + } + if (nread == 0) + { + throw std::runtime_error("socket closed while receiving large payload"); + } + done += static_cast(nread); + } + ::close(fd); + } + catch (...) + { + thread_error = std::current_exception(); + } + }); + + IpiSocket socket; + socket.connect(server.address()); + socket.write_doubles(expected); + socket.close(); + + peer.join(); + rethrow_thread_error(thread_error); + EXPECT_EQ(0, std::memcmp(received.data(), expected.data(), received.size())); +} + +TEST(IpiSocketTest, ReadDoublesRejectsByteCountOverflow) +{ + IpiSocket socket; + const std::size_t count = std::numeric_limits::max() / sizeof(double) + 1; + + try + { + static_cast(socket.read_doubles(count)); + FAIL() << "overflowing double payload size should throw"; + } + catch (const std::overflow_error& exc) + { + EXPECT_NE(std::string::npos, std::string(exc.what()).find(std::to_string(count))); + } + catch (...) + { + FAIL() << "overflowing double payload size should throw std::overflow_error"; + } +} + +TEST(IpiSocketTest, WriteStringSendsExactBytesWithoutTerminator) +{ + UnixSocketServer server; + const std::string expected = "{\"scf_converged\":false}"; + std::vector received(expected.size()); + std::exception_ptr thread_error; + std::thread peer([&]() { + try + { + const int fd = server.accept_once(); + recv_all(fd, received.data(), received.size()); + ::close(fd); + } + catch (...) + { + thread_error = std::current_exception(); + } + }); + + IpiSocket socket; + socket.connect(server.address()); + socket.write_string(expected); + socket.close(); + + peer.join(); + rethrow_thread_error(thread_error); + EXPECT_EQ(expected, std::string(received.begin(), received.end())); +}