Skip to content

Commit 69e58f9

Browse files
committed
Read port from Puma socket
1 parent d53d915 commit 69e58f9

2 files changed

Lines changed: 62 additions & 4 deletions

File tree

lib/tidewave.rb

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -167,7 +167,11 @@ def origin_allowed_path?(path)
167167
end
168168

169169
def local_port(request)
170-
request.get_header("SERVER_PORT").to_i.nonzero?
170+
sock = request.env["puma.socket"]
171+
return unless sock
172+
173+
addr = sock.respond_to?(:local_address) ? sock.local_address : sock.to_io.local_address
174+
addr.ip? ? addr.ip_port : nil
171175
end
172176

173177
def valid_client_ip?(request)

test/tidewave_test.rb

Lines changed: 57 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -103,7 +103,8 @@ def test_config_endpoint_returns_json
103103
app,
104104
path: "/tidewave/config",
105105
host: "example.test:3000",
106-
server_port: "4000"
106+
server_port: "4000",
107+
puma_socket: fake_socket(port: 5000)
107108
)
108109

109110
assert_equal 200, status
@@ -116,7 +117,7 @@ def test_config_endpoint_returns_json
116117
assert_equal Tidewave::VERSION, payload["tidewave_version"]
117118
assert_equal({ "id" => "dashbit" }, payload["team"])
118119
assert_equal "demo-app", payload["project_name"]
119-
assert_equal 4000, payload["local_port"]
120+
assert_equal 5000, payload["local_port"]
120121
end
121122

122123
def test_config_endpoint_includes_orm_adapter_when_configured
@@ -132,6 +133,28 @@ def test_config_endpoint_includes_orm_adapter_when_configured
132133
assert_equal "sequel", JSON.parse(body)["orm_adapter"]
133134
end
134135

136+
def test_config_endpoint_returns_nil_local_port_for_missing_invalid_or_zero_puma_socket_port
137+
app = Tidewave.new(@downstream_app, allow_remote_access: true, project_name: "demo-app")
138+
139+
[ nil, fake_socket(port: nil, ip: false) ].each do |puma_socket|
140+
_status, _headers, body = perform_request(app, path: "/tidewave/config", puma_socket: puma_socket)
141+
142+
assert_nil JSON.parse(body)["local_port"]
143+
end
144+
end
145+
146+
def test_config_endpoint_reads_local_port_from_socket_io_fallback
147+
app = Tidewave.new(@downstream_app, allow_remote_access: true, project_name: "demo-app")
148+
149+
_status, _headers, body = perform_request(
150+
app,
151+
path: "/tidewave/config",
152+
puma_socket: fake_socket(port: 5001, direct_local_address: false)
153+
)
154+
155+
assert_equal 5001, JSON.parse(body)["local_port"]
156+
end
157+
135158
def test_project_name_is_required
136159
error = assert_raises(ArgumentError) do
137160
Tidewave.new(@downstream_app, allow_remote_access: true)
@@ -205,7 +228,7 @@ def test_logs_security_rejections
205228

206229
private
207230

208-
def perform_request(app, path:, method: "GET", body: nil, remote_addr: "127.0.0.1", origin: nil, forwarded_for: nil, host: nil, server_port: nil)
231+
def perform_request(app, path:, method: "GET", body: nil, remote_addr: "127.0.0.1", origin: nil, forwarded_for: nil, host: nil, server_port: nil, puma_socket: nil)
209232
env = Rack::MockRequest.env_for(path,
210233
method: method,
211234
input: body.to_s,
@@ -215,11 +238,42 @@ def perform_request(app, path:, method: "GET", body: nil, remote_addr: "127.0.0.
215238
env["HTTP_X_FORWARDED_FOR"] = forwarded_for if forwarded_for
216239
env["HTTP_HOST"] = host if host
217240
env["SERVER_PORT"] = server_port if server_port
241+
env["puma.socket"] = puma_socket if puma_socket
218242

219243
status, headers, response = app.call(env)
220244
[ status, headers, collect_body(response) ]
221245
end
222246

247+
def fake_socket(port:, ip: true, direct_local_address: true)
248+
addr = Struct.new(:port, :ip) do
249+
def ip?
250+
ip
251+
end
252+
253+
def ip_port
254+
port
255+
end
256+
end.new(port, ip)
257+
258+
if direct_local_address
259+
Struct.new(:addr) do
260+
def local_address
261+
addr
262+
end
263+
end.new(addr)
264+
else
265+
Struct.new(:addr) do
266+
def to_io
267+
Struct.new(:addr) do
268+
def local_address
269+
addr
270+
end
271+
end.new(addr)
272+
end
273+
end.new(addr)
274+
end
275+
end
276+
223277
def collect_body(response)
224278
body = +""
225279
response.each { |part| body << part }

0 commit comments

Comments
 (0)