mirror of
https://github.com/qdm12/gluetun.git
synced 2026-07-21 09:56:31 +02:00
Compare commits
13 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| d7892a2ffe | |||
| 52a41cb891 | |||
| 6c76273ef6 | |||
| 366062dc12 | |||
| 8abb05567c | |||
| a53a0267e4 | |||
| 4e986c8af7 | |||
| 6d84462f00 | |||
| acab89b91a | |||
| 48c1f2bf6a | |||
| c599e7fd2c | |||
| ff6e45fae0 | |||
| 17f24343d6 |
+15
-11
@@ -28,6 +28,10 @@ on:
|
||||
- go.mod
|
||||
- go.sum
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
verify:
|
||||
runs-on: ubuntu-latest
|
||||
@@ -37,7 +41,7 @@ jobs:
|
||||
env:
|
||||
DOCKER_BUILDKIT: "1"
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: reviewdog/action-misspell@v1
|
||||
with:
|
||||
@@ -75,7 +79,7 @@ jobs:
|
||||
actions: read
|
||||
contents: read
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
with:
|
||||
@@ -101,7 +105,7 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
environment: secrets
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- run: docker build -t qmcgaw/gluetun .
|
||||
|
||||
@@ -129,12 +133,12 @@ jobs:
|
||||
secrets.PROTONVPN_OPENVPN_PASSWORD }}" | ./ci/runner
|
||||
protonvpn-openvpn-port-forwarding
|
||||
|
||||
- name:
|
||||
Run Gluetun container with Private Internet Access OpenVPN and port
|
||||
forwarding configuration
|
||||
run: echo -e "${{ secrets.PRIVATEINTERNETACCESS_OPENVPN_USER }}\n${{
|
||||
secrets.PRIVATEINTERNETACCESS_OPENVPN_PASSWORD }}" | ./ci/runner
|
||||
private-internet-access-openvpn-port-forwarding
|
||||
# - name:
|
||||
# Run Gluetun container with Private Internet Access OpenVPN and port
|
||||
# forwarding configuration
|
||||
# run: echo -e "${{ secrets.PRIVATEINTERNETACCESS_OPENVPN_USER }}\n${{
|
||||
# secrets.PRIVATEINTERNETACCESS_OPENVPN_PASSWORD }}" | ./ci/runner
|
||||
# private-internet-access-openvpn-port-forwarding
|
||||
|
||||
- name: Run Gluetun container with AirVPN Wireguard configuration
|
||||
run: echo -e "${{ secrets.AIRVPN_WIREGUARD_PRIVATE_KEY }}\n${{
|
||||
@@ -153,7 +157,7 @@ jobs:
|
||||
contents: read
|
||||
security-events: write
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v7
|
||||
- uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version-file: go.mod
|
||||
@@ -179,7 +183,7 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
environment: secrets
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
# extract metadata (tags, labels) for Docker
|
||||
# https://github.com/docker/metadata-action
|
||||
|
||||
@@ -11,7 +11,7 @@ jobs:
|
||||
issues: write
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v7
|
||||
- uses: crazy-max/ghaction-github-labeler@v6
|
||||
with:
|
||||
yaml-file: .github/labels.yml
|
||||
|
||||
@@ -11,6 +11,10 @@ on:
|
||||
- "**.md"
|
||||
- .github/workflows/markdown.yml
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
markdown:
|
||||
runs-on: ubuntu-latest
|
||||
@@ -19,7 +23,7 @@ jobs:
|
||||
contents: read
|
||||
environment: secrets
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/checkout@v7
|
||||
|
||||
- uses: DavidAnson/markdownlint-cli2-action@v22
|
||||
with:
|
||||
|
||||
@@ -12,6 +12,10 @@ formatters:
|
||||
- builtin$
|
||||
- examples$
|
||||
|
||||
run:
|
||||
build-tags:
|
||||
- integration
|
||||
|
||||
linters:
|
||||
settings:
|
||||
misspell:
|
||||
|
||||
+1
-1
@@ -276,7 +276,7 @@ ENV VPN_SERVICE_PROVIDER=pia \
|
||||
PUID=1000 \
|
||||
PGID=1000
|
||||
ENTRYPOINT ["/gluetun-entrypoint"]
|
||||
EXPOSE 8000/tcp 8888/tcp 8388/tcp 8388/udp 1080/tcp
|
||||
EXPOSE 8000/tcp 8888/tcp 8388/tcp 8388/udp 1080/tcp 1080/udp
|
||||
HEALTHCHECK --interval=5s --timeout=5s --start-period=10s --retries=3 CMD /gluetun-entrypoint healthcheck
|
||||
ARG TARGETPLATFORM
|
||||
RUN apk add --no-cache --update -l wget && \
|
||||
|
||||
@@ -73,7 +73,7 @@ Lightweight swiss-army-knife-like VPN client to multiple VPN service providers
|
||||
- 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 Shadowsocks proxy server (protocol based on SOCKS5 with an encryption layer, tunnels TCP+UDP)
|
||||
- Built in Socks5 proxy server (tunnels TCP) - partial credits to @angelakis and @adjscent
|
||||
- Built in Socks5 proxy server (tunnels TCP+UDP) - partial credits to @angelakis and @adjscent
|
||||
- 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 LAN devices to it](https://github.com/qdm12/gluetun-wiki/blob/main/setup/connect-a-lan-device-to-gluetun.md)
|
||||
|
||||
@@ -9,8 +9,9 @@ import (
|
||||
)
|
||||
|
||||
// Start launches a command and streams stdout and stderr to channels.
|
||||
// All the channels returned are ready only and won't be closed
|
||||
// if the command fails later.
|
||||
// stdoutLines and stderrLines channels will be closed when there is no more
|
||||
// output to read, in order for the caller to catch all lines even after the
|
||||
// command has finished. The waitError channel returned will never be closed.
|
||||
func (c *Cmder) Start(cmd *exec.Cmd) (
|
||||
stdoutLines, stderrLines <-chan string,
|
||||
waitError <-chan error, startErr error,
|
||||
@@ -38,6 +39,7 @@ func start(cmd execCmd) (stdoutLines, stderrLines <-chan string,
|
||||
if err != nil {
|
||||
_ = stdout.Close()
|
||||
<-stdoutDone
|
||||
close(stdoutLinesCh)
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
go streamToChannel(stderrReady, stderrDone, stderr, stderrLinesCh)
|
||||
@@ -45,9 +47,11 @@ func start(cmd execCmd) (stdoutLines, stderrLines <-chan string,
|
||||
err = cmd.Start()
|
||||
if err != nil {
|
||||
_ = stdout.Close()
|
||||
_ = stderr.Close()
|
||||
<-stdoutDone
|
||||
close(stdoutLinesCh)
|
||||
_ = stderr.Close()
|
||||
<-stderrDone
|
||||
close(stderrLinesCh)
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
|
||||
@@ -55,8 +59,10 @@ func start(cmd execCmd) (stdoutLines, stderrLines <-chan string,
|
||||
go func() {
|
||||
err := cmd.Wait()
|
||||
<-stdoutDone
|
||||
<-stderrDone
|
||||
close(stdoutLinesCh)
|
||||
_ = stdout.Close()
|
||||
<-stderrDone
|
||||
close(stderrLinesCh)
|
||||
_ = stderr.Close()
|
||||
waitErrorCh <- err
|
||||
}()
|
||||
|
||||
@@ -89,30 +89,48 @@ func Test_start(t *testing.T) {
|
||||
|
||||
require.NoError(t, err)
|
||||
|
||||
var stdoutIndex, stderrIndex int
|
||||
|
||||
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)
|
||||
collectAndCheckChannels(t, stdoutLines, stderrLines, waitError,
|
||||
testCase.stdout, testCase.stderr, testCase.waitErr)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -18,31 +18,39 @@ func (c *Cmder) RunAndLog(ctx context.Context, command string, logger Logger) (e
|
||||
return err
|
||||
}
|
||||
|
||||
streamCtx, streamCancel := context.WithCancel(context.Background())
|
||||
streamDone := make(chan struct{})
|
||||
go streamLines(streamCtx, streamDone, logger, stdout, stderr)
|
||||
go streamLines(streamDone, logger, stdout, stderr)
|
||||
|
||||
err = <-waitError
|
||||
streamCancel()
|
||||
<-streamDone
|
||||
return err
|
||||
}
|
||||
|
||||
func streamLines(ctx context.Context, done chan<- struct{},
|
||||
logger Logger, stdout, stderr <-chan string,
|
||||
func streamLines(done chan<- struct{}, logger Logger,
|
||||
stdout, stderr <-chan string,
|
||||
) {
|
||||
defer close(done)
|
||||
|
||||
var line string
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case line = <-stdout:
|
||||
logger.Info(line)
|
||||
case line = <-stderr:
|
||||
logger.Error(line)
|
||||
case line, ok := <-stdout:
|
||||
if ok {
|
||||
logger.Info(line)
|
||||
break
|
||||
}
|
||||
if stderr == nil {
|
||||
return
|
||||
}
|
||||
stdout = nil
|
||||
case line, ok := <-stderr:
|
||||
if ok {
|
||||
logger.Error(line)
|
||||
break
|
||||
}
|
||||
if stdout == nil {
|
||||
return
|
||||
}
|
||||
stderr = nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -29,19 +29,16 @@ func (r *Runner) Run(ctx context.Context, errCh chan<- error, ready chan<- struc
|
||||
return
|
||||
}
|
||||
|
||||
streamCtx, streamCancel := context.WithCancel(context.Background())
|
||||
streamDone := make(chan struct{})
|
||||
go streamLines(streamCtx, streamDone, r.logger,
|
||||
go streamLines(streamDone, r.logger,
|
||||
stdoutLines, stderrLines, ready)
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
<-waitError
|
||||
streamCancel()
|
||||
<-streamDone
|
||||
errCh <- ctx.Err()
|
||||
case err := <-waitError:
|
||||
streamCancel()
|
||||
<-streamDone
|
||||
errCh <- err
|
||||
}
|
||||
|
||||
@@ -1,26 +1,37 @@
|
||||
package openvpn
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func streamLines(ctx context.Context, done chan<- struct{},
|
||||
func streamLines(done chan<- struct{},
|
||||
logger Logger, stdout, stderr <-chan string,
|
||||
tunnelReady chan<- struct{},
|
||||
) {
|
||||
defer close(done)
|
||||
|
||||
var line string
|
||||
|
||||
for {
|
||||
var line string
|
||||
var ok bool
|
||||
errLine := false
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case line = <-stdout:
|
||||
case line = <-stderr:
|
||||
errLine = true
|
||||
case line, ok = <-stdout:
|
||||
if ok {
|
||||
break
|
||||
}
|
||||
if stderr == nil {
|
||||
return
|
||||
}
|
||||
stdout = nil
|
||||
case line, ok = <-stderr:
|
||||
if ok {
|
||||
errLine = true
|
||||
break
|
||||
}
|
||||
if stdout == nil {
|
||||
return
|
||||
}
|
||||
stderr = nil
|
||||
}
|
||||
line, level := processLogLine(line)
|
||||
if line == "" {
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"strings"
|
||||
|
||||
"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/models"
|
||||
"github.com/qdm12/gluetun/internal/provider/utils"
|
||||
@@ -65,7 +66,11 @@ func modifyConfig(lines []string, connection models.Connection,
|
||||
}
|
||||
|
||||
// Add values
|
||||
modified = append(modified, "proto "+connection.Protocol)
|
||||
protocol := 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, "dev "+settings.Interface)
|
||||
modified = append(modified, "mute-replay-warnings")
|
||||
|
||||
+74
-65
@@ -21,8 +21,8 @@ type server struct {
|
||||
listening atomic.Bool
|
||||
socksConnCtx context.Context //nolint:containedctx
|
||||
socksConnCancel context.CancelFunc
|
||||
done <-chan struct{}
|
||||
stopping atomic.Bool
|
||||
done <-chan error
|
||||
stopCh chan<- struct{}
|
||||
}
|
||||
|
||||
func newServer(settings Settings) *server {
|
||||
@@ -58,9 +58,11 @@ func (s *server) Start(ctx context.Context) (runErr <-chan error, err error) {
|
||||
ready := make(chan struct{})
|
||||
runErrCh := make(chan error)
|
||||
runErr = runErrCh
|
||||
done := make(chan struct{})
|
||||
done := make(chan error)
|
||||
s.done = done
|
||||
go s.runServer(ready, runErrCh, done)
|
||||
stop := make(chan struct{})
|
||||
s.stopCh = stop
|
||||
go s.runServer(ready, runErrCh, stop, done)
|
||||
select {
|
||||
case <-ready:
|
||||
case <-ctx.Done():
|
||||
@@ -71,78 +73,85 @@ func (s *server) Start(ctx context.Context) (runErr <-chan error, err error) {
|
||||
}
|
||||
|
||||
func (s *server) runServer(ready chan<- struct{},
|
||||
runErrCh chan<- error, done chan<- struct{},
|
||||
runErrCh chan<- error, stop <-chan struct{}, done chan<- error,
|
||||
) {
|
||||
close(ready)
|
||||
defer close(done)
|
||||
wg := new(sync.WaitGroup)
|
||||
defer wg.Wait()
|
||||
|
||||
wg.Go(func() {
|
||||
err := s.udpRouter.run(s.socksConnCtx)
|
||||
if err != nil {
|
||||
if !s.stopping.Load() {
|
||||
_ = s.stop()
|
||||
runErrCh <- fmt.Errorf("running UDP router: %w", err)
|
||||
}
|
||||
}
|
||||
})
|
||||
udpErrCh := make(chan error)
|
||||
go func() {
|
||||
udpErrCh <- s.udpRouter.run(s.socksConnCtx)
|
||||
}()
|
||||
|
||||
dialer := &net.Dialer{}
|
||||
for {
|
||||
connection, err := s.tcpListener.Accept()
|
||||
if err != nil {
|
||||
if !s.stopping.Load() {
|
||||
_ = s.stop()
|
||||
runErrCh <- fmt.Errorf("accepting connection: %w", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
wg.Add(1)
|
||||
go func(ctx context.Context, connection net.Conn,
|
||||
dialer *net.Dialer, wg *sync.WaitGroup,
|
||||
) {
|
||||
defer wg.Done()
|
||||
socksConn := &socksConn{
|
||||
dialer: dialer,
|
||||
username: s.username,
|
||||
password: s.password,
|
||||
clientConn: connection,
|
||||
udpRouter: s.udpRouter,
|
||||
logger: s.logger,
|
||||
}
|
||||
err := socksConn.run(ctx)
|
||||
tcpErrCh := make(chan error)
|
||||
go func() {
|
||||
var wg sync.WaitGroup
|
||||
defer wg.Wait()
|
||||
|
||||
dialer := &net.Dialer{}
|
||||
for {
|
||||
connection, err := s.tcpListener.Accept()
|
||||
if err != nil {
|
||||
s.logger.Infof("running socks connection: %s", err)
|
||||
s.socksConnCancel() // stop ongoing TCP socks connections - no impact on UDP
|
||||
tcpErrCh <- fmt.Errorf("accepting connection: %w", err)
|
||||
return
|
||||
}
|
||||
}(s.socksConnCtx, connection, dialer, wg)
|
||||
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 {
|
||||
case <-stop:
|
||||
s.listening.Store(false)
|
||||
var errs []error
|
||||
err := s.tcpListener.Close()
|
||||
if err != nil {
|
||||
errs = append(errs, fmt.Errorf("closing TCP listener: %w", err))
|
||||
}
|
||||
// stop ongoing TCP socks connections. This impacts the udpRouter run error when it is being closed.
|
||||
s.socksConnCancel()
|
||||
<-tcpErrCh // wait for TCP server to stop
|
||||
err = s.udpRouter.close()
|
||||
if err != nil {
|
||||
errs = append(errs, fmt.Errorf("closing UDP router: %w", err))
|
||||
}
|
||||
<-udpErrCh // wait for UDP router to stop
|
||||
if len(errs) > 0 {
|
||||
// Only write to the done channel if the [server.Stop] method is waiting to read from it
|
||||
done <- errors.Join(errs...)
|
||||
}
|
||||
// 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.
|
||||
case err := <-udpErrCh:
|
||||
_ = s.tcpListener.Close() // stop accepting new TCP connections
|
||||
s.socksConnCancel() // stop ongoing TCP socks connections
|
||||
<-tcpErrCh // wait for TCP server to stop
|
||||
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) {
|
||||
s.stopping.Store(true)
|
||||
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
|
||||
close(s.stopCh)
|
||||
return <-s.done
|
||||
}
|
||||
|
||||
func (s *server) listeningAddress() net.Addr {
|
||||
|
||||
@@ -0,0 +1,112 @@
|
||||
//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")
|
||||
}
|
||||
+20
-19
@@ -10,6 +10,7 @@ import (
|
||||
"net/netip"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
var (
|
||||
@@ -210,13 +211,6 @@ func (c *socksConn) handleUDPAssociateRequest(ctx context.Context,
|
||||
return fmt.Errorf("getting udp association addresses: %w", err)
|
||||
}
|
||||
|
||||
err = c.encodeSuccessResponse(c.clientConn, socksVersion, succeeded,
|
||||
bindAddrType, bindAddress, bindPort)
|
||||
if err != nil {
|
||||
c.encodeFailedResponse(c.clientConn, socksVersion, generalServerFailure)
|
||||
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)
|
||||
@@ -224,21 +218,28 @@ func (c *socksConn) handleUDPAssociateRequest(ctx context.Context,
|
||||
}
|
||||
defer c.udpRouter.unregisterAssociation(association)
|
||||
|
||||
err = c.encodeSuccessResponse(c.clientConn, socksVersion, succeeded,
|
||||
bindAddrType, bindAddress, bindPort)
|
||||
if err != nil {
|
||||
c.encodeFailedResponse(c.clientConn, socksVersion, generalServerFailure)
|
||||
return fmt.Errorf("writing successful %s response: %w", udpAssociate, err)
|
||||
}
|
||||
|
||||
associationCtx, associationCancel := context.WithCancel(ctx)
|
||||
defer associationCancel()
|
||||
|
||||
handlerDone := make(chan struct{})
|
||||
go func() {
|
||||
defer close(handlerDone)
|
||||
c.udpRouter.runAssociationHandler(associationCtx, association)
|
||||
}()
|
||||
var wg sync.WaitGroup
|
||||
|
||||
go func() {
|
||||
wg.Go(func() {
|
||||
c.udpRouter.runAssociationHandler(associationCtx, association)
|
||||
})
|
||||
|
||||
wg.Go(func() {
|
||||
_, _ = io.Copy(io.Discard, c.clientConn)
|
||||
associationCancel()
|
||||
}()
|
||||
})
|
||||
<-associationCtx.Done()
|
||||
<-handlerDone
|
||||
wg.Wait()
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -255,10 +256,10 @@ func udpAssociateExpectedClientEndpoint(request request) (expectedAddrPort netip
|
||||
}
|
||||
return netip.AddrPortFrom(netip.Addr{}, request.port), nil
|
||||
case domainName:
|
||||
if request.destination != "" || request.port != 0 {
|
||||
return netip.AddrPort{}, fmt.Errorf("domain name is not supported for UDP associate destination")
|
||||
}
|
||||
return netip.AddrPort{}, nil
|
||||
// For UDP associate, client endpoint matching is based on observed UDP source
|
||||
// address/port. A hostname is not directly matchable at this stage, so we
|
||||
// ignore the domain name request destination entirely.
|
||||
return netip.AddrPortFrom(netip.Addr{}, request.port), nil
|
||||
default:
|
||||
return netip.AddrPort{}, fmt.Errorf("address type %d is not supported", request.addressType)
|
||||
}
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/netip"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
@@ -703,6 +704,70 @@ 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) {
|
||||
t.Parallel()
|
||||
testCases := map[string]struct {
|
||||
|
||||
@@ -48,7 +48,7 @@ func newUDPRouter(ctx context.Context, address string, logger Logger) (router *u
|
||||
listener: listener,
|
||||
bufferPool: sync.Pool{
|
||||
New: func() any {
|
||||
return bytes.NewBuffer(make([]byte, pooledUDPPacketBufferCapacity))
|
||||
return bytes.NewBuffer(make([]byte, 0, pooledUDPPacketBufferCapacity))
|
||||
},
|
||||
},
|
||||
nextAssociationID: 1,
|
||||
@@ -76,7 +76,7 @@ func (r *udpRouter) registerAssociation(controlConn net.Conn, expectedAddrPort n
|
||||
r.mutex.Lock()
|
||||
defer r.mutex.Unlock()
|
||||
|
||||
const udpPacketChannelBuffer = 2
|
||||
const udpPacketChannelBuffer = 64
|
||||
associationID := r.nextAssociationID
|
||||
r.nextAssociationID++
|
||||
|
||||
@@ -111,10 +111,6 @@ func (r *udpRouter) unregisterAssociation(association udpAssociation) {
|
||||
delete(r.clientAddrPortToAssociation, clientAddrPort)
|
||||
}
|
||||
|
||||
if association.clientAddrPort.IsValid() {
|
||||
delete(r.clientAddrPortToAssociation, association.clientAddrPort)
|
||||
}
|
||||
|
||||
pendingAssociations := r.clientIPToPendingAssociations[association.controlConnAddr]
|
||||
for i, pendingAssociation := range pendingAssociations {
|
||||
if pendingAssociation.id == association.id {
|
||||
@@ -163,8 +159,9 @@ func (r *udpRouter) run(ctx context.Context) error {
|
||||
|
||||
func (r *udpRouter) routePacket(sourceAddrPort netip.AddrPort, packet *bytes.Buffer) error {
|
||||
r.mutex.Lock()
|
||||
defer r.mutex.Unlock()
|
||||
association, packetFromClient := r.findClientAssociation(sourceAddrPort)
|
||||
r.mutex.Unlock()
|
||||
|
||||
if !packetFromClient {
|
||||
r.bufferPool.Put(packet)
|
||||
return nil
|
||||
@@ -334,6 +331,23 @@ func (r *udpRouter) writeClientPacketToDestination(ctx context.Context,
|
||||
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)
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolving destination UDP address: %w", err)
|
||||
|
||||
@@ -19,6 +19,8 @@ import (
|
||||
)
|
||||
|
||||
func Test_udpRouter_ResolveGithubFromCloudflareDNS(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := t.Context()
|
||||
var cancel context.CancelFunc
|
||||
deadline, hasDeadline := ctx.Deadline()
|
||||
@@ -107,7 +109,7 @@ func Test_udpRouter_ResolveGithubFromCloudflareDNS(t *testing.T) {
|
||||
assert.NoError(t, err, "closing client UDP connection")
|
||||
})
|
||||
|
||||
queryID := uint16(rand.Uint())
|
||||
queryID := uint16(rand.Uint32()) //nolint:gosec
|
||||
dnsRequest := &dns.Msg{
|
||||
MsgHdr: dns.MsgHdr{
|
||||
Id: queryID,
|
||||
@@ -156,7 +158,7 @@ func Test_udpRouter_ResolveGithubFromCloudflareDNS(t *testing.T) {
|
||||
assert.Equal(t, dns.RcodeSuccess, dnsResponse.Rcode)
|
||||
require.NotEmpty(t, dnsResponse.Question)
|
||||
assert.Equal(t, dns.Fqdn("github.com"), dnsResponse.Question[0].Name)
|
||||
assert.Equal(t, uint16(dns.TypeA), dnsResponse.Question[0].Qtype)
|
||||
assert.Equal(t, dns.TypeA, dnsResponse.Question[0].Qtype)
|
||||
assert.NotEmpty(t, dnsResponse.Answer)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user