fix: store prelude state on endpoints

This commit is contained in:
Frog Rocky committed 2026-06-01 16:30:56 +02:00
1 parent 756d3dd6bd
commit 948c555794
4 files changed
+188 -69

No files matched your search

+4 -44
View File
@@ -129,32 +129,12 @@ type PreludeUDPConn struct {
rulesArr [5]Rules
junkCount int
junkGen *junkGenerator
statesMu sync.Mutex
states map[string]*PreludeState
state PreludeState
resendInterval time.Duration
}
func (c *PreludeUDPConn) preludeState(addr net.Addr) *PreludeState {
key := ""
if addr != nil {
key = addr.String()
}
c.statesMu.Lock()
defer c.statesMu.Unlock()
if c.states == nil {
c.states = make(map[string]*PreludeState)
}
state := c.states[key]
if state == nil {
state = new(PreludeState)
c.states[key] = state
}
return state
}
func (c *PreludeUDPConn) WriteMsgUDP(b, oob []byte, addr *net.UDPAddr) (n, oobn int, err error) {
state := c.preludeState(addr)
state := &c.state
if state.ClaimSend(time.Now(), c.resendInterval) {
buf := c.pool.Get()
ctx := writeContext{
@@ -344,37 +324,17 @@ type PreludeBatchConn struct {
rulesArr [5]Rules
junkCount int
junkGen *junkGenerator
statesMu sync.Mutex
states map[string]*PreludeState
state PreludeState
resendInterval time.Duration
}
func (c *PreludeBatchConn) preludeState(addr net.Addr) *PreludeState {
key := ""
if addr != nil {
key = addr.String()
}
c.statesMu.Lock()
defer c.statesMu.Unlock()
if c.states == nil {
c.states = make(map[string]*PreludeState)
}
state := c.states[key]
if state == nil {
state = new(PreludeState)
c.states[key] = state
}
return state
}
func (c *PreludeBatchConn) WriteBatch(ms []ipv4.Message, flags int) (n int, err error) {
if len(ms) == 0 {
return c.BatchConn.WriteBatch(ms, flags)
}
preludeMsg := &ms[0]
state := c.preludeState(preludeMsg.Addr)
state := &c.state
if state.ClaimSend(time.Now(), c.resendInterval) {
ctx := writeContext{
FlexBuffer: WrapFlexBuffer(nil),
+53
View File
@@ -7,3 +7,56 @@ type PreludeEndpoint interface {
PreludeState() *conceal.PreludeState
ResetPreludeState()
}
type wrappedEndpoint interface {
Endpoint
UnwrapEndpoint() Endpoint
}
type preludeEndpoint struct {
Endpoint
prelude conceal.PreludeState
}
func wrapPreludeEndpoint(ep Endpoint) Endpoint {
if ep == nil {
return nil
}
if _, ok := ep.(PreludeEndpoint); ok {
return ep
}
if _, ok := ep.(wrappedEndpoint); ok {
return ep
}
return &preludeEndpoint{Endpoint: ep}
}
func unwrapEndpoint(ep Endpoint) Endpoint {
for {
wrapped, ok := ep.(wrappedEndpoint)
if !ok {
return ep
}
ep = wrapped.UnwrapEndpoint()
}
}
func (e *preludeEndpoint) UnwrapEndpoint() Endpoint {
return e.Endpoint
}
func (e *preludeEndpoint) ClearSrc() {
e.Endpoint.ClearSrc()
e.ResetPreludeState()
}
func (e *preludeEndpoint) PreludeState() *conceal.PreludeState {
return &e.prelude
}
func (e *preludeEndpoint) ResetPreludeState() {
e.prelude.Reset()
}
var _ PreludeEndpoint = (*preludeEndpoint)(nil)
var _ wrappedEndpoint = (*preludeEndpoint)(nil)
+13 -25
View File
@@ -32,7 +32,6 @@ type ConcealBind struct {
fallbackPort uint16
fallbackSessions map[string]*concealBindFallbackSession
preludeStates map[string]*conceal.PreludeState
pipeline atomic.Pointer[conceal.UDPDatagramPipeline]
}
@@ -90,6 +89,7 @@ func (b *ConcealBind) wrapReceiveFunc(fn ReceiveFunc) ReceiveFunc {
pipeline := b.currentPipeline()
if pipeline == nil || !pipeline.Active() {
wrapPreludeEndpoints(eps[:n])
return n, nil
}
@@ -115,11 +115,18 @@ func (b *ConcealBind) wrapReceiveFunc(fn ReceiveFunc) ReceiveFunc {
}
sizes[i] = size
eps[i] = wrapPreludeEndpoint(eps[i])
}
return n, nil
}
}
func wrapPreludeEndpoints(eps []Endpoint) {
for i, ep := range eps {
eps[i] = wrapPreludeEndpoint(ep)
}
}
func (b *ConcealBind) Close() error {
b.mu.Lock()
b.closeFallbackSessionsLocked()
@@ -138,19 +145,7 @@ func (b *ConcealBind) preludeState(ep Endpoint) *conceal.PreludeState {
if preludeEP, ok := ep.(PreludeEndpoint); ok {
return preludeEP.PreludeState()
}
key := ep.DstToString()
b.mu.Lock()
defer b.mu.Unlock()
if b.preludeStates == nil {
b.preludeStates = make(map[string]*conceal.PreludeState)
}
state := b.preludeStates[key]
if state == nil {
state = new(conceal.PreludeState)
b.preludeStates[key] = state
}
return state
return nil
}
func (b *ConcealBind) resetPreludeState(ep Endpoint) {
@@ -159,15 +154,7 @@ func (b *ConcealBind) resetPreludeState(ep Endpoint) {
}
if preludeEP, ok := ep.(PreludeEndpoint); ok {
preludeEP.ResetPreludeState()
return
}
key := ep.DstToString()
b.mu.Lock()
if state := b.preludeStates[key]; state != nil {
state.Reset()
}
b.mu.Unlock()
}
func (b *ConcealBind) Send(bufs [][]byte, ep Endpoint) error {
@@ -177,7 +164,7 @@ func (b *ConcealBind) Send(bufs [][]byte, ep Endpoint) error {
pipeline := b.currentPipeline()
if pipeline == nil || !pipeline.Active() {
return b.inner.Send(bufs, ep)
return b.inner.Send(bufs, unwrapEndpoint(ep))
}
batchSize := b.inner.BatchSize()
@@ -210,7 +197,7 @@ func (b *ConcealBind) Send(bufs [][]byte, ep Endpoint) error {
if len(batch) == 0 {
return nil
}
err := b.inner.Send(batch, ep)
err := b.inner.Send(batch, unwrapEndpoint(ep))
putRetained()
clearBatch()
return err
@@ -267,6 +254,7 @@ func (b *ConcealBind) ParseEndpoint(s string) (Endpoint, error) {
if err != nil {
return nil, err
}
ep = wrapPreludeEndpoint(ep)
b.resetPreludeState(ep)
return ep, nil
}
@@ -411,7 +399,7 @@ func (s *concealBindFallbackSession) relay() {
if err != nil {
return
}
if err := s.parent.inner.Send([][]byte{buf[:n]}, s.ep); err != nil {
if err := s.parent.inner.Send([][]byte{buf[:n]}, unwrapEndpoint(s.ep)); err != nil {
return
}
}
+118
View File
@@ -92,6 +92,29 @@ func TestConcealBindNoOpWithoutOpts(t *testing.T) {
}
}
func TestConcealBindNoOpUnwrapsNonPreludeEndpoint(t *testing.T) {
inner := &rawPacketBind{
fakePacketBind: fakePacketBind{batchSize: 3},
}
bind := NewConcealBind(inner)
endpoint, err := bind.ParseEndpoint("127.0.0.1:51820")
if err != nil {
t.Fatalf("parse endpoint: %v", err)
}
if _, ok := endpoint.(PreludeEndpoint); !ok {
t.Fatalf("parsed endpoint does not carry prelude state")
}
transport := makeTransportPacket()
if err := bind.Send([][]byte{transport}, endpoint); err != nil {
t.Fatalf("send: %v", err)
}
if inner.badEndpoint {
t.Fatalf("inner bind received wrapped endpoint, want raw endpoint")
}
}
func TestConcealBindSendBatchesActivePipeline(t *testing.T) {
inner := &fakePacketBind{batchSize: 3}
bind := NewConcealBind(inner)
@@ -268,6 +291,52 @@ func TestConcealBindPreludePositiveResendIntervalForcesResend(t *testing.T) {
}
}
func TestConcealBindWrapsNonPreludeEndpointAndUnwrapsForInnerBind(t *testing.T) {
inner := &rawPacketBind{
fakePacketBind: fakePacketBind{batchSize: 8},
}
bind := NewConcealBind(inner)
bind.SetFramedOpts(conceal.FramedOpts{H4: mustHeader(t, "779")})
bind.SetPreludeOpts(conceal.PreludeOpts{
ResendInterval: 0,
RulesArr: [5]conceal.Rules{
mustParseRules(t, "<b 0xaabb>"),
},
})
endpoint, err := bind.ParseEndpoint("127.0.0.1:51820")
if err != nil {
t.Fatalf("parse endpoint: %v", err)
}
if _, ok := endpoint.(PreludeEndpoint); !ok {
t.Fatalf("parsed endpoint does not carry prelude state")
}
transport := makeTransportPacket()
if err := bind.Send([][]byte{transport}, endpoint); err != nil {
t.Fatalf("first send: %v", err)
}
if err := bind.Send([][]byte{transport}, endpoint); err != nil {
t.Fatalf("second send: %v", err)
}
if inner.badEndpoint {
t.Fatalf("inner bind received wrapped endpoint, want raw endpoint")
}
wirePackets := flattenSendCalls(inner.sendCalls)
if len(wirePackets) != 3 {
t.Fatalf("wire packet count = %d, want 3", len(wirePackets))
}
if !bytes.Equal(wirePackets[0], []byte{0xaa, 0xbb}) {
t.Fatalf("wire prelude = %x, want aabb", wirePackets[0])
}
for i, call := range inner.sendCalls {
if _, ok := call.endpoint.(*rawPacketEndpoint); !ok {
t.Fatalf("send call %d endpoint = %T, want *rawPacketEndpoint", i, call.endpoint)
}
}
}
func TestConcealBindSendAndReceive(t *testing.T) {
senderInner := &fakePacketBind{batchSize: 4}
sender := newTestConcealBind(t, senderInner)
@@ -665,6 +734,53 @@ func (e *fakePacketEndpoint) SrcIP() netip.Addr {
return netip.Addr{}
}
type rawPacketBind struct {
fakePacketBind
badEndpoint bool
}
func (b *rawPacketBind) Send(bufs [][]byte, ep Endpoint) error {
if _, ok := ep.(*rawPacketEndpoint); !ok {
b.badEndpoint = true
}
return b.fakePacketBind.Send(bufs, ep)
}
func (b *rawPacketBind) ParseEndpoint(s string) (Endpoint, error) {
addr, err := netip.ParseAddrPort(s)
if err != nil {
return nil, err
}
return &rawPacketEndpoint{addr: addr}, nil
}
type rawPacketEndpoint struct {
addr netip.AddrPort
}
func (e *rawPacketEndpoint) ClearSrc() {}
func (e *rawPacketEndpoint) SrcToString() string {
return ""
}
func (e *rawPacketEndpoint) DstToString() string {
return e.addr.String()
}
func (e *rawPacketEndpoint) DstToBytes() []byte {
out, _ := e.addr.MarshalBinary()
return out
}
func (e *rawPacketEndpoint) DstIP() netip.Addr {
return e.addr.Addr()
}
func (e *rawPacketEndpoint) SrcIP() netip.Addr {
return netip.Addr{}
}
func flattenSendCalls(calls []fakeSendCall) [][]byte {
var out [][]byte
for _, call := range calls {
@@ -691,5 +807,7 @@ func wireHeader(t *testing.T, packet []byte) uint32 {
}
var _ Bind = (*fakePacketBind)(nil)
var _ Bind = (*rawPacketBind)(nil)
var _ Endpoint = (*fakePacketEndpoint)(nil)
var _ Endpoint = (*rawPacketEndpoint)(nil)
var _ PreludeEndpoint = (*fakePacketEndpoint)(nil)