Refactor to proxy router

This commit is contained in:
wweir
2020-02-05 18:00:09 +08:00
parent 8cd4fe57ba
commit fd5c34dca1
40 changed files with 538 additions and 1777 deletions
-98
View File
@@ -1,98 +0,0 @@
package proxy
import (
"net"
"github.com/golang/glog"
"github.com/wweir/sower/proxy/parser"
"github.com/wweir/sower/proxy/shadow"
"github.com/wweir/sower/proxy/socks5"
"github.com/wweir/sower/proxy/transport"
)
func StartClient(tran transport.Transport, isSocks5 bool, server, cipher, password, listenIP string) {
conn80 := listenLocal(listenIP, "80")
conn443 := listenLocal(listenIP, "443")
var isHttp bool
var conn net.Conn
glog.Infoln("Client started.")
for {
select {
case conn = <-conn80:
isHttp = true
case conn = <-conn443:
isHttp = false
}
resolveAddr(&server)
glog.V(1).Infof("new conn from (%s) to (%s)", conn.RemoteAddr(), server)
rc, err := tran.Dial(server)
if err != nil {
conn.Close()
glog.Errorln(err)
continue
}
switch {
case isSocks5 && isHttp:
c, host, port, err := parser.ParseHttpAddr(conn)
if err != nil {
c.Close()
rc.Close()
glog.Errorln(err)
continue
}
conn = c
rc = socks5.ToSocks5(rc, host, port)
case isSocks5 && !isHttp:
c, host, err := parser.ParseHttpsHost(conn)
if err != nil {
c.Close()
rc.Close()
glog.Errorln(err)
continue
}
conn = c
rc = socks5.ToSocks5(rc, host, "443")
case !isSocks5 && isHttp:
rc = shadow.Shadow(rc, cipher, password)
rc = parser.NewHttpConn(rc)
case !isSocks5 && !isHttp:
rc = shadow.Shadow(rc, cipher, password)
rc = parser.NewHttpsConn(rc, "443")
}
go relay(conn, rc)
}
}
func listenLocal(listenIP string, port string) <-chan net.Conn {
connCh := make(chan net.Conn, 10)
go func() {
ln, err := net.Listen("tcp", net.JoinHostPort(listenIP, port))
if err != nil {
glog.Fatalln(err)
}
for {
conn, err := ln.Accept()
if err != nil {
glog.Errorln("accept", listenIP+port, "fail:", err)
continue
}
conn.(*net.TCPConn).SetKeepAlive(true)
connCh <- conn
}
}()
glog.Infoln("listening port:", port)
return connCh
}
-127
View File
@@ -1,127 +0,0 @@
package proxy
import (
"context"
"crypto/tls"
"io"
"net"
"net/http"
"net/url"
"time"
"github.com/golang/glog"
"github.com/wweir/sower/proxy/parser"
"github.com/wweir/sower/proxy/shadow"
"github.com/wweir/sower/proxy/socks5"
"github.com/wweir/sower/proxy/transport"
)
func StartHttpProxy(tran transport.Transport, isSocks5 bool, server, cipher, password, addr string) {
srv := &http.Server{
Addr: addr,
Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resolveAddr(&server)
if r.Method == http.MethodConnect {
httpsProxy(w, r, tran, isSocks5, server, cipher, password)
} else {
httpProxy(w, r, tran, isSocks5, server, cipher, password)
}
}),
// Disable HTTP/2.
TLSNextProto: map[string]func(*http.Server, *tls.Conn, http.Handler){},
IdleTimeout: 90 * time.Second,
}
glog.Fatalln(srv.ListenAndServe())
}
func httpProxy(w http.ResponseWriter, r *http.Request,
tran transport.Transport, isSocks5 bool, server, cipher, password string) {
roundTripper := &http.Transport{
MaxIdleConns: 100,
IdleConnTimeout: 90 * time.Second,
TLSHandshakeTimeout: 10 * time.Second,
ExpectContinueTimeout: 1 * time.Second,
}
if isSocks5 {
roundTripper.Proxy = func(*http.Request) (*url.URL, error) {
return url.Parse("socks5://" + server)
}
} else {
roundTripper.DialContext = func(context.Context, string, string) (net.Conn, error) {
conn, err := tran.Dial(server)
if err != nil {
return nil, err
}
conn = shadow.Shadow(conn, cipher, password)
return parser.NewHttpConn(conn), nil
}
}
resp, err := roundTripper.RoundTrip(r)
if err != nil {
http.Error(w, err.Error(), http.StatusServiceUnavailable)
glog.Errorln("serve https proxy, get remote data:", err)
return
}
defer resp.Body.Close()
for k, vs := range resp.Header {
for _, v := range vs {
w.Header().Add(k, v)
}
}
w.WriteHeader(resp.StatusCode)
io.Copy(w, resp.Body)
}
func httpsProxy(w http.ResponseWriter, r *http.Request,
tran transport.Transport, isSocks5 bool, server, cipher, password string) {
// local conn
conn, _, err := w.(http.Hijacker).Hijack()
if err != nil {
http.Error(w, err.Error(), http.StatusServiceUnavailable)
return
}
conn.(*net.TCPConn).SetKeepAlive(true)
if _, err := conn.Write([]byte(r.Proto + " 200 Connection established\r\n\r\n")); err != nil {
http.Error(w, err.Error(), http.StatusServiceUnavailable)
conn.Close()
glog.Errorln("serve https proxy, write data fail:", err)
return
}
// remote conn
rc, err := tran.Dial(server)
if err != nil {
http.Error(w, err.Error(), http.StatusServiceUnavailable)
conn.Close()
glog.Errorln("serve https proxy, dial remote fail:", err)
return
}
host, port, err := net.SplitHostPort(r.Host)
if err != nil {
http.Error(w, err.Error(), http.StatusServiceUnavailable)
conn.Close()
glog.Errorln("serve https proxy, dial remote fail:", err)
return
}
if isSocks5 {
rc = socks5.ToSocks5(rc, host, port)
} else {
rc = shadow.Shadow(rc, cipher, password)
rc = parser.NewHttpsConn(rc, port)
}
relay(rc, conn)
}
-12
View File
@@ -1,12 +0,0 @@
// Package parser transter conn to be a parser conn
//
// init request payload:
// <type>(1) + <size>(2))(+Overhead) + <data>(size+Overhead)
// data definition:
// 0x00(any): [size](1) + [addr:port] + content
// 0x01(http): content
// 0x02(https): [port](2) + content
//
// init response payload:
// ([status code](2) + <size>(2))(+Overhead) + <content>(size+Overhead)
package parser
-163
View File
@@ -1,163 +0,0 @@
package parser
import (
"bufio"
"io"
"net"
"net/http"
"strconv"
"strings"
"github.com/pkg/errors"
"github.com/wweir/sower/util"
)
const (
OTHER byte = iota
HTTP
HTTPS
)
// Write Addr
type conn struct {
typ byte
domain string
port string
init bool
net.Conn
}
func NewOtherConn(c net.Conn, domain, port string) net.Conn {
return &conn{
typ: OTHER,
domain: domain,
port: port,
init: true,
Conn: c,
}
}
func NewHttpConn(c net.Conn) net.Conn {
return &conn{
typ: HTTP,
init: true,
Conn: c,
}
}
func NewHttpsConn(c net.Conn, port string) net.Conn {
return &conn{
typ: HTTPS,
port: port,
init: true,
Conn: c,
}
}
func (c *conn) Write(b []byte) (n int, err error) {
if c.init {
var pkg []byte
var prefixLen int
switch c.typ {
case OTHER:
// type + domain + ':' + port + data
prefixLen = 1 + len(c.domain) + 1 + len(c.port)
pkg = make([]byte, 0, prefixLen+len(b))
pkg = append(pkg, OTHER)
pkg = append(pkg, byte(len(c.domain)+1+len(c.port)))
pkg = append(pkg, []byte(c.domain+":"+c.port)...)
case HTTP:
// type + data
prefixLen = 1
pkg = make([]byte, 0, prefixLen+len(b))
pkg = append(pkg, HTTP)
case HTTPS:
// type + port + data
prefixLen = 1 + 2
pkg = make([]byte, 0, prefixLen+len(b))
pkg = append(pkg, HTTPS)
port, _ := strconv.Atoi(c.port)
pkg = append(pkg, byte(port>>8), byte(port))
}
c.init = false
n, err := c.Conn.Write(append(pkg, b...))
// n should larger than prefix length, if not, err is not nil
return n - prefixLen, err
}
return c.Conn.Write(b)
}
// Read Addr
func ParseAddr(conn net.Conn) (net.Conn, string, string, error) {
buf := make([]byte, 1)
if _, err := io.ReadFull(conn, buf); err != nil {
return conn, "", "", err
}
switch buf[0] {
case OTHER:
if _, err := io.ReadFull(conn, buf); err != nil {
return conn, "", "", err
}
buf = make([]byte, int(buf[0]))
if _, err := io.ReadFull(conn, buf); err != nil {
return conn, "", "", err
}
addr := string(buf)
if idx := strings.LastIndex(addr, ":"); idx != -1 {
return conn, addr[:idx], addr[idx+1:], nil
}
return conn, "", "", errors.New("invalid payload")
case HTTP:
return ParseHttpAddr(conn)
case HTTPS:
buf = make([]byte, 2)
if _, err := io.ReadFull(conn, buf); err != nil {
return conn, "", "", err
}
port := strconv.Itoa(int(buf[0])<<8 + int(buf[1]))
conn, domain, err := ParseHttpsHost(conn)
return conn, domain, port, err
default:
return conn, "", "", errors.Errorf("not supported type (%v)", buf[0])
}
}
func ParseHttpAddr(conn net.Conn) (net.Conn, string, string, error) {
teeConn := &util.TeeConn{Conn: conn}
teeConn.StartOrReset()
defer teeConn.Stop()
b := bufio.NewReader(teeConn)
resp, err := http.ReadRequest(b)
if err != nil {
return teeConn, "", "", err
}
if idx := strings.LastIndex(resp.Host, ":"); idx != -1 {
return teeConn, resp.Host[:idx], resp.Host[idx+1:], nil
}
return teeConn, resp.Host, "80", nil
}
func ParseHttpsHost(conn net.Conn) (net.Conn, string, error) {
teeConn := &util.TeeConn{Conn: conn}
teeConn.StartOrReset()
defer teeConn.Stop()
domain, _, err := extractSNI(teeConn)
if err != nil {
return teeConn, "", err
} else if domain == "" {
return teeConn, "", errors.New("ClientHello did not present an SNI extension")
}
return teeConn, domain, nil
}
-97
View File
@@ -1,97 +0,0 @@
package parser
import (
"bufio"
"bytes"
"io/ioutil"
"net"
"net/http"
"testing"
"github.com/wweir/sower/proxy/shadow"
"github.com/wweir/sower/util"
)
func TestParseAddr1(t *testing.T) {
c1, c2 := net.Pipe()
go func() {
c1 = NewHttpConn(c1)
req, _ := http.NewRequest("GET", "http://wweir.cc", bytes.NewReader([]byte{1, 2, 3}))
req.Write(c1)
}()
c2, host, port, err := ParseAddr(c2)
if err != nil || host != "wweir.cc" || port != "80" {
t.Error(err, host, port)
}
req, err := http.ReadRequest(bufio.NewReader(c2))
if err != nil {
t.Error(err)
}
data, err := ioutil.ReadAll(req.Body)
if err != nil || len(data) != 3 || data[0] != 1 {
t.Error(err, data)
}
}
func TestParseAddr2(t *testing.T) {
c1, c2 := net.Pipe()
go func() {
c1 = NewHttpsConn(c1, "443")
c1.Write(util.HTTPS.PingMsg("wweir.cc"))
}()
_, host, port, err := ParseAddr(c2)
if err != nil || host != "wweir.cc" || port != "443" {
t.Error(err, host, port)
}
}
func TestParseAddr3(t *testing.T) {
c1, c2 := net.Pipe()
go func() {
c1 = NewOtherConn(c1, "wweir.cc", "1080")
c1.Write(util.HTTPS.PingMsg("wweir.cc"))
}()
_, host, port, err := ParseAddr(c2)
if err != nil || host != "wweir.cc" || port != "1080" {
t.Error(err, host, port)
}
}
func TestParseAddr4(t *testing.T) {
c1, c2 := net.Pipe()
go func() {
c1 = shadow.Shadow(c1, "AES_128_GCM", "12345678")
c1 = NewHttpConn(c1)
req, _ := http.NewRequest("GET", "http://wweir.cc", bytes.NewReader([]byte{1, 2, 3}))
req.Write(c1)
}()
c2 = shadow.Shadow(c2, "AES_128_GCM", "12345678")
c2, host, port, err := ParseAddr(c2)
if err != nil || host != "wweir.cc" || port != "80" {
t.Error(err, host, port)
}
req, err := http.ReadRequest(bufio.NewReader(c2))
if err != nil {
t.Error(err)
}
data, err := ioutil.ReadAll(req.Body)
if err != nil || len(data) != 3 || data[0] != 1 {
t.Error(err, data)
}
}
-232
View File
@@ -1,232 +0,0 @@
// Copyright 2016 Google Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package parser
import (
"encoding/binary"
"errors"
"fmt"
"io"
)
func extractSNI(r io.Reader) (string, int, error) {
handshake, tlsver, err := handshakeRecord(r)
if err != nil {
return "", 0, fmt.Errorf("reading TLS record: %s", err)
}
sni, err := parseHello(handshake)
if err != nil {
return "", 0, fmt.Errorf("reading ClientHello: %s", err)
}
if len(sni) == 0 {
// ClientHello did not present an SNI extension. Valid packet,
// no hostname.
return "", tlsver, nil
}
hostname, err := parseSNI(sni)
if err != nil {
return "", 0, fmt.Errorf("parsing SNI extension: %s", err)
}
return hostname, tlsver, nil
}
// Extract the indicated hostname, if any, from the given SNI
// extension bytes.
func parseSNI(b []byte) (string, error) {
b, _, err := vector(b, 2)
if err != nil {
return "", err
}
var ret []byte
for len(b) >= 3 {
typ := b[0]
ret, b, err = vector(b[1:], 2)
if err != nil {
return "", fmt.Errorf("truncated SNI extension")
}
if typ == sniHostnameID {
return string(ret), nil
}
}
if len(b) != 0 {
return "", fmt.Errorf("trailing garbage at end of SNI extension")
}
// No DNS-based SNI present.
return "", nil
}
const sniExtensionID = 0
const sniHostnameID = 0
// Parse a TLS handshake record as a ClientHello message and extract
// the SNI extension bytes, if any.
func parseHello(b []byte) ([]byte, error) {
if len(b) == 0 {
return nil, errors.New("zero length handshake record")
}
if b[0] != 1 {
return nil, fmt.Errorf("non-ClientHello handshake record type %d", b[0])
}
// We're expecting a stricter TLS parser to run after we've
// proxied, so we ignore any trailing bytes that might be present
// (e.g. another handshake message).
b, _, err := vector(b[1:], 3)
if err != nil {
return nil, fmt.Errorf("reading ClientHello: %s", err)
}
// ClientHello must be at least 34 bytes to reach the first vector
// length byte. The actual minimal size is larger than that, but
// vector() will correctly handle truncated packets.
if len(b) < 34 {
return nil, errors.New("ClientHello packet too short")
}
if b[0] != 3 {
return nil, fmt.Errorf("ClientHello has unsupported version %d.%d", b[0], b[1])
}
switch b[1] {
case 1, 2, 3:
// TLS 1.0, TLS 1.1, TLS 1.2
default:
return nil, fmt.Errorf("TLS record has unsupported version %d.%d", b[0], b[1])
}
// Skip over version and random struct
b = b[34:]
// We don't technically care about SessionID, but we care that the
// framing is well-formed all the way up to the SNI field, so that
// we are sure that we're pulling the same SNI bytes as the
// eventual TLS implementation.
vec, b, err := vector(b, 1)
if err != nil {
return nil, fmt.Errorf("reading ClientHello SessionID: %s", err)
}
if len(vec) > 32 {
return nil, fmt.Errorf("ClientHello SessionID too long (%db)", len(vec))
}
// Likewise, we're just checking the bare minimum of framing.
vec, b, err = vector(b, 2)
if err != nil {
return nil, fmt.Errorf("reading ClientHello CipherSuites: %s", err)
}
if len(vec) < 2 || len(vec)%2 != 0 {
return nil, fmt.Errorf("ClientHello CipherSuites invalid length %d", len(vec))
}
vec, b, err = vector(b, 1)
if err != nil {
return nil, fmt.Errorf("reading ClientHello CompressionMethods: %s", err)
}
if len(vec) < 1 {
return nil, fmt.Errorf("ClientHello CompressionMethods invalid length %d", len(vec))
}
// Finally, we reach the extensions.
if len(b) == 0 {
// No extensions. This is not an error, it just means we have
// no SNI payload.
return nil, nil
}
b, vec, err = vector(b, 2)
if err != nil {
return nil, fmt.Errorf("reading ClientHello extensions: %s", err)
}
if len(vec) != 0 {
return nil, fmt.Errorf("%d bytes of trailing garbage in ClientHello", len(vec))
}
for len(b) >= 4 {
typ := binary.BigEndian.Uint16(b[:2])
vec, b, err = vector(b[2:], 2)
if err != nil {
return nil, fmt.Errorf("reading ClientHello extension %d: %s", typ, err)
}
if typ == sniExtensionID {
// Found the SNI extension, return its payload. We don't
// care about anything in the packet beyond this point.
return vec, nil
}
}
if len(b) != 0 {
return nil, fmt.Errorf("%d bytes of trailing garbage in ClientHello", len(b))
}
// Successfully parsed all extensions, but there was no SNI.
return nil, nil
}
const maxTLSRecordLength = 16384
// Read one TLS record, which must be for the handshake protocol, from r.
func handshakeRecord(r io.Reader) ([]byte, int, error) {
var hdr struct {
Type uint8
Major, Minor uint8
Length uint16
}
if err := binary.Read(r, binary.BigEndian, &hdr); err != nil {
return nil, 0, fmt.Errorf("reading TLS record header: %s", err)
}
if hdr.Type != 22 {
return nil, 0, fmt.Errorf("TLS record is not a handshake")
}
if hdr.Major != 3 {
return nil, 0, fmt.Errorf("TLS record has unsupported version %d.%d", hdr.Major, hdr.Minor)
}
switch hdr.Minor {
case 1, 2, 3:
// TLS 1.0, TLS 1.1, TLS 1.2
default:
return nil, 0, fmt.Errorf("TLS record has unsupported version %d.%d", hdr.Major, hdr.Minor)
}
if hdr.Length > maxTLSRecordLength {
return nil, 0, fmt.Errorf("TLS record length is greater than %d", maxTLSRecordLength)
}
ret := make([]byte, hdr.Length)
if _, err := io.ReadFull(r, ret); err != nil {
return nil, 0, err
}
return ret, int(hdr.Minor), nil
}
func vector(b []byte, lenBytes int) ([]byte, []byte, error) {
if len(b) < lenBytes {
return nil, nil, errors.New("not enough space in packet for vector")
}
var l int
for _, b := range b[:lenBytes] {
l = (l << 8) + int(b)
}
if len(b) < l+lenBytes {
return nil, nil, errors.New("not enough space in packet for vector")
}
return b[lenBytes : l+lenBytes], b[l+lenBytes:], nil
}
-51
View File
@@ -1,51 +0,0 @@
package proxy
import (
"net"
"strings"
"github.com/golang/glog"
"github.com/wweir/sower/proxy/parser"
"github.com/wweir/sower/proxy/shadow"
"github.com/wweir/sower/proxy/transport"
)
func StartServer(tran transport.Transport, port, cipher, password string) {
if port == "" {
glog.Fatalln("port must set")
}
if !strings.HasPrefix(port, ":") {
port = ":" + port
}
connCh, err := tran.Listen(port)
if err != nil {
glog.Fatalf("listen %v fail: %s", port, err)
}
glog.Infoln("Server started.")
for {
go handle(<-connCh, cipher, password)
}
}
func handle(conn net.Conn, cipher, password string) {
conn = shadow.Shadow(conn, cipher, password)
conn, host, port, err := parser.ParseAddr(conn)
if err != nil {
conn.Close()
glog.Warningln(err)
return
}
glog.V(1).Infof("new conn from %s to %s:%s", conn.RemoteAddr(), host, port)
rc, err := net.Dial("tcp", net.JoinHostPort(host, port))
if err != nil {
conn.Close()
glog.Warningln(err)
return
}
rc.(*net.TCPConn).SetKeepAlive(true)
relay(rc, conn)
}
-76
View File
@@ -1,76 +0,0 @@
package shadow
import (
"crypto/aes"
"crypto/cipher"
"github.com/pkg/errors"
"golang.org/x/crypto/chacha20poly1305"
)
//go:generate stringer -type=typ $GOFILE
type typ int
const (
AES_128_GCM typ = iota
AES_192_GCM
AES_256_GCM
CHACHA20_IETF_POLY1305
XCHACHA20_IETF_POLY1305
cipherEnd
)
func ListCiphers() []string {
list := make([]string, 0, int(cipherEnd))
for i := typ(0); i < cipherEnd; i++ {
list = append(list, i.String())
}
return list
}
func pickCipher(typ, password string) (cipher.AEAD, error) {
var blockSize int
switch typ {
case AES_128_GCM.String():
blockSize = 16
case AES_192_GCM.String():
blockSize = 24
case AES_256_GCM.String():
blockSize = 32
case CHACHA20_IETF_POLY1305.String():
return chacha20poly1305.New(genKey(password, 256))
case XCHACHA20_IETF_POLY1305.String():
return chacha20poly1305.NewX(genKey(password, 256))
default:
return nil, errors.New("do not support cipher type: " + typ)
}
// aes gcm
block, err := aes.NewCipher(genKey(password, blockSize))
if err != nil {
return nil, errors.Wrap(err, "password")
}
aead, err := cipher.NewGCM(block)
if err != nil {
return nil, errors.Wrap(err, AES_128_GCM.String())
}
return aead, nil
}
func genKey(filler string, size int) []byte {
res := make([]byte, size)
if filler == "" {
panic("password should not be empty")
}
fillerByte := []byte(filler)
length := len(fillerByte)
for i := 0; ; i++ {
if copy(res[i*length:], fillerByte) != length {
return res
}
}
}
-5
View File
@@ -1,5 +0,0 @@
// Package shadow transter conn to be a crypto conn
// support aead mode only
// data payload:
// <size>(2+Overhead) + <content>(size+Overhead)
package shadow
-111
View File
@@ -1,111 +0,0 @@
package shadow
import (
"crypto/cipher"
"encoding/binary"
"io"
"math/rand"
"net"
)
const MAX_SIZE = 0xFFFF
type conn struct {
maxSize int
aead cipher.AEAD
encryptNonce func() []byte
decryptNonce func() []byte
writeBuf []byte
readBuf []byte
readOffset int
net.Conn
}
func (c *conn) Read(b []byte) (n int, err error) {
// read from buffer
if c.readOffset != 0 {
dataSize := len(c.readBuf) - c.aead.Overhead()
n = copy(b, c.readBuf[c.readOffset:dataSize])
c.readOffset += n
if c.readOffset == dataSize {
c.readOffset = 0
}
return
}
// read from conn
dataSize := 0
{ //read data size
c.readBuf = make([]byte, 2+c.aead.Overhead())
if _, err = io.ReadFull(c.Conn, c.readBuf); err != nil {
return
}
if _, err = c.aead.Open(c.readBuf[:0], c.decryptNonce(), c.readBuf, nil); err != nil {
return
}
dataSize = int(c.readBuf[0])<<8 + int(c.readBuf[1])
}
{ // read data
c.readBuf = make([]byte, dataSize+c.aead.Overhead())
if _, err = io.ReadFull(c.Conn, c.readBuf); err != nil {
return
}
if _, err = c.aead.Open(c.readBuf[:0], c.decryptNonce(), c.readBuf, nil); err != nil {
return
}
}
// buffer extra data
if n = copy(b, c.readBuf[:dataSize]); n < dataSize {
c.readOffset = n
}
return
}
func (c *conn) Write(b []byte) (n int, err error) {
bLen := len(b)
dataSize := MAX_SIZE - (2 + c.aead.Overhead()) - c.aead.Overhead()
if bLen < c.maxSize {
dataSize = bLen
}
// BigEndian
c.writeBuf[0], c.writeBuf[1] = byte(dataSize>>8), byte(dataSize)
c.aead.Seal(c.writeBuf[:0], c.encryptNonce(), c.writeBuf[:2], nil)
c.aead.Seal(c.writeBuf[:2+c.aead.Overhead()], c.encryptNonce(), b[:dataSize], nil)
_, err = c.Conn.Write(c.writeBuf[:dataSize+(2+c.aead.Overhead())+c.aead.Overhead()])
if err != nil {
return 0, err
}
return dataSize, err
}
func Shadow(c net.Conn, cipher, password string) net.Conn {
aead, err := pickCipher(cipher, password)
if err != nil {
panic(err)
}
return &conn{
maxSize: MAX_SIZE - (2 - aead.Overhead()) - aead.Overhead(),
aead: aead,
encryptNonce: newNonce(password, aead.NonceSize()),
decryptNonce: newNonce(password, aead.NonceSize()),
writeBuf: make([]byte, 0xFFFF),
Conn: c,
}
}
func newNonce(password string, size int) func() []byte {
num, _ := binary.Varint([]byte(password))
rnd := rand.New(rand.NewSource(num))
buf := make([]byte, size)
return func() []byte {
rnd.Read(buf)
return buf
}
}
-22
View File
@@ -1,22 +0,0 @@
package shadow
import (
"net"
"testing"
)
func TestShadow(t *testing.T) {
c1, c2 := net.Pipe()
go func() {
conn := Shadow(c1, "AES_128_GCM", "12345678")
conn.Write([]byte{1, 2})
}()
conn := Shadow(c2, "AES_128_GCM", "12345678")
buf := make([]byte, 3)
n, _ := conn.Read(buf)
if n!=2|| buf[0] != 1 || buf[1] != 2 {
t.Error(buf)
}
}
-16
View File
@@ -1,16 +0,0 @@
// Code generated by "stringer -type=typ cipher.go"; DO NOT EDIT.
package shadow
import "strconv"
const _typ_name = "AES_128_GCMAES_192_GCMAES_256_GCMCHACHA20_IETF_POLY1305XCHACHA20_IETF_POLY1305cipherEnd"
var _typ_index = [...]uint8{0, 11, 22, 33, 55, 78, 87}
func (i typ) String() string {
if i < 0 || i >= typ(len(_typ_index)-1) {
return "typ(" + strconv.FormatInt(int64(i), 10) + ")"
}
return _typ_name[_typ_index[i]:_typ_index[i+1]]
}
-150
View File
@@ -1,150 +0,0 @@
package socks5
import (
"encoding/binary"
"io"
"net"
"strconv"
"github.com/pkg/errors"
)
func ToSocks5(c net.Conn, domain, port string) net.Conn {
num, _ := strconv.Atoi(port)
bytes := []byte{byte(num >> 8), byte(num)}
return &conn{init: make(chan struct{}), Conn: c, domain: domain, port: bytes}
}
type conn struct {
init chan struct{}
domain string
port []byte
net.Conn
}
func (c *conn) Read(b []byte) (n int, err error) {
<-c.init
return c.Conn.Read(b)
}
func (c *conn) Write(b []byte) (n int, err error) {
select {
case <-c.init:
return c.Conn.Write(b)
default:
}
{
req := &authReq{
VER: 5,
NMETHODS: 1,
METHODS: [1]byte{0}, // NO AUTHENTICATION REQUIRED
}
if err := binary.Write(c.Conn, binary.BigEndian, req); err != nil {
return 0, err
}
}
{
resp := &authResp{}
if err := binary.Read(c.Conn, binary.BigEndian, resp); err != nil {
return 0, err
}
}
{
req := &request{
req: req{
VER: 5, // socks5
CMD: 1, // CONNECT
RSV: 0, // RESERVED
ATYP: 3, // DOMAINNAME
},
DST_ADDR: append([]byte{byte(len(c.domain))}, []byte(c.domain)...),
DST_PORT: c.port,
}
if _, err := c.Conn.Write(req.Bytes()); err != nil {
return 0, err
}
}
{
resp := &response{}
if err := binary.Read(c.Conn, binary.BigEndian, &(resp.resp)); err != nil {
return 0, err
}
switch resp.REP {
case 0x00:
default:
return 0, errors.Errorf("socks5 handshake fail, return code: %d", resp.REP)
}
switch resp.ATYP {
case 0x01: // IPv4
resp.DST_ADDR = make([]byte, net.IPv4len)
if _, err := io.ReadFull(c.Conn, resp.DST_ADDR); err != nil {
return 0, err
}
case 0x03: // domain name
if _, err := io.ReadFull(c.Conn, resp.DST_ADDR[:1]); err != nil {
return 0, err
}
if _, err := io.ReadFull(c.Conn, resp.DST_ADDR[1:1+int(resp.DST_ADDR[0])]); err != nil {
return 0, err
}
case 0x04:
resp.DST_ADDR = make([]byte, net.IPv6len)
if _, err := io.ReadFull(c.Conn, resp.DST_ADDR); err != nil {
return 0, err
}
}
resp.DST_PORT = make([]byte, 2)
if _, err := io.ReadFull(c.Conn, resp.DST_PORT); err != nil {
return 0, err
}
}
close(c.init)
return c.Conn.Write(b)
}
type authReq struct {
VER byte
NMETHODS byte
METHODS [1]byte // 1 to 255, fix to no authentication
}
type authResp struct {
VER byte
METHOD byte
}
type request struct {
req
DST_ADDR []byte // first byte is length
DST_PORT []byte // two bytes
}
type req struct {
VER byte
CMD byte
RSV byte
ATYP byte
}
func (r *request) Bytes() []byte {
out := []byte{r.VER, r.CMD, r.RSV, r.ATYP}
out = append(out, r.DST_ADDR...)
return append(out, r.DST_PORT...)
}
type response struct {
resp
DST_ADDR []byte // first byte is length
DST_PORT []byte // two bytes
}
type resp struct {
VER byte
REP byte
RSV byte
ATYP byte
}
-115
View File
@@ -1,115 +0,0 @@
package transport
import (
"net"
"github.com/golang/glog"
"github.com/pkg/errors"
kcp "github.com/xtaci/kcp-go"
)
type kcpTran struct {
client
server
}
type client struct {
DataShard int
ParityShard int
DSCP int
SockBuf int
AckNodelay bool
NoDelay int
Interval int
Resend int
NoCongestion int
SndWnd int
RcvWnd int
MTU int
}
type server struct {
DataShard int
ParityShard int
DSCP int
SockBuf int
}
func init() {
transports["KCP"] = &kcpTran{
client: client{
DataShard: 10,
ParityShard: 3,
DSCP: 0,
SockBuf: 4194304,
NoDelay: 0,
Interval: 50,
Resend: 0,
NoCongestion: 0,
SndWnd: 0,
RcvWnd: 0,
MTU: 1350,
},
server: server{
DataShard: 10,
ParityShard: 3,
DSCP: 0,
SockBuf: 4194304,
},
}
}
func (c *client) Dial(server string) (net.Conn, error) {
conn, err := kcp.DialWithOptions(server, nil, c.DataShard, c.ParityShard)
if err != nil {
return nil, errors.Wrap(err, "dial")
}
conn.SetStreamMode(true)
conn.SetWriteDelay(false)
conn.SetNoDelay(c.NoDelay, c.Interval, c.Resend, c.NoCongestion)
conn.SetWindowSize(c.SndWnd, c.RcvWnd)
conn.SetMtu(c.MTU)
conn.SetACKNoDelay(c.AckNodelay)
if err := conn.SetDSCP(c.DSCP); err != nil {
return nil, errors.Wrap(err, "SetDSCP")
}
if err := conn.SetReadBuffer(c.SockBuf); err != nil {
return nil, errors.Wrap(err, "SetReadBuffer")
}
if err := conn.SetWriteBuffer(c.SockBuf); err != nil {
return nil, errors.Wrap(err, "SetWriteBuffer")
}
return conn, nil
}
func (s *server) Listen(port string) (<-chan net.Conn, error) {
ln, err := kcp.ListenWithOptions(port, nil, s.DataShard, s.ParityShard)
if err != nil {
return nil, err
}
if err := ln.SetDSCP(s.DSCP); err != nil {
return nil, errors.Wrap(err, "SetDSCP")
}
if err := ln.SetReadBuffer(s.SockBuf); err != nil {
return nil, errors.Wrap(err, "SetReadBuffer")
}
if err := ln.SetWriteBuffer(s.SockBuf); err != nil {
return nil, errors.Wrap(err, "SetWriteBuffer")
}
connCh := make(chan net.Conn)
go func() {
for {
conn, err := ln.AcceptKCP()
if err != nil {
glog.Fatalln("KCP listen:", err)
}
connCh <- conn
}
}()
return connCh, nil
}
-137
View File
@@ -1,137 +0,0 @@
package transport
import (
"context"
"crypto/rand"
"crypto/rsa"
"crypto/tls"
"crypto/x509"
"encoding/pem"
"math/big"
"net"
"time"
"github.com/golang/glog"
quic "github.com/lucas-clemente/quic-go"
"github.com/pkg/errors"
"github.com/wweir/sower/util"
)
type quicTran struct {
clientConf *quic.Config
sess quic.Session
serverConf *quic.Config
}
func init() {
transports["QUIC"] = &quicTran{
clientConf: &quic.Config{
HandshakeTimeout: time.Second,
KeepAlive: true,
IdleTimeout: time.Minute,
},
serverConf: &quic.Config{
MaxIncomingStreams: 1024,
},
}
}
func (c *quicTran) Dial(server string) (net.Conn, error) {
if c.sess == nil {
if sess, err := quic.DialAddr(server, &tls.Config{InsecureSkipVerify: true}, c.clientConf); err != nil {
return nil, errors.Wrap(err, "session")
} else {
go func() {
<-sess.Context().Done()
sess.Close()
c.sess = nil
}()
c.sess = sess
}
}
var stream quic.Stream
if err := util.WithTimeout(func() (err error) {
if stream, err = c.sess.OpenStream(); err != nil {
c.sess = nil
}
return
}, time.Second); err != nil {
return nil, errors.Wrap(err, "stream")
}
return &streamConn{
Stream: stream,
sess: c.sess,
}, nil
}
type streamConn struct {
quic.Stream
sess quic.Session
}
func (s *streamConn) LocalAddr() net.Addr {
return s.sess.LocalAddr()
}
func (s *streamConn) RemoteAddr() net.Addr {
return s.sess.RemoteAddr()
}
func mockTlsPem() *tls.Config {
key, err := rsa.GenerateKey(rand.Reader, 1024)
if err != nil {
glog.Fatalln(err)
}
template := x509.Certificate{SerialNumber: big.NewInt(1)}
certDER, err := x509.CreateCertificate(rand.Reader, &template, &template, &key.PublicKey, key)
if err != nil {
glog.Fatalln(err)
}
certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER})
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(key)})
tlsCert, err := tls.X509KeyPair(certPEM, keyPEM)
if err != nil {
glog.Fatalln(err)
}
return &tls.Config{Certificates: []tls.Certificate{tlsCert}}
}
func (s *quicTran) Listen(port string) (<-chan net.Conn, error) {
ln, err := quic.ListenAddr(port, mockTlsPem(), s.serverConf)
if err != nil {
return nil, errors.WithStack(err)
}
connCh := make(chan net.Conn)
go func() {
for {
sess, err := ln.Accept(context.Background())
if err != nil {
glog.Fatalln(err)
}
go accept(sess, connCh)
}
}()
return connCh, nil
}
func accept(sess quic.Session, connCh chan<- net.Conn) {
glog.V(1).Infoln("new session from ", sess.RemoteAddr())
defer sess.Close()
for {
stream, err := sess.AcceptStream(context.Background())
if err != nil {
glog.Errorln(err)
return
}
connCh <- &streamConn{stream, sess}
}
}
-58
View File
@@ -1,58 +0,0 @@
package transport
import (
"net"
"time"
"github.com/golang/glog"
)
type tcp struct {
DialTimeout time.Duration
isSocks5 bool
}
func init() {
transports["TCP"] = &tcp{
DialTimeout: 5 * time.Second,
}
transports["SOCKS5"] = &tcp{
DialTimeout: 5 * time.Second,
isSocks5: true,
}
}
func (t *tcp) Dial(server string) (net.Conn, error) {
conn, err := net.DialTimeout("tcp", server, t.DialTimeout)
if err != nil {
return nil, err
}
conn.(*net.TCPConn).SetKeepAlive(true)
return conn, nil
}
func (t *tcp) Listen(port string) (<-chan net.Conn, error) {
if t.isSocks5 {
panic("not support run as socks5 server")
}
ln, err := net.Listen("tcp", port)
if err != nil {
return nil, err
}
connCh := make(chan net.Conn)
go func() {
for {
conn, err := ln.Accept()
if err != nil {
glog.Fatalln("TCP listen:", err)
}
conn.(*net.TCPConn).SetKeepAlive(true)
connCh <- conn
}
}()
return connCh, nil
}
-30
View File
@@ -1,30 +0,0 @@
package transport
import (
"net"
"github.com/pkg/errors"
)
type Transport interface {
Dial(server string) (net.Conn, error)
Listen(port string) (<-chan net.Conn, error)
}
var transports = map[string]Transport{}
func ListTransports() []string {
list := make([]string, 0, len(transports))
for key := range transports {
list = append(list, key)
}
return list
}
func GetTransport(netType string) (Transport, error) {
tran, ok := transports[netType]
if !ok {
return nil, errors.New("invalid net type: " + netType)
}
return tran, nil
}
-53
View File
@@ -1,53 +0,0 @@
package proxy
import (
"io"
"net"
"sync"
"sync/atomic"
"time"
"github.com/golang/glog"
)
// race safe
var resolved = false
func resolveAddr(server *string) {
if !resolved {
if addr, err := net.ResolveTCPAddr("tcp", *server); err != nil {
glog.Errorln(err)
} else {
glog.Infof("remote server (%s)=>(%s)", *server, addr)
*server = addr.String()
resolved = true
}
}
}
func relay(conn1, conn2 net.Conn) {
wg := &sync.WaitGroup{}
exitFlag := new(int32)
wg.Add(2)
go redirect(conn2, conn1, wg, exitFlag)
redirect(conn1, conn2, wg, exitFlag)
wg.Wait()
}
func redirect(dst, src net.Conn, wg *sync.WaitGroup, exitFlag *int32) {
if _, err := io.Copy(dst, src); err != nil {
glog.V(1).Infof("%s<>%s -> %s<>%s: %s", src.RemoteAddr(), src.LocalAddr(), dst.LocalAddr(), dst.RemoteAddr(), err)
}
if atomic.CompareAndSwapInt32(exitFlag, 0, 1) {
// wakeup blocked goroutine
now := time.Now()
src.SetDeadline(now)
dst.SetDeadline(now)
} else {
src.Close()
dst.Close()
}
wg.Done()
}