// Package gcp provides Google Cloud Platform infrastructure components for Vertex AI Endpoints.
package gcp
import (
"fmt"
namer "github.com/davidmontoyago/commodity-namer"
vertexmodeldeployment "github.com/davidmontoyago/pulumi-gcp-vertex-model-deployment/sdk/go/pulumi-gcp-vertex-model-deployment/resources"
"github.com/pulumi/pulumi-gcp/sdk/v8/go/gcp/artifactregistry"
"github.com/pulumi/pulumi-gcp/sdk/v8/go/gcp/projects"
"github.com/pulumi/pulumi-gcp/sdk/v8/go/gcp/serviceaccount"
"github.com/pulumi/pulumi-gcp/sdk/v8/go/gcp/storage"
"github.com/pulumi/pulumi-gcp/sdk/v8/go/gcp/vertex"
"github.com/pulumi/pulumi/sdk/v3/go/pulumi"
)
// AIEndpoint represents an GCP Vertex AI Endpoint running on the internet.
type AIEndpoint struct {
pulumi.ResourceState
namer.Namer
Project string
Region string
ModelImageURL pulumi.StringOutput
ModelDir string
ModelPredictionInputSchemaPath string
ModelPredictionOutputSchemaPath string
ModelPredictionBehaviorSchemaPath string
ModelBucketBasePath string
ModelDisplayName pulumi.StringOutput
ModelCommandArgs pulumi.StringArrayOutput
EnvVars pulumi.StringMapOutput
MachineType pulumi.StringOutput
AcceleratorType pulumi.StringOutput
AcceleratorCount pulumi.IntOutput
MinReplicaCount pulumi.IntOutput
MaxReplicaCount pulumi.IntOutput
EndpointDisplayName pulumi.StringOutput
ContainerPort pulumi.IntOutput
HealthRoute pulumi.StringOutput
PredictRoute pulumi.StringOutput
EnableAccessLogging pulumi.BoolOutput
DisableContainerLogging pulumi.BoolOutput
EnableSpotVMs pulumi.BoolOutput
DeletionProtection pulumi.BoolOutput
Labels map[string]string
name string
// Core resources
modelServiceAccount *serviceaccount.Account
endpoint *vertex.AiEndpoint
artifactsBucket *storage.Bucket
modelDeployment *vertexmodeldeployment.VertexModelDeployment
uploadedModelArtifacts pulumi.StringArrayOutput
// IAM bindings for the model service account
iamMembers []*projects.IAMMember
bucketIAMMembers []*storage.BucketIAMMember
registryIAMAccess *artifactregistry.RepositoryIamMember
}
// NewAIEndpoint creates a new AIEndpoint instance with the provided configuration.
func NewAIEndpoint(ctx *pulumi.Context, name string, args *AIEndpointArgs, opts ...pulumi.ResourceOption) (*AIEndpoint, error) {
if args.Project == "" {
return nil, fmt.Errorf("project is required")
}
if args.Region == "" {
return nil, fmt.Errorf("region is required")
}
if args.ModelBucketBasePath == "" {
args.ModelBucketBasePath = "model"
}
aiEndpoint := &AIEndpoint{
Namer: namer.New(name, namer.WithReplace()),
Project: args.Project,
Region: args.Region,
// Default to the latest TensorFlow 2.15 CPU prediction container
ModelImageURL: setDefaultString(args.ModelImageURL, "us-docker.pkg.dev/vertex-ai/prediction/tf2-cpu.2-15:latest"),
ModelDir: args.ModelDir,
ModelPredictionInputSchemaPath: args.ModelPredictionInputSchemaPath,
ModelPredictionOutputSchemaPath: args.ModelPredictionOutputSchemaPath,
ModelPredictionBehaviorSchemaPath: args.ModelPredictionBehaviorSchemaPath,
ModelBucketBasePath: args.ModelBucketBasePath,
ModelDisplayName: setDefaultString(args.ModelDisplayName, name+"-model"),
ModelCommandArgs: setDefaultStringArray(args.ModelCommandArgs, []string{}),
EnvVars: toPulumiStringMap(args.EnvVars).ToStringMapOutput(),
MachineType: setDefaultString(args.MachineType, "n1-highmem-4"),
AcceleratorType: setDefaultString(args.AcceleratorType, "ACCELERATOR_TYPE_UNSPECIFIED"),
AcceleratorCount: setDefaultInt(args.AcceleratorCount, 1),
EndpointDisplayName: setDefaultString(args.EndpointDisplayName, name),
ContainerPort: setDefaultInt(args.ContainerPort, 8080),
HealthRoute: setDefaultString(args.HealthRoute, "/health"),
PredictRoute: setDefaultString(args.PredictRoute, "/predict"),
MinReplicaCount: setDefaultInt(args.MinReplicaCount, 1),
MaxReplicaCount: setDefaultInt(args.MaxReplicaCount, 3),
EnableAccessLogging: setDefaultBool(args.EnableAccessLogging, false),
DisableContainerLogging: setDefaultBool(args.DisableContainerLogging, false),
EnableSpotVMs: setDefaultBool(args.EnableSpotVMs, false),
DeletionProtection: setDefaultBool(args.DeletionProtection, false),
Labels: args.Labels,
name: name,
}
err := ctx.RegisterComponentResource("pulumi-vertex-endpoint:gcp:AIEndpoint", name, aiEndpoint, opts...)
if err != nil {
return nil, fmt.Errorf("failed to register component resource: %w", err)
}
// Deploy the infrastructure
err = aiEndpoint.deploy(ctx, args)
if err != nil {
return nil, fmt.Errorf("failed to deploy AI endpoint: %w", err)
}
outputs := pulumi.Map{
"vertex_ai_endpoint_model_service_account_email": aiEndpoint.modelServiceAccount.Email,
"vertex_ai_endpoint_id": aiEndpoint.endpoint.ID(),
"vertex_ai_endpoint_name": aiEndpoint.endpoint.Name,
"vertex_ai_endpoint_model_image_url": aiEndpoint.modelDeployment.ModelImageUrl,
"vertex_ai_endpoint_model_deployment_id": aiEndpoint.modelDeployment.ID(),
"vertex_ai_endpoint_model_artifacts_bucket_uri": aiEndpoint.modelDeployment.ModelArtifactsBucketUri,
"vertex_ai_endpoint_deployed_model_id": aiEndpoint.modelDeployment.DeployedModelId,
"vertex_ai_endpoint_model_prediction_input_schema_uri": aiEndpoint.modelDeployment.ModelPredictionInputSchemaUri,
"vertex_ai_endpoint_model_prediction_output_schema_uri": aiEndpoint.modelDeployment.ModelPredictionOutputSchemaUri,
"vertex_ai_endpoint_model_prediction_behavior_schema_uri": aiEndpoint.modelDeployment.ModelPredictionBehaviorSchemaUri,
}
if aiEndpoint.artifactsBucket != nil {
outputs["vertex_ai_endpoint_artifacts_bucket_name"] = aiEndpoint.artifactsBucket.Name
outputs["vertex_ai_endpoint_uploaded_model_files"] = aiEndpoint.uploadedModelArtifacts
}
err = ctx.RegisterResourceOutputs(aiEndpoint, outputs)
if err != nil {
return nil, fmt.Errorf("failed to register resource outputs: %w", err)
}
return aiEndpoint, nil
}
// deploy provisions all the resources for the Vertex AI Endpoint.
func (v *AIEndpoint) deploy(ctx *pulumi.Context, args *AIEndpointArgs) error {
// Create service account for the model deployment
modelServiceAccount, err := v.createModelServiceAccount(ctx)
if err != nil {
return fmt.Errorf("failed to create model service account: %w", err)
}
v.modelServiceAccount = modelServiceAccount
// Grant necessary IAM roles to the model service account
iamMembers, err := v.grantModelIAMRoles(ctx, modelServiceAccount.Email)
if err != nil {
return fmt.Errorf("failed to grant model IAM roles: %w", err)
}
v.iamMembers = iamMembers
if args.EnablePrivateRegistryAccess {
registryIAMAccess, err := v.grantRegistryIAMAccess(ctx, modelServiceAccount.Email)
if err != nil {
return fmt.Errorf("failed to grant registry IAM access: %w", err)
}
v.registryIAMAccess = registryIAMAccess
}
// the endpoint named will be attached to the model deployment
endpoint, err := v.createEndpoint(ctx)
if err != nil {
return fmt.Errorf("failed to create endpoint: %w", err)
}
v.endpoint = endpoint
modelArtifactsRequired := args.ModelDir != ""
var modelArtifactsURI pulumi.StringOutput
var uploadedObjects []pulumi.Resource
if modelArtifactsRequired {
// Upload model artifacts (including schemas) to bucket
modelArtifactsURI, uploadedObjects, err = v.uploadModelToBucket(ctx, args.ModelDir, args.ModelBucketBasePath, args.Labels)
if err != nil {
return fmt.Errorf("failed to upload model to bucket: %w", err)
}
// Collect uploaded object names for tracking
uploadedObjectNames := pulumi.StringArray{}
for _, resource := range uploadedObjects {
if bucketObject, ok := resource.(*storage.BucketObject); ok {
uploadedObjectNames = append(uploadedObjectNames, bucketObject.Name.ApplyT(func(name string) string {
return name
}).(pulumi.StringOutput))
}
}
v.uploadedModelArtifacts = uploadedObjectNames.ToStringArrayOutput()
// allow the model service account to access the model bucket
bucketIamMembers, err := v.grantModelBucketIAMAccess(ctx, modelServiceAccount.Email)
if err != nil {
return fmt.Errorf("failed to grant model bucket IAM access: %w", err)
}
v.bucketIAMMembers = bucketIamMembers
}
// Deploy the model using https://github.com/davidmontoyago/pulumi-gcp-vertex-model-deployment
modelDeployment, err := v.deployModel(ctx, endpoint.Name, modelServiceAccount.Email,
modelArtifactsRequired, modelArtifactsURI, uploadedObjects)
if err != nil {
return fmt.Errorf("failed to deploy model /o\\: %w", err)
}
v.modelDeployment = modelDeployment
return nil
}
// Getter methods for accessing internal resources
// GetModelServiceAccount returns the model service account resource.
func (v *AIEndpoint) GetModelServiceAccount() *serviceaccount.Account {
return v.modelServiceAccount
}
// GetModel returns the Vertex AI Model resource.
// Note: In this initial version, models are not yet implemented.
func (v *AIEndpoint) GetModel() interface{} {
return nil
}
// GetModelDeployment returns the Vertex AI Model Deployment resource.
func (v *AIEndpoint) GetModelDeployment() *vertexmodeldeployment.VertexModelDeployment {
return v.modelDeployment
}
// GetEndpoint returns the Vertex AI Endpoint resource.
func (v *AIEndpoint) GetEndpoint() *vertex.AiEndpoint {
return v.endpoint
}
// GetIAMMembers returns the IAM member resources.
func (v *AIEndpoint) GetIAMMembers() []*projects.IAMMember {
return v.iamMembers
}
// GetUploadedModelArtifacts returns the array of uploaded model artifact names.
func (v *AIEndpoint) GetUploadedModelArtifacts() pulumi.StringArrayOutput {
return v.uploadedModelArtifacts
}
// Package config provides an environment config helper
package config
import (
"fmt"
"log"
"github.com/kelseyhightower/envconfig"
"github.com/pulumi/pulumi/sdk/v3/go/pulumi"
"github.com/davidmontoyago/pulumi-gcp-ai-endpoint/pkg/vertex/gcp"
)
// Config allows setting the vertex endpoint configuration via environment variables
type Config struct {
GCPProject string `envconfig:"GCP_PROJECT" required:"true"`
GCPRegion string `envconfig:"GCP_REGION" required:"true"`
ModelImageURL string `envconfig:"MODEL_IMAGE_URL" default:"us-docker.pkg.dev/vertex-ai/prediction/tf2-cpu.2-15:latest"`
ModelDir string `envconfig:"MODEL_DIR" default:""`
ModelPredictionInputSchemaPath string `envconfig:"MODEL_PREDICTION_INPUT_SCHEMA_PATH" default:""`
ModelPredictionOutputSchemaPath string `envconfig:"MODEL_PREDICTION_OUTPUT_SCHEMA_PATH" default:""`
ModelPredictionBehaviorSchemaPath string `envconfig:"MODEL_PREDICTION_BEHAVIOR_SCHEMA_PATH" default:""`
ModelBucketBasePath string `envconfig:"MODEL_BUCKET_BASE_PATH" default:"model"`
ModelDisplayName string `envconfig:"MODEL_DISPLAY_NAME" default:""`
MachineType string `envconfig:"MACHINE_TYPE" default:"n1-standard-2"`
AcceleratorType string `envconfig:"ACCELERATOR_TYPE" default:"ACCELERATOR_TYPE_UNSPECIFIED"`
AcceleratorCount int `envconfig:"ACCELERATOR_COUNT" default:"1"`
DeletionProtection bool `envconfig:"DELETION_PROTECTION" default:"false"`
EndpointDisplayName string `envconfig:"ENDPOINT_DISPLAY_NAME" default:""`
ContainerPort int `envconfig:"CONTAINER_PORT" default:"8080"`
HealthRoute string `envconfig:"HEALTH_ROUTE" default:"/health"`
PredictRoute string `envconfig:"PREDICT_ROUTE" default:"/predict"`
MinReplicaCount int `envconfig:"MIN_REPLICA_COUNT" default:"1"`
MaxReplicaCount int `envconfig:"MAX_REPLICA_COUNT" default:"3"`
EnableAccessLogging bool `envconfig:"ENABLE_ACCESS_LOGGING" default:"false"`
DisableContainerLogging bool `envconfig:"DISABLE_CONTAINER_LOGGING" default:"false"`
EnableSpotVMs bool `envconfig:"ENABLE_SPOT_VMS" default:"false"`
EnablePrivateRegistryAccess bool `envconfig:"ENABLE_PRIVATE_REGISTRY_ACCESS" default:"false"`
}
// LoadConfig loads configuration from environment variables
// All required environment variables must be set or will cause an error
func LoadConfig() (*Config, error) {
var config Config
err := envconfig.Process("", &config)
if err != nil {
return nil, fmt.Errorf("failed to load configuration from environment variables: %w", err)
}
log.Printf("Configuration loaded successfully:")
log.Printf(" GCP Project: %s", config.GCPProject)
log.Printf(" GCP Region: %s", config.GCPRegion)
log.Printf(" Model Dir: %s", config.ModelDir)
log.Printf(" Model Prediction Input Schema Path: %s", config.ModelPredictionInputSchemaPath)
log.Printf(" Model Prediction Output Schema Path: %s", config.ModelPredictionOutputSchemaPath)
log.Printf(" Model Prediction Behavior Schema Path: %s", config.ModelPredictionBehaviorSchemaPath)
log.Printf(" Model Bucket Base Path: %s", config.ModelBucketBasePath)
log.Printf(" Model Image URL: %s", config.ModelImageURL)
log.Printf(" Machine Type: %s", config.MachineType)
log.Printf(" Accelerator Type: %s", config.AcceleratorType)
log.Printf(" Accelerator Count: %d", config.AcceleratorCount)
log.Printf(" Deletion Protection: %t", config.DeletionProtection)
log.Printf(" Endpoint Display Name: %s", config.EndpointDisplayName)
log.Printf(" Model Display Name: %s", config.ModelDisplayName)
log.Printf(" Container Port: %d", config.ContainerPort)
log.Printf(" Health Route: %s", config.HealthRoute)
log.Printf(" Predict Route: %s", config.PredictRoute)
log.Printf(" Min Replica Count: %d", config.MinReplicaCount)
log.Printf(" Max Replica Count: %d", config.MaxReplicaCount)
log.Printf(" Enable Access Logging: %t", config.EnableAccessLogging)
log.Printf(" Disable Container Logging: %t", config.DisableContainerLogging)
log.Printf(" Enable Spot VMs: %t", config.EnableSpotVMs)
log.Printf(" Enable Private Registry Access: %t", config.EnablePrivateRegistryAccess)
return &config, nil
}
// ToAIEndpointArgs converts the config to AIEndpointArgs for use with the Pulumi component
func (c *Config) ToAIEndpointArgs() *gcp.AIEndpointArgs {
args := &gcp.AIEndpointArgs{
Project: c.GCPProject,
Region: c.GCPRegion,
ModelDir: c.ModelDir,
ModelPredictionInputSchemaPath: c.ModelPredictionInputSchemaPath,
ModelPredictionOutputSchemaPath: c.ModelPredictionOutputSchemaPath,
ModelBucketBasePath: c.ModelBucketBasePath,
ModelImageURL: pulumi.String(c.ModelImageURL),
MachineType: pulumi.String(c.MachineType),
AcceleratorType: pulumi.String(c.AcceleratorType),
AcceleratorCount: pulumi.Int(c.AcceleratorCount),
DeletionProtection: pulumi.Bool(c.DeletionProtection),
ContainerPort: pulumi.Int(c.ContainerPort),
HealthRoute: pulumi.String(c.HealthRoute),
PredictRoute: pulumi.String(c.PredictRoute),
MinReplicaCount: pulumi.Int(c.MinReplicaCount),
MaxReplicaCount: pulumi.Int(c.MaxReplicaCount),
EnablePrivateRegistryAccess: c.EnablePrivateRegistryAccess,
EnableAccessLogging: pulumi.Bool(c.EnableAccessLogging),
DisableContainerLogging: pulumi.Bool(c.DisableContainerLogging),
EnableSpotVMs: pulumi.Bool(c.EnableSpotVMs),
}
// Set optional fields only if provided
if c.EndpointDisplayName != "" {
args.EndpointDisplayName = pulumi.String(c.EndpointDisplayName)
}
if c.ModelDisplayName != "" {
args.ModelDisplayName = pulumi.String(c.ModelDisplayName)
}
if c.ModelPredictionBehaviorSchemaPath != "" {
args.ModelPredictionBehaviorSchemaPath = c.ModelPredictionBehaviorSchemaPath
}
return args
}
package gcp
import (
"fmt"
"github.com/pulumi/pulumi-gcp/sdk/v8/go/gcp/projects"
"github.com/pulumi/pulumi-gcp/sdk/v8/go/gcp/vertex"
"github.com/pulumi/pulumi/sdk/v3/go/pulumi"
)
// createEndpoint creates a Vertex AI Endpoint.
func (v *AIEndpoint) createEndpoint(ctx *pulumi.Context) (*vertex.AiEndpoint, error) {
aiService, err := projects.NewService(ctx, v.NewResourceName("aiplatform", "service", 63), &projects.ServiceArgs{
Project: pulumi.String(v.Project),
Service: pulumi.String("aiplatform.googleapis.com"),
},
pulumi.RetainOnDelete(true),
pulumi.Parent(v),
)
if err != nil {
return nil, fmt.Errorf("failed to enable AI platform service: %w", err)
}
endpoint, err := vertex.NewAiEndpoint(ctx, v.NewResourceName("ai", "endpoint", 63), &vertex.AiEndpointArgs{
Project: pulumi.String(v.Project),
Region: pulumi.String(v.Region),
Location: pulumi.String(v.Region),
DisplayName: v.EndpointDisplayName,
Description: pulumi.String("Vertex AI endpoint for real-time predictions"),
Labels: toPulumiStringMap(v.Labels),
},
pulumi.Parent(v),
pulumi.DependsOn([]pulumi.Resource{aiService}),
)
if err != nil {
return nil, fmt.Errorf("failed to create endpoint: %w", err)
}
return endpoint, nil
}
package gcp
import (
"fmt"
"mime"
"os"
"path/filepath"
"strings"
"github.com/pulumi/pulumi-gcp/sdk/v8/go/gcp/storage"
"github.com/pulumi/pulumi/sdk/v3/go/pulumi"
)
// uploadModelToBucket creates a bucket for model artifacts and uploads the model directory.
// It returns the GCS URI of the uploaded model artifacts and the uploaded objects for dependency tracking.
func (v *AIEndpoint) uploadModelToBucket(ctx *pulumi.Context, modelDir string, modelBucketBasePath string, labels map[string]string) (pulumi.StringOutput, []pulumi.Resource, error) {
// Create the bucket for model artifacts
bucketName := v.NewResourceName("model", "bucket", 63)
// Merge default labels with provided labels
bucketLabels := pulumi.StringMap{
"purpose": pulumi.String("model-storage"),
}
// Add user-provided labels
for key, value := range labels {
bucketLabels[key] = pulumi.String(value)
}
artifactsBucket, err := storage.NewBucket(ctx, bucketName, &storage.BucketArgs{
Name: pulumi.String(bucketName),
Location: pulumi.String(v.Region),
Project: pulumi.String(v.Project),
ForceDestroy: pulumi.Bool(true), // Model data is part of the pipeline, safe to implode.
// Enable Uniform Bucket Level Access (UBLA) for enhanced security
// This is required for SBOMs and prevents ACL-based access control
UniformBucketLevelAccess: pulumi.Bool(true),
Versioning: &storage.BucketVersioningArgs{
Enabled: pulumi.Bool(true), // Enable versioning for audit trail
},
Labels: bucketLabels,
}, pulumi.Parent(v))
if err != nil {
return pulumi.StringOutput{}, nil, fmt.Errorf("failed to create artifacts bucket: %w", err)
}
v.artifactsBucket = artifactsBucket
// No luck with https://github.com/pulumi/pulumi-synced-folder /o\
// Upload the model artifacts
uploadedObjects, err := v.uploadDirectoryToModelBucket(ctx, modelDir, modelBucketBasePath)
if err != nil {
return pulumi.StringOutput{}, nil, fmt.Errorf("failed to upload model artifacts: %w", err)
}
modelArtifactsURI := pulumi.Sprintf("gs://%s/%s", artifactsBucket.Name, modelBucketBasePath)
return modelArtifactsURI, uploadedObjects, nil
}
// uploadDirectoryToModelBucket traverses a directory and uploads all files to a GCS bucket.
func (v *AIEndpoint) uploadDirectoryToModelBucket(ctx *pulumi.Context, localDir, baseObjectPath string) ([]pulumi.Resource, error) {
var bucketObjects []*storage.BucketObject
err := filepath.Walk(localDir, func(filePath string, info os.FileInfo, err error) error {
if err != nil {
return fmt.Errorf("error walking path %s: %w", filePath, err)
}
// Skip directories
if info.IsDir() {
return nil
}
// Skip hidden files and system files
if strings.HasPrefix(info.Name(), ".") {
return nil
}
// Calculate relative path from the base directory to preserve directory structure
relPath, err := filepath.Rel(localDir, filePath)
if err != nil {
return fmt.Errorf("error calculating relative path: %w", err)
}
// Convert to GCS object key (this preserves the original filename and path structure)
gcsObjectName := strings.ReplaceAll(relPath, string(filepath.Separator), "/")
// Detect content type
contentType := detectContentType(filePath)
// Create a unique resource name by replacing path separators with hyphens
resourceName := fmt.Sprintf("file-%s", strings.ReplaceAll(gcsObjectName, "/", "-"))
resourceName = strings.ReplaceAll(resourceName, ".", "-")
// Prepend the base object path if provided
if baseObjectPath != "" {
gcsObjectName = filepath.Join(baseObjectPath, gcsObjectName)
gcsObjectName = strings.ReplaceAll(gcsObjectName, string(filepath.Separator), "/")
}
// Create BucketObject resource
bucketObject, err := storage.NewBucketObject(ctx, resourceName, &storage.BucketObjectArgs{
Name: pulumi.String(gcsObjectName),
Bucket: v.artifactsBucket.Name,
Source: pulumi.NewFileAsset(filePath),
ContentType: pulumi.String(contentType),
}, pulumi.Parent(v))
if err != nil {
return fmt.Errorf("error creating bucket object for %s: %w", filePath, err)
}
bucketObjects = append(bucketObjects, bucketObject)
return nil
})
if err != nil {
return nil, fmt.Errorf("error uploading directory %s: %w", localDir, err)
}
uploadedResources := make([]pulumi.Resource, len(bucketObjects))
for i, bucketObject := range bucketObjects {
uploadedResources[i] = bucketObject
}
return uploadedResources, nil
}
// detectContentType determines the MIME type of a file based on its extension
func detectContentType(filePath string) string {
ext := filepath.Ext(filePath)
contentType := mime.TypeByExtension(ext)
if contentType == "" {
// Default to binary if type cannot be determined
contentType = "application/octet-stream"
}
return contentType
}
package gcp
import (
"fmt"
vertexmodeldeployment "github.com/davidmontoyago/pulumi-gcp-vertex-model-deployment/sdk/go/pulumi-gcp-vertex-model-deployment/resources"
"github.com/pulumi/pulumi/sdk/v3/go/pulumi"
)
// deployModel deploys the model to the Vertex AI Endpoint.
func (v *AIEndpoint) deployModel(ctx *pulumi.Context, endpointName pulumi.StringOutput,
serviceAccountEmail pulumi.StringOutput,
modelArtifactsRequired bool,
modelArtifactsURI pulumi.StringOutput,
uploadedObjects []pulumi.Resource) (*vertexmodeldeployment.VertexModelDeployment, error) {
modelDeploymentArgs := &vertexmodeldeployment.VertexModelDeploymentArgs{
ProjectId: pulumi.String(v.Project),
Region: pulumi.String(v.Region),
ModelImageUrl: v.ModelImageURL,
Args: v.ModelCommandArgs,
Env: v.EnvVars,
Port: v.ContainerPort,
EndpointModelDeployment: &vertexmodeldeployment.EndpointModelDeploymentArgsArgs{
EndpointId: endpointName,
MachineType: v.MachineType,
AcceleratorType: v.AcceleratorType,
AcceleratorCount: v.AcceleratorCount,
MinReplicas: v.MinReplicaCount,
MaxReplicas: v.MaxReplicaCount,
EnableAccessLogging: v.EnableAccessLogging,
DisableContainerLogging: v.DisableContainerLogging,
EnableSpotVMs: v.EnableSpotVMs,
},
ServiceAccount: serviceAccountEmail,
}
// Include dependencies on both the bucket permissions and uploaded model artifacts if any
dependencies := []pulumi.Resource{v.modelServiceAccount}
for _, bucketIAMMember := range v.bucketIAMMembers {
dependencies = append(dependencies, bucketIAMMember)
}
if modelArtifactsRequired {
modelDeploymentArgs.ModelArtifactsBucketUri = modelArtifactsURI
if v.ModelPredictionInputSchemaPath != "" {
modelDeploymentArgs.ModelPredictionInputSchemaUri = pulumi.Sprintf("%s/%s", modelArtifactsURI, v.ModelPredictionInputSchemaPath)
}
if v.ModelPredictionOutputSchemaPath != "" {
modelDeploymentArgs.ModelPredictionOutputSchemaUri = pulumi.Sprintf("%s/%s", modelArtifactsURI, v.ModelPredictionOutputSchemaPath)
}
if v.ModelPredictionBehaviorSchemaPath != "" {
modelDeploymentArgs.ModelPredictionBehaviorSchemaUri = pulumi.Sprintf("%s/%s", modelArtifactsURI, v.ModelPredictionBehaviorSchemaPath)
}
dependencies = append(dependencies, v.artifactsBucket)
if len(uploadedObjects) > 0 {
dependencies = append(dependencies, uploadedObjects...)
}
}
if v.registryIAMAccess != nil {
dependencies = append(dependencies, v.registryIAMAccess)
}
modelDeployment, err := vertexmodeldeployment.NewVertexModelDeployment(ctx,
v.NewResourceName("vertex-model-deployment", "", 63),
modelDeploymentArgs,
pulumi.Parent(v),
pulumi.DependsOn(dependencies),
)
if err != nil {
return nil, fmt.Errorf("failed to create model deployment for endpoint: %w", err)
}
return modelDeployment, nil
}
package gcp
import (
"fmt"
"strings"
"github.com/pulumi/pulumi-gcp/sdk/v8/go/gcp/artifactregistry"
"github.com/pulumi/pulumi-gcp/sdk/v8/go/gcp/projects"
"github.com/pulumi/pulumi-gcp/sdk/v8/go/gcp/serviceaccount"
"github.com/pulumi/pulumi-gcp/sdk/v8/go/gcp/storage"
"github.com/pulumi/pulumi/sdk/v3/go/pulumi"
)
// createModelServiceAccount creates a service account for Vertex AI operations.
func (v *AIEndpoint) createModelServiceAccount(ctx *pulumi.Context) (*serviceaccount.Account, error) {
accountID := v.NewResourceName("model-sa", "service-account", 30)
return serviceaccount.NewAccount(ctx, v.NewResourceName("model-sa", "service-account", 63), &serviceaccount.AccountArgs{
Project: pulumi.String(v.Project),
AccountId: pulumi.String(accountID),
DisplayName: pulumi.Sprintf("%s Vertex AI Service Account", v.EndpointDisplayName),
Description: pulumi.String("Service account for deployed model operations"),
}, pulumi.Parent(v))
}
// grantModelIAMRoles grants necessary IAM roles to the model service account.
func (v *AIEndpoint) grantModelIAMRoles(ctx *pulumi.Context, serviceAccountEmail pulumi.StringOutput) ([]*projects.IAMMember, error) {
// IAM roles specific to what the deployed model needs to operate
roles := []string{
"roles/storage.bucketViewer", // List and get buckets
"roles/logging.logWriter", // For writing logs during prediction
"roles/monitoring.metricWriter", // For writing custom metrics
"roles/aiplatform.user", // For accessing Vertex AI resources
}
iamMembers := make([]*projects.IAMMember, len(roles))
for roleIndex, role := range roles {
bindingName := v.NewResourceName(fmt.Sprintf("model-sa-iam-%s", role), "", 63)
member, err := projects.NewIAMMember(ctx, bindingName, &projects.IAMMemberArgs{
Project: pulumi.String(v.Project),
Role: pulumi.String(role),
Member: pulumi.Sprintf("serviceAccount:%s", serviceAccountEmail),
}, pulumi.Parent(v))
if err != nil {
return nil, fmt.Errorf("failed to create IAM member for role %s: %w", role, err)
}
iamMembers[roleIndex] = member
}
return iamMembers, nil
}
// grantRegistryIAMAccess grants the SA access to the registry source of the model server docker image.
func (v *AIEndpoint) grantRegistryIAMAccess(ctx *pulumi.Context, serviceAccountEmail pulumi.StringOutput) (*artifactregistry.RepositoryIamMember, error) {
modelImageRepoName := v.ModelImageURL.ApplyT(func(url string) string {
return strings.Split(url, "/")[2]
}).(pulumi.StringOutput)
bindingName := v.NewResourceName("model-registry-access", "iam-member", 63)
repoMember, err := artifactregistry.NewRepositoryIamMember(ctx, bindingName, &artifactregistry.RepositoryIamMemberArgs{
Repository: modelImageRepoName,
Location: pulumi.String(v.Region),
Project: pulumi.String(v.Project),
Role: pulumi.String("roles/artifactregistry.reader"),
Member: pulumi.Sprintf("serviceAccount:%s", serviceAccountEmail),
}, pulumi.Parent(v))
if err != nil {
return nil, fmt.Errorf("failed to grant registry IAM access: %w", err)
}
return repoMember, nil
}
func (v *AIEndpoint) grantModelBucketIAMAccess(ctx *pulumi.Context, serviceAccountEmail pulumi.StringOutput) ([]*storage.BucketIAMMember, error) {
bindingName := v.NewResourceName("model-bucket-access", "iam-member", 63)
bucketMember, err := storage.NewBucketIAMMember(ctx, bindingName, &storage.BucketIAMMemberArgs{
Bucket: v.artifactsBucket.Name,
Role: pulumi.String("roles/storage.objectViewer"),
Member: pulumi.Sprintf("serviceAccount:%s", serviceAccountEmail),
}, pulumi.Parent(v))
if err != nil {
return nil, fmt.Errorf("failed to grant model bucket IAM access: %w", err)
}
return []*storage.BucketIAMMember{bucketMember}, nil
}
package gcp
import (
"github.com/pulumi/pulumi/sdk/v3/go/pulumi"
)
// Helper functions for setting default values
func setDefaultString(input pulumi.StringInput, defaultValue string) pulumi.StringOutput {
if input == nil {
return pulumi.String(defaultValue).ToStringOutput()
}
return input.ToStringOutput()
}
func setDefaultInt(input pulumi.IntInput, defaultValue int) pulumi.IntOutput {
if input == nil {
return pulumi.Int(defaultValue).ToIntOutput()
}
return input.ToIntOutput()
}
func setDefaultStringArray(input []string, defaultValue []string) pulumi.StringArrayOutput {
if input == nil {
input = defaultValue
}
result := make(pulumi.StringArray, 0, len(input))
for _, v := range input {
result = append(result, pulumi.String(v))
}
return result.ToStringArrayOutput()
}
func setDefaultBool(input pulumi.BoolInput, defaultValue bool) pulumi.BoolOutput {
if input == nil {
return pulumi.Bool(defaultValue).ToBoolOutput()
}
return input.ToBoolOutput()
}
// toPulumiStringMap converts a Go map[string]string to pulumi.StringMap.
func toPulumiStringMap(input map[string]string) pulumi.StringMap {
if input == nil {
return pulumi.StringMap{}
}
result := make(pulumi.StringMap)
for k, v := range input {
result[k] = pulumi.String(v)
}
return result
}