Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
.ijwb/

bazel-*
MODULE.bazel.lock
examples/MODULE.bazel.lock
Expand Down
13 changes: 12 additions & 1 deletion cmd/svcinit/BUILD.bazel
Original file line number Diff line number Diff line change
@@ -1,9 +1,11 @@
load("@rules_go//go:def.bzl", "go_binary", "go_library")
load("@rules_go//go:def.bzl", "go_binary", "go_library", "go_test")

go_library(
name = "svcinit_lib",
srcs = [
"main.go",
"reserve_reusable_port_unix.go",
"reserve_reusable_port_windows.go",
"set_sockopts_for_port_assignment_unix.go",
"set_sockopts_for_port_assignment_windows.go",
],
Expand Down Expand Up @@ -49,10 +51,19 @@ go_library(
"@rules_go//go/platform:solaris": [
"@org_golang_x_sys//unix",
],
"@rules_go//go/platform:windows": [
"@org_golang_x_sys//windows",
],
"//conditions:default": [],
}),
)

go_test(
name = "svcinit_test",
srcs = ["reserve_reusable_port_test.go"],
embed = [":svcinit_lib"],
)

go_binary(
name = "svcinit",
data = ["//cmd/get_assigned_port"],
Expand Down
98 changes: 64 additions & 34 deletions cmd/svcinit/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import (
"encoding/json"
"errors"
"fmt"
"io"
"log"
"maps"
"math"
Expand Down Expand Up @@ -124,8 +125,9 @@ func main() {
listener, err := net.Listen("tcp", "127.0.0.1:0")
must(err)

ports, err := assignPorts(unversionedSpecs)
ports, reservedPorts, err := assignPorts(unversionedSpecs)
must(err)
defer closeReservedPorts(reservedPorts)

svcctlPort := listener.Addr().(*net.TCPAddr).Port
svcctlPortStr := strconv.Itoa(svcctlPort)
Expand Down Expand Up @@ -374,9 +376,10 @@ func readServiceSpecs(
func assignPorts(
serviceSpecs map[string]svclib.ServiceSpec,
) (
svclib.Ports, error,
svclib.Ports, map[string][]io.Closer, error,
) {
var toClose []net.Listener
var toClose []io.Closer
reservedPorts := map[string][]io.Closer{}
ports := svclib.Ports{}

for label, spec := range serviceSpecs {
Expand All @@ -386,34 +389,49 @@ func assignPorts(
}

// Note, this can cause collisions. So be careful!
// To avoid port collisions, set the `so_reuseport_aware` option on the service definition
// and use the SO_REUSEPORT socket option in your services.
// To avoid port collisions, set so_reuseport_aware on the service definition
// and use SO_REUSEPORT on Unix or SO_REUSEADDR on Windows in your services.
for portName, port := range namedPorts {
// We do a bit of a dance here to set SO_LINGER to 0. For details, see
// https://stackoverflow.com/questions/71975992/what-really-is-the-linger-time-that-can-be-set-with-so-linger-on-sockets
lc := net.ListenConfig{
Control: func(network, address string, conn syscall.RawConn) error {
var setSockoptErr error
err := conn.Control(func(fd uintptr) {
setSockoptErr = setSockoptsForPortAssignment(fd, &syscall.Linger{
Onoff: 1,
Linger: 0,
var reservedPort io.Closer
var err error
if spec.SoReuseportAware {
requestedPort, parseErr := strconv.Atoi(port)
if parseErr != nil || requestedPort < 0 || requestedPort > 65535 {
return nil, nil, fmt.Errorf("invalid port %q for %s", port, label)
}
reservedPort, port, err = reserveReusablePort(requestedPort)
if err != nil {
return nil, nil, err
}
} else {
// We do a bit of a dance here to set SO_LINGER to 0. For details, see
// https://stackoverflow.com/questions/71975992/what-really-is-the-linger-time-that-can-be-set-with-so-linger-on-sockets
lc := net.ListenConfig{
Control: func(network, address string, conn syscall.RawConn) error {
var setSockoptErr error
err := conn.Control(func(fd uintptr) {
setSockoptErr = setSockoptsForPortAssignment(fd, &syscall.Linger{
Onoff: 1,
Linger: 0,
})
})
})
if err != nil {
return err
}
return setSockoptErr
},
}
if err != nil {
return err
}
return setSockoptErr
},
}

listener, err := lc.Listen(context.Background(), "tcp", "127.0.0.1:"+port)
if err != nil {
return nil, err
}
_, port, err = net.SplitHostPort(listener.Addr().String())
if err != nil {
return nil, err
listener, listenErr := lc.Listen(context.Background(), "tcp", "127.0.0.1:"+port)
if listenErr != nil {
return nil, nil, listenErr
}
_, port, err = net.SplitHostPort(listener.Addr().String())
if err != nil {
listener.Close()
return nil, nil, err
}
reservedPort = listener
}

qualifiedPortName := label
Expand Down Expand Up @@ -442,15 +460,17 @@ func assignPorts(
}

if !spec.SoReuseportAware {
toClose = append(toClose, listener)
toClose = append(toClose, reservedPort)
} else {
reservedPorts[label] = append(reservedPorts[label], reservedPort)
}
}
}

for _, listener := range toClose {
err := listener.Close()
for _, reservedPort := range toClose {
err := reservedPort.Close()
if err != nil {
return nil, err
return nil, nil, err
}
}

Expand Down Expand Up @@ -481,10 +501,20 @@ func assignPorts(

serializedPorts, err := ports.Marshal()
if err != nil {
return nil, err
return nil, nil, err
}
os.Setenv("ASSIGNED_PORTS", string(serializedPorts))
return ports, nil
return ports, reservedPorts, nil
}

func closeReservedPorts(reservedPorts map[string][]io.Closer) {
for label, ports := range reservedPorts {
for _, port := range ports {
if err := port.Close(); err != nil {
log.Printf("failed to close reusable port reservation for %s: %v\n", label, err)
}
}
}
}

func augmentServiceSpecs(
Expand Down
70 changes: 70 additions & 0 deletions cmd/svcinit/reserve_reusable_port_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,70 @@
package main

import (
"context"
"net"
"syscall"
"testing"
"time"
)

func TestReusablePortReservation(t *testing.T) {
reservation, port, err := reserveReusablePort(0)
if err != nil {
t.Fatal(err)
}
defer reservation.Close()

unawareListener, err := net.Listen("tcp4", "127.0.0.1:"+port)
if err == nil {
unawareListener.Close()
t.Fatal("listener without a reusable-port option unexpectedly claimed the reserved port")
}

lc := net.ListenConfig{
Control: func(network, address string, conn syscall.RawConn) error {
var setSockoptErr error
err := conn.Control(func(fd uintptr) {
setSockoptErr = setSockoptsForPortAssignment(fd, &syscall.Linger{
Onoff: 1,
Linger: 0,
})
})
if err != nil {
return err
}
return setSockoptErr
},
}
listener, err := lc.Listen(context.Background(), "tcp4", "127.0.0.1:"+port)
if err != nil {
t.Fatalf("listen on reserved port: %v", err)
}
defer listener.Close()

acceptErr := make(chan error, 1)
go func() {
conn, err := listener.Accept()
if err != nil {
acceptErr <- err
return
}
conn.Close()
acceptErr <- nil
}()

conn, err := net.DialTimeout("tcp4", "127.0.0.1:"+port, time.Second)
if err != nil {
t.Fatalf("dial service listener: %v", err)
}
conn.Close()

select {
case err := <-acceptErr:
if err != nil {
t.Fatalf("accept from service listener: %v", err)
}
case <-time.After(time.Second):
t.Fatal("connection was not accepted by the service listener")
}
}
56 changes: 56 additions & 0 deletions cmd/svcinit/reserve_reusable_port_unix.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,56 @@
//go:build unix

package main

import (
"fmt"
"io"
"os"
"strconv"
"syscall"
)

func reserveReusablePort(port int) (io.Closer, string, error) {
fd, err := syscall.Socket(syscall.AF_INET, syscall.SOCK_STREAM, syscall.IPPROTO_TCP)
if err != nil {
return nil, "", fmt.Errorf("socket: %w", err)
}
syscall.CloseOnExec(fd)

file := os.NewFile(uintptr(fd), "rules_itest_reserved_reuseport")
success := false
defer func() {
if !success {
file.Close()
}
}()

// Do not set SO_REUSEADDR here. Go TCP listeners enable it by default on
// Linux, where it would allow an unaware listener to claim a bind-only
// reservation. SO_REUSEPORT alone allows the aware service to share it.
if err := setSockoptsForPortAssignment(uintptr(fd), &syscall.Linger{
Onoff: 1,
Linger: 0,
}); err != nil {
return nil, "", fmt.Errorf("set reusable reservation socket options: %w", err)
}

if err := syscall.Bind(fd, &syscall.SockaddrInet4{
Port: port,
Addr: [4]byte{127, 0, 0, 1},
}); err != nil {
return nil, "", fmt.Errorf("bind reusable reservation socket: %w", err)
}

addr, err := syscall.Getsockname(fd)
if err != nil {
return nil, "", fmt.Errorf("getsockname reusable reservation socket: %w", err)
}
tcpAddr, ok := addr.(*syscall.SockaddrInet4)
if !ok {
return nil, "", fmt.Errorf("getsockname returned %T, expected *syscall.SockaddrInet4", addr)
}

success = true
return file, strconv.Itoa(tcpAddr.Port), nil
}
67 changes: 67 additions & 0 deletions cmd/svcinit/reserve_reusable_port_windows.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,67 @@
//go:build windows

package main

import (
"fmt"
"io"
"strconv"
"syscall"

"golang.org/x/sys/windows"
)

type windowsPortReservation struct {
socket windows.Handle
}

func (r *windowsPortReservation) Close() error {
return windows.Closesocket(r.socket)
}

func reserveReusablePort(port int) (io.Closer, string, error) {
socket, err := windows.WSASocket(
windows.AF_INET,
windows.SOCK_STREAM,
windows.IPPROTO_TCP,
nil,
0,
windows.WSA_FLAG_OVERLAPPED|windows.WSA_FLAG_NO_HANDLE_INHERIT,
)
if err != nil {
return nil, "", fmt.Errorf("socket: %w", err)
}

success := false
defer func() {
if !success {
windows.Closesocket(socket)
}
}()

if err := setSockoptsForPortAssignment(uintptr(socket), &syscall.Linger{
Onoff: 1,
Linger: 0,
}); err != nil {
return nil, "", fmt.Errorf("set reusable reservation socket options: %w", err)
}

if err := windows.Bind(socket, &windows.SockaddrInet4{
Port: port,
Addr: [4]byte{127, 0, 0, 1},
}); err != nil {
return nil, "", fmt.Errorf("bind reusable reservation socket: %w", err)
}

addr, err := windows.Getsockname(socket)
if err != nil {
return nil, "", fmt.Errorf("getsockname reusable reservation socket: %w", err)
}
tcpAddr, ok := addr.(*windows.SockaddrInet4)
if !ok {
return nil, "", fmt.Errorf("getsockname returned %T, expected *windows.SockaddrInet4", addr)
}

success = true
return &windowsPortReservation{socket: socket}, strconv.Itoa(tcpAddr.Port), nil
}
Loading
Loading