agents: rebuild project

Kim committed Mar 6, 2026 at 12:20 UTC d37426854c8054244164414fcfed4eb609cc1f6b
61 files changed +3696 -10780
cmd/demo-app/main.go
+58 -81
@@ -1,6 +1,7 @@
1 package main
2
3 import (
4 + "context"
5 "embed"
6 "encoding/base64"
7 "encoding/json"
@@ -14,11 +15,9 @@ import (
15 "syscall"
16 "time"
17
17 - "github.com/rs/zerolog/log"
18 "golang.org/x/net/websocket"
19
20 "gosuda.org/portal/sdk"
21 - "gosuda.org/portal/types"
21 )
22
23 //go:embed static
@@ -42,156 +41,134 @@ func main() {
41 flag.IntVar(&flagPort, "port", 8092, "local demo HTTP port")
42 flag.StringVar(&flagName, "name", "demo-app", "backend display name")
43 flag.StringVar(&flagDesc, "description", "Portal demo connectivity app", "lease description")
45 - flag.StringVar(&flagTags, "tags", "demo,connectivity,activity,cloud,sun,moning", "comma-separated lease tags")
44 + flag.StringVar(&flagTags, "tags", "demo,connectivity,activity,cloud,sun,morning", "comma-separated lease tags")
45 flag.StringVar(&flagOwner, "owner", "PortalApp Developer", "lease owner")
46 flag.BoolVar(&flagHide, "hide", false, "hide this lease from listings")
47 flag.Parse()
48
49 if err := runDemo(); err != nil {
51 - log.Fatal().Err(err).Msg("execute demo command")
50 + fmt.Fprintf(os.Stderr, "demo command failed: %v\n", err)
51 + os.Exit(1)
52 }
53 }
54
55 func runDemo() error {
56 - // 1) Create SDK client and connect to relay(s)
57 - opts := []sdk.ClientOption{sdk.WithBootstrapServers([]string{flagServerURL})}
58 - sdkClient, err := sdk.NewClient(opts...)
56 + sdkClient, err := sdk.NewClient(sdk.ClientConfig{RelayURL: flagServerURL})
57 if err != nil {
58 return fmt.Errorf("new client: %w", err)
59 }
60 defer sdkClient.Close()
61
64 - // 2) Register lease
65 - // Create base64 data URI from embedded thumbnail
62 thumbnailDataURI := "data:image/png;base64," + base64.StdEncoding.EncodeToString(thumbnailPNG)
63
68 - listener, err := sdkClient.Listen(
69 - flagName,
70 - types.WithDescription(flagDesc),
71 - types.WithTags(strings.Split(flagTags, ",")),
72 - types.WithOwner(flagOwner),
73 - types.WithThumbnail(thumbnailDataURI),
74 - types.WithHide(flagHide),
75 - )
64 + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
65 + defer stop()
66 +
67 + listener, err := sdkClient.Listen(ctx, sdk.ListenRequest{
68 + Name: flagName,
69 + Metadata: sdk.LeaseMetadata{
70 + Description: flagDesc,
71 + Tags: splitCSV(flagTags),
72 + Owner: flagOwner,
73 + Thumbnail: thumbnailDataURI,
74 + Hide: flagHide,
75 + },
76 + })
77 if err != nil {
78 return fmt.Errorf("listen: %w", err)
79 }
80 defer listener.Close()
81
81 - // 4) Setup HTTP handler
82 mux := http.NewServeMux()
83
84 - // Serve static files from embedded filesystem
84 staticFS, err := fs.Sub(staticFiles, "static")
85 if err != nil {
86 return fmt.Errorf("create static fs: %w", err)
87 }
88 mux.Handle("/", http.FileServer(http.FS(staticFS)))
89
91 - // Simple HTTP ping endpoint for connectivity checks
90 mux.HandleFunc("/api/ping", func(w http.ResponseWriter, _ *http.Request) {
91 w.Header().Set("Content-Type", "application/json")
92 resp := map[string]any{
93 "message": "pong",
94 "time": time.Now().UTC().Format(time.RFC3339),
95 }
98 - if err := json.NewEncoder(w).Encode(resp); err != nil {
99 - log.Error().Err(err).Msg("write ping response")
100 - }
96 + _ = json.NewEncoder(w).Encode(resp)
97 })
98
103 - // WebSocket echo endpoint
99 mux.Handle("/ws", websocket.Handler(func(conn *websocket.Conn) {
100 defer conn.Close()
101 for {
102 var msg string
103 if err := websocket.Message.Receive(conn, &msg); err != nil {
109 - if err.Error() != "EOF" {
110 - log.Error().Err(err).Msg("websocket read error")
111 - }
104 break
105 }
114 - log.Debug().Str("msg", msg).Msg("websocket received")
106 if err := websocket.Message.Send(conn, "echo: "+msg); err != nil {
116 - log.Error().Err(err).Msg("websocket write error")
107 break
108 }
109 }
110 }))
111
122 - // Test endpoint for multiple Set-Cookie headers
123 - // Note: HttpOnly cookies cannot be set via Service Worker (browser security limitation)
112 mux.HandleFunc("/api/test-cookies", func(w http.ResponseWriter, _ *http.Request) {
125 - http.SetCookie(w, &http.Cookie{
126 - Name: "session_id",
127 - Value: "abc123",
128 - Path: "/",
129 - MaxAge: 3600,
130 - })
131 - http.SetCookie(w, &http.Cookie{
132 - Name: "auth_token",
133 - Value: "secret456",
134 - Path: "/",
135 - MaxAge: 3600,
136 - })
137 - http.SetCookie(w, &http.Cookie{
138 - Name: "csrf_token",
139 - Value: "xyz789",
140 - Path: "/",
141 - MaxAge: 3600,
142 - })
143 - http.SetCookie(w, &http.Cookie{
144 - Name: "user_pref",
145 - Value: "dark_mode",
146 - Path: "/",
147 - MaxAge: 86400,
148 - })
113 + for _, cookie := range []*http.Cookie{
114 + {Name: "session_id", Value: "abc123", Path: "/", MaxAge: 3600},
115 + {Name: "auth_token", Value: "secret456", Path: "/", MaxAge: 3600},
116 + {Name: "csrf_token", Value: "xyz789", Path: "/", MaxAge: 3600},
117 + {Name: "user_pref", Value: "dark_mode", Path: "/", MaxAge: 86400},
118 + } {
119 + http.SetCookie(w, cookie)
120 + }
121 w.Header().Set("Content-Type", "application/json")
150 - if encodeErr := json.NewEncoder(w).Encode(map[string]any{
122 + _ = json.NewEncoder(w).Encode(map[string]any{
123 "message": "4 cookies set: session_id, auth_token, csrf_token, user_pref",
152 - }); encodeErr != nil {
153 - log.Error().Err(encodeErr).Msg("write test-cookies response")
154 - }
124 + })
125 })
126
157 - // 5) Serve HTTP over relay listener
158 - log.Info().Msgf("[demo] serving HTTP over relay; lease=%s", flagName)
159 -
160 - // Also serve on local port for direct testing
127 + localAddr := fmt.Sprintf(":%d", flagPort)
128 go func() {
162 - localAddr := fmt.Sprintf(":%d", flagPort)
163 - log.Info().Msgf("[demo] also serving on local port %s for direct testing", localAddr)
129 localSrv := &http.Server{
130 Addr: localAddr,
131 Handler: mux,
132 ReadHeaderTimeout: 5 * time.Second,
133 }
169 - if err := localSrv.ListenAndServe(); err != nil {
170 - log.Error().Err(err).Msg("local http serve error")
171 - }
134 + _ = localSrv.ListenAndServe()
135 }()
136
174 - srvErr := make(chan error, 1)
175 - go func() {
176 - relaySrv := &http.Server{
177 - Handler: mux,
178 - ReadHeaderTimeout: 5 * time.Second,
179 - }
180 - srvErr <- relaySrv.Serve(listener)
181 - }()
137 + relaySrv := &http.Server{
138 + Handler: mux,
139 + ReadHeaderTimeout: 5 * time.Second,
140 + }
141
142 sig := make(chan os.Signal, 1)
143 signal.Notify(sig, syscall.SIGINT, syscall.SIGTERM)
144
145 + errCh := make(chan error, 1)
146 + go func() {
147 + errCh <- relaySrv.Serve(listener)
148 + }()
149 +
150 select {
151 case <-sig:
188 - log.Info().Msg("[demo] shutting down...")
189 - case err := <-srvErr:
190 - if err != nil {
191 - log.Error().Err(err).Msg("[demo] http serve error")
152 + case err := <-errCh:
153 + if err != nil && err != http.ErrServerClosed {
154 + return err
155 }
156 }
157
195 - log.Info().Msg("[demo] shutdown complete")
158 return nil
159 }
160 +
161 +func splitCSV(raw string) []string {
162 + if strings.TrimSpace(raw) == "" {
163 + return nil
164 + }
165 + parts := strings.Split(raw, ",")
166 + out := make([]string, 0, len(parts))
167 + for _, part := range parts {
168 + part = strings.TrimSpace(part)
169 + if part != "" {
170 + out = append(out, part)
171 + }
172 + }
173 + return out
174 +}
cmd/demo-app/static/index.html
+51 -125
@@ -1,131 +1,57 @@
1 -<!DOCTYPE html>
1 +<!doctype html>
2 <html lang="en">
3 -
3 <head>
5 - <meta charset="UTF-8">
6 - <meta name="viewport" content="width=device-width, initial-scale=1.0, maximum-scale=1.0, user-scalable=no">
7 - <meta name="apple-mobile-web-app-capable" content="yes">
8 - <meta name="mobile-web-app-capable" content="yes">
9 - <title>Portal Demo Connectivity</title>
10 - <link rel="stylesheet" href="style.css">
4 + <meta charset="utf-8">
5 + <meta name="viewport" content="width=device-width, initial-scale=1">
6 + <title>Portal Demo App</title>
7 + <link rel="stylesheet" href="/style.css">
8 </head>
12 -
9 <body>
14 - <div class="container">
15 - <h1>Portal Demo Connectivity</h1>
16 -
17 - <div class="toolbar">
18 - <button id="httpPingBtn">HTTP Ping</button>
19 - <button id="testCookiesBtn">Test Cookies</button>
20 - <button id="wsConnectBtn">WS Connect</button>
21 - <button id="wsSendBtn" disabled>WS Send "hello"</button>
22 - </div>
23 -
24 - <div class="canvas-wrapper">
25 - <pre id="log" class="log"></pre>
26 - </div>
27 -
28 - <div id="status" class="status">Idle</div>
29 - </div>
30 -
31 - <script>
32 - const statusEl = document.getElementById('status');
33 - const logEl = document.getElementById('log');
34 - const httpPingBtn = document.getElementById('httpPingBtn');
35 - const wsConnectBtn = document.getElementById('wsConnectBtn');
36 - const wsSendBtn = document.getElementById('wsSendBtn');
37 -
38 - let ws = null;
39 -
40 - function log(message) {
41 - const time = new Date().toISOString();
42 - logEl.textContent += `[${time}] ${message}\n`;
43 - logEl.scrollTop = logEl.scrollHeight;
44 - }
45 -
46 - httpPingBtn.addEventListener('click', async () => {
47 - statusEl.textContent = 'HTTP: Pinging...';
48 - try {
49 - const res = await fetch('/api/ping');
50 - const json = await res.json();
51 - statusEl.textContent = 'HTTP: OK';
52 - log(`HTTP /api/ping -> ${JSON.stringify(json)}`);
53 - } catch (err) {
54 - statusEl.textContent = 'HTTP: Error';
55 - log(`HTTP error: ${err}`);
56 - }
57 - });
58 -
59 - document.getElementById('testCookiesBtn').addEventListener('click', async () => {
60 - statusEl.textContent = 'Cookies: Testing...';
61 - try {
62 - const res = await fetch('/api/test-cookies');
63 - const json = await res.json();
64 - statusEl.textContent = 'Cookies: OK';
65 - log(`HTTP /api/test-cookies -> ${JSON.stringify(json)}`);
66 -
67 - // Log response headers
68 - log(`Response headers:`);
69 - for (const [key, value] of res.headers.entries()) {
70 - log(` ${key}: ${value}`);
71 - }
72 -
73 - // Log current cookies
74 - log(`Current document.cookie: ${document.cookie || '(empty)'}`);
75 - } catch (err) {
76 - statusEl.textContent = 'Cookies: Error';
77 - log(`HTTP error: ${err}`);
78 - }
79 - });
80 -
81 - wsConnectBtn.addEventListener('click', () => {
82 - if (ws && ws.readyState === WebSocket.OPEN) {
83 - ws.close();
84 - return;
85 - }
86 -
87 - const protocol = window.location.protocol === 'https:' ? 'wss:' : 'ws:';
88 - const basePath = location.pathname.endsWith('/') ? location.pathname : (location.pathname + '/');
89 - const url = protocol + '//' + window.location.host + basePath + 'ws';
90 -
91 - statusEl.textContent = 'WS: Connecting...';
92 - log(`WS connecting to ${url}`);
93 -
94 - ws = new WebSocket(url);
95 -
96 - ws.onopen = () => {
97 - statusEl.textContent = 'WS: Connected';
98 - log('WS connected');
99 - wsSendBtn.disabled = false;
100 - wsConnectBtn.textContent = 'WS Disconnect';
101 - };
102 -
103 - ws.onclose = () => {
104 - statusEl.textContent = 'WS: Disconnected';
105 - log('WS disconnected');
106 - wsSendBtn.disabled = true;
107 - wsConnectBtn.textContent = 'WS Connect';
108 - };
109 -
110 - ws.onerror = (err) => {
111 - log(`WS error: ${err.message || err}`);
112 - };
113 -
114 - ws.onmessage = (event) => {
115 - log(`WS recv: ${event.data}`);
116 - };
117 - });
118 -
119 - wsSendBtn.addEventListener('click', () => {
120 - if (!ws || ws.readyState !== WebSocket.OPEN) {
121 - log('WS send skipped: not connected');
122 - return;
123 - }
124 - const msg = 'hello';
125 - ws.send(msg);
126 - log(`WS send: ${msg}`);
127 - });
128 - </script>
10 + <main class="shell">
11 + <header>
12 + <p class="eyebrow">Portal Demo</p>
13 + <h1>Relay-backed demo application</h1>
14 + <p>Use this page to verify static assets, cookies, fetches, and websocket echo.</p>
15 + </header>
16 +
17 + <section class="card">
18 + <h2>HTTP</h2>
19 + <button id="ping">Fetch /api/ping</button>
20 + <button id="cookies">Set cookies</button>
21 + <pre id="output">Ready.</pre>
22 + </section>
23 +
24 + <section class="card">
25 + <h2>WebSocket</h2>
26 + <input id="ws-input" value="hello portal">
27 + <button id="ws-send">Send</button>
28 + <pre id="ws-output">Disconnected.</pre>
29 + </section>
30 + </main>
31 +
32 + <script>
33 + const output = document.getElementById('output');
34 + const wsOutput = document.getElementById('ws-output');
35 +
36 + document.getElementById('ping').addEventListener('click', async () => {
37 + const res = await fetch('/api/ping');
38 + output.textContent = JSON.stringify(await res.json(), null, 2);
39 + });
40 +
41 + document.getElementById('cookies').addEventListener('click', async () => {
42 + const res = await fetch('/api/test-cookies');
43 + output.textContent = JSON.stringify(await res.json(), null, 2);
44 + });
45 +
46 + const proto = location.protocol === 'https:' ? 'wss://' : 'ws://';
47 + const ws = new WebSocket(proto + location.host + '/ws');
48 + ws.onopen = () => { wsOutput.textContent = 'Connected.'; };
49 + ws.onmessage = (event) => { wsOutput.textContent = event.data; };
50 + ws.onclose = () => { wsOutput.textContent = 'Disconnected.'; };
51 +
52 + document.getElementById('ws-send').addEventListener('click', () => {
53 + ws.send(document.getElementById('ws-input').value);
54 + });
55 + </script>
56 </body>
130 -
57 </html>
cmd/demo-app/static/style.css
+31 -70
@@ -1,85 +1,46 @@
1 -* {
2 - margin: 0;
3 - padding: 0;
4 - box-sizing: border-box;
1 +:root {
2 + color-scheme: light;
3 + font-family: "Georgia", serif;
4 + background: #f6f1e8;
5 + color: #1f1a17;
6 }
7
8 body {
8 - font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, sans-serif;
9 - background: #f5f5f5;
10 - min-height: 100vh;
11 - display: flex;
12 - align-items: center;
13 - justify-content: center;
14 - padding: 16px;
15 -}
16 -
17 -.container {
18 - width: 100%;
19 - max-width: 520px;
20 - background: #ffffff;
21 - border-radius: 8px;
22 - box-shadow: 0 4px 16px rgba(0, 0, 0, 0.08);
23 - padding: 16px;
24 - display: flex;
25 - flex-direction: column;
9 + margin: 0;
10 }
11
28 -h1 {
29 - font-size: 18px;
30 - margin-bottom: 12px;
31 - color: #333333;
32 - text-align: center;
12 +.shell {
13 + max-width: 960px;
14 + margin: 0 auto;
15 + padding: 48px 20px 80px;
16 }
17
35 -.toolbar {
36 - display: flex;
37 - flex-wrap: wrap;
38 - gap: 8px;
39 - margin-bottom: 12px;
40 - justify-content: center;
18 +.eyebrow {
19 + text-transform: uppercase;
20 + letter-spacing: 0.18em;
21 + font-size: 12px;
22 + color: #8f5e3b;
23 }
24
43 -.toolbar button {
44 - padding: 6px 12px;
45 - font-size: 13px;
46 - border-radius: 4px;
47 - border: 1px solid #d0d7de;
48 - background: #ffffff;
49 - cursor: pointer;
50 - transition: background 0.15s ease, border-color 0.15s ease;
25 +.card {
26 + background: #fff;
27 + border: 1px solid #dccdb8;
28 + padding: 20px;
29 + margin-top: 18px;
30 + box-shadow: 0 12px 28px rgba(0, 0, 0, 0.06);
31 }
32
53 -.toolbar button:hover {
54 - background: #f3f4f6;
55 - border-color: #c3ccd6;
33 +button,
34 +input {
35 + font: inherit;
36 + margin-right: 8px;
37 + margin-bottom: 8px;
38 + padding: 10px 12px;
39 }
40
58 -.canvas-wrapper {
59 - border-radius: 4px;
60 - border: 1px solid #e1e4e8;
61 - background: #111827;
62 - padding: 8px;
63 - height: 220px;
41 +pre {
42 + background: #201a17;
43 + color: #f7efe4;
44 + padding: 14px;
45 overflow: auto;
46 }
66 -
67 -.log {
68 - width: 100%;
69 - height: 100%;
70 - border: none;
71 - background: transparent;
72 - color: #e5e7eb;
73 - font-family: ui-monospace, SFMono-Regular, Menlo, Monaco, Consolas,
74 - "Liberation Mono", "Courier New", monospace;
75 - font-size: 12px;
76 - white-space: pre-wrap;
77 - word-break: break-all;
78 -}
79 -
80 -.status {
81 - margin-top: 10px;
82 - font-size: 13px;
83 - color: #4b5563;
84 - text-align: right;
85 -}
cmd/portal-tunnel/README.md
+18 -75
@@ -1,91 +1,34 @@
1 # Portal-tunnel
2
3 -Portal-tunnel is a tunneling tool that connects your local services to a [portal network](github.com/gosuda/portal), without any additional configuration.
3 +Portal-tunnel connects a local service to a Portal relay with the legacy CLI shape restored on top of the new core.
4
5 ## Usage
6
7 -You can run the tunnel using command-line flags or a configuration file.
8 -
9 -### Run binary
10 -
7 ```bash
12 -./bin/portal-tunnel --host localhost:8080 \
13 - --relay https://portal.gosuda.org,https://portal.thumbgo.kr,https://portal.iwanhae.kr \
14 - --name <service> \
8 +./portal-tunnel --host localhost:8080 \
9 + --relay https://portal.example.com \
10 + --name myapp \
11 --description "Service description" \
12 --tags tag1,tag2 \
17 - --thumbnail https://example.com/thumb.png
13 + --thumbnail https://example.com/thumb.png \
14 + --owner "Portal Operator"
15 ```
16
20 -### Transport Model
21 -
22 -Portal tunnel always runs in TLS reverse-connect mode:
23 -
24 -- Reverse admission requires HTTPS relay endpoints.
25 -- Tunnel-side TLS uses keyless signing with auto-discovered signer materials.
26 -- Traffic is proxied from tunnel to local `--host` over TCP.
27 -- Public access is `https://<service>.<portal-root-host>/`.
28 -
29 -### Control-Plane Admission
30 -
31 -Portal tunnel authenticates control-plane operations with lease token headers.
32 -
33 -- No client certificate setup is required for `/sdk/*` requests.
34 -- The relay enforces token and policy checks before accepting reverse connections.
35 -
17 ## Flags
18
19 ```text
39 -Usage:
40 - portal-tunnel [OPTIONS] [ARGUMENTS]
41 -
42 -Options:
43 - --relay Portal relay server API URLs (comma-separated, https only) [default: https://localhost:4017] [env: RELAYS]
44 - --host Target host to proxy to (host:port or URL) [env: APP_HOST]
45 - --name Service name [env: APP_NAME]
46 - --description Service description metadata [env: APP_DESCRIPTION]
47 - --tags Service tags metadata (comma-separated) [env: APP_TAGS]
48 - --thumbnail Service thumbnail URL metadata [env: APP_THUMBNAIL]
49 - --owner Service owner metadata [env: APP_OWNER]
50 - --hide Hide service from discovery (metadata) [env: APP_HIDE]
51 - -h, --help Print this help message and exit
20 +--relay Portal relay server API URLs (comma-separated, https only) [env: RELAYS]
21 +--host Target host to proxy to (host:port or URL) [env: APP_HOST]
22 +--name Service name [env: APP_NAME]
23 +--description Service description metadata [env: APP_DESCRIPTION]
24 +--tags Service tags metadata (comma-separated) [env: APP_TAGS]
25 +--thumbnail Service thumbnail URL metadata [env: APP_THUMBNAIL]
26 +--owner Service owner metadata [env: APP_OWNER]
27 +--hide Hide service from discovery [env: APP_HIDE]
28 ```
29
54 -## Examples
55 -
56 -### Quick Start (HTTP)
30 +## Notes
31
58 -```bash
59 -# macOS/Linux
60 -curl -fsSL https://portal.example.com/tunnel | APP_HOST=localhost:3000 APP_NAME=myapp sh
61 -
62 -# Windows PowerShell
63 -$env:APP_HOST="localhost:3000"; $env:APP_NAME="myapp"; irm https://portal.example.com/tunnel | iex
64 -```
65 -
66 -Installer integrity policy:
67 -
68 -- The installer downloads `BIN_URL` and `BIN_URL.sha256`.
69 -- SHA256 verification is mandatory and fail-closed.
70 -- Missing, malformed, or mismatched checksums abort startup with a remediation hint.
71 -
72 -### Production
73 -
74 -```bash
75 -export RELAYS=https://portal.example.com
76 -export APP_HOST=localhost:3000
77 -export APP_NAME=myapp
78 -
79 -./bin/portal-tunnel
80 -```
81 -
82 -When the local service is unreachable, the tunnel returns an HTTP 503 "Service Unavailable" page to the browser.
83 -
84 -### Multiple Relays (High Availability)
85 -
86 -```bash
87 -./bin/portal-tunnel \
88 - --host localhost:3000 \
89 - --name myapp \
90 - --relay https://portal1.example.com,https://portal2.example.com
91 -```
32 +- The current runtime accepts multiple relay URLs but uses the first one.
33 +- If tenant TLS material is not provided internally, the SDK generates a self-signed certificate for the registered hostnames.
34 +- When the local service is unreachable, the tunnel returns an HTTP 503 page.
cmd/portal-tunnel/main.go
+87 -79
@@ -6,6 +6,7 @@ import (
6 "flag"
7 "fmt"
8 "io"
9 + "log"
10 "net"
11 "os"
12 "os/signal"
@@ -14,11 +15,8 @@ import (
15 "syscall"
16 "time"
17
17 - "github.com/rs/zerolog"
18 - "github.com/rs/zerolog/log"
19 -
18 + "gosuda.org/portal/portal"
19 "gosuda.org/portal/sdk"
21 - "gosuda.org/portal/types"
20 )
21
22 var (
@@ -33,8 +31,6 @@ var (
31 )
32
33 func main() {
36 - log.Logger = log.Output(zerolog.ConsoleWriter{Out: os.Stdout, TimeFormat: time.RFC3339})
37 -
34 defaultRelayURLs := os.Getenv("RELAYS")
35 if defaultRelayURLs == "" {
36 defaultRelayURLs = "https://localhost:4017"
@@ -43,18 +39,15 @@ func main() {
39 flag.StringVar(&flagRelayURLs, "relay", defaultRelayURLs, "Portal relay server API URLs (comma-separated, https only) [env: RELAYS]")
40 flag.StringVar(&flagHost, "host", os.Getenv("APP_HOST"), "Target host to proxy to (host:port or URL) [env: APP_HOST]")
41 flag.StringVar(&flagName, "name", os.Getenv("APP_NAME"), "Service name [env: APP_NAME]")
46 -
42 flag.StringVar(&flagDesc, "description", os.Getenv("APP_DESCRIPTION"), "Service description metadata [env: APP_DESCRIPTION]")
43 flag.StringVar(&flagTags, "tags", os.Getenv("APP_TAGS"), "Service tags metadata (comma-separated) [env: APP_TAGS]")
44 flag.StringVar(&flagThumbnail, "thumbnail", os.Getenv("APP_THUMBNAIL"), "Service thumbnail URL metadata [env: APP_THUMBNAIL]")
45 flag.StringVar(&flagOwner, "owner", os.Getenv("APP_OWNER"), "Service owner metadata [env: APP_OWNER]")
51 -
52 - defaultHide := os.Getenv("APP_HIDE") == "true"
53 - flag.BoolVar(&flagHide, "hide", defaultHide, "Hide service from discovery (metadata) [env: APP_HIDE]")
46 + flag.BoolVar(&flagHide, "hide", os.Getenv("APP_HIDE") == "true", "Hide service from discovery (metadata) [env: APP_HIDE]")
47 flag.Parse()
48
49 if err := runTunnel(); err != nil {
57 - log.Error().Err(err).Msg("Exited with error")
50 + log.Printf("Exited with error: %v", err)
51 os.Exit(1)
52 }
53 }
@@ -63,34 +56,39 @@ func runTunnel() error {
56 ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
57 defer stop()
58
66 - relayURLs := types.ParseURLs(flagRelayURLs)
59 + relayURLs := parseURLs(flagRelayURLs)
60 if len(relayURLs) == 0 {
61 return errors.New("no relay URLs provided")
62 }
70 - relayURLs, err := normalizeRelayURLsForReverseConnect(relayURLs)
63 + relayURL, err := normalizeRelayURLsForReverseConnect(relayURLs)
64 if err != nil {
65 return err
66 }
67
75 - log.Info().Msgf("Local service is reachable at %s", flagHost)
76 - log.Info().Msg("Starting Portal Tunnel...")
77 - log.Info().Msgf(" Local: %s", flagHost)
78 - log.Info().Msgf(" Relays: %s", strings.Join(relayURLs, ", "))
68 + log.Printf("Local service is reachable at %s", flagHost)
69 + log.Printf("Starting Portal Tunnel...")
70 + log.Printf(" Local: %s", flagHost)
71 + log.Printf(" Relays: %s", strings.Join(relayURLs, ", "))
72 + if len(relayURLs) > 1 {
73 + log.Printf(" Note: current runtime uses the first relay URL only: %s", relayURL)
74 + }
75
80 - opts := []sdk.ClientOption{sdk.WithBootstrapServers(relayURLs)}
81 - sdkClient, err := sdk.NewClient(opts...)
76 + sdkClient, err := sdk.NewClient(sdk.ClientConfig{RelayURL: relayURL})
77 if err != nil {
78 return fmt.Errorf("service %s: failed to create client: %w", flagName, err)
79 }
85 -
86 - listener, err := sdkClient.Listen(
87 - flagName,
88 - types.WithDescription(flagDesc),
89 - types.WithTags(types.ParseURLs(flagTags)),
90 - types.WithOwner(flagOwner),
91 - types.WithThumbnail(flagThumbnail),
92 - types.WithHide(flagHide),
93 - )
80 + defer sdkClient.Close()
81 +
82 + listener, err := sdkClient.Listen(ctx, sdk.ListenRequest{
83 + Name: flagName,
84 + Metadata: sdk.LeaseMetadata{
85 + Description: flagDesc,
86 + Tags: parseURLs(flagTags),
87 + Owner: flagOwner,
88 + Thumbnail: flagThumbnail,
89 + Hide: flagHide,
90 + },
91 + })
92 if err != nil {
93 return fmt.Errorf("service %s: failed to register service: %w", flagName, err)
94 }
@@ -101,13 +99,14 @@ func runTunnel() error {
99 _ = listener.Close()
100 }()
101
104 - log.Info().Msg("")
105 - log.Info().Msg("Access via:")
106 - log.Info().Msgf("- Relay: %s", relayURLs[0])
107 - if leaseAware, ok := listener.(interface{ LeaseID() string }); ok {
108 - log.Info().Msgf("- Lease ID: %s", leaseAware.LeaseID())
102 + log.Printf("")
103 + log.Printf("Access via:")
104 + log.Printf("- Relay: %s", relayURL)
105 + log.Printf("- Lease ID: %s", listener.LeaseID())
106 + for _, publicURL := range listener.PublicURLs() {
107 + log.Printf("- Public: %s", publicURL)
108 }
110 - log.Info().Str("service", flagName).Msg("")
109 + log.Printf("")
110
111 connCount := 0
112 var connWG sync.WaitGroup
@@ -116,7 +115,7 @@ loop:
115 for {
116 select {
117 case <-ctx.Done():
119 - log.Info().Msg("[tunnel] shutting down...")
118 + log.Printf("[tunnel] shutting down...")
119 break loop
120 default:
121 }
@@ -124,28 +123,28 @@ loop:
123 relayConn, err := listener.Accept()
124 if err != nil {
125 if errors.Is(err, net.ErrClosed) {
127 - log.Info().Msg("[tunnel] listener closed")
126 + log.Printf("[tunnel] listener closed")
127 break loop
128 }
129 select {
130 case <-ctx.Done():
131 break loop
132 default:
134 - log.Error().Err(err).Msg("Failed to accept connection")
133 + log.Printf("Failed to accept connection: %v", err)
134 continue
135 }
136 }
137
138 connCount++
140 - log.Info().Msgf("→ [#%d] New connection from %s", connCount, relayConn.RemoteAddr())
139 + log.Printf("[#%d] New connection from %s", connCount, relayConn.RemoteAddr())
140
141 connWG.Add(1)
142 go func(relayConn net.Conn) {
143 defer connWG.Done()
144 if err := proxyConnection(ctx, flagHost, relayConn); err != nil {
146 - log.Error().Str("proxy", "TLS→TCP").Err(err).Msg("Proxy error")
145 + log.Printf("Proxy error: %v", err)
146 }
148 - log.Info().Str("proxy", "TLS→TCP").Msg("Connection closed")
147 + log.Printf("Connection closed")
148 }(relayConn)
149 }
150
@@ -158,23 +157,15 @@ loop:
157 select {
158 case <-done:
159 case <-time.After(5 * time.Second):
161 - log.Warn().Msg("[tunnel] shutdown timeout, some connections still active")
160 + log.Printf("[tunnel] shutdown timeout, some connections still active")
161 }
162
164 - log.Info().Msg("[tunnel] shutdown complete")
163 + log.Printf("[tunnel] shutdown complete")
164 return nil
165 }
166
168 -func normalizeRelayURLsForReverseConnect(relayURLs []string) ([]string, error) {
169 - normalized := make([]string, 0, len(relayURLs))
170 - for _, relayURL := range relayURLs {
171 - normalizedURL, err := types.NormalizeRelayAPIURL(relayURL)
172 - if err != nil {
173 - return nil, fmt.Errorf("invalid relay URL %q: %w", relayURL, err)
174 - }
175 - normalized = append(normalized, normalizedURL)
176 - }
177 - return normalized, nil
167 +func normalizeRelayURLsForReverseConnect(relayURLs []string) (string, error) {
168 + return portal.NormalizeRelayURL(relayURLs[0])
169 }
170
171 var bufferPool = sync.Pool{
@@ -187,7 +178,7 @@ var bufferPool = sync.Pool{
178 func proxyConnection(ctx context.Context, localAddr string, relayConn net.Conn) error {
179 defer relayConn.Close()
180
190 - targetAddr, err := types.NormalizeTargetAddr(localAddr)
181 + targetAddr, err := normalizeTargetAddr(localAddr)
182 if err != nil {
183 return fmt.Errorf("invalid --host value %q: %w", localAddr, err)
184 }
@@ -195,28 +186,18 @@ func proxyConnection(ctx context.Context, localAddr string, relayConn net.Conn)
186 dialer := &net.Dialer{Timeout: 5 * time.Second}
187 localConn, err := dialer.DialContext(ctx, "tcp", targetAddr)
188 if err != nil {
198 - log.Debug().
199 - Str("addr", targetAddr).
200 - Err(err).
201 - Msg("Local service unavailable")
189 return writeEmptyHTTPResponse(relayConn)
190 }
191 defer localConn.Close()
192
206 - log.Info().Str("addr", targetAddr).Msg("Connected to local service")
207 -
193 errCh := make(chan error, 2)
194 stopCh := make(chan struct{})
195
196 go func() {
197 select {
198 case <-ctx.Done():
214 - if closeErr := relayConn.Close(); closeErr != nil {
215 - log.Debug().Err(closeErr).Msg("failed to close relay connection on shutdown")
216 - }
217 - if closeErr := localConn.Close(); closeErr != nil {
218 - log.Debug().Err(closeErr).Msg("failed to close local connection on shutdown")
219 - }
199 + _ = relayConn.Close()
200 + _ = localConn.Close()
201 case <-stopCh:
202 }
203 }()
@@ -225,13 +206,8 @@ func proxyConnection(ctx context.Context, localAddr string, relayConn net.Conn)
206 bufPtr := bufferPool.Get().(*[]byte)
207 defer bufferPool.Put(bufPtr)
208 _, err := io.CopyBuffer(localConn, relayConn, *bufPtr)
228 - if err != nil {
229 - log.Debug().Err(err).Msg("relay->local copy ended")
230 - }
209 if tcpConn, ok := localConn.(*net.TCPConn); ok {
232 - if closeErr := tcpConn.CloseWrite(); closeErr != nil {
233 - log.Debug().Err(closeErr).Msg("failed to close local write side")
234 - }
210 + _ = tcpConn.CloseWrite()
211 }
212 errCh <- err
213 }()
@@ -240,13 +216,7 @@ func proxyConnection(ctx context.Context, localAddr string, relayConn net.Conn)
216 bufPtr := bufferPool.Get().(*[]byte)
217 defer bufferPool.Put(bufPtr)
218 _, err := io.CopyBuffer(relayConn, localConn, *bufPtr)
243 - if err != nil {
244 - log.Debug().Err(err).Msg("local->relay copy ended")
245 - }
246 - // TLS does not support half-close; full close unblocks the relay→local goroutine.
247 - if closeErr := relayConn.Close(); closeErr != nil {
248 - log.Debug().Err(closeErr).Msg("failed to close relay conn after local->relay copy")
249 - }
219 + _ = relayConn.Close()
220 errCh <- err
221 }()
222
@@ -266,7 +236,7 @@ func writeEmptyHTTPResponse(conn net.Conn) error {
236 <html>
237 <head><title>Service Unavailable</title></head>
238 <body style="font-family:sans-serif;text-align:center;padding:50px;">
269 -<h1>🔌 Service Unavailable</h1>
239 +<h1>Service Unavailable</h1>
240 <p>The local service is not currently running.</p>
241 <p>Please start your local application and refresh this page.</p>
242 </body>
@@ -279,3 +249,41 @@ func writeEmptyHTTPResponse(conn net.Conn) error {
249 _, err := conn.Write([]byte(response))
250 return err
251 }
252 +
253 +func parseURLs(raw string) []string {
254 + if strings.TrimSpace(raw) == "" {
255 + return nil
256 + }
257 + parts := strings.Split(raw, ",")
258 + out := make([]string, 0, len(parts))
259 + for _, part := range parts {
260 + part = strings.TrimSpace(part)
261 + if part != "" {
262 + out = append(out, part)
263 + }
264 + }
265 + return out
266 +}
267 +
268 +func normalizeTargetAddr(raw string) (string, error) {
269 + raw = strings.TrimSpace(raw)
270 + if raw == "" {
271 + return "", errors.New("target address is required")
272 + }
273 + if strings.Contains(raw, "://") {
274 + if strings.HasPrefix(strings.ToLower(raw), "http://") {
275 + raw = strings.TrimPrefix(raw, "http://")
276 + }
277 + if strings.HasPrefix(strings.ToLower(raw), "https://") {
278 + raw = strings.TrimPrefix(raw, "https://")
279 + }
280 + raw = strings.TrimSuffix(raw, "/")
281 + }
282 + if _, _, err := net.SplitHostPort(raw); err == nil {
283 + return raw, nil
284 + }
285 + if strings.Count(raw, ":") == 0 {
286 + return net.JoinHostPort(raw, "80"), nil
287 + }
288 + return "", fmt.Errorf("invalid target address %q", raw)
289 +}
cmd/relay-server/admin.go
+47 -75
@@ -1,103 +1,75 @@
1 package main
2
3 import (
4 + "encoding/base64"
5 + "fmt"
6 "net/http"
7 "strings"
8
9 "gosuda.org/portal/portal"
8 - portaladmin "gosuda.org/portal/portal/admin"
9 - "gosuda.org/portal/portal/policy"
10 )
11
12 -// Admin is a thin adapter between relay-server wiring and portal/admin handlers.
12 type Admin struct {
14 - service *portaladmin.Service
15 - handler *portaladmin.Handler
13 + secret string
14 + trustProxy bool
15 + frontend *Frontend
16 + server *portal.Server
17 }
18
18 -func NewAdmin(frontend *Frontend, authManager *policy.Authenticator, portalURL string, trustProxy bool) *Admin {
19 - service := portaladmin.NewService(authManager)
20 - normalizedPortalURL := strings.TrimSpace(portalURL)
21 - admin := &Admin{
22 - service: service,
19 +func NewAdmin(secret string, trustProxy bool, frontend *Frontend) *Admin {
20 + return &Admin{
21 + secret: strings.TrimSpace(secret),
22 + trustProxy: trustProxy,
23 + frontend: frontend,
24 }
24 -
25 - serveStatic := func(w http.ResponseWriter, r *http.Request, appPath string, serv *portal.RelayServer) {
26 - if frontend == nil {
27 - http.NotFound(w, r)
28 - return
29 - }
30 - frontend.ServeAppStatic(w, r, appPath, serv)
31 - }
32 -
33 - admin.handler = portaladmin.NewHandler(portaladmin.HandlerConfig{
34 - Service: service,
35 - TrustProxy: trustProxy,
36 - ServeAppStatic: serveStatic,
37 - ListLeases: func(serv *portal.RelayServer) any {
38 - return convertLeaseEntriesToRows(serv, admin, true, normalizedPortalURL)
39 - },
40 - DecodeLeaseID: decodeLeaseID,
41 - IsSecureRequest: isSecureRequestWithPolicy,
42 - WriteAPIData: writeAPIData,
43 - WriteAPIOK: writeAPIOK,
44 - WriteAPIError: writeAPIError,
45 - WriteAPIErrorWithData: writeAPIErrorWithData,
46 - })
47 -
48 - return admin
25 }
26
51 -// GetApproveManager exposes the approval manager.
52 -func (a *Admin) GetApproveManager() *policy.Approver {
53 - if a == nil || a.service == nil {
54 - return nil
55 - }
56 - return a.service.GetApproveManager()
27 +func (a *Admin) Bind(server *portal.Server) {
28 + a.server = server
29 }
30
59 -// GetBPSManager exposes the BPS manager.
60 -func (a *Admin) GetBPSManager() *policy.RateLimiter {
61 - if a == nil || a.service == nil {
62 - return nil
31 +func (a *Admin) HandleAdminRequest(w http.ResponseWriter, r *http.Request) {
32 + if !a.authorize(r) {
33 + w.Header().Set("WWW-Authenticate", `Bearer realm="portal-admin"`)
34 + http.Error(w, "unauthorized", http.StatusUnauthorized)
35 + return
36 }
64 - return a.service.GetBPSManager()
65 -}
37
67 -// GetIPManager exposes the IP manager.
68 -func (a *Admin) GetIPManager() *policy.IPFilter {
69 - if a == nil || a.service == nil {
70 - return nil
38 + switch strings.TrimSuffix(r.URL.Path, "/") {
39 + case "/admin":
40 + a.handleAdminIndex(w)
41 + case "/admin/leases":
42 + writeJSON(w, http.StatusOK, convertLeaseEntriesToRows(a.server, true, a.frontend.portalURL))
43 + default:
44 + http.NotFound(w, r)
45 }
72 - return a.service.GetIPManager()
46 }
47
75 -func (a *Admin) SetSettingsPath(path string) {
76 - if a == nil || a.service == nil {
77 - return
78 - }
79 - a.service.SetSettingsPath(path)
48 +func (a *Admin) handleAdminIndex(w http.ResponseWriter) {
49 + rows := convertLeaseEntriesToRows(a.server, true, a.frontend.portalURL)
50 + w.Header().Set("Content-Type", "text/html; charset=utf-8")
51 + _, _ = fmt.Fprintf(w, `<!doctype html><html><body><h1>Portal Admin</h1><p>%d leases</p><p><a href="/admin/leases">JSON lease list</a></p></body></html>`, len(rows))
52 }
53
82 -func (a *Admin) SaveSettings(serv *portal.RelayServer) {
83 - if a == nil || a.service == nil {
84 - return
54 +func (a *Admin) authorize(r *http.Request) bool {
55 + if a.secret == "" {
56 + return true
57 }
86 - a.service.SaveSettings(serv)
87 -}
88 -
89 -func (a *Admin) LoadSettings(serv *portal.RelayServer) {
90 - if a == nil || a.service == nil {
91 - return
58 + if subtleValueMatch(strings.TrimSpace(r.URL.Query().Get("key")), a.secret) {
59 + return true
60 }
93 - a.service.LoadSettings(serv)
94 -}
95 -
96 -// HandleAdminRequest routes /admin/* requests.
97 -func (a *Admin) HandleAdminRequest(w http.ResponseWriter, r *http.Request, serv *portal.RelayServer) {
98 - if a == nil || a.handler == nil {
99 - writeAPIError(w, http.StatusInternalServerError, "admin_handler_unavailable", "admin handler unavailable")
100 - return
61 + auth := strings.TrimSpace(r.Header.Get("Authorization"))
62 + if strings.HasPrefix(strings.ToLower(auth), "bearer ") && subtleValueMatch(strings.TrimSpace(auth[7:]), a.secret) {
63 + return true
64 + }
65 + if strings.HasPrefix(strings.ToLower(auth), "basic ") {
66 + raw, err := base64.StdEncoding.DecodeString(strings.TrimSpace(auth[6:]))
67 + if err == nil {
68 + parts := strings.SplitN(string(raw), ":", 2)
69 + if len(parts) == 2 && subtleValueMatch(parts[1], a.secret) {
70 + return true
71 + }
72 + }
73 }
102 - a.handler.HandleAdminRequest(w, r, serv)
74 + return false
75 }
cmd/relay-server/frontend.go
+88 -124
@@ -1,6 +1,7 @@
1 package main
2
3 import (
4 + "embed"
5 "encoding/json"
6 "html"
7 "io/fs"
@@ -9,8 +10,6 @@ import (
10 "strings"
11 "sync"
12
12 - "github.com/rs/zerolog/log"
13 -
13 "gosuda.org/portal/portal"
14 )
15
@@ -19,12 +18,13 @@ type readDirFileFS interface {
18 fs.ReadDirFS
19 }
20
22 -// Frontend handles serving embedded frontend assets and SSR.
21 +//go:embed dist/*
22 +var embeddedDistFS embed.FS
23 +
24 type Frontend struct {
24 - distFS readDirFileFS
25 - admin *Admin
26 - // portalURL is injected runtime config used for OG metadata and SSR links.
25 + distFS readDirFileFS
26 portalURL string
27 + server *portal.Server
28
29 cachedPortalHTML []byte
30 cachedPortalHTMLOnce sync.Once
@@ -32,74 +32,115 @@ type Frontend struct {
32
33 func NewFrontend(portalURL string) *Frontend {
34 return &Frontend{
35 - distFS: distFS,
35 + distFS: embeddedDistFS,
36 portalURL: strings.TrimSpace(portalURL),
37 }
38 }
39
40 -// SetAdmin attaches an Admin instance. Frontend methods tolerate nil admin.
41 -func (f *Frontend) SetAdmin(admin *Admin) {
42 - f.admin = admin
40 +func (f *Frontend) Bind(server *portal.Server) {
41 + f.server = server
42 }
43
45 -func (f *Frontend) initPortalHTMLCache() error {
46 - var err error
47 - f.cachedPortalHTML, err = f.distFS.ReadFile("dist/app/portal.html")
48 - return err
44 +func (f *Frontend) ServeAsset(w http.ResponseWriter, r *http.Request, assetPath, contentType string) {
45 + assetPath, ok := cleanFrontendPath(assetPath)
46 + if !ok {
47 + http.NotFound(w, r)
48 + return
49 + }
50 +
51 + fullPath := path.Join("dist", "app", assetPath)
52 + data, err := f.distFS.ReadFile(fullPath)
53 + if err != nil {
54 + http.NotFound(w, r)
55 + return
56 + }
57 + if contentType == "" {
58 + contentType = getContentType(path.Ext(assetPath))
59 + }
60 + if contentType != "" {
61 + w.Header().Set("Content-Type", contentType)
62 + }
63 + w.Header().Set("Cache-Control", "public, max-age=3600")
64 + w.WriteHeader(http.StatusOK)
65 + _, _ = w.Write(data)
66 }
67
51 -func (f *Frontend) ServeAsset(mux *http.ServeMux, route, assetPath, contentType string) {
52 - mux.HandleFunc(route, func(w http.ResponseWriter, r *http.Request) {
53 - fullPath := path.Join("dist", "app", assetPath)
54 - b, err := f.distFS.ReadFile(fullPath)
55 - if err != nil {
68 +func (f *Frontend) ServeAppStatic(w http.ResponseWriter, r *http.Request, appPath string) {
69 + appPath, ok := cleanFrontendPath(appPath)
70 + if !ok {
71 + http.NotFound(w, r)
72 + return
73 + }
74 + if appPath == "" {
75 + f.servePortalHTMLWithSSR(w)
76 + return
77 + }
78 +
79 + fullPath := path.Join("dist", "app", appPath)
80 + data, err := f.distFS.ReadFile(fullPath)
81 + if err != nil {
82 + if path.Ext(appPath) != "" {
83 http.NotFound(w, r)
84 return
85 }
59 - if contentType != "" {
60 - w.Header().Set("Content-Type", contentType)
61 - }
62 - w.WriteHeader(http.StatusOK)
63 - _, _ = w.Write(b)
64 - })
86 + f.servePortalHTMLWithSSR(w)
87 + return
88 + }
89 +
90 + if contentType := getContentType(path.Ext(appPath)); contentType != "" {
91 + w.Header().Set("Content-Type", contentType)
92 + }
93 + w.Header().Set("Cache-Control", "public, max-age=3600")
94 + w.WriteHeader(http.StatusOK)
95 + _, _ = w.Write(data)
96 }
97
67 -// servePortalHTMLWithSSR serves portal.html with SSR data injection.
68 -func (f *Frontend) servePortalHTMLWithSSR(w http.ResponseWriter, r *http.Request, serv *portal.RelayServer) {
69 - setCORSHeaders(w)
98 +func cleanFrontendPath(raw string) (string, bool) {
99 + raw = strings.TrimSpace(raw)
100 + if raw == "" {
101 + return "", true
102 + }
103
71 - // Initialize cache on first use
104 + cleaned := strings.TrimPrefix(path.Clean("/"+raw), "/")
105 + if cleaned == "." || cleaned == "" {
106 + return "", true
107 + }
108 + if cleaned == ".." || strings.HasPrefix(cleaned, "../") {
109 + return "", false
110 + }
111 + return cleaned, true
112 +}
113 +
114 +func (f *Frontend) servePortalHTMLWithSSR(w http.ResponseWriter) {
115 f.cachedPortalHTMLOnce.Do(func() {
73 - if err := f.initPortalHTMLCache(); err != nil {
74 - log.Error().Err(err).Msg("Failed to cache portal.html")
75 - }
116 + f.cachedPortalHTML, _ = f.distFS.ReadFile("dist/app/portal.html")
117 })
118
78 - if f.cachedPortalHTML == nil {
79 - http.NotFound(w, r)
119 + if len(f.cachedPortalHTML) == 0 {
120 + http.NotFound(w, nil)
121 return
122 }
123
83 - // Inject SSR data into cached template
84 - injectedHTML := f.injectServerData(string(f.cachedPortalHTML), serv)
124 + htmlContent := string(f.cachedPortalHTML)
125 + htmlContent = f.injectServerData(htmlContent)
126 + htmlContent = f.injectOGMetadata(htmlContent, "", "", "")
127
86 - // Inject OG metadata (defaults for main app)
87 - injectedHTML = f.injectOGMetadata(injectedHTML, "", "", "")
88 -
89 - // Set headers
128 w.Header().Set("Content-Type", "text/html; charset=utf-8")
129 w.Header().Set("Cache-Control", "no-cache, must-revalidate")
92 -
93 - // Send response
130 w.WriteHeader(http.StatusOK)
95 - if _, err := w.Write([]byte(injectedHTML)); err != nil {
96 - log.Debug().Err(err).Msg("failed to write portal HTML response")
97 - }
131 + _, _ = w.Write([]byte(htmlContent))
132 +}
133
99 - log.Debug().Msg("Served portal.html with SSR data")
134 +func (f *Frontend) injectServerData(htmlContent string) string {
135 + rows := convertLeaseEntriesToRows(f.server, false, f.portalURL)
136 + jsonData, err := json.Marshal(rows)
137 + if err != nil {
138 + jsonData = []byte("[]")
139 + }
140 + ssrScript := `<script id="__SSR_DATA__" type="application/json">` + string(jsonData) + `</script>`
141 + return strings.Replace(htmlContent, "</head>", ssrScript+"\n</head>", 1)
142 }
143
102 -// injectOGMetadata replaces OG placeholders with actual values.
144 func (f *Frontend) injectOGMetadata(htmlContent, title, description, imageURL string) string {
145 if title == "" {
146 title = "Portal Proxy Gateway"
@@ -108,7 +149,6 @@ func (f *Frontend) injectOGMetadata(htmlContent, title, description, imageURL st
149 description = "Transform your local services into web-accessible endpoints. Instant access from anywhere."
150 }
151 if imageURL == "" {
111 - // Use absolute URL if possible
152 base := strings.TrimSuffix(f.portalURL, "/")
153 if !strings.HasPrefix(base, "http") {
154 base = "https://" + base
@@ -121,81 +161,5 @@ func (f *Frontend) injectOGMetadata(htmlContent, title, description, imageURL st
161 "[%OG_DESCRIPTION%]", html.EscapeString(description),
162 "[%OG_IMAGE_URL%]", html.EscapeString(imageURL),
163 )
124 -
164 return replacer.Replace(htmlContent)
165 }
127 -
128 -// injectServerData injects server data into HTML for SSR.
129 -func (f *Frontend) injectServerData(htmlContent string, serv *portal.RelayServer) string {
130 - // Get server data from lease manager
131 - rows := []leaseRow{}
132 - if f.admin != nil {
133 - rows = convertLeaseEntriesToRows(serv, f.admin, false, f.portalURL)
134 - }
135 -
136 - // Marshal to JSON
137 - jsonData, err := json.Marshal(rows)
138 - if err != nil {
139 - log.Error().Err(err).Msg("Failed to marshal server data for SSR")
140 - jsonData = []byte("[]")
141 - }
142 -
143 - // Create SSR script tag
144 - ssrScript := `<script id="__SSR_DATA__" type="application/json">` + string(jsonData) + `</script>`
145 -
146 - // Inject before </head> tag
147 - injected := strings.Replace(htmlContent, "</head>", ssrScript+"\n</head>", 1)
148 -
149 - log.Debug().
150 - Int("rows", len(rows)).
151 - Int("jsonSize", len(jsonData)).
152 - Msg("Injected SSR data into HTML")
153 -
154 - return injected
155 -}
156 -
157 -// ServeAppStatic serves static files for app UI (React app) from embedded FS.
158 -// Falls back to portal.html with SSR when path is root or file not found.
159 -func (f *Frontend) ServeAppStatic(w http.ResponseWriter, r *http.Request, appPath string, serv *portal.RelayServer) {
160 - // Prevent directory traversal
161 - if strings.Contains(appPath, "..") {
162 - http.Error(w, "Invalid path", http.StatusBadRequest)
163 - return
164 - }
165 -
166 - setCORSHeaders(w)
167 -
168 - // If path is empty or "/", serve portal.html with SSR
169 - if appPath == "" || appPath == "/" {
170 - f.servePortalHTMLWithSSR(w, r, serv)
171 - return
172 - }
173 -
174 - // Try to read from embedded FS
175 - fullPath := path.Join("dist", "app", appPath)
176 - data, err := f.distFS.ReadFile(fullPath)
177 - if err != nil {
178 - // File not found - fallback to portal.html with SSR for SPA routing
179 - log.Debug().Err(err).Str("path", appPath).Msg("app static file not found, falling back to SSR")
180 - f.servePortalHTMLWithSSR(w, r, serv)
181 - return
182 - }
183 -
184 - // Set content type based on extension
185 - ext := path.Ext(appPath)
186 - contentType := getContentType(ext)
187 - if contentType != "" {
188 - w.Header().Set("Content-Type", contentType)
189 - }
190 -
191 - w.Header().Set("Cache-Control", "public, max-age=3600")
192 - w.WriteHeader(http.StatusOK)
193 - if _, err := w.Write(data); err != nil {
194 - log.Debug().Err(err).Str("path", appPath).Msg("failed to write app static response")
195 - }
196 -
197 - log.Debug().
198 - Str("path", appPath).
199 - Int("size", len(data)).
200 - Msg("served app static file")
201 -}
cmd/relay-server/http_helpers.go
+42 -99
@@ -1,18 +1,20 @@
1 package main
2
3 import (
4 - "encoding/base64"
4 + "crypto/subtle"
5 "encoding/json"
6 + "html"
7 + "mime"
8 "net/http"
9 "strings"
8 -
9 - "github.com/rs/zerolog/log"
10 -
11 - "gosuda.org/portal/portal/keyless"
12 - "gosuda.org/portal/portal/policy"
13 - "gosuda.org/portal/types"
10 )
11
12 +func writeJSON(w http.ResponseWriter, status int, data any) {
13 + w.Header().Set("Content-Type", "application/json")
14 + w.WriteHeader(status)
15 + _ = json.NewEncoder(w).Encode(data)
16 +}
17 +
18 func isSecureRequestWithPolicy(r *http.Request, trustProxyHeaders bool) bool {
19 if r == nil {
20 return false
@@ -20,115 +22,56 @@ func isSecureRequestWithPolicy(r *http.Request, trustProxyHeaders bool) bool {
22 if r.TLS != nil {
23 return true
24 }
23 - if !trustProxyHeaders || !policy.IsTrustedProxyRemoteAddr(r.RemoteAddr) {
25 + if !trustProxyHeaders {
26 return false
27 }
26 - if hasForwardedToken(r.Header.Get("X-Forwarded-Proto"), "https") {
27 - return true
28 - }
29 - return hasForwardedToken(r.Header.Get("X-Forwarded-Ssl"), "on")
28 + return strings.EqualFold(strings.TrimSpace(r.Header.Get("X-Forwarded-Proto")), "https")
29 }
30
32 -func hasForwardedToken(raw, target string) bool {
33 - for token := range strings.SplitSeq(raw, ",") {
34 - if strings.EqualFold(strings.TrimSpace(token), target) {
35 - return true
36 - }
37 - }
38 - return false
31 +func hasPathPrefix(path, prefix string) bool {
32 + return strings.HasPrefix(strings.TrimSpace(path), prefix)
33 }
34
41 -func isWebSocketUpgrade(req *http.Request) bool {
42 - if req == nil {
35 +func trimPathPrefix(path, prefix string) string {
36 + return strings.TrimPrefix(strings.TrimSpace(path), prefix)
37 +}
38 +
39 +func subtleValueMatch(left, right string) bool {
40 + if left == "" || right == "" {
41 return false
42 }
45 - return hasForwardedToken(req.Header.Get("Upgrade"), "websocket")
43 + return subtle.ConstantTimeCompare([]byte(left), []byte(right)) == 1
44 +}
45 +
46 +func escapeHTML(value string) string {
47 + return html.EscapeString(value)
48 }
49
48 -// getContentType returns the MIME type for a file extension.
50 func getContentType(ext string) string {
50 - switch ext {
51 - case ".html":
52 - return "text/html; charset=utf-8"
53 - case ".js":
54 - return "application/javascript"
55 - case ".json":
56 - return "application/json"
57 - case ".wasm":
58 - return "application/wasm"
51 + ext = strings.TrimSpace(ext)
52 + if ext == "" {
53 + return ""
54 + }
55 + if contentType := mime.TypeByExtension(ext); contentType != "" {
56 + return contentType
57 + }
58 +
59 + switch strings.ToLower(ext) {
60 + case ".js", ".mjs":
61 + return "text/javascript; charset=utf-8"
62 case ".css":
60 - return "text/css"
61 - case ".mp4":
62 - return "video/mp4"
63 + return "text/css; charset=utf-8"
64 case ".svg":
65 return "image/svg+xml"
65 - case ".png":
66 - return "image/png"
66 case ".ico":
67 return "image/x-icon"
68 + case ".jpg", ".jpeg":
69 + return "image/jpeg"
70 + case ".png":
71 + return "image/png"
72 + case ".json", ".webmanifest":
73 + return "application/json; charset=utf-8"
74 default:
75 return ""
76 }
77 }
73 -
74 -// setCORSHeaders sets permissive CORS headers for GET/OPTIONS and common headers.
75 -func setCORSHeaders(w http.ResponseWriter) {
76 - w.Header().Set("Access-Control-Allow-Origin", "*")
77 - w.Header().Set("Access-Control-Allow-Methods", "GET, OPTIONS")
78 - w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Accept, Accept-Encoding")
79 -}
80 -
81 -func writeAPIData(w http.ResponseWriter, status int, data any) {
82 - w.Header().Set("Content-Type", "application/json")
83 - w.WriteHeader(status)
84 - if err := json.NewEncoder(w).Encode(types.APIEnvelope{
85 - OK: true,
86 - Data: data,
87 - }); err != nil {
88 - log.Error().Err(err).Msg("[HTTP] Failed to encode API success response")
89 - }
90 -}
91 -
92 -func writeAPIOK(w http.ResponseWriter, status int) {
93 - w.Header().Set("Content-Type", "application/json")
94 - w.WriteHeader(status)
95 - if err := json.NewEncoder(w).Encode(types.APIEnvelope{OK: true}); err != nil {
96 - log.Error().Err(err).Msg("[HTTP] Failed to encode API success response")
97 - }
98 -}
99 -
100 -func writeAPIError(w http.ResponseWriter, status int, code, message string) {
101 - writeAPIErrorWithData(w, status, code, message, nil)
102 -}
103 -
104 -func writeAPIErrorWithData(w http.ResponseWriter, status int, code, message string, data any) {
105 - w.Header().Set("Content-Type", "application/json")
106 - w.WriteHeader(status)
107 - if err := json.NewEncoder(w).Encode(types.APIEnvelope{
108 - OK: false,
109 - Data: data,
110 - Error: &types.APIError{
111 - Code: code,
112 - Message: message,
113 - },
114 - }); err != nil {
115 - log.Error().Err(err).Msg("[HTTP] Failed to encode API error response")
116 - }
117 -}
118 -
119 -func writeSignError(w http.ResponseWriter, status int, message string) {
120 - w.Header().Set("Content-Type", "application/json")
121 - w.WriteHeader(status)
122 - _ = json.NewEncoder(w).Encode(keyless.ErrorResponse{Error: message})
123 -}
124 -
125 -func decodeLeaseID(encoded string) (string, bool) {
126 - idBytes, err := base64.URLEncoding.DecodeString(encoded)
127 - if err != nil {
128 - idBytes, err = base64.RawURLEncoding.DecodeString(encoded)
129 - if err != nil {
130 - return "", false
131 - }
132 - }
133 - return string(idBytes), true
134 -}
cmd/relay-server/lease_rows.go
+45 -183
@@ -3,43 +3,54 @@ package main
3 import (
4 "encoding/json"
5 "fmt"
6 + "strings"
7 "time"
8
8 - "github.com/rs/zerolog/log"
9 -
9 "gosuda.org/portal/portal"
11 - "gosuda.org/portal/portal/policy"
12 - "gosuda.org/portal/types"
13 -)
14 -
15 -const (
16 - leaseConnectedWindow = 15 * time.Second
17 - staleLeaseHideWindow = 3 * time.Minute
10 )
11
20 -// leaseRow represents a lease entry for display in admin UI and frontend.
12 type leaseRow struct {
22 - TTL string
23 - Metadata string
24 - Kind string
25 - IP string
26 - DNS string
27 - LastSeen string
28 - LastSeenISO string
29 - FirstSeenISO string
30 - Name string
31 - Peer string
32 - Link string
33 - BPS int64
34 - Hide bool
35 - StaleRed bool
36 - IsApproved bool
37 - IsDenied bool
38 - Connected bool
39 - IsIPBanned bool
13 + TTL string `json:"ttl"`
14 + Metadata string `json:"metadata"`
15 + Kind string `json:"kind"`
16 + DNS string `json:"dns"`
17 + Name string `json:"name"`
18 + Peer string `json:"peer"`
19 + Link string `json:"link"`
20 + Hide bool `json:"hide"`
21 + Connected bool `json:"connected"`
22 +}
23 +
24 +func convertLeaseEntriesToRows(serv *portal.Server, includeHidden bool, portalURL string) []leaseRow {
25 + if serv == nil {
26 + return nil
27 + }
28 + snapshots := serv.ListLeases()
29 + rows := make([]leaseRow, 0, len(snapshots))
30 + for _, snapshot := range snapshots {
31 + if !includeHidden && snapshot.Metadata.Hide {
32 + continue
33 + }
34 + metadataJSON, _ := json.Marshal(snapshot.Metadata)
35 + host := ""
36 + if len(snapshot.Hostnames) > 0 {
37 + host = snapshot.Hostnames[0]
38 + }
39 + rows = append(rows, leaseRow{
40 + TTL: formatDuration(time.Until(snapshot.ExpiresAt)),
41 + Metadata: string(metadataJSON),
42 + Kind: "https",
43 + DNS: host,
44 + Name: snapshot.Name,
45 + Peer: snapshot.ID,
46 + Link: leaseLink(host, portalURL),
47 + Hide: snapshot.Metadata.Hide,
48 + Connected: snapshot.Ready > 0,
49 + })
50 + }
51 + return rows
52 }
53
42 -// formatDuration formats a duration for TTL display.
54 func formatDuration(d time.Duration) string {
55 if d <= 0 {
56 return ""
@@ -53,159 +64,10 @@ func formatDuration(d time.Duration) string {
64 return fmt.Sprintf("%.0fs", d.Seconds())
65 }
66
56 -// formatLastSeen formats a duration since last seen.
57 -func formatLastSeen(d time.Duration) string {
58 - if d >= time.Hour {
59 - h := int(d / time.Hour)
60 - m := int((d % time.Hour) / time.Minute)
61 - if m > 0 {
62 - return fmt.Sprintf("%dh %dm", h, m)
63 - }
64 - return fmt.Sprintf("%dh", h)
65 - }
66 - if d >= time.Minute {
67 - m := int(d / time.Minute)
68 - s := int((d % time.Minute) / time.Second)
69 - if s > 0 {
70 - return fmt.Sprintf("%dm %ds", m, s)
71 - }
72 - return fmt.Sprintf("%dm", m)
73 - }
74 - return fmt.Sprintf("%ds", int(d/time.Second))
75 -}
76 -
77 -func isLeaseConnected(since time.Duration) bool {
78 - return since < leaseConnectedWindow
79 -}
80 -
81 -// fromLeaseEntry populates the leaseRow from a LeaseEntry with common fields.
82 -func (r *leaseRow) fromLeaseEntry(entry *types.LeaseEntry, admin *Admin, portalURL string) {
83 - lease := entry.Lease
84 - identityID := lease.ID
85 - since := max(time.Since(entry.LastSeen), 0)
86 - connected := isLeaseConnected(since)
87 -
88 - name := lease.Name
89 - if name == "" {
90 - name = "(unnamed)"
91 - }
92 -
93 - kind := "http"
94 - if lease.TLS {
95 - kind = "https"
96 - }
97 -
98 - dnsLabel := identityID
99 - if len(dnsLabel) > 8 {
100 - dnsLabel = dnsLabel[:8] + "..."
101 - }
102 -
103 - var bps int64
104 - if admin != nil {
105 - if bpsMgr := admin.GetBPSManager(); bpsMgr != nil {
106 - bps = bpsMgr.GetBPSLimit(identityID)
107 - }
108 - }
109 -
110 - metadata := lease.Metadata
111 - metadataStr := ""
112 - if b, err := json.Marshal(metadata); err == nil {
113 - metadataStr = string(b)
114 - } else {
115 - log.Warn().Err(err).Str("lease_id", identityID).Msg("[leaseRow] Failed to marshal lease metadata")
116 - }
117 -
118 - r.Peer = identityID
119 - r.Name = name
120 - r.Kind = kind
121 - r.Connected = connected
122 - r.DNS = dnsLabel
123 - r.LastSeen = formatLastSeen(since)
124 - r.LastSeenISO = entry.LastSeen.UTC().Format(time.RFC3339)
125 - r.FirstSeenISO = entry.FirstSeen.UTC().Format(time.RFC3339)
126 - r.TTL = formatDuration(time.Until(entry.Lease.Expires))
127 - linkLabel := identityID
128 - if normalized, ok := types.NormalizeServiceName(lease.Name); ok {
129 - linkLabel = normalized
130 - } else if normalized, ok := types.NormalizeServiceName(identityID); ok {
131 - linkLabel = normalized
132 - }
133 -
134 - publicHost := types.PortalRootHost(portalURL)
135 - if publicHost == "" {
136 - publicHost = types.PortalHostPort(portalURL)
137 - }
138 - if linkLabel != "" && publicHost != "" {
139 - r.Link = fmt.Sprintf("//%s.%s/", linkLabel, publicHost)
140 - } else {
141 - r.Link = ""
142 - }
143 - r.StaleRed = !connected && since >= leaseConnectedWindow
144 - r.Hide = metadata.Hide
145 - r.Metadata = metadataStr
146 - r.BPS = bps
147 -
148 - if admin != nil {
149 - if approveMgr := admin.GetApproveManager(); approveMgr != nil {
150 - r.IsApproved = approveMgr.GetApprovalMode() == policy.ModeAuto || approveMgr.IsLeaseApproved(identityID)
151 - r.IsDenied = approveMgr.IsLeaseDenied(identityID)
152 - }
153 -
154 - if ipMgr := admin.GetIPManager(); ipMgr != nil {
155 - r.IP = ipMgr.GetLeaseIP(identityID)
156 - if r.IP != "" {
157 - r.IsIPBanned = ipMgr.IsIPBanned(r.IP)
158 - }
159 - }
160 - }
161 -}
162 -
163 -// convertLeaseEntriesToRows converts LeaseEntry data to leaseRow format.
164 -// If forAdmin is true, includes all leases with admin-only fields.
165 -// If forAdmin is false, filters out banned, unapproved, hidden, and stale leases.
166 -func convertLeaseEntriesToRows(serv *portal.RelayServer, admin *Admin, forAdmin bool, portalURL string) []leaseRow {
167 - leaseEntries := serv.GetLeaseManager().GetAllLeaseEntries()
168 - rows := []leaseRow{}
169 - now := time.Now()
170 -
171 - bannedList := serv.GetLeaseManager().GetBannedLeases()
172 - bannedMap := make(map[string]struct{}, len(bannedList))
173 - for _, b := range bannedList {
174 - bannedMap[b] = struct{}{}
175 - }
176 -
177 - for _, entry := range leaseEntries {
178 - if now.After(entry.Lease.Expires) {
179 - continue
180 - }
181 -
182 - identityID := entry.Lease.ID
183 - metadata := entry.Lease.Metadata
184 -
185 - if !forAdmin {
186 - if _, banned := bannedMap[identityID]; banned {
187 - continue
188 - }
189 - if admin != nil {
190 - approveManager := admin.GetApproveManager()
191 - if approveManager.GetApprovalMode() == policy.ModeManual && !approveManager.IsLeaseApproved(identityID) {
192 - continue
193 - }
194 - }
195 - if metadata.Hide {
196 - continue
197 - }
198 - since := max(now.Sub(entry.LastSeen), 0)
199 - connected := isLeaseConnected(since)
200 - if !connected && since >= staleLeaseHideWindow {
201 - continue
202 - }
203 - }
204 -
205 - var row leaseRow
206 - row.fromLeaseEntry(entry, admin, portalURL)
207 - rows = append(rows, row)
67 +func leaseLink(host, portalURL string) string {
68 + host = strings.TrimSpace(host)
69 + if host == "" {
70 + return ""
71 }
209 -
210 - return rows
72 + return "https://" + host + "/"
73 }
cmd/relay-server/main.go
+25 -157
@@ -1,30 +1,18 @@
1 package main
2
3 import (
4 - "context"
4 "flag"
5 "fmt"
7 - "net"
6 + "log"
7 "os"
9 - "os/signal"
8 "strings"
11 - "syscall"
12 - "time"
13 -
14 - "github.com/rs/zerolog"
15 - "github.com/rs/zerolog/log"
16 -
17 - "gosuda.org/portal/portal"
18 - "gosuda.org/portal/portal/policy"
19 - "gosuda.org/portal/portal/sni"
20 - "gosuda.org/portal/types"
9 )
10
11 const (
12 defaultAPIPort = 4017
13 defaultSNIPort = 443
14 defaultPortalURL = "https://localhost:4017"
27 - defaultKeylessDir = "/etc/portal/keyless"
15 + defaultKeylessDir = ".portal-certs"
16 )
17
18 type relayServerConfig struct {
@@ -40,8 +28,6 @@ type relayServerConfig struct {
28 }
29
30 func main() {
43 - log.Logger = log.Output(zerolog.ConsoleWriter{Out: os.Stdout, TimeFormat: time.RFC3339})
44 -
31 cfg := relayServerConfig{}
32
33 portalURL := strings.TrimSuffix(trimmedEnv("PORTAL_URL"), "/")
@@ -50,9 +36,9 @@ func main() {
36 }
37 bootstrapsCSV := trimmedEnv("BOOTSTRAP_URIS")
38 if bootstrapsCSV == "" {
53 - bootstrapsCSV = types.DefaultBootstrapFrom(portalURL)
39 + bootstrapsCSV = portalURL
40 }
55 - sniPort := types.ParsePortNumber(os.Getenv("SNI_PORT"), defaultSNIPort)
41 + sniPort := parsePortNumber(os.Getenv("SNI_PORT"), defaultSNIPort)
42 keylessDir := trimmedEnv("KEYLESS_DIR")
43 if keylessDir == "" {
44 keylessDir = defaultKeylessDir
@@ -73,150 +59,20 @@ func main() {
59 flag.StringVar(&cfg.CloudflareToken, "cloudflare-token", cloudflareToken, "Cloudflare DNS API token (Zone:Read + DNS:Edit) (env: CLOUDFLARE_TOKEN)")
60 flag.Parse()
61
76 - cfg.Bootstraps = types.ParseURLs(bootstrapsCSV)
77 - parsedTrustedProxyCIDRs, err := policy.ParseTrustedProxyCIDRs(cfg.TrustedProxyCIDRs)
78 - if err != nil {
79 - log.Fatal().Err(err).Msg("parse trusted proxy CIDRs")
80 - }
81 - policy.SetTrustedProxyCIDRs(parsedTrustedProxyCIDRs)
82 - if err := runServer(cfg); err != nil {
83 - log.Fatal().Err(err).Msg("execute root command")
84 - }
85 -}
86 -
87 -func runServer(cfg relayServerConfig) error {
88 - ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
89 - defer stop()
90 - sniListenAddr := fmt.Sprintf(":%d", cfg.SNIPort)
91 -
92 - log.Info().
93 - Str("portal_base_url", cfg.PortalURL).
94 - Bool("trust_proxy_headers", cfg.TrustProxyHeaders).
95 - Str("trusted_proxy_cidrs", cfg.TrustedProxyCIDRs).
96 - Strs("bootstrap_uris", cfg.Bootstraps).
97 - Msg("[server] frontend configuration")
98 -
99 - rootHost := types.PortalRootHost(cfg.PortalURL)
100 - apiUpstreamAddr := types.LoopbackForwardAddr(fmt.Sprintf(":%d", cfg.AdminPort))
101 - serv, err := portal.NewRelayServer(ctx, cfg.Bootstraps, sniListenAddr, rootHost, cfg.KeylessDir, cfg.CloudflareToken)
102 - if err != nil {
103 - return fmt.Errorf("create relay server: %w", err)
104 - }
105 -
106 - frontend := NewFrontend(cfg.PortalURL)
107 - authManager := policy.NewAuthenticator(cfg.AdminSecretKey)
108 - admin := NewAdmin(frontend, authManager, cfg.PortalURL, cfg.TrustProxyHeaders)
109 - frontend.SetAdmin(admin)
110 -
111 - // Load persisted admin settings (ban list, BPS limits, IP bans)
112 - admin.LoadSettings(serv)
113 - ipMgr := admin.GetIPManager()
114 - serv.GetLeaseManager().SetOnLeaseDeleted(func(leaseID string) {
115 - leaseID = strings.TrimSpace(leaseID)
116 - if leaseID == "" {
117 - return
118 - }
119 -
120 - serv.GetReverseHub().DropLease(leaseID)
121 - if sniRouter := serv.GetSNIRouter(); sniRouter != nil {
122 - sniRouter.UnregisterRouteByLeaseID(leaseID)
123 - }
124 - if ipMgr != nil {
125 - ipMgr.RemoveLeaseIP(leaseID)
126 - }
127 - })
128 - if ipMgr != nil {
129 - serv.GetReverseHub().SetIPBanChecker(func(ip string) bool {
130 - return policy.IsIPBannedByPolicy(ipMgr, ip)
131 - })
132 - serv.GetReverseHub().SetOnAccepted(func(leaseID, ip string) {
133 - if strings.TrimSpace(leaseID) == "" || strings.TrimSpace(ip) == "" {
134 - return
135 - }
136 - ipMgr.RegisterLeaseIP(leaseID, ip)
137 - })
138 - }
139 -
140 - // Set up SNI connection callback to route to tunnel backends
141 - serv.GetSNIRouter().SetConnectionCallback(func(clientConn net.Conn, route *sni.Route) {
142 - leaseID := ""
143 - if route != nil {
144 - leaseID = strings.TrimSpace(route.LeaseID)
145 - }
146 - if leaseID == "" {
147 - logSNIRouteWarning(route, nil, "[SNI] Missing lease id in route; dropping connection")
148 - closeSNIClientConn(clientConn, route)
149 - return
150 - }
151 -
152 - if _, ok := serv.GetLeaseManager().GetLeaseByID(leaseID); !ok {
153 - logSNIRouteWarning(route, nil, "[SNI] Lease not active; dropping connection and unregistering route")
154 - serv.GetSNIRouter().UnregisterRouteByLeaseID(leaseID)
155 - closeSNIClientConn(clientConn, route)
156 - return
157 - }
158 -
159 - // Get BPS manager for rate limiting
160 - bpsManager := admin.GetBPSManager()
161 -
162 - reverseConn, err := serv.GetReverseHub().AcquireForTLS(leaseID, portal.TLSAcquireWait)
163 - if err != nil {
164 - logSNIRouteWarning(route, err, "[SNI] Reverse tunnel unavailable")
165 - closeSNIClientConn(clientConn, route)
166 - return
167 - }
168 -
169 - // SNI path is reverse-only (NAT-friendly): relay never dials app directly.
170 - policy.EstablishRelayWithBPS(clientConn, reverseConn.Conn, leaseID, bpsManager)
171 - reverseConn.Close()
172 - })
173 -
174 - serv.ConfigurePortalRootFallback(rootHost, apiUpstreamAddr)
175 -
176 - if err := serv.Start(); err != nil {
177 - return fmt.Errorf("start relay server: %w", err)
62 + cfg.Bootstraps = parseURLs(bootstrapsCSV)
63 + if len(cfg.Bootstraps) == 0 {
64 + cfg.Bootstraps = []string{cfg.PortalURL}
65 }
179 -
180 - apiServ := serveAPI(fmt.Sprintf(":%d", cfg.AdminPort), serv, admin, frontend, cfg, stop)
181 -
182 - <-ctx.Done()
183 - log.Info().Msg("[server] shutting down...")
184 -
185 - // Stop relay first: drain idle reverse conns + close active SNI conns
186 - // so that all HTTP handlers blocked on HandleConnect.Wait() can return.
187 - serv.Stop()
188 -
189 - shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
190 - defer cancel()
191 - if apiServ != nil {
192 - if err := apiServ.Shutdown(shutdownCtx); err != nil {
193 - log.Error().Err(err).Msg("[server] http server shutdown error")
194 - }
66 + if cfg.PortalURL == "" {
67 + cfg.PortalURL = cfg.Bootstraps[0]
68 }
69
197 - log.Info().Msg("[server] shutdown complete")
198 - return nil
199 -}
70 + log.Printf("[server] portal base url %s", cfg.PortalURL)
71 + log.Printf("[server] bootstraps %s", strings.Join(cfg.Bootstraps, ", "))
72
201 -func closeSNIClientConn(clientConn net.Conn, route *sni.Route) {
202 - if err := clientConn.Close(); err != nil {
203 - withSNIRouteFields(log.Debug().Err(err), route).Msg("[SNI] failed to close client connection")
204 - }
205 -}
206 -
207 -func logSNIRouteWarning(route *sni.Route, err error, msg string) {
208 - event := log.Warn()
209 - if err != nil {
210 - event = event.Err(err)
211 - }
212 - withSNIRouteFields(event, route).Msg(msg)
213 -}
214 -
215 -func withSNIRouteFields(event *zerolog.Event, route *sni.Route) *zerolog.Event {
216 - if route == nil {
217 - return event
73 + if err := runServer(cfg); err != nil {
74 + log.Fatalf("execute root command: %v", err)
75 }
219 - return event.Str("lease_id", route.LeaseID).Str("sni", route.SNI)
76 }
77
78 func trimmedEnv(name string) string {
@@ -227,3 +83,15 @@ func parseBoolEnv(name string) bool {
83 raw := trimmedEnv(name)
84 return strings.EqualFold(raw, "true") || raw == "1"
85 }
86 +
87 +func parsePortNumber(raw string, fallback int) int {
88 + raw = strings.TrimSpace(raw)
89 + if raw == "" {
90 + return fallback
91 + }
92 + var port int
93 + if _, err := fmt.Sscanf(raw, "%d", &port); err != nil || port < 1 || port > 65535 {
94 + return fallback
95 + }
96 + return port
97 +}
cmd/relay-server/registry.go
+17 -238
@@ -1,247 +1,26 @@
1 package main
2
3 -import (
4 - "encoding/json"
5 - "net/http"
6 - "strings"
7 -
8 - "github.com/rs/zerolog/log"
9 -
10 - "gosuda.org/portal/portal"
11 - "gosuda.org/portal/portal/policy"
12 - "gosuda.org/portal/types"
13 -)
14 -
15 -const sdkRequestBodyLimitBytes = 4 << 20 // 4 MiB
16 -
17 -// SDKRegistry handles HTTP API for client lease registration.
18 -type SDKRegistry struct {
19 - ipManager *policy.IPFilter
20 - portalURL string
21 - trustProxyHeaders bool
22 -}
23 -
24 -// HandleSDKRequest routes /sdk/* requests.
25 -func (r *SDKRegistry) HandleSDKRequest(w http.ResponseWriter, req *http.Request, serv *portal.RelayServer) {
26 - if serv == nil {
27 - writeAPIError(w, http.StatusInternalServerError, "registry_unavailable", "registry service unavailable")
28 - return
29 - }
30 -
31 - path := strings.TrimSuffix(req.URL.Path, "/")
32 - switch path {
33 - case types.PathSDKRegister:
34 - r.handleRegister(w, req, serv)
35 - case types.PathSDKUnregister:
36 - r.handleUnregister(w, req, serv)
37 - case types.PathSDKRenew:
38 - r.handleRenew(w, req, serv)
39 - case types.PathSDKDomain:
40 - r.handleDomain(w, serv)
41 - case types.PathSDKConnect:
42 - r.handleConnect(w, req, serv)
43 - default:
44 - http.NotFound(w, req)
45 - }
46 -}
47 -
48 -func (r *SDKRegistry) handleConnect(w http.ResponseWriter, req *http.Request, serv *portal.RelayServer) {
49 - if req.Method != http.MethodGet {
50 - w.Header().Set("Allow", http.MethodGet)
51 - writeAPIError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
52 - return
53 - }
54 - if isWebSocketUpgrade(req) {
55 - writeAPIError(w, http.StatusBadRequest, "unsupported_transport", "websocket transport is not supported")
56 - return
57 - }
58 -
59 - admission, ok := r.admitControlPlane(
60 - w,
61 - req,
62 - serv,
63 - req.URL.Query().Get("lease_id"),
64 - req.Header.Get(types.ReverseConnectTokenHeader),
65 - true,
66 - )
67 - if !ok {
68 - return
69 - }
70 -
71 - hijacker, ok := w.(http.Hijacker)
72 - if !ok {
73 - writeAPIError(w, http.StatusInternalServerError, "hijacker_unavailable", "server does not support connection hijacking")
74 - return
75 - }
76 - conn, rw, err := hijacker.Hijack()
77 - if err != nil {
78 - writeAPIError(w, http.StatusInternalServerError, "hijack_failed", "failed to hijack connection")
79 - return
80 - }
81 - if _, err := rw.WriteString("HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: keep-alive\r\n\r\n"); err != nil {
82 - if closeErr := conn.Close(); closeErr != nil {
83 - log.Debug().Err(closeErr).Msg("[Registry] failed to close hijacked connection after write failure")
84 - }
85 - return
86 - }
87 - if err := rw.Flush(); err != nil {
88 - if closeErr := conn.Close(); closeErr != nil {
89 - log.Debug().Err(closeErr).Msg("[Registry] failed to close hijacked connection after flush failure")
3 +import "strings"
4 +
5 +func parseURLs(raw string) []string {
6 + if strings.TrimSpace(raw) == "" {
7 + return nil
8 + }
9 + parts := strings.Split(raw, ",")
10 + out := make([]string, 0, len(parts))
11 + for _, part := range parts {
12 + part = strings.TrimSpace(part)
13 + if part != "" {
14 + out = append(out, part)
15 }
91 - return
92 - }
93 -
94 - serv.HandleRegistryConnect(conn, admission)
95 -}
96 -
97 -// handleRegister handles SDK lease registration requests.
98 -func (r *SDKRegistry) handleRegister(w http.ResponseWriter, req *http.Request, serv *portal.RelayServer) {
99 - if !r.requireMethod(w, req, http.MethodPost) {
100 - return
101 - }
102 -
103 - var registerReq types.RegisterRequest
104 - if !r.decodeRequestBody(w, req, &registerReq, "[Registry] Failed to decode registration request") {
105 - return
106 - }
107 -
108 - admission, ok := r.admitControlPlane(
109 - w,
110 - req,
111 - serv,
112 - registerReq.LeaseID,
113 - registerReq.ReverseToken,
114 - false,
115 - )
116 - if !ok {
117 - return
118 - }
119 -
120 - registerResp, apiErr := serv.RegisterLease(portal.RegistryRegisterInput{
121 - LeaseID: admission.LeaseID,
122 - ReverseToken: admission.ReverseToken,
123 - Name: registerReq.Name,
124 - Metadata: &registerReq.Metadata,
125 - TLS: registerReq.TLS,
126 - PortalURL: r.portalURL,
127 - })
128 - if !writeRegistryError(w, apiErr) {
129 - return
130 - }
131 -
132 - writeAPIData(w, http.StatusOK, registerResp)
133 -}
134 -
135 -// handleUnregister handles SDK lease unregistration requests.
136 -func (r *SDKRegistry) handleUnregister(w http.ResponseWriter, req *http.Request, serv *portal.RelayServer) {
137 - if !r.requireMethod(w, req, http.MethodPost) {
138 - return
139 - }
140 -
141 - var unregisterReq types.UnregisterRequest
142 - if !r.decodeRequestBody(w, req, &unregisterReq, "[Registry] Failed to decode unregistration request") {
143 - return
144 - }
145 -
146 - admission, ok := r.admitControlPlane(
147 - w,
148 - req,
149 - serv,
150 - unregisterReq.LeaseID,
151 - unregisterReq.ReverseToken,
152 - true,
153 - )
154 - if !ok {
155 - return
156 - }
157 -
158 - serv.UnregisterLease(admission.LeaseID)
159 - writeAPIOK(w, http.StatusOK)
160 -}
161 -
162 -// handleRenew handles SDK lease renewal requests (keepalive).
163 -func (r *SDKRegistry) handleRenew(w http.ResponseWriter, req *http.Request, serv *portal.RelayServer) {
164 - if !r.requireMethod(w, req, http.MethodPost) {
165 - return
166 - }
167 -
168 - var renewReq types.RenewRequest
169 - if !r.decodeRequestBody(w, req, &renewReq, "[Registry] Failed to decode renewal request") {
170 - return
171 - }
172 -
173 - admission, ok := r.admitControlPlane(
174 - w,
175 - req,
176 - serv,
177 - renewReq.LeaseID,
178 - renewReq.ReverseToken,
179 - true,
180 - )
181 - if !ok {
182 - return
183 - }
184 -
185 - if !writeRegistryError(w, serv.RenewLease(admission.Entry)) {
186 - return
187 - }
188 - writeAPIOK(w, http.StatusOK)
189 -}
190 -
191 -// handleDomain returns the relay's base domain for TLS certificate construction.
192 -func (r *SDKRegistry) handleDomain(w http.ResponseWriter, serv *portal.RelayServer) {
193 - domainResp, apiErr := serv.RegistryDomain()
194 - if !writeRegistryError(w, apiErr) {
195 - return
196 - }
197 - writeAPIData(w, http.StatusOK, domainResp)
198 -}
199 -
200 -func (r *SDKRegistry) admitControlPlane(
201 - w http.ResponseWriter,
202 - req *http.Request,
203 - serv *portal.RelayServer,
204 - rawLeaseID, rawToken string,
205 - requireExistingLease bool,
206 -) (portal.RegistryAdmissionResult, bool) {
207 - clientIP := policy.ExtractClientIP(req, r.trustProxyHeaders)
208 - admission, apiErr := serv.AdmitControlPlane(portal.RegistryAdmissionInput{
209 - RawLeaseID: rawLeaseID,
210 - RawReverseToken: rawToken,
211 - ClientIP: clientIP,
212 - IsClientIPBanned: policy.IsIPBannedByPolicy(r.ipManager, clientIP),
213 - RequireExisting: requireExistingLease,
214 - })
215 - if !writeRegistryError(w, apiErr) {
216 - return portal.RegistryAdmissionResult{}, false
217 - }
218 - return admission, true
219 -}
220 -
221 -func (r *SDKRegistry) requireMethod(w http.ResponseWriter, req *http.Request, method string) bool {
222 - if req.Method == method {
223 - return true
224 - }
225 -
226 - w.Header().Set("Allow", method)
227 - writeAPIError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
228 - return false
229 -}
230 -
231 -func (r *SDKRegistry) decodeRequestBody(w http.ResponseWriter, req *http.Request, dst any, logMessage string) bool {
232 - req.Body = http.MaxBytesReader(w, req.Body, sdkRequestBodyLimitBytes)
233 - if err := json.NewDecoder(req.Body).Decode(dst); err != nil {
234 - log.Error().Err(err).Msg(logMessage)
235 - writeAPIError(w, http.StatusBadRequest, "invalid_request", "invalid request body")
236 - return false
16 }
238 - return true
17 + return out
18 }
19
241 -func writeRegistryError(w http.ResponseWriter, apiErr *types.APIError) bool {
242 - if apiErr == nil {
20 +func isRelayControlPlanePath(path string) bool {
21 + switch strings.TrimSpace(path) {
22 + case "/sdk/register", "/sdk/connect", "/sdk/renew", "/sdk/unregister", "/sdk/domain":
23 return true
24 }
245 - writeAPIError(w, apiErr.StatusCode, apiErr.Code, apiErr.Message)
246 - return false
25 + return strings.HasPrefix(path, "/sdk/")
26 }
cmd/relay-server/serve.go
+105 -216
@@ -2,254 +2,143 @@ package main
2
3 import (
4 "context"
5 - "crypto/tls"
6 - "embed"
7 - "encoding/json"
8 - "errors"
5 "fmt"
6 + "log"
7 "net"
8 "net/http"
12 - "strconv"
9 + "os"
10 + "os/signal"
11 "strings"
14 - "time"
15 -
16 - "github.com/rs/zerolog/log"
12 + "syscall"
13
14 "gosuda.org/portal/portal"
19 - "gosuda.org/portal/portal/keyless"
20 - "gosuda.org/portal/portal/policy"
21 - "gosuda.org/portal/types"
15 + "gosuda.org/portal/portal/acme"
16 )
17
24 -const defaultHTTPSPort = "443"
25 -
26 -//go:embed dist/*
27 -var distFS embed.FS
18 +func runServer(cfg relayServerConfig) error {
19 + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
20 + defer stop()
21
29 -// serveAPI builds the admin/API mux and returns the server.
30 -func serveAPI(addr string, serv *portal.RelayServer, admin *Admin, frontend *Frontend, cfg relayServerConfig, cancel context.CancelFunc) *http.Server {
31 - if addr == "" {
32 - addr = ":0"
22 + if len(cfg.Bootstraps) > 0 && cfg.PortalURL == "" {
23 + cfg.PortalURL = cfg.Bootstraps[0]
24 }
34 -
35 - // Create app UI mux
36 - appMux := http.NewServeMux()
37 -
38 - // Serve favicons (ico/png/svg) from dist/app
39 - frontend.ServeAsset(appMux, "/favicon.ico", "favicon.ico", "image/x-icon")
40 - frontend.ServeAsset(appMux, "/favicon.png", "favicon.png", "image/png")
41 - frontend.ServeAsset(appMux, "/favicon.svg", "favicon.svg", "image/svg+xml")
42 -
43 - // Portal app assets (JS, CSS, etc.) - served from /app/
44 - appMux.HandleFunc(types.PathAppPrefix, func(w http.ResponseWriter, r *http.Request) {
45 - setCORSHeaders(w)
46 - if r.Method == http.MethodOptions {
47 - w.WriteHeader(http.StatusOK)
48 - return
49 - }
50 - p := strings.TrimPrefix(r.URL.Path, types.PathAppPrefix)
51 - frontend.ServeAppStatic(w, r, p, serv)
52 - })
53 -
54 - // Tunnel installer script and binaries
55 - appMux.HandleFunc(types.PathTunnelScript, func(w http.ResponseWriter, r *http.Request) {
56 - serveTunnelScript(w, r, cfg.PortalURL)
57 - })
58 - appMux.HandleFunc(types.PathTunnelBinary, func(w http.ResponseWriter, r *http.Request) {
59 - serveTunnelBinary(w, r)
25 + rootHost := portal.PortalRootHost(cfg.PortalURL)
26 + apiListenAddr := fmt.Sprintf(":%d", cfg.AdminPort)
27 + sniListenAddr := fmt.Sprintf(":%d", cfg.SNIPort)
28 +
29 + acmeManager, err := acme.NewManager(acme.Config{
30 + BaseDomain: rootHost,
31 + KeyDir: cfg.KeylessDir,
32 + CloudflareToken: cfg.CloudflareToken,
33 })
61 -
62 - // SDK registry API for /sdk/* endpoints
63 - var sdkIPManager *policy.IPFilter
64 - if admin != nil {
65 - sdkIPManager = admin.GetIPManager()
66 - }
67 - registry := &SDKRegistry{
68 - ipManager: sdkIPManager,
69 - portalURL: cfg.PortalURL,
70 - trustProxyHeaders: cfg.TrustProxyHeaders,
34 + if err != nil {
35 + return fmt.Errorf("create acme manager: %w", err)
36 }
72 - appMux.HandleFunc(types.PathSDKPrefix, func(w http.ResponseWriter, r *http.Request) {
73 - registry.HandleSDKRequest(w, r, serv)
74 - })
75 -
76 - // Keyless signer endpoint.
77 - appMux.HandleFunc(types.PathKeylessSign, func(w http.ResponseWriter, r *http.Request) {
78 - handleKeylessSign(w, r, serv.GetKeylessSigner())
79 - })
80 -
81 - // App UI index page - serve React frontend with SSR (delegates to serveAppStatic)
82 - appMux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
83 - // serveAppStatic handles both "/" and 404 fallback with SSR
84 - p := strings.TrimPrefix(r.URL.Path, "/")
85 - frontend.ServeAppStatic(w, r, p, serv)
86 - })
37
88 - appMux.HandleFunc(types.PathHealthz, func(w http.ResponseWriter, _ *http.Request) {
89 - w.WriteHeader(http.StatusOK)
90 - if _, err := w.Write([]byte("{\"status\":\"ok\"}")); err != nil {
91 - log.Debug().Err(err).Msg("[healthz] failed to write response")
92 - }
93 - })
94 -
95 - // Admin API
96 - appMux.HandleFunc(types.PathAdminPrefix+"/", func(w http.ResponseWriter, r *http.Request) {
97 - admin.HandleAdminRequest(w, r, serv)
98 - })
99 -
100 - // Create the main handler
101 - appDomain := types.DefaultAppPattern(cfg.PortalURL)
102 - handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
103 - // Handle subdomain requests
104 - if types.IsSubdomain(appDomain, r.Host) {
105 - log.Debug().
106 - Str("host", r.Host).
107 - Str("url", r.URL.String()).
108 - Msg("[server] handling subdomain request")
109 - // TLS-enabled subdomains should terminate on SNI passthrough.
110 - // Redirect only insecure requests; secure requests here would loop.
111 - if !isSecureRequestWithPolicy(r, cfg.TrustProxyHeaders) {
112 - log.Debug().Str("host", r.Host).Msg("[server] redirecting to HTTPS")
113 - redirectToHTTPS(w, r, serv.GetSNIRouter().GetAddr())
114 - return
115 - }
38 + certFile, keyFile, err := acmeManager.EnsureCertificate(ctx)
39 + if err != nil {
40 + return fmt.Errorf("ensure relay certificate: %w", err)
41 + }
42
117 - log.Warn().Str("host", r.Host).Msg("[server] tls subdomain reached admin listener without SNI route")
118 - http.Error(w, "tls-enabled subdomain must be served via SNI route", http.StatusMisdirectedRequest)
119 - return
120 - }
121 - appMux.ServeHTTP(w, r)
43 + frontend := NewFrontend(cfg.PortalURL)
44 + admin := NewAdmin(cfg.AdminSecretKey, cfg.TrustProxyHeaders, frontend)
45 +
46 + server, err := portal.NewServer(portal.ServerConfig{
47 + PortalURL: cfg.PortalURL,
48 + APIListenAddr: apiListenAddr,
49 + SNIListenAddr: sniListenAddr,
50 + RootHost: rootHost,
51 + RootFallbackAddr: loopbackAddr(apiListenAddr),
52 + APITLS: portal.TLSMaterialConfig{
53 + CertPEM: mustRead(certFile),
54 + KeyPEM: mustRead(keyFile),
55 + },
56 + APIHandlerWrapper: serveAPI(frontend, admin, cfg),
57 })
58 + if err != nil {
59 + return fmt.Errorf("create relay server: %w", err)
60 + }
61
124 - // Add security headers middleware
125 - secureHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
126 - w.Header().Set("X-Content-Type-Options", "nosniff")
127 - w.Header().Set("X-Frame-Options", "DENY")
128 - w.Header().Set("Referrer-Policy", "strict-origin-when-cross-origin")
129 - handler.ServeHTTP(w, r)
130 - })
62 + frontend.Bind(server)
63 + admin.Bind(server)
64
132 - srv := &http.Server{
133 - Addr: addr,
134 - Handler: secureHandler,
135 - ReadHeaderTimeout: 5 * time.Second,
136 - TLSNextProto: make(map[string]func(*http.Server, *tls.Conn, http.Handler)),
137 - }
138 - acmeManager := serv.GetACMEManager()
139 - rootHost := types.PortalRootHost(cfg.PortalURL)
140 - srv.TLSConfig = &tls.Config{
141 - MinVersion: tls.VersionTLS12,
142 - ClientAuth: tls.NoClientCert,
143 - GetCertificate: func(hello *tls.ClientHelloInfo) (*tls.Certificate, error) {
144 - serverName := strings.TrimSpace(strings.ToLower(hello.ServerName))
145 - if serverName != "" && !strings.EqualFold(serverName, rootHost) {
146 - return nil, fmt.Errorf("acme certificate is only served for portal root host %q", rootHost)
147 - }
148 - certFile, keyFile := acmeManager.TLSFiles()
149 - cert, err := tls.LoadX509KeyPair(certFile, keyFile)
150 - if err != nil {
151 - return nil, fmt.Errorf("load acme certificate: %w", err)
152 - }
153 - return &cert, nil
154 - },
65 + if err := server.Start(ctx); err != nil {
66 + return fmt.Errorf("start relay server: %w", err)
67 }
68 + acmeManager.Start(ctx)
69 + defer acmeManager.Stop()
70
157 - go func() {
158 - log.Info().Str("addr", addr).Msg("[server] https api enabled via ACME")
159 - err := srv.ListenAndServeTLS("", "")
160 - if err != nil && err != http.ErrServerClosed {
161 - log.Error().Err(err).Msg("[server] http error")
162 - cancel()
163 - }
164 - }()
71 + log.Printf("[server] https api enabled via ACME/self-signed on %s", loopbackAddr(server.APIAddr()))
72 + log.Printf("[server] sni router listening on %s", server.SNIAddr())
73 + log.Printf("[server] root host %s", rootHost)
74
166 - return srv
75 + return server.Wait()
76 }
77
169 -// redirectToHTTPS redirects the request to HTTPS using the configured SNI port.
170 -func redirectToHTTPS(w http.ResponseWriter, r *http.Request, sniListenAddr string) {
171 - host := strings.TrimSpace(r.Host)
172 - if h, _, err := net.SplitHostPort(host); err == nil {
173 - host = h
174 - }
175 -
176 - // Extract port from sniListenAddr (e.g., ":443", "443", "example.com:443")
177 - port := defaultHTTPSPort
178 - if raw := strings.TrimSpace(sniListenAddr); raw != "" {
179 - switch {
180 - case strings.HasPrefix(raw, ":"):
181 - port = strings.TrimPrefix(raw, ":")
182 - case strings.Count(raw, ":") == 0:
183 - port = raw
184 - default:
185 - if _, p, err := net.SplitHostPort(raw); err == nil {
186 - port = p
78 +func serveAPI(frontend *Frontend, admin *Admin, cfg relayServerConfig) func(http.Handler) http.Handler {
79 + return func(base http.Handler) http.Handler {
80 + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
81 + switch {
82 + case isRelayControlPlanePath(r.URL.Path):
83 + base.ServeHTTP(w, r)
84 + case r.URL.Path == "/healthz":
85 + base.ServeHTTP(w, r)
86 + case isFrontendRootAssetPath(r.URL.Path):
87 + frontend.ServeAsset(w, r, strings.TrimPrefix(r.URL.Path, "/"), "")
88 + case hasPathPrefix(r.URL.Path, "/assets/"):
89 + frontend.ServeAsset(w, r, strings.TrimPrefix(r.URL.Path, "/"), "")
90 + case r.URL.Path == "/" || r.URL.Path == "/app" || r.URL.Path == "/app/":
91 + frontend.ServeAppStatic(w, r, "")
92 + case hasPathPrefix(r.URL.Path, "/app/"):
93 + frontend.ServeAppStatic(w, r, trimPathPrefix(r.URL.Path, "/app/"))
94 + case r.URL.Path == "/admin" || r.URL.Path == "/admin/":
95 + admin.HandleAdminRequest(w, r)
96 + case hasPathPrefix(r.URL.Path, "/admin/"):
97 + admin.HandleAdminRequest(w, r)
98 + case r.URL.Path == "/tunnel":
99 + serveTunnelScript(w, r, cfg.PortalURL)
100 + case hasPathPrefix(r.URL.Path, "/tunnel/bin/"):
101 + serveTunnelBinary(w, r)
102 + default:
103 + base.ServeHTTP(w, r)
104 }
188 - }
189 - if n, err := strconv.Atoi(port); err != nil || n < 1 || n > 65535 {
190 - port = defaultHTTPSPort
191 - }
192 - }
193 -
194 - if port != defaultHTTPSPort {
195 - host = net.JoinHostPort(host, port)
105 + })
106 }
197 -
198 - target := "https://" + host + r.URL.Path
199 - if r.URL.RawQuery != "" {
200 - target += "?" + r.URL.RawQuery
201 - }
202 - http.Redirect(w, r, target, http.StatusPermanentRedirect)
107 }
108
205 -func handleKeylessSign(w http.ResponseWriter, r *http.Request, signer *keyless.Signer) {
206 - if signer == nil {
207 - writeSignError(w, http.StatusNotFound, "keyless signer is disabled")
208 - return
109 +func isFrontendRootAssetPath(requestPath string) bool {
110 + switch requestPath {
111 + case "/favicon.ico",
112 + "/favicon.svg",
113 + "/favicon-96x96.png",
114 + "/apple-touch-icon.png",
115 + "/web-app-manifest-192x192.png",
116 + "/web-app-manifest-512x512.png",
117 + "/portal.jpg":
118 + return true
119 + default:
120 + return false
121 }
122 +}
123
211 - if r.Method != http.MethodPost {
212 - w.Header().Set("Allow", http.MethodPost)
213 - writeSignError(w, http.StatusMethodNotAllowed, "method not allowed")
214 - return
124 +func loopbackAddr(addr string) string {
125 + host, port, err := net.SplitHostPort(addr)
126 + if err != nil {
127 + return addr
128 }
216 -
217 - if ct := r.Header.Get("Content-Type"); ct != "" && !strings.HasPrefix(ct, "application/json") {
218 - writeSignError(w, http.StatusUnsupportedMediaType, "content type must be application/json")
219 - return
129 + if host == "" || host == "0.0.0.0" || host == "::" {
130 + host = "127.0.0.1"
131 }
132 + return net.JoinHostPort(host, port)
133 +}
134
222 - r.Body = http.MaxBytesReader(w, r.Body, 1<<16)
223 - defer r.Body.Close()
224 -
225 - var req keyless.SignRequest
226 - if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
227 - writeSignError(w, http.StatusBadRequest, "invalid json body")
228 - return
135 +func mustRead(path string) []byte {
136 + if path == "" {
137 + log.Fatal("missing required PEM path")
138 }
230 -
231 - resp, err := signer.Sign(r.Context(), &req)
139 + data, err := os.ReadFile(path)
140 if err != nil {
233 - status := http.StatusInternalServerError
234 - switch {
235 - case errors.Is(err, keyless.ErrSignerDisabled):
236 - status = http.StatusNotFound
237 - case errors.Is(err, keyless.ErrInvalidArgument):
238 - status = http.StatusBadRequest
239 - case errors.Is(err, keyless.ErrPermissionDenied):
240 - status = http.StatusForbidden
241 - }
242 - msg := err.Error()
243 - if status == http.StatusInternalServerError {
244 - msg = "internal signing error"
245 - }
246 - writeSignError(w, status, msg)
247 - return
248 - }
249 -
250 - w.Header().Set("Content-Type", "application/json")
251 - if err := json.NewEncoder(w).Encode(resp); err != nil {
252 - log.Error().Err(err).Msg("[signer] failed to encode sign response")
253 - writeSignError(w, http.StatusInternalServerError, "failed to encode response")
141 + log.Fatalf("read %s: %v", path, err)
142 }
143 + return data
144 }
cmd/relay-server/tunnel.go
+31 -52
@@ -5,9 +5,9 @@ import (
5 "encoding/hex"
6 "fmt"
7 "net/http"
8 + "os"
9 + "path/filepath"
10 "strings"
9 -
10 - "github.com/rs/zerolog/log"
11 )
12
13 const tunnelScriptTemplate = `#!/usr/bin/env sh
@@ -163,108 +163,87 @@ try {
163 `
164
165 func serveTunnelScript(w http.ResponseWriter, r *http.Request, portalURL string) {
166 - setCORSHeaders(w)
166 if r.Method != http.MethodGet && r.Method != http.MethodHead {
167 w.Header().Set("Allow", http.MethodGet+", "+http.MethodHead)
168 http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
169 return
170 }
171
173 - targetOS := r.URL.Query().Get("os")
174 - var isWindows bool
175 - if targetOS != "" {
176 - isWindows = strings.EqualFold(targetOS, "windows")
177 - } else {
178 - // Fallback: check User-Agent
179 - ua := strings.ToLower(r.UserAgent())
180 - isWindows = strings.Contains(ua, "windows")
181 - }
182 -
183 - var script string
184 - var contentType string
185 - var filename string
186 -
172 + isWindows := strings.EqualFold(strings.TrimSpace(r.URL.Query().Get("os")), "windows")
173 + script := fmt.Sprintf(tunnelScriptTemplate, portalURL)
174 + contentType := "text/x-shellscript"
175 + filename := "tunnel.sh"
176 if isWindows {
177 script = fmt.Sprintf(tunnelPowerShellScriptTemplate, portalURL)
189 - contentType = "text/plain" // or application/x-powershell
178 + contentType = "text/plain; charset=utf-8"
179 filename = "tunnel.ps1"
191 - } else {
192 - script = fmt.Sprintf(tunnelScriptTemplate, portalURL)
193 - contentType = "text/x-shellscript"
194 - filename = "tunnel.sh"
180 }
181
182 w.Header().Set("Content-Type", contentType)
198 - w.Header().Set("Cache-Control", "no-cache, no-store, must-revalidate")
183 w.Header().Set("Content-Disposition", fmt.Sprintf("inline; filename=\"%s\"", filename))
200 - w.WriteHeader(http.StatusOK)
184 if r.Method == http.MethodGet {
202 - if _, err := w.Write([]byte(script)); err != nil {
203 - log.Debug().Err(err).Msg("failed to write tunnel script")
204 - }
185 + _, _ = w.Write([]byte(script))
186 }
187 }
188
189 func serveTunnelBinary(w http.ResponseWriter, r *http.Request) {
209 - setCORSHeaders(w)
190 if r.Method != http.MethodGet && r.Method != http.MethodHead {
191 w.Header().Set("Allow", http.MethodGet+", "+http.MethodHead)
192 http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
193 return
194 }
195
216 - slug := strings.TrimPrefix(r.URL.Path, "/tunnel/bin/")
217 - slug = strings.Trim(slug, "/")
196 + slug := strings.Trim(strings.TrimPrefix(r.URL.Path, "/tunnel/bin/"), "/")
197 checksumRequest := strings.HasSuffix(slug, ".sha256")
198 if checksumRequest {
199 slug = strings.TrimSuffix(slug, ".sha256")
200 }
201
223 - path, ok := tunnelBinaryAssetBySlug[slug]
224 - if !ok || strings.TrimSpace(path) == "" {
202 + path, ok := tunnelBinaryPathBySlug(slug)
203 + if !ok {
204 http.NotFound(w, r)
205 return
206 }
207
229 - data, err := distFS.ReadFile(path)
208 + data, err := os.ReadFile(path)
209 if err != nil {
231 - log.Error().Err(err).Str("path", path).Msg("failed to read embedded tunnel binary")
210 http.NotFound(w, r)
211 return
212 }
235 -
213 sum := sha256.Sum256(data)
214 checksumHex := hex.EncodeToString(sum[:])
215
216 if checksumRequest {
217 w.Header().Set("Content-Type", "text/plain; charset=utf-8")
241 - w.Header().Set("Cache-Control", "public, max-age=600")
242 - w.WriteHeader(http.StatusOK)
218 if r.Method == http.MethodGet {
244 - if _, err := fmt.Fprintf(w, "%s portal-tunnel-%s\n", checksumHex, slug); err != nil {
245 - log.Debug().Err(err).Str("slug", slug).Msg("failed to write tunnel checksum")
246 - }
219 + _, _ = fmt.Fprintf(w, "%s portal-tunnel-%s\n", checksumHex, slug)
220 }
221 return
222 }
223
224 w.Header().Set("Content-Type", "application/octet-stream")
252 - w.Header().Set("Content-Disposition", fmt.Sprintf("attachment; filename=\"portal-tunnel-%s\"", slug))
253 - w.Header().Set("Cache-Control", "public, max-age=600")
225 w.Header().Set("X-Checksum-Sha256", checksumHex)
255 - w.WriteHeader(http.StatusOK)
226 if r.Method == http.MethodGet {
257 - if _, err := w.Write(data); err != nil {
258 - log.Debug().Err(err).Str("slug", slug).Msg("failed to write tunnel binary")
227 + _, _ = w.Write(data)
228 + }
229 +}
230 +
231 +func tunnelBinaryPathBySlug(slug string) (string, bool) {
232 + candidates := []string{
233 + filepath.Join("dist", "tunnel", tunnelBinaryName(slug)),
234 + filepath.Join("bin", tunnelBinaryName(slug)),
235 + }
236 + for _, candidate := range candidates {
237 + if _, err := os.Stat(candidate); err == nil {
238 + return candidate, true
239 }
240 }
241 + return "", false
242 }
243
263 -var tunnelBinaryAssetBySlug = map[string]string{
264 - "linux-amd64": "dist/tunnel/portal-tunnel-linux-amd64",
265 - "linux-arm64": "dist/tunnel/portal-tunnel-linux-arm64",
266 - "darwin-amd64": "dist/tunnel/portal-tunnel-darwin-amd64",
267 - "darwin-arm64": "dist/tunnel/portal-tunnel-darwin-arm64",
268 - "windows-amd64": "dist/tunnel/portal-tunnel-windows-amd64.exe",
269 - "windows-arm64": "dist/tunnel/portal-tunnel-windows-arm64.exe",
244 +func tunnelBinaryName(slug string) string {
245 + if strings.HasPrefix(slug, "windows-") {
246 + return "portal-tunnel-" + slug + ".exe"
247 + }
248 + return "portal-tunnel-" + slug
249 }
docs/greenfield-raw-tcp-sni-keyless.md new
+457
@@ -0,0 +1,457 @@
1 +# Greenfield Design: Raw TCP Reverse-Connect + SNI Passthrough + Keyless TLS
2 +
3 +## Status
4 +
5 +Greenfield design. This document is intentionally not constrained by the current implementation, test suite, or ADR set. It describes the replacement system as if built from scratch.
6 +
7 +## Scope
8 +
9 +Keep only these product properties:
10 +
11 +- Raw TCP reverse-connect from backend to relay
12 +- SNI-based tenant routing
13 +- Tenant TLS passthrough end-to-end
14 +- Keyless TLS for relay-owned admin/API TLS
15 +
16 +Everything else is redesignable.
17 +
18 +## Goals
19 +
20 +- One clear state owner per lease
21 +- No fixed worker pools
22 +- No shared global reverse queue logic beyond lease lookup
23 +- Deterministic connection lifecycle
24 +- Easy to instrument and debug
25 +- Backpressure and failure behavior defined up front
26 +
27 +## Non-Goals
28 +
29 +- Backward compatibility with the current internal implementation
30 +- Multiple transport modes
31 +- Relay-side tenant TLS termination
32 +- WebSocket support
33 +- HTTP/2 support
34 +- HTTP/1 and HTTP/2 abstraction at the tunnel layer
35 +
36 +## High-Level Model
37 +
38 +The system has four runtime components:
39 +
40 +1. `ControlPlane`
41 + Validates registration, renewal, unregister, and reverse-connect admission.
42 +
43 +2. `RouteTable`
44 + Maps SNI hostnames to `lease_id`.
45 +
46 +3. `LeaseBroker`
47 + One broker per lease. Owns all idle reverse connections for that lease.
48 +
49 +4. `LeaseAgent`
50 + One agent per SDK listener. Keeps a small target number of idle reverse connections available at the relay.
51 +
52 +The key simplification is this:
53 +
54 +- Relay global state only resolves `lease_id -> LeaseBroker`
55 +- All reverse connection lifecycle for a lease is owned by that broker
56 +- SDK global state only owns `listener -> LeaseAgent`
57 +- All reverse session lifecycle for a listener is owned by that agent
58 +
59 +## Connection Roles
60 +
61 +### 1. Browser to Relay
62 +
63 +- TCP to relay SNI listener
64 +- Relay peeks ClientHello
65 +- SNI resolves to `lease_id`
66 +- Relay claims one idle reverse connection from that lease broker
67 +- Relay writes `TLSStartMarker`
68 +- Relay bridges raw TCP in both directions
69 +
70 +Relay never terminates tenant TLS.
71 +
72 +### 2. SDK to Relay
73 +
74 +- TCP + TLS to admin/API listener
75 +- HTTP `GET /sdk/connect?lease_id=...`
76 +- Reverse token in header
77 +- On success, connection becomes an idle reverse session owned by the lease broker
78 +
79 +### 3. Admin/API TLS
80 +
81 +- Relay terminates root-domain TLS only
82 +- Private key stays behind keyless signer
83 +
84 +## HTTP Version Policy
85 +
86 +Use HTTP/1.1 only everywhere HTTP exists in the system.
87 +
88 +### Relay-Owned HTTP
89 +
90 +- Admin/API listener is HTTP/1.1 only
91 +- `/sdk/connect` is HTTP/1.1 only
92 +- No HTTP/2 on relay listeners
93 +
94 +Reason:
95 +
96 +- `/sdk/connect` depends on connection hijacking semantics
97 +- HTTP/2 adds stream multiplexing with no value for this design
98 +- HTTP/2 makes connection behavior harder to reason about operationally
99 +
100 +### Tenant-Facing HTTP
101 +
102 +Tenant traffic still uses raw TCP + TLS passthrough, so the relay does not enforce HTTP version directly.
103 +
104 +Instead, the tenant-side TLS terminator used by the SDK/tunnel must advertise only:
105 +
106 +- `http/1.1`
107 +
108 +It must not advertise:
109 +
110 +- `h2`
111 +
112 +Reason:
113 +
114 +- HTTP/2 connection coalescing can reuse one TLS connection across multiple subdomains when certificate coverage overlaps
115 +- This design routes once per TCP/TLS connection using SNI
116 +- If multiple origins ride the same HTTP/2 connection, later requests can bypass per-origin SNI routing decisions
117 +- Shared wildcard certificates across tenant subdomains make that risk worse
118 +
119 +This greenfield design therefore chooses a hard rule:
120 +
121 +- SNI passthrough routing plus shared subdomain certificate coverage implies HTTP/1.1 only
122 +
123 +If future work wants tenant HTTP/2, it must first remove cross-tenant certificate overlap or redesign routing away from single-handshake SNI ownership.
124 +
125 +## Per-Lease Relay Design
126 +
127 +`LeaseBroker` is the only owner of reverse-session state for a lease.
128 +
129 +```text
130 +LeaseBroker
131 + lease_id
132 + state: active | dropped | stopped
133 + ready queue: bounded FIFO of idle reverse sessions
134 + metrics: ready_count, claimed_count, dropped_count, last_claim_at
135 +```
136 +
137 +### LeaseBroker API
138 +
139 +- `Offer(session) error`
140 +- `Claim(ctx) (*ReverseSession, error)`
141 +- `Drop()`
142 +- `Reset()`
143 +- `Stop()`
144 +
145 +### Rules
146 +
147 +- `Offer` rejects immediately if broker is dropped or stopped
148 +- `Claim` blocks until a valid idle session is available or timeout/cancel fires
149 +- `Drop` drains ready queue and closes all idle sessions
150 +- `Reset` reopens a dropped lease after successful re-registration
151 +- `Stop` is terminal and used only for process shutdown
152 +
153 +No separate global `dropped` map exists outside the broker.
154 +
155 +## Reverse Session Design
156 +
157 +`ReverseSession` is a single reverse TCP connection plus its state.
158 +
159 +```text
160 +states:
161 + connecting
162 + admitted
163 + idle
164 + claimed
165 + bridged
166 + closed
167 +```
168 +
169 +### State Transitions
170 +
171 +- SDK connects: `connecting -> admitted`
172 +- Broker accepts idle session: `admitted -> idle`
173 +- SNI path claims it: `idle -> claimed`
174 +- Relay writes `TLSStartMarker`: `claimed -> bridged`
175 +- Any close/error: `* -> closed`
176 +
177 +### Session Rules
178 +
179 +- Keepalive is allowed only in `idle`
180 +- Control marker write is allowed only in `claimed`
181 +- Session close is idempotent
182 +- Close unblocks any waiter
183 +
184 +## SDK Design
185 +
186 +`LeaseAgent` replaces fixed reverse worker pools.
187 +
188 +```text
189 +LeaseAgent
190 + lease
191 + relay_url
192 + target_ready
193 + current_sessions
194 + paused: bool
195 + state: running | paused | stopped
196 +```
197 +
198 +### LeaseAgent Behavior
199 +
200 +- Maintain `target_ready` idle reverse sessions
201 +- Default `target_ready = 1`
202 +- Optional burst mode may raise to `2` or `4` based on recent claim rate
203 +- No fixed fan-out like `16 workers`
204 +
205 +### LeaseAgent Loop
206 +
207 +1. If `current_sessions < target_ready`, open one reverse session
208 +2. Complete reverse-connect handshake
209 +3. Wait for `TLSStartMarker`
210 +4. Run backend TLS handshake locally
211 +5. Deliver accepted connection to app `Accept()`
212 +6. Decrement active session count
213 +7. Replenish one new idle session
214 +
215 +This is slot-based, not worker-based.
216 +
217 +## Protocol
218 +
219 +Keep the binary markers minimal:
220 +
221 +- `0x00`: idle keepalive
222 +- `0x02`: activate TLS passthrough
223 +
224 +No non-TLS tenant mode.
225 +
226 +### Reverse Connect Admission
227 +
228 +Admission order remains strict:
229 +
230 +1. IP policy
231 +2. Lease existence/state
232 +3. Reverse token
233 +
234 +Only admitted sessions can enter `LeaseBroker.Offer`.
235 +
236 +## API Surface
237 +
238 +Keep the existing control-plane shape, simplified internally:
239 +
240 +- `POST /sdk/register`
241 +- `GET /sdk/connect?lease_id=...`
242 +- `POST /sdk/renew`
243 +- `POST /sdk/unregister`
244 +- `GET /sdk/domain`
245 +- `POST /v1/sign`
246 +
247 +All HTTP endpoints in this list are HTTP/1.1 only.
248 +
249 +### Register
250 +
251 +- Requires `lease_id`, `name`, `reverse_token`, `tls=true`
252 +- Creates or resets lease broker
253 +- Registers route in `RouteTable`
254 +
255 +### Connect
256 +
257 +- Admits request
258 +- Hijacks TCP connection
259 +- Creates `ReverseSession`
260 +- Offers it into the lease broker
261 +- Blocks until session is claimed or closed
262 +
263 +### Renew
264 +
265 +- Extends lease TTL
266 +- Keeps broker active
267 +
268 +### Unregister
269 +
270 +- Drops broker
271 +- Removes route
272 +- Deletes lease record
273 +
274 +## Routing
275 +
276 +`RouteTable` owns only routing data:
277 +
278 +- exact match
279 +- single-label wildcard match
280 +- portal-root fallback
281 +
282 +It does not own reverse session state.
283 +
284 +On tenant SNI hit:
285 +
286 +1. Resolve `lease_id`
287 +2. Lookup broker
288 +3. `Claim(ctx)`
289 +4. Write `TLSStartMarker`
290 +5. Bridge raw TCP
291 +
292 +## Shutdown Semantics
293 +
294 +### Lease Drop
295 +
296 +- New reverse offers rejected
297 +- Ready queue drained
298 +- In-flight bridged sessions continue until close
299 +
300 +### Process Stop
301 +
302 +- Control-plane listener stops admitting new reverse sessions
303 +- SNI listener stops admitting new browser sessions
304 +- All brokers enter stopped state
305 +- Idle sessions close immediately
306 +- Bridged sessions close on listener shutdown or peer close
307 +
308 +## Failure Handling
309 +
310 +### Reverse Connect Rejection
311 +
312 +Fatal:
313 +
314 +- banned IP
315 +- missing lease
316 +- invalid token
317 +- unsupported transport
318 +
319 +Transient:
320 +
321 +- relay restart
322 +- temporary dial failure
323 +- temporary route churn
324 +
325 +### SDK Agent Rules
326 +
327 +- Fatal rejection pauses the agent
328 +- Successful renew/register clears pause
329 +- Transient failure retries with bounded backoff
330 +
331 +## Backpressure
332 +
333 +Per-lease ready queue is bounded.
334 +
335 +Recommended defaults:
336 +
337 +- `target_ready = 1`
338 +- `ready_queue_capacity = 8`
339 +- temporary burst increase to `2` or `4`
340 +
341 +If the ready queue is full:
342 +
343 +- Evict the oldest idle session
344 +- Never evict a claimed or bridged session
345 +
346 +## Observability
347 +
348 +Every reverse session gets a `conn_id`.
349 +
350 +Required structured logs:
351 +
352 +- reverse admitted
353 +- reverse offered
354 +- reverse claimed
355 +- TLS marker sent
356 +- bridge started
357 +- bridge ended
358 +- session evicted
359 +- lease dropped
360 +- lease reset
361 +- fatal agent pause
362 +
363 +Core metrics:
364 +
365 +- ready sessions per lease
366 +- claim wait duration
367 +- reverse connect failures by reason
368 +- session lifetime by state
369 +- first-claim latency after register
370 +
371 +## Package Shape
372 +
373 +Suggested greenfield package split:
374 +
375 +```text
376 +portal/controlplane
377 +portal/routing
378 +portal/broker
379 +portal/session
380 +portal/keyless
381 +sdk/agent
382 +sdk/listener
383 +```
384 +
385 +### Ownership
386 +
387 +- `controlplane`: admission and API handlers
388 +- `routing`: SNI lookup only
389 +- `broker`: lease-local ready queue and lifecycle
390 +- `session`: reverse session state machine
391 +- `sdk/agent`: maintain target idle sessions
392 +- `sdk/listener`: app-facing `net.Listener`
393 +
394 +## Minimal Relay Pseudocode
395 +
396 +```text
397 +on_sni_client(server_name):
398 + lease_id = route_table.lookup(server_name)
399 + broker = brokers.get(lease_id)
400 + session = broker.claim(ctx)
401 + session.send_tls_start()
402 + bridge(client_conn, session.conn)
403 +```
404 +
405 +```text
406 +on_sdk_connect(lease_id, token, conn):
407 + admit(lease_id, token, client_ip)
408 + session = new_reverse_session(conn)
409 + broker = brokers.get_or_create(lease_id)
410 + broker.offer(session)
411 + wait_until_claimed_or_closed(session)
412 +```
413 +
414 +## Minimal SDK Pseudocode
415 +
416 +```text
417 +run_agent():
418 + while running:
419 + if active_sessions < target_ready:
420 + spawn open_one_session()
421 + wait_for_signal()
422 +```
423 +
424 +```text
425 +open_one_session():
426 + conn = dial_relay()
427 + do_tls()
428 + do_http_connect()
429 + wait_for_tls_start_marker()
430 + tls_server_handshake()
431 + deliver_to_accept_queue()
432 +```
433 +
434 +## Why This Is Simpler
435 +
436 +- Lease-local ownership replaces mixed global and per-connection state
437 +- Slot-based replenishment replaces fixed worker pools
438 +- Reverse lifecycle becomes a state machine instead of loosely coupled goroutines
439 +- Routing, admission, and reverse session storage are separated cleanly
440 +
441 +## Recommended Implementation Order
442 +
443 +1. Build `LeaseBroker` and `ReverseSession` as isolated packages
444 +2. Build SDK `LeaseAgent` with target-ready semantics
445 +3. Replace relay reverse path behind feature flag or alternate binary
446 +4. Reconnect SNI router to broker claims
447 +5. Reconnect control-plane handlers
448 +6. Add observability before traffic testing
449 +
450 +## Cutover Strategy
451 +
452 +Best option: separate greenfield binary or branch, not incremental mutation of the current reverse path.
453 +
454 +Reason:
455 +
456 +- The main value of this design is deleting hidden lifecycle coupling
457 +- Partial migration would preserve too much of the old complexity
go.mod
+4 -7
@@ -5,20 +5,17 @@ go 1.26.0
5 require (
6 github.com/go-acme/lego/v4 v4.32.0
7 github.com/gosuda/keyless_tls v0.0.1-0.20260304212324-7733f8366abc
8 - github.com/rs/zerolog v1.34.0
9 - golang.org/x/net v0.51.0
8 + golang.org/x/net v0.50.0
9 + golang.org/x/sync v0.19.0
10 )
11
12 require (
13 github.com/cenkalti/backoff/v5 v5.0.3 // indirect
14 github.com/go-jose/go-jose/v4 v4.1.3 // indirect
15 - github.com/mattn/go-colorable v0.1.14 // indirect
16 - github.com/mattn/go-isatty v0.0.20 // indirect
15 github.com/miekg/dns v1.1.72 // indirect
16 golang.org/x/crypto v0.48.0 // indirect
19 - golang.org/x/mod v0.33.0 // indirect
20 - golang.org/x/sync v0.19.0 // indirect
17 + golang.org/x/mod v0.32.0 // indirect
18 golang.org/x/sys v0.41.0 // indirect
19 golang.org/x/text v0.34.0 // indirect
23 - golang.org/x/tools v0.42.0 // indirect
20 + golang.org/x/tools v0.41.0 // indirect
21 )
go.sum
+6 -22
@@ -1,50 +1,34 @@
1 github.com/cenkalti/backoff/v5 v5.0.3 h1:ZN+IMa753KfX5hd8vVaMixjnqRZ3y8CuJKRKj1xcsSM=
2 github.com/cenkalti/backoff/v5 v5.0.3/go.mod h1:rkhZdG3JZukswDf7f0cwqPNk4K0sa+F97BxZthm/crw=
3 -github.com/coreos/go-systemd/v22 v22.5.0/go.mod h1:Y58oyj3AT4RCenI/lSvhwexgC+NSVTIJ3seZv2GcEnc=
3 github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1VwoXQT9A3Wy9MM3WgvqSxFWenqJduM=
4 github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
5 github.com/go-acme/lego/v4 v4.32.0 h1:z7Ss7aa1noabhKj+DBzhNCO2SM96xhE3b0ucVW3x8Tc=
6 github.com/go-acme/lego/v4 v4.32.0/go.mod h1:lI2fZNdgeM/ymf9xQ9YKbgZm6MeDuf91UrohMQE4DhI=
7 github.com/go-jose/go-jose/v4 v4.1.3 h1:CVLmWDhDVRa6Mi/IgCgaopNosCaHz7zrMeF9MlZRkrs=
8 github.com/go-jose/go-jose/v4 v4.1.3/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08=
10 -github.com/godbus/dbus/v5 v5.0.4/go.mod h1:xhWf0FNVPg57R7Z0UbKHbJfkEywrmjJnf7w5xrFpKfA=
9 github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
10 github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
11 github.com/gosuda/keyless_tls v0.0.1-0.20260304212324-7733f8366abc h1:aS9LQ35x6EtrGKCmOWRj6Y9aQ2l5hP8dVva4oxB9VEg=
12 github.com/gosuda/keyless_tls v0.0.1-0.20260304212324-7733f8366abc/go.mod h1:BOhUZgiAAQzxKO3QcC4fCXgd/+lqxgIu1OyIYTqtta8=
15 -github.com/mattn/go-colorable v0.1.13/go.mod h1:7S9/ev0klgBDR4GtXTXX8a3vIGJpMovkB8vQcUbaXHg=
16 -github.com/mattn/go-colorable v0.1.14 h1:9A9LHSqF/7dyVVX6g0U9cwm9pG3kP9gSzcuIPHPsaIE=
17 -github.com/mattn/go-colorable v0.1.14/go.mod h1:6LmQG8QLFO4G5z1gPvYEzlUgJ2wF+stgPZH1UqBm1s8=
18 -github.com/mattn/go-isatty v0.0.16/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/yFXSvRLM=
19 -github.com/mattn/go-isatty v0.0.19/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
20 -github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
21 -github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
13 github.com/miekg/dns v1.1.72 h1:vhmr+TF2A3tuoGNkLDFK9zi36F2LS+hKTRW0Uf8kbzI=
14 github.com/miekg/dns v1.1.72/go.mod h1:+EuEPhdHOsfk6Wk5TT2CzssZdqkmFhf8r+aVyDEToIs=
24 -github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
15 github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 h1:Jamvg5psRIccs7FGNTlIRMkT8wgtp5eCXdBlqhYGL6U=
16 github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
27 -github.com/rs/xid v1.6.0/go.mod h1:7XoLgs4eV+QndskICGsho+ADou8ySMSjJKDIan90Nz0=
28 -github.com/rs/zerolog v1.34.0 h1:k43nTLIwcTVQAncfCw4KZ2VY6ukYoZaBPNOE8txlOeY=
29 -github.com/rs/zerolog v1.34.0/go.mod h1:bJsvje4Z08ROH4Nhs5iH600c3IkWhwp44iRc54W6wYQ=
17 github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
18 github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
19 golang.org/x/crypto v0.48.0 h1:/VRzVqiRSggnhY7gNRxPauEQ5Drw9haKdM0jqfcCFts=
20 golang.org/x/crypto v0.48.0/go.mod h1:r0kV5h3qnFPlQnBSrULhlsRfryS2pmewsg+XfMgkVos=
34 -golang.org/x/mod v0.33.0 h1:tHFzIWbBifEmbwtGz65eaWyGiGZatSrT9prnU8DbVL8=
35 -golang.org/x/mod v0.33.0/go.mod h1:swjeQEj+6r7fODbD2cqrnje9PnziFuw4bmLbBZFrQ5w=
36 -golang.org/x/net v0.51.0 h1:94R/GTO7mt3/4wIKpcR5gkGmRLOuE/2hNGeWq/GBIFo=
37 -golang.org/x/net v0.51.0/go.mod h1:aamm+2QF5ogm02fjy5Bb7CQ0WMt1/WVM7FtyaTLlA9Y=
21 +golang.org/x/mod v0.32.0 h1:9F4d3PHLljb6x//jOyokMv3eX+YDeepZSEo3mFJy93c=
22 +golang.org/x/mod v0.32.0/go.mod h1:SgipZ/3h2Ci89DlEtEXWUk/HteuRin+HHhN+WbNhguU=
23 +golang.org/x/net v0.50.0 h1:ucWh9eiCGyDR3vtzso0WMQinm2Dnt8cFMuQa9K33J60=
24 +golang.org/x/net v0.50.0/go.mod h1:UgoSli3F/pBgdJBHCTc+tp3gmrU4XswgGRgtnwWTfyM=
25 golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
26 golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
40 -golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
41 -golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
42 -golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
27 golang.org/x/sys v0.41.0 h1:Ivj+2Cp/ylzLiEU89QhWblYnOE9zerudt9Ftecq2C6k=
28 golang.org/x/sys v0.41.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
29 golang.org/x/text v0.34.0 h1:oL/Qq0Kdaqxa1KbNeMKwQq0reLCCaFtqu2eNuSeNHbk=
30 golang.org/x/text v0.34.0/go.mod h1:homfLqTYRFyVYemLBFl5GgL/DWEiH5wcsQ5gSh1yziA=
47 -golang.org/x/tools v0.42.0 h1:uNgphsn75Tdz5Ji2q36v/nsFSfR/9BRFvqhGBaJGd5k=
48 -golang.org/x/tools v0.42.0/go.mod h1:Ma6lCIwGZvHK6XtgbswSoWroEkhugApmsXyrUmBhfr0=
31 +golang.org/x/tools v0.41.0 h1:a9b8iMweWG+S0OBnlU36rzLp20z1Rp10w+IY2czHTQc=
32 +golang.org/x/tools v0.41.0/go.mod h1:XSY6eDqxVNiYgezAVqqCeihT4j1U2CCsqvH3WhQpnlg=
33 gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
34 gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
internal/testutil/cert.go new
+57
@@ -0,0 +1,57 @@
1 +package testutil
2 +
3 +import (
4 + "crypto/ecdsa"
5 + "crypto/elliptic"
6 + "crypto/rand"
7 + "crypto/x509"
8 + "crypto/x509/pkix"
9 + "encoding/pem"
10 + "math/big"
11 + "net"
12 + "time"
13 +)
14 +
15 +func SelfSignedCertPEM(host string) (certPEM, keyPEM []byte, err error) {
16 + priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
17 + if err != nil {
18 + return nil, nil, err
19 + }
20 +
21 + serialLimit := new(big.Int).Lsh(big.NewInt(1), 128)
22 + serial, err := rand.Int(rand.Reader, serialLimit)
23 + if err != nil {
24 + return nil, nil, err
25 + }
26 +
27 + template := &x509.Certificate{
28 + SerialNumber: serial,
29 + Subject: pkix.Name{
30 + CommonName: host,
31 + },
32 + NotBefore: time.Now().Add(-time.Hour),
33 + NotAfter: time.Now().Add(24 * time.Hour),
34 + KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment,
35 + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
36 + BasicConstraintsValid: true,
37 + }
38 +
39 + if ip := net.ParseIP(host); ip != nil {
40 + template.IPAddresses = []net.IP{ip}
41 + } else {
42 + template.DNSNames = []string{host}
43 + }
44 +
45 + der, err := x509.CreateCertificate(rand.Reader, template, template, &priv.PublicKey, priv)
46 + if err != nil {
47 + return nil, nil, err
48 + }
49 +
50 + certPEM = pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})
51 + keyBytes, err := x509.MarshalECPrivateKey(priv)
52 + if err != nil {
53 + return nil, nil, err
54 + }
55 + keyPEM = pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyBytes})
56 + return certPEM, keyPEM, nil
57 +}
portal/acme/acme.go
+200 -323
@@ -12,19 +12,18 @@ import (
12 "encoding/pem"
13 "errors"
14 "fmt"
15 + "net"
16 "os"
17 "path/filepath"
18 "strings"
19 "sync"
20 + "time"
21
22 "github.com/go-acme/lego/v4/certcrypto"
23 "github.com/go-acme/lego/v4/certificate"
24 lego "github.com/go-acme/lego/v4/lego"
25 "github.com/go-acme/lego/v4/providers/dns/cloudflare"
26 "github.com/go-acme/lego/v4/registration"
25 - "github.com/rs/zerolog/log"
26 -
27 - "gosuda.org/portal/types"
27 )
28
29 const (
@@ -35,309 +34,221 @@ const (
34 defaultACMEEmailPrefix = "acme@"
35 )
36
38 -type certTarget struct {
39 - Name string
40 - KeyFile string
41 - CertFile string
42 - Domains []string
37 +type Config struct {
38 + BaseDomain string
39 + KeyDir string
40 + CloudflareToken string
41 +}
42 +
43 +type Manager struct {
44 + cfg Config
45 +
46 + mu sync.RWMutex
47 + startOnce sync.Once
48 + stopOnce sync.Once
49 + stopCh chan struct{}
50 + wg sync.WaitGroup
51 }
52
53 type provisionConfig struct {
46 - TargetName string
54 KeyFile string
55 CertFile string
49 - Email string
56 AccountKeyFile string
57 RegistrationFile string
58 + Email string
59 CloudflareToken string
60 Domains []string
61 }
62
56 -type Config struct {
57 - BaseDomain string
58 - KeyDir string
59 - CloudflareToken string
63 +type acmeUser struct {
64 + Key crypto.PrivateKey
65 + Registration *registration.Resource
66 + Email string
67 }
68
62 -type Manager struct {
63 - stopCh chan struct{}
64 - cfg Config
65 - waitGroup sync.WaitGroup
66 - mu sync.RWMutex
67 - startOnce sync.Once
68 - stopOnce sync.Once
69 -}
69 +func NewManager(cfg Config) (*Manager, error) {
70 + cfg.BaseDomain = normalizeHost(cfg.BaseDomain)
71 + cfg.KeyDir = strings.TrimSpace(cfg.KeyDir)
72 + cfg.CloudflareToken = strings.TrimSpace(cfg.CloudflareToken)
73
71 -func NewManager(ctx context.Context, cfg Config) (*Manager, string, error) {
72 - if strings.TrimSpace(cfg.KeyDir) == "" {
73 - return nil, "", nil
74 + if cfg.KeyDir == "" {
75 + return nil, errors.New("acme key directory is required")
76 + }
77 + if cfg.BaseDomain == "" {
78 + return nil, errors.New("acme base domain is required")
79 }
80
76 - manager := &Manager{
77 - cfg: Config{
78 - BaseDomain: cfg.BaseDomain,
79 - KeyDir: cfg.KeyDir,
80 - CloudflareToken: cfg.CloudflareToken,
81 - },
81 + return &Manager{
82 + cfg: cfg,
83 stopCh: make(chan struct{}),
83 - }
84 + }, nil
85 +}
86
85 - generated, err := EnsureLocalDevelopmentCertificate(manager.cfg.KeyDir, manager.cfg.BaseDomain)
86 - if err != nil {
87 - return nil, "", fmt.Errorf("ensure local development certificate: %w", err)
88 - }
89 - if generated {
90 - log.Info().
91 - Str("base_host", manager.cfg.BaseDomain).
92 - Str("key_file", manager.SigningKeyFile()).
93 - Msg("[signer] generated self-signed localhost development certificate")
87 +func (m *Manager) EnsureCertificate(ctx context.Context) (string, string, error) {
88 + if m == nil {
89 + return "", "", errors.New("acme manager is nil")
90 }
91
96 - keyFile, err := manager.PrepareSigningKey(ctx)
97 - if err != nil {
98 - return nil, "", err
92 + if isLocalhost(m.cfg.BaseDomain) {
93 + if err := ensureLocalDevelopmentCertificate(m.cfg.KeyDir, m.cfg.BaseDomain); err != nil {
94 + return "", "", err
95 + }
96 + return m.TLSFiles()
97 }
100 - return manager, keyFile, nil
101 -}
98
103 -// PrepareSigningKey resolves the signer key path and runs ACME provisioning when required.
104 -// It encapsulates local-host detection and ACME enable/disable policy.
105 -func (m *Manager) PrepareSigningKey(ctx context.Context) (string, error) {
106 - if m == nil {
107 - return "", errors.New("acme manager is nil")
99 + if m.cfg.CloudflareToken == "" {
100 + return "", "", errors.New("cloudflare token is required for non-local relay certificates")
101 }
102
110 - keyDir := strings.TrimSpace(m.cfg.KeyDir)
111 - baseDomain := strings.TrimSpace(m.cfg.BaseDomain)
112 - cloudflareToken := strings.TrimSpace(m.cfg.CloudflareToken)
113 - isLocalBaseHost := types.IsLocalhost(baseDomain)
103 + if err := EnsureDNSRecords(ctx, m.cfg.BaseDomain, m.cfg.CloudflareToken); err != nil {
104 + return "", "", fmt.Errorf("ensure dns records: %w", err)
105 + }
106
115 - shouldEnsureWithACME := keyDir != "" && cloudflareToken != "" && baseDomain != "" && !isLocalBaseHost
116 - if shouldEnsureWithACME {
117 - keyFile, err := m.EnsureSigningKey(ctx)
118 - if err != nil {
119 - return "", fmt.Errorf("ensure keyless signing key: %w", err)
107 + certFile, keyFile, err := m.TLSFiles()
108 + if err == nil {
109 + covered, err := certCoversDomains(certFile, certificateDomains(m.cfg.BaseDomain))
110 + if err == nil && covered {
111 + return certFile, keyFile, nil
112 }
121 - return keyFile, nil
113 }
114
124 - log.Info().
125 - Bool("has_key_dir", keyDir != "").
126 - Bool("has_cloudflare_token", cloudflareToken != "").
127 - Bool("has_base_domain", baseDomain != "").
128 - Bool("is_local_base_host", isLocalBaseHost).
129 - Msg("[signer] ACME issuance disabled (requires key directory, Cloudflare token, and base domain)")
130 -
131 - if keyDir == "" {
132 - return "", nil
115 + if err := m.provision(ctx); err != nil {
116 + return "", "", err
117 }
134 - return m.SigningKeyFile(), nil
118 + return m.TLSFiles()
119 }
120
137 -func (m *Manager) keyDir() string {
138 - if m == nil {
139 - return ""
121 +func (m *Manager) Start(ctx context.Context) {
122 + if m == nil || isLocalhost(m.cfg.BaseDomain) || m.cfg.CloudflareToken == "" {
123 + return
124 }
141 - return m.cfg.KeyDir
125 +
126 + m.startOnce.Do(func() {
127 + m.wg.Add(1)
128 + go m.renewalLoop(ctx)
129 + })
130 }
131
144 -// SigningKeyFile returns the unified signer key path under configured key directory.
145 -func (m *Manager) SigningKeyFile() string {
132 +func (m *Manager) Stop() {
133 if m == nil {
147 - return ""
134 + return
135 }
149 - keyDir := m.keyDir()
150 - if keyDir == "" {
151 - return ""
152 - }
153 - return keyPath(keyDir)
154 -}
155 -
156 -type acmeUser struct {
157 - Key crypto.PrivateKey
158 - Registration *registration.Resource
159 - Email string
136 + m.stopOnce.Do(func() {
137 + close(m.stopCh)
138 + })
139 + m.wg.Wait()
140 }
141
162 -func (u *acmeUser) GetEmail() string {
163 - if u == nil {
164 - return ""
142 +func (m *Manager) TLSFiles() (string, string, error) {
143 + if m == nil {
144 + return "", "", errors.New("acme manager is nil")
145 }
166 - return u.Email
167 -}
168 -
169 -func (u *acmeUser) GetRegistration() *registration.Resource {
170 - if u == nil {
171 - return nil
146 + certFile := filepath.Join(m.cfg.KeyDir, fullChainFileName)
147 + keyFile := filepath.Join(m.cfg.KeyDir, keyFileName)
148 + if !fileExists(certFile) || !fileExists(keyFile) {
149 + return "", "", errors.New("relay certificate files do not exist")
150 }
173 - return u.Registration
151 + return certFile, keyFile, nil
152 }
153
176 -func (u *acmeUser) GetPrivateKey() crypto.PrivateKey {
177 - if u == nil {
178 - return nil
154 +func (m *Manager) provision(ctx context.Context) error {
155 + cfg := provisionConfig{
156 + KeyFile: filepath.Join(m.cfg.KeyDir, keyFileName),
157 + CertFile: filepath.Join(m.cfg.KeyDir, fullChainFileName),
158 + AccountKeyFile: filepath.Join(m.cfg.KeyDir, accountKeyFileName),
159 + RegistrationFile: filepath.Join(m.cfg.KeyDir, registrationFileName),
160 + Email: defaultACMEEmailPrefix + m.cfg.BaseDomain,
161 + CloudflareToken: m.cfg.CloudflareToken,
162 + Domains: certificateDomains(m.cfg.BaseDomain),
163 }
180 - return u.Key
181 -}
164
183 -// EnsureSigningKey provisions a keyless signing key via ACME DNS-01 when missing.
184 -func (m *Manager) EnsureSigningKey(ctx context.Context) (string, error) {
185 - if m == nil {
186 - return "", errors.New("acme manager is nil")
187 - }
188 - configuredKeyDir := m.keyDir()
189 - if configuredKeyDir == "" {
190 - return "", nil
165 + for _, path := range []string{cfg.KeyFile, cfg.CertFile, cfg.AccountKeyFile, cfg.RegistrationFile} {
166 + if err := ensureParentDir(path); err != nil {
167 + return err
168 + }
169 }
170 if err := ctx.Err(); err != nil {
193 - return "", fmt.Errorf("keyless provisioning canceled: %w", err)
171 + return fmt.Errorf("acme provisioning canceled: %w", err)
172 }
173
196 - baseDomain := m.cfg.BaseDomain
197 - if baseDomain == "" {
198 - return "", errors.New("base domain is required for ACME provisioning")
174 + client, _, err := newClient(cfg)
175 + if err != nil {
176 + return err
177 }
178
201 - targets, err := buildCertTargets(baseDomain, configuredKeyDir)
179 + obtained, err := client.Certificate.Obtain(certificate.ObtainRequest{
180 + Domains: cfg.Domains,
181 + Bundle: true,
182 + })
183 if err != nil {
203 - return "", err
184 + return fmt.Errorf("obtain certificate: %w", err)
185 }
205 - signerKeyFile := keyPath(configuredKeyDir)
206 -
207 - missingTargets := make([]certTarget, 0, len(targets))
208 - for _, target := range targets {
209 - if !fileExists(target.KeyFile) || !fileExists(target.CertFile) {
210 - missingTargets = append(missingTargets, target)
211 - continue
212 - }
213 -
214 - covered, coverageErr := certCoversDomains(target.CertFile, target.Domains)
215 - if coverageErr != nil {
216 - log.Warn().
217 - Err(coverageErr).
218 - Str("target", target.Name).
219 - Str("cert_file", target.CertFile).
220 - Msg("[signer] failed to validate existing certificate; re-issuing via ACME")
221 - missingTargets = append(missingTargets, target)
222 - continue
223 - }
224 - if !covered {
225 - log.Warn().
226 - Str("target", target.Name).
227 - Str("cert_file", target.CertFile).
228 - Strs("required_domains", target.Domains).
229 - Msg("[signer] existing certificate does not cover required domains; re-issuing via ACME")
230 - missingTargets = append(missingTargets, target)
231 - continue
232 - }
186 + if len(obtained.Certificate) == 0 || len(obtained.PrivateKey) == 0 {
187 + return errors.New("acme obtain response missing certificate or private key")
188 }
234 - if len(missingTargets) == 0 {
235 - return signerKeyFile, nil
189 +
190 + if err := writeFileAtomic(cfg.CertFile, obtained.Certificate, 0o644); err != nil {
191 + return fmt.Errorf("write certificate chain: %w", err)
192 }
237 - if !hasCloudflareToken(m.cfg.CloudflareToken) {
238 - if !fileExists(signerKeyFile) {
239 - log.Warn().
240 - Str("key_file", signerKeyFile).
241 - Msg("[signer] keyless key file is missing and Cloudflare credentials are not set; signer will stay disabled")
242 - }
243 - for _, target := range missingTargets {
244 - log.Warn().
245 - Str("target", target.Name).
246 - Str("key_file", target.KeyFile).
247 - Str("cert_file", target.CertFile).
248 - Msg("[signer] ACME target is missing and Cloudflare credentials are not set")
249 - }
250 - return signerKeyFile, nil
193 + if err := writeFileAtomic(cfg.KeyFile, obtained.PrivateKey, 0o600); err != nil {
194 + return fmt.Errorf("write private key: %w", err)
195 }
196 + return nil
197 +}
198
253 - for _, target := range missingTargets {
254 - cfg, buildErr := buildProvisionConfig(baseDomain, target, m.cfg.CloudflareToken)
255 - if buildErr != nil {
256 - return "", buildErr
257 - }
258 - log.Info().
259 - Str("target", cfg.TargetName).
260 - Strs("domains", cfg.Domains).
261 - Str("key_file", cfg.KeyFile).
262 - Str("cert_file", cfg.CertFile).
263 - Msg("[signer] ACME target is missing or invalid; issuing certificate with ACME DNS-01 via Cloudflare")
264 -
265 - if err := m.provisionCertificate(cfg); err != nil {
266 - return "", err
199 +func (m *Manager) renewalLoop(ctx context.Context) {
200 + defer m.wg.Done()
201 +
202 + ticker := time.NewTicker(24 * time.Hour)
203 + defer ticker.Stop()
204 +
205 + for {
206 + select {
207 + case <-ctx.Done():
208 + return
209 + case <-m.stopCh:
210 + return
211 + case <-ticker.C:
212 + if !m.shouldRenew() {
213 + continue
214 + }
215 + renewCtx, cancel := context.WithTimeout(ctx, 2*time.Minute)
216 + _ = m.provision(renewCtx)
217 + cancel()
218 }
219 }
269 - return signerKeyFile, nil
220 }
221
272 -// TLSFiles returns the unified fullchain and private key file paths when both exist.
273 -func (m *Manager) TLSFiles() (string, string) {
274 - if m == nil {
275 - return "", ""
276 - }
277 - keyDir := m.keyDir()
278 - if keyDir == "" {
279 - return "", ""
280 - }
281 -
282 - keyFile := keyPath(keyDir)
283 - certFile := fullChainPath(keyDir)
284 - if fileExists(certFile) && fileExists(keyFile) {
285 - return certFile, keyFile
286 - }
222 +func (m *Manager) shouldRenew() bool {
223 + m.mu.RLock()
224 + defer m.mu.RUnlock()
225
288 - return "", ""
226 + certFile := filepath.Join(m.cfg.KeyDir, fullChainFileName)
227 + needsRenewal, err := certNeedsRenewal(certFile, certificateDomains(m.cfg.BaseDomain))
228 + return err == nil && needsRenewal
229 }
230
291 -func buildProvisionConfig(baseDomain string, target certTarget, cloudflareToken string) (provisionConfig, error) {
292 - keyDir := filepath.Dir(target.KeyFile)
293 - accountKeyFile := filepath.Join(keyDir, accountKeyFileName)
294 - registrationFile := filepath.Join(keyDir, registrationFileName)
295 - email := defaultACMEEmailPrefix + baseDomain
296 -
297 - if target.KeyFile == "" || target.CertFile == "" || len(target.Domains) == 0 {
298 - return provisionConfig{}, errors.New("invalid ACME target")
299 - }
300 - if _, err := resolveDomain(baseDomain); err != nil {
301 - return provisionConfig{}, err
302 - }
303 -
304 - return provisionConfig{
305 - TargetName: target.Name,
306 - KeyFile: target.KeyFile,
307 - CertFile: target.CertFile,
308 - Email: email,
309 - Domains: target.Domains,
310 - AccountKeyFile: accountKeyFile,
311 - RegistrationFile: registrationFile,
312 - CloudflareToken: cloudflareToken,
313 - }, nil
231 +func certificateDomains(baseDomain string) []string {
232 + return []string{baseDomain, "*." + baseDomain}
233 }
234
316 -func resolveDomain(baseDomain string) (string, error) {
317 - if baseDomain == "" {
318 - return "", errors.New("base domain is required")
235 +func certNeedsRenewal(certFile string, domains []string) (bool, error) {
236 + certPEM, err := os.ReadFile(certFile)
237 + if err != nil {
238 + return false, err
239 }
320 - return baseDomain, nil
321 -}
322 -
323 -func buildCertTargets(baseDomain, configuredKeyDir string) ([]certTarget, error) {
324 - base, err := resolveDomain(baseDomain)
240 + cert, err := ParseCertificatePEM(certPEM)
241 if err != nil {
326 - return nil, err
242 + return false, err
243 }
328 - keyDir := configuredKeyDir
329 - if keyDir == "" {
330 - return nil, errors.New("key directory is required")
244 + if time.Until(cert.NotAfter) < 30*24*time.Hour {
245 + return true, nil
246 }
332 -
333 - return []certTarget{
334 - {
335 - Name: "unified",
336 - KeyFile: keyPath(keyDir),
337 - CertFile: fullChainPath(keyDir),
338 - Domains: []string{"*." + base, base},
339 - },
340 - }, nil
247 + covered, err := certCoversDomains(certFile, domains)
248 + if err != nil {
249 + return false, err
250 + }
251 + return !covered, nil
252 }
253
254 func certCoversDomains(certFile string, domains []string) (bool, error) {
@@ -345,43 +256,32 @@ func certCoversDomains(certFile string, domains []string) (bool, error) {
256 if err != nil {
257 return false, err
258 }
348 -
259 cert, err := ParseCertificatePEM(certPEM)
260 if err != nil {
261 return false, err
262 }
353 -
263 for _, domain := range domains {
355 - if after, ok := strings.CutPrefix(domain, "*."); ok {
356 - probeHost := "acme-probe." + after
357 - if err := cert.VerifyHostname(probeHost); err != nil {
358 - return false, err
264 + if wildcardDomain, ok := strings.CutPrefix(domain, "*."); ok {
265 + if err := cert.VerifyHostname("probe." + wildcardDomain); err != nil {
266 + return false, nil
267 }
268 continue
269 }
362 -
270 if err := cert.VerifyHostname(domain); err != nil {
364 - return false, err
271 + return false, nil
272 }
273 }
367 -
274 return true, nil
275 }
276
371 -func (m *Manager) provisionCertificate(cfg provisionConfig) error {
372 - for _, path := range []string{cfg.KeyFile, cfg.CertFile, cfg.AccountKeyFile, cfg.RegistrationFile} {
373 - if err := ensureParentDir(path); err != nil {
374 - return err
375 - }
376 - }
377 -
277 +func newClient(cfg provisionConfig) (*lego.Client, *acmeUser, error) {
278 accountKey, err := loadOrCreateAccountKey(cfg.AccountKeyFile)
279 if err != nil {
380 - return fmt.Errorf("load ACME account key: %w", err)
280 + return nil, nil, fmt.Errorf("load acme account key: %w", err)
281 }
282 accountReg, err := loadRegistration(cfg.RegistrationFile)
283 if err != nil {
384 - return fmt.Errorf("load ACME registration: %w", err)
284 + return nil, nil, fmt.Errorf("load acme registration: %w", err)
285 }
286
287 user := &acmeUser{
@@ -389,13 +289,14 @@ func (m *Manager) provisionCertificate(cfg provisionConfig) error {
289 Key: accountKey,
290 Registration: accountReg,
291 }
292 +
293 clientConfig := lego.NewConfig(user)
294 clientConfig.CADirURL = lego.LEDirectoryProduction
295 clientConfig.Certificate.KeyType = certcrypto.RSA2048
296
297 client, err := lego.NewClient(clientConfig)
298 if err != nil {
398 - return fmt.Errorf("create ACME client: %w", err)
299 + return nil, nil, fmt.Errorf("create acme client: %w", err)
300 }
301
302 cfConfig := cloudflare.NewDefaultConfig()
@@ -403,64 +304,34 @@ func (m *Manager) provisionCertificate(cfg provisionConfig) error {
304
305 provider, err := cloudflare.NewDNSProviderConfig(cfConfig)
306 if err != nil {
406 - return fmt.Errorf("create Cloudflare DNS provider: %w", err)
307 + return nil, nil, fmt.Errorf("create cloudflare dns provider: %w", err)
308 }
408 - err = client.Challenge.SetDNS01Provider(provider)
409 - if err != nil {
410 - return fmt.Errorf("set DNS-01 challenge provider: %w", err)
309 + if err := client.Challenge.SetDNS01Provider(provider); err != nil {
310 + return nil, nil, fmt.Errorf("set dns01 provider: %w", err)
311 }
312
313 if user.Registration == nil {
414 - reg, regErr := client.Registration.Register(registration.RegisterOptions{
415 - TermsOfServiceAgreed: true,
416 - })
417 - if regErr != nil {
418 - return fmt.Errorf("register ACME account: %w", regErr)
314 + reg, err := client.Registration.Register(registration.RegisterOptions{TermsOfServiceAgreed: true})
315 + if err != nil {
316 + return nil, nil, fmt.Errorf("register acme account: %w", err)
317 }
318 user.Registration = reg
421 - if saveErr := saveRegistration(cfg.RegistrationFile, reg); saveErr != nil {
422 - return fmt.Errorf("persist ACME registration: %w", saveErr)
319 + if err := saveRegistration(cfg.RegistrationFile, reg); err != nil {
320 + return nil, nil, fmt.Errorf("persist acme registration: %w", err)
321 }
322 }
323
426 - obtained, err := client.Certificate.Obtain(certificate.ObtainRequest{
427 - Domains: cfg.Domains,
428 - Bundle: true,
429 - })
430 - if err != nil {
431 - return fmt.Errorf("obtain certificate: %w", err)
432 - }
433 - if len(obtained.PrivateKey) == 0 {
434 - return errors.New("ACME response did not include private key")
435 - }
436 - if len(obtained.Certificate) == 0 {
437 - return errors.New("ACME response did not include certificate chain")
438 - }
439 -
440 - if err := writeFileAtomic(cfg.KeyFile, obtained.PrivateKey, 0o600); err != nil {
441 - return fmt.Errorf("write keyless private key: %w", err)
442 - }
443 - if err := writeFileAtomic(cfg.CertFile, obtained.Certificate, 0o644); err != nil {
444 - return fmt.Errorf("write keyless certificate chain: %w", err)
445 - }
446 -
447 - log.Info().
448 - Str("target", cfg.TargetName).
449 - Str("key_file", cfg.KeyFile).
450 - Str("cert_file", cfg.CertFile).
451 - Strs("domains", cfg.Domains).
452 - Msg("[signer] ACME certificate issued and stored")
453 - return nil
324 + return client, user, nil
325 }
326
327 +func (u *acmeUser) GetEmail() string { return u.Email }
328 +func (u *acmeUser) GetRegistration() *registration.Resource { return u.Registration }
329 +func (u *acmeUser) GetPrivateKey() crypto.PrivateKey { return u.Key }
330 +
331 func loadOrCreateAccountKey(path string) (crypto.PrivateKey, error) {
332 keyPEM, err := os.ReadFile(path)
333 if err == nil {
459 - key, parseErr := parsePEMPrivateKey(keyPEM)
460 - if parseErr != nil {
461 - return nil, parseErr
462 - }
463 - return key, nil
334 + return parsePEMPrivateKey(keyPEM)
335 }
336 if !errors.Is(err, os.ErrNotExist) {
337 return nil, err
@@ -468,18 +339,15 @@ func loadOrCreateAccountKey(path string) (crypto.PrivateKey, error) {
339
340 key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
341 if err != nil {
471 - return nil, fmt.Errorf("generate ACME account key: %w", err)
342 + return nil, fmt.Errorf("generate account key: %w", err)
343 }
344 pkcs8, err := x509.MarshalPKCS8PrivateKey(key)
345 if err != nil {
475 - return nil, fmt.Errorf("marshal ACME account key: %w", err)
346 + return nil, fmt.Errorf("marshal account key: %w", err)
347 }
477 - pemData := pem.EncodeToMemory(&pem.Block{
478 - Type: "PRIVATE KEY",
479 - Bytes: pkcs8,
480 - })
481 - if err := writeFileAtomic(path, pemData, 0o600); err != nil {
482 - return nil, fmt.Errorf("persist ACME account key: %w", err)
348 + keyPEM = pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: pkcs8})
349 + if err := writeFileAtomic(path, keyPEM, 0o600); err != nil {
350 + return nil, fmt.Errorf("persist account key: %w", err)
351 }
352 return key, nil
353 }
@@ -487,9 +355,8 @@ func loadOrCreateAccountKey(path string) (crypto.PrivateKey, error) {
355 func parsePEMPrivateKey(keyPEM []byte) (crypto.PrivateKey, error) {
356 block, _ := pem.Decode(keyPEM)
357 if block == nil {
490 - return nil, errors.New("invalid private key PEM")
358 + return nil, errors.New("invalid private key pem")
359 }
492 -
360 if key, err := x509.ParsePKCS8PrivateKey(block.Bytes); err == nil {
361 switch typed := key.(type) {
362 case *ecdsa.PrivateKey:
@@ -538,10 +405,7 @@ func ensureParentDir(path string) error {
405 if dir == "" || dir == "." {
406 return nil
407 }
541 - if err := os.MkdirAll(dir, 0o700); err != nil {
542 - return fmt.Errorf("create directory %s: %w", dir, err)
543 - }
544 - return nil
408 + return os.MkdirAll(dir, 0o700)
409 }
410
411 func fileExists(path string) bool {
@@ -562,9 +426,7 @@ func writeFileAtomic(path string, data []byte, mode os.FileMode) error {
426 return err
427 }
428 tmpName := tmp.Name()
565 - defer func() {
566 - _ = os.Remove(tmpName)
567 - }()
429 + defer func() { _ = os.Remove(tmpName) }()
430
431 if _, err := tmp.Write(data); err != nil {
432 _ = tmp.Close()
@@ -583,14 +445,29 @@ func writeFileAtomic(path string, data []byte, mode os.FileMode) error {
445 return os.Chmod(path, mode)
446 }
447
586 -func hasCloudflareToken(cloudflareToken string) bool {
587 - return cloudflareToken != ""
448 +func ParseCertificatePEM(pemData []byte) (*x509.Certificate, error) {
449 + block, _ := pem.Decode(pemData)
450 + if block == nil {
451 + return nil, errors.New("no pem block found")
452 + }
453 + return x509.ParseCertificate(block.Bytes)
454 }
455
590 -func fullChainPath(keyDir string) string {
591 - return filepath.Join(keyDir, fullChainFileName)
456 +func normalizeHost(host string) string {
457 + host = strings.ToLower(strings.TrimSpace(host))
458 + host = strings.TrimPrefix(host, "*.")
459 + host = strings.TrimSuffix(host, ".")
460 + return host
461 }
462
594 -func keyPath(keyDir string) string {
595 - return filepath.Join(keyDir, keyFileName)
463 +func isLocalhost(host string) bool {
464 + host = normalizeHost(host)
465 + switch host {
466 + case "", "localhost":
467 + return true
468 + }
469 + if ip := net.ParseIP(host); ip != nil {
470 + return ip.IsLoopback()
471 + }
472 + return strings.HasSuffix(host, ".localhost")
473 }
portal/acme/acme_test.go new
+35
@@ -0,0 +1,35 @@
1 +package acme
2 +
3 +import (
4 + "context"
5 + "testing"
6 +)
7 +
8 +func TestEnsureCertificateGeneratesLocalDevelopmentMaterial(t *testing.T) {
9 + t.Parallel()
10 +
11 + keyDir := t.TempDir()
12 + manager, err := NewManager(Config{
13 + BaseDomain: "localhost",
14 + KeyDir: keyDir,
15 + })
16 + if err != nil {
17 + t.Fatalf("NewManager() error = %v", err)
18 + }
19 +
20 + certFile, keyFile, err := manager.EnsureCertificate(context.Background())
21 + if err != nil {
22 + t.Fatalf("EnsureCertificate() error = %v", err)
23 + }
24 + if certFile == "" || keyFile == "" {
25 + t.Fatalf("EnsureCertificate() = %q, %q, want certificate paths", certFile, keyFile)
26 + }
27 +
28 + covered, err := certCoversDomains(certFile, []string{"localhost"})
29 + if err != nil {
30 + t.Fatalf("certCoversDomains() error = %v", err)
31 + }
32 + if !covered {
33 + t.Fatal("certCoversDomains() = false, want true")
34 + }
35 +}
portal/acme/dnsrecord.go
+22 -81
@@ -12,21 +12,15 @@ import (
12 "net/url"
13 "strings"
14 "time"
15 -
16 - "github.com/rs/zerolog/log"
17 -
18 - "gosuda.org/portal/types"
15 )
16
17 const (
18 cfAPIBase = "https://api.cloudflare.com/client/v4"
19 publicIPURL = "https://api4.ipify.org"
20 dnsHTTPTimeout = 15 * time.Second
25 - dnsAutoTTL = 1 // Cloudflare "automatic" TTL
21 + dnsAutoTTL = 1
22 )
23
28 -// Cloudflare API response types.
29 -
24 type cfError struct {
25 Message string `json:"message"`
26 Code int `json:"code"`
@@ -42,8 +36,6 @@ type cfDNSRecord struct {
36 Type string `json:"type"`
37 Name string `json:"name"`
38 Content string `json:"content"`
45 - TTL int `json:"ttl"`
46 - Proxied bool `json:"proxied"`
39 }
40
41 type cfZonesResult struct {
@@ -64,14 +56,11 @@ type cfRecordResult struct {
56 Success bool `json:"success"`
57 }
58
67 -// EnsureDNSRecords creates or updates Cloudflare A records for the base domain
68 -// and its wildcard subdomain, pointing to the server's detected public IP.
69 -// Skips silently when baseDomain is empty, localhost, or cloudflareToken is missing.
59 func EnsureDNSRecords(ctx context.Context, baseDomain, cloudflareToken string) error {
71 - baseDomain = strings.TrimSpace(baseDomain)
60 + baseDomain = normalizeHost(baseDomain)
61 cloudflareToken = strings.TrimSpace(cloudflareToken)
62
74 - if baseDomain == "" || cloudflareToken == "" || types.IsLocalhost(baseDomain) {
63 + if baseDomain == "" || cloudflareToken == "" || isLocalhost(baseDomain) {
64 return nil
65 }
66
@@ -80,30 +69,22 @@ func EnsureDNSRecords(ctx context.Context, baseDomain, cloudflareToken string) e
69
70 publicIP, err := detectPublicIP(ctx)
71 if err != nil {
83 - return fmt.Errorf("detect public IP: %w", err)
72 + return fmt.Errorf("detect public ip: %w", err)
73 }
74
86 - log.Info().
87 - Str("public_ip", publicIP).
88 - Str("base_domain", baseDomain).
89 - Msg("[DNS] detected server public IP")
90 -
75 zoneID, err := findZoneID(ctx, cloudflareToken, baseDomain)
76 if err != nil {
93 - return fmt.Errorf("find Cloudflare zone for %s: %w", baseDomain, err)
77 + return fmt.Errorf("find cloudflare zone: %w", err)
78 }
79
96 - targets := []string{baseDomain, "*." + baseDomain}
97 - for _, name := range targets {
80 + for _, name := range []string{baseDomain, "*." + baseDomain} {
81 if err := ensureARecord(ctx, cloudflareToken, zoneID, name, publicIP); err != nil {
82 return fmt.Errorf("ensure A record for %s: %w", name, err)
83 }
84 }
102 -
85 return nil
86 }
87
106 -// detectPublicIP fetches the server's public IPv4 address from an external service.
88 func detectPublicIP(ctx context.Context) (string, error) {
89 ctx, cancel := context.WithTimeout(ctx, dnsHTTPTimeout)
90 defer cancel()
@@ -112,7 +93,6 @@ func detectPublicIP(ctx context.Context) (string, error) {
93 if err != nil {
94 return "", err
95 }
115 -
96 resp, err := http.DefaultClient.Do(req)
97 if err != nil {
98 return "", err
@@ -123,72 +103,49 @@ func detectPublicIP(ctx context.Context) (string, error) {
103 if err != nil {
104 return "", err
105 }
126 -
106 ip := strings.TrimSpace(string(body))
107 parsed := net.ParseIP(ip)
129 - if parsed == nil {
130 - return "", fmt.Errorf("invalid IP address: %q", ip)
131 - }
132 - if parsed.To4() == nil {
133 - return "", fmt.Errorf("expected IPv4 address, got: %q", ip)
108 + if parsed == nil || parsed.To4() == nil {
109 + return "", fmt.Errorf("invalid ipv4 address: %q", ip)
110 }
135 -
111 return ip, nil
112 }
113
139 -// findZoneID looks up the Cloudflare zone ID by progressively stripping
140 -// subdomain labels from the given domain (e.g., portal.example.com → example.com).
114 func findZoneID(ctx context.Context, token, domain string) (string, error) {
115 parts := strings.Split(domain, ".")
143 - for i := range len(parts) - 1 {
116 + for i := 0; i < len(parts)-1; i++ {
117 candidate := strings.Join(parts[i:], ".")
145 -
118 zones, err := cfListZones(ctx, token, candidate)
119 if err != nil {
120 return "", err
121 }
150 - for _, z := range zones {
151 - if strings.EqualFold(z.Name, candidate) {
152 - log.Debug().
153 - Str("zone", z.Name).
154 - Str("zone_id", z.ID).
155 - Msg("[DNS] found Cloudflare zone")
156 - return z.ID, nil
122 + for _, zone := range zones {
123 + if strings.EqualFold(zone.Name, candidate) {
124 + return zone.ID, nil
125 }
126 }
127 }
160 -
161 - return "", fmt.Errorf("no Cloudflare zone found for domain %s", domain)
128 + return "", fmt.Errorf("no cloudflare zone found for %s", domain)
129 }
130
164 -// ensureARecord creates or updates a single A record.
165 -// If the record exists with the correct IP and proxy-off, it is left untouched.
131 func ensureARecord(ctx context.Context, token, zoneID, name, ip string) error {
132 records, err := cfListDNSRecords(ctx, token, zoneID, name, "A")
133 if err != nil {
134 return err
135 }
136
172 - for _, r := range records {
173 - if !strings.EqualFold(r.Name, name) {
137 + for _, record := range records {
138 + if !strings.EqualFold(record.Name, name) {
139 continue
140 }
176 - if r.Content == ip && !r.Proxied {
177 - log.Info().
178 - Str("name", name).
179 - Str("ip", ip).
180 - Msg("[DNS] A record already up to date")
141 + if record.Content == ip {
142 return nil
143 }
183 - // Record exists but IP or proxy status differs — update it.
184 - return cfUpdateDNSRecord(ctx, token, zoneID, r.ID, name, ip)
144 + return cfUpdateDNSRecord(ctx, token, zoneID, record.ID, name, ip)
145 }
186 -
146 return cfCreateDNSRecord(ctx, token, zoneID, name, ip)
147 }
148
190 -// ── Cloudflare API helpers ──────────────────────────────────────────
191 -
149 func cfListZones(ctx context.Context, token, name string) ([]cfZone, error) {
150 u, _ := url.Parse(cfAPIBase + "/zones")
151 q := u.Query()
@@ -224,7 +181,6 @@ func cfListDNSRecords(ctx context.Context, token, zoneID, name, recordType strin
181
182 func cfCreateDNSRecord(ctx context.Context, token, zoneID, name, ip string) error {
183 endpoint := fmt.Sprintf("%s/zones/%s/dns_records", cfAPIBase, zoneID)
227 -
184 body := map[string]any{
185 "type": "A",
186 "name": name,
@@ -240,17 +196,11 @@ func cfCreateDNSRecord(ctx context.Context, token, zoneID, name, ip string) erro
196 if !out.Success {
197 return cfErrs(out.Errors)
198 }
243 -
244 - log.Info().
245 - Str("name", name).
246 - Str("ip", ip).
247 - Msg("[DNS] created A record")
199 return nil
200 }
201
202 func cfUpdateDNSRecord(ctx context.Context, token, zoneID, recordID, name, ip string) error {
203 endpoint := fmt.Sprintf("%s/zones/%s/dns_records/%s", cfAPIBase, zoneID, recordID)
253 -
204 body := map[string]any{
205 "type": "A",
206 "name": name,
@@ -266,16 +216,9 @@ func cfUpdateDNSRecord(ctx context.Context, token, zoneID, recordID, name, ip st
216 if !out.Success {
217 return cfErrs(out.Errors)
218 }
269 -
270 - log.Info().
271 - Str("name", name).
272 - Str("ip", ip).
273 - Msg("[DNS] updated A record")
219 return nil
220 }
221
277 -// ── HTTP transport ──────────────────────────────────────────────────
278 -
222 func cfGet(ctx context.Context, token, rawURL string, out any) error {
223 req, err := http.NewRequestWithContext(ctx, http.MethodGet, rawURL, nil)
224 if err != nil {
@@ -289,7 +232,6 @@ func cfGet(ctx context.Context, token, rawURL string, out any) error {
232 return err
233 }
234 defer resp.Body.Close()
292 -
235 return json.NewDecoder(resp.Body).Decode(out)
236 }
237
@@ -311,17 +253,16 @@ func cfMutate(ctx context.Context, method, token, rawURL string, body any, out a
253 return err
254 }
255 defer resp.Body.Close()
314 -
256 return json.NewDecoder(resp.Body).Decode(out)
257 }
258
259 func cfErrs(errs []cfError) error {
260 if len(errs) == 0 {
320 - return errors.New("cloudflare API request failed")
261 + return errors.New("cloudflare api request failed")
262 }
322 - msgs := make([]string, 0, len(errs))
323 - for _, e := range errs {
324 - msgs = append(msgs, fmt.Sprintf("[%d] %s", e.Code, e.Message))
263 + messages := make([]string, 0, len(errs))
264 + for _, cfErr := range errs {
265 + messages = append(messages, fmt.Sprintf("[%d] %s", cfErr.Code, cfErr.Message))
266 }
326 - return fmt.Errorf("cloudflare API: %s", strings.Join(msgs, "; "))
267 + return errors.New(strings.Join(messages, "; "))
268 }
portal/acme/local.go
+26 -77
@@ -10,54 +10,40 @@ import (
10 "fmt"
11 "math/big"
12 "net"
13 - "strings"
13 + "path/filepath"
14 "time"
15 -
16 - "gosuda.org/portal/types"
15 )
16
17 const localDevelopmentCertificateTTL = 3650 * 24 * time.Hour
18
21 -// EnsureLocalDevelopmentCertificate ensures keyless TLS materials exist for localhost-style development.
22 -// It only acts when baseHost points to localhost/loopback semantics.
23 -func EnsureLocalDevelopmentCertificate(keyDir, baseHost string) (bool, error) {
24 - keyDir = strings.TrimSpace(keyDir)
25 - if keyDir == "" {
26 - return false, nil
27 - }
28 -
29 - baseHost = normalizeLocalDevelopmentHost(baseHost)
30 - if !types.IsLocalhost(baseHost) {
31 - return false, nil
32 - }
33 -
19 +func ensureLocalDevelopmentCertificate(keyDir, baseHost string) error {
20 domains := localDevelopmentDomains(baseHost)
35 - keyFile := keyPath(keyDir)
36 - certFile := fullChainPath(keyDir)
21 + keyFile := filepath.Join(keyDir, keyFileName)
22 + certFile := filepath.Join(keyDir, fullChainFileName)
23
24 if fileExists(keyFile) && fileExists(certFile) {
25 covered, err := certCoversDomains(certFile, domains)
26 if err == nil && covered {
41 - return false, nil
27 + return nil
28 }
29 }
30
31 if err := ensureParentDir(keyFile); err != nil {
46 - return false, err
32 + return err
33 }
34 if err := ensureParentDir(certFile); err != nil {
49 - return false, err
35 + return err
36 }
37
38 privateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
39 if err != nil {
54 - return false, fmt.Errorf("generate local development signing key: %w", err)
40 + return fmt.Errorf("generate local dev private key: %w", err)
41 }
42
43 serialLimit := new(big.Int).Lsh(big.NewInt(1), 128)
44 serialNumber, err := rand.Int(rand.Reader, serialLimit)
45 if err != nil {
60 - return false, fmt.Errorf("generate local development certificate serial: %w", err)
46 + return fmt.Errorf("generate local dev certificate serial: %w", err)
47 }
48
49 now := time.Now().UTC()
@@ -75,80 +61,43 @@ func EnsureLocalDevelopmentCertificate(keyDir, baseHost string) (bool, error) {
61 IsCA: true,
62 }
63
78 - dnsNames := make(map[string]struct{}, len(domains))
79 - ipAddresses := make(map[string]net.IP)
64 for _, domain := range domains {
65 if ip := net.ParseIP(domain); ip != nil {
82 - ipAddresses[ip.String()] = ip
83 - continue
84 - }
85 - domain = strings.TrimSpace(domain)
86 - if domain == "" {
66 + template.IPAddresses = append(template.IPAddresses, ip)
67 continue
68 }
89 - dnsNames[domain] = struct{}{}
90 - }
91 -
92 - for dnsName := range dnsNames {
93 - template.DNSNames = append(template.DNSNames, dnsName)
94 - }
95 - for _, ipAddress := range ipAddresses {
96 - template.IPAddresses = append(template.IPAddresses, ipAddress)
69 + template.DNSNames = append(template.DNSNames, domain)
70 }
71
72 certDER, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey)
73 if err != nil {
101 - return false, fmt.Errorf("create local development certificate: %w", err)
74 + return fmt.Errorf("create local dev certificate: %w", err)
75 }
76
77 certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER})
105 - privateKeyDER, err := x509.MarshalPKCS8PrivateKey(privateKey)
78 + keyDER, err := x509.MarshalPKCS8PrivateKey(privateKey)
79 if err != nil {
107 - return false, fmt.Errorf("marshal local development private key: %w", err)
80 + return fmt.Errorf("marshal local dev private key: %w", err)
81 }
109 - privateKeyPEM := pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: privateKeyDER})
82 + keyPEM := pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: keyDER})
83
111 - if err := writeFileAtomic(keyFile, privateKeyPEM, 0o600); err != nil {
112 - return false, fmt.Errorf("write local development private key: %w", err)
113 - }
84 if err := writeFileAtomic(certFile, certPEM, 0o644); err != nil {
115 - return false, fmt.Errorf("write local development certificate: %w", err)
85 + return fmt.Errorf("write local dev certificate: %w", err)
86 }
117 -
118 - return true, nil
119 -}
120 -
121 -func normalizeLocalDevelopmentHost(host string) string {
122 - host = strings.ToLower(strings.TrimSpace(host))
123 - host = strings.TrimPrefix(strings.TrimSuffix(host, "."), "*.")
124 - return host
87 + if err := writeFileAtomic(keyFile, keyPEM, 0o600); err != nil {
88 + return fmt.Errorf("write local dev private key: %w", err)
89 + }
90 + return nil
91 }
92
93 func localDevelopmentDomains(baseHost string) []string {
128 - baseHost = normalizeLocalDevelopmentHost(baseHost)
129 - if baseHost == "" {
130 - return []string{"localhost", "*.localhost", "127.0.0.1", "::1"}
131 - }
132 -
94 + baseHost = normalizeHost(baseHost)
95 domains := []string{"localhost", "*.localhost", "127.0.0.1", "::1"}
134 - domains = append(domains, baseHost)
135 -
136 - if net.ParseIP(baseHost) == nil {
137 - domains = append(domains, "*."+baseHost)
138 - }
139 -
140 - seen := make(map[string]struct{}, len(domains))
141 - out := make([]string, 0, len(domains))
142 - for _, domain := range domains {
143 - domain = strings.TrimSpace(domain)
144 - if domain == "" {
145 - continue
146 - }
147 - if _, ok := seen[domain]; ok {
148 - continue
96 + if baseHost != "" && baseHost != "localhost" {
97 + domains = append(domains, baseHost)
98 + if net.ParseIP(baseHost) == nil {
99 + domains = append(domains, "*."+baseHost)
100 }
150 - seen[domain] = struct{}{}
151 - out = append(out, domain)
101 }
153 - return out
102 + return domains
103 }
portal/acme/renew.go deleted
-280
@@ -1,280 +0,0 @@
1 -package acme
2 -
3 -import (
4 - "context"
5 - "crypto/x509"
6 - "encoding/pem"
7 - "errors"
8 - "fmt"
9 - "os"
10 - "time"
11 -
12 - "github.com/go-acme/lego/v4/certcrypto"
13 - "github.com/go-acme/lego/v4/certificate"
14 - lego "github.com/go-acme/lego/v4/lego"
15 - "github.com/go-acme/lego/v4/providers/dns/cloudflare"
16 - "github.com/rs/zerolog/log"
17 -)
18 -
19 -const (
20 - // RenewalCheckInterval is how often to check if renewal is needed.
21 - RenewalCheckInterval = 24 * time.Hour
22 -
23 - // RenewalThreshold is how long before expiration to renew.
24 - RenewalThreshold = 30 * 24 * time.Hour
25 -
26 - // RenewalOperationTimeout bounds a single renewal attempt.
27 - RenewalOperationTimeout = 2 * time.Minute
28 -)
29 -
30 -// Start begins the certificate renewal loop. It checks periodically if the
31 -// certificate needs renewal and renews it automatically.
32 -func (m *Manager) Start(ctx context.Context) {
33 - if m == nil || m.cfg.KeyDir == "" || !hasCloudflareToken(m.cfg.CloudflareToken) {
34 - return
35 - }
36 -
37 - m.startOnce.Do(func() {
38 - m.waitGroup.Add(1)
39 - go m.renewalLoop(ctx)
40 - })
41 -}
42 -
43 -// Stop stops the renewal loop.
44 -func (m *Manager) Stop() {
45 - if m == nil {
46 - return
47 - }
48 - m.stopOnce.Do(func() {
49 - close(m.stopCh)
50 - })
51 - m.waitGroup.Wait()
52 -}
53 -
54 -func (m *Manager) renewalLoop(ctx context.Context) {
55 - defer m.waitGroup.Done()
56 -
57 - ticker := time.NewTicker(RenewalCheckInterval)
58 - defer ticker.Stop()
59 -
60 - for {
61 - select {
62 - case <-m.stopCh:
63 - return
64 - case <-ctx.Done():
65 - return
66 - case <-ticker.C:
67 - if m.shouldRenew() {
68 - renewCtx, cancel := context.WithTimeout(ctx, RenewalOperationTimeout)
69 - err := m.renewCertificate(renewCtx)
70 - cancel()
71 - if err != nil {
72 - log.Error().Err(err).Msg("[acme] certificate renewal failed")
73 - }
74 - }
75 - }
76 - }
77 -}
78 -
79 -// shouldRenew checks if the certificate needs renewal.
80 -func (m *Manager) shouldRenew() bool {
81 - m.mu.RLock()
82 - defer m.mu.RUnlock()
83 -
84 - targets, err := buildCertTargets(m.cfg.BaseDomain, m.cfg.KeyDir)
85 - if err != nil {
86 - log.Debug().Err(err).Msg("[acme] cannot build certificate targets for renewal check")
87 - return false
88 - }
89 -
90 - for _, target := range targets {
91 - needsRenewal, checkErr := certNeedsRenewal(target.CertFile, target.Domains)
92 - if checkErr != nil {
93 - log.Debug().
94 - Err(checkErr).
95 - Str("target", target.Name).
96 - Str("cert_file", target.CertFile).
97 - Msg("[acme] cannot read certificate for renewal check")
98 - continue
99 - }
100 - if needsRenewal {
101 - return true
102 - }
103 - }
104 - return false
105 -}
106 -
107 -func certNeedsRenewal(certFile string, domains []string) (bool, error) {
108 - certPEM, err := os.ReadFile(certFile)
109 - if err != nil {
110 - return false, err
111 - }
112 -
113 - cert, err := ParseCertificatePEM(certPEM)
114 - if err != nil {
115 - return false, err
116 - }
117 -
118 - timeUntilExpiry := time.Until(cert.NotAfter)
119 - needsRenewal := timeUntilExpiry < RenewalThreshold
120 - if !needsRenewal {
121 - covered, coverageErr := certCoversDomains(certFile, domains)
122 - if coverageErr != nil {
123 - return false, coverageErr
124 - }
125 - if !covered {
126 - log.Warn().
127 - Str("cert_file", certFile).
128 - Strs("required_domains", domains).
129 - Msg("[acme] certificate does not cover required domains; renewal required")
130 - needsRenewal = true
131 - }
132 - }
133 - if needsRenewal {
134 - log.Info().
135 - Time("not_after", cert.NotAfter).
136 - Dur("time_remaining", timeUntilExpiry).
137 - Str("cert_file", certFile).
138 - Msg("[acme] certificate needs renewal")
139 - } else {
140 - log.Debug().
141 - Time("not_after", cert.NotAfter).
142 - Dur("time_remaining", timeUntilExpiry).
143 - Str("cert_file", certFile).
144 - Msg("[acme] certificate does not need renewal")
145 - }
146 - return needsRenewal, nil
147 -}
148 -
149 -// renewCertificate renews the certificate via ACME.
150 -func (m *Manager) renewCertificate(ctx context.Context) error {
151 - m.mu.Lock()
152 - defer m.mu.Unlock()
153 -
154 - if err := ctx.Err(); err != nil {
155 - return fmt.Errorf("renewal canceled: %w", err)
156 - }
157 -
158 - if m.cfg.BaseDomain == "" {
159 - return errors.New("base domain not configured")
160 - }
161 -
162 - targets, err := buildCertTargets(m.cfg.BaseDomain, m.cfg.KeyDir)
163 - if err != nil {
164 - return fmt.Errorf("build ACME targets: %w", err)
165 - }
166 -
167 - for _, target := range targets {
168 - needsRenewal, checkErr := certNeedsRenewal(target.CertFile, target.Domains)
169 - if checkErr != nil {
170 - continue
171 - }
172 - if !needsRenewal {
173 - continue
174 - }
175 -
176 - cfg, cfgErr := buildProvisionConfig(m.cfg.BaseDomain, target, m.cfg.CloudflareToken)
177 - if cfgErr != nil {
178 - return fmt.Errorf("build provision config: %w", cfgErr)
179 - }
180 -
181 - log.Info().
182 - Str("target", cfg.TargetName).
183 - Strs("domains", cfg.Domains).
184 - Str("cert_file", cfg.CertFile).
185 - Msg("[acme] renewing certificate")
186 -
187 - if err := m.doRenew(cfg); err != nil {
188 - return fmt.Errorf("renew certificate for target %s: %w", cfg.TargetName, err)
189 - }
190 -
191 - log.Info().
192 - Str("target", cfg.TargetName).
193 - Str("cert_file", cfg.CertFile).
194 - Msg("[acme] certificate renewed successfully")
195 - }
196 - return nil
197 -}
198 -
199 -func (m *Manager) doRenew(cfg provisionConfig) error {
200 - accountKey, err := loadOrCreateAccountKey(cfg.AccountKeyFile)
201 - if err != nil {
202 - return fmt.Errorf("load ACME account key: %w", err)
203 - }
204 - accountReg, err := loadRegistration(cfg.RegistrationFile)
205 - if err != nil {
206 - return fmt.Errorf("load ACME registration: %w", err)
207 - }
208 -
209 - user := &acmeUser{
210 - Email: cfg.Email,
211 - Key: accountKey,
212 - Registration: accountReg,
213 - }
214 - clientConfig := lego.NewConfig(user)
215 - clientConfig.CADirURL = lego.LEDirectoryProduction
216 - clientConfig.Certificate.KeyType = certcrypto.RSA2048
217 -
218 - client, err := lego.NewClient(clientConfig)
219 - if err != nil {
220 - return fmt.Errorf("create ACME client: %w", err)
221 - }
222 -
223 - cfConfig := cloudflare.NewDefaultConfig()
224 - cfConfig.AuthToken = cfg.CloudflareToken
225 -
226 - provider, err := cloudflare.NewDNSProviderConfig(cfConfig)
227 - if err != nil {
228 - return fmt.Errorf("create Cloudflare DNS provider: %w", err)
229 - }
230 - err = client.Challenge.SetDNS01Provider(provider)
231 - if err != nil {
232 - return fmt.Errorf("set DNS-01 challenge provider: %w", err)
233 - }
234 -
235 - certFile := cfg.CertFile
236 - certPEM, err := os.ReadFile(certFile)
237 - if err != nil {
238 - return fmt.Errorf("read certificate for renewal: %w", err)
239 - }
240 -
241 - keyPEM, err := os.ReadFile(cfg.KeyFile)
242 - if err != nil {
243 - return fmt.Errorf("read private key for renewal: %w", err)
244 - }
245 -
246 - renewed, err := client.Certificate.RenewWithOptions(certificate.Resource{
247 - Domain: cfg.Domains[0],
248 - Certificate: certPEM,
249 - PrivateKey: keyPEM,
250 - }, &certificate.RenewOptions{
251 - Bundle: true,
252 - })
253 - if err != nil {
254 - return fmt.Errorf("ACME renew: %w", err)
255 - }
256 - if len(renewed.Certificate) == 0 {
257 - return errors.New("ACME renewal response did not include certificate chain")
258 - }
259 -
260 - if err := writeFileAtomic(cfg.CertFile, renewed.Certificate, 0o644); err != nil {
261 - return fmt.Errorf("write renewed certificate chain: %w", err)
262 - }
263 -
264 - if len(renewed.PrivateKey) > 0 {
265 - if err := writeFileAtomic(cfg.KeyFile, renewed.PrivateKey, 0o600); err != nil {
266 - return fmt.Errorf("write renewed private key: %w", err)
267 - }
268 - }
269 -
270 - return nil
271 -}
272 -
273 -// ParseCertificatePEM parses a PEM-encoded certificate.
274 -func ParseCertificatePEM(pemData []byte) (*x509.Certificate, error) {
275 - block, _ := pem.Decode(pemData)
276 - if block == nil {
277 - return nil, errors.New("no PEM block found")
278 - }
279 - return x509.ParseCertificate(block.Bytes)
280 -}
portal/admin/handler.go deleted
-538
@@ -1,538 +0,0 @@
1 -package admin
2 -
3 -import (
4 - "encoding/base64"
5 - "encoding/json"
6 - "errors"
7 - "net"
8 - "net/http"
9 - "strings"
10 -
11 - "github.com/rs/zerolog/log"
12 -
13 - "gosuda.org/portal/portal"
14 - "gosuda.org/portal/portal/policy"
15 - "gosuda.org/portal/types"
16 -)
17 -
18 -var errInvalidLeaseID = errors.New("invalid lease ID")
19 -
20 -type HandlerConfig struct {
21 - Service *Service
22 - ServeAppStatic func(http.ResponseWriter, *http.Request, string, *portal.RelayServer)
23 - ListLeases func(*portal.RelayServer) any
24 - Stats func(*portal.RelayServer) map[string]any
25 - DecodeLeaseID func(string) (string, bool)
26 - IsSecureRequest func(*http.Request, bool) bool
27 - WriteAPIData func(http.ResponseWriter, int, any)
28 - WriteAPIOK func(http.ResponseWriter, int)
29 - WriteAPIError func(http.ResponseWriter, int, string, string)
30 - WriteAPIErrorWithData func(http.ResponseWriter, int, string, string, any)
31 - TrustProxy bool
32 -}
33 -
34 -// Handler routes /admin/* HTTP requests and delegates policy mutations to Service.
35 -type Handler struct {
36 - service *Service
37 - serveAppStatic func(http.ResponseWriter, *http.Request, string, *portal.RelayServer)
38 - listLeases func(*portal.RelayServer) any
39 - stats func(*portal.RelayServer) map[string]any
40 - decodeLeaseID func(string) (string, bool)
41 - isSecureRequest func(*http.Request, bool) bool
42 - writeAPIData func(http.ResponseWriter, int, any)
43 - writeAPIOK func(http.ResponseWriter, int)
44 - writeAPIError func(http.ResponseWriter, int, string, string)
45 - writeAPIErrorWithData func(http.ResponseWriter, int, string, string, any)
46 - trustProxy bool
47 -}
48 -
49 -func NewHandler(cfg HandlerConfig) *Handler {
50 - h := &Handler{
51 - service: cfg.Service,
52 - trustProxy: cfg.TrustProxy,
53 - serveAppStatic: cfg.ServeAppStatic,
54 - listLeases: cfg.ListLeases,
55 - stats: cfg.Stats,
56 - decodeLeaseID: cfg.DecodeLeaseID,
57 - isSecureRequest: cfg.IsSecureRequest,
58 - writeAPIData: cfg.WriteAPIData,
59 - writeAPIOK: cfg.WriteAPIOK,
60 - writeAPIError: cfg.WriteAPIError,
61 - writeAPIErrorWithData: cfg.WriteAPIErrorWithData,
62 - }
63 -
64 - if h.serveAppStatic == nil {
65 - h.serveAppStatic = func(w http.ResponseWriter, r *http.Request, _ string, _ *portal.RelayServer) {
66 - http.NotFound(w, r)
67 - }
68 - }
69 - if h.listLeases == nil {
70 - h.listLeases = func(_ *portal.RelayServer) any { return []any{} }
71 - }
72 - if h.stats == nil {
73 - h.stats = func(serv *portal.RelayServer) map[string]any {
74 - count := 0
75 - if serv != nil && serv.GetLeaseManager() != nil {
76 - count = len(serv.GetLeaseManager().GetAllLeaseEntries())
77 - }
78 - return map[string]any{
79 - "leases_count": count,
80 - }
81 - }
82 - }
83 - if h.decodeLeaseID == nil {
84 - h.decodeLeaseID = decodeLeaseIDFallback
85 - }
86 - if h.isSecureRequest == nil {
87 - h.isSecureRequest = func(r *http.Request, _ bool) bool {
88 - return r != nil && r.TLS != nil
89 - }
90 - }
91 - if h.writeAPIData == nil {
92 - h.writeAPIData = func(w http.ResponseWriter, status int, data any) {
93 - writeDefaultEnvelope(w, status, types.APIEnvelope{OK: true, Data: data})
94 - }
95 - }
96 - if h.writeAPIOK == nil {
97 - h.writeAPIOK = func(w http.ResponseWriter, status int) {
98 - writeDefaultEnvelope(w, status, types.APIEnvelope{OK: true})
99 - }
100 - }
101 - if h.writeAPIError == nil {
102 - h.writeAPIError = func(w http.ResponseWriter, status int, code, message string) {
103 - writeDefaultEnvelope(w, status, types.APIEnvelope{
104 - OK: false,
105 - Error: &types.APIError{
106 - Code: code,
107 - Message: message,
108 - },
109 - })
110 - }
111 - }
112 - if h.writeAPIErrorWithData == nil {
113 - h.writeAPIErrorWithData = func(w http.ResponseWriter, status int, code, message string, data any) {
114 - writeDefaultEnvelope(w, status, types.APIEnvelope{
115 - OK: false,
116 - Data: data,
117 - Error: &types.APIError{
118 - Code: code,
119 - Message: message,
120 - },
121 - })
122 - }
123 - }
124 -
125 - return h
126 -}
127 -
128 -func writeDefaultEnvelope(w http.ResponseWriter, status int, envelope types.APIEnvelope) {
129 - w.Header().Set("Content-Type", "application/json")
130 - w.WriteHeader(status)
131 - if err := json.NewEncoder(w).Encode(envelope); err != nil {
132 - log.Error().Err(err).Msg("[Admin] Failed to encode API envelope")
133 - }
134 -}
135 -
136 -// HandleAdminRequest routes /admin/* requests.
137 -func (h *Handler) HandleAdminRequest(w http.ResponseWriter, r *http.Request, serv *portal.RelayServer) {
138 - if h.service == nil {
139 - h.writeAPIError(w, http.StatusInternalServerError, "admin_service_unavailable", "admin service unavailable")
140 - return
141 - }
142 -
143 - route := strings.Trim(strings.TrimPrefix(r.URL.Path, types.PathAdminPrefix), "/")
144 -
145 - // Public routes (no authentication required)
146 - switch {
147 - case route == "login" && r.Method == http.MethodPost:
148 - h.handleLogin(w, r)
149 - return
150 - case route == "login":
151 - h.serveAppStatic(w, r, "", serv)
152 - return
153 - case route == "logout" && r.Method == http.MethodPost:
154 - h.handleLogout(w, r)
155 - return
156 - case route == "auth/status" && r.Method == http.MethodGet:
157 - h.handleAuthStatus(w, r)
158 - return
159 - }
160 -
161 - // Protected routes - require authentication
162 - if !h.service.IsAuthenticated(r) {
163 - // For page requests (no specific route), show login page.
164 - if route == "" {
165 - h.serveAppStatic(w, r, "", serv)
166 - return
167 - }
168 - // For API requests, return 401 envelope.
169 - h.writeAPIError(w, http.StatusUnauthorized, "unauthorized", "unauthorized")
170 - return
171 - }
172 -
173 - switch {
174 - case route == "":
175 - h.serveAppStatic(w, r, "", serv)
176 - case route == "leases" && r.Method == http.MethodGet:
177 - h.writeAPIData(w, http.StatusOK, h.listLeases(serv))
178 - case route == "leases/banned" && r.Method == http.MethodGet:
179 - h.writeAPIData(w, http.StatusOK, serv.GetLeaseManager().GetBannedLeases())
180 - case route == "stats" && r.Method == http.MethodGet:
181 - h.writeAPIData(w, http.StatusOK, h.stats(serv))
182 - case route == "settings" && r.Method == http.MethodGet:
183 - h.handleGetSettings(w)
184 - case route == "settings/approval-mode":
185 - h.handleApprovalModeRequest(w, r, serv)
186 - case strings.HasPrefix(route, "leases/"):
187 - if !h.handleLeaseActionRouteRequest(w, r, serv, route) {
188 - http.NotFound(w, r)
189 - }
190 - case strings.HasPrefix(route, "ips/") && strings.HasSuffix(route, "/ban"):
191 - h.handleIPBanRequest(w, r, serv, route)
192 - default:
193 - http.NotFound(w, r)
194 - }
195 -}
196 -
197 -func (h *Handler) handleLogin(w http.ResponseWriter, r *http.Request) {
198 - authManager := h.service.GetAuthManager()
199 - clientIP := policy.ExtractClientIP(r, h.trustProxy)
200 -
201 - // Check if IP is locked.
202 - if authManager.IsIPLocked(clientIP) {
203 - remaining := authManager.GetLockRemainingSeconds(clientIP)
204 - h.writeAPIErrorWithData(
205 - w,
206 - http.StatusTooManyRequests,
207 - "auth_locked",
208 - "Too many failed attempts. Please try again later.",
209 - types.AdminLoginResponse{
210 - Locked: true,
211 - RemainingSeconds: remaining,
212 - },
213 - )
214 - return
215 - }
216 -
217 - var req types.AdminLoginRequest
218 - r.Body = http.MaxBytesReader(w, r.Body, 1<<16)
219 - if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
220 - h.writeAPIError(w, http.StatusBadRequest, "invalid_request", "invalid request body")
221 - return
222 - }
223 -
224 - if !authManager.ValidateKey(req.Key) {
225 - // Record failed attempt.
226 - nowLocked := authManager.RecordFailedLogin(clientIP)
227 - log.Warn().Str("ip", clientIP).Bool("now_locked", nowLocked).Msg("[Admin] Failed login attempt")
228 -
229 - response := types.AdminLoginResponse{
230 - Locked: nowLocked,
231 - }
232 - if nowLocked {
233 - response.RemainingSeconds = authManager.GetLockRemainingSeconds(clientIP)
234 - }
235 - h.writeAPIErrorWithData(w, http.StatusUnauthorized, "invalid_key", "Invalid key", response)
236 - return
237 - }
238 -
239 - // Successful login.
240 - authManager.ResetFailedLogin(clientIP)
241 - token := authManager.CreateSession()
242 - secureCookie := h.isSecureRequest(r, h.trustProxy)
243 -
244 - http.SetCookie(w, &http.Cookie{
245 - Name: CookieName,
246 - Value: token,
247 - Path: "/admin",
248 - HttpOnly: true,
249 - Secure: secureCookie,
250 - SameSite: http.SameSiteStrictMode,
251 - MaxAge: 86400, // 24 hours
252 - })
253 -
254 - log.Info().Str("ip", clientIP).Msg("[Admin] Successful login")
255 - h.writeAPIData(w, http.StatusOK, types.AdminLoginResponse{Success: true})
256 -}
257 -
258 -func (h *Handler) handleApprovalModeRequest(w http.ResponseWriter, r *http.Request, serv *portal.RelayServer) {
259 - approveManager := h.service.GetApproveManager()
260 - switch r.Method {
261 - case http.MethodGet:
262 - h.writeAPIData(w, http.StatusOK, types.AdminApprovalModeResponse{
263 - ApprovalMode: string(approveManager.GetApprovalMode()),
264 - })
265 - case http.MethodPost:
266 - var req types.AdminApprovalModeRequest
267 - r.Body = http.MaxBytesReader(w, r.Body, 1<<16)
268 - if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
269 - h.writeAPIError(w, http.StatusBadRequest, "invalid_request", "invalid request body")
270 - return
271 - }
272 - mode := policy.Mode(req.Mode)
273 - if mode != policy.ModeAuto && mode != policy.ModeManual {
274 - h.writeAPIError(w, http.StatusBadRequest, "invalid_mode", "invalid mode (must be 'auto' or 'manual')")
275 - return
276 - }
277 - _ = approveManager.SetApprovalMode(mode) // mode already validated above
278 - h.service.SaveSettings(serv)
279 - log.Info().Str("mode", string(mode)).Msg("[Admin] Approval mode changed")
280 - h.writeAPIData(w, http.StatusOK, types.AdminApprovalModeResponse{
281 - ApprovalMode: string(mode),
282 - })
283 - default:
284 - h.writeAPIError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
285 - }
286 -}
287 -
288 -func (h *Handler) handleLogout(w http.ResponseWriter, r *http.Request) {
289 - authManager := h.service.GetAuthManager()
290 - cookie, err := r.Cookie(CookieName)
291 - if err == nil && cookie.Value != "" {
292 - authManager.DeleteSession(cookie.Value)
293 - }
294 - secureCookie := h.isSecureRequest(r, h.trustProxy)
295 -
296 - http.SetCookie(w, &http.Cookie{
297 - Name: CookieName,
298 - Value: "",
299 - Path: "/admin",
300 - HttpOnly: true,
301 - Secure: secureCookie,
302 - SameSite: http.SameSiteStrictMode,
303 - MaxAge: -1, // Delete cookie.
304 - })
305 -
306 - h.writeAPIOK(w, http.StatusOK)
307 -}
308 -
309 -func (h *Handler) handleAuthStatus(w http.ResponseWriter, r *http.Request) {
310 - h.writeAPIData(w, http.StatusOK, types.AdminAuthStatusResponse{
311 - Authenticated: h.service.IsAuthenticated(r),
312 - AuthEnabled: h.service.AuthEnabled(),
313 - })
314 -}
315 -
316 -func (h *Handler) parseLeaseActionRoute(route string) (leaseID, action string, err error) {
317 - parts := strings.Split(route, "/")
318 - if len(parts) != 3 || parts[0] != "leases" {
319 - return "", "", errors.New("route not found")
320 - }
321 -
322 - action = parts[2]
323 - switch action {
324 - case "ban", "bps", "approve", "deny":
325 - default:
326 - return "", "", errors.New("route not found")
327 - }
328 -
329 - var ok bool
330 - leaseID, ok = h.decodeLeaseID(parts[1])
331 - if !ok {
332 - return "", action, errInvalidLeaseID
333 - }
334 -
335 - return leaseID, action, nil
336 -}
337 -
338 -func (h *Handler) handleLeaseActionRouteRequest(w http.ResponseWriter, r *http.Request, serv *portal.RelayServer, route string) bool {
339 - leaseID, action, err := h.parseLeaseActionRoute(route)
340 - if err != nil {
341 - if errors.Is(err, errInvalidLeaseID) {
342 - h.writeAPIError(w, http.StatusBadRequest, "invalid_lease_id", "invalid lease ID")
343 - return true
344 - }
345 - return false
346 - }
347 -
348 - switch action {
349 - case "ban":
350 - h.handleLeaseBanRequest(w, r, serv, leaseID)
351 - case "bps":
352 - h.handleLeaseBPSRequest(w, r, serv, leaseID)
353 - case "approve":
354 - h.handleLeaseApproveRequest(w, r, serv, leaseID)
355 - case "deny":
356 - h.handleLeaseDenyRequest(w, r, serv, leaseID)
357 - default:
358 - return false
359 - }
360 -
361 - return true
362 -}
363 -
364 -// handleLeaseToggleRequest is a generic helper for lease toggle actions (ban, approve, deny).
365 -func (h *Handler) handleLeaseToggleRequest(
366 - w http.ResponseWriter,
367 - r *http.Request,
368 - serv *portal.RelayServer,
369 - leaseID string,
370 - onPost func(),
371 - onDelete func(),
372 - logMsgPost string,
373 - logMsgDelete string,
374 -) {
375 - if strings.TrimSpace(leaseID) == "" {
376 - h.writeAPIError(w, http.StatusBadRequest, "invalid_lease_id", "invalid lease ID")
377 - return
378 - }
379 -
380 - switch r.Method {
381 - case http.MethodPost:
382 - onPost()
383 - h.service.SaveSettings(serv)
384 - if logMsgPost != "" {
385 - log.Info().Str("lease_id", leaseID).Msg(logMsgPost)
386 - }
387 - h.writeAPIOK(w, http.StatusOK)
388 - case http.MethodDelete:
389 - onDelete()
390 - h.service.SaveSettings(serv)
391 - if logMsgDelete != "" {
392 - log.Info().Str("lease_id", leaseID).Msg(logMsgDelete)
393 - }
394 - h.writeAPIOK(w, http.StatusOK)
395 - default:
396 - h.writeAPIError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
397 - }
398 -}
399 -
400 -func (h *Handler) handleLeaseBanRequest(w http.ResponseWriter, r *http.Request, serv *portal.RelayServer, leaseID string) {
401 - h.handleLeaseToggleRequest(
402 - w, r, serv, leaseID,
403 - func() { serv.GetLeaseManager().BanLease(leaseID) },
404 - func() { serv.GetLeaseManager().UnbanLease(leaseID) },
405 - "", "",
406 - )
407 -}
408 -
409 -func (h *Handler) handleGetSettings(w http.ResponseWriter) {
410 - approveManager := h.service.GetApproveManager()
411 - h.writeAPIData(w, http.StatusOK, types.AdminSettingsResponse{
412 - ApprovalMode: string(approveManager.GetApprovalMode()),
413 - ApprovedLeases: approveManager.GetApprovedLeases(),
414 - DeniedLeases: approveManager.GetDeniedLeases(),
415 - })
416 -}
417 -
418 -func (h *Handler) handleLeaseApproveRequest(w http.ResponseWriter, r *http.Request, serv *portal.RelayServer, leaseID string) {
419 - approveManager := h.service.GetApproveManager()
420 - h.handleLeaseToggleRequest(
421 - w, r, serv, leaseID,
422 - func() {
423 - approveManager.ApproveLease(leaseID)
424 - approveManager.UndenyLease(leaseID) // Remove from denied if exists.
425 - },
426 - func() { approveManager.RevokeLease(leaseID) },
427 - "[Admin] Lease approved",
428 - "[Admin] Lease approval revoked",
429 - )
430 -}
431 -
432 -func (h *Handler) handleLeaseDenyRequest(w http.ResponseWriter, r *http.Request, serv *portal.RelayServer, leaseID string) {
433 - approveManager := h.service.GetApproveManager()
434 - h.handleLeaseToggleRequest(
435 - w, r, serv, leaseID,
436 - func() { approveManager.DenyLease(leaseID) },
437 - func() { approveManager.UndenyLease(leaseID) },
438 - "[Admin] Lease denied",
439 - "[Admin] Lease denial removed",
440 - )
441 -}
442 -
443 -func (h *Handler) handleLeaseBPSRequest(w http.ResponseWriter, r *http.Request, serv *portal.RelayServer, leaseID string) {
444 - if strings.TrimSpace(leaseID) == "" {
445 - h.writeAPIError(w, http.StatusBadRequest, "invalid_lease_id", "invalid lease ID")
446 - return
447 - }
448 -
449 - bpsManager := h.service.GetBPSManager()
450 - switch r.Method {
451 - case http.MethodPost:
452 - var req types.AdminBPSRequest
453 - r.Body = http.MaxBytesReader(w, r.Body, 1<<16)
454 - if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
455 - h.writeAPIError(w, http.StatusBadRequest, "invalid_request", "invalid request body")
456 - return
457 - }
458 - if bpsManager == nil {
459 - h.writeAPIError(w, http.StatusInternalServerError, "bps_manager_unavailable", "bps manager not initialized")
460 - return
461 - }
462 - oldBPS := bpsManager.GetBPSLimit(leaseID)
463 - bpsManager.SetBPSLimit(leaseID, req.BPS)
464 - log.Info().
465 - Str("lease_id", leaseID).
466 - Int64("old_bps", oldBPS).
467 - Int64("new_bps", req.BPS).
468 - Msg("[Admin] BPS limit updated")
469 - h.service.SaveSettings(serv)
470 - h.writeAPIOK(w, http.StatusOK)
471 - case http.MethodDelete:
472 - if bpsManager == nil {
473 - h.writeAPIError(w, http.StatusInternalServerError, "bps_manager_unavailable", "bps manager not initialized")
474 - return
475 - }
476 - oldBPS := bpsManager.GetBPSLimit(leaseID)
477 - bpsManager.SetBPSLimit(leaseID, 0)
478 - log.Info().
479 - Str("lease_id", leaseID).
480 - Int64("old_bps", oldBPS).
481 - Msg("[Admin] BPS limit removed (now unlimited)")
482 - h.service.SaveSettings(serv)
483 - h.writeAPIOK(w, http.StatusOK)
484 - default:
485 - h.writeAPIError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
486 - }
487 -}
488 -
489 -func (h *Handler) handleIPBanRequest(w http.ResponseWriter, r *http.Request, serv *portal.RelayServer, route string) {
490 - // Route format: ips/{ip}/ban.
491 - parts := strings.Split(route, "/")
492 - if len(parts) != 3 {
493 - http.NotFound(w, r)
494 - return
495 - }
496 -
497 - ip := parts[1]
498 - if ip == "" {
499 - h.writeAPIError(w, http.StatusBadRequest, "invalid_ip", "invalid IP address")
500 - return
501 - }
502 - if net.ParseIP(ip) == nil {
503 - h.writeAPIError(w, http.StatusBadRequest, "invalid_ip", "invalid IP address")
504 - return
505 - }
506 -
507 - ipManager := h.service.GetIPManager()
508 - if ipManager == nil {
509 - h.writeAPIError(w, http.StatusInternalServerError, "ip_manager_unavailable", "ip manager not initialized")
510 - return
511 - }
512 -
513 - switch r.Method {
514 - case http.MethodPost:
515 - ipManager.BanIP(ip)
516 - h.service.SaveSettings(serv)
517 - log.Info().Str("ip", ip).Msg("[Admin] IP banned")
518 - h.writeAPIOK(w, http.StatusOK)
519 - case http.MethodDelete:
520 - ipManager.UnbanIP(ip)
521 - h.service.SaveSettings(serv)
522 - log.Info().Str("ip", ip).Msg("[Admin] IP unbanned")
523 - h.writeAPIOK(w, http.StatusOK)
524 - default:
525 - h.writeAPIError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
526 - }
527 -}
528 -
529 -func decodeLeaseIDFallback(encoded string) (string, bool) {
530 - idBytes, err := base64.URLEncoding.DecodeString(encoded)
531 - if err != nil {
532 - idBytes, err = base64.RawURLEncoding.DecodeString(encoded)
533 - if err != nil {
534 - return "", false
535 - }
536 - }
537 - return string(idBytes), true
538 -}
portal/admin/handler_test.go deleted
-176
@@ -1,176 +0,0 @@
1 -package admin
2 -
3 -import (
4 - "encoding/json"
5 - "net/http"
6 - "net/http/httptest"
7 - "strings"
8 - "testing"
9 -
10 - "gosuda.org/portal/portal"
11 - "gosuda.org/portal/portal/policy"
12 - "gosuda.org/portal/types"
13 -)
14 -
15 -func TestHandleAdminRequestLoginSuccessSetsSessionCookie(t *testing.T) {
16 - service := NewService(policy.NewAuthenticator("test-secret"))
17 - handler := newTestHandler(t, service, true, nil)
18 -
19 - req := httptest.NewRequest(http.MethodPost, types.PathAdminPrefix+"/login", strings.NewReader(`{"key":"test-secret"}`))
20 - req.RemoteAddr = "203.0.113.10:1234"
21 - rec := httptest.NewRecorder()
22 -
23 - handler.HandleAdminRequest(rec, req, nil)
24 -
25 - if rec.Code != http.StatusOK {
26 - t.Fatalf("expected status %d, got %d", http.StatusOK, rec.Code)
27 - }
28 - envelope := decodeEnvelope(t, rec)
29 - if !envelope.OK {
30 - t.Fatalf("expected OK response, got %+v", envelope)
31 - }
32 -
33 - cookies := rec.Result().Cookies()
34 - if len(cookies) == 0 {
35 - t.Fatalf("expected admin session cookie")
36 - }
37 - cookie := cookies[0]
38 - if cookie.Name != CookieName {
39 - t.Fatalf("expected cookie name %q, got %q", CookieName, cookie.Name)
40 - }
41 - if cookie.Path != "/admin" {
42 - t.Fatalf("expected cookie path /admin, got %q", cookie.Path)
43 - }
44 - if !cookie.HttpOnly {
45 - t.Fatalf("expected HttpOnly cookie")
46 - }
47 - if !cookie.Secure {
48 - t.Fatalf("expected Secure cookie")
49 - }
50 - if cookie.MaxAge != 86400 {
51 - t.Fatalf("expected MaxAge 86400, got %d", cookie.MaxAge)
52 - }
53 -}
54 -
55 -func TestHandleAdminRequestProtectedRouteUnauthorized(t *testing.T) {
56 - service := NewService(policy.NewAuthenticator("test-secret"))
57 - handler := newTestHandler(t, service, false, func(_ *portal.RelayServer) any {
58 - t.Fatalf("list leases should not be called for unauthorized request")
59 - return nil
60 - })
61 -
62 - req := httptest.NewRequest(http.MethodGet, types.PathAdminPrefix+"/leases", nil)
63 - rec := httptest.NewRecorder()
64 -
65 - handler.HandleAdminRequest(rec, req, nil)
66 -
67 - if rec.Code != http.StatusUnauthorized {
68 - t.Fatalf("expected status %d, got %d", http.StatusUnauthorized, rec.Code)
69 - }
70 - envelope := decodeEnvelope(t, rec)
71 - if envelope.OK || envelope.Error == nil || envelope.Error.Code != "unauthorized" {
72 - t.Fatalf("expected unauthorized API error, got %+v", envelope)
73 - }
74 -}
75 -
76 -func TestHandleAdminRequestApprovalModeInvalidMode(t *testing.T) {
77 - service := NewService(policy.NewAuthenticator("test-secret"))
78 - handler := newTestHandler(t, service, false, nil)
79 - token := service.GetAuthManager().CreateSession()
80 -
81 - req := httptest.NewRequest(http.MethodPost, types.PathAdminPrefix+"/settings/approval-mode", strings.NewReader(`{"mode":"invalid"}`))
82 - req.AddCookie(&http.Cookie{Name: CookieName, Value: token})
83 - rec := httptest.NewRecorder()
84 -
85 - handler.HandleAdminRequest(rec, req, nil)
86 -
87 - if rec.Code != http.StatusBadRequest {
88 - t.Fatalf("expected status %d, got %d", http.StatusBadRequest, rec.Code)
89 - }
90 - envelope := decodeEnvelope(t, rec)
91 - if envelope.OK || envelope.Error == nil || envelope.Error.Code != "invalid_mode" {
92 - t.Fatalf("expected invalid_mode API error, got %+v", envelope)
93 - }
94 -}
95 -
96 -func TestHandleAdminRequestLeaseActionInvalidLeaseID(t *testing.T) {
97 - service := NewService(policy.NewAuthenticator("test-secret"))
98 - handler := newTestHandler(t, service, false, nil)
99 - token := service.GetAuthManager().CreateSession()
100 -
101 - req := httptest.NewRequest(http.MethodPost, types.PathAdminPrefix+"/leases/not!base64/ban", nil)
102 - req.AddCookie(&http.Cookie{Name: CookieName, Value: token})
103 - rec := httptest.NewRecorder()
104 -
105 - handler.HandleAdminRequest(rec, req, nil)
106 -
107 - if rec.Code != http.StatusBadRequest {
108 - t.Fatalf("expected status %d, got %d", http.StatusBadRequest, rec.Code)
109 - }
110 - envelope := decodeEnvelope(t, rec)
111 - if envelope.OK || envelope.Error == nil || envelope.Error.Code != "invalid_lease_id" {
112 - t.Fatalf("expected invalid_lease_id API error, got %+v", envelope)
113 - }
114 -}
115 -
116 -func newTestHandler(t *testing.T, service *Service, secure bool, listLeases func(*portal.RelayServer) any) *Handler {
117 - t.Helper()
118 -
119 - if listLeases == nil {
120 - listLeases = func(_ *portal.RelayServer) any { return []any{} }
121 - }
122 -
123 - return NewHandler(HandlerConfig{
124 - Service: service,
125 - TrustProxy: false,
126 - ServeAppStatic: func(w http.ResponseWriter, _ *http.Request, _ string, _ *portal.RelayServer) {
127 - w.WriteHeader(http.StatusOK)
128 - },
129 - ListLeases: listLeases,
130 - IsSecureRequest: func(_ *http.Request, _ bool) bool {
131 - return secure
132 - },
133 - WriteAPIData: writeTestAPIData,
134 - WriteAPIOK: writeTestAPIOK,
135 - WriteAPIError: writeTestAPIError,
136 - WriteAPIErrorWithData: writeTestAPIErrorWithData,
137 - })
138 -}
139 -
140 -func decodeEnvelope(t *testing.T, rec *httptest.ResponseRecorder) types.APIEnvelope {
141 - t.Helper()
142 - var envelope types.APIEnvelope
143 - if err := json.Unmarshal(rec.Body.Bytes(), &envelope); err != nil {
144 - t.Fatalf("failed to decode API envelope: %v", err)
145 - }
146 - return envelope
147 -}
148 -
149 -func writeTestAPIData(w http.ResponseWriter, status int, data any) {
150 - w.Header().Set("Content-Type", "application/json")
151 - w.WriteHeader(status)
152 - _ = json.NewEncoder(w).Encode(types.APIEnvelope{OK: true, Data: data})
153 -}
154 -
155 -func writeTestAPIOK(w http.ResponseWriter, status int) {
156 - w.Header().Set("Content-Type", "application/json")
157 - w.WriteHeader(status)
158 - _ = json.NewEncoder(w).Encode(types.APIEnvelope{OK: true})
159 -}
160 -
161 -func writeTestAPIError(w http.ResponseWriter, status int, code, message string) {
162 - writeTestAPIErrorWithData(w, status, code, message, nil)
163 -}
164 -
165 -func writeTestAPIErrorWithData(w http.ResponseWriter, status int, code, message string, data any) {
166 - w.Header().Set("Content-Type", "application/json")
167 - w.WriteHeader(status)
168 - _ = json.NewEncoder(w).Encode(types.APIEnvelope{
169 - OK: false,
170 - Data: data,
171 - Error: &types.APIError{
172 - Code: code,
173 - Message: message,
174 - },
175 - })
176 -}
portal/admin/service.go deleted
-223
@@ -1,223 +0,0 @@
1 -package admin
2 -
3 -import (
4 - "encoding/json"
5 - "net/http"
6 - "os"
7 - "path/filepath"
8 - "sync"
9 -
10 - "github.com/rs/zerolog/log"
11 -
12 - "gosuda.org/portal/portal"
13 - "gosuda.org/portal/portal/policy"
14 -)
15 -
16 -const CookieName = "portal_admin"
17 -
18 -// Service manages admin policy state and persistence.
19 -type Service struct {
20 - approveManager *policy.Approver
21 - bpsManager *policy.RateLimiter
22 - ipManager *policy.IPFilter
23 - authManager *policy.Authenticator
24 - settingsPath string
25 - settingsMu sync.Mutex
26 -}
27 -
28 -// settings stores persistent admin configuration.
29 -type settings struct {
30 - BannedLeases []string `json:"banned_leases"`
31 - BPSLimits map[string]int64 `json:"bps_limits"`
32 - ApprovalMode policy.Mode `json:"approval_mode"`
33 - ApprovedLeases []string `json:"approved_leases,omitempty"`
34 - DeniedLeases []string `json:"denied_leases,omitempty"`
35 - BannedIPs []string `json:"banned_ips,omitempty"`
36 -}
37 -
38 -func NewService(authManager *policy.Authenticator) *Service {
39 - bpsManager := policy.NewRateLimiter()
40 -
41 - return &Service{
42 - settingsPath: "admin_settings.json",
43 - approveManager: policy.NewApprover(),
44 - bpsManager: bpsManager,
45 - ipManager: policy.NewIPFilter(),
46 - authManager: authManager,
47 - }
48 -}
49 -
50 -func (s *Service) authUnavailable() bool {
51 - return s == nil || s.authManager == nil || !s.authManager.HasSecretKey()
52 -}
53 -
54 -func (s *Service) GetApproveManager() *policy.Approver {
55 - if s == nil {
56 - return nil
57 - }
58 - return s.approveManager
59 -}
60 -
61 -func (s *Service) GetBPSManager() *policy.RateLimiter {
62 - if s == nil {
63 - return nil
64 - }
65 - return s.bpsManager
66 -}
67 -
68 -func (s *Service) GetIPManager() *policy.IPFilter {
69 - if s == nil {
70 - return nil
71 - }
72 - return s.ipManager
73 -}
74 -
75 -func (s *Service) GetAuthManager() *policy.Authenticator {
76 - if s == nil {
77 - return nil
78 - }
79 - return s.authManager
80 -}
81 -
82 -func (s *Service) SetSettingsPath(path string) {
83 - if s == nil {
84 - return
85 - }
86 - s.settingsMu.Lock()
87 - defer s.settingsMu.Unlock()
88 - s.settingsPath = path
89 -}
90 -
91 -func (s *Service) SaveSettings(serv *portal.RelayServer) {
92 - if s == nil || serv == nil {
93 - return
94 - }
95 -
96 - s.settingsMu.Lock()
97 - defer s.settingsMu.Unlock()
98 -
99 - lm := serv.GetLeaseManager()
100 - banned := lm.GetBannedLeases()
101 -
102 - bpsLimits := map[string]int64{}
103 - if s.bpsManager != nil {
104 - bpsLimits = s.bpsManager.GetAllBPSLimits()
105 - }
106 -
107 - var bannedIPs []string
108 - if s.ipManager != nil {
109 - bannedIPs = s.ipManager.GetBannedIPs()
110 - }
111 -
112 - payload := settings{
113 - BannedLeases: banned,
114 - BPSLimits: bpsLimits,
115 - ApprovalMode: s.approveManager.GetApprovalMode(),
116 - ApprovedLeases: s.approveManager.GetApprovedLeases(),
117 - DeniedLeases: s.approveManager.GetDeniedLeases(),
118 - BannedIPs: bannedIPs,
119 - }
120 -
121 - data, err := json.MarshalIndent(payload, "", " ")
122 - if err != nil {
123 - log.Error().Err(err).Msg("[Admin] Failed to marshal admin settings")
124 - return
125 - }
126 -
127 - dir := filepath.Dir(s.settingsPath)
128 - if dir != "" && dir != "." {
129 - if err := os.MkdirAll(dir, 0755); err != nil {
130 - log.Error().Err(err).Msg("[Admin] Failed to create settings directory")
131 - return
132 - }
133 - }
134 -
135 - if err := os.WriteFile(s.settingsPath, data, 0600); err != nil {
136 - log.Error().Err(err).Msg("[Admin] Failed to save admin settings")
137 - return
138 - }
139 -
140 - log.Debug().Str("path", s.settingsPath).Msg("[Admin] Saved admin settings")
141 -}
142 -
143 -func (s *Service) LoadSettings(serv *portal.RelayServer) {
144 - if s == nil || serv == nil {
145 - return
146 - }
147 -
148 - s.settingsMu.Lock()
149 - defer s.settingsMu.Unlock()
150 -
151 - data, err := os.ReadFile(s.settingsPath)
152 - if err != nil {
153 - if os.IsNotExist(err) {
154 - log.Debug().Msg("[Admin] No admin settings file found, starting fresh")
155 - return
156 - }
157 - log.Error().Err(err).Msg("[Admin] Failed to read admin settings")
158 - return
159 - }
160 -
161 - var payload settings
162 - if err := json.Unmarshal(data, &payload); err != nil {
163 - log.Error().Err(err).Msg("[Admin] Failed to parse admin settings")
164 - return
165 - }
166 -
167 - lm := serv.GetLeaseManager()
168 -
169 - for _, leaseID := range payload.BannedLeases {
170 - lm.BanLease(leaseID)
171 - }
172 -
173 - for leaseID, bps := range payload.BPSLimits {
174 - if s.bpsManager != nil {
175 - s.bpsManager.SetBPSLimit(leaseID, bps)
176 - }
177 - }
178 -
179 - if payload.ApprovalMode != "" {
180 - if err := s.approveManager.SetApprovalMode(payload.ApprovalMode); err != nil {
181 - log.Warn().Err(err).Str("mode", string(payload.ApprovalMode)).Msg("[Admin] Ignoring invalid approval mode from settings")
182 - }
183 - }
184 -
185 - for _, leaseID := range payload.ApprovedLeases {
186 - s.approveManager.ApproveLease(leaseID)
187 - }
188 -
189 - for _, leaseID := range payload.DeniedLeases {
190 - s.approveManager.DenyLease(leaseID)
191 - }
192 -
193 - if s.ipManager != nil && len(payload.BannedIPs) > 0 {
194 - s.ipManager.SetBannedIPs(payload.BannedIPs)
195 - }
196 -
197 - log.Info().
198 - Int("banned_count", len(payload.BannedLeases)).
199 - Int("bps_limits_count", len(payload.BPSLimits)).
200 - Str("approval_mode", string(s.approveManager.GetApprovalMode())).
201 - Int("approved_count", len(payload.ApprovedLeases)).
202 - Int("denied_count", len(payload.DeniedLeases)).
203 - Int("banned_ips_count", len(payload.BannedIPs)).
204 - Msg("[Admin] Loaded admin settings")
205 -}
206 -
207 -// IsAuthenticated checks if the request has a valid admin session.
208 -func (s *Service) IsAuthenticated(r *http.Request) bool {
209 - if s.authUnavailable() {
210 - return false
211 - }
212 -
213 - cookie, err := r.Cookie(CookieName)
214 - if err != nil {
215 - return false
216 - }
217 -
218 - return s.authManager.ValidateSession(cookie.Value)
219 -}
220 -
221 -func (s *Service) AuthEnabled() bool {
222 - return !s.authUnavailable()
223 -}
portal/admin/service_test.go deleted
-75
@@ -1,75 +0,0 @@
1 -package admin
2 -
3 -import (
4 - "context"
5 - "path/filepath"
6 - "slices"
7 - "testing"
8 -
9 - "gosuda.org/portal/portal"
10 - "gosuda.org/portal/portal/policy"
11 -)
12 -
13 -func TestServiceSaveLoadSettingsRoundTrip(t *testing.T) {
14 - service := NewService(policy.NewAuthenticator("test-secret"))
15 - settingsPath := filepath.Join(t.TempDir(), "admin_settings.json")
16 - service.SetSettingsPath(settingsPath)
17 -
18 - sourceServer := mustNewTestRelayServer(t)
19 - sourceServer.GetLeaseManager().BanLease("lease-ban")
20 - service.GetBPSManager().SetBPSLimit("lease-bps", 4096)
21 - if err := service.GetApproveManager().SetApprovalMode(policy.ModeManual); err != nil {
22 - t.Fatalf("SetApprovalMode: %v", err)
23 - }
24 - service.GetApproveManager().ApproveLease("lease-approved")
25 - service.GetApproveManager().DenyLease("lease-denied")
26 - service.GetIPManager().BanIP("203.0.113.10")
27 -
28 - service.SaveSettings(sourceServer)
29 -
30 - targetServer := mustNewTestRelayServer(t)
31 - loaded := NewService(policy.NewAuthenticator("test-secret"))
32 - loaded.SetSettingsPath(settingsPath)
33 - loaded.LoadSettings(targetServer)
34 -
35 - if !contains(targetServer.GetLeaseManager().GetBannedLeases(), "lease-ban") {
36 - t.Fatalf("expected banned lease to be restored")
37 - }
38 - if got := loaded.GetBPSManager().GetBPSLimit("lease-bps"); got != 4096 {
39 - t.Fatalf("expected BPS limit 4096, got %d", got)
40 - }
41 - if loaded.GetApproveManager().GetApprovalMode() != policy.ModeManual {
42 - t.Fatalf("expected approval mode manual")
43 - }
44 - if !loaded.GetApproveManager().IsLeaseApproved("lease-approved") {
45 - t.Fatalf("expected approved lease to be restored")
46 - }
47 - if !loaded.GetApproveManager().IsLeaseDenied("lease-denied") {
48 - t.Fatalf("expected denied lease to be restored")
49 - }
50 - if !loaded.GetIPManager().IsIPBanned("203.0.113.10") {
51 - t.Fatalf("expected banned IP to be restored")
52 - }
53 -}
54 -
55 -func mustNewTestRelayServer(t *testing.T) *portal.RelayServer {
56 - t.Helper()
57 -
58 - server, err := portal.NewRelayServer(
59 - context.Background(),
60 - nil,
61 - ":0",
62 - "localhost",
63 - t.TempDir(),
64 - "",
65 - )
66 - if err != nil {
67 - t.Fatalf("create relay server: %v", err)
68 - }
69 - t.Cleanup(server.Stop)
70 - return server
71 -}
72 -
73 -func contains(values []string, target string) bool {
74 - return slices.Contains(values, target)
75 -}
portal/api.go new
+89
@@ -0,0 +1,89 @@
1 +package portal
2 +
3 +import (
4 + "encoding/json"
5 + "net/http"
6 + "time"
7 +)
8 +
9 +const (
10 + HeaderReverseToken = "X-Portal-Token"
11 + MarkerKeepalive = byte(0x00)
12 + MarkerTLSStart = byte(0x02)
13 +)
14 +
15 +type APIEnvelope struct {
16 + OK bool `json:"ok"`
17 + Data any `json:"data,omitempty"`
18 + Error *APIError `json:"error,omitempty"`
19 +}
20 +
21 +type APIError struct {
22 + Code string `json:"code"`
23 + Message string `json:"message"`
24 +}
25 +
26 +type LeaseMetadata struct {
27 + Description string `json:"description,omitempty"`
28 + Tags []string `json:"tags,omitempty"`
29 + Owner string `json:"owner,omitempty"`
30 + Thumbnail string `json:"thumbnail,omitempty"`
31 + Hide bool `json:"hide,omitempty"`
32 +}
33 +
34 +type RegisterRequest struct {
35 + Name string `json:"name"`
36 + Hostnames []string `json:"hostnames,omitempty"`
37 + Metadata LeaseMetadata `json:"metadata,omitempty"`
38 + ReverseToken string `json:"reverse_token"`
39 + TLS bool `json:"tls"`
40 + TTLSeconds int `json:"ttl_seconds,omitempty"`
41 +}
42 +
43 +type RegisterResponse struct {
44 + LeaseID string `json:"lease_id"`
45 + Hostnames []string `json:"hostnames"`
46 + Metadata LeaseMetadata `json:"metadata,omitempty"`
47 + ExpiresAt time.Time `json:"expires_at"`
48 + ConnectURL string `json:"connect_url"`
49 +}
50 +
51 +type RenewRequest struct {
52 + LeaseID string `json:"lease_id"`
53 + ReverseToken string `json:"reverse_token"`
54 + TTLSeconds int `json:"ttl_seconds,omitempty"`
55 +}
56 +
57 +type RenewResponse struct {
58 + LeaseID string `json:"lease_id"`
59 + ExpiresAt time.Time `json:"expires_at"`
60 +}
61 +
62 +type UnregisterRequest struct {
63 + LeaseID string `json:"lease_id"`
64 + ReverseToken string `json:"reverse_token"`
65 +}
66 +
67 +type DomainResponse struct {
68 + RootHost string `json:"root_host"`
69 + SuggestedHostname string `json:"suggested_hostname"`
70 +}
71 +
72 +func writeAPIData(w http.ResponseWriter, status int, data any) {
73 + w.Header().Set("Content-Type", "application/json")
74 + w.WriteHeader(status)
75 + _ = json.NewEncoder(w).Encode(APIEnvelope{OK: true, Data: data})
76 +}
77 +
78 +func writeAPIOK(w http.ResponseWriter, status int) {
79 + writeAPIData(w, status, map[string]any{})
80 +}
81 +
82 +func writeAPIError(w http.ResponseWriter, status int, code, message string) {
83 + w.Header().Set("Content-Type", "application/json")
84 + w.WriteHeader(status)
85 + _ = json.NewEncoder(w).Encode(APIEnvelope{
86 + OK: false,
87 + Error: &APIError{Code: code, Message: message},
88 + })
89 +}
portal/broker.go new
+316
@@ -0,0 +1,316 @@
1 +package portal
2 +
3 +import (
4 + "context"
5 + "errors"
6 + "fmt"
7 + "net"
8 + "sync"
9 + "time"
10 +)
11 +
12 +var (
13 + errLeaseDropped = errors.New("lease dropped")
14 + errLeaseStopped = errors.New("lease stopped")
15 + errBrokerFull = errors.New("broker ready queue full")
16 +)
17 +
18 +type brokerState int
19 +
20 +const (
21 + brokerStateActive brokerState = iota
22 + brokerStateDropped
23 + brokerStateStopped
24 +)
25 +
26 +type leaseBroker struct {
27 + leaseID string
28 + idleInterval time.Duration
29 + readyLimit int
30 +
31 + mu sync.Mutex
32 + ready []*reverseSession
33 + state brokerState
34 + notify chan struct{}
35 +}
36 +
37 +func newLeaseBroker(leaseID string, idleInterval time.Duration, readyLimit int) *leaseBroker {
38 + return &leaseBroker{
39 + leaseID: leaseID,
40 + idleInterval: idleInterval,
41 + readyLimit: readyLimit,
42 + notify: make(chan struct{}, 1),
43 + }
44 +}
45 +
46 +func (b *leaseBroker) Offer(session *reverseSession) error {
47 + if session == nil {
48 + return errors.New("reverse session is required")
49 + }
50 +
51 + b.mu.Lock()
52 + defer b.mu.Unlock()
53 +
54 + switch b.state {
55 + case brokerStateDropped:
56 + return errLeaseDropped
57 + case brokerStateStopped:
58 + return errLeaseStopped
59 + }
60 +
61 + if b.readyLimit > 0 && len(b.ready) >= b.readyLimit {
62 + return errBrokerFull
63 + }
64 +
65 + session.StartIdle()
66 + b.ready = append(b.ready, session)
67 + b.signalLocked()
68 + go b.watchSession(session)
69 + return nil
70 +}
71 +
72 +func (b *leaseBroker) Claim(ctx context.Context) (*reverseSession, error) {
73 + for {
74 + b.mu.Lock()
75 + switch b.state {
76 + case brokerStateDropped:
77 + b.mu.Unlock()
78 + return nil, errLeaseDropped
79 + case brokerStateStopped:
80 + b.mu.Unlock()
81 + return nil, errLeaseStopped
82 + }
83 +
84 + if len(b.ready) > 0 {
85 + session := b.ready[0]
86 + b.ready = b.ready[1:]
87 + b.mu.Unlock()
88 +
89 + if session.IsClosed() {
90 + continue
91 + }
92 + if err := session.Activate(); err != nil {
93 + _ = session.Close()
94 + continue
95 + }
96 + return session, nil
97 + }
98 + b.mu.Unlock()
99 +
100 + select {
101 + case <-ctx.Done():
102 + return nil, ctx.Err()
103 + case <-b.notify:
104 + }
105 + }
106 +}
107 +
108 +func (b *leaseBroker) Drop() {
109 + b.transition(brokerStateDropped)
110 +}
111 +
112 +func (b *leaseBroker) Reset() {
113 + b.mu.Lock()
114 + defer b.mu.Unlock()
115 + if b.state == brokerStateDropped {
116 + b.state = brokerStateActive
117 + b.signalLocked()
118 + }
119 +}
120 +
121 +func (b *leaseBroker) Stop() {
122 + b.transition(brokerStateStopped)
123 +}
124 +
125 +func (b *leaseBroker) ReadyCount() int {
126 + b.mu.Lock()
127 + defer b.mu.Unlock()
128 + return len(b.ready)
129 +}
130 +
131 +func (b *leaseBroker) transition(state brokerState) {
132 + b.mu.Lock()
133 + sessions := b.ready
134 + b.ready = nil
135 + b.state = state
136 + b.signalLocked()
137 + b.mu.Unlock()
138 +
139 + for _, session := range sessions {
140 + _ = session.Close()
141 + }
142 +}
143 +
144 +func (b *leaseBroker) watchSession(session *reverseSession) {
145 + <-session.Done()
146 + b.mu.Lock()
147 + defer b.mu.Unlock()
148 + for i := range b.ready {
149 + if b.ready[i] == session {
150 + b.ready = append(b.ready[:i], b.ready[i+1:]...)
151 + break
152 + }
153 + }
154 + b.signalLocked()
155 +}
156 +
157 +func (b *leaseBroker) signalLocked() {
158 + select {
159 + case b.notify <- struct{}{}:
160 + default:
161 + }
162 +}
163 +
164 +type reverseSessionState int
165 +
166 +const (
167 + reverseSessionAdmitted reverseSessionState = iota
168 + reverseSessionIdle
169 + reverseSessionClaimed
170 + reverseSessionClosed
171 +)
172 +
173 +type reverseSession struct {
174 + conn net.Conn
175 + idleInterval time.Duration
176 +
177 + mu sync.Mutex
178 + state reverseSessionState
179 + keepaliveStop chan struct{}
180 + keepaliveDone chan struct{}
181 + done chan struct{}
182 + closeOnce sync.Once
183 +}
184 +
185 +func newReverseSession(conn net.Conn, idleInterval time.Duration) *reverseSession {
186 + return &reverseSession{
187 + conn: conn,
188 + idleInterval: idleInterval,
189 + state: reverseSessionAdmitted,
190 + done: make(chan struct{}),
191 + }
192 +}
193 +
194 +func (s *reverseSession) Conn() net.Conn {
195 + return s.conn
196 +}
197 +
198 +func (s *reverseSession) Done() <-chan struct{} {
199 + return s.done
200 +}
201 +
202 +func (s *reverseSession) IsClosed() bool {
203 + select {
204 + case <-s.done:
205 + return true
206 + default:
207 + return false
208 + }
209 +}
210 +
211 +func (s *reverseSession) StartIdle() {
212 + s.mu.Lock()
213 + if s.state != reverseSessionAdmitted {
214 + s.mu.Unlock()
215 + return
216 + }
217 + s.state = reverseSessionIdle
218 + stop := make(chan struct{})
219 + done := make(chan struct{})
220 + s.keepaliveStop = stop
221 + s.keepaliveDone = done
222 + s.mu.Unlock()
223 +
224 + go s.runKeepalive(stop, done)
225 +}
226 +
227 +func (s *reverseSession) Activate() error {
228 + s.mu.Lock()
229 + if s.state != reverseSessionIdle {
230 + state := s.state
231 + s.mu.Unlock()
232 + return fmt.Errorf("session not idle: %d", state)
233 + }
234 + stop := s.keepaliveStop
235 + done := s.keepaliveDone
236 + s.keepaliveStop = nil
237 + s.keepaliveDone = nil
238 + s.state = reverseSessionClaimed
239 + s.mu.Unlock()
240 +
241 + if stop != nil {
242 + close(stop)
243 + }
244 + if done != nil {
245 + <-done
246 + }
247 +
248 + s.mu.Lock()
249 + defer s.mu.Unlock()
250 + if s.state == reverseSessionClosed {
251 + return net.ErrClosed
252 + }
253 + _ = s.conn.SetWriteDeadline(time.Now().Add(defaultSessionWriteLimit))
254 + _, err := s.conn.Write([]byte{MarkerTLSStart})
255 + _ = s.conn.SetWriteDeadline(time.Time{})
256 + if err != nil {
257 + go s.Close()
258 + }
259 + return err
260 +}
261 +
262 +func (s *reverseSession) Close() error {
263 + var err error
264 + s.closeOnce.Do(func() {
265 + s.mu.Lock()
266 + stop := s.keepaliveStop
267 + done := s.keepaliveDone
268 + s.keepaliveStop = nil
269 + s.keepaliveDone = nil
270 + s.state = reverseSessionClosed
271 + conn := s.conn
272 + s.mu.Unlock()
273 +
274 + if stop != nil {
275 + close(stop)
276 + }
277 + if done != nil {
278 + <-done
279 + }
280 +
281 + err = conn.Close()
282 + close(s.done)
283 + })
284 + return err
285 +}
286 +
287 +func (s *reverseSession) runKeepalive(stop <-chan struct{}, done chan<- struct{}) {
288 + defer close(done)
289 +
290 + ticker := time.NewTicker(s.idleInterval)
291 + defer ticker.Stop()
292 +
293 + for {
294 + select {
295 + case <-stop:
296 + return
297 + case <-s.done:
298 + return
299 + case <-ticker.C:
300 + }
301 +
302 + s.mu.Lock()
303 + if s.state != reverseSessionIdle {
304 + s.mu.Unlock()
305 + return
306 + }
307 + _ = s.conn.SetWriteDeadline(time.Now().Add(defaultSessionWriteLimit))
308 + _, err := s.conn.Write([]byte{MarkerKeepalive})
309 + _ = s.conn.SetWriteDeadline(time.Time{})
310 + s.mu.Unlock()
311 + if err != nil {
312 + _ = s.Close()
313 + return
314 + }
315 + }
316 +}
portal/broker_test.go new
+77
@@ -0,0 +1,77 @@
1 +package portal
2 +
3 +import (
4 + "context"
5 + "io"
6 + "net"
7 + "testing"
8 + "time"
9 +)
10 +
11 +func TestLeaseBrokerClaimActivatesTLSMarker(t *testing.T) {
12 + t.Parallel()
13 +
14 + serverConn, clientConn := net.Pipe()
15 + t.Cleanup(func() {
16 + _ = serverConn.Close()
17 + _ = clientConn.Close()
18 + })
19 +
20 + broker := newLeaseBroker("lease-test", time.Hour, 2)
21 + session := newReverseSession(serverConn, time.Hour)
22 + if err := broker.Offer(session); err != nil {
23 + t.Fatalf("Offer() error = %v", err)
24 + }
25 +
26 + markerCh := make(chan byte, 1)
27 + errCh := make(chan error, 1)
28 + go func() {
29 + var marker [1]byte
30 + if _, err := io.ReadFull(clientConn, marker[:]); err != nil {
31 + errCh <- err
32 + return
33 + }
34 + markerCh <- marker[0]
35 + }()
36 +
37 + claimCtx, cancel := context.WithTimeout(context.Background(), time.Second)
38 + defer cancel()
39 +
40 + claimed, err := broker.Claim(claimCtx)
41 + if err != nil {
42 + t.Fatalf("Claim() error = %v", err)
43 + }
44 + if claimed != session {
45 + t.Fatalf("Claim() returned unexpected session")
46 + }
47 +
48 + select {
49 + case err := <-errCh:
50 + t.Fatalf("ReadFull() error = %v", err)
51 + case marker := <-markerCh:
52 + if marker != MarkerTLSStart {
53 + t.Fatalf("marker = 0x%02x, want 0x%02x", marker, MarkerTLSStart)
54 + }
55 + case <-time.After(time.Second):
56 + t.Fatal("timed out waiting for activation marker")
57 + }
58 +}
59 +
60 +func TestLeaseBrokerDropClosesIdleSessions(t *testing.T) {
61 + t.Parallel()
62 +
63 + serverConn, clientConn := net.Pipe()
64 + broker := newLeaseBroker("lease-test", time.Hour, 2)
65 + session := newReverseSession(serverConn, time.Hour)
66 + if err := broker.Offer(session); err != nil {
67 + t.Fatalf("Offer() error = %v", err)
68 + }
69 +
70 + broker.Drop()
71 +
72 + buf := make([]byte, 1)
73 + _ = clientConn.SetReadDeadline(time.Now().Add(time.Second))
74 + if _, err := clientConn.Read(buf); err == nil {
75 + t.Fatal("Read() succeeded, want connection close")
76 + }
77 +}
portal/helpers.go new
+167
@@ -0,0 +1,167 @@
1 +package portal
2 +
3 +import (
4 + "crypto/rand"
5 + "encoding/hex"
6 + "fmt"
7 + "net"
8 + "net/url"
9 + "strings"
10 + "time"
11 +)
12 +
13 +const (
14 + defaultLeaseTTL = 2 * time.Minute
15 + defaultClaimTimeout = 10 * time.Second
16 + defaultIdleKeepalive = 15 * time.Second
17 + defaultReadyQueueLimit = 8
18 + defaultClientHelloWait = 2 * time.Second
19 + defaultControlBodyLimit = 32 << 10
20 + defaultSessionWriteLimit = 5 * time.Second
21 +)
22 +
23 +func PortalRootHost(portalURL string) string {
24 + u, err := url.Parse(strings.TrimSpace(portalURL))
25 + if err != nil || u.Host == "" {
26 + return ""
27 + }
28 + return normalizeHostname(u.Hostname())
29 +}
30 +
31 +func NormalizeRelayURL(raw string) (string, error) {
32 + u, err := url.Parse(strings.TrimSpace(raw))
33 + if err != nil {
34 + return "", fmt.Errorf("parse relay url: %w", err)
35 + }
36 + if !strings.EqualFold(u.Scheme, "https") {
37 + return "", fmt.Errorf("relay url must use https: %q", raw)
38 + }
39 + if u.Host == "" {
40 + return "", fmt.Errorf("relay url host is empty: %q", raw)
41 + }
42 + u.Path = strings.TrimRight(u.Path, "/")
43 + u.RawQuery = ""
44 + u.Fragment = ""
45 + return u.String(), nil
46 +}
47 +
48 +func normalizeHostname(host string) string {
49 + host = strings.TrimSpace(strings.ToLower(host))
50 + host = strings.TrimSuffix(host, ".")
51 + return host
52 +}
53 +
54 +func sanitizeLabel(name string) string {
55 + name = strings.ToLower(strings.TrimSpace(name))
56 + var b strings.Builder
57 + lastHyphen := false
58 + for _, r := range name {
59 + switch {
60 + case r >= 'a' && r <= 'z':
61 + b.WriteRune(r)
62 + lastHyphen = false
63 + case r >= '0' && r <= '9':
64 + b.WriteRune(r)
65 + lastHyphen = false
66 + default:
67 + if b.Len() == 0 || lastHyphen {
68 + continue
69 + }
70 + b.WriteByte('-')
71 + lastHyphen = true
72 + }
73 + }
74 + s := strings.Trim(b.String(), "-")
75 + if s == "" {
76 + return "app"
77 + }
78 + return s
79 +}
80 +
81 +func suggestHostname(name, rootHost string) string {
82 + label := sanitizeLabel(name)
83 + if rootHost == "" {
84 + return label
85 + }
86 + return label + "." + rootHost
87 +}
88 +
89 +func hostnameMatchesWildcard(pattern, host string) bool {
90 + if !strings.HasPrefix(pattern, "*.") {
91 + return false
92 + }
93 + suffix := strings.TrimPrefix(pattern, "*.")
94 + parts := strings.Split(host, ".")
95 + if len(parts) < 2 {
96 + return false
97 + }
98 + return normalizeHostname("*."+strings.Join(parts[1:], ".")) == normalizeHostname(pattern) && strings.Count(suffix, ".")+1 == len(parts)-1
99 +}
100 +
101 +func randomID(prefix string) string {
102 + buf := make([]byte, 8)
103 + if _, err := rand.Read(buf); err != nil {
104 + panic(err)
105 + }
106 + return prefix + hex.EncodeToString(buf)
107 +}
108 +
109 +func randomToken() string {
110 + return randomID("tok_")
111 +}
112 +
113 +func durationOrDefault(v, fallback time.Duration) time.Duration {
114 + if v > 0 {
115 + return v
116 + }
117 + return fallback
118 +}
119 +
120 +func intOrDefault(v, fallback int) int {
121 + if v > 0 {
122 + return v
123 + }
124 + return fallback
125 +}
126 +
127 +func normalizeMetadata(meta LeaseMetadata) LeaseMetadata {
128 + meta.Description = strings.TrimSpace(meta.Description)
129 + meta.Owner = strings.TrimSpace(meta.Owner)
130 + meta.Thumbnail = strings.TrimSpace(meta.Thumbnail)
131 + meta.Tags = normalizeTags(meta.Tags)
132 + return meta
133 +}
134 +
135 +func normalizeTags(tags []string) []string {
136 + if len(tags) == 0 {
137 + return nil
138 + }
139 + seen := make(map[string]struct{}, len(tags))
140 + out := make([]string, 0, len(tags))
141 + for _, tag := range tags {
142 + tag = strings.TrimSpace(tag)
143 + if tag == "" {
144 + continue
145 + }
146 + if _, ok := seen[tag]; ok {
147 + continue
148 + }
149 + seen[tag] = struct{}{}
150 + out = append(out, tag)
151 + }
152 + if len(out) == 0 {
153 + return nil
154 + }
155 + return out
156 +}
157 +
158 +func hostPortOrLoopback(addr string) string {
159 + host, port, err := net.SplitHostPort(addr)
160 + if err != nil {
161 + return addr
162 + }
163 + if host == "" || host == "::" || host == "0.0.0.0" {
164 + host = "127.0.0.1"
165 + }
166 + return net.JoinHostPort(host, port)
167 +}
portal/keyless/client.go deleted
-208
@@ -1,208 +0,0 @@
1 -package keyless
2 -
3 -import (
4 - "context"
5 - "crypto/tls"
6 - "crypto/x509"
7 - "encoding/pem"
8 - "errors"
9 - "fmt"
10 - "net"
11 - "net/url"
12 - "strings"
13 - "time"
14 -
15 - "github.com/rs/zerolog/log"
16 -
17 - keylesstls "github.com/gosuda/keyless_tls/keyless"
18 -
19 - "gosuda.org/portal/types"
20 -)
21 -
22 -// BuildClientTLSConfig builds a keyless TLS server config for tunnel-side TLS termination.
23 -// It returns the TLS config and a close callback for signer resources.
24 -func BuildClientTLSConfig(relayAddr, keylessServerName, domain string) (*tls.Config, func(), error) {
25 - if keylessServerName == "" {
26 - return nil, nil, errors.New("keyless server name is required")
27 - }
28 - if domain == "" {
29 - return nil, nil, errors.New("tls domain is required")
30 - }
31 - certPEM, rootCAPEM, err := ResolveMaterials(
32 - context.Background(),
33 - relayAddr,
34 - keylessServerName,
35 - nil,
36 - nil,
37 - )
38 - if err != nil {
39 - return nil, nil, fmt.Errorf("prepare keyless materials: %w", err)
40 - }
41 -
42 - if verifyErr := VerifyCertificateHostname(certPEM, domain); verifyErr != nil {
43 - return nil, nil, fmt.Errorf("keyless certificate does not cover %s: %w", domain, verifyErr)
44 - }
45 -
46 - remoteSigner, err := keylesstls.NewRemoteSigner(keylesstls.RemoteSignerConfig{
47 - Endpoint: relayAddr,
48 - ServerName: keylessServerName,
49 - KeyID: RelayKeyID,
50 - RootCAPEM: rootCAPEM,
51 - }, certPEM)
52 - if err != nil {
53 - return nil, nil, fmt.Errorf("create keyless remote signer: %w", err)
54 - }
55 -
56 - tlsConfig, err := keylesstls.NewServerTLSConfig(keylesstls.ServerTLSConfig{
57 - CertPEM: certPEM,
58 - Signer: remoteSigner,
59 - })
60 - if err != nil {
61 - _ = remoteSigner.Close()
62 - return nil, nil, fmt.Errorf("create keyless TLS config: %w", err)
63 - }
64 - tlsConfig.NextProtos = []string{"http/1.1"}
65 -
66 - return tlsConfig, func() { _ = remoteSigner.Close() }, nil
67 -}
68 -
69 -// ResolveMaterials prepares certificate chain and root CAs for keyless TLS mode.
70 -func ResolveMaterials(
71 - ctx context.Context,
72 - keylessEndpoint string,
73 - keylessServerName string,
74 - inlineCertPEM []byte,
75 - inlineRootCAPEM []byte,
76 -) ([]byte, []byte, error) {
77 - certPEM := append([]byte(nil), inlineCertPEM...)
78 - rootCAPEM := append([]byte(nil), inlineRootCAPEM...)
79 -
80 - // If both are explicitly provided, no need for endpoint fetch.
81 - if len(certPEM) > 0 && len(rootCAPEM) > 0 {
82 - return certPEM, rootCAPEM, nil
83 - }
84 -
85 - chainFromEndpoint, err := FetchEndpointCertificateChain(ctx, keylessEndpoint, keylessServerName)
86 - if err != nil && len(certPEM) == 0 {
87 - return nil, nil, fmt.Errorf("auto-discover certificate chain from signer endpoint: %w", err)
88 - }
89 - if err != nil {
90 - log.Debug().Err(err).Msg("[SDK] Failed to fetch cert from endpoint, using inline materials")
91 - }
92 -
93 - if len(certPEM) == 0 {
94 - certPEM = chainFromEndpoint
95 - }
96 - if len(certPEM) == 0 {
97 - return nil, nil, errors.New("keyless certificate chain is required")
98 - }
99 -
100 - if len(rootCAPEM) == 0 && len(chainFromEndpoint) > 0 {
101 - rootCAPEM = append([]byte(nil), chainFromEndpoint...)
102 - }
103 - if len(rootCAPEM) == 0 {
104 - rootCAPEM = append([]byte(nil), certPEM...)
105 - }
106 -
107 - return certPEM, rootCAPEM, nil
108 -}
109 -
110 -// VerifyCertificateHostname checks whether the leaf cert covers hostname.
111 -func VerifyCertificateHostname(certPEM []byte, hostname string) error {
112 - _, leaf, err := ParseCertificateChainPEM(certPEM)
113 - if err != nil {
114 - return err
115 - }
116 - return leaf.VerifyHostname(hostname)
117 -}
118 -
119 -// ParseCertificateChainPEM parses PEM cert chain and returns DER chain + leaf.
120 -func ParseCertificateChainPEM(certPEM []byte) ([][]byte, *x509.Certificate, error) {
121 - if len(certPEM) == 0 {
122 - return nil, nil, errors.New("certificate PEM is empty")
123 - }
124 -
125 - var chain [][]byte
126 - rest := certPEM
127 - for {
128 - block, next := pem.Decode(rest)
129 - if block == nil {
130 - break
131 - }
132 - if block.Type == "CERTIFICATE" {
133 - chain = append(chain, block.Bytes)
134 - }
135 - rest = next
136 - }
137 - if len(chain) == 0 {
138 - return nil, nil, errors.New("no certificate blocks found")
139 - }
140 -
141 - leaf, err := x509.ParseCertificate(chain[0])
142 - if err != nil {
143 - return nil, nil, fmt.Errorf("parse leaf certificate: %w", err)
144 - }
145 -
146 - return chain, leaf, nil
147 -}
148 -
149 -// FetchEndpointCertificateChain fetches peer cert chain from signer endpoint.
150 -func FetchEndpointCertificateChain(ctx context.Context, endpoint string, serverName string) ([]byte, error) {
151 - raw := endpoint
152 - if raw == "" {
153 - return nil, errors.New("endpoint is required")
154 - }
155 - if !strings.Contains(raw, "://") {
156 - raw = "https://" + raw
157 - }
158 -
159 - u, err := url.Parse(raw)
160 - if err != nil {
161 - return nil, fmt.Errorf("parse endpoint URL: %w", err)
162 - }
163 - if u.Scheme == "http" {
164 - return nil, errors.New("http signer endpoint does not expose TLS certificate chain (use https endpoint)")
165 - }
166 -
167 - host := u.Hostname()
168 - if host == "" {
169 - return nil, errors.New("endpoint hostname is empty")
170 - }
171 - port := u.Port()
172 - if port == "" {
173 - port = "443"
174 - }
175 - if serverName == "" {
176 - serverName = host
177 - }
178 -
179 - dialer := &net.Dialer{Timeout: 5 * time.Second}
180 - rawConn, err := dialer.DialContext(ctx, "tcp", net.JoinHostPort(host, port))
181 - if err != nil {
182 - return nil, fmt.Errorf("dial signer endpoint: %w", err)
183 - }
184 -
185 - tlsConn := tls.Client(rawConn, &tls.Config{
186 - MinVersion: tls.VersionTLS12,
187 - ServerName: serverName,
188 - InsecureSkipVerify: types.IsLocalhost(host),
189 - })
190 - defer tlsConn.Close()
191 - if err := tlsConn.HandshakeContext(ctx); err != nil {
192 - return nil, fmt.Errorf("TLS handshake with signer endpoint: %w", err)
193 - }
194 -
195 - peerCerts := tlsConn.ConnectionState().PeerCertificates
196 - if len(peerCerts) == 0 {
197 - return nil, errors.New("no peer certificates from signer endpoint")
198 - }
199 -
200 - var chainPEM []byte
201 - for _, cert := range peerCerts {
202 - chainPEM = append(chainPEM, pem.EncodeToMemory(&pem.Block{
203 - Type: "CERTIFICATE",
204 - Bytes: cert.Raw,
205 - })...)
206 - }
207 - return chainPEM, nil
208 -}
portal/keyless/signer.go deleted
-77
@@ -1,77 +0,0 @@
1 -package keyless
2 -
3 -import (
4 - "context"
5 - "errors"
6 - "fmt"
7 - "os"
8 - "time"
9 -
10 - ksigner "github.com/gosuda/keyless_tls/relay/signer"
11 - "github.com/gosuda/keyless_tls/relay/signrpc"
12 -)
13 -
14 -const (
15 - RelayKeyID = "relay-cert"
16 - defaultAllowedSkew = 30 * time.Second
17 -)
18 -
19 -var (
20 - ErrSignerDisabled = errors.New("keyless signer is disabled")
21 - ErrInvalidArgument = ksigner.ErrInvalidArgument
22 - ErrPermissionDenied = ksigner.ErrPermissionDenied
23 -)
24 -
25 -type SignRequest = signrpc.SignRequest
26 -type SignResponse = signrpc.SignResponse
27 -type ErrorResponse = signrpc.ErrorResponse
28 -
29 -// Signer serves the keyless signing endpoint used by tunnel keyless mode.
30 -type Signer struct {
31 - service *ksigner.Service
32 - keyID string
33 -}
34 -
35 -func NewSigner(KeyFile string) (*Signer, error) {
36 - keyPEM, err := os.ReadFile(KeyFile)
37 - if err != nil {
38 - if errors.Is(err, os.ErrNotExist) {
39 - return nil, nil
40 - }
41 - return nil, fmt.Errorf("read keyless signing key: %w", err)
42 - }
43 -
44 - signingKey, err := ksigner.ParsePrivateKeyPEM(keyPEM)
45 - if err != nil {
46 - return nil, fmt.Errorf("parse keyless signing key: %w", err)
47 - }
48 -
49 - store := ksigner.NewStaticKeyStore()
50 - if err := store.Put(RelayKeyID, signingKey); err != nil {
51 - return nil, fmt.Errorf("register keyless signing key: %w", err)
52 - }
53 -
54 - svc := &ksigner.Service{
55 - Store: store,
56 - AllowedSkew: defaultAllowedSkew,
57 - }
58 -
59 - return &Signer{
60 - service: svc,
61 - keyID: RelayKeyID,
62 - }, nil
63 -}
64 -
65 -func (s *Signer) KeyID() string {
66 - if s == nil {
67 - return ""
68 - }
69 - return s.keyID
70 -}
71 -
72 -func (s *Signer) Sign(ctx context.Context, req *SignRequest) (*SignResponse, error) {
73 - if s == nil || s.service == nil {
74 - return nil, ErrSignerDisabled
75 - }
76 - return s.service.Sign(ctx, req)
77 -}
portal/lease.go deleted
-207
@@ -1,207 +0,0 @@
1 -package portal
2 -
3 -import (
4 - "sync"
5 - "time"
6 -
7 - "gosuda.org/portal/types"
8 -)
9 -
10 -type LeaseManager struct {
11 - leases map[string]*types.LeaseEntry
12 - stopCh chan struct{}
13 - bannedLeases map[string]struct{}
14 - onLeaseDeleted func(string)
15 - ttlInterval time.Duration
16 - leasesLock sync.RWMutex
17 - startOnce sync.Once
18 - stopOnce sync.Once
19 -}
20 -
21 -func NewLeaseManager(ttlInterval time.Duration) *LeaseManager {
22 - return &LeaseManager{
23 - leases: make(map[string]*types.LeaseEntry),
24 - stopCh: make(chan struct{}),
25 - ttlInterval: ttlInterval,
26 - bannedLeases: make(map[string]struct{}),
27 - }
28 -}
29 -
30 -func (lm *LeaseManager) Start() {
31 - lm.startOnce.Do(func() {
32 - go lm.ttlWorker()
33 - })
34 -}
35 -
36 -func (lm *LeaseManager) Stop() {
37 - lm.stopOnce.Do(func() {
38 - close(lm.stopCh)
39 - })
40 -}
41 -
42 -func (lm *LeaseManager) ttlWorker() {
43 - ticker := time.NewTicker(lm.ttlInterval)
44 - defer ticker.Stop()
45 -
46 - for {
47 - select {
48 - case <-ticker.C:
49 - lm.cleanupExpiredLeases()
50 - case <-lm.stopCh:
51 - return
52 - }
53 - }
54 -}
55 -
56 -func (lm *LeaseManager) cleanupExpiredLeases() {
57 - lm.leasesLock.Lock()
58 -
59 - now := time.Now()
60 - expired := make([]string, 0)
61 - for id, lease := range lm.leases {
62 - if now.After(lease.Lease.Expires) {
63 - delete(lm.leases, id)
64 - expired = append(expired, id)
65 - }
66 - }
67 - callback := lm.onLeaseDeleted
68 - lm.leasesLock.Unlock()
69 -
70 - if callback == nil {
71 - return
72 - }
73 - for _, id := range expired {
74 - callback(id)
75 - }
76 -}
77 -
78 -func (lm *LeaseManager) UpdateLease(lease *types.Lease) bool {
79 - lm.leasesLock.Lock()
80 - defer lm.leasesLock.Unlock()
81 -
82 - identityID := lease.ID
83 -
84 - // Check if lease is already expired
85 - if time.Now().After(lease.Expires) {
86 - return false
87 - }
88 -
89 - // policy checks
90 - if _, banned := lm.bannedLeases[identityID]; banned {
91 - return false
92 - }
93 -
94 - // Check for name conflicts (only if name is not empty)
95 - if lease.Name != "" && lease.Name != "(unnamed)" {
96 - for existingID, existingEntry := range lm.leases {
97 - // Skip if it's the same identity (updating own lease)
98 - if existingID == identityID {
99 - continue
100 - }
101 - // Check if another identity is using the same name
102 - if existingEntry.Lease.Name == lease.Name {
103 - // Name conflict with a different identity
104 - return false
105 - }
106 - }
107 - }
108 -
109 - var firstSeen time.Time
110 - if existing, exists := lm.leases[identityID]; exists {
111 - firstSeen = existing.FirstSeen
112 - }
113 - if firstSeen.IsZero() {
114 - firstSeen = time.Now()
115 - }
116 -
117 - lm.leases[identityID] = &types.LeaseEntry{
118 - Lease: lease,
119 - LastSeen: time.Now(),
120 - FirstSeen: firstSeen,
121 - }
122 -
123 - return true
124 -}
125 -
126 -func (lm *LeaseManager) DeleteLease(leaseID string) bool {
127 - lm.leasesLock.Lock()
128 - if _, exists := lm.leases[leaseID]; exists {
129 - delete(lm.leases, leaseID)
130 - callback := lm.onLeaseDeleted
131 - lm.leasesLock.Unlock()
132 - if callback != nil {
133 - callback(leaseID)
134 - }
135 - return true
136 - }
137 - lm.leasesLock.Unlock()
138 - return false
139 -}
140 -
141 -func (lm *LeaseManager) SetOnLeaseDeleted(callback func(string)) {
142 - lm.leasesLock.Lock()
143 - defer lm.leasesLock.Unlock()
144 - lm.onLeaseDeleted = callback
145 -}
146 -
147 -func (lm *LeaseManager) GetLeaseByID(leaseID string) (*types.LeaseEntry, bool) {
148 - lm.leasesLock.RLock()
149 - defer lm.leasesLock.RUnlock()
150 -
151 - // Check if banned
152 - if _, banned := lm.bannedLeases[leaseID]; banned {
153 - return nil, false
154 - }
155 -
156 - lease, exists := lm.leases[leaseID]
157 - if !exists {
158 - return nil, false
159 - }
160 -
161 - // Check if lease is expired
162 - if time.Now().After(lease.Lease.Expires) {
163 - return nil, false
164 - }
165 -
166 - return lease, true
167 -}
168 -
169 -// GetAllLeaseEntries returns all lease entries from the lease manager.
170 -func (lm *LeaseManager) GetAllLeaseEntries() []*types.LeaseEntry {
171 - lm.leasesLock.RLock()
172 - defer lm.leasesLock.RUnlock()
173 -
174 - now := time.Now()
175 - var entries []*types.LeaseEntry
176 -
177 - for _, entry := range lm.leases {
178 - if now.Before(entry.Lease.Expires) {
179 - entries = append(entries, entry)
180 - }
181 - }
182 -
183 - return entries
184 -}
185 -
186 -// BanLease adds a lease ID to the denylist.
187 -func (lm *LeaseManager) BanLease(leaseID string) {
188 - lm.leasesLock.Lock()
189 - lm.bannedLeases[leaseID] = struct{}{}
190 - lm.leasesLock.Unlock()
191 -}
192 -
193 -func (lm *LeaseManager) UnbanLease(leaseID string) {
194 - lm.leasesLock.Lock()
195 - delete(lm.bannedLeases, leaseID)
196 - lm.leasesLock.Unlock()
197 -}
198 -
199 -func (lm *LeaseManager) GetBannedLeases() []string {
200 - lm.leasesLock.RLock()
201 - defer lm.leasesLock.RUnlock()
202 - banned := make([]string, 0, len(lm.bannedLeases))
203 - for id := range lm.bannedLeases {
204 - banned = append(banned, id)
205 - }
206 - return banned
207 -}
portal/lease_test.go deleted
-92
@@ -1,92 +0,0 @@
1 -package portal
2 -
3 -import (
4 - "slices"
5 - "testing"
6 - "time"
7 -
8 - "gosuda.org/portal/types"
9 -)
10 -
11 -func TestLeaseManagerDeleteLeaseInvokesCallback(t *testing.T) {
12 - lm := NewLeaseManager(time.Second)
13 -
14 - var deleted []string
15 - lm.SetOnLeaseDeleted(func(id string) {
16 - deleted = append(deleted, id)
17 - })
18 -
19 - lease := &types.Lease{
20 - ID: "lease-1",
21 - Name: "app-1",
22 - Expires: time.Now().Add(30 * time.Second),
23 - }
24 - if !lm.UpdateLease(lease) {
25 - t.Fatalf("expected lease update success")
26 - }
27 -
28 - if !lm.DeleteLease("lease-1") {
29 - t.Fatalf("expected lease deletion success")
30 - }
31 - if !slices.Contains(deleted, "lease-1") {
32 - t.Fatalf("expected callback with lease-1, got %v", deleted)
33 - }
34 -}
35 -
36 -func TestLeaseManagerCleanupExpiredLeasesInvokesCallback(t *testing.T) {
37 - lm := NewLeaseManager(time.Second)
38 -
39 - var deleted []string
40 - lm.SetOnLeaseDeleted(func(id string) {
41 - deleted = append(deleted, id)
42 - })
43 -
44 - lm.leases["expired-1"] = &types.LeaseEntry{
45 - Lease: &types.Lease{
46 - ID: "expired-1",
47 - Name: "expired",
48 - Expires: time.Now().Add(-1 * time.Second),
49 - },
50 - }
51 - lm.leases["active-1"] = &types.LeaseEntry{
52 - Lease: &types.Lease{
53 - ID: "active-1",
54 - Name: "active",
55 - Expires: time.Now().Add(30 * time.Second),
56 - },
57 - }
58 -
59 - lm.cleanupExpiredLeases()
60 -
61 - if !slices.Contains(deleted, "expired-1") {
62 - t.Fatalf("expected callback with expired-1, got %v", deleted)
63 - }
64 - if _, ok := lm.leases["expired-1"]; ok {
65 - t.Fatal("expected expired-1 removed")
66 - }
67 - if _, ok := lm.leases["active-1"]; !ok {
68 - t.Fatal("expected active-1 to remain")
69 - }
70 -}
71 -
72 -func TestLeaseManagerStopIsIdempotent(_ *testing.T) {
73 - lm := NewLeaseManager(10 * time.Millisecond)
74 -
75 - lm.Start()
76 - lm.Stop()
77 - lm.Stop()
78 -}
79 -
80 -func TestLeaseManagerGetBannedLeasesReturnsPlainLeaseIDs(t *testing.T) {
81 - lm := NewLeaseManager(time.Second)
82 - lm.BanLease("lease-a")
83 - lm.BanLease("lease-b")
84 -
85 - got := lm.GetBannedLeases()
86 - slices.Sort(got)
87 -
88 - want := []string{"lease-a", "lease-b"}
89 - if !slices.Equal(got, want) {
90 - t.Fatalf("GetBannedLeases() = %v, want %v", got, want)
91 - }
92 -}
portal/policy/approver.go deleted
-106
@@ -1,106 +0,0 @@
1 -package policy
2 -
3 -import (
4 - "fmt"
5 - "sync"
6 -)
7 -
8 -// Mode represents the approval mode for new connections.
9 -type Mode string
10 -
11 -const (
12 - ModeAuto Mode = "auto"
13 - ModeManual Mode = "manual"
14 -)
15 -
16 -// Approver manages approval/denial state for leases.
17 -type Approver struct {
18 - approvedLeases map[string]struct{}
19 - deniedLeases map[string]struct{}
20 - approvalMode Mode
21 - mu sync.RWMutex
22 -}
23 -
24 -func NewApprover() *Approver {
25 - return &Approver{
26 - approvalMode: ModeAuto,
27 - approvedLeases: make(map[string]struct{}),
28 - deniedLeases: make(map[string]struct{}),
29 - }
30 -}
31 -
32 -func (m *Approver) GetApprovalMode() Mode {
33 - m.mu.RLock()
34 - defer m.mu.RUnlock()
35 - return m.approvalMode
36 -}
37 -
38 -func (m *Approver) SetApprovalMode(mode Mode) error {
39 - if mode != ModeAuto && mode != ModeManual {
40 - return fmt.Errorf("invalid approval mode: %q", mode)
41 - }
42 - m.mu.Lock()
43 - defer m.mu.Unlock()
44 - m.approvalMode = mode
45 - return nil
46 -}
47 -
48 -func (m *Approver) IsLeaseApproved(leaseID string) bool {
49 - m.mu.RLock()
50 - defer m.mu.RUnlock()
51 - _, ok := m.approvedLeases[leaseID]
52 - return ok
53 -}
54 -
55 -func (m *Approver) ApproveLease(leaseID string) {
56 - m.mu.Lock()
57 - defer m.mu.Unlock()
58 - m.approvedLeases[leaseID] = struct{}{}
59 - delete(m.deniedLeases, leaseID)
60 -}
61 -
62 -func (m *Approver) RevokeLease(leaseID string) {
63 - m.mu.Lock()
64 - defer m.mu.Unlock()
65 - delete(m.approvedLeases, leaseID)
66 -}
67 -
68 -func (m *Approver) GetApprovedLeases() []string {
69 - m.mu.RLock()
70 - defer m.mu.RUnlock()
71 - result := make([]string, 0, len(m.approvedLeases))
72 - for id := range m.approvedLeases {
73 - result = append(result, id)
74 - }
75 - return result
76 -}
77 -
78 -func (m *Approver) IsLeaseDenied(leaseID string) bool {
79 - m.mu.RLock()
80 - defer m.mu.RUnlock()
81 - _, ok := m.deniedLeases[leaseID]
82 - return ok
83 -}
84 -
85 -func (m *Approver) DenyLease(leaseID string) {
86 - m.mu.Lock()
87 - defer m.mu.Unlock()
88 - m.deniedLeases[leaseID] = struct{}{}
89 - delete(m.approvedLeases, leaseID)
90 -}
91 -
92 -func (m *Approver) UndenyLease(leaseID string) {
93 - m.mu.Lock()
94 - defer m.mu.Unlock()
95 - delete(m.deniedLeases, leaseID)
96 -}
97 -
98 -func (m *Approver) GetDeniedLeases() []string {
99 - m.mu.RLock()
100 - defer m.mu.RUnlock()
101 - result := make([]string, 0, len(m.deniedLeases))
102 - for id := range m.deniedLeases {
103 - result = append(result, id)
104 - }
105 - return result
106 -}
portal/policy/authenticator.go deleted
-278
@@ -1,278 +0,0 @@
1 -package policy
2 -
3 -import (
4 - "crypto/rand"
5 - "crypto/subtle"
6 - "encoding/hex"
7 - "fmt"
8 - "io"
9 - "sort"
10 - "sync"
11 - "time"
12 -
13 - "github.com/rs/zerolog/log"
14 -)
15 -
16 -const (
17 - maxFailedAttempts = 3
18 - lockDuration = 1 * time.Minute
19 - sessionDuration = 24 * time.Hour
20 - failedLoginRetention = 15 * time.Minute
21 - failedLoginSweepWindow = 1 * time.Minute
22 - maxFailedLoginEntries = 4096
23 -)
24 -
25 -// Authenticator manages admin authentication with rate limiting.
26 -type Authenticator struct {
27 - lastSweepAt time.Time
28 - failedLogins map[string]*loginAttempt
29 - sessions map[string]time.Time
30 - secretKey string
31 - mu sync.RWMutex
32 -}
33 -
34 -type loginAttempt struct {
35 - lockedAt time.Time
36 - lastSeenAt time.Time
37 - count int
38 -}
39 -
40 -// NewAuthenticator creates a new Authenticator with the given secret key.
41 -func NewAuthenticator(secretKey string) *Authenticator {
42 - if secretKey == "" {
43 - randomBytes := make([]byte, 16)
44 - if _, err := rand.Read(randomBytes); err != nil {
45 - log.Fatal().Err(err).Msg("[server] failed to generate random admin secret key")
46 - }
47 - secretKey = hex.EncodeToString(randomBytes)
48 - log.Warn().Int("key_length", len(secretKey)).Msg("[server] auto-generated ADMIN_SECRET_KEY (set ADMIN_SECRET_KEY env to use your own)")
49 - } else {
50 - log.Info().Int("key_length", len(secretKey)).Msg("[server] admin authentication enabled")
51 - }
52 -
53 - return &Authenticator{
54 - secretKey: secretKey,
55 - failedLogins: make(map[string]*loginAttempt),
56 - sessions: make(map[string]time.Time),
57 - }
58 -}
59 -
60 -// IsIPLocked checks if an IP is currently locked out.
61 -func (m *Authenticator) IsIPLocked(ip string) bool {
62 - m.mu.RLock()
63 - defer m.mu.RUnlock()
64 -
65 - attempt, exists := m.failedLogins[ip]
66 - if !exists {
67 - return false
68 - }
69 - return lockRemaining(attempt, time.Now()) > 0
70 -}
71 -
72 -// GetLockRemainingSeconds returns the remaining seconds until the IP is unlocked.
73 -func (m *Authenticator) GetLockRemainingSeconds(ip string) int {
74 - m.mu.RLock()
75 - defer m.mu.RUnlock()
76 -
77 - attempt, exists := m.failedLogins[ip]
78 - if !exists {
79 - return 0
80 - }
81 - return int(lockRemaining(attempt, time.Now()).Seconds())
82 -}
83 -
84 -// RecordFailedLogin records a failed login attempt and returns true if the IP is now locked.
85 -func (m *Authenticator) RecordFailedLogin(ip string) bool {
86 - m.mu.Lock()
87 - defer m.mu.Unlock()
88 -
89 - now := time.Now()
90 - m.maybeSweepFailedLoginsLocked(now)
91 -
92 - attempt, exists := m.failedLogins[ip]
93 - if !exists {
94 - attempt = &loginAttempt{}
95 - m.failedLogins[ip] = attempt
96 - }
97 -
98 - // Reset if lock has expired
99 - if attempt.count >= maxFailedAttempts && now.Sub(attempt.lockedAt) >= lockDuration {
100 - attempt.count = 0
101 - }
102 -
103 - attempt.count++
104 - attempt.lastSeenAt = now
105 -
106 - locked := false
107 - if attempt.count >= maxFailedAttempts {
108 - attempt.lockedAt = now
109 - locked = true
110 - }
111 -
112 - m.enforceFailedLoginCapLocked()
113 - return locked
114 -}
115 -
116 -// ResetFailedLogin resets the failed login count for an IP.
117 -func (m *Authenticator) ResetFailedLogin(ip string) {
118 - m.mu.Lock()
119 - defer m.mu.Unlock()
120 -
121 - delete(m.failedLogins, ip)
122 -}
123 -
124 -// ValidateKey checks if the provided key matches the secret key.
125 -func (m *Authenticator) ValidateKey(key string) bool {
126 - if m.secretKey == "" {
127 - return false
128 - }
129 - return subtle.ConstantTimeCompare([]byte(key), []byte(m.secretKey)) == 1
130 -}
131 -
132 -// HasSecretKey returns true if a secret key is configured.
133 -func (m *Authenticator) HasSecretKey() bool {
134 - return m.secretKey != ""
135 -}
136 -
137 -// CreateSession creates a new session and returns the token.
138 -func (m *Authenticator) CreateSession() string {
139 - token, err := generateToken()
140 - if err != nil {
141 - log.Fatal().Err(err).Msg("[server] failed to generate secure admin session token")
142 - }
143 -
144 - m.mu.Lock()
145 - defer m.mu.Unlock()
146 -
147 - m.sessions[token] = time.Now().Add(sessionDuration)
148 -
149 - // Clean up expired sessions
150 - m.cleanupExpiredSessions()
151 -
152 - return token
153 -}
154 -
155 -// ValidateSession checks if a session token is valid.
156 -func (m *Authenticator) ValidateSession(token string) bool {
157 - if token == "" {
158 - return false
159 - }
160 -
161 - m.mu.RLock()
162 - defer m.mu.RUnlock()
163 -
164 - expiry, exists := m.sessions[token]
165 - if !exists {
166 - return false
167 - }
168 -
169 - return time.Now().Before(expiry)
170 -}
171 -
172 -// DeleteSession removes a session.
173 -func (m *Authenticator) DeleteSession(token string) {
174 - m.mu.Lock()
175 - defer m.mu.Unlock()
176 -
177 - delete(m.sessions, token)
178 -}
179 -
180 -// cleanupExpiredSessions removes expired sessions (must be called with lock held).
181 -func (m *Authenticator) cleanupExpiredSessions() {
182 - now := time.Now()
183 - for token, expiry := range m.sessions {
184 - if now.After(expiry) {
185 - delete(m.sessions, token)
186 - }
187 - }
188 -}
189 -
190 -// generateToken generates a secure random token.
191 -func generateToken() (string, error) {
192 - return generateTokenFromReader(rand.Reader)
193 -}
194 -
195 -func generateTokenFromReader(reader io.Reader) (string, error) {
196 - bytes := make([]byte, 32)
197 - if _, err := io.ReadFull(reader, bytes); err != nil {
198 - return "", fmt.Errorf("read random session token bytes: %w", err)
199 - }
200 - return hex.EncodeToString(bytes), nil
201 -}
202 -
203 -func (m *Authenticator) maybeSweepFailedLoginsLocked(now time.Time) {
204 - if !m.lastSweepAt.IsZero() && now.Sub(m.lastSweepAt) < failedLoginSweepWindow {
205 - return
206 - }
207 -
208 - m.sweepExpiredFailedLoginsLocked(now)
209 - m.lastSweepAt = now
210 -}
211 -
212 -func (m *Authenticator) sweepExpiredFailedLoginsLocked(now time.Time) {
213 - for ip, attempt := range m.failedLogins {
214 - if attempt == nil {
215 - delete(m.failedLogins, ip)
216 - continue
217 - }
218 -
219 - lastSeenAt := attempt.lastSeenAt
220 - if lastSeenAt.IsZero() {
221 - lastSeenAt = attempt.lockedAt
222 - }
223 - if lastSeenAt.IsZero() || now.Sub(lastSeenAt) >= failedLoginRetention {
224 - delete(m.failedLogins, ip)
225 - }
226 - }
227 -}
228 -
229 -func (m *Authenticator) enforceFailedLoginCapLocked() {
230 - if len(m.failedLogins) <= maxFailedLoginEntries {
231 - return
232 - }
233 -
234 - type failedEntry struct {
235 - lastSeenAt time.Time
236 - ip string
237 - }
238 -
239 - entries := make([]failedEntry, 0, len(m.failedLogins))
240 - for ip, attempt := range m.failedLogins {
241 - if attempt == nil {
242 - entries = append(entries, failedEntry{ip: ip})
243 - continue
244 - }
245 -
246 - lastSeenAt := attempt.lastSeenAt
247 - if lastSeenAt.IsZero() {
248 - lastSeenAt = attempt.lockedAt
249 - }
250 -
251 - entries = append(entries, failedEntry{
252 - ip: ip,
253 - lastSeenAt: lastSeenAt,
254 - })
255 - }
256 -
257 - sort.Slice(entries, func(i, j int) bool {
258 - return entries[i].lastSeenAt.Before(entries[j].lastSeenAt)
259 - })
260 -
261 - overflow := len(m.failedLogins) - maxFailedLoginEntries
262 - for i := range overflow {
263 - delete(m.failedLogins, entries[i].ip)
264 - }
265 -}
266 -
267 -func lockRemaining(attempt *loginAttempt, now time.Time) time.Duration {
268 - if attempt == nil || attempt.count < maxFailedAttempts {
269 - return 0
270 - }
271 -
272 - remaining := lockDuration - now.Sub(attempt.lockedAt)
273 - if remaining <= 0 {
274 - return 0
275 - }
276 -
277 - return remaining
278 -}
portal/policy/authenticator_test.go deleted
-123
@@ -1,123 +0,0 @@
1 -package policy
2 -
3 -import (
4 - "bytes"
5 - "errors"
6 - "fmt"
7 - "strings"
8 - "testing"
9 - "time"
10 -
11 - "github.com/rs/zerolog"
12 - "github.com/rs/zerolog/log"
13 -)
14 -
15 -type failingReader struct{}
16 -
17 -func (failingReader) Read(_ []byte) (int, error) {
18 - return 0, errors.New("rng unavailable")
19 -}
20 -
21 -func TestGenerateTokenFromReaderFailsClosed(t *testing.T) {
22 - t.Parallel()
23 -
24 - token, err := generateTokenFromReader(failingReader{})
25 - if err == nil {
26 - t.Fatal("expected an error when entropy source fails")
27 - }
28 - if token != "" {
29 - t.Fatalf("expected empty token on entropy failure, got %q", token)
30 - }
31 -}
32 -
33 -func TestRecordFailedLoginSweepsExpiredEntries(t *testing.T) {
34 - t.Parallel()
35 -
36 - m := NewAuthenticator("test-secret")
37 - now := time.Now()
38 -
39 - m.mu.Lock()
40 - m.failedLogins["expired-entry"] = &loginAttempt{
41 - count: 1,
42 - lastSeenAt: now.Add(-failedLoginRetention - time.Second),
43 - }
44 - m.lastSweepAt = now.Add(-failedLoginSweepWindow - time.Second)
45 - m.mu.Unlock()
46 -
47 - locked := m.RecordFailedLogin("active-entry")
48 - if locked {
49 - t.Fatal("first failed login attempt should not lock the IP")
50 - }
51 -
52 - m.mu.RLock()
53 - defer m.mu.RUnlock()
54 -
55 - if len(m.failedLogins) != 1 {
56 - t.Fatalf("expected 1 retained entry after sweep, got %d", len(m.failedLogins))
57 - }
58 - if _, exists := m.failedLogins["expired-entry"]; exists {
59 - t.Fatal("expired failed login entry should have been removed")
60 - }
61 - if _, exists := m.failedLogins["active-entry"]; !exists {
62 - t.Fatal("active failed login entry should be retained")
63 - }
64 -}
65 -
66 -func TestRecordFailedLoginEnforcesEntryCap(t *testing.T) {
67 - t.Parallel()
68 -
69 - m := NewAuthenticator("test-secret")
70 - base := time.Now().Add(-2 * time.Minute)
71 -
72 - m.mu.Lock()
73 - for i := range maxFailedLoginEntries {
74 - key := fmt.Sprintf("old-%05d", i)
75 - m.failedLogins[key] = &loginAttempt{
76 - count: 1,
77 - lastSeenAt: base.Add(time.Duration(i) * time.Millisecond),
78 - }
79 - }
80 - m.lastSweepAt = time.Now()
81 - m.mu.Unlock()
82 -
83 - m.RecordFailedLogin("new-entry")
84 -
85 - m.mu.RLock()
86 - defer m.mu.RUnlock()
87 -
88 - if len(m.failedLogins) != maxFailedLoginEntries {
89 - t.Fatalf("expected failed login map cap of %d, got %d", maxFailedLoginEntries, len(m.failedLogins))
90 - }
91 - if _, exists := m.failedLogins["new-entry"]; !exists {
92 - t.Fatal("new failed login entry should be retained after eviction")
93 - }
94 - if _, exists := m.failedLogins["old-00000"]; exists {
95 - t.Fatal("oldest failed login entry should be evicted when cap is exceeded")
96 - }
97 -}
98 -
99 -func TestAuthenticatorDoesNotLogPlaintextSecretsOrSessionToken(t *testing.T) {
100 - const secretKey = "super-secret-admin-key"
101 -
102 - var buf bytes.Buffer
103 - originalLogger := log.Logger
104 - log.Logger = zerolog.New(&buf)
105 - t.Cleanup(func() {
106 - log.Logger = originalLogger
107 - })
108 -
109 - m := NewAuthenticator(secretKey)
110 - logOutput := buf.String()
111 - if strings.Contains(logOutput, secretKey) {
112 - t.Fatalf("expected auth manager logs to omit plaintext secret key, got %q", logOutput)
113 - }
114 -
115 - buf.Reset()
116 - token := m.CreateSession()
117 - if token == "" {
118 - t.Fatal("expected non-empty session token")
119 - }
120 - if strings.Contains(buf.String(), token) {
121 - t.Fatalf("expected auth manager logs to omit plaintext session token, got %q", buf.String())
122 - }
123 -}
portal/policy/ip_filter.go deleted
-301
@@ -1,301 +0,0 @@
1 -package policy
2 -
3 -import (
4 - "fmt"
5 - "net"
6 - "net/http"
7 - "slices"
8 - "strings"
9 - "sync"
10 -)
11 -
12 -// IPFilter manages IP-based bans and lease-to-IP mapping.
13 -type IPFilter struct {
14 - bannedIPs map[string]struct{}
15 - leaseToIP map[string]string
16 - ipToLeases map[string][]string
17 - mu sync.RWMutex
18 -}
19 -
20 -var (
21 - trustedProxyMu sync.RWMutex
22 - trustedProxyCIDRs []*net.IPNet
23 -)
24 -
25 -const (
26 - xForwardedForHeader = "X-Forwarded-For"
27 - xRealIPHeader = "X-Real-IP"
28 -)
29 -
30 -// NewIPFilter creates a new IP filter.
31 -func NewIPFilter() *IPFilter {
32 - return &IPFilter{
33 - bannedIPs: make(map[string]struct{}),
34 - leaseToIP: make(map[string]string),
35 - ipToLeases: make(map[string][]string),
36 - }
37 -}
38 -
39 -// SetTrustedProxyCIDRs configures which remote peers can supply trusted forwarded headers.
40 -func SetTrustedProxyCIDRs(cidrs []*net.IPNet) {
41 - trustedProxyMu.Lock()
42 - defer trustedProxyMu.Unlock()
43 -
44 - if len(cidrs) == 0 {
45 - trustedProxyCIDRs = nil
46 - return
47 - }
48 -
49 - trustedProxyCIDRs = append(make([]*net.IPNet, 0, len(cidrs)), cidrs...)
50 -}
51 -
52 -// ParseTrustedProxyCIDRs parses a comma-separated CIDR allowlist for trusted proxy peers.
53 -// Empty input returns nil, nil.
54 -func ParseTrustedProxyCIDRs(raw string) ([]*net.IPNet, error) {
55 - raw = strings.TrimSpace(raw)
56 - if raw == "" {
57 - return nil, nil
58 - }
59 -
60 - parts := strings.Split(raw, ",")
61 - cidrs := make([]*net.IPNet, 0, len(parts))
62 - seen := make(map[string]struct{}, len(parts))
63 - for _, part := range parts {
64 - candidate := strings.TrimSpace(part)
65 - if candidate == "" {
66 - continue
67 - }
68 -
69 - _, network, err := net.ParseCIDR(candidate)
70 - if err != nil {
71 - return nil, fmt.Errorf("invalid trusted proxy CIDR %q: %w", candidate, err)
72 - }
73 -
74 - networkKey := network.String()
75 - if _, exists := seen[networkKey]; exists {
76 - continue
77 - }
78 -
79 - seen[networkKey] = struct{}{}
80 - cidrs = append(cidrs, network)
81 - }
82 -
83 - return cidrs, nil
84 -}
85 -
86 -// IsTrustedProxyRemoteAddr reports whether a remote peer is in the trusted proxy allowlist.
87 -func IsTrustedProxyRemoteAddr(remoteAddr string) bool {
88 - remoteIP := parseRemoteAddrIP(remoteAddr)
89 - if remoteIP == nil {
90 - return false
91 - }
92 -
93 - trustedProxyMu.RLock()
94 - defer trustedProxyMu.RUnlock()
95 -
96 - for _, network := range trustedProxyCIDRs {
97 - if network != nil && network.Contains(remoteIP) {
98 - return true
99 - }
100 - }
101 -
102 - return false
103 -}
104 -
105 -// BanIP adds an IP to the ban list.
106 -func (m *IPFilter) BanIP(ip string) {
107 - m.mu.Lock()
108 - defer m.mu.Unlock()
109 - m.bannedIPs[ip] = struct{}{}
110 -}
111 -
112 -// UnbanIP removes an IP from the ban list.
113 -func (m *IPFilter) UnbanIP(ip string) {
114 - m.mu.Lock()
115 - defer m.mu.Unlock()
116 - delete(m.bannedIPs, ip)
117 -}
118 -
119 -// IsIPBanned checks if an IP is banned.
120 -func (m *IPFilter) IsIPBanned(ip string) bool {
121 - m.mu.RLock()
122 - defer m.mu.RUnlock()
123 - _, banned := m.bannedIPs[ip]
124 - return banned
125 -}
126 -
127 -// IsIPBannedByPolicy applies shared runtime policy rules before checking the ban map.
128 -func IsIPBannedByPolicy(ipFilter *IPFilter, candidate string) bool {
129 - if ipFilter == nil {
130 - return false
131 - }
132 - candidate = strings.TrimSpace(candidate)
133 - if candidate == "" {
134 - return false
135 - }
136 - return ipFilter.IsIPBanned(candidate)
137 -}
138 -
139 -// GetBannedIPs returns all banned IPs.
140 -func (m *IPFilter) GetBannedIPs() []string {
141 - m.mu.RLock()
142 - defer m.mu.RUnlock()
143 - result := make([]string, 0, len(m.bannedIPs))
144 - for ip := range m.bannedIPs {
145 - result = append(result, ip)
146 - }
147 - return result
148 -}
149 -
150 -// SetBannedIPs sets the banned IPs list (for loading from settings).
151 -func (m *IPFilter) SetBannedIPs(ips []string) {
152 - m.mu.Lock()
153 - defer m.mu.Unlock()
154 - m.bannedIPs = make(map[string]struct{}, len(ips))
155 - for _, ip := range ips {
156 - m.bannedIPs[ip] = struct{}{}
157 - }
158 -}
159 -
160 -// RegisterLeaseIP associates a lease ID with an IP address.
161 -func (m *IPFilter) RegisterLeaseIP(leaseID, ip string) {
162 - m.mu.Lock()
163 - defer m.mu.Unlock()
164 - if leaseID == "" || ip == "" {
165 - return
166 - }
167 -
168 - if oldIP, exists := m.leaseToIP[leaseID]; exists {
169 - if oldIP == ip {
170 - // Already registered; avoid duplicate lease entries per IP.
171 - return
172 - }
173 - m.removeLeaseFromIP(leaseID, oldIP)
174 - }
175 -
176 - // Defensively avoid duplicates if state was previously inconsistent.
177 - if slices.Contains(m.ipToLeases[ip], leaseID) {
178 - m.leaseToIP[leaseID] = ip
179 - return
180 - }
181 -
182 - m.leaseToIP[leaseID] = ip
183 - m.ipToLeases[ip] = append(m.ipToLeases[ip], leaseID)
184 -}
185 -
186 -// removeLeaseFromIP removes a lease from IP's lease list (must hold lock).
187 -func (m *IPFilter) removeLeaseFromIP(leaseID, ip string) {
188 - leases := m.ipToLeases[ip]
189 - for i, id := range leases {
190 - if id == leaseID {
191 - m.ipToLeases[ip] = append(leases[:i], leases[i+1:]...)
192 - break
193 - }
194 - }
195 - if len(m.ipToLeases[ip]) == 0 {
196 - delete(m.ipToLeases, ip)
197 - }
198 -}
199 -
200 -// GetLeaseIP returns the IP address for a lease ID.
201 -func (m *IPFilter) GetLeaseIP(leaseID string) string {
202 - m.mu.RLock()
203 - defer m.mu.RUnlock()
204 - return m.leaseToIP[leaseID]
205 -}
206 -
207 -// GetIPLeases returns all lease IDs for an IP.
208 -func (m *IPFilter) GetIPLeases(ip string) []string {
209 - m.mu.RLock()
210 - defer m.mu.RUnlock()
211 - result := make([]string, len(m.ipToLeases[ip]))
212 - copy(result, m.ipToLeases[ip])
213 - return result
214 -}
215 -
216 -// RemoveLeaseIP removes lease-to-IP mapping for a lease ID.
217 -func (m *IPFilter) RemoveLeaseIP(leaseID string) {
218 - m.mu.Lock()
219 - defer m.mu.Unlock()
220 -
221 - ip, exists := m.leaseToIP[leaseID]
222 - if !exists {
223 - return
224 - }
225 - delete(m.leaseToIP, leaseID)
226 - m.removeLeaseFromIP(leaseID, ip)
227 -}
228 -
229 -func normalizeClientIPCandidate(raw string) string {
230 - candidate := strings.TrimSpace(raw)
231 - if candidate == "" {
232 - return ""
233 - }
234 -
235 - if ip := net.ParseIP(candidate); ip != nil {
236 - return candidate
237 - }
238 -
239 - host, _, err := net.SplitHostPort(candidate)
240 - if err != nil {
241 - return ""
242 - }
243 - host = strings.TrimSpace(host)
244 - if host == "" || net.ParseIP(host) == nil {
245 - return ""
246 - }
247 - return host
248 -}
249 -
250 -func parseRemoteAddrIP(remoteAddr string) net.IP {
251 - remoteAddr = strings.TrimSpace(remoteAddr)
252 - if remoteAddr == "" {
253 - return nil
254 - }
255 -
256 - host := remoteAddr
257 - if parsedHost, _, err := net.SplitHostPort(remoteAddr); err == nil {
258 - host = parsedHost
259 - }
260 -
261 - return net.ParseIP(strings.TrimSpace(host))
262 -}
263 -
264 -// ExtractClientIP extracts the client IP from an HTTP request.
265 -// Forwarded headers are trusted only when trustProxyHeaders is true and peer is trusted.
266 -func ExtractClientIP(r *http.Request, trustProxyHeaders bool) string {
267 - if r == nil {
268 - return ""
269 - }
270 -
271 - if trustProxyHeaders && IsTrustedProxyRemoteAddr(r.RemoteAddr) {
272 - // Check X-Forwarded-For header first (for proxied requests).
273 - if xff := r.Header.Get(xForwardedForHeader); xff != "" {
274 - // X-Forwarded-For can contain multiple IPs, take the first one.
275 - if before, _, ok := strings.Cut(xff, ","); ok {
276 - if ip := normalizeClientIPCandidate(before); ip != "" {
277 - return ip
278 - }
279 - } else if ip := normalizeClientIPCandidate(xff); ip != "" {
280 - return ip
281 - }
282 - }
283 -
284 - // Check X-Real-IP header.
285 - if xri := r.Header.Get(xRealIPHeader); xri != "" {
286 - if ip := normalizeClientIPCandidate(xri); ip != "" {
287 - return ip
288 - }
289 - }
290 - }
291 -
292 - // Fall back to RemoteAddr.
293 - ip, _, err := net.SplitHostPort(r.RemoteAddr)
294 - if err != nil {
295 - return strings.TrimSpace(r.RemoteAddr)
296 - }
297 - if normalized := normalizeClientIPCandidate(ip); normalized != "" {
298 - return normalized
299 - }
300 - return strings.TrimSpace(ip)
301 -}
portal/policy/ip_filter_test.go deleted
-118
@@ -1,118 +0,0 @@
1 -package policy
2 -
3 -import (
4 - "net"
5 - "net/http"
6 - "net/http/httptest"
7 - "testing"
8 -)
9 -
10 -func mustCIDR(t *testing.T, raw string) *net.IPNet {
11 - t.Helper()
12 -
13 - _, network, err := net.ParseCIDR(raw)
14 - if err != nil {
15 - t.Fatalf("parse CIDR %q: %v", raw, err)
16 - }
17 - return network
18 -}
19 -
20 -func TestExtractClientIPTrustsForwardedHeadersOnlyFromTrustedProxy(t *testing.T) {
21 - SetTrustedProxyCIDRs([]*net.IPNet{mustCIDR(t, "10.0.0.0/8")})
22 - t.Cleanup(func() {
23 - SetTrustedProxyCIDRs(nil)
24 - })
25 -
26 - trustedReq := httptest.NewRequest(http.MethodGet, "http://localhost", nil)
27 - trustedReq.RemoteAddr = "10.1.2.3:45000"
28 - trustedReq.Header.Set("X-Forwarded-For", "203.0.113.10, 10.1.2.3")
29 -
30 - if got := ExtractClientIP(trustedReq, true); got != "203.0.113.10" {
31 - t.Fatalf("expected forwarded client IP from trusted proxy, got %q", got)
32 - }
33 -
34 - untrustedReq := httptest.NewRequest(http.MethodGet, "http://localhost", nil)
35 - untrustedReq.RemoteAddr = "198.51.100.5:45000"
36 - untrustedReq.Header.Set("X-Forwarded-For", "203.0.113.10")
37 -
38 - if got := ExtractClientIP(untrustedReq, true); got != "198.51.100.5" {
39 - t.Fatalf("expected remote IP fallback for untrusted proxy, got %q", got)
40 - }
41 -}
42 -
43 -func TestExtractClientIPDoesNotTrustHeadersWithoutAllowlist(t *testing.T) {
44 - SetTrustedProxyCIDRs(nil)
45 -
46 - req := httptest.NewRequest(http.MethodGet, "http://localhost", nil)
47 - req.RemoteAddr = "10.1.2.3:45000"
48 - req.Header.Set("X-Real-IP", "203.0.113.77")
49 -
50 - if got := ExtractClientIP(req, true); got != "10.1.2.3" {
51 - t.Fatalf("expected remote IP when allowlist is empty, got %q", got)
52 - }
53 -}
54 -
55 -func TestIsTrustedProxyRemoteAddr(t *testing.T) {
56 - SetTrustedProxyCIDRs([]*net.IPNet{
57 - mustCIDR(t, "10.0.0.0/8"),
58 - mustCIDR(t, "2001:db8::/32"),
59 - })
60 - t.Cleanup(func() {
61 - SetTrustedProxyCIDRs(nil)
62 - })
63 -
64 - if !IsTrustedProxyRemoteAddr("10.9.8.7:443") {
65 - t.Fatal("expected IPv4 remote to match trusted CIDR")
66 - }
67 - if !IsTrustedProxyRemoteAddr("[2001:db8::1]:443") {
68 - t.Fatal("expected IPv6 remote to match trusted CIDR")
69 - }
70 - if IsTrustedProxyRemoteAddr("198.51.100.2:443") {
71 - t.Fatal("did not expect non-allowlisted remote to be trusted")
72 - }
73 -}
74 -
75 -func TestIsIPBannedByPolicy(t *testing.T) {
76 - ipFilter := NewIPFilter()
77 - ipFilter.BanIP("203.0.113.22")
78 -
79 - tests := []struct {
80 - name string
81 - filter *IPFilter
82 - candidate string
83 - want bool
84 - }{
85 - {
86 - name: "nil manager",
87 - filter: nil,
88 - candidate: "203.0.113.22",
89 - want: false,
90 - },
91 - {
92 - name: "empty candidate",
93 - filter: ipFilter,
94 - candidate: " ",
95 - want: false,
96 - },
97 - {
98 - name: "trimmed banned ip",
99 - filter: ipFilter,
100 - candidate: " 203.0.113.22 ",
101 - want: true,
102 - },
103 - {
104 - name: "not banned ip",
105 - filter: ipFilter,
106 - candidate: "203.0.113.99",
107 - want: false,
108 - },
109 - }
110 -
111 - for _, tt := range tests {
112 - t.Run(tt.name, func(t *testing.T) {
113 - if got := IsIPBannedByPolicy(tt.filter, tt.candidate); got != tt.want {
114 - t.Fatalf("IsIPBannedByPolicy(%q)=%v, want %v", tt.candidate, got, tt.want)
115 - }
116 - })
117 - }
118 -}
portal/policy/rate_limiter.go deleted
-518
@@ -1,518 +0,0 @@
1 -package policy
2 -
3 -import (
4 - "io"
5 - "maps"
6 - "net"
7 - "sync"
8 - "sync/atomic"
9 - "time"
10 -
11 - "github.com/rs/zerolog/log"
12 -)
13 -
14 -// RateLimiter manages per-lease bytes-per-second rate limiting.
15 -type RateLimiter struct {
16 - bpsLimits map[string]int64
17 - bpsBuckets map[string]*Bucket
18 - mu sync.Mutex
19 -}
20 -
21 -// NewRateLimiter creates a new BPS manager.
22 -func NewRateLimiter() *RateLimiter {
23 - return &RateLimiter{
24 - bpsLimits: make(map[string]int64),
25 - bpsBuckets: make(map[string]*Bucket),
26 - }
27 -}
28 -
29 -// SetBPSLimit sets the BPS limit for a lease.
30 -func (m *RateLimiter) SetBPSLimit(leaseID string, bps int64) {
31 - m.mu.Lock()
32 - defer m.mu.Unlock()
33 - if bps <= 0 {
34 - delete(m.bpsLimits, leaseID)
35 - delete(m.bpsBuckets, leaseID)
36 - return
37 - }
38 - m.bpsLimits[leaseID] = bps
39 - // Reset bucket to apply new rate
40 - delete(m.bpsBuckets, leaseID)
41 -}
42 -
43 -// GetBPSLimit returns the BPS limit for a lease (0 = unlimited).
44 -func (m *RateLimiter) GetBPSLimit(leaseID string) int64 {
45 - m.mu.Lock()
46 - defer m.mu.Unlock()
47 - if v, ok := m.bpsLimits[leaseID]; ok {
48 - return v
49 - }
50 - return 0
51 -}
52 -
53 -// GetAllBPSLimits returns a copy of all BPS limits.
54 -func (m *RateLimiter) GetAllBPSLimits() map[string]int64 {
55 - m.mu.Lock()
56 - defer m.mu.Unlock()
57 - result := make(map[string]int64, len(m.bpsLimits))
58 - maps.Copy(result, m.bpsLimits)
59 - return result
60 -}
61 -
62 -// GetBucket returns a rate limit bucket for a lease, creating one if needed.
63 -func (m *RateLimiter) GetBucket(leaseID string) *Bucket {
64 - m.mu.Lock()
65 - defer m.mu.Unlock()
66 -
67 - bps, ok := m.bpsLimits[leaseID]
68 - if !ok || bps <= 0 {
69 - return nil // No limit
70 - }
71 -
72 - if bucket, exists := m.bpsBuckets[leaseID]; exists {
73 - return bucket
74 - }
75 -
76 - // Create new bucket
77 - bucket := NewBucket(bps, bps)
78 - m.bpsBuckets[leaseID] = bucket
79 - log.Debug().
80 - Str("lease_id", leaseID).
81 - Int64("bps", bps).
82 - Msg("[BPS] Created rate limit bucket")
83 - return bucket
84 -}
85 -
86 -// CleanupLease removes BPS data for a lease.
87 -func (m *RateLimiter) CleanupLease(leaseID string) {
88 - m.mu.Lock()
89 - defer m.mu.Unlock()
90 - delete(m.bpsLimits, leaseID)
91 - delete(m.bpsBuckets, leaseID)
92 -}
93 -
94 -// Copy copies data with rate limiting.
95 -func (m *RateLimiter) Copy(dst io.Writer, src io.Reader, leaseID string) (int64, error) {
96 - bucket := m.GetBucket(leaseID)
97 - return Copy(dst, src, bucket)
98 -}
99 -
100 -// CopyInterruptible copies data with rate limiting, interruptible via done channel.
101 -func (m *RateLimiter) CopyInterruptible(dst io.Writer, src io.Reader, leaseID string, done <-chan struct{}) (int64, error) {
102 - bucket := m.GetBucket(leaseID)
103 - return CopyInterruptible(dst, src, bucket, done)
104 -}
105 -
106 -// EstablishRelayWithBPS sets up bidirectional relay with BPS limiting.
107 -// In the new TLS passthrough architecture, this uses net.Conn.
108 -func EstablishRelayWithBPS(clientConn, leaseConn net.Conn, leaseID string, bpsManager *RateLimiter) {
109 - log.Info().
110 - Str("lease_id", leaseID).
111 - Msg("[Relay] Starting relay connection")
112 -
113 - defer func() {
114 - log.Info().
115 - Str("lease_id", leaseID).
116 - Msg("[Relay] Relay connection closed")
117 - }()
118 -
119 - var wg sync.WaitGroup
120 - wg.Add(2)
121 -
122 - // done is closed when either copy direction finishes, unblocking
123 - // any rate-limit sleep in the other direction.
124 - done := make(chan struct{})
125 - var doneOnce sync.Once
126 - closeDone := func() { doneOnce.Do(func() { close(done) }) }
127 -
128 - if bpsManager == nil {
129 - go func() {
130 - defer wg.Done()
131 - _, _ = io.Copy(leaseConn, clientConn)
132 - closeDone()
133 - _ = leaseConn.Close()
134 - }()
135 -
136 - go func() {
137 - defer wg.Done()
138 - _, _ = io.Copy(clientConn, leaseConn)
139 - closeDone()
140 - _ = clientConn.Close()
141 - }()
142 - } else {
143 - bpsLimit := bpsManager.GetBPSLimit(leaseID)
144 - log.Info().
145 - Str("lease_id", leaseID).
146 - Int64("bps_limit", bpsLimit).
147 - Msg("[Relay] Starting relay connection with rate limit")
148 -
149 - // Client -> Lease
150 - go func() {
151 - defer wg.Done()
152 - _, _ = bpsManager.CopyInterruptible(leaseConn, clientConn, leaseID, done)
153 - closeDone()
154 - if err := leaseConn.Close(); err != nil {
155 - log.Debug().Err(err).Str("lease_id", leaseID).Msg("[Relay] failed to close lease connection")
156 - }
157 - }()
158 -
159 - // Lease -> Client
160 - go func() {
161 - defer wg.Done()
162 - _, _ = bpsManager.CopyInterruptible(clientConn, leaseConn, leaseID, done)
163 - closeDone()
164 - if err := clientConn.Close(); err != nil {
165 - log.Debug().Err(err).Str("lease_id", leaseID).Msg("[Relay] failed to close client connection")
166 - }
167 - }()
168 - }
169 -
170 - wg.Wait()
171 -}
172 -
173 -// Bucket is a thread-safe rate limiter that supports multiple concurrent connections
174 -// sharing the same bandwidth limit. It uses a token bucket algorithm where tokens
175 -// represent bytes, and the bucket refills at the configured rate.
176 -type Bucket struct {
177 - lastRefill time.Time
178 - rateBps int64
179 - tokens float64
180 - maxTokens float64
181 - totalBytes int64
182 - totalWaited int64
183 - throttleHits int64
184 - mu sync.Mutex
185 -}
186 -
187 -// NewBucket creates a limiter for rateBps with burst bytes.
188 -// Multiple connections can share this bucket for fair bandwidth distribution.
189 -func NewBucket(rateBps int64, burst int64) *Bucket {
190 - if rateBps <= 0 {
191 - return nil
192 - }
193 - if burst <= 0 {
194 - burst = rateBps // default burst = 1 second worth
195 - }
196 - return &Bucket{
197 - rateBps: rateBps,
198 - tokens: float64(burst), // start with full burst
199 - maxTokens: float64(burst),
200 - lastRefill: time.Now(),
201 - }
202 -}
203 -
204 -// Take requests n bytes from the bucket. If not enough tokens are available,
205 -// it waits until sufficient tokens accumulate. This ensures fair distribution
206 -// among multiple concurrent connections sharing the same bucket.
207 -func (b *Bucket) Take(n int64) {
208 - if b == nil || n <= 0 {
209 - return
210 - }
211 -
212 - needed := float64(n)
213 -
214 - for {
215 - b.mu.Lock()
216 -
217 - // Refill tokens based on elapsed time
218 - now := time.Now()
219 - elapsed := now.Sub(b.lastRefill).Seconds()
220 - b.tokens += elapsed * float64(b.rateBps)
221 - if b.tokens > b.maxTokens {
222 - b.tokens = b.maxTokens
223 - }
224 - b.lastRefill = now
225 -
226 - if b.tokens >= needed {
227 - // Enough tokens available, consume them
228 - b.tokens -= needed
229 - b.mu.Unlock()
230 - atomic.AddInt64(&b.totalBytes, n)
231 - return
232 - }
233 -
234 - // Calculate how long to wait for enough tokens
235 - deficit := needed - b.tokens
236 - waitTime := time.Duration(deficit / float64(b.rateBps) * float64(time.Second))
237 -
238 - // Take whatever tokens are available now
239 - if b.tokens > 0 {
240 - needed -= b.tokens
241 - b.tokens = 0
242 - }
243 -
244 - b.mu.Unlock()
245 -
246 - // Wait for tokens to accumulate
247 - if waitTime > 0 {
248 - atomic.AddInt64(&b.throttleHits, 1)
249 - atomic.AddInt64(&b.totalWaited, int64(waitTime))
250 - log.Debug().
251 - Int64("bytes_requested", n).
252 - Int64("rate_bps", b.rateBps).
253 - Dur("wait_time", waitTime).
254 - Msg("[RateLimit] Throttling - waiting for bandwidth")
255 - time.Sleep(waitTime)
256 - }
257 - }
258 -}
259 -
260 -// TakeInterruptible requests n bytes from the bucket, like Take, but the wait
261 -// can be interrupted via the done channel. Returns false if interrupted.
262 -func (b *Bucket) TakeInterruptible(n int64, done <-chan struct{}) bool {
263 - if b == nil || n <= 0 {
264 - return true
265 - }
266 -
267 - needed := float64(n)
268 -
269 - for {
270 - b.mu.Lock()
271 -
272 - now := time.Now()
273 - elapsed := now.Sub(b.lastRefill).Seconds()
274 - b.tokens += elapsed * float64(b.rateBps)
275 - if b.tokens > b.maxTokens {
276 - b.tokens = b.maxTokens
277 - }
278 - b.lastRefill = now
279 -
280 - if b.tokens >= needed {
281 - b.tokens -= needed
282 - b.mu.Unlock()
283 - atomic.AddInt64(&b.totalBytes, n)
284 - return true
285 - }
286 -
287 - deficit := needed - b.tokens
288 - waitTime := time.Duration(deficit / float64(b.rateBps) * float64(time.Second))
289 -
290 - if b.tokens > 0 {
291 - needed -= b.tokens
292 - b.tokens = 0
293 - }
294 -
295 - b.mu.Unlock()
296 -
297 - if waitTime > 0 {
298 - atomic.AddInt64(&b.throttleHits, 1)
299 - atomic.AddInt64(&b.totalWaited, int64(waitTime))
300 - log.Debug().
301 - Int64("bytes_requested", n).
302 - Int64("rate_bps", b.rateBps).
303 - Dur("wait_time", waitTime).
304 - Msg("[RateLimit] Throttling - waiting for bandwidth")
305 - timer := time.NewTimer(waitTime)
306 - select {
307 - case <-done:
308 - timer.Stop()
309 - return false
310 - case <-timer.C:
311 - }
312 - }
313 - }
314 -}
315 -
316 -// TakeWithTimeout requests n bytes but returns false if it would take longer
317 -// than maxWait to acquire them. Returns true if tokens were acquired.
318 -func (b *Bucket) TakeWithTimeout(n int64, maxWait time.Duration) bool {
319 - if b == nil || n <= 0 {
320 - return true
321 - }
322 -
323 - deadline := time.Now().Add(maxWait)
324 - needed := float64(n)
325 -
326 - for {
327 - b.mu.Lock()
328 -
329 - // Refill tokens
330 - now := time.Now()
331 - elapsed := now.Sub(b.lastRefill).Seconds()
332 - b.tokens += elapsed * float64(b.rateBps)
333 - if b.tokens > b.maxTokens {
334 - b.tokens = b.maxTokens
335 - }
336 - b.lastRefill = now
337 -
338 - if b.tokens >= needed {
339 - b.tokens -= needed
340 - b.mu.Unlock()
341 - atomic.AddInt64(&b.totalBytes, n)
342 - return true
343 - }
344 -
345 - // Check if we have time to wait
346 - deficit := needed - b.tokens
347 - waitTime := time.Duration(deficit / float64(b.rateBps) * float64(time.Second))
348 -
349 - if now.Add(waitTime).After(deadline) {
350 - b.mu.Unlock()
351 - return false // Would take too long
352 - }
353 -
354 - if b.tokens > 0 {
355 - needed -= b.tokens
356 - b.tokens = 0
357 - }
358 -
359 - b.mu.Unlock()
360 -
361 - if waitTime > 0 {
362 - atomic.AddInt64(&b.throttleHits, 1)
363 - atomic.AddInt64(&b.totalWaited, int64(waitTime))
364 - time.Sleep(waitTime)
365 - }
366 - }
367 -}
368 -
369 -// Available returns the current number of available tokens (bytes).
370 -func (b *Bucket) Available() float64 {
371 - if b == nil {
372 - return 0
373 - }
374 - b.mu.Lock()
375 - defer b.mu.Unlock()
376 -
377 - // Refill first
378 - now := time.Now()
379 - elapsed := now.Sub(b.lastRefill).Seconds()
380 - b.tokens += elapsed * float64(b.rateBps)
381 - if b.tokens > b.maxTokens {
382 - b.tokens = b.maxTokens
383 - }
384 - b.lastRefill = now
385 -
386 - return b.tokens
387 -}
388 -
389 -// Rate returns the configured rate in bytes per second.
390 -func (b *Bucket) Rate() int64 {
391 - if b == nil {
392 - return 0
393 - }
394 - return b.rateBps
395 -}
396 -
397 -// Stats returns current statistics.
398 -func (b *Bucket) Stats() (totalBytes, throttleHits int64, totalWaited time.Duration) {
399 - return atomic.LoadInt64(&b.totalBytes),
400 - atomic.LoadInt64(&b.throttleHits),
401 - time.Duration(atomic.LoadInt64(&b.totalWaited))
402 -}
403 -
404 -// internal buffer pool for Copy - 64KB reduces Take() call frequency and lock contention
405 -// Using *[]byte to avoid interface boxing allocation in sync.Pool.
406 -var bufPool = sync.Pool{New: func() any {
407 - b := make([]byte, 64*1024)
408 - return &b
409 -}}
410 -
411 -// Copy copies from src to dst, enforcing the provided byte-rate bucket if not nil.
412 -// Multiple Copy calls sharing the same bucket will fairly share the bandwidth.
413 -// Returns bytes written and any copy error encountered.
414 -func Copy(dst io.Writer, src io.Reader, b *Bucket) (int64, error) {
415 - if b == nil {
416 - return io.Copy(dst, src)
417 - }
418 - buf := *bufPool.Get().(*[]byte)
419 - defer bufPool.Put(&buf)
420 -
421 - var total int64
422 - startTime := time.Now()
423 -
424 - for {
425 - nr, er := src.Read(buf)
426 - if nr > 0 {
427 - // Take rate limit BEFORE writing - this delays the write if needed
428 - // Multiple connections sharing this bucket will wait fairly
429 - b.Take(int64(nr))
430 - nw, ew := dst.Write(buf[:nr])
431 - if nw > 0 {
432 - total += int64(nw)
433 - }
434 - if ew != nil {
435 - logCopyStats(b, total, startTime)
436 - return total, ew
437 - }
438 - if nr != nw {
439 - logCopyStats(b, total, startTime)
440 - return total, io.ErrShortWrite
441 - }
442 - }
443 - if er != nil {
444 - if er == io.EOF {
445 - break
446 - }
447 - logCopyStats(b, total, startTime)
448 - return total, er
449 - }
450 - }
451 - logCopyStats(b, total, startTime)
452 - return total, nil
453 -}
454 -
455 -// CopyInterruptible copies from src to dst with rate limiting, interruptible via done channel.
456 -// When done is closed, the current rate-limit wait is interrupted and the copy returns.
457 -func CopyInterruptible(dst io.Writer, src io.Reader, b *Bucket, done <-chan struct{}) (int64, error) {
458 - if b == nil {
459 - return io.Copy(dst, src)
460 - }
461 - buf := *bufPool.Get().(*[]byte)
462 - defer bufPool.Put(&buf)
463 -
464 - var total int64
465 - startTime := time.Now()
466 -
467 - for {
468 - nr, er := src.Read(buf)
469 - if nr > 0 {
470 - if !b.TakeInterruptible(int64(nr), done) {
471 - logCopyStats(b, total, startTime)
472 - return total, nil
473 - }
474 - nw, ew := dst.Write(buf[:nr])
475 - if nw > 0 {
476 - total += int64(nw)
477 - }
478 - if ew != nil {
479 - logCopyStats(b, total, startTime)
480 - return total, ew
481 - }
482 - if nr != nw {
483 - logCopyStats(b, total, startTime)
484 - return total, io.ErrShortWrite
485 - }
486 - }
487 - if er != nil {
488 - if er == io.EOF {
489 - break
490 - }
491 - logCopyStats(b, total, startTime)
492 - return total, er
493 - }
494 - }
495 - logCopyStats(b, total, startTime)
496 - return total, nil
497 -}
498 -
499 -// logCopyStats logs summary statistics when copy completes.
500 -func logCopyStats(b *Bucket, totalBytes int64, startTime time.Time) {
501 - if b == nil || totalBytes == 0 {
502 - return
503 - }
504 - elapsed := time.Since(startTime)
505 - if elapsed > 0 {
506 - actualBps := float64(totalBytes) / elapsed.Seconds()
507 - throttleHits := atomic.LoadInt64(&b.throttleHits)
508 - totalWaited := time.Duration(atomic.LoadInt64(&b.totalWaited))
509 - log.Debug().
510 - Int64("total_bytes", totalBytes).
511 - Int64("rate_limit_bps", b.rateBps).
512 - Float64("actual_bps", actualBps).
513 - Dur("elapsed", elapsed).
514 - Int64("throttle_hits", throttleHits).
515 - Dur("total_waited", totalWaited).
516 - Msg("[RateLimit] Copy completed")
517 - }
518 -}
portal/registry.go deleted
-248
@@ -1,248 +0,0 @@
1 -package portal
2 -
3 -import (
4 - "crypto/subtle"
5 - "fmt"
6 - "net"
7 - "net/http"
8 - "strings"
9 - "time"
10 -
11 - "github.com/rs/zerolog/log"
12 -
13 - "gosuda.org/portal/types"
14 -)
15 -
16 -// DefaultLeaseTTL defines the default lease lifetime across relay components.
17 -const DefaultLeaseTTL = 30 * time.Second
18 -
19 -// RegistryAdmissionInput describes runtime context for control-plane admission checks.
20 -type RegistryAdmissionInput struct {
21 - RawLeaseID string
22 - RawReverseToken string
23 - ClientIP string
24 - IsClientIPBanned bool
25 - RequireExisting bool
26 -}
27 -
28 -// RegistryAdmissionResult returns normalized, validated admission context.
29 -type RegistryAdmissionResult struct {
30 - Entry *types.LeaseEntry
31 - LeaseID string
32 - ReverseToken string
33 - ClientIP string
34 -}
35 -
36 -// RegistryRegisterInput describes a lease registration request.
37 -type RegistryRegisterInput struct {
38 - LeaseID string
39 - ReverseToken string
40 - Name string
41 - Metadata *types.Metadata
42 - PortalURL string
43 - TLS bool
44 -}
45 -
46 -// AdmitControlPlane validates and normalizes control-plane credentials before SDK operations.
47 -func (g *RelayServer) AdmitControlPlane(input RegistryAdmissionInput) (RegistryAdmissionResult, *types.APIError) {
48 - leaseID, reverseToken := normalizeRegistryCredentials(input.RawLeaseID, input.RawReverseToken)
49 - if err := validateRegistryCredentials(leaseID, reverseToken); err != nil {
50 - return RegistryAdmissionResult{}, err
51 - }
52 -
53 - if input.IsClientIPBanned {
54 - return RegistryAdmissionResult{}, registryAPIError(http.StatusForbidden, "ip_banned", "ip is banned")
55 - }
56 -
57 - if g == nil || g.leaseManager == nil {
58 - return RegistryAdmissionResult{}, registryAPIError(http.StatusInternalServerError, "registry_unavailable", "registry service unavailable")
59 - }
60 -
61 - entry, exists := g.leaseManager.GetLeaseByID(leaseID)
62 - if input.RequireExisting && !exists {
63 - return RegistryAdmissionResult{}, registryAPIError(http.StatusNotFound, "lease_not_found", "lease not found")
64 - }
65 -
66 - if exists && !matchLeaseToken(entry.Lease.ReverseToken, reverseToken) {
67 - return RegistryAdmissionResult{}, registryAPIError(http.StatusUnauthorized, "unauthorized", "unauthorized reverse connect")
68 - }
69 -
70 - return RegistryAdmissionResult{
71 - LeaseID: leaseID,
72 - ReverseToken: reverseToken,
73 - ClientIP: strings.TrimSpace(input.ClientIP),
74 - Entry: entry,
75 - }, nil
76 -}
77 -
78 -// RegisterLease creates a new lease and associated SNI route.
79 -func (g *RelayServer) RegisterLease(input RegistryRegisterInput) (types.RegisterResponse, *types.APIError) {
80 - if g == nil || g.leaseManager == nil || g.reverseHub == nil || g.sniRouter == nil {
81 - return types.RegisterResponse{}, registryAPIError(http.StatusInternalServerError, "registry_unavailable", "registry service unavailable")
82 - }
83 -
84 - name := strings.TrimSpace(input.Name)
85 - if !types.IsValidServiceName(name) {
86 - return types.RegisterResponse{}, registryAPIError(http.StatusBadRequest, "invalid_name", "name must be a DNS label (letters, digits, hyphen; no dots or underscores)")
87 - }
88 - if !input.TLS {
89 - return types.RegisterResponse{}, registryAPIError(http.StatusBadRequest, "tls_required", "tls must be enabled")
90 - }
91 -
92 - metadata := types.Metadata{}
93 - if input.Metadata != nil {
94 - metadata = *input.Metadata
95 - }
96 -
97 - lease := &types.Lease{
98 - ID: input.LeaseID,
99 - Name: name,
100 - Metadata: metadata,
101 - Expires: time.Now().Add(DefaultLeaseTTL),
102 - TLS: true,
103 - ReverseToken: input.ReverseToken,
104 - }
105 -
106 - if !g.leaseManager.UpdateLease(lease) {
107 - return types.RegisterResponse{}, registryAPIError(http.StatusConflict, "lease_rejected", "failed to register lease (name conflict or policy violation)")
108 - }
109 - g.reverseHub.ClearDropped(input.LeaseID)
110 -
111 - sniName := types.BuildSNIName(name, g.BaseHost)
112 - if sniName == "" {
113 - g.leaseManager.DeleteLease(input.LeaseID)
114 - return types.RegisterResponse{}, registryAPIError(http.StatusInternalServerError, "sni_name_invalid", "failed to build SNI route name")
115 - }
116 - if err := g.sniRouter.RegisterRoute(sniName, input.LeaseID, name); err != nil {
117 - g.leaseManager.DeleteLease(input.LeaseID)
118 - return types.RegisterResponse{}, registryAPIError(http.StatusInternalServerError, "sni_register_failed", fmt.Sprintf("failed to register SNI route: %v", err))
119 - }
120 -
121 - log.Info().
122 - Str("lease_id", input.LeaseID).
123 - Str("name", name).
124 - Bool("tls", true).
125 - Msg("[Registry] Lease registered")
126 -
127 - return types.RegisterResponse{
128 - LeaseID: input.LeaseID,
129 - PublicURL: types.ServicePublicURL(strings.TrimSpace(input.PortalURL), name),
130 - Success: true,
131 - }, nil
132 -}
133 -
134 -// UnregisterLease removes lease state, route state, and reverse-connection state.
135 -func (g *RelayServer) UnregisterLease(leaseID string) {
136 - leaseID = strings.TrimSpace(leaseID)
137 - if leaseID == "" || g == nil {
138 - return
139 - }
140 -
141 - if g.leaseManager != nil && g.leaseManager.DeleteLease(leaseID) {
142 - // DeleteLease callback (handleLeaseDeleted) already called
143 - // DropLease + UnregisterRouteByLeaseID, so we're done.
144 - log.Info().
145 - Str("lease_id", leaseID).
146 - Msg("[Registry] Lease unregistered")
147 - return
148 - }
149 -
150 - // Lease was already removed (e.g. TTL expiry) — defensive cleanup.
151 - if g.sniRouter != nil {
152 - g.sniRouter.UnregisterRouteByLeaseID(leaseID)
153 - }
154 - if g.reverseHub != nil {
155 - g.reverseHub.DropLease(leaseID)
156 - }
157 -}
158 -
159 -// RenewLease extends lease expiry and opportunistically refreshes SNI routing.
160 -func (g *RelayServer) RenewLease(entry *types.LeaseEntry) *types.APIError {
161 - if entry == nil || entry.Lease == nil {
162 - return registryAPIError(http.StatusNotFound, "lease_not_found", "lease not found")
163 - }
164 - if g == nil || g.leaseManager == nil || g.sniRouter == nil {
165 - return registryAPIError(http.StatusInternalServerError, "registry_unavailable", "registry service unavailable")
166 - }
167 -
168 - entry.Lease.Expires = time.Now().Add(DefaultLeaseTTL)
169 - if !g.leaseManager.UpdateLease(entry.Lease) {
170 - return registryAPIError(http.StatusInternalServerError, "renew_failed", "failed to renew lease")
171 - }
172 -
173 - sniName := types.BuildSNIName(entry.Lease.Name, g.BaseHost)
174 - if sniName == "" {
175 - log.Warn().
176 - Str("lease_id", entry.Lease.ID).
177 - Str("name", entry.Lease.Name).
178 - Str("base_host", g.BaseHost).
179 - Msg("[Registry] Skipping SNI route refresh due to invalid SNI name")
180 - return nil
181 - }
182 - if err := g.sniRouter.RegisterRoute(sniName, entry.Lease.ID, entry.Lease.Name); err != nil {
183 - log.Warn().
184 - Err(err).
185 - Str("lease_id", entry.Lease.ID).
186 - Str("name", entry.Lease.Name).
187 - Msg("[Registry] Failed to refresh SNI route on renew")
188 - }
189 - return nil
190 -}
191 -
192 -// RegistryDomain returns the configured relay base domain.
193 -func (g *RelayServer) RegistryDomain() (types.DomainResponse, *types.APIError) {
194 - if g == nil {
195 - return types.DomainResponse{}, registryAPIError(http.StatusServiceUnavailable, "base_domain_missing", "base domain not configured")
196 - }
197 - baseHost := strings.TrimSpace(g.BaseHost)
198 - if baseHost == "" {
199 - return types.DomainResponse{}, registryAPIError(http.StatusServiceUnavailable, "base_domain_missing", "base domain not configured")
200 - }
201 - return types.DomainResponse{
202 - Success: true,
203 - BaseDomain: baseHost,
204 - }, nil
205 -}
206 -
207 -// HandleRegistryConnect admits reverse traffic into the reverse hub.
208 -func (g *RelayServer) HandleRegistryConnect(conn net.Conn, admission RegistryAdmissionResult) {
209 - if g == nil || g.reverseHub == nil {
210 - if conn != nil {
211 - _ = conn.Close()
212 - }
213 - return
214 - }
215 - g.reverseHub.HandleConnect(conn, admission.LeaseID, admission.ReverseToken, admission.ClientIP)
216 -}
217 -
218 -// matchLeaseToken compares lease-bound values in constant time.
219 -func matchLeaseToken(expected, provided string) bool {
220 - expected = strings.TrimSpace(expected)
221 - provided = strings.TrimSpace(provided)
222 - if expected == "" || provided == "" {
223 - return false
224 - }
225 - return subtle.ConstantTimeCompare([]byte(expected), []byte(provided)) == 1
226 -}
227 -
228 -func normalizeRegistryCredentials(rawLeaseID, rawReverseToken string) (leaseID, reverseToken string) {
229 - return strings.TrimSpace(rawLeaseID), strings.TrimSpace(rawReverseToken)
230 -}
231 -
232 -func validateRegistryCredentials(leaseID, reverseToken string) *types.APIError {
233 - if leaseID == "" {
234 - return registryAPIError(http.StatusBadRequest, "missing_lease_id", "lease_id is required")
235 - }
236 - if reverseToken == "" {
237 - return registryAPIError(http.StatusBadRequest, "missing_reverse_token", "reverse_token is required")
238 - }
239 - return nil
240 -}
241 -
242 -func registryAPIError(statusCode int, code, message string) *types.APIError {
243 - return &types.APIError{
244 - StatusCode: statusCode,
245 - Code: code,
246 - Message: message,
247 - }
248 -}
portal/registry_test.go deleted
-146
@@ -1,146 +0,0 @@
1 -package portal
2 -
3 -import (
4 - "testing"
5 - "time"
6 -
7 - "gosuda.org/portal/portal/sni"
8 - "gosuda.org/portal/types"
9 -)
10 -
11 -func newTestRegistryRelay(baseHost string) *RelayServer {
12 - s := &RelayServer{
13 - BaseHost: baseHost,
14 - leaseManager: NewLeaseManager(DefaultLeaseTTL),
15 - reverseHub: NewReverseHub(),
16 - sniRouter: sni.NewRouter(":0"),
17 - }
18 - s.bindLeaseLifecycleHooks()
19 - return s
20 -}
21 -
22 -func newTestLease(id, name, token string) *types.Lease {
23 - return &types.Lease{
24 - ID: id,
25 - Name: name,
26 - ReverseToken: token,
27 - TLS: true,
28 - Expires: time.Now().Add(2 * time.Minute),
29 - }
30 -}
31 -
32 -func TestMatchLeaseToken(t *testing.T) {
33 - t.Parallel()
34 -
35 - tests := []struct {
36 - name string
37 - expected string
38 - provided string
39 - want bool
40 - }{
41 - {name: "exact match", expected: "token-1", provided: "token-1", want: true},
42 - {name: "trimmed match", expected: " token-1 ", provided: "\ttoken-1\n", want: true},
43 - {name: "mismatch", expected: "token-1", provided: "token-2", want: false},
44 - {name: "empty expected", expected: "", provided: "token-1", want: false},
45 - {name: "empty provided", expected: "token-1", provided: " ", want: false},
46 - }
47 -
48 - for _, tt := range tests {
49 - t.Run(tt.name, func(t *testing.T) {
50 - t.Parallel()
51 - if got := matchLeaseToken(tt.expected, tt.provided); got != tt.want {
52 - t.Fatalf("matchLeaseToken(%q, %q)=%t, want %t", tt.expected, tt.provided, got, tt.want)
53 - }
54 - })
55 - }
56 -}
57 -
58 -func TestRegisterLease(t *testing.T) {
59 - t.Parallel()
60 -
61 - serv := newTestRegistryRelay("example.com")
62 -
63 - resp, apiErr := serv.RegisterLease(RegistryRegisterInput{
64 - LeaseID: "lease-1",
65 - ReverseToken: "token-1",
66 - Name: "demo",
67 - TLS: true,
68 - PortalURL: "https://portal.example.com",
69 - })
70 - if apiErr != nil {
71 - t.Fatalf("RegisterLease returned error: %+v", apiErr)
72 - }
73 - if !resp.Success {
74 - t.Fatal("expected success response")
75 - }
76 - if _, ok := serv.leaseManager.GetLeaseByID("lease-1"); !ok {
77 - t.Fatal("expected lease to be persisted")
78 - }
79 - sniName := types.BuildSNIName("demo", "example.com")
80 - if _, ok := serv.sniRouter.GetRoute(sniName); !ok {
81 - t.Fatalf("expected SNI route %q to be registered", sniName)
82 - }
83 -
84 - _, apiErr = serv.RegisterLease(RegistryRegisterInput{
85 - LeaseID: "lease-2",
86 - ReverseToken: "token-2",
87 - Name: "demo2",
88 - TLS: false,
89 - })
90 - if apiErr == nil || apiErr.Code != "tls_required" {
91 - t.Fatalf("expected tls_required error, got %+v", apiErr)
92 - }
93 -}
94 -
95 -func TestRenewAndUnregisterLease(t *testing.T) {
96 - t.Parallel()
97 -
98 - serv := newTestRegistryRelay("example.com")
99 - lease := newTestLease("lease-1", "demo", "token-1")
100 - if !serv.leaseManager.UpdateLease(lease) {
101 - t.Fatal("failed to seed lease")
102 - }
103 - if err := serv.sniRouter.RegisterRoute(types.BuildSNIName("demo", "example.com"), "lease-1", "demo"); err != nil {
104 - t.Fatalf("seed route: %v", err)
105 - }
106 - entry, _ := serv.leaseManager.GetLeaseByID("lease-1")
107 - oldExpires := entry.Lease.Expires
108 -
109 - if apiErr := serv.RenewLease(entry); apiErr != nil {
110 - t.Fatalf("RenewLease returned error: %+v", apiErr)
111 - }
112 - if entry.Lease.Expires.Equal(oldExpires) {
113 - t.Fatalf("renewed expiry did not change: %v", entry.Lease.Expires)
114 - }
115 - remaining := time.Until(entry.Lease.Expires)
116 - if remaining < 20*time.Second || remaining > 40*time.Second {
117 - t.Fatalf("renewed expiry remaining=%v, want around %v", remaining, DefaultLeaseTTL)
118 - }
119 -
120 - serv.UnregisterLease("lease-1")
121 - if _, ok := serv.leaseManager.GetLeaseByID("lease-1"); ok {
122 - t.Fatal("expected lease to be removed")
123 - }
124 - if _, ok := serv.sniRouter.GetRouteByLeaseID("lease-1"); ok {
125 - t.Fatal("expected SNI route to be removed")
126 - }
127 -}
128 -
129 -func TestRegistryDomain(t *testing.T) {
130 - t.Parallel()
131 -
132 - serv := newTestRegistryRelay("example.com")
133 - resp, apiErr := serv.RegistryDomain()
134 - if apiErr != nil {
135 - t.Fatalf("RegistryDomain returned error: %+v", apiErr)
136 - }
137 - if !resp.Success || resp.BaseDomain != "example.com" {
138 - t.Fatalf("unexpected domain response: %+v", resp)
139 - }
140 -
141 - serv.BaseHost = ""
142 - _, apiErr = serv.RegistryDomain()
143 - if apiErr == nil || apiErr.Code != "base_domain_missing" {
144 - t.Fatalf("expected base_domain_missing error, got %+v", apiErr)
145 - }
146 -}
portal/relay.go deleted
-231
@@ -1,231 +0,0 @@
1 -package portal
2 -
3 -import (
4 - "context"
5 - "fmt"
6 - "net"
7 - "strings"
8 - "sync"
9 - "time"
10 -
11 - "github.com/rs/zerolog/log"
12 -
13 - "gosuda.org/portal/portal/acme"
14 - "gosuda.org/portal/portal/keyless"
15 - "gosuda.org/portal/portal/sni"
16 -)
17 -
18 -type RelayServer struct {
19 - leaseManager *LeaseManager
20 - reverseHub *ReverseHub
21 - sniRouter *sni.Router
22 - acmeManager *acme.Manager
23 - keylessSigner *keyless.Signer
24 - BaseHost string
25 - address []string
26 - stopOnce sync.Once
27 -}
28 -
29 -// NewRelayServer creates a new relay server.
30 -func NewRelayServer(
31 - ctx context.Context,
32 - address []string,
33 - sniPort string,
34 - baseHost string,
35 - keylessDir string,
36 - cloudflareToken string,
37 -) (*RelayServer, error) {
38 - server := &RelayServer{
39 - BaseHost: baseHost,
40 - address: address,
41 - leaseManager: NewLeaseManager(DefaultLeaseTTL),
42 - reverseHub: NewReverseHub(),
43 - sniRouter: sni.NewRouter(sniPort),
44 - }
45 -
46 - // Auto-register DNS A records in Cloudflare (best-effort, non-fatal).
47 - if err := acme.EnsureDNSRecords(ctx, baseHost, cloudflareToken); err != nil {
48 - log.Warn().Err(err).Msg("[DNS] failed to auto-register DNS records; continuing without")
49 - }
50 -
51 - acmeManager, keyFile, err := acme.NewManager(ctx, acme.Config{
52 - BaseDomain: baseHost,
53 - KeyDir: keylessDir,
54 - CloudflareToken: cloudflareToken,
55 - })
56 - if err != nil {
57 - return nil, err
58 - }
59 - server.acmeManager = acmeManager
60 -
61 - signer, err := keyless.NewSigner(keyFile)
62 - if err != nil {
63 - return nil, fmt.Errorf("configure keyless signer: %w", err)
64 - }
65 - server.keylessSigner = signer
66 - if signer != nil {
67 - log.Info().
68 - Str("key_id", signer.KeyID()).
69 - Msg("[signer] keyless signer enabled at /v1/sign")
70 - }
71 -
72 - server.bindLeaseLifecycleHooks()
73 - server.bindReverseConnectAuthorizer()
74 - return server, nil
75 -}
76 -
77 -func (g *RelayServer) bindLeaseLifecycleHooks() {
78 - if g == nil || g.leaseManager == nil {
79 - return
80 - }
81 - g.leaseManager.SetOnLeaseDeleted(g.handleLeaseDeleted)
82 -}
83 -
84 -func (g *RelayServer) handleLeaseDeleted(leaseID string) {
85 - leaseID = strings.TrimSpace(leaseID)
86 - if leaseID == "" {
87 - return
88 - }
89 -
90 - if g.reverseHub != nil {
91 - g.reverseHub.DropLease(leaseID)
92 - }
93 - if g.sniRouter != nil {
94 - g.sniRouter.UnregisterRouteByLeaseID(leaseID)
95 - }
96 -}
97 -
98 -func (g *RelayServer) bindReverseConnectAuthorizer() {
99 - if g == nil || g.reverseHub == nil {
100 - return
101 - }
102 - g.reverseHub.SetAuthorizer(g.authorizeReverseConnect)
103 -}
104 -
105 -func (g *RelayServer) authorizeReverseConnect(leaseID, token string) bool {
106 - if g == nil || g.leaseManager == nil {
107 - return false
108 - }
109 -
110 - entry, ok := g.leaseManager.GetLeaseByID(leaseID)
111 - if !ok || entry == nil || entry.Lease == nil {
112 - return false
113 - }
114 -
115 - return matchLeaseToken(entry.Lease.ReverseToken, token)
116 -}
117 -
118 -// GetLeaseManager returns the lease manager instance.
119 -func (g *RelayServer) GetLeaseManager() *LeaseManager {
120 - return g.leaseManager
121 -}
122 -
123 -// GetReverseHub returns the reverse hub instance.
124 -func (g *RelayServer) GetReverseHub() *ReverseHub {
125 - return g.reverseHub
126 -}
127 -
128 -// GetSNIRouter returns the SNI router instance.
129 -func (g *RelayServer) GetSNIRouter() *sni.Router {
130 - return g.sniRouter
131 -}
132 -
133 -// GetKeylessSigner returns relay keyless signer when configured.
134 -func (g *RelayServer) GetKeylessSigner() *keyless.Signer {
135 - return g.keylessSigner
136 -}
137 -
138 -// GetACMEManager returns relay ACME manager.
139 -func (g *RelayServer) GetACMEManager() *acme.Manager {
140 - return g.acmeManager
141 -}
142 -
143 -// ConfigurePortalRootFallback forwards unmatched root-domain SNI traffic to the provided upstream listener.
144 -func (g *RelayServer) ConfigurePortalRootFallback(rootSNI, upstreamAddr string) {
145 - if g == nil || g.sniRouter == nil {
146 - return
147 - }
148 -
149 - rootSNI = strings.TrimSpace(rootSNI)
150 - if rootSNI == "" {
151 - return
152 - }
153 -
154 - upstreamAddr = strings.TrimSpace(upstreamAddr)
155 - if upstreamAddr == "" {
156 - log.Warn().
157 - Msg("[RelayServer] root-domain SNI fallback upstream is empty; fallback disabled")
158 - return
159 - }
160 -
161 - g.sniRouter.SetNoRouteHandler(func(clientConn net.Conn, serverName string) bool {
162 - return g.handleRootFallback(clientConn, serverName, rootSNI, upstreamAddr)
163 - })
164 -}
165 -
166 -func (g *RelayServer) handleRootFallback(clientConn net.Conn, serverName, rootSNI, upstreamAddr string) bool {
167 - if !strings.EqualFold(strings.TrimSpace(serverName), rootSNI) {
168 - return false
169 - }
170 -
171 - dialer := &net.Dialer{Timeout: 5 * time.Second}
172 - upstreamConn, err := dialer.DialContext(context.Background(), "tcp", upstreamAddr)
173 - if err != nil {
174 - log.Warn().
175 - Err(err).
176 - Str("sni", serverName).
177 - Str("upstream", upstreamAddr).
178 - Msg("[SNI] failed to forward root domain to admin/API listener")
179 - if closeErr := clientConn.Close(); closeErr != nil {
180 - log.Debug().Err(closeErr).Str("sni", serverName).Msg("[SNI] failed to close client connection")
181 - }
182 - return true
183 - }
184 -
185 - log.Debug().
186 - Str("sni", serverName).
187 - Str("upstream", upstreamAddr).
188 - Msg("[SNI] forwarding root domain to admin/API listener")
189 - sni.BridgeConnections(clientConn, upstreamConn)
190 - return true
191 -}
192 -
193 -// Start starts the relay server.
194 -func (g *RelayServer) Start() error {
195 - g.leaseManager.Start()
196 -
197 - if err := g.sniRouter.Start(); err != nil {
198 - log.Error().Err(err).Str("addr", g.sniRouter.GetAddr()).Msg("[RelayServer] Failed to start SNI router")
199 - return err
200 - }
201 - log.Info().Str("addr", g.sniRouter.GetAddr()).Msg("[RelayServer] SNI router started")
202 -
203 - // Start ACME renewal loop
204 - if g.acmeManager != nil {
205 - g.acmeManager.Start(context.Background())
206 - }
207 -
208 - log.Info().Msg("[RelayServer] Started")
209 - return nil
210 -}
211 -
212 -// Stop stops the relay server.
213 -func (g *RelayServer) Stop() {
214 - g.stopOnce.Do(func() {
215 - if g.leaseManager != nil {
216 - g.leaseManager.Stop()
217 - }
218 - if g.reverseHub != nil {
219 - g.reverseHub.Shutdown()
220 - }
221 - if g.sniRouter != nil {
222 - if err := g.sniRouter.Stop(); err != nil {
223 - log.Warn().Err(err).Msg("[RelayServer] Failed to stop SNI router")
224 - }
225 - }
226 - if g.acmeManager != nil {
227 - g.acmeManager.Stop()
228 - }
229 - log.Info().Msg("[RelayServer] Stopped")
230 - })
231 -}
portal/relay_test.go deleted
-128
@@ -1,128 +0,0 @@
1 -package portal
2 -
3 -import (
4 - "context"
5 - "net"
6 - "strings"
7 - "testing"
8 - "time"
9 -
10 - "gosuda.org/portal/types"
11 -)
12 -
13 -func TestRelayServerReverseHubAuthorizerTrimsToken(t *testing.T) {
14 - serv, err := NewRelayServer(context.Background(), nil, ":0", "example.com", "", "")
15 - if err != nil {
16 - t.Fatalf("create relay server: %v", err)
17 - }
18 -
19 - lease := &types.Lease{
20 - ID: "lease-authorizer-trim",
21 - Name: "tenant",
22 - TLS: true,
23 - ReverseToken: " reverse-token ",
24 - Expires: time.Now().Add(time.Minute),
25 - }
26 - if ok := serv.GetLeaseManager().UpdateLease(lease); !ok {
27 - t.Fatal("failed to register lease in lease manager")
28 - }
29 -
30 - hub := serv.GetReverseHub()
31 - if !hub.isAuthorized(lease.ID, "reverse-token") {
32 - t.Fatal("expected authorizer to accept trimmed reverse token")
33 - }
34 - if !hub.isAuthorized(lease.ID, " reverse-token ") {
35 - t.Fatal("expected authorizer to accept token with surrounding whitespace")
36 - }
37 - if hub.isAuthorized(lease.ID, "wrong-token") {
38 - t.Fatal("expected authorizer to reject wrong token")
39 - }
40 -}
41 -
42 -func TestRelayServerHandleLeaseDeletedDropsRouteAndPool(t *testing.T) {
43 - serv, err := NewRelayServer(context.Background(), nil, ":0", "example.com", "", "")
44 - if err != nil {
45 - t.Fatalf("create relay server: %v", err)
46 - }
47 -
48 - const (
49 - leaseID = "lease-delete-hook"
50 - leaseSNI = "tenant.example.com"
51 - )
52 -
53 - if err := serv.GetSNIRouter().RegisterRoute(leaseSNI, leaseID, "tenant"); err != nil {
54 - t.Fatalf("register route: %v", err)
55 - }
56 -
57 - local, peer := net.Pipe()
58 - defer func() { _ = peer.Close() }()
59 - conn := NewReverseConn(local)
60 - defer conn.Close()
61 -
62 - if ok := serv.GetReverseHub().Offer(leaseID, conn); !ok {
63 - t.Fatal("offer failed")
64 - }
65 -
66 - serv.handleLeaseDeleted(" " + leaseID + " ")
67 -
68 - if _, ok := serv.GetSNIRouter().GetRoute(leaseSNI); ok {
69 - t.Fatal("expected route to be removed when lease is deleted")
70 - }
71 -
72 - _, acquireErr := serv.GetReverseHub().AcquireForTLS(leaseID, 100*time.Millisecond)
73 - if acquireErr == nil {
74 - t.Fatal("expected reverse pool to be dropped when lease is deleted")
75 - }
76 - if !strings.Contains(acquireErr.Error(), "no tunnel available") {
77 - t.Fatalf("unexpected acquire error: %v", acquireErr)
78 - }
79 -}
80 -
81 -func TestRelayServerHandleRootFallbackRequiresRootHostMatch(t *testing.T) {
82 - serv, err := NewRelayServer(context.Background(), nil, ":0", "example.com", "", "")
83 - if err != nil {
84 - t.Fatalf("create relay server: %v", err)
85 - }
86 -
87 - client, peer := net.Pipe()
88 - defer func() {
89 - _ = client.Close()
90 - _ = peer.Close()
91 - }()
92 -
93 - handled := serv.handleRootFallback(client, "tenant.example.com", "portal.example.com", "127.0.0.1:4017")
94 - if handled {
95 - t.Fatal("expected non-root SNI to bypass fallback handler")
96 - }
97 -}
98 -
99 -func TestRelayServerHandleRootFallbackClosesClientOnDialFailure(t *testing.T) {
100 - serv, err := NewRelayServer(context.Background(), nil, ":0", "example.com", "", "")
101 - if err != nil {
102 - t.Fatalf("create relay server: %v", err)
103 - }
104 -
105 - client, peer := net.Pipe()
106 - defer func() { _ = peer.Close() }()
107 -
108 - handled := serv.handleRootFallback(client, "portal.example.com", "portal.example.com", "invalid-upstream")
109 - if !handled {
110 - t.Fatal("expected root SNI to be handled by fallback path")
111 - }
112 -
113 - _ = peer.SetReadDeadline(time.Now().Add(250 * time.Millisecond))
114 - var b [1]byte
115 - if _, err := peer.Read(b[:]); err == nil {
116 - t.Fatal("expected client connection to be closed after fallback dial failure")
117 - }
118 -}
119 -
120 -func TestRelayServerStopIsIdempotent(t *testing.T) {
121 - serv, err := NewRelayServer(context.Background(), nil, ":0", "example.com", "", "")
122 - if err != nil {
123 - t.Fatalf("create relay server: %v", err)
124 - }
125 -
126 - serv.Stop()
127 - serv.Stop()
128 -}
portal/reverse_hub.go deleted
-453
@@ -1,453 +0,0 @@
1 -package portal
2 -
3 -import (
4 - "fmt"
5 - "net"
6 - "strings"
7 - "sync"
8 - "sync/atomic"
9 - "time"
10 -
11 - "github.com/rs/zerolog/log"
12 -
13 - "gosuda.org/portal/types"
14 -)
15 -
16 -const (
17 - // QueueSize is the maximum number of pending reverse connections per lease.
18 - QueueSize = 64
19 -
20 - // DefaultAcquireTimeout is the default timeout for acquiring a reverse connection.
21 - DefaultAcquireTimeout = 2 * time.Second
22 -
23 - // TLSAcquireWait is the timeout for TLS passthrough connections.
24 - TLSAcquireWait = 2 * time.Second
25 -
26 - // AuthFailureDelay is the delay before closing unauthorized connections (rate limiting).
27 - AuthFailureDelay = 2 * time.Second
28 -
29 - // controlWriteTimeout bounds control-marker writes on reverse connections.
30 - controlWriteTimeout = 2 * time.Second
31 -
32 - // ReverseIdleKeepaliveInterval sends an idle keepalive byte to reduce
33 - // reverse connection disconnections from intermediate idle timeouts.
34 - ReverseIdleKeepaliveInterval = 25 * time.Second
35 -)
36 -
37 -// ReverseConn wraps a net.Conn with lifecycle management for the connection pool.
38 -type ReverseConn struct {
39 - Conn net.Conn
40 - done chan struct{}
41 - active chan struct{}
42 - once sync.Once
43 - // activateOnce ensures active channel is closed exactly once.
44 - activateOnce sync.Once
45 - // writeMu serializes writes while the connection is idle.
46 - writeMu sync.Mutex
47 - // closed tracks local close to help queue consumers skip stale entries.
48 - closed atomic.Bool
49 -}
50 -
51 -// NewReverseConn creates a new pooled connection.
52 -func NewReverseConn(conn net.Conn) *ReverseConn {
53 - return &ReverseConn{
54 - Conn: conn,
55 - done: make(chan struct{}),
56 - active: make(chan struct{}),
57 - }
58 -}
59 -
60 -// Close closes the connection and signals completion.
61 -func (c *ReverseConn) Close() {
62 - c.closed.Store(true)
63 - _ = c.Conn.Close()
64 - c.once.Do(func() {
65 - close(c.done)
66 - })
67 -}
68 -
69 -// Wait blocks until the connection is closed.
70 -func (c *ReverseConn) Wait() {
71 - <-c.done
72 -}
73 -
74 -func (c *ReverseConn) IsClosed() bool {
75 - return c == nil || c.closed.Load()
76 -}
77 -
78 -func (c *ReverseConn) Activate() {
79 - if c == nil {
80 - return
81 - }
82 - c.activateOnce.Do(func() {
83 - close(c.active)
84 - })
85 -}
86 -
87 -func (c *ReverseConn) WriteControlByte(marker byte, timeout time.Duration) error {
88 - if c == nil || c.Conn == nil {
89 - return net.ErrClosed
90 - }
91 - c.writeMu.Lock()
92 - defer c.writeMu.Unlock()
93 -
94 - if timeout > 0 {
95 - _ = c.Conn.SetWriteDeadline(time.Now().Add(timeout))
96 - defer c.Conn.SetWriteDeadline(time.Time{})
97 - }
98 - _, err := c.Conn.Write([]byte{marker})
99 - return err
100 -}
101 -
102 -type ReverseHub struct {
103 - pools map[string]chan *ReverseConn
104 - dropped map[string]struct{}
105 - authorizer func(leaseID, token string) bool
106 - ipBanChecker func(ip string) bool
107 - onAccepted func(leaseID, ip string)
108 - stopCh chan struct{}
109 - mu sync.RWMutex
110 - stopOnce sync.Once
111 -}
112 -
113 -// NewReverseHub creates a new reverse connection hub.
114 -func NewReverseHub() *ReverseHub {
115 - return &ReverseHub{
116 - pools: make(map[string]chan *ReverseConn),
117 - dropped: make(map[string]struct{}),
118 - stopCh: make(chan struct{}),
119 - }
120 -}
121 -
122 -func (h *ReverseHub) getOrCreatePool(leaseID string) chan *ReverseConn {
123 - select {
124 - case <-h.stopCh:
125 - return nil
126 - default:
127 - }
128 -
129 - h.mu.Lock()
130 - defer h.mu.Unlock()
131 -
132 - if _, dropped := h.dropped[leaseID]; dropped {
133 - return nil
134 - }
135 -
136 - if pool, ok := h.pools[leaseID]; ok {
137 - return pool
138 - }
139 - pool := make(chan *ReverseConn, QueueSize)
140 - h.pools[leaseID] = pool
141 - return pool
142 -}
143 -
144 -func (h *ReverseHub) getPool(leaseID string) (chan *ReverseConn, bool) {
145 - h.mu.RLock()
146 - defer h.mu.RUnlock()
147 - pool, ok := h.pools[leaseID]
148 - return pool, ok
149 -}
150 -
151 -// SetAuthorizer sets the authentication function for new connections.
152 -func (h *ReverseHub) SetAuthorizer(authorizer func(leaseID, token string) bool) {
153 - h.mu.Lock()
154 - defer h.mu.Unlock()
155 - h.authorizer = authorizer
156 -}
157 -
158 -// SetIPBanChecker sets optional IP ban check for reverse connections.
159 -func (h *ReverseHub) SetIPBanChecker(checker func(ip string) bool) {
160 - h.mu.Lock()
161 - defer h.mu.Unlock()
162 - h.ipBanChecker = checker
163 -}
164 -
165 -// SetOnAccepted sets optional callback for authorized reverse connections.
166 -func (h *ReverseHub) SetOnAccepted(onAccepted func(leaseID, ip string)) {
167 - h.mu.Lock()
168 - defer h.mu.Unlock()
169 - h.onAccepted = onAccepted
170 -}
171 -
172 -func (h *ReverseHub) isAuthorized(leaseID, token string) bool {
173 - h.mu.RLock()
174 - authorizer := h.authorizer
175 - h.mu.RUnlock()
176 - if authorizer == nil {
177 - return false
178 - }
179 - return authorizer(leaseID, token)
180 -}
181 -
182 -func (h *ReverseHub) isIPBanned(ip string) bool {
183 - h.mu.RLock()
184 - checker := h.ipBanChecker
185 - h.mu.RUnlock()
186 - ip = strings.TrimSpace(ip)
187 - if checker == nil || ip == "" {
188 - return false
189 - }
190 - return checker(ip)
191 -}
192 -
193 -func (h *ReverseHub) notifyAccepted(leaseID, ip string) {
194 - h.mu.RLock()
195 - onAccepted := h.onAccepted
196 - h.mu.RUnlock()
197 - if onAccepted == nil {
198 - return
199 - }
200 - onAccepted(leaseID, ip)
201 -}
202 -
203 -func (h *ReverseHub) Offer(leaseID string, conn *ReverseConn) bool {
204 - leaseID = strings.TrimSpace(leaseID)
205 - if leaseID == "" || conn == nil || conn.Conn == nil {
206 - return false
207 - }
208 -
209 - pool := h.getOrCreatePool(leaseID)
210 - if pool == nil {
211 - return false
212 - }
213 -
214 - for range QueueSize + 1 {
215 - select {
216 - case pool <- conn:
217 - // Re-check: DropLease may have run between getOrCreatePool and send.
218 - h.mu.RLock()
219 - _, dropped := h.dropped[leaseID]
220 - h.mu.RUnlock()
221 - if dropped {
222 - // Pool was drained by DropLease; conn may still be in channel.
223 - // Close it defensively — DropLease drain will also close if it gets it.
224 - conn.Close()
225 - return false
226 - }
227 - return true
228 - default:
229 - }
230 -
231 - // Pool full: evict one oldest entry and retry.
232 - select {
233 - case old := <-pool:
234 - if old != nil {
235 - old.Close()
236 - }
237 - default:
238 - }
239 - }
240 -
241 - return false
242 -}
243 -
244 -func (h *ReverseHub) AcquireForTLS(leaseID string, timeout time.Duration) (*ReverseConn, error) {
245 - leaseID = strings.TrimSpace(leaseID)
246 -
247 - pool, ok := h.getPool(leaseID)
248 - if !ok {
249 - return nil, fmt.Errorf("no tunnel available for lease %s", leaseID)
250 - }
251 -
252 - if timeout <= 0 {
253 - timeout = DefaultAcquireTimeout
254 - }
255 -
256 - timer := time.NewTimer(timeout)
257 - defer timer.Stop()
258 -
259 - deadline := time.Now().Add(timeout)
260 - for {
261 - remaining := time.Until(deadline)
262 - if remaining <= 0 {
263 - return nil, fmt.Errorf("tunnel acquisition timeout for lease %s", leaseID)
264 - }
265 - if !timer.Stop() {
266 - select {
267 - case <-timer.C:
268 - default:
269 - }
270 - }
271 - timer.Reset(remaining)
272 -
273 - select {
274 - case conn := <-pool:
275 - if conn == nil || conn.IsClosed() {
276 - continue
277 - }
278 - // Stop idle keepalive and signal tunnel worker to release this connection.
279 - conn.Activate()
280 - err := conn.WriteControlByte(types.TLSStartMarker, controlWriteTimeout)
281 - if err == nil {
282 - return conn, nil
283 - }
284 -
285 - log.Warn().
286 - Err(err).
287 - Str("lease_id", leaseID).
288 - Msg("[ReverseHub] Failed to send TLS start marker; retrying with new connection")
289 - conn.Close()
290 - continue
291 - case <-timer.C:
292 - return nil, fmt.Errorf("tunnel acquisition timeout for lease %s", leaseID)
293 - }
294 - }
295 -}
296 -
297 -func (h *ReverseHub) DropLease(leaseID string) {
298 - leaseID = strings.TrimSpace(leaseID)
299 - if leaseID == "" {
300 - return
301 - }
302 -
303 - h.mu.Lock()
304 - pool, ok := h.pools[leaseID]
305 - if ok {
306 - delete(h.pools, leaseID)
307 - }
308 - h.dropped[leaseID] = struct{}{}
309 - h.mu.Unlock()
310 -
311 - if !ok {
312 - return
313 - }
314 -
315 - // Drain and close pending connections
316 - for {
317 - select {
318 - case conn := <-pool:
319 - if conn != nil {
320 - conn.Close()
321 - }
322 - default:
323 - return
324 - }
325 - }
326 -}
327 -
328 -// ClearDropped removes a lease from the dropped set, allowing it to be re-registered.
329 -// This should be called when a lease is re-registered after being dropped.
330 -func (h *ReverseHub) ClearDropped(leaseID string) {
331 - leaseID = strings.TrimSpace(leaseID)
332 - if leaseID == "" {
333 - return
334 - }
335 -
336 - h.mu.Lock()
337 - delete(h.dropped, leaseID)
338 - h.mu.Unlock()
339 -}
340 -
341 -func (h *ReverseHub) HandleConnect(conn net.Conn, leaseID, token, remoteIP string) {
342 - if conn == nil {
343 - return
344 - }
345 -
346 - leaseID = strings.TrimSpace(leaseID)
347 - token = strings.TrimSpace(token)
348 - remoteIP = strings.TrimSpace(remoteIP)
349 -
350 - if leaseID == "" {
351 - log.Warn().Msg("[ReverseHub] Missing lease_id on reverse connect")
352 - h.rejectConn(conn, "[ReverseHub] failed to close unauthorized reverse connection")
353 - return
354 - }
355 -
356 - if h.isIPBanned(remoteIP) {
357 - log.Warn().
358 - Str("lease_id", leaseID).
359 - Str("ip", remoteIP).
360 - Msg("[ReverseHub] IP banned for reverse connect")
361 - h.rejectConn(conn, "[ReverseHub] failed to close banned reverse connection")
362 - return
363 - }
364 -
365 - if !h.isAuthorized(leaseID, token) {
366 - log.Warn().Str("lease_id", leaseID).Msg("[ReverseHub] Unauthorized reverse connect")
367 - h.rejectConn(conn, "[ReverseHub] failed to close unauthorized reverse connection")
368 - return
369 - }
370 -
371 - h.notifyAccepted(leaseID, remoteIP)
372 - reverseConn := NewReverseConn(conn)
373 - if !h.Offer(leaseID, reverseConn) {
374 - log.Warn().Str("lease_id", leaseID).Msg("[ReverseHub] Connection pool full for lease")
375 - reverseConn.Close()
376 - return
377 - }
378 -
379 - h.keepAliveWhileIdle(reverseConn, leaseID)
380 -
381 - // Wait until the connection is used and closed
382 - reverseConn.Wait()
383 -}
384 -
385 -func (h *ReverseHub) rejectConn(conn net.Conn, debugCloseMessage string) {
386 - time.Sleep(AuthFailureDelay)
387 - h.closeConn(conn, debugCloseMessage)
388 -}
389 -
390 -func (h *ReverseHub) closeConn(conn net.Conn, debugCloseMessage string) {
391 - if conn == nil {
392 - return
393 - }
394 - if err := conn.Close(); err != nil {
395 - log.Debug().Err(err).Msg(debugCloseMessage)
396 - }
397 -}
398 -
399 -// Shutdown closes the stop channel and drains all pools, causing idle
400 -// HandleConnect goroutines to unblock and return.
401 -func (h *ReverseHub) Shutdown() {
402 - h.stopOnce.Do(func() {
403 - close(h.stopCh)
404 -
405 - h.mu.Lock()
406 - pools := h.pools
407 - h.pools = make(map[string]chan *ReverseConn)
408 - h.mu.Unlock()
409 -
410 - for _, pool := range pools {
411 - drainPool(pool)
412 - }
413 - })
414 -}
415 -
416 -func drainPool(pool chan *ReverseConn) {
417 - for {
418 - select {
419 - case conn := <-pool:
420 - if conn != nil {
421 - conn.Close()
422 - }
423 - default:
424 - return
425 - }
426 - }
427 -}
428 -
429 -func (h *ReverseHub) keepAliveWhileIdle(conn *ReverseConn, leaseID string) {
430 - ticker := time.NewTicker(ReverseIdleKeepaliveInterval)
431 - defer ticker.Stop()
432 -
433 - for {
434 - select {
435 - case <-h.stopCh:
436 - conn.Close()
437 - return
438 - case <-conn.done:
439 - return
440 - case <-conn.active:
441 - return
442 - case <-ticker.C:
443 - if err := conn.WriteControlByte(types.ReverseKeepaliveMarker, controlWriteTimeout); err != nil {
444 - log.Debug().
445 - Err(err).
446 - Str("lease_id", leaseID).
447 - Msg("[ReverseHub] Idle keepalive write failed")
448 - conn.Close()
449 - return
450 - }
451 - }
452 - }
453 -}
portal/reverse_hub_test.go deleted
-446
@@ -1,446 +0,0 @@
1 -package portal
2 -
3 -import (
4 - "io"
5 - "net"
6 - "strings"
7 - "testing"
8 - "time"
9 -
10 - "gosuda.org/portal/types"
11 -)
12 -
13 -func TestReverseHubAuthorization(t *testing.T) {
14 - hub := NewReverseHub()
15 -
16 - if hub.isAuthorized("lease-1", "token-1") {
17 - t.Fatal("expected unauthorized when authorizer is not configured")
18 - }
19 -
20 - hub.SetAuthorizer(func(leaseID, token string) bool {
21 - return leaseID == "lease-1" && token == "token-1"
22 - })
23 -
24 - if !hub.isAuthorized("lease-1", "token-1") {
25 - t.Fatal("expected authorized")
26 - }
27 - if hub.isAuthorized("lease-1", "wrong-token") {
28 - t.Fatal("expected unauthorized for wrong token")
29 - }
30 -}
31 -
32 -func TestReverseHubIsIPBannedTrimsInput(t *testing.T) {
33 - hub := NewReverseHub()
34 - hub.SetIPBanChecker(func(ip string) bool {
35 - return ip == "203.0.113.50"
36 - })
37 -
38 - if !hub.isIPBanned(" 203.0.113.50 ") {
39 - t.Fatal("expected trimmed IP to be checked as banned")
40 - }
41 -}
42 -
43 -func TestReverseHubIsIPBannedSkipsEmptyInput(t *testing.T) {
44 - hub := NewReverseHub()
45 - called := false
46 - hub.SetIPBanChecker(func(string) bool {
47 - called = true
48 - return true
49 - })
50 -
51 - if hub.isIPBanned(" ") {
52 - t.Fatal("expected whitespace IP to be treated as not banned")
53 - }
54 - if called {
55 - t.Fatal("expected checker to be skipped for empty IP candidate")
56 - }
57 -}
58 -
59 -func TestReverseHubOfferRejectsInvalidInput(t *testing.T) {
60 - hub := NewReverseHub()
61 -
62 - if ok := hub.Offer("", nil); ok {
63 - t.Fatal("expected offer with empty lease and nil connection to fail")
64 - }
65 - if ok := hub.Offer(" ", nil); ok {
66 - t.Fatal("expected offer with whitespace lease and nil connection to fail")
67 - }
68 -}
69 -
70 -func TestHandleConnectTrimsLeaseIDAndToken(t *testing.T) {
71 - hub := NewReverseHub()
72 - leaseID := "lease-connect-trim"
73 - token := "token-connect-trim"
74 - hub.SetAuthorizer(func(gotLeaseID, gotToken string) bool {
75 - return gotLeaseID == leaseID && gotToken == token
76 - })
77 -
78 - local, peer := net.Pipe()
79 - defer func() {
80 - _ = peer.Close()
81 - }()
82 -
83 - done := make(chan struct{})
84 - go func() {
85 - hub.HandleConnect(local, " "+leaseID+" ", " "+token+" ", " 127.0.0.1 ")
86 - close(done)
87 - }()
88 -
89 - markerRead := make(chan byte, 1)
90 - readErr := make(chan error, 1)
91 - go func() {
92 - var b [1]byte
93 - _, err := io.ReadFull(peer, b[:])
94 - if err != nil {
95 - readErr <- err
96 - return
97 - }
98 - markerRead <- b[0]
99 - }()
100 -
101 - got, err := hub.AcquireForTLS(leaseID, 500*time.Millisecond)
102 - deadline := time.Now().Add(500 * time.Millisecond)
103 - for err != nil {
104 - if time.Now().After(deadline) {
105 - t.Fatalf("AcquireForTLS failed: %v", err)
106 - }
107 - time.Sleep(10 * time.Millisecond)
108 - got, err = hub.AcquireForTLS(leaseID, 25*time.Millisecond)
109 - }
110 - if got == nil {
111 - t.Fatal("AcquireForTLS returned nil connection")
112 - }
113 -
114 - select {
115 - case err := <-readErr:
116 - t.Fatalf("failed to read marker: %v", err)
117 - case marker := <-markerRead:
118 - if marker != types.TLSStartMarker {
119 - t.Fatalf("unexpected marker: %d", marker)
120 - }
121 - case <-time.After(500 * time.Millisecond):
122 - t.Fatal("timed out waiting for start marker")
123 - }
124 -
125 - got.Close()
126 -
127 - select {
128 - case <-done:
129 - case <-time.After(500 * time.Millisecond):
130 - t.Fatal("HandleConnect did not return after connection close")
131 - }
132 -}
133 -
134 -func TestAcquireForTLSSendsStartMarker(t *testing.T) {
135 - hub := NewReverseHub()
136 - leaseID := "lease-tls-marker"
137 -
138 - local, peer := net.Pipe()
139 - defer func() {
140 - _ = peer.Close()
141 - }()
142 - conn := NewReverseConn(local)
143 - defer conn.Close()
144 -
145 - if ok := hub.Offer(leaseID, conn); !ok {
146 - t.Fatal("offer failed")
147 - }
148 -
149 - markerRead := make(chan byte, 1)
150 - readErr := make(chan error, 1)
151 - go func() {
152 - var b [1]byte
153 - _, err := io.ReadFull(peer, b[:])
154 - if err != nil {
155 - readErr <- err
156 - return
157 - }
158 - markerRead <- b[0]
159 - }()
160 -
161 - got, err := hub.AcquireForTLS(leaseID, 500*time.Millisecond)
162 - if err != nil {
163 - t.Fatalf("AcquireForTLS failed: %v", err)
164 - }
165 - if got != conn {
166 - t.Fatal("AcquireForTLS returned unexpected connection")
167 - }
168 -
169 - select {
170 - case err := <-readErr:
171 - t.Fatalf("failed to read marker: %v", err)
172 - case b := <-markerRead:
173 - if b != types.TLSStartMarker {
174 - t.Fatalf("unexpected marker: %d", b)
175 - }
176 - case <-time.After(500 * time.Millisecond):
177 - t.Fatal("timed out waiting for start marker")
178 - }
179 -}
180 -
181 -func TestAcquireForTLSPollLoopSendsStartMarker(t *testing.T) {
182 - hub := NewReverseHub()
183 - leaseID := "lease-http-marker"
184 -
185 - local, peer := net.Pipe()
186 - defer func() {
187 - _ = peer.Close()
188 - }()
189 - conn := NewReverseConn(local)
190 - defer conn.Close()
191 -
192 - if ok := hub.Offer(leaseID, conn); !ok {
193 - t.Fatal("offer failed")
194 - }
195 -
196 - markerRead := make(chan byte, 1)
197 - readErr := make(chan error, 1)
198 - go func() {
199 - var b [1]byte
200 - _, err := io.ReadFull(peer, b[:])
201 - if err != nil {
202 - readErr <- err
203 - return
204 - }
205 - markerRead <- b[0]
206 - }()
207 -
208 - var (
209 - got *ReverseConn
210 - err error
211 - )
212 - deadline := time.Now().Add(500 * time.Millisecond)
213 - for {
214 - got, err = hub.AcquireForTLS(leaseID, 25*time.Millisecond)
215 - if err == nil {
216 - break
217 - }
218 - if time.Now().After(deadline) {
219 - t.Fatalf("AcquireForTLS failed: %v", err)
220 - }
221 - time.Sleep(10 * time.Millisecond)
222 - }
223 - if got != conn {
224 - t.Fatal("AcquireForTLS returned unexpected connection")
225 - }
226 -
227 - select {
228 - case err := <-readErr:
229 - t.Fatalf("failed to read marker: %v", err)
230 - case b := <-markerRead:
231 - if b != types.TLSStartMarker {
232 - t.Fatalf("unexpected marker: %d", b)
233 - }
234 - case <-time.After(500 * time.Millisecond):
235 - t.Fatal("timed out waiting for start marker")
236 - }
237 -}
238 -
239 -func TestHandleConnectOffersAuthorizedConn(t *testing.T) {
240 - hub := NewReverseHub()
241 - leaseID := "lease-connect"
242 - token := "reverse-token"
243 - accepted := make(chan struct{}, 1)
244 - hub.SetAuthorizer(func(gotLeaseID, gotToken string) bool {
245 - return gotLeaseID == leaseID && gotToken == token
246 - })
247 - hub.SetOnAccepted(func(gotLeaseID, _ string) {
248 - if gotLeaseID != leaseID {
249 - return
250 - }
251 - select {
252 - case accepted <- struct{}{}:
253 - default:
254 - }
255 - })
256 -
257 - local, peer := net.Pipe()
258 - defer func() {
259 - _ = peer.Close()
260 - }()
261 -
262 - done := make(chan struct{})
263 - go func() {
264 - hub.HandleConnect(local, leaseID, token, "127.0.0.1")
265 - close(done)
266 - }()
267 -
268 - select {
269 - case <-accepted:
270 - case <-time.After(500 * time.Millisecond):
271 - t.Fatal("HandleConnect did not reach accepted state")
272 - }
273 -
274 - markerRead := make(chan byte, 1)
275 - readErr := make(chan error, 1)
276 - go func() {
277 - var b [1]byte
278 - _, err := io.ReadFull(peer, b[:])
279 - if err != nil {
280 - readErr <- err
281 - return
282 - }
283 - markerRead <- b[0]
284 - }()
285 -
286 - got, err := hub.AcquireForTLS(leaseID, 500*time.Millisecond)
287 - if err != nil {
288 - t.Fatalf("AcquireForTLS failed: %v", err)
289 - }
290 - if got == nil {
291 - t.Fatal("AcquireForTLS returned nil connection")
292 - }
293 -
294 - select {
295 - case err := <-readErr:
296 - t.Fatalf("failed to read start marker: %v", err)
297 - case b := <-markerRead:
298 - if b != types.TLSStartMarker {
299 - t.Fatalf("unexpected marker: %d", b)
300 - }
301 - case <-time.After(500 * time.Millisecond):
302 - t.Fatal("timed out waiting for start marker")
303 - }
304 -
305 - got.Close()
306 -
307 - select {
308 - case <-done:
309 - case <-time.After(500 * time.Millisecond):
310 - t.Fatal("HandleConnect did not return after connection close")
311 - }
312 -}
313 -
314 -func TestHandleConnectUnauthorizedHonorsAuthDelay(t *testing.T) {
315 - hub := NewReverseHub()
316 - leaseID := "lease-auth-delay"
317 - hub.SetAuthorizer(func(gotLeaseID, gotToken string) bool {
318 - return gotLeaseID == leaseID && gotToken == "expected-token"
319 - })
320 -
321 - local, peer := net.Pipe()
322 - defer func() {
323 - _ = peer.Close()
324 - }()
325 -
326 - done := make(chan struct{})
327 - go func() {
328 - hub.HandleConnect(local, leaseID, "wrong-token", "127.0.0.1")
329 - close(done)
330 - }()
331 -
332 - select {
333 - case <-done:
334 - t.Fatal("HandleConnect returned before auth delay elapsed")
335 - case <-time.After(AuthFailureDelay / 2):
336 - }
337 -
338 - select {
339 - case <-done:
340 - case <-time.After(AuthFailureDelay + 500*time.Millisecond):
341 - t.Fatal("HandleConnect did not return after auth delay")
342 - }
343 -
344 - _ = peer.SetReadDeadline(time.Now().Add(250 * time.Millisecond))
345 - var b [1]byte
346 - if _, err := peer.Read(b[:]); err == nil {
347 - t.Fatal("expected connection to be closed after unauthorized connect")
348 - }
349 -}
350 -
351 -func TestHandleConnectBannedIPHonorsAuthDelay(t *testing.T) {
352 - hub := NewReverseHub()
353 - leaseID := "lease-banned-ip"
354 - token := "token-ok"
355 - blockedIP := "203.0.113.50"
356 -
357 - hub.SetAuthorizer(func(gotLeaseID, gotToken string) bool {
358 - return gotLeaseID == leaseID && gotToken == token
359 - })
360 - hub.SetIPBanChecker(func(ip string) bool {
361 - return ip == blockedIP
362 - })
363 -
364 - local, peer := net.Pipe()
365 - defer func() {
366 - _ = peer.Close()
367 - }()
368 -
369 - done := make(chan struct{})
370 - go func() {
371 - hub.HandleConnect(local, leaseID, token, blockedIP)
372 - close(done)
373 - }()
374 -
375 - select {
376 - case <-done:
377 - t.Fatal("HandleConnect returned before auth delay elapsed for banned IP")
378 - case <-time.After(AuthFailureDelay / 2):
379 - }
380 -
381 - select {
382 - case <-done:
383 - case <-time.After(AuthFailureDelay + 500*time.Millisecond):
384 - t.Fatal("HandleConnect did not return after banned-IP auth delay")
385 - }
386 -
387 - _, err := hub.AcquireForTLS(leaseID, 100*time.Millisecond)
388 - if err == nil {
389 - t.Fatal("expected no reverse tunnel for banned IP")
390 - }
391 - if !strings.Contains(err.Error(), "no tunnel available") {
392 - t.Fatalf("unexpected acquire error after banned IP connect: %v", err)
393 - }
394 -
395 - _ = peer.SetReadDeadline(time.Now().Add(250 * time.Millisecond))
396 - var b [1]byte
397 - if _, err := peer.Read(b[:]); err == nil {
398 - t.Fatal("expected banned reverse connection to be closed")
399 - }
400 -}
401 -
402 -func TestDropLeaseCleansPoolImmediately(t *testing.T) {
403 - hub := NewReverseHub()
404 - leaseID := "lease-drop"
405 -
406 - local, peer := net.Pipe()
407 - defer func() {
408 - _ = peer.Close()
409 - }()
410 - conn := NewReverseConn(local)
411 - defer conn.Close()
412 -
413 - if ok := hub.Offer(leaseID, conn); !ok {
414 - t.Fatal("offer failed")
415 - }
416 -
417 - start := time.Now()
418 - hub.DropLease(leaseID)
419 -
420 - _, err := hub.AcquireForTLS(leaseID, 2*time.Second)
421 - if err == nil {
422 - t.Fatal("expected acquire to fail after lease drop")
423 - }
424 - if !strings.Contains(err.Error(), "no tunnel available") {
425 - t.Fatalf("unexpected error after lease drop: %v", err)
426 - }
427 - if elapsed := time.Since(start); elapsed > 250*time.Millisecond {
428 - t.Fatalf("expected immediate cleanup after drop, acquire took %v", elapsed)
429 - }
430 -
431 - otherLocal, otherPeer := net.Pipe()
432 - defer func() {
433 - _ = otherPeer.Close()
434 - }()
435 - otherConn := NewReverseConn(otherLocal)
436 - defer otherConn.Close()
437 - if ok := hub.Offer(leaseID, otherConn); ok {
438 - t.Fatal("expected dropped lease to reject new offered connections")
439 - }
440 -
441 - _ = peer.SetReadDeadline(time.Now().Add(250 * time.Millisecond))
442 - var b [1]byte
443 - if _, err := peer.Read(b[:]); err == nil {
444 - t.Fatal("expected dropped pooled connection to be closed")
445 - }
446 -}
portal/routing.go new
+68
@@ -0,0 +1,68 @@
1 +package portal
2 +
3 +import (
4 + "strings"
5 + "sync"
6 +)
7 +
8 +type routeTable struct {
9 + mu sync.RWMutex
10 + exact map[string]string
11 +}
12 +
13 +func newRouteTable() *routeTable {
14 + return &routeTable{exact: make(map[string]string)}
15 +}
16 +
17 +func (t *routeTable) Set(host, leaseID string) {
18 + host = normalizeHostname(host)
19 + if host == "" {
20 + return
21 + }
22 + t.mu.Lock()
23 + defer t.mu.Unlock()
24 + t.exact[host] = leaseID
25 +}
26 +
27 +func (t *routeTable) Delete(host string) {
28 + host = normalizeHostname(host)
29 + if host == "" {
30 + return
31 + }
32 + t.mu.Lock()
33 + defer t.mu.Unlock()
34 + delete(t.exact, host)
35 +}
36 +
37 +func (t *routeTable) DeleteLease(hosts []string) {
38 + t.mu.Lock()
39 + defer t.mu.Unlock()
40 + for _, host := range hosts {
41 + delete(t.exact, normalizeHostname(host))
42 + }
43 +}
44 +
45 +func (t *routeTable) Lookup(host string) (string, bool) {
46 + host = normalizeHostname(host)
47 + if host == "" {
48 + return "", false
49 + }
50 +
51 + t.mu.RLock()
52 + defer t.mu.RUnlock()
53 +
54 + if leaseID, ok := t.exact[host]; ok {
55 + return leaseID, true
56 + }
57 +
58 + parts := stringsSplit(host, ".")
59 + if len(parts) < 3 {
60 + return "", false
61 + }
62 + wildcard := "*." + stringsJoin(parts[1:], ".")
63 + leaseID, ok := t.exact[wildcard]
64 + return leaseID, ok
65 +}
66 +
67 +func stringsSplit(s, sep string) []string { return strings.Split(s, sep) }
68 +func stringsJoin(parts []string, sep string) string { return strings.Join(parts, sep) }
portal/server.go new
+708
@@ -0,0 +1,708 @@
1 +package portal
2 +
3 +import (
4 + "context"
5 + "crypto/subtle"
6 + "crypto/tls"
7 + "encoding/json"
8 + "errors"
9 + "fmt"
10 + "io"
11 + "net"
12 + "net/http"
13 + "strings"
14 + "sync"
15 + "time"
16 +
17 + "github.com/gosuda/keyless_tls/relay/l4"
18 + "golang.org/x/sync/errgroup"
19 +)
20 +
21 +type ServerConfig struct {
22 + PortalURL string
23 + APIListenAddr string
24 + SNIListenAddr string
25 + RootHost string
26 + RootFallbackAddr string
27 + LeaseTTL time.Duration
28 + ClaimTimeout time.Duration
29 + IdleKeepaliveInterval time.Duration
30 + ReadyQueueLimit int
31 + ClientHelloTimeout time.Duration
32 + APITLS TLSMaterialConfig
33 + APIHandlerWrapper func(http.Handler) http.Handler
34 +}
35 +
36 +type Server struct {
37 + cfg ServerConfig
38 +
39 + apiServer *http.Server
40 + apiListener net.Listener
41 + sniListener net.Listener
42 + apiTLSClose io.Closer
43 +
44 + ctx context.Context
45 + cancel context.CancelFunc
46 + group *errgroup.Group
47 +
48 + routes *routeTable
49 +
50 + mu sync.RWMutex
51 + leases map[string]*leaseRecord
52 +
53 + shutdownOnce sync.Once
54 +}
55 +
56 +type leaseRecord struct {
57 + ID string
58 + Name string
59 + Hostnames []string
60 + Metadata LeaseMetadata
61 + ReverseToken string
62 + ExpiresAt time.Time
63 + Broker *leaseBroker
64 +}
65 +
66 +type LeaseSnapshot struct {
67 + ID string
68 + Name string
69 + Hostnames []string
70 + Metadata LeaseMetadata
71 + ExpiresAt time.Time
72 + Ready int
73 +}
74 +
75 +func NewServer(cfg ServerConfig) (*Server, error) {
76 + if cfg.APIListenAddr == "" {
77 + cfg.APIListenAddr = ":4017"
78 + }
79 + if cfg.SNIListenAddr == "" {
80 + cfg.SNIListenAddr = ":443"
81 + }
82 + cfg.LeaseTTL = durationOrDefault(cfg.LeaseTTL, defaultLeaseTTL)
83 + cfg.ClaimTimeout = durationOrDefault(cfg.ClaimTimeout, defaultClaimTimeout)
84 + cfg.IdleKeepaliveInterval = durationOrDefault(cfg.IdleKeepaliveInterval, defaultIdleKeepalive)
85 + cfg.ReadyQueueLimit = intOrDefault(cfg.ReadyQueueLimit, defaultReadyQueueLimit)
86 + cfg.ClientHelloTimeout = durationOrDefault(cfg.ClientHelloTimeout, defaultClientHelloWait)
87 + if cfg.RootHost == "" {
88 + cfg.RootHost = PortalRootHost(cfg.PortalURL)
89 + }
90 + if cfg.RootHost == "" {
91 + return nil, errors.New("root host is required")
92 + }
93 + if len(cfg.APITLS.CertPEM) == 0 {
94 + return nil, errors.New("api tls certificate is required")
95 + }
96 + if len(cfg.APITLS.KeyPEM) == 0 && cfg.APITLS.Keyless == nil {
97 + return nil, errors.New("api tls key or keyless signer is required")
98 + }
99 +
100 + return &Server{
101 + cfg: cfg,
102 + routes: newRouteTable(),
103 + leases: make(map[string]*leaseRecord),
104 + }, nil
105 +}
106 +
107 +func (s *Server) Start(ctx context.Context) error {
108 + if s.group != nil {
109 + return errors.New("server already started")
110 + }
111 +
112 + apiListener, err := net.Listen("tcp", s.cfg.APIListenAddr)
113 + if err != nil {
114 + return fmt.Errorf("listen api: %w", err)
115 + }
116 + sniListener, err := net.Listen("tcp", s.cfg.SNIListenAddr)
117 + if err != nil {
118 + _ = apiListener.Close()
119 + return fmt.Errorf("listen sni: %w", err)
120 + }
121 +
122 + serverCtx, cancel := context.WithCancel(ctx)
123 + group, groupCtx := errgroup.WithContext(serverCtx)
124 +
125 + apiServer := &http.Server{
126 + Handler: s.wrapAPIHandler(s.apiHandler()),
127 + ReadHeaderTimeout: 10 * time.Second,
128 + TLSNextProto: make(map[string]func(*http.Server, *tls.Conn, http.Handler)),
129 + }
130 + apiCloser, err := attachAPITLS(apiServer, s.cfg.APITLS)
131 + if err != nil {
132 + _ = apiListener.Close()
133 + _ = sniListener.Close()
134 + cancel()
135 + return fmt.Errorf("configure api tls: %w", err)
136 + }
137 +
138 + s.apiListener = tls.NewListener(apiListener, apiServer.TLSConfig)
139 + s.sniListener = sniListener
140 + s.apiServer = apiServer
141 + s.apiTLSClose = apiCloser
142 + s.ctx = groupCtx
143 + s.cancel = cancel
144 + s.group = group
145 +
146 + group.Go(s.runAPIServer)
147 + group.Go(s.runSNIListener)
148 + group.Go(s.runLeaseJanitor)
149 + group.Go(s.watchContext)
150 + return nil
151 +}
152 +
153 +func (s *Server) Wait() error {
154 + if s.group == nil {
155 + return nil
156 + }
157 + return s.group.Wait()
158 +}
159 +
160 +func (s *Server) Shutdown(ctx context.Context) error {
161 + var shutdownErr error
162 + s.shutdownOnce.Do(func() {
163 + if s.cancel != nil {
164 + s.cancel()
165 + }
166 +
167 + s.mu.Lock()
168 + for _, lease := range s.leases {
169 + lease.Broker.Stop()
170 + }
171 + s.mu.Unlock()
172 +
173 + if s.sniListener != nil {
174 + if err := s.sniListener.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
175 + shutdownErr = err
176 + }
177 + }
178 + if s.apiServer != nil {
179 + if err := s.apiServer.Shutdown(ctx); err != nil && shutdownErr == nil {
180 + shutdownErr = err
181 + }
182 + }
183 + if s.apiTLSClose != nil {
184 + _ = s.apiTLSClose.Close()
185 + }
186 + })
187 + return shutdownErr
188 +}
189 +
190 +func (s *Server) APIAddr() string {
191 + if s.apiListener == nil {
192 + return ""
193 + }
194 + return s.apiListener.Addr().String()
195 +}
196 +
197 +func (s *Server) SNIAddr() string {
198 + if s.sniListener == nil {
199 + return ""
200 + }
201 + return s.sniListener.Addr().String()
202 +}
203 +
204 +func (s *Server) GetLease(leaseID string) (LeaseSnapshot, bool) {
205 + s.mu.RLock()
206 + record, ok := s.leases[strings.TrimSpace(leaseID)]
207 + s.mu.RUnlock()
208 + if !ok {
209 + return LeaseSnapshot{}, false
210 + }
211 + return LeaseSnapshot{
212 + ID: record.ID,
213 + Name: record.Name,
214 + Hostnames: append([]string(nil), record.Hostnames...),
215 + Metadata: record.Metadata,
216 + ExpiresAt: record.ExpiresAt,
217 + Ready: record.Broker.ReadyCount(),
218 + }, true
219 +}
220 +
221 +func (s *Server) ListLeases() []LeaseSnapshot {
222 + s.mu.RLock()
223 + defer s.mu.RUnlock()
224 +
225 + out := make([]LeaseSnapshot, 0, len(s.leases))
226 + for _, record := range s.leases {
227 + out = append(out, LeaseSnapshot{
228 + ID: record.ID,
229 + Name: record.Name,
230 + Hostnames: append([]string(nil), record.Hostnames...),
231 + Metadata: record.Metadata,
232 + ExpiresAt: record.ExpiresAt,
233 + Ready: record.Broker.ReadyCount(),
234 + })
235 + }
236 + return out
237 +}
238 +
239 +func (s *Server) apiHandler() http.Handler {
240 + mux := http.NewServeMux()
241 + mux.HandleFunc("/", s.handleRoot)
242 + mux.HandleFunc("/healthz", s.handleHealthz)
243 + mux.HandleFunc("/sdk/domain", s.handleDomain)
244 + mux.HandleFunc("/sdk/register", s.handleRegister)
245 + mux.HandleFunc("/sdk/renew", s.handleRenew)
246 + mux.HandleFunc("/sdk/unregister", s.handleUnregister)
247 + mux.HandleFunc("/sdk/connect", s.handleConnect)
248 + return mux
249 +}
250 +
251 +func (s *Server) handleRoot(w http.ResponseWriter, _ *http.Request) {
252 + writeAPIData(w, http.StatusOK, map[string]any{
253 + "service": "portal-relay",
254 + "root": s.cfg.RootHost,
255 + })
256 +}
257 +
258 +func (s *Server) handleHealthz(w http.ResponseWriter, _ *http.Request) {
259 + writeAPIData(w, http.StatusOK, map[string]any{"status": "ok"})
260 +}
261 +
262 +func (s *Server) handleDomain(w http.ResponseWriter, r *http.Request) {
263 + if r.Method != http.MethodGet {
264 + writeAPIError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
265 + return
266 + }
267 + name := r.URL.Query().Get("name")
268 + writeAPIData(w, http.StatusOK, DomainResponse{
269 + RootHost: s.cfg.RootHost,
270 + SuggestedHostname: suggestHostname(name, s.cfg.RootHost),
271 + })
272 +}
273 +
274 +func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) {
275 + if r.Method != http.MethodPost {
276 + writeAPIError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
277 + return
278 + }
279 + var req RegisterRequest
280 + if err := decodeJSONBody(w, r, &req); err != nil {
281 + writeAPIError(w, http.StatusBadRequest, "invalid_json", err.Error())
282 + return
283 + }
284 + resp, err := s.registerLease(req)
285 + if err != nil {
286 + status, code := http.StatusBadRequest, "invalid_request"
287 + if errors.Is(err, errHostnameConflict) {
288 + status, code = http.StatusConflict, "hostname_conflict"
289 + }
290 + writeAPIError(w, status, code, err.Error())
291 + return
292 + }
293 + writeAPIData(w, http.StatusCreated, resp)
294 +}
295 +
296 +func (s *Server) handleRenew(w http.ResponseWriter, r *http.Request) {
297 + if r.Method != http.MethodPost {
298 + writeAPIError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
299 + return
300 + }
301 + var req RenewRequest
302 + if err := decodeJSONBody(w, r, &req); err != nil {
303 + writeAPIError(w, http.StatusBadRequest, "invalid_json", err.Error())
304 + return
305 + }
306 + resp, err := s.renewLease(req)
307 + if err != nil {
308 + status, code := http.StatusBadRequest, "invalid_request"
309 + if errors.Is(err, errLeaseNotFound) {
310 + status, code = http.StatusNotFound, "lease_not_found"
311 + }
312 + if errors.Is(err, errUnauthorized) {
313 + status, code = http.StatusForbidden, "unauthorized"
314 + }
315 + writeAPIError(w, status, code, err.Error())
316 + return
317 + }
318 + writeAPIData(w, http.StatusOK, resp)
319 +}
320 +
321 +func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
322 + if r.Method != http.MethodPost {
323 + writeAPIError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
324 + return
325 + }
326 + var req UnregisterRequest
327 + if err := decodeJSONBody(w, r, &req); err != nil {
328 + writeAPIError(w, http.StatusBadRequest, "invalid_json", err.Error())
329 + return
330 + }
331 + if err := s.unregisterLease(req); err != nil {
332 + status, code := http.StatusBadRequest, "invalid_request"
333 + if errors.Is(err, errLeaseNotFound) {
334 + status, code = http.StatusNotFound, "lease_not_found"
335 + }
336 + if errors.Is(err, errUnauthorized) {
337 + status, code = http.StatusForbidden, "unauthorized"
338 + }
339 + writeAPIError(w, status, code, err.Error())
340 + return
341 + }
342 + writeAPIOK(w, http.StatusOK)
343 +}
344 +
345 +func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
346 + if r.Method != http.MethodGet {
347 + writeAPIError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
348 + return
349 + }
350 + if r.ProtoMajor != 1 {
351 + writeAPIError(w, http.StatusHTTPVersionNotSupported, "http11_only", "reverse connect requires HTTP/1.1")
352 + return
353 + }
354 +
355 + leaseID := strings.TrimSpace(r.URL.Query().Get("lease_id"))
356 + token := strings.TrimSpace(r.Header.Get(HeaderReverseToken))
357 + lease, err := s.lookupLeaseByID(leaseID, token)
358 + if err != nil {
359 + status, code := http.StatusForbidden, "unauthorized"
360 + if errors.Is(err, errLeaseNotFound) {
361 + status, code = http.StatusNotFound, "lease_not_found"
362 + }
363 + writeAPIError(w, status, code, err.Error())
364 + return
365 + }
366 +
367 + hijacker, ok := w.(http.Hijacker)
368 + if !ok {
369 + writeAPIError(w, http.StatusInternalServerError, "hijack_unsupported", "hijacking is not supported")
370 + return
371 + }
372 +
373 + conn, rw, err := hijacker.Hijack()
374 + if err != nil {
375 + writeAPIError(w, http.StatusInternalServerError, "hijack_failed", err.Error())
376 + return
377 + }
378 +
379 + if _, err := rw.WriteString("HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: keep-alive\r\n\r\n"); err != nil {
380 + _ = conn.Close()
381 + return
382 + }
383 + if err := rw.Flush(); err != nil {
384 + _ = conn.Close()
385 + return
386 + }
387 +
388 + session := newReverseSession(conn, s.cfg.IdleKeepaliveInterval)
389 + if err := lease.Broker.Offer(session); err != nil {
390 + _ = session.Close()
391 + }
392 +}
393 +
394 +func (s *Server) registerLease(req RegisterRequest) (RegisterResponse, error) {
395 + if strings.TrimSpace(req.Name) == "" {
396 + return RegisterResponse{}, errors.New("name is required")
397 + }
398 + if strings.TrimSpace(req.ReverseToken) == "" {
399 + return RegisterResponse{}, errors.New("reverse token is required")
400 + }
401 + if !req.TLS {
402 + return RegisterResponse{}, errors.New("tls must be true")
403 + }
404 +
405 + hostnames := normalizeHostnames(req.Hostnames)
406 + if len(hostnames) == 0 {
407 + hostnames = []string{suggestHostname(req.Name, s.cfg.RootHost)}
408 + }
409 +
410 + s.mu.Lock()
411 + defer s.mu.Unlock()
412 +
413 + for _, host := range hostnames {
414 + if owner := s.findLeaseByHostnameLocked(host); owner != nil {
415 + return RegisterResponse{}, fmt.Errorf("%w: %s", errHostnameConflict, host)
416 + }
417 + }
418 +
419 + ttl := s.cfg.LeaseTTL
420 + if req.TTLSeconds > 0 {
421 + ttl = time.Duration(req.TTLSeconds) * time.Second
422 + }
423 +
424 + leaseID := randomID("lease_")
425 + expiresAt := time.Now().Add(ttl)
426 + record := &leaseRecord{
427 + ID: leaseID,
428 + Name: strings.TrimSpace(req.Name),
429 + Hostnames: hostnames,
430 + Metadata: normalizeMetadata(req.Metadata),
431 + ReverseToken: req.ReverseToken,
432 + ExpiresAt: expiresAt,
433 + Broker: newLeaseBroker(leaseID, s.cfg.IdleKeepaliveInterval, s.cfg.ReadyQueueLimit),
434 + }
435 +
436 + s.leases[leaseID] = record
437 + for _, host := range hostnames {
438 + s.routes.Set(host, leaseID)
439 + }
440 +
441 + return RegisterResponse{
442 + LeaseID: leaseID,
443 + Hostnames: append([]string(nil), hostnames...),
444 + Metadata: record.Metadata,
445 + ExpiresAt: expiresAt,
446 + ConnectURL: s.connectURL(),
447 + }, nil
448 +}
449 +
450 +func (s *Server) renewLease(req RenewRequest) (RenewResponse, error) {
451 + s.mu.Lock()
452 + defer s.mu.Unlock()
453 +
454 + record, ok := s.leases[strings.TrimSpace(req.LeaseID)]
455 + if !ok {
456 + return RenewResponse{}, errLeaseNotFound
457 + }
458 + if !tokenMatches(record.ReverseToken, req.ReverseToken) {
459 + return RenewResponse{}, errUnauthorized
460 + }
461 +
462 + ttl := s.cfg.LeaseTTL
463 + if req.TTLSeconds > 0 {
464 + ttl = time.Duration(req.TTLSeconds) * time.Second
465 + }
466 + record.ExpiresAt = time.Now().Add(ttl)
467 + record.Broker.Reset()
468 + return RenewResponse{LeaseID: record.ID, ExpiresAt: record.ExpiresAt}, nil
469 +}
470 +
471 +func (s *Server) unregisterLease(req UnregisterRequest) error {
472 + s.mu.Lock()
473 + record, ok := s.leases[strings.TrimSpace(req.LeaseID)]
474 + if !ok {
475 + s.mu.Unlock()
476 + return errLeaseNotFound
477 + }
478 + if !tokenMatches(record.ReverseToken, req.ReverseToken) {
479 + s.mu.Unlock()
480 + return errUnauthorized
481 + }
482 + delete(s.leases, record.ID)
483 + s.mu.Unlock()
484 +
485 + s.routes.DeleteLease(record.Hostnames)
486 + record.Broker.Drop()
487 + return nil
488 +}
489 +
490 +func (s *Server) lookupLeaseByID(leaseID, token string) (*leaseRecord, error) {
491 + s.mu.RLock()
492 + record, ok := s.leases[strings.TrimSpace(leaseID)]
493 + s.mu.RUnlock()
494 + if !ok {
495 + return nil, errLeaseNotFound
496 + }
497 + if time.Now().After(record.ExpiresAt) {
498 + return nil, errLeaseNotFound
499 + }
500 + if !tokenMatches(record.ReverseToken, token) {
501 + return nil, errUnauthorized
502 + }
503 + return record, nil
504 +}
505 +
506 +func (s *Server) findLeaseByHostnameLocked(host string) *leaseRecord {
507 + host = normalizeHostname(host)
508 + for _, lease := range s.leases {
509 + for _, candidate := range lease.Hostnames {
510 + if normalizeHostname(candidate) == host {
511 + return lease
512 + }
513 + }
514 + }
515 + return nil
516 +}
517 +
518 +func (s *Server) runAPIServer() error {
519 + err := s.apiServer.Serve(s.apiListener)
520 + if err == nil || errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
521 + return nil
522 + }
523 + return err
524 +}
525 +
526 +func (s *Server) runSNIListener() error {
527 + for {
528 + conn, err := s.sniListener.Accept()
529 + if err != nil {
530 + if s.ctx.Err() != nil || errors.Is(err, net.ErrClosed) {
531 + return nil
532 + }
533 + return err
534 + }
535 + go s.handleSNIConn(conn)
536 + }
537 +}
538 +
539 +func (s *Server) handleSNIConn(conn net.Conn) {
540 + clientHello, wrappedConn, err := l4.InspectClientHello(conn, s.cfg.ClientHelloTimeout)
541 + if err != nil {
542 + _ = wrappedConn.Close()
543 + return
544 + }
545 +
546 + serverName := normalizeHostname(clientHello.ServerName)
547 + if serverName == "" {
548 + _ = wrappedConn.Close()
549 + return
550 + }
551 +
552 + if serverName == s.cfg.RootHost && s.cfg.RootFallbackAddr != "" {
553 + s.bridgeToFallback(wrappedConn)
554 + return
555 + }
556 +
557 + leaseID, ok := s.routes.Lookup(serverName)
558 + if !ok {
559 + _ = wrappedConn.Close()
560 + return
561 + }
562 +
563 + s.mu.RLock()
564 + record := s.leases[leaseID]
565 + s.mu.RUnlock()
566 + if record == nil || time.Now().After(record.ExpiresAt) {
567 + _ = wrappedConn.Close()
568 + return
569 + }
570 +
571 + claimCtx, cancel := context.WithTimeout(s.ctx, s.cfg.ClaimTimeout)
572 + defer cancel()
573 +
574 + session, err := record.Broker.Claim(claimCtx)
575 + if err != nil {
576 + _ = wrappedConn.Close()
577 + return
578 + }
579 +
580 + bridgeConns(wrappedConn, session.Conn())
581 +}
582 +
583 +func (s *Server) bridgeToFallback(conn net.Conn) {
584 + upstream, err := net.DialTimeout("tcp", hostPortOrLoopback(s.cfg.RootFallbackAddr), 5*time.Second)
585 + if err != nil {
586 + _ = conn.Close()
587 + return
588 + }
589 + bridgeConns(conn, upstream)
590 +}
591 +
592 +func (s *Server) runLeaseJanitor() error {
593 + ticker := time.NewTicker(5 * time.Second)
594 + defer ticker.Stop()
595 +
596 + for {
597 + select {
598 + case <-s.ctx.Done():
599 + return nil
600 + case <-ticker.C:
601 + s.cleanupExpiredLeases()
602 + }
603 + }
604 +}
605 +
606 +func (s *Server) cleanupExpiredLeases() {
607 + now := time.Now()
608 +
609 + s.mu.Lock()
610 + expired := make([]*leaseRecord, 0)
611 + for leaseID, lease := range s.leases {
612 + if now.After(lease.ExpiresAt) {
613 + expired = append(expired, lease)
614 + delete(s.leases, leaseID)
615 + }
616 + }
617 + s.mu.Unlock()
618 +
619 + for _, lease := range expired {
620 + s.routes.DeleteLease(lease.Hostnames)
621 + lease.Broker.Drop()
622 + }
623 +}
624 +
625 +func (s *Server) watchContext() error {
626 + <-s.ctx.Done()
627 + shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
628 + defer cancel()
629 + return s.Shutdown(shutdownCtx)
630 +}
631 +
632 +func (s *Server) connectURL() string {
633 + base := strings.TrimRight(s.cfg.PortalURL, "/")
634 + if base == "" && s.apiListener != nil {
635 + return "https://" + hostPortOrLoopback(s.apiListener.Addr().String()) + "/sdk/connect"
636 + }
637 + return base + "/sdk/connect"
638 +}
639 +
640 +func (s *Server) wrapAPIHandler(base http.Handler) http.Handler {
641 + if s.cfg.APIHandlerWrapper == nil {
642 + return base
643 + }
644 + return s.cfg.APIHandlerWrapper(base)
645 +}
646 +
647 +var (
648 + errLeaseNotFound = errors.New("lease not found")
649 + errUnauthorized = errors.New("unauthorized")
650 + errHostnameConflict = errors.New("hostname already registered")
651 +)
652 +
653 +func decodeJSONBody(w http.ResponseWriter, r *http.Request, dst any) error {
654 + r.Body = http.MaxBytesReader(w, r.Body, defaultControlBodyLimit)
655 + defer r.Body.Close()
656 + return json.NewDecoder(r.Body).Decode(dst)
657 +}
658 +
659 +func normalizeHostnames(hosts []string) []string {
660 + seen := make(map[string]struct{}, len(hosts))
661 + out := make([]string, 0, len(hosts))
662 + for _, host := range hosts {
663 + host = normalizeHostname(host)
664 + if host == "" {
665 + continue
666 + }
667 + if _, ok := seen[host]; ok {
668 + continue
669 + }
670 + seen[host] = struct{}{}
671 + out = append(out, host)
672 + }
673 + return out
674 +}
675 +
676 +func tokenMatches(expected, actual string) bool {
677 + if len(expected) == 0 || len(actual) == 0 {
678 + return false
679 + }
680 + return subtle.ConstantTimeCompare([]byte(expected), []byte(actual)) == 1
681 +}
682 +
683 +func bridgeConns(left, right net.Conn) {
684 + defer left.Close()
685 + defer right.Close()
686 +
687 + var group errgroup.Group
688 + group.Go(func() error {
689 + _, err := io.Copy(right, left)
690 + closeWrite(right)
691 + return err
692 + })
693 + group.Go(func() error {
694 + _, err := io.Copy(left, right)
695 + closeWrite(left)
696 + return err
697 + })
698 + _ = group.Wait()
699 +}
700 +
701 +func closeWrite(conn net.Conn) {
702 + type closeWriter interface {
703 + CloseWrite() error
704 + }
705 + if cw, ok := conn.(closeWriter); ok {
706 + _ = cw.CloseWrite()
707 + }
708 +}
portal/sni/parser.go deleted
-322
@@ -1,322 +0,0 @@
1 -// Package sni provides TLS ClientHello parsing to extract SNI (Server Name Indication).
2 -// This is used for routing TLS connections in the TLS passthrough architecture.
3 -package sni
4 -
5 -import (
6 - "bytes"
7 - "encoding/binary"
8 - "errors"
9 - "fmt"
10 - "io"
11 -)
12 -
13 -const (
14 - // TLS wire constants for ClientHello/SNI parsing.
15 - tlsRecordContentTypeHandshake = byte(0x16)
16 - tlsHandshakeTypeClientHello = byte(0x01)
17 - tlsExtensionTypeServerName = uint16(0x0000)
18 - tlsServerNameTypeHostName = byte(0x00)
19 -)
20 -
21 -var (
22 - // ErrInvalidTLSRecord is returned when the TLS record is malformed.
23 - ErrInvalidTLSRecord = errors.New("invalid TLS record")
24 - // ErrNotClientHello is returned when the record is not a ClientHello.
25 - ErrNotClientHello = errors.New("not a ClientHello message")
26 - // ErrNoSNI is returned when the ClientHello doesn't contain SNI.
27 - ErrNoSNI = errors.New("no SNI found in ClientHello")
28 - // ErrInvalidSNI is returned when the SNI hostname is invalid.
29 - ErrInvalidSNI = errors.New("invalid SNI hostname")
30 -)
31 -
32 -// ExtractSNI extracts the SNI hostname from a TLS ClientHello message.
33 -// It reads from the provided reader and returns the SNI hostname.
34 -// The reader should be positioned at the start of the TLS record.
35 -func ExtractSNI(r io.Reader) (string, error) {
36 - // Read the TLS record header (5 bytes)
37 - // ContentType (1) + Version (2) + Length (2)
38 - header := make([]byte, 5)
39 - if _, err := io.ReadFull(r, header); err != nil {
40 - return "", fmt.Errorf("reading TLS header: %w", err)
41 - }
42 -
43 - // Check ContentType (handshake)
44 - if header[0] != tlsRecordContentTypeHandshake {
45 - return "", ErrNotClientHello
46 - }
47 -
48 - // Read the handshake message length
49 - recordLen := binary.BigEndian.Uint16(header[3:5])
50 - if recordLen < 4 {
51 - return "", ErrInvalidTLSRecord
52 - }
53 -
54 - // Read the full handshake message
55 - record := make([]byte, recordLen)
56 - if _, err := io.ReadFull(r, record); err != nil {
57 - return "", fmt.Errorf("reading TLS record: %w", err)
58 - }
59 -
60 - // Parse the handshake message
61 - return parseHandshake(record)
62 -}
63 -
64 -// parseHandshake parses a TLS handshake message and extracts SNI.
65 -func parseHandshake(data []byte) (string, error) {
66 - if len(data) < 4 {
67 - return "", ErrInvalidTLSRecord
68 - }
69 -
70 - // HandshakeType (1) + Length (3)
71 - handshakeType := data[0]
72 - declaredLen := int(data[1])<<16 | int(data[2])<<8 | int(data[3])
73 - if declaredLen < 0 || declaredLen > len(data)-4 {
74 - return "", ErrInvalidTLSRecord
75 - }
76 -
77 - // Check if it's a ClientHello.
78 - if handshakeType != tlsHandshakeTypeClientHello {
79 - return "", ErrNotClientHello
80 - }
81 -
82 - // Skip handshake header (4 bytes)
83 - return parseClientHello(data[4 : 4+declaredLen])
84 -}
85 -
86 -// parseClientHello parses a ClientHello message and extracts SNI.
87 -func parseClientHello(data []byte) (string, error) {
88 - // client_version(2) + random(32) + session_id_len(1)
89 - if len(data) < 35 {
90 - return "", ErrInvalidTLSRecord
91 - }
92 -
93 - offset := 0
94 -
95 - // Client Version (2 bytes)
96 - offset += 2
97 -
98 - // Random (32 bytes)
99 - offset += 32
100 -
101 - if offset >= len(data) {
102 - return "", ErrInvalidTLSRecord
103 - }
104 -
105 - // Session ID Length (1 byte) + Session ID
106 - sessionIDLen := int(data[offset])
107 - offset += 1 + sessionIDLen
108 -
109 - if offset > len(data) {
110 - return "", ErrInvalidTLSRecord
111 - }
112 -
113 - // Cipher Suites Length (2 bytes) + Cipher Suites
114 - if offset+2 > len(data) {
115 - return "", ErrInvalidTLSRecord
116 - }
117 - cipherSuitesLen := int(binary.BigEndian.Uint16(data[offset : offset+2]))
118 - offset += 2 + cipherSuitesLen
119 -
120 - if offset > len(data) {
121 - return "", ErrInvalidTLSRecord
122 - }
123 -
124 - // Compression Methods Length (1 byte) + Compression Methods
125 - if offset+1 > len(data) {
126 - return "", ErrInvalidTLSRecord
127 - }
128 - compressionMethodsLen := int(data[offset])
129 - offset += 1 + compressionMethodsLen
130 -
131 - if offset > len(data) {
132 - return "", ErrInvalidTLSRecord
133 - }
134 -
135 - // Extensions Length (2 bytes)
136 - if offset+2 > len(data) {
137 - return "", ErrNoSNI
138 - }
139 - extensionsLen := int(binary.BigEndian.Uint16(data[offset : offset+2]))
140 - offset += 2
141 -
142 - if extensionsLen == 0 || offset+extensionsLen > len(data) {
143 - return "", ErrNoSNI
144 - }
145 -
146 - // Parse extensions
147 - extensions := data[offset : offset+extensionsLen]
148 - return parseExtensions(extensions)
149 -}
150 -
151 -// parseExtensions parses TLS extensions and extracts SNI.
152 -func parseExtensions(data []byte) (string, error) {
153 - offset := 0
154 -
155 - for offset < len(data) {
156 - if offset+4 > len(data) {
157 - return "", ErrInvalidTLSRecord
158 - }
159 -
160 - // Extension Type (2 bytes)
161 - extType := binary.BigEndian.Uint16(data[offset : offset+2])
162 - offset += 2
163 -
164 - // Extension Length (2 bytes)
165 - extLen := int(binary.BigEndian.Uint16(data[offset : offset+2]))
166 - offset += 2
167 -
168 - if offset+extLen > len(data) {
169 - return "", ErrInvalidTLSRecord
170 - }
171 -
172 - // Extension Type server_name (SNI)
173 - if extType == tlsExtensionTypeServerName {
174 - return parseSNIExtension(data[offset : offset+extLen])
175 - }
176 -
177 - offset += extLen
178 - }
179 -
180 - return "", ErrNoSNI
181 -}
182 -
183 -// parseSNIExtension parses the SNI extension and returns the hostname.
184 -func parseSNIExtension(data []byte) (string, error) {
185 - if len(data) < 2 {
186 - return "", ErrNoSNI
187 - }
188 -
189 - // SNI List Length (2 bytes)
190 - listLen := int(binary.BigEndian.Uint16(data[0:2]))
191 - if listLen == 0 || 2+listLen > len(data) {
192 - return "", ErrNoSNI
193 - }
194 -
195 - offset := 2
196 - end := 2 + listLen
197 -
198 - for offset < end {
199 - if offset+3 > end {
200 - return "", ErrInvalidTLSRecord
201 - }
202 -
203 - // Name Type (1 byte)
204 - nameType := data[offset]
205 - offset++
206 -
207 - // Name Length (2 bytes)
208 - nameLen := int(binary.BigEndian.Uint16(data[offset : offset+2]))
209 - offset += 2
210 -
211 - if offset+nameLen > end {
212 - return "", ErrInvalidTLSRecord
213 - }
214 -
215 - // Name Type host_name
216 - if nameType == tlsServerNameTypeHostName {
217 - if nameLen == 0 {
218 - return "", ErrNoSNI
219 - }
220 - hostname := string(data[offset : offset+nameLen])
221 - if !isValidSNIHostname(hostname) {
222 - return "", ErrInvalidSNI
223 - }
224 - return hostname, nil
225 - }
226 -
227 - offset += nameLen
228 - }
229 -
230 - return "", ErrNoSNI
231 -}
232 -
233 -// isValidSNIHostname validates that a hostname is a valid DNS name per RFC 1035 and RFC 1123.
234 -// - Total length must not exceed 253 characters
235 -// - Labels must be 1-63 characters
236 -// - Labels can contain a-z, A-Z, 0-9, and hyphen
237 -// - Labels cannot start or end with hyphen
238 -// - No null bytes or other control characters.
239 -func isValidSNIHostname(hostname string) bool {
240 - if len(hostname) == 0 || len(hostname) > 253 {
241 - return false
242 - }
243 -
244 - // Check for null bytes and other control characters
245 - for i := range len(hostname) {
246 - if hostname[i] < 0x20 || hostname[i] > 0x7E {
247 - return false
248 - }
249 - }
250 -
251 - start := 0
252 - for i := 0; i <= len(hostname); i++ {
253 - if i == len(hostname) || hostname[i] == '.' {
254 - label := hostname[start:i]
255 - if len(label) == 0 || len(label) > 63 {
256 - return false
257 - }
258 - // Check label characters
259 - for j, c := range []byte(label) {
260 - // Allow a-z, A-Z, 0-9, and hyphen
261 - if (c < 'a' || c > 'z') && (c < 'A' || c > 'Z') && (c < '0' || c > '9') && c != '-' {
262 - return false
263 - }
264 - // Label cannot start or end with hyphen
265 - if c == '-' && (j == 0 || j == len(label)-1) {
266 - return false
267 - }
268 - }
269 - start = i + 1
270 - }
271 - }
272 -
273 - return true
274 -}
275 -
276 -// PeekSNI peeks at the SNI from a connection without consuming the data.
277 -// It returns the SNI and a new reader that includes the peeked data.
278 -// This is useful for routing connections before fully reading them.
279 -func PeekSNI(r io.Reader, bufSize int) (string, io.Reader, error) {
280 - if bufSize < 5 {
281 - return "", nil, fmt.Errorf("peek buffer too small: %d", bufSize)
282 - }
283 -
284 - // Read TLS record header first so we only read the exact record size.
285 - header := make([]byte, 5)
286 - if _, err := io.ReadFull(r, header); err != nil {
287 - return "", nil, fmt.Errorf("peeking TLS header: %w", err)
288 - }
289 - if header[0] != tlsRecordContentTypeHandshake {
290 - reader := io.MultiReader(bytes.NewReader(header), r)
291 - return "", reader, ErrNotClientHello
292 - }
293 -
294 - recordLen := int(binary.BigEndian.Uint16(header[3:5]))
295 - if recordLen <= 0 {
296 - reader := io.MultiReader(bytes.NewReader(header), r)
297 - return "", reader, ErrInvalidTLSRecord
298 - }
299 -
300 - totalLen := 5 + recordLen
301 - if totalLen > bufSize {
302 - reader := io.MultiReader(bytes.NewReader(header), r)
303 - return "", reader, fmt.Errorf("TLS record too large for peek buffer: need %d bytes, have %d", totalLen, bufSize)
304 - }
305 -
306 - buf := make([]byte, totalLen)
307 - copy(buf, header)
308 - if _, err := io.ReadFull(r, buf[5:]); err != nil {
309 - reader := io.MultiReader(bytes.NewReader(buf[:5]), r)
310 - return "", reader, fmt.Errorf("peeking TLS record: %w", err)
311 - }
312 -
313 - // Create a reader that includes the peeked data.
314 - reader := io.MultiReader(bytes.NewReader(buf), r)
315 -
316 - sni, err := ExtractSNI(bytes.NewReader(buf))
317 - if err != nil {
318 - return "", reader, err
319 - }
320 -
321 - return sni, reader, nil
322 -}
portal/sni/parser_test.go deleted
-165
@@ -1,165 +0,0 @@
1 -package sni
2 -
3 -import (
4 - "bytes"
5 - "encoding/binary"
6 - "errors"
7 - "io"
8 - "strings"
9 - "testing"
10 -)
11 -
12 -func TestExtractSNI(t *testing.T) {
13 - clientHello := buildClientHello("example.com", true)
14 -
15 - sni, err := ExtractSNI(bytes.NewReader(clientHello))
16 - if err != nil {
17 - t.Fatalf("ExtractSNI failed: %v", err)
18 - }
19 - if sni != "example.com" {
20 - t.Errorf("Expected SNI 'example.com', got '%s'", sni)
21 - }
22 -}
23 -
24 -func TestExtractSNI_NoSNI(t *testing.T) {
25 - clientHello := buildClientHello("", false)
26 -
27 - _, err := ExtractSNI(bytes.NewReader(clientHello))
28 - if !errors.Is(err, ErrNoSNI) {
29 - t.Errorf("Expected ErrNoSNI, got: %v", err)
30 - }
31 -}
32 -
33 -func TestExtractSNI_NotClientHello(t *testing.T) {
34 - serverHello := buildTLSRecord(0x02, nil)
35 -
36 - _, err := ExtractSNI(bytes.NewReader(serverHello))
37 - if !errors.Is(err, ErrNotClientHello) {
38 - t.Errorf("Expected ErrNotClientHello, got: %v", err)
39 - }
40 -}
41 -
42 -func TestPeekSNI(t *testing.T) {
43 - clientHello := buildClientHello("example.com", true)
44 -
45 - sni, reader, err := PeekSNI(bytes.NewReader(clientHello), 4096)
46 - if err != nil {
47 - t.Fatalf("PeekSNI failed: %v", err)
48 - }
49 - if sni != "example.com" {
50 - t.Errorf("Expected SNI 'example.com', got '%s'", sni)
51 - }
52 -
53 - // Verify we can still read the full ClientHello from the returned reader.
54 - buf, err := io.ReadAll(reader)
55 - if err != nil {
56 - t.Fatalf("Reading from returned reader failed: %v", err)
57 - }
58 - if !bytes.Equal(buf, clientHello) {
59 - t.Error("Returned reader doesn't contain the full ClientHello")
60 - }
61 -}
62 -
63 -func TestExtractSNI_TruncatedClientHello(t *testing.T) {
64 - record := buildTLSRecord(0x01, make([]byte, 34)) // 1 byte short for session_id_len field access
65 - _, err := ExtractSNI(bytes.NewReader(record))
66 - if !errors.Is(err, ErrInvalidTLSRecord) {
67 - t.Fatalf("expected ErrInvalidTLSRecord, got: %v", err)
68 - }
69 -}
70 -
71 -func TestExtractSNI_InvalidHandshakeLength(t *testing.T) {
72 - record := buildTLSRecord(0x01, []byte{0x03, 0x03, 0x00, 0x00, 0x00})
73 - // Corrupt handshake declared length to exceed available bytes
74 - record[6] = 0x00
75 - record[7] = 0x01
76 - record[8] = 0x00
77 -
78 - _, err := ExtractSNI(bytes.NewReader(record))
79 - if !errors.Is(err, ErrInvalidTLSRecord) {
80 - t.Fatalf("expected ErrInvalidTLSRecord, got: %v", err)
81 - }
82 -}
83 -
84 -func TestPeekSNI_NotClientHelloPreservesData(t *testing.T) {
85 - payload := []byte("plaintext")
86 - buf := append([]byte{0x17, 0x03, 0x03, 0x00, byte(len(payload))}, payload...)
87 -
88 - _, reader, err := PeekSNI(bytes.NewReader(buf), 4096)
89 - if !errors.Is(err, ErrNotClientHello) {
90 - t.Fatalf("expected ErrNotClientHello, got: %v", err)
91 - }
92 -
93 - got, readErr := io.ReadAll(reader)
94 - if readErr != nil {
95 - t.Fatalf("failed to read returned reader: %v", readErr)
96 - }
97 - if !bytes.Equal(got, buf) {
98 - t.Fatalf("returned reader did not preserve original bytes")
99 - }
100 -}
101 -
102 -func TestPeekSNI_RecordTooLarge(t *testing.T) {
103 - clientHello := buildClientHello(strings.Repeat("a", 10)+".example.com", true)
104 -
105 - _, reader, err := PeekSNI(bytes.NewReader(clientHello), 64)
106 - if err == nil || !strings.Contains(err.Error(), "too large") {
107 - t.Fatalf("expected record too large error, got: %v", err)
108 - }
109 -
110 - got, readErr := io.ReadAll(reader)
111 - if readErr != nil {
112 - t.Fatalf("failed to read returned reader: %v", readErr)
113 - }
114 - if !bytes.Equal(got, clientHello) {
115 - t.Fatalf("returned reader did not preserve original bytes")
116 - }
117 -}
118 -
119 -func buildClientHello(sni string, includeSNI bool) []byte {
120 - body := make([]byte, 0, 128)
121 - body = append(body, 0x03, 0x03) // TLS 1.2
122 - body = append(body, make([]byte, 32)...)
123 - body = append(body, 0x00) // Session ID length
124 - body = append(body, 0x00, 0x02) // Cipher Suites length
125 - body = append(body, 0x00, 0x2f) // TLS_RSA_WITH_AES_128_CBC_SHA
126 - body = append(body, 0x01, 0x00) // Compression methods
127 - extensions := make([]byte, 0, 64) // Extensions
128 - if includeSNI {
129 - host := []byte(sni)
130 - sniData := make([]byte, 2+1+2+len(host)) // list_len + name_type + name_len + host
131 - binary.BigEndian.PutUint16(sniData[0:2], uint16(1+2+len(host)))
132 - sniData[2] = 0x00 // host_name
133 - binary.BigEndian.PutUint16(sniData[3:5], uint16(len(host)))
134 - copy(sniData[5:], host)
135 -
136 - ext := make([]byte, 4+len(sniData)) // ext_type + ext_len + ext_data
137 - binary.BigEndian.PutUint16(ext[0:2], 0x0000)
138 - binary.BigEndian.PutUint16(ext[2:4], uint16(len(sniData)))
139 - copy(ext[4:], sniData)
140 - extensions = append(extensions, ext...)
141 - }
142 - body = append(body, byte(len(extensions)>>8), byte(len(extensions)))
143 - body = append(body, extensions...)
144 -
145 - return buildTLSRecord(0x01, body)
146 -}
147 -
148 -func buildTLSRecord(handshakeType byte, handshakeBody []byte) []byte {
149 - handshake := make([]byte, 4+len(handshakeBody))
150 - handshake[0] = handshakeType
151 - handshakeLen := len(handshakeBody)
152 - handshake[1] = byte(handshakeLen >> 16)
153 - handshake[2] = byte(handshakeLen >> 8)
154 - handshake[3] = byte(handshakeLen)
155 - copy(handshake[4:], handshakeBody)
156 -
157 - record := make([]byte, 5+len(handshake))
158 - record[0] = 0x16 // Handshake
159 - record[1] = 0x03 // TLS 1.x
160 - record[2] = 0x01 // TLS 1.0 record version
161 - recordLen := len(handshake)
162 - binary.BigEndian.PutUint16(record[3:5], uint16(recordLen))
163 - copy(record[5:], handshake)
164 - return record
165 -}
portal/sni/router.go deleted
-432
@@ -1,432 +0,0 @@
1 -// Package sni provides TLS SNI-based TCP routing for the Portal relay.
2 -package sni
3 -
4 -import (
5 - "context"
6 - "errors"
7 - "fmt"
8 - "io"
9 - "net"
10 - "strings"
11 - "sync"
12 - "time"
13 -
14 - "github.com/rs/zerolog/log"
15 -)
16 -
17 -var (
18 - // ErrNoRoute is returned when no route is found for the SNI.
19 - ErrNoRoute = errors.New("no route found for SNI")
20 - // ErrRouterClosed is returned when the router is closed.
21 - ErrRouterClosed = errors.New("router is closed")
22 -)
23 -
24 -const (
25 - // maxTLSRecordSize is TLS plaintext limit (16KB) plus allowance for overhead.
26 - // Using this avoids dropping valid large ClientHello messages.
27 - maxTLSRecordSize = 16*1024 + 2048
28 -)
29 -
30 -// Route represents a registered route.
31 -type Route struct {
32 - SNI string
33 - LeaseID string
34 - LeaseName string
35 -}
36 -
37 -// Router handles SNI-based TCP routing.
38 -type Router struct {
39 - listener net.Listener
40 - routes map[string]*Route
41 - leases map[string]*Route
42 - conns map[net.Conn]struct{}
43 - onConnection func(conn net.Conn, route *Route)
44 - onNoRoute func(conn net.Conn, sni string) bool
45 - stopCh chan struct{}
46 - addr string
47 - wg sync.WaitGroup
48 - mu sync.RWMutex
49 - stopOnce sync.Once
50 -}
51 -
52 -// NewRouter creates a new SNI router.
53 -func NewRouter(addr string) *Router {
54 - return &Router{
55 - addr: addr,
56 - routes: make(map[string]*Route),
57 - leases: make(map[string]*Route),
58 - conns: make(map[net.Conn]struct{}),
59 - stopCh: make(chan struct{}),
60 - }
61 -}
62 -
63 -// GetAddr returns the listen address.
64 -func (r *Router) GetAddr() string {
65 - return r.addr
66 -}
67 -
68 -// SetConnectionCallback sets the callback for new connections.
69 -func (r *Router) SetConnectionCallback(cb func(conn net.Conn, route *Route)) {
70 - r.mu.Lock()
71 - defer r.mu.Unlock()
72 - r.onConnection = cb
73 -}
74 -
75 -// SetNoRouteHandler sets the callback for unmatched SNI connections.
76 -// Return true when the callback handled the connection lifecycle.
77 -func (r *Router) SetNoRouteHandler(cb func(conn net.Conn, sni string) bool) {
78 - r.mu.Lock()
79 - defer r.mu.Unlock()
80 - r.onNoRoute = cb
81 -}
82 -
83 -// RegisterRoute registers a new route for an SNI.
84 -func (r *Router) RegisterRoute(sni, leaseID, leaseName string) error {
85 - r.mu.Lock()
86 - defer r.mu.Unlock()
87 -
88 - select {
89 - case <-r.stopCh:
90 - return ErrRouterClosed
91 - default:
92 - }
93 -
94 - sni = strings.ToLower(strings.TrimSpace(sni))
95 - if sni == "" {
96 - return errors.New("sni is required")
97 - }
98 -
99 - // Remove previous SNI entry when a lease is re-registered with a new name.
100 - if oldRoute, ok := r.leases[leaseID]; ok && oldRoute.SNI != sni {
101 - delete(r.routes, oldRoute.SNI)
102 - }
103 -
104 - // Warn if SNI is already registered to a different lease (should not happen if lease manager is consistent)
105 - if existingRoute, ok := r.routes[sni]; ok && existingRoute.LeaseID != leaseID {
106 - log.Warn().
107 - Str("sni", sni).
108 - Str("existing_lease_id", existingRoute.LeaseID).
109 - Str("new_lease_id", leaseID).
110 - Msg("[SNI] SNI already registered to different lease; overwriting")
111 - delete(r.leases, existingRoute.LeaseID)
112 - }
113 -
114 - route := &Route{
115 - SNI: sni,
116 - LeaseID: leaseID,
117 - LeaseName: leaseName,
118 - }
119 -
120 - r.routes[sni] = route
121 - r.leases[leaseID] = route
122 -
123 - log.Info().
124 - Str("sni", sni).
125 - Str("lease_id", leaseID).
126 - Msg("[SNI] Route registered")
127 -
128 - return nil
129 -}
130 -
131 -// UnregisterRoute removes a route for an SNI.
132 -func (r *Router) UnregisterRoute(sni string) {
133 - r.mu.Lock()
134 - defer r.mu.Unlock()
135 -
136 - sni = strings.ToLower(strings.TrimSpace(sni))
137 -
138 - if route, ok := r.routes[sni]; ok {
139 - delete(r.routes, sni)
140 - delete(r.leases, route.LeaseID)
141 - log.Info().
142 - Str("sni", sni).
143 - Str("lease_id", route.LeaseID).
144 - Msg("[SNI] Route unregistered")
145 - }
146 -}
147 -
148 -// UnregisterRouteByLeaseID removes a route by lease ID.
149 -func (r *Router) UnregisterRouteByLeaseID(leaseID string) {
150 - r.mu.Lock()
151 - defer r.mu.Unlock()
152 -
153 - if route, ok := r.leases[leaseID]; ok {
154 - delete(r.routes, route.SNI)
155 - delete(r.leases, leaseID)
156 - log.Info().
157 - Str("sni", route.SNI).
158 - Str("lease_id", leaseID).
159 - Msg("[SNI] Route unregistered")
160 - }
161 -}
162 -
163 -// GetRoute returns the route for an SNI.
164 -func (r *Router) GetRoute(sni string) (*Route, bool) {
165 - r.mu.RLock()
166 - defer r.mu.RUnlock()
167 -
168 - sni = strings.ToLower(strings.TrimSpace(sni))
169 -
170 - // Try exact match first
171 - if route, ok := r.routes[sni]; ok {
172 - return route, true
173 - }
174 -
175 - // Try wildcard match (e.g., *.example.com matches foo.example.com)
176 - // TLS wildcards only match a single DNS label, so for "foo.bar.example.com"
177 - // we only check "*.bar.example.com", not "*.example.com"
178 - parts := strings.Split(sni, ".")
179 - if len(parts) >= 2 {
180 - // Only check the immediate parent wildcard
181 - wildcard := "*." + strings.Join(parts[1:], ".")
182 - if route, ok := r.routes[wildcard]; ok {
183 - return route, true
184 - }
185 - }
186 -
187 - return nil, false
188 -}
189 -
190 -// GetRouteByLeaseID returns the route for a lease ID.
191 -func (r *Router) GetRouteByLeaseID(leaseID string) (*Route, bool) {
192 - r.mu.RLock()
193 - defer r.mu.RUnlock()
194 -
195 - route, ok := r.leases[leaseID]
196 - return route, ok
197 -}
198 -
199 -// GetAllRoutes returns all registered routes.
200 -func (r *Router) GetAllRoutes() []*Route {
201 - r.mu.RLock()
202 - defer r.mu.RUnlock()
203 -
204 - routes := make([]*Route, 0, len(r.routes))
205 - for _, route := range r.routes {
206 - routes = append(routes, route)
207 - }
208 - return routes
209 -}
210 -
211 -// Start starts the SNI router on the configured address.
212 -func (r *Router) Start() error {
213 - listenConfig := &net.ListenConfig{}
214 - listener, err := listenConfig.Listen(context.Background(), "tcp", r.addr)
215 - if err != nil {
216 - return fmt.Errorf("failed to listen on %s: %w", r.addr, err)
217 - }
218 -
219 - r.mu.Lock()
220 - r.listener = listener
221 - r.mu.Unlock()
222 -
223 - log.Info().
224 - Str("addr", r.addr).
225 - Msg("[SNI] Router started")
226 -
227 - r.wg.Add(1)
228 - go r.acceptLoop(listener)
229 -
230 - return nil
231 -}
232 -
233 -// Stop stops the SNI router.
234 -func (r *Router) Stop() error {
235 - r.stopOnce.Do(func() {
236 - close(r.stopCh)
237 -
238 - r.mu.Lock()
239 - if r.listener != nil {
240 - closeWithDebugLog(r.listener, "[SNI] failed to close listener")
241 - }
242 - for conn := range r.conns {
243 - closeWithDebugLog(conn, "[SNI] failed to close tracked connection")
244 - }
245 - r.mu.Unlock()
246 - })
247 -
248 - r.wg.Wait()
249 - log.Info().Msg("[SNI] Router stopped")
250 - return nil
251 -}
252 -
253 -// Addr returns the router's listen address.
254 -func (r *Router) Addr() net.Addr {
255 - r.mu.RLock()
256 - defer r.mu.RUnlock()
257 -
258 - if r.listener != nil {
259 - return r.listener.Addr()
260 - }
261 - return nil
262 -}
263 -
264 -// acceptLoop accepts incoming connections.
265 -func (r *Router) acceptLoop(listener net.Listener) {
266 - defer r.wg.Done()
267 -
268 - for {
269 - conn, err := listener.Accept()
270 - if err != nil {
271 - select {
272 - case <-r.stopCh:
273 - return
274 - default:
275 - log.Error().Err(err).Msg("[SNI] Accept error")
276 - continue
277 - }
278 - }
279 -
280 - r.wg.Add(1)
281 - go r.handleConnection(conn)
282 - }
283 -}
284 -
285 -// handleConnection handles a single connection.
286 -func (r *Router) handleConnection(clientConn net.Conn) {
287 - defer r.wg.Done()
288 -
289 - // Set a deadline for reading the ClientHello
290 - if err := clientConn.SetReadDeadline(time.Now().Add(5 * time.Second)); err != nil {
291 - log.Debug().Err(err).Msg("[SNI] failed to set read deadline")
292 - }
293 -
294 - // Peek at the SNI from the ClientHello
295 - sni, peekedReader, err := PeekSNI(clientConn, maxTLSRecordSize)
296 - if err != nil {
297 - log.Error().
298 - Err(err).
299 - Str("remote", clientConn.RemoteAddr().String()).
300 - Msg("[SNI] Failed to extract SNI")
301 - closeWithDebugLog(clientConn, "[SNI] failed to close client connection")
302 - return
303 - }
304 -
305 - // Clear the deadline
306 - if err := clientConn.SetReadDeadline(time.Time{}); err != nil {
307 - log.Debug().Err(err).Msg("[SNI] failed to clear read deadline")
308 - }
309 -
310 - // Wrap the connection so callbacks can still read the peeked bytes.
311 - wrappedConn := &peekedConn{
312 - Conn: clientConn,
313 - reader: peekedReader,
314 - }
315 -
316 - // Find the route
317 - route, ok := r.GetRoute(sni)
318 - if !ok {
319 - onNoRoute := r.getNoRouteHandler()
320 - if onNoRoute != nil && onNoRoute(wrappedConn, sni) {
321 - return
322 - }
323 -
324 - log.Warn().
325 - Str("sni", sni).
326 - Str("remote", clientConn.RemoteAddr().String()).
327 - Msg("[SNI] No route found")
328 - closeWithDebugLog(clientConn, "[SNI] failed to close unrouted client connection")
329 - return
330 - }
331 -
332 - log.Debug().
333 - Str("sni", sni).
334 - Str("lease_id", route.LeaseID).
335 - Str("remote", clientConn.RemoteAddr().String()).
336 - Msg("[SNI] Route found")
337 -
338 - // Call the connection callback if set
339 - onConnection := r.getConnectionHandler()
340 -
341 - if onConnection != nil {
342 - r.mu.Lock()
343 - r.conns[wrappedConn] = struct{}{}
344 - r.mu.Unlock()
345 -
346 - onConnection(wrappedConn, route)
347 -
348 - r.mu.Lock()
349 - delete(r.conns, wrappedConn)
350 - r.mu.Unlock()
351 - return
352 - }
353 -
354 - // No callback configured - close connection
355 - log.Warn().
356 - Str("sni", sni).
357 - Msg("[SNI] No connection callback configured, closing connection")
358 - closeWithDebugLog(clientConn, "[SNI] failed to close client connection")
359 -}
360 -
361 -func closeWithDebugLog(closer io.Closer, msg string) {
362 - if closer == nil {
363 - return
364 - }
365 - if err := closer.Close(); err != nil {
366 - log.Debug().Err(err).Msg(msg)
367 - }
368 -}
369 -
370 -func (r *Router) getNoRouteHandler() func(conn net.Conn, sni string) bool {
371 - r.mu.RLock()
372 - defer r.mu.RUnlock()
373 - return r.onNoRoute
374 -}
375 -
376 -func (r *Router) getConnectionHandler() func(conn net.Conn, route *Route) {
377 - r.mu.RLock()
378 - defer r.mu.RUnlock()
379 - return r.onConnection
380 -}
381 -
382 -// BridgeConnections bridges two connections.
383 -func BridgeConnections(conn1, conn2 net.Conn) {
384 - defer conn1.Close()
385 - defer conn2.Close()
386 -
387 - errCh := make(chan error, 2)
388 -
389 - // Conn1 -> Conn2
390 - go func() {
391 - _, err := io.Copy(conn2, conn1)
392 - errCh <- err
393 - closeWithDebugLog(conn2, "[SNI] failed to close bridged connection 2")
394 - }()
395 -
396 - // Conn2 -> Conn1
397 - go func() {
398 - _, err := io.Copy(conn1, conn2)
399 - errCh <- err
400 - closeWithDebugLog(conn1, "[SNI] failed to close bridged connection 1")
401 - }()
402 -
403 - // Wait for either direction to close
404 - <-errCh
405 -}
406 -
407 -// ExtractSNIFromConnection extracts SNI from a connection without consuming data.
408 -// It returns the SNI and a wrapped connection that includes the peeked data.
409 -func ExtractSNIFromConnection(conn net.Conn, bufSize int) (string, net.Conn, error) {
410 - sni, reader, err := PeekSNI(conn, bufSize)
411 - if err != nil {
412 - return "", nil, err
413 - }
414 -
415 - // Wrap the connection to include the peeked data.
416 - wrappedConn := &peekedConn{
417 - Conn: conn,
418 - reader: reader,
419 - }
420 -
421 - return sni, wrappedConn, nil
422 -}
423 -
424 -// peekedConn wraps a net.Conn to include peeked data.
425 -type peekedConn struct {
426 - net.Conn
427 - reader io.Reader
428 -}
429 -
430 -func (c *peekedConn) Read(p []byte) (int, error) {
431 - return c.reader.Read(p)
432 -}
portal/sni/router_test.go deleted
-347
@@ -1,347 +0,0 @@
1 -package sni
2 -
3 -import (
4 - "errors"
5 - "net"
6 - "testing"
7 - "time"
8 -)
9 -
10 -func TestRouter_RegisterRoute(t *testing.T) {
11 - router := NewRouter("")
12 -
13 - // Test basic registration
14 - err := router.RegisterRoute("example.com", "lease-1", "test")
15 - if err != nil {
16 - t.Fatalf("failed to register route: %v", err)
17 - }
18 -
19 - // Test duplicate registration by same lease (should succeed - update)
20 - err = router.RegisterRoute("example.com", "lease-1", "test")
21 - if err != nil {
22 - t.Fatalf("failed to update route: %v", err)
23 - }
24 -
25 - route, ok := router.GetRoute("example.com")
26 - if !ok {
27 - t.Fatal("route not found")
28 - }
29 - if route.LeaseID != "lease-1" {
30 - t.Errorf("expected lease-1, got %s", route.LeaseID)
31 - }
32 -}
33 -
34 -func TestRouter_UnregisterRoute(t *testing.T) {
35 - router := NewRouter("")
36 -
37 - err := router.RegisterRoute("example.com", "lease-1", "test")
38 - if err != nil {
39 - t.Fatalf("failed to register route: %v", err)
40 - }
41 -
42 - router.UnregisterRoute("example.com")
43 -
44 - _, ok := router.GetRoute("example.com")
45 - if ok {
46 - t.Error("expected route to be unregistered")
47 - }
48 -}
49 -
50 -func TestRouter_UnregisterRouteByLeaseID(t *testing.T) {
51 - router := NewRouter("")
52 -
53 - err := router.RegisterRoute("example.com", "lease-1", "test")
54 - if err != nil {
55 - t.Fatalf("failed to register route: %v", err)
56 - }
57 -
58 - router.UnregisterRouteByLeaseID("lease-1")
59 -
60 - _, ok := router.GetRoute("example.com")
61 - if ok {
62 - t.Error("expected route to be unregistered")
63 - }
64 -}
65 -
66 -func TestRouter_GetRoute_Wildcard(t *testing.T) {
67 - router := NewRouter("")
68 -
69 - // Register wildcard route
70 - err := router.RegisterRoute("*.example.com", "lease-1", "test")
71 - if err != nil {
72 - t.Fatalf("failed to register wildcard route: %v", err)
73 - }
74 -
75 - tests := []struct {
76 - sni string
77 - wantName string
78 - wantOK bool
79 - }{
80 - {"foo.example.com", "*.example.com", true}, // should match
81 - {"bar.example.com", "*.example.com", true}, // should match
82 - {"example.com", "", false}, // should NOT match (no subdomain)
83 - {"foo.bar.example.com", "", false}, // should NOT match (TLS wildcard only matches one level)
84 - {"other.com", "", false}, // should NOT match
85 - }
86 -
87 - for _, tt := range tests {
88 - t.Run(tt.sni, func(t *testing.T) {
89 - route, ok := router.GetRoute(tt.sni)
90 - if ok != tt.wantOK {
91 - t.Errorf("GetRoute(%q) = %v, want %v", tt.sni, ok, tt.wantOK)
92 - return
93 - }
94 - if ok && route.SNI != tt.wantName {
95 - t.Errorf("GetRoute(%q) matched %q, want %q", tt.sni, route.SNI, tt.wantName)
96 - }
97 - })
98 - }
99 -}
100 -
101 -func TestRouter_GetRoute_ExactBeforeWildcard(t *testing.T) {
102 - router := NewRouter("")
103 -
104 - // Register both exact and wildcard routes
105 - err := router.RegisterRoute("*.example.com", "lease-1", "wildcard")
106 - if err != nil {
107 - t.Fatalf("failed to register wildcard route: %v", err)
108 - }
109 -
110 - err = router.RegisterRoute("specific.example.com", "lease-2", "specific")
111 - if err != nil {
112 - t.Fatalf("failed to register specific route: %v", err)
113 - }
114 -
115 - // Exact match should take precedence
116 - route, ok := router.GetRoute("specific.example.com")
117 - if !ok {
118 - t.Fatal("route not found")
119 - }
120 - if route.LeaseID != "lease-2" {
121 - t.Errorf("expected lease-2 (exact match), got %s", route.LeaseID)
122 - }
123 -
124 - // Other subdomains should match wildcard
125 - route, ok = router.GetRoute("other.example.com")
126 - if !ok {
127 - t.Fatal("route not found")
128 - }
129 - if route.LeaseID != "lease-1" {
130 - t.Errorf("expected lease-1 (wildcard match), got %s", route.LeaseID)
131 - }
132 -}
133 -
134 -func TestRouter_GetRouteByLeaseID(t *testing.T) {
135 - router := NewRouter("")
136 -
137 - err := router.RegisterRoute("example.com", "lease-1", "test")
138 - if err != nil {
139 - t.Fatalf("failed to register route: %v", err)
140 - }
141 -
142 - route, ok := router.GetRouteByLeaseID("lease-1")
143 - if !ok {
144 - t.Fatal("route not found by lease ID")
145 - }
146 - if route.SNI != "example.com" {
147 - t.Errorf("expected SNI example.com, got %s", route.SNI)
148 - }
149 -
150 - _, ok = router.GetRouteByLeaseID("nonexistent")
151 - if ok {
152 - t.Error("expected route not found for nonexistent lease ID")
153 - }
154 -}
155 -
156 -func TestRouter_GetAllRoutes(t *testing.T) {
157 - router := NewRouter("")
158 -
159 - _ = router.RegisterRoute("example.com", "lease-1", "test1")
160 - _ = router.RegisterRoute("other.com", "lease-2", "test2")
161 -
162 - routes := router.GetAllRoutes()
163 - if len(routes) != 2 {
164 - t.Errorf("expected 2 routes, got %d", len(routes))
165 - }
166 -}
167 -
168 -func TestRouter_CaseInsensitive(t *testing.T) {
169 - router := NewRouter("")
170 -
171 - err := router.RegisterRoute("Example.COM", "lease-1", "test")
172 - if err != nil {
173 - t.Fatalf("failed to register route: %v", err)
174 - }
175 -
176 - // Should find with different case
177 - route, ok := router.GetRoute("EXAMPLE.com")
178 - if !ok {
179 - t.Fatal("route not found with different case")
180 - }
181 - if route.SNI != "example.com" {
182 - t.Errorf("expected normalized SNI example.com, got %s", route.SNI)
183 - }
184 -}
185 -
186 -func TestRouter_LeaseRename(t *testing.T) {
187 - router := NewRouter("")
188 -
189 - // Register with name1
190 - err := router.RegisterRoute("name1.example.com", "lease-1", "name1")
191 - if err != nil {
192 - t.Fatalf("failed to register route: %v", err)
193 - }
194 -
195 - // Same lease re-registers with name2
196 - err = router.RegisterRoute("name2.example.com", "lease-1", "name2")
197 - if err != nil {
198 - t.Fatalf("failed to re-register route: %v", err)
199 - }
200 -
201 - // Old name should be gone
202 - _, ok := router.GetRoute("name1.example.com")
203 - if ok {
204 - t.Error("old route should be removed")
205 - }
206 -
207 - // New name should exist
208 - route, ok := router.GetRoute("name2.example.com")
209 - if !ok {
210 - t.Fatal("new route not found")
211 - }
212 - if route.LeaseID != "lease-1" {
213 - t.Errorf("expected lease-1, got %s", route.LeaseID)
214 - }
215 -}
216 -
217 -func TestRouter_Stop(t *testing.T) {
218 - router := NewRouter("")
219 -
220 - err := router.RegisterRoute("example.com", "lease-1", "test")
221 - if err != nil {
222 - t.Fatalf("failed to register route: %v", err)
223 - }
224 -
225 - // Stop should not panic
226 - err = router.Stop()
227 - if err != nil {
228 - t.Errorf("stop failed: %v", err)
229 - }
230 -
231 - // Registration after stop should fail
232 - err = router.RegisterRoute("other.com", "lease-2", "test2")
233 - if !errors.Is(err, ErrRouterClosed) {
234 - t.Errorf("expected ErrRouterClosed, got %v", err)
235 - }
236 -}
237 -
238 -func TestRouter_HandleConnectionNoRouteHandlerHandled(t *testing.T) {
239 - router := NewRouter("")
240 - noRouteCalls := make(chan string, 1)
241 - router.SetNoRouteHandler(func(_ net.Conn, sni string) bool {
242 - select {
243 - case noRouteCalls <- sni:
244 - default:
245 - }
246 - return true
247 - })
248 -
249 - client, server := net.Pipe()
250 - defer func() {
251 - _ = client.Close()
252 - _ = server.Close()
253 - }()
254 -
255 - done := make(chan struct{})
256 - router.wg.Add(1)
257 - go func() {
258 - router.handleConnection(server)
259 - close(done)
260 - }()
261 -
262 - if _, err := client.Write(buildClientHello("tenant.example.com", true)); err != nil {
263 - t.Fatalf("write client hello: %v", err)
264 - }
265 -
266 - select {
267 - case gotSNI := <-noRouteCalls:
268 - if gotSNI != "tenant.example.com" {
269 - t.Fatalf("no-route handler sni=%q, want %q", gotSNI, "tenant.example.com")
270 - }
271 - case <-time.After(500 * time.Millisecond):
272 - t.Fatal("no-route handler was not called")
273 - }
274 -
275 - select {
276 - case <-done:
277 - case <-time.After(500 * time.Millisecond):
278 - t.Fatal("handleConnection did not return after handled no-route callback")
279 - }
280 -
281 - _ = client.SetReadDeadline(time.Now().Add(75 * time.Millisecond))
282 - var b [1]byte
283 - _, err := client.Read(b[:])
284 - if err == nil {
285 - t.Fatal("expected read timeout while connection remains open")
286 - }
287 - var netErr net.Error
288 - if !errors.As(err, &netErr) || !netErr.Timeout() {
289 - t.Fatalf("expected timeout error to indicate open connection, got: %v", err)
290 - }
291 -}
292 -
293 -func TestRouter_HandleConnectionNoRouteHandlerDeclined(t *testing.T) {
294 - router := NewRouter("")
295 - noRouteCalls := make(chan string, 1)
296 - router.SetNoRouteHandler(func(_ net.Conn, sni string) bool {
297 - select {
298 - case noRouteCalls <- sni:
299 - default:
300 - }
301 - return false
302 - })
303 -
304 - client, server := net.Pipe()
305 - defer func() {
306 - _ = client.Close()
307 - _ = server.Close()
308 - }()
309 -
310 - done := make(chan struct{})
311 - router.wg.Add(1)
312 - go func() {
313 - router.handleConnection(server)
314 - close(done)
315 - }()
316 -
317 - if _, err := client.Write(buildClientHello("tenant.example.com", true)); err != nil {
318 - t.Fatalf("write client hello: %v", err)
319 - }
320 -
321 - select {
322 - case <-noRouteCalls:
323 - case <-time.After(500 * time.Millisecond):
324 - t.Fatal("no-route handler was not called")
325 - }
326 -
327 - select {
328 - case <-done:
329 - case <-time.After(500 * time.Millisecond):
330 - t.Fatal("handleConnection did not return after declined no-route callback")
331 - }
332 -
333 - _ = client.SetReadDeadline(time.Now().Add(250 * time.Millisecond))
334 - var b [1]byte
335 - _, err := client.Read(b[:])
336 - if err == nil {
337 - t.Fatal("expected closed connection after declined no-route callback")
338 - }
339 - var netErr net.Error
340 - if errors.As(err, &netErr) && netErr.Timeout() {
341 - t.Fatalf("expected closed-connection error, got timeout: %v", err)
342 - }
343 -}
344 -
345 -func TestCloseWithDebugLogNilCloser(_ *testing.T) {
346 - closeWithDebugLog(nil, "noop")
347 -}
portal/tls.go new
+62
@@ -0,0 +1,62 @@
1 +package portal
2 +
3 +import (
4 + "crypto/tls"
5 + "fmt"
6 + "io"
7 + "net/http"
8 +
9 + keylesslib "github.com/gosuda/keyless_tls/keyless"
10 +)
11 +
12 +type TLSMaterialConfig struct {
13 + CertPEM []byte
14 + KeyPEM []byte
15 + Keyless *RemoteSignerConfig
16 +}
17 +
18 +type RemoteSignerConfig struct {
19 + Endpoint string
20 + ServerName string
21 + KeyID string
22 + ClientCertPEM []byte
23 + ClientKeyPEM []byte
24 + RootCAPEM []byte
25 +}
26 +
27 +func attachAPITLS(server *http.Server, cfg TLSMaterialConfig) (io.Closer, error) {
28 + if server == nil {
29 + return nil, fmt.Errorf("http server is required")
30 + }
31 + if cfg.Keyless != nil {
32 + remoteSigner, err := keylesslib.AttachToHTTPServer(server, keylesslib.HTTPServerAttachConfig{
33 + CertPEM: cfg.CertPEM,
34 + RemoteSigner: keylesslib.RemoteSignerConfig{
35 + Endpoint: cfg.Keyless.Endpoint,
36 + ServerName: cfg.Keyless.ServerName,
37 + KeyID: cfg.Keyless.KeyID,
38 + ClientCertPEM: cfg.Keyless.ClientCertPEM,
39 + ClientKeyPEM: cfg.Keyless.ClientKeyPEM,
40 + RootCAPEM: cfg.Keyless.RootCAPEM,
41 + },
42 + NextProtos: []string{"http/1.1"},
43 + MinTLSVersion: tls.VersionTLS12,
44 + })
45 + if err != nil {
46 + return nil, err
47 + }
48 + return remoteSigner, nil
49 + }
50 +
51 + cert, err := tls.X509KeyPair(cfg.CertPEM, cfg.KeyPEM)
52 + if err != nil {
53 + return nil, fmt.Errorf("parse api tls key pair: %w", err)
54 + }
55 +
56 + server.TLSConfig = &tls.Config{
57 + MinVersion: tls.VersionTLS12,
58 + NextProtos: []string{"http/1.1"},
59 + Certificates: []tls.Certificate{cert},
60 + }
61 + return nil, nil
62 +}
sdk/client.go
+287 -135
@@ -1,205 +1,357 @@
1 -// Package sdk provides a client for registering leases with the Portal relay.
1 package sdk
2
3 import (
4 + "bufio"
5 + "bytes"
6 + "context"
7 "crypto/rand"
8 "crypto/tls"
9 + "crypto/x509"
10 "encoding/hex"
8 - "errors"
11 + "encoding/json"
12 "fmt"
13 + "io"
14 "net"
15 + "net/http"
16 "net/url"
12 - "sync"
17 + "strings"
18 "time"
19
15 - "github.com/rs/zerolog/log"
16 -
17 - "gosuda.org/portal/portal/keyless"
18 - "gosuda.org/portal/types"
19 -)
20 -
21 -// SDK-specific errors.
22 -var (
23 - ErrNoAvailableRelay = errors.New("no available relay")
24 - ErrInvalidName = errors.New("lease name must be a DNS label (letters, digits, hyphen; no dots or underscores)")
20 + "gosuda.org/portal/portal"
21 )
22
27 -// ClientConfig configures the SDK client.
23 type ClientConfig struct {
29 - BootstrapServers []string
30 - ReverseDialTimeout time.Duration // Reverse connect dial timeout (default: 5 seconds)
24 + RelayURL string
25 + RootCAPEM []byte
26 + InsecureSkipVerify bool
27 + DialTimeout time.Duration
28 + RequestTimeout time.Duration
29 + HandshakeTimeout time.Duration
30 + LeaseTTL time.Duration
31 + RenewBefore time.Duration
32 + ReadyTarget int
33 }
34
33 -// ClientOption configures ClientConfig.
34 -type ClientOption func(*ClientConfig)
35 +type Client struct {
36 + baseURL *url.URL
37 + httpClient *http.Client
38 + rawTLSConfig *tls.Config
39 + dialTimeout time.Duration
40 + handshakeTimeout time.Duration
41 + leaseTTL time.Duration
42 + renewBefore time.Duration
43 + readyTarget int
44 +}
45
36 -// WithBootstrapServers sets the bootstrap relay servers.
37 -func WithBootstrapServers(servers []string) ClientOption {
38 - return func(c *ClientConfig) {
39 - c.BootstrapServers = servers
46 +func NewClient(cfg ClientConfig) (*Client, error) {
47 + relayURL, err := portal.NormalizeRelayURL(cfg.RelayURL)
48 + if err != nil {
49 + return nil, err
50 + }
51 + baseURL, err := url.Parse(relayURL)
52 + if err != nil {
53 + return nil, fmt.Errorf("parse relay url: %w", err)
54 }
41 -}
55
43 -// WithReverseDialTimeout sets the reverse dial timeout.
44 -func WithReverseDialTimeout(timeout time.Duration) ClientOption {
45 - return func(c *ClientConfig) {
46 - c.ReverseDialTimeout = timeout
56 + rootCAs, err := buildRootCAs(cfg.RootCAPEM)
57 + if err != nil {
58 + return nil, err
59 + }
60 +
61 + cfg.DialTimeout = durationOrDefault(cfg.DialTimeout, 5*time.Second)
62 + cfg.RequestTimeout = durationOrDefault(cfg.RequestTimeout, 15*time.Second)
63 + cfg.HandshakeTimeout = durationOrDefault(cfg.HandshakeTimeout, 15*time.Second)
64 + cfg.LeaseTTL = durationOrDefault(cfg.LeaseTTL, 2*time.Minute)
65 + cfg.RenewBefore = durationOrDefault(cfg.RenewBefore, 30*time.Second)
66 + cfg.ReadyTarget = intOrDefault(cfg.ReadyTarget, 1)
67 +
68 + baseTLS := &tls.Config{
69 + MinVersion: tls.VersionTLS12,
70 + ServerName: baseURL.Hostname(),
71 + RootCAs: rootCAs,
72 + InsecureSkipVerify: cfg.InsecureSkipVerify,
73 + NextProtos: []string{"http/1.1"},
74 + }
75 +
76 + transport := &http.Transport{
77 + TLSClientConfig: baseTLS.Clone(),
78 + ForceAttemptHTTP2: false,
79 }
80 +
81 + return &Client{
82 + baseURL: baseURL,
83 + httpClient: &http.Client{
84 + Transport: transport,
85 + Timeout: cfg.RequestTimeout,
86 + },
87 + rawTLSConfig: baseTLS,
88 + dialTimeout: cfg.DialTimeout,
89 + handshakeTimeout: cfg.HandshakeTimeout,
90 + leaseTTL: cfg.LeaseTTL,
91 + renewBefore: cfg.RenewBefore,
92 + readyTarget: cfg.ReadyTarget,
93 + }, nil
94 }
95
50 -// Client is a minimal client for lease registration with the relay.
51 -type Client struct {
52 - config *ClientConfig
53 - mu sync.Mutex
96 +func (c *Client) Close() {
97 + if c == nil || c.httpClient == nil {
98 + return
99 + }
100 + if transport, ok := c.httpClient.Transport.(*http.Transport); ok {
101 + transport.CloseIdleConnections()
102 + }
103 }
104
56 -// NewClient creates a new SDK client.
57 -func NewClient(opt ...ClientOption) (*Client, error) {
58 - config := &ClientConfig{
59 - BootstrapServers: []string{},
60 - ReverseDialTimeout: 5 * time.Second,
105 +func (c *Client) Listen(ctx context.Context, req ListenRequest) (*Listener, error) {
106 + if strings.TrimSpace(req.Name) == "" {
107 + return nil, fmt.Errorf("listener name is required")
108 }
109
63 - for _, o := range opt {
64 - o(config)
110 + reverseToken := strings.TrimSpace(req.ReverseToken)
111 + if reverseToken == "" {
112 + reverseToken = randomToken()
113 }
114
67 - return &Client{config: config}, nil
68 -}
115 + readyTarget := intOrDefault(req.ReadyTarget, c.readyTarget)
116 + leaseTTL := durationOrDefault(req.LeaseTTL, c.leaseTTL)
117 +
118 + registerReq := portal.RegisterRequest{
119 + Name: req.Name,
120 + Hostnames: append([]string(nil), req.Hostnames...),
121 + Metadata: req.Metadata,
122 + ReverseToken: reverseToken,
123 + TLS: true,
124 + TTLSeconds: int(leaseTTL / time.Second),
125 + }
126
70 -// Listen creates a listener and registers it with the relay.
71 -// In reverse-connect mode (TCP tunnel + TLS SNI routing), this registers the
72 -// lease and returns a listener that accepts relay-proxied connections.
73 -func (c *Client) Listen(name string, options ...types.MetadataOption) (net.Listener, error) {
74 - c.mu.Lock()
75 - defer c.mu.Unlock()
127 + var registerResp portal.RegisterResponse
128 + if err := c.doJSON(ctx, http.MethodPost, "/sdk/register", registerReq, &registerResp); err != nil {
129 + return nil, err
130 + }
131
77 - if name == "" {
78 - return nil, errors.New("name is required")
132 + var (
133 + tlsConf *tls.Config
134 + tlsCloser io.Closer
135 + err error
136 + )
137 + if len(req.TLS.CertPEM) > 0 || len(req.TLS.KeyPEM) > 0 || req.TLS.Keyless != nil {
138 + tlsConf, tlsCloser, err = buildTenantTLSConfig(req.TLS)
139 + if err != nil {
140 + _ = c.unregisterLease(context.Background(), registerResp.LeaseID, reverseToken)
141 + return nil, err
142 + }
143 + } else {
144 + tlsConf, err = buildAutoTenantTLSConfig(registerResp.Hostnames)
145 + if err != nil {
146 + _ = c.unregisterLease(context.Background(), registerResp.LeaseID, reverseToken)
147 + return nil, err
148 + }
149 }
80 - if !types.IsValidServiceName(name) {
81 - return nil, ErrInvalidName
150 +
151 + listenerCtx, cancel := context.WithCancel(ctx)
152 + l := &Listener{
153 + client: c,
154 + ctx: listenerCtx,
155 + cancel: cancel,
156 + leaseID: registerResp.LeaseID,
157 + hostnames: append([]string(nil), registerResp.Hostnames...),
158 + metadata: registerResp.Metadata,
159 + reverseToken: reverseToken,
160 + leaseTTL: leaseTTL,
161 + readyTarget: readyTarget,
162 + tlsConfig: tlsConf,
163 + tlsCloser: tlsCloser,
164 + accepted: make(chan net.Conn, max(readyTarget*2, 1)),
165 + signal: make(chan struct{}, 1),
166 }
167
84 - relayAddrs, err := types.NormalizeRelayAPIURLs(c.config.BootstrapServers)
168 + go l.runSupervisor()
169 + go l.runRenewLoop()
170 + l.notify()
171 + return l, nil
172 +}
173 +
174 +func (c *Client) doJSON(ctx context.Context, method, path string, payload any, out any) error {
175 + var body io.Reader
176 + if payload != nil {
177 + buf, err := json.Marshal(payload)
178 + if err != nil {
179 + return fmt.Errorf("marshal payload: %w", err)
180 + }
181 + body = bytes.NewReader(buf)
182 + }
183 +
184 + req, err := http.NewRequestWithContext(ctx, method, c.resolve(path), body)
185 if err != nil {
86 - return nil, ErrNoAvailableRelay
186 + return err
187 }
188 + req.Header.Set("Content-Type", "application/json")
189
89 - lease, err := c.newLease(name, options...)
190 + resp, err := c.httpClient.Do(req)
191 if err != nil {
91 - return nil, err
192 + return err
193 }
194 + defer resp.Body.Close()
195
94 - listeners := make([]net.Listener, 0, len(relayAddrs))
95 - closeActiveListeners := func() {
96 - for _, listener := range listeners {
97 - _ = listener.Close()
196 + var envelope apiEnvelope
197 + if err := json.NewDecoder(resp.Body).Decode(&envelope); err != nil {
198 + return fmt.Errorf("decode response: %w", err)
199 + }
200 + if !envelope.OK {
201 + if envelope.Error == nil {
202 + return fmt.Errorf("api request failed with status %d", resp.StatusCode)
203 }
204 + return fmt.Errorf("%s: %s", envelope.Error.Code, envelope.Error.Message)
205 }
206 + if out == nil {
207 + return nil
208 + }
209 + return json.Unmarshal(envelope.Data, out)
210 +}
211
101 - runCloseFns := func(closeFns []func()) {
102 - for _, closeFn := range closeFns {
103 - if closeFn != nil {
104 - closeFn()
105 - }
106 - }
212 +func (c *Client) renewLease(ctx context.Context, leaseID, reverseToken string, ttl time.Duration) error {
213 + return c.doJSON(ctx, http.MethodPost, "/sdk/renew", portal.RenewRequest{
214 + LeaseID: leaseID,
215 + ReverseToken: reverseToken,
216 + TTLSeconds: int(ttl / time.Second),
217 + }, &portal.RenewResponse{})
218 +}
219 +
220 +func (c *Client) unregisterLease(ctx context.Context, leaseID, reverseToken string) error {
221 + return c.doJSON(ctx, http.MethodPost, "/sdk/unregister", portal.UnregisterRequest{
222 + LeaseID: leaseID,
223 + ReverseToken: reverseToken,
224 + }, nil)
225 +}
226 +
227 +func (c *Client) openReverseSession(ctx context.Context, leaseID, reverseToken string) (net.Conn, error) {
228 + dialer := &tls.Dialer{
229 + NetDialer: &net.Dialer{Timeout: c.dialTimeout},
230 + Config: c.rawTLSConfig.Clone(),
231 }
232
109 - for _, relayAddr := range relayAddrs {
110 - tlsConfig, listenerCloseFns, tlsErr := c.buildTLSConfig(relayAddr, name)
111 - if tlsErr != nil {
112 - closeActiveListeners()
113 - return nil, tlsErr
114 - }
233 + conn, err := dialer.DialContext(ctx, "tcp", ensurePort(c.baseURL.Host))
234 + if err != nil {
235 + return nil, err
236 + }
237
116 - leaseCopy := *lease
117 - listener, listenerErr := NewListener(relayAddr, &leaseCopy, tlsConfig, 0, c.config.ReverseDialTimeout, listenerCloseFns...)
118 - if listenerErr != nil {
119 - runCloseFns(listenerCloseFns)
120 - closeActiveListeners()
121 - return nil, fmt.Errorf("create relay listener: %w", listenerErr)
122 - }
238 + connectURL, err := url.Parse(c.resolve("/sdk/connect"))
239 + if err != nil {
240 + _ = conn.Close()
241 + return nil, err
242 + }
243 + query := connectURL.Query()
244 + query.Set("lease_id", leaseID)
245 + connectURL.RawQuery = query.Encode()
246 +
247 + req := &http.Request{
248 + Method: http.MethodGet,
249 + URL: connectURL,
250 + Host: c.baseURL.Host,
251 + Header: make(http.Header),
252 + }
253 + req.Header.Set(portal.HeaderReverseToken, reverseToken)
254 + req.Header.Set("Connection", "keep-alive")
255
124 - if startErr := listener.Start(); startErr != nil {
125 - _ = listener.Close()
126 - closeActiveListeners()
127 - return nil, fmt.Errorf("start relay listener: %w", startErr)
128 - }
256 + if err := req.Write(conn); err != nil {
257 + _ = conn.Close()
258 + return nil, err
259 + }
260
130 - listeners = append(listeners, listener)
261 + reader := bufio.NewReader(conn)
262 + resp, err := http.ReadResponse(reader, req)
263 + if err != nil {
264 + _ = conn.Close()
265 + return nil, err
266 }
267 + defer resp.Body.Close()
268
133 - listener := net.Listener(newMultiRelayListener(lease.ID, listeners))
134 - if len(listeners) == 1 {
135 - listener = listeners[0]
269 + if resp.StatusCode != http.StatusOK {
270 + body, _ := io.ReadAll(io.LimitReader(resp.Body, 8<<10))
271 + _ = conn.Close()
272 + return nil, fmt.Errorf("reverse connect failed: %s", strings.TrimSpace(string(body)))
273 }
274
138 - log.Info().
139 - Str("lease_id", lease.ID).
140 - Str("name", name).
141 - Bool("tls", true).
142 - Msg("[SDK] Lease registered with TLS")
275 + return wrapBufferedConn(conn, reader), nil
276 +}
277
144 - return listener, nil
278 +func (c *Client) resolve(path string) string {
279 + ref, _ := url.Parse(path)
280 + return c.baseURL.ResolveReference(ref).String()
281 }
282
147 -func (c *Client) newLease(name string, options ...types.MetadataOption) (*types.Lease, error) {
148 - var metadata types.Metadata
149 - for _, option := range options {
150 - option(&metadata)
151 - }
283 +type apiEnvelope struct {
284 + OK bool `json:"ok"`
285 + Data json.RawMessage `json:"data"`
286 + Error *portal.APIError `json:"error"`
287 +}
288
153 - idBytes := make([]byte, 16)
154 - if _, err := rand.Read(idBytes); err != nil {
155 - return nil, fmt.Errorf("generate lease ID: %w", err)
289 +func buildRootCAs(rootCAPEM []byte) (*x509.CertPool, error) {
290 + if len(rootCAPEM) == 0 {
291 + return nil, nil
292 }
293 + pool := x509.NewCertPool()
294 + if !pool.AppendCertsFromPEM(rootCAPEM) {
295 + return nil, fmt.Errorf("failed to parse relay root ca")
296 + }
297 + return pool, nil
298 +}
299
158 - tokenBytes := make([]byte, 16)
159 - if _, err := rand.Read(tokenBytes); err != nil {
160 - return nil, fmt.Errorf("generate reverse token: %w", err)
300 +func durationOrDefault(v, fallback time.Duration) time.Duration {
301 + if v > 0 {
302 + return v
303 }
304 + return fallback
305 +}
306
163 - lease := &types.Lease{
164 - ID: hex.EncodeToString(idBytes),
165 - Name: name,
166 - TLS: true,
167 - ReverseToken: hex.EncodeToString(tokenBytes),
168 - Metadata: types.Metadata{
169 - Description: metadata.Description,
170 - Tags: metadata.Tags,
171 - Thumbnail: metadata.Thumbnail,
172 - Owner: metadata.Owner,
173 - Hide: metadata.Hide,
174 - },
175 - Expires: time.Now().Add(30 * time.Second),
307 +func intOrDefault(v, fallback int) int {
308 + if v > 0 {
309 + return v
310 }
177 - return lease, nil
311 + return fallback
312 }
313
180 -func (c *Client) buildTLSConfig(relayAddr, leaseName string) (*tls.Config, []func(), error) {
181 - parsed, err := url.Parse(relayAddr)
182 - if err != nil {
183 - return nil, nil, fmt.Errorf("invalid relay address: %s, %w", relayAddr, err)
314 +func max(a, b int) int {
315 + if a > b {
316 + return a
317 }
185 - keylessServerName := parsed.Hostname()
186 - if keylessServerName == "" {
187 - return nil, nil, fmt.Errorf("relay hostname is required: %s", relayAddr)
318 + return b
319 +}
320 +
321 +func randomToken() string {
322 + buf := make([]byte, 8)
323 + if _, err := rand.Read(buf); err != nil {
324 + panic(err)
325 }
189 - baseHost := types.PortalRootHost(relayAddr)
190 - if baseHost == "" {
191 - return nil, nil, fmt.Errorf("keyless base host is required for relay %s", relayAddr)
326 + return "tok_" + hex.EncodeToString(buf)
327 +}
328 +
329 +func ensurePort(host string) string {
330 + if _, _, err := net.SplitHostPort(host); err == nil {
331 + return host
332 }
193 - domain := leaseName + "." + baseHost
333 + return net.JoinHostPort(host, "443")
334 +}
335
195 - tlsConfig, closeFn, err := keyless.BuildClientTLSConfig(relayAddr, keylessServerName, domain)
196 - if err != nil {
197 - return nil, nil, err
336 +type bufferedConn struct {
337 + net.Conn
338 + reader *bytes.Reader
339 +}
340 +
341 +func wrapBufferedConn(conn net.Conn, reader *bufio.Reader) net.Conn {
342 + if reader == nil || reader.Buffered() == 0 {
343 + return conn
344 }
199 - return tlsConfig, []func(){closeFn}, nil
345 + buf := make([]byte, reader.Buffered())
346 + if _, err := io.ReadFull(reader, buf); err != nil {
347 + return conn
348 + }
349 + return &bufferedConn{Conn: conn, reader: bytes.NewReader(buf)}
350 }
351
202 -// Close is a no-op kept for caller compatibility.
203 -func (c *Client) Close() error {
204 - return nil
352 +func (c *bufferedConn) Read(p []byte) (int, error) {
353 + if c.reader != nil && c.reader.Len() > 0 {
354 + return c.reader.Read(p)
355 + }
356 + return c.Conn.Read(p)
357 }
sdk/listener.go
+133 -755
@@ -1,858 +1,236 @@
1 package sdk
2
3 import (
4 - "bufio"
5 - "bytes"
4 "context"
5 "crypto/tls"
8 - "encoding/json"
6 "errors"
7 "fmt"
8 "io"
9 "net"
13 - "net/http"
14 - "net/url"
15 - "strings"
10 "sync"
11 "time"
12
19 - "github.com/rs/zerolog/log"
20 -
21 - "gosuda.org/portal/types"
22 -)
23 -
24 -const (
25 - relayKeepaliveInterval = 10 * time.Second
26 - reverseReadTimeout = 1 * time.Second
27 - defaultReverseWorkers = 16
28 - defaultReverseDialTimeout = 5 * time.Second
29 - defaultTLSHandshakeTimeout = 10 * time.Second
13 + "gosuda.org/portal/portal"
14 )
15
32 -var fatalReverseConnectRejectionCodes = map[string]struct{}{
33 - "ip_banned": {},
34 - "lease_not_found": {},
35 - "method_not_allowed": {},
36 - "missing_lease_id": {},
37 - "missing_reverse_token": {},
38 - "tls_required": {},
39 - "unauthorized": {},
40 - "unsupported_transport": {},
41 -}
42 -
43 -type reverseConnectRejectionError struct {
44 - code string
45 - detail string
46 - statusCode int
47 -}
48 -
49 -func (e *reverseConnectRejectionError) Error() string {
50 - if e == nil {
51 - return "reverse connect rejected"
52 - }
53 - if e.detail == "" {
54 - return fmt.Sprintf("reverse connect rejected: status=%d", e.statusCode)
55 - }
56 - return fmt.Sprintf("reverse connect rejected: status=%d error=%s", e.statusCode, e.detail)
57 -}
16 +type LeaseMetadata = portal.LeaseMetadata
17
59 -func (e *reverseConnectRejectionError) IsFatal() bool {
60 - if e == nil {
61 - return false
62 - }
63 - if _, ok := fatalReverseConnectRejectionCodes[e.code]; ok {
64 - return true
65 - }
66 - switch e.statusCode {
67 - case http.StatusBadRequest,
68 - http.StatusUnauthorized,
69 - http.StatusForbidden,
70 - http.StatusNotFound,
71 - http.StatusMethodNotAllowed,
72 - http.StatusUpgradeRequired:
73 - return true
74 - default:
75 - return false
76 - }
18 +type ListenRequest struct {
19 + Name string
20 + Hostnames []string
21 + Metadata LeaseMetadata
22 + ReverseToken string
23 + ReadyTarget int
24 + LeaseTTL time.Duration
25 + TLS portal.TLSMaterialConfig
26 }
27
79 -// Listener is a net.Listener backed by relay tunnel registration.
80 -// The relay connects to this listener after SNI routing resolves the lease.
28 type Listener struct {
82 - tlsConfig *tls.Config
83 - lease *types.Lease
84 - httpClient *http.Client
85 - stopCh chan struct{}
86 - acceptCh chan net.Conn
87 - relayAddr string
88 - closeFns []func()
89 - wg sync.WaitGroup
90 - reverseWorkers int
91 - reverseDialTimeout time.Duration
92 - mu sync.RWMutex
93 - closeOnce sync.Once
94 - closed bool
29 + client *Client
30 + ctx context.Context
31 + cancel context.CancelFunc
32 + leaseID string
33 + hostnames []string
34 + metadata LeaseMetadata
35 + reverseToken string
36 + leaseTTL time.Duration
37 + readyTarget int
38 + tlsConfig *tls.Config
39 + tlsCloser io.Closer
40 +
41 + accepted chan net.Conn
42 + signal chan struct{}
43 +
44 + mu sync.Mutex
45 + activeSessions int
46 + closeOnce sync.Once
47 }
48
97 -var _ net.Listener = (*Listener)(nil)
98 -
99 -// NewListener creates a relay-backed listener.
100 -// If tlsConfig is provided, reverse workers complete TLS handshakes before enqueueing connections.
101 -func NewListener(relayAddr string, lease *types.Lease, tlsConfig *tls.Config, reverseWorkers int, reverseDialTimeout time.Duration, closeFns ...func()) (*Listener, error) {
102 - if lease == nil {
103 - return nil, errors.New("lease is required")
104 - }
105 - if lease.ID == "" {
106 - return nil, errors.New("lease ID is required")
107 - }
108 - if lease.Name == "" {
109 - return nil, errors.New("lease name is required")
110 - }
111 - if lease.ReverseToken == "" {
112 - return nil, errors.New("lease reverse token is required")
113 - }
114 - if tlsConfig == nil {
115 - return nil, errors.New("tls config is required")
116 - }
117 - apiURL, err := types.NormalizeRelayAPIURL(relayAddr)
118 - if err != nil {
119 - return nil, err
120 - }
121 - host := types.PortalRootHost(apiURL)
122 - clientTransport := http.DefaultTransport.(*http.Transport).Clone()
123 - transportTLSConfig := &tls.Config{
124 - MinVersion: tls.VersionTLS12,
125 - ServerName: host,
126 - InsecureSkipVerify: types.IsLocalhost(host),
127 - }
128 - clientTransport.TLSClientConfig = transportTLSConfig
129 -
130 - if reverseWorkers <= 0 {
131 - reverseWorkers = defaultReverseWorkers
132 - }
133 - if reverseDialTimeout <= 0 {
134 - reverseDialTimeout = defaultReverseDialTimeout
135 - }
136 - lease.TLS = true
137 -
138 - return &Listener{
139 - relayAddr: apiURL,
140 - lease: lease,
141 - httpClient: &http.Client{
142 - Timeout: 10 * time.Second,
143 - Transport: clientTransport,
144 - },
145 - tlsConfig: tlsConfig,
146 - closeFns: closeFns,
147 - stopCh: make(chan struct{}),
148 - acceptCh: make(chan net.Conn, 128),
149 - reverseWorkers: reverseWorkers,
150 - reverseDialTimeout: reverseDialTimeout,
151 - }, nil
152 -}
153 -
154 -// Start registers the lease with relay and starts reverse workers.
155 -func (l *Listener) Start() error {
156 - l.mu.Lock()
157 - if l.closed {
158 - l.mu.Unlock()
159 - return net.ErrClosed
160 - }
161 - l.mu.Unlock()
162 -
163 - if err := l.registerWithRelay(); err != nil {
164 - return fmt.Errorf("register lease with relay: %w", err)
165 - }
166 -
167 - l.wg.Add(1)
168 - go l.keepaliveLoop()
169 - for i := range l.reverseWorkers {
170 - l.wg.Add(1)
171 - go l.reverseAcceptWorker(i)
172 - }
173 -
174 - log.Info().
175 - Str("lease_id", l.lease.ID).
176 - Str("name", l.lease.Name).
177 - Str("relay", l.relayAddr).
178 - Int("reverse_workers", l.reverseWorkers).
179 - Msg("[SDK] Relay listener started")
180 -
181 - return nil
182 -}
183 -
184 -// Accept waits for the next connection from relay.
185 -// Reverse workers deliver ready connections to acceptCh.
49 func (l *Listener) Accept() (net.Conn, error) {
187 - l.mu.RLock()
188 - closed := l.closed
189 - l.mu.RUnlock()
190 - if closed {
191 - return nil, net.ErrClosed
192 - }
193 -
194 - var conn net.Conn
50 select {
196 - case <-l.stopCh:
51 + case <-l.ctx.Done():
52 return nil, net.ErrClosed
198 - case conn = <-l.acceptCh:
53 + case conn := <-l.accepted:
54 if conn == nil {
55 return nil, net.ErrClosed
56 }
57 + return conn, nil
58 }
203 - return conn, nil
59 }
60
206 -// Close unregisters lease from relay.
61 func (l *Listener) Close() error {
208 - var retErr error
62 + var closeErr error
63 l.closeOnce.Do(func() {
210 - close(l.stopCh)
211 -
212 - l.mu.Lock()
213 - l.closed = true
214 - l.mu.Unlock()
64 + l.cancel()
65
216 - l.wg.Wait()
217 - for _, closeFn := range l.closeFns {
218 - if closeFn != nil {
219 - closeFn()
220 - }
66 + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
67 + defer cancel()
68 + if err := l.client.unregisterLease(ctx, l.leaseID, l.reverseToken); err != nil {
69 + closeErr = err
70 }
222 -
223 - if err := l.unregisterFromRelay(); err != nil {
224 - log.Warn().Err(err).Str("lease_id", l.lease.ID).Msg("[SDK] Failed to unregister lease")
225 - retErr = err
71 + if l.tlsCloser != nil {
72 + _ = l.tlsCloser.Close()
73 }
74 })
228 -
229 - return retErr
75 + return closeErr
76 }
77
232 -// Addr returns a dummy address (connections come via reverse tunnel).
78 func (l *Listener) Addr() net.Addr {
234 - return &net.TCPAddr{IP: net.IPv4(0, 0, 0, 0), Port: 0}
79 + return listenerAddr("portal:" + l.leaseID)
80 }
81
237 -// LeaseID returns lease ID registered to relay.
82 func (l *Listener) LeaseID() string {
239 - return l.lease.ID
83 + return l.leaseID
84 }
85
242 -func (l *Listener) keepaliveLoop() {
243 - defer l.wg.Done()
86 +func (l *Listener) Hostnames() []string {
87 + return append([]string(nil), l.hostnames...)
88 +}
89
245 - ticker := time.NewTicker(relayKeepaliveInterval)
246 - defer ticker.Stop()
90 +func (l *Listener) Metadata() LeaseMetadata {
91 + return l.metadata
92 +}
93
248 - for {
249 - select {
250 - case <-l.stopCh:
251 - return
252 - case <-ticker.C:
253 - if err := l.sendKeepalive(); err != nil {
254 - if isLeaseNotFoundError(err) {
255 - if rerr := l.registerWithRelay(); rerr != nil {
256 - log.Warn().
257 - Err(rerr).
258 - Str("lease_id", l.lease.ID).
259 - Msg("[SDK] Relay keepalive failed and re-register failed")
260 - } else {
261 - log.Info().
262 - Str("lease_id", l.lease.ID).
263 - Str("name", l.lease.Name).
264 - Msg("[SDK] Lease re-registered after relay reset")
265 - }
266 - continue
267 - }
268 - log.Warn().Err(err).Str("lease_id", l.lease.ID).Msg("[SDK] Relay keepalive failed")
269 - }
270 - }
94 +func (l *Listener) PublicURLs() []string {
95 + urls := make([]string, 0, len(l.hostnames))
96 + for _, host := range l.hostnames {
97 + urls = append(urls, "https://"+host)
98 }
99 + return urls
100 }
101
274 -func (l *Listener) reverseAcceptWorker(workerID int) {
275 - defer l.wg.Done()
276 -
102 +func (l *Listener) runSupervisor() {
103 for {
104 select {
279 - case <-l.stopCh:
105 + case <-l.ctx.Done():
106 return
281 - default:
107 + case <-l.signal:
108 }
109
284 - conn, err := l.openReverseConnection()
285 - if err != nil {
286 - var rejectionErr *reverseConnectRejectionError
287 - if errors.As(err, &rejectionErr) && rejectionErr.IsFatal() {
288 - event := log.Error().
289 - Err(err).
290 - Str("lease_id", l.lease.ID).
291 - Int("worker_id", workerID).
292 - Int("status_code", rejectionErr.statusCode)
293 - if rejectionErr.code != "" {
294 - event = event.Str("relay_error_code", rejectionErr.code)
295 - }
296 - event.Msg("[SDK] Fatal reverse connect rejection; stopping worker")
297 - return
298 - }
299 - select {
300 - case <-l.stopCh:
301 - return
302 - case <-time.After(500 * time.Millisecond):
303 - }
304 - continue
110 + for l.reserveSessionSlot() {
111 + go l.runSession()
112 }
306 -
307 - err = l.waitForReverseStart(conn, types.TLSStartMarker)
308 - if err != nil {
309 - if closeErr := conn.Close(); closeErr != nil {
310 - log.Debug().Err(closeErr).Msg("[SDK] failed to close reverse connection")
311 - }
312 - if errors.Is(err, net.ErrClosed) {
313 - return
314 - }
315 - if errors.Is(err, io.EOF) {
316 - continue
317 - }
318 - log.Debug().
319 - Err(err).
320 - Str("lease_id", l.lease.ID).
321 - Int("worker_id", workerID).
322 - Msg("[SDK] Reverse wait failed")
323 - continue
324 - }
325 -
326 - conn, err = l.prepareAcceptedConnection(conn)
327 - if err != nil {
328 - if errors.Is(err, net.ErrClosed) {
329 - return
330 - }
331 - if errors.Is(err, io.EOF) {
332 - continue
333 - }
334 - log.Debug().
335 - Err(err).
336 - Str("lease_id", l.lease.ID).
337 - Int("worker_id", workerID).
338 - Msg("[SDK] Reverse connection preparation failed")
339 - continue
340 - }
341 -
342 - select {
343 - case <-l.stopCh:
344 - if closeErr := conn.Close(); closeErr != nil {
345 - log.Debug().Err(closeErr).Msg("[SDK] failed to close reverse connection on shutdown")
346 - }
347 - return
348 - case l.acceptCh <- conn:
349 - }
350 - }
351 -}
352 -
353 -func (l *Listener) openReverseConnection() (net.Conn, error) {
354 - if l.isStopping() {
355 - return nil, net.ErrClosed
356 - }
357 - connectURL, err := relayConnectURL(l.relayAddr, l.lease.ID, l.lease.ReverseToken)
358 - if err != nil {
359 - return nil, err
360 - }
361 - u, err := url.Parse(connectURL)
362 - if err != nil {
363 - return nil, fmt.Errorf("parse reverse connect URL: %w", err)
364 - }
365 - address := u.Host
366 - if address == "" {
367 - return nil, errors.New("reverse connect URL missing host")
368 - }
369 - if u.Scheme != "https" {
370 - return nil, errors.New("reverse connect must use https scheme")
371 - }
372 - if _, _, splitErr := net.SplitHostPort(address); splitErr != nil {
373 - address = net.JoinHostPort(address, "443")
374 - }
375 -
376 - timeout := l.reverseSetupTimeout()
377 - ctx, cancel := l.newStopAwareContext(timeout)
378 - defer cancel()
379 - dialer := &net.Dialer{
380 - Timeout: timeout,
381 - }
382 - rawConn, err := dialer.DialContext(ctx, "tcp", address)
383 - if err != nil {
384 - if l.isStopping() || errors.Is(err, context.Canceled) {
385 - return nil, net.ErrClosed
386 - }
387 - return nil, fmt.Errorf("dial reverse tcp: %w", err)
388 - }
389 - stopConnWatch := l.closeConnOnStop(rawConn)
390 - defer stopConnWatch()
391 -
392 - serverName := u.Hostname()
393 - if serverName == "" {
394 - _ = rawConn.Close()
395 - return nil, errors.New("reverse connect URL missing TLS server name")
396 - }
397 - reverseTLSConfig := &tls.Config{
398 - MinVersion: tls.VersionTLS12,
399 - ServerName: serverName,
400 - InsecureSkipVerify: types.IsLocalhost(serverName),
401 - }
402 - tlsConn := tls.Client(rawConn, reverseTLSConfig)
403 - err = tlsConn.HandshakeContext(ctx)
404 - if err != nil {
405 - _ = rawConn.Close()
406 - if l.isStopping() || errors.Is(err, context.Canceled) || errors.Is(err, net.ErrClosed) {
407 - return nil, net.ErrClosed
408 - }
409 - return nil, fmt.Errorf("reverse TLS handshake: %w", err)
113 }
411 - conn := net.Conn(tlsConn)
412 -
413 - err = l.writeReverseConnectRequest(conn, u)
414 - if err != nil {
415 - _ = conn.Close()
416 - return nil, err
417 - }
418 - reader, err := l.readReverseConnectResponse(conn)
419 - if err != nil {
420 - _ = conn.Close()
421 - return nil, err
422 - }
423 - if l.isStopping() {
424 - _ = conn.Close()
425 - return nil, net.ErrClosed
426 - }
427 -
428 - return &bufferedConn{Conn: conn, reader: reader}, nil
114 }
115
431 -func (l *Listener) prepareAcceptedConnection(conn net.Conn) (net.Conn, error) {
432 - l.mu.RLock()
433 - tlsConfig := l.tlsConfig
434 - l.mu.RUnlock()
435 - if tlsConfig == nil {
436 - return conn, nil
116 +func (l *Listener) runRenewLoop() {
117 + interval := l.leaseTTL / 2
118 + if interval <= 0 {
119 + interval = 30 * time.Second
120 }
438 -
439 - tlsConn := tls.Server(conn, tlsConfig)
440 - handshakeCtx, cancel := l.newStopAwareContext(defaultTLSHandshakeTimeout)
441 - defer cancel()
442 - if err := tlsConn.HandshakeContext(handshakeCtx); err != nil {
443 - _ = conn.Close()
444 - if l.isStopping() || errors.Is(err, context.Canceled) || errors.Is(err, net.ErrClosed) {
445 - return nil, net.ErrClosed
446 - }
447 - return nil, fmt.Errorf("TLS handshake failed: %w", err)
448 - }
449 - return tlsConn, nil
450 -}
451 -
452 -func buildReverseConnectRequest(u *url.URL, reverseToken string) (*http.Request, error) {
453 - if u == nil {
454 - return nil, errors.New("reverse connect URL is required")
455 - }
456 - if u.Host == "" {
457 - return nil, errors.New("reverse connect URL missing host")
121 + if l.client.renewBefore > 0 && l.leaseTTL > l.client.renewBefore {
122 + interval = l.leaseTTL - l.client.renewBefore
123 }
459 -
460 - token := strings.TrimSpace(reverseToken)
461 - if token == "" {
462 - return nil, errors.New("reverse token is required")
124 + if interval <= 0 {
125 + interval = 30 * time.Second
126 }
127
465 - requestPath := u.EscapedPath()
466 - if requestPath == "" {
467 - requestPath = "/"
468 - }
469 -
470 - requestURL := &url.URL{
471 - Path: requestPath,
472 - RawPath: u.RawPath,
473 - RawQuery: u.RawQuery,
474 - }
475 - req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, requestURL.String(), nil)
476 - if err != nil {
477 - return nil, fmt.Errorf("build reverse connect request: %w", err)
478 - }
479 - req.Host = u.Host
480 - req.Header.Set(types.ReverseConnectTokenHeader, token)
481 - req.Header.Set("Connection", "keep-alive")
482 - return req, nil
483 -}
484 -
485 -func (l *Listener) writeReverseConnectRequest(conn net.Conn, u *url.URL) error {
486 - timeout := l.reverseSetupTimeout()
487 - if err := conn.SetWriteDeadline(time.Now().Add(timeout)); err != nil {
488 - return fmt.Errorf("set reverse connect write deadline: %w", err)
489 - }
490 - defer func() {
491 - _ = conn.SetWriteDeadline(time.Time{})
492 - }()
128 + ticker := time.NewTicker(interval)
129 + defer ticker.Stop()
130
494 - req, err := buildReverseConnectRequest(u, l.lease.ReverseToken)
495 - if err != nil {
496 - return err
497 - }
498 - if err := req.Write(conn); err != nil {
499 - if l.isStopping() || errors.Is(err, net.ErrClosed) {
500 - return net.ErrClosed
131 + for {
132 + select {
133 + case <-l.ctx.Done():
134 + return
135 + case <-ticker.C:
136 + ctx, cancel := context.WithTimeout(l.ctx, 10*time.Second)
137 + _ = l.client.renewLease(ctx, l.leaseID, l.reverseToken, l.leaseTTL)
138 + cancel()
139 }
502 - return fmt.Errorf("write reverse connect request: %w", err)
140 }
504 - return nil
141 }
142
507 -func (l *Listener) readReverseConnectResponse(conn net.Conn) (*bufio.Reader, error) {
508 - timeout := l.reverseSetupTimeout()
509 - if err := conn.SetReadDeadline(time.Now().Add(timeout)); err != nil {
510 - return nil, fmt.Errorf("set reverse connect read deadline: %w", err)
511 - }
512 - defer func() {
513 - _ = conn.SetReadDeadline(time.Time{})
514 - }()
143 +func (l *Listener) runSession() {
144 + defer l.releaseSessionSlot()
145
516 - reader := bufio.NewReader(conn)
517 - resp, err := http.ReadResponse(reader, &http.Request{Method: http.MethodGet})
146 + conn, err := l.client.openReverseSession(l.ctx, l.leaseID, l.reverseToken)
147 if err != nil {
519 - if l.isStopping() || errors.Is(err, net.ErrClosed) {
520 - return nil, net.ErrClosed
521 - }
522 - return nil, fmt.Errorf("read reverse connect response: %w", err)
523 - }
524 - defer resp.Body.Close()
525 - if resp.StatusCode != http.StatusOK {
526 - body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
527 - code, detail := parseReverseConnectRejection(body)
528 - if detail == "" {
529 - detail = strings.TrimSpace(http.StatusText(resp.StatusCode))
530 - }
531 - return nil, &reverseConnectRejectionError{
532 - statusCode: resp.StatusCode,
533 - code: code,
534 - detail: detail,
535 - }
148 + sleepOrDone(l.ctx, time.Second)
149 + return
150 }
537 - return reader, nil
538 -}
539 -
540 -func parseReverseConnectRejection(body []byte) (string, string) {
541 - trimmedBody := strings.TrimSpace(string(body))
542 - if trimmedBody == "" {
543 - return "", ""
544 - }
545 -
546 - var envelope types.APIRawEnvelope
547 - if err := json.Unmarshal(body, &envelope); err != nil || envelope.Error == nil {
548 - return "", trimmedBody
549 - }
550 -
551 - code := strings.TrimSpace(envelope.Error.Code)
552 - message := strings.TrimSpace(envelope.Error.Message)
553 - switch {
554 - case message != "" && code != "":
555 - return code, fmt.Sprintf("%s (code=%s)", message, code)
556 - case message != "":
557 - return code, message
558 - case code != "":
559 - return code, code
560 - default:
561 - return "", trimmedBody
562 - }
563 -}
564 -
565 -func (l *Listener) reverseSetupTimeout() time.Duration {
566 - if l.reverseDialTimeout <= 0 {
567 - return defaultReverseDialTimeout
568 - }
569 - return l.reverseDialTimeout
570 -}
151
572 -func (l *Listener) newStopAwareContext(timeout time.Duration) (context.Context, context.CancelFunc) {
573 - if timeout <= 0 {
574 - timeout = defaultReverseDialTimeout
575 - }
576 - ctx, cancel := context.WithTimeout(context.Background(), timeout)
577 - go func() {
578 - select {
579 - case <-l.stopCh:
580 - cancel()
581 - case <-ctx.Done():
582 - }
583 - }()
584 - return ctx, cancel
585 -}
586 -
587 -func (l *Listener) closeConnOnStop(conn net.Conn) func() {
588 - done := make(chan struct{})
589 - go func() {
590 - select {
591 - case <-l.stopCh:
592 - _ = conn.Close()
593 - case <-done:
152 + if err := l.awaitActivation(conn); err != nil {
153 + _ = conn.Close()
154 + if !errors.Is(err, context.Canceled) && !errors.Is(err, net.ErrClosed) {
155 + sleepOrDone(l.ctx, time.Second)
156 }
595 - }()
596 - return func() {
597 - close(done)
598 - }
599 -}
600 -
601 -func (l *Listener) isStopping() bool {
602 - select {
603 - case <-l.stopCh:
604 - return true
605 - default:
606 - return false
157 }
158 }
159
610 -func (l *Listener) waitForReverseStart(conn net.Conn, expectedMarker byte) error {
160 +func (l *Listener) awaitActivation(conn net.Conn) error {
161 var marker [1]byte
162 for {
613 - _ = conn.SetReadDeadline(time.Now().Add(reverseReadTimeout))
614 - _, err := io.ReadFull(conn, marker[:])
615 - if err == nil {
616 - _ = conn.SetReadDeadline(time.Time{})
617 - if marker[0] == types.ReverseKeepaliveMarker {
618 - continue
619 - }
620 - if marker[0] == expectedMarker {
621 - return nil
622 - }
623 - return fmt.Errorf("invalid reverse marker: %d", marker[0])
624 - }
625 -
626 - var netErr net.Error
627 - if errors.As(err, &netErr) && netErr.Timeout() {
628 - select {
629 - case <-l.stopCh:
630 - return net.ErrClosed
631 - default:
632 - continue
633 - }
163 + _ = conn.SetReadDeadline(time.Now().Add(2 * l.client.handshakeTimeout))
164 + if _, err := io.ReadFull(conn, marker[:]); err != nil {
165 + return err
166 }
167 + _ = conn.SetReadDeadline(time.Time{})
168
636 - select {
637 - case <-l.stopCh:
638 - return net.ErrClosed
169 + switch marker[0] {
170 + case portal.MarkerKeepalive:
171 + continue
172 + case portal.MarkerTLSStart:
173 + return l.activate(conn)
174 default:
640 - return err
175 + return fmt.Errorf("unexpected reverse marker: 0x%02x", marker[0])
176 }
177 }
178 }
179
645 -func (l *Listener) registerWithRelay() error {
646 - reqBody := types.RegisterRequest{
647 - LeaseID: l.lease.ID,
648 - Name: l.lease.Name,
649 - Metadata: l.lease.Metadata,
650 - TLS: l.lease.TLS,
651 - ReverseToken: l.lease.ReverseToken,
652 - }
653 -
654 - return l.postJSON(types.PathSDKRegister, reqBody)
655 -}
656 -
657 -func (l *Listener) unregisterFromRelay() error {
658 - reqBody := types.UnregisterRequest{
659 - LeaseID: l.lease.ID,
660 - ReverseToken: l.lease.ReverseToken,
661 - }
662 - return l.postJSON(types.PathSDKUnregister, reqBody)
663 -}
664 -
665 -func (l *Listener) sendKeepalive() error {
666 - reqBody := types.RenewRequest{
667 - LeaseID: l.lease.ID,
668 - ReverseToken: l.lease.ReverseToken,
669 - }
670 - return l.postJSON(types.PathSDKRenew, reqBody)
671 -}
672 -
673 -func (l *Listener) postJSON(path string, body any) error {
674 - payload, err := json.Marshal(body)
675 - if err != nil {
676 - return fmt.Errorf("marshal request: %w", err)
677 - }
678 -
679 - endpoint := strings.TrimSuffix(l.relayAddr, "/") + path
680 - req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, endpoint, bytes.NewReader(payload))
681 - if err != nil {
682 - return fmt.Errorf("build POST %s request: %w", path, err)
683 - }
684 - req.Header.Set("Content-Type", "application/json")
685 -
686 - resp, err := l.httpClient.Do(req)
687 - if err != nil {
688 - return fmt.Errorf("POST %s: %w", path, err)
689 - }
690 - defer resp.Body.Close()
691 -
692 - data, _ := io.ReadAll(resp.Body)
693 - if len(data) == 0 {
694 - if resp.StatusCode >= http.StatusOK && resp.StatusCode < http.StatusMultipleChoices {
695 - return nil
696 - }
697 - return fmt.Errorf("POST %s failed: status=%d", path, resp.StatusCode)
698 - }
699 -
700 - var envelope types.APIRawEnvelope
701 - if err := json.Unmarshal(data, &envelope); err != nil {
702 - return fmt.Errorf("POST %s failed: invalid API envelope: %w", path, err)
703 - }
704 - if envelope.OK {
705 - if resp.StatusCode >= http.StatusOK && resp.StatusCode < http.StatusMultipleChoices {
706 - return nil
707 - }
708 - return fmt.Errorf("POST %s failed: status=%d body=%s", path, resp.StatusCode, strings.TrimSpace(string(data)))
180 +func (l *Listener) activate(conn net.Conn) error {
181 + tlsConn := tls.Server(conn, l.tlsConfig.Clone())
182 + handshakeCtx, cancel := context.WithTimeout(l.ctx, l.client.handshakeTimeout)
183 + defer cancel()
184 + if err := tlsConn.HandshakeContext(handshakeCtx); err != nil {
185 + return err
186 }
187
711 - msg := ""
712 - if envelope.Error != nil {
713 - msg = strings.TrimSpace(envelope.Error.Message)
714 - }
715 - if msg == "" {
716 - msg = strings.TrimSpace(string(data))
717 - }
718 - if msg == "" {
719 - msg = fmt.Sprintf("status=%d", resp.StatusCode)
720 - }
721 - if envelope.Error != nil && strings.TrimSpace(envelope.Error.Code) != "" {
722 - return fmt.Errorf("POST %s rejected: %s (code=%s)", path, msg, strings.TrimSpace(envelope.Error.Code))
188 + select {
189 + case <-l.ctx.Done():
190 + _ = tlsConn.Close()
191 + return l.ctx.Err()
192 + case l.accepted <- tlsConn:
193 + return nil
194 }
724 - return fmt.Errorf("POST %s rejected: %s", path, msg)
195 }
196
727 -func isLeaseNotFoundError(err error) bool {
728 - if err == nil {
197 +func (l *Listener) reserveSessionSlot() bool {
198 + l.mu.Lock()
199 + defer l.mu.Unlock()
200 + if l.ctx.Err() != nil {
201 return false
202 }
731 - return strings.Contains(strings.ToLower(err.Error()), "lease not found")
732 -}
733 -
734 -func relayConnectURL(relayAddr, leaseID, token string) (string, error) {
735 - if strings.TrimSpace(leaseID) == "" {
736 - return "", errors.New("leaseID is required")
737 - }
738 - if strings.TrimSpace(token) == "" {
739 - return "", errors.New("reverse token is required")
740 - }
741 -
742 - u, err := url.Parse(relayAddr)
743 - if err != nil {
744 - return "", fmt.Errorf("parse relay URL: %w", err)
745 - }
746 - if u.Scheme != "https" {
747 - return "", fmt.Errorf("unsupported relay URL scheme: %q (use https)", u.Scheme)
203 + if l.activeSessions >= l.readyTarget {
204 + return false
205 }
749 - u.Path = types.PathSDKConnect
750 - q := u.Query()
751 - q.Set("lease_id", leaseID)
752 - u.RawQuery = q.Encode()
753 - u.Fragment = ""
754 - return u.String(), nil
755 -}
756 -
757 -type bufferedConn struct {
758 - net.Conn
759 - reader *bufio.Reader
760 -}
761 -
762 -func (c *bufferedConn) Read(p []byte) (int, error) {
763 - return c.reader.Read(p)
764 -}
765 -
766 -type multiRelayListener struct {
767 - acceptCh chan net.Conn
768 - stopCh chan struct{}
769 - leaseID string
770 - listeners []net.Listener
771 - wg sync.WaitGroup
772 - closeOnce sync.Once
206 + l.activeSessions++
207 + return true
208 }
209
775 -func newMultiRelayListener(leaseID string, listeners []net.Listener) *multiRelayListener {
776 - m := &multiRelayListener{
777 - leaseID: leaseID,
778 - listeners: listeners,
779 - acceptCh: make(chan net.Conn, 128),
780 - stopCh: make(chan struct{}),
781 - }
782 -
783 - for i, listener := range listeners {
784 - m.wg.Add(1)
785 - go m.forwardAccept(i, listener)
786 - }
787 -
788 - return m
210 +func (l *Listener) releaseSessionSlot() {
211 + l.mu.Lock()
212 + l.activeSessions--
213 + l.mu.Unlock()
214 + l.notify()
215 }
216
791 -func (m *multiRelayListener) forwardAccept(index int, listener net.Listener) {
792 - defer m.wg.Done()
793 -
794 - for {
795 - conn, err := listener.Accept()
796 - if err != nil {
797 - select {
798 - case <-m.stopCh:
799 - return
800 - default:
801 - }
802 - if errors.Is(err, net.ErrClosed) {
803 - return
804 - }
805 - log.Debug().
806 - Err(err).
807 - Int("relay_index", index).
808 - Msg("[SDK] relay listener accept failed")
809 - continue
810 - }
811 -
812 - select {
813 - case <-m.stopCh:
814 - _ = conn.Close()
815 - return
816 - case m.acceptCh <- conn:
817 - }
217 +func (l *Listener) notify() {
218 + select {
219 + case l.signal <- struct{}{}:
220 + default:
221 }
222 }
223
821 -func (m *multiRelayListener) Accept() (net.Conn, error) {
224 +func sleepOrDone(ctx context.Context, d time.Duration) {
225 + timer := time.NewTimer(d)
226 + defer timer.Stop()
227 select {
823 - case <-m.stopCh:
824 - return nil, net.ErrClosed
825 - case conn := <-m.acceptCh:
826 - if conn == nil {
827 - return nil, net.ErrClosed
828 - }
829 - return conn, nil
228 + case <-ctx.Done():
229 + case <-timer.C:
230 }
231 }
232
833 -func (m *multiRelayListener) Close() error {
834 - var retErr error
835 - m.closeOnce.Do(func() {
836 - close(m.stopCh)
837 -
838 - for _, listener := range m.listeners {
839 - if err := listener.Close(); err != nil && retErr == nil {
840 - retErr = err
841 - }
842 - }
233 +type listenerAddr string
234
844 - m.wg.Wait()
845 - })
846 - return retErr
847 -}
848 -
849 -func (m *multiRelayListener) Addr() net.Addr {
850 - if len(m.listeners) > 0 {
851 - return m.listeners[0].Addr()
852 - }
853 - return &net.TCPAddr{IP: net.IPv4(0, 0, 0, 0), Port: 0}
854 -}
855 -
856 -func (m *multiRelayListener) LeaseID() string {
857 - return m.leaseID
858 -}
235 +func (a listenerAddr) Network() string { return "portal" }
236 +func (a listenerAddr) String() string { return string(a) }
sdk/listener_test.go
+199 -584
@@ -1,655 +1,270 @@
1 package sdk
2
3 import (
4 + "bufio"
5 + "context"
6 "crypto/tls"
7 "errors"
8 "fmt"
9 "io"
10 "net"
11 "net/http"
10 - "net/url"
12 "strings"
13 "testing"
14 "time"
15
15 - "gosuda.org/portal/types"
16 + "gosuda.org/portal/internal/testutil"
17 + "gosuda.org/portal/portal"
18 )
19
18 -func TestNewListener_Succeeds(t *testing.T) {
20 +func TestListenerEndToEndTLSHTTP(t *testing.T) {
21 t.Parallel()
22
21 - relayAddr := "https://localhost:4017"
22 - lease := &types.Lease{
23 - ID: "test-lease",
24 - Name: "test-app",
25 - ReverseToken: "test-token",
26 - }
27 - tlsConfig := &tls.Config{MinVersion: tls.VersionTLS12}
28 - listener, err := NewListener(relayAddr, lease, tlsConfig, 0, 0)
23 + apiCertPEM, apiKeyPEM, err := testutil.SelfSignedCertPEM("127.0.0.1")
24 if err != nil {
30 - t.Fatalf("NewListener failed: %v", err)
31 - }
32 - if listener == nil {
33 - t.Fatal("NewListener returned nil listener")
34 - }
35 - defer listener.Close()
36 -}
37 -
38 -func TestNormalizeRelayAPIURL(t *testing.T) {
39 - t.Parallel()
40 -
41 - tests := []struct {
42 - name string
43 - in string
44 - want string
45 - wantErr bool
46 - }{
47 - {name: "localhost subdomain to localhost", in: "https://demo-app.localhost:4017", want: "https://localhost:4017"},
48 - {name: "http base rejected", in: "http://example.com", wantErr: true},
49 - {name: "https base", in: "https://example.com/", want: "https://example.com"},
50 - {name: "bare host", in: "localhost:4017", want: "https://localhost:4017"},
51 - {name: "invalid ws scheme", in: "ws://localhost:4017", wantErr: true},
52 - {name: "invalid wss scheme", in: "wss://example.com", wantErr: true},
53 - {name: "invalid relay path", in: "https://localhost:4017/relay", wantErr: true},
54 - {name: "invalid scheme", in: "ftp://example.com", wantErr: true},
55 - {name: "empty", in: "", wantErr: true},
56 - }
57 -
58 - for _, tt := range tests {
59 - t.Run(tt.name, func(t *testing.T) {
60 - t.Parallel()
61 -
62 - got, err := types.NormalizeRelayAPIURL(tt.in)
63 - if tt.wantErr {
64 - if err == nil {
65 - t.Fatalf("expected error for input %q, got none", tt.in)
66 - }
67 - return
68 - }
69 - if err != nil {
70 - t.Fatalf("unexpected error for input %q: %v", tt.in, err)
71 - }
72 - if got != tt.want {
73 - t.Fatalf("types.NormalizeRelayAPIURL(%q) = %q, want %q", tt.in, got, tt.want)
74 - }
75 - })
76 - }
77 -}
78 -
79 -func TestRelayConnectURL(t *testing.T) {
80 - t.Parallel()
81 -
82 - tests := []struct {
83 - name string
84 - relayAddr string
85 - wantScheme string
86 - wantHost string
87 - wantErr bool
88 - }{
89 - {
90 - name: "http relay URL rejected",
91 - relayAddr: "http://localhost:4017",
92 - wantErr: true,
93 - },
94 - {
95 - name: "https relay URL",
96 - relayAddr: "https://relay.example.com",
97 - wantScheme: "https",
98 - wantHost: "relay.example.com",
99 - },
100 - }
101 - for _, tt := range tests {
102 - t.Run(tt.name, func(t *testing.T) {
103 - t.Parallel()
104 -
105 - got, err := relayConnectURL(tt.relayAddr, "lease-1", "token-1")
106 - if tt.wantErr {
107 - if err == nil {
108 - t.Fatalf("expected error for relay %q", tt.relayAddr)
109 - }
110 - return
111 - }
112 - if err != nil {
113 - t.Fatalf("unexpected error: %v", err)
114 - }
115 -
116 - parsed, err := url.Parse(got)
117 - if err != nil {
118 - t.Fatalf("parse URL: %v", err)
119 - }
120 - if parsed.Scheme != tt.wantScheme {
121 - t.Fatalf("unexpected URL scheme: got %q want %q", parsed.Scheme, tt.wantScheme)
122 - }
123 - if parsed.Host != tt.wantHost {
124 - t.Fatalf("unexpected URL host: got %q want %q", parsed.Host, tt.wantHost)
125 - }
126 - if parsed.Path != types.PathSDKConnect {
127 - t.Fatalf("unexpected URL path: got %q want %q", parsed.Path, types.PathSDKConnect)
128 - }
129 - if parsed.Query().Get("lease_id") != "lease-1" {
130 - t.Fatalf("missing lease_id in URL: %q", got)
131 - }
132 - if parsed.Query().Get("token") != "" {
133 - t.Fatalf("token must not be present in URL query: %q", got)
134 - }
135 - })
25 + t.Fatalf("SelfSignedCertPEM(api) error = %v", err)
26 }
137 -
138 - if _, err := relayConnectURL("http://localhost:4017", "", "token-1"); err == nil {
139 - t.Fatal("expected error for empty lease ID")
140 - }
141 - if _, err := relayConnectURL("http://localhost:4017", "lease-1", ""); err == nil {
142 - t.Fatal("expected error for empty token")
143 - }
144 - if _, err := relayConnectURL("ws://localhost:4017", "lease-1", "token-1"); err == nil {
145 - t.Fatal("expected error for unsupported ws scheme")
146 - }
147 -}
148 -
149 -func TestBuildReverseConnectRequest(t *testing.T) {
150 - t.Parallel()
151 -
152 - connectURL, err := relayConnectURL("https://relay.example.com", "lease-1", "token-1")
27 + tenantHost := "app.portal.test"
28 + tenantCertPEM, tenantKeyPEM, err := testutil.SelfSignedCertPEM(tenantHost)
29 if err != nil {
154 - t.Fatalf("relayConnectURL returned error: %v", err)
155 - }
156 -
157 - u, err := url.Parse(connectURL)
30 + t.Fatalf("SelfSignedCertPEM(tenant) error = %v", err)
31 + }
32 +
33 + relay, err := portal.NewServer(portal.ServerConfig{
34 + PortalURL: "https://127.0.0.1",
35 + APIListenAddr: "127.0.0.1:0",
36 + SNIListenAddr: "127.0.0.1:0",
37 + RootHost: "portal.test",
38 + RootFallbackAddr: "127.0.0.1:1",
39 + APITLS: portal.TLSMaterialConfig{
40 + CertPEM: apiCertPEM,
41 + KeyPEM: apiKeyPEM,
42 + },
43 + })
44 if err != nil {
159 - t.Fatalf("parse connect URL: %v", err)
160 - }
161 -
162 - req, err := buildReverseConnectRequest(u, " token-1 ")
45 + t.Fatalf("NewServer() error = %v", err)
46 + }
47 +
48 + ctx, cancel := context.WithCancel(context.Background())
49 + defer cancel()
50 + if err := relay.Start(ctx); err != nil {
51 + t.Fatalf("Start() error = %v", err)
52 + }
53 + t.Cleanup(func() {
54 + shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 5*time.Second)
55 + defer shutdownCancel()
56 + _ = relay.Shutdown(shutdownCtx)
57 + _ = relay.Wait()
58 + })
59 +
60 + client, err := NewClient(ClientConfig{
61 + RelayURL: "https://" + relay.APIAddr(),
62 + InsecureSkipVerify: true,
63 + ReadyTarget: 1,
64 + })
65 if err != nil {
164 - t.Fatalf("buildReverseConnectRequest returned error: %v", err)
165 - }
166 -
167 - if req.Method != http.MethodGet {
168 - t.Fatalf("unexpected request method: got %q want %q", req.Method, http.MethodGet)
169 - }
170 - if req.Host != "relay.example.com" {
171 - t.Fatalf("unexpected host header: got %q want %q", req.Host, "relay.example.com")
172 - }
173 - if req.URL.Path != types.PathSDKConnect {
174 - t.Fatalf("unexpected request path: got %q want %q", req.URL.Path, types.PathSDKConnect)
175 - }
176 - if req.URL.Query().Get("lease_id") != "lease-1" {
177 - t.Fatalf("unexpected lease_id query: %q", req.URL.Query().Get("lease_id"))
178 - }
179 - if req.URL.Query().Get("token") != "" {
180 - t.Fatalf("token must not be present in query: %q", req.URL.RawQuery)
181 - }
182 - if got := req.Header.Get(types.ReverseConnectTokenHeader); got != "token-1" {
183 - t.Fatalf("unexpected reverse token header: got %q want %q", got, "token-1")
184 - }
185 -}
186 -
187 -func TestOpenReverseConnection_RejectsNonHTTPSRelay(t *testing.T) {
188 - t.Parallel()
189 - l := &Listener{
190 - relayAddr: "http://localhost:4017",
191 - lease: &types.Lease{ID: "lease-1", ReverseToken: "token-1"},
192 - reverseDialTimeout: 2 * time.Second,
193 - stopCh: make(chan struct{}),
194 - }
195 -
196 - conn, err := l.openReverseConnection()
66 + t.Fatalf("NewClient() error = %v", err)
67 + }
68 + t.Cleanup(client.Close)
69 +
70 + listener, err := client.Listen(ctx, ListenRequest{
71 + Name: "demo",
72 + Hostnames: []string{tenantHost},
73 + Metadata: LeaseMetadata{
74 + Description: "demo description",
75 + Tags: []string{"demo", "test", "demo"},
76 + Owner: "portal",
77 + Thumbnail: "https://example.test/thumb.png",
78 + Hide: true,
79 + },
80 + TLS: portal.TLSMaterialConfig{
81 + CertPEM: tenantCertPEM,
82 + KeyPEM: tenantKeyPEM,
83 + },
84 + })
85 if err != nil {
198 - if !strings.Contains(err.Error(), "https") {
199 - t.Fatalf("expected https scheme error, got: %v", err)
200 - }
201 - return
86 + t.Fatalf("Listen() error = %v", err)
87 }
203 - _ = conn.Close()
204 - t.Fatal("expected openReverseConnection to reject non-https relay")
205 -}
206 -
207 -func TestOpenReverseConnection_StopUnblocksTLSHandshake(t *testing.T) {
208 - t.Parallel()
88 + t.Cleanup(func() { _ = listener.Close() })
89
210 - ln, err := net.Listen("tcp", "127.0.0.1:0")
211 - if err != nil {
212 - t.Fatalf("listen: %v", err)
90 + if listener.Metadata().Description != "demo description" {
91 + t.Fatalf("listener.Metadata().Description = %q", listener.Metadata().Description)
92 }
214 - defer ln.Close()
215 -
216 - accepted := make(chan struct{}, 1)
217 - go func() {
218 - conn, acceptErr := ln.Accept()
219 - if acceptErr != nil {
220 - return
93 + if got, ok := relay.GetLease(listener.LeaseID()); !ok {
94 + t.Fatalf("GetLease(%q) = not found", listener.LeaseID())
95 + } else {
96 + if got.Metadata.Owner != "portal" {
97 + t.Fatalf("GetLease().Metadata.Owner = %q", got.Metadata.Owner)
98 }
222 - defer conn.Close()
223 - accepted <- struct{}{}
224 - buf := make([]byte, 1)
225 - _, _ = conn.Read(buf)
226 - }()
227 - l := &Listener{
228 - relayAddr: "https://" + ln.Addr().String(),
229 - lease: &types.Lease{ID: "lease-1", ReverseToken: "token-1"},
230 - reverseDialTimeout: 5 * time.Second,
231 - stopCh: make(chan struct{}),
232 - }
233 -
234 - done := make(chan error, 1)
235 - go func() {
236 - _, openErr := l.openReverseConnection()
237 - done <- openErr
238 - }()
239 -
240 - select {
241 - case <-accepted:
242 - case <-time.After(1 * time.Second):
243 - t.Fatal("timed out waiting for reverse dial accept")
244 - }
245 -
246 - close(l.stopCh)
247 -
248 - select {
249 - case openErr := <-done:
250 - if openErr == nil {
251 - t.Fatal("expected stop-aware openReverseConnection error")
99 + if len(got.Metadata.Tags) != 2 {
100 + t.Fatalf("GetLease().Metadata.Tags = %v, want deduped tags", got.Metadata.Tags)
101 }
253 - if !errors.Is(openErr, net.ErrClosed) {
254 - t.Fatalf("expected net.ErrClosed, got: %v", openErr)
102 + if !got.Metadata.Hide {
103 + t.Fatal("GetLease().Metadata.Hide = false, want true")
104 }
256 - case <-time.After(1 * time.Second):
257 - t.Fatal("openReverseConnection did not unblock after stop")
105 }
259 -}
260 -
261 -func TestWriteReverseConnectRequest_RespectsWriteDeadline(t *testing.T) {
262 - t.Parallel()
263 -
264 - local, peer := net.Pipe()
265 - defer local.Close()
266 - defer peer.Close()
106
268 - requestURL, err := url.Parse("https://relay.example.com" + types.PathSDKConnect + "?lease_id=lease-1")
269 - if err != nil {
270 - t.Fatalf("parse request URL: %v", err)
107 + httpDone := make(chan error, 1)
108 + server := &http.Server{
109 + Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
110 + w.Header().Set("Content-Type", "text/plain")
111 + _, _ = io.WriteString(w, "hello over relay\n")
112 + }),
113 + ReadHeaderTimeout: 5 * time.Second,
114 }
272 - l := &Listener{
273 - lease: &types.Lease{ReverseToken: "token-1"},
274 - reverseDialTimeout: 25 * time.Millisecond,
275 - stopCh: make(chan struct{}),
276 - }
277 -
278 - errCh := make(chan error, 1)
115 go func() {
280 - errCh <- l.writeReverseConnectRequest(local, requestURL)
116 + httpDone <- server.Serve(listener)
117 }()
282 -
283 - select {
284 - case writeErr := <-errCh:
285 - if writeErr == nil {
286 - t.Fatal("expected write deadline error")
287 - }
288 - var netErr net.Error
289 - if !errors.As(writeErr, &netErr) || !netErr.Timeout() {
290 - t.Fatalf("expected timeout error, got: %v", writeErr)
118 + t.Cleanup(func() {
119 + shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 5*time.Second)
120 + defer shutdownCancel()
121 + _ = server.Shutdown(shutdownCtx)
122 + select {
123 + case err := <-httpDone:
124 + if err != nil && err != http.ErrServerClosed && !errors.Is(err, net.ErrClosed) {
125 + t.Fatalf("server.Serve() error = %v", err)
126 + }
127 + case <-time.After(2 * time.Second):
128 }
292 - case <-time.After(500 * time.Millisecond):
293 - t.Fatal("timed out waiting for write result")
294 - }
295 -}
296 -
297 -func TestReadReverseConnectResponse_RespectsReadDeadline(t *testing.T) {
298 - t.Parallel()
299 -
300 - local, peer := net.Pipe()
301 - defer local.Close()
302 - defer peer.Close()
303 - l := &Listener{
304 - reverseDialTimeout: 25 * time.Millisecond,
305 - stopCh: make(chan struct{}),
306 - }
307 -
308 - errCh := make(chan error, 1)
309 - go func() {
310 - _, readErr := l.readReverseConnectResponse(local)
311 - errCh <- readErr
312 - }()
129 + })
130
314 - select {
315 - case readErr := <-errCh:
316 - if readErr == nil {
317 - t.Fatal("expected read deadline error")
131 + deadline := time.Now().Add(10 * time.Second)
132 + for {
133 + body, err := doTenantRequest(relay.SNIAddr(), tenantHost, "/")
134 + if err == nil {
135 + if !strings.Contains(body, "hello over relay") {
136 + t.Fatalf("body = %q, want relay payload", body)
137 + }
138 + return
139 }
319 - var netErr net.Error
320 - if !errors.As(readErr, &netErr) || !netErr.Timeout() {
321 - t.Fatalf("expected timeout error, got: %v", readErr)
140 + if time.Now().After(deadline) {
141 + t.Fatalf("doTenantRequest() last error = %v", err)
142 }
323 - case <-time.After(500 * time.Millisecond):
324 - t.Fatal("timed out waiting for read result")
143 + time.Sleep(100 * time.Millisecond)
144 }
145 }
146
328 -func TestParseReverseConnectRejection(t *testing.T) {
147 +func TestListenerEndToEndTLSHTTP_AutoSelfSigned(t *testing.T) {
148 t.Parallel()
149
331 - tests := []struct {
332 - name string
333 - body string
334 - code string
335 - want string
336 - }{
337 - {
338 - name: "envelope with code and message",
339 - body: `{"ok":false,"error":{"code":"ip_banned","message":"ip is banned"}}`,
340 - code: "ip_banned",
341 - want: "ip is banned (code=ip_banned)",
342 - },
343 - {
344 - name: "envelope with message only",
345 - body: `{"ok":false,"error":{"code":"","message":"missing lease_id"}}`,
346 - code: "",
347 - want: "missing lease_id",
348 - },
349 - {
350 - name: "plain text body",
351 - body: " unauthorized reverse connect ",
352 - code: "",
353 - want: "unauthorized reverse connect",
354 - },
355 - {
356 - name: "empty body",
357 - body: " ",
358 - code: "",
359 - want: "",
150 + apiCertPEM, apiKeyPEM, err := testutil.SelfSignedCertPEM("127.0.0.1")
151 + if err != nil {
152 + t.Fatalf("SelfSignedCertPEM(api) error = %v", err)
153 + }
154 + tenantHost := "auto.portal.test"
155 +
156 + relay, err := portal.NewServer(portal.ServerConfig{
157 + PortalURL: "https://127.0.0.1",
158 + APIListenAddr: "127.0.0.1:0",
159 + SNIListenAddr: "127.0.0.1:0",
160 + RootHost: "portal.test",
161 + RootFallbackAddr: "127.0.0.1:1",
162 + APITLS: portal.TLSMaterialConfig{
163 + CertPEM: apiCertPEM,
164 + KeyPEM: apiKeyPEM,
165 },
166 + })
167 + if err != nil {
168 + t.Fatalf("NewServer() error = %v", err)
169 + }
170 +
171 + ctx, cancel := context.WithCancel(context.Background())
172 + defer cancel()
173 + if err := relay.Start(ctx); err != nil {
174 + t.Fatalf("Start() error = %v", err)
175 + }
176 + t.Cleanup(func() {
177 + shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 5*time.Second)
178 + defer shutdownCancel()
179 + _ = relay.Shutdown(shutdownCtx)
180 + _ = relay.Wait()
181 + })
182 +
183 + client, err := NewClient(ClientConfig{
184 + RelayURL: "https://" + relay.APIAddr(),
185 + InsecureSkipVerify: true,
186 + ReadyTarget: 1,
187 + })
188 + if err != nil {
189 + t.Fatalf("NewClient() error = %v", err)
190 }
191 + t.Cleanup(client.Close)
192
363 - for _, tt := range tests {
364 - t.Run(tt.name, func(t *testing.T) {
365 - t.Parallel()
366 - code, detail := parseReverseConnectRejection([]byte(tt.body))
367 - if code != tt.code {
368 - t.Fatalf("parseReverseConnectRejection(%q) code=%q, want %q", tt.body, code, tt.code)
369 - }
370 - if detail != tt.want {
371 - t.Fatalf("parseReverseConnectRejection(%q) detail=%q, want %q", tt.body, detail, tt.want)
372 - }
373 - })
193 + listener, err := client.Listen(ctx, ListenRequest{
194 + Name: "auto-demo",
195 + Hostnames: []string{tenantHost},
196 + })
197 + if err != nil {
198 + t.Fatalf("Listen() error = %v", err)
199 }
375 -}
376 -
377 -func TestReadReverseConnectResponseParsesEnvelopeError(t *testing.T) {
378 - t.Parallel()
200 + t.Cleanup(func() { _ = listener.Close() })
201
380 - local, peer := net.Pipe()
381 - defer local.Close()
382 - defer peer.Close()
383 - l := &Listener{
384 - reverseDialTimeout: 500 * time.Millisecond,
385 - stopCh: make(chan struct{}),
202 + httpDone := make(chan error, 1)
203 + server := &http.Server{
204 + Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
205 + _, _ = io.WriteString(w, "auto tls ok\n")
206 + }),
207 + ReadHeaderTimeout: 5 * time.Second,
208 }
387 -
388 - errCh := make(chan error, 1)
209 go func() {
390 - _, err := l.readReverseConnectResponse(local)
391 - errCh <- err
210 + httpDone <- server.Serve(listener)
211 }()
393 -
394 - body := `{"ok":false,"error":{"code":"ip_banned","message":"ip is banned"}}`
395 - response := fmt.Sprintf(
396 - "HTTP/1.1 403 Forbidden\r\nContent-Type: application/json\r\nContent-Length: %d\r\n\r\n%s",
397 - len(body),
398 - body,
399 - )
400 - if _, err := io.WriteString(peer, response); err != nil {
401 - t.Fatalf("write response: %v", err)
402 - }
403 -
404 - select {
405 - case err := <-errCh:
406 - if err == nil {
407 - t.Fatal("expected readReverseConnectResponse to fail for 403 response")
408 - }
409 - var rejectionErr *reverseConnectRejectionError
410 - if !errors.As(err, &rejectionErr) {
411 - t.Fatalf("expected reverseConnectRejectionError, got: %T %v", err, err)
412 - }
413 - if rejectionErr.statusCode != http.StatusForbidden {
414 - t.Fatalf("rejection statusCode=%d, want %d", rejectionErr.statusCode, http.StatusForbidden)
415 - }
416 - if rejectionErr.code != "ip_banned" {
417 - t.Fatalf("rejection code=%q, want %q", rejectionErr.code, "ip_banned")
418 - }
419 - if rejectionErr.detail != "ip is banned (code=ip_banned)" {
420 - t.Fatalf("rejection detail=%q, want %q", rejectionErr.detail, "ip is banned (code=ip_banned)")
421 - }
422 - if !rejectionErr.IsFatal() {
423 - t.Fatalf("expected ip_banned rejection to be fatal: %+v", rejectionErr)
424 - }
425 - if strings.Contains(err.Error(), `{\"ok\":false`) {
426 - t.Fatalf("expected formatted error detail instead of raw JSON: %v", err)
212 + t.Cleanup(func() {
213 + shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 5*time.Second)
214 + defer shutdownCancel()
215 + _ = server.Shutdown(shutdownCtx)
216 + select {
217 + case err := <-httpDone:
218 + if err != nil && err != http.ErrServerClosed && !errors.Is(err, net.ErrClosed) {
219 + t.Fatalf("server.Serve() error = %v", err)
220 + }
221 + case <-time.After(2 * time.Second):
222 }
428 - case <-time.After(1 * time.Second):
429 - t.Fatal("timed out waiting for readReverseConnectResponse error")
430 - }
431 -}
223 + })
224
433 -func TestReverseConnectRejectionErrorIsFatal(t *testing.T) {
434 - t.Parallel()
435 -
436 - var nilErr *reverseConnectRejectionError
437 - if nilErr.IsFatal() {
438 - t.Fatal("nil rejection error must not be fatal")
439 - }
440 -
441 - tests := []struct {
442 - err *reverseConnectRejectionError
443 - name string
444 - want bool
445 - }{
446 - {
447 - name: "fatal by known code",
448 - err: &reverseConnectRejectionError{
449 - statusCode: http.StatusTooManyRequests,
450 - code: "ip_banned",
451 - },
452 - want: true,
453 - },
454 - {
455 - name: "fatal by status code",
456 - err: &reverseConnectRejectionError{
457 - statusCode: http.StatusUpgradeRequired,
458 - },
459 - want: true,
460 - },
461 - {
462 - name: "transient retry status",
463 - err: &reverseConnectRejectionError{
464 - statusCode: http.StatusServiceUnavailable,
465 - },
466 - want: false,
467 - },
468 - {
469 - name: "unknown code and status",
470 - err: &reverseConnectRejectionError{
471 - statusCode: http.StatusInternalServerError,
472 - code: "unexpected_failure",
473 - },
474 - want: false,
475 - },
476 - }
477 -
478 - for _, tt := range tests {
479 - t.Run(tt.name, func(t *testing.T) {
480 - t.Parallel()
481 - if got := tt.err.IsFatal(); got != tt.want {
482 - t.Fatalf("IsFatal()=%t, want %t for %+v", got, tt.want, tt.err)
225 + deadline := time.Now().Add(10 * time.Second)
226 + for {
227 + body, err := doTenantRequest(relay.SNIAddr(), tenantHost, "/")
228 + if err == nil {
229 + if !strings.Contains(body, "auto tls ok") {
230 + t.Fatalf("body = %q, want auto tls payload", body)
231 }
484 - })
485 - }
486 -}
487 -
488 -func TestWaitForReverseStart_HTTPMode(t *testing.T) {
489 - t.Parallel()
490 - l := &Listener{stopCh: make(chan struct{})}
491 - local, peer := net.Pipe()
492 - defer local.Close()
493 - defer peer.Close()
494 -
495 - done := make(chan error, 1)
496 - go func() {
497 - done <- l.waitForReverseStart(local, types.TLSStartMarker)
498 - }()
499 -
500 - _, err := peer.Write([]byte{types.TLSStartMarker})
501 - if err != nil {
502 - t.Fatalf("write marker: %v", err)
503 - }
504 -
505 - select {
506 - case err := <-done:
507 - if err != nil {
508 - t.Fatalf("waitForReverseStart failed: %v", err)
232 + return
233 }
510 - case <-time.After(500 * time.Millisecond):
511 - t.Fatal("timed out waiting for marker")
512 - }
513 -}
514 -
515 -func TestWaitForReverseStart_TLSMode(t *testing.T) {
516 - t.Parallel()
517 - l := &Listener{
518 - stopCh: make(chan struct{}),
519 - tlsConfig: &tls.Config{MinVersion: tls.VersionTLS12},
520 - }
521 - local, peer := net.Pipe()
522 - defer local.Close()
523 - defer peer.Close()
524 -
525 - done := make(chan error, 1)
526 - go func() {
527 - done <- l.waitForReverseStart(local, types.TLSStartMarker)
528 - }()
529 -
530 - _, err := peer.Write([]byte{types.TLSStartMarker})
531 - if err != nil {
532 - t.Fatalf("write marker: %v", err)
533 - }
534 -
535 - select {
536 - case err := <-done:
537 - if err != nil {
538 - t.Fatalf("waitForReverseStart failed: %v", err)
234 + if time.Now().After(deadline) {
235 + t.Fatalf("doTenantRequest() last error = %v", err)
236 }
540 - case <-time.After(500 * time.Millisecond):
541 - t.Fatal("timed out waiting for marker")
237 + time.Sleep(100 * time.Millisecond)
238 }
239 }
240
545 -func TestWaitForReverseStart_IgnoresKeepaliveMarker(t *testing.T) {
546 - t.Parallel()
547 - l := &Listener{stopCh: make(chan struct{})}
548 - local, peer := net.Pipe()
549 - defer local.Close()
550 - defer peer.Close()
551 -
552 - done := make(chan error, 1)
553 - go func() {
554 - done <- l.waitForReverseStart(local, types.TLSStartMarker)
555 - }()
556 -
557 - _, err := peer.Write([]byte{types.ReverseKeepaliveMarker})
241 +func doTenantRequest(addr, host, path string) (string, error) {
242 + conn, err := tls.Dial("tcp", addr, &tls.Config{
243 + ServerName: host,
244 + InsecureSkipVerify: true,
245 + NextProtos: []string{"http/1.1"},
246 + })
247 if err != nil {
559 - t.Fatalf("write keepalive marker: %v", err)
560 - }
561 - _, err = peer.Write([]byte{types.TLSStartMarker})
562 - if err != nil {
563 - t.Fatalf("write start marker: %v", err)
564 - }
565 -
566 - select {
567 - case err := <-done:
568 - if err != nil {
569 - t.Fatalf("waitForReverseStart failed: %v", err)
570 - }
571 - case <-time.After(500 * time.Millisecond):
572 - t.Fatal("timed out waiting for marker")
248 + return "", err
249 }
574 -}
250 + defer conn.Close()
251
576 -func TestWaitForReverseStart_TLSRejectsHTTPMarker(t *testing.T) {
577 - t.Parallel()
578 - l := &Listener{
579 - stopCh: make(chan struct{}),
580 - tlsConfig: &tls.Config{MinVersion: tls.VersionTLS12},
252 + if _, err := fmt.Fprintf(conn, "GET %s HTTP/1.1\r\nHost: %s\r\nConnection: close\r\n\r\n", path, host); err != nil {
253 + return "", err
254 }
582 - local, peer := net.Pipe()
583 - defer local.Close()
584 - defer peer.Close()
255
586 - done := make(chan error, 1)
587 - go func() {
588 - done <- l.waitForReverseStart(local, types.TLSStartMarker)
589 - }()
590 -
591 - _, err := peer.Write([]byte{types.NonTLSStartMarker})
256 + resp, err := http.ReadResponse(bufio.NewReader(conn), &http.Request{Method: http.MethodGet})
257 if err != nil {
593 - t.Fatalf("write marker: %v", err)
594 - }
595 -
596 - select {
597 - case err := <-done:
598 - if err == nil {
599 - t.Fatal("expected invalid marker error")
600 - }
601 - case <-time.After(500 * time.Millisecond):
602 - t.Fatal("timed out waiting for marker")
258 + return "", err
259 }
604 -}
260 + defer resp.Body.Close()
261
606 -func TestWaitForReverseStart_HTTPRejectsTLSMarker(t *testing.T) {
607 - t.Parallel()
608 - l := &Listener{stopCh: make(chan struct{})}
609 - local, peer := net.Pipe()
610 - defer local.Close()
611 - defer peer.Close()
612 -
613 - done := make(chan error, 1)
614 - go func() {
615 - done <- l.waitForReverseStart(local, types.NonTLSStartMarker)
616 - }()
617 -
618 - _, err := peer.Write([]byte{types.TLSStartMarker})
262 + body, err := io.ReadAll(resp.Body)
263 if err != nil {
620 - t.Fatalf("write marker: %v", err)
621 - }
622 -
623 - select {
624 - case err := <-done:
625 - if err == nil {
626 - t.Fatal("expected invalid marker error")
627 - }
628 - case <-time.After(500 * time.Millisecond):
629 - t.Fatal("timed out waiting for marker")
264 + return "", err
265 }
631 -}
632 -
633 -func TestWaitForReverseStart_StopCancelsWait(t *testing.T) {
634 - t.Parallel()
635 - l := &Listener{stopCh: make(chan struct{})}
636 - local, peer := net.Pipe()
637 - defer local.Close()
638 -
639 - done := make(chan error, 1)
640 - go func() {
641 - done <- l.waitForReverseStart(local, types.TLSStartMarker)
642 - }()
643 -
644 - close(l.stopCh)
645 - _ = peer.Close()
646 -
647 - select {
648 - case err := <-done:
649 - if !errors.Is(err, net.ErrClosed) {
650 - t.Fatalf("expected net.ErrClosed when listener stops, got: %v", err)
651 - }
652 - case <-time.After(500 * time.Millisecond):
653 - t.Fatal("waitForReverseStart did not stop after cancellation")
266 + if resp.StatusCode != http.StatusOK {
267 + return "", fmt.Errorf("status %d: %s", resp.StatusCode, string(body))
268 }
269 + return string(body), nil
270 }
sdk/tls.go new
+138
@@ -0,0 +1,138 @@
1 +package sdk
2 +
3 +import (
4 + "crypto/ecdsa"
5 + "crypto/elliptic"
6 + "crypto/rand"
7 + "crypto/tls"
8 + "crypto/x509"
9 + "crypto/x509/pkix"
10 + "encoding/pem"
11 + "fmt"
12 + "io"
13 + "math/big"
14 + "net"
15 + "strings"
16 + "time"
17 +
18 + keylesslib "github.com/gosuda/keyless_tls/keyless"
19 +
20 + "gosuda.org/portal/portal"
21 +)
22 +
23 +func buildTenantTLSConfig(cfg portal.TLSMaterialConfig) (*tls.Config, io.Closer, error) {
24 + if len(cfg.CertPEM) == 0 {
25 + return nil, nil, fmt.Errorf("tenant certificate is required")
26 + }
27 + if cfg.Keyless != nil {
28 + remoteSigner, err := keylesslib.NewRemoteSigner(keylesslib.RemoteSignerConfig{
29 + Endpoint: cfg.Keyless.Endpoint,
30 + ServerName: cfg.Keyless.ServerName,
31 + KeyID: cfg.Keyless.KeyID,
32 + ClientCertPEM: cfg.Keyless.ClientCertPEM,
33 + ClientKeyPEM: cfg.Keyless.ClientKeyPEM,
34 + RootCAPEM: cfg.Keyless.RootCAPEM,
35 + }, cfg.CertPEM)
36 + if err != nil {
37 + return nil, nil, err
38 + }
39 + tlsConf, err := keylesslib.NewServerTLSConfig(keylesslib.ServerTLSConfig{
40 + CertPEM: cfg.CertPEM,
41 + Signer: remoteSigner,
42 + NextProtos: []string{"http/1.1"},
43 + MinVersion: tls.VersionTLS12,
44 + })
45 + if err != nil {
46 + _ = remoteSigner.Close()
47 + return nil, nil, err
48 + }
49 + return tlsConf, remoteSigner, nil
50 + }
51 +
52 + cert, err := tls.X509KeyPair(cfg.CertPEM, cfg.KeyPEM)
53 + if err != nil {
54 + return nil, nil, fmt.Errorf("parse tenant tls key pair: %w", err)
55 + }
56 +
57 + return &tls.Config{
58 + MinVersion: tls.VersionTLS12,
59 + NextProtos: []string{"http/1.1"},
60 + Certificates: []tls.Certificate{cert},
61 + }, nil, nil
62 +}
63 +
64 +func buildAutoTenantTLSConfig(hostnames []string) (*tls.Config, error) {
65 + certPEM, keyPEM, err := selfSignedTenantCert(hostnames)
66 + if err != nil {
67 + return nil, err
68 + }
69 + tlsConf, _, err := buildTenantTLSConfig(portal.TLSMaterialConfig{
70 + CertPEM: certPEM,
71 + KeyPEM: keyPEM,
72 + })
73 + return tlsConf, err
74 +}
75 +
76 +func selfSignedTenantCert(hostnames []string) (certPEM, keyPEM []byte, err error) {
77 + priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
78 + if err != nil {
79 + return nil, nil, fmt.Errorf("generate self-signed tenant key: %w", err)
80 + }
81 +
82 + serialLimit := new(big.Int).Lsh(big.NewInt(1), 128)
83 + serial, err := rand.Int(rand.Reader, serialLimit)
84 + if err != nil {
85 + return nil, nil, fmt.Errorf("generate self-signed tenant serial: %w", err)
86 + }
87 +
88 + template := &x509.Certificate{
89 + SerialNumber: serial,
90 + Subject: pkix.Name{
91 + CommonName: firstHostname(hostnames),
92 + },
93 + NotBefore: time.Now().Add(-1 * time.Hour),
94 + NotAfter: time.Now().Add(30 * 24 * time.Hour),
95 + KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment,
96 + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
97 + BasicConstraintsValid: true,
98 + }
99 +
100 + for _, host := range hostnames {
101 + host = strings.TrimSpace(host)
102 + if host == "" {
103 + continue
104 + }
105 + if ip := net.ParseIP(host); ip != nil {
106 + template.IPAddresses = append(template.IPAddresses, ip)
107 + continue
108 + }
109 + template.DNSNames = append(template.DNSNames, host)
110 + }
111 + if len(template.DNSNames) == 0 && len(template.IPAddresses) == 0 {
112 + template.DNSNames = []string{"localhost"}
113 + template.Subject.CommonName = "localhost"
114 + }
115 +
116 + der, err := x509.CreateCertificate(rand.Reader, template, template, &priv.PublicKey, priv)
117 + if err != nil {
118 + return nil, nil, fmt.Errorf("create self-signed tenant certificate: %w", err)
119 + }
120 +
121 + certPEM = pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})
122 + keyDER, err := x509.MarshalECPrivateKey(priv)
123 + if err != nil {
124 + return nil, nil, fmt.Errorf("marshal self-signed tenant key: %w", err)
125 + }
126 + keyPEM = pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER})
127 + return certPEM, keyPEM, nil
128 +}
129 +
130 +func firstHostname(hostnames []string) string {
131 + for _, host := range hostnames {
132 + host = strings.TrimSpace(host)
133 + if host != "" {
134 + return host
135 + }
136 + }
137 + return "localhost"
138 +}
types/api.go deleted
-153
@@ -1,153 +0,0 @@
1 -package types
2 -
3 -import "encoding/json"
4 -
5 -// API path constants for Portal relay server.
6 -
7 -// SDK API paths for lease registration and tunnel connections.
8 -const (
9 - PathSDKPrefix = "/sdk/"
10 - PathSDKRegister = "/sdk/register"
11 - PathSDKUnregister = "/sdk/unregister"
12 - PathSDKRenew = "/sdk/renew"
13 - PathSDKDomain = "/sdk/domain"
14 - PathSDKConnect = "/sdk/connect"
15 -
16 - // Admin API paths.
17 - PathAdminPrefix = "/admin"
18 - PathAdminLogin = "/admin/login"
19 - PathAdminLogout = "/admin/logout"
20 - PathAdminAuthStatus = "/admin/auth/status"
21 - PathAdminLeases = "/admin/leases"
22 - PathAdminLeasesBanned = "/admin/leases/banned"
23 - PathAdminStats = "/admin/stats"
24 - PathAdminSettings = "/admin/settings"
25 - PathAdminApprovalMode = "/admin/settings/approval-mode"
26 -
27 - // Keyless API paths.
28 - PathKeylessSign = "/v1/sign"
29 -
30 - // Health check path.
31 - PathHealthz = "/healthz"
32 -
33 - // Tunnel installer paths.
34 - PathTunnelScript = "/tunnel"
35 - PathTunnelBinary = "/tunnel/bin/"
36 -
37 - // App static assets path prefix.
38 - PathAppPrefix = "/app/"
39 -)
40 -
41 -// Client API types for /sdk/* endpoints.
42 -
43 -// APIError is the normalized API error payload.
44 -type APIError struct {
45 - Code string `json:"code"`
46 - Message string `json:"message"`
47 - StatusCode int `json:"-"` // HTTP status code, not serialized
48 -}
49 -
50 -// APIEnvelope is the canonical response wrapper for relay APIs.
51 -type APIEnvelope struct {
52 - Data any `json:"data,omitempty"`
53 - Error *APIError `json:"error,omitempty"`
54 - OK bool `json:"ok"`
55 -}
56 -
57 -// APIRawEnvelope is the decoding-friendly envelope with raw data payload.
58 -type APIRawEnvelope struct {
59 - Error *APIError `json:"error,omitempty"`
60 - Data json.RawMessage `json:"data,omitempty"`
61 - OK bool `json:"ok"`
62 -}
63 -
64 -// RegisterRequest is the lease registration request.
65 -type RegisterRequest struct {
66 - LeaseID string `json:"lease_id"`
67 - Name string `json:"name"`
68 - ReverseToken string `json:"reverse_token"`
69 - Metadata Metadata `json:"metadata"`
70 - TLS bool `json:"tls"`
71 -}
72 -
73 -// RegisterResponse is the lease registration response.
74 -type RegisterResponse struct {
75 - Message string `json:"message,omitempty"`
76 - LeaseID string `json:"lease_id,omitempty"`
77 - PublicURL string `json:"public_url,omitempty"`
78 - Success bool `json:"success"`
79 -}
80 -
81 -// UnregisterRequest is the lease unregistration request.
82 -type UnregisterRequest struct {
83 - LeaseID string `json:"lease_id"`
84 - ReverseToken string `json:"reverse_token"`
85 -}
86 -
87 -// RenewRequest is the lease renewal request.
88 -type RenewRequest struct {
89 - LeaseID string `json:"lease_id"`
90 - ReverseToken string `json:"reverse_token"`
91 -}
92 -
93 -// APIResponse is a generic Client API response.
94 -type APIResponse struct {
95 - Message string `json:"message,omitempty"`
96 - Success bool `json:"success"`
97 -}
98 -
99 -// DomainResponse is the Client domain discovery response.
100 -type DomainResponse struct {
101 - Message string `json:"message,omitempty"`
102 - BaseDomain string `json:"base_domain,omitempty"`
103 - Success bool `json:"success"`
104 -}
105 -
106 -// Admin API types for /admin/* endpoints.
107 -
108 -// AdminLoginRequest is the admin login request body.
109 -type AdminLoginRequest struct {
110 - Key string `json:"key"`
111 -}
112 -
113 -// AdminLoginResponse is the admin login response.
114 -type AdminLoginResponse struct {
115 - Error string `json:"error,omitempty"`
116 - RemainingSeconds int `json:"remaining_seconds,omitempty"`
117 - Success bool `json:"success"`
118 - Locked bool `json:"locked,omitempty"`
119 -}
120 -
121 -// AdminAuthStatusResponse is the admin auth status response.
122 -type AdminAuthStatusResponse struct {
123 - Authenticated bool `json:"authenticated"`
124 - AuthEnabled bool `json:"auth_enabled"`
125 -}
126 -
127 -// AdminSettingsResponse is the admin settings response.
128 -type AdminSettingsResponse struct {
129 - ApprovalMode string `json:"approval_mode"`
130 - ApprovedLeases []string `json:"approved_leases"`
131 - DeniedLeases []string `json:"denied_leases"`
132 -}
133 -
134 -// AdminApprovalModeRequest is the request to change approval mode.
135 -type AdminApprovalModeRequest struct {
136 - Mode string `json:"mode"`
137 -}
138 -
139 -// AdminApprovalModeResponse is the approval mode response.
140 -type AdminApprovalModeResponse struct {
141 - ApprovalMode string `json:"approval_mode"`
142 -}
143 -
144 -// AdminBPSRequest is the request to set BPS limit for a lease.
145 -type AdminBPSRequest struct {
146 - BPS int64 `json:"bps"`
147 -}
148 -
149 -// AdminStatsResponse is the admin stats response.
150 -type AdminStatsResponse struct {
151 - Uptime string `json:"uptime"`
152 - LeasesCount int `json:"leases_count"`
153 -}
types/netutil.go deleted
-452
@@ -1,452 +0,0 @@
1 -package types
2 -
3 -import (
4 - "errors"
5 - "fmt"
6 - "net"
7 - "net/url"
8 - "strconv"
9 - "strings"
10 -)
11 -
12 -const defaultBootstrapURL = "https://localhost:4017"
13 -
14 -func normalizeRootHost(raw string) string {
15 - normalized := strings.ToLower(strings.TrimSpace(raw))
16 - normalized = strings.TrimPrefix(strings.TrimSuffix(normalized, "."), "*.")
17 - return normalized
18 -}
19 -
20 -// IsLocalhost reports whether host resolves to localhost/loopback semantics.
21 -// It accepts bare hosts, host:port forms, and bracketed IPv6 literals.
22 -func IsLocalhost(host string) bool {
23 - normalized := strings.ToLower(strings.TrimSpace(host))
24 - if normalized == "" {
25 - return false
26 - }
27 -
28 - if parsedHost, _, err := net.SplitHostPort(normalized); err == nil {
29 - normalized = parsedHost
30 - }
31 - normalized = strings.TrimPrefix(strings.TrimSuffix(normalized, "."), "*.")
32 - normalized = strings.TrimPrefix(strings.TrimSuffix(normalized, "]"), "[")
33 -
34 - if normalized == "localhost" || strings.HasSuffix(normalized, ".localhost") {
35 - return true
36 - }
37 -
38 - if ip := net.ParseIP(normalized); ip != nil {
39 - return ip.IsLoopback()
40 - }
41 - return false
42 -}
43 -
44 -// StripScheme removes http:// or https:// prefix from a string.
45 -func StripScheme(s string) string {
46 - s = strings.TrimSpace(s)
47 - s = strings.TrimSuffix(s, "/")
48 - s = strings.TrimPrefix(s, "http://")
49 - s = strings.TrimPrefix(s, "https://")
50 - return s
51 -}
52 -
53 -// StripWildcard removes *. prefix from a domain pattern.
54 -func StripWildcard(s string) string {
55 - s = strings.TrimSpace(s)
56 - s = strings.TrimPrefix(s, "*.")
57 - return s
58 -}
59 -
60 -// StripPort removes a trailing :port from a host string if present.
61 -func StripPort(s string) string {
62 - if s == "" {
63 - return s
64 - }
65 - if idx := strings.LastIndexByte(s, ':'); idx >= 0 && idx+1 < len(s) {
66 - port := s[idx+1:]
67 - digits := true
68 - for _, ch := range port {
69 - if ch < '0' || ch > '9' {
70 - digits = false
71 - break
72 - }
73 - }
74 - if digits {
75 - return s[:idx]
76 - }
77 - }
78 - return s
79 -}
80 -
81 -// IsSubdomain reports whether host matches the given domain pattern.
82 -// Pattern can be a wildcard like "*.example.com" or exact domain.
83 -func IsSubdomain(domain, host string) bool {
84 - if host == "" || domain == "" {
85 - return false
86 - }
87 -
88 - h := strings.ToLower(StripPort(StripScheme(host)))
89 - d := strings.ToLower(StripPort(StripScheme(domain)))
90 -
91 - if strings.HasPrefix(d, "*.") {
92 - suffix := d[1:]
93 - return len(h) > len(suffix) && strings.HasSuffix(h, suffix)
94 - }
95 -
96 - if h == d {
97 - return true
98 - }
99 -
100 - return strings.HasSuffix(h, "."+d)
101 -}
102 -
103 -func parsePortalAddress(raw, fallbackScheme string) (scheme, rootHost, hostPort string, ok bool) {
104 - normalized := strings.TrimSpace(raw)
105 - if normalized == "" {
106 - return "", "", "", false
107 - }
108 -
109 - if fallbackScheme == "" {
110 - fallbackScheme = "https"
111 - }
112 - fallbackScheme = strings.ToLower(strings.TrimSpace(fallbackScheme))
113 -
114 - if !strings.Contains(normalized, "://") {
115 - normalized = fallbackScheme + "://" + normalized
116 - }
117 -
118 - parsed, err := url.Parse(normalized)
119 - if err != nil || parsed.Hostname() == "" {
120 - return "", "", "", false
121 - }
122 -
123 - rootHost = normalizeRootHost(parsed.Hostname())
124 - if rootHost == "" {
125 - return "", "", "", false
126 - }
127 -
128 - scheme = strings.ToLower(strings.TrimSpace(parsed.Scheme))
129 - if scheme == "" {
130 - scheme = fallbackScheme
131 - }
132 -
133 - if port := strings.TrimSpace(parsed.Port()); port != "" {
134 - hostPort = net.JoinHostPort(rootHost, port)
135 - } else {
136 - hostPort = rootHost
137 - }
138 -
139 - return scheme, rootHost, strings.ToLower(strings.TrimSpace(hostPort)), true
140 -}
141 -
142 -// NormalizeServiceName canonicalizes and validates a service/lease name for DNS usage.
143 -// Valid names are a single DNS label: [a-z0-9]([a-z0-9-]{0,61}[a-z0-9])?
144 -func NormalizeServiceName(name string) (string, bool) {
145 - normalized := strings.ToLower(strings.TrimSpace(name))
146 - normalized = strings.TrimPrefix(normalized, "*.")
147 - normalized = strings.TrimSuffix(normalized, ".")
148 -
149 - if normalized == "" || len(normalized) > 63 {
150 - return "", false
151 - }
152 - if strings.Contains(normalized, ".") || normalized[0] == '-' || normalized[len(normalized)-1] == '-' {
153 - return "", false
154 - }
155 - for _, ch := range normalized {
156 - switch {
157 - case ch >= 'a' && ch <= 'z':
158 - case ch >= '0' && ch <= '9':
159 - case ch == '-':
160 - default:
161 - return "", false
162 - }
163 - }
164 - return normalized, true
165 -}
166 -
167 -// IsValidServiceName reports whether a service name can be used as a DNS label.
168 -func IsValidServiceName(name string) bool {
169 - _, ok := NormalizeServiceName(name)
170 - return ok
171 -}
172 -
173 -// DefaultAppPattern builds a wildcard subdomain pattern from a base portal URL or host.
174 -func DefaultAppPattern(base string) string {
175 - if strings.TrimSpace(base) == "" {
176 - return "*.localhost:4017"
177 - }
178 - _, _, hostPort, ok := parsePortalAddress(base, "https")
179 - if !ok || hostPort == "" {
180 - return "*.localhost:4017"
181 - }
182 - return "*." + hostPort
183 -}
184 -
185 -// ServicePublicURL returns a service URL derived from portalURL and service name.
186 -func ServicePublicURL(portalURL, serviceName string) string {
187 - normalizedName, ok := NormalizeServiceName(serviceName)
188 - if !ok {
189 - return ""
190 - }
191 -
192 - scheme, rootHost, _, ok := parsePortalAddress(portalURL, "https")
193 - if !ok || rootHost == "" {
194 - return ""
195 - }
196 - return fmt.Sprintf("%s://%s.%s", scheme, normalizedName, rootHost)
197 -}
198 -
199 -// PortalHostPort returns normalized host[:port] from a portal URL-like input.
200 -func PortalHostPort(portalURL string) string {
201 - _, _, hostPort, ok := parsePortalAddress(portalURL, "https")
202 - if !ok {
203 - return ""
204 - }
205 - return hostPort
206 -}
207 -
208 -// PortalRootHost extracts the root hostname from a portal URL.
209 -func PortalRootHost(portalURL string) string {
210 - _, rootHost, _, ok := parsePortalAddress(portalURL, "https")
211 - if !ok {
212 - return ""
213 - }
214 - return rootHost
215 -}
216 -
217 -// DefaultBootstrapFrom derives a relay API bootstrap URL from a base portal URL or host.
218 -func DefaultBootstrapFrom(base string) string {
219 - base = strings.TrimSpace(base)
220 - if base == "" {
221 - return defaultBootstrapURL
222 - }
223 -
224 - if !strings.Contains(base, "://") {
225 - base = "https://" + base
226 - }
227 -
228 - u, err := url.Parse(strings.TrimSuffix(base, "/"))
229 - if err != nil || u.Host == "" {
230 - return defaultBootstrapURL
231 - }
232 - if u.Scheme != "http" && u.Scheme != "https" {
233 - return defaultBootstrapURL
234 - }
235 - if p := strings.TrimSpace(u.Path); p != "" && p != "/" {
236 - return defaultBootstrapURL
237 - }
238 -
239 - u.Scheme = "https"
240 - u.Path = ""
241 - u.RawQuery = ""
242 - u.Fragment = ""
243 - return strings.TrimSuffix(u.String(), "/")
244 -}
245 -
246 -// LeaseNameFromHost extracts the lease name from a subdomain host.
247 -func LeaseNameFromHost(host, appURL string) (string, bool) {
248 - if !IsSubdomain(appURL, host) {
249 - return "", false
250 - }
251 -
252 - normalizedHost := strings.ToLower(strings.TrimSpace(StripPort(host)))
253 - baseHost := strings.ToLower(strings.TrimSpace(
254 - StripPort(StripWildcard(StripScheme(appURL))),
255 - ))
256 -
257 - if normalizedHost == "" || baseHost == "" || normalizedHost == baseHost {
258 - return "", false
259 - }
260 -
261 - suffix := "." + baseHost
262 - if !strings.HasSuffix(normalizedHost, suffix) {
263 - return "", false
264 - }
265 -
266 - leaseName := strings.TrimSuffix(normalizedHost, suffix)
267 - normalizedLeaseName, ok := NormalizeServiceName(leaseName)
268 - if !ok {
269 - return "", false
270 - }
271 -
272 - return normalizedLeaseName, true
273 -}
274 -
275 -// BuildSNIName constructs the SNI hostname for a lease.
276 -func BuildSNIName(leaseName, baseHost string) string {
277 - normalizedLeaseName, ok := NormalizeServiceName(leaseName)
278 - if !ok {
279 - return ""
280 - }
281 -
282 - normalizedBaseHost := PortalRootHost(baseHost)
283 - if normalizedBaseHost == "" {
284 - normalizedBaseHost = normalizeRootHost(baseHost)
285 - }
286 - if normalizedBaseHost == "" {
287 - return ""
288 - }
289 -
290 - return normalizedLeaseName + "." + normalizedBaseHost
291 -}
292 -
293 -// ParseURLs splits a comma-separated string into a list of trimmed, non-empty URLs.
294 -func ParseURLs(raw string) []string {
295 - raw = strings.TrimSpace(raw)
296 - if raw == "" {
297 - return nil
298 - }
299 - parts := strings.Split(raw, ",")
300 - out := make([]string, 0, len(parts))
301 - for _, p := range parts {
302 - p = strings.TrimSpace(p)
303 - if p != "" {
304 - out = append(out, p)
305 - }
306 - }
307 - return out
308 -}
309 -
310 -// ParsePortNumber parses a port number from a string, returning fallback on error.
311 -// The raw value may optionally start with a colon prefix.
312 -func ParsePortNumber(raw string, fallback int) int {
313 - value := strings.TrimSpace(raw)
314 - if value == "" {
315 - return fallback
316 - }
317 - value = strings.TrimPrefix(value, ":")
318 - port, err := strconv.Atoi(value)
319 - if err != nil || port < 1 || port > 65535 {
320 - return fallback
321 - }
322 - return port
323 -}
324 -
325 -// LoopbackForwardAddr converts a listen address to a loopback forward address.
326 -// For example, ":4017" becomes "127.0.0.1:4017".
327 -func LoopbackForwardAddr(listenAddr string) string {
328 - raw := strings.TrimSpace(listenAddr)
329 - if raw == "" {
330 - return ""
331 - }
332 -
333 - var port string
334 - switch {
335 - case strings.HasPrefix(raw, ":"):
336 - port = strings.TrimPrefix(raw, ":")
337 - case strings.Count(raw, ":") == 0:
338 - port = raw
339 - default:
340 - _, p, err := net.SplitHostPort(raw)
341 - if err != nil {
342 - return ""
343 - }
344 - port = p
345 - }
346 -
347 - portNum, err := strconv.Atoi(port)
348 - if err != nil || portNum < 1 || portNum > 65535 {
349 - return ""
350 - }
351 -
352 - return net.JoinHostPort("127.0.0.1", strconv.Itoa(portNum))
353 -}
354 -
355 -// NormalizeTargetAddr normalizes a target address for dialing.
356 -// If the input is a URL (e.g., "http://localhost:8080"), it extracts the host.
357 -// Otherwise, it returns the trimmed input.
358 -func NormalizeTargetAddr(raw string) (string, error) {
359 - raw = strings.TrimSpace(raw)
360 - if raw == "" {
361 - return "", errors.New("empty host")
362 - }
363 -
364 - // Treat plain host[:port] as a dial target without URL parsing.
365 - if !strings.Contains(raw, "://") {
366 - return raw, nil
367 - }
368 -
369 - u, err := url.Parse(raw)
370 - if err != nil {
371 - return "", fmt.Errorf("parse target URL: %w", err)
372 - }
373 - if strings.TrimSpace(u.Host) == "" {
374 - return "", errors.New("missing host in URL")
375 - }
376 - return u.Host, nil
377 -}
378 -
379 -// NormalizeRelayAPIURL normalizes a relay API URL.
380 -// It accepts host:port input (defaults to https), validates the scheme,
381 -// normalizes localhost hostnames, and removes path/query/fragment.
382 -func NormalizeRelayAPIURL(raw string) (string, error) {
383 - raw = strings.TrimSpace(raw)
384 - if raw == "" {
385 - return "", errors.New("empty relay URL")
386 - }
387 -
388 - // Accept host:port input.
389 - if !strings.Contains(raw, "://") {
390 - raw = "https://" + raw
391 - }
392 -
393 - u, err := url.Parse(raw)
394 - if err != nil {
395 - return "", fmt.Errorf("parse relay URL: %w", err)
396 - }
397 - if u.Host == "" {
398 - return "", fmt.Errorf("relay URL missing host: %q", raw)
399 - }
400 -
401 - if host := strings.ToLower(strings.TrimSpace(u.Hostname())); strings.HasSuffix(host, ".localhost") {
402 - port := u.Port()
403 - if port != "" {
404 - u.Host = net.JoinHostPort("localhost", port)
405 - } else {
406 - u.Host = "localhost"
407 - }
408 - }
409 -
410 - switch u.Scheme {
411 - case "https":
412 - default:
413 - return "", fmt.Errorf("unsupported relay URL scheme: %q (use https)", u.Scheme)
414 - }
415 -
416 - if p := strings.TrimSpace(u.Path); p != "" && p != "/" {
417 - return "", fmt.Errorf("relay URL must not include path: %q", raw)
418 - }
419 -
420 - u.RawQuery = ""
421 - u.Fragment = ""
422 - u.Path = ""
423 -
424 - return strings.TrimSuffix(u.String(), "/"), nil
425 -}
426 -
427 -// NormalizeRelayAPIURLs normalizes a list of relay API URLs, deduplicating results.
428 -// Returns an error if no valid URLs remain after normalization.
429 -func NormalizeRelayAPIURLs(bootstrapServers []string) ([]string, error) {
430 - if len(bootstrapServers) == 0 {
431 - return nil, errors.New("no available relay")
432 - }
433 -
434 - seen := make(map[string]struct{}, len(bootstrapServers))
435 - out := make([]string, 0, len(bootstrapServers))
436 - for _, relay := range bootstrapServers {
437 - normalized, err := NormalizeRelayAPIURL(relay)
438 - if err != nil {
439 - continue
440 - }
441 - if _, exists := seen[normalized]; exists {
442 - continue
443 - }
444 - seen[normalized] = struct{}{}
445 - out = append(out, normalized)
446 - }
447 -
448 - if len(out) == 0 {
449 - return nil, errors.New("no available relay")
450 - }
451 - return out, nil
452 -}
types/netutil_test.go deleted
-309
@@ -1,309 +0,0 @@
1 -package types
2 -
3 -import (
4 - "strings"
5 - "testing"
6 -)
7 -
8 -func TestNormalizeTargetAddr(t *testing.T) {
9 - t.Parallel()
10 -
11 - tests := []struct {
12 - name string
13 - in string
14 - want string
15 - wantErr bool
16 - }{
17 - {
18 - name: "host and port",
19 - in: "localhost:3000",
20 - want: "localhost:3000",
21 - },
22 - {
23 - name: "url with scheme",
24 - in: "http://localhost:3000",
25 - want: "localhost:3000",
26 - },
27 - {
28 - name: "url missing host",
29 - in: "http:///only-path",
30 - wantErr: true,
31 - },
32 - {
33 - name: "empty",
34 - in: " ",
35 - wantErr: true,
36 - },
37 - }
38 -
39 - for _, tt := range tests {
40 - t.Run(tt.name, func(t *testing.T) {
41 - t.Parallel()
42 -
43 - got, err := NormalizeTargetAddr(tt.in)
44 - if tt.wantErr {
45 - if err == nil {
46 - t.Fatalf("expected error for input %q", tt.in)
47 - }
48 - return
49 - }
50 -
51 - if err != nil {
52 - t.Fatalf("unexpected error for input %q: %v", tt.in, err)
53 - }
54 -
55 - if got != tt.want {
56 - t.Fatalf("NormalizeTargetAddr(%q) = %q, want %q", tt.in, got, tt.want)
57 - }
58 - })
59 - }
60 -}
61 -
62 -func TestNormalizeServiceName(t *testing.T) {
63 - t.Parallel()
64 -
65 - tests := []struct {
66 - name string
67 - in string
68 - want string
69 - ok bool
70 - }{
71 - {name: "simple", in: "my-app", want: "my-app", ok: true},
72 - {name: "trim and lowercase", in: " My-App ", want: "my-app", ok: true},
73 - {name: "strip wildcard and trailing dot", in: "*.Service.", want: "service", ok: true},
74 - {name: "empty", in: " ", ok: false},
75 - {name: "contains dot", in: "api.v1", ok: false},
76 - {name: "contains underscore", in: "api_v1", ok: false},
77 - {name: "leading hyphen", in: "-api", ok: false},
78 - {name: "trailing hyphen", in: "api-", ok: false},
79 - {name: "too long", in: strings.Repeat("a", 64), ok: false},
80 - {name: "contains spaces", in: "api v1", ok: false},
81 - }
82 -
83 - for _, tt := range tests {
84 - t.Run(tt.name, func(t *testing.T) {
85 - t.Parallel()
86 -
87 - got, ok := NormalizeServiceName(tt.in)
88 - if ok != tt.ok {
89 - t.Fatalf("NormalizeServiceName(%q) ok=%v, want %v", tt.in, ok, tt.ok)
90 - }
91 - if got != tt.want {
92 - t.Fatalf("NormalizeServiceName(%q)=%q, want %q", tt.in, got, tt.want)
93 - }
94 - })
95 - }
96 -}
97 -
98 -func TestPortalHostDerivationConsistency(t *testing.T) {
99 - t.Parallel()
100 -
101 - portalURL := "https://relay.edge.example.com:8443/path"
102 -
103 - if got := PortalRootHost(portalURL); got != "relay.edge.example.com" {
104 - t.Fatalf("PortalRootHost(%q)=%q, want %q", portalURL, got, "relay.edge.example.com")
105 - }
106 - if got := PortalHostPort(portalURL); got != "relay.edge.example.com:8443" {
107 - t.Fatalf("PortalHostPort(%q)=%q, want %q", portalURL, got, "relay.edge.example.com:8443")
108 - }
109 - if got := DefaultAppPattern(portalURL); got != "*.relay.edge.example.com:8443" {
110 - t.Fatalf("DefaultAppPattern(%q)=%q, want %q", portalURL, got, "*.relay.edge.example.com:8443")
111 - }
112 - if got := BuildSNIName("Api-Gateway", portalURL); got != "api-gateway.relay.edge.example.com" {
113 - t.Fatalf("BuildSNIName()=%q, want %q", got, "api-gateway.relay.edge.example.com")
114 - }
115 -}
116 -
117 -func TestServicePublicURL(t *testing.T) {
118 - t.Parallel()
119 -
120 - tests := []struct {
121 - name string
122 - portalURL string
123 - service string
124 - want string
125 - }{
126 - {
127 - name: "preserve explicit scheme and root host",
128 - portalURL: "https://portal.example.com:4017/admin",
129 - service: "my-app",
130 - want: "https://my-app.portal.example.com",
131 - },
132 - {
133 - name: "default scheme for host-only portal URL",
134 - portalURL: "portal.example.com",
135 - service: "My-App",
136 - want: "https://my-app.portal.example.com",
137 - },
138 - {
139 - name: "invalid service returns empty",
140 - portalURL: "https://portal.example.com",
141 - service: "my_app",
142 - want: "",
143 - },
144 - {
145 - name: "invalid portal returns empty",
146 - portalURL: "",
147 - service: "my-app",
148 - want: "",
149 - },
150 - }
151 -
152 - for _, tt := range tests {
153 - t.Run(tt.name, func(t *testing.T) {
154 - t.Parallel()
155 -
156 - got := ServicePublicURL(tt.portalURL, tt.service)
157 - if got != tt.want {
158 - t.Fatalf("ServicePublicURL(%q, %q)=%q, want %q", tt.portalURL, tt.service, got, tt.want)
159 - }
160 - })
161 - }
162 -}
163 -
164 -func TestNonApexPortalRoundTrip(t *testing.T) {
165 - t.Parallel()
166 -
167 - const (
168 - portalURL = "https://portal.edge.example.com:8443"
169 - leaseNameInput = "My-App"
170 - expectedLeaseName = "my-app"
171 - expectedHost = "my-app.portal.edge.example.com"
172 - expectedPublicURL = "https://my-app.portal.edge.example.com"
173 - )
174 -
175 - if got := BuildSNIName(leaseNameInput, portalURL); got != expectedHost {
176 - t.Fatalf("BuildSNIName(%q, %q)=%q, want %q", leaseNameInput, portalURL, got, expectedHost)
177 - }
178 - if got := ServicePublicURL(portalURL, leaseNameInput); got != expectedPublicURL {
179 - t.Fatalf("ServicePublicURL(%q, %q)=%q, want %q", portalURL, leaseNameInput, got, expectedPublicURL)
180 - }
181 -
182 - appURLs := []string{
183 - portalURL,
184 - "portal.edge.example.com:8443",
185 - "*.portal.edge.example.com:8443",
186 - }
187 - for _, appURL := range appURLs {
188 - t.Run(appURL, func(t *testing.T) {
189 - t.Parallel()
190 -
191 - gotLeaseName, ok := LeaseNameFromHost(expectedHost, appURL)
192 - if !ok {
193 - t.Fatalf("LeaseNameFromHost(%q, %q) expected success", expectedHost, appURL)
194 - }
195 - if gotLeaseName != expectedLeaseName {
196 - t.Fatalf("LeaseNameFromHost(%q, %q)=%q, want %q", expectedHost, appURL, gotLeaseName, expectedLeaseName)
197 - }
198 - })
199 - }
200 -
201 - if gotLeaseName, ok := LeaseNameFromHost("portal.edge.example.com", portalURL); ok {
202 - t.Fatalf("LeaseNameFromHost() unexpected success for apex host, got %q", gotLeaseName)
203 - }
204 -}
205 -
206 -func TestIsValidServiceName(t *testing.T) {
207 - t.Parallel()
208 -
209 - if !IsValidServiceName("my-app") {
210 - t.Fatalf("expected valid service name")
211 - }
212 - if IsValidServiceName("my_app") {
213 - t.Fatalf("expected underscore to be invalid for DNS-safe service names")
214 - }
215 -}
216 -
217 -func TestParsePortalAddressPreservesFullRootHost(t *testing.T) {
218 - t.Parallel()
219 -
220 - scheme, rootHost, hostPort, ok := parsePortalAddress("https://portal.edge.example.com:8443/path", "http")
221 - if !ok {
222 - t.Fatalf("expected parsePortalAddress success")
223 - }
224 - if scheme != "https" {
225 - t.Fatalf("scheme=%q, want https", scheme)
226 - }
227 - if rootHost != "portal.edge.example.com" {
228 - t.Fatalf("rootHost=%q, want portal.edge.example.com", rootHost)
229 - }
230 - if hostPort != "portal.edge.example.com:8443" {
231 - t.Fatalf("hostPort=%q, want portal.edge.example.com:8443", hostPort)
232 - }
233 -}
234 -
235 -func TestParsePortalAddressHostOnlyUsesFallbackScheme(t *testing.T) {
236 - t.Parallel()
237 -
238 - scheme, rootHost, hostPort, ok := parsePortalAddress("portal.edge.example.com:9443", "http")
239 - if !ok {
240 - t.Fatalf("expected parsePortalAddress success")
241 - }
242 - if scheme != "http" {
243 - t.Fatalf("scheme=%q, want http", scheme)
244 - }
245 - if rootHost != "portal.edge.example.com" {
246 - t.Fatalf("rootHost=%q, want portal.edge.example.com", rootHost)
247 - }
248 - if hostPort != "portal.edge.example.com:9443" {
249 - t.Fatalf("hostPort=%q, want portal.edge.example.com:9443", hostPort)
250 - }
251 -}
252 -
253 -func TestParsePortalAddressNormalizesTrailingDotAndCase(t *testing.T) {
254 - t.Parallel()
255 -
256 - scheme, rootHost, hostPort, ok := parsePortalAddress("HTTPS://Portal.Edge.Example.COM.:7443", "http")
257 - if !ok {
258 - t.Fatalf("expected parsePortalAddress success")
259 - }
260 - if scheme != "https" {
261 - t.Fatalf("scheme=%q, want https", scheme)
262 - }
263 - if rootHost != "portal.edge.example.com" {
264 - t.Fatalf("rootHost=%q, want portal.edge.example.com", rootHost)
265 - }
266 - if hostPort != "portal.edge.example.com:7443" {
267 - t.Fatalf("hostPort=%q, want portal.edge.example.com:7443", hostPort)
268 - }
269 -}
270 -
271 -func TestBuildSNINameFallbackNormalizesRootHost(t *testing.T) {
272 - t.Parallel()
273 -
274 - got := BuildSNIName("Api-Gateway", " *.Portal.Edge.Example.COM. ")
275 - want := "api-gateway.portal.edge.example.com"
276 - if got != want {
277 - t.Fatalf("BuildSNIName fallback=%q, want %q", got, want)
278 - }
279 -}
280 -
281 -func TestIsLocalhost(t *testing.T) {
282 - t.Parallel()
283 -
284 - tests := []struct {
285 - name string
286 - host string
287 - want bool
288 - }{
289 - {name: "localhost", host: "localhost", want: true},
290 - {name: "localhost with port", host: "localhost:4017", want: true},
291 - {name: "subdomain localhost", host: "portal.localhost", want: true},
292 - {name: "ipv4 loopback", host: "127.0.0.1", want: true},
293 - {name: "ipv4 loopback with port", host: "127.0.0.1:4017", want: true},
294 - {name: "ipv6 loopback", host: "::1", want: true},
295 - {name: "ipv6 loopback with port", host: "[::1]:4017", want: true},
296 - {name: "public host", host: "example.com", want: false},
297 - {name: "public ip", host: "8.8.8.8", want: false},
298 - }
299 -
300 - for _, tt := range tests {
301 - t.Run(tt.name, func(t *testing.T) {
302 - t.Parallel()
303 -
304 - if got := IsLocalhost(tt.host); got != tt.want {
305 - t.Fatalf("IsLocalhostHost(%q)=%v, want %v", tt.host, got, tt.want)
306 - }
307 - })
308 - }
309 -}
types/types.go deleted
-70
@@ -1,70 +0,0 @@
1 -package types
2 -
3 -import "time"
4 -
5 -const (
6 - // Reverse-connect protocol markers and headers.
7 - ReverseKeepaliveMarker = byte(0x00)
8 - NonTLSStartMarker = byte(0x01)
9 - TLSStartMarker = byte(0x02)
10 - ReverseConnectTokenHeader = "X-Portal-Reverse-Token"
11 -)
12 -
13 -// Lease represents a registered service.
14 -type Lease struct {
15 - Expires time.Time `json:"expires"`
16 - ID string `json:"id"`
17 - Name string `json:"name"`
18 - ReverseToken string `json:"-"`
19 - Metadata Metadata `json:"metadata"`
20 - TLS bool `json:"tls"`
21 -}
22 -
23 -// LeaseEntry represents a registered lease with expiration tracking.
24 -type LeaseEntry struct {
25 - Lease *Lease
26 - LastSeen time.Time
27 - FirstSeen time.Time
28 -}
29 -
30 -// Metadata holds service metadata for a lease.
31 -type Metadata struct {
32 - Description string `json:"description,omitempty"`
33 - Thumbnail string `json:"thumbnail,omitempty"`
34 - Owner string `json:"owner,omitempty"`
35 - Tags []string `json:"tags,omitempty"`
36 - Hide bool `json:"hide,omitempty"`
37 -}
38 -
39 -// MetadataOption configures Metadata.
40 -type MetadataOption func(*Metadata)
41 -
42 -func WithDescription(description string) MetadataOption {
43 - return func(m *Metadata) {
44 - m.Description = description
45 - }
46 -}
47 -
48 -func WithTags(tags []string) MetadataOption {
49 - return func(m *Metadata) {
50 - m.Tags = tags
51 - }
52 -}
53 -
54 -func WithThumbnail(thumbnail string) MetadataOption {
55 - return func(m *Metadata) {
56 - m.Thumbnail = thumbnail
57 - }
58 -}
59 -
60 -func WithOwner(owner string) MetadataOption {
61 - return func(m *Metadata) {
62 - m.Owner = owner
63 - }
64 -}
65 -
66 -func WithHide(hide bool) MetadataOption {
67 - return func(m *Metadata) {
68 - m.Hide = hide
69 - }
70 -}