feat: support multi relay in sdk, tidy demo app

rabbitprincess committed Mar 7, 2026 at 18:59 UTC afb17d8b1cfa1d2ce7609f2d3a7405ab539b4eda
15 files changed +902 -529
cmd/demo-app/handler.go new
+61
@@ -0,0 +1,61 @@
1 +package main
2 +
3 +import (
4 + "embed"
5 + "encoding/json"
6 + "io/fs"
7 + "net/http"
8 + "time"
9 +
10 + "golang.org/x/net/websocket"
11 +)
12 +
13 +//go:embed static
14 +var staticFiles embed.FS
15 +
16 +func newHandler() http.Handler {
17 + staticFS, _ := fs.Sub(staticFiles, "static")
18 +
19 + mux := http.NewServeMux()
20 + mux.Handle("/", http.FileServer(http.FS(staticFS)))
21 + mux.HandleFunc("/api/ping", handlePing)
22 + mux.Handle("/ws", websocket.Handler(handleWebSocket))
23 + mux.HandleFunc("/api/test-cookies", handleCookies)
24 + return mux
25 +}
26 +
27 +func handlePing(w http.ResponseWriter, _ *http.Request) {
28 + w.Header().Set("Content-Type", "application/json")
29 + _ = json.NewEncoder(w).Encode(map[string]any{
30 + "message": "pong",
31 + "time": time.Now().UTC().Format(time.RFC3339),
32 + })
33 +}
34 +
35 +func handleWebSocket(conn *websocket.Conn) {
36 + defer conn.Close()
37 + for {
38 + var msg string
39 + if err := websocket.Message.Receive(conn, &msg); err != nil {
40 + return
41 + }
42 + if err := websocket.Message.Send(conn, "echo: "+msg); err != nil {
43 + return
44 + }
45 + }
46 +}
47 +
48 +func handleCookies(w http.ResponseWriter, _ *http.Request) {
49 + for _, cookie := range []*http.Cookie{
50 + {Name: "session_id", Value: "abc123", Path: "/", MaxAge: 3600},
51 + {Name: "auth_token", Value: "secret456", Path: "/", MaxAge: 3600},
52 + {Name: "csrf_token", Value: "xyz789", Path: "/", MaxAge: 3600},
53 + {Name: "user_pref", Value: "dark_mode", Path: "/", MaxAge: 86400},
54 + } {
55 + http.SetCookie(w, cookie)
56 + }
57 + w.Header().Set("Content-Type", "application/json")
58 + _ = json.NewEncoder(w).Encode(map[string]any{
59 + "message": "4 cookies set: session_id, auth_token, csrf_token, user_pref",
60 + })
61 +}
cmd/demo-app/main.go
+18 -114
@@ -2,34 +2,22 @@ package main
2
3 import (
4 "context"
5 - "embed"
5 + _ "embed"
6 "encoding/base64"
7 - "encoding/json"
8 - "errors"
7 "flag"
8 "fmt"
11 - "io/fs"
12 - "net/http"
9 "os"
10 "os/signal"
15 - "strings"
11 "syscall"
12 "time"
13
14 "github.com/rs/zerolog"
15 "github.com/rs/zerolog/log"
21 - "golang.org/x/net/websocket"
16
17 "github.com/gosuda/portal/v2/sdk"
18 "github.com/gosuda/portal/v2/types"
19 )
20
27 -//go:embed static
28 -var staticFiles embed.FS
29 -
30 -//go:embed static/thumbnail.png
31 -var thumbnailPNG []byte
32 -
21 var (
22 flagServerURL string
23 flagPort int
@@ -38,19 +26,24 @@ var (
26 flagTags string
27 flagOwner string
28 flagHide bool
29 +
30 + //go:embed static/thumbnail.png
31 + thumbnailPNG []byte
32 + flagThumbnail = "data:image/png;base64," + base64.StdEncoding.EncodeToString(thumbnailPNG)
33 )
34
35 func main() {
36 log.Logger = log.Output(zerolog.ConsoleWriter{Out: os.Stdout, TimeFormat: time.RFC3339})
37 logger := log.With().Str("component", "demo-app").Logger()
38
47 - flag.StringVar(&flagServerURL, "server-url", "https://localhost:4017", "relay API URL (https only)")
39 + flag.StringVar(&flagServerURL, "server-url", "https://localhost:4017", "relay API URLs (comma-separated, https only)")
40 flag.IntVar(&flagPort, "port", 8092, "local demo HTTP port")
41 flag.StringVar(&flagName, "name", "demo-app", "backend display name")
42 flag.StringVar(&flagDesc, "description", "Portal demo connectivity app", "lease description")
43 flag.StringVar(&flagTags, "tags", "demo,connectivity,activity,cloud,sun,morning", "comma-separated lease tags")
44 flag.StringVar(&flagOwner, "owner", "PortalApp Developer", "lease owner")
45 flag.BoolVar(&flagHide, "hide", false, "hide this lease from listings")
46 +
47 flag.Parse()
48
49 if err := runDemo(); err != nil {
@@ -62,24 +55,22 @@ func main() {
55 func runDemo() error {
56 logger := log.With().Str("component", "demo-app").Logger()
57
65 - sdkClient, err := sdk.NewClient(sdk.ClientConfig{RelayURL: flagServerURL})
58 + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
59 + defer stop()
60 +
61 + sdkClient, err := sdk.NewClient(sdk.ClientConfig{RelayURLs: sdk.SplitCSV(flagServerURL)})
62 if err != nil {
63 return fmt.Errorf("new client: %w", err)
64 }
65 defer sdkClient.Close()
66
71 - thumbnailDataURI := "data:image/png;base64," + base64.StdEncoding.EncodeToString(thumbnailPNG)
72 -
73 - ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
74 - defer stop()
75 -
67 listener, err := sdkClient.Listen(ctx, sdk.ListenRequest{
68 Name: flagName,
69 Metadata: types.LeaseMetadata{
70 Description: flagDesc,
80 - Tags: splitCSV(flagTags),
71 + Tags: sdk.SplitCSV(flagTags),
72 Owner: flagOwner,
82 - Thumbnail: thumbnailDataURI,
73 + Thumbnail: flagThumbnail,
74 Hide: flagHide,
75 },
76 })
@@ -89,106 +80,19 @@ func runDemo() error {
80 defer listener.Close()
81
82 logger.Info().
92 - Str("lease_id", listener.LeaseID()).
83 Strs("public_urls", listener.PublicURLs()).
84 Int("local_port", flagPort).
85 Msg("demo app registered with relay")
86
97 - mux := http.NewServeMux()
98 -
99 - staticFS, err := fs.Sub(staticFiles, "static")
100 - if err != nil {
101 - return fmt.Errorf("create static fs: %w", err)
87 + if err := sdk.RunHTTPApp(ctx, listener, newHandler(), sdk.HTTPServeOptions{
88 + LocalAddr: fmt.Sprintf(":%d", flagPort),
89 + }); err != nil {
90 + return err
91 }
103 - mux.Handle("/", http.FileServer(http.FS(staticFS)))
104 -
105 - mux.HandleFunc("/api/ping", func(w http.ResponseWriter, _ *http.Request) {
106 - w.Header().Set("Content-Type", "application/json")
107 - resp := map[string]any{
108 - "message": "pong",
109 - "time": time.Now().UTC().Format(time.RFC3339),
110 - }
111 - _ = json.NewEncoder(w).Encode(resp)
112 - })
113 -
114 - mux.Handle("/ws", websocket.Handler(func(conn *websocket.Conn) {
115 - defer conn.Close()
116 - for {
117 - var msg string
118 - if err := websocket.Message.Receive(conn, &msg); err != nil {
119 - break
120 - }
121 - if err := websocket.Message.Send(conn, "echo: "+msg); err != nil {
122 - break
123 - }
124 - }
125 - }))
126 -
127 - mux.HandleFunc("/api/test-cookies", func(w http.ResponseWriter, _ *http.Request) {
128 - for _, cookie := range []*http.Cookie{
129 - {Name: "session_id", Value: "abc123", Path: "/", MaxAge: 3600},
130 - {Name: "auth_token", Value: "secret456", Path: "/", MaxAge: 3600},
131 - {Name: "csrf_token", Value: "xyz789", Path: "/", MaxAge: 3600},
132 - {Name: "user_pref", Value: "dark_mode", Path: "/", MaxAge: 86400},
133 - } {
134 - http.SetCookie(w, cookie)
135 - }
136 - w.Header().Set("Content-Type", "application/json")
137 - _ = json.NewEncoder(w).Encode(map[string]any{
138 - "message": "4 cookies set: session_id, auth_token, csrf_token, user_pref",
139 - })
140 - })
141 -
142 - localAddr := fmt.Sprintf(":%d", flagPort)
143 - go func() {
144 - localSrv := &http.Server{
145 - Addr: localAddr,
146 - Handler: mux,
147 - ReadHeaderTimeout: 5 * time.Second,
148 - }
149 - logger.Info().Str("addr", localAddr).Msg("demo app listening locally")
150 - if err := localSrv.ListenAndServe(); err != nil && err != http.ErrServerClosed {
151 - logger.Error().Err(err).Str("addr", localAddr).Msg("demo local server stopped")
152 - }
153 - }()
154 -
155 - relaySrv := &http.Server{
156 - Handler: mux,
157 - ReadHeaderTimeout: 5 * time.Second,
158 - }
159 -
160 - sig := make(chan os.Signal, 1)
161 - signal.Notify(sig, syscall.SIGINT, syscall.SIGTERM)
92
163 - errCh := make(chan error, 1)
164 - go func() {
165 - errCh <- relaySrv.Serve(listener)
166 - }()
167 -
168 - select {
169 - case <-sig:
93 + if ctx.Err() != nil {
94 logger.Info().Msg("demo app shutting down")
171 - case err := <-errCh:
172 - if err != nil && !errors.Is(err, http.ErrServerClosed) {
173 - return err
174 - }
95 }
176 -
96 logger.Info().Msg("demo app shutdown complete")
97 return nil
98 }
180 -
181 -func splitCSV(raw string) []string {
182 - if strings.TrimSpace(raw) == "" {
183 - return nil
184 - }
185 - parts := strings.Split(raw, ",")
186 - out := make([]string, 0, len(parts))
187 - for _, part := range parts {
188 - part = strings.TrimSpace(part)
189 - if part != "" {
190 - out = append(out, part)
191 - }
192 - }
193 - return out
194 -}
cmd/portal-tunnel/README.md
+1
@@ -30,6 +30,7 @@ Portal-tunnel connects a local service to a Portal relay with the legacy CLI sha
30 ## Notes
31
32 - Multiple relay URLs are registered independently. Each relay gets its own lease ID and public URLs.
33 +- Portal-tunnel now consumes one aggregate SDK listener, so the CLI no longer manages per-relay listener loops itself.
34 - Startup is fail-fast: if any configured relay cannot register, the tunnel exits instead of partially publishing.
35 - Tenant TLS is provisioned automatically through the relay keyless signer. The SDK fetches the relay certificate chain and uses `/v1/sign` for remote signing.
36 - When the local service is unreachable, the tunnel returns an HTTP 503 page.
cmd/portal-tunnel/main.go
+47 -30
@@ -7,13 +7,13 @@ import (
7 "fmt"
8 "os"
9 "os/signal"
10 - "sync"
10 "sync/atomic"
11 "syscall"
12 "time"
13
14 "github.com/rs/zerolog"
15 "github.com/rs/zerolog/log"
16 + "golang.org/x/sync/errgroup"
17
18 "github.com/gosuda/portal/v2/sdk"
19 "github.com/gosuda/portal/v2/types"
@@ -62,16 +62,18 @@ func runTunnel() error {
62 ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
63 defer stop()
64
65 - relayURLs, err := normalizeRelayURLs(flagRelayURLs)
66 - if err != nil {
67 - return err
68 - }
65 + relayURLs := sdk.SplitCSV(flagRelayURLs)
66 if len(relayURLs) == 0 {
67 return errors.New("no relay URLs provided")
68 }
69
70 + targetAddr, err := normalizeTargetAddr(flagHost)
71 + if err != nil {
72 + return fmt.Errorf("invalid --host value %q: %w", flagHost, err)
73 + }
74 +
75 logger.Info().
74 - Str("local", flagHost).
76 + Str("local", targetAddr).
77 Int("relay_count", len(relayURLs)).
78 Strs("relays", relayURLs).
79 Msg("starting portal tunnel")
@@ -80,59 +82,74 @@ func runTunnel() error {
82 Name: flagName,
83 Metadata: types.LeaseMetadata{
84 Description: flagDesc,
83 - Tags: splitCSV(flagTags),
85 + Tags: sdk.SplitCSV(flagTags),
86 Owner: flagOwner,
87 Thumbnail: flagThumbnail,
88 Hide: flagHide,
89 },
90 }
91
90 - runtimes, err := startRelayRuntimes(ctx, relayURLs, listenReq)
92 + client, err := sdk.NewClient(sdk.ClientConfig{RelayURLs: relayURLs})
93 if err != nil {
92 - return fmt.Errorf("service %s: failed to start relays: %w", flagName, err)
94 + return fmt.Errorf("service %s: failed to create client: %w", flagName, err)
95 }
96 + defer client.Close()
97
95 - var connWG sync.WaitGroup
96 - var connCount atomic.Int64
97 - relayDone := make(chan relayLoopResult, len(runtimes))
98 + listener, err := client.Listen(ctx, listenReq)
99 + if err != nil {
100 + return fmt.Errorf("service %s: failed to start listener: %w", flagName, err)
101 + }
102 + defer listener.Close()
103
99 - for _, runtime := range runtimes {
104 + for _, entry := range listener.Entries() {
105 logger.Info().
101 - Str("relay", runtime.relayURL).
102 - Str("lease_id", runtime.listener.LeaseID()).
103 - Strs("public_urls", runtime.listener.PublicURLs()).
106 + Str("relay", entry.RelayURL).
107 + Str("lease_id", entry.LeaseID).
108 + Strs("public_urls", entry.PublicURLs()).
109 Msg("relay tunnel ready")
105 - go runtime.run(ctx, flagHost, &connWG, &connCount, relayDone)
110 }
111
108 - waitErr := waitForRelayLoops(ctx, relayDone, len(runtimes))
109 - if waitErr != nil {
110 - stop()
111 - }
112 - closeErr := closeRelayRuntimes(runtimes)
112 + var connGroup errgroup.Group
113 + var connCount atomic.Int64
114 + group, groupCtx := errgroup.WithContext(ctx)
115 + group.Go(func() error {
116 + if err := runProxyLoop(groupCtx, listener, targetAddr, &connGroup, &connCount); err != nil {
117 + return fmt.Errorf("relay accept loop: %w", err)
118 + }
119 + return nil
120 + })
121 + group.Go(func() error {
122 + <-groupCtx.Done()
123 + if err := listener.Close(); err != nil {
124 + return fmt.Errorf("listener close: %w", err)
125 + }
126 + return nil
127 + })
128 +
129 + waitErr := group.Wait()
130 if waitErr != nil {
131 logger.Error().Err(waitErr).Msg("relay supervisor exited with error")
132 }
116 - if closeErr != nil {
117 - logger.Error().Err(closeErr).Msg("relay shutdown failed")
118 - }
133
134 if ctx.Err() != nil {
135 logger.Info().Msg("tunnel shutting down")
136 }
137
124 - done := make(chan struct{})
138 + done := make(chan error, 1)
139 go func() {
126 - connWG.Wait()
127 - close(done)
140 + done <- connGroup.Wait()
141 }()
142
143 select {
131 - case <-done:
144 + case err := <-done:
145 + if err != nil {
146 + logger.Error().Err(err).Msg("proxy connection group failed")
147 + waitErr = errors.Join(waitErr, err)
148 + }
149 case <-time.After(5 * time.Second):
150 logger.Warn().Msg("tunnel shutdown timeout; connections still active")
151 }
152
153 logger.Info().Msg("tunnel shutdown complete")
137 - return errors.Join(waitErr, closeErr)
154 + return waitErr
155 }
cmd/portal-tunnel/relays.go
+40 -180
@@ -13,194 +13,51 @@ import (
13 "time"
14
15 "github.com/rs/zerolog/log"
16 + "golang.org/x/sync/errgroup"
17
18 "github.com/gosuda/portal/v2/sdk"
19 )
20
20 -type relayRuntime struct {
21 - client *sdk.Client
22 - listener *sdk.Listener
23 - relayURL string
24 -}
25 -
26 -type relayLoopResult struct {
27 - err error
28 - leaseID string
29 - relayURL string
21 +var bufferPool = sync.Pool{
22 + New: func() any {
23 + b := make([]byte, 64*1024)
24 + return &b
25 + },
26 }
27
32 -func (r *relayRuntime) run(ctx context.Context, localAddr string, connWG *sync.WaitGroup, connCount *atomic.Int64, done chan<- relayLoopResult) {
33 - logger := log.With().
34 - Str("component", "portal-tunnel").
35 - Str("relay", r.relayURL).
36 - Str("lease_id", r.listener.LeaseID()).
37 - Logger()
38 -
39 - var runErr error
40 - defer func() {
41 - done <- relayLoopResult{
42 - leaseID: r.listener.LeaseID(),
43 - relayURL: r.relayURL,
44 - err: runErr,
45 - }
46 - }()
28 +func runProxyLoop(ctx context.Context, listener *sdk.Listener, targetAddr string, connGroup *errgroup.Group, connCount *atomic.Int64) error {
29 + logger := log.With().Str("component", "portal-tunnel").Logger()
30
31 for {
49 - relayConn, err := r.listener.Accept()
32 + relayConn, entry, err := listener.AcceptEntry()
33 if err != nil {
34 if errors.Is(err, net.ErrClosed) || errors.Is(err, context.Canceled) || ctx.Err() != nil {
52 - return
35 + return nil
36 }
54 - runErr = err
55 - return
37 + return err
38 }
39
40 connID := connCount.Add(1)
41 logger.Info().
42 Int64("conn_id", connID).
43 Str("remote_addr", relayConn.RemoteAddr().String()).
44 + Str("relay", entry.RelayURL).
45 + Str("lease_id", entry.LeaseID).
46 Msg("accepted relay connection")
47
64 - connWG.Add(1)
65 - go func(connID int64, relayConn net.Conn) {
66 - defer connWG.Done()
67 - if err := proxyConnection(ctx, localAddr, relayConn); err != nil {
48 + connGroup.Go(func() error {
49 + if err := proxyConnection(ctx, targetAddr, relayConn); err != nil {
50 logger.Error().Err(err).Int64("conn_id", connID).Msg("proxy connection failed")
51 }
52 logger.Info().Int64("conn_id", connID).Msg("proxy connection closed")
71 - }(connID, relayConn)
72 - }
73 -}
74 -
75 -func startRelayRuntimes(ctx context.Context, relayURLs []string, req sdk.ListenRequest) ([]*relayRuntime, error) {
76 - runtimes := make([]*relayRuntime, 0, len(relayURLs))
77 - for _, relayURL := range relayURLs {
78 - client, err := sdk.NewClient(sdk.ClientConfig{RelayURL: relayURL})
79 - if err != nil {
80 - _ = closeRelayRuntimes(runtimes)
81 - return nil, fmt.Errorf("create relay client %s: %w", relayURL, err)
82 - }
83 -
84 - listener, err := client.Listen(ctx, req)
85 - if err != nil {
86 - client.Close()
87 - _ = closeRelayRuntimes(runtimes)
88 - return nil, fmt.Errorf("register relay lease %s: %w", relayURL, err)
89 - }
90 -
91 - runtimes = append(runtimes, &relayRuntime{
92 - relayURL: relayURL,
93 - client: client,
94 - listener: listener,
53 + return nil
54 })
55 }
97 - return runtimes, nil
98 -}
99 -
100 -func closeRelayRuntimes(runtimes []*relayRuntime) error {
101 - var closeErr error
102 - for _, runtime := range runtimes {
103 - if runtime == nil {
104 - continue
105 - }
106 - if runtime.listener != nil {
107 - closeErr = errors.Join(closeErr, runtime.listener.Close())
108 - }
109 - if runtime.client != nil {
110 - runtime.client.Close()
111 - }
112 - }
113 - return closeErr
114 -}
115 -
116 -func waitForRelayLoops(ctx context.Context, done <-chan relayLoopResult, relayCount int) error {
117 - logger := log.With().Str("component", "portal-tunnel").Logger()
118 - active := relayCount
119 -
120 - for active > 0 {
121 - result := <-done
122 - active--
123 -
124 - switch {
125 - case result.err != nil:
126 - logger.Error().
127 - Err(result.err).
128 - Str("relay", result.relayURL).
129 - Str("lease_id", result.leaseID).
130 - Int("remaining_relays", active).
131 - Msg("relay accept loop stopped")
132 - case ctx.Err() != nil:
133 - logger.Info().
134 - Str("relay", result.relayURL).
135 - Str("lease_id", result.leaseID).
136 - Int("remaining_relays", active).
137 - Msg("relay accept loop stopped during shutdown")
138 - default:
139 - logger.Warn().
140 - Str("relay", result.relayURL).
141 - Str("lease_id", result.leaseID).
142 - Int("remaining_relays", active).
143 - Msg("relay accept loop stopped")
144 - }
145 - }
146 -
147 - select {
148 - case <-ctx.Done():
149 - return nil
150 - default:
151 - }
152 - return errors.New("all relay listeners stopped")
56 }
57
155 -func normalizeRelayURLs(raw string) ([]string, error) {
156 - seen := make(map[string]struct{})
157 - var relayURLs []string
158 - for _, relayURL := range splitCSV(raw) {
159 - normalized, err := normalizeRelayURL(relayURL)
160 - if err != nil {
161 - return nil, err
162 - }
163 - if _, ok := seen[normalized]; ok {
164 - continue
165 - }
166 - seen[normalized] = struct{}{}
167 - relayURLs = append(relayURLs, normalized)
168 - }
169 - return relayURLs, nil
170 -}
171 -
172 -func normalizeRelayURL(raw string) (string, error) {
173 - u, err := url.Parse(strings.TrimSpace(raw))
174 - if err != nil {
175 - return "", fmt.Errorf("parse relay url: %w", err)
176 - }
177 - if !strings.EqualFold(u.Scheme, "https") {
178 - return "", fmt.Errorf("relay url must use https: %q", raw)
179 - }
180 - if u.Host == "" {
181 - return "", fmt.Errorf("relay url host is empty: %q", raw)
182 - }
183 - u.Path = strings.TrimRight(u.Path, "/")
184 - u.RawQuery = ""
185 - u.Fragment = ""
186 - return u.String(), nil
187 -}
188 -
189 -var bufferPool = sync.Pool{
190 - New: func() any {
191 - b := make([]byte, 64*1024)
192 - return &b
193 - },
194 -}
195 -
196 -func proxyConnection(ctx context.Context, localAddr string, relayConn net.Conn) error {
58 +func proxyConnection(ctx context.Context, targetAddr string, relayConn net.Conn) error {
59 defer relayConn.Close()
60
199 - targetAddr, err := normalizeTargetAddr(localAddr)
200 - if err != nil {
201 - return fmt.Errorf("invalid --host value %q: %w", localAddr, err)
202 - }
203 -
61 dialer := &net.Dialer{Timeout: 5 * time.Second}
62 localConn, err := dialer.DialContext(ctx, "tcp", targetAddr)
63 if err != nil {
@@ -271,40 +128,43 @@ func writeEmptyHTTPResponse(conn net.Conn) error {
128 return err
129 }
130
274 -func splitCSV(raw string) []string {
275 - if strings.TrimSpace(raw) == "" {
276 - return nil
277 - }
278 - parts := strings.Split(raw, ",")
279 - out := make([]string, 0, len(parts))
280 - for _, part := range parts {
281 - part = strings.TrimSpace(part)
282 - if part != "" {
283 - out = append(out, part)
284 - }
285 - }
286 - return out
287 -}
288 -
131 func normalizeTargetAddr(raw string) (string, error) {
132 raw = strings.TrimSpace(raw)
133 if raw == "" {
134 return "", errors.New("target address is required")
135 }
136 +
137 if strings.Contains(raw, "://") {
295 - if strings.HasPrefix(strings.ToLower(raw), "http://") {
296 - raw = strings.TrimPrefix(raw, "http://")
138 + targetURL, err := url.Parse(raw)
139 + if err != nil {
140 + return "", fmt.Errorf("parse target url: %w", err)
141 + }
142 + if !strings.EqualFold(targetURL.Scheme, "http") && !strings.EqualFold(targetURL.Scheme, "https") {
143 + return "", fmt.Errorf("unsupported target url scheme %q", targetURL.Scheme)
144 + }
145 + if targetURL.Host == "" {
146 + return "", errors.New("target url host is empty")
147 + }
148 + if targetURL.Path != "" && targetURL.Path != "/" {
149 + return "", errors.New("target url path is not supported")
150 }
298 - if strings.HasPrefix(strings.ToLower(raw), "https://") {
299 - raw = strings.TrimPrefix(raw, "https://")
151 + if targetURL.RawQuery != "" {
152 + return "", errors.New("target url query is not supported")
153 }
301 - raw = strings.TrimSuffix(raw, "/")
154 + if targetURL.Fragment != "" {
155 + return "", errors.New("target url fragment is not supported")
156 + }
157 + raw = targetURL.Host
158 }
159 +
160 if _, _, err := net.SplitHostPort(raw); err == nil {
161 return raw, nil
162 }
163 if strings.Count(raw, ":") == 0 {
164 return net.JoinHostPort(raw, "80"), nil
165 }
166 + if ip := net.ParseIP(raw); ip != nil {
167 + return net.JoinHostPort(raw, "80"), nil
168 + }
169 return "", fmt.Errorf("invalid target address %q", raw)
170 }
cmd/portal-tunnel/relays_test.go new
+44
@@ -0,0 +1,44 @@
1 +package main
2 +
3 +import "testing"
4 +
5 +func TestNormalizeTargetAddr(t *testing.T) {
6 + t.Parallel()
7 +
8 + tests := []struct {
9 + name string
10 + raw string
11 + want string
12 + wantErr bool
13 + }{
14 + {name: "host and port", raw: "localhost:8080", want: "localhost:8080"},
15 + {name: "host only", raw: "localhost", want: "localhost:80"},
16 + {name: "http url", raw: "http://localhost:8080", want: "localhost:8080"},
17 + {name: "https url", raw: "https://127.0.0.1", want: "127.0.0.1:80"},
18 + {name: "ipv6 host", raw: "::1", want: "[::1]:80"},
19 + {name: "url with path", raw: "http://localhost:8080/app", wantErr: true},
20 + {name: "url with query", raw: "http://localhost:8080/?x=1", wantErr: true},
21 + {name: "empty", raw: " ", wantErr: true},
22 + }
23 +
24 + for _, tt := range tests {
25 + tt := tt
26 + t.Run(tt.name, func(t *testing.T) {
27 + t.Parallel()
28 +
29 + got, err := normalizeTargetAddr(tt.raw)
30 + if tt.wantErr {
31 + if err == nil {
32 + t.Fatalf("normalizeTargetAddr(%q) error = nil, want error", tt.raw)
33 + }
34 + return
35 + }
36 + if err != nil {
37 + t.Fatalf("normalizeTargetAddr(%q) error = %v", tt.raw, err)
38 + }
39 + if got != tt.want {
40 + t.Fatalf("normalizeTargetAddr(%q) = %q, want %q", tt.raw, got, tt.want)
41 + }
42 + })
43 + }
44 +}
docs/architecture.md
+9 -6
@@ -54,13 +54,16 @@ That distinction matters because `/sdk/connect` stops being ordinary HTTP once h
54
55 ### SDK (`sdk/`)
56
57 -- `Client`: validates relay URL, owns HTTP client and raw TLS dial config
58 -- `Listener`: registers a lease, maintains `readyTarget` reverse sessions, renews lease TTL, and yields accepted tenant TLS connections
57 +- `Client`: validates one or more relay URLs and owns per-relay HTTP client and raw TLS dial config
58 +- `Listener`: registers one lease per relay, maintains per-entry `readyTarget` reverse sessions, renews lease TTLs, and yields accepted tenant TLS connections through one aggregate listener surface
59 +- Default app flow is `RelayURLs -> NewClient -> Listen -> PublicURLs -> http.Server.Serve(listener)`
60 +- `helper.go`: optional `RunHTTPApp` helper for serving one handler on both a local HTTP port and the relay listener
61 +- Relay-aware entry inspection is reserved for advanced callers such as `portal-tunnel`
62 - Tenant TLS is created automatically through the relay keyless signer; callers do not provide a local self-signed fallback path
63
64 ### Tunnel (`cmd/portal-tunnel`)
65
63 -- Registers a lease through the SDK
66 +- Creates one SDK client, registers one lease per relay through the SDK, and consumes one aggregate listener
67 - Accepts claimed tenant connections from the relay
68 - Proxies raw TCP to a local `--host`
69 - Returns an HTTP 503 response when the local target is unavailable
@@ -69,9 +72,9 @@ That distinction matters because `/sdk/connect` stops being ordinary HTTP once h
72
73 ### Raw reverse transport (`TLS=true` only)
74
72 -1. SDK/tunnel registers a lease with `POST /sdk/register`.
73 -2. SDK opens one or more reverse sessions with `GET /sdk/connect?lease_id=...`.
74 -3. Relay hijacks each `/sdk/connect` request and places the connection in the per-lease broker ready queue.
75 +1. SDK/tunnel registers one lease per relay with `POST /sdk/register`.
76 +2. SDK opens one or more reverse sessions per registered lease with `GET /sdk/connect?lease_id=...`.
77 +3. Each relay hijacks `/sdk/connect` requests and places the connection in the per-lease broker ready queue.
78 4. While idle, the relay writes `0x00` keepalive markers.
79 5. A browser connects to the relay SNI listener.
80 6. Relay extracts SNI from ClientHello, resolves a lease, and claims one ready reverse session.
portal/helpers.go
-51
@@ -3,13 +3,10 @@ package portal
3 import (
4 "crypto/rand"
5 "encoding/hex"
6 - "fmt"
6 "net"
7 "net/url"
8 "strings"
9 "time"
11 -
12 - "github.com/gosuda/portal/v2/types"
10 )
11
12 const (
@@ -30,23 +27,6 @@ func PortalRootHost(portalURL string) string {
27 return normalizeHostname(u.Hostname())
28 }
29
33 -func NormalizeRelayURL(raw string) (string, error) {
34 - u, err := url.Parse(strings.TrimSpace(raw))
35 - if err != nil {
36 - return "", fmt.Errorf("parse relay url: %w", err)
37 - }
38 - if !strings.EqualFold(u.Scheme, "https") {
39 - return "", fmt.Errorf("relay url must use https: %q", raw)
40 - }
41 - if u.Host == "" {
42 - return "", fmt.Errorf("relay url host is empty: %q", raw)
43 - }
44 - u.Path = strings.TrimRight(u.Path, "/")
45 - u.RawQuery = ""
46 - u.Fragment = ""
47 - return u.String(), nil
48 -}
49 -
30 func normalizeHostname(host string) string {
31 host = strings.TrimSpace(strings.ToLower(host))
32 host = strings.TrimSuffix(host, ".")
@@ -110,37 +90,6 @@ func intOrDefault(v, fallback int) int {
90 return fallback
91 }
92
113 -func normalizeMetadata(meta types.LeaseMetadata) types.LeaseMetadata {
114 - meta.Description = strings.TrimSpace(meta.Description)
115 - meta.Owner = strings.TrimSpace(meta.Owner)
116 - meta.Thumbnail = strings.TrimSpace(meta.Thumbnail)
117 - meta.Tags = normalizeTags(meta.Tags)
118 - return meta
119 -}
120 -
121 -func normalizeTags(tags []string) []string {
122 - if len(tags) == 0 {
123 - return nil
124 - }
125 - seen := make(map[string]struct{}, len(tags))
126 - out := make([]string, 0, len(tags))
127 - for _, tag := range tags {
128 - tag = strings.TrimSpace(tag)
129 - if tag == "" {
130 - continue
131 - }
132 - if _, ok := seen[tag]; ok {
133 - continue
134 - }
135 - seen[tag] = struct{}{}
136 - out = append(out, tag)
137 - }
138 - if len(out) == 0 {
139 - return nil
140 - }
141 - return out
142 -}
143 -
93 func HostPortOrLoopback(addr string) string {
94 host, port, err := net.SplitHostPort(addr)
95 if err != nil {
portal/server.go
+1 -1
@@ -485,7 +485,7 @@ func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (type
485 ID: leaseID,
486 Name: strings.TrimSpace(req.Name),
487 Hostnames: hostnames,
488 - Metadata: normalizeMetadata(req.Metadata),
488 + Metadata: req.Metadata,
489 ReverseToken: req.ReverseToken,
490 ExpiresAt: expiresAt,
491 FirstSeenAt: now,
sdk/client.go
+155 -86
@@ -22,65 +22,112 @@ import (
22 "github.com/gosuda/portal/v2/types"
23 )
24
25 +const (
26 + defaultDialTimeout = 5 * time.Second
27 + defaultRequestTimeout = 15 * time.Second
28 + defaultHandshakeTimeout = 15 * time.Second
29 + defaultLeaseTTL = 2 * time.Minute
30 + defaultRenewBefore = 30 * time.Second
31 + defaultReadyTarget = 1
32 +)
33 +
34 +// ClientConfig configures the SDK client.
35 type ClientConfig struct {
26 - RelayURL string
27 - RootCAPEM []byte
28 - InsecureSkipVerify bool
29 - DialTimeout time.Duration
30 - RequestTimeout time.Duration
31 - HandshakeTimeout time.Duration
32 - LeaseTTL time.Duration
33 - RenewBefore time.Duration
34 - ReadyTarget int
36 + RelayURLs []string
37 + RootCAPEM []byte
38 }
39
40 type Client struct {
38 - baseURL *url.URL
39 - httpClient *http.Client
40 - rawTLSConfig *tls.Config
41 - dialTimeout time.Duration
42 - handshakeTimeout time.Duration
43 - leaseTTL time.Duration
44 - renewBefore time.Duration
45 - readyTarget int
41 + clients []*relayClient
42 +}
43 +
44 +type relayClient struct {
45 + baseURL *url.URL
46 + httpClient *http.Client
47 + rawTLSConfig *tls.Config
48 }
49
50 func NewClient(cfg ClientConfig) (*Client, error) {
49 - baseURL, err := url.Parse(strings.TrimSpace(cfg.RelayURL))
51 + relayURLs, err := normalizeRelayURLs(cfg.RelayURLs)
52 if err != nil {
51 - return nil, fmt.Errorf("parse relay url: %w", err)
53 + return nil, err
54 }
53 - if !strings.EqualFold(baseURL.Scheme, "https") {
54 - return nil, fmt.Errorf("relay url must use https: %q", cfg.RelayURL)
55 +
56 + clients := make([]*relayClient, 0, len(relayURLs))
57 + for _, relayURL := range relayURLs {
58 + client, err := newRelayClient(cfg, relayURL)
59 + if err != nil {
60 + for _, existing := range clients {
61 + existing.Close()
62 + }
63 + return nil, err
64 + }
65 + clients = append(clients, client)
66 }
56 - if baseURL.Host == "" {
57 - return nil, fmt.Errorf("relay url host is empty: %q", cfg.RelayURL)
67 +
68 + return &Client{clients: clients}, nil
69 +}
70 +
71 +func (c *Client) Close() {
72 + if c == nil {
73 + return
74 }
59 - baseURL.Path = strings.TrimRight(baseURL.Path, "/")
60 - baseURL.RawQuery = ""
61 - baseURL.Fragment = ""
75
63 - if cfg.DialTimeout <= 0 {
64 - cfg.DialTimeout = 5 * time.Second
76 + for _, client := range c.clients {
77 + client.Close()
78 }
66 - if cfg.RequestTimeout <= 0 {
67 - cfg.RequestTimeout = 15 * time.Second
79 +}
80 +
81 +func (c *Client) Listen(ctx context.Context, req ListenRequest) (*Listener, error) {
82 + if strings.TrimSpace(req.Name) == "" {
83 + return nil, errors.New("listener name is required")
84 }
69 - if cfg.HandshakeTimeout <= 0 {
70 - cfg.HandshakeTimeout = 15 * time.Second
85 + if len(c.clients) == 0 {
86 + return nil, errors.New("no relay urls configured")
87 }
72 - if cfg.LeaseTTL <= 0 {
73 - cfg.LeaseTTL = 2 * time.Minute
88 +
89 + listenerCtx, cancel := context.WithCancel(ctx)
90 + listener := &Listener{
91 + baseContext: func() context.Context { return listenerCtx },
92 + ctxDone: listenerCtx.Done(),
93 + cancel: cancel,
94 }
75 - if cfg.RenewBefore <= 0 {
76 - cfg.RenewBefore = 30 * time.Second
95 +
96 + entries := make([]*listenerLease, 0, len(c.clients))
97 + acceptedCap := 0
98 + for _, client := range c.clients {
99 + entry, entryAcceptedCap, err := client.listenEntry(listener, req)
100 + if err != nil {
101 + cancel()
102 + return nil, errors.Join(err, closeListenerEntries(entries))
103 + }
104 + entries = append(entries, entry)
105 + acceptedCap += entryAcceptedCap
106 + }
107 +
108 + if acceptedCap <= 0 {
109 + acceptedCap = len(entries)
110 }
78 - if cfg.ReadyTarget <= 0 {
79 - cfg.ReadyTarget = 1
111 +
112 + listener.accepted = make(chan acceptedConn, acceptedCap)
113 + listener.entries = entries
114 + for _, entry := range entries {
115 + go entry.runSupervisor()
116 + go entry.runRenewLoop()
117 + entry.notify()
118 }
119
82 - if len(cfg.RootCAPEM) == 0 && !cfg.InsecureSkipVerify && isLocalRelayHost(baseURL.Hostname()) {
83 - bootstrapCtx, cancel := context.WithTimeout(context.Background(), cfg.DialTimeout+cfg.HandshakeTimeout)
120 + return listener, nil
121 +}
122 +
123 +func newRelayClient(cfg ClientConfig, relayURL string) (*relayClient, error) {
124 + baseURL, err := url.Parse(relayURL)
125 + if err != nil {
126 + return nil, fmt.Errorf("parse relay url: %w", err)
127 + }
128 +
129 + if len(cfg.RootCAPEM) == 0 && isLocalRelayHost(baseURL.Hostname()) {
130 + bootstrapCtx, cancel := context.WithTimeout(context.Background(), defaultDialTimeout+defaultHandshakeTimeout)
131 defer cancel()
132
133 _, rootCAPEM, bootstrapErr := keyless.ResolveMaterials(bootstrapCtx, baseURL.String(), baseURL.Hostname())
@@ -96,11 +143,10 @@ func NewClient(cfg ClientConfig) (*Client, error) {
143 }
144
145 baseTLS := &tls.Config{
99 - MinVersion: tls.VersionTLS12,
100 - ServerName: baseURL.Hostname(),
101 - RootCAs: rootCAs,
102 - InsecureSkipVerify: cfg.InsecureSkipVerify,
103 - NextProtos: []string{"http/1.1"},
146 + MinVersion: tls.VersionTLS12,
147 + ServerName: baseURL.Hostname(),
148 + RootCAs: rootCAs,
149 + NextProtos: []string{"http/1.1"},
150 }
151
152 transport := &http.Transport{
@@ -108,22 +154,55 @@ func NewClient(cfg ClientConfig) (*Client, error) {
154 ForceAttemptHTTP2: false,
155 }
156
111 - return &Client{
157 + return &relayClient{
158 baseURL: baseURL,
159 httpClient: &http.Client{
160 Transport: transport,
115 - Timeout: cfg.RequestTimeout,
161 + Timeout: defaultRequestTimeout,
162 },
117 - rawTLSConfig: baseTLS,
118 - dialTimeout: cfg.DialTimeout,
119 - handshakeTimeout: cfg.HandshakeTimeout,
120 - leaseTTL: cfg.LeaseTTL,
121 - renewBefore: cfg.RenewBefore,
122 - readyTarget: cfg.ReadyTarget,
163 + rawTLSConfig: baseTLS,
164 }, nil
165 }
166
126 -func (c *Client) Close() {
167 +func normalizeRelayURLs(rawURLs []string) ([]string, error) {
168 + if len(rawURLs) == 0 {
169 + return nil, errors.New("relay url is required")
170 + }
171 +
172 + seen := make(map[string]struct{}, len(rawURLs))
173 + relayURLs := make([]string, 0, len(rawURLs))
174 + for _, raw := range rawURLs {
175 + normalized, err := normalizeRelayURL(raw)
176 + if err != nil {
177 + return nil, err
178 + }
179 + if _, ok := seen[normalized]; ok {
180 + continue
181 + }
182 + seen[normalized] = struct{}{}
183 + relayURLs = append(relayURLs, normalized)
184 + }
185 + return relayURLs, nil
186 +}
187 +
188 +func normalizeRelayURL(raw string) (string, error) {
189 + baseURL, err := url.Parse(strings.TrimSpace(raw))
190 + if err != nil {
191 + return "", fmt.Errorf("parse relay url: %w", err)
192 + }
193 + if !strings.EqualFold(baseURL.Scheme, "https") {
194 + return "", fmt.Errorf("relay url must use https: %q", raw)
195 + }
196 + if baseURL.Host == "" {
197 + return "", fmt.Errorf("relay url host is empty: %q", raw)
198 + }
199 + baseURL.Path = strings.TrimRight(baseURL.Path, "/")
200 + baseURL.RawQuery = ""
201 + baseURL.Fragment = ""
202 + return baseURL.String(), nil
203 +}
204 +
205 +func (c *relayClient) Close() {
206 if c == nil || c.httpClient == nil {
207 return
208 }
@@ -132,11 +211,7 @@ func (c *Client) Close() {
211 }
212 }
213
135 -func (c *Client) Listen(ctx context.Context, req ListenRequest) (*Listener, error) {
136 - if strings.TrimSpace(req.Name) == "" {
137 - return nil, errors.New("listener name is required")
138 - }
139 -
214 +func (c *relayClient) listenEntry(listener *Listener, req ListenRequest) (*listenerLease, int, error) {
215 reverseToken := strings.TrimSpace(req.ReverseToken)
216 if reverseToken == "" {
217 reverseToken = randomToken()
@@ -144,11 +219,11 @@ func (c *Client) Listen(ctx context.Context, req ListenRequest) (*Listener, erro
219
220 readyTarget := req.ReadyTarget
221 if readyTarget <= 0 {
147 - readyTarget = c.readyTarget
222 + readyTarget = defaultReadyTarget
223 }
224 leaseTTL := req.LeaseTTL
225 if leaseTTL <= 0 {
151 - leaseTTL = c.leaseTTL
226 + leaseTTL = defaultLeaseTTL
227 }
228 acceptedCap := max(readyTarget*2, 1)
229
@@ -162,41 +237,35 @@ func (c *Client) Listen(ctx context.Context, req ListenRequest) (*Listener, erro
237 }
238
239 var registerResp types.RegisterResponse
165 - if err := c.doJSON(ctx, http.MethodPost, types.PathSDKRegister, registerReq, &registerResp); err != nil {
166 - return nil, err
240 + if err := c.doJSON(listener.baseContext(), http.MethodPost, types.PathSDKRegister, registerReq, &registerResp); err != nil {
241 + return nil, 0, err
242 }
243
244 tlsConf, tlsCloser, err := keyless.BuildClientTLSConfig(c.baseURL.String(), registerResp.Hostnames)
245 if err != nil {
246 _ = c.unregisterLease(context.Background(), registerResp.LeaseID, reverseToken)
172 - return nil, err
247 + return nil, 0, err
248 }
249
175 - listenerCtx, cancel := context.WithCancel(ctx)
176 - l := &Listener{
177 - client: c,
178 - baseContext: func() context.Context { return listenerCtx },
179 - ctxDone: listenerCtx.Done(),
180 - cancel: cancel,
181 - leaseID: registerResp.LeaseID,
182 - hostnames: append([]string(nil), registerResp.Hostnames...),
183 - metadata: registerResp.Metadata,
250 + return &listenerLease{
251 + parent: listener,
252 + client: c,
253 + info: ListenerEntry{
254 + RelayURL: c.baseURL.String(),
255 + LeaseID: registerResp.LeaseID,
256 + Hostnames: append([]string(nil), registerResp.Hostnames...),
257 + Metadata: registerResp.Metadata,
258 + },
259 reverseToken: reverseToken,
260 leaseTTL: leaseTTL,
261 readyTarget: readyTarget,
262 tlsConfig: tlsConf,
263 tlsCloser: tlsCloser,
189 - accepted: make(chan net.Conn, acceptedCap),
264 signal: make(chan struct{}, 1),
191 - }
192 -
193 - go l.runSupervisor()
194 - go l.runRenewLoop()
195 - l.notify()
196 - return l, nil
265 + }, acceptedCap, nil
266 }
267
199 -func (c *Client) doJSON(ctx context.Context, method, path string, payload any, out any) error {
268 +func (c *relayClient) doJSON(ctx context.Context, method, path string, payload any, out any) error {
269 var body io.Reader
270 if payload != nil {
271 buf, err := json.Marshal(payload)
@@ -234,7 +303,7 @@ func (c *Client) doJSON(ctx context.Context, method, path string, payload any, o
303 return json.Unmarshal(envelope.Data, out)
304 }
305
237 -func (c *Client) renewLease(ctx context.Context, leaseID, reverseToken string, ttl time.Duration) error {
306 +func (c *relayClient) renewLease(ctx context.Context, leaseID, reverseToken string, ttl time.Duration) error {
307 return c.doJSON(ctx, http.MethodPost, types.PathSDKRenew, types.RenewRequest{
308 LeaseID: leaseID,
309 ReverseToken: reverseToken,
@@ -242,16 +311,16 @@ func (c *Client) renewLease(ctx context.Context, leaseID, reverseToken string, t
311 }, &types.RenewResponse{})
312 }
313
245 -func (c *Client) unregisterLease(ctx context.Context, leaseID, reverseToken string) error {
314 +func (c *relayClient) unregisterLease(ctx context.Context, leaseID, reverseToken string) error {
315 return c.doJSON(ctx, http.MethodPost, types.PathSDKUnregister, types.UnregisterRequest{
316 LeaseID: leaseID,
317 ReverseToken: reverseToken,
318 }, nil)
319 }
320
252 -func (c *Client) openReverseSession(ctx context.Context, leaseID, reverseToken string) (net.Conn, error) {
321 +func (c *relayClient) openReverseSession(ctx context.Context, leaseID, reverseToken string) (net.Conn, error) {
322 dialer := &tls.Dialer{
254 - NetDialer: &net.Dialer{Timeout: c.dialTimeout},
323 + NetDialer: &net.Dialer{Timeout: defaultDialTimeout},
324 Config: c.rawTLSConfig.Clone(),
325 }
326
@@ -300,7 +369,7 @@ func (c *Client) openReverseSession(ctx context.Context, leaseID, reverseToken s
369 return wrapBufferedConn(conn, reader), nil
370 }
371
303 -func (c *Client) resolve(path string) string {
372 +func (c *relayClient) resolve(path string) string {
373 ref, _ := url.Parse(path)
374 return c.baseURL.ResolveReference(ref).String()
375 }
sdk/client_test.go
+43 -2
@@ -14,13 +14,17 @@ func TestNewClientAutoTrustsLocalhostRelayCertificate(t *testing.T) {
14 }))
15 defer server.Close()
16
17 - client, err := NewClient(ClientConfig{RelayURL: server.URL})
17 + client, err := NewClient(ClientConfig{RelayURLs: []string{server.URL}})
18 if err != nil {
19 t.Fatalf("NewClient() error = %v", err)
20 }
21 defer client.Close()
22
23 - resp, err := client.httpClient.Get(server.URL)
23 + if len(client.clients) != 1 {
24 + t.Fatalf("client count = %d, want 1", len(client.clients))
25 + }
26 +
27 + resp, err := client.clients[0].httpClient.Get(server.URL)
28 if err != nil {
29 t.Fatalf("httpClient.Get() error = %v", err)
30 }
@@ -30,3 +34,40 @@ func TestNewClientAutoTrustsLocalhostRelayCertificate(t *testing.T) {
34 t.Fatalf("status = %d, want %d", resp.StatusCode, http.StatusOK)
35 }
36 }
37 +
38 +func TestNewClientSupportsDedupedRelayURLs(t *testing.T) {
39 + t.Parallel()
40 +
41 + serverA := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
42 + w.WriteHeader(http.StatusOK)
43 + }))
44 + defer serverA.Close()
45 +
46 + serverB := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
47 + w.WriteHeader(http.StatusOK)
48 + }))
49 + defer serverB.Close()
50 +
51 + client, err := NewClient(ClientConfig{
52 + RelayURLs: []string{serverA.URL, serverB.URL},
53 + })
54 + if err != nil {
55 + t.Fatalf("NewClient() error = %v", err)
56 + }
57 + defer client.Close()
58 +
59 + if len(client.clients) != 2 {
60 + t.Fatalf("client count = %d, want 2", len(client.clients))
61 + }
62 +
63 + for i, relayClient := range client.clients {
64 + resp, err := relayClient.httpClient.Get(relayClient.baseURL.String())
65 + if err != nil {
66 + t.Fatalf("client[%d].httpClient.Get() error = %v", i, err)
67 + }
68 + _ = resp.Body.Close()
69 + if resp.StatusCode != http.StatusOK {
70 + t.Fatalf("client[%d] status = %d, want %d", i, resp.StatusCode, http.StatusOK)
71 + }
72 + }
73 +}
sdk/helper.go new
+100
@@ -0,0 +1,100 @@
1 +package sdk
2 +
3 +import (
4 + "context"
5 + "errors"
6 + "fmt"
7 + "net"
8 + "net/http"
9 + "strings"
10 + "time"
11 +
12 + "golang.org/x/sync/errgroup"
13 +)
14 +
15 +const defaultHTTPShutdownTimeout = 5 * time.Second
16 +
17 +type HTTPServeOptions struct {
18 + LocalAddr string
19 + ReadHeaderTimeout time.Duration
20 +}
21 +
22 +// RunHTTPApp serves one handler on the relay listener and, optionally, on a
23 +// local HTTP address for app-local access.
24 +func RunHTTPApp(ctx context.Context, relayListener net.Listener, handler http.Handler, opts HTTPServeOptions) error {
25 + readHeaderTimeout := opts.ReadHeaderTimeout
26 + if readHeaderTimeout <= 0 {
27 + readHeaderTimeout = defaultRequestTimeout
28 + }
29 +
30 + relaySrv := &http.Server{
31 + Handler: handler,
32 + ReadHeaderTimeout: readHeaderTimeout,
33 + }
34 +
35 + var localSrv *http.Server
36 + if opts.LocalAddr != "" {
37 + localSrv = &http.Server{
38 + Addr: opts.LocalAddr,
39 + Handler: handler,
40 + ReadHeaderTimeout: readHeaderTimeout,
41 + }
42 + }
43 +
44 + group, groupCtx := errgroup.WithContext(ctx)
45 + if localSrv != nil {
46 + group.Go(func() error {
47 + if err := localSrv.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
48 + return fmt.Errorf("serve local http: %w", err)
49 + }
50 + return nil
51 + })
52 + }
53 + group.Go(func() error {
54 + if err := relaySrv.Serve(relayListener); err != nil && !errors.Is(err, http.ErrServerClosed) {
55 + return fmt.Errorf("serve relay http: %w", err)
56 + }
57 + return nil
58 + })
59 + group.Go(func() error {
60 + <-groupCtx.Done()
61 +
62 + shutdownCtx, cancel := context.WithTimeout(context.Background(), defaultHTTPShutdownTimeout)
63 + defer cancel()
64 +
65 + var localErr error
66 + if localSrv != nil {
67 + localErr = localSrv.Shutdown(shutdownCtx)
68 + if errors.Is(localErr, http.ErrServerClosed) {
69 + localErr = nil
70 + }
71 + }
72 +
73 + relayErr := relaySrv.Shutdown(shutdownCtx)
74 + if errors.Is(relayErr, http.ErrServerClosed) {
75 + relayErr = nil
76 + }
77 +
78 + return errors.Join(localErr, relayErr)
79 + })
80 +
81 + return group.Wait()
82 +}
83 +
84 +// SplitCSV splits a comma-separated string, trimming whitespace and dropping
85 +// empty entries.
86 +func SplitCSV(raw string) []string {
87 + if strings.TrimSpace(raw) == "" {
88 + return nil
89 + }
90 +
91 + parts := strings.Split(raw, ",")
92 + out := make([]string, 0, len(parts))
93 + for _, part := range parts {
94 + part = strings.TrimSpace(part)
95 + if part != "" {
96 + out = append(out, part)
97 + }
98 + }
99 + return out
100 +}
sdk/helper_test.go new
+119
@@ -0,0 +1,119 @@
1 +package sdk
2 +
3 +import (
4 + "context"
5 + "io"
6 + "net"
7 + "net/http"
8 + "testing"
9 + "time"
10 +)
11 +
12 +func TestRunHTTPAppRelayOnly(t *testing.T) {
13 + t.Parallel()
14 +
15 + listener, err := net.Listen("tcp", "127.0.0.1:0")
16 + if err != nil {
17 + t.Fatalf("Listen() error = %v", err)
18 + }
19 + defer listener.Close()
20 +
21 + ctx, cancel := context.WithCancel(context.Background())
22 + defer cancel()
23 +
24 + errCh := make(chan error, 1)
25 + go func() {
26 + errCh <- RunHTTPApp(ctx, listener, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
27 + _, _ = io.WriteString(w, "ok")
28 + }), HTTPServeOptions{})
29 + }()
30 +
31 + waitForHTTP(t, "http://"+listener.Addr().String())
32 + cancel()
33 +
34 + select {
35 + case err := <-errCh:
36 + if err != nil {
37 + t.Fatalf("RunHTTPApp() error = %v", err)
38 + }
39 + case <-time.After(3 * time.Second):
40 + t.Fatal("RunHTTPApp() did not exit after context cancellation")
41 + }
42 +}
43 +
44 +func TestRunHTTPAppLocalAndRelay(t *testing.T) {
45 + t.Parallel()
46 +
47 + relayListener, err := net.Listen("tcp", "127.0.0.1:0")
48 + if err != nil {
49 + t.Fatalf("Listen() error = %v", err)
50 + }
51 + defer relayListener.Close()
52 +
53 + localListener, err := net.Listen("tcp", "127.0.0.1:0")
54 + if err != nil {
55 + t.Fatalf("Listen() error = %v", err)
56 + }
57 + localAddr := localListener.Addr().String()
58 + _ = localListener.Close()
59 +
60 + ctx, cancel := context.WithCancel(context.Background())
61 + defer cancel()
62 +
63 + errCh := make(chan error, 1)
64 + go func() {
65 + errCh <- RunHTTPApp(ctx, relayListener, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
66 + _, _ = io.WriteString(w, "ok")
67 + }), HTTPServeOptions{
68 + LocalAddr: localAddr,
69 + })
70 + }()
71 +
72 + waitForHTTP(t, "http://"+relayListener.Addr().String())
73 + waitForHTTP(t, "http://"+localAddr)
74 + cancel()
75 +
76 + select {
77 + case err := <-errCh:
78 + if err != nil {
79 + t.Fatalf("RunHTTPApp() error = %v", err)
80 + }
81 + case <-time.After(3 * time.Second):
82 + t.Fatal("RunHTTPApp() did not exit after context cancellation")
83 + }
84 +}
85 +
86 +func TestSplitCSV(t *testing.T) {
87 + t.Parallel()
88 +
89 + got := SplitCSV(" a, ,b,c ,, d ")
90 + want := []string{"a", "b", "c", "d"}
91 +
92 + if len(got) != len(want) {
93 + t.Fatalf("SplitCSV() len = %d, want %d", len(got), len(want))
94 + }
95 + for i := range want {
96 + if got[i] != want[i] {
97 + t.Fatalf("SplitCSV()[%d] = %q, want %q", i, got[i], want[i])
98 + }
99 + }
100 +}
101 +
102 +func waitForHTTP(t *testing.T, rawURL string) {
103 + t.Helper()
104 +
105 + client := &http.Client{Timeout: 200 * time.Millisecond}
106 + deadline := time.Now().Add(3 * time.Second)
107 + for time.Now().Before(deadline) {
108 + resp, err := client.Get(rawURL)
109 + if err == nil {
110 + _ = resp.Body.Close()
111 + if resp.StatusCode == http.StatusOK {
112 + return
113 + }
114 + }
115 + time.Sleep(50 * time.Millisecond)
116 + }
117 +
118 + t.Fatalf("timed out waiting for %s", rawURL)
119 +}
sdk/listener.go
+146 -59
@@ -10,6 +10,8 @@ import (
10 "sync"
11 "time"
12
13 + "golang.org/x/sync/errgroup"
14 +
15 "github.com/gosuda/portal/v2/types"
16 )
17
@@ -22,35 +24,71 @@ type ListenRequest struct {
24 LeaseTTL time.Duration
25 }
26
27 +type ListenerEntry struct {
28 + RelayURL string
29 + LeaseID string
30 + Hostnames []string
31 + Metadata types.LeaseMetadata
32 +}
33 +
34 +// PublicURLs returns the HTTPS URLs exposed by one relay-specific lease.
35 +func (e ListenerEntry) PublicURLs() []string {
36 + urls := make([]string, 0, len(e.Hostnames))
37 + for _, host := range e.Hostnames {
38 + urls = append(urls, "https://"+host)
39 + }
40 + return urls
41 +}
42 +
43 +func (e ListenerEntry) clone() ListenerEntry {
44 + e.Hostnames = append([]string(nil), e.Hostnames...)
45 + return e
46 +}
47 +
48 type Listener struct {
49 + baseContext func() context.Context
50 + ctxDone <-chan struct{}
51 + cancel context.CancelFunc
52 + accepted chan acceptedConn
53 + entries []*listenerLease
54 + closeOnce sync.Once
55 +}
56 +
57 +type listenerLease struct {
58 tlsCloser io.Closer
59 tlsConfig *tls.Config
28 - baseContext func() context.Context
29 - ctxDone <-chan struct{}
30 - cancel context.CancelFunc
31 - client *Client
60 + parent *Listener
61 + client *relayClient
62 signal chan struct{}
33 - accepted chan net.Conn
34 - leaseID string
63 + info ListenerEntry
64 reverseToken string
36 - hostnames []string
37 - metadata types.LeaseMetadata
65 readyTarget int
66 leaseTTL time.Duration
67 activeSessions int
41 - closeOnce sync.Once
68 mu sync.Mutex
69 }
70
71 +type acceptedConn struct {
72 + conn net.Conn
73 + entry ListenerEntry
74 +}
75 +
76 func (l *Listener) Accept() (net.Conn, error) {
77 + conn, _, err := l.AcceptEntry()
78 + return conn, err
79 +}
80 +
81 +// AcceptEntry returns the next accepted connection plus relay-specific lease
82 +// metadata for callers that need to distinguish which relay claimed it.
83 +func (l *Listener) AcceptEntry() (net.Conn, ListenerEntry, error) {
84 select {
85 case <-l.ctxDone:
48 - return nil, net.ErrClosed
49 - case conn := <-l.accepted:
50 - if conn == nil {
51 - return nil, net.ErrClosed
86 + return nil, ListenerEntry{}, net.ErrClosed
87 + case accepted := <-l.accepted:
88 + if accepted.conn == nil {
89 + return nil, ListenerEntry{}, net.ErrClosed
90 }
53 - return conn, nil
91 + return accepted.conn, accepted.entry.clone(), nil
92 }
93 }
94
@@ -58,47 +96,79 @@ func (l *Listener) Close() error {
96 var closeErr error
97 l.closeOnce.Do(func() {
98 l.cancel()
61 -
62 - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
63 - defer cancel()
64 - if err := l.client.unregisterLease(ctx, l.leaseID, l.reverseToken); err != nil {
65 - closeErr = err
66 - }
67 - if l.tlsCloser != nil {
68 - _ = l.tlsCloser.Close()
69 - }
99 + closeErr = closeListenerEntries(l.entries)
100 })
101 return closeErr
102 }
103
104 func (l *Listener) Addr() net.Addr {
75 - return listenerAddr("portal:" + l.leaseID)
76 -}
77 -
78 -func (l *Listener) LeaseID() string {
79 - return l.leaseID
105 + if entry, ok := l.singleEntry(); ok {
106 + return listenerAddr("portal:" + entry.LeaseID)
107 + }
108 + return listenerAddr("portal:multi")
109 }
110
82 -func (l *Listener) Hostnames() []string {
83 - return append([]string(nil), l.hostnames...)
111 +// Entries returns relay-specific lease details for advanced multi-relay callers.
112 +func (l *Listener) Entries() []ListenerEntry {
113 + entries := make([]ListenerEntry, 0, len(l.entries))
114 + for _, entry := range l.entries {
115 + entries = append(entries, entry.info.clone())
116 + }
117 + return entries
118 }
119
86 -func (l *Listener) Metadata() types.LeaseMetadata {
87 - return l.metadata
120 +func (l *Listener) singleEntry() (ListenerEntry, bool) {
121 + if len(l.entries) != 1 {
122 + return ListenerEntry{}, false
123 + }
124 + return l.entries[0].info.clone(), true
125 }
126
127 +// PublicURLs returns all public HTTPS URLs exposed by the listener.
128 func (l *Listener) PublicURLs() []string {
91 - urls := make([]string, 0, len(l.hostnames))
92 - for _, host := range l.hostnames {
93 - urls = append(urls, "https://"+host)
129 + var urls []string
130 + for _, entry := range l.entries {
131 + urls = append(urls, entry.info.PublicURLs()...)
132 }
133 return urls
134 }
135
98 -func (l *Listener) runSupervisor() {
136 +func closeListenerEntries(entries []*listenerLease) error {
137 + if len(entries) == 0 {
138 + return nil
139 + }
140 +
141 + var closeErr error
142 + var mu sync.Mutex
143 + var group errgroup.Group
144 + group.SetLimit(min(len(entries), 4))
145 +
146 + for _, entry := range entries {
147 + if entry == nil {
148 + continue
149 + }
150 + group.Go(func() error {
151 + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
152 + defer cancel()
153 +
154 + if err := entry.close(ctx); err != nil {
155 + mu.Lock()
156 + closeErr = errors.Join(closeErr, err)
157 + mu.Unlock()
158 + }
159 + return nil
160 + })
161 + }
162 +
163 + _ = group.Wait()
164 +
165 + return closeErr
166 +}
167 +
168 +func (l *listenerLease) runSupervisor() {
169 for {
170 select {
101 - case <-l.ctxDone:
171 + case <-l.parent.ctxDone:
172 return
173 case <-l.signal:
174 }
@@ -109,13 +179,13 @@ func (l *Listener) runSupervisor() {
179 }
180 }
181
112 -func (l *Listener) runRenewLoop() {
182 +func (l *listenerLease) runRenewLoop() {
183 interval := l.leaseTTL / 2
184 if interval <= 0 {
185 interval = 30 * time.Second
186 }
117 - if l.client.renewBefore > 0 && l.leaseTTL > l.client.renewBefore {
118 - interval = l.leaseTTL - l.client.renewBefore
187 + if defaultRenewBefore > 0 && l.leaseTTL > defaultRenewBefore {
188 + interval = l.leaseTTL - defaultRenewBefore
189 }
190 if interval <= 0 {
191 interval = 30 * time.Second
@@ -126,21 +196,21 @@ func (l *Listener) runRenewLoop() {
196
197 for {
198 select {
129 - case <-l.ctxDone:
199 + case <-l.parent.ctxDone:
200 return
201 case <-ticker.C:
202 ctx, cancel := context.WithTimeout(l.context(), 10*time.Second)
133 - _ = l.client.renewLease(ctx, l.leaseID, l.reverseToken, l.leaseTTL)
203 + _ = l.client.renewLease(ctx, l.info.LeaseID, l.reverseToken, l.leaseTTL)
204 cancel()
205 }
206 }
207 }
208
139 -func (l *Listener) runSession() {
209 +func (l *listenerLease) runSession() {
210 defer l.releaseSessionSlot()
211
212 sessionCtx := l.context()
143 - conn, err := l.client.openReverseSession(sessionCtx, l.leaseID, l.reverseToken)
213 + conn, err := l.client.openReverseSession(sessionCtx, l.info.LeaseID, l.reverseToken)
214 if err != nil {
215 sleepOrDone(sessionCtx, time.Second)
216 return
@@ -154,10 +224,10 @@ func (l *Listener) runSession() {
224 }
225 }
226
157 -func (l *Listener) awaitActivation(conn net.Conn) error {
227 +func (l *listenerLease) awaitActivation(conn net.Conn) error {
228 var marker [1]byte
229 for {
160 - _ = conn.SetReadDeadline(time.Now().Add(2 * l.client.handshakeTimeout))
230 + _ = conn.SetReadDeadline(time.Now().Add(2 * defaultHandshakeTimeout))
231 if _, err := io.ReadFull(conn, marker[:]); err != nil {
232 return err
233 }
@@ -174,24 +244,24 @@ func (l *Listener) awaitActivation(conn net.Conn) error {
244 }
245 }
246
177 -func (l *Listener) activate(conn net.Conn) error {
247 +func (l *listenerLease) activate(conn net.Conn) error {
248 tlsConn := tls.Server(conn, l.tlsConfig.Clone())
179 - handshakeCtx, cancel := context.WithTimeout(l.context(), l.client.handshakeTimeout)
249 + handshakeCtx, cancel := context.WithTimeout(l.context(), defaultHandshakeTimeout)
250 defer cancel()
251 if err := tlsConn.HandshakeContext(handshakeCtx); err != nil {
252 return err
253 }
254
255 select {
186 - case <-l.ctxDone:
256 + case <-l.parent.ctxDone:
257 _ = tlsConn.Close()
258 return l.context().Err()
189 - case l.accepted <- tlsConn:
259 + case l.parent.accepted <- acceptedConn{conn: tlsConn, entry: l.info.clone()}:
260 return nil
261 }
262 }
263
194 -func (l *Listener) reserveSessionSlot() bool {
264 +func (l *listenerLease) reserveSessionSlot() bool {
265 l.mu.Lock()
266 defer l.mu.Unlock()
267 if l.isClosed() {
@@ -204,20 +274,37 @@ func (l *Listener) reserveSessionSlot() bool {
274 return true
275 }
276
207 -func (l *Listener) releaseSessionSlot() {
277 +func (l *listenerLease) releaseSessionSlot() {
278 l.mu.Lock()
279 l.activeSessions--
280 l.mu.Unlock()
281 l.notify()
282 }
283
214 -func (l *Listener) notify() {
284 +func (l *listenerLease) notify() {
285 select {
286 case l.signal <- struct{}{}:
287 default:
288 }
289 }
290
291 +func (l *listenerLease) close(ctx context.Context) error {
292 + if l == nil {
293 + return nil
294 + }
295 +
296 + var closeErr error
297 + if l.client != nil {
298 + if err := l.client.unregisterLease(ctx, l.info.LeaseID, l.reverseToken); err != nil {
299 + closeErr = errors.Join(closeErr, err)
300 + }
301 + }
302 + if l.tlsCloser != nil {
303 + closeErr = errors.Join(closeErr, l.tlsCloser.Close())
304 + }
305 + return closeErr
306 +}
307 +
308 func sleepOrDone(ctx context.Context, d time.Duration) {
309 timer := time.NewTimer(d)
310 defer timer.Stop()
@@ -232,21 +319,21 @@ type listenerAddr string
319 func (a listenerAddr) Network() string { return "portal" }
320 func (a listenerAddr) String() string { return string(a) }
321
235 -func (l *Listener) context() context.Context {
236 - if l.baseContext != nil {
237 - if ctx := l.baseContext(); ctx != nil {
322 +func (l *listenerLease) context() context.Context {
323 + if l.parent != nil && l.parent.baseContext != nil {
324 + if ctx := l.parent.baseContext(); ctx != nil {
325 return ctx
326 }
327 }
328 return context.Background()
329 }
330
244 -func (l *Listener) isClosed() bool {
245 - if l.ctxDone == nil {
331 +func (l *listenerLease) isClosed() bool {
332 + if l.parent == nil || l.parent.ctxDone == nil {
333 return false
334 }
335 select {
249 - case <-l.ctxDone:
336 + case <-l.parent.ctxDone:
337 return true
338 default:
339 return false
sdk/listener_test.go new
+118
@@ -0,0 +1,118 @@
1 +package sdk
2 +
3 +import (
4 + "net"
5 + "testing"
6 +)
7 +
8 +func TestListenerSingleEntryAccessors(t *testing.T) {
9 + t.Parallel()
10 +
11 + listener := &Listener{
12 + entries: []*listenerLease{
13 + {
14 + info: ListenerEntry{
15 + RelayURL: "https://relay.example.com",
16 + LeaseID: "lease-1",
17 + Hostnames: []string{"app.relay.example.com"},
18 + },
19 + },
20 + },
21 + }
22 +
23 + entry, ok := listener.singleEntry()
24 + if !ok {
25 + t.Fatal("singleEntry() ok = false, want true")
26 + }
27 + if entry.LeaseID != "lease-1" {
28 + t.Fatalf("singleEntry().LeaseID = %q, want %q", entry.LeaseID, "lease-1")
29 + }
30 +
31 + publicURLs := listener.PublicURLs()
32 + if len(publicURLs) != 1 || publicURLs[0] != "https://app.relay.example.com" {
33 + t.Fatalf("PublicURLs() = %#v, want [https://app.relay.example.com]", publicURLs)
34 + }
35 +}
36 +
37 +func TestListenerMultiEntryAccessors(t *testing.T) {
38 + t.Parallel()
39 +
40 + listener := &Listener{
41 + entries: []*listenerLease{
42 + {
43 + info: ListenerEntry{
44 + RelayURL: "https://relay-a.example.com",
45 + LeaseID: "lease-a",
46 + Hostnames: []string{"a.example.com"},
47 + },
48 + },
49 + {
50 + info: ListenerEntry{
51 + RelayURL: "https://relay-b.example.com",
52 + LeaseID: "lease-b",
53 + Hostnames: []string{"b.example.com"},
54 + },
55 + },
56 + },
57 + }
58 +
59 + if _, ok := listener.singleEntry(); ok {
60 + t.Fatal("singleEntry() ok = true, want false")
61 + }
62 +
63 + entries := listener.Entries()
64 + if len(entries) != 2 {
65 + t.Fatalf("Entries() len = %d, want 2", len(entries))
66 + }
67 +
68 + publicURLs := listener.PublicURLs()
69 + if len(publicURLs) != 2 {
70 + t.Fatalf("PublicURLs() len = %d, want 2", len(publicURLs))
71 + }
72 +}
73 +
74 +func TestListenerAcceptEntry(t *testing.T) {
75 + t.Parallel()
76 +
77 + done := make(chan struct{})
78 + serverConn1, clientConn1 := net.Pipe()
79 + defer clientConn1.Close()
80 + serverConn2, clientConn2 := net.Pipe()
81 + defer clientConn2.Close()
82 +
83 + listener := &Listener{
84 + ctxDone: done,
85 + accepted: make(chan acceptedConn, 1),
86 + }
87 + listener.accepted <- acceptedConn{
88 + conn: serverConn1,
89 + entry: ListenerEntry{
90 + RelayURL: "https://relay.example.com",
91 + LeaseID: "lease-1",
92 + Hostnames: []string{"app.relay.example.com"},
93 + },
94 + }
95 +
96 + conn, entry, err := listener.AcceptEntry()
97 + if err != nil {
98 + t.Fatalf("AcceptEntry() error = %v", err)
99 + }
100 + defer conn.Close()
101 +
102 + if conn != serverConn1 {
103 + t.Fatal("AcceptEntry() did not return the original connection")
104 + }
105 + if entry.LeaseID != "lease-1" {
106 + t.Fatalf("AcceptEntry().LeaseID = %q, want %q", entry.LeaseID, "lease-1")
107 + }
108 +
109 + listener.accepted <- acceptedConn{conn: serverConn2}
110 + plainConn, err := listener.Accept()
111 + if err != nil {
112 + t.Fatalf("Accept() error = %v", err)
113 + }
114 + defer plainConn.Close()
115 + if plainConn != serverConn2 {
116 + t.Fatal("Accept() did not return the original connection")
117 + }
118 +}