remote-task-excutor-cli/cmd/run.go

212 lines
5.9 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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")
}