Files
sower/proxy/parser/parse.go
T
2019-04-16 08:04:37 +08:00

164 lines
3.3 KiB
Go

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
}