commit 23199581da51acc3bfda52ed73289d3e1f98bd8c Author: ston1th Date: Mon Mar 26 21:07:28 2018 +0200 initial commit diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..9a63461 --- /dev/null +++ b/.gitignore @@ -0,0 +1,2 @@ +client_* +server_* diff --git a/build.sh b/build.sh new file mode 100755 index 0000000..05800ed --- /dev/null +++ b/build.sh @@ -0,0 +1,14 @@ +#!/bin/sh + +os="linux windows" +arch="386 amd64" +files="client server" + +for o in ${os[@]}; do + [ "${o}" == "windows" ] && ext=".exe" || ext= + for a in ${arch[@]}; do + for f in ${files[@]}; do + GOOS=${o} GOARCH=${a} go build -ldflags '-s -w' -o ${f}_${o}_${a}${ext} ${f}.go + done + done +done diff --git a/client.go b/client.go new file mode 100644 index 0000000..a438baa --- /dev/null +++ b/client.go @@ -0,0 +1,114 @@ +package main + +import ( + "bytes" + "flag" + "io" + "io/ioutil" + "log" + "net" + "net/http" + "os" + "time" +) + +func makeRead(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 = makeRead(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 = makeRead(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 ( + server string + listen string + connect string + ) + flag.StringVar(&server, "server", "", "server addr") + flag.StringVar(&listen, "listen", ":8080", "client listen address (use '-' for stdin/stdout)") + flag.StringVar(&connect, "connect", "", "remote connect addr") + if err := client(server, listen, connect); err != nil { + panic(err) + } +} diff --git a/main.go b/main.go new file mode 100644 index 0000000..68ff54a --- /dev/null +++ b/main.go @@ -0,0 +1,247 @@ +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) + } + } +} diff --git a/server.go b/server.go new file mode 100644 index 0000000..e9d29dc --- /dev/null +++ b/server.go @@ -0,0 +1,143 @@ +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) + } +}