cf5a2ed15f854bc2bbc6db11f5c2b970353b9db2

Author
Toby Padilla <toby@charm.sh>
Committer
Toby Padilla <toby@charm.sh>
Date

Message

Use middleware for session handling

Diff

  1diff --git a/main.go b/main.go
  2index 46830542f28470051d12e571a464322d88dffef2..143b8c85058209c2d8379a67a2ddbb67668a7bff 100644
  3--- a/main.go
  4+++ b/main.go
  5@@ -17,7 +17,7 @@ func main() {
  6 	if err != nil {
  7 		panic(err)
  8 	}
  9-	s, err := NewServer(cfg.Port, cfg.KeyPath, tui.SessionHandler)
 10+	s, err := NewServer(cfg.Port, cfg.KeyPath, LoggingMiddleware(), BubbleTeaMiddleware(tui.SessionHandler))
 11 	if err != nil {
 12 		panic(err)
 13 	}
 14diff --git a/server.go b/server.go
 15index efbaed0031c2213751c7538725edfea97bde6ce4..ea735c41c1304a1eeb46ae294fb6dbba8f263403 100644
 16--- a/server.go
 17+++ b/server.go
 18@@ -12,22 +12,44 @@ import (
 19 	gossh "golang.org/x/crypto/ssh"
 20 )
 21 
 22-type SessionHandler func(ssh.Session) (tea.Model, error)
 23+type Middleware func(ssh.Handler) ssh.Handler
 24 
 25-type Server struct {
 26-	server  *ssh.Server
 27-	key     gossh.PublicKey
 28-	handler SessionHandler
 29+func LoggingMiddleware() Middleware {
 30+	return func(sh ssh.Handler) ssh.Handler {
 31+		return func(s ssh.Session) {
 32+			hpk := s.PublicKey() != nil
 33+			log.Printf("%s connect %v %v\n", s.RemoteAddr().String(), hpk, s.Command())
 34+			sh(s)
 35+			log.Printf("%s disconnect %v %v\n", s.RemoteAddr().String(), hpk, s.Command())
 36+		}
 37+	}
 38 }
 39 
 40-func NewServer(port int, keyPath string, handler SessionHandler) (*Server, error) {
 41-	s := &Server{
 42-		server:  &ssh.Server{},
 43-		handler: handler,
 44+func BubbleTeaMiddleware(bth func(ssh.Session) tea.Model) Middleware {
 45+	return func(sh ssh.Handler) ssh.Handler {
 46+		return func(s ssh.Session) {
 47+			m := bth(s)
 48+			if m != nil {
 49+				p := tea.NewProgram(m, tea.WithAltScreen(), tea.WithInput(s), tea.WithOutput(s))
 50+				err := p.Start()
 51+				if err != nil {
 52+					log.Printf("%s error %v: %s\n", s.RemoteAddr().String(), s.Command(), err)
 53+				}
 54+			}
 55+			sh(s)
 56+		}
 57 	}
 58+}
 59+
 60+type Server struct {
 61+	server *ssh.Server
 62+	key    gossh.PublicKey
 63+}
 64+
 65+func NewServer(port int, keyPath string, mw ...Middleware) (*Server, error) {
 66+	s := &Server{server: &ssh.Server{}}
 67 	s.server.Version = "OpenSSH_7.6p1"
 68 	s.server.Addr = fmt.Sprintf(":%d", port)
 69-	s.server.Handler = s.sessionHandler
 70 	s.server.PasswordHandler = s.passHandler
 71 	s.server.PublicKeyHandler = s.authHandler
 72 	kps := strings.Split(keyPath, string(filepath.Separator))
 73@@ -42,28 +64,15 @@ func NewServer(port int, keyPath string, handler SessionHandler) (*Server, error
 74 	if err != nil {
 75 		return nil, err
 76 	}
 77+	h := func(s ssh.Session) {}
 78+	for _, m := range mw {
 79+		h = m(h)
 80+	}
 81+	s.server.Handler = h
 82 	return s, nil
 83 }
 84 
 85 func (srv *Server) sessionHandler(s ssh.Session) {
 86-	hpk := s.PublicKey() != nil
 87-	log.Printf("%s connect %v %v\n", s.RemoteAddr().String(), hpk, s.Command())
 88-	m, err := srv.handler(s)
 89-	if err != nil {
 90-		log.Printf("%s error %v %s\n", s.RemoteAddr().String(), hpk, err)
 91-		s.Exit(1)
 92-		return
 93-	}
 94-	if m != nil {
 95-		p := tea.NewProgram(m, tea.WithAltScreen(), tea.WithInput(s), tea.WithOutput(s))
 96-		err = p.Start()
 97-		if err != nil {
 98-			log.Printf("%s error %v %s\n", s.RemoteAddr().String(), hpk, err)
 99-			s.Exit(1)
100-			return
101-		}
102-	}
103-	log.Printf("%s disconnect %v %v\n", s.RemoteAddr().String(), hpk, s.Command())
104 }
105 
106 func (srv *Server) authHandler(ctx ssh.Context, key ssh.PublicKey) bool {
107diff --git a/tui/model.go b/tui/model.go
108index 4a0e1d2366b58bf3745a7ebbbf6ba5d4b77630b2..e29f72f5574a96cde86a196eead3488f5b3d5c66 100644
109--- a/tui/model.go
110+++ b/tui/model.go
111@@ -24,12 +24,12 @@ func (e errMsg) Error() string {
112 	return e.err.Error()
113 }
114 
115-func SessionHandler(s ssh.Session) (tea.Model, error) {
116+func SessionHandler(s ssh.Session) tea.Model {
117 	pty, changes, active := s.Pty()
118 	if !active {
119-		return nil, fmt.Errorf("you need to do this from a terminal with PTY support")
120+		return nil
121 	}
122-	return NewModel(pty.Window.Width, pty.Window.Height, changes), nil
123+	return NewModel(pty.Window.Width, pty.Window.Height, changes)
124 }
125 
126 type Model struct {