| @@ -18,6 +18,7 @@ import ( |
| 18 | 18 | "strconv" |
| 19 | 19 | "strings" |
| 20 | 20 | "sync" |
| 21 | "sync/atomic" |
| 21 | 22 | "time" |
| 22 | 23 | |
| 23 | 24 | "golang.org/x/crypto/ssh" |
| @@ -37,10 +38,21 @@ type Server struct { |
| 37 | 38 | sshCfg *ssh.ServerConfig |
| 38 | 39 | authLimiter *rateLimiter |
| 39 | 40 | sessions sync.WaitGroup // accepted connections still being served |
| 41 | mu sync.Mutex |
| 42 | conns map[*conn]struct{} |
| 43 | } |
| 44 | |
| 45 | // conn is one accepted connection and how many sessions it is running. |
| 46 | // A CLI's shared connection sits idle between commands; on shutdown an |
| 47 | // idle connection is closed at once and only a session mid-command is |
| 48 | // waited for (#141). |
| 49 | type conn struct { |
| 50 | net net.Conn |
| 51 | active atomic.Int32 |
| 40 | 52 | } |
| 41 | 53 | |
| 42 | 54 | func New(cfg config.Config, st *store.Store) (*Server, error) { |
| 43 | | s := &Server{cfg: cfg, st: st, authLimiter: newRateLimiter(cfg.Limits.SSHAuthRate, time.Minute)} |
| 55 | s := &Server{cfg: cfg, st: st, authLimiter: newRateLimiter(cfg.Limits.SSHAuthRate, time.Minute), conns: map[*conn]struct{}{}} |
| 44 | 56 | |
| 45 | 57 | sc := &ssh.ServerConfig{ |
| 46 | 58 | PublicKeyCallback: s.authenticate, |
| @@ -116,6 +128,13 @@ func (s *Server) authenticate(meta ssh.ConnMetadata, pub ssh.PublicKey) (*ssh.Pe |
| 116 | 128 | } |
| 117 | 129 | fp := ssh.FingerprintSHA256(pub) |
| 118 | 130 | key, err := s.st.SSHKeyByFingerprint(fp) |
| 131 | if err != nil && !errors.Is(err, store.ErrNotFound) { |
| 132 | // The store, not the key, failed. Neither a failure against the |
| 133 | // limiter nor "unknown key": a busy database during a restart |
| 134 | // would otherwise lock every client out for a minute. |
| 135 | slog.Error("ssh auth: key lookup", "err", err) |
| 136 | return nil, fmt.Errorf("authentication temporarily unavailable") |
| 137 | } |
| 119 | 138 | if err != nil { |
| 120 | 139 | if s.cfg.Registration.Mode != "closed" { |
| 121 | 140 | return &ssh.Permissions{Extensions: map[string]string{ |
| @@ -138,22 +157,38 @@ func (s *Server) authenticate(meta ssh.ConnMetadata, pub ssh.PublicKey) (*ssh.Pe |
| 138 | 157 | // Serve accepts connections on ln until it is closed. |
| 139 | 158 | func (s *Server) Serve(ln net.Listener) error { |
| 140 | 159 | for { |
| 141 | | conn, err := ln.Accept() |
| 160 | nc, err := ln.Accept() |
| 142 | 161 | if err != nil { |
| 143 | 162 | return err |
| 144 | 163 | } |
| 164 | c := &conn{net: nc} |
| 165 | s.mu.Lock() |
| 166 | s.conns[c] = struct{}{} |
| 167 | s.mu.Unlock() |
| 145 | 168 | s.sessions.Add(1) |
| 146 | 169 | go func() { |
| 147 | 170 | defer s.sessions.Done() |
| 148 | | s.handleConn(conn) |
| 171 | defer func() { |
| 172 | s.mu.Lock() |
| 173 | delete(s.conns, c) |
| 174 | s.mu.Unlock() |
| 175 | }() |
| 176 | s.handleConn(c) |
| 149 | 177 | }() |
| 150 | 178 | } |
| 151 | 179 | } |
| 152 | 180 | |
| 153 | | // Shutdown waits for every accepted connection to finish, or for ctx. The |
| 154 | | // caller closes the listener first; a push in flight completes rather |
| 155 | | // than being cut mid-pack. |
| 181 | // Shutdown closes every idle connection, then waits for the ones with a |
| 182 | // session running, or for ctx. The caller closes the listener first; a |
| 183 | // push in flight completes rather than being cut mid-pack. |
| 156 | 184 | func (s *Server) Shutdown(ctx context.Context) error { |
| 185 | s.mu.Lock() |
| 186 | for c := range s.conns { |
| 187 | if c.active.Load() == 0 { |
| 188 | c.net.Close() |
| 189 | } |
| 190 | } |
| 191 | s.mu.Unlock() |
| 157 | 192 | done := make(chan struct{}) |
| 158 | 193 | go func() { |
| 159 | 194 | s.sessions.Wait() |
| @@ -167,9 +202,9 @@ func (s *Server) Shutdown(ctx context.Context) error { |
| 167 | 202 | } |
| 168 | 203 | } |
| 169 | 204 | |
| 170 | | func (s *Server) handleConn(conn net.Conn) { |
| 171 | | defer conn.Close() |
| 172 | | sconn, chans, reqs, err := ssh.NewServerConn(conn, s.sshCfg) |
| 205 | func (s *Server) handleConn(c *conn) { |
| 206 | defer c.net.Close() |
| 207 | sconn, chans, reqs, err := ssh.NewServerConn(c.net, s.sshCfg) |
| 173 | 208 | if err != nil { |
| 174 | 209 | return |
| 175 | 210 | } |
| @@ -185,7 +220,11 @@ func (s *Server) handleConn(conn net.Conn) { |
| 185 | 220 | if err != nil { |
| 186 | 221 | continue |
| 187 | 222 | } |
| 188 | | go s.handleSession(sconn, ch, chReqs) |
| 223 | c.active.Add(1) |
| 224 | go func() { |
| 225 | defer c.active.Add(-1) |
| 226 | s.handleSession(sconn, ch, chReqs) |
| 227 | }() |
| 189 | 228 | } |
| 190 | 229 | } |
| 191 | 230 | |