diff --git a/conn/bind_std.go b/conn/bind_std.go index c437da3..0a15de0 100644 --- a/conn/bind_std.go +++ b/conn/bind_std.go @@ -23,6 +23,12 @@ import ( "golang.org/x/net/ipv6" ) +type EgressProvider interface { + SetEgressPort(port uint16) bool + LookupEgress(destination netip.AddrPort) *net.UDPConn + ReceiveEgress(buffer []byte) (int, netip.AddrPort, error) +} + var _ Bind = (*StdNetBind)(nil) // StdNetBind implements Bind for all platforms. While Windows has its own Bind @@ -32,6 +38,7 @@ var _ Bind = (*StdNetBind)(nil) // proposal in https://github.com/golang/go/issues/45886#issuecomment-1218301564. type StdNetBind struct { externalControl control.Func + egressProvider EgressProvider reservedForEndpoint map[netip.AddrPort][3]uint8 mu sync.Mutex // protects all fields except as specified @@ -257,10 +264,29 @@ again: if len(fns) == 0 { return nil, 0, syscall.EAFNOSUPPORT } + if s.egressProvider != nil { + s.egressProvider.SetEgressPort(uint16(port)) + fns = append(fns, func(bufs [][]byte, sizes []int, endpoints []Endpoint) (int, error) { + dataLength, source, err := s.egressProvider.ReceiveEgress(bufs[0]) + if err != nil { + return 0, err + } + sizes[0] = dataLength + if dataLength > 3 { + common.ClearArray(bufs[0][1:4]) + } + endpoints[0] = &StdNetEndpoint{AddrPort: source} + return 1, nil + }) + } return fns, uint16(port), nil } +func (s *StdNetBind) SetEgressProvider(provider EgressProvider) { + s.egressProvider = provider +} + func (s *StdNetBind) putMessages(msgs *[]ipv6.Message) { for i := range *msgs { buffers := (*msgs)[i].Buffers @@ -371,6 +397,9 @@ func (s *StdNetBind) Close() error { s.mu.Lock() defer s.mu.Unlock() + if s.egressProvider != nil { + s.egressProvider.SetEgressPort(0) + } var err1, err2 error if s.ipv4 != nil { err1 = s.ipv4.Close() @@ -415,13 +444,14 @@ func (s *StdNetBind) Send(bufs [][]byte, endpoint Endpoint, offset int) error { } bufs = bufs[IdealBatchSize:] } + standardEndpoint := endpoint.(*StdNetEndpoint) s.mu.Lock() blackhole := s.blackhole4 conn := s.ipv4 offload := s.ipv4TxOffload br := batchWriter(s.ipv4PC) is6 := false - if endpoint.DstIP().Is6() { + if standardEndpoint.DstIP().Is6() { blackhole = s.blackhole6 conn = s.ipv6 br = s.ipv6PC @@ -442,30 +472,42 @@ func (s *StdNetBind) Send(bufs [][]byte, endpoint Endpoint, offset int) error { ua := s.udpAddrPool.Get().(*net.UDPAddr) defer s.udpAddrPool.Put(ua) if is6 { - as16 := endpoint.DstIP().As16() + as16 := standardEndpoint.DstIP().As16() copy(ua.IP, as16[:]) ua.IP = ua.IP[:16] } else { - as4 := endpoint.DstIP().As4() + as4 := standardEndpoint.DstIP().As4() copy(ua.IP, as4[:]) ua.IP = ua.IP[:4] } - ua.Port = int(endpoint.(*StdNetEndpoint).Port()) + ua.Port = int(standardEndpoint.Port()) var ( retried bool err error ) for _, buf := range bufs { if len(buf) > offset+3 { - reserved, loaded := s.reservedForEndpoint[endpoint.(*StdNetEndpoint).AddrPort] + reserved, loaded := s.reservedForEndpoint[standardEndpoint.AddrPort] if loaded { copy(buf[offset+1:offset+4], reserved[:]) } } } + if s.egressProvider != nil { + memberConn := s.egressProvider.LookupEgress(standardEndpoint.AddrPort) + if memberConn != nil { + for _, buf := range bufs { + _, err = memberConn.WriteToUDPAddrPort(buf[offset:], standardEndpoint.AddrPort) + if err != nil { + return err + } + } + return nil + } + } retry: if offload { - n := coalesceMessages(ua, endpoint.(*StdNetEndpoint), bufs, offset, *msgs, setGSOSize) + n := coalesceMessages(ua, standardEndpoint, bufs, offset, *msgs, setGSOSize) err = s.send(conn, br, (*msgs)[:n]) if err != nil && offload && errShouldDisableUDPGSO(err) { offload = false @@ -483,7 +525,7 @@ retry: for i := range bufs { (*msgs)[i].Addr = ua (*msgs)[i].Buffers[0] = bufs[i][offset:] - setSrcControl(&(*msgs)[i].OOB, endpoint.(*StdNetEndpoint)) + setSrcControl(&(*msgs)[i].OOB, standardEndpoint) } err = s.send(conn, br, (*msgs)[:len(bufs)]) }