diff --git a/README.md b/README.md index 3eee4f2..837857a 100644 --- a/README.md +++ b/README.md @@ -135,7 +135,7 @@ Allowlists for both requests and bind mount restrictions can be specified for pa 1. Set `-proxycontainername` or the environment variable `SP_PROXYCONTAINERNAME` to the name of the socket proxy container. 2. Make sure that each container that will use the socket proxy is in a Docker network that the socket proxy container is also in. -3. Use the same regex syntax for request allowlists and for bind mount restrictions that were discussed earlier, but for labels on each container that will use the socket proxy. Each label name will have the prefix of `socket-proxy.allow.`, with `socket-proxy.allow.bindmountfrom` for bind mount restrictions. For example: +3. Use the same regex syntax for request allowlists and for bind mount restrictions that were discussed earlier, but for labels on each container that will use the socket proxy. Each label name has the prefix `.allow.`; by default this is `socket-proxy.allow.`, with `socket-proxy.allow.bindmountfrom` for bind mount restrictions. Set `-dockerlabelprefix` or `SP_DOCKERLABELPREFIX` when multiple socket proxies share a Docker daemon. For example, `-dockerlabelprefix=traefik-socket-proxy` uses labels beginning with `traefik-socket-proxy.allow.`. ```yaml services: @@ -254,6 +254,7 @@ socket-proxy can be configured via command-line parameters or via environment va | `-watchdoginterval` | `SP_WATCHDOGINTERVAL` | `0` | Check for socket availability every x seconds (disable checks, if not set or value is 0) | | `-proxysocketendpoint` | `SP_PROXYSOCKETENDPOINT` | (not set) | Proxy to the given unix socket instead of a TCP port | | `-proxysocketendpointfilemode` | `SP_PROXYSOCKETENDPOINTFILEMODE` | `0600` | Explicitly set the file mode for the filtered unix socket endpoint (only useful with `-proxysocketendpoint`) | +| `-dockerlabelprefix` | `SP_DOCKERLABELPREFIX` | `socket-proxy` | Specifies the prefix before `.allow.` in Docker container labels used for per-container allowlists. It must contain only lowercase letters, digits, dots, and hyphens. For example, `-dockerlabelprefix=traefik-socket-proxy` recognizes `traefik-socket-proxy.allow.get`. | | `-proxycontainername` | `SP_PROXYCONTAINERNAME` | (not set) | Provides the name of the socket proxy container to enable per-container allowlists specified by Docker container labels (not available with `-proxysocketendpoint`) | ### Changelog diff --git a/cmd/socket-proxy/handlehttprequest.go b/cmd/socket-proxy/handlehttprequest.go index f366135..a70461b 100644 --- a/cmd/socket-proxy/handlehttprequest.go +++ b/cmd/socket-proxy/handlehttprequest.go @@ -74,6 +74,17 @@ func determineAllowList(r *http.Request) (config.AllowList, bool) { slog.Warn("cannot get valid IP address for client allowlist check", "reason", err, "method", r.Method, "URL", r.URL, "client", r.RemoteAddr) // #nosec G706 - structured logging (slog) safely encodes values } if !allowedIP { + // A container can send its first request before Docker's start event has + // updated the per-container allowlist. Refresh it once before denying + // the request, while still failing closed if Docker cannot be queried. + if cfg.ProxyContainerName != "" { + refreshedAllowLists, err := cfg.RefreshAllowLists(r.Context()) + if err != nil { + slog.Warn("failed to refresh per-container allowlists", "error", err) + } else if allowList, found := refreshedAllowLists.FindByIP(clientIPStr); found { + return allowList, true + } + } return config.AllowList{}, false } } diff --git a/internal/config/config.go b/internal/config/config.go index f5212d5..503f55b 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -1,7 +1,6 @@ package config import ( - "context" "errors" "flag" "fmt" @@ -15,20 +14,11 @@ import ( "slices" "strconv" "strings" - "sync" - "time" - - "github.com/wollomatic/socket-proxy/internal/docker/api/types/container" - "github.com/wollomatic/socket-proxy/internal/docker/api/types/events" - "github.com/wollomatic/socket-proxy/internal/docker/api/types/filters" - "github.com/wollomatic/socket-proxy/internal/docker/client" ) -const allowedDockerLabelPrefix = "socket-proxy.allow." - const ( defaultAllowFrom = "127.0.0.1/32" // allowed IPs to connect to the proxy - defaultAllowHealthcheck = false // allow health check requests (HEAD http://127.0.0.1:55555/health) + defaultAllowHealthcheck = false // allow health check requests (HEAD http://localhost:55555/health) defaultLogJSON = false // if true, log in JSON format defaultLogLevel = "INFO" // log level as string defaultListenIP = "127.0.0.1" // ip address to bind the server to @@ -40,10 +30,12 @@ const ( defaultProxySocketEndpoint = "" // empty string means no socket listener, but regular TCP listener defaultProxySocketEndpointFileMode = uint(0o600) // set the file mode of the unix socket endpoint defaultAllowBindMountFrom = "" // empty string means no bind mount restrictions + defaultDockerLabelPrefix = "socket-proxy" // prefix before .allow. in per-container allowlist labels defaultProxyContainerName = "" // socket-proxy Docker container name (empty string disables container labels for allowlists) ) type Config struct { + allowListsRefresh allowListsRefreshState AllowLists *AllowListRegistry AllowFrom []string AllowHealthcheck bool @@ -59,13 +51,6 @@ type Config struct { ProxyContainerName string } -type AllowListRegistry struct { - mutex sync.RWMutex // mutex to control read/write of byIP - networks []string // names of networks in which socket proxy access is allowed for non-default allowlists - Default AllowList // default allowlist - byIP map[string]AllowList // map container IP address to allowlist for that container -} - type AllowList struct { ID string // Container ID (empty for the default allowlist) ContainerName string // Container name (empty for the default allowlist) @@ -91,6 +76,8 @@ var supportedHTTPMethods = []string{ http.MethodOptions, } +var dockerLabelPrefixRegexp = regexp.MustCompile(`^[a-z0-9]+(?:[.-][a-z0-9]+)*$`) + // InitConfig reads configuration from environment variables and command-line // flags, validates the resulting values, and returns the initialized Config. func InitConfig() (*Config, error) { @@ -102,6 +89,7 @@ func InitConfig() (*Config, error) { logLevel string endpointFileMode uint allowBindMountFromString string + dockerLabelPrefix string defaultAllowFromValue = defaultAllowFrom defaultAllowHealthcheckValue = defaultAllowHealthcheck defaultLogJSONValue = defaultLogJSON @@ -115,6 +103,7 @@ func InitConfig() (*Config, error) { defaultProxySocketEndpointValue = defaultProxySocketEndpoint defaultProxySocketEndpointFileModeValue = defaultProxySocketEndpointFileMode defaultAllowBindMountFromValue = defaultAllowBindMountFrom + defaultDockerLabelPrefixValue = defaultDockerLabelPrefix defaultProxyContainerNameValue = defaultProxyContainerName ) @@ -171,6 +160,9 @@ func InitConfig() (*Config, error) { if val, ok := os.LookupEnv("SP_ALLOWBINDMOUNTFROM"); ok && val != "" { defaultAllowBindMountFromValue = val } + if val, ok := os.LookupEnv("SP_DOCKERLABELPREFIX"); ok && val != "" { + defaultDockerLabelPrefixValue = val + } if val, ok := os.LookupEnv("SP_PROXYCONTAINERNAME"); ok && val != "" { defaultProxyContainerNameValue = val } @@ -189,7 +181,7 @@ func InitConfig() (*Config, error) { } flag.StringVar(&allowFromString, "allowfrom", defaultAllowFromValue, "allowed IPs or hostname to connect to the proxy") - flag.BoolVar(&cfg.AllowHealthcheck, "allowhealthcheck", defaultAllowHealthcheckValue, "allow health check requests (HEAD http://127.0.0.1:55555/health)") + flag.BoolVar(&cfg.AllowHealthcheck, "allowhealthcheck", defaultAllowHealthcheckValue, "allow health check requests (HEAD http://localhost:55555/health)") flag.BoolVar(&cfg.LogJSON, "logjson", defaultLogJSONValue, "log in JSON format (otherwise log in plain text") flag.StringVar(&listenIP, "listenip", defaultListenIPValue, "ip address to listen on") flag.StringVar(&logLevel, "loglevel", defaultLogLevelValue, "set log level: DEBUG, INFO, WARN, ERROR") @@ -201,12 +193,17 @@ func InitConfig() (*Config, error) { flag.StringVar(&cfg.ProxySocketEndpoint, "proxysocketendpoint", defaultProxySocketEndpointValue, "unix socket endpoint (if set, used instead of the TCP listener)") flag.UintVar(&endpointFileMode, "proxysocketendpointfilemode", defaultProxySocketEndpointFileModeValue, "set the file mode of the unix socket endpoint") flag.StringVar(&allowBindMountFromString, "allowbindmountfrom", defaultAllowBindMountFromValue, "allowed directories for bind mounts (comma-separated)") + flag.StringVar(&dockerLabelPrefix, "dockerlabelprefix", defaultDockerLabelPrefixValue, "prefix before .allow. in Docker container allowlist labels") flag.StringVar(&cfg.ProxyContainerName, "proxycontainername", defaultProxyContainerNameValue, "socket-proxy Docker container name") for i := range methodAllowLists { flag.Var(&methodAllowLists[i].regexStrings, "allow"+methodAllowLists[i].method, "regex for "+methodAllowLists[i].method+" requests (not set means method is not allowed)") } flag.Parse() + if !dockerLabelPrefixRegexp.MatchString(dockerLabelPrefix) { + return nil, fmt.Errorf("invalid dockerlabelprefix %q: use lowercase letters, digits, dots, and hyphens only; dots and hyphens must separate alphanumeric parts", dockerLabelPrefix) + } + // init allowlist registry to configure default allowlist cfg.AllowLists = &AllowListRegistry{} @@ -292,256 +289,11 @@ func InitConfig() (*Config, error) { return nil, err } } + allowedDockerLabelPrefix = dockerLabelPrefix + ".allow." return &cfg, nil } -// UpdateAllowLists populates the byIP allowlists then keeps them updated -func (cfg *Config) UpdateAllowLists() { - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - - dockerClient, err := client.NewClientWithOpts( - client.WithHost("unix://"+cfg.SocketPath), - client.WithAPIVersionNegotiation(), - ) - if err != nil { - slog.Error("failed to create Docker client", "error", err) - return - } - defer func(dockerClient *client.Client) { - err := dockerClient.Close() - if err != nil { - slog.Error("failed to close Docker client", "error", err) - } - }(dockerClient) - - err = cfg.AllowLists.initByIP(ctx, dockerClient) - if err != nil { - slog.Error("failed to initialise non-default allowlists", "error", err) - return - } - slog.Debug("initialised non-default allowlists") - - filter := filters.NewArgs() - filter.Add("type", "container") - filter.Add("event", "start") - filter.Add("event", "restart") - filter.Add("event", "die") - eventsChan, errChan := dockerClient.Events(ctx, events.ListOptions{Filters: filter}) - slog.Debug("subscribed to Docker event stream to update allowlists") - - // print non-default request allowlists - cfg.AllowLists.PrintByIP(cfg.LogJSON) - - // handle Docker events to update allowlists - for { - select { - case event, ok := <-eventsChan: - if !ok { - slog.Info("Docker event stream closed") - return - } - containerName := eventContainerName(event) - slog.Debug("received Docker container event", "action", event.Action, "container", containerName) - addedIPs, removedIPs, updateErr := cfg.AllowLists.updateFromEvent(ctx, dockerClient, event) - if updateErr != nil { - slog.Warn("failed to update allowlists from container event", "error", updateErr) - continue - } - for _, ip := range addedIPs { - cfg.AllowLists.mutex.RLock() - allowList, found := cfg.AllowLists.byIP[ip] - cfg.AllowLists.mutex.RUnlock() - if found { - allowList.Print(ip, cfg.LogJSON) - } - } - for _, ip := range removedIPs { - slog.Info("removed allowlist for container", "container", containerName, "ip", ip) - } - case err := <-errChan: - if err != nil { - slog.Error("received error from Docker event stream", "error", err) - return - } - } - } -} - -// PrintNetworks prints the allowed networks -func (allowLists *AllowListRegistry) PrintNetworks() { - if len(allowLists.networks) > 0 { - slog.Info("socket proxy networks detected", "socketproxynetworks", allowLists.networks) - } else { - // we only log this on DEBUG level because the socket proxy networks are used for per-container allowlists - slog.Debug("no socket proxy networks detected") - } -} - -// PrintDefault prints the default allowlist -func (allowLists *AllowListRegistry) PrintDefault(logJSON bool) { - allowLists.Default.Print("", logJSON) -} - -// PrintByIP prints the non-default allowlists -func (allowLists *AllowListRegistry) PrintByIP(logJSON bool) { - allowLists.mutex.RLock() - defer allowLists.mutex.RUnlock() - for ip, allowList := range allowLists.byIP { - allowList.Print(ip, logJSON) - } -} - -// FindByIP returns the allowlist corresponding to the given IP address if found -func (allowLists *AllowListRegistry) FindByIP(ip string) (AllowList, bool) { - allowLists.mutex.RLock() - defer allowLists.mutex.RUnlock() - allowList, found := allowLists.byIP[ip] - return allowList, found -} - -// initialise allowlist registry byIP allowlists -func (allowLists *AllowListRegistry) initByIP(ctx context.Context, dockerClient *client.Client) error { - filter := filters.NewArgs() - for _, network := range allowLists.networks { - filter.Add("network", network) - } - containers, err := dockerClient.ContainerList(ctx, container.ListOptions{Filters: filter}) - if err != nil { - return err - } - - allowLists.mutex.Lock() - defer allowLists.mutex.Unlock() - - allowLists.byIP = make(map[string]AllowList) - - for _, cntr := range containers { - allowedRequests, allowedBindMounts, err := extractLabelData(cntr) - if err != nil { - allowLists.byIP = nil - return err - } - - if len(allowedRequests) > 0 || len(allowedBindMounts) > 0 { - for networkID, cntrNetwork := range cntr.NetworkSettings.Networks { - if slices.Contains(allowLists.networks, networkID) { - allowList := AllowList{ - ID: cntr.ID, - ContainerName: containerName(cntr), - AllowedRequests: allowedRequests, - AllowedBindMounts: allowedBindMounts, - } - - if len(cntrNetwork.IPAddress) > 0 { - allowLists.byIP[cntrNetwork.IPAddress] = allowList - } - if len(cntrNetwork.GlobalIPv6Address) > 0 { - allowLists.byIP[cntrNetwork.GlobalIPv6Address] = allowList - } - } - } - } - } - - return nil -} - -// update the allowlist registry based on the Docker event -func (allowLists *AllowListRegistry) updateFromEvent( - ctx context.Context, dockerClient *client.Client, event events.Message, -) ([]string, []string, error) { - containerID := event.Actor.ID - var ( - addedIPs []string - removedIPs []string - err error - ) - - switch event.Action { - case "start", "restart": - addedIPs, err = allowLists.add(ctx, dockerClient, containerID) - if err != nil { - return nil, nil, err - } - case "die": - removedIPs = allowLists.remove(containerID) - } - return addedIPs, removedIPs, nil -} - -// add the allowlist for the container with the given ID to the allowlist registry -// if it has at least one socket-proxy allow label and is in a same network as the socket-proxy -func (allowLists *AllowListRegistry) add( - ctx context.Context, dockerClient *client.Client, containerID string, -) ([]string, error) { - filter := filters.NewArgs() - filter.Add("id", containerID) - for _, network := range allowLists.networks { - filter.Add("network", network) - } - containers, err := dockerClient.ContainerList(ctx, container.ListOptions{Filters: filter}) - if err != nil { - return nil, err - } - if len(containers) == 0 { - slog.Debug("container is not in a network with socket-proxy or may have stopped", "id", shortContainerID(containerID)) - return nil, nil - } - cntr := containers[0] - - allowedRequests, allowedBindMounts, err := extractLabelData(cntr) - if err != nil { - return nil, err - } - - var ips []string - if len(allowedRequests) > 0 || len(allowedBindMounts) > 0 { - allowList := AllowList{ - ID: cntr.ID, - ContainerName: containerName(cntr), - AllowedRequests: allowedRequests, - AllowedBindMounts: allowedBindMounts, - } - - allowLists.mutex.Lock() - defer allowLists.mutex.Unlock() - - for networkID, cntrNetwork := range cntr.NetworkSettings.Networks { - if slices.Contains(allowLists.networks, networkID) { - ipv4Address := cntrNetwork.IPAddress - if len(ipv4Address) > 0 { - allowLists.byIP[ipv4Address] = allowList - ips = append(ips, ipv4Address) - } - ipv6Address := cntrNetwork.GlobalIPv6Address - if len(ipv6Address) > 0 { - allowLists.byIP[ipv6Address] = allowList - ips = append(ips, ipv6Address) - } - } - } - } - - return ips, nil -} - -// remove allowlists having the given container ID from the allowlist registry -func (allowLists *AllowListRegistry) remove(containerID string) []string { - allowLists.mutex.Lock() - defer allowLists.mutex.Unlock() - - var removedIPs []string - for ip, allowList := range allowLists.byIP { - if allowList.ID == containerID { - delete(allowLists.byIP, ip) - removedIPs = append(removedIPs, ip) - } - } - return removedIPs -} - // Print prints the allowlist, including the IP address of the associated container if it is not empty, // and in JSON format if logJSON is true func (allowList AllowList) Print(ip string, logJSON bool) { @@ -596,33 +348,6 @@ func (allowList AllowList) Print(ip string, logJSON bool) { } } -// containerName returns Docker's container name without its leading slash. -// It falls back to the short container ID for unusual responses without a name. -func containerName(cntr container.Summary) string { - for _, name := range cntr.Names { - if name = strings.TrimPrefix(name, "/"); name != "" { - return name - } - } - return shortContainerID(cntr.ID) -} - -// eventContainerName returns the name Docker includes with container events. -// It falls back to the short container ID when the event has no name. -func eventContainerName(event events.Message) string { - if name := event.Actor.Attributes["name"]; name != "" { - return name - } - return shortContainerID(event.Actor.ID) -} - -func shortContainerID(id string) string { - if len(id) > 12 { - return id[:12] - } - return id -} - // compile allowed requests regex pattern func compileRegexp(regex, method, configLocation string) (*regexp.Regexp, error) { r, err := regexp.Compile("^" + regex + "$") @@ -662,80 +387,3 @@ func parseAllowedBindMounts(allowBindMountFromString string) ([]string, error) { } return allowedBindMounts, nil } - -// return list of docker networks that the socket-proxy container is in -func listSocketProxyNetworks(socketPath, proxyContainerName string) ([]string, error) { - cntr, err := getSocketProxyContainerSummary(socketPath, proxyContainerName) - if err != nil { - return nil, err - } - - networks := make([]string, 0, len(cntr.NetworkSettings.Networks)) - for networkID := range cntr.NetworkSettings.Networks { - networks = append(networks, networkID) - } - return networks, nil -} - -// return Docker container summary for the socket proxy container -func getSocketProxyContainerSummary(socketPath, proxyContainerName string) (container.Summary, error) { - const maxTries = 3 - - dockerClient, err := client.NewClientWithOpts( - client.WithHost("unix://"+socketPath), - client.WithAPIVersionNegotiation(), - ) - if err != nil { - return container.Summary{}, err - } - defer func(dockerClient *client.Client) { - err := dockerClient.Close() - if err != nil { - slog.Error("failed to close Docker client", "error", err) - } - }(dockerClient) - - ctx := context.Background() - filter := filters.NewArgs() - filter.Add("name", proxyContainerName) - var containers []container.Summary - for i := 1; i <= maxTries; i++ { - containers, err = dockerClient.ContainerList(ctx, container.ListOptions{Filters: filter}) - if err != nil { - return container.Summary{}, err - } - if len(containers) > 0 { - return containers[0], nil - } - if i < maxTries { - time.Sleep(time.Duration(i) * time.Second) - } - } - return container.Summary{}, fmt.Errorf("socket-proxy container \"%s\" was not found after %d attempts; verify the container name is correct and the container is running", proxyContainerName, maxTries) -} - -// extract Docker container allowlist label data from the container summary -func extractLabelData(cntr container.Summary) (map[string][]*regexp.Regexp, []string, error) { - allowedRequests := make(map[string][]*regexp.Regexp) - var allowedBindMounts []string - for labelName, labelValue := range cntr.Labels { - if strings.HasPrefix(labelName, allowedDockerLabelPrefix) && labelValue != "" { - allowSpec := strings.ToUpper(strings.TrimPrefix(labelName, allowedDockerLabelPrefix)) - method, _, _ := strings.Cut(allowSpec, ".") - if slices.Contains(supportedHTTPMethods, method) { - r, err := compileRegexp(labelValue, method, "docker container label") - if err != nil { - return nil, nil, err - } - allowedRequests[method] = append(allowedRequests[method], r) - } else if allowSpec == "BINDMOUNTFROM" { - var err error - allowedBindMounts, err = parseAllowedBindMounts(labelValue) - if err != nil { - return nil, nil, err - } - } - } - } - return allowedRequests, allowedBindMounts, nil -} diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 9b3c8a5..0dc774d 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -4,14 +4,8 @@ import ( "flag" "math" "os" - "reflect" - "regexp" - "sort" "strconv" "testing" - - "github.com/wollomatic/socket-proxy/internal/docker/api/types/container" - "github.com/wollomatic/socket-proxy/internal/docker/api/types/events" ) func resetFlagsForTest(t *testing.T, args []string) func() { @@ -19,6 +13,7 @@ func resetFlagsForTest(t *testing.T, args []string) func() { prevCommandLine := flag.CommandLine prevArgs := os.Args + prevDockerLabelPrefix := allowedDockerLabelPrefix flag.CommandLine = flag.NewFlagSet(args[0], flag.ContinueOnError) flag.CommandLine.SetOutput(os.Stderr) @@ -27,164 +22,7 @@ func resetFlagsForTest(t *testing.T, args []string) func() { return func() { flag.CommandLine = prevCommandLine os.Args = prevArgs - } -} - -func Test_extractLabelData(t *testing.T) { - tests := []struct { - name string // description of this test case - // Named input parameters for target function. - cntr container.Summary - want map[string][]*regexp.Regexp - want2 []string - wantErr bool - }{ - { - name: "valid labels with multiple methods and regexes", - cntr: container.Summary{ - Labels: map[string]string{ - "socket-proxy.allow.get.0": "regex1", - "socket-proxy.allow.get.1": "regex2", - "socket-proxy.allow.post": "regex3", - }, - }, - want: map[string][]*regexp.Regexp{ - "GET": {regexp.MustCompile("^regex1$"), regexp.MustCompile("^regex2$")}, - "POST": {regexp.MustCompile("^regex3$")}, - }, - want2: nil, - wantErr: false, - }, - { - name: "invalid regex in label value", - cntr: container.Summary{ - Labels: map[string]string{ - "socket-proxy.allow.get": "invalid[regex", - }, - }, - want: nil, - want2: nil, - wantErr: true, - }, - { - name: "non-allow labels are ignored", - cntr: container.Summary{ - Labels: map[string]string{ - "socket-proxy.allow.get": "regex1", - "other.label": "value", - }, - }, - want: map[string][]*regexp.Regexp{ - "GET": {regexp.MustCompile("^regex1$")}, - }, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got, got2, gotErr := extractLabelData(tt.cntr) - if gotErr != nil { - if !tt.wantErr { - t.Errorf("extractLabelData() failed: %v", gotErr) - } - return - } - if tt.wantErr { - t.Fatal("extractLabelData() succeeded unexpectedly") - } - if !regexMapsEqual(got, tt.want) { - t.Errorf("extractLabelData() = %v, want %v", got, tt.want) - } - if !reflect.DeepEqual(got2, tt.want2) { - t.Errorf("extractLabelData() = %v, want %v", got2, tt.want2) - } - }) - } -} - -func regexMapsEqual(a, b map[string][]*regexp.Regexp) bool { - if len(a) != len(b) { - return false - } - for method, aRegexes := range a { - bRegexes, ok := b[method] - if !ok || len(aRegexes) != len(bRegexes) { - return false - } - aRegexStrings := make([]string, 0, len(aRegexes)) - for _, ar := range aRegexes { - aRegexStrings = append(aRegexStrings, ar.String()) - } - bRegexStrings := make([]string, 0, len(bRegexes)) - for _, br := range bRegexes { - bRegexStrings = append(bRegexStrings, br.String()) - } - sort.Strings(aRegexStrings) - sort.Strings(bRegexStrings) - for i, ar := range aRegexStrings { - if ar != bRegexStrings[i] { - return false - } - } - } - return true -} - -func TestContainerName(t *testing.T) { - tests := []struct { - name string - cntr container.Summary - want string - }{ - { - name: "uses the first Docker container name", - cntr: container.Summary{ID: "0123456789abcdef", Names: []string{"/traefik", "/ignored"}}, - want: "traefik", - }, - { - name: "falls back to the short ID when Docker provides no name", - cntr: container.Summary{ID: "0123456789abcdef"}, - want: "0123456789ab", - }, - { - name: "does not panic for a short fallback ID", - cntr: container.Summary{ID: "short"}, - want: "short", - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - if got := containerName(tt.cntr); got != tt.want { - t.Errorf("containerName() = %q, want %q", got, tt.want) - } - }) - } -} - -func TestEventContainerName(t *testing.T) { - tests := []struct { - name string - event events.Message - want string - }{ - { - name: "uses the container name from event attributes", - event: events.Message{Actor: events.Actor{ID: "0123456789abcdef", Attributes: map[string]string{"name": "traefik"}}}, - want: "traefik", - }, - { - name: "falls back to the short ID when the event has no name", - event: events.Message{Actor: events.Actor{ID: "0123456789abcdef"}}, - want: "0123456789ab", - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - if got := eventContainerName(tt.event); got != tt.want { - t.Errorf("eventContainerName() = %q, want %q", got, tt.want) - } - }) + allowedDockerLabelPrefix = prevDockerLabelPrefix } } @@ -210,6 +48,38 @@ func TestInitConfig_AllowMethodFlagOverridesEnv(t *testing.T) { } } +func TestInitConfig_DockerLabelPrefixFlagOverridesEnv(t *testing.T) { + t.Setenv("SP_DOCKERLABELPREFIX", "from-env") + restore := resetFlagsForTest(t, []string{"socket-proxy", "-dockerlabelprefix=gameserver-socket-proxy"}) + defer restore() + + _, err := InitConfig() + if err != nil { + t.Fatalf("InitConfig() error = %v", err) + } + + if got, want := allowedDockerLabelPrefix, "gameserver-socket-proxy.allow."; got != want { + t.Errorf("allowedDockerLabelPrefix = %q, want %q", got, want) + } +} + +func TestInitConfig_InvalidDockerLabelPrefix(t *testing.T) { + for _, prefix := range []string{ + "Traefik", + "traefik..proxy", + "traefik-", + } { + t.Run(prefix, func(t *testing.T) { + restore := resetFlagsForTest(t, []string{"socket-proxy", "-dockerlabelprefix=" + prefix}) + defer restore() + + if _, err := InitConfig(); err == nil { + t.Fatalf("InitConfig() with dockerlabelprefix %q unexpectedly succeeded", prefix) + } + }) + } +} + func TestInitConfig_ShutdownGraceTimeTooLarge(t *testing.T) { restore := resetFlagsForTest(t, []string{ "socket-proxy", diff --git a/internal/config/dockerlabels.go b/internal/config/dockerlabels.go new file mode 100644 index 0000000..f4f720b --- /dev/null +++ b/internal/config/dockerlabels.go @@ -0,0 +1,502 @@ +package config + +import ( + "context" + "fmt" + "log/slog" + "regexp" + "slices" + "strings" + "sync" + "time" + + "github.com/wollomatic/socket-proxy/internal/docker/api/types/container" + "github.com/wollomatic/socket-proxy/internal/docker/api/types/events" + "github.com/wollomatic/socket-proxy/internal/docker/api/types/filters" + "github.com/wollomatic/socket-proxy/internal/docker/client" +) + +const ( + allowListsRefreshCooldown = 50 * time.Millisecond + allowListsRefreshRetryBackoff = 10 * time.Millisecond + defaultAllowListsRefreshTimeout = 10 * time.Second +) + +var allowedDockerLabelPrefix = defaultDockerLabelPrefix + ".allow." + +type allowListsRefreshState struct { + mutex sync.Mutex + inFlight *allowListsRefresh + completed time.Time + lastResult *AllowListRegistry + lastError error + timeout time.Duration // zero uses defaultAllowListsRefreshTimeout +} + +type allowListsRefresh struct { + done chan struct{} + allowLists *AllowListRegistry + err error +} + +type AllowListRegistry struct { + mutex sync.RWMutex // mutex to control read/write of byIP and revision + revision uint64 // generation used to detect updates during Docker snapshots + networks []string // names of networks in which socket proxy access is allowed for non-default allowlists + Default AllowList // default allowlist + byIP map[string]AllowList // map container IP address to allowlist for that container +} + +// UpdateAllowLists populates the byIP allowlists then keeps them updated. +func (cfg *Config) UpdateAllowLists() { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + dockerClient, err := client.NewClientWithOpts( + client.WithHost("unix://"+cfg.SocketPath), + client.WithAPIVersionNegotiation(), + ) + if err != nil { + slog.Error("failed to create Docker client", "error", err) + return + } + defer func(dockerClient *client.Client) { + err := dockerClient.Close() + if err != nil { + slog.Error("failed to close Docker client", "error", err) + } + }(dockerClient) + + err = cfg.AllowLists.initByIP(ctx, dockerClient) + if err != nil { + slog.Error("failed to initialise non-default allowlists", "error", err) + return + } + slog.Debug("initialised non-default allowlists") + + filter := filters.NewArgs() + filter.Add("type", "container") + filter.Add("event", "start") + filter.Add("event", "restart") + filter.Add("event", "die") + eventsChan, errChan := dockerClient.Events(ctx, events.ListOptions{Filters: filter}) + slog.Debug("subscribed to Docker event stream to update allowlists") + + // print non-default request allowlists + cfg.AllowLists.PrintByIP(cfg.LogJSON) + + // handle Docker events to update allowlists + for { + select { + case event, ok := <-eventsChan: + if !ok { + slog.Info("Docker event stream closed") + return + } + containerName := eventContainerName(event) + slog.Debug("received Docker container event", "action", event.Action, "container", containerName) + addedIPs, removedIPs, updateErr := cfg.AllowLists.updateFromEvent(ctx, dockerClient, event) + if updateErr != nil { + slog.Warn("failed to update allowlists from container event", "error", updateErr) + continue + } + for _, ip := range addedIPs { + cfg.AllowLists.mutex.RLock() + allowList, found := cfg.AllowLists.byIP[ip] + cfg.AllowLists.mutex.RUnlock() + if found { + allowList.Print(ip, cfg.LogJSON) + } + } + for _, ip := range removedIPs { + slog.Info("removed allowlist for container", "container", containerName, "ip", ip) + } + case err := <-errChan: + if err != nil { + slog.Error("received error from Docker event stream", "error", err) + return + } + } + } +} + +// RefreshAllowLists updates the per-container allowlists from Docker. Concurrent +// callers share an in-flight refresh, and recently completed results are reused +// to avoid repeatedly scanning Docker during bursts of rejected requests. +func (cfg *Config) RefreshAllowLists(ctx context.Context) (*AllowListRegistry, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + + state := &cfg.allowListsRefresh + state.mutex.Lock() + if refresh := state.inFlight; refresh != nil { + state.mutex.Unlock() + return waitForAllowListsRefresh(ctx, refresh) + } + if time.Since(state.completed) < allowListsRefreshCooldown { + allowLists, err := state.lastResult, state.lastError + state.mutex.Unlock() + return allowLists, err + } + refresh := &allowListsRefresh{done: make(chan struct{})} + state.inFlight = refresh + timeout := state.timeout + if timeout <= 0 { + timeout = defaultAllowListsRefreshTimeout + } + state.mutex.Unlock() + + refreshCtx, cancel := context.WithTimeout(context.Background(), timeout) + go func() { + defer cancel() + cfg.refreshAllowLists(refreshCtx, refresh) + }() + return waitForAllowListsRefresh(ctx, refresh) +} + +func (cfg *Config) refreshAllowLists(ctx context.Context, refresh *allowListsRefresh) { + dockerClient, err := client.NewClientWithOpts( + client.WithHost("unix://"+cfg.SocketPath), + client.WithAPIVersionNegotiation(), + ) + if err == nil { + err = cfg.AllowLists.initByIP(ctx, dockerClient) + if closeErr := dockerClient.Close(); closeErr != nil { + slog.Error("failed to close Docker client", "error", closeErr) + } + } + + state := &cfg.allowListsRefresh + state.mutex.Lock() + refresh.allowLists = cfg.AllowLists + refresh.err = err + state.lastResult = refresh.allowLists + state.lastError = err + state.completed = time.Now() + state.inFlight = nil + close(refresh.done) + state.mutex.Unlock() +} + +func waitForAllowListsRefresh(ctx context.Context, refresh *allowListsRefresh) (*AllowListRegistry, error) { + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-refresh.done: + return refresh.allowLists, refresh.err + } +} + +// PrintNetworks prints the allowed networks. +func (allowLists *AllowListRegistry) PrintNetworks() { + if len(allowLists.networks) > 0 { + slog.Info("socket proxy networks detected", "socketproxynetworks", allowLists.networks) + } else { + // we only log this on DEBUG level because the socket proxy networks are used for per-container allowlists + slog.Debug("no socket proxy networks detected") + } +} + +// PrintDefault prints the default allowlist. +func (allowLists *AllowListRegistry) PrintDefault(logJSON bool) { + allowLists.Default.Print("", logJSON) +} + +// PrintByIP prints the non-default allowlists. +func (allowLists *AllowListRegistry) PrintByIP(logJSON bool) { + allowLists.mutex.RLock() + defer allowLists.mutex.RUnlock() + for ip, allowList := range allowLists.byIP { + allowList.Print(ip, logJSON) + } +} + +// FindByIP returns the allowlist corresponding to the given IP address if found. +func (allowLists *AllowListRegistry) FindByIP(ip string) (AllowList, bool) { + allowLists.mutex.RLock() + defer allowLists.mutex.RUnlock() + allowList, found := allowLists.byIP[ip] + return allowList, found +} + +// initialise allowlist registry byIP allowlists +func (allowLists *AllowListRegistry) initByIP(ctx context.Context, dockerClient *client.Client) error { + filter := filters.NewArgs() + for _, network := range allowLists.networks { + filter.Add("network", network) + } + + for { + if err := ctx.Err(); err != nil { + return err + } + + allowLists.mutex.RLock() + snapshotRevision := allowLists.revision + allowLists.mutex.RUnlock() + + containers, err := dockerClient.ContainerList(ctx, container.ListOptions{Filters: filter}) + if err != nil { + return err + } + + byIP, err := buildByIP(containers, allowLists.networks) + if err != nil { + return err + } + + allowLists.mutex.Lock() + if allowLists.revision == snapshotRevision { + allowLists.byIP = byIP + allowLists.revision++ + allowLists.mutex.Unlock() + return nil + } + allowLists.mutex.Unlock() + + select { + case <-ctx.Done(): + return ctx.Err() + case <-time.After(allowListsRefreshRetryBackoff): + } + } +} + +func buildByIP(containers []container.Summary, networks []string) (map[string]AllowList, error) { + byIP := make(map[string]AllowList) + for _, cntr := range containers { + allowedRequests, allowedBindMounts, err := extractLabelData(cntr) + if err != nil { + return nil, err + } + + if len(allowedRequests) > 0 || len(allowedBindMounts) > 0 { + for networkID, cntrNetwork := range cntr.NetworkSettings.Networks { + if slices.Contains(networks, networkID) { + allowList := AllowList{ + ID: cntr.ID, + ContainerName: containerName(cntr), + AllowedRequests: allowedRequests, + AllowedBindMounts: allowedBindMounts, + } + + if len(cntrNetwork.IPAddress) > 0 { + byIP[cntrNetwork.IPAddress] = allowList + } + if len(cntrNetwork.GlobalIPv6Address) > 0 { + byIP[cntrNetwork.GlobalIPv6Address] = allowList + } + } + } + } + } + + return byIP, nil +} + +// update the allowlist registry based on the Docker event +func (allowLists *AllowListRegistry) updateFromEvent( + ctx context.Context, dockerClient *client.Client, event events.Message, +) ([]string, []string, error) { + allowLists.mutex.Lock() + allowLists.revision++ + allowLists.mutex.Unlock() + + containerID := event.Actor.ID + var ( + addedIPs []string + removedIPs []string + err error + ) + + switch event.Action { + case "start", "restart": + addedIPs, err = allowLists.add(ctx, dockerClient, containerID) + if err != nil { + return nil, nil, err + } + case "die": + removedIPs = allowLists.remove(containerID) + } + return addedIPs, removedIPs, nil +} + +// add the allowlist for the container with the given ID to the allowlist registry +// if it has at least one socket-proxy allow label and is in a same network as the socket-proxy +func (allowLists *AllowListRegistry) add( + ctx context.Context, dockerClient *client.Client, containerID string, +) ([]string, error) { + filter := filters.NewArgs() + filter.Add("id", containerID) + for _, network := range allowLists.networks { + filter.Add("network", network) + } + containers, err := dockerClient.ContainerList(ctx, container.ListOptions{Filters: filter}) + if err != nil { + return nil, err + } + if len(containers) == 0 { + slog.Debug("container is not in a network with socket-proxy or may have stopped", "id", shortContainerID(containerID)) + return nil, nil + } + cntr := containers[0] + + allowedRequests, allowedBindMounts, err := extractLabelData(cntr) + if err != nil { + return nil, err + } + + var ips []string + if len(allowedRequests) > 0 || len(allowedBindMounts) > 0 { + allowList := AllowList{ + ID: cntr.ID, + ContainerName: containerName(cntr), + AllowedRequests: allowedRequests, + AllowedBindMounts: allowedBindMounts, + } + + allowLists.mutex.Lock() + defer allowLists.mutex.Unlock() + + if allowLists.byIP == nil { + allowLists.byIP = make(map[string]AllowList) + } + + for networkID, cntrNetwork := range cntr.NetworkSettings.Networks { + if slices.Contains(allowLists.networks, networkID) { + ipv4Address := cntrNetwork.IPAddress + if len(ipv4Address) > 0 { + allowLists.byIP[ipv4Address] = allowList + ips = append(ips, ipv4Address) + } + ipv6Address := cntrNetwork.GlobalIPv6Address + if len(ipv6Address) > 0 { + allowLists.byIP[ipv6Address] = allowList + ips = append(ips, ipv6Address) + } + } + } + } + + return ips, nil +} + +// remove allowlists having the given container ID from the allowlist registry +func (allowLists *AllowListRegistry) remove(containerID string) []string { + allowLists.mutex.Lock() + defer allowLists.mutex.Unlock() + + var removedIPs []string + for ip, allowList := range allowLists.byIP { + if allowList.ID == containerID { + delete(allowLists.byIP, ip) + removedIPs = append(removedIPs, ip) + } + } + return removedIPs +} + +// containerName returns Docker's container name without its leading slash. +// It falls back to the short container ID for unusual responses without a name. +func containerName(cntr container.Summary) string { + for _, name := range cntr.Names { + if name = strings.TrimPrefix(name, "/"); name != "" { + return name + } + } + return shortContainerID(cntr.ID) +} + +// eventContainerName returns the name Docker includes with container events. +// It falls back to the short container ID when the event has no name. +func eventContainerName(event events.Message) string { + if name := event.Actor.Attributes["name"]; name != "" { + return name + } + return shortContainerID(event.Actor.ID) +} + +func shortContainerID(id string) string { + if len(id) > 12 { + return id[:12] + } + return id +} + +// return list of docker networks that the socket proxy container is in +func listSocketProxyNetworks(socketPath, proxyContainerName string) ([]string, error) { + cntr, err := getSocketProxyContainerSummary(socketPath, proxyContainerName) + if err != nil { + return nil, err + } + + networks := make([]string, 0, len(cntr.NetworkSettings.Networks)) + for networkID := range cntr.NetworkSettings.Networks { + networks = append(networks, networkID) + } + return networks, nil +} + +// return Docker container summary for the socket proxy container +func getSocketProxyContainerSummary(socketPath, proxyContainerName string) (container.Summary, error) { + const maxTries = 3 + + dockerClient, err := client.NewClientWithOpts( + client.WithHost("unix://"+socketPath), + client.WithAPIVersionNegotiation(), + ) + if err != nil { + return container.Summary{}, err + } + defer func(dockerClient *client.Client) { + err := dockerClient.Close() + if err != nil { + slog.Error("failed to close Docker client", "error", err) + } + }(dockerClient) + + ctx := context.Background() + filter := filters.NewArgs() + filter.Add("name", proxyContainerName) + var containers []container.Summary + for i := 1; i <= maxTries; i++ { + containers, err = dockerClient.ContainerList(ctx, container.ListOptions{Filters: filter}) + if err != nil { + return container.Summary{}, err + } + if len(containers) > 0 { + return containers[0], nil + } + if i < maxTries { + time.Sleep(time.Duration(i) * time.Second) + } + } + return container.Summary{}, fmt.Errorf("socket-proxy container \"%s\" was not found after %d attempts; verify the container name is correct and the container is running", proxyContainerName, maxTries) +} + +// extract Docker container allowlist label data from the container summary +func extractLabelData(cntr container.Summary) (map[string][]*regexp.Regexp, []string, error) { + allowedRequests := make(map[string][]*regexp.Regexp) + var allowedBindMounts []string + for labelName, labelValue := range cntr.Labels { + if strings.HasPrefix(labelName, allowedDockerLabelPrefix) && labelValue != "" { + allowSpec := strings.ToUpper(strings.TrimPrefix(labelName, allowedDockerLabelPrefix)) + method, _, _ := strings.Cut(allowSpec, ".") + if slices.Contains(supportedHTTPMethods, method) { + r, err := compileRegexp(labelValue, method, "docker container label") + if err != nil { + return nil, nil, err + } + allowedRequests[method] = append(allowedRequests[method], r) + } else if allowSpec == "BINDMOUNTFROM" { + var err error + allowedBindMounts, err = parseAllowedBindMounts(labelValue) + if err != nil { + return nil, nil, err + } + } + } + } + return allowedRequests, allowedBindMounts, nil +} diff --git a/internal/config/dockerlabels_test.go b/internal/config/dockerlabels_test.go new file mode 100644 index 0000000..04dddf3 --- /dev/null +++ b/internal/config/dockerlabels_test.go @@ -0,0 +1,617 @@ +package config + +import ( + "context" + "encoding/json" + "errors" + "net" + "net/http" + "path/filepath" + "reflect" + "regexp" + "sort" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/wollomatic/socket-proxy/internal/docker/api/types/container" + "github.com/wollomatic/socket-proxy/internal/docker/api/types/events" + "github.com/wollomatic/socket-proxy/internal/docker/api/types/network" +) + +func Test_extractLabelData(t *testing.T) { + tests := []struct { + name string // description of this test case + // Named input parameters for target function. + cntr container.Summary + prefix string + want map[string][]*regexp.Regexp + want2 []string + wantErr bool + }{ + { + name: "valid labels with multiple methods and regexes", + cntr: container.Summary{ + Labels: map[string]string{ + "socket-proxy.allow.get.0": "regex1", + "socket-proxy.allow.get.1": "regex2", + "socket-proxy.allow.post": "regex3", + }, + }, + want: map[string][]*regexp.Regexp{ + "GET": {regexp.MustCompile("^regex1$"), regexp.MustCompile("^regex2$")}, + "POST": {regexp.MustCompile("^regex3$")}, + }, + want2: nil, + wantErr: false, + }, + { + name: "invalid regex in label value", + cntr: container.Summary{ + Labels: map[string]string{ + "socket-proxy.allow.get": "invalid[regex", + }, + }, + want: nil, + want2: nil, + wantErr: true, + }, + { + name: "custom label prefix ignores the default prefix", + cntr: container.Summary{ + Labels: map[string]string{ + "socket-proxy.allow.get": "default", + "gameserver-socket-proxy.allow.get": "custom", + }, + }, + prefix: "gameserver-socket-proxy", + want: map[string][]*regexp.Regexp{ + "GET": {regexp.MustCompile("^custom$")}, + }, + }, + { + name: "non-allow labels are ignored", + cntr: container.Summary{ + Labels: map[string]string{ + "socket-proxy.allow.get": "regex1", + "other.label": "value", + }, + }, + want: map[string][]*regexp.Regexp{ + "GET": {regexp.MustCompile("^regex1$")}, + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + previousPrefix := allowedDockerLabelPrefix + defer func() { allowedDockerLabelPrefix = previousPrefix }() + allowedDockerLabelPrefix = defaultDockerLabelPrefix + ".allow." + if tt.prefix != "" { + allowedDockerLabelPrefix = tt.prefix + ".allow." + } + got, got2, gotErr := extractLabelData(tt.cntr) + if gotErr != nil { + if !tt.wantErr { + t.Errorf("extractLabelData() failed: %v", gotErr) + } + return + } + if tt.wantErr { + t.Fatal("extractLabelData() succeeded unexpectedly") + } + if !regexMapsEqual(got, tt.want) { + t.Errorf("extractLabelData() = %v, want %v", got, tt.want) + } + if !reflect.DeepEqual(got2, tt.want2) { + t.Errorf("extractLabelData() = %v, want %v", got2, tt.want2) + } + }) + } +} + +func TestContainerName(t *testing.T) { + tests := []struct { + name string + cntr container.Summary + want string + }{ + { + name: "uses the first Docker container name", + cntr: container.Summary{ID: "0123456789abcdef", Names: []string{"/traefik", "/ignored"}}, + want: "traefik", + }, + { + name: "falls back to the short ID when Docker provides no name", + cntr: container.Summary{ID: "0123456789abcdef"}, + want: "0123456789ab", + }, + { + name: "does not panic for a short fallback ID", + cntr: container.Summary{ID: "short"}, + want: "short", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := containerName(tt.cntr); got != tt.want { + t.Errorf("containerName() = %q, want %q", got, tt.want) + } + }) + } +} + +func TestEventContainerName(t *testing.T) { + tests := []struct { + name string + event events.Message + want string + }{ + { + name: "uses the container name from event attributes", + event: events.Message{Actor: events.Actor{ID: "0123456789abcdef", Attributes: map[string]string{"name": "traefik"}}}, + want: "traefik", + }, + { + name: "falls back to the short ID when the event has no name", + event: events.Message{Actor: events.Actor{ID: "0123456789abcdef"}}, + want: "0123456789ab", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := eventContainerName(tt.event); got != tt.want { + t.Errorf("eventContainerName() = %q, want %q", got, tt.want) + } + }) + } +} + +func regexMapsEqual(a, b map[string][]*regexp.Regexp) bool { + if len(a) != len(b) { + return false + } + for method, aRegexes := range a { + bRegexes, ok := b[method] + if !ok || len(aRegexes) != len(bRegexes) { + return false + } + aRegexStrings := make([]string, 0, len(aRegexes)) + for _, ar := range aRegexes { + aRegexStrings = append(aRegexStrings, ar.String()) + } + bRegexStrings := make([]string, 0, len(bRegexes)) + for _, br := range bRegexes { + bRegexStrings = append(bRegexStrings, br.String()) + } + sort.Strings(aRegexStrings) + sort.Strings(bRegexStrings) + for i, ar := range aRegexStrings { + if ar != bRegexStrings[i] { + return false + } + } + } + return true +} + +func TestRefreshAllowLists(t *testing.T) { + socketPath := filepath.Join(t.TempDir(), "docker.sock") + listener, err := (&net.ListenConfig{}).Listen(context.Background(), "unix", socketPath) + if err != nil { + t.Fatalf("listen on Docker test socket: %v", err) + } + t.Cleanup(func() { _ = listener.Close() }) + + go func() { + _ = http.Serve(listener, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/_ping": + w.Header().Set("Api-Version", "1.51") + case "/v1.51/containers/json": + if err := json.NewEncoder(w).Encode([]container.Summary{{ + ID: "container-id", + Labels: map[string]string{ + "socket-proxy.allow.get": "/version", + }, + NetworkSettings: &container.NetworkSettingsSummary{Networks: map[string]*network.EndpointSettings{ + "proxy-network": {IPAddress: "172.20.0.2"}, + }}, + }}); err != nil { + t.Errorf("encode container list: %v", err) + } + default: + http.NotFound(w, r) + } + })) + }() + + cfg := Config{ + SocketPath: socketPath, + AllowLists: &AllowListRegistry{ + networks: []string{"proxy-network"}, + }, + } + refreshedAllowLists, err := cfg.RefreshAllowLists(context.Background()) + if err != nil { + t.Fatalf("RefreshAllowLists() error = %v", err) + } + + allowList, found := refreshedAllowLists.FindByIP("172.20.0.2") + if !found { + t.Fatal("allowlist was not added for container IP") + } + if !matchAny(allowList.AllowedRequests[http.MethodGet], "/version") { + t.Fatal("refreshed allowlist does not contain the container label") + } +} + +func TestRefreshAllowListsErrorPreservesRegistry(t *testing.T) { + socketPath := filepath.Join(t.TempDir(), "docker.sock") + listener, err := (&net.ListenConfig{}).Listen(context.Background(), "unix", socketPath) + if err != nil { + t.Fatalf("listen on Docker test socket: %v", err) + } + t.Cleanup(func() { _ = listener.Close() }) + + go func() { + _ = http.Serve(listener, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/_ping": + w.Header().Set("Api-Version", "1.51") + case "/v1.51/containers/json": + if err := json.NewEncoder(w).Encode([]container.Summary{{ + ID: "container-id", + Labels: map[string]string{ + "socket-proxy.allow.get": "invalid[regex", + }, + NetworkSettings: &container.NetworkSettingsSummary{Networks: map[string]*network.EndpointSettings{ + "proxy-network": {IPAddress: "172.20.0.3"}, + }}, + }}); err != nil { + t.Errorf("encode container list: %v", err) + } + default: + http.NotFound(w, r) + } + })) + }() + + const initialRevision = 7 + cfg := Config{ + SocketPath: socketPath, + AllowLists: &AllowListRegistry{ + revision: initialRevision, + networks: []string{"proxy-network"}, + byIP: map[string]AllowList{ + "172.20.0.2": {ID: "existing-container-id"}, + }, + }, + } + if _, err := cfg.RefreshAllowLists(context.Background()); err == nil { + t.Fatal("RefreshAllowLists() unexpectedly succeeded") + } + + allowList, found := cfg.AllowLists.FindByIP("172.20.0.2") + if !found || allowList.ID != "existing-container-id" { + t.Fatal("RefreshAllowLists() modified the existing allowlist after an error") + } + cfg.AllowLists.mutex.RLock() + revision := cfg.AllowLists.revision + cfg.AllowLists.mutex.RUnlock() + if revision != initialRevision { + t.Fatalf("revision = %d, want %d", revision, initialRevision) + } +} + +func TestRefreshDoesNotOverwriteNewerEventUpdate(t *testing.T) { + socketPath := filepath.Join(t.TempDir(), "docker.sock") + listener, err := (&net.ListenConfig{}).Listen(context.Background(), "unix", socketPath) + if err != nil { + t.Fatalf("listen on Docker test socket: %v", err) + } + t.Cleanup(func() { _ = listener.Close() }) + + requestStarted := make(chan struct{}) + releaseRequest := make(chan struct{}) + var ( + containerListRequests atomic.Int32 + startOnce sync.Once + releaseOnce sync.Once + ) + release := func() { + releaseOnce.Do(func() { close(releaseRequest) }) + } + t.Cleanup(release) + + go func() { + _ = http.Serve(listener, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/_ping": + w.Header().Set("Api-Version", "1.51") + case "/v1.51/containers/json": + if containerListRequests.Add(1) != 1 { + if err := json.NewEncoder(w).Encode([]container.Summary{}); err != nil { + t.Errorf("encode container list: %v", err) + } + return + } + startOnce.Do(func() { close(requestStarted) }) + <-releaseRequest + if err := json.NewEncoder(w).Encode([]container.Summary{{ + ID: "stopped-container-id", + Labels: map[string]string{ + "socket-proxy.allow.get": "/version", + }, + NetworkSettings: &container.NetworkSettingsSummary{Networks: map[string]*network.EndpointSettings{ + "proxy-network": {IPAddress: "172.20.0.2"}, + }}, + }}); err != nil { + t.Errorf("encode container list: %v", err) + } + default: + http.NotFound(w, r) + } + })) + }() + + cfg := Config{ + SocketPath: socketPath, + AllowLists: &AllowListRegistry{ + networks: []string{"proxy-network"}, + byIP: map[string]AllowList{ + "172.20.0.2": {ID: "stopped-container-id"}, + }, + }, + } + refreshResult := make(chan error, 1) + go func() { + _, refreshErr := cfg.RefreshAllowLists(context.Background()) + refreshResult <- refreshErr + }() + + select { + case <-requestStarted: + case <-time.After(time.Second): + t.Fatal("timed out waiting for Docker refresh to start") + } + + _, removedIPs, err := cfg.AllowLists.updateFromEvent(context.Background(), nil, events.Message{ + Action: events.ActionDie, + Actor: events.Actor{ID: "stopped-container-id"}, + }) + if err != nil { + t.Fatalf("updateFromEvent() error = %v", err) + } + if !reflect.DeepEqual(removedIPs, []string{"172.20.0.2"}) { + t.Fatalf("removed IPs = %v, want [172.20.0.2]", removedIPs) + } + + release() + select { + case err := <-refreshResult: + if err != nil { + t.Fatalf("RefreshAllowLists() error = %v", err) + } + case <-time.After(time.Second): + t.Fatal("timed out waiting for Docker refresh retry") + } + if got := containerListRequests.Load(); got != 2 { + t.Fatalf("Docker container list requests = %d, want 2", got) + } + if _, found := cfg.AllowLists.FindByIP("172.20.0.2"); found { + t.Fatal("refresh restored the allowlist removed by a newer event") + } +} + +func TestRefreshRevisionRetriesRespectTimeout(t *testing.T) { + socketPath := filepath.Join(t.TempDir(), "docker.sock") + listener, err := (&net.ListenConfig{}).Listen(context.Background(), "unix", socketPath) + if err != nil { + t.Fatalf("listen on Docker test socket: %v", err) + } + t.Cleanup(func() { _ = listener.Close() }) + + const refreshTimeout = 50 * time.Millisecond + var containerListRequests atomic.Int32 + cfg := Config{ + SocketPath: socketPath, + AllowLists: &AllowListRegistry{ + byIP: map[string]AllowList{ + "172.20.0.2": {ID: "existing-container-id"}, + }, + }, + allowListsRefresh: allowListsRefreshState{ + timeout: refreshTimeout, + }, + } + + go func() { + _ = http.Serve(listener, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/_ping": + w.Header().Set("Api-Version", "1.51") + case "/v1.51/containers/json": + containerListRequests.Add(1) + cfg.AllowLists.mutex.Lock() + cfg.AllowLists.revision++ + cfg.AllowLists.mutex.Unlock() + if err := json.NewEncoder(w).Encode([]container.Summary{}); err != nil { + t.Errorf("encode container list: %v", err) + } + default: + http.NotFound(w, r) + } + })) + }() + + if _, err := cfg.RefreshAllowLists(context.Background()); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("RefreshAllowLists() error = %v, want context.DeadlineExceeded", err) + } + maxRequests := int32(refreshTimeout/allowListsRefreshRetryBackoff) + 1 + if got := containerListRequests.Load(); got > maxRequests { + t.Fatalf("Docker container list requests = %d, want at most %d", got, maxRequests) + } + if _, found := cfg.AllowLists.FindByIP("172.20.0.2"); !found { + t.Fatal("timed-out refresh replaced the existing allowlist") + } +} + +func TestRefreshCoalescing(t *testing.T) { + socketPath := filepath.Join(t.TempDir(), "docker.sock") + listener, err := (&net.ListenConfig{}).Listen(context.Background(), "unix", socketPath) + if err != nil { + t.Fatalf("listen on Docker test socket: %v", err) + } + t.Cleanup(func() { _ = listener.Close() }) + + var ( + containerListRequests atomic.Int32 + requestStarted = make(chan struct{}) + releaseRequest = make(chan struct{}) + startOnce sync.Once + releaseOnce sync.Once + ) + release := func() { + releaseOnce.Do(func() { close(releaseRequest) }) + } + t.Cleanup(release) + + go func() { + _ = http.Serve(listener, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/_ping": + w.Header().Set("Api-Version", "1.51") + case "/v1.51/containers/json": + containerListRequests.Add(1) + startOnce.Do(func() { close(requestStarted) }) + <-releaseRequest + http.Error(w, "Docker unavailable", http.StatusServiceUnavailable) + default: + http.NotFound(w, r) + } + })) + }() + + cfg := Config{ + SocketPath: socketPath, + AllowLists: &AllowListRegistry{}, + } + firstResult := make(chan error, 1) + go func() { + _, refreshErr := cfg.RefreshAllowLists(context.Background()) + firstResult <- refreshErr + }() + + select { + case <-requestStarted: + case <-time.After(time.Second): + t.Fatal("timed out waiting for Docker refresh to start") + } + + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond) + defer cancel() + if _, err := cfg.RefreshAllowLists(ctx); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("RefreshAllowLists() error = %v, want context.DeadlineExceeded", err) + } + + const callers = 8 + results := make(chan error, callers) + for range callers { + go func() { + _, refreshErr := cfg.RefreshAllowLists(context.Background()) + results <- refreshErr + }() + } + release() + + if err := <-firstResult; err == nil { + t.Fatal("first RefreshAllowLists() unexpectedly succeeded") + } + for range callers { + if err := <-results; err == nil { + t.Fatal("shared RefreshAllowLists() unexpectedly succeeded") + } + } + if got := containerListRequests.Load(); got != 1 { + t.Fatalf("Docker container list requests = %d, want 1", got) + } + + if _, err := cfg.RefreshAllowLists(context.Background()); err == nil { + t.Fatal("cached RefreshAllowLists() unexpectedly succeeded") + } + if got := containerListRequests.Load(); got != 1 { + t.Fatalf("Docker container list requests during cooldown = %d, want 1", got) + } +} + +func TestRefreshTimeoutClearsInFlightRefresh(t *testing.T) { + socketPath := filepath.Join(t.TempDir(), "docker.sock") + listener, err := (&net.ListenConfig{}).Listen(context.Background(), "unix", socketPath) + if err != nil { + t.Fatalf("listen on Docker test socket: %v", err) + } + t.Cleanup(func() { _ = listener.Close() }) + + var containerListRequests atomic.Int32 + firstRequestCanceled := make(chan struct{}) + go func() { + _ = http.Serve(listener, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/_ping": + w.Header().Set("Api-Version", "1.51") + case "/v1.51/containers/json": + if containerListRequests.Add(1) == 1 { + <-r.Context().Done() + close(firstRequestCanceled) + return + } + if err := json.NewEncoder(w).Encode([]container.Summary{}); err != nil { + t.Errorf("encode container list: %v", err) + } + default: + http.NotFound(w, r) + } + })) + }() + + cfg := Config{ + SocketPath: socketPath, + AllowLists: &AllowListRegistry{}, + allowListsRefresh: allowListsRefreshState{ + timeout: 100 * time.Millisecond, + }, + } + if _, err := cfg.RefreshAllowLists(context.Background()); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("RefreshAllowLists() error = %v, want context.DeadlineExceeded", err) + } + select { + case <-firstRequestCanceled: + case <-time.After(time.Second): + t.Fatal("timed out waiting for Docker request cancellation") + } + + // Bypass the result cooldown so the retry exercises the cleared in-flight state. + cfg.allowListsRefresh.mutex.Lock() + cfg.allowListsRefresh.completed = time.Time{} + cfg.allowListsRefresh.mutex.Unlock() + + if _, err := cfg.RefreshAllowLists(context.Background()); err != nil { + t.Fatalf("retry RefreshAllowLists() error = %v", err) + } + if got := containerListRequests.Load(); got != 2 { + t.Fatalf("Docker container list requests = %d, want 2", got) + } +} + +func matchAny(regexes []*regexp.Regexp, value string) bool { + for _, regex := range regexes { + if regex.MatchString(value) { + return true + } + } + return false +}