Add EgressProvider
This commit is contained in:
parent
2c27bbf4f9
commit
6f5e8b1947
1 changed files with 49 additions and 7 deletions
|
|
@ -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)])
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue