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) } } }