peripconn.go 2.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143
  1. package fasthttp
  2. import (
  3. "crypto/tls"
  4. "encoding/binary"
  5. "net"
  6. "sync"
  7. )
  8. type perIPConnCounter struct {
  9. perIPConnPool sync.Pool
  10. perIPTLSConnPool sync.Pool
  11. m map[uint32]int
  12. lock sync.Mutex
  13. }
  14. func (cc *perIPConnCounter) Register(ip uint32) int {
  15. cc.lock.Lock()
  16. if cc.m == nil {
  17. cc.m = make(map[uint32]int)
  18. }
  19. n := cc.m[ip] + 1
  20. cc.m[ip] = n
  21. cc.lock.Unlock()
  22. return n
  23. }
  24. func (cc *perIPConnCounter) Unregister(ip uint32) {
  25. cc.lock.Lock()
  26. defer cc.lock.Unlock()
  27. if cc.m == nil {
  28. // developer safeguard
  29. panic("BUG: perIPConnCounter.Register() wasn't called")
  30. }
  31. n := max(cc.m[ip]-1, 0)
  32. cc.m[ip] = n
  33. }
  34. type perIPConn struct {
  35. net.Conn
  36. perIPConnCounter *perIPConnCounter
  37. ip uint32
  38. lock sync.Mutex
  39. }
  40. type perIPTLSConn struct {
  41. *tls.Conn
  42. perIPConnCounter *perIPConnCounter
  43. ip uint32
  44. lock sync.Mutex
  45. }
  46. func acquirePerIPConn(conn net.Conn, ip uint32, counter *perIPConnCounter) net.Conn {
  47. if tlsConn, ok := conn.(*tls.Conn); ok {
  48. v := counter.perIPTLSConnPool.Get()
  49. if v == nil {
  50. return &perIPTLSConn{
  51. perIPConnCounter: counter,
  52. Conn: tlsConn,
  53. ip: ip,
  54. }
  55. }
  56. c := v.(*perIPTLSConn)
  57. c.Conn = tlsConn
  58. c.ip = ip
  59. return c
  60. }
  61. v := counter.perIPConnPool.Get()
  62. if v == nil {
  63. return &perIPConn{
  64. perIPConnCounter: counter,
  65. Conn: conn,
  66. ip: ip,
  67. }
  68. }
  69. c := v.(*perIPConn)
  70. c.Conn = conn
  71. c.ip = ip
  72. return c
  73. }
  74. func (c *perIPConn) Close() error {
  75. c.lock.Lock()
  76. cc := c.Conn
  77. c.Conn = nil
  78. c.lock.Unlock()
  79. if cc == nil {
  80. return nil
  81. }
  82. err := cc.Close()
  83. c.perIPConnCounter.Unregister(c.ip)
  84. c.perIPConnCounter.perIPConnPool.Put(c)
  85. return err
  86. }
  87. func (c *perIPTLSConn) Close() error {
  88. c.lock.Lock()
  89. cc := c.Conn
  90. c.Conn = nil
  91. c.lock.Unlock()
  92. if cc == nil {
  93. return nil
  94. }
  95. err := cc.Close()
  96. c.perIPConnCounter.Unregister(c.ip)
  97. c.perIPConnCounter.perIPTLSConnPool.Put(c)
  98. return err
  99. }
  100. func getUint32IP(c net.Conn) uint32 {
  101. return ip2uint32(getConnIP4(c))
  102. }
  103. func getConnIP4(c net.Conn) net.IP {
  104. addr := c.RemoteAddr()
  105. ipAddr, ok := addr.(*net.TCPAddr)
  106. if !ok {
  107. return net.IPv4zero
  108. }
  109. return ipAddr.IP.To4()
  110. }
  111. func ip2uint32(ip net.IP) uint32 {
  112. if len(ip) != 4 {
  113. return 0
  114. }
  115. return uint32(ip[0])<<24 | uint32(ip[1])<<16 | uint32(ip[2])<<8 | uint32(ip[3])
  116. }
  117. func uint322ip(ip uint32) net.IP {
  118. b := make(net.IP, net.IPv4len)
  119. binary.BigEndian.PutUint32(b, ip)
  120. return b
  121. }