gotun/main.go
2018-03-26 21:07:28 +02:00

247 lines
4.4 KiB
Go

package main
import (
"bytes"
"crypto/rand"
"encoding/hex"
"flag"
"io"
"io/ioutil"
"log"
"net"
"net/http"
"os"
"sync"
"time"
)
func genKey() string {
b := make([]byte, 8)
rand.Read(b)
return hex.EncodeToString(b)
}
type srv struct {
sync.Mutex
conns map[string]*conn
}
func (s *srv) Conn(key string) *conn {
s.Lock()
defer s.Unlock()
return s.conns[key]
}
type conn struct {
R *buf
W *buf
C net.Conn
}
func newConn(c net.Conn) *conn {
return &conn{newBuf(), newBuf(), c}
}
type buf struct {
sync.Mutex
b *bytes.Buffer
}
func newBuf() *buf {
return &buf{b: bytes.NewBuffer(nil)}
}
func (b *buf) Read(buf []byte) (int, error) {
b.Lock()
defer b.Unlock()
return b.b.Read(buf)
}
func (b *buf) Write(buf []byte) (int, error) {
b.Lock()
defer b.Unlock()
return b.b.Write(buf)
}
func (s *srv) NewConn(key, addr string) error {
c, err := net.Dial("tcp", addr)
if err != nil {
return err
}
conn := newConn(c)
go func() {
for {
if conn.R == nil {
return
}
_, err := io.Copy(conn.R, c)
if err == io.EOF {
return
}
}
}()
go func() {
for {
if conn.W == nil {
return
}
io.Copy(c, conn.W)
}
}()
s.Lock()
defer s.Unlock()
s.conns[key] = conn
return nil
}
func (s *srv) Remove(key string) {
s.Lock()
defer s.Unlock()
conn := s.conns[key]
if conn != nil {
conn.R = nil
conn.W = nil
conn.C.Close()
}
delete(s.conns, key)
}
func (s *srv) ServeHTTP(w http.ResponseWriter, r *http.Request) {
p := r.URL.Path
switch p {
case "/new":
addr, err := ioutil.ReadAll(r.Body)
if err != nil {
http.Error(w, err.Error(), http.StatusInternalServerError)
return
}
key := genKey()
err = s.NewConn(key, string(addr))
if err != nil {
http.Error(w, err.Error(), http.StatusInternalServerError)
return
}
w.Write([]byte(key))
case "/close":
s.Remove(r.Header.Get("Key"))
default:
conn := s.Conn(r.Header.Get("Key"))
_, err := io.Copy(conn.W, r.Body)
log.Println("client", err)
r.Body.Close()
w.Header().Set("Content-Type", "application/octet-stream")
_, err = io.Copy(w, conn.R)
log.Println("remote", err)
}
}
func NewRead(r io.Reader) (chan []byte, chan error) {
read := make(chan []byte)
cerr := make(chan error)
go func() {
for {
b := make([]byte, 1024)
n, err := r.Read(b)
if err == io.EOF {
cerr <- err
return
}
if err != nil {
continue
}
if n > 0 {
read <- b[0:n]
}
}
}()
return read, cerr
}
func client(server, listen, connect string) error {
var (
read chan []byte
cerr chan error
write io.Writer
)
switch listen {
case "-":
read, cerr = NewRead(os.Stdin)
write = os.Stdout
default:
l, err := net.Listen("tcp", listen)
if err != nil {
return err
}
conn, err := l.Accept()
if err != nil {
return err
}
read, cerr = NewRead(conn)
write = conn
}
buf := bytes.NewBuffer([]byte(connect))
resp, err := http.Post("http://"+server+"/new", "text/plain", buf)
if err != nil {
return err
}
bkey, _ := ioutil.ReadAll(resp.Body)
resp.Body.Close()
key := string(bkey)
tick := time.NewTicker(time.Millisecond * 250)
buf.Reset()
cl := &http.Client{Transport: &http.Transport{MaxIdleConnsPerHost: 1}}
req, _ := http.NewRequest("POST", "http://"+server+"/", nil)
req.Header.Set("Content-Type", "application/octet-stream")
req.Header.Set("Key", key)
for {
select {
case <-tick.C:
b := bytes.NewBuffer(nil)
buf.WriteTo(b)
req.Body = ioutil.NopCloser(b)
resp, err := cl.Do(req)
if err != nil {
log.Println(err.Error())
continue
}
_, err = io.Copy(write, resp.Body)
if err != nil {
log.Println(err.Error())
}
resp.Body.Close()
case b := <-read:
buf.Write(b)
case <-cerr:
req, _ := http.NewRequest("GET", "http://"+server+"/close", nil)
req.Header.Set("Key", key)
cl.Do(req)
return nil
}
}
return nil
}
func main() {
var (
mode bool
server string
listen string
connect string
)
flag.BoolVar(&mode, "client", false, "enable client")
flag.StringVar(&server, "server", "", "server addr")
flag.StringVar(&listen, "listen", ":8080", "server/client listen addr (use 'std' for stdin/stdout)")
flag.StringVar(&connect, "connect", "", "remote connect addr")
flag.Parse()
switch mode {
case true:
if err := client(server, listen, connect); err != nil {
panic(err)
}
case false:
if err := http.ListenAndServe(listen, &srv{conns: make(map[string]*conn)}); err != nil {
panic(err)
}
}
}