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, ®isterReq, "[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: ®isterReq.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, ®isterResp); 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
-}