From e10c83963e665557ebaf2377c98c3ba373db8d3d Mon Sep 17 00:00:00 2001 From: Mattt Zmuda Date: Thu, 20 Jun 2024 10:48:24 -0700 Subject: [PATCH 1/2] Update to replicate-go v0.20.0 --- go.mod | 2 +- go.sum | 4 ++-- internal/cmd/model/create.go | 3 ++- 3 files changed, 5 insertions(+), 4 deletions(-) diff --git a/go.mod b/go.mod index b7bb526..142efc0 100644 --- a/go.mod +++ b/go.mod @@ -13,7 +13,7 @@ require ( github.com/cli/browser v1.3.0 github.com/getkin/kin-openapi v0.125.0 github.com/mattn/go-isatty v0.0.20 - github.com/replicate/replicate-go v0.18.1 + github.com/replicate/replicate-go v0.20.0 github.com/schollz/progressbar/v3 v3.14.4 github.com/spf13/cobra v1.8.0 github.com/stretchr/testify v1.9.0 diff --git a/go.sum b/go.sum index e14332e..ee2702e 100644 --- a/go.sum +++ b/go.sum @@ -78,8 +78,8 @@ github.com/perimeterx/marshmallow v1.1.5 h1:a2LALqQ1BlHM8PZblsDdidgv1mWi1DgC2UmX github.com/perimeterx/marshmallow v1.1.5/go.mod h1:dsXbUu8CRzfYP5a87xpp0xq9S3u0Vchtcl8we9tYaXw= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/replicate/replicate-go v0.18.1 h1:4zduLVJxdQAoyl7zKj1e2nxwJVMcT6O/sXe6/eUEtns= -github.com/replicate/replicate-go v0.18.1/go.mod h1:D2x8SztjeUKcaYnSgVu3H2DechufLJWZJB4+TLA3Rag= +github.com/replicate/replicate-go v0.20.0 h1:ujksgJCyJMuRdXjtoRe6wA08NmCTS16LP/x7UtvSLRE= +github.com/replicate/replicate-go v0.20.0/go.mod h1:D2x8SztjeUKcaYnSgVu3H2DechufLJWZJB4+TLA3Rag= github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc= github.com/rivo/uniseg v0.4.7 h1:WUdvkW8uEhrYfLC4ZzdpI2ztxP1I582+49Oc5Mq64VQ= github.com/rivo/uniseg v0.4.7/go.mod h1:FN3SvrM+Zdj16jyLfmOkMNblXMcoc8DfTHruCPUcx88= diff --git a/internal/cmd/model/create.go b/internal/cmd/model/create.go index 1e7dde9..416a193 100644 --- a/internal/cmd/model/create.go +++ b/internal/cmd/model/create.go @@ -1,6 +1,7 @@ package model import ( + "encoding/json" "fmt" "github.com/cli/browser" @@ -59,7 +60,7 @@ var createCmd = &cobra.Command{ } if flags.Changed("json") || !util.IsTTY() { - bytes, err := model.MarshalJSON() + bytes, err := json.MarshalIndent(model, "", " ") if err != nil { return fmt.Errorf("failed to serialize model: %w", err) } From 91ae336ac064fd4cfb2124eb026880e57e270b3a Mon Sep 17 00:00:00 2001 From: Mattt Zmuda Date: Thu, 20 Jun 2024 10:48:35 -0700 Subject: [PATCH 2/2] Add deployments subcommand --- cmd/replicate/main.go | 2 + internal/cmd/deployment/create.go | 114 +++++++++++++++++++++++++ internal/cmd/deployment/list.go | 135 ++++++++++++++++++++++++++++++ internal/cmd/deployment/root.go | 39 +++++++++ internal/cmd/deployment/schema.go | 103 +++++++++++++++++++++++ internal/cmd/deployment/show.go | 93 ++++++++++++++++++++ internal/cmd/deployment/update.go | 123 +++++++++++++++++++++++++++ 7 files changed, 609 insertions(+) create mode 100644 internal/cmd/deployment/create.go create mode 100644 internal/cmd/deployment/list.go create mode 100644 internal/cmd/deployment/root.go create mode 100644 internal/cmd/deployment/schema.go create mode 100644 internal/cmd/deployment/show.go create mode 100644 internal/cmd/deployment/update.go diff --git a/cmd/replicate/main.go b/cmd/replicate/main.go index 2d6860e..d2bd215 100644 --- a/cmd/replicate/main.go +++ b/cmd/replicate/main.go @@ -9,6 +9,7 @@ import ( "github.com/replicate/cli/internal/cmd" "github.com/replicate/cli/internal/cmd/account" "github.com/replicate/cli/internal/cmd/auth" + "github.com/replicate/cli/internal/cmd/deployment" "github.com/replicate/cli/internal/cmd/hardware" "github.com/replicate/cli/internal/cmd/model" "github.com/replicate/cli/internal/cmd/prediction" @@ -41,6 +42,7 @@ func init() { model.RootCmd, prediction.RootCmd, training.RootCmd, + deployment.RootCmd, hardware.RootCmd, cmd.ScaffoldCmd, } { diff --git a/internal/cmd/deployment/create.go b/internal/cmd/deployment/create.go new file mode 100644 index 0000000..988ca20 --- /dev/null +++ b/internal/cmd/deployment/create.go @@ -0,0 +1,114 @@ +package deployment + +import ( + "encoding/json" + "fmt" + + "github.com/cli/browser" + "github.com/replicate/replicate-go" + "github.com/spf13/cobra" + + "github.com/replicate/cli/internal/client" + "github.com/replicate/cli/internal/identifier" + "github.com/replicate/cli/internal/util" +) + +// createCmd represents the create command +var createCmd = &cobra.Command{ + Use: "create <[owner/]name> [flags]", + Short: "Create a new deployment", + Example: `replicate deployment create text-to-image --model=stability-ai/sdxl --hardware=gpu-a100-large`, + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + r8, err := client.NewClient() + if err != nil { + return err + } + + opts := &replicate.CreateDeploymentOptions{} + + opts.Name = args[0] + + flags := cmd.Flags() + + modelFlag, _ := flags.GetString("model") + id, err := identifier.ParseIdentifier(modelFlag) + if err != nil { + return fmt.Errorf("expected /[:version] but got %s", args[0]) + } + opts.Model = fmt.Sprintf("%s/%s", id.Owner, id.Name) + if id.Version != "" { + opts.Version = id.Version + } else { + model, err := r8.GetModel(cmd.Context(), id.Owner, id.Name) + if err != nil { + return fmt.Errorf("failed to get model: %w", err) + } + opts.Version = model.LatestVersion.ID + } + + opts.Hardware, _ = flags.GetString("hardware") + + flagMap := map[string]*int{ + "min-instances": &opts.MinInstances, + "max-instances": &opts.MaxInstances, + } + for flagName, optPtr := range flagMap { + if flags.Changed(flagName) { + value, _ := flags.GetInt(flagName) + *optPtr = value + } + } + + deployment, err := r8.CreateDeployment(cmd.Context(), *opts) + if err != nil { + return fmt.Errorf("failed to create deployment: %w", err) + } + + if flags.Changed("json") || !util.IsTTY() { + bytes, err := json.MarshalIndent(deployment, "", " ") + if err != nil { + return fmt.Errorf("failed to serialize model: %w", err) + } + fmt.Println(string(bytes)) + return nil + } + + url := fmt.Sprintf("https://replicate.com/deployments/%s/%s", deployment.Owner, deployment.Name) + if flags.Changed("web") { + if util.IsTTY() { + fmt.Println("Opening in browser...") + } + + err := browser.OpenURL(url) + if err != nil { + return fmt.Errorf("failed to open browser: %w", err) + } + + return nil + } + + fmt.Printf("Deployment created: %s\n", url) + + return nil + }, +} + +func init() { + addCreateFlags(createCmd) +} + +func addCreateFlags(cmd *cobra.Command) { + cmd.Flags().String("model", "", "Model to deploy") + _ = cmd.MarkFlagRequired("model") + + cmd.Flags().String("hardware", "", "SKU of the hardware to run the model") + _ = cmd.MarkFlagRequired("hardware") + + cmd.Flags().Int("min-instances", 0, "Minimum number of instances to run the model") + cmd.Flags().Int("max-instances", 0, "Maximum number of instances to run the model") + + cmd.Flags().Bool("json", false, "Emit JSON") + cmd.Flags().Bool("web", false, "View on web") + cmd.MarkFlagsMutuallyExclusive("json", "web") +} diff --git a/internal/cmd/deployment/list.go b/internal/cmd/deployment/list.go new file mode 100644 index 0000000..9207ab5 --- /dev/null +++ b/internal/cmd/deployment/list.go @@ -0,0 +1,135 @@ +package deployment + +import ( + "encoding/json" + "fmt" + "os/exec" + "strconv" + + "github.com/spf13/cobra" + + "github.com/replicate/cli/internal/client" + "github.com/replicate/cli/internal/util" + + "github.com/charmbracelet/bubbles/table" + tea "github.com/charmbracelet/bubbletea" + "github.com/charmbracelet/lipgloss" +) + +var baseStyle = lipgloss.NewStyle(). + BorderStyle(lipgloss.NormalBorder()). + BorderForeground(lipgloss.Color("240")) + +type model struct { + table table.Model +} + +func (m model) Init() tea.Cmd { return nil } + +func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { + var cmd tea.Cmd + switch msg := msg.(type) { //nolint:gocritic + case tea.KeyMsg: + switch msg.String() { + case "esc": + if m.table.Focused() { + m.table.Blur() + } else { + m.table.Focus() + } + case "q", "ctrl+c": + return m, tea.Quit + case "enter": + selected := m.table.SelectedRow() + if len(selected) == 0 { + return m, nil + } + url := fmt.Sprintf("https://replicate.com/deployments/%s", selected[0]) + return m, tea.ExecProcess(exec.Command("open", url), nil) + } + } + m.table, cmd = m.table.Update(msg) + return m, cmd +} + +func (m model) View() string { + return baseStyle.Render(m.table.View()) + "\n" +} + +var listCmd = &cobra.Command{ + Use: "list", + Short: "List deployments", + Example: "replicate deployment list", + RunE: func(cmd *cobra.Command, _ []string) error { + ctx := cmd.Context() + + r8, err := client.NewClient() + if err != nil { + return err + } + + deployments, err := r8.ListDeployments(ctx) + if err != nil { + return fmt.Errorf("failed to get deployments: %w", err) + } + + if cmd.Flags().Changed("json") || !util.IsTTY() { + bytes, err := json.MarshalIndent(deployments, "", " ") + if err != nil { + return fmt.Errorf("failed to marshal deployments: %w", err) + } + fmt.Println(string(bytes)) + return nil + } + + columns := []table.Column{ + {Title: "Name", Width: 20}, + {Title: "Release #", Width: 10}, + {Title: "Model Version", Width: 60}, + } + + rows := []table.Row{} + + for _, deployment := range deployments.Results { + rows = append(rows, table.Row{ + deployment.Owner + "/" + deployment.Name, + strconv.Itoa(deployment.CurrentRelease.Number), + fmt.Sprintf("%s:%s", deployment.CurrentRelease.Model, deployment.CurrentRelease.Version), + }) + } + + t := table.New( + table.WithColumns(columns), + table.WithRows(rows), + table.WithFocused(true), + table.WithHeight(30), + ) + + s := table.DefaultStyles() + s.Header = s.Header. + BorderStyle(lipgloss.NormalBorder()). + BorderForeground(lipgloss.Color("240")). + BorderBottom(true). + Bold(false) + s.Selected = s.Selected. + Foreground(lipgloss.Color("229")). + Background(lipgloss.Color("57")). + Bold(false) + t.SetStyles(s) + + m := model{t} + if _, err := tea.NewProgram(m).Run(); err != nil { + return err + } + + return nil + }, +} + +func init() { + addListFlags(listCmd) +} + +func addListFlags(cmd *cobra.Command) { + cmd.Flags().Bool("json", false, "Emit JSON") +} diff --git a/internal/cmd/deployment/root.go b/internal/cmd/deployment/root.go new file mode 100644 index 0000000..061883b --- /dev/null +++ b/internal/cmd/deployment/root.go @@ -0,0 +1,39 @@ +package deployment + +import ( + "github.com/spf13/cobra" +) + +var RootCmd = &cobra.Command{ + Use: "deployments [subcommand]", + Short: "Interact with deployments", + Aliases: []string{"deployments", "d"}, +} + +func init() { + RootCmd.AddGroup(&cobra.Group{ + ID: "subcommand", + Title: "Subcommands:", + }) + for _, cmd := range []*cobra.Command{ + listCmd, + showCmd, + schemaCmd, + createCmd, + updateCmd, + } { + RootCmd.AddCommand(cmd) + cmd.GroupID = "subcommand" + } + + // RootCmd.AddGroup(&cobra.Group{ + // ID: "alias", + // Title: "Alias commands:", + // }) + // for _, cmd := range []*cobra.Command{ + // runCmd, + // } { + // RootCmd.AddCommand(cmd) + // cmd.GroupID = "alias" + // } +} diff --git a/internal/cmd/deployment/schema.go b/internal/cmd/deployment/schema.go new file mode 100644 index 0000000..9dc254f --- /dev/null +++ b/internal/cmd/deployment/schema.go @@ -0,0 +1,103 @@ +package deployment + +import ( + "encoding/json" + "fmt" + + "github.com/replicate/cli/internal/client" + "github.com/replicate/cli/internal/identifier" + "github.com/replicate/cli/internal/util" + + "github.com/replicate/replicate-go" + "github.com/spf13/cobra" +) + +var schemaCmd = &cobra.Command{ + Use: "schema <[owner/]name>", + Short: "Show the inputs and outputs of a deployment", + Example: `replicate deployment schema acme/text-to-image`, + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + id, err := identifier.ParseIdentifier(args[0]) + if err != nil { + return fmt.Errorf("invalid model specified: %s", args[0]) + } + + ctx := cmd.Context() + + r8, err := client.NewClient() + if err != nil { + return err + } + + deployment, err := r8.GetDeployment(ctx, id.Owner, id.Name) + if err != nil { + return fmt.Errorf("failed to get deployment: %w", err) + } + + if deployment.CurrentRelease.Version == "" { + return fmt.Errorf("deployment %s has no current release", args[0]) + } + + version, err := r8.GetModelVersion(ctx, id.Owner, id.Name, deployment.CurrentRelease.Version) + if err != nil { + return fmt.Errorf("failed to get model version of current release: %w", err) + } + + if cmd.Flags().Changed("json") || !util.IsTTY() { + bytes, err := json.MarshalIndent(version.OpenAPISchema, "", " ") + if err != nil { + return fmt.Errorf("failed to serialize schema: %w", err) + } + fmt.Println(string(bytes)) + + return nil + } + + return printModelVersionSchema(version) + }, +} + +// TODO: move this to util package +func printModelVersionSchema(version *replicate.ModelVersion) error { + inputSchema, outputSchema, err := util.GetSchemas(*version) + if err != nil { + return fmt.Errorf("failed to get schemas: %w", err) + } + + if inputSchema != nil { + fmt.Println("Inputs:") + + for _, propName := range util.SortedKeys(inputSchema.Properties) { + prop, ok := inputSchema.Properties[propName] + if !ok { + continue + } + + description := prop.Value.Description + if prop.Value.Enum != nil { + for _, enum := range prop.Value.Enum { + description += fmt.Sprintf("\n- %s", enum) + } + } + + fmt.Printf("- %s: %s (type: %s)\n", propName, description, prop.Value.Type) + } + fmt.Println() + } + + if outputSchema != nil { + fmt.Println("Output:") + fmt.Printf("- type: %s\n", outputSchema.Type) + if outputSchema.Type.Is("array") { + fmt.Printf("- items: %s %s\n", outputSchema.Items.Value.Type, outputSchema.Items.Value.Format) + } + fmt.Println() + } + + return nil +} + +func init() { + schemaCmd.Flags().Bool("json", false, "Emit JSON") +} diff --git a/internal/cmd/deployment/show.go b/internal/cmd/deployment/show.go new file mode 100644 index 0000000..4d58f4f --- /dev/null +++ b/internal/cmd/deployment/show.go @@ -0,0 +1,93 @@ +package deployment + +import ( + "encoding/json" + "fmt" + "strings" + + "github.com/cli/browser" + "github.com/spf13/cobra" + + "github.com/replicate/cli/internal/client" + "github.com/replicate/cli/internal/identifier" + "github.com/replicate/cli/internal/util" +) + +var showCmd = &cobra.Command{ + Use: "show <[owner/]name> [flags]", + Short: "Show a deployment", + Example: "replicate deployment show acme/text-to-image", + Args: cobra.ExactArgs(1), + Aliases: []string{"view"}, + RunE: func(cmd *cobra.Command, args []string) error { + ctx := cmd.Context() + + r8, err := client.NewClient() + if err != nil { + return err + } + + name := args[0] + if !strings.Contains(name, "/") { + account, err := r8.GetCurrentAccount(ctx) + if err != nil { + return fmt.Errorf("failed to get current account: %w", err) + } + name = fmt.Sprintf("%s/%s", account.Username, name) + } + id, err := identifier.ParseIdentifier(name) + if err != nil { + return fmt.Errorf("invalid deployment specified: %s", name) + } + + if cmd.Flags().Changed("web") { + if util.IsTTY() { + fmt.Println("Opening in browser...") + } + + url := fmt.Sprintf("https://replicate.com/deployments/%s/%s", id.Owner, id.Name) + err := browser.OpenURL(url) + if err != nil { + return fmt.Errorf("failed to open browser: %w", err) + } + + return nil + } + + deployment, err := r8.GetDeployment(ctx, id.Owner, id.Name) + if err != nil { + return fmt.Errorf("failed to get deployment: %w", err) + } + + if cmd.Flags().Changed("json") || !util.IsTTY() { + bytes, err := json.MarshalIndent(deployment, "", " ") + if err != nil { + return fmt.Errorf("failed to marshal model: %w", err) + } + fmt.Println(string(bytes)) + return nil + } + + if id.Version != "" { + fmt.Println("Ignoring specified version", id.Version) + } + + fmt.Printf("%s/%s\n", deployment.Owner, deployment.Name) + fmt.Println() + fmt.Printf("Release #%d\n", deployment.CurrentRelease.Number) + fmt.Println("Model:", deployment.CurrentRelease.Model) + fmt.Println("Version:", deployment.CurrentRelease.Version) + fmt.Println("Hardware:", deployment.CurrentRelease.Configuration.Hardware) + fmt.Println("Min instances:", deployment.CurrentRelease.Configuration.MinInstances) + fmt.Println("Max instances:", deployment.CurrentRelease.Configuration.MaxInstances) + + return nil + }, +} + +func init() { + showCmd.Flags().Bool("json", false, "Emit JSON") + showCmd.Flags().Bool("web", false, "Open in web browser") + + showCmd.MarkFlagsMutuallyExclusive("json", "web") +} diff --git a/internal/cmd/deployment/update.go b/internal/cmd/deployment/update.go new file mode 100644 index 0000000..26a1c45 --- /dev/null +++ b/internal/cmd/deployment/update.go @@ -0,0 +1,123 @@ +package deployment + +import ( + "encoding/json" + "fmt" + "strings" + + "github.com/cli/browser" + "github.com/replicate/replicate-go" + "github.com/spf13/cobra" + + "github.com/replicate/cli/internal/client" + "github.com/replicate/cli/internal/identifier" + "github.com/replicate/cli/internal/util" +) + +// updateCmd represents the create command +var updateCmd = &cobra.Command{ + Use: "update <[owner/]name> [flags]", + Short: "Update an existing deployment", + Example: `replicate deployment update acme/text-to-image --max-instances=2`, + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + r8, err := client.NewClient() + if err != nil { + return err + } + + name := args[0] + if !strings.Contains(name, "/") { + account, err := r8.GetCurrentAccount(cmd.Context()) + if err != nil { + return fmt.Errorf("failed to get current account: %w", err) + } + name = fmt.Sprintf("%s/%s", account.Username, name) + } + deploymentID, err := identifier.ParseIdentifier(name) + if err != nil { + return fmt.Errorf("invalid deployment specified: %s", name) + } + + opts := &replicate.UpdateDeploymentOptions{} + + flags := cmd.Flags() + + if flags.Changed("version") { + value, _ := flags.GetString("version") + var version string + if strings.Contains(value, ":") { + modelID, err := identifier.ParseIdentifier(value) + if err != nil { + return fmt.Errorf("invalid model version specified: %s", value) + } + version = modelID.Version + } else { + version = value + } + opts.Version = &version + } + + if flags.Changed("hardware") { + value, _ := flags.GetString("hardware") + opts.Hardware = &value + } + + if flags.Changed("min-instances") { + value, _ := flags.GetInt("min-instances") + opts.MinInstances = &value + } + + if flags.Changed("max-instances") { + value, _ := flags.GetInt("max-instances") + opts.MaxInstances = &value + } + + deployment, err := r8.UpdateDeployment(cmd.Context(), deploymentID.Owner, deploymentID.Name, *opts) + if err != nil { + return fmt.Errorf("failed to update deployment: %w", err) + } + + if flags.Changed("json") || !util.IsTTY() { + bytes, err := json.MarshalIndent(deployment, "", " ") + if err != nil { + return fmt.Errorf("failed to serialize model: %w", err) + } + fmt.Println(string(bytes)) + return nil + } + + url := fmt.Sprintf("https://replicate.com/deployments/%s/%s", deployment.Owner, deployment.Name) + if flags.Changed("web") { + if util.IsTTY() { + fmt.Println("Opening in browser...") + } + + err := browser.OpenURL(url) + if err != nil { + return fmt.Errorf("failed to open browser: %w", err) + } + + return nil + } + + fmt.Printf("Deployment updated: %s\n", url) + + return nil + }, +} + +func init() { + addUpdateFlags(updateCmd) +} + +func addUpdateFlags(cmd *cobra.Command) { + cmd.Flags().String("version", "", "Version of the model to deploy") + cmd.Flags().String("hardware", "", "SKU of the hardware to run the model") + cmd.Flags().Int("min-instances", 0, "Minimum number of instances to run the model") + cmd.Flags().Int("max-instances", 0, "Maximum number of instances to run the model") + + cmd.Flags().Bool("json", false, "Emit JSON") + cmd.Flags().Bool("web", false, "View on web") + cmd.MarkFlagsMutuallyExclusive("json", "web") +}