Skip to content

Instantly share code, notes, and snippets.

@kung-foo
Created May 26, 2016 19:43
Show Gist options
  • Select an option

  • Save kung-foo/0945e9efa2d46e2bbe0ca28e6527b086 to your computer and use it in GitHub Desktop.

Select an option

Save kung-foo/0945e9efa2d46e2bbe0ca28e6527b086 to your computer and use it in GitHub Desktop.
package monitor
import (
"net"
"net/http"
"sync"
"sync/atomic"
)
type NetworkMonitor struct {
mu sync.RWMutex
conns map[*TrackingConn]struct{}
bytesRead int64
bytesWritten int64
}
func NewNetworkMonitor() *NetworkMonitor {
return &NetworkMonitor{
conns: make(map[*TrackingConn]struct{}, 0),
}
}
func (nm *NetworkMonitor) NewMonitoredHTTPClient() *http.Client {
return &http.Client{
Transport: &http.Transport{
Dial: nm.Dialer,
},
}
}
func (nm *NetworkMonitor) Dialer(network, addr string) (net.Conn, error) {
conn, err := net.Dial(network, addr)
if err != nil {
return nil, err
}
tc := &TrackingConn{
Conn: conn,
onClose: nm.removeConnection,
}
nm.mu.Lock()
nm.conns[tc] = struct{}{}
nm.mu.Unlock()
return tc, nil
}
func (nm *NetworkMonitor) BytesRead() (n int64) {
nm.mu.RLock()
defer nm.mu.RUnlock()
n = nm.bytesRead
for c := range nm.conns {
n += c.BytesRead()
}
return
}
func (nm *NetworkMonitor) BytesWritten() (n int64) {
nm.mu.RLock()
defer nm.mu.RUnlock()
n = nm.bytesWritten
for c := range nm.conns {
n += c.BytesWritten()
}
return
}
func (nm *NetworkMonitor) removeConnection(conn *TrackingConn) {
nm.mu.Lock()
defer nm.mu.Unlock()
nm.bytesRead += conn.BytesRead()
nm.bytesWritten += conn.BytesWritten()
delete(nm.conns, conn)
}
type TrackingConn struct {
net.Conn
bytesRead int64
bytesWritten int64
onClose func(*TrackingConn)
}
func (tc *TrackingConn) Close() (err error) {
err = tc.Conn.Close()
tc.onClose(tc)
return
}
func (tc *TrackingConn) Read(b []byte) (n int, err error) {
n, err = tc.Conn.Read(b)
if n > 0 {
atomic.AddInt64(&tc.bytesRead, int64(n))
}
return
}
func (tc *TrackingConn) BytesRead() int64 {
return atomic.LoadInt64(&tc.bytesRead)
}
func (tc *TrackingConn) Write(b []byte) (n int, err error) {
n, err = tc.Conn.Write(b)
if n > 0 {
atomic.AddInt64(&tc.bytesWritten, int64(n))
}
return
}
func (tc *TrackingConn) BytesWritten() int64 {
return atomic.LoadInt64(&tc.bytesWritten)
}
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment