SharedMemoryCleaner.java

/*
 * junixsocket
 *
 * Copyright 2009-2026 Christian Kohlschütter
 *
 * Licensed under the Apache License, Version 2.0 (the "License");
 * you may not use this file except in compliance with the License.
 * You may obtain a copy of the License at
 *
 *     http://www.apache.org/licenses/LICENSE-2.0
 *
 * Unless required by applicable law or agreed to in writing, software
 * distributed under the License is distributed on an "AS IS" BASIS,
 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
 * See the License for the specific language governing permissions and
 * limitations under the License.
 */
package org.newsclub.net.unix.memory;

import java.io.FileDescriptor;
import java.io.IOException;
import java.lang.foreign.Arena;
import java.lang.foreign.MemorySegment;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Map.Entry;
import java.util.Objects;
import java.util.WeakHashMap;

import org.newsclub.net.unix.CleanableState;
import org.newsclub.net.unix.MemoryImplUtilInternal;

import com.kohlschutter.annotations.compiletime.SuppressFBWarnings;

class SharedMemoryCleaner extends CleanableState {
  private final Map<MemorySegment, Integer> segments = new LinkedHashMap<>();
  private final Map<Futex, Futex> futexes = new WeakHashMap<>();
  private Arena arena;
  private final boolean closeArena;
  private MemorySegment arenaSegment;
  final FileDescriptor fd;

  SharedMemoryCleaner(Arena arena, Object observed, FileDescriptor fd) {
    super(observed);
    if (arena == null) {
      this.arena = null;
      this.closeArena = true;
    } else {
      this.arena = arena;
      this.closeArena = false;
    }
    this.fd = fd;
  }

  void registerMemorySegment(MemorySegment ms, int duplicates) {
    synchronized (segments) {
      segments.put(ms, duplicates);
    }
  }

  @Override
  @SuppressWarnings("PMD.CognitiveComplexity")
  @SuppressFBWarnings("USO_UNSAFE_OBJECT_SYNCHRONIZATION")
  protected synchronized void doClean() throws IOException {
    if (!SharedMemory.isUtilLoaded()) {
      // Nothing to do
      return;
    }

    Map<FileDescriptor, Long> map = SharedMemory.FD_MEMORY;
    if (map != null && fd != null) {
      synchronized (map) {
        map.remove(fd);
      }
    }

    synchronized (futexes) {
      if (!futexes.isEmpty()) {
        for (Futex f : futexes.keySet()) {
          try {
            f.tryWake(true); // unblock waiting threads
            f.close();
          } catch (Exception e) { // NOPMD
            // ignore
          }
        }
        futexes.clear();
      }
    }

    if (closeArena && arena != null) {
      arena.close();
    }

    MemoryImplUtilInternal util = SharedMemory.getUtil();

    IOException exc = null;
    if (fd != null && fd.valid()) {
      try {
        util.close(fd);
      } catch (IOException e) {
        exc = e;
      }
    }

    List<Entry<MemorySegment, Integer>> list;
    synchronized (segments) {
      list = new ArrayList<>(segments.entrySet()).reversed();
      segments.clear();
    }
    for (Map.Entry<MemorySegment, Integer> en : list) {
      MemorySegment ms = en.getKey();
      int duplicates = en.getValue();

      long addr = ms.address();
      long length = ms.byteSize();
      if (ms.scope().isAlive()) {
        util.madvise(addr, length, MemoryImplUtilInternal.MADV_FREE_NOW, true);
        continue;
      }

      try {
        util.unmap(addr, length, duplicates, false);
      } catch (IOException e) {
        if (exc == null) {
          exc = e;
        } else {
          exc.addSuppressed(e);
        }
      }
    }

    if (exc != null) {
      throw exc;
    }
  }

  public boolean isCovered(MemorySegment segment) {
    Objects.requireNonNull(segment);
    long start = segment.address();
    long end = start + segment.byteSize();
    synchronized (segments) {
      for (MemorySegment ms : segments.keySet()) {
        if (ms == segment) { // NOPMD
          return true;
        }
        long addr = ms.address();
        if (start >= addr && end <= addr + ms.byteSize()) {
          return true;
        }
      }
    }
    return false;
  }

  public synchronized MemorySegment getArenaSegment() {
    if (this.arenaSegment == null) {
      this.arenaSegment = getArena().allocate(0);
    }
    return arenaSegment;
  }

  public void registerFutex(Futex futex) {
    synchronized (futexes) {
      futexes.put(futex, futex);
    }
  }

  public void checkCovered(MemorySegment addr) throws IOException {
    if (!isCovered(addr)) {
      throw new IOException("Not a MemorySegment of ours");
    }
  }

  public synchronized Arena getArena() {
    if (this.arena == null) {
      this.arena = Arena.ofShared();
    }
    return this.arena;
  }
}