package main import ( "bytes" "crypto/rand" "encoding/hex" "flag" "io" "io/ioutil" "log" "net" "net/http" "sync" ) 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 main() { var listen string flag.StringVar(&listen, "listen", ":8080", "server listen address") flag.Parse() if err := http.ListenAndServe(listen, &srv{conns: make(map[string]*conn)}); err != nil { panic(err) } }