test: add local multi-hop relay harness

Hee Sung Son committed May 13, 2026 at 14:11 UTC 3c828057905db29b46bc70f215f3f57c2fa5f2d4
11 files changed +625 -29
Makefile
+6 -1
@@ -1,4 +1,4 @@
1 -.PHONY: help install fmt vet lint lint-auto test vuln tidy all run build build-frontend build-docs build-tunnel build-server clean load-test
1 +.PHONY: help install fmt vet lint lint-auto test test-local-multihop vuln tidy all run build build-frontend build-docs build-tunnel build-server clean load-test
2
3 .DEFAULT_GOAL := help
4
@@ -15,6 +15,8 @@ help:
15 @echo " make install - Install Go developer tools used by this repo"
16 @echo " make fmt - Apply gofmt/goimports"
17 @echo " make lint-auto - Run autofix lint/format pipeline"
18 + @echo " make test - Run Go tests, including the local multi-hop harness"
19 + @echo " make test-local-multihop - Run the focused local multi-hop relay harness"
20 @echo " make build - Build everything (frontend, tunnel, server)"
21 @echo " make build-frontend - Build React frontend (Tailwind CSS 4)"
22 @echo " make build-docs - Build documentation site (SvelteKit)"
@@ -46,6 +48,9 @@ lint-auto:
48 test:
49 go test -v -coverprofile=coverage.out $(GO_PACKAGES)
50
51 +test-local-multihop:
52 + go test -v ./portal -run 'TestLocalCluster'
53 +
54 vuln:
55 govulncheck $(GO_PACKAGES)
56
docs/src/routes/architecture/+page.md
+8
@@ -300,6 +300,14 @@ Result: raw public UDP exposure with an internal QUIC datagram backhaul. UDP and
300 - The overlay peer API is plain HTTP on the WireGuard network, not public Internet HTTP. It serves the same discovery payload shape used by public `/discovery`.
301 - Overlay failure affects inter-relay discovery, mesh synchronization, and multi-hop relay forwarding. Direct tenant TLS routing, keyless TLS, register/renew/connect, and public UDP ingress do not depend on the WireGuard transport path.
302
303 +### Local multi-hop test harness
304 +
305 +`portal/multihop_local_test.go` verifies the multi-hop relay path inside one `go test` process. It starts local relay servers on `127.0.0.1:0`, seeds signed local discovery descriptors, and uses test fakes only for the WireGuard transport layer.
306 +
307 +Use `make test-local-multihop` for the focused local harness. The same tests are also included in `make test` through the `./portal/...` package set. Use `go test ./portal ./sdk` when checking the narrower portal and SDK interaction without the full repository test suite.
308 +
309 +The harness exercises SDK expose/register logic, `/sdk/hop`, `RelaySet`, signed descriptors, route registration, SNI ingress, hop token matching, and `bridgeLeaseConn`. It intentionally excludes public DNS, ACME provider side effects, public registry bootstrap, and real WireGuard devices. `localRelaySpec` provides relay-level server and descriptor mutation hooks for future policy tests, and the fake overlay records synced peers so fake hop streams only open after the `/sdk/hop` overlay sync path has run.
310 +
311 ## Control Plane Flow
312
313 ### 1. Register
portal/api_server.go
+3 -3
@@ -332,7 +332,7 @@ func (s *Server) handleRegisterChallenge(w http.ResponseWriter, r *http.Request)
332 Path: types.PathSDKRegister,
333 }).String()
334
335 - if strings.TrimSpace(req.HopToken) != "" && s.overlay == nil {
335 + if strings.TrimSpace(req.HopToken) != "" && !s.hasHopTransport() {
336 utils.WriteAPIError(w, http.StatusServiceUnavailable, types.APIErrorCodeFeatureUnavailable, errFeatureUnavailable.Error())
337 return
338 }
@@ -406,7 +406,7 @@ func (s *Server) handleHop(w http.ResponseWriter, r *http.Request) {
406 utils.MethodNotAllowedError().Write(w)
407 return
408 }
409 - if s.overlay == nil || s.relaySet == nil {
409 + if !s.hasHopTransport() || !s.hasOverlayRuntime() || s.relaySet == nil {
410 utils.WriteAPIError(w, http.StatusServiceUnavailable, types.APIErrorCodeFeatureUnavailable, errFeatureUnavailable.Error())
411 return
412 }
@@ -459,7 +459,7 @@ func (s *Server) handleHop(w http.ResponseWriter, r *http.Request) {
459 utils.InvalidRequestError(fmt.Errorf("forward relay: %w", err)).Write(w)
460 return
461 }
462 - if err := s.overlay.Sync(s.relaySet.OverlayPeerStates()); err != nil {
462 + if err := s.syncOverlayPeers(s.relaySet.OverlayPeerStates()); err != nil {
463 utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, err.Error())
464 return
465 }
portal/discovery/refresher.go
+20 -1
@@ -161,10 +161,26 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
161 }
162 continue
163 }
164 + client := r.httpClient
165 + var closeClient func()
166 + if utils.IsLocalRelayHost(baseURL.Hostname()) {
167 + _, localClient, transport, err := utils.NewHTTPTLSClient(ctx, baseURL, defaultRequestTimeout)
168 + if err != nil {
169 + if recoveryFailures > 0 {
170 + r.logDiscoveryFailure(relayURL, relayURL, recoveryFailures, err)
171 + }
172 + continue
173 + }
174 + client = localClient
175 + closeClient = transport.CloseIdleConnections
176 + }
177
178 startedAt := time.Now()
179 var resp types.DiscoveryResponse
167 - if err := utils.HTTPDoAPIPath(ctx, r.httpClient, baseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
180 + if err := utils.HTTPDoAPIPath(ctx, client, baseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
181 + if closeClient != nil {
182 + closeClient()
183 + }
184 if ctx.Err() != nil {
185 return ctx.Err()
186 }
@@ -173,6 +189,9 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
189 }
190 continue
191 }
192 + if closeClient != nil {
193 + closeClient()
194 + }
195 measuredAt := time.Now().UTC()
196
197 if _, err := r.relaySet.ApplyRelayDiscoveryResponse(relayURL, resp, measuredAt); err != nil {
portal/lease.go
+1
@@ -438,6 +438,7 @@ func (r *leaseRegistry) Renew(req types.RenewRequest, clientIP string) (types.Re
438 if strings.TrimSpace(reportedIP) != "" {
439 record.ReportedIP = reportedIP
440 }
441 + record.Metadata = req.Metadata.Copy()
442 r.policy.IPFilter().RegisterIdentityIP(leaseKey, clientIP)
443 identity := record.Identity
444 r.mu.Unlock()
portal/multihop_local_test.go new
+459
@@ -0,0 +1,459 @@
1 +package portal
2 +
3 +import (
4 + "bufio"
5 + "context"
6 + "crypto/tls"
7 + "errors"
8 + "fmt"
9 + "io"
10 + "net"
11 + "net/http"
12 + "slices"
13 + "sync"
14 + "testing"
15 + "time"
16 +
17 + "github.com/gosuda/portal-tunnel/v2/portal/acme"
18 + "github.com/gosuda/portal-tunnel/v2/portal/auth"
19 + "github.com/gosuda/portal-tunnel/v2/portal/discovery"
20 + "github.com/gosuda/portal-tunnel/v2/portal/overlay"
21 + "github.com/gosuda/portal-tunnel/v2/sdk"
22 + "github.com/gosuda/portal-tunnel/v2/types"
23 + "github.com/gosuda/portal-tunnel/v2/utils"
24 +)
25 +
26 +type localRelayCluster struct {
27 + relays []*localRelay
28 + byIP map[string]*localRelay
29 +}
30 +
31 +type localRelaySpec struct {
32 + Name string
33 + ServerMutator func(*Server)
34 + DescriptorMutator func(*types.RelayDescriptor)
35 +}
36 +
37 +type localRelay struct {
38 + name string
39 + server *Server
40 + apiURL string
41 + sniAddr string
42 + overlayIP string
43 + overlay *fakeOverlay
44 + descriptorMutator func(*types.RelayDescriptor)
45 +}
46 +
47 +type fakeOverlay struct {
48 + cfg overlay.Config
49 + mu sync.RWMutex
50 + syncedIPs map[string]struct{}
51 +}
52 +
53 +func (o *fakeOverlay) Config() overlay.Config {
54 + return o.cfg.Copy()
55 +}
56 +
57 +func (o *fakeOverlay) Sync(states []discovery.RelayState) error {
58 + synced := make(map[string]struct{}, len(states))
59 + for _, state := range states {
60 + overlayIP, err := utils.DeriveWireGuardOverlayIPv4(state.Descriptor.WireGuardPublicKey)
61 + if err != nil {
62 + return err
63 + }
64 + synced[overlayIP] = struct{}{}
65 + }
66 + o.mu.Lock()
67 + o.syncedIPs = synced
68 + o.mu.Unlock()
69 + return nil
70 +}
71 +
72 +func (o *fakeOverlay) hasSyncedPeer(overlayIPv4 string) bool {
73 + o.mu.RLock()
74 + defer o.mu.RUnlock()
75 + _, ok := o.syncedIPs[overlayIPv4]
76 + return ok
77 +}
78 +
79 +type fakeHopMux struct {
80 + cluster *localRelayCluster
81 + overlay *fakeOverlay
82 +}
83 +
84 +func (m *fakeHopMux) OpenStream(ctx context.Context, overlayIPv4, token string) (net.Conn, error) {
85 + if m == nil || m.cluster == nil {
86 + return nil, errors.New("fake hop mux is not connected to a cluster")
87 + }
88 + if m.overlay == nil || !m.overlay.hasSyncedPeer(overlayIPv4) {
89 + return nil, fmt.Errorf("fake hop target %q was not synced to overlay", overlayIPv4)
90 + }
91 + target := m.cluster.byIP[overlayIPv4]
92 + if target == nil {
93 + return nil, fmt.Errorf("fake hop target %q not found", overlayIPv4)
94 + }
95 + left, right := net.Pipe()
96 + go target.bridgeFakeHop(ctx, right, token)
97 + return left, nil
98 +}
99 +
100 +func (r *localRelay) bridgeFakeHop(ctx context.Context, conn net.Conn, token string) {
101 + r.server.registry.mu.RLock()
102 + record := r.server.registry.recordByHopToken(token, time.Now())
103 + r.server.registry.mu.RUnlock()
104 + if record == nil {
105 + _ = conn.Close()
106 + return
107 + }
108 + if err := r.server.bridgeLeaseConn(ctx, conn, record); err != nil {
109 + _ = conn.Close()
110 + }
111 +}
112 +
113 +func newLocalRelayCluster(t *testing.T, names ...string) *localRelayCluster {
114 + t.Helper()
115 + specs := make([]localRelaySpec, 0, len(names))
116 + for _, name := range names {
117 + specs = append(specs, localRelaySpec{Name: name})
118 + }
119 + return newLocalRelayClusterFromSpecs(t, specs...)
120 +}
121 +
122 +func newLocalRelayClusterFromSpecs(t *testing.T, specs ...localRelaySpec) *localRelayCluster {
123 + t.Helper()
124 + if len(specs) == 0 {
125 + t.Fatal("local relay cluster requires at least one relay")
126 + }
127 +
128 + ctx, cancel := context.WithCancel(context.Background())
129 + cluster := &localRelayCluster{
130 + relays: make([]*localRelay, 0, len(specs)),
131 + byIP: make(map[string]*localRelay, len(specs)),
132 + }
133 + t.Cleanup(func() {
134 + cancel()
135 + for _, relay := range cluster.relays {
136 + _ = relay.server.Shutdown(context.Background())
137 + if err := relay.server.Wait(); err != nil {
138 + t.Fatalf("relay %s Wait() error = %v", relay.name, err)
139 + }
140 + }
141 + })
142 +
143 + for _, spec := range specs {
144 + relay := startLocalRelay(t, ctx, spec)
145 + cluster.relays = append(cluster.relays, relay)
146 + }
147 + cluster.seedDiscovery(t)
148 + return cluster
149 +}
150 +
151 +func startLocalRelay(t *testing.T, ctx context.Context, spec localRelaySpec) *localRelay {
152 + t.Helper()
153 + name := spec.Name
154 + if name == "" {
155 + t.Fatal("local relay spec name is required")
156 + }
157 + server, err := NewServer(ServerConfig{
158 + PortalURL: "https://localhost:4017",
159 + IdentityPath: tempIdentityPath(t),
160 + ACME: acme.Config{KeyDir: t.TempDir()},
161 + APIListenAddr: "127.0.0.1:0",
162 + SNIListenAddr: "127.0.0.1:0",
163 + })
164 + if err != nil {
165 + t.Fatalf("NewServer(%s) error = %v", name, err)
166 + }
167 + if err := server.Start(ctx, nil); err != nil {
168 + t.Fatalf("Start(%s) error = %v", name, err)
169 + }
170 +
171 + apiPort := mustPort(t, server.apiListener.Addr().String())
172 + sniPort := mustPort(t, server.sniListener.Addr().String())
173 + apiURL := "https://localhost:" + apiPort
174 + server.cfg.PortalURL = apiURL
175 + server.cfg.SNIPort = mustAtoi(t, sniPort)
176 + server.registry.sniPort = server.cfg.SNIPort
177 + if spec.ServerMutator != nil {
178 + spec.ServerMutator(server)
179 + }
180 +
181 + wgPrivate, err := utils.GenerateWireGuardPrivateKey()
182 + if err != nil {
183 + t.Fatalf("GenerateWireGuardPrivateKey(%s) error = %v", name, err)
184 + }
185 + wgPublic, err := utils.WireGuardPublicKeyFromPrivate(wgPrivate)
186 + if err != nil {
187 + t.Fatalf("WireGuardPublicKeyFromPrivate(%s) error = %v", name, err)
188 + }
189 + fakeOverlay := &fakeOverlay{cfg: overlay.Config{
190 + PublicKey: wgPublic,
191 + ListenPort: overlay.DefaultListenPort,
192 + }}
193 + fakeHopMux := &fakeHopMux{overlay: fakeOverlay}
194 + server.testHooks = &serverTestHooks{
195 + overlayConfig: fakeOverlay.Config,
196 + syncOverlayPeers: fakeOverlay.Sync,
197 + openHopStream: fakeHopMux.OpenStream,
198 + }
199 + server.cfg.DiscoveryEnabled = true
200 +
201 + overlayIP, err := utils.DeriveWireGuardOverlayIPv4(wgPublic)
202 + if err != nil {
203 + t.Fatalf("DeriveWireGuardOverlayIPv4(%s) error = %v", name, err)
204 + }
205 + return &localRelay{
206 + name: name,
207 + server: server,
208 + apiURL: apiURL,
209 + sniAddr: net.JoinHostPort("127.0.0.1", sniPort),
210 + overlayIP: overlayIP,
211 + overlay: fakeOverlay,
212 + descriptorMutator: spec.DescriptorMutator,
213 + }
214 +}
215 +
216 +func (c *localRelayCluster) seedDiscovery(t *testing.T) {
217 + t.Helper()
218 + urls := make([]string, 0, len(c.relays))
219 + descriptors := make([]types.RelayDescriptor, 0, len(c.relays))
220 + for _, relay := range c.relays {
221 + urls = append(urls, relay.apiURL)
222 + desc, err := relay.server.newSelfDescriptor(time.Now())
223 + if err != nil {
224 + t.Fatalf("newSelfDescriptor(%s) error = %v", relay.name, err)
225 + }
226 + if relay.descriptorMutator != nil {
227 + relay.descriptorMutator(&desc)
228 + desc, err = auth.SignRelayDescriptor(desc, relay.server.identity.PrivateKey)
229 + if err != nil {
230 + t.Fatalf("SignRelayDescriptor(%s) error = %v", relay.name, err)
231 + }
232 + }
233 + descriptors = append(descriptors, desc)
234 + c.byIP[relay.overlayIP] = relay
235 + }
236 + for _, relay := range c.relays {
237 + relay.server.relaySet = discovery.NewRelaySet(urls)
238 + fakeHopMux := &fakeHopMux{cluster: c, overlay: relay.overlay}
239 + relay.server.testHooks.openHopStream = fakeHopMux.OpenStream
240 + changed, err := relay.server.relaySet.ApplyRelayDiscoveryResponse(relay.apiURL, types.DiscoveryResponse{
241 + ProtocolVersion: types.DiscoveryVersion,
242 + GeneratedAt: time.Now().UTC(),
243 + Relays: descriptors,
244 + }, time.Now())
245 + if err != nil {
246 + t.Fatalf("ApplyRelayDiscoveryResponse(%s) error = %v", relay.name, err)
247 + }
248 + if !changed {
249 + t.Fatalf("ApplyRelayDiscoveryResponse(%s) changed = false, want true", relay.name)
250 + }
251 + }
252 +}
253 +
254 +func (c *localRelayCluster) relay(index int) *localRelay {
255 + return c.relays[index]
256 +}
257 +
258 +func (c *localRelayCluster) relayURLs() []string {
259 + out := make([]string, 0, len(c.relays))
260 + for _, relay := range c.relays {
261 + out = append(out, relay.apiURL)
262 + }
263 + return out
264 +}
265 +
266 +func (c *localRelayCluster) exposeMultiHop(t *testing.T, ctx context.Context, name string) *sdk.Exposure {
267 + t.Helper()
268 + exposure, err := sdk.Expose(ctx, sdk.ExposeConfig{
269 + Name: name,
270 + TargetAddr: "127.0.0.1:1",
271 + MultiHop: c.relayURLs(),
272 + BanMITM: false,
273 + })
274 + if err != nil {
275 + t.Fatalf("sdk.Expose() error = %v", err)
276 + }
277 + t.Cleanup(func() {
278 + _ = exposure.Close()
279 + })
280 + return exposure
281 +}
282 +
283 +func TestLocalClusterExplicitMultiHopRegistersRoutes(t *testing.T) {
284 + cluster := newLocalRelayCluster(t, "entry", "middle", "exit")
285 + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Second)
286 + defer cancel()
287 +
288 + exposure := cluster.exposeMultiHop(t, ctx, "local-hop")
289 + waitForLocalHopRoutes(t, cluster, "local-hop.localhost")
290 +
291 + if got := exposure.ActiveRelayURLs(); len(got) != 1 || got[0] != cluster.relay(2).apiURL {
292 + t.Fatalf("ActiveRelayURLs() = %v, want exit relay %q", got, cluster.relay(2).apiURL)
293 + }
294 + entryRecord, ok := cluster.relay(0).server.registry.Lookup("local-hop.localhost")
295 + if !ok {
296 + t.Fatal("entry relay public hostname lookup failed")
297 + }
298 + if _, _, hasNext := entryRecord.nextHop(); !hasNext {
299 + t.Fatal("entry relay route has no next hop")
300 + }
301 + if middle := firstRecord(cluster.relay(1).server, func(record *leaseRecord) bool {
302 + return record.isHopMiddle()
303 + }); middle == nil {
304 + t.Fatal("middle relay hop route missing")
305 + }
306 + if exit := firstRecord(cluster.relay(2).server, func(record *leaseRecord) bool {
307 + return record.isHopExit() && record.stream != nil
308 + }); exit == nil {
309 + t.Fatal("exit relay stream lease missing")
310 + }
311 +}
312 +
313 +func TestLocalClusterPublicIngressTraversesFakeHopChain(t *testing.T) {
314 + cluster := newLocalRelayCluster(t, "entry", "middle", "exit")
315 + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Second)
316 + defer cancel()
317 +
318 + exposure := cluster.exposeMultiHop(t, ctx, "local-http")
319 + handlerReady := make(chan error, 1)
320 + go func() {
321 + handlerReady <- exposure.RunHTTP(ctx, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
322 + _, _ = io.WriteString(w, "local multi-hop ok")
323 + }), "")
324 + }()
325 + waitForLocalHopRoutes(t, cluster, "local-http.localhost")
326 +
327 + body := requestThroughEntry(t, cluster.relay(0), "local-http.localhost")
328 + if body != "local multi-hop ok" {
329 + t.Fatalf("public ingress body = %q, want %q", body, "local multi-hop ok")
330 + }
331 +
332 + cancel()
333 + if err := <-handlerReady; err != nil && !errors.Is(err, context.Canceled) {
334 + t.Fatalf("RunHTTP() error = %v", err)
335 + }
336 +}
337 +
338 +func TestLocalClusterDiscoveryStaysLocal(t *testing.T) {
339 + cluster := newLocalRelayCluster(t, "entry", "middle", "exit")
340 + want := cluster.relayURLs()
341 + slices.Sort(want)
342 +
343 + for _, relay := range cluster.relays {
344 + self, err := relay.server.newSelfDescriptor(time.Now())
345 + if err != nil {
346 + t.Fatalf("newSelfDescriptor(%s) error = %v", relay.name, err)
347 + }
348 + descriptors := relay.server.relaySet.Descriptors(self)
349 + got := make([]string, 0, len(descriptors))
350 + for _, desc := range descriptors {
351 + got = append(got, desc.APIHTTPSAddr)
352 + }
353 + slices.Sort(got)
354 + if !slices.Equal(got, want) {
355 + t.Fatalf("discovery relays for %s = %v, want local relays %v", relay.name, got, want)
356 + }
357 + }
358 +}
359 +
360 +func waitForLocalHopRoutes(t *testing.T, cluster *localRelayCluster, publicHostname string) {
361 + t.Helper()
362 + eventually(t, 10*time.Second, func() (bool, string) {
363 + if _, ok := cluster.relay(0).server.registry.Lookup(publicHostname); !ok {
364 + return false, "entry public route missing"
365 + }
366 + if firstRecord(cluster.relay(1).server, func(record *leaseRecord) bool { return record.isHopMiddle() }) == nil {
367 + return false, "middle hop route missing"
368 + }
369 + if firstRecord(cluster.relay(2).server, func(record *leaseRecord) bool {
370 + return record.isHopExit() && record.stream != nil && record.stream.ReadyCount() > 0
371 + }) == nil {
372 + return false, "exit stream lease not ready"
373 + }
374 + return true, ""
375 + })
376 +}
377 +
378 +func firstRecord(server *Server, match func(*leaseRecord) bool) *leaseRecord {
379 + server.registry.mu.RLock()
380 + defer server.registry.mu.RUnlock()
381 + for _, record := range server.registry.records {
382 + if record != nil && match(record) {
383 + return record
384 + }
385 + }
386 + return nil
387 +}
388 +
389 +func requestThroughEntry(t *testing.T, entry *localRelay, serverName string) string {
390 + t.Helper()
391 + conn, err := tls.DialWithDialer(&net.Dialer{Timeout: 5 * time.Second}, "tcp", entry.sniAddr, &tls.Config{
392 + MinVersion: tls.VersionTLS12,
393 + ServerName: serverName,
394 + InsecureSkipVerify: true,
395 + NextProtos: []string{"http/1.1"},
396 + })
397 + if err != nil {
398 + t.Fatalf("dial entry SNI listener: %v", err)
399 + }
400 + defer conn.Close()
401 +
402 + req, err := http.NewRequest(http.MethodGet, "https://"+serverName+"/", nil)
403 + if err != nil {
404 + t.Fatalf("NewRequest() error = %v", err)
405 + }
406 + if err := req.Write(conn); err != nil {
407 + t.Fatalf("write HTTP request: %v", err)
408 + }
409 + resp, err := http.ReadResponse(bufioNewReader(conn), req)
410 + if err != nil {
411 + t.Fatalf("read HTTP response: %v", err)
412 + }
413 + defer resp.Body.Close()
414 + body, err := io.ReadAll(resp.Body)
415 + if err != nil {
416 + t.Fatalf("read HTTP response body: %v", err)
417 + }
418 + if resp.StatusCode != http.StatusOK {
419 + t.Fatalf("GET through entry status = %d body=%q, want 200", resp.StatusCode, string(body))
420 + }
421 + return string(body)
422 +}
423 +
424 +func eventually(t *testing.T, timeout time.Duration, check func() (bool, string)) {
425 + t.Helper()
426 + deadline := time.Now().Add(timeout)
427 + var last string
428 + for time.Now().Before(deadline) {
429 + ok, reason := check()
430 + if ok {
431 + return
432 + }
433 + last = reason
434 + time.Sleep(25 * time.Millisecond)
435 + }
436 + t.Fatalf("condition not met within %s: %s", timeout, last)
437 +}
438 +
439 +func mustPort(t *testing.T, addr string) string {
440 + t.Helper()
441 + _, port, err := net.SplitHostPort(addr)
442 + if err != nil {
443 + t.Fatalf("SplitHostPort(%q) error = %v", addr, err)
444 + }
445 + return port
446 +}
447 +
448 +func mustAtoi(t *testing.T, raw string) int {
449 + t.Helper()
450 + var out int
451 + if _, err := fmt.Sscanf(raw, "%d", &out); err != nil {
452 + t.Fatalf("parse int %q: %v", raw, err)
453 + }
454 + return out
455 +}
456 +
457 +func bufioNewReader(conn net.Conn) *bufio.Reader {
458 + return bufio.NewReader(conn)
459 +}
portal/server.go
+73 -6
@@ -129,6 +129,13 @@ type Server struct {
129 relaySet *discovery.RelaySet
130 announceLimiter *discovery.AnnounceLimiter
131 registry *leaseRegistry
132 + testHooks *serverTestHooks
133 +}
134 +
135 +type serverTestHooks struct {
136 + overlayConfig func() overlay.Config
137 + syncOverlayPeers func([]discovery.RelayState) error
138 + openHopStream func(context.Context, string, string) (net.Conn, error)
139 }
140
141 func NewServer(cfg ServerConfig) (*Server, error) {
@@ -553,9 +560,21 @@ func (s *Server) bridgeLeaseConn(ctx context.Context, conn net.Conn, record *lea
560
561 openCtx, cancel := context.WithTimeout(ctx, defaultClaimTimeout)
562 defer cancel()
556 - next, err := s.overlay.OpenHopStream(openCtx, overlayIPv4, forwardToken)
557 - if err != nil {
558 - return fmt.Errorf("open next hop stream: %w", err)
563 + var next net.Conn
564 + var lastErr error
565 + for {
566 + var err error
567 + next, err = s.openHopStream(openCtx, overlayIPv4, forwardToken)
568 + if err == nil {
569 + break
570 + }
571 + lastErr = err
572 + if errors.Is(err, net.ErrClosed) {
573 + return fmt.Errorf("open next hop stream: %w", err)
574 + }
575 + if !utils.SleepOrDone(openCtx, defaultHopOpenRetryWait) {
576 + return fmt.Errorf("open next hop stream within %s: %w", defaultClaimTimeout, errors.Join(lastErr, openCtx.Err()))
577 + }
578 }
579 s.proxy.bridge(conn, next, "", nil)
580 return nil
@@ -576,6 +595,53 @@ func (s *Server) bridgeLeaseConn(ctx context.Context, conn net.Conn, record *lea
595 return nil
596 }
597
598 +func (s *Server) hasHopTransport() bool {
599 + return s != nil && (s.overlay != nil || (s.testHooks != nil && s.testHooks.openHopStream != nil))
600 +}
601 +
602 +func (s *Server) hasOverlayRuntime() bool {
603 + return s != nil && (s.overlay != nil || (s.testHooks != nil && s.testHooks.overlayConfig != nil && s.testHooks.syncOverlayPeers != nil))
604 +}
605 +
606 +func (s *Server) currentOverlayConfig() (overlay.Config, bool) {
607 + if s == nil {
608 + return overlay.Config{}, false
609 + }
610 + if s.overlay != nil {
611 + return s.overlay.Config(), true
612 + }
613 + if s.testHooks != nil && s.testHooks.overlayConfig != nil {
614 + return s.testHooks.overlayConfig(), true
615 + }
616 + return overlay.Config{}, false
617 +}
618 +
619 +func (s *Server) syncOverlayPeers(states []discovery.RelayState) error {
620 + if s == nil {
621 + return errors.New("server is unavailable")
622 + }
623 + if s.overlay != nil {
624 + return s.overlay.Sync(states)
625 + }
626 + if s.testHooks != nil && s.testHooks.syncOverlayPeers != nil {
627 + return s.testHooks.syncOverlayPeers(states)
628 + }
629 + return errors.New("relay overlay is unavailable")
630 +}
631 +
632 +func (s *Server) openHopStream(ctx context.Context, overlayIPv4, token string) (net.Conn, error) {
633 + if s == nil {
634 + return nil, errors.New("server is unavailable")
635 + }
636 + if s.overlay != nil {
637 + return s.overlay.OpenHopStream(ctx, overlayIPv4, token)
638 + }
639 + if s.testHooks != nil && s.testHooks.openHopStream != nil {
640 + return s.testHooks.openHopStream(ctx, overlayIPv4, token)
641 + }
642 + return nil, errors.New("relay overlay is unavailable")
643 +}
644 +
645 func (s *Server) runRegistryJanitor(ctx context.Context, interval time.Duration) error {
646 if interval <= 0 {
647 return errors.New("janitor interval must be positive")
@@ -751,10 +817,11 @@ func (s *Server) newSelfDescriptor(now time.Time) (types.RelayDescriptor, error)
817
818 var wireGuardPublicKey string
819 var wireGuardPort int
754 - if s.overlay != nil {
755 - cfg := s.overlay.Config()
820 + supportsOverlay := false
821 + if cfg, ok := s.currentOverlayConfig(); ok {
822 wireGuardPublicKey = cfg.PublicKey
823 wireGuardPort = cfg.ListenPort
824 + supportsOverlay = true
825 }
826
827 return auth.SignRelayDescriptor(types.RelayDescriptor{
@@ -765,7 +832,7 @@ func (s *Server) newSelfDescriptor(now time.Time) (types.RelayDescriptor, error)
832 APIHTTPSAddr: s.cfg.PortalURL,
833 WireGuardPublicKey: wireGuardPublicKey,
834 WireGuardPort: wireGuardPort,
768 - SupportsOverlay: s.overlay != nil,
835 + SupportsOverlay: supportsOverlay,
836 SupportsUDP: s.cfg.UDPEnabled && s.quicBackhaul != nil,
837 SupportsTCP: s.cfg.TCPEnabled,
838 ActiveConnections: s.proxy.activeConnectionCount(),
sdk/api_client.go
+11 -5
@@ -237,16 +237,22 @@ func (l *listener) registerLease(ctx context.Context, ttl time.Duration, udpEnab
237
238 func (l *listener) renewRegisteredLease(ctx context.Context, ttl time.Duration, accessToken string) (types.RenewResponse, error) {
239 var resp types.RenewResponse
240 - if err := utils.HTTPDoAPIPath(ctx, l.httpClient, l.relayURL, http.MethodPost, types.PathSDKRenew, types.RenewRequest{
241 - AccessToken: accessToken,
242 - TTL: int(ttl / time.Second),
243 - ReportedIP: utils.ResolvePublicIP(ctx),
244 - }, nil, &resp); err != nil {
240 + req := newRenewRequest(ttl, accessToken, utils.ResolvePublicIP(ctx), l.metadata)
241 + if err := utils.HTTPDoAPIPath(ctx, l.httpClient, l.relayURL, http.MethodPost, types.PathSDKRenew, req, nil, &resp); err != nil {
242 return types.RenewResponse{}, err
243 }
244 return resp, nil
245 }
246
247 +func newRenewRequest(ttl time.Duration, accessToken, reportedIP string, metadata types.LeaseMetadata) types.RenewRequest {
248 + return types.RenewRequest{
249 + AccessToken: accessToken,
250 + TTL: int(ttl / time.Second),
251 + ReportedIP: reportedIP,
252 + Metadata: metadata.Copy(),
253 + }
254 +}
255 +
256 func (l *listener) unregisterLease(ctx context.Context, accessToken string, hopRoutes []types.HopRoute) error {
257 hopErr := l.unregisterHopRoutes(ctx, hopRoutes)
258 err := utils.HTTPDoAPIPath(ctx, l.httpClient, l.relayURL, http.MethodPost, types.PathSDKUnregister, types.UnregisterRequest{
sdk/api_client_test.go new
+39
@@ -0,0 +1,39 @@
1 +package sdk
2 +
3 +import (
4 + "testing"
5 + "time"
6 +
7 + "github.com/gosuda/portal-tunnel/v2/types"
8 +)
9 +
10 +func TestNewRenewRequestIncludesMetadata(t *testing.T) {
11 + t.Parallel()
12 +
13 + metadata := types.LeaseMetadata{
14 + Description: "live app",
15 + Tags: []string{"demo", "live"},
16 + Owner: "ops",
17 + Thumbnail: "https://example.com/thumb.png",
18 + Hide: true,
19 + }
20 +
21 + req := newRenewRequest(2*time.Minute, "token", "203.0.113.10", metadata)
22 + if req.AccessToken != "token" {
23 + t.Fatalf("AccessToken = %q, want token", req.AccessToken)
24 + }
25 + if req.TTL != 120 {
26 + t.Fatalf("TTL = %d, want 120", req.TTL)
27 + }
28 + if req.ReportedIP != "203.0.113.10" {
29 + t.Fatalf("ReportedIP = %q, want 203.0.113.10", req.ReportedIP)
30 + }
31 + if got := req.Metadata; got.Description != metadata.Description || got.Owner != metadata.Owner || got.Thumbnail != metadata.Thumbnail || got.Hide != metadata.Hide || len(got.Tags) != 2 || got.Tags[0] != "demo" || got.Tags[1] != "live" {
32 + t.Fatalf("Metadata = %#v, want %#v", got, metadata)
33 + }
34 +
35 + metadata.Tags[0] = "mutated"
36 + if req.Metadata.Tags[0] != "demo" {
37 + t.Fatalf("Metadata tags alias input slice: got %q", req.Metadata.Tags[0])
38 + }
39 +}
sdk/expose.go
+1 -10
@@ -641,16 +641,7 @@ func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
641 multiHop = append([]string(nil), e.multiHop...)
642 explicitRelays := append([]string(nil), e.explicitRelays...)
643 if len(multiHop) > 0 {
644 - listenerRelayURLs = e.relaySet.PriorityRelays(discovery.ClientState{
645 - ExplicitRelayURLs: explicitRelays,
646 - MaxActiveRelays: e.maxActiveRelays,
647 - RequireUDP: e.udpEnabled,
648 - RequireTCP: e.tcpEnabled,
649 - LocalAddress: e.identity.Address,
650 - })
651 - if exitRelayURL := multiHop[len(multiHop)-1]; !slices.Contains(listenerRelayURLs, exitRelayURL) {
652 - listenerRelayURLs = append(listenerRelayURLs, exitRelayURL)
653 - }
644 + listenerRelayURLs = []string{multiHop[len(multiHop)-1]}
645 } else if e.multiHopDepth > 1 {
646 multiHop = e.relaySet.PriorityMultiHop(discovery.ClientState{
647 MultiHopDepth: e.multiHopDepth,
types/api.go
+4 -3
@@ -110,9 +110,10 @@ type DiscoveryAnnounceResponse struct {
110 }
111
112 type RenewRequest struct {
113 - AccessToken string `json:"access_token"`
114 - TTL int `json:"ttl,omitempty"`
115 - ReportedIP string `json:"reported_ip,omitempty"`
113 + AccessToken string `json:"access_token"`
114 + TTL int `json:"ttl,omitempty"`
115 + ReportedIP string `json:"reported_ip,omitempty"`
116 + Metadata LeaseMetadata `json:"metadata,omitempty"`
117 }
118
119 type RenewResponse struct {