// Copyright (c) 2026 Petr BalvĂ­n (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 }