577 lines
17 KiB
Go
577 lines
17 KiB
Go
// 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
|
|
}
|