diff --git a/kcp-go/LICENSE b/kcp-go/LICENSE old mode 100755 new mode 100644 diff --git a/kcp-go/crypt.go b/kcp-go/crypt.go old mode 100755 new mode 100644 diff --git a/kcp-go/entropy.go b/kcp-go/entropy.go old mode 100755 new mode 100644 index eec960f..156c1cd --- a/kcp-go/entropy.go +++ b/kcp-go/entropy.go @@ -8,6 +8,7 @@ import ( "io" ) +// Entropy defines a entropy source type Entropy interface { Init() Fill(nonce []byte) diff --git a/kcp-go/fec.go b/kcp-go/fec.go old mode 100755 new mode 100644 diff --git a/kcp-go/kcp.go b/kcp-go/kcp.go old mode 100755 new mode 100644 index 29d0d72..6bfb04e --- a/kcp-go/kcp.go +++ b/kcp-go/kcp.go @@ -104,6 +104,7 @@ type segment struct { xmit uint32 resendts uint32 fastack uint32 + acked uint32 // mark if the seg has acked data []byte } @@ -181,8 +182,11 @@ func (kcp *KCP) newSegment(size int) (seg segment) { } // delSegment recycles a KCP segment -func (kcp *KCP) delSegment(seg segment) { - xmitBuf.Put(seg.data) +func (kcp *KCP) delSegment(seg *segment) { + if seg.data != nil { + xmitBuf.Put(seg.data) + seg.data = nil + } } // PeekSize checks the size of next message in the recv queue @@ -238,7 +242,7 @@ func (kcp *KCP) Recv(buffer []byte) (n int) { buffer = buffer[len(seg.data):] n += len(seg.data) count++ - kcp.delSegment(*seg) + kcp.delSegment(seg) if seg.frg == 0 { break } @@ -382,10 +386,8 @@ func (kcp *KCP) parse_ack(sn uint32) { for k := range kcp.snd_buf { seg := &kcp.snd_buf[k] if sn == seg.sn { - kcp.delSegment(*seg) - copy(kcp.snd_buf[k:], kcp.snd_buf[k+1:]) - kcp.snd_buf[len(kcp.snd_buf)-1] = segment{} - kcp.snd_buf = kcp.snd_buf[:len(kcp.snd_buf)-1] + seg.acked = 1 + kcp.delSegment(seg) break } if _itimediff(sn, seg.sn) < 0 { @@ -394,7 +396,7 @@ func (kcp *KCP) parse_ack(sn uint32) { } } -func (kcp *KCP) parse_fastack(sn uint32) { +func (kcp *KCP) parse_fastack(sn, ts uint32) { if _itimediff(sn, kcp.snd_una) < 0 || _itimediff(sn, kcp.snd_nxt) >= 0 { return } @@ -403,7 +405,7 @@ func (kcp *KCP) parse_fastack(sn uint32) { seg := &kcp.snd_buf[k] if _itimediff(sn, seg.sn) < 0 { break - } else if sn != seg.sn { + } else if sn != seg.sn && _itimediff(seg.ts, ts) <= 0 { seg.fastack++ } } @@ -414,7 +416,7 @@ func (kcp *KCP) parse_una(una uint32) { for k := range kcp.snd_buf { seg := &kcp.snd_buf[k] if _itimediff(una, seg.sn) > 0 { - kcp.delSegment(*seg) + kcp.delSegment(seg) count++ } else { break @@ -430,12 +432,12 @@ func (kcp *KCP) ack_push(sn, ts uint32) { kcp.acklist = append(kcp.acklist, ackItem{sn, ts}) } -func (kcp *KCP) parse_data(newseg segment) { +// returns true if data has repeated +func (kcp *KCP) parse_data(newseg segment) bool { sn := newseg.sn if _itimediff(sn, kcp.rcv_nxt+kcp.rcv_wnd) >= 0 || _itimediff(sn, kcp.rcv_nxt) < 0 { - kcp.delSegment(newseg) - return + return true } n := len(kcp.rcv_buf) - 1 @@ -445,7 +447,6 @@ func (kcp *KCP) parse_data(newseg segment) { seg := &kcp.rcv_buf[i] if seg.sn == sn { repeat = true - atomic.AddUint64(&DefaultSnmp.RepeatSegs, 1) break } if _itimediff(sn, seg.sn) > 0 { @@ -455,6 +456,11 @@ func (kcp *KCP) parse_data(newseg segment) { } if !repeat { + // replicate the content if it's new + dataCopy := xmitBuf.Get().([]byte)[:len(newseg.data)] + copy(dataCopy, newseg.data) + newseg.data = dataCopy + if insert_idx == n+1 { kcp.rcv_buf = append(kcp.rcv_buf, newseg) } else { @@ -462,8 +468,6 @@ func (kcp *KCP) parse_data(newseg segment) { copy(kcp.rcv_buf[insert_idx+1:], kcp.rcv_buf[insert_idx:]) kcp.rcv_buf[insert_idx] = newseg } - } else { - kcp.delSegment(newseg) } // move available data from rcv_buf -> rcv_queue @@ -481,6 +485,8 @@ func (kcp *KCP) parse_data(newseg segment) { kcp.rcv_queue = append(kcp.rcv_queue, kcp.rcv_buf[:count]...) kcp.rcv_buf = kcp.remove_front(kcp.rcv_buf, count) } + + return repeat } // Input when you received a low level packet (eg. UDP packet), call it @@ -491,8 +497,7 @@ func (kcp *KCP) Input(data []byte, regular, ackNoDelay bool) int { return -1 } - var maxack uint32 - var lastackts uint32 + var latest uint32 // the latest ack packet var flag int var inSegs uint64 @@ -535,20 +540,15 @@ func (kcp *KCP) Input(data []byte, regular, ackNoDelay bool) int { if cmd == IKCP_CMD_ACK { kcp.parse_ack(sn) - kcp.shrink_buf() - if flag == 0 { - flag = 1 - maxack = sn - lastackts = ts - } else if _itimediff(sn, maxack) > 0 { - maxack = sn - lastackts = ts - } + kcp.parse_fastack(sn, ts) + flag |= 1 + latest = ts } else if cmd == IKCP_CMD_PUSH { + repeat := true if _itimediff(sn, kcp.rcv_nxt+kcp.rcv_wnd) < 0 { kcp.ack_push(sn, ts) if _itimediff(sn, kcp.rcv_nxt) >= 0 { - seg := kcp.newSegment(int(length)) + var seg segment seg.conv = conv seg.cmd = cmd seg.frg = frg @@ -556,12 +556,11 @@ func (kcp *KCP) Input(data []byte, regular, ackNoDelay bool) int { seg.ts = ts seg.sn = sn seg.una = una - copy(seg.data, data[:length]) - kcp.parse_data(seg) - } else { - atomic.AddUint64(&DefaultSnmp.RepeatSegs, 1) + seg.data = data[:length] // delayed data copying + repeat = kcp.parse_data(seg) } - } else { + } + if regular && repeat { atomic.AddUint64(&DefaultSnmp.RepeatSegs, 1) } } else if cmd == IKCP_CMD_WASK { @@ -579,32 +578,36 @@ func (kcp *KCP) Input(data []byte, regular, ackNoDelay bool) int { } atomic.AddUint64(&DefaultSnmp.InSegs, inSegs) + // update rtt with the latest ts + // ignore the FEC packet if flag != 0 && regular { - kcp.parse_fastack(maxack) current := currentMs() - if _itimediff(current, lastackts) >= 0 { - kcp.update_ack(_itimediff(current, lastackts)) + if _itimediff(current, latest) >= 0 { + kcp.update_ack(_itimediff(current, latest)) } } - if _itimediff(kcp.snd_una, snd_una) > 0 { - if kcp.cwnd < kcp.rmt_wnd { - mss := kcp.mss - if kcp.cwnd < kcp.ssthresh { - kcp.cwnd++ - kcp.incr += mss - } else { - if kcp.incr < mss { - kcp.incr = mss - } - kcp.incr += (mss*mss)/kcp.incr + (mss / 16) - if (kcp.cwnd+1)*mss <= kcp.incr { + // cwnd update when packet arrived + if kcp.nocwnd == 0 { + if _itimediff(kcp.snd_una, snd_una) > 0 { + if kcp.cwnd < kcp.rmt_wnd { + mss := kcp.mss + if kcp.cwnd < kcp.ssthresh { kcp.cwnd++ + kcp.incr += mss + } else { + if kcp.incr < mss { + kcp.incr = mss + } + kcp.incr += (mss*mss)/kcp.incr + (mss / 16) + if (kcp.cwnd+1)*mss <= kcp.incr { + kcp.cwnd++ + } + } + if kcp.cwnd > kcp.rmt_wnd { + kcp.cwnd = kcp.rmt_wnd + kcp.incr = kcp.rmt_wnd * mss } - } - if kcp.cwnd > kcp.rmt_wnd { - kcp.cwnd = kcp.rmt_wnd - kcp.incr = kcp.rmt_wnd * mss } } } @@ -722,7 +725,6 @@ func (kcp *KCP) flush(ackOnly bool) uint32 { kcp.snd_buf = append(kcp.snd_buf, newseg) kcp.snd_nxt++ newSegsCount++ - kcp.snd_queue[k].data = nil } if newSegsCount > 0 { kcp.snd_queue = kcp.remove_front(kcp.snd_queue, newSegsCount) @@ -743,6 +745,9 @@ func (kcp *KCP) flush(ackOnly bool) uint32 { for k := range ref { segment := &ref[k] needsend := false + if segment.acked == 1 { + continue + } if segment.xmit == 0 { // initial transmit needsend = true segment.rto = kcp.rx_rto @@ -774,6 +779,7 @@ func (kcp *KCP) flush(ackOnly bool) uint32 { } if needsend { + current = currentMs() // time update for a blocking call segment.xmit++ segment.ts = current segment.wnd = seg.wnd @@ -784,7 +790,6 @@ func (kcp *KCP) flush(ackOnly bool) uint32 { if size+need > int(kcp.mtu) { kcp.output(buffer, size) - current = currentMs() // time update for a blocking call ptr = buffer } @@ -826,31 +831,34 @@ func (kcp *KCP) flush(ackOnly bool) uint32 { atomic.AddUint64(&DefaultSnmp.RetransSegs, sum) } - // update ssthresh - // rate halving, https://tools.ietf.org/html/rfc6937 - if change > 0 { - inflight := kcp.snd_nxt - kcp.snd_una - kcp.ssthresh = inflight / 2 - if kcp.ssthresh < IKCP_THRESH_MIN { - kcp.ssthresh = IKCP_THRESH_MIN + // cwnd update + if kcp.nocwnd == 0 { + // update ssthresh + // rate halving, https://tools.ietf.org/html/rfc6937 + if change > 0 { + inflight := kcp.snd_nxt - kcp.snd_una + kcp.ssthresh = inflight / 2 + if kcp.ssthresh < IKCP_THRESH_MIN { + kcp.ssthresh = IKCP_THRESH_MIN + } + kcp.cwnd = kcp.ssthresh + resent + kcp.incr = kcp.cwnd * kcp.mss } - kcp.cwnd = kcp.ssthresh + resent - kcp.incr = kcp.cwnd * kcp.mss - } - // congestion control, https://tools.ietf.org/html/rfc5681 - if lost > 0 { - kcp.ssthresh = cwnd / 2 - if kcp.ssthresh < IKCP_THRESH_MIN { - kcp.ssthresh = IKCP_THRESH_MIN + // congestion control, https://tools.ietf.org/html/rfc5681 + if lost > 0 { + kcp.ssthresh = cwnd / 2 + if kcp.ssthresh < IKCP_THRESH_MIN { + kcp.ssthresh = IKCP_THRESH_MIN + } + kcp.cwnd = 1 + kcp.incr = kcp.mss } - kcp.cwnd = 1 - kcp.incr = kcp.mss - } - if kcp.cwnd < 1 { - kcp.cwnd = 1 - kcp.incr = kcp.mss + if kcp.cwnd < 1 { + kcp.cwnd = 1 + kcp.incr = kcp.mss + } } return uint32(minrto) @@ -1000,9 +1008,5 @@ func (kcp *KCP) WaitSnd() int { // remove front n elements from queue func (kcp *KCP) remove_front(q []segment, n int) []segment { newn := copy(q, q[n:]) - gc := q[newn:] - for k := range gc { - gc[k].data = nil // de-ref data - } return q[:newn] } diff --git a/kcp-go/sess.go b/kcp-go/sess.go index 0853b5b..4bded00 100755 --- a/kcp-go/sess.go +++ b/kcp-go/sess.go @@ -3,14 +3,16 @@ package kcp import ( "crypto/rand" "encoding/binary" - "github.com/pkg/errors" - "golang.org/x/net/ipv4" "hash/crc32" "log" "net" "sync" "sync/atomic" "time" + + "github.com/pkg/errors" + "golang.org/x/net/ipv4" + "golang.org/x/net/ipv6" ) type errTimeout struct { @@ -39,9 +41,6 @@ const ( // accept backlog acceptBacklog = 128 - - // prerouting(to session) queue - qlen = 128 ) const ( @@ -94,7 +93,8 @@ type ( die chan struct{} // notify current session has Closed chReadEvent chan struct{} // notify Read() can be called without blocking chWriteEvent chan struct{} // notify Write() can be called without blocking - chErrorEvent chan error // notify Read() have an error + chReadError chan error // notify PacketConn.Read() have an error + chWriteError chan error // notify PacketConn.Write() have an error // nonce generator nonce Entropy @@ -120,7 +120,8 @@ func newUDPSession(conv uint32, dataShards, parityShards int, l *Listener, conn sess.nonce.Init() sess.chReadEvent = make(chan struct{}, 1) sess.chWriteEvent = make(chan struct{}, 1) - sess.chErrorEvent = make(chan error, 1) + sess.chReadError = make(chan error, 1) + sess.chWriteError = make(chan error, 1) sess.remote = remote sess.conn = conn sess.l = l @@ -182,6 +183,7 @@ func (s *UDPSession) Read(b []byte) (n int, err error) { n = copy(b, s.bufptr) s.bufptr = s.bufptr[n:] s.mu.Unlock() + atomic.AddUint64(&DefaultSnmp.BytesReceived, uint64(n)) return n, nil } @@ -191,10 +193,10 @@ func (s *UDPSession) Read(b []byte) (n int, err error) { } if size := s.kcp.PeekSize(); size > 0 { // peek data size from kcp - atomic.AddUint64(&DefaultSnmp.BytesReceived, uint64(size)) if len(b) >= size { // receive data into 'b' directly s.kcp.Recv(b) s.mu.Unlock() + atomic.AddUint64(&DefaultSnmp.BytesReceived, uint64(size)) return size, nil } @@ -209,6 +211,7 @@ func (s *UDPSession) Read(b []byte) (n int, err error) { n = copy(b, s.recvbuf) // copy to 'b' s.bufptr = s.recvbuf[n:] // pointer update s.mu.Unlock() + atomic.AddUint64(&DefaultSnmp.BytesReceived, uint64(n)) return n, nil } @@ -232,7 +235,7 @@ func (s *UDPSession) Read(b []byte) (n int, err error) { case <-s.chReadEvent: case <-c: case <-s.die: - case err = <-s.chErrorEvent: + case err = <-s.chReadError: if timeout != nil { timeout.Stop() } @@ -296,6 +299,11 @@ func (s *UDPSession) Write(b []byte) (n int, err error) { case <-s.chWriteEvent: case <-c: case <-s.die: + case err = <-s.chWriteError: + if timeout != nil { + timeout.Stop() + } + return n, err } if timeout != nil { @@ -321,7 +329,7 @@ func (s *UDPSession) Close() error { s.isClosed = true atomic.AddUint64(&DefaultSnmp.CurrEstab, ^uint64(0)) if s.l == nil { // client socket close - //return s.conn.Close() + return s.conn.Close() } return nil } @@ -425,10 +433,11 @@ func (s *UDPSession) SetDSCP(dscp int) error { s.mu.Lock() defer s.mu.Unlock() if s.l == nil { - if nc, ok := s.conn.(*connectedUDPConn); ok { - return ipv4.NewConn(nc.UDPConn).SetTOS(dscp << 2) - } else if nc, ok := s.conn.(net.Conn); ok { - return ipv4.NewConn(nc).SetTOS(dscp << 2) + if nc, ok := s.conn.(net.Conn); ok { + if err := ipv4.NewConn(nc).SetTOS(dscp << 2); err != nil { + return ipv6.NewConn(nc).SetTrafficClass(dscp) + } + return nil } } return errors.New(errInvalidOperation) @@ -502,6 +511,8 @@ func (s *UDPSession) output(buf []byte) { if n, err := s.conn.WriteTo(ext, s.remote); err == nil { nbytes += n npkts++ + } else { + s.notifyWriteError(err) } } @@ -509,6 +520,8 @@ func (s *UDPSession) output(buf []byte) { if n, err := s.conn.WriteTo(ecc[k], s.remote); err == nil { nbytes += n npkts++ + } else { + s.notifyWriteError(err) } } atomic.AddUint64(&DefaultSnmp.OutPkts, uint64(npkts)) @@ -544,6 +557,13 @@ func (s *UDPSession) notifyWriteEvent() { } } +func (s *UDPSession) notifyWriteError(err error) { + select { + case s.chWriteError <- err: + default: + } +} + func (s *UDPSession) kcpInput(data []byte) { var kcpInErrors, fecErrs, fecRecovered, fecParityShards uint64 @@ -627,64 +647,52 @@ func (s *UDPSession) kcpInput(data []byte) { } } -func (s *UDPSession) receiver(ch chan<- inPacket) { - for { - data := xmitBuf.Get().([]byte)[:mtuLimit] - if n, from, err := s.conn.ReadFrom(data); err == nil && n >= s.headerSize+IKCP_OVERHEAD { - select { - case ch <- inPacket{from, data[:n]}: - case <-s.die: - return - } - } else if err != nil { - s.chErrorEvent <- err - return - } else { - atomic.AddUint64(&DefaultSnmp.InErrs, 1) - } - } -} - // the read loop for a client session func (s *UDPSession) readLoop() { - chPacket := make(chan inPacket, qlen) - go s.receiver(chPacket) + buf := make([]byte, mtuLimit) + var src string firstPacket := true for { - select { - case p := <-chPacket: - raw := p.data - data := p.data - from := p.from - dataValid := false - if firstPacket{ - log.Println("firstPacket from", from.String()) - } - if s.block != nil { - s.block.Decrypt(data, data) - data = data[nonceSize:] - checksum := crc32.ChecksumIEEE(data[crcSize:]) - if checksum == binary.LittleEndian.Uint32(data) { - data = data[crcSize:] - dataValid = true - } else { - atomic.AddUint64(&DefaultSnmp.InCsumErrors, 1) - } - } else if s.block == nil { - dataValid = true + if n, addr, err := s.conn.ReadFrom(buf); err == nil { + // make sure the packet is from the same source + if src == "" { // set source address + src = addr.String() + } else if addr.String() != src { + atomic.AddUint64(&DefaultSnmp.InErrs, 1) + continue } - if dataValid { - if firstPacket{ - log.Println("firstPacket valided", from.String()) - firstPacket = false - //remote upd ip may change - s.remote = from + if n >= s.headerSize+IKCP_OVERHEAD { + data := buf[:n] + dataValid := false + if s.block != nil { + s.block.Decrypt(data, data) + data = data[nonceSize:] + checksum := crc32.ChecksumIEEE(data[crcSize:]) + if checksum == binary.LittleEndian.Uint32(data) { + data = data[crcSize:] + dataValid = true + } else { + atomic.AddUint64(&DefaultSnmp.InCsumErrors, 1) + } + } else if s.block == nil { + dataValid = true } - s.kcpInput(data) + + if dataValid { + if firstPacket{ + log.Println("firstPacket valided from", addr.String()) + firstPacket = false + //remote upd ip may change + s.remote = addr + } + s.kcpInput(data) + } + } else { + atomic.AddUint64(&DefaultSnmp.InErrs, 1) } - xmitBuf.Put(raw) - case <-s.die: + } else { + s.chReadError <- err return } } @@ -700,19 +708,14 @@ type ( conn net.PacketConn // the underlying packet connection sessions map[string]*UDPSession // all sessions accepted by this Listener - chAccepts chan *UDPSession // Listen() backlog - chSessionClosed chan net.Addr // session close queue - headerSize int // the additional header to a KCP frame - die chan struct{} // notify the listener has closed - rd atomic.Value // read deadline for Accept() + sessionLock sync.Mutex + chAccepts chan *UDPSession // Listen() backlog + chSessionClosed chan net.Addr // session close queue + headerSize int // the additional header to a KCP frame + die chan struct{} // notify the listener has closed + rd atomic.Value // read deadline for Accept() wd atomic.Value } - - // a incoming packet definition - inPacket struct { - from net.Addr - data []byte - } ) // monitor incoming data for all connections of server @@ -720,93 +723,77 @@ func (l *Listener) monitor() { // a cache for session object last used var lastAddr string var lastSession *UDPSession - - chPacket := make(chan inPacket, qlen) - go l.receiver(chPacket) + buf := make([]byte, mtuLimit) for { - select { - case p := <-chPacket: - raw := p.data - data := p.data - from := p.from - dataValid := false - if l.block != nil { - l.block.Decrypt(data, data) - data = data[nonceSize:] - checksum := crc32.ChecksumIEEE(data[crcSize:]) - if checksum == binary.LittleEndian.Uint32(data) { - data = data[crcSize:] + if n, from, err := l.conn.ReadFrom(buf); err == nil { + if n >= l.headerSize+IKCP_OVERHEAD { + data := buf[:n] + dataValid := false + if l.block != nil { + l.block.Decrypt(data, data) + data = data[nonceSize:] + checksum := crc32.ChecksumIEEE(data[crcSize:]) + if checksum == binary.LittleEndian.Uint32(data) { + data = data[crcSize:] + dataValid = true + } else { + atomic.AddUint64(&DefaultSnmp.InCsumErrors, 1) + } + } else if l.block == nil { dataValid = true - } else { - atomic.AddUint64(&DefaultSnmp.InCsumErrors, 1) - } - } else if l.block == nil { - dataValid = true - } - - if dataValid { - addr := from.String() - var s *UDPSession - var ok bool - - // the packets received from an address always come in batch, - // cache the session for next packet, without querying map. - if addr == lastAddr { - s, ok = lastSession, true - } else if s, ok = l.sessions[addr]; ok { - lastSession = s - lastAddr = addr } - if !ok { // new session - if len(l.chAccepts) < cap(l.chAccepts) { // do not let the new sessions overwhelm accept queue - var conv uint32 - convValid := false - if l.fecDecoder != nil { - isfec := binary.LittleEndian.Uint16(data[4:]) - if isfec == typeData { - conv = binary.LittleEndian.Uint32(data[fecHeaderSizePlus2:]) + if dataValid { + addr := from.String() + var s *UDPSession + var ok bool + + // the packets received from an address always come in batch, + // cache the session for next packet, without querying map. + if addr == lastAddr { + s, ok = lastSession, true + } else { + l.sessionLock.Lock() + if s, ok = l.sessions[addr]; ok { + lastSession = s + lastAddr = addr + } + l.sessionLock.Unlock() + } + + if !ok { // new session + if len(l.chAccepts) < cap(l.chAccepts) { // do not let the new sessions overwhelm accept queue + var conv uint32 + convValid := false + if l.fecDecoder != nil { + isfec := binary.LittleEndian.Uint16(data[4:]) + if isfec == typeData { + conv = binary.LittleEndian.Uint32(data[fecHeaderSizePlus2:]) + convValid = true + } + } else { + conv = binary.LittleEndian.Uint32(data) convValid = true } - } else { - conv = binary.LittleEndian.Uint32(data) - convValid = true - } - if convValid { // creates a new session only if the 'conv' field in kcp is accessible - s := newUDPSession(conv, l.dataShards, l.parityShards, l, l.conn, from, l.block) - s.kcpInput(data) - l.sessions[addr] = s - l.chAccepts <- s + if convValid { // creates a new session only if the 'conv' field in kcp is accessible + s := newUDPSession(conv, l.dataShards, l.parityShards, l, l.conn, from, l.block) + s.kcpInput(data) + l.sessionLock.Lock() + l.sessions[addr] = s + l.sessionLock.Unlock() + l.chAccepts <- s + } } + } else { + s.kcpInput(data) } - } else { - s.kcpInput(data) } + } else { + atomic.AddUint64(&DefaultSnmp.InErrs, 1) } - - xmitBuf.Put(raw) - case deadlink := <-l.chSessionClosed: - delete(l.sessions, deadlink.String()) - case <-l.die: - return - } - } -} - -func (l *Listener) receiver(ch chan<- inPacket) { - for { - data := xmitBuf.Get().([]byte)[:mtuLimit] - if n, from, err := l.conn.ReadFrom(data); err == nil && n >= l.headerSize+IKCP_OVERHEAD { - select { - case ch <- inPacket{from, data[:n]}: - case <-l.die: - return - } - } else if err != nil { - return } else { - atomic.AddUint64(&DefaultSnmp.InErrs, 1) + return } } } @@ -830,7 +817,10 @@ func (l *Listener) SetWriteBuffer(bytes int) error { // SetDSCP sets the 6bit DSCP field of IP header func (l *Listener) SetDSCP(dscp int) error { if nc, ok := l.conn.(net.Conn); ok { - return ipv4.NewConn(nc).SetTOS(dscp << 2) + if err := ipv4.NewConn(nc).SetTOS(dscp << 2); err != nil { + return ipv6.NewConn(nc).SetTrafficClass(dscp) + } + return nil } return errors.New(errInvalidOperation) } @@ -883,13 +873,14 @@ func (l *Listener) Close() error { } // closeSession notify the listener that a session has closed -func (l *Listener) closeSession(remote net.Addr) bool { - select { - case l.chSessionClosed <- remote: +func (l *Listener) closeSession(remote net.Addr) (ret bool) { + l.sessionLock.Lock() + defer l.sessionLock.Unlock() + if _, ok := l.sessions[remote.String()]; ok { + delete(l.sessions, remote.String()) return true - case <-l.die: - return false } + return false } // Addr returns the listener's network address, The Addr returned is shared by all invocations of Addr, so do not modify it. @@ -943,17 +934,22 @@ func Dial(raddr string) (net.Conn, error) { return DialWithOptions(raddr, nil, 0 // DialWithOptions connects to the remote address "raddr" on the network "udp" with packet encryption func DialWithOptions(raddr string, block BlockCrypt, dataShards, parityShards int) (*UDPSession, error) { + // network type detection udpaddr, err := net.ResolveUDPAddr("udp", raddr) if err != nil { return nil, errors.Wrap(err, "net.ResolveUDPAddr") } + network := "udp4" + if udpaddr.IP.To4() == nil { + network = "udp" + } - udpconn, err := net.DialUDP("udp", nil, udpaddr) + conn, err := net.ListenUDP(network, nil) if err != nil { return nil, errors.Wrap(err, "net.DialUDP") } - return NewConn(raddr, block, dataShards, parityShards, &connectedUDPConn{udpconn}) + return NewConn(raddr, block, dataShards, parityShards, conn) } // NewConn establishes a session and talks KCP protocol over a packet connection. @@ -967,6 +963,7 @@ func NewConn(raddr string, block BlockCrypt, dataShards, parityShards int, conn binary.Read(rand.Reader, binary.LittleEndian, &convid) return newUDPSession(convid, dataShards, parityShards, nil, conn, udpaddr, block), nil } + func NewP2pConn(udpConn net.PacketConn, raddr string, block BlockCrypt, dataShards, parityShards int) (*UDPSession, error){ udpaddr, err := net.ResolveUDPAddr("udp", raddr) if err != nil { @@ -975,16 +972,9 @@ func NewP2pConn(udpConn net.PacketConn, raddr string, block BlockCrypt, dataShar return newUDPSession(0x1, dataShards, parityShards, nil, udpConn, udpaddr, block), nil } + // monotonic reference time point var refTime time.Time = time.Now() // currentMs returns current elasped monotonic milliseconds since program startup func currentMs() uint32 { return uint32(time.Now().Sub(refTime) / time.Millisecond) } - -// connectedUDPConn is a wrapper for net.UDPConn which converts WriteTo syscalls -// to Write syscalls that are 4 times faster on some OS'es. This should only be -// used for connections that were produced by a net.Dial* call. -type connectedUDPConn struct{ *net.UDPConn } - -// WriteTo redirects all writes to the Write syscall, which is 4 times faster. -func (c *connectedUDPConn) WriteTo(b []byte, addr net.Addr) (int, error) { return c.Write(b) } diff --git a/kcp-go/snmp.go b/kcp-go/snmp.go old mode 100755 new mode 100644 diff --git a/kcp-go/updater.go b/kcp-go/updater.go old mode 100755 new mode 100644 diff --git a/p2pclient/main.go b/p2pclient/main.go index 62abf09..faa8ee7 100755 --- a/p2pclient/main.go +++ b/p2pclient/main.go @@ -85,12 +85,12 @@ func main() { }, cli.StringFlag{ Name: "listentcp,l", - Value: ":12948", + Value: ":2022", Usage: "local listen address", }, cli.StringFlag{ Name: "remoteudp, r", - Value: "vps:29900", + Value: "127.0.0.1:4000", Usage: "kcp server address", }, cli.StringFlag{ @@ -301,8 +301,11 @@ func main() { go tcpListener(chTCPConn, &config) for{ peerAddr, err := getPeerAddr(&config) + time.Sleep(2*time.Second) if err == nil{ p2pHandle(&config, peerAddr, chTCPConn) + }else{ + time.Sleep(5*time.Second) } } } @@ -400,13 +403,8 @@ func getPeerAddr(config *Config)(string, error){ isServer = pair_s peerAddr = mess.Data log.Println("peer addr is ", peerAddr) - n, err := kcpConn.Write(finMess) - if err != nil { - log.Println("kcpConn.Write", err) - return "", err - } - log.Println("writen ", n) - time.Sleep(1*time.Second) + kcpConn.Write(finMess) + case "fin": return peerAddr, nil } } diff --git a/p2pserver/config.go b/p2pserver/config.go index 417df9c..f3a1f84 100755 --- a/p2pserver/config.go +++ b/p2pserver/config.go @@ -8,7 +8,7 @@ import ( // Config for server type Config struct { Listen string `json:"listen"` - Key string `json:"key"` + Passwd string `json:"passwd"` Crypt string `json:"crypt"` Mode string `json:"mode"` MTU int `json:"mtu"` diff --git a/p2pserver/main.go b/p2pserver/main.go index b95cc58..8092364 100755 --- a/p2pserver/main.go +++ b/p2pserver/main.go @@ -6,10 +6,9 @@ import ( "encoding/csv" "encoding/json" "fmt" + "github.com/hikaricai/p2p_tun/kcp-go" "log" "math/rand" - "net/http" - _ "net/http/pprof" "os" "sync" "sync/atomic" @@ -19,7 +18,6 @@ import ( "path/filepath" - "github.com/hikaricai/p2p_tun/kcp-go" "github.com/urfave/cli" ) @@ -54,7 +52,7 @@ func main() { Usage: "kcp server listen address", }, cli.StringFlag{ - Name: "key", + Name: "passwd", Value: "1234", Usage: "pre-shared secret between client and server", EnvVar: "KCPTUN_KEY", @@ -170,7 +168,7 @@ func main() { myApp.Action = func(c *cli.Context) error { config := Config{} config.Listen = c.String("listen") - config.Key = c.String("key") + config.Passwd = c.String("passwd") config.Crypt = c.String("crypt") config.Mode = c.String("mode") config.MTU = c.Int("mtu") @@ -220,7 +218,7 @@ func main() { log.Println("version:", VERSION) log.Println("initiating key derivation") - pass := pbkdf2.Key([]byte(config.Key), []byte(SALT), 4096, 32, sha1.New) + pass := pbkdf2.Key([]byte(config.Passwd), []byte(SALT), 4096, 32, sha1.New) var block kcp.BlockCrypt switch config.Crypt { case "sm4": @@ -281,9 +279,6 @@ func main() { } go snmpLogger(config.SnmpLog, config.SnmpPeriod) - if config.Pprof { - go http.ListenAndServe(":6060", nil) - } for { log.Println("listening new kcp") @@ -310,20 +305,26 @@ type DigHoleMess struct { Data string } -type Peer struct { +type P2PSession struct { addr string chPair chan string + conn1 *kcp.UDPSession + conn2 *kcp.UDPSession + chFin chan struct{} } -var keymap = make(map[string]*Peer) +var keymap = make(map[string]*P2PSession) var keymu sync.Mutex +var session *P2PSession; + func handleClient(conn *kcp.UDPSession) { reader := bufio.NewReader(conn) defer conn.Close() var dataReady int32 - var chThreadDie chan struct{}= make(chan struct{}) + var chThreadDie = make(chan struct{}) + var ok bool; defer close(chThreadDie) go timeout(conn, &dataReady, chThreadDie) for { @@ -333,31 +334,34 @@ func handleClient(conn *kcp.UDPSession) { return } var mess DigHoleMess - json.Unmarshal([]byte(line), &mess) + err = json.Unmarshal([]byte(line), &mess) + if err != nil{ + continue + } switch mess.Cmd { case "login": remoteAddr := conn.RemoteAddr().String() log.Println("login from ", remoteAddr) key := mess.Data - log.Println("key is ", key) + keymu.Lock() - peer, ok := keymap[key] + session, ok = keymap[key] if ok { - peerAddr := peer.addr + peerAddr := session.addr log.Println("find peer and addr is", peerAddr) delete(keymap, key) - keymu.Unlock() - peer.chPair <- remoteAddr + session.conn2 = conn + session.chPair <- remoteAddr jsonPairMess := phaseJsonMess("pair_c", peerAddr) conn.Write(jsonPairMess) } else { log.Println("no peer, registed") - peer := Peer{remoteAddr, make(chan string)} - keymap[key] = &peer - keymu.Unlock() - go notifyAddr(conn, &peer, chThreadDie) + session = &P2PSession{remoteAddr, make(chan string), conn, nil,make(chan struct{})} + keymap[key] = session + go p2pSessionHandler(session, chThreadDie) } + keymu.Unlock() case "ping": atomic.StoreInt32(&dataReady, 1) jsonPingMess := phaseJsonMess("ping", "hello") @@ -365,8 +369,7 @@ func handleClient(conn *kcp.UDPSession) { log.Println("rcv ping from ", conn.RemoteAddr().String()) case "fin": log.Println("fin from", conn.RemoteAddr().String()) - time.Sleep(time.Second) - return + session.chFin <- struct{}{} } } } @@ -389,15 +392,26 @@ func timeout(conn *kcp.UDPSession, dataReady *int32, chThreadDie chan struct{}){ } } -func notifyAddr(conn *kcp.UDPSession, peer *Peer, chThreadDie chan struct{}){ +func p2pSessionHandler(session *P2PSession, chThreadDie chan struct{}){ + finCnt :=0 for { select { case <-chThreadDie: return - case peerAddr := <-peer.chPair: + case peerAddr := <-session.chPair: jsonPairMess := phaseJsonMess("pair_s", peerAddr) - conn.Write(jsonPairMess) - return + session.conn1.Write(jsonPairMess) + case <-session.chFin: + finCnt++ + if finCnt == 2{ + jsonPairMess := phaseJsonMess("fin", "bye") + session.conn1.Write(jsonPairMess) + session.conn2.Write(jsonPairMess) + time.Sleep(time.Second) + session.conn1.Close() + session.conn2.Close() + return + } } } } diff --git a/smux/LICENSE b/smux/LICENSE deleted file mode 100755 index eed41ac..0000000 --- a/smux/LICENSE +++ /dev/null @@ -1,21 +0,0 @@ -MIT License - -Copyright (c) 2016-2017 Daniel Fu - -Permission is hereby granted, free of charge, to any person obtaining a copy -of this software and associated documentation files (the "Software"), to deal -in the Software without restriction, including without limitation the rights -to use, copy, modify, merge, publish, distribute, sublicense, and/or sell -copies of the Software, and to permit persons to whom the Software is -furnished to do so, subject to the following conditions: - -The above copyright notice and this permission notice shall be included in all -copies or substantial portions of the Software. - -THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR -IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, -FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE -AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER -LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, -OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE -SOFTWARE. diff --git a/smux/frame.go b/smux/frame.go deleted file mode 100755 index 71d3d44..0000000 --- a/smux/frame.go +++ /dev/null @@ -1,60 +0,0 @@ -package smux - -import ( - "encoding/binary" - "fmt" -) - -const ( - version = 1 -) - -const ( // cmds - cmdSYN byte = iota // stream open - cmdFIN // stream close, a.k.a EOF mark - cmdPSH // data push - cmdNOP // no operation -) - -const ( - sizeOfVer = 1 - sizeOfCmd = 1 - sizeOfLength = 2 - sizeOfSid = 4 - headerSize = sizeOfVer + sizeOfCmd + sizeOfSid + sizeOfLength -) - -// Frame defines a packet from or to be multiplexed into a single connection -type Frame struct { - ver byte - cmd byte - sid uint32 - data []byte -} - -func newFrame(cmd byte, sid uint32) Frame { - return Frame{ver: version, cmd: cmd, sid: sid} -} - -type rawHeader [headerSize]byte - -func (h rawHeader) Version() byte { - return h[0] -} - -func (h rawHeader) Cmd() byte { - return h[1] -} - -func (h rawHeader) Length() uint16 { - return binary.LittleEndian.Uint16(h[2:]) -} - -func (h rawHeader) StreamID() uint32 { - return binary.LittleEndian.Uint32(h[4:]) -} - -func (h rawHeader) String() string { - return fmt.Sprintf("Version:%d Cmd:%d StreamID:%d Length:%d", - h.Version(), h.Cmd(), h.StreamID(), h.Length()) -} diff --git a/smux/mux.go b/smux/mux.go deleted file mode 100755 index 3cc8f11..0000000 --- a/smux/mux.go +++ /dev/null @@ -1,80 +0,0 @@ -package smux - -import ( - "fmt" - "io" - "time" - - "github.com/pkg/errors" -) - -// Config is used to tune the Smux session -type Config struct { - // KeepAliveInterval is how often to send a NOP command to the remote - KeepAliveInterval time.Duration - - // KeepAliveTimeout is how long the session - // will be closed if no data has arrived - KeepAliveTimeout time.Duration - - // MaxFrameSize is used to control the maximum - // frame size to sent to the remote - MaxFrameSize int - - // MaxReceiveBuffer is used to control the maximum - // number of data in the buffer pool - MaxReceiveBuffer int -} - -// DefaultConfig is used to return a default configuration -func DefaultConfig() *Config { - return &Config{ - KeepAliveInterval: 10 * time.Second, - KeepAliveTimeout: 30 * time.Second, - MaxFrameSize: 32768, - MaxReceiveBuffer: 4194304, - } -} - -// VerifyConfig is used to verify the sanity of configuration -func VerifyConfig(config *Config) error { - if config.KeepAliveInterval == 0 { - return errors.New("keep-alive interval must be positive") - } - if config.KeepAliveTimeout < config.KeepAliveInterval { - return fmt.Errorf("keep-alive timeout must be larger than keep-alive interval") - } - if config.MaxFrameSize <= 0 { - return errors.New("max frame size must be positive") - } - if config.MaxFrameSize > 65535 { - return errors.New("max frame size must not be larger than 65535") - } - if config.MaxReceiveBuffer <= 0 { - return errors.New("max receive buffer must be positive") - } - return nil -} - -// Server is used to initialize a new server-side connection. -func Server(conn io.ReadWriteCloser, config *Config) (*Session, error) { - if config == nil { - config = DefaultConfig() - } - if err := VerifyConfig(config); err != nil { - return nil, err - } - return newSession(config, conn, false), nil -} - -// Client is used to initialize a new client-side connection. -func Client(conn io.ReadWriteCloser, config *Config) (*Session, error) { - if config == nil { - config = DefaultConfig() - } - - if err := VerifyConfig(config); err != nil { - return nil, err - } - return newSession(config, conn, true), nil -} diff --git a/smux/session.go b/smux/session.go deleted file mode 100755 index c29634e..0000000 --- a/smux/session.go +++ /dev/null @@ -1,350 +0,0 @@ -package smux - -import ( - "encoding/binary" - "io" - "sync" - "sync/atomic" - "time" - - "github.com/pkg/errors" -) - -const ( - defaultAcceptBacklog = 1024 -) - -const ( - errBrokenPipe = "broken pipe" - errInvalidProtocol = "invalid protocol version" - errGoAway = "stream id overflows, should start a new connection" -) - -type writeRequest struct { - frame Frame - result chan writeResult -} - -type writeResult struct { - n int - err error -} - -// Session defines a multiplexed connection for streams -type Session struct { - conn io.ReadWriteCloser - - config *Config - nextStreamID uint32 // next stream identifier - nextStreamIDLock sync.Mutex - - bucket int32 // token bucket - bucketNotify chan struct{} // used for waiting for tokens - - streams map[uint32]*Stream // all streams in this session - streamLock sync.Mutex // locks streams - - die chan struct{} // flag session has died - dieLock sync.Mutex - chAccepts chan *Stream - - dataReady int32 // flag data has arrived - - goAway int32 // flag id exhausted - - deadline atomic.Value - - writes chan writeRequest -} - -func newSession(config *Config, conn io.ReadWriteCloser, client bool) *Session { - s := new(Session) - s.die = make(chan struct{}) - s.conn = conn - s.config = config - s.streams = make(map[uint32]*Stream) - s.chAccepts = make(chan *Stream, defaultAcceptBacklog) - s.bucket = int32(config.MaxReceiveBuffer) - s.bucketNotify = make(chan struct{}, 1) - s.writes = make(chan writeRequest) - - if client { - s.nextStreamID = 1 - } else { - s.nextStreamID = 0 - } - go s.recvLoop() - go s.sendLoop() - go s.keepalive() - return s -} - -// OpenStream is used to create a new stream -func (s *Session) OpenStream() (*Stream, error) { - if s.IsClosed() { - return nil, errors.New(errBrokenPipe) - } - - // generate stream id - s.nextStreamIDLock.Lock() - if s.goAway > 0 { - s.nextStreamIDLock.Unlock() - return nil, errors.New(errGoAway) - } - - s.nextStreamID += 2 - sid := s.nextStreamID - if sid == sid%2 { // stream-id overflows - s.goAway = 1 - s.nextStreamIDLock.Unlock() - return nil, errors.New(errGoAway) - } - s.nextStreamIDLock.Unlock() - - stream := newStream(sid, s.config.MaxFrameSize, s) - - if _, err := s.writeFrame(newFrame(cmdSYN, sid)); err != nil { - return nil, errors.Wrap(err, "writeFrame") - } - - s.streamLock.Lock() - s.streams[sid] = stream - s.streamLock.Unlock() - return stream, nil -} - -// AcceptStream is used to block until the next available stream -// is ready to be accepted. -func (s *Session) AcceptStream() (*Stream, error) { - var deadline <-chan time.Time - if d, ok := s.deadline.Load().(time.Time); ok && !d.IsZero() { - timer := time.NewTimer(time.Until(d)) - defer timer.Stop() - deadline = timer.C - } - select { - case stream := <-s.chAccepts: - return stream, nil - case <-deadline: - return nil, errTimeout - case <-s.die: - return nil, errors.New(errBrokenPipe) - } -} - -// Close is used to close the session and all streams. -func (s *Session) Close() (err error) { - s.dieLock.Lock() - - select { - case <-s.die: - s.dieLock.Unlock() - return errors.New(errBrokenPipe) - default: - close(s.die) - s.dieLock.Unlock() - s.streamLock.Lock() - for k := range s.streams { - s.streams[k].sessionClose() - } - s.streamLock.Unlock() - s.notifyBucket() - return s.conn.Close() - } -} - -// notifyBucket notifies recvLoop that bucket is available -func (s *Session) notifyBucket() { - select { - case s.bucketNotify <- struct{}{}: - default: - } -} - -// IsClosed does a safe check to see if we have shutdown -func (s *Session) IsClosed() bool { - select { - case <-s.die: - return true - default: - return false - } -} - -// NumStreams returns the number of currently open streams -func (s *Session) NumStreams() int { - if s.IsClosed() { - return 0 - } - s.streamLock.Lock() - defer s.streamLock.Unlock() - return len(s.streams) -} - -// SetDeadline sets a deadline used by Accept* calls. -// A zero time value disables the deadline. -func (s *Session) SetDeadline(t time.Time) error { - s.deadline.Store(t) - return nil -} - -// notify the session that a stream has closed -func (s *Session) streamClosed(sid uint32) { - s.streamLock.Lock() - if n := s.streams[sid].recycleTokens(); n > 0 { // return remaining tokens to the bucket - if atomic.AddInt32(&s.bucket, int32(n)) > 0 { - s.notifyBucket() - } - } - delete(s.streams, sid) - s.streamLock.Unlock() -} - -// returnTokens is called by stream to return token after read -func (s *Session) returnTokens(n int) { - if atomic.AddInt32(&s.bucket, int32(n)) > 0 { - s.notifyBucket() - } -} - -// session read a frame from underlying connection -// it's data is pointed to the input buffer -func (s *Session) readFrame(buffer []byte) (f Frame, err error) { - var hdr rawHeader - if _, err := io.ReadFull(s.conn, hdr[:]); err != nil { - return f, errors.Wrap(err, "readFrame") - } - - if hdr.Version() != version { - return f, errors.New(errInvalidProtocol) - } - - f.ver = hdr.Version() - f.cmd = hdr.Cmd() - f.sid = hdr.StreamID() - if length := hdr.Length(); length > 0 { - f.data = buffer[:length] - if _, err := io.ReadFull(s.conn, f.data); err != nil { - return f, errors.Wrap(err, "readFrame") - } - } - return f, nil -} - -// recvLoop keeps on reading from underlying connection if tokens are available -func (s *Session) recvLoop() { - buffer := make([]byte, 1<<16) - for { - for atomic.LoadInt32(&s.bucket) <= 0 && !s.IsClosed() { - <-s.bucketNotify - } - - if f, err := s.readFrame(buffer); err == nil { - atomic.StoreInt32(&s.dataReady, 1) - - switch f.cmd { - case cmdNOP: - case cmdSYN: - s.streamLock.Lock() - if _, ok := s.streams[f.sid]; !ok { - stream := newStream(f.sid, s.config.MaxFrameSize, s) - s.streams[f.sid] = stream - select { - case s.chAccepts <- stream: - case <-s.die: - } - } - s.streamLock.Unlock() - case cmdFIN: - s.streamLock.Lock() - if stream, ok := s.streams[f.sid]; ok { - stream.markRST() - stream.notifyReadEvent() - } - s.streamLock.Unlock() - case cmdPSH: - s.streamLock.Lock() - if stream, ok := s.streams[f.sid]; ok { - atomic.AddInt32(&s.bucket, -int32(len(f.data))) - stream.pushBytes(f.data) - stream.notifyReadEvent() - } - s.streamLock.Unlock() - default: - s.Close() - return - } - } else { - s.Close() - return - } - } -} - -func (s *Session) keepalive() { - tickerPing := time.NewTicker(s.config.KeepAliveInterval) - tickerTimeout := time.NewTicker(s.config.KeepAliveTimeout) - defer tickerPing.Stop() - defer tickerTimeout.Stop() - for { - select { - case <-tickerPing.C: - s.writeFrame(newFrame(cmdNOP, 0)) - s.notifyBucket() // force a signal to the recvLoop - case <-tickerTimeout.C: - if !atomic.CompareAndSwapInt32(&s.dataReady, 1, 0) { - s.Close() - return - } - case <-s.die: - return - } - } -} - -func (s *Session) sendLoop() { - buf := make([]byte, (1<<16)+headerSize) - for { - select { - case <-s.die: - return - case request := <-s.writes: - buf[0] = request.frame.ver - buf[1] = request.frame.cmd - binary.LittleEndian.PutUint16(buf[2:], uint16(len(request.frame.data))) - binary.LittleEndian.PutUint32(buf[4:], request.frame.sid) - copy(buf[headerSize:], request.frame.data) - n, err := s.conn.Write(buf[:headerSize+len(request.frame.data)]) - - n -= headerSize - if n < 0 { - n = 0 - } - - result := writeResult{ - n: n, - err: err, - } - - request.result <- result - close(request.result) - } - } -} - -// writeFrame writes the frame to the underlying connection -// and returns the number of bytes written if successful -func (s *Session) writeFrame(f Frame) (n int, err error) { - req := writeRequest{ - frame: f, - result: make(chan writeResult, 1), - } - select { - case <-s.die: - return 0, errors.New(errBrokenPipe) - case s.writes <- req: - } - - result := <-req.result - return result.n, result.err -} diff --git a/smux/stream.go b/smux/stream.go deleted file mode 100755 index 2ce00d2..0000000 --- a/smux/stream.go +++ /dev/null @@ -1,262 +0,0 @@ -package smux - -import ( - "bytes" - "io" - "net" - "sync" - "sync/atomic" - "time" - - "github.com/pkg/errors" -) - -// Stream implements net.Conn -type Stream struct { - id uint32 - rstflag int32 - sess *Session - buffer bytes.Buffer - bufferLock sync.Mutex - frameSize int - chReadEvent chan struct{} // notify a read event - die chan struct{} // flag the stream has closed - dieLock sync.Mutex - readDeadline atomic.Value - writeDeadline atomic.Value -} - -// newStream initiates a Stream struct -func newStream(id uint32, frameSize int, sess *Session) *Stream { - s := new(Stream) - s.id = id - s.chReadEvent = make(chan struct{}, 1) - s.frameSize = frameSize - s.sess = sess - s.die = make(chan struct{}) - return s -} - -// ID returns the unique stream ID. -func (s *Stream) ID() uint32 { - return s.id -} - -// Read implements net.Conn -func (s *Stream) Read(b []byte) (n int, err error) { - if len(b) == 0 { - select { - case <-s.die: - return 0, errors.New(errBrokenPipe) - default: - return 0, nil - } - } - - var deadline <-chan time.Time - if d, ok := s.readDeadline.Load().(time.Time); ok && !d.IsZero() { - timer := time.NewTimer(time.Until(d)) - defer timer.Stop() - deadline = timer.C - } - -READ: - s.bufferLock.Lock() - n, _ = s.buffer.Read(b) - s.bufferLock.Unlock() - - if n > 0 { - s.sess.returnTokens(n) - return n, nil - } else if atomic.LoadInt32(&s.rstflag) == 1 { - _ = s.Close() - return 0, io.EOF - } - - select { - case <-s.chReadEvent: - goto READ - case <-deadline: - return n, errTimeout - case <-s.die: - return 0, errors.New(errBrokenPipe) - } -} - -// Write implements net.Conn -func (s *Stream) Write(b []byte) (n int, err error) { - var deadline <-chan time.Time - if d, ok := s.writeDeadline.Load().(time.Time); ok && !d.IsZero() { - timer := time.NewTimer(time.Until(d)) - defer timer.Stop() - deadline = timer.C - } - - select { - case <-s.die: - return 0, errors.New(errBrokenPipe) - default: - } - - frames := s.split(b, cmdPSH, s.id) - sent := 0 - for k := range frames { - req := writeRequest{ - frame: frames[k], - result: make(chan writeResult, 1), - } - - select { - case s.sess.writes <- req: - case <-s.die: - return sent, errors.New(errBrokenPipe) - case <-deadline: - return sent, errTimeout - } - - select { - case result := <-req.result: - sent += result.n - if result.err != nil { - return sent, result.err - } - case <-s.die: - return sent, errors.New(errBrokenPipe) - case <-deadline: - return sent, errTimeout - } - } - return sent, nil -} - -// Close implements net.Conn -func (s *Stream) Close() error { - s.dieLock.Lock() - - select { - case <-s.die: - s.dieLock.Unlock() - return errors.New(errBrokenPipe) - default: - close(s.die) - s.dieLock.Unlock() - s.sess.streamClosed(s.id) - _, err := s.sess.writeFrame(newFrame(cmdFIN, s.id)) - return err - } -} - -// SetReadDeadline sets the read deadline as defined by -// net.Conn.SetReadDeadline. -// A zero time value disables the deadline. -func (s *Stream) SetReadDeadline(t time.Time) error { - s.readDeadline.Store(t) - return nil -} - -// SetWriteDeadline sets the write deadline as defined by -// net.Conn.SetWriteDeadline. -// A zero time value disables the deadline. -func (s *Stream) SetWriteDeadline(t time.Time) error { - s.writeDeadline.Store(t) - return nil -} - -// SetDeadline sets both read and write deadlines as defined by -// net.Conn.SetDeadline. -// A zero time value disables the deadlines. -func (s *Stream) SetDeadline(t time.Time) error { - if err := s.SetReadDeadline(t); err != nil { - return err - } - if err := s.SetWriteDeadline(t); err != nil { - return err - } - return nil -} - -// session closes the stream -func (s *Stream) sessionClose() { - s.dieLock.Lock() - defer s.dieLock.Unlock() - - select { - case <-s.die: - default: - close(s.die) - } -} - -// LocalAddr satisfies net.Conn interface -func (s *Stream) LocalAddr() net.Addr { - if ts, ok := s.sess.conn.(interface { - LocalAddr() net.Addr - }); ok { - return ts.LocalAddr() - } - return nil -} - -// RemoteAddr satisfies net.Conn interface -func (s *Stream) RemoteAddr() net.Addr { - if ts, ok := s.sess.conn.(interface { - RemoteAddr() net.Addr - }); ok { - return ts.RemoteAddr() - } - return nil -} - -// pushBytes a slice into buffer -func (s *Stream) pushBytes(p []byte) { - s.bufferLock.Lock() - s.buffer.Write(p) - s.bufferLock.Unlock() -} - -// recycleTokens transform remaining bytes to tokens(will truncate buffer) -func (s *Stream) recycleTokens() (n int) { - s.bufferLock.Lock() - n = s.buffer.Len() - s.buffer.Reset() - s.bufferLock.Unlock() - return -} - -// split large byte buffer into smaller frames, reference only -func (s *Stream) split(bts []byte, cmd byte, sid uint32) []Frame { - frames := make([]Frame, 0, len(bts)/s.frameSize+1) - for len(bts) > s.frameSize { - frame := newFrame(cmd, sid) - frame.data = bts[:s.frameSize] - bts = bts[s.frameSize:] - frames = append(frames, frame) - } - if len(bts) > 0 { - frame := newFrame(cmd, sid) - frame.data = bts - frames = append(frames, frame) - } - return frames -} - -// notify read event -func (s *Stream) notifyReadEvent() { - select { - case s.chReadEvent <- struct{}{}: - default: - } -} - -// mark this stream has been reset -func (s *Stream) markRST() { - atomic.StoreInt32(&s.rstflag, 1) -} - -var errTimeout error = &timeoutError{} - -type timeoutError struct{} - -func (e *timeoutError) Error() string { return "i/o timeout" } -func (e *timeoutError) Timeout() bool { return true } -func (e *timeoutError) Temporary() bool { return true }