Squashed commit of the following:

commit 58d2c9d3468aa1c04f8cb60f0680186b6330804c
Author: xtaci <daniel820313@gmail.com>
Date:   Thu Dec 26 17:40:36 2019 +0800

    combine smux v1 & v2
This commit is contained in:
xtaci
2019-12-26 17:45:07 +08:00
parent 56b45e84f7
commit 34bc48fd64
5 changed files with 52 additions and 145 deletions
+32 -61
View File
@@ -17,7 +17,6 @@ import (
kcp "github.com/xtaci/kcp-go"
"github.com/xtaci/kcptun/generic"
"github.com/xtaci/smux"
smuxv2 "github.com/xtaci/smux/v2"
)
const (
@@ -33,14 +32,14 @@ const (
var VERSION = "SELFBUILD"
// handleClient aggregates connection p1 on mux with 'writeLock'
func handleClient(mux generic.Mux, p1 net.Conn, ctrl *generic.CopyControl, quiet bool) {
func handleClient(session *smux.Session, p1 net.Conn, ctrl *generic.CopyControl, quiet bool) {
logln := func(v ...interface{}) {
if !quiet {
log.Println(v...)
}
}
defer p1.Close()
p2, err := mux.Open()
p2, err := session.OpenStream()
if err != nil {
logln(err)
return
@@ -48,27 +47,15 @@ func handleClient(mux generic.Mux, p1 net.Conn, ctrl *generic.CopyControl, quiet
defer p2.Close()
if s2, ok := p2.(generic.Stream); ok {
logln("stream opened", "in:", p1.RemoteAddr(), "out:", fmt.Sprint(s2.RemoteAddr(), "(", s2.ID(), ")"))
defer logln("stream closed", "in:", p1.RemoteAddr(), "out:", fmt.Sprint(s2.RemoteAddr(), "(", s2.ID(), ")"))
}
logln("stream opened", "in:", p1.RemoteAddr(), "out:", fmt.Sprint(p2.RemoteAddr(), "(", p2.ID(), ")"))
defer logln("stream closed", "in:", p1.RemoteAddr(), "out:", fmt.Sprint(p2.RemoteAddr(), "(", p2.ID(), ")"))
// start tunnel & wait for tunnel termination
streamCopy := func(dst io.Writer, src io.ReadCloser) {
if _, err := generic.Copy(dst, src, ctrl); err != nil {
if s2, ok := p2.(generic.Stream); ok {
// verbose error handling
cause := err
if e, ok := err.(interface{ Cause() error }); ok {
cause = e.Cause()
}
switch cause {
case smux.ErrInvalidProtocol:
log.Println("smux version:1", err, "in:", p1.RemoteAddr(), "out:", fmt.Sprint(s2.RemoteAddr(), "(", s2.ID(), ")"))
case smuxv2.ErrInvalidProtocol:
log.Println("smux version:2", err, "in:", p1.RemoteAddr(), "out:", fmt.Sprint(s2.RemoteAddr(), "(", s2.ID(), ")"))
}
// report protocol error
if err == smux.ErrInvalidProtocol {
log.Println("smux", err, "in:", p1.RemoteAddr(), "out:", fmt.Sprint(p2.RemoteAddr(), "(", p2.ID(), ")"))
}
}
p1.Close()
@@ -377,7 +364,7 @@ func main() {
block, _ = kcp.NewAESBlockCrypt(pass)
}
createConn := func() (generic.Mux, error) {
createConn := func() (*smux.Session, error) {
kcpconn, err := dial(&config, block)
if err != nil {
return nil, errors.Wrap(err, "dial()")
@@ -399,47 +386,31 @@ func main() {
log.Println("SetWriteBuffer:", err)
}
log.Println("smux version:", config.SmuxVer, "on connection:", kcpconn.LocalAddr(), "->", kcpconn.RemoteAddr())
switch config.SmuxVer {
case 1:
smuxConfig := smux.DefaultConfig()
smuxConfig.MaxReceiveBuffer = config.SmuxBuf
smuxConfig.KeepAliveInterval = time.Duration(config.KeepAlive) * time.Second
smuxConfig := smux.DefaultConfig()
smuxConfig.Version = config.SmuxVer
smuxConfig.MaxReceiveBuffer = config.SmuxBuf
smuxConfig.MaxStreamBuffer = config.StreamBuf
smuxConfig.KeepAliveInterval = time.Duration(config.KeepAlive) * time.Second
// stream multiplex
var session *smux.Session
if config.NoComp {
session, err = smux.Client(kcpconn, smuxConfig)
} else {
session, err = smux.Client(generic.NewCompStream(kcpconn), smuxConfig)
}
if err != nil {
return nil, errors.Wrap(err, "createConn()")
}
return session, nil
case 2:
smuxConfig := smuxv2.DefaultConfig()
smuxConfig.MaxReceiveBuffer = config.SmuxBuf
smuxConfig.MaxStreamBuffer = config.StreamBuf
smuxConfig.KeepAliveInterval = time.Duration(config.KeepAlive) * time.Second
// stream multiplex
var session *smuxv2.Session
if config.NoComp {
session, err = smuxv2.Client(kcpconn, smuxConfig)
} else {
session, err = smuxv2.Client(generic.NewCompStream(kcpconn), smuxConfig)
}
if err != nil {
return nil, errors.Wrap(err, "createConn()")
}
return session, nil
default:
panic("incorrect smux version")
if err := smux.VerifyConfig(smuxConfig); err != nil {
log.Fatalf("%+v", err)
}
// stream multiplex
var session *smux.Session
if config.NoComp {
session, err = smux.Client(kcpconn, smuxConfig)
} else {
session, err = smux.Client(generic.NewCompStream(kcpconn), smuxConfig)
}
if err != nil {
return nil, errors.Wrap(err, "createConn()")
}
return session, nil
}
// wait until a connection is ready
waitConn := func() generic.Mux {
waitConn := func() *smux.Session {
for {
if session, err := createConn(); err == nil {
return session
@@ -452,7 +423,7 @@ func main() {
numconn := uint16(config.Conn)
muxes := make([]struct {
session generic.Mux
session *smux.Session
ttl time.Time
ctrl *generic.CopyControl // for control of memory in copying
}, numconn)
@@ -463,7 +434,7 @@ func main() {
muxes[k].ctrl = &generic.CopyControl{Buffer: make([]byte, bufSize)}
}
chScavenger := make(chan generic.Mux, 128)
chScavenger := make(chan *smux.Session, 128)
go scavenger(chScavenger, config.ScavengeTTL)
go generic.SnmpLogger(config.SnmpLog, config.SnmpPeriod)
rr := uint16(0)
@@ -490,11 +461,11 @@ func main() {
}
type scavengeSession struct {
session generic.Mux
session *smux.Session
ts time.Time
}
func scavenger(ch chan generic.Mux, ttl int) {
func scavenger(ch chan *smux.Session, ttl int) {
ticker := time.NewTicker(time.Second)
defer ticker.Stop()
var sessionList []scavengeSession
-21
View File
@@ -1,21 +0,0 @@
package generic
import (
"io"
"net"
)
type Mux interface {
Open() (io.ReadWriteCloser, error)
Accept() (io.ReadWriteCloser, error)
IsClosed() bool
NumStreams() int
RemoteAddr() net.Addr
Close() error
}
type Stream interface {
io.ReadWriteCloser
ID() uint32
RemoteAddr() net.Addr
}
+1 -2
View File
@@ -13,8 +13,7 @@ require (
github.com/urfave/cli v1.21.0
github.com/xtaci/kcp-go v5.4.20+incompatible
github.com/xtaci/lossyconn v0.0.0-20190602105132-8df528c0c9ae // indirect
github.com/xtaci/smux v1.4.8
github.com/xtaci/smux/v2 v2.0.18
github.com/xtaci/smux v1.5.4
github.com/xtaci/tcpraw v1.2.25
golang.org/x/crypto v0.0.0-20191206172530-e9b2fee46413
golang.org/x/net v0.0.0-20191209160850-c0dbc17a3553 // indirect
+2 -12
View File
@@ -23,18 +23,8 @@ github.com/xtaci/kcp-go v5.4.20+incompatible h1:TN1uey3Raw0sTz0Fg8GkfM0uH3YwzhnZ
github.com/xtaci/kcp-go v5.4.20+incompatible/go.mod h1:bN6vIwHQbfHaHtFpEssmWsN45a+AZwO7eyRCmEIbtvE=
github.com/xtaci/lossyconn v0.0.0-20190602105132-8df528c0c9ae h1:J0GxkO96kL4WF+AIT3M4mfUVinOCPgf2uUWYFUzN0sM=
github.com/xtaci/lossyconn v0.0.0-20190602105132-8df528c0c9ae/go.mod h1:gXtu8J62kEgmN++bm9BVICuT/e8yiLI2KFobd/TRFsE=
github.com/xtaci/smux v1.4.6 h1:p9e/qj3Bj0zUT8qJWdmAZfmx5lOcZh0vLL0bQ8jnA7M=
github.com/xtaci/smux v1.4.6/go.mod h1:LuA3S0xssf4fmGRJ7ow3EehgmDUzib4EcobaFNKvlMA=
github.com/xtaci/smux v1.4.7 h1:ew5LGDZJWBuhwDq3cAqCQx4jRLKm54Y0SkObNQJJKu8=
github.com/xtaci/smux v1.4.7/go.mod h1:LuA3S0xssf4fmGRJ7ow3EehgmDUzib4EcobaFNKvlMA=
github.com/xtaci/smux v1.4.8 h1:QzYkAtRqlBqRJqDVOt/0Qj9O91YdxjEZppQt5aJrOAU=
github.com/xtaci/smux v1.4.8/go.mod h1:LuA3S0xssf4fmGRJ7ow3EehgmDUzib4EcobaFNKvlMA=
github.com/xtaci/smux/v2 v2.0.16 h1:2pGGbkFKTaMHIctYaovpwRpgdwWYy/6ZPaOQo00VW08=
github.com/xtaci/smux/v2 v2.0.16/go.mod h1:Iqy5a3Gax2p7WCKHOHkSNo/COthNFXd3/vqrcKNtzqI=
github.com/xtaci/smux/v2 v2.0.17 h1:baE6Dek0lkTZjofAFZxrd+LYOxN92GTxBRFRR6nKcXU=
github.com/xtaci/smux/v2 v2.0.17/go.mod h1:Iqy5a3Gax2p7WCKHOHkSNo/COthNFXd3/vqrcKNtzqI=
github.com/xtaci/smux/v2 v2.0.18 h1:Sa+W8IMR0dv3Tj9XxDUqovEFYGtU8Upydmmpez7IZgA=
github.com/xtaci/smux/v2 v2.0.18/go.mod h1:Iqy5a3Gax2p7WCKHOHkSNo/COthNFXd3/vqrcKNtzqI=
github.com/xtaci/smux v1.5.4 h1:XpNDiuRLq9qDfoe7U2ROl363fvD42DCYmJaBI3i3TKM=
github.com/xtaci/smux v1.5.4/go.mod h1:OMlQbT5vcgl2gb49mFkYo6SMf+zP3rcjcwQz7ZU7IGY=
github.com/xtaci/tcpraw v1.2.25 h1:VDlqo0op17JeXBM6e2G9ocCNLOJcw9mZbobMbJjo0vk=
github.com/xtaci/tcpraw v1.2.25/go.mod h1:dKyZ2V75s0cZ7cbgJYdxPvms7af0joIeOyx1GgJQbLk=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
+17 -49
View File
@@ -19,7 +19,6 @@ import (
kcp "github.com/xtaci/kcp-go"
"github.com/xtaci/kcptun/generic"
"github.com/xtaci/smux"
smuxv2 "github.com/xtaci/smux/v2"
"github.com/xtaci/tcpraw"
)
@@ -45,48 +44,30 @@ func handleMux(conn net.Conn, config *Config) {
log.Println("smux version:", config.SmuxVer, "on connection:", conn.LocalAddr(), "->", conn.RemoteAddr())
// stream multiplex
var muxer generic.Mux
switch config.SmuxVer {
case 1:
smuxConfig := smux.DefaultConfig()
smuxConfig.MaxReceiveBuffer = config.SmuxBuf
smuxConfig.KeepAliveInterval = time.Duration(config.KeepAlive) * time.Second
smuxConfig := smux.DefaultConfig()
smuxConfig.Version = config.SmuxVer
smuxConfig.MaxReceiveBuffer = config.SmuxBuf
smuxConfig.MaxStreamBuffer = config.StreamBuf
smuxConfig.KeepAliveInterval = time.Duration(config.KeepAlive) * time.Second
mux, err := smux.Server(conn, smuxConfig)
if err != nil {
log.Println(err)
return
}
defer mux.Close()
muxer = mux
case 2:
smuxConfig := smuxv2.DefaultConfig()
smuxConfig.MaxReceiveBuffer = config.SmuxBuf
smuxConfig.MaxStreamBuffer = config.StreamBuf
smuxConfig.KeepAliveInterval = time.Duration(config.KeepAlive) * time.Second
mux, err := smuxv2.Server(conn, smuxConfig)
if err != nil {
log.Println(err)
return
}
defer mux.Close()
muxer = mux
default:
panic("incorrect smux version")
mux, err := smux.Server(conn, smuxConfig)
if err != nil {
log.Println(err)
return
}
defer mux.Close()
// copy to stream control
copyControl := &generic.CopyControl{Buffer: make([]byte, bufSize)}
for {
stream, err := muxer.Accept()
stream, err := mux.AcceptStream()
if err != nil {
log.Println(err)
return
}
go func(p1 io.ReadWriteCloser) {
go func(p1 *smux.Stream) {
var p2 net.Conn
var err error
if !isUnix {
@@ -105,7 +86,7 @@ func handleMux(conn net.Conn, config *Config) {
}
}
func handleClient(p1 io.ReadWriteCloser, p2 net.Conn, ctrl *generic.CopyControl, quiet bool) {
func handleClient(p1 *smux.Stream, p2 net.Conn, ctrl *generic.CopyControl, quiet bool) {
logln := func(v ...interface{}) {
if !quiet {
log.Println(v...)
@@ -115,27 +96,14 @@ func handleClient(p1 io.ReadWriteCloser, p2 net.Conn, ctrl *generic.CopyControl,
defer p1.Close()
defer p2.Close()
if s1, ok := p1.(generic.Stream); ok {
logln("stream opened", "in:", fmt.Sprint(s1.RemoteAddr(), "(", s1.ID(), ")"), "out:", p2.RemoteAddr())
defer logln("stream closed", "in:", fmt.Sprint(s1.RemoteAddr(), "(", s1.ID(), ")"), "out:", p2.RemoteAddr())
}
logln("stream opened", "in:", fmt.Sprint(p1.RemoteAddr(), "(", p1.ID(), ")"), "out:", p2.RemoteAddr())
defer logln("stream closed", "in:", fmt.Sprint(p1.RemoteAddr(), "(", p1.ID(), ")"), "out:", p2.RemoteAddr())
// start tunnel & wait for tunnel termination
streamCopy := func(dst io.Writer, src io.ReadCloser) {
if _, err := generic.Copy(dst, src, ctrl); err != nil {
if s1, ok := p1.(generic.Stream); ok {
// verbose error handling
cause := err
if e, ok := err.(interface{ Cause() error }); ok {
cause = e.Cause()
}
switch cause {
case smux.ErrInvalidProtocol:
log.Println("smux version:1", err, "in:", fmt.Sprint(s1.RemoteAddr(), "(", s1.ID(), ")"), "out:", p2.RemoteAddr())
case smuxv2.ErrInvalidProtocol:
log.Println("smux version:2", err, "in:", fmt.Sprint(s1.RemoteAddr(), "(", s1.ID(), ")"), "out:", p2.RemoteAddr())
}
if err == smux.ErrInvalidProtocol {
log.Println("smux", err, "in:", fmt.Sprint(p1.RemoteAddr(), "(", p1.ID(), ")"), "out:", p2.RemoteAddr())
}
}
p1.Close()