From 769688052303d96b304452dcdc646faca0e3f802 Mon Sep 17 00:00:00 2001 From: Mason Wheeler Date: Wed, 23 Sep 2026 10:41:26 -0700 Subject: [PATCH] feat(sfcompute): opt in to public IPs and configurable ingress --- v1/providers/sfcomputev2/README.md | 35 +++ v1/providers/sfcomputev2/api_client.go | 55 ++-- v1/providers/sfcomputev2/api_client_test.go | 10 +- v1/providers/sfcomputev2/brev_constants.go | 4 +- v1/providers/sfcomputev2/capabilities.go | 26 +- v1/providers/sfcomputev2/client.go | 35 +-- v1/providers/sfcomputev2/firewall.go | 236 +++++++++++++++++ .../sfcomputev2/firewall_cleanup_test.go | 142 +++++++++++ v1/providers/sfcomputev2/firewall_test.go | 238 ++++++++++++++++++ v1/providers/sfcomputev2/instance.go | 65 ++++- v1/providers/sfcomputev2/instancetype.go | 8 + 11 files changed, 802 insertions(+), 52 deletions(-) create mode 100644 v1/providers/sfcomputev2/README.md create mode 100644 v1/providers/sfcomputev2/firewall.go create mode 100644 v1/providers/sfcomputev2/firewall_cleanup_test.go create mode 100644 v1/providers/sfcomputev2/firewall_test.go diff --git a/v1/providers/sfcomputev2/README.md b/v1/providers/sfcomputev2/README.md new file mode 100644 index 0000000..655ea01 --- /dev/null +++ b/v1/providers/sfcomputev2/README.md @@ -0,0 +1,35 @@ +# SFCompute V2 networking + +Configurable ingress is optional. Existing credentials continue using the SSH +proxy and do not create firewalls. To request public IPv4 for new instances: + +```go +credential := v2.NewSFCCredentialV2(refID, apiKey, organization, workspace) +credential.EnableConfigurableFirewall = true +``` + +The equivalent JSON field is `"enable_configurable_firewall": true`. Enable it +only after the SFCompute migration, API, and capacity reconciler have finished +deploying. Roll out this Brev integration afterward. The pool must advertise +`public_ipv4_skus`; the adapter selects only those SKUs. Missing support returns +an error before creating an instance or firewall. + +The API key needs firewall read, create, write, and delete permissions in the +pool's workspace, in addition to the existing instance and pool permissions. +Each new instance gets its own firewall. TCP ports 22 and 2222 remain open for +SSH; additional port ranges apply to TCP and UDP. Sources must be public IPv4 +CIDRs or `0.0.0.0/0`. Narrowing outbound traffic is unsupported. SFCompute's +firewall quotas apply, including the limit of 99 custom firewalls per workspace. + +`GetInstance` and `ListInstances` return the public IP, SSH port 22, and stable +ingress rule IDs for revocation. Firewall edits preserve unrelated rules and +retry concurrent changes through `/integrations/brev/v1/firewalls`. Rule updates +require the version returned by that integration endpoint. Termination deletes +the attached firewall only when its ID matches the instance's Brev ownership tag. +A cleanup failure is returned to the caller; retrying termination retries +firewall deletion. Cleanup uses a separate bounded context +so cancellation of a create request does not cancel its firewall cleanup. + +Turning the credential option off affects new instances. Existing public-IP +instances remain accessible and their firewalls can still be managed. Existing +SSH-proxy instances cannot acquire public networking through a rule update. diff --git a/v1/providers/sfcomputev2/api_client.go b/v1/providers/sfcomputev2/api_client.go index 922ddaa..fd81b9a 100644 --- a/v1/providers/sfcomputev2/api_client.go +++ b/v1/providers/sfcomputev2/api_client.go @@ -31,6 +31,8 @@ type createInstanceRequest struct { CloudInitUserData *string `json:"cloud_init_user_data,omitempty"` Tags map[string]string `json:"tags,omitempty"` PreviewEnableInfiniband bool `json:"_preview_enable_infiniband"` + EnablePublicIPv4 bool `json:"enable_public_ipv4,omitempty"` + Firewall string `json:"firewall,omitempty"` } type instanceStatus string @@ -47,12 +49,15 @@ type instanceSKUSummary struct { } type instanceResponse struct { - ID string `json:"id"` - Name string `json:"name"` - Status instanceStatus `json:"status"` - InstanceSKU *instanceSKUSummary `json:"instance_sku"` - CreatedAt int64 `json:"created_at"` - Tags map[string]string `json:"tags"` + ID string `json:"id"` + Name string `json:"name"` + Status instanceStatus `json:"status"` + InstanceSKU *instanceSKUSummary `json:"instance_sku"` + CreatedAt int64 `json:"created_at"` + Tags map[string]string `json:"tags"` + EnablePublicIPv4 bool `json:"enable_public_ipv4"` + Firewall string `json:"firewall"` + PublicIP string `json:"public_ip"` } type listInstancesResponse struct { @@ -92,6 +97,7 @@ type allocationSchedule struct { type poolResponse struct { AllocationSchedule allocationSchedule `json:"allocation_schedule"` + PublicIPv4SKUs *[]string `json:"public_ipv4_skus,omitempty"` } type apiError struct { @@ -153,8 +159,16 @@ func (c *apiClient) listInstances(ctx context.Context, workspace, pool string) ( } func (c *apiClient) terminateInstance(ctx context.Context, id string) error { + _, err := c.terminateInstanceWithResponse(ctx, id) + return err +} + +func (c *apiClient) terminateInstanceWithResponse(ctx context.Context, id string) (*instanceResponse, error) { var response instanceResponse - return c.do(ctx, http.MethodPost, "/instances/"+url.PathEscape(id)+"/terminate", nil, nil, &response) + if err := c.do(ctx, http.MethodPost, "/instances/"+url.PathEscape(id)+"/terminate", nil, nil, &response); err != nil { + return nil, err + } + return &response, nil } func (c *apiClient) getSSHInfo(ctx context.Context, id string) (*instanceSSHInfo, error) { @@ -181,11 +195,19 @@ func (c *apiClient) do( requestBody any, responseBody any, ) error { + _, err := c.doRequest(ctx, method, brevAPIPath+path, query, requestBody, responseBody, nil) + return err +} + +func (c *apiClient) doRequest( + ctx context.Context, method, path string, query url.Values, + requestBody, responseBody any, headers http.Header, +) (http.Header, error) { var body io.Reader if requestBody != nil { encoded, err := json.Marshal(requestBody) if err != nil { - return err + return nil, err } body = bytes.NewReader(encoded) } @@ -193,34 +215,37 @@ func (c *apiClient) do( request, err := http.NewRequestWithContext( ctx, method, - strings.TrimRight(c.baseURL, "/")+brevAPIPath+path, + strings.TrimRight(c.baseURL, "/")+path, body, ) if err != nil { - return err + return nil, err } request.URL.RawQuery = query.Encode() request.Header.Set("Accept", "application/json") request.Header.Set("Authorization", "Bearer "+c.apiKey) + for key, values := range headers { + request.Header[key] = values + } if requestBody != nil { request.Header.Set("Content-Type", "application/json") } response, err := c.httpClient.Do(request) if err != nil { - return err + return nil, err } defer func() { _ = response.Body.Close() }() responseBytes, err := io.ReadAll(response.Body) if err != nil { - return err + return nil, err } if response.StatusCode < http.StatusOK || response.StatusCode >= http.StatusMultipleChoices { - return &apiError{statusCode: response.StatusCode, body: string(responseBytes)} + return response.Header, &apiError{statusCode: response.StatusCode, body: string(responseBytes)} } if responseBody == nil || len(responseBytes) == 0 { - return nil + return response.Header, nil } - return json.Unmarshal(responseBytes, responseBody) + return response.Header, json.Unmarshal(responseBytes, responseBody) } diff --git a/v1/providers/sfcomputev2/api_client_test.go b/v1/providers/sfcomputev2/api_client_test.go index 8e13315..eb44fb2 100644 --- a/v1/providers/sfcomputev2/api_client_test.go +++ b/v1/providers/sfcomputev2/api_client_test.go @@ -17,7 +17,7 @@ func TestAPIClientUsesBrevContract(t *testing.T) { require.Equal(t, "Bearer api-key", request.Header.Get("Authorization")) switch request.Method + " " + request.URL.Path { - case "POST /integrations/brev/v1/instances": + case createInstanceRoute: var body map[string]any require.NoError(t, json.NewDecoder(request.Body).Decode(&body)) require.Equal(t, "sfc:pool:account:workspace:default", body["pool"]) @@ -29,7 +29,7 @@ func TestAPIClientUsesBrevContract(t *testing.T) { require.Equal(t, "brev-ref", tags[tagKeyRefID]) require.Equal(t, false, body["_preview_enable_infiniband"]) writeJSON(t, writer, instanceResponse{ID: "inst_created", Status: instanceStatusAwaitingAllocation}) - case "GET /integrations/brev/v1/instances": + case listInstancesRoute: require.Equal(t, "sfc:workspace:account:workspace", request.URL.Query().Get("workspace")) require.Equal(t, []string{"sfc:pool:account:workspace:default"}, request.URL.Query()["pool"]) require.Equal(t, "200", request.URL.Query().Get("limit")) @@ -43,13 +43,13 @@ func TestAPIClientUsesBrevContract(t *testing.T) { } require.Equal(t, "next-page", request.URL.Query().Get("starting_after")) writeJSON(t, writer, listInstancesResponse{Data: []instanceResponse{{ID: "inst_listed_2"}}}) - case "GET /integrations/brev/v1/instances/inst_test": + case getTestInstanceRoute: writeJSON(t, writer, instanceResponse{ID: "inst_test", Status: instanceStatusRunning}) case "GET /integrations/brev/v1/instances/inst_test/ssh": writeJSON(t, writer, instanceSSHInfo{Hostname: "192.0.2.1", Port: 22}) - case "POST /integrations/brev/v1/instances/inst_test/terminate": + case terminateTestInstanceRoute: writeJSON(t, writer, instanceResponse{ID: "inst_test", Status: instanceStatusTerminated}) - case "GET /integrations/brev/v1/pools/sfc:pool:account:workspace:default": + case getTestPoolRoute: writeJSON(t, writer, poolResponse{AllocationSchedule: allocationSchedule{ ByInstanceSKU: map[string][]scheduleEntry{"is_sku": {{StartAt: 0, NodeCount: 1}}}, }}) diff --git a/v1/providers/sfcomputev2/brev_constants.go b/v1/providers/sfcomputev2/brev_constants.go index 2e238fb..db678c4 100644 --- a/v1/providers/sfcomputev2/brev_constants.go +++ b/v1/providers/sfcomputev2/brev_constants.go @@ -6,10 +6,10 @@ import "fmt" const ( defaultSSHUsername = "ubuntu" - // Internal tag keys written to every SFCompute V2 instance. These are stripped from - // v1.Instance.Tags on read so they don't surface as user-facing tags. + // Provider metadata is stripped from v1.Instance.Tags on read. tagKeyCloudCredRefID = "brev-cloud-cred-ref-id" //nolint:gosec // not a secret tagKeyRefID = "brev-ref-id" + tagKeyFirewallID = "brev-firewall-id" // Brev environment config for SFCompute V2. brevDefaultImageResourcePath = "sfc:image:sfcompute:public:ubuntu-24.04.4-cuda-12.8" diff --git a/v1/providers/sfcomputev2/capabilities.go b/v1/providers/sfcomputev2/capabilities.go index e9b62d6..2f9d527 100644 --- a/v1/providers/sfcomputev2/capabilities.go +++ b/v1/providers/sfcomputev2/capabilities.go @@ -15,10 +15,28 @@ func getSFCCapabilitiesV2() v1.Capabilities { } } -func (c *SFCClientV2) GetCapabilities(_ context.Context) (v1.Capabilities, error) { - return getSFCCapabilitiesV2(), nil +func (c *SFCClientV2) GetCapabilities(ctx context.Context) (v1.Capabilities, error) { + capabilities := getSFCCapabilitiesV2() + if !c.enableConfigurableFirewall { + return capabilities, nil + } + pool, err := c.client.getPool(ctx, c.GetDefaultPoolResourcePath()) + if err != nil { + return nil, err + } + if pool.PublicIPv4SKUs != nil && len(*pool.PublicIPv4SKUs) > 0 { + capabilities = append(capabilities, v1.CapabilityModifyFirewall) + } + return capabilities, nil } -func (c *SFCCredentialV2) GetCapabilities(_ context.Context) (v1.Capabilities, error) { - return getSFCCapabilitiesV2(), nil +func (c *SFCCredentialV2) GetCapabilities(ctx context.Context) (v1.Capabilities, error) { + if !c.EnableConfigurableFirewall { + return getSFCCapabilitiesV2(), nil + } + client, err := c.MakeClient(ctx, "") + if err != nil { + return nil, err + } + return client.GetCapabilities(ctx) } diff --git a/v1/providers/sfcomputev2/client.go b/v1/providers/sfcomputev2/client.go index 19b5d59..a1b6269 100644 --- a/v1/providers/sfcomputev2/client.go +++ b/v1/providers/sfcomputev2/client.go @@ -10,10 +10,11 @@ const CloudProviderID = "sfcompute" // SFCCredentialV2 holds authentication details for a Brev-managed SFCompute V2 account. type SFCCredentialV2 struct { - RefID string - APIKey string `json:"api_key"` - Organization string `json:"organization"` - Workspace string `json:"workspace"` + RefID string + APIKey string `json:"api_key"` + Organization string `json:"organization"` + Workspace string `json:"workspace"` + EnableConfigurableFirewall bool `json:"enable_configurable_firewall,omitempty"` } var _ v1.CloudCredential = &SFCCredentialV2{} @@ -45,12 +46,13 @@ func (c *SFCCredentialV2) GetTenantID() (string, error) { type SFCClientV2 struct { v1.NotImplCloudClient - refID string - organization string - workspace string - location string - client *apiClient - logger v1.Logger + refID string + organization string + workspace string + location string + client *apiClient + logger v1.Logger + enableConfigurableFirewall bool } var _ v1.CloudClient = &SFCClientV2{} @@ -65,12 +67,13 @@ func WithLogger(logger v1.Logger) SFCClientV2Option { func (c *SFCCredentialV2) MakeClientWithOptions(_ context.Context, location string, opts ...SFCClientV2Option) (v1.CloudClient, error) { sfcClient := &SFCClientV2{ - refID: c.RefID, - organization: c.Organization, - workspace: c.Workspace, - location: location, - client: newAPIClient(c.APIKey), - logger: &v1.NoopLogger{}, + refID: c.RefID, + organization: c.Organization, + workspace: c.Workspace, + location: location, + client: newAPIClient(c.APIKey), + logger: &v1.NoopLogger{}, + enableConfigurableFirewall: c.EnableConfigurableFirewall, } for _, opt := range opts { diff --git a/v1/providers/sfcomputev2/firewall.go b/v1/providers/sfcomputev2/firewall.go new file mode 100644 index 0000000..803a9e5 --- /dev/null +++ b/v1/providers/sfcomputev2/firewall.go @@ -0,0 +1,236 @@ +package v2 + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "fmt" + "net/http" + "net/netip" + "net/url" + "slices" + "sort" + "strconv" + "strings" + "time" + + v1 "github.com/brevdev/cloud/v1" + "github.com/google/uuid" +) + +const firewallVersionHeader = "X-SFC-Firewall-Version" + +type firewallRule struct { + Direction string `json:"direction"` + Protocol string `json:"protocol"` + Port string `json:"port,omitempty"` + Source string `json:"source"` +} + +type firewallResponse struct { + ID string `json:"id"` + Rules []firewallRule `json:"rules"` +} + +func isAPIStatus(err error, status int) bool { + var responseErr *apiError + return errors.As(err, &responseErr) && responseErr.statusCode == status +} + +func sshFirewallRules() []firewallRule { + return []firewallRule{ + {Direction: "ingress", Protocol: "tcp", Port: "22", Source: "0.0.0.0/0"}, + {Direction: "ingress", Protocol: "tcp", Port: "2222", Source: "0.0.0.0/0"}, + } +} + +func expandFirewallRules(rules v1.FirewallRules) ([]firewallRule, error) { + for _, rule := range rules.EgressRules { + if !isUnrestrictedEgress(rule) { + return nil, fmt.Errorf("SFCompute does not support restricting outbound traffic") + } + } + var result []firewallRule + for _, rule := range rules.IngressRules { + if rule.FromPort < 0 || rule.ToPort > 65535 || rule.FromPort > rule.ToPort { + return nil, fmt.Errorf("invalid ingress port range %d-%d", rule.FromPort, rule.ToPort) + } + port := strconv.Itoa(int(rule.FromPort)) + if rule.FromPort != rule.ToPort { + port += "-" + strconv.Itoa(int(rule.ToPort)) + } + sources := rule.IPRanges + if len(sources) == 0 { + sources = []string{"0.0.0.0/0"} + } + for _, source := range sources { + prefix, err := netip.ParsePrefix(source) + if err != nil || !prefix.Addr().Is4() { + return nil, fmt.Errorf("ingress source must be an IPv4 CIDR: %q", source) + } + // Brev's provider contract has no protocol field; a port rule covers TCP and UDP. + for _, protocol := range []string{"tcp", "udp"} { + result = append(result, firewallRule{ + Direction: "ingress", Protocol: protocol, Port: port, Source: prefix.Masked().String(), + }) + } + } + } + return result, nil +} + +func isUnrestrictedEgress(rule v1.FirewallRule) bool { + return rule.FromPort == 0 && rule.ToPort == 65535 && + (len(rule.IPRanges) == 0 || slices.Equal(rule.IPRanges, []string{"0.0.0.0/0"})) +} + +func deduplicateRules(rules []firewallRule) []firewallRule { + result := make([]firewallRule, 0, len(rules)) + seen := make(map[firewallRule]bool) + for _, rule := range rules { + if !seen[rule] { + result = append(result, rule) + seen[rule] = true + } + } + return result +} + +func (c *SFCClientV2) createInstanceFirewall(ctx context.Context, rules v1.FirewallRules) (*firewallResponse, error) { + expanded, err := expandFirewallRules(rules) + if err != nil { + return nil, err + } + request := struct { + Name string `json:"name"` + Workspace string `json:"workspace"` + Rules []firewallRule `json:"rules"` + }{ + Name: "brev-" + uuid.NewString(), Workspace: c.GetWorkspaceResourcePath(), + Rules: deduplicateRules(append(sshFirewallRules(), expanded...)), + } + var response firewallResponse + _, err = c.client.doRequest(ctx, http.MethodPost, "/integrations/brev/v1/firewalls", nil, request, &response, nil) + if err != nil { + return nil, err + } + if response.ID == "" { + return nil, fmt.Errorf("SFCompute returned a firewall without an ID") + } + return &response, nil +} + +func (c *SFCClientV2) cleanupFirewall(ctx context.Context, id string, prior error) error { + // A canceled create request must not cancel cleanup of its separate firewall. + ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + defer cancel() + _, err := c.client.doRequest(ctx, http.MethodDelete, "/integrations/brev/v1/firewalls/"+url.PathEscape(id), nil, nil, nil, nil) + if isAPIStatus(err, http.StatusNotFound) { + err = nil + } + if err != nil { + err = fmt.Errorf("clean up firewall %s: %w", id, err) + } + return errors.Join(prior, err) +} + +func (c *SFCClientV2) getFirewall(ctx context.Context, id string) (*firewallResponse, string, error) { + var response firewallResponse + headers, err := c.client.doRequest(ctx, http.MethodGet, "/integrations/brev/v1/firewalls/"+url.PathEscape(id), nil, nil, &response, nil) + if err != nil { + return nil, "", err + } + return &response, headers.Get(firewallVersionHeader), nil +} + +func ingressRuleID(rule firewallRule) string { + hash := sha256.Sum256([]byte(rule.Port + "|" + rule.Source)) + return "sfc-ingress-" + hex.EncodeToString(hash[:]) +} + +func (c *SFCClientV2) loadInstanceFirewall(ctx context.Context, source *instanceResponse, instance *v1.Instance) error { + if source.Firewall == "" || source.Status == instanceStatusTerminated { + return nil + } + firewall, _, err := c.getFirewall(ctx, source.Firewall) + if err != nil { + return err + } + rules := make(map[string]v1.FirewallRule) + for _, rule := range firewall.Rules { + if rule.Direction != "ingress" || rule.Port == "" || slices.Contains(sshFirewallRules(), rule) { + continue + } + ports := strings.SplitN(rule.Port, "-", 2) + from, err := strconv.ParseInt(ports[0], 10, 32) + if err != nil { + return fmt.Errorf("invalid firewall port from SFCompute: %w", err) + } + to := from + if len(ports) == 2 { + to, err = strconv.ParseInt(ports[1], 10, 32) + if err != nil { + return fmt.Errorf("invalid firewall port from SFCompute: %w", err) + } + } + id := ingressRuleID(rule) + rules[id] = v1.FirewallRule{ID: id, FromPort: int32(from), ToPort: int32(to), IPRanges: []string{rule.Source}} + } + ids := make([]string, 0, len(rules)) + for id := range rules { + ids = append(ids, id) + } + sort.Strings(ids) + for _, id := range ids { + instance.FirewallRules.IngressRules = append(instance.FirewallRules.IngressRules, rules[id]) + } + return nil +} + +func (c *SFCClientV2) AddFirewallRulesToInstance(ctx context.Context, args v1.AddFirewallRulesToInstanceArgs) error { + rules, err := expandFirewallRules(args.FirewallRules) + if err != nil { + return err + } + return c.updateInstanceFirewall(ctx, args.InstanceID, func(current []firewallRule) []firewallRule { + return deduplicateRules(append(current, rules...)) + }) +} + +func (c *SFCClientV2) RevokeSecurityGroupRules(ctx context.Context, args v1.RevokeSecurityGroupRuleArgs) error { + return c.updateInstanceFirewall(ctx, args.InstanceID, func(current []firewallRule) []firewallRule { + return slices.DeleteFunc(current, func(rule firewallRule) bool { + return rule.Direction == "ingress" && !slices.Contains(sshFirewallRules(), rule) && + slices.Contains(args.SecurityGroupRuleIDs, ingressRuleID(rule)) + }) + }) +} + +func (c *SFCClientV2) updateInstanceFirewall(ctx context.Context, id v1.CloudProviderInstanceID, change func([]firewallRule) []firewallRule) error { + instance, err := c.client.getInstance(ctx, string(id)) + if err != nil { + return err + } + if !instance.EnablePublicIPv4 || instance.Firewall == "" || instance.Status == instanceStatusTerminated { + return fmt.Errorf("configurable firewall requires a live instance created with public networking: %w", v1.ErrNotImplemented) + } + for range 5 { + firewall, version, err := c.getFirewall(ctx, instance.Firewall) + if err != nil { + return err + } + if version == "" { + return fmt.Errorf("SFCompute API does not support conditional firewall updates") + } + request := struct { + Rules []firewallRule `json:"rules"` + }{Rules: change(firewall.Rules)} + _, err = c.client.doRequest(ctx, http.MethodPut, "/integrations/brev/v1/firewalls/"+url.PathEscape(instance.Firewall), nil, + request, nil, http.Header{firewallVersionHeader: []string{version}}) + if !isAPIStatus(err, http.StatusConflict) { + return err + } + } + return fmt.Errorf("firewall changed during five update attempts; retry the operation") +} diff --git a/v1/providers/sfcomputev2/firewall_cleanup_test.go b/v1/providers/sfcomputev2/firewall_cleanup_test.go new file mode 100644 index 0000000..d610a2f --- /dev/null +++ b/v1/providers/sfcomputev2/firewall_cleanup_test.go @@ -0,0 +1,142 @@ +package v2 + +import ( + "context" + "net/http" + "testing" + + v1 "github.com/brevdev/cloud/v1" + "github.com/stretchr/testify/require" +) + +func TestNetworkingCleansUpRejectedOrUnacknowledgedCreates(t *testing.T) { + t.Parallel() + for _, rejected := range []bool{false, true} { + name := "unacknowledged" + if rejected { + name = "rejected" + } + t.Run(name, func(t *testing.T) { + t.Parallel() + terminated, deleted := false, false + client := networkingTestClient(t, true, func(w http.ResponseWriter, r *http.Request) { + switch r.Method + " " + r.URL.Path { + case getTestPoolRoute: + writeJSON(t, w, networkingPool()) + case listInstancesRoute: + writeJSON(t, w, listInstancesResponse{}) + case createFirewallRoute: + writeJSON(t, w, firewallResponse{ID: "frwl_test"}) + case createInstanceRoute: + if rejected { + w.WriteHeader(http.StatusUnprocessableEntity) + return + } + writeJSON(t, w, instanceResponse{ID: "inst_test", Status: instanceStatusAwaitingAllocation}) + case terminateTestInstanceRoute: + terminated = true + writeJSON(t, w, instanceResponse{ID: "inst_test", Status: instanceStatusTerminated}) + case deleteTestFirewallRoute: + deleted = true + w.WriteHeader(http.StatusNoContent) + default: + t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path) + http.NotFound(w, r) + } + }) + _, err := client.CreateInstance(context.Background(), v1.CreateInstanceAttrs{RefID: "test"}) + require.Error(t, err) + require.True(t, deleted) + require.Equal(t, !rejected, terminated) + }) + } +} + +func TestLegacyTerminationDoesNotRequireFirewallPermissions(t *testing.T) { + t.Parallel() + client := networkingTestClient(t, false, func(w http.ResponseWriter, r *http.Request) { + require.Equal(t, http.MethodPost, r.Method) + require.Equal(t, "/integrations/brev/v1/instances/inst_test/terminate", r.URL.Path) + writeJSON(t, w, instanceResponse{ID: "inst_test", Status: instanceStatusTerminated}) + }) + require.NoError(t, client.TerminateInstance(context.Background(), "inst_test")) +} + +func TestCanceledCreateStillCleansUpFirewall(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + deleted := false + client := networkingTestClient(t, true, func(w http.ResponseWriter, r *http.Request) { + switch r.Method + " " + r.URL.Path { + case getTestPoolRoute: + writeJSON(t, w, networkingPool()) + case listInstancesRoute: + writeJSON(t, w, listInstancesResponse{}) + case createFirewallRoute: + writeJSON(t, w, firewallResponse{ID: "frwl_test"}) + case createInstanceRoute: + cancel() + w.WriteHeader(http.StatusUnprocessableEntity) + case deleteTestFirewallRoute: + deleted = true + w.WriteHeader(http.StatusNoContent) + default: + t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path) + http.NotFound(w, r) + } + }) + _, err := client.CreateInstance(ctx, v1.CreateInstanceAttrs{RefID: "test"}) + require.Error(t, err) + require.True(t, deleted, "cleanup must survive cancellation of the create request") +} + +func TestTerminationRetriesFailedFirewallCleanup(t *testing.T) { + t.Parallel() + deletes := 0 + client := networkingTestClient(t, false, func(w http.ResponseWriter, r *http.Request) { + switch r.Method + " " + r.URL.Path { + case terminateTestInstanceRoute: + writeJSON(t, w, instanceResponse{ + ID: "inst_test", Status: instanceStatusTerminated, + EnablePublicIPv4: true, Firewall: "frwl_test", + Tags: map[string]string{tagKeyFirewallID: "frwl_test"}, + }) + case deleteTestFirewallRoute: + deletes++ + if deletes == 1 { + w.WriteHeader(http.StatusServiceUnavailable) + } else { + w.WriteHeader(http.StatusNoContent) + } + default: + t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path) + http.NotFound(w, r) + } + }) + require.ErrorContains(t, client.TerminateInstance(context.Background(), "inst_test"), "frwl_test") + require.NoError(t, client.TerminateInstance(context.Background(), "inst_test")) + require.Equal(t, 2, deletes) +} + +func TestTerminationPreservesUnmanagedFirewalls(t *testing.T) { + t.Parallel() + for _, ownedFirewall := range []string{"", "frwl_previous"} { + t.Run("owned="+ownedFirewall, func(t *testing.T) { + t.Parallel() + client := networkingTestClient(t, false, func(w http.ResponseWriter, r *http.Request) { + if r.Method+" "+r.URL.Path != terminateTestInstanceRoute { + t.Errorf("must not delete an unmanaged firewall: %s %s", r.Method, r.URL.Path) + http.NotFound(w, r) + return + } + writeJSON(t, w, instanceResponse{ + ID: "inst_test", Status: instanceStatusTerminated, + EnablePublicIPv4: true, Firewall: "frwl_test", + Tags: map[string]string{tagKeyFirewallID: ownedFirewall}, + }) + }) + require.NoError(t, client.TerminateInstance(context.Background(), "inst_test")) + }) + } +} diff --git a/v1/providers/sfcomputev2/firewall_test.go b/v1/providers/sfcomputev2/firewall_test.go new file mode 100644 index 0000000..987a92b --- /dev/null +++ b/v1/providers/sfcomputev2/firewall_test.go @@ -0,0 +1,238 @@ +package v2 + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "testing" + + v1 "github.com/brevdev/cloud/v1" + "github.com/stretchr/testify/require" +) + +const ( + terminateTestInstanceRoute = "POST /integrations/brev/v1/instances/inst_test/terminate" + createInstanceRoute = "POST /integrations/brev/v1/instances" + listInstancesRoute = "GET /integrations/brev/v1/instances" + getTestInstanceRoute = "GET /integrations/brev/v1/instances/inst_test" + createFirewallRoute = "POST /integrations/brev/v1/firewalls" + deleteTestFirewallRoute = "DELETE /integrations/brev/v1/firewalls/frwl_test" + getTestPoolRoute = "GET /integrations/brev/v1/pools/sfc:pool:account:workspace:default" +) + +func networkingTestClient(t *testing.T, enabled bool, handler http.HandlerFunc) *SFCClientV2 { + t.Helper() + server := httptest.NewServer(handler) + t.Cleanup(server.Close) + credential := NewSFCCredentialV2("cred", "key", "account", "workspace") + credential.EnableConfigurableFirewall = enabled + client, err := credential.MakeClient(context.Background(), sfcLocation) + require.NoError(t, err) + sfc, ok := client.(*SFCClientV2) + require.True(t, ok) + sfc.client.baseURL = server.URL + return sfc +} + +func networkingPool() poolResponse { + return poolResponse{ + PublicIPv4SKUs: pointerTo([]string{"is_public"}), + AllocationSchedule: allocationSchedule{ByInstanceSKU: map[string][]scheduleEntry{ + "is_legacy": {{StartAt: 0, NodeCount: 2}}, + "is_public": {{StartAt: 0, NodeCount: 2}}, + }}, + } +} + +func TestNetworkingRequiresOptInAndServerSupport(t *testing.T) { + t.Parallel() + for _, enabled := range []bool{false, true} { + t.Run(fmt.Sprint(enabled), func(t *testing.T) { + t.Parallel() + created := false + client := networkingTestClient(t, enabled, func(w http.ResponseWriter, r *http.Request) { + switch r.Method + " " + r.URL.Path { + case getTestPoolRoute: + pool := networkingPool() + pool.PublicIPv4SKUs = nil + writeJSON(t, w, pool) + case listInstancesRoute: + writeJSON(t, w, listInstancesResponse{}) + case createInstanceRoute: + created = true + var body map[string]any + require.NoError(t, json.NewDecoder(r.Body).Decode(&body)) + require.NotContains(t, body, "enable_public_ipv4") + require.NotContains(t, body, "firewall") + writeJSON(t, w, instanceResponse{ID: "inst_test", Status: instanceStatusAwaitingAllocation}) + default: + t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path) + http.NotFound(w, r) + } + }) + _, err := client.CreateInstance(context.Background(), v1.CreateInstanceAttrs{RefID: "test"}) + if enabled { + require.ErrorContains(t, err, "not enabled") + require.False(t, created) + } else { + require.NoError(t, err) + require.True(t, created) + } + }) + } +} + +func TestPublicInstanceLifecycle(t *testing.T) { + t.Parallel() + firewall := firewallResponse{ID: "frwl_test"} + var instanceTags map[string]string + deleted := false + client := networkingTestClient(t, true, func(w http.ResponseWriter, r *http.Request) { + switch r.Method + " " + r.URL.Path { + case getTestPoolRoute: + writeJSON(t, w, networkingPool()) + case listInstancesRoute: + writeJSON(t, w, listInstancesResponse{}) + case createFirewallRoute: + var body struct { + Workspace string `json:"workspace"` + Rules []firewallRule `json:"rules"` + } + require.NoError(t, json.NewDecoder(r.Body).Decode(&body)) + require.Equal(t, "sfc:workspace:account:workspace", body.Workspace) + firewall.Rules = body.Rules + require.Len(t, firewall.Rules, 4) + require.Contains(t, firewall.Rules, firewallRule{"ingress", "udp", "8000-8010", "8.8.8.0/24"}) + writeJSON(t, w, firewall) + case createInstanceRoute: + var body createInstanceRequest + require.NoError(t, json.NewDecoder(r.Body).Decode(&body)) + require.Equal(t, "is_public", body.InstanceSKU) + require.True(t, body.EnablePublicIPv4) + require.Equal(t, firewall.ID, body.Firewall) + require.Equal(t, firewall.ID, body.Tags[tagKeyFirewallID]) + instanceTags = body.Tags + writeJSON(t, w, instanceResponse{ID: "inst_test", EnablePublicIPv4: true, Firewall: firewall.ID, Tags: instanceTags}) + case getTestInstanceRoute: + writeJSON(t, w, instanceResponse{ + ID: "inst_test", Status: instanceStatusRunning, + EnablePublicIPv4: true, Firewall: firewall.ID, PublicIP: "192.0.2.10", + Tags: instanceTags, + }) + case "GET /integrations/brev/v1/firewalls/frwl_test": + writeJSON(t, w, firewall) + case terminateTestInstanceRoute: + writeJSON(t, w, instanceResponse{ + ID: "inst_test", Status: instanceStatusTerminated, + EnablePublicIPv4: true, Firewall: firewall.ID, Tags: instanceTags, + }) + case deleteTestFirewallRoute: + deleted = true + w.WriteHeader(http.StatusNoContent) + default: + t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path) + http.NotFound(w, r) + } + }) + ctx := context.Background() + capabilities, err := client.GetCapabilities(ctx) + require.NoError(t, err) + require.Contains(t, capabilities, v1.CapabilityModifyFirewall) + _, err = client.CreateInstance(ctx, v1.CreateInstanceAttrs{RefID: "test", FirewallRules: v1.FirewallRules{ + IngressRules: []v1.FirewallRule{{FromPort: 8000, ToPort: 8010, IPRanges: []string{"8.8.8.1/24"}}}, + }}) + require.NoError(t, err) + instance, err := client.GetInstance(ctx, "inst_test") + require.NoError(t, err) + require.Equal(t, "192.0.2.10", instance.PublicIP) + require.Equal(t, 22, instance.SSHPort) + require.Len(t, instance.FirewallRules.IngressRules, 1) + require.NotEmpty(t, instance.FirewallRules.IngressRules[0].ID) + require.NotContains(t, instance.Tags, tagKeyFirewallID) + require.NoError(t, client.TerminateInstance(ctx, "inst_test")) + require.True(t, deleted) +} + +func TestFirewallUpdateRetriesWithoutLosingConcurrentRules(t *testing.T) { + t.Parallel() + firewall := firewallResponse{ID: "frwl_test", Rules: sshFirewallRules()} + version := 1 + updates := 0 + concurrentRule := firewallRule{"ingress", "tcp", "9000", "8.8.8.0/24"} + client := networkingTestClient(t, true, func(w http.ResponseWriter, r *http.Request) { + switch r.Method + " " + r.URL.Path { + case getTestInstanceRoute: + writeJSON(t, w, instanceResponse{ID: "inst_test", EnablePublicIPv4: true, Firewall: firewall.ID}) + case "GET /integrations/brev/v1/firewalls/frwl_test": + w.Header().Set(firewallVersionHeader, fmt.Sprint(version)) + writeJSON(t, w, firewall) + case "PUT /integrations/brev/v1/firewalls/frwl_test": + updates++ + require.Equal(t, fmt.Sprint(version), r.Header.Get(firewallVersionHeader)) + if updates == 1 { + firewall.Rules = append(firewall.Rules, concurrentRule) + version++ + w.WriteHeader(http.StatusConflict) + return + } + var body firewallResponse + require.NoError(t, json.NewDecoder(r.Body).Decode(&body)) + firewall.Rules = body.Rules + version++ + writeJSON(t, w, firewall) + default: + t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path) + http.NotFound(w, r) + } + }) + ctx := context.Background() + require.NoError(t, client.AddFirewallRulesToInstance(ctx, v1.AddFirewallRulesToInstanceArgs{ + InstanceID: "inst_test", FirewallRules: v1.FirewallRules{IngressRules: []v1.FirewallRule{{FromPort: 8080, ToPort: 8080}}}, + })) + require.Equal(t, 2, updates) + require.Contains(t, firewall.Rules, concurrentRule) + require.Len(t, firewall.Rules, 5) + id := ingressRuleID(firewallRule{Port: "8080", Source: "0.0.0.0/0"}) + require.NoError(t, client.RevokeSecurityGroupRules(ctx, v1.RevokeSecurityGroupRuleArgs{ + InstanceID: "inst_test", SecurityGroupRuleIDs: []string{id}, + })) + require.Len(t, firewall.Rules, 3) + require.Contains(t, firewall.Rules, concurrentRule) + for _, rule := range sshFirewallRules() { + require.Contains(t, firewall.Rules, rule) + } +} + +func TestNetworkingRejectsInvalidRules(t *testing.T) { + t.Parallel() + for _, rule := range []v1.FirewallRule{ + {FromPort: -1, ToPort: 22}, + {FromPort: 20, ToPort: 10}, + {FromPort: 1, ToPort: 65536}, + {FromPort: 80, ToPort: 80, IPRanges: []string{"::/0"}}, + {FromPort: 80, ToPort: 80, IPRanges: []string{"192.0.2.1"}}, + } { + _, err := expandFirewallRules(v1.FirewallRules{IngressRules: []v1.FirewallRule{rule}}) + require.Error(t, err) + } + _, err := expandFirewallRules(v1.FirewallRules{EgressRules: []v1.FirewallRule{{FromPort: 80, ToPort: 80}}}) + require.ErrorContains(t, err, "outbound") +} + +func TestFirewallUpdateRequiresServerVersion(t *testing.T) { + t.Parallel() + client := networkingTestClient(t, true, func(w http.ResponseWriter, r *http.Request) { + require.Equal(t, http.MethodGet, r.Method) + if r.URL.Path == "/integrations/brev/v1/instances/inst_test" { + writeJSON(t, w, instanceResponse{ID: "inst_test", EnablePublicIPv4: true, Firewall: "frwl_test"}) + } else { + writeJSON(t, w, firewallResponse{ID: "frwl_test", Rules: sshFirewallRules()}) + } + }) + err := client.AddFirewallRulesToInstance(context.Background(), v1.AddFirewallRulesToInstanceArgs{ + InstanceID: "inst_test", FirewallRules: v1.FirewallRules{IngressRules: []v1.FirewallRule{{FromPort: 8080, ToPort: 8080}}}, + }) + require.ErrorContains(t, err, "conditional firewall updates") +} diff --git a/v1/providers/sfcomputev2/instance.go b/v1/providers/sfcomputev2/instance.go index 34cf70a..217f2a1 100644 --- a/v1/providers/sfcomputev2/instance.go +++ b/v1/providers/sfcomputev2/instance.go @@ -46,7 +46,7 @@ func (c *SFCClientV2) CreateInstance(ctx context.Context, attrs v1.CreateInstanc v1.LogField("location", attrs.Location), ) - tags := make(map[string]string, len(attrs.Tags)+2) + tags := make(map[string]string, len(attrs.Tags)+3) maps.Copy(tags, attrs.Tags) tags[tagKeyCloudCredRefID] = c.refID tags[tagKeyRefID] = attrs.RefID @@ -68,13 +68,34 @@ func (c *SFCClientV2) CreateInstance(ctx context.Context, attrs v1.CreateInstanc if name := makeSFCName(attrs.RefID, attrs.Tags); sfcNamePattern.MatchString(name) { req.Name = &name } + if c.enableConfigurableFirewall { + firewall, err := c.createInstanceFirewall(ctx, attrs.FirewallRules) + if err != nil { + return nil, errors.WrapAndTrace(err) + } + req.EnablePublicIPv4 = true + req.Firewall = firewall.ID + req.Tags[tagKeyFirewallID] = firewall.ID + } resp, err := c.client.createInstance(ctx, req) if err != nil { + if req.Firewall != "" { + err = c.cleanupFirewall(ctx, req.Firewall, err) + } return nil, errors.WrapAndTrace(err) } if resp == nil { return nil, errors.WrapAndTrace(fmt.Errorf("no instance returned from create")) } + if req.EnablePublicIPv4 && (!resp.EnablePublicIPv4 || resp.Firewall != req.Firewall) { + err := fmt.Errorf("SFCompute API did not accept the requested public networking") + cleanupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + defer cancel() + if terminateErr := c.client.terminateInstance(cleanupCtx, resp.ID); terminateErr != nil && !isAPIStatus(terminateErr, http.StatusNotFound) { + return nil, errors.WrapAndTrace(fmt.Errorf("%w; instance cleanup failed: %w", err, terminateErr)) + } + return nil, errors.WrapAndTrace(c.cleanupFirewall(ctx, req.Firewall, err)) + } instance, err := c.sfcInstanceToBrevInstance(resp, nil) if err != nil { @@ -107,15 +128,21 @@ func (c *SFCClientV2) GetInstance(ctx context.Context, id v1.CloudProviderInstan return nil, errors.WrapAndTrace(fmt.Errorf("instance %s not found", id)) } - sshInfo, err := c.getSSHInfo(ctx, string(id), resp.Status) - if err != nil { - return nil, errors.WrapAndTrace(err) + var sshInfo *instanceSSHInfo + if !resp.EnablePublicIPv4 { + sshInfo, err = c.getSSHInfo(ctx, string(id), resp.Status) + if err != nil { + return nil, errors.WrapAndTrace(err) + } } instance, err := c.sfcInstanceToBrevInstance(resp, sshInfo) if err != nil { return nil, errors.WrapAndTrace(err) } + if err := c.loadInstanceFirewall(ctx, resp, instance); err != nil { + return nil, errors.WrapAndTrace(err) + } c.logger.Debug(ctx, "sfcv2: GetInstance end", v1.LogField("instanceID", id), @@ -146,7 +173,11 @@ func (c *SFCClientV2) ListInstances(ctx context.Context, args v1.ListInstancesAr continue } - sshInfo, err := c.getSSHInfo(ctx, inst.ID, inst.Status) + var sshInfo *instanceSSHInfo + var err error + if !inst.EnablePublicIPv4 { + sshInfo, err = c.getSSHInfo(ctx, inst.ID, inst.Status) + } if err != nil { c.logger.Error(ctx, err, v1.LogField("msg", "sfcv2: ListInstances skipping instance due to SSH error"), @@ -163,6 +194,9 @@ func (c *SFCClientV2) ListInstances(ctx context.Context, args v1.ListInstancesAr ) continue } + if err := c.loadInstanceFirewall(ctx, &inst, brevInst); err != nil { + return nil, errors.WrapAndTrace(err) + } instances = append(instances, *brevInst) } @@ -178,9 +212,13 @@ func (c *SFCClientV2) TerminateInstance(ctx context.Context, id v1.CloudProvider v1.LogField("instanceID", id), ) - if err := c.client.terminateInstance(ctx, string(id)); err != nil { + instance, err := c.client.terminateInstanceWithResponse(ctx, string(id)) + if err != nil { return normalizeTerminateInstanceError(err) } + if instance.EnablePublicIPv4 && instance.Firewall != "" && instance.Tags[tagKeyFirewallID] == instance.Firewall { + return c.cleanupFirewall(ctx, instance.Firewall, nil) + } c.logger.Debug(ctx, "sfcv2: TerminateInstance end", v1.LogField("instanceID", id), @@ -226,13 +264,20 @@ func (c *SFCClientV2) sfcInstanceToBrevInstance(inst *instanceResponse, sshInfo userTags := make(v1.Tags) for k, v := range tags { switch k { - case tagKeyCloudCredRefID, tagKeyRefID: + case tagKeyCloudCredRefID, tagKeyRefID, tagKeyFirewallID: default: userTags[k] = v } } status := sfcStatusToLifecycleStatus(inst.Status) + hostname, sshPort := sshInfo.GetHostname(), int(sshInfo.GetPort()) + if inst.EnablePublicIPv4 { + hostname, sshPort = inst.PublicIP, 22 + if hostname == "" && status == v1.LifecycleStatusRunning { + status = v1.LifecycleStatusPending + } + } diskInt64, err := h100InstanceTypeMetadata.diskBytes.ByteCountInUnitInt64(v1.Gibibyte) if err != nil { @@ -244,10 +289,10 @@ func (c *SFCClientV2) sfcInstanceToBrevInstance(inst *instanceResponse, sshInfo Name: inst.Name, CloudID: v1.CloudProviderInstanceID(inst.ID), RefID: tags[tagKeyRefID], - PublicDNS: sshInfo.GetHostname(), - PublicIP: sshInfo.GetHostname(), + PublicDNS: hostname, + PublicIP: hostname, SSHUser: defaultSSHUsername, - SSHPort: int(sshInfo.GetPort()), + SSHPort: sshPort, CreatedAt: time.Unix(inst.CreatedAt, 0), DiskSize: diskSize, DiskSizeBytes: h100InstanceTypeMetadata.diskBytes, diff --git a/v1/providers/sfcomputev2/instancetype.go b/v1/providers/sfcomputev2/instancetype.go index 6b2fd39..31023db 100644 --- a/v1/providers/sfcomputev2/instancetype.go +++ b/v1/providers/sfcomputev2/instancetype.go @@ -3,6 +3,7 @@ package v2 import ( "context" "fmt" + "slices" "sort" "time" @@ -118,6 +119,7 @@ func (c *SFCClientV2) GetInstanceTypes(ctx context.Context, args v1.GetInstanceT } instanceType := buildInstanceType(h100InstanceTypeMetadata, true) + instanceType.CanModifyFirewallRules = c.enableConfigurableFirewall if !v1.IsSelectedByArgs(instanceType, args) { return []v1.InstanceType{}, nil @@ -145,10 +147,16 @@ func (c *SFCClientV2) skuFreeCapacity(ctx context.Context) (map[string]int, erro if poolResp == nil { return map[string]int{}, nil } + if c.enableConfigurableFirewall && poolResp.PublicIPv4SKUs == nil { + return nil, fmt.Errorf("configurable firewall is not enabled by the SFCompute API") + } now := time.Now().Unix() free := make(map[string]int) for skuID, schedule := range poolResp.AllocationSchedule.ByInstanceSKU { + if c.enableConfigurableFirewall && !slices.Contains(*poolResp.PublicIPv4SKUs, skuID) { + continue + } free[skuID] = currentScheduleAllocation(schedule, now) }