Files

577 lines
17 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"errors"
"io"
"net"
"sync"
"time"
"sourcedock.dev/petrbalvin/tensor/internal/base"
)
// Default bounds of a networked world: one collective's wait, one
// frame's payload, and the number of frames a link may queue. All
// three exist so that a stuck or hostile peer is an error, never a
// hang and never an allocation.
const (
defaultTimeout = 10 * time.Minute
defaultMaxMessage = int64(16) << 30
linkQueue = 64
// pendingCap bounds how many unmatched frames one recvFrom may
// hold: a peer flooding frames nobody asked for is an error,
// never an unbounded allocator.
pendingCap = 4096
)
// Options bounds a networked world. An in-process world from Launch
// takes none of them: its links are channels, and progress is the
// program's own business, as it is in MPI.
type Options struct {
// Timeout bounds one collective's wait on the network: dialing,
// the handshake and every send and receive carry it as a
// deadline, refreshed each time a frame moves. Zero means the
// default of ten minutes; a negative value means no deadline at
// all.
Timeout time.Duration
// MaxMessage is the largest frame payload the world accepts, in
// bytes. A peer announcing more is refused before any allocation.
// Zero means the default of 16 GiB.
MaxMessage int64
}
func (o Options) timeout() time.Duration {
switch {
case o.Timeout > 0:
return o.Timeout
case o.Timeout < 0:
return 0 // no deadline
default:
return defaultTimeout
}
}
func (o Options) maxMessage() int64 {
if o.MaxMessage > 0 {
return o.MaxMessage
}
return defaultMaxMessage
}
// World is one rank's end of an SPMD world: its place in it, the links
// to the other ranks and the collectives. A World is driven by one
// goroutine: like an MPI rank, it never runs two collectives at once.
type World struct {
rank int
size int
timeout time.Duration
maxMessage int64
networked bool
peers []*peer // peers[r] is the link to rank r; nil for this rank
closers []io.Closer
// The hub's readers and writers. The writers are waited on before
// an orderly close, so the last collective's frames are on the
// wire before the connections go away; the readers are waited on
// after, once those connections have broken their blocking reads.
drainWg sync.WaitGroup
pumpWg sync.WaitGroup
done chan struct{} // closed when the world failed or left
closeDone sync.Once
failErr error
}
// Rank returns this rank's index, from 0 to Size-1.
func (w *World) Rank() int { return w.rank }
// Size returns the number of ranks in the world.
func (w *World) Size() int { return w.size }
// Close ends a networked world: whatever the collectives queued is
// written to the wire, then the connections close and the other ranks
// see the departure as their next receive failing. On an in-process
// world it is a no-op, because Launch tears the world down.
func (w *World) Close() error {
w.leave()
w.drainWg.Wait()
var first error
for _, c := range w.closers {
if err := c.Close(); err != nil && first == nil {
first = err
}
}
w.pumpWg.Wait()
return first
}
// leave closes done without recording a failure: the ordinary exit of
// the rank's program.
func (w *World) leave() {
w.closeDone.Do(func() { close(w.done) })
}
// fail records the world's first failure, closes done so that every
// wait on this world wakes, and closes the networked links. It returns
// the failure an outside caller should see.
func (w *World) fail(err error) error {
w.closeDone.Do(func() {
w.failErr = err
close(w.done)
for _, c := range w.closers {
c.Close()
}
})
return w.status()
}
// status is the entry check every public operation makes: a world that
// failed, or a rank whose program already returned, answers with an
// error and never with data.
func (w *World) status() error {
select {
case <-w.done:
if w.failErr != nil {
return base.Errf("spmd: rank %d world is in a failed state: %v", w.rank, w.failErr)
}
return base.Errf("spmd: rank %d world has left", w.rank)
default:
return nil
}
}
// deadline is the absolute time one collective may run to on a
// networked world, refreshed each time a frame moves; the zero time
// means no deadline, which is what an in-process world always gets.
func (w *World) deadline() time.Time {
if !w.networked || w.timeout <= 0 {
return time.Time{}
}
return time.Now().Add(w.timeout)
}
// depart closes the peer's gone channel: the peer closed its side of
// the link, so no further frame will ever come from it.
func (p *peer) depart() {
select {
case <-p.gone:
default:
close(p.gone)
}
}
// hasLeft reports whether the peer closed its side of the link.
func (p *peer) hasLeft() bool {
if p.gone == nil {
return false
}
select {
case <-p.gone:
return true
default:
return false
}
}
// sendTo carries one frame to rank r. On an in-process world it walks
// the direct channel; on a networked world it walks this rank's own
// link to the hub, because rank 0 routes every frame by its
// destination.
func (w *World) sendTo(r int, tag uint8, payload []byte) error {
if err := w.status(); err != nil {
return err
}
m := message{tag: tag, from: w.rank, dest: r, data: payload}
p := w.peers[r]
if p != nil && p.outbox != nil {
select {
case p.outbox <- m:
// On a buffered networked link a queued frame is not yet
// a delivered frame, so a peer that left between the two
// dropped it. An in-process handoff is the delivery
// itself: the receiver taking the frame and then leaving
// is its own healthy business.
if w.networked && p.hasLeft() {
return w.fail(base.Errf("spmd: rank %d sending to rank %d: rank %d left the world", w.rank, r, r))
}
return nil
case <-p.gone:
return w.fail(base.Errf("spmd: rank %d sending to rank %d: rank %d left the world", w.rank, r, r))
case <-w.done:
return w.status()
}
}
// A networked rank that is not the hub owns one connection, and
// every word it sends rides it; the hub reads the destination.
if err := w.peers[0].conn.writeFrame(m, w.deadline()); err != nil {
return w.fail(base.Errf("spmd: rank %d sending to rank %d: %v", w.rank, r, err))
}
return nil
}
// recvFrom returns the payload of the next frame rank r sent under the
// wanted tag, holding frames that arrived earlier for other ranks or
// tags until the collective asks for them: no answer ever depends on
// the arrival order. Any failure fails the world.
func (w *World) recvFrom(r int, tag uint8) ([]byte, error) {
if err := w.status(); err != nil {
return nil, err
}
p := w.peers[r]
if w.networked && w.rank != 0 {
// One stream carries every rank's words to a rank that is not
// the hub, so one pending stash serves them all.
p = w.peers[0]
}
for i, m := range p.pending {
if m.from == r && m.tag == tag {
p.pending = append(p.pending[:i], p.pending[i+1:]...)
return m.data, nil
}
}
for {
var m message
var err error
switch {
case p.inbox != nil:
select {
case m = <-p.inbox:
case <-p.gone:
// The peer left: whatever it queued before leaving is
// still in the buffer and still counts; only an empty
// buffer means the frames will never come.
for {
select {
case m = <-p.inbox:
if m.from == r && m.tag == tag {
return m.data, nil
}
p.pending = append(p.pending, m)
continue
default:
}
return nil, w.fail(base.Errf("spmd: rank %d receiving from rank %d: rank %d left the world", w.rank, r, r))
}
case <-w.done:
return nil, w.status()
}
default:
// This rank drives its single connection; the hub has
// already routed whatever was not for it.
m, err = p.conn.readFrame(w.maxMessage, w.deadline())
if err != nil {
return nil, w.fail(base.Errf("spmd: rank %d receiving from rank %d: %v", w.rank, r, w.readErr(err)))
}
if m.dest != w.rank {
return nil, w.fail(base.Errf("spmd: rank %d got a frame addressed to rank %d", w.rank, m.dest))
}
}
if m.from == r && m.tag == tag {
return m.data, nil
}
p.pending = append(p.pending, m)
if len(p.pending) >= pendingCap {
return nil, w.fail(base.Errf("spmd: rank %d holds %d unmatched frames against rank %d", w.rank, len(p.pending), r))
}
}
}
// readErr names a networked read failure for what it is: a peer that
// closed or dropped its connection.
func (w *World) readErr(err error) error {
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) {
return errors.New("the peer closed its connection")
}
return err
}
// pump reads one networked hub link for the world's lifetime: frames
// for rank 0 join the link's inbox, frames for anybody else join that
// rank's outbox unchanged. The hub is a post office, never an
// interpreter: what a frame carries is the collectives' business. A
// frame bound for another rank reads into a buffer from the routed
// frame pool, whose ownership rides the message through the outbox
// channel to the link's one drain. A peer that closes its connection
// has left the world, which is its own business too; only a broken or
// unreadable link fails this world.
func (w *World) pump(p *peer) {
for {
m, err := p.conn.readFrameInto(w.maxMessage, w.deadline(), takeFrameBuffer)
if err != nil {
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) || errors.Is(err, net.ErrClosed) {
p.depart()
return
}
w.fail(base.Errf("spmd: rank 0 reading from rank %d: %v", p.rank, w.readErr(err)))
return
}
if m.from != p.rank || m.dest < 0 || m.dest >= w.size {
w.fail(base.Errf("spmd: rank 0 got a mislabelled frame from rank %d", p.rank))
return
}
var out chan message
if m.dest == 0 {
out = p.inbox
} else {
out = w.peers[m.dest].outbox
}
select {
case out <- m:
case <-w.done:
return
}
}
}
// drain writes one networked hub link for the world's lifetime: it is
// the only goroutine that ever writes the connection, so the frames
// the collective logic and the routed traffic send share one ordered
// stream without a lock. It is also the single consumer of the
// routed frames the pump queues, which makes it the one place a
// pooled payload is returned: writeFrame is the payload's last
// reader, so the buffer goes back the moment the write returns,
// whatever the answer was. When the world ends it writes out whatever
// the last collectives queued before it leaves, so an orderly close
// never drops a frame that was sent. A write against a link whose peer
// already left is not failed, because the departure was the peer's own
// clean act; any other write error is, named at once rather than left
// to surface later as a hang.
func (w *World) drain(p *peer) {
write := func(m message) bool {
err := p.conn.writeFrame(m, w.deadline())
if m.pooled {
routedFrames.retire(m.data)
}
if err == nil {
return true
}
if p.hasLeft() || w.ended() {
return false
}
w.fail(base.Errf("spmd: rank 0 writing to rank %d: %v", p.rank, err))
return false
}
for {
select {
case m := <-p.outbox:
if !write(m) {
return
}
case <-w.done:
for {
select {
case m := <-p.outbox:
if !write(m) {
return
}
default:
return
}
}
}
}
}
// ended reports whether the world's done channel has closed.
func (w *World) ended() bool {
select {
case <-w.done:
return true
default:
return false
}
}
// startPumps launches the hub's readers and writers, one of each per
// link. They live as long as the world does.
func (w *World) startPumps() {
for r := 1; r < w.size; r++ {
p := w.peers[r]
w.pumpWg.Go(func() { w.pump(p) })
w.drainWg.Go(func() { w.drain(p) })
}
}
// Launch runs the same function on size ranks of one process, one
// goroutine per rank, over in-process links: the same collectives, the
// same answers and the same rules as a networked world, which makes it
// the development and test surface of the package. The errors of the
// ranks that failed come back joined in rank order, so the report is
// deterministic and the rank that caused the trouble is in it; a rank
// whose function panics fails its world, and the panic comes back as
// that rank's error.
func Launch(size int, fn func(w *World) error) error {
if size < 1 {
return base.Errf("spmd: a world needs at least one rank, got %d", size)
}
worlds := make([]*World, size)
for r := range worlds {
worlds[r] = &World{
rank: r,
size: size,
done: make(chan struct{}),
peers: make([]*peer, size),
}
}
// One channel per direction of every pair: the sender's outbox is
// the receiver's inbox, so a frame crosses without a middleman and
// a receive wakes the moment its rank's world leaves.
for r := range worlds {
for q := r + 1; q < size; q++ {
rToQ := make(chan message, linkQueue)
qToR := make(chan message, linkQueue)
worlds[r].peers[q] = &peer{rank: q, inbox: qToR, outbox: rToQ, gone: worlds[q].done}
worlds[q].peers[r] = &peer{rank: r, inbox: rToQ, outbox: qToR, gone: worlds[r].done}
}
}
errs := make([]error, size)
var wg sync.WaitGroup
for r := range worlds {
wg.Go(func() {
defer func() {
if p := recover(); p != nil {
errs[r] = base.Errf("spmd: rank %d panicked: %v", r, p)
worlds[r].fail(errs[r])
}
worlds[r].leave()
}()
errs[r] = fn(worlds[r])
})
}
wg.Wait()
return errors.Join(errs...)
}
// Listen builds the rank 0 end of a networked world: it listens on the
// address until every other rank has joined, assigning ranks in dial
// order. Connections that do not say the spmd handshake are closed
// and skipped, so they cannot take a rank's place.
func Listen(addr string, size int, opts Options) (*World, error) {
ln, err := net.Listen("tcp", addr)
if err != nil {
return nil, base.Errf("spmd: listening on %s: %v", addr, err)
}
return listen(ln, size, opts)
}
// listen assembles the rank 0 world on a ready listener; the tests use
// it to hand over a listener whose address they already know.
func listen(ln net.Listener, size int, opts Options) (*World, error) {
if size < 1 {
ln.Close()
return nil, base.Errf("spmd: a world needs at least one rank, got %d", size)
}
timeout := opts.timeout()
w := &World{
rank: 0,
size: size,
timeout: timeout,
maxMessage: opts.maxMessage(),
networked: true,
done: make(chan struct{}),
peers: make([]*peer, size),
}
w.closers = append(w.closers, ln)
if tcp, ok := ln.(*net.TCPListener); ok && timeout > 0 {
tcp.SetDeadline(time.Now().Add(timeout))
}
deadline := w.deadline()
for joined := 1; joined < size; joined++ {
conn, err := ln.Accept()
if err != nil {
return nil, w.fail(base.Errf("spmd: rank 0 accepting rank %d on %s: %v", joined, ln.Addr(), err))
}
if err := readHello(conn, deadline); err != nil {
// A connection that does not say the handshake is not one
// of ours; it takes no rank's place.
conn.Close()
joined--
continue
}
if err := sendWelcome(conn, size, joined); err != nil {
conn.Close()
return nil, w.fail(base.Errf("spmd: rank 0 welcoming rank %d: %v", joined, err))
}
w.peers[joined] = &peer{
rank: joined,
conn: newTCPLink(conn),
inbox: make(chan message, linkQueue),
outbox: make(chan message, linkQueue),
gone: make(chan struct{}),
}
w.closers = append(w.closers, conn)
}
// Every rank has its link; nobody else joins this world. The
// listener was closers[0], and its job is done.
ln.Close()
w.closers = w.closers[1:]
w.startPumps()
return w, nil
}
// Join builds the other ranks' end of a networked world: it dials the
// listening rank 0, which answers with the world's size and this
// connection's rank. The world is a star over rank 0's listener, so
// one address is the whole world's knowledge; rank 0 routes every
// frame to its destination.
func Join(addr string, opts Options) (*World, error) {
timeout := opts.timeout()
d := net.Dialer{Timeout: timeout}
conn, err := d.Dial("tcp", addr)
if err != nil {
return nil, base.Errf("spmd: rank dialling %s: %v", addr, err)
}
w := &World{
timeout: timeout,
maxMessage: opts.maxMessage(),
networked: true,
done: make(chan struct{}),
closers: []io.Closer{conn},
}
if err := sendHello(conn); err != nil {
return nil, w.fail(base.Errf("spmd: rank saying hello to %s: %v", addr, err))
}
size, rank, err := readWelcome(conn, w.deadline())
if err != nil {
return nil, w.fail(base.Errf("spmd: rank reading %s's welcome: %v", addr, err))
}
if rank == 0 {
return nil, w.fail(base.Errf("spmd: the listener at %s answered a joining connection with rank 0", addr))
}
w.rank, w.size = rank, size
w.peers = make([]*peer, size)
w.peers[0] = &peer{rank: 0, conn: newTCPLink(conn)}
return w, nil
}
// Barrier blocks until every rank of the world has reached it. It
// carries no data and no arithmetic, so there is nothing in it to be
// anything but deterministic.
func (w *World) Barrier() error {
if err := w.status(); err != nil {
return err
}
if w.rank == 0 {
for r := 1; r < w.size; r++ {
if _, err := w.recvFrom(r, tagBarrier); err != nil {
return err
}
}
for r := 1; r < w.size; r++ {
if err := w.sendTo(r, tagBarrierAck, nil); err != nil {
return err
}
}
return nil
}
if err := w.sendTo(0, tagBarrier, nil); err != nil {
return err
}
_, err := w.recvFrom(0, tagBarrierAck)
return err
}