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 {