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

143 lines
2.3 KiB
Go

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