mirror of
https://github.com/amnezia-vpn/amneziawg-go.git
synced 2026-10-02 21:36:06 +03:00
fix: store prelude state on endpoints
This commit is contained in:
1 parent
756d3dd6bd
commit
948c555794
4 files changed
+188
-69
No files matched your search
+4
-44
@@ -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),
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
Reference in new issue
Block a user