diff --git a/tls.go b/tls.go index c6862b9..afb99eb 100644 --- a/tls.go +++ b/tls.go @@ -40,54 +40,31 @@ import ( "golang.org/x/crypto/hkdf" ) -type ReaderConn struct { - Conn net.Conn - Reader *bytes.Reader - Written int - Closed bool +type WeakConn struct { + net.Conn } -func (c *ReaderConn) Read(b []byte) (int, error) { - if c.Closed { - return 0, errors.New("Closed") - } - n, err := c.Reader.Read(b) - if err == io.EOF { - return n, errors.New("io.EOF") // prevent looping - } - return n, err +func (c *WeakConn) Read(b []byte) (int, error) { + return 0, fmt.Errorf("Read(%v)", len(b)) } -func (c *ReaderConn) Write(b []byte) (int, error) { - if c.Closed { - return 0, errors.New("Closed") - } - c.Written += len(b) - return len(b), nil +func (c *WeakConn) Write(b []byte) (int, error) { + return 0, fmt.Errorf("Write(%v)", len(b)) } -func (c *ReaderConn) Close() error { - c.Closed = true +func (c *WeakConn) Close() error { + return fmt.Errorf("Close()") +} + +func (c *WeakConn) SetDeadline(t time.Time) error { return nil } -func (c *ReaderConn) LocalAddr() net.Addr { - return c.Conn.LocalAddr() -} - -func (c *ReaderConn) RemoteAddr() net.Addr { - return c.Conn.RemoteAddr() -} - -func (c *ReaderConn) SetDeadline(t time.Time) error { +func (c *WeakConn) SetReadDeadline(t time.Time) error { return nil } -func (c *ReaderConn) SetReadDeadline(t time.Time) error { - return nil -} - -func (c *ReaderConn) SetWriteDeadline(t time.Time) error { +func (c *WeakConn) SetWriteDeadline(t time.Time) error { return nil } @@ -175,7 +152,7 @@ func Server(ctx context.Context, conn net.Conn, config *Config) (*Conn, error) { done = true break } - if copying || len(c2sSaved) > size || len(s2cSaved) > 0 { // follow; too long; unexpected + if len(c2sSaved) > size || copying { // too long; follow break } if clientHelloLen == 0 && len(c2sSaved) > recordHeaderLen { @@ -191,19 +168,12 @@ func Server(ctx context.Context, conn net.Conn, config *Config) (*Conn, error) { mutex.Unlock() continue } - if len(c2sSaved) > clientHelloLen { // unexpected - break - } - readerConn := &ReaderConn{ - Conn: conn, - Reader: bytes.NewReader(c2sSaved), - } hs.c = &Conn{ - conn: readerConn, - config: config, + conn: &WeakConn{conn}, + config: config, + rawInput: *bytes.NewBuffer(c2sSaved), } - hs.clientHello, err = hs.c.readClientHello(context.Background()) - if err != nil || readerConn.Reader.Len() > 0 || readerConn.Written > 0 || readerConn.Closed { + if hs.clientHello, err = hs.c.readClientHello(context.Background()); err != nil { break } if hs.c.vers != VersionTLS13 || !config.ServerNames[hs.clientHello.serverName] { @@ -260,11 +230,8 @@ func Server(ctx context.Context, conn net.Conn, config *Config) (*Conn, error) { } break } - if done { - mutex.Unlock() - } else { - copying = true - mutex.Unlock() + mutex.Unlock() + if !done { io.CopyBuffer(target, underlying, buf) } waitGroup.Done() @@ -289,15 +256,11 @@ func Server(ctx context.Context, conn net.Conn, config *Config) (*Conn, error) { mutex.Lock() s2cSaved = append(s2cSaved, buf[:n]...) if hs.c == nil || hs.c.conn != conn { + copying = true if _, err = conn.Write(buf[:n]); err != nil { done = true - break } - if copying || len(s2cSaved) > size { // follow; too long - break - } - mutex.Unlock() - continue + break } done = true // special if len(s2cSaved) > size { @@ -386,11 +349,8 @@ func Server(ctx context.Context, conn net.Conn, config *Config) (*Conn, error) { handled = true break } - if done { - mutex.Unlock() - } else { - copying = true - mutex.Unlock() + mutex.Unlock() + if !done { io.CopyBuffer(underlying, target, buf) } waitGroup.Done()