diff --git a/main.go b/main.go index 3488355..244cd48 100644 --- a/main.go +++ b/main.go @@ -38,12 +38,22 @@ var ( Envar("MTG_PORT"). Default("3128"). Uint16() - readTimeout = app.Flag("read-timeout", "Socket read timeout"). + statsIP = app.Flag("stats-ip", "Which IP bind stats server to"). + Short('t'). + Envar("MTG_STATS_IP"). + Default("127.0.0.1"). + IP() + statsPort = app.Flag("stats-port", "Which port bind stats to."). + Short('q'). + Envar("MTG_STATS_PORT"). + Default("3129"). + Uint16() + readTimeout = app.Flag("read-timeout", "Socket read timeout."). Short('r'). Envar("MTG_READ_TIMEOUT"). Default("30s"). Duration() - writeTimeout = app.Flag("write-timeout", "Socket write timeout"). + writeTimeout = app.Flag("write-timeout", "Socket write timeout."). Short('w'). Envar("MTG_WRITE_TIMEOUT"). Default("30s"). @@ -95,8 +105,12 @@ func main() { )).Sugar() printURLs() + + stat := proxy.NewStats() + go stat.Serve(*statsIP, *statsPort) + srv := proxy.NewServer(*bindIP, int(*bindPort), secretBytes, logger, - *readTimeout, *writeTimeout) + *readTimeout, *writeTimeout, stat) if err := srv.Serve(); err != nil { logger.Fatal(err.Error()) } diff --git a/proxy/server.go b/proxy/server.go index 7be96d4..cea4d60 100644 --- a/proxy/server.go +++ b/proxy/server.go @@ -25,6 +25,7 @@ type Server struct { ctx context.Context readTimeout time.Duration writeTimeout time.Duration + stats *Stats } func (s *Server) Serve() error { @@ -50,6 +51,8 @@ func (s *Server) Addr() string { func (s *Server) accept(conn net.Conn) { defer conn.Close() + defer s.stats.closeConnection() + s.stats.newConnection() ctx, cancel := context.WithCancel(context.Background()) socketID := s.makeSocketID() @@ -102,6 +105,7 @@ func (s *Server) makeSocketID() string { func (s *Server) getClientStream(conn net.Conn, ctx context.Context, cancel context.CancelFunc, socketID string) (io.ReadWriteCloser, int16, error) { wConn := newTimeoutReadWriteCloser(conn, s.readTimeout, s.writeTimeout) + wConn = newTrafficReadWriteCloser(wConn, s.stats.addIncomingTraffic, s.stats.addOutgoingTraffic) frame, err := obfuscated2.ExtractFrame(wConn) if err != nil { return nil, 0, errors.Annotate(err, "Cannot create client stream") @@ -113,7 +117,7 @@ func (s *Server) getClientStream(conn net.Conn, ctx context.Context, cancel cont } wConn = newLogReadWriteCloser(wConn, s.logger, socketID, "client") - wConn = newCipherReadWriteCloser(conn, obfs2) + wConn = newCipherReadWriteCloser(wConn, obfs2) wConn = newCtxReadWriteCloser(wConn, ctx, cancel) return wConn, dc, nil @@ -125,13 +129,14 @@ func (s *Server) getTelegramStream(dc int16, ctx context.Context, cancel context return nil, errors.Annotate(err, "Cannot dial") } wConn := newTimeoutReadWriteCloser(socket, s.readTimeout, s.writeTimeout) + wConn = newTrafficReadWriteCloser(wConn, s.stats.addIncomingTraffic, s.stats.addOutgoingTraffic) obfs2, frame := obfuscated2.MakeTelegramObfuscated2Frame() if n, err := socket.Write(frame); err != nil || n != len(frame) { return nil, errors.Annotate(err, "Cannot write hadnshake frame") } - wConn = newLogReadWriteCloser(socket, s.logger, socketID, "telegram") + wConn = newLogReadWriteCloser(wConn, s.logger, socketID, "telegram") wConn = newCipherReadWriteCloser(wConn, obfs2) wConn = newCtxReadWriteCloser(wConn, ctx, cancel) @@ -140,13 +145,11 @@ func (s *Server) getTelegramStream(dc int16, ctx context.Context, cancel context func (s *Server) pipe(wait *sync.WaitGroup, reader io.Reader, writer io.Writer) { defer wait.Done() - - buf := make([]byte, bufferSize) - io.CopyBuffer(writer, reader, buf) + io.Copy(writer, reader) } func NewServer(ip net.IP, port int, secret []byte, logger *zap.SugaredLogger, - readTimeout, writeTimeout time.Duration) *Server { + readTimeout, writeTimeout time.Duration, stat *Stats) *Server { return &Server{ ip: ip, port: port, @@ -155,5 +158,6 @@ func NewServer(ip net.IP, port int, secret []byte, logger *zap.SugaredLogger, logger: logger, readTimeout: readTimeout, writeTimeout: writeTimeout, + stats: stat, } } diff --git a/proxy/stats.go b/proxy/stats.go new file mode 100644 index 0000000..02016d1 --- /dev/null +++ b/proxy/stats.go @@ -0,0 +1,58 @@ +package proxy + +import ( + "encoding/json" + "net" + "net/http" + "strconv" + "sync/atomic" + "time" +) + +type statsUptime time.Time + +func (s statsUptime) MarshalJSON() ([]byte, error) { + uptime := int(time.Since(time.Time(s)).Seconds()) + return []byte(strconv.Itoa(uptime)), nil +} + +type Stats struct { + AllConnections uint64 `json:"all_connections"` + ActiveConnections uint32 `json:"active_connections"` + Traffic struct { + Incoming uint64 `json:"incoming"` + Outgoing uint64 `json:"outgoing"` + } `json:"traffic"` + Uptime statsUptime `json:"uptime"` +} + +func (s *Stats) newConnection() { + atomic.AddUint64(&s.AllConnections, 1) + atomic.AddUint32(&s.ActiveConnections, 1) +} + +func (s *Stats) closeConnection() { + atomic.AddUint32(&s.ActiveConnections, ^uint32(0)) +} + +func (s *Stats) addIncomingTraffic(n int) { + atomic.AddUint64(&s.Traffic.Incoming, uint64(n)) +} + +func (s *Stats) addOutgoingTraffic(n int) { + atomic.AddUint64(&s.Traffic.Outgoing, uint64(n)) +} + +func (s *Stats) Serve(host net.IP, port uint16) { + http.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(s) + }) + + addr := net.JoinHostPort(host.String(), strconv.Itoa(int(port))) + http.ListenAndServe(addr, nil) +} + +func NewStats() *Stats { + return &Stats{Uptime: statsUptime(time.Now())} +} diff --git a/proxy/trafficrwc.go b/proxy/trafficrwc.go new file mode 100644 index 0000000..6afd300 --- /dev/null +++ b/proxy/trafficrwc.go @@ -0,0 +1,33 @@ +package proxy + +import "io" + +type TrafficReadWriteCloser struct { + conn io.ReadWriteCloser + readCallback func(int) + writeCallback func(int) +} + +func (t *TrafficReadWriteCloser) Read(p []byte) (n int, err error) { + n, err = t.conn.Read(p) + t.readCallback(n) + return +} + +func (t *TrafficReadWriteCloser) Write(p []byte) (n int, err error) { + n, err = t.conn.Write(p) + t.writeCallback(n) + return +} + +func (t *TrafficReadWriteCloser) Close() error { + return t.conn.Close() +} + +func newTrafficReadWriteCloser(conn io.ReadWriteCloser, readCallback, writeCallback func(int)) io.ReadWriteCloser { + return &TrafficReadWriteCloser{ + conn: conn, + readCallback: readCallback, + writeCallback: writeCallback, + } +}