Просмотр исходного кода

cleanly close out session on disconnect

Mike 3 лет назад
Родитель
Сommit
04ebc297f3
5 измененных файлов с 66 добавлено и 43 удалено
  1. 2 1
      cmd/main.go
  2. 4 4
      oscar/buddy.go
  3. 6 6
      oscar/icbm.go
  4. 32 29
      oscar/protocol.go
  5. 22 3
      oscar/session.go

+ 2 - 1
cmd/main.go

@@ -109,6 +109,8 @@ func handleAuthConnection(sm *oscar.SessionManager, conn net.Conn) {
 }
 
 func handleBOSConnection(sm *oscar.SessionManager, fm *oscar.FeedbagStore, conn net.Conn, seq *uint32) {
+	defer conn.Close()
+
 	fmt.Println("VerifyLogin...")
 	sess, err := oscar.VerifyLogin(sm, conn, seq)
 	if err != nil {
@@ -127,7 +129,6 @@ func handleBOSConnection(sm *oscar.SessionManager, fm *oscar.FeedbagStore, conn
 	if err := oscar.ReadBos(sm, sess, fm, conn, seq); err != nil && err != io.EOF {
 		if err != io.EOF {
 			fmt.Println(err.Error())
-			os.Exit(1)
 		}
 	}
 }

+ 4 - 4
oscar/buddy.go

@@ -173,11 +173,11 @@ func NotifyArrival(sess *Session, sm *SessionManager, fm *FeedbagStore) error {
 			},
 		}
 
-		adjSess.MsgChan <- &XMessage{
+		adjSess.SendMessage(&XMessage{
 			flap:      flap,
 			snacFrame: snacFrameOut,
 			snacOut:   snacPayloadOut,
-		}
+		})
 	}
 	return nil
 }
@@ -207,11 +207,11 @@ func NotifyDeparture(sess *Session, sm *SessionManager, fm *FeedbagStore) error
 			},
 		}
 
-		adjSess.MsgChan <- &XMessage{
+		adjSess.SendMessage(&XMessage{
 			flap:      flap,
 			snacFrame: snacFrameOut,
 			snacOut:   snacPayloadOut,
-		}
+		})
 	}
 	return nil
 }

+ 6 - 6
oscar/icbm.go

@@ -277,7 +277,7 @@ func SendAndReceiveChannelMsgTohost(sm *SessionManager, fm *FeedbagStore, sess *
 		mm.snacOut.(*snacClientIM).TLVs = append(mm.snacOut.(*snacClientIM).TLVs, t)
 	}
 
-	session.MsgChan <- mm
+	session.SendMessage(mm)
 
 	snacFrameOut := snacFrame{
 		foodGroup: ICBM,
@@ -403,7 +403,7 @@ func SendIM(sm *SessionManager, sess *Session, destScreenName string, msg string
 		return err
 	}
 
-	session.MsgChan <- &XMessage{
+	session.SendMessage(&XMessage{
 		flap: &flapFrame{
 			startMarker: 42,
 			frameType:   2,
@@ -428,7 +428,7 @@ func SendIM(sm *SessionManager, sess *Session, destScreenName string, msg string
 				},
 			},
 		},
-	}
+	})
 
 	return nil
 }
@@ -538,7 +538,7 @@ func SendAndReceiveClientEvent(sm *SessionManager, fm *FeedbagStore, sess *Sessi
 
 	snacPayloadIn.screenName = sess.ScreenName
 
-	session.MsgChan <- &XMessage{
+	session.SendMessage(&XMessage{
 		flap: &flapFrame{
 			startMarker: 42,
 			frameType:   2,
@@ -548,7 +548,7 @@ func SendAndReceiveClientEvent(sm *SessionManager, fm *FeedbagStore, sess *Sessi
 			subGroup:  ICBMClientEvent,
 		},
 		snacOut: snacPayloadIn,
-	}
+	})
 
 	return nil
 }
@@ -653,7 +653,7 @@ func SendAndReceiveEvilRequest(sm *SessionManager, fm *FeedbagStore, sess *Sessi
 		mm.snacOut.(*snacEvilNotification).screenName = sess.ScreenName
 	}
 
-	recipSess.MsgChan <- mm
+	recipSess.SendMessage(mm)
 
 	return NotifyArrival(recipSess, sm, fm)
 }

+ 32 - 29
oscar/protocol.go

@@ -457,75 +457,78 @@ const (
 	FlapFrameKeepAlive       = 0x05
 )
 
