mirror of
https://github.com/wweir/sower.git
synced 2024-04-21 12:42:15 +00:00
Refactor to proxy router
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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]]
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
Reference in New Issue
Block a user