Compare commits

..

3 Commits

Author SHA1 Message Date
Quentin McGaw 10bcad1be6 pr feedback 2026-05-28 14:43:20 +00:00
Quentin McGaw a25df5fe82 add integration test 2026-05-27 15:17:25 +00:00
Quentin McGaw 7a72a5373b initial 2026-05-27 00:53:42 +00:00
18 changed files with 160 additions and 420 deletions
+11 -15
View File
@@ -28,10 +28,6 @@ on:
- go.mod - go.mod
- go.sum - go.sum
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs: jobs:
verify: verify:
runs-on: ubuntu-latest runs-on: ubuntu-latest
@@ -41,7 +37,7 @@ jobs:
env: env:
DOCKER_BUILDKIT: "1" DOCKER_BUILDKIT: "1"
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- uses: reviewdog/action-misspell@v1 - uses: reviewdog/action-misspell@v1
with: with:
@@ -79,7 +75,7 @@ jobs:
actions: read actions: read
contents: read contents: read
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- uses: actions/setup-go@v6 - uses: actions/setup-go@v6
with: with:
@@ -105,7 +101,7 @@ jobs:
runs-on: ubuntu-latest runs-on: ubuntu-latest
environment: secrets environment: secrets
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- run: docker build -t qmcgaw/gluetun . - run: docker build -t qmcgaw/gluetun .
@@ -133,12 +129,12 @@ jobs:
secrets.PROTONVPN_OPENVPN_PASSWORD }}" | ./ci/runner secrets.PROTONVPN_OPENVPN_PASSWORD }}" | ./ci/runner
protonvpn-openvpn-port-forwarding protonvpn-openvpn-port-forwarding
# - name: - name:
# Run Gluetun container with Private Internet Access OpenVPN and port Run Gluetun container with Private Internet Access OpenVPN and port
# forwarding configuration forwarding configuration
# run: echo -e "${{ secrets.PRIVATEINTERNETACCESS_OPENVPN_USER }}\n${{ run: echo -e "${{ secrets.PRIVATEINTERNETACCESS_OPENVPN_USER }}\n${{
# secrets.PRIVATEINTERNETACCESS_OPENVPN_PASSWORD }}" | ./ci/runner secrets.PRIVATEINTERNETACCESS_OPENVPN_PASSWORD }}" | ./ci/runner
# private-internet-access-openvpn-port-forwarding private-internet-access-openvpn-port-forwarding
- name: Run Gluetun container with AirVPN Wireguard configuration - name: Run Gluetun container with AirVPN Wireguard configuration
run: echo -e "${{ secrets.AIRVPN_WIREGUARD_PRIVATE_KEY }}\n${{ run: echo -e "${{ secrets.AIRVPN_WIREGUARD_PRIVATE_KEY }}\n${{
@@ -157,7 +153,7 @@ jobs:
contents: read contents: read
security-events: write security-events: write
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- uses: actions/setup-go@v6 - uses: actions/setup-go@v6
with: with:
go-version-file: go.mod go-version-file: go.mod
@@ -183,7 +179,7 @@ jobs:
runs-on: ubuntu-latest runs-on: ubuntu-latest
environment: secrets environment: secrets
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
# extract metadata (tags, labels) for Docker # extract metadata (tags, labels) for Docker
# https://github.com/docker/metadata-action # https://github.com/docker/metadata-action
+1 -1
View File
@@ -11,7 +11,7 @@ jobs:
issues: write issues: write
runs-on: ubuntu-latest runs-on: ubuntu-latest
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- uses: crazy-max/ghaction-github-labeler@v6 - uses: crazy-max/ghaction-github-labeler@v6
with: with:
yaml-file: .github/labels.yml yaml-file: .github/labels.yml
+1 -5
View File
@@ -11,10 +11,6 @@ on:
- "**.md" - "**.md"
- .github/workflows/markdown.yml - .github/workflows/markdown.yml
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs: jobs:
markdown: markdown:
runs-on: ubuntu-latest runs-on: ubuntu-latest
@@ -23,7 +19,7 @@ jobs:
contents: read contents: read
environment: secrets environment: secrets
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- uses: DavidAnson/markdownlint-cli2-action@v22 - uses: DavidAnson/markdownlint-cli2-action@v22
with: with:
-4
View File
@@ -12,10 +12,6 @@ formatters:
- builtin$ - builtin$
- examples$ - examples$
run:
build-tags:
- integration
linters: linters:
settings: settings:
misspell: misspell:
+1 -1
View File
@@ -276,7 +276,7 @@ ENV VPN_SERVICE_PROVIDER=pia \
PUID=1000 \ PUID=1000 \
PGID=1000 PGID=1000
ENTRYPOINT ["/gluetun-entrypoint"] ENTRYPOINT ["/gluetun-entrypoint"]
EXPOSE 8000/tcp 8888/tcp 8388/tcp 8388/udp 1080/tcp 1080/udp EXPOSE 8000/tcp 8888/tcp 8388/tcp 8388/udp 1080/tcp
HEALTHCHECK --interval=5s --timeout=5s --start-period=10s --retries=3 CMD /gluetun-entrypoint healthcheck HEALTHCHECK --interval=5s --timeout=5s --start-period=10s --retries=3 CMD /gluetun-entrypoint healthcheck
ARG TARGETPLATFORM ARG TARGETPLATFORM
RUN apk add --no-cache --update -l wget && \ RUN apk add --no-cache --update -l wget && \
+1 -1
View File
@@ -73,7 +73,7 @@ Lightweight swiss-army-knife-like VPN client to multiple VPN service providers
- Choose the vpn network protocol, `udp` or `tcp` - Choose the vpn network protocol, `udp` or `tcp`
- Built in firewall kill switch to allow traffic only with needed the VPN servers and LAN devices - Built in firewall kill switch to allow traffic only with needed the VPN servers and LAN devices
- Built in Shadowsocks proxy server (protocol based on SOCKS5 with an encryption layer, tunnels TCP+UDP) - Built in Shadowsocks proxy server (protocol based on SOCKS5 with an encryption layer, tunnels TCP+UDP)
- Built in Socks5 proxy server (tunnels TCP+UDP) - partial credits to @angelakis and @adjscent - Built in Socks5 proxy server (tunnels TCP) - partial credits to @angelakis and @adjscent
- Built in HTTP proxy (tunnels HTTP and HTTPS through TCP) - Built in HTTP proxy (tunnels HTTP and HTTPS through TCP)
- [Connect other containers to it](https://github.com/qdm12/gluetun-wiki/blob/main/setup/connect-a-container-to-gluetun.md) - [Connect other containers to it](https://github.com/qdm12/gluetun-wiki/blob/main/setup/connect-a-container-to-gluetun.md)
- [Connect LAN devices to it](https://github.com/qdm12/gluetun-wiki/blob/main/setup/connect-a-lan-device-to-gluetun.md) - [Connect LAN devices to it](https://github.com/qdm12/gluetun-wiki/blob/main/setup/connect-a-lan-device-to-gluetun.md)
+4 -10
View File
@@ -9,9 +9,8 @@ import (
) )
// Start launches a command and streams stdout and stderr to channels. // Start launches a command and streams stdout and stderr to channels.
// stdoutLines and stderrLines channels will be closed when there is no more // All the channels returned are ready only and won't be closed
// output to read, in order for the caller to catch all lines even after the // if the command fails later.
// command has finished. The waitError channel returned will never be closed.
func (c *Cmder) Start(cmd *exec.Cmd) ( func (c *Cmder) Start(cmd *exec.Cmd) (
stdoutLines, stderrLines <-chan string, stdoutLines, stderrLines <-chan string,
waitError <-chan error, startErr error, waitError <-chan error, startErr error,
@@ -39,7 +38,6 @@ func start(cmd execCmd) (stdoutLines, stderrLines <-chan string,
if err != nil { if err != nil {
_ = stdout.Close() _ = stdout.Close()
<-stdoutDone <-stdoutDone
close(stdoutLinesCh)
return nil, nil, nil, err return nil, nil, nil, err
} }
go streamToChannel(stderrReady, stderrDone, stderr, stderrLinesCh) go streamToChannel(stderrReady, stderrDone, stderr, stderrLinesCh)
@@ -47,11 +45,9 @@ func start(cmd execCmd) (stdoutLines, stderrLines <-chan string,
err = cmd.Start() err = cmd.Start()
if err != nil { if err != nil {
_ = stdout.Close() _ = stdout.Close()
<-stdoutDone
close(stdoutLinesCh)
_ = stderr.Close() _ = stderr.Close()
<-stdoutDone
<-stderrDone <-stderrDone
close(stderrLinesCh)
return nil, nil, nil, err return nil, nil, nil, err
} }
@@ -59,10 +55,8 @@ func start(cmd execCmd) (stdoutLines, stderrLines <-chan string,
go func() { go func() {
err := cmd.Wait() err := cmd.Wait()
<-stdoutDone <-stdoutDone
close(stdoutLinesCh)
_ = stdout.Close()
<-stderrDone <-stderrDone
close(stderrLinesCh) _ = stdout.Close()
_ = stderr.Close() _ = stderr.Close()
waitErrorCh <- err waitErrorCh <- err
}() }()
+24 -42
View File
@@ -89,48 +89,30 @@ func Test_start(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
collectAndCheckChannels(t, stdoutLines, stderrLines, waitError, var stdoutIndex, stderrIndex int
testCase.stdout, testCase.stderr, testCase.waitErr)
done := false
for !done {
select {
case line := <-stdoutLines:
assert.Equal(t, testCase.stdout[stdoutIndex], line)
stdoutIndex++
case line := <-stderrLines:
assert.Equal(t, testCase.stderr[stderrIndex], line)
stderrIndex++
case err := <-waitError:
if testCase.waitErr != nil {
require.Error(t, err)
assert.Equal(t, testCase.waitErr.Error(), err.Error())
} else {
assert.NoError(t, err)
}
done = true
}
}
assert.Equal(t, len(testCase.stdout), stdoutIndex)
assert.Equal(t, len(testCase.stderr), stderrIndex)
}) })
} }
} }
func collectAndCheckChannels(t *testing.T, stdoutLines, stderrLines <-chan string,
waitError <-chan error, expectedStdout, expectedStderr []string, expectedWaitErr error,
) {
t.Helper()
stdoutIndex := 0
stderrIndex := 0
done := false
for !done {
select {
case line, ok := <-stdoutLines:
if !ok {
stdoutLines = nil
continue
}
assert.Equal(t, expectedStdout[stdoutIndex], line)
stdoutIndex++
case line, ok := <-stderrLines:
if !ok {
stderrLines = nil
continue
}
assert.Equal(t, expectedStderr[stderrIndex], line)
stderrIndex++
case err := <-waitError:
if expectedWaitErr != nil {
require.Error(t, err)
assert.Equal(t, expectedWaitErr.Error(), err.Error())
} else {
assert.NoError(t, err)
}
done = true
}
}
assert.Equal(t, len(expectedStdout), stdoutIndex)
assert.Equal(t, len(expectedStderr), stderrIndex)
}
+13 -21
View File
@@ -18,39 +18,31 @@ func (c *Cmder) RunAndLog(ctx context.Context, command string, logger Logger) (e
return err return err
} }
streamCtx, streamCancel := context.WithCancel(context.Background())
streamDone := make(chan struct{}) streamDone := make(chan struct{})
go streamLines(streamDone, logger, stdout, stderr) go streamLines(streamCtx, streamDone, logger, stdout, stderr)
err = <-waitError err = <-waitError
streamCancel()
<-streamDone <-streamDone
return err return err
} }
func streamLines(done chan<- struct{}, logger Logger, func streamLines(ctx context.Context, done chan<- struct{},
stdout, stderr <-chan string, logger Logger, stdout, stderr <-chan string,
) { ) {
defer close(done) defer close(done)
var line string
for { for {
select { select {
case line, ok := <-stdout: case <-ctx.Done():
if ok { return
logger.Info(line) case line = <-stdout:
break logger.Info(line)
} case line = <-stderr:
if stderr == nil { logger.Error(line)
return
}
stdout = nil
case line, ok := <-stderr:
if ok {
logger.Error(line)
break
}
if stdout == nil {
return
}
stderr = nil
} }
} }
} }
+4 -1
View File
@@ -29,16 +29,19 @@ func (r *Runner) Run(ctx context.Context, errCh chan<- error, ready chan<- struc
return return
} }
streamCtx, streamCancel := context.WithCancel(context.Background())
streamDone := make(chan struct{}) streamDone := make(chan struct{})
go streamLines(streamDone, r.logger, go streamLines(streamCtx, streamDone, r.logger,
stdoutLines, stderrLines, ready) stdoutLines, stderrLines, ready)
select { select {
case <-ctx.Done(): case <-ctx.Done():
<-waitError <-waitError
streamCancel()
<-streamDone <-streamDone
errCh <- ctx.Err() errCh <- ctx.Err()
case err := <-waitError: case err := <-waitError:
streamCancel()
<-streamDone <-streamDone
errCh <- err errCh <- err
} }
+9 -20
View File
@@ -1,37 +1,26 @@
package openvpn package openvpn
import ( import (
"context"
"strings" "strings"
) )
func streamLines(done chan<- struct{}, func streamLines(ctx context.Context, done chan<- struct{},
logger Logger, stdout, stderr <-chan string, logger Logger, stdout, stderr <-chan string,
tunnelReady chan<- struct{}, tunnelReady chan<- struct{},
) { ) {
defer close(done) defer close(done)
var line string
for { for {
var line string
var ok bool
errLine := false errLine := false
select { select {
case line, ok = <-stdout: case <-ctx.Done():
if ok { return
break case line = <-stdout:
} case line = <-stderr:
if stderr == nil { errLine = true
return
}
stdout = nil
case line, ok = <-stderr:
if ok {
errLine = true
break
}
if stdout == nil {
return
}
stderr = nil
} }
line, level := processLogLine(line) line, level := processLogLine(line)
if line == "" { if line == "" {
+1 -6
View File
@@ -6,7 +6,6 @@ import (
"strings" "strings"
"github.com/qdm12/gluetun/internal/configuration/settings" "github.com/qdm12/gluetun/internal/configuration/settings"
"github.com/qdm12/gluetun/internal/constants"
"github.com/qdm12/gluetun/internal/constants/openvpn" "github.com/qdm12/gluetun/internal/constants/openvpn"
"github.com/qdm12/gluetun/internal/models" "github.com/qdm12/gluetun/internal/models"
"github.com/qdm12/gluetun/internal/provider/utils" "github.com/qdm12/gluetun/internal/provider/utils"
@@ -66,11 +65,7 @@ func modifyConfig(lines []string, connection models.Connection,
} }
// Add values // Add values
protocol := connection.Protocol modified = append(modified, "proto "+connection.Protocol)
if protocol == constants.TCP {
protocol = "tcp-client"
}
modified = append(modified, "proto "+protocol)
modified = append(modified, fmt.Sprintf("remote %s %d", connection.IP, connection.Port)) modified = append(modified, fmt.Sprintf("remote %s %d", connection.IP, connection.Port))
modified = append(modified, "dev "+settings.Interface) modified = append(modified, "dev "+settings.Interface)
modified = append(modified, "mute-replay-warnings") modified = append(modified, "mute-replay-warnings")
+63 -72
View File
@@ -21,8 +21,8 @@ type server struct {
listening atomic.Bool listening atomic.Bool
socksConnCtx context.Context //nolint:containedctx socksConnCtx context.Context //nolint:containedctx
socksConnCancel context.CancelFunc socksConnCancel context.CancelFunc
done <-chan error done <-chan struct{}
stopCh chan<- struct{} stopping atomic.Bool
} }
func newServer(settings Settings) *server { func newServer(settings Settings) *server {
@@ -58,11 +58,9 @@ func (s *server) Start(ctx context.Context) (runErr <-chan error, err error) {
ready := make(chan struct{}) ready := make(chan struct{})
runErrCh := make(chan error) runErrCh := make(chan error)
runErr = runErrCh runErr = runErrCh
done := make(chan error) done := make(chan struct{})
s.done = done s.done = done
stop := make(chan struct{}) go s.runServer(ready, runErrCh, done)
s.stopCh = stop
go s.runServer(ready, runErrCh, stop, done)
select { select {
case <-ready: case <-ready:
case <-ctx.Done(): case <-ctx.Done():
@@ -73,85 +71,78 @@ func (s *server) Start(ctx context.Context) (runErr <-chan error, err error) {
} }
func (s *server) runServer(ready chan<- struct{}, func (s *server) runServer(ready chan<- struct{},
runErrCh chan<- error, stop <-chan struct{}, done chan<- error, runErrCh chan<- error, done chan<- struct{},
) { ) {
close(ready) close(ready)
defer close(done) defer close(done)
wg := new(sync.WaitGroup)
defer wg.Wait()
udpErrCh := make(chan error) wg.Go(func() {
go func() { err := s.udpRouter.run(s.socksConnCtx)
udpErrCh <- s.udpRouter.run(s.socksConnCtx) if err != nil {
}() if !s.stopping.Load() {
_ = s.stop()
tcpErrCh := make(chan error) runErrCh <- fmt.Errorf("running UDP router: %w", err)
go func() {
var wg sync.WaitGroup
defer wg.Wait()
dialer := &net.Dialer{}
for {
connection, err := s.tcpListener.Accept()
if err != nil {
s.socksConnCancel() // stop ongoing TCP socks connections - no impact on UDP
tcpErrCh <- fmt.Errorf("accepting connection: %w", err)
return
} }
wg.Go(func() {
connection := connection // capture loop variable
socksConn := &socksConn{
dialer: dialer,
username: s.username,
password: s.password,
clientConn: connection,
udpRouter: s.udpRouter,
logger: s.logger,
}
err := socksConn.run(s.socksConnCtx)
if err != nil {
s.logger.Infof("running socks connection: %s", err)
}
})
} }
}() })
select { dialer := &net.Dialer{}
case <-stop: for {
s.listening.Store(false) connection, err := s.tcpListener.Accept()
var errs []error
err := s.tcpListener.Close()
if err != nil { if err != nil {
errs = append(errs, fmt.Errorf("closing TCP listener: %w", err)) if !s.stopping.Load() {
_ = s.stop()
runErrCh <- fmt.Errorf("accepting connection: %w", err)
}
return
} }
// stop ongoing TCP socks connections. This impacts the udpRouter run error when it is being closed. wg.Add(1)
s.socksConnCancel() go func(ctx context.Context, connection net.Conn,
<-tcpErrCh // wait for TCP server to stop dialer *net.Dialer, wg *sync.WaitGroup,
err = s.udpRouter.close() ) {
if err != nil { defer wg.Done()
errs = append(errs, fmt.Errorf("closing UDP router: %w", err)) socksConn := &socksConn{
} dialer: dialer,
<-udpErrCh // wait for UDP router to stop username: s.username,
if len(errs) > 0 { password: s.password,
// Only write to the done channel if the [server.Stop] method is waiting to read from it clientConn: connection,
done <- errors.Join(errs...) udpRouter: s.udpRouter,
} logger: s.logger,
// If no error, the done channel is closed so the error is effectively `nil` }
// Note: do NOT write an error the runError channel, since we are stopping the server gracefully. err := socksConn.run(ctx)
case err := <-udpErrCh: if err != nil {
_ = s.tcpListener.Close() // stop accepting new TCP connections s.logger.Infof("running socks connection: %s", err)
s.socksConnCancel() // stop ongoing TCP socks connections }
<-tcpErrCh // wait for TCP server to stop }(s.socksConnCtx, connection, dialer, wg)
runErrCh <- fmt.Errorf("running UDP router: %w", err)
case err := <-tcpErrCh:
s.socksConnCancel()
_ = s.udpRouter.close() // stop UDP router
<-udpErrCh // wait for UDP router to stop
runErrCh <- fmt.Errorf("running TCP server: %w", err)
} }
} }
func (s *server) Stop() (err error) { func (s *server) Stop() (err error) {
close(s.stopCh) s.stopping.Store(true)
return <-s.done err = s.stop()
<-s.done // wait for run goroutine to finish
s.stopping.Store(false)
return err
}
func (s *server) stop() error {
s.listening.Store(false)
var errs []error
err := s.tcpListener.Close()
if err != nil {
errs = append(errs, fmt.Errorf("closing TCP listener: %w", err))
}
err = s.udpRouter.close()
if err != nil {
errs = append(errs, fmt.Errorf("closing UDP router: %w", err))
}
s.socksConnCancel() // stop ongoing socks connections
if len(errs) > 0 {
return errors.Join(errs...)
}
return nil
} }
func (s *server) listeningAddress() net.Addr { func (s *server) listeningAddress() net.Addr {
-112
View File
@@ -1,112 +0,0 @@
//go:build integration
package socks5
import (
"math/rand/v2"
"net"
"testing"
"time"
"github.com/miekg/dns"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func Test_Server_UDPResolution(t *testing.T) {
t.Parallel()
ctx := t.Context()
server := newServer(Settings{
Address: "127.0.0.1:0",
Logger: noopLogger{},
})
runErr, err := server.Start(ctx)
require.NoError(t, err, "starting SOCKS5 server")
const timeout = 3 * time.Second
// Connect to the SOCKS5 server via TCP to negotiate UDP associate
dialer := &net.Dialer{Timeout: timeout}
tcpConn, err := dialer.DialContext(ctx, "tcp", server.listeningAddress().String())
require.NoError(t, err, "tcp connecting to SOCKS5 server")
t.Cleanup(func() { tcpConn.Close() })
negotiateSOCKS5(t, tcpConn, "", "")
// UDP Associate Command: [VERSION (5), CMD (3 = UDP ASSOC), RSV (0), ATYP (1 = IPv4), ADDR (0.0.0.0), PORT (0)]
_, err = tcpConn.Write([]byte{5, 3, 0, 1, 0, 0, 0, 0, 0, 0})
require.NoError(t, err, "sending UDP ASSOC request")
relayAddressString, err := readSOCKS5ResponseAddress(t, tcpConn)
require.NoError(t, err, "reading UDP ASSOC reply")
relayAddress, err := net.ResolveUDPAddr("udp", relayAddressString)
require.NoError(t, err, "resolving udp relay address")
// Dial the relay using IPv4 so source IP family matches the control connection.
udpConn, err := net.DialUDP("udp4", nil, relayAddress)
require.NoError(t, err, "dialing UDP relay")
t.Cleanup(func() { _ = udpConn.Close() })
queryID := uint16(rand.Uint32()) //nolint:gosec
dnsRequest := &dns.Msg{
MsgHdr: dns.MsgHdr{
Id: queryID,
RecursionDesired: true,
},
Question: []dns.Question{{
Name: dns.Fqdn("github.com"),
Qtype: dns.TypeA,
Qclass: dns.ClassINET,
}},
}
dnsQuery, err := dnsRequest.Pack()
require.NoError(t, err)
// Encapsulate DNS payload into SOCKS5 UDP Request Header
// [RSV (0,0), FRAG (0), ATYP (1 = IPv4), DST.ADDR (1.1.1.1), DST.PORT (53)]
packet := append([]byte{0, 0, 0, 1, 1, 1, 1, 1, 0, 53}, dnsQuery...)
// Send encapsulated packet to the proxy's UDP relay address
_, err = udpConn.Write(packet)
require.NoError(t, err, "sending UDP packet to relay")
// Read response from the proxy relay
err = udpConn.SetReadDeadline(time.Now().Add(timeout))
require.NoError(t, err, "setting read deadline on UDP connection")
buffer := make([]byte, 2048)
n, err := udpConn.Read(buffer)
require.NoError(t, err, "receiving UDP response from relay")
const minimumHeaderSize = 10
require.GreaterOrEqual(t, n, minimumHeaderSize, "received UDP packet too short to contain valid SOCKS5 header")
// Verify header layout and slice out the raw DNS response
// Header format: RSV(2) FRAG(1) ATYP(1) DST.ADDR(variable) DST.PORT(2)
atyp := buffer[3]
var headerSize int
switch atyp {
case 1: // IPv4
headerSize = 10
case 3: // Domain name
headerSize = 4 + 1 + int(buffer[4]) + 2
case 4: // IPv6
headerSize = 22
default:
t.Fatalf("Unknown ATYP in SOCKS5 UDP header: %d", atyp)
}
dnsResponse := new(dns.Msg)
err = dnsResponse.Unpack(buffer[headerSize:n])
require.NoError(t, err, "unpacking DNS response from SOCKS5 UDP packet")
assert.Equal(t, queryID, dnsResponse.Id, "DNS response ID should match query ID")
select {
case err := <-runErr:
require.NoError(t, err, "SOCKS5 server run error")
default:
}
err = server.Stop()
require.NoError(t, err, "stopping SOCKS5 server")
}
+18 -19
View File
@@ -10,7 +10,6 @@ import (
"net/netip" "net/netip"
"strconv" "strconv"
"strings" "strings"
"sync"
) )
var ( var (
@@ -211,13 +210,6 @@ func (c *socksConn) handleUDPAssociateRequest(ctx context.Context,
return fmt.Errorf("getting udp association addresses: %w", err) return fmt.Errorf("getting udp association addresses: %w", err)
} }
association, err := c.udpRouter.registerAssociation(c.clientConn, expectedAddrPort)
if err != nil {
c.encodeFailedResponse(c.clientConn, socksVersion, generalServerFailure)
return fmt.Errorf("registering udp association: %w", err)
}
defer c.udpRouter.unregisterAssociation(association)
err = c.encodeSuccessResponse(c.clientConn, socksVersion, succeeded, err = c.encodeSuccessResponse(c.clientConn, socksVersion, succeeded,
bindAddrType, bindAddress, bindPort) bindAddrType, bindAddress, bindPort)
if err != nil { if err != nil {
@@ -225,21 +217,28 @@ func (c *socksConn) handleUDPAssociateRequest(ctx context.Context,
return fmt.Errorf("writing successful %s response: %w", udpAssociate, err) return fmt.Errorf("writing successful %s response: %w", udpAssociate, err)
} }
association, err := c.udpRouter.registerAssociation(c.clientConn, expectedAddrPort)
if err != nil {
c.encodeFailedResponse(c.clientConn, socksVersion, generalServerFailure)
return fmt.Errorf("registering udp association: %w", err)
}
defer c.udpRouter.unregisterAssociation(association)
associationCtx, associationCancel := context.WithCancel(ctx) associationCtx, associationCancel := context.WithCancel(ctx)
defer associationCancel() defer associationCancel()
var wg sync.WaitGroup handlerDone := make(chan struct{})
go func() {
wg.Go(func() { defer close(handlerDone)
c.udpRouter.runAssociationHandler(associationCtx, association) c.udpRouter.runAssociationHandler(associationCtx, association)
}) }()
wg.Go(func() { go func() {
_, _ = io.Copy(io.Discard, c.clientConn) _, _ = io.Copy(io.Discard, c.clientConn)
associationCancel() associationCancel()
}) }()
<-associationCtx.Done() <-associationCtx.Done()
wg.Wait() <-handlerDone
return nil return nil
} }
@@ -256,10 +255,10 @@ func udpAssociateExpectedClientEndpoint(request request) (expectedAddrPort netip
} }
return netip.AddrPortFrom(netip.Addr{}, request.port), nil return netip.AddrPortFrom(netip.Addr{}, request.port), nil
case domainName: case domainName:
// For UDP associate, client endpoint matching is based on observed UDP source if request.destination != "" || request.port != 0 {
// address/port. A hostname is not directly matchable at this stage, so we return netip.AddrPort{}, fmt.Errorf("domain name is not supported for UDP associate destination")
// ignore the domain name request destination entirely. }
return netip.AddrPortFrom(netip.Addr{}, request.port), nil return netip.AddrPort{}, nil
default: default:
return netip.AddrPort{}, fmt.Errorf("address type %d is not supported", request.addressType) return netip.AddrPort{}, fmt.Errorf("address type %d is not supported", request.addressType)
} }
-65
View File
@@ -8,7 +8,6 @@ import (
"fmt" "fmt"
"io" "io"
"net" "net"
"net/netip"
"strconv" "strconv"
"strings" "strings"
"testing" "testing"
@@ -704,70 +703,6 @@ func Test_decodeRequest(t *testing.T) {
} }
} }
func Test_udpAssociateExpectedClientEndpoint(t *testing.T) {
t.Parallel()
testCases := map[string]struct {
request request
expected netip.AddrPort
expectedErr string
}{
"ipv4_endpoint": {
request: request{
addressType: ipv4,
destination: "192.0.2.10",
port: 5555,
},
expected: netip.MustParseAddrPort("192.0.2.10:5555"),
},
"ipv4_unspecified_address": {
request: request{
addressType: ipv4,
destination: "0.0.0.0",
port: 6000,
},
expected: netip.AddrPortFrom(netip.Addr{}, 6000),
},
"domain_name_with_port": {
request: request{
addressType: domainName,
destination: "client.example",
port: 7000,
},
expected: netip.AddrPortFrom(netip.Addr{}, 7000),
},
"domain_name_without_port": {
request: request{
addressType: domainName,
destination: "client.example",
},
expected: netip.AddrPort{},
},
"unsupported_address_type": {
request: request{
addressType: 255,
},
expectedErr: "address type 255 is not supported",
},
}
for name, testCase := range testCases {
t.Run(name, func(t *testing.T) {
t.Parallel()
result, err := udpAssociateExpectedClientEndpoint(testCase.request)
if testCase.expectedErr != "" {
assert.ErrorContains(t, err, testCase.expectedErr)
return
}
assert.NoError(t, err)
assert.Equal(t, testCase.expected, result)
})
}
}
func Test_verifyFirstNegotiation(t *testing.T) { func Test_verifyFirstNegotiation(t *testing.T) {
t.Parallel() t.Parallel()
testCases := map[string]struct { testCases := map[string]struct {
+7 -21
View File
@@ -48,7 +48,7 @@ func newUDPRouter(ctx context.Context, address string, logger Logger) (router *u
listener: listener, listener: listener,
bufferPool: sync.Pool{ bufferPool: sync.Pool{
New: func() any { New: func() any {
return bytes.NewBuffer(make([]byte, 0, pooledUDPPacketBufferCapacity)) return bytes.NewBuffer(make([]byte, pooledUDPPacketBufferCapacity))
}, },
}, },
nextAssociationID: 1, nextAssociationID: 1,
@@ -76,7 +76,7 @@ func (r *udpRouter) registerAssociation(controlConn net.Conn, expectedAddrPort n
r.mutex.Lock() r.mutex.Lock()
defer r.mutex.Unlock() defer r.mutex.Unlock()
const udpPacketChannelBuffer = 64 const udpPacketChannelBuffer = 2
associationID := r.nextAssociationID associationID := r.nextAssociationID
r.nextAssociationID++ r.nextAssociationID++
@@ -111,6 +111,10 @@ func (r *udpRouter) unregisterAssociation(association udpAssociation) {
delete(r.clientAddrPortToAssociation, clientAddrPort) delete(r.clientAddrPortToAssociation, clientAddrPort)
} }
if association.clientAddrPort.IsValid() {
delete(r.clientAddrPortToAssociation, association.clientAddrPort)
}
pendingAssociations := r.clientIPToPendingAssociations[association.controlConnAddr] pendingAssociations := r.clientIPToPendingAssociations[association.controlConnAddr]
for i, pendingAssociation := range pendingAssociations { for i, pendingAssociation := range pendingAssociations {
if pendingAssociation.id == association.id { if pendingAssociation.id == association.id {
@@ -159,9 +163,8 @@ func (r *udpRouter) run(ctx context.Context) error {
func (r *udpRouter) routePacket(sourceAddrPort netip.AddrPort, packet *bytes.Buffer) error { func (r *udpRouter) routePacket(sourceAddrPort netip.AddrPort, packet *bytes.Buffer) error {
r.mutex.Lock() r.mutex.Lock()
defer r.mutex.Unlock()
association, packetFromClient := r.findClientAssociation(sourceAddrPort) association, packetFromClient := r.findClientAssociation(sourceAddrPort)
r.mutex.Unlock()
if !packetFromClient { if !packetFromClient {
r.bufferPool.Put(packet) r.bufferPool.Put(packet)
return nil return nil
@@ -331,23 +334,6 @@ func (r *udpRouter) writeClientPacketToDestination(ctx context.Context,
return fmt.Errorf("decoding UDP datagram: %w", err) return fmt.Errorf("decoding UDP datagram: %w", err)
} }
host, portStr, err := net.SplitHostPort(destination)
if err != nil {
return fmt.Errorf("splitting destination host and port: %w", err)
}
if _, err := netip.ParseAddr(host); err != nil { // domain name
addrs, err := net.DefaultResolver.LookupHost(ctx, host)
if err != nil {
return fmt.Errorf("resolving destination host: %w", err)
}
if len(addrs) == 0 {
return fmt.Errorf("resolving destination host: no addresses found for %q", host)
}
destination = net.JoinHostPort(addrs[0], portStr)
}
resolvedDestinationUDPAddress, err := net.ResolveUDPAddr("udp", destination) resolvedDestinationUDPAddress, err := net.ResolveUDPAddr("udp", destination)
if err != nil { if err != nil {
return fmt.Errorf("resolving destination UDP address: %w", err) return fmt.Errorf("resolving destination UDP address: %w", err)
@@ -19,8 +19,6 @@ import (
) )
func Test_udpRouter_ResolveGithubFromCloudflareDNS(t *testing.T) { func Test_udpRouter_ResolveGithubFromCloudflareDNS(t *testing.T) {
t.Parallel()
ctx := t.Context() ctx := t.Context()
var cancel context.CancelFunc var cancel context.CancelFunc
deadline, hasDeadline := ctx.Deadline() deadline, hasDeadline := ctx.Deadline()
@@ -109,7 +107,7 @@ func Test_udpRouter_ResolveGithubFromCloudflareDNS(t *testing.T) {
assert.NoError(t, err, "closing client UDP connection") assert.NoError(t, err, "closing client UDP connection")
}) })
queryID := uint16(rand.Uint32()) //nolint:gosec queryID := uint16(rand.Uint())
dnsRequest := &dns.Msg{ dnsRequest := &dns.Msg{
MsgHdr: dns.MsgHdr{ MsgHdr: dns.MsgHdr{
Id: queryID, Id: queryID,
@@ -158,7 +156,7 @@ func Test_udpRouter_ResolveGithubFromCloudflareDNS(t *testing.T) {
assert.Equal(t, dns.RcodeSuccess, dnsResponse.Rcode) assert.Equal(t, dns.RcodeSuccess, dnsResponse.Rcode)
require.NotEmpty(t, dnsResponse.Question) require.NotEmpty(t, dnsResponse.Question)
assert.Equal(t, dns.Fqdn("github.com"), dnsResponse.Question[0].Name) assert.Equal(t, dns.Fqdn("github.com"), dnsResponse.Question[0].Name)
assert.Equal(t, dns.TypeA, dnsResponse.Question[0].Qtype) assert.Equal(t, uint16(dns.TypeA), dnsResponse.Question[0].Qtype)
assert.NotEmpty(t, dnsResponse.Answer) assert.NotEmpty(t, dnsResponse.Answer)
require.NoError(t, err) require.NoError(t, err)
} }