返回 DeepSeek-Reasonix
set_test.go
根目录 / internal / remote / forward / set_test.go
1 package forward
2
3 import (
4 "errors"
5 "fmt"
6 "io"
7 "net"
8 "testing"
9 "time"
10
11 "golang.org/x/crypto/ssh"
12
13 "reasonix/internal/remote/sshtest"
14 )
15
16 // dialSSHClient connects to the sshtest server as a real ssh client.
17 func dialSSHClient(t *testing.T, srv *sshtest.Server) *ssh.Client {
18 t.Helper()
19 cfg := &ssh.ClientConfig{
20 User: "test",
21 HostKeyCallback: ssh.InsecureIgnoreHostKey(),
22 Timeout: 5 * time.Second,
23 }
24 cl, err := ssh.Dial("tcp", srv.Addr, cfg)
25 if err != nil {
26 t.Fatalf("ssh dial: %v", err)
27 }
28 t.Cleanup(func() { cl.Close() })
29 return cl
30 }
31
32 // echoServer starts a local TCP echo server for -L target testing.
33 func echoServer(t *testing.T) string {
34 t.Helper()
35 ln, err := net.Listen("tcp", "127.0.0.1:0")
36 if err != nil {
37 t.Fatal(err)
38 }
39 t.Cleanup(func() { ln.Close() })
40 go func() {
41 for {
42 c, err := ln.Accept()
43 if err != nil {
44 return
45 }
46 go func() { _, _ = io.Copy(c, c); c.Close() }()
47 }
48 }()
49 return ln.Addr().String()
50 }
51
52 func TestLocalForwardEndToEnd(t *testing.T) {
53 srv := sshtest.Start(t, sshtest.Options{})
54 cl := dialSSHClient(t, srv)
55 target := echoServer(t)
56
57 set := NewSet(nil)
58 defer set.Close()
59 if err := set.Attach(cl); err != nil {
60 t.Fatalf("attach: %v", err)
61 }
62 bound, err := set.Add(Spec{Direction: Local, BindAddr: "127.0.0.1:0", TargetAddr: target})
63 if err != nil {
64 t.Fatalf("add local forward: %v", err)
65 }
66 if bound == "" {
67 t.Fatal("no bound address")
68 }
69
70 conn, err := net.Dial("tcp", bound)
71 if err != nil {
72 t.Fatalf("dial forward: %v", err)
73 }
74 defer conn.Close()
75 if _, err := conn.Write([]byte("ping")); err != nil {
76 t.Fatal(err)
77 }
78 buf := make([]byte, 4)
79 _ = conn.SetReadDeadline(time.Now().Add(5 * time.Second))
80 if _, err := io.ReadFull(conn, buf); err != nil {
81 t.Fatalf("read echo: %v", err)
82 }
83 if string(buf) != "ping" {
84 t.Fatalf("echo = %q, want ping", buf)
85 }
86 }
87
88 func TestLocalListenerPersistsAcrossReattach(t *testing.T) {
89 srv := sshtest.Start(t, sshtest.Options{})
90 target := echoServer(t)
91
92 set := NewSet(nil)
93 defer set.Close()
94 cl1 := dialSSHClient(t, srv)
95 if err := set.Attach(cl1); err != nil {
96 t.Fatal(err)
97 }
98 bound, err := set.Add(Spec{Direction: Local, BindAddr: "127.0.0.1:0", TargetAddr: target})
99 if err != nil {
100 t.Fatal(err)
101 }
102
103 // Simulate a connection drop then reconnect on a new client.
104 set.Detach()
105 cl2 := dialSSHClient(t, srv)
106 if err := set.Attach(cl2); err != nil {
107 t.Fatal(err)
108 }
109
110 // The bound address must be unchanged (listener stayed open).
111 entries := set.List()
112 if len(entries) != 1 || entries[0].BoundAddr != bound {
113 t.Fatalf("bound address changed across reattach: %+v (was %s)", entries, bound)
114 }
115
116 // And traffic works again through the new connection.
117 conn, err := net.Dial("tcp", bound)
118 if err != nil {
119 t.Fatalf("dial after reattach: %v", err)
120 }
121 defer conn.Close()
122 _, _ = conn.Write([]byte("pong"))
123 buf := make([]byte, 4)
124 _ = conn.SetReadDeadline(time.Now().Add(5 * time.Second))
125 if _, err := io.ReadFull(conn, buf); err != nil {
126 t.Fatalf("read echo after reattach: %v", err)
127 }
128 if string(buf) != "pong" {
129 t.Fatalf("echo = %q", buf)
130 }
131 }
132
133 func TestDuplicateForwardRejected(t *testing.T) {
134 set := NewSet(nil)
135 defer set.Close()
136 spec := Spec{Name: "web", Direction: Local, BindAddr: "127.0.0.1:0", TargetAddr: "svc:80"}
137 if _, err := set.Add(spec); err != nil {
138 t.Fatal(err)
139 }
140 if _, err := set.Add(spec); !errors.Is(err, ErrDuplicateForward) {
141 t.Fatalf("second add err = %v, want ErrDuplicateForward", err)
142 }
143 }
144
145 func TestReplaceSwapsLiveForwardAfterReplacementStarts(t *testing.T) {
146 srv := sshtest.Start(t, sshtest.Options{})
147 cl := dialSSHClient(t, srv)
148 firstTarget := echoServer(t)
149 secondTarget := echoServer(t)
150 set := NewSet(nil)
151 defer set.Close()
152 if err := set.Attach(cl); err != nil {
153 t.Fatal(err)
154 }
155 firstBound, err := set.Add(Spec{Name: "serve", Direction: Local, BindAddr: "127.0.0.1:0", TargetAddr: firstTarget})
156 if err != nil {
157 t.Fatal(err)
158 }
159 secondBound, err := set.Replace(Spec{Name: "serve", Direction: Local, BindAddr: "127.0.0.1:0", TargetAddr: secondTarget})
160 if err != nil {
161 t.Fatal(err)
162 }
163 if firstBound == secondBound {
164 t.Fatalf("replacement reused old listener %q", firstBound)
165 }
166 entries := set.List()
167 if len(entries) != 1 || entries[0].Spec.TargetAddr != secondTarget || !entries[0].Up {
168 t.Fatalf("replacement registry = %+v", entries)
169 }
170 if conn, err := net.DialTimeout("tcp", firstBound, 100*time.Millisecond); err == nil {
171 _ = conn.Close()
172 t.Fatalf("old listener %q is still accepting", firstBound)
173 }
174 conn, err := net.Dial("tcp", secondBound)
175 if err != nil {
176 t.Fatalf("dial replacement: %v", err)
177 }
178 _ = conn.Close()
179 }
180
181 func TestReplaceFailurePreservesExistingForward(t *testing.T) {
182 srv := sshtest.Start(t, sshtest.Options{})
183 cl := dialSSHClient(t, srv)
184 target := echoServer(t)
185 occupied, err := net.Listen("tcp", "127.0.0.1:0")
186 if err != nil {
187 t.Fatal(err)
188 }
189 defer occupied.Close()
190 set := NewSet(nil)
191 defer set.Close()
192 if err := set.Attach(cl); err != nil {
193 t.Fatal(err)
194 }
195 bound, err := set.Add(Spec{Name: "serve", Direction: Local, BindAddr: "127.0.0.1:0", TargetAddr: target})
196 if err != nil {
197 t.Fatal(err)
198 }
199 if _, err := set.Replace(Spec{Name: "serve", Direction: Local, BindAddr: occupied.Addr().String(), TargetAddr: "other:80"}); err == nil {
200 t.Fatal("Replace unexpectedly bound an occupied address")
201 }
202 entries := set.List()
203 if len(entries) != 1 || entries[0].BoundAddr != bound || entries[0].Spec.TargetAddr != target || !entries[0].Up {
204 t.Fatalf("failed replacement disturbed existing forward: %+v", entries)
205 }
206 }
207
208 func TestReplaceWhileDetachedPreservesExistingForward(t *testing.T) {
209 set := NewSet(nil)
210 defer set.Close()
211 old := Spec{Name: "serve", Direction: Local, BindAddr: "127.0.0.1:0", TargetAddr: "old:80"}
212 if _, err := set.Add(old); err != nil {
213 t.Fatal(err)
214 }
215 if _, err := set.Replace(Spec{Name: "serve", Direction: Local, BindAddr: "127.0.0.1:0", TargetAddr: "new:80"}); !errors.Is(err, ErrNotAttached) {
216 t.Fatalf("Replace error = %v, want ErrNotAttached", err)
217 }
218 entries := set.List()
219 if len(entries) != 1 || entries[0].Spec.TargetAddr != old.TargetAddr {
220 t.Fatalf("detached replacement disturbed old forward: %+v", entries)
221 }
222 }
223
224 func TestBindBusyReported(t *testing.T) {
225 srv := sshtest.Start(t, sshtest.Options{})
226 cl := dialSSHClient(t, srv)
227 // Occupy a port.
228 occupied, err := net.Listen("tcp", "127.0.0.1:0")
229 if err != nil {
230 t.Fatal(err)
231 }
232 defer occupied.Close()
233 busyAddr := occupied.Addr().String()
234
235 set := NewSet(nil)
236 defer set.Close()
237 if err := set.Attach(cl); err != nil {
238 t.Fatal(err)
239 }
240 _, err = set.Add(Spec{Direction: Local, BindAddr: busyAddr, TargetAddr: "svc:80"})
241 if err == nil {
242 t.Fatal("expected bind-busy error")
243 }
244 if !errors.Is(err, ErrBindBusy) {
245 t.Fatalf("err = %v, want ErrBindBusy", err)
246 }
247 }
248
249 func TestRemoteForwardEndToEnd(t *testing.T) {
250 srv := sshtest.Start(t, sshtest.Options{})
251 cl := dialSSHClient(t, srv)
252 target := echoServer(t)
253
254 events := make(chan Event, 8)
255 set := NewSet(func(e Event) { events <- e })
256 defer set.Close()
257 if err := set.Attach(cl); err != nil {
258 t.Fatal(err)
259 }
260 // -R: sshtest listens on its side and forwards back to our local target.
261 if _, err := set.Add(Spec{Direction: Remote, BindAddr: "127.0.0.1:0", TargetAddr: target}); err != nil {
262 t.Fatalf("add remote forward: %v", err)
263 }
264
265 // Find the remote bound address from the registry.
266 var bound string
267 deadline := time.After(5 * time.Second)
268 for bound == "" {
269 select {
270 case <-deadline:
271 t.Fatal("remote forward never came up")
272 default:
273 }
274 for _, e := range set.List() {
275 if e.Up {
276 bound = e.BoundAddr
277 }
278 }
279 if bound == "" {
280 time.Sleep(20 * time.Millisecond)
281 }
282 }
283
284 conn, err := net.Dial("tcp", bound)
285 if err != nil {
286 t.Fatalf("dial remote-forward bind: %v", err)
287 }
288 defer conn.Close()
289 _, _ = conn.Write([]byte("rrrr"))
290 buf := make([]byte, 4)
291 _ = conn.SetReadDeadline(time.Now().Add(5 * time.Second))
292 if _, err := io.ReadFull(conn, buf); err != nil {
293 t.Fatalf("read echo via -R: %v", err)
294 }
295 if string(buf) != "rrrr" {
296 t.Fatalf("echo = %q", buf)
297 }
298 }
299
300 // A remote listener that dies under an attached Set must clear Up. Leaving the
301 // entry marked Up with a dead BoundAddr makes credential-proxy ensure reuse the
302 // stale port and the remote serve dials connection refused forever.
303 func TestRemoteForwardAcceptExitMarksDown(t *testing.T) {
304 srv := sshtest.Start(t, sshtest.Options{})
305 cl := dialSSHClient(t, srv)
306 target := echoServer(t)
307
308 set := NewSet(nil)
309 defer set.Close()
310 if err := set.Attach(cl); err != nil {
311 t.Fatal(err)
312 }
313 if _, err := set.Add(Spec{Name: "cred", Direction: Remote, BindAddr: "127.0.0.1:0", TargetAddr: target}); err != nil {
314 t.Fatalf("add remote forward: %v", err)
315 }
316 deadline := time.After(5 * time.Second)
317 for {
318 entries := set.List()
319 if len(entries) == 1 && entries[0].Up && entries[0].BoundAddr != "" {
320 break
321 }
322 select {
323 case <-deadline:
324 t.Fatalf("remote forward never came up: %+v", entries)
325 case <-time.After(20 * time.Millisecond):
326 }
327 }
328
329 // Close the SSH client without Detach: Accept returns, and the registry
330 // must not keep advertising the dead listener as Up.
331 _ = cl.Close()
332 deadline = time.After(5 * time.Second)
333 for {
334 entries := set.List()
335 if len(entries) == 1 && !entries[0].Up {
336 return
337 }
338 select {
339 case <-deadline:
340 t.Fatalf("dead remote forward still Up: %+v", entries)
341 case <-time.After(20 * time.Millisecond):
342 }
343 }
344 }
345
346 var _ = fmt.Sprintf
347
347 lines GO