From 8d01bb82358921f798033200ef7a0c4c84228a4d Mon Sep 17 00:00:00 2001 From: Mason Wheeler Date: Wed, 29 Jul 2026 18:10:52 -0700 Subject: [PATCH 1/3] refactor(sfcompute): use Brev integration API --- v1/providers/sfcomputev2/brev_constants.go | 28 -- v1/providers/sfcomputev2/client.go | 11 +- v1/providers/sfcomputev2/instance.go | 222 +++++----------- v1/providers/sfcomputev2/instancetype.go | 248 +++++------------- .../sfcomputev2/integration_client.go | 234 +++++++++++++++++ .../sfcomputev2/integration_client_test.go | 215 +++++++++++++++ v1/providers/sfcomputev2/validation_test.go | 2 +- 7 files changed, 591 insertions(+), 369 deletions(-) delete mode 100644 v1/providers/sfcomputev2/brev_constants.go create mode 100644 v1/providers/sfcomputev2/integration_client.go create mode 100644 v1/providers/sfcomputev2/integration_client_test.go diff --git a/v1/providers/sfcomputev2/brev_constants.go b/v1/providers/sfcomputev2/brev_constants.go deleted file mode 100644 index 2e238fb..0000000 --- a/v1/providers/sfcomputev2/brev_constants.go +++ /dev/null @@ -1,28 +0,0 @@ -package v2 - -import "fmt" - -// Package-internal constants — SSH defaults and internal tag keys. -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. - tagKeyCloudCredRefID = "brev-cloud-cred-ref-id" //nolint:gosec // not a secret - tagKeyRefID = "brev-ref-id" - - // Brev environment config for SFCompute V2. - brevDefaultImageResourcePath = "sfc:image:sfcompute:public:ubuntu-24.04.4-cuda-12.8" -) - -func (c *SFCClientV2) GetDefaultPoolResourcePath() string { - return fmt.Sprintf("sfc:pool:%s:%s:default", c.organization, c.workspace) -} - -func (c *SFCClientV2) GetWorkspaceResourcePath() string { - return fmt.Sprintf("sfc:workspace:%s:%s", c.organization, c.workspace) -} - -func (c *SFCClientV2) GetDefaultImageResourcePath() string { - return brevDefaultImageResourcePath -} diff --git a/v1/providers/sfcomputev2/client.go b/v1/providers/sfcomputev2/client.go index 309450e..5867a64 100644 --- a/v1/providers/sfcomputev2/client.go +++ b/v1/providers/sfcomputev2/client.go @@ -4,7 +4,6 @@ import ( "context" v1 "github.com/brevdev/cloud/v1" - sfc "github.com/sfcompute/sfc-go" ) const CloudProviderID = "sfcompute" @@ -50,7 +49,7 @@ type SFCClientV2 struct { organization string workspace string location string - client *sfc.SDK + client *integrationClient logger v1.Logger } @@ -64,13 +63,19 @@ func WithLogger(logger v1.Logger) SFCClientV2Option { } } +func WithAPIURL(apiURL string) SFCClientV2Option { + return func(c *SFCClientV2) { + c.client.baseURL = apiURL + } +} + 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: sfc.New(sfc.WithSecurity(c.APIKey)), + client: newIntegrationClient(c.APIKey), logger: &v1.NoopLogger{}, } diff --git a/v1/providers/sfcomputev2/instance.go b/v1/providers/sfcomputev2/instance.go index f3c1076..ce8c172 100644 --- a/v1/providers/sfcomputev2/instance.go +++ b/v1/providers/sfcomputev2/instance.go @@ -2,8 +2,6 @@ package v2 import ( "context" - "encoding/base64" - "fmt" "maps" "regexp" "slices" @@ -12,23 +10,16 @@ import ( "github.com/alecthomas/units" "github.com/brevdev/cloud/internal/errors" v1 "github.com/brevdev/cloud/v1" - "github.com/sfcompute/sfc-go/models/components" - "github.com/sfcompute/sfc-go/models/operations" - "github.com/sfcompute/sfc-go/optionalnullable" ) -// SFC instance names must match `[a-zA-Z0-9][a-zA-Z0-9._-]{0,254}`: start with an -// alphanumeric character, then alphanumerics/dot/underscore/hyphen, max 255 chars. +const defaultSSHUsername = "ubuntu" + var ( sfcNamePattern = regexp.MustCompile(`^[a-zA-Z0-9][a-zA-Z0-9._-]{0,254}$`) sfcNameDisallowed = regexp.MustCompile(`[^a-zA-Z0-9._-]`) sfcNameLeading = regexp.MustCompile(`^[^a-zA-Z0-9]+`) ) -// sanitizeSFCName coerces a requested instance name into SFC's required format: -// disallowed characters are replaced with '-', leading non-alphanumeric characters are -// dropped (SFC requires an alphanumeric first character), and the result is truncated to -// the 255-char max. Returns "" if no usable characters remain. func sanitizeSFCName(name string) string { name = sfcNameDisallowed.ReplaceAllString(name, "-") name = sfcNameLeading.ReplaceAllString(name, "") @@ -44,84 +35,51 @@ func (c *SFCClientV2) CreateInstance(ctx context.Context, attrs v1.CreateInstanc v1.LogField("location", attrs.Location), ) - tags := make(map[string]string, len(attrs.Tags)+2) - maps.Copy(tags, attrs.Tags) - tags[tagKeyCloudCredRefID] = c.refID - tags[tagKeyRefID] = attrs.RefID - - // Spread instances across every SKU in the capacity rather than piling onto one. - sku, err := c.selectAvailableSku(ctx) - if err != nil { - return nil, errors.WrapAndTrace(err) + request := integrationCreateInstanceRequest{ + Workspace: c.workspace, + RefID: attrs.RefID, + CloudCredentialRefID: c.refID, + InstanceType: attrs.InstanceType, + SSHPublicKey: attrs.PublicKey, + Tags: maps.Clone(attrs.Tags), } - - cloudInit := sshKeyCloudInit(attrs.PublicKey) - req := components.CreateInstanceRequest{ - Pool: c.GetDefaultPoolResourcePath(), - Image: c.GetDefaultImageResourcePath(), - InstanceSku: sku, - CloudInitUserData: &cloudInit, - Tags: optionalnullable.From(&tags), - } - // name is optional; sanitize the requested name to SFC's format and send it only if - // something valid remains. Otherwise omit it — identity is preserved in the tags above. if name := sanitizeSFCName(attrs.Name); sfcNamePattern.MatchString(name) { - req.Name = optionalnullable.From(&name) + request.Name = &name } - resp, err := c.client.Instances.Create(ctx, req) + + response, err := c.client.createInstance(ctx, request) if err != nil { - return nil, errors.WrapAndTrace(err) - } - if resp.InstanceResponse == nil { - return nil, errors.WrapAndTrace(fmt.Errorf("no instance returned from create")) + return nil, wrapIntegrationError(err) } - - instance, err := c.sfcInstanceToBrevInstance(resp.InstanceResponse, nil) + instance, err := integrationInstanceToBrevInstance(response) if err != nil { return nil, errors.WrapAndTrace(err) } c.logger.Debug(ctx, "sfcv2: CreateInstance end", - v1.LogField("instanceID", resp.InstanceResponse.ID), - v1.LogField("instanceSku", sku), + v1.LogField("instanceID", response.ID), ) - return instance, nil } -func sshKeyCloudInit(sshKey string) string { - script := fmt.Sprintf("#cloud-config\nssh_authorized_keys:\n - %s", sshKey) - return base64.StdEncoding.EncodeToString([]byte(script)) -} - func (c *SFCClientV2) GetInstance(ctx context.Context, id v1.CloudProviderInstanceID) (*v1.Instance, error) { c.logger.Debug(ctx, "sfcv2: GetInstance start", v1.LogField("instanceID", id), ) - resp, err := c.client.Instances.Fetch(ctx, string(id)) + response, err := c.client.getInstance(ctx, c.workspace, string(id)) if err != nil { - return nil, errors.WrapAndTrace(err) - } - if resp.InstanceResponse == nil { - return nil, errors.WrapAndTrace(fmt.Errorf("instance %s not found", id)) + return nil, wrapIntegrationError(err) } - - sshInfo, err := c.getSSHInfo(ctx, string(id), resp.InstanceResponse.Status) - if err != nil { - return nil, errors.WrapAndTrace(err) - } - - instance, err := c.sfcInstanceToBrevInstance(resp.InstanceResponse, sshInfo) + instance, err := integrationInstanceToBrevInstance(response) if err != nil { return nil, errors.WrapAndTrace(err) } c.logger.Debug(ctx, "sfcv2: GetInstance end", v1.LogField("instanceID", id), - v1.LogField("status", resp.InstanceResponse.Status), + v1.LogField("status", response.Status), ) - return instance, nil } @@ -130,49 +88,27 @@ func (c *SFCClientV2) ListInstances(ctx context.Context, args v1.ListInstancesAr v1.LogField("location", c.location), ) - poolID := c.GetDefaultPoolResourcePath() - resp, err := c.client.Instances.List(ctx, operations.ListInstancesRequest{ - Workspace: c.GetWorkspaceResourcePath(), - Pool: []string{poolID}, - }) + response, err := c.client.listInstances(ctx, c.workspace) if err != nil { - return nil, errors.WrapAndTrace(err) - } - if resp.ListInstancesResponse == nil { - return []v1.Instance{}, nil + return nil, wrapIntegrationError(err) } - var instances []v1.Instance - for _, inst := range resp.ListInstancesResponse.Data { - // Filter by instance IDs if specified. - if len(args.InstanceIDs) > 0 && !slices.Contains(args.InstanceIDs, v1.CloudProviderInstanceID(inst.ID)) { + instances := make([]v1.Instance, 0, len(response)) + for i := range response { + if len(args.InstanceIDs) > 0 && + !slices.Contains(args.InstanceIDs, v1.CloudProviderInstanceID(response[i].ID)) { continue } - - sshInfo, err := c.getSSHInfo(ctx, inst.ID, inst.Status) + instance, err := integrationInstanceToBrevInstance(&response[i]) if err != nil { - c.logger.Error(ctx, err, - v1.LogField("msg", "sfcv2: ListInstances skipping instance due to SSH error"), - v1.LogField("instanceID", inst.ID), - ) - continue + return nil, errors.WrapAndTrace(err) } - - brevInst, err := c.sfcInstanceToBrevInstance(&inst, sshInfo) - if err != nil { - c.logger.Error(ctx, err, - v1.LogField("msg", "sfcv2: ListInstances skipping instance due to conversion error"), - v1.LogField("instanceID", inst.ID), - ) - continue - } - instances = append(instances, *brevInst) + instances = append(instances, *instance) } c.logger.Debug(ctx, "sfcv2: ListInstances end", v1.LogField("instance count", len(instances)), ) - return instances, nil } @@ -181,99 +117,81 @@ func (c *SFCClientV2) TerminateInstance(ctx context.Context, id v1.CloudProvider v1.LogField("instanceID", id), ) - _, err := c.client.Instances.TerminateInstance(ctx, string(id)) - if err != nil { - return errors.WrapAndTrace(err) + if _, err := c.client.terminateInstance(ctx, c.workspace, string(id)); err != nil { + return wrapIntegrationError(err) } c.logger.Debug(ctx, "sfcv2: TerminateInstance end", v1.LogField("instanceID", id), ) - return nil } -func (c *SFCClientV2) getSSHInfo(ctx context.Context, id string, status components.InstanceStatus) (*components.InstanceSSHInfo, error) { - if status != components.InstanceStatusRunning { - return nil, nil - } - - resp, err := c.client.Instances.GetSSHInfoForInstance(ctx, id) +func integrationInstanceToBrevInstance(instance *integrationInstance) (*v1.Instance, error) { + diskBytes := v1.NewBytes(1500, v1.Gigabyte) + diskGB, err := diskBytes.ByteCountInUnitInt64(v1.Gibibyte) if err != nil { - return nil, errors.WrapAndTrace(err) - } - if resp.InstanceSSHInfo == nil { - return nil, nil - } - - return resp.InstanceSSHInfo, nil -} - -func (c *SFCClientV2) sfcInstanceToBrevInstance(inst *components.InstanceResponse, sshInfo *components.InstanceSSHInfo) (*v1.Instance, error) { - tags, _ := inst.GetTags().GetOrZero() - - cloudCredRefID := tags[tagKeyCloudCredRefID] - if cloudCredRefID == "" { - cloudCredRefID = c.refID + return nil, err } - userTags := make(v1.Tags) - for k, v := range tags { - switch k { - case tagKeyCloudCredRefID, tagKeyRefID: - default: - userTags[k] = v + sshUser := defaultSSHUsername + var publicHost string + var sshPort int + if instance.Connection != nil { + publicHost = instance.Connection.Hostname + sshPort = instance.Connection.Port + if instance.Connection.Username != "" { + sshUser = instance.Connection.Username } } - status := sfcStatusToLifecycleStatus(inst.Status) - - diskInt64, err := h100InstanceTypeMetadata.diskBytes.ByteCountInUnitInt64(v1.Gibibyte) - if err != nil { - return nil, err - } - diskSize := units.Base2Bytes(diskInt64 * int64(units.Gibibyte)) - return &v1.Instance{ - Name: inst.Name, - CloudID: v1.CloudProviderInstanceID(inst.ID), - RefID: tags[tagKeyRefID], - PublicDNS: sshInfo.GetHostname(), - PublicIP: sshInfo.GetHostname(), - SSHUser: defaultSSHUsername, - SSHPort: int(sshInfo.GetPort()), - CreatedAt: time.Unix(inst.CreatedAt, 0), - DiskSize: diskSize, - DiskSizeBytes: h100InstanceTypeMetadata.diskBytes, + Name: instance.Name, + CloudID: v1.CloudProviderInstanceID(instance.ID), + RefID: instance.RefID, + PublicDNS: publicHost, + PublicIP: publicHost, + SSHUser: sshUser, + SSHPort: sshPort, + CreatedAt: time.Unix(instance.CreatedAt, 0), + DiskSize: units.Base2Bytes(diskGB * int64(units.Gibibyte)), + DiskSizeBytes: diskBytes, Status: v1.Status{ - LifecycleStatus: status, + LifecycleStatus: integrationStatusToLifecycleStatus(instance.Status), }, - InstanceTypeID: h100InstanceTypeMetadata.instanceTypeID, - InstanceType: h100InstanceType, + InstanceTypeID: integrationInstanceTypeID(instance.InstanceType), + InstanceType: instance.InstanceType, Location: sfcLocation, Spot: false, Stoppable: false, Rebootable: false, - CloudCredRefID: cloudCredRefID, - Tags: userTags, + CloudCredRefID: instance.CloudCredentialRefID, + Tags: maps.Clone(instance.Tags), }, nil } -func sfcStatusToLifecycleStatus(status components.InstanceStatus) v1.LifecycleStatus { +func integrationStatusToLifecycleStatus(status string) v1.LifecycleStatus { switch status { - case components.InstanceStatusAwaitingAllocation: - return v1.LifecycleStatusPending - case components.InstanceStatusRunning: + case "running": return v1.LifecycleStatusRunning - case components.InstanceStatusTerminated: + case "terminated": return v1.LifecycleStatusTerminated - case components.InstanceStatusFailed: + case "failed": return v1.LifecycleStatusFailed default: return v1.LifecycleStatusPending } } +func wrapIntegrationError(err error) error { + var apiError *integrationAPIError + if errors.As(err, &apiError) && + (apiError.Code == "instance_not_found" || apiError.StatusCode == 404) { + return errors.WrapAndTrace(errors.Join(v1.ErrResourceNotFound, err)) + } + return errors.WrapAndTrace(err) +} + func (c *SFCClientV2) RebootInstance(_ context.Context, _ v1.CloudProviderInstanceID) error { return v1.ErrNotImplemented } diff --git a/v1/providers/sfcomputev2/instancetype.go b/v1/providers/sfcomputev2/instancetype.go index bda56e4..f50a3ae 100644 --- a/v1/providers/sfcomputev2/instancetype.go +++ b/v1/providers/sfcomputev2/instancetype.go @@ -3,105 +3,85 @@ package v2 import ( "context" "fmt" - "sort" "time" "github.com/alecthomas/units" "github.com/bojanz/currency" "github.com/brevdev/cloud/internal/errors" v1 "github.com/brevdev/cloud/v1" - "github.com/sfcompute/sfc-go/models/components" - "github.com/sfcompute/sfc-go/models/operations" ) const ( h100InstanceType = "h100.ib" - sfcVCPU = 112 - sfcGPUCount = 8 sfcLocation = "sfc" diskTypeSSD = "ssd" - formFactorSXM5 = "sxm5" ) -type sfcInstanceTypeMetadata struct { - diskBytes v1.Bytes - memoryBytes v1.Bytes - gpuVRAM v1.Bytes - vcpu int32 - gpuCount int32 - gpuManufacturer v1.Manufacturer - architecture v1.Architecture - deployTime time.Duration - price currency.Amount - instanceTypeID v1.InstanceTypeID +var h100InstanceTypeID = integrationInstanceTypeID(h100InstanceType) + +func integrationInstanceTypeID(instanceType string) v1.InstanceTypeID { + return v1.MakeGenericInstanceTypeID(v1.InstanceType{ + Location: sfcLocation, + Type: instanceType, + }) } -var h100InstanceTypeMetadata = func() sfcInstanceTypeMetadata { - price, err := currency.NewAmount("24.99", "USD") +func buildInstanceType(source integrationInstanceType) (v1.InstanceType, error) { + price, err := currency.NewAmount(source.PriceDollarsPerHour, "USD") if err != nil { - panic(err) + return v1.InstanceType{}, fmt.Errorf("parse SFC instance price: %w", err) } - m := sfcInstanceTypeMetadata{ - diskBytes: v1.NewBytes(1500, v1.Gigabyte), - memoryBytes: v1.NewBytes(960, v1.Gigabyte), - gpuVRAM: v1.NewBytes(80, v1.Gigabyte), - vcpu: sfcVCPU, - gpuCount: sfcGPUCount, - gpuManufacturer: v1.ManufacturerNVIDIA, - architecture: v1.ArchitectureX86_64, - deployTime: 14 * time.Minute, - price: price, - } - - // Compute the instance type ID from a representative InstanceType so it matches - // what Brev expects when validating or storing the type. - it := buildInstanceType(m, true) - m.instanceTypeID = it.ID - return m -}() -func buildInstanceType(m sfcInstanceTypeMetadata, isAvailable bool) v1.InstanceType { - ramInt64, _ := m.memoryBytes.ByteCountInUnitInt64(v1.Gibibyte) - ram := units.Base2Bytes(ramInt64 * int64(units.Gibibyte)) - - vramInt64, _ := m.gpuVRAM.ByteCountInUnitInt64(v1.Gibibyte) - vram := units.Base2Bytes(vramInt64 * int64(units.Gibibyte)) - - diskInt64, _ := m.diskBytes.ByteCountInUnitInt64(v1.Gibibyte) - diskSize := units.Base2Bytes(diskInt64 * int64(units.Gibibyte)) + memoryBytes := v1.NewBytes(v1.BytesValue(source.MemoryGB), v1.Gigabyte) + memoryGB, err := memoryBytes.ByteCountInUnitInt64(v1.Gibibyte) + if err != nil { + return v1.InstanceType{}, err + } + gpuMemoryBytes := v1.NewBytes(v1.BytesValue(source.GPUMemoryGB), v1.Gigabyte) + gpuMemoryGB, err := gpuMemoryBytes.ByteCountInUnitInt64(v1.Gibibyte) + if err != nil { + return v1.InstanceType{}, err + } + diskBytes := v1.NewBytes(v1.BytesValue(source.DiskGB), v1.Gigabyte) + diskGB, err := diskBytes.ByteCountInUnitInt64(v1.Gibibyte) + if err != nil { + return v1.InstanceType{}, err + } + deployTime := time.Duration(source.EstimatedDeploySeconds) * time.Second - it := v1.InstanceType{ - IsAvailable: isAvailable, - Type: h100InstanceType, - Memory: ram, - MemoryBytes: m.memoryBytes, - VCPU: m.vcpu, + return v1.InstanceType{ + ID: integrationInstanceTypeID(source.ID), + IsAvailable: source.AvailableCount > 0, + Type: source.ID, + Memory: units.Base2Bytes(memoryGB * int64(units.Gibibyte)), + MemoryBytes: memoryBytes, + VCPU: source.VCPU, Location: sfcLocation, Stoppable: false, Rebootable: false, IsContainer: false, Provider: CloudProviderID, - BasePrice: &m.price, - EstimatedDeployTime: &m.deployTime, + BasePrice: &price, + EstimatedDeployTime: &deployTime, SupportedGPUs: []v1.GPU{{ - Count: m.gpuCount, - Type: "H100", - Manufacturer: m.gpuManufacturer, - Name: "H100", - Memory: vram, - MemoryBytes: m.gpuVRAM, - NetworkDetails: formFactorSXM5, + Count: source.GPUCount, + Type: source.GPUType, + Manufacturer: v1.ManufacturerNVIDIA, + Name: source.GPUType, + Memory: units.Base2Bytes(gpuMemoryGB * int64(units.Gibibyte)), + MemoryBytes: gpuMemoryBytes, + NetworkDetails: source.GPUNetworkDetails, }}, SupportedStorage: []v1.Storage{{ Type: diskTypeSSD, Count: 1, - Size: diskSize, - SizeBytes: m.diskBytes, + Size: units.Base2Bytes(diskGB * int64(units.Gibibyte)), + SizeBytes: diskBytes, }}, - SupportedArchitectures: []v1.Architecture{m.architecture}, - } - it.ID = v1.MakeGenericInstanceTypeID(it) - return it + SupportedArchitectures: []v1.Architecture{ + v1.GetArchitecture(source.Architecture), + }, + }, nil } func (c *SFCClientV2) GetInstanceTypes(ctx context.Context, args v1.GetInstanceTypeArgs) ([]v1.InstanceType, error) { @@ -109,131 +89,29 @@ func (c *SFCClientV2) GetInstanceTypes(ctx context.Context, args v1.GetInstanceT v1.LogField("location", c.location), ) - available, err := c.availableSlots(ctx) - if err != nil { - return nil, errors.WrapAndTrace(err) - } - - if available <= 0 { - c.logger.Debug(ctx, "sfcv2: GetInstanceTypes no available slots") - return []v1.InstanceType{}, nil - } - - instanceType := buildInstanceType(h100InstanceTypeMetadata, true) - - if !v1.IsSelectedByArgs(instanceType, args) { - return []v1.InstanceType{}, nil - } - - c.logger.Debug(ctx, "sfcv2: GetInstanceTypes end", - v1.LogField("available slots", available), - ) - - return []v1.InstanceType{instanceType}, nil -} - -// skuFreeCapacity returns, per instance SKU in the configured capacity, how many more -// instances can be created on that SKU right now: the SKU's current node allocation minus the -// number of non-terminated instances already on it. Counts are clamped at zero. Reads the -// per-SKU allocation from the capacity schedule and the per-SKU consumption from the instance -// list, issuing exactly two API calls. -func (c *SFCClientV2) skuFreeCapacity(ctx context.Context) (map[string]int, error) { - poolID := c.GetDefaultPoolResourcePath() - - poolResp, err := c.client.Pools.Fetch(ctx, poolID, nil) - if err != nil { - return nil, errors.WrapAndTrace(err) - } - if poolResp.PoolResponse == nil { - return map[string]int{}, nil - } - - now := time.Now().Unix() - free := make(map[string]int) - for skuID, schedule := range poolResp.PoolResponse.AllocationSchedule.ByInstanceSku { - free[skuID] = currentScheduleAllocation(schedule, now) - } - - resp, err := c.client.Instances.List(ctx, operations.ListInstancesRequest{ - Workspace: c.GetWorkspaceResourcePath(), - Pool: []string{poolID}, - }) + response, err := c.client.listInstanceTypes(ctx, c.workspace) if err != nil { - return nil, errors.WrapAndTrace(err) - } - if resp.ListInstancesResponse != nil { - for _, inst := range resp.ListInstancesResponse.Data { - // Every non-terminated instance occupies a slot on its SKU, including failed ones. - if inst.Status == components.InstanceStatusTerminated { - continue - } - sku, ok := inst.GetInstanceSku().Get() - if !ok || sku == nil { - continue - } - free[sku.ID]-- - } - } - - for skuID, n := range free { - free[skuID] = max(n, 0) + return nil, wrapIntegrationError(err) } - return free, nil -} -// currentScheduleAllocation returns the NodeCount from the schedule entry whose -// [StartAt, EndAt) range is currently in effect. EndAt is null only on the final, unbounded -// entry. Returns 0 if no entry is in effect. -func currentScheduleAllocation(schedule []components.ScheduleEntry, now int64) int { - for _, entry := range schedule { - if entry.StartAt > now { + instanceTypes := make([]v1.InstanceType, 0, len(response)) + for _, source := range response { + if source.AvailableCount <= 0 { continue } - // A set, non-null EndAt bounds the range; the final entry's null EndAt is unbounded. - if endAt, ok := entry.EndAt.Get(); ok && endAt != nil && now >= *endAt { - continue + instanceType, err := buildInstanceType(source) + if err != nil { + return nil, errors.WrapAndTrace(err) } - return entry.NodeCount - } - return 0 -} - -// availableSlots returns how many more instances can be created in the configured capacity, -// summed across every SKU. -func (c *SFCClientV2) availableSlots(ctx context.Context) (int, error) { - free, err := c.skuFreeCapacity(ctx) - if err != nil { - return 0, errors.WrapAndTrace(err) - } - total := 0 - for _, n := range free { - total += n - } - return total, nil -} - -// selectAvailableSku returns an instance SKU in the configured capacity that still has a free -// node. SKUs are considered in sorted order so selection is deterministic; since the goal is to -// fully consume every SKU and the order doesn't matter, this drains one SKU before moving to the -// next. Returns an error if no SKU has free capacity. -func (c *SFCClientV2) selectAvailableSku(ctx context.Context) (string, error) { - free, err := c.skuFreeCapacity(ctx) - if err != nil { - return "", errors.WrapAndTrace(err) - } - - skuIDs := make([]string, 0, len(free)) - for skuID := range free { - skuIDs = append(skuIDs, skuID) - } - sort.Strings(skuIDs) - - for _, skuID := range skuIDs { - if free[skuID] > 0 { - return skuID, nil + if v1.IsSelectedByArgs(instanceType, args) { + instanceTypes = append(instanceTypes, instanceType) } } - return "", fmt.Errorf("no instance SKU with available capacity in %s", c.GetDefaultPoolResourcePath()) + + c.logger.Debug(ctx, "sfcv2: GetInstanceTypes end", + v1.LogField("instance type count", len(instanceTypes)), + ) + return instanceTypes, nil } func (c *SFCClientV2) GetLocations(_ context.Context, _ v1.GetLocationsArgs) ([]v1.Location, error) { diff --git a/v1/providers/sfcomputev2/integration_client.go b/v1/providers/sfcomputev2/integration_client.go new file mode 100644 index 0000000..b48e9ea --- /dev/null +++ b/v1/providers/sfcomputev2/integration_client.go @@ -0,0 +1,234 @@ +package v2 + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "net/url" + "strings" + "time" +) + +const ( + defaultAPIURL = "https://api.sfcompute.com" + integrationAPI = "/integrations/brev/v1" +) + +type integrationClient struct { + baseURL string + apiKey string + httpClient *http.Client +} + +type integrationInstanceType struct { + ID string `json:"id"` + AvailableCount int64 `json:"available_count"` + VCPU int32 `json:"vcpu"` + MemoryGB int64 `json:"memory_gb"` + DiskGB int64 `json:"disk_gb"` + GPUCount int32 `json:"gpu_count"` + GPUType string `json:"gpu_type"` + GPUMemoryGB int64 `json:"gpu_memory_gb"` + GPUNetworkDetails string `json:"gpu_network_details"` + Architecture string `json:"architecture"` + PriceDollarsPerHour string `json:"price_dollars_per_hour"` + EstimatedDeploySeconds int64 `json:"estimated_deploy_seconds"` +} + +type integrationInstanceTypesResponse struct { + Data []integrationInstanceType `json:"data"` +} + +type integrationCreateInstanceRequest struct { + Workspace string `json:"workspace"` + RefID string `json:"ref_id"` + CloudCredentialRefID string `json:"cloud_credential_ref_id"` + Name *string `json:"name,omitempty"` + InstanceType string `json:"instance_type"` + SSHPublicKey string `json:"ssh_public_key"` + Tags map[string]string `json:"tags,omitempty"` +} + +type integrationInstance struct { + ID string `json:"id"` + Name string `json:"name"` + RefID string `json:"ref_id"` + CloudCredentialRefID string `json:"cloud_credential_ref_id"` + InstanceType string `json:"instance_type"` + Status string `json:"status"` + CreatedAt int64 `json:"created_at"` + Tags map[string]string `json:"tags"` + Connection *integrationInstanceConnection `json:"connection"` +} + +type integrationInstanceConnection struct { + Hostname string `json:"hostname"` + Port int `json:"port"` + Username string `json:"username"` +} + +type integrationInstancesResponse struct { + Data []integrationInstance `json:"data"` +} + +type integrationAPIError struct { + StatusCode int + Type string + Code string + Message string +} + +type integrationErrorResponse struct { + Error struct { + Type string `json:"type"` + Message string `json:"message"` + Details []struct { + Code string `json:"code"` + Message string `json:"message"` + } `json:"details"` + } `json:"error"` +} + +func (e *integrationAPIError) Error() string { + if e.Code != "" { + return fmt.Sprintf("SFC integration API returned %d (%s): %s", e.StatusCode, e.Code, e.Message) + } + return fmt.Sprintf("SFC integration API returned %d: %s", e.StatusCode, e.Message) +} + +func newIntegrationClient(apiKey string) *integrationClient { + return &integrationClient{ + baseURL: strings.TrimRight(defaultAPIURL, "/"), + apiKey: apiKey, + httpClient: &http.Client{ + Timeout: 30 * time.Second, + }, + } +} + +func (c *integrationClient) listInstanceTypes(ctx context.Context, workspace string) ([]integrationInstanceType, error) { + var response integrationInstanceTypesResponse + err := c.do(ctx, http.MethodGet, c.workspacePath("/instance_types", workspace), nil, &response) + return response.Data, err +} + +func (c *integrationClient) createInstance( + ctx context.Context, + request integrationCreateInstanceRequest, +) (*integrationInstance, error) { + var response integrationInstance + if err := c.do(ctx, http.MethodPost, integrationAPI+"/instances", request, &response); err != nil { + return nil, err + } + return &response, nil +} + +func (c *integrationClient) listInstances(ctx context.Context, workspace string) ([]integrationInstance, error) { + var response integrationInstancesResponse + err := c.do(ctx, http.MethodGet, c.workspacePath("/instances", workspace), nil, &response) + return response.Data, err +} + +func (c *integrationClient) getInstance(ctx context.Context, workspace, id string) (*integrationInstance, error) { + var response integrationInstance + path := integrationAPI + "/instances/" + url.PathEscape(id) + if err := c.do(ctx, http.MethodGet, addWorkspace(path, workspace), nil, &response); err != nil { + return nil, err + } + return &response, nil +} + +func (c *integrationClient) terminateInstance(ctx context.Context, workspace, id string) (*integrationInstance, error) { + var response integrationInstance + path := integrationAPI + "/instances/" + url.PathEscape(id) + "/terminate" + if err := c.do(ctx, http.MethodPost, addWorkspace(path, workspace), nil, &response); err != nil { + return nil, err + } + return &response, nil +} + +func (c *integrationClient) workspacePath(path, workspace string) string { + return addWorkspace(integrationAPI+path, workspace) +} + +func addWorkspace(path, workspace string) string { + query := url.Values{} + query.Set("workspace", workspace) + return path + "?" + query.Encode() +} + +func (c *integrationClient) do( + ctx context.Context, + method string, + path string, + body any, + response any, +) error { + var requestBody io.Reader + if body != nil { + encoded, err := json.Marshal(body) + if err != nil { + return fmt.Errorf("encode SFC integration request: %w", err) + } + requestBody = bytes.NewReader(encoded) + } + + request, err := http.NewRequestWithContext( + ctx, + method, + strings.TrimRight(c.baseURL, "/")+path, + requestBody, + ) + if err != nil { + return fmt.Errorf("create SFC integration request: %w", err) + } + request.Header.Set("Accept", "application/json") + request.Header.Set("Authorization", "Bearer "+c.apiKey) + if body != nil { + request.Header.Set("Content-Type", "application/json") + } + + result, err := c.httpClient.Do(request) + if err != nil { + return fmt.Errorf("call SFC integration API: %w", err) + } + defer func() { _ = result.Body.Close() }() + + if result.StatusCode < http.StatusOK || result.StatusCode >= http.StatusMultipleChoices { + return decodeIntegrationError(result) + } + if response == nil { + _, err = io.Copy(io.Discard, result.Body) + return err + } + if err := json.NewDecoder(result.Body).Decode(response); err != nil { + return fmt.Errorf("decode SFC integration response: %w", err) + } + return nil +} + +func decodeIntegrationError(response *http.Response) error { + var body integrationErrorResponse + if err := json.NewDecoder(io.LimitReader(response.Body, 1<<20)).Decode(&body); err != nil { + return &integrationAPIError{ + StatusCode: response.StatusCode, + Message: http.StatusText(response.StatusCode), + } + } + + apiError := &integrationAPIError{ + StatusCode: response.StatusCode, + Type: body.Error.Type, + Message: body.Error.Message, + } + if len(body.Error.Details) > 0 { + apiError.Code = body.Error.Details[0].Code + if apiError.Message == "" { + apiError.Message = body.Error.Details[0].Message + } + } + return apiError +} diff --git a/v1/providers/sfcomputev2/integration_client_test.go b/v1/providers/sfcomputev2/integration_client_test.go new file mode 100644 index 0000000..055af3d --- /dev/null +++ b/v1/providers/sfcomputev2/integration_client_test.go @@ -0,0 +1,215 @@ +package v2 + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "sync/atomic" + "testing" + + v1 "github.com/brevdev/cloud/v1" + "github.com/stretchr/testify/require" +) + +const testInstanceJSON = `{ + "id": "inst_test", + "name": "training-node", + "ref_id": "brev-instance-ref", + "cloud_credential_ref_id": "credential-ref", + "instance_type": "h100.ib", + "status": "running", + "created_at": 1785283200, + "tags": {"team": "training"}, + "connection": { + "hostname": "203.0.113.10", + "port": 2222, + "username": "ubuntu" + } +}` + +func newTestClient(t *testing.T, handler http.HandlerFunc) *SFCClientV2 { + t.Helper() + + server := httptest.NewServer(handler) + t.Cleanup(server.Close) + + credential := NewSFCCredentialV2( + "credential-ref", + "test-token", + "sfcompute", + "brev-production", + ) + client, err := credential.MakeClientWithOptions( + context.Background(), + sfcLocation, + WithAPIURL(server.URL), + ) + require.NoError(t, err) + + sfcClient, ok := client.(*SFCClientV2) + require.True(t, ok) + return sfcClient +} + +func requireIntegrationRequest(t *testing.T, request *http.Request, method, path string) { + t.Helper() + require.Equal(t, method, request.Method) + require.Equal(t, path, request.URL.Path) + require.Equal(t, "Bearer test-token", request.Header.Get("Authorization")) +} + +func TestCreateInstanceUsesIntegrationContract(t *testing.T) { + t.Parallel() + + client := newTestClient(t, func(writer http.ResponseWriter, request *http.Request) { + requireIntegrationRequest(t, request, http.MethodPost, integrationAPI+"/instances") + + var body map[string]any + require.NoError(t, json.NewDecoder(request.Body).Decode(&body)) + require.Equal(t, "brev-production", body["workspace"]) + require.Equal(t, "brev-instance-ref", body["ref_id"]) + require.Equal(t, "credential-ref", body["cloud_credential_ref_id"]) + require.Equal(t, "h100.ib", body["instance_type"]) + require.Equal(t, "ssh-ed25519 YWJj", body["ssh_public_key"]) + require.NotContains(t, body, "pool") + require.NotContains(t, body, "instance_sku") + require.NotContains(t, body, "image") + require.NotContains(t, body, "cloud_init_user_data") + + writer.Header().Set("Content-Type", "application/json") + writer.WriteHeader(http.StatusCreated) + _, _ = writer.Write([]byte(testInstanceJSON)) + }) + + instance, err := client.CreateInstance(context.Background(), v1.CreateInstanceAttrs{ + Name: "training node", + RefID: "brev-instance-ref", + InstanceType: h100InstanceType, + PublicKey: "ssh-ed25519 YWJj", + Tags: v1.Tags{"team": "training"}, + }) + require.NoError(t, err) + require.Equal(t, v1.CloudProviderInstanceID("inst_test"), instance.CloudID) + require.Equal(t, "203.0.113.10", instance.PublicIP) + require.Equal(t, 2222, instance.SSHPort) + require.Equal(t, v1.LifecycleStatusRunning, instance.Status.LifecycleStatus) + require.Equal(t, h100InstanceTypeID, instance.InstanceTypeID) +} + +func TestListInstancesUsesOneRequest(t *testing.T) { + t.Parallel() + + var requestCount atomic.Int32 + client := newTestClient(t, func(writer http.ResponseWriter, request *http.Request) { + requestCount.Add(1) + requireIntegrationRequest(t, request, http.MethodGet, integrationAPI+"/instances") + require.Equal(t, "brev-production", request.URL.Query().Get("workspace")) + + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{"object":"list","data":[` + testInstanceJSON + `]}`)) + }) + + instances, err := client.ListInstances(context.Background(), v1.ListInstancesArgs{ + InstanceIDs: []v1.CloudProviderInstanceID{"inst_test"}, + }) + require.NoError(t, err) + require.Len(t, instances, 1) + require.Equal(t, int32(1), requestCount.Load()) + require.Equal(t, "203.0.113.10", instances[0].PublicDNS) +} + +func TestGetAndTerminateInstance(t *testing.T) { + t.Parallel() + + var methods []string + client := newTestClient(t, func(writer http.ResponseWriter, request *http.Request) { + methods = append(methods, request.Method) + require.Equal(t, "brev-production", request.URL.Query().Get("workspace")) + + switch request.Method { + case http.MethodGet: + requireIntegrationRequest(t, request, http.MethodGet, integrationAPI+"/instances/inst_test") + case http.MethodPost: + requireIntegrationRequest(t, request, http.MethodPost, integrationAPI+"/instances/inst_test/terminate") + default: + t.Fatalf("unexpected method %s", request.Method) + } + + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(testInstanceJSON)) + }) + + instance, err := client.GetInstance(context.Background(), "inst_test") + require.NoError(t, err) + require.Equal(t, "inst_test", string(instance.CloudID)) + require.NoError(t, client.TerminateInstance(context.Background(), "inst_test")) + require.Equal(t, []string{http.MethodGet, http.MethodPost}, methods) +} + +func TestGetInstanceTypesUsesServerMetadata(t *testing.T) { + t.Parallel() + + client := newTestClient(t, func(writer http.ResponseWriter, request *http.Request) { + requireIntegrationRequest(t, request, http.MethodGet, integrationAPI+"/instance_types") + require.Equal(t, "brev-production", request.URL.Query().Get("workspace")) + + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{ + "object": "list", + "data": [{ + "id": "h100.ib", + "available_count": 2, + "vcpu": 112, + "memory_gb": 960, + "disk_gb": 1500, + "gpu_count": 8, + "gpu_type": "H100", + "gpu_memory_gb": 80, + "gpu_network_details": "SXM5", + "architecture": "x86_64", + "price_dollars_per_hour": "24.99", + "estimated_deploy_seconds": 840 + }] + }`)) + }) + + instanceTypes, err := client.GetInstanceTypes(context.Background(), v1.GetInstanceTypeArgs{}) + require.NoError(t, err) + require.Len(t, instanceTypes, 1) + require.Equal(t, h100InstanceTypeID, instanceTypes[0].ID) + require.Equal(t, h100InstanceType, instanceTypes[0].Type) + require.Equal(t, int32(112), instanceTypes[0].VCPU) + require.Equal(t, int32(8), instanceTypes[0].SupportedGPUs[0].Count) + require.Equal(t, "SXM5", instanceTypes[0].SupportedGPUs[0].NetworkDetails) + require.Equal(t, 14*60, int(instanceTypes[0].EstimatedDeployTime.Seconds())) +} + +func TestIntegrationErrorsPreserveCodes(t *testing.T) { + t.Parallel() + + client := newTestClient(t, func(writer http.ResponseWriter, request *http.Request) { + requireIntegrationRequest(t, request, http.MethodGet, integrationAPI+"/instances/missing") + writer.Header().Set("Content-Type", "application/json") + writer.WriteHeader(http.StatusNotFound) + _, _ = writer.Write([]byte(`{ + "error": { + "type": "not_found", + "message": "Brev instance not found", + "details": [{ + "code": "instance_not_found", + "message": "instance is not managed by Brev" + }] + } + }`)) + }) + + _, err := client.GetInstance(context.Background(), "missing") + require.Error(t, err) + require.ErrorIs(t, err, v1.ErrResourceNotFound) + + var apiError *integrationAPIError + require.True(t, errors.As(err, &apiError)) + require.Equal(t, "instance_not_found", apiError.Code) +} diff --git a/v1/providers/sfcomputev2/validation_test.go b/v1/providers/sfcomputev2/validation_test.go index 310e45c..c22a085 100644 --- a/v1/providers/sfcomputev2/validation_test.go +++ b/v1/providers/sfcomputev2/validation_test.go @@ -14,7 +14,7 @@ func TestValidationFunctions(t *testing.T) { config := validation.ProviderConfig{ Credential: NewSFCCredentialV2("validation-test", getAPIKey(), getOrganization(), getWorkspace()), StableIDs: []v1.InstanceTypeID{ - h100InstanceTypeMetadata.instanceTypeID, + h100InstanceTypeID, }, } From 6e4ee192e07607889aa38b2cc3ee8484e2adcfaf Mon Sep 17 00:00:00 2001 From: Mason Wheeler Date: Wed, 29 Jul 2026 23:09:46 -0700 Subject: [PATCH 2/3] fix(sfcompute): preserve integration lifecycle semantics --- go.mod | 2 - go.sum | 4 - v1/providers/sfcomputev2/capabilities.go | 1 + v1/providers/sfcomputev2/instance.go | 29 +++-- .../sfcomputev2/integration_client.go | 47 +++++++- .../sfcomputev2/integration_client_test.go | 113 +++++++++++++++++- 6 files changed, 176 insertions(+), 20 deletions(-) diff --git a/go.mod b/go.mod index b44c82e..0a3771c 100644 --- a/go.mod +++ b/go.mod @@ -20,7 +20,6 @@ require ( github.com/nebius/gosdk v0.2.22 github.com/pkg/errors v0.9.1 github.com/sfcompute/nodes-go v0.1.0-alpha.4 - github.com/sfcompute/sfc-go v0.1.0-preview.3 github.com/stretchr/testify v1.11.1 golang.org/x/crypto v0.52.0 golang.org/x/text v0.37.0 @@ -84,7 +83,6 @@ require ( github.com/sirupsen/logrus v1.9.3 // indirect github.com/spf13/afero v1.15.0 // indirect github.com/spf13/pflag v1.0.10 // indirect - github.com/spyzhov/ajson v0.8.0 // indirect github.com/tidwall/gjson v1.18.0 // indirect github.com/tidwall/match v1.1.1 // indirect github.com/tidwall/pretty v1.2.1 // indirect diff --git a/go.sum b/go.sum index 61bb01b..f320ec4 100644 --- a/go.sum +++ b/go.sum @@ -156,16 +156,12 @@ github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0t github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc= github.com/sfcompute/nodes-go v0.1.0-alpha.4 h1:oFBWcMPSpqLYm/NDs5I1jTvzgx9rsXDL9Ghsm30Hc0Q= github.com/sfcompute/nodes-go v0.1.0-alpha.4/go.mod h1:nUviHgK+Fgt2hDFcRL3M8VoyiypC8fc0dsY8C30QU8M= -github.com/sfcompute/sfc-go v0.1.0-preview.3 h1:azKThmbm9ljQ+z8RP4039XwV4bJMTcYKNpKcxrpNf5A= -github.com/sfcompute/sfc-go v0.1.0-preview.3/go.mod h1:SDgYqB2R6gFM+bzLBeF/Fb+J1HHaTlDuStSkiFuMWDU= github.com/sirupsen/logrus v1.9.3 h1:dueUQJ1C2q9oE3F7wvmSGAaVtTmUizReu6fjN8uqzbQ= github.com/sirupsen/logrus v1.9.3/go.mod h1:naHLuLoDiP4jHNo9R0sCBMtWGeIprob74mVsIT4qYEQ= github.com/spf13/afero v1.15.0 h1:b/YBCLWAJdFWJTN9cLhiXXcD7mzKn9Dm86dNnfyQw1I= github.com/spf13/afero v1.15.0/go.mod h1:NC2ByUVxtQs4b3sIUphxK0NioZnmxgyCrfzeuq8lxMg= github.com/spf13/pflag v1.0.10 h1:4EBh2KAYBwaONj6b2Ye1GiHfwjqyROoF4RwYO+vPwFk= github.com/spf13/pflag v1.0.10/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= -github.com/spyzhov/ajson v0.8.0 h1:sFXyMbi4Y/BKjrsfkUZHSjA2JM1184enheSjjoT/zCc= -github.com/spyzhov/ajson v0.8.0/go.mod h1:63V+CGM6f1Bu/p4nLIN8885ojBdt88TbLoSFzyqMuVA= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw= github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= diff --git a/v1/providers/sfcomputev2/capabilities.go b/v1/providers/sfcomputev2/capabilities.go index e9b62d6..1ed2377 100644 --- a/v1/providers/sfcomputev2/capabilities.go +++ b/v1/providers/sfcomputev2/capabilities.go @@ -9,6 +9,7 @@ import ( func getSFCCapabilitiesV2() v1.Capabilities { return v1.Capabilities{ v1.CapabilityCreateInstance, + v1.CapabilityCreateIdempotentInstance, v1.CapabilityTerminateInstance, v1.CapabilityCreateTerminateInstance, v1.CapabilityTags, diff --git a/v1/providers/sfcomputev2/instance.go b/v1/providers/sfcomputev2/instance.go index ce8c172..15b26cf 100644 --- a/v1/providers/sfcomputev2/instance.go +++ b/v1/providers/sfcomputev2/instance.go @@ -88,7 +88,11 @@ func (c *SFCClientV2) ListInstances(ctx context.Context, args v1.ListInstancesAr v1.LogField("location", c.location), ) - response, err := c.client.listInstances(ctx, c.workspace) + instanceIDs := make([]string, len(args.InstanceIDs)) + for i, id := range args.InstanceIDs { + instanceIDs[i] = string(id) + } + response, err := c.client.listInstances(ctx, c.workspace, instanceIDs) if err != nil { return nil, wrapIntegrationError(err) } @@ -128,7 +132,7 @@ func (c *SFCClientV2) TerminateInstance(ctx context.Context, id v1.CloudProvider } func integrationInstanceToBrevInstance(instance *integrationInstance) (*v1.Instance, error) { - diskBytes := v1.NewBytes(1500, v1.Gigabyte) + diskBytes := v1.NewBytes(v1.BytesValue(instance.DiskGB), v1.Gigabyte) diskGB, err := diskBytes.ByteCountInUnitInt64(v1.Gibibyte) if err != nil { return nil, err @@ -137,13 +141,20 @@ func integrationInstanceToBrevInstance(instance *integrationInstance) (*v1.Insta sshUser := defaultSSHUsername var publicHost string var sshPort int - if instance.Connection != nil { + hasConnection := instance.Connection != nil && + instance.Connection.Hostname != "" && + instance.Connection.Port > 0 + if hasConnection { publicHost = instance.Connection.Hostname sshPort = instance.Connection.Port if instance.Connection.Username != "" { sshUser = instance.Connection.Username } } + status := instance.Status + if status == "running" && !hasConnection { + status = "pending" + } return &v1.Instance{ Name: instance.Name, @@ -157,7 +168,7 @@ func integrationInstanceToBrevInstance(instance *integrationInstance) (*v1.Insta DiskSize: units.Base2Bytes(diskGB * int64(units.Gibibyte)), DiskSizeBytes: diskBytes, Status: v1.Status{ - LifecycleStatus: integrationStatusToLifecycleStatus(instance.Status), + LifecycleStatus: integrationStatusToLifecycleStatus(status), }, InstanceTypeID: integrationInstanceTypeID(instance.InstanceType), InstanceType: instance.InstanceType, @@ -185,9 +196,13 @@ func integrationStatusToLifecycleStatus(status string) v1.LifecycleStatus { func wrapIntegrationError(err error) error { var apiError *integrationAPIError - if errors.As(err, &apiError) && - (apiError.Code == "instance_not_found" || apiError.StatusCode == 404) { - return errors.WrapAndTrace(errors.Join(v1.ErrResourceNotFound, err)) + if errors.As(err, &apiError) { + switch apiError.Code { + case "instance_not_found": + return errors.WrapAndTrace(errors.Join(v1.ErrResourceNotFound, err)) + case "capacity_exhausted": + return errors.WrapAndTrace(errors.Join(v1.ErrInsufficientResources, err)) + } } return errors.WrapAndTrace(err) } diff --git a/v1/providers/sfcomputev2/integration_client.go b/v1/providers/sfcomputev2/integration_client.go index b48e9ea..604f002 100644 --- a/v1/providers/sfcomputev2/integration_client.go +++ b/v1/providers/sfcomputev2/integration_client.go @@ -13,8 +13,9 @@ import ( ) const ( - defaultAPIURL = "https://api.sfcompute.com" - integrationAPI = "/integrations/brev/v1" + defaultAPIURL = "https://api.sfcompute.com" + integrationAPI = "/integrations/brev/v1" + maxInstanceIDsPerGet = 200 ) type integrationClient struct { @@ -58,6 +59,7 @@ type integrationInstance struct { RefID string `json:"ref_id"` CloudCredentialRefID string `json:"cloud_credential_ref_id"` InstanceType string `json:"instance_type"` + DiskGB int64 `json:"disk_gb"` Status string `json:"status"` CreatedAt int64 `json:"created_at"` Tags map[string]string `json:"tags"` @@ -126,9 +128,46 @@ func (c *integrationClient) createInstance( return &response, nil } -func (c *integrationClient) listInstances(ctx context.Context, workspace string) ([]integrationInstance, error) { +func (c *integrationClient) listInstances( + ctx context.Context, + workspace string, + instanceIDs []string, +) ([]integrationInstance, error) { + if len(instanceIDs) == 0 { + return c.listInstanceBatch(ctx, workspace, nil) + } + + instances := make([]integrationInstance, 0, len(instanceIDs)) + for start := 0; start < len(instanceIDs); start += maxInstanceIDsPerGet { + end := min(start+maxInstanceIDsPerGet, len(instanceIDs)) + batch, err := c.listInstanceBatch(ctx, workspace, instanceIDs[start:end]) + if err != nil { + return nil, err + } + instances = append(instances, batch...) + } + return instances, nil +} + +func (c *integrationClient) listInstanceBatch( + ctx context.Context, + workspace string, + instanceIDs []string, +) ([]integrationInstance, error) { + query := url.Values{} + query.Set("workspace", workspace) + for _, id := range instanceIDs { + query.Add("id", id) + } + var response integrationInstancesResponse - err := c.do(ctx, http.MethodGet, c.workspacePath("/instances", workspace), nil, &response) + err := c.do( + ctx, + http.MethodGet, + integrationAPI+"/instances?"+query.Encode(), + nil, + &response, + ) return response.Data, err } diff --git a/v1/providers/sfcomputev2/integration_client_test.go b/v1/providers/sfcomputev2/integration_client_test.go index 055af3d..7cac1f5 100644 --- a/v1/providers/sfcomputev2/integration_client_test.go +++ b/v1/providers/sfcomputev2/integration_client_test.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "errors" + "fmt" "net/http" "net/http/httptest" "sync/atomic" @@ -19,6 +20,7 @@ const testInstanceJSON = `{ "ref_id": "brev-instance-ref", "cloud_credential_ref_id": "credential-ref", "instance_type": "h100.ib", + "disk_gb": 2048, "status": "running", "created_at": 1785283200, "tags": {"team": "training"}, @@ -96,6 +98,7 @@ func TestCreateInstanceUsesIntegrationContract(t *testing.T) { require.Equal(t, 2222, instance.SSHPort) require.Equal(t, v1.LifecycleStatusRunning, instance.Status.LifecycleStatus) require.Equal(t, h100InstanceTypeID, instance.InstanceTypeID) + require.True(t, instance.DiskSizeBytes.Equal(v1.NewBytes(2048, v1.Gigabyte))) } func TestListInstancesUsesOneRequest(t *testing.T) { @@ -106,13 +109,14 @@ func TestListInstancesUsesOneRequest(t *testing.T) { requestCount.Add(1) requireIntegrationRequest(t, request, http.MethodGet, integrationAPI+"/instances") require.Equal(t, "brev-production", request.URL.Query().Get("workspace")) + require.Equal(t, []string{"inst_test", "inst_other"}, request.URL.Query()["id"]) writer.Header().Set("Content-Type", "application/json") _, _ = writer.Write([]byte(`{"object":"list","data":[` + testInstanceJSON + `]}`)) }) instances, err := client.ListInstances(context.Background(), v1.ListInstancesArgs{ - InstanceIDs: []v1.CloudProviderInstanceID{"inst_test"}, + InstanceIDs: []v1.CloudProviderInstanceID{"inst_test", "inst_other"}, }) require.NoError(t, err) require.Len(t, instances, 1) @@ -120,6 +124,51 @@ func TestListInstancesUsesOneRequest(t *testing.T) { require.Equal(t, "203.0.113.10", instances[0].PublicDNS) } +func TestListInstancesChunksLargeIDFilters(t *testing.T) { + t.Parallel() + + var batchSizes []int + client := newTestClient(t, func(writer http.ResponseWriter, request *http.Request) { + requireIntegrationRequest(t, request, http.MethodGet, integrationAPI+"/instances") + batchSizes = append(batchSizes, len(request.URL.Query()["id"])) + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write([]byte(`{"object":"list","data":[]}`)) + }) + + instanceIDs := make([]string, maxInstanceIDsPerGet+1) + for i := range instanceIDs { + instanceIDs[i] = fmt.Sprintf("inst_%d", i) + } + instances, err := client.client.listInstances( + context.Background(), + "brev-production", + instanceIDs, + ) + require.NoError(t, err) + require.Empty(t, instances) + require.Equal(t, []int{maxInstanceIDsPerGet, 1}, batchSizes) +} + +func TestRunningInstanceWithoutConnectionRemainsPending(t *testing.T) { + t.Parallel() + + client := newTestClient(t, func(writer http.ResponseWriter, request *http.Request) { + requireIntegrationRequest(t, request, http.MethodGet, integrationAPI+"/instances/inst_test") + + var response map[string]any + require.NoError(t, json.Unmarshal([]byte(testInstanceJSON), &response)) + response["connection"] = nil + writer.Header().Set("Content-Type", "application/json") + require.NoError(t, json.NewEncoder(writer).Encode(response)) + }) + + instance, err := client.GetInstance(context.Background(), "inst_test") + require.NoError(t, err) + require.Equal(t, v1.LifecycleStatusPending, instance.Status.LifecycleStatus) + require.Empty(t, instance.PublicIP) + require.Zero(t, instance.SSHPort) +} + func TestGetAndTerminateInstance(t *testing.T) { t.Parallel() @@ -167,7 +216,7 @@ func TestGetInstanceTypesUsesServerMetadata(t *testing.T) { "gpu_count": 8, "gpu_type": "H100", "gpu_memory_gb": 80, - "gpu_network_details": "SXM5", + "gpu_network_details": "sxm5", "architecture": "x86_64", "price_dollars_per_hour": "24.99", "estimated_deploy_seconds": 840 @@ -182,7 +231,7 @@ func TestGetInstanceTypesUsesServerMetadata(t *testing.T) { require.Equal(t, h100InstanceType, instanceTypes[0].Type) require.Equal(t, int32(112), instanceTypes[0].VCPU) require.Equal(t, int32(8), instanceTypes[0].SupportedGPUs[0].Count) - require.Equal(t, "SXM5", instanceTypes[0].SupportedGPUs[0].NetworkDetails) + require.Equal(t, "sxm5", instanceTypes[0].SupportedGPUs[0].NetworkDetails) require.Equal(t, 14*60, int(instanceTypes[0].EstimatedDeployTime.Seconds())) } @@ -213,3 +262,61 @@ func TestIntegrationErrorsPreserveCodes(t *testing.T) { require.True(t, errors.As(err, &apiError)) require.Equal(t, "instance_not_found", apiError.Code) } + +func TestGenericNotFoundIsNotAnInstanceNotFound(t *testing.T) { + t.Parallel() + + client := newTestClient(t, func(writer http.ResponseWriter, request *http.Request) { + requireIntegrationRequest(t, request, http.MethodGet, integrationAPI+"/instances/inst_test") + writer.Header().Set("Content-Type", "application/json") + writer.WriteHeader(http.StatusNotFound) + _, _ = writer.Write([]byte(`{ + "error": { + "type": "not_found", + "message": "Brev integration not found", + "details": [{ + "code": "integration_not_enabled", + "message": "the Brev integration is not enabled" + }] + } + }`)) + }) + + _, err := client.GetInstance(context.Background(), "inst_test") + require.Error(t, err) + require.NotErrorIs(t, err, v1.ErrResourceNotFound) +} + +func TestCapacityExhaustedMapsToInsufficientResources(t *testing.T) { + t.Parallel() + + client := newTestClient(t, func(writer http.ResponseWriter, request *http.Request) { + requireIntegrationRequest(t, request, http.MethodPost, integrationAPI+"/instances") + writer.Header().Set("Content-Type", "application/json") + writer.WriteHeader(http.StatusServiceUnavailable) + _, _ = writer.Write([]byte(`{ + "error": { + "type": "service_unavailable", + "message": "no Brev capacity is currently available", + "details": [{ + "code": "capacity_exhausted", + "message": "all eligible H100 capacity is in use" + }] + } + }`)) + }) + + _, err := client.CreateInstance(context.Background(), v1.CreateInstanceAttrs{ + RefID: "brev-instance-ref", + InstanceType: h100InstanceType, + PublicKey: "ssh-ed25519 YWJj", + }) + require.Error(t, err) + require.ErrorIs(t, err, v1.ErrInsufficientResources) +} + +func TestCapabilitiesIncludeIdempotentCreate(t *testing.T) { + t.Parallel() + + require.True(t, getSFCCapabilitiesV2().IsCapable(v1.CapabilityCreateIdempotentInstance)) +} From 6489c8addcfacfce9d175bbc95346f7897f9ffd3 Mon Sep 17 00:00:00 2001 From: Mason Wheeler Date: Thu, 30 Jul 2026 13:56:29 -0700 Subject: [PATCH 3/3] fix(sfcompute): complete integration list contract --- v1/providers/sfcomputev2/capabilities.go | 1 - v1/providers/sfcomputev2/instance.go | 28 +++++- .../sfcomputev2/integration_client.go | 40 ++++++-- .../sfcomputev2/integration_client_test.go | 96 +++++++++++++++++++ 4 files changed, 150 insertions(+), 15 deletions(-) diff --git a/v1/providers/sfcomputev2/capabilities.go b/v1/providers/sfcomputev2/capabilities.go index 1ed2377..d0fbd14 100644 --- a/v1/providers/sfcomputev2/capabilities.go +++ b/v1/providers/sfcomputev2/capabilities.go @@ -12,7 +12,6 @@ func getSFCCapabilitiesV2() v1.Capabilities { v1.CapabilityCreateIdempotentInstance, v1.CapabilityTerminateInstance, v1.CapabilityCreateTerminateInstance, - v1.CapabilityTags, } } diff --git a/v1/providers/sfcomputev2/instance.go b/v1/providers/sfcomputev2/instance.go index 15b26cf..a04c2f8 100644 --- a/v1/providers/sfcomputev2/instance.go +++ b/v1/providers/sfcomputev2/instance.go @@ -99,14 +99,13 @@ func (c *SFCClientV2) ListInstances(ctx context.Context, args v1.ListInstancesAr instances := make([]v1.Instance, 0, len(response)) for i := range response { - if len(args.InstanceIDs) > 0 && - !slices.Contains(args.InstanceIDs, v1.CloudProviderInstanceID(response[i].ID)) { - continue - } instance, err := integrationInstanceToBrevInstance(&response[i]) if err != nil { return nil, errors.WrapAndTrace(err) } + if !matchesListArgs(*instance, args) { + continue + } instances = append(instances, *instance) } @@ -116,6 +115,27 @@ func (c *SFCClientV2) ListInstances(ctx context.Context, args v1.ListInstancesAr return instances, nil } +func matchesListArgs(instance v1.Instance, args v1.ListInstancesArgs) bool { + if len(args.InstanceIDs) > 0 && !slices.Contains(args.InstanceIDs, instance.CloudID) { + return false + } + if len(args.Locations) > 0 && + !args.Locations.IsAll() && + !args.Locations.IsAllowed(instance.Location) { + return false + } + for key, values := range args.TagFilters { + value, ok := instance.Tags[key] + if !ok { + return false + } + if len(values) > 0 && !slices.Contains(values, value) { + return false + } + } + return true +} + func (c *SFCClientV2) TerminateInstance(ctx context.Context, id v1.CloudProviderInstanceID) error { c.logger.Debug(ctx, "sfcv2: TerminateInstance start", v1.LogField("instanceID", id), diff --git a/v1/providers/sfcomputev2/integration_client.go b/v1/providers/sfcomputev2/integration_client.go index 604f002..6bfade7 100644 --- a/v1/providers/sfcomputev2/integration_client.go +++ b/v1/providers/sfcomputev2/integration_client.go @@ -73,7 +73,9 @@ type integrationInstanceConnection struct { } type integrationInstancesResponse struct { - Data []integrationInstance `json:"data"` + Cursor string `json:"cursor"` + HasMore bool `json:"has_more"` + Data []integrationInstance `json:"data"` } type integrationAPIError struct { @@ -156,19 +158,37 @@ func (c *integrationClient) listInstanceBatch( ) ([]integrationInstance, error) { query := url.Values{} query.Set("workspace", workspace) + query.Set("limit", fmt.Sprint(maxInstanceIDsPerGet)) for _, id := range instanceIDs { query.Add("id", id) } - var response integrationInstancesResponse - err := c.do( - ctx, - http.MethodGet, - integrationAPI+"/instances?"+query.Encode(), - nil, - &response, - ) - return response.Data, err + var instances []integrationInstance + var cursor string + for { + if cursor != "" { + query.Set("starting_after", cursor) + } + + var response integrationInstancesResponse + if err := c.do( + ctx, + http.MethodGet, + integrationAPI+"/instances?"+query.Encode(), + nil, + &response, + ); err != nil { + return nil, err + } + instances = append(instances, response.Data...) + if !response.HasMore { + return instances, nil + } + if response.Cursor == "" || response.Cursor == cursor { + return nil, fmt.Errorf("SFC integration API returned an invalid pagination cursor") + } + cursor = response.Cursor + } } func (c *integrationClient) getInstance(ctx context.Context, workspace, id string) (*integrationInstance, error) { diff --git a/v1/providers/sfcomputev2/integration_client_test.go b/v1/providers/sfcomputev2/integration_client_test.go index 7cac1f5..138c7d2 100644 --- a/v1/providers/sfcomputev2/integration_client_test.go +++ b/v1/providers/sfcomputev2/integration_client_test.go @@ -110,6 +110,7 @@ func TestListInstancesUsesOneRequest(t *testing.T) { requireIntegrationRequest(t, request, http.MethodGet, integrationAPI+"/instances") require.Equal(t, "brev-production", request.URL.Query().Get("workspace")) require.Equal(t, []string{"inst_test", "inst_other"}, request.URL.Query()["id"]) + require.Equal(t, "200", request.URL.Query().Get("limit")) writer.Header().Set("Content-Type", "application/json") _, _ = writer.Write([]byte(`{"object":"list","data":[` + testInstanceJSON + `]}`)) @@ -149,6 +150,100 @@ func TestListInstancesChunksLargeIDFilters(t *testing.T) { require.Equal(t, []int{maxInstanceIDsPerGet, 1}, batchSizes) } +func TestListInstancesFollowsPagination(t *testing.T) { + t.Parallel() + + var requestCount atomic.Int32 + client := newTestClient(t, func(writer http.ResponseWriter, request *http.Request) { + requireIntegrationRequest(t, request, http.MethodGet, integrationAPI+"/instances") + page := requestCount.Add(1) + + var instance map[string]any + require.NoError(t, json.Unmarshal([]byte(testInstanceJSON), &instance)) + if page == 1 { + require.Empty(t, request.URL.Query().Get("starting_after")) + instance["id"] = "inst_first" + instance["ref_id"] = "ref-first" + require.NoError(t, json.NewEncoder(writer).Encode(map[string]any{ + "object": "list", + "cursor": "nodec_next", + "has_more": true, + "data": []any{instance}, + })) + return + } + + require.Equal(t, int32(2), page) + require.Equal(t, "nodec_next", request.URL.Query().Get("starting_after")) + instance["id"] = "inst_second" + instance["ref_id"] = "ref-second" + require.NoError(t, json.NewEncoder(writer).Encode(map[string]any{ + "object": "list", + "cursor": "nodec_done", + "has_more": false, + "data": []any{instance}, + })) + }) + + instances, err := client.ListInstances(context.Background(), v1.ListInstancesArgs{}) + require.NoError(t, err) + require.Len(t, instances, 2) + require.Equal(t, v1.CloudProviderInstanceID("inst_first"), instances[0].CloudID) + require.Equal(t, v1.CloudProviderInstanceID("inst_second"), instances[1].CloudID) + require.Equal(t, int32(2), requestCount.Load()) +} + +func TestListInstancesRejectsInvalidPagination(t *testing.T) { + t.Parallel() + + client := newTestClient(t, func(writer http.ResponseWriter, _ *http.Request) { + _, _ = writer.Write([]byte(`{ + "object": "list", + "cursor": null, + "has_more": true, + "data": [] + }`)) + }) + + _, err := client.ListInstances(context.Background(), v1.ListInstancesArgs{}) + require.ErrorContains(t, err, "SFC integration API returned an invalid pagination cursor") +} + +func TestMatchesListArgs(t *testing.T) { + t.Parallel() + + instance := v1.Instance{ + CloudID: "inst_test", + Location: sfcLocation, + Tags: v1.Tags{ + "team": "training", + "env": "test", + }, + } + + require.True(t, matchesListArgs(instance, v1.ListInstancesArgs{ + InstanceIDs: []v1.CloudProviderInstanceID{"inst_test"}, + Locations: v1.LocationsFilter{sfcLocation}, + TagFilters: map[string][]string{"team": {"training"}}, + })) + require.True(t, matchesListArgs(instance, v1.ListInstancesArgs{ + Locations: v1.LocationsFilter{"all"}, + TagFilters: map[string][]string{"env": nil}, + })) + require.False(t, matchesListArgs(instance, v1.ListInstancesArgs{ + InstanceIDs: []v1.CloudProviderInstanceID{"inst_other"}, + })) + require.False(t, matchesListArgs(instance, v1.ListInstancesArgs{ + Locations: v1.LocationsFilter{"other"}, + })) + require.False(t, matchesListArgs(instance, v1.ListInstancesArgs{ + TagFilters: map[string][]string{"team": {"batch"}}, + })) + require.False(t, matchesListArgs(instance, v1.ListInstancesArgs{ + TagFilters: map[string][]string{"missing": nil}, + })) +} + func TestRunningInstanceWithoutConnectionRemainsPending(t *testing.T) { t.Parallel() @@ -319,4 +414,5 @@ func TestCapabilitiesIncludeIdempotentCreate(t *testing.T) { t.Parallel() require.True(t, getSFCCapabilitiesV2().IsCapable(v1.CapabilityCreateIdempotentInstance)) + require.False(t, getSFCCapabilitiesV2().IsCapable(v1.CapabilityTags)) }