mirror of
https://github.com/lwch/natpass.git
synced 2024-04-21 12:41:54 +00:00
1. 增加代码注释
2. 增加通用encoding层
This commit is contained in:
@@ -0,0 +1,21 @@
|
||||
package encoding
|
||||
|
||||
import "io"
|
||||
|
||||
// Codec format data to []byte, decode data from []byte
|
||||
type Codec interface {
|
||||
// Marshal format data to []byte
|
||||
Marshal(interface{}) ([]byte, error)
|
||||
// Unmarshal decode data from []byte
|
||||
Unmarshal([]byte, interface{}) error
|
||||
}
|
||||
|
||||
// Compressor compressor interface
|
||||
type Compressor interface {
|
||||
// Compress get compress writer
|
||||
Compress(io.Writer) (io.WriteCloser, error)
|
||||
// Decompress get decompress reader
|
||||
Decompress(io.Reader) (io.ReadCloser, error)
|
||||
// SetLevel set compress level
|
||||
SetLevel(int) error
|
||||
}
|
||||
@@ -0,0 +1,88 @@
|
||||
package gzip
|
||||
|
||||
import (
|
||||
"compress/gzip"
|
||||
"fmt"
|
||||
"io"
|
||||
"sync"
|
||||
|
||||
"github.com/lwch/natpass/code/network/encoding"
|
||||
"github.com/lwch/runtime"
|
||||
)
|
||||
|
||||
type writer struct {
|
||||
*gzip.Writer
|
||||
pool *sync.Pool
|
||||
}
|
||||
|
||||
// Close close write and put writer to pool
|
||||
func (w *writer) Close() error {
|
||||
w.pool.Put(w)
|
||||
return w.Writer.Close()
|
||||
}
|
||||
|
||||
type reader struct {
|
||||
*gzip.Reader
|
||||
pool *sync.Pool
|
||||
}
|
||||
|
||||
// Close close reader and put reader to pool
|
||||
func (r *reader) Close() error {
|
||||
r.pool.Put(r)
|
||||
return r.Reader.Close()
|
||||
}
|
||||
|
||||
type compressor struct {
|
||||
level int
|
||||
poolWriter [gzip.BestCompression]sync.Pool
|
||||
poolReader sync.Pool
|
||||
}
|
||||
|
||||
// New create compressor
|
||||
func New(level ...int) (encoding.Compressor, error) {
|
||||
if len(level) > 0 {
|
||||
if level[0] < 0 || level[0] > gzip.BestCompression {
|
||||
return nil, fmt.Errorf("invalid gzip compress level: %d", level[0])
|
||||
}
|
||||
} else {
|
||||
level = append(level, 6)
|
||||
}
|
||||
ret := new(compressor)
|
||||
ret.level = level[0]
|
||||
for i := 0; i < gzip.BestCompression; i++ {
|
||||
ret.poolWriter[i].New = func() interface{} {
|
||||
w, err := gzip.NewWriterLevel(io.Discard, i)
|
||||
runtime.Assert(err)
|
||||
return &writer{Writer: w, pool: &ret.poolWriter[i]}
|
||||
}
|
||||
}
|
||||
ret.poolReader.New = func() interface{} {
|
||||
r, err := gzip.NewReader(io.NopCloser(nil))
|
||||
runtime.Assert(err)
|
||||
return &reader{Reader: r, pool: &ret.poolReader}
|
||||
}
|
||||
return ret, nil
|
||||
}
|
||||
|
||||
// Compress gzip compress
|
||||
func (c *compressor) Compress(w io.Writer) (io.WriteCloser, error) {
|
||||
pw := c.poolWriter[c.level].Get().(*writer)
|
||||
pw.Writer.Reset(w)
|
||||
return pw, nil
|
||||
}
|
||||
|
||||
// Decompress gzip decompress
|
||||
func (c *compressor) Decompress(r io.Reader) (io.ReadCloser, error) {
|
||||
pr := c.poolReader.Get().(*reader)
|
||||
pr.Reader.Reset(r)
|
||||
return pr, nil
|
||||
}
|
||||
|
||||
// SetLevel set compress level
|
||||
func (c *compressor) SetLevel(level int) error {
|
||||
if level < 0 || level > gzip.BestCompression {
|
||||
return fmt.Errorf("invalid gzip compress level: %d", level)
|
||||
}
|
||||
c.level = level
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
package proto
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/lwch/natpass/code/network/encoding"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
type codec struct{}
|
||||
|
||||
// New create protobuf codec
|
||||
func New() encoding.Codec {
|
||||
return &codec{}
|
||||
}
|
||||
|
||||
// Marshal protobuf marshal
|
||||
func (*codec) Marshal(v interface{}) ([]byte, error) {
|
||||
vv, ok := v.(proto.Message)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("invalid value type, want proto.Message, got %T", v)
|
||||
}
|
||||
return proto.Marshal(vv)
|
||||
}
|
||||
|
||||
// Unmarshal protobuf unmarshal
|
||||
func (*codec) Unmarshal(data []byte, v interface{}) error {
|
||||
vv, ok := v.(proto.Message)
|
||||
if !ok {
|
||||
return fmt.Errorf("invalid value type, want proto.Message, got %T", v)
|
||||
}
|
||||
return proto.Unmarshal(data, vv)
|
||||
}
|
||||
+106
-34
@@ -13,7 +13,8 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/lwch/logging"
|
||||
"google.golang.org/protobuf/proto"
|
||||
"github.com/lwch/natpass/code/network/encoding"
|
||||
"github.com/lwch/natpass/code/network/encoding/proto"
|
||||
)
|
||||
|
||||
var errTooLong = errors.New("too long")
|
||||
@@ -22,12 +23,13 @@ var errTimeout = errors.New("timeout")
|
||||
|
||||
// Conn network connection
|
||||
type Conn struct {
|
||||
c net.Conn
|
||||
lockRead sync.Mutex
|
||||
sizeRead [6]byte
|
||||
chWrite chan []byte
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
c net.Conn
|
||||
lockRead sync.Mutex
|
||||
chWrite chan []byte
|
||||
codec encoding.Codec
|
||||
compressor encoding.Compressor
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
// NewConn create connection
|
||||
@@ -36,6 +38,7 @@ func NewConn(c net.Conn) *Conn {
|
||||
conn := &Conn{
|
||||
c: c,
|
||||
chWrite: make(chan []byte, 1024),
|
||||
codec: proto.New(),
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
}
|
||||
@@ -43,66 +46,135 @@ func NewConn(c net.Conn) *Conn {
|
||||
return conn
|
||||
}
|
||||
|
||||
// SetCompressor set compressor
|
||||
func (c *Conn) SetCompressor(cp encoding.Compressor) *Conn {
|
||||
c.compressor = cp
|
||||
return c
|
||||
}
|
||||
|
||||
// SetCodec set codec
|
||||
func (c *Conn) SetCodec(cc encoding.Codec) *Conn {
|
||||
c.codec = cc
|
||||
return c
|
||||
}
|
||||
|
||||
// Close close connection
|
||||
func (c *Conn) Close() {
|
||||
c.c.Close()
|
||||
c.cancel()
|
||||
}
|
||||
|
||||
func (c *Conn) read(timeout time.Duration) (uint32, uint16, []byte, error) {
|
||||
type header struct {
|
||||
Size uint16
|
||||
Checksum uint32
|
||||
}
|
||||
|
||||
func (c *Conn) read(timeout time.Duration) ([]byte, error) {
|
||||
c.lockRead.Lock()
|
||||
defer c.lockRead.Unlock()
|
||||
c.c.SetReadDeadline(time.Now().Add(timeout))
|
||||
_, err := io.ReadFull(c.c, c.sizeRead[:])
|
||||
var hdr header
|
||||
err := binary.Read(c.c, binary.BigEndian, &hdr)
|
||||
if err != nil {
|
||||
return 0, 0, nil, err
|
||||
return nil, err
|
||||
}
|
||||
size := binary.BigEndian.Uint16(c.sizeRead[:])
|
||||
enc := binary.BigEndian.Uint32(c.sizeRead[2:])
|
||||
buf := make([]byte, size)
|
||||
buf := make([]byte, hdr.Size)
|
||||
_, err = io.ReadFull(c.c, buf)
|
||||
if err != nil {
|
||||
return 0, 0, nil, err
|
||||
return nil, err
|
||||
}
|
||||
return enc, size, buf, nil
|
||||
if crc32.ChecksumIEEE(buf) != hdr.Checksum {
|
||||
return nil, errChecksum
|
||||
}
|
||||
return buf, nil
|
||||
}
|
||||
|
||||
func (c *Conn) unserialize(data []byte) (*Msg, error) {
|
||||
if c.compressor != nil {
|
||||
dec, err := c.compressor.Decompress(bytes.NewReader(data))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var buffer bytes.Buffer
|
||||
_, err = io.Copy(&buffer, dec)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
data = buffer.Bytes()
|
||||
}
|
||||
var msg Msg
|
||||
err := c.codec.Unmarshal(data, &msg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &msg, nil
|
||||
}
|
||||
|
||||
// ReadMessage read message with timeout
|
||||
func (c *Conn) ReadMessage(timeout time.Duration) (*Msg, uint16, error) {
|
||||
enc, size, buf, err := c.read(timeout)
|
||||
buf, err := c.read(timeout)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
if crc32.ChecksumIEEE(buf) != enc {
|
||||
return nil, 0, errChecksum
|
||||
}
|
||||
var msg Msg
|
||||
err = proto.Unmarshal(buf, &msg)
|
||||
msg, err := c.unserialize(buf)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return &msg, size, nil
|
||||
return msg, uint16(len(buf)), nil
|
||||
}
|
||||
|
||||
func (c *Conn) serialize(msg *Msg) ([]byte, error) {
|
||||
data, err := c.codec.Marshal(msg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if c.compressor != nil {
|
||||
var buffer bytes.Buffer
|
||||
enc, err := c.compressor.Compress(&buffer)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
_, err = io.Copy(enc, bytes.NewReader(data))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buffer.Bytes(), nil
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
|
||||
func (c *Conn) write(data []byte, timeout time.Duration) error {
|
||||
hdr := header{
|
||||
Size: uint16(len(data)),
|
||||
Checksum: crc32.ChecksumIEEE(data),
|
||||
}
|
||||
var buffer bytes.Buffer
|
||||
err := binary.Write(&buffer, binary.BigEndian, hdr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = io.Copy(&buffer, bytes.NewReader(data))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
select {
|
||||
case c.chWrite <- buffer.Bytes():
|
||||
return nil
|
||||
case <-time.After(timeout):
|
||||
return errTimeout
|
||||
}
|
||||
}
|
||||
|
||||
// WriteMessage write message with timeout
|
||||
func (c *Conn) WriteMessage(m *Msg, timeout time.Duration) error {
|
||||
data, err := proto.Marshal(m)
|
||||
func (c *Conn) WriteMessage(msg *Msg, timeout time.Duration) error {
|
||||
data, err := c.serialize(msg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(data) > math.MaxUint16 {
|
||||
return errTooLong
|
||||
}
|
||||
buf := make([]byte, len(data)+len(c.sizeRead))
|
||||
binary.BigEndian.PutUint16(buf, uint16(len(data)))
|
||||
binary.BigEndian.PutUint32(buf[2:], crc32.ChecksumIEEE(data))
|
||||
copy(buf[len(c.sizeRead):], data)
|
||||
select {
|
||||
case c.chWrite <- buf:
|
||||
return nil
|
||||
case <-time.After(timeout):
|
||||
return errTimeout
|
||||
}
|
||||
return c.write(data, timeout)
|
||||
}
|
||||
|
||||
// RemoteAddr get connection remote address
|
||||
|
||||
Reference in New Issue
Block a user