Add EgressProvider

This commit is contained in:
世界 2026-07-17 10:40:26 +08:00
parent 2c27bbf4f9
commit 6f5e8b1947
No known key found for this signature in database
GPG key ID: CD109927C34A63C4

View file

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