From a77d04567a0b725ac2ceabada2d4dda0224f3c09 Mon Sep 17 00:00:00 2001 From: "rui.zheng" Date: Fri, 25 Sep 2015 17:56:52 +0800 Subject: [PATCH] add Selector interface --- conn.go | 55 ++++++++++++++++++++++++++----------------------------- 1 file changed, 26 insertions(+), 29 deletions(-) diff --git a/conn.go b/conn.go index 1d574d5..d20ca2b 100644 --- a/conn.go +++ b/conn.go @@ -8,19 +8,18 @@ import ( "time" ) -type Config struct { - Methods []uint8 - SelectMethod func(methods ...uint8) uint8 - MethodSelected func(method uint8, conn net.Conn) (net.Conn, error) -} - -func defaultConfig() *Config { - return &Config{} +type Selector interface { + // return supported methods + Methods() []uint8 + // select method + Select(methods ...uint8) (method uint8) + // on method selected + OnSelected(method uint8, conn net.Conn) (net.Conn, error) } type Conn struct { c net.Conn - config *Config + selector Selector method uint8 isClient bool handshaked bool @@ -28,18 +27,18 @@ type Conn struct { handshakeErr error } -func ClientConn(conn net.Conn, config *Config) *Conn { +func ClientConn(conn net.Conn, selector Selector) *Conn { return &Conn{ c: conn, - config: config, + selector: selector, isClient: true, } } -func ServerConn(conn net.Conn, config *Config) *Conn { +func ServerConn(conn net.Conn, selector Selector) *Conn { return &Conn{ - c: conn, - config: config, + c: conn, + selector: selector, } } @@ -64,11 +63,13 @@ func (conn *Conn) Handleshake() error { } func (conn *Conn) clientHandshake() error { - if conn.config == nil { - conn.config = defaultConfig() - } + var methods []uint8 + var nm int - nm := len(conn.config.Methods) + if conn.selector != nil { + methods = conn.selector.Methods() + } + nm = len(methods) if nm == 0 { nm = 1 } @@ -76,7 +77,7 @@ func (conn *Conn) clientHandshake() error { b := make([]byte, 2+nm) b[0] = Ver5 b[1] = uint8(nm) - copy(b[2:], conn.config.Methods) + copy(b[2:], methods) if _, err := conn.c.Write(b); err != nil { return err @@ -90,8 +91,8 @@ func (conn *Conn) clientHandshake() error { return ErrBadVersion } - if conn.config.MethodSelected != nil { - c, err := conn.config.MethodSelected(b[1], conn.c) + if conn.selector != nil { + c, err := conn.selector.OnSelected(b[1], conn.c) if err != nil { return err } @@ -104,26 +105,22 @@ func (conn *Conn) clientHandshake() error { } func (conn *Conn) serverHandshake() error { - if conn.config == nil { - conn.config = defaultConfig() - } - methods, err := ReadMethods(conn.c) if err != nil { return err } method := MethodNoAuth - if conn.config.SelectMethod != nil { - method = conn.config.SelectMethod(methods...) + if conn.selector != nil { + method = conn.selector.Select(methods...) } if _, err := conn.c.Write([]byte{Ver5, method}); err != nil { return err } - if conn.config.MethodSelected != nil { - c, err := conn.config.MethodSelected(method, conn.c) + if conn.selector != nil { + c, err := conn.selector.OnSelected(method, conn.c) if err != nil { return err }