212 lines
5.9 KiB
Go
212 lines
5.9 KiB
Go
package cmd
|
||
|
||
import (
|
||
"context"
|
||
"encoding/json"
|
||
"fmt"
|
||
"github.com/spf13/cobra"
|
||
"log"
|
||
"os"
|
||
"remote-task-excutor-cli/pkg/config"
|
||
"remote-task-excutor-cli/pkg/models"
|
||
"remote-task-excutor-cli/pkg/runners"
|
||
"remote-task-excutor-cli/pkg/service/auth"
|
||
"remote-task-excutor-cli/pkg/service/preparation"
|
||
"remote-task-excutor-cli/pkg/service/task"
|
||
"strings"
|
||
"time"
|
||
)
|
||
|
||
var runCmd = &cobra.Command{
|
||
Use: "run",
|
||
Short: "Run an remote ML task",
|
||
Long: `Execute an remote ML task with specified parameters`,
|
||
RunE: func(cmd *cobra.Command, args []string) error {
|
||
// 创建参数结构体
|
||
fmt.Println("run called, args is", args)
|
||
runConfig, err := parseParams()
|
||
if err != nil {
|
||
log.Fatalf("Failed to parse parameters: %v", err)
|
||
}
|
||
printParams(runConfig)
|
||
// 验证参数
|
||
if err := models.ValidateConfig(runConfig); err != nil {
|
||
log.Fatalf("Validation error: %v", err)
|
||
}
|
||
|
||
ctx, cancel := context.WithCancel(context.Background())
|
||
defer cancel()
|
||
// 1. 配置Nacos连接
|
||
nacosCfg := config.NacosConfig{
|
||
Host: config.Host,
|
||
Port: config.Port,
|
||
NamespaceID: config.Namespace,
|
||
Group: config.Group,
|
||
DataID: config.DataID,
|
||
LogDir: "/var/log/nacos",
|
||
CacheDir: "/tmp/nacos/cache",
|
||
}
|
||
|
||
// 2. 创建Nacos配置读取器(包含同步配置加载)
|
||
reader, err := config.NewConfigReader(ctx, nacosCfg)
|
||
if err != nil {
|
||
log.Fatalf("初始化配置读取器失败: %v", err)
|
||
}
|
||
defer reader.Close()
|
||
|
||
// 3. 直接获取配置(无需等待)
|
||
apiCfg, err := reader.GetAPIConfig()
|
||
if err != nil {
|
||
log.Fatalf("获取API配置失败: %v", err)
|
||
}
|
||
fmt.Println("remote task executor service config:", apiCfg)
|
||
// 这里添加实际的工作流执行代码
|
||
|
||
authService := auth.NewTokenService(apiCfg)
|
||
preparationService := preparation.NewPreparationService(apiCfg.BaseURL, 4*time.Hour)
|
||
taskService := task.NewTaskService(apiCfg.BaseURL, 10*time.Second)
|
||
runner := runners.NewTaskRunner(authService, preparationService, taskService)
|
||
|
||
if err := runner.RunTask(ctx, runConfig); err != nil {
|
||
fmt.Println("Workflow failed with error:", err)
|
||
return err
|
||
}
|
||
|
||
fmt.Println("Workflow started successfully!")
|
||
return nil
|
||
},
|
||
}
|
||
|
||
func printParams(config *models.RunConfig) {
|
||
// 执行工作流
|
||
fmt.Println("Starting remote ML task with configuration:")
|
||
fmt.Printf("Code Config: %s\n", config.CodeConfig)
|
||
fmt.Printf("Resource: %s\n", config.Resource)
|
||
fmt.Printf("Image: %d\n", config.Image)
|
||
fmt.Printf("Command: %s\n", config.Command)
|
||
fmt.Printf("Dataset: %s\n", config.Dataset)
|
||
fmt.Printf("Model Name: %s\n", config.ModelName)
|
||
fmt.Printf("Run Args: %v\n", config.RunArgs)
|
||
}
|
||
|
||
// 定义命令行参数变量
|
||
var (
|
||
codeConfig string
|
||
resource string
|
||
image int
|
||
command string
|
||
runArgs string
|
||
dataset string
|
||
modelName string
|
||
taskOutput string
|
||
resourceType string
|
||
)
|
||
|
||
func parseParams() (*models.RunConfig, error) {
|
||
var c models.CodeConfig
|
||
var r models.ResourceConfig
|
||
var d models.DataResourceConfig
|
||
var m models.DataResourceConfig
|
||
var rs []string
|
||
|
||
if err := json.Unmarshal([]byte(codeConfig), &c); err != nil {
|
||
return nil, fmt.Errorf("failed to unmarshal codeConfig:%s, error:%v", codeConfig, err)
|
||
}
|
||
// 生成临时路径
|
||
tempDir, err := os.MkdirTemp("", "code-*")
|
||
if err != nil {
|
||
return nil, fmt.Errorf("failed to create temp dir:%s, error:%v", tempDir, err)
|
||
}
|
||
c.MountPath = tempDir
|
||
if err := json.Unmarshal([]byte(resource), &r); err != nil {
|
||
return nil, fmt.Errorf("failed to unmarshal resourceConfig:%s, error:%v", resource, err)
|
||
}
|
||
|
||
if err := json.Unmarshal([]byte(dataset), &d); err != nil {
|
||
return nil, fmt.Errorf("failed to unmarshal datasetConfig:%s, error:%v", dataset, err)
|
||
}
|
||
|
||
if modelName != "" {
|
||
if err := json.Unmarshal([]byte(modelName), &m); err != nil {
|
||
return nil, fmt.Errorf("failed to unmarshal commandConfig:%s, error:%v", modelName, err)
|
||
}
|
||
}
|
||
|
||
if err := json.Unmarshal([]byte(runArgs), &rs); err != nil {
|
||
return nil, fmt.Errorf("failed to unmarshal runArgsConfig:%s, error:%v", runArgs, err)
|
||
}
|
||
|
||
params, err := ParseRunArgs(rs)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("failed to parse runArgsConfig:%s, error:%v", runArgs, err)
|
||
}
|
||
|
||
return &models.RunConfig{
|
||
CodeConfig: c,
|
||
Resource: r,
|
||
Image: image,
|
||
Command: command,
|
||
Dataset: d,
|
||
ModelName: m,
|
||
RunArgs: params,
|
||
TaskOutput: taskOutput,
|
||
}, nil
|
||
}
|
||
|
||
// ParseRunArgs 解析运行参数字符串为键值对
|
||
func ParseRunArgs(input []string) (map[string]string, error) {
|
||
result := make(map[string]string)
|
||
|
||
if len(input) == 0 {
|
||
return result, nil
|
||
}
|
||
|
||
for _, i := range input {
|
||
kv := strings.Split(i, "=")
|
||
if len(kv) != 2 {
|
||
return nil, fmt.Errorf("failed to parse runArgs: %s", i)
|
||
}
|
||
key := strings.TrimPrefix(strings.TrimSpace(kv[0]), "--")
|
||
value := strings.TrimSpace(kv[1])
|
||
result[key] = value
|
||
}
|
||
|
||
return result, nil
|
||
}
|
||
|
||
func init() {
|
||
rootCmd.AddCommand(runCmd)
|
||
|
||
// 添加参数标志
|
||
runCmd.Flags().StringVarP(&codeConfig, "code_config", "c", "",
|
||
"代码配置,支持私有仓库和公有仓库,私有仓库填写ssh地址,公有仓库填写https git地址")
|
||
|
||
runCmd.Flags().StringVarP(&resource, "resource", "r", "",
|
||
"资源规格 (required)")
|
||
|
||
runCmd.Flags().IntVarP(&image, "image", "i", 0,
|
||
"运行镜像 (required)")
|
||
runCmd.Flags().StringVarP(&resourceType, "resource_type", "t", "",
|
||
"资源类型")
|
||
runCmd.Flags().StringVarP(&command, "command", "m", "",
|
||
"启动命令 (required)")
|
||
|
||
runCmd.Flags().StringVarP(&runArgs, "run_args", "a", "",
|
||
"运行参数 (格式: key1=value1,key2=value2)")
|
||
|
||
runCmd.Flags().StringVarP(&dataset, "dataset", "d", "",
|
||
"选择数据集 (required)")
|
||
|
||
runCmd.Flags().StringVarP(&modelName, "model_name", "n", "",
|
||
"选择模型")
|
||
runCmd.Flags().StringVarP(&taskOutput, "task_output", "o", "",
|
||
"任务输出目录")
|
||
|
||
// 设置必需参数
|
||
runCmd.MarkFlagRequired("resource")
|
||
runCmd.MarkFlagRequired("image")
|
||
runCmd.MarkFlagRequired("command")
|
||
runCmd.MarkFlagRequired("dataset")
|
||
runCmd.MarkFlagRequired("task_output")
|
||
}
|