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