| @@ -18,6 +18,7 @@ import ( |
| 18 | "strconv" |
18 | "strconv" |
| 19 | "strings" |
19 | "strings" |
| 20 | "sync" |
20 | "sync" |
| |
21 | "sync/atomic" |
| 21 | "time" |
22 | "time" |
| 22 | |
23 | |
| 23 | "golang.org/x/crypto/ssh" |
24 | "golang.org/x/crypto/ssh" |
| @@ -37,10 +38,21 @@ type Server struct { |
| 37 | sshCfg *ssh.ServerConfig |
38 | sshCfg *ssh.ServerConfig |
| 38 | authLimiter *rateLimiter |
39 | authLimiter *rateLimiter |
| 39 | sessions sync.WaitGroup // accepted connections still being served |
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 | func New(cfg config.Config, st *store.Store) (*Server, error) { |
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 | sc := &ssh.ServerConfig{ |
57 | sc := &ssh.ServerConfig{ |
| 46 | PublicKeyCallback: s.authenticate, |
58 | PublicKeyCallback: s.authenticate, |
| @@ -116,6 +128,13 @@ func (s *Server) authenticate(meta ssh.ConnMetadata, pub ssh.PublicKey) (*ssh.Pe |
| 116 | } |
128 | } |
| 117 | fp := ssh.FingerprintSHA256(pub) |
129 | fp := ssh.FingerprintSHA256(pub) |
| 118 | key, err := s.st.SSHKeyByFingerprint(fp) |
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 | if err != nil { |
138 | if err != nil { |
| 120 | if s.cfg.Registration.Mode != "closed" { |
139 | if s.cfg.Registration.Mode != "closed" { |
| 121 | return &ssh.Permissions{Extensions: map[string]string{ |
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 | // Serve accepts connections on ln until it is closed. |
157 | // Serve accepts connections on ln until it is closed. |
| 139 | func (s *Server) Serve(ln net.Listener) error { |
158 | func (s *Server) Serve(ln net.Listener) error { |
| 140 | for { |
159 | for { |
| 141 | conn, err := ln.Accept() |
160 | nc, err := ln.Accept() |
| 142 | if err != nil { |
161 | if err != nil { |
| 143 | return err |
162 | return err |
| 144 | } |
163 | } |
| |
164 | c := &conn{net: nc} |
| |
165 | s.mu.Lock() |
| |
166 | s.conns[c] = struct{}{} |
| |
167 | s.mu.Unlock() |
| 145 | s.sessions.Add(1) |
168 | s.sessions.Add(1) |
| 146 | go func() { |
169 | go func() { |
| 147 | defer s.sessions.Done() |
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 |
181 | // Shutdown closes every idle connection, then waits for the ones with a |
| 154 | // caller closes the listener first; a push in flight completes rather |
182 | // session running, or for ctx. The caller closes the listener first; a |
| 155 | // than being cut mid-pack. |
183 | // push in flight completes rather than being cut mid-pack. |
| 156 | func (s *Server) Shutdown(ctx context.Context) error { |
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 | done := make(chan struct{}) |
192 | done := make(chan struct{}) |
| 158 | go func() { |
193 | go func() { |
| 159 | s.sessions.Wait() |
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) { |
205 | func (s *Server) handleConn(c *conn) { |
| 171 | defer conn.Close() |
206 | defer c.net.Close() |
| 172 | sconn, chans, reqs, err := ssh.NewServerConn(conn, s.sshCfg) |
207 | sconn, chans, reqs, err := ssh.NewServerConn(c.net, s.sshCfg) |
| 173 | if err != nil { |
208 | if err != nil { |
| 174 | return |
209 | return |
| 175 | } |
210 | } |
| @@ -185,7 +220,11 @@ func (s *Server) handleConn(conn net.Conn) { |
| 185 | if err != nil { |
220 | if err != nil { |
| 186 | continue |
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 | |