返回 DeepSeek-Reasonix
sshtest.go
根目录 / internal / remote / sshtest / sshtest.go
1 // Package sshtest is an in-process SSH server for exercising the remote
2 // module without a real sshd. It supports publickey and password auth, session
3 // exec with scripted responses, direct-tcpip (for -L forwards), tcpip-forward
4 // (for -R forwards), and an SFTP subsystem via pkg/sftp's server. It is
5 // test-only.
6 package sshtest
7
8 import (
9 "errors"
10 "fmt"
11 "io"
12 "net"
13 "sync"
14 "testing"
15 "time"
16
17 "github.com/pkg/sftp"
18 "golang.org/x/crypto/ssh"
19 )
20
21 // Server is a running in-process SSH server.
22 type Server struct {
23 Addr string
24 HostKey ssh.Signer
25 config *ssh.ServerConfig
26 listener net.Listener
27 execFunc func(cmd string) (stdout string, stderr string, exit int)
28 sftpRoot string
29 enableSFT bool
30
31 mu sync.Mutex
32 conns []net.Conn
33 listeners []net.Listener
34 wg sync.WaitGroup
35
36 // stdinMu guards lastStdin, the session script the client streams after
37 // the exec request (models sshd's forwarding to the remote process's
38 // stdin). Test code reads it to assert the payload arrived intact.
39 stdinMu sync.Mutex
40 lastStdin []byte
41 }
42
43 // Options configures a test server.
44 type Options struct {
45 // HostKeys, when non-empty, are offered by the server. The first key is
46 // also exposed as Server.HostKey. Empty generates one ed25519 key.
47 HostKeys []ssh.Signer
48 // Password, when non-empty, enables password auth accepting (any user,
49 // this password).
50 Password string
51 // AuthorizedKey, when set, enables publickey auth accepting this key.
52 AuthorizedKey ssh.PublicKey
53 // Exec handles `exec` requests; nil => a default echoing the command.
54 Exec func(cmd string) (stdout string, stderr string, exit int)
55 // SFTPRoot enables the SFTP subsystem rooted at this directory.
56 SFTPRoot string
57 }
58
59 // Start launches a server on 127.0.0.1:0.
60 func Start(t *testing.T, opts Options) *Server {
61 t.Helper()
62 hostKeys := opts.HostKeys
63 if len(hostKeys) == 0 {
64 hostKey, err := generateHostKey()
65 if err != nil {
66 t.Fatalf("host key: %v", err)
67 }
68 hostKeys = []ssh.Signer{hostKey}
69 }
70 cfg := &ssh.ServerConfig{}
71 for _, hostKey := range hostKeys {
72 if hostKey == nil {
73 t.Fatal("host key must not be nil")
74 }
75 cfg.AddHostKey(hostKey)
76 }
77 if opts.Password != "" {
78 cfg.PasswordCallback = func(conn ssh.ConnMetadata, pass []byte) (*ssh.Permissions, error) {
79 if string(pass) == opts.Password {
80 return &ssh.Permissions{}, nil
81 }
82 return nil, errors.New("bad password")
83 }
84 }
85 if opts.AuthorizedKey != nil {
86 want := opts.AuthorizedKey.Marshal()
87 cfg.PublicKeyCallback = func(conn ssh.ConnMetadata, key ssh.PublicKey) (*ssh.Permissions, error) {
88 if string(key.Marshal()) == string(want) {
89 return &ssh.Permissions{}, nil
90 }
91 return nil, errors.New("unknown key")
92 }
93 }
94 if opts.Password == "" && opts.AuthorizedKey == nil {
95 cfg.NoClientAuth = true
96 }
97
98 ln, err := net.Listen("tcp", "127.0.0.1:0")
99 if err != nil {
100 t.Fatalf("listen: %v", err)
101 }
102 s := &Server{
103 Addr: ln.Addr().String(),
104 HostKey: hostKeys[0],
105 config: cfg,
106 listener: ln,
107 execFunc: opts.Exec,
108 sftpRoot: opts.SFTPRoot,
109 enableSFT: opts.SFTPRoot != "",
110 }
111 s.wg.Add(1)
112 go s.serve()
113 t.Cleanup(s.Close)
114 return s
115 }
116
117 // Close stops the server and all active connections.
118 func (s *Server) Close() {
119 _ = s.listener.Close()
120 s.mu.Lock()
121 for _, c := range s.conns {
122 _ = c.Close()
123 }
124 for _, ln := range s.listeners {
125 _ = ln.Close()
126 }
127 s.conns = nil
128 s.listeners = nil
129 s.mu.Unlock()
130 s.wg.Wait()
131 }
132
133 // DropConnections closes every currently-open client connection without
134 // stopping the server, simulating a network drop so a supervised Client must
135 // reconnect.
136 func (s *Server) DropConnections() {
137 s.mu.Lock()
138 conns := s.conns
139 s.conns = nil
140 s.mu.Unlock()
141 for _, c := range conns {
142 _ = c.Close()
143 }
144 }
145
146 func (s *Server) serve() {
147 defer s.wg.Done()
148 for {
149 nConn, err := s.listener.Accept()
150 if err != nil {
151 return
152 }
153 s.mu.Lock()
154 s.conns = append(s.conns, nConn)
155 s.mu.Unlock()
156 s.wg.Go(func() {
157 s.handleConn(nConn)
158 })
159 }
160 }
161
162 func (s *Server) handleConn(nConn net.Conn) {
163 sshConn, chans, reqs, err := ssh.NewServerConn(nConn, s.config)
164 if err != nil {
165 return
166 }
167 defer sshConn.Close()
168 go s.handleGlobalRequests(sshConn, reqs)
169 for newCh := range chans {
170 switch newCh.ChannelType() {
171 case "session":
172 go s.handleSession(newCh)
173 case "direct-tcpip":
174 go s.handleDirectTCPIP(newCh)
175 default:
176 _ = newCh.Reject(ssh.UnknownChannelType, "unsupported")
177 }
178 }
179 }
180
181 func (s *Server) handleGlobalRequests(conn *ssh.ServerConn, reqs <-chan *ssh.Request) {
182 for req := range reqs {
183 switch req.Type {
184 case "keepalive@openssh.com":
185 if req.WantReply {
186 _ = req.Reply(true, nil)
187 }
188 case "tcpip-forward":
189 s.handleTCPIPForward(conn, req)
190 case "cancel-tcpip-forward":
191 if req.WantReply {
192 _ = req.Reply(true, nil)
193 }
194 default:
195 if req.WantReply {
196 _ = req.Reply(false, nil)
197 }
198 }
199 }
200 }
201
202 func (s *Server) handleSession(newCh ssh.NewChannel) {
203 ch, reqs, err := newCh.Accept()
204 if err != nil {
205 return
206 }
207 defer ch.Close()
208 for req := range reqs {
209 switch req.Type {
210 case "exec":
211 cmd := parseStringPayload(req.Payload)
212 // The client streams the script after the exec request; model the
213 // same ordering: wait for the first data frame, then run the
214 // handler. Reading until EOF would deadlock.
215 first := make(chan struct{})
216 go func() {
217 buf := make([]byte, 64*1024)
218 n, _ := ch.Read(buf)
219 if n > 0 {
220 s.stdinMu.Lock()
221 s.lastStdin = append([]byte(nil), buf[:n]...)
222 s.stdinMu.Unlock()
223 }
224 close(first)
225 }()
226 if req.WantReply {
227 _ = req.Reply(true, nil)
228 }
229 select {
230 case <-first:
231 case <-time.After(500 * time.Millisecond):
232 }
233 s.runExec(ch, cmd)
234 return
235 case "subsystem":
236 name := parseStringPayload(req.Payload)
237 if name == "sftp" && s.enableSFT {
238 if req.WantReply {
239 _ = req.Reply(true, nil)
240 }
241 s.runSFTP(ch)
242 return
243 }
244 if req.WantReply {
245 _ = req.Reply(false, nil)
246 }
247 case "shell", "pty-req", "env":
248 if req.WantReply {
249 _ = req.Reply(true, nil)
250 }
251 default:
252 if req.WantReply {
253 _ = req.Reply(false, nil)
254 }
255 }
256 }
257 }
258
259 // LastStdin returns the script streamed by the most recent exec request,
260 // once the session channel has closed (the same completion point a remote
261 // process observes). Test code polls it to assert the payload arrived intact.
262 func (s *Server) LastStdin() string {
263 s.stdinMu.Lock()
264 defer s.stdinMu.Unlock()
265 return string(s.lastStdin)
266 }
267
268 func (s *Server) runExec(ch ssh.Channel, cmd string) {
269 stdout, stderr, exit := "", "", 0
270 if s.execFunc != nil {
271 stdout, stderr, exit = s.execFunc(cmd)
272 } else {
273 stdout = cmd
274 }
275 _, _ = io.WriteString(ch, stdout)
276 if stderr != "" {
277 _, _ = io.WriteString(ch.Stderr(), stderr)
278 }
279 sendExitStatus(ch, exit)
280 }
281
282 func (s *Server) runSFTP(ch ssh.Channel) {
283 var server *sftp.Server
284 var err error
285 if s.sftpRoot != "" {
286 server, err = sftp.NewServer(ch, sftp.WithServerWorkingDirectory(s.sftpRoot))
287 } else {
288 server, err = sftp.NewServer(ch)
289 }
290 if err != nil {
291 return
292 }
293 _ = server.Serve()
294 _ = server.Close()
295 }
296
297 // handleDirectTCPIP implements -L forwards: dial the requested target and
298 // splice.
299 func (s *Server) handleDirectTCPIP(newCh ssh.NewChannel) {
300 var payload struct {
301 HostToConnect string
302 PortToConnect uint32
303 OriginatorHost string
304 OriginatorPort uint32
305 }
306 if err := ssh.Unmarshal(newCh.ExtraData(), &payload); err != nil {
307 _ = newCh.Reject(ssh.ConnectionFailed, "bad payload")
308 return
309 }
310 target := net.JoinHostPort(payload.HostToConnect, fmt.Sprintf("%d", payload.PortToConnect))
311 dst, err := net.Dial("tcp", target)
312 if err != nil {
313 _ = newCh.Reject(ssh.ConnectionFailed, err.Error())
314 return
315 }
316 ch, reqs, err := newCh.Accept()
317 if err != nil {
318 _ = dst.Close()
319 return
320 }
321 go ssh.DiscardRequests(reqs)
322 splice(ch, dst)
323 }
324
325 // handleTCPIPForward implements -R forwards: listen locally on the server and
326 // open a forwarded-tcpip channel back to the client for each accepted conn.
327 func (s *Server) handleTCPIPForward(conn *ssh.ServerConn, req *ssh.Request) {
328 var payload struct {
329 BindAddr string
330 BindPort uint32
331 }
332 if err := ssh.Unmarshal(req.Payload, &payload); err != nil {
333 if req.WantReply {
334 _ = req.Reply(false, nil)
335 }
336 return
337 }
338 ln, err := net.Listen("tcp", net.JoinHostPort(payload.BindAddr, fmt.Sprintf("%d", payload.BindPort)))
339 if err != nil {
340 if req.WantReply {
341 _ = req.Reply(false, nil)
342 }
343 return
344 }
345 boundPort := uint32(ln.Addr().(*net.TCPAddr).Port)
346 s.mu.Lock()
347 s.listeners = append(s.listeners, ln)
348 s.mu.Unlock()
349 if req.WantReply {
350 _ = req.Reply(true, ssh.Marshal(struct{ Port uint32 }{boundPort}))
351 }
352 go func() {
353 for {
354 c, err := ln.Accept()
355 if err != nil {
356 return
357 }
358 go func() {
359 origPort := uint32(1)
360 if ta, ok := c.RemoteAddr().(*net.TCPAddr); ok && ta.Port > 0 {
361 origPort = uint32(ta.Port)
362 }
363 msg := struct {
364 ConnHost string
365 ConnPort uint32
366 OrigHost string
367 OrigPort uint32
368 }{payload.BindAddr, boundPort, "127.0.0.1", origPort}
369 ch, reqs, err := conn.OpenChannel("forwarded-tcpip", ssh.Marshal(msg))
370 if err != nil {
371 _ = c.Close()
372 return
373 }
374 go ssh.DiscardRequests(reqs)
375 splice(ch, c)
376 }()
377 }
378 }()
379 }
380
381 func splice(a io.ReadWriteCloser, b net.Conn) {
382 done := make(chan struct{}, 2)
383 go func() { _, _ = io.Copy(a, b); done <- struct{}{} }()
384 go func() { _, _ = io.Copy(b, a); done <- struct{}{} }()
385 <-done
386 _ = a.Close()
387 _ = b.Close()
388 }
389
390 func parseStringPayload(p []byte) string {
391 if len(p) < 4 {
392 return ""
393 }
394 n := int(p[0])<<24 | int(p[1])<<16 | int(p[2])<<8 | int(p[3])
395 if 4+n > len(p) {
396 return ""
397 }
398 return string(p[4 : 4+n])
399 }
400
401 func sendExitStatus(ch ssh.Channel, code int) {
402 _, _ = ch.SendRequest("exit-status", false, ssh.Marshal(struct{ Status uint32 }{uint32(code)}))
403 }
404
404 lines GO