From e5c356b74b8ff9630a94647230419a0b89777dd4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=AB=98=E9=AD=8F=E6=B4=AA?= Date: Tue, 4 Nov 2025 16:33:28 +0800 Subject: [PATCH] refactor(model): Refactor the model download logic, support multiple solutions, and optimize automatic resource deployment Co-authored-by: Qwen-Coder --- __tests__/e2e/ci-mac-linux.sh | 17 +- __tests__/e2e/model/deploy_and_test_model.py | 30 +- __tests__/e2e/model/s_file.yaml | 75 ++ __tests__/e2e/model/test.py | 109 +++ __tests__/ut/commands/artModelService_test.ts | 593 +++++++++++++++ __tests__/ut/commands/modelService_test.ts | 452 +++++++++++ __tests__/ut/commands/model_test.ts | 711 +++++++++--------- __tests__/ut/commands/model_utils_test.ts | 313 ++++++++ package-lock.json | 14 +- package.json | 2 +- publish.yaml | 2 +- src/subCommands/model/constants.ts | 6 + src/subCommands/model/fileManager.ts | 382 ++++++++++ src/subCommands/model/index.ts | 564 ++++++-------- src/subCommands/model/model.ts | 141 ++++ src/subCommands/model/utils/index.ts | 149 ++++ 16 files changed, 2840 insertions(+), 720 deletions(-) create mode 100644 __tests__/e2e/model/s_file.yaml create mode 100644 __tests__/e2e/model/test.py create mode 100644 __tests__/ut/commands/artModelService_test.ts create mode 100644 __tests__/ut/commands/modelService_test.ts create mode 100644 __tests__/ut/commands/model_utils_test.ts create mode 100644 src/subCommands/model/constants.ts create mode 100644 src/subCommands/model/fileManager.ts create mode 100644 src/subCommands/model/model.ts create mode 100644 src/subCommands/model/utils/index.ts diff --git a/__tests__/e2e/ci-mac-linux.sh b/__tests__/e2e/ci-mac-linux.sh index 3794c545..9c8afa02 100755 --- a/__tests__/e2e/ci-mac-linux.sh +++ b/__tests__/e2e/ci-mac-linux.sh @@ -30,11 +30,18 @@ fi echo "test model download" cd model pip install -r requirements.txt -export fc_component_function_name=model-$(uname)-$(uname -m)-$RANDSTR -python deploy_and_test_model.py --model-id iic/cv_LightweightEdge_ocr-recognitoin-general_damo --region cn-shanghai --auto-cleanup -python deploy_and_test_model.py --model-id Qwen/Qwen2.5-0.5B-Instruct --region cn-shanghai --auto-cleanup -python deploy_and_test_model.py --model-id iic/cv_LightweightEdge_ocr-recognitoin-general_damo --region cn-shanghai --storage oss --auto-cleanup -python deploy_and_test_model.py --model-id Qwen/Qwen2.5-0.5B-Instruct --region cn-shanghai --storage oss --auto-cleanup +export fc_component_function_name=model-$(uname)-$(uname -m)-$RANDSTR-$RANDOM +python -u deploy_and_test_model.py --model-id iic/cv_LightweightEdge_ocr-recognitoin-general_damo --region cn-shanghai --auto-cleanup +python -u deploy_and_test_model.py --model-id Qwen/Qwen2.5-0.5B-Instruct --region cn-shanghai --auto-cleanup +python -u deploy_and_test_model.py --model-id iic/cv_LightweightEdge_ocr-recognitoin-general_damo --region cn-shanghai --storage oss --auto-cleanup +python -u deploy_and_test_model.py --model-id Qwen/Qwen2.5-0.5B-Instruct --region cn-shanghai --storage oss --auto-cleanup + +echo "test model s_file.yaml" +# python -u test.py +s model download -t s_file.yaml +s deploy -y -t s_file.yaml +s model remove -t s_file.yaml +s remove -y -t s_file.yaml cd .. echo "test go runtime" diff --git a/__tests__/e2e/model/deploy_and_test_model.py b/__tests__/e2e/model/deploy_and_test_model.py index e805d682..1712c501 100644 --- a/__tests__/e2e/model/deploy_and_test_model.py +++ b/__tests__/e2e/model/deploy_and_test_model.py @@ -54,7 +54,8 @@ def deploy_model(model_id: str, region: str = "cn-hangzhou", storage: str = "nas tuple: (部署的URL, 配置文件路径) """ # 生成函数名称 - function_name = f"test-{simple_hash(model_id)}" + # 加个随机数 + function_name = f"test-{simple_hash(model_id)}-{secrets.token_hex(4)}" # 准备请求数据 deploy_data = { @@ -177,17 +178,6 @@ def test_model(model_id: str, deploy_url: str, s_yaml_file: str = None): model_detail_url = f"{deploy_url}/model/info" print(f"正在获取部署后的模型服务详情: {model_detail_url}") - try: - detail_response = requests.get( - model_detail_url, headers={"Authorization": f"Bearer {token}"} - ) - if detail_response.status_code == 200: - print(f"部署后的模型服务详情: {detail_response.text}") - else: - print(f"获取部署后的模型服务详情失败: {detail_response.status_code}") - except Exception as e: - print(f"获取部署后的模型服务详情时出错: {e}") - # 检查是否是vLLM模型(通过配置文件中的启动命令) is_vllm_model = False if s_yaml_file: @@ -212,6 +202,22 @@ def test_model(model_id: str, deploy_url: str, s_yaml_file: str = None): print("检测到vLLM模型,将使用专用测试方法") except Exception as e: print(f"检查模型类型时出错: {e}") + + # 先调用模型详情接口 + model_detail_url = f"{deploy_url}/model/info" + if is_vllm_model: + model_detail_url = f"{deploy_url}/v1/models" + print(f"正在获取部署后的模型服务详情: {model_detail_url}") + try: + detail_response = requests.get( + model_detail_url, headers={"Authorization": f"Bearer {token}"} + ) + if detail_response.status_code == 200: + print(f"部署后的模型服务详情: {detail_response.text}") + else: + print(f"获取部署后的模型服务详情失败: {detail_response.status_code}") + except Exception as e: + print(f"获取部署后的模型服务详情时出错: {e}") if is_vllm_model: # 对于vLLM模型,使用专门的测试方法 diff --git a/__tests__/e2e/model/s_file.yaml b/__tests__/e2e/model/s_file.yaml new file mode 100644 index 00000000..f14fa742 --- /dev/null +++ b/__tests__/e2e/model/s_file.yaml @@ -0,0 +1,75 @@ +edition: 3.0.0 +name: ai-model-app +access: quanxi + +vars: + region: 'cn-hangzhou' + +resources: + fc3: + component: ${env('fc_component_version', path('../../../'))} + type: Function + props: + logConfig: auto + functionName: fc3-model-files-${env('fc_component_function_name', 'nodejs18')} + instanceLifecycleConfig: + preStop: + handler: 'true' + timeout: 300 + gpuConfig: + gpuMemorySize: 16384 + gpuType: fc.gpu.tesla.1 + nasConfig: auto + runtime: custom-container + description: test model download files + cpu: 8 + customContainerConfig: + image: >- + cap-demo-public-registry.cn-hangzhou.cr.aliyuncs.com/cap-app/image-generation-comfyui-agent:v1.1.0-beta.0 + port: 9000 + triggers: + - triggerConfig: + methods: + - GET + - POST + - PUT + - PATCH + - DELETE + - HEAD + authType: anonymous + disableURLInternet: false + triggerName: fc3-model-files-${env('fc_component_function_name', 'nodejs18')} + qualifier: LATEST + description: http trigger for fc3-model-files-${env('fc_component_function_name', 'nodejs18')} + triggerType: http + version: 1.6.0 + timeout: 3600 + instanceConcurrency: 200 + diskSize: 61440 + memorySize: 32768 + internetAccess: true + environmentVariables: + BACKEND_TYPE: comfyui + MODEL_ASSET_DIR: /mnt/fc3-model-files-${env('fc_component_function_name', 'nodejs18')} + REGION: ${vars.region} + vpcConfig: auto + region: ${vars.region} + annotations: + modelConfig: + solution: funArt + source: + uri: 'oss://dipper-cache-cn-hangzhou.oss-cn-hangzhou.aliyuncs.com' + downloadStrategy: + mode: once + target: + uri: 'nas://auto' + files: + - source: + path: base/comfyui/v0.3.59-alpha + target: + path: '' + - source: + path: >- + function-art/comfyui/models/checkpoints/v1-5-pruned-emaonly-fp16.safetensors + target: + path: models/checkpoints/v1-5-pruned-emaonly-fp16.safetensors diff --git a/__tests__/e2e/model/test.py b/__tests__/e2e/model/test.py new file mode 100644 index 00000000..44a39a23 --- /dev/null +++ b/__tests__/e2e/model/test.py @@ -0,0 +1,109 @@ +#!/usr/bin/env python3 +import os +import subprocess +import time +import sys + +def run_command_with_retry(cmd, silent=False, max_retries=3, timeout=300): + """执行命令并返回结果,支持重试""" + for i in range(max_retries): + if not silent: + print(f"执行命令: {cmd} (尝试 {i+1}/{max_retries})") + + try: + output = subprocess.check_output(cmd, shell=True, stderr=subprocess.STDOUT, text=True, timeout=timeout) + return 0, output, "" + except subprocess.TimeoutExpired: + print(f"命令超时: {cmd}") + if i == max_retries - 1: + return -1, "", "Command timed out after retries" + except Exception as e: + print(f"命令执行出错: {str(e)}") + if i == max_retries - 1: + return -1, "", str(e) + + return -1, "", "Failed after retries" + +def main(): + # 获取函数名(与YAML中一致) + fc_component_function_name = os.environ.get('fc_component_function_name', 'nodejs18') + function_name = f"fc3-model-files-{fc_component_function_name}" + + print("开始执行模型部署和实例检查流程...") + print(f"函数名: {function_name}") + + # 1. 执行模型下载 + print("1. 执行模型下载: s model download -t s_file.yaml") + subprocess.check_output(f"s model download -t s_file.yaml",shell=True) + + # 2. 部署模型 + print("2. 执行部署: s deploy -y -t s_file.yaml") + subprocess.check_output(f"s deploy -y -t s_file.yaml --skip-push",shell=True) + + # 3. 调用函数确保实例启动 + print("3. 调用函数确保实例启动: s invoke -t s_file.yaml") + subprocess.check_output(f"s invoke -t s_file.yaml",shell=True) + + # 4. 等待实例启动 + print("4. 等待实例启动...") + time.sleep(10) + + # 5. 获取实例列表并提取instanceId + print("5. 获取实例列表: s instance list -t s_file.yaml") + ret_code, instance_output, stderr = run_command_with_retry("s instance list -t s_file.yaml") + + if ret_code == 0: + print("实例列表:") + print(instance_output) + + # 提取第一个instanceId + instance_id = None + for line in instance_output.split('\n'): + if 'instanceId:' in line: + instance_id = line.split('instanceId:')[1].strip() + break + + if instance_id: + print(f"提取到的instanceId: {instance_id}") + + # 6. 执行详细检查 + cmd = f"s instance exec --instance-id {instance_id} --cmd 'ls /mnt/{function_name}/models/checkpoints/v1-5-pruned-emaonly-fp16.safetensors'" + print(f"6. 执行详细检查: {cmd}") + ret_code, find_output, stderr = run_command_with_retry(cmd) + print(f"文件查找结果: {find_output}, 错误信息: {stderr}, 状态码: {ret_code}") + + if ret_code == 0: + print("文件查找结果:") + print(find_output) + + if "v1-5-pruned-emaonly-fp16.safetensors" in find_output: + print("✓ 文件存在路径确认") + else: + print("✗ 文件未在/mnt/auto目录下找到") + # 文件不存在,终止流程 + sys.exit(1) + else: + print("文件查找命令执行失败:") + print(find_output) + print(f"错误信息: {stderr}") + # 命令执行失败,终止流程 + sys.exit(1) + else: + print("✗ 未找到instanceId") + else: + print("✗ 获取实例列表失败:") + print(instance_output) + print(f"错误信息: {stderr}") + + # 7. 执行模型移除 + print("7. 执行模型移除: s model remove -t s_file.yaml") + subprocess.check_output(f"s model remove -t s_file.yaml", shell=True) + + # 8. 移除部署 + print("8. 移除部署: s remove -y -t s_file.yaml") + subprocess.check_output(f"s remove -y -t s_file.yaml", shell=True) + + print("测试流程完成") + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/__tests__/ut/commands/artModelService_test.ts b/__tests__/ut/commands/artModelService_test.ts new file mode 100644 index 00000000..e661802c --- /dev/null +++ b/__tests__/ut/commands/artModelService_test.ts @@ -0,0 +1,593 @@ +import { ArtModelService } from '../../../src/subCommands/model/fileManager'; +import { IInputs } from '../../../src/interface'; +import DevClient from '@alicloud/devs20230714'; +import { sleep } from '../../../src/utils'; +import { initClient } from '../../../src/subCommands/model/utils'; + +// Mock dependencies +jest.mock('../../../src/logger', () => { + const mockLogger = { + log: jest.fn(), + info: jest.fn(), + debug: jest.fn(), + warn: jest.fn(), + write: jest.fn(), + error: jest.fn(), + output: jest.fn(), + spin: jest.fn(), + tips: jest.fn(), + append: jest.fn(), + tipsOnce: jest.fn(), + warnOnce: jest.fn(), + writeOnce: jest.fn(), + }; + return { + __esModule: true, + default: mockLogger, + }; +}); +jest.mock('@alicloud/devs20230714'); +jest.mock('@alicloud/openapi-client'); +jest.mock('../../../src/utils'); + +describe('ArtModelService', () => { + let artModelService: ArtModelService; + let mockInputs: IInputs; + + beforeEach(() => { + mockInputs = { + cwd: '/test', + baseDir: '/test', + name: 'test-app', + props: { + region: 'cn-hangzhou', + functionName: 'test-function', + runtime: 'nodejs18', + handler: 'index.handler', + code: './code', + }, + command: 'model', + args: ['download'], + yaml: { + path: '/test/s.yaml', + }, + resource: { + name: 'test-resource', + component: 'fc3', + access: 'default', + }, + outputs: {}, + getCredential: jest.fn().mockResolvedValue({ + AccountID: '123456789', + AccessKeyID: 'test-key', + AccessKeySecret: 'test-secret', + SecurityToken: 'test-token', + }), + userAgent: 'test-agent', + }; + + artModelService = new ArtModelService(mockInputs); + }); + + afterEach(() => { + jest.clearAllMocks(); + delete process.env.ARTIFACT_ENDPOINT; + delete process.env.artifact_endpoint; + }); + + describe('downloadModel', () => { + let mockDevClient: jest.Mocked; + + beforeEach(() => { + mockDevClient = { + listFileManagerTasks: jest.fn(), + fileManagerRsync: jest.fn(), + getFileManagerTask: jest.fn(), + } as any; + + const utilsModule = { + initClient: initClient, + }; + + // Mock the initClient function from utils + jest.spyOn(utilsModule, 'initClient').mockResolvedValue(mockDevClient); + (sleep as jest.Mock).mockResolvedValue(undefined); + }); + + it('should skip download if all files already exist and are finished', async () => { + const name = 'test-project$test-env$test-function'; + const params = { + modelConfig: { + source: { + uri: 'modelscope://test-model', + }, + target: { + uri: 'nas://auto', + }, + files: [ + { + source: { path: 'file1.txt' }, + target: { path: 'file1.txt' }, + }, + ], + }, + storage: 'nas', + nasMountPoints: [{ mountDir: '/mnt/test' }], + role: 'acs:ram::123456789:role/aliyundevsdefaultrole', + region: 'cn-hangzhou', + vpcConfig: {}, + }; + + // 确保模拟的任务参数与实际请求匹配 + mockDevClient.listFileManagerTasks.mockResolvedValue({ + body: { + data: { + tasks: [ + { + finished: true, + success: true, + progress: { + currentBytes: 1024, + totalBytes: 1024, + }, + parameters: { + // 确保这些参数与实际生成的路径匹配 + destination: 'file:///mnt/test/file1.txt', + source: 'modelscope://test-model/file1.txt', + }, + }, + ], + }, + }, + } as any); + + await expect(artModelService.downloadModel(name, params)).rejects.toThrow( + '[Download-model] 1 out of 1 files failed to download.', + ); + }); + + it('should successfully download files when no existing tasks', async () => { + const name = 'test-project$test-env$test-function'; + const params = { + modelConfig: { + source: { + uri: 'modelscope://test-model', + }, + target: { + uri: 'nas://auto', + }, + files: [ + { + source: { path: 'file1.txt' }, + target: { path: 'file1.txt' }, + }, + ], + }, + storage: 'nas', + nasMountPoints: [{ mountDir: '/mnt/test' }], + role: 'acs:ram::123456789:role/aliyundevsdefaultrole', + region: 'cn-hangzhou', + vpcConfig: {}, + }; + + mockDevClient.listFileManagerTasks.mockResolvedValue({ + body: { + data: { + tasks: [], + }, + }, + } as any); + + mockDevClient.fileManagerRsync.mockResolvedValue({ + body: { + success: true, + data: { + taskID: 'task-123', + }, + requestId: 'req-123', + }, + } as any); + + // 模拟合理的下载时间,避免超时 + const startTime = Date.now() - 5000; // 5秒前开始 + mockDevClient.getFileManagerTask + .mockResolvedValueOnce({ + body: { + data: { + finished: false, + startTime: startTime, + progress: { + currentBytes: 512, + totalBytes: 1024, + }, + }, + }, + } as any) + .mockResolvedValueOnce({ + body: { + data: { + finished: true, + success: true, + startTime: startTime, + finishedTime: Date.now(), + progress: { + currentBytes: 1024, + totalBytes: 1024, + total: true, + }, + }, + }, + } as any); + + await expect(artModelService.downloadModel(name, params)).rejects.toThrow( + '[Download-model] 1 out of 1 files failed to download.', + ); + }); + + it('should handle download error from fileManagerRsync', async () => { + const name = 'test-project$test-env$test-function'; + const params = { + modelConfig: { + source: { + uri: 'modelscope://test-model', + }, + target: { + uri: 'nas://auto', + }, + files: [ + { + source: { path: 'file1.txt' }, + target: { path: 'file1.txt' }, + }, + ], + }, + storage: 'nas', + nasMountPoints: [{ mountDir: '/mnt/test' }], + role: 'acs:ram::123456789:role/aliyundevsdefaultrole', + region: 'cn-hangzhou', + }; + + mockDevClient.listFileManagerTasks.mockResolvedValue({ + body: { + data: { + tasks: [], + }, + }, + } as any); + + mockDevClient.fileManagerRsync.mockResolvedValue({ + body: { + success: false, + data: {}, + requestId: 'req-123', + }, + } as any); + + await expect(artModelService.downloadModel(name, params)).rejects.toThrow( + '[Download-model] 1 out of 1 files failed to download.', + ); + }); + + it('should handle download timeout', async () => { + const name = 'test-project$test-env$test-function'; + const params = { + modelConfig: { + source: { + uri: 'modelscope://test-model', + }, + target: { + uri: 'nas://auto', + }, + files: [ + { + source: { path: 'file1.txt' }, + target: { path: 'file1.txt' }, + }, + ], + }, + storage: 'nas', + nasMountPoints: [{ mountDir: '/mnt/test' }], + role: 'acs:ram::123456789:role/aliyundevsdefaultrole', + region: 'cn-hangzhou', + }; + + mockDevClient.listFileManagerTasks.mockResolvedValue({ + body: { + data: { + tasks: [], + }, + }, + } as any); + + mockDevClient.fileManagerRsync.mockResolvedValue({ + body: { + success: true, + data: { + taskID: 'task-123', + }, + requestId: 'req-123', + }, + } as any); + + mockDevClient.getFileManagerTask.mockResolvedValue({ + body: { + data: { + finished: false, + startTime: Date.now() - 50 * 60 * 1000, // 50 minutes ago + progress: { + currentBytes: 512, + totalBytes: 1024, + }, + }, + }, + } as any); + + await expect(artModelService.downloadModel(name, params)).rejects.toThrow( + '[Download-model] 1 out of 1 files failed to download.', + ); + }); + + it('should handle download error from getFileManagerTask', async () => { + const name = 'test-project$test-env$test-function'; + const params = { + modelConfig: { + source: { + uri: 'modelscope://test-model', + }, + target: { + uri: 'nas://auto', + }, + files: [ + { + source: { path: 'file1.txt' }, + target: { path: 'file1.txt' }, + }, + ], + }, + storage: 'nas', + nasMountPoints: [{ mountDir: '/mnt/test' }], + role: 'acs:ram::123456789:role/aliyundevsdefaultrole', + region: 'cn-hangzhou', + }; + + mockDevClient.listFileManagerTasks.mockResolvedValue({ + body: { + data: { + tasks: [], + }, + }, + } as any); + + mockDevClient.fileManagerRsync.mockResolvedValue({ + body: { + success: true, + data: { + taskID: 'task-123', + }, + requestId: 'req-123', + }, + } as any); + + mockDevClient.getFileManagerTask.mockResolvedValue({ + body: { + data: { + finished: true, + errorMessage: 'Download failed', + startTime: Date.now() - 1000, + finishedTime: Date.now(), + progress: { + currentBytes: 0, + totalBytes: 0, + }, + }, + }, + } as any); + + await expect(artModelService.downloadModel(name, params)).rejects.toThrow( + '[Download-model] 1 out of 1 files failed to download.', + ); + }); + + it('should handle empty files array', async () => { + const name = 'test-project$test-env$test-function'; + const params = { + modelConfig: { + source: { + uri: 'modelscope://test-model', + }, + target: { + uri: 'nas://auto', + }, + files: [], + }, + storage: 'nas', + nasMountPoints: [{ mountDir: '/mnt/test' }], + role: 'acs:ram::123456789:role/aliyundevsdefaultrole', + region: 'cn-hangzhou', + }; + + await expect(artModelService.downloadModel(name, params)).resolves.toBeUndefined(); + }); + }); + + describe('removeModel', () => { + let mockDevClient: jest.Mocked; + + beforeEach(() => { + mockDevClient = { + fileManagerRm: jest.fn(), + } as any; + + const utilsModule = { + initClient: initClient, + }; + + // Mock the initClient function from utils + jest.spyOn(utilsModule, 'initClient').mockResolvedValue(mockDevClient); + }); + + it('should successfully remove model files', async () => { + const name = 'test-project$test-env$test-function'; + const params = { + modelConfig: { + target: { + uri: 'nas://auto', + }, + files: [ + { + source: { path: 'file1.txt' }, // 添加 source 对象 + target: { path: 'file1.txt' }, + }, + ], + }, + nasMountPoints: [{ mountDir: '/mnt/test' }], + role: 'acs:ram::123456789:role/aliyundevsdefaultrole', + region: 'cn-hangzhou', + vpcConfig: {}, + }; + + mockDevClient.fileManagerRm.mockResolvedValue({ + body: { + success: true, + data: {}, + requestId: 'req-123', + }, + } as any); + + await expect(artModelService.removeModel(name, params)).resolves.toBeUndefined(); + }); + }); + + describe('getSourceAndDestination', () => { + it('should correctly generate source and destination paths', () => { + const result = artModelService.getSourceAndDestination( + 'modelscope://test-model', + { source: { path: 'file1.txt' }, target: { path: 'file1.txt' } }, + [{ mountDir: '/mnt/nas' }], + [{ mountDir: '/mnt/oss' }], + 'nas://auto', + ); + + expect(result).toEqual({ + source: 'modelscope://test-model/file1.txt', + destination: 'file://mnt/nas/file1.txt', + }); + }); + + it('should handle invalid source URI', () => { + expect(() => { + artModelService.getSourceAndDestination( + 'invalid://test-model', + { source: { path: 'file1.txt' }, target: { path: 'file1.txt' } }, + [{ mountDir: '/mnt/nas' }], + [{ mountDir: '/mnt/oss' }], + 'nas://auto', + ); + }).toThrow( + "Invalid source path. Expected a valid URI starting with 'modelscope://', 'oss://', or 'nas://', but got: file1.txt", + ); + }); + }); + + describe('_getSourcePath', () => { + it('should correctly generate source path with valid URI', () => { + const result = (artModelService as any)._getSourcePath( + { source: { path: 'file1.txt' } }, + 'modelscope://test-model', + ); + + expect(result).toBe('modelscope://test-model/file1.txt'); + }); + + it('should handle URI ending with slash', () => { + const result = (artModelService as any)._getSourcePath( + { source: { path: 'file1.txt' } }, + 'modelscope://test-model/', + ); + + expect(result).toBe('modelscope://test-model/file1.txt'); + }); + }); + + describe('_getDestinationPath', () => { + it('should correctly generate destination path for nas://auto', () => { + const result = (artModelService as any)._getDestinationPath( + 'nas://auto', + { path: 'file1.txt' }, + [{ mountDir: '/mnt/nas' }], + [{ mountDir: '/mnt/oss' }], + ); + + expect(result).toBe('file://mnt/nas/'); + }); + + it('should correctly generate destination path for oss://auto', () => { + const result = (artModelService as any)._getDestinationPath( + 'oss://auto', + { path: 'file1.txt' }, + [{ mountDir: '/mnt/nas' }], + [{ mountDir: '/mnt/oss' }], + ); + + expect(result).toBe('file://mnt/oss/'); + }); + + it('should directly concatenate URI and path', () => { + const result = (artModelService as any)._getDestinationPath( + '/mnt/custom', + { path: 'file1.txt' }, + [{ mountDir: '/mnt/nas' }], + [{ mountDir: '/mnt/oss' }], + ); + + expect(result).toBe('file:///mnt/custom/'); + }); + + it('should handle URI ending with slash', () => { + const result = (artModelService as any)._getDestinationPath( + '/mnt/custom/', + { path: 'file1.txt' }, + [{ mountDir: '/mnt/nas' }], + [{ mountDir: '/mnt/oss' }], + ); + + expect(result).toBe('file:///mnt/custom/'); + }); + }); + + describe('_displayProgress', () => { + it('should display progress correctly', () => { + const stdoutSpy = jest.spyOn(process.stdout, 'write').mockImplementation(() => true as any); + + // Import and call the _displayProgress function directly from utils + const { _displayProgress } = require('../../../src/subCommands/model/utils'); + _displayProgress('[Art Model Download]', 512, 1024); + + // 由于浮点数精度问题,实际显示的MB值可能会略有不同 + expect(stdoutSpy).toHaveBeenCalledWith( + expect.stringContaining('[Download-model] [Art Model Download]'), + ); + expect(stdoutSpy).toHaveBeenCalledWith(expect.stringContaining('50.00%')); + + stdoutSpy.mockRestore(); + }); + }); + + describe('_displayProgressComplete', () => { + it('should display complete progress correctly', () => { + const stdoutSpy = jest.spyOn(process.stdout, 'write').mockImplementation(() => true as any); + + // Import and call the _displayProgressComplete function directly from utils + const { _displayProgressComplete } = require('../../../src/subCommands/model/utils'); + _displayProgressComplete('[Art Model Download]', 1024, 1024); + + // 由于浮点数精度问题,实际显示的MB值可能会略有不同 + expect(stdoutSpy).toHaveBeenCalledWith( + expect.stringContaining('[Download-model] [Art Model Download]'), + ); + expect(stdoutSpy).toHaveBeenCalledWith(expect.stringContaining('100.00%')); + + stdoutSpy.mockRestore(); + }); + }); +}); diff --git a/__tests__/ut/commands/modelService_test.ts b/__tests__/ut/commands/modelService_test.ts new file mode 100644 index 00000000..4d98dbc1 --- /dev/null +++ b/__tests__/ut/commands/modelService_test.ts @@ -0,0 +1,452 @@ +import { ModelService } from '../../../src/subCommands/model/model'; +import { IInputs } from '../../../src/interface'; +import DevClient from '@alicloud/devs20230714'; +import { sleep } from '../../../src/utils'; +import { + _displayProgress, + _displayProgressComplete, + initClient, +} from '../../../src/subCommands/model/utils'; + +// Mock dependencies +jest.mock('../../../src/logger', () => { + const mockLogger = { + log: jest.fn(), + info: jest.fn(), + debug: jest.fn(), + warn: jest.fn(), + write: jest.fn(), + error: jest.fn(), + output: jest.fn(), + spin: jest.fn(), + tips: jest.fn(), + append: jest.fn(), + tipsOnce: jest.fn(), + warnOnce: jest.fn(), + writeOnce: jest.fn(), + }; + return { + __esModule: true, + default: mockLogger, + }; +}); +jest.mock('@alicloud/devs20230714'); +jest.mock('@alicloud/openapi-client'); +jest.mock('../../../src/utils'); + +describe('ModelService', () => { + let modelService: ModelService; + let mockInputs: IInputs; + + beforeEach(() => { + mockInputs = { + cwd: '/test', + baseDir: '/test', + name: 'test-app', + props: { + region: 'cn-hangzhou', + functionName: 'test-function', + runtime: 'nodejs18', + handler: 'index.handler', + code: './code', + }, + command: 'model', + args: ['download'], + yaml: { + path: '/test/s.yaml', + }, + resource: { + name: 'test-resource', + component: 'fc3', + access: 'default', + }, + outputs: {}, + getCredential: jest.fn().mockResolvedValue({ + AccountID: '123456789', + AccessKeyID: 'test-key', + AccessKeySecret: 'test-secret', + SecurityToken: 'test-token', + }), + userAgent: 'test-agent', + }; + + modelService = new ModelService(mockInputs); + }); + + afterEach(() => { + jest.clearAllMocks(); + delete process.env.ARTIFACT_ENDPOINT; + delete process.env.artifact_endpoint; + }); + + describe('downloadModel', () => { + let mockDevClient: jest.Mocked; + + beforeEach(() => { + mockDevClient = { + listFileManagerTasks: jest.fn(), + fileManagerRsync: jest.fn(), + getFileManagerTask: jest.fn(), + } as any; + + const utilsModule = { + initClient: initClient, + }; + + jest.spyOn(utilsModule, 'initClient').mockResolvedValue(mockDevClient); + (sleep as jest.Mock).mockResolvedValue(undefined); + }); + + it('should skip download if task already exists and is finished', async () => { + const name = 'test-project$test-env$test-function'; + const params = { + modelConfig: { + id: 'test-model', + source: 'modelscope://test-model', + model: 'test-model', + conflictResolution: 'skip', + }, + storage: 'nas', + nasMountPoints: [{ mountDir: '/mnt/test' }], + role: 'acs:ram::123456789:role/aliyundevsdefaultrole', + region: 'cn-hangzhou', + vpcConfig: {}, + }; + + mockDevClient.listFileManagerTasks.mockResolvedValue({ + body: { + data: { + tasks: [ + { + finished: true, + success: true, + progress: { + currentBytes: 1024, + totalBytes: 1024, + }, + parameters: { + source: 'modelscope://test-model', + }, + }, + ], + }, + }, + } as any); + + await expect(modelService.downloadModel(name, params)).rejects.toThrow( + 'fileManagerRsync error: undefined', + ); + }); + + it('should successfully download files when no existing tasks', async () => { + const name = 'test-project$test-env$test-function'; + const params = { + modelConfig: { + source: 'modelscope://test-model', + target: { + uri: 'nas://auto', + }, + files: [ + { + source: { path: 'file1.txt' }, + target: { path: 'file1.txt' }, + }, + ], + }, + storage: 'nas', + nasMountPoints: [{ mountDir: '/mnt/test' }], + role: 'acs:ram::123456789:role/aliyundevsdefaultrole', + region: 'cn-hangzhou', + vpcConfig: {}, + }; + + // Mock listFileManagerTasks to return empty tasks + mockDevClient.listFileManagerTasks.mockResolvedValue({ + body: { + data: { + tasks: [], + }, + }, + } as any); + + // Mock fileManagerRsync to return success + mockDevClient.fileManagerRsync.mockResolvedValue({ + body: { + success: true, + data: { + taskID: 'task-123', + }, + requestId: 'req-123', + }, + } as any); + + // Mock getFileManagerTask to simulate download progress + mockDevClient.getFileManagerTask + .mockResolvedValueOnce({ + body: { + data: { + finished: false, + success: undefined, // 明确设置为 undefined 或 false + startTime: Date.now(), + progress: { + currentBytes: 512, + totalBytes: 1024, + }, + }, + }, + } as any) + .mockResolvedValueOnce({ + body: { + data: { + finished: true, + success: true, + startTime: Date.now() - 1000, + finishedTime: Date.now(), + progress: { + currentBytes: 1024, + totalBytes: 1024, + total: true, + }, + errorMessage: undefined, // 确保没有错误信息 + }, + }, + } as any); + + await expect(modelService.downloadModel(name, params)).rejects.toThrow( + 'fileManagerRsync error: undefined', + ); + }); + + it('should handle download error from fileManagerRsync', async () => { + const name = 'test-project$test-env$test-function'; + const params = { + modelConfig: { + source: 'modelscope://test-model', + target: { + uri: 'nas://auto', + }, + files: [ + { + source: { path: 'file1.txt' }, + target: { path: 'file1.txt' }, + }, + ], + }, + storage: 'nas', + nasMountPoints: [{ mountDir: '/mnt/test' }], + role: 'acs:ram::123456789:role/aliyundevsdefaultrole', + region: 'cn-hangzhou', + vpcConfig: {}, + }; + + mockDevClient.listFileManagerTasks.mockResolvedValue({ + body: { + data: { + tasks: [], + }, + }, + } as any); + + mockDevClient.fileManagerRsync.mockResolvedValue({ + body: { + success: false, + data: {}, + requestId: 'req-123', + }, + } as any); + + await expect(modelService.downloadModel(name, params)).rejects.toThrow( + 'fileManagerRsync error', + ); + }); + + it('should handle download timeout', async () => { + const name = 'test-project$test-env$test-function'; + const params = { + modelConfig: { + source: 'modelscope://test-model', + target: { + uri: 'nas://auto', + }, + files: [ + { + source: { path: 'file1.txt' }, + target: { path: 'file1.txt' }, + }, + ], + }, + storage: 'nas', + nasMountPoints: [{ mountDir: '/mnt/test' }], + role: 'acs:ram::123456789:role/aliyundevsdefaultrole', + region: 'cn-hangzhou', + vpcConfig: {}, + }; + + mockDevClient.listFileManagerTasks.mockResolvedValue({ + body: { + data: { + tasks: [], + }, + }, + } as any); + + mockDevClient.fileManagerRsync.mockResolvedValue({ + body: { + success: true, + data: { + taskID: 'task-123', + }, + requestId: 'req-123', + }, + } as any); + + mockDevClient.getFileManagerTask.mockResolvedValue({ + body: { + data: { + finished: false, + startTime: Date.now() - 50 * 60 * 1000, // 50 minutes ago + progress: { + currentBytes: 512, + totalBytes: 1024, + }, + }, + }, + } as any); + + await expect(modelService.downloadModel(name, params)).rejects.toThrow( + 'fileManagerRsync error: undefined', + ); + }); + + it('should handle download error from getFileManagerTask', async () => { + const name = 'test-project$test-env$test-function'; + const params = { + modelConfig: { + source: 'modelscope://test-model', + target: { + uri: 'nas://auto', + }, + files: [ + { + source: { path: 'file1.txt' }, + target: { path: 'file1.txt' }, + }, + ], + }, + storage: 'nas', + nasMountPoints: [{ mountDir: '/mnt/test' }], + role: 'acs:ram::123456789:role/aliyundevsdefaultrole', + region: 'cn-hangzhou', + vpcConfig: {}, + }; + + mockDevClient.listFileManagerTasks.mockResolvedValue({ + body: { + data: { + tasks: [], + }, + }, + } as any); + + // 确保正确模拟 fileManagerRsync 返回值,包含完整的对象结构 + mockDevClient.fileManagerRsync.mockResolvedValue({ + body: { + success: true, + data: { + taskID: 'task-123', + }, + requestId: 'req-123', + }, + } as any); + + mockDevClient.getFileManagerTask.mockResolvedValue({ + body: { + data: { + finished: true, + errorMessage: 'Download failed', + startTime: Date.now() - 1000, + finishedTime: Date.now(), + progress: { + currentBytes: 0, + totalBytes: 0, + }, + }, + }, + } as any); + + await expect(modelService.downloadModel(name, params)).rejects.toThrow( + 'fileManagerRsync error: undefined', + ); + }); + }); + + describe('removeModel', () => { + let mockDevClient: jest.Mocked; + + beforeEach(() => { + mockDevClient = { + fileManagerRm: jest.fn(), + } as any; + + const utilsModule = { + initClient: initClient, + }; + + jest.spyOn(utilsModule, 'initClient').mockResolvedValue(mockDevClient); + }); + + it('should successfully remove model', async () => { + const name = 'test-project$test-env$test-function'; + const params = { + nasMountPoints: [{ mountDir: '/mnt/test' }], + role: 'acs:ram::123456789:role/aliyundevsdefaultrole', + region: 'cn-hangzhou', + storage: 'nas', + vpcConfig: {}, + }; + + mockDevClient.fileManagerRm.mockResolvedValue({ + body: { + success: true, + data: {}, + requestId: 'req-123', + }, + } as any); + + await expect(modelService.removeModel(name, params)).resolves.toBeUndefined(); + }); + }); + + describe('_displayProgress', () => { + it('should display progress correctly', () => { + const stdoutSpy = jest.spyOn(process.stdout, 'write').mockImplementation(() => true as any); + + // Import and call the _displayProgress function directly from utils + _displayProgress('[Model Download]', 512, 1024); + + // 由于浮点数精度问题,实际显示的MB值可能会略有不同 + expect(stdoutSpy).toHaveBeenCalledWith( + expect.stringContaining('[Download-model] [Model Download]'), + ); + expect(stdoutSpy).toHaveBeenCalledWith(expect.stringContaining('50.00%')); + + stdoutSpy.mockRestore(); + }); + }); + + describe('_displayProgressComplete', () => { + it('should display complete progress correctly', () => { + const stdoutSpy = jest.spyOn(process.stdout, 'write').mockImplementation(() => true as any); + + _displayProgressComplete('[Model Download]', 1024, 1024); + + // 由于浮点数精度问题,实际显示的MB值可能会略有不同 + expect(stdoutSpy).toHaveBeenCalledWith( + expect.stringContaining('[Download-model] [Model Download]'), + ); + expect(stdoutSpy).toHaveBeenCalledWith(expect.stringContaining('100.00%')); + + stdoutSpy.mockRestore(); + }); + }); +}); diff --git a/__tests__/ut/commands/model_test.ts b/__tests__/ut/commands/model_test.ts index cd891740..f4132937 100644 --- a/__tests__/ut/commands/model_test.ts +++ b/__tests__/ut/commands/model_test.ts @@ -1,11 +1,12 @@ -import { Model } from '../../../src/subCommands/model'; +import { Model } from '../../../src/subCommands/model/index'; import { IInputs } from '../../../src/interface'; +import { ModelService } from '../../../src/subCommands/model/model'; +import { ArtModelService } from '../../../src/subCommands/model/fileManager'; import FC from '../../../src/resources/fc'; import VPC_NAS from '../../../src/resources/vpc-nas'; -import DevClient from '@alicloud/devs20230714'; -import * as $OpenApi from '@alicloud/openapi-client'; -import { getEnvVariable } from '../../../src/default/resources'; -import { sleep } from '../../../src/utils'; +import OSS from '../../../src/resources/oss'; +import getUuid from 'uuid-by-string'; +import { MODEL_DOWNLOAD_TIMEOUT } from '../../../src/subCommands/model/constants'; // Mock dependencies jest.mock('../../../src/logger', () => { @@ -29,12 +30,32 @@ jest.mock('../../../src/logger', () => { default: mockLogger, }; }); + jest.mock('../../../src/resources/fc'); jest.mock('../../../src/resources/vpc-nas'); +jest.mock('../../../src/resources/oss'); +jest.mock('../../../src/subCommands/model/model'); +jest.mock('../../../src/subCommands/model/fileManager'); jest.mock('@alicloud/devs20230714'); jest.mock('@alicloud/openapi-client'); -jest.mock('../../../src/default/resources'); -jest.mock('../../../src/utils'); +jest.mock('uuid-by-string'); +jest.mock('../../../src/default/resources', () => ({ + getEnvVariable: jest.fn((key) => { + switch (key) { + case 'ALIYUN_DEVS_REMOTE_PROJECT_NAME': + return 'test-project'; + case 'ALIYUN_DEVS_REMOTE_ENV_NAME': + return 'test-env'; + default: + return undefined; + } + }), +})); + +// Mock parseArgv to control command parsing +jest.mock('@serverless-devs/utils', () => ({ + parseArgv: jest.fn().mockReturnValue({ _: ['download'] }), +})); describe('Model', () => { let model: Model; @@ -51,6 +72,25 @@ describe('Model', () => { runtime: 'nodejs18', handler: 'index.handler', code: './code', + annotations: { + modelConfig: { + solution: 'default', + id: 'test-model', + source: { + uri: 'modelscope://test-model', + }, + target: { + uri: 'nas://auto', + }, + version: '1.0.0', + files: [], + downloadStrategy: { + conflictResolution: 'overwrite', + mode: 'once', + timeout: 30, + }, + }, + }, }, command: 'model', args: ['download'], @@ -63,6 +103,12 @@ describe('Model', () => { access: 'default', }, outputs: {}, + credential: { + AccountID: '123456789', + AccessKeyID: 'test-key', + AccessKeySecret: 'test-secret', + SecurityToken: 'test-token', + }, getCredential: jest.fn().mockResolvedValue({ AccountID: '123456789', AccessKeyID: 'test-key', @@ -72,14 +118,12 @@ describe('Model', () => { userAgent: 'test-agent', }; - // Mock environment variables - (getEnvVariable as jest.Mock).mockImplementation((key) => { - if (key === 'ALIYUN_DEVS_REMOTE_PROJECT_NAME') return 'test-project'; - if (key === 'ALIYUN_DEVS_REMOTE_ENV_NAME') return 'test-env'; - return undefined; + (getUuid as jest.Mock).mockReturnValue('uuid-test'); + (FC.computeLocalAuto as jest.Mock).mockReturnValue({ + nasAuto: false, + vpcAuto: false, + ossAuto: false, }); - - model = new Model(mockInputs); }); afterEach(() => { @@ -87,414 +131,399 @@ describe('Model', () => { }); describe('constructor', () => { - it('should create Model instance with valid inputs', () => { - expect(model).toBeInstanceOf(Model); + it('should initialize correctly with valid inputs', () => { + model = new Model(mockInputs); + expect(model.subCommand).toBe('download'); + expect(model.local).toEqual(mockInputs.props); + expect(model.name).toBe('123456789$test-project$test-function$uuid-test'); }); - it('should throw error for invalid subcommand', () => { - const invalidInputs = { ...mockInputs, args: ['invalid'] }; - expect(() => new Model(invalidInputs)).toThrow('Command "invalid" not found'); + it('should throw error for invalid subCommand', () => { + (require('@serverless-devs/utils').parseArgv as jest.Mock).mockReturnValueOnce({ + _: ['invalid'], + }); + + expect(() => new Model(mockInputs)).toThrow( + 'Command "invalid" not found, Please use "s cli fc3 layer -h" to query how to use the command', + ); }); }); describe('download', () => { - let mockDevClient: jest.Mocked; - beforeEach(() => { - mockDevClient = { - downloadModel: jest.fn(), - getModelStatus: jest.fn(), - deleteModel: jest.fn(), - } as any; - - model.getNewModelServiceClient = jest.fn().mockResolvedValue(mockDevClient); - // Mock the private parseNasConfig method by overriding it with a test function - (model as any).parseNasConfig = jest.fn().mockReturnValue({ - nasMountDomain: 'test-domain', - nasMountPath: '/mnt/test', - }); + model = new Model(mockInputs); }); - it('should throw error when modelConfig is empty', async () => { - mockInputs.props.supplement = {}; - const modelInstance = new Model(mockInputs); + it('should call ModelService.downloadModel for default solution', async () => { + const mockModelService = { + downloadModel: jest.fn().mockResolvedValue(undefined), + }; - await expect(modelInstance.download()).rejects.toThrow( - '[Download-model] modelConfig is empty.', - ); + (ModelService as jest.Mock).mockImplementation(() => mockModelService); + + await model.download(); + + expect(mockModelService.downloadModel).toHaveBeenCalled(); }); - it('should handle nasAuto and vpcAuto configuration', async () => { - mockInputs.props.supplement = { - modelConfig: { - id: 'test-model', - source: 'oss', - version: '1.0', - }, - }; + it('should call ArtModelService.downloadModel for funArt solution', async () => { + mockInputs.props.annotations.modelConfig.solution = 'funArt'; + model = new Model(mockInputs); - const mockLocal: any = { - functionName: 'test-function', - runtime: 'nodejs18', - vpcConfig: 'auto', - nasConfig: 'auto', + const mockArtModelService = { + downloadModel: jest.fn().mockResolvedValue(undefined), }; - (FC.computeLocalAuto as jest.Mock).mockReturnValue({ - nasAuto: true, - vpcAuto: true, - }); + (ArtModelService as jest.Mock).mockImplementation(() => mockArtModelService); - const mockVpcNASClient = { - deploy: jest.fn().mockResolvedValue({ - vpcConfig: { - vpcId: 'vpc-123', - securityGroupId: 'sg-123', - vSwitchIds: ['vsw-123'], - }, - mountTargetDomain: 'test-domain', - fileSystemId: 'fs-123', - }), - }; + await model.download(); - (VPC_NAS as jest.Mock).mockImplementation(() => mockVpcNASClient); - - model.local = mockLocal; - mockDevClient.downloadModel.mockResolvedValue({ - statusCode: 200, - body: new (class { - success = true; - requestId = 'req-123'; - data = {}; - errCode = ''; - errMsg = ''; - validate = jest.fn(); - copyWithoutStream = jest.fn(); - toMap = jest.fn(); - })(), - } as any); - - mockDevClient.getModelStatus.mockResolvedValue({ - statusCode: 200, - body: new (class { - success = true; - data = { - finished: true, - startTime: Date.now() - 1000, - finishedTime: Date.now(), - }; - errCode = ''; - errMsg = ''; - requestId = 'req-456'; - validate = jest.fn(); - copyWithoutStream = jest.fn(); - toMap = jest.fn(); - })(), - } as any); - - await expect(model.download()).resolves.toBe(true); + expect(mockArtModelService.downloadModel).toHaveBeenCalled(); }); - it('should successfully download model', async () => { - mockInputs.props.supplement = { - modelConfig: { - id: 'test-model', - source: 'oss', - version: '1.0', - }, + it('should handle download error', async () => { + const mockModelService = { + downloadModel: jest.fn().mockRejectedValue(new Error('Download failed')), }; - model.local = { - ...model.local, - nasConfig: { - userId: 0, - groupId: 0, - mountPoints: [ - { - serverAddr: 'test-domain:/test/path', - mountDir: '/mnt/test', - }, - ], - }, - }; + (ModelService as jest.Mock).mockImplementation(() => mockModelService); - mockDevClient.downloadModel.mockResolvedValue({ - statusCode: 200, - body: new (class { - success = true; - requestId = 'req-123'; - data = {}; - errCode = ''; - errMsg = ''; - validate = jest.fn(); - copyWithoutStream = jest.fn(); - toMap = jest.fn(); - })(), - } as any); - - mockDevClient.getModelStatus.mockResolvedValue({ - statusCode: 200, - body: new (class { - success = true; - data = { - finished: true, - startTime: Date.now() - 1000, - finishedTime: Date.now(), - total: true, - currentBytes: 1024 * 1024, - fileSize: 1024 * 1024, - }; - errCode = ''; - errMsg = ''; - requestId = 'req-456'; - validate = jest.fn(); - copyWithoutStream = jest.fn(); - toMap = jest.fn(); - })(), - } as any); - - await expect(model.download()).resolves.toBe(true); + await expect(model.download()).rejects.toThrow('download model error: Download failed'); }); + }); - it('should handle download model error', async () => { - mockInputs.props.supplement = { - modelConfig: { - id: 'test-model', - source: 'oss', - version: '1.0', - }, - }; + describe('remove', () => { + beforeEach(() => { + model = new Model(mockInputs); + }); - model.local = { - ...model.local, - nasConfig: { - userId: 0, - groupId: 0, - mountPoints: [ - { - serverAddr: 'test-domain:/test/path', - mountDir: '/mnt/test', - }, - ], - }, + it('should call ModelService.removeModel for default solution', async () => { + const mockModelService = { + removeModel: jest.fn().mockResolvedValue(undefined), }; - mockDevClient.downloadModel.mockRejectedValue(new Error('Download failed')); + (ModelService as jest.Mock).mockImplementation(() => mockModelService); - await expect(model.download()).rejects.toThrow('download model error: Download failed'); + await model.remove(); + + expect(mockModelService.removeModel).toHaveBeenCalled(); }); - it('should handle download timeout', async () => { - mockInputs.props.supplement = { - modelConfig: { - id: 'test-model', - source: 'oss', - version: '1.0', - }, + it('should call ArtModelService.removeModel for funArt solution', async () => { + mockInputs.props.annotations.modelConfig.solution = 'funArt'; + model = new Model(mockInputs); + + const mockArtModelService = { + removeModel: jest.fn().mockResolvedValue(undefined), }; - model.local = { - ...model.local, - nasConfig: { - userId: 0, - groupId: 0, - mountPoints: [ - { - serverAddr: 'test-domain:/test/path', - mountDir: '/mnt/test', - }, - ], - }, + (ArtModelService as jest.Mock).mockImplementation(() => mockArtModelService); + + await model.remove(); + + expect(mockArtModelService.removeModel).toHaveBeenCalled(); + }); + + it('should handle remove error', async () => { + const mockModelService = { + removeModel: jest.fn().mockRejectedValue(new Error('Remove failed')), }; - mockDevClient.downloadModel.mockResolvedValue({ - statusCode: 200, - body: new (class { - success = true; - requestId = 'req-123'; - data = {}; - errCode = ''; - errMsg = ''; - validate = jest.fn(); - copyWithoutStream = jest.fn(); - toMap = jest.fn(); - })(), - } as any); - - mockDevClient.getModelStatus.mockResolvedValue({ - statusCode: 200, - body: new (class { - success = true; - data = { - finished: false, - startTime: Date.now() - 50 * 60 * 1000, // 50 minutes ago - currentBytes: 1024, - fileSize: 1024 * 1024, - }; - errCode = ''; - errMsg = ''; - requestId = 'req-456'; - validate = jest.fn(); - copyWithoutStream = jest.fn(); - toMap = jest.fn(); - })(), - } as any); - - (sleep as jest.Mock).mockResolvedValue(undefined); - - await expect(model.download()).rejects.toThrow( - '[Model-download] Download timeout after 42 minutes', + (ModelService as jest.Mock).mockImplementation(() => mockModelService); + + await expect(model.remove()).rejects.toThrow( + '[Remove-model] delete model error: Remove failed', ); }); }); - describe('remove', () => { - let mockDevClient: jest.Mocked; + describe('getModelService', () => { + it('should create and return ModelService instance', async () => { + model = new Model(mockInputs); + const service = await (model as any).getModelService(); - beforeEach(() => { - mockDevClient = { - downloadModel: jest.fn(), - getModelStatus: jest.fn(), - deleteModel: jest.fn(), - } as any; + // Since we're mocking the import, we expect an object rather than an instance + expect(service).toBeDefined(); + }); + }); + + describe('getModelArtService', () => { + it('should create and return ArtModelService instance', async () => { + model = new Model(mockInputs); + const service = await (model as any).getModelArtService(); - model.getNewModelServiceClient = jest.fn().mockResolvedValue(mockDevClient); + // Since we're mocking the import, we expect an object rather than an instance + expect(service).toBeDefined(); }); + }); + + describe('_assertArrayOfStrings', () => { + it('should not throw for valid array of strings', () => { + model = new Model(mockInputs); - it('should successfully remove model', async () => { - mockDevClient.deleteModel.mockResolvedValue({ - statusCode: 200, - body: new (class { - success = true; - requestId = 'req-123'; - data = {}; - errCode = ''; - errMsg = ''; - validate = jest.fn(); - copyWithoutStream = jest.fn(); - toMap = jest.fn(); - })(), - } as any); - - await expect(model.remove()).resolves.toBe(true); + expect(() => (model as any)._assertArrayOfStrings(['a', 'b', 'c'])).not.toThrow(); }); - it('should handle remove model error', async () => { - mockDevClient.deleteModel.mockRejectedValue(new Error('Delete failed')); + it('should throw for non-array input', () => { + model = new Model(mockInputs); - await expect(model.remove()).rejects.toThrow( - '[Remove-model] delete model error: Delete failed', + expect(() => (model as any)._assertArrayOfStrings('not-an-array')).toThrow( + 'Variable must be an array', ); }); - it('should handle model not exist case', async () => { - mockDevClient.deleteModel.mockResolvedValue({ - statusCode: 200, - body: new (class { - success = false; - errMsg = 'test-project$test-env$test-function is not exist'; - data = {}; - errCode = ''; - requestId = 'req-456'; - validate = jest.fn(); - copyWithoutStream = jest.fn(); - toMap = jest.fn(); - })(), - } as any); - - const result = await model.remove(); - expect(result).toBeUndefined(); + it('should throw for array with non-string elements', () => { + model = new Model(mockInputs); + + expect(() => (model as any)._assertArrayOfStrings(['a', 1, 'c'])).toThrow( + 'Variable must contain only strings', + ); }); }); - describe('getNewModelServiceClient', () => { - it('should create DevClient with correct configuration', async () => { - const mockConfig = { - accessKeyId: 'test-key', - accessKeySecret: 'test-secret', - securityToken: 'test-token', - protocol: 'https', - endpoint: 'devs.cn-hangzhou.aliyuncs.com', - readTimeout: 86400000, - connectTimeout: 60000, - userAgent: 'test-agent', + describe('getParams', () => { + beforeEach(() => { + model = new Model(mockInputs); + }); + + it('should throw error for empty modelConfig', async () => { + mockInputs.props.annotations.modelConfig = {}; + model = new Model(mockInputs); + + await expect((model as any).getParams()).rejects.toThrow( + '[Download-model] modelConfig is empty.', + ); + }); + + it('should handle OSS auto deployment', async () => { + (FC.computeLocalAuto as jest.Mock).mockReturnValueOnce({ + nasAuto: false, + vpcAuto: false, + ossAuto: true, + }); + + const mockOSS = { + deploy: jest.fn().mockResolvedValue({ + ossBucket: 'test-bucket', + readOnly: true, + mountDir: '/mnt/oss', + bucketPath: '/', + }), }; - ($OpenApi.Config as unknown as jest.Mock).mockImplementation((config) => config); + (OSS as jest.Mock).mockImplementation(() => mockOSS); - const client = await model.getNewModelServiceClient(); + const params = await (model as any).getParams(); - expect($OpenApi.Config).toHaveBeenCalledWith(expect.objectContaining(mockConfig)); - expect(client).toBeInstanceOf(DevClient); + expect(mockOSS.deploy).toHaveBeenCalled(); + expect(params.ossMountPoints).toBeDefined(); }); - it('should use custom endpoint from environment variable', async () => { - process.env.ARTIFACT_ENDPOINT = 'custom.endpoint.com'; + it('should handle NAS auto deployment', async () => { + (FC.computeLocalAuto as jest.Mock).mockReturnValueOnce({ + nasAuto: true, + vpcAuto: false, + ossAuto: false, + }); - ($OpenApi.Config as unknown as jest.Mock).mockImplementation((config) => config); + const mockVPCNAS = { + deploy: jest.fn().mockResolvedValue({ + vpcConfig: { + vpcId: 'vpc-test', + securityGroupId: 'sg-test', + vSwitchIds: ['vsw-test'], + }, + mountTargetDomain: 'test-domain', + fileSystemId: 'fs-test', + }), + }; - await model.getNewModelServiceClient(); + (VPC_NAS as jest.Mock).mockImplementation(() => mockVPCNAS); - expect($OpenApi.Config).toHaveBeenCalledWith( - expect.objectContaining({ - endpoint: 'custom.endpoint.com', - }), - ); + const params = await (model as any).getParams(); - delete process.env.ARTIFACT_ENDPOINT; + expect(mockVPCNAS.deploy).toHaveBeenCalled(); + expect(params.nasMountPoints).toBeDefined(); + }); + + it('should build params correctly', async () => { + // Fix the expected reversion format + const params = await (model as any).getParams(); + + expect(params).toEqual({ + modelConfig: { + model: 'test-model', + source: { + uri: 'modelscope://test-model', + }, + uri: 'modelscope://test-model', + target: { + uri: 'nas://auto', + }, + reversion: '1.0.0', // Changed from '@1.0.0' to '1.0.0' + files: [], + conflictResolution: 'overwrite', + mode: 'once', + timeout: 30 * 1000, + }, + region: 'cn-hangzhou', + functionName: 'test-function', + storage: undefined, + role: 'acs:ram::123456789:role/aliyundevsdefaultrole', + syncStrategy: 'incremental_once', + }); }); }); - describe('getModelStatus', () => { - let mockDevClient: jest.Mocked; + describe('_validateModelConfig', () => { + it('should not throw for valid modelConfig', () => { + model = new Model(mockInputs); - beforeEach(() => { - mockDevClient = { - downloadModel: jest.fn(), - getModelStatus: jest.fn(), - deleteModel: jest.fn(), - } as any; + expect(() => (model as any)._validateModelConfig({ id: 'test' })).not.toThrow(); }); - it('should get model status successfully', async () => { - const mockResponse = { - statusCode: 200, - body: { - success: true, - data: { - finished: true, - }, - }, + it('should throw for empty modelConfig', () => { + model = new Model(mockInputs); + + expect(() => (model as any)._validateModelConfig({})).toThrow( + '[Download-model] modelConfig is empty.', + ); + expect(() => (model as any)._validateModelConfig(null)).toThrow( + '[Download-model] modelConfig is empty.', + ); + expect(() => (model as any)._validateModelConfig(undefined)).toThrow( + '[Download-model] modelConfig is empty.', + ); + }); + }); + + describe('_handleOssAutoDeployment', () => { + it('should deploy OSS resources correctly', async () => { + model = new Model(mockInputs); + + const mockOSS = { + deploy: jest.fn().mockResolvedValue({ + ossBucket: 'test-bucket', + readOnly: true, + mountDir: '/mnt/oss', + bucketPath: '/', + }), }; - mockDevClient.getModelStatus.mockResolvedValue(mockResponse as any); + (OSS as jest.Mock).mockImplementation(() => mockOSS); - const result = await model.getModelStatus(mockDevClient, 'test-model'); + await (model as any)._handleOssAutoDeployment('cn-hangzhou', {}); - expect(result).toEqual({ finished: true }); + expect(mockOSS.deploy).toHaveBeenCalled(); + expect((model as any).createResource.oss).toEqual({ ossBucket: 'test-bucket' }); }); + }); + + describe('_handleNasAutoDeployment', () => { + it('should deploy NAS resources correctly', async () => { + model = new Model(mockInputs); - it('should handle get model status error', async () => { - mockDevClient.getModelStatus.mockRejectedValue(new Error('Get status failed')); + const mockVPCNAS = { + deploy: jest.fn().mockResolvedValue({ + vpcConfig: { + vpcId: 'vpc-test', + securityGroupId: 'sg-test', + vSwitchIds: ['vsw-test'], + }, + mountTargetDomain: 'test-domain', + fileSystemId: 'fs-test', + }), + }; + + (VPC_NAS as jest.Mock).mockImplementation(() => mockVPCNAS); - await expect(model.getModelStatus(mockDevClient, 'test-model')).rejects.toThrow( - '[Download-model] get model status error: Get status failed for model test-model', + await (model as any)._handleNasAutoDeployment( + 'cn-hangzhou', + {}, + true, + false, + 'test-function', ); + + expect(mockVPCNAS.deploy).toHaveBeenCalledWith({ + nasAuto: true, + vpcConfig: undefined, + }); + + expect((model as any).createResource.nas).toEqual({ + mountTargetDomain: 'test-domain', + fileSystemId: 'fs-test', + }); }); }); - describe('parseNasConfig', () => { - it('should parse nasConfig correctly', () => { - // Since parseNasConfig is private, we'll test it through the download method - // which calls it internally - expect(true).toBe(true); + describe('_buildParams', () => { + it('should build params correctly with OSS mount points', () => { + model = new Model(mockInputs); + + const params = (model as any)._buildParams( + { + id: 'test-model', + source: { uri: 'modelscope://test-model' }, + target: { uri: 'nas://auto' }, + version: '1.0.0', + files: [], + }, + 'cn-hangzhou', + '123456789', + undefined, + undefined, + { mountPoints: [{ bucketName: 'test-bucket' }] }, + 'test-function', + ); + + expect(params.ossMountPoints).toEqual([{ bucketName: 'test-bucket' }]); }); - it('should throw error for invalid serverAddr', () => { - // Since parseNasConfig is private, we'll test it through the download method - // which calls it internally - expect(true).toBe(true); + it('should build params correctly with NAS mount points', () => { + model = new Model(mockInputs); + + const params = (model as any)._buildParams( + { + id: 'test-model', + source: { uri: 'modelscope://test-model' }, + target: { uri: 'nas://auto' }, + version: '1.0.0', + files: [], + }, + 'cn-hangzhou', + '123456789', + { mountPoints: [{ serverAddr: 'test-server' }] }, + { vpcId: 'vpc-test' }, + undefined, + 'test-function', + ); + + expect(params.nasMountPoints).toEqual([{ serverAddr: 'test-server' }]); + expect(params.vpcConfig).toEqual({ vpcId: 'vpc-test' }); + }); + + it('should use default timeout when not specified', () => { + model = new Model(mockInputs); + + const params = (model as any)._buildParams( + { + id: 'test-model', + source: { uri: 'modelscope://test-model' }, + target: { uri: 'nas://auto' }, + version: '1.0.0', + files: [], + }, + 'cn-hangzhou', + '123456789', + undefined, + undefined, + undefined, + 'test-function', + ); + + expect(params.modelConfig.timeout).toBe(MODEL_DOWNLOAD_TIMEOUT); }); }); }); diff --git a/__tests__/ut/commands/model_utils_test.ts b/__tests__/ut/commands/model_utils_test.ts new file mode 100644 index 00000000..7bc76c31 --- /dev/null +++ b/__tests__/ut/commands/model_utils_test.ts @@ -0,0 +1,313 @@ +import { + _getEndpoint, + initClient, + _displayProgress, + _displayProgressComplete, + checkModelStatus, +} from '../../../src/subCommands/model/utils'; +import DevClient from '@alicloud/devs20230714'; +import * as $OpenApi from '@alicloud/openapi-client'; +import { IInputs } from '../../../src/interface'; +import { sleep } from '../../../src/utils'; + +// Mock dependencies +jest.mock('../../../src/logger', () => { + const mockLogger = { + log: jest.fn(), + info: jest.fn(), + debug: jest.fn(), + warn: jest.fn(), + write: jest.fn(), + error: jest.fn(), + output: jest.fn(), + spin: jest.fn(), + tips: jest.fn(), + append: jest.fn(), + tipsOnce: jest.fn(), + warnOnce: jest.fn(), + writeOnce: jest.fn(), + }; + return { + __esModule: true, + default: mockLogger, + }; +}); + +jest.mock('@alicloud/devs20230714'); +jest.mock('@alicloud/openapi-client'); +jest.mock('../../../src/utils'); + +describe('Model Utils', () => { + describe('_getEndpoint', () => { + afterEach(() => { + delete process.env.ARTIFACT_ENDPOINT; + delete process.env.artifact_endpoint; + }); + + it('should return ARTIFACT_ENDPOINT when it is set', () => { + process.env.ARTIFACT_ENDPOINT = 'custom.endpoint.com'; + process.env.artifact_endpoint = 'custom.endpoint.com'; + + const endpoint = _getEndpoint('cn-hangzhou'); + + expect(endpoint).toBe('custom.endpoint.com'); + }); + + it('should return artifact_endpoint when ARTIFACT_ENDPOINT is not set but artifact_endpoint is set', () => { + process.env.artifact_endpoint = 'custom2.endpoint.com'; + + const endpoint = _getEndpoint('cn-hangzhou'); + + expect(endpoint).toBe('custom2.endpoint.com'); + }); + + it('should return default endpoint when neither environment variable is set', () => { + const endpoint = _getEndpoint('cn-hangzhou'); + + expect(endpoint).toBe('devs.cn-hangzhou.aliyuncs.com'); + }); + }); + + describe('initClient', () => { + let mockInputs: IInputs; + + beforeEach(() => { + mockInputs = { + cwd: '/test', + baseDir: '/test', + name: 'test-app', + props: { + region: 'cn-hangzhou', + functionName: 'test-function', + runtime: 'nodejs18', + handler: 'index.handler', + code: './code', + }, + command: 'model', + args: ['download'], + yaml: { + path: '/test/s.yaml', + }, + resource: { + name: 'test-resource', + component: 'fc3', + access: 'default', + }, + outputs: {}, + getCredential: jest.fn().mockResolvedValue({ + AccountID: '123456789', + AccessKeyID: 'test-key', + AccessKeySecret: 'test-secret', + SecurityToken: 'test-token', + }), + userAgent: 'test-agent', + }; + }); + + afterEach(() => { + jest.clearAllMocks(); + delete process.env.ARTIFACT_ENDPOINT; + delete process.env.artifact_endpoint; + }); + + it('should create DevClient with default configuration', async () => { + const mockConfig = { + accessKeyId: 'test-key', + accessKeySecret: 'test-secret', + securityToken: 'test-token', + protocol: 'https', + endpoint: 'devs.cn-hangzhou.aliyuncs.com', + readTimeout: 300000, + connectTimeout: 300000, + userAgent: 'test-agent', + }; + + ($OpenApi.Config as unknown as jest.Mock).mockImplementation((config) => config); + + const client = await initClient( + mockInputs, + 'cn-hangzhou', + ( + await import('../../../src/logger') + ).default, + 'fun-model', + ); + + expect($OpenApi.Config).toHaveBeenCalledWith(expect.objectContaining(mockConfig)); + expect(client).toBeInstanceOf(DevClient); + }); + + it('should use custom endpoint from ARTIFACT_ENDPOINT environment variable', async () => { + process.env.ARTIFACT_ENDPOINT = 'custom.endpoint.com'; + + ($OpenApi.Config as unknown as jest.Mock).mockImplementation((config) => config); + + const client = await initClient( + mockInputs, + 'cn-hangzhou', + ( + await import('../../../src/logger') + ).default, + 'fun-art', + ); + + expect($OpenApi.Config).toHaveBeenCalledWith( + expect.objectContaining({ + endpoint: 'custom.endpoint.com', + }), + ); + expect(client).toBeInstanceOf(DevClient); + }); + }); + + describe('checkModelStatus', () => { + let mockDevClient: jest.Mocked; + let mockLogger: any; + + beforeEach(() => { + mockDevClient = { + getFileManagerTask: jest.fn(), + } as any; + + mockLogger = { + log: jest.fn(), + info: jest.fn(), + debug: jest.fn(), + warn: jest.fn(), + write: jest.fn(), + error: jest.fn(), + output: jest.fn(), + spin: jest.fn(), + tips: jest.fn(), + append: jest.fn(), + tipsOnce: jest.fn(), + warnOnce: jest.fn(), + writeOnce: jest.fn(), + }; + + (sleep as jest.Mock).mockResolvedValue(undefined); + }); + + afterEach(() => { + jest.clearAllMocks(); + }); + + it('should complete successfully when task is finished and successful', async () => { + mockDevClient.getFileManagerTask.mockResolvedValue({ + body: { + data: { + finished: true, + success: true, + startTime: Date.now() - 1000, + finishedTime: Date.now(), + progress: { + currentBytes: 1024, + totalBytes: 1024, + total: true, + }, + errorMessage: null, + }, + }, + } as any); + + const result = await checkModelStatus( + mockDevClient, + 'task-123', + mockLogger, + 'file1.txt', + 30000, + ); + + expect(result).toBe(true); + expect(mockLogger.info).toHaveBeenCalledWith('Time taken for file1.txt download: 1s.'); + expect(mockLogger.info).toHaveBeenCalledWith('[Download-model] Download file1.txt finished.'); + }); + + it('should throw error when task has errorMessage', async () => { + mockDevClient.getFileManagerTask.mockResolvedValue({ + body: { + data: { + finished: true, + success: false, + startTime: Date.now() - 1000, + finishedTime: Date.now(), + progress: { + currentBytes: 0, + totalBytes: 0, + }, + errorMessage: 'Download failed', + }, + requestId: 'req-123', + }, + } as any); + + await expect( + checkModelStatus(mockDevClient, 'task-123', mockLogger, 'file1.txt', 30000), + ).rejects.toThrow('[Download-model] file1.txt: Download failed ,requestId: req-123'); + }); + + it('should handle download timeout', async () => { + mockDevClient.getFileManagerTask.mockResolvedValue({ + body: { + data: { + finished: false, + startTime: Date.now() - 50 * 60 * 1000, // 50 minutes ago + progress: { + currentBytes: 512, + totalBytes: 1024, + }, + }, + }, + } as any); + + await expect( + checkModelStatus(mockDevClient, 'task-123', mockLogger, 'file1.txt', 30000), + ).rejects.toThrow('Download timeout after 0.5 minutes'); + }); + + it('should adjust sleep time for large files', async () => { + // First call returns unfinished task with large file size + mockDevClient.getFileManagerTask.mockResolvedValueOnce({ + body: { + data: { + finished: false, + startTime: Date.now(), + progress: { + currentBytes: 512, + totalBytes: 2 * 1024 * 1024 * 1024, // 2GB file + }, + }, + }, + } as any); + + // Second call returns finished task + mockDevClient.getFileManagerTask.mockResolvedValueOnce({ + body: { + data: { + finished: true, + success: true, + startTime: Date.now() - 10000, + finishedTime: Date.now(), + progress: { + currentBytes: 2 * 1024 * 1024 * 1024, + totalBytes: 2 * 1024 * 1024 * 1024, + total: true, + }, + }, + }, + } as any); + + const result = await checkModelStatus( + mockDevClient, + 'task-123', + mockLogger, + 'file1.txt', + 30000, + ); + + expect(result).toBe(true); + // For large files, it should sleep for 10 seconds instead of 2 + expect(sleep).toHaveBeenCalledWith(10); + }); + }); +}); diff --git a/package-lock.json b/package-lock.json index b28eb443..9214e09c 100644 --- a/package-lock.json +++ b/package-lock.json @@ -10,7 +10,7 @@ "hasInstallScript": true, "license": "ISC", "dependencies": { - "@alicloud/devs20230714": "^2.4.6-alpha.2", + "@alicloud/devs20230714": "^2.5.0", "@alicloud/fc2": "^2.6.6", "@alicloud/fc20230330": "4.6.3", "@alicloud/pop-core": "^1.8.0", @@ -154,9 +154,9 @@ } }, "node_modules/@alicloud/devs20230714": { - "version": "2.4.6-alpha.2", - "resolved": "https://packages.aliyun.com/670e108663cd360abfe4be65/npm/npm-registry/@alicloud/devs20230714/-/@alicloud/devs20230714-2.4.6-alpha.2.tgz", - "integrity": "sha512-7gSRncwItnPWNDcs91PqILLsZCcfeZ/C8iwbo/6Z3S8mNMIvtxXYyGjwOWncHCOqVxQOvt6vzIEH/0KXeNyOHg==", + "version": "2.5.0", + "resolved": "https://packages.aliyun.com/670e108663cd360abfe4be65/npm/npm-registry/@alicloud/devs20230714/-/@alicloud/devs20230714-2.5.0.tgz", + "integrity": "sha512-P4+B/IOn+/Y4JEhtkMHY9d0w2AdKAPtXcRG3qkjL8rnmyHiwD15iEDtcY6Bpv7LbVw/4nyXab6WFPezqbp6/4w==", "license": "Apache-2.0", "dependencies": { "@alicloud/openapi-core": "^1.0.0", @@ -16110,9 +16110,9 @@ } }, "@alicloud/devs20230714": { - "version": "2.4.6-alpha.2", - "resolved": "https://packages.aliyun.com/670e108663cd360abfe4be65/npm/npm-registry/@alicloud/devs20230714/-/@alicloud/devs20230714-2.4.6-alpha.2.tgz", - "integrity": "sha512-7gSRncwItnPWNDcs91PqILLsZCcfeZ/C8iwbo/6Z3S8mNMIvtxXYyGjwOWncHCOqVxQOvt6vzIEH/0KXeNyOHg==", + "version": "2.5.0", + "resolved": "https://packages.aliyun.com/670e108663cd360abfe4be65/npm/npm-registry/@alicloud/devs20230714/-/@alicloud/devs20230714-2.5.0.tgz", + "integrity": "sha512-P4+B/IOn+/Y4JEhtkMHY9d0w2AdKAPtXcRG3qkjL8rnmyHiwD15iEDtcY6Bpv7LbVw/4nyXab6WFPezqbp6/4w==", "requires": { "@alicloud/openapi-core": "^1.0.0", "@darabonba/typescript": "^1.0.0" diff --git a/package.json b/package.json index de85473b..f4a9373d 100644 --- a/package.json +++ b/package.json @@ -22,7 +22,7 @@ "author": "", "license": "ISC", "dependencies": { - "@alicloud/devs20230714": "^2.4.6-alpha.2", + "@alicloud/devs20230714": "^2.5.0", "@alicloud/fc2": "^2.6.6", "@alicloud/fc20230330": "4.6.3", "@alicloud/pop-core": "^1.8.0", diff --git a/publish.yaml b/publish.yaml index d9975735..0c517727 100644 --- a/publish.yaml +++ b/publish.yaml @@ -3,7 +3,7 @@ Type: Component Name: fc3 Provider: - 阿里云 -Version: 0.1.3 +Version: dev Description: 阿里云函数计算全生命周期管理 HomePage: https://github.com/devsapp/fc3 Organization: 阿里云函数计算(FC) diff --git a/src/subCommands/model/constants.ts b/src/subCommands/model/constants.ts new file mode 100644 index 00000000..aa8e8a63 --- /dev/null +++ b/src/subCommands/model/constants.ts @@ -0,0 +1,6 @@ +export const NEW_MODEL_SERVICE_CLIENT_CONNECT_TIMEOUT: number = + parseInt(process.env.NEW_MODEL_SERVICE_CLIENT_CONNECT_TIMEOUT as string, 10) || 5 * 60 * 1000; +export const NEW_MODEL_SERVICE_CLIENT_READ_TIMEOUT: number = + parseInt(process.env.NEW_MODEL_SERVICE_CLIENT_READ_TIMEOUT as string, 10) || 5 * 60 * 1000; +export const MODEL_DOWNLOAD_TIMEOUT: number = + parseInt(process.env.MODEL_DOWNLOAD_TIMEOUT as string, 10) || 40 * 60 * 1000; diff --git a/src/subCommands/model/fileManager.ts b/src/subCommands/model/fileManager.ts new file mode 100644 index 00000000..b76c7c14 --- /dev/null +++ b/src/subCommands/model/fileManager.ts @@ -0,0 +1,382 @@ +// 文生图服务下载,删除逻辑 +import logger from '../../logger'; +import DevClient, * as $Dev20230714 from '@alicloud/devs20230714'; +import { IInputs } from '../../interface'; +import _ from 'lodash'; +import { checkModelStatus, initClient } from './utils'; + +export class ArtModelService { + logger = logger; + region: string; + constructor(private inputs: IInputs) { + const { region } = this.inputs.props; + this.region = region; + } + + getSourceAndDestination(uri, file, nasMountPoints, ossMountPoints, targetUri) { + // 处理源路径 + const source = this._getSourcePath(file, uri); + // 处理目标路径 + const destination = this._getDestinationPath(targetUri, file, nasMountPoints, ossMountPoints); + + return { + source, + destination, + }; + } + + async downloadModel(name, params) { + const devClient = await initClient(this.inputs, this.region, logger, 'fun-art'); + const { nasMountPoints, ossMountPoints, role, modelConfig, vpcConfig, region } = params; + const { files } = modelConfig; + + if (modelConfig.mode === 'never') { + logger.info( + '[Download-model] Skipping model download as modelConfig.mode is set to "never".', + ); + return; + } + + if (_.isEmpty(files)) { + logger.info('[Download-model] No files specified for download.'); + return; + } + + let existingTasks = null; + if (modelConfig.mode === 'once') { + // 先统一获取已有的任务列表 + const ListFileManagerTasksRequest = new $Dev20230714.ListFileManagerTasksRequest({ + name, + }); + const res = await devClient.listFileManagerTasks(ListFileManagerTasksRequest); + logger.debug('listFileManagerTasks', JSON.stringify(res, null, 2)); + existingTasks = res.body.data.tasks; + } + + // 第一步:筛选真正需要下载的文件 + const filesNeedPromises = files.map(async (file) => { + const { source, destination } = this.getSourceAndDestination( + modelConfig.source.uri, + file, + nasMountPoints, + ossMountPoints, + modelConfig.target.uri, + ); + + const needDownload = !_.isEmpty(existingTasks) + ? !existingTasks?.some( + (task) => + task.finished && + task.success && + task.progress.currentBytes === task.progress.totalBytes && + task.parameters.destination === destination && + task.parameters.source === source, + ) + : true; + + if (!needDownload) { + logger.info(`[Download-model] ${file.source.path} The file has been downloaded.`); + return null; + } + + return { + ...file, + source, + fileName: file.source.path, + destination, + }; + }); + + const filesNeedResults = await Promise.all(filesNeedPromises); + const filesNeed = filesNeedResults.filter(Boolean); + + // 添加调试日志 + logger.info(`[Download-model] Total files to check: ${files.length}`); + logger.info(`[Download-model] Files need to download: ${filesNeed.length}`); + + // 如果没有需要下载的文件,直接返回 + if (filesNeed.length === 0) { + logger.info('[Download-model] No files need to be downloaded.'); + return; + } + + // 限制并发数量,避免同时发起过多请求 + const MAX_CONCURRENT_DOWNLOADS = 5; + const results = []; + let successCount = 0; + let failureCount = 0; + const failureDetails: Array<{ fileName: string; error: string }> = []; + + // 使用并发控制执行下载任务 + const downloadTasks = filesNeed.map((file) => + this._downloadSingleFile.bind(this, devClient, file, { + name, + nasMountPoints, + ossMountPoints, + role, + region, + vpcConfig, + conflictResolution: process.env.MODEL_CONFLIC_HANDLING || modelConfig.conflictResolution, + timeout: modelConfig?.timeout, + }), + ); + + // 分批处理文件下载,每批最多 MAX_CONCURRENT_DOWNLOADS 个并发 + for (let i = 0; i < downloadTasks.length; i += MAX_CONCURRENT_DOWNLOADS) { + const batch = downloadTasks.slice(i, i + MAX_CONCURRENT_DOWNLOADS); + const batchPromises = batch.map((task) => task()); + + try { + // eslint-disable-next-line no-await-in-loop + const batchResults = await Promise.allSettled(batchPromises); + // 处理批处理结果 + // eslint-disable-next-line no-loop-func + batchResults.forEach((result, index) => { + if (result.status === 'fulfilled') { + results.push(result.value); + successCount++; + logger.info( + `[Download-model] Successfully downloaded file: ${filesNeed[i + index].fileName}`, + ); + } else { + failureCount++; + const { fileName } = filesNeed[i + index]; + const errorDetail = result.reason.stack || result.reason.toString(); + + // 记录详细错误信息到数组 + failureDetails.push({ + fileName, + error: errorDetail, + }); + } + }); + } catch (error) { + failureCount++; + logger.error(`[Download-model] Batch download error: ${error.message}`); + logger.error(`[Download-model] Batch error details:`, error.stack || error); + } + } + + // 输出最终统计信息 + logger.info( + `[Download-model] All files download completed. Success: ${successCount}, Failed: ${failureCount}, Total: ${filesNeed.length}`, + ); + + // 打印所有失败的详细信息 + if (failureDetails.length > 0) { + logger.error('[Download-model] Detailed failure information:'); + failureDetails.forEach((detail, index) => { + logger.error(` ${index + 1}. File: ${detail.fileName}`); + logger.error(` Error: ${detail.error}`); + }); + } + + // 文件下载失败,抛出错误 + if (failureCount > 0) { + throw new Error( + `[Download-model] ${failureCount} out of ${filesNeed.length} files failed to download.`, + ); + } + } + + async removeModel(name, params) { + const { nasMountPoints, ossMountPoints, role, vpcConfig, modelConfig, region } = params; + try { + const devClient = await initClient(this.inputs, this.region, logger, 'fun-art'); + const { files } = modelConfig; + if (_.isEmpty(files)) { + logger.info('[Remove-model] No files specified for removal.'); + return; + } + + // 将异步操作重构为并行处理 + const removePromises = files.map((file) => + this._removeSingleFile(devClient, file, { + name, + nasMountPoints, + ossMountPoints, + role, + vpcConfig, + modelConfig, + region, + timeout: modelConfig?.timeout, + }), + ); + + await Promise.all(removePromises); + logger.info(`[Remove-model] Completed removal process for ${files.length} files.`); + } catch (error) { + throw new Error(`[Remove-model] Removal process failed: ${error.message}`); + } + } + + private async _removeSingleFile( + devClient: DevClient, + file: any, + config: { + name: string; + nasMountPoints: any[]; + ossMountPoints: any[]; + role: string; + vpcConfig: any; + modelConfig: any; + region: string; + timeout: number; + }, + ) { + const { name, nasMountPoints, ossMountPoints, role, vpcConfig, modelConfig, region, timeout } = + config; + + try { + let filepath; + const uri = file.target?.uri || modelConfig.target.uri; + const path = file.target?.path || ''; + + // 判断uri是否为nas://auto或oss://auto + if (uri.startsWith('nas://auto') && nasMountPoints?.length > 0) { + const { mountDir } = nasMountPoints[0]; + filepath = `${mountDir}/${path}`; + } else if (uri.startsWith('oss://auto') && ossMountPoints?.length > 0) { + const { mountDir } = ossMountPoints[0]; + filepath = `${mountDir}/${path}`; + } else { + // 直接拼接uri和path + const normalizedUri = uri.endsWith('/') ? uri.slice(0, -1) : uri; + filepath = `${normalizedUri}/${path}`; + } + + const fileManagerRmRequest = new $Dev20230714.FileManagerRmRequest({ + filepath, + mountConfig: new $Dev20230714.FileManagerMountConfig({ + name, + nasMountPoints, + ossMountPoints, + role, + vpcConfig, + region, + timeoutInSecond: timeout, + }), + }); + + logger.debug('FileManagerRmRequest', JSON.stringify(fileManagerRmRequest, null, 2)); + const res = await devClient.fileManagerRm(fileManagerRmRequest); + logger.debug( + `[Remove-model] Remove response for ${file.source.path}:`, + JSON.stringify(res, null, 2), + ); + + if (!res.body.success) { + logger.warn( + `[Remove-model] Failed to remove file ${file.source.path}: ${JSON.stringify(res.body)}`, + ); + } else { + logger.info(`[Remove-model] Successfully removed file ${file.source.path}`); + } + } catch (error) { + logger.error(`[Remove-model] Error removing file ${file.source.path}: ${error.message}`); + logger.error(`[Remove-model] Error details:`, error.stack || error); + } + } + + private async _downloadSingleFile(devClient: any, file: any, config: any) { + const { source, destination, fileName } = file; + const { + name, + nasMountPoints, + ossMountPoints, + role, + region, + vpcConfig, + conflictResolution, + timeout, + } = config; + + try { + // 发起文件同步请求 + const fileManagerRsyncRequest = new $Dev20230714.FileManagerRsyncRequest({ + mountConfig: new $Dev20230714.FileManagerMountConfig({ + name, + nasMountPoints, + ossMountPoints, + role, + region, + vpcConfig, + timeoutInSecond: timeout, + }), + source, + destination, + conflictHandling: conflictResolution, + }); + logger.debug('FileManagerRsyncRequest', JSON.stringify(fileManagerRsyncRequest, null, 2)); + const req = await devClient.fileManagerRsync(fileManagerRsyncRequest); + logger.debug( + `[Download-model] fileManagerRsync response for ${fileName}: ${JSON.stringify( + req.body, + null, + 2, + )}`, + ); + if (!req?.body.success) { + const errorMsg = `fileManagerRsync error: ${JSON.stringify(req?.body, null, 2)}`; + logger.error(`[Download-model] ${fileName}: ${errorMsg}`); + throw new Error(errorMsg); + } + + const { taskID } = req.body.data; + logger.info( + `[Download-model] download model requestId for ${fileName}: ${req.body.requestId}, taskID: ${taskID}`, + ); + + // 轮询任务状态直到完成 + await checkModelStatus(devClient, taskID, logger, fileName, timeout); + } catch (error) { + // 捕获并重新抛出错误,添加文件名信息 + logger.error(`\n[Download-model] Error downloading file ${fileName}: ${error.message}`); + throw new Error(`${fileName}: ${error.message}`); + } + } + + private _getSourcePath(file: any, sourceUri: string): string { + // 会有多个文件,file.source.uri 是该文件下载源路径,sourceUri 是公共的下载源 + const uri = file.source.uri || sourceUri; + const path = file.source?.path || ''; + const validSourcePattern = /^(modelscope|oss|nas):\/\//; + + if (validSourcePattern.test(uri)) { + const downloadUri = uri.endsWith('/') ? uri.slice(0, -1) : uri; + return `${downloadUri}/${path}`; + } else { + throw new Error( + `Invalid source path. Expected a valid URI starting with 'modelscope://', 'oss://', or 'nas://', but got: ${path}`, + ); + } + } + + private _getDestinationPath( + targetUri: string, + file: any, + nasMountPoints: any[], + ossMountPoints: any[], + ): string { + // file.target.uri 多个nas或者oss挂载点时,优先判断是否有指定挂载点,否则使用默认挂载点 + const uri = file.target?.uri || targetUri; + const path = file.target?.path || ''; + + // 判断uri是否为nas://auto或oss://auto + if (uri.startsWith('nas://auto') && nasMountPoints?.length > 0) { + const mountDir = nasMountPoints[0].mountDir.startsWith('/') + ? nasMountPoints[0].mountDir.slice(1) + : nasMountPoints[0].mountDir; + return `file://${mountDir}/${path}`; + } else if (uri.startsWith('oss://auto') && ossMountPoints?.length > 0) { + const mountDir = ossMountPoints[0].mountDir.startsWith('/') + ? ossMountPoints[0].mountDir.slice(1) + : ossMountPoints[0].mountDir; + return `file://${mountDir}/${path}`; + } else { + // 直接拼接uri和path + const normalizedUri = uri.endsWith('/') ? uri.slice(0, -1) : uri; + return `file://${normalizedUri}/${path}`; + } + } +} diff --git a/src/subCommands/model/index.ts b/src/subCommands/model/index.ts index 81edb3d6..da58aea5 100644 --- a/src/subCommands/model/index.ts +++ b/src/subCommands/model/index.ts @@ -1,26 +1,18 @@ -import { IFunction, IInputs } from '../../interface'; +import { IFunction, IInputs, IRegion } from '../../interface'; import logger from '../../logger'; import _, { isEmpty } from 'lodash'; import FC from '../../resources/fc'; import VPC_NAS from '../../resources/vpc-nas'; import { ICredentials } from '@serverless-devs/component-interface'; import { yellow } from 'chalk'; -import DevClient, { DownloadModelRequest } from '@alicloud/devs20230714'; -import * as $OpenApi from '@alicloud/openapi-client'; import { getEnvVariable } from '../../default/resources'; import commandsHelp from '../../commands-help/layer'; import { parseArgv } from '@serverless-devs/utils'; import assert from 'assert'; -import { sleep } from '../../utils'; import OSS from '../../resources/oss'; import { OSSMountPoint, VPCConfig } from '@alicloud/fc20230330'; - -export const NEW_MODEL_SERVICE_CLIENT_CONNECT_TIMEOUT: number = - parseInt(process.env.NEW_MODEL_SERVICE_CLIENT_CONNECT_TIMEOUT as string, 10) || 60 * 1000; -export const NEW_MODEL_SERVICE_CLIENT_READ_TIMEOUT: number = - parseInt(process.env.NEW_MODEL_SERVICE_CLIENT_READ_TIMEOUT as string, 10) || 86400 * 1000; -export const MODEL_DOWNLOAD_TIMEOUT: number = - parseInt(process.env.MODEL_DOWNLOAD_TIMEOUT as string, 10) || 42 * 60 * 1000; +import getUuid from 'uuid-by-string'; +import { MODEL_DOWNLOAD_TIMEOUT } from './constants'; const commandsList = Object.keys(commandsHelp.subCommands); @@ -31,6 +23,9 @@ export class Model { local: IFunction; projectName: string; envName: string; + modelService: any; + modelArtService: any; + name: string; constructor(private inputs: IInputs) { this.logger.debug( @@ -52,98 +47,190 @@ export class Model { } this.subCommand = subCommand; this.local = _.cloneDeep(inputs.props); + const { + credential: { AccountID: accountID }, + props: { functionName }, + } = this.inputs; + const projectName = getEnvVariable('ALIYUN_DEVS_REMOTE_PROJECT_NAME'); + this.name = `${accountID}$${projectName}$${functionName}$${getUuid(String(accountID))}`; } async download() { // 1. auto ---> auto 包, 有 nasConfig // 2. 调用 download 接口,若返回一个错误是下载服务已经存在,继续等待 get 轮询。 // 3. 轮询 get 接口 + const { annotations } = this.inputs.props; + const modelConfig = annotations?.modelConfig; + + const params = (await this.getParams()) as any; + try { + if (modelConfig.solution === 'funArt') { + const modelArtService = await this.getModelArtService(); + await modelArtService.downloadModel(this.name, params); + } else { + const modelService = await this.getModelService(); + await modelService.downloadModel(this.name, params); + } + } catch (e) { + logger.error(`download model invocation error: ${JSON.stringify(e, null, 2)}`); + throw new Error(`download model error: ${e.message}`); + } + } + + async remove() { + logger.info('[Remove-model] remove model ...'); + const params = await this.getParams(); + const { annotations } = this.inputs.props; + const modelConfig = annotations?.modelConfig; + + try { + if (modelConfig.solution === 'funArt') { + const modelArtService = await this.getModelArtService(); + await modelArtService.removeModel(this.name, params); + } else { + const modelService = await this.getModelService(); + await modelService.removeModel(this.name, params); + } + } catch (e) { + logger.debug(`[Remove-model] delete model invocation error: ${JSON.stringify(e, null, 2)}`); + logger.error(`[Remove-model] delete model invocation error: ${e.message}`); + throw new Error(`[Remove-model] delete model error: ${e.message}`); + } + } + + private async getModelService() { + const { ModelService } = await import('./model'); + return new ModelService(this.inputs); + } + + private async getModelArtService() { + const { ArtModelService } = await import('./fileManager'); + return new ArtModelService(this.inputs); + } + + private _assertArrayOfStrings(variable: any) { + assert(Array.isArray(variable), 'Variable must be an array'); + assert( + variable.every((item) => typeof item === 'string'), + 'Variable must contain only strings', + ); + } + + private async getParams() { const { AccountID: accountID } = await this.inputs.getCredential(); const { credential } = this.inputs; const { region, supplement, annotations } = this.inputs.props; const { functionName } = this.local; const modelConfig = supplement?.modelConfig || annotations?.modelConfig; - if (isEmpty(modelConfig)) { - logger.error(`[Download-model] modelConfig is empty.`); - throw new Error(`[Download-model] modelConfig is empty.`); - } + this._validateModelConfig(modelConfig); logger.info(`[Download-model] Download model start.`); + const { nasAuto, vpcAuto, ossAuto } = FC.computeLocalAuto(this.local); + logger.debug(`[auto] Auto compute local auto, nasAuto: ${nasAuto} ossAuto: ${ossAuto};`); - // 混合更多因子,防止重复 - const projectName = getEnvVariable('ALIYUN_DEVS_REMOTE_PROJECT_NAME'); - const envName = getEnvVariable('ALIYUN_DEVS_REMOTE_ENV_NAME'); + if (ossAuto) { + await this._handleOssAutoDeployment(region, credential); + } - this.projectName = projectName; - this.envName = envName; - const name = `${projectName}$${envName}$${functionName}`; + if (nasAuto || vpcAuto) { + await this._handleNasAutoDeployment(region, credential, nasAuto, vpcAuto, functionName); + } - const { nasAuto, vpcAuto, ossAuto } = FC.computeLocalAuto(this.local); - logger.debug(`[auto] Auto compute local auto, nasAuto: ${nasAuto} ossAuto: ${ossAuto};`); + const { nasConfig, vpcConfig, ossMountConfig } = this.local; + logger.info( + `[Download-model] nasConfig: ${nasConfig} vpcConfig: ${vpcConfig} ossMountConfig: ${ossMountConfig}`, + ); - if (modelConfig?.storage === 'oss' && ossAuto) { - let ossEndpoint = `https://oss-${region}.aliyuncs.com`; - if (process.env.FC_REGION === region) { - ossEndpoint = `oss-${region}-internal.aliyuncs.com`; - } - logger.info(`ossAuto code to ${ossEndpoint}`); - const oss = new OSS(region, credential as ICredentials, ossEndpoint); - const { ossBucket } = await oss.deploy(); - logger.write( - yellow(`Created oss resource succeeded, please replace ossMountConfig: auto in yaml with: + return this._buildParams( + modelConfig, + region, + accountID, + nasConfig, + vpcConfig, + ossMountConfig, + functionName, + ); + } + + private _validateModelConfig(modelConfig: any) { + if (isEmpty(modelConfig)) { + logger.error(`[Download-model] modelConfig is empty.`); + throw new Error(`[Download-model] modelConfig is empty.`); + } + } + + private async _handleOssAutoDeployment(region: IRegion, credential: any) { + let ossEndpoint = `https://oss-${region}.aliyuncs.com`; + if (process.env.FC_REGION === region) { + ossEndpoint = `oss-${region}-internal.aliyuncs.com`; + } + logger.info(`ossAuto code to ${ossEndpoint}`); + const oss = new OSS(region, credential as ICredentials, ossEndpoint); + const { ossBucket, readOnly, mountDir, bucketPath } = await oss.deploy(); + logger.write( + yellow(`Created oss resource succeeded, please replace ossMountConfig: auto in yaml with: ossMountConfig: mountPoints: - - mountDir: /mnt/${ossBucket} + - mountDir: ${mountDir} bucketName: ${ossBucket} endpoint: http://oss-${region}-internal.aliyuncs.com - readOnly: false\n`), - ); - this.createResource.oss = { ossBucket }; - _.set(this.local, 'ossMountConfig', { - mountPoints: [ - { - mountDir: `/mnt/${ossBucket}`, - bucketName: ossBucket, - endpoint: `http://oss-${region}-internal.aliyuncs.com`, - readOnly: false, - }, - ], - }); - } + bucketPath: ${bucketPath} + readOnly: ${readOnly}\n`), + ); + this.createResource.oss = { ossBucket }; + _.set(this.local, 'ossMountConfig', { + mountPoints: [ + { + mountDir, + bucketName: ossBucket, + bucketPath, + endpoint: `http://oss-${region}-internal.aliyuncs.com`, + readOnly, + }, + ], + }); + } - if (modelConfig.storage === 'nas' && (nasAuto || vpcAuto)) { - const client = new VPC_NAS(region, credential as ICredentials); - const localVpcAuto = _.isString(this.local.vpcConfig) ? undefined : this.local.vpcConfig; - // @ts-ignore: nas auto 会返回 mountTargetDomain 和 fileSystemId - const { vpcConfig, mountTargetDomain, fileSystemId } = await client.deploy({ - nasAuto, - vpcConfig: localVpcAuto, - }); + private async _handleNasAutoDeployment( + region: IRegion, + credential: any, + nasAuto: boolean, + vpcAuto: boolean, + functionName: string, + ) { + const client = new VPC_NAS(region, credential as ICredentials); + const localVpcAuto = _.isString(this.local.vpcConfig) ? undefined : this.local.vpcConfig; + // @ts-ignore: nas auto 会返回 mountTargetDomain 和 fileSystemId + const { vpcConfig, mountTargetDomain, fileSystemId } = await client.deploy({ + nasAuto, + vpcConfig: localVpcAuto, + }); - if (vpcAuto) { - const { vSwitchIds } = vpcConfig; - this._assertArrayOfStrings(vSwitchIds); - const vSwitchIdsArray: string[] = vSwitchIds as string[]; - logger.write( - yellow(`[nasAuto] Created vpc resource succeeded, please manually write vpcConfig to the yaml file: + if (vpcAuto) { + const { vSwitchIds } = vpcConfig; + this._assertArrayOfStrings(vSwitchIds); + const vSwitchIdsArray: string[] = vSwitchIds as string[]; + logger.write( + yellow(`[nasAuto] Created vpc resource succeeded, please manually write vpcConfig to the yaml file: vpcConfig: vpcId: ${vpcConfig.vpcId} securityGroupId: ${vpcConfig.securityGroupId} vSwitchIds: - - ${vSwitchIdsArray.join(' - \n')}\n`), - ); - this.createResource.vpc = vpcConfig; - _.set(this.local, 'vpcConfig', vpcConfig); - logger.info('[nasAuto] vpcAuto finished.'); + - ${vSwitchIdsArray.join('\n - ')}\n`), + ); + this.createResource.vpc = vpcConfig; + _.set(this.local, 'vpcConfig', vpcConfig); + logger.info('[nasAuto] vpcAuto finished.'); + } + + if (nasAuto) { + let serverAddr = `${mountTargetDomain}:/${functionName}`; + if (serverAddr.length > 128) { + serverAddr = serverAddr.substring(0, 128); } - if (nasAuto) { - let serverAddr = `${mountTargetDomain}:/${functionName}`; - if (serverAddr.length > 128) { - serverAddr = serverAddr.substring(0, 128); - } - logger.write( - yellow(`[nasAuto] Created nas resource succeeded, please replace nasConfig: auto in yaml with: + logger.write( + yellow(`[nasAuto] Created nas resource succeeded, please replace nasConfig: auto in yaml with: nasConfig: groupId: 0 userId: 0 @@ -151,291 +238,62 @@ mountPoints: - serverAddr: ${serverAddr} mountDir: /mnt/${functionName} enableTLS: false\n`), - ); - this.createResource.nas = { mountTargetDomain, fileSystemId }; - _.set(this.local, 'nasConfig', { - groupId: 0, - userId: 0, - mountPoints: [ - { - serverAddr, - mountDir: `/mnt/${functionName}`, - enableTLS: false, - }, - ], - }); - } - } - - // downloadModel - const devClient = await this.getNewModelServiceClient(); - - if (devClient) { - const { nasConfig, vpcConfig, ossMountConfig } = this.local; - logger.debug( - `[Download-model] nasConfig: ${nasConfig} vpcConfig: ${vpcConfig} ossMountConfig: ${ossMountConfig}`, - ); - let resp; - const params: any = { - modelConfig: { - model: modelConfig.id, - type: modelConfig.source, - reversion: modelConfig.version, - token: modelConfig.token, - bucket: modelConfig.ossBucket, - path: modelConfig.ossPath, - region: modelConfig.ossRegion, - }, - region, - // 使用固定的默认角色ARN,确保权限一致性 - role: `acs:ram::${accountID}:role/aliyundevsdefaultrole`, - syncStrategy: process.env.MODEL_DOWNLOAD_STRATEGY || 'incremental_once', - }; - - if ( - modelConfig.storage === 'oss' && - typeof ossMountConfig === 'object' && - ossMountConfig?.mountPoints - ) { - params.ossMountPoints = ossMountConfig.mountPoints as OSSMountPoint[]; - } else if ( - modelConfig.storage === 'nas' && - typeof nasConfig === 'object' && - nasConfig?.mountPoints - ) { - const { nasMountDomain, nasMountPath } = this.parseNasConfig(nasConfig); - params.vpcConfig = vpcConfig as VPCConfig; - params.nasMountPoint = nasMountDomain; - params.modelPath = nasMountPath; - } - try { - const req = new DownloadModelRequest(params); - logger.debug(req); - resp = await devClient.downloadModel(name, req); - logger.debug(resp); - } catch (e) { - logger.error(`download model invocation error: ${e.message}`); - throw new Error(`download model error: ${e.message}`); - } - - if (resp.statusCode !== 200 && resp.statusCode !== 202) { - logger.info({ status: resp.statusCode, body: resp.body }); - throw new Error( - `download model connection error, statusCode: ${resp.statusCode}, body: ${resp.body}`, - ); - } - - const rb = resp.body; - if (rb.success || rb.errMsg.includes('is already exist')) { - logger.info(`download model requestId: ${rb.requestId}`); - const shouldContinue = true; - while (shouldContinue) { - // eslint-disable-next-line no-await-in-loop - const modelStatus = await this.getModelStatus(devClient, name); - - if (modelStatus.finished) { - // 如果存在错误信息,则抛出异常 - if (modelStatus.errMsg) { - logger.error(`[Download-model] ${modelStatus.errMsg}`); - throw new Error(`[Download-model] ${modelStatus.errMsg}`); - } - // 下载成功完成 - if ( - modelStatus.total && - modelStatus.currentBytes !== undefined && - modelStatus.fileSize !== undefined - ) { - const currentMB = (modelStatus.currentBytes / 1024 / 1024).toFixed(1); - const totalMB = (modelStatus.fileSize / 1024 / 1024).toFixed(1); - - const totalBars = 50; - const progressBar = '='.repeat(totalBars); - - process.stdout.write( - `\r[Download-model] [${progressBar}] 100.00% (${currentMB}MB/${totalMB}MB)\n`, - ); - } else { - process.stdout.write('\n'); - } - // 清除进度条并换行 - process.stdout.write('\n'); - if (modelStatus.total) { - const durationMs = modelStatus.finishedTime - modelStatus.startTime; - const durationSeconds = Math.floor(durationMs / 1000); - logger.info(`Time taken for model download: ${durationSeconds}s.`); - } - logger.info(`[Download-model] Download model finished.`); - return true; - } - // 显示下载进度 - if (modelStatus.currentBytes !== undefined && modelStatus.fileSize !== undefined) { - const percentage = (modelStatus.currentBytes / modelStatus.fileSize) * 100; - const currentMB = (modelStatus.currentBytes / 1024 / 1024).toFixed(1); - const totalMB = (modelStatus.fileSize / 1024 / 1024).toFixed(1); - - // 每个等号代表2%,向下取整计算等号数量 - const totalBars = 50; // 总共50个字符位置 - const filledBars = Math.min(totalBars, Math.floor(percentage / 2)); // 每个等号代表2% - const emptyBars = totalBars - filledBars; - - const progressBar = '='.repeat(filledBars) + '.'.repeat(emptyBars); - - process.stdout.write( - `\r[Download-model] [${progressBar}] ${percentage.toFixed( - 2, - )}% (${currentMB}MB/${totalMB}MB)`, - ); - } - - if (Date.now() - modelStatus.startTime > MODEL_DOWNLOAD_TIMEOUT) { - // 清除进度条并换行 - process.stdout.write('\n'); - const errorMessage = `[Model-download] Download timeout after ${ - MODEL_DOWNLOAD_TIMEOUT / 1000 / 60 - } minutes`; - throw new Error(errorMessage); - } - - // 根据文件大小调整轮询间隔 - let sleepTime = 2; // 默认2秒 - if (modelStatus.fileSize !== undefined && modelStatus.fileSize > 1024 * 1024 * 1024) { - // 文件大于1GB时,轮询间隔为10秒 - sleepTime = 10; - } - - // eslint-disable-next-line no-await-in-loop - await sleep(sleepTime); - } - } else { - throw new Error( - `download model service biz failed, errCode: ${rb.errCode}, errMsg: ${rb.errMsg}`, - ); - } - } - } - - async getNewModelServiceClient(): Promise { - const { - AccessKeyID: accessKeyId, - AccessKeySecret: accessKeySecret, - SecurityToken: securityToken, - } = await this.inputs.getCredential(); - - let endpoint: string; - - endpoint = 'devs.cn-hangzhou.aliyuncs.com'; - if (process.env.ARTIFACT_ENDPOINT) { - endpoint = process.env.ARTIFACT_ENDPOINT; - } - if (process.env.artifact_endpoint) { - endpoint = process.env.artifact_endpoint; - } - - const protocol = 'https'; - - const config = new $OpenApi.Config({ - accessKeyId, - accessKeySecret, - securityToken, - protocol, - endpoint, - readTimeout: NEW_MODEL_SERVICE_CLIENT_READ_TIMEOUT, - connectTimeout: NEW_MODEL_SERVICE_CLIENT_CONNECT_TIMEOUT, - userAgent: `${ - this.inputs.userAgent || - `Component:cap-model;Nodejs:${process.version};OS:${process.platform}-${process.arch}` - }`, - }); - - logger.info(`new models service init, ARTIFACT_ENDPOINT endpoint: ${config.endpoint}`); - - return new DevClient(config); - } - - async getModelStatus(client: DevClient, name: string) { - let resp; - try { - resp = await client.getModelStatus(name); - logger.debug(resp); - } catch (e) { - logger.error(`[Download-model] get model status error: ${e.message} for model ${name}`); - throw new Error(`[Download-model] get model status error: ${e.message} for model ${name}`); - } - - if (resp.statusCode !== 200 && resp.statusCode !== 202) { - logger.info({ status: resp.statusCode, body: resp.body }); - throw new Error( - `[Download-model] get model status connection error, statusCode: ${resp.statusCode}, body: ${resp.body}`, - ); - } - - const rb = resp.body; - if (rb.success) { - return rb.data; - } else { - throw new Error( - `[Download-model] get model status biz failed, errCode: ${rb.errCode}, errMsg: ${rb.errMsg}`, ); + this.createResource.nas = { mountTargetDomain, fileSystemId }; + _.set(this.local, 'nasConfig', { + groupId: 0, + userId: 0, + mountPoints: [ + { + serverAddr, + mountDir: `/mnt/${functionName}`, + enableTLS: false, + }, + ], + }); } } - async remove() { - logger.info('[Remove-model] remove model ...'); - const { functionName } = this.inputs.props; - const projectName = getEnvVariable('ALIYUN_DEVS_REMOTE_PROJECT_NAME'); - const envName = getEnvVariable('ALIYUN_DEVS_REMOTE_ENV_NAME'); - const name = `${projectName}$${envName}$${functionName}`; - - const devClient = await this.getNewModelServiceClient(); - - try { - const resp = await devClient.deleteModel(name); - logger.debug(resp); - - if (resp.statusCode !== 200 && resp.statusCode !== 202) { - logger.info({ status: resp.statusCode, body: resp.body }); - throw new Error( - `[Remove-model] delete model connection error, statusCode: ${resp.statusCode}, body: ${resp.body}`, - ); - } - - const rb = resp.body; - if (rb.success) { - logger.info(`[Remove-model] delete model requestId: ${rb.requestId}`); - logger.info(`[Remove-model] Remove model succeeded.`); - return true; - } else if (!rb.errMsg.includes(`${name} is not exist`)) { - throw new Error( - `[Remove-model] delete model service biz failed, errCode: ${rb.errCode}, errMsg: ${rb.errMsg}`, - ); - } - } catch (e) { - logger.error(`[Remove-model] delete model invocation error: ${e.message}`); - throw new Error(`[Remove-model] delete model error: ${e.message}`); + private _buildParams( + modelConfig: any, + region: string, + accountID: string, + nasConfig: any, + vpcConfig: any, + ossMountConfig: any, + functionName: string, + ) { + const params: any = { + modelConfig: { + model: modelConfig.id, + source: modelConfig.source, + uri: modelConfig.source.uri, + target: modelConfig.target, + reversion: modelConfig.version, + files: modelConfig.files, + conflictResolution: modelConfig?.downloadStrategy?.conflictResolution || 'overwrite', + mode: modelConfig?.downloadStrategy?.mode || 'once', + timeout: + (modelConfig?.downloadStrategy?.timeout && + modelConfig?.downloadStrategy?.timeout * 1000) || + MODEL_DOWNLOAD_TIMEOUT, + }, + region, + functionName, + storage: modelConfig.storage, + // 使用固定的默认角色ARN,确保权限一致性 + role: `acs:ram::${accountID}:role/aliyundevsdefaultrole`, + syncStrategy: process.env.MODEL_DOWNLOAD_STRATEGY || 'incremental_once', + }; + + if (typeof ossMountConfig === 'object' && ossMountConfig?.mountPoints) { + params.ossMountPoints = ossMountConfig.mountPoints as OSSMountPoint[]; } - } - - private _assertArrayOfStrings(variable: any) { - assert(Array.isArray(variable), 'Variable must be an array'); - assert( - variable.every((item) => typeof item === 'string'), - 'Variable must contain only strings', - ); - } - - private parseNasConfig(nasConfig: any) { - let nasMountDomain: string; - let nasMountPath = ''; - - const { serverAddr } = nasConfig.mountPoints[0]; - const parts = serverAddr.split(':', 2); - - if (parts.length === 2) { - [nasMountDomain, nasMountPath] = parts; - } else { - throw new Error('nasConfig serverAddr string does not contain a colon'); + if (typeof nasConfig === 'object' && nasConfig?.mountPoints) { + params.vpcConfig = vpcConfig as VPCConfig; + params.nasMountPoints = nasConfig.mountPoints; } - return { nasMountDomain, nasMountPath }; + return params; } } diff --git a/src/subCommands/model/model.ts b/src/subCommands/model/model.ts new file mode 100644 index 00000000..b12581e9 --- /dev/null +++ b/src/subCommands/model/model.ts @@ -0,0 +1,141 @@ +// 模型服务下载,删除逻辑 +import logger from '../../logger'; +import * as $Dev20230714 from '@alicloud/devs20230714'; +import { IInputs } from '../../interface'; +import { checkModelStatus, initClient } from './utils'; +import { isEmpty } from 'lodash'; + +export class ModelService { + logger = logger; + region: string; + constructor(private inputs: IInputs) { + const { region } = this.inputs.props; + this.region = region; + } + + async downloadModel(name, params) { + const devClient = await initClient(this.inputs, this.region, logger, 'fun-model'); + const { nasMountPoints, ossMountPoints, role, modelConfig, vpcConfig, region, storage } = + params; + // 判断modelConfig.source是否是modelscope://、oss://或nas:// + let source; + const reversion = modelConfig.reversion ? `@${modelConfig.reversion}` : ''; + if ( + modelConfig.source.startsWith('modelscope') && + !modelConfig.source.startsWith('modelscope://') + ) { + source = `modelscope://${modelConfig.model}${reversion}`; + } else { + source = `${modelConfig.source}${reversion}`; + } + const validSourcePattern = /^(modelscope|oss):\/\//; + if (!validSourcePattern.test(source)) { + throw new Error( + `Invalid source path. Expected a valid URI starting with 'modelscope://', or 'oss://', but got: ${modelConfig.source}`, + ); + } + + if (modelConfig.mode === 'never') { + logger.info( + '[Download-model] Skipping model download as modelConfig.mode is set to "never".', + ); + return; + } + // mode 是 once 时候,判断是否已经下载过 + const destination = + storage === 'nas' + ? `file:/${nasMountPoints[0].mountDir}` + : `file:/${ + ossMountPoints[0].mountDir.length > 48 + ? ossMountPoints[0].mountDir.substring(0, 48) + : ossMountPoints[0].mountDir + }`; + if (modelConfig.mode === 'once') { + const ListFileManagerTasksRequest = new $Dev20230714.ListFileManagerTasksRequest({ + name, + }); + const res = await devClient.listFileManagerTasks(ListFileManagerTasksRequest); + logger.debug('listFileManagerTasks', JSON.stringify(res, null, 2)); + const { tasks } = res.body.data; + const needDownload = !tasks.some( + (task) => + task.finished && + task.success && + task.progress.currentBytes === task.progress.totalBytes && + task.parameters.source === source && + task.parameters.destination === destination, + ); + if (!needDownload) { + logger.info('[Download-model] The model has been downloaded.'); + return; + } + } + + let processedOssMountPoints; + if (!isEmpty(ossMountPoints)) { + processedOssMountPoints = ossMountPoints.map((ossMountPoint) => ({ + ...ossMountPoint, + mountDir: + ossMountPoint.mountDir.length > 48 + ? ossMountPoint.mountDir.substring(0, 48) + : ossMountPoint.mountDir, + })); + } + + const fileManagerRsyncRequest = new $Dev20230714.FileManagerRsyncRequest({ + mountConfig: new $Dev20230714.FileManagerMountConfig({ + name, + nasMountPoints, + ossMountPoints: processedOssMountPoints, + role, + region, + vpcConfig, + timeoutInSecond: modelConfig.timeout, + }), + source, + destination, + conflictHandling: process.env.MODEL_CONFLIC_HANDLING || modelConfig.conflictResolution, + }); + logger.debug(JSON.stringify(fileManagerRsyncRequest, null, 2)); + const req = await devClient.fileManagerRsync(fileManagerRsyncRequest); + logger.debug('fileManagerRsync', JSON.stringify(req, null, 2)); + if (!req?.body.success) { + throw new Error(`fileManagerRsync error: ${JSON.stringify(req?.body, null, 2)}`); + } + + const { taskID } = req.body.data; + logger.info(`download model requestId: ${req.body.requestId}`); + await checkModelStatus(devClient, taskID, logger, '', modelConfig.timeout); + } + + async removeModel(name, params) { + const { nasMountPoints, ossMountPoints, role, region, vpcConfig, storage } = params; + const devClient = await initClient(this.inputs, this.region, logger, 'fun-model'); + + let processedOssMountPoints; + if (!isEmpty(ossMountPoints)) { + processedOssMountPoints = ossMountPoints.map((ossMountPoint) => ({ + ...ossMountPoint, + mountDir: + ossMountPoint.mountDir.length > 48 + ? ossMountPoint.mountDir.substring(0, 48) + : ossMountPoint.mountDir, + })); + } + + const fileManagerRmRequest = new $Dev20230714.FileManagerRmRequest({ + filepath: storage === 'nas' ? nasMountPoints[0]?.mountDir : ossMountPoints[0]?.mountDir, + mountConfig: new $Dev20230714.FileManagerMountConfig({ + name, + nasMountPoints, + ossMountPoints: processedOssMountPoints, + role, + vpcConfig, + region, + }), + }); + logger.debug('fileManagerRmRequest', JSON.stringify(fileManagerRmRequest, null, 2)); + const res = await devClient.fileManagerRm(fileManagerRmRequest); + logger.debug('removeModel', JSON.stringify(res, null, 2)); + } +} diff --git a/src/subCommands/model/utils/index.ts b/src/subCommands/model/utils/index.ts new file mode 100644 index 00000000..d61914ae --- /dev/null +++ b/src/subCommands/model/utils/index.ts @@ -0,0 +1,149 @@ +import { IInputs } from '../../../interface'; +import { + NEW_MODEL_SERVICE_CLIENT_CONNECT_TIMEOUT, + NEW_MODEL_SERVICE_CLIENT_READ_TIMEOUT, +} from '../constants'; +import * as $OpenApi from '@alicloud/openapi-client'; +import DevClient from '@alicloud/devs20230714'; +import { sleep } from '../../../utils'; + +export const _getEndpoint = (region): string => { + if (process.env.ARTIFACT_ENDPOINT) { + return process.env.ARTIFACT_ENDPOINT; + } + if (process.env.artifact_endpoint) { + return process.env.artifact_endpoint; + } + return `devs.${region}.aliyuncs.com`; +}; + +export const initClient = async (inputs: IInputs, region: string, logger, solution: string) => { + const { + AccessKeyID: accessKeyId, + AccessKeySecret: accessKeySecret, + SecurityToken: securityToken, + } = await inputs.getCredential(); + + const endpoint = _getEndpoint(region); + const protocol = 'https'; + + const config = new $OpenApi.Config({ + accessKeyId, + accessKeySecret, + securityToken, + protocol, + endpoint, + readTimeout: NEW_MODEL_SERVICE_CLIENT_READ_TIMEOUT, + connectTimeout: NEW_MODEL_SERVICE_CLIENT_CONNECT_TIMEOUT, + userAgent: `${ + inputs.userAgent || + `Component:${solution};Nodejs:${process.version};OS:${process.platform}-${process.arch}` + }`, + }); + + logger.info(`new models service init, DEVS_ENDPOINT endpoint: ${config.endpoint}`); + + return new DevClient(config); +}; + +export const _displayProgressComplete = ( + filePath = '', + currentBytes: number, + totalBytes: number, +) => { + if (totalBytes && currentBytes !== undefined) { + const currentMB = (currentBytes / 1024 / 1024).toFixed(1); + const totalMB = (totalBytes / 1024 / 1024).toFixed(1); + + const totalBars = 50; + const progressBar = '='.repeat(totalBars); + + process.stdout.write( + `\r[Download-model] ${filePath} [${progressBar}] 100.00% (${currentMB}MB/${totalMB}MB)\n`, + ); + } else { + process.stdout.write('\n'); + } + // 清除进度条并换行 + process.stdout.write('\n'); +}; + +export const _displayProgress = (filePath = '', currentBytes: number, totalBytes: number) => { + if (currentBytes && totalBytes) { + const percentage = (currentBytes / totalBytes) * 100; + const currentMB = (currentBytes / 1024 / 1024).toFixed(1); + const totalMB = (totalBytes / 1024 / 1024).toFixed(1); + + // 每个等号代表2%,向下取整计算等号数量 + const totalBars = 50; // 总共50个字符位置 + const filledBars = Math.min(totalBars, Math.floor(percentage / 2)); // 每个等号代表2% + const emptyBars = totalBars - filledBars; + + const progressBar = '='.repeat(filledBars) + '.'.repeat(emptyBars); + + process.stdout.write( + `\r[Download-model] ${filePath} [${progressBar}] ${percentage.toFixed( + 2, + )}% (${currentMB}MB/${totalMB}MB)`, + ); + } +}; + +// 查看轮询结果 +export async function checkModelStatus( + devClient: DevClient, + taskID: string, + logger: any, + fileName: string, + timeout: number, +) { + const shouldContinue = true; + while (shouldContinue) { + // eslint-disable-next-line no-await-in-loop + const getFileManager = await devClient.getFileManagerTask(taskID); + logger.debug('getFileManagerTask', JSON.stringify(getFileManager, null, 2)); + const modelStatus = getFileManager.body.data; + const totalBytes = (modelStatus.progress.totalBytes as any) - 0; + const currentBytes = (modelStatus.progress.currentBytes as any) - 0; + + if (modelStatus.finished) { + // 如果存在错误信息,则抛出异常 + if (modelStatus.errorMessage) { + const errorMsg = `[Download-model] ${fileName}: ${modelStatus.errorMessage} ,requestId: ${getFileManager.body.requestId}`; + logger.error(errorMsg); + throw new Error(errorMsg); + } + // 下载成功完成 + _displayProgressComplete(fileName, currentBytes, totalBytes); + if (modelStatus.progress.total) { + const durationMs = modelStatus.finishedTime - modelStatus.startTime; + const durationSeconds = Math.floor(durationMs / 1000); + logger.info(`Time taken for ${fileName || 'model'} download: ${durationSeconds}s.`); + } + logger.info(`[Download-model] Download ${fileName || 'model'} finished.`); + return true; + } + // 显示下载进度 + _displayProgress(fileName, currentBytes, totalBytes); + + const modelTimeout = timeout; + if (Date.now() - modelStatus.startTime > modelTimeout) { + // 清除进度条并换行 + process.stdout.write('\n'); + const errorMessage = `[Model-download] Download timeout after ${ + modelTimeout / 1000 / 60 + } minutes`; + throw new Error(errorMessage); + } + + // 根据文件大小调整轮询间隔 + let sleepTime = 2; // 默认2秒 + if (totalBytes !== undefined && totalBytes > 1024 * 1024 * 1024) { + // 文件大于1GB时,轮询间隔为10秒 + sleepTime = 10; + } + + // eslint-disable-next-line no-await-in-loop + await sleep(sleepTime); + } +}