-func streamFromConn(rw io.ReadWriter, in chan IncomingMessage) {
+func readIncomingRequests(rw io.Reader, msCh chan IncomingMessage, errCh chan error) {
+	defer close(msCh)
+	defer close(errCh)
+
 	for {
 		flap := &flapFrame{}
 		if err := flap.read(rw); err != nil {
-			if err != io.EOF {
-				panic("temporary panic: " + err.Error())
-			} else {
-				break
-			}
+			errCh <- err
+			return
 		}
 
 		switch flap.frameType {
 		case FlapFrameSignon:
-			panic("shouldn't get FlapFrameSignon")
+			errCh <- errors.New("shouldn't get FlapFrameSignon")
+			return
 		case FlapFrameData:
 			b := make([]byte, flap.payloadLength)
 			if _, err := rw.Read(b); err != nil {
-				if err != io.EOF {
-					panic("temporary panic: " + err.Error())
-				} else {
-					break
-				}
+				errCh <- err
+				return
 			}
 
 			buf := bytes.NewBuffer(b)
 
 			snac := &snacFrame{}
 			if err := snac.read(buf); err != nil {
-				if err != io.EOF {
-					panic("temporary panic: " + err.Error())
-				} else {
-					break
-				}
+				errCh <- err
+				return
 			}
 
-			in <- IncomingMessage{
+			msCh <- IncomingMessage{
 				flap: flap,
 				snac: snac,
 				buf:  buf,
 			}
 		case FlapFrameError:
-			panic(fmt.Sprintf("got FlapFrameError: %v", flap))
+			errCh <- fmt.Errorf("got FlapFrameError: %v", flap)
+			return
 		case FlapFrameSignoff:
-			panic(fmt.Sprintf("got signoff: %v", flap))
+			errCh <- fmt.Errorf("got signoff: %v", flap)
+			return
 		case FlapFrameKeepAlive:
 			fmt.Println("keepalive heartbeat")
+			return
 		default:
-			panic(fmt.Sprintf("unknown frame type: %v", flap))
+			errCh <- fmt.Errorf("unknown frame type: %v", flap)
+			return
 		}
 	}
 }
 
-func ReadBos(sm *SessionManager, sess *Session, fm *FeedbagStore, rw io.ReadWriter, sequence *uint32) error {
-	in := make(chan IncomingMessage)
-	defer close(in)
+func ReadBos(sm *SessionManager, sess *Session, fm *FeedbagStore, rw io.ReadWriteCloser, sequence *uint32) error {
+	defer rw.Close()
+	defer sess.Close()
 
-	go streamFromConn(rw, in)
+	// buffered so that the go routine has room to exit
+	msgCh := make(chan IncomingMessage, 1)
+	errCh := make(chan error, 1)
+	go readIncomingRequests(rw, msgCh, errCh)
 
 	for {
 		select {
-		case m := <-in:
-			err := routeIncomingRequests(sm, sess, fm, rw, sequence, m.snac, m.flap, m.buf)
-			if err != nil {
+		case m := <-msgCh:
+			if err := routeIncomingRequests(sm, sess, fm, rw, sequence, m.snac, m.flap, m.buf); err != nil {
 				return err
 			}
-		case m := <-sess.MsgChan:
+		case m := <-sess.RecvMessage():
 			if err := writeOutSNAC(nil, m.flap, m.snacFrame, m.snacOut, sequence, rw); err != nil {
 				panic("error handling handleXMessage: " + err.Error())
 			}
+		case err := <-errCh:
+			return err
 		}
 	}
 }

+ 22 - 3
oscar/session.go

@@ -12,7 +12,8 @@ var errSessNotFound = errors.New("session was not found")
 type Session struct {
 	ID          string
 	ScreenName  string
-	MsgChan     chan *XMessage
+	msgCh       chan *XMessage
+	stopCh      chan struct{}
 	Mutex       sync.RWMutex
 	Warning     uint16
 	AwayMessage string
@@ -71,6 +72,23 @@ func (s *Session) GetWarning() uint16 {
 	return w
 }
 
+func (s *Session) RecvMessage() chan *XMessage {
+	return s.msgCh
+}
+
+func (s *Session) SendMessage(msg *XMessage) {
+	select {
+	case <-s.stopCh:
+		return
+	case s.msgCh <- msg:
+	}
+}
+
+func (s *Session) Close() {
+	fmt.Println("closing out session")
+	close(s.stopCh)
+}
+
 type SessionManager struct {
 	store    map[string]*Session
 	mapMutex sync.RWMutex
@@ -122,8 +140,9 @@ func (s *SessionManager) NewSession() (*Session, error) {
 		return nil, err
 	}
 	sess := &Session{
-		ID:      id.String(),
-		MsgChan: make(chan *XMessage, 1),
+		ID:     id.String(),
+		msgCh:  make(chan *XMessage, 1),
+		stopCh: make(chan struct{}),
 	}
 	s.store[sess.ID] = sess
 	return sess, nil