add buffer pools

This commit is contained in:
rui.zheng
2016-02-04 13:45:26 +08:00
parent 35f5f2d12c
commit 4c9cdf0bb9
2 changed files with 124 additions and 68 deletions
+1 -1
View File
@@ -1,4 +1,4 @@
gosocks5 gosocks5
======== ========
golang and socks 5 golang and SOCKSV5
+123 -67
View File
@@ -4,7 +4,7 @@
package gosocks5 package gosocks5
import ( import (
"bytes" //"bytes"
"encoding/binary" "encoding/binary"
"errors" "errors"
"fmt" "fmt"
@@ -12,6 +12,7 @@ import (
//"log" //"log"
"net" "net"
"strconv" "strconv"
"sync"
) )
const ( const (
@@ -61,16 +62,33 @@ var (
ErrAuthFailure = errors.New("Auth failure") ErrAuthFailure = errors.New("Auth failure")
) )
// buffer pools
var (
sPool = sync.Pool{
New: func() interface{} {
return make([]byte, 576)
},
} // small buff pool
lPool = sync.Pool{
New: func() interface{} {
return make([]byte, 64*1024+262)
},
} // large buff pool for udp
)
/* /*
Method selection Method selection
+----+----------+----------+ +----+----------+----------+
|VER | NMETHODS | METHODS | |VER | NMETHODS | METHODS |
+----+----------+----------+ +----+----------+----------+
| 1 | 1 | 1 to 255 | | 1 | 1 | 1 to 255 |
+----+----------+----------+ +----+----------+----------+
*/ */
func ReadMethods(r io.Reader) ([]uint8, error) { func ReadMethods(r io.Reader) ([]uint8, error) {
b := make([]byte, 257) //b := make([]byte, 257)
b := sPool.Get().([]byte)
defer sPool.Put(b)
n, err := io.ReadAtLeast(r, b, 2) n, err := io.ReadAtLeast(r, b, 2)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -91,7 +109,10 @@ func ReadMethods(r io.Reader) ([]uint8, error) {
} }
} }
return b[2:length], nil methods := make([]byte, int(b[1]))
copy(methods, b[2:length])
return methods, nil
} }
func WriteMethod(method uint8, w io.Writer) error { func WriteMethod(method uint8, w io.Writer) error {
@@ -101,11 +122,11 @@ func WriteMethod(method uint8, w io.Writer) error {
/* /*
Username/Password authentication request Username/Password authentication request
+----+------+----------+------+----------+ +----+------+----------+------+----------+
|VER | ULEN | UNAME | PLEN | PASSWD | |VER | ULEN | UNAME | PLEN | PASSWD |
+----+------+----------+------+----------+ +----+------+----------+------+----------+
| 1 | 1 | 1 to 255 | 1 | 1 to 255 | | 1 | 1 | 1 to 255 | 1 | 1 to 255 |
+----+------+----------+------+----------+ +----+------+----------+------+----------+
*/ */
type UserPassRequest struct { type UserPassRequest struct {
Version byte Version byte
@@ -122,7 +143,10 @@ func NewUserPassRequest(ver byte, u, p string) *UserPassRequest {
} }
func ReadUserPassRequest(r io.Reader) (*UserPassRequest, error) { func ReadUserPassRequest(r io.Reader) (*UserPassRequest, error) {
b := make([]byte, 513) // b := make([]byte, 513)
b := sPool.Get().([]byte)
defer sPool.Put(b)
n, err := io.ReadAtLeast(r, b, 2) n, err := io.ReadAtLeast(r, b, 2)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -159,7 +183,10 @@ func ReadUserPassRequest(r io.Reader) (*UserPassRequest, error) {
} }
func (req *UserPassRequest) Write(w io.Writer) error { func (req *UserPassRequest) Write(w io.Writer) error {
b := make([]byte, 513) // b := make([]byte, 513)
b := sPool.Get().([]byte)
defer sPool.Put(b)
b[0] = req.Version b[0] = req.Version
ulen := len(req.Username) ulen := len(req.Username)
b[1] = byte(ulen) b[1] = byte(ulen)
@@ -183,11 +210,11 @@ func (req *UserPassRequest) String() string {
/* /*
Username/Password authentication response Username/Password authentication response
+----+--------+ +----+--------+
|VER | STATUS | |VER | STATUS |
+----+--------+ +----+--------+
| 1 | 1 | | 1 | 1 |
+----+--------+ +----+--------+
*/ */
type UserPassResponse struct { type UserPassResponse struct {
Version byte Version byte
@@ -202,7 +229,10 @@ func NewUserPassResponse(ver, status byte) *UserPassResponse {
} }
func ReadUserPassResponse(r io.Reader) (*UserPassResponse, error) { func ReadUserPassResponse(r io.Reader) (*UserPassResponse, error) {
b := make([]byte, 2) // b := make([]byte, 2)
b := sPool.Get().([]byte)
defer sPool.Put(b)
if _, err := io.ReadFull(r, b); err != nil { if _, err := io.ReadFull(r, b); err != nil {
return nil, err return nil, err
} }
@@ -229,6 +259,14 @@ func (res *UserPassResponse) String() string {
res.Version, res.Status) res.Version, res.Status)
} }
/*
Address
+------+----------+----------+
| ATYP | ADDR | PORT |
+------+----------+----------+
| 1 | Variable | 2 |
+------+----------+----------+
*/
type Addr struct { type Addr struct {
Type uint8 Type uint8
Host string Host string
@@ -287,11 +325,11 @@ func (addr *Addr) String() string {
/* /*
The SOCKSv5 request The SOCKSv5 request
+----+-----+-------+------+----------+----------+ +----+-----+-------+------+----------+----------+
|VER | CMD | RSV | ATYP | DST.ADDR | DST.PORT | |VER | CMD | RSV | ATYP | DST.ADDR | DST.PORT |
+----+-----+-------+------+----------+----------+ +----+-----+-------+------+----------+----------+
| 1 | 1 | X'00' | 1 | Variable | 2 | | 1 | 1 | X'00' | 1 | Variable | 2 |
+----+-----+-------+------+----------+----------+ +----+-----+-------+------+----------+----------+
*/ */
type Request struct { type Request struct {
Cmd uint8 Cmd uint8
@@ -306,7 +344,10 @@ func NewRequest(cmd uint8, addr *Addr) *Request {
} }
func ReadRequest(r io.Reader) (*Request, error) { func ReadRequest(r io.Reader) (*Request, error) {
b := make([]byte, 262) // b := make([]byte, 262)
b := sPool.Get().([]byte)
defer sPool.Put(b)
n, err := io.ReadAtLeast(r, b, 5) n, err := io.ReadAtLeast(r, b, 5)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -348,11 +389,13 @@ func ReadRequest(r io.Reader) (*Request, error) {
} }
func (r *Request) Write(w io.Writer) (err error) { func (r *Request) Write(w io.Writer) (err error) {
b := make([]byte, 262) //b := make([]byte, 262)
b := sPool.Get().([]byte)
defer sPool.Put(b)
b[0] = Ver5 b[0] = Ver5
b[1] = r.Cmd b[1] = r.Cmd
// b[2] = 0 //rsv b[2] = 0 //rsv
b[3] = AddrIPv4 // default b[3] = AddrIPv4 // default
length := 10 length := 10
@@ -375,11 +418,11 @@ func (r *Request) String() string {
/* /*
The SOCKSv5 reply The SOCKSv5 reply
+----+-----+-------+------+----------+----------+ +----+-----+-------+------+----------+----------+
|VER | REP | RSV | ATYP | BND.ADDR | BND.PORT | |VER | REP | RSV | ATYP | BND.ADDR | BND.PORT |
+----+-----+-------+------+----------+----------+ +----+-----+-------+------+----------+----------+
| 1 | 1 | X'00' | 1 | Variable | 2 | | 1 | 1 | X'00' | 1 | Variable | 2 |
+----+-----+-------+------+----------+----------+ +----+-----+-------+------+----------+----------+
*/ */
type Reply struct { type Reply struct {
Rep uint8 Rep uint8
@@ -394,7 +437,10 @@ func NewReply(rep uint8, addr *Addr) *Reply {
} }
func ReadReply(r io.Reader) (*Reply, error) { func ReadReply(r io.Reader) (*Reply, error) {
b := make([]byte, 262) // b := make([]byte, 262)
b := sPool.Get().([]byte)
defer sPool.Put(b)
n, err := io.ReadAtLeast(r, b, 5) n, err := io.ReadAtLeast(r, b, 5)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -437,11 +483,13 @@ func ReadReply(r io.Reader) (*Reply, error) {
} }
func (r *Reply) Write(w io.Writer) (err error) { func (r *Reply) Write(w io.Writer) (err error) {
b := make([]byte, 262) // b := make([]byte, 262)
b := sPool.Get().([]byte)
defer sPool.Put(b)
b[0] = Ver5 b[0] = Ver5
b[1] = r.Rep b[1] = r.Rep
// b[2] = 0 //rsv b[2] = 0 //rsv
b[3] = AddrIPv4 // default b[3] = AddrIPv4 // default
length := 10 length := 10
@@ -465,11 +513,11 @@ func (r *Reply) String() string {
/* /*
UDP request UDP request
+----+------+------+----------+----------+----------+ +----+------+------+----------+----------+----------+
|RSV | FRAG | ATYP | DST.ADDR | DST.PORT | DATA | |RSV | FRAG | ATYP | DST.ADDR | DST.PORT | DATA |
+----+------+------+----------+----------+----------+ +----+------+------+----------+----------+----------+
| 2 | 1 | 1 | Variable | 2 | Variable | | 2 | 1 | 1 | Variable | 2 | Variable |
+----+------+------+----------+----------+----------+ +----+------+------+----------+----------+----------+
*/ */
type UDPHeader struct { type UDPHeader struct {
Rsv uint16 Rsv uint16
@@ -485,6 +533,23 @@ func NewUDPHeader(rsv uint16, frag uint8, addr *Addr) *UDPHeader {
} }
} }
func (h *UDPHeader) Write(w io.Writer) error {
b := sPool.Get().([]byte)
defer sPool.Put(b)
binary.BigEndian.PutUint16(b[:2], h.Rsv)
b[2] = h.Frag
addr := h.Addr
if addr == nil {
addr = &Addr{}
}
length, _ := addr.Encode(b[3:])
_, err := w.Write(b[:3+length])
return err
}
func (h *UDPHeader) String() string { func (h *UDPHeader) String() string {
return fmt.Sprintf("%d %d %d %s", return fmt.Sprintf("%d %d %d %s",
h.Rsv, h.Frag, h.Addr.Type, h.Addr.String()) h.Rsv, h.Frag, h.Addr.Type, h.Addr.String())
@@ -503,7 +568,10 @@ func NewUDPDatagram(header *UDPHeader, data []byte) *UDPDatagram {
} }
func ReadUDPDatagram(r io.Reader) (*UDPDatagram, error) { func ReadUDPDatagram(r io.Reader) (*UDPDatagram, error) {
b := make([]byte, 65797) // b := make([]byte, 65797)
b := lPool.Get().([]byte)
defer lPool.Put(b)
n, err := io.ReadAtLeast(r, b, 5) n, err := io.ReadAtLeast(r, b, 5)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -526,7 +594,7 @@ func ReadUDPDatagram(r io.Reader) (*UDPDatagram, error) {
default: default:
return nil, ErrBadAddrType return nil, ErrBadAddrType
} }
// extended feature, for udp over tcp // extended feature, for udp over tcp, using reserved field for data length
dlen := int(header.Rsv) dlen := int(header.Rsv)
if n < hlen+dlen { if n < hlen+dlen {
if _, err := io.ReadFull(r, b[n:hlen+dlen]); err != nil { if _, err := io.ReadFull(r, b[n:hlen+dlen]); err != nil {
@@ -540,38 +608,26 @@ func ReadUDPDatagram(r io.Reader) (*UDPDatagram, error) {
return nil, err return nil, err
} }
data := make([]byte, dlen)
copy(data, b[hlen:n])
d := &UDPDatagram{ d := &UDPDatagram{
Header: header, Header: header,
Data: b[hlen:n], Data: data,
} }
return d, nil return d, nil
} }
func (d *UDPDatagram) Write(w io.Writer) error { func (d *UDPDatagram) Write(w io.Writer) error {
buffer := &bytes.Buffer{} h := d.Header
if h == nil {
b := make([]byte, 259) h = &UDPHeader{}
if d.Header != nil {
binary.BigEndian.PutUint16(b[:2], d.Header.Rsv)
buffer.Write(b[:2])
buffer.WriteByte(d.Header.Frag)
b[0] = AddrIPv4
b[1] = 0
length := 7
if d.Header.Addr != nil {
length, _ = d.Header.Addr.Encode(b)
}
buffer.Write(b[:length])
} else {
b[3] = AddrIPv4
buffer.Write(b[:10])
} }
if err := h.Write(w); err != nil {
buffer.Write(d.Data) return err
_, err := w.Write(buffer.Bytes()) }
_, err := w.Write(d.Data)
return err return err
} }