diff --git a/pkg/apis/llm/comfyui_const.go b/pkg/apis/llm/comfyui_const.go new file mode 100644 index 0000000000..7c16ee6079 --- /dev/null +++ b/pkg/apis/llm/comfyui_const.go @@ -0,0 +1,24 @@ +package llm + +const ( + LLM_COMFYUI_BASE_PATH = "/root/ComfyUI" + LLM_COMFYUI_MODELS_PATH = LLM_COMFYUI_BASE_PATH + "/models" + LLM_COMFYUI_HF_ENDPOINT = LLM_VLLM_HF_ENDPOINT + LLM_COMFYUI_CHECKPOINTS_DIR = "checkpoints" + LLM_COMFYUI_LORAS_DIR = "loras" + LLM_COMFYUI_VAE_DIR = "vae" + LLM_COMFYUI_CONTROLNET_DIR = "controlnet" + LLM_COMFYUI_CLIP_DIR = "clip" + LLM_COMFYUI_TEXT_ENCODERS_DIR = "text_encoders" + LLM_COMFYUI_CLIP_VISION_DIR = "clip_vision" + LLM_COMFYUI_DIFFUSION_MODELS_DIR = "diffusion_models" + LLM_COMFYUI_EMBEDDINGS_DIR = "embeddings" + LLM_COMFYUI_IPADAPTER_DIR = "ipadapter" + LLM_COMFYUI_STYLE_MODELS_DIR = "style_models" + LLM_COMFYUI_UNET_DIR = "unet" + LLM_COMFYUI_UPSCALE_MODELS_DIR = "upscale_models" + LLM_COMFYUI_HF_MODELS_DIR = LLM_COMFYUI_CHECKPOINTS_DIR + LLM_COMFYUI_HF_MODELS_PATH = LLM_COMFYUI_MODELS_PATH + "/" + LLM_COMFYUI_HF_MODELS_DIR + LLM_COMFYUI_MODELS_VOLUME_SUBDIR = "storage-models/models" + LLM_COMFYUI_STORAGE_VOLUME_SUBDIR = "storage" +) diff --git a/pkg/apis/llm/instantmodel.go b/pkg/apis/llm/instantmodel.go index 1d86e6d11d..fffa2cdd75 100644 --- a/pkg/apis/llm/instantmodel.go +++ b/pkg/apis/llm/instantmodel.go @@ -4,6 +4,10 @@ import ( "yunion.io/x/onecloud/pkg/apis" ) +const ( + InstantModelSourceHuggingFace = "huggingface" +) + type InstantModelListInput struct { apis.SharableVirtualResourceListInput apis.EnabledResourceBaseListInput @@ -22,6 +26,9 @@ type InstantModelImportInput struct { ModelName string `json:"model_name"` ModelTag string `json:"model_tag"` LlmType LLMContainerType `json:"llm_type"` + Source string `json:"source,omitempty"` + RepoId string `json:"repo_id,omitempty"` + Revision string `json:"revision,omitempty"` } type InstantModelCreateInput struct { @@ -31,6 +38,9 @@ type InstantModelCreateInput struct { LlmType LLMContainerType `json:"llm_type"` ModelName string `json:"model_name"` ModelTag string `json:"model_tag"` + Source string `json:"source,omitempty"` + RepoId string `json:"repo_id,omitempty"` + Revision string `json:"revision,omitempty"` ImageId string `json:"image_id"` Size int64 `json:"size"` ModelId string `json:"model_id"` diff --git a/pkg/apis/llm/instantmodel_huggingface.go b/pkg/apis/llm/instantmodel_huggingface.go index 75d615adf6..fb2b9e940c 100644 --- a/pkg/apis/llm/instantmodel_huggingface.go +++ b/pkg/apis/llm/instantmodel_huggingface.go @@ -1,12 +1,18 @@ package llm type InstantModelHuggingFaceSearchInput struct { - Q string `json:"q"` - Author string `json:"author,omitempty"` - Filter []string `json:"filter,omitempty"` - Direction int `json:"direction,omitempty"` - Limit int `json:"limit,omitempty"` - Sort string `json:"sort,omitempty"` + Q string `json:"q"` + Author string `json:"author,omitempty"` + Filter []string `json:"filter,omitempty"` + Limit int `json:"limit,omitempty"` + Sort string `json:"sort,omitempty"` + Cursor string `json:"cursor,omitempty"` +} + +type InstantModelHuggingFaceSearchOutput struct { + Data []InstantModelHuggingFaceSearchResult `json:"data"` + NextCursor string `json:"next_cursor,omitempty"` + HasMore bool `json:"has_more"` } type InstantModelHuggingFaceRepoInfoInput struct { diff --git a/pkg/apis/llm/llm_container.go b/pkg/apis/llm/llm_container.go index 3931be8b5f..a70e4724ae 100644 --- a/pkg/apis/llm/llm_container.go +++ b/pkg/apis/llm/llm_container.go @@ -26,12 +26,30 @@ var ( string(LLM_CONTAINER_OPENCLAW), string(LLM_CONTAINER_HERMES_AGENT), ) + LLM_INSTANT_MODEL_TYPES = sets.NewString( + string(LLM_CONTAINER_OLLAMA), + string(LLM_CONTAINER_VLLM), + string(LLM_CONTAINER_COMFYUI), + string(LLM_CONTAINER_OPENCLAW), + ) ) func IsLLMContainerType(t string) bool { return LLM_CONTAINER_TYPES.Has(t) } +func IsLLMInstantModelType(t string) bool { + return LLM_INSTANT_MODEL_TYPES.Has(t) +} + +func GetLLMInstantModelContainerType(t LLMContainerType) LLMContainerType { + return t +} + +func IsLLMInstantModelCompatible(instantModelType LLMContainerType, containerType LLMContainerType) bool { + return GetLLMInstantModelContainerType(instantModelType) == containerType +} + type LLMContainerCreateInput struct { apis.VirtualResourceCreateInput LLMId string `json:"llm_id"` diff --git a/pkg/apis/llm/llm_instant_model.go b/pkg/apis/llm/llm_instant_model.go index 08c8759edc..2b1afd5640 100644 --- a/pkg/apis/llm/llm_instant_model.go +++ b/pkg/apis/llm/llm_instant_model.go @@ -14,8 +14,9 @@ type LLMInternalInstantMdlInfo struct { type LLMSaveInstantModelInput struct { apis.ProjectizedResourceCreateInput - ModelId string `json:"model_id"` - ModelFullName string `json:"model_full_name"` + ModelId string `json:"model_id"` + ModelFullName string `json:"model_full_name"` + Mounts []string `json:"mounts"` InstantModelId string `json:"instant_model_id"` diff --git a/pkg/llm/drivers/llm_container/base_driver.go b/pkg/llm/drivers/llm_container/base_driver.go index 3c2e1fa68c..220b440dc4 100644 --- a/pkg/llm/drivers/llm_container/base_driver.go +++ b/pkg/llm/drivers/llm_container/base_driver.go @@ -58,7 +58,7 @@ func (b *baseDriver) ValidateLLMSkuCreateData(ctx context.Context, userCred mccl return nil, errors.Wrapf(err, "validate mounted model %s", mdl) } instantModle := instMdl.(*models.SInstantModel) - if instantModle.LlmType != input.LLMType { + if !api.IsLLMInstantModelCompatible(api.LLMContainerType(instantModle.LlmType), api.LLMContainerType(input.LLMType)) { return nil, errors.Wrapf(httperrors.ErrInvalidStatus, "mounted model %s is not of type %s", mdl, input.LLMType) } input.MountedModels[i] = instantModle.GetId() @@ -90,7 +90,7 @@ func (b *baseDriver) ValidateLLMSkuUpdateData(ctx context.Context, userCred mccl return nil, errors.Wrapf(err, "validate mounted model %s", mdl) } instantModle := instMdl.(*models.SInstantModel) - if instantModle.LlmType != sku.LLMType { + if !api.IsLLMInstantModelCompatible(api.LLMContainerType(instantModle.LlmType), api.LLMContainerType(sku.LLMType)) { return nil, errors.Wrapf(httperrors.ErrInvalidStatus, "mounted model %s is not of type %s", mdl, sku.LLMType) } mountedModels[i] = instantModle.GetId() diff --git a/pkg/llm/drivers/llm_container/comfyui.go b/pkg/llm/drivers/llm_container/comfyui.go index ef8d5e4ecc..a859a7c1e0 100644 --- a/pkg/llm/drivers/llm_container/comfyui.go +++ b/pkg/llm/drivers/llm_container/comfyui.go @@ -2,6 +2,19 @@ package llm_container import ( "context" + "encoding/base64" + "encoding/json" + "fmt" + "net/url" + "os" + "path" + "path/filepath" + "sort" + "strconv" + "strings" + + "yunion.io/x/log" + "yunion.io/x/pkg/errors" commonapi "yunion.io/x/onecloud/pkg/apis" computeapi "yunion.io/x/onecloud/pkg/apis/compute" @@ -37,6 +50,15 @@ func (c *comfyui) GetEffectiveSpec(llm *models.SLLM, sku *models.SLLMSku) interf } func (c *comfyui) GetContainerSpec(ctx context.Context, llm *models.SLLM, image *models.SLLMImage, sku *models.SLLMSku, props []string, devices []computeapi.SIsolatedDevice, diskId string) *computeapi.PodContainerCreateInput { + var postOverlays []*commonapi.ContainerVolumeMountDiskPostOverlay + if llm != nil { + var err error + postOverlays, err = llm.GetMountedModelsPostOverlay() + if err != nil { + log.Errorf("GetMountedModelsPostOverlay failed %s", err) + } + } + spec := computeapi.ContainerSpec{ ContainerSpec: commonapi.ContainerSpec{ Image: image.ToContainerImage(), @@ -79,7 +101,7 @@ func (c *comfyui) GetContainerSpec(ctx context.Context, llm *models.SLLM, image { Disk: &commonapi.ContainerVolumeMountDisk{ Index: &diskIndex, - SubDirectory: "storage", + SubDirectory: api.LLM_COMFYUI_STORAGE_VOLUME_SUBDIR, }, Type: commonapi.CONTAINER_VOLUME_MOUNT_TYPE_DISK, MountPath: "/root", @@ -87,10 +109,12 @@ func (c *comfyui) GetContainerSpec(ctx context.Context, llm *models.SLLM, image { Disk: &commonapi.ContainerVolumeMountDisk{ Index: &diskIndex, - SubDirectory: "storage-models/models", + SubDirectory: api.LLM_COMFYUI_MODELS_VOLUME_SUBDIR, + PostOverlay: postOverlays, }, - Type: commonapi.CONTAINER_VOLUME_MOUNT_TYPE_DISK, - MountPath: "/root/ComfyUI/models", + Type: commonapi.CONTAINER_VOLUME_MOUNT_TYPE_DISK, + MountPath: api.LLM_COMFYUI_MODELS_PATH, + Propagation: commonapi.MOUNTPROPAGATION_PROPAGATION_HOST_TO_CONTAINER, }, { Disk: &commonapi.ContainerVolumeMountDisk{ @@ -114,7 +138,7 @@ func (c *comfyui) GetContainerSpec(ctx context.Context, llm *models.SLLM, image SubDirectory: "storage-user/input", }, Type: commonapi.CONTAINER_VOLUME_MOUNT_TYPE_DISK, - MountPath: "/root/ComfyUI/input", + MountPath: path.Join(api.LLM_COMFYUI_BASE_PATH, "input"), }, { Disk: &commonapi.ContainerVolumeMountDisk{ @@ -122,7 +146,7 @@ func (c *comfyui) GetContainerSpec(ctx context.Context, llm *models.SLLM, image SubDirectory: "storage-user/output", }, Type: commonapi.CONTAINER_VOLUME_MOUNT_TYPE_DISK, - MountPath: "/root/ComfyUI/output", + MountPath: path.Join(api.LLM_COMFYUI_BASE_PATH, "output"), }, { Disk: &commonapi.ContainerVolumeMountDisk{ @@ -130,7 +154,7 @@ func (c *comfyui) GetContainerSpec(ctx context.Context, llm *models.SLLM, image SubDirectory: "storage-user/workflows", }, Type: commonapi.CONTAINER_VOLUME_MOUNT_TYPE_DISK, - MountPath: "/root/ComfyUI/user/default/workflows", + MountPath: path.Join(api.LLM_COMFYUI_BASE_PATH, "user/default/workflows"), }, } spec.VolumeMounts = append(spec.VolumeMounts, ctrVols...) @@ -151,23 +175,138 @@ func (c *comfyui) GetLLMAccessUrlInfo(ctx context.Context, userCred mcclient.Tok } func (c *comfyui) GetProbedInstantModelsExt(ctx context.Context, userCred mcclient.TokenCredential, llm *models.SLLM, mdlIds ...string) (map[string]api.LLMInternalInstantMdlInfo, error) { - return nil, nil + lc, err := llm.GetLLMContainer() + if err != nil { + return nil, errors.Wrap(err, "get llm container") + } + + cmd := buildComfyUIProbeCommand() + output, err := exec(ctx, lc.CmpId, cmd, 10) + if err != nil { + return make(map[string]api.LLMInternalInstantMdlInfo), nil + } + + modelsMap := make(map[string]api.LLMInternalInstantMdlInfo) + lines := strings.Split(strings.TrimSpace(output), "\n") + for _, line := range lines { + fields := strings.Fields(line) + if len(fields) < 2 { + continue + } + sizeKB, _ := strconv.ParseInt(fields[0], 10, 64) + fullPath := path.Clean(fields[1]) + key, info, ok := buildComfyUIProbedModelInfo(fullPath, sizeKB*1024) + if !ok { + continue + } + modelType := getComfyUILLMTypeByModelPath(fullPath) + instMdl, _ := models.GetInstantModelManager().FindInstantModelByMountAndLLMType(fullPath, modelType, true) + if instMdl == nil { + instMdl, _ = models.GetInstantModelManager().FindInstantModelByLLMType(info.ModelId, info.Tag, modelType, true) + } + if instMdl == nil && info.Tag == resolveHfdRevision("") { + instMdl, _ = models.GetInstantModelManager().FindInstantModelByLLMType(info.ModelId, "", modelType, true) + } + if instMdl != nil { + key = instMdl.Id + info.ModelId = instMdl.ModelId + info.Name = instMdl.ModelName + info.Tag = instMdl.ModelTag + if len(mdlIds) > 0 && !containsString(mdlIds, key) { + continue + } + } else if len(mdlIds) > 0 { + continue + } + modelsMap[key] = info + } + return modelsMap, nil } func (c *comfyui) DetectModelPaths(ctx context.Context, userCred mcclient.TokenCredential, llm *models.SLLM, pkgInfo api.LLMInternalInstantMdlInfo) ([]string, error) { - return nil, nil + lc, err := llm.GetLLMContainer() + if err != nil { + return nil, errors.Wrap(err, "get llm container") + } + + instMdl, _ := models.GetInstantModelManager().FindInstantModelByLLMType(pkgInfo.ModelId, pkgInfo.Tag, api.LLM_CONTAINER_COMFYUI, true) + if instMdl == nil && pkgInfo.Tag == resolveHfdRevision("") { + instMdl, _ = models.GetInstantModelManager().FindInstantModelByLLMType(pkgInfo.ModelId, "", api.LLM_CONTAINER_COMFYUI, true) + } + if instMdl != nil && len(instMdl.Mounts) > 0 { + paths := make([]string, 0, len(instMdl.Mounts)) + for _, modelPath := range instMdl.Mounts { + checkCmd := fmt.Sprintf("[ -e %s ] && echo 'EXIST' || echo 'MISSING'", shellQuoteSingle(modelPath)) + output, err := exec(ctx, lc.CmpId, checkCmd, 10) + if err != nil { + return nil, errors.Wrap(err, "failed to check file existence") + } + if strings.Contains(output, "EXIST") { + paths = append(paths, modelPath) + } + } + if len(paths) > 0 { + return paths, nil + } + } + + candidates := []string{} + for _, modelTypeDir := range getComfyUIModelTypeDirs() { + candidates = append(candidates, path.Join(api.LLM_COMFYUI_MODELS_PATH, modelTypeDir, buildComfyUIModelDirName(pkgInfo.ModelId, pkgInfo.Tag))) + if pkgInfo.Name != "" { + candidates = append(candidates, path.Join(api.LLM_COMFYUI_MODELS_PATH, modelTypeDir, filepath.Base(pkgInfo.Name))) + } + } + for _, modelPath := range candidates { + checkCmd := fmt.Sprintf("[ -e %s ] && echo 'EXIST' || echo 'MISSING'", shellQuoteSingle(modelPath)) + output, err := exec(ctx, lc.CmpId, checkCmd, 10) + if err != nil { + return nil, errors.Wrap(err, "failed to check file existence") + } + if strings.Contains(output, "EXIST") { + return []string{modelPath}, nil + } + } + return nil, errors.Errorf("model directory missing for %s:%s", pkgInfo.ModelId, pkgInfo.Tag) } func (c *comfyui) GetImageInternalPathMounts(sApp *models.SInstantModel) map[string]string { - return nil + res := make(map[string]string) + for _, mount := range sApp.Mounts { + relPath, ok := getComfyUIModelMountRelPath(mount) + if !ok { + continue + } + res[relPath] = path.Join(api.LLM_COMFYUI_MODELS_VOLUME_SUBDIR, relPath) + } + return res } func (c *comfyui) GetSaveDirectories(sApp *models.SInstantModel) (string, []string, error) { - return "", nil, nil + var filteredMounts []string + for _, mount := range sApp.Mounts { + relPath, ok := getComfyUIModelMountRelPath(mount) + if !ok { + continue + } + filteredMounts = append(filteredMounts, relPath) + } + return "", filteredMounts, nil } func (c *comfyui) ValidateMounts(mounts []string, mdlName string, mdlTag string) ([]string, error) { - return nil, nil + out := make([]string, 0, len(mounts)) + for _, mount := range mounts { + cleanMount := path.Clean(strings.TrimSpace(mount)) + if cleanMount == "." || cleanMount == "/" { + continue + } + if _, ok := getComfyUIModelMountRelPath(cleanMount); !ok { + return nil, errors.Errorf("invalid comfyui model mount %q: must be under %s", mount, api.LLM_COMFYUI_MODELS_PATH) + } + out = append(out, cleanMount) + } + return out, nil } func (c *comfyui) CheckDuplicateMounts(errStr string, dupIndex int) string { @@ -175,15 +314,55 @@ func (c *comfyui) CheckDuplicateMounts(errStr string, dupIndex int) string { } func (c *comfyui) GetInstantModelIdByPostOverlay(postOverlay *commonapi.ContainerVolumeMountDiskPostOverlay, mdlNameToId map[string]string) string { + if postOverlay == nil { + return "" + } + findByPath := func(p string) string { + modelName, modelTag, ok := parseComfyUIModelDirName(p) + if !ok { + return "" + } + return mdlNameToId[modelName+":"+modelTag] + } + if postOverlay.Image != nil { + for k, v := range postOverlay.Image.PathMap { + if mdlId := findByPath(k); mdlId != "" { + return mdlId + } + if mdlId := findByPath(v); mdlId != "" { + return mdlId + } + } + } + for _, hostLowerDir := range postOverlay.HostLowerDir { + if mdlId := findByPath(hostLowerDir); mdlId != "" { + return mdlId + } + } return "" } func (c *comfyui) GetDirPostOverlay(dir api.LLMMountDirInfo) *commonapi.ContainerVolumeMountDiskPostOverlay { - return nil + uid := int64(1000) + gid := int64(1000) + ov := dir.ToOverlay() + ov.FsUser = &uid + ov.FsGroup = &gid + return &ov } func (c *comfyui) PreInstallModel(ctx context.Context, userCred mcclient.TokenCredential, llm *models.SLLM, instMdl *models.SLLMInstantModel) error { - return nil + lc, err := llm.GetLLMContainer() + if err != nil { + return errors.Wrap(err, "get llm container") + } + dirs := make([]string, 0, len(getComfyUIModelTypeDirs())) + for _, modelTypeDir := range getComfyUIModelTypeDirs() { + dirs = append(dirs, shellQuoteSingle(path.Join(api.LLM_COMFYUI_MODELS_PATH, modelTypeDir))) + } + cmd := fmt.Sprintf("mkdir -p %s", strings.Join(dirs, " ")) + _, err = exec(ctx, lc.CmpId, cmd, 10) + return err } func (c *comfyui) InstallModel(ctx context.Context, userCred mcclient.TokenCredential, llm *models.SLLM, dirs []string, mdlIds []string) error { @@ -195,5 +374,501 @@ func (c *comfyui) UninstallModel(ctx context.Context, userCred mcclient.TokenCre } func (c *comfyui) DownloadModel(ctx context.Context, userCred mcclient.TokenCredential, llm *models.SLLM, tmpDir string, modelName string, modelTag string) (string, []string, error) { - return "", nil, nil + if strings.TrimSpace(tmpDir) == "" { + return "", nil, errors.Error("tmpDir is empty") + } + if strings.TrimSpace(modelName) == "" { + return "", nil, errors.Error("modelName is empty") + } + + rev := resolveHfdRevision(modelTag) + apiURL := fmt.Sprintf("%s/api/models/%s?revision=%s", api.LLM_COMFYUI_HF_ENDPOINT, escapeURLPathPreserveSlash(modelName), url.QueryEscape(rev)) + log.Infof("Downloading HF model for ComfyUI via HF Mirror API: %s", func() string { + b, _ := json.Marshal(map[string]string{ + "model": modelName, + "revision": rev, + "dir": tmpDir, + "endpoint": api.LLM_COMFYUI_HF_ENDPOINT, + "api": apiURL, + }) + return string(b) + }()) + metaBody, err := llm.HttpGet(ctx, apiURL) + if err != nil { + return "", nil, errors.Wrapf(err, "fetch hf model metadata failed: %s", apiURL) + } + meta := hfModelAPIResponse{} + if err := json.Unmarshal(metaBody, &meta); err != nil { + return "", nil, errors.Wrap(err, "unmarshal hf model metadata") + } + if len(meta.Siblings) == 0 { + return "", nil, errors.Errorf("hf model metadata has no siblings: %s", apiURL) + } + + filenames := make([]string, 0, len(meta.Siblings)) + for _, s := range meta.Siblings { + filenames = append(filenames, s.RFilename) + } + targets, mounts, err := buildComfyUIDownloadTargets(tmpDir, modelName, filenames) + if err != nil { + return "", nil, err + } + + for _, target := range targets { + if isNonEmptyFile(target.LocalPath) { + continue + } + if err := os.MkdirAll(filepath.Dir(target.LocalPath), 0755); err != nil { + return "", nil, errors.Wrapf(err, "mkdir for %s", target.LocalPath) + } + fileURL := fmt.Sprintf("%s/%s/resolve/%s/%s", api.LLM_COMFYUI_HF_ENDPOINT, escapeURLPathPreserveSlash(modelName), url.PathEscape(rev), escapeURLPathPreserveSlash(target.Source)) + if err := llm.HttpDownloadFile(ctx, fileURL, target.LocalPath); err != nil { + return "", nil, errors.Wrapf(err, "download file failed: %s -> %s", fileURL, target.LocalPath) + } + } + + return modelName, mounts, nil +} + +func getComfyUIModelMountRelPath(mount string) (string, bool) { + cleanMount := path.Clean(strings.TrimSpace(mount)) + basePath := path.Clean(api.LLM_COMFYUI_MODELS_PATH) + if cleanMount == basePath || !strings.HasPrefix(cleanMount, basePath+"/") { + return "", false + } + relPath := strings.TrimPrefix(cleanMount, basePath+"/") + if relPath == "" || relPath == "." { + return "", false + } + return relPath, true +} + +func buildComfyUIModelDirName(modelName string, modelTag string) string { + modelDir, err := normalizeComfyUIRepoIDPath(modelName) + if err != nil { + return strings.TrimSpace(modelName) + } + return modelDir +} + +func buildComfyUIModelPath(modelName string, modelTag string, llmType api.LLMContainerType) (string, error) { + modelTypeDir, err := getComfyUIModelDirByLLMType(llmType) + if err != nil { + return "", err + } + modelDir, err := normalizeComfyUIRepoIDPath(modelName) + if err != nil { + return "", err + } + return path.Join(api.LLM_COMFYUI_MODELS_PATH, modelTypeDir, modelDir), nil +} + +func parseComfyUIModelDirName(dirName string) (string, string, bool) { + baseName := path.Base(strings.TrimSpace(dirName)) + if modelName, modelTag, ok := parseComfyUILegacyHFModelDirName(baseName); ok { + return modelName, modelTag, true + } + if modelName, modelTag, ok := parseVLLMModelDirName(baseName); ok { + return modelName, modelTag, true + } + repoID, ok := extractComfyUIRepoIDPath(dirName) + if !ok { + return "", "", false + } + return repoID, resolveHfdRevision(""), true +} + +func buildComfyUIProbedModelInfo(modelPath string, sizeBytes int64) (string, api.LLMInternalInstantMdlInfo, bool) { + modelName, modelTag, ok := parseComfyUIModelDirName(modelPath) + if !ok { + return "", api.LLMInternalInstantMdlInfo{}, false + } + return modelName + ":" + modelTag, api.LLMInternalInstantMdlInfo{ + Name: modelName, + Tag: modelTag, + ModelId: modelName, + Size: sizeBytes, + }, true +} + +type comfyUIDownloadTarget struct { + Source string + LocalPath string + Mount string +} + +func buildComfyUIDownloadTargets(tmpDir string, repoID string, filenames []string) ([]comfyUIDownloadTarget, []string, error) { + splitRepo := hasComfyUISplitModelFiles(filenames) + candidates := make([]comfyUIDownloadCandidate, 0, len(filenames)) + hasTopLevelByDir := make(map[string]bool) + for _, filename := range filenames { + src, ok := cleanComfyUIHFFilename(filename) + if !ok { + continue + } + modelDir, ok := classifyComfyUIModelFile(src, splitRepo) + if !ok { + continue + } + candidate := comfyUIDownloadCandidate{ + Source: src, + ModelDir: modelDir, + Target: path.Base(src), + TopLevel: path.Dir(src) == ".", + IsGeneric: isComfyUIGenericDiffusersFile(src), + } + candidates = append(candidates, candidate) + if candidate.TopLevel { + hasTopLevelByDir[modelDir] = true + } + } + + targets := make([]comfyUIDownloadTarget, 0, len(candidates)) + mountSet := make(map[string]struct{}) + usedTargets := make(map[string]struct{}) + for _, candidate := range candidates { + if candidate.IsGeneric && !candidate.TopLevel && hasTopLevelByDir[candidate.ModelDir] { + continue + } + targetName := disambiguateComfyUIDownloadTarget(candidate, usedTargets) + relPath := path.Join(candidate.ModelDir, targetName) + mount := path.Join(api.LLM_COMFYUI_MODELS_PATH, relPath) + targets = append(targets, comfyUIDownloadTarget{ + Source: candidate.Source, + LocalPath: filepath.Join(tmpDir, filepath.FromSlash(relPath)), + Mount: mount, + }) + mountSet[mount] = struct{}{} + } + if len(targets) == 0 { + return nil, nil, errors.Errorf("no supported comfyui model files found in %s", repoID) + } + mounts := make([]string, 0, len(mountSet)) + for mount := range mountSet { + mounts = append(mounts, mount) + } + sort.Strings(mounts) + return targets, mounts, nil +} + +type comfyUIDownloadCandidate struct { + Source string + ModelDir string + Target string + TopLevel bool + IsGeneric bool +} + +func isComfyUIGenericDiffusersFile(filename string) bool { + base := strings.ToLower(path.Base(filename)) + switch base { + case "diffusion_pytorch_model.safetensors", "diffusion_pytorch_model.bin", "model.safetensors", "pytorch_model.bin": + return true + default: + return strings.HasPrefix(base, "model-") && (strings.HasSuffix(base, ".safetensors") || strings.HasSuffix(base, ".bin")) + } +} + +func disambiguateComfyUIDownloadTarget(candidate comfyUIDownloadCandidate, used map[string]struct{}) string { + target := candidate.Target + key := path.Join(candidate.ModelDir, target) + if _, ok := used[key]; !ok { + used[key] = struct{}{} + return target + } + + parent := path.Base(path.Dir(candidate.Source)) + if parent != "." && parent != "/" && parent != "" { + target = sanitizeComfyUIDownloadTargetPart(parent) + "-" + candidate.Target + key = path.Join(candidate.ModelDir, target) + if _, ok := used[key]; !ok { + used[key] = struct{}{} + return target + } + } + + ext := path.Ext(candidate.Target) + name := strings.TrimSuffix(candidate.Target, ext) + for i := 2; ; i++ { + target = fmt.Sprintf("%s-%d%s", name, i, ext) + key = path.Join(candidate.ModelDir, target) + if _, ok := used[key]; !ok { + used[key] = struct{}{} + return target + } + } +} + +func sanitizeComfyUIDownloadTargetPart(part string) string { + part = strings.TrimSpace(part) + var b strings.Builder + for _, r := range part { + switch { + case r >= 'a' && r <= 'z': + b.WriteRune(r) + case r >= 'A' && r <= 'Z': + b.WriteRune(r) + case r >= '0' && r <= '9': + b.WriteRune(r) + case r == '-' || r == '_' || r == '.': + b.WriteRune(r) + default: + b.WriteRune('-') + } + } + if b.Len() == 0 { + return "model" + } + return b.String() +} + +func cleanComfyUIHFFilename(filename string) (string, bool) { + filename = strings.TrimSpace(filename) + if filename == "" { + return "", false + } + cleanName := path.Clean(filename) + if cleanName == "." || cleanName == ".." || strings.HasPrefix(cleanName, "../") || strings.HasPrefix(cleanName, "/") { + return "", false + } + return cleanName, true +} + +func hasComfyUISplitModelFiles(filenames []string) bool { + for _, filename := range filenames { + cleanName, ok := cleanComfyUIHFFilename(filename) + if !ok || !isComfyUIModelFile(cleanName) { + continue + } + lower := strings.ToLower(cleanName) + base := path.Base(lower) + if hasComfyUIPathPart(lower, "clip") || + hasComfyUIPathPart(lower, api.LLM_COMFYUI_TEXT_ENCODERS_DIR) || + hasComfyUIPathPartPrefix(lower, "text_encoder") || + hasComfyUIPathPart(lower, "vae") || + hasComfyUIPathPart(lower, "unet") || + hasComfyUIPathPart(lower, "transformer") || + hasComfyUIPathPart(lower, api.LLM_COMFYUI_DIFFUSION_MODELS_DIR) || + strings.Contains(base, "t5xxl") || + strings.HasPrefix(base, "qwen") || + strings.HasPrefix(base, "clip_") || + base == "ae.safetensors" || + strings.Contains(base, "vae") { + return true + } + } + return false +} + +func classifyComfyUIModelFile(filename string, splitRepo bool) (string, bool) { + cleanName, ok := cleanComfyUIHFFilename(filename) + if !ok || !isComfyUIModelFile(cleanName) { + return "", false + } + lower := strings.ToLower(cleanName) + base := path.Base(lower) + switch { + case strings.Contains(lower, "ip-adapter") || strings.Contains(lower, "ipadapter"): + return api.LLM_COMFYUI_IPADAPTER_DIR, true + case strings.Contains(lower, "controlnet") || strings.Contains(lower, "control-net"): + return api.LLM_COMFYUI_CONTROLNET_DIR, true + case strings.Contains(lower, "lora") || strings.Contains(lower, "lycoris"): + return api.LLM_COMFYUI_LORAS_DIR, true + case strings.Contains(lower, "clip_vision") || strings.Contains(lower, "clip-vision"): + return api.LLM_COMFYUI_CLIP_VISION_DIR, true + case hasComfyUIPathPart(lower, "clip") || + hasComfyUIPathPart(lower, api.LLM_COMFYUI_TEXT_ENCODERS_DIR) || + hasComfyUIPathPartPrefix(lower, "text_encoder") || + strings.Contains(base, "t5xxl") || + strings.HasPrefix(base, "qwen") || + strings.HasPrefix(base, "umt5") || + strings.HasPrefix(base, "clip_"): + return api.LLM_COMFYUI_TEXT_ENCODERS_DIR, true + case hasComfyUIPathPart(lower, "vae") || + strings.Contains(lower, "autoencoder") || + base == "ae.safetensors" || + strings.Contains(base, "vae"): + return api.LLM_COMFYUI_VAE_DIR, true + case strings.Contains(lower, "upscale") || + strings.Contains(lower, "realesrgan") || + strings.Contains(lower, "ultrasharp") || + strings.Contains(lower, "esrgan"): + return api.LLM_COMFYUI_UPSCALE_MODELS_DIR, true + case hasComfyUIPathPart(lower, api.LLM_COMFYUI_EMBEDDINGS_DIR) || strings.Contains(lower, "embedding"): + return api.LLM_COMFYUI_EMBEDDINGS_DIR, true + case hasComfyUIPathPart(lower, api.LLM_COMFYUI_STYLE_MODELS_DIR) || strings.Contains(lower, "style_model"): + return api.LLM_COMFYUI_STYLE_MODELS_DIR, true + case hasComfyUIPathPart(lower, "unet") || + hasComfyUIPathPart(lower, "transformer") || + hasComfyUIPathPart(lower, api.LLM_COMFYUI_DIFFUSION_MODELS_DIR) || + strings.Contains(base, "unet") || + strings.Contains(base, "diffusion") || + (splitRepo && strings.Contains(base, "flux")): + return api.LLM_COMFYUI_DIFFUSION_MODELS_DIR, true + default: + return api.LLM_COMFYUI_CHECKPOINTS_DIR, true + } +} + +func isComfyUIModelFile(filename string) bool { + ext := strings.ToLower(path.Ext(filename)) + switch ext { + case ".safetensors", ".ckpt", ".pt", ".pth", ".bin", ".gguf": + return true + default: + return false + } +} + +func hasComfyUIPathPart(filename string, part string) bool { + for _, item := range strings.Split(filename, "/") { + if item == part { + return true + } + } + return false +} + +func hasComfyUIPathPartPrefix(filename string, prefix string) bool { + for _, item := range strings.Split(filename, "/") { + if strings.HasPrefix(item, prefix) { + return true + } + } + return false +} + +func parseComfyUILegacyHFModelDirName(dirName string) (string, string, bool) { + encoded, ok := strings.CutPrefix(strings.TrimSpace(dirName), "hf-") + if !ok { + return "", "", false + } + repoPart, tagPart, ok := strings.Cut(encoded, "--") + if !ok || repoPart == "" || tagPart == "" { + return "", "", false + } + repoID, err := base64.RawURLEncoding.DecodeString(repoPart) + if err != nil { + return "", "", false + } + tag, err := base64.RawURLEncoding.DecodeString(tagPart) + if err != nil { + return "", "", false + } + repoIDStr := strings.TrimSpace(string(repoID)) + tagStr := strings.TrimSpace(string(tag)) + if repoIDStr == "" || tagStr == "" { + return "", "", false + } + return repoIDStr, tagStr, true +} + +func normalizeComfyUIRepoIDPath(repoID string) (string, error) { + repoID = strings.TrimSpace(repoID) + if repoID == "" { + return "", errors.Error("repo_id is empty") + } + repoPath := path.Clean(repoID) + if repoPath == "." || repoPath == ".." || strings.HasPrefix(repoPath, "../") || strings.HasPrefix(repoPath, "/") { + return "", errors.Errorf("invalid repo_id path %q", repoID) + } + for _, part := range strings.Split(repoPath, "/") { + if part == "" || part == "." || part == ".." { + return "", errors.Errorf("invalid repo_id path %q", repoID) + } + } + return repoPath, nil +} + +func extractComfyUIRepoIDPath(modelPath string) (string, bool) { + cleanPath := path.Clean(strings.TrimSpace(modelPath)) + if cleanPath == "." || cleanPath == "/" { + return "", false + } + for _, modelTypeDir := range getComfyUIModelTypeDirs() { + prefixes := []string{ + path.Join(api.LLM_COMFYUI_MODELS_PATH, modelTypeDir), + path.Join("/", modelTypeDir), + path.Join(api.LLM_COMFYUI_MODELS_VOLUME_SUBDIR, modelTypeDir), + } + for _, prefix := range prefixes { + if cleanPath == prefix || !strings.HasPrefix(cleanPath, prefix+"/") { + continue + } + repoPath := strings.TrimPrefix(cleanPath, prefix+"/") + if normalized, err := normalizeComfyUIRepoIDPath(repoPath); err == nil { + return normalized, true + } + return "", false + } + } + normalized, err := normalizeComfyUIRepoIDPath(cleanPath) + if err != nil { + return "", false + } + return normalized, true +} + +func getComfyUIModelDirByLLMType(llmType api.LLMContainerType) (string, error) { + switch llmType { + case api.LLM_CONTAINER_COMFYUI: + return api.LLM_COMFYUI_CHECKPOINTS_DIR, nil + default: + return "", errors.Errorf("unsupported comfyui llm_type %q", llmType) + } +} + +func getComfyUILLMTypeByModelDir(modelDir string) api.LLMContainerType { + return api.LLM_CONTAINER_COMFYUI +} + +func getComfyUILLMTypeByModelPath(modelPath string) api.LLMContainerType { + cleanPath := path.Clean(strings.TrimSpace(modelPath)) + for _, modelTypeDir := range getComfyUIModelTypeDirs() { + prefixes := []string{ + path.Join(api.LLM_COMFYUI_MODELS_PATH, modelTypeDir), + path.Join("/", modelTypeDir), + path.Join(api.LLM_COMFYUI_MODELS_VOLUME_SUBDIR, modelTypeDir), + } + for _, prefix := range prefixes { + if cleanPath == prefix || strings.HasPrefix(cleanPath, prefix+"/") { + return getComfyUILLMTypeByModelDir(modelTypeDir) + } + } + } + return getComfyUILLMTypeByModelDir(path.Base(path.Dir(cleanPath))) +} + +func getComfyUIModelTypeDirs() []string { + return []string{ + api.LLM_COMFYUI_CHECKPOINTS_DIR, + api.LLM_COMFYUI_LORAS_DIR, + api.LLM_COMFYUI_VAE_DIR, + api.LLM_COMFYUI_CONTROLNET_DIR, + api.LLM_COMFYUI_EMBEDDINGS_DIR, + api.LLM_COMFYUI_UPSCALE_MODELS_DIR, + api.LLM_COMFYUI_TEXT_ENCODERS_DIR, + api.LLM_COMFYUI_CLIP_DIR, + api.LLM_COMFYUI_CLIP_VISION_DIR, + api.LLM_COMFYUI_UNET_DIR, + api.LLM_COMFYUI_DIFFUSION_MODELS_DIR, + api.LLM_COMFYUI_IPADAPTER_DIR, + "diffusers", + api.LLM_COMFYUI_STYLE_MODELS_DIR, + } +} + +func buildComfyUIProbeCommand() string { + filePatterns := make([]string, 0, len(getComfyUIModelTypeDirs())) + dirPatterns := make([]string, 0, len(getComfyUIModelTypeDirs())) + for _, modelTypeDir := range getComfyUIModelTypeDirs() { + modelDir := shellQuoteSingle(path.Join(api.LLM_COMFYUI_MODELS_PATH, modelTypeDir)) + filePatterns = append(filePatterns, modelDir+"/*") + dirPatterns = append(dirPatterns, modelDir+"/*@*/", modelDir+"/hf-*/", modelDir+"/*/*/") + } + return fmt.Sprintf( + "for f in %s; do [ -f \"$f\" ] && case \"$f\" in *.safetensors|*.ckpt|*.pt|*.pth|*.bin|*.gguf) du -sk \"$f\";; esac; done; for d in %s; do [ -d \"$d\" ] && du -sk \"$d\"; done", + strings.Join(filePatterns, " "), + strings.Join(dirPatterns, " "), + ) } diff --git a/pkg/llm/drivers/llm_container/huggingface_model_dir.go b/pkg/llm/drivers/llm_container/huggingface_model_dir.go new file mode 100644 index 0000000000..7bef0bdd3a --- /dev/null +++ b/pkg/llm/drivers/llm_container/huggingface_model_dir.go @@ -0,0 +1,39 @@ +package llm_container + +import ( + "net/url" + "strings" +) + +const hfModelDirSeparator = "@" + +func buildVLLMModelDirName(modelName string, modelTag string) string { + modelName = strings.TrimSpace(modelName) + modelTag = resolveHfdRevision(modelTag) + return url.PathEscape(modelName) + hfModelDirSeparator + url.PathEscape(modelTag) +} + +func parseVLLMModelDirName(dirName string) (string, string, bool) { + namePart, tagPart, ok := strings.Cut(strings.TrimSpace(dirName), hfModelDirSeparator) + if !ok || namePart == "" || tagPart == "" { + return "", "", false + } + modelName, err := url.PathUnescape(namePart) + if err != nil || modelName == "" { + return "", "", false + } + modelTag, err := url.PathUnescape(tagPart) + if err != nil || modelTag == "" { + return "", "", false + } + return modelName, modelTag, true +} + +func containsString(items []string, target string) bool { + for _, item := range items { + if item == target { + return true + } + } + return false +} diff --git a/pkg/llm/models/instantmodel.go b/pkg/llm/models/instantmodel.go index 692505c13a..e82ec267e2 100644 --- a/pkg/llm/models/instantmodel.go +++ b/pkg/llm/models/instantmodel.go @@ -313,9 +313,10 @@ func (man *SInstantModelManager) ValidateCreateData( return input, errors.Wrap(err, "SSharableVirtualResourceBaseManager.ValidateCreateData") } - if !apis.IsLLMContainerType(string(input.LlmType)) { + if !apis.IsLLMInstantModelType(string(input.LlmType)) { return input, errors.Wrapf(httperrors.ErrInvalidFormat, "invalid llm_type %s", input.LlmType) } + input = normalizeInstantModelCreateInput(input) if len(input.ImageId) > 0 { img, err := fetchImage(ctx, userCred, input.ImageId) @@ -346,10 +347,13 @@ func (man *SInstantModelManager) ValidateCreateData( if err != nil { return input, errors.Wrap(err, "GetLLMContainerInstantModelDriver") } - _, err = drv.ValidateMounts(input.Mounts, input.ModelName, input.ModelTag) + input.Mounts, err = drv.ValidateMounts(input.Mounts, input.ModelName, input.ModelTag) if err != nil { return input, errors.Wrap(err, "validateMounts") } + if len(input.Mounts) == 0 { + return input, errors.Wrap(errors.ErrEmpty, "empty mounts") + } } input.Enabled = nil @@ -421,11 +425,7 @@ func (model *SInstantModel) PostCreate( return } if input.ImageId == "" && (input.DoNotImport == nil || !*input.DoNotImport) { - model.startImportTask(ctx, userCred, apis.InstantModelImportInput{ - LlmType: input.LlmType, - ModelName: input.ModelName, - ModelTag: input.ModelTag, - }) + model.startImportTask(ctx, userCred, buildInstantModelImportInputFromCreate(input)) } } @@ -566,6 +566,56 @@ func (man *SInstantModelManager) FindInstantModel(mdlId, tag string, isEnabled b return &mdls[0], nil } +func (man *SInstantModelManager) FindInstantModelByLLMType(mdlId, tag string, llmType apis.LLMContainerType, isEnabled bool) (*SInstantModel, error) { + q := man.Query().Equals("model_id", mdlId).Equals("llm_type", string(llmType)).Equals("status", imageapi.IMAGE_STATUS_ACTIVE) + if isEnabled { + q = q.IsTrue("enabled") + } + q = q.Desc("created_at") + + mdls := make([]SInstantModel, 0) + err := db.FetchModelObjects(man, q, &mdls) + if err != nil { + return nil, errors.Wrap(err, "FetchModelObjects") + } + if len(mdls) == 0 { + return nil, nil + } + if len(tag) > 0 { + for i := range mdls { + if mdls[i].ModelTag == tag { + return &mdls[i], nil + } + } + } + return &mdls[0], nil +} + +func (man *SInstantModelManager) FindInstantModelByMountAndLLMType(mount string, llmType apis.LLMContainerType, isEnabled bool) (*SInstantModel, error) { + q := man.Query().Equals("llm_type", string(llmType)).Equals("status", imageapi.IMAGE_STATUS_ACTIVE).Contains("mounts", mount) + if isEnabled { + q = q.IsTrue("enabled") + } + q = q.Desc("created_at") + + mdls := make([]SInstantModel, 0) + err := db.FetchModelObjects(man, q, &mdls) + if err != nil { + return nil, errors.Wrap(err, "FetchModelObjects") + } + if len(mdls) == 0 { + return nil, nil + } + for i := range mdls { + for _, mdlMount := range mdls[i].Mounts { + if mdlMount == mount { + return &mdls[i], nil + } + } + } + return nil, nil +} + func (model *SInstantModel) PerformEnable( ctx context.Context, userCred mcclient.TokenCredential, diff --git a/pkg/llm/models/instantmodel_huggingface.go b/pkg/llm/models/instantmodel_huggingface.go index f02e2f3222..ed27628845 100644 --- a/pkg/llm/models/instantmodel_huggingface.go +++ b/pkg/llm/models/instantmodel_huggingface.go @@ -22,6 +22,7 @@ import ( const ( huggingFaceMirrorEndpoint = "https://hf-mirror.com" huggingFaceImportMode = "snapshot" + huggingFaceSortDirection = -1 ) type huggingFaceSearchItem struct { @@ -76,13 +77,14 @@ func (man *SInstantModelManager) getPropertyHuggingFaceSearch(ctx context.Contex } input.Author = strings.TrimSpace(input.Author) input.Sort = strings.TrimSpace(input.Sort) + input.Cursor = strings.TrimSpace(input.Cursor) for i := range input.Filter { input.Filter[i] = strings.TrimSpace(input.Filter[i]) } searchURL := buildHuggingFaceSearchURL(input) - body, err := huggingFaceHTTPGet(ctx, searchURL) + body, header, err := huggingFaceHTTPGet(ctx, searchURL) if err != nil { return nil, errors.Wrap(err, "huggingFaceHTTPGet") } @@ -91,7 +93,12 @@ func (man *SInstantModelManager) getPropertyHuggingFaceSearch(ctx context.Contex if err := json.Unmarshal(body, &items); err != nil { return nil, errors.Wrap(err, "json.Unmarshal") } - return jsonutils.Marshal(normalizeHuggingFaceSearchResults(items)), nil + nextCursor := getHuggingFaceNextCursor(header) + return jsonutils.Marshal(apis.InstantModelHuggingFaceSearchOutput{ + Data: normalizeHuggingFaceSearchResults(items), + NextCursor: nextCursor, + HasMore: nextCursor != "", + }), nil } func (man *SInstantModelManager) GetPropertyHuggingfaceRepoInfo(ctx context.Context, userCred mcclient.TokenCredential, query jsonutils.JSONObject) (jsonutils.JSONObject, error) { @@ -213,13 +220,46 @@ func buildHuggingFaceSearchURL(input apis.InstantModelHuggingFaceSearchInput) st if input.Sort != "" { queryParts = append(queryParts, fmt.Sprintf("sort=%s", url.QueryEscape(input.Sort))) } - if input.Direction != 0 { - queryParts = append(queryParts, fmt.Sprintf("direction=%d", input.Direction)) + queryParts = append(queryParts, fmt.Sprintf("direction=%d", huggingFaceSortDirection)) + if input.Cursor != "" { + queryParts = append(queryParts, fmt.Sprintf("cursor=%s", url.QueryEscape(input.Cursor))) } queryParts = append(queryParts, fmt.Sprintf("limit=%d", input.Limit)) return fmt.Sprintf("%s/api/models?%s", huggingFaceMirrorEndpoint, strings.Join(queryParts, "&")) } +func getHuggingFaceNextCursor(header http.Header) string { + for _, linkHeader := range header.Values("Link") { + for _, link := range strings.Split(linkHeader, ",") { + link = strings.TrimSpace(link) + if !isHuggingFaceNextLink(link) { + continue + } + start := strings.Index(link, "<") + end := strings.Index(link, ">") + if start < 0 || end <= start+1 { + continue + } + nextURL, err := url.Parse(link[start+1 : end]) + if err != nil { + continue + } + return strings.TrimSpace(nextURL.Query().Get("cursor")) + } + } + return "" +} + +func isHuggingFaceNextLink(link string) bool { + for _, part := range strings.Split(link, ";") { + part = strings.TrimSpace(part) + if strings.EqualFold(part, `rel="next"`) || strings.EqualFold(part, "rel=next") { + return true + } + } + return false +} + func isHuggingFaceGated(v interface{}) bool { switch value := v.(type) { case bool: @@ -240,25 +280,25 @@ func hasTag(tags []string, target string) bool { return false } -func huggingFaceHTTPGet(ctx context.Context, reqURL string) ([]byte, error) { +func huggingFaceHTTPGet(ctx context.Context, reqURL string) ([]byte, http.Header, error) { req, err := http.NewRequestWithContext(ctx, http.MethodGet, reqURL, nil) if err != nil { - return nil, errors.Wrap(err, "http.NewRequestWithContext") + return nil, nil, errors.Wrap(err, "http.NewRequestWithContext") } client := &http.Client{Timeout: 60 * time.Second} resp, err := client.Do(req) if err != nil { - return nil, errors.Wrap(err, "client.Do") + return nil, nil, errors.Wrap(err, "client.Do") } defer resp.Body.Close() if resp.StatusCode != http.StatusOK { - return nil, errors.Errorf("unexpected status code: %d", resp.StatusCode) + return nil, nil, errors.Errorf("unexpected status code: %d", resp.StatusCode) } body, err := io.ReadAll(resp.Body) if err != nil { - return nil, errors.Wrap(err, "io.ReadAll") + return nil, nil, errors.Wrap(err, "io.ReadAll") } - return body, nil + return body, resp.Header, nil } func getHuggingFaceRepoInfo(ctx context.Context, repoID string, revision string) (huggingFaceRepoInfoResponse, error) { @@ -266,7 +306,7 @@ func getHuggingFaceRepoInfo(ctx context.Context, repoID string, revision string) if revision != "" { repoURL = fmt.Sprintf("%s?revision=%s", repoURL, url.QueryEscape(revision)) } - body, err := huggingFaceHTTPGet(ctx, repoURL) + body, _, err := huggingFaceHTTPGet(ctx, repoURL) if err != nil { return huggingFaceRepoInfoResponse{}, errors.Wrap(err, "huggingFaceHTTPGet") } diff --git a/pkg/llm/models/instantmodel_huggingface_import.go b/pkg/llm/models/instantmodel_huggingface_import.go new file mode 100644 index 0000000000..ebe6ebb178 --- /dev/null +++ b/pkg/llm/models/instantmodel_huggingface_import.go @@ -0,0 +1,104 @@ +package models + +import ( + "strings" + + "yunion.io/x/jsonutils" + + imageapi "yunion.io/x/onecloud/pkg/apis/image" + apis "yunion.io/x/onecloud/pkg/apis/llm" +) + +const defaultHuggingFaceRevision = "main" + +func normalizeInstantModelSource(source string, repoID string) string { + source = strings.TrimSpace(source) + if strings.EqualFold(source, apis.InstantModelSourceHuggingFace) || (source == "" && strings.TrimSpace(repoID) != "") { + return apis.InstantModelSourceHuggingFace + } + return source +} + +func resolveImportRepoAndRevision(input apis.InstantModelImportInput) (string, string, string) { + repoID := strings.TrimSpace(input.RepoId) + source := normalizeInstantModelSource(input.Source, repoID) + revision := strings.TrimSpace(input.Revision) + if repoID == "" { + repoID = strings.TrimSpace(input.ModelName) + } + if revision == "" { + revision = strings.TrimSpace(input.ModelTag) + } + if source == apis.InstantModelSourceHuggingFace && repoID != "" && revision == "" { + revision = defaultHuggingFaceRevision + } + return source, repoID, revision +} + +func buildInstantModelImportInputFromCreate(input apis.InstantModelCreateInput) apis.InstantModelImportInput { + importInput := apis.InstantModelImportInput{ + Source: input.Source, + RepoId: input.RepoId, + Revision: input.Revision, + ModelName: input.ModelName, + ModelTag: input.ModelTag, + LlmType: input.LlmType, + } + source, repoID, revision := resolveImportRepoAndRevision(importInput) + importInput.Source = source + importInput.RepoId = repoID + importInput.Revision = revision + importInput.ModelName = repoID + importInput.ModelTag = revision + return importInput +} + +func normalizeInstantModelCreateInput(input apis.InstantModelCreateInput) apis.InstantModelCreateInput { + importInput := buildInstantModelImportInputFromCreate(input) + input.Source = importInput.Source + input.RepoId = importInput.RepoId + input.Revision = importInput.Revision + input.ModelName = importInput.ModelName + input.ModelTag = importInput.ModelTag + return input +} + +func buildInstantModelImageProperties(input apis.InstantModelImportInput, repoID string, resolvedRevision string) map[string]string { + source, _, requestedRevision := resolveImportRepoAndRevision(input) + properties := map[string]string{ + "llm_type": string(input.LlmType), + } + if input.ModelName != "" { + properties["model_name"] = input.ModelName + } + if input.ModelTag != "" { + properties["model_tag"] = input.ModelTag + } + if repoID != "" { + properties["model_name"] = repoID + properties["source_repo_id"] = repoID + } + if requestedRevision != "" { + properties["model_tag"] = requestedRevision + properties["source_requested_revision"] = requestedRevision + } + if source != "" { + properties["source"] = source + } + if resolvedRevision != "" { + properties["source_resolved_revision"] = resolvedRevision + } + return properties +} + +func withInstantModelPostOverlayImageProperties(properties map[string]string, pathMap map[string]string) map[string]string { + if len(pathMap) == 0 { + return properties + } + if properties == nil { + properties = make(map[string]string) + } + properties[imageapi.IMAGE_INTERNAL_PATH_MAP] = jsonutils.Marshal(pathMap).String() + properties[imageapi.IMAGE_USED_BY_POST_OVERLAY] = "true" + return properties +} diff --git a/pkg/llm/models/llm_instant_model_sync.go b/pkg/llm/models/llm_instant_model_sync.go index 8ceb8bbca7..c01681cbfd 100644 --- a/pkg/llm/models/llm_instant_model_sync.go +++ b/pkg/llm/models/llm_instant_model_sync.go @@ -25,21 +25,33 @@ import ( "yunion.io/x/onecloud/pkg/util/logclient" ) +func getInstantModelPostOverlayVolumeMountIndex(drv ILLMContainerInstantModelDriver) int { + if drv.GetType() == apis.LLM_CONTAINER_COMFYUI { + return 1 + } + return 0 +} + func (llm *SLLM) getMountedInstantModels(ctx context.Context, probedExt map[string]apis.LLMInternalInstantMdlInfo) (map[string]struct{}, error) { container, err := llm.GetLLMSContainer(ctx) if err != nil { return nil, errors.Wrap(err, "GetSContainer") } + drv, err := GetLLMContainerInstantModelDriver(llm.GetLLMContainerDriver().GetType()) + if err != nil { + return nil, nil + } + volumeMountIndex := getInstantModelPostOverlayVolumeMountIndex(drv) if container.Spec == nil { return nil, errors.Wrap(errors.ErrEmpty, "no Spec") } - if len(container.Spec.VolumeMounts) == 0 { + if volumeMountIndex < 0 || volumeMountIndex >= len(container.Spec.VolumeMounts) { return nil, errors.Wrap(errors.ErrEmpty, "no VolumeMounts") } - if container.Spec.VolumeMounts[0].Disk == nil { + if container.Spec.VolumeMounts[volumeMountIndex].Disk == nil { return nil, errors.Wrap(errors.ErrEmpty, "no Disk") } - if len(container.Spec.VolumeMounts[0].Disk.PostOverlay) == 0 { + if len(container.Spec.VolumeMounts[volumeMountIndex].Disk.PostOverlay) == 0 { return nil, nil } mdlNameToId := make(map[string]string) @@ -47,14 +59,22 @@ func (llm *SLLM) getMountedInstantModels(ctx context.Context, probedExt map[stri mdlNameToId[model.Name+":"+model.Tag] = mdlId } mdlMap := make(map[string]struct{}) - postOverlays := container.Spec.VolumeMounts[0].Disk.PostOverlay - drv, err := GetLLMContainerInstantModelDriver(llm.GetLLMContainerDriver().GetType()) - if err != nil { - return nil, nil - } + postOverlays := container.Spec.VolumeMounts[volumeMountIndex].Disk.PostOverlay for i := range postOverlays { postOverlay := postOverlays[i] - mdlId := drv.GetInstantModelIdByPostOverlay(postOverlay, mdlNameToId) + mdlId := "" + if postOverlay.Image != nil { + instMdl, err := GetInstantModelManager().findInstantModelByImageId(postOverlay.Image.Id) + if err != nil { + return nil, errors.Wrapf(err, "findInstantModelByImageId %s", postOverlay.Image.Id) + } + if instMdl != nil { + mdlId = instMdl.Id + } + } + if mdlId == "" { + mdlId = drv.GetInstantModelIdByPostOverlay(postOverlay, mdlNameToId) + } if mdlId != "" { mdlMap[mdlId] = struct{}{} } @@ -277,7 +297,7 @@ func (llm *SLLM) PerformQuickModels(ctx context.Context, userCred mcclient.Token input.Models[i].LlmType = mdl.LlmType } } - if !apis.IsLLMContainerType(input.Models[i].LlmType) || apis.LLMContainerType(input.Models[i].LlmType) != llm.GetLLMContainerDriver().GetType() { + if !apis.IsLLMInstantModelType(input.Models[i].LlmType) || !apis.IsLLMInstantModelCompatible(apis.LLMContainerType(input.Models[i].LlmType), llm.GetLLMContainerDriver().GetType()) { errs = append(errs, errors.Wrapf(httperrors.ErrInvalidStatus, "model %s is not of type %s", input.Models[i].Id, llm.GetLLMContainerDriver().GetType())) } } @@ -423,7 +443,11 @@ func (llm *SLLM) RequestUnmountModel(ctx context.Context, userCred mcclient.Toke } var unmountOverlays []*commonapi.ContainerVolumeMountDiskPostOverlay - existingOverlays := container.Spec.VolumeMounts[0].Disk.PostOverlay + volumeMountIndex := getInstantModelPostOverlayVolumeMountIndex(drv) + if volumeMountIndex < 0 || volumeMountIndex >= len(container.Spec.VolumeMounts) || container.Spec.VolumeMounts[volumeMountIndex].Disk == nil { + return nil, nil, errors.Wrap(errors.ErrEmpty, "no instant model volume mount") + } + existingOverlays := container.Spec.VolumeMounts[volumeMountIndex].Disk.PostOverlay for i := range existingOverlays { eOverlay := existingOverlays[i] @@ -518,8 +542,12 @@ func (llm *SLLM) containerUnmountPaths(ctx context.Context, userCred mcclient.To return errors.Wrapf(errors.ErrInvalidStatus, "cannot unmount post path in status %s", ctr.Status) } + drv, err := GetLLMContainerInstantModelDriver(llm.GetLLMContainerDriver().GetType()) + if err != nil { + return nil + } params := computeapi.ContainerVolumeMountRemovePostOverlayInput{ - Index: 0, + Index: getInstantModelPostOverlayVolumeMountIndex(drv), PostOverlay: overlays, UseLazy: true, ClearLayers: true, @@ -560,8 +588,12 @@ func (llm *SLLM) containerMountPaths(ctx context.Context, userCred mcclient.Toke if !computeapi.ContainerFinalStatus.Has(ctr.Status) { return errors.Wrapf(errors.ErrInvalidStatus, "cannot mount post path in status %s", ctr.Status) } + drv, err := GetLLMContainerInstantModelDriver(llm.GetLLMContainerDriver().GetType()) + if err != nil { + return nil + } params := computeapi.ContainerVolumeMountAddPostOverlayInput{ - Index: 0, + Index: getInstantModelPostOverlayVolumeMountIndex(drv), PostOverlay: overlays, } _, err = compute.Containers.PerformAction(s, ctr.Id, "add-volume-mount-post-overlay", jsonutils.Marshal(params)) diff --git a/pkg/llm/models/llm_save_instant_model.go b/pkg/llm/models/llm_save_instant_model.go index ec16ed2104..36b027de9d 100644 --- a/pkg/llm/models/llm_save_instant_model.go +++ b/pkg/llm/models/llm_save_instant_model.go @@ -6,6 +6,7 @@ import ( "io" "net/http" "os" + "strings" "time" "yunion.io/x/jsonutils" @@ -41,29 +42,53 @@ func (llm *SLLM) PerformSaveInstantModel( return nil, httperrors.NewInvalidStatusError("LLM is not running") } - mdlInfos, err := llm.getProbedInstantModelsExt(ctx, userCred) - if err != nil { - return nil, errors.Wrap(err, "getProbedPackagesExt") - } - var mdlInfo *api.LLMInternalInstantMdlInfo - for _, info := range mdlInfos { - if info.ModelId == input.ModelId { - mdlInfo = &info - break + modelId := strings.TrimSpace(input.ModelId) + if modelId == "" { + return nil, httperrors.NewMissingParameterError("model_id") + } + mountDirs := make([]string, 0) + if len(input.Mounts) > 0 { + drv, err := GetLLMContainerInstantModelDriver(llm.GetLLMContainerDriver().GetType()) + if err != nil { + return nil, errors.Wrap(err, "GetLLMContainerInstantModelDriver") + } + mountDirs, err = drv.ValidateMounts(input.Mounts, "", "") + if err != nil { + return nil, errors.Wrap(err, "validateMounts") + } + if len(mountDirs) == 0 { + return nil, errors.Wrap(errors.ErrEmpty, "empty mounts") + } + } else { + mdlInfos, err := llm.getProbedInstantModelsExt(ctx, userCred) + if err != nil { + return nil, errors.Wrap(err, "getProbedPackagesExt") } - } - if mdlInfo == nil { - return nil, httperrors.NewBadRequestError("ModelId %s not found", input.ModelId) - } - mountDirs, err := llm.detectModelPaths(ctx, userCred, *mdlInfo) - if err != nil { - return nil, errors.Wrap(err, "detectModelPaths") + for _, info := range mdlInfos { + if info.ModelId == input.ModelId { + mdlInfo = &info + break + } + } + if mdlInfo == nil { + return nil, httperrors.NewBadRequestError("ModelId %s not found", input.ModelId) + } + + mountDirs, err = llm.detectModelPaths(ctx, userCred, *mdlInfo) + if err != nil { + return nil, errors.Wrap(err, "detectModelPaths") + } + modelId = mdlInfo.ModelId } if len(input.ModelFullName) == 0 { - input.ModelFullName = fmt.Sprintf("%s-%s", mdlInfo.Name+":"+mdlInfo.Tag, time.Now().Format("060102")) + if mdlInfo != nil { + input.ModelFullName = fmt.Sprintf("%s-%s", mdlInfo.Name+":"+mdlInfo.Tag, time.Now().Format("060102")) + } else { + input.ModelFullName = fmt.Sprintf("%s-%s", modelId, time.Now().Format("060102")) + } } var ownerId mcclient.IIdentityProvider @@ -97,16 +122,24 @@ func (llm *SLLM) PerformSaveInstantModel( modelName, modelTag, _ := llm.GetLargeLanguageModelName(input.ModelFullName) if len(modelName) == 0 { - modelName = mdlInfo.Name + if mdlInfo != nil { + modelName = mdlInfo.Name + } else { + modelName = modelId + } } if len(modelTag) == 0 { - modelTag = mdlInfo.Tag + if mdlInfo != nil { + modelTag = mdlInfo.Tag + } else { + modelTag = "main" + } } drv := llm.GetLLMContainerDriver() instantModelCreateInput := api.InstantModelCreateInput{ LlmType: drv.GetType(), - ModelId: mdlInfo.ModelId, + ModelId: modelId, ModelName: modelName, ModelTag: modelTag, Mounts: mountDirs, @@ -154,7 +187,7 @@ func (llm *SLLM) DoSaveModelImage(ctx context.Context, userCred mcclient.TokenCr saveImageInput := computeapi.ContainerSaveVolumeMountToImageInput{ GenerateName: input.ModelFullName, Notes: fmt.Sprintf("instance model image for %s(%s)", instantModel.ModelId, instantModel.ModelName+":"+instantModel.ModelTag), - Index: 0, + Index: getInstantModelSaveVolumeMountIndex(drv), Dirs: saveDirs, UsedByPostOverlay: true, DirPrefix: prefix, @@ -184,6 +217,10 @@ func (llm *SLLM) DoSaveModelImage(ctx context.Context, userCred mcclient.TokenCr return nil } +func getInstantModelSaveVolumeMountIndex(drv ILLMContainerInstantModelDriver) int { + return getInstantModelPostOverlayVolumeMountIndex(drv) +} + func (llm *SLLM) StartSaveModelImageTask(ctx context.Context, userCred mcclient.TokenCredential, input api.LLMSaveInstantModelInput) (*taskman.STask, error) { llm.SetStatus(ctx, userCred, api.LLM_STATUS_START_SAVE_MODEL, "StartSaveModelImageTask") diff --git a/pkg/mcclient/options/llm/instantmodel.go b/pkg/mcclient/options/llm/instantmodel.go index ae95752385..aebf8f2e97 100644 --- a/pkg/mcclient/options/llm/instantmodel.go +++ b/pkg/mcclient/options/llm/instantmodel.go @@ -29,11 +29,14 @@ func (o *LLMInstantModelShowOptions) Params() (jsonutils.JSONObject, error) { type LLMInstantModelCreateOptions struct { options.BaseCreateOptions - LLM_TYPE string `help:"llm container type" choices:"ollama|vllm" json:"llm_type"` + LLM_TYPE string `help:"llm instant model type" choices:"ollama|vllm|comfyui" json:"llm_type"` MODEL_NAME string `json:"model_name"` MODEL_TAG string `json:"model_tag"` - ImageId string `json:"image_id"` + Source string `help:"model source, e.g. huggingface" json:"source"` + RepoId string `help:"huggingface repo id, e.g. Qwen/Qwen3-8B" json:"repo_id"` + Revision string `help:"huggingface revision, e.g. main or refs/pr/7" json:"revision"` + ImageId string `json:"image_id"` Mounts []string `json:"mounts"` } @@ -63,9 +66,11 @@ func (o *LLMInstantModelDeleteOptions) Params() (jsonutils.JSONObject, error) { } type LLMInstantModelImportOptions struct { - LLM_TYPE string `help:"llm container type" choices:"ollama|vllm" json:"llm_type"` + LLM_TYPE string `help:"llm instant model type" choices:"ollama|vllm|comfyui" json:"llm_type"` MODEL_NAME string `help:"model name to import, e.g. qwen3 or Qwen/Qwen3-VL-8B-Instruct" json:"model_name"` MODEL_TAG string `help:"model tag to import, e.g. 8b" json:"model_tag"` + REPO_ID string `help:"huggingface repo id, e.g. Qwen/Qwen3-8B" json:"repo_id"` + REVISION string `help:"huggingface revision, e.g. main or refs/pr/7" json:"revision"` } func (o *LLMInstantModelImportOptions) Params() (jsonutils.JSONObject, error) { @@ -73,6 +78,11 @@ func (o *LLMInstantModelImportOptions) Params() (jsonutils.JSONObject, error) { ModelName: o.MODEL_NAME, ModelTag: o.MODEL_TAG, LlmType: api.LLMContainerType(o.LLM_TYPE), + RepoId: o.REPO_ID, + Revision: o.REVISION, + } + if o.REPO_ID != "" { + input.Source = api.InstantModelSourceHuggingFace } return jsonutils.Marshal(input), nil } diff --git a/pkg/mcclient/options/llm/instantmodel_huggingface.go b/pkg/mcclient/options/llm/instantmodel_huggingface.go index 68c14d3b7c..858a43757c 100644 --- a/pkg/mcclient/options/llm/instantmodel_huggingface.go +++ b/pkg/mcclient/options/llm/instantmodel_huggingface.go @@ -3,12 +3,12 @@ package llm import "yunion.io/x/jsonutils" type LLMInstantModelHuggingFaceSearchOptions struct { - Q string `help:"huggingface query string" json:"q"` - Author string `help:"filter by author or organization" json:"author"` - Filter []string `help:"filter by tags, e.g. text-generation or pytorch" json:"filter"` - Direction int `help:"sort direction, e.g. -1 for descending or 1 for ascending" json:"direction"` - Limit int `help:"max number of search results" json:"limit"` - Sort string `help:"sort order, e.g. downloads|likes|updated" json:"sort"` + Q string `help:"huggingface query string" json:"q"` + Author string `help:"filter by author or organization" json:"author"` + Filter []string `help:"filter by tags, e.g. text-generation or pytorch" json:"filter"` + Limit int `help:"max number of search results" json:"limit"` + Sort string `help:"sort order, e.g. downloads|likes|updated" json:"sort"` + Cursor string `help:"cursor returned by previous huggingface-search response" json:"cursor"` } func (o *LLMInstantModelHuggingFaceSearchOptions) Params() (jsonutils.JSONObject, error) { diff --git a/pkg/mcclient/options/llm/llm.go b/pkg/mcclient/options/llm/llm.go index 80f3cf9d64..b7c0021ab5 100644 --- a/pkg/mcclient/options/llm/llm.go +++ b/pkg/mcclient/options/llm/llm.go @@ -209,8 +209,9 @@ func (opts *LLMProviderModelsOptions) Params() (jsonutils.JSONObject, error) { type LLMSaveInstantModelOptions struct { LLMIdOptions - MODEL_ID string `help:"llm model id, e.g. 500a1f067a9f"` - Name string `help:"instant app name, e.g. qwen3:8b"` + MODEL_ID string `help:"llm model id, e.g. 500a1f067a9f"` + Name string `help:"instant app name, e.g. qwen3:8b"` + Mount []string `help:"model file or directory path to package; repeat to package multiple ComfyUI files together" json:"-"` AutoRestart bool } @@ -219,6 +220,7 @@ func (opts *LLMSaveInstantModelOptions) Params() (jsonutils.JSONObject, error) { input := api.LLMSaveInstantModelInput{ ModelId: opts.MODEL_ID, ModelFullName: opts.Name, + Mounts: opts.Mount, // AutoRestart: opts.AutoRestart, } return jsonutils.Marshal(input), nil