mirror of
https://github.com/amnezia-vpn/amneziawg-go.git
synced 2026-10-02 21:36:06 +03:00
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:
1 parent
8d9bf0ed57
commit
f8b2d2ba5a
17 files changed
+881
-74
No files matched your search
@@ -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
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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, ×tamp); 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
|
||||
}
|
||||
@@ -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]) {
|
||||
|
||||
@@ -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
@@ -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,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
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in new issue
Block a user