Skip to content

Commit 5bf5d98

Browse files
fix(tun): close TUN gracefully on client shutdown
1 parent 619a6f8 commit 5bf5d98

2 files changed

Lines changed: 56 additions & 11 deletions

File tree

app/cmd/client.go

Lines changed: 22 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -879,6 +879,10 @@ func runClient(v *viper.Viper) {
879879

880880
// Register modes
881881
var runner clientModeRunner
882+
var tunDone <-chan struct{}
883+
884+
tunCtx, cancelTun := context.WithCancel(context.Background())
885+
defer cancelTun()
882886
if config.SOCKS5 != nil {
883887
runner.Add("SOCKS5 server", func() error {
884888
return clientSOCKS5(*config.SOCKS5, c)
@@ -915,9 +919,14 @@ func runClient(v *viper.Viper) {
915919
})
916920
}
917921
if config.TUN != nil {
922+
done := make(chan struct{})
923+
tunDone = done
924+
918925
runner.Add("TUN", func() error {
919-
return clientTUN(*config.TUN, c)
926+
defer close(done)
927+
return clientTUN(tunCtx, *config.TUN, c)
920928
})
929+
921930
}
922931

923932
signalChan := make(chan os.Signal, 1)
@@ -932,11 +941,20 @@ func runClient(v *viper.Viper) {
932941
select {
933942
case <-signalChan:
934943
logger.Info("received signal, shutting down gracefully")
944+
cancelTun()
945+
if tunDone != nil {
946+
<-tunDone
947+
}
935948
case r := <-runnerChan:
936949
if r.OK {
937950
logger.Info(r.Msg)
938951
} else {
939-
_ = c.Close() // Close the client here as Fatal will exit the program without running defer
952+
cancelTun()
953+
954+
if tunDone != nil {
955+
<-tunDone
956+
}
957+
_ = c.Close()
940958
if r.Err != nil {
941959
logger.Fatal(r.Msg, zap.Error(r.Err))
942960
} else {
@@ -1149,7 +1167,7 @@ func clientTCPRedirect(config tcpRedirectConfig, c client.Client) error {
11491167
return p.ListenAndServe(laddr)
11501168
}
11511169

1152-
func clientTUN(config tunConfig, c client.Client) error {
1170+
func clientTUN(ctx context.Context, config tunConfig, c client.Client) error {
11531171
supportedPlatforms := []string{"linux", "darwin", "windows", "android"}
11541172
if !slices.Contains(supportedPlatforms, runtime.GOOS) {
11551173
logger.Error("TUN is not supported on this platform", zap.String("platform", runtime.GOOS))
@@ -1232,7 +1250,7 @@ func clientTUN(config tunConfig, c client.Client) error {
12321250
}
12331251
}
12341252
logger.Info("TUN listening", zap.String("interface", config.Name))
1235-
return server.Serve()
1253+
return server.Serve(ctx)
12361254
}
12371255

12381256
// parseServerAddrString parses server address string.

app/internal/tun/server.go

Lines changed: 34 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -2,10 +2,12 @@ package tun
22

33
import (
44
"context"
5+
"errors"
56
"fmt"
67
"io"
78
"net"
89
"net/netip"
10+
"sync"
911

1012
tun "github.com/apernet/sing-tun"
1113
"github.com/sagernet/sing/common/buf"
@@ -48,7 +50,7 @@ type EventLogger interface {
4850
UDPError(addr string, err error)
4951
}
5052

51-
func (s *Server) Serve() error {
53+
func (s *Server) Serve(ctx context.Context) error {
5254
if !isIPv6Supported() {
5355
s.Logger.Warn("tun-pre-check", zap.String("msg", "IPv6 is not supported or enabled on this system, TUN device is created without IPv6 support."))
5456
s.Inet6Address = nil
@@ -74,14 +76,23 @@ func (s *Server) Serve() error {
7476
if err != nil {
7577
return fmt.Errorf("failed to create tun interface: %w", err)
7678
}
77-
defer tunIf.Close()
79+
var closeTunOnce sync.Once
7880

81+
closeTun := func() {
82+
closeTunOnce.Do(func() {
83+
_ = tunIf.Close()
84+
})
85+
}
86+
defer closeTun()
7987
tunStack, err := tun.NewSystem(tun.StackOptions{
80-
Context: context.Background(),
88+
Context: ctx,
8189
Tun: tunIf,
8290
TunOptions: tunOpts,
8391
UDPTimeout: s.Timeout,
84-
Handler: &tunHandler{s},
92+
Handler: &tunHandler{
93+
Server: s,
94+
shutdownCtx: ctx,
95+
},
8596
Logger: &singLogger{
8697
tag: "tun-stack",
8798
zapLogger: s.Logger,
@@ -93,11 +104,25 @@ func (s *Server) Serve() error {
93104
return fmt.Errorf("failed to create tun stack: %w", err)
94105
}
95106
defer tunStack.Close()
96-
return tunStack.(tun.StackRunner).Run()
107+
108+
stopClose := context.AfterFunc(ctx, closeTun)
109+
defer stopClose()
110+
err = tunStack.(tun.StackRunner).Run()
111+
if ctx.Err() != nil {
112+
return nil
113+
}
114+
115+
return err
97116
}
98117

99118
type tunHandler struct {
100119
*Server
120+
shutdownCtx context.Context
121+
}
122+
123+
func (t *tunHandler) isShutdownCancellation(err error) bool {
124+
return errors.Is(err, context.Canceled) &&
125+
errors.Is(t.shutdownCtx.Err(), context.Canceled)
101126
}
102127

103128
var _ tun.Handler = (*tunHandler)(nil)
@@ -110,10 +135,11 @@ func (t *tunHandler) NewConnection(ctx context.Context, conn net.Conn, m metadat
110135
}
111136
var closeErr error
112137
defer func() {
113-
if t.EventLogger != nil {
138+
if t.EventLogger != nil && !t.isShutdownCancellation(closeErr) {
114139
t.EventLogger.TCPError(addr, reqAddr, closeErr)
115140
}
116141
}()
142+
117143
rc, err := t.HyClient.TCP(reqAddr)
118144
if err != nil {
119145
closeErr = err
@@ -147,10 +173,11 @@ func (t *tunHandler) NewPacketConnection(ctx context.Context, conn network.Packe
147173
}
148174
var closeErr error
149175
defer func() {
150-
if t.EventLogger != nil {
176+
if t.EventLogger != nil && !t.isShutdownCancellation(closeErr) {
151177
t.EventLogger.UDPError(addr, closeErr)
152178
}
153179
}()
180+
154181
rc, err := t.HyClient.UDP()
155182
if err != nil {
156183
closeErr = err

0 commit comments

Comments
 (0)