feat: add conceal fallback port proxying

Add fallback_port support so invalid TCP/UDP conceal traffic is proxied to a local fallback service instead of being dropped.

Also add typed conceal format errors, preserve consumed bytes for TCP fallback replay, wire fallback config through UAPI/binds/Outline, and cover TCP, UDP, and UAPI behavior with tests.
This commit is contained in:
Frog Rocky authored and Yaroslav Gurov committed 2026-05-11 16:59:04 +02:00
1 parent 8d9bf0ed57
commit f8b2d2ba5a
17 files changed
+881 -74

No files matched your search

+7
View File
@@ -117,6 +117,13 @@ Value as a tag-sequence, as described in Tag format paragraph
> [!NOTE]
> Although it is not required, `Network = tcp` would highly-likely require `FormatIn` and `FormatOut` parameters since stream connection does not have fixed-sized messages
### Fallback
**`FallbackPort: int`**
- A local fallback service port in range `1-65535`.
- If configured and incoming data does not match the configured AWG/conceal formats, AWG proxies that TCP stream or UDP sender to a loopback service on `FallbackPort` and relays the response back to the original sender.
- Absence means fallback is disabled.
### Message format
**`FormatIn: string of tags`**
+22 -13
View File
@@ -241,30 +241,30 @@ func (e *frameEncoding) IsInitiationRecord(b []byte) bool {
return e.recordKind(b) == frameRecordInitiation
}
func (e *frameEncoding) Decode(b []byte) int {
func (e *frameEncoding) Decode(b []byte) (int, error) {
switch e.recordKind(b) {
case frameRecordInitiation:
if e.compat {
return decodeOneCompat(b, e.header.initial, e.padding.initial)
return decodeOneCompat(b, e.header.initial, e.padding.initial), nil
}
return decodeOne(b, e.header.initial, e.padding.initial, WireguardMsgInitiationType)
return decodeOne(b, e.header.initial, e.padding.initial, WireguardMsgInitiationType), nil
case frameRecordResponse:
if e.compat {
return decodeOneCompat(b, e.header.response, e.padding.response)
return decodeOneCompat(b, e.header.response, e.padding.response), nil
}
return decodeOne(b, e.header.response, e.padding.response, WireguardMsgResponseType)
return decodeOne(b, e.header.response, e.padding.response, WireguardMsgResponseType), nil
case frameRecordCookie:
if e.compat {
return decodeOneCompat(b, e.header.cookie, e.padding.cookie)
return decodeOneCompat(b, e.header.cookie, e.padding.cookie), nil
}
return decodeOne(b, e.header.cookie, e.padding.cookie, WireguardMsgCookieReplyType)
return decodeOne(b, e.header.cookie, e.padding.cookie, WireguardMsgCookieReplyType), nil
case frameRecordTransport:
if e.compat {
return decodeOneCompat(b, e.header.transport, e.padding.transport)
return decodeOneCompat(b, e.header.transport, e.padding.transport), nil
}
return decodeOne(b, e.header.transport, e.padding.transport, WireguardMsgTransportType)
return decodeOne(b, e.header.transport, e.padding.transport, WireguardMsgTransportType), nil
default:
return len(b)
return 0, NewFormatError(b, errInvalidData)
}
}
@@ -289,7 +289,13 @@ type FramedConn struct {
func (c *FramedConn) Read(b []byte) (n int, err error) {
n, err = c.Conn.Read(b)
n = c.enc.Decode(b[:n])
if n > 0 {
var decodeErr error
n, decodeErr = c.enc.Decode(b[:n])
if decodeErr != nil {
return 0, decodeErr
}
}
return n, err
}
@@ -328,7 +334,7 @@ func (c *FramedUDPConn) ReadMsgUDP(b, oob []byte) (n, oobn, flags int, addr *net
if err != nil {
return 0, 0, 0, nil, err
}
n = c.enc.Decode(b[:n])
n, err = c.enc.Decode(b[:n])
return n, oobn, flags, addr, err
}
@@ -370,7 +376,10 @@ func (c *FramedBatchConn) ReadBatch(ms []ipv4.Message, flags int) (n int, err er
for i := range ms[:n] {
b := ms[i].Buffers[0][:ms[i].N]
ms[i].N = c.enc.Decode(b)
ms[i].N, err = c.enc.Decode(b)
if err != nil {
return 0, err
}
}
return n, nil
+13
View File
@@ -174,9 +174,22 @@ func (c *PreludeConn) Read(b []byte) (n int, err error) {
c.seenValid = true
return n, nil
}
if c.isPreludeRecord(b[:n]) {
continue
}
return 0, NewFormatError(b[:n], errInvalidData)
}
}
func (c *PreludeConn) isPreludeRecord(b []byte) bool {
for _, rules := range c.rulesArr {
if rules.Match(b, c.pool) {
return true
}
}
return false
}
func (c *PreludeConn) Write(b []byte) (n int, err error) {
if c.recordEncoding.IsInitiationRecord(b) {
if err := c.writePreludeRecords(); err != nil {
+45
View File
@@ -0,0 +1,45 @@
package conceal
import (
"bytes"
"errors"
)
var ErrFormat = errors.New("conceal format error")
type FormatError struct {
Data []byte
Err error
}
func NewFormatError(data []byte, err error) *FormatError {
return &FormatError{
Data: bytes.Clone(data),
Err: err,
}
}
func (e *FormatError) Error() string {
if e == nil {
return ErrFormat.Error()
}
if e.Err == nil {
return ErrFormat.Error()
}
return ErrFormat.Error() + ": " + e.Err.Error()
}
func (e *FormatError) Unwrap() []error {
if e == nil || e.Err == nil {
return []error{ErrFormat}
}
return []error{ErrFormat, e.Err}
}
func FormatErrorData(err error) []byte {
var formatErr *FormatError
if errors.As(err, &formatErr) {
return bytes.Clone(formatErr.Data)
}
return nil
}
+44
View File
@@ -0,0 +1,44 @@
package conceal
import (
"bytes"
"errors"
"sync"
"testing"
)
func TestMasqueradeConnReadRecordReturnsFormatErrorWithData(t *testing.T) {
rules, err := ParseRules("<b 0xfeed><dz be 2><d>")
if err != nil {
t.Fatalf("parse rules: %v", err)
}
pool := sync.Pool{
New: func() any {
return make([]byte, 65535)
},
}
conn, ok := NewMasqueradeConn(newBenchmarkStreamConn([]byte("GET /")), &pool, MasqueradeOpts{
RulesIn: rules,
})
if !ok {
t.Fatal("expected masquerade connection")
}
buf := make([]byte, 64)
n, err := conn.ReadRecord(buf)
if !errors.Is(err, ErrFormat) {
t.Fatalf("ReadRecord error = %v, want ErrFormat", err)
}
if n != 0 {
t.Fatalf("ReadRecord n = %d, want 0", n)
}
var formatErr *FormatError
if !errors.As(err, &formatErr) {
t.Fatalf("ReadRecord error type = %T, want *FormatError", err)
}
if !bytes.Equal(formatErr.Data, []byte("GE")) {
t.Fatalf("format error data = %q, want %q", formatErr.Data, "GE")
}
}
+77 -12
View File
@@ -22,6 +22,11 @@ type readContext struct {
FlexBuffer
*BufferPool
nextDataSize int
formatData []byte
}
func (ctx *readContext) rememberRead(b []byte) {
ctx.formatData = append(ctx.formatData, b...)
}
type writeContext struct {
@@ -55,14 +60,39 @@ func (r Rules) Write(w io.Writer, ctx *writeContext) error {
}
func (r Rules) Read(rd io.Reader, ctx *readContext) error {
formatDataStart := len(ctx.formatData)
for _, rule := range r {
if err := rule.Read(rd, ctx); err != nil {
if errors.Is(err, ErrFormat) {
return err
}
if errors.Is(err, errInvalidData) {
return NewFormatError(ctx.formatData[formatDataStart:], err)
}
return err
}
}
return nil
}
func (r Rules) Match(b []byte, pool *BufferPool) bool {
if r == nil {
return false
}
tmp := pool.Get()
defer pool.Put(tmp)
reader := newSliceReader(b)
ctx := readContext{
FlexBuffer: WrapFlexBuffer(tmp),
BufferPool: pool,
}
if err := r.Read(&reader, &ctx); err != nil {
return false
}
return len(reader.buf) == 0
}
func buildBytesRule(val string) (Rule, error) {
val = strings.TrimPrefix(val, "0x")
@@ -100,7 +130,9 @@ func (r *bytesRule) Read(rd io.Reader, ctx *readContext) error {
defer ctx.Put(tmp)
buf := tmp[:len(r.data)]
if _, err := io.ReadFull(rd, buf); err != nil {
n, err := io.ReadFull(rd, buf)
ctx.rememberRead(buf[:n])
if err != nil {
return err
}
@@ -144,7 +176,9 @@ func (r *randRule) Read(rd io.Reader, ctx *readContext) error {
defer ctx.Put(tmp)
buf := tmp[:r.length]
if _, err := io.ReadFull(rd, buf); err != nil {
n, err := io.ReadFull(rd, buf)
ctx.rememberRead(buf[:n])
if err != nil {
return err
}
@@ -191,7 +225,9 @@ func (r *randDigitRule) Read(rd io.Reader, ctx *readContext) error {
defer ctx.Put(tmp)
buf := tmp[:r.length]
if _, err := io.ReadFull(rd, buf); err != nil {
n, err := io.ReadFull(rd, buf)
ctx.rememberRead(buf[:n])
if err != nil {
return err
}
@@ -242,7 +278,9 @@ func (r *randCharRule) Read(rd io.Reader, ctx *readContext) error {
defer ctx.Put(tmp)
buf := tmp[:r.length]
if _, err := io.ReadFull(rd, buf); err != nil {
n, err := io.ReadFull(rd, buf)
ctx.rememberRead(buf[:n])
if err != nil {
return err
}
@@ -271,10 +309,16 @@ func (r *timestampRule) Write(w io.Writer, ctx *writeContext) error {
}
func (r *timestampRule) Read(rd io.Reader, ctx *readContext) error {
var timestamp uint32
if err := binary.Read(rd, binary.BigEndian, &timestamp); err != nil {
tmp := ctx.Get()
defer ctx.Put(tmp)
buf := tmp[:4]
n, err := io.ReadFull(rd, buf)
ctx.rememberRead(buf[:n])
if err != nil {
return err
}
_ = binary.BigEndian.Uint32(buf)
// TODO: check timestamp?
@@ -389,7 +433,12 @@ func (r *dataSizeRule) Read(rd io.Reader, ctx *readContext) error {
switch r.format {
case NumFormatBE:
buf := tmp[:r.length]
if _, err := io.ReadFull(rd, buf); err != nil {
n, err := io.ReadFull(rd, buf)
ctx.rememberRead(buf[:n])
if err != nil {
if errors.Is(err, io.ErrShortBuffer) {
return errInvalidData
}
return err
}
var size int
@@ -401,7 +450,12 @@ func (r *dataSizeRule) Read(rd io.Reader, ctx *readContext) error {
case NumFormatLE:
buf := tmp[:r.length]
if _, err := io.ReadFull(rd, buf); err != nil {
n, err := io.ReadFull(rd, buf)
ctx.rememberRead(buf[:n])
if err != nil {
if errors.Is(err, io.ErrShortBuffer) {
return errInvalidData
}
return err
}
var size int
@@ -413,25 +467,35 @@ func (r *dataSizeRule) Read(rd io.Reader, ctx *readContext) error {
case NumFormatAscii:
n, err := ReadUntil(rd, tmp, r.end)
if err == nil {
ctx.rememberRead(tmp[:n+1])
} else {
ctx.rememberRead(tmp[:n])
}
if err != nil {
return err
}
size64, err := strconv.ParseInt(string(tmp[:n]), 10, 32)
if err != nil {
return err
return fmt.Errorf("%w: %v", errInvalidData, err)
}
ctx.nextDataSize = int(size64)
case NumFormatHex:
n, err := ReadUntil(rd, tmp, r.end)
if err == nil {
ctx.rememberRead(tmp[:n+1])
} else {
ctx.rememberRead(tmp[:n])
}
if err != nil {
return err
}
size64, err := strconv.ParseInt(string(tmp[:n]), 16, 32)
if err != nil {
return err
return fmt.Errorf("%w: %v", errInvalidData, err)
}
ctx.nextDataSize = int(size64)
}
@@ -462,9 +526,10 @@ func (r *dataRule) Write(w io.Writer, ctx *writeContext) error {
func (r *dataRule) Read(rd io.Reader, ctx *readContext) error {
buf := ctx.PushTail(ctx.nextDataSize)
if buf == nil {
return io.ErrShortBuffer
return errInvalidData
}
_, err := io.ReadFull(rd, buf)
n, err := io.ReadFull(rd, buf)
ctx.rememberRead(buf[:n])
return err
}
+5 -1
View File
@@ -94,7 +94,11 @@ func (p *UDPDatagramPipeline) DecodeInPlace(buf []byte, n int) (int, bool) {
}
if p.framingActive {
n = p.framing.Decode(buf[:n])
var err error
n, err = p.framing.Decode(buf[:n])
if err != nil {
return 0, false
}
}
if !p.classifier.IsValid(buf[:n]) {
+8
View File
@@ -26,6 +26,7 @@ var (
_ Framable = (*StdNetBind)(nil)
_ Preludable = (*StdNetBind)(nil)
_ Masqueradable = (*StdNetBind)(nil)
_ Fallbackable = (*StdNetBind)(nil)
)
// StdNetBind implements Bind for all platforms. While Windows has its own Bind
@@ -55,6 +56,7 @@ type StdNetBind struct {
framedOpts conceal.FramedOpts
preludeOpts conceal.PreludeOpts
masqueradeOpts conceal.MasqueradeOpts
fallbackPort uint16
}
func NewStdNetBind() Bind {
@@ -326,11 +328,17 @@ func (s *StdNetBind) Close() error {
var err1, err2 error
if s.ipv4 != nil {
if closer, ok := s.ipv4PC.(fallbackSessionCloser); ok {
closer.closeFallbackSessions()
}
err1 = s.ipv4.Close()
s.ipv4 = nil
s.ipv4PC = nil
}
if s.ipv6 != nil {
if closer, ok := s.ipv6PC.(fallbackSessionCloser); ok {
closer.closeFallbackSessions()
}
err2 = s.ipv6.Close()
s.ipv6 = nil
s.ipv6PC = nil
+24 -9
View File
@@ -2,6 +2,7 @@ package conn
import (
"context"
"errors"
"fmt"
"io"
"net"
@@ -17,6 +18,7 @@ var (
_ Framable = (*BindStream)(nil)
_ Preludable = (*BindStream)(nil)
_ Masqueradable = (*BindStream)(nil)
_ Fallbackable = (*BindStream)(nil)
)
type streamPacketQueue struct {
@@ -54,6 +56,7 @@ type BindStream struct {
framedOpts conceal.FramedOpts
preludeOpts conceal.PreludeOpts
masqueradeOpts conceal.MasqueradeOpts
fallbackPort uint16
}
func (b *BindStream) readFaucet() ReceiveFunc {
@@ -81,6 +84,11 @@ func (b *BindStream) readStream(ep *streamEndpoint) {
sp := b.streamPacketPool.Get().(*streamPacketQueue)
n, err := ep.conn.Read(sp.buf[:])
if err != nil {
b.streamPacketPool.Put(sp)
if b.fallbackPort != 0 && errors.Is(err, conceal.ErrFormat) {
b.proxyStreamFallback(ep, conceal.FormatErrorData(err))
return
}
ep.Close()
return
}
@@ -100,19 +108,19 @@ func (b *BindStream) accept(listener net.Listener) {
default:
}
conn, err := listener.Accept()
rawConn, err := listener.Accept()
if err != nil {
// log this error somewhere
break
}
conn = b.upgradeConn(conn)
conn := b.upgradeConn(rawConn)
b.wg.Add(1)
go b.handleAccepted(conn)
go b.handleAccepted(rawConn, conn)
}
}
func (b *BindStream) handleAccepted(conn net.Conn) {
func (b *BindStream) handleAccepted(rawConn, conn net.Conn) {
defer b.wg.Done()
ap, err := netip.ParseAddrPort(conn.RemoteAddr().String())
@@ -121,8 +129,9 @@ func (b *BindStream) handleAccepted(conn net.Conn) {
}
ep := &streamEndpoint{
conn: conn,
dst: ap,
conn: conn,
rawConn: rawConn,
dst: ap,
}
b.wg.Add(1)
@@ -140,13 +149,14 @@ func (b *BindStream) dial(ep *streamEndpoint) error {
return nil
}
conn, err := b.dialer.DialContext(b.ctx, "tcp", ep.DstToString())
rawConn, err := b.dialer.DialContext(b.ctx, "tcp", ep.DstToString())
if err != nil {
return fmt.Errorf("failed to dial context: %v", err)
}
conn = b.upgradeConn(conn)
conn := b.upgradeConn(rawConn)
ep.conn = conn
ep.rawConn = rawConn
b.wg.Add(1)
go func() {
@@ -256,7 +266,8 @@ func (b *BindStream) BatchSize() int {
var _ Endpoint = (*streamEndpoint)(nil)
type streamEndpoint struct {
conn net.Conn
conn net.Conn
rawConn net.Conn
dst netip.AddrPort
mutex sync.Mutex
@@ -270,6 +281,10 @@ func (e *streamEndpoint) Close() {
e.conn.Close()
e.conn = nil
}
if e.rawConn != nil {
e.rawConn.Close()
e.rawConn = nil
}
}
func (e *streamEndpoint) DstToString() string {
+18
View File
@@ -18,6 +18,10 @@ type Masqueradable interface {
SetMasqueradeOpts(opts conceal.MasqueradeOpts)
}
type Fallbackable interface {
SetFallbackPort(port uint16)
}
type concealStage string
const (
@@ -117,6 +121,9 @@ func (b *StdNetBind) upgradeUDPConn(conn UDPConn) UDPConn {
}
}
}
if b.fallbackPort != 0 {
conn = newFallbackUDPConn(conn, origin, b.fallbackPort)
}
return conn
}
@@ -138,6 +145,9 @@ func (b *StdNetBind) upgradePacketConn(conn LinuxPacketConn) LinuxPacketConn {
}
}
}
if b.fallbackPort != 0 {
conn = newFallbackBatchConn(conn, origin, b.fallbackPort)
}
return conn
}
@@ -153,6 +163,10 @@ func (b *StdNetBind) SetMasqueradeOpts(opts conceal.MasqueradeOpts) {
b.masqueradeOpts = opts
}
func (b *StdNetBind) SetFallbackPort(port uint16) {
b.fallbackPort = port
}
func (b *BindStream) upgradeConn(conn net.Conn) net.Conn {
var recordConn conceal.StreamRecordConn
for _, stage := range b.streamConcealPipeline().stages {
@@ -186,3 +200,7 @@ func (b *BindStream) SetPreludeOpts(opts conceal.PreludeOpts) {
func (b *BindStream) SetMasqueradeOpts(opts conceal.MasqueradeOpts) {
b.masqueradeOpts = opts
}
func (b *BindStream) SetFallbackPort(port uint16) {
b.fallbackPort = port
}
+6 -18
View File
@@ -3,6 +3,7 @@ package conn
import (
"bytes"
"encoding/binary"
"errors"
"net"
"slices"
"sync"
@@ -228,7 +229,7 @@ func TestBindStreamTCPPreludeDropsInjectedDecoysOnRead(t *testing.T) {
}
}
func TestBindStreamTCPPreludeDropsLeadingInvalidRecords(t *testing.T) {
func TestBindStreamTCPPreludeRejectsLeadingUnknownRecords(t *testing.T) {
senderRaw, receiverRaw := net.Pipe()
defer senderRaw.Close()
defer receiverRaw.Close()
@@ -239,29 +240,16 @@ func TestBindStreamTCPPreludeDropsLeadingInvalidRecords(t *testing.T) {
recordWriter := mustRecordConn(t, senderRaw, &bind.bufferPool, conceal.MasqueradeOpts{
RulesOut: mustParseRules(t, "<dz be 2><d>"),
})
framedWriter, ok := conceal.NewFramedConn(recordWriter, &bind.bufferPool, bind.framedOpts)
if !ok {
t.Fatal("expected framed writer")
}
initiation := makeInitiationPacket()
writeErr := make(chan error, 1)
go func() {
if _, err := recordWriter.WriteRecord([]byte{0xde, 0xad}); err != nil {
writeErr <- err
return
}
if _, err := recordWriter.WriteRecord([]byte{0xbe, 0xef, 0x01}); err != nil {
writeErr <- err
return
}
_, err := framedWriter.Write(initiation)
_, err := recordWriter.WriteRecord([]byte{0xde, 0xad})
writeErr <- err
}()
got := readPacket(t, receiver, len(initiation))
if !bytes.Equal(got, initiation) {
t.Fatalf("read-side did not recover after invalid leading records")
buf := make([]byte, 16)
if _, err := receiver.Read(buf); !errors.Is(err, conceal.ErrFormat) {
t.Fatalf("receiver error = %v, want ErrFormat", err)
}
if err := <-writeErr; err != nil {
+328
View File
@@ -0,0 +1,328 @@
package conn
import (
"errors"
"io"
"net"
"strconv"
"sync"
"github.com/amnezia-vpn/amneziawg-go/conceal"
"golang.org/x/net/ipv4"
)
func fallbackTCPAddress(addr net.Addr, port uint16) string {
host := "127.0.0.1"
if tcpAddr, ok := addr.(*net.TCPAddr); ok && len(tcpAddr.IP) > 0 && tcpAddr.IP.To4() == nil {
host = "::1"
}
return net.JoinHostPort(host, strconv.Itoa(int(port)))
}
func fallbackUDPAddress(addr net.Addr, port uint16) *net.UDPAddr {
ip := net.IPv4(127, 0, 0, 1)
if udpAddr, ok := addr.(*net.UDPAddr); ok && len(udpAddr.IP) > 0 && udpAddr.IP.To4() == nil {
ip = net.IPv6loopback
}
return &net.UDPAddr{IP: ip, Port: int(port)}
}
func (b *BindStream) proxyStreamFallback(ep *streamEndpoint, first []byte) {
rawConn := ep.rawConn
if rawConn == nil {
rawConn = ep.conn
}
if rawConn == nil {
return
}
fallbackConn, err := b.dialer.DialContext(b.ctx, "tcp", fallbackTCPAddress(rawConn.RemoteAddr(), b.fallbackPort))
if err != nil {
ep.Close()
return
}
if len(first) > 0 {
if _, err := fallbackConn.Write(first); err != nil {
fallbackConn.Close()
ep.Close()
return
}
}
b.wg.Add(1)
go func() {
defer b.wg.Done()
_, _ = io.Copy(rawConn, fallbackConn)
rawConn.Close()
fallbackConn.Close()
}()
_, _ = io.Copy(fallbackConn, rawConn)
rawConn.Close()
fallbackConn.Close()
}
type fallbackUDPConn struct {
UDPConn
origin UDPConn
port uint16
mu sync.Mutex
sessions map[string]*fallbackUDPSession
}
func newFallbackUDPConn(conn, origin UDPConn, port uint16) UDPConn {
return &fallbackUDPConn{
UDPConn: conn,
origin: origin,
port: port,
sessions: make(map[string]*fallbackUDPSession),
}
}
func (c *fallbackUDPConn) ReadMsgUDP(b, oob []byte) (n, oobn, flags int, addr *net.UDPAddr, err error) {
for {
n, oobn, flags, addr, err = c.UDPConn.ReadMsgUDP(b, oob)
if err == nil {
return n, oobn, flags, addr, nil
}
if !errors.Is(err, conceal.ErrFormat) || addr == nil {
return n, oobn, flags, addr, err
}
if data := conceal.FormatErrorData(err); len(data) > 0 {
_ = c.forward(addr, data)
}
}
}
func (c *fallbackUDPConn) Close() error {
c.mu.Lock()
for key, session := range c.sessions {
session.close()
delete(c.sessions, key)
}
c.mu.Unlock()
return c.UDPConn.Close()
}
func (c *fallbackUDPConn) forward(remote *net.UDPAddr, data []byte) error {
session, err := c.session(remote)
if err != nil {
return err
}
_, err = session.conn.Write(data)
return err
}
func (c *fallbackUDPConn) session(remote *net.UDPAddr) (*fallbackUDPSession, error) {
key := remote.String()
c.mu.Lock()
if session := c.sessions[key]; session != nil {
c.mu.Unlock()
return session, nil
}
c.mu.Unlock()
fallbackConn, err := net.DialUDP("udp", nil, fallbackUDPAddress(remote, c.port))
if err != nil {
return nil, err
}
session := &fallbackUDPSession{
parent: c,
remote: &net.UDPAddr{
IP: append(net.IP(nil), remote.IP...),
Port: remote.Port,
Zone: remote.Zone,
},
conn: fallbackConn,
}
c.mu.Lock()
if existing := c.sessions[key]; existing != nil {
c.mu.Unlock()
fallbackConn.Close()
return existing, nil
}
c.sessions[key] = session
c.mu.Unlock()
go session.relay()
return session, nil
}
func (c *fallbackUDPConn) removeSession(remote string) {
c.mu.Lock()
delete(c.sessions, remote)
c.mu.Unlock()
}
type fallbackUDPSession struct {
parent *fallbackUDPConn
remote *net.UDPAddr
conn *net.UDPConn
once sync.Once
}
func (s *fallbackUDPSession) relay() {
defer s.parent.removeSession(s.remote.String())
buf := make([]byte, 65535)
for {
n, err := s.conn.Read(buf)
if err != nil {
return
}
_, _, _ = s.parent.origin.WriteMsgUDP(buf[:n], nil, s.remote)
}
}
func (s *fallbackUDPSession) close() {
s.once.Do(func() {
s.conn.Close()
})
}
type fallbackBatchConn struct {
LinuxPacketConn
origin LinuxPacketConn
port uint16
mu sync.Mutex
sessions map[string]*fallbackBatchSession
}
type fallbackSessionCloser interface {
closeFallbackSessions()
}
func newFallbackBatchConn(conn, origin LinuxPacketConn, port uint16) LinuxPacketConn {
return &fallbackBatchConn{
LinuxPacketConn: conn,
origin: origin,
port: port,
sessions: make(map[string]*fallbackBatchSession),
}
}
func (c *fallbackBatchConn) ReadBatch(ms []ipv4.Message, flags int) (n int, err error) {
for {
n, err = c.LinuxPacketConn.ReadBatch(ms, flags)
if err == nil {
return n, nil
}
if !errors.Is(err, conceal.ErrFormat) {
return n, err
}
data := conceal.FormatErrorData(err)
remote := firstBatchUDPAddr(ms)
if len(data) == 0 || remote == nil {
return n, err
}
_ = c.forward(remote, data)
}
}
func (c *fallbackBatchConn) closeFallbackSessions() {
c.mu.Lock()
for key, session := range c.sessions {
session.close()
delete(c.sessions, key)
}
c.mu.Unlock()
}
func (c *fallbackBatchConn) forward(remote *net.UDPAddr, data []byte) error {
session, err := c.session(remote)
if err != nil {
return err
}
_, err = session.conn.Write(data)
return err
}
func (c *fallbackBatchConn) session(remote *net.UDPAddr) (*fallbackBatchSession, error) {
key := remote.String()
c.mu.Lock()
if session := c.sessions[key]; session != nil {
c.mu.Unlock()
return session, nil
}
c.mu.Unlock()
fallbackConn, err := net.DialUDP("udp", nil, fallbackUDPAddress(remote, c.port))
if err != nil {
return nil, err
}
session := &fallbackBatchSession{
parent: c,
remote: &net.UDPAddr{
IP: append(net.IP(nil), remote.IP...),
Port: remote.Port,
Zone: remote.Zone,
},
conn: fallbackConn,
}
c.mu.Lock()
if existing := c.sessions[key]; existing != nil {
c.mu.Unlock()
fallbackConn.Close()
return existing, nil
}
c.sessions[key] = session
c.mu.Unlock()
go session.relay()
return session, nil
}
func (c *fallbackBatchConn) removeSession(remote string) {
c.mu.Lock()
delete(c.sessions, remote)
c.mu.Unlock()
}
type fallbackBatchSession struct {
parent *fallbackBatchConn
remote *net.UDPAddr
conn *net.UDPConn
once sync.Once
}
func (s *fallbackBatchSession) relay() {
defer s.parent.removeSession(s.remote.String())
buf := make([]byte, 65535)
for {
n, err := s.conn.Read(buf)
if err != nil {
return
}
msg := ipv4.Message{
Buffers: [][]byte{buf[:n]},
Addr: s.remote,
}
_, _ = s.parent.origin.WriteBatch([]ipv4.Message{msg}, 0)
}
}
func (s *fallbackBatchSession) close() {
s.once.Do(func() {
s.conn.Close()
})
}
func firstBatchUDPAddr(ms []ipv4.Message) *net.UDPAddr {
for i := range ms {
if ms[i].N == 0 || ms[i].Addr == nil {
continue
}
if addr, ok := ms[i].Addr.(*net.UDPAddr); ok {
return addr
}
}
return nil
}
+198
View File
@@ -0,0 +1,198 @@
package conn
import (
"bytes"
"io"
"net"
"strconv"
"sync"
"testing"
"time"
"github.com/amnezia-vpn/amneziawg-go/conceal"
)
func TestBindStreamFallbackProxiesFormatErrorToTCPPort(t *testing.T) {
fallbackListener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("listen fallback: %v", err)
}
defer fallbackListener.Close()
fallbackPort := fallbackListener.Addr().(*net.TCPAddr).Port
probe := []byte("GET / HTTP/1.1\r\n")
response := []byte("HTTP/1.1 200 OK\r\n\r\n")
received := make(chan []byte, 1)
served := make(chan error, 1)
go func() {
conn, err := fallbackListener.Accept()
if err != nil {
served <- err
return
}
defer conn.Close()
buf := make([]byte, len(probe))
if _, err := io.ReadFull(conn, buf); err != nil {
served <- err
return
}
received <- bytes.Clone(buf)
_, err = conn.Write(response)
served <- err
}()
bind := NewBindStream()
bind.SetMasqueradeOpts(conceal.MasqueradeOpts{
RulesIn: mustParseRules(t, "<b 0xfeed><dz be 2><d>"),
RulesOut: mustParseRules(t, "<b 0xfeed><dz be 2><d>"),
})
bind.SetFallbackPort(uint16(fallbackPort))
listenPort := freeTCPPort(t)
_, actualPort, err := bind.Open(listenPort)
if err != nil {
t.Fatalf("open bind stream: %v", err)
}
defer bind.Close()
client, err := net.Dial("tcp", net.JoinHostPort("127.0.0.1", strconv.Itoa(int(actualPort))))
if err != nil {
t.Fatalf("dial bind stream: %v", err)
}
defer client.Close()
if err := client.SetDeadline(time.Now().Add(2 * time.Second)); err != nil {
t.Fatalf("set deadline: %v", err)
}
if _, err := client.Write(probe[:2]); err != nil {
t.Fatalf("write initial probe bytes: %v", err)
}
if _, err := client.Write(probe[2:]); err != nil {
t.Fatalf("write remaining probe bytes: %v", err)
}
gotResponse := make([]byte, len(response))
if _, err := io.ReadFull(client, gotResponse); err != nil {
t.Fatalf("read fallback response: %v", err)
}
if !bytes.Equal(gotResponse, response) {
t.Fatalf("fallback response = %q, want %q", gotResponse, response)
}
select {
case gotProbe := <-received:
if !bytes.Equal(gotProbe, probe) {
t.Fatalf("fallback received = %q, want %q", gotProbe, probe)
}
case <-time.After(2 * time.Second):
t.Fatal("fallback service did not receive probe")
}
if err := <-served; err != nil {
t.Fatalf("fallback service failed: %v", err)
}
}
func TestFallbackUDPConnProxiesFormatErrorToUDPPort(t *testing.T) {
fallbackServer, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
if err != nil {
t.Fatalf("listen fallback udp: %v", err)
}
defer fallbackServer.Close()
probe := []byte("GE")
response := []byte("udp-ok")
received := make(chan []byte, 1)
served := make(chan error, 1)
go func() {
buf := make([]byte, 64)
n, addr, err := fallbackServer.ReadFromUDP(buf)
if err != nil {
served <- err
return
}
received <- bytes.Clone(buf[:n])
_, err = fallbackServer.WriteToUDP(response, addr)
served <- err
}()
raw, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
if err != nil {
t.Fatalf("listen awg udp: %v", err)
}
defer raw.Close()
pool := sync.Pool{
New: func() any {
return make([]byte, 65535)
},
}
masquerade, ok := conceal.NewMasqueradeUDPConn(raw, &pool, conceal.MasqueradeOpts{
RulesIn: mustParseRules(t, "<b 0xfeed><dz be 2><d>"),
})
if !ok {
t.Fatal("expected masquerade udp conn")
}
fallback := newFallbackUDPConn(masquerade, raw, uint16(fallbackServer.LocalAddr().(*net.UDPAddr).Port))
readDone := make(chan error, 1)
go func() {
buf := make([]byte, 64)
_, _, _, _, err := fallback.ReadMsgUDP(buf, nil)
readDone <- err
}()
client, err := net.DialUDP("udp4", nil, raw.LocalAddr().(*net.UDPAddr))
if err != nil {
t.Fatalf("dial awg udp: %v", err)
}
defer client.Close()
if err := client.SetDeadline(time.Now().Add(2 * time.Second)); err != nil {
t.Fatalf("set udp deadline: %v", err)
}
if _, err := client.Write(probe); err != nil {
t.Fatalf("write udp probe: %v", err)
}
gotResponse := make([]byte, len(response))
if _, err := io.ReadFull(client, gotResponse); err != nil {
t.Fatalf("read udp fallback response: %v", err)
}
if !bytes.Equal(gotResponse, response) {
t.Fatalf("udp fallback response = %q, want %q", gotResponse, response)
}
select {
case gotProbe := <-received:
if !bytes.Equal(gotProbe, probe) {
t.Fatalf("udp fallback received = %q, want %q", gotProbe, probe)
}
case <-time.After(2 * time.Second):
t.Fatal("udp fallback service did not receive probe")
}
if err := <-served; err != nil {
t.Fatalf("udp fallback service failed: %v", err)
}
raw.Close()
select {
case <-readDone:
case <-time.After(2 * time.Second):
t.Fatal("fallback ReadMsgUDP did not exit after close")
}
}
func freeTCPPort(t *testing.T) uint16 {
t.Helper()
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("allocate tcp port: %v", err)
}
defer listener.Close()
return uint16(listener.Addr().(*net.TCPAddr).Port)
}
+5
View File
@@ -50,6 +50,7 @@ type Device struct {
framedOpts conceal.FramedOpts
preludeOpts conceal.PreludeOpts
masqueradeOpts conceal.MasqueradeOpts
fallbackPort uint16
}
staticIdentity struct {
@@ -516,6 +517,10 @@ func (device *Device) BindUpdate() error {
masqueradable.SetMasqueradeOpts(netc.masqueradeOpts)
}
if fallbackable, ok := underlying.(conn.Fallbackable); ok {
fallbackable.SetFallbackPort(netc.fallbackPort)
}
recvFns, netc.port, err = netc.bind.Open(netc.port)
if err != nil {
netc.port = 0
+20
View File
@@ -160,6 +160,10 @@ func (device *Device) IpcGetOperation(w io.Writer) error {
sendf("format_out=%s", device.net.masqueradeOpts.RulesOut.Spec())
}
if device.net.fallbackPort != 0 {
sendf("fallback_port=%d", device.net.fallbackPort)
}
sendf("header_compat=%s", strconv.FormatBool(device.net.framedOpts.HeaderCompat))
for _, peer := range device.peers.keyMap {
@@ -614,6 +618,22 @@ func (device *Device) handleDeviceLine(key, value string) error {
return ipcErrorf(ipc.IpcErrorPortInUse, "failed to set header_compat: %w", err)
}
case "fallback_port":
port, err := strconv.ParseUint(value, 10, 16)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse fallback_port: %w", err)
}
if port == 0 {
return ipcErrorf(ipc.IpcErrorInvalid, "fallback_port must be in range 1-65535")
}
device.log.Verbosef("UAPI: Updating fallback_port")
device.net.fallbackPort = uint16(port)
if err := device.BindUpdate(); err != nil {
return ipcErrorf(ipc.IpcErrorPortInUse, "failed to set fallback_port: %w", err)
}
default:
return ipcErrorf(ipc.IpcErrorInvalid, "invalid UAPI device key: %v", key)
}
+35
View File
@@ -0,0 +1,35 @@
package device
import (
"bytes"
"strings"
"testing"
"github.com/amnezia-vpn/amneziawg-go/conn/bindtest"
"github.com/amnezia-vpn/amneziawg-go/tun/tuntest"
)
func TestDeviceFallbackPortUAPI(t *testing.T) {
tunDevice := tuntest.NewChannelTUN()
binds := bindtest.NewChannelBinds()
device := NewDevice(tunDevice.TUN(), binds[0], NewLogger(LogLevelError, ""))
defer device.Close()
if err := device.IpcSet(uapiCfg("fallback_port", "8082")); err != nil {
t.Fatalf("set fallback_port: %v", err)
}
var got bytes.Buffer
if err := device.IpcGetOperation(&got); err != nil {
t.Fatalf("get uapi: %v", err)
}
if !strings.Contains(got.String(), "fallback_port=8082\n") {
t.Fatalf("IpcGetOperation output missing fallback_port: %q", got.String())
}
for _, value := range []string{"0", "65536"} {
if err := device.IpcSet(uapiCfg("fallback_port", value)); err == nil {
t.Fatalf("fallback_port=%s accepted, want error", value)
}
}
}
+26 -21
View File
@@ -16,27 +16,28 @@ import (
)
type DeviceConfig struct {
PrivateKey string `yaml:"private_key"`
Address []string `yaml:"address"`
Dns []string `yaml:"dns"`
Mtu int `yaml:"mtu,omitempty"`
Jc int `yaml:"jc,omitempty"`
Jmin int `yaml:"jmin,omitempty"`
Jmax int `yaml:"jmax,omitempty"`
S1 int `yaml:"s1,omitempty"`
S2 int `yaml:"s2,omitempty"`
S3 int `yaml:"s3,omitempty"`
S4 int `yaml:"s4,omitempty"`
H1 string `yaml:"h1,omitempty"`
H2 string `yaml:"h2,omitempty"`
H3 string `yaml:"h3,omitempty"`
H4 string `yaml:"h4,omitempty"`
I1 string `yaml:"i1,omitempty"`
I2 string `yaml:"i2,omitempty"`
I3 string `yaml:"i3,omitempty"`
I4 string `yaml:"i4,omitempty"`
I5 string `yaml:"i5,omitempty"`
Peers []PeerConfig `yaml:"peers,omitempty"`
PrivateKey string `yaml:"private_key"`
Address []string `yaml:"address"`
Dns []string `yaml:"dns"`
Mtu int `yaml:"mtu,omitempty"`
Jc int `yaml:"jc,omitempty"`
Jmin int `yaml:"jmin,omitempty"`
Jmax int `yaml:"jmax,omitempty"`
S1 int `yaml:"s1,omitempty"`
S2 int `yaml:"s2,omitempty"`
S3 int `yaml:"s3,omitempty"`
S4 int `yaml:"s4,omitempty"`
H1 string `yaml:"h1,omitempty"`
H2 string `yaml:"h2,omitempty"`
H3 string `yaml:"h3,omitempty"`
H4 string `yaml:"h4,omitempty"`
I1 string `yaml:"i1,omitempty"`
I2 string `yaml:"i2,omitempty"`
I3 string `yaml:"i3,omitempty"`
I4 string `yaml:"i4,omitempty"`
I5 string `yaml:"i5,omitempty"`
FallbackPort int `yaml:"fallback_port,omitempty"`
Peers []PeerConfig `yaml:"peers,omitempty"`
}
type PeerConfig struct {
@@ -136,6 +137,10 @@ func genIpcString(cfg *DeviceConfig) (string, error) {
b.WriteString("\ni5=")
b.WriteString(cfg.I5)
}
if cfg.FallbackPort != 0 {
b.WriteString("\nfallback_port=")
b.WriteString(strconv.Itoa(cfg.FallbackPort))
}
for _, peer := range cfg.Peers {
publicKeyBytes, err := base64.StdEncoding.DecodeString(peer.PublicKey)