diff --git a/.gitignore b/.gitignore index ef5e183bd..6ba7b5381 100644 --- a/.gitignore +++ b/.gitignore @@ -161,3 +161,11 @@ test_cookbook/ /test*.py swanlog/ tests/server/config/_generated_e2e.yaml + +# Full field-level contract surface: a generated artifact (8k+ lines, unreviewable diff). +# Regenerate via `python -m tests.server.contract.update_baseline`; do not commit it. +# NOTE: client_api_routes.json is the compact, COMMITTED guard -- do not ignore that one. +tests/server/contract/client_api_baseline.json + +# Redis dump file produced by a local redis-server (test infra), never source. +*.rdb diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index d254bd537..52548fad3 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -3,42 +3,42 @@ repos: rev: 7.3.0 hooks: - id: flake8 - exclude: ^(examples/|cookbook/|client_tools/|src/twinkle_client/|tests/) + exclude: ^(examples/|cookbook/|client_tools/|tests/) - repo: https://github.com/PyCQA/isort rev: 7.0.0 hooks: - id: isort - exclude: ^(examples/|cookbook/|client_tools/|src/twinkle_client/|tests/) + exclude: ^(examples/|cookbook/|client_tools/|tests/) - repo: https://github.com/google/yapf rev: v0.43.0 hooks: - id: yapf - exclude: ^(examples/|cookbook/|client_tools/|src/twinkle_client/|tests/) + exclude: ^(examples/|cookbook/|client_tools/|tests/) - repo: https://github.com/asottile/pyupgrade rev: v3.19.1 hooks: - id: pyupgrade args: [--py38-plus] - exclude: ^(examples/|cookbook/|client_tools/|src/twinkle_client/|tests/) + exclude: ^(examples/|cookbook/|client_tools/|tests/) - repo: https://github.com/pre-commit/pre-commit-hooks rev: v6.0.0 hooks: - id: trailing-whitespace - exclude: ^(client_tools/|src/twinkle_client/|tests/) + exclude: ^(client_tools/|tests/) - id: check-yaml - exclude: ^(client_tools/|src/twinkle_client/|tests/) + exclude: ^(client_tools/|tests/) - id: end-of-file-fixer - exclude: ^(client_tools/|src/twinkle_client/|tests/) + exclude: ^(client_tools/|tests/) - id: requirements-txt-fixer - exclude: ^(client_tools/|src/twinkle_client/|tests/) + exclude: ^(client_tools/|tests/) - id: double-quote-string-fixer - exclude: ^(client_tools/|src/twinkle_client/|tests/) + exclude: ^(client_tools/|tests/) - id: check-merge-conflict - exclude: ^(client_tools/|src/twinkle_client/|tests/) + exclude: ^(client_tools/|tests/) - id: mixed-line-ending args: ["--fix=lf"] - exclude: ^(client_tools/|src/twinkle_client/|tests/) + exclude: ^(client_tools/|tests/) diff --git a/README.md b/README.md index fc64af16f..6f4b05a03 100644 --- a/README.md +++ b/README.md @@ -256,7 +256,7 @@ from twinkle import init_tinker_client from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset, DatasetMeta from twinkle.preprocessor import SelfCognitionProcessor -from twinkle.server.common import input_feature_to_datum +from twinkle.server.model.tinker_datum import input_feature_to_datum base_model = 'ms://Qwen/Qwen3.8-27B' base_url='your-base-url' diff --git a/README_ZH.md b/README_ZH.md index d7b3d66fa..3d5f05fac 100644 --- a/README_ZH.md +++ b/README_ZH.md @@ -245,7 +245,7 @@ from twinkle import init_tinker_client from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset, DatasetMeta from twinkle.preprocessor import SelfCognitionProcessor -from twinkle.server.common import input_feature_to_datum +from twinkle.server.model.tinker_datum import input_feature_to_datum base_model = 'ms://Qwen/Qwen3.8-27B' base_url='your-base-url' diff --git a/cookbook/client/async_rl/client_orchestrated_grpo.py b/cookbook/client/async_rl/client_orchestrated_grpo.py index 0f3162569..cce5ccfd1 100644 --- a/cookbook/client/async_rl/client_orchestrated_grpo.py +++ b/cookbook/client/async_rl/client_orchestrated_grpo.py @@ -15,9 +15,10 @@ from twinkle.dataset import Dataset, DatasetMeta from twinkle.preprocessor.llm import GSM8KProcessor from twinkle.reward import GSM8KAccuracyReward -from twinkle_client import DataPlaneClient, init_twinkle_client +from twinkle import init_twinkle_client +from twinkle_client import DataPlaneClient from twinkle_client.async_rl import Worker, WorkerPipeline -from twinkle_client.common.json_utils import json_safe +from twinkle.protocol.json_utils import json_safe from twinkle_client.model import MultiLoraTransformersModel from twinkle_client.sampler import vLLMSampler diff --git a/cookbook/client/async_rl/server_config.yaml b/cookbook/client/async_rl/server_config.yaml index 103ea523f..d627521da 100644 --- a/cookbook/client/async_rl/server_config.yaml +++ b/cookbook/client/async_rl/server_config.yaml @@ -19,8 +19,7 @@ telemetry: otlp_endpoint: http://localhost:4317 persistence: - mode: file - file_path: /tmp/twinkle_state.json + mode: memory applications: @@ -42,9 +41,6 @@ applications: target_ongoing_requests: 128 ray_actor_options: num_cpus: 0.1 - runtime_env: - env_vars: - TWINKLE_FAIL_FAST: "0" # TransferQueue-backed DataRef service. - name: data-plane @@ -95,7 +91,6 @@ applications: runtime_env: env_vars: TWINKLE_TRUST_REMOTE_CODE: "1" - TWINKLE_FAIL_FAST: "0" # A second GPU hosts vLLM and loads the same local base model. - name: sampler-Qwen3.5-4B @@ -133,7 +128,6 @@ applications: runtime_env: env_vars: TWINKLE_TRUST_REMOTE_CODE: "1" - TWINKLE_FAIL_FAST: "0" - name: processor route_prefix: /api/v1/processor @@ -155,6 +149,3 @@ applications: target_ongoing_requests: 128 ray_actor_options: num_cpus: 0.1 - runtime_env: - env_vars: - TWINKLE_FAIL_FAST: "0" diff --git a/cookbook/client/server/megatron/server_config.yaml b/cookbook/client/server/megatron/server_config.yaml index 24132a07e..680e2db8a 100644 --- a/cookbook/client/server/megatron/server_config.yaml +++ b/cookbook/client/server/megatron/server_config.yaml @@ -18,10 +18,9 @@ telemetry: # Top-level placement makes the launcher propagate this to every Ray worker # via env vars, so the configured backend is used regardless of which # deployment initializes the ServerState actor first. -# mode: memory | file | redis +# mode: memory | redis # memory: requires an initialized Ray runtime (the launcher handles this # automatically; standalone scripts must call ray.init() first) -# file_path: required for `file` mode # redis_url / key_prefix: required for `redis` mode persistence: mode: redis @@ -54,7 +53,6 @@ applications: env_vars: TWINKLE_TRUST_REMOTE_CODE: "0" TWINKLE_LONG_POLL_TIMEOUT: "120" - TWINKLE_FAIL_FAST: "0" # 3. Sampler Service - Runs inference / sampling using vLLM engine # Used for generating text from the model (e.g., evaluating LoRA results). @@ -98,7 +96,6 @@ applications: env_vars: TWINKLE_TRUST_REMOTE_CODE: "0" TWINKLE_LONG_POLL_TIMEOUT: "120" - TWINKLE_FAIL_FAST: "0" # 2. Model Service - Hosts the base model for training. # Config: PP=2 x DP=2 on 4 GPUs, ~27GB weights/GPU, comfortable for LoRA training @@ -139,4 +136,3 @@ applications: env_vars: TWINKLE_TRUST_REMOTE_CODE: "0" TWINKLE_LONG_POLL_TIMEOUT: "120" - TWINKLE_FAIL_FAST: "0" diff --git a/cookbook/client/server/megatron/server_config_4b.yaml b/cookbook/client/server/megatron/server_config_4b.yaml index 9bdbd5e72..7eed4699d 100644 --- a/cookbook/client/server/megatron/server_config_4b.yaml +++ b/cookbook/client/server/megatron/server_config_4b.yaml @@ -31,9 +31,6 @@ applications: target_ongoing_requests: 128 # Target concurrent requests per replica ray_actor_options: num_cpus: 0.1 # CPU resources allocated to this actor - runtime_env: - env_vars: - TWINKLE_FAIL_FAST: "0" # 2. Model Service (commented out) - Would host the base model for training. # Uncomment and configure if you need a training model worker. @@ -71,7 +68,6 @@ applications: runtime_env: env_vars: TWINKLE_TRUST_REMOTE_CODE: "0" - TWINKLE_FAIL_FAST: "0" # 3. Sampler Service - Runs inference / sampling using vLLM engine # Used for generating text from the model (e.g., evaluating LoRA results). @@ -109,7 +105,6 @@ applications: runtime_env: env_vars: TWINKLE_TRUST_REMOTE_CODE: "0" - TWINKLE_FAIL_FAST: "0" # 4. Processor Service - name: processor @@ -132,6 +127,3 @@ applications: target_ongoing_requests: 128 ray_actor_options: num_cpus: 0.1 - runtime_env: - env_vars: - TWINKLE_FAIL_FAST: "0" diff --git a/cookbook/client/server/transformer/server_config.yaml b/cookbook/client/server/transformer/server_config.yaml index d3ddb2adb..7c8050c76 100644 --- a/cookbook/client/server/transformer/server_config.yaml +++ b/cookbook/client/server/transformer/server_config.yaml @@ -18,14 +18,12 @@ telemetry: # Top-level placement makes the launcher propagate this to every Ray worker # via env vars, so the configured backend is used regardless of which # deployment initializes the ServerState actor first. -# mode: memory | file | redis +# mode: memory | redis # memory: requires an initialized Ray runtime (the launcher handles this # automatically; standalone scripts must call ray.init() first) -# file_path: required for `file` mode # redis_url / key_prefix: required for `redis` mode persistence: - mode: file - file_path: /tmp/twinkle_state.json + mode: memory # Applications: each entry defines a service component deployed on the server applications: @@ -49,9 +47,6 @@ applications: target_ongoing_requests: 128 # Target concurrent requests per replica ray_actor_options: num_cpus: 0.1 # CPU resources allocated to this actor - runtime_env: - env_vars: - TWINKLE_FAIL_FAST: "0" # 2. Model Service - Hosts the base model for training. - name: models-Qwen3.5-4B @@ -85,7 +80,6 @@ applications: runtime_env: env_vars: TWINKLE_TRUST_REMOTE_CODE: "1" - TWINKLE_FAIL_FAST: "0" # 3. Sampler Service - Runs inference / sampling using vLLM engine # Used for generating text from the model (e.g., evaluating LoRA results). @@ -122,7 +116,6 @@ applications: runtime_env: env_vars: TWINKLE_TRUST_REMOTE_CODE: "1" - TWINKLE_FAIL_FAST: "0" # 4. Processor Service - name: processor @@ -145,6 +138,3 @@ applications: target_ongoing_requests: 128 ray_actor_options: num_cpus: 0.1 - runtime_env: - env_vars: - TWINKLE_FAIL_FAST: "0" diff --git a/cookbook/client/tinker/dpo.py b/cookbook/client/tinker/dpo.py index 6091e8080..3f6e25d05 100644 --- a/cookbook/client/tinker/dpo.py +++ b/cookbook/client/tinker/dpo.py @@ -23,11 +23,12 @@ import swanlab from tinker import types -from twinkle import init_tinker_client, get_logger +from twinkle import get_logger +from twinkle import init_tinker_client from twinkle.dataset import Dataset, DatasetMeta, LazyDataset from twinkle.dataloader import DataLoader from twinkle.preprocessor import EmojiDPOProcessor -from twinkle.server.common import input_feature_to_datum +from twinkle.server.model.tinker_datum import input_feature_to_datum logger = get_logger() diff --git a/cookbook/client/tinker/multi_modal.py b/cookbook/client/tinker/multi_modal.py index dae6a4dd9..594c4bb8c 100644 --- a/cookbook/client/tinker/multi_modal.py +++ b/cookbook/client/tinker/multi_modal.py @@ -30,7 +30,7 @@ from twinkle.preprocessor import Preprocessor from twinkle.dataset import DatasetMeta, LazyDataset from twinkle.dataloader import DataLoader -from twinkle.server.common import input_feature_to_datum # Key: converts InputFeature -> Datum +from twinkle.server.model.tinker_datum import input_feature_to_datum # Key: converts InputFeature -> Datum from twinkle import get_logger logger = get_logger() diff --git a/cookbook/client/tinker/self_cognition.py b/cookbook/client/tinker/self_cognition.py index 200f7f13a..3a1cd75be 100644 --- a/cookbook/client/tinker/self_cognition.py +++ b/cookbook/client/tinker/self_cognition.py @@ -16,7 +16,7 @@ from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset, DatasetMeta from twinkle.preprocessor import SelfCognitionProcessor -from twinkle.server.common import input_feature_to_datum +from twinkle.server.model.tinker_datum import input_feature_to_datum # Initialize the Tinker client before importing ServiceClient init_tinker_client() diff --git a/cookbook/client/tinker/upload_to_hub.py b/cookbook/client/tinker/upload_to_hub.py index da39cc021..d35822144 100644 --- a/cookbook/client/tinker/upload_to_hub.py +++ b/cookbook/client/tinker/upload_to_hub.py @@ -6,9 +6,11 @@ # # How it works: # 1. The server submits the upload as a background task and returns a -# request_id immediately, so the HTTP call never times out. -# 2. The client polls /upload_status/{request_id} every few seconds and -# blocks until the upload completes or raises on failure. +# Task_Envelope with a request_id immediately, so the HTTP call never times out. +# 2. The client's future layer long-polls /twinkle/retrieve_future and blocks +# until the upload reaches a terminal state, raising on failure. +# (`upload_to_hub` keeps its `poll_interval` / `async_upload` arguments for +# signature compatibility; both are deprecated and have no effect.) # # Prerequisites: # - Server must be running (see server.py / server_config.yaml) @@ -20,7 +22,8 @@ import os -from twinkle import get_logger, init_twinkle_client +from twinkle import get_logger +from twinkle import init_twinkle_client from twinkle_client.model import MultiLoraTransformersModel logger = get_logger() diff --git a/cookbook/client/twinkle/dpo.py b/cookbook/client/twinkle/dpo.py index b7fedcd83..869229211 100644 --- a/cookbook/client/twinkle/dpo.py +++ b/cookbook/client/twinkle/dpo.py @@ -6,24 +6,19 @@ # Step 1: Load environment variables from a .env file (e.g., API tokens) import dotenv -import os -from typing import Any, Dict, List - -dotenv.load_dotenv('.env') import numpy as np +import os import torch from peft import LoraConfig +from typing import Any, Dict, List from twinkle import get_logger -from twinkle.dataset import Dataset, DatasetMeta -from twinkle_client import init_twinkle_client +from twinkle import init_twinkle_client from twinkle.dataloader import DataLoader -from twinkle_client.model import MultiLoraTransformersModel -from twinkle.loss import DPOLoss -from twinkle.metric import DPOMetric +from twinkle.dataset import Dataset, DatasetMeta from twinkle.preprocessor import EmojiDPOProcessor -from twinkle.processor import InputProcessor +dotenv.load_dotenv('.env') logger = get_logger() # Configuration (direct values, not from env) @@ -68,11 +63,9 @@ def create_dpo_dataset(): dataset = Dataset(DatasetMeta(dataset_id, data_slice=range(100))) dataset.set_template('Qwen3_5Template', model_id=f'ms://{base_model}', max_length=max_length) dataset.map( - EmojiDPOProcessor, - init_args={ + EmojiDPOProcessor, init_args={ 'system': system_prompt, - } - ) + }) # DPO preprocessor returns {'positive': [...], 'negative': [...]} # batch_encode handles this format automatically dataset.encode() @@ -121,7 +114,7 @@ def train(): # Step 5: Configure the model # Create a multi-LoRA Transformers model pointing to the base model on ModelScope - model = MultiLoraTransformersModel(model_id=f'ms://{base_model}') + model = client.model(f'ms://{base_model}') # Define LoRA configuration: apply low-rank adapters to all linear layers lora_config = LoraConfig( @@ -162,7 +155,7 @@ def train(): optim_step = 0 max_steps = len(dataloader) logger.info(f'Starting LoRA DPO training: loss_type={loss_type}, beta={dpo_beta}, lr={learning_rate}') - logger.info(f'Using base model (disable_lora=True) as reference model') + logger.info('Using base model (disable_lora=True) as reference model') for batch in dataloader: # batch is List[Dict] with 'positive' and 'negative' keys @@ -199,7 +192,6 @@ def train(): # model.upload_to_hub( # checkpoint_dir=twinkle_path, # hub_model_id=hub_model_id, - # async_upload=False # ) # logger.info(f"Uploaded checkpoint to hub: {hub_model_id}") diff --git a/cookbook/client/twinkle/embedding.py b/cookbook/client/twinkle/embedding.py index 304a321fa..7c286a57a 100644 --- a/cookbook/client/twinkle/embedding.py +++ b/cookbook/client/twinkle/embedding.py @@ -24,18 +24,15 @@ # megatron mrope model gets valid positions (transformers derives them internally). import dotenv - -dotenv.load_dotenv('.env') - import os -from typing import Any, Dict, List - from peft import LoraConfig +from typing import Any, Dict, List -from twinkle import get_logger, init_twinkle_client +from twinkle import get_logger +from twinkle import init_twinkle_client from twinkle.template import Qwen3_5Template -from twinkle_client.model import MultiLoraTransformersModel +dotenv.load_dotenv('.env') logger = get_logger() # ========== Configuration ========== @@ -84,13 +81,13 @@ def build_minibatch(tokenizer) -> List[Dict[str, Any]]: def train(): # Step 1: connect to the running Twinkle server. - init_twinkle_client( + client = init_twinkle_client( base_url=os.environ.get('TWINKLE_SERVER_URL', 'http://localhost:8000'), api_key=os.environ.get('TWINKLE_SERVER_TOKEN', 'EMPTY_TOKEN'), ) # Step 2: build the client model with a fresh LoRA adapter. - model = MultiLoraTransformersModel(model_id=MODEL_ID) + model = client.model(MODEL_ID) model.add_adapter_to_model(ADAPTER_NAME, LoraConfig(target_modules='all-linear')) model.set_template('Qwen3_5Template', model_id=MODEL_ID) diff --git a/cookbook/client/twinkle/multi_modal.py b/cookbook/client/twinkle/multi_modal.py index 3a5a45b68..1cea9bb5e 100644 --- a/cookbook/client/twinkle/multi_modal.py +++ b/cookbook/client/twinkle/multi_modal.py @@ -6,22 +6,19 @@ # Step 1: Load environment variables from a .env file (e.g., API tokens) import dotenv -import os -from twinkle.data_format import Trajectory, Message -from twinkle.preprocessor import Preprocessor - -dotenv.load_dotenv('.env') import numpy as np +import os import torch from peft import LoraConfig from twinkle import get_logger -from twinkle.dataset import DatasetMeta -from twinkle_client import init_twinkle_client +from twinkle import init_twinkle_client +from twinkle.data_format import Message, Trajectory from twinkle.dataloader import DataLoader -from twinkle.dataset import LazyDataset -from twinkle_client.model import MultiLoraTransformersModel +from twinkle.dataset import DatasetMeta, LazyDataset +from twinkle.preprocessor import Preprocessor +dotenv.load_dotenv('.env') logger = get_logger() base_model = os.environ.get('TWINKLE_MODEL_ID', 'Qwen/Qwen3.5-4B') @@ -58,12 +55,10 @@ def __call__(self, rows): return rows def preprocess(self, row) -> Trajectory: - return Trajectory( - messages=[ - Message(role='user', content='Using LaTeX to perform OCR on the image.', images=[row['image']]), - Message(role='assistant', content=row['text']), - ] - ) + return Trajectory(messages=[ + Message(role='user', content='Using LaTeX to perform OCR on the image.', images=[row['image']]), + Message(role='assistant', content=row['text']), + ]) def train(): @@ -87,7 +82,7 @@ def train(): # Step 5: Configure the model # Create a multi-LoRA Transformers model pointing to the base model on ModelScope - model = MultiLoraTransformersModel(model_id=f'ms://{base_model}') + model = client.model(f'ms://{base_model}') # Define LoRA configuration: apply low-rank adapters to all linear layers lora_config = LoraConfig(target_modules='all-linear') @@ -160,7 +155,6 @@ def train(): # model.upload_to_hub( # checkpoint_dir=twinkle_path, # hub_model_id=hub_model_id, - # async_upload=False # ) # logger.info(f"Uploaded checkpoint to hub: {hub_model_id}") diff --git a/cookbook/client/twinkle/multi_turn_rollout.py b/cookbook/client/twinkle/multi_turn_rollout.py index 199dca385..942f572b6 100644 --- a/cookbook/client/twinkle/multi_turn_rollout.py +++ b/cookbook/client/twinkle/multi_turn_rollout.py @@ -21,21 +21,19 @@ # ``model.save(is_sampler=True)`` and point the sampler at the saved adapter; that # sync is intentionally omitted here to keep the rollout example focused. -import os -from typing import Any, Dict, List, Tuple - import dotenv +import os from peft import LoraConfig +from typing import Any, Dict, List, Tuple -from twinkle import get_logger, init_twinkle_client +from twinkle import get_logger +from twinkle import init_twinkle_client from twinkle.advantage import GRPOAdvantage from twinkle.data_format import SamplingParams from twinkle.template import Qwen3_5Template from twinkle_agentic.envs import EnvTool, OpenEnv from twinkle_agentic.tools.tool_manager import ToolManager -from twinkle_client.model import MultiLoraTransformersModel from twinkle_client.rollout import ClientMultiTurnRollout -from twinkle_client.sampler import vLLMSampler dotenv.load_dotenv('.env') @@ -45,7 +43,7 @@ BASE_MODEL = os.environ.get('TWINKLE_MODEL_ID', 'Qwen/Qwen3.5-4B') MODEL_ID = f'ms://{BASE_MODEL}' ADAPTER_NAME = 'default' -NUM_GENERATIONS = 2 # GRPO group size (rollout runs num_samples=1 per trajectory) +NUM_GENERATIONS = 2 # GRPO group size (rollout runs num_samples=1 per trajectory) BATCH_SIZE = 2 MAX_NEW_TOKENS = 512 MAX_TURNS = 4 @@ -81,7 +79,8 @@ Your goal is to win the game by getting as close to 21 as possible without going over. -Use the `play` tool to choose either `hit` or `stand`. Reason briefly before each action. Once the environment reports that the game is over, give a short final answer without calling another tool.""" +Use the `play` tool to choose either `hit` or `stand`. Reason briefly before each action. +Once the environment reports that the game is over, give a short final answer without calling another tool.""" def blackjack_action_mapper(tool_name: str, arguments: Dict[str, Any]) -> Dict[str, Any]: @@ -105,8 +104,7 @@ def create_env_tool(env: OpenEnv) -> EnvTool: def prepare_trajectories( - n_trajectories: int, -) -> Tuple[List[Dict[str, Any]], List[ToolManager], List[List[EnvTool]], List[OpenEnv]]: + n_trajectories: int, ) -> Tuple[List[Dict[str, Any]], List[ToolManager], List[List[EnvTool]], List[OpenEnv]]: """Create and reset one independent OpenEnv instance per trajectory.""" trajectories = [] tool_managers = [] @@ -127,10 +125,17 @@ def prepare_trajectories( tool_manager = ToolManager(env_tools) trajectories.append({ 'messages': [ - {'role': 'system', 'content': SYSTEM_PROMPT}, - {'role': 'user', 'content': initial_observation}, + { + 'role': 'system', + 'content': SYSTEM_PROMPT + }, + { + 'role': 'user', + 'content': initial_observation + }, ], - 'tools': tool_manager.tool_infos(), + 'tools': + tool_manager.tool_infos(), }) tool_managers.append(tool_manager) env_tools_list.append(env_tools) @@ -149,13 +154,13 @@ def extract_rewards(env_tools_list: List[List[EnvTool]]) -> List[float]: def train(): # Step 1: connect to the running Twinkle server. - init_twinkle_client( + client = init_twinkle_client( base_url=os.environ.get('TWINKLE_SERVER_URL', 'http://localhost:8000'), api_key=os.environ.get('TWINKLE_SERVER_TOKEN', 'EMPTY_TOKEN'), ) # Step 2: training model (GRPO), mirroring the ray-local example's config. - model = MultiLoraTransformersModel(model_id=MODEL_ID) + model = client.model(MODEL_ID) model.add_adapter_to_model(ADAPTER_NAME, LoraConfig(target_modules='all-linear', r=16, lora_alpha=32)) model.set_loss('GRPOLoss', epsilon=0.2) model.set_optimizer('Adam', lr=LEARNING_RATE) @@ -163,7 +168,7 @@ def train(): model.set_template('Qwen3_5Template', model_id=MODEL_ID, enable_thinking=False) # Step 3: client sampler (HTTP). - sampler = vLLMSampler(model_id=MODEL_ID) + sampler = client.sampler(MODEL_ID) sampler.set_template('Qwen3_5Template', model_id=MODEL_ID, enable_thinking=False) # Step 4: multi-turn rollout. Each call receives trajectory-bound ToolManagers diff --git a/cookbook/client/twinkle/sample.py b/cookbook/client/twinkle/sample.py index 5bbfd424c..7e999462b 100644 --- a/cookbook/client/twinkle/sample.py +++ b/cookbook/client/twinkle/sample.py @@ -10,16 +10,13 @@ # Step 1: Load environment variables from a .env file (e.g., API tokens) import dotenv - -dotenv.load_dotenv('.env') - import os from transformers import AutoTokenizer from twinkle import get_logger -from twinkle_client import init_twinkle_client -from twinkle_client.sampler import vLLMSampler +from twinkle import init_twinkle_client +dotenv.load_dotenv('.env') logger = get_logger() MODEL_ID = os.environ.get('TWINKLE_MODEL_ID', 'Qwen/Qwen3.5-4B') @@ -31,6 +28,7 @@ # Example: ADAPTER_URI = 'twinkle://20260301_142318-Qwen_Qwen3-4B-199d2cdb/weights/twinkle-lora-0' + def sample(): # Step 2: Initialize the Twinkle client to communicate with the remote server. client = init_twinkle_client( @@ -39,7 +37,7 @@ def sample(): ) # Step 3: Create the sampler client pointing to the model on the server - sampler = vLLMSampler(model_id=MODEL_ID) + sampler = client.sampler(MODEL_ID) # Step 4: Set the chat template so the sampler can encode Trajectory inputs sampler.set_template('Qwen3_5Template', model_id=MODEL_ID) diff --git a/cookbook/client/twinkle/self_cognition.py b/cookbook/client/twinkle/self_cognition.py index f5250b290..76f442359 100644 --- a/cookbook/client/twinkle/self_cognition.py +++ b/cookbook/client/twinkle/self_cognition.py @@ -6,19 +6,15 @@ # Step 1: Load environment variables from a .env file (e.g., API tokens) import dotenv - -dotenv.load_dotenv('.env') - import os from peft import LoraConfig from twinkle import get_logger -from twinkle.dataset import DatasetMeta from twinkle import init_twinkle_client from twinkle.dataloader import DataLoader -from twinkle.dataset import Dataset -from twinkle_client.model import MultiLoraTransformersModel +from twinkle.dataset import Dataset, DatasetMeta +dotenv.load_dotenv('.env') logger = get_logger() base_model = os.environ.get('TWINKLE_MODEL_ID', 'Qwen/Qwen3.5-4B') @@ -26,7 +22,6 @@ api_key = os.environ.get('TWINKLE_SERVER_TOKEN', 'EMPTY_TOKEN') save_dir = '/tmp/twinkle_sft_output' - # Step 2: Initialize the Twinkle client to communicate with the remote server. # - base_url: the address of the running Twinkle server # - api_key: authentication token (loaded from environment variable) @@ -74,7 +69,7 @@ def train(): # Step 5: Configure the model # Create a multi-LoRA Transformers model pointing to the base model on ModelScope - model = MultiLoraTransformersModel(model_id=f'ms://{base_model}') + model = client.model(f'ms://{base_model}') # Define LoRA configuration: apply low-rank adapters to all linear layers lora_config = LoraConfig(target_modules='all-linear') @@ -160,7 +155,6 @@ def train(): # model.upload_to_hub( # checkpoint_dir=twinkle_path, # hub_model_id=hub_model_id, - # async_upload=False # ) # logger.info(f"Uploaded checkpoint to hub: {hub_model_id}") diff --git a/cookbook/client/twinkle/short_math_grpo.py b/cookbook/client/twinkle/short_math_grpo.py index 1e90d38ce..21494ab06 100644 --- a/cookbook/client/twinkle/short_math_grpo.py +++ b/cookbook/client/twinkle/short_math_grpo.py @@ -20,30 +20,24 @@ # Requires both model and sampler services to be configured. import dotenv - -dotenv.load_dotenv('.env') - import gc import os import re -from peft import LoraConfig -from typing import List, Tuple, Dict, Any - import swanlab +from peft import LoraConfig +from typing import Any, Dict, List, Tuple from twinkle import get_logger -from twinkle.reward import GSM8KAccuracyReward -from twinkle.reward.base import Reward -from twinkle.advantage import GRPOAdvantage -from twinkle.dataset import DatasetMeta -from twinkle.metric import CompletionRewardMetric from twinkle import init_twinkle_client +from twinkle.advantage import GRPOAdvantage from twinkle.dataloader import DataLoader -from twinkle.dataset import Dataset +from twinkle.dataset import Dataset, DatasetMeta +from twinkle.metric import CompletionRewardMetric from twinkle.preprocessor.llm import GSM8KProcessor -from twinkle_client.model import MultiLoraTransformersModel -from twinkle_client.sampler import vLLMSampler +from twinkle.reward import GSM8KAccuracyReward +from twinkle.reward.base import Reward +dotenv.load_dotenv('.env') logger = get_logger() @@ -64,10 +58,7 @@ def __call__(self, trajectories: List[Dict[str, Any]], **kwargs) -> List[float]: completion = msg.get('content', '') break - has_answer = bool( - re.search(r'\\boxed\{[^}]+\}', completion) - or re.search(r'####\s*[\-\d,\.]+', completion) - ) + has_answer = bool(re.search(r'\\boxed\{[^}]+\}', completion) or re.search(r'####\s*[\-\d,\.]+', completion)) if not has_answer: rewards.append(0.0) @@ -79,6 +70,7 @@ def __call__(self, trajectories: List[Dict[str, Any]], **kwargs) -> List[float]: rewards.append(max(0.0, 1.0 - (length - 200) / 3000)) return rewards + # ========== Configuration ========== BASE_MODEL = os.environ.get('TWINKLE_MODEL_ID', 'Qwen/Qwen3.5-4B') MODEL_ID = f'ms://{BASE_MODEL}' @@ -96,20 +88,20 @@ def __call__(self, trajectories: List[Dict[str, Any]], **kwargs) -> List[float]: SWANLAB_PROJECT = 'twinkle-grpo' SWANLAB_EXPERIMENT_NAME = 'short-math-grpo' - SYSTEM_PROMPT = ('You are a helpful math assistant. Solve the problem with minimal but correct reasoning ' 'and put your final answer within \\boxed{}.') + def create_gsm8k_dataset(): - dataset = Dataset(DatasetMeta('ms://modelscope/gsm8k', subset_name='main', split='train', data_slice=range(DATA_NUM))) + dataset = Dataset( + DatasetMeta('ms://modelscope/gsm8k', subset_name='main', split='train', data_slice=range(DATA_NUM))) dataset.set_template('Qwen3_5Template', model_id=MODEL_ID, max_length=2048, enable_thinking=False) dataset.map(GSM8KProcessor(system=SYSTEM_PROMPT)) dataset.encode(add_generation_prompt=True) return dataset -def compute_rewards( - trajectories: List[Dict[str, Any]], -) -> Tuple[List[float], List[float], List[float]]: + +def compute_rewards(trajectories: List[Dict[str, Any]], ) -> Tuple[List[float], List[float], List[float]]: accuracy_reward_fn = GSM8KAccuracyReward() brevity_reward_fn = GSM8KBrevityReward() @@ -151,7 +143,7 @@ def train(): dataloader = DataLoader(dataset=dataset, batch_size=BATCH_SIZE, num_workers=0) # Step 3: Configure the training model - model = MultiLoraTransformersModel(model_id=MODEL_ID) + model = client.model(MODEL_ID) lora_config = LoraConfig( target_modules='all-linear', @@ -182,7 +174,7 @@ def train(): model.set_template('Qwen3_5Template', model_id=MODEL_ID) # Step 4: Configure the sampler - sampler = vLLMSampler(model_id=MODEL_ID) + sampler = client.sampler(MODEL_ID) sampler.set_template('Qwen3_5Template', model_id=MODEL_ID) # Step 5: Setup metrics and advantage function @@ -241,9 +233,7 @@ def train(): # ========== 3. Compute rewards ========== - total_rewards, brevity_rewards, accuracy_rewards = compute_rewards( - all_input_data - ) + total_rewards, brevity_rewards, accuracy_rewards = compute_rewards(all_input_data) metrics.accumulate( completion_lengths=all_completion_lengths, rewards={ @@ -253,7 +243,6 @@ def train(): }, ) - # ========== 4. Compute advantages ========== advantages = advantage_fn( total_rewards, diff --git a/cookbook/client/twinkle/upload_to_hub.py b/cookbook/client/twinkle/upload_to_hub.py index 2c780e256..7a2bc2fd2 100644 --- a/cookbook/client/twinkle/upload_to_hub.py +++ b/cookbook/client/twinkle/upload_to_hub.py @@ -6,23 +6,23 @@ # # How it works: # 1. The server submits the upload as a background task and returns a -# request_id immediately, so the HTTP call never times out. -# 2. The client polls /upload_status/{request_id} every few seconds and -# blocks until the upload completes or raises on failure. +# Task_Envelope with a request_id immediately, so the HTTP call never times out. +# 2. The client's future layer long-polls /twinkle/retrieve_future and blocks +# until the upload reaches a terminal state, raising on failure. +# (`upload_to_hub` keeps its `poll_interval` / `async_upload` arguments for +# signature compatibility; both are deprecated and have no effect.) # # Prerequisites: # - Server must be running (see server.py / server_config.yaml) # - A ModelScope API token with write access to the target repository import dotenv - -dotenv.load_dotenv('.env') - import os -from twinkle import get_logger, init_twinkle_client -from twinkle_client.model import MultiLoraTransformersModel +from twinkle import get_logger +from twinkle import init_twinkle_client +dotenv.load_dotenv('.env') logger = get_logger() # ── Configuration ───────────────────────────────────────────────────────────── @@ -43,10 +43,10 @@ def upload(): # Step 1: Initialize the Twinkle client - init_twinkle_client(base_url=base_url, api_key=api_key) + client = init_twinkle_client(base_url=base_url, api_key=api_key) # Step 2: Create the model client (no training state needed for upload) - model = MultiLoraTransformersModel(model_id=f'ms://{base_model}') + model = client.model(f'ms://{base_model}') # Step 3: Upload checkpoint to ModelScope Hub. # The client polls for completion automatically; progress is printed to stdout. diff --git a/docs/source_en/Usage Guide/Embedding-Training.md b/docs/source_en/Usage Guide/Embedding-Training.md index d0517e189..4198f6991 100644 --- a/docs/source_en/Usage Guide/Embedding-Training.md +++ b/docs/source_en/Usage Guide/Embedding-Training.md @@ -128,7 +128,7 @@ The difference from the bare library is that the client passes **class-name stri ```python from peft import LoraConfig -from twinkle_client import init_twinkle_client +from twinkle import init_twinkle_client from twinkle_client.model import MultiLoraTransformersModel # --- Connect to the running Twinkle server --- diff --git a/docs/source_en/Usage Guide/Introduction-with-Qwen3.5.md b/docs/source_en/Usage Guide/Introduction-with-Qwen3.5.md index 22dc07dfc..e61ed9197 100644 --- a/docs/source_en/Usage Guide/Introduction-with-Qwen3.5.md +++ b/docs/source_en/Usage Guide/Introduction-with-Qwen3.5.md @@ -368,7 +368,7 @@ from peft import LoraConfig from twinkle import get_logger from twinkle.dataset import DatasetMeta -from twinkle_client import init_twinkle_client +from twinkle import init_twinkle_client from twinkle_client.dataloader import DataLoader from twinkle_client.dataset import Dataset from twinkle_client.model import MultiLoraTransformersModel @@ -458,7 +458,7 @@ from twinkle import init_tinker_client from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset, DatasetMeta from twinkle.preprocessor import SelfCognitionProcessor -from twinkle.server.common import input_feature_to_datum +from twinkle.server.model.tinker_datum import input_feature_to_datum # Initialize Tinker client (must be called before importing ServiceClient) init_tinker_client() diff --git a/docs/source_en/Usage Guide/Quick-Start.md b/docs/source_en/Usage Guide/Quick-Start.md index 69953a9a1..0b722de00 100644 --- a/docs/source_en/Usage Guide/Quick-Start.md +++ b/docs/source_en/Usage Guide/Quick-Start.md @@ -514,7 +514,7 @@ from twinkle import get_logger from twinkle.advantage import GRPOAdvantage from twinkle.dataset import DatasetMeta from twinkle.metric import CompletionRewardMetric -from twinkle_client import init_twinkle_client +from twinkle import init_twinkle_client from twinkle_client.dataloader import DataLoader from twinkle_client.dataset import Dataset from twinkle_client.model import MultiLoraTransformersModel @@ -761,7 +761,7 @@ from tinker import ServiceClient from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset, DatasetMeta from twinkle.preprocessor import SelfCognitionProcessor -from twinkle.server.common import input_feature_to_datum +from twinkle.server.model.tinker_datum import input_feature_to_datum # The base model to fine-tune / evaluate base_model = 'ms://Qwen/Qwen3.5-4B' diff --git a/docs/source_en/Usage Guide/Server and Client/Server.md b/docs/source_en/Usage Guide/Server and Client/Server.md index f31209b07..771e56249 100644 --- a/docs/source_en/Usage Guide/Server and Client/Server.md +++ b/docs/source_en/Usage Guide/Server and Client/Server.md @@ -183,10 +183,10 @@ telemetry: otlp_endpoint: http://localhost:4317 # Persistence: storage backend for ServerState (sessions, models, futures, etc.) -# mode: memory | file | redis +# mode: memory | redis persistence: - mode: file - file_path: /tmp/twinkle_state.json + mode: redis + redis_url: redis://localhost:6379/0 # Application list: Each entry defines a service component deployed on the Server applications: @@ -350,7 +350,7 @@ The difference from the Megatron backend is only in the `backend` parameter of t | `proxy_location` | HTTP proxy location (`EveryNode` or `HeadOnly`) | | `http_options` | HTTP listener config (`host`, `port`) | | `telemetry` | Observability config (`enabled`, `otlp_endpoint`) | -| `persistence` | State persistence config (`mode`, `file_path`, `redis_url`) | +| `persistence` | State persistence config (`mode`, `redis_url`) | | `applications` | Application component list | > The config file uses strict validation (`extra='forbid'`). Any misspelled field name will be rejected before startup. Use `twinkle-server check-config -c xxx.yaml` to detect errors early. @@ -418,8 +418,7 @@ Storage backend for ServerState (sessions, models, futures, etc.). | Field | Type | Default | Description | |-------|------|---------|-------------| -| `mode` | str | `memory` | `memory` / `file` / `redis` | -| `file_path` | str | — | Required for `file` mode, JSON file path | +| `mode` | str | `memory` | `memory` / `redis` | | `redis_url` | str | — | Required for `redis` mode, e.g. `redis://localhost:6379` | | `key_prefix` | str | `""` | Optional global key prefix | @@ -450,3 +449,33 @@ twinkle-server check-config -c server_config.yaml | `use_megatron: false` | `backend: transformers` | Additionally, this refactor introduces two new top-level fields — `telemetry` and `persistence` — which did not exist before. Add them as needed. + +## Execution time bounds + +Every backend call has a finite time bound. `T` is the effective task execution +timeout: it equals `execution_timeout`, or `3600s` when that setting is `0`. +`asyncio.wait_for` uses `T`. The Ray wait uses `R`, which is a method's explicit +constant timeout when present and otherwise `T`. The default `T` is `1800s`. + +Two distinct bounds follow, and they must not be collapsed into one number: + +| Bound | Expression | Meaning | +|-------|------------|---------| +| Record-terminal bound | `queue_timeout + T` | After this, a task's future record is guaranteed to be in a terminal state (`completed`/`failed`). Use it for alerting thresholds and client polling total-timeout. | +| Resource-release bound | `Collect_Width × R` from execution start, or `queue_timeout + Collect_Width × R` from submission | After this, the executor thread and the in-flight model-actor call for that task are guaranteed to have finished. Use it for capacity planning. | + +`Collect_Width = len(self._actors) = world_size = tp × pp × dp` — the number of +futures each `remote_function` collection waits on per call. Evidence: +`LazyCollect._get_result` iterates `self._futures`, which come from +`_get_workers(self._actors, execute)` (`infra/__init__.py`), covering every actor — +not just the data-parallel width. On a `tp=8` deployment the execution-start +resource-release bound is therefore `8 × R`, not `R`. + +After the task record becomes terminal, the per-replica Admission_Gate can remain +closed for at most `max(0, Collect_Width × R − T)`: the record is already terminal, +but a leaked executor thread may still hold the gate until its `ray.get` returns or +raises. During that window newly arriving tasks fail fast with a `server`/503 error. + +Each persisted future stores its immutable `absolute_deadline` when it is created. +Cleanup therefore reaches the same decision regardless of which deployment process +holds the cleanup lease. diff --git a/docs/source_en/Usage Guide/Server and Client/Tinker-Compatible-Client.md b/docs/source_en/Usage Guide/Server and Client/Tinker-Compatible-Client.md index a530174e3..df9d1169e 100644 --- a/docs/source_en/Usage Guide/Server and Client/Tinker-Compatible-Client.md +++ b/docs/source_en/Usage Guide/Server and Client/Tinker-Compatible-Client.md @@ -45,7 +45,7 @@ from twinkle import init_tinker_client from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset, DatasetMeta from twinkle.preprocessor import SelfCognitionProcessor -from twinkle.server.common import input_feature_to_datum +from twinkle.server.model.tinker_datum import input_feature_to_datum # Step 1: Initialize Tinker client before importing ServiceClient init_tinker_client() diff --git a/docs/source_en/Usage Guide/Server and Client/Twinkle-Client.md b/docs/source_en/Usage Guide/Server and Client/Twinkle-Client.md index 598f215ef..3f3cd365f 100644 --- a/docs/source_en/Usage Guide/Server and Client/Twinkle-Client.md +++ b/docs/source_en/Usage Guide/Server and Client/Twinkle-Client.md @@ -5,7 +5,7 @@ Twinkle Client is the native client, designed with the philosophy: **Change `fro ## Initialization ```python -from twinkle_client import init_twinkle_client +from twinkle import init_twinkle_client # Initialize client, connect to Twinkle Server client = init_twinkle_client( @@ -38,7 +38,7 @@ latest_path = client.get_latest_checkpoint_path(run_id='xxx') ## Migrating from Local Code to Remote -Migration is very simple, just replace the import path from `twinkle` to `twinkle_client`: +Keep the data-processing and training-loop code, then create the remote model through `client.model(...)` after initializing the remote client: ```python # Local training code (original) @@ -50,10 +50,13 @@ from twinkle.model import MultiLoraTransformersModel # DataLoader and Dataset can be imported from either local twinkle or remote twinkle_client from twinkle.dataloader import DataLoader # or: from twinkle_client.dataloader import DataLoader from twinkle.dataset import Dataset # or: from twinkle_client.dataset import Dataset -from twinkle_client.model import MultiLoraTransformersModel +from twinkle import init_twinkle_client + +client = init_twinkle_client(base_url=base_url, api_key=api_key) +model = client.model(f'ms://{base_model}') ``` -Training loops, data processing, and other logic do not need any modifications. +Training loops and data processing do not need any modifications. Prefer `client.model(...)` over constructing `MultiLoraTransformersModel(...)` directly so the model wrapper explicitly reuses the current client's transport, session, and authentication context. ## Complete Training Example (Transformers Backend) @@ -64,12 +67,11 @@ dotenv.load_dotenv('.env') from peft import LoraConfig from twinkle import get_logger from twinkle.dataset import DatasetMeta -from twinkle_client import init_twinkle_client +from twinkle import init_twinkle_client # DataLoader and Dataset can be imported from either local twinkle or remote twinkle_client from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset -from twinkle_client.model import MultiLoraTransformersModel logger = get_logger() @@ -118,8 +120,8 @@ dataset.encode(batched=True) # Create DataLoader dataloader = DataLoader(dataset=dataset, batch_size=4) -# Step 4: Configure model -model = MultiLoraTransformersModel(model_id=f'ms://{base_model}') +# Step 4: Create the remote model bound to the current client +model = client.model(f'ms://{base_model}') # Configure LoRA: apply low-rank adapters to all linear layers lora_config = LoraConfig(target_modules='all-linear') @@ -171,12 +173,14 @@ for epoch in range(3): logger.info(f'Saved checkpoint: {twinkle_path}') # Step 8: Upload to ModelScope Hub (optional) +# The server always uploads in the background; this call waits through the future layer +# until upload completion or failure. Do not pass async_upload or poll_interval: they +# remain only for old-call compatibility, are deprecated, and have no effect. # YOUR_USER_NAME = "your_username" # hub_model_id = f'{YOUR_USER_NAME}/twinkle-self-cognition' # model.upload_to_hub( # checkpoint_dir=twinkle_path, # hub_model_id=hub_model_id, -# async_upload=False # ) ``` @@ -223,17 +227,16 @@ from twinkle.advantage import GRPOAdvantage from twinkle.data_format import SamplingParams from twinkle.template import Qwen3_5Template from twinkle_agentic.tools.tool_manager import ToolManager -from twinkle_client.model import MultiLoraTransformersModel from twinkle_client.rollout import ClientMultiTurnRollout from twinkle_client.sampler import vLLMSampler MODEL_ID = 'ms://Qwen/Qwen3.5-4B' NUM_GENERATIONS = 2 # GRPO group size (rollout samples num_samples=1 per trajectory) -init_twinkle_client(base_url='http://127.0.0.1:8000', api_key='EMPTY_TOKEN') +client = init_twinkle_client(base_url='http://127.0.0.1:8000', api_key='EMPTY_TOKEN') -# Training model (GRPO) -model = MultiLoraTransformersModel(model_id=MODEL_ID) +# Training model (GRPO): bind the current transport, session, and authentication context +model = client.model(MODEL_ID) model.add_adapter_to_model('default', LoraConfig(target_modules='all-linear', r=16, lora_alpha=32)) model.set_loss('GRPOLoss', epsilon=0.2) model.set_optimizer('Adam', lr=1e-5) diff --git a/docs/source_en/Usage Guide/Train-as-a-Service.md b/docs/source_en/Usage Guide/Train-as-a-Service.md index 286b52001..c07e072e1 100644 --- a/docs/source_en/Usage Guide/Train-as-a-Service.md +++ b/docs/source_en/Usage Guide/Train-as-a-Service.md @@ -24,11 +24,11 @@ Sample code: import os from tqdm import tqdm from tinker import types -from twinkle_client import init_tinker_client +from twinkle import init_tinker_client from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset, DatasetMeta from twinkle.preprocessor import SelfCognitionProcessor -from twinkle.server.common import input_feature_to_datum +from twinkle.server.model.tinker_datum import input_feature_to_datum base_model = 'ms://Qwen/Qwen3.8-27B' base_url='https://www.modelscope.cn/twinkle' diff --git "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/Embedding\350\256\255\347\273\203.md" "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/Embedding\350\256\255\347\273\203.md" index 7e3829ef7..b3b5d3084 100644 --- "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/Embedding\350\256\255\347\273\203.md" +++ "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/Embedding\350\256\255\347\273\203.md" @@ -128,7 +128,7 @@ set_processor('InputProcessor') ```python from peft import LoraConfig -from twinkle_client import init_twinkle_client +from twinkle import init_twinkle_client from twinkle_client.model import MultiLoraTransformersModel # --- Connect to the running Twinkle server --- diff --git "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/Qwen3.5\346\234\200\344\275\263\345\256\236\350\267\265.md" "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/Qwen3.5\346\234\200\344\275\263\345\256\236\350\267\265.md" index c80b1e15d..d6fc45d2b 100644 --- "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/Qwen3.5\346\234\200\344\275\263\345\256\236\350\267\265.md" +++ "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/Qwen3.5\346\234\200\344\275\263\345\256\236\350\267\265.md" @@ -368,7 +368,7 @@ from peft import LoraConfig from twinkle import get_logger from twinkle.dataset import DatasetMeta -from twinkle_client import init_twinkle_client +from twinkle import init_twinkle_client from twinkle_client.dataloader import DataLoader from twinkle_client.dataset import Dataset from twinkle_client.model import MultiLoraTransformersModel @@ -458,7 +458,7 @@ from twinkle import init_tinker_client from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset, DatasetMeta from twinkle.preprocessor import SelfCognitionProcessor -from twinkle.server.common import input_feature_to_datum +from twinkle.server.model.tinker_datum import input_feature_to_datum # 初始化 Tinker 客户端(必须在导入 ServiceClient 之前) init_tinker_client() diff --git "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\345\277\253\351\200\237\345\274\200\345\247\213.md" "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\345\277\253\351\200\237\345\274\200\345\247\213.md" index 0cfaa4bba..2efa4cc8e 100644 --- "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\345\277\253\351\200\237\345\274\200\345\247\213.md" +++ "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\345\277\253\351\200\237\345\274\200\345\247\213.md" @@ -515,7 +515,7 @@ from twinkle import get_logger from twinkle.advantage import GRPOAdvantage from twinkle.dataset import DatasetMeta from twinkle.metric import CompletionRewardMetric -from twinkle_client import init_twinkle_client +from twinkle import init_twinkle_client from twinkle_client.dataloader import DataLoader from twinkle_client.dataset import Dataset from twinkle_client.model import MultiLoraTransformersModel @@ -762,7 +762,7 @@ from tinker import ServiceClient from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset, DatasetMeta from twinkle.preprocessor import SelfCognitionProcessor -from twinkle.server.common import input_feature_to_datum +from twinkle.server.model.tinker_datum import input_feature_to_datum # The base model to fine-tune / evaluate base_model = 'Qwen/Qwen3.5-4B' diff --git "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/Tinker\345\205\274\345\256\271\345\256\242\346\210\267\347\253\257.md" "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/Tinker\345\205\274\345\256\271\345\256\242\346\210\267\347\253\257.md" index a1f7e064f..8690ff055 100644 --- "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/Tinker\345\205\274\345\256\271\345\256\242\346\210\267\347\253\257.md" +++ "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/Tinker\345\205\274\345\256\271\345\256\242\346\210\267\347\253\257.md" @@ -45,7 +45,7 @@ from twinkle import init_tinker_client from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset, DatasetMeta from twinkle.preprocessor import SelfCognitionProcessor -from twinkle.server.common import input_feature_to_datum +from twinkle.server.model.tinker_datum import input_feature_to_datum # Step 1: 在导入 ServiceClient 之前,先初始化 Tinker 客户端 init_tinker_client() diff --git "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/Twinkle\345\256\242\346\210\267\347\253\257.md" "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/Twinkle\345\256\242\346\210\267\347\253\257.md" index 668156d75..ed60681df 100644 --- "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/Twinkle\345\256\242\346\210\267\347\253\257.md" +++ "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/Twinkle\345\256\242\346\210\267\347\253\257.md" @@ -5,7 +5,7 @@ Twinkle Client 是原生客户端,设计理念是:**将 `from twinkle import ## 初始化 ```python -from twinkle_client import init_twinkle_client +from twinkle import init_twinkle_client # 初始化客户端,连接到 Twinkle Server client = init_twinkle_client( @@ -38,7 +38,7 @@ latest_path = client.get_latest_checkpoint_path(run_id='xxx') ## 从本地代码迁移到远端 -迁移非常简单,只需将 import 路径从 `twinkle` 替换为 `twinkle_client`: +迁移时保留数据处理和训练循环;初始化远端客户端后,通过 `client.model(...)` 创建显式绑定该客户端的远端模型: ```python # 本地训练代码(原始) @@ -50,10 +50,13 @@ from twinkle.model import MultiLoraTransformersModel # DataLoader 和 Dataset 使用本地 twinkle 或远端 twinkle_client 均可 from twinkle.dataloader import DataLoader # 或 from twinkle_client.dataloader import DataLoader from twinkle.dataset import Dataset # 或 from twinkle_client.dataset import Dataset -from twinkle_client.model import MultiLoraTransformersModel +from twinkle import init_twinkle_client + +client = init_twinkle_client(base_url=base_url, api_key=api_key) +model = client.model(f'ms://{base_model}') ``` -训练循环、数据处理等逻辑完全不需要修改。 +训练循环、数据处理等逻辑完全不需要修改;推荐使用 `client.model(...)`,而不是直接构造 `MultiLoraTransformersModel(...)`,以确保模型包装器显式复用当前客户端的 transport、会话与认证上下文。 ## 完整训练示例(Transformers 后端) @@ -64,12 +67,11 @@ dotenv.load_dotenv('.env') from peft import LoraConfig from twinkle import get_logger from twinkle.dataset import DatasetMeta -from twinkle_client import init_twinkle_client +from twinkle import init_twinkle_client # DataLoader 和 Dataset 使用本地 twinkle 或远端 twinkle_client 均可 from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset -from twinkle_client.model import MultiLoraTransformersModel logger = get_logger() @@ -118,8 +120,8 @@ dataset.encode(batched=True) # 创建 DataLoader dataloader = DataLoader(dataset=dataset, batch_size=4) -# Step 4: 配置模型 -model = MultiLoraTransformersModel(model_id=f'ms://{base_model}') +# Step 4: 通过当前 client 创建并绑定远端模型 +model = client.model(f'ms://{base_model}') # 配置 LoRA:对所有线性层应用低秩适配器 lora_config = LoraConfig(target_modules='all-linear') @@ -171,12 +173,13 @@ for epoch in range(3): logger.info(f'Saved checkpoint: {twinkle_path}') # Step 8: 上传到 ModelScope Hub(可选) +# 服务端始终在后台执行上传;当前调用会通过 future layer 等待上传完成或抛出失败。 +# 不要传 async_upload 或 poll_interval:两者仅为兼容旧调用保留,已废弃且无效果。 # YOUR_USER_NAME = "your_username" # hub_model_id = f'{YOUR_USER_NAME}/twinkle-self-cognition' # model.upload_to_hub( # checkpoint_dir=twinkle_path, # hub_model_id=hub_model_id, -# async_upload=False # ) ``` @@ -223,17 +226,16 @@ from twinkle.advantage import GRPOAdvantage from twinkle.data_format import SamplingParams from twinkle.template import Qwen3_5Template from twinkle_agentic.tools.tool_manager import ToolManager -from twinkle_client.model import MultiLoraTransformersModel from twinkle_client.rollout import ClientMultiTurnRollout from twinkle_client.sampler import vLLMSampler MODEL_ID = 'ms://Qwen/Qwen3.5-4B' NUM_GENERATIONS = 2 # GRPO group size(rollout 每条采样 num_samples=1) -init_twinkle_client(base_url='http://127.0.0.1:8000', api_key='EMPTY_TOKEN') +client = init_twinkle_client(base_url='http://127.0.0.1:8000', api_key='EMPTY_TOKEN') -# 训练模型(GRPO) -model = MultiLoraTransformersModel(model_id=MODEL_ID) +# 训练模型(GRPO):通过当前 client 显式绑定 transport、会话与认证上下文 +model = client.model(MODEL_ID) model.add_adapter_to_model('default', LoraConfig(target_modules='all-linear', r=16, lora_alpha=32)) model.set_loss('GRPOLoss', epsilon=0.2) model.set_optimizer('Adam', lr=1e-5) diff --git "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/\346\234\215\345\212\241\347\253\257.md" "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/\346\234\215\345\212\241\347\253\257.md" index a4df7a2da..02467834d 100644 --- "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/\346\234\215\345\212\241\347\253\257.md" +++ "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\346\234\215\345\212\241\347\253\257\345\222\214\345\256\242\346\210\267\347\253\257/\346\234\215\345\212\241\347\253\257.md" @@ -183,10 +183,10 @@ telemetry: otlp_endpoint: http://localhost:4317 # 持久化:ServerState 的存储后端(sessions、models、futures 等) -# mode: memory | file | redis +# mode: memory | redis persistence: - mode: file - file_path: /tmp/twinkle_state.json + mode: redis + redis_url: redis://localhost:6379/0 # 应用列表:每个条目定义一个部署在 Server 上的服务组件 applications: @@ -350,7 +350,7 @@ Transformers 后端与 Megatron 后端的区别仅在 Model 服务的 `backend` | `proxy_location` | HTTP 代理位置(`EveryNode` 或 `HeadOnly`) | | `http_options` | HTTP 监听配置(`host`、`port`) | | `telemetry` | 可观测性配置(`enabled`、`otlp_endpoint`) | -| `persistence` | 状态持久化配置(`mode`、`file_path`、`redis_url`) | +| `persistence` | 状态持久化配置(`mode`、`redis_url`) | | `applications` | 应用组件列表 | > 配置文件启用了严格校验(`extra='forbid'`),任何拼写错误的字段名都会在启动前报错。可使用 `twinkle-server check-config -c xxx.yaml` 提前检测。 @@ -418,8 +418,7 @@ ServerState(sessions、models、futures 等)的存储后端。 | 字段 | 类型 | 默认值 | 说明 | |------|------|--------|------| -| `mode` | str | `memory` | `memory` / `file` / `redis` | -| `file_path` | str | — | `file` 模式必填,JSON 文件路径 | +| `mode` | str | `memory` | `memory` / `redis` | | `redis_url` | str | — | `redis` 模式必填,如 `redis://localhost:6379` | | `key_prefix` | str | `""` | 可选的全局 key 前缀 | @@ -450,3 +449,20 @@ twinkle-server check-config -c server_config.yaml | `use_megatron: false` | `backend: transformers` | 此外本次重构新增了 `telemetry` 和 `persistence` 两个顶层字段(旧版本中不存在),可按需添加。 + +## 执行时间上界 + +每一次 backend 调用都存在有限时间上界。`T` 是任务的有效 execution timeout:等于 task-queue 配置中的 `execution_timeout`;配置为 `0` 时取 `3600` 秒。`asyncio.wait_for` 使用 `T`。Ray 等待使用 `R`:方法显式声明 timeout 时取该常量,否则取 `T`。`T` 的默认值为 `1800` 秒。 + +由此派生出两个**不同**的上界,二者不得合成一个数: + +| 上界 | 表达式 | 含义 | +|------|--------|------| +| 记录终态上界 | `queue_timeout + T` | 超过它后,任务的 future 记录必处于终态(`completed`/`failed`)。用于设置告警阈值与客户端轮询总超时。 | +| 资源释放上界 | 从执行开始为 `Collect_Width × R`;从提交开始为 `queue_timeout + Collect_Width × R` | 超过它后,该任务占用的 executor 线程与 model actor 在飞调用必已结束。用于容量规划。 | + +`Collect_Width = len(self._actors) = world_size = tp × pp × dp`——即每次 `remote_function` 结果收集所等待的 future 个数。证据:`LazyCollect._get_result` 遍历的 `self._futures` 来自 `_get_workers(self._actors, execute)`(`infra/__init__.py`),覆盖全部 actor,而非 data-parallel 宽度。因此在 `tp=8` 的部署上,从执行开始的资源释放上界是 `8 × R` 而非 `R`。 + +任务记录进入终态后,per-replica 准入闸门额外保持关闭的最长时长为 `max(0, Collect_Width × R − T)`:此时记录已是终态,但泄漏的 executor 线程可能仍持有闸门,直到其 `ray.get` 返回或抛出。在该窗口内新到达的任务会以 `server`/503 错误快速失败。 + +每条持久化 future 在创建时写入不可变的 `absolute_deadline`,因此无论哪个 deployment 进程持有 cleanup lease,清理结果都由任务自身契约决定。 diff --git "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\350\256\255\347\273\203\346\234\215\345\212\241.md" "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\350\256\255\347\273\203\346\234\215\345\212\241.md" index c57cf066b..fb515944a 100644 --- "a/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\350\256\255\347\273\203\346\234\215\345\212\241.md" +++ "b/docs/source_zh/\344\275\277\347\224\250\346\214\207\345\274\225/\350\256\255\347\273\203\346\234\215\345\212\241.md" @@ -25,11 +25,11 @@ import os from tqdm import tqdm from tinker import types -from twinkle_client import init_tinker_client +from twinkle import init_tinker_client from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset, DatasetMeta from twinkle.preprocessor import SelfCognitionProcessor -from twinkle.server.common import input_feature_to_datum +from twinkle.server.model.tinker_datum import input_feature_to_datum base_model = 'ms://Qwen/Qwen3.8-27B' base_url='https://www.modelscope.cn/twinkle' diff --git a/notebook/dpo.ipynb b/notebook/dpo.ipynb index d2a13cfa5..9d575c099 100644 --- a/notebook/dpo.ipynb +++ b/notebook/dpo.ipynb @@ -140,11 +140,11 @@ "\n", "from tinker import types\n", "from getpass import getpass\n", - "from twinkle import init_tinker_client, get_logger\n", + "from twinkle import get_logger\nfrom twinkle import init_tinker_client\n", "from twinkle.dataset import Dataset, DatasetMeta, LazyDataset\n", "from twinkle.dataloader import DataLoader\n", "from twinkle.preprocessor import EmojiDPOProcessor\n", - "from twinkle.server.common import input_feature_to_datum\n", + "from twinkle.server.model.tinker_datum import input_feature_to_datum\n", "\n", "logger = get_logger()\n", "\n", @@ -478,7 +478,7 @@ "source": [ "# 推理示例(使用线上服务,无需本地 GPU)\n", "from tinker import types\n", - "from twinkle import init_tinker_client, get_logger\n", + "from twinkle import get_logger\nfrom twinkle import init_tinker_client\n", "from twinkle.data_format import Message, Trajectory\n", "from twinkle.template import Template\n", "\n", diff --git a/notebook/multi_modal.ipynb b/notebook/multi_modal.ipynb index 9bb52467f..17221cf98 100644 --- a/notebook/multi_modal.ipynb +++ b/notebook/multi_modal.ipynb @@ -124,7 +124,7 @@ "from twinkle.data_format import Trajectory, Message\n", "from twinkle.preprocessor import Preprocessor\n", "from twinkle.dataset import DatasetMeta\n", - "from twinkle_client import init_twinkle_client\n", + "from twinkle import init_twinkle_client\n", "from twinkle.dataloader import DataLoader\n", "from twinkle.dataset import LazyDataset\n", "from twinkle_client.model import MultiLoraTransformersModel\n", @@ -407,7 +407,7 @@ "source": [ "# 推理示例(使用线上服务,无需本地 GPU)\n", "from tinker import types\n", - "from twinkle import init_tinker_client, get_logger\n", + "from twinkle import get_logger\nfrom twinkle import init_tinker_client\n", "from twinkle.data_format import Message, Trajectory\n", "from twinkle.template import Qwen3_5Template\n", "\n", diff --git a/notebook/sample.ipynb b/notebook/sample.ipynb index 5c3c18b1b..22d56cb00 100644 --- a/notebook/sample.ipynb +++ b/notebook/sample.ipynb @@ -102,7 +102,7 @@ "source": [ "from tinker import types\n", "from getpass import getpass\n", - "from twinkle import init_tinker_client, get_logger\n", + "from twinkle import get_logger\nfrom twinkle import init_tinker_client\n", "from twinkle.data_format import Message, Trajectory\n", "from twinkle.template import Template\n", "\n", diff --git a/notebook/self_cognition.ipynb b/notebook/self_cognition.ipynb index 5e3383f13..d6c3b1090 100644 --- a/notebook/self_cognition.ipynb +++ b/notebook/self_cognition.ipynb @@ -124,7 +124,7 @@ "from twinkle.dataloader import DataLoader\n", "from twinkle.dataset import Dataset, DatasetMeta\n", "from twinkle.preprocessor import SelfCognitionProcessor\n", - "from twinkle.server.common import input_feature_to_datum\n", + "from twinkle.server.model.tinker_datum import input_feature_to_datum\n", "from getpass import getpass" ] }, diff --git a/notebook/short_math_grpo.ipynb b/notebook/short_math_grpo.ipynb index a05766789..2d40fe912 100644 --- a/notebook/short_math_grpo.ipynb +++ b/notebook/short_math_grpo.ipynb @@ -136,7 +136,8 @@ "from typing import List, Tuple, Dict, Any\n", "\n", "from getpass import getpass\n", - "from twinkle import get_logger, init_twinkle_client\n", + "from twinkle import get_logger\n", + "from twinkle import init_twinkle_client\n", "from twinkle.reward.base import Reward\n", "from twinkle.advantage import GRPOAdvantage\n", "from twinkle.dataset import DatasetMeta, Dataset\n", @@ -573,7 +574,8 @@ "source": [ "# 推理示例(使用线上服务,无需本地 GPU)\n", "from tinker import types\n", - "from twinkle import init_tinker_client, get_logger\n", + "from twinkle import get_logger\n", + "from twinkle import init_tinker_client\n", "from twinkle.data_format import Message, Trajectory\n", "from twinkle.template import Template\n", "\n", diff --git a/pyproject.toml b/pyproject.toml index 4b52c9c97..32ea8e643 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -58,7 +58,8 @@ server = [ test = [ "hypothesis>=6.0", "pytest", - "pytest-asyncio" + "pytest-asyncio", + "grimp>=3.0" ] docs = [ "sphinx>=5.3.0,<6.0.0", @@ -87,5 +88,6 @@ build-backend = "setuptools.build_meta" where = ["src"] [tool.setuptools.package-data] +"twinkle_client" = ["py.typed"] "twinkle_client.skills.bundled" = ["*.md"] "twinkle.kernel.ops.dsv4_sas_li.aclnn" = ["*.h", "*.cpp", "**/*.cpp"] diff --git a/setup.cfg b/setup.cfg index 811fd55cb..0454a8446 100644 --- a/setup.cfg +++ b/setup.cfg @@ -24,6 +24,9 @@ max-line-length = 120 select = B,E,F,P,T4,W,B9 ignore = F401,F403,F405,F821,W503,E251,W504,E126,E125 exclude = docs/src,*.pyi,.git,peft.py +per-file-ignores = + # Embedded LLM prompt/tool text with intentionally long lines that must stay verbatim. + src/twinkle_client/auto/agent/monitor.py:E501 [darglint] ignore=DAR101 diff --git a/src/twinkle/__init__.py b/src/twinkle/__init__.py index f64917a5e..af38ebfbd 100644 --- a/src/twinkle/__init__.py +++ b/src/twinkle/__init__.py @@ -1,10 +1,32 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any + +from ._lazy_module import _LazyModule # noqa + + +def init_tinker_client(**kwargs) -> None: + """Initialize the Tinker-compatible client without eager client imports.""" + from twinkle_client import init_tinker_client as _init_tinker_client + return _init_tinker_client(**kwargs) + + +def init_twinkle_client( + base_url: str | None = None, + api_key: str | None = None, + session_heartbeat_interval: int = 10, + **kwargs, +) -> Any: + """Initialize the Twinkle client without eager client imports.""" + from twinkle_client import init_twinkle_client as _init_twinkle_client + return _init_twinkle_client( + base_url=base_url, + api_key=api_key, + session_heartbeat_interval=session_heartbeat_interval, + **kwargs, + ) -from .utils.import_utils import _LazyModule # noqa if TYPE_CHECKING: - from twinkle_client import init_tinker_client, init_twinkle_client from .infra import get_device_placement, initialize, is_master, remote_class, remote_function from .utils import (GPU, NPU, DeviceGroup, DeviceMesh, Platform, Plugin, check_unsafe, exists, find_free_port, find_node_ip, framework_util, get_logger, requires, torch_util, trust_remote_code) @@ -21,8 +43,6 @@ import sys - from twinkle_client import init_tinker_client, init_twinkle_client - sys.modules[__name__] = _LazyModule( __name__, globals()['__file__'], @@ -30,6 +50,6 @@ module_spec=__spec__, # noqa extra_objects={ 'init_tinker_client': init_tinker_client, - 'init_twinkle_client': init_twinkle_client + 'init_twinkle_client': init_twinkle_client, }, ) diff --git a/src/twinkle/_lazy_module.py b/src/twinkle/_lazy_module.py new file mode 100644 index 000000000..1a7c2d260 --- /dev/null +++ b/src/twinkle/_lazy_module.py @@ -0,0 +1,61 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Dependency-light lazy module implementation used by package initializers.""" +import importlib +import os +from itertools import chain +from types import ModuleType +from typing import Any + + +class _LazyModule(ModuleType): + """ + Module class that surfaces all objects but only performs associated imports when the objects are requested. + """ + + # Very heavily inspired by optuna.integration._IntegrationModule + # https://github.com/optuna/optuna/blob/master/optuna/integration/__init__.py + def __init__(self, name, module_file, import_structure, module_spec=None, extra_objects=None): + super().__init__(name) + self._modules = set(import_structure.keys()) + self._class_to_module = {} + for key, values in import_structure.items(): + for value in values: + self._class_to_module[value] = key + # Needed for autocompletion in an IDE + self.__all__ = list(import_structure.keys()) + list(chain(*import_structure.values())) + self.__file__ = module_file + self.__spec__ = module_spec + self.__path__ = [os.path.dirname(module_file)] + self._objects = {} if extra_objects is None else extra_objects + self._name = name + self._import_structure = import_structure + + # Needed for autocompletion in an IDE + def __dir__(self): + result = super().__dir__() + # The elements of self.__all__ that are submodules may or may not be in the dir already, depending on whether + # they have been accessed or not. So we only add the elements of self.__all__ that are not already in the dir. + for attr in self.__all__: + if attr not in result: + result.append(attr) + return result + + def __getattr__(self, name: str) -> Any: + if name in self._objects: + return self._objects[name] + if name in self._modules: + value = self._get_module(name) + elif name in self._class_to_module.keys(): + module = self._get_module(self._class_to_module[name]) + value = getattr(module, name) + else: + raise AttributeError(f'module {self.__name__} has no attribute {name}') + + setattr(self, name, value) + return value + + def _get_module(self, module_name: str): + return importlib.import_module('.' + module_name, self.__name__) + + def __reduce__(self): + return self.__class__, (self._name, self.__file__, self._import_structure) diff --git a/src/twinkle/data_format/__init__.py b/src/twinkle/data_format/__init__.py index 5db25a2b5..01a9b629a 100644 --- a/src/twinkle/data_format/__init__.py +++ b/src/twinkle/data_format/__init__.py @@ -1,6 +1,9 @@ # Copyright (c) ModelScope Contributors. All rights reserved. +from .encoding import ENCODED_INPUT_KEYS, is_encoded from .input_feature import InputFeature from .message import Message, Tool, ToolCall from .output import LossOutput, ModelOutput from .sampling import SampledSequence, SampleResponse, SamplingMask, SamplingParams +from .tq_fields import (REQUIRED_MODEL_INPUT_FIELDS, ROLLOUT_TRAIN_FIELDS, TRANSFORMERS_INPUT_FIELDS, + columns_to_tq_fields, rows_to_tq_fields) from .trajectory import Trajectory, attach_user_data, pack_user_data, pack_value, user_data_get diff --git a/src/twinkle/data_format/encoding.py b/src/twinkle/data_format/encoding.py new file mode 100644 index 000000000..82de46808 --- /dev/null +++ b/src/twinkle/data_format/encoding.py @@ -0,0 +1,32 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""The single definition of "is this entry already encoded model input?". + +One predicate, one place. Before this module the same rule existed three times +(``MegatronModel._not_encoded``, ``TransformersModel._not_encoded``, +``Sampler._not_encoded``), one of them carrying a comment that it was "aligned +with" another -- an invariant only a human could maintain. The wire schema needs +the same rule to decide whether an ``inputs`` entry is an ``InputFeature`` or a +``Trajectory``, so a fourth copy would have made divergence a matter of time. + +``input_embedding`` matters as much as ``input_ids``: a batch carrying only +embeddings is already encoded, and misreading it as a ``Trajectory`` sends it +through ``template.batch_encode``, which fails far away from the cause. +""" +from __future__ import annotations + +from collections.abc import Mapping +from typing import Any + +# Presence of any of these keys means the entry carries encoded model input. +ENCODED_INPUT_KEYS: tuple[str, ...] = ('input_ids', 'input_embedding') + + +def is_encoded(entry: Any) -> bool: + """True when ``entry`` is an already-encoded ``InputFeature``-shaped mapping. + + A non-mapping is not encoded -- callers that need a type error raise it + themselves; this predicate answers only the classification question. + """ + if not isinstance(entry, Mapping): + return False + return any(key in entry for key in ENCODED_INPUT_KEYS) diff --git a/src/twinkle/tq_utils.py b/src/twinkle/data_format/tq_fields.py similarity index 51% rename from src/twinkle/tq_utils.py rename to src/twinkle/data_format/tq_fields.py index f7c34f073..a1eff271c 100644 --- a/src/twinkle/tq_utils.py +++ b/src/twinkle/data_format/tq_fields.py @@ -1,10 +1,45 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -"""Small TransferQueue packing helpers shared by both async-RL modes.""" +"""TransferQueue field packing and the field-name schema both async-RL modes share. + +Lives in ``data_format/`` because "rows/columns -> TensorDict" is a data-format +conversion, alongside ``input_feature`` / ``trajectory`` / ``encoding`` / ``message`` / +``output`` / ``sampling``. It used to sit at the package root as ``twinkle/tq_utils.py`` +-- the only domain module there, unreachable via ``twinkle.``, under an unexplained +abbreviation and a ``_utils`` suffix that undersold what it does (it validates field +consistency and raises). + +The field lists moved here from ``twinkle_agentic/async_rl/tq_utils.py``: the schema and +the packing logic belong in one file, and that shim existed only to re-export this +module. This deliberately means ``twinkle`` holds the RL training field names +(``logprobs`` / ``rewards`` / ``advantages`` / ``returns``) while ``twinkle_agentic`` +only consumes them. + +``torch`` / ``tensordict`` stay inside the functions: ``tensordict`` arrives with +``TransferQueue``, which is only in the ``async-rl`` extra, so this module must import +cleanly without it. Do NOT hoist them. +""" from __future__ import annotations from numbers import Number from typing import Any +TRANSFORMERS_INPUT_FIELDS = ( + 'input_ids', + 'labels', + 'attention_mask', + 'position_ids', + 'cu_seqlens', + 'completion_mask', + 'pixel_values', + 'image_grid_thw', + 'video_pixel_values', + 'video_grid_thw', + 'input_features', + 'feature_attention_mask', +) +REQUIRED_MODEL_INPUT_FIELDS = ('input_ids', 'labels', 'attention_mask', 'position_ids') +ROLLOUT_TRAIN_FIELDS = (*TRANSFORMERS_INPUT_FIELDS, 'logprobs', 'rewards', 'advantages', 'returns') + def rows_to_tq_fields(rows: list[dict[str, Any]]): from tensordict import TensorDict diff --git a/src/twinkle/dataset/base.py b/src/twinkle/dataset/base.py index c0fceb52d..b064df274 100644 --- a/src/twinkle/dataset/base.py +++ b/src/twinkle/dataset/base.py @@ -11,8 +11,6 @@ from typing import Any, Callable, Dict, List, Optional, Type, Union import twinkle -from twinkle import preprocessor -from twinkle.hub import HubOperation from twinkle.infra import remote_class, remote_function from twinkle.preprocessor import DataFilter, Preprocessor from twinkle.template import Template @@ -197,6 +195,7 @@ def _load_dataset(dataset_meta: DatasetMeta, **kwargs): kwargs['na_filter'] = False dataset = load_dataset(file_type, **load_kwargs, **kwargs) else: + from twinkle.hub import HubOperation dataset = HubOperation.load_dataset(dataset_id, subset_name, split, **kwargs) # fix: Some dataset sources return DatasetDict instead of Dataset, which breaks downstream select/map calls. diff --git a/src/twinkle/infra/__init__.py b/src/twinkle/infra/__init__.py index 4758b3415..6d6ae8d07 100644 --- a/src/twinkle/infra/__init__.py +++ b/src/twinkle/infra/__init__.py @@ -10,7 +10,7 @@ from twinkle.notifier import Notifier, notify_exception from twinkle.utils import DeviceGroup, DeviceMesh, Platform, check_unsafe, framework_util, get_logger, requires -from .collectors import collect_tensor_dict +from .collectors import collect_tensor_dict as collect_tensor_dict logger = get_logger() @@ -530,7 +530,7 @@ def _run_continous_work(self, func_name: str, execute_method, workers, args, kwa try: ordered: List[Any] = [None] * batch_len for _, indices, ref in submitted: - part = ray.get(ref, timeout=ray_get_timeout) if ray_get_timeout else ray.get(ref) + part = ray.get(ref, timeout=ray_get_timeout) if ray_get_timeout is not None else ray.get(ref) if not isinstance(part, (list, tuple)) or len(part) != len(indices): raise TypeError(f'{func_name}: enable_continous_work needs one result per request, but a worker given ' f'{len(indices)} request(s) returned {type(part).__name__} of length ' @@ -740,7 +740,6 @@ def _get_device_mesh_param(args, kwargs): def _prepare_lazy_collect(args, kwargs): # if a worker received an actor handle, # lazy collect should be false to prevent any outer function receives an object ref - from ._ray import RayHelper if not os.environ.get('WORKER_NAME'): # If this is a driver return args, kwargs @@ -996,7 +995,10 @@ def remote_function(dispatch: Union[Literal['slice', 'all', 'slice_dp', 'last_pp sync: If True, use synchronous execution (execute_all_sync) instead of async. Required for methods with NCCL collective operations (e.g., Megatron forward_backward). lazy_collect: Do lazy collect, this boolean value decides whether this function needs lazy collect. If setting to None, it will follow the global setting. - timeout: Timeout in seconds for ray.get() when collecting results. Instance attribute ``_ray_get_timeout`` overrides this. + timeout: Timeout in seconds for ray.get() when collecting results. The decorator's + explicitly declared value takes priority; the instance attribute ``_ray_get_timeout`` + is the fallback for methods that declare none (``timeout if timeout is not None + else instance``). enable_continous_work: Route each request to the least busy worker instead of slicing the batch over all of them, and return the results in the caller's order. This is what lets a batch smaller than the worker @@ -1044,7 +1046,14 @@ def wrapper(self, *args, **kwargs) -> T1: else: # This is the driver from ._ray import RayHelper - execute_method = RayHelper.execute_all_async if not sync else RayHelper.execute_all_sync + + # Resolve the effective ray.get timeout before choosing execute_method: + # the decorator's explicit value wins, the instance attribute is the + # fallback. ``is not None`` (not ``or``) so that a decorator ``timeout=0`` + # is honored instead of falling back to unbounded waiting. + _rgt = timeout if timeout is not None else getattr(self, '_ray_get_timeout', None) + execute_method = RayHelper.execute_all_async if not sync else functools.partial( + RayHelper.execute_all_sync, timeout=_rgt) # Only classes whose workers run methods side by side need # this; elsewhere Ray already orders calls per actor. _concurrent_actor = bool(getattr(self, '_max_concurrency', None)) @@ -1060,8 +1069,7 @@ def wrapper(self, *args, **kwargs) -> T1: _batch_len = _cw_batch_len(args, kwargs) if _batch_len: return _run_continous_work(self, func.__name__, execute_method, _workers, args, kwargs, - _batch_len, - getattr(self, '_ray_get_timeout', None) or timeout) + _batch_len, _rgt) if RayHelper.has_ref(args, kwargs): # If has any object-ref, dispatch in worker, because we don't know the structure in the ref. # for example, dataloader returns any data list. @@ -1079,7 +1087,6 @@ def wrapper(self, *args, **kwargs) -> T1: # busy. _tracked_refs = _cw_register(self, func.__name__, result) if _concurrent_actor else [] # This is a result future, call it to get the actual result - _rgt = getattr(self, '_ray_get_timeout', None) or timeout result_func = RayHelper.do_get_and_collect_func( _collect_func, collect, result, device_mesh, timeout=_rgt) _local_lazy_collect = _lazy_collect @@ -1090,13 +1097,13 @@ def wrapper(self, *args, **kwargs) -> T1: if func.__name__ == '__len__': # Get the first result and ignore the `lazy_collect` import ray - return ray.get(result[0]) + return ray.get(result[0], timeout=_rgt) if func.__name__ == '__next__': import ray for _res in result: # raise when any worker raises StopIteration - stop = ray.get(_res[1]) + stop = ray.get(_res[1], timeout=_rgt) if stop: raise StopIteration() result = [_res[0] for _res in result] diff --git a/src/twinkle/infra/_ray/ray_helper.py b/src/twinkle/infra/_ray/ray_helper.py index ffd4e1a42..4cc5f6a66 100644 --- a/src/twinkle/infra/_ray/ray_helper.py +++ b/src/twinkle/infra/_ray/ray_helper.py @@ -137,10 +137,17 @@ def is_worker(): return RayHelper.ray_inited() and ray._private.worker.global_worker.mode == ray._private.worker.WORKER_MODE @staticmethod - def execute_all_sync(method_name: str, workers_and_args: List[Tuple[Any, List[Any], Dict[str, Any]]]): - """Execute method and return results.""" + def execute_all_sync(method_name: str, workers_and_args: List[Tuple[Any, List[Any], Dict[str, Any]]], timeout=None): + """Execute method and return results. + + ``timeout`` is passed to ``ray.get(list, timeout=)``, whose semantics are + the **total** wall-clock time to collect the whole list -- different from + ``LazyCollect``'s per-future timing (see ``do_get_and_collect_func``). + The two paths are each bounded on their own; the total-time semantics here + are strictly tighter. + """ import ray - return ray.get(RayHelper.execute_all_async(method_name, workers_and_args)) + return ray.get(RayHelper.execute_all_async(method_name, workers_and_args), timeout=timeout) @staticmethod def execute_all_async(method_name: str, workers_and_args: List[Tuple[Any, List[Any], Dict[str, Any]]]): diff --git a/src/twinkle/loss/grpo.py b/src/twinkle/loss/grpo.py index e146edee6..2bf8d39fc 100644 --- a/src/twinkle/loss/grpo.py +++ b/src/twinkle/loss/grpo.py @@ -202,14 +202,16 @@ def _pad_and_align_to_batch( elif n_sample == n_pos: # Response-only form (e.g. old_logps from vLLM). result[i, pos] = sample - elif n_sample >= seq_len: - # Full-sequence form (e.g. ref_logps right-padded with ignore-value). - result[i, pos] = sample[:seq_len][mask[i]] + elif n_pos == 0 or (n_sample > 0 and pos[-1].item() < n_sample): + # Variable-length full-sequence form. The processor right-pads the + # batch, but per-sample RL fields from Tinker remain unpadded. They + # are valid when every selected mask position exists in this row. + result[i, pos] = sample[pos] else: raise AssertionError(f'data/mask length mismatch at sample {i}: ' f'n_pos={n_pos}, n_sample={n_sample}, seq_len={seq_len} ' - '(expected n_sample == n_pos for response-only form, ' - 'or n_sample >= seq_len for full-sequence form)') + '(expected n_sample == n_pos for response-only form, or all masked positions ' + 'to exist in the per-sample full-sequence form)') return result diff --git a/src/twinkle/model/__init__.py b/src/twinkle/model/__init__.py index 2b367bbd4..02f9eefd9 100644 --- a/src/twinkle/model/__init__.py +++ b/src/twinkle/model/__init__.py @@ -1,7 +1,7 @@ # Copyright (c) ModelScope Contributors. All rights reserved. from typing import TYPE_CHECKING -from twinkle.utils.import_utils import _LazyModule +from twinkle._lazy_module import _LazyModule if TYPE_CHECKING: from .base import TwinkleModel diff --git a/src/twinkle/model/megatron/__init__.py b/src/twinkle/model/megatron/__init__.py index 0f4625667..e7e611d29 100644 --- a/src/twinkle/model/megatron/__init__.py +++ b/src/twinkle/model/megatron/__init__.py @@ -6,7 +6,7 @@ # Follow the same LazyModule approach as `twinkle.model`: only import when those symbols are actually accessed. from typing import TYPE_CHECKING -from twinkle.utils.import_utils import _LazyModule +from twinkle._lazy_module import _LazyModule if TYPE_CHECKING: from .megatron import MegatronModel, MegatronStrategy diff --git a/src/twinkle/model/megatron/megatron.py b/src/twinkle/model/megatron/megatron.py index 851529c6a..f321192a7 100644 --- a/src/twinkle/model/megatron/megatron.py +++ b/src/twinkle/model/megatron/megatron.py @@ -25,7 +25,7 @@ import twinkle.patch from twinkle import DeviceMesh, Platform, remote_class, remote_function, requires, torch_util from twinkle.checkpoint_engine.mixin import CheckpointEngineMixin -from twinkle.data_format import InputFeature, ModelOutput, Trajectory +from twinkle.data_format import InputFeature, ModelOutput, Trajectory, is_encoded from twinkle.hub import HubOperation from twinkle.infra import collect_tensor_dict from twinkle.loss import CrossEntropyLoss, Loss @@ -36,7 +36,6 @@ from twinkle.processor import InputProcessor from twinkle.template import Template from twinkle.utils import construct_class, get_logger, selective_log_softmax -from twinkle.utils.nccl_safe import _is_fail_fast from ._mindspeed_runtime import ensure_mindspeed_adaptor_patched from .strategy import MegatronStrategy @@ -202,7 +201,7 @@ def _get_default_group(self): @staticmethod def _not_encoded(inputs): assert isinstance(inputs, dict) - return 'input_ids' not in inputs and 'input_embedding' not in inputs + return not is_encoded(inputs) @staticmethod def _slice_value_for_microbatch(value, mb_start: int, mb_end: int, micro_batch_size: int): @@ -420,62 +419,47 @@ def forward_step_func(data_iterator, model): embeddings = None _loss_instance = loss_instance is_last_pp = mpu.is_pipeline_last_stage(False, unwrapped_model.vp_stage) - try: - if task == 'embedding': - # MegatronEmbeddingPatch already pooled output to [n_seqs, hidden] on last PP stage. - if is_last_pp: - embeddings = output_tensor - elif labels is not None and is_last_pp: - _loss_require_logps = getattr(_loss_instance, 'require_logps', True) - _loss_require_entropy = getattr(_loss_instance, 'require_entropy', False) - _packed = batch.get('packed_seq_params') - cu_seqlens_q = getattr(_packed, 'cu_seqlens_q', None) if _packed is not None else None - if _loss_require_logps: - loss_mask = (labels != -100).bool() - masked_labels = labels.clone() - masked_labels[~loss_mask] = 0 - output_tensor.div_(temperature) - if _loss_require_entropy: - logps, entropies = selective_log_softmax(output_tensor, masked_labels, return_entropy=True) - else: - logps = selective_log_softmax(output_tensor, masked_labels) - # Reconstruct full-length tensors from CP-split shards - logps = processor.postprocess_tensor_cp(logps, cu_seqlens=cu_seqlens_q) - if entropies is not None: - entropies = processor.postprocess_tensor_cp(entropies, cu_seqlens=cu_seqlens_q) - batch['labels'] = processor.postprocess_tensor_cp(labels, cu_seqlens=cu_seqlens_q) - if completion_mask is not None: - # Same index space as labels, so it needs the same CP reassembly. - batch['completion_mask'] = processor.postprocess_tensor_cp( - completion_mask, cu_seqlens=cu_seqlens_q) - if 'position_ids' in batch: - pos = batch['position_ids'] - if pos.dim() == 3: - pos = pos[0] # [2/3, 1, seq] → [1, seq] - batch['position_ids'] = processor.postprocess_tensor_cp(pos, cu_seqlens=cu_seqlens_q) - # Unpack packed sequences into per-sequence batch format - _outputs = {'logps': logps} + if task == 'embedding': + # MegatronEmbeddingPatch already pooled output to [n_seqs, hidden] on last PP stage. + if is_last_pp: + embeddings = output_tensor + elif labels is not None and is_last_pp: + _loss_require_logps = getattr(_loss_instance, 'require_logps', True) + _loss_require_entropy = getattr(_loss_instance, 'require_entropy', False) + _packed = batch.get('packed_seq_params') + cu_seqlens_q = getattr(_packed, 'cu_seqlens_q', None) if _packed is not None else None + if _loss_require_logps: + loss_mask = (labels != -100).bool() + masked_labels = labels.clone() + masked_labels[~loss_mask] = 0 + output_tensor.div_(temperature) + if _loss_require_entropy: + logps, entropies = selective_log_softmax(output_tensor, masked_labels, return_entropy=True) + else: + logps = selective_log_softmax(output_tensor, masked_labels) + # Reconstruct full-length tensors from CP-split shards + logps = processor.postprocess_tensor_cp(logps, cu_seqlens=cu_seqlens_q) if entropies is not None: - _outputs['entropies'] = entropies - if hasattr(_loss_instance, 'require_logits') and _loss_instance.require_logits: - _outputs['logits'] = output_tensor - batch, _outputs = processor.unpack_packed_sequences(batch, _outputs) - logps = _outputs['logps'] - entropies = _outputs.get('entropies', None) - unpacked_logits = _outputs.get('logits', None) - except Exception as e: - # Data processing error (e.g. unpack_packed_sequences dimension mismatch). - # Must catch here inside the scheduler to prevent exception escaping - # and breaking PP P2P communication → NCCL hang. - if _is_fail_fast(): - raise - logger.warning('[nccl_safe] forward_step_func data processing error: ' - '%s: %s', - type(e).__name__, e) - logps = None - unpacked_logits = None - entropies = None - embeddings = None + entropies = processor.postprocess_tensor_cp(entropies, cu_seqlens=cu_seqlens_q) + batch['labels'] = processor.postprocess_tensor_cp(labels, cu_seqlens=cu_seqlens_q) + if completion_mask is not None: + # Same index space as labels, so it needs the same CP reassembly. + batch['completion_mask'] = processor.postprocess_tensor_cp(completion_mask, cu_seqlens=cu_seqlens_q) + if 'position_ids' in batch: + pos = batch['position_ids'] + if pos.dim() == 3: + pos = pos[0] # [2/3, 1, seq] → [1, seq] + batch['position_ids'] = processor.postprocess_tensor_cp(pos, cu_seqlens=cu_seqlens_q) + # Unpack packed sequences into per-sequence batch format + _outputs = {'logps': logps} + if entropies is not None: + _outputs['entropies'] = entropies + if hasattr(_loss_instance, 'require_logits') and _loss_instance.require_logits: + _outputs['logits'] = output_tensor + batch, _outputs = processor.unpack_packed_sequences(batch, _outputs) + logps = _outputs['logps'] + entropies = _outputs.get('entropies', None) + unpacked_logits = _outputs.get('logits', None) return output_tensor, partial( post_loss_function, inputs=batch, @@ -883,7 +867,7 @@ def clip_grad_and_step(self, max_grad_norm: float = 1.0, norm_type=2, **kwargs): self.zero_grad(**kwargs) self.lr_step(**kwargs) - @remote_function(dispatch='all', collect='first', sync=True) + @remote_function(dispatch='all', collect='first', sync=True, timeout=3600) def save(self, name: Optional[str] = None, output_dir: Optional[str] = None, @@ -1486,7 +1470,7 @@ def _patch_adapter(self, adapter_name: str, config_or_dir: Union[PeftConfig, str self._default_tokenizer = self.optimizer_group[adapter_name].template.processor self.active_group = adapter_name - @remote_function(dispatch='all', sync=True) + @remote_function(dispatch='all', sync=True, timeout=3600) def add_adapter_to_model( self, adapter_name: str, diff --git a/src/twinkle/model/megatron/multi_lora_megatron.py b/src/twinkle/model/megatron/multi_lora_megatron.py index ebda91501..8af62df4b 100644 --- a/src/twinkle/model/megatron/multi_lora_megatron.py +++ b/src/twinkle/model/megatron/multi_lora_megatron.py @@ -291,7 +291,7 @@ def _load_multi_lora_optimizer(self, checkpoint_dir: str, adapter_name: str = '' if optimizer_config is not None and 'iteration' in state_dict: optimizer_config.cur_step = state_dict['iteration'] - @remote_function(dispatch='all', collect='first', sync=True) + @remote_function(dispatch='all', collect='first', sync=True, timeout=3600) def save(self, name, output_dir: Optional[str] = None, interval=1, **kwargs): adapter_name = kwargs.pop('adapter_name', None) self._check_adapter_valid(adapter_name) @@ -372,7 +372,7 @@ def load(self, name: str, output_dir: Optional[str] = None, **kwargs): if dist.is_initialized(): dist.barrier() - @remote_function(dispatch='all', collect='first', sync=True) + @remote_function(dispatch='all', collect='first', sync=True, timeout=3600) def resume_from_checkpoint(self, checkpoint_dir, *, resume_only_model=False, **kwargs): adapter_name = kwargs.pop('adapter_name', None) self._check_adapter_valid(adapter_name) @@ -403,7 +403,7 @@ def get_state_dict(self, **kwargs): self._check_adapter_valid(kwargs.get('adapter_name')) return self.multi_adapter.get_state_dict(**kwargs) - @remote_function(dispatch='all', sync=True) + @remote_function(dispatch='all', sync=True, timeout=3600) def add_adapter_to_model( self, adapter_name: str, diff --git a/src/twinkle/model/multi_lora.py b/src/twinkle/model/multi_lora.py index ff7766102..70a77dca0 100644 --- a/src/twinkle/model/multi_lora.py +++ b/src/twinkle/model/multi_lora.py @@ -151,6 +151,19 @@ def deactivate_adapter(self): def patch_target_parameters(self, module, target_parameters): self.target_parameter_manager.patch(module, target_parameters) + @staticmethod + def _autocast_adapter_dtype(peft_model, adapter_name: str) -> None: + """Apply PEFT's own adapter dtype autocast to one preallocated slot. + + ``PeftModel.__init__`` runs ``_cast_adapter_dtype`` (fp16/bf16 -> fp32) for the + slot created at construction time, but ``PeftModel.add_adapter`` does not. Calling + it explicitly keeps every preallocated slot on the same dtype PEFT would produce, + instead of forcing fp32 ourselves. It is a no-op for an fp32 base. + """ + base = getattr(peft_model, 'base_model', None) + if base is not None and hasattr(base, '_cast_adapter_dtype'): + base._cast_adapter_dtype(adapter_name=adapter_name, autocast_adapter_dtype=True) + @contextmanager def adapter(self, tenant_adapter_name: str, disable_lora: bool = False): self.activate_adapter(tenant_adapter_name) @@ -193,12 +206,14 @@ def _after(_module): _before(_module) else: _before(self.module) - yield adapter_name - if isinstance(self.module, list): - for _module in self.module: - _after(_module) - else: - _after(self.module) + try: + yield adapter_name + finally: + if isinstance(self.module, list): + for _module in self.module: + _after(_module) + else: + _after(self.module) # self.deactivate_adapter() def check_length( @@ -510,6 +525,7 @@ def patch(self, def _patch_peft(_module): if isinstance(_module, PeftModel): _module.add_adapter(lora_tenant.adapter_name, config, low_cpu_mem_usage=low_cpu_mem_usage) + self._autocast_adapter_dtype(_module, lora_tenant.adapter_name) else: _peft_model: PeftModel = get_peft_model( _module, config, lora_tenant.adapter_name, low_cpu_mem_usage=low_cpu_mem_usage) @@ -526,6 +542,7 @@ def _patch_megatron(_module): _config = deepcopy(config) if isinstance(_module, PeftModel): _module.add_adapter(lora_tenant.adapter_name, _config, low_cpu_mem_usage=low_cpu_mem_usage) + self._autocast_adapter_dtype(_module, lora_tenant.adapter_name) else: # TODO first wrap needs parse target_modules, need to fix later if _config.target_modules: diff --git a/src/twinkle/model/optimizer_group.py b/src/twinkle/model/optimizer_group.py index f5177d672..150e694bc 100644 --- a/src/twinkle/model/optimizer_group.py +++ b/src/twinkle/model/optimizer_group.py @@ -48,12 +48,6 @@ class BaseOptimizerGroup: _device_mesh: DeviceMesh = None _last_grad_norm: float = 0.0 - def __setattr__(self, name, value): - if name == 'loss_instance' and value is not None: - from twinkle.utils.nccl_safe import safe_loss - value = safe_loss(value) - super().__setattr__(name, value) - def do_grad_sync(self, gradient_accumulation_steps: Optional[int] = None) -> bool: if gradient_accumulation_steps is None: gradient_accumulation_steps = self.gradient_accumulation_steps diff --git a/src/twinkle/model/transformers/moe/expert_parallel.py b/src/twinkle/model/transformers/moe/expert_parallel.py index 218e7b337..46717a8c2 100644 --- a/src/twinkle/model/transformers/moe/expert_parallel.py +++ b/src/twinkle/model/transformers/moe/expert_parallel.py @@ -234,7 +234,7 @@ def forward(hidden_states: torch.Tensor, *args, **kwargs): else: raise ValueError(f'Unsupported hidden_states ndim: {hidden_states.ndim}') - # R2 / R3 routing replay: pass block-level replay state + # Pass block-level routing replay state. from .router_replay import get_replay_state replay_state = get_replay_state(block_name) diff --git a/src/twinkle/model/transformers/transformers.py b/src/twinkle/model/transformers/transformers.py index 0087cc7f4..f7332eec1 100644 --- a/src/twinkle/model/transformers/transformers.py +++ b/src/twinkle/model/transformers/transformers.py @@ -27,7 +27,7 @@ from twinkle import DeviceMesh, Platform, remote_class, remote_function from twinkle.checkpoint_engine import CheckpointEngine from twinkle.checkpoint_engine.mixin import CheckpointEngineMixin -from twinkle.data_format import InputFeature, ModelOutput, Trajectory +from twinkle.data_format import InputFeature, ModelOutput, Trajectory, is_encoded from twinkle.hub import HubOperation from twinkle.infra import collect_tensor_dict from twinkle.loss import CrossEntropyLoss, Loss @@ -395,7 +395,7 @@ def _get_default_group(self): @staticmethod def _not_encoded(inputs): assert isinstance(inputs, dict) - return 'input_ids' not in inputs and 'input_embedding' not in inputs + return not is_encoded(inputs) def _lazy_wrap_model(self): if not self._model_wrapped: @@ -1480,7 +1480,8 @@ def _load_optimizer(self, checkpoint_dir, **kwargs): state_dict = torch.load(scheduler_path, map_location='cpu', weights_only=True) optimizer_config.lr_scheduler.load_state_dict(state_dict) - def _ensure_lora_dtype(self, model): + @staticmethod + def _ensure_lora_dtype(model): """Force LoRA parameters to use the same dtype as base model for FSDP2 compatibility.""" base_dtype = None is_npu_device = Platform.device_prefix() == 'npu' @@ -1557,7 +1558,7 @@ def _restore_training_state(self, checkpoint_dir, *, adapter_name=''): return trainer_state - @remote_function(dispatch='all', collect='first', sync=True) + @remote_function(dispatch='all', collect='first', sync=True, timeout=3600) def resume_from_checkpoint(self, checkpoint_dir, *, resume_only_model=False, **kwargs): adapter_name = kwargs.get('adapter_name', '') diff --git a/src/twinkle/server/utils/ray_serve_patch.py b/src/twinkle/patch/ray_serve.py similarity index 97% rename from src/twinkle/server/utils/ray_serve_patch.py rename to src/twinkle/patch/ray_serve.py index 69c385dd8..dabae7542 100644 --- a/src/twinkle/server/utils/ray_serve_patch.py +++ b/src/twinkle/patch/ray_serve.py @@ -138,4 +138,4 @@ def get_runtime_env_for_patches() -> dict: Returns: dict: Ray runtime_env configuration """ - return {'worker_process_setup_hook': ('twinkle.server.utils.ray_serve_patch._apply_patch_in_worker_process')} + return {'worker_process_setup_hook': ('twinkle.patch.ray_serve._apply_patch_in_worker_process')} diff --git a/src/twinkle/protocol/__init__.py b/src/twinkle/protocol/__init__.py new file mode 100644 index 000000000..9b1c3b8c0 --- /dev/null +++ b/src/twinkle/protocol/__init__.py @@ -0,0 +1,22 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Wire contract shared by the Twinkle server and client.""" +from .headers import (H_AUTH, H_AUTH_TWINKLE, H_MULTIPLEX, H_MULTIPLEX_LEGACY, H_REQUEST_ID, H_REQUEST_ID_LEGACY, + build_routing_headers) +from .json_utils import json_safe +from .serialize import deserialize_object, serialize_object +from .types import * # noqa: F403 +from .types import __all__ as _TYPES_ALL + +__all__ = [ + *_TYPES_ALL, + 'H_AUTH', + 'H_AUTH_TWINKLE', + 'H_MULTIPLEX', + 'H_MULTIPLEX_LEGACY', + 'H_REQUEST_ID', + 'H_REQUEST_ID_LEGACY', + 'build_routing_headers', + 'json_safe', + 'serialize_object', + 'deserialize_object', +] diff --git a/src/twinkle_client/http/headers.py b/src/twinkle/protocol/headers.py similarity index 100% rename from src/twinkle_client/http/headers.py rename to src/twinkle/protocol/headers.py diff --git a/src/twinkle_client/common/json_utils.py b/src/twinkle/protocol/json_utils.py similarity index 99% rename from src/twinkle_client/common/json_utils.py rename to src/twinkle/protocol/json_utils.py index 51c039c15..e0517a1d4 100644 --- a/src/twinkle_client/common/json_utils.py +++ b/src/twinkle/protocol/json_utils.py @@ -4,10 +4,8 @@ from collections.abc import Mapping from numbers import Number -from typing import Any - from pydantic import BaseModel - +from typing import Any _PRIMITIVE_TYPES = (str, Number, bool, bytes, type(None)) diff --git a/src/twinkle_client/common/serialize.py b/src/twinkle/protocol/serialize.py similarity index 72% rename from src/twinkle_client/common/serialize.py rename to src/twinkle/protocol/serialize.py index 42a27beb3..227b0b883 100644 --- a/src/twinkle_client/common/serialize.py +++ b/src/twinkle/protocol/serialize.py @@ -1,22 +1,36 @@ # Copyright (c) ModelScope Contributors. All rights reserved. +"""Tagged serialization for non-JSON domain objects used by the wire contract. + +Heavy domain types are resolved lazily so importing :mod:`twinkle.protocol` does +not load PEFT or the dataset implementation. +""" import json from dataclasses import fields +from functools import lru_cache from numbers import Number -from peft import LoraConfig from pydantic import BaseModel from typing import Any, Mapping -from twinkle.dataset import DatasetMeta - -supported_types = { - DatasetMeta, - LoraConfig, -} - -primitive_types = (str, Number, bool, bytes, type(None)) +primitive_types = (str, Number, bool, type(None)) container_types = (Mapping, list, tuple, set, frozenset) basic_types = (*primitive_types, *container_types) -_DATASET_META_FIELDS = {field.name for field in fields(DatasetMeta)} + + +@lru_cache(maxsize=1) +def _dataset_meta_cls(): + from twinkle.dataset import DatasetMeta + return DatasetMeta + + +@lru_cache(maxsize=1) +def _dataset_meta_fields() -> frozenset[str]: + return frozenset(field.name for field in fields(_dataset_meta_cls())) + + +@lru_cache(maxsize=1) +def _lora_config_cls(): + from peft import LoraConfig + return LoraConfig def _serialize_data_slice(data_slice): @@ -45,13 +59,15 @@ def _deserialize_data_slice(data_slice): raise ValueError(f'Unsupported data_slice type: {slice_type}') -def serialize_object(obj) -> str: - if isinstance(obj, DatasetMeta): - data = {name: getattr(obj, name) for name in _DATASET_META_FIELDS} +def serialize_object(obj) -> Any: + if isinstance(obj, (bytes, bytearray, memoryview)): + raise TypeError(f'Unsupported binary object: {type(obj).__name__}') + if isinstance(obj, _dataset_meta_cls()): + data = {name: getattr(obj, name) for name in _dataset_meta_fields()} data['data_slice'] = _serialize_data_slice(data.get('data_slice')) data['_TWINKLE_TYPE_'] = 'DatasetMeta' return json.dumps(data, ensure_ascii=False) - elif isinstance(obj, LoraConfig): + elif isinstance(obj, _lora_config_cls()): filtered_dict = {} for _subkey, _subvalue in obj.__dict__.items(): if isinstance(_subvalue, basic_types) and not _subkey.startswith('_'): @@ -82,11 +98,12 @@ def deserialize_object(data: str) -> Any: if '_TWINKLE_TYPE_' in data: _type = data.pop('_TWINKLE_TYPE_') if _type == 'DatasetMeta': - data = {key: value for key, value in data.items() if key in _DATASET_META_FIELDS} + fields_set = _dataset_meta_fields() + data = {key: value for key, value in data.items() if key in fields_set} data['data_slice'] = _deserialize_data_slice(data.get('data_slice')) - return DatasetMeta(**data) + return _dataset_meta_cls()(**data) elif _type == 'LoraConfig': - return LoraConfig(**data) + return _lora_config_cls()(**data) else: raise ValueError(f'Unsupported type: {_type}') else: diff --git a/src/twinkle/protocol/types/__init__.py b/src/twinkle/protocol/types/__init__.py new file mode 100644 index 000000000..01b27257d --- /dev/null +++ b/src/twinkle/protocol/types/__init__.py @@ -0,0 +1,156 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +# yapf: disable +from .base import (BACKEND_ONLY_KEY, DataModel, FieldRole, ResponseModel, StrictRequest, backend_kwarg, backend_only, + fields_with_role, passthrough, read_backend_only, read_field_role) +from .checkpoint import ResolvedLoadPath +from .component import (DataAppendRequest, DataGetRequest, DataPlaneSampleRequest, DataPutRequest, DataRef, + DataReleaseRequest, DataRowsResponse, UnloadAdapterPathsRequest) +from .data import (CORE_INPUT_KEYS, VLM_TENSOR_FIELDS, WireInputBatch, WireInputFeature, WireInputs, WireMessage, + WireTrajectory, declared_wire_keys, export_batch) +from .lifecycle import TERMINAL_STATUSES, CancelRequest, CancelResponse, RetrieveFutureRequest, TaskEnvelope, TaskStatus +from .model import (AdapterRequest, AddAdapterRequest, AddMetricRequest, AddMetricResponse, ApplyPatchRequest, + ApplyPatchResponse, BackwardResponse, CalculateLossResponse, CalculateMetricRequest, + CalculateMetricResponse, ClipGradAndStepRequest, ClipGradAndStepResponse, ClipGradNormRequest, + ClipGradNormResponse, CreateRequest, CreateResponse, DataPlaneForwardOnlyRequest, + DataPlaneForwardRequest, ForwardBackwardResponse, ForwardBackwardTaskRequest, ForwardOnlyRequest, + ForwardRequest, ForwardResponse, GetTrainConfigsResponse, LoadRequest, LoadResponse, LrStepRequest, + LrStepResponse, ModelResult, OkResponse, ResumeFromCheckpointRequest, SaveRequest, SaveResponse, + SetLossRequest, SetLossResponse, SetLrSchedulerRequest, SetLrSchedulerResponse, SetOptimizerRequest, + SetOptimizerResponse, SetProcessorRequest, SetProcessorResponse, SetTemplateRequest, + SetTemplateResponse, StepRequest, StepResponse, TrainingProgressResponse, UploadToHubRequest, + ZeroGradResponse) +from .processor import (ProcessorCallRequest, ProcessorCallResponse, ProcessorCreateRequest, ProcessorCreateResponse, + ProcessorHeartbeatRequest, ProcessorHeartbeatResponse) +from .sampler import (SampledSequenceModel, SamplerAddAdapterRequest, SamplerAddAdapterResponse, SamplerCreateResponse, + SampleRequest, SampleResponseModel, SampleResponseModelList, SamplerSetTemplateRequest, + SamplerSetTemplateResponse) +from .server import (CapacityInfoResponse, CheckpointPathResponse, ClientFeatures, DeleteCheckpointResponse, + GetServerCapabilitiesResponse, HealthResponse, ProtocolLimits, SupportedModel, WeightsInfoRequest) +from .session import CreateSessionRequest, CreateSessionResponse, SessionHeartbeatRequest, SessionHeartbeatResponse +from .training import (Checkpoint, CheckpointsListResponse, CreateModelRequest, Cursor, LoraConfig, + ParsedCheckpointTwinklePath, TrainingRun, TrainingRunsResponse, WeightsInfoResponse) + +# yapf: enable + +__all__ = [ + 'BACKEND_ONLY_KEY', + 'DataModel', + 'FieldRole', + 'ResponseModel', + 'StrictRequest', + 'backend_kwarg', + 'backend_only', + 'fields_with_role', + 'passthrough', + 'read_backend_only', + 'read_field_role', + 'ResolvedLoadPath', + 'DataAppendRequest', + 'DataGetRequest', + 'DataPlaneSampleRequest', + 'DataPutRequest', + 'DataRef', + 'DataReleaseRequest', + 'DataRowsResponse', + 'UnloadAdapterPathsRequest', + 'CORE_INPUT_KEYS', + 'VLM_TENSOR_FIELDS', + 'WireInputBatch', + 'WireInputFeature', + 'WireInputs', + 'WireMessage', + 'WireTrajectory', + 'declared_wire_keys', + 'export_batch', + 'TERMINAL_STATUSES', + 'CancelRequest', + 'CancelResponse', + 'RetrieveFutureRequest', + 'TaskEnvelope', + 'TaskStatus', + 'AdapterRequest', + 'AddAdapterRequest', + 'AddMetricRequest', + 'AddMetricResponse', + 'ApplyPatchRequest', + 'ApplyPatchResponse', + 'BackwardResponse', + 'CalculateLossResponse', + 'CalculateMetricRequest', + 'CalculateMetricResponse', + 'ClipGradAndStepRequest', + 'ClipGradAndStepResponse', + 'ClipGradNormRequest', + 'ClipGradNormResponse', + 'CreateRequest', + 'CreateResponse', + 'DataPlaneForwardOnlyRequest', + 'DataPlaneForwardRequest', + 'ForwardBackwardResponse', + 'ForwardBackwardTaskRequest', + 'ForwardOnlyRequest', + 'ForwardRequest', + 'ForwardResponse', + 'GetTrainConfigsResponse', + 'LoadRequest', + 'LoadResponse', + 'LrStepRequest', + 'LrStepResponse', + 'ModelResult', + 'OkResponse', + 'ResumeFromCheckpointRequest', + 'SaveRequest', + 'SaveResponse', + 'SetLossRequest', + 'SetLossResponse', + 'SetLrSchedulerRequest', + 'SetLrSchedulerResponse', + 'SetOptimizerRequest', + 'SetOptimizerResponse', + 'SetProcessorRequest', + 'SetProcessorResponse', + 'SetTemplateRequest', + 'SetTemplateResponse', + 'StepRequest', + 'StepResponse', + 'TrainingProgressResponse', + 'UploadToHubRequest', + 'ZeroGradResponse', + 'ProcessorCallRequest', + 'ProcessorCallResponse', + 'ProcessorCreateRequest', + 'ProcessorCreateResponse', + 'ProcessorHeartbeatRequest', + 'ProcessorHeartbeatResponse', + 'SampledSequenceModel', + 'SamplerAddAdapterRequest', + 'SamplerAddAdapterResponse', + 'SamplerCreateResponse', + 'SampleRequest', + 'SampleResponseModel', + 'SampleResponseModelList', + 'SamplerSetTemplateRequest', + 'SamplerSetTemplateResponse', + 'CapacityInfoResponse', + 'CheckpointPathResponse', + 'ClientFeatures', + 'DeleteCheckpointResponse', + 'GetServerCapabilitiesResponse', + 'HealthResponse', + 'ProtocolLimits', + 'SupportedModel', + 'WeightsInfoRequest', + 'CreateSessionRequest', + 'CreateSessionResponse', + 'SessionHeartbeatRequest', + 'SessionHeartbeatResponse', + 'Checkpoint', + 'CheckpointsListResponse', + 'CreateModelRequest', + 'Cursor', + 'LoraConfig', + 'ParsedCheckpointTwinklePath', + 'TrainingRun', + 'TrainingRunsResponse', + 'WeightsInfoResponse', +] diff --git a/src/twinkle/protocol/types/base.py b/src/twinkle/protocol/types/base.py new file mode 100644 index 000000000..d428afc81 --- /dev/null +++ b/src/twinkle/protocol/types/base.py @@ -0,0 +1,169 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Shared pydantic base classes, field roles, and the naming rulings for the wire contract. + +This module is a public contract carrier imported across packages (Twinkle_Server +reverse-imports ``twinkle.protocol.types``); it therefore intentionally carries **no** +underscore prefix. + +Naming rulings (authoritative for all three split specs; kept in code, not only in +the spec, so a later reader cannot merge these away): + +1. Schema_Module modules imported across packages do NOT use an underscore prefix. + ``base.py`` / ``errors.py`` / ``lifecycle.py`` / ``data.py`` are public-contract + carriers; an underscore means "package-private", and a cross-package import of a + private module is a violation. Modules used only inside Twinkle_Client (never + imported by Twinkle_Server) are exempt. +2. New twinkle-native request models do NOT reuse a class name already present in + ``tinker.types``. Known collision to avoid: ``ForwardBackwardRequest``. Two + handlers import ``types`` from twinkle_client and from tinker respectively; a + same-named model is distinguished only by the import alias and is easy to + misread in a review diff. +3. The field expressing a failure-semantic category is named ``error_code``, NOT + ``status_code`` -- an execution-time failure is delivered with HTTP 200, so the + value is systematically unequal to the response status code. +4. A closed value set on a wire field is declared as ``Literal`` / enum, never a + bare ``str`` (see ``QueueStateLiteral`` in ``errors.py``). + +Field roles +----------- +Every declared request field has exactly one role, and the role -- not the field +name -- decides whether it reaches the backend: + +- ``Control`` (the default): consumed by the handler itself, or passed explicitly + as a named argument. ``inputs``, ``adapter_name``, ``seq_id`` and the data-plane + reference fields are control fields. Forwarding them again through + ``**backend_kwargs`` would either duplicate a keyword argument or leak a + protocol field into a backend signature. +- ``BackendKwarg``: a user-facing backend parameter. Forwarded when, and only + when, its value is not ``None``. +- ``Passthrough``: a declared dict whose *keys* are dynamic. Its contents are + flattened into the backend kwargs, and its keys are exempt from + ``extra='forbid'`` (that setting constrains the model's own field set, not the + inside of a declared dict). The keys are forwarded as given -- see + :func:`passthrough` for why they are not checked against the target's signature. + +The role lives in the field's ``json_schema_extra`` so that one declaration site +carries it -- there is deliberately no second per-endpoint parameter table. +""" +from __future__ import annotations + +from enum import StrEnum +from pydantic import BaseModel, ConfigDict, Field +from pydantic.fields import FieldInfo +from typing import Any, Optional + + +class StrictRequest(BaseModel): + """Request bodies. Typos fail loudly.""" + + model_config = ConfigDict(frozen=True, extra='forbid') + + +class ResponseModel(BaseModel): + """Response bodies. An old client tolerates new server fields.""" + + model_config = ConfigDict(frozen=True, extra='ignore') + + +class DataModel(BaseModel): + """Data-plane models (InputFeature / Trajectory on the wire). + + ``extra='allow'``, not ``forbid`` and not ``ignore``, and the difference is + load-bearing. A user's ``Preprocessor`` / ``Template`` routinely leaves extra + columns on an entry (whatever ``dataset.map`` did not drop). Rejecting those + would break a large number of existing datasets -- but *ignoring* them is just + as wrong, because the entry is then re-exported to the backend without them, + silently dropping data the caller sent. ``allow`` keeps unknown keys on the + model so :meth:`export` can hand them back. + + Do NOT "fix" this to ``StrictRequest`` or to ``extra='ignore'``. Strictness on + the data plane belongs to the *declared* fields (the ones Twinkle_Core reads), + which carry strict types; it does not belong to the field set. + """ + + model_config = ConfigDict(frozen=True, extra='allow') + + +class FieldRole(StrEnum): + """How a declared request field relates to the backend call.""" + + Control = 'control' + BackendKwarg = 'backend_kwarg' + Passthrough = 'passthrough' + + +# Keys under which field metadata is stored in a field's ``json_schema_extra``. +FIELD_ROLE_KEY = 'twinkle_field_role' +BACKEND_ONLY_KEY = 'twinkle_backend_only' + + +def _with_extra(field_kwargs: dict[str, Any], **extra: Any) -> FieldInfo: + merged = dict(field_kwargs.pop('json_schema_extra', None) or {}) + merged.update(extra) + return Field(json_schema_extra=merged, **field_kwargs) + + +def backend_kwarg(*backends: str, **field_kwargs: Any) -> FieldInfo: + """Declare a field as a backend keyword argument. + + With no ``backends`` the field applies to every backend. Naming one or more + restricts it: sending a non-``None`` value to a deployment running a different + backend is rejected before the task is enqueued. + + A restricted field MUST be ``Optional[...] = None``. Giving it a backend's + constant default would make it carry a non-``None`` value on every deployment, + so the "non-``None`` on the wrong backend" check would reject every request. + """ + return _with_extra(field_kwargs, **{ + FIELD_ROLE_KEY: FieldRole.BackendKwarg.value, + BACKEND_ONLY_KEY: tuple(backends) or None, + }) + + +def backend_only(*backends: str, **field_kwargs: Any) -> FieldInfo: + """A backend keyword argument restricted to the given backend(s). + + Kept as a named entry point because "this parameter only exists on megatron" + is the property a reader looks for at the declaration site; it is + :func:`backend_kwarg` with a non-empty backend tuple, not a second mechanism. + """ + if not backends: + raise ValueError('backend_only() requires at least one backend; use backend_kwarg() for an unrestricted field') + return backend_kwarg(*backends, **field_kwargs) + + +def passthrough(**field_kwargs: Any) -> FieldInfo: + """Declare a dict field whose keys are dynamic backend parameters. + + Its contents are forwarded to the backend as given. They are deliberately not + checked against the target's signature: a plugin routinely reads a real parameter + straight out of ``**kwargs`` (``InputProcessor`` does this with ``padding_side``), + and ``inspect.signature`` cannot see such a read -- so any such check rejects valid + requests. A misspelled plugin argument therefore still surfaces from the plugin. + """ + field_kwargs.setdefault('default_factory', dict) + return _with_extra(field_kwargs, **{FIELD_ROLE_KEY: FieldRole.Passthrough.value}) + + +def _read_extra(field_info: FieldInfo, key: str) -> Any: + extra = getattr(field_info, 'json_schema_extra', None) + if isinstance(extra, dict): + return extra.get(key) + return None + + +def read_field_role(field_info: FieldInfo) -> FieldRole: + """The field's role; ``Control`` when undeclared.""" + value = _read_extra(field_info, FIELD_ROLE_KEY) + return FieldRole(value) if value is not None else FieldRole.Control + + +def read_backend_only(field_info: FieldInfo) -> tuple[str, ...] | None: + """Return the backend tuple a field was restricted to, or ``None`` if unrestricted.""" + value = _read_extra(field_info, BACKEND_ONLY_KEY) + return tuple(value) if value else None + + +def fields_with_role(model_cls: type[BaseModel], role: FieldRole) -> dict[str, FieldInfo]: + """The model's declared fields carrying ``role``, in declaration order.""" + return {name: info for name, info in model_cls.model_fields.items() if read_field_role(info) is role} diff --git a/src/twinkle_client/types/checkpoint.py b/src/twinkle/protocol/types/checkpoint.py similarity index 100% rename from src/twinkle_client/types/checkpoint.py rename to src/twinkle/protocol/types/checkpoint.py diff --git a/src/twinkle_client/types/component.py b/src/twinkle/protocol/types/component.py similarity index 53% rename from src/twinkle_client/types/component.py rename to src/twinkle/protocol/types/component.py index d7e9ec2a0..4751774ef 100644 --- a/src/twinkle_client/types/component.py +++ b/src/twinkle/protocol/types/component.py @@ -2,13 +2,21 @@ """Protocol types for directly orchestrating asynchronous server components.""" from __future__ import annotations -from typing import Any +from pydantic import BaseModel, Field, JsonValue, model_validator +from typing import Any, Optional -from pydantic import BaseModel, Field, model_validator +from .base import ResponseModel, StrictRequest +from .data import WireInputBatch class DataRef(BaseModel): - """Opaque reference to rows stored in the server-side TransferQueue.""" + """Opaque reference to rows stored in the server-side TransferQueue. + + A value carried inside other bodies rather than a body of its own, and it is + round-tripped by the client, so it keeps the plain base. No wire schema is applied + to what it points at: the rows never travel in the request body, so the data-plane + constraints would be a category error here. + """ ref_id: str size: int @@ -17,53 +25,57 @@ class DataRef(BaseModel): num_tokens: int = 0 -class DataPutRequest(BaseModel): +class DataPutRequest(StrictRequest): rows: list[dict[str, Any]] kind: str = 'data' tags: list[dict[str, Any]] | None = None -class DataGetRequest(BaseModel): +class DataGetRequest(StrictRequest): ref: DataRef fields: list[str] | None = None include_tags: bool = False -class DataAppendRequest(BaseModel): +class DataAppendRequest(StrictRequest): ref: DataRef rows: list[dict[str, Any]] tags: list[dict[str, Any]] | None = None -class DataReleaseRequest(BaseModel): +class DataReleaseRequest(StrictRequest): ref: DataRef -class DataRowsResponse(BaseModel): +class DataRowsResponse(ResponseModel): rows: list[dict[str, Any]] tags: list[dict[str, Any]] = Field(default_factory=list) -class DataPlaneSampleRequest(BaseModel): - inputs: Any = None +class DataPlaneSampleRequest(StrictRequest): + """Body of ``POST /twinkle/sample_to_data_plane``. + + Exactly one input source: inline entries (wire-validated) or a ``DataRef``. + """ + + inputs: WireInputBatch | None = None input_ref: DataRef | None = None - sampling_params: dict[str, Any] | None = None + sampling_params: dict[str, JsonValue] | None = None adapter_name: str = '' adapter_uri: str | None = None policy_version: int | None = None group_ids: list[str] | None = None - num_samples: int = 1 + num_samples: int = Field(default=1, ge=1) @model_validator(mode='after') - def validate_input(self) -> 'DataPlaneSampleRequest': + def validate_input(self) -> DataPlaneSampleRequest: if (self.inputs is None) == (self.input_ref is None): raise ValueError('exactly one of inputs and input_ref must be provided') if self.group_ids is not None and self.inputs is not None: - size = len(self.inputs) if isinstance(self.inputs, list) else 1 - if len(self.group_ids) != size: + if len(self.group_ids) != len(self.inputs): raise ValueError('group_ids must contain one value per sampler input') return self -class UnloadAdapterPathsRequest(BaseModel): +class UnloadAdapterPathsRequest(StrictRequest): adapter_paths: list[str] diff --git a/src/twinkle/protocol/types/data.py b/src/twinkle/protocol/types/data.py new file mode 100644 index 000000000..b7936dbcd --- /dev/null +++ b/src/twinkle/protocol/types/data.py @@ -0,0 +1,216 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Wire schema for the inline ``inputs`` data plane. + +These models are the *declared* type of the ``inputs`` request field, so FastAPI +validates a batch during body parsing -- before the handler runs, before a future +record exists, and before anything reaches a GPU. The seam that hands data to the +backend (``twinkle.server.lifecycle.submit.to_backend_inputs``) therefore only +exports an already-valid object; it is not the first place a malformed batch is +noticed. + +Two asymmetries are deliberate: + +- **Strict on declared fields, open on the field set.** Every field Twinkle_Core + reads is declared with a strict type (``StrictInt`` leaves reject ``true`` and + ``1.0``), while unknown JSON-native keys are kept and re-exported: a user's + preprocessor may leave extra columns on an entry and dropping them would lose + data the caller sent. See :class:`~twinkle.protocol.types.base.DataModel`. +- **Shallowest-first unions.** Nesting depth encodes tensor rank here, so a rank + range needs a union. Declaring the deepest branch first is a large, silent + pessimisation: given a 2-D input the 3-D branch does not fail at element 0, it + descends into every row and records one error per element. Measured on pydantic + 2.13.4 with a 1024 x 8192 (8.4M element) 2-D ``input_ids``: + + List[List[int]] single type 0.141 s + Union[3D, 2D, 1D] deepest-first 5.368 s <-- 41x worse + Union[1D, 2D, 3D] shallowest-first 0.131 s + + Do NOT reorder these to deepest-first. ``test_wire_schema.py`` asserts the + declared depth sequence is strictly increasing, so the ordering is checked + structurally rather than by a flaky timing benchmark. + +Known technical debt: encoding tensor shape in JSON nesting depth is what forces +the unions and makes validation cost scale with element count. The target shape is +a flat ``{dtype, shape, data}`` tensor envelope (which is what ``tinker`` uses). +Migrating is a breaking wire change and is out of scope here; do not paper over it +with hand-written Python-level depth checks or extra union branches, which would +add to the debt rather than pay it down. +""" +from __future__ import annotations + +from collections.abc import Mapping +from pydantic import BeforeValidator, Field, StrictInt, model_validator +from typing import Annotated, Any, List, Literal, Optional, Union + +from twinkle.data_format.encoding import ENCODED_INPUT_KEYS +from twinkle.protocol.types.base import DataModel + +# --------------------------------------------------------------------------- # +# Leaf types. Shallowest-first, and ``StrictInt`` wherever the values come from a +# tensor's ``tolist()`` -- there, a bool or a float is an upstream defect. Lax +# ``int`` coercion would turn ``[true, false]`` into ``[1, 0]``. +# +# ``union_mode='left_to_right'`` makes the declared order load-bearing instead of +# leaving branch selection to pydantic's heuristics. +# --------------------------------------------------------------------------- # + +_LEFT_TO_RIGHT = Field(union_mode='left_to_right') + +Ints1to2 = Annotated[Union[list[StrictInt], list[list[StrictInt]]], _LEFT_TO_RIGHT] +Ints1to3 = Annotated[Union[list[StrictInt], list[list[StrictInt]], list[list[list[StrictInt]]]], _LEFT_TO_RIGHT] +Ints3 = list[list[list[StrictInt]]] + +_Number = Union[StrictInt, float] +Numbers1to2 = Annotated[Union[list[_Number], list[list[_Number]]], _LEFT_TO_RIGHT] +Numbers1to4 = Annotated[Union[list[_Number], list[list[_Number]], list[list[list[_Number]]], + list[list[list[list[_Number]]]]], _LEFT_TO_RIGHT] + +# Media references travel as strings on the wire (local path, ``http(s)://`` URL, or +# a ``data:`` base64 URI). ``PIL.Image`` / raw ``bytes`` / ``np.ndarray`` are valid in +# the in-process training path but are not JSON, so they are not declared here. +MediaList = list[str] + +# The VLM tensor fields batched by concatenation rather than padding. Declared here +# because this module must stay free of Twinkle_Core's heavyweight imports; a +# consistency test asserts this set equals ``InputProcessor.VLM_CONCAT_FIELDS``, so a +# future addition there fails loudly instead of being silently dropped on the wire. +VLM_TENSOR_FIELDS: frozenset[str] = frozenset({ + 'pixel_values', + 'image_grid_thw', + 'pixel_values_videos', + 'video_grid_thw', + 'input_features', + 'input_features_mask', + 'feature_attention_mask', + 'grid_thws', +}) + + +class WireMessage(DataModel): + """One conversation turn, as sent over HTTP.""" + + role: Literal['system', 'user', 'assistant', 'tool'] | None = None + type: str | None = None + content: str | list[dict[str, Any]] | None = None + tool_calls: list[dict[str, Any]] | None = None + tool_call_id: str | None = None + reasoning_content: str | None = None + images: MediaList | None = None + videos: MediaList | None = None + audios: MediaList | None = None + + +class WireInputFeature(DataModel): + """An already-encoded entry: token ids (or embeddings) plus aligned tensors.""" + + input_ids: Ints1to2 | None = None + input_embedding: Numbers1to2 | None = None + attention_mask: Ints1to2 | None = None + labels: Ints1to2 | None = None + completion_mask: Ints1to2 | None = None + # 1-D standard encoding, 2-D Qwen-VL mrope ``[3, T]``, 3-D megatron ``[3, 1, N]``. + position_ids: Ints1to3 | None = None + # Exactly ``[seq_len, num_layers, topk]``. + routed_experts: Ints3 | None = None + length: StrictInt | None = None + + # VLM tensors: float values are normal here, so no strict-int leaves. + pixel_values: Numbers1to4 | None = None + image_grid_thw: Numbers1to4 | None = None + pixel_values_videos: Numbers1to4 | None = None + video_grid_thw: Numbers1to4 | None = None + input_features: Numbers1to4 | None = None + input_features_mask: Numbers1to4 | None = None + feature_attention_mask: Numbers1to4 | None = None + grid_thws: Numbers1to4 | None = None + + @model_validator(mode='after') + def require_encoded_key(self) -> WireInputFeature: + """At least one of the encoded-input keys must be present. + + Declared as a model validator rather than by making ``input_ids`` required: + an embedding-only batch is legitimately encoded, and this is the same rule + the backends apply (:data:`ENCODED_INPUT_KEYS`). + """ + if all(getattr(self, key, None) is None for key in ENCODED_INPUT_KEYS): + raise ValueError(f'an encoded entry requires one of {list(ENCODED_INPUT_KEYS)}') + return self + + +class WireTrajectory(DataModel): + """A not-yet-encoded entry: messages the server template will encode.""" + + messages: list[WireMessage] + images: MediaList | None = None + videos: MediaList | None = None + audios: MediaList | None = None + tools: list[dict[str, Any]] | None = None + # ``List[Tuple[str, str]]`` on the wire: the PyArrow-stable encoding of the + # user-data pairs attached by ``twinkle.data_format.attach_user_data``. + user_data: list[tuple[str, str]] | None = None + + +# A batch is homogeneous: every entry is encoded, or none is. Expressed as a union of +# *lists* rather than a list of unions, so a mixed batch fails to match either branch +# instead of being silently accepted and blowing up inside the backend. Order is +# encoded-first, matching ``is_encoded``: a trajectory has neither encoded key, so it +# cannot satisfy ``WireInputFeature``'s validator. +WireInputs = Union[list[WireInputFeature], list[WireTrajectory]] + + +def _as_batch(value: Any) -> Any: + """Accept a single entry where a batch is expected. + + Callers have always been allowed to pass one mapping instead of a one-element + list; normalising here keeps that while letting the declared type stay a batch, + so downstream code has exactly one shape to handle. + """ + return [value] if isinstance(value, Mapping) else value + + +#: The declared type of an inline ``inputs`` request field. +WireInputBatch = Annotated[WireInputs, BeforeValidator(_as_batch)] + +# Every ``inputs`` key Twinkle_Core reads. Maintained by hand on purpose: an AST scan +# would have to follow aliases (``inp = inputs[i]`` then ``inp.get('x')``), i.e. do a +# local data-flow analysis, and its false negatives would *silently* disable the +# consistency check that is this schema's only safety net against a dropped field. +# When adding a read of a new ``inputs`` key, add it here. +CORE_INPUT_KEYS: frozenset[str] = frozenset({ + 'input_ids', + 'input_embedding', + 'attention_mask', + 'labels', + 'completion_mask', + 'position_ids', + 'routed_experts', + 'length', + 'messages', + 'images', + 'videos', + 'audios', + 'tools', + 'user_data', +}) | VLM_TENSOR_FIELDS + + +def declared_wire_keys() -> frozenset[str]: + """Union of the field names declared across the wire input models.""" + return frozenset(WireInputFeature.model_fields) | frozenset(WireTrajectory.model_fields) + + +def export(entry: WireInputFeature | WireTrajectory) -> dict[str, Any]: + """Render a validated entry as the plain dict the backend consumes. + + ``exclude_none=True`` is required, not cosmetic: Twinkle_Core branches on key + *presence* in many places (``is_encoded``, ``inputs.pop('labels', None)``, the + VLM concat fields), so emitting unset optionals as ``None`` would change + behaviour. Unknown keys the caller sent are preserved -- that is why + :class:`DataModel` uses ``extra='allow'``. + """ + return entry.model_dump(exclude_none=True) + + +def export_batch(entries: list[Any]) -> list[dict[str, Any]]: + """Export a validated batch, leaving already-plain entries untouched.""" + return [export(entry) if isinstance(entry, DataModel) else entry for entry in entries] diff --git a/src/twinkle/protocol/types/errors.py b/src/twinkle/protocol/types/errors.py new file mode 100644 index 000000000..b896e4eed --- /dev/null +++ b/src/twinkle/protocol/types/errors.py @@ -0,0 +1,68 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Structured failure payload -- the single representation of a failure. + +Twinkle <-> tinker exception mapping (verified, kept here so a future new exception +class can be lined up against its tinker counterpart): + +- Tinker 0.29.0 ``RequestFailedError`` (``tinker/_exceptions.py``; carries + ``message`` / ``request_id`` / ``category``) is the "the task completed in a failed + terminal state" exception. Its wire values are ``unknown`` / ``server`` / + ``user``, matching :class:`ErrorCategory`; legacy TitleCase values are normalized. + +``twinkle_client/utils/patch_tinker.py`` shows two SDKs can coexist in one process, +so a semantically-equal but differently-named exception must be lookup-able. +""" +from __future__ import annotations + +from enum import StrEnum +from pydantic import Field, field_validator, model_validator +from typing import Any, Literal, Optional + +from .base import ResponseModel + +# Closed value set, kept in sync with the server-side ``QueueState`` enum values +# (a consistency test asserts the two sets are equal). Wire fields carrying a queue +# state declare this alias, never a bare ``str`` (naming ruling 4). +QueueStateLiteral = Literal['active', 'paused_rate_limit', 'paused_capacity', 'unknown'] + + +class ErrorCategory(StrEnum): + """Error attribution. Matches tinker's ``RequestErrorCategory``.""" + + Unknown = 'unknown' + Server = 'server' + User = 'user' + + +class ErrorPayload(ResponseModel): + """The single representation of a failure, on the wire and in state. + + ``error_code``, not ``status_code``: once server-request-lifecycle lands, an + execution-time failure is delivered with HTTP 200, so this value is + *systematically* unequal to the response status code. Keeping the name + ``status_code`` would make every reader misparse it once. The 400-599 range is + kept to reuse HTTP's semantic space, not to align with response codes. + + Inherits ``ResponseModel`` (``extra='ignore'``), so a future added field does + not make an old client fail to parse it. + """ + + error: str = Field(max_length=1024) + category: ErrorCategory + error_code: int = Field(ge=400, le=599) + request_id: str + traceback: str | None = Field(default=None, max_length=65536) + details: list[dict[str, Any]] | None = None + + @field_validator('category', mode='before') + @classmethod + def normalize_legacy_category(cls, value: Any) -> Any: + if isinstance(value, str): + return value.lower() + return value + + @model_validator(mode='after') + def traceback_is_server_only(self) -> ErrorPayload: + if self.traceback is not None and self.category is not ErrorCategory.Server: + raise ValueError('traceback is only valid for server errors') + return self diff --git a/src/twinkle/protocol/types/lifecycle.py b/src/twinkle/protocol/types/lifecycle.py new file mode 100644 index 000000000..32555e24b --- /dev/null +++ b/src/twinkle/protocol/types/lifecycle.py @@ -0,0 +1,83 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""The request-lifecycle wire model: one envelope for submit and retrieve. + +This module is a public-contract carrier imported across packages (Twinkle_Server +reverse-imports ``twinkle.protocol.types``); per the naming rulings in ``base.py`` it +therefore intentionally carries **no** underscore prefix. +""" +from __future__ import annotations + +from typing import Any, Literal, Optional + +from .base import ResponseModel, StrictRequest +from .errors import ErrorPayload, QueueStateLiteral + +# The lifecycle status of a queued task. Kept in sync with the server-side +# ``TaskStatus`` enum values (a consistency test asserts the two sets are equal). +TaskStatus = Literal['pending', 'queued', 'running', 'completed', 'failed', 'cancelled'] + +# The two states past which a task never changes again. ``frozenset`` so a caller +# cannot mutate the shared set. +TERMINAL_STATUSES: frozenset[str] = frozenset({'completed', 'failed', 'cancelled'}) + + +class RetrieveFutureRequest(StrictRequest): + """Body of ``POST /twinkle/retrieve_future``. + + ``request_id`` is the only field: the caller already knows which adapter it + targeted, so no ``model_id`` is needed to correlate the reply. A brand-new + endpoint with no legacy clients, so it takes the strict base (unknown fields + fail loudly) rather than tolerating extras. + """ + + request_id: str + + +class CancelRequest(StrictRequest): + """Body of ``POST /twinkle/cancel``: best-effort cancel of a not-yet-started task. + + New endpoint with no legacy clients, so it takes the strict base (unknown fields + fail loudly). + """ + + request_id: str + + +class CancelResponse(ResponseModel): + """Reply to ``POST /twinkle/cancel``. + + ``cancelled`` is True only when the task is in the terminal ``cancelled`` state + after the attempt; ``state`` is its status afterwards (``cancelled`` / ``running`` + / ``completed`` / ``failed`` / ``not_found``). A running or already-terminal task + is never interrupted -- cancel only drops tasks that have not started. + """ + + cancelled: bool + state: str + + +class TaskEnvelope(ResponseModel): + """The one lifecycle reply, shared by Submit_Endpoint and Retrieve_Endpoint. + + Success and failure live in *different* fields, and both endpoints fill the + same field for the same meaning. That is the whole point: if failure rode in + ``result`` on submit but in ``error`` on retrieve, a task that fails inside the + Inline_Fast_Path window -- which is exactly where ``step`` / ``zero_grad`` / + ``set_loss`` fail -- would have its payload read from the wrong place and + silently dropped. + + ``result`` is ``Optional[Any]`` rather than each endpoint's concrete response + model: this is the lifecycle-layer model, not a per-endpoint generic. + Deserialization to the concrete model is done by the Client_Future_Layer once + it holds a terminal envelope, since it knows the caller's expected type. + + No ``model_id``: the caller already knows which adapter it targeted; + ``request_id`` is the only key needed to correlate a reply. + """ + + request_id: str + status: TaskStatus + result: Any | None = None # set iff status == 'completed' + error: ErrorPayload | None = None # set iff status == 'failed' + queue_state: QueueStateLiteral | None = None + queue_state_reason: str | None = None diff --git a/src/twinkle/protocol/types/model.py b/src/twinkle/protocol/types/model.py new file mode 100644 index 000000000..f2e6bd976 --- /dev/null +++ b/src/twinkle/protocol/types/model.py @@ -0,0 +1,423 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Request / response models for the twinkle-native model endpoints. + +One declaration per endpoint, shared by Twinkle_Client and the server handler, so +there is a single answer to "what may this endpoint receive". Every field carries a +role (see :mod:`twinkle.protocol.types.base`): + +- plain fields are **control** fields: the handler consumes them or passes them as a + named argument, and they are never re-forwarded through ``**backend_kwargs``; +- :func:`backend_kwarg` / :func:`backend_only` fields are forwarded to the backend + when their value is not ``None``; +- :func:`passthrough` fields are declared dicts whose keys are dynamic (plugin + constructor / loss arguments) and are flattened into the backend kwargs. + +Requests are strict: an unknown top-level field is a typo and fails with 422 before +the task is enqueued. That is only safe because dynamic parameters have a declared +home -- the passthrough regions -- so strictness never blocks a legitimate +user-supplied argument. +""" +from __future__ import annotations + +from pydantic import Field, JsonValue, field_validator, model_validator +from typing import Any, Dict, List, Optional, Union + +from .base import ResponseModel, StrictRequest, backend_kwarg, backend_only, passthrough +from .component import DataRef +from .data import WireInputBatch + + +class CreateRequest(StrictRequest): + """Body of ``POST /twinkle/create``: a session-establishing no-op.""" + + +# --------------------------------------------------------------------------- # +# Control-plane requests +# --------------------------------------------------------------------------- # + + +class AdapterRequest(StrictRequest): + """The shared shape of an adapter-scoped operation. + + ``seq_id`` is an idempotency key, not a backend parameter: the submit shell + claims ``(session, adapter, seq_id)`` before enqueueing so a retried + gradient-mutating call is applied at most once. It must stay a declared field -- + under ``extra='forbid'`` an undeclared ``seq_id`` would be rejected outright, + which would silently disable that dedup. + """ + + adapter_name: str + seq_id: int | None = None + gradient_accumulation_steps: int | None = backend_kwarg(default=None, ge=1) + + +class StepRequest(AdapterRequest): + """Body of ``POST /twinkle/step``.""" + + optim_params: dict[str, JsonValue] | None = backend_kwarg(default=None) + + +class LrStepRequest(AdapterRequest): + """Body of ``POST /twinkle/lr_step``.""" + + # ``OptimizerParamScheduler.step(increment=...)``; the transformers scheduler has + # no equivalent knob. + increment: int | None = backend_only('megatron', default=None, ge=0) + + +class ClipGradNormRequest(AdapterRequest): + """Body of ``POST /twinkle/clip_grad_norm``. + + Bound to its own model rather than the bare :class:`AdapterRequest`: the endpoint + has always read these two values, and sharing a model with the parameterless ops + meant the schema could not say so. + """ + + max_grad_norm: float = Field(default=1.0, gt=0) + norm_type: int = Field(default=2, gt=0) + + +class ClipGradAndStepRequest(ClipGradNormRequest): + """Body of ``POST /twinkle/clip_grad_and_step``.""" + + optim_params: dict[str, JsonValue] | None = backend_kwarg(default=None) + + +class CalculateMetricRequest(StrictRequest): + """Body of ``POST /twinkle/calculate_metric``.""" + + adapter_name: str + is_training: bool = True + + +# --------------------------------------------------------------------------- # +# Inline forward family +# +# Three endpoints, three models. They were one shared model, which meant the schema +# could not express that only the gradient-mutating variants take ``seq_id`` or that +# ``forward_only`` does not need an adapter -- the handler had to carry that +# knowledge instead, as a second source of truth. +# +# ``ForwardBackwardRequest`` is deliberately NOT the name of the fwd-bwd model: +# ``tinker.types.ForwardBackwardRequest`` already exists and two handlers import a +# module named ``types`` from each package, so a same-named model would be +# distinguishable only by import alias. +# --------------------------------------------------------------------------- # + + +class _InlineForwardBase(StrictRequest): + """Fields common to the three inline forward endpoints.""" + + inputs: WireInputBatch + task: str | None = backend_kwarg(default=None) + temperature: float | None = backend_kwarg(default=None, gt=0) + return_logits: bool | None = backend_kwarg(default=None) + micro_batch_size: int | None = backend_kwarg(default=None, ge=1) + gradient_accumulation_steps: int | None = backend_kwarg(default=None, ge=1) + # Read only by the transformers backend. + sampling_masks: JsonValue | None = backend_only('transformers', default=None) + router_replay_action: str | None = backend_only('transformers', default=None) + # Loss inputs (``advantages`` / ``old_logps`` / ``ref_outputs`` / ...). Their key + # set is decided by the configured Loss, so they get a declared dict rather than + # top-level fields; the flattening in ``backend_kwargs`` keeps the backend call + # shape identical to before. + loss_kwargs: dict[str, JsonValue] = passthrough() + + +class ForwardRequest(_InlineForwardBase): + """Body of ``POST /twinkle/forward``: keeps the graph, mutates no gradients.""" + + adapter_name: str + disable_lora: bool | None = backend_kwarg(default=None) + + +class ForwardOnlyRequest(_InlineForwardBase): + """Body of ``POST /twinkle/forward_only``: no graph, no gradients. + + An existing ``adapter_name`` supplies the template and adapter context, even when + ``disable_lora`` requests base-weight inference. There is no ``seq_id`` because + nothing is mutated to be idempotent about. + """ + + adapter_name: str + disable_lora: bool | None = backend_kwarg(default=None) + + +class ForwardBackwardTaskRequest(_InlineForwardBase): + """Body of ``POST /twinkle/forward_backward``: accumulates gradients.""" + + adapter_name: str + seq_id: int | None = None + sync_gradients: bool | None = backend_kwarg(default=None) + loss_scale: float | None = backend_kwarg(default=None) + + +# --------------------------------------------------------------------------- # +# Data-plane forward family +# +# ``input_refs`` / ``input_field`` / ``kwarg_fields`` are control fields: the handler +# resolves them into rows and bound kwargs. No wire schema applies to a ``DataRef`` -- +# it is an opaque handle and the rows it points at never travel in this body. +# --------------------------------------------------------------------------- # + + +class DataPlaneForwardRequest(StrictRequest): + """Body of the ``*_from_data_plane`` forward endpoints.""" + + input_refs: list[DataRef] = Field(min_length=1) + input_field: str | None = None + # Values are *field paths*, not parameter values, so this is not a passthrough + # region: nothing in it is forwarded verbatim. + kwarg_fields: dict[str, str] = Field(default_factory=dict) + adapter_name: str + seq_id: int | None = None + task: str | None = backend_kwarg(default=None) + temperature: float | None = backend_kwarg(default=None, gt=0) + return_logits: bool | None = backend_kwarg(default=None) + disable_lora: bool | None = backend_kwarg(default=None) + micro_batch_size: int | None = backend_kwarg(default=None, ge=1) + gradient_accumulation_steps: int | None = backend_kwarg(default=None, ge=1) + loss_kwargs: dict[str, JsonValue] = passthrough() + + +class DataPlaneForwardOnlyRequest(StrictRequest): + """Body of ``POST /twinkle/forward_only_from_data_plane``. + + This endpoint is read-only, so it deliberately has no ``seq_id`` idempotency + key. Its fields are declared directly rather than inherited from the + gradient-mutating data-plane request. + """ + + input_refs: list[DataRef] = Field(min_length=1) + input_field: str | None = None + kwarg_fields: dict[str, str] = Field(default_factory=dict) + adapter_name: str + task: str | None = backend_kwarg(default=None) + temperature: float | None = backend_kwarg(default=None, gt=0) + return_logits: bool | None = backend_kwarg(default=None) + disable_lora: bool | None = backend_kwarg(default=None) + micro_batch_size: int | None = backend_kwarg(default=None, ge=1) + gradient_accumulation_steps: int | None = backend_kwarg(default=None, ge=1) + loss_kwargs: dict[str, JsonValue] = passthrough() + output_ref: DataRef | None = None + output_fields: dict[str, str] = Field(default_factory=dict) + + @model_validator(mode='after') + def validate_output(self) -> DataPlaneForwardOnlyRequest: + if (self.output_ref is None) != (len(self.output_fields) == 0): + raise ValueError('output_ref and output_fields must be configured together') + return self + + +# --------------------------------------------------------------------------- # +# Plugin setters +# +# Each takes the plugin identifier as a control field (the handler passes it +# positionally) plus one passthrough region for the plugin's constructor arguments. +# The passthrough keys are forwarded to the plugin as given -- there is no spelling +# check against a sibling ``target``: signature reflection cannot see a parameter a +# plugin reads straight out of ``**kwargs`` (``InputProcessor`` does this with +# ``padding_side``), so any such check rejects valid requests. A misspelt argument +# therefore surfaces from the plugin itself. +# --------------------------------------------------------------------------- # + + +class SetLossRequest(StrictRequest): + loss_cls: str + adapter_name: str + init_kwargs: dict[str, JsonValue] = passthrough() + + +class SetOptimizerRequest(StrictRequest): + optimizer_cls: str + adapter_name: str + init_kwargs: dict[str, JsonValue] = passthrough() + + +class SetLrSchedulerRequest(StrictRequest): + scheduler_cls: str + adapter_name: str + init_kwargs: dict[str, JsonValue] = passthrough() + + +class SetTemplateRequest(StrictRequest): + """Body of ``POST /twinkle/set_template``. + + No top-level ``model_id``: the backend always overrides it with its own + ``tokenizer_id``, so a declared field would advertise a parameter that has no + effect. Callers that pass ``model_id`` reach the template constructor through + ``init_kwargs`` like any other template argument. + """ + + template_cls: str + adapter_name: str + init_kwargs: dict[str, JsonValue] = passthrough() + + +class SetProcessorRequest(StrictRequest): + processor_cls: str + adapter_name: str + init_kwargs: dict[str, JsonValue] = passthrough() + + +class AddMetricRequest(StrictRequest): + metric_cls: str + adapter_name: str + is_training: bool | None = None + init_kwargs: dict[str, JsonValue] = passthrough() + + +class ApplyPatchRequest(StrictRequest): + patch_cls: str + adapter_name: str + init_kwargs: dict[str, JsonValue] = passthrough() + + +# --------------------------------------------------------------------------- # +# Checkpoint I/O and adapter lifecycle +# --------------------------------------------------------------------------- # + + +class SaveRequest(StrictRequest): + adapter_name: str + name: str | None = None + save_optimizer: bool = False + is_sampler: bool = False # If True, delete existing sampler weights before saving + consumed_train_samples: int | None = backend_kwarg(default=None, ge=0) + merge_lora: bool | None = backend_only('megatron', default=None) + + +class LoadRequest(StrictRequest): + adapter_name: str + name: str + load_optimizer: bool = False + no_load_optim: bool | None = backend_only('megatron', default=None) + no_load_rng: bool | None = backend_only('megatron', default=None) + strict: bool | None = backend_only('transformers', default=None) + + +class ResumeFromCheckpointRequest(StrictRequest): + """Body of ``POST /twinkle/resume_from_checkpoint``.""" + + name: str + adapter_name: str = '' + resume_only_model: bool = False + + +class AddAdapterRequest(StrictRequest): + adapter_name: str + # ``config`` is None for full-parameter training (no LoRA adapter) and a + # serialized LoraConfig string for LoRA training. + config: str | None = None + save_dir: str | None = None + gradient_accumulation_steps: int | None = backend_kwarg(default=None, ge=1) + init_kwargs: dict[str, JsonValue] = passthrough() + + +class UploadToHubRequest(StrictRequest): + """Body of ``POST /twinkle/upload_to_hub``. + + No ``async_upload``: the server always runs the upload as a background task and + the client waits through the future layer, so the flag could only ever be ignored. + """ + + checkpoint_dir: str | dict[str, Any] + hub_model_id: str + hub_token: str | None = None + + @field_validator('checkpoint_dir', mode='before') + @classmethod + def extract_checkpoint_dir(cls, v): + """Accept a ``save`` response dict and take its twinkle path. + + Raises a validation error -- not ``KeyError`` -- when the key is absent, so a + wrong-shaped dict is a 422 naming the missing key instead of a 500. + """ + if isinstance(v, dict): + if 'twinkle_path' not in v: + raise ValueError("checkpoint_dir dict must contain 'twinkle_path'") + return v['twinkle_path'] + return v + + +# --------------------------------------------------------------------------- # +# Response models +# --------------------------------------------------------------------------- # + + +class OkResponse(ResponseModel): + """Response for endpoints whose underlying method returns None.""" + status: str = 'ok' + + +class ModelResult(ResponseModel): + """Generic result wrapper; ``ModelResult`` is the retained historical public name.""" + result: Any + + +# --- Result-bearing responses --- + + +class ForwardResponse(ResponseModel): + """Response for /forward and /forward_only endpoints (returns ModelOutput).""" + result: Any + + +class ForwardBackwardResponse(ResponseModel): + """Response for /forward_backward endpoint (returns ModelOutput).""" + result: Any + + +class CalculateLossResponse(ResponseModel): + """Response for /calculate_loss endpoint (returns float).""" + result: float + + +class ClipGradNormResponse(ResponseModel): + """Response for /clip_grad_norm endpoint (returns float as str).""" + result: str + + +class GetTrainConfigsResponse(ResponseModel): + """Response for /get_train_configs endpoint (returns str).""" + result: str + + +class CalculateMetricResponse(ResponseModel): + """Response for /calculate_metric endpoint (returns Dict).""" + result: dict[str, Any] + + +class SaveResponse(ResponseModel): + """Response for /save endpoint (returns twinkle path + checkpoint dir).""" + twinkle_path: str + checkpoint_dir: str | None = None + + +class TrainingProgressResponse(ResponseModel): + """Response for /resume_from_checkpoint endpoint.""" + result: dict[str, Any] + + +# --- Void responses (return None → OkResponse) --- + +BackwardResponse = OkResponse +StepResponse = OkResponse +ZeroGradResponse = OkResponse +LrStepResponse = OkResponse +SetLossResponse = OkResponse +SetOptimizerResponse = OkResponse +SetLrSchedulerResponse = OkResponse +LoadResponse = OkResponse +SetTemplateResponse = OkResponse +SetProcessorResponse = OkResponse +ClipGradAndStepResponse = OkResponse +ApplyPatchResponse = OkResponse +AddMetricResponse = OkResponse + +# --- Other responses --- + + +class CreateResponse(ResponseModel): + """Response for /create endpoint.""" + status: str = 'ok' diff --git a/src/twinkle/protocol/types/processor.py b/src/twinkle/protocol/types/processor.py new file mode 100644 index 000000000..de7de0d96 --- /dev/null +++ b/src/twinkle/protocol/types/processor.py @@ -0,0 +1,55 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Request / response models for the twinkle processor endpoints. + +The processor surface is a generic RPC bridge: ``create`` names a class to build and +``call`` names a method to invoke, both with caller-supplied arguments whose names are +only known to the target. Those arguments therefore live in declared passthrough +dicts (``init_kwargs`` / ``call_kwargs``) rather than being spread over the top level, +which is what lets the envelope itself be strict -- a misspelt ``processor_id`` or +``function`` now fails instead of being silently treated as an argument. + +No ``target``: the callable is resolved from ``processor_type`` + ``class_type`` (or a +live instance plus ``function``), not from a single sibling field, so the passthrough +keys are forwarded unchecked. Claiming otherwise would need a second resolution path +that guesses. + +Class names are prefixed with ``Processor`` to avoid collisions when importing from +``twinkle.protocol.types`` alongside ``model.py``. +""" +from __future__ import annotations + +from pydantic import JsonValue +from typing import Any, Dict + +from .base import ResponseModel, StrictRequest, passthrough + + +class ProcessorCreateRequest(StrictRequest): + processor_type: str + class_type: str + init_kwargs: dict[str, JsonValue] = passthrough() + + +class ProcessorHeartbeatRequest(StrictRequest): + processor_id: str + + +class ProcessorCallRequest(StrictRequest): + processor_id: str + function: str + call_kwargs: dict[str, JsonValue] = passthrough() + + +class ProcessorCreateResponse(ResponseModel): + """Response body for the /create endpoint.""" + processor_id: str + + +class ProcessorHeartbeatResponse(ResponseModel): + """Response body for the /heartbeat endpoint.""" + status: str = 'ok' + + +class ProcessorCallResponse(ResponseModel): + """Response body for the /call endpoint.""" + result: Any diff --git a/src/twinkle/protocol/types/sampler.py b/src/twinkle/protocol/types/sampler.py new file mode 100644 index 000000000..b89602b68 --- /dev/null +++ b/src/twinkle/protocol/types/sampler.py @@ -0,0 +1,94 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Request / response models for the twinkle-native sampler endpoints. + +Shared by the server handler and the twinkle client. Field roles follow +:mod:`twinkle.protocol.types.base`; the sampler handlers pass everything they need +explicitly, so these requests carry control fields and -- for the template setter -- +one passthrough region, and no free-floating backend kwargs. + +Class names carry a ``Sampler`` prefix wherever ``model.py`` already owns the bare +name (``AddAdapterRequest``, ``SetTemplateRequest``, ``CreateResponse`` and their +responses). The two modules describe *different* endpoints with different field sets; +a shared bare name is distinguished only by an import alias and, when a handler does +``import twinkle.protocol.types as types``, silently resolves to whichever module the +package ``__init__`` re-exported first -- which is how the sampler endpoints once +bound ``model.py``'s schema. Prefixing at the definition site removes the ambiguity, +matching :mod:`twinkle.protocol.types.processor`. +""" +from __future__ import annotations + +from pydantic import Field, JsonValue +from typing import Any, Dict, List, Literal, Optional, Tuple + +from .base import ResponseModel, StrictRequest, passthrough +from .data import WireInputBatch + +StopReason = Literal['length', 'stop', 'abort', 'error'] + + +class SampleRequest(StrictRequest): + """Request body for the ``/sample`` and ``/sample_stream`` endpoints. + + ``num_samples`` is not a top-level field: it is a sampling parameter and + ``SamplingParams.from_dict(sampling_params)`` is the one place sampling + parameters are built. A second, top-level spelling would be a second source of + truth for the same value. + """ + + inputs: WireInputBatch = Field(..., description='Trajectory or InputFeature entries to sample from') + sampling_params: dict[str, JsonValue] | None = Field( + None, description='Sampling parameters (max_tokens, temperature, num_samples, etc.)') + adapter_name: str = Field('', description='Adapter name for LoRA inference') + adapter_uri: str | None = Field(None, description='Adapter URI (twinkle:// path or local path) for LoRA inference') + + +class SampledSequenceModel(ResponseModel): + """A single sampled sequence, mirroring twinkle.data_format.SampledSequence.""" + stop_reason: StopReason = Field(..., description="Stop reason: 'length' or 'stop'") + tokens: list[int] = Field(..., description='Token IDs of the sampled sequence') + logprobs: list[list[tuple[int, float]] | None] | None = Field(None, description='Per-token log-probabilities') + decoded: str | None = Field(None, description='Decoded text of the sampled sequence') + new_input_feature: dict[str, Any] | None = Field( + None, description='Updated InputFeature after sampling (input_ids, labels, etc.)') + + +class SampleResponseModel(ResponseModel): + """Mirroring twinkle.data_format.SampleResponse.""" + sequences: list[SampledSequenceModel] = Field(..., description='List of sampled sequences') + prompt_token_ids: list[int] | None = Field(None, description='Token IDs of the prompt the sequences continue') + prompt_logprobs: list[float | None] | None = None + topk_prompt_logprobs: list[list[tuple[int, float]] | None] | None = None + + +class SampleResponseModelList(ResponseModel): + """Response body for the /sample endpoint""" + samples: list[SampleResponseModel] = Field(..., description='List of sample responses') + + +class SamplerSetTemplateRequest(StrictRequest): + """Request body for the sampler ``/set_template`` endpoint.""" + template_cls: str = Field(..., description="Template class name (e.g. 'Template')") + adapter_name: str = Field('', description='Adapter name to associate the template with') + init_kwargs: dict[str, JsonValue] = passthrough() + + +class SamplerSetTemplateResponse(ResponseModel): + """Response body for the sampler /set_template endpoint.""" + status: str = 'ok' + + +class SamplerAddAdapterRequest(StrictRequest): + """Request body for the ``/add_adapter_to_sampler`` endpoint.""" + adapter_name: str = Field(..., description='Name of the adapter to add') + config: Any = Field(..., description='LoRA configuration dict') + + +class SamplerAddAdapterResponse(ResponseModel): + """Response body for the /add_adapter_to_sampler endpoint.""" + status: str = 'ok' + adapter_name: str + + +class SamplerCreateResponse(ResponseModel): + """Response body for the sampler /create endpoint.""" + status: str = 'ok' diff --git a/src/twinkle/protocol/types/server.py b/src/twinkle/protocol/types/server.py new file mode 100644 index 000000000..dd335d585 --- /dev/null +++ b/src/twinkle/protocol/types/server.py @@ -0,0 +1,63 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Shared Pydantic response models for the twinkle server health/error endpoints.""" +from pydantic import BaseModel, Field +from typing import List + +from .base import ResponseModel, StrictRequest + + +class SupportedModel(BaseModel): + """Information about a supported model. + + A nested value inside a response, not a response body of its own, so it keeps the + plain base -- the strict/ignore split applies to what crosses the wire as a whole. + """ + model_name: str + + +class ClientFeatures(ResponseModel): + task_envelope: bool = True + cancel: bool = False + data_plane: bool = False + full_training: bool = False + batch_retrieve: bool = False + + +class ProtocolLimits(ResponseModel): + long_poll_timeout_seconds: float | None = None + max_payload_bytes: int | None = None + max_batch_size: int | None = None + + +class GetServerCapabilitiesResponse(ResponseModel): + """Versioned Twinkle-native capabilities with old-server defaults.""" + supported_models: List[SupportedModel] + protocol_version: int = 1 + features: ClientFeatures = Field(default_factory=ClientFeatures) + limits: ProtocolLimits = Field(default_factory=ProtocolLimits) + + +class HealthResponse(ResponseModel): + status: str + + +class DeleteCheckpointResponse(ResponseModel): + success: bool + message: str + + +class WeightsInfoRequest(StrictRequest): + twinkle_path: str + + +class CheckpointPathResponse(ResponseModel): + """Response body for the /checkpoint_path endpoint.""" + path: str + twinkle_path: str + + +class CapacityInfoResponse(ResponseModel): + """Response body for the /capacity_info endpoint.""" + max_loras: int + used_loras: int + free_loras: int diff --git a/src/twinkle_client/types/session.py b/src/twinkle/protocol/types/session.py similarity index 62% rename from src/twinkle_client/types/session.py rename to src/twinkle/protocol/types/session.py index f6b1adb72..07e940ded 100644 --- a/src/twinkle_client/types/session.py +++ b/src/twinkle/protocol/types/session.py @@ -1,24 +1,24 @@ # Copyright (c) ModelScope Contributors. All rights reserved. """Pydantic models for twinkle session management endpoints.""" -from pydantic import BaseModel from typing import Any, Dict, Optional +from .base import ResponseModel, StrictRequest -class CreateSessionRequest(BaseModel): + +class CreateSessionRequest(StrictRequest): """Request body for POST /twinkle/create_session.""" - metadata: Optional[Dict[str, Any]] = None + metadata: dict[str, Any] | None = None -class CreateSessionResponse(BaseModel): +class CreateSessionResponse(ResponseModel): """Response body for POST /twinkle/create_session.""" session_id: str -class SessionHeartbeatRequest(BaseModel): +class SessionHeartbeatRequest(StrictRequest): """Request body for POST /twinkle/session_heartbeat.""" session_id: str -class SessionHeartbeatResponse(BaseModel): +class SessionHeartbeatResponse(ResponseModel): """Response body for POST /twinkle/session_heartbeat.""" - pass diff --git a/src/twinkle_client/types/training.py b/src/twinkle/protocol/types/training.py similarity index 90% rename from src/twinkle_client/types/training.py rename to src/twinkle/protocol/types/training.py index bab59ac63..6cdad789c 100644 --- a/src/twinkle_client/types/training.py +++ b/src/twinkle/protocol/types/training.py @@ -3,12 +3,14 @@ Shared Pydantic models for twinkle training runs and checkpoints. These types are used both by twinkle_client (as request/response shapes) -and by twinkle.server.common.io_utils (as persistence models). +and by twinkle.server.checkpoint.twinkle (as persistence / response models). """ from datetime import datetime from pydantic import BaseModel from typing import Any, Dict, List, Optional +from .base import ResponseModel + class Cursor(BaseModel): limit: int @@ -49,12 +51,12 @@ class TrainingRun(BaseModel): user_metadata: Optional[Dict[str, Any]] = None -class TrainingRunsResponse(BaseModel): +class TrainingRunsResponse(ResponseModel): training_runs: List[TrainingRun] cursor: Cursor -class CheckpointsListResponse(BaseModel): +class CheckpointsListResponse(ResponseModel): checkpoints: List[Checkpoint] cursor: Optional[Cursor] = None @@ -68,7 +70,7 @@ class ParsedCheckpointTwinklePath(BaseModel): checkpoint_id: str -class WeightsInfoResponse(BaseModel): +class WeightsInfoResponse(ResponseModel): """Twinkle weights info response.""" training_run_id: str base_model: str diff --git a/src/twinkle/sampler/base.py b/src/twinkle/sampler/base.py index a756a93ce..63fa344e2 100644 --- a/src/twinkle/sampler/base.py +++ b/src/twinkle/sampler/base.py @@ -5,7 +5,7 @@ import twinkle from twinkle import remote_function -from twinkle.data_format import InputFeature, SampleResponse, SamplingParams, Trajectory +from twinkle.data_format import InputFeature, SampleResponse, SamplingParams, Trajectory, is_encoded from twinkle.patch import Patch from twinkle.template import Template from twinkle.utils import construct_class @@ -51,10 +51,11 @@ def apply_patch(self, patch_cls: Union[Patch, Type[Patch], str], **kwargs) -> No def _not_encoded(inputs: Any) -> bool: """Check if inputs are not yet encoded (i.e., is Trajectory, not InputFeature). - Aligned with TransformersModel._not_encoded for consistency. + Delegates to the single shared predicate so the three backends and the wire + schema cannot drift apart. """ assert isinstance(inputs, dict), f'Expected dict, got {type(inputs)}' - return 'input_ids' not in inputs and 'input_embedding' not in inputs + return not is_encoded(inputs) def _is_trajectory(self, inputs: Any) -> bool: """Check if inputs are Trajectory type (not encoded).""" diff --git a/src/twinkle/sampler/vllm_sampler/vllm_sampler.py b/src/twinkle/sampler/vllm_sampler/vllm_sampler.py index 7877ddb55..6438fa660 100644 --- a/src/twinkle/sampler/vllm_sampler/vllm_sampler.py +++ b/src/twinkle/sampler/vllm_sampler/vllm_sampler.py @@ -251,6 +251,14 @@ async def _sample_single( else: feat['input_ids'] = response.prompt_token_ids feat['labels'] = [-100] * len(response.prompt_token_ids) + # A sampling prompt (e.g. a tinker ModelInput) carries input_ids but no labels; + # concat_input_feature would then derive a zero-length prefix completion_mask + # and raise. The prompt is pure context, so materialise aligned all-context + # labels when they are missing or length-mismatched. Present, aligned labels are + # left untouched (preserving provenance), and the logprobs-only path -- which + # never concatenates -- is not touched. + if not logprobs_only and 'input_ids' in feat and len(feat.get('labels') or []) != len(feat['input_ids']): + feat['labels'] = [-100] * len(feat['input_ids']) sequences = [] for seq in response.sequences: if logprobs_only: @@ -495,7 +503,7 @@ def unload_adapter_paths(self, adapter_paths: list[str]) -> None: """Unload policy snapshots from vLLM and clear cached requests.""" self._run_in_loop(self.engine.unload_lora_paths(adapter_paths)) - @remote_function(dispatch='all', collect='first', lazy_collect=False) + @remote_function(dispatch='all', collect='first', lazy_collect=False, timeout=3600) def load_full_weights_from_path(self, path: Optional[str] = None) -> int: """Load a full (non-LoRA) HF checkpoint into the engine's base model. diff --git a/src/twinkle/server/checkpoint/__init__.py b/src/twinkle/server/checkpoint/__init__.py index e8bfccd1e..473083671 100644 --- a/src/twinkle/server/checkpoint/__init__.py +++ b/src/twinkle/server/checkpoint/__init__.py @@ -1,5 +1,5 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -"""Checkpoint subsystem (TIER 2 consolidation). +"""Checkpoint subsystem. Top-level package consolidating the checkpoint base classes (split from the former 1017-line ``utils/checkpoint_base.py``, now deleted) with the concrete @@ -10,7 +10,7 @@ from twinkle.server.checkpoint import ( create_checkpoint_manager, create_training_run_manager, BaseCheckpointManager, BaseTrainingRunManager, BaseFileManager, - validate_user_path, validate_ownership, _resolve_client_save_dir, + validate_user_path, validate_ownership, TRAIN_RUN_INFO_FILENAME, TWINKLE_DEFAULT_SAVE_DIR, ) """ @@ -30,5 +30,4 @@ 'BaseTrainingRunManager', 'validate_user_path', 'validate_ownership', - '_resolve_client_save_dir', ] diff --git a/src/twinkle/server/checkpoint/checkpoint_manager.py b/src/twinkle/server/checkpoint/checkpoint_manager.py index 038bc3093..8b0210082 100644 --- a/src/twinkle/server/checkpoint/checkpoint_manager.py +++ b/src/twinkle/server/checkpoint/checkpoint_manager.py @@ -1,7 +1,5 @@ # Copyright (c) ModelScope Contributors. All rights reserved. """Abstract base checkpoint manager. - -Relocated from ``utils/checkpoint_base.py`` (TIER 2 consolidation). No logic change. """ from __future__ import annotations @@ -16,7 +14,7 @@ from twinkle import get_logger from twinkle.hub import HubOperation -from twinkle_client.types import ResolvedLoadPath +from twinkle.protocol.types import ResolvedLoadPath from .paths import CHECKPOINT_INFO_FILENAME, validate_user_path from .training_run_manager import BaseFileManager, BaseTrainingRunManager @@ -29,7 +27,6 @@ class BaseCheckpointManager(BaseFileManager, ABC): Subclasses must implement: - path_prefix property - - path_field_name property - _create_checkpoint method - _parse_checkpoint method - _create_checkpoints_response method @@ -54,12 +51,6 @@ def path_prefix(self) -> str: """Return the path prefix (e.g., 'twinkle://').""" pass - @property - @abstractmethod - def path_field_name(self) -> str: - """Return the field name for the path (e.g., 'twinkle_path' or 'tinker_path').""" - pass - @abstractmethod def _create_checkpoint(self, checkpoint_id: str, diff --git a/src/twinkle/server/checkpoint/models.py b/src/twinkle/server/checkpoint/models.py index fdb885c20..fa42bc28b 100644 --- a/src/twinkle/server/checkpoint/models.py +++ b/src/twinkle/server/checkpoint/models.py @@ -1,8 +1,6 @@ # Copyright (c) ModelScope Contributors. All rights reserved. """Internal Pydantic base specs used as type constraints for the generic checkpoint / training-run managers. - -Relocated from ``utils/checkpoint_base.py`` (TIER 2 consolidation). No logic change. """ from __future__ import annotations diff --git a/src/twinkle/server/checkpoint/paths.py b/src/twinkle/server/checkpoint/paths.py index 202f1747f..d22d4c413 100644 --- a/src/twinkle/server/checkpoint/paths.py +++ b/src/twinkle/server/checkpoint/paths.py @@ -1,8 +1,6 @@ # Copyright (c) ModelScope Contributors. All rights reserved. """Path constants, token hashing, client-save-dir resolution, and permission helpers for the checkpoint subsystem. - -Relocated from ``utils/checkpoint_base.py`` (TIER 2 consolidation). No logic change. """ from __future__ import annotations diff --git a/src/twinkle/server/checkpoint/tinker.py b/src/twinkle/server/checkpoint/tinker.py index 518a60d7c..ac72e46e8 100644 --- a/src/twinkle/server/checkpoint/tinker.py +++ b/src/twinkle/server/checkpoint/tinker.py @@ -8,7 +8,9 @@ from tinker import types as tinker_types from typing import Any, Dict, List, Optional -from twinkle.server.checkpoint import TRAIN_RUN_INFO_FILENAME, BaseCheckpointManager, BaseTrainingRunManager +from twinkle.server.checkpoint.checkpoint_manager import BaseCheckpointManager +from twinkle.server.checkpoint.paths import TRAIN_RUN_INFO_FILENAME +from twinkle.server.checkpoint.training_run_manager import BaseTrainingRunManager class TinkerTrainingRunManager(BaseTrainingRunManager): @@ -73,10 +75,6 @@ class TinkerCheckpointManager(BaseCheckpointManager): def path_prefix(self) -> str: return 'twinkle://' - @property - def path_field_name(self) -> str: - return 'tinker_path' - def _create_checkpoint(self, checkpoint_id, checkpoint_type, @@ -130,6 +128,3 @@ def _create_parsed_path(self, path, training_run_id, checkpoint_type, def _create_weights_info(self, run_info: dict[str, Any]) -> tinker_types.WeightsInfoResponse: return tinker_types.WeightsInfoResponse(**run_info) - - def parse_tinker_path(self, tinker_path: str) -> tinker_types.ParsedCheckpointTinkerPath | None: - return self.parse_path(tinker_path) diff --git a/src/twinkle/server/checkpoint/training_run_manager.py b/src/twinkle/server/checkpoint/training_run_manager.py index 6526fb860..7075268e4 100644 --- a/src/twinkle/server/checkpoint/training_run_manager.py +++ b/src/twinkle/server/checkpoint/training_run_manager.py @@ -1,7 +1,5 @@ # Copyright (c) ModelScope Contributors. All rights reserved. """Base file manager and abstract training-run manager. - -Relocated from ``utils/checkpoint_base.py`` (TIER 2 consolidation). No logic change. """ from __future__ import annotations diff --git a/src/twinkle/server/checkpoint/twinkle.py b/src/twinkle/server/checkpoint/twinkle.py index c9036c6ab..e4fedfc96 100644 --- a/src/twinkle/server/checkpoint/twinkle.py +++ b/src/twinkle/server/checkpoint/twinkle.py @@ -2,16 +2,17 @@ """ Twinkle-specific checkpoint and training-run managers. -Uses ``twinkle_client.types.training`` models for all serialization and response construction. +Uses ``twinkle.protocol.types.training`` models for all serialization and response construction. """ from datetime import datetime from typing import Any, Dict, List, Optional -from twinkle.server.checkpoint import (TRAIN_RUN_INFO_FILENAME, BaseCheckpointManager, BaseTrainingRunManager, - validate_ownership) -from twinkle_client.types.training import (Checkpoint, CheckpointsListResponse, CreateModelRequest, Cursor, - ParsedCheckpointTwinklePath, TrainingRun, TrainingRunsResponse, - WeightsInfoResponse) +from twinkle.protocol.types.training import (Checkpoint, CheckpointsListResponse, CreateModelRequest, Cursor, + ParsedCheckpointTwinklePath, TrainingRun, TrainingRunsResponse, + WeightsInfoResponse) +from twinkle.server.checkpoint.checkpoint_manager import BaseCheckpointManager +from twinkle.server.checkpoint.paths import TRAIN_RUN_INFO_FILENAME, validate_ownership +from twinkle.server.checkpoint.training_run_manager import BaseTrainingRunManager class TwinkleTrainingRunManager(BaseTrainingRunManager): @@ -64,10 +65,6 @@ class TwinkleCheckpointManager(BaseCheckpointManager): def path_prefix(self) -> str: return 'twinkle://' - @property - def path_field_name(self) -> str: - return 'twinkle_path' - def _create_checkpoint(self, checkpoint_id, checkpoint_type, diff --git a/src/twinkle/server/common/__init__.py b/src/twinkle/server/common/__init__.py deleted file mode 100644 index 5fa5109d2..000000000 --- a/src/twinkle/server/common/__init__.py +++ /dev/null @@ -1,10 +0,0 @@ -# Copyright (c) ModelScope Contributors. All rights reserved. -from .datum import datum_to_input_feature, extract_rl_features_for_loss, input_feature_to_datum -from .router import StickyLoraRequestRouter - -__all__ = [ - 'datum_to_input_feature', - 'extract_rl_features_for_loss', - 'input_feature_to_datum', - 'StickyLoraRequestRouter', -] diff --git a/src/twinkle/server/config/application_spec.py b/src/twinkle/server/config/application_spec.py index 5e6fa89bf..397fe8b0e 100644 --- a/src/twinkle/server/config/application_spec.py +++ b/src/twinkle/server/config/application_spec.py @@ -11,10 +11,23 @@ """ from __future__ import annotations +import os from pydantic import BaseModel, ConfigDict, Field, model_validator from typing import Any, Literal -from twinkle.server.utils.task_queue.config import TaskQueueConfig +from twinkle.server.task_queue.config import TaskQueueConfig + +# Env var keys the launcher sets from the gateway ``server_config`` so that any +# Ray worker (model / sampler / processor), not just the gateway, applies the +# configured ServerState policy instead of ``get_server_state``'s hardcoded +# defaults. Only the four *policy* fields are propagated; ``actor_name`` is a +# per-process cache key, not a cross-worker policy, so it is excluded. +SERVER_STATE_ENV_KEYS: tuple[str, ...] = ( + 'TWINKLE_SERVER_STATE_EXPIRATION_TIMEOUT', + 'TWINKLE_SERVER_STATE_CLEANUP_INTERVAL', + 'TWINKLE_SERVER_STATE_PER_TOKEN_MODEL_LIMIT', + 'TWINKLE_SERVER_STATE_METRICS_UPDATE_INTERVAL', +) # ---------- shared helpers ------------------------------------------------- # @@ -72,7 +85,7 @@ class SamplerArgs(_ArgsBase): nproc_per_node: int = 1 device_group: dict[str, Any] device_mesh: dict[str, Any] - sampler_type: Literal['mock', 'vllm', 'vllm_async', 'torch'] + sampler_type: Literal['mock', 'vllm', 'vllm_async'] engine_args: dict[str, Any] | None = None queue_config: TaskQueueConfig = Field(default_factory=TaskQueueConfig) data_plane_url: str | None = None @@ -95,6 +108,47 @@ class ServerStateArgs(_ArgsBase): metrics_update_interval: float | None = None actor_name: str | None = None + def to_env_vars(self) -> dict[str, str]: + """Serialize the policy fields to env vars for Ray-worker propagation. + + Mirrors :meth:`PersistenceConfig.to_env_vars`. Only the four policy + fields are emitted (``actor_name`` is a per-process cache key, not a + policy); unset (``None``) fields are skipped so the worker falls back to + ``ServerState``'s own defaults. + """ + env: dict[str, str] = {} + if self.expiration_timeout is not None: + env['TWINKLE_SERVER_STATE_EXPIRATION_TIMEOUT'] = str(self.expiration_timeout) + if self.cleanup_interval is not None: + env['TWINKLE_SERVER_STATE_CLEANUP_INTERVAL'] = str(self.cleanup_interval) + if self.per_token_model_limit is not None: + env['TWINKLE_SERVER_STATE_PER_TOKEN_MODEL_LIMIT'] = str(self.per_token_model_limit) + if self.metrics_update_interval is not None: + env['TWINKLE_SERVER_STATE_METRICS_UPDATE_INTERVAL'] = str(self.metrics_update_interval) + return env + + @classmethod + def from_env(cls) -> ServerStateArgs | None: + """Reconstruct policy fields from launcher-set env vars. + + Returns ``None`` when no ``TWINKLE_SERVER_STATE_*`` key is set, so a + caller can distinguish "no env-configured policy" from "explicitly + configured to a default value". Only policy fields are populated; + ``actor_name`` stays ``None``. + """ + exp = os.environ.get('TWINKLE_SERVER_STATE_EXPIRATION_TIMEOUT') + clean = os.environ.get('TWINKLE_SERVER_STATE_CLEANUP_INTERVAL') + limit = os.environ.get('TWINKLE_SERVER_STATE_PER_TOKEN_MODEL_LIMIT') + interval = os.environ.get('TWINKLE_SERVER_STATE_METRICS_UPDATE_INTERVAL') + if exp is None and clean is None and limit is None and interval is None: + return None + return cls( + expiration_timeout=float(exp) if exp is not None else None, + cleanup_interval=float(clean) if clean is not None else None, + per_token_model_limit=int(limit) if limit is not None else None, + metrics_update_interval=float(interval) if interval is not None else None, + ) + class ServerArgs(_ArgsBase): """Args for the gateway ``server`` deployment.""" @@ -106,12 +160,17 @@ class ServerArgs(_ArgsBase): class ProcessorArgs(_ArgsBase): - """Args for the ``processor`` deployment.""" + """Args for the ``processor`` deployment. + + A processor deployment has no task queue, so there is deliberately no + ``queue_config`` field here: with ``extra='forbid'`` a YAML that sets + ``processor.args.queue_config`` now fails validation with the offending path + instead of being silently ignored. + """ ncpu_proc_per_node: int | None = None device_group: dict[str, Any] | None = None device_mesh: dict[str, Any] | None = None - queue_config: TaskQueueConfig = Field(default_factory=TaskQueueConfig) class DataPlaneArgs(_ArgsBase): diff --git a/src/twinkle/server/utils/backend_dispatch.py b/src/twinkle/server/config/backend_dispatch.py similarity index 96% rename from src/twinkle/server/utils/backend_dispatch.py rename to src/twinkle/server/config/backend_dispatch.py index 65613e570..b4da33f90 100644 --- a/src/twinkle/server/utils/backend_dispatch.py +++ b/src/twinkle/server/config/backend_dispatch.py @@ -1,5 +1,5 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -"""Generic validate-then-dispatch helper for backend selectors (TIER 1, R12). +"""Generic validate-then-dispatch helper for backend selectors. The Model backend selector (``mock | transformers | megatron``) and the Sampler type selector (``mock | vllm | torch``) share one validate-then-dispatch shape. diff --git a/src/twinkle/server/config/persistence.py b/src/twinkle/server/config/persistence.py index c61c18c1c..cfb77192e 100644 --- a/src/twinkle/server/config/persistence.py +++ b/src/twinkle/server/config/persistence.py @@ -11,7 +11,6 @@ # ServerState actor first. PERSISTENCE_ENV_KEYS: tuple[str, ...] = ( 'TWINKLE_PERSISTENCE_MODE', - 'TWINKLE_PERSISTENCE_FILE_PATH', 'TWINKLE_PERSISTENCE_REDIS_URL', 'TWINKLE_PERSISTENCE_KEY_PREFIX', ) @@ -22,16 +21,13 @@ class PersistenceConfig(BaseModel): model_config = ConfigDict(extra='forbid') - mode: Literal['memory', 'file', 'redis'] = 'memory' - file_path: str | None = None # required for file mode + mode: Literal['memory', 'redis'] = 'memory' redis_url: str | None = None # required for redis mode key_prefix: str = '' # optional global key prefix def to_env_vars(self) -> dict[str, str]: """Serialize this config to env var key/value pairs for worker propagation.""" env: dict[str, str] = {'TWINKLE_PERSISTENCE_MODE': self.mode} - if self.file_path: - env['TWINKLE_PERSISTENCE_FILE_PATH'] = self.file_path if self.redis_url: env['TWINKLE_PERSISTENCE_REDIS_URL'] = self.redis_url if self.key_prefix: @@ -50,7 +46,6 @@ def from_env(cls) -> PersistenceConfig | None: return None return cls( mode=mode, - file_path=os.environ.get('TWINKLE_PERSISTENCE_FILE_PATH'), redis_url=os.environ.get('TWINKLE_PERSISTENCE_REDIS_URL'), key_prefix=os.environ.get('TWINKLE_PERSISTENCE_KEY_PREFIX', ''), ) diff --git a/src/twinkle/server/config/server_config.py b/src/twinkle/server/config/server_config.py index 15206c7fc..82bac4d44 100644 --- a/src/twinkle/server/config/server_config.py +++ b/src/twinkle/server/config/server_config.py @@ -19,7 +19,7 @@ from typing import Any from twinkle.server.exceptions import ConfigParseError -from twinkle.server.utils.task_queue.config import TaskQueueConfig +from twinkle.server.task_queue.config import TaskQueueConfig from .application_spec import ApplicationSpec, HttpOptions from .persistence import PersistenceConfig from .telemetry import TelemetryConfig @@ -71,8 +71,6 @@ def from_yaml(cls, path: str | Path) -> ServerConfig: def _validate_cross_field(self) -> ServerConfig: if self.persistence.mode == 'redis' and not self.persistence.redis_url: raise ValueError("persistence.redis_url is required when persistence.mode == 'redis'", ) - if self.persistence.mode == 'file' and not self.persistence.file_path: - raise ValueError("persistence.file_path is required when persistence.mode == 'file'", ) return self # ---- round-trip / serialization -------------------------------------- # diff --git a/src/twinkle/server/data_plane/handlers.py b/src/twinkle/server/data_plane/handlers.py index aaf214ed7..508b4dde4 100644 --- a/src/twinkle/server/data_plane/handlers.py +++ b/src/twinkle/server/data_plane/handlers.py @@ -5,7 +5,7 @@ from fastapi import Depends, FastAPI from typing import TYPE_CHECKING -import twinkle_client.types as types +import twinkle.protocol.types as types if TYPE_CHECKING: from .app import DataPlaneManagement diff --git a/src/twinkle/server/data_plane/proxy.py b/src/twinkle/server/data_plane/proxy.py index 0aa2ae7f2..3d1aa97d4 100644 --- a/src/twinkle/server/data_plane/proxy.py +++ b/src/twinkle/server/data_plane/proxy.py @@ -5,8 +5,8 @@ import httpx from typing import Any -from twinkle_client.http.headers import build_routing_headers -from twinkle_client.types.component import DataRef +from twinkle.protocol.headers import build_routing_headers +from twinkle.protocol.types.component import DataRef class DataPlaneProxy: diff --git a/src/twinkle/server/data_plane/store.py b/src/twinkle/server/data_plane/store.py index 81fe7fa60..291a6deb0 100644 --- a/src/twinkle/server/data_plane/store.py +++ b/src/twinkle/server/data_plane/store.py @@ -5,9 +5,9 @@ import uuid from typing import Any -from twinkle.tq_utils import rows_to_tq_fields -from twinkle_client.common.json_utils import json_safe -from twinkle_client.types.component import DataRef +from twinkle.data_format import rows_to_tq_fields +from twinkle.protocol.json_utils import json_safe +from twinkle.protocol.types.component import DataRef def _keys(ref: DataRef) -> list[str]: diff --git a/src/twinkle/server/deployment.py b/src/twinkle/server/deployment.py index ccd700964..45ca43fbc 100644 --- a/src/twinkle/server/deployment.py +++ b/src/twinkle/server/deployment.py @@ -31,13 +31,17 @@ from collections.abc import Awaitable, Callable from contextlib import asynccontextmanager from fastapi import FastAPI, Request +from fastapi.exceptions import RequestValidationError from fastapi.responses import JSONResponse from ray import serve from typing import Any -from twinkle.server.telemetry.middleware import create_metrics_middleware +from twinkle.protocol.types.errors import ErrorCategory +from twinkle.server.exceptions import TwinkleServerError +from twinkle.server.middleware.auth import verify_request_token +from twinkle.server.task_errors import build_error_payload +from twinkle.server.telemetry.http_middleware import create_metrics_middleware from twinkle.server.telemetry.tracing import create_tracing_middleware -from twinkle.server.utils.validation import verify_request_token from twinkle.utils.logger import get_logger logger = get_logger() @@ -47,6 +51,86 @@ OnShutdown = Callable[[Any], Awaitable[None]] +async def twinkle_server_error_handler(request: Request, exc: TwinkleServerError) -> JSONResponse: + """Map a TwinkleServerError to a structured response, fields at the top level. + + Status code is the exception's ``error_code``; the body is an ``ErrorPayload`` + (``error`` / ``category`` / ``error_code`` / ``request_id``) placed at the top + level rather than nested under ``detail``. A Decision_Boundary-left rejection + (``category=user``) carries no traceback. + """ + request_id = getattr(request.state, 'request_id', None) or '' + payload = build_error_payload( + str(exc) or exc.__class__.__name__, + category=exc.category, + error_code=exc.error_code, + request_id=request_id, + ) + return JSONResponse(status_code=exc.error_code, content=payload.model_dump(mode='json', exclude_none=True)) + + +# A body can produce hundreds of errors (one per element of a mis-typed tensor), and a +# response listing all of them helps nobody while costing bandwidth on every retry. +_MAX_VALIDATION_DETAILS = 20 + + +def _validation_detail(error: dict[str, Any]) -> dict[str, Any]: + """One pydantic error as a JSON-safe detail entry.""" + location = [str(part) for part in error.get('loc', ())] + return { + 'field': location[-1] if location else '', + 'path': '.'.join(location), + 'type': error.get('type', ''), + 'message': error.get('msg', ''), + } + + +def _validation_summary(errors: list[dict[str, Any]]) -> str: + fields = [] + for error in errors: + path = '.'.join(str(part) for part in error.get('loc', ())) + if path and path not in fields: + fields.append(path) + shown = ', '.join(fields[:_MAX_VALIDATION_DETAILS]) or 'request body' + suffix = '' if len(fields) <= _MAX_VALIDATION_DETAILS else f' (+{len(fields) - _MAX_VALIDATION_DETAILS} more)' + return f'Request body validation failed for: {shown}{suffix}' + + +def _validation_mentions_unknown_field(errors: list[dict[str, Any]]) -> bool: + return any(error.get('type') == 'extra_forbidden' for error in errors) + + +async def validation_error_handler(request: Request, exc: RequestValidationError) -> JSONResponse: + """Map a body validation failure to a 422 carrying an ``ErrorPayload``. + + Lives next to ``twinkle_server_error_handler`` because both are the same concern -- + the wire shape of a failure -- and both are registered by ``build_deployment_app`` + for all four deployments. It used to sit in ``validation/``, whose ``__init__`` + docstring admitted the fit was awkward ("another half of the story"). + + FastAPI's default handler answers with ``{"detail": [...]}``, a second error shape + on the wire; this makes the Model, Sampler and Processor deployments answer + identically to every other twinkle failure. The per-field ``details`` name the + offending field, its path, and why it was rejected; there is no traceback because a + rejected body is the caller's problem, not a crash. + """ + errors = list(exc.errors()) + message = _validation_summary(errors) + if _validation_mentions_unknown_field(errors): + # An unknown top-level field is what an older client looks like against a newer + # server, so say so instead of leaving the caller to infer it from a field list. + message += ('. Unknown fields are rejected; if this worked before, upgrade ' + 'twinkle-kit on the client to match the server version.') + payload = build_error_payload( + message, + category=ErrorCategory.User, + error_code=422, + request_id=getattr(request.state, 'request_id', None) or '', + details=[_validation_detail(error) for error in errors[:_MAX_VALIDATION_DETAILS]], + ) + return JSONResponse(status_code=422, content=payload.model_dump(mode='json', exclude_none=True)) + + def get_servable() -> Any: """The single definition of the servable-object accessor used by every builder. @@ -74,17 +158,20 @@ def build_deployment_app( shutdown → ``on_shutdown(get_servable())`` (best-effort) then ``flush_telemetry_safely()`` so buffered OTLP batches flush on graceful replica termination; - 2. [if ``attach_cleanup_middleware``] the gateway-only lazy-cleanup + 2. the ``TwinkleServerError`` and ``RequestValidationError`` handlers, so a + rejected request body carries the same ``ErrorPayload`` shape as any other + failure; + 3. [if ``attach_cleanup_middleware``] the gateway-only lazy-cleanup middleware (registered first ⇒ innermost), since the Gateway has no per-handler hook; - 3. ``catch_unhandled_exceptions`` middleware, inside auth/tracing/metrics + 4. ``catch_unhandled_exceptions`` middleware, inside auth/tracing/metrics and outside cleanup/routes; - 4. ``verify_token`` middleware; - 5. ``create_tracing_middleware(component)``; - 6. ``create_metrics_middleware(component)``; - 7. [if ``attach_replica_id_header``] replica-id response header middleware + 5. ``verify_token`` middleware; + 6. ``create_tracing_middleware(component)``; + 7. ``create_metrics_middleware(component)``; + 8. [if ``attach_replica_id_header``] replica-id response header middleware (registered last ⇒ outermost); - 8. ``register_routes(app, get_servable)``. + 9. ``register_routes(app, get_servable)``. Args: component: ``'Gateway' | 'Model' | 'Sampler' | 'Processor'`` — used as @@ -123,6 +210,12 @@ async def lifespan(app: FastAPI): app = FastAPI(lifespan=lifespan, **(fastapi_kwargs or {})) + app.add_exception_handler(TwinkleServerError, twinkle_server_error_handler) + # Request-body validation failures answer with the same ``ErrorPayload`` shape as + # every other error, registered here so all deployments behave identically rather + # than each app keeping (or forgetting) its own copy. + app.add_exception_handler(RequestValidationError, validation_error_handler) + # Registration order matters: FastAPI runs middleware LIFO, so the LAST # registered wraps the outermost layer. Register cleanup (if any) first so # it stays innermost, then the exception boundary, auth, tracing, metrics, @@ -142,10 +235,21 @@ async def ensure_state_cleanup_started(request: Request, call_next): async def catch_unhandled_exceptions(request: Request, call_next): try: return await call_next(request) - except Exception: - error = traceback.format_exc() - logger.error(error) - return JSONResponse(status_code=500, content={'detail': error}) + except Exception as exc: + tb = traceback.format_exc() + logger.error(tb) + # Unify the last-resort 500 with the rest of the wire: an + # ``ErrorPayload`` body (Server category keeps the traceback) instead + # of the legacy ``{'detail': }`` shape. + request_id = getattr(request.state, 'request_id', None) or '' + payload = build_error_payload( + str(exc) or exc.__class__.__name__, + category=ErrorCategory.Server, + error_code=500, + request_id=request_id, + traceback_text=tb, + ) + return JSONResponse(status_code=500, content=payload.model_dump(mode='json', exclude_none=True)) @app.middleware('http') async def verify_token(request: Request, call_next): @@ -199,31 +303,6 @@ def bind_deployment( return deployment_cls.options(**deploy_options).bind(*bind_args, **(bind_kwargs or {})) -def init_twinkle_runtime( - is_mock: bool, - nproc_per_node: int, - device_group: Any, - device_mesh_dict: dict[str, Any], -) -> Any | None: - """Initialize the Twinkle distributed runtime and build a DeviceMesh. - - Shared by ModelManagement and SamplerManagement ``__init__``. - Returns ``None`` for mock backends (CPU-only, no device mesh). - """ - import twinkle - from twinkle import DeviceMesh - - if is_mock: - twinkle.initialize( - mode='ray', nproc_per_node=nproc_per_node, ncpu_proc_per_node=1, groups=[device_group], lazy_collect=False) - return None - - twinkle.initialize(mode='ray', nproc_per_node=nproc_per_node, groups=[device_group], lazy_collect=False) - if 'mesh_dim_names' in device_mesh_dict: - return DeviceMesh(**device_mesh_dict) - return DeviceMesh.from_sizes(**device_mesh_dict) - - class LazyCleanupMixin: """Single source of the lazy first-request ServerState cleanup-start behavior. @@ -231,13 +310,22 @@ class LazyCleanupMixin: ``_ensure_state_cleanup_started`` methods collapse into one. The method name is preserved so existing call sites (``_on_request_start``, ``_ensure_sticky``, the Gateway cleanup middleware) are unchanged. + + ``_state_cleanup_started`` is declared here with a class-level default so the four + deployment classes need no ``__init__`` change and the ``getattr`` fallback can go: + only ``GatewayServer`` used to initialise it explicitly, and Model/Sampler/Processor + relied on ``getattr``'s default. The first write of ``True`` shadows the + class attribute on the instance -- expected, since the class attribute is only a + default. """ + _state_cleanup_started: bool = False + async def _ensure_state_cleanup_started(self) -> None: - if getattr(self, '_state_cleanup_started', False): + if self._state_cleanup_started: return try: - # Idempotent via ServerState's own ``_cleanup_running`` guard. + # Idempotent via the ResourceCleanupCoordinator's own start guard. await self.state.start_cleanup_task() except Exception as e: logger.warning(f'Failed to start ServerState cleanup task: {e}') diff --git a/src/twinkle/server/exceptions.py b/src/twinkle/server/exceptions.py index dc58ad81d..b8985879d 100644 --- a/src/twinkle/server/exceptions.py +++ b/src/twinkle/server/exceptions.py @@ -1,16 +1,46 @@ -"""Twinkle Server unified exception hierarchy.""" +"""Twinkle Server unified exception hierarchy. + +Every exception carries an ``error_code`` (an HTTP-status-shaped int in 400-599) and +a ``category`` (:class:`ErrorCategory`). A single ``TwinkleServerError`` exception +handler (see the gateway/model/sampler apps) reads these two attributes to build a +structured response whose fields sit at the top level of the body -- not nested +under ``detail``. +""" from __future__ import annotations +from twinkle.protocol.types.errors import ErrorCategory + class TwinkleServerError(Exception): - """Base class for all Twinkle Server exceptions.""" - pass + """Base class for all Twinkle Server exceptions. + + ``error_code`` / ``category`` are class-level defaults a subclass overrides; an + instance may also override them via keyword to avoid a subclass per status code. + """ + + error_code: int = 500 + category: ErrorCategory = ErrorCategory.Server + + def __init__( + self, + message: str = '', + *, + error_code: int | None = None, + category: ErrorCategory | None = None, + ) -> None: + super().__init__(message) + if error_code is not None: + self.error_code = error_code + if category is not None: + self.category = category class StateBackendError(TwinkleServerError): """State backend operation failed (connection lost, timeout, data serialization error, etc.).""" - pass + + error_code = 500 + category = ErrorCategory.Server class ConfigError(TwinkleServerError): @@ -22,6 +52,9 @@ class ConfigError(TwinkleServerError): re-running the server. """ + error_code = 500 + category = ErrorCategory.Server + def __init__( self, field: str, @@ -45,21 +78,99 @@ class ConfigParseError(TwinkleServerError): value violates a field/cross-field rule) and from ``FileNotFoundError`` (which signals that the source could not be read at all). """ - pass + + error_code = 500 + category = ErrorCategory.Server class ResourceExhaustedError(TwinkleServerError): """Resource exhausted — queue full, insufficient memory, connection pool exhausted, etc.""" - pass + + error_code = 503 + category = ErrorCategory.Server + + +class EndpointUnavailableError(TwinkleServerError): + """The endpoint is not implemented by this deployment's backend.""" + + error_code = 501 + category = ErrorCategory.Server + + +class RequestRejectedError(TwinkleServerError): + """Decision_Boundary-left failure: rejectable from the request body, deployment + config, and loaded schema alone, so it is returned with a real HTTP status code + and writes NO future record. + + Named ``RequestRejectedError`` rather than ``RequestValidationError`` to avoid a + collision with ``fastapi.exceptions.RequestValidationError``. The default is a + 400/User rejection; the subclasses below pin the specific status codes from the + Decision_Boundary placement table. + """ + + error_code = 400 + category = ErrorCategory.User -class FullModeBusyError(TwinkleServerError): +class TrainModeMismatchError(RequestRejectedError): + """The request's train mode does not match the deployment's (LoRA vs full).""" + + error_code = 400 + category = ErrorCategory.User + + +class InputTokensExceededError(RequestRejectedError): + """The request's input token count exceeds ``max_input_tokens``.""" + + error_code = 422 + category = ErrorCategory.User + + +class BatchSizeError(RequestRejectedError): + """Batch size is incompatible with the data world size (too small / not a multiple).""" + + error_code = 422 + category = ErrorCategory.User + + +class RateLimitExceededError(RequestRejectedError): + """The request or token rate exceeds the configured limit.""" + + error_code = 429 + category = ErrorCategory.User + + +class ResourceQuotaExceededError(RequestRejectedError): + """The caller exhausted a configured per-token resource quota.""" + + error_code = 429 + category = ErrorCategory.User + + +class ResourceNotFoundError(RequestRejectedError): + """A well-formed request names a resource (adapter / session) that is absent. + + Distinct from a malformed request (400): the request itself is valid but the + referenced resource does not exist or is expiring, so it is a 404 on the + Decision_Boundary left. Raised (never ``assert``-ed) so the check survives + ``python -O`` and is classified as user-facing rather than a 500. + """ + + error_code = 404 + category = ErrorCategory.User + + +class FullModeBusyError(RequestRejectedError): """A full-parameter (exclusive) model deployment already has a holder. Full-parameter training rewrites the shared base-model weights, so a single - deployment can only host one training task at a time. + deployment can only host one training task at a time. It is a request rejection + (a second tenant is turned away), hence a 409 on the Decision_Boundary left. """ + error_code = 409 + category = ErrorCategory.User + def __init__(self, current_holder: str) -> None: self.current_holder = current_holder super().__init__('This deployment runs in full-parameter (exclusive) mode and is already ' diff --git a/src/twinkle/server/gateway/app.py b/src/twinkle/server/gateway/app.py index bdf020655..0c54f506b 100644 --- a/src/twinkle/server/gateway/app.py +++ b/src/twinkle/server/gateway/app.py @@ -11,14 +11,14 @@ from fastapi import FastAPI, HTTPException from typing import Any -import twinkle_client.types as types +import twinkle.protocol.types as types from twinkle.server.deployment import LazyCleanupMixin, bind_deployment, build_deployment_app from twinkle.server.state import get_server_state from twinkle.utils.logger import get_logger from .openai_handlers import _register_openai_routes from .proxy import ServiceProxy -from .tinker_handlers import _register_tinker_routes -from .twinkle_handlers import _register_twinkle_routes +from .tinker_handlers import _register_gateway_tinker_routes +from .twinkle_handlers import _register_gateway_twinkle_routes logger = get_logger() @@ -44,6 +44,18 @@ def __init__(self, self._supported_model_names = frozenset(m.model_name for m in self.supported_models) self._modelscope_config_lock = asyncio.Lock() self._state_cleanup_started = False + # Per-instance, not module-level: a process-global set never got cleared + # across replica rebuilds and made the OpenAI template tests non-isolatable. + self._template_initialized: set[str] = set() + + @property + def supported_model_names(self) -> frozenset[str]: + """Base-model names this gateway accepts. + + Public because ``openai_handlers._resolve_base_model`` reads it; it used to + reach into ``gateway._supported_model_names`` directly. + """ + return self._supported_model_names @staticmethod def _normalize_server_state_args(server_config: Any) -> dict[str, Any]: @@ -116,8 +128,8 @@ def build_gateway_app(deploy_options: dict[str, Any], # because it has no per-handler request hook, so the lazy-cleanup middleware # must cover every route (and stays innermost). def register_routes(app: FastAPI, get_self: Any) -> None: - _register_tinker_routes(app, get_self) - _register_twinkle_routes(app, get_self) + _register_gateway_tinker_routes(app, get_self) + _register_gateway_twinkle_routes(app, get_self) _register_openai_routes(app, get_self) async def _on_shutdown(servable: Any) -> None: diff --git a/src/twinkle/server/gateway/openai_handlers.py b/src/twinkle/server/gateway/openai_handlers.py index cb33f1617..aeed51954 100644 --- a/src/twinkle/server/gateway/openai_handlers.py +++ b/src/twinkle/server/gateway/openai_handlers.py @@ -15,7 +15,7 @@ from fastapi.responses import JSONResponse, StreamingResponse from typing import TYPE_CHECKING, Any -from twinkle_client.http.headers import H_AUTH, H_AUTH_TWINKLE, build_routing_headers +from twinkle.protocol.headers import H_AUTH, H_AUTH_TWINKLE, build_routing_headers if TYPE_CHECKING: from .app import GatewayServer @@ -24,6 +24,7 @@ from twinkle.server.utils import get_template_for_model from twinkle.utils.logger import get_logger +from . import routes from .openai_bridge import make_error, translate_chat_request, translate_response, translate_stream_chunk logger = get_logger() @@ -87,7 +88,7 @@ async def chat_completions( # Non-streaming: proxy to /twinkle/sample, translate response response = await self.proxy.proxy_request( request, - endpoint='twinkle/sample', + endpoint=routes.TWINKLE_SAMPLE, base_model=base_model, service_type='sampler', body_override=body_bytes, @@ -121,7 +122,7 @@ async def _sse_generator(): try: async for line in self.proxy.proxy_request_stream( request, - endpoint='twinkle/sample_stream', + endpoint=routes.TWINKLE_SAMPLE_STREAM, base_model=base_model, service_type='sampler', body_override=body_bytes, @@ -195,12 +196,12 @@ async def _resolve_base_model(gateway: GatewayServer, model: str) -> str | None: pass # Check if it's directly a supported base model - if model in gateway._supported_model_names: + if model in gateway.supported_model_names: return model # Fallback: if there's exactly one supported model, use it - if len(gateway._supported_model_names) == 1: - return next(iter(gateway._supported_model_names)) + if len(gateway.supported_model_names) == 1: + return next(iter(gateway.supported_model_names)) return None @@ -211,10 +212,6 @@ def _build_sticky_headers(sticky_key: str, request: Request) -> dict[str, str]: return build_routing_headers(sticky_key, auth) -# Per-process cache; each Ray Serve worker holds its own instance. -_template_initialized: set[str] = set() - - async def _ensure_template( gateway: GatewayServer, base_model: str, @@ -223,10 +220,10 @@ async def _ensure_template( ) -> None: """Ensure the sampler has a chat template set for encoding Trajectory inputs. - Called once per base_model (cached in-process). On failure, logs a warning - but doesn't block — the sampler will return its own error if needed. + Called once per base_model (cached on the ``GatewayServer`` instance). On failure, + logs a warning but doesn't block -- the sampler will return its own error if needed. """ - if base_model in _template_initialized: + if base_model in gateway._template_initialized: return template_cls = get_template_for_model(base_model) @@ -239,14 +236,14 @@ async def _ensure_template( try: resp = await gateway.proxy.proxy_request( request, - endpoint='twinkle/set_template', + endpoint=routes.TWINKLE_SET_TEMPLATE, base_model=base_model, service_type='sampler', body_override=set_template_body, extra_headers=sticky_headers, ) if resp.status_code == 200: - _template_initialized.add(base_model) + gateway._template_initialized.add(base_model) else: logger.warning('set_template failed: %s', resp.body.decode()[:200]) except Exception as e: diff --git a/src/twinkle/server/gateway/proxy.py b/src/twinkle/server/gateway/proxy.py index 043e421ea..664d68d23 100644 --- a/src/twinkle/server/gateway/proxy.py +++ b/src/twinkle/server/gateway/proxy.py @@ -10,11 +10,14 @@ import httpx from fastapi import Request, Response +from fastapi.responses import JSONResponse from typing import Any +from twinkle.protocol.headers import H_MULTIPLEX, H_MULTIPLEX_LEGACY, H_REQUEST_ID, H_REQUEST_ID_LEGACY +from twinkle.protocol.types.errors import ErrorCategory, ErrorPayload from twinkle.server.telemetry.tracing import inject_context from twinkle.utils.logger import get_logger -from twinkle_client.http.headers import H_MULTIPLEX, H_MULTIPLEX_LEGACY, H_REQUEST_ID, H_REQUEST_ID_LEGACY +from . import routes logger = get_logger() @@ -56,7 +59,6 @@ def _build_target_url(self, service_type: str, base_model: str, endpoint: str) - Returns: Complete target URL for the internal service """ - prefix = self.route_prefix.rstrip('/') if self.route_prefix else '' host = self.http_options.get('host', 'localhost') port = self.http_options.get('port', 8000) @@ -64,7 +66,7 @@ def _build_target_url(self, service_type: str, base_model: str, endpoint: str) - host = 'localhost' base_url = f'http://{host}:{port}' - return f'{base_url}{prefix}/{service_type}/{base_model}/{endpoint}' + return f'{base_url}{routes.target_url(self.route_prefix, service_type, base_model, endpoint)}' def _prepare_headers(self, request_headers) -> dict[str, str]: """Prepare headers for proxying by removing problematic headers.""" @@ -144,7 +146,19 @@ async def proxy_request( ) except Exception as e: logger.error('Proxy error: %s', str(e), exc_info=True) - return Response(content=f'Proxy Error: {str(e)}', status_code=502) + # The gateway could not reach the upstream deployment. Return the + # unified ``ErrorPayload`` (502/Server) instead of a plain-text body + # so every gateway failure has the same wire shape. Upstream error + # responses are passed through unchanged above, preserving their own + # ErrorPayload body. + request_id = request.headers.get(H_REQUEST_ID) or request.headers.get(H_REQUEST_ID_LEGACY) or '' + payload = ErrorPayload( + error=f'Proxy Error: {str(e)}', + category=ErrorCategory.Server, + error_code=502, + request_id=request_id, + ) + return JSONResponse(status_code=502, content=payload.model_dump(mode='json', exclude_none=True)) async def proxy_request_stream( self, @@ -203,7 +217,7 @@ async def proxy_to_model(self, request: Request, endpoint: str, base_model: str) endpoint: The tinker endpoint name (e.g., 'create_model', 'forward') base_model: The base model name for routing """ - return await self.proxy_request(request, f'tinker/{endpoint}', base_model, 'model') + return await self.proxy_request(request, routes.tinker_endpoint(endpoint), base_model, 'model') async def proxy_to_sampler(self, request: Request, endpoint: str, base_model: str) -> Response: """Proxy request to sampler's tinker endpoint (/tinker/). @@ -213,4 +227,4 @@ async def proxy_to_sampler(self, request: Request, endpoint: str, base_model: st endpoint: The tinker endpoint name (e.g., 'asample') base_model: The base model name for routing """ - return await self.proxy_request(request, f'tinker/{endpoint}', base_model, 'sampler') + return await self.proxy_request(request, routes.tinker_endpoint(endpoint), base_model, 'sampler') diff --git a/src/twinkle/server/gateway/routes.py b/src/twinkle/server/gateway/routes.py new file mode 100644 index 000000000..0cececc05 --- /dev/null +++ b/src/twinkle/server/gateway/routes.py @@ -0,0 +1,26 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Route-path constants shared by the gateway's proxy and its OpenAI bridge. + +Single source for the ``route_prefix`` convention these three places must agree on; they +previously agreed by having the same string typed out in three files. Placed on +the consumer side (gateway) rather than in ``launcher`` -- launcher *produces* the +``route_prefix``, gateway consumes it -- to avoid a new ``gateway -> launcher`` edge. +""" +from __future__ import annotations + +# Downstream endpoint paths the gateway proxies to (relative to the service's route prefix). +TWINKLE_SAMPLE = 'twinkle/sample' +TWINKLE_SAMPLE_STREAM = 'twinkle/sample_stream' +TWINKLE_SET_TEMPLATE = 'twinkle/set_template' + +TINKER_PREFIX = 'tinker' + + +def tinker_endpoint(endpoint: str) -> str: + """``tinker/`` -- the shape ``proxy_to_model`` / ``proxy_to_sampler`` build.""" + return f'{TINKER_PREFIX}/{endpoint}' + + +def target_url(route_prefix: str, service_type: str, base_model: str, endpoint: str) -> str: + """The one definition of ``{route_prefix}/{service_type}/{base_model}/{endpoint}``.""" + return f'{route_prefix.rstrip("/")}/{service_type}/{base_model}/{endpoint}' diff --git a/src/twinkle/server/gateway/tinker_handlers.py b/src/twinkle/server/gateway/tinker_handlers.py index 9ef98ce15..93809d780 100644 --- a/src/twinkle/server/gateway/tinker_handlers.py +++ b/src/twinkle/server/gateway/tinker_handlers.py @@ -2,13 +2,11 @@ """ Tinker-compatible gateway handlers. -All endpoints are prefixed /* and registered via _register_tinker_routes(app, self_fn). +All endpoints are prefixed /* and registered via _register_gateway_tinker_routes(app, self_fn). self_fn is injected via FastAPI Depends to obtain the GatewayServer instance at request time. """ from __future__ import annotations -import asyncio -import os from collections.abc import Callable from fastapi import Depends, FastAPI, HTTPException, Request, Response from tinker import types @@ -19,14 +17,55 @@ from twinkle.hub import HubOperation from twinkle.server.checkpoint import create_checkpoint_manager, create_training_run_manager -from twinkle.server.utils.task_queue import QueueState -from twinkle.server.utils.validation import get_token_from_request +from twinkle.server.middleware.auth import get_token_from_request +from twinkle.server.state.models import FutureFailureRecord +from twinkle.server.task_errors import trim_traceback from twinkle.utils.logger import get_logger +from .use_cases import delete_checkpoint +from .use_cases import get_training_run as get_training_run_use_case +from .use_cases import get_weights_info, list_checkpoints, list_training_runs, poll_future logger = get_logger() - -def _register_tinker_routes(app: FastAPI, self_fn: Callable[[], GatewayServer]) -> None: +# Keys must equal ``state.models.FAILURE_REASON_CODES`` (guarded by test_envelope). +_TINKER_FAILURE_WIRE: dict[str, tuple[int, str]] = { + 'invalid_request': (400, 'user'), + 'request_rejected': (400, 'user'), + 'resource_not_found': (404, 'user'), + 'full_mode_busy': (409, 'user'), + 'input_tokens_exceeded': (422, 'user'), + 'batch_size_invalid': (422, 'user'), + 'rate_limit_exceeded': (429, 'user'), + 'resource_quota_exceeded': (429, 'user'), + 'cancelled': (499, 'user'), + 'endpoint_unavailable': (501, 'server'), + 'state_contention': (503, 'server'), + 'backend_gate_unavailable': (503, 'server'), + 'orphaned_replica': (503, 'server'), + 'execution_timeout': (504, 'server'), + 'deadline_exceeded': (500, 'server'), + 'internal_error': (500, 'server'), +} + + +def _tinker_error_from_failure(stored: Any, *, request_id: str) -> dict[str, Any]: + """Map domain failure state to the existing Tinker-compatible error body.""" + failure = FutureFailureRecord.model_validate(stored) + error_code, category = _TINKER_FAILURE_WIRE.get(failure.reason_code, (500, 'server')) + payload: dict[str, Any] = { + 'error': failure.message[:1024], + 'category': category, + 'error_code': error_code, + 'request_id': request_id, + } + if failure.details is not None: + payload['details'] = failure.details + if failure.diagnostic and category == 'server': + payload['traceback'] = trim_traceback(failure.diagnostic) + return payload + + +def _register_gateway_tinker_routes(app: FastAPI, self_fn: Callable[[], GatewayServer]) -> None: """Register all /* Tinker routes on the given FastAPI app. self_fn is a zero-argument callable that returns the current GatewayServer @@ -43,7 +82,7 @@ async def get_server_capabilities( request: Request, self: GatewayServer = Depends(self_fn), ) -> types.GetServerCapabilitiesResponse: - # Convert twinkle_client.types.SupportedModel to tinker.types.SupportedModel + # Convert twinkle.protocol.types.SupportedModel to tinker.types.SupportedModel tinker_supported_models = [types.SupportedModel(model_name=m.model_name) for m in self.supported_models] return types.GetServerCapabilitiesResponse(supported_models=tinker_supported_models) @@ -82,50 +121,24 @@ async def retrieve_future(request: Request, self: GatewayServer = Depends(self_fn)) -> Any: """Retrieve the result of an async task with long polling.""" request_id = body.request_id - max_wait = float(os.environ.get('TWINKLE_LONG_POLL_TIMEOUT', '30')) - poll_interval = float(os.environ.get('TWINKLE_POLL_INTERVAL', '0.5')) - start = asyncio.get_running_loop().time() - - while True: - record = await self.state.get_future(request_id) - + outcome = await poll_future(self.state, request_id) + record = outcome.record + if outcome.timed_out: + response_data: dict[str, Any] = {'type': 'try_again'} if record is not None: - status = record.get('status') - if status not in ('pending', 'queued', 'running', 'rate_limited'): - break - - # ``record is None`` here means the future hasn't been written yet - # (cross-replica visibility lag) — fold into the long-poll loop - # rather than short-circuit ``try_again``: returning immediately - # lets the SDK hammer this endpoint at ~150 Hz. - if asyncio.get_running_loop().time() - start >= max_wait: - response_data: dict[str, Any] = {'type': 'try_again'} - if record is not None: - if queue_state := record.get('queue_state'): - response_data['queue_state'] = queue_state - if queue_state_reason := record.get('queue_state_reason'): - response_data['queue_state_reason'] = queue_state_reason - return response_data - - await asyncio.sleep(poll_interval) + if queue_state := record.get('queue_state'): + response_data['queue_state'] = queue_state + if queue_state_reason := record.get('queue_state_reason'): + response_data['queue_state_reason'] = queue_state_reason + return response_data status = record.get('status') - - if status == 'rate_limited': - return { - 'type': 'try_again', - 'queue_state': QueueState.PAUSED_RATE_LIMIT.value, - 'queue_state_reason': record.get('reason', 'Rate limit exceeded') - } - if status == 'failed': - result = record.get('result', {}) - return {'error': result.get('error', 'Unknown error'), 'category': result.get('category', 'Server')} + return _tinker_error_from_failure(record.get('failure'), request_id=request_id) result = record.get('result') if result is None: raise HTTPException(status_code=500, detail='Task completed but no result found') - if hasattr(result, 'model_dump'): return result.model_dump() return result @@ -135,14 +148,12 @@ async def retrieve_future(request: Request, @app.get('/training_runs') async def get_training_runs(request: Request, limit: int = 20, offset: int = 0) -> types.TrainingRunsResponse: token = get_token_from_request(request) - training_run_manager = create_training_run_manager(token, client_type='tinker') - return training_run_manager.list_runs(limit=limit, offset=offset) + return list_training_runs(token, 'tinker', limit=limit, offset=offset) @app.get('/training_runs/{run_id}') async def get_training_run(request: Request, run_id: str) -> types.TrainingRun: token = get_token_from_request(request) - training_run_manager = create_training_run_manager(token, client_type='tinker') - run = training_run_manager.get(run_id) + run = get_training_run_use_case(token, 'tinker', run_id) if not run: raise HTTPException(status_code=404, detail=f'Training run {run_id} not found') return run @@ -150,8 +161,7 @@ async def get_training_run(request: Request, run_id: str) -> types.TrainingRun: @app.get('/training_runs/{run_id}/checkpoints') async def get_run_checkpoints(request: Request, run_id: str) -> types.CheckpointsListResponse: token = get_token_from_request(request) - checkpoint_manager = create_checkpoint_manager(token, client_type='tinker') - response = checkpoint_manager.list_checkpoints(run_id) + response = list_checkpoints(token, 'tinker', run_id) if not response: raise HTTPException(status_code=404, detail=f'Training run {run_id} not found') return response @@ -159,8 +169,7 @@ async def get_run_checkpoints(request: Request, run_id: str) -> types.Checkpoint @app.delete('/training_runs/{run_id}/checkpoints/{checkpoint_id:path}') async def delete_run_checkpoint(request: Request, run_id: str, checkpoint_id: str) -> Any: token = get_token_from_request(request) - checkpoint_manager = create_checkpoint_manager(token, client_type='tinker') - success = checkpoint_manager.delete(run_id, checkpoint_id) + success = delete_checkpoint(token, 'tinker', run_id, checkpoint_id) if not success: raise HTTPException(status_code=404, detail=f'Checkpoint {checkpoint_id} not found for run {run_id}') return None @@ -168,9 +177,8 @@ async def delete_run_checkpoint(request: Request, run_id: str, checkpoint_id: st @app.post('/weights_info') async def weights_info(request: Request, body: dict[str, Any]) -> types.WeightsInfoResponse: token = get_token_from_request(request) - checkpoint_manager = create_checkpoint_manager(token, client_type='tinker') tinker_path = body.get('tinker_path') - response = checkpoint_manager.get_weights_info(tinker_path) + response = get_weights_info(token, 'tinker', tinker_path) if not response: raise HTTPException(status_code=404, detail=f'Weights at {tinker_path} not found') return response diff --git a/src/twinkle/server/gateway/twinkle_handlers.py b/src/twinkle/server/gateway/twinkle_handlers.py index c3de4d8a5..3df44ffc8 100644 --- a/src/twinkle/server/gateway/twinkle_handlers.py +++ b/src/twinkle/server/gateway/twinkle_handlers.py @@ -2,26 +2,32 @@ """ Twinkle-native gateway handlers. -All endpoints are prefixed /twinkle/* and registered via _register_twinkle_routes(app, self_fn). +All endpoints are prefixed /twinkle/* and registered via _register_gateway_twinkle_routes(app, self_fn). """ from __future__ import annotations from collections.abc import Callable -from fastapi import Depends, FastAPI, HTTPException, Request +from fastapi import Depends, FastAPI, Request from typing import TYPE_CHECKING if TYPE_CHECKING: from .app import GatewayServer -import twinkle_client.types as types +import twinkle.protocol.types as types from twinkle.server.checkpoint import create_checkpoint_manager, create_training_run_manager, validate_user_path -from twinkle.server.utils.validation import get_token_from_request +from twinkle.server.exceptions import RequestRejectedError, ResourceNotFoundError +from twinkle.server.lifecycle.envelope import envelope_from_record +from twinkle.server.lifecycle.poll_config import long_poll_window +from twinkle.server.middleware.auth import get_token_from_request from twinkle.utils.logger import get_logger +from .use_cases import delete_checkpoint +from .use_cases import get_training_run as get_training_run_use_case +from .use_cases import get_weights_info, list_checkpoints, list_training_runs, poll_future logger = get_logger() -def _register_twinkle_routes(app: FastAPI, self_fn: Callable[[], GatewayServer]) -> None: +def _register_gateway_twinkle_routes(app: FastAPI, self_fn: Callable[[], GatewayServer]) -> None: """Register all /twinkle/* routes on the given FastAPI app.""" @app.get('/twinkle/capacity_info', response_model=types.CapacityInfoResponse) @@ -79,7 +85,18 @@ async def get_server_capabilities( request: Request, self: GatewayServer = Depends(self_fn), ) -> types.GetServerCapabilitiesResponse: - return types.GetServerCapabilitiesResponse(supported_models=self.supported_models) + return types.GetServerCapabilitiesResponse( + supported_models=self.supported_models, + protocol_version=1, + features=types.ClientFeatures( + task_envelope=True, + cancel=True, + data_plane=True, + full_training=True, + batch_retrieve=False, + ), + limits=types.ProtocolLimits(long_poll_timeout_seconds=long_poll_window()), + ) @app.post('/twinkle/create_session', response_model=types.CreateSessionResponse) async def create_session( @@ -98,31 +115,69 @@ async def session_heartbeat( ) -> types.SessionHeartbeatResponse: alive = await self.state.touch_session(body.session_id) if not alive: - raise HTTPException(status_code=404, detail='Unknown session') + raise ResourceNotFoundError('Unknown session') return types.SessionHeartbeatResponse() + @app.post('/twinkle/retrieve_future', response_model=types.TaskEnvelope) + async def retrieve_future( + request: Request, + body: types.RetrieveFutureRequest, + self: GatewayServer = Depends(self_fn), + ) -> types.TaskEnvelope: + """Long-poll a twinkle-native task to a terminal state. + + Returns 200 for every outcome except a request_id that stayed invisible for + a whole window -- the HTTP call succeeded, it successfully reported the + task's state. Unlike the tinker endpoint next door, ``completed`` with a + null result is a valid success (step / zero_grad / lr_step all return None), + so this handler never raises the tinker endpoint's + ``HTTPException(500, 'Task completed but no result found')``. + + A fixed interval, not exponential backoff: measured on real hardware, a + 0.05->1.0s doubling schedule is ~22% SLOWER per step because its interval + grows fastest across the 0.5-1.2s band where data-plane tasks actually + finish. See ``poll_config`` for the numbers. + """ + request_id = body.request_id + outcome = await poll_future(self.state, request_id) + if outcome.record is None: + raise ResourceNotFoundError(f'request_id {request_id} not found or expired') + return envelope_from_record(request_id, outcome.record) + + @app.post('/twinkle/cancel', response_model=types.CancelResponse) + async def cancel_future( + request: Request, + body: types.CancelRequest, + self: GatewayServer = Depends(self_fn), + ) -> types.CancelResponse: + """Best-effort cancel of a not-yet-started task. + + Drops the task from the compute queue only if it has not begun running; a + running or already-terminal task is reported but never interrupted, so cancel + can never corrupt in-flight GPU/optimizer state. + """ + result = await self.state.cancel_future(body.request_id) + return types.CancelResponse(**result) + @app.get('/twinkle/training_runs', response_model=types.TrainingRunsResponse) async def get_training_runs(request: Request, limit: int = 20, offset: int = 0) -> types.TrainingRunsResponse: token = get_token_from_request(request) - training_run_manager = create_training_run_manager(token, client_type='twinkle') - return training_run_manager.list_runs(limit=limit, offset=offset) + return list_training_runs(token, 'twinkle', limit=limit, offset=offset) @app.get('/twinkle/training_runs/{run_id}', response_model=types.TrainingRun) async def get_training_run(request: Request, run_id: str) -> types.TrainingRun: token = get_token_from_request(request) - training_run_manager = create_training_run_manager(token, client_type='twinkle') - run = training_run_manager.get_with_permission(run_id) + run = get_training_run_use_case(token, 'twinkle', run_id, check_permission=True) if not run: - raise HTTPException(status_code=404, detail=f'Training run {run_id} not found or access denied') + raise ResourceNotFoundError(f'Training run {run_id} not found or access denied') return run @app.get('/twinkle/training_runs/{run_id}/checkpoints', response_model=types.CheckpointsListResponse) async def get_run_checkpoints(request: Request, run_id: str) -> types.CheckpointsListResponse: token = get_token_from_request(request) - checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') - response = checkpoint_manager.list_checkpoints(run_id) + response = list_checkpoints(token, 'twinkle', run_id) if response is None: - raise HTTPException(status_code=404, detail=f'Training run {run_id} not found or access denied') + raise ResourceNotFoundError(f'Training run {run_id} not found or access denied') return response @app.delete( @@ -133,22 +188,20 @@ async def delete_run_checkpoint(request: Request, run_id: str, token = get_token_from_request(request) if not validate_user_path(token, checkpoint_id): - raise HTTPException(status_code=400, detail='Invalid checkpoint path: path traversal not allowed') + raise RequestRejectedError('Invalid checkpoint path: path traversal not allowed') - checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') - success = checkpoint_manager.delete(run_id, checkpoint_id) + success = delete_checkpoint(token, 'twinkle', run_id, checkpoint_id) if not success: - raise HTTPException(status_code=404, detail=f'Checkpoint {checkpoint_id} not found or access denied') + raise ResourceNotFoundError(f'Checkpoint {checkpoint_id} not found or access denied') return types.DeleteCheckpointResponse(success=True, message=f'Checkpoint {checkpoint_id} deleted successfully') @app.post('/twinkle/weights_info', response_model=types.WeightsInfoResponse) async def weights_info(request: Request, body: types.WeightsInfoRequest) -> types.WeightsInfoResponse: token = get_token_from_request(request) - checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') - response = checkpoint_manager.get_weights_info(body.twinkle_path) + response = get_weights_info(token, 'twinkle', body.twinkle_path) if response is None: - raise HTTPException(status_code=404, detail=f'Weights at {body.twinkle_path} not found or access denied') + raise ResourceNotFoundError(f'Weights at {body.twinkle_path} not found or access denied') return response @app.get('/twinkle/checkpoint_path/{run_id}/{checkpoint_id:path}', response_model=types.CheckpointPathResponse) @@ -156,18 +209,18 @@ async def get_checkpoint_path(request: Request, run_id: str, checkpoint_id: str) token = get_token_from_request(request) if not validate_user_path(token, checkpoint_id): - raise HTTPException(status_code=400, detail='Invalid checkpoint path: path traversal not allowed') + raise RequestRejectedError('Invalid checkpoint path: path traversal not allowed') training_run_manager = create_training_run_manager(token, client_type='twinkle') checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') run = training_run_manager.get(run_id) if not run: - raise HTTPException(status_code=404, detail=f'Training run {run_id} not found or access denied') + raise ResourceNotFoundError(f'Training run {run_id} not found or access denied') checkpoint = checkpoint_manager.get(run_id, checkpoint_id) if not checkpoint: - raise HTTPException(status_code=404, detail=f'Checkpoint {checkpoint_id} not found') + raise ResourceNotFoundError(f'Checkpoint {checkpoint_id} not found') ckpt_dir = checkpoint_manager.get_ckpt_dir(run_id, checkpoint_id) return types.CheckpointPathResponse(path=str(ckpt_dir), twinkle_path=checkpoint.twinkle_path) diff --git a/src/twinkle/server/gateway/use_cases.py b/src/twinkle/server/gateway/use_cases.py new file mode 100644 index 000000000..26691014b --- /dev/null +++ b/src/twinkle/server/gateway/use_cases.py @@ -0,0 +1,64 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Gateway use cases that do real assembly or control flow. + +What is *not* here is the point: ``create_session`` / ``touch_session`` were one-line +forwards to ``state``, so "is it in this file?" told a reader nothing. Now the file holds +only the long-poll loop and the five checkpoint use cases, which wire up +``create_*_manager(token, client_type)`` -- i.e. things a handler cannot express in one +line. The four handlers that call ``self.state`` directly (``get_capacity_info``, +``cancel_future``, ``get_cleanup_stats``, ``get_model_metadata``) deliberately stay +direct: wrapping them for symmetry would add forwarding, not structure. +""" +from __future__ import annotations + +import asyncio +from dataclasses import dataclass +from typing import Any + +from twinkle.server.checkpoint import create_checkpoint_manager, create_training_run_manager +from twinkle.server.lifecycle.poll_config import long_poll_window, retrieve_poll_interval + +_TERMINAL_STATUSES = frozenset({'completed', 'failed', 'cancelled'}) + + +@dataclass(frozen=True, slots=True) +class FuturePollResult: + record: dict[str, Any] | None + timed_out: bool + + +async def poll_future(state: Any, request_id: str) -> FuturePollResult: + """Long-poll one canonical future record without constructing wire responses.""" + deadline = asyncio.get_running_loop().time() + long_poll_window() + interval = retrieve_poll_interval() + record = None + while True: + record = await state.get_future(request_id) + if record is not None and record.get('status') in _TERMINAL_STATUSES: + return FuturePollResult(record=record, timed_out=False) + if asyncio.get_running_loop().time() >= deadline: + return FuturePollResult(record=record, timed_out=True) + await asyncio.sleep(interval) + + +def list_training_runs(token: str, client_type: str, *, limit: int, offset: int) -> Any: + return create_training_run_manager(token, client_type=client_type).list_runs(limit=limit, offset=offset) + + +def get_training_run(token: str, client_type: str, run_id: str, *, check_permission: bool = False) -> Any | None: + manager = create_training_run_manager(token, client_type=client_type) + if check_permission and hasattr(manager, 'get_with_permission'): + return manager.get_with_permission(run_id) + return manager.get(run_id) + + +def list_checkpoints(token: str, client_type: str, run_id: str) -> Any | None: + return create_checkpoint_manager(token, client_type=client_type).list_checkpoints(run_id) + + +def delete_checkpoint(token: str, client_type: str, run_id: str, checkpoint_id: str) -> bool: + return create_checkpoint_manager(token, client_type=client_type).delete(run_id, checkpoint_id) + + +def get_weights_info(token: str, client_type: str, path: str) -> Any | None: + return create_checkpoint_manager(token, client_type=client_type).get_weights_info(path) diff --git a/src/twinkle/server/launcher/builder_registry.py b/src/twinkle/server/launcher/builder_registry.py index a18be38b1..9c81d07e0 100644 --- a/src/twinkle/server/launcher/builder_registry.py +++ b/src/twinkle/server/launcher/builder_registry.py @@ -1,9 +1,6 @@ # Copyright (c) ModelScope Contributors. All rights reserved. """``import_path`` → deployment-builder resolution. -Extracted from the former single-file ``launcher.py`` (TIER 3 same-named-package -decomposition). No logic change. - The operator-facing YAML ``import_path`` literals (``"server"``, ``"model"``, ``"sampler"``, ``"processor"``, ``"data_plane"``) resolve to internal builder functions. The function selected by the ``"server"`` literal was renamed diff --git a/src/twinkle/server/launcher/env_propagation.py b/src/twinkle/server/launcher/env_propagation.py index 3b0c9b956..c478444d9 100644 --- a/src/twinkle/server/launcher/env_propagation.py +++ b/src/twinkle/server/launcher/env_propagation.py @@ -1,9 +1,8 @@ # Copyright (c) ModelScope Contributors. All rights reserved. """Collection of telemetry / persistence env vars for propagation to Ray workers. -Extracted from the former single-file ``launcher.py`` (TIER 3 same-named-package -decomposition). No logic change. These vars are read inside each Ray Serve -worker process — telemetry by ``ensure_telemetry_initialized()`` and persistence +These variables are read inside each Ray Serve worker process — telemetry by +``ensure_telemetry_initialized()`` and persistence by ``PersistenceConfig.from_env()`` — so the chosen backend / telemetry config is independent of deployment startup order. """ @@ -21,10 +20,6 @@ 'TWINKLE_MODEL_ID_ALIASES', ) -# NCCL-safe env var keys: controls fault tolerance behavior in distributed -# training (safe_loss / @nccl_safe). Must reach model worker actors. -NCCL_SAFE_ENV_KEYS: tuple[str, ...] = ('TWINKLE_FAIL_FAST', ) - def build_telemetry_env_vars() -> dict[str, str]: """Collect telemetry env vars from ``os.environ`` for worker propagation.""" @@ -37,9 +32,15 @@ def build_persistence_env_vars() -> dict[str, str]: return {k: os.environ[k] for k in PERSISTENCE_ENV_KEYS if k in os.environ} -def build_nccl_safe_env_vars() -> dict[str, str]: - """Collect NCCL-safe env vars from ``os.environ`` for worker propagation.""" - return {k: os.environ[k] for k in NCCL_SAFE_ENV_KEYS if k in os.environ} +def build_server_state_env_vars() -> dict[str, str]: + """Collect ServerState-policy env vars from ``os.environ`` for worker propagation. + + Read inside each worker by ``ServerStateArgs.from_env()`` (via + ``get_server_state``) so the configured quota / expiry / metrics interval is + applied everywhere, not only in the gateway that first built the state. + """ + from twinkle.server.config.application_spec import SERVER_STATE_ENV_KEYS + return {k: os.environ[k] for k in SERVER_STATE_ENV_KEYS if k in os.environ} def build_propagated_env_vars() -> dict[str, str]: @@ -47,5 +48,5 @@ def build_propagated_env_vars() -> dict[str, str]: merged: dict[str, str] = {} merged.update(build_telemetry_env_vars()) merged.update(build_persistence_env_vars()) - merged.update(build_nccl_safe_env_vars()) + merged.update(build_server_state_env_vars()) return merged diff --git a/src/twinkle/server/launcher/server_launcher.py b/src/twinkle/server/launcher/server_launcher.py index 4c64c1208..75fd0bd52 100644 --- a/src/twinkle/server/launcher/server_launcher.py +++ b/src/twinkle/server/launcher/server_launcher.py @@ -18,9 +18,9 @@ from twinkle import get_logger from twinkle.hub.model_alias import MODEL_ID_ALIASES_ENV, build_model_alias_map +from twinkle.patch.ray_serve import apply_ray_serve_patches, get_runtime_env_for_patches from twinkle.server.config import ServerConfig from twinkle.server.config.application_spec import ApplicationSpec -from twinkle.server.utils.ray_serve_patch import apply_ray_serve_patches, get_runtime_env_for_patches from .builder_registry import get_builders, resolve_builder from .env_propagation import build_propagated_env_vars @@ -223,6 +223,23 @@ def launch(self) -> None: os.environ[k] = v logger.info(f'Persistence backend configured: mode={persistence.mode}') + # Export the gateway ``server`` application's ServerState policy (quota / + # expiry / cleanup / metrics interval) to env vars for the same reason: + # so every worker's first ``get_server_state()`` applies the configured + # values instead of the hardcoded defaults. Without this the model worker + # that enforces ``per_token_model_limit`` runs on the default (30), + # silently ignoring the YAML value. + server_specs = [a for a in self.config.applications if a.import_path == 'server'] + if len(server_specs) > 1: + logger.warning(f'{len(server_specs)} "server" applications declared; using the first ' + 'for ServerState policy env propagation.') + if server_specs: + server_state_env = server_specs[0].args.server_config.to_env_vars() + for k, v in server_state_env.items(): + os.environ[k] = v + if server_state_env: + logger.info(f'ServerState policy exported to worker env: {server_state_env}') + model_alias_map = build_model_alias_map(self.config.applications) if model_alias_map: os.environ[MODEL_ID_ALIASES_ENV] = json.dumps(model_alias_map, ensure_ascii=False) diff --git a/src/twinkle/server/lifecycle/__init__.py b/src/twinkle/server/lifecycle/__init__.py new file mode 100644 index 000000000..b77b0c75f --- /dev/null +++ b/src/twinkle/server/lifecycle/__init__.py @@ -0,0 +1,7 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Server-side request-lifecycle package (server-request-lifecycle spec). + +Holds the pieces shared by the Submit_Endpoint and Retrieve_Endpoint: the single +FutureRecord -> TaskEnvelope mapping point, the poll/window configuration, and the +Submit_Endpoint shell. +""" diff --git a/src/twinkle/server/lifecycle/envelope.py b/src/twinkle/server/lifecycle/envelope.py new file mode 100644 index 000000000..25343c9b0 --- /dev/null +++ b/src/twinkle/server/lifecycle/envelope.py @@ -0,0 +1,97 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""The one place a FutureRecord becomes a TaskEnvelope. + +Both the Submit_Endpoint and the Retrieve_Endpoint go through this function so a +``failed`` status always lands in ``error`` and never in ``result``. Duplicating +this mapping per endpoint is how a task that failed inside the Inline_Fast_Path +window loses its payload. +""" +from __future__ import annotations + +from typing import Any + +from twinkle.protocol.types.errors import ErrorCategory, ErrorPayload +from twinkle.protocol.types.lifecycle import TaskEnvelope +from twinkle.server.state.models import FutureFailureRecord +from twinkle.server.task_errors import trim_traceback + +# Keys must equal ``state.models.FAILURE_REASON_CODES`` (guarded by test_envelope). +_FAILURE_WIRE: dict[str, tuple[int, ErrorCategory]] = { + 'invalid_request': (400, ErrorCategory.User), + 'request_rejected': (400, ErrorCategory.User), + 'resource_not_found': (404, ErrorCategory.User), + 'full_mode_busy': (409, ErrorCategory.User), + 'input_tokens_exceeded': (422, ErrorCategory.User), + 'batch_size_invalid': (422, ErrorCategory.User), + 'rate_limit_exceeded': (429, ErrorCategory.User), + 'resource_quota_exceeded': (429, ErrorCategory.User), + 'cancelled': (499, ErrorCategory.User), + 'endpoint_unavailable': (501, ErrorCategory.Server), + 'state_contention': (503, ErrorCategory.Server), + 'backend_gate_unavailable': (503, ErrorCategory.Server), + 'orphaned_replica': (503, ErrorCategory.Server), + 'execution_timeout': (504, ErrorCategory.Server), + 'deadline_exceeded': (500, ErrorCategory.Server), + 'internal_error': (500, ErrorCategory.Server), +} + + +def error_payload_from_failure(stored: Any, *, request_id: str) -> ErrorPayload: + """Map one protocol-independent failure record to Twinkle's wire model.""" + failure = FutureFailureRecord.model_validate(stored) + error_code, category = _FAILURE_WIRE.get( + failure.reason_code, + (500, ErrorCategory.Server), + ) + diagnostic = failure.diagnostic if category is ErrorCategory.Server else None + return ErrorPayload( + error=failure.message[:1024], + category=category, + error_code=error_code, + request_id=request_id, + traceback=trim_traceback(diagnostic) if diagnostic else None, + details=failure.details, + ) + + +def envelope_from_record( + request_id: str, + record: dict[str, Any] | None, + *, + fallback_status: str = 'pending', +) -> TaskEnvelope: + """Map a stored ``FutureRecord`` dict to the wire ``TaskEnvelope``. + + Failed and cancelled records carry a protocol-independent ``failure`` field. + This is the only place that maps those domain reasons to Twinkle's + ``ErrorPayload`` status/category vocabulary. Legacy failures embedded in + ``result`` are intentionally unsupported because that format was never merged. + + ``completed`` with ``result is None`` is a valid success (``step`` / + ``zero_grad`` / ``lr_step`` all return ``None``); it is NOT treated as a + failure. The tinker endpoint's ``HTTPException(500, 'Task completed but no + result found')`` is a bug that this function deliberately does not copy. + + ``failed`` and ``cancelled`` both carry an ``ErrorPayload`` in ``error`` (the + cancel payload is stored the same way a failure payload is), so the client can + distinguish them by ``status`` while reading one field. + """ + record = record or {} + status = record.get('status', fallback_status) + common = dict( + queue_state=record.get('queue_state'), + queue_state_reason=record.get('queue_state_reason'), + ) + if status in ('failed', 'cancelled'): + return TaskEnvelope( + request_id=request_id, + status=status, + error=error_payload_from_failure(record.get('failure'), request_id=request_id), + **common, + ) + return TaskEnvelope( + request_id=request_id, + status=status, + result=record.get('result') if status == 'completed' else None, + **common, + ) diff --git a/src/twinkle/server/lifecycle/poll_config.py b/src/twinkle/server/lifecycle/poll_config.py new file mode 100644 index 000000000..6aa6c47f7 --- /dev/null +++ b/src/twinkle/server/lifecycle/poll_config.py @@ -0,0 +1,79 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""The single declaration point for Long_Poll_Window and the retrieve poll interval. + +Both retrieve endpoints -- twinkle's ``POST /twinkle/retrieve_future`` and tinker's +``POST /retrieve_future`` -- read their timing from here so there is one source of truth, +and both long-poll at the same fixed interval. + +Why fixed and not exponential backoff (measured, and it overturned the original design): +the spec's argument for backoff was "a 5ms zero_grad that missed the Inline_Fast_Path +window should not wait out a full 500ms tick". On real hardware that case does not exist +-- control-plane ops measure 0.00-0.09s and are absorbed by the 50ms inline window, so +they never reach this endpoint at all. What does reach it is ``forward_backward`` at +0.52-0.65s, and there a 0.05->1.0s doubling schedule checks at 0.05/0.10/0.20/0.40/0.80, +landing on 0.80 for the whole cluster, whereas a fixed 0.5s checks at 0.05/0.55/1.05 and +catches most of it at 0.55. Backoff measured ~22% SLOWER per step (0.980s vs 0.800s mean) +because its interval grows fastest exactly across the band where real tasks finish. + +A denser ceiling (~0.2s) would beat both, at 2.5x the poll rate against a shared state +backend. That is a tuning knob, not a correctness one, and it is not worth optimising for +small-model step times: at production scale a data-plane call runs for minutes and any of +these granularities is noise. +""" +from __future__ import annotations + +import os + +from twinkle.utils.logger import get_logger + +logger = get_logger() + +# Documented assumption for a typical ingress / L7 gateway idle-connection limit. +# It is NOT hard-coded into any decision logic -- it is only the threshold at which +# ``long_poll_window()`` warns that a configured window is likely to be cut off. +# Operators can override the real limit for their deployment. +_ASSUMED_GATEWAY_IDLE_LIMIT = 60.0 + +# Default Long_Poll_Window. 30 < assumed gateway limit (60) and 30 < client HTTP +# timeout (90), so a retrieve request that waits a full window survives the gateway. +_DEFAULT_LONG_POLL_TIMEOUT = 30.0 + +# Fixed poll interval for BOTH retrieve endpoints (env ``TWINKLE_POLL_INTERVAL``, default +# 0.5s). Declared here so neither endpoint reads ``os.environ`` on its own. See the module +# docstring for the measurement that rejected exponential backoff. +_DEFAULT_POLL_INTERVAL = 0.5 + +# The last window value we warned about. ``long_poll_window()`` is on the hot path of both +# retrieve endpoints -- not just startup -- so an unguarded warning would fire on every +# retrieve request (roughly twice a second during training) to say something that only +# needs saying once. Keyed on the value rather than a bare bool so that a *changed* +# misconfiguration warns again instead of being swallowed by the first one. +_warned_window: float | None = None + + +def long_poll_window() -> float: + """Return the Long_Poll_Window in seconds (env ``TWINKLE_LONG_POLL_TIMEOUT``). + + Warns **once per configured value** when that value is at least the assumed gateway + idle limit: the retrieve endpoint is itself served through the gateway, so a window + past the gateway's limit would recreate the connection-cut problem this spec removes. + + The env var is re-read on every call (callers may change it, and tests do), so the + warning needs its own de-duplication -- this function runs per retrieve request, not + only at startup. + """ + global _warned_window + value = float(os.environ.get('TWINKLE_LONG_POLL_TIMEOUT', str(_DEFAULT_LONG_POLL_TIMEOUT))) + if value >= _ASSUMED_GATEWAY_IDLE_LIMIT and value != _warned_window: + _warned_window = value + logger.warning( + '[poll_config] TWINKLE_LONG_POLL_TIMEOUT=%.1fs is >= the assumed gateway idle limit ' + '(%.1fs). The retrieve endpoint is served through the gateway too, so a window this ' + 'large may be cut off mid-request. Lower it or confirm your gateway idle limit.', value, + _ASSUMED_GATEWAY_IDLE_LIMIT) + return value + + +def retrieve_poll_interval() -> float: + """Fixed poll interval shared by both retrieve endpoints (env ``TWINKLE_POLL_INTERVAL``).""" + return float(os.environ.get('TWINKLE_POLL_INTERVAL', str(_DEFAULT_POLL_INTERVAL))) diff --git a/src/twinkle/server/lifecycle/protocols.py b/src/twinkle/server/lifecycle/protocols.py new file mode 100644 index 000000000..c05ea12bb --- /dev/null +++ b/src/twinkle/server/lifecycle/protocols.py @@ -0,0 +1,60 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Static declarations of what a deployment class must provide. + +Not a new layer and not a base class: these are ``typing.Protocol`` declarations that +turn the implicit host contract of ``run_submit`` / ``input_metrics`` / the queue mixins +into something a type checker can check. Today those requirements are satisfied by duck +typing and ``getattr`` fallbacks, so a missing attribute surfaces at run time -- sometimes +inside a coroutine sitting in the compute queue. + +Deliberately two layers, matching the two real shapes: every queued deployment +(Gateway/Model/Sampler/Processor) satisfies ``QueuedDeployment``; only ``ModelManagement`` +satisfies ``DataParallelDeployment`` (it is the one with ``data_world_size``). +``DataPlaneManagement`` satisfies neither -- it has no ``state`` and opts out of the +cleanup middleware via ``attach_cleanup_middleware=False``; the gate below is what makes +that fact visible statically instead of only via that boolean. +""" +from __future__ import annotations + +from fastapi import Request +from typing import Any, Protocol, runtime_checkable + +from twinkle.server.lifecycle.envelope import TaskEnvelope +from twinkle.server.state import ServerState +from twinkle.server.task_queue.config import TaskQueueConfig + + +@runtime_checkable +class QueuedDeployment(Protocol): + """A deployment that admits requests through the compute queue.""" + + state: ServerState + replica_id: str + + @property + def task_queue_config(self) -> TaskQueueConfig: + ... + + async def _on_request_start(self, request: Request) -> str: + ... + + def assert_resource_exists(self, resource_id: str | None) -> None: + ... + + async def _peek_terminal(self, request_id: str, *, fallback_status: str) -> TaskEnvelope: + ... + + async def submit_and_peek(self, *args: Any, **kwargs: Any) -> TaskEnvelope: + ... + + async def call_backend(self, fn: Any, /, *args: Any, admit: bool = True, **kwargs: Any) -> Any: + ... + + +@runtime_checkable +class DataParallelDeployment(QueuedDeployment, Protocol): + """A queued deployment that also shards a batch across data-parallel ranks.""" + + @property + def data_world_size(self) -> int: + ... diff --git a/src/twinkle/server/lifecycle/submit.py b/src/twinkle/server/lifecycle/submit.py new file mode 100644 index 000000000..b1b409a47 --- /dev/null +++ b/src/twinkle/server/lifecycle/submit.py @@ -0,0 +1,220 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Submit_Endpoint shell and the named seams every queued handler shares. + +The Inline_Fast_Path wait itself (``submit_and_peek``) lives on +:class:`~twinkle.server.task_queue.mixin.TaskQueueMixin`, since it operates on +queue state; this module owns the request-shaped pieces around it. +""" +from __future__ import annotations + +import uuid +from collections.abc import Callable, Coroutine +from fastapi import Request +from typing import TYPE_CHECKING, Any + +from twinkle.data_format import InputFeature, Trajectory, is_encoded +from twinkle.protocol.types.base import FieldRole, fields_with_role +from twinkle.protocol.types.data import export_batch +from twinkle.protocol.types.lifecycle import TaskEnvelope +from twinkle.server.middleware.auth import get_session_id_from_request +from twinkle.server.validation import assert_request_supported + +if TYPE_CHECKING: + from twinkle.server.lifecycle.protocols import DataParallelDeployment, QueuedDeployment + +# --------------------------------------------------------------------------- # +# Named seams shared by every queued twinkle-native handler. +# --------------------------------------------------------------------------- # + + +def to_backend_inputs(inputs: Any, *, single: bool = False) -> Any: + """Seam A: export wire-validated ``inputs`` as the objects the backend consumes. + + This is an *export*, not a validation step. The request model declares ``inputs`` + as :data:`~twinkle.protocol.types.data.WireInputBatch`, so a malformed batch is + already rejected during FastAPI body parsing -- before a future record exists and + before anything reaches a GPU. Validating here instead would put the first check + inside the queued task, where a rejection has already cost an enqueue. + + Entries arrive as wire models and are exported with ``exclude_none`` semantics, so + unset optional fields stay absent (Twinkle_Core branches on key presence) and + unknown keys the caller sent are preserved. ``InputFeature`` / ``Trajectory`` are + ``TypedDict``s, so constructing them is a plain dict build. + + With ``single=True`` exactly one object is returned (the streaming path accepts + only one input) and a batch of any other size is a ``ValueError``. Plain dicts pass + through unchanged: the data-plane path resolves rows itself and never goes through + the wire schema. + """ + entries = export_batch(inputs) if isinstance(inputs, list) else inputs + if single: + if isinstance(entries, list): + if len(entries) != 1: + raise ValueError('Streaming only supports a single input') + entries = entries[0] + if isinstance(entries, dict): + return _as_backend_entry(entries) + return entries + if isinstance(entries, list): + return [_as_backend_entry(entry) if isinstance(entry, dict) else entry for entry in entries] + if isinstance(entries, dict): + return [_as_backend_entry(entries)] + return entries + + +def _as_backend_entry(entry: dict[str, Any]) -> Any: + """One exported entry as its ``TypedDict`` shape.""" + return InputFeature(**entry) if is_encoded(entry) else Trajectory(**entry) + + +def backend_kwargs(body: Any) -> dict[str, Any]: + """Seam B: the keyword arguments forwarded to the backend call. + + Exactly two sources, both declared on the request model (see + :mod:`twinkle.protocol.types.base`): + + 1. fields whose role is ``BackendKwarg``, included iff their value is not ``None``; + 2. the contents of every ``Passthrough`` field, flattened. + + Control fields are never forwarded. That exclusion is the point of the field roles: + forwarding *all* declared non-``None`` fields would re-send ``inputs`` / + ``adapter_name`` / ``seq_id``, which the handlers already pass explicitly -- a + duplicate keyword argument at best, and a protocol field leaking into a backend + signature at worst. + """ + model_cls = type(body) + kwargs: dict[str, Any] = {} + for name in fields_with_role(model_cls, FieldRole.BackendKwarg): + value = getattr(body, name, None) + if value is not None: + kwargs[name] = value + for name in fields_with_role(model_cls, FieldRole.Passthrough): + region = getattr(body, name, None) or {} + overlap = set(region) & set(kwargs) + if overlap: + raise ValueError(f'{name} collides with declared backend parameters: {", ".join(sorted(overlap))}') + kwargs.update(region) + return kwargs + + +def input_metrics(self: DataParallelDeployment, body: Any, *, data_parallel: bool = False) -> dict[str, Any]: + """Seam C: scheduling metrics (input_tokens, and batch_size/data_world_size). + + Reads validated wire models, so no isinstance guards: ``inputs`` is a list and + ``input_ids`` is either absent or a list of ints. + """ + inputs = body.inputs + input_tokens = sum(len(getattr(entry, 'input_ids', None) or ()) for entry in inputs) + metrics: dict[str, Any] = {'input_tokens': input_tokens} + if data_parallel: + metrics['batch_size'] = len(inputs) + metrics['data_world_size'] = self.data_world_size + return metrics + + +def resolve_twinkle_adapter_name(request: Request, adapter_name: str | None) -> str | None: + """Build a stable per-session adapter name, falling back to request_id for older clients.""" + if adapter_name is None or adapter_name == '': + return None + owner_id = get_session_id_from_request(request) or request.state.request_id + return owner_id + '-' + adapter_name + + +async def run_submit( + self: QueuedDeployment, + request: Request, + body: Any, + *, + task_type: str, + backend_call: Callable[..., Coroutine], + metrics: Callable[[Any, Any], dict[str, Any]] | None = None, + assert_resource: bool = True, + capability: str | None = None, +) -> TaskEnvelope: + """The common Submit_Endpoint judgment sequence, called by every queued + twinkle-native handler instead of being repeated in each. + + Order is load-bearing: request start -> adapter resolution -> preflight -> + ``submit_and_peek`` (whose ``schedule_task`` runs its own resource preflight). + Every admission check runs before any state write, so a rejected request writes + nothing. + + ``assert_request_supported`` is the one place the request is checked against *this + deployment*: a parameter that only exists on the other backend and an endpoint this + backend does not implement are decided here -- before the seq claim and before the + enqueue, so a rejected request runs on zero data-parallel ranks. Putting these checks + in the queued task instead would let an incompatible request cost a full GPU dispatch. + Passthrough keys are forwarded unjudged (no spelling check): see + :mod:`twinkle.server.validation.backend_compat`. + + A plain helper, not a signature-rewriting decorator: each handler keeps its natural + FastAPI signature so the app stays shallow enough for Ray Serve to cloudpickle (a + signature-patching wrapper once deepened the route graph past CPython's C-stack + recursion guard during ``serve.ingress``). + + Four queued endpoints deliberately do NOT route through here and call + ``submit_and_peek`` / ``submit_background_and_peek`` directly: model + ``add_adapter_to_model`` (creates the adapter, so the resource assertion cannot + apply and it owns the train_mode/full-mode checks), model ``upload_to_hub`` (pure + I/O, background task), and sampler ``sample`` / ``sample_to_data_plane`` (no adapter + semantics). Zero-write and admission still hold for them because both live inside + ``schedule_task`` -> ``_perform_preflight_checks``, not in this shell. + + ``backend_call(self, body, adapter_name, token)`` runs the endpoint-specific call + and returns the JSON-safe task result. ``metrics(self, body)`` supplies scheduling + kwargs; omit it for control-plane ops. ``assert_resource`` guards on the adapter + existing before the work runs; set it False for endpoints that create/drop it. + ``capability`` names the backend capability the endpoint needs, when the endpoint is + not implemented by every backend. + """ + token = await self._on_request_start(request) + adapter_name = resolve_twinkle_adapter_name(request, body.adapter_name) + + # ---- Preflight: decidable from the body plus this deployment's backend, so it + # runs before any state write and before the enqueue. ---- + assert_request_supported(self, body, capability=capability) + + schedule_kwargs = metrics(self, body) if metrics is not None else {} + + async def _task(): + if assert_resource: + self.assert_resource_exists(adapter_name) + return await backend_call(self, body, adapter_name, token) + + # ---- Idempotent dedup: a client-supplied seq_id makes a retried stateful op + # apply at most once. Claim (session_id, adapter, seq_id) -> request_id atomically + # before enqueue; a hit returns the original task's envelope instead of re-enqueuing. + # Only grad-mutating client calls set seq_id, so other endpoints skip this. + # + # The adapter must be part of the key. Each client model object owns its own seq + # counter starting at 1, while session_id is process-global -- so two adapters + # trained from one process would collide on (session, seq) and the second + # forward_backward would be dropped as a duplicate AND handed the first adapter's + # loss. That is a silent wrong-result bug, i.e. the exact failure this dedup + # exists to prevent, one level up. ---- + request_id = f'req_{uuid.uuid4().hex}' + seq_id = getattr(body, 'seq_id', None) + dedup_key = None + if seq_id is not None: + session_id = get_session_id_from_request(request) or request.state.request_id + dedup_key = f'seq::{session_id}::{adapter_name or "-"}::{seq_id}' + ttl = int(self.task_queue_config.effective_execution_timeout) + 60 + prior_request_id = await self.state.claim_seq(dedup_key, request_id, ttl) + if prior_request_id is not None: + return await self._peek_terminal(prior_request_id, fallback_status='pending') + + # ---- Decision_Boundary: preflight (in schedule_task) then peek ---- + try: + return await self.submit_and_peek( + _task, model_id=adapter_name, token=token, task_type=task_type, request_id=request_id, **schedule_kwargs) + except Exception: + # Release the seq claim only when the task never made it onto the queue -- + # decided by whether a future record exists, NOT by the exception type. A + # preflight rejection raises before any record is written, so releasing lets + # a retry re-enqueue. But a failure *after* enqueue (e.g. a transient state + # error inside the peek) leaves a live task that will still run; releasing + # there would let a retry enqueue a duplicate -> the exact double-apply this + # dedup prevents. When unsure (record exists), keep the claim. + if dedup_key is not None and await self.state.get_future(request_id) is None: + await self.state.release_seq(dedup_key) + raise diff --git a/src/twinkle/server/middleware/__init__.py b/src/twinkle/server/middleware/__init__.py new file mode 100644 index 000000000..9bb98c0ac --- /dev/null +++ b/src/twinkle/server/middleware/__init__.py @@ -0,0 +1,6 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""HTTP middleware for the server deployments. + +Sits next to ``deployment.py``'s middleware stack: ``auth.verify_request_token`` is the +token-verification middleware registered by ``build_deployment_app``. +""" diff --git a/src/twinkle/server/utils/validation.py b/src/twinkle/server/middleware/auth.py similarity index 97% rename from src/twinkle/server/utils/validation.py rename to src/twinkle/server/middleware/auth.py index ef779cd27..c3b7f2456 100644 --- a/src/twinkle/server/utils/validation.py +++ b/src/twinkle/server/middleware/auth.py @@ -3,7 +3,7 @@ from fastapi.responses import JSONResponse from typing import Any -from twinkle_client.http.headers import H_AUTH, H_AUTH_TWINKLE, H_REQUEST_ID +from twinkle.protocol.headers import H_AUTH, H_AUTH_TWINKLE, H_REQUEST_ID _OPENAI_COMPAT_SUFFIXES = ('/chat/completions', '/models') diff --git a/src/twinkle/server/model/app.py b/src/twinkle/server/model/app.py index b3589ffee..e55e4a8d1 100644 --- a/src/twinkle/server/model/app.py +++ b/src/twinkle/server/model/app.py @@ -7,30 +7,30 @@ """ from __future__ import annotations -from fastapi import FastAPI, Request +import asyncio +from fastapi import FastAPI, HTTPException, Request from ray import serve from ray.serve.config import RequestRouterConfig from typing import Any from twinkle import DeviceGroup -from twinkle.server.common.router import StickyLoraRequestRouter -from twinkle.server.deployment import LazyCleanupMixin, bind_deployment, build_deployment_app, init_twinkle_runtime +from twinkle.server.config.backend_dispatch import BackendSelector +from twinkle.server.deployment import LazyCleanupMixin, bind_deployment, build_deployment_app from twinkle.server.exceptions import FullModeBusyError +from twinkle.server.middleware.auth import get_token_from_request +from twinkle.server.model.routing import StickyLoraRequestRouter +from twinkle.server.runtime import init_twinkle_runtime +from twinkle.server.session_resource import AdapterManagerMixin from twinkle.server.state import ServerState, get_server_state +from twinkle.server.task_queue import TaskQueueConfig, TaskQueueMixin from twinkle.server.utils import wrap_builder_with_device_group_env -from twinkle.server.utils.backend_dispatch import BackendSelector -from twinkle.server.utils.lifecycle import AdapterManagerMixin -from twinkle.server.utils.task_queue import TaskQueueConfig, TaskQueueMixin -from twinkle.server.utils.validation import get_token_from_request from twinkle.utils.logger import get_logger -from .tinker_handlers import _register_tinker_routes -from .twinkle_handlers import _register_twinkle_routes +from .tinker_handlers import _register_model_tinker_routes +from .twinkle_handlers import _register_model_twinkle_routes logger = get_logger() -# ``FullModeBusyError`` lives in ``twinkle.server.exceptions``; re-exported here -# for backwards compatibility with callers importing it from this module. -__all__ = ['FullModeBusyError', 'ModelManagement', 'build_model_app'] +__all__ = ['ModelManagement', 'build_model_app'] # Ctor kwargs consumed by the MultiLora wrappers' signatures but unknown to the # plain (full-parameter) model classes, where **kwargs flows into HF @@ -131,9 +131,18 @@ async def __init__(self, from twinkle.server.data_plane import DataPlaneProxy self.data_plane = DataPlaneProxy(data_plane_url) self._replica_registered = False - - # Initialize mixins - self._init_task_queue(queue_config, deployment_name='Model') + self._model_unhealthy = False + self._health_probe_task: asyncio.Task | None = None + + actors = getattr(self.model, '_actors', None) + self._init_task_queue( + queue_config, + deployment_name='Model', + enable_admission_gate=True, + on_backend_timeout=self._probe_after_timeout, + collect_width=len(actors) if actors else 1, + ) + self.model._ray_get_timeout = self.task_queue_config.effective_execution_timeout self._init_adapter_manager(**(adapter_config or {})) await self._register_replica_on_startup() # Note: countdown task is started lazily in _ensure_sticky() @@ -152,6 +161,7 @@ async def _register_replica_on_startup(self) -> None: """Register this replica's capacity before Ray Serve marks it ready.""" if not self._replica_registered: await self.state.register_replica(self.replica_id, self.max_loras) + await self.state.touch_replica_last_seen(self.replica_id) self._replica_registered = True @serve.multiplexed(max_num_models_per_replica=5) @@ -166,9 +176,11 @@ async def _ensure_sticky(self): async def _on_request_start(self, request: Request) -> str: await self._ensure_sticky() + await self.state.touch_replica_last_seen(self.replica_id) await self._ensure_state_cleanup_started() - token = get_token_from_request(request) - return token + if self._model_unhealthy: + raise HTTPException(status_code=503, detail='Model actors are unavailable') + return get_token_from_request(request) async def shutdown(self) -> None: """Explicit async cleanup — called via FastAPI shutdown event.""" @@ -176,32 +188,49 @@ async def shutdown(self) -> None: await self.state.unregister_replica(self.replica_id) except Exception: pass + await self.shutdown_task_queue() await self.data_plane.close() - def check_model_health(self) -> dict: - """Probe model actors liveness via a lightweight ping. - - Returns a dict with 'healthy' (bool) and 'detail' (str). - If the model actors are dead (e.g. OOM/SIGSEGV), the ping call - will raise RayActorError, signalling the watchdog to restart. - """ + async def _run_model_health_probe(self) -> dict: try: - result = self.model.ping() + result = await self.call_backend(self.model.ping, admit=False) if result is True: + self._model_unhealthy = False return {'healthy': True, 'detail': 'model actors alive'} + self._model_unhealthy = True return {'healthy': False, 'detail': f'unexpected ping result: {result}'} except Exception as e: + self._model_unhealthy = True return {'healthy': False, 'detail': f'model actor unreachable: {e}'} + async def check_model_health(self) -> dict: + """Run one coalesced actor probe outside the event loop.""" + current = getattr(self, '_health_probe_task', None) + if current is None or current.done(): + self._health_probe_task = asyncio.create_task(self._run_model_health_probe()) + return await asyncio.shield(self._health_probe_task) + + def mark_unhealthy(self) -> None: + """Flag the deployment unhealthy; /healthz returns 503 until a probe recovers it.""" + self._model_unhealthy = True + + async def _probe_after_timeout(self) -> None: + """Fired by ComputeWorker on a backend timeout: probe and log liveness.""" + result = await self.check_model_health() + logger.warning('[Model] post-timeout liveness probe: %s', result) + async def _cleanup_adapter(self, adapter_name: str) -> None: if self.get_resource_info(adapter_name): - self.clear_resource_state(adapter_name) if self.train_mode == 'full': # No PEFT adapter to remove; restore clean base weights so the # next tenant does not inherit this tenant's trained weights. - self.model.reload_initial_weights() + # Takes the Admission_Gate: this path is driven by the background + # countdown and never enters Task_Queue, so the gate is what keeps + # it from colliding with an in-flight training call. + await self.call_backend(self.model.reload_initial_weights) else: - self.model.remove_adapter(adapter_name) + await self.call_backend(self.model.remove_adapter, adapter_name) + self.clear_resource_state(adapter_name) self.unregister_resource(adapter_name) await self.state.unload_model(adapter_name) @@ -232,9 +261,9 @@ def assert_full_mode_available(self, adapter_name: str | None = None) -> None: """ if not self.is_full_mode: return - for rid, info in self._resource_records.items(): - if rid != adapter_name and not info.get('expiring'): - raise FullModeBusyError(rid) + holder = self.find_active_resource(exclude=adapter_name) + if holder is not None: + raise FullModeBusyError(holder) async def _on_adapter_expired(self, adapter_name: str) -> None: self.fail_pending_tasks_for_model(adapter_name, reason='Adapter expired') @@ -280,8 +309,8 @@ def build_model_app(model_id: str, # teardown via ``on_shutdown`` and its sticky-LoRA router via # ``request_router_config``. def register_routes(app: FastAPI, get_self: Any) -> None: - _register_tinker_routes(app, get_self) - _register_twinkle_routes(app, get_self) + _register_model_tinker_routes(app, get_self) + _register_model_twinkle_routes(app, get_self) async def _on_shutdown(servable: Any) -> None: await servable.shutdown() diff --git a/src/twinkle/server/model/backends/megatron_model.py b/src/twinkle/server/model/backends/megatron_model.py index 56f9acc98..b06228248 100644 --- a/src/twinkle/server/model/backends/megatron_model.py +++ b/src/twinkle/server/model/backends/megatron_model.py @@ -11,15 +11,14 @@ """ import torch from tinker import types -from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union from twinkle import remote_class, remote_function from twinkle.data_format import InputFeature, Trajectory from twinkle.infra import collect_tensor_dict from twinkle.model.megatron import MegatronModel, MultiLoraMegatronModel -from twinkle.server.common.datum import datum_to_input_feature, extract_rl_features_for_loss from twinkle.server.model.backends.common import (TwinkleCompatModelBase, clean_metrics, collect_forward_backward_results, to_cpu_safe_output) +from twinkle.server.model.tinker_datum import datum_to_input_feature, extract_rl_features_for_loss from twinkle.utils.nccl_safe import nccl_safe_megatron @@ -33,7 +32,7 @@ class in the MRO. For full-parameter training the ``adapter_name`` is the """ @remote_function(dispatch='slice_dp', collect=collect_forward_backward_results, sync=True) - @nccl_safe_megatron(tinker=True) + @nccl_safe_megatron def tinker_forward_backward(self, *, inputs: list[types.Datum], adapter_name: str, loss_fn: str, **kwargs): """Combined forward and backward pass.""" self._tinker_setup_loss(loss_fn, inputs, adapter_name, kwargs) @@ -54,7 +53,7 @@ def tinker_forward_backward(self, *, inputs: list[types.Datum], adapter_name: st return [results, loss] @remote_function(dispatch='slice_dp', collect=collect_forward_backward_results) - @nccl_safe_megatron(tinker=True) + @nccl_safe_megatron def tinker_forward_only(self, *, inputs: list[types.Datum], adapter_name: str = None, **kwargs): """Forward pass without gradient computation.""" template = self.get_template(adapter_name) @@ -102,26 +101,24 @@ def tinker_calculate_metric(self, is_training, **kwargs): metric = super().calculate_metric(is_training, **kwargs) return clean_metrics(metric) - @remote_function(dispatch='all', sync=True) - def tinker_load(self, checkpoint_dir: str, **kwargs): - """Load checkpoint with token-based isolation support.""" - token = kwargs.pop('token', None) - if not token: - raise ValueError('Token is required for loading checkpoints') - from twinkle.server.checkpoint import create_checkpoint_manager - checkpoint_manager = create_checkpoint_manager(token, client_type='tinker') - resolved = checkpoint_manager.resolve_load_path(checkpoint_dir) - if resolved.is_twinkle_path: - return super().load(name=resolved.checkpoint_name, output_dir=resolved.checkpoint_dir, **kwargs) - else: - return super().load(name=resolved.checkpoint_name, **kwargs) + @remote_function(dispatch='all', sync=True, timeout=3600) + def tinker_load(self, *, checkpoint_name: str, output_dir: str | None = None, **kwargs): + """Load a checkpoint from an already-resolved location. + + Path resolution (token isolation, twinkle-vs-external path shapes) belongs to the + handler layer: it is a server storage policy, and this class runs inside a Ray + actor as a compute backend. + """ + if output_dir is not None: + return super().load(name=checkpoint_name, output_dir=output_dir, **kwargs) + return super().load(name=checkpoint_name, **kwargs) # ------------------------------------------------------------------ # Twinkle-native methods (InputFeature/Trajectory-based I/O) # ------------------------------------------------------------------ @remote_function(dispatch='slice_dp', collect=collect_tensor_dict) - @nccl_safe_megatron(forward_only=True) + @nccl_safe_megatron def forward_only(self, *, inputs: InputFeature | list[InputFeature] | Trajectory | list[Trajectory], **kwargs): """Forward-only for twinkle-native clients (InputFeature/Trajectory I/O).""" output = super().forward_only(inputs=inputs, **kwargs) @@ -135,7 +132,7 @@ def forward_backward(self, *, inputs: InputFeature | list[InputFeature] | Trajec output = super().forward_backward(inputs=inputs, **kwargs) return to_cpu_safe_output(output) - @remote_function(collect='first', lazy_collect=False) + @remote_function(collect='first', lazy_collect=False, sync=True, timeout=4) def ping(self) -> bool: """Lightweight liveness probe for watchdog health checks.""" return True diff --git a/src/twinkle/server/model/backends/mock_model.py b/src/twinkle/server/model/backends/mock_model.py index 7bf9866cc..d204090a8 100644 --- a/src/twinkle/server/model/backends/mock_model.py +++ b/src/twinkle/server/model/backends/mock_model.py @@ -140,7 +140,7 @@ def calculate_metric(self, *args: Any, **kwargs: Any) -> dict[str, float]: return {'loss': 0.5, 'grad_norm': 0.1} @remote_function() - def tinker_load(self, checkpoint_dir: str, **kwargs: Any) -> None: + def tinker_load(self, *, checkpoint_name: str, output_dir: str | None = None, **kwargs: Any) -> None: return None # ----- Configuration setters ----------------------------------------- # @@ -240,7 +240,7 @@ def remove_adapter(self, adapter_name: str) -> None: def has_adapter(self, adapter_name: str) -> bool: return adapter_name in self._adapters - @remote_function(collect='first', lazy_collect=False) + @remote_function(collect='first', lazy_collect=False, sync=True, timeout=4) def ping(self) -> bool: """Lightweight liveness probe for watchdog health checks.""" return True diff --git a/src/twinkle/server/model/backends/transformers_model.py b/src/twinkle/server/model/backends/transformers_model.py index 8dc503bb0..c36e53c89 100644 --- a/src/twinkle/server/model/backends/transformers_model.py +++ b/src/twinkle/server/model/backends/transformers_model.py @@ -13,17 +13,15 @@ (InputFeature/Trajectory-based I/O) via /twinkle/* endpoints. """ from tinker import types -from typing import List, Union from twinkle import remote_class, remote_function from twinkle.data_format import InputFeature, Trajectory from twinkle.infra import collect_tensor_dict from twinkle.model import MultiLoraTransformersModel from twinkle.model.transformers import TransformersModel -from twinkle.server.common.datum import datum_to_input_feature, extract_rl_features_for_loss from twinkle.server.model.backends.common import (TwinkleCompatModelBase, clean_metrics, collect_forward_backward_results, to_cpu_safe_output) -from twinkle.utils.nccl_safe import nccl_safe +from twinkle.server.model.tinker_datum import datum_to_input_feature, extract_rl_features_for_loss class _TransformersTinkerCompatMixin(TwinkleCompatModelBase): @@ -48,7 +46,6 @@ def tinker_forward_only(self, *, inputs: list[types.Datum], adapter_name: str = return [results, 0.0] @remote_function(dispatch='slice_dp', collect=collect_forward_backward_results) - @nccl_safe(tinker=True) def tinker_forward_backward(self, *, inputs: list[types.Datum], adapter_name: str, loss_fn: str, **kwargs): self._tinker_setup_loss(loss_fn, inputs, adapter_name, kwargs) template = self.get_template(adapter_name) @@ -83,18 +80,18 @@ def tinker_calculate_metric(self, is_training, **kwargs): return clean_metrics(metric) @remote_function() - def tinker_load(self, checkpoint_dir: str, **kwargs): - """Load checkpoint with token-based isolation support.""" - token = kwargs.pop('token', None) - if not token: - raise ValueError('Token is required for loading checkpoints') - from twinkle.server.checkpoint import create_checkpoint_manager - checkpoint_manager = create_checkpoint_manager(token, client_type='tinker') - resolved = checkpoint_manager.resolve_load_path(checkpoint_dir) - if resolved.is_twinkle_path: - return super().load(name=resolved.checkpoint_name, output_dir=resolved.checkpoint_dir, **kwargs) - else: - return super().load(name=resolved.checkpoint_name, **kwargs) + def tinker_load(self, *, checkpoint_name: str, output_dir: str | None = None, **kwargs): + """Load a checkpoint from an already-resolved location. + + Path resolution (token isolation, twinkle-vs-external path shapes) belongs to the + handler layer: it is a server storage policy, and this class runs inside a Ray + actor as a compute backend. Keeping it here meant a checkpoint-layout change had + to touch GPU-side code -- the hardest layer to test -- and the resolution block + was duplicated verbatim in the megatron backend. + """ + if output_dir is not None: + return super().load(name=checkpoint_name, output_dir=output_dir, **kwargs) + return super().load(name=checkpoint_name, **kwargs) # ------------------------------------------------------------------ # Twinkle-native methods (InputFeature/Trajectory-based I/O) @@ -107,14 +104,13 @@ def forward_only(self, *, inputs: InputFeature | list[InputFeature] | Trajectory return to_cpu_safe_output(output) @remote_function(dispatch='slice_dp', collect=collect_tensor_dict) - @nccl_safe def forward_backward(self, *, inputs: InputFeature | list[InputFeature] | Trajectory | list[Trajectory], **kwargs): """Forward+backward for twinkle-native clients (InputFeature/Trajectory I/O).""" self._normalize_ref_outputs(kwargs) output = super().forward_backward(inputs=inputs, **kwargs) return to_cpu_safe_output(output) - @remote_function(collect='first', lazy_collect=False) + @remote_function(collect='first', lazy_collect=False, sync=True, timeout=4) def ping(self) -> bool: """Lightweight liveness probe for watchdog health checks.""" return True diff --git a/src/twinkle/server/model/utils.py b/src/twinkle/server/model/data_plane_inputs.py similarity index 98% rename from src/twinkle/server/model/utils.py rename to src/twinkle/server/model/data_plane_inputs.py index aed4b3c20..fd0c6e223 100644 --- a/src/twinkle/server/model/utils.py +++ b/src/twinkle/server/model/data_plane_inputs.py @@ -5,7 +5,7 @@ import asyncio from typing import Any -from twinkle_client.common.json_utils import json_safe +from twinkle.protocol.json_utils import json_safe def model_result_rows(result: Any, batch_size: int) -> list[dict[str, Any]]: diff --git a/src/twinkle/server/common/router.py b/src/twinkle/server/model/routing.py similarity index 91% rename from src/twinkle/server/common/router.py rename to src/twinkle/server/model/routing.py index e0ccef069..02b7669a3 100644 --- a/src/twinkle/server/common/router.py +++ b/src/twinkle/server/model/routing.py @@ -1,8 +1,5 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -# Moved from tinker/common/router.py — logic unchanged. -from ray.serve.request_router import (FIFOMixin, MultiplexMixin, PendingRequest, ReplicaID, ReplicaResult, - RequestRouter, RunningReplica) -from typing import Dict, List, Optional +from ray.serve.request_router import FIFOMixin, MultiplexMixin, PendingRequest, ReplicaID, RequestRouter, RunningReplica from twinkle.server.state import ServerState, get_server_state from twinkle.utils.logger import get_logger diff --git a/src/twinkle/server/common/datum.py b/src/twinkle/server/model/tinker_datum.py similarity index 100% rename from src/twinkle/server/common/datum.py rename to src/twinkle/server/model/tinker_datum.py diff --git a/src/twinkle/server/model/tinker_handlers.py b/src/twinkle/server/model/tinker_handlers.py index c3df70efc..5eca201b0 100644 --- a/src/twinkle/server/model/tinker_handlers.py +++ b/src/twinkle/server/model/tinker_handlers.py @@ -1,31 +1,33 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -""" -Tinker-compatible model handler mixin. +"""Tinker-compatible routes for the Model deployment. -All endpoints are prefixed /tinker/... and use schedule_task() returning UntypedAPIFuture. -self_fn is injected via FastAPI Depends to obtain the ModelManagement instance at request time. +Registered by ``_register_model_tinker_routes(app, self_fn)`` -- module-level route +registration closing over ``self_fn`` via ``Depends``, not a mixin: there is no +inheritance relationship with the deployment class. All endpoints are prefixed +/tinker/... and use schedule_task() returning UntypedAPIFuture. ``self_fn`` is injected +via FastAPI Depends to obtain the ModelManagement instance at request time. """ from __future__ import annotations import traceback from collections.abc import Callable from fastapi import Depends, FastAPI, Request -from peft import LoraConfig from tinker import types -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING if TYPE_CHECKING: from .app import ModelManagement from twinkle.server.checkpoint import create_checkpoint_manager, create_training_run_manager from twinkle.server.exceptions import FullModeBusyError +from twinkle.server.task_queue.types import UserTaskError from twinkle.server.utils import get_template_for_model from twinkle.utils.logger import get_logger logger = get_logger() -def _register_tinker_routes(app: FastAPI, self_fn: Callable[[], ModelManagement]) -> None: +def _register_model_tinker_routes(app: FastAPI, self_fn: Callable[[], ModelManagement]) -> None: """Register all /tinker/* routes on the given FastAPI app. self_fn is a zero-argument callable that returns the current ModelManagement @@ -45,17 +47,11 @@ async def _create_adapter(): try: # Validate lora_config against the deployment's train_mode up front. if self.is_full_mode and body.lora_config: - return types.RequestFailedResponse( - error='This deployment runs in full-parameter (exclusive) mode; do not pass ' - 'lora_config. Use create_full_training_client (or omit lora_config).', - category=types.RequestErrorCategory.User, - ) + raise UserTaskError('This deployment runs in full-parameter (exclusive) mode; do not pass ' + 'lora_config. Use create_full_training_client (or omit lora_config).') if (not self.is_full_mode) and (not body.lora_config): - return types.RequestFailedResponse( - error='This deployment runs in LoRA mode; lora_config is required. ' - 'Use create_lora_training_client.', - category=types.RequestErrorCategory.User, - ) + raise UserTaskError('This deployment runs in LoRA mode; lora_config is required. ' + 'Use create_lora_training_client.') # Exclusive full-parameter training: reject early (before touching # state) if another tenant already holds the deployment. if self.is_full_mode: @@ -69,34 +65,33 @@ async def _create_adapter(): template = get_template_for_model(self.base_model) if self.is_full_mode: self.register_resource(adapter_name, token, session_id=body.session_id) - self.model.set_template(template, adapter_name=model_adapter, model_id=self.base_model) - self.model.set_processor('InputProcessor', adapter_name=model_adapter) - self.model.set_optimizer('Adam', adapter_name=model_adapter) + await self.call_backend( + self.model.set_template, template, adapter_name=model_adapter, model_id=self.base_model) + await self.call_backend(self.model.set_processor, 'InputProcessor', adapter_name=model_adapter) + await self.call_backend(self.model.set_optimizer, 'Adam', adapter_name=model_adapter) self.set_resource_state(adapter_name, 'grad_ready', False) else: - # TODO: Make LoraConfig more flexible + from peft import LoraConfig lora_cfg = LoraConfig(r=body.lora_config.rank, target_modules='all-linear') self.register_resource(adapter_name, token, session_id=body.session_id) - self.model.add_adapter_to_model(adapter_name=adapter_name, config_or_dir=lora_cfg) - self.model.set_template(template, adapter_name=adapter_name, model_id=self.base_model) - self.model.set_processor('InputProcessor', adapter_name=adapter_name) - self.model.set_optimizer('Adam', adapter_name=adapter_name) + await self.call_backend( + self.model.add_adapter_to_model, adapter_name=adapter_name, config_or_dir=lora_cfg) + await self.call_backend( + self.model.set_template, template, adapter_name=adapter_name, model_id=self.base_model) + await self.call_backend(self.model.set_processor, 'InputProcessor', adapter_name=adapter_name) + await self.call_backend(self.model.set_optimizer, 'Adam', adapter_name=adapter_name) self.set_resource_state(adapter_name, 'grad_ready', False) training_run_manager = create_training_run_manager(token, client_type='tinker') training_run_manager.save(_model_id, body) return types.CreateModelResponse(model_id=_model_id) except FullModeBusyError as e: - # Nothing was registered yet (check runs before register_model). - return types.RequestFailedResponse(error=str(e), category=types.RequestErrorCategory.User) + raise UserTaskError(str(e)) from e except Exception: if _model_id: adapter_name = self.get_adapter_name(adapter_name=_model_id) await self._cleanup_adapter(adapter_name) logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise return await self.schedule_task(_create_adapter, token=token, task_type='create_model') @@ -147,8 +142,8 @@ async def _do_forward(): model_adapter = self.resolve_model_adapter_name(adapter_name) datum_list = body.forward_input.data loss_fn_config = body.forward_input.loss_fn_config or {} - output, loss = self.model.tinker_forward_only( - inputs=datum_list, adapter_name=model_adapter, **loss_fn_config) + output, loss = await self.call_backend( + self.model.tinker_forward_only, inputs=datum_list, adapter_name=model_adapter, **loss_fn_config) return types.ForwardBackwardOutput( loss_fn_output_type='CrossEntropyLossReturn', loss_fn_outputs=output, @@ -156,10 +151,7 @@ async def _do_forward(): ) except Exception: logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise datum_list = body.forward_input.data input_tokens = sum(len(d.model_input.to_ints()) for d in datum_list) @@ -190,8 +182,12 @@ async def _do_forward_backward(): datum_list = body.forward_backward_input.data loss_fn = body.forward_backward_input.loss_fn loss_fn_config = body.forward_backward_input.loss_fn_config or {} - output, loss = self.model.tinker_forward_backward( - inputs=datum_list, adapter_name=model_adapter, loss_fn=loss_fn, **loss_fn_config) + output, loss = await self.call_backend( + self.model.tinker_forward_backward, + inputs=datum_list, + adapter_name=model_adapter, + loss_fn=loss_fn, + **loss_fn_config) output_type = ('ImportanceSamplingLossReturn' if loss_fn == 'importance_sampling' else 'CrossEntropyLossReturn') self.set_resource_state(adapter_name, 'grad_ready', True) @@ -202,10 +198,7 @@ async def _do_forward_backward(): ) except Exception: logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise datum_list = body.forward_backward_input.data input_tokens = sum(len(d.model_input.to_ints()) for d in datum_list) @@ -238,16 +231,15 @@ async def _do_optim(): if not self.get_resource_state(adapter_name, 'grad_ready', False): raise RuntimeError(f'No accumulated gradients for adapter={adapter_name}; ' 'call forward_backward before optim_step') - self.model.tinker_step(adam_params=body.adam_params, adapter_name=model_adapter) + await self.call_backend( + self.model.tinker_step, adam_params=body.adam_params, adapter_name=model_adapter) self.set_resource_state(adapter_name, 'grad_ready', False) - metrics = self.model.tinker_calculate_metric(is_training=True, adapter_name=model_adapter) + metrics = await self.call_backend( + self.model.tinker_calculate_metric, is_training=True, adapter_name=model_adapter) return types.OptimStepResponse(metrics=metrics) except Exception: logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise return await self.schedule_task(_do_optim, model_id=body.model_id, token=token, task_type='optim_step') @@ -267,16 +259,17 @@ async def _do_save(): checkpoint_manager = create_checkpoint_manager(token, client_type='tinker') checkpoint_name = checkpoint_manager.get_ckpt_name(body.path) save_dir = checkpoint_manager.get_save_dir(model_id=body.model_id, is_sampler=False) - self.model.save( - name=checkpoint_name, output_dir=save_dir, adapter_name=model_adapter, save_optimizer=True) + await self.call_backend( + self.model.save, + name=checkpoint_name, + output_dir=save_dir, + adapter_name=model_adapter, + save_optimizer=True) tinker_path = checkpoint_manager.save(body.model_id, name=checkpoint_name, is_sampler=False) return types.SaveWeightsResponse(path=tinker_path, type='save_weights') except Exception: logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise return await self.schedule_task(_do_save, model_id=body.model_id, token=token, task_type='save_weights') @@ -298,7 +291,8 @@ async def _do_save_for_sampler(): # Must save the checkpoint in the twinkle format before calling model.save() tinker_path = checkpoint_manager.save(body.model_id, name=checkpoint_name, is_sampler=True) logger.info(f'Saving weights to {save_dir}') - self.model.save( + await self.call_backend( + self.model.save, name='latest', output_dir=save_dir, adapter_name=self.resolve_model_adapter_name(adapter_name), @@ -320,10 +314,7 @@ async def _do_save_for_sampler(): path=tinker_path, sampling_session_id=sampling_session_id) except Exception: logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise return await self.schedule_task( _do_save_for_sampler, model_id=body.model_id, token=token, task_type='save_weights_for_sampler') @@ -341,18 +332,21 @@ async def _do_load(): assert self.model is not None, 'Model not loaded, please load model first' adapter_name = self.get_adapter_name(adapter_name=body.model_id) self.assert_resource_exists(adapter_name) - self.model.tinker_load( - checkpoint_dir=body.path, + # Path resolution (token isolation, twinkle-vs-external path shapes) is a + # server storage policy and belongs in the handler, not in the GPU-side + # backend actor. The backend receives an already-resolved location. + checkpoint_manager = create_checkpoint_manager(token, client_type='tinker') + resolved = checkpoint_manager.resolve_load_path(body.path) + await self.call_backend( + self.model.tinker_load, + checkpoint_name=resolved.checkpoint_name, + output_dir=resolved.checkpoint_dir if resolved.is_twinkle_path else None, load_optimizer=body.optimizer, - adapter_name=self.resolve_model_adapter_name(adapter_name), - token=token) + adapter_name=self.resolve_model_adapter_name(adapter_name)) self.set_resource_state(adapter_name, 'grad_ready', False) return types.LoadWeightsResponse(path=body.path, type='load_weights') except Exception: logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise return await self.schedule_task(_do_load, model_id=body.model_id, token=token, task_type='load_weights') diff --git a/src/twinkle/server/model/twinkle_handlers.py b/src/twinkle/server/model/twinkle_handlers.py index 582aa5c00..a63a5bfbe 100644 --- a/src/twinkle/server/model/twinkle_handlers.py +++ b/src/twinkle/server/model/twinkle_handlers.py @@ -1,64 +1,58 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -""" -Twinkle-native model handler mixin. - -All endpoints are prefixed /twinkle/... and use schedule_task_and_wait() returning -results directly (synchronous from the client's perspective). -self_fn is injected via FastAPI Depends to obtain the ModelManagement instance at request time. +"""Twinkle-native routes for the Model deployment. + +Registered by ``_register_model_twinkle_routes(app, self_fn)`` -- module-level route +registration closing over ``self_fn`` via ``Depends``, not a mixin: there is no +inheritance relationship with the deployment class. All queued endpoints are prefixed +/twinkle/... and return a Task_Envelope via the shared ``run_submit`` judgment sequence: +the handler submits work and returns immediately, and the client's Client_Future_Layer +resolves the envelope to a terminal state. ``self_fn`` is injected via FastAPI Depends to +obtain the ModelManagement instance at request time. """ from __future__ import annotations -import asyncio import torch -import traceback from collections.abc import Callable from fastapi import Depends, FastAPI, HTTPException, Request from pathlib import Path -from peft import LoraConfig -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING if TYPE_CHECKING: from .app import ModelManagement -import twinkle_client.types as types -from twinkle.data_format import InputFeature, Trajectory +import twinkle.protocol.types as types +from twinkle.protocol.serialize import deserialize_object from twinkle.server.checkpoint import (_resolve_client_save_dir, create_checkpoint_manager, create_training_run_manager, validate_user_path) -from twinkle.server.exceptions import FullModeBusyError -from twinkle.server.model.utils import (data_plane_request_shape, merge_forward_kwargs, resolve_data_plane_model_inputs, - select_output_rows) -from twinkle.server.utils.validation import get_session_id_from_request +from twinkle.server.exceptions import RequestRejectedError, TrainModeMismatchError +from twinkle.server.lifecycle.submit import (backend_kwargs, input_metrics, resolve_twinkle_adapter_name, run_submit, + to_backend_inputs) +from twinkle.server.middleware.auth import get_session_id_from_request +from twinkle.server.model.data_plane_inputs import (data_plane_request_shape, merge_forward_kwargs, + resolve_data_plane_model_inputs, select_output_rows) +from twinkle.server.validation import BackendCapability from twinkle.utils.logger import get_logger -from twinkle_client.common.serialize import deserialize_object logger = get_logger() -def _parse_inputs(inputs: Any): - """Convert raw dict/list inputs to InputFeature or Trajectory objects.""" - if isinstance(inputs, list) and inputs: - first = inputs[0] - if isinstance(first, dict) and 'input_ids' in first: - return [InputFeature(**item) for item in inputs] - else: - return [Trajectory(**item) for item in inputs] - elif isinstance(inputs, dict): - if 'input_ids' in inputs: - return [InputFeature(**inputs)] - else: - return [Trajectory(**inputs)] - return inputs +def _dp_metrics(self, body): + """Scheduling metrics for inline data-parallel endpoints (forward / forward_backward).""" + return input_metrics(self, body, data_parallel=True) + +def _tokens_only_metrics(self, body): + """Scheduling metrics for inline non-data-parallel endpoints (forward_only).""" + return input_metrics(self, body, data_parallel=False) -def _get_twinkle_adapter_name(request: Request, adapter_name: str | None) -> str | None: - """Build a stable per-session adapter name, falling back to request_id for older clients.""" - if adapter_name is None or adapter_name == '': - return None - owner_id = get_session_id_from_request(request) or request.state.request_id - return owner_id + '-' + adapter_name +def _data_plane_metrics(self, body): + """Scheduling metrics derived from DataRef shape for *_from_data_plane endpoints.""" + input_tokens, batch_size = data_plane_request_shape(body) + return {'input_tokens': input_tokens, 'batch_size': batch_size, 'data_world_size': self.data_world_size} -def _register_twinkle_routes(app: FastAPI, self_fn: Callable[[], ModelManagement]) -> None: + +def _register_model_twinkle_routes(app: FastAPI, self_fn: Callable[[], ModelManagement]) -> None: """Register all /twinkle/* routes on the given FastAPI app. self_fn is a zero-argument callable that returns the current ModelManagement @@ -71,437 +65,400 @@ async def model_healthz( self: ModelManagement = Depends(self_fn), ) -> dict: """Deep health probe: pings underlying model actors to verify liveness.""" - result = self.check_model_health() - if not result['healthy']: + result = await self.check_model_health() + if self._model_unhealthy or not result['healthy']: from fastapi.responses import JSONResponse return JSONResponse(status_code=503, content=result) return result - async def run_task(coro): - """Await a schedule_task_and_wait coroutine and surface any exception as a - structured HTTP 500 response so the client receives the full traceback instead - of an opaque connection-level error. - - Note: HTTPException is re-raised directly to preserve its status code and detail. - """ - try: - return await coro - except HTTPException: - raise # Re-raise HTTPException directly to preserve status code - except Exception: - logger.error(traceback.format_exc()) - raise HTTPException(status_code=500, detail=traceback.format_exc()) - @app.post('/twinkle/create', response_model=types.CreateResponse) async def create(request: Request, body: types.CreateRequest, self: ModelManagement = Depends(self_fn)) -> types.CreateResponse: await self._on_request_start(request) return types.CreateResponse() - @app.post('/twinkle/forward', response_model=types.ForwardResponse) + # ------------------------------------------------------------------ # + # Inline data / forward family + # ------------------------------------------------------------------ # + + @app.post('/twinkle/forward', response_model=types.TaskEnvelope) async def forward(request: Request, body: types.ForwardRequest, - self: ModelManagement = Depends(self_fn)) -> types.ForwardResponse: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - inputs = _parse_inputs(body.inputs) - ret = self.model.forward( - inputs=inputs, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + async def _call(self, body, adapter_name, token): + ret = await self.call_backend( + self.model.forward, + inputs=to_backend_inputs(body.inputs), + adapter_name=self.resolve_model_adapter_name(adapter_name), + **backend_kwargs(body)) return {'result': ret} - inputs_list = body.inputs if isinstance(body.inputs, list) else [body.inputs] - input_tokens = sum(len(inp.get('input_ids', [])) if isinstance(inp, dict) else 0 for inp in inputs_list) - batch_size = len(inputs_list) - return await run_task( - self.schedule_task_and_wait( - _task, - model_id=adapter_name, - token=token, - input_tokens=input_tokens, - batch_size=batch_size, - data_world_size=self.data_world_size, - task_type='forward', - )) + return await run_submit( + self, + request, + body, + task_type='forward', + backend_call=_call, + metrics=_dp_metrics, + capability=BackendCapability.Forward) - @app.post('/twinkle/forward_from_data_plane', response_model=types.ForwardResponse) - async def forward_from_data_plane( - request: Request, - body: types.DataPlaneForwardRequest, - self: ModelManagement = Depends(self_fn), - ) -> types.ForwardResponse: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + @app.post('/twinkle/forward_only', response_model=types.TaskEnvelope) + async def forward_only( + request: Request, body: types.ForwardOnlyRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - raw_inputs, field_kwargs = await resolve_data_plane_model_inputs(body, self.data_plane) - kwargs = merge_forward_kwargs(body.model_extra or {}, field_kwargs) - ret = self.model.forward( - inputs=_parse_inputs(raw_inputs), - adapter_name=adapter_name, - **kwargs, - ) + async def _call(self, body, adapter_name, token): + ret = await self.call_backend( + self.model.forward_only, + inputs=to_backend_inputs(body.inputs), + adapter_name=self.resolve_model_adapter_name(adapter_name), + **backend_kwargs(body)) return {'result': ret} - input_tokens, batch_size = data_plane_request_shape(body) - return await run_task( - self.schedule_task_and_wait( - _task, - model_id=adapter_name, - token=token, - input_tokens=input_tokens, - batch_size=batch_size, - data_world_size=self.data_world_size, - task_type='forward_from_data_plane', - )) + return await run_submit( + self, request, body, task_type='forward_only', backend_call=_call, metrics=_tokens_only_metrics) - @app.post('/twinkle/remove_adapter') - async def remove_adapter( - request: Request, - body: types.AdapterRequest, - self: ModelManagement = Depends(self_fn), - ) -> dict[str, str]: - """Release a drained tenant's in-memory training adapter.""" - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + @app.post('/twinkle/forward_backward', response_model=types.TaskEnvelope) + async def forward_backward( + request: Request, body: types.ForwardBackwardTaskRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - await self._cleanup_adapter(adapter_name) - return {'status': 'ok'} + async def _call(self, body, adapter_name, token): - return await run_task( - self.schedule_task_and_wait( - _task, - model_id=adapter_name, - token=token, - task_type='remove_adapter', - )) + def first_element(data): + while isinstance(data, list): + if len(data) == 0: + return None + data = data[0] + return data - @app.post('/twinkle/forward_only', response_model=types.ForwardResponse) - async def forward_only( - request: Request, - body: types.ForwardOnlyRequest, - self: ModelManagement = Depends(self_fn), - ) -> types.ForwardResponse: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) - - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - inputs = _parse_inputs(body.inputs) - ret = self.model.forward_only( - inputs=inputs, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + all_inputs = to_backend_inputs(body.inputs) + for inputs in all_inputs: + for key in inputs: + if isinstance(inputs[key], list) and isinstance(first_element(inputs[key]), (int, float)): + inputs[key] = torch.tensor(inputs[key]) + ret = await self.call_backend( + self.model.forward_backward, + inputs=all_inputs, + adapter_name=self.resolve_model_adapter_name(adapter_name), + **backend_kwargs(body)) return {'result': ret} - inputs_list = body.inputs if isinstance(body.inputs, list) else [body.inputs] - input_tokens = sum(len(inp.get('input_ids', [])) if isinstance(inp, dict) else 0 for inp in inputs_list) - return await run_task( - self.schedule_task_and_wait( - _task, - model_id=adapter_name, - token=token, - input_tokens=input_tokens, - task_type='forward_only', - )) + return await run_submit( + self, request, body, task_type='forward_backward', backend_call=_call, metrics=_dp_metrics) - @app.post('/twinkle/forward_only_from_data_plane', response_model=types.ForwardResponse) - async def forward_only_from_data_plane( - request: Request, - body: types.DataPlaneForwardOnlyRequest, - self: ModelManagement = Depends(self_fn), - ) -> types.ForwardResponse: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + @app.post('/twinkle/calculate_loss', response_model=types.TaskEnvelope) + async def calculate_loss( + request: Request, body: types.AdapterRequest, self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - raw_inputs, field_kwargs = await resolve_data_plane_model_inputs(body, self.data_plane) - inputs = _parse_inputs(raw_inputs) - kwargs = merge_forward_kwargs(body.model_extra or {}, field_kwargs) - ret = self.model.forward_only(inputs=inputs, adapter_name=adapter_name, **kwargs) - if body.output_ref is not None: - rows = select_output_rows( - ret, - batch_size=len(inputs), - output_fields=body.output_fields, - ) - output_ref = await self.data_plane.append(body.output_ref, rows) - return {'result': output_ref.model_dump()} + async def _call(self, body, adapter_name, token): + ret = await self.call_backend( + self.model.calculate_loss, + adapter_name=self.resolve_model_adapter_name(adapter_name), + **backend_kwargs(body)) return {'result': ret} - input_tokens, batch_size = data_plane_request_shape(body) - return await run_task( - self.schedule_task_and_wait( - _task, - model_id=adapter_name, - token=token, - input_tokens=input_tokens, - batch_size=batch_size, - data_world_size=self.data_world_size, - task_type='forward_only_from_data_plane', - )) + return await run_submit( + self, + request, + body, + task_type='calculate_loss', + backend_call=_call, + capability=BackendCapability.CalculateLoss) - @app.post('/twinkle/calculate_loss', response_model=types.CalculateLossResponse) - async def calculate_loss( - request: Request, - body: types.AdapterRequest, - self: ModelManagement = Depends(self_fn), - ) -> types.CalculateLossResponse: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + @app.post('/twinkle/backward', response_model=types.TaskEnvelope) + async def backward(request: Request, body: types.AdapterRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - ret = self.model.calculate_loss(adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) - return {'result': ret} + async def _call(self, body, adapter_name, token): + await self.call_backend( + self.model.backward, adapter_name=self.resolve_model_adapter_name(adapter_name), **backend_kwargs(body)) - return await run_task( - self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='calculate_loss')) + return await run_submit( + self, request, body, task_type='backward', backend_call=_call, capability=BackendCapability.Backward) - @app.post('/twinkle/backward') - async def backward(request: Request, body: types.AdapterRequest, self: ModelManagement = Depends(self_fn)) -> None: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + # ------------------------------------------------------------------ # + # Data-plane forward family (DataRef inputs; response only enters the contract) + # ------------------------------------------------------------------ # - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - self.model.backward(adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + @app.post('/twinkle/forward_from_data_plane', response_model=types.TaskEnvelope) + async def forward_from_data_plane( + request: Request, body: types.DataPlaneForwardRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='backward')) + async def _call(self, body, adapter_name, token): + raw_inputs, field_kwargs = await resolve_data_plane_model_inputs(body, self.data_plane) + kwargs = merge_forward_kwargs(backend_kwargs(body), field_kwargs) + ret = await self.call_backend( + self.model.forward, + inputs=to_backend_inputs(raw_inputs), + adapter_name=self.resolve_model_adapter_name(adapter_name), + **kwargs) + return {'result': ret} - @app.post('/twinkle/forward_backward', response_model=types.ForwardBackwardResponse) - async def forward_backward( - request: Request, - body: types.ForwardRequest, - self: ModelManagement = Depends(self_fn), - ) -> types.ForwardBackwardResponse: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + return await run_submit( + self, + request, + body, + task_type='forward_from_data_plane', + backend_call=_call, + metrics=_data_plane_metrics, + capability=BackendCapability.Forward) - def first_element(data): - while isinstance(data, list): - if len(data) == 0: - return None - data = data[0] - return data + @app.post('/twinkle/forward_only_from_data_plane', response_model=types.TaskEnvelope) + async def forward_only_from_data_plane( + request: Request, body: types.DataPlaneForwardOnlyRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - all_inputs = _parse_inputs(body.inputs) - for inputs in all_inputs: - for key in inputs: - if isinstance(inputs[key], list) and isinstance(first_element(inputs[key]), (int, float)): - inputs[key] = torch.tensor(inputs[key]) - ret = self.model.forward_backward( - inputs=all_inputs, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + async def _call(self, body, adapter_name, token): + raw_inputs, field_kwargs = await resolve_data_plane_model_inputs(body, self.data_plane) + inputs = to_backend_inputs(raw_inputs) + kwargs = merge_forward_kwargs(backend_kwargs(body), field_kwargs) + ret = await self.call_backend( + self.model.forward_only, + inputs=inputs, + adapter_name=self.resolve_model_adapter_name(adapter_name), + **kwargs) + if body.output_ref is not None: + rows = select_output_rows(ret, batch_size=len(inputs), output_fields=body.output_fields) + output_ref = await self.data_plane.append(body.output_ref, rows) + return {'result': output_ref.model_dump()} return {'result': ret} - inputs_list = body.inputs if isinstance(body.inputs, list) else [body.inputs] - input_tokens = sum(len(inp.get('input_ids', [])) if isinstance(inp, dict) else 0 for inp in inputs_list) - batch_size = len(inputs_list) - return await run_task( - self.schedule_task_and_wait( - _task, - model_id=adapter_name, - token=token, - input_tokens=input_tokens, - batch_size=batch_size, - data_world_size=self.data_world_size, - task_type='forward_backward', - )) + return await run_submit( + self, + request, + body, + task_type='forward_only_from_data_plane', + backend_call=_call, + metrics=_data_plane_metrics) - @app.post('/twinkle/forward_backward_from_data_plane', response_model=types.ForwardBackwardResponse) + @app.post('/twinkle/forward_backward_from_data_plane', response_model=types.TaskEnvelope) async def forward_backward_from_data_plane( - request: Request, - body: types.DataPlaneForwardRequest, - self: ModelManagement = Depends(self_fn), - ) -> types.ForwardBackwardResponse: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + request: Request, body: types.DataPlaneForwardRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) + async def _call(self, body, adapter_name, token): raw_inputs, field_kwargs = await resolve_data_plane_model_inputs(body, self.data_plane) - kwargs = merge_forward_kwargs(body.model_extra or {}, field_kwargs) - ret = self.model.forward_backward( - inputs=_parse_inputs(raw_inputs), - adapter_name=adapter_name, - **kwargs, - ) + kwargs = merge_forward_kwargs(backend_kwargs(body), field_kwargs) + ret = await self.call_backend( + self.model.forward_backward, + inputs=to_backend_inputs(raw_inputs), + adapter_name=self.resolve_model_adapter_name(adapter_name), + **kwargs) return {'result': ret} - input_tokens, batch_size = data_plane_request_shape(body) - return await run_task( - self.schedule_task_and_wait( - _task, - model_id=adapter_name, - token=token, - input_tokens=input_tokens, - batch_size=batch_size, - data_world_size=self.data_world_size, - task_type='forward_backward_from_data_plane', - )) + return await run_submit( + self, + request, + body, + task_type='forward_backward_from_data_plane', + backend_call=_call, + metrics=_data_plane_metrics) - @app.post('/twinkle/clip_grad_norm', response_model=types.ClipGradNormResponse) + # ------------------------------------------------------------------ # + # Optimizer / control plane + # ------------------------------------------------------------------ # + + @app.post('/twinkle/clip_grad_norm', response_model=types.TaskEnvelope) async def clip_grad_norm( - request: Request, - body: types.AdapterRequest, - self: ModelManagement = Depends(self_fn), - ) -> types.ClipGradNormResponse: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + request: Request, body: types.ClipGradNormRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - ret = self.model.clip_grad_norm(adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + async def _call(self, body, adapter_name, token): + ret = await self.call_backend( + self.model.clip_grad_norm, + max_grad_norm=body.max_grad_norm, + norm_type=body.norm_type, + adapter_name=self.resolve_model_adapter_name(adapter_name), + **backend_kwargs(body)) return {'result': str(ret)} - return await run_task( - self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='clip_grad_norm')) + return await run_submit(self, request, body, task_type='clip_grad_norm', backend_call=_call) - @app.post('/twinkle/step') - async def step(request: Request, body: types.AdapterRequest, self: ModelManagement = Depends(self_fn)) -> None: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + @app.post('/twinkle/step', response_model=types.TaskEnvelope) + async def step(request: Request, body: types.StepRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - self.model.step(adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + async def _call(self, body, adapter_name, token): + await self.call_backend( + self.model.step, adapter_name=self.resolve_model_adapter_name(adapter_name), **backend_kwargs(body)) - await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='step')) + return await run_submit(self, request, body, task_type='step', backend_call=_call) - @app.post('/twinkle/zero_grad') - async def zero_grad(request: Request, body: types.AdapterRequest, self: ModelManagement = Depends(self_fn)) -> None: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + @app.post('/twinkle/zero_grad', response_model=types.TaskEnvelope) + async def zero_grad(request: Request, body: types.AdapterRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - self.model.zero_grad(adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + async def _call(self, body, adapter_name, token): + await self.call_backend( + self.model.zero_grad, + adapter_name=self.resolve_model_adapter_name(adapter_name), + **backend_kwargs(body)) - await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='zero_grad')) + return await run_submit(self, request, body, task_type='zero_grad', backend_call=_call) - @app.post('/twinkle/lr_step') - async def lr_step(request: Request, body: types.AdapterRequest, self: ModelManagement = Depends(self_fn)) -> None: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + @app.post('/twinkle/lr_step', response_model=types.TaskEnvelope) + async def lr_step(request: Request, body: types.LrStepRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - self.model.lr_step(adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + async def _call(self, body, adapter_name, token): + await self.call_backend( + self.model.lr_step, adapter_name=self.resolve_model_adapter_name(adapter_name), **backend_kwargs(body)) - await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='lr_step')) + return await run_submit(self, request, body, task_type='lr_step', backend_call=_call) - @app.post('/twinkle/clip_grad_and_step') + @app.post('/twinkle/clip_grad_and_step', response_model=types.TaskEnvelope) async def clip_grad_and_step( - request: Request, - body: types.ClipGradAndStepRequest, - self: ModelManagement = Depends(self_fn), - ) -> None: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + request: Request, body: types.ClipGradAndStepRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - self.model.clip_grad_and_step( + async def _call(self, body, adapter_name, token): + await self.call_backend( + self.model.clip_grad_and_step, max_grad_norm=body.max_grad_norm, norm_type=body.norm_type, adapter_name=self.resolve_model_adapter_name(adapter_name), - **extra_kwargs, - ) + **backend_kwargs(body)) - await run_task( - self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='clip_grad_and_step')) + return await run_submit(self, request, body, task_type='clip_grad_and_step', backend_call=_call) - @app.post('/twinkle/get_train_configs', response_model=types.GetTrainConfigsResponse) + @app.post('/twinkle/get_train_configs', response_model=types.TaskEnvelope) async def get_train_configs( - request: Request, - body: types.AdapterRequest, - self: ModelManagement = Depends(self_fn), - ) -> types.GetTrainConfigsResponse: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + request: Request, body: types.AdapterRequest, self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - ret = self.model.get_train_configs( - adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + async def _call(self, body, adapter_name, token): + ret = await self.call_backend( + self.model.get_train_configs, + adapter_name=self.resolve_model_adapter_name(adapter_name), + **backend_kwargs(body)) return {'result': ret} - return await run_task( - self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='get_train_configs')) + return await run_submit(self, request, body, task_type='get_train_configs', backend_call=_call) - @app.post('/twinkle/set_loss') - async def set_loss(request: Request, body: types.SetLossRequest, self: ModelManagement = Depends(self_fn)) -> None: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + @app.post('/twinkle/set_loss', response_model=types.TaskEnvelope) + async def set_loss(request: Request, body: types.SetLossRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - self.model.set_loss( - body.loss_cls, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + async def _call(self, body, adapter_name, token): + await self.call_backend( + self.model.set_loss, + body.loss_cls, + adapter_name=self.resolve_model_adapter_name(adapter_name), + **backend_kwargs(body)) - await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='set_loss')) + return await run_submit(self, request, body, task_type='set_loss', backend_call=_call) - @app.post('/twinkle/set_optimizer') + @app.post('/twinkle/set_optimizer', response_model=types.TaskEnvelope) async def set_optimizer( - request: Request, - body: types.SetOptimizerRequest, - self: ModelManagement = Depends(self_fn), - ) -> None: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + request: Request, body: types.SetOptimizerRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - self.model.set_optimizer( - body.optimizer_cls, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + async def _call(self, body, adapter_name, token): + await self.call_backend( + self.model.set_optimizer, + body.optimizer_cls, + adapter_name=self.resolve_model_adapter_name(adapter_name), + **backend_kwargs(body)) - await run_task( - self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='set_optimizer')) + return await run_submit(self, request, body, task_type='set_optimizer', backend_call=_call) - @app.post('/twinkle/set_lr_scheduler') + @app.post('/twinkle/set_lr_scheduler', response_model=types.TaskEnvelope) async def set_lr_scheduler( - request: Request, - body: types.SetLrSchedulerRequest, - self: ModelManagement = Depends(self_fn), - ) -> None: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + request: Request, body: types.SetLrSchedulerRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - self.model.set_lr_scheduler( - body.scheduler_cls, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + async def _call(self, body, adapter_name, token): + await self.call_backend( + self.model.set_lr_scheduler, + body.scheduler_cls, + adapter_name=self.resolve_model_adapter_name(adapter_name), + **backend_kwargs(body)) + + return await run_submit(self, request, body, task_type='set_lr_scheduler', backend_call=_call) + + @app.post('/twinkle/set_template', response_model=types.TaskEnvelope) + async def set_template( + request: Request, body: types.SetTemplateRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: + + async def _call(self, body, adapter_name, token): + await self.call_backend( + self.model.set_template, + body.template_cls, + adapter_name=self.resolve_model_adapter_name(adapter_name), + **backend_kwargs(body)) + + return await run_submit(self, request, body, task_type='set_template', backend_call=_call) + + @app.post('/twinkle/set_processor', response_model=types.TaskEnvelope) + async def set_processor( + request: Request, body: types.SetProcessorRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: + + async def _call(self, body, adapter_name, token): + await self.call_backend( + self.model.set_processor, + body.processor_cls, + adapter_name=self.resolve_model_adapter_name(adapter_name), + **backend_kwargs(body)) + + return await run_submit(self, request, body, task_type='set_processor', backend_call=_call) - await run_task( - self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='set_lr_scheduler')) + @app.post('/twinkle/add_metric', response_model=types.TaskEnvelope) + async def add_metric(request: Request, body: types.AddMetricRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - @app.post('/twinkle/save', response_model=types.SaveResponse) + async def _call(self, body, adapter_name, token): + metric_cls = deserialize_object(body.metric_cls) + await self.call_backend( + self.model.add_metric, + metric_cls, + is_training=body.is_training, + adapter_name=self.resolve_model_adapter_name(adapter_name), + **backend_kwargs(body)) + + return await run_submit(self, request, body, task_type='add_metric', backend_call=_call) + + @app.post('/twinkle/apply_patch', response_model=types.TaskEnvelope) + async def apply_patch( + request: Request, body: types.ApplyPatchRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: + + async def _call(self, body, adapter_name, token): + patch_cls = deserialize_object(body.patch_cls) + await self.call_backend( + self.model.apply_patch, + patch_cls, + adapter_name=self.resolve_model_adapter_name(adapter_name), + **backend_kwargs(body)) + + return await run_submit(self, request, body, task_type='apply_patch', backend_call=_call) + + @app.post('/twinkle/calculate_metric', response_model=types.TaskEnvelope) + async def calculate_metric( + request: Request, body: types.CalculateMetricRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: + + async def _call(self, body, adapter_name, token): + ret = await self.call_backend( + self.model.calculate_metric, + is_training=body.is_training, + adapter_name=self.resolve_model_adapter_name(adapter_name), + **backend_kwargs(body)) + return {'result': ret} + + return await run_submit(self, request, body, task_type='calculate_metric', backend_call=_call) + + # ------------------------------------------------------------------ # + # Checkpoint I/O (need the caller token) + # ------------------------------------------------------------------ # + + @app.post('/twinkle/save', response_model=types.TaskEnvelope) async def save(request: Request, body: types.SaveRequest, - self: ModelManagement = Depends(self_fn)) -> types.SaveResponse: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} + async def _call(self, body, adapter_name, token): checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') checkpoint_name = checkpoint_manager.get_ckpt_name(body.name) save_dir = checkpoint_manager.get_save_dir(model_id=adapter_name, is_sampler=body.is_sampler) @@ -510,157 +467,114 @@ async def _task(): model_id=adapter_name, name=checkpoint_name, is_sampler=body.is_sampler) # For sampler weights the actual data is always written to 'latest/'. model_save_name = 'latest' if body.is_sampler else checkpoint_name - checkpoint_dir = self.model.save( + checkpoint_dir = await self.call_backend( + self.model.save, name=model_save_name, output_dir=save_dir, adapter_name=self.resolve_model_adapter_name(adapter_name), save_optimizer=body.save_optimizer, - **extra_kwargs) + **backend_kwargs(body)) return {'twinkle_path': twinkle_path, 'checkpoint_dir': checkpoint_dir} - return await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='save')) + return await run_submit(self, request, body, task_type='save', backend_call=_call) - @app.post('/twinkle/load') - async def load(request: Request, body: types.LoadRequest, self: ModelManagement = Depends(self_fn)) -> None: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + @app.post('/twinkle/load', response_model=types.TaskEnvelope) + async def load(request: Request, body: types.LoadRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} + async def _call(self, body, adapter_name, token): checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') resolved = checkpoint_manager.resolve_load_path(body.name) - self.model.load( + await self.call_backend( + self.model.load, name=resolved.checkpoint_name, output_dir=resolved.checkpoint_dir, adapter_name=self.resolve_model_adapter_name(adapter_name), load_optimizer=body.load_optimizer, token=token, - **extra_kwargs) + **backend_kwargs(body)) - await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='load')) + return await run_submit(self, request, body, task_type='load', backend_call=_call) - @app.post('/twinkle/resume_from_checkpoint', response_model=types.TrainingProgressResponse) + @app.post('/twinkle/resume_from_checkpoint', response_model=types.TaskEnvelope) async def resume_from_checkpoint( - request: Request, - body: types.ResumeFromCheckpointRequest, - self: ModelManagement = Depends(self_fn), - ) -> types.TrainingProgressResponse: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + request: Request, body: types.ResumeFromCheckpointRequest, + self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: - async def _task(): - self.assert_resource_exists(adapter_name) + async def _call(self, body, adapter_name, token): checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') resolved = checkpoint_manager.resolve_load_path(body.name) checkpoint_dir = ( Path(resolved.checkpoint_dir, resolved.checkpoint_name).as_posix() if resolved.checkpoint_dir else body.name) - ret = self.model.resume_from_checkpoint( + ret = await self.call_backend( + self.model.resume_from_checkpoint, checkpoint_dir, resume_only_model=body.resume_only_model, - adapter_name=self.resolve_model_adapter_name(adapter_name), - ) + adapter_name=self.resolve_model_adapter_name(adapter_name)) return {'result': ret} - return await run_task(self.schedule_task_and_wait(_task, task_type='resume')) + return await run_submit(self, request, body, task_type='resume', backend_call=_call) - @app.post('/twinkle/upload_to_hub', response_model=types.UploadToHubResponse) - async def upload_to_hub( - request: Request, - body: types.UploadToHubRequest, - self: ModelManagement = Depends(self_fn), - ) -> types.UploadToHubResponse: - token = await self._on_request_start(request) + # ------------------------------------------------------------------ # + # Adapter lifecycle (create / drop the adapter itself: no resource assert) + # ------------------------------------------------------------------ # - async def _task(): - if body.checkpoint_dir.startswith('twinkle://'): - checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') - parsed = checkpoint_manager.parse_twinkle_path(body.checkpoint_dir) - if not parsed: - raise ValueError(f'Invalid twinkle path format: {body.checkpoint_dir}') - checkpoint_id = parsed.checkpoint_id - model_id_to_load = parsed.training_run_id - checkpoint = checkpoint_manager.get(model_id_to_load, checkpoint_id) - if not checkpoint: - raise ValueError(f'Checkpoint not found or access denied: {body.checkpoint_dir}') - checkpoint_dir = str( - checkpoint_manager.get_ckpt_dir(model_id=model_id_to_load, checkpoint_id=checkpoint_id)) - else: - checkpoint_dir = body.checkpoint_dir - # Run blocking upload in thread pool so the event loop is not blocked. - # async_upload is intentionally ignored here: the task queue + client polling - # already provide the fire-and-forget / wait semantics without holding the - # HTTP connection open for the full duration of the upload. - await asyncio.to_thread( - self.model.upload_to_hub, - checkpoint_dir=checkpoint_dir, - hub_model_id=body.hub_model_id, - hub_token=body.hub_token or token, - async_upload=False, - ) + @app.post('/twinkle/remove_adapter', response_model=types.TaskEnvelope) + async def remove_adapter( + request: Request, body: types.AdapterRequest, self: ModelManagement = Depends(self_fn)) -> types.TaskEnvelope: + """Release a drained tenant's in-memory training adapter.""" + + async def _call(self, body, adapter_name, token): + await self._cleanup_adapter(adapter_name) + return {'status': 'ok'} - future_ref = await self.schedule_background_task(_task, task_type='upload_to_hub') - request_id = future_ref.get('request_id') - if request_id is None: - raise HTTPException(status_code=500, detail=f'Upload task scheduling failed: {future_ref}') - return types.UploadToHubResponse(request_id=request_id) + return await run_submit( + self, request, body, task_type='remove_adapter', backend_call=_call, assert_resource=False) - @app.get('/twinkle/upload_status/{request_id}', response_model=types.UploadStatusResponse) - async def upload_status( - request: Request, - request_id: str, - self: ModelManagement = Depends(self_fn), - ) -> types.UploadStatusResponse: - await self._on_request_start(request) - record = await self.state.get_future(request_id) - if record is None: - raise HTTPException(status_code=404, detail=f'Upload task not found: {request_id}') - status = record.get('status', 'unknown') - error = None - if status == 'failed': - error = record.get('result', {}).get('error', 'Unknown error') - return types.UploadStatusResponse(request_id=request_id, status=status, error=error) - - @app.post('/twinkle/add_adapter_to_model', response_model=types.AddAdapterResponse) + @app.post('/twinkle/add_adapter_to_model', response_model=types.TaskEnvelope) async def add_adapter_to_model( request: Request, body: types.AddAdapterRequest, self: ModelManagement = Depends(self_fn), - ) -> types.AddAdapterResponse: - assert body.adapter_name, 'You need to specify a valid `adapter_name`' + ) -> types.TaskEnvelope: + # This endpoint creates the adapter, so it cannot use the standard resource + # assertion. The Decision_Boundary left checks (train_mode 400 / full-mode + # 409) run here, before any state write, raising RequestRejectedError + # subclasses (zero future writes). + # + # Raised, not asserted: a missing adapter_name is decidable from the request body + # alone, so it owes the caller a real 400. A bare `assert` would surface as a 500 + # ('the server broke') and would vanish entirely under `python -O`, letting an + # empty adapter_name through to the backend. + if not body.adapter_name: + raise RequestRejectedError('`adapter_name` is required and must be non-empty.') token = await self._on_request_start(request) if not validate_user_path(token, body.adapter_name): - raise HTTPException(status_code=400, detail=f'Invalid adapter_name: {body.adapter_name}') - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) + raise RequestRejectedError(f'Invalid adapter_name: {body.adapter_name}') + adapter_name = resolve_twinkle_adapter_name(request, body.adapter_name) session_id = get_session_id_from_request(request) try: resolved_save_dir = _resolve_client_save_dir(body.save_dir).as_posix() if body.save_dir else None except ValueError as exc: - raise HTTPException(status_code=400, detail=str(exc)) + raise RequestRejectedError(str(exc)) from exc - async def _task(): - config = deserialize_object(body.config) - extra_kwargs = body.model_extra or {} - training_run_manager = create_training_run_manager(token, client_type='twinkle') + config = deserialize_object(body.config) - # Validate the supplied config against the deployment's train_mode. - if self.is_full_mode and config is not None: - raise HTTPException( - status_code=400, - detail='This deployment runs in full-parameter (exclusive) mode; pass config=None ' - '(do not send a LoraConfig).') - if (not self.is_full_mode) and config is None: - raise HTTPException( - status_code=400, detail='This deployment runs in LoRA mode; a LoraConfig is required.') - - # In full mode ensure the exclusive deployment is free before touching state. - if self.is_full_mode: - try: - self.assert_full_mode_available(adapter_name) - except FullModeBusyError as e: - raise HTTPException(status_code=409, detail=str(e)) + # ---- Decision_Boundary left: validate against the deployment's train_mode ---- + if self.is_full_mode and config is not None: + raise TrainModeMismatchError('This deployment runs in full-parameter (exclusive) mode; pass ' + 'config=None (do not send a LoraConfig).') + if (not self.is_full_mode) and config is None: + raise TrainModeMismatchError('This deployment runs in LoRA mode; a LoraConfig is required.') + if self.is_full_mode: + # Raises FullModeBusyError (409) if another tenant holds the exclusive deployment. + self.assert_full_mode_available(adapter_name) + async def _task(): + from peft import LoraConfig + extra_kwargs = backend_kwargs(body) + training_run_manager = create_training_run_manager(token, client_type='twinkle') lora_config = None if isinstance(config, LoraConfig): lora_config = types.LoraConfig(rank=config.r, train_unembed=False, train_mlp=True, train_attn=True) @@ -682,7 +596,7 @@ async def _task(): # No PEFT adapter to add; the default optimizer group is used. self.set_resource_state(adapter_name, 'grad_ready', False) else: - self.model.add_adapter_to_model(adapter_name, config, **extra_kwargs) + await self.call_backend(self.model.add_adapter_to_model, adapter_name, config, **extra_kwargs) except Exception: self.unregister_resource(adapter_name) await self.state.unload_model(adapter_name) @@ -690,118 +604,40 @@ async def _task(): training_run_manager.save(adapter_name, run_config) return {'status': 'ok', 'adapter_name': adapter_name} - return await run_task( - self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='add_adapter_to_model')) - - @app.post('/twinkle/apply_patch') - async def apply_patch( - request: Request, - body: types.ApplyPatchRequest, - self: ModelManagement = Depends(self_fn), - ) -> None: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) - - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - patch_cls = deserialize_object(body.patch_cls) - self.model.apply_patch( - patch_cls, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) + return await self.submit_and_peek(_task, model_id=adapter_name, token=token, task_type='add_adapter_to_model') - await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='apply_patch')) + # ------------------------------------------------------------------ # + # Hub upload (pure I/O -> background task; state-tracked via Retrieve_Endpoint) + # ------------------------------------------------------------------ # - @app.post('/twinkle/add_metric') - async def add_metric( - request: Request, - body: types.AddMetricRequest, - self: ModelManagement = Depends(self_fn), - ) -> None: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) - - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - metric_cls = deserialize_object(body.metric_cls) - self.model.add_metric( - metric_cls, - is_training=body.is_training, - adapter_name=self.resolve_model_adapter_name(adapter_name), - **extra_kwargs) - - await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='add_metric')) - - @app.post('/twinkle/set_template') - async def set_template( - request: Request, - body: types.SetTemplateRequest, - self: ModelManagement = Depends(self_fn), - ) -> None: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) - - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - self.model.set_template( - body.template_cls, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) - - await run_task(self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='set_template')) - - @app.post('/twinkle/set_processor') - async def set_processor( - request: Request, - body: types.SetProcessorRequest, - self: ModelManagement = Depends(self_fn), - ) -> None: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) - - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - self.model.set_processor( - body.processor_cls, adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) - - await run_task( - self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='set_processor')) - - @app.post('/twinkle/calculate_metric', response_model=types.CalculateMetricResponse) - async def calculate_metric( - request: Request, - body: types.CalculateMetricRequest, - self: ModelManagement = Depends(self_fn), - ) -> types.CalculateMetricResponse: - token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) - - async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - ret = self.model.calculate_metric( - is_training=body.is_training, - adapter_name=self.resolve_model_adapter_name(adapter_name), - **extra_kwargs) - return {'result': ret} - - return await run_task( - self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='calculate_metric')) - - @app.post('/twinkle/get_state_dict', response_model=types.GetStateDictResponse) - async def get_state_dict( + @app.post('/twinkle/upload_to_hub', response_model=types.TaskEnvelope) + async def upload_to_hub( request: Request, - body: types.GetStateDictRequest, + body: types.UploadToHubRequest, self: ModelManagement = Depends(self_fn), - ) -> types.GetStateDictResponse: + ) -> types.TaskEnvelope: token = await self._on_request_start(request) - adapter_name = _get_twinkle_adapter_name(request, body.adapter_name) async def _task(): - self.assert_resource_exists(adapter_name) - extra_kwargs = body.model_extra or {} - ret = self.model.get_state_dict(adapter_name=self.resolve_model_adapter_name(adapter_name), **extra_kwargs) - return {'result': ret} + if body.checkpoint_dir.startswith('twinkle://'): + checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') + parsed = checkpoint_manager.parse_twinkle_path(body.checkpoint_dir) + if not parsed: + raise ValueError(f'Invalid twinkle path format: {body.checkpoint_dir}') + checkpoint = checkpoint_manager.get(parsed.training_run_id, parsed.checkpoint_id) + if not checkpoint: + raise ValueError(f'Checkpoint not found or access denied: {body.checkpoint_dir}') + checkpoint_dir = str( + checkpoint_manager.get_ckpt_dir( + model_id=parsed.training_run_id, checkpoint_id=parsed.checkpoint_id)) + else: + checkpoint_dir = body.checkpoint_dir + await self.call_backend( + self.model.upload_to_hub, + checkpoint_dir=checkpoint_dir, + hub_model_id=body.hub_model_id, + hub_token=body.hub_token or token, + async_upload=False, + ) - return await run_task( - self.schedule_task_and_wait(_task, model_id=adapter_name, token=token, task_type='get_state_dict')) + return await self.submit_background_and_peek(_task, task_type='upload_to_hub') diff --git a/src/twinkle/server/processor/app.py b/src/twinkle/server/processor/app.py index 9f6b88966..b471cede8 100644 --- a/src/twinkle/server/processor/app.py +++ b/src/twinkle/server/processor/app.py @@ -18,11 +18,11 @@ from ray import serve from typing import Any -import twinkle -from twinkle import DeviceGroup, DeviceMesh, get_logger +from twinkle import DeviceGroup, get_logger from twinkle.server.deployment import LazyCleanupMixin, bind_deployment, build_deployment_app +from twinkle.server.runtime import init_twinkle_runtime +from twinkle.server.session_resource import ProcessorManagerMixin from twinkle.server.state import ServerState, get_server_state -from twinkle.server.utils.lifecycle import ProcessorManagerMixin from .twinkle_handlers import _register_processor_routes logger = get_logger() @@ -36,7 +36,7 @@ class ProcessorManagement(LazyCleanupMixin, ProcessorManagerMixin): Lifecycle is handled by ProcessorManagerMixin: - Processors are registered with a session ID on creation. - - A background thread expires processors whose session has timed out. + - A background task expires processors whose session has timed out. - Per-user processor limit is enforced at registration. - Sticky session routing ensures session requests hit the same replica. """ @@ -48,16 +48,13 @@ def __init__(self, nproc_per_node: int = 1, processor_config: dict[str, Any] | None = None): self.device_group = DeviceGroup(**device_group) - twinkle.initialize( - mode='ray', + self.device_mesh = init_twinkle_runtime( + is_mock=False, nproc_per_node=nproc_per_node, - groups=[self.device_group], - lazy_collect=False, - ncpu_proc_per_node=ncpu_proc_per_node) - if 'mesh_dim_names' in device_mesh: - self.device_mesh = DeviceMesh(**device_mesh) - else: - self.device_mesh = DeviceMesh.from_sizes(**device_mesh) + device_group=self.device_group, + device_mesh_dict=device_mesh, + ncpu_proc_per_node=ncpu_proc_per_node, + ) # processor objects keyed by processor_id self.resource_dict: dict[str, Any] = {} @@ -82,10 +79,17 @@ async def _ensure_sticky(self): self._ensure_countdown_started() await self._ensure_state_cleanup_started() - def _on_processor_expired(self, processor_id: str) -> None: - """Called by the countdown thread when a processor's session expires.""" + async def _on_processor_expired(self, processor_id: str) -> None: + """Remove the local processor and release its shared quota lease.""" + info = self.get_resource_info(processor_id) self.resource_dict.pop(processor_id, None) self.unregister_resource(processor_id) + if info is not None: + await self.state.release_processor_quota( + info['token'], + processor_id, + lease_seconds=self._processor_quota_lease_seconds, + ) def build_processor_app(ncpu_proc_per_node: int, diff --git a/src/twinkle/server/processor/twinkle_handlers.py b/src/twinkle/server/processor/twinkle_handlers.py index d9da305ba..5de7814ec 100644 --- a/src/twinkle/server/processor/twinkle_handlers.py +++ b/src/twinkle/server/processor/twinkle_handlers.py @@ -1,10 +1,11 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -""" -Processor management handler mixin. +"""Processor management routes for the Processor deployment. -All endpoints are prefixed /twinkle/... and handle processor lifecycle -(create, call). self_fn is injected via FastAPI Depends to obtain the -ProcessorManagement instance at request time. +Registered by ``_register_processor_routes(app, self_fn)`` -- module-level route +registration closing over ``self_fn`` via ``Depends``, not a mixin: there is no +inheritance relationship with the deployment class. All endpoints are prefixed +/twinkle/... and handle processor lifecycle (create, call). ``self_fn`` is injected via +FastAPI Depends to obtain the ProcessorManagement instance at request time. """ from __future__ import annotations @@ -18,12 +19,12 @@ if TYPE_CHECKING: from .app import ProcessorManagement -import twinkle_client.types as types +import twinkle.protocol.types as types +from twinkle.protocol.serialize import deserialize_object +from twinkle.server.middleware.auth import get_session_id_from_request, get_token_from_request from twinkle.server.telemetry.correlation import SESSION_ID, TOKEN_ID from twinkle.server.telemetry.tracing import traced_operation -from twinkle.server.utils.validation import get_session_id_from_request, get_token_from_request from twinkle.utils.logger import get_logger -from twinkle_client.common.serialize import deserialize_object logger = get_logger() @@ -45,7 +46,7 @@ async def create( processor_type_name = body.processor_type class_type = body.class_type - _kwargs = body.model_extra or {} + _kwargs = dict(body.init_kwargs) assert processor_type_name in _PROCESSOR_TYPES, f'Invalid processor type: {processor_type_name}' processor_module = importlib.import_module(f'twinkle.{processor_type_name}') @@ -55,39 +56,63 @@ async def create( session_id = get_session_id_from_request(request) processor_id = str(uuid.uuid4().hex) - # Register for lifecycle tracking (enforces per-user limit) - self.register_resource(processor_id, token, session_id) - - _kwargs.pop('remote_group', None) - _kwargs.pop('device_mesh', None) - - resolved_kwargs = {} - for key, value in _kwargs.items(): - if isinstance(value, str) and value.startswith('pid:'): - ref_id = value[4:] - resolved_kwargs[key] = self.resource_dict[ref_id] - else: - value = deserialize_object(value) - resolved_kwargs[key] = value - - # Run processor instantiation in a thread to avoid blocking the event loop, - # which would starve the session-liveness coroutines submitted by the - # countdown thread via asyncio.run_coroutine_threadsafe. - _remote_group = self.device_group.name - _device_mesh = self.device_mesh - - def _do_create(): - return getattr(processor_module, class_type)( - remote_group=_remote_group, device_mesh=_device_mesh, instance_id=processor_id, **resolved_kwargs) - - # Span the primary processor.create op with token + session correlation. - with traced_operation( - f'processor.create.{processor_type_name}.{class_type}', attrs={ - TOKEN_ID: token, - SESSION_ID: session_id, - }): - processor = await asyncio.get_running_loop().run_in_executor(None, _do_create) - self.resource_dict[processor_id] = processor + # Reject malformed registrations before touching shared quota state. + self._validate_registration(processor_id, token, session_id) + await self.state.reserve_processor_quota( + token, + processor_id, + session_id, + limit=self._per_token_processor_limit, + lease_seconds=self._processor_quota_lease_seconds, + ) + + try: + self.register_resource(processor_id, token, session_id) + + _kwargs.pop('remote_group', None) + _kwargs.pop('device_mesh', None) + + resolved_kwargs = {} + for key, value in _kwargs.items(): + if isinstance(value, str) and value.startswith('pid:'): + ref_id = value[4:] + resolved_kwargs[key] = self.resource_dict[ref_id] + else: + value = deserialize_object(value) + resolved_kwargs[key] = value + + # Run processor instantiation in a thread to avoid blocking the event loop, + # which would starve the session-liveness coroutines submitted by the + # countdown thread via asyncio.run_coroutine_threadsafe. + _remote_group = self.device_group.name + _device_mesh = self.device_mesh + + def _do_create(): + return getattr(processor_module, class_type)( + remote_group=_remote_group, device_mesh=_device_mesh, instance_id=processor_id, **resolved_kwargs) + + # Span the primary processor.create op with token + session correlation. + with traced_operation( + f'processor.create.{processor_type_name}.{class_type}', + attrs={ + TOKEN_ID: token, + SESSION_ID: session_id, + }): + processor = await asyncio.get_running_loop().run_in_executor(None, _do_create) + self.resource_dict[processor_id] = processor + except Exception: + self.resource_dict.pop(processor_id, None) + self.unregister_resource(processor_id) + try: + await self.state.release_processor_quota( + token, + processor_id, + lease_seconds=self._processor_quota_lease_seconds, + ) + except Exception as release_error: + logger.warning('Failed to release processor quota after create failure for %s: %r', processor_id, + release_error) + raise return types.ProcessorCreateResponse(processor_id='pid:' + processor_id) @app.post('/twinkle/call', response_model=types.ProcessorCallResponse) @@ -98,7 +123,7 @@ async def call( processor_id = body.processor_id function_name = body.function - _kwargs = body.model_extra or {} + _kwargs = dict(body.call_kwargs) processor_id = processor_id[4:] self.assert_resource_exists(processor_id) processor = self.resource_dict.get(processor_id) diff --git a/src/twinkle/server/runtime.py b/src/twinkle/server/runtime.py new file mode 100644 index 000000000..1a6d8e131 --- /dev/null +++ b/src/twinkle/server/runtime.py @@ -0,0 +1,52 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Twinkle distributed-runtime initialisation for the deployment classes. + +Moved out of ``deployment.py``: that module is deployment-*construction* +infrastructure -- the FastAPI scaffold, the middleware stack, the ``serve.ingress`` +chain -- and ``twinkle.initialize`` + ``DeviceMesh`` construction is neither. It lived +there only because Model and Sampler both needed it; now Processor reuses it too, so its +formerly-inlined copy is gone. +""" +from __future__ import annotations + +from typing import Any + + +def init_twinkle_runtime( + is_mock: bool, + nproc_per_node: int, + device_group: Any, + device_mesh_dict: dict[str, Any], + *, + ncpu_proc_per_node: int | None = None, +) -> Any | None: + """Initialize the Twinkle distributed runtime and build a DeviceMesh. + + Shared by ModelManagement, SamplerManagement and ProcessorManagement ``__init__``. + Returns ``None`` for mock backends (CPU-only, no device mesh). + + ``ncpu_proc_per_node`` is forwarded only when provided; Model/Sampler leave it unset + (preserving their prior behaviour), while Processor passes its own value -- reusing + this function without that parameter would have silently changed the processor's CPU + process count. + """ + import twinkle + from twinkle import DeviceMesh + + if is_mock: + twinkle.initialize( + mode='ray', nproc_per_node=nproc_per_node, ncpu_proc_per_node=1, groups=[device_group], lazy_collect=False) + return None + + init_kwargs: dict[str, Any] = { + 'mode': 'ray', + 'nproc_per_node': nproc_per_node, + 'groups': [device_group], + 'lazy_collect': False, + } + if ncpu_proc_per_node is not None: + init_kwargs['ncpu_proc_per_node'] = ncpu_proc_per_node + twinkle.initialize(**init_kwargs) + if 'mesh_dim_names' in device_mesh_dict: + return DeviceMesh(**device_mesh_dict) + return DeviceMesh.from_sizes(**device_mesh_dict) diff --git a/src/twinkle/server/sampler/app.py b/src/twinkle/server/sampler/app.py index 8941a40bf..b10c5894b 100644 --- a/src/twinkle/server/sampler/app.py +++ b/src/twinkle/server/sampler/app.py @@ -7,18 +7,19 @@ """ from __future__ import annotations -import asyncio from fastapi import FastAPI, Request from ray import serve from typing import Any from twinkle import DeviceGroup -from twinkle.server.deployment import LazyCleanupMixin, bind_deployment, build_deployment_app, init_twinkle_runtime +from twinkle.server.config.backend_dispatch import BackendSelector +from twinkle.server.deployment import LazyCleanupMixin, bind_deployment, build_deployment_app +from twinkle.server.middleware.auth import get_token_from_request +from twinkle.server.runtime import init_twinkle_runtime from twinkle.server.state import ServerState, get_server_state +from twinkle.server.task_queue.config import TaskQueueConfig +from twinkle.server.task_queue.mixin import TaskQueueMixin from twinkle.server.utils import wrap_builder_with_device_group_env -from twinkle.server.utils.backend_dispatch import BackendSelector -from twinkle.server.utils.task_queue import TaskQueueConfig, TaskQueueMixin -from twinkle.server.utils.validation import get_token_from_request from twinkle.utils.logger import get_logger from .tinker_handlers import _register_tinker_sampler_routes from .twinkle_handlers import _register_twinkle_sampler_routes @@ -45,12 +46,6 @@ def _make_vllm_async_sampler(kw: dict[str, Any]) -> Any: return VLLMSamplerTQ(**kw, context_manager=None) -def _make_torch_sampler(kw: dict[str, Any]) -> Any: - from twinkle.sampler import TorchSampler # type: ignore[attr-defined] - - return TorchSampler(**kw) - - # Single validate-then-dispatch selector for the sampler backend. SAMPLER_SELECTOR = BackendSelector( 'sampler_type', @@ -58,7 +53,6 @@ def _make_torch_sampler(kw: dict[str, Any]) -> Any: 'mock': _make_mock_sampler, 'vllm': _make_vllm_sampler, 'vllm_async': _make_vllm_async_sampler, - 'torch': _make_torch_sampler, }, ) @@ -77,7 +71,7 @@ class SamplerManagement(LazyCleanupMixin, TaskQueueMixin): """Unified sampler management service. Manages: - - vLLM or Torch sampler initialization and lifecycle + - mock or vLLM sampler initialization and lifecycle - Tinker inference requests (/tinker/asample) with rate limiting via TaskQueueMixin - Twinkle inference requests (/twinkle/*) calling sampler directly - Template configuration for trajectory encoding @@ -103,12 +97,12 @@ def __init__(self, self.sampler_type = sampler_type self.model_id = model_id replica_context = serve.get_replica_context() - replica_id = replica_context.replica_id.unique_id + self.replica_id = replica_context.replica_id.unique_id sampler_kwargs: dict[str, Any] = { 'model_id': model_id, 'remote_group': self.device_group.name, - 'instance_id': replica_id, + 'instance_id': self.replica_id, } if sampler_type != 'mock': sampler_kwargs.update( @@ -127,14 +121,25 @@ def __init__(self, from twinkle.server.data_plane import DataPlaneProxy self.data_plane = DataPlaneProxy(data_plane_url) - # Initialize task queue mixin - self._init_task_queue(queue_config, deployment_name='Sampler') + actors = getattr(self.sampler, '_actors', None) + self._init_task_queue( + queue_config, + deployment_name='Sampler', + collect_width=len(actors) if actors else 1, + ) + self.sampler._ray_get_timeout = self.task_queue_config.effective_execution_timeout async def shutdown(self) -> None: - cancel_all = getattr(self.sampler, 'cancel_all_generations', None) - if callable(cancel_all): - await asyncio.to_thread(cancel_all) - await self.data_plane.close() + try: + cancel_all = getattr(self.sampler, 'cancel_all_generations', None) + if callable(cancel_all): + await self.call_backend(cancel_all) + finally: + try: + await self.state.unregister_replica(self.replica_id) + finally: + await self.shutdown_task_queue() + await self.data_plane.close() @serve.multiplexed(max_num_models_per_replica=5) async def _sticky_entry(self, sticky_key: str): @@ -146,9 +151,9 @@ async def _ensure_sticky(self): async def _on_request_start(self, request: Request) -> str: await self._ensure_sticky() + await self.state.touch_replica_last_seen(self.replica_id) await self._ensure_state_cleanup_started() - token = get_token_from_request(request) - return token + return get_token_from_request(request) def build_sampler_app(model_id: str, @@ -172,7 +177,7 @@ def build_sampler_app(model_id: str, device_group: Device group configuration dict device_mesh: Device mesh configuration dict for parallelism deploy_options: Ray Serve deployment options - sampler_type: Sampler selector — ``mock`` | ``vllm`` | ``vllm_async`` | ``torch``. + sampler_type: Sampler selector — ``mock`` | ``vllm`` | ``vllm_async``. Validated up front; bad values raise :class:`ConfigError` before any side effect. engine_args: Additional engine arguments for the sampler diff --git a/src/twinkle/server/sampler/backends/__init__.py b/src/twinkle/server/sampler/backends/__init__.py index 82c3255cb..c2caf8e9e 100644 --- a/src/twinkle/server/sampler/backends/__init__.py +++ b/src/twinkle/server/sampler/backends/__init__.py @@ -1,15 +1,11 @@ -STREAM_SENTINEL = '__STREAM_END__' +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Sampler backend implementations. +The cross-process streaming bridge lives in ``streaming`` and is re-exported here +so a sibling backend module can import it from ``.streaming`` without importing +this package root, while external consumers keep using +``from twinkle.server.sampler.backends import stream_to_queue``. +""" +from .streaming import STREAM_SENTINEL, stream_to_queue -def stream_to_queue(sampler, queue, inputs, sampling_params=None, adapter_name='', adapter_path=None): - """Push streaming deltas from *sampler* to a cross-process Ray queue. - - Works with any object that exposes a ``sample_stream`` iterator. - """ - try: - for delta, reason in sampler.sample_stream(inputs, sampling_params, adapter_name, adapter_path): - queue.put((delta, reason)) - except Exception as e: - queue.put(e) - finally: - queue.put(STREAM_SENTINEL) +__all__ = ['STREAM_SENTINEL', 'stream_to_queue'] diff --git a/src/twinkle/server/sampler/backends/mock_sampler.py b/src/twinkle/server/sampler/backends/mock_sampler.py index 2d5e5930d..4d5179033 100644 --- a/src/twinkle/server/sampler/backends/mock_sampler.py +++ b/src/twinkle/server/sampler/backends/mock_sampler.py @@ -168,7 +168,7 @@ def sample_stream( def sample_stream_to_queue(self, queue, inputs, sampling_params=None, adapter_name='', adapter_path=None): """Push streaming deltas to a cross-process Ray queue.""" - from . import stream_to_queue + from .streaming import stream_to_queue stream_to_queue(self, queue, inputs, sampling_params, adapter_name, adapter_path) @remote_function() diff --git a/src/twinkle/server/sampler/backends/streaming.py b/src/twinkle/server/sampler/backends/streaming.py new file mode 100644 index 000000000..dc61ff7e7 --- /dev/null +++ b/src/twinkle/server/sampler/backends/streaming.py @@ -0,0 +1,26 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Cross-process streaming bridge for sampler backends. + +Lives in its own module (rather than the ``backends`` package ``__init__``) so +that a backend implementation (e.g. ``mock_sampler``) can import it without +importing its own package root — which would be a package-initialisation-order +dependency. ``backends/__init__`` re-exports these names, so external consumers +(``from twinkle.server.sampler.backends import stream_to_queue``) are unchanged. +""" +from __future__ import annotations + +STREAM_SENTINEL = '__STREAM_END__' + + +def stream_to_queue(sampler, queue, inputs, sampling_params=None, adapter_name='', adapter_path=None): + """Push streaming deltas from *sampler* to a cross-process Ray queue. + + Works with any object that exposes a ``sample_stream`` iterator. + """ + try: + for delta, reason in sampler.sample_stream(inputs, sampling_params, adapter_name, adapter_path): + queue.put((delta, reason)) + except Exception as e: + queue.put(e) + finally: + queue.put(STREAM_SENTINEL) diff --git a/src/twinkle/server/sampler/tinker_handlers.py b/src/twinkle/server/sampler/tinker_handlers.py index de43b4b93..153a3939f 100644 --- a/src/twinkle/server/sampler/tinker_handlers.py +++ b/src/twinkle/server/sampler/tinker_handlers.py @@ -1,8 +1,10 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -""" -Tinker-compatible sampler handler mixin. +"""Tinker-compatible routes for the Sampler deployment. -Provides POST /tinker/asample using schedule_task() returning UntypedAPIFuture. +Registered by ``_register_tinker_sampler_routes(app, self_fn)`` -- module-level route +registration closing over ``self_fn`` via ``Depends``, not a mixin: there is no +inheritance relationship with the deployment class. Provides POST /tinker/asample using +schedule_task() returning UntypedAPIFuture. """ from __future__ import annotations @@ -18,12 +20,30 @@ from twinkle.data_format import SamplingParams from twinkle.server.checkpoint import create_checkpoint_manager +from twinkle.server.sampler.weights import resolve_sampler_weights +from twinkle.server.task_queue.types import UserTaskError from twinkle.server.utils import get_template_for_model from twinkle.utils.logger import get_logger logger = get_logger() +def _sampled_sequence(*, stop_reason, tokens, logprobs): + return types.SampledSequence( + stop_reason=stop_reason, + tokens=tokens, + logprobs=logprobs, + ) + + +def _sample_response(*, sequences, prompt_logprobs, topk_prompt_logprobs): + return types.SampleResponse( + sequences=sequences, + prompt_logprobs=prompt_logprobs, + topk_prompt_logprobs=topk_prompt_logprobs, + ) + + def _register_tinker_sampler_routes(app: FastAPI, self_fn: Callable[[], SamplerManagement]) -> None: """Register the tinker sampler route on the given FastAPI app. @@ -52,9 +72,12 @@ async def _do_sample(): # Set template for sampler based on model type template = get_template_for_model(self.model_id) - self.sampler.set_template(template, model_id=self.model_id) - # Reset prefix cache for new weights - self.sampler.reset_prefix_cache() + await self.call_backend(self.sampler.set_template, template, model_id=self.model_id) + # Reset prefix cache unconditionally on every tinker request (by + # design): the tinker dialect does not signal whether weights + # changed, so it always invalidates. This differs from the twinkle + # endpoints, which reset only when an adapter_uri is supplied. + await self.call_backend(self.sampler.reset_prefix_cache) # Get model_path from body or sampling session model_path = body.model_path @@ -71,10 +94,7 @@ async def _do_sample(): # Base-model sampling is valid when no model_path was provided. if adapter_uri and not os.path.exists(adapter_uri): - return types.RequestFailedResponse( - error=f'Adapter URI {model_path} does not exist. Please check the model_path.', - category=types.RequestErrorCategory.User, - ) + raise UserTaskError(f'Adapter URI {model_path} does not exist. Please check the model_path.') # Convert tinker SamplingParams to twinkle SamplingParams if needed sampling_params = None @@ -85,20 +105,19 @@ async def _do_sample(): top_p=body.sampling_params.top_p, top_k=body.sampling_params.top_k, stop=body.sampling_params.stop, + # tinker 0.16.1 has no SamplingParams.logprobs field, but its + # SampledSequence contract and GRPO training require one + # chosen-token logprob per generated token. + logprobs=1, ) - # A resolved checkpoint is either a LoRA adapter dir (has - # adapter_config.json) or a full-parameter HF checkpoint. Full - # checkpoints are loaded into the sampler base model instead of - # being passed as a LoRA adapter. - lora_path = None - if adapter_uri: - if os.path.exists(os.path.join(adapter_uri, 'adapter_config.json')): - lora_path = adapter_uri - else: - self.sampler.load_full_weights_from_path(adapter_uri) - - responses = self.sampler.sample( + # LoRA adapter dir vs full-parameter checkpoint (shared helper); + # a full checkpoint is loaded into the base model and yields no + # LoRA path. + lora_path = await resolve_sampler_weights(self, adapter_uri) + + responses = await self.call_backend( + self.sampler.sample, inputs=[prompt_inputs] * body.num_samples, sampling_params=sampling_params, adapter_path=lora_path, @@ -116,25 +135,26 @@ async def _do_sample(): flattened = [float(lp_list[0][1]) for lp_list in seq.logprobs if lp_list] except (IndexError, TypeError): flattened = [] - if flattened and len(flattened) == len(seq.logprobs): + if len(flattened) == len(seq.tokens): logprobs = flattened + else: + raise RuntimeError( + f'Sampler returned {len(flattened)} logprobs for {len(seq.tokens)} generated ' + 'tokens; refusing to return a misaligned Tinker SampledSequence.') tinker_sequences.append( - types.SampledSequence( + _sampled_sequence( stop_reason=seq.stop_reason, tokens=list(seq.tokens), logprobs=logprobs, )) - return types.SampleResponse( + return _sample_response( sequences=tinker_sequences, prompt_logprobs=responses[0].prompt_logprobs, topk_prompt_logprobs=responses[0].topk_prompt_logprobs, ) except Exception: logger.error(traceback.format_exc()) - return types.RequestFailedResponse( - error=traceback.format_exc(), - category=types.RequestErrorCategory.Server, - ) + raise input_tokens = len(body.prompt.to_ints()) return await self.schedule_task( diff --git a/src/twinkle/server/sampler/twinkle_handlers.py b/src/twinkle/server/sampler/twinkle_handlers.py index 609c25f21..2aca4b22b 100644 --- a/src/twinkle/server/sampler/twinkle_handlers.py +++ b/src/twinkle/server/sampler/twinkle_handlers.py @@ -1,8 +1,10 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -""" -Twinkle-native sampler handler mixin. +"""Twinkle-native routes for the Sampler deployment. -Provides /twinkle/* sampler endpoints. +Registered by ``_register_twinkle_sampler_routes(app, self_fn)`` -- module-level route +registration closing over ``self_fn`` via ``Depends``, not a mixin: there is no +inheritance relationship with the deployment class. Provides /twinkle/* sampler +endpoints. """ from __future__ import annotations @@ -11,24 +13,27 @@ import traceback import uuid from collections.abc import Callable -from fastapi import Depends, FastAPI, HTTPException, Request +from fastapi import Depends, FastAPI, Request from fastapi.responses import StreamingResponse from typing import TYPE_CHECKING -from twinkle_client.common.serialize import deserialize_object - if TYPE_CHECKING: from .app import SamplerManagement import numpy as np -import twinkle_client.types as types -from twinkle.data_format import InputFeature, SamplingParams, Trajectory -from twinkle.server.telemetry.correlation import MODEL_ID, TOKEN_ID +import twinkle.protocol.types as types +from twinkle.data_format import SamplingParams +from twinkle.protocol.json_utils import json_safe +from twinkle.protocol.serialize import deserialize_object +from twinkle.protocol.types import sampler as sampler_types +from twinkle.server.exceptions import EndpointUnavailableError, RequestRejectedError +from twinkle.server.lifecycle.submit import backend_kwargs, resolve_twinkle_adapter_name, to_backend_inputs +from twinkle.server.sampler.weights import resolve_sampler_weights +from twinkle.server.task_errors import task_error_payload +from twinkle.server.telemetry.correlation import MODEL_ID from twinkle.server.telemetry.tracing import traced_operation -from twinkle.server.utils.validation import get_session_id_from_request from twinkle.utils.logger import get_logger -from twinkle_client.common.json_utils import json_safe logger = get_logger() @@ -51,14 +56,6 @@ def _serialize_input_feature(feature: dict) -> dict: return result -def _get_twinkle_sampler_adapter_name(request: Request, adapter_name: str | None) -> str | None: - """Build a stable per-session adapter name, falling back to request_id for older clients.""" - if adapter_name is None or adapter_name == '': - return None - owner_id = get_session_id_from_request(request) or request.state.request_id - return owner_id + '-' + adapter_name - - def _build_rollout_rows_and_tags( sample_models: list[types.SampleResponseModel], *, @@ -132,31 +129,55 @@ def _submission_states(value) -> list[dict]: return value if isinstance(value, list) else [value] -async def _await_generation( - sampler, - submission_id: str, -): - """Poll an admitted generation without occupying the sampler admission queue.""" - collected = False +async def _stream_queue(q, sentinel, request_id: str, total_timeout: float, single_get_timeout: float = 60.0): + loop = asyncio.get_running_loop() + start = loop.time() try: + while True: + remaining = total_timeout - (loop.time() - start) + if remaining <= 0: + payload = task_error_payload( + 'sample_stream exceeded the execution time bound', request_id=request_id, error_code=504) + yield json.dumps(payload) + '\n' + break + try: + item = await asyncio.wait_for( + loop.run_in_executor(None, q.get), timeout=min(single_get_timeout, remaining)) + except asyncio.TimeoutError: + payload = task_error_payload( + 'sample_stream timed out waiting for the next token', request_id=request_id, error_code=504) + yield json.dumps(payload) + '\n' + break + if item == sentinel: + break + if isinstance(item, Exception): + payload = task_error_payload(f'{type(item).__name__}: {item}', request_id=request_id, error_code=500) + yield json.dumps(payload) + '\n' + break + delta, reason = item + yield json.dumps({'delta': delta, 'finish_reason': reason}) + '\n' + finally: + try: + q.shutdown(force=True) + except Exception: + pass + + +async def _await_generation(service: SamplerManagement, submission_id: str, timeout: float): + """Poll one admitted generation through the backend boundary.""" + collected = False + + async def poll(): + nonlocal collected poll_interval = 0.01 while True: try: - states = _submission_states(await asyncio.to_thread(sampler.get_generation_status, submission_id)) + states = _submission_states(await service.call_backend(service.sampler.get_generation_status, + submission_id)) except Exception as error: - # A pending read-only actor call can be cancelled by Ray while - # the generation submitted just above remains alive. Treating - # that as a generation failure makes the finally block discard - # otherwise valid rollout work. Retry only Ray's explicit task - # cancellation; actor death and application errors must still - # propagate immediately. from ray.exceptions import TaskCancelledError if not isinstance(error, TaskCancelledError): raise - logger.warning( - 'Generation status poll was cancelled; retrying submission %s', - submission_id, - ) await asyncio.sleep(poll_interval) poll_interval = min(poll_interval * 1.5, 0.25) continue @@ -168,15 +189,19 @@ async def _await_generation( error = failed.get('error') or failed.get('status', 'unknown failure') raise RuntimeError(f'generation {submission_id} failed: {error}') if states and all(state.get('status') == 'completed' for state in states): - responses = await asyncio.to_thread(sampler.collect_generation, submission_id) + responses = await service.call_backend(service.sampler.collect_generation, submission_id) collected = True return responses await asyncio.sleep(poll_interval) poll_interval = min(poll_interval * 1.5, 0.25) + + try: + return await asyncio.wait_for(poll(), timeout=timeout) finally: if not collected: try: - await asyncio.to_thread(sampler.cancel_generation, submission_id) + await asyncio.wait_for( + service.call_backend(service.sampler.cancel_generation, submission_id), timeout=4.0) except Exception: logger.warning('Failed to cancel generation %s', submission_id, exc_info=True) @@ -188,30 +213,15 @@ def _register_twinkle_sampler_routes(app: FastAPI, self_fn: Callable[[], Sampler It is wired in via Depends so it is resolved lazily at request time. """ - async def run_task(coro): - """Await a schedule_task_and_wait coroutine and surface any exception as a - structured HTTP 500 response so the client receives the full traceback instead - of an opaque connection-level error. - - Note: HTTPException is re-raised directly to preserve its status code and detail. - """ - try: - return await coro - except HTTPException: - raise - except Exception: - logger.error(traceback.format_exc()) - raise HTTPException(status_code=500, detail=traceback.format_exc()) - - @app.post('/twinkle/create', response_model=types.CreateResponse) - async def create(request: Request, self: SamplerManagement = Depends(self_fn)) -> types.CreateResponse: + @app.post('/twinkle/create', response_model=sampler_types.SamplerCreateResponse) + async def create( + request: Request, self: SamplerManagement = Depends(self_fn)) -> sampler_types.SamplerCreateResponse: """Health check / session creation endpoint.""" - return types.CreateResponse() + return sampler_types.SamplerCreateResponse() - @app.post('/twinkle/sample', response_model=types.SampleResponseModelList) - async def sample( - request: Request, body: types.SampleRequest, - self: SamplerManagement = Depends(self_fn)) -> types.SampleResponseModelList: + @app.post('/twinkle/sample', response_model=types.TaskEnvelope) + async def sample(request: Request, body: types.SampleRequest, + self: SamplerManagement = Depends(self_fn)) -> types.TaskEnvelope: """Sample completions from the model. Supports Trajectory or InputFeature inputs, with optional LoRA adapter. @@ -222,36 +232,18 @@ async def _task(): # Resolve adapter adapter_path = None adapter_name = body.adapter_name or '' - full_adapter_name = _get_twinkle_sampler_adapter_name(request, adapter_name) or '' + full_adapter_name = resolve_twinkle_adapter_name(request, adapter_name) or '' if body.adapter_uri: - import os - from twinkle.server.checkpoint import create_checkpoint_manager checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') _, resolved_uri = checkpoint_manager.parse_adapter_uri(body.adapter_uri) - # Reset prefix cache only when new weights are loaded - self.sampler.reset_prefix_cache() - # LoRA adapter dir (has adapter_config.json) vs full-parameter - # HF checkpoint. Full checkpoints replace the sampler base model. - if resolved_uri and os.path.exists(os.path.join(resolved_uri, 'adapter_config.json')): - adapter_path = resolved_uri - elif resolved_uri: - self.sampler.load_full_weights_from_path(resolved_uri) - - # Parse inputs - inputs = body.inputs - if isinstance(inputs, list) and inputs: - first = inputs[0] - if isinstance(first, dict) and 'input_ids' in first: - inputs = [InputFeature(**item) for item in inputs] - else: - inputs = [Trajectory(**item) for item in inputs] - elif isinstance(inputs, dict): - if 'input_ids' in inputs: - inputs = [InputFeature(**inputs)] - else: - inputs = [Trajectory(**inputs)] + # Reset prefix cache only when new weights are loaded. + await self.call_backend(self.sampler.reset_prefix_cache) + adapter_path = await resolve_sampler_weights(self, resolved_uri) + + # Parse inputs (shared seam; batch form) + inputs = to_backend_inputs(body.inputs) # Build sampling params params = None @@ -259,62 +251,54 @@ async def _task(): params = SamplingParams.from_dict(body.sampling_params) # Sample - responses = self.sampler.sample( + responses = await self.call_backend( + self.sampler.sample, inputs, params, adapter_name=full_adapter_name, adapter_path=adapter_path, ) - return types.SampleResponseModelList(samples=_to_sample_response_models(responses)) - - # Calculate metrics for queue scheduling - inputs_list = body.inputs if isinstance(body.inputs, list) else [body.inputs] - input_tokens = sum(len(inp.get('input_ids', [])) if isinstance(inp, dict) else 0 for inp in inputs_list) - return await run_task( - self.schedule_task_and_wait( - _task, - token=token, - input_tokens=input_tokens, - task_type='sample', - )) + return types.SampleResponseModelList(samples=_to_sample_response_models(responses)).model_dump() + + # Calculate metrics for queue scheduling. The body is wire-validated, so the + # entries are models and ``input_ids`` is absent or a list of ints. + input_tokens = sum(len(getattr(entry, 'input_ids', None) or ()) for entry in body.inputs) + return await self.submit_and_peek(_task, token=token, input_tokens=input_tokens, task_type='sample') - @app.post('/twinkle/sample_to_data_plane', response_model=types.DataRef) + @app.post('/twinkle/sample_to_data_plane', response_model=types.TaskEnvelope) async def sample_to_data_plane( request: Request, body: types.DataPlaneSampleRequest, self: SamplerManagement = Depends(self_fn), - ) -> types.DataRef: - """Generate a complete group, store it server-side, and return its DataRef.""" + ) -> types.TaskEnvelope: + """Generate a complete group, store it server-side, and return a Task_Envelope + whose result is the stored group's DataRef.""" token = await self._on_request_start(request) if not self.data_plane.enabled: - raise HTTPException(status_code=503, detail='sample_to_data_plane requires data_plane_url') + raise EndpointUnavailableError('sample_to_data_plane requires data_plane_url') if not callable(getattr(self.sampler, 'submit_generation', None)): - raise HTTPException(status_code=503, detail='sampler_type must be vllm_async') + raise EndpointUnavailableError('sampler_type must be vllm_async') adapter_path = None - full_adapter_name = _get_twinkle_sampler_adapter_name(request, body.adapter_name) or '' + full_adapter_name = resolve_twinkle_adapter_name(request, body.adapter_name) or '' if body.adapter_uri: from twinkle.server.checkpoint import create_checkpoint_manager checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') _, adapter_path = checkpoint_manager.parse_adapter_uri(body.adapter_uri) inputs = (await self.data_plane.get(body.input_ref) if body.input_ref is not None else body.inputs) - if isinstance(inputs, list) and inputs: - first = inputs[0] - if isinstance(first, dict) and 'input_ids' in first: - inputs = [InputFeature(**item) for item in inputs] - else: - inputs = [Trajectory(**item) for item in inputs] - elif isinstance(inputs, dict): - inputs = [InputFeature(**inputs)] if 'input_ids' in inputs else [Trajectory(**inputs)] + inputs = to_backend_inputs(inputs) params_dict = dict(body.sampling_params or {}) params_dict['num_samples'] = body.num_samples params = SamplingParams.from_dict(params_dict) submission_id = uuid.uuid4().hex - async def _admit(): - await asyncio.to_thread( + async def _generate_and_store(): + # vLLM async engine owns generation concurrency, so the whole + # admit -> await -> store sequence runs as one background future + # (outside the serial compute queue) and its result is the DataRef. + await self.call_backend( self.sampler.submit_generation, submission_id, inputs, @@ -322,33 +306,18 @@ async def _admit(): adapter_name=full_adapter_name, adapter_path=adapter_path, ) - return submission_id - - inline_inputs = body.inputs if isinstance(body.inputs, list) else [body.inputs] - input_tokens = ( - body.input_ref.num_tokens if body.input_ref is not None else sum( - len(item.get('input_ids', [])) for item in inline_inputs if isinstance(item, dict))) - await run_task( - self.schedule_task_and_wait( - _admit, - model_id=full_adapter_name or None, - token=token, - input_tokens=input_tokens, - task_type='sample_admission', - )) + responses = await _await_generation(self, submission_id, self.task_queue_config.effective_execution_timeout) + rows, tags = _build_rollout_rows_and_tags( + _to_sample_response_models(responses), + group_ids=body.group_ids, + policy_version=body.policy_version, + adapter_uri=body.adapter_uri, + ) + ref = await self.data_plane.put([json_safe(item) for item in rows], kind='rollout', tags=tags) + return ref.model_dump() - responses = await _await_generation(self.sampler, submission_id) - rows, tags = _build_rollout_rows_and_tags( - _to_sample_response_models(responses), - group_ids=body.group_ids, - policy_version=body.policy_version, - adapter_uri=body.adapter_uri, - ) - return await self.data_plane.put( - [json_safe(item) for item in rows], - kind='rollout', - tags=tags, - ) + return await self.submit_background_and_peek( + _generate_and_store, model_id=full_adapter_name or None, task_type='sample_to_data_plane') @app.post('/twinkle/unload_adapter_paths') async def unload_adapter_paths( @@ -367,38 +336,41 @@ async def unload_adapter_paths( resolved_paths.append(adapter_path) unload = getattr(self.sampler, 'unload_adapter_paths', None) if unload is not None: - unload(resolved_paths) + await self.call_backend(unload, resolved_paths) return {'status': 'ok'} - @app.post('/twinkle/set_template', response_model=types.SetTemplateResponse) + @app.post('/twinkle/set_template', response_model=sampler_types.SamplerSetTemplateResponse) async def set_template( request: Request, - body: types.SetTemplateRequest, + body: sampler_types.SamplerSetTemplateRequest, self: SamplerManagement = Depends(self_fn), - ) -> types.SetTemplateResponse: + ) -> sampler_types.SamplerSetTemplateResponse: """Set the chat template for encoding Trajectory inputs.""" - extra_kwargs = body.model_extra or {} with traced_operation('sampler.set_template'): - self.sampler.set_template(body.template_cls, **extra_kwargs) - return types.SetTemplateResponse() + await self.call_backend(self.sampler.set_template, body.template_cls, **backend_kwargs(body)) + return sampler_types.SamplerSetTemplateResponse() - @app.post('/twinkle/add_adapter_to_sampler', response_model=types.AddAdapterResponse) + @app.post('/twinkle/add_adapter_to_sampler', response_model=sampler_types.SamplerAddAdapterResponse) async def add_adapter_to_sampler( request: Request, - body: types.AddAdapterRequest, + body: sampler_types.SamplerAddAdapterRequest, self: SamplerManagement = Depends(self_fn), - ) -> types.AddAdapterResponse: + ) -> sampler_types.SamplerAddAdapterResponse: """Add a LoRA adapter to the sampler.""" - assert body.adapter_name, 'You need to specify a valid `adapter_name`' - full_adapter_name = _get_twinkle_sampler_adapter_name(request, body.adapter_name) + # Raised, not asserted: decidable from the request body alone, so it owes the caller + # a real 400 rather than an AssertionError surfacing as a 500 -- and a bare assert + # would vanish under `python -O`, letting an empty adapter_name reach the backend. + if not body.adapter_name: + raise RequestRejectedError('`adapter_name` is required and must be non-empty.') + full_adapter_name = resolve_twinkle_adapter_name(request, body.adapter_name) from peft import LoraConfig config = LoraConfig(**body.config) if isinstance(body.config, dict) else body.config with traced_operation('sampler.add_adapter_to_sampler', attrs={MODEL_ID: self.model_id}): - self.sampler.add_adapter_to_sampler(full_adapter_name, config) + await self.call_backend(self.sampler.add_adapter_to_sampler, full_adapter_name, config) - return types.AddAdapterResponse(adapter_name=full_adapter_name) + return sampler_types.SamplerAddAdapterResponse(adapter_name=full_adapter_name) @app.post('/twinkle/apply_patch') async def apply_patch( @@ -406,10 +378,9 @@ async def apply_patch( body: types.ApplyPatchRequest, self: SamplerManagement = Depends(self_fn), ) -> None: - extra_kwargs = body.model_extra or {} patch_cls = deserialize_object(body.patch_cls) with traced_operation('sampler.apply_patch'): - self.sampler.apply_patch(patch_cls, **extra_kwargs) + await self.call_backend(self.sampler.apply_patch, patch_cls, **backend_kwargs(body)) @app.post('/twinkle/sample_stream') async def sample_stream( @@ -429,32 +400,22 @@ async def sample_stream( adapter_path = None adapter_name = body.adapter_name or '' - full_adapter_name = _get_twinkle_sampler_adapter_name(request, adapter_name) or '' + full_adapter_name = resolve_twinkle_adapter_name(request, adapter_name) or '' if body.adapter_uri: - import os - from twinkle.server.checkpoint import create_checkpoint_manager checkpoint_manager = create_checkpoint_manager(token, client_type='twinkle') _, resolved_uri = checkpoint_manager.parse_adapter_uri(body.adapter_uri) - self.sampler.reset_prefix_cache() - if resolved_uri and os.path.exists(os.path.join(resolved_uri, 'adapter_config.json')): - adapter_path = resolved_uri - elif resolved_uri: - self.sampler.load_full_weights_from_path(resolved_uri) - - inputs = body.inputs - if isinstance(inputs, list): - if len(inputs) != 1: - raise HTTPException(status_code=400, detail='Streaming only supports a single input') - inputs = inputs[0] - if isinstance(inputs, dict): - if 'input_ids' in inputs: - inputs_parsed = InputFeature(**inputs) - else: - inputs_parsed = Trajectory(**inputs) - else: - inputs_parsed = inputs + await self.call_backend(self.sampler.reset_prefix_cache) + adapter_path = await resolve_sampler_weights(self, resolved_uri) + + # Streaming accepts exactly one input; the shared seam enforces that and + # returns a single parsed object. Its ValueError maps to the same 400 this + # endpoint has always returned. + try: + inputs_parsed = to_backend_inputs(body.inputs, single=True) + except ValueError as e: + raise RequestRejectedError(str(e)) params = None if body.sampling_params: @@ -464,8 +425,17 @@ async def sample_stream( from .backends import STREAM_SENTINEL + request_id = f'req_{uuid.uuid4().hex}' + actors = self.sampler._actors + if not actors: + + async def _no_actor_generator(): + payload = task_error_payload('No available sampler actor', request_id=request_id, error_code=503) + yield json.dumps(payload) + '\n' + + return StreamingResponse(_no_actor_generator(), media_type='application/x-ndjson') q = Queue(maxsize=128) - actor = self.sampler._actors[0] + actor = actors[0] actor.sample_stream_to_queue.remote( q, inputs_parsed, @@ -474,16 +444,12 @@ async def sample_stream( adapter_path=adapter_path, ) - async def _stream_generator(): - loop = asyncio.get_event_loop() - while True: - item = await loop.run_in_executor(None, q.get) - if item == STREAM_SENTINEL: - break - if isinstance(item, Exception): - yield json.dumps({'error': str(item)}) + '\n' - break - delta, reason = item - yield json.dumps({'delta': delta, 'finish_reason': reason}) + '\n' - - return StreamingResponse(_stream_generator(), media_type='application/x-ndjson') + return StreamingResponse( + _stream_queue( + q, + STREAM_SENTINEL, + request_id, + self.task_queue_config.effective_execution_timeout, + ), + media_type='application/x-ndjson', + ) diff --git a/src/twinkle/server/sampler/weights.py b/src/twinkle/server/sampler/weights.py new file mode 100644 index 000000000..2e5005df8 --- /dev/null +++ b/src/twinkle/server/sampler/weights.py @@ -0,0 +1,35 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Shared sampler weight-resolution rule (F004 / P004). + +The "a resolved checkpoint is either a LoRA adapter dir or a full-parameter HF +checkpoint" rule was re-derived from a filesystem probe in three sampler handlers +(``sample`` / ``sample_stream`` on the twinkle dialect and ``asample`` on tinker), +and the copies had begun to diverge. It lives here so the storage-layout decision +has a single owner. + +Prefix-cache invalidation is deliberately NOT handled here: each caller keeps its +own ``reset_prefix_cache`` policy (the tinker endpoint resets unconditionally on +every request; the twinkle endpoints reset only when an ``adapter_uri`` is +present), which is an observable behaviour difference this helper must not erase. +""" +from __future__ import annotations + +import os +from typing import Any + + +async def resolve_sampler_weights(service: Any, resolved_uri: str | None) -> str | None: + """Resolve a checkpoint path into a LoRA adapter path (or load full weights). + + Returns the LoRA ``adapter_path`` when ``resolved_uri`` is a directory holding + an ``adapter_config.json``. Otherwise the path is a full-parameter checkpoint: + it is loaded into the sampler base model via ``load_full_weights_from_path`` and + ``None`` is returned (no LoRA adapter to pass). ``None``/empty input returns + ``None`` unchanged. + """ + if not resolved_uri: + return None + if os.path.exists(os.path.join(resolved_uri, 'adapter_config.json')): + return resolved_uri + await service.call_backend(service.sampler.load_full_weights_from_path, resolved_uri) + return None diff --git a/src/twinkle/server/session_resource/__init__.py b/src/twinkle/server/session_resource/__init__.py new file mode 100644 index 000000000..fdb7ebdc1 --- /dev/null +++ b/src/twinkle/server/session_resource/__init__.py @@ -0,0 +1,14 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Session-bound resource lifecycle utilities (adapters / processors). + +Named ``session_resource`` (not ``lifecycle``) to avoid colliding with +``twinkle.server.lifecycle``, which is the *request* lifecycle (submit / retrieve +/ envelope). This package is the *resource* lifecycle: registration, heartbeat and +session-driven expiration of session-bound resources. +""" + +from .adapter import AdapterManagerMixin +from .base import SessionResourceMixin +from .processor import ProcessorManagerMixin + +__all__ = ['AdapterManagerMixin', 'ProcessorManagerMixin', 'SessionResourceMixin'] diff --git a/src/twinkle/server/utils/lifecycle/adapter.py b/src/twinkle/server/session_resource/adapter.py similarity index 63% rename from src/twinkle/server/utils/lifecycle/adapter.py rename to src/twinkle/server/session_resource/adapter.py index b23b84bb4..392a05bf1 100644 --- a/src/twinkle/server/utils/lifecycle/adapter.py +++ b/src/twinkle/server/session_resource/adapter.py @@ -11,7 +11,7 @@ """ from __future__ import annotations -from typing import Any +from abc import abstractmethod from twinkle.utils.logger import get_logger from .base import SessionResourceMixin @@ -29,9 +29,9 @@ class AdapterManagerMixin(SessionResourceMixin): 1. Call _init_adapter_manager() in __init__ 2. Override _on_adapter_expired() to customize expiration handling - Attributes: - _adapter_timeout: Session inactivity timeout in seconds used to determine if a session is alive. - _adapter_max_lifetime: Maximum lifetime in seconds for any adapter, regardless of session liveness. + The inactivity timeout / max lifetime are stored on the base mixin as + ``_resource_timeout`` / ``_resource_max_lifetime`` (set via + ``_init_adapter_manager``). """ # Set resource type for logging @@ -57,38 +57,19 @@ def _init_adapter_manager( resource_max_lifetime=adapter_max_lifetime, ) - @property - def _adapter_timeout(self) -> float: - """Adapter timeout for backward compatibility.""" - return self._resource_timeout - - @property - def _adapter_max_lifetime(self) -> float | None: - """Adapter max lifetime for backward compatibility.""" - return self._resource_max_lifetime - - @property - def _adapter_records(self) -> dict[str, dict[str, Any]]: - """Adapter records for backward compatibility.""" - return self._resource_records - async def _on_resource_expired(self, resource_id: str) -> None: - """Internal hook called by base class. Delegates to _on_adapter_expired.""" + """Base-class expiry hook; forwards to the domain hook ``_on_adapter_expired``. + + ``_on_adapter_expired`` is the supported extension point: the adapter-domain + name is kept deliberately so subclass authors override a method named for + adapters rather than the generic base-class resource hook. + """ await self._on_adapter_expired(resource_id) + @abstractmethod async def _on_adapter_expired(self, adapter_name: str) -> None: - """Hook method called when an adapter expires. - - This method must be overridden by inheriting classes to handle - adapter expiration logic. The base implementation raises NotImplementedError. - - Args: - adapter_name: Name of the expired adapter. - - Raises: - NotImplementedError: If not overridden by inheriting class. - """ - raise NotImplementedError(f'_on_adapter_expired must be implemented by {self.__class__.__name__}') + """Hook method called when an adapter expires.""" + ... @staticmethod def get_adapter_name(adapter_name: str) -> str: @@ -103,7 +84,3 @@ def get_adapter_name(adapter_name: str) -> str: The adapter name to use """ return adapter_name - - def stop_adapter_countdown(self) -> None: - """Stop the background countdown task.""" - self.stop_resource_countdown() diff --git a/src/twinkle/server/utils/lifecycle/base.py b/src/twinkle/server/session_resource/base.py similarity index 74% rename from src/twinkle/server/utils/lifecycle/base.py rename to src/twinkle/server/session_resource/base.py index 394cd2200..4294da510 100644 --- a/src/twinkle/server/utils/lifecycle/base.py +++ b/src/twinkle/server/session_resource/base.py @@ -9,18 +9,19 @@ import asyncio import time -from abc import abstractmethod +from abc import ABC, abstractmethod from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from twinkle.server.state import ServerState +from twinkle.server.exceptions import RequestRejectedError, ResourceNotFoundError from twinkle.utils.logger import get_logger logger = get_logger() -class SessionResourceMixin: +class SessionResourceMixin(ABC): """Base mixin for managing session-bound resources with automatic expiration. This mixin tracks resources and automatically expires them when their @@ -62,37 +63,47 @@ def _init_resource_manager( self._resource_max_lifetime = resource_max_lifetime # Resource lifecycle tracking - # Dict mapping resource_id -> - # {'token': str, 'session_id': str, 'created_at': float, 'state': dict, 'expiring': bool} + # Dict mapping resource_id -> lifecycle metadata, including the last time + # the shared state backend confirmed that the owning session was alive. self._resource_records: dict[str, dict[str, Any]] = {} # Countdown task self._resource_countdown_running = False self._countdown_task: asyncio.Task | None = None - async def _is_session_alive(self, session_id: str) -> bool: - """Check if a session is still alive via state proxy. + async def _is_session_alive(self, session_id: str, resource_id: str, record: dict[str, Any]) -> bool: + """Check session liveness with a bounded fail-open window. - Args: - session_id: Session ID to check - - Returns: - True if session is alive, False if expired or not found + ``record`` is the caller-held lifecycle dict (the same object stored in + ``_resource_records``), passed in so a concurrent ``unregister_resource`` + cannot turn a by-key lookup here into a ``KeyError`` mid-sweep. """ if not session_id: - return True # No session association means always alive + raise RuntimeError(f'registered {self._resource_type} {resource_id} has no session_id') try: last_heartbeat = await self.state.get_session_last_heartbeat(session_id) - except Exception as e: - logger.warning(f'[{self._resource_type}Manager] Failed to check session liveness: {e}') - return True # Assume alive on error + except Exception as exc: + elapsed = time.time() - record['last_liveness_confirmed_at'] + logger.warning( + '[%sManager] Session liveness probe failed for %s; bounded fail-open ' + 'elapsed=%.3fs timeout=%.3fs error=%r', + self._resource_type, + resource_id, + elapsed, + self._resource_timeout, + exc, + ) + return elapsed < self._resource_timeout if last_heartbeat is None: - return False # Session doesn't exist + return False - # Check if session has timed out - return (time.time() - last_heartbeat) < self._resource_timeout + now = time.time() + alive = (now - last_heartbeat) < self._resource_timeout + if alive: + record['last_liveness_confirmed_at'] = now + return alive def _validate_registration(self, resource_id: str, token: str, session_id: str) -> None: """Validate before registering a resource. Override for custom validation. @@ -103,11 +114,11 @@ def _validate_registration(self, resource_id: str, token: str, session_id: str) session_id: Session ID Raises: - ValueError: If validation fails - RuntimeError: If resource limit is reached + RequestRejectedError: If the required session ID is absent. """ if not session_id: - raise ValueError(f'session_id must be provided when registering {self._resource_type} {resource_id}') + raise RequestRejectedError( + f'session_id must be provided when registering {self._resource_type} {resource_id}') def _create_resource_record(self, token: str, session_id: str) -> dict[str, Any]: """Create a new resource record. Override to add custom fields. @@ -119,12 +130,14 @@ def _create_resource_record(self, token: str, session_id: str) -> dict[str, Any] Returns: Resource record dict """ + now = time.time() return { 'token': token, 'session_id': session_id, - 'created_at': time.time(), + 'created_at': now, 'state': {}, 'expiring': False, + 'last_liveness_confirmed_at': now, } def register_resource(self, resource_id: str, token: str, session_id: str) -> None: @@ -136,8 +149,7 @@ def register_resource(self, resource_id: str, token: str, session_id: str) -> No session_id: Session ID to associate with this resource. Raises: - ValueError: If session_id is None or empty. - RuntimeError: If custom validation fails (e.g., limit reached). + RequestRejectedError: If session_id is None or empty. """ self._validate_registration(resource_id, token, session_id) @@ -193,16 +205,6 @@ def get_resource_state(self, resource_id: str, key: str, default: Any = None) -> state = info.get('state') or {} return state.get(key, default) - def pop_resource_state(self, resource_id: str, key: str, default: Any = None) -> Any: - """Pop a per-resource state value.""" - info = self._resource_records.get(resource_id) - if info is None: - return default - state = info.get('state') - if not isinstance(state, dict): - return default - return state.pop(key, default) - def clear_resource_state(self, resource_id: str) -> None: """Clear all per-resource state values.""" info = self._resource_records.get(resource_id) @@ -210,15 +212,41 @@ def clear_resource_state(self, resource_id: str) -> None: return info['state'] = {} + def find_active_resource(self, exclude: str | None = None) -> str | None: + """Return the id of an active (registered, not expiring) resource, if any. + + "Active" means present in the records and not marked ``expiring``. + ``exclude`` (a resource id) is skipped so a caller can ask "is any *other* + resource active?" — e.g. a tenant re-issuing against its own resource must + not count itself. Returns the first matching id, or ``None`` when none is + active. This is the owner-side query that lets callers avoid reaching into + the private ``_resource_records`` dict. + """ + for rid, info in self._resource_records.items(): + if rid != exclude and not info.get('expiring'): + return rid + return None + def assert_resource_exists(self, resource_id: str) -> None: """Validate that a resource exists and is not expiring. Raises: - AssertionError: If resource not found or expiring. + ResourceNotFoundError: 404/User — resource absent or expiring. Raised + (not ``assert``-ed) so the check is classified as a user-facing + 404 rather than collapsing to a 500, and so it survives + ``python -O`` (which strips ``assert`` statements). """ info = self._resource_records.get(resource_id) - assert resource_id and info is not None and not info.get('expiring'), \ - f'{self._resource_type} {resource_id} not found' + if not (resource_id and info is not None and not info.get('expiring')): + raise ResourceNotFoundError(f'{self._resource_type} {resource_id} not found') + + async def _on_resource_liveness_confirmed(self, resource_id: str) -> bool: + """Refresh any resource-specific lease after a successful session probe. + + Returns ``False`` when a resource-specific lease has already been lost and + the local resource must be expired to avoid running outside its quota. + """ + return True @abstractmethod async def _on_resource_expired(self, resource_id: str) -> None: @@ -268,12 +296,9 @@ async def _resource_countdown_loop(self) -> None: expired_resources.append((resource_id, token, session_id)) continue - try: - session_alive = await self._is_session_alive(session_id) - except Exception as e: - logger.warning(f'[{self._resource_type}Manager] Failed to check session liveness ' - f'for {resource_id}: {type(e).__name__}: {e}') - continue + session_alive = await self._is_session_alive(session_id, resource_id, info) + if session_alive: + session_alive = await self._on_resource_liveness_confirmed(resource_id) session_expired = not session_alive logger.debug(f'[{self._resource_type}Manager] {self._resource_type} {resource_id} session check ' f'(session_id={session_id}, session_alive={not session_expired})') @@ -315,10 +340,6 @@ def _ensure_countdown_started(self) -> None: self._countdown_task = asyncio.create_task(self._resource_countdown_loop()) logger.debug(f'[{self._resource_type}Manager] Countdown task started') - async def _async_ensure_countdown_started(self) -> None: - """Async version for convenience.""" - self._ensure_countdown_started() - def stop_resource_countdown(self) -> None: """Stop the background countdown task.""" if self._resource_countdown_running: diff --git a/src/twinkle/server/session_resource/processor.py b/src/twinkle/server/session_resource/processor.py new file mode 100644 index 000000000..69a6192e3 --- /dev/null +++ b/src/twinkle/server/session_resource/processor.py @@ -0,0 +1,92 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +""" +Processor Lifecycle Manager Mixin for Twinkle Server. + +Mirrors AdapterManagerMixin but adds a global per-token processor limit. +Sessions are tracked via session ID; processors expire when their session expires. +""" +from __future__ import annotations + +from abc import abstractmethod + +from twinkle.utils.logger import get_logger +from .base import SessionResourceMixin + +logger = get_logger() + + +class ProcessorManagerMixin(SessionResourceMixin): + """Mixin for processor lifecycle management with session-based expiration. + + Mirrors AdapterManagerMixin with an additional per-token processor limit. + + Inheriting classes should: + 1. Call _init_processor_manager() in __init__ + 2. Override _on_processor_expired() to handle cleanup + + The inactivity timeout is stored on the base mixin as ``_resource_timeout`` + (set via ``_init_processor_manager``); ``_per_token_processor_limit`` caps the + active processors per user token. + """ + + # Set resource type for logging + _resource_type = 'Processor' + + def _init_processor_manager( + self, + processor_timeout: float = 1800.0, + per_token_processor_limit: int = 20, + ) -> None: + """Initialize the processor manager. + + Args: + processor_timeout: Timeout in seconds to determine if a session is alive. + Default is 1800.0 (30 minutes). + per_token_processor_limit: Maximum active processors per user token. + Default is 20. + """ + self._init_resource_manager( + resource_timeout=processor_timeout, + resource_max_lifetime=None, # No max lifetime for processors + ) + self._per_token_processor_limit = per_token_processor_limit + # The countdown runs every 10 seconds. A 30-second lease tolerates two + # missed renewals while still bounding stale reservations after a crash. + self._processor_quota_lease_seconds = 30.0 + + async def _on_resource_liveness_confirmed(self, resource_id: str) -> bool: + """Renew this processor's shared quota lease after a healthy probe.""" + info = self._resource_records.get(resource_id) + if info is None: + return False + try: + renewed = await self.state.renew_processor_quota( + info['token'], + resource_id, + lease_seconds=self._processor_quota_lease_seconds, + ) + except Exception as exc: + # Keep the local processor during a transient backend outage. Once the + # backend recovers, a lost/expired reservation returns False and the + # countdown loop removes the unaccounted local resource. + logger.warning('[ProcessorManager] Failed to renew quota lease for %s: %r', resource_id, exc) + return True + if not renewed: + logger.warning('[ProcessorManager] Quota lease for %s was lost; expiring local processor', resource_id) + return renewed + + async def _on_resource_expired(self, resource_id: str) -> None: + """Base-class expiry hook; forwards to the domain hook ``_on_processor_expired``. + + ``_on_processor_expired`` is the supported extension point: the + processor-domain name is kept deliberately so subclass authors override a + method named for processors rather than the generic base-class hook. It is + ``async`` to match the sibling ``AdapterManagerMixin._on_adapter_expired`` + contract, so both resource kinds expose the same extension-point shape. + """ + await self._on_processor_expired(resource_id) + + @abstractmethod + async def _on_processor_expired(self, processor_id: str) -> None: + """Hook called when a processor's session expires.""" + ... diff --git a/src/twinkle/server/state/__init__.py b/src/twinkle/server/state/__init__.py index 1b6ca9ba5..5f9a47549 100644 --- a/src/twinkle/server/state/__init__.py +++ b/src/twinkle/server/state/__init__.py @@ -1,8 +1,8 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -from twinkle.server.config.persistence import PersistenceConfig from .backend import create_backend from .base import BaseManager -from .config_manager import ConfigManager +from .cleanup_coordinator import ResourceCleanupCoordinator +from .count_publisher import ResourceCountPublisher from .future_manager import FutureManager from .model_manager import ModelManager from .models import FutureRecord, ModelRecord, SamplingSessionRecord, SessionRecord @@ -24,13 +24,12 @@ 'ModelManager', 'SamplingSessionManager', 'FutureManager', - 'ConfigManager', # Server state 'ServerState', 'ReplicaRegistry', + 'ResourceCleanupCoordinator', + 'ResourceCountPublisher', 'get_server_state', - 'reset_server_state_cache', # Persistence backend factory - 'PersistenceConfig', 'create_backend', ] diff --git a/src/twinkle/server/state/backend/__init__.py b/src/twinkle/server/state/backend/__init__.py index 6d46eda0b..420bcef2a 100644 --- a/src/twinkle/server/state/backend/__init__.py +++ b/src/twinkle/server/state/backend/__init__.py @@ -1,7 +1,6 @@ from twinkle.server.config.persistence import PersistenceConfig from .base import StateBackend from .factory import create_backend -from .file_backend import FileBackend from .redis_backend import RedisBackend # NOTE: ``RayActorBackend`` is intentionally NOT imported here. It top-level @@ -10,7 +9,6 @@ # from ``.memory_backend`` — as ``create_backend`` does lazily for memory mode. __all__ = [ 'StateBackend', - 'FileBackend', 'RedisBackend', 'PersistenceConfig', 'create_backend', diff --git a/src/twinkle/server/state/backend/base.py b/src/twinkle/server/state/backend/base.py index e50f0f110..b65f7d2b2 100644 --- a/src/twinkle/server/state/backend/base.py +++ b/src/twinkle/server/state/backend/base.py @@ -4,16 +4,29 @@ from collections.abc import Callable from typing import Any +from twinkle.protocol.types.errors import ErrorCategory +from twinkle.server.exceptions import StateBackendError -class ConcurrencyError(RuntimeError): - """Raised by ``StateBackend.update_atomic`` when contention exhausts retries.""" + +class ConcurrencyError(StateBackendError): + """Atomic state update failed after exhausting the backend retry budget.""" + + error_code = 503 + category = ErrorCategory.Server class StateBackend(ABC): - """Unified interface for state storage backends. + """Unified interface for the Ray-actor and Redis state backends. - All state management operations go through this interface, supporting - multiple backend implementations (memory, file, Redis). + ``close()`` releases only this process's backend handle; shared state must + survive. ``key_prefix`` has intentionally different physical meanings: + Redis prepends it to stored keys while Ray uses it only to namespace the + detached actor. Logical keys returned by ``keys()`` omit the Redis prefix, + so physical keyspaces cannot be copied directly between modes. + + Only ``*`` wildcard matching is portable. Ray's ``fnmatch`` additionally + accepts ``?`` and character classes, while Redis uses its native glob. + ``health_check()`` always returns a strict bool and absorbs transport errors. """ @abstractmethod @@ -72,7 +85,7 @@ async def update_atomic( - The read/transform/write triple is atomic against concurrent callers on the same backend; Redis-backed implementations use WATCH+MULTI+EXEC and may raise :class:`ConcurrencyError` after exhausting internal - retries (default 3). + retries (currently 16). ``transform`` must be picklable when running against a Ray-backed backend — pass module-level functions wrapped with ``functools.partial``, @@ -80,17 +93,14 @@ async def update_atomic( """ ... + @abstractmethod async def mget(self, keys: list[str]) -> list[Any | None]: - """Batch-read multiple keys. Returns values in the same order as *keys*. - - Default implementation falls back to serial ``get()`` calls. - Backends should override for efficiency (e.g. Redis MGET). - """ - return [await self.get(key) for key in keys] + """Batch-read multiple keys, preserving input order.""" + ... @abstractmethod async def close(self) -> None: - """Close backend connection / release resources.""" + """Release this process's handle without changing shared state.""" ... @abstractmethod diff --git a/src/twinkle/server/state/backend/factory.py b/src/twinkle/server/state/backend/factory.py index be24c7824..842a67dfe 100644 --- a/src/twinkle/server/state/backend/factory.py +++ b/src/twinkle/server/state/backend/factory.py @@ -26,17 +26,10 @@ def create_backend(config: PersistenceConfig | None = None) -> StateBackend: match config.mode: case 'memory': - # Deferred import: RayActorBackend pulls in ``ray``, which is an - # optional dependency. Importing it lazily means callers that - # never select memory mode (e.g. file/redis users) do not need - # ray installed just to load this factory. + # Deferred import keeps the module-level import graph light; the Twinkle + # server always runs on Ray Serve, so ``ray`` is available here. from .memory_backend import RayActorBackend return RayActorBackend(key_prefix=config.key_prefix) - case 'file': - if not config.file_path: - raise ValueError('file_path is required for file persistence mode') - from .file_backend import FileBackend - return FileBackend(config.file_path) case 'redis': if not config.redis_url: raise ValueError('redis_url is required for redis persistence mode') diff --git a/src/twinkle/server/state/backend/file_backend.py b/src/twinkle/server/state/backend/file_backend.py deleted file mode 100644 index 42f94e430..000000000 --- a/src/twinkle/server/state/backend/file_backend.py +++ /dev/null @@ -1,235 +0,0 @@ -from __future__ import annotations - -import asyncio -import fcntl -import json -import os -import tempfile -import time -from collections.abc import Callable -from contextlib import contextmanager -from fnmatch import fnmatch -from typing import Any - -from .base import StateBackend - - -class FileBackend(StateBackend): - """Local JSON file-based persistent state backend. - - Storage format is a single JSON file: - ``{key: {"value": ..., "expire_at": float|null}}``. - File I/O is wrapped with ``asyncio.to_thread`` to avoid blocking the - event loop. Every operation that reads-or-writes goes through a sibling - ``.lock`` file held with ``fcntl.LOCK_EX`` so concurrent processes and - coroutines all serialize on the same critical section — that is the only - way ``update_atomic`` can give a meaningful atomicity guarantee against a - concurrent ``set`` / ``delete`` on the same key. - """ - - def __init__(self, file_path: str) -> None: - self._file_path = file_path - self._lock_path = f'{file_path}.lock' - self._init_file() - - def _init_file(self) -> None: - """Auto-create the data file, the lock file, and any missing parent dir.""" - dir_path = os.path.dirname(self._file_path) - if dir_path and not os.path.exists(dir_path): - os.makedirs(dir_path, exist_ok=True) - if not os.path.exists(self._file_path): - with open(self._file_path, 'w', encoding='utf-8') as f: - json.dump({}, f) - # Touch the lock file so flock has a stable inode across processes. - if not os.path.exists(self._lock_path): - with open(self._lock_path, 'a', encoding='utf-8'): - pass - - # ----- lock + file primitives ---------------------------------------- # - - @contextmanager - def _locked(self): - """Hold an exclusive flock on the sibling lock file for the block.""" - with open(self._lock_path, 'a+', encoding='utf-8') as lock_f: - fcntl.flock(lock_f.fileno(), fcntl.LOCK_EX) - try: - yield - finally: - fcntl.flock(lock_f.fileno(), fcntl.LOCK_UN) - - def _load_sync(self) -> dict[str, dict[str, Any]]: - try: - with open(self._file_path, encoding='utf-8') as f: - return json.load(f) - except (json.JSONDecodeError, FileNotFoundError): - return {} - - def _save_sync(self, data: dict[str, dict[str, Any]]) -> None: - """Write temp file then atomic-replace. Caller must hold ``_locked``.""" - # Drop expired entries on the write path so the file never grows - # unbounded with stale keys. - now = time.time() - data = {k: v for k, v in data.items() if v.get('expire_at') is None or v['expire_at'] > now} - - dir_path = os.path.dirname(self._file_path) or '.' - fd = tempfile.NamedTemporaryFile( - mode='w', - suffix='.tmp', - dir=dir_path, - delete=False, - encoding='utf-8', - ) - try: - json.dump(data, fd, ensure_ascii=False) - fd.flush() - os.fsync(fd.fileno()) - fd.close() - os.replace(fd.name, self._file_path) - except BaseException: - if os.path.exists(fd.name): - os.unlink(fd.name) - raise - - def _is_expired(self, entry: dict[str, Any]) -> bool: - expire_at = entry.get('expire_at') - return expire_at is not None and time.time() >= expire_at - - # ----- public API: every op runs under one lock ---------------------- # - - def _set_sync(self, key: str, value: Any, ttl: int | None) -> None: - with self._locked(): - data = self._load_sync() - expire_at = (time.time() + ttl) if ttl is not None else None - data[key] = {'value': value, 'expire_at': expire_at} - self._save_sync(data) - - async def set(self, key: str, value: Any, ttl: int | None = None) -> None: - await asyncio.to_thread(self._set_sync, key, value, ttl) - - def _get_sync(self, key: str) -> Any | None: - with self._locked(): - data = self._load_sync() - entry = data.get(key) - if entry is None: - return None - if self._is_expired(entry): - del data[key] - self._save_sync(data) - return None - return entry['value'] - - async def get(self, key: str) -> Any | None: - return await asyncio.to_thread(self._get_sync, key) - - def _delete_sync(self, key: str) -> None: - with self._locked(): - data = self._load_sync() - if key in data: - del data[key] - self._save_sync(data) - - async def delete(self, key: str) -> None: - await asyncio.to_thread(self._delete_sync, key) - - def _exists_sync(self, key: str) -> bool: - with self._locked(): - data = self._load_sync() - entry = data.get(key) - if entry is None: - return False - if self._is_expired(entry): - del data[key] - self._save_sync(data) - return False - return True - - async def exists(self, key: str) -> bool: - return await asyncio.to_thread(self._exists_sync, key) - - def _keys_sync(self, pattern: str) -> list[str]: - with self._locked(): - data = self._load_sync() - result: list[str] = [] - expired_keys: list[str] = [] - for key, entry in data.items(): - if self._is_expired(entry): - expired_keys.append(key) - continue - if fnmatch(key, pattern): - result.append(key) - if expired_keys: - for key in expired_keys: - del data[key] - self._save_sync(data) - return result - - async def keys(self, pattern: str) -> list[str]: - return await asyncio.to_thread(self._keys_sync, pattern) - - async def count(self, pattern: str) -> int: - return len(await self.keys(pattern)) - - def _set_nx_sync(self, key: str, value: Any, ttl: int | None) -> bool: - with self._locked(): - data = self._load_sync() - entry = data.get(key) - if entry is not None and not self._is_expired(entry): - return False - expire_at = (time.time() + ttl) if ttl is not None else None - data[key] = {'value': value, 'expire_at': expire_at} - self._save_sync(data) - return True - - async def set_nx(self, key: str, value: Any, ttl: int | None = None) -> bool: - return await asyncio.to_thread(self._set_nx_sync, key, value, ttl) - - def _update_atomic_sync( - self, - key: str, - transform: Callable[[Any | None], Any | None], - ttl: int | None, - ) -> Any | None: - with self._locked(): - data = self._load_sync() - entry = data.get(key) - current = None if (entry is None or self._is_expired(entry)) else entry['value'] - new_value = transform(current) - if new_value is None: - return current - expire_at = (time.time() + ttl) if ttl is not None else None - data[key] = {'value': new_value, 'expire_at': expire_at} - self._save_sync(data) - return new_value - - async def update_atomic( - self, - key: str, - transform: Callable[[Any | None], Any | None], - ttl: int | None = None, - ) -> Any | None: - return await asyncio.to_thread(self._update_atomic_sync, key, transform, ttl) - - def _mget_sync(self, keys: list[str]) -> list[Any | None]: - with self._locked(): - data = self._load_sync() - results: list[Any | None] = [] - for key in keys: - entry = data.get(key) - if entry is None or self._is_expired(entry): - results.append(None) - else: - results.append(entry['value']) - return results - - async def mget(self, keys: list[str]) -> list[Any | None]: - return await asyncio.to_thread(self._mget_sync, keys) - - async def close(self) -> None: - """File backend has no persistent connection — nothing to release.""" - pass - - async def health_check(self) -> bool: - try: - return os.access(self._file_path, os.W_OK) - except OSError: - return False diff --git a/src/twinkle/server/state/backend/memory_backend.py b/src/twinkle/server/state/backend/memory_backend.py index 4b7c2d960..433f8baa4 100644 --- a/src/twinkle/server/state/backend/memory_backend.py +++ b/src/twinkle/server/state/backend/memory_backend.py @@ -103,7 +103,8 @@ async def mget(self, keys: list[str]) -> list[Any | None]: results.append(value) return results - async def close(self) -> None: + async def flush_all(self) -> None: + """Destructively clear shared state for tests and explicit admin flows.""" self._store.clear() async def health_check(self) -> bool: @@ -123,8 +124,8 @@ class RayActorBackend(StateBackend): def __init__(self, key_prefix: str = '') -> None: if not ray.is_initialized(): raise RuntimeError('RayActorBackend requires an initialized Ray runtime — call ' - 'ray.init() first, switch persistence to "file"/"redis", or ' - 'rely on the deployment launcher to start Ray.') + 'ray.init() first, switch persistence to "redis", or rely on ' + 'the deployment launcher to start Ray.') name = _actor_name(key_prefix) try: self._actor = ray.get_actor(name) @@ -169,13 +170,13 @@ async def mget(self, keys: list[str]) -> list[Any | None]: return await self._actor.mget.remote(keys) async def close(self) -> None: - await self._actor.close.remote() + """Release this process's actor handle without clearing shared state.""" + self._actor = None async def health_check(self) -> bool: try: - return await self._actor.health_check.remote() - except ray.exceptions.RayActorError: - # The actor crashed (OOM, node died). Don't silently re-create - # it — that would lose all in-memory state. Let readiness probes - # see False and the deployment owner decide to restart. + return bool(await self._actor.health_check.remote()) + except Exception: + # The actor crashed (OOM, node died) or this local handle was closed. + # Do not silently recreate it because that could hide state loss. return False diff --git a/src/twinkle/server/state/backend/redis_backend.py b/src/twinkle/server/state/backend/redis_backend.py index cc11ac3ff..e6769551a 100644 --- a/src/twinkle/server/state/backend/redis_backend.py +++ b/src/twinkle/server/state/backend/redis_backend.py @@ -157,6 +157,6 @@ async def close(self) -> None: async def health_check(self) -> bool: """Check if Redis is healthy and available.""" try: - return await self._client.ping() + return bool(await self._client.ping()) except Exception: return False diff --git a/src/twinkle/server/state/base.py b/src/twinkle/server/state/base.py index 8cb055ae4..bfc238416 100644 --- a/src/twinkle/server/state/base.py +++ b/src/twinkle/server/state/base.py @@ -2,7 +2,7 @@ from __future__ import annotations import time -from abc import ABC, abstractmethod +from abc import ABC from datetime import datetime, timezone from pydantic import BaseModel from typing import Generic, TypeVar @@ -17,8 +17,8 @@ class BaseManager(ABC, Generic[T]): """Abstract base class for resource managers using StateBackend. - Provides common async CRUD operations and timestamp parsing. - Subclasses must implement `cleanup_expired`. + Provides common async CRUD operations and timestamp parsing. Cleanup is + deliberately not polymorphic because each manager needs different inputs. """ def __init__(self, backend: StateBackend, key_prefix: str, record_type: type[T], expiration_timeout: float): @@ -76,20 +76,6 @@ async def get_all(self) -> dict[str, T]: result[resource_id] = self._record_type.model_validate(data) return result - # ----- Cleanup ----- - - @abstractmethod - async def cleanup_expired(self, cutoff_time: float, **kwargs) -> int: - """Remove all records older than cutoff_time. - - Args: - cutoff_time: Unix timestamp; records with activity before this are removed. - - Returns: - Number of records removed. - """ - ... - # ----- Helpers ----- def _parse_timestamp(self, timestamp_str: str) -> float: diff --git a/src/twinkle/server/state/cleanup_coordinator.py b/src/twinkle/server/state/cleanup_coordinator.py new file mode 100644 index 000000000..c6278c248 --- /dev/null +++ b/src/twinkle/server/state/cleanup_coordinator.py @@ -0,0 +1,247 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Cleanup orchestration and cleanup-leader election. + +Extracted from ``ServerState`` because these ~120 lines need only the +managers' ``count()`` / concrete ``cleanup_expired`` plus the shared +``StateBackend``, yet used to live beside ~25 CRUD facade methods, so any change +here meant re-reading all ~660 lines to confirm nothing else was touched. + +Leader election lives here rather than in ``ServerState`` because it exists *for* +the cleanup loop: exactly one holder of the lease runs the cascade, so four Ray +Serve workers do not each sweep (and do not each publish resource counts). This +module deliberately does NOT import ``telemetry``: the resource-count publishing +is driven through the ``on_become_leader`` / ``on_lose_leader`` callbacks so the +telemetry dependency stays isolated to ``count_publisher``. +""" +from __future__ import annotations + +import asyncio +import functools +import time +import uuid +from collections.abc import Awaitable, Callable, Mapping +from typing import Any + +from twinkle.utils.logger import get_logger +from .backend import StateBackend +from .base import BaseManager + +logger = get_logger() + +# ---------- Cleanup-leader election ------------------------------------------ +# +# Every Ray Serve worker creates its own ``ServerState``; without coordination +# each one would run the periodic cleanup and metrics-publish loop, so a single +# Twinkle deployment would multiply the work and inflate every gauge by the +# worker count. We elect one leader per backend by racing for a TTL-scoped key +# inside the shared StateBackend: the winner runs cleanup + publishes metrics, +# the others stay quiet. + +LEADER_KEY = 'cleanup_leader' # actual backend key: 'cleanup_leader' +LEASE_TTL = 30 # seconds — leader loses the lease after this without a renew +LEASE_RENEW = 10 # seconds — must be < LEASE_TTL/2 so two missed renews still beat the TTL + + +def _renew_if_owner(current: str | None, *, owner: str) -> str | None: + """``update_atomic`` transform: only re-write the lease if it is still mine.""" + if current == owner: + return owner + return None + + +class ResourceCleanupCoordinator: + """Owns the cleanup loop, the cascaded expiry across managers, and leader election.""" + + def __init__( + self, + backend: StateBackend, + managers: Mapping[str, BaseManager], + *, + cleanup_interval: float, + expiration_timeout: float, + sweep_processor_quotas: Callable[[], Awaitable[None]], + on_become_leader: Callable[[], Awaitable[None]] | None = None, + on_lose_leader: Callable[[], Awaitable[None]] | None = None, + ) -> None: + self._backend = backend + self._managers = dict(managers) + self._cleanup_interval = float(cleanup_interval) + self._expiration_timeout = float(expiration_timeout) + self._sweep_processor_quotas = sweep_processor_quotas + self._on_become_leader = on_become_leader + self._on_lose_leader = on_lose_leader + + self._cleanup_task: asyncio.Task | None = None + self._cleanup_running = False + self._leader_id = uuid.uuid4().hex + self._is_leader = False + self._leader_task: asyncio.Task | None = None + self._leader_running = False + + @property + def is_leader(self) -> bool: + return self._is_leader + + # ----- Resource cleanup ----- + + async def cleanup_expired_resources(self) -> dict[str, int]: + """Clean up expired sessions, models, sampling_sessions, and futures. + + Sessions expire based on last_heartbeat (or created_at). Models and sampling + sessions are also cascade-expired when their owning session expires. Futures + expire based on updated_at (or created_at). + """ + current_time = time.time() + cutoff_time = current_time - self._expiration_timeout + + session_mgr = self._managers['sessions'] + model_mgr = self._managers['models'] + sampling_mgr = self._managers['sampling_sessions'] + future_mgr = self._managers['futures'] + + # Determine expired sessions and remove them in a SINGLE pass, then cascade + # the SAME set to dependent resources. Using one authoritative set closes the + # TOCTOU window where a session touched mid-cleanup could survive removal while + # its children were cascade-deleted. + expired_session_ids, sessions_removed = await session_mgr.collect_and_remove_expired(cutoff_time) + + models_removed = await model_mgr.cleanup_expired(cutoff_time, expired_session_ids=expired_session_ids) + samplings_removed = await sampling_mgr.cleanup_expired(cutoff_time, expired_session_ids=expired_session_ids) + + alive_replica_ids = await model_mgr.get_alive_replica_ids(self._expiration_timeout) + futures_removed = await future_mgr.cleanup_expired(cutoff_time, alive_replica_ids=alive_replica_ids) + await self._sweep_processor_quotas() + + return { + 'sessions': sessions_removed, + 'models': models_removed, + 'sampling_sessions': samplings_removed, + 'futures': futures_removed, + } + + async def _cleanup_loop(self) -> None: + """Background task that periodically cleans up expired resources. + + Gated by leader election — non-leader workers skip the actual cleanup so the + same backend isn't swept 4x by 4 deployment workers. + """ + while self._cleanup_running: + try: + await asyncio.sleep(self._cleanup_interval) + if not self._is_leader: + continue + stats = await self.cleanup_expired_resources() + if any(stats.values()): + logger.debug(f'[ServerState Cleanup] Removed expired resources: {stats}') + except asyncio.CancelledError: + break + except Exception as e: + logger.warning(f'[ServerState Cleanup] Error during cleanup: {e}') + continue + + # ----- Leader election ----- + + async def _leader_loop(self) -> None: + """Acquire and renew the cleanup-leader lease every LEASE_RENEW seconds.""" + await self._try_acquire_or_renew() # Race for leadership at startup + while self._leader_running: + try: + await asyncio.sleep(LEASE_RENEW) + await self._try_acquire_or_renew() + except asyncio.CancelledError: + break + except Exception as e: + logger.warning(f'[ServerState Leader] renew error: {e}') + continue + + async def _try_acquire_or_renew(self) -> None: + was_leader = self._is_leader + try: + if self._is_leader: + val = await self._backend.update_atomic( + LEADER_KEY, + functools.partial(_renew_if_owner, owner=self._leader_id), + ttl=LEASE_TTL, + ) + self._is_leader = (val == self._leader_id) + else: + self._is_leader = await self._backend.set_nx(LEADER_KEY, self._leader_id, ttl=LEASE_TTL) + except Exception as e: + logger.warning(f'[ServerState Leader] backend error during election: {e}') + self._is_leader = False + if was_leader: + # Our renewal failed but our lease value may still be sitting in the + # backend, so a plain ``set_nx`` would keep returning False for up to + # LEASE_TTL and leadership would stall unclaimed. Best-effort delete + # ONLY when we were the leader (never steal a lease another replica + # legitimately holds), swallowing errors so a delete failure cannot + # escape the election loop. The next tick can then re-acquire. + try: + await self._backend.delete(LEADER_KEY) + except Exception: + pass + + if self._is_leader and not was_leader: + logger.info(f'[ServerState] became cleanup leader (id={self._leader_id[:8]})') + if self._on_become_leader is not None: + await self._on_become_leader() + elif not self._is_leader and was_leader: + logger.warning(f'[ServerState] lost cleanup leadership (id={self._leader_id[:8]})') + if self._on_lose_leader is not None: + await self._on_lose_leader() + + # ----- Lifecycle ----- + + async def start(self) -> bool: + """Start the background cleanup + leader-election tasks. + + Idempotent: returns ``False`` if already running. The guard lives here (not in + ``ServerState``) so the implementation and its guard can never drift apart into + a double-start. + """ + if self._cleanup_running: + return False + # Rebuild in-memory indexes from backend data before the loops start. + await self._managers['models'].rebuild_indexes() + self._cleanup_running = True + self._cleanup_task = asyncio.create_task(self._cleanup_loop()) + self._leader_running = True + self._leader_task = asyncio.create_task(self._leader_loop()) + return True + + async def stop(self) -> bool: + """Stop the background cleanup + leader-election tasks. Returns ``False`` if not running.""" + if not self._cleanup_running: + return False + self._cleanup_running = False + if self._cleanup_task: + self._cleanup_task.cancel() + self._cleanup_task = None + self._leader_running = False + if self._leader_task: + self._leader_task.cancel() + self._leader_task = None + if self._is_leader: + # Release callback registration; the lease itself expires on its own TTL — + # update_atomic can't express "atomic delete", so we accept a short outage + # where the gauge reads 0 between leaders. + if self._on_lose_leader is not None: + await self._on_lose_leader() + self._is_leader = False + return True + + async def get_cleanup_stats(self) -> dict[str, Any]: + """Get current cleanup configuration and resource counts.""" + return { + 'expiration_timeout': self._expiration_timeout, + 'cleanup_interval': self._cleanup_interval, + 'cleanup_running': self._cleanup_running, + 'is_leader': self._is_leader, + 'leader_id': self._leader_id, + 'resource_counts': { + 'sessions': await self._managers['sessions'].count(), + 'models': await self._managers['models'].count(), + 'sampling_sessions': await self._managers['sampling_sessions'].count(), + 'futures': await self._managers['futures'].count(), + }, + } diff --git a/src/twinkle/server/state/config_manager.py b/src/twinkle/server/state/config_manager.py deleted file mode 100644 index b5a3185a8..000000000 --- a/src/twinkle/server/state/config_manager.py +++ /dev/null @@ -1,88 +0,0 @@ -# Copyright (c) ModelScope Contributors. All rights reserved. -from __future__ import annotations - -from typing import Any - -from .backend import StateBackend - -# Key prefix used to namespace configuration entries inside the backend. -_CONFIG_PREFIX = 'config::' -_CONFIG_PATTERN = f'{_CONFIG_PREFIX}*' - - -class ConfigManager: - """ - Manages key-value configuration entries via a :class:`StateBackend`. - - Configuration entries have no expiry; they persist until explicitly removed - or cleared. This manager does not inherit from BaseManager because config - values are arbitrary Python objects rather than Pydantic models, and all - storage is delegated to the injected backend. - - Methods are ``async`` because :class:`StateBackend` operations are async. - Atomicity for read-modify-write entries comes from the backend's own - primitives (``set_nx`` / ``update_atomic``), not from any single-threaded - actor assumption — each worker holds its own ``ConfigManager`` bound to the - shared backend, so no additional locking is layered on top of the backend. - """ - - def __init__(self, backend: StateBackend) -> None: - self._backend = backend - - @staticmethod - def _make_key(key: str) -> str: - return f'{_CONFIG_PREFIX}{key}' - - # ----- CRUD ----- - - async def add(self, key: str, value: Any) -> None: - """Add or overwrite a configuration value.""" - await self._backend.set(self._make_key(key), value) - - async def add_or_get(self, key: str, value: Any) -> Any: - """Add a value if the key does not exist; otherwise return the existing value. - - Args: - key: Configuration key. - value: Value to store if the key is absent. - - Returns: - The existing or newly stored value. - """ - backend_key = self._make_key(key) - existing = await self._backend.get(backend_key) - if existing is not None: - return existing - # Use set_nx for atomicity within a single backend; if another - # writer already populated the key we return the winning value. - if await self._backend.set_nx(backend_key, value): - return value - return await self._backend.get(backend_key) - - async def get(self, key: str) -> Any | None: - """Return the configuration value for key, or None.""" - return await self._backend.get(self._make_key(key)) - - async def pop(self, key: str) -> Any | None: - """Remove and return the configuration value for key, or None. - - Note: get-then-delete is not atomic (TOCTOU); a concurrent pop may - return the same value twice. This is acceptable for config entries - where double-return is harmless. - """ - backend_key = self._make_key(key) - value = await self._backend.get(backend_key) - if value is None: - return None - await self._backend.delete(backend_key) - return value - - async def clear(self) -> None: - """Remove all configuration entries.""" - keys = await self._backend.keys(_CONFIG_PATTERN) - for backend_key in keys: - await self._backend.delete(backend_key) - - async def count(self) -> int: - """Return the number of stored configuration entries.""" - return await self._backend.count(_CONFIG_PATTERN) diff --git a/src/twinkle/server/state/count_publisher.py b/src/twinkle/server/state/count_publisher.py new file mode 100644 index 000000000..aa09c2533 --- /dev/null +++ b/src/twinkle/server/state/count_publisher.py @@ -0,0 +1,78 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Resource-count publishing for the cleanup leader. + +This is the single file in ``state/`` that imports the metrics registry. +Pulling ``_metrics_publish_loop`` out of ``ServerState`` shrinks the persistence +layer's dependency on the observability layer to this one module. + +Only the cleanup leader publishes: ``ServerState`` drives this through the +coordinator's ``on_become_leader`` / ``on_lose_leader`` callbacks, so four Ray Serve +workers do not multiply the gauges -- +``test_lgtm_telemetry.py::test_active_sessions_no_4x_inflation`` guards it. +""" +from __future__ import annotations + +import asyncio +from collections.abc import Sequence + +from twinkle.server.telemetry import MetricsRegistry +from twinkle.utils.logger import get_logger +from .base import BaseManager + +logger = get_logger() + + +class ResourceCountPublisher: + """Owns the periodic push of resource counts into the ``MetricsRegistry`` cache. + + The ObservableGauges registered by :class:`MetricsRegistry` read the cache at OTEL + export time and report whatever was pushed last, so this loop is the single writer + of those four gauges. + """ + + def __init__(self, managers: Sequence[tuple[str, BaseManager]], *, interval: float) -> None: + self._managers = tuple(managers) + self._interval = float(interval) + self._task: asyncio.Task | None = None + self._running = False + + async def start(self) -> None: + """Start pushing counts. Idempotent — a second call while running is a no-op.""" + if self._task is not None and not self._task.done(): + return + self._running = True + self._task = asyncio.create_task(self._publish_loop()) + + async def stop(self) -> None: + self._running = False + if self._task is not None: + self._task.cancel() + try: + await self._task + except (asyncio.CancelledError, Exception): + pass + self._task = None + + def clear(self) -> None: + """Zero this worker's resource-gauge cache. + + Called on leadership loss: after the publish loop is cancelled it never + overwrites the cache again, so without zeroing the stale worker would keep + emitting its last counts forever. The new leader publishes the authoritative + counts from its own process. + """ + MetricsRegistry.get().clear_resource_counts() + + async def _publish_loop(self) -> None: + """Push resource counts into the MetricsRegistry cache every N seconds.""" + registry = MetricsRegistry.get() + while self._running: + try: + await asyncio.sleep(self._interval) + for name, mgr in self._managers: + registry.set_resource_count(name, await mgr.count()) + except asyncio.CancelledError: + break + except Exception as e: + logger.debug(f'[ResourceCountPublisher] Error publishing metrics: {e}') + continue diff --git a/src/twinkle/server/state/future_manager.py b/src/twinkle/server/state/future_manager.py index 5adf89ccb..54bca4c70 100644 --- a/src/twinkle/server/state/future_manager.py +++ b/src/twinkle/server/state/future_manager.py @@ -2,27 +2,34 @@ from __future__ import annotations import functools -from datetime import datetime +import time from typing import Any +from twinkle.utils.logger import get_logger from .backend.base import StateBackend from .base import BaseManager -from .models import FutureRecord +from .models import FutureFailureRecord, FutureRecord, _now_iso + +logger = get_logger() # Status sets used by the do-not-regress guard inside the atomic transform. -_TERMINAL_STATUSES = frozenset({'completed', 'failed'}) +_TERMINAL_STATUSES = frozenset({'completed', 'failed', 'cancelled'}) _NON_TERMINAL_STATUSES = frozenset({'pending', 'queued', 'running'}) def _future_record_transform( existing: dict | None, *, + request_id: str, new_status: str, model_id: str | None, reason: str | None, result: Any, + failure: dict[str, Any] | None, queue_state: str | None, queue_state_reason: str | None, + replica_id: str | None, + absolute_deadline: float | None, now: str, ) -> dict | None: """Atomic transform body for :meth:`FutureManager.store_status`. @@ -30,12 +37,16 @@ def _future_record_transform( Module-level so it remains picklable when forwarded across the Ray actor boundary (closures and lambdas cannot be). - Drops the write entirely (returns ``None``) when ``new_status`` would - regress a terminal status — the StateBackend.update_atomic contract treats - a ``None`` return as "keep the current value", which is what stops stale - retries from clobbering a freshly committed terminal state. + A record already in a terminal state is never overwritten (returns ``None``, + which ``update_atomic`` treats as "keep current value"). A write of a + *different* terminal state is logged; a write of the *same* terminal state is + dropped silently (State_Backend idempotent retries produce these and they + indicate no defect). """ - if (existing is not None and existing.get('status') in _TERMINAL_STATUSES and new_status in _NON_TERMINAL_STATUSES): + existing_status = existing.get('status') if existing is not None else None + if existing_status in _TERMINAL_STATUSES: + if new_status != existing_status: + logger.warning('future %s already terminal as %r; refusing %r', request_id, existing_status, new_status) return None if existing is None: @@ -44,8 +55,11 @@ def _future_record_transform( model_id=model_id, reason=reason, result=result, + failure=failure, queue_state=queue_state, queue_state_reason=queue_state_reason, + replica_id=replica_id, + absolute_deadline=absolute_deadline, created_at=now, updated_at=now, ) @@ -55,9 +69,17 @@ def _future_record_transform( updated['status'] = new_status updated['model_id'] = model_id updated['updated_at'] = now + # replica_id is set at creation and is deliberately NOT overwritten here. if reason is not None: updated['reason'] = reason - if result is not None: + if new_status in _TERMINAL_STATUSES: + if new_status in ('failed', 'cancelled'): + updated['result'] = None + updated['failure'] = failure + else: + updated['failure'] = None + updated['result'] = result + elif result is not None: updated['result'] = result if queue_state is not None: updated['queue_state'] = queue_state @@ -66,11 +88,28 @@ def _future_record_transform( return updated -class FutureManager(BaseManager[FutureRecord]): - """Manages async task futures / request statuses. +_CANCELLABLE_STATUSES = frozenset({'pending', 'queued'}) + - Expiry is based on `updated_at` (falls back to `created_at`). +def _cancel_if_not_started_transform(existing: dict | None, *, failure: dict, now: str) -> dict | None: + """Atomic transform: cancel iff the task has not started (pending/queued). + + Returns ``None`` (no change) for running/terminal/missing records so a task + already executing is never interrupted -- cancel is best-effort on the queue. """ + status = existing.get('status') if existing is not None else None + if status not in _CANCELLABLE_STATUSES: + return None + updated = dict(existing) + updated['status'] = 'cancelled' + updated['result'] = None + updated['failure'] = failure + updated['updated_at'] = now + return updated + + +class FutureManager(BaseManager[FutureRecord]): + """Manage future state, terminal retention, and immutable task deadlines.""" def __init__(self, backend: StateBackend, expiration_timeout: float) -> None: super().__init__(backend, 'future::', FutureRecord, expiration_timeout) @@ -84,8 +123,11 @@ async def store_status( model_id: str | None, reason: str | None = None, result: Any = None, + failure: FutureFailureRecord | None = None, queue_state: str | None = None, queue_state_reason: str | None = None, + replica_id: str | None = None, + absolute_deadline: float | None = None, ) -> None: """Create or update a future record with the latest status. @@ -96,42 +138,125 @@ async def store_status( If the result object has a ``model_dump`` method (i.e. it is a Pydantic model) it is serialized to a plain dict before storage. """ + is_failure = status in ('failed', 'cancelled') + if is_failure and failure is None: + raise ValueError(f'{status} future requires a FutureFailureRecord') + if is_failure and result is not None: + raise ValueError(f'{status} future cannot carry result') + if not is_failure and failure is not None: + raise ValueError(f'{status} future cannot carry failure') if result is not None and hasattr(result, 'model_dump'): result = result.model_dump() + failure_data = failure.model_dump() if failure is not None else None - now = datetime.now().isoformat() + now = _now_iso() await self._backend.update_atomic( self._make_key(request_id), functools.partial( _future_record_transform, + request_id=request_id, new_status=status, model_id=model_id, reason=reason, result=result, + failure=failure_data, queue_state=queue_state, queue_state_reason=queue_state_reason, + replica_id=replica_id, + absolute_deadline=absolute_deadline, now=now, ), ) + async def cancel_if_pending(self, request_id: str) -> str | None: + """Cancel a task iff it has not started; return the resulting status. + + Writes a terminal ``cancelled`` record with a domain failure only when + the current status is pending/queued -- a running task is left alone. + Returns the record's status after the attempt, or ``None`` if there is no + record for ``request_id``. + """ + failure = FutureFailureRecord( + reason_code='cancelled', + message='Task cancelled by client', + attribution='user', + ) + result = await self._backend.update_atomic( + self._make_key(request_id), + functools.partial(_cancel_if_not_started_transform, failure=failure.model_dump(), now=_now_iso()), + ) + return result.get('status') if result else None + # ----- Cleanup ----- - async def cleanup_expired(self, cutoff_time: float, **kwargs) -> int: - """Remove futures whose last update is older than cutoff_time. + async def cleanup_expired( + self, + cutoff_time: float, + *, + alive_replica_ids: set[str] | None = None, + ) -> int: + """Expire future records without ever deleting a non-terminal one. + + Processing matrix: + + | status | replica alive | past deadline | action | + |--------------|---------------|---------------|-------------------| + | Terminal | — | ts < cutoff | delete | + | non-Terminal | yes | no | keep (untouched) | + | non-Terminal | yes | yes | write ``failed`` | + | non-Terminal | no | — | write ``failed`` | Args: - cutoff_time: Unix timestamp threshold. + cutoff_time: Unix timestamp; terminal records older than it are deleted. + alive_replica_ids: replicas currently considered alive. ``None`` disables + the orphan check (every non-terminal record is treated as owned). Returns: - Number of futures removed. + Number of terminal records removed (records written ``failed`` are not + counted here; they are removed on a later pass once terminal). """ all_records = await self.get_all() - expired_ids = [] + now = time.time() + expired_ids: list[str] = [] for request_id, record in all_records.items(): - timestamp_str = record.updated_at or record.created_at - timestamp = self._parse_timestamp(timestamp_str) - if timestamp < cutoff_time: - expired_ids.append(request_id) + if record.status in _TERMINAL_STATUSES: + timestamp = self._parse_timestamp(record.updated_at or record.created_at) + if timestamp < cutoff_time: + expired_ids.append(request_id) + continue + + # Non-terminal records are never deleted -- only ever written ``failed``. + # replica_id None (pre-upgrade) => ownership unknown => treated as alive. + replica_id = record.replica_id + replica_alive = (replica_id is None or alive_replica_ids is None or replica_id in alive_replica_ids) + if not replica_alive: + await self.store_status( + request_id, + 'failed', + record.model_id, + failure=FutureFailureRecord( + reason_code='orphaned_replica', + message='The replica that owned this task is no longer available.', + attribution='server', + ), + replica_id=replica_id, + ) + continue + deadline = record.absolute_deadline + if deadline is None: + deadline = self._parse_timestamp(record.created_at) + self.expiration_timeout + if now > deadline: + await self.store_status( + request_id, + 'failed', + record.model_id, + failure=FutureFailureRecord( + reason_code='deadline_exceeded', + message='Task exceeded the absolute survival bound without reaching a terminal state.', + attribution='server', + ), + replica_id=replica_id, + ) for request_id in expired_ids: await self.remove(request_id) diff --git a/src/twinkle/server/state/model_manager.py b/src/twinkle/server/state/model_manager.py index ae75621c4..aaf520fca 100644 --- a/src/twinkle/server/state/model_manager.py +++ b/src/twinkle/server/state/model_manager.py @@ -11,7 +11,9 @@ from __future__ import annotations import functools +import time +from twinkle.server.exceptions import ResourceQuotaExceededError from .backend.base import StateBackend from .base import BaseManager from .models import ModelRecord @@ -36,6 +38,19 @@ def _counter_delta_transform(existing: object, *, delta: int) -> int: return new if new > 0 else 0 +async def _remove_with_record(manager: ModelManager, model_id: str, record: ModelRecord) -> bool: + """Remove a known record without exposing it in the public method signature.""" + removed = await BaseManager.remove(manager, model_id) + if not removed: + return False + if record.token: + await manager._backend.update_atomic( + manager._token_count_key(record.token), + functools.partial(_counter_delta_transform, delta=-1), + ) + return True + + class ModelManager(BaseManager[ModelRecord]): """Manages registered models with backend-derived per-token / per-replica indexes. @@ -101,6 +116,28 @@ async def unregister_replica(self, replica_id: str) -> None: await self.remove(model_id) await self._replicas.unregister(replica_id) + async def touch_replica_last_seen(self, replica_id: str) -> None: + """Refresh a replica's liveness timestamp.""" + await self._replicas.touch_last_seen(replica_id) + + async def get_alive_replica_ids(self, liveness_threshold: float) -> set[str]: + """Return replicas considered alive. + + A replica is alive when it has a ``last_seen`` within ``liveness_threshold``, + OR when it has a ``max_loras`` entry but no ``last_seen`` yet (registered + before this spec / before its first request -- treated as alive so an + upgrade does not orphan in-flight tasks). + """ + registered = await self._replicas.get_all() + last_seen = await self._replicas.get_all_last_seen() + now = time.time() + alive: set[str] = set() + for rid in set(registered) | set(last_seen): + ls = last_seen.get(rid) + if (ls is None and rid in registered) or (ls is not None and (now - ls) <= liveness_threshold): + alive.add(rid) + return alive + async def get_available_replica_ids(self, candidate_ids: list[str]) -> list[str]: """Return the subset of ``candidate_ids`` that still have capacity. @@ -133,12 +170,13 @@ async def add(self, model_id: str, record: ModelRecord) -> None: record is written, so two concurrent adds with the same token cannot both observe ``limit - 1`` and both succeed (the prior count-then-add race). If the increment would exceed the limit, it is rolled back and a - ``RuntimeError`` is raised; if the record write fails, the increment is - rolled back too so the counter never drifts above the real model count. + ``ResourceQuotaExceededError`` is raised; if the record write fails, the + increment is rolled back too so the counter never drifts above the real + model count. Raises: - RuntimeError: when adding ``record`` would exceed - ``per_token_model_limit`` for ``record.token``. + ResourceQuotaExceededError: when adding ``record`` would exceed the + configured per-token model quota. """ token = record.token if not token: @@ -157,7 +195,8 @@ async def add(self, model_id: str, record: ModelRecord) -> None: # Roll the speculative increment back and reject. ``new_count - 1`` # is the count that was already present before this add. await self._backend.update_atomic(key, functools.partial(_counter_delta_transform, delta=-1)) - raise RuntimeError(f'Model limit exceeded: {new_count - 1}/{self._per_token_model_limit} models') + raise ResourceQuotaExceededError( + f'Model quota exceeded for this token: {new_count - 1}/{self._per_token_model_limit} models') try: await super().add(model_id, record) @@ -166,28 +205,18 @@ async def add(self, model_id: str, record: ModelRecord) -> None: await self._backend.update_atomic(key, functools.partial(_counter_delta_transform, delta=-1)) raise - async def remove(self, model_id: str, *, _record: ModelRecord | None = None) -> bool: - """Remove a record by ID, decrementing its owning token's counter. - - When the caller already holds the record (e.g. from a prior ``get_all``), - pass it via ``_record`` to skip the redundant backend fetch. - """ - record = _record or await self.get(model_id) + async def remove(self, model_id: str) -> bool: + """Remove a record by ID and decrement its token quota counter.""" + record = await self.get(model_id) if record is None: return False - await super().remove(model_id) - if record.token: - await self._backend.update_atomic( - self._token_count_key(record.token), - functools.partial(_counter_delta_transform, delta=-1), - ) - return True + return await _remove_with_record(self, model_id, record) # ----- Cleanup -------------------------------------------------------- # - async def cleanup_expired(self, cutoff_time: float, expired_session_ids: list[str] | None = None, **kwargs) -> int: + async def cleanup_expired(self, cutoff_time: float, expired_session_ids: list[str]) -> int: """Remove models older than ``cutoff_time`` or whose owning session expired.""" - session_set = set(expired_session_ids or []) + session_set = set(expired_session_ids) all_records = await self.get_all() expired_ids: list[str] = [] for model_id, record in all_records.items(): @@ -198,17 +227,11 @@ async def cleanup_expired(self, cutoff_time: float, expired_session_ids: list[st if created_at < cutoff_time: expired_ids.append(model_id) for model_id in expired_ids: - await self.remove(model_id, _record=all_records[model_id]) + await _remove_with_record(self, model_id, all_records[model_id]) return len(expired_ids) # ----- Backend-derived helpers --------------------------------------- # - async def _count_models_for_token(self, token: str | None) -> int: - if not token: - return 0 - all_records = await self.get_all() - return sum(1 for r in all_records.values() if r.token == token) - async def _models_for_replica(self, replica_id: str) -> list[str]: all_records = await self.get_all() return [mid for mid, r in all_records.items() if r.replica_id == replica_id] diff --git a/src/twinkle/server/state/models.py b/src/twinkle/server/state/models.py index 71279b894..0364aa6ce 100644 --- a/src/twinkle/server/state/models.py +++ b/src/twinkle/server/state/models.py @@ -2,13 +2,16 @@ from __future__ import annotations import time -from datetime import datetime -from pydantic import BaseModel, Field -from typing import Any +from datetime import datetime, timezone +from pydantic import BaseModel, Field, model_validator +from typing import Any, Literal def _now_iso() -> str: - return datetime.now().isoformat() + # UTC-aware so _parse_timestamp (which reads timestamps back as UTC) agrees with + # it and with time.time(); a naive local string would be misread as UTC and skew + # every expiry comparison by the host's UTC offset. + return datetime.now(timezone.utc).isoformat() class SessionRecord(BaseModel): @@ -44,6 +47,39 @@ class SamplingSessionRecord(BaseModel): created_at: str = Field(default_factory=_now_iso) +class FutureFailureRecord(BaseModel): + """Protocol-independent reason an asynchronous task failed.""" + + reason_code: str + message: str + attribution: Literal['user', 'server'] + details: list[dict[str, Any]] | None = None + diagnostic: str | None = None + + +# Canonical set of protocol-independent failure reason codes. The Twinkle and Tinker +# gateways each map this same key set to their own wire vocabularies; keeping the set +# in one place lets a consistency test catch a map that drifts out of coverage. +FAILURE_REASON_CODES: frozenset[str] = frozenset({ + 'invalid_request', + 'request_rejected', + 'resource_not_found', + 'full_mode_busy', + 'input_tokens_exceeded', + 'batch_size_invalid', + 'rate_limit_exceeded', + 'resource_quota_exceeded', + 'cancelled', + 'endpoint_unavailable', + 'state_contention', + 'backend_gate_unavailable', + 'orphaned_replica', + 'execution_timeout', + 'deadline_exceeded', + 'internal_error', +}) + + class FutureRecord(BaseModel): """Represents an async task future / request status.""" @@ -51,7 +87,22 @@ class FutureRecord(BaseModel): model_id: str | None = None reason: str | None = None result: Any = None + failure: FutureFailureRecord | None = None queue_state: str | None = None queue_state_reason: str | None = None + # Replica ownership and deadline are fixed when the record is created. + replica_id: str | None = None + absolute_deadline: float | None = None created_at: str = Field(default_factory=_now_iso) updated_at: str = Field(default_factory=_now_iso) + + @model_validator(mode='after') + def result_and_failure_match_status(self) -> FutureRecord: + is_failure = self.status in ('failed', 'cancelled') + if is_failure and self.failure is None: + raise ValueError(f'{self.status} future requires failure') + if is_failure and self.result is not None: + raise ValueError(f'{self.status} future cannot carry result') + if not is_failure and self.failure is not None: + raise ValueError(f'{self.status} future cannot carry failure') + return self diff --git a/src/twinkle/server/state/replica_registry.py b/src/twinkle/server/state/replica_registry.py index b2e11d13b..ae1118dd4 100644 --- a/src/twinkle/server/state/replica_registry.py +++ b/src/twinkle/server/state/replica_registry.py @@ -1,28 +1,28 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -"""Backend-backed registry of replica capacity. +"""Backend-backed registry of replica capacity and liveness. -Each entry persists to ``replica::::max_loras`` in the configured -:class:`StateBackend` (Redis or the actor-wrapped RayActorBackend), so every -Ray Serve worker sees one consistent view of the cluster's capacity even -though each worker holds its own ``ServerState`` instance. - -The registry knows *only* about declared capacity. The current loaded-model -count is derived by querying the persisted ``model::*`` records directly — -nothing here caches that count, so concurrent writes from different workers -cannot drift into an inconsistent local index. +Capacity and ``last_seen`` use separate keys so sampler liveness does not alter +the model-capacity data shape. """ from __future__ import annotations +import time + from .backend.base import StateBackend REPLICA_PREFIX = 'replica::' _MAX_LORAS_SUFFIX = '::max_loras' +_LAST_SEEN_SUFFIX = '::last_seen' def _make_key(replica_id: str) -> str: return f'{REPLICA_PREFIX}{replica_id}{_MAX_LORAS_SUFFIX}' +def _last_seen_key(replica_id: str) -> str: + return f'{REPLICA_PREFIX}{replica_id}{_LAST_SEEN_SUFFIX}' + + def _replica_id_from_key(key: str) -> str | None: if not key.startswith(REPLICA_PREFIX) or not key.endswith(_MAX_LORAS_SUFFIX): return None @@ -30,7 +30,7 @@ def _replica_id_from_key(key: str) -> str | None: class ReplicaRegistry: - """Read/write replica capacity through the shared :class:`StateBackend`.""" + """Read/write replica capacity and liveness through the shared backend.""" def __init__(self, backend: StateBackend) -> None: self._backend = backend @@ -42,16 +42,26 @@ async def register(self, replica_id: str, max_loras: int) -> None: async def unregister(self, replica_id: str) -> None: """Remove the capacity entry for ``replica_id`` (idempotent).""" await self._backend.delete(_make_key(replica_id)) + await self._backend.delete(_last_seen_key(replica_id)) + + async def touch_last_seen(self, replica_id: str) -> None: + """Refresh the replica's liveness timestamp (separate key from max_loras).""" + await self._backend.set(_last_seen_key(replica_id), time.time()) - async def get_max_loras(self, replica_id: str) -> int | None: - """Return the declared capacity, or ``None`` if the replica is unknown.""" - value = await self._backend.get(_make_key(replica_id)) - if value is None: - return None - try: - return int(value) - except (TypeError, ValueError): - return None + async def get_all_last_seen(self) -> dict[str, float]: + """Return every replica's last-seen timestamp.""" + keys = await self._backend.keys(f'{REPLICA_PREFIX}*{_LAST_SEEN_SUFFIX}') + out: dict[str, float] = {} + for key in keys: + if not key.startswith(REPLICA_PREFIX) or not key.endswith(_LAST_SEEN_SUFFIX): + continue + rid = key[len(REPLICA_PREFIX):-len(_LAST_SEEN_SUFFIX)] + value = await self._backend.get(key) + try: + out[rid] = float(value) + except (TypeError, ValueError): + continue + return out async def get_all(self) -> dict[str, int]: """Return every registered replica's declared capacity.""" diff --git a/src/twinkle/server/state/sampling_manager.py b/src/twinkle/server/state/sampling_manager.py index 7dd535a5e..a76f24561 100644 --- a/src/twinkle/server/state/sampling_manager.py +++ b/src/twinkle/server/state/sampling_manager.py @@ -18,20 +18,20 @@ def __init__(self, backend: StateBackend, expiration_timeout: float) -> None: # ----- Cleanup ----- - async def cleanup_expired(self, cutoff_time: float, expired_session_ids: list[str] | None = None, **kwargs) -> int: + async def cleanup_expired(self, cutoff_time: float, expired_session_ids: list[str]) -> int: """Remove sampling sessions that are older than cutoff_time, or whose owning session has already been expired. Args: cutoff_time: Unix timestamp threshold. - expired_session_ids: Optional list of session IDs that have just - been expired; any sampling session belonging to one of these + expired_session_ids: Session IDs that have just been expired; any + sampling session belonging to one of these sessions will also be removed regardless of its own age. Returns: Number of sampling sessions removed. """ - session_set = set(expired_session_ids or []) + session_set = set(expired_session_ids) all_records = await self.get_all() expired_ids = [] diff --git a/src/twinkle/server/state/server_state.py b/src/twinkle/server/state/server_state.py index 1fb548f54..dd926aa57 100644 --- a/src/twinkle/server/state/server_state.py +++ b/src/twinkle/server/state/server_state.py @@ -1,8 +1,8 @@ # Copyright (c) ModelScope Contributors. All rights reserved. from __future__ import annotations -import asyncio import functools +import math import re import time import uuid @@ -10,58 +10,109 @@ from typing import Any from twinkle.server.config.persistence import PersistenceConfig -from twinkle.server.telemetry import MetricsRegistry +from twinkle.server.exceptions import ResourceQuotaExceededError from twinkle.server.telemetry.correlation import (BASE_MODEL, MODEL_ID, REPLICA_ID, SAMPLING_SESSION_ID, SESSION_ID, TOKEN_ID) from twinkle.server.telemetry.tracing import traced_operation from twinkle.utils.logger import get_logger from .backend import StateBackend from .backend.factory import create_backend -from .config_manager import ConfigManager +from .cleanup_coordinator import ResourceCleanupCoordinator +from .count_publisher import ResourceCountPublisher from .future_manager import FutureManager from .model_manager import ModelManager -from .models import ModelRecord, SamplingSessionRecord, SessionRecord +from .models import FutureFailureRecord, ModelRecord, SamplingSessionRecord, SessionRecord from .sampling_manager import SamplingSessionManager from .session_manager import SessionManager logger = get_logger() -# ---------- Cleanup-leader election ------------------------------------------ -# -# Every Ray Serve worker creates its own ``ServerState``; without coordination -# each one would run the periodic cleanup and metrics-publish loop, so a -# single Twinkle deployment would multiply the work and inflate every gauge -# by the worker count. We elect one leader per backend by racing for a TTL- -# scoped key inside the shared StateBackend: the winner runs cleanup + -# publishes metrics, the others stay quiet. +_PROCESSOR_QUOTA_PREFIX = 'processor_quota::' + + +def _clean_processor_reservations(existing: Any, *, now: float) -> dict[str, dict[str, Any]]: + """Return only well-formed processor reservations whose leases are active.""" + if not isinstance(existing, dict): + return {} + active: dict[str, dict[str, Any]] = {} + for processor_id, reservation in existing.items(): + if not isinstance(processor_id, str) or not isinstance(reservation, dict): + continue + expires_at = reservation.get('lease_expires_at') + if isinstance(expires_at, (int, float)) and float(expires_at) > now: + active[processor_id] = dict(reservation) + return active + + +def _reserve_processor_transform( + existing: Any, + *, + processor_id: str, + session_id: str, + now: float, + lease_seconds: float, + limit: int, +) -> dict[str, dict[str, Any]]: + reservations = _clean_processor_reservations(existing, now=now) + if processor_id in reservations or len(reservations) < limit: + reservations[processor_id] = { + 'session_id': session_id, + 'lease_expires_at': now + lease_seconds, + } + return reservations + + +def _renew_processor_transform( + existing: Any, + *, + processor_id: str, + now: float, + lease_seconds: float, +) -> dict[str, dict[str, Any]]: + reservations = _clean_processor_reservations(existing, now=now) + reservation = reservations.get(processor_id) + if reservation is not None: + reservation['lease_expires_at'] = now + lease_seconds + return reservations -LEADER_KEY = 'cleanup_leader' # actual backend key: 'cleanup_leader' -LEASE_TTL = 30 # seconds — leader loses the lease after this without a renew -LEASE_RENEW = 10 # seconds — must be < LEASE_TTL/2 so two missed renews still beat the TTL +def _release_processor_transform( + existing: Any, + *, + processor_id: str, + now: float, +) -> dict[str, dict[str, Any]]: + reservations = _clean_processor_reservations(existing, now=now) + reservations.pop(processor_id, None) + return reservations -def _renew_if_owner(current: str | None, *, owner: str) -> str | None: - """``update_atomic`` transform: only re-write the lease if it is still mine.""" - if current == owner: - return owner - return None + +def _sweep_processor_transform(existing: Any, *, now: float) -> dict[str, dict[str, Any]]: + return _clean_processor_reservations(existing, now=now) class ServerState: """Unified server state management class. - Composes five resource managers: + Composes four resource managers: - :class:`SessionManager` — client sessions - :class:`ModelManager` — registered models - :class:`SamplingSessionManager` — sampling sessions - :class:`FutureManager` — async task futures - - :class:`ConfigManager` — key-value configuration Each Ray Serve worker owns one process-local instance, bound directly to a - shared :class:`StateBackend`. The cleanup loop is started from the - deployment's FastAPI ``lifespan`` startup hook and only runs in the worker - that wins the cleanup-leader lease — see :meth:`_leader_loop`. + shared :class:`StateBackend`. + + Cleanup start-up: NOT from FastAPI lifespan startup. Ray Serve binds + ``servable_object`` *after* lifespan startup, so ``start_cleanup_task()`` is + lazy-started on the first request by + ``deployment.LazyCleanupMixin._ensure_state_cleanup_started`` (the reason is + spelled out in ``build_deployment_app``'s lifespan comment). Cleanup + orchestration and cleanup-leader election live in + :class:`ResourceCleanupCoordinator`; resource-count publishing lives in + :class:`ResourceCountPublisher`. ``start_cleanup_task`` is idempotent, which is + what makes per-request invocation safe. """ def __init__( @@ -80,24 +131,44 @@ def __init__( self._model_mgr = ModelManager(self._backend, expiration_timeout, per_token_model_limit) self._sampling_mgr = SamplingSessionManager(self._backend, expiration_timeout) self._future_mgr = FutureManager(self._backend, expiration_timeout) - self._config_mgr = ConfigManager(self._backend) self.expiration_timeout = expiration_timeout self.cleanup_interval = cleanup_interval - self._cleanup_task: asyncio.Task | None = None - self._cleanup_running = False - - # Leader election + metrics-publish loop state. ``metrics_update_interval`` - # is a typed parameter (a misspelled key now fails loudly rather than - # being silently ignored); it controls how often the leader pushes counts - # into the MetricsRegistry cache. - self._leader_id = uuid.uuid4().hex - self._is_leader = False - self._leader_task: asyncio.Task | None = None - self._leader_running = False - self._metrics_publish_task: asyncio.Task | None = None - self._metrics_publish_running = False - self._metrics_update_interval: float = float(metrics_update_interval) + + # Cleanup orchestration, leader election and resource-count publishing are + # delegated. ``metrics_update_interval`` controls how often the leader pushes + # counts into the MetricsRegistry cache. ``on_become_leader`` / + # ``on_lose_leader`` wire leader identity to the publisher's lifecycle so only + # the leader publishes, while keeping the coordinator itself free + # of any telemetry dependency. + _managers = { + 'sessions': self._session_mgr, + 'models': self._model_mgr, + 'sampling_sessions': self._sampling_mgr, + 'futures': self._future_mgr, + } + self._count_publisher = ResourceCountPublisher( + [ + ('active_sessions', self._session_mgr), + ('active_models', self._model_mgr), + ('active_sampling_sessions', self._sampling_mgr), + ('active_futures', self._future_mgr), + ], + interval=float(metrics_update_interval), + ) + self._cleanup = ResourceCleanupCoordinator( + self._backend, + _managers, + cleanup_interval=cleanup_interval, + expiration_timeout=expiration_timeout, + sweep_processor_quotas=self.sweep_processor_quotas, + on_become_leader=self._count_publisher.start, + on_lose_leader=self._on_lose_leader, + ) + + async def _on_lose_leader(self) -> None: + await self._count_publisher.stop() + self._count_publisher.clear() async def get_capacity_info(self) -> dict[str, int]: return await self._model_mgr.get_capacity_info() @@ -237,6 +308,84 @@ async def get_available_replica_ids(self, candidate_ids: list[str]) -> list[str] """ return await self._model_mgr.get_available_replica_ids(candidate_ids) + # ----- Processor Quota Management ----- + + @staticmethod + def _processor_quota_key(token: str) -> str: + return f'{_PROCESSOR_QUOTA_PREFIX}{token}' + + @staticmethod + def _processor_quota_ttl(lease_seconds: float) -> int: + # The key-level TTL is only stale-key hygiene. Individual entries carry + # their own deadlines and are cleaned atomically on every operation. + return max(1, math.ceil(lease_seconds * 2)) + + async def reserve_processor_quota( + self, + token: str, + processor_id: str, + session_id: str, + *, + limit: int, + lease_seconds: float, + ) -> None: + """Atomically reserve one cluster-wide processor slot for ``token``.""" + now = time.time() + reservations = await self._backend.update_atomic( + self._processor_quota_key(token), + functools.partial( + _reserve_processor_transform, + processor_id=processor_id, + session_id=session_id, + now=now, + lease_seconds=lease_seconds, + limit=limit, + ), + ttl=self._processor_quota_ttl(lease_seconds), + ) + if not isinstance(reservations, dict) or processor_id not in reservations: + raise ResourceQuotaExceededError(f'Per-user processor quota ({limit}) reached for token {token[:8]}...') + + async def renew_processor_quota( + self, + token: str, + processor_id: str, + *, + lease_seconds: float, + ) -> bool: + """Renew an existing reservation; never recreate an expired lease.""" + now = time.time() + reservations = await self._backend.update_atomic( + self._processor_quota_key(token), + functools.partial( + _renew_processor_transform, + processor_id=processor_id, + now=now, + lease_seconds=lease_seconds, + ), + ttl=self._processor_quota_ttl(lease_seconds), + ) + return isinstance(reservations, dict) and processor_id in reservations + + async def release_processor_quota(self, token: str, processor_id: str, *, lease_seconds: float = 30.0) -> None: + """Idempotently release a processor reservation.""" + await self._backend.update_atomic( + self._processor_quota_key(token), + functools.partial(_release_processor_transform, processor_id=processor_id, now=time.time()), + ttl=self._processor_quota_ttl(lease_seconds), + ) + + async def sweep_processor_quotas(self, *, lease_seconds: float = 30.0) -> None: + """Remove expired leases from every persisted processor quota map.""" + now = time.time() + keys = await self._backend.keys(f'{_PROCESSOR_QUOTA_PREFIX}*') + for key in keys: + await self._backend.update_atomic( + key, + functools.partial(_sweep_processor_transform, now=now), + ttl=self._processor_quota_ttl(lease_seconds), + ) + # ----- Sampling Session Management ----- async def create_sampling_session(self, payload: dict[str, Any], sampling_session_id: str | None = None) -> str: @@ -280,6 +429,35 @@ async def get_future(self, request_id: str) -> dict[str, Any] | None: record = await self._future_mgr.get(request_id) return record.model_dump() if record is not None else None + async def claim_seq(self, dedup_key: str, request_id: str, ttl: int) -> str | None: + """Idempotency claim for a client seq_id. + + Atomically records ``dedup_key -> request_id`` if unseen and returns ``None`` + (caller proceeds to enqueue). If the key already exists, returns the prior + ``request_id`` so the caller can return that task's envelope instead of + enqueuing a duplicate. ``ttl`` bounds the dedup window. + """ + if await self._backend.set_nx(dedup_key, request_id, ttl=ttl): + return None + return await self._backend.get(dedup_key) + + async def release_seq(self, dedup_key: str) -> None: + """Drop a seq dedup claim (used when the claimed request never enqueued, e.g. + preflight rejected it) so a retry can be admitted rather than see a phantom.""" + await self._backend.delete(dedup_key) + + async def cancel_future(self, request_id: str) -> dict[str, Any]: + """Best-effort cancel: drop the task iff it has not started running. + + Returns ``{'cancelled': bool, 'state': str}`` where ``state`` is the task's + status after the attempt (``cancelled`` if just dropped or already cancelled, + ``running``/``completed``/``failed`` if too late, ``not_found`` if unknown). + """ + status = await self._future_mgr.cancel_if_pending(request_id) + if status is None: + return {'cancelled': False, 'state': 'not_found'} + return {'cancelled': status == 'cancelled', 'state': status} + async def store_future_status( self, request_id: str, @@ -287,25 +465,28 @@ async def store_future_status( model_id: str | None, reason: str | None = None, result: Any = None, + failure: FutureFailureRecord | None = None, queue_state: str | None = None, queue_state_reason: str | None = None, + replica_id: str | None = None, + absolute_deadline: float | None = None, ) -> None: - """Store task status with optional result. + """Store task status with either a success result or domain failure. Supports the full task lifecycle: - PENDING: Task created, waiting to be processed - QUEUED: Task in queue waiting for execution - RUNNING: Task currently executing - COMPLETED: Task completed successfully (result required) - - FAILED: Task failed with error (result contains error payload) - - RATE_LIMITED: Task rejected due to rate limiting (reason required) + - FAILED: Task failed (failure contains protocol-independent details) Args: request_id: Unique identifier for the request. - status: Task status string (pending/queued/running/completed/failed/rate_limited). + status: Task status string (pending/queued/running/completed/failed). model_id: Optional associated model_id. - reason: Optional reason string (used for rate_limited status). - result: Optional result data (used for completed/failed status). + reason: Optional reason string. + result: Optional success result data. + failure: Optional protocol-independent failure record. queue_state: Optional queue state for tinker client (active/paused_rate_limit/paused_capacity). queue_state_reason: Optional reason for the queue state. """ @@ -315,255 +496,58 @@ async def store_future_status( model_id=model_id, reason=reason, result=result, + failure=failure, queue_state=queue_state, queue_state_reason=queue_state_reason, + replica_id=replica_id, + absolute_deadline=absolute_deadline, ) - # ----- Configuration Management ----- - - async def add_config(self, key: str, value: Any) -> None: - """Add or overwrite a configuration value.""" - await self._config_mgr.add(key, value) - - async def add_or_get_config(self, key: str, value: Any) -> Any: - """Add a config value if absent; otherwise return the existing value.""" - return await self._config_mgr.add_or_get(key, value) - - async def get_config(self, key: str) -> Any | None: - """Return the configuration value for key, or None.""" - return await self._config_mgr.get(key) - - async def pop_config(self, key: str) -> Any | None: - """Remove and return the configuration value for key, or None.""" - return await self._config_mgr.pop(key) - - async def clear_config(self) -> None: - """Remove all configuration entries.""" - await self._config_mgr.clear() - - async def count_config(self) -> int: - """Return the number of stored configuration entries.""" - return await self._config_mgr.count() - # ----- Resource Cleanup ----- async def cleanup_expired_resources(self) -> dict[str, int]: """Clean up expired sessions, models, sampling_sessions, and futures. - Sessions expire based on last_heartbeat (or created_at). Models and - sampling sessions are also cascade-expired when their owning session - expires. Futures expire based on updated_at (or created_at). - - Returns: - Dict with counts of cleaned up resources by type. + Delegates to :class:`ResourceCleanupCoordinator`. """ - current_time = time.time() - cutoff_time = current_time - self.expiration_timeout - - # Determine expired sessions and remove them in a SINGLE pass, then - # cascade the SAME set to dependent resources. Using one authoritative - # set (rather than a separate expiry scan followed by a second scan in - # cleanup) closes the TOCTOU window where a session touched mid-cleanup - # could survive removal while its children were cascade-deleted. - expired_session_ids, sessions_removed = await self._session_mgr.collect_and_remove_expired(cutoff_time) - - models_removed = await self._model_mgr.cleanup_expired(cutoff_time, expired_session_ids=expired_session_ids) - samplings_removed = await self._sampling_mgr.cleanup_expired( - cutoff_time, expired_session_ids=expired_session_ids) - futures_removed = await self._future_mgr.cleanup_expired(cutoff_time) - - return { - 'sessions': sessions_removed, - 'models': models_removed, - 'sampling_sessions': samplings_removed, - 'futures': futures_removed, - } + return await self._cleanup.cleanup_expired_resources() - async def _cleanup_loop(self) -> None: - """Background task that periodically cleans up expired resources. + async def touch_replica_last_seen(self, replica_id: str) -> None: + """Refresh a replica's liveness timestamp in the shared registry.""" + await self._model_mgr.touch_replica_last_seen(replica_id) - Gated by leader election — non-leader workers skip the actual cleanup - so the same backend isn't swept 4× by 4 deployment workers. - """ - while self._cleanup_running: - try: - await asyncio.sleep(self.cleanup_interval) - if not self._is_leader: - continue - stats = await self.cleanup_expired_resources() - if any(stats.values()): - logger.debug(f'[ServerState Cleanup] Removed expired resources: {stats}') - except asyncio.CancelledError: - break - except Exception as e: - logger.warning(f'[ServerState Cleanup] Error during cleanup: {e}') - continue - - # ----- Leader election + metrics publish ----- - - async def _leader_loop(self) -> None: - """Acquire and renew the cleanup-leader lease every LEASE_RENEW seconds.""" - await self._try_acquire_or_renew() # Race for leadership at startup - while self._leader_running: - try: - await asyncio.sleep(LEASE_RENEW) - await self._try_acquire_or_renew() - except asyncio.CancelledError: - break - except Exception as e: - logger.warning(f'[ServerState Leader] renew error: {e}') - continue + # ----- Cleanup + leader election (delegated to ResourceCleanupCoordinator) ----- - async def _try_acquire_or_renew(self) -> None: - was_leader = self._is_leader - try: - if self._is_leader: - val = await self._backend.update_atomic( - LEADER_KEY, - functools.partial(_renew_if_owner, owner=self._leader_id), - ttl=LEASE_TTL, - ) - self._is_leader = (val == self._leader_id) - else: - self._is_leader = await self._backend.set_nx(LEADER_KEY, self._leader_id, ttl=LEASE_TTL) - except Exception as e: - logger.warning(f'[ServerState Leader] backend error during election: {e}') - self._is_leader = False - if was_leader: - # Our renewal failed but our lease value may still be sitting in - # the backend, so a plain ``set_nx`` would keep returning False - # for up to LEASE_TTL and leadership would stall unclaimed. - # Best-effort delete ONLY when we were the leader (never steal a - # lease another replica legitimately holds), swallowing errors so - # a delete failure cannot escape the election loop. The next tick - # can then re-acquire immediately. - try: - await self._backend.delete(LEADER_KEY) - except Exception: - pass - - if self._is_leader and not was_leader: - await self._on_become_leader() - elif not self._is_leader and was_leader: - await self._on_lose_leader() - - async def _on_become_leader(self) -> None: - logger.info(f'[ServerState] became cleanup leader (id={self._leader_id[:8]})') - # Start pushing resource counts so the four ObservableGauges in the - # MetricsRegistry have a single source of truth across deployments. - if self._metrics_publish_task is None or self._metrics_publish_task.done(): - self._metrics_publish_running = True - self._metrics_publish_task = asyncio.create_task(self._metrics_publish_loop()) + @property + def _is_leader(self) -> bool: + return self._cleanup._is_leader - async def _on_lose_leader(self) -> None: - logger.warning(f'[ServerState] lost cleanup leadership (id={self._leader_id[:8]})') - self._metrics_publish_running = False - if self._metrics_publish_task is not None: - self._metrics_publish_task.cancel() - try: - await self._metrics_publish_task - except (asyncio.CancelledError, Exception): - pass - self._metrics_publish_task = None - # Clear this worker's resource-gauge cache after cancelling the publish - # task. Across replicas the old leader's MetricsRegistry cache lives in - # a different process, and after handover its publish loop is cancelled - # and never overwrites the cache again — so without this zeroing the - # stale worker would keep emitting its last counts forever. The new - # leader publishes the authoritative counts from its own process. - MetricsRegistry.get().clear_resource_counts() - - async def _metrics_publish_loop(self) -> None: - """Push resource counts into the MetricsRegistry cache every N seconds. - - Only runs while this ``ServerState`` holds the cleanup-leader lease. - The ObservableGauges registered by :class:`MetricsRegistry` read the - cache at OTEL export time and report whatever was pushed last. - """ - registry = MetricsRegistry.get() - sources = ( - ('active_sessions', self._session_mgr), - ('active_models', self._model_mgr), - ('active_sampling_sessions', self._sampling_mgr), - ('active_futures', self._future_mgr), - ) - while self._metrics_publish_running: - try: - await asyncio.sleep(self._metrics_update_interval) - if not self._is_leader: - continue - for name, mgr in sources: - registry.set_resource_count(name, await mgr.count()) - except asyncio.CancelledError: - break - except Exception as e: - logger.debug(f'[ServerState] Error publishing metrics: {e}') - continue + @property + def _leader_id(self) -> str: + return self._cleanup._leader_id + + async def _try_acquire_or_renew(self) -> None: + await self._cleanup._try_acquire_or_renew() async def start_cleanup_task(self) -> bool: """Start the background cleanup + leader-election tasks. - Returns: - True if tasks were started, False if already running. + Returns True if tasks were started, False if already running. The + idempotency guard lives inside the coordinator's ``start`` so a + per-request lazy invocation cannot double-start. """ - if self._cleanup_running: - return False - # Rebuild in-memory indexes from backend data - await self._rebuild_indexes() - self._cleanup_running = True - self._cleanup_task = asyncio.create_task(self._cleanup_loop()) - self._leader_running = True - self._leader_task = asyncio.create_task(self._leader_loop()) - return True - - async def _rebuild_indexes(self) -> None: - """Rebuild in-memory indexes from backend data after startup.""" - # Rebuild model indexes - await self._model_mgr.rebuild_indexes() + return await self._cleanup.start() async def stop_cleanup_task(self) -> bool: """Stop the background cleanup + leader-election tasks. - Returns: - True if tasks were stopped, False if not running. + Returns True if tasks were stopped, False if not running. """ - if not self._cleanup_running: - return False - self._cleanup_running = False - if self._cleanup_task: - self._cleanup_task.cancel() - self._cleanup_task = None - self._leader_running = False - if self._leader_task: - self._leader_task.cancel() - self._leader_task = None - if self._is_leader: - # Release callback registration; the lease itself expires on its own - # TTL — update_atomic can't express "atomic delete", so we accept a - # short outage where the gauge reads 0 between leaders. - await self._on_lose_leader() - self._is_leader = False - return True + return await self._cleanup.stop() async def get_cleanup_stats(self) -> dict[str, Any]: - """Get current cleanup configuration and resource counts. - - Returns: - Dict with cleanup configuration and task status. - """ - return { - 'expiration_timeout': self.expiration_timeout, - 'cleanup_interval': self.cleanup_interval, - 'cleanup_running': self._cleanup_running, - 'is_leader': self._is_leader, - 'leader_id': self._leader_id, - 'resource_counts': { - 'sessions': await self._session_mgr.count(), - 'models': await self._model_mgr.count(), - 'sampling_sessions': await self._sampling_mgr.count(), - 'futures': await self._future_mgr.count(), - }, - } + """Get current cleanup configuration and resource counts.""" + return await self._cleanup.get_cleanup_stats() # --------------------------------------------------------------------------- @@ -573,30 +557,38 @@ async def get_cleanup_stats(self) -> dict[str, Any]: # Each Ray Serve worker binds one ``ServerState`` instance to the shared # ``StateBackend`` for the lifetime of the process — the cleanup loop and # leader-election loop are started exactly once per worker (see -# ``start_cleanup_task``). Callers use ``actor_name`` as the cache key purely -# for per-process deduplication; cross-worker coordination happens inside the -# shared backend, not in this dict. +# ``start_cleanup_task``). Callers use ``cache_key`` purely for per-process +# deduplication; cross-worker coordination happens inside the shared backend, +# not in this dict. _PROCESS_STATE_CACHE: dict[str, ServerState] = {} +# ServerState policy defaults. Used when neither an explicit argument nor a +# launcher-propagated env var (``ServerStateArgs.from_env``) supplies a value. +_DEFAULT_EXPIRATION_TIMEOUT = 86400.0 # 24 hours in seconds +_DEFAULT_CLEANUP_INTERVAL = 3600.0 # 1 hour in seconds +_DEFAULT_PER_TOKEN_MODEL_LIMIT = 30 +_DEFAULT_METRICS_UPDATE_INTERVAL = 15.0 + -def get_server_state(actor_name: str = 'twinkle_server_state', +def get_server_state(cache_key: str = 'twinkle_server_state', backend: StateBackend | None = None, persistence_config: PersistenceConfig | None = None, - expiration_timeout: float = 86400.0, - cleanup_interval: float = 3600.0, - per_token_model_limit: int = 30, - metrics_update_interval: float = 15.0) -> ServerState: + expiration_timeout: float | None = None, + cleanup_interval: float | None = None, + per_token_model_limit: int | None = None, + metrics_update_interval: float | None = None) -> ServerState: """Return a process-local :class:`ServerState` bound directly to the backend. - Within one process the same ``actor_name`` returns the same cached instance + Within one process the same ``cache_key`` returns the same cached instance so repeated callers share one ``ServerState`` and the cleanup loop is started exactly once. Cross-worker consistency comes from the shared :class:`StateBackend` rather than from any singleton in this process. Args: - actor_name: Cache key for the per-process ``ServerState`` instance. - The legacy parameter name is kept for call-site compatibility. + cache_key: Cache key for the per-process ``ServerState`` instance. + (Formerly ``actor_name``; the parameter never carried actor + semantics.) backend: Optional :class:`StateBackend` to inject. When ``None`` a backend is built from ``persistence_config`` (or env vars) via :func:`create_backend`. @@ -613,10 +605,32 @@ def get_server_state(actor_name: str = 'twinkle_server_state', if backend is None and persistence_config is None: persistence_config = PersistenceConfig.from_env() - cached = _PROCESS_STATE_CACHE.get(actor_name) + cached = _PROCESS_STATE_CACHE.get(cache_key) if cached is not None: return cached + # Resolve the ServerState policy: an explicit argument wins, else the + # launcher-propagated env (so a non-gateway worker honours the operator's + # YAML instead of the hardcoded default), else the module default. + from twinkle.server.config.application_spec import ServerStateArgs + env_policy = ServerStateArgs.from_env() + + def _resolve(explicit, env_value, default): + if explicit is not None: + return explicit + if env_value is not None: + return env_value + return default + + expiration_timeout = _resolve(expiration_timeout, getattr(env_policy, 'expiration_timeout', None), + _DEFAULT_EXPIRATION_TIMEOUT) + cleanup_interval = _resolve(cleanup_interval, getattr(env_policy, 'cleanup_interval', None), + _DEFAULT_CLEANUP_INTERVAL) + per_token_model_limit = _resolve(per_token_model_limit, getattr(env_policy, 'per_token_model_limit', None), + _DEFAULT_PER_TOKEN_MODEL_LIMIT) + metrics_update_interval = _resolve(metrics_update_interval, getattr(env_policy, 'metrics_update_interval', None), + _DEFAULT_METRICS_UPDATE_INTERVAL) + state = ServerState( backend=backend, persistence_config=persistence_config, @@ -625,7 +639,11 @@ def get_server_state(actor_name: str = 'twinkle_server_state', per_token_model_limit=per_token_model_limit, metrics_update_interval=metrics_update_interval, ) - _PROCESS_STATE_CACHE[actor_name] = state + _PROCESS_STATE_CACHE[cache_key] = state + logger.info( + 'ServerState policy in effect: per_token_model_limit=%s expiration_timeout=%s ' + 'cleanup_interval=%s metrics_update_interval=%s (resolution: explicit>env>default)', per_token_model_limit, + expiration_timeout, cleanup_interval, metrics_update_interval) # Cleanup task is started by the deployment's FastAPI ``lifespan`` hook # via ``await state.start_cleanup_task()`` — that's the single async # entry point each worker has, so we don't need any sync-context diff --git a/src/twinkle/server/state/session_manager.py b/src/twinkle/server/state/session_manager.py index 1bc901b0c..0e4beeb37 100644 --- a/src/twinkle/server/state/session_manager.py +++ b/src/twinkle/server/state/session_manager.py @@ -97,12 +97,3 @@ async def remove_many(self, ids: list[str]) -> int: if await self.remove(session_id): removed += 1 return removed - - async def cleanup_expired(self, cutoff_time: float, **kwargs) -> int: - """Remove sessions whose last activity is older than ``cutoff_time``. - - Returns: - Number of sessions removed. - """ - _, removed = await self.collect_and_remove_expired(cutoff_time) - return removed diff --git a/src/twinkle/server/task_errors.py b/src/twinkle/server/task_errors.py new file mode 100644 index 000000000..cd94d7b70 --- /dev/null +++ b/src/twinkle/server/task_errors.py @@ -0,0 +1,65 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Helpers for constructing protocol-layer ``ErrorPayload`` values.""" +from __future__ import annotations + +from typing import Any + +from twinkle.protocol.types.errors import ErrorCategory, ErrorPayload + +_ERROR_MAX = 1024 +_TRACEBACK_MAX = 65536 +_TRUNCATION_MARKER = '...[traceback truncated, tail kept]...\n' + + +def trim_traceback(text: str) -> str: + """Keep the tail of an over-long traceback (innermost frames are densest).""" + if len(text) <= _TRACEBACK_MAX: + return text + keep = _TRACEBACK_MAX - len(_TRUNCATION_MARKER) + return _TRUNCATION_MARKER + text[-keep:] + + +def build_error_payload( + error: str, + *, + request_id: str, + error_code: int = 500, + category: ErrorCategory | str = ErrorCategory.Server, + traceback_text: str | None = None, + details: list[dict[str, Any]] | None = None, +) -> ErrorPayload: + """Build a validated error payload with bounded diagnostic text.""" + if isinstance(category, str): + category = ErrorCategory(category.lower()) + tb = trim_traceback(traceback_text) if category is ErrorCategory.Server and traceback_text else None + lines = str(error).splitlines() + summary = (lines[0] if lines else '')[:_ERROR_MAX] + return ErrorPayload( + error=summary, + category=category, + error_code=error_code, + request_id=request_id, + traceback=tb, + details=details, + ) + + +def task_error_payload( + error: str, + *, + request_id: str, + error_code: int = 500, + category: ErrorCategory | str = ErrorCategory.Server, + traceback_text: str | None = None, + details: list[dict[str, Any]] | None = None, +) -> dict[str, Any]: + """Build a JSON-safe wire payload for direct or streaming responses.""" + payload = build_error_payload( + error, + request_id=request_id, + error_code=error_code, + category=category, + traceback_text=traceback_text, + details=details, + ) + return payload.model_dump(mode='json', exclude_none=True) diff --git a/src/twinkle/server/utils/task_queue/__init__.py b/src/twinkle/server/task_queue/__init__.py similarity index 84% rename from src/twinkle/server/utils/task_queue/__init__.py rename to src/twinkle/server/task_queue/__init__.py index 5c90d3181..bf5391611 100644 --- a/src/twinkle/server/utils/task_queue/__init__.py +++ b/src/twinkle/server/task_queue/__init__.py @@ -2,7 +2,7 @@ """ Task Queue package. -Public exports (backward-compatible with the former task_queue.py module): +Public exports: - TaskStatus - task lifecycle enum - QueueState - queue state enum for tinker client compatibility - TaskQueueConfig - queue and rate-limit configuration dataclass @@ -12,12 +12,13 @@ from .config import TaskQueueConfig from .mixin import TaskQueueMixin from .rate_limiter import RateLimiter -from .types import QueuedTask, QueueState, TaskStatus +from .types import QueuedTask, QueueState, TaskStatus, UserTaskError from .worker import ComputeWorker __all__ = [ 'TaskStatus', 'QueueState', + 'UserTaskError', 'QueuedTask', 'TaskQueueConfig', 'TaskQueueMixin', diff --git a/src/twinkle/server/task_queue/backend_gate.py b/src/twinkle/server/task_queue/backend_gate.py new file mode 100644 index 000000000..82d17c1e0 --- /dev/null +++ b/src/twinkle/server/task_queue/backend_gate.py @@ -0,0 +1,121 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""The blocking-backend-call boundary and the per-replica admission gate. + +Conceptually unrelated to "task queue": this is where an event-loop coroutine hands +work to a thread and waits. It was welded onto ``TaskQueueMixin`` because +``tests/server/static/test_no_direct_backend_call.py`` forces every backend call +through ``call_backend``, which made that mixin the sole path to the backend -- one +static check binding two concepts. + +This class is a pure callable wrapper: it receives an already-bound ``fn`` from the +caller and only does ``executor.submit(functools.partial(fn, ...))``. It never writes a +``self.model.`` attribute chain and never holds a backend reference, so it is +invisible to the AST scan in ``test_no_direct_backend_call.py`` and needs no exemption +entry. +""" +from __future__ import annotations + +import asyncio +import contextlib +import functools +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from typing import Any + +from .types import BackendBusyError + + +class BackendGate: + """Owns the two executors, the admission lock, and the poison event. + + Construction must happen inside a running event loop: ``asyncio.Lock()`` and + ``asyncio.Event()`` bind to the loop that creates them. This holds today because + ``_init_task_queue`` is called from Ray Serve's ``async def __init__``. + """ + + def __init__(self, *, enable_admission_gate: bool = False) -> None: + self._executor = ThreadPoolExecutor(thread_name_prefix='twinkle-backend') + self._probe_executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix='twinkle-backend-probe') + self._admission: asyncio.Lock | None = asyncio.Lock() if enable_admission_gate else None + self._poisoned = asyncio.Event() + + async def _acquire(self, gate: asyncio.Lock) -> None: + if self._poisoned.is_set(): + raise BackendBusyError('This replica is waiting for a timed-out backend call to exit.') + if not gate.locked(): + await gate.acquire() + else: + acquire_task = asyncio.create_task(gate.acquire()) + poison_task = asyncio.create_task(self._poisoned.wait()) + try: + done, _ = await asyncio.wait((acquire_task, poison_task), return_when=asyncio.FIRST_COMPLETED) + except asyncio.CancelledError: + acquire_task.cancel() + poison_task.cancel() + await asyncio.gather(acquire_task, poison_task, return_exceptions=True) + if acquire_task.done() and not acquire_task.cancelled() and acquire_task.result(): + gate.release() + raise + if poison_task in done and self._poisoned.is_set(): + if acquire_task.done() and not acquire_task.cancelled() and acquire_task.result(): + gate.release() + else: + acquire_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await acquire_task + raise BackendBusyError('This replica is waiting for a timed-out backend call to exit.') + poison_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await poison_task + await acquire_task + if self._poisoned.is_set(): + gate.release() + raise BackendBusyError('This replica is waiting for a timed-out backend call to exit.') + + async def call(self, fn: Callable[..., Any], /, *args: Any, admit: bool = True, **kwargs: Any) -> Any: + """Run one backend call outside the event loop. + + Normal model calls serialize through the admission gate. If the awaiting task + times out while its thread is still running, the gate is poisoned: waiters fail + immediately until that thread exits. Health probes bypass the gate and use a + reserved executor thread. Sampler deployments disable the gate because their + backend owns request concurrency. + """ + loop = asyncio.get_running_loop() + gate = self._admission if admit else None + if gate is not None: + await self._acquire(gate) + + executor = self._executor if admit else self._probe_executor + try: + concurrent_future = executor.submit(functools.partial(fn, *args, **kwargs)) + except Exception: + if gate is not None and gate.locked(): + gate.release() + raise + + if gate is not None: + + def release_gate(_future) -> None: + + def release() -> None: + self._poisoned.clear() + if gate.locked(): + gate.release() + + with contextlib.suppress(RuntimeError): + loop.call_soon_threadsafe(release) + + concurrent_future.add_done_callback(release_gate) + + try: + return await asyncio.wrap_future(concurrent_future, loop=loop) + except asyncio.CancelledError: + if gate is not None and concurrent_future.running(): + self._poisoned.set() + raise + + def shutdown(self) -> None: + # Do not wait on threads that may be leaked on a timed-out backend call. + self._executor.shutdown(wait=False, cancel_futures=True) + self._probe_executor.shutdown(wait=False, cancel_futures=True) diff --git a/src/twinkle/server/task_queue/config.py b/src/twinkle/server/task_queue/config.py new file mode 100644 index 000000000..449c694cc --- /dev/null +++ b/src/twinkle/server/task_queue/config.py @@ -0,0 +1,72 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +""" +Task queue configuration. + +Provides TaskQueueConfig (Pydantic) for controlling rate limits, timeouts, +and queue behavior. Constraints are validated at construction time so an +invalid YAML/dict value is rejected before the deployment reaches a ready +state. +""" +from __future__ import annotations + +from pydantic import BaseModel, ConfigDict, Field + +# Finite bounds used when configuration omits a limit and by long-running methods. +_ZERO_EXECUTION_TIMEOUT_FALLBACK: float = 3600.0 +_MAX_DECLARED_BACKEND_TIMEOUT: float = 3600.0 +_ABSOLUTE_TTL_MULTIPLIER: int = 2 + + +class TaskQueueConfig(BaseModel): + """Configuration for task queue and rate limiting. + + Attributes: + rps_limit: Maximum requests per second per user token, per replica. ``0`` disables. + tps_limit: Maximum input tokens per second per user token, per replica. ``0`` disables. + window_seconds: Sliding window for rate-limit calculations. Must be > 0. + queue_timeout: Maximum time a task can wait in queue (seconds). + execution_timeout: Maximum time a task can execute (seconds). ``0`` means "no + configured limit"; a finite bound of 3600s is substituted instead of + unbounded waiting (see ``effective_execution_timeout``). + enabled: Whether rate limiting is enabled. + token_cleanup_multiplier: Multiplier for token cleanup threshold. + token_cleanup_interval: How often to run cleanup task (seconds). + max_input_tokens: Maximum allowed input tokens per request. + inline_fast_path_timeout: Upper bound (seconds) on how long submit briefly + polls the record so a millisecond-scale control-plane op (step / zero_grad) + completes in a single HTTP round trip instead of forcing a retrieve. + Must remain < Long_Poll_Window. + """ + + model_config = ConfigDict(extra='forbid') + + rps_limit: float = Field(default=100.0, ge=0) + tps_limit: float = Field(default=16000.0, ge=0) + window_seconds: float = Field(default=1.0, gt=0) + queue_timeout: float = Field(default=300.0, ge=0) + execution_timeout: float = Field(default=1800.0, ge=0) + enabled: bool = True + token_cleanup_multiplier: float = Field(default=10.0, ge=0) + token_cleanup_interval: float = Field(default=60.0, ge=0) + max_input_tokens: int = Field(default=16000, ge=1) + inline_fast_path_timeout: float = Field(default=0.05, gt=0) + + @property + def effective_execution_timeout(self) -> float: + """The single source of the execution time bound. + + ``0`` is not rejected (that would fail existing deployments); it is read + as "no configured limit" and replaced by a finite fallback so the bound + is always positive. This value feeds both ``_ray_get_timeout`` and the + ComputeWorker's ``asyncio.wait_for`` -- there is no second, independently + configurable timeout. + """ + if self.execution_timeout > 0: + return self.execution_timeout + return _ZERO_EXECUTION_TIMEOUT_FALLBACK + + def absolute_future_ttl(self, collect_width: int) -> float: + """Conservative lifetime for a non-terminal future record.""" + ray_timeout = max(self.effective_execution_timeout, _MAX_DECLARED_BACKEND_TIMEOUT) + resource_bound = max(1, collect_width) * ray_timeout + return _ABSOLUTE_TTL_MULTIPLIER * (self.queue_timeout + resource_bound) diff --git a/src/twinkle/server/utils/task_queue/mixin.py b/src/twinkle/server/task_queue/mixin.py similarity index 54% rename from src/twinkle/server/utils/task_queue/mixin.py rename to src/twinkle/server/task_queue/mixin.py index dcdf805ce..31bdff261 100644 --- a/src/twinkle/server/utils/task_queue/mixin.py +++ b/src/twinkle/server/task_queue/mixin.py @@ -14,9 +14,14 @@ from collections.abc import Callable, Coroutine from typing import TYPE_CHECKING, Any -from twinkle.server.telemetry.middleware import get_task_metrics -from twinkle.server.utils.task_errors import task_error_payload +from twinkle.protocol.types.lifecycle import TERMINAL_STATUSES, TaskEnvelope +from twinkle.server.exceptions import BatchSizeError, ConfigError, InputTokensExceededError, RateLimitExceededError +from twinkle.server.lifecycle.envelope import envelope_from_record +from twinkle.server.lifecycle.poll_config import long_poll_window +from twinkle.server.state.models import FutureFailureRecord +from twinkle.server.telemetry.metrics import get_task_metrics from twinkle.utils.logger import get_logger +from .backend_gate import BackendGate from .config import TaskQueueConfig from .rate_limiter import RateLimiter from .types import QueuedTask, QueueState, TaskStatus @@ -33,7 +38,7 @@ class TaskQueueMixin: Execution paths --------------- - 1. Compute queue (schedule_task / schedule_task_and_wait): + 1. Compute queue (schedule_task / submit_and_peek): Single background worker, serial execution, round-robin across queues. Use for GPU operations: forward, backward, step, save, load, etc. @@ -49,17 +54,52 @@ class TaskQueueMixin: """ state: ServerState + replica_id: str - def _init_task_queue(self, config: TaskQueueConfig | None = None, deployment_name: str = '') -> None: + def _init_task_queue( + self, + config: TaskQueueConfig | None = None, + deployment_name: str = '', + *, + enable_admission_gate: bool = False, + on_backend_timeout: Callable[[], Coroutine[Any, Any, None]] | None = None, + collect_width: int = 1, + ) -> None: """Initialise the task queue, rate limiter, and compute worker. ``config`` must be a typed :class:`TaskQueueConfig` (the launcher passes the instance straight through). ``None`` constructs a default config. + + ``enable_admission_gate`` turns on the per-replica Admission_Gate + (:meth:`call_backend`). ``ModelManagement`` enables it; ``SamplerManagement`` + does not (vllm sampler owns its own concurrency and the weight-update / + generation mutual exclusion is covered by infra ``_cw_barrier``). + + ``on_backend_timeout`` runs after a backend timeout. ``collect_width`` is the + number of actor results a backend call may collect and determines the persisted + future deadline. """ self._task_queue_config = config if config is not None else TaskQueueConfig() + if self._task_queue_config.execution_timeout == 0: + logger.warning( + '[TaskQueue] execution_timeout=0: a finite %.0fs bound has replaced unbounded waiting ' + '(deployment=%s).', self._task_queue_config.effective_execution_timeout, deployment_name or 'unknown') + # The Inline_Fast_Path window must stay strictly under Long_Poll_Window: a + # submit that peeks longer than a retrieve would wait makes no sense (D4). + # Raised, not asserted -- `python -O` strips asserts and would drop this + # invariant silently. + _inline = self._task_queue_config.inline_fast_path_timeout + _window = long_poll_window() + if _inline >= _window: + raise ConfigError( + 'inline_fast_path_timeout', + _inline, + message=(f'inline_fast_path_timeout ({_inline}s) must be < Long_Poll_Window ' + f'({_window}s); lower it or raise TWINKLE_LONG_POLL_TIMEOUT.')) self._deployment_name = deployment_name self._task_metrics = get_task_metrics(deployment_name) if deployment_name else None + self._future_absolute_ttl = self._task_queue_config.absolute_future_ttl(collect_width) self._rate_limiter = RateLimiter( rps_limit=self._task_queue_config.rps_limit, @@ -69,6 +109,7 @@ def _init_task_queue(self, config: TaskQueueConfig | None = None, deployment_nam token_cleanup_interval=self._task_queue_config.token_cleanup_interval, active_tokens_gauge=self._task_metrics.rate_limiter_active_tokens if self._task_metrics else None, deployment_name=deployment_name, + replica_id=self.replica_id, ) self._rate_limiter.start_cleanup_task() @@ -77,10 +118,36 @@ def _init_task_queue(self, config: TaskQueueConfig | None = None, deployment_nam config=self._task_queue_config, task_metrics=self._task_metrics, deployment_name=deployment_name, + on_backend_timeout=on_backend_timeout, ) + self._backend_gate = BackendGate(enable_admission_gate=enable_admission_gate) self._event_loop: asyncio.AbstractEventLoop | None = None + @property + def task_queue_config(self) -> TaskQueueConfig: + """The deployment's validated queue config. + + Public because four modules outside this file read it; it was + ``_task_queue_config`` (name-private, interface-public), which meant the + Host_Protocol could not declare it without declaring a private name. + """ + return self._task_queue_config + + async def call_backend(self, fn: Callable[..., Any], /, *args: Any, admit: bool = True, **kwargs: Any) -> Any: + """Run one backend call outside the event loop. + + Single delegation to :class:`BackendGate`. The name and call shape are + load-bearing: ``test_no_direct_backend_call.py`` keys on ``self.call_backend`` + call sites, so keeping it here means the check's predicate and its exemption + table need no change. + """ + return await self._backend_gate.call(fn, *args, admit=admit, **kwargs) + + def _future_deadline(self) -> float: + ttl = getattr(self, '_future_absolute_ttl', self._task_queue_config.absolute_future_ttl(1)) + return time.time() + ttl + @staticmethod def _queue_key(model_id: str | None, token: str | None) -> str: if model_id: @@ -91,67 +158,43 @@ def _queue_key(model_id: str | None, token: str | None) -> str: async def _perform_preflight_checks( self, - request_id: str, model_id: str | None, token: str | None, input_tokens: int, batch_size: int | None = None, data_world_size: int | None = None, batch_size_multiple: int | None = None, - persist_failure: bool = True, - ) -> dict[str, Any] | None: + ) -> None: """Run rate-limit and validation checks before queuing a task. - Returns None if all checks pass, or an error-response dict on failure. + Returns ``None`` when every check passes. On failure it RAISES a + ``RequestRejectedError`` subclass -- the Decision_Boundary is this line, and + raising before any ``store_future_status`` call is what guarantees zero + future writes for a rejected request. It writes no FAILED + record and returns no ``_error`` marker. """ if not token or not self._task_queue_config.enabled: - return None - - async def reject(error_msg: str, queue_state: str) -> dict[str, Any]: - error_payload = {'error': error_msg, 'category': 'User'} - if persist_failure: - await self.state.store_future_status( - request_id, - TaskStatus.FAILED.value, - model_id, - result=error_payload, - queue_state=queue_state, - queue_state_reason=error_msg, - ) - return {'request_id': request_id, 'model_id': model_id} - # Private marker consumed by schedule_task_and_wait(). It is not - # returned by the public polling-style schedule_task() API. - return { - 'request_id': request_id, - 'model_id': model_id, - '_error': error_msg, - } + return if input_tokens > self._task_queue_config.max_input_tokens: - error_msg = (f'Input tokens ({input_tokens}) exceed maximum allowed ' - f'({self._task_queue_config.max_input_tokens})') - return await reject(error_msg, QueueState.UNKNOWN.value) + raise InputTokensExceededError(f'Input tokens ({input_tokens}) exceed maximum allowed ' + f'({self._task_queue_config.max_input_tokens})') if batch_size is not None and data_world_size is not None: if batch_size < data_world_size: - error_msg = (f'Batch size {batch_size} must be >= data world size {data_world_size}') - return await reject(error_msg, QueueState.UNKNOWN.value) + raise BatchSizeError(f'Batch size {batch_size} must be >= data world size {data_world_size}') if batch_size_multiple is not None: required_multiple = data_world_size * batch_size_multiple if batch_size % required_multiple != 0: - error_msg = (f'Batch size {batch_size} must be divisible by {required_multiple} ' - f'so each data-parallel shard gets a multiple of ' - f'{batch_size_multiple} examples') - return await reject(error_msg, QueueState.UNKNOWN.value) + raise BatchSizeError(f'Batch size {batch_size} must be divisible by {required_multiple} ' + f'so each data-parallel shard gets a multiple of ' + f'{batch_size_multiple} examples') allowed, reason = await self._rate_limiter.check_and_record(token, input_tokens) if not allowed: if self._task_metrics: self._task_metrics.rate_limit_rejections.inc(tags={'deployment': self._deployment_name}) - error_msg = f'Rate limit exceeded: {reason}' - return await reject(error_msg, QueueState.PAUSED_RATE_LIMIT.value) - - return None + raise RateLimitExceededError(f'Rate limit exceeded: {reason}') async def _schedule_task( self, @@ -163,42 +206,39 @@ async def _schedule_task( data_world_size: int | None = None, batch_size_multiple: int | None = None, task_type: str | None = None, - *, - completion: asyncio.Future[Any] | None = None, - persist_status: bool, + request_id: str | None = None, ) -> dict[str, Any]: - """Common enqueue path for polling and in-process wait callers.""" - request_id = f'req_{uuid.uuid4().hex}' + """Common enqueue path. Always persists status: the future record is the + single delivery channel for both result and failure.""" + request_id = request_id or f'req_{uuid.uuid4().hex}' - preflight_result = await self._perform_preflight_checks( - request_id=request_id, + # Decision_Boundary: raises RequestRejectedError before any state write. + await self._perform_preflight_checks( model_id=model_id, token=token, input_tokens=input_tokens, batch_size=batch_size, data_world_size=data_world_size, batch_size_multiple=batch_size_multiple, - persist_failure=persist_status, ) - if preflight_result is not None: - return preflight_result if self._event_loop is None: self._event_loop = asyncio.get_running_loop() - if persist_status: - await self.state.store_future_status( - request_id, - TaskStatus.PENDING.value, - model_id, - queue_state=QueueState.ACTIVE.value, - ) + await self.state.store_future_status( + request_id, + TaskStatus.PENDING.value, + model_id, + queue_state=QueueState.ACTIVE.value, + replica_id=getattr(self, 'replica_id', None), + absolute_deadline=self._future_deadline(), + ) queue_key = self._queue_key(model_id=model_id, token=token) self._compute_worker.ensure_queue_registered(queue_key) await self._compute_worker.ensure_started() - q = self._compute_worker.task_queues[queue_key] + q = self._compute_worker.get_queue(queue_key) await q.put( QueuedTask( request_id=request_id, @@ -208,25 +248,18 @@ async def _schedule_task( input_tokens=input_tokens, task_type=task_type, created_at=time.monotonic(), - completion=completion, - persist_status=persist_status, )) - if persist_status: - await self.state.store_future_status( - request_id, - TaskStatus.QUEUED.value, - model_id, - queue_state=QueueState.ACTIVE.value, - ) + await self.state.store_future_status( + request_id, + TaskStatus.QUEUED.value, + model_id, + queue_state=QueueState.ACTIVE.value, + ) logger.info(f'[TaskQueue] Task {request_id} queued, type={task_type or "unknown"}, ' f'model_id={model_id}, queue_key={queue_key}, ' f'queue_depth={q.qsize()}, input_tokens={input_tokens}') - self._compute_worker.new_task_event.set() - - if self._task_metrics: - total_depth = sum(q.qsize() for q in self._compute_worker.task_queues.values()) - self._task_metrics.queue_depth.set(total_depth, tags={'deployment': self._deployment_name}) + self._compute_worker.notify_new_task() return {'request_id': request_id, 'model_id': model_id} @@ -240,6 +273,7 @@ async def schedule_task( data_world_size: int | None = None, batch_size_multiple: int | None = None, task_type: str | None = None, + request_id: str | None = None, ) -> dict[str, Any]: """Schedule a GPU compute task through the serial compute queue. @@ -268,46 +302,66 @@ async def schedule_task( data_world_size=data_world_size, batch_size_multiple=batch_size_multiple, task_type=task_type, - persist_status=True, + request_id=request_id, ) - async def schedule_task_and_wait( + # Poll cadence *inside* the Inline_Fast_Path window. Much smaller than the window + # itself so a task that finishes early is noticed promptly; the window bound + # (config.inline_fast_path_timeout) is what actually caps submit latency. + _INLINE_FAST_PATH_POLL = 0.005 + + async def _peek_terminal(self, request_id: str, *, fallback_status: str) -> TaskEnvelope: + """Poll a record up to the Inline_Fast_Path window; return a terminal envelope + if it settled, else a non-terminal envelope with ``fallback_status``.""" + deadline = time.monotonic() + self._task_queue_config.inline_fast_path_timeout + record = None + while time.monotonic() < deadline: + record = await self.state.get_future(request_id) + if record is not None and record.get('status') in TERMINAL_STATUSES: + return envelope_from_record(request_id, record) + await asyncio.sleep(self._INLINE_FAST_PATH_POLL) + return envelope_from_record(request_id, record, fallback_status=fallback_status) + + async def submit_and_peek( self, coro_factory: Callable[[], Coroutine], + *, model_id: str | None = None, token: str | None = None, - input_tokens: int = 0, - batch_size: int | None = None, - data_world_size: int | None = None, - batch_size_multiple: int | None = None, task_type: str | None = None, - ) -> Any: - """Schedule a compute task and block until it completes. + request_id: str | None = None, + **schedule_kwargs: Any, + ) -> TaskEnvelope: + """Enqueue a task, then briefly wait so a fast op finishes in one round trip. + + Exceeding the window is not an error: the caller polls + ``/twinkle/retrieve_future`` instead. That is what makes this loop + fundamentally different from the deleted in-process blocking wait -- it owes + nothing to failure handling, so it needs no terminal write, no missing-record + branch, and no race with the worker. A terminal record inside the window is + returned as a terminal envelope (success or failure); a window that elapses + still non-terminal returns a non-terminal envelope carrying ``queue_state``. + """ + ref = await self.schedule_task( + coro_factory, model_id=model_id, token=token, task_type=task_type, request_id=request_id, **schedule_kwargs) + return await self._peek_terminal(ref['request_id'], fallback_status='pending') - Twinkle-side counterpart to schedule_task(). Enqueues the task through - the same serial worker but delivers the result through an in-process - Future. Large model outputs therefore never enter ServerState. + async def submit_background_and_peek( + self, + coro_factory: Callable[[], Coroutine], + *, + model_id: str | None = None, + task_type: str | None = None, + ) -> TaskEnvelope: + """Fire-and-forget variant of :meth:`submit_and_peek` for pure-I/O tasks. - Raises: - RuntimeError: If the task fails or scheduling is rejected. + Uses ``schedule_background_task`` (outside the serial compute queue) but returns + the same Task_Envelope, so ``upload_to_hub`` shares the retrieve/future machinery + instead of its own status endpoint. The task is already RUNNING on return, so the + non-terminal fallback is ``running``. """ - completion = asyncio.get_running_loop().create_future() - task_ref = await self._schedule_task( - coro_factory, - model_id=model_id, - token=token, - input_tokens=input_tokens, - batch_size=batch_size, - data_world_size=data_world_size, - batch_size_multiple=batch_size_multiple, - task_type=task_type, - completion=completion, - persist_status=False, - ) - if error := task_ref.get('_error'): - completion.cancel() - raise RuntimeError(error) - return await completion + ref = await self.schedule_background_task(coro_factory, model_id=model_id, task_type=task_type) + return await self._peek_terminal(ref['request_id'], fallback_status='running') async def schedule_background_task( self, @@ -342,6 +396,8 @@ async def schedule_background_task( TaskStatus.RUNNING.value, model_id, queue_state=QueueState.ACTIVE.value, + replica_id=getattr(self, 'replica_id', None), + absolute_deadline=self._future_deadline(), ) async def _run() -> None: @@ -355,13 +411,18 @@ async def _run() -> None: queue_state=QueueState.ACTIVE.value, ) logger.info(f'[TaskQueue] Background task {request_id} completed, type={task_type or "unknown"}') - except Exception: - error_payload = task_error_payload(traceback.format_exc()) + except Exception as exc: + failure = FutureFailureRecord( + reason_code='internal_error', + message=f'{type(exc).__name__}: {exc}'[:1024], + attribution='server', + diagnostic=traceback.format_exc(), + ) await self.state.store_future_status( request_id, TaskStatus.FAILED.value, model_id, - result=error_payload, + failure=failure, queue_state=QueueState.ACTIVE.value, ) logger.error(f'[TaskQueue] Background task {request_id} FAILED, type={task_type or "unknown"}:\n' @@ -385,32 +446,9 @@ def _schedule() -> None: self._event_loop.call_soon_threadsafe(_schedule) - def get_queue_stats(self) -> dict[str, Any]: - """Return current compute queue statistics.""" - return { - 'queue_size': - sum(q.qsize() for q in self._compute_worker.task_queues.values()), - 'queue_count': - len(self._compute_worker.task_queues), - 'worker_running': (self._compute_worker._worker_task is not None - and not self._compute_worker._worker_task.done()), - 'rate_limit_config': { - 'rps_limit': self._task_queue_config.rps_limit, - 'tps_limit': self._task_queue_config.tps_limit, - 'enabled': self._task_queue_config.enabled, - }, - } - - def get_rate_limit_stats(self, token: str) -> dict[str, Any]: - """Return rate-limiting stats for a user token.""" - return self._rate_limiter.get_stats(token) - - def get_rate_limiter_memory_stats(self) -> dict[str, Any]: - """Return memory usage statistics from the rate limiter.""" - return self._rate_limiter.get_memory_stats() - async def shutdown_task_queue(self) -> None: """Gracefully shut down the compute queue and release resources.""" await self._rate_limiter.stop_cleanup_task() await self._compute_worker.stop() + self._backend_gate.shutdown() logger.debug('[TaskQueue] Task queue shutdown complete') diff --git a/src/twinkle/server/utils/task_queue/rate_limiter.py b/src/twinkle/server/task_queue/rate_limiter.py similarity index 80% rename from src/twinkle/server/utils/task_queue/rate_limiter.py rename to src/twinkle/server/task_queue/rate_limiter.py index 8f1dbe97e..eebb72a3a 100644 --- a/src/twinkle/server/utils/task_queue/rate_limiter.py +++ b/src/twinkle/server/task_queue/rate_limiter.py @@ -10,7 +10,6 @@ import asyncio import time -from typing import Any from twinkle.utils.logger import get_logger @@ -43,21 +42,25 @@ def __init__( token_cleanup_interval: float = 60.0, active_tokens_gauge=None, deployment_name: str = '', + replica_id: str = '', ): """Initialize the rate limiter. Args: - rps_limit: Maximum requests per second per user token. - tps_limit: Maximum input tokens per second per user token. + rps_limit: Maximum requests per second per user token, per replica. + ``0`` disables the RPS limit. + tps_limit: Maximum input tokens per second per user token, per replica. + ``0`` disables the TPS limit. window_seconds: Time window for rate limiting (default 1.0s). token_cleanup_multiplier: Multiplier for token cleanup threshold. Tokens inactive for window_seconds * token_cleanup_multiplier will be removed. Default is 10.0 (10x the window). token_cleanup_interval: How often to run the cleanup task in seconds. Default is 60.0 (every minute). - active_tokens_gauge: Optional gauge adapter (see twinkle.server.telemetry.middleware) + active_tokens_gauge: Optional gauge adapter (see twinkle.server.telemetry.metrics) for tracking the active token count. deployment_name: Deployment name for metrics labels. + replica_id: Replica identifier for metrics labels. """ self.rps_limit = rps_limit self.tps_limit = tps_limit @@ -79,7 +82,11 @@ def __init__( # Metrics gauge for active token count self._active_tokens_gauge = active_tokens_gauge - self._deployment_name = deployment_name + self._metric_tags = {} + if deployment_name: + self._metric_tags['deployment'] = deployment_name + if replica_id: + self._metric_tags['replica'] = replica_id def _cleanup_old_requests(self, token: str, current_time: float) -> None: """Remove requests outside the sliding window.""" @@ -120,8 +127,7 @@ async def _cleanup_inactive_tokens(self) -> None: f'Active tokens remaining: {len(self._token_requests)}') if self._active_tokens_gauge is not None: - tags = {'deployment': self._deployment_name} if self._deployment_name else {} - self._active_tokens_gauge.set(len(self._token_requests), tags=tags) + self._active_tokens_gauge.set(len(self._token_requests), tags=self._metric_tags) except asyncio.CancelledError: logger.debug('[RateLimiter] Cleanup task cancelled') @@ -165,45 +171,13 @@ async def check_and_record(self, token: str, input_tokens: int) -> tuple[bool, s request_count = len(requests) token_count = sum(count for _, count in requests) - if request_count >= self.rps_limit: + if self.rps_limit > 0 and request_count >= self.rps_limit: return False, f'RPS limit exceeded: {request_count}/{self.rps_limit} requests/s' - if token_count + input_tokens > self.tps_limit: + if self.tps_limit > 0 and token_count + input_tokens > self.tps_limit: return False, f'TPS limit exceeded: {token_count + input_tokens}/{self.tps_limit} tokens/s' self._token_requests[token].append((current_time, input_tokens)) if self._active_tokens_gauge is not None: - tags = {'deployment': self._deployment_name} if self._deployment_name else {} - self._active_tokens_gauge.set(len(self._token_requests), tags=tags) + self._active_tokens_gauge.set(len(self._token_requests), tags=self._metric_tags) return True, None - - def get_stats(self, token: str) -> dict[str, Any]: - """Get current rate limiting stats for a token.""" - current_time = time.time() - self._cleanup_old_requests(token, current_time) - - if token in self._token_requests: - self._last_activity[token] = current_time - - requests = self._token_requests.get(token, []) - request_count = len(requests) - token_count = sum(count for _, count in requests) - - return { - 'current_rps': request_count, - 'current_tps': token_count, - 'rps_limit': self.rps_limit, - 'tps_limit': self.tps_limit, - 'rps_available': self.rps_limit - request_count, - 'tps_available': self.tps_limit - token_count, - } - - def get_memory_stats(self) -> dict[str, Any]: - """Get memory usage statistics for monitoring.""" - return { - 'active_tokens': len(self._token_requests), - 'tracked_tokens': len(self._last_activity), - 'cleanup_threshold_seconds': self.window_seconds * self.token_cleanup_multiplier, - 'cleanup_interval_seconds': self.token_cleanup_interval, - 'cleanup_task_running': self._cleanup_started and self._cleanup_task and not self._cleanup_task.done(), - } diff --git a/src/twinkle/server/utils/task_queue/types.py b/src/twinkle/server/task_queue/types.py similarity index 71% rename from src/twinkle/server/utils/task_queue/types.py rename to src/twinkle/server/task_queue/types.py index daf8d2bb2..9085ff8ec 100644 --- a/src/twinkle/server/utils/task_queue/types.py +++ b/src/twinkle/server/task_queue/types.py @@ -9,11 +9,9 @@ """ from __future__ import annotations -import asyncio from collections.abc import Callable, Coroutine from dataclasses import dataclass from enum import Enum -from typing import Any class TaskStatus(Enum): @@ -23,7 +21,21 @@ class TaskStatus(Enum): RUNNING = 'running' # Task currently executing COMPLETED = 'completed' # Task completed successfully FAILED = 'failed' # Task failed with error - RATE_LIMITED = 'rate_limited' # Task rejected due to rate limiting + CANCELLED = 'cancelled' # Task cancelled by the client before it started running + + +class UserTaskError(ValueError): + """A queued operation rejected because of caller input or usage.""" + + +class BackendBusyError(RuntimeError): + """Raised when the per-replica Admission_Gate is held by a leaked backend call. + + A new backend call arriving while the gate is closed (its holder is a call that + already exceeded ``asyncio.wait_for`` but whose executor thread has not yet + returned) fails fast with this error instead of queueing behind it. The worker + maps it to ``ErrorPayload(category='server', error_code=503)``. + """ class QueueState(Enum): @@ -48,10 +60,3 @@ class QueuedTask: input_tokens: int task_type: str | None created_at: float - first_rate_limited_at: float | None = None - # ``schedule_task_and_wait`` is an in-process request/response path. Its - # potentially large result is delivered through this Future instead of - # being persisted in ServerState merely for the same process to read it - # back. Polling-style ``schedule_task`` leaves this as ``None``. - completion: asyncio.Future[Any] | None = None - persist_status: bool = True diff --git a/src/twinkle/server/utils/task_queue/worker.py b/src/twinkle/server/task_queue/worker.py similarity index 57% rename from src/twinkle/server/utils/task_queue/worker.py rename to src/twinkle/server/task_queue/worker.py index fdbb36d16..aaa9a6b90 100644 --- a/src/twinkle/server/utils/task_queue/worker.py +++ b/src/twinkle/server/task_queue/worker.py @@ -12,22 +12,57 @@ import time import traceback from collections import deque -from typing import TYPE_CHECKING, Any, Deque - +from typing import TYPE_CHECKING, Any, Callable, Deque + +from twinkle.protocol.types.errors import ErrorCategory +from twinkle.server.exceptions import (BatchSizeError, EndpointUnavailableError, FullModeBusyError, + InputTokensExceededError, RateLimitExceededError, RequestRejectedError, + ResourceNotFoundError, ResourceQuotaExceededError, StateBackendError, + TrainModeMismatchError, TwinkleServerError) +from twinkle.server.state.backend.base import ConcurrencyError +from twinkle.server.state.models import FutureFailureRecord from twinkle.server.telemetry.correlation import MODEL_ID, TOKEN_ID from twinkle.server.telemetry.tracing import traced_operation -from twinkle.server.utils.task_errors import task_error_payload from twinkle.utils.logger import get_logger from .config import TaskQueueConfig -from .types import QueuedTask, QueueState, TaskStatus +from .types import BackendBusyError, QueuedTask, QueueState, TaskStatus, UserTaskError if TYPE_CHECKING: from twinkle.server.state import ServerState - from twinkle.server.telemetry.middleware import TaskMetrics + from twinkle.server.telemetry.metrics import TaskMetrics logger = get_logger() +def _reason_code_for_server_error(exc: TwinkleServerError) -> str: + """Map typed execution failures to protocol-independent reason codes.""" + mappings: tuple[tuple[type[TwinkleServerError], str], ...] = ( + (TrainModeMismatchError, 'invalid_request'), + (ResourceNotFoundError, 'resource_not_found'), + (FullModeBusyError, 'full_mode_busy'), + (InputTokensExceededError, 'input_tokens_exceeded'), + (BatchSizeError, 'batch_size_invalid'), + (RateLimitExceededError, 'rate_limit_exceeded'), + (ResourceQuotaExceededError, 'resource_quota_exceeded'), + (EndpointUnavailableError, 'endpoint_unavailable'), + (ConcurrencyError, 'state_contention'), + (StateBackendError, 'backend_gate_unavailable'), + (RequestRejectedError, 'request_rejected'), + ) + for error_type, reason_code in mappings: + if isinstance(exc, error_type): + return reason_code + return 'internal_error' + + +# Ray_Get_Timeout is classified the same as asyncio.TimeoutError: 504/Server. +try: + from ray.exceptions import GetTimeoutError as _RayGetTimeout + _TIMEOUT_EXCEPTIONS: tuple[type[BaseException], ...] = (asyncio.TimeoutError, _RayGetTimeout) +except Exception: # pragma: no cover - ray always present in server runtime + _TIMEOUT_EXCEPTIONS = (asyncio.TimeoutError, ) + + class ComputeWorker: """Serial background worker that processes GPU compute tasks. @@ -46,11 +81,14 @@ def __init__( config: TaskQueueConfig, task_metrics: TaskMetrics | None, deployment_name: str, + on_backend_timeout: Callable[[], Any] | None = None, ) -> None: self._state = state self._config = config self._task_metrics = task_metrics self._deployment_name = deployment_name + # Optional coroutine-returning callback fired on a backend timeout. + self._on_backend_timeout = on_backend_timeout self.task_queues: dict[str, asyncio.Queue] = {} self.queue_order: Deque[str] = deque() @@ -90,6 +128,29 @@ def ensure_queue_registered(self, queue_key: str) -> None: if queue_key not in self.queue_order: self.queue_order.append(queue_key) + def get_queue(self, queue_key: str) -> asyncio.Queue: + """Return the registered queue for ``queue_key``. + + Callers must have registered it first via :meth:`ensure_queue_registered`. + Exposed so producers (the mixin) enqueue through a method instead of + indexing the worker's internal ``task_queues`` container directly. + """ + return self.task_queues[queue_key] + + def total_queued(self) -> int: + """Total number of pending tasks across all per-key queues.""" + return sum(q.qsize() for q in self.task_queues.values()) + + def notify_new_task(self) -> None: + """Wake the worker loop and refresh the queue-depth gauge. + + Producers (the mixin) call this instead of reaching into the worker's + ``new_task_event`` directly; it is also the single write point of the + queue-depth gauge, so the gauge has one writer rather than two. + """ + self.new_task_event.set() + self._record_queue_depth() + # ------------------------------------------------------------------ # Metrics helpers # ------------------------------------------------------------------ @@ -112,6 +173,12 @@ def _record_execution_time(self, task_type: str, exec_time: float) -> None: 'task_type': task_type, }) + def _record_queue_depth(self) -> None: + """Single writer of the queue-depth gauge.""" + if self._task_metrics: + total_depth = sum(qq.qsize() for qq in self.task_queues.values()) + self._task_metrics.queue_depth.set(total_depth, tags={'deployment': self._deployment_name}) + def _record_queue_metrics(self, task_type: str, queue_wait: float) -> None: """Observe queue wait time and update current queue depth if metrics are enabled.""" if self._task_metrics: @@ -120,39 +187,36 @@ def _record_queue_metrics(self, task_type: str, queue_wait: float) -> None: 'deployment': self._deployment_name, 'task_type': task_type, }) - total_depth = sum(qq.qsize() for qq in self.task_queues.values()) - self._task_metrics.queue_depth.set(total_depth, tags={'deployment': self._deployment_name}) + self._record_queue_depth() # ------------------------------------------------------------------ - @staticmethod - def _complete_result(task: QueuedTask, result: Any) -> None: - if task.completion is not None and not task.completion.done(): - task.completion.set_result(result) - - @staticmethod - def _complete_error(task: QueuedTask, error: str) -> None: - if task.completion is not None and not task.completion.done(): - task.completion.set_exception(RuntimeError(error)) - async def _store_task_failed( self, task: QueuedTask, error: str, queue_state: str, queue_state_reason: str | None = None, + *, + reason_code: str = 'internal_error', + attribution: str = 'server', + traceback_text: str | None = None, ) -> None: - """Store FAILED status with a standardised error payload.""" - if task.persist_status: - await self._state.store_future_status( - task.request_id, - TaskStatus.FAILED.value, - task.model_id, - result=task_error_payload(error), - queue_state=queue_state, - queue_state_reason=queue_state_reason, - ) - self._complete_error(task, error) + """Store FAILED status with protocol-independent failure details.""" + failure = FutureFailureRecord( + reason_code=reason_code, + message=(error.splitlines() or [''])[0][:1024], + attribution='user' if attribution == 'user' else 'server', + diagnostic=traceback_text, + ) + await self._state.store_future_status( + task.request_id, + TaskStatus.FAILED.value, + task.model_id, + failure=failure, + queue_state=queue_state, + queue_state_reason=queue_state_reason, + ) async def fail_queue_tasks(self, queue_key: str, reason: str) -> None: """Drain a queue and mark all pending tasks as FAILED.""" @@ -203,13 +267,12 @@ async def _execute_task(self, task: QueuedTask, queue_key: str, q: asyncio.Queue Handles execution timeout, general exceptions, and always calls q.task_done() in the finally block. """ - if task.persist_status: - await self._state.store_future_status( - task.request_id, - TaskStatus.RUNNING.value, - task.model_id, - queue_state=QueueState.ACTIVE.value, - ) + await self._state.store_future_status( + task.request_id, + TaskStatus.RUNNING.value, + task.model_id, + queue_state=QueueState.ACTIVE.value, + ) task_type = task.task_type or 'unknown' exec_start = time.monotonic() @@ -232,36 +295,87 @@ async def _execute_task(self, task: QueuedTask, queue_key: str, q: asyncio.Queue f'type={task_type}, queue_key={queue_key}') with traced_operation(handler_span_name, attrs=handler_attrs): coro = task.coro_factory() - if self._config.execution_timeout > 0: - result = await asyncio.wait_for(coro, timeout=self._config.execution_timeout) - else: - result = await coro + # effective_execution_timeout is always positive (0 -> finite fallback), + # so wait_for is always in effect. + result = await asyncio.wait_for(coro, timeout=self._config.effective_execution_timeout) exec_time = time.monotonic() - exec_start logger.info(f'[ComputeWorker] Task {task.request_id} completed in {exec_time:.2f}s, type={task_type}') - if task.persist_status: - await self._state.store_future_status( - task.request_id, - TaskStatus.COMPLETED.value, - task.model_id, - result=result, - queue_state=QueueState.ACTIVE.value, - ) - self._complete_result(task, result) - except asyncio.TimeoutError: + await self._state.store_future_status( + task.request_id, + TaskStatus.COMPLETED.value, + task.model_id, + result=result, + queue_state=QueueState.ACTIVE.value, + ) + except _TIMEOUT_EXCEPTIONS: task_status = 'timeout' exec_time = time.monotonic() - exec_start - error = (f'Execution timeout exceeded: {self._config.execution_timeout}s, ' + error = (f'Backend call timed out (bound {self._config.effective_execution_timeout}s), ' f'actual execution time: {exec_time:.2f}s') logger.error(f'[ComputeWorker] Task {task.request_id} TIMEOUT after {exec_time:.2f}s, ' f'type={task_type}, queue_key={queue_key}') - await self._store_task_failed(task, error, QueueState.ACTIVE.value) - except Exception: + # asyncio.TimeoutError and Ray_Get_Timeout are 504/Server. + await self._store_task_failed(task, error, QueueState.ACTIVE.value, reason_code='execution_timeout') + # Probe actor liveness after a timeout so an operator learns the replica's + # state without waiting for a second request to also time out. + if self._on_backend_timeout is not None: + try: + await self._on_backend_timeout() + except Exception: + logger.error(f'[ComputeWorker] backend-timeout probe failed:\n{traceback.format_exc(limit=3)}') + except UserTaskError as exc: + task_status = 'failed' + exec_time = time.monotonic() - exec_start + await self._store_task_failed( + task, + f'{type(exc).__name__}: {exc}', + QueueState.UNKNOWN.value, + reason_code='request_rejected', + attribution='user', + ) + except BackendBusyError as exc: task_status = 'failed' exec_time = time.monotonic() - exec_start - error = traceback.format_exc() + error = str(exc) + logger.error(f'[ComputeWorker] Task {task.request_id} REFUSED (admission gate held) after ' + f'{exec_time:.2f}s, type={task_type}, queue_key={queue_key}') + # Gate held by a leaked timed-out call -> 503/Server. + await self._store_task_failed(task, error, QueueState.ACTIVE.value, reason_code='backend_gate_unavailable') + except TwinkleServerError as exc: + # A typed server error carries its own status + category (e.g. + # ResourceNotFoundError = 404/User from a deferred + # assert_resource_exists). Honour them instead of collapsing every + # such failure to 500/Server. Only Server-category errors keep a + # traceback; a User rejection does not. + task_status = 'failed' + exec_time = time.monotonic() - exec_start + is_server = exc.category is ErrorCategory.Server + if is_server: + logger.error(f'[ComputeWorker] Task {task.request_id} FAILED after {exec_time:.2f}s, ' + f'type={task_type}:\n{traceback.format_exc(limit=3)}') + await self._store_task_failed( + task, + f'{type(exc).__name__}: {exc}', + QueueState.ACTIVE.value, + reason_code=_reason_code_for_server_error(exc), + attribution=exc.category.value, + traceback_text=traceback.format_exc() if is_server else None, + ) + except Exception as exc: + task_status = 'failed' + exec_time = time.monotonic() - exec_start + # error is a single-line summary; the full traceback goes only to the + # traceback field, never into `error`. + error = f'{type(exc).__name__}: {exc}' logger.error(f'[ComputeWorker] Task {task.request_id} FAILED after {exec_time:.2f}s, ' f'type={task_type}:\n{traceback.format_exc(limit=3)}') - await self._store_task_failed(task, error, QueueState.ACTIVE.value) + await self._store_task_failed( + task, + error, + QueueState.ACTIVE.value, + reason_code='internal_error', + traceback_text=traceback.format_exc(), + ) finally: q.task_done() self._record_execution_time(task_type, exec_time) @@ -296,12 +410,29 @@ async def _try_run_one(self) -> bool: await self._fail_timed_out_task(task, queue_wait, q) continue # try the next queue + # A record already in a Terminal_State (e.g. written 'failed' by the + # state-hygiene orphan handling) must not be executed again. + if await self._is_record_terminal(task.request_id): + logger.info(f'[ComputeWorker] Task {task.request_id} already terminal on dequeue; skipping.') + q.task_done() + continue + # Execute the task (serial: stops after the first execution) await self._execute_task(task, queue_key, q) return True return False + async def _is_record_terminal(self, request_id: str) -> bool: + """True if the future record already holds a Terminal_State.""" + try: + record = await self._state.get_future(request_id) + except Exception: + return False + if not record: + return False + return record.get('status') in (TaskStatus.COMPLETED.value, TaskStatus.FAILED.value, TaskStatus.CANCELLED.value) + # ------------------------------------------------------------------ # Main worker loop # ------------------------------------------------------------------ diff --git a/src/twinkle/server/telemetry/__init__.py b/src/twinkle/server/telemetry/__init__.py index 21803431f..ded3157e4 100644 --- a/src/twinkle/server/telemetry/__init__.py +++ b/src/twinkle/server/telemetry/__init__.py @@ -1,4 +1,3 @@ -from twinkle.server.config.telemetry import TelemetryConfig from .metrics import MetricsRegistry from .provider import get_meter, init_telemetry, shutdown_telemetry from .tracing import extract_context, get_current_span, get_tracer, inject_context @@ -6,7 +5,6 @@ __all__ = [ 'MetricsRegistry', - 'TelemetryConfig', 'get_meter', 'init_telemetry', 'shutdown_telemetry', diff --git a/src/twinkle/server/telemetry/http_middleware.py b/src/twinkle/server/telemetry/http_middleware.py new file mode 100644 index 000000000..beeee4e51 --- /dev/null +++ b/src/twinkle/server/telemetry/http_middleware.py @@ -0,0 +1,56 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""FastAPI HTTP request-metrics middleware. + +Split out of the former ``middleware.py``: the metric adapters and +containers live in ``metrics.py`` next to the ``MetricsRegistry``; this file holds only +the HTTP middleware factory, so a file named for HTTP middleware contains HTTP middleware +and the queue code no longer imports a module called ``middleware`` just to reach +``get_task_metrics``. +""" +from __future__ import annotations + +import time +from collections.abc import Callable +from typing import Any + +from twinkle.server.telemetry.metrics import MetricsRegistry + + +def create_metrics_middleware(deployment: str) -> Callable: + """Return a FastAPI ``http`` middleware that records request metrics. + + Usage inside a ``build_*_app()`` function:: + + from twinkle.server.telemetry.http_middleware import create_metrics_middleware + from twinkle.server.telemetry.tracing import create_tracing_middleware + + app.middleware('http')(verify_token) + app.middleware('http')(create_tracing_middleware("Model")) + app.middleware('http')(create_metrics_middleware("Model")) # outermost + + FastAPI executes middleware in LIFO order, so the **last** middleware + registered is the outermost wrapper. Register metrics last so its + latency observation covers the full request path including tracing + overhead and authentication. + """ + + async def metrics_middleware(request: Any, call_next: Callable) -> Any: + start = time.monotonic() + response = await call_next(request) + elapsed = time.monotonic() - start + status = str(response.status_code) + method = request.scope['route'].path if 'route' in request.scope else request.url.path + m = MetricsRegistry.get().request_metrics(deployment) + m.requests_total.inc(tags={ + 'deployment': deployment, + 'method': method, + 'status': status, + }) + m.request_duration_seconds.observe( + elapsed, tags={ + 'deployment': deployment, + 'method': method, + }) + return response + + return metrics_middleware diff --git a/src/twinkle/server/telemetry/metrics.py b/src/twinkle/server/telemetry/metrics.py index 20518d363..586277708 100644 --- a/src/twinkle/server/telemetry/metrics.py +++ b/src/twinkle/server/telemetry/metrics.py @@ -1,7 +1,25 @@ -"""Twinkle Server metrics registry — low-invasiveness facade over OpenTelemetry metrics.""" +"""Twinkle Server metrics registry — low-invasiveness facade over OpenTelemetry metrics. + +Besides the :class:`MetricsRegistry` that declares the raw OTEL instruments, this module +holds the legacy-API adapter classes (``_Counter`` / ``_Histogram`` / ``_Gauge``), the +structured containers (:class:`TaskMetrics` / ``_RequestMetrics``) and ``get_task_metrics``. +Keeping the metric types and the registry in one file (with the HTTP middleware factory in +``http_middleware.py``) avoids a ``metrics <-> adapters`` import cycle. + +Per-deployment adapters are cached on the *registry instance* +(``_task_metrics`` / ``_request_metrics``), not at module level. That is load-bearing: the +adapters hold bound instrument objects, and ``reset()`` swaps the singleton precisely in +order to rebind them to a real MeterProvider (``worker_init.ensure_telemetry_initialized`` +does ``init_telemetry()`` -> ``reset()``). A module-level cache survived that swap, so +anything that called ``get_task_metrics`` before ``init_telemetry`` cached NoOp instruments +for the life of the process. +""" from __future__ import annotations +from pydantic import BaseModel, ConfigDict +from typing import Any + from .provider import get_meter try: @@ -16,6 +34,91 @@ ('active_futures', 'twinkle.futures.active', 'Number of pending futures/tasks'), ) +# --------------------------------------------------------------------------- +# Adapter classes – wrap OTEL instruments to expose the legacy Ray-style API +# (``.inc(tags=...)`` / ``.set(value, tags=...)`` / ``.observe(value, tags=...)``) +# while delegating all measurements to OpenTelemetry. +# --------------------------------------------------------------------------- + + +class _Counter: + """Adapter mapping ``.inc(value, tags=...)`` to ``otel_counter.add()``.""" + + def __init__(self, instrument: Any) -> None: + self._instrument = instrument + + def inc(self, value: float = 1.0, tags: dict[str, str] | None = None) -> None: + self._instrument.add(value, attributes=tags or {}) + + +class _Histogram: + """Adapter mapping ``.observe(value, tags=...)`` to ``otel_histogram.record()``.""" + + def __init__(self, instrument: Any) -> None: + self._instrument = instrument + + def observe(self, value: float, tags: dict[str, str] | None = None) -> None: + self._instrument.record(value, attributes=tags or {}) + + +class _Gauge: + """Adapter mapping ``.set(value, tags=...)`` onto an OTEL UpDownCounter. + + OpenTelemetry up/down counters take *deltas*, not absolute values, so we + track the last reported value per attribute combination and emit the + incremental change. State is held per adapter instance (= per deployment), + keyed by the frozen attribute tuple. + """ + + def __init__(self, instrument: Any) -> None: + self._instrument = instrument + self._last: dict[tuple, float] = {} + + def set(self, value: float, tags: dict[str, str] | None = None) -> None: + attrs = tags or {} + key = tuple(sorted(attrs.items())) + last = self._last.get(key, 0.0) + delta = value - last + if delta != 0: + self._instrument.add(delta, attributes=attrs) + self._last[key] = value + + +# --------------------------------------------------------------------------- +# Pydantic containers for structured metric access +# --------------------------------------------------------------------------- + + +class TaskMetrics(BaseModel): + """Task queue metrics container. + + Attributes: + queue_depth: Current number of queued tasks (gauge). + tasks_total: Total task completions (counter). + execution_seconds: Pure task execution time in seconds (histogram). + queue_wait_seconds: Time from enqueue to execution start (histogram). + rate_limit_rejections: Total rate-limit rejections (counter). + rate_limiter_active_tokens: Tokens tracked by rate limiter (gauge). + """ + + model_config = ConfigDict(arbitrary_types_allowed=True) + + queue_depth: _Gauge + tasks_total: _Counter + execution_seconds: _Histogram + queue_wait_seconds: _Histogram + rate_limit_rejections: _Counter + rate_limiter_active_tokens: _Gauge + + +class _RequestMetrics(BaseModel): + """HTTP request metrics container (internal).""" + + model_config = ConfigDict(arbitrary_types_allowed=True) + + requests_total: _Counter + request_duration_seconds: _Histogram + class MetricsRegistry: """Centrally declares all metrics. Business code retrieves singleton via MetricsRegistry.get(). @@ -85,6 +188,11 @@ def __init__(self) -> None: description=description, ) + # Per-deployment adapter caches held on the instance so ``reset()`` (which + # swaps this singleton to rebind instruments) invalidates them. + self._task_metrics: dict[str, TaskMetrics] = {} + self._request_metrics: dict[str, _RequestMetrics] = {} + def _make_gauge_callback(self, name: str): """Build the sync OTEL callback that reads ``_resource_cache[name]``.""" @@ -93,6 +201,32 @@ def _callback(options): # noqa: ARG001 -- OTEL signature return _callback + # ----- Per-deployment adapter accessors ----- + + def task_metrics(self, deployment: str) -> TaskMetrics: + """Return (or build) the per-deployment task-queue metric adapters.""" + cached = self._task_metrics.get(deployment) + if cached is None: + cached = self._task_metrics[deployment] = TaskMetrics( + queue_depth=_Gauge(self.queue_depth), + tasks_total=_Counter(self.tasks_total), + execution_seconds=_Histogram(self.task_execution_seconds), + queue_wait_seconds=_Histogram(self.task_wait_seconds), + rate_limit_rejections=_Counter(self.rate_limit_rejections), + rate_limiter_active_tokens=_Gauge(self.rate_limiter_active_tokens), + ) + return cached + + def request_metrics(self, deployment: str) -> _RequestMetrics: + """Return (or build) the per-deployment HTTP request metric adapters.""" + cached = self._request_metrics.get(deployment) + if cached is None: + cached = self._request_metrics[deployment] = _RequestMetrics( + requests_total=_Counter(self.requests_total), + request_duration_seconds=_Histogram(self.request_duration_seconds), + ) + return cached + # ----- Push API for the cleanup leader ----- def set_resource_count(self, name: str, value: int) -> None: @@ -120,3 +254,17 @@ def get(cls) -> MetricsRegistry: def reset(cls) -> None: """Reset singleton (for testing or telemetry re-initialization).""" cls._instance = None + + +def get_task_metrics(deployment: str) -> TaskMetrics: + """Return the per-deployment task-queue metric adapters. + + Signature unchanged (``_init_task_queue`` needs no edit); the adapters are now + cached on the ``MetricsRegistry`` instance so ``reset()`` invalidates them. + """ + return MetricsRegistry.get().task_metrics(deployment) + + +def get_request_metrics(deployment: str) -> _RequestMetrics: + """Return the per-deployment HTTP request metric adapters.""" + return MetricsRegistry.get().request_metrics(deployment) diff --git a/src/twinkle/server/telemetry/middleware.py b/src/twinkle/server/telemetry/middleware.py deleted file mode 100644 index f6140da67..000000000 --- a/src/twinkle/server/telemetry/middleware.py +++ /dev/null @@ -1,208 +0,0 @@ -# Copyright (c) ModelScope Contributors. All rights reserved. -""" -Central metrics module for Twinkle server observability. - -This module is a **back-compat keyword shim plus a real ``_Gauge`` adapter** over -the OpenTelemetry instruments declared in -:class:`twinkle.server.telemetry.metrics.MetricsRegistry`. ``_Counter`` and -``_Histogram`` are thin pass-throughs whose only role is to accept the legacy -Ray-style ``tags=`` keyword and forward it as OTEL's ``attributes=``; ``_Gauge`` -does real work — it translates the legacy ``set(value)`` API onto OTEL's -delta-based UpDownCounter by tracking the last reported value per attribute set. -Routing every measurement through OTEL while preserving the legacy call API -means existing call sites do not need to change. - -Public entry-points (unchanged signatures): - -* ``create_metrics_middleware(deployment)`` – FastAPI HTTP middleware -* ``get_task_metrics(deployment)`` – task-queue / rate-limit gauges -""" -from __future__ import annotations - -import time -from collections.abc import Callable -from pydantic import BaseModel, ConfigDict -from typing import Any - -from twinkle.server.telemetry import MetricsRegistry -from twinkle.utils.logger import get_logger - -logger = get_logger() - -# Per-process caches; each Ray Serve worker holds its own instance. -_task_metrics_cache: dict[str, TaskMetrics] = {} -_request_metrics_cache: dict[str, _RequestMetrics] = {} - -# --------------------------------------------------------------------------- -# Adapter classes – wrap OTEL instruments to expose the legacy Ray-style API -# (``.inc(tags=...)`` / ``.set(value, tags=...)`` / ``.observe(value, tags=...)``) -# while delegating all measurements to OpenTelemetry. -# --------------------------------------------------------------------------- - - -class _Counter: - """Adapter mapping ``.inc(value, tags=...)`` to ``otel_counter.add()``.""" - - def __init__(self, instrument: Any) -> None: - self._instrument = instrument - - def inc(self, value: float = 1.0, tags: dict[str, str] | None = None) -> None: - self._instrument.add(value, attributes=tags or {}) - - -class _Histogram: - """Adapter mapping ``.observe(value, tags=...)`` to ``otel_histogram.record()``.""" - - def __init__(self, instrument: Any) -> None: - self._instrument = instrument - - def observe(self, value: float, tags: dict[str, str] | None = None) -> None: - self._instrument.record(value, attributes=tags or {}) - - -class _Gauge: - """Adapter mapping ``.set(value, tags=...)`` onto an OTEL UpDownCounter. - - OpenTelemetry up/down counters take *deltas*, not absolute values, so we - track the last reported value per attribute combination and emit the - incremental change. State is held per adapter instance (= per deployment), - keyed by the frozen attribute tuple. - """ - - def __init__(self, instrument: Any) -> None: - self._instrument = instrument - self._last: dict[tuple, float] = {} - - def set(self, value: float, tags: dict[str, str] | None = None) -> None: - attrs = tags or {} - key = tuple(sorted(attrs.items())) - last = self._last.get(key, 0.0) - delta = value - last - if delta != 0: - self._instrument.add(delta, attributes=attrs) - self._last[key] = value - - -# --------------------------------------------------------------------------- -# Pydantic containers for structured metric access -# --------------------------------------------------------------------------- - - -class TaskMetrics(BaseModel): - """Task queue metrics container. - - Attributes: - queue_depth: Current number of queued tasks (gauge). - tasks_total: Total task completions (counter). - execution_seconds: Pure task execution time in seconds (histogram). - queue_wait_seconds: Time from enqueue to execution start (histogram). - rate_limit_rejections: Total rate-limit rejections (counter). - rate_limiter_active_tokens: Tokens tracked by rate limiter (gauge). - """ - - model_config = ConfigDict(arbitrary_types_allowed=True) - - queue_depth: _Gauge - tasks_total: _Counter - execution_seconds: _Histogram - queue_wait_seconds: _Histogram - rate_limit_rejections: _Counter - rate_limiter_active_tokens: _Gauge - - -class _RequestMetrics(BaseModel): - """HTTP request metrics container (internal).""" - - model_config = ConfigDict(arbitrary_types_allowed=True) - - requests_total: _Counter - request_duration_seconds: _Histogram - - -# --------------------------------------------------------------------------- -# A. Request-level metrics (FastAPI middleware) -# --------------------------------------------------------------------------- - - -def _get_request_metrics(deployment: str) -> _RequestMetrics: - """Return (or create) per-deployment HTTP request metric adapters.""" - if deployment in _request_metrics_cache: - return _request_metrics_cache[deployment] - - reg = MetricsRegistry.get() - metrics = _RequestMetrics( - requests_total=_Counter(reg.requests_total), - request_duration_seconds=_Histogram(reg.request_duration_seconds), - ) - _request_metrics_cache[deployment] = metrics - return metrics - - -def create_metrics_middleware(deployment: str) -> Callable: - """Return a FastAPI ``http`` middleware that records request metrics. - - Usage inside a ``build_*_app()`` function:: - - from twinkle.server.telemetry.middleware import create_metrics_middleware - from twinkle.server.telemetry.tracing import create_tracing_middleware - - app.middleware('http')(verify_token) - app.middleware('http')(create_tracing_middleware("Model")) - app.middleware('http')(create_metrics_middleware("Model")) # outermost - - FastAPI executes middleware in LIFO order, so the **last** middleware - registered is the outermost wrapper. Register metrics last so its - latency observation covers the full request path including tracing - overhead and authentication. - """ - - async def metrics_middleware(request: Any, call_next: Callable) -> Any: - start = time.monotonic() - response = await call_next(request) - elapsed = time.monotonic() - start - status = str(response.status_code) - method = request.scope['route'].path if 'route' in request.scope else request.url.path - m = _get_request_metrics(deployment) - m.requests_total.inc(tags={ - 'deployment': deployment, - 'method': method, - 'status': status, - }) - m.request_duration_seconds.observe( - elapsed, tags={ - 'deployment': deployment, - 'method': method, - }) - return response - - return metrics_middleware - - -# --------------------------------------------------------------------------- -# B. Task-queue metrics -# --------------------------------------------------------------------------- - - -def get_task_metrics(deployment: str) -> TaskMetrics: - """Return (or create) per-deployment task-queue metric adapters. - - Returns a :class:`TaskMetrics` container of adapter objects; the - adapters delegate every measurement to the OTEL instruments held by - :class:`twinkle.server.telemetry.metrics.MetricsRegistry`. A separate - adapter instance is cached per deployment so that gauge-state tracking - (last value per attribute set) stays isolated. - """ - if deployment in _task_metrics_cache: - return _task_metrics_cache[deployment] - - reg = MetricsRegistry.get() - metrics = TaskMetrics( - queue_depth=_Gauge(reg.queue_depth), - tasks_total=_Counter(reg.tasks_total), - execution_seconds=_Histogram(reg.task_execution_seconds), - queue_wait_seconds=_Histogram(reg.task_wait_seconds), - rate_limit_rejections=_Counter(reg.rate_limit_rejections), - rate_limiter_active_tokens=_Gauge(reg.rate_limiter_active_tokens), - ) - _task_metrics_cache[deployment] = metrics - return metrics diff --git a/src/twinkle/server/telemetry/tracing.py b/src/twinkle/server/telemetry/tracing.py index 4f3b64732..18e351cf0 100644 --- a/src/twinkle/server/telemetry/tracing.py +++ b/src/twinkle/server/telemetry/tracing.py @@ -9,7 +9,6 @@ try: from opentelemetry import trace - from opentelemetry.context import Context from opentelemetry.propagate import extract, inject _OTEL_AVAILABLE = True except Exception: diff --git a/src/twinkle/server/telemetry/worker_init.py b/src/twinkle/server/telemetry/worker_init.py index dcf53b907..c3aebb75c 100644 --- a/src/twinkle/server/telemetry/worker_init.py +++ b/src/twinkle/server/telemetry/worker_init.py @@ -37,8 +37,9 @@ def ensure_telemetry_initialized() -> None: return try: - from twinkle.server.telemetry import TelemetryConfig, init_telemetry + from twinkle.server.config.telemetry import TelemetryConfig from twinkle.server.telemetry.metrics import MetricsRegistry + from twinkle.server.telemetry.provider import init_telemetry config = TelemetryConfig( enabled=True, @@ -82,7 +83,7 @@ def flush_telemetry_safely() -> None: so every error here is swallowed. """ try: - from twinkle.server.telemetry import shutdown_telemetry + from twinkle.server.telemetry.provider import shutdown_telemetry shutdown_telemetry() except Exception as e: # pragma: no cover - defensive logger.warning(f'Telemetry shutdown failed: {e}') diff --git a/src/twinkle/server/utils/__init__.py b/src/twinkle/server/utils/__init__.py index d6d484ab4..ab27cf29f 100644 --- a/src/twinkle/server/utils/__init__.py +++ b/src/twinkle/server/utils/__init__.py @@ -1,5 +1,12 @@ # Copyright (c) ModelScope Contributors. All rights reserved. +"""Stateless server-side helpers. + +Deliberately re-exports only ``device_utils`` and ``template_utils``: the five call sites +that import through this bucket all want just those two (``get_template_for_model``, a +52-line string map, and ``wrap_builder_with_device_group_env``), while the code that +actually needs the queue / session-resource machinery imports the full path. Re-exporting +the mixins pulled the whole OpenTelemetry SDK into every one of those five importers +(``task_queue.mixin`` -> ``telemetry`` -> the OpenTelemetry SDK + OTLP exporter). +""" from .device_utils import auto_fill_device_group_visible_devices, wrap_builder_with_device_group_env -from .lifecycle import AdapterManagerMixin, ProcessorManagerMixin, SessionResourceMixin -from .task_queue import QueueState, RateLimiter, TaskQueueConfig, TaskQueueMixin, TaskStatus from .template_utils import get_template_for_model diff --git a/src/twinkle/server/utils/lifecycle/__init__.py b/src/twinkle/server/utils/lifecycle/__init__.py deleted file mode 100644 index ea5740004..000000000 --- a/src/twinkle/server/utils/lifecycle/__init__.py +++ /dev/null @@ -1,8 +0,0 @@ -# Copyright (c) ModelScope Contributors. All rights reserved. -"""Lifecycle management utilities for session-bound resources.""" - -from .adapter import AdapterManagerMixin -from .base import SessionResourceMixin -from .processor import ProcessorManagerMixin - -__all__ = ['AdapterManagerMixin', 'ProcessorManagerMixin', 'SessionResourceMixin'] diff --git a/src/twinkle/server/utils/lifecycle/processor.py b/src/twinkle/server/utils/lifecycle/processor.py deleted file mode 100644 index 0a86309d9..000000000 --- a/src/twinkle/server/utils/lifecycle/processor.py +++ /dev/null @@ -1,109 +0,0 @@ -# Copyright (c) ModelScope Contributors. All rights reserved. -""" -Processor Lifecycle Manager Mixin for Twinkle Server. - -Mirrors AdapterManagerMixin but adds a global per-token processor limit. -Sessions are tracked via session ID; processors expire when their session expires. -""" -from __future__ import annotations - -import time -from typing import Any - -from twinkle.utils.logger import get_logger -from .base import SessionResourceMixin - -logger = get_logger() - - -class ProcessorManagerMixin(SessionResourceMixin): - """Mixin for processor lifecycle management with session-based expiration. - - Mirrors AdapterManagerMixin with an additional per-token processor limit. - - Inheriting classes should: - 1. Call _init_processor_manager() in __init__ - 2. Override _on_processor_expired() to handle cleanup - - Attributes: - _processor_timeout: Session inactivity timeout in seconds. - _per_token_processor_limit: Maximum active processors per user token. - """ - - # Set resource type for logging - _resource_type = 'Processor' - - def _init_processor_manager( - self, - processor_timeout: float = 1800.0, - per_token_processor_limit: int = 20, - ) -> None: - """Initialize the processor manager. - - Args: - processor_timeout: Timeout in seconds to determine if a session is alive. - Default is 1800.0 (30 minutes). - per_token_processor_limit: Maximum active processors per user token. - Default is 20. - """ - self._init_resource_manager( - resource_timeout=processor_timeout, - resource_max_lifetime=None, # No max lifetime for processors - ) - self._per_token_processor_limit = per_token_processor_limit - - @property - def _processor_timeout(self) -> float: - """Processor timeout for backward compatibility.""" - return self._resource_timeout - - @property - def _processor_records(self) -> dict[str, dict[str, Any]]: - """Processor records for backward compatibility.""" - return self._resource_records - - def _validate_registration(self, resource_id: str, token: str, session_id: str) -> None: - """Validate before registering a processor. Checks per-token limit. - - Args: - resource_id: Processor identifier - token: User token - session_id: Session ID - - Raises: - ValueError: If session_id is empty. - RuntimeError: If per-token limit is reached. - """ - super()._validate_registration(resource_id, token, session_id) - - current_count = sum(1 for info in self._resource_records.values() if info.get('token') == token) - if current_count >= self._per_token_processor_limit: - raise RuntimeError(f'Per-user processor limit ({self._per_token_processor_limit}) reached ' - f'for token {token[:8]}...') - - def _create_resource_record(self, token: str, session_id: str) -> dict[str, Any]: - """Create a new processor record without state field.""" - return { - 'token': token, - 'session_id': session_id, - 'created_at': time.time(), - 'expiring': False, - } - - async def _on_resource_expired(self, resource_id: str) -> None: - """Internal hook called by base class. Delegates to _on_processor_expired.""" - self._on_processor_expired(resource_id) - - def _on_processor_expired(self, processor_id: str) -> None: - """Hook called when a processor's session expires. - - Must be overridden by inheriting classes. - - Raises: - NotImplementedError: If not overridden. - """ - raise NotImplementedError(f'_on_processor_expired must be implemented by {self.__class__.__name__}') - - def stop_processor_countdown(self) -> None: - """Stop the background countdown task.""" - self.stop_resource_countdown() diff --git a/src/twinkle/server/utils/task_errors.py b/src/twinkle/server/utils/task_errors.py deleted file mode 100644 index 478745c03..000000000 --- a/src/twinkle/server/utils/task_errors.py +++ /dev/null @@ -1,5 +0,0 @@ -# Copyright (c) ModelScope Contributors. All rights reserved. - - -def task_error_payload(error: str) -> dict[str, str]: - return {'error': error, 'category': 'Server'} diff --git a/src/twinkle/server/utils/task_queue/config.py b/src/twinkle/server/utils/task_queue/config.py deleted file mode 100644 index 79d62095a..000000000 --- a/src/twinkle/server/utils/task_queue/config.py +++ /dev/null @@ -1,40 +0,0 @@ -# Copyright (c) ModelScope Contributors. All rights reserved. -""" -Task queue configuration. - -Provides TaskQueueConfig (Pydantic) for controlling rate limits, timeouts, -and queue behavior. Constraints are validated at construction time so an -invalid YAML/dict value is rejected before the deployment reaches a ready -state. -""" -from __future__ import annotations - -from pydantic import BaseModel, ConfigDict, Field - - -class TaskQueueConfig(BaseModel): - """Configuration for task queue and rate limiting. - - Attributes: - rps_limit: Maximum requests per second per user token. ``0`` disables. - tps_limit: Maximum input tokens per second per user token. ``0`` disables. - window_seconds: Sliding window for rate-limit calculations. Must be > 0. - queue_timeout: Maximum time a task can wait in queue (seconds). - execution_timeout: Maximum time a task can execute (seconds). 0 means no limit. - enabled: Whether rate limiting is enabled. - token_cleanup_multiplier: Multiplier for token cleanup threshold. - token_cleanup_interval: How often to run cleanup task (seconds). - max_input_tokens: Maximum allowed input tokens per request. - """ - - model_config = ConfigDict(extra='forbid') - - rps_limit: float = Field(default=100.0, ge=0) - tps_limit: float = Field(default=16000.0, ge=0) - window_seconds: float = Field(default=1.0, gt=0) - queue_timeout: float = Field(default=300.0, ge=0) - execution_timeout: float = Field(default=120.0, ge=0) - enabled: bool = True - token_cleanup_multiplier: float = Field(default=10.0, ge=0) - token_cleanup_interval: float = Field(default=60.0, ge=0) - max_input_tokens: int = Field(default=16000, ge=1) diff --git a/src/twinkle/server/validation/__init__.py b/src/twinkle/server/validation/__init__.py new file mode 100644 index 000000000..36335a7a7 --- /dev/null +++ b/src/twinkle/server/validation/__init__.py @@ -0,0 +1,26 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Request checks that are decidable before a task is enqueued. + +Two checks behind one entry point (:func:`assert_request_supported`), both running in +the synchronous request path: + +- the endpoint exists on this deployment's backend (else 501); +- no declared parameter belongs to a different backend (else 422). + +Their shared property is what makes them belong together: each is answerable from the +request body plus this deployment's configuration, so a rejection costs one HTTP round +trip and runs the backend method on **zero** data-parallel ranks. Discovering the same +problems from a backend exception instead means the failure surfaces inside the NCCL +critical section, where some ranks have already done work. + +Both read *declared* metadata, so neither can reject a valid request. A third, +heuristic check over passthrough key spellings was implemented and removed for failing +that bar -- see :mod:`.backend_compat` for the case that killed it. +""" +from .backend_compat import BackendCapability, assert_request_supported, resolve_backend + +__all__ = [ + 'BackendCapability', + 'assert_request_supported', + 'resolve_backend', +] diff --git a/src/twinkle/server/validation/backend_compat.py b/src/twinkle/server/validation/backend_compat.py new file mode 100644 index 000000000..3c3d4b427 --- /dev/null +++ b/src/twinkle/server/validation/backend_compat.py @@ -0,0 +1,121 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Reject a request this deployment cannot serve, before it is enqueued. + +Two questions, both answerable from the request body plus the deployment's own +configuration: + +1. does the body set a parameter that only exists on the *other* backend? +2. does the endpoint exist on this backend at all? + +Both used to be answered by the backend raising during execution -- which means after +the task was enqueued, a future record written, and the call fanned out to every +data-parallel rank. Answering them here makes the number of ranks that ran the backend +method for a rejected request exactly zero. + +**Why there is no spelling check over passthrough keys.** One was implemented here and +removed after it rejected a working call: ``set_processor('InputProcessor', +padding_side='right')``, where ``padding_side`` is a real parameter that +``InputProcessor`` reads via ``kwargs.get('padding_side')`` and therefore never declares. +It scored 0.75 against the declared ``padding_free``. That is not a tunable threshold +problem -- ``inspect.signature`` cannot see a ``**kwargs`` read at all, so "misspelled" +and "read out of ``**kwargs``" are indistinguishable to it, and any threshold catching +``bate`` -> ``beta`` also catches this. A check that rejects valid requests is worse than +no check, so the passthrough region is forwarded as given and a misspelled plugin +argument still surfaces from the plugin itself. +""" +from __future__ import annotations + +from enum import StrEnum +from typing import Any, Optional + +from twinkle.protocol.types.base import FieldRole, fields_with_role, read_backend_only +from twinkle.server.exceptions import EndpointUnavailableError, RequestRejectedError +from twinkle.utils.logger import get_logger + + +class BackendCapability(StrEnum): + """An endpoint that not every backend implements. + + Only the gradient-path splits belong here. Megatron fuses forward and backward, so + it raises ``NotImplementedError('Megatron only supports forward_backward and + forward_only')`` for each of these three; naming them makes that a 501 with the + alternative endpoints in the message instead of a 500 carrying a backend traceback. + """ + + Forward = 'forward' + Backward = 'backward' + CalculateLoss = 'calculate_loss' + + +# Capabilities a backend does NOT provide. Absent backends support everything; ``mock`` +# is deliberately absent because it is a test double that accepts every call. +_UNSUPPORTED: dict[str, frozenset[str]] = { + 'megatron': frozenset({BackendCapability.Forward, BackendCapability.Backward, BackendCapability.CalculateLoss}), +} + +_ALTERNATIVES = 'use `forward_backward` (training) or `forward_only` (inference) instead' + + +def resolve_backend(service: Any) -> str | None: + """This deployment's declared backend, or ``None`` when it has no backend concept. + + Read from the deployment's own configuration (``ModelManagement.backend``), never + inferred from the wrapper class name, the request body, or a backend exception: those + are all restatements of the same fact one step further from the source, and the + sampler deployment has no ``backend`` at all. + """ + backend = getattr(service, 'backend', None) + if backend is None: + # Not an error: Sampler deployments have no ``backend`` attribute at all. But a + # silent ``None`` meant the whole backend-compat preflight vanished with no + # trace, so the skip is now observable. ``debug`` not ``warning``: + # for Sampler the skip is normal and happens every request. + get_logger().debug('backend-compat preflight skipped: %s exposes no ``backend`` attribute', + type(service).__name__) + return None + return backend if isinstance(backend, str) else None + + +def assert_endpoint_available(service: Any, capability: str | None) -> None: + """Raise 501 when this deployment's backend does not implement ``capability``.""" + if capability is None: + return + backend = resolve_backend(service) + if backend is not None and capability in _UNSUPPORTED.get(backend, frozenset()): + raise EndpointUnavailableError(f'`{capability}` is not available on the {backend} backend; {_ALTERNATIVES}.') + + +def assert_backend_fields(service: Any, body: Any) -> None: + """Raise 422 for a field restricted to a backend other than this deployment's. + + Only a field carrying a non-``None`` value is checked. A restricted field is always + ``Optional[...] = None``, precisely so that "not sent" and "sent to the wrong + backend" stay distinguishable -- were it given the backend's own default, every + request on the other half of the fleet would be rejected. + """ + backend = resolve_backend(service) + if backend is None or backend == 'mock': + return + offenders = [] + for name, info in fields_with_role(type(body), FieldRole.BackendKwarg).items(): + allowed = read_backend_only(info) + if allowed and backend not in allowed and getattr(body, name, None) is not None: + offenders.append((name, allowed)) + if offenders: + details = '; '.join(f'`{name}` is only supported on {"/".join(allowed)}' for name, allowed in offenders) + raise RequestRejectedError( + f'This deployment runs the {backend} backend. {details}. Remove the parameter or target a ' + f'deployment running a supporting backend.', + error_code=422) + + +def assert_request_supported(service: Any, body: Any, *, capability: str | None = None) -> None: + """The single preflight entry point, called from ``run_submit`` before the enqueue. + + Both checks read *declared* metadata, so neither can produce a false positive. A + spelling heuristic over passthrough keys was tried here and removed: signature + reflection cannot distinguish a misspelling from a parameter a target reads straight + out of ``**kwargs``, so it rejected working calls. See the module docstring. + """ + assert_endpoint_available(service, capability) + assert_backend_fields(service, body) diff --git a/src/twinkle/utils/import_utils.py b/src/twinkle/utils/import_utils.py index d460521d1..e394aa5e4 100644 --- a/src/twinkle/utils/import_utils.py +++ b/src/twinkle/utils/import_utils.py @@ -1,13 +1,7 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -import importlib import importlib.metadata -import importlib.util -import os from functools import lru_cache -from itertools import chain from packaging.requirements import Requirement -from types import ModuleType -from typing import Any @lru_cache @@ -31,57 +25,3 @@ def exists(package: str): return True except ImportError: return False - - -class _LazyModule(ModuleType): - """ - Module class that surfaces all objects but only performs associated imports when the objects are requested. - """ - - # Very heavily inspired by optuna.integration._IntegrationModule - # https://github.com/optuna/optuna/blob/master/optuna/integration/__init__.py - def __init__(self, name, module_file, import_structure, module_spec=None, extra_objects=None): - super().__init__(name) - self._modules = set(import_structure.keys()) - self._class_to_module = {} - for key, values in import_structure.items(): - for value in values: - self._class_to_module[value] = key - # Needed for autocompletion in an IDE - self.__all__ = list(import_structure.keys()) + list(chain(*import_structure.values())) - self.__file__ = module_file - self.__spec__ = module_spec - self.__path__ = [os.path.dirname(module_file)] - self._objects = {} if extra_objects is None else extra_objects - self._name = name - self._import_structure = import_structure - - # Needed for autocompletion in an IDE - def __dir__(self): - result = super().__dir__() - # The elements of self.__all__ that are submodules may or may not be in the dir already, depending on whether - # they have been accessed or not. So we only add the elements of self.__all__ that are not already in the dir. - for attr in self.__all__: - if attr not in result: - result.append(attr) - return result - - def __getattr__(self, name: str) -> Any: - if name in self._objects: - return self._objects[name] - if name in self._modules: - value = self._get_module(name) - elif name in self._class_to_module.keys(): - module = self._get_module(self._class_to_module[name]) - value = getattr(module, name) - else: - raise AttributeError(f'module {self.__name__} has no attribute {name}') - - setattr(self, name, value) - return value - - def _get_module(self, module_name: str): - return importlib.import_module('.' + module_name, self.__name__) - - def __reduce__(self): - return self.__class__, (self._name, self.__file__, self._import_structure) diff --git a/src/twinkle/utils/nccl_safe.py b/src/twinkle/utils/nccl_safe.py index 8bafeb1f2..a7c7eef0c 100644 --- a/src/twinkle/utils/nccl_safe.py +++ b/src/twinkle/utils/nccl_safe.py @@ -1,339 +1,73 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -"""NCCL-safe utilities for production distributed training. - -Provides three layers of protection to prevent NCCL hangs: - -Layer 1 - safe_loss(): - Wraps loss instances to catch computation errors and return - graph-connected zero loss (ensures FSDP ReduceScatter can proceed). - -Layer 2 - @nccl_safe decorator: - Wraps forward_backward methods to ensure backward() always executes - after forward() has started, even if intermediate code raises. - -Layer 3 - @nccl_safe_megatron decorator: - Wraps Megatron backend methods (forward_only, forward_backward) where - the entire function body involves NCCL communication (sync=True). - Catches pre-communication errors (e.g. data preprocessing failures) - that would otherwise leave other DP ranks waiting at a collective. - -Controlled by environment variable: - TWINKLE_FAIL_FAST=1 (default, development): all protection is transparent, - exceptions propagate normally. - TWINKLE_FAIL_FAST=0 (production): protection activated, exceptions in - NCCL-critical sections are caught and handled gracefully. +"""NCCL critical-section failure logging. + +Single responsibility: inside a Megatron NCCL-critical method, log a +rank-attributed failure and then re-raise it unchanged. + +This module does NOT prevent asymmetric-failure blocking -- nothing at this layer +can. A rank that swallows its exception still does not enter the collective, so the +other ranks stay blocked regardless. The time bound for an asymmetric failure comes +from Ray_Get_Timeout (the effective execution timeout applied per future), not from +this decorator. Diagnosability is the only reason this wrapper exists. + +Coverage removed together with the former Layer 1 (the loss-instance wrapper) and +Layer 2 (the forward/backward decorator) silent degradation: under FSDP, the window +between ``calculate_loss``'s loss call and its surrounding bookkeeping (metric +accumulation, ``status.num_tokens``), between the three calls inside a +``forward_backward`` body, and numerical problems inside a loss (NaN, shape mismatch) +may each constitute a "forward ran, backward did not" asymmetric-failure window. That +window is no longer covered by any silent degradation; its time bound is the two +bounds documented for the task queue (record-terminal = ``queue_timeout + T``; +resource-release = ``Collect_Width * T``). """ import functools -import os -from twinkle.data_format import LossOutput -from twinkle.loss import Loss from twinkle.utils.logger import get_logger logger = get_logger() +# Errors are logged with at most this many trailing characters of traceback. +_TRACEBACK_LIMIT = 8192 -def _is_fail_fast() -> bool: - """Check if fail-fast mode is enabled (default: enabled). - - Returns True (fail-fast/development mode) unless TWINKLE_FAIL_FAST - is explicitly set to a falsy value. - """ - val = os.getenv('TWINKLE_FAIL_FAST', '1').upper() - return val not in ('0', 'NO', 'FALSE', 'OFF') - - -# ─── Layer 1: safe_loss ──────────────────────────────────────────────────── +def _global_rank() -> int: + """Best-effort global rank for failure attribution; -1 if unavailable.""" + try: + from twinkle.utils import Platform + return Platform.get_rank() + except Exception: + return -1 -def safe_loss(loss_instance): - """Wrap loss instance for production graceful degradation. - Always wraps the loss instance (idempotent). The fail-fast check is deferred - to call time so that TWINKLE_FAIL_FAST can be set after wrapping (e.g. in - Ray actor processes where env vars may not be inherited from the launcher). +def nccl_safe_megatron(func): + """Log a rank-attributed failure inside the NCCL critical section, then re-raise. - When TWINKLE_FAIL_FAST=1 (default, development): wrapper is transparent, - exceptions propagate normally. - When TWINKLE_FAIL_FAST=0 (production): wrapper catches exceptions and - returns a graph-connected zero loss (ensures FSDP ReduceScatter proceeds). - - Idempotent: already-wrapped instances are returned as-is. - """ - if getattr(loss_instance, '_nccl_safe_wrapped', False): - return loss_instance - return SafeLossWrapper(loss_instance) - - -class SafeLossWrapper(Loss): - """Loss subclass that catches computation errors and returns graph-connected zero loss. - - Inherits from :class:`twinkle.loss.Loss` so ``isinstance(wrapper, Loss)`` - assertions in the training pipeline continue to pass. + This decorator does *not* prevent asymmetric-failure blocking -- nothing at this + layer can. A rank that swallows its exception still does not enter the collective. + The time bound for that case comes from Ray_Get_Timeout. Diagnosability is the only + reason this wrapper still exists. Its behavior is unconditional: no environment + variable or config switch affects it, and it returns no degraded value. """ - def __init__(self, loss_instance): - super().__init__() - self._loss_instance = loss_instance - self.require_logps = getattr(loss_instance, 'require_logps', True) - self.require_entropy = getattr(loss_instance, 'require_entropy', False) - self.require_logits = getattr(loss_instance, 'require_logits', False) - self.enable_sampling_replay = getattr(loss_instance, 'enable_sampling_replay', False) - self.require_values = getattr(loss_instance, 'require_values', False) - self.reduction = getattr(loss_instance, 'reduction', 'mean') - self._nccl_safe_wrapped = True - - def __call__(self, inputs, outputs, **kwargs): - if _is_fail_fast(): - return self._loss_instance(inputs, outputs, **kwargs) + @functools.wraps(func) + def wrapper(self, *args, **kwargs): try: - return self._loss_instance(inputs, outputs, **kwargs) - except Exception as e: + return func(self, *args, **kwargs) + except Exception as exc: import traceback - logger.warning('[nccl_safe] Loss computation skipped due to error: ' - '%s: %s\n%s', - type(e).__name__, e, traceback.format_exc()) - return _zero_loss(outputs) - - def micro_batch_scale(self, inputs, indices): - """Preserve the wrapped loss's micro-batch reduction semantics.""" - return self._loss_instance.micro_batch_scale(inputs, indices) - - -def _zero_loss(outputs) -> 'LossOutput': - """Create a graph-connected zero loss for FSDP compatibility. - - Finds a gradient-bearing tensor from outputs to maintain graph connectivity, - ensuring backward hooks (ReduceScatter) fire. - """ - import torch - if isinstance(outputs, dict): - for key in ('logps', 'values', 'logits', 'loss'): - t = outputs.get(key) - if t is not None and isinstance(t, torch.Tensor) and t.requires_grad: - return LossOutput(loss=(t.flatten()[:1] * 0).sum(), num_tokens=0) - # Fallback: standalone zero tensor (may not trigger FSDP hooks) - device = 'cpu' - if isinstance(outputs, dict): - for v in outputs.values(): - if hasattr(v, 'device'): - device = v.device - break - return LossOutput(loss=torch.zeros((), device=device, requires_grad=True), num_tokens=0) - - -# ─── Layer 2: @nccl_safe decorator ────────────────────────────────────────── - - -def nccl_safe(func=None, *, tinker=False): - """Decorator ensuring backward() executes if forward() has already run. - - Detects forward completion by comparing train_status.outputs before/after - the wrapped function call. If an exception occurs after forward has run - but before backward completes, forces a zero-gradient backward pass to - prevent NCCL hang (other ranks waiting for ReduceScatter). - - Args: - func: The function to decorate (when used without arguments). - tinker: If True, fallback returns ``[[], 0.0]`` (tinker format). - If False, fallback returns outputs dict with ``loss=0.0``. - - Usage:: - - @remote_function(dispatch='slice_dp', collect=...) - @nccl_safe(tinker=True) - def tinker_forward_backward(self, *, inputs, adapter_name, ...): - # method body completely unchanged - ... - - @remote_function(dispatch='slice_dp', collect=...) - @nccl_safe - def forward_backward(self, *, inputs, **kwargs): - # method body completely unchanged - ... - """ - - def decorator(fn): - - @functools.wraps(fn) - def wrapper(self, *args, **kwargs): - if _is_fail_fast(): - return fn(self, *args, **kwargs) - - # Extract adapter_name for state tracking - adapter_name = kwargs.get('adapter_name') - if adapter_name is None and hasattr(self, '_get_default_group'): - adapter_name = self._get_default_group() - - og = self.optimizer_group.get(adapter_name) if adapter_name else None - if og is None: - # Cannot track state without optimizer group, passthrough - return fn(self, *args, **kwargs) - - # Snapshot state before call to detect forward completion - outputs_before = og.train_status.outputs - - try: - return fn(self, *args, **kwargs) - except Exception as e: - outputs_after = og.train_status.outputs - forward_ran = (outputs_after is not None and outputs_after is not outputs_before) - - if not forward_ran: - # Pre-forward failure: no NCCL ops started, safe to propagate - raise - - # Forward completed. Check if backward already ran. - # TransformersModel.backward() clears loss_value to None. - backward_done = (og.train_status.loss_value is None) - - if backward_done: - # Post-backward failure (e.g. output formatting) - # No NCCL hang risk, just return gracefully - logger.warning(f'[nccl_safe] Post-backward error (no NCCL risk): ' - f'{type(e).__name__}: {e}') - else: - # CRITICAL: forward ran but backward didn't → NCCL hang risk! - logger.warning(f'[nccl_safe] Forcing zero backward to prevent NCCL hang: ' - f'{type(e).__name__}: {e}') - _force_zero_backward(self, og, adapter_name, kwargs) - - # Return fallback result - if tinker: - return [[], 0.0] - outputs_after['loss'] = 0.0 - return outputs_after - - return wrapper - - if func is not None: - # @nccl_safe without arguments - return decorator(func) - # @nccl_safe(tinker=True) with arguments - return decorator - - -def _iter_model_params(model): - """Iterate parameters from ``model.model``, supporting single model or list of models.""" - raw_model = getattr(model, 'model', None) - if raw_model is None: - return iter([]) - if isinstance(raw_model, (list, tuple)): - for m in raw_model: - yield from m.parameters() - else: - yield from raw_model.parameters() - - -def _force_zero_backward(model, og, adapter_name, kwargs): - """Force a zero-gradient backward pass to prevent NCCL hang. - - Creates a graph-connected zero loss tensor and calls backward(), - ensuring FSDP ReduceScatter hooks fire on all ranks. - """ - import torch - - outputs = og.train_status.outputs - - # Find a graph-connected tensor for zero loss - zero_loss = None - if outputs is not None and isinstance(outputs, dict): - for key in ('logps', 'values', 'logits', 'loss'): - t = outputs.get(key) - if t is not None and isinstance(t, torch.Tensor) and t.requires_grad: - zero_loss = (t.flatten()[:1] * 0).sum() - break - - if zero_loss is None: - # Fallback: use first model parameter to maintain graph connectivity. - # Do NOT detach() the parameter -- the zero loss must remain connected - # to the model's autograd graph so FSDP ReduceScatter hooks fire. - # Use lazy iteration to avoid materializing the full parameter list. - try: - param = next((p for p in _iter_model_params(model) if p.requires_grad), None) - if param is not None: - zero_loss = (param.flatten()[0] * 0).sum() + rank = _global_rank() + context = f'twinkle backend method={func.__name__}, global_rank={rank}' + if hasattr(exc, 'add_note'): + exc.add_note(context) + elif exc.args: + exc.args = (f'{exc.args[0]} [{context}]', *exc.args[1:]) else: - zero_loss = torch.zeros((), device='cuda', requires_grad=True) - except Exception: - zero_loss = torch.zeros((), device='cuda', requires_grad=True) - - og.train_status.loss_value = zero_loss - - # Call backward with minimal kwargs - bwd_kwargs = {'adapter_name': adapter_name} - gas = kwargs.get('gradient_accumulation_steps') - if gas is not None: - bwd_kwargs['gradient_accumulation_steps'] = gas - model.backward(**bwd_kwargs) - - -# ─── Layer 3: @nccl_safe_megatron decorator ────────────────────────────────── - - -def nccl_safe_megatron(func=None, *, tinker=False, forward_only=False): - """Decorator for Megatron backend methods where the entire body is NCCL-critical. - - Unlike @nccl_safe (which detects forward/backward boundaries), this decorator - treats the **entire function** as a NCCL-critical section. In Megatron, - forward_only and forward_backward both call get_forward_backward_func() which - requires all DP ranks to enter synchronously. If one rank fails during data - preprocessing (before entering Megatron's scheduler), other ranks will hang - waiting for the collective. - - This decorator catches ALL exceptions (when TWINKLE_FAIL_FAST=0) and returns - a safe fallback value, preventing NCCL hang from asymmetric failures. - - Args: - func: The function to decorate (when used without arguments). - tinker: If True, fallback returns ``[[], 0.0]`` (tinker format). - forward_only: If True, fallback returns empty dict ``{}`` (forward_only format). - - Usage:: - - @remote_function(dispatch='slice_dp', collect=..., sync=True) - @nccl_safe_megatron - def forward_backward(self, *, inputs, **kwargs): - ... - - @remote_function(dispatch='slice_dp', collect=...) - @nccl_safe_megatron(forward_only=True) - def forward_only(self, *, inputs, **kwargs): - ... - - @remote_function(dispatch='slice_dp', collect=..., sync=True) - @nccl_safe_megatron(tinker=True) - def tinker_forward_backward(self, *, inputs, **kwargs): - ... - """ - - def decorator(fn): - - @functools.wraps(fn) - def wrapper(self, *args, **kwargs): - if _is_fail_fast(): - return fn(self, *args, **kwargs) - - try: - return fn(self, *args, **kwargs) - except Exception as e: - import traceback - logger.warning(f'[nccl_safe_megatron] Exception in Megatron method ' - f'{fn.__name__}: {type(e).__name__}: {e}\n' - f'{traceback.format_exc()}') - - # Return safe fallback to prevent NCCL hang on other ranks - if tinker: - return [[], 0.0] - if forward_only: - return {} - # forward_backward fallback: return dict with loss=0.0 - return {'loss': 0.0} - - return wrapper - - if func is not None: - # @nccl_safe_megatron without arguments - return decorator(func) - # @nccl_safe_megatron(tinker=True) with arguments - return decorator + exc.args = (context, ) + tb = traceback.format_exc() + if len(tb) > _TRACEBACK_LIMIT: + tb = tb[-_TRACEBACK_LIMIT:] + logger.error('[nccl_safe_megatron] %s in %s on global rank %s:\n%s', + type(exc).__name__, func.__name__, rank, tb) + raise + + return wrapper diff --git a/src/twinkle_agentic/async_rl/data_plane.py b/src/twinkle_agentic/async_rl/data_plane.py index da9015f0c..58423cedc 100644 --- a/src/twinkle_agentic/async_rl/data_plane.py +++ b/src/twinkle_agentic/async_rl/data_plane.py @@ -5,9 +5,10 @@ from typing import Any, Sequence +from twinkle.data_format import (REQUIRED_MODEL_INPUT_FIELDS, ROLLOUT_TRAIN_FIELDS, columns_to_tq_fields, + rows_to_tq_fields) from .native_tq import (AsyncTQClient, append_fields, batch_size_for_groups, clear_partition, fetch_ready_batch, metadata_size, preallocate_partition, set_sample_tags, split_batch_meta) -from .tq_utils import REQUIRED_MODEL_INPUT_FIELDS, ROLLOUT_TRAIN_FIELDS, columns_to_tq_fields, rows_to_tq_fields from .types import ClaimedBatch, LoraContext, PartitionAdmission, PreparedPartition, PromptGroup, RolloutOutput _REQUIRED_ROLLOUT_FIELDS = frozenset((*REQUIRED_MODEL_INPUT_FIELDS, 'logprobs', 'rewards')) diff --git a/src/twinkle_agentic/async_rl/pipeline.py b/src/twinkle_agentic/async_rl/pipeline.py index 75a78e023..09d5846e9 100644 --- a/src/twinkle_agentic/async_rl/pipeline.py +++ b/src/twinkle_agentic/async_rl/pipeline.py @@ -637,7 +637,7 @@ def _train_batch_with_config( *, model_data_parallel_size: int = 1, ) -> dict[str, Any]: - from .tq_utils import REQUIRED_MODEL_INPUT_FIELDS + from twinkle.data_format import REQUIRED_MODEL_INPUT_FIELDS size = int(data.batch_size[0]) inputs = [{name: data[name][index] for name in REQUIRED_MODEL_INPUT_FIELDS} for index in range(size)] diff --git a/src/twinkle_agentic/async_rl/tq_utils.py b/src/twinkle_agentic/async_rl/tq_utils.py deleted file mode 100644 index 55284567e..000000000 --- a/src/twinkle_agentic/async_rl/tq_utils.py +++ /dev/null @@ -1,29 +0,0 @@ -# Copyright (c) ModelScope Contributors. All rights reserved. -from __future__ import annotations - -from twinkle.tq_utils import columns_to_tq_fields, rows_to_tq_fields - -TRANSFORMERS_INPUT_FIELDS = ( - 'input_ids', - 'labels', - 'attention_mask', - 'position_ids', - 'cu_seqlens', - 'completion_mask', - 'pixel_values', - 'image_grid_thw', - 'video_pixel_values', - 'video_grid_thw', - 'input_features', - 'feature_attention_mask', -) -REQUIRED_MODEL_INPUT_FIELDS = ('input_ids', 'labels', 'attention_mask', 'position_ids') -ROLLOUT_TRAIN_FIELDS = (*TRANSFORMERS_INPUT_FIELDS, 'logprobs', 'rewards', 'advantages', 'returns') - -__all__ = [ - 'ROLLOUT_TRAIN_FIELDS', - 'REQUIRED_MODEL_INPUT_FIELDS', - 'TRANSFORMERS_INPUT_FIELDS', - 'columns_to_tq_fields', - 'rows_to_tq_fields', -] diff --git a/src/twinkle_agentic/preprocessor/intent_classifier.py b/src/twinkle_agentic/preprocessor/intent_classifier.py index 6d1b21c24..35e1c10f8 100644 --- a/src/twinkle_agentic/preprocessor/intent_classifier.py +++ b/src/twinkle_agentic/preprocessor/intent_classifier.py @@ -39,7 +39,7 @@ r'\\times|\\div|\\pm|\\leq|\\geq|\\neq|\\approx|\\equiv|' r'\\infty|\\pi|\\alpha|\\beta|\\gamma|\\theta|\\lambda|\\mu|\\sigma|\\prod|\\to|\\rightarrow|' r'\\\[.+?\\\]|' - # R1-distill writes math in plain Unicode without $...$; catch operators, Greek, sub/super digits, fractions. + # Also match plain-Unicode operators, Greek letters, super/subscripts, and fractions. r'[×÷±°∑∏∫√∂∇∞∈∋⊂⊃⊆⊇≤≥≠≈≡≅∝⇒⇔]|' r'[α-ωΔΘΛΞΠΣΦΨΩ]|' r'[⁰¹²³⁴-⁹₀-₉]|' @@ -355,8 +355,8 @@ class IntentClassifier(Preprocessor): Pure-heuristic, no LLM. Each intent is a pluggable :class:`IntentDetector`; pass ``detectors=[...]`` to extend or override. - R3: this is an *annotator* — by default it never drops rows - (``drop_no_key_rounds=False``); rows with no detected key round are simply + This is an *annotator*: by default it never drops rows + (``drop_no_key_rounds=False``), and rows with no detected key round are simply tagged ``INTENT_OTHER``. Set ``drop_no_key_rounds=True`` to also filter. Annotates per row:: @@ -366,7 +366,7 @@ class IntentClassifier(Preprocessor): ('intents', dict[str, str])] # per-round intent """ - # R4: default to the detectors with a live downstream consumer. The heavier + # Default to detectors with a live downstream consumer. The heavier # heuristics (ComplexLogic / Reasoning / UserDissatisfaction) are kept as # importable classes but dropped from the default set — their outputs had no # active consumer. Pass ``detectors=[...]`` to re-enable them. diff --git a/src/twinkle_client/__init__.py b/src/twinkle_client/__init__.py index bb13f19ad..1177403b5 100644 --- a/src/twinkle_client/__init__.py +++ b/src/twinkle_client/__init__.py @@ -1,8 +1,11 @@ # Copyright (c) ModelScope Contributors. All rights reserved. from __future__ import annotations -from typing import Optional, TYPE_CHECKING + +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: + from .data_plane import DataPlaneClient + from .http import ClientContext, ClientTransport from .manager import TwinkleClient @@ -21,7 +24,7 @@ def init_tinker_client(**kwargs) -> None: Example:: - >>> from twinkle_client import init_tinker_client + >>> from twinkle import init_tinker_client >>> init_tinker_client() >>> from tinker import ServiceClient >>> client = ServiceClient(base_url='http://localhost:8000', api_key='your_token') @@ -35,11 +38,11 @@ def init_tinker_client(**kwargs) -> None: def init_twinkle_client( - base_url: Optional[str] = None, - api_key: Optional[str] = None, + base_url: str | None = None, + api_key: str | None = None, session_heartbeat_interval: int = 10, **kwargs, -) -> 'TwinkleClient': +) -> TwinkleClient: """ Initialize a Twinkle client. @@ -64,7 +67,7 @@ def init_twinkle_client( An initialised :class:`~twinkle_client.manager.TwinkleClient` instance. """ from .manager import TwinkleClient - return TwinkleClient( + return TwinkleClient.connect( base_url=base_url, api_key=api_key, session_heartbeat_interval=session_heartbeat_interval, @@ -72,6 +75,17 @@ def init_twinkle_client( ) -from .data_plane import DataPlaneClient +def __getattr__(name: str) -> Any: + if name == 'DataPlaneClient': + from .data_plane import DataPlaneClient + value = DataPlaneClient + elif name in {'ClientContext', 'ClientTransport'}: + from .http import ClientContext, ClientTransport + value = {'ClientContext': ClientContext, 'ClientTransport': ClientTransport}[name] + else: + raise AttributeError(f'module {__name__!r} has no attribute {name!r}') + globals()[name] = value + return value + -__all__ = ['DataPlaneClient', 'init_tinker_client', 'init_twinkle_client'] +__all__ = ['ClientContext', 'ClientTransport', 'DataPlaneClient', 'init_tinker_client', 'init_twinkle_client'] diff --git a/src/twinkle_client/_future.py b/src/twinkle_client/_future.py new file mode 100644 index 000000000..b8fc5a3b4 --- /dev/null +++ b/src/twinkle_client/_future.py @@ -0,0 +1,179 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Client_Future_Layer: the one polling implementation in Twinkle_Client. + +Public client methods keep their synchronous signatures by calling :func:`resolve`; +no future object is ever exposed. Private module (underscore name) because it is +never imported by Twinkle_Server. +""" +from __future__ import annotations + +import logging +import requests +import time +from typing import Any, Optional + +from twinkle.protocol.types.lifecycle import TERMINAL_STATUSES, TaskEnvelope +from twinkle_client.exceptions import TaskCancelledError, TaskFailedError, TaskRecordLostError, TaskWaitTimeoutError +from twinkle_client.http import ClientTransport +from twinkle_client.http.context import capture_transport + +logger = logging.getLogger('twinkle_client') + +# An independent constant, NOT derived from any server-side timeout: the server +# guarantees a task reaches a terminal state, so this is only "how long the client is +# willing to wait". Deliberately different from the server's execution_timeout fallback +# so a reader does not think the two are related. +_DEFAULT_TOTAL_TIMEOUT = 7200.0 + +# A 404 is Retrieve_Endpoint's verdict after a whole Long_Poll_Window, but state +# jitter (Redis blip, actor restart) can hide an in-flight record for one window, +# so bound-retry before declaring the record lost. +_NOT_FOUND_RETRY_MAX = 3 + +# 5xx / 408 / connection errors are retried with exponential backoff. +_TRANSPORT_RETRY_MAX = 5 + + +def _retrieve_url(transport: ClientTransport) -> str: + return transport.url('twinkle/retrieve_future') + + +def _cancel_url(transport: ClientTransport) -> str: + return transport.url('twinkle/cancel') + + +def _best_effort_cancel(request_id: str, transport: ClientTransport) -> None: + """Ask the server to drop a task when the caller abandons the wait (e.g. Ctrl-C). + + Never raises: a failed cancel must not mask the original interrupt. The server + only drops not-yet-started tasks, so a running task is unaffected. + """ + try: + transport.post(_cancel_url(transport), json_data={'request_id': request_id}, timeout=2) + except BaseException as e: # noqa: BLE001 - best effort; never mask the interrupt + logger.debug('[future] best-effort cancel of %s failed: %s', request_id, e) + + +def _post_retrieve(request_id: str, transport: ClientTransport) -> TaskEnvelope: + """POST one retrieve and parse the reply into a TaskEnvelope. + + Raises ``requests.HTTPError`` (a :class:`TwinkleHTTPError` after the client + error-parsing change lands) on a non-2xx response. + """ + response = transport.post(_retrieve_url(transport), json_data={'request_id': request_id}) + return TaskEnvelope.model_validate(response.json()) + + +def _status_of(error: requests.HTTPError) -> int | None: + status = getattr(error, 'status_code', None) + if status is None and getattr(error, 'response', None) is not None: + status = error.response.status_code + return status + + +def _is_retryable(status: int) -> bool: + return status == 408 or 500 <= status <= 599 + + +def _log_queue_state(reply: TaskEnvelope) -> None: + if reply.queue_state and reply.queue_state != 'active': + logger.info('[future] task %s waiting: queue_state=%s reason=%s', reply.request_id, reply.queue_state, + reply.queue_state_reason) + + +def _unwrap(env: TaskEnvelope, model_cls) -> Any: + """Turn a terminal TaskEnvelope into a return value or an exception. + + Takes the envelope *whole* rather than destructured fields: the submit and + retrieve paths must not be able to pass different subsets. A signature like + ``(status, result, request_id, model_cls, error=None)`` would let the submit + path simply never pass ``error`` -- and then every failure completing inside + the Inline_Fast_Path window would raise 'no recorded payload' while its real + payload sat unread. + """ + if env.status == 'failed': + p = env.error + raise TaskFailedError( + p.error, + category=p.category.value, + request_id=env.request_id, + error_code=p.error_code, + details=p.details, + ) + if env.status == 'cancelled': + p = env.error + raise TaskCancelledError( + p.error if p is not None else 'Task cancelled', + request_id=env.request_id, + error_code=p.error_code if p is not None else None, + ) + return model_cls.model_validate(env.result) if model_cls is not None else env.result + + +def resolve( + submit: TaskEnvelope, + *, + model_cls, + total_timeout: float = _DEFAULT_TOTAL_TIMEOUT, + transport: ClientTransport | None = None, +) -> Any: + """Block until ``submit``'s task reaches a terminal state, then return its result. + + A terminal submit envelope is unwrapped directly, issuing no Retrieve_Endpoint + request at all (single round trip for control-plane ops). Otherwise the same + ``_unwrap`` is applied to each retrieve reply, so a failure inside the + Inline_Fast_Path window and one observed via retrieve take an identical path. + + The main loop never sleeps: waiting is delegated to Retrieve_Endpoint's + long-poll. Only transport retries back off. + """ + if submit.status in TERMINAL_STATUSES: + return _unwrap(submit, model_cls) # same call as the retrieve path + + resolved_transport = capture_transport(transport) + deadline = time.monotonic() + total_timeout + transport_failures = not_found_count = 0 + try: + while True: + if time.monotonic() >= deadline: + raise TaskWaitTimeoutError(request_id=submit.request_id, waited=total_timeout) + try: + reply = _post_retrieve(submit.request_id, resolved_transport) + transport_failures = not_found_count = 0 + except (requests.ConnectionError, requests.Timeout): + transport_failures += 1 + if transport_failures > _TRANSPORT_RETRY_MAX: + raise + time.sleep(min(2**transport_failures, 30)) + continue + except requests.HTTPError as e: + status = _status_of(e) + if status == 404: + not_found_count += 1 + if not_found_count > _NOT_FOUND_RETRY_MAX: + raise TaskRecordLostError(request_id=submit.request_id) from e + continue + if status is None or not _is_retryable(status): + raise + transport_failures += 1 + if transport_failures > _TRANSPORT_RETRY_MAX: + raise + time.sleep(min(2**transport_failures, 30)) + continue + if reply.status in TERMINAL_STATUSES: + return _unwrap(reply, model_cls) # same call as the submit path + _log_queue_state(reply) + except (KeyboardInterrupt, SystemExit): + # Caller abandoned the wait: best-effort ask the server to drop the task if it + # has not started, then re-raise so the interrupt is never swallowed. + _best_effort_cancel(submit.request_id, resolved_transport) + raise + + +def resolve_response(response, model_cls, *, transport: ClientTransport | None = None) -> Any: + """Resolve a Submit_Endpoint response using the submitter's transport.""" + return resolve( + TaskEnvelope.model_validate(response.json()), + model_cls=model_cls, + transport=transport, + ) diff --git a/src/twinkle_client/_request_builder.py b/src/twinkle_client/_request_builder.py new file mode 100644 index 000000000..3afc8b5d1 --- /dev/null +++ b/src/twinkle_client/_request_builder.py @@ -0,0 +1,120 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Build a twinkle-native request body from a caller's keyword arguments. + +Every public client method used to hand-assemble a ``json_data`` dict. That made the +request schema a second, implicit source of truth: a field the server declared but the +client never sent, or a name only one side spelled correctly, was invisible until a +request failed on the wire. Here the request *model* is the only source of truth -- +the client instantiates it, so the same rules that guard the server also guard the +caller, in-process and without a round trip. + +Routing a caller's ``**kwargs`` needs exactly three rules, all read off the model: + +1. the name is a declared field -> assign it to that field; +2. it is not declared and the model has exactly one passthrough region -> put it there; +3. it is not declared and the model has none, or more than one -> raise. + +Rule 3 does not guess. A processor call has both ``init_kwargs`` (the constructor's +arguments) and ``call_kwargs`` (the invoked method's); no rule based on the name alone +can tell which one a caller meant, and choosing wrong sends a valid argument to the +wrong callable -- a silently wrong result rather than an error. +""" +from __future__ import annotations + +from dataclasses import asdict, is_dataclass +from pydantic import BaseModel +from typing import Any, Mapping + +from twinkle.protocol.json_utils import json_safe +from twinkle.protocol.types.base import FieldRole, fields_with_role +from twinkle_client.exceptions import TwinkleClientValidationError + + +def to_wire_value(value: Any) -> Any: + """Recursively convert a caller value to its JSON wire representation. + + ``ClientTransport.post`` and schema-driven request construction both use this + function. ``post_model`` remains a separate Pydantic JSON entry point and + deliberately excludes unset optionals. + """ + if isinstance(value, (str, int, float, bool, type(None))): + return value + if isinstance(value, (bytes, bytearray, memoryview)): + raise TwinkleClientValidationError('Binary values are not supported by the JSON transport') + component_id = getattr(value, 'processor_id', None) + if isinstance(component_id, str): + return component_id + if isinstance(value, Mapping): + return {str(key): to_wire_value(item) for key, item in value.items()} + if isinstance(value, (list, tuple, set, frozenset)): + return [to_wire_value(item) for item in value] + + from peft import LoraConfig + + from twinkle.dataset import DatasetMeta + if isinstance(value, (DatasetMeta, LoraConfig)): + from twinkle.protocol.serialize import serialize_object + return serialize_object(value) + if is_dataclass(value) and not isinstance(value, type): + return to_wire_value(asdict(value)) + if isinstance(value, BaseModel): + return value.model_dump(mode='json') + return json_safe(value) + + +def build_request(model_cls: type[BaseModel], /, **values: Any) -> BaseModel: + """Instantiate ``model_cls`` from caller arguments, routing undeclared names. + + ``None`` values for undeclared names are dropped rather than routed: client methods + pass optional arguments unconditionally, and forwarding an explicit ``None`` into a + passthrough region would hand the backend a null it never had before. + + Raises: + TwinkleClientValidationError: an argument has no field and no unambiguous + passthrough region, or it collides with an explicitly passed region key. + pydantic.ValidationError: the assembled body violates the model. + """ + declared = model_cls.model_fields + regions = fields_with_role(model_cls, FieldRole.Passthrough) + + body: dict[str, Any] = {} + routed: dict[str, Any] = {} + for name, value in values.items(): + if name in declared: + body[name] = to_wire_value(value) if value is not None else None + elif value is None: + continue + elif len(regions) == 1: + routed[name] = to_wire_value(value) + elif not regions: + raise TwinkleClientValidationError( + f'{model_cls.__name__} has no field {name!r} and no passthrough region to put it in; ' + f'known fields: {sorted(declared)}') + else: + raise TwinkleClientValidationError( + f'{model_cls.__name__} has no field {name!r} and more than one passthrough region, so its ' + f'target is ambiguous. Pass it inside one of: {sorted(regions)}') + + if routed: + region = next(iter(regions)) + explicit = body.get(region) or {} + if not isinstance(explicit, Mapping): + raise TwinkleClientValidationError(f'{region} must be a mapping, got {type(explicit).__name__}') + collisions = set(explicit) & set(routed) + if collisions: + raise TwinkleClientValidationError( + f'these arguments were passed both directly and inside {region}: {sorted(collisions)}') + body[region] = {**explicit, **routed} + + return model_cls(**body) + + +def request_json(body: BaseModel) -> str: + """Serialize a request body once. + + ``exclude_none=True`` keeps unset optionals off the wire, which is what lets the + server treat "absent" and "not requested" as the same thing rather than maintaining + a second set of defaults. One pydantic-core pass, not a Python-level walk followed + by ``json.dumps``. + """ + return body.model_dump_json(exclude_none=True) diff --git a/src/twinkle_client/async_rl/workers.py b/src/twinkle_client/async_rl/workers.py index 8d105bcf9..a9014fd07 100644 --- a/src/twinkle_client/async_rl/workers.py +++ b/src/twinkle_client/async_rl/workers.py @@ -37,10 +37,7 @@ def __init__(self, workers: Sequence[Worker]) -> None: raise ValueError(f'worker names must be unique, got {names}') async def run(self) -> None: - tasks = { - asyncio.create_task(worker.run(), name=worker.name): worker - for worker in self.workers - } + tasks = {asyncio.create_task(worker.run(), name=worker.name): worker for worker in self.workers} try: done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_EXCEPTION) failure = next( diff --git a/src/twinkle_client/auto/__init__.py b/src/twinkle_client/auto/__init__.py index 3916c91e2..27a42cc67 100644 --- a/src/twinkle_client/auto/__init__.py +++ b/src/twinkle_client/auto/__init__.py @@ -18,12 +18,16 @@ def _configure_logging(verbose: bool = False) -> None: from logging.handlers import RotatingFileHandler handler = RotatingFileHandler( - _LOG_FILE, maxBytes=5 * 1024 * 1024, backupCount=3, encoding='utf-8', + _LOG_FILE, + maxBytes=5 * 1024 * 1024, + backupCount=3, + encoding='utf-8', ) - handler.setFormatter(logging.Formatter( - '%(asctime)s [%(levelname)s] %(name)s: %(message)s', - datefmt='%Y-%m-%d %H:%M:%S', - )) + handler.setFormatter( + logging.Formatter( + '%(asctime)s [%(levelname)s] %(name)s: %(message)s', + datefmt='%Y-%m-%d %H:%M:%S', + )) level = logging.DEBUG if verbose else logging.INFO twinkle_logger = logging.getLogger('twinkle') @@ -45,6 +49,7 @@ def main(argv: list[str] | None = None) -> int: """ import sys import typer + from twinkle.version import __version__ app = typer.Typer( @@ -61,43 +66,51 @@ def _version_callback(value: bool) -> None: @app.command() def launch( run_id: str | None = typer.Option( - None, '--run-id', '-r', + None, + '--run-id', + '-r', envvar='TWINKLE_AUTO_RUN_ID', help='Attach to an existing training run by ID.', ), llm_base_url: str = typer.Option( - 'http://localhost:11434/v1', '--llm-base-url', + 'http://localhost:11434/v1', + '--llm-base-url', envvar='TWINKLE_LLM_BASE_URL', help='LLM API base URL.', ), llm_model: str = typer.Option( - 'qwen3.5', '--llm-model', + 'qwen3.5', + '--llm-model', envvar='TWINKLE_LLM_MODEL', help='LLM model name.', ), llm_api_key: str = typer.Option( - 'not-needed', '--llm-api-key', + 'not-needed', + '--llm-api-key', envvar='TWINKLE_LLM_API_KEY', help='LLM API key.', ), verbose: bool = typer.Option( - False, '--verbose', '-v', + False, + '--verbose', + '-v', envvar='TWINKLE_AUTO_VERBOSE', help='Enable verbose (DEBUG) logging.', ), version: bool = typer.Option( - False, '--version', '-V', - callback=_version_callback, is_eager=True, + False, + '--version', + '-V', + callback=_version_callback, + is_eager=True, help='Show version and exit.', ), ) -> None: """Launch Twinkle Auto.""" _configure_logging(verbose=verbose) logger = get_logger() - logger.info( - f'Auto starting — model={llm_model}, base_url={llm_base_url}, ' - f'run_id={run_id}, log_file={_LOG_FILE}' - ) + logger.info(f'Auto starting — model={llm_model}, base_url={llm_base_url}, ' + f'run_id={run_id}, log_file={_LOG_FILE}') from twinkle_client.auto.app import TwinkleAuto diff --git a/src/twinkle_client/auto/agent/core.py b/src/twinkle_client/auto/agent/core.py index a3c3c5b1c..88c387f31 100644 --- a/src/twinkle_client/auto/agent/core.py +++ b/src/twinkle_client/auto/agent/core.py @@ -5,15 +5,18 @@ import asyncio import json -from twinkle.utils.logger import get_logger -from typing import Any, Callable +from typing import TYPE_CHECKING, Any, Callable +from twinkle.utils.logger import get_logger from twinkle_client.auto.agent.prompts import SYSTEM_PROMPT from twinkle_client.auto.agent.tools import TOOL_SCHEMAS, ToolExecutor from twinkle_client.auto.connection import LocalConnection logger = get_logger() +if TYPE_CHECKING: + from openai import AsyncOpenAI + class AgentLoop: """Async tool-calling agent loop using OpenAI-compatible API. @@ -28,7 +31,7 @@ class AgentLoop: def __init__( self, connection: LocalConnection, - llm_client: 'AsyncOpenAI', + llm_client: AsyncOpenAI, llm_model: str, skills_prompt: str = '', ): @@ -41,7 +44,10 @@ def __init__( if skills_prompt: full_prompt = f'{SYSTEM_PROMPT}\n\n{skills_prompt}' self.history: list[dict[str, Any]] = [ - {'role': 'system', 'content': full_prompt}, + { + 'role': 'system', + 'content': full_prompt + }, ] async def send( @@ -94,7 +100,8 @@ async def send( except json.JSONDecodeError as e: logger.error(f'Tool {func_name}: invalid JSON args: {e}\n raw={raw_args[:500]}') args = {} - logger.info(f'Executing tool: {func_name}({", ".join(f"{k}={v!r}" for k, v in list(args.items())[:5])})') + logger.info( + f'Executing tool: {func_name}({", ".join(f"{k}={v!r}" for k, v in list(args.items())[:5])})') result = await self._tool_executor.execute(func_name, args) logger.debug(f'Tool {func_name} result ({len(result)} chars): {result[:300]}') self.history.append({ @@ -156,7 +163,10 @@ async def _call_llm_stream( tool_calls_map[idx] = { 'id': '', 'type': 'function', - 'function': {'name': '', 'arguments': ''}, + 'function': { + 'name': '', + 'arguments': '' + }, } tc = tool_calls_map[idx] if tc_delta.id: diff --git a/src/twinkle_client/auto/agent/monitor.py b/src/twinkle_client/auto/agent/monitor.py index b25169f00..34f2f415e 100644 --- a/src/twinkle_client/auto/agent/monitor.py +++ b/src/twinkle_client/auto/agent/monitor.py @@ -19,13 +19,16 @@ import re import time from pathlib import Path -from typing import Any, Callable +from typing import TYPE_CHECKING, Any, Callable from twinkle.utils.logger import get_logger from twinkle_client.auto.connection import LocalConnection logger = get_logger() +if TYPE_CHECKING: + from openai import AsyncOpenAI + # Maximum auto-fix attempts per run (prevent infinite retry loops) _MAX_FIX_ATTEMPTS = 3 @@ -95,7 +98,7 @@ def __init__( self, connection: LocalConnection, on_message: Callable[[str], None], - llm_client: 'AsyncOpenAI', + llm_client: AsyncOpenAI, llm_model: str = 'qwen3.5', poll_interval: float = 30.0, ): @@ -305,8 +308,14 @@ async def _ask_llm(self, snapshot: dict[str, Any]) -> str | None: response = await self._client.chat.completions.create( model=self.llm_model, messages=[ - {'role': 'system', 'content': MONITOR_SYSTEM_PROMPT + extra}, - {'role': 'user', 'content': user_content}, + { + 'role': 'system', + 'content': MONITOR_SYSTEM_PROMPT + extra + }, + { + 'role': 'user', + 'content': user_content + }, ], temperature=0.3, max_tokens=4096, @@ -340,10 +349,8 @@ async def _apply_fix(self, run_id: str, diagnosis: str, fixed_script: str) -> No """Apply auto-fix: update script + resume training.""" attempts = self._fix_attempts.get(run_id, 0) if attempts >= _MAX_FIX_ATTEMPTS: - self.on_message( - f'[Monitor] 已达最大自动修复次数 ({_MAX_FIX_ATTEMPTS}),不再尝试。' - '请手动检查或输入指令。' - ) + self.on_message(f'[Monitor] 已达最大自动修复次数 ({_MAX_FIX_ATTEMPTS}),不再尝试。' + '请手动检查或输入指令。') return self.on_message(f'[Monitor] 检测到问题,正在自动修复 (第{attempts + 1}次)...\n诊断: {diagnosis}') @@ -390,7 +397,7 @@ def _parse_fix_response(response: str) -> tuple[str, str]: else: # Fallback: text before the python block before = response[:response.find('```python')] - lines = [l.strip() for l in before.splitlines() if l.strip() and not l.startswith('```')] + lines = [line.strip() for line in before.splitlines() if line.strip() and not line.startswith('```')] diagnosis = lines[-1] if lines else 'Auto-fix applied' return diagnosis, fixed_script diff --git a/src/twinkle_client/auto/agent/search_tools.py b/src/twinkle_client/auto/agent/search_tools.py new file mode 100644 index 000000000..4bbb652da --- /dev/null +++ b/src/twinkle_client/auto/agent/search_tools.py @@ -0,0 +1,63 @@ +# Copyright (c) Twinkle Contributors. All rights reserved. +"""Private ModelScope Hub search tools used by ``ToolExecutor``.""" +from __future__ import annotations + +import asyncio + + +class _SearchTools: + + async def _tool_search_datasets(self, query: str, limit: int = 5) -> dict: + """Search ModelScope for datasets.""" + return await self._search_hub('datasets', query, limit) + + async def _tool_search_models(self, query: str, limit: int = 5) -> dict: + """Search ModelScope for models.""" + return await self._search_hub('models', query, limit) + + async def _search_hub(self, resource_type: str, query: str, limit: int) -> dict: + """Unified ModelScope Hub search for models or datasets.""" + + def _search(): + if resource_type == 'datasets': + return self._search_datasets_impl(query, limit) + else: + return self._search_models_impl(query, limit) + + try: + items = await asyncio.get_event_loop().run_in_executor(None, _search) + return {'query': query, 'results': items} + except Exception as e: + return {'error': f'{resource_type.title()} search failed: {e}'} + + @staticmethod + def _search_datasets_impl(query: str, limit: int) -> list[dict]: + """Search datasets via ModelScope SDK (new API).""" + from modelscope.hub.api import HubApi + api = HubApi() + result = api.list_datasets('', search=query, page_size=limit) + datasets = result.get('datasets', []) + return [{'id': d.get('id', ''), 'name': d.get('display_name', d.get('id', ''))} for d in datasets] + + @staticmethod + def _search_models_impl(query: str, limit: int) -> list[dict]: + """Search models via ModelScope HTTP API (SDK doesn't support search).""" + import requests + resp = requests.put( + 'https://modelscope.cn/api/v1/models/', + json={ + 'Name': query, + 'PageSize': limit, + 'PageNumber': 1 + }, + timeout=15, + ) + resp.raise_for_status() + data = resp.json() + if not data.get('Success'): + raise RuntimeError(data.get('Message', 'Unknown error')) + models = data.get('Data', {}).get('Models', []) + return [{ + 'id': f"{m.get('Path', '')}/{m.get('Name', '')}", + 'name': m.get('ChineseName') or m.get('Name', ''), + } for m in models] diff --git a/src/twinkle_client/auto/agent/server_tools.py b/src/twinkle_client/auto/agent/server_tools.py new file mode 100644 index 000000000..d73f5c941 --- /dev/null +++ b/src/twinkle_client/auto/agent/server_tools.py @@ -0,0 +1,803 @@ +# Copyright (c) Twinkle Contributors. All rights reserved. +"""Private server lifecycle and cluster tools used by ``ToolExecutor``.""" +from __future__ import annotations + +import asyncio +import json +import os + + +class _ServerTools: + """Server-side operations mixed into ``ToolExecutor``. + + ``ToolExecutor`` owns the URL value; declaring it here makes that host + requirement visible without assigning runtime state in the mixin. + """ + + _server_url: str | None + + async def _check_server_health(self, url: str) -> bool: + """Check if Twinkle Server is reachable (non-blocking).""" + import urllib.error + import urllib.request + + def _probe(): + try: + req = urllib.request.Request(f'{url}/api/v1/healthz', method='GET') + urllib.request.urlopen(req, timeout=3) + return True + except (urllib.error.URLError, OSError): + # Try a simpler connectivity check + try: + urllib.request.urlopen(url, timeout=3) + return True + except (urllib.error.URLError, OSError): + return False + + return await asyncio.get_event_loop().run_in_executor(None, _probe) + + async def _tool_start_server( + self, + model_id: str, + train_gpus: int | None = None, + port: int = 8000, + backend: str = 'transformers', + samplers: list[dict] | None = None, + ) -> dict: + """Start Ray cluster + Twinkle Server. Idempotent. Supports multi-model.""" + server_url = self._server_url or os.environ.get('TWINKLE_SERVER_URL') or f'http://localhost:{port}' + + # Idempotent: skip if already running + if await self._check_server_health(server_url): + self._server_url = server_url + return {'status': 'already_running', 'server_url': server_url} + + def _start(): + sampler_list = samplers or [] + + # Step 1: Detect hardware & compute GPU partition + total_hw_gpus = self._detect_gpu_count() + if total_hw_gpus == 0: + return {'status': 'error', 'error': 'No GPUs detected. Cannot start training server.'} + + alloc = self._compute_gpu_allocation(sampler_list, train_gpus, total_hw_gpus) + if 'error' in alloc: + return {'status': 'error', 'error': alloc['error']} + t_gpus, sampler_gpu_total = alloc['train_gpus'], alloc['sampler_gpus'] + + # Step 2: Generate server_config.yaml + config_path = self._generate_server_config( + model_id=model_id, + train_gpus=t_gpus, + port=port, + backend=backend, + samplers=sampler_list, + ) + + # Step 3: Start Ray cluster (multi-node GPU partitioning) + ray_err = self._start_ray_cluster(t_gpus, sampler_gpu_total) + if ray_err: + return {'status': 'error', 'error': ray_err} + + # Step 4: Launch Twinkle Server process + proc, log_path, err = self._launch_server_process(config_path) + if err: + return {'status': 'error', 'error': err} + + # Step 5: Wait for readiness (healthz + sampler engine) + return self._wait_server_ready( + server_url=server_url, + proc=proc, + log_path=log_path, + sampler_list=sampler_list, + model_id=model_id, + t_gpus=t_gpus, + backend=backend, + config_path=config_path, + ) + + result = await asyncio.get_event_loop().run_in_executor(None, _start) + if result.get('status') in ('started', 'already_running'): + self._server_url = server_url + return result + + @staticmethod + def _detect_gpu_count() -> int: + """Detect total hardware GPU count via nvidia-smi.""" + import subprocess as _sp + try: + r = _sp.run( + ['nvidia-smi', '--query-gpu=index', '--format=csv,noheader'], + capture_output=True, + text=True, + timeout=10, + ) + if r.returncode == 0: + return len([ln for ln in r.stdout.strip().split('\n') if ln.strip()]) + except (FileNotFoundError, OSError): + pass + return 0 + + @staticmethod + def _compute_gpu_allocation( + sampler_list: list[dict], + train_gpus: int | None, + total_hw_gpus: int, + ) -> dict: + """Compute GPU partition: {train_gpus, sampler_gpus} or {error}.""" + sampler_gpu_total = 0 + for s in sampler_list: + s_tp = s.get('tp', 1) + s_dp, s_gpus = s.get('dp'), s.get('gpus') + if s_gpus is not None: + sampler_gpu_total += s_gpus + elif s_dp is not None: + sampler_gpu_total += s_tp * s_dp + else: + sampler_gpu_total += s_tp # default dp=1 + + t_gpus = train_gpus if train_gpus is not None else max(1, total_hw_gpus - sampler_gpu_total) + needed = t_gpus + sampler_gpu_total + if needed > total_hw_gpus: + return { + 'error': (f'Requested {needed} GPUs (train={t_gpus}, samplers={sampler_gpu_total}) ' + f'but only {total_hw_gpus} available.'), + } + return {'train_gpus': t_gpus, 'sampler_gpus': sampler_gpu_total} + + @staticmethod + def _start_ray_cluster(train_gpus: int, sampler_gpus: int) -> str | None: + """Start Ray multi-node cluster with GPU partitioning. + + Each role gets its own Ray node with dedicated CUDA_VISIBLE_DEVICES + so GPUs are indexed from 0 within each node. This prevents the + GPU ID mapping issues that occur with a single-node setup. + + On a single machine, multiple raylets need separate --temp-dir to + avoid being detected as "already running". + + Returns an error message on failure, or None on success. + """ + import subprocess as _sp + import tempfile + from pathlib import Path + + _sp.run(['ray', 'stop', '--force'], capture_output=True, timeout=15) + + # Create unique temp dirs so each `ray start` spawns a separate raylet + ray_base = Path(tempfile.gettempdir()) / 'twinkle_ray' + ray_base.mkdir(parents=True, exist_ok=True) + + def _ray_node( + devices: str, + num_gpus: int, + *, + head: bool = False, + node_name: str = 'worker', + ) -> str | None: + env = os.environ.copy() + env['CUDA_VISIBLE_DEVICES'] = devices + temp_dir = str(ray_base / node_name) + cmd = ['ray', 'start', f'--temp-dir={temp_dir}'] + if head: + cmd += ['--head', '--port=6379', '--disable-usage-stats', '--include-dashboard=false'] + else: + cmd += ['--address=127.0.0.1:6379'] + cmd.append(f'--num-gpus={num_gpus}') + r = _sp.run(cmd, capture_output=True, text=True, timeout=30, env=env) + if r.returncode != 0 and 'already' not in r.stderr.lower(): + return r.stderr.strip() + return None + + # Head node — training model GPUs + model_devices = ','.join(str(i) for i in range(train_gpus)) + err = _ray_node(model_devices, train_gpus, head=True, node_name='head') + if err: + return f'Ray head start failed: {err}' + + # GPU Worker node — sampler GPUs + if sampler_gpus > 0: + sampler_devices = ','.join(str(i) for i in range(train_gpus, train_gpus + sampler_gpus)) + err = _ray_node(sampler_devices, sampler_gpus, node_name='gpu_worker') + if err: + return f'Ray GPU worker start failed: {err}' + + # CPU Worker node — processor (no GPU) + _ray_node('', 0, node_name='cpu_worker') + return None + + @staticmethod + def _launch_server_process(config_path: str) -> tuple: + """Launch Twinkle Server as a detached background process. + + Returns (proc, log_path, error). On success error is None. + """ + import subprocess as _sp + from pathlib import Path + + log_dir = Path.home() / '.cache' / 'twinkle' + log_dir.mkdir(parents=True, exist_ok=True) + log_path = str(log_dir / 'server.log') + log_file = open(log_path, 'w') + + cmd = ['python', '-m', 'twinkle.server', 'launch', '--config', config_path] + try: + proc = _sp.Popen( + cmd, + stdout=log_file, + stderr=_sp.STDOUT, + start_new_session=True, + ) + except OSError as e: + log_file.close() + return None, log_path, f'Failed to start Twinkle server: {e}' + return proc, log_path, None + + @staticmethod + def _wait_server_ready( + server_url: str, + proc, + log_path: str, + sampler_list: list[dict], + model_id: str, + t_gpus: int, + backend: str, + config_path: str, + ) -> dict: + """Poll server until healthy (healthz + sampler engine ready).""" + import time + import urllib.error + import urllib.request + + timeout_s = 120 if sampler_list else 60 + needed = t_gpus + sum(s.get('gpus') or (s.get('tp', 1) * s.get('dp', 1)) for s in sampler_list) + + for _ in range(timeout_s): + time.sleep(1) + if proc.poll() is not None: + # Server died — read log tail to diagnose + log_tail = _ServerTools._read_log_tail(log_path, max_chars=2000) + error_msg = (f'Server exited immediately (code={proc.returncode}). ' + f'Model: {model_id}, GPUs: {t_gpus}, Samplers: {len(sampler_list)}.\n' + f'--- server.log tail ---\n{log_tail}') + return { + 'status': 'error', + 'error': error_msg, + 'log_path': log_path, + 'hint': 'Check if required packages are installed (pip install -e ".[all]").', + } + try: + urllib.request.urlopen(f'{server_url}/api/v1/healthz', timeout=2) + except (OSError, Exception): + continue + + # healthz OK — additionally wait for sampler vLLM engines + if sampler_list and not _ServerTools._probe_sampler_ready(server_url, sampler_list, model_id): + return { + 'status': 'started', + 'warning': 'Server is up but sampler may still be loading.', + 'server_url': server_url, + 'server_pid': proc.pid, + 'model_id': model_id, + 'log_path': log_path, + } + + return { + 'status': 'started', + 'server_url': server_url, + 'server_pid': proc.pid, + 'model_id': model_id, + 'train_gpus': t_gpus, + 'backend': backend, + 'samplers': [s.get('model_id') for s in sampler_list], + 'total_gpus_used': needed, + 'config_path': config_path, + 'log_path': log_path, + } + + return { + 'status': 'timeout', + 'error': 'Health check did not pass within timeout. Models may still be loading.', + 'server_pid': proc.pid, + 'log_path': log_path, + } + + @staticmethod + def _read_log_tail(log_path: str, max_chars: int = 2000) -> str: + """Read the tail of a log file for error diagnosis.""" + try: + with open(log_path, errors='replace') as f: + content = f.read() + if len(content) <= max_chars: + return content.strip() + return content[-max_chars:].strip() + except OSError: + return '(could not read log file)' + + @staticmethod + def _probe_sampler_ready(server_url: str, sampler_list: list[dict], fallback_model_id: str) -> bool: + """Probe sampler route up to 90s to confirm vLLM engine is loaded.""" + import time + import urllib.error + import urllib.request + + s_mid = sampler_list[0].get('model_id', fallback_model_id) + probe_url = f'{server_url}/api/v1/sampler/{s_mid}/twinkle/create' + + for _ in range(90): + try: + req = urllib.request.Request( + probe_url, + method='POST', + data=b'{}', + headers={'Content-Type': 'application/json'}, + ) + urllib.request.urlopen(req, timeout=5) + return True # non-error response = ready + except urllib.error.HTTPError as e: + if e.code < 500: + return True # 4xx = actor alive, just bad request + time.sleep(1) # 5xx = still loading + except (OSError, Exception): + time.sleep(1) + return False + + @staticmethod + def _generate_server_config( + model_id: str, + train_gpus: int, + port: int = 8000, + backend: str = 'transformers', + samplers: list[dict] | None = None, + ) -> str: + """Generate a server_config.yaml from template and return its path. + + Supports multi-model topology: + - 1 training model (student) + - N sampler/teacher models (for RL/OPD) + - 1 processor service + """ + import yaml + from pathlib import Path + + sampler_list = samplers or [] + + # Sanitize model name for use in route/names + def _short(mid: str) -> str: + return mid.split('/')[-1] if '/' in mid else mid + + model_short = _short(model_id) + + # Collect all model IDs for supported_models + all_model_ids = [model_id] + [s['model_id'] for s in sampler_list] + + # === Build applications list === + applications = [] + + # 1. API Gateway + applications.append({ + 'name': + 'server', + 'route_prefix': + '/api/v1', + 'import_path': + 'server', + 'args': { + 'server_config': { + 'per_token_model_limit': 3 + }, + 'supported_models': all_model_ids, + }, + 'deployments': [{ + 'name': 'TinkerCompatServer', + 'max_ongoing_requests': 50, + 'autoscaling_config': { + 'min_replicas': 1, + 'max_replicas': 1, + 'target_ongoing_requests': 128, + }, + 'ray_actor_options': { + 'num_cpus': 0.1 + }, + }], + }) + + # 2. Build GPU-requiring applications (model + samplers), + # then sort by GPU count DESCENDING before appending. + # Largest PG deploys first → it has the fewest node choices → + # avoids GPU scheduling deadlock on single-machine multi-node. + gpu_apps: list[tuple[int, dict]] = [] # (gpu_count, app_config) + + # 2a. Training model worker (student) + gpu_apps.append(( + train_gpus, + { + 'name': + f'models-{model_short}', + 'route_prefix': + f'/api/v1/model/{model_id}', + 'import_path': + 'model', + 'args': { + 'backend': backend, + 'model_id': f'ms://{model_id}', + 'max_length': 500000, # total tokens per forward pass (must match max_input_tokens) + 'nproc_per_node': train_gpus, + 'device_group': { + 'name': 'model', + 'ranks': train_gpus, + 'device_type': 'cuda', + }, + 'device_mesh': { + 'device_type': 'cuda', + 'dp_size': train_gpus, + }, + 'queue_config': { + 'rps_limit': 100, + 'tps_limit': 100000, + 'max_input_tokens': 500000, + }, + 'adapter_config': { + 'adapter_timeout': 600, + }, + }, + 'deployments': [{ + 'name': 'ModelManagement', + 'autoscaling_config': { + 'min_replicas': 1, + 'max_replicas': 1, + 'target_ongoing_requests': 16, + }, + 'ray_actor_options': { + 'num_cpus': 0.1, + 'runtime_env': { + 'env_vars': { + 'TWINKLE_TRUST_REMOTE_CODE': '1' + }, + }, + }, + }], + })) + + # 2b. Sampler/teacher models + sampler_name_count: dict[str, int] = {} + for sampler_cfg in sampler_list: + s_model_id = sampler_cfg['model_id'] + s_short = _short(s_model_id) + + # Deduplicate names when multiple samplers share the same short name + sampler_name_count[s_short] = sampler_name_count.get(s_short, 0) + 1 + if sampler_name_count[s_short] > 1: + s_name = f'sampler-{s_short}-{sampler_name_count[s_short]}' + else: + s_name = f'sampler-{s_short}' + + s_engine = sampler_cfg.get('engine', 'vllm') + s_max_len = sampler_cfg.get('max_model_len', 16000) + + # Compute tp / dp / total GPUs: + # tp = tensor parallelism (GPUs per vLLM process, for large models) + # dp = data parallelism (number of independent inference replicas) + # total GPUs = tp * dp + s_tp = sampler_cfg.get('tp', 1) + s_dp = sampler_cfg.get('dp', None) + s_gpus = sampler_cfg.get('gpus', None) + + if s_dp is not None and s_gpus is not None: + # Both specified: validate consistency + s_tp = s_gpus // s_dp if s_tp == 1 else s_tp + elif s_gpus is not None: + # Only total GPUs specified: derive dp + s_dp = max(1, s_gpus // s_tp) + elif s_dp is not None: + # Only dp specified: derive total + s_gpus = s_tp * s_dp + else: + # Nothing specified: default to 1 GPU (tp=1, dp=1) + s_dp = 1 + s_gpus = s_tp * s_dp + + s_total_gpus = s_tp * s_dp + + # Build device_mesh: include tp_size when tp>1 so that + # world_size = tp*dp and slice_dp dispatch computes correct + # rank_stride for DP data sharding. + mesh_config: dict = {'device_type': 'cuda', 'dp_size': s_dp} + if s_tp > 1: + mesh_config['tp_size'] = s_tp + + sampler_app: dict = { + 'name': + s_name, + 'route_prefix': + f'/api/v1/sampler/{s_model_id}', + 'import_path': + 'sampler', + 'args': { + 'model_id': f'ms://{s_model_id}', + 'nproc_per_node': s_total_gpus, + 'sampler_type': s_engine, + 'device_group': { + 'name': s_name, + 'ranks': s_total_gpus, + 'device_type': 'cuda', + 'gpus_per_worker': s_tp, + }, + 'device_mesh': mesh_config, + 'queue_config': { + 'rps_limit': 100, + 'tps_limit': 100000, + }, + }, + 'deployments': [{ + 'name': 'SamplerManagement', + 'autoscaling_config': { + 'min_replicas': 1, + 'max_replicas': 1, + 'target_ongoing_requests': 16, + }, + 'ray_actor_options': { + 'num_cpus': 0.1, + 'runtime_env': { + 'env_vars': { + 'TWINKLE_TRUST_REMOTE_CODE': '1' + }, + }, + }, + }], + } + + # Add engine-specific args + if s_engine == 'vllm': + engine_args = { + 'max_model_len': s_max_len, + 'gpu_memory_utilization': 0.85, + 'enable_lora': True, + 'logprobs_mode': 'processed_logprobs', + } + # Set tensor_parallel_size when tp > 1 + if s_tp > 1: + engine_args['tensor_parallel_size'] = s_tp + sampler_app['args']['engine_args'] = engine_args + + gpu_apps.append((s_total_gpus, sampler_app)) + + # 3. Sort GPU apps by GPU count DESCENDING, then append in order. + # Largest PG deploys first → claims the largest node → avoids deadlock. + gpu_apps.sort(key=lambda x: x[0], reverse=True) + for _, app_cfg in gpu_apps: + applications.append(app_cfg) + + # 4. Processor service + applications.append({ + 'name': + 'processor', + 'route_prefix': + '/api/v1/processor', + 'import_path': + 'processor', + 'args': { + 'ncpu_proc_per_node': 2, + 'device_group': { + 'name': 'processor', + 'ranks': 2, + 'device_type': 'CPU', + }, + 'device_mesh': { + 'device_type': 'CPU', + 'dp_size': 2, + }, + }, + 'deployments': [{ + 'name': 'ProcessorManagement', + 'autoscaling_config': { + 'min_replicas': 1, + 'max_replicas': 1, + 'target_ongoing_requests': 128, + }, + 'ray_actor_options': { + 'num_cpus': 0.1 + }, + }], + }) + + # === Assemble final config === + config = { + 'proxy_location': 'EveryNode', + 'http_options': { + 'host': '0.0.0.0', + 'port': port, + }, + 'applications': applications, + } + + # Write to ~/.cache/twinkle/server_config.yaml + config_dir = Path.home() / '.cache' / 'twinkle' + config_dir.mkdir(parents=True, exist_ok=True) + config_path = config_dir / 'server_config.yaml' + with open(config_path, 'w') as f: + yaml.dump(config, f, default_flow_style=False, allow_unicode=True) + + return str(config_path) + + async def _tool_shutdown_server(self) -> dict: + """Shut down Twinkle Server and Ray cluster. DESTROYS GPU model state.""" + import subprocess as _sp + + def _shutdown(): + results = {} + + # 1. Try `serve shutdown` to cleanly stop Ray Serve deployments + try: + r = _sp.run(['serve', 'shutdown', '-y'], capture_output=True, text=True, timeout=30) + results['serve_shutdown'] = 'ok' if r.returncode == 0 else r.stderr.strip() + except (FileNotFoundError, OSError) as e: + results['serve_shutdown'] = f'skipped: {e}' + + # 2. Kill any remaining twinkle.server processes + try: + _sp.run(['pkill', '-f', 'twinkle.server'], capture_output=True, timeout=5) + except (FileNotFoundError, OSError): + pass + + # 3. Stop Ray cluster + try: + r = _sp.run(['ray', 'stop', '--force'], capture_output=True, text=True, timeout=15) + results['ray_stop'] = 'ok' if r.returncode == 0 else r.stderr.strip() + except (FileNotFoundError, OSError) as e: + results['ray_stop'] = f'failed: {e}' + + results['status'] = 'shutdown_complete' + results['warning'] = 'All GPU model state has been released.' + return results + + return await asyncio.get_event_loop().run_in_executor(None, _shutdown) + + async def _tool_list_supported_models(self, base_url: str | None = None) -> dict: + """Query the Twinkle server for supported models.""" + url = base_url or self._server_url or os.environ.get('TWINKLE_SERVER_URL') or 'http://localhost:8000' + + def _query(): + # Use a lightweight HTTP GET instead of init_twinkle_client() which + # creates a session + heartbeat thread that would leak since we never + # call close(). + import urllib.error + import urllib.request + + endpoint = f'{url}/api/v1/twinkle/get_server_capabilities' + req = urllib.request.Request(endpoint, method='GET') + resp = urllib.request.urlopen(req, timeout=10) + data = json.loads(resp.read().decode()) + models = data.get('supported_models', []) + # Each model entry may be a dict with 'model_name' or a plain string + model_names = [] + for m in models: + if isinstance(m, dict): + model_names.append(m.get('model_name', '')) + else: + model_names.append(str(m)) + return { + 'base_url': url, + 'supported_models': model_names, + } + + try: + return await asyncio.get_event_loop().run_in_executor(None, _query) + except Exception as e: + return {'error': f'Failed to query {url}: {e}'} + + async def _tool_get_cluster_info(self) -> dict: + """Query cluster resources: try Ray first, fall back to nvidia-smi.""" + + def _query(): + # 1. Try connecting to an existing Ray cluster + ray_info = self._try_ray_cluster() + if ray_info is not None: + ray_info['ray_active'] = True + return ray_info + + # 2. Ray not available — fall back to nvidia-smi + nvidia_info = self._try_nvidia_smi() + nvidia_info['ray_active'] = False + nvidia_info['hint'] = ('Ray cluster is not running. To use distributed training, ' + 'start Ray first: `ray start --head --num-gpus=N` or use ' + 'the server mode run.sh script.') + return nvidia_info + + return await asyncio.get_event_loop().run_in_executor(None, _query) + + @staticmethod + def _try_ray_cluster() -> dict | None: + """Attempt to query an existing Ray cluster. Returns None if unavailable.""" + try: + import ray + except ImportError: + return None + + import logging as _logging + + try: + if not ray.is_initialized(): + ray.init( + address='auto', + ignore_reinit_error=True, + _timeout_s=5, + logging_level=_logging.ERROR, + configure_logging=False, + ) + + resources = ray.cluster_resources() + available = ray.available_resources() + nodes = ray.nodes() + gpu_total = resources.get('GPU', 0) + gpu_available = available.get('GPU', 0) + gpu_types = set() + for node in nodes: + for key in node.get('Resources', {}): + if key.startswith('accelerator_type:'): + gpu_types.add(key.split(':', 1)[1]) + return { + 'num_nodes': len([n for n in nodes if n.get('Alive')]), + 'gpu_total': int(gpu_total), + 'gpu_available': int(gpu_available), + 'gpu_types': sorted(gpu_types) if gpu_types else ['unknown'], + 'cpu_total': resources.get('CPU', 0), + 'memory_bytes': resources.get('memory', 0), + } + except Exception: + try: + import ray as _ray + if _ray.is_initialized(): + _ray.shutdown() + except Exception: + pass + return None + + @staticmethod + def _try_nvidia_smi() -> dict: + """Parse nvidia-smi output for local GPU info.""" + import subprocess as _sp + + try: + result = _sp.run( + [ + 'nvidia-smi', '--query-gpu=index,name,memory.total,memory.free,utilization.gpu', + '--format=csv,noheader,nounits' + ], + capture_output=True, + text=True, + timeout=10, + ) + if result.returncode != 0: + return {'error': f'nvidia-smi failed: {result.stderr.strip()}', 'gpu_total': 0} + + gpus = [] + for line in result.stdout.strip().split('\n'): + if not line.strip(): + continue + parts = [p.strip() for p in line.split(',')] + if len(parts) >= 5: + try: + gpus.append({ + 'index': int(parts[0]), + 'name': parts[1], + 'memory_total_mb': int(parts[2]), + 'memory_free_mb': int(parts[3]), + 'utilization_pct': int(parts[4]) if parts[4].isdigit() else 0, + }) + except (ValueError, IndexError): + # Skip lines with unparseable values (e.g. [N/A]) + continue + + gpu_types = sorted({g['name'] for g in gpus}) + return { + 'gpu_total': len(gpus), + 'gpu_available': len([g for g in gpus if g['utilization_pct'] < 10]), + 'gpu_types': gpu_types if gpu_types else ['none'], + 'gpus': gpus, + 'source': 'nvidia-smi', + } + except FileNotFoundError: + return {'error': 'nvidia-smi not found (no NVIDIA GPU or driver not installed)', 'gpu_total': 0} + except Exception as e: + return {'error': f'nvidia-smi query failed: {e}', 'gpu_total': 0} diff --git a/src/twinkle_client/auto/agent/tool_schemas.py b/src/twinkle_client/auto/agent/tool_schemas.py new file mode 100644 index 000000000..7db691fa9 --- /dev/null +++ b/src/twinkle_client/auto/agent/tool_schemas.py @@ -0,0 +1,338 @@ +# Copyright (c) Twinkle Contributors. All rights reserved. +"""OpenAI function-calling schemas exposed by the auto agent.""" +from __future__ import annotations + +from typing import Any + +TOOL_SCHEMAS: list[dict[str, Any]] = [ + { + 'type': 'function', + 'function': { + 'name': 'list_training_runs', + 'description': 'List all active and historical training runs.', + 'parameters': { + 'type': 'object', + 'properties': {}, + 'required': [] + }, + }, + }, + { + 'type': 'function', + 'function': { + 'name': 'get_training_status', + 'description': 'Get detailed status and recent metrics for a training run.', + 'parameters': { + 'type': 'object', + 'properties': { + 'run_id': { + 'type': 'string', + 'description': 'Training run ID.' + }, + }, + 'required': ['run_id'], + }, + }, + }, + { + 'type': 'function', + 'function': { + 'name': + 'start_server', + 'description': ('Start Ray cluster and Twinkle Server. MUST be called before start_training. ' + 'Idempotent: skips if server is already reachable. ' + 'Supports multi-model deployments: one training model + N sampler/teacher models. ' + 'Automatically generates server_config.yaml from parameters.'), + 'parameters': { + 'type': 'object', + 'properties': { + 'model_id': { + 'type': 'string', + 'description': 'Student/training model ID (e.g. "Qwen/Qwen3.5-4B").', + }, + 'train_gpus': { + 'type': 'integer', + 'description': 'GPUs for the training model. Default: auto-detect remaining GPUs.', + }, + 'backend': { + 'type': 'string', + 'enum': ['transformers', 'megatron'], + 'description': 'Training model backend. Default: transformers.', + }, + 'samplers': { + 'type': + 'array', + 'description': ('List of sampler/teacher models for RL/OPD. Each entry deploys ' + 'an inference service (vLLM or torch). Omit for simple SFT.'), + 'items': { + 'type': 'object', + 'properties': { + 'model_id': { + 'type': 'string', + 'description': 'Teacher/reference model ID (e.g. "Qwen/Qwen3.5-72B").', + }, + 'gpus': { + 'type': 'integer', + 'description': + 'Total number of GPUs for this sampler. Default: 1. Must equal tp * dp.', + }, + 'tp': { + 'type': + 'integer', + 'description': + ('Tensor parallelism size (GPUs per vLLM worker process). ' + 'Use tp>1 for large models that do not fit on a single GPU. Default: 1.'), + }, + 'dp': { + 'type': + 'integer', + 'description': ('Data parallelism size (number of independent inference replicas). ' + 'If not specified, computed as gpus // tp. Default: 1.'), + }, + 'engine': { + 'type': 'string', + 'enum': ['vllm', 'torch'], + 'description': 'Inference engine. Default: vllm.', + }, + 'max_model_len': { + 'type': 'integer', + 'description': 'Max sequence length for inference. Default: 16000.', + }, + }, + 'required': ['model_id'], + }, + }, + 'port': { + 'type': 'integer', + 'description': 'HTTP port for server. Default: 8000.', + }, + }, + 'required': ['model_id'], + }, + }, + }, + { + 'type': 'function', + 'function': { + 'name': + 'shutdown_server', + 'description': ('Shut down Twinkle Server and Ray cluster. WARNING: This releases all GPU resources ' + 'and DESTROYS model state held in server memory. Only call when training is truly ' + 'finished and you no longer need the server. Model weights/optimizer state in GPU ' + 'will be LOST unless a checkpoint was explicitly saved.'), + 'parameters': { + 'type': 'object', + 'properties': {}, + 'required': [] + }, + }, + }, + { + 'type': 'function', + 'function': { + 'name': + 'start_training', + 'description': ('Create a new training run: write the client script, launch it, and start monitoring. ' + 'REQUIRES: Twinkle Server must be running (call start_server first). ' + 'The client script connects to the server — server holds model state in GPU memory. ' + 'Kill client = pause (state preserved). Re-launch client = resume.'), + 'parameters': { + 'type': 'object', + 'properties': { + 'run_id': { + 'type': 'string', + 'description': 'Unique run ID (e.g., "grpo-gsm8k").' + }, + 'script_content': { + 'type': 'string', + 'description': 'Full Python source code of the training script.' + }, + 'model_id': { + 'type': 'string', + 'description': 'Model identifier for metadata (e.g., "Qwen/Qwen3.5-4B").' + }, + }, + 'required': ['run_id', 'script_content'], + }, + }, + }, + { + 'type': 'function', + 'function': { + 'name': 'select_run', + 'description': 'Switch to monitor a different training run. Updates connection context.', + 'parameters': { + 'type': 'object', + 'properties': { + 'run_id': { + 'type': 'string', + 'description': 'Training run ID to monitor.' + }, + }, + 'required': ['run_id'], + }, + }, + }, + { + 'type': 'function', + 'function': { + 'name': + 'pause_training', + 'description': ('Pause training by killing the client process (SIGKILL). ' + 'Server retains all state — call resume_training to continue.'), + 'parameters': { + 'type': 'object', + 'properties': { + 'run_id': { + 'type': 'string', + 'description': 'Training run ID to pause.' + }, + }, + 'required': ['run_id'], + }, + }, + }, + { + 'type': 'function', + 'function': { + 'name': 'resume_training', + 'description': 'Resume a paused training run by re-launching the client script. Server state is preserved.', + 'parameters': { + 'type': 'object', + 'properties': { + 'run_id': { + 'type': 'string', + 'description': 'Training run ID to resume.' + }, + }, + 'required': ['run_id'], + }, + }, + }, + { + 'type': 'function', + 'function': { + 'name': + 'stop_training', + 'description': ('Gracefully stop the training client (SIGTERM). The script saves a checkpoint ' + 'before exiting. Server retains model/optimizer state in GPU memory — ' + 'use resume_training to continue. Similar to pause_training but with checkpoint save. ' + 'To fully release GPU resources, use shutdown_server.'), + 'parameters': { + 'type': 'object', + 'properties': { + 'run_id': { + 'type': 'string', + 'description': 'Training run ID to stop.' + }, + }, + 'required': ['run_id'], + }, + }, + }, + { + 'type': 'function', + 'function': { + 'name': + 'update_script', + 'description': + ('Update the training script for a run. Archives the current train.py as train_v{N}.py ' + 'and writes the new version. Use after diagnosing a script error, then call resume_training.'), + 'parameters': { + 'type': 'object', + 'properties': { + 'run_id': { + 'type': 'string', + 'description': 'Training run ID.' + }, + 'script_content': { + 'type': 'string', + 'description': 'Full Python source code of the new training script.' + }, + }, + 'required': ['run_id', 'script_content'], + }, + }, + }, + { + 'type': 'function', + 'function': { + 'name': + 'list_supported_models', + 'description': ('Query the Twinkle server for its list of supported base models. ' + 'Always call this before writing a training script to verify model availability.'), + 'parameters': { + 'type': 'object', + 'properties': { + 'base_url': { + 'type': + 'string', + 'description': + 'Server base URL. Default: http://localhost:8000. Cloud: http://www.modelscope.cn/twinkle', + }, + }, + 'required': [], + }, + }, + }, + { + 'type': 'function', + 'function': { + 'name': 'search_datasets', + 'description': 'Search ModelScope Hub for datasets matching a query.', + 'parameters': { + 'type': 'object', + 'properties': { + 'query': { + 'type': 'string', + 'description': 'Search query for datasets.' + }, + 'limit': { + 'type': 'integer', + 'description': 'Max results (default 5).' + }, + }, + 'required': ['query'], + }, + }, + }, + { + 'type': 'function', + 'function': { + 'name': 'search_models', + 'description': 'Search ModelScope Hub for models matching a query.', + 'parameters': { + 'type': 'object', + 'properties': { + 'query': { + 'type': 'string', + 'description': 'Search query for models.' + }, + 'limit': { + 'type': 'integer', + 'description': 'Max results (default 5).' + }, + }, + 'required': ['query'], + }, + }, + }, + { + 'type': 'function', + 'function': { + 'name': + 'get_cluster_info', + 'description': ('Get cluster GPU resource info for planning training parallelism. ' + 'First attempts to query a running Ray cluster; if Ray is not available, ' + 'falls back to nvidia-smi for local GPU discovery. ' + 'The result indicates whether Ray is active — if not, the training script ' + 'should either start a local Ray cluster itself or the user should launch ' + 'Ray manually (see server mode run.sh).'), + 'parameters': { + 'type': 'object', + 'properties': {}, + 'required': [] + }, + }, + }, +] diff --git a/src/twinkle_client/auto/agent/tools.py b/src/twinkle_client/auto/agent/tools.py index 3689d59eb..40b08f7ca 100644 --- a/src/twinkle_client/auto/agent/tools.py +++ b/src/twinkle_client/auto/agent/tools.py @@ -3,304 +3,20 @@ from __future__ import annotations -import asyncio import json import os from typing import Any, Callable from twinkle.utils.logger import get_logger from twinkle_client.auto.connection import LocalConnection +from .search_tools import _SearchTools +from .server_tools import _ServerTools +from .tool_schemas import TOOL_SCHEMAS logger = get_logger() -# ────────────────────────────────────────────────────────────────────────────── -# Tool schemas (OpenAI function calling format) -# ────────────────────────────────────────────────────────────────────────────── -TOOL_SCHEMAS: list[dict[str, Any]] = [ - { - 'type': 'function', - 'function': { - 'name': 'list_training_runs', - 'description': 'List all active and historical training runs.', - 'parameters': {'type': 'object', 'properties': {}, 'required': []}, - }, - }, - { - 'type': 'function', - 'function': { - 'name': 'get_training_status', - 'description': 'Get detailed status and recent metrics for a training run.', - 'parameters': { - 'type': 'object', - 'properties': { - 'run_id': {'type': 'string', 'description': 'Training run ID.'}, - }, - 'required': ['run_id'], - }, - }, - }, - { - 'type': 'function', - 'function': { - 'name': 'start_server', - 'description': ( - 'Start Ray cluster and Twinkle Server. MUST be called before start_training. ' - 'Idempotent: skips if server is already reachable. ' - 'Supports multi-model deployments: one training model + N sampler/teacher models. ' - 'Automatically generates server_config.yaml from parameters.' - ), - 'parameters': { - 'type': 'object', - 'properties': { - 'model_id': { - 'type': 'string', - 'description': 'Student/training model ID (e.g. "Qwen/Qwen3.5-4B").', - }, - 'train_gpus': { - 'type': 'integer', - 'description': 'GPUs for the training model. Default: auto-detect remaining GPUs.', - }, - 'backend': { - 'type': 'string', - 'enum': ['transformers', 'megatron'], - 'description': 'Training model backend. Default: transformers.', - }, - 'samplers': { - 'type': 'array', - 'description': ( - 'List of sampler/teacher models for RL/OPD. Each entry deploys ' - 'an inference service (vLLM or torch). Omit for simple SFT.' - ), - 'items': { - 'type': 'object', - 'properties': { - 'model_id': { - 'type': 'string', - 'description': 'Teacher/reference model ID (e.g. "Qwen/Qwen3.5-72B").', - }, - 'gpus': { - 'type': 'integer', - 'description': 'Total number of GPUs for this sampler. Default: 1. Must equal tp * dp.', - }, - 'tp': { - 'type': 'integer', - 'description': ( - 'Tensor parallelism size (GPUs per vLLM worker process). ' - 'Use tp>1 for large models that do not fit on a single GPU. Default: 1.' - ), - }, - 'dp': { - 'type': 'integer', - 'description': ( - 'Data parallelism size (number of independent inference replicas). ' - 'If not specified, computed as gpus // tp. Default: 1.' - ), - }, - 'engine': { - 'type': 'string', - 'enum': ['vllm', 'torch'], - 'description': 'Inference engine. Default: vllm.', - }, - 'max_model_len': { - 'type': 'integer', - 'description': 'Max sequence length for inference. Default: 16000.', - }, - }, - 'required': ['model_id'], - }, - }, - 'port': { - 'type': 'integer', - 'description': 'HTTP port for server. Default: 8000.', - }, - }, - 'required': ['model_id'], - }, - }, - }, - { - 'type': 'function', - 'function': { - 'name': 'shutdown_server', - 'description': ( - 'Shut down Twinkle Server and Ray cluster. WARNING: This releases all GPU resources ' - 'and DESTROYS model state held in server memory. Only call when training is truly ' - 'finished and you no longer need the server. Model weights/optimizer state in GPU ' - 'will be LOST unless a checkpoint was explicitly saved.' - ), - 'parameters': {'type': 'object', 'properties': {}, 'required': []}, - }, - }, - { - 'type': 'function', - 'function': { - 'name': 'start_training', - 'description': ( - 'Create a new training run: write the client script, launch it, and start monitoring. ' - 'REQUIRES: Twinkle Server must be running (call start_server first). ' - 'The client script connects to the server — server holds model state in GPU memory. ' - 'Kill client = pause (state preserved). Re-launch client = resume.' - ), - 'parameters': { - 'type': 'object', - 'properties': { - 'run_id': {'type': 'string', 'description': 'Unique run ID (e.g., "grpo-gsm8k").'}, - 'script_content': {'type': 'string', 'description': 'Full Python source code of the training script.'}, - 'model_id': {'type': 'string', 'description': 'Model identifier for metadata (e.g., "Qwen/Qwen3.5-4B").'}, - }, - 'required': ['run_id', 'script_content'], - }, - }, - }, - { - 'type': 'function', - 'function': { - 'name': 'select_run', - 'description': 'Switch to monitor a different training run. Updates connection context.', - 'parameters': { - 'type': 'object', - 'properties': { - 'run_id': {'type': 'string', 'description': 'Training run ID to monitor.'}, - }, - 'required': ['run_id'], - }, - }, - }, - { - 'type': 'function', - 'function': { - 'name': 'pause_training', - 'description': 'Pause training by killing the client process (SIGKILL). Server retains all state — call resume_training to continue.', - 'parameters': { - 'type': 'object', - 'properties': { - 'run_id': {'type': 'string', 'description': 'Training run ID to pause.'}, - }, - 'required': ['run_id'], - }, - }, - }, - { - 'type': 'function', - 'function': { - 'name': 'resume_training', - 'description': 'Resume a paused training run by re-launching the client script. Server state is preserved.', - 'parameters': { - 'type': 'object', - 'properties': { - 'run_id': {'type': 'string', 'description': 'Training run ID to resume.'}, - }, - 'required': ['run_id'], - }, - }, - }, - { - 'type': 'function', - 'function': { - 'name': 'stop_training', - 'description': ( - 'Gracefully stop the training client (SIGTERM). The script saves a checkpoint ' - 'before exiting. Server retains model/optimizer state in GPU memory — ' - 'use resume_training to continue. Similar to pause_training but with checkpoint save. ' - 'To fully release GPU resources, use shutdown_server.' - ), - 'parameters': { - 'type': 'object', - 'properties': { - 'run_id': {'type': 'string', 'description': 'Training run ID to stop.'}, - }, - 'required': ['run_id'], - }, - }, - }, - { - 'type': 'function', - 'function': { - 'name': 'update_script', - 'description': 'Update the training script for a run. Archives the current train.py as train_v{N}.py and writes the new version. Use after diagnosing a script error, then call resume_training.', - 'parameters': { - 'type': 'object', - 'properties': { - 'run_id': {'type': 'string', 'description': 'Training run ID.'}, - 'script_content': {'type': 'string', 'description': 'Full Python source code of the new training script.'}, - }, - 'required': ['run_id', 'script_content'], - }, - }, - }, - { - 'type': 'function', - 'function': { - 'name': 'list_supported_models', - 'description': 'Query the Twinkle server for its list of supported base models. Always call this before writing a training script to verify model availability.', - 'parameters': { - 'type': 'object', - 'properties': { - 'base_url': { - 'type': 'string', - 'description': 'Server base URL. Default: http://localhost:8000. Cloud: http://www.modelscope.cn/twinkle', - }, - }, - 'required': [], - }, - }, - }, - { - 'type': 'function', - 'function': { - 'name': 'search_datasets', - 'description': 'Search ModelScope Hub for datasets matching a query.', - 'parameters': { - 'type': 'object', - 'properties': { - 'query': {'type': 'string', 'description': 'Search query for datasets.'}, - 'limit': {'type': 'integer', 'description': 'Max results (default 5).'}, - }, - 'required': ['query'], - }, - }, - }, - { - 'type': 'function', - 'function': { - 'name': 'search_models', - 'description': 'Search ModelScope Hub for models matching a query.', - 'parameters': { - 'type': 'object', - 'properties': { - 'query': {'type': 'string', 'description': 'Search query for models.'}, - 'limit': {'type': 'integer', 'description': 'Max results (default 5).'}, - }, - 'required': ['query'], - }, - }, - }, - - { - 'type': 'function', - 'function': { - 'name': 'get_cluster_info', - 'description': ( - 'Get cluster GPU resource info for planning training parallelism. ' - 'First attempts to query a running Ray cluster; if Ray is not available, ' - 'falls back to nvidia-smi for local GPU discovery. ' - 'The result indicates whether Ray is active — if not, the training script ' - 'should either start a local Ray cluster itself or the user should launch ' - 'Ray manually (see server mode run.sh).' - ), - 'parameters': {'type': 'object', 'properties': {}, 'required': []}, - }, - }, -] - - -# ────────────────────────────────────────────────────────────────────────────── -# Tool executor -# ────────────────────────────────────────────────────────────────────────────── - - -class ToolExecutor: +class ToolExecutor(_ServerTools, _SearchTools): """Executes agent tool calls against the local connection.""" def __init__(self, connection: LocalConnection): @@ -325,11 +41,7 @@ async def execute(self, name: str, arguments: dict[str, Any]) -> str: def _resolve_server_url(self) -> str: """Resolve server URL: instance state > env var > default.""" - return ( - self._server_url - or os.environ.get('TWINKLE_SERVER_URL') - or 'http://localhost:8000' - ) + return (self._server_url or os.environ.get('TWINKLE_SERVER_URL') or 'http://localhost:8000') async def _tool_list_training_runs(self) -> list[dict]: return self.connection.list_training_runs() @@ -345,12 +57,12 @@ async def _tool_start_training(self, run_id: str, script_content: str, model_id: server_url = self._resolve_server_url() if not await self._check_server_health(server_url): return { - 'status': 'error', - 'run_id': run_id, - 'error': ( - f'Twinkle Server is not reachable at {server_url}. ' - 'Call start_server first to launch Ray cluster and Twinkle Server.' - ), + 'status': + 'error', + 'run_id': + run_id, + 'error': (f'Twinkle Server is not reachable at {server_url}. ' + 'Call start_server first to launch Ray cluster and Twinkle Server.'), } result = self.connection.start_training(run_id, script_content, model_id) actual_run_id = result.get('run_id', run_id) @@ -377,813 +89,3 @@ async def _tool_stop_training(self, run_id: str) -> dict: async def _tool_update_script(self, run_id: str, script_content: str) -> dict: return self.connection.update_script(run_id, script_content) - - # ── Server lifecycle ── - - async def _check_server_health(self, url: str) -> bool: - """Check if Twinkle Server is reachable (non-blocking).""" - import urllib.request - import urllib.error - - def _probe(): - try: - req = urllib.request.Request(f'{url}/api/v1/healthz', method='GET') - urllib.request.urlopen(req, timeout=3) - return True - except (urllib.error.URLError, OSError): - # Try a simpler connectivity check - try: - urllib.request.urlopen(url, timeout=3) - return True - except (urllib.error.URLError, OSError): - return False - - return await asyncio.get_event_loop().run_in_executor(None, _probe) - - # ── Server startup pipeline ── - - async def _tool_start_server( - self, - model_id: str, - train_gpus: int | None = None, - port: int = 8000, - backend: str = 'transformers', - samplers: list[dict] | None = None, - ) -> dict: - """Start Ray cluster + Twinkle Server. Idempotent. Supports multi-model.""" - server_url = self._server_url or os.environ.get('TWINKLE_SERVER_URL') or f'http://localhost:{port}' - - # Idempotent: skip if already running - if await self._check_server_health(server_url): - self._server_url = server_url - return {'status': 'already_running', 'server_url': server_url} - - def _start(): - sampler_list = samplers or [] - - # Step 1: Detect hardware & compute GPU partition - total_hw_gpus = self._detect_gpu_count() - if total_hw_gpus == 0: - return {'status': 'error', 'error': 'No GPUs detected. Cannot start training server.'} - - alloc = self._compute_gpu_allocation(sampler_list, train_gpus, total_hw_gpus) - if 'error' in alloc: - return {'status': 'error', 'error': alloc['error']} - t_gpus, sampler_gpu_total = alloc['train_gpus'], alloc['sampler_gpus'] - - # Step 2: Generate server_config.yaml - config_path = self._generate_server_config( - model_id=model_id, train_gpus=t_gpus, - port=port, backend=backend, samplers=sampler_list, - ) - - # Step 3: Start Ray cluster (multi-node GPU partitioning) - ray_err = self._start_ray_cluster(t_gpus, sampler_gpu_total) - if ray_err: - return {'status': 'error', 'error': ray_err} - - # Step 4: Launch Twinkle Server process - proc, log_path, err = self._launch_server_process(config_path) - if err: - return {'status': 'error', 'error': err} - - # Step 5: Wait for readiness (healthz + sampler engine) - return self._wait_server_ready( - server_url=server_url, proc=proc, log_path=log_path, - sampler_list=sampler_list, model_id=model_id, - t_gpus=t_gpus, backend=backend, config_path=config_path, - ) - - result = await asyncio.get_event_loop().run_in_executor(None, _start) - if result.get('status') in ('started', 'already_running'): - self._server_url = server_url - return result - - # ── Server startup helpers ── - - @staticmethod - def _detect_gpu_count() -> int: - """Detect total hardware GPU count via nvidia-smi.""" - import subprocess as _sp - try: - r = _sp.run( - ['nvidia-smi', '--query-gpu=index', '--format=csv,noheader'], - capture_output=True, text=True, timeout=10, - ) - if r.returncode == 0: - return len([ln for ln in r.stdout.strip().split('\n') if ln.strip()]) - except (FileNotFoundError, OSError): - pass - return 0 - - @staticmethod - def _compute_gpu_allocation( - sampler_list: list[dict], - train_gpus: int | None, - total_hw_gpus: int, - ) -> dict: - """Compute GPU partition: {train_gpus, sampler_gpus} or {error}.""" - sampler_gpu_total = 0 - for s in sampler_list: - s_tp = s.get('tp', 1) - s_dp, s_gpus = s.get('dp'), s.get('gpus') - if s_gpus is not None: - sampler_gpu_total += s_gpus - elif s_dp is not None: - sampler_gpu_total += s_tp * s_dp - else: - sampler_gpu_total += s_tp # default dp=1 - - t_gpus = train_gpus if train_gpus is not None else max(1, total_hw_gpus - sampler_gpu_total) - needed = t_gpus + sampler_gpu_total - if needed > total_hw_gpus: - return { - 'error': ( - f'Requested {needed} GPUs (train={t_gpus}, samplers={sampler_gpu_total}) ' - f'but only {total_hw_gpus} available.' - ), - } - return {'train_gpus': t_gpus, 'sampler_gpus': sampler_gpu_total} - - @staticmethod - def _start_ray_cluster(train_gpus: int, sampler_gpus: int) -> str | None: - """Start Ray multi-node cluster with GPU partitioning. - - Each role gets its own Ray node with dedicated CUDA_VISIBLE_DEVICES - so GPUs are indexed from 0 within each node. This prevents the - GPU ID mapping issues that occur with a single-node setup. - - On a single machine, multiple raylets need separate --temp-dir to - avoid being detected as "already running". - - Returns an error message on failure, or None on success. - """ - import subprocess as _sp - import tempfile - from pathlib import Path - - _sp.run(['ray', 'stop', '--force'], capture_output=True, timeout=15) - - # Create unique temp dirs so each `ray start` spawns a separate raylet - ray_base = Path(tempfile.gettempdir()) / 'twinkle_ray' - ray_base.mkdir(parents=True, exist_ok=True) - - def _ray_node( - devices: str, num_gpus: int, *, - head: bool = False, node_name: str = 'worker', - ) -> str | None: - env = os.environ.copy() - env['CUDA_VISIBLE_DEVICES'] = devices - temp_dir = str(ray_base / node_name) - cmd = ['ray', 'start', f'--temp-dir={temp_dir}'] - if head: - cmd += ['--head', '--port=6379', '--disable-usage-stats', '--include-dashboard=false'] - else: - cmd += ['--address=127.0.0.1:6379'] - cmd.append(f'--num-gpus={num_gpus}') - r = _sp.run(cmd, capture_output=True, text=True, timeout=30, env=env) - if r.returncode != 0 and 'already' not in r.stderr.lower(): - return r.stderr.strip() - return None - - # Head node — training model GPUs - model_devices = ','.join(str(i) for i in range(train_gpus)) - err = _ray_node(model_devices, train_gpus, head=True, node_name='head') - if err: - return f'Ray head start failed: {err}' - - # GPU Worker node — sampler GPUs - if sampler_gpus > 0: - sampler_devices = ','.join(str(i) for i in range(train_gpus, train_gpus + sampler_gpus)) - err = _ray_node(sampler_devices, sampler_gpus, node_name='gpu_worker') - if err: - return f'Ray GPU worker start failed: {err}' - - # CPU Worker node — processor (no GPU) - _ray_node('', 0, node_name='cpu_worker') - return None - - @staticmethod - def _launch_server_process(config_path: str) -> tuple: - """Launch Twinkle Server as a detached background process. - - Returns (proc, log_path, error). On success error is None. - """ - import subprocess as _sp - from pathlib import Path - - log_dir = Path.home() / '.cache' / 'twinkle' - log_dir.mkdir(parents=True, exist_ok=True) - log_path = str(log_dir / 'server.log') - log_file = open(log_path, 'w') - - cmd = ['python', '-m', 'twinkle.server', 'launch', '--config', config_path] - try: - proc = _sp.Popen( - cmd, stdout=log_file, stderr=_sp.STDOUT, - start_new_session=True, - ) - except OSError as e: - log_file.close() - return None, log_path, f'Failed to start Twinkle server: {e}' - return proc, log_path, None - - @staticmethod - def _wait_server_ready( - server_url: str, - proc, - log_path: str, - sampler_list: list[dict], - model_id: str, - t_gpus: int, - backend: str, - config_path: str, - ) -> dict: - """Poll server until healthy (healthz + sampler engine ready).""" - import time - import urllib.request - import urllib.error - - timeout_s = 120 if sampler_list else 60 - needed = t_gpus + sum( - s.get('gpus') or (s.get('tp', 1) * s.get('dp', 1)) for s in sampler_list - ) - - for _ in range(timeout_s): - time.sleep(1) - if proc.poll() is not None: - # Server died — read log tail to diagnose - log_tail = ToolExecutor._read_log_tail(log_path, max_chars=2000) - error_msg = ( - f'Server exited immediately (code={proc.returncode}). ' - f'Model: {model_id}, GPUs: {t_gpus}, Samplers: {len(sampler_list)}.\n' - f'--- server.log tail ---\n{log_tail}' - ) - return { - 'status': 'error', - 'error': error_msg, - 'log_path': log_path, - 'hint': 'Check if required packages are installed (pip install -e ".[all]").', - } - try: - urllib.request.urlopen(f'{server_url}/api/v1/healthz', timeout=2) - except (OSError, Exception): - continue - - # healthz OK — additionally wait for sampler vLLM engines - if sampler_list and not ToolExecutor._probe_sampler_ready(server_url, sampler_list, model_id): - return { - 'status': 'started', - 'warning': 'Server is up but sampler may still be loading.', - 'server_url': server_url, 'server_pid': proc.pid, - 'model_id': model_id, 'log_path': log_path, - } - - return { - 'status': 'started', - 'server_url': server_url, 'server_pid': proc.pid, - 'model_id': model_id, 'train_gpus': t_gpus, - 'backend': backend, - 'samplers': [s.get('model_id') for s in sampler_list], - 'total_gpus_used': needed, - 'config_path': config_path, 'log_path': log_path, - } - - return { - 'status': 'timeout', - 'error': 'Health check did not pass within timeout. Models may still be loading.', - 'server_pid': proc.pid, 'log_path': log_path, - } - - @staticmethod - def _read_log_tail(log_path: str, max_chars: int = 2000) -> str: - """Read the tail of a log file for error diagnosis.""" - try: - with open(log_path, 'r', errors='replace') as f: - content = f.read() - if len(content) <= max_chars: - return content.strip() - return content[-max_chars:].strip() - except OSError: - return '(could not read log file)' - - @staticmethod - def _probe_sampler_ready(server_url: str, sampler_list: list[dict], fallback_model_id: str) -> bool: - """Probe sampler route up to 90s to confirm vLLM engine is loaded.""" - import time - import urllib.request - import urllib.error - - s_mid = sampler_list[0].get('model_id', fallback_model_id) - probe_url = f'{server_url}/api/v1/sampler/{s_mid}/twinkle/create' - - for _ in range(90): - try: - req = urllib.request.Request( - probe_url, method='POST', data=b'{}', - headers={'Content-Type': 'application/json'}, - ) - urllib.request.urlopen(req, timeout=5) - return True # non-error response = ready - except urllib.error.HTTPError as e: - if e.code < 500: - return True # 4xx = actor alive, just bad request - time.sleep(1) # 5xx = still loading - except (OSError, Exception): - time.sleep(1) - return False - - @staticmethod - def _generate_server_config( - model_id: str, - train_gpus: int, - port: int = 8000, - backend: str = 'transformers', - samplers: list[dict] | None = None, - ) -> str: - """Generate a server_config.yaml from template and return its path. - - Supports multi-model topology: - - 1 training model (student) - - N sampler/teacher models (for RL/OPD) - - 1 processor service - """ - from pathlib import Path - import yaml - - sampler_list = samplers or [] - - # Sanitize model name for use in route/names - def _short(mid: str) -> str: - return mid.split('/')[-1] if '/' in mid else mid - - model_short = _short(model_id) - - # Collect all model IDs for supported_models - all_model_ids = [model_id] + [s['model_id'] for s in sampler_list] - - # === Build applications list === - applications = [] - - # 1. API Gateway - applications.append({ - 'name': 'server', - 'route_prefix': '/api/v1', - 'import_path': 'server', - 'args': { - 'server_config': {'per_token_model_limit': 3}, - 'supported_models': all_model_ids, - }, - 'deployments': [{ - 'name': 'TinkerCompatServer', - 'max_ongoing_requests': 50, - 'autoscaling_config': { - 'min_replicas': 1, - 'max_replicas': 1, - 'target_ongoing_requests': 128, - }, - 'ray_actor_options': {'num_cpus': 0.1}, - }], - }) - - # 2. Build GPU-requiring applications (model + samplers), - # then sort by GPU count DESCENDING before appending. - # Largest PG deploys first → it has the fewest node choices → - # avoids GPU scheduling deadlock on single-machine multi-node. - gpu_apps: list[tuple[int, dict]] = [] # (gpu_count, app_config) - - # 2a. Training model worker (student) - gpu_apps.append((train_gpus, { - 'name': f'models-{model_short}', - 'route_prefix': f'/api/v1/model/{model_id}', - 'import_path': 'model', - 'args': { - 'backend': backend, - 'model_id': f'ms://{model_id}', - 'max_length': 500000, # total tokens per forward pass (must match max_input_tokens) - 'nproc_per_node': train_gpus, - 'device_group': { - 'name': 'model', - 'ranks': train_gpus, - 'device_type': 'cuda', - }, - 'device_mesh': { - 'device_type': 'cuda', - 'dp_size': train_gpus, - }, - 'queue_config': { - 'rps_limit': 100, - 'tps_limit': 100000, - 'max_input_tokens': 500000, - }, - 'adapter_config': { - 'adapter_timeout': 600, - }, - }, - 'deployments': [{ - 'name': 'ModelManagement', - 'autoscaling_config': { - 'min_replicas': 1, - 'max_replicas': 1, - 'target_ongoing_requests': 16, - }, - 'ray_actor_options': { - 'num_cpus': 0.1, - 'runtime_env': { - 'env_vars': {'TWINKLE_TRUST_REMOTE_CODE': '1'}, - }, - }, - }], - })) - - # 2b. Sampler/teacher models - sampler_name_count: dict[str, int] = {} - for sampler_cfg in sampler_list: - s_model_id = sampler_cfg['model_id'] - s_short = _short(s_model_id) - - # Deduplicate names when multiple samplers share the same short name - sampler_name_count[s_short] = sampler_name_count.get(s_short, 0) + 1 - if sampler_name_count[s_short] > 1: - s_name = f'sampler-{s_short}-{sampler_name_count[s_short]}' - else: - s_name = f'sampler-{s_short}' - - s_engine = sampler_cfg.get('engine', 'vllm') - s_max_len = sampler_cfg.get('max_model_len', 16000) - - # Compute tp / dp / total GPUs: - # tp = tensor parallelism (GPUs per vLLM process, for large models) - # dp = data parallelism (number of independent inference replicas) - # total GPUs = tp * dp - s_tp = sampler_cfg.get('tp', 1) - s_dp = sampler_cfg.get('dp', None) - s_gpus = sampler_cfg.get('gpus', None) - - if s_dp is not None and s_gpus is not None: - # Both specified: validate consistency - s_tp = s_gpus // s_dp if s_tp == 1 else s_tp - elif s_gpus is not None: - # Only total GPUs specified: derive dp - s_dp = max(1, s_gpus // s_tp) - elif s_dp is not None: - # Only dp specified: derive total - s_gpus = s_tp * s_dp - else: - # Nothing specified: default to 1 GPU (tp=1, dp=1) - s_dp = 1 - s_gpus = s_tp * s_dp - - s_total_gpus = s_tp * s_dp - - # Build device_mesh: include tp_size when tp>1 so that - # world_size = tp*dp and slice_dp dispatch computes correct - # rank_stride for DP data sharding. - mesh_config: dict = {'device_type': 'cuda', 'dp_size': s_dp} - if s_tp > 1: - mesh_config['tp_size'] = s_tp - - sampler_app: dict = { - 'name': s_name, - 'route_prefix': f'/api/v1/sampler/{s_model_id}', - 'import_path': 'sampler', - 'args': { - 'model_id': f'ms://{s_model_id}', - 'nproc_per_node': s_total_gpus, - 'sampler_type': s_engine, - 'device_group': { - 'name': s_name, - 'ranks': s_total_gpus, - 'device_type': 'cuda', - 'gpus_per_worker': s_tp, - }, - 'device_mesh': mesh_config, - 'queue_config': { - 'rps_limit': 100, - 'tps_limit': 100000, - }, - }, - 'deployments': [{ - 'name': 'SamplerManagement', - 'autoscaling_config': { - 'min_replicas': 1, - 'max_replicas': 1, - 'target_ongoing_requests': 16, - }, - 'ray_actor_options': { - 'num_cpus': 0.1, - 'runtime_env': { - 'env_vars': {'TWINKLE_TRUST_REMOTE_CODE': '1'}, - }, - }, - }], - } - - # Add engine-specific args - if s_engine == 'vllm': - engine_args = { - 'max_model_len': s_max_len, - 'gpu_memory_utilization': 0.85, - 'enable_lora': True, - 'logprobs_mode': 'processed_logprobs', - } - # Set tensor_parallel_size when tp > 1 - if s_tp > 1: - engine_args['tensor_parallel_size'] = s_tp - sampler_app['args']['engine_args'] = engine_args - - gpu_apps.append((s_total_gpus, sampler_app)) - - # 3. Sort GPU apps by GPU count DESCENDING, then append in order. - # Largest PG deploys first → claims the largest node → avoids deadlock. - gpu_apps.sort(key=lambda x: x[0], reverse=True) - for _, app_cfg in gpu_apps: - applications.append(app_cfg) - - # 4. Processor service - applications.append({ - 'name': 'processor', - 'route_prefix': '/api/v1/processor', - 'import_path': 'processor', - 'args': { - 'ncpu_proc_per_node': 2, - 'device_group': { - 'name': 'processor', - 'ranks': 2, - 'device_type': 'CPU', - }, - 'device_mesh': { - 'device_type': 'CPU', - 'dp_size': 2, - }, - }, - 'deployments': [{ - 'name': 'ProcessorManagement', - 'autoscaling_config': { - 'min_replicas': 1, - 'max_replicas': 1, - 'target_ongoing_requests': 128, - }, - 'ray_actor_options': {'num_cpus': 0.1}, - }], - }) - - # === Assemble final config === - config = { - 'proxy_location': 'EveryNode', - 'http_options': { - 'host': '0.0.0.0', - 'port': port, - }, - 'applications': applications, - } - - # Write to ~/.cache/twinkle/server_config.yaml - config_dir = Path.home() / '.cache' / 'twinkle' - config_dir.mkdir(parents=True, exist_ok=True) - config_path = config_dir / 'server_config.yaml' - with open(config_path, 'w') as f: - yaml.dump(config, f, default_flow_style=False, allow_unicode=True) - - return str(config_path) - - async def _tool_shutdown_server(self) -> dict: - """Shut down Twinkle Server and Ray cluster. DESTROYS GPU model state.""" - import subprocess as _sp - - def _shutdown(): - results = {} - - # 1. Try `serve shutdown` to cleanly stop Ray Serve deployments - try: - r = _sp.run(['serve', 'shutdown', '-y'], capture_output=True, text=True, timeout=30) - results['serve_shutdown'] = 'ok' if r.returncode == 0 else r.stderr.strip() - except (FileNotFoundError, OSError) as e: - results['serve_shutdown'] = f'skipped: {e}' - - # 2. Kill any remaining twinkle.server processes - try: - _sp.run(['pkill', '-f', 'twinkle.server'], capture_output=True, timeout=5) - except (FileNotFoundError, OSError): - pass - - # 3. Stop Ray cluster - try: - r = _sp.run(['ray', 'stop', '--force'], capture_output=True, text=True, timeout=15) - results['ray_stop'] = 'ok' if r.returncode == 0 else r.stderr.strip() - except (FileNotFoundError, OSError) as e: - results['ray_stop'] = f'failed: {e}' - - results['status'] = 'shutdown_complete' - results['warning'] = 'All GPU model state has been released.' - return results - - return await asyncio.get_event_loop().run_in_executor(None, _shutdown) - - # ── Server queries ── - - async def _tool_list_supported_models(self, base_url: str | None = None) -> dict: - """Query the Twinkle server for supported models.""" - url = base_url or self._resolve_server_url() - - def _query(): - # Use a lightweight HTTP GET instead of init_twinkle_client() which - # creates a session + heartbeat thread that would leak since we never - # call close(). - import urllib.request - import urllib.error - - endpoint = f'{url}/api/v1/twinkle/get_server_capabilities' - req = urllib.request.Request(endpoint, method='GET') - resp = urllib.request.urlopen(req, timeout=10) - data = json.loads(resp.read().decode()) - models = data.get('supported_models', []) - # Each model entry may be a dict with 'model_name' or a plain string - model_names = [] - for m in models: - if isinstance(m, dict): - model_names.append(m.get('model_name', '')) - else: - model_names.append(str(m)) - return { - 'base_url': url, - 'supported_models': model_names, - } - - try: - return await asyncio.get_event_loop().run_in_executor(None, _query) - except Exception as e: - return {'error': f'Failed to query {url}: {e}'} - - async def _tool_search_datasets(self, query: str, limit: int = 5) -> dict: - """Search ModelScope for datasets.""" - return await self._search_hub('datasets', query, limit) - - async def _tool_search_models(self, query: str, limit: int = 5) -> dict: - """Search ModelScope for models.""" - return await self._search_hub('models', query, limit) - - async def _search_hub(self, resource_type: str, query: str, limit: int) -> dict: - """Unified ModelScope Hub search for models or datasets.""" - - def _search(): - if resource_type == 'datasets': - return self._search_datasets_impl(query, limit) - else: - return self._search_models_impl(query, limit) - - try: - items = await asyncio.get_event_loop().run_in_executor(None, _search) - return {'query': query, 'results': items} - except Exception as e: - return {'error': f'{resource_type.title()} search failed: {e}'} - - @staticmethod - def _search_datasets_impl(query: str, limit: int) -> list[dict]: - """Search datasets via ModelScope SDK (new API).""" - from modelscope.hub.api import HubApi - api = HubApi() - result = api.list_datasets('', search=query, page_size=limit) - datasets = result.get('datasets', []) - return [ - {'id': d.get('id', ''), 'name': d.get('display_name', d.get('id', ''))} - for d in datasets - ] - - @staticmethod - def _search_models_impl(query: str, limit: int) -> list[dict]: - """Search models via ModelScope HTTP API (SDK doesn't support search).""" - import requests - resp = requests.put( - 'https://modelscope.cn/api/v1/models/', - json={'Name': query, 'PageSize': limit, 'PageNumber': 1}, - timeout=15, - ) - resp.raise_for_status() - data = resp.json() - if not data.get('Success'): - raise RuntimeError(data.get('Message', 'Unknown error')) - models = data.get('Data', {}).get('Models', []) - return [ - { - 'id': f"{m.get('Path', '')}/{m.get('Name', '')}", - 'name': m.get('ChineseName') or m.get('Name', ''), - } - for m in models - ] - - # ── Cluster info ── - - async def _tool_get_cluster_info(self) -> dict: - """Query cluster resources: try Ray first, fall back to nvidia-smi.""" - - def _query(): - # 1. Try connecting to an existing Ray cluster - ray_info = self._try_ray_cluster() - if ray_info is not None: - ray_info['ray_active'] = True - return ray_info - - # 2. Ray not available — fall back to nvidia-smi - nvidia_info = self._try_nvidia_smi() - nvidia_info['ray_active'] = False - nvidia_info['hint'] = ( - 'Ray cluster is not running. To use distributed training, ' - 'start Ray first: `ray start --head --num-gpus=N` or use ' - 'the server mode run.sh script.' - ) - return nvidia_info - - return await asyncio.get_event_loop().run_in_executor(None, _query) - - @staticmethod - def _try_ray_cluster() -> dict | None: - """Attempt to query an existing Ray cluster. Returns None if unavailable.""" - try: - import ray - except ImportError: - return None - - import logging as _logging - - try: - if not ray.is_initialized(): - ray.init( - address='auto', - ignore_reinit_error=True, - _timeout_s=5, - logging_level=_logging.ERROR, - configure_logging=False, - ) - - resources = ray.cluster_resources() - available = ray.available_resources() - nodes = ray.nodes() - gpu_total = resources.get('GPU', 0) - gpu_available = available.get('GPU', 0) - gpu_types = set() - for node in nodes: - for key in node.get('Resources', {}): - if key.startswith('accelerator_type:'): - gpu_types.add(key.split(':', 1)[1]) - return { - 'num_nodes': len([n for n in nodes if n.get('Alive')]), - 'gpu_total': int(gpu_total), - 'gpu_available': int(gpu_available), - 'gpu_types': sorted(gpu_types) if gpu_types else ['unknown'], - 'cpu_total': resources.get('CPU', 0), - 'memory_bytes': resources.get('memory', 0), - } - except Exception: - try: - import ray as _ray - if _ray.is_initialized(): - _ray.shutdown() - except Exception: - pass - return None - - @staticmethod - def _try_nvidia_smi() -> dict: - """Parse nvidia-smi output for local GPU info.""" - import subprocess as _sp - - try: - result = _sp.run( - ['nvidia-smi', '--query-gpu=index,name,memory.total,memory.free,utilization.gpu', - '--format=csv,noheader,nounits'], - capture_output=True, text=True, timeout=10, - ) - if result.returncode != 0: - return {'error': f'nvidia-smi failed: {result.stderr.strip()}', 'gpu_total': 0} - - gpus = [] - for line in result.stdout.strip().split('\n'): - if not line.strip(): - continue - parts = [p.strip() for p in line.split(',')] - if len(parts) >= 5: - try: - gpus.append({ - 'index': int(parts[0]), - 'name': parts[1], - 'memory_total_mb': int(parts[2]), - 'memory_free_mb': int(parts[3]), - 'utilization_pct': int(parts[4]) if parts[4].isdigit() else 0, - }) - except (ValueError, IndexError): - # Skip lines with unparseable values (e.g. [N/A]) - continue - - gpu_types = sorted(set(g['name'] for g in gpus)) - return { - 'gpu_total': len(gpus), - 'gpu_available': len([g for g in gpus if g['utilization_pct'] < 10]), - 'gpu_types': gpu_types if gpu_types else ['none'], - 'gpus': gpus, - 'source': 'nvidia-smi', - } - except FileNotFoundError: - return {'error': 'nvidia-smi not found (no NVIDIA GPU or driver not installed)', 'gpu_total': 0} - except Exception as e: - return {'error': f'nvidia-smi query failed: {e}', 'gpu_total': 0} diff --git a/src/twinkle_client/auto/app.py b/src/twinkle_client/auto/app.py index cf0b8a067..48d091306 100644 --- a/src/twinkle_client/auto/app.py +++ b/src/twinkle_client/auto/app.py @@ -43,7 +43,7 @@ class TwinkleAuto: def __init__( self, - run_id: Optional[str] = None, + run_id: str | None = None, llm_base_url: str = 'http://localhost:11434/v1', llm_model: str = 'qwen3.5', llm_api_key: str = 'not-needed', @@ -64,12 +64,12 @@ def run(self) -> None: pass async def _main(self) -> None: + from openai import AsyncOpenAI + from twinkle_client.auto.agent.core import AgentLoop from twinkle_client.auto.agent.monitor import TrainingMonitor from twinkle_client.auto.connection import LocalConnection - from openai import AsyncOpenAI - # Connection self._connection = LocalConnection() if self.run_id: @@ -125,9 +125,7 @@ async def _chat_loop(self) -> None: loop = asyncio.get_event_loop() while True: try: - user_input = await loop.run_in_executor( - None, lambda: input(f'{_GREEN}You:{_RESET} ') - ) + user_input = await loop.run_in_executor(None, lambda: input(f'{_GREEN}You:{_RESET} ')) except (KeyboardInterrupt, EOFError): break diff --git a/src/twinkle_client/auto/connection.py b/src/twinkle_client/auto/connection.py index 2734cf186..cb34bc977 100644 --- a/src/twinkle_client/auto/connection.py +++ b/src/twinkle_client/auto/connection.py @@ -20,7 +20,6 @@ from __future__ import annotations import json -from twinkle.utils.logger import get_logger import os import re import shutil @@ -30,6 +29,8 @@ from pathlib import Path from typing import Any +from twinkle.utils.logger import get_logger + logger = get_logger() DEFAULT_BASE_DIR = Path.home() / '.cache' / 'twinkle' @@ -186,7 +187,11 @@ def _launch_script(self, run_id: str) -> dict[str, Any]: error_msg = output_file.read_text().strip()[-500:] if output_file.exists() else '' meta['status'] = 'error' self._write_meta(run_id, meta) - return {'status': 'error', 'run_id': run_id, 'error': error_msg or f'Process exited immediately (code={retcode})'} + return { + 'status': 'error', + 'run_id': run_id, + 'error': error_msg or f'Process exited immediately (code={retcode})' + } meta['pid'] = proc.pid meta['status'] = 'running' diff --git a/src/twinkle_client/auto/runtime.py b/src/twinkle_client/auto/runtime.py index 17cc9d5fa..b3abf1eea 100644 --- a/src/twinkle_client/auto/runtime.py +++ b/src/twinkle_client/auto/runtime.py @@ -35,12 +35,11 @@ import sys import time from pathlib import Path -from typing import Any, TYPE_CHECKING +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: - from twinkle_client.model import MultiLoraTransformersModel from twinkle.dataloader import DataLoader - + from twinkle_client.model import MultiLoraTransformersModel DEFAULT_BASE_DIR = Path.home() / '.cache' / 'twinkle' @@ -68,9 +67,7 @@ def __init__(self, run_id: str | None = None, base_dir: Path | str | None = None if run_id is None: run_id = os.environ.get('TWINKLE_RUN_ID', '') if not run_id: - raise ValueError( - 'run_id must be provided or TWINKLE_RUN_ID env var must be set' - ) + raise ValueError('run_id must be provided or TWINKLE_RUN_ID env var must be set') self.run_id = run_id self.run_dir = self.base_dir / run_id @@ -247,8 +244,8 @@ def finish(self, status: str = 'completed') -> None: def register_graceful_shutdown( self, - model: 'MultiLoraTransformersModel', - dataloader: 'DataLoader | None' = None, + model: MultiLoraTransformersModel, + dataloader: DataLoader | None = None, checkpoint_name: str = 'interrupted', ) -> None: """Register SIGTERM handler for graceful shutdown with checkpoint. @@ -270,6 +267,7 @@ def register_graceful_shutdown( rt.register_graceful_shutdown(model, dataloader) # ... training loop ... """ + def _shutdown_handler(signum, frame): self.log('SIGTERM received, saving checkpoint before exit...') try: diff --git a/src/twinkle_client/common/__init__.py b/src/twinkle_client/common/__init__.py new file mode 100644 index 000000000..ec8001aaa --- /dev/null +++ b/src/twinkle_client/common/__init__.py @@ -0,0 +1,7 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Client-internal helpers shared across the twinkle_client subpackages. + +Regular package (carries this ``__init__``) so ``component_rpc`` is included by +``setuptools.packages.find`` in a built wheel; a namespace-only directory would +be dropped from the distribution. +""" diff --git a/src/twinkle_client/common/component_rpc.py b/src/twinkle_client/common/component_rpc.py new file mode 100644 index 000000000..65b877455 --- /dev/null +++ b/src/twinkle_client/common/component_rpc.py @@ -0,0 +1,42 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Shared transport-bound plumbing for processor-backed component clients.""" +from __future__ import annotations + +from typing import Any + +from twinkle.protocol.types.processor import (ProcessorCallRequest, ProcessorCallResponse, ProcessorCreateRequest, + ProcessorCreateResponse) +from twinkle_client._request_builder import build_request +from twinkle_client.http import ClientTransport +from twinkle_client.http.client import DEFAULT_TIMEOUT +from twinkle_client.http.context import capture_transport + + +def create_remote_component( + processor_type: str, + class_type: str, + *, + transport: ClientTransport | None = None, + **init_kwargs: Any, +) -> str: + """Create a server-side component using one captured transport.""" + resolved = capture_transport(transport) + body = build_request(ProcessorCreateRequest, processor_type=processor_type, class_type=class_type, **init_kwargs) + response = resolved.post_model(resolved.url('processor/twinkle/create'), body) + return ProcessorCreateResponse(**response.json()).processor_id + + +def call_remote_component( + processor_id: str, + function: str, + http_timeout: Any = DEFAULT_TIMEOUT, + /, + *, + transport: ClientTransport | None = None, + **call_kwargs: Any, +) -> Any: + """Invoke one server-side component using its owner's transport.""" + resolved = capture_transport(transport) + body = build_request(ProcessorCallRequest, processor_id=processor_id, function=function, **call_kwargs) + response = resolved.post_model(resolved.url('processor/twinkle/call'), body, timeout=http_timeout) + return ProcessorCallResponse(**response.json()).result diff --git a/src/twinkle_client/common/remote_component.py b/src/twinkle_client/common/remote_component.py new file mode 100644 index 000000000..a5017f78b --- /dev/null +++ b/src/twinkle_client/common/remote_component.py @@ -0,0 +1,45 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Shared transport binding for processor-backed remote component wrappers.""" +from __future__ import annotations + +from typing import Any + +from twinkle_client.common.component_rpc import call_remote_component, create_remote_component +from twinkle_client.http import ClientTransport +from twinkle_client.http.context import capture_transport + + +class RemoteComponent: + """Bind one remote component id to one captured transport. + + The class intentionally has no ``__init__``: concrete wrappers retain + ownership of their domain-specific constructor and MRO. + """ + + _transport: ClientTransport + processor_id: str + + def _bind_remote( + self, + processor_type: str, + class_type: str, + *, + transport: ClientTransport | None = None, + **kwargs: Any, + ) -> None: + self._transport = capture_transport(transport) + self.processor_id = create_remote_component( + processor_type, + class_type, + transport=self._transport, + **kwargs, + ) + + def _call(self, function: str, *args: Any, **kwargs: Any) -> Any: + return call_remote_component( + self.processor_id, + function, + *args, + transport=self._transport, + **kwargs, + ) diff --git a/src/twinkle_client/data_plane.py b/src/twinkle_client/data_plane.py index 281df8315..eec097165 100644 --- a/src/twinkle_client/data_plane.py +++ b/src/twinkle_client/data_plane.py @@ -6,10 +6,10 @@ from collections.abc import Callable from typing import Any, TypeVar -from twinkle_client.common.json_utils import json_safe -from twinkle_client.http import get_base_url, http_post -from twinkle_client.types.component import DataRef, DataRowsResponse - +from twinkle.protocol.json_utils import json_safe +from twinkle.protocol.types.component import DataRef, DataRowsResponse +from twinkle_client.http import ClientTransport +from twinkle_client.http.context import capture_transport _T = TypeVar('_T') @@ -21,8 +21,9 @@ async def _call_in_thread(func: Callable[..., _T], /, *args: Any, **kwargs: Any) class DataPlaneClient: - def __init__(self, server_url: str | None = None): - self.server_url = (server_url or f'{get_base_url()}/data-plane').rstrip('/') + def __init__(self, server_url: str | None = None, *, transport: ClientTransport | None = None): + self._transport = capture_transport(transport) + self.server_url = (server_url or self._transport.url('data-plane')).rstrip('/') def put( self, @@ -31,9 +32,13 @@ def put( kind: str = 'data', tags: list[dict[str, Any]] | None = None, ) -> DataRef: - response = http_post( + response = self._transport.post( f'{self.server_url}/twinkle/put', - json_data={'rows': json_safe(rows), 'kind': kind, 'tags': json_safe(tags)}, + json_data={ + 'rows': json_safe(rows), + 'kind': kind, + 'tags': json_safe(tags) + }, ) response.raise_for_status() return DataRef(**response.json()) @@ -51,9 +56,12 @@ async def aput( return await _call_in_thread(self.put, rows, kind=kind, tags=tags) def get(self, ref: DataRef, *, fields: list[str] | None = None) -> list[dict[str, Any]]: - response = http_post( + response = self._transport.post( f'{self.server_url}/twinkle/get', - json_data={'ref': ref.model_dump(), 'fields': fields}, + json_data={ + 'ref': ref.model_dump(), + 'fields': fields + }, ) response.raise_for_status() return DataRowsResponse(**response.json()).rows @@ -64,9 +72,13 @@ def get_batch( *, fields: list[str] | None = None, ) -> DataRowsResponse: - response = http_post( + response = self._transport.post( f'{self.server_url}/twinkle/get', - json_data={'ref': ref.model_dump(), 'fields': fields, 'include_tags': True}, + json_data={ + 'ref': ref.model_dump(), + 'fields': fields, + 'include_tags': True + }, ) response.raise_for_status() return DataRowsResponse(**response.json()) @@ -94,7 +106,7 @@ def append( *, tags: list[dict[str, Any]] | None = None, ) -> DataRef: - response = http_post( + response = self._transport.post( f'{self.server_url}/twinkle/append', json_data={ 'ref': ref.model_dump(), @@ -118,7 +130,7 @@ async def aappend( return await _call_in_thread(self.append, ref, rows, tags=tags) def release(self, ref: DataRef) -> None: - response = http_post( + response = self._transport.post( f'{self.server_url}/twinkle/release', json_data={'ref': ref.model_dump()}, ) diff --git a/src/twinkle_client/dataloader/__init__.py b/src/twinkle_client/dataloader/__init__.py index b94ed5b8d..3db1e11c0 100644 --- a/src/twinkle_client/dataloader/__init__.py +++ b/src/twinkle_client/dataloader/__init__.py @@ -1 +1,3 @@ from .dataloader import DataLoader + +__all__ = ['DataLoader'] diff --git a/src/twinkle_client/dataloader/dataloader.py b/src/twinkle_client/dataloader/dataloader.py index 86e2bbf42..21cd5cdbb 100644 --- a/src/twinkle_client/dataloader/dataloader.py +++ b/src/twinkle_client/dataloader/dataloader.py @@ -1,115 +1,50 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +from __future__ import annotations -from typing import Callable, Type, Union -from twinkle_client.http import http_post -from twinkle.dataset import Dataset -from twinkle.processor import InputProcessor +from typing import TYPE_CHECKING, Callable, Type, Union -class DataLoader(object): - """Client wrapper for DataLoader that calls server HTTP endpoints.""" +from twinkle_client.common.remote_component import RemoteComponent +from twinkle_client.http import ClientTransport + +if TYPE_CHECKING: + from twinkle.processor import InputProcessor + from twinkle_client.dataset import Dataset - def __init__(self, dataset: Union[Dataset, Callable], **kwargs): - from twinkle_client.http import get_base_url - self.server_url = f'{get_base_url()}/processor/twinkle' - response = http_post( - url=f'{self.server_url}/create', - json_data={ - 'processor_type': 'dataloader', - 'class_type': 'DataLoader', - **{'dataset': dataset}, **kwargs - } - ) - response.raise_for_status() - self.processor_id = response.json()['processor_id'] +class DataLoader(RemoteComponent): + """Client wrapper for DataLoader that calls server HTTP endpoints.""" + + def __init__( + self, + dataset: Dataset | Callable, + *, + transport: ClientTransport | None = None, + **kwargs, + ): + dataset_transport = getattr(dataset, '_transport', None) + if transport is not None and dataset_transport is not None and transport is not dataset_transport: + raise ValueError('DataLoader and its remote Dataset must use the same ClientTransport') + self._bind_remote( + 'dataloader', 'DataLoader', dataset=dataset, transport=transport or dataset_transport, **kwargs) - def __len__(self): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__len__', - **{}, - } - ) - response.raise_for_status() - return response.json()["result"] - + return self._call('__len__') - def set_processor(self, processor_cls: Union[Type[InputProcessor], str, InputProcessor, Callable], **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'set_processor', - **{'processor_cls': processor_cls}, - **kwargs - } - ) - response.raise_for_status() - return response.json()["result"] - + def set_processor(self, processor_cls: type[InputProcessor] | str | InputProcessor | Callable, **kwargs): + return self._call('set_processor', processor_cls=processor_cls, **kwargs) def __iter__(self): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__iter__', - **{}, - } - ) - response.raise_for_status() + self._call('__iter__') return self - + def __next__(self): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__next__', - } - ) - response.raise_for_status() - return response.json()["result"] - + return self._call('__next__') def skip_consumed_samples(self, consumed_train_samples: int): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'skip_consumed_samples', - **{'consumed_train_samples': consumed_train_samples}, - } - ) - response.raise_for_status() - return response.json()["result"] - + return self._call('skip_consumed_samples', consumed_train_samples=consumed_train_samples) def resume_from_checkpoint(self, consumed_train_samples, **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'resume_from_checkpoint', - **{'consumed_train_samples': consumed_train_samples}, - **kwargs - } - ) - response.raise_for_status() - return response.json()["result"] - + return self._call('resume_from_checkpoint', consumed_train_samples=consumed_train_samples, **kwargs) def get_state(self): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'get_state', - **{}, - } - ) - response.raise_for_status() - return response.json()["result"] - \ No newline at end of file + return self._call('get_state') diff --git a/src/twinkle_client/dataset/__init__.py b/src/twinkle_client/dataset/__init__.py index ba37b1fe5..23f7ed855 100644 --- a/src/twinkle_client/dataset/__init__.py +++ b/src/twinkle_client/dataset/__init__.py @@ -3,3 +3,5 @@ from .iterable_packing_dataset import IterablePackingDataset from .lazy_dataset import LazyDataset from .packing_dataset import PackingDataset + +__all__ = ['Dataset', 'IterableDataset', 'IterablePackingDataset', 'LazyDataset', 'PackingDataset'] diff --git a/src/twinkle_client/dataset/base.py b/src/twinkle_client/dataset/base.py index bec2d4309..c7a7b8a92 100644 --- a/src/twinkle_client/dataset/base.py +++ b/src/twinkle_client/dataset/base.py @@ -1,191 +1,71 @@ - +# Copyright (c) ModelScope Contributors. All rights reserved. from typing import Any, Callable, Dict, Optional, Type, Union -from twinkle_client.http import http_post -from twinkle.dataset import Dataset + from twinkle.dataset import DatasetMeta -from twinkle.preprocessor import DataFilter -from twinkle.preprocessor import Preprocessor +from twinkle.preprocessor import DataFilter, Preprocessor from twinkle.template import Template +from twinkle_client.common.remote_component import RemoteComponent +from twinkle_client.http import ClientTransport + -class Dataset(object): +class Dataset(RemoteComponent): """Client wrapper for Dataset that calls server HTTP endpoints.""" - def __init__(self, dataset_meta: DatasetMeta = None, **kwargs): - from twinkle_client.http import get_base_url - - self.server_url = f'{get_base_url()}/processor/twinkle' - response = http_post( - url=f'{self.server_url}/create', - json_data={ - 'processor_type': 'dataset', - 'class_type': 'Dataset', - **{'dataset_meta': dataset_meta}, **kwargs - } - ) - response.raise_for_status() - self.processor_id = response.json()['processor_id'] - - + def __init__( + self, + dataset_meta: DatasetMeta = None, + *, + transport: ClientTransport | None = None, + **kwargs, + ): + self._bind_remote('dataset', 'Dataset', dataset_meta=dataset_meta, transport=transport, **kwargs) + def set_template(self, template_func: Union[Template, Type[Template], str], **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'set_template', - **{'template_func': template_func}, - **kwargs - } - ) - response.raise_for_status() - return response.json()["result"] - + return self._call('set_template', template_func=template_func, **kwargs) def encode(self, add_generation_prompt: bool = False, timeout: Optional[int] = 600, **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'encode', - **{'add_generation_prompt': add_generation_prompt}, - **kwargs - }, - timeout=timeout - ) - response.raise_for_status() - return response.json()["result"] - + return self._call('encode', timeout, add_generation_prompt=add_generation_prompt, **kwargs) def check(self, **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'check', - **{}, - **kwargs - } - ) - response.raise_for_status() - return response.json()["result"] - + return self._call('check', **kwargs) def cast_column(self, column: str, decode: bool = True): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'cast_column', - **{'column': column, 'decode': decode}, - } - ) - response.raise_for_status() - return response.json()["result"] - - - def map(self, preprocess_func: Union[Preprocessor, Callable, str, Type[Preprocessor]], dataset_meta: DatasetMeta = None, init_args: Dict[str, Any] = None, **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'map', - **{'preprocess_func': preprocess_func, 'dataset_meta': dataset_meta, 'init_args': init_args}, - **kwargs - } - ) - response.raise_for_status() - return response.json()["result"] - - - def filter(self, filter_func: Union[Callable, str, Type[DataFilter], DataFilter], dataset_meta: DatasetMeta = None, init_args: Dict[str, Any] = None, **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'filter', - **{'filter_func': filter_func, 'dataset_meta': dataset_meta, 'init_args': init_args}, - **kwargs - } - ) - response.raise_for_status() - return response.json()["result"] - + return self._call('cast_column', column=column, decode=decode) + + def map(self, + preprocess_func: Union[Preprocessor, Callable, str, Type[Preprocessor]], + dataset_meta: DatasetMeta = None, + init_args: Dict[str, Any] = None, + **kwargs): + return self._call( + 'map', preprocess_func=preprocess_func, dataset_meta=dataset_meta, init_args=init_args, **kwargs) + + def filter(self, + filter_func: Union[Callable, str, Type[DataFilter], DataFilter], + dataset_meta: DatasetMeta = None, + init_args: Dict[str, Any] = None, + **kwargs): + return self._call('filter', filter_func=filter_func, dataset_meta=dataset_meta, init_args=init_args, **kwargs) def add_dataset(self, dataset_meta: DatasetMeta, **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'add_dataset', - **{'dataset_meta': dataset_meta}, - **kwargs - } - ) - response.raise_for_status() - return response.json()["result"] - - - def mix_dataset(self, interleave = True): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'mix_dataset', - **{'interleave': interleave}, - } - ) - response.raise_for_status() - return response.json()["result"] - - - def save_as(self, output_path: str, format: Optional[str] = None, batch_size: int = 1000, mode: str = 'immediate', **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'save_as', - **{'output_path': output_path, 'format': format, 'batch_size': batch_size, 'mode': mode}, - **kwargs - } - ) - response.raise_for_status() - return response.json()["result"] - + return self._call('add_dataset', dataset_meta=dataset_meta, **kwargs) + + def mix_dataset(self, interleave=True): + return self._call('mix_dataset', interleave=interleave) + + def save_as(self, + output_path: str, + format: Optional[str] = None, + batch_size: int = 1000, + mode: str = 'immediate', + **kwargs): + return self._call('save_as', output_path=output_path, format=format, batch_size=batch_size, mode=mode, **kwargs) def flush_save(self): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'flush_save', - **{}, - } - ) - response.raise_for_status() - return response.json()["result"] - + return self._call('flush_save') def __getitem__(self, idx): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__getitem__', - **{'idx': idx}, - } - ) - response.raise_for_status() - return response.json()["result"] - + return self._call('__getitem__', idx=idx) def __len__(self): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__len__', - **{}, - } - ) - response.raise_for_status() - return response.json()["result"] - \ No newline at end of file + return self._call('__len__') diff --git a/src/twinkle_client/dataset/iterable_dataset.py b/src/twinkle_client/dataset/iterable_dataset.py index 0646ea33b..d3b5eeced 100644 --- a/src/twinkle_client/dataset/iterable_dataset.py +++ b/src/twinkle_client/dataset/iterable_dataset.py @@ -1,88 +1,33 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +from torch.utils.data import IterableDataset as TorchIterableDataset -from twinkle_client.http import http_post -from twinkle.dataset import Dataset from twinkle.dataset import DatasetMeta -from torch.utils.data import IterableDataset +from twinkle_client.common.remote_component import RemoteComponent +from twinkle_client.http import ClientTransport -class IterableDataset(IterableDataset): - """Client wrapper for IterableDataset that calls server HTTP endpoints.""" - def __init__(self, dataset_meta: DatasetMeta = None, **kwargs): - from twinkle_client.http import get_base_url +class IterableDataset(TorchIterableDataset, RemoteComponent): + """Remote iterable backed by one server-side cursor. - self.server_url = f'{get_base_url()}/processor/twinkle' - response = http_post( - url=f'{self.server_url}/create', - json_data={ - 'processor_type': 'dataset', - 'class_type': 'IterableDataset', - **{'dataset_meta': dataset_meta}, **kwargs - } - ) - response.raise_for_status() - self.processor_id = response.json()['processor_id'] + Iteration is stateful and does not support concurrent or repeated iteration + over the same wrapper instance. + """ - - def add_dataset(self, dataset_meta: DatasetMeta, **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'add_dataset', - **{'dataset_meta': dataset_meta}, - **kwargs - } - ) - response.raise_for_status() - return response.json()["result"] - - - def __len__(self): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__len__', - **{}, - } - ) - response.raise_for_status() - return response.json()["result"] - + def __init__( + self, + dataset_meta: DatasetMeta = None, + *, + transport: ClientTransport | None = None, + **kwargs, + ): + self._bind_remote('dataset', 'IterableDataset', dataset_meta=dataset_meta, transport=transport, **kwargs) - def __getitem__(self, idx): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__getitem__', - **{'idx': idx}, - } - ) - response.raise_for_status() - return response.json()["result"] - + def add_dataset(self, dataset_meta: DatasetMeta, **kwargs): + return self._call('add_dataset', dataset_meta=dataset_meta, **kwargs) def __iter__(self): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__iter__', - **{}, - } - ) - response.raise_for_status() + self._call('__iter__') return self - + def __next__(self): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__next__', - } - ) - response.raise_for_status() - return response.json()["result"] - \ No newline at end of file + return self._call('__next__') diff --git a/src/twinkle_client/dataset/iterable_packing_dataset.py b/src/twinkle_client/dataset/iterable_packing_dataset.py index f4b6b5d1e..5fc91174f 100644 --- a/src/twinkle_client/dataset/iterable_packing_dataset.py +++ b/src/twinkle_client/dataset/iterable_packing_dataset.py @@ -1,77 +1,46 @@ - +# Copyright (c) ModelScope Contributors. All rights reserved. +from torch.utils.data import IterableDataset as TorchIterableDataset from typing import Type, Union -from twinkle_client.http import http_post -from twinkle.dataset import Dataset + from twinkle.dataset import DatasetMeta from twinkle.template import Template -from torch.utils.data import IterableDataset +from twinkle_client.common.remote_component import RemoteComponent +from twinkle_client.http import ClientTransport -class IterablePackingDataset(IterableDataset): - """Client wrapper for IterablePackingDataset that calls server HTTP endpoints.""" - def __init__(self, dataset_meta: DatasetMeta = None, packing_interval: int = 128, packing_num_proc: int = 1, cyclic: bool = False, **kwargs): - from twinkle_client.http import get_base_url +class IterablePackingDataset(TorchIterableDataset, RemoteComponent): + """Remote packing iterable backed by one non-reentrant server cursor.""" - self.server_url = f'{get_base_url()}/processor/twinkle' - response = http_post( - url=f'{self.server_url}/create', - json_data={ - 'processor_type': 'dataset', - 'class_type': 'IterablePackingDataset', - **{'dataset_meta': dataset_meta, 'packing_interval': packing_interval, 'packing_num_proc': packing_num_proc, 'cyclic': cyclic}, **kwargs - } + def __init__( + self, + dataset_meta: DatasetMeta = None, + packing_interval: int = 128, + packing_num_proc: int = 1, + cyclic: bool = False, + *, + transport: ClientTransport | None = None, + **kwargs, + ): + self._bind_remote( + 'dataset', + 'IterablePackingDataset', + dataset_meta=dataset_meta, + packing_interval=packing_interval, + packing_num_proc=packing_num_proc, + cyclic=cyclic, + transport=transport, + **kwargs, ) - response.raise_for_status() - self.processor_id = response.json()['processor_id'] - def set_template(self, template_cls: Union[Type[Template], str, Template], **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'set_template', - **{'template_cls': template_cls}, - **kwargs - } - ) - response.raise_for_status() - return response.json()["result"] - + return self._call('set_template', template_cls=template_cls, **kwargs) def pack_dataset(self): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'pack_dataset', - **{}, - } - ) - response.raise_for_status() - return response.json()["result"] - + return self._call('pack_dataset') def __iter__(self): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__iter__', - **{}, - } - ) - response.raise_for_status() + self._call('__iter__') return self - + def __next__(self): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__next__', - } - ) - response.raise_for_status() - return response.json()["result"] - \ No newline at end of file + return self._call('__next__') diff --git a/src/twinkle_client/dataset/lazy_dataset.py b/src/twinkle_client/dataset/lazy_dataset.py index 7bf49c70e..866fd64ec 100644 --- a/src/twinkle_client/dataset/lazy_dataset.py +++ b/src/twinkle_client/dataset/lazy_dataset.py @@ -1,137 +1,17 @@ - -from typing import Any, Callable, Dict, Optional, Type, Union -from twinkle_client.http import http_post -from twinkle.dataset import Dataset +# Copyright (c) ModelScope Contributors. All rights reserved. from twinkle.dataset import DatasetMeta -from twinkle.preprocessor import DataFilter -from twinkle.preprocessor import Preprocessor +from twinkle_client.http import ClientTransport from .base import Dataset + class LazyDataset(Dataset): """Client wrapper for LazyDataset that calls server HTTP endpoints.""" - def __init__(self, dataset_meta: DatasetMeta = None, **kwargs): - from twinkle_client.http import get_base_url - - self.server_url = f'{get_base_url()}/processor/twinkle' - response = http_post( - url=f'{self.server_url}/create', - json_data={ - 'processor_type': 'dataset', - 'class_type': 'LazyDataset', - **{'dataset_meta': dataset_meta}, **kwargs - } - ) - response.raise_for_status() - self.processor_id = response.json()['processor_id'] - - - def map(self, preprocess_func: Union[Preprocessor, Callable, str, Type[Preprocessor]], dataset_meta: DatasetMeta = None, init_args: Dict[str, Any] = None, **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'map', - **{'preprocess_func': preprocess_func, 'dataset_meta': dataset_meta, 'init_args': init_args}, - **kwargs - } - ) - response.raise_for_status() - return response.json()["result"] - - - def filter(self, filter_func: Union[Callable, str, Type[DataFilter], DataFilter], dataset_meta: DatasetMeta = None, init_args: Dict[str, Any] = None, **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'filter', - **{'filter_func': filter_func, 'dataset_meta': dataset_meta, 'init_args': init_args}, - **kwargs - } - ) - response.raise_for_status() - return response.json()["result"] - - - def add_dataset(self, dataset_meta: DatasetMeta, **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'add_dataset', - **{'dataset_meta': dataset_meta}, - **kwargs - } - ) - response.raise_for_status() - return response.json()["result"] - - - def mix_dataset(self, interleave = True): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'mix_dataset', - **{'interleave': interleave}, - } - ) - response.raise_for_status() - return response.json()["result"] - - - def encode(self, add_generation_prompt: bool = False, timeout: Optional[int] = 600, **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'encode', - **{'add_generation_prompt': add_generation_prompt}, - **kwargs - }, - timeout=timeout - ) - response.raise_for_status() - return response.json()["result"] - - - def check(self, **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'check', - **{}, - **kwargs - } - ) - response.raise_for_status() - return response.json()["result"] - - - def __getitem__(self, idx): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__getitem__', - **{'idx': idx}, - } - ) - response.raise_for_status() - return response.json()["result"] - - - def __len__(self): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__len__', - **{}, - } - ) - response.raise_for_status() - return response.json()["result"] - \ No newline at end of file + def __init__( + self, + dataset_meta: DatasetMeta = None, + *, + transport: ClientTransport | None = None, + **kwargs, + ): + self._bind_remote('dataset', 'LazyDataset', dataset_meta=dataset_meta, transport=transport, **kwargs) diff --git a/src/twinkle_client/dataset/packing_dataset.py b/src/twinkle_client/dataset/packing_dataset.py index 855777674..bf5c014e2 100644 --- a/src/twinkle_client/dataset/packing_dataset.py +++ b/src/twinkle_client/dataset/packing_dataset.py @@ -1,63 +1,34 @@ - -from twinkle_client.http import http_post -from twinkle.dataset import Dataset +# Copyright (c) ModelScope Contributors. All rights reserved. from twinkle.dataset import DatasetMeta +from twinkle_client.http import ClientTransport from .base import Dataset + class PackingDataset(Dataset): """Client wrapper for PackingDataset that calls server HTTP endpoints.""" - def __init__(self, dataset_meta: DatasetMeta = None, packing_num_proc: int = 1, **kwargs): - from twinkle_client.http import get_base_url - - self.server_url = f'{get_base_url()}/processor/twinkle' - response = http_post( - url=f'{self.server_url}/create', - json_data={ - 'processor_type': 'dataset', - 'class_type': 'PackingDataset', - **{'dataset_meta': dataset_meta, 'packing_num_proc': packing_num_proc}, **kwargs - } + def __init__( + self, + dataset_meta: DatasetMeta = None, + packing_num_proc: int = 1, + *, + transport: ClientTransport | None = None, + **kwargs, + ): + self._bind_remote( + 'dataset', + 'PackingDataset', + dataset_meta=dataset_meta, + packing_num_proc=packing_num_proc, + transport=transport, + **kwargs, ) - response.raise_for_status() - self.processor_id = response.json()['processor_id'] - def pack_dataset(self): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': 'pack_dataset', - **{}, - } - ) - response.raise_for_status() - return response.json()["result"] - + return self._call('pack_dataset') def __getitem__(self, index): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__getitem__', - **{'index': index}, - } - ) - response.raise_for_status() - return response.json()["result"] - + return self._call('__getitem__', index=index) def __len__(self): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__len__', - **{}, - } - ) - response.raise_for_status() - return response.json()["result"] - \ No newline at end of file + return self._call('__len__') diff --git a/src/twinkle_client/exceptions.py b/src/twinkle_client/exceptions.py new file mode 100644 index 000000000..dd56d6e35 --- /dev/null +++ b/src/twinkle_client/exceptions.py @@ -0,0 +1,137 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Twinkle_Client exceptions for the request lifecycle. + +Two axes, kept deliberately distinct: + +- Transport / HTTP failures (:class:`TwinkleHTTPError`) inherit ``requests.HTTPError`` + so existing ``except requests.HTTPError`` clauses keep working. They carry the + server's ``error_code`` / ``category`` when the response body had them. +- Task-outcome and polling failures (:class:`TaskFailedError`, + :class:`TaskWaitTimeoutError`, :class:`TaskRecordLostError`) do NOT inherit + ``requests.HTTPError``: a task that reaches a ``failed`` terminal state is + delivered over HTTP 200, so it is not an HTTP-level error. +- Request-construction failures (:class:`TwinkleClientValidationError`) happen before + any HTTP call is made. +""" +from __future__ import annotations + +import requests +from typing import Any, Optional + +from twinkle.protocol.types.errors import ErrorCategory + + +class TwinkleClientValidationError(ValueError): + """A caller argument could not be placed in the request model, in-process. + + Distinct from ``pydantic.ValidationError``, which reports a *field* that failed + validation. This one is raised **before** the model is constructed, when an + argument has no field to go into at all -- so it cannot be expressed as a field + error. Either way no HTTP request is sent. + + A ``ValueError`` subclass so callers already catching ``ValueError`` around request + construction keep working. + """ + + +class TwinkleHTTPError(requests.HTTPError): + """An HTTP 4xx/5xx (other than 410) from a twinkle endpoint. + + Inherits ``requests.HTTPError`` so callers already catching that keep working. + ``status_code`` is the HTTP status; ``error_code`` / ``category`` come from the + server's structured error body when present. ``category`` always uses the + lowercase :class:`ErrorCategory` wire value. + """ + + def __init__( + self, + *args: Any, + status_code: int | None = None, + error_code: int | None = None, + category: str = ErrorCategory.Unknown.value, + request_id: str | None = None, + details: list[dict[str, Any]] | None = None, + traceback: str | None = None, + **kwargs: Any, + ) -> None: + super().__init__(*args, **kwargs) + self.status_code = status_code + self.error_code = error_code + self.category = category + self.request_id = request_id + self.details = details + self.traceback = traceback + + +class TaskFailedError(Exception): + """A task reached the ``failed`` terminal state (delivered over HTTP 200). + + The twinkle counterpart of tinker's ``RequestFailedError`` + (``tinker/_exceptions.py``; carries ``message`` / ``request_id`` / ``category``). + Deliberately NOT a ``requests.HTTPError`` subclass: the HTTP call succeeded, it + is the *task* that failed, so this is not an HTTP-level error. + """ + + def __init__( + self, + error: str, + *, + category: str, + request_id: str, + error_code: int | None = None, + details: list[dict[str, Any]] | None = None, + ) -> None: + super().__init__(error) + self.error = error + self.category = category + self.request_id = request_id + self.error_code = error_code + self.details = details + + +class TaskCancelledError(Exception): + """A task reached the ``cancelled`` terminal state (delivered over HTTP 200). + + Distinct from :class:`TaskFailedError`: the task did not fail, it was cancelled + before it started running (client cancel). Not a ``requests.HTTPError`` -- the + HTTP call succeeded; the task was simply dropped. + """ + + def __init__( + self, + error: str, + *, + request_id: str, + error_code: int | None = None, + ) -> None: + super().__init__(error) + self.error = error + self.request_id = request_id + self.error_code = error_code + + +class TaskWaitTimeoutError(Exception): + """The Client_Future_Layer stopped polling after ``total_timeout`` seconds. + + "I am not waiting any longer" -- distinct from :class:`TaskRecordLostError`, + which indicates a state-layer problem. The task itself is guaranteed to reach a + terminal state by the server; this only means the client gave up. + """ + + def __init__(self, *, request_id: str, waited: float) -> None: + super().__init__(f'Timed out after {waited:.1f}s waiting for task {request_id}') + self.request_id = request_id + self.waited = waited + + +class TaskRecordLostError(Exception): + """Retrieve_Endpoint returned 404 for a whole run of consecutive attempts. + + Distinct from :class:`TaskWaitTimeoutError`: a 404 run points at the state layer + (Redis blip, actor restart) rather than a slow task, so the operator response + differs. + """ + + def __init__(self, *, request_id: str) -> None: + super().__init__(f'Task record for {request_id} was not found after repeated retries') + self.request_id = request_id diff --git a/src/twinkle_client/http/__init__.py b/src/twinkle_client/http/__init__.py index e36ce1e27..700e46acb 100644 --- a/src/twinkle_client/http/__init__.py +++ b/src/twinkle_client/http/__init__.py @@ -1,19 +1,8 @@ -from .http_utils import http_delete, http_get, http_post -from .utils import (TWINKLE_SERVER_TOKEN, TWINKLE_SERVER_URL, get_api_key, get_base_url, get_request_id, - get_session_id, set_api_key, set_base_url, set_request_id, set_session_id) +"""Public HTTP transport API.""" +from .client import ClientTransport +from .context import ClientContext __all__ = [ - 'http_get', - 'http_post', - 'http_delete', - 'TWINKLE_SERVER_URL', - 'TWINKLE_SERVER_TOKEN', - 'set_base_url', - 'get_base_url', - 'set_api_key', - 'get_api_key', - 'set_session_id', - 'get_session_id', - 'set_request_id', - 'get_request_id', + 'ClientContext', + 'ClientTransport', ] diff --git a/src/twinkle_client/http/client.py b/src/twinkle_client/http/client.py new file mode 100644 index 000000000..a209bc09b --- /dev/null +++ b/src/twinkle_client/http/client.py @@ -0,0 +1,242 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Instance-owned HTTP transport.""" +from __future__ import annotations + +import requests +from collections.abc import Mapping +from requests.adapters import HTTPAdapter +from typing import Any +from urllib3.util.retry import Retry + +from twinkle.protocol.headers import build_routing_headers +from twinkle.protocol.types.errors import ErrorCategory, ErrorPayload +from twinkle_client._request_builder import to_wire_value +from twinkle_client.exceptions import TwinkleClientValidationError, TwinkleHTTPError +from .context import ClientContext, capture_transport + +# Must be greater than the server long-poll window and below common gateway idle limits. +_HTTP_TIMEOUT = 90 +DEFAULT_TIMEOUT = object() + + +def _handle_response(response: requests.Response) -> requests.Response: + if response.status_code == 410: + try: + detail = response.json().get('detail', 'Iterator exhausted') + except Exception: + detail = response.text or 'Iterator exhausted' + raise StopIteration(detail) + + if response.ok: + return response + + try: + body = response.json() + except Exception: + body = None + payload: ErrorPayload | None = None + if isinstance(body, dict): + try: + payload = ErrorPayload.model_validate(body) + except Exception: + payload = None + + if payload is not None: + summary = payload.error or response.text + category = payload.category.value + error_code = payload.error_code + request_id = payload.request_id + details = payload.details + traceback_text = payload.traceback + elif isinstance(body, dict): + summary = body.get('detail') or response.text + category = ErrorCategory.Unknown.value + error_code = request_id = details = traceback_text = None + else: + summary = response.text + category = ErrorCategory.Unknown.value + error_code = request_id = details = traceback_text = None + + message = f'{response.status_code} Error for url: {response.url}\nServer detail:\n{summary}' + raise TwinkleHTTPError( + message, + response=response, + status_code=response.status_code, + error_code=error_code, + category=category, + request_id=request_id, + details=details, + traceback=traceback_text, + ) + + +class ClientTransport: + """The sole request-time owner of URL, identity, headers, and HTTP resources. + + ``post`` recursively converts arbitrary JSON-like values through + :func:`to_wire_value`; ``post_model`` serializes a validated Pydantic model + directly and excludes ``None`` fields. The two entry points intentionally + remain distinct. + + The adapter retries only idempotent GET/DELETE requests. POST is never + transparently replayed because control-plane calls such as ``create_session`` + do not carry a deduplication key; read-only future retrieval handles retries + explicitly in the future layer. + """ + + def __init__( + self, + context: ClientContext, + *, + session: requests.Session | None = None, + timeout: float = _HTTP_TIMEOUT, + pool_maxsize: int = 32, + ) -> None: + if pool_maxsize < 2: + raise ValueError('pool_maxsize must be at least 2') + self._context = context + self._session = session or requests.Session() + if session is None: + retry = Retry( + total=3, + connect=3, + read=0, + status=3, + backoff_factor=0.25, + status_forcelist=(408, 429, 500, 502, 503, 504), + allowed_methods=frozenset({'GET', 'DELETE'}), + respect_retry_after_header=True, + raise_on_status=False, + ) + adapter = HTTPAdapter( + max_retries=retry, + pool_connections=pool_maxsize, + pool_maxsize=pool_maxsize, + pool_block=True, + ) + self._session.mount('http://', adapter) + self._session.mount('https://', adapter) + self._timeout = timeout + self._closed = False + self._published = False + self._capabilities: object | None = None + + @property + def context(self) -> ClientContext: + return self._context + + @property + def closed(self) -> bool: + return self._closed + + def bind_context(self, context: ClientContext) -> None: + """Replace provisional identity before the transport is published to wrappers.""" + self._ensure_open() + if self._published: + raise RuntimeError('Cannot rebind a published ClientTransport') + self._context = context + + def _mark_published(self) -> None: + self._published = True + + @property + def cached_capabilities(self) -> object | None: + return self._capabilities + + @cached_capabilities.setter + def cached_capabilities(self, value: object) -> None: + self._capabilities = value + + def url(self, path_or_url: str = '') -> str: + if path_or_url.startswith(('http://', 'https://')): + return path_or_url + if not path_or_url: + return self._context.base_url + return f'{self._context.base_url}/{path_or_url.lstrip("/")}' + + def _headers(self, additional_headers: Mapping[str, str] | None = None) -> dict[str, str]: + headers = build_routing_headers(self._context.routing_id, f'Bearer {self._context.api_key}') + if self._context.session_id: + headers['X-Twinkle-Session-Id'] = self._context.session_id + if additional_headers: + headers.update(additional_headers) + return headers + + def _request_timeout(self, timeout: object) -> float | None: + return self._timeout if timeout is DEFAULT_TIMEOUT else timeout # type: ignore[return-value] + + def _ensure_open(self) -> None: + if self._closed: + raise RuntimeError('ClientTransport is closed') + + def get(self, + path_or_url: str = '', + *, + params: Mapping[str, Any] | None = None, + headers: Mapping[str, str] | None = None, + timeout: float | None | object = DEFAULT_TIMEOUT) -> requests.Response: + self._ensure_open() + response = self._session.get( + self.url(path_or_url), + headers=self._headers(headers), + params=to_wire_value(params or {}), + timeout=self._request_timeout(timeout), + ) + return _handle_response(response) + + def post(self, + path_or_url: str = '', + *, + json_data: Mapping[str, Any] | None = None, + data: Any = None, + headers: Mapping[str, str] | None = None, + timeout: float | None | object = DEFAULT_TIMEOUT) -> requests.Response: + self._ensure_open() + if isinstance(data, (bytes, bytearray, memoryview)): + raise TwinkleClientValidationError('Binary request bodies are not supported by this transport') + response = self._session.post( + self.url(path_or_url), + headers=self._headers(headers), + json=to_wire_value(json_data or {}), + data=data, + timeout=self._request_timeout(timeout), + ) + return _handle_response(response) + + def post_model(self, + path_or_url: str, + body: Any, + *, + headers: Mapping[str, str] | None = None, + timeout: float | None | object = DEFAULT_TIMEOUT) -> requests.Response: + from twinkle_client._request_builder import request_json + self._ensure_open() + request_headers = {'content-type': 'application/json', **dict(headers or {})} + response = self._session.post( + self.url(path_or_url), + headers=self._headers(request_headers), + data=request_json(body), + timeout=self._request_timeout(timeout), + ) + return _handle_response(response) + + def delete(self, + path_or_url: str = '', + *, + params: Mapping[str, Any] | None = None, + headers: Mapping[str, str] | None = None, + timeout: float | None | object = DEFAULT_TIMEOUT) -> requests.Response: + self._ensure_open() + response = self._session.delete( + self.url(path_or_url), + headers=self._headers(headers), + params=to_wire_value(params or {}), + timeout=self._request_timeout(timeout), + ) + return _handle_response(response) + + def close(self) -> None: + if self._closed: + return + self._closed = True + self._session.close() diff --git a/src/twinkle_client/http/context.py b/src/twinkle_client/http/context.py new file mode 100644 index 000000000..58bb8d29a --- /dev/null +++ b/src/twinkle_client/http/context.py @@ -0,0 +1,99 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Immutable client identity and the compatibility default-transport registry.""" +from __future__ import annotations + +import logging +import os +import threading +import uuid +from dataclasses import dataclass +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from .client import ClientTransport + +TWINKLE_SERVER_URL = os.environ.get('TWINKLE_SERVER_URL', 'http://127.0.0.1:8000') +TWINKLE_SERVER_TOKEN = os.environ.get('TWINKLE_SERVER_TOKEN', 'EMPTY_TOKEN') + +logger = logging.getLogger('twinkle_client') + + +def _normalize_base_url(base_url: str) -> str: + base_url = base_url.rstrip('/') + return base_url if base_url.endswith('/api/v1') else f'{base_url}/api/v1' + + +@dataclass(frozen=True, slots=True) +class ClientContext: + """A resolved request identity captured by one transport.""" + + base_url: str + api_key: str + session_id: str | None = None + routing_id: str = '' + + def __post_init__(self) -> None: + object.__setattr__(self, 'base_url', _normalize_base_url(self.base_url)) + if not self.routing_id: + object.__setattr__(self, 'routing_id', uuid.uuid4().hex) + + +_default_lock = threading.RLock() +_default_transport: ClientTransport | None = None + + +def _new_env_transport() -> ClientTransport: + from .client import ClientTransport + return ClientTransport( + ClientContext( + base_url=os.environ.get('TWINKLE_SERVER_URL', TWINKLE_SERVER_URL), + api_key=os.environ.get('TWINKLE_SERVER_TOKEN', TWINKLE_SERVER_TOKEN), + )) + + +def capture_transport(explicit: ClientTransport | None = None) -> ClientTransport: + """Return an explicit transport or capture the current compatibility default.""" + global _default_transport + if explicit is not None: + if explicit.closed: + raise RuntimeError('Cannot capture a closed ClientTransport') + explicit._mark_published() + return explicit + with _default_lock: + if _default_transport is None or _default_transport.closed: + _default_transport = _new_env_transport() + logger.info('No explicit Twinkle client configured; using %s', _default_transport.context.base_url) + _default_transport._mark_published() + return _default_transport + + +def set_default_transport(transport: ClientTransport) -> None: + if transport.closed: + raise RuntimeError('Cannot register a closed ClientTransport') + global _default_transport + with _default_lock: + if _default_transport is not None and _default_transport is not transport and not _default_transport.closed: + logger.warning( + 'Replacing default Twinkle transport %s with %s; existing wrappers retain the old transport', + _default_transport.context.base_url, + transport.context.base_url, + ) + transport._mark_published() + _default_transport = transport + + +def clear_default_transport(transport: ClientTransport) -> None: + global _default_transport + with _default_lock: + if _default_transport is transport: + _default_transport = None + + +# Private compatibility seam for the deferred Tinker monkey patch. New Twinkle code +# must use ClientTransport directly; these names are intentionally not re-exported. +def get_api_key() -> str: + return capture_transport().context.api_key + + +def get_request_id() -> str: + return capture_transport().context.routing_id diff --git a/src/twinkle_client/http/http_utils.py b/src/twinkle_client/http/http_utils.py deleted file mode 100644 index 6aac84a1a..000000000 --- a/src/twinkle_client/http/http_utils.py +++ /dev/null @@ -1,185 +0,0 @@ -import requests -from typing import Any, Callable, Dict, Mapping, Optional - -from .headers import build_routing_headers -from .utils import get_api_key, get_base_url, get_request_id, get_session_id - - -def _build_headers(additional_headers: Optional[Dict[str, str]] = None) -> Dict[str, str]: - """ - Build HTTP headers with request ID and authorization. - - Args: - additional_headers: Additional headers to include - - Returns: - Dictionary of headers - """ - headers = build_routing_headers(get_request_id(), 'Bearer ' + get_api_key()) - - if session_id := get_session_id(): - headers['X-Twinkle-Session-Id'] = session_id - - if additional_headers: - headers.update(additional_headers) - - return headers - - -def _serialize_params(params: Dict[str, Any]) -> Dict[str, Any]: - """ - Serialize parameters, handling special objects like processors. - - Args: - params: Parameters to serialize - - Returns: - Serialized parameters dictionary - """ - serialized = {} - for key, value in params.items(): - if hasattr(value, 'processor_id'): - serialized[key] = value.processor_id - elif hasattr(value, '__dict__'): - from twinkle_client.common.serialize import serialize_object - serialized[key] = serialize_object(value) - else: - serialized[key] = value - return serialized - - -def _handle_response(response: requests.Response) -> requests.Response: - """ - Handle common response processing. - - Args: - response: Response object - - Returns: - Response object - - Raises: - StopIteration: When server returns HTTP 410 (iterator exhausted) - requests.HTTPError: When server returns a 4xx/5xx error, with the - server-side ``detail`` field (full traceback) included in the - exception message so callers don't need to inspect the response body. - """ - # Convert HTTP 410 Gone to StopIteration - # This indicates an iterator has been exhausted - if response.status_code == 410: - raise StopIteration(response.json().get('detail', 'Iterator exhausted')) - - if not response.ok: - try: - detail = response.json().get('detail', response.text) - except Exception: - detail = response.text - http_error_msg = ( - f'{response.status_code} Error for url: {response.url}\n' - f'Server detail:\n{detail}' - ) - raise requests.HTTPError(http_error_msg, response=response) - - return response - - -def http_get( - url: Optional[str] = None, - params: Optional[Dict[str, Any]] = {}, - additional_headers: Optional[Dict[str, str]] = {}, - timeout: int = 600, -) -> requests.Response: - """ - Send HTTP GET request with required headers. - - Args: - url: The target URL - params: Query parameters - additional_headers: Additional headers to include - timeout: Request timeout in seconds - - Returns: - requests.Response object - """ - url = url or get_base_url() - headers = _build_headers(additional_headers) - serialized_params = _serialize_params(params) - - response = requests.get( - url, - headers=headers, - params=serialized_params, - timeout=timeout, - ) - - return _handle_response(response) - - -def http_post( - url: Optional[str] = None, - json_data: Optional[Dict[str, Any]] = {}, - data: Optional[Any] = {}, - additional_headers: Optional[Dict[str, str]] = {}, - timeout: Optional[int] = 600, -) -> requests.Response: - """ - Send HTTP POST request with required headers. - - Args: - url: The target URL - json_data: JSON data to send in request body - data: Form data or raw data to send in request body - additional_headers: Additional headers to include - timeout: Request timeout in seconds; None disables the timeout. - - Returns: - requests.Response object - - Raises: - StopIteration: When server returns HTTP 410 (iterator exhausted) - """ - url = url or get_base_url() - headers = _build_headers(additional_headers) - serialized_json = _serialize_params(json_data) - - response = requests.post( - url, - headers=headers, - json=serialized_json, - data=data, - timeout=timeout, - ) - - return _handle_response(response) - - -def http_delete( - url: Optional[str] = None, - params: Optional[Dict[str, Any]] = {}, - additional_headers: Optional[Dict[str, str]] = {}, - timeout: int = 600, -) -> requests.Response: - """ - Send HTTP DELETE request with required headers. - - Args: - url: The target URL - params: Query parameters - additional_headers: Additional headers to include - timeout: Request timeout in seconds - - Returns: - requests.Response object - """ - url = url or get_base_url() - headers = _build_headers(additional_headers) - serialized_params = _serialize_params(params) - - response = requests.delete( - url, - headers=headers, - params=serialized_params, - timeout=timeout, - ) - - return _handle_response(response) diff --git a/src/twinkle_client/http/utils.py b/src/twinkle_client/http/utils.py deleted file mode 100644 index f5b348358..000000000 --- a/src/twinkle_client/http/utils.py +++ /dev/null @@ -1,64 +0,0 @@ -import os -import uuid -from datetime import datetime -from typing import Optional - -TWINKLE_SERVER_URL = os.environ.get('TWINKLE_SERVER_URL', 'http://127.0.0.1:8000') -TWINKLE_SERVER_TOKEN = os.environ.get('TWINKLE_SERVER_TOKEN', 'EMPTY_TOKEN') - -# Global variables for configuration -_base_url: Optional[str] = None -_api_key: Optional[str] = None -_session_id: Optional[str] = None -_request_id: Optional[str] = None - - -def set_base_url(url: str): - """Set the base URL for HTTP requests.""" - global _base_url - _base_url = url.rstrip('/') - - -def get_base_url() -> str: - """Get the current base URL.""" - base_url = _base_url or TWINKLE_SERVER_URL - if not base_url.endswith('/api/v1'): - base_url += '/api/v1' - return base_url - - -def set_api_key(api_key: str): - """Set the API key for HTTP requests.""" - global _api_key - _api_key = api_key - - -def get_api_key() -> str: - """Get the current API key.""" - return _api_key or TWINKLE_SERVER_TOKEN - - -def set_session_id(session_id: str): - """Set the session ID.""" - global _session_id - _session_id = session_id - - -def get_session_id() -> Optional[str]: - """Get the current session ID.""" - return _session_id - - -def set_request_id(request_id: str): - """Set the global request ID for HTTP requests (shared across all threads).""" - global _request_id - _request_id = request_id - - -def get_request_id() -> str: - """Get the global request ID or generate and cache a new one.""" - global _request_id - if _request_id is not None: - return _request_id - _request_id = datetime.now().strftime('%Y%m%d_%H%M%S') + '-' + str(uuid.uuid4().hex)[0:8] - return _request_id diff --git a/src/twinkle_client/manager.py b/src/twinkle_client/manager.py index 12257d63c..e87ff2292 100644 --- a/src/twinkle_client/manager.py +++ b/src/twinkle_client/manager.py @@ -2,93 +2,123 @@ from __future__ import annotations import atexit +import os import threading -from typing import Any, Dict, List, Optional, Tuple +from dataclasses import replace +from typing import Any + from twinkle import get_logger -from twinkle_client.types.server import (CapacityInfoResponse, DeleteCheckpointResponse, GetServerCapabilitiesResponse) -from twinkle_client.types.session import (CreateSessionRequest, CreateSessionResponse, SessionHeartbeatRequest, - SessionHeartbeatResponse) -from twinkle_client.types.training import (Checkpoint, Cursor, ParsedCheckpointTwinklePath, TrainingRun, - TrainingRunsResponse, WeightsInfoResponse) -from .http import get_api_key, get_base_url, http_delete, http_get, http_post, set_api_key, set_base_url, set_session_id +from twinkle.protocol.types.server import CapacityInfoResponse, DeleteCheckpointResponse, GetServerCapabilitiesResponse +from twinkle.protocol.types.session import CreateSessionRequest, CreateSessionResponse, SessionHeartbeatRequest +from twinkle.protocol.types.training import (Checkpoint, Cursor, ParsedCheckpointTwinklePath, TrainingRun, + WeightsInfoResponse) +from twinkle_client.exceptions import TwinkleHTTPError +from twinkle_client.http import ClientContext, ClientTransport +from twinkle_client.http.context import (TWINKLE_SERVER_TOKEN, TWINKLE_SERVER_URL, clear_default_transport, + set_default_transport) logger = get_logger() -class TwinkleClientError(Exception): - """Base exception for TwinkleManager errors.""" - pass - class TwinkleClient: - """ - Client manager for interacting with Twinkle REST API. - - On initialization this client: - - Sets the base_url and api_key into the shared context so that all other - client objects (MultiLoraTransformersModel, vLLMSampler, processor clients) - automatically pick up the same configuration. - - Creates a server-side session and stores the session_id in context so that - every outgoing HTTP request carries it in the ``X-Twinkle-Session-Id`` header. - - Starts a lightweight background thread that touches the session every - ``session_heartbeat_interval`` seconds to keep it alive. - - Args: - base_url: Base URL of the Twinkle server (e.g. "http://localhost:8000"). - Falls back to the ``TWINKLE_SERVER_URL`` environment variable. - api_key: API key for authentication. Falls back to the - ``TWINKLE_SERVER_TOKEN`` environment variable. - route_prefix: API route prefix (default: "/twinkle"). - session_heartbeat_interval: Seconds between session touch calls (default: 30). - session_metadata: Optional metadata dict stored with the session on the server. + """Owner of one connected transport, remote session, and heartbeat thread. + + Use :meth:`connect` (normally through ``init_twinkle_client``) for remote I/O. + ``__init__`` only accepts already-established state, which keeps partial + connection failures from publishing a client or leaking a heartbeat thread. """ def __init__( self, - base_url: Optional[str] = None, - api_key: Optional[str] = None, - route_prefix: Optional[str] = '/twinkle', + *, + transport: ClientTransport, + heartbeat_transport: ClientTransport | None = None, + route_prefix: str = '/twinkle', session_heartbeat_interval: int = 10, - session_metadata: Optional[Dict[str, Any]] = None, - ): - # Resolve and store config, then propagate to context so all generated - # client objects that call get_base_url() / get_api_key() get these values. - if base_url: - set_base_url(base_url) - if api_key: - set_api_key(api_key) - - self.base_url = get_base_url() - self.api_key = get_api_key() + ) -> None: + """Build an already-connected client without performing remote I/O.""" + self._transport = transport + self._heartbeat_transport = heartbeat_transport or transport + self.base_url = transport.context.base_url + self.api_key = transport.context.api_key self.route_prefix = route_prefix.rstrip('/') if route_prefix else '' - - # Create a server-side session. - self._session_id: str = self.create_session(session_metadata) - set_session_id(self._session_id) - - # Start background session-touch thread. + self._session_id = transport.context.session_id self._heartbeat_interval = session_heartbeat_interval self._stop_event = threading.Event() + self._heartbeat_thread: threading.Thread | None = None + self._close_lock = threading.Lock() + self._closed = False + + @classmethod + def connect( + cls, + base_url: str | None = None, + api_key: str | None = None, + route_prefix: str | None = '/twinkle', + session_heartbeat_interval: int = 10, + session_metadata: dict[str, Any] | None = None, + ) -> TwinkleClient: + """Create the remote session and atomically publish a connected client.""" + context = ClientContext( + base_url=base_url or os.environ.get('TWINKLE_SERVER_URL', TWINKLE_SERVER_URL), + api_key=api_key or os.environ.get('TWINKLE_SERVER_TOKEN', TWINKLE_SERVER_TOKEN), + ) + transport = ClientTransport(context) + prefix = route_prefix.rstrip('/') if route_prefix else '' + client = None + heartbeat_transport = None + try: + response = transport.post( + f'{context.base_url}{prefix}/create_session', + json_data=CreateSessionRequest(metadata=session_metadata).model_dump(), + ) + session_id = CreateSessionResponse.model_validate(response.json()).session_id + transport.bind_context(replace(context, session_id=session_id)) + heartbeat_transport = ClientTransport(transport.context) + client = cls( + transport=transport, + heartbeat_transport=heartbeat_transport, + route_prefix=prefix, + session_heartbeat_interval=session_heartbeat_interval, + ) + set_default_transport(transport) + client._start_heartbeat() + atexit.register(client.close) + return client + except BaseException: + if client is None: + if heartbeat_transport is not None: + heartbeat_transport.close() + transport.close() + else: + client.close() + raise + + @property + def transport(self) -> ClientTransport: + return self._transport + + def _start_heartbeat(self) -> None: self._heartbeat_thread = threading.Thread( target=self._touch_session_loop, daemon=True, name='TwinkleSessionHeartbeat', ) self._heartbeat_thread.start() - atexit.register(self.close) def get_capacity_info(self) -> CapacityInfoResponse: """ Get the server's global LoRA capacity information. Returns: - :class:`~twinkle_client.types.server.CapacityInfoResponse` with + :class:`~twinkle.protocol.types.server.CapacityInfoResponse` with ``max_loras``, ``used_loras``, and ``free_loras`` fields. Raises: - TwinkleClientError: If the request fails. + TwinkleHTTPError: If the request fails. """ - response = http_get(self._get_url('/capacity_info')) - data = self._handle_response(response) + response = self._transport.get(self._get_url('/capacity_info')) + data = response.json() return CapacityInfoResponse(**data) # ------------------------------------------------------------------ @@ -99,18 +129,7 @@ def _get_url(self, endpoint: str) -> str: """Construct full URL for an endpoint.""" return f'{self.base_url}{self.route_prefix}{endpoint}' - def _handle_response(self, response, expected_code: int = 200) -> dict[str, Any]: - """Handle HTTP response and raise appropriate errors.""" - if response.status_code != expected_code: - try: - error_data = response.json() - detail = error_data.get('detail', str(error_data)) - except Exception: - detail = response.text - raise TwinkleClientError(f'Request failed with status {response.status_code}: {detail}') - return response.json() - - def create_session(self, metadata: Optional[Dict[str, Any]] = None) -> str: + def create_session(self, metadata: dict[str, Any] | None = None) -> str: """ Create a server-side session. @@ -121,13 +140,12 @@ def create_session(self, metadata: Optional[Dict[str, Any]] = None) -> str: The session ID string. Raises: - TwinkleClientError: If the session creation request fails. + TwinkleHTTPError: If the session creation request fails. """ - resp = http_post( + resp = self._transport.post( self._get_url('/create_session'), json_data=CreateSessionRequest(metadata=metadata).model_dump(), ) - resp.raise_for_status() return CreateSessionResponse(**resp.json()).session_id def _touch_session_loop(self) -> None: @@ -144,12 +162,11 @@ def _touch_session_loop(self) -> None: success = False try: logger.debug(f'[TwinkleClient] Touching session (session={self._session_id})...') - resp = http_post( + self._heartbeat_transport.post( self._get_url('/session_heartbeat'), json_data=SessionHeartbeatRequest(session_id=self._session_id).model_dump(), timeout=min(self._heartbeat_interval, 10), ) - resp.raise_for_status() success = True except Exception as e: logger.error(f'[TwinkleClient] Session heartbeat error: {e}') @@ -160,10 +177,38 @@ def _touch_session_loop(self) -> None: self._stop_event.wait(timeout=sleep_time) def close(self) -> None: - """Stop the background heartbeat thread and clear session context.""" + """Stop owned resources exactly once without affecting another client.""" + with self._close_lock: + if self._closed: + return + self._closed = True self._stop_event.set() - if self._heartbeat_thread.is_alive(): - self._heartbeat_thread.join(timeout=2) + if self._heartbeat_thread is not None and self._heartbeat_thread.is_alive(): + self._heartbeat_thread.join(timeout=max(2, min(self._heartbeat_interval, 10))) + clear_default_transport(self._transport) + if self._heartbeat_transport is not self._transport: + self._heartbeat_transport.close() + self._transport.close() + try: + atexit.unregister(self.close) + except Exception: + pass + + def __enter__(self) -> TwinkleClient: + return self + + def __exit__(self, exc_type, exc_value, traceback) -> None: + self.close() + + def model(self, model_id: str, **kwargs: Any): + """Create a remote training model bound explicitly to this client.""" + from twinkle_client.model import MultiLoraTransformersModel + return MultiLoraTransformersModel(model_id, transport=self._transport, **kwargs) + + def sampler(self, model_id: str, **kwargs: Any): + """Create a remote sampler bound explicitly to this client.""" + from twinkle_client.sampler import vLLMSampler + return vLLMSampler(model_id, transport=self._transport, **kwargs) # ------------------------------------------------------------------ # Health Check @@ -177,7 +222,7 @@ def health_check(self) -> bool: True if server is healthy, False otherwise. """ try: - response = http_get(self._get_url('/healthz')) + response = self._transport.get(self._get_url('/healthz')) return response.status_code == 200 except Exception: return False @@ -187,21 +232,25 @@ def get_server_capabilities(self) -> GetServerCapabilitiesResponse: Get the server's supported models and capabilities. Returns: - :class:`~twinkle_client.types.server.GetServerCapabilitiesResponse` with + :class:`~twinkle.protocol.types.server.GetServerCapabilitiesResponse` with ``supported_models`` field containing a list of supported model names. Raises: - TwinkleClientError: If the request fails. + TwinkleHTTPError: If the request fails. """ - response = http_get(self._get_url('/get_server_capabilities')) - data = self._handle_response(response) - return GetServerCapabilitiesResponse(**data) + cached = self._transport.cached_capabilities + if isinstance(cached, GetServerCapabilitiesResponse): + return cached + response = self._transport.get(self._get_url('/get_server_capabilities')) + capabilities = GetServerCapabilitiesResponse.model_validate(response.json()) + self._transport.cached_capabilities = capabilities + return capabilities # ------------------------------------------------------------------ # Training Runs # ------------------------------------------------------------------ - def list_training_runs(self, limit: int = 20, offset: int = 0, all_users: bool = False) -> List[TrainingRun]: + def list_training_runs(self, limit: int = 20, offset: int = 0, all_users: bool = False) -> list[TrainingRun]: """ List training runs. @@ -213,17 +262,17 @@ def list_training_runs(self, limit: int = 20, offset: int = 0, all_users: bool = all_users: If True, return all runs (if permission allows). Returns: - List of :class:`~twinkle_client.types.training.TrainingRun` objects. + List of :class:`~twinkle.protocol.types.training.TrainingRun` objects. Raises: - TwinkleClientError: If the request fails. + TwinkleHTTPError: If the request fails. """ - params: Dict[str, Any] = {'limit': limit, 'offset': offset} + params: dict[str, Any] = {'limit': limit, 'offset': offset} if all_users: params['all_users'] = 'true' - response = http_get(self._get_url('/training_runs'), params=params) - data = self._handle_response(response) + response = self._transport.get(self._get_url('/training_runs'), params=params) + data = response.json() return [TrainingRun(**r) for r in data.get('training_runs', [])] @@ -232,7 +281,7 @@ def list_training_runs_with_cursor( limit: int = 20, offset: int = 0, all_users: bool = False, - ) -> Tuple[List[TrainingRun], Cursor]: + ) -> tuple[list[TrainingRun], Cursor]: """ List training runs with pagination info. @@ -245,14 +294,14 @@ def list_training_runs_with_cursor( Tuple of (list of TrainingRun, Cursor with pagination info). Raises: - TwinkleClientError: If the request fails. + TwinkleHTTPError: If the request fails. """ - params: Dict[str, Any] = {'limit': limit, 'offset': offset} + params: dict[str, Any] = {'limit': limit, 'offset': offset} if all_users: params['all_users'] = 'true' - response = http_get(self._get_url('/training_runs'), params=params) - data = self._handle_response(response) + response = self._transport.get(self._get_url('/training_runs'), params=params) + data = response.json() runs = [TrainingRun(**r) for r in data.get('training_runs', [])] cursor = Cursor(**data.get('cursor', {})) @@ -266,20 +315,20 @@ def get_training_run(self, run_id: str) -> TrainingRun: run_id: The training run identifier. Returns: - :class:`~twinkle_client.types.training.TrainingRun` object with run details. + :class:`~twinkle.protocol.types.training.TrainingRun` object with run details. Raises: - TwinkleClientError: If run not found or access denied. + TwinkleHTTPError: If run not found or access denied. """ - response = http_get(self._get_url(f'/training_runs/{run_id}')) - data = self._handle_response(response) + response = self._transport.get(self._get_url(f'/training_runs/{run_id}')) + data = response.json() return TrainingRun(**data) # ------------------------------------------------------------------ # Checkpoints # ------------------------------------------------------------------ - def list_checkpoints(self, run_id: str) -> List[Checkpoint]: + def list_checkpoints(self, run_id: str) -> list[Checkpoint]: """ List checkpoints for a training run. @@ -287,13 +336,13 @@ def list_checkpoints(self, run_id: str) -> List[Checkpoint]: run_id: The training run identifier. Returns: - List of :class:`~twinkle_client.types.training.Checkpoint` objects. + List of :class:`~twinkle.protocol.types.training.Checkpoint` objects. Raises: - TwinkleClientError: If run not found or access denied. + TwinkleHTTPError: If run not found or access denied. """ - response = http_get(self._get_url(f'/training_runs/{run_id}/checkpoints')) - data = self._handle_response(response) + response = self._transport.get(self._get_url(f'/training_runs/{run_id}/checkpoints')) + data = response.json() return [Checkpoint(**c) for c in data.get('checkpoints', [])] def get_checkpoint_path(self, run_id: str, checkpoint_id: str) -> ParsedCheckpointTwinklePath: @@ -305,14 +354,14 @@ def get_checkpoint_path(self, run_id: str, checkpoint_id: str) -> ParsedCheckpoi checkpoint_id: The checkpoint identifier (e.g. "weights/20240101_120000"). Returns: - :class:`~twinkle_client.types.training.ParsedCheckpointTwinklePath` with + :class:`~twinkle.protocol.types.training.ParsedCheckpointTwinklePath` with ``path`` (filesystem) and ``twinkle_path`` fields. Raises: - TwinkleClientError: If checkpoint not found or access denied. + TwinkleHTTPError: If checkpoint not found or access denied. """ - response = http_get(self._get_url(f'/checkpoint_path/{run_id}/{checkpoint_id}')) - data = self._handle_response(response) + response = self._transport.get(self._get_url(f'/checkpoint_path/{run_id}/{checkpoint_id}')) + data = response.json() return ParsedCheckpointTwinklePath( path=data.get('path', ''), twinkle_path=data.get('twinkle_path', ''), @@ -333,7 +382,7 @@ def get_checkpoint_twinkle_path(self, run_id: str, checkpoint_id: str) -> str: Twinkle path string (e.g. "twinkle://run_id/weights/checkpoint_name"). Raises: - TwinkleClientError: If checkpoint not found or access denied. + TwinkleHTTPError: If checkpoint not found or access denied. """ return self.get_checkpoint_path(run_id, checkpoint_id).twinkle_path @@ -346,14 +395,14 @@ def delete_checkpoint(self, run_id: str, checkpoint_id: str) -> DeleteCheckpoint checkpoint_id: The checkpoint identifier. Returns: - :class:`~twinkle_client.types.server.DeleteCheckpointResponse` indicating success. + :class:`~twinkle.protocol.types.server.DeleteCheckpointResponse` indicating success. Raises: - TwinkleClientError: If checkpoint not found or access denied. + TwinkleHTTPError: If checkpoint not found or access denied. """ url = self._get_url(f'/training_runs/{run_id}/checkpoints/{checkpoint_id}') - response = http_delete(url) - data = self._handle_response(response) + response = self._transport.delete(url) + data = response.json() return DeleteCheckpointResponse(**data) # ------------------------------------------------------------------ @@ -368,21 +417,21 @@ def get_weights_info(self, twinkle_path: str) -> WeightsInfoResponse: twinkle_path: The twinkle:// path to the weights. Returns: - :class:`~twinkle_client.types.training.WeightsInfoResponse` with fields: + :class:`~twinkle.protocol.types.training.WeightsInfoResponse` with fields: ``training_run_id``, ``base_model``, ``model_owner``, ``is_lora``, ``lora_rank``. Raises: - TwinkleClientError: If weights not found or access denied. + TwinkleHTTPError: If weights not found or access denied. """ - response = http_post(self._get_url('/weights_info'), json_data={'twinkle_path': twinkle_path}) - data = self._handle_response(response) + response = self._transport.post(self._get_url('/weights_info'), json_data={'twinkle_path': twinkle_path}) + data = response.json() return WeightsInfoResponse(**data) # ------------------------------------------------------------------ # Convenience Methods # ------------------------------------------------------------------ - def get_latest_checkpoint_path(self, run_id: str) -> Optional[str]: + def get_latest_checkpoint_path(self, run_id: str) -> str | None: """ Get the filesystem path to the latest checkpoint for a training run. @@ -395,7 +444,7 @@ def get_latest_checkpoint_path(self, run_id: str) -> Optional[str]: Filesystem path string to the latest checkpoint, or ``None`` if none exist. Raises: - TwinkleClientError: If run not found or access denied. + TwinkleHTTPError: If run not found or access denied. """ checkpoints = self.list_checkpoints(run_id) if not checkpoints: @@ -403,7 +452,7 @@ def get_latest_checkpoint_path(self, run_id: str) -> Optional[str]: latest = checkpoints[-1] return self.get_checkpoint_path(run_id, latest.checkpoint_id).path - def find_training_run_by_model(self, base_model: str) -> List[TrainingRun]: + def find_training_run_by_model(self, base_model: str) -> list[TrainingRun]: """ Find training runs for a specific base model. @@ -411,7 +460,7 @@ def find_training_run_by_model(self, base_model: str) -> List[TrainingRun]: base_model: The base model name to search for. Returns: - List of :class:`~twinkle_client.types.training.TrainingRun` objects + List of :class:`~twinkle.protocol.types.training.TrainingRun` objects matching the base model. """ all_runs = self.list_training_runs(limit=100) diff --git a/src/twinkle_client/model/__init__.py b/src/twinkle_client/model/__init__.py index 94e3538af..b4c8012f0 100644 --- a/src/twinkle_client/model/__init__.py +++ b/src/twinkle_client/model/__init__.py @@ -1 +1,3 @@ from .multi_lora_transformers import MultiLoraTransformersModel + +__all__ = ['MultiLoraTransformersModel'] diff --git a/src/twinkle_client/model/multi_lora_transformers.py b/src/twinkle_client/model/multi_lora_transformers.py index 3471ca6c0..a9fbfd8ed 100644 --- a/src/twinkle_client/model/multi_lora_transformers.py +++ b/src/twinkle_client/model/multi_lora_transformers.py @@ -1,30 +1,32 @@ -from typing import Any, Dict, Optional +from __future__ import annotations + +import itertools +import logging +import threading +from collections.abc import Mapping from pathlib import Path -import time -from twinkle_client.http import http_get, http_post -from twinkle_client.common.json_utils import json_safe -from twinkle_client.types.component import DataRef -from twinkle_client.types.model import ( - CalculateLossResponse, - CalculateMetricResponse, - ClipGradNormResponse, - ForwardBackwardResponse, - ForwardResponse, - GetStateDictResponse, - GetTrainConfigsResponse, - SaveResponse, - TrainingProgressResponse, -) - - -def _data_ref_payload(inputs: DataRef | list[DataRef]) -> dict[str, Any]: +from typing import TYPE_CHECKING, Any, Dict, Optional + +if TYPE_CHECKING: + from peft import LoraConfig + +from twinkle.protocol.types import model as model_types +from twinkle.protocol.types.component import DataRef +from twinkle_client._request_builder import build_request +from twinkle_client.http import ClientTransport +from twinkle_client.http.context import capture_transport + +logger = logging.getLogger('twinkle_client') + + +def _data_refs(inputs: DataRef | list[DataRef]) -> list[dict[str, Any]]: """Encode one or more opaque references for a DataPlane model endpoint.""" refs = [inputs] if isinstance(inputs, DataRef) else list(inputs) if not refs: raise ValueError('at least one DataRef is required') if not all(isinstance(item, DataRef) for item in refs): raise TypeError('data-plane model inputs must contain only DataRef values') - return {'input_refs': [item.model_dump() for item in refs]} + return [item.model_dump() for item in refs] class MultiLoraTransformersModel: @@ -32,69 +34,187 @@ class MultiLoraTransformersModel: This client manages adapters and sends training/inference requests to the model server. The server-side session (managed by TwinkleClient) keeps the model alive. + + Every method builds its endpoint's request model rather than a dict, so a + misspelled or wrongly-typed argument fails here -- in the caller's own stack trace, + with no request sent. Arguments that are not declared fields (loss inputs, plugin + constructor arguments) are routed into that model's passthrough region, so public + signatures stay ``**kwargs`` and callers are unchanged. """ - def __init__(self, model_id: str, **kwargs): - """Initialize model client.""" - from twinkle_client.http import get_base_url - self.server_url = get_base_url() + def __init__( + self, + model_id: str, + *, + transport: ClientTransport | None = None, + **kwargs, + ): + """Initialize a model wrapper bound to one immutable request identity.""" + self._transport = capture_transport(transport) kwargs.pop('data_plane_url', None) if '://' in model_id: model_id = model_id.split('://')[1] self.model_id = model_id - self.server_url = f'{self.server_url}/model/{model_id}/twinkle' + self.server_url = f'{self._transport.context.base_url}/model/{model_id}/twinkle' self.adapter_name = None - response = http_post( - url=f'{self.server_url}/create', - ) - response.raise_for_status() + # Per-client monotonic sequence for idempotent dedup of stateful training ops: + # the server dedups on (session_id, seq_id) so a retried grad/step call is + # applied at most once. Reserved once per call and reused on retry. + self._seq_counter = itertools.count(1) + self._seq_lock = threading.Lock() + # The server-side component is created lazily on first use rather than in + # __init__, so constructing the wrapper performs no network I/O. + self._created = False + self._create_lock = threading.Lock() + + # ------------------------------------------------------------------ # + # Request plumbing + # ------------------------------------------------------------------ # + + def _ensure_created(self) -> None: + """Create the server-side model component once, on first use. + + Deferred out of ``__init__`` so construction has no side effect: a failed + ``create`` surfaces from the first operation instead of leaving a + half-initialised object published. Idempotent and thread-safe. + """ + if self._created: + return + with self._create_lock: + if self._created: + return + self._transport.post(f'{self.server_url}/create') + self._created = True + + def _submit(self, endpoint: str, model_cls, response_cls, **values): + """Build, send, and resolve one twinkle-native request.""" + self._ensure_created() + body = build_request(model_cls, **values) + response = self._transport.post_model(f'{self.server_url}/{endpoint}', body) + return self._await_task(response, response_cls) + + def _await_task(self, response, model_cls): + """Resolve a Submit_Endpoint response through the Client_Future_Layer. + + Blocks until the task is terminal and returns the deserialized ``model_cls`` + result (or ``None``), raising ``TaskFailedError`` on a failed terminal state. + Keeps every public method's synchronous signature unchanged. + """ + from twinkle_client._future import resolve_response + return resolve_response(response, model_cls, transport=self._transport) + + def _next_seq_id(self) -> int: + """Reserve the next monotonic seq_id for a stateful op (dedup key with session).""" + with self._seq_lock: + return next(self._seq_counter) - def add_adapter_to_model(self, adapter_name: str, config: Optional[Dict[str, Any]] = None, **kwargs) -> None: + # ------------------------------------------------------------------ # + # Adapter lifecycle + # ------------------------------------------------------------------ # + + def add_adapter_to_model( + self, + adapter_name: str, + config: LoraConfig | Mapping[str, Any] | None = None, + **kwargs, + ) -> None: """Add a new adapter to the model. Pass a peft ``LoraConfig`` (or its dict form) for LoRA training against a LoRA-mode deployment. Pass ``config=None`` for full-parameter training against a ``train_mode: full`` deployment. """ - save_dir = kwargs.get('save_dir') + if isinstance(config, Mapping): + from peft import LoraConfig + + config = LoraConfig(**config) + save_dir = kwargs.pop('save_dir', None) if save_dir: - kwargs['save_dir'] = Path(save_dir).expanduser().resolve().as_posix() - response = http_post( - url=f'{self.server_url}/add_adapter_to_model', - json_data={'adapter_name': adapter_name, 'config': config, **kwargs} - ) - response.raise_for_status() + save_dir = Path(save_dir).expanduser().resolve().as_posix() + self._submit( + 'add_adapter_to_model', + model_types.AddAdapterRequest, + None, + adapter_name=adapter_name, + config=config, + save_dir=save_dir, + **kwargs) self.adapter_name = adapter_name def remove_adapter(self, adapter_name: str | None = None) -> None: """Release one client-owned adapter from the training component.""" name = adapter_name or self.adapter_name - response = http_post( - url=f'{self.server_url}/remove_adapter', - json_data={'adapter_name': name}, - ) - response.raise_for_status() + self._submit('remove_adapter', model_types.AdapterRequest, None, adapter_name=name) if name == self.adapter_name: self.adapter_name = None - def forward(self, inputs: Any, **kwargs) -> ForwardResponse: + # ------------------------------------------------------------------ # + # Inline forward family + # ------------------------------------------------------------------ # + + def forward(self, inputs: Any, **kwargs) -> model_types.ForwardResponse: """Execute forward pass on inline model inputs.""" - response = http_post( - url=f'{self.server_url}/forward', - json_data={'inputs': inputs, 'adapter_name': self.adapter_name, **kwargs}, - ) - response.raise_for_status() - return ForwardResponse(**response.json()) - - def forward_only(self, inputs: Any, **kwargs) -> ForwardResponse: + return self._submit( + 'forward', + model_types.ForwardRequest, + model_types.ForwardResponse, + inputs=inputs, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + **kwargs) + + def forward_only(self, inputs: Any, **kwargs) -> model_types.ForwardResponse: """Execute forward pass without gradient computation on inline inputs.""" - response = http_post( - url=f'{self.server_url}/forward_only', - json_data={'inputs': inputs, 'adapter_name': self.adapter_name, **kwargs}, - ) - response.raise_for_status() - return ForwardResponse(**response.json()) + return self._submit( + 'forward_only', + model_types.ForwardOnlyRequest, + model_types.ForwardResponse, + inputs=inputs, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + **kwargs) + + def forward_backward(self, inputs: Any, **kwargs) -> model_types.ForwardBackwardResponse: + """Execute combined forward and backward pass on inline inputs.""" + return self._submit( + 'forward_backward', + model_types.ForwardBackwardTaskRequest, + model_types.ForwardBackwardResponse, + inputs=inputs, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + seq_id=self._next_seq_id(), + **kwargs) + + def calculate_loss(self, **kwargs) -> model_types.CalculateLossResponse: + """Calculate loss from model outputs.""" + return self._submit( + 'calculate_loss', + model_types.AdapterRequest, + model_types.CalculateLossResponse, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + **kwargs) + + def get_train_configs(self, **kwargs) -> model_types.GetTrainConfigsResponse: + """Get training configs.""" + return self._submit( + 'get_train_configs', + model_types.AdapterRequest, + model_types.GetTrainConfigsResponse, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + **kwargs) + + def backward(self, **kwargs) -> None: + """Execute backward pass.""" + self._submit( + 'backward', + model_types.AdapterRequest, + None, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + seq_id=self._next_seq_id(), + **kwargs) + + # ------------------------------------------------------------------ # + # Data-plane forward family + # ------------------------------------------------------------------ # def forward_from_data_plane( self, @@ -103,20 +223,17 @@ def forward_from_data_plane( input_field: str | None = None, kwarg_fields: dict[str, str] | None = None, **kwargs, - ) -> ForwardResponse: + ) -> model_types.ForwardResponse: """Execute forward using rows referenced from the server DataPlane.""" - response = http_post( - url=f'{self.server_url}/forward_from_data_plane', - json_data={ - **_data_ref_payload(inputs), - 'adapter_name': self.adapter_name, - 'input_field': input_field, - 'kwarg_fields': kwarg_fields or {}, - **json_safe(kwargs), - }, - ) - response.raise_for_status() - return ForwardResponse(**response.json()) + return self._submit( + 'forward_from_data_plane', + model_types.DataPlaneForwardRequest, + model_types.ForwardResponse, + input_refs=_data_refs(inputs), + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + input_field=input_field, + kwarg_fields=kwarg_fields or {}, + **kwargs) def forward_only_from_data_plane( self, @@ -127,62 +244,23 @@ def forward_only_from_data_plane( output_ref: DataRef | None = None, output_fields: dict[str, str] | None = None, **kwargs, - ) -> ForwardResponse | DataRef: + ) -> model_types.ForwardResponse | DataRef: """Execute forward-only using DataPlane rows and optionally append outputs.""" - body = { - **_data_ref_payload(inputs), - 'adapter_name': self.adapter_name, - 'input_field': input_field, - 'kwarg_fields': kwarg_fields or {}, - 'output_ref': output_ref.model_dump() if output_ref is not None else None, - 'output_fields': output_fields or {}, - **json_safe(kwargs), - } - response = http_post( - url=f'{self.server_url}/forward_only_from_data_plane', - json_data=body, - ) - response.raise_for_status() - result = ForwardResponse(**response.json()) + result = self._submit( + 'forward_only_from_data_plane', + model_types.DataPlaneForwardOnlyRequest, + model_types.ForwardResponse, + input_refs=_data_refs(inputs), + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + input_field=input_field, + kwarg_fields=kwarg_fields or {}, + output_ref=output_ref.model_dump() if output_ref is not None else None, + output_fields=output_fields or {}, + **kwargs) if output_ref is not None: return DataRef(**result.result) return result - def calculate_loss(self, **kwargs) -> CalculateLossResponse: - """Calculate loss from model outputs.""" - response = http_post( - url=f'{self.server_url}/calculate_loss', - json_data={'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() - return CalculateLossResponse(**response.json()) - - def get_train_configs(self, **kwargs) -> GetTrainConfigsResponse: - """Get training configs.""" - response = http_post( - url=f'{self.server_url}/get_train_configs', - json_data={'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() - return GetTrainConfigsResponse(**response.json()) - - def backward(self, **kwargs) -> None: - """Execute backward pass.""" - response = http_post( - url=f'{self.server_url}/backward', - json_data={'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() - - def forward_backward(self, inputs: Any, **kwargs) -> ForwardBackwardResponse: - """Execute combined forward and backward pass on inline inputs.""" - response = http_post( - url=f'{self.server_url}/forward_backward', - json_data={'inputs': inputs, 'adapter_name': self.adapter_name, **kwargs}, - ) - response.raise_for_status() - return ForwardBackwardResponse(**response.json()) - def forward_backward_from_data_plane( self, inputs: DataRef | list[DataRef], @@ -190,208 +268,233 @@ def forward_backward_from_data_plane( input_field: str | None = None, kwarg_fields: dict[str, str] | None = None, **kwargs, - ) -> ForwardBackwardResponse: + ) -> model_types.ForwardBackwardResponse: """Execute forward/backward using rows referenced from the server DataPlane.""" - response = http_post( - url=f'{self.server_url}/forward_backward_from_data_plane', - json_data={ - **_data_ref_payload(inputs), - 'adapter_name': self.adapter_name, - 'input_field': input_field, - 'kwarg_fields': kwarg_fields or {}, - **json_safe(kwargs), - }, - ) - response.raise_for_status() - return ForwardBackwardResponse(**response.json()) + return self._submit( + 'forward_backward_from_data_plane', + model_types.DataPlaneForwardRequest, + model_types.ForwardBackwardResponse, + input_refs=_data_refs(inputs), + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + input_field=input_field, + kwarg_fields=kwarg_fields or {}, + seq_id=self._next_seq_id(), + **kwargs) + + # ------------------------------------------------------------------ # + # Optimizer / scheduler steps + # ------------------------------------------------------------------ # def step(self, **kwargs) -> None: """Execute optimizer step.""" - response = http_post( - url=f'{self.server_url}/step', - json_data={'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() + self._submit( + 'step', + model_types.StepRequest, + None, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + seq_id=self._next_seq_id(), + **kwargs) def zero_grad(self, **kwargs) -> None: """Zero out gradients.""" - response = http_post( - url=f'{self.server_url}/zero_grad', - json_data={'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() + self._submit( + 'zero_grad', + model_types.AdapterRequest, + None, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + **kwargs) def lr_step(self, **kwargs) -> None: """Execute learning rate scheduler step.""" - response = http_post( - url=f'{self.server_url}/lr_step', - json_data={'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() - - def clip_grad_norm(self, max_grad_norm: float = 1.0, norm_type: int = 2, **kwargs) -> ClipGradNormResponse: + self._submit( + 'lr_step', + model_types.LrStepRequest, + None, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + seq_id=self._next_seq_id(), + **kwargs) + + def clip_grad_norm(self, + max_grad_norm: float = 1.0, + norm_type: int = 2, + **kwargs) -> model_types.ClipGradNormResponse: """Clip gradient norm.""" - response = http_post( - url=f'{self.server_url}/clip_grad_norm', - json_data={'max_grad_norm': max_grad_norm, 'norm_type': norm_type, 'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() - return ClipGradNormResponse(**response.json()) + return self._submit( + 'clip_grad_norm', + model_types.ClipGradNormRequest, + model_types.ClipGradNormResponse, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + max_grad_norm=max_grad_norm, + norm_type=norm_type, + **kwargs) def clip_grad_and_step(self, max_grad_norm: float = 1.0, norm_type: int = 2, **kwargs) -> None: """Clip gradient norm and execute optimizer step in one call.""" - response = http_post( - url=f'{self.server_url}/clip_grad_and_step', - json_data={'max_grad_norm': max_grad_norm, 'norm_type': norm_type, 'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() + self._submit( + 'clip_grad_and_step', + model_types.ClipGradAndStepRequest, + None, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + max_grad_norm=max_grad_norm, + norm_type=norm_type, + seq_id=self._next_seq_id(), + **kwargs) + + # ------------------------------------------------------------------ # + # Plugin setters + # ------------------------------------------------------------------ # def set_loss(self, loss_cls: str, **kwargs) -> None: """Set the loss function.""" - response = http_post( - url=f'{self.server_url}/set_loss', - json_data={'loss_cls': loss_cls, 'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() + self._submit( + 'set_loss', + model_types.SetLossRequest, + None, + loss_cls=loss_cls, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + **kwargs) def set_optimizer(self, optimizer_cls: str, **kwargs) -> None: """Set the optimizer.""" - response = http_post( - url=f'{self.server_url}/set_optimizer', - json_data={'optimizer_cls': optimizer_cls, 'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() + self._submit( + 'set_optimizer', + model_types.SetOptimizerRequest, + None, + optimizer_cls=optimizer_cls, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + **kwargs) def set_lr_scheduler(self, scheduler_cls: str, **kwargs) -> None: """Set the learning rate scheduler.""" - response = http_post( - url=f'{self.server_url}/set_lr_scheduler', - json_data={'scheduler_cls': scheduler_cls, 'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() - - def save(self, name: str, **kwargs) -> SaveResponse: - """Save model checkpoint.""" - response = http_post( - url=f'{self.server_url}/save', - json_data={'name': name, 'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() - return SaveResponse(**response.json()) - - def load(self, name: str, **kwargs) -> None: - """Load model checkpoint.""" - response = http_post( - url=f'{self.server_url}/load', - json_data={'name': name, 'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() - - def resume_from_checkpoint(self, name: str, *, resume_only_model: bool = False, **kwargs) -> Dict[str, Any]: - response = http_post( - url=f'{self.server_url}/resume_from_checkpoint', - json_data={'name': name, 'adapter_name': self.adapter_name, - 'resume_only_model': resume_only_model, **kwargs} - ) - response.raise_for_status() - return TrainingProgressResponse(**response.json()).result - - def apply_patch(self, patch_cls: str, **kwargs) -> None: - """Apply a patch to the model.""" - response = http_post( - url=f'{self.server_url}/apply_patch', - json_data={'patch_cls': patch_cls, 'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() - - def add_metric(self, metric_cls: str, is_training: Optional[bool] = None, **kwargs) -> None: - """Add a metric to the model.""" - response = http_post( - url=f'{self.server_url}/add_metric', - json_data={'metric_cls': metric_cls, 'is_training': is_training, 'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() + self._submit( + 'set_lr_scheduler', + model_types.SetLrSchedulerRequest, + None, + scheduler_cls=scheduler_cls, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + **kwargs) def set_template(self, template_cls: str, **kwargs) -> None: - """Set the template for data processing.""" - response = http_post( - url=f'{self.server_url}/set_template', - json_data={'template_cls': template_cls, 'adapter_name': self.adapter_name, 'model_id': self.model_id, **kwargs} - ) - response.raise_for_status() + """Set the template for data processing. + + ``model_id`` is not injected here: the backend always overrides it with its own + tokenizer id, so sending it made the request advertise a parameter that had no + effect. A caller that passes it explicitly still reaches the template + constructor through the passthrough region. + """ + self._submit( + 'set_template', + model_types.SetTemplateRequest, + None, + template_cls=template_cls, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + **kwargs) def set_processor(self, processor_cls: str, **kwargs) -> None: """Set the input processor.""" - response = http_post( - url=f'{self.server_url}/set_processor', - json_data={'processor_cls': processor_cls, 'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() + self._submit( + 'set_processor', + model_types.SetProcessorRequest, + None, + processor_cls=processor_cls, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + **kwargs) + + def add_metric(self, metric_cls: str, is_training: bool | None = None, **kwargs) -> None: + """Add a metric to the model.""" + self._submit( + 'add_metric', + model_types.AddMetricRequest, + None, + metric_cls=metric_cls, + is_training=is_training, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + **kwargs) - def calculate_metric(self, is_training: bool = True, **kwargs) -> CalculateMetricResponse: + def apply_patch(self, patch_cls: str, **kwargs) -> None: + """Apply a patch to the model.""" + self._submit( + 'apply_patch', + model_types.ApplyPatchRequest, + None, + patch_cls=patch_cls, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + **kwargs) + + def calculate_metric(self, is_training: bool = True, **kwargs) -> model_types.CalculateMetricResponse: """Calculate metrics from model outputs.""" - response = http_post( - url=f'{self.server_url}/calculate_metric', - json_data={'is_training': is_training, 'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() - return CalculateMetricResponse(**response.json()) - - def get_state_dict(self, **kwargs) -> GetStateDictResponse: - """Get model state dictionary.""" - response = http_post( - url=f'{self.server_url}/get_state_dict', - json_data={'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() - return GetStateDictResponse(**response.json()) + return self._submit( + 'calculate_metric', + model_types.CalculateMetricRequest, + model_types.CalculateMetricResponse, + is_training=is_training, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + **kwargs) + + # ------------------------------------------------------------------ # + # Checkpoint I/O + # ------------------------------------------------------------------ # + + def save(self, name: str, **kwargs) -> model_types.SaveResponse: + """Save model checkpoint.""" + return self._submit( + 'save', + model_types.SaveRequest, + model_types.SaveResponse, + name=name, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + **kwargs) + + def load(self, name: str, **kwargs) -> None: + """Load model checkpoint.""" + self._submit( + 'load', + model_types.LoadRequest, + None, + name=name, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + **kwargs) + + def resume_from_checkpoint(self, name: str, *, resume_only_model: bool = False, **kwargs) -> dict[str, Any]: + """Resume weights (and optionally optimizer state) from a checkpoint.""" + progress = self._submit( + 'resume_from_checkpoint', + model_types.ResumeFromCheckpointRequest, + model_types.TrainingProgressResponse, + name=name, + adapter_name=kwargs.pop('adapter_name', self.adapter_name), + resume_only_model=resume_only_model, + **kwargs) + return progress.result def upload_to_hub( self, checkpoint_dir: str, hub_model_id: str, - hub_token: Optional[str] = None, + hub_token: str | None = None, async_upload: bool = True, poll_interval: float = 5.0, ) -> None: """Upload model checkpoint to hub. - Submits the upload task to the server and polls for completion. - Blocks until the upload finishes or raises on failure. + Submits the upload task and blocks (via the Client_Future_Layer) until it + finishes, raising ``TaskFailedError`` on failure. Args: checkpoint_dir: The directory path of the checkpoint to upload. hub_model_id: The hub model id. hub_token: The hub token (optional). async_upload: Deprecated, has no effect. The server always runs the - upload in the background and the client polls for completion. - poll_interval: Seconds between status poll requests (default: 5). + upload in the background and the client waits via the future layer. + poll_interval: Deprecated, has no effect. Pacing is now owned by the + server-side long-poll of the Retrieve_Endpoint. """ - response = http_post( - url=f'{self.server_url}/upload_to_hub', - json_data={ - 'checkpoint_dir': checkpoint_dir, - 'hub_model_id': hub_model_id, - 'hub_token': hub_token, - } - ) - response.raise_for_status() - request_id = response.json().get('request_id') - if not request_id: - return - - print(f'[upload_to_hub] Upload started (task {request_id}), waiting for completion...') - while True: - status_resp = http_get(url=f'{self.server_url}/upload_status/{request_id}') - status_resp.raise_for_status() - data = status_resp.json() - status = data.get('status', 'unknown') - if status == 'completed': - print(f'[upload_to_hub] Upload completed successfully.') - return - elif status == 'failed': - error = data.get('error', 'Unknown error') - raise RuntimeError(f'[upload_to_hub] Upload failed: {error}') - else: - print(f'[upload_to_hub] Status: {status}...') - time.sleep(poll_interval) + logger.info('[upload_to_hub] submitting upload, waiting for completion...') + self._submit( + 'upload_to_hub', + model_types.UploadToHubRequest, + None, + checkpoint_dir=checkpoint_dir, + hub_model_id=hub_model_id, + hub_token=hub_token) + logger.info('[upload_to_hub] upload completed successfully.') diff --git a/src/twinkle_client/processor/__init__.py b/src/twinkle_client/processor/__init__.py deleted file mode 100644 index 677da196c..000000000 --- a/src/twinkle_client/processor/__init__.py +++ /dev/null @@ -1 +0,0 @@ -from .base import InputProcessor diff --git a/src/twinkle_client/processor/base.py b/src/twinkle_client/processor/base.py deleted file mode 100644 index 377865274..000000000 --- a/src/twinkle_client/processor/base.py +++ /dev/null @@ -1,38 +0,0 @@ - -from typing import List, Literal, Optional, Union -from twinkle_client.http import http_post -from twinkle import DeviceMesh -from twinkle.data_format import InputFeature - -class InputProcessor(object): - """Client wrapper for InputProcessor that calls server HTTP endpoints.""" - - def __init__(self, device_mesh: Optional[DeviceMesh] = None, padding_free: bool = False, framework: Literal['transformers', 'megatron'] = 'transformers', **kwargs): - from twinkle_client.http import get_base_url - - self.server_url = f'{get_base_url()}/processor/twinkle' - response = http_post( - url=f'{self.server_url}/create', - json_data={ - 'processor_type': 'processor', - 'class_type': 'InputProcessor', - **{'device_mesh': device_mesh, 'padding_free': padding_free, 'framework': framework}, **kwargs - } - ) - response.raise_for_status() - self.processor_id = response.json()['processor_id'] - - - def __call__(self, inputs: Union[InputFeature, List[InputFeature]], **kwargs): - response = http_post( - url=f'{self.server_url}/call', - json_data={ - 'processor_id': self.processor_id, - 'function': '__call__', - **{'inputs': inputs}, - **kwargs - } - ) - response.raise_for_status() - return response.json()["result"] - \ No newline at end of file diff --git a/src/twinkle_client/py.typed b/src/twinkle_client/py.typed new file mode 100644 index 000000000..e13a090dc --- /dev/null +++ b/src/twinkle_client/py.typed @@ -0,0 +1 @@ +# PEP 561 marker diff --git a/src/twinkle_client/rollout/multi_turn.py b/src/twinkle_client/rollout/multi_turn.py index c20597af2..911783a9e 100644 --- a/src/twinkle_client/rollout/multi_turn.py +++ b/src/twinkle_client/rollout/multi_turn.py @@ -28,6 +28,16 @@ from twinkle_client.sampler import vLLMSampler +@dataclasses.dataclass +class _RolloutState: + pifs: list[dict[str, Any]] + all_logprobs: list[list[Any]] + stop_reasons: list[str | None] + turns: list[int] + truncated: list[bool] + done: list[bool] + + class ClientMultiTurnRollout: """Agentic multi-turn rollout with tool use, driven over HTTP. @@ -104,167 +114,150 @@ def __call__(self, trajectories: List[Trajectory], **kwargs) -> List[Trajectory] if n == 0: return [] - sampling_params = self._as_sampling_params_dict( - kwargs.get('sampling_params', self.sampling_params)) - tool_managers = self._resolve_tool_managers( - kwargs.get('tool_manager', self.tool_manager), n) - - # 1. Encode each trajectory once; ``pifs[i]`` is the live per-turn - # state for trajectory ``i``. ``vLLMSampler.sample`` is responsible for - # JSON-serialising the feature (ndarray / tensor -> list) before the - # HTTP POST, so no conversion is needed here. - pifs: List[Dict[str, Any]] = [] - for traj in trajectories: - pif = self.template.encode(traj, add_generation_prompt=True) - pif.setdefault('messages', list(traj.get('messages', []))) - pifs.append(pif) - - all_logprobs: List[List[Any]] = [[] for _ in range(n)] - stop_reasons: List[Optional[str]] = [None] * n - turns: List[int] = [0] * n - truncated: List[bool] = [False] * n - done: List[bool] = [False] * n + sampling_params = self._as_sampling_params_dict(kwargs.get('sampling_params', self.sampling_params)) + tool_managers = self._resolve_tool_managers(kwargs.get('tool_manager', self.tool_manager), n) + state = self._initialize_state(trajectories) for _ in range(self.max_turns): - active = [i for i in range(n) if not done[i]] + active = [index for index in range(n) if not state.done[index]] if not active: break - # 2. One batched HTTP sample call for all currently-live - # trajectories. No device_mesh / min_batch_size padding: an HTTP - # client has no Ray DP ranks to align against. - # - # Passthrough contract: ``vLLMSampler.sample()`` may raise - # network / timeout / HTTP errors (e.g. requests exceptions). We - # deliberately do NOT wrap this call in try/except -- such errors - # propagate unchanged to the caller so ret/backoff policy stays an - # upstream concern (retry/backoff) and failures are never - # silently swallowed. - batch_pifs = [pifs[i] for i in active] - resps = self.sampler.sample(batch_pifs, sampling_params=sampling_params) - - pending_bridges: List[tuple] = [] # (global_idx, tool_messages) - for local_idx, global_idx in enumerate(active): - turns[global_idx] += 1 - seq = resps[local_idx].sequences[0] - - # ``new_input_feature`` is the running pif for the next round; - # the /twinkle/sample response contract guarantees it is set and - # carries ``input_ids``. A missing feature makes the next round - # impossible, so raise a batch/trajectory-indexed RuntimeError. - if seq.new_input_feature is None or 'input_ids' not in seq.new_input_feature: - raise RuntimeError( - f'Sampler returned a sequence without new_input_feature.input_ids at ' - f'batch index {local_idx} (trajectory {global_idx}); ' - f'cannot continue multi-turn.') - - pifs[global_idx] = dict(seq.new_input_feature) - # Per-round logprobs/token alignment guard: each sampled token - # must carry exactly one logprob entry. Mirrors the core-lib - # ``len(seq.logprobs) != len(seq.tokens)`` semantic so client and - # Ray paths cannot drift on this invariant. - if seq.logprobs is not None: - if len(seq.logprobs) != len(seq.tokens): - raise RuntimeError( - f'logprobs length ({len(seq.logprobs)}) does not match sampled ' - f'token count ({len(seq.tokens)}) at turn {turns[global_idx]} ' - f'(trajectory {global_idx})') - all_logprobs[global_idx].extend(seq.logprobs) - stop_reasons[global_idx] = seq.stop_reason - - # 3. Termination conditions. - # Cut off at ``max_tokens``: truncated, same as the max_turns and - # length-cap cases below, and same as ``MultiTurnRollout`` and - # ``ApiMultiTurnRollout``. Tool calls in the cut reply are still - # not dispatched. - if seq.stop_reason == 'length': - truncated[global_idx] = True - done[global_idx] = True - continue - - # 3a. Sequence-length cap. - if (self.max_trajectory_tokens is not None and len( - pifs[global_idx].get('input_ids') or []) >= self.max_trajectory_tokens): - truncated[global_idx] = True - done[global_idx] = True - continue - - # 3b. Parse tool calls from the freshly sampled assistant turn. - _msgs = pifs[global_idx].get('messages') or [] - _last_msg = _msgs[-1] if _msgs else None - tool_calls = (_last_msg.get('tool_calls') if isinstance(_last_msg, dict) else None) - if not tool_calls: - tool_calls = self.template.parse_tool_call(seq.decoded or '') - if not tool_calls: - done[global_idx] = True - continue - - # 3c. Hit the turn cap while still wanting to call a tool: force - # truncation. Also covers the ``max_turns == 1`` edge, where - # the very first sampled turn trips this branch. - if turns[global_idx] >= self.max_turns: - truncated[global_idx] = True - stop_reasons[global_idx] = 'max_turns' - done[global_idx] = True - continue - - # 4. Dispatch tools for this trajectory via its ToolManager. - tool_manager = tool_managers[global_idx] - if tool_manager is None: - raise ValueError( - f'trajectory {global_idx} produced tool_calls but no tool_manager ' - f'was provided (at construction time or as a per-call kwarg).') - tool_messages = [{ - 'role': 'tool', - 'content': tool_manager(tc), - } for tc in tool_calls] - pending_bridges.append((global_idx, tool_messages)) - - # Stitch bridge tokens (tool turns + next generation prompt) for - # every trajectory with outstanding tool turns. Reuses the shared - # pure function so client and core-lib paths cannot drift. - for global_idx, tool_messages in pending_bridges: - extended = extend_with_bridge(pifs[global_idx], tool_messages, self.template) - if extended is None: - # Trajectory exceeded max_length (truncation strategy 'delete'). - truncated[global_idx] = True - done[global_idx] = True - else: - pifs[global_idx] = extended - - # 4b. Final logprobs/labels alignment guard. For every trajectory that - # collected logprobs, the total logprob count must equal the number - # of trainable positions (labels != -100) in the final pif. This is - # the same invariant grpo._pad_and_align_to_batch relies on; a - # mismatch would silently corrupt GRPO old_logps alignment, so we - # fail loudly with the specific numbers. - for i in range(n): - if not all_logprobs[i]: + # One batched HTTP call for all live trajectories. Network and timeout + # errors intentionally propagate unchanged so retry policy stays upstream. + responses = self.sampler.sample( + [state.pifs[index] for index in active], + sampling_params=sampling_params, + ) + pending_bridges = self._process_responses(active, responses, state, tool_managers) + self._apply_bridges(state, pending_bridges) + + self._validate_logprob_alignment(state.pifs, state.all_logprobs) + return self._build_outputs( + trajectories, + state.pifs, + state.all_logprobs, + state.turns, + state.stop_reasons, + state.truncated, + ) + + # ------------------------------------------------------------------ private + + def _initialize_state(self, trajectories: List[Trajectory]) -> _RolloutState: + pifs: list[dict[str, Any]] = [] + for trajectory in trajectories: + pif = self.template.encode(trajectory, add_generation_prompt=True) + pif.setdefault('messages', list(trajectory.get('messages', []))) + pifs.append(pif) + size = len(trajectories) + return _RolloutState( + pifs=pifs, + all_logprobs=[[] for _ in range(size)], + stop_reasons=[None] * size, + turns=[0] * size, + truncated=[False] * size, + done=[False] * size, + ) + + def _process_responses(self, active, responses, state: _RolloutState, + tool_managers) -> list[tuple[int, list[dict]]]: + pending_bridges = [] + for local_index, global_index in enumerate(active): + sequence = responses[local_index].sequences[0] + tool_messages = self._process_sequence(local_index, global_index, sequence, state, + tool_managers[global_index]) + if tool_messages is not None: + pending_bridges.append((global_index, tool_messages)) + return pending_bridges + + def _process_sequence(self, local_index, global_index, sequence, state: _RolloutState, + tool_manager: ToolManager | None) -> list[dict] | None: + state.turns[global_index] += 1 + if sequence.new_input_feature is None or 'input_ids' not in sequence.new_input_feature: + raise RuntimeError(f'Sampler returned a sequence without new_input_feature.input_ids at ' + f'batch index {local_index} (trajectory {global_index}); ' + f'cannot continue multi-turn.') + + state.pifs[global_index] = dict(sequence.new_input_feature) + if sequence.logprobs is not None: + if len(sequence.logprobs) != len(sequence.tokens): + raise RuntimeError(f'logprobs length ({len(sequence.logprobs)}) does not match sampled ' + f'token count ({len(sequence.tokens)}) at turn {state.turns[global_index]} ' + f'(trajectory {global_index})') + state.all_logprobs[global_index].extend(sequence.logprobs) + state.stop_reasons[global_index] = sequence.stop_reason + + if sequence.stop_reason == 'length' or self._at_token_limit(state.pifs[global_index]): + state.truncated[global_index] = True + state.done[global_index] = True + return None + + messages = state.pifs[global_index].get('messages') or [] + last_message = messages[-1] if messages else None + tool_calls = last_message.get('tool_calls') if isinstance(last_message, dict) else None + tool_calls = tool_calls or self.template.parse_tool_call(sequence.decoded or '') + if not tool_calls: + state.done[global_index] = True + return None + if state.turns[global_index] >= self.max_turns: + state.truncated[global_index] = True + state.stop_reasons[global_index] = 'max_turns' + state.done[global_index] = True + return None + if tool_manager is None: + raise ValueError(f'trajectory {global_index} produced tool_calls but no tool_manager ' + f'was provided (at construction time or as a per-call kwarg).') + return [{'role': 'tool', 'content': tool_manager(tool_call)} for tool_call in tool_calls] + + def _at_token_limit(self, pif: dict[str, Any]) -> bool: + return self.max_trajectory_tokens is not None and len(pif.get('input_ids') or []) >= self.max_trajectory_tokens + + def _apply_bridges(self, state: _RolloutState, pending_bridges: list[tuple[int, list[dict]]]) -> None: + for global_index, tool_messages in pending_bridges: + extended = extend_with_bridge(state.pifs[global_index], tool_messages, self.template) + if extended is None: + state.truncated[global_index] = True + state.done[global_index] = True + else: + state.pifs[global_index] = extended + + @staticmethod + def _validate_logprob_alignment(pifs: List[Dict[str, Any]], all_logprobs: List[List[Any]]) -> None: + """Reject output that would corrupt downstream GRPO old-logprob alignment.""" + for index, logprobs in enumerate(all_logprobs): + if not logprobs: continue - labels_i = pifs[i].get('labels') or [] - trainable_i = sum(1 for label in labels_i if label != -100) - if len(all_logprobs[i]) != trainable_i: - raise RuntimeError(f'logprobs/labels misaligned for trajectory {i}: ' - f'{len(all_logprobs[i])} logprobs vs {trainable_i} ' + labels = pifs[index].get('labels') or [] + trainable = sum(1 for label in labels if label != -100) + if len(logprobs) != trainable: + raise RuntimeError(f'logprobs/labels misaligned for trajectory {index}: ' + f'{len(logprobs)} logprobs vs {trainable} ' f'trainable labels (labels != -100). This invariant is ' f'required by grpo._pad_and_align_to_batch; a mismatch ' f'would silently corrupt GRPO old_logps alignment.') - # 5. Merge pif fields into each trajectory dict at TOP LEVEL, preserving - # input length and order. - outs: List[Trajectory] = [] - for i, traj in enumerate(trajectories): - out = dict(traj) - out.update(pifs[i]) - out['messages'] = list(pifs[i].get('messages') or out.get('messages', [])) - out['logprobs'] = all_logprobs[i] if all_logprobs[i] else None - out['turns'] = turns[i] - out['stop_reason'] = stop_reasons[i] - out['truncated'] = truncated[i] - outs.append(out) - return outs - - # ------------------------------------------------------------------ private + @staticmethod + def _build_outputs( + trajectories: List[Trajectory], + pifs: List[Dict[str, Any]], + all_logprobs: List[List[Any]], + turns: List[int], + stop_reasons: List[Optional[str]], + truncated: List[bool], + ) -> List[Trajectory]: + """Merge final per-trajectory state while preserving input order.""" + outputs: List[Trajectory] = [] + for index, trajectory in enumerate(trajectories): + output = dict(trajectory) + output.update(pifs[index]) + output['messages'] = list(pifs[index].get('messages') or output.get('messages', [])) + output['logprobs'] = all_logprobs[index] or None + output['turns'] = turns[index] + output['stop_reason'] = stop_reasons[index] + output['truncated'] = truncated[index] + outputs.append(output) + return outputs @staticmethod def _as_sampling_params_dict(sampling_params) -> Optional[Dict[str, Any]]: diff --git a/src/twinkle_client/sampler/__init__.py b/src/twinkle_client/sampler/__init__.py index 06b961b3c..f0961b1e0 100644 --- a/src/twinkle_client/sampler/__init__.py +++ b/src/twinkle_client/sampler/__init__.py @@ -1 +1,3 @@ from .vllm_sampler import vLLMSampler + +__all__ = ['vLLMSampler'] diff --git a/src/twinkle_client/sampler/vllm_sampler.py b/src/twinkle_client/sampler/vllm_sampler.py index 4e271ce7a..caa5cc3fa 100644 --- a/src/twinkle_client/sampler/vllm_sampler.py +++ b/src/twinkle_client/sampler/vllm_sampler.py @@ -1,12 +1,17 @@ import asyncio from dataclasses import asdict -from typing import Any, Dict, List, Optional, Union -from twinkle_client.http import http_post -from twinkle_client.types.sampler import AddAdapterResponse, SampleResponseModel, SetTemplateResponse from peft import PeftConfig -from twinkle.data_format import Trajectory, InputFeature, SamplingParams -from twinkle_client.common.json_utils import json_safe -from twinkle_client.types.component import DataRef +from typing import Any, Dict, List, Optional, Union + +from twinkle.data_format import InputFeature, SamplingParams, Trajectory +from twinkle.protocol.json_utils import json_safe +from twinkle.protocol.types.component import DataPlaneSampleRequest, DataRef, UnloadAdapterPathsRequest +from twinkle.protocol.types.sampler import (SamplerAddAdapterRequest, SamplerAddAdapterResponse, SampleRequest, + SampleResponseModel, SampleResponseModelList, SamplerSetTemplateRequest, + SamplerSetTemplateResponse) +from twinkle_client._request_builder import build_request +from twinkle_client.http import ClientTransport +from twinkle_client.http.context import capture_transport # Intentionally does NOT subclass ``twinkle.sampler.base.Sampler``: importing @@ -31,35 +36,38 @@ class vLLMSampler: The server-side session (managed by TwinkleClient) keeps the sampler alive. """ - def __init__(self, model_id: str, **kwargs): - """Create the sampler instance on server.""" - from twinkle_client.http import get_base_url - self.server_url = get_base_url() + def __init__( + self, + model_id: str, + *, + transport: ClientTransport | None = None, + **kwargs, + ): + """Create the sampler instance on server with one captured transport.""" + self._transport = capture_transport(transport) from twinkle_client.data_plane import DataPlaneClient - self.data_plane = DataPlaneClient(kwargs.pop('data_plane_url', None)) + self.data_plane = DataPlaneClient(kwargs.pop('data_plane_url', None), transport=self._transport) self.adapter_name = None if '://' in model_id: model_id = model_id.split('://')[1] self.model_id = model_id - self.server_url = f'{self.server_url}/sampler/{model_id}/twinkle' - response = http_post( - url=f'{self.server_url}/create', - json_data=kwargs - ) - response.raise_for_status() + self.server_url = f'{self._transport.context.base_url}/sampler/{model_id}/twinkle' + self._transport.post(f'{self.server_url}/create', json_data=kwargs) + + def _await_task(self, response, model_cls): + """Resolve a Submit_Endpoint response through the Client_Future_Layer.""" + from twinkle_client._future import resolve_response + return resolve_response(response, model_cls, transport=self._transport) - def add_adapter_to_sampler(self, adapter_name: str, config: PeftConfig, **kwargs) -> AddAdapterResponse: + def add_adapter_to_sampler(self, adapter_name: str, config: PeftConfig, **kwargs) -> SamplerAddAdapterResponse: """Add a new adapter to the sampler.""" if isinstance(config, PeftConfig): config = config.__dict__ - response = http_post( - url=f'{self.server_url}/add_adapter_to_sampler', - json_data={'adapter_name': adapter_name, 'config': config, **kwargs} - ) - response.raise_for_status() + body = build_request(SamplerAddAdapterRequest, adapter_name=adapter_name, config=config, **kwargs) + response = self._transport.post_model(f'{self.server_url}/add_adapter_to_sampler', body) self.adapter_name = adapter_name - return AddAdapterResponse(**response.json()) + return SamplerAddAdapterResponse(**response.json()) def sample( self, @@ -81,26 +89,27 @@ def sample( Returns: SampleResponseModel with 'sequences' list, each containing tokens, logprobs, stop_reason. """ + response = self._transport.post_model( + f'{self.server_url}/sample', + build_request( + SampleRequest, + inputs=_json_safe(inputs), + sampling_params=self._sampling_params(sampling_params, num_samples), + adapter_name=adapter_name, + adapter_uri=adapter_uri)) + return self._await_task(response, SampleResponseModelList).samples + + @staticmethod + def _sampling_params(sampling_params: Optional[Union[SamplingParams, Dict[str, Any]]], + num_samples: int) -> Dict[str, Any]: + """Normalise sampling parameters into the single dict the server builds from.""" if isinstance(sampling_params, SamplingParams): sampling_params = asdict(sampling_params) else: sampling_params = dict(sampling_params or {}) if num_samples != 1 and sampling_params.setdefault('num_samples', num_samples) != num_samples: raise ValueError('num_samples conflicts with sampling_params.num_samples') - json_data = { - 'inputs': _json_safe(inputs), - 'sampling_params': _json_safe(sampling_params), - 'adapter_name': adapter_name, - } - if adapter_uri is not None: - json_data['adapter_uri'] = adapter_uri - - response = http_post( - url=f'{self.server_url}/sample', - json_data=json_data - ) - response.raise_for_status() - return [SampleResponseModel(**r) for r in response.json()['samples']] + return _json_safe(sampling_params) def sample_to_data_plane( self, @@ -114,22 +123,18 @@ def sample_to_data_plane( num_samples: int = 1, ) -> DataRef: """Generate complete prompt groups and keep their rows in the server DataPlane.""" - body = { - 'sampling_params': sampling_params, - 'adapter_name': adapter_name, - 'adapter_uri': adapter_uri, - 'policy_version': policy_version, - 'group_ids': group_ids, - 'num_samples': num_samples, - } - body['input_ref' if isinstance(inputs, DataRef) else 'inputs'] = ( - inputs.model_dump() if isinstance(inputs, DataRef) else _json_safe(inputs)) - response = http_post( - url=f'{self.server_url}/sample_to_data_plane', - json_data=json_safe(body), - ) - response.raise_for_status() - return DataRef(**response.json()) + source = ({'input_ref': inputs.model_dump()} if isinstance(inputs, DataRef) else {'inputs': _json_safe(inputs)}) + body = build_request( + DataPlaneSampleRequest, + sampling_params=_json_safe(sampling_params) if sampling_params else None, + adapter_name=adapter_name, + adapter_uri=adapter_uri, + policy_version=policy_version, + group_ids=group_ids, + num_samples=num_samples, + **source) + response = self._transport.post_model(f'{self.server_url}/sample_to_data_plane', body) + return self._await_task(response, DataRef) async def asample( self, @@ -174,25 +179,17 @@ async def asample_to_data_plane( def unload_adapter_paths(self, adapter_paths: list[str]) -> None: """Evict policy snapshots that are no longer referenced by this client.""" - response = http_post( - url=f'{self.server_url}/unload_adapter_paths', - json_data={'adapter_paths': adapter_paths}, - ) - response.raise_for_status() + self._transport.post_model(f'{self.server_url}/unload_adapter_paths', + build_request(UnloadAdapterPathsRequest, adapter_paths=adapter_paths)) - def set_template(self, template_cls: str, adapter_name: str = '', **kwargs) -> SetTemplateResponse: + def set_template(self, template_cls: str, adapter_name: str = '', **kwargs) -> SamplerSetTemplateResponse: """Set the template for encoding trajectories.""" - response = http_post( - url=f'{self.server_url}/set_template', - json_data={'template_cls': template_cls, 'adapter_name': adapter_name, **kwargs} - ) - response.raise_for_status() - return SetTemplateResponse(**response.json()) - + body = build_request(SamplerSetTemplateRequest, template_cls=template_cls, adapter_name=adapter_name, **kwargs) + response = self._transport.post_model(f'{self.server_url}/set_template', body) + return SamplerSetTemplateResponse(**response.json()) + def apply_patch(self, patch_cls: str, **kwargs) -> None: """Apply a patch to the model.""" - response = http_post( - url=f'{self.server_url}/apply_patch', - json_data={'patch_cls': patch_cls, 'adapter_name': self.adapter_name, **kwargs} - ) - response.raise_for_status() + from twinkle.protocol.types.model import ApplyPatchRequest + body = build_request(ApplyPatchRequest, patch_cls=patch_cls, adapter_name=self.adapter_name or '', **kwargs) + self._transport.post_model(f'{self.server_url}/apply_patch', body) diff --git a/src/twinkle_client/skills/base.py b/src/twinkle_client/skills/base.py index cf2bad788..1a64c7d24 100644 --- a/src/twinkle_client/skills/base.py +++ b/src/twinkle_client/skills/base.py @@ -14,10 +14,11 @@ from __future__ import annotations import dataclasses -from twinkle.utils.logger import get_logger from abc import ABC, abstractmethod from pathlib import Path +from twinkle.utils.logger import get_logger + logger = get_logger() # File stems to skip when scanning for skill markdown files diff --git a/src/twinkle_client/skills/bundled/twinkle-training.md b/src/twinkle_client/skills/bundled/twinkle-training.md index da582d3aa..281ae4342 100644 --- a/src/twinkle_client/skills/bundled/twinkle-training.md +++ b/src/twinkle_client/skills/bundled/twinkle-training.md @@ -828,7 +828,7 @@ rt.finish(status='completed') ### Example 5: Sampling / Inference Only ```python -from twinkle_client import init_twinkle_client +from twinkle import init_twinkle_client from twinkle_client.sampler import vLLMSampler # 1. Init diff --git a/src/twinkle_client/skills/manager.py b/src/twinkle_client/skills/manager.py index 34ceb9c56..6de4484d1 100644 --- a/src/twinkle_client/skills/manager.py +++ b/src/twinkle_client/skills/manager.py @@ -4,7 +4,6 @@ from __future__ import annotations from twinkle.utils.logger import get_logger - from twinkle_client.skills.base import Skill, SkillProvider logger = get_logger() @@ -72,10 +71,8 @@ def format_for_prompt(self) -> str: sections: list[str] = [] sections.append('# Available Skills') sections.append('') - sections.append( - 'The following skills provide you with specialized knowledge and capabilities. ' - 'Use them to better assist the user.' - ) + sections.append('The following skills provide you with specialized knowledge and capabilities. ' + 'Use them to better assist the user.') sections.append('') for skill in self._skills: diff --git a/src/twinkle_client/skills/modelscope_provider.py b/src/twinkle_client/skills/modelscope_provider.py index 9fe27f9c4..3e83c0013 100644 --- a/src/twinkle_client/skills/modelscope_provider.py +++ b/src/twinkle_client/skills/modelscope_provider.py @@ -4,9 +4,9 @@ from __future__ import annotations import asyncio -from twinkle.utils.logger import get_logger from pathlib import Path +from twinkle.utils.logger import get_logger from twinkle_client.skills.base import SkillProvider logger = get_logger() @@ -42,7 +42,11 @@ async def fetch(self) -> None: if (repo_dir / '.git').exists(): proc = await asyncio.create_subprocess_exec( - 'git', '-C', str(repo_dir), 'pull', '--ff-only', + 'git', + '-C', + str(repo_dir), + 'pull', + '--ff-only', stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.PIPE, ) @@ -52,8 +56,14 @@ async def fetch(self) -> None: else: self.cache_dir.mkdir(parents=True, exist_ok=True) proc = await asyncio.create_subprocess_exec( - 'git', 'clone', '--depth', '1', '--branch', self._branch, - self._repo_url, str(repo_dir), + 'git', + 'clone', + '--depth', + '1', + '--branch', + self._branch, + self._repo_url, + str(repo_dir), stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.PIPE, ) diff --git a/src/twinkle_client/types/__init__.py b/src/twinkle_client/types/__init__.py deleted file mode 100644 index 1c25324a1..000000000 --- a/src/twinkle_client/types/__init__.py +++ /dev/null @@ -1,106 +0,0 @@ -# Copyright (c) ModelScope Contributors. All rights reserved. -from .model import ( - AddAdapterRequest, - AddAdapterResponse, - AddMetricRequest, - AddMetricResponse, - AdapterRequest, - ApplyPatchRequest, - ApplyPatchResponse, - BackwardResponse, - CalculateLossResponse, - CalculateMetricRequest, - CalculateMetricResponse, - ClipGradAndStepRequest, - ClipGradAndStepResponse, - ClipGradNormResponse, - CreateRequest, - CreateResponse, - DataPlaneForwardOnlyRequest, - DataPlaneForwardRequest, - ForwardBackwardResponse, - ForwardOnlyRequest, - ForwardRequest, - ForwardResponse, - GetStateDictRequest, - GetStateDictResponse, - GetTrainConfigsResponse, - LoadRequest, - LoadResponse, - LrStepResponse, - ModelResult, - OkResponse, - ResumeFromCheckpointRequest, - SaveRequest, - SaveResponse, - SetLossRequest, - SetLossResponse, - SetLrSchedulerRequest, - SetLrSchedulerResponse, - SetOptimizerRequest, - SetOptimizerResponse, - SetProcessorRequest, - SetProcessorResponse, - SetTemplateRequest, - SetTemplateResponse, - StepResponse, - TrainingProgressResponse, - UploadToHubRequest, - UploadToHubResponse, - UploadStatusResponse, - ZeroGradResponse, -) -from .processor import ( - ProcessorCallRequest, - ProcessorCallResponse, - ProcessorCreateRequest, - ProcessorCreateResponse, - ProcessorHeartbeatRequest, - ProcessorHeartbeatResponse, -) -from .sampler import ( - AddAdapterRequest as SamplerAddAdapterRequest, - AddAdapterResponse, - CreateResponse as SamplerCreateResponse, - SampledSequenceModel, - SampleRequest, - SampleResponseModel, - SampleResponseModelList, - SetTemplateRequest as SamplerSetTemplateRequest, - SetTemplateResponse as SamplerSetTemplateResponse, -) -from .server import ( - CheckpointPathResponse, - DeleteCheckpointResponse, - ErrorResponse, - GetServerCapabilitiesResponse, - HealthResponse, - SupportedModel, - WeightsInfoRequest, - WeightsInfoResponse as ServerWeightsInfoResponse, - CapacityInfoResponse, -) -from .session import CreateSessionRequest, CreateSessionResponse, SessionHeartbeatRequest, SessionHeartbeatResponse -from .training import ( - Checkpoint, - CheckpointsListResponse, - CreateModelRequest, - Cursor, - LoraConfig, - ParsedCheckpointTwinklePath, - TrainingRun, - TrainingRunsResponse, - WeightsInfoResponse, -) - -from .checkpoint import ResolvedLoadPath -from .component import ( - DataAppendRequest, - DataGetRequest, - DataPlaneSampleRequest, - DataPutRequest, - DataRef, - DataReleaseRequest, - DataRowsResponse, - UnloadAdapterPathsRequest, -) diff --git a/src/twinkle_client/types/model.py b/src/twinkle_client/types/model.py deleted file mode 100644 index 3d4f4cf45..000000000 --- a/src/twinkle_client/types/model.py +++ /dev/null @@ -1,353 +0,0 @@ -# Copyright (c) ModelScope Contributors. All rights reserved. -""" -Pydantic request/response models for twinkle model management endpoints. - -These models are used by both the server-side handler and the twinkle client. -""" -from pydantic import BaseModel, Field, field_validator, model_validator -from typing import Any, Dict, List, Optional, Union - -from .component import DataRef - - -class CreateRequest(BaseModel): - - class Config: - extra = 'allow' - - -class ForwardRequest(BaseModel): - inputs: Any - adapter_name: str - - class Config: - extra = 'allow' - - -class ForwardOnlyRequest(BaseModel): - inputs: Any - adapter_name: Optional[str] = None - - class Config: - extra = 'allow' - - -class DataPlaneForwardRequest(BaseModel): - input_refs: List[DataRef] = Field(min_length=1) - input_field: str | None = None - kwarg_fields: Dict[str, str] = Field(default_factory=dict) - adapter_name: str - - class Config: - extra = 'allow' - - -class DataPlaneForwardOnlyRequest(DataPlaneForwardRequest): - output_ref: DataRef | None = None - output_fields: Dict[str, str] = Field(default_factory=dict) - - @model_validator(mode='after') - def validate_output(self) -> 'DataPlaneForwardOnlyRequest': - if (self.output_ref is None) != (len(self.output_fields) == 0): - raise ValueError('output_ref and output_fields must be configured together') - return self - - -class AdapterRequest(BaseModel): - adapter_name: str - - class Config: - extra = 'allow' - - -class SetLossRequest(BaseModel): - loss_cls: str - adapter_name: str - - class Config: - extra = 'allow' - - -class SetOptimizerRequest(BaseModel): - optimizer_cls: str - adapter_name: str - - class Config: - extra = 'allow' - - -class SetLrSchedulerRequest(BaseModel): - scheduler_cls: str - adapter_name: str - - class Config: - extra = 'allow' - - -class SaveRequest(BaseModel): - adapter_name: str - save_optimizer: bool = False - name: Optional[str] = None - is_sampler: bool = False # If True, delete existing sampler weights before saving - - class Config: - extra = 'allow' - - -class UploadToHubRequest(BaseModel): - checkpoint_dir: Union[str, Dict] - hub_model_id: str - hub_token: Optional[str] = None - async_upload: bool = False - - @field_validator('checkpoint_dir', mode='before') - @classmethod - def extract_checkpoint_dir(cls, v): - if isinstance(v, dict): - return v['twinkle_path'] - return v - - class Config: - extra = 'allow' - - -class LoadRequest(BaseModel): - adapter_name: str - load_optimizer: bool = False - name: str - - class Config: - extra = 'allow' - - -class ResumeFromCheckpointRequest(BaseModel): - """Request for /resume_from_checkpoint endpoint.""" - name: str - adapter_name: str = '' - resume_only_model: bool = False - - class Config: - extra = 'allow' - - -class AddAdapterRequest(BaseModel): - adapter_name: str - # ``config`` is None for full-parameter training (no LoRA adapter) and a - # serialized LoraConfig string for LoRA training. - config: Optional[str] = None - save_dir: Optional[str] = None - - class Config: - extra = 'allow' - - -class SetTemplateRequest(BaseModel): - template_cls: str - adapter_name: str - - class Config: - extra = 'allow' - - -class SetProcessorRequest(BaseModel): - processor_cls: str - adapter_name: str - - class Config: - extra = 'allow' - - -class CalculateMetricRequest(BaseModel): - adapter_name: str - is_training: bool = True - - class Config: - extra = 'allow' - - -class GetStateDictRequest(BaseModel): - adapter_name: str - - class Config: - extra = 'allow' - - -class ClipGradAndStepRequest(BaseModel): - adapter_name: str - max_grad_norm: float = 1.0 - norm_type: int = 2 - - class Config: - extra = 'allow' - - -class ApplyPatchRequest(BaseModel): - patch_cls: str - adapter_name: str - - class Config: - extra = 'allow' - - -class AddMetricRequest(BaseModel): - metric_cls: str - adapter_name: str - is_training: Optional[bool] = None - - class Config: - extra = 'allow' - - -# --------------------------------------------------------------------------- -# Response models -# --------------------------------------------------------------------------- - - -class OkResponse(BaseModel): - """Response for endpoints whose underlying method returns None.""" - status: str = 'ok' - - -class ModelResult(BaseModel): - """Generic single-value result wrapper returned by result-bearing endpoints.""" - result: Any - - -# --- Result-bearing responses --- - -class ForwardResponse(BaseModel): - """Response for /forward and /forward_only endpoints (returns ModelOutput).""" - result: Any - - -class ForwardBackwardResponse(BaseModel): - """Response for /forward_backward endpoint (returns ModelOutput).""" - result: Any - - -class CalculateLossResponse(BaseModel): - """Response for /calculate_loss endpoint (returns float).""" - result: float - - -class ClipGradNormResponse(BaseModel): - """Response for /clip_grad_norm endpoint (returns float as str).""" - result: str - - -class GetTrainConfigsResponse(BaseModel): - """Response for /get_train_configs endpoint (returns str).""" - result: str - - -class GetStateDictResponse(BaseModel): - """Response for /get_state_dict endpoint (returns Dict).""" - result: Dict[str, Any] - - -class CalculateMetricResponse(BaseModel): - """Response for /calculate_metric endpoint (returns Dict).""" - result: Dict[str, Any] - - -class SaveResponse(BaseModel): - """Response for /save endpoint (returns twinkle path + checkpoint dir).""" - twinkle_path: str - checkpoint_dir: Optional[str] = None - - -class TrainingProgressResponse(BaseModel): - """Response for /resume_from_checkpoint endpoint.""" - result: Dict[str, Any] - - -# --- Void responses (return None → OkResponse) --- - -class BackwardResponse(OkResponse): - """Response for /backward endpoint.""" - pass - - -class StepResponse(OkResponse): - """Response for /step (optimizer step) endpoint.""" - pass - - -class ZeroGradResponse(OkResponse): - """Response for /zero_grad endpoint.""" - pass - - -class LrStepResponse(OkResponse): - """Response for /lr_step endpoint.""" - pass - - -class SetLossResponse(OkResponse): - """Response for /set_loss endpoint.""" - pass - - -class SetOptimizerResponse(OkResponse): - """Response for /set_optimizer endpoint.""" - pass - - -class SetLrSchedulerResponse(OkResponse): - """Response for /set_lr_scheduler endpoint.""" - pass - - -class LoadResponse(OkResponse): - """Response for /load endpoint.""" - pass - - -class SetTemplateResponse(OkResponse): - """Response for /set_template endpoint.""" - pass - - -class SetProcessorResponse(OkResponse): - """Response for /set_processor endpoint.""" - pass - - -class UploadToHubResponse(BaseModel): - """Response for /upload_to_hub endpoint.""" - request_id: str - - -class UploadStatusResponse(BaseModel): - """Response for /upload_status/{request_id} endpoint.""" - request_id: str - status: str # pending / queued / running / completed / failed - error: Optional[str] = None - - -class ClipGradAndStepResponse(OkResponse): - """Response for /clip_grad_and_step endpoint.""" - pass - - -class ApplyPatchResponse(OkResponse): - """Response for /apply_patch endpoint.""" - pass - - -class AddMetricResponse(OkResponse): - """Response for /add_metric endpoint.""" - pass - - -# --- Other responses --- - -class CreateResponse(BaseModel): - """Response for /create endpoint.""" - status: str = 'ok' - - -class AddAdapterResponse(BaseModel): - """Response for /add_adapter_to_model endpoint.""" - status: str = 'ok' - adapter_name: str diff --git a/src/twinkle_client/types/processor.py b/src/twinkle_client/types/processor.py deleted file mode 100644 index fe8674ce1..000000000 --- a/src/twinkle_client/types/processor.py +++ /dev/null @@ -1,46 +0,0 @@ -# Copyright (c) ModelScope Contributors. All rights reserved. -""" -Pydantic request/response models for twinkle processor endpoints. - -These models are used by both the server-side handler and the twinkle client. - -Note: Class names are prefixed with 'Processor' to avoid name collisions when -importing from twinkle_client.types alongside model.py classes. -""" -from pydantic import BaseModel -from typing import Any - - -class ProcessorCreateRequest(BaseModel): - processor_type: str - class_type: str - - class Config: - extra = 'allow' - - -class ProcessorHeartbeatRequest(BaseModel): - processor_id: str - - -class ProcessorCallRequest(BaseModel): - processor_id: str - function: str - - class Config: - extra = 'allow' - - -class ProcessorCreateResponse(BaseModel): - """Response body for the /create endpoint.""" - processor_id: str - - -class ProcessorHeartbeatResponse(BaseModel): - """Response body for the /heartbeat endpoint.""" - status: str = 'ok' - - -class ProcessorCallResponse(BaseModel): - """Response body for the /call endpoint.""" - result: Any diff --git a/src/twinkle_client/types/sampler.py b/src/twinkle_client/types/sampler.py deleted file mode 100644 index 47ce1bdd6..000000000 --- a/src/twinkle_client/types/sampler.py +++ /dev/null @@ -1,76 +0,0 @@ -# Copyright (c) ModelScope Contributors. All rights reserved. -""" -Pydantic request/response models for twinkle sampler endpoints. - -These models are used by both the server-side handler and the twinkle client. -""" -from pydantic import BaseModel, Field -from typing import Any, Dict, List, Literal, Optional, Tuple - -StopReason = Literal['length', 'stop', 'abort', 'error'] - - -class SampleRequest(BaseModel): - """Request body for the /sample endpoint.""" - inputs: Any = Field(..., description='List of Trajectory or InputFeature dicts') - sampling_params: Optional[Dict[str, Any]] = Field( - None, description='Sampling parameters (max_tokens, temperature, num_samples, etc.)') - adapter_name: str = Field('', description='Adapter name for LoRA inference') - adapter_uri: Optional[str] = Field( - None, description='Adapter URI (twinkle:// path or local path) for LoRA inference') - - -class SampledSequenceModel(BaseModel): - """A single sampled sequence, mirroring twinkle.data_format.SampledSequence.""" - stop_reason: StopReason = Field(..., description="Stop reason: 'length' or 'stop'") - tokens: List[int] = Field(..., description='Token IDs of the sampled sequence') - logprobs: Optional[List[Optional[List[Tuple[int, float]]]]] = Field(None, description='Per-token log-probabilities') - decoded: Optional[str] = Field(None, description='Decoded text of the sampled sequence') - new_input_feature: Optional[Dict[str, Any]] = Field( - None, description='Updated InputFeature after sampling (input_ids, labels, etc.)') - - -class SampleResponseModel(BaseModel): - """Mirroring twinkle.data_format.SampleResponse.""" - sequences: List[SampledSequenceModel] = Field( - ..., description='List of sampled sequences') - prompt_token_ids: Optional[List[int]] = Field( - None, description='Token IDs of the prompt the sequences continue') - prompt_logprobs: Optional[List[Optional[float]]] = None - topk_prompt_logprobs: Optional[List[Optional[List[Tuple[int, float]]]]] = None - - -class SampleResponseModelList(BaseModel): - """Response body for the /sample endpoint""" - samples: List[SampleResponseModel] = Field(..., description='List of sample responses') - - -class SetTemplateRequest(BaseModel): - """Request body for the /set_template endpoint.""" - template_cls: str = Field(..., description="Template class name (e.g. 'Template')") - adapter_name: str = Field('', description='Adapter name to associate the template with') - - class Config: - extra = 'allow' - - -class SetTemplateResponse(BaseModel): - """Response body for the /set_template endpoint.""" - status: str = 'ok' - - -class AddAdapterRequest(BaseModel): - """Request body for the /add_adapter_to_sampler endpoint.""" - adapter_name: str = Field(..., description='Name of the adapter to add') - config: Any = Field(..., description='LoRA configuration dict') - - -class AddAdapterResponse(BaseModel): - """Response body for the /add_adapter_to_sampler endpoint.""" - status: str = 'ok' - adapter_name: str - - -class CreateResponse(BaseModel): - """Response body for the /create endpoint.""" - status: str = 'ok' diff --git a/src/twinkle_client/types/server.py b/src/twinkle_client/types/server.py deleted file mode 100644 index 1c7c992d7..000000000 --- a/src/twinkle_client/types/server.py +++ /dev/null @@ -1,49 +0,0 @@ -# Copyright (c) ModelScope Contributors. All rights reserved. -"""Shared Pydantic response models for the twinkle server health/error endpoints.""" -from pydantic import BaseModel -from typing import Any, List, Optional - - -class SupportedModel(BaseModel): - """Information about a supported model.""" - model_name: str - - -class GetServerCapabilitiesResponse(BaseModel): - """Response body for the /get_server_capabilities endpoint.""" - supported_models: List[SupportedModel] - - -class HealthResponse(BaseModel): - status: str - - -class DeleteCheckpointResponse(BaseModel): - success: bool - message: str - - -class ErrorResponse(BaseModel): - detail: str - - -class WeightsInfoRequest(BaseModel): - twinkle_path: str - - -class WeightsInfoResponse(BaseModel): - """Response body for the /weights_info endpoint.""" - weights_info: Any - - -class CheckpointPathResponse(BaseModel): - """Response body for the /checkpoint_path endpoint.""" - path: str - twinkle_path: str - - -class CapacityInfoResponse(BaseModel): - """Response body for the /capacity_info endpoint.""" - max_loras: int - used_loras: int - free_loras: int diff --git a/src/twinkle_client/utils/__init__.py b/src/twinkle_client/utils/__init__.py new file mode 100644 index 000000000..bcd3b883f --- /dev/null +++ b/src/twinkle_client/utils/__init__.py @@ -0,0 +1,7 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Client-internal utilities. + +Regular package (carries this ``__init__``) so ``patch_tinker`` is included by +``setuptools.packages.find`` in a built wheel; a namespace-only directory would be +dropped from the distribution. +""" diff --git a/src/twinkle_client/utils/patch_tinker.py b/src/twinkle_client/utils/patch_tinker.py index ed9abaa82..7b4088151 100644 --- a/src/twinkle_client/utils/patch_tinker.py +++ b/src/twinkle_client/utils/patch_tinker.py @@ -12,8 +12,8 @@ import os from typing import TYPE_CHECKING, Any, Dict, Mapping, Optional, Union -from twinkle_client.http.headers import build_routing_headers -from twinkle_client.http.utils import get_api_key, get_request_id +from twinkle.protocol.headers import build_routing_headers +from twinkle_client.http.context import get_api_key, get_request_id _patched = False _loss_fn_config_patched = False @@ -57,9 +57,8 @@ def _patched_async_tinker_init( if api_key is None: api_key = os.environ.get('TWINKLE_SERVER_TOKEN') if api_key is None: - raise TinkerError( - 'The api_key client option must be set either by passing api_key to the client or by setting the TWINKLE_SERVER_TOKEN environment variable' - ) + raise TinkerError('The api_key client option must be set either by passing api_key to the client ' + 'or by setting the TWINKLE_SERVER_TOKEN environment variable') # REMOVED: api_key 'tml-' prefix validation # Original code: # if not api_key.startswith("tml-"): @@ -120,6 +119,7 @@ def _patched_from_tinker_path(cls, tinker_path: str) -> Any: def _make_patched_service_client_init(original): + def _patched_service_client_init(self, user_metadata=None, **kwargs): """Patched version of ServiceClient.__init__ that injects Twinkle-specific headers.""" # Resolve api_key with the same priority order used by AsyncTinker: @@ -147,8 +147,8 @@ def _create_full_training_client_submit(self, base_model, seed=None, user_metada the training loop (forward_backward / optim_step / save_weights / ...) is identical to the LoRA path. """ - from tinker.lib.public_interfaces import service_client as _sc from tinker.lib.internal_client_holder import ClientConnectionPoolType + from tinker.lib.public_interfaces import service_client as _sc session_id = self.holder.get_session_id() model_seq_id = self.holder.get_training_client_id() @@ -184,8 +184,7 @@ async def _create_full_training_client_async(): def _create_full_training_client(self, base_model, seed=None, user_metadata=None): """Create a full-parameter (non-LoRA) training client (blocking).""" - return _create_full_training_client_submit( - self, base_model, seed=seed, user_metadata=user_metadata).result() + return _create_full_training_client_submit(self, base_model, seed=seed, user_metadata=user_metadata).result() async def _create_full_training_client_async(self, base_model, seed=None, user_metadata=None): diff --git a/tests/data_format/test_tq_fields.py b/tests/data_format/test_tq_fields.py new file mode 100644 index 000000000..e882f65c8 --- /dev/null +++ b/tests/data_format/test_tq_fields.py @@ -0,0 +1,59 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Tests for ``twinkle.data_format.tq_fields``. + +The packing helpers carry field-consistency validation and a numeric/non-numeric +type branch but previously had zero test coverage. These cover the three branches +required by (empty rows, inconsistent fields, numeric+non-numeric mix), plus a +subprocess assertion (not depending on the ``async-rl`` extra) that guards: the +module must import cleanly without ``tensordict``. +""" +import subprocess +import sys + +import pytest + + +def test_rows_to_tq_fields_empty_rows(): + pytest.importorskip('tensordict') + from twinkle.data_format import rows_to_tq_fields + + packed = rows_to_tq_fields([]) + assert packed.batch_size[0] == 0 + + +def test_rows_to_tq_fields_rejects_inconsistent_fields(): + """Rows with differing key sets must raise, not silently pack a ragged TensorDict.""" + pytest.importorskip('tensordict') + from twinkle.data_format import rows_to_tq_fields + + with pytest.raises(ValueError): + rows_to_tq_fields([{'input_ids': [1]}, {'input_ids': [2], 'labels': [3]}]) + + +def test_columns_to_tq_fields_mixes_numeric_and_non_numeric(): + """Numeric columns go through ``torch.tensor``; the rest through ``NonTensorStack``.""" + pytest.importorskip('tensordict') + import torch + + from twinkle.data_format import columns_to_tq_fields + + packed = columns_to_tq_fields({'scores': [1, 2], 'names': ['a', 'b']}, 2) + assert packed.batch_size[0] == 2 + assert isinstance(packed['scores'], torch.Tensor) + assert list(packed['names']) == ['a', 'b'] + + +def test_tq_fields_imports_without_tensordict(): + """The module must import cleanly without the async-rl extra installed. + + Subprocess with ``tensordict`` blocked from ``sys.modules``, asserting that importing + the module (and reading the constants) does not touch it -- function-level imports are + what make that true, so hoisting them would break this. + """ + code = ( + 'import sys;' + "sys.modules['tensordict'] = None;" + 'from twinkle.data_format import tq_fields, ROLLOUT_TRAIN_FIELDS;' + "assert 'input_ids' in ROLLOUT_TRAIN_FIELDS" + ) + subprocess.run([sys.executable, '-c', code], check=True) diff --git a/tests/infra/test_ray_get_timeout.py b/tests/infra/test_ray_get_timeout.py new file mode 100644 index 000000000..e82bfc44b --- /dev/null +++ b/tests/infra/test_ray_get_timeout.py @@ -0,0 +1,137 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Unit tests for the Sync_Dispatch_Path time bound. + +These exercise only ``twinkle.infra`` against a plain sleeping Ray actor. They +depend on neither GPU, Megatron, nor any ``src/twinkle/server/**`` component. +""" +from __future__ import annotations + +import pytest + +ray = pytest.importorskip('ray') + +import twinkle.infra as infra # noqa: E402 +from twinkle.infra import remote_function # noqa: E402 +from twinkle.infra._ray.ray_helper import RayHelper # noqa: E402 + + +@ray.remote +class _Sleeper: + """A plain Ray actor whose only method sleeps for a caller-supplied time.""" + + def slow(self, seconds: float): + import time + time.sleep(seconds) + return seconds + + def slow_batch(self, seconds: list[float]): + import time + time.sleep(seconds[0]) + return seconds + + def _twinkle_async_slow_batch(self, seconds: list[float]): + return self.slow_batch(seconds) + + +@pytest.fixture(scope='module', autouse=True) +def _ray_and_ray_mode(): + """Bring up Ray and put infra into 'ray' mode for the driver-side path.""" + ray.init(ignore_reinit_error=True, num_cpus=2, logging_level='ERROR') + prev_mode = infra._mode + infra._mode = 'ray' + try: + yield + finally: + infra._mode = prev_mode + + +def _make_driver(): + """A minimal stand-in for a remote_class handle: one actor, no concurrency.""" + driver = type('Driver', (), {})() + driver._actors = [_Sleeper.remote()] + driver._max_concurrency = None + return driver + + +def test_execute_all_sync_times_out(_ray_and_ray_mode): + """``execute_all_sync(timeout=)`` raises when the remote does not return in time.""" + actor = _Sleeper.remote() + workers_and_args = [(actor, [3.0], {})] + with pytest.raises(ray.exceptions.GetTimeoutError): + RayHelper.execute_all_sync('slow', workers_and_args, timeout=0.5) + + +def test_execute_all_sync_returns_within_timeout(_ray_and_ray_mode): + actor = _Sleeper.remote() + workers_and_args = [(actor, [0.1], {})] + assert RayHelper.execute_all_sync('slow', workers_and_args, timeout=10.0) == [0.1] + + +def test_decorator_timeout_takes_priority_over_instance(): + """A small decorator timeout wins over a large instance ``_ray_get_timeout``.""" + + def slow(self, seconds): # body runs in the worker, not here + return seconds + + wrapped = remote_function(dispatch='all', collect='first', sync=True, timeout=0.5)(slow) + driver = _make_driver() + driver._ray_get_timeout = 100.0 # would allow the call if it were consulted + with pytest.raises(ray.exceptions.GetTimeoutError): + wrapped(driver, 3.0) + + +def test_decorator_timeout_wins_when_larger_than_instance(): + """The decorator value wins even when it is the *larger* of the two. + + A large decorator timeout with a tiny instance value must NOT time out -- + proving the instance value is ignored when the decorator declares one. + """ + + def slow(self, seconds): + return seconds + + wrapped = remote_function(dispatch='all', collect='first', sync=True, timeout=100.0)(slow) + driver = _make_driver() + driver._ray_get_timeout = 0.3 # would time out if it were consulted + result = wrapped(driver, 1.0) + # sync collect may hand back a lazy-collect callable; resolving it must not time out. + assert (result() if callable(result) else result) == 1.0 + + +def test_instance_timeout_is_fallback_when_decorator_absent(): + """With no decorator timeout, the instance ``_ray_get_timeout`` applies.""" + + def slow(self, seconds): + return seconds + + wrapped = remote_function(dispatch='all', collect='first', sync=True)(slow) + driver = _make_driver() + driver._ray_get_timeout = 0.3 + with pytest.raises(ray.exceptions.GetTimeoutError): + wrapped(driver, 2.0) + + +def test_continuous_work_timeout_zero_is_not_treated_as_falsy(): + + def slow_batch(self, seconds): + return seconds + + wrapped = remote_function( + dispatch='all', collect='first', timeout=0, enable_continous_work=True)(slow_batch) + driver = _make_driver() + driver._ray_get_timeout = 100.0 + with pytest.raises(ray.exceptions.GetTimeoutError): + wrapped(driver, [1.0]) + + +def test_decorator_timeout_zero_is_not_treated_as_falsy(): + """Timeout=0 means 'time out immediately', not 'fall back to unbounded'.""" + + def slow(self, seconds): + return seconds + + wrapped = remote_function(dispatch='all', collect='first', sync=True, timeout=0)(slow) + driver = _make_driver() + driver._ray_get_timeout = 100.0 # the old ``or`` bug would fall back here + with pytest.raises(ray.exceptions.GetTimeoutError): + wrapped(driver, 1.0) diff --git a/tests/loss/test_grpo_gkd.py b/tests/loss/test_grpo_gkd.py index 3d0f55120..1f9daa707 100644 --- a/tests/loss/test_grpo_gkd.py +++ b/tests/loss/test_grpo_gkd.py @@ -82,6 +82,27 @@ def test_grpo_list_advantages(self): result = loss_fn(inputs, outputs, old_logps=old_logps, advantages=adv_list) assert torch.isfinite(result['loss']) + def test_pad_variable_length_full_sequence_rows(self): + """Unpadded per-sample rows align against a right-padded batch mask.""" + mask = torch.tensor([ + [False, True, True, False, False], + [False, False, True, True, True], + ]) + rows = [ + [-9.0, 0.1, 0.2], + [-9.0, -9.0, 0.3, 0.4, 0.5], + ] + + got = GRPOLoss()._pad_and_align_to_batch(rows, mask, mask.device, torch.float32) + + assert torch.equal(got[0], torch.tensor([0.0, 0.1, 0.2, 0.0, 0.0])) + assert torch.equal(got[1], torch.tensor([0.0, 0.0, 0.3, 0.4, 0.5])) + + def test_pad_rejects_full_sequence_missing_a_masked_position(self): + mask = torch.tensor([[False, False, True, True, True]]) + with pytest.raises(AssertionError, match='all masked positions'): + GRPOLoss()._pad_and_align_to_batch([[0.1, 0.2, 0.3, 0.4]], mask, mask.device, torch.float32) + def test_grpo_weights_sequences_equally(self): labels = torch.tensor([ [1, -100, -100], diff --git a/tests/model/test_micro_batch.py b/tests/model/test_micro_batch.py index 8d21c1dbb..ec4083e96 100644 --- a/tests/model/test_micro_batch.py +++ b/tests/model/test_micro_batch.py @@ -8,7 +8,6 @@ from twinkle.model.micro_batch import MicroBatchConfig, plan_micro_batches from twinkle.model.transformers.transformers import TransformersModel from twinkle.processor import InputProcessor -from twinkle.utils.nccl_safe import safe_loss @pytest.mark.parametrize('packing_algorithm', ['ffd', 'kk']) @@ -72,17 +71,6 @@ def test_sample_mean_and_token_sum_micro_batch_scales(): assert CrossEntropyLoss(reduction='sum').micro_batch_scale(inputs, [0]) == 1.0 -def test_safe_loss_preserves_wrapped_micro_batch_scale(): - inputs = [ - {'labels': [1, -100]}, - {'labels': [2, 3]}, - {'labels': [4, -100]}, - {'labels': [5, 6]}, - ] - - assert safe_loss(GRPOLoss()).micro_batch_scale(inputs, [0, 2]) == .5 - - def test_loss_without_micro_batch_semantics_fails_when_split(): with pytest.raises(NotImplementedError, match='does not support micro-batching'): Loss().micro_batch_scale([{}, {}], [0]) diff --git a/tests/model/test_multi_lora_dtype.py b/tests/model/test_multi_lora_dtype.py index 611008a40..e8d403e0c 100644 --- a/tests/model/test_multi_lora_dtype.py +++ b/tests/model/test_multi_lora_dtype.py @@ -17,6 +17,6 @@ def test_multi_lora_dtype_matches_bf16_base_before_fsdp_wrap(): assert {param.dtype for name, param in model.named_parameters() if 'lora_' in name} == {torch.float32} - TransformersModel._ensure_lora_dtype(None, model) + TransformersModel._ensure_lora_dtype(model) assert {param.dtype for name, param in model.named_parameters() if 'lora_' in name} == {torch.bfloat16} diff --git a/tests/model/test_multi_lora_target_parameters.py b/tests/model/test_multi_lora_target_parameters.py index b28ef6b36..7e23369a0 100644 --- a/tests/model/test_multi_lora_target_parameters.py +++ b/tests/model/test_multi_lora_target_parameters.py @@ -1,14 +1,15 @@ import copy +import pytest import sys -import types - import torch +import types from peft import LoraConfig, get_peft_model from peft.utils import set_peft_model_state_dict from torch import nn print(f"sys.path: {sys.path}") + class FakePackedExperts(nn.Module): def __init__(self, num_experts=2, hidden=4, intermediate=6, *, is_transposed=False): @@ -45,6 +46,42 @@ def forward(self, x, expert_idx=0): return self.mlp.experts(x, expert_idx=expert_idx) +class FakePeftModule: + + def __init__(self, peft_config, active_adapter): + self.peft_config = peft_config + self.active_adapter = active_adapter + + +def test_save_context_restores_peft_state_when_load_fails(): + from twinkle.model.multi_lora import LoraTenant, MultiLora + + original_config = {'lora_0': object(), 'lora_1': object()} + modules = [ + FakePeftModule(original_config, 'lora_1'), + FakePeftModule(original_config, 'lora_0'), + ] + multi_lora = MultiLora(max_loras=2, max_r=4) + multi_lora.module = modules + multi_lora.loras = [ + LoraTenant( + index=0, + adapter_name='lora_0', + config=_make_target_cfg(), + tenant_adapter_name='tenant', + tenant_config=_make_target_cfg(), + ) + ] + + with pytest.raises(RuntimeError, match='load failed'): + with multi_lora.save_context('tenant') as adapter_name: + assert adapter_name == 'lora_0' + raise RuntimeError('load failed') + + assert [module.peft_config for module in modules] == [original_config, original_config] + assert [module.active_adapter for module in modules] == ['lora_1', 'lora_0'] + + def test_peft_target_parameter_key_shapes_for_3d_experts(): model = FakeModel() cfg = LoraConfig( @@ -58,11 +95,13 @@ def test_peft_target_parameter_key_shapes_for_3d_experts(): state = peft_model.state_dict() lora_shapes = {key: tuple(state[key].shape) for key in state if "lora_" in key} + # PEFT 0.18.1 follows Linear's (out_features, in_features) convention for + # target parameters: A projects input->rank and B projects rank->output. assert lora_shapes == { - "base_model.model.mlp.experts.base_layer.lora_A.default.weight": (4, 12), - "base_model.model.mlp.experts.base_layer.lora_B.default.weight": (4, 4), - "base_model.model.mlp.experts.lora_A.default.weight": (4, 4), - "base_model.model.mlp.experts.lora_B.default.weight": (6, 4), + "base_model.model.mlp.experts.base_layer.lora_A.default.weight": (4, 4), + "base_model.model.mlp.experts.base_layer.lora_B.default.weight": (12, 4), + "base_model.model.mlp.experts.lora_A.default.weight": (4, 6), + "base_model.model.mlp.experts.lora_B.default.weight": (4, 4), } @@ -85,10 +124,7 @@ def test_target_parameter_multi_lora_updates_only_active_adapter(): manager.acquire("adapter_a", "lora_0", _make_target_cfg(r=2)) manager.acquire("adapter_b", "lora_1", _make_target_cfg(r=2)) - params_before = { - name: param.detach().clone() - for name, param in manager.named_slot_parameters("adapter_b") - } + params_before = {name: param.detach().clone() for name, param in manager.named_slot_parameters("adapter_b")} opt = torch.optim.SGD(manager.parameters_for_tenant("adapter_a"), lr=0.1) with manager.adapter("adapter_a"): @@ -116,8 +152,7 @@ def test_multilora_releases_target_parameter_slot_to_initial_weights(): initial_a = { name: param.detach().clone() - for name, param in multi_lora.target_parameter_manager.named_slot_parameters("adapter_a") - if ".lora_A." in name + for name, param in multi_lora.target_parameter_manager.named_slot_parameters("adapter_a") if ".lora_A." in name } with torch.no_grad(): @@ -133,7 +168,8 @@ def test_multilora_releases_target_parameter_slot_to_initial_weights(): else: assert torch.count_nonzero(param.detach()) == 0 -# Note: PEFT (Parameter-Efficient Fine-Tuning) does not natively support + +# Note: PEFT (Parameter-Efficient Fine-Tuning) does not natively support # installing multiple LoRA slots on target parameters. # def test_target_parameter_state_dict_loads_with_peft(): # from twinkle.model.multi_lora_target_parameters import TargetParameterLoraManager @@ -259,11 +295,3 @@ def test_multilora_transformers_installs_target_parameters_once(): pass else: raise AssertionError("different target_parameters should be rejected") - -# Run in the local environment. -if __name__ == "__main__": - assert test_peft_target_parameter_key_shapes_for_3d_experts() == True - assert test_target_parameter_multi_lora_updates_only_active_adapter() == True - assert test_multilora_releases_target_parameter_slot_to_initial_weights() == True - assert test_multilora_state_dict_round_trips_target_parameters() == True - assert test_multilora_transformers_installs_target_parameters_once() == True \ No newline at end of file diff --git a/tests/sampler/test_vllm_startup_lock.py b/tests/sampler/test_vllm_startup_lock.py index d07ad74a4..06e5dc47c 100644 --- a/tests/sampler/test_vllm_startup_lock.py +++ b/tests/sampler/test_vllm_startup_lock.py @@ -1,6 +1,5 @@ import multiprocessing import os - import pytest from twinkle.utils.parallel import PosixFileLock @@ -9,7 +8,7 @@ def _hold_startup_lock(lock_path: str, acquired, release) -> None: with PosixFileLock(lock_path): acquired.set() - if not release.wait(timeout=5): + if not release.wait(timeout=30): raise TimeoutError('test did not release vLLM startup lock') @@ -33,21 +32,23 @@ def test_vllm_engine_startup_is_serialized(tmp_path): try: first.start() - assert first_acquired.wait(timeout=5) + assert first_acquired.wait(timeout=30) second.start() - assert second_started.wait(timeout=5) + assert second_started.wait(timeout=30) assert not second_acquired.wait(timeout=0.2) release_first.set() - assert second_acquired.wait(timeout=5) + assert second_acquired.wait(timeout=30) finally: release_first.set() for process in (first, second): - process.join(timeout=5) + if process.pid is None: + continue + process.join(timeout=30) if process.is_alive(): process.terminate() - process.join(timeout=5) + process.join(timeout=30) assert first.exitcode == 0 assert second.exitcode == 0 diff --git a/tests/server/config/server_config_4b_e2e.yaml b/tests/server/config/server_config_4b_e2e.yaml index 6a9dfd317..f32eec8b4 100644 --- a/tests/server/config/server_config_4b_e2e.yaml +++ b/tests/server/config/server_config_4b_e2e.yaml @@ -36,6 +36,7 @@ applications: backend: transformers model_id: "ms://Qwen/Qwen3.5-4B" max_length: 10240 + max_loras: 10 nproc_per_node: 2 device_group: name: model diff --git a/tests/server/config/server_config_4b_e2e_megatron.yaml b/tests/server/config/server_config_4b_e2e_megatron.yaml index a6f193501..4d5e3b600 100644 --- a/tests/server/config/server_config_4b_e2e_megatron.yaml +++ b/tests/server/config/server_config_4b_e2e_megatron.yaml @@ -36,6 +36,7 @@ applications: backend: megatron model_id: "ms://Qwen/Qwen3.5-4B" max_length: 10240 + max_loras: 10 nproc_per_node: 4 device_group: name: model diff --git a/tests/server/config/test_server_config.py b/tests/server/config/test_server_config.py index bebd63c9f..b5bacf821 100644 --- a/tests/server/config/test_server_config.py +++ b/tests/server/config/test_server_config.py @@ -9,6 +9,8 @@ """ from __future__ import annotations +import re + import pytest import yaml from hypothesis import given, settings @@ -19,15 +21,12 @@ from twinkle.server.config import ApplicationSpec, ServerConfig from twinkle.server.exceptions import ConfigParseError from twinkle.server.launcher import ServerLauncher +from twinkle.server.sampler.app import SAMPLER_SELECTOR, build_sampler_app # ---------- minimal valid config strategy ---------------------------------- # _PERSISTENCE_VARIANTS = st.one_of( st.fixed_dictionaries({'mode': st.just('memory')}), - st.fixed_dictionaries({ - 'mode': st.just('file'), - 'file_path': st.just('/tmp/state.json') - }), st.fixed_dictionaries({ 'mode': st.just('redis'), 'redis_url': st.just('redis://localhost:6379/0') @@ -86,13 +85,6 @@ def test_redis_mode_missing_url() -> None: assert 'persistence.redis_url' in msg or 'redis_url' in msg -def test_file_mode_missing_path() -> None: - with pytest.raises(ValidationError) as exc: - ServerConfig.model_validate({'persistence': {'mode': 'file'}}) - msg = str(exc.value) - assert 'persistence.file_path' in msg or 'file_path' in msg - - @settings(max_examples=100) @given(bad_backend=st.text(min_size=1, max_size=8).filter(lambda s: s not in ('mock', 'transformers', 'megatron'))) def test_bad_backend_names_field(bad_backend: str) -> None: @@ -126,6 +118,26 @@ def test_nested_field_constraint_violation_named(bad_max_input_tokens: int) -> N assert any('max_input_tokens' in err['loc'] for err in errors) +def test_torch_sampler_is_rejected_during_config_validation() -> None: + with pytest.raises(ValidationError): + ApplicationSpec.model_validate({ + 'name': 'sampler', + 'import_path': 'sampler', + 'args': { + 'model_id': 'm', + 'device_group': {}, + 'device_mesh': {}, + 'sampler_type': 'torch', + }, + }) + + +def test_sampler_docstring_values_match_selector() -> None: + doc = build_sampler_app.__doc__ or '' + line = next(line for line in doc.splitlines() if 'sampler_type:' in line) + assert set(re.findall(r'``(\w+)``', line)) == set(SAMPLER_SELECTOR.builders) + + # ---------- round-trip fidelity ----------------------------------------- # @@ -257,8 +269,34 @@ def test_data_plane_application_uses_its_own_strict_args_schema() -> None: ApplicationSpec.model_validate({ 'name': 'data-plane', 'import_path': 'data_plane', - 'args': {'unknown': True}, + 'args': { + 'unknown': True + }, + }) + + +def test_processor_queue_config_is_rejected() -> None: + # A processor deployment has no task queue; queue_config was silently ignored + # before and now fails validation (F013 / P009) naming the offending field. + ApplicationSpec.model_validate({ + 'name': 'processor', + 'import_path': 'processor', + 'args': { + 'ncpu_proc_per_node': 1 + }, + }) + with pytest.raises(ValidationError) as exc: + ApplicationSpec.model_validate({ + 'name': 'processor', + 'import_path': 'processor', + 'args': { + 'ncpu_proc_per_node': 1, + 'queue_config': { + 'rps_limit': 4 + } + }, }) + assert 'queue_config' in str(exc.value) def test_cookbook_examples_load() -> None: diff --git a/tests/server/conftest.py b/tests/server/conftest.py index 4ac4d24fd..dfcc7f2e7 100644 --- a/tests/server/conftest.py +++ b/tests/server/conftest.py @@ -1,5 +1,5 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -"""Shared Ray runtime + per-test isolation for ``tests/server`` (state, cli, ...). +"""Shared Ray runtime, per-test isolation, and evidence boundaries. ``RayActorBackend`` is a forwarding wrapper around a detached Ray actor; instantiating one without an initialized Ray runtime raises @@ -12,6 +12,15 @@ of the actor wrapper. To keep tests independent we clear that actor's store before each test function. Tests that pin a non-default ``key_prefix`` get their own actor; this fixture intentionally leaves those alone. + +Evidence boundary: every mock-model backend method accepts +``**kwargs`` without argument validation, and the mock enters no real collective. +A mock-backed test therefore proves neither request/argument validation nor NCCL +behavior (asymmetric failure, collective mis-pairing, ReduceScatter, etc.). It may +prove only backend dispatch, task-queue behavior, timeout/admission mechanisms, +and event-loop responsiveness. Validation and NCCL claims require the GPU-gated +``test_nccl_safe_*_e2e.py`` tests against a real server. The contract suite covers +all five apps; Tinker compatibility follows the pinned 0.16.1 SDK wire values. """ from __future__ import annotations @@ -40,7 +49,7 @@ def _reset_canonical_state_actor(): """Clear the canonical state actor's store before each test function. Hypothesis property tests reuse the function scope across examples and - so should call ``backend.close()`` themselves to reset between examples. + so should call the actor's explicit ``flush_all()`` test hook themselves. """ import ray @@ -53,7 +62,7 @@ def _reset_canonical_state_actor(): actor = None if actor is not None: try: - ray.get(actor.close.remote()) + ray.get(actor.flush_all.remote()) except Exception: pass yield diff --git a/tests/server/contract/client_api_baseline.json b/tests/server/contract/client_api_baseline.json deleted file mode 100644 index 65f9db147..000000000 --- a/tests/server/contract/client_api_baseline.json +++ /dev/null @@ -1,1119 +0,0 @@ -{ - "data_plane": { - "paths": { - "/twinkle/append": { - "POST": { - "operationId": "append_twinkle_append_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/get": { - "POST": { - "operationId": "get_twinkle_get_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/put": { - "POST": { - "operationId": "put_twinkle_put_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/release": { - "POST": { - "operationId": "release_twinkle_release_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - } - } - }, - "gateway": { - "paths": { - "/asample": { - "POST": { - "operationId": "asample_asample_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/chat/completions": { - "POST": { - "operationId": "chat_completions_chat_completions_post", - "parameters": [], - "responses": [ - "200" - ] - } - }, - "/create_model": { - "POST": { - "operationId": "create_model_create_model_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/create_sampling_session": { - "POST": { - "operationId": "create_sampling_session_create_sampling_session_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/create_session": { - "POST": { - "operationId": "create_session_create_session_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/forward": { - "POST": { - "operationId": "forward_forward_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/forward_backward": { - "POST": { - "operationId": "forward_backward_forward_backward_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/get_info": { - "POST": { - "operationId": "get_info_get_info_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/get_server_capabilities": { - "GET": { - "operationId": "get_server_capabilities_get_server_capabilities_get", - "parameters": [], - "responses": [ - "200" - ] - } - }, - "/healthz": { - "GET": { - "operationId": "healthz_healthz_get", - "parameters": [], - "responses": [ - "200" - ] - } - }, - "/load_weights": { - "POST": { - "operationId": "load_weights_load_weights_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/models": { - "GET": { - "operationId": "list_models_models_get", - "parameters": [], - "responses": [ - "200" - ] - } - }, - "/optim_step": { - "POST": { - "operationId": "optim_step_optim_step_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/retrieve_future": { - "POST": { - "operationId": "retrieve_future_retrieve_future_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/save_weights": { - "POST": { - "operationId": "save_weights_save_weights_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/save_weights_for_sampler": { - "POST": { - "operationId": "save_weights_for_sampler_save_weights_for_sampler_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/session_heartbeat": { - "POST": { - "operationId": "session_heartbeat_session_heartbeat_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/telemetry": { - "POST": { - "operationId": "telemetry_telemetry_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/training_runs": { - "GET": { - "operationId": "get_training_runs_training_runs_get", - "parameters": [ - { - "in": "query", - "name": "limit", - "required": false, - "schema": { - "default": 20, - "title": "Limit", - "type": "integer" - } - }, - { - "in": "query", - "name": "offset", - "required": false, - "schema": { - "default": 0, - "title": "Offset", - "type": "integer" - } - } - ], - "responses": [ - "200", - "422" - ] - } - }, - "/training_runs/{run_id}": { - "GET": { - "operationId": "get_training_run_training_runs__run_id__get", - "parameters": [ - { - "in": "path", - "name": "run_id", - "required": true, - "schema": { - "title": "Run Id", - "type": "string" - } - } - ], - "responses": [ - "200", - "422" - ] - } - }, - "/training_runs/{run_id}/checkpoints": { - "GET": { - "operationId": "get_run_checkpoints_training_runs__run_id__checkpoints_get", - "parameters": [ - { - "in": "path", - "name": "run_id", - "required": true, - "schema": { - "title": "Run Id", - "type": "string" - } - } - ], - "responses": [ - "200", - "422" - ] - } - }, - "/training_runs/{run_id}/checkpoints/{checkpoint_id}": { - "DELETE": { - "operationId": "delete_run_checkpoint_training_runs__run_id__checkpoints__checkpoint_id__delete", - "parameters": [ - { - "in": "path", - "name": "run_id", - "required": true, - "schema": { - "title": "Run Id", - "type": "string" - } - }, - { - "in": "path", - "name": "checkpoint_id", - "required": true, - "schema": { - "title": "Checkpoint Id", - "type": "string" - } - } - ], - "responses": [ - "200", - "422" - ] - } - }, - "/training_runs/{run_id}/checkpoints/{checkpoint_id}/publish": { - "POST": { - "operationId": "publish_checkpoint_training_runs__run_id__checkpoints__checkpoint_id__publish_post", - "parameters": [ - { - "in": "path", - "name": "run_id", - "required": true, - "schema": { - "title": "Run Id", - "type": "string" - } - }, - { - "in": "path", - "name": "checkpoint_id", - "required": true, - "schema": { - "title": "Checkpoint Id", - "type": "string" - } - } - ], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/capacity_info": { - "GET": { - "operationId": "get_capacity_info_twinkle_capacity_info_get", - "parameters": [], - "responses": [ - "200" - ] - } - }, - "/twinkle/checkpoint_path/{run_id}/{checkpoint_id}": { - "GET": { - "operationId": "get_checkpoint_path_twinkle_checkpoint_path__run_id___checkpoint_id__get", - "parameters": [ - { - "in": "path", - "name": "run_id", - "required": true, - "schema": { - "title": "Run Id", - "type": "string" - } - }, - { - "in": "path", - "name": "checkpoint_id", - "required": true, - "schema": { - "title": "Checkpoint Id", - "type": "string" - } - } - ], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/create_session": { - "POST": { - "operationId": "create_session_twinkle_create_session_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/get_server_capabilities": { - "GET": { - "operationId": "get_server_capabilities_twinkle_get_server_capabilities_get", - "parameters": [], - "responses": [ - "200" - ] - } - }, - "/twinkle/healthz": { - "GET": { - "operationId": "healthz_twinkle_healthz_get", - "parameters": [], - "responses": [ - "200" - ] - } - }, - "/twinkle/healthz/deep": { - "GET": { - "operationId": "healthz_deep_twinkle_healthz_deep_get", - "parameters": [], - "responses": [ - "200" - ] - } - }, - "/twinkle/session_heartbeat": { - "POST": { - "operationId": "session_heartbeat_twinkle_session_heartbeat_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/status": { - "GET": { - "operationId": "status_twinkle_status_get", - "parameters": [], - "responses": [ - "200" - ] - } - }, - "/twinkle/training_runs": { - "GET": { - "operationId": "get_training_runs_twinkle_training_runs_get", - "parameters": [ - { - "in": "query", - "name": "limit", - "required": false, - "schema": { - "default": 20, - "title": "Limit", - "type": "integer" - } - }, - { - "in": "query", - "name": "offset", - "required": false, - "schema": { - "default": 0, - "title": "Offset", - "type": "integer" - } - } - ], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/training_runs/{run_id}": { - "GET": { - "operationId": "get_training_run_twinkle_training_runs__run_id__get", - "parameters": [ - { - "in": "path", - "name": "run_id", - "required": true, - "schema": { - "title": "Run Id", - "type": "string" - } - } - ], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/training_runs/{run_id}/checkpoints": { - "GET": { - "operationId": "get_run_checkpoints_twinkle_training_runs__run_id__checkpoints_get", - "parameters": [ - { - "in": "path", - "name": "run_id", - "required": true, - "schema": { - "title": "Run Id", - "type": "string" - } - } - ], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/training_runs/{run_id}/checkpoints/{checkpoint_id}": { - "DELETE": { - "operationId": "delete_run_checkpoint_twinkle_training_runs__run_id__checkpoints__checkpoint_id__delete", - "parameters": [ - { - "in": "path", - "name": "run_id", - "required": true, - "schema": { - "title": "Run Id", - "type": "string" - } - }, - { - "in": "path", - "name": "checkpoint_id", - "required": true, - "schema": { - "title": "Checkpoint Id", - "type": "string" - } - } - ], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/weights_info": { - "POST": { - "operationId": "weights_info_twinkle_weights_info_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/unload_model": { - "POST": { - "operationId": "unload_model_unload_model_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/weights_info": { - "POST": { - "operationId": "weights_info_weights_info_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - } - } - }, - "model": { - "paths": { - "/healthz": { - "GET": { - "operationId": "healthz_healthz_get", - "parameters": [], - "responses": [ - "200" - ] - } - }, - "/tinker/create_model": { - "POST": { - "operationId": "create_model_tinker_create_model_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/tinker/forward": { - "POST": { - "operationId": "forward_tinker_forward_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/tinker/forward_backward": { - "POST": { - "operationId": "forward_backward_tinker_forward_backward_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/tinker/get_info": { - "POST": { - "operationId": "get_info_tinker_get_info_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/tinker/load_weights": { - "POST": { - "operationId": "load_weights_tinker_load_weights_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/tinker/optim_step": { - "POST": { - "operationId": "optim_step_tinker_optim_step_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/tinker/save_weights": { - "POST": { - "operationId": "save_weights_tinker_save_weights_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/tinker/save_weights_for_sampler": { - "POST": { - "operationId": "save_weights_for_sampler_tinker_save_weights_for_sampler_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/tinker/unload_model": { - "POST": { - "operationId": "unload_model_tinker_unload_model_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/add_adapter_to_model": { - "POST": { - "operationId": "add_adapter_to_model_twinkle_add_adapter_to_model_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/add_metric": { - "POST": { - "operationId": "add_metric_twinkle_add_metric_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/apply_patch": { - "POST": { - "operationId": "apply_patch_twinkle_apply_patch_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/backward": { - "POST": { - "operationId": "backward_twinkle_backward_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/calculate_loss": { - "POST": { - "operationId": "calculate_loss_twinkle_calculate_loss_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/calculate_metric": { - "POST": { - "operationId": "calculate_metric_twinkle_calculate_metric_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/clip_grad_and_step": { - "POST": { - "operationId": "clip_grad_and_step_twinkle_clip_grad_and_step_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/clip_grad_norm": { - "POST": { - "operationId": "clip_grad_norm_twinkle_clip_grad_norm_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/create": { - "POST": { - "operationId": "create_twinkle_create_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/forward": { - "POST": { - "operationId": "forward_twinkle_forward_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/forward_from_data_plane": { - "POST": { - "operationId": "forward_from_data_plane_twinkle_forward_from_data_plane_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/forward_backward": { - "POST": { - "operationId": "forward_backward_twinkle_forward_backward_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/forward_backward_from_data_plane": { - "POST": { - "operationId": "forward_backward_from_data_plane_twinkle_forward_backward_from_data_plane_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/forward_only": { - "POST": { - "operationId": "forward_only_twinkle_forward_only_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/forward_only_from_data_plane": { - "POST": { - "operationId": "forward_only_from_data_plane_twinkle_forward_only_from_data_plane_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/get_state_dict": { - "POST": { - "operationId": "get_state_dict_twinkle_get_state_dict_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/get_train_configs": { - "POST": { - "operationId": "get_train_configs_twinkle_get_train_configs_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/load": { - "POST": { - "operationId": "load_twinkle_load_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/lr_step": { - "POST": { - "operationId": "lr_step_twinkle_lr_step_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/remove_adapter": { - "POST": { - "operationId": "remove_adapter_twinkle_remove_adapter_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/resume_from_checkpoint": { - "POST": { - "operationId": "resume_from_checkpoint_twinkle_resume_from_checkpoint_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/save": { - "POST": { - "operationId": "save_twinkle_save_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/set_loss": { - "POST": { - "operationId": "set_loss_twinkle_set_loss_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/set_lr_scheduler": { - "POST": { - "operationId": "set_lr_scheduler_twinkle_set_lr_scheduler_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/set_optimizer": { - "POST": { - "operationId": "set_optimizer_twinkle_set_optimizer_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/set_processor": { - "POST": { - "operationId": "set_processor_twinkle_set_processor_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/set_template": { - "POST": { - "operationId": "set_template_twinkle_set_template_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/step": { - "POST": { - "operationId": "step_twinkle_step_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/upload_status/{request_id}": { - "GET": { - "operationId": "upload_status_twinkle_upload_status__request_id__get", - "parameters": [ - { - "in": "path", - "name": "request_id", - "required": true, - "schema": { - "title": "Request Id", - "type": "string" - } - } - ], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/upload_to_hub": { - "POST": { - "operationId": "upload_to_hub_twinkle_upload_to_hub_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/zero_grad": { - "POST": { - "operationId": "zero_grad_twinkle_zero_grad_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - } - } - }, - "processor": { - "paths": { - "/twinkle/call": { - "POST": { - "operationId": "call_twinkle_call_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/create": { - "POST": { - "operationId": "create_twinkle_create_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - } - } - }, - "sampler": { - "paths": { - "/tinker/asample": { - "POST": { - "operationId": "asample_tinker_asample_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/add_adapter_to_sampler": { - "POST": { - "operationId": "add_adapter_to_sampler_twinkle_add_adapter_to_sampler_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/apply_patch": { - "POST": { - "operationId": "apply_patch_twinkle_apply_patch_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/create": { - "POST": { - "operationId": "create_twinkle_create_post", - "parameters": [], - "responses": [ - "200" - ] - } - }, - "/twinkle/sample": { - "POST": { - "operationId": "sample_twinkle_sample_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/sample_to_data_plane": { - "POST": { - "operationId": "sample_to_data_plane_twinkle_sample_to_data_plane_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/sample_stream": { - "POST": { - "operationId": "sample_stream_twinkle_sample_stream_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/set_template": { - "POST": { - "operationId": "set_template_twinkle_set_template_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - }, - "/twinkle/unload_adapter_paths": { - "POST": { - "operationId": "unload_adapter_paths_twinkle_unload_adapter_paths_post", - "parameters": [], - "responses": [ - "200", - "422" - ] - } - } - } - } -} diff --git a/tests/server/contract/client_api_harness.py b/tests/server/contract/client_api_harness.py index 0496a1b07..ab82e0582 100644 --- a/tests/server/contract/client_api_harness.py +++ b/tests/server/contract/client_api_harness.py @@ -2,10 +2,10 @@ """ Client-API contract harness. -Builds the four FastAPI apps used by the Ray Serve deployments (Gateway, Model, -Sampler, Processor) by registering their route-registration helpers against a -fresh FastAPI instance, then extracts the client-facing surface (route paths, -HTTP methods, and request/response schemas) as a stable JSON dict. +Builds the five FastAPI apps used by the Ray Serve deployments (Data Plane, +Gateway, Model, Sampler, Processor) by registering their route-registration helpers against a +fresh FastAPI instance, then extracts route paths, methods, parameters, and +recursive request/response type shapes as a stable JSON dict. Used to: - snapshot the current surface into ``client_api_baseline.json`` before the @@ -25,12 +25,18 @@ """ from __future__ import annotations +import dataclasses import json -from collections.abc import Callable +import re +import sys +import types as pytypes +from collections.abc import Callable, Mapping, Sequence +from enum import Enum from fastapi import FastAPI -from fastapi.openapi.utils import get_openapi +from fastapi.routing import APIRoute from pathlib import Path -from typing import Any +from pydantic import BaseModel +from typing import Annotated, Any, Literal, Union, get_args, get_origin, get_type_hints # ----- App build helpers --------------------------------------------------- # @@ -39,25 +45,33 @@ def _noop_self() -> None: return None +def build_data_plane_app() -> FastAPI: + from twinkle.server.data_plane.handlers import register_data_plane_routes + + app = FastAPI() + register_data_plane_routes(app, _noop_self) + return app + + def build_gateway_app() -> FastAPI: from twinkle.server.gateway.openai_handlers import _register_openai_routes - from twinkle.server.gateway.tinker_handlers import _register_tinker_routes - from twinkle.server.gateway.twinkle_handlers import _register_twinkle_routes + from twinkle.server.gateway.tinker_handlers import _register_gateway_tinker_routes + from twinkle.server.gateway.twinkle_handlers import _register_gateway_twinkle_routes app = FastAPI() - _register_tinker_routes(app, _noop_self) - _register_twinkle_routes(app, _noop_self) + _register_gateway_tinker_routes(app, _noop_self) + _register_gateway_twinkle_routes(app, _noop_self) _register_openai_routes(app, _noop_self) return app def build_model_app() -> FastAPI: - from twinkle.server.model.tinker_handlers import _register_tinker_routes - from twinkle.server.model.twinkle_handlers import _register_twinkle_routes + from twinkle.server.model.tinker_handlers import _register_model_tinker_routes + from twinkle.server.model.twinkle_handlers import _register_model_twinkle_routes app = FastAPI() - _register_tinker_routes(app, _noop_self) - _register_twinkle_routes(app, _noop_self) + _register_model_tinker_routes(app, _noop_self) + _register_model_twinkle_routes(app, _noop_self) return app @@ -80,6 +94,7 @@ def build_processor_app() -> FastAPI: APP_BUILDERS: dict[str, Callable[[], FastAPI]] = { + 'data_plane': build_data_plane_app, 'gateway': build_gateway_app, 'model': build_model_app, 'sampler': build_sampler_app, @@ -91,41 +106,96 @@ def build_processor_app() -> FastAPI: _HTTP_METHODS = {'GET', 'POST', 'PUT', 'PATCH', 'DELETE'} -def _extract_app_surface(app: FastAPI) -> dict[str, Any]: - """Return a SLIM client-contract view of ``app``'s OpenAPI surface. - - Snapshots, per path and HTTP method, only the stable client-facing contract: - the ``operationId``, the ``parameters``, and the set of response status - codes. The full ``components.schemas`` body and per-operation ``requestBody`` - schema are intentionally NOT snapshotted — they churn on Pydantic / FastAPI - version bumps without representing a real client-contract change. Route - paths, HTTP methods, and response status codes remain frozen. - """ - spec = get_openapi( - title='contract', - version='0.0.0', - routes=app.routes, - ) +def _type_contract(annotation: Any, seen: frozenset[str] = frozenset()) -> Any: + """Build a stable field-level schema for Pydantic models and SDK dataclasses.""" + if annotation is None or annotation is type(None): + return {'type': 'null'} + if annotation is Any: + return {} + origin = get_origin(annotation) + args = get_args(annotation) + if origin is Annotated: + return _type_contract(args[0], seen) + if origin in (Union, pytypes.UnionType): + return {'anyOf': [_type_contract(arg, seen) for arg in args]} + if origin in (list, set, tuple, Sequence): + return {'type': 'array', 'items': _type_contract(args[0], seen) if args else {}} + if origin in (dict, Mapping): + return {'type': 'object', 'additionalProperties': _type_contract(args[1], seen) if len(args) > 1 else {}} + if origin is Literal: + return {'enum': list(args)} + if isinstance(annotation, type) and issubclass(annotation, Enum): + return {'enum': [item.value for item in annotation]} + if isinstance(annotation, type) and issubclass(annotation, BaseModel): + return annotation.model_json_schema() + if isinstance(annotation, type) and dataclasses.is_dataclass(annotation): + name = f'{annotation.__module__}.{annotation.__qualname__}' + if name in seen: + return {'$ref': name} + module = sys.modules.get(annotation.__module__) + try: + hints = get_type_hints(annotation, globalns=vars(module) if module else None) + except (NameError, TypeError): + hints = annotation.__annotations__ + properties = {} + required = [] + for field in dataclasses.fields(annotation): + if not field.init or field.name.startswith('_'): + continue + properties[field.name] = _type_contract(hints.get(field.name, Any), seen | {name}) + if field.default is dataclasses.MISSING and field.default_factory is dataclasses.MISSING: + required.append(field.name) + result = {'type': 'object', 'properties': properties} + if required: + result['required'] = required + return result + primitive = {str: 'string', int: 'integer', float: 'number', bool: 'boolean'} + if annotation in primitive: + return {'type': primitive[annotation]} + return {'pythonType': getattr(annotation, '__qualname__', repr(annotation))} + + +def _parameter_contract(field: Any) -> dict[str, Any]: + field_info = field.field_info + return { + 'name': field.alias, + 'required': bool(field_info.is_required()), + 'schema': _type_contract(field_info.annotation), + } + +def _extract_app_surface(app: FastAPI) -> dict[str, Any]: + """Return every route's complete request and response type shape.""" paths: dict[str, dict[str, Any]] = {} - for path, ops in (spec.get('paths') or {}).items(): - clean_ops: dict[str, Any] = {} - for method, op in ops.items(): - if method.upper() not in _HTTP_METHODS: - continue - clean_ops[method.upper()] = { - 'operationId': op.get('operationId'), - 'parameters': op.get('parameters', []), - 'responses': sorted((op.get('responses') or {}).keys()), + for route in app.routes: + if not isinstance(route, APIRoute): + continue + extra_responses = {} + for status, response in route.responses.items(): + extra_responses[str(status)] = { + 'description': response.get('description'), + 'content': response.get('content'), + 'model': _type_contract(response.get('model')) if response.get('model') else None, } - if clean_ops: - paths[path] = clean_ops - + operation = { + 'operationId': route.operation_id or route.name, + 'body': [_parameter_contract(field) for field in route.dependant.body_params], + 'path': [_parameter_contract(field) for field in route.dependant.path_params], + 'query': [_parameter_contract(field) for field in route.dependant.query_params], + 'headers': [_parameter_contract(field) for field in route.dependant.header_params], + 'cookies': [_parameter_contract(field) for field in route.dependant.cookie_params], + 'response': _type_contract(route.response_model), + 'responses': extra_responses, + 'statusCode': route.status_code or 200, + } + for method in sorted(route.methods & _HTTP_METHODS): + client_path = _client_path(route.path) + paths.setdefault(client_path, {})[method] = operation return {'paths': paths} def extract_full_surface() -> dict[str, Any]: - """Build all four apps and return a per-app contract surface dict.""" + """Build all five apps and return a per-app contract surface dict.""" surface: dict[str, Any] = {} for name, builder in APP_BUILDERS.items(): app = builder() @@ -133,10 +203,65 @@ def extract_full_surface() -> dict[str, Any]: return surface -# ----- Baseline I/O -------------------------------------------------------- # +def _client_path(route_path: str) -> str: + """Strip FastAPI path-converter suffixes so ``{id:path}`` reads as ``{id}``.""" + return re.sub(r'{([^}:]+):[^}]+}', r'{\1}', route_path) + + +def _model_name(annotation: Any) -> str | None: + """A stable, human-readable name for a request/response model annotation.""" + if annotation is None: + return None + return getattr(annotation, '__qualname__', None) or repr(annotation) + +def extract_route_inventory() -> dict[str, dict[str, Any]]: + """A compact, reviewable projection of the wire surface. + + One entry per route -- ``" " -> {response, body, statusCode}`` -- + naming the model classes instead of inlining their field schemas. + + This is the projection that gets committed. A route appearing, disappearing, or + changing its response/body model shows up as a few readable lines in a PR diff, + whereas the full field-level surface is ~8k lines and nobody reads that diff. + The trade-off is explicit: this catches route-level and model-level changes, not + field-level drift inside a model. + """ + inventory: dict[str, dict[str, Any]] = {} + for name, builder in APP_BUILDERS.items(): + app = builder() + routes: dict[str, Any] = {} + for route in app.routes: + if not isinstance(route, APIRoute): + continue + client_path = _client_path(route.path) + body = [_model_name(field.field_info.annotation) for field in route.dependant.body_params] + for method in sorted(route.methods & _HTTP_METHODS): + routes[f'{method} {client_path}'] = { + 'response': _model_name(route.response_model), + 'body': body, + 'statusCode': route.status_code or 200, + } + inventory[name] = routes + return inventory + + +# ----- Snapshot I/O -------------------------------------------------------- # + +# The full field-level surface (:func:`extract_full_surface`). A GENERATED artifact, +# deliberately NOT committed: an 8k-line diff on every intentional wire change is noise +# nobody reads. Being regenerated from the code under test, it cannot by itself detect an +# unintended change -- ``ROUTES_PATH`` is the guard that can. Keep that asymmetry in mind +# before treating a green baseline test as evidence of anything. BASELINE_PATH = Path(__file__).parent / 'client_api_baseline.json' +# The compact route inventory (:func:`extract_route_inventory`). COMMITTED to git: this +# is the actual regression guard, so it has to stay tracked for the guard to mean +# anything. +ROUTES_PATH = Path(__file__).parent / 'client_api_routes.json' + +_REGEN_HINT = 'Regenerate with: python -m tests.server.contract.update_baseline' + def write_baseline(path: Path | None = None) -> Path: """Snapshot the current client-API surface to ``client_api_baseline.json``.""" @@ -146,6 +271,28 @@ def write_baseline(path: Path | None = None) -> Path: return p +def write_route_inventory(path: Path | None = None) -> Path: + """Snapshot the compact route inventory to ``client_api_routes.json``.""" + p = Path(path) if path is not None else ROUTES_PATH + p.write_text(json.dumps(extract_route_inventory(), indent=2, sort_keys=True) + '\n') + return p + + def load_baseline(path: Path | None = None) -> dict[str, Any]: + """Load the generated full surface, failing with a fix hint rather than a bare OSError.""" p = Path(path) if path is not None else BASELINE_PATH + if not p.is_file(): + raise FileNotFoundError(f'Contract baseline missing: {p}\n' + f'It is a generated artifact and is deliberately not committed. ' + f'{_REGEN_HINT}') + return json.loads(p.read_text()) + + +def load_route_inventory(path: Path | None = None) -> dict[str, Any]: + """Load the committed route inventory, failing loudly if it went missing.""" + p = Path(path) if path is not None else ROUTES_PATH + if not p.is_file(): + raise FileNotFoundError(f'Committed route inventory missing: {p}\n' + f'This file IS tracked by git -- restore it instead of regenerating ' + f'blindly, or the guard silently becomes a tautology. {_REGEN_HINT}') return json.loads(p.read_text()) diff --git a/tests/server/contract/client_api_routes.json b/tests/server/contract/client_api_routes.json new file mode 100644 index 000000000..ec8f176e7 --- /dev/null +++ b/tests/server/contract/client_api_routes.json @@ -0,0 +1,628 @@ +{ + "data_plane": { + "POST /twinkle/append": { + "body": [ + "DataAppendRequest" + ], + "response": "DataRef", + "statusCode": 200 + }, + "POST /twinkle/get": { + "body": [ + "DataGetRequest" + ], + "response": "DataRowsResponse", + "statusCode": 200 + }, + "POST /twinkle/put": { + "body": [ + "DataPutRequest" + ], + "response": "DataRef", + "statusCode": 200 + }, + "POST /twinkle/release": { + "body": [ + "DataReleaseRequest" + ], + "response": "dict", + "statusCode": 200 + } + }, + "gateway": { + "DELETE /training_runs/{run_id}/checkpoints/{checkpoint_id}": { + "body": [], + "response": "Any", + "statusCode": 200 + }, + "DELETE /twinkle/training_runs/{run_id}/checkpoints/{checkpoint_id}": { + "body": [], + "response": "DeleteCheckpointResponse", + "statusCode": 200 + }, + "GET /get_server_capabilities": { + "body": [], + "response": "GetServerCapabilitiesResponse", + "statusCode": 200 + }, + "GET /healthz": { + "body": [], + "response": "HealthResponse", + "statusCode": 200 + }, + "GET /models": { + "body": [], + "response": null, + "statusCode": 200 + }, + "GET /training_runs": { + "body": [], + "response": "TrainingRunsResponse", + "statusCode": 200 + }, + "GET /training_runs/{run_id}": { + "body": [], + "response": "TrainingRun", + "statusCode": 200 + }, + "GET /training_runs/{run_id}/checkpoints": { + "body": [], + "response": "CheckpointsListResponse", + "statusCode": 200 + }, + "GET /twinkle/capacity_info": { + "body": [], + "response": "CapacityInfoResponse", + "statusCode": 200 + }, + "GET /twinkle/checkpoint_path/{run_id}/{checkpoint_id}": { + "body": [], + "response": "CheckpointPathResponse", + "statusCode": 200 + }, + "GET /twinkle/get_server_capabilities": { + "body": [], + "response": "GetServerCapabilitiesResponse", + "statusCode": 200 + }, + "GET /twinkle/healthz": { + "body": [], + "response": "HealthResponse", + "statusCode": 200 + }, + "GET /twinkle/healthz/deep": { + "body": [], + "response": "dict", + "statusCode": 200 + }, + "GET /twinkle/status": { + "body": [], + "response": "dict", + "statusCode": 200 + }, + "GET /twinkle/training_runs": { + "body": [], + "response": "TrainingRunsResponse", + "statusCode": 200 + }, + "GET /twinkle/training_runs/{run_id}": { + "body": [], + "response": "TrainingRun", + "statusCode": 200 + }, + "GET /twinkle/training_runs/{run_id}/checkpoints": { + "body": [], + "response": "CheckpointsListResponse", + "statusCode": 200 + }, + "POST /asample": { + "body": [ + "SampleRequest" + ], + "response": "Any", + "statusCode": 200 + }, + "POST /chat/completions": { + "body": [], + "response": null, + "statusCode": 200 + }, + "POST /create_model": { + "body": [ + "CreateModelRequest" + ], + "response": "Any", + "statusCode": 200 + }, + "POST /create_sampling_session": { + "body": [ + "CreateSamplingSessionRequest" + ], + "response": "CreateSamplingSessionResponse", + "statusCode": 200 + }, + "POST /create_session": { + "body": [ + "CreateSessionRequest" + ], + "response": "CreateSessionResponse", + "statusCode": 200 + }, + "POST /forward": { + "body": [ + "ForwardRequest" + ], + "response": "Any", + "statusCode": 200 + }, + "POST /forward_backward": { + "body": [ + "ForwardBackwardRequest" + ], + "response": "Any", + "statusCode": 200 + }, + "POST /get_info": { + "body": [ + "GetInfoRequest" + ], + "response": "Any", + "statusCode": 200 + }, + "POST /load_weights": { + "body": [ + "LoadWeightsRequest" + ], + "response": "Any", + "statusCode": 200 + }, + "POST /optim_step": { + "body": [ + "OptimStepRequest" + ], + "response": "Any", + "statusCode": 200 + }, + "POST /retrieve_future": { + "body": [ + "FutureRetrieveRequest" + ], + "response": "Any", + "statusCode": 200 + }, + "POST /save_weights": { + "body": [ + "SaveWeightsRequest" + ], + "response": "Any", + "statusCode": 200 + }, + "POST /save_weights_for_sampler": { + "body": [ + "SaveWeightsForSamplerRequest" + ], + "response": "Any", + "statusCode": 200 + }, + "POST /session_heartbeat": { + "body": [ + "SessionHeartbeatRequest" + ], + "response": "SessionHeartbeatResponse", + "statusCode": 200 + }, + "POST /telemetry": { + "body": [ + "TelemetrySendRequest" + ], + "response": "TelemetryResponse", + "statusCode": 200 + }, + "POST /training_runs/{run_id}/checkpoints/{checkpoint_id}/publish": { + "body": [], + "response": null, + "statusCode": 200 + }, + "POST /twinkle/cancel": { + "body": [ + "CancelRequest" + ], + "response": "CancelResponse", + "statusCode": 200 + }, + "POST /twinkle/create_session": { + "body": [ + "CreateSessionRequest" + ], + "response": "CreateSessionResponse", + "statusCode": 200 + }, + "POST /twinkle/retrieve_future": { + "body": [ + "RetrieveFutureRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/session_heartbeat": { + "body": [ + "SessionHeartbeatRequest" + ], + "response": "SessionHeartbeatResponse", + "statusCode": 200 + }, + "POST /twinkle/weights_info": { + "body": [ + "WeightsInfoRequest" + ], + "response": "WeightsInfoResponse", + "statusCode": 200 + }, + "POST /unload_model": { + "body": [ + "UnloadModelRequest" + ], + "response": "Any", + "statusCode": 200 + }, + "POST /weights_info": { + "body": [ + "dict" + ], + "response": "WeightsInfoResponse", + "statusCode": 200 + } + }, + "model": { + "GET /healthz": { + "body": [], + "response": "dict", + "statusCode": 200 + }, + "POST /tinker/create_model": { + "body": [ + "CreateModelRequest" + ], + "response": "UntypedAPIFuture", + "statusCode": 200 + }, + "POST /tinker/forward": { + "body": [ + "ForwardRequest" + ], + "response": "UntypedAPIFuture", + "statusCode": 200 + }, + "POST /tinker/forward_backward": { + "body": [ + "ForwardBackwardRequest" + ], + "response": "UntypedAPIFuture", + "statusCode": 200 + }, + "POST /tinker/get_info": { + "body": [ + "GetInfoRequest" + ], + "response": "GetInfoResponse", + "statusCode": 200 + }, + "POST /tinker/load_weights": { + "body": [ + "LoadWeightsRequest" + ], + "response": "UntypedAPIFuture", + "statusCode": 200 + }, + "POST /tinker/optim_step": { + "body": [ + "OptimStepRequest" + ], + "response": "UntypedAPIFuture", + "statusCode": 200 + }, + "POST /tinker/save_weights": { + "body": [ + "SaveWeightsRequest" + ], + "response": "UntypedAPIFuture", + "statusCode": 200 + }, + "POST /tinker/save_weights_for_sampler": { + "body": [ + "SaveWeightsForSamplerRequest" + ], + "response": "UntypedAPIFuture", + "statusCode": 200 + }, + "POST /tinker/unload_model": { + "body": [ + "UnloadModelRequest" + ], + "response": "UntypedAPIFuture", + "statusCode": 200 + }, + "POST /twinkle/add_adapter_to_model": { + "body": [ + "AddAdapterRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/add_metric": { + "body": [ + "AddMetricRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/apply_patch": { + "body": [ + "ApplyPatchRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/backward": { + "body": [ + "AdapterRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/calculate_loss": { + "body": [ + "AdapterRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/calculate_metric": { + "body": [ + "CalculateMetricRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/clip_grad_and_step": { + "body": [ + "ClipGradAndStepRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/clip_grad_norm": { + "body": [ + "ClipGradNormRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/create": { + "body": [ + "CreateRequest" + ], + "response": "CreateResponse", + "statusCode": 200 + }, + "POST /twinkle/forward": { + "body": [ + "ForwardRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/forward_backward": { + "body": [ + "ForwardBackwardTaskRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/forward_backward_from_data_plane": { + "body": [ + "DataPlaneForwardRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/forward_from_data_plane": { + "body": [ + "DataPlaneForwardRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/forward_only": { + "body": [ + "ForwardOnlyRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/forward_only_from_data_plane": { + "body": [ + "DataPlaneForwardOnlyRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/get_train_configs": { + "body": [ + "AdapterRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/load": { + "body": [ + "LoadRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/lr_step": { + "body": [ + "LrStepRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/remove_adapter": { + "body": [ + "AdapterRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/resume_from_checkpoint": { + "body": [ + "ResumeFromCheckpointRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/save": { + "body": [ + "SaveRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/set_loss": { + "body": [ + "SetLossRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/set_lr_scheduler": { + "body": [ + "SetLrSchedulerRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/set_optimizer": { + "body": [ + "SetOptimizerRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/set_processor": { + "body": [ + "SetProcessorRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/set_template": { + "body": [ + "SetTemplateRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/step": { + "body": [ + "StepRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/upload_to_hub": { + "body": [ + "UploadToHubRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/zero_grad": { + "body": [ + "AdapterRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + } + }, + "processor": { + "POST /twinkle/call": { + "body": [ + "ProcessorCallRequest" + ], + "response": "ProcessorCallResponse", + "statusCode": 200 + }, + "POST /twinkle/create": { + "body": [ + "ProcessorCreateRequest" + ], + "response": "ProcessorCreateResponse", + "statusCode": 200 + } + }, + "sampler": { + "POST /tinker/asample": { + "body": [ + "SampleRequest" + ], + "response": "UntypedAPIFuture", + "statusCode": 200 + }, + "POST /twinkle/add_adapter_to_sampler": { + "body": [ + "SamplerAddAdapterRequest" + ], + "response": "SamplerAddAdapterResponse", + "statusCode": 200 + }, + "POST /twinkle/apply_patch": { + "body": [ + "ApplyPatchRequest" + ], + "response": null, + "statusCode": 200 + }, + "POST /twinkle/create": { + "body": [], + "response": "SamplerCreateResponse", + "statusCode": 200 + }, + "POST /twinkle/sample": { + "body": [ + "SampleRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/sample_stream": { + "body": [ + "SampleRequest" + ], + "response": null, + "statusCode": 200 + }, + "POST /twinkle/sample_to_data_plane": { + "body": [ + "DataPlaneSampleRequest" + ], + "response": "TaskEnvelope", + "statusCode": 200 + }, + "POST /twinkle/set_template": { + "body": [ + "SamplerSetTemplateRequest" + ], + "response": "SamplerSetTemplateResponse", + "statusCode": 200 + }, + "POST /twinkle/unload_adapter_paths": { + "body": [ + "UnloadAdapterPathsRequest" + ], + "response": "dict", + "statusCode": 200 + } + } +} diff --git a/tests/server/contract/test_client_api_contract.py b/tests/server/contract/test_client_api_contract.py new file mode 100644 index 000000000..f543ee436 --- /dev/null +++ b/tests/server/contract/test_client_api_contract.py @@ -0,0 +1,74 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Client-API wire-surface guards. + +Two guards with different strengths, kept apart on purpose: + +1. :func:`test_route_inventory_matches_committed_snapshot` compares the live route + inventory against ``client_api_routes.json``, which **is** committed. This is the + real guard: an unintended route addition/removal, or a changed response/body model, + fails here and shows up as a readable diff in the PR. +2. :func:`test_full_surface_extraction_is_self_consistent` only exercises the + field-level extractor. ``client_api_baseline.json`` is a generated, gitignored + artifact, so comparing against it cannot detect drift -- it would be comparing the + code to itself. The test is therefore scoped to what it can honestly assert: that + extraction runs, covers all five apps, and round-trips through JSON. + +Plus the load-bearing structural invariants of the request-lifecycle refactor. +""" +from __future__ import annotations + +import json + +import pytest + +from tests.server.contract.client_api_harness import (extract_full_surface, extract_route_inventory, + load_route_inventory) + +_APPS = {'data_plane', 'gateway', 'model', 'processor', 'sampler'} + + +def test_route_inventory_matches_committed_snapshot(): + current = extract_route_inventory() + committed = load_route_inventory() + assert set(current) == _APPS + + diffs = [] + for app in sorted(set(current) | set(committed)): + cur_routes, old_routes = current.get(app, {}), committed.get(app, {}) + for key in sorted(set(cur_routes) | set(old_routes)): + if cur_routes.get(key) != old_routes.get(key): + diffs.append(f' [{app}] {key}: committed={old_routes.get(key)} current={cur_routes.get(key)}') + assert not diffs, ('Client-facing route surface differs from the committed inventory.\n' + 'If the change is intentional, regenerate and review the diff:\n' + ' python -m tests.server.contract.update_baseline\n' + '\n'.join(diffs)) + + +def test_full_surface_extraction_is_self_consistent(): + # Scoped to what a self-generated snapshot can prove: the extractor works. + surface = extract_full_surface() + assert set(surface) == _APPS + for app, contract in surface.items(): + assert contract['paths'], f'{app} exposed no routes' + assert json.loads(json.dumps(surface, sort_keys=True)) == surface + + +def test_schedule_task_and_wait_removed(): + # Future records replace the in-process blocking wait. + from twinkle.server.task_queue.mixin import TaskQueueMixin + assert not hasattr(TaskQueueMixin, 'schedule_task_and_wait') + assert hasattr(TaskQueueMixin, 'submit_and_peek') + + +def test_new_client_types_importable(): + # The only permitted client-side additions. + import twinkle.protocol.types.base as base + import twinkle.protocol.types.errors as errors + + for symbol in ('StrictRequest', 'ResponseModel', 'DataModel', 'backend_only'): + assert hasattr(base, symbol) + for symbol in ('ErrorPayload', 'ErrorCategory', 'QueueStateLiteral'): + assert hasattr(errors, symbol) + + +if __name__ == '__main__': + raise SystemExit(pytest.main([__file__, '-v'])) diff --git a/tests/server/contract/test_error_wire.py b/tests/server/contract/test_error_wire.py new file mode 100644 index 000000000..41b394e42 --- /dev/null +++ b/tests/server/contract/test_error_wire.py @@ -0,0 +1,38 @@ +from __future__ import annotations + +from fastapi import FastAPI +from fastapi.testclient import TestClient +from tinker.types import RequestFailedResponse + +from twinkle.server.gateway.tinker_handlers import _register_gateway_tinker_routes + + +class _State: + + async def get_future(self, request_id: str): + return { + 'status': 'failed', + 'failure': { + 'reason_code': 'execution_timeout', + 'message': 'backend timed out', + 'attribution': 'server', + }, + } + + +class _Gateway: + state = _State() + + +def test_retrieve_future_returns_parseable_error_payload(): + app = FastAPI() + _register_gateway_tinker_routes(app, lambda: _Gateway()) + + response = TestClient(app).post('/retrieve_future', json={'request_id': 'req-1'}) + + assert response.status_code == 200 + body = response.json() + assert body['error_code'] == 504 + assert body['request_id'] == 'req-1' + parsed = RequestFailedResponse.model_validate(body) + assert parsed.category.value == 'server' diff --git a/tests/server/contract/test_protocol_migration.py b/tests/server/contract/test_protocol_migration.py new file mode 100644 index 000000000..44bab4116 --- /dev/null +++ b/tests/server/contract/test_protocol_migration.py @@ -0,0 +1,35 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Guards for the wire-contract migration to :mod:`twinkle.protocol`.""" +from __future__ import annotations + +_MIGRATED_NAMES = frozenset( + """ +AdapterRequest AddAdapterRequest AddMetricRequest AddMetricResponse ApplyPatchRequest ApplyPatchResponse +BACKEND_ONLY_KEY BackwardResponse CORE_INPUT_KEYS CalculateLossResponse CalculateMetricRequest +CalculateMetricResponse CancelRequest CancelResponse CapacityInfoResponse Checkpoint CheckpointPathResponse +CheckpointsListResponse ClientFeatures ClipGradAndStepRequest ClipGradAndStepResponse ClipGradNormRequest +ClipGradNormResponse CreateModelRequest CreateRequest CreateResponse CreateSessionRequest CreateSessionResponse Cursor +DataAppendRequest DataGetRequest DataModel DataPlaneForwardOnlyRequest DataPlaneForwardRequest DataPlaneSampleRequest +DataPutRequest DataRef DataReleaseRequest DataRowsResponse DeleteCheckpointResponse FieldRole ForwardBackwardResponse +ForwardBackwardTaskRequest ForwardOnlyRequest ForwardRequest ForwardResponse GetServerCapabilitiesResponse +GetTrainConfigsResponse HealthResponse LoadRequest LoadResponse LoraConfig LrStepRequest LrStepResponse ModelResult +OkResponse ParsedCheckpointTwinklePath ProcessorCallRequest ProcessorCallResponse ProcessorCreateRequest +ProcessorCreateResponse ProcessorHeartbeatRequest ProcessorHeartbeatResponse ProtocolLimits ResolvedLoadPath +ResponseModel ResumeFromCheckpointRequest RetrieveFutureRequest SampleRequest SampleResponseModel +SampleResponseModelList SampledSequenceModel SamplerAddAdapterRequest SamplerAddAdapterResponse SamplerCreateResponse +SamplerSetTemplateRequest SamplerSetTemplateResponse SaveRequest SaveResponse SessionHeartbeatRequest +SessionHeartbeatResponse SetLossRequest SetLossResponse SetLrSchedulerRequest SetLrSchedulerResponse SetOptimizerRequest +SetOptimizerResponse SetProcessorRequest SetProcessorResponse SetTemplateRequest SetTemplateResponse StepRequest +StepResponse StrictRequest SupportedModel TERMINAL_STATUSES TaskEnvelope TaskStatus TrainingProgressResponse TrainingRun +TrainingRunsResponse UnloadAdapterPathsRequest UploadToHubRequest VLM_TENSOR_FIELDS WeightsInfoRequest +WeightsInfoResponse WireInputBatch WireInputFeature WireInputs WireMessage WireTrajectory ZeroGradResponse backend_kwarg +backend_only declared_wire_keys export_batch fields_with_role passthrough read_backend_only read_field_role +""".split()) + + +def test_protocol_exports_match_pre_migration_snapshot() -> None: + import twinkle.protocol.types as types + + assert set(types.__all__) == _MIGRATED_NAMES + assert len(types.__all__) == 120 + assert 'ErrorResponse' not in types.__all__ diff --git a/tests/server/contract/update_baseline.py b/tests/server/contract/update_baseline.py index 45609bfb7..f5b470973 100644 --- a/tests/server/contract/update_baseline.py +++ b/tests/server/contract/update_baseline.py @@ -1,21 +1,32 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -"""Regenerate the client-API contract baseline. +"""Manual contract-maintenance helper; it is not a pytest test or CI entry point. + +Regenerate the client-API contract snapshots. Run with:: python -m tests.server.contract.update_baseline -Only invoke after confirming that the current client-facing surface has been -intentionally changed and approved as part of this refactor. +Two artifacts, with deliberately different git treatment: + +- ``client_api_baseline.json`` -- the full field-level surface. NOT committed + (gitignored): an 8k-line diff on every intentional wire change is noise nobody + reads, and being regenerated from the code under test it proves nothing on its own. +- ``client_api_routes.json`` -- the compact route inventory. COMMITTED, and it is the + actual regression guard. Review its diff: every added/removed route and every + changed response/body model appears there as a readable line. + +Only regenerate after confirming that the current client-facing surface changed +intentionally. """ from __future__ import annotations -from tests.server.contract.client_api_harness import write_baseline +from tests.server.contract.client_api_harness import write_baseline, write_route_inventory def main() -> None: - path = write_baseline() - print(f'Wrote baseline: {path}') + print(f'Wrote generated baseline (not committed): {write_baseline()}') + print(f'Wrote route inventory (COMMIT THIS): {write_route_inventory()}') if __name__ == '__main__': diff --git a/tests/server/data_plane/test_proxy.py b/tests/server/data_plane/test_proxy.py index 550998db2..2956ea327 100644 --- a/tests/server/data_plane/test_proxy.py +++ b/tests/server/data_plane/test_proxy.py @@ -3,8 +3,8 @@ import pytest from twinkle.server.data_plane.proxy import DataPlaneProxy -from twinkle_client.http.headers import H_AUTH, H_AUTH_TWINKLE, H_REQUEST_ID -from twinkle_client.types import DataRef +from twinkle.protocol.headers import H_AUTH, H_AUTH_TWINKLE, H_REQUEST_ID +from twinkle.protocol.types import DataRef class _Response: diff --git a/tests/server/data_plane/test_store.py b/tests/server/data_plane/test_store.py index 3c7ec1633..8ebd315d6 100644 --- a/tests/server/data_plane/test_store.py +++ b/tests/server/data_plane/test_store.py @@ -8,7 +8,11 @@ @pytest.mark.asyncio async def test_data_ref_round_trip_append_release_and_ref_isolation(monkeypatch) -> None: - import transfer_queue as tq + # TransferQueue ships in the `async-rl` extra, not in `server` / `client`, so an + # environment installed without that extra must SKIP here rather than fail. A bare + # import turned an absent optional dependency into a permanently red test, which + # teaches readers to ignore red. + tq = pytest.importorskip('transfer_queue') records = {} @@ -99,7 +103,7 @@ async def kv_list(partition_id): @pytest.mark.asyncio async def test_append_rejects_row_count_mismatch() -> None: - from twinkle_client.types import DataRef + from twinkle.protocol.types import DataRef store = TQDataRefStore.__new__(TQDataRefStore) ref = DataRef(ref_id='r', size=2, fields=['x']) @@ -108,7 +112,7 @@ async def test_append_rejects_row_count_mismatch() -> None: def test_partition_is_stable_and_scoped_by_data_ref() -> None: - from twinkle_client.types import DataRef + from twinkle.protocol.types import DataRef first = DataRef(ref_id='a', size=1, fields=['x']) same = DataRef(ref_id='a', size=99, fields=['other']) diff --git a/tests/server/fixtures/server_config_mock.yaml b/tests/server/fixtures/server_config_mock.yaml index 39833f1c0..67303ed5c 100644 --- a/tests/server/fixtures/server_config_mock.yaml +++ b/tests/server/fixtures/server_config_mock.yaml @@ -12,10 +12,9 @@ http_options: persistence: # Gateway and Model run as separate Ray Serve replicas (separate - # processes); the tinker future flow needs cross-process visibility. - # ``memory`` mode is per-process; file (or redis) is required. - mode: file - file_path: /tmp/twinkle_state_mock.json + # processes); the tinker future flow needs cross-process visibility, which + # ``memory`` provides via a detached Ray named actor. + mode: memory applications: diff --git a/tests/server/gateway/test_openai_handlers.py b/tests/server/gateway/test_openai_handlers.py index 8a98afbbe..048a4ad61 100644 --- a/tests/server/gateway/test_openai_handlers.py +++ b/tests/server/gateway/test_openai_handlers.py @@ -17,18 +17,10 @@ # ---------- Fixtures ------------------------------------------------------- # -@pytest.fixture(autouse=True) -def _reset_template_cache(): - from twinkle.server.gateway.openai_handlers import _template_initialized - _template_initialized.clear() - yield - _template_initialized.clear() - - @pytest.fixture def mock_gateway(): """Build a minimal FastAPI app with OpenAI routes and a mock GatewayServer.""" - import twinkle_client.types as types + import twinkle.protocol.types as types from twinkle.server.gateway.openai_handlers import _register_openai_routes mock_state = AsyncMock() @@ -42,7 +34,10 @@ def mock_gateway(): mock_self.state = mock_state mock_self.proxy = mock_proxy mock_self.supported_models = [types.SupportedModel(model_name='Qwen/Qwen3.5-4B')] - mock_self._supported_model_names = frozenset(['Qwen/Qwen3.5-4B']) + mock_self.supported_model_names = frozenset(['Qwen/Qwen3.5-4B']) + # Per-instance template cache; a fresh mock per test isolates it, so the + # former module-global clear fixture is no longer needed. + mock_self._template_initialized = set() app = FastAPI() _register_openai_routes(app, lambda: mock_self) @@ -125,7 +120,7 @@ def test_missing_messages_returns_400(self, mock_gateway): def test_model_not_found_returns_404(self, mock_gateway): mock_self, app = mock_gateway mock_self.supported_models = [] # No supported models - mock_self._supported_model_names = frozenset() + mock_self.supported_model_names = frozenset() client = TestClient(app) resp = client.post( diff --git a/tests/server/gateway/test_proxy.py b/tests/server/gateway/test_proxy.py new file mode 100644 index 000000000..9b03e4a98 --- /dev/null +++ b/tests/server/gateway/test_proxy.py @@ -0,0 +1,66 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Directed tests for ``gateway/proxy.py``. + +Covers the route-URL construction now sourced from ``gateway.routes``, the +``H_MULTIPLEX`` header compatibility in ``_prepare_headers``, and the 502 ``ErrorPayload`` +fallback when the upstream is unreachable. +""" +from __future__ import annotations + +import json + +import pytest +from starlette.requests import Request +from unittest.mock import AsyncMock + +from twinkle.server.gateway.proxy import ServiceProxy +from twinkle.protocol.headers import H_MULTIPLEX, H_MULTIPLEX_LEGACY, H_REQUEST_ID + + +def _make_request(headers: list[tuple[bytes, bytes]] | None = None) -> Request: + scope = { + 'type': 'http', + 'method': 'POST', + 'headers': headers or [], + 'query_string': b'', + 'path': '/', + } + return Request(scope) + + +@pytest.mark.parametrize( + 'route_prefix,host,expected', + [ + ('/api/v1', 'localhost', 'http://localhost:8000/api/v1/model/Qwen/tinker/forward'), + ('/api/v1/', 'localhost', 'http://localhost:8000/api/v1/model/Qwen/tinker/forward'), + ('', 'localhost', 'http://localhost:8000/model/Qwen/tinker/forward'), + ('/api/v1', '0.0.0.0', 'http://localhost:8000/api/v1/model/Qwen/tinker/forward'), + ], +) +def test_build_target_url(route_prefix, host, expected): + proxy = ServiceProxy(http_options={'host': host, 'port': 8000}, route_prefix=route_prefix) + assert proxy._build_target_url('model', 'Qwen', 'tinker/forward') == expected + + +def test_prepare_headers_sets_multiplex_from_request_id(): + proxy = ServiceProxy(http_options={}, route_prefix='/api/v1') + headers = proxy._prepare_headers({H_REQUEST_ID: 'req-1'}) + assert headers.get(H_MULTIPLEX) == 'req-1' + assert headers.get(H_MULTIPLEX_LEGACY) == 'req-1' + # ``host`` / ``content-length`` are stripped before forwarding. + assert 'host' not in {k.lower() for k in headers} + + +@pytest.mark.asyncio +async def test_proxy_request_502_fallback_returns_error_payload(): + proxy = ServiceProxy(http_options={'host': 'localhost', 'port': 8000}, route_prefix='/api/v1') + proxy.client.request = AsyncMock(side_effect=RuntimeError('upstream down')) + request = _make_request(headers=[(H_REQUEST_ID.encode(), b'req-9')]) + + response = await proxy.proxy_request(request, 'tinker/forward', 'Qwen', 'model', body_override=b'{}') + + assert response.status_code == 502 + payload = json.loads(response.body) + assert payload['error_code'] == 502 + assert payload['category'] == 'server' + assert payload['request_id'] == 'req-9' diff --git a/tests/server/integration/e2e_helpers.py b/tests/server/integration/e2e_helpers.py index 9d3073bad..2eca95f0d 100644 --- a/tests/server/integration/e2e_helpers.py +++ b/tests/server/integration/e2e_helpers.py @@ -22,7 +22,9 @@ MODEL_ID = f'ms://{BASE_MODEL}' BASE_URL = os.environ.get('TWINKLE_SERVER_URL', 'http://localhost:9000') API_KEY = 'EMPTY_API_KEY' -TIMEOUT = 120 # seconds per operation before declaring hang +TIMEOUT = float(os.environ.get('TWINKLE_TEST_OPERATION_TIMEOUT', '120')) +# Per-operation hang threshold. PPU Megatron cold JIT can exceed 300s; callers may +# raise this without weakening the default CI/GPU bound. GRADIENT_ACCUMULATION_STEPS = 2 # Megatron requires GA >= 2 @@ -66,6 +68,17 @@ def log(msg: str) -> None: # Dataset Factories # ═══════════════════════════════════════════════════════════════════════════ + +def _local_arrow_dataset(path: str, data_slice): + """Load selected rows from cached Arrow without hub metadata access.""" + from datasets import Dataset as HFDataset + from twinkle.dataset import Dataset, DatasetMeta + + source = HFDataset.from_file(path) + indices = [index % len(source) for index in data_slice] + return Dataset(DatasetMeta(data=source.select(indices))) + + def create_sft_dataset(data_slice=range(100)): """Create SelfCognition SFT dataset (small slice for speed).""" from twinkle.dataloader import DataLoader @@ -83,7 +96,11 @@ def create_dpo_dataset(data_slice=range(50)): from twinkle.dataset import Dataset, DatasetMeta from twinkle.preprocessor import EmojiDPOProcessor - dataset = Dataset(DatasetMeta('ms://hjh0119/shareAI-Llama3-DPO-zh-en-emoji', data_slice=data_slice)) + local_arrow = os.environ.get('TWINKLE_TEST_DPO_ARROW') + if local_arrow: + dataset = _local_arrow_dataset(local_arrow, data_slice) + else: + dataset = Dataset(DatasetMeta('ms://hjh0119/shareAI-Llama3-DPO-zh-en-emoji', data_slice=data_slice)) dataset.set_template('Qwen3_5Template', model_id=MODEL_ID, max_length=1024) dataset.map(EmojiDPOProcessor, init_args={'system': 'You are a helpful assistant.'}) dataset.encode() @@ -97,7 +114,11 @@ def create_grpo_dataset(data_slice=range(50)): system_prompt = ('You are a helpful math assistant. Solve the problem with minimal but correct reasoning ' 'and put your final answer within \\boxed{}.') - dataset = Dataset(DatasetMeta('ms://modelscope/gsm8k', subset_name='main', split='train', data_slice=data_slice)) + local_arrow = os.environ.get('TWINKLE_TEST_GRPO_ARROW') + if local_arrow: + dataset = _local_arrow_dataset(local_arrow, data_slice) + else: + dataset = Dataset(DatasetMeta('ms://modelscope/gsm8k', subset_name='main', split='train', data_slice=data_slice)) dataset.set_template('Qwen3_5Template', model_id=MODEL_ID, max_length=2048, enable_thinking=False) dataset.map(GSM8KProcessor(system=system_prompt)) dataset.encode(add_generation_prompt=True) diff --git a/tests/server/integration/test_actor_recovery.py b/tests/server/integration/test_actor_recovery.py new file mode 100644 index 000000000..16527bbae --- /dev/null +++ b/tests/server/integration/test_actor_recovery.py @@ -0,0 +1,86 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Post-timeout liveness probe and health status bit. + +Binds the real ``ModelManagement`` health methods onto a minimal harness with a +toggleable mock ``ping`` and a direct ``call_backend``. No GPU/Ray/full server. +""" +from __future__ import annotations + +import pytest +from fastapi import FastAPI + +from twinkle.server.model.app import ModelManagement +from twinkle.server.model.twinkle_handlers import _register_model_twinkle_routes + + +class _MockModel: + + def __init__(self) -> None: + self.alive = True + + def ping(self) -> bool: + if not self.alive: + raise RuntimeError('actor unreachable (simulated)') + return True + + +class _HealthHarness: + # Reuse the real implementations under test. + _run_model_health_probe = ModelManagement._run_model_health_probe + check_model_health = ModelManagement.check_model_health + mark_unhealthy = ModelManagement.mark_unhealthy + _probe_after_timeout = ModelManagement._probe_after_timeout + + def __init__(self, model: _MockModel) -> None: + self.model = model + self._model_unhealthy = False + + async def call_backend(self, fn, /, *args, admit: bool = True, **kwargs): + return fn(*args, **kwargs) + + +@pytest.mark.asyncio +async def test_timeout_probe_marks_unhealthy_then_recovers(): + model = _MockModel() + h = _HealthHarness(model) + + # Healthy at first. + result = await h.check_model_health() + assert result['healthy'] is True + assert h._model_unhealthy is False + + # A backend timeout fires the probe while the actor is unreachable. + model.alive = False + await h._probe_after_timeout() + assert h._model_unhealthy is True # /healthz would return 503 + + # Actor recovers; one successful probe clears the bit (no restart needed). + model.alive = True + result = await h.check_model_health() + assert result['healthy'] is True + assert h._model_unhealthy is False + + +@pytest.mark.asyncio +async def test_health_route_returns_503_when_probe_fails(): + model = _MockModel() + model.alive = False + harness = _HealthHarness(model) + app = FastAPI() + _register_model_twinkle_routes(app, lambda: harness) + route = next(route for route in app.routes if getattr(route, 'path', None) == '/healthz') + + response = await route.endpoint(object(), harness) + + assert response.status_code == 503 + + +@pytest.mark.asyncio +async def test_mark_unhealthy_is_cleared_by_successful_probe(): + h = _HealthHarness(_MockModel()) + h.mark_unhealthy() + assert h._model_unhealthy is True + + result = await h.check_model_health() + assert result['healthy'] is True + assert h._model_unhealthy is False diff --git a/tests/server/integration/test_blocking_boundary.py b/tests/server/integration/test_blocking_boundary.py new file mode 100644 index 000000000..7d35e229d --- /dev/null +++ b/tests/server/integration/test_blocking_boundary.py @@ -0,0 +1,215 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Blocking_Call_Boundary integration tests. + +These exercise the real ``TaskQueueMixin.call_backend`` through a minimal harness +that sets only the two attributes it uses (a dedicated executor and the optional +Admission_Gate), constructed exactly as ``_init_task_queue`` does. The backend is a +deliberately slow plain callable -- no GPU, Megatron, or Ray involved. +""" +from __future__ import annotations + +import asyncio +import httpx +import pytest +import threading +import time +from concurrent.futures import ThreadPoolExecutor +from fastapi import FastAPI +from fastapi.responses import JSONResponse + +ray = pytest.importorskip('ray') + +from twinkle.server.task_queue.backend_gate import BackendGate # noqa: E402 +from twinkle.server.task_queue.mixin import TaskQueueMixin # noqa: E402 +from twinkle.server.task_queue.types import BackendBusyError # noqa: E402 + + +class _Harness(TaskQueueMixin): + """Minimal holder exposing the real call_backend with a chosen gate setting.""" + + def __init__(self, gate_enabled: bool, *, max_workers: int | None = None) -> None: + # ``call_backend`` delegates to the extracted ``BackendGate``; build one and + # alias its internals so the assertions below still read the same names. + self._backend_gate = BackendGate(enable_admission_gate=gate_enabled) + if max_workers is not None: + self._backend_gate._executor.shutdown(wait=False) + self._backend_gate._executor = ThreadPoolExecutor( + max_workers=max_workers, thread_name_prefix='twinkle-backend') + self._backend_executor = self._backend_gate._executor + self._backend_probe_executor = self._backend_gate._probe_executor + self._backend_admission = self._backend_gate._admission + self._backend_poisoned = self._backend_gate._poisoned + + def close(self) -> None: + self._backend_gate.shutdown() + + +@pytest.mark.asyncio +async def test_healthz_style_probe_responsive_during_slow_backend(): + """While a slow backend call is in flight, an admit=False probe + (as /healthz uses) returns well within 5 seconds.""" + h = _Harness(gate_enabled=True) + try: + slow = asyncio.create_task(h.call_backend(lambda: time.sleep(3.0))) + await asyncio.sleep(0.05) # let the slow call take the gate + a thread + + loop = asyncio.get_running_loop() + start = loop.time() + probe = await h.call_backend(lambda: 'pong', admit=False) # no gate, like the ping probe + elapsed = loop.time() - start + + assert probe == 'pong' + assert elapsed < 5.0 + await slow + finally: + h.close() + + +@pytest.mark.asyncio +async def test_normal_gate_contention_waits_instead_of_failing(): + h = _Harness(gate_enabled=True) + try: + first = asyncio.create_task(h.call_backend(lambda: (time.sleep(0.2), 'first')[1])) + await asyncio.sleep(0.05) + second = asyncio.create_task(h.call_backend(lambda: 'second')) + assert await first == 'first' + assert await second == 'second' + finally: + h.close() + + +@pytest.mark.asyncio +async def test_cancelled_gate_waiter_does_not_steal_lock(): + h = _Harness(gate_enabled=True) + release = threading.Event() + + def wait_for_release(): + while not release.is_set(): + time.sleep(0.01) + + try: + first = asyncio.create_task(h.call_backend(wait_for_release)) + await asyncio.sleep(0.05) + waiter = asyncio.create_task(h.call_backend(lambda: 'cancelled')) + await asyncio.sleep(0.05) + waiter.cancel() + with pytest.raises(asyncio.CancelledError): + await waiter + release.set() + await first + assert await h.call_backend(lambda: 'next') == 'next' + finally: + release.set() + h.close() + + +@pytest.mark.asyncio +async def test_cancelled_queued_backend_call_releases_gate(): + h = _Harness(gate_enabled=True, max_workers=1) + release_worker = threading.Event() + occupied = h._backend_executor.submit(release_worker.wait) + try: + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(h.call_backend(lambda: 'never-started'), timeout=0.05) + await asyncio.sleep(0) + assert not h._backend_admission.locked() + release_worker.set() + occupied.result(timeout=5) + assert await h.call_backend(lambda: 'next') == 'next' + finally: + release_worker.set() + h.close() + + +@pytest.mark.asyncio +async def test_gate_held_by_leaked_call_fast_fails_next_task(): + """A call that outlives its wait_for keeps the gate; the next + admitting call fails fast with BackendBusyError instead of entering the backend.""" + h = _Harness(gate_enabled=True) + entered = {'count': 0} + + def slow(): + time.sleep(1.5) + + def would_enter_backend(): + entered['count'] += 1 + return 'should-not-run' + + try: + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(h.call_backend(slow), timeout=0.3) + + # The leaked thread still holds the gate. + with pytest.raises(BackendBusyError): + await h.call_backend(would_enter_backend) + assert entered['count'] == 0 # never reached the backend + + # After the leaked thread truly finishes, the gate frees on its own. + await asyncio.sleep(1.6) + assert await h.call_backend(would_enter_backend) == 'should-not-run' + assert entered['count'] == 1 + finally: + h.close() + + +@pytest.mark.asyncio +async def test_probe_times_out_while_same_serial_actor_is_busy(): + + @ray.remote + class SerialActor: + + def slow(self): + time.sleep(1.0) + + def ping(self): + return True + + started_ray = not ray.is_initialized() + if started_ray: + ray.init(num_cpus=1, logging_level='ERROR') + actor = SerialActor.remote() + h = _Harness(gate_enabled=True) + app = FastAPI() + + @app.get('/healthz') + async def healthz(): + try: + await h.call_backend(lambda: ray.get(actor.ping.remote(), timeout=0.2), admit=False) + return {'healthy': True} + except ray.exceptions.GetTimeoutError: + return JSONResponse(status_code=503, content={'healthy': False}) + + try: + slow = asyncio.create_task(h.call_backend(lambda: ray.get(actor.slow.remote(), timeout=10))) + await asyncio.sleep(0.1) + start = time.monotonic() + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url='http://test') as client: + response = await client.get('/healthz') + assert response.status_code == 503 + assert time.monotonic() - start < 5 + await slow + finally: + h.close() + ray.kill(actor) + if started_ray: + ray.shutdown() + + +@pytest.mark.asyncio +async def test_sampler_without_gate_runs_two_calls_concurrently(): + """Case 3 / opt-in: with the gate disabled (SamplerManagement), two backend + calls are in flight at once rather than serialized.""" + h = _Harness(gate_enabled=False) + try: + loop = asyncio.get_running_loop() + start = loop.time() + results = await asyncio.gather( + h.call_backend(lambda: (time.sleep(1.0), 'a')[1]), + h.call_backend(lambda: (time.sleep(1.0), 'b')[1]), + ) + elapsed = loop.time() - start + + assert sorted(results) == ['a', 'b'] + assert elapsed < 1.8 # concurrent, not ~2.0s serialized + finally: + h.close() diff --git a/tests/server/integration/test_dpo_e2e.py b/tests/server/integration/test_dpo_e2e.py index 035ae188f..7a69cbcab 100644 --- a/tests/server/integration/test_dpo_e2e.py +++ b/tests/server/integration/test_dpo_e2e.py @@ -182,7 +182,7 @@ def test_dpo_tinker(): """ from tinker import types from twinkle.dataloader import DataLoader - from twinkle.server.common import input_feature_to_datum + from twinkle.server.model.tinker_datum import input_feature_to_datum backend = get_backend() log(f'=== test_dpo_tinker [backend={backend}] ===') diff --git a/tests/server/integration/test_full_cycle_e2e.py b/tests/server/integration/test_full_cycle_e2e.py index 383633e90..6fa31f006 100644 --- a/tests/server/integration/test_full_cycle_e2e.py +++ b/tests/server/integration/test_full_cycle_e2e.py @@ -44,7 +44,8 @@ reason='Set TWINKLE_TEST_GPU_E2E=1 to run real GPU E2E tests (requires running server)', ) -from twinkle import get_logger, init_twinkle_client # noqa: E402 +from twinkle import get_logger # noqa: E402 +from twinkle import init_twinkle_client # noqa: E402 from twinkle.dataloader import DataLoader # noqa: E402 from twinkle.dataset import Dataset, DatasetMeta # noqa: E402 from twinkle_client.model import MultiLoraTransformersModel # noqa: E402 diff --git a/tests/server/integration/test_full_param_e2e.py b/tests/server/integration/test_full_param_e2e.py index b328b8a5e..857fbc711 100644 --- a/tests/server/integration/test_full_param_e2e.py +++ b/tests/server/integration/test_full_param_e2e.py @@ -41,11 +41,12 @@ reason='Set TWINKLE_TEST_GPU_E2E=1 to run real GPU E2E tests (requires running server)', ) -from twinkle import get_logger, init_tinker_client # noqa: E402 +from twinkle import get_logger # noqa: E402 +from twinkle import init_tinker_client # noqa: E402 from twinkle.dataloader import DataLoader # noqa: E402 from twinkle.dataset import Dataset, DatasetMeta # noqa: E402 from twinkle.preprocessor import SelfCognitionProcessor # noqa: E402 -from twinkle.server.common import input_feature_to_datum # noqa: E402 +from twinkle.server.model.tinker_datum import input_feature_to_datum # noqa: E402 init_tinker_client() diff --git a/tests/server/integration/test_mock_mode_startup.py b/tests/server/integration/test_mock_mode_startup.py index be8610ced..d5f706061 100644 --- a/tests/server/integration/test_mock_mode_startup.py +++ b/tests/server/integration/test_mock_mode_startup.py @@ -211,7 +211,7 @@ def test_mock_mode_reaches_ready_under_30s_and_is_deterministic(ray_cluster) -> def _exercise_twinkle_clients(base: str) -> None: - from twinkle_client import init_twinkle_client + from twinkle import init_twinkle_client from twinkle_client.model import MultiLoraTransformersModel from twinkle_client.sampler import vLLMSampler @@ -254,8 +254,6 @@ def _exercise_twinkle_clients(base: str) -> None: assert isinstance(metric.result, dict) cfgs = model.get_train_configs() assert isinstance(cfgs.result, str) - state = model.get_state_dict() - assert isinstance(state.result, dict) save_resp = model.save(name='step-1') assert save_resp.twinkle_path and save_resp.twinkle_path.startswith('twinkle://') @@ -310,7 +308,7 @@ def _exercise_tinker_client(base: str) -> None: import os from tinker import ServiceClient, types - from twinkle_client import init_tinker_client + from twinkle import init_tinker_client # patch_tinker injects Twinkle's auth + Ray Serve multiplex headers and # lifts tinker's ``tml-`` api-key prefix check so EMPTY_TOKEN passes. diff --git a/tests/server/integration/test_nccl_safe_tinker_e2e.py b/tests/server/integration/test_nccl_safe_tinker_e2e.py index aa94f99c4..23259c7c8 100644 --- a/tests/server/integration/test_nccl_safe_tinker_e2e.py +++ b/tests/server/integration/test_nccl_safe_tinker_e2e.py @@ -1,16 +1,15 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -"""Real E2E test for NCCL-safe fault tolerance via Tinker client path. +"""Real E2E test for loud-failure semantics via the Tinker client path. -Exercises the /tinker/forward_backward endpoint through the upstream Tinker SDK. -All adversarial scenarios verify that safe_loss catches errors gracefully without -NCCL hang or model state corruption. +Exercises ``/tinker/forward_backward`` through the upstream Tinker SDK. The +invariant under test (post silent-degradation removal): a request whose loss +computation fails does NOT come back as a silent zero-loss success -- it enters a +failed terminal state -- and a subsequent valid request on the same deployment +still succeeds. Prerequisites: - 1. Ray cluster running with GPUs (2 for model DP/TP, optionally 1 for sampler) - 2. Twinkle server started with TWINKLE_FAIL_FAST=0 - -Usage (direct): - python tests/server/integration/test_nccl_safe_tinker_e2e.py + 1. Ray cluster running with GPUs (2 for model DP/TP) + 2. Twinkle server started with queue_config.execution_timeout=30 Usage (pytest, requires TWINKLE_TEST_GPU_E2E=1): TWINKLE_TEST_GPU_E2E=1 pytest tests/server/integration/test_nccl_safe_tinker_e2e.py -v @@ -18,10 +17,7 @@ from __future__ import annotations import os -import sys import time -import logging -import traceback import numpy as np import pytest @@ -31,409 +27,105 @@ reason='Set TWINKLE_TEST_GPU_E2E=1 to run real GPU E2E tests (requires running server)', ) -logging.basicConfig(level=logging.INFO, format='%(asctime)s [%(levelname)s] %(message)s') -logger = logging.getLogger(__name__) - - -def log(msg): - """Print + flush to avoid log suppression by init_tinker_client().""" - print(f'[E2E-Tinker] {msg}', flush=True) - - BASE_MODEL = 'Qwen/Qwen3.5-4B' SERVER_URL = os.environ.get('TWINKLE_SERVER_URL', 'http://localhost:9000') -TIMEOUT = 120 - +EXECUTION_TIMEOUT = float(os.environ.get('TWINKLE_TEST_EXECUTION_TIMEOUT', '30')) +TIMEOUT = EXECUTION_TIMEOUT + 15 +# The `global_rank=` attribution is added by `nccl_safe_megatron`, which decorates +# only the Megatron backend; the transformers backend carries no such annotation. +# Gate the rank-attribution assertion on the backend so this file is safe under +# TWINKLE_TEST_BACKEND=transformers. +BACKEND = os.environ.get('TWINKLE_TEST_BACKEND', 'megatron') -def wait_for_server(url, timeout=300): - """Wait for Twinkle server to become ready.""" - import requests - start = time.time() - while time.time() - start < timeout: - try: - resp = requests.get(f'{url}/-/routes', timeout=5) - if resp.status_code == 200: - elapsed = int(time.time() - start) - log(f'Server is ready (waited {elapsed}s)') - return True - except Exception: - pass - time.sleep(5) - raise TimeoutError(f'Server not ready after {timeout}s') - -def init_client(): - """Initialize Tinker client and create training client.""" +def _init_client(): os.environ['TINKER_BASE_URL'] = SERVER_URL os.environ['TWINKLE_SERVER_TOKEN'] = 'EMPTY_TOKEN' - - from twinkle_client import init_tinker_client + from twinkle import init_tinker_client init_tinker_client() - from tinker import ServiceClient - service_client = ServiceClient() - training_client = service_client.create_lora_training_client(base_model=BASE_MODEL, rank=16) - log('Training client created successfully') - return training_client + return ServiceClient().create_lora_training_client(base_model=BASE_MODEL, rank=16) -def make_datum(seq_len=32, completion_len=16, *, bad_logprobs_len=None, include_advantages=True): - """Construct a Datum for GRPO training.""" +def _make_datum(seq_len=64, completion_len=32, *, bad_logprobs_len=None): from tinker import types - prompt_len = seq_len - completion_len input_tokens = list(range(1, seq_len + 1)) target_tokens = [0] * prompt_len + list(range(100, 100 + completion_len)) weights = [0] * prompt_len + [1] * completion_len - - if bad_logprobs_len is not None: - logprobs_values = np.random.randn(bad_logprobs_len).astype(np.float32) - padded_logprobs = [0.0] * prompt_len + logprobs_values.tolist() - else: - logprobs_values = np.random.randn(completion_len).astype(np.float32) - padded_logprobs = [0.0] * prompt_len + logprobs_values.tolist() - - loss_fn_inputs = { - 'target_tokens': target_tokens, - 'weights': weights, - 'logprobs': types.TensorData.from_numpy(np.array(padded_logprobs, dtype=np.float32)), - } - - if include_advantages: - advantage = float(np.random.randn()) - padded_advantages = [0.0] * prompt_len + [advantage] * completion_len - loss_fn_inputs['advantages'] = types.TensorData.from_numpy( - np.array(padded_advantages, dtype=np.float32)) - + n = bad_logprobs_len if bad_logprobs_len is not None else completion_len + padded_logprobs = [0.0] * prompt_len + np.random.randn(n).astype(np.float32).tolist() + advantage = float(np.random.randn()) return types.Datum( model_input=types.ModelInput.from_ints(input_tokens), - loss_fn_inputs=loss_fn_inputs, + loss_fn_inputs={ + 'target_tokens': target_tokens, + 'weights': weights, + 'logprobs': types.TensorData.from_numpy(np.array(padded_logprobs, dtype=np.float32)), + 'advantages': types.TensorData.from_numpy( + np.array([0.0] * prompt_len + [advantage] * completion_len, dtype=np.float32)), + }, ) -def run_forward_backward(training_client, datums, test_name, expect_success=True): - """Run forward_backward and return (success, result, elapsed_seconds).""" - log(f'[{test_name}] Sending {len(datums)} datums...') - start = time.time() - try: - result = training_client.forward_backward(datums, 'importance_sampling').result() - elapsed = time.time() - start - log(f'[{test_name}] Completed in {elapsed:.1f}s') - if hasattr(result, 'metrics') and result.metrics: - loss_avg = result.metrics.get('loss:avg', 'N/A') - log(f'[{test_name}] loss:avg = {loss_avg}') - return True, result, elapsed - except Exception as e: - elapsed = time.time() - start - log(f'[{test_name}] FAILED in {elapsed:.1f}s: {type(e).__name__}: {e}') - if elapsed > TIMEOUT: - log(f'[{test_name}] TIMEOUT! This suggests NCCL hang!') - return False, None, elapsed - - -def do_optim_step(training_client, test_name): - """Run optimizer step.""" +def _assert_recovery_terminal(tc) -> None: + """Require success on Megatron; Transformers may fail loudly after DDP poisoning.""" from tinker import types - try: - training_client.optim_step(types.AdamParams(learning_rate=1e-5)).result() - log(f'[{test_name}] optim_step OK') - return True - except Exception as e: - log(f'[{test_name}] optim_step FAILED: {e}') - return False - - -# ═══════════════════════════════════════════════════════════════════════════ -# Test Scenarios (19 tests) -# ═══════════════════════════════════════════════════════════════════════════ - -def test_1_normal_grpo(tc): - datums = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, result, elapsed = run_forward_backward(tc, datums, 'TEST-1-NORMAL') - assert ok and elapsed < TIMEOUT - do_optim_step(tc, 'TEST-1-NORMAL') - return True - -def test_2_bad_old_logps(tc): - datums = [ - make_datum(seq_len=64, completion_len=32), - make_datum(seq_len=64, completion_len=32, bad_logprobs_len=5), - make_datum(seq_len=64, completion_len=32), - make_datum(seq_len=64, completion_len=32, bad_logprobs_len=99), - ] - ok, result, elapsed = run_forward_backward(tc, datums, 'TEST-2-BAD-LOGPS') - if not ok: - return elapsed < TIMEOUT - assert elapsed < TIMEOUT - do_optim_step(tc, 'TEST-2-BAD-LOGPS') - return True - -def test_3_recovery(tc): - datums = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, datums, 'TEST-3-RECOVERY') - assert ok and elapsed < TIMEOUT - do_optim_step(tc, 'TEST-3-RECOVERY') - return True - -def test_4_no_advantages(tc): - datums = [make_datum(seq_len=64, completion_len=32, include_advantages=False) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, datums, 'TEST-4-NO-ADV') - assert ok and elapsed < TIMEOUT - do_optim_step(tc, 'TEST-4-NO-ADV') - return True + from tinker._exceptions import RequestFailedError -def test_5_consecutive_bad(tc): - for i in range(5): - datums = [make_datum(seq_len=64, completion_len=32, bad_logprobs_len=3+i) for _ in range(4)] - _, _, elapsed = run_forward_backward(tc, datums, f'TEST-5-{i+1}') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, f'TEST-5-{i+1}') - return True + request = tc.forward_backward([_make_datum() for _ in range(4)], 'importance_sampling') + if BACKEND == 'megatron': + assert request.result() is not None + tc.optim_step(types.AdamParams(learning_rate=1e-5)).result() + return + try: + assert request.result(timeout=TIMEOUT) is not None + tc.optim_step(types.AdamParams(learning_rate=1e-5)).result() + except RequestFailedError as exc: + assert exc.category is types.RequestErrorCategory.Server -def test_6_nan_logprobs(tc): - from tinker import types - datums = [] - for _ in range(4): - d = make_datum(seq_len=64, completion_len=32) - d.loss_fn_inputs['logprobs'] = types.TensorData.from_numpy( - np.array([float('nan')] * 64, dtype=np.float32)) - datums.append(d) - _, _, elapsed = run_forward_backward(tc, datums, 'TEST-6-NAN') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-6-NAN') - return True -def test_7_inf_logprobs(tc): - from tinker import types - datums = [] - for _ in range(4): - d = make_datum(seq_len=64, completion_len=32) - inf_arr = np.full(64, float('inf'), dtype=np.float32) - inf_arr[::2] = float('-inf') - d.loss_fn_inputs['logprobs'] = types.TensorData.from_numpy(inf_arr) - datums.append(d) - _, _, elapsed = run_forward_backward(tc, datums, 'TEST-7-INF') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-7-INF') - return True +def test_failure_is_terminal_then_valid_request_succeeds(): + """A malformed request fails loudly (terminal), a subsequent valid one succeeds. -def test_8_extreme_advantages(tc): + Replaces the former assertion "failure degraded to zero loss and training + continued". If the recovery request does not reach a terminal success, that is + recorded as evidence that actor recovery and admission gate are + necessary, not optional. + """ from tinker import types - datums = [] - for i in range(4): - d = make_datum(seq_len=64, completion_len=32) - val = 1e30 if i % 2 == 0 else -1e30 - adv = np.full(64, 0.0, dtype=np.float32) - adv[32:] = val - d.loss_fn_inputs['advantages'] = types.TensorData.from_numpy(adv) - datums.append(d) - _, _, elapsed = run_forward_backward(tc, datums, 'TEST-8-EXTREME-ADV') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-8-EXTREME-ADV') - return True - -def test_9_zero_completion(tc): - from tinker import types - datums = [] - for _ in range(4): - d = types.Datum( - model_input=types.ModelInput.from_ints(list(range(1, 65))), - loss_fn_inputs={ - 'target_tokens': [0]*64, 'weights': [0]*64, - 'logprobs': types.TensorData.from_numpy(np.zeros(64, dtype=np.float32)), - 'advantages': types.TensorData.from_numpy(np.zeros(64, dtype=np.float32)), - }, - ) - datums.append(d) - _, _, elapsed = run_forward_backward(tc, datums, 'TEST-9-ZERO-COMPL') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-9-ZERO-COMPL') - return True - -def test_10_partial_advantages(tc): - datums = [ - make_datum(seq_len=64, completion_len=32, include_advantages=True), - make_datum(seq_len=64, completion_len=32, include_advantages=False), - make_datum(seq_len=64, completion_len=32, include_advantages=True), - make_datum(seq_len=64, completion_len=32, include_advantages=False), - ] - _, _, elapsed = run_forward_backward(tc, datums, 'TEST-10-PARTIAL-ADV') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-10-PARTIAL-ADV') - return True - -def test_11_mixed_seq_lengths(tc): - datums = [ - make_datum(seq_len=32, completion_len=16), - make_datum(seq_len=128, completion_len=64), - make_datum(seq_len=48, completion_len=24), - make_datum(seq_len=96, completion_len=48), - ] - _, _, elapsed = run_forward_backward(tc, datums, 'TEST-11-MIXED') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-11-MIXED') - return True + from tinker._exceptions import RequestFailedError + tc = _init_client() -def test_12_all_bad(tc): - datums = [make_datum(seq_len=64, completion_len=32, bad_logprobs_len=i) for i in range(4)] - _, _, elapsed = run_forward_backward(tc, datums, 'TEST-12-ALL-BAD') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-12-ALL-BAD') - return True - -def test_13_forward_only_then_train(tc): - datums_infer = [make_datum(seq_len=64, completion_len=32, include_advantages=False) for _ in range(4)] + # Deliberately malformed: logprobs length inconsistent with the completion. + bad = [_make_datum(bad_logprobs_len=5) for _ in range(4)] start = time.time() - try: - tc.forward(datums_infer).result() - except Exception: - if time.time() - start >= TIMEOUT: - return False - datums_train = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, datums_train, 'TEST-13-TRAIN') - if not ok or elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-13-TRAIN') - return True - -def test_14_rapid_bad_good(tc): - for i in range(5): - bad = [make_datum(seq_len=64, completion_len=32, bad_logprobs_len=i+1) for _ in range(4)] - _, _, elapsed = run_forward_backward(tc, bad, f'TEST-14-BAD-{i+1}') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, f'TEST-14-BAD-{i+1}') - good = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, good, f'TEST-14-GOOD-{i+1}') - if not ok or elapsed >= TIMEOUT: - return False - do_optim_step(tc, f'TEST-14-GOOD-{i+1}') - return True + with pytest.raises(RequestFailedError) as caught: + tc.forward_backward(bad, 'importance_sampling').result(timeout=TIMEOUT) + assert caught.value.category is types.RequestErrorCategory.Server + assert time.time() - start < TIMEOUT, 'malformed request must fail fast, not hang (NCCL)' -def test_15_final_health(tc): - datums = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, datums, 'TEST-15-FINAL') - assert ok and elapsed < TIMEOUT - do_optim_step(tc, 'TEST-15-FINAL') - return True + # Megatron must recover successfully. Tinker's Transformers path executes + # forward/loss/backward separately; after a mid-iteration failure, the test only + # guarantees that the next request reaches a terminal state. + _assert_recovery_terminal(tc) -def test_16_large_batch(tc): - datums = [make_datum(seq_len=64, completion_len=32) for _ in range(16)] - ok, _, elapsed = run_forward_backward(tc, datums, 'TEST-16-LARGE') - if elapsed >= TIMEOUT: - return False - assert ok - do_optim_step(tc, 'TEST-16-LARGE') - return True - -def test_17_single_datum(tc): - # With dp_size=2 + nproc_per_node=2, minimum batch must be >= data_world_size - datums = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, datums, 'TEST-17-SMALL') - if elapsed >= TIMEOUT: - return False - assert ok - do_optim_step(tc, 'TEST-17-SMALL') - return True - -def test_18_save_after_error(tc): - bad = [make_datum(seq_len=64, completion_len=32, bad_logprobs_len=2) for _ in range(4)] - _, _, elapsed = run_forward_backward(tc, bad, 'TEST-18-ERR') - if elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-18-ERR') - try: - tc.save_weights_for_sampler().result() - except Exception: - pass - good = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, good, 'TEST-18-POST') - if not ok or elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-18-POST') - return True - -def test_19_consecutive_optim_steps(tc): - datums = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, datums, 'TEST-19-BASE') - assert ok and elapsed < TIMEOUT - for i in range(3): - do_optim_step(tc, f'TEST-19-STEP-{i+1}') - datums = [make_datum(seq_len=64, completion_len=32) for _ in range(4)] - ok, _, elapsed = run_forward_backward(tc, datums, 'TEST-19-VERIFY') - if not ok or elapsed >= TIMEOUT: - return False - do_optim_step(tc, 'TEST-19-VERIFY') - return True - - -ALL_TESTS = [ - ('TEST-1: Normal GRPO Training', test_1_normal_grpo), - ('TEST-2: Bad old_logps (original bug)', test_2_bad_old_logps), - ('TEST-3: Recovery after error', test_3_recovery), - ('TEST-4: No advantages (zero loss)', test_4_no_advantages), - ('TEST-5: Consecutive bad batches', test_5_consecutive_bad), - ('TEST-6: NaN logprobs', test_6_nan_logprobs), - ('TEST-7: +Inf/-Inf logprobs', test_7_inf_logprobs), - ('TEST-8: Extreme advantages (1e30)', test_8_extreme_advantages), - ('TEST-9: Zero completion tokens', test_9_zero_completion), - ('TEST-10: Partial advantages (ragged)', test_10_partial_advantages), - ('TEST-11: Mixed sequence lengths', test_11_mixed_seq_lengths), - ('TEST-12: All datums bad (100%)', test_12_all_bad), - ('TEST-13: forward_only then train', test_13_forward_only_then_train), - ('TEST-14: Rapid bad->good alternation', test_14_rapid_bad_good), - ('TEST-15: Final health check', test_15_final_health), - ('TEST-16: Large batch (16 datums)', test_16_large_batch), - ('TEST-17: Single datum batch', test_17_single_datum), - ('TEST-18: Save after error', test_18_save_after_error), - ('TEST-19: Consecutive optim_steps', test_19_consecutive_optim_steps), -] - - -def main(): - log('=' * 60) - log('NCCL-Safe E2E Test - Tinker Client Path') - log('=' * 60) - log(f'Server URL: {SERVER_URL}') - log(f'Base Model: {BASE_MODEL}') - log(f'TWINKLE_FAIL_FAST = {os.getenv("TWINKLE_FAIL_FAST", "1 (default)")}') - - wait_for_server(SERVER_URL) - tc = init_client() - - results = [] - for name, test_fn in ALL_TESTS: - log(f'\n{"=" * 60}\n{name}\n{"=" * 60}') - try: - passed = test_fn(tc) - results.append((name, 'PASS' if passed else 'FAIL')) - log(f'[{name}] {"PASS" if passed else "FAIL"}') - except Exception as e: - log(f'{name}: EXCEPTION: {e}') - traceback.print_exc() - results.append((name, 'FAIL')) - - log(f'\n{"=" * 60}\nRESULTS SUMMARY\n{"=" * 60}') - all_passed = all(s == 'PASS' for _, s in results) - for name, status in results: - log(f' [{status}] {name}') - log(f'\n{"ALL" if all_passed else "SOME"} {len(results)} TESTS {"PASSED" if all_passed else "FAILED"}!') - return 0 if all_passed else 1 - - -def test_nccl_safe_tinker_e2e(): - """Pytest-collected entry point.""" - rc = main() - assert rc == 0, 'Some Tinker NCCL-safe E2E tests failed' +def test_partial_rank_failure_is_terminal_then_recovers(): + from tinker import types + from tinker._exceptions import RequestFailedError + tc = _init_client() -if __name__ == '__main__': - sys.exit(main()) + batch = [_make_datum() for _ in range(4)] + batch[0] = _make_datum(bad_logprobs_len=5) + start = time.time() + with pytest.raises(RequestFailedError) as caught: + tc.forward_backward(batch, 'importance_sampling').result(timeout=TIMEOUT) + assert caught.value.category is types.RequestErrorCategory.Server + # Megatron attributes the failure to a global rank via nccl_safe_megatron; the + # transformers backend has no such annotation ( removed its old decorator). + if BACKEND == 'megatron': + assert 'global_rank=' in str(caught.value) + assert time.time() - start < TIMEOUT + + _assert_recovery_terminal(tc) diff --git a/tests/server/integration/test_nccl_safe_twinkle_e2e.py b/tests/server/integration/test_nccl_safe_twinkle_e2e.py index c9cce48a6..705e2ec34 100644 --- a/tests/server/integration/test_nccl_safe_twinkle_e2e.py +++ b/tests/server/integration/test_nccl_safe_twinkle_e2e.py @@ -1,16 +1,14 @@ # Copyright (c) ModelScope Contributors. All rights reserved. -"""Real E2E test for NCCL-safe fault tolerance via Twinkle client path. +"""Real E2E test for loud-failure semantics via the Twinkle-native client path. -Exercises the /twinkle/forward_backward endpoint through the Twinkle SDK -(init_twinkle_client + MultiLoraTransformersModel). This is a SEPARATE code -path from the Tinker SDK (/tinker/forward_backward). +Exercises ``/twinkle/forward_backward`` through the Twinkle client. The invariant +under test (post silent-degradation removal): a request whose loss computation fails +does NOT come back as a silent zero-loss success -- it fails loudly -- and a +subsequent valid request on the same deployment still succeeds. Prerequisites: - 1. Ray cluster running with GPUs (2 for model DP/TP, optionally 1 for sampler) - 2. Twinkle server started with TWINKLE_FAIL_FAST=0 - -Usage (direct): - python tests/server/integration/test_nccl_safe_twinkle_e2e.py + 1. Ray cluster running with GPUs (2 for model DP/TP) + 2. Twinkle server started with queue_config.execution_timeout=30 Usage (pytest, requires TWINKLE_TEST_GPU_E2E=1): TWINKLE_TEST_GPU_E2E=1 pytest tests/server/integration/test_nccl_safe_twinkle_e2e.py -v @@ -18,11 +16,7 @@ from __future__ import annotations import os -import sys import time -import logging -import traceback -from typing import Any, Dict, List import numpy as np import pytest @@ -32,314 +26,82 @@ reason='Set TWINKLE_TEST_GPU_E2E=1 to run real GPU E2E tests (requires running server)', ) -logging.basicConfig(level=logging.INFO, format='%(asctime)s [%(levelname)s] %(message)s') -logger = logging.getLogger(__name__) - - -def log(msg): - print(f'[E2E-Twinkle] {msg}', flush=True) - - BASE_MODEL = 'Qwen/Qwen3.5-4B' SERVER_URL = os.environ.get('TWINKLE_SERVER_URL', 'http://localhost:9000') -TIMEOUT = 120 -ADAPTER_NAME = 'nccl-safe-test' +EXECUTION_TIMEOUT = float(os.environ.get('TWINKLE_TEST_EXECUTION_TIMEOUT', '30')) +TIMEOUT = EXECUTION_TIMEOUT + 15 +ADAPTER_NAME = 'loud-failure-test' +# The `global_rank=` attribution is added by `nccl_safe_megatron`, which decorates +# only the Megatron backend; the transformers backend's forward_backward carries no +# such annotation. Gate the rank-attribution assertion on the backend so this file is +# safe to run under the integration-e2e SKILL's TWINKLE_TEST_BACKEND=transformers path. +BACKEND = os.environ.get('TWINKLE_TEST_BACKEND', 'megatron') -def wait_for_server(url, timeout=300): - """Wait for Twinkle server to become ready.""" - import requests - start = time.time() - while time.time() - start < timeout: - try: - resp = requests.get(f'{url}/-/routes', timeout=5) - if resp.status_code == 200: - log(f'Server is ready (waited {int(time.time() - start)}s)') - return True - except Exception: - pass - time.sleep(5) - raise TimeoutError(f'Server not ready after {timeout}s') - - -def init_client(): - """Initialize Twinkle client and configure model for GRPO training.""" - from twinkle_client import init_twinkle_client - from twinkle_client.model import MultiLoraTransformersModel +def _init_client(): from peft import LoraConfig + from twinkle import init_twinkle_client + from twinkle_client.model import MultiLoraTransformersModel init_twinkle_client(base_url=SERVER_URL, api_key='EMPTY_TOKEN') - model = MultiLoraTransformersModel(model_id=f'ms://{BASE_MODEL}') model.add_adapter_to_model( adapter_name=ADAPTER_NAME, config=LoraConfig(r=16, target_modules=['q_proj', 'v_proj']), - gradient_accumulation_steps=1, + # GA>=2 (repo convention, see e2e_helpers): with GA=1 every backward syncs + # DDP immediately, so a mid-iteration failure can leave the reducer + # half-finished and poison the next request. GA=2 runs accumulation steps + # under no_sync, keeping the recovery request clean. + gradient_accumulation_steps=2, ) model.set_loss('GRPOLoss', init_args={'epsilon': 0.2}) model.set_optimizer('Adam', lr=1e-5) model.set_template('Qwen3_5Template') model.set_processor('InputProcessor', padding_side='right') - log('Twinkle client + model configured successfully') return model -def make_input_features( - batch_size=4, seq_len=64, completion_len=32, *, - bad_old_logps_len=None, include_advantages=True, - nan_old_logps=False, extreme_advantages=None, all_labels_masked=False, -): - """Construct InputFeature list + old_logps + advantages for GRPO.""" +def _make_inputs(batch_size=4, seq_len=64, completion_len=32, *, bad_old_logps_len=None): prompt_len = seq_len - completion_len - input_features = [] - old_logps_list = [] - advantages_list = [] - - for i in range(batch_size): - input_ids = list(range(1, seq_len + 1)) - labels = [-100] * seq_len if all_labels_masked else ( - [-100] * prompt_len + list(range(100, 100 + completion_len))) - input_features.append({ - 'input_ids': input_ids, - 'labels': labels, + features, old_logps, advantages = [], [], [] + for _ in range(batch_size): + features.append({ + 'input_ids': list(range(1, seq_len + 1)), + 'labels': [-100] * prompt_len + list(range(100, 100 + completion_len)), 'attention_mask': [1] * seq_len, 'position_ids': list(range(seq_len)), }) - - if bad_old_logps_len is not None: - logps = np.random.randn(bad_old_logps_len).tolist() - elif nan_old_logps: - logps = [float('nan')] * completion_len - else: - logps = np.random.randn(completion_len).tolist() - old_logps_list.append(logps) - - if extreme_advantages is not None: - advantages_list.append(extreme_advantages if i % 2 == 0 else -extreme_advantages) - else: - advantages_list.append(float(np.random.randn())) - - old_logps = old_logps_list if include_advantages else None - advantages = advantages_list if include_advantages else None - return input_features, old_logps, advantages - - -def run_forward_backward(model, inputs, old_logps, advantages, test_name): - """Run forward_backward and return (success, result, elapsed_seconds).""" - log(f'[{test_name}] Sending {len(inputs)} input features...') - start = time.time() - try: - kwargs: Dict[str, Any] = {} - if old_logps is not None: - kwargs['old_logps'] = old_logps - if advantages is not None: - kwargs['advantages'] = advantages - - result = model.forward_backward(inputs=inputs, **kwargs) - elapsed = time.time() - start - log(f'[{test_name}] Completed in {elapsed:.1f}s') - if hasattr(result, 'result') and result.result is not None: - log(f'[{test_name}] result = {result.result}') - return True, result, elapsed - except Exception as e: - elapsed = time.time() - start - log(f'[{test_name}] FAILED in {elapsed:.1f}s: {type(e).__name__}: {e}') - if elapsed > TIMEOUT: - log(f'[{test_name}] TIMEOUT! This suggests NCCL hang!') - return False, None, elapsed - - -def do_optim_step(model, test_name): - """Run clip_grad_and_step.""" - try: - model.clip_grad_and_step() - log(f'[{test_name}] clip_grad_and_step OK') - return True - except Exception as e: - log(f'[{test_name}] clip_grad_and_step FAILED: {e}') - return False + n = bad_old_logps_len if bad_old_logps_len is not None else completion_len + old_logps.append(np.random.randn(n).tolist()) + advantages.append(float(np.random.randn())) + return features, old_logps, advantages -# ═══════════════════════════════════════════════════════════════════════════ -# Test Scenarios (12 tests) -# ═══════════════════════════════════════════════════════════════════════════ +def test_failure_is_terminal_then_valid_request_succeeds(): + """A malformed request fails loudly, a subsequent valid one succeeds. -def test_1_normal_grpo(m): - inputs, old_logps, adv = make_input_features(batch_size=4) - ok, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, 'TEST-1-NORMAL') - assert ok and elapsed < TIMEOUT - do_optim_step(m, 'TEST-1-NORMAL') - return True + Replaces the former assertion "failure degraded to zero loss and training + continued". If the recovery request does not reach a terminal success, that is + recorded as evidence that actor recovery and admission gate are + necessary, not optional. + """ + model = _init_client() -def test_2_bad_old_logps(m): - inputs, old_logps, adv = make_input_features(batch_size=4, bad_old_logps_len=5) - ok, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, 'TEST-2-BAD-LOGPS') - if not ok: - return elapsed < TIMEOUT - assert elapsed < TIMEOUT - do_optim_step(m, 'TEST-2-BAD-LOGPS') - return True - -def test_3_recovery(m): - inputs, old_logps, adv = make_input_features(batch_size=4) - ok, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, 'TEST-3-RECOVERY') - assert ok and elapsed < TIMEOUT - do_optim_step(m, 'TEST-3-RECOVERY') - return True - -def test_4_nan_old_logps(m): - inputs, old_logps, adv = make_input_features(batch_size=4, nan_old_logps=True) - _, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, 'TEST-4-NAN') - if elapsed >= TIMEOUT: - return False - do_optim_step(m, 'TEST-4-NAN') - return True - -def test_5_extreme_advantages(m): - inputs, old_logps, adv = make_input_features(batch_size=4, extreme_advantages=1e30) - _, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, 'TEST-5-EXTREME') - if elapsed >= TIMEOUT: - return False - do_optim_step(m, 'TEST-5-EXTREME') - return True - -def test_6_all_labels_masked(m): - inputs, old_logps, adv = make_input_features(batch_size=4, all_labels_masked=True) - _, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, 'TEST-6-MASKED') - if elapsed >= TIMEOUT: - return False - do_optim_step(m, 'TEST-6-MASKED') - return True - -def test_7_consecutive_bad(m): - for i in range(5): - inputs, old_logps, adv = make_input_features(batch_size=4, bad_old_logps_len=i+1) - _, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, f'TEST-7-{i+1}') - if elapsed >= TIMEOUT: - return False - do_optim_step(m, f'TEST-7-{i+1}') - return True - -def test_8_rapid_bad_good(m): - for i in range(5): - bad_in, bad_lp, bad_adv = make_input_features(batch_size=4, bad_old_logps_len=i+1) - _, _, elapsed = run_forward_backward(m, bad_in, bad_lp, bad_adv, f'TEST-8-BAD-{i+1}') - if elapsed >= TIMEOUT: - return False - do_optim_step(m, f'TEST-8-BAD-{i+1}') - good_in, good_lp, good_adv = make_input_features(batch_size=4) - ok, _, elapsed = run_forward_backward(m, good_in, good_lp, good_adv, f'TEST-8-GOOD-{i+1}') - if not ok or elapsed >= TIMEOUT: - return False - do_optim_step(m, f'TEST-8-GOOD-{i+1}') - return True - -def test_9_final_health(m): - inputs, old_logps, adv = make_input_features(batch_size=4) - ok, _, elapsed = run_forward_backward(m, inputs, old_logps, adv, 'TEST-9-FINAL') - assert ok and elapsed < TIMEOUT - do_optim_step(m, 'TEST-9-FINAL') - return True - -def test_10_gradient_accumulation_error(m): - inputs, lp, adv = make_input_features(batch_size=4) - ok, _, elapsed = run_forward_backward(m, inputs, lp, adv, 'TEST-10-GA1') - if not ok or elapsed >= TIMEOUT: - return False - bad_in, bad_lp, bad_adv = make_input_features(batch_size=4, bad_old_logps_len=3) - _, _, elapsed = run_forward_backward(m, bad_in, bad_lp, bad_adv, 'TEST-10-GA2-BAD') - if elapsed >= TIMEOUT: - return False - inputs, lp, adv = make_input_features(batch_size=4) - ok, _, elapsed = run_forward_backward(m, inputs, lp, adv, 'TEST-10-GA3') - if not ok or elapsed >= TIMEOUT: - return False - do_optim_step(m, 'TEST-10-GA') - return True - -def test_11_forward_only_then_train(m): - inputs, _, _ = make_input_features(batch_size=4, include_advantages=False) + bad_features, bad_old_logps, bad_adv = _make_inputs(bad_old_logps_len=5) start = time.time() - try: - m.forward_only(inputs=inputs) - except Exception: - if time.time() - start >= TIMEOUT: - return False - train_in, lp, adv = make_input_features(batch_size=4) - ok, _, elapsed = run_forward_backward(m, train_in, lp, adv, 'TEST-11-TRAIN') - if not ok or elapsed >= TIMEOUT: - return False - do_optim_step(m, 'TEST-11-TRAIN') - return True - -def test_12_mixed_seq_lengths(m): - all_inputs, all_lp, all_adv = [], [], [] - for sl, cl in [(32, 16), (128, 64), (48, 24), (96, 48)]: - feats, lp, adv = make_input_features(batch_size=1, seq_len=sl, completion_len=cl) - all_inputs.extend(feats) - if lp: - all_lp.extend(lp) - if adv: - all_adv.extend(adv) - _, _, elapsed = run_forward_backward(m, all_inputs, all_lp, all_adv, 'TEST-12-MIXED') - if elapsed >= TIMEOUT: - return False - do_optim_step(m, 'TEST-12-MIXED') - return True - - -ALL_TESTS = [ - ('TEST-1: Normal GRPO Training', test_1_normal_grpo), - ('TEST-2: Bad old_logps (original bug)', test_2_bad_old_logps), - ('TEST-3: Recovery after error', test_3_recovery), - ('TEST-4: NaN old_logps', test_4_nan_old_logps), - ('TEST-5: Extreme advantages (1e30)', test_5_extreme_advantages), - ('TEST-6: All labels masked (-100)', test_6_all_labels_masked), - ('TEST-7: Consecutive bad batches', test_7_consecutive_bad), - ('TEST-8: Rapid bad->good', test_8_rapid_bad_good), - ('TEST-9: Final health check', test_9_final_health), - ('TEST-10: Gradient accumulation error', test_10_gradient_accumulation_error), - ('TEST-11: forward_only then train', test_11_forward_only_then_train), - ('TEST-12: Mixed sequence lengths', test_12_mixed_seq_lengths), -] - - -def main(): - log('=' * 60) - log('NCCL-Safe E2E Test - Twinkle Client Path') - log('=' * 60) - log(f'Server URL: {SERVER_URL}') - log(f'Base Model: {BASE_MODEL}') - log(f'TWINKLE_FAIL_FAST = {os.getenv("TWINKLE_FAIL_FAST", "1 (default)")}') - - wait_for_server(SERVER_URL) - m = init_client() - - results = [] - for name, test_fn in ALL_TESTS: - log(f'\n{"=" * 60}\n{name}\n{"=" * 60}') - try: - passed = test_fn(m) - results.append((name, 'PASS' if passed else 'FAIL')) - log(f'[{name}] {"PASS" if passed else "FAIL"}') - except Exception as e: - log(f'{name}: EXCEPTION: {e}') - traceback.print_exc() - results.append((name, 'FAIL')) - - log(f'\n{"=" * 60}\nRESULTS SUMMARY\n{"=" * 60}') - all_passed = all(s == 'PASS' for _, s in results) - for name, status in results: - log(f' [{status}] {name}') - log(f'\n{"ALL" if all_passed else "SOME"} {len(results)} TESTS {"PASSED" if all_passed else "FAILED"}!') - return 0 if all_passed else 1 - - -def test_nccl_safe_twinkle_e2e(): - """Pytest-collected entry point.""" - rc = main() - assert rc == 0, 'Some Twinkle NCCL-safe E2E tests failed' - - -if __name__ == '__main__': - sys.exit(main()) + with pytest.raises(Exception) as caught: + model.forward_backward( + inputs=bad_features, adapter_name=ADAPTER_NAME, old_logps=bad_old_logps, advantages=bad_adv) + message = str(caught.value) + # The failure must be loud and descriptive (not a silent zero-loss success): + # the deliberate old_logps/completion length mismatch surfaces on both backends. + assert 'mismatch' in message, message + # Megatron additionally attributes the failure to a global rank via nccl_safe_megatron. + if BACKEND == 'megatron': + assert 'global_rank=' in message, message + assert time.time() - start < TIMEOUT, 'malformed request must fail fast, not hang (NCCL)' + + good_features, good_old_logps, good_adv = _make_inputs() + result = model.forward_backward( + inputs=good_features, adapter_name=ADAPTER_NAME, old_logps=good_old_logps, advantages=good_adv) + assert result is not None diff --git a/tests/server/integration/test_sft_e2e.py b/tests/server/integration/test_sft_e2e.py index 40c794b11..16a014b8c 100644 --- a/tests/server/integration/test_sft_e2e.py +++ b/tests/server/integration/test_sft_e2e.py @@ -5,6 +5,12 @@ - Twinkle client x (transformers | megatron) - Tinker client x (transformers | megatron) +Each test also verifies save-LoRA + resume-training succeeds: after the training +loop it saves a checkpoint (Twinkle ``model.save``; Tinker ``save_state``), resumes +(Twinkle ``resume_from_checkpoint``; Tinker +``create_training_client_from_state_with_optimizer``), and runs a few more steps +that must complete without timeout. + Backend selection via env var TWINKLE_TEST_BACKEND (default: transformers). ## How to run @@ -45,6 +51,7 @@ create_tinker_training_client, create_twinkle_sft_model, get_backend, + init_tinker_client_session, init_twinkle_client_session, log, wait_for_server, @@ -52,6 +59,7 @@ # ── Configuration ── SFT_TRAIN_STEPS = 20 # 20 steps ensures enough training for both backends +SFT_RESUME_STEPS = 3 # post-resume steps that must run without timeout # ═══════════════════════════════════════════════════════════════════════════ @@ -107,7 +115,32 @@ def test_sft_twinkle(): # Assertions — both backends should report real loss via calculate_metric assert len(losses) >= 4, f'Expected at least 4 logged losses, got {len(losses)}' assert_loss_decreases(losses, 'sft_twinkle') - log(f'test_sft_twinkle PASSED (backend={backend})') + + # ── Save LoRA + resume training (must succeed) ── + save_resp = model.save( + name='sft-twinkle-resume', + save_optimizer=True, + consumed_train_samples=dataloader.get_state()['consumed_train_samples'], + ) + ckpt = save_resp.twinkle_path + assert ckpt, 'save() did not return a twinkle_path' + log(f'saved LoRA checkpoint: {ckpt}') + + progress = model.resume_from_checkpoint(ckpt) + log(f'resumed from checkpoint: {progress}') + + resume_loader = DataLoader(dataset=create_sft_dataset(), batch_size=4) + resumed = 0 + for step, batch in enumerate(resume_loader): + if step >= SFT_RESUME_STEPS: + break + t0 = time.time() + model.forward_backward(inputs=batch) + model.clip_grad_and_step() + assert_no_timeout(time.time() - t0, f'sft_twinkle resume step {step}') + resumed += 1 + assert resumed == SFT_RESUME_STEPS, f'expected {SFT_RESUME_STEPS} post-resume steps, ran {resumed}' + log(f'test_sft_twinkle PASSED (backend={backend}) [+save LoRA +resume]') # ═══════════════════════════════════════════════════════════════════════════ @@ -123,7 +156,7 @@ def test_sft_tinker(): """ from tinker import types from twinkle.dataloader import DataLoader - from twinkle.server.common import input_feature_to_datum + from twinkle.server.model.tinker_datum import input_feature_to_datum backend = get_backend() log(f'=== test_sft_tinker [backend={backend}] ===') @@ -171,7 +204,31 @@ def test_sft_tinker(): # Assertions assert len(losses) >= 4, f'Expected at least 4 logged losses, got {len(losses)}' assert_loss_decreases(losses, 'sft_tinker') - log(f'test_sft_tinker PASSED (backend={backend})') + + # ── Save state + resume training (must succeed) ── + save_result = training_client.save_state('sft-tinker-resume').result() + state_path = save_result.path + assert state_path, 'save_state() did not return a path' + log(f'saved tinker state: {state_path}') + + # Resume restores both weights and optimizer state into a fresh client. + service_client = init_tinker_client_session() + resumed_client = service_client.create_training_client_from_state_with_optimizer(path=state_path) + log('resumed tinker training client from saved state') + + resume_loader = DataLoader(dataset=create_sft_dataset(), batch_size=4) + resumed = 0 + for step, batch in enumerate(resume_loader): + if step >= SFT_RESUME_STEPS: + break + input_datums = [input_feature_to_datum(input_feature) for input_feature in batch] + t0 = time.time() + resumed_client.forward_backward(input_datums, 'cross_entropy').result() + resumed_client.optim_step(types.AdamParams(learning_rate=1e-4)).result() + assert_no_timeout(time.time() - t0, f'sft_tinker resume step {step}') + resumed += 1 + assert resumed == SFT_RESUME_STEPS, f'expected {SFT_RESUME_STEPS} post-resume steps, ran {resumed}' + log(f'test_sft_tinker PASSED (backend={backend}) [+save state +resume]') # ── Direct execution ── diff --git a/tests/server/lifecycle/__init__.py b/tests/server/lifecycle/__init__.py new file mode 100644 index 000000000..85b3e739d --- /dev/null +++ b/tests/server/lifecycle/__init__.py @@ -0,0 +1 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. diff --git a/tests/server/lifecycle/test_envelope.py b/tests/server/lifecycle/test_envelope.py new file mode 100644 index 000000000..fdc588bd8 --- /dev/null +++ b/tests/server/lifecycle/test_envelope.py @@ -0,0 +1,78 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Unit tests for the single FutureRecord -> TaskEnvelope mapping point.""" +from __future__ import annotations + +import pytest +from pydantic import ValidationError + +from twinkle.server.lifecycle.envelope import envelope_from_record + + +def test_wire_maps_cover_exactly_the_canonical_reason_codes(): + """The Twinkle and Tinker failure maps must both stay in lockstep with the + canonical reason-code set, so a newly added domain reason cannot silently + fall back to 500 on one protocol face.""" + from twinkle.server.gateway.tinker_handlers import _TINKER_FAILURE_WIRE + from twinkle.server.lifecycle.envelope import _FAILURE_WIRE + from twinkle.server.state.models import FAILURE_REASON_CODES + + assert set(_FAILURE_WIRE) == FAILURE_REASON_CODES + assert set(_TINKER_FAILURE_WIRE) == FAILURE_REASON_CODES + + +def test_completed_with_none_result_is_a_success_not_a_failure(): + """`completed` + `result is None` is a valid success.""" + env = envelope_from_record('req-1', {'status': 'completed', 'result': None}) + assert env.status == 'completed' + assert env.result is None + assert env.error is None + + +def test_completed_carries_result_and_no_error(): + env = envelope_from_record('req-1', {'status': 'completed', 'result': {'loss': 0.5}}) + assert env.result == {'loss': 0.5} + assert env.error is None + + +def test_domain_failure_maps_to_twinkle_error_payload(): + env = envelope_from_record( + 'req-9', { + 'status': 'failed', + 'failure': { + 'reason_code': 'execution_timeout', + 'message': 'backend timed out', + 'attribution': 'server', + 'diagnostic': 'full traceback', + }, + }) + assert env.status == 'failed' + assert env.result is None + assert env.error is not None + assert env.error.error == 'backend timed out' + assert env.error.category.value == 'server' + assert env.error.error_code == 504 + assert env.error.request_id == 'req-9' + assert env.error.traceback == 'full traceback' + + +def test_legacy_failure_in_result_is_not_accepted(): + with pytest.raises(ValidationError): + envelope_from_record( + 'req-old', {'status': 'failed', 'result': {'error': 'boom', 'category': 'server'}}) + + +def test_non_terminal_record_carries_queue_state_and_no_payload(): + env = envelope_from_record( + 'req-3', {'status': 'running', 'queue_state': 'active', 'queue_state_reason': 'x'}) + assert env.status == 'running' + assert env.result is None + assert env.error is None + assert env.queue_state == 'active' + assert env.queue_state_reason == 'x' + + +def test_missing_record_falls_back_to_pending(): + env = envelope_from_record('req-4', None) + assert env.status == 'pending' + assert env.result is None + assert env.error is None diff --git a/tests/server/lifecycle/test_envelope_coverage.py b/tests/server/lifecycle/test_envelope_coverage.py new file mode 100644 index 000000000..039327cb1 --- /dev/null +++ b/tests/server/lifecycle/test_envelope_coverage.py @@ -0,0 +1,66 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Task_Envelope coverage check. + +Walks the model and sampler route tables and asserts that every twinkle-native +POST route that enters the Task_Queue declares ``response_model = TaskEnvelope``. +The exemption list (endpoints that do NOT enter the queue, plus the streaming +endpoint) is declared here, in one place. +""" +from __future__ import annotations + +from fastapi.routing import APIRoute + +from tests.server.contract.client_api_harness import build_model_app, build_sampler_app +from twinkle.protocol.types.lifecycle import TaskEnvelope + +# The single exemption declaration, keyed BY APP. A flat path set would be wrong: +# ``/twinkle/set_template`` and ``/twinkle/apply_patch`` exist on both apps, but only the +# sampler's bypass the queue -- the model's are queued and must return a Task_Envelope. +# Sharing one set silently exempted the model's two and left a hole in this guard. +_EXEMPT_BY_APP = { + 'model': { + # health/session bootstrap only, no queue + '/twinkle/create', + }, + 'sampler': { + # direct call_backend, no queue + '/twinkle/create', + '/twinkle/set_template', + '/twinkle/add_adapter_to_sampler', + '/twinkle/apply_patch', + '/twinkle/unload_adapter_paths', + # the one streaming exception + '/twinkle/sample_stream', + }, +} + + +def _queued_twinkle_post_routes(app, exempt): + for route in app.routes: + if not isinstance(route, APIRoute): + continue + if 'POST' not in route.methods: + continue + if not route.path.startswith('/twinkle/'): + continue + if route.path in exempt: + continue + yield route + + +def test_every_queued_twinkle_route_returns_task_envelope(): + violations = [] + for app_name, app in (('model', build_model_app()), ('sampler', build_sampler_app())): + for route in _queued_twinkle_post_routes(app, _EXEMPT_BY_APP[app_name]): + if route.response_model is not TaskEnvelope: + violations.append((app_name, route.path, route.response_model)) + assert violations == [], f'queued routes not returning TaskEnvelope: {violations}' + + +def test_model_side_set_template_and_apply_patch_are_not_exempt(): + # Regression guard for the hole above: these two are queued on the model app, so + # they must be covered by the assertion rather than skipped by a shared path set. + assert '/twinkle/set_template' not in _EXEMPT_BY_APP['model'] + assert '/twinkle/apply_patch' not in _EXEMPT_BY_APP['model'] + covered = {route.path for route in _queued_twinkle_post_routes(build_model_app(), _EXEMPT_BY_APP['model'])} + assert {'/twinkle/set_template', '/twinkle/apply_patch'} <= covered diff --git a/tests/server/lifecycle/test_preflight_rejection.py b/tests/server/lifecycle/test_preflight_rejection.py new file mode 100644 index 000000000..74639c137 --- /dev/null +++ b/tests/server/lifecycle/test_preflight_rejection.py @@ -0,0 +1,132 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Decision_Boundary tests: preflight rejects with real status codes and zero writes. + +Covers the case where a rejected request writes no future record, plus the +TwinkleServerError handler wire shape. No Ray or GPU is involved: the +task queue is driven with a spy state that counts ``store_future_status`` calls. +""" +from __future__ import annotations + +import pytest +from fastapi import FastAPI, Request +from fastapi.testclient import TestClient + +from twinkle.server.deployment import twinkle_server_error_handler +from twinkle.server.exceptions import (BatchSizeError, InputTokensExceededError, RateLimitExceededError, + RequestRejectedError, TwinkleServerError) +from twinkle.server.task_queue.config import TaskQueueConfig +from twinkle.server.task_queue.mixin import TaskQueueMixin + + +class _SpyState: + """Counts every future write so a rejected request can be proven to write nothing.""" + + def __init__(self): + self.store_calls = 0 + + async def store_future_status(self, *args, **kwargs): + self.store_calls += 1 + + async def get_future(self, request_id): + return None + + +class _Harness(TaskQueueMixin): + + def __init__(self, **config_kwargs): + self.state = _SpyState() + self.replica_id = 'test-replica' + self._init_task_queue(TaskQueueConfig(**config_kwargs), deployment_name='test') + + +async def _noop(): + return None + + +@pytest.mark.asyncio +async def test_input_tokens_rejection_is_422_and_zero_writes(): + """An over-limit request raises 422 and writes no record.""" + h = _Harness(enabled=True, max_input_tokens=10) + try: + with pytest.raises(InputTokensExceededError) as exc: + await h.schedule_task(lambda: _noop(), model_id='m', token='tok', input_tokens=999, task_type='forward') + assert exc.value.error_code == 422 + assert exc.value.category.value == 'user' + assert h.state.store_calls == 0 + finally: + await h.shutdown_task_queue() + + +@pytest.mark.asyncio +async def test_batch_size_rejection_is_422_and_zero_writes(): + h = _Harness(enabled=True, max_input_tokens=100000) + try: + with pytest.raises(BatchSizeError): + await h.schedule_task( + lambda: _noop(), model_id='m', token='tok', input_tokens=1, + batch_size=1, data_world_size=4, task_type='forward') + assert h.state.store_calls == 0 + finally: + await h.shutdown_task_queue() + + +@pytest.mark.asyncio +async def test_rate_limit_rejection_is_429_and_zero_writes(): + h = _Harness(enabled=True, rps_limit=1, tps_limit=1000000, window_seconds=100, max_input_tokens=100000) + try: + # First call is admitted (it enqueues -> writes); reset the counter and + # assert the rejected second call (same window, rps=1) writes nothing. + await h.schedule_task(lambda: _noop(), model_id='m', token='tok', input_tokens=1, task_type='forward') + h.state.store_calls = 0 + with pytest.raises(RateLimitExceededError) as exc: + await h.schedule_task(lambda: _noop(), model_id='m', token='tok', input_tokens=1, task_type='forward') + assert exc.value.error_code == 429 + assert h.state.store_calls == 0 + finally: + await h.shutdown_task_queue() + + +@pytest.mark.asyncio +async def test_disabled_queue_skips_preflight(): + """The 'no token or queue disabled' short circuit is preserved.""" + h = _Harness(enabled=False, max_input_tokens=10) + try: + ref = await h.schedule_task(lambda: _noop(), model_id='m', token='tok', input_tokens=999, task_type='forward') + assert 'request_id' in ref # not rejected: enqueued normally + finally: + await h.shutdown_task_queue() + + +def test_error_handler_puts_fields_at_top_level(): + """The handler returns error_code as the status and fields at top level.""" + app = FastAPI() + app.add_exception_handler(TwinkleServerError, twinkle_server_error_handler) + + @app.get('/boom') + async def boom(request: Request): + raise RequestRejectedError('nope', error_code=409) + + resp = TestClient(app, raise_server_exceptions=False).get('/boom') + assert resp.status_code == 409 + body = resp.json() + assert 'detail' not in body # not nested under detail + assert body['error'] == 'nope' + assert body['category'] == 'user' + assert body['error_code'] == 409 + + +def test_error_handler_bounds_overlong_domain_error(): + app = FastAPI() + app.add_exception_handler(TwinkleServerError, twinkle_server_error_handler) + + @app.get('/boom') + async def boom(request: Request): + raise RequestRejectedError(f'bad request {"X" * 2048}\ninternal detail', error_code=409) + + response = TestClient(app, raise_server_exceptions=False).get('/boom') + assert response.status_code == 409 + body = response.json() + assert len(body['error']) == 1024 + assert '\n' not in body['error'] + assert body['category'] == 'user' + assert 'traceback' not in body diff --git a/tests/server/lifecycle/test_protocols.py b/tests/server/lifecycle/test_protocols.py new file mode 100644 index 000000000..4f670a5be --- /dev/null +++ b/tests/server/lifecycle/test_protocols.py @@ -0,0 +1,26 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""``run_submit`` must stay a free function. + +Reverting it to a decorator would rewrite its signature and deepen the route graph, which +once hit the CPython C-stack recursion limit in ``serve.ingress``'s cloudpickle phase. +Adding a ``self: QueuedDeployment`` annotation is a zero-runtime-cost change, so +the function form must be unchanged. +""" +import inspect + +from twinkle.server.lifecycle.protocols import DataParallelDeployment, QueuedDeployment +from twinkle.server.lifecycle.submit import input_metrics, run_submit + + +def test_run_submit_is_a_free_function(): + assert inspect.isfunction(run_submit) + assert inspect.iscoroutinefunction(run_submit) + assert inspect.isfunction(input_metrics) + + +def test_host_protocols_are_two_layers(): + # DataParallelDeployment is the strictly stronger contract (adds data_world_size). + # (``issubclass`` is avoided: runtime_checkable Protocols with data members raise.) + assert QueuedDeployment in DataParallelDeployment.__mro__ + assert hasattr(DataParallelDeployment, 'data_world_size') + assert not hasattr(QueuedDeployment, 'data_world_size') diff --git a/tests/server/lifecycle/test_retrieve_endpoint.py b/tests/server/lifecycle/test_retrieve_endpoint.py new file mode 100644 index 000000000..84b8767e4 --- /dev/null +++ b/tests/server/lifecycle/test_retrieve_endpoint.py @@ -0,0 +1,97 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Wire tests for the twinkle Retrieve_Endpoint. + +These use a fake state and FastAPI's TestClient; no Ray runtime is needed, so they +live outside the state-actor fixtures. +""" +from __future__ import annotations + +import time + +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from twinkle.server.deployment import twinkle_server_error_handler +from twinkle.server.exceptions import TwinkleServerError +from twinkle.server.gateway.twinkle_handlers import _register_gateway_twinkle_routes + + +class _State: + """A fake ServerState whose get_future returns a fixed record (or None).""" + + def __init__(self, record): + self._record = record + + async def get_future(self, request_id: str): + return self._record + + +class _Gateway: + + def __init__(self, record): + self.state = _State(record) + + +def _client(record) -> TestClient: + app = FastAPI() + app.add_exception_handler(TwinkleServerError, twinkle_server_error_handler) + _register_gateway_twinkle_routes(app, lambda: _Gateway(record)) + return TestClient(app) + + +def test_completed_with_null_result_returns_200_and_null(monkeypatch): + """Completed + result=None is 200 with result null, not 500.""" + client = _client({'status': 'completed', 'result': None}) + resp = client.post('/twinkle/retrieve_future', json={'request_id': 'req-1'}) + assert resp.status_code == 200 + body = resp.json() + assert body['status'] == 'completed' + assert body['result'] is None + assert body['error'] is None + + +def test_domain_failure_returns_200_and_valid_envelope(): + client = _client({ + 'status': 'failed', + 'failure': { + 'reason_code': 'internal_error', + 'message': 'boom', + 'attribution': 'server', + }, + }) + resp = client.post('/twinkle/retrieve_future', json={'request_id': 'req-2'}) + assert resp.status_code == 200 + body = resp.json() + assert body['status'] == 'failed' + assert body['result'] is None + assert body['error']['error'] == 'boom' + assert body['error']['category'] == 'server' + assert body['error']['error_code'] == 500 + assert body['error']['request_id'] == 'req-2' + + +def test_always_missing_record_404s_only_after_the_full_window(monkeypatch): + """A request_id that never appears returns 404, and only after waiting a window.""" + monkeypatch.setenv('TWINKLE_LONG_POLL_TIMEOUT', '0.3') + client = _client(None) + start = time.monotonic() + resp = client.post('/twinkle/retrieve_future', json={'request_id': 'ghost'}) + waited = time.monotonic() - start + assert resp.status_code == 404 + assert 'ghost' in resp.json()['error'] + assert resp.json()['category'] == 'user' + assert resp.json()['error_code'] == 404 + # It must fold the missing record into the wait loop, not short-circuit. + assert waited >= 0.3 + + +def test_terminal_record_returns_immediately(monkeypatch): + """A record already terminal must not wait out the window.""" + monkeypatch.setenv('TWINKLE_LONG_POLL_TIMEOUT', '30') + client = _client({'status': 'completed', 'result': {'ok': True}}) + start = time.monotonic() + resp = client.post('/twinkle/retrieve_future', json={'request_id': 'req-5'}) + waited = time.monotonic() - start + assert resp.status_code == 200 + assert resp.json()['result'] == {'ok': True} + assert waited < 5.0 diff --git a/tests/server/lifecycle/test_run_submit_dedup.py b/tests/server/lifecycle/test_run_submit_dedup.py new file mode 100644 index 000000000..f9e91d76c --- /dev/null +++ b/tests/server/lifecycle/test_run_submit_dedup.py @@ -0,0 +1,111 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""run_submit seq_id dedup: the release-on-failure decision must be driven by whether +a future record exists (i.e. whether the task was enqueued), NOT by exception type. + +The load-bearing case: if submit_and_peek raises *after* the task was enqueued (e.g. a +transient state error inside the inline peek), the seq claim must be KEPT -- releasing +it would let a retry enqueue a duplicate, the exact double-apply the dedup prevents. +""" +from __future__ import annotations + +from types import SimpleNamespace + +import pytest + +from twinkle.server.lifecycle.submit import run_submit +from twinkle.protocol.types.model import ForwardBackwardTaskRequest + + +class _FakeState: + def __init__(self, record_after_claim): + self._record_after_claim = record_after_claim + self.claimed = {} + self.released = [] + + async def claim_seq(self, dedup_key, request_id, ttl): + # Unseen -> claim it and let the caller proceed (returns None). + self.claimed[dedup_key] = request_id + return None + + async def get_future(self, request_id): + # Simulates whether a PENDING/QUEUED record was written (task enqueued). + return self._record_after_claim + + async def release_seq(self, dedup_key): + self.released.append(dedup_key) + + +class _FakeManagement: + def __init__(self, record_after_claim): + self.state = _FakeState(record_after_claim) + self.task_queue_config = SimpleNamespace(effective_execution_timeout=60.0) + # A real deployment declares its backend; preflight reads it from here. + self.backend = 'transformers' + self.data_world_size = 1 + + async def _on_request_start(self, request): + return 'token' + + async def submit_and_peek(self, *args, **kwargs): + # Fail *after* the (simulated) enqueue -- e.g. a state blip during the peek. + raise RuntimeError('transient state error during peek') + + +def _request(): + return SimpleNamespace(state=SimpleNamespace(session_id='sess-1', request_id='rq-1')) + + +def _body(adapter_name: str = 'ad', seq_id: int = 7) -> ForwardBackwardTaskRequest: + """A real request model, not a stand-in. + + ``run_submit`` now reads field roles off the body to build the backend kwargs and to + run preflight, so a ``SimpleNamespace`` would exercise a shape production never sees. + """ + return ForwardBackwardTaskRequest(inputs=[{'input_ids': [1, 2]}], adapter_name=adapter_name, seq_id=seq_id) + + +async def _call(self, body, adapter_name, token): # pragma: no cover - never invoked + return {'ok': True} + + +@pytest.mark.asyncio +async def test_release_kept_when_task_already_enqueued(): + # A record exists (task enqueued) -> peek error must NOT release the claim. + mgmt = _FakeManagement(record_after_claim={'status': 'queued'}) + with pytest.raises(RuntimeError): + await run_submit(mgmt, _request(), _body(), task_type='forward_backward', backend_call=_call) + assert mgmt.state.released == [], 'claim wrongly released for an already-enqueued task' + assert 'seq::sess-1::sess-1-ad::7' in mgmt.state.claimed + + +@pytest.mark.asyncio +async def test_release_when_never_enqueued(): + # No record (e.g. preflight rejected before any write) -> release so a retry can re-enqueue. + mgmt = _FakeManagement(record_after_claim=None) + with pytest.raises(RuntimeError): + await run_submit(mgmt, _request(), _body(), task_type='forward_backward', backend_call=_call) + assert mgmt.state.released == ['seq::sess-1::sess-1-ad::7'], 'claim should be released when nothing enqueued' + + +@pytest.mark.asyncio +async def test_dedup_key_is_scoped_per_adapter(): + """Two adapters in ONE session must not collide on the same seq_id. + + Every client model object owns its own seq counter starting at 1 while ``session_id`` + is process-global, so multi-LoRA training from one process issues seq_id=1 twice. If + the adapter were missing from the key, the second adapter's forward_backward would be + swallowed as a duplicate and handed the first adapter's loss -- a silent wrong result. + """ + mgmt = _FakeManagement(record_after_claim={'status': 'queued'}) + for adapter in ('lora-A', 'lora-B'): + with pytest.raises(RuntimeError): + await run_submit( + mgmt, + _request(), + _body(adapter_name=adapter, seq_id=1), + task_type='forward_backward', + backend_call=_call) + + claimed = set(mgmt.state.claimed) + assert claimed == {'seq::sess-1::sess-1-lora-A::1', 'seq::sess-1::sess-1-lora-B::1'}, ( + f'adapters collided on one dedup key: {claimed}') diff --git a/tests/server/lifecycle/test_static_guards.py b/tests/server/lifecycle/test_static_guards.py new file mode 100644 index 000000000..730c4ef7f --- /dev/null +++ b/tests/server/lifecycle/test_static_guards.py @@ -0,0 +1,111 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Static / structural guards for the lifecycle refactor. + +- the deleted symbols occur zero times under ``src/twinkle/**``. +- ``TaskEnvelope`` has exactly one construction site. +- The client-side invariant: the client HTTP timeout is <= 120 and strictly + greater than the server Long_Poll_Window. +- The task status set has two independent declarations that must not drift. + +These exist because the spec states most of its guarantees in prose. A prose claim that +nothing checks decays into a false claim -- as happened with "a consistency test asserts +the two sets are equal", which was written in a docstring while no such test existed. +""" +from __future__ import annotations + +import pytest +from pathlib import Path + +_REPO_ROOT = Path(__file__).resolve().parents[3] +_SRC = _REPO_ROOT / 'src' / 'twinkle' + +# Symbols the refactor removed. A wildcard search (not a per-file list) must find +# each of them zero times across the whole server tree. +_FORBIDDEN_SYMBOLS = ( + 'schedule_task_and_wait', + 'run_task', + 'persist_status', + '_complete_result', + '_complete_error', +) + + +@pytest.mark.parametrize('symbol', _FORBIDDEN_SYMBOLS) +def test_deleted_symbol_has_zero_occurrences(symbol): + hits = [] + for path in _SRC.rglob('*.py'): + text = path.read_text(encoding='utf-8') + if symbol in text: + hits.append(str(path.relative_to(_REPO_ROOT))) + assert hits == [], f'{symbol!r} still occurs in: {hits}' + + +def test_client_http_timeout_bounds(): + from twinkle.server.lifecycle.poll_config import long_poll_window + from twinkle_client.http.client import _HTTP_TIMEOUT + + assert _HTTP_TIMEOUT <= 120 + assert _HTTP_TIMEOUT > long_poll_window() + + +def test_task_envelope_has_exactly_one_construction_site(): + """'s structural precondition: one mapping point, mechanically enforced. + + ``envelope_from_record`` is the only place a FutureRecord becomes a TaskEnvelope, so + ``failed`` always lands in ``error`` and never in ``result`` regardless of which + endpoint answered. That was previously only a docstring claim -- a handler building a + ``TaskEnvelope(...)`` itself would silently reintroduce the exact defect the single + mapping point exists to prevent (a failure inside the Inline_Fast_Path window losing + its payload), and every existing test would still pass. + """ + sites = [] + for path in _SRC.rglob('*.py'): + text = path.read_text(encoding='utf-8') + for lineno, line in enumerate(text.splitlines(), start=1): + if 'TaskEnvelope(' in line and 'class TaskEnvelope' not in line: + sites.append(f'{path.relative_to(_REPO_ROOT)}:{lineno}') + + offenders = [site for site in sites if 'server/lifecycle/envelope.py' not in site] + assert offenders == [], ('TaskEnvelope must only be constructed in lifecycle/envelope.py ' + f'(via envelope_from_record); found: {offenders}') + assert sites, 'expected to find the construction sites inside envelope.py' + + +def test_server_task_status_enum_matches_client_literal(): + """The two independent declarations of the task status set must not drift. + + ``twinkle.protocol.types.lifecycle.TaskStatus`` (a Literal on the wire model) and the + server's ``TaskStatus`` enum are declared separately. ``envelope_from_record`` copies + ``record['status']`` straight into ``TaskEnvelope.status``, so a value the server can + write but the Literal does not list would fail pydantic validation *while serialising + the response* -- i.e. a 500 from retrieve for a task that actually finished. + + The client module's comment claimed such a test existed; it did not. This is it. + """ + from typing import get_args + + from twinkle.server.task_queue.types import TaskStatus as ServerTaskStatus + from twinkle.protocol.types.lifecycle import TERMINAL_STATUSES + from twinkle.protocol.types.lifecycle import TaskStatus as WireTaskStatus + + server_values = {member.value for member in ServerTaskStatus} + wire_values = set(get_args(WireTaskStatus)) + assert server_values == wire_values, (f'task status sets drifted: server-only={server_values - wire_values}, ' + f'wire-only={wire_values - server_values}') + assert TERMINAL_STATUSES <= wire_values, 'TERMINAL_STATUSES must be a subset of the declared statuses' + + +def test_client_future_layer_is_not_imported_by_the_server(): + """``_future.py`` carries an underscore because the dependency runs one way only. + + The server reverse-imports ``twinkle.protocol.types`` (the shared wire contract), but the + client's polling layer is private to the client. An import in the other direction would + make the server depend on client retry policy, which its own long-poll already owns. + Another claim that lived only in a docstring. + """ + offenders = [] + for path in _SRC.rglob('*.py'): + text = path.read_text(encoding='utf-8') + if 'twinkle_client._future' in text or 'from twinkle_client import _future' in text: + offenders.append(str(path.relative_to(_REPO_ROOT))) + assert offenders == [], f'server must not import the client future layer: {offenders}' diff --git a/tests/server/lifecycle/test_submit_peek_e2e.py b/tests/server/lifecycle/test_submit_peek_e2e.py new file mode 100644 index 000000000..c87aba599 --- /dev/null +++ b/tests/server/lifecycle/test_submit_peek_e2e.py @@ -0,0 +1,82 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""End-to-end proof of the Inline_Fast_Path + Client_Future_Layer seam. + +A minimal harness drives the *real* ``submit_and_peek`` against a real compute +worker and a real (memory) ServerState -- no HTTP, no GPU. The resulting envelope +is round-tripped through model_dump/model_validate (simulating the wire) and fed to +the real client ``resolve``, so this covers the exact submit -> client path. +""" +from __future__ import annotations + +import pytest + +ray = pytest.importorskip('ray') + +from twinkle.server.state import ServerState # noqa: E402 +from twinkle.server.task_queue.config import TaskQueueConfig # noqa: E402 +from twinkle.server.task_queue.mixin import TaskQueueMixin # noqa: E402 +from twinkle_client import _future # noqa: E402 +from twinkle_client.exceptions import TaskFailedError # noqa: E402 +from twinkle.protocol.types.lifecycle import TaskEnvelope # noqa: E402 + + +class _Harness(TaskQueueMixin): + """Real task queue + real state, with a window generous enough to be deterministic.""" + + def __init__(self) -> None: + self.state = ServerState() + self.replica_id = 'test-replica' + # enabled=False skips rate limiting; a 5s window makes "task finishes inside + # the window" deterministic for a trivial in-process coroutine. + self._init_task_queue( + TaskQueueConfig(enabled=False, inline_fast_path_timeout=5.0), deployment_name='test') + + +def _across_the_wire(env: TaskEnvelope) -> TaskEnvelope: + return TaskEnvelope.model_validate(env.model_dump(mode='json')) + + +@pytest.mark.asyncio +async def test_window_completed_task_is_single_round_trip(monkeypatch): + """A task terminal within the window makes the client issue zero retrieves.""" + h = _Harness() + + async def _ok(): + return None + + try: + env = await h.submit_and_peek(lambda: _ok(), task_type='step') + assert env.status == 'completed' + assert env.result is None + + monkeypatch.setattr(_future, '_post_retrieve', + lambda _r: pytest.fail('completed submit must not poll retrieve')) + assert _future.resolve(_across_the_wire(env), model_cls=None) is None + finally: + await h.shutdown_task_queue() + + +@pytest.mark.asyncio +async def test_window_failed_task_surfaces_payload_as_taskfailed(monkeypatch): + """A failure inside the window reaches the client via the submit + response and is raised as TaskFailedError with its payload intact.""" + h = _Harness() + + async def _boom(): + raise ValueError('kaboom') + + try: + env = await h.submit_and_peek(lambda: _boom(), task_type='step') + assert env.status == 'failed' + assert env.error is not None + assert 'kaboom' in env.error.error + + monkeypatch.setattr(_future, '_post_retrieve', + lambda _r: pytest.fail('failed submit must not poll retrieve')) + with pytest.raises(TaskFailedError) as exc: + _future.resolve(_across_the_wire(env), model_cls=None) + assert 'kaboom' in exc.value.error + assert exc.value.category == 'server' + assert exc.value.error_code == 500 + finally: + await h.shutdown_task_queue() diff --git a/tests/server/lifecycle/test_timing_bounds.py b/tests/server/lifecycle/test_timing_bounds.py new file mode 100644 index 000000000..b2f5f5464 --- /dev/null +++ b/tests/server/lifecycle/test_timing_bounds.py @@ -0,0 +1,139 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Timing guards for the lifecycle constants. + +- a single Submit_Endpoint's server-side duration is bounded by the + Inline_Fast_Path window + 1s and is INDEPENDENT of how long the task itself runs. This + is the spec's core benefit claim and was previously the only property with no automated + guard. +- The retrieve poll interval is a single shared declaration, is strictly + inside the Long_Poll_Window, and is deliberately FIXED (see the measurement recorded in + ``poll_config`` and in the test below). +""" +from __future__ import annotations + +import asyncio +import time +from pathlib import Path + +import pytest + +ray = pytest.importorskip('ray') + +from twinkle.server.lifecycle.poll_config import long_poll_window, retrieve_poll_interval # noqa: E402 +from twinkle.server.state import ServerState # noqa: E402 +from twinkle.server.task_queue.config import TaskQueueConfig # noqa: E402 +from twinkle.server.task_queue.mixin import TaskQueueMixin # noqa: E402 + +_WINDOW = 0.05 + + +class _Harness(TaskQueueMixin): + def __init__(self) -> None: + self.state = ServerState() + self.replica_id = 'test-replica' + self._init_task_queue( + TaskQueueConfig(enabled=False, inline_fast_path_timeout=_WINDOW), deployment_name='test') + + +@pytest.mark.asyncio +async def test_submit_duration_is_bounded_and_task_duration_independent(): + """Submit returns on the window, not on task completion.""" + h = _Harness() + + async def fast(): + return None + + async def slow(): + await asyncio.sleep(5.0) + return {'loss': 1.0} + + try: + # Warm up first: the very first submit pays for creating the detached state actor + # and starting the ComputeWorker (~10s on a cold Ray). That cost is not part of the + # submit path this test is about, and timing it made the test pass or fail depending + # on whether an earlier test in the directory had already warmed the backend. + await h.submit_and_peek(lambda: fast(), task_type='warmup') + + started = time.monotonic() + env = await h.submit_and_peek(lambda: slow(), task_type='forward_backward') + elapsed = time.monotonic() - started + + # Bounded by the window + 1s even though the task needs 5s. + assert elapsed < _WINDOW + 1.0, f'submit took {elapsed:.3f}s, expected < {_WINDOW + 1.0}s' + # 5s task cannot have finished, so the envelope must be non-terminal. + assert env.status not in ('completed', 'failed', 'cancelled'), env.status + assert env.result is None and env.error is None + assert env.request_id + finally: + await h.shutdown_task_queue() + + +def test_poll_interval_satisfies_the_constant_chain(): + interval, window = retrieve_poll_interval(), long_poll_window() + assert interval > 0, 'a non-positive interval would busy-spin the state backend' + assert interval < window, 'a single poll step must not exceed the whole window' + + +def test_both_retrieve_endpoints_share_one_interval_declaration(): + """One declaration point, and no endpoint reading os.environ on its own. + + Also pins the measured decision: a FIXED interval, not exponential backoff. The + backoff variant was implemented, measured on real PPU hardware, and reverted -- + control-plane ops (0.00-0.09s) never reach retrieve because the 50ms inline window + absorbs them, while ``forward_backward`` lands at 0.52-0.65s where a 0.05->1.0s + doubling schedule checks at 0.80 instead of the fixed schedule's 0.55, costing ~22% + per step. If backoff is ever reintroduced, it needs a ceiling around 0.2s and an + explicit decision to pay 2.5x the poll rate. + """ + import twinkle.server.gateway.tinker_handlers as tinker_h + import twinkle.server.gateway.twinkle_handlers as twinkle_h + from twinkle.server.lifecycle import poll_config + + assert not hasattr(poll_config, 'initial_poll_interval'), 'backoff was reverted by measurement' + assert not hasattr(poll_config, 'max_poll_interval'), 'backoff was reverted by measurement' + + for module in (tinker_h, twinkle_h): + source = Path(module.__file__).read_text(encoding='utf-8') + assert 'TWINKLE_POLL_INTERVAL' not in source, f'{module.__name__} reads the env var directly' + assert 'TWINKLE_LONG_POLL_TIMEOUT' not in source, f'{module.__name__} reads the env var directly' + + +def test_gateway_guard_warns_once_per_value_not_once_per_request(monkeypatch): + """The poll-window guard must be audible but not spam. + + ``long_poll_window()`` runs on the hot path of both retrieve endpoints, not only at + startup, so an unguarded warning would repeat on every retrieve request (~2/s during + training). It must fire for a misconfigured value, stay silent on repeats of that same + value, and speak up again when the value changes to a different offending one. + + Counts calls on a stub logger rather than using ``caplog``: ``get_logger()`` sets + ``propagate = False``, so records never reach the root handler caplog installs. + """ + from twinkle.server.lifecycle import poll_config + + class _Counter: + def __init__(self): + self.warnings = 0 + + def warning(self, *_args, **_kwargs): + self.warnings += 1 + + counter = _Counter() + monkeypatch.setattr(poll_config, 'logger', counter) + monkeypatch.setattr(poll_config, '_warned_window', None, raising=False) + + def warnings_while(value: str, calls: int) -> int: + monkeypatch.setenv('TWINKLE_LONG_POLL_TIMEOUT', value) + before = counter.warnings + for _ in range(calls): + assert poll_config.long_poll_window() == float(value) + return counter.warnings - before + + assert warnings_while('90', calls=5) == 1, 'an over-the-gateway-limit window must warn exactly once' + assert warnings_while('90', calls=5) == 0, 'repeats of the same value must stay silent' + assert warnings_while('120', calls=3) == 1, 'a different offending value must warn again' + assert warnings_while('30', calls=5) == 0, 'a safe window must never warn' + + +if __name__ == '__main__': + raise SystemExit(pytest.main([__file__, '-v'])) diff --git a/tests/server/lifecycle/test_tinker_retrieve_regression.py b/tests/server/lifecycle/test_tinker_retrieve_regression.py new file mode 100644 index 000000000..4af4611f7 --- /dev/null +++ b/tests/server/lifecycle/test_tinker_retrieve_regression.py @@ -0,0 +1,67 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Tinker /retrieve_future wire regression. + +The tinker endpoint's response shape and status-code semantics must be unchanged by +this spec, across all three shapes: ``try_again`` / ``{error, category}`` / bare +result. It shares the poll_config window but keeps its own wire contract. +""" +from __future__ import annotations + +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from twinkle.server.gateway.tinker_handlers import _register_gateway_tinker_routes + + +class _State: + + def __init__(self, record): + self._record = record + + async def get_future(self, request_id): + return self._record + + +class _Gateway: + + def __init__(self, record): + self.state = _State(record) + + +def _client(record): + app = FastAPI() + _register_gateway_tinker_routes(app, lambda: _Gateway(record)) + return TestClient(app) + + +def test_try_again_shape_for_non_terminal(monkeypatch): + monkeypatch.setenv('TWINKLE_LONG_POLL_TIMEOUT', '0.2') + monkeypatch.setenv('TWINKLE_POLL_INTERVAL', '0.05') + resp = _client({'status': 'running', 'queue_state': 'active'}).post( + '/retrieve_future', json={'request_id': 'r'}) + assert resp.status_code == 200 + assert resp.json()['type'] == 'try_again' + + +def test_error_category_shape_for_failed(): + record = { + 'status': 'failed', + 'failure': { + 'reason_code': 'internal_error', + 'message': 'boom', + 'attribution': 'server', + }, + } + resp = _client(record).post('/retrieve_future', json={'request_id': 'r'}) + assert resp.status_code == 200 + body = resp.json() + assert body['error'] == 'boom' + assert body['category'] == 'server' + assert 'type' not in body + + +def test_bare_result_shape_for_completed(): + resp = _client({'status': 'completed', 'result': {'foo': 'bar'}}).post( + '/retrieve_future', json={'request_id': 'r'}) + assert resp.status_code == 200 + assert resp.json() == {'foo': 'bar'} diff --git a/tests/server/lifecycle/test_to_backend_inputs.py b/tests/server/lifecycle/test_to_backend_inputs.py new file mode 100644 index 000000000..4be370375 --- /dev/null +++ b/tests/server/lifecycle/test_to_backend_inputs.py @@ -0,0 +1,48 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Characterization tests for the shared ``to_backend_inputs`` seam (F003 / P003). + +Pins the input-shape rules that the sampler handlers used to re-implement inline, +so the three call sites (``sample`` / ``sample_to_data_plane`` batch form, and +``sample_stream`` single form) now share one definition. +""" +from __future__ import annotations + +import pytest + +from twinkle.data_format import InputFeature, Trajectory +from twinkle.server.lifecycle.submit import to_backend_inputs + +_IF = {'input_ids': [1, 2, 3]} +_TRAJ = {'messages': [{'role': 'user', 'content': 'hi'}]} + + +def test_batch_list_of_input_features(): + # Each element is parsed with InputFeature (dict with input_ids). + assert to_backend_inputs([_IF, _IF]) == [InputFeature(**_IF), InputFeature(**_IF)] + + +def test_batch_list_of_trajectories(): + # A dict without input_ids is parsed as a Trajectory. + assert to_backend_inputs([_TRAJ]) == [Trajectory(**_TRAJ)] + + +def test_batch_single_dict_becomes_one_element_list(): + assert to_backend_inputs(_IF) == [InputFeature(**_IF)] + assert to_backend_inputs(_TRAJ) == [Trajectory(**_TRAJ)] + + +def test_batch_passthrough_for_non_list_non_dict(): + sentinel = object() + assert to_backend_inputs(sentinel) is sentinel + + +def test_single_returns_one_object_not_a_list(): + out = to_backend_inputs([_IF], single=True) + assert not isinstance(out, list) + assert out == InputFeature(**_IF) + assert to_backend_inputs(_TRAJ, single=True) == Trajectory(**_TRAJ) + + +def test_single_rejects_multi_element_list(): + with pytest.raises(ValueError, match='single input'): + to_backend_inputs([_IF, _IF], single=True) diff --git a/tests/server/model/test_mock_model.py b/tests/server/model/test_mock_model.py index 5c24450f4..5564d4344 100644 --- a/tests/server/model/test_mock_model.py +++ b/tests/server/model/test_mock_model.py @@ -47,7 +47,6 @@ 'save', 'load', 'resume_from_checkpoint', - 'get_state_dict', 'get_train_configs', 'add_adapter', 'add_adapter_to_model', diff --git a/tests/server/model/test_replica_lifecycle.py b/tests/server/model/test_replica_lifecycle.py index e42193c35..84b961391 100644 --- a/tests/server/model/test_replica_lifecycle.py +++ b/tests/server/model/test_replica_lifecycle.py @@ -12,12 +12,17 @@ class _CapacityState: def __init__(self) -> None: self.capacities: dict[str, int] = {} + self.last_seen: set[str] = set() async def register_replica(self, replica_id: str, max_loras: int) -> None: self.capacities[replica_id] = max_loras async def unregister_replica(self, replica_id: str) -> None: self.capacities.pop(replica_id, None) + self.last_seen.discard(replica_id) + + async def touch_replica_last_seen(self, replica_id: str) -> None: + self.last_seen.add(replica_id) async def get_capacity_info(self) -> dict[str, int]: max_loras = sum(self.capacities.values()) @@ -31,6 +36,7 @@ def _make_lifecycle_manager(state: _CapacityState, replica_id: str, max_loras: i manager.max_loras = max_loras manager._replica_registered = False manager.data_plane = SimpleNamespace(close=AsyncMock()) + manager.shutdown_task_queue = AsyncMock() return manager @@ -42,6 +48,7 @@ async def test_replica_lifecycle_updates_shared_capacity() -> None: await first._register_replica_on_startup() assert await state.get_capacity_info() == {'max_loras': 3, 'used_loras': 0, 'free_loras': 3} + assert state.last_seen == {'replica-1'} await second._register_replica_on_startup() assert await state.get_capacity_info() == {'max_loras': 6, 'used_loras': 0, 'free_loras': 6} @@ -52,9 +59,10 @@ async def test_replica_lifecycle_updates_shared_capacity() -> None: @pytest.mark.asyncio async def test_async_constructor_registers_replica_before_ready() -> None: - state = SimpleNamespace(register_replica=AsyncMock()) + state = SimpleNamespace(register_replica=AsyncMock(), touch_replica_last_seen=AsyncMock()) replica_context = SimpleNamespace(replica_id=SimpleNamespace(unique_id='replica-1')) manager = ModelManagement.__new__(ModelManagement) + manager._task_queue_config = SimpleNamespace(effective_execution_timeout=1800.0) with patch('twinkle.server.model.app.DeviceGroup', return_value=SimpleNamespace(name='group')), \ patch('twinkle.server.model.app.init_twinkle_runtime', return_value=None), \ @@ -75,4 +83,5 @@ async def test_async_constructor_registers_replica_before_ready() -> None: ) state.register_replica.assert_awaited_once_with('replica-1', 3) + state.touch_replica_last_seen.assert_awaited_once_with('replica-1') assert manager._replica_registered is True diff --git a/tests/server/model/test_tinker_compat_output.py b/tests/server/model/test_tinker_compat_output.py index 2041aefa9..61c075570 100644 --- a/tests/server/model/test_tinker_compat_output.py +++ b/tests/server/model/test_tinker_compat_output.py @@ -1,7 +1,7 @@ import torch from tinker import types -from twinkle.server.common.datum import extract_rl_features_for_loss +from twinkle.server.model.tinker_datum import extract_rl_features_for_loss from twinkle.server.model.backends.common import TwinkleCompatModelBase diff --git a/tests/server/model/test_tinker_handlers.py b/tests/server/model/test_tinker_handlers.py index 474ce700f..7e78f306a 100644 --- a/tests/server/model/test_tinker_handlers.py +++ b/tests/server/model/test_tinker_handlers.py @@ -1,10 +1,10 @@ import pytest -from unittest.mock import AsyncMock, MagicMock, patch from fastapi import FastAPI from starlette.requests import Request from tinker import types +from unittest.mock import AsyncMock, MagicMock, patch -from twinkle.server.model.tinker_handlers import _register_tinker_routes +from twinkle.server.model.tinker_handlers import _register_model_tinker_routes class _DummyManagement: @@ -29,7 +29,7 @@ def _datum(): async def test_tinker_dpo_forward_backward_requires_per_dp_pairs(): management = _DummyManagement() app = FastAPI() - _register_tinker_routes(app, lambda: management) + _register_model_tinker_routes(app, lambda: management) body = types.ForwardBackwardRequest( model_id='model1', @@ -73,9 +73,11 @@ def assert_resource_exists(self, adapter_name): pass async def schedule_task(self, task, **kwargs): - # Actually execute the task to test response logic return await task() + async def call_backend(self, fn, /, *args, **kwargs): + return fn(*args, **kwargs) + @pytest.mark.asyncio @patch('twinkle.server.model.tinker_handlers.create_checkpoint_manager') @@ -89,7 +91,7 @@ async def test_save_weights_for_sampler_path_mode_returns_path(mock_create_ckpt_ management = _SaveWeightsDummyManagement() app = FastAPI() - _register_tinker_routes(app, lambda: management) + _register_model_tinker_routes(app, lambda: management) body = types.SaveWeightsForSamplerRequest( model_id='model1', @@ -117,7 +119,7 @@ async def test_save_weights_for_sampler_session_mode_returns_none_path(mock_crea management = _SaveWeightsDummyManagement() app = FastAPI() - _register_tinker_routes(app, lambda: management) + _register_model_tinker_routes(app, lambda: management) body = types.SaveWeightsForSamplerRequest( model_id='model1', diff --git a/tests/server/model/test_twinkle_async_inputs.py b/tests/server/model/test_twinkle_async_inputs.py index 7a13b73c5..ed89c3280 100644 --- a/tests/server/model/test_twinkle_async_inputs.py +++ b/tests/server/model/test_twinkle_async_inputs.py @@ -4,18 +4,27 @@ from fastapi import FastAPI from starlette.requests import Request -import twinkle_client.types as types -from twinkle.server.model.twinkle_handlers import _register_twinkle_routes -from twinkle.server.model.utils import model_result_rows +import twinkle.protocol.types as types +from twinkle.server.model.data_plane_inputs import model_result_rows +from twinkle.server.model.twinkle_handlers import _register_model_twinkle_routes def test_model_result_rows_keeps_one_output_row_per_sample() -> None: assert model_result_rows( - {'logps': [[-1.0], [-2.0]], 'loss': 0.25}, + { + 'logps': [[-1.0], [-2.0]], + 'loss': 0.25 + }, batch_size=2, ) == [ - {'logps': [-1.0], 'loss': 0.25}, - {'logps': [-2.0], 'loss': 0.25}, + { + 'logps': [-1.0], + 'loss': 0.25 + }, + { + 'logps': [-2.0], + 'loss': 0.25 + }, ] @@ -29,12 +38,16 @@ def __init__(self): self.data_plane = self self.rows = { 'data-a': [{ - 'train_input': {'input_ids': [index]}, + 'train_input': { + 'input_ids': [index] + }, 'sampled_logprobs': [-0.1], 'advantage': 1.0, } for index in range(4)], 'data-b': [{ - 'train_input': {'input_ids': [index]}, + 'train_input': { + 'input_ids': [index] + }, 'sampled_logprobs': [-0.2], 'advantage': -1.0, } for index in range(4, 8)], @@ -59,20 +72,23 @@ async def get(self, ref, *, fields=None): return rows return [{field: row[field] for field in fields} for row in rows] - async def schedule_task_and_wait(self, task, **kwargs): - self.scheduled.append(kwargs) - return await task() + async def submit_and_peek(self, coro_factory, *, model_id=None, token=None, task_type=None, **schedule_kwargs): + self.scheduled.append(schedule_kwargs) + result = await coro_factory() + from twinkle.protocol.types.lifecycle import TaskEnvelope + return TaskEnvelope(request_id='req-test', status='completed', result=result) + + async def call_backend(self, fn, /, *args, admit=True, **kwargs): + return fn(*args, **kwargs) @pytest.mark.asyncio async def test_forward_backward_resolves_multiple_data_refs_and_field_kwargs() -> None: management = _SchedulingManagement() app = FastAPI() - _register_twinkle_routes(app, lambda: management) + _register_model_twinkle_routes(app, lambda: management) route = next( - route for route in app.routes - if getattr(route, 'path', None) == '/twinkle/forward_backward_from_data_plane' - ) + route for route in app.routes if getattr(route, 'path', None) == '/twinkle/forward_backward_from_data_plane') request = Request({'type': 'http', 'headers': []}) request.state.session_id = 'session' body = types.DataPlaneForwardRequest( @@ -115,11 +131,9 @@ async def test_forward_backward_binds_nested_dpo_ref_logps_without_coercion() -> }, ] app = FastAPI() - _register_twinkle_routes(app, lambda: management) + _register_model_twinkle_routes(app, lambda: management) route = next( - route for route in app.routes - if getattr(route, 'path', None) == '/twinkle/forward_backward_from_data_plane' - ) + route for route in app.routes if getattr(route, 'path', None) == '/twinkle/forward_backward_from_data_plane') request = Request({'type': 'http', 'headers': []}) request.state.session_id = 'session' body = types.DataPlaneForwardRequest( diff --git a/tests/server/sampler/test_mock_sampler.py b/tests/server/sampler/test_mock_sampler.py index 8efae14b5..f9a2806df 100644 --- a/tests/server/sampler/test_mock_sampler.py +++ b/tests/server/sampler/test_mock_sampler.py @@ -115,6 +115,12 @@ def test_mock_dispatch_returns_mock_sampler() -> None: def test_explicit_async_vllm_uses_non_blocking_sampler(monkeypatch) -> None: + # Nothing here needs TransferQueue, but importing it is unavoidable: the + # `twinkle_agentic.async_rl` package eagerly pulls in `native_tq`, which subclasses + # `transfer_queue.GRPOGroupNSampler` at module level. That optional dependency lives + # in the `async-rl` extra, so skip when it is absent instead of failing a sampler test + # for a data-path dependency it does not use. + pytest.importorskip('transfer_queue') from twinkle_agentic.async_rl import vllm_sampler_tq as module captured = {} diff --git a/tests/server/sampler/test_resolve_sampler_weights.py b/tests/server/sampler/test_resolve_sampler_weights.py new file mode 100644 index 000000000..1d98c76c9 --- /dev/null +++ b/tests/server/sampler/test_resolve_sampler_weights.py @@ -0,0 +1,53 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Tests for the shared sampler weight-resolution helper (F004 / P004).""" +from __future__ import annotations + +import pytest + +from twinkle.server.sampler.weights import resolve_sampler_weights + + +class _FakeSampler: + + def load_full_weights_from_path(self, path): + return None + + +class _FakeService: + """Minimal stand-in exposing the two attributes the helper touches.""" + + def __init__(self): + self.sampler = _FakeSampler() + self.full_weight_loads: list[str] = [] + + async def call_backend(self, fn, *args, **kwargs): + # The helper only routes load_full_weights_from_path through call_backend. + if getattr(fn, '__name__', None) == 'load_full_weights_from_path': + self.full_weight_loads.append(args[0]) + return fn(*args, **kwargs) + + +@pytest.mark.asyncio +async def test_lora_dir_returns_adapter_path_and_loads_no_full_weights(tmp_path): + (tmp_path / 'adapter_config.json').write_text('{}') + svc = _FakeService() + result = await resolve_sampler_weights(svc, str(tmp_path)) + assert result == str(tmp_path) + assert svc.full_weight_loads == [] + + +@pytest.mark.asyncio +async def test_full_checkpoint_loads_weights_and_returns_none(tmp_path): + # A directory without adapter_config.json is a full-parameter checkpoint. + svc = _FakeService() + result = await resolve_sampler_weights(svc, str(tmp_path)) + assert result is None + assert svc.full_weight_loads == [str(tmp_path)] + + +@pytest.mark.asyncio +async def test_empty_uri_is_a_noop(): + svc = _FakeService() + assert await resolve_sampler_weights(svc, None) is None + assert await resolve_sampler_weights(svc, '') is None + assert svc.full_weight_loads == [] diff --git a/tests/server/sampler/test_stream_guarantees.py b/tests/server/sampler/test_stream_guarantees.py new file mode 100644 index 000000000..d14895751 --- /dev/null +++ b/tests/server/sampler/test_stream_guarantees.py @@ -0,0 +1,113 @@ +from __future__ import annotations + +import asyncio +import json +import threading +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from fastapi import FastAPI + +from twinkle.server.sampler.app import SamplerManagement +from twinkle.server.sampler.twinkle_handlers import _await_generation, _register_twinkle_sampler_routes, _stream_queue + + +class _BlockingQueue: + + def __init__(self) -> None: + self.released = threading.Event() + self.get_exited = threading.Event() + self.closed = False + + def get(self): + self.released.wait() + self.get_exited.set() + return 'sentinel' + + def shutdown(self, *, force: bool) -> None: + assert force is True + self.closed = True + self.released.set() + + +@pytest.mark.asyncio +async def test_sampler_request_refreshes_replica_liveness(): + service = SamplerManagement.__new__(SamplerManagement) + service.replica_id = 'sampler-replica' + service.state = SimpleNamespace(touch_replica_last_seen=AsyncMock()) + service._ensure_sticky = AsyncMock() + service._ensure_state_cleanup_started = AsyncMock() + request = SimpleNamespace( + headers={'Authorization': 'Bearer token'}, state=SimpleNamespace(token='token')) + + assert await service._on_request_start(request) == 'token' + service.state.touch_replica_last_seen.assert_awaited_once_with('sampler-replica') + + +@pytest.mark.asyncio +async def test_stream_without_actor_returns_structured_error(): + service = SimpleNamespace( + sampler=SimpleNamespace(_actors=[]), + _on_request_start=AsyncMock(return_value='token'), + ) + app = FastAPI() + _register_twinkle_sampler_routes(app, lambda: service) + route = next(route for route in app.routes if getattr(route, 'path', None) == '/twinkle/sample_stream') + request = SimpleNamespace(state=SimpleNamespace(request_id='request')) + body = SimpleNamespace(adapter_name='', adapter_uri=None, inputs={'input_ids': [1]}, sampling_params=None) + + response = await route.endpoint(request, body, service) + chunks = [chunk async for chunk in response.body_iterator] + payload = json.loads(chunks[0]) + assert payload['category'] == 'server' + assert payload['error_code'] == 503 + assert payload['request_id'].startswith('req_') + + +@pytest.mark.asyncio +async def test_stream_timeout_returns_error_payload_and_closes_queue(): + queue = _BlockingQueue() + chunks = [ + chunk async for chunk in _stream_queue( + queue, + sentinel='sentinel', + request_id='req-stream', + total_timeout=0.05, + single_get_timeout=0.05, + ) + ] + + payload = json.loads(chunks[0]) + assert payload['category'] == 'server' + assert payload['error_code'] == 504 + assert payload['request_id'] == 'req-stream' + assert queue.closed is True + assert queue.get_exited.wait(timeout=5) + + +class _GenerationService: + + def __init__(self) -> None: + self.cancelled = False + self.sampler = SimpleNamespace( + get_generation_status=lambda _submission_id: {'status': 'running'}, + collect_generation=lambda _submission_id: [], + cancel_generation=self._cancel, + ) + + async def call_backend(self, fn, /, *args, **kwargs): + return fn(*args, **kwargs) + + def _cancel(self, _submission_id: str) -> None: + self.cancelled = True + + +@pytest.mark.asyncio +async def test_generation_poll_has_total_timeout_and_cancels(): + service = _GenerationService() + + with pytest.raises(asyncio.TimeoutError): + await _await_generation(service, 'submission', timeout=0.05) + + assert service.cancelled is True diff --git a/tests/server/sampler/test_tinker_handlers.py b/tests/server/sampler/test_tinker_handlers.py index 1e338594d..e6fe1596c 100644 --- a/tests/server/sampler/test_tinker_handlers.py +++ b/tests/server/sampler/test_tinker_handlers.py @@ -20,6 +20,7 @@ class _DummySampler: def __init__(self): self.adapter_paths = [] + self.sampling_params = [] def set_template(self, *args, **kwargs): return None @@ -29,6 +30,7 @@ def reset_prefix_cache(self): def sample(self, inputs, sampling_params=None, adapter_name='', *, adapter_path=None, **kwargs): self.adapter_paths.append(adapter_path) + self.sampling_params.append(sampling_params) return [ SampleResponse( sequences=[SampledSequence( @@ -52,6 +54,9 @@ async def _on_request_start(self, request): async def schedule_task(self, task, **kwargs): return await task() + async def call_backend(self, fn, /, *args, **kwargs): + return fn(*args, **kwargs) + @pytest.mark.asyncio async def test_tinker_asample_allows_base_model_session_without_model_path(): @@ -71,4 +76,6 @@ async def test_tinker_asample_allows_base_model_session_without_model_path(): response = await route.endpoint(request, body, management) assert isinstance(response, types.SampleResponse) + assert response.sequences[0].tokens == [1, 2] assert management.sampler.adapter_paths == [None] + assert management.sampler.sampling_params[0].logprobs == 1 diff --git a/tests/server/sampler/test_twinkle_async_rows.py b/tests/server/sampler/test_twinkle_async_rows.py index 30b7016b2..4df488130 100644 --- a/tests/server/sampler/test_twinkle_async_rows.py +++ b/tests/server/sampler/test_twinkle_async_rows.py @@ -1,10 +1,12 @@ from __future__ import annotations +from types import SimpleNamespace + import pytest from fastapi import FastAPI from starlette.requests import Request -import twinkle_client.types as types +import twinkle.protocol.types as types from twinkle.data_format import SampledSequence, SampleResponse from twinkle.server.sampler.twinkle_handlers import ( _register_twinkle_sampler_routes, @@ -67,13 +69,19 @@ def __init__(self): self.enabled = True self.scheduled = [] self.put_rows = None + self.task_queue_config = SimpleNamespace(effective_execution_timeout=60.0) async def _on_request_start(self, _request): return 'token' - async def schedule_task_and_wait(self, task, **kwargs): - self.scheduled.append(kwargs) - return await task() + async def submit_background_and_peek(self, coro_factory, *, model_id=None, task_type=None): + self.scheduled.append({'model_id': model_id, 'task_type': task_type}) + result = await coro_factory() + from twinkle.protocol.types.lifecycle import TaskEnvelope + return TaskEnvelope(request_id='req-test', status='completed', result=result) + + async def call_backend(self, fn, /, *args, admit=True, **kwargs): + return fn(*args, **kwargs) def submit_generation(self, submission_id, inputs, params, **kwargs): self.submission_id = submission_id @@ -124,14 +132,15 @@ async def test_sample_to_data_plane_returns_ref_after_short_admission() -> None: sampling_params={'max_tokens': 4}, ) - ref = await route.endpoint(request, body, management) + env = await route.endpoint(request, body, management) + ref = types.DataRef(**env.result) assert ref.ref_id == 'rollout-ref' + # The whole admit -> generate -> store runs as one background future now + # (vLLM owns generation concurrency), so there is a single scheduled task. assert management.scheduled == [{ 'model_id': 'session-adapter', - 'token': 'token', - 'input_tokens': 1, - 'task_type': 'sample_admission', + 'task_type': 'sample_to_data_plane', }] assert management.put_rows == [{ 'train_input': {'input_ids': [1, 7], 'labels': [-100, 7]}, diff --git a/tests/server/session_resource/test_contract.py b/tests/server/session_resource/test_contract.py new file mode 100644 index 000000000..477b509cf --- /dev/null +++ b/tests/server/session_resource/test_contract.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +import asyncio +from contextlib import suppress +from unittest import mock + +import pytest + +from twinkle.server.exceptions import RequestRejectedError +from twinkle.server.session_resource.adapter import AdapterManagerMixin +from twinkle.server.session_resource.base import SessionResourceMixin +from twinkle.server.session_resource.processor import ProcessorManagerMixin + + +class _State: + + def __init__(self, outcomes: list[object] | None = None) -> None: + self.outcomes = list(outcomes or []) + + async def get_session_last_heartbeat(self, session_id: str) -> float | None: + outcome = self.outcomes.pop(0) + if isinstance(outcome, Exception): + raise outcome + return outcome # type: ignore[return-value] + + +class _ResourceManager(SessionResourceMixin): + + def __init__(self, state: _State, timeout: float = 10.0) -> None: + self.state = state + self.expired: list[str] = [] + self._init_resource_manager(resource_timeout=timeout) + + async def _on_resource_expired(self, resource_id: str) -> None: + self.expired.append(resource_id) + + +class _MissingBaseHook(SessionResourceMixin): + pass + + +class _MissingAdapterHook(AdapterManagerMixin): + pass + + +class _MissingProcessorHook(ProcessorManagerMixin): + pass + + +def test_all_resource_expiry_hooks_are_abstract() -> None: + for cls in (_MissingBaseHook, _MissingAdapterHook, _MissingProcessorHook): + with pytest.raises(TypeError): + cls() + + +def test_registration_requires_session_id() -> None: + manager = _ResourceManager(_State()) + with pytest.raises(RequestRejectedError): + manager.register_resource('r1', 'token', '') + + +@pytest.mark.asyncio +async def test_liveness_failure_has_hard_upper_bound_and_recovery_refreshes() -> None: + state = _State([RuntimeError('down'), 109.0, RuntimeError('down'), RuntimeError('down')]) + manager = _ResourceManager(state, timeout=10.0) + with mock.patch('twinkle.server.session_resource.base.time.time', return_value=100.0): + manager.register_resource('r1', 'token', 'session') + record = manager.get_resource_info('r1') + + with mock.patch('twinkle.server.session_resource.base.time.time', return_value=105.0): + assert await manager._is_session_alive('session', 'r1', record) is True + with mock.patch('twinkle.server.session_resource.base.time.time', return_value=110.0): + assert await manager._is_session_alive('session', 'r1', record) is True + assert manager.get_resource_info('r1')['last_liveness_confirmed_at'] == 110.0 + with mock.patch('twinkle.server.session_resource.base.time.time', return_value=115.0): + assert await manager._is_session_alive('session', 'r1', record) is True + with mock.patch('twinkle.server.session_resource.base.time.time', return_value=120.0): + assert await manager._is_session_alive('session', 'r1', record) is False + + +@pytest.mark.asyncio +async def test_countdown_restart_preserves_confirmation_time() -> None: + manager = _ResourceManager(_State([100.0])) + with mock.patch('twinkle.server.session_resource.base.time.time', return_value=100.0): + manager.register_resource('r1', 'token', 'session') + confirmed_at = manager.get_resource_info('r1')['last_liveness_confirmed_at'] + + manager._ensure_countdown_started() + first_task = manager._countdown_task + manager.stop_resource_countdown() + with suppress(asyncio.CancelledError): + await first_task + + manager._ensure_countdown_started() + second_task = manager._countdown_task + assert second_task is not first_task + assert manager.get_resource_info('r1')['last_liveness_confirmed_at'] == confirmed_at + manager.stop_resource_countdown() + with suppress(asyncio.CancelledError): + await second_task + + +def test_worker_restart_does_not_restore_local_resources() -> None: + old = _ResourceManager(_State()) + old.register_resource('r1', 'token', 'session') + restarted = _ResourceManager(_State()) + assert restarted.get_resource_info('r1') is None diff --git a/tests/server/start_e2e_server.py b/tests/server/start_e2e_server.py index 0e8939ed0..4af2e2b85 100644 --- a/tests/server/start_e2e_server.py +++ b/tests/server/start_e2e_server.py @@ -1,4 +1,6 @@ -"""One-click: restart Ray cluster + launch Twinkle server + wait until ready. +"""Manual PPU E2E helper; it is not a pytest test or CI entry point. + +One-click: restart Ray cluster + launch Twinkle server + wait until ready. Usage: python start_e2e_server.py # default config (transformers LoRA) @@ -22,8 +24,8 @@ import requests # ── Paths ── -RAY = "/mnt/nas2/anaconda3/envs/tinker_myl/bin/ray" -PYTHON = "/mnt/nas2/anaconda3/envs/tinker_myl/bin/python" +RAY = "/mnt/nas2/anaconda3/envs/twinkle_ppu_vllm/bin/ray" +PYTHON = "/mnt/nas2/anaconda3/envs/twinkle_ppu_vllm/bin/python" WORKDIR = "/mnt/nas2/yunlin.myl/twinkle" DEFAULT_CONFIG = "tests/server/config/server_config_4b_e2e.yaml" RAY_TEMP_DIR = "/mnt/nas2/yunlin.myl/ray_logs" @@ -136,14 +138,14 @@ def restart_ray(): run(f"{RAY} stop --force", check=False) time.sleep(2) - # Head node: GPU 0,1,2,3 (4 GPUs for model PP=2 x DP=2) + # Head node: 4 GPUs for model PP=2 x DP=2 (skip busy card 2). run(f"{RAY} start --head --port=6379 --num-gpus=4 " f"--disable-usage-stats --temp-dir={RAY_TEMP_DIR}", - env={"CUDA_VISIBLE_DEVICES": "0,1,2,3"}) + env={"CUDA_VISIBLE_DEVICES": "0,1,3,4"}) - # Worker: GPU 4 (1 GPU for sampler) + # Worker: 1 GPU for sampler. run(f"{RAY} start --address=127.0.0.1:6379 --num-gpus=1", - env={"CUDA_VISIBLE_DEVICES": "4"}) + env={"CUDA_VISIBLE_DEVICES": "5"}) # CPU-only worker (processor + server) run(f"{RAY} start --address=127.0.0.1:6379 --num-gpus=0", diff --git a/tests/server/state/fake_backend.py b/tests/server/state/fake_backend.py new file mode 100644 index 000000000..eaac98833 --- /dev/null +++ b/tests/server/state/fake_backend.py @@ -0,0 +1,85 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""In-process ``StateBackend`` test double. + +Both shipped backends need infrastructure -- ``memory`` starts a detached Ray actor, +``redis`` needs a server -- so state-logic tests that only care about +``FutureManager`` / manager semantics use this dict-backed fake instead. Running on +a single event loop, ``set_nx`` / ``update_atomic`` are atomic simply by not +awaiting between read and write, mirroring the real backends' guarantee. +""" +from __future__ import annotations + +import time +from collections.abc import Callable +from fnmatch import fnmatch +from typing import Any + +from twinkle.server.state.backend.base import StateBackend + + +class FakeBackend(StateBackend): + """Dict-backed backend with TTL support, for tests only.""" + + def __init__(self) -> None: + self._store: dict[str, tuple[Any, float | None]] = {} + + def _is_expired(self, key: str) -> bool: + entry = self._store.get(key) + if entry is None: + return True + _, expire_at = entry + if expire_at is not None and time.time() >= expire_at: + del self._store[key] + return True + return False + + async def set(self, key: str, value: Any, ttl: int | None = None) -> None: + self._store[key] = (value, (time.time() + ttl) if ttl is not None else None) + + async def get(self, key: str) -> Any | None: + if self._is_expired(key): + return None + return self._store[key][0] + + async def mget(self, keys: list[str]) -> list[Any | None]: + return [await self.get(key) for key in keys] + + async def delete(self, key: str) -> None: + self._store.pop(key, None) + + async def exists(self, key: str) -> bool: + return not self._is_expired(key) + + async def keys(self, pattern: str) -> list[str]: + return [k for k in list(self._store) if not self._is_expired(k) and fnmatch(k, pattern)] + + async def count(self, pattern: str) -> int: + return len(await self.keys(pattern)) + + async def set_nx(self, key: str, value: Any, ttl: int | None = None) -> bool: + if not self._is_expired(key): + return False + await self.set(key, value, ttl) + return True + + async def update_atomic( + self, + key: str, + transform: Callable[[Any | None], Any | None], + ttl: int | None = None, + ) -> Any | None: + current = await self.get(key) + updated = transform(current) + if updated is None: + return current + await self.set(key, updated, ttl) + return updated + + async def flush_all(self) -> None: + self._store.clear() + + async def close(self) -> None: + pass + + async def health_check(self) -> bool: + return True diff --git a/tests/server/state/test_error_payload.py b/tests/server/state/test_error_payload.py new file mode 100644 index 000000000..417470e1e --- /dev/null +++ b/tests/server/state/test_error_payload.py @@ -0,0 +1,110 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Tests for direct/streaming ErrorPayload construction and Tinker parsing.""" +from __future__ import annotations + +import pytest +from pydantic import ValidationError + +from twinkle.server.task_errors import build_error_payload, task_error_payload +from twinkle.protocol.types.errors import ErrorCategory, ErrorPayload + + +def test_overlong_traceback_is_trimmed_tail_kept_with_marker(): + long_tb = 'X' * 10 + ('line\n' * 40000) # well over 65536 chars + assert len(long_tb) > 65536 + + payload = task_error_payload( + 'RuntimeError: boom', request_id='req_1', error_code=500, traceback_text=long_tb) + + tb = payload['traceback'] + assert tb is not None + assert len(tb) <= 65536 + assert 'truncated' in tb # truncation marker present + assert tb.endswith('line\n') # tail preserved + + +def test_task_error_payload_shapes_and_sanitizes_errors(): + payload = task_error_payload( + 'RuntimeError: boom\n File "/server/path.py", line 1', request_id='req_1', error_code=500) + + assert payload == { + 'error': 'RuntimeError: boom', + 'category': ErrorCategory.Server.value, + 'error_code': 500, + 'request_id': 'req_1', + } + + +def test_build_error_payload_bounds_text_and_preserves_metadata(): + details = [{'field': 'adapter_name'}] + payload = build_error_payload( + f'bad input {"X" * 2048}\nignored second line', + request_id='req_build', + error_code=422, + category='user', + traceback_text='server stack', + details=details, + ) + + assert isinstance(payload, ErrorPayload) + assert len(payload.error) == 1024 + assert '\n' not in payload.error + assert payload.category is ErrorCategory.User + assert payload.error_code == 422 + assert payload.request_id == 'req_build' + assert payload.traceback is None + assert payload.details == details + assert task_error_payload( + f'bad input {"X" * 2048}\nignored second line', + request_id='req_build', + error_code=422, + category='user', + traceback_text='server stack', + details=details, + ) == payload.model_dump(mode='json', exclude_none=True) + + +def test_user_category_carries_no_traceback(): + payload = task_error_payload( + 'invalid field', request_id='req_2', error_code=422, + category=ErrorCategory.User, traceback_text='Traceback (most recent call last): ...') + + assert payload['category'] == ErrorCategory.User.value + assert 'traceback' not in payload + + +def test_error_category_matches_tinker_wire_values(): + from tinker.types import RequestErrorCategory + + assert {item.value for item in RequestErrorCategory} == {item.value for item in ErrorCategory} + + +def test_tinker_sdk_parses_six_field_like_two_field(): + """Tinker's RequestFailedResponse ignores extra fields, so a six-field + payload parses equal to a two-field one on the declared fields. + + tinker's RequestErrorCategory values are lowercase ('server'), so the payloads + here use that value; the point under test is that the four extra fields are + ignored, not the category spelling.""" + from tinker.types import RequestFailedResponse + + two = {'error': 'boom', 'category': 'server'} + six = task_error_payload('boom', request_id='req_9', error_code=504) + + parsed_six = RequestFailedResponse.model_validate(six) + parsed_two = RequestFailedResponse.model_validate(two) + + assert parsed_six.error == parsed_two.error + assert parsed_six.category == parsed_two.category + + +@pytest.mark.parametrize('category', [ErrorCategory.User, ErrorCategory.Unknown]) +def test_non_server_traceback_is_rejected(category): + with pytest.raises(ValidationError): + ErrorPayload( + error='bad input', + category=category, + error_code=400, + request_id='req_11', + traceback='server stack', + ) diff --git a/tests/server/state/test_future_lifecycle.py b/tests/server/state/test_future_lifecycle.py new file mode 100644 index 000000000..9728780f3 --- /dev/null +++ b/tests/server/state/test_future_lifecycle.py @@ -0,0 +1,249 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""State-hygiene tests for FutureManager cleanup and the do-not-regress guard. + +Both shipped backends need infrastructure (``memory`` starts a detached Ray actor, +``redis`` needs a server), +so these pure ``FutureManager`` semantics run against the dict-backed fake below. +""" +from __future__ import annotations + +import time +from collections.abc import Callable +from fnmatch import fnmatch +from typing import Any +from unittest import mock + +import pytest + +from twinkle.server.state.backend.base import StateBackend +from twinkle.server.state.future_manager import FutureManager +from twinkle.server.state.models import FutureFailureRecord + + +class _FakeBackend(StateBackend): + """Dict-backed StateBackend. Single event loop, so ``update_atomic`` is atomic + simply by not awaiting between read and write -- the real backends' guarantee.""" + + def __init__(self) -> None: + self._store: dict[str, tuple[Any, float | None]] = {} + + def _is_expired(self, key: str) -> bool: + entry = self._store.get(key) + if entry is None: + return True + _, expire_at = entry + if expire_at is not None and time.time() >= expire_at: + del self._store[key] + return True + return False + + async def set(self, key: str, value: Any, ttl: int | None = None) -> None: + self._store[key] = (value, (time.time() + ttl) if ttl is not None else None) + + async def get(self, key: str) -> Any | None: + return None if self._is_expired(key) else self._store[key][0] + + async def mget(self, keys: list[str]) -> list[Any | None]: + return [await self.get(key) for key in keys] + + async def delete(self, key: str) -> None: + self._store.pop(key, None) + + async def exists(self, key: str) -> bool: + return not self._is_expired(key) + + async def keys(self, pattern: str) -> list[str]: + return [k for k in list(self._store) if not self._is_expired(k) and fnmatch(k, pattern)] + + async def count(self, pattern: str) -> int: + return len(await self.keys(pattern)) + + async def set_nx(self, key: str, value: Any, ttl: int | None = None) -> bool: + if not self._is_expired(key): + return False + await self.set(key, value, ttl) + return True + + async def update_atomic( + self, + key: str, + transform: Callable[[Any | None], Any | None], + ttl: int | None = None, + ) -> Any | None: + current = await self.get(key) + updated = transform(current) + if updated is None: + return current + await self.set(key, updated, ttl) + return updated + + async def close(self) -> None: + pass + + async def health_check(self) -> bool: + return True + + +@pytest.fixture +def manager(): + return FutureManager(_FakeBackend(), expiration_timeout=300.0) + + +async def _store(manager, request_id, status, *, replica_id=None, absolute_deadline=None): + failure = None + if status in ('failed', 'cancelled'): + failure = FutureFailureRecord( + reason_code='internal_error', message='boom', attribution='server') + await manager.store_status( + request_id, + status, + model_id='m1', + failure=failure, + replica_id=replica_id, + absolute_deadline=absolute_deadline, + ) + + +@pytest.mark.asyncio +async def test_non_terminal_with_live_replica_is_kept(manager): + await _store(manager, 'r1', 'running', replica_id='replica-A') + removed = await manager.cleanup_expired( + cutoff_time=time.time() + 10, alive_replica_ids={'replica-A'}) + assert removed == 0 + rec = await manager.get('r1') + assert rec is not None and rec.status == 'running' + + +@pytest.mark.asyncio +async def test_non_terminal_orphan_is_failed_not_deleted(manager): + await _store(manager, 'r2', 'running', replica_id='dead-replica') + await manager.cleanup_expired(cutoff_time=time.time() + 10, alive_replica_ids={'replica-A'}) + rec = await manager.get('r2') + assert rec is not None # NOT deleted + assert rec.status == 'failed' + assert rec.result is None + assert rec.failure.reason_code == 'orphaned_replica' + assert rec.failure.attribution == 'server' + + +@pytest.mark.asyncio +async def test_non_terminal_past_absolute_deadline_is_failed(manager): + await _store(manager, 'r3', 'running', replica_id='replica-A', absolute_deadline=time.time() - 1) + await manager.cleanup_expired(cutoff_time=time.time() + 10, alive_replica_ids={'replica-A'}) + rec = await manager.get('r3') + assert rec is not None and rec.status == 'failed' + assert rec.failure.reason_code == 'deadline_exceeded' + + +@pytest.mark.asyncio +async def test_record_without_deadline_uses_expiration_timeout(manager): + await _store(manager, 'without-deadline', 'running', replica_id=None) + with mock.patch('twinkle.server.state.future_manager.time.time', return_value=time.time() + 301): + await manager.cleanup_expired(cutoff_time=time.time() + 10, alive_replica_ids=set()) + rec = await manager.get('without-deadline') + assert rec is not None and rec.status == 'failed' + assert rec.failure.reason_code == 'deadline_exceeded' + + +@pytest.mark.asyncio +async def test_terminal_expired_is_deleted(manager): + await _store(manager, 'r4', 'completed', replica_id='replica-A') + removed = await manager.cleanup_expired( + cutoff_time=time.time() + 10, alive_replica_ids={'replica-A'}) + assert removed == 1 + assert await manager.get('r4') is None + + +@pytest.mark.asyncio +async def test_terminal_to_terminal_different_is_refused_and_warns(manager): + await _store(manager, 'r5', 'failed', replica_id='replica-A') + with mock.patch('twinkle.server.state.future_manager.logger') as log: + await manager.store_status('r5', 'completed', model_id='m1') + rec = await manager.get('r5') + assert rec.status == 'failed' # not overwritten + assert log.warning.called + + +@pytest.mark.asyncio +async def test_terminal_to_terminal_same_is_dropped_without_warning(manager): + await _store(manager, 'r6', 'completed', replica_id='replica-A') + with mock.patch('twinkle.server.state.future_manager.logger') as log: + await manager.store_status('r6', 'completed', model_id='m1') + rec = await manager.get('r6') + assert rec.status == 'completed' + assert not log.warning.called + + +@pytest.mark.asyncio +async def test_replica_id_and_deadline_set_at_creation_not_overwritten(manager): + deadline = time.time() + 100 + await _store(manager, 'r7', 'pending', replica_id='replica-A', absolute_deadline=deadline) + await manager.store_status( + 'r7', 'running', model_id='m1', replica_id='replica-B', absolute_deadline=time.time() + 999) + rec = await manager.get('r7') + assert rec.replica_id == 'replica-A' + assert rec.absolute_deadline == deadline + + +@pytest.mark.asyncio +async def test_stored_timestamps_align_with_wall_clock_regardless_of_host_tz(manager): + """Writer (_now_iso), reader (_parse_timestamp) and time.time() must agree. + + A record written now must parse to within a second of time.time() on any host, + not skewed by the host's UTC offset (the former naive-local / read-as-UTC bug). + """ + before = time.time() + await _store(manager, 'r8', 'running', replica_id='replica-A') + after = time.time() + rec = await manager.get('r8') + parsed = manager._parse_timestamp(rec.created_at) + assert before - 1 <= parsed <= after + 1 + + +@pytest.mark.asyncio +async def test_claim_seq_dedups_then_release_readmits(): + from twinkle.server.state.server_state import ServerState + state = ServerState(backend=_FakeBackend()) + # First claim of a (session, seq_id) is unseen -> None, caller proceeds to enqueue. + assert await state.claim_seq('seq::s1::1', 'reqA', ttl=60) is None + # A duplicate claim returns the original request_id -> caller returns its envelope. + assert await state.claim_seq('seq::s1::1', 'reqB', ttl=60) == 'reqA' + # Releasing (e.g. the original was preflight-rejected) re-admits the same seq_id. + await state.release_seq('seq::s1::1') + assert await state.claim_seq('seq::s1::1', 'reqC', ttl=60) is None + # Different session with the same seq_id never collides. + assert await state.claim_seq('seq::s2::1', 'reqD', ttl=60) is None + + +@pytest.mark.asyncio +async def test_cancel_drops_pending_but_never_running(): + from twinkle.server.state.server_state import ServerState + state = ServerState(backend=_FakeBackend()) + # pending -> cancel drops it to a terminal domain failure. + await state.store_future_status('rp', 'pending', 'm1') + assert await state.cancel_future('rp') == {'cancelled': True, 'state': 'cancelled'} + rec = await state.get_future('rp') + assert rec['status'] == 'cancelled' + assert rec['result'] is None + assert rec['failure']['reason_code'] == 'cancelled' + assert rec['failure']['attribution'] == 'user' + # running -> cancel is a no-op; in-flight work is never interrupted. + await state.store_future_status('rr', 'running', 'm1') + assert await state.cancel_future('rr') == {'cancelled': False, 'state': 'running'} + # unknown request_id -> not_found. + assert await state.cancel_future('nope') == {'cancelled': False, 'state': 'not_found'} + + +def test_cancelled_record_maps_to_error_envelope(): + from twinkle.server.lifecycle.envelope import envelope_from_record + rec = { + 'status': 'cancelled', + 'failure': { + 'reason_code': 'cancelled', + 'message': 'Task cancelled by client', + 'attribution': 'user', + }, + } + env = envelope_from_record('rc', rec) + assert env.status == 'cancelled' + assert env.error is not None and env.error.error_code == 499 diff --git a/tests/server/state/test_leader_election.py b/tests/server/state/test_leader_election.py index f1f8955c3..792989c30 100644 --- a/tests/server/state/test_leader_election.py +++ b/tests/server/state/test_leader_election.py @@ -14,7 +14,7 @@ from twinkle.server.state import ServerState from twinkle.server.state.backend.memory_backend import RayActorBackend -from twinkle.server.state.server_state import LEADER_KEY, LEASE_RENEW +from twinkle.server.state.cleanup_coordinator import LEADER_KEY, LEASE_RENEW from twinkle.server.telemetry import MetricsRegistry diff --git a/tests/server/state/test_managers.py b/tests/server/state/test_managers.py index 1b9016be1..242dc3508 100644 --- a/tests/server/state/test_managers.py +++ b/tests/server/state/test_managers.py @@ -7,6 +7,7 @@ from datetime import datetime, timezone from unittest import mock +from twinkle.server.exceptions import ResourceQuotaExceededError from twinkle.server.state import ServerState from twinkle.server.state.backend.memory_backend import RayActorBackend from twinkle.server.state.future_manager import FutureManager @@ -14,6 +15,7 @@ from twinkle.server.state.models import FutureRecord, ModelRecord, SamplingSessionRecord, SessionRecord from twinkle.server.state.sampling_manager import SamplingSessionManager from twinkle.server.state.session_manager import SessionManager +from .fake_backend import FakeBackend # ============================================================ # SessionManager Tests @@ -97,7 +99,7 @@ async def test_cleanup_expired(self, manager): await manager.add('new_sess', new_record) cutoff = now - 500 - removed_count = await manager.cleanup_expired(cutoff) + _, removed_count = await manager.collect_and_remove_expired(cutoff) assert removed_count == 1 assert await manager.get('old_sess') is None assert await manager.get('new_sess') is not None @@ -110,7 +112,7 @@ async def test_cleanup_expired_uses_created_at_fallback(self, manager): await manager.add('old_sess', record) cutoff = time.time() - 100 - removed_count = await manager.cleanup_expired(cutoff) + _, removed_count = await manager.collect_and_remove_expired(cutoff) assert removed_count == 1 @@ -148,11 +150,11 @@ async def test_remove(self, manager): @pytest.mark.asyncio async def test_token_limit_enforced(self, manager): - """Adding more models than per_token_model_limit should raise RuntimeError.""" + """Adding more models than the per-token quota raises a user quota error.""" for i in range(3): await manager.add(f'm{i}', ModelRecord(token='tok1')) - with pytest.raises(RuntimeError, match='Model limit exceeded'): + with pytest.raises(ResourceQuotaExceededError, match='Model quota exceeded'): await manager.add('m3', ModelRecord(token='tok1')) @pytest.mark.asyncio @@ -172,6 +174,12 @@ async def test_replica_registration(self, manager): assert info['used_loras'] == 0 assert info['free_loras'] == 5 + @pytest.mark.asyncio + async def test_liveness_only_replica_is_alive(self, manager): + await manager.touch_replica_last_seen('sampler-replica') + alive = await manager.get_alive_replica_ids(liveness_threshold=60) + assert 'sampler-replica' in alive + @pytest.mark.asyncio async def test_capacity_info_after_add(self, manager): await manager.register_replica('r1', max_loras=3) @@ -196,9 +204,10 @@ async def test_indexes_derived_from_backend(self, manager): avail = await manager.get_available_replica_ids(['r1', 'r2']) assert avail == ['r1', 'r2'] - # Per-token count enforces the limit using the persisted records. - count = await manager._count_models_for_token('tok1') - assert count == 2 + # The public add path enforces the per-token quota from persisted counts. + await manager.add('m3', ModelRecord(token='tok1')) + with pytest.raises(ResourceQuotaExceededError, match='Model quota exceeded'): + await manager.add('m4', ModelRecord(token='tok1')) @pytest.mark.asyncio async def test_cascade_cleanup_by_session(self, manager): @@ -253,7 +262,7 @@ async def try_add(i: int) -> None: try: await state.register_model({'base_model': 'b'}, token='tok', model_id=f'm{i}') results.append(True) - except RuntimeError: + except ResourceQuotaExceededError: results.append(False) await asyncio.gather(*(try_add(i) for i in range(n))) @@ -270,7 +279,7 @@ async def test_remove_frees_token_slot(self): state = ServerState(backend=backend, per_token_model_limit=1) await state.register_model({'base_model': 'b'}, token='tok', model_id='m1') - with pytest.raises(RuntimeError): + with pytest.raises(ResourceQuotaExceededError): await state.register_model({'base_model': 'b'}, token='tok', model_id='m2') assert await state.unload_model('m1') is True @@ -289,10 +298,57 @@ async def test_rebuild_indexes_recovers_counter_from_records(self): await state._model_mgr.rebuild_indexes() await state.register_model({'base_model': 'b'}, token='tok', model_id='m3') - with pytest.raises(RuntimeError): + with pytest.raises(ResourceQuotaExceededError): await state.register_model({'base_model': 'b'}, token='tok', model_id='m4') +# ============================================================ +# Cluster-wide Processor Quota Tests +# ============================================================ + + +class TestProcessorQuota: + + @pytest.mark.asyncio + async def test_two_server_states_share_one_token_limit(self): + backend = FakeBackend() + states = [ServerState(backend=backend), ServerState(backend=backend)] + + async def reserve(i: int) -> bool: + try: + await states[i % 2].reserve_processor_quota( + 'token', f'p{i}', f's{i}', limit=3, lease_seconds=30.0) + return True + except ResourceQuotaExceededError: + return False + + accepted = await asyncio.gather(*(reserve(i) for i in range(12))) + assert sum(accepted) == 3 + + @pytest.mark.asyncio + async def test_reservation_is_idempotent_and_release_frees_slot(self): + state = ServerState(backend=FakeBackend()) + await state.reserve_processor_quota('token', 'p1', 's1', limit=1, lease_seconds=30.0) + await state.reserve_processor_quota('token', 'p1', 's1', limit=1, lease_seconds=30.0) + with pytest.raises(ResourceQuotaExceededError) as exc: + await state.reserve_processor_quota('token', 'p2', 's2', limit=1, lease_seconds=30.0) + assert exc.value.error_code == 429 + assert exc.value.category.value == 'user' + + await state.release_processor_quota('token', 'p1') + await state.release_processor_quota('token', 'p1') + await state.reserve_processor_quota('token', 'p2', 's2', limit=1, lease_seconds=30.0) + + @pytest.mark.asyncio + async def test_expired_worker_lease_is_reclaimed(self): + state = ServerState(backend=FakeBackend()) + with mock.patch('twinkle.server.state.server_state.time.time', return_value=100.0): + await state.reserve_processor_quota('token', 'dead', 's1', limit=1, lease_seconds=10.0) + with mock.patch('twinkle.server.state.server_state.time.time', return_value=111.0): + assert await state.renew_processor_quota('token', 'dead', lease_seconds=10.0) is False + await state.reserve_processor_quota('token', 'replacement', 's2', limit=1, lease_seconds=10.0) + + # ============================================================ # SamplingSessionManager Tests # ============================================================ @@ -328,7 +384,7 @@ async def test_cleanup_expired_by_age(self, manager): await manager.add('samp_new', record2) cutoff = time.time() - 100 - removed = await manager.cleanup_expired(cutoff) + removed = await manager.cleanup_expired(cutoff, []) assert removed == 1 assert await manager.get('samp_old') is None assert await manager.get('samp_new') is not None diff --git a/tests/server/state/test_redis_integration.py b/tests/server/state/test_redis_integration.py index e730703a1..ed6d5fd4f 100644 --- a/tests/server/state/test_redis_integration.py +++ b/tests/server/state/test_redis_integration.py @@ -132,55 +132,16 @@ async def test_model_write_visible(make_state) -> None: @pytest.mark.asyncio -async def test_session_and_config(make_state) -> None: +async def test_session_shared_across_states(make_state) -> None: a = make_state() b = make_state() sid = await a.create_session({'session_id': f'sess-{uuid.uuid4().hex[:6]}'}) assert await b.get_session_last_heartbeat(sid) is not None - await a.add_config('feature_flag', {'value': 42}) - assert await b.get_config('feature_flag') == {'value': 42} - # ---------- Concurrent-write consistency --------------------------------- # -@pytest.mark.asyncio -async def test_concurrent_config_writes_no_torn_records(make_state) -> None: - """Many concurrent writes of distinct keys complete and every record - equals one of the writes (no torn / partial value).""" - a = make_state() - b = make_state() - n = 40 - payload = {f'k-{i}': {'idx': i, 'note': 'x' * 32} for i in range(n)} - - async def writer(state: ServerState, items: dict) -> None: - await asyncio.gather(*(state.add_config(k, v) for k, v in items.items())) - - half = list(payload.items())[:n // 2] - other = list(payload.items())[n // 2:] - await asyncio.gather(writer(a, dict(half)), writer(b, dict(other))) - - # Every key must read back equal to its expected payload from either side. - for k, v in payload.items(): - assert await a.get_config(k) == v, k - assert await b.get_config(k) == v, k - - -@pytest.mark.asyncio -async def test_concurrent_same_key_lands_one_of_committed(make_state) -> None: - """Two writers race on the same key — final value equals one of the - writes; no torn record.""" - a = make_state() - b = make_state() - write_a = {'who': 'a', 'payload': list(range(8))} - write_b = {'who': 'b', 'payload': list(range(8, 16))} - - await asyncio.gather(a.add_config('contended', write_a), b.add_config('contended', write_b)) - final = await a.get_config('contended') - assert final in (write_a, write_b) - - @pytest.mark.asyncio async def test_concurrent_replica_registration(make_state) -> None: a = make_state() @@ -198,9 +159,9 @@ async def test_concurrent_replica_registration(make_state) -> None: # ---------- Manager-level atomic-update guarantees ----------------------- # # # These tests pin the contract that the manager-level RMW paths -# (``SessionManager.touch``, ``ConfigManager.add_or_get``, -# ``FutureManager.store_status``) now go through ``StateBackend.update_atomic`` -# or ``set_nx``, so a concurrent retry cannot lose a freshly committed write. +# (``SessionManager.touch`` and ``FutureManager.store_status``) go through +# ``StateBackend.update_atomic``, so a concurrent retry cannot lose a freshly +# committed write. @pytest.mark.asyncio @@ -241,29 +202,6 @@ async def hammer(state: ServerState) -> None: assert final >= start, 'final heartbeat predates the test start — every write was lost' -@pytest.mark.asyncio -async def test_concurrent_add_or_get_consistent_value(make_state) -> None: - """``ConfigManager.add_or_get`` is implemented on top of ``set_nx``, - which is atomic in Redis. Two writers racing distinct values for the - same key must return the *same* committed value.""" - a = make_state() - b = make_state() - key = f'cfg-{uuid.uuid4().hex[:6]}' - write_a = {'who': 'a'} - write_b = {'who': 'b'} - - got_a, got_b = await asyncio.gather( - a.add_or_get_config(key, write_a), - b.add_or_get_config(key, write_b), - ) - # Both calls must observe the same committed value — that's the whole - # point of the SETNX-backed contract. - assert got_a == got_b - final = await a.get_config(key) - assert final == got_a - assert final in (write_a, write_b) - - @pytest.mark.asyncio async def test_concurrent_future_update_no_state_regression(make_state) -> None: """Once a future is recorded as ``completed`` a concurrent ``pending`` diff --git a/tests/server/state/test_update_atomic.py b/tests/server/state/test_update_atomic.py index 012e7f237..d03b3a79e 100644 --- a/tests/server/state/test_update_atomic.py +++ b/tests/server/state/test_update_atomic.py @@ -1,7 +1,6 @@ -"""Cross-backend tests for ``StateBackend.update_atomic`` and ``set_nx(ttl)``. +"""Cross-backend tests for the Ray-actor and Redis state backends. -Exercises five contracts against each of the three production backends -(Memory, File, Redis): +Exercises atomic updates, leases, close semantics, and logical key prefixes: - read-transform-write returns the new value - ``transform`` returning ``None`` is a no-op and returns the existing value - ``ttl`` shapes the new value's expiry @@ -15,14 +14,15 @@ import asyncio import functools +import json import os import pytest import pytest_asyncio -import tempfile import uuid from typing import Any -from twinkle.server.state.backend.file_backend import FileBackend +from twinkle.server.deployment import twinkle_server_error_handler +from twinkle.server.state.backend.base import ConcurrencyError from twinkle.server.state.backend.memory_backend import RayActorBackend REDIS_URL = os.environ.get('TWINKLE_TEST_REDIS_URL', 'redis://localhost:6379/0') @@ -71,13 +71,6 @@ def _replace_with(current: Any | None, *, value: Any) -> Any: return value -def _file_backend() -> FileBackend: - f = tempfile.NamedTemporaryFile(suffix='.json', delete=False) - f.close() - os.unlink(f.name) - return FileBackend(f.name) - - def _redis_backend(): from twinkle.server.state.backend.redis_backend import RedisBackend @@ -89,12 +82,6 @@ def memory_backend() -> RayActorBackend: return RayActorBackend() -@pytest.fixture -def file_backend(): - backend = _file_backend() - yield backend - - @pytest_asyncio.fixture async def redis_backend(): backend = _redis_backend() @@ -105,6 +92,43 @@ async def redis_backend(): await backend.close() +@pytest_asyncio.fixture(params=['memory', 'redis']) +async def backend_pair(request): + prefix = f'twinkle-contract-{uuid.uuid4().hex[:8]}::' + if request.param == 'redis': + if not _REDIS_AVAILABLE_AT_COLLECTION: + pytest.skip(f'Redis at {REDIS_URL} unreachable') + from twinkle.server.state.backend.redis_backend import RedisBackend + first = RedisBackend(REDIS_URL, key_prefix=prefix) + second = RedisBackend(REDIS_URL, key_prefix=prefix) + else: + first = RayActorBackend(key_prefix=prefix) + second = RayActorBackend(key_prefix=prefix) + yield first, second + for key in await second.keys('*'): + await second.delete(key) + await first.close() + await second.close() + + +# ---------- Shared backend contract -------------------------------------- # + + +@pytest.mark.asyncio +async def test_close_releases_handle_without_dropping_shared_state(backend_pair) -> None: + first, second = backend_pair + await first.set('shared', {'value': 1}) + await first.close() + assert await second.get('shared') == {'value': 1} + + +@pytest.mark.asyncio +async def test_key_prefix_is_hidden_from_logical_keys(backend_pair) -> None: + first, _ = backend_pair + await first.set('session::one', 1) + assert await first.keys('session::*') == ['session::one'] + + # ---------- Memory backend ----------------------------------------------- # @@ -158,40 +182,6 @@ async def test_memory_set_nx_with_ttl(memory_backend) -> None: assert await memory_backend.set_nx('lease', 'next', ttl=1) is True -# ---------- File backend ------------------------------------------------- # - - -@pytest.mark.asyncio -async def test_file_update_atomic_read_transform_write(file_backend) -> None: - await file_backend.set('k', 5) - result = await file_backend.update_atomic('k', functools.partial(_increment_or_init, delta=3)) - assert result == 8 - assert await file_backend.get('k') == 8 - - -@pytest.mark.asyncio -async def test_file_update_atomic_none_is_noop(file_backend) -> None: - await file_backend.set('k', 42) - result = await file_backend.update_atomic('k', _no_op) - assert result == 42 - - -@pytest.mark.asyncio -async def test_file_update_atomic_respects_ttl(file_backend) -> None: - await file_backend.update_atomic('leased', functools.partial(_replace_with, value='holder'), ttl=1) - assert await file_backend.get('leased') == 'holder' - await asyncio.sleep(1.1) - assert await file_backend.get('leased') is None - - -@pytest.mark.asyncio -async def test_file_set_nx_with_ttl(file_backend) -> None: - assert await file_backend.set_nx('lease', 'owner', ttl=1) is True - assert await file_backend.set_nx('lease', 'other', ttl=1) is False - await asyncio.sleep(1.1) - assert await file_backend.set_nx('lease', 'next', ttl=1) is True - - # ---------- Redis backend ------------------------------------------------ # @@ -232,3 +222,39 @@ async def test_redis_update_atomic_respects_ttl(redis_backend) -> None: assert await redis_backend.get('leased') == 'holder' await asyncio.sleep(1.5) assert await redis_backend.get('leased') is None + + +@_redis_skip +@pytest.mark.asyncio +async def test_redis_retry_exhaustion_maps_to_http_503(redis_backend) -> None: + import redis + from starlette.requests import Request + + sync_client = redis.Redis.from_url(REDIS_URL, decode_responses=True) + real_key = redis_backend._make_key('always-contended') + calls = 0 + + def force_watch_conflict(current: Any | None) -> int: + nonlocal calls + calls += 1 + sync_client.set(real_key, json.dumps(calls)) + return int(current or 0) + 1 + + try: + with pytest.raises(ConcurrencyError) as caught: + await redis_backend.update_atomic('always-contended', force_watch_conflict) + finally: + sync_client.close() + + assert calls == 16 + request = Request({'type': 'http', 'method': 'POST', 'path': '/', 'headers': []}) + request.state.request_id = 'req-contention' + response = await twinkle_server_error_handler(request, caught.value) + payload = json.loads(response.body) + assert response.status_code == 503 + assert payload == { + 'error': "update_atomic exhausted 16 retries on key 'always-contended'", + 'category': 'server', + 'error_code': 503, + 'request_id': 'req-contention', + } diff --git a/tests/server/static/__init__.py b/tests/server/static/__init__.py new file mode 100644 index 000000000..85b3e739d --- /dev/null +++ b/tests/server/static/__init__.py @@ -0,0 +1 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. diff --git a/tests/server/static/backend_call_exemptions.py b/tests/server/static/backend_call_exemptions.py new file mode 100644 index 000000000..1d0307ffe --- /dev/null +++ b/tests/server/static/backend_call_exemptions.py @@ -0,0 +1,32 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Shared exemption list for the "no direct backend call" static checks. + +This file is the SINGLE source of allowed Blocking_Backend_Call bypasses. It is +consumed by this spec's check (``test_no_direct_backend_call.py``) and is intended +to be consumed unchanged by the ``server-request-lifecycle`` spec's equivalent +check -- there must be exactly one physical copy, not one per spec. + +Each entry is ``(module_relpath, function_name)`` where ``module_relpath`` is +relative to ``src/twinkle/server`` and ``function_name`` is the innermost enclosing +function of the exempted call. + +The allowed exemptions are: + +- the ray ``Queue.get`` inside ``sample_stream``'s ``_stream_queue``: it bridges the + sampler actor's process boundary and is bounded by a dedicated double-timeout, + not by ``call_backend``; +- the ``.sample_stream_to_queue.remote(...)`` call inside ``sample_stream`` + itself: streaming generation must keep producing while the HTTP response streams, + so it cannot use ``call_backend`` as-is and carries its own double timeout. This is + a ``remote_function`` bypass that the guard now *detects* (via backend-derived + local tracking) and that is *explicitly* accepted here — replacing the previous + "No remote_function call is exempt" claim, which was true only because the guard + could not see this shape. +""" +from __future__ import annotations + +# (module_relpath under src/twinkle/server, innermost enclosing function name) +BACKEND_CALL_EXEMPTIONS: frozenset[tuple[str, str]] = frozenset({ + ('sampler/twinkle_handlers.py', '_stream_queue'), + ('sampler/twinkle_handlers.py', 'sample_stream'), +}) diff --git a/tests/server/static/test_adapter_name_mapping.py b/tests/server/static/test_adapter_name_mapping.py new file mode 100644 index 000000000..6856f099e --- /dev/null +++ b/tests/server/static/test_adapter_name_mapping.py @@ -0,0 +1,71 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Static check: model backend calls map the tenant adapter name (F002 / P002). + +``run_submit`` hands each ``backend_call`` a *tenant-scoped* adapter name +(``owner_id-``). Before that reaches a backend model method it must be +translated by ``ModelManagement.resolve_model_adapter_name`` — which returns +``''`` in full-parameter mode (the empty-string default optimizer group) and the +name unchanged in LoRA mode. The inline forward/backward/step endpoints did this; +the ``*_from_data_plane`` endpoints originally passed the raw tenant name, so a +full-mode data-plane request drove the wrong optimizer group. + +This guard fails if any ``self.call_backend(...)`` in ``model/twinkle_handlers.py`` +passes ``adapter_name`` as the bare local (i.e. unmapped). Passing the raw name +*positionally* (as ``add_adapter_to_model`` does when creating a tenant adapter) +is intentionally not a keyword and is therefore not flagged. +""" +from __future__ import annotations + +import ast +import pathlib + +import twinkle + +_HANDLERS = pathlib.Path(twinkle.__file__).resolve().parent / 'server' / 'model' / 'twinkle_handlers.py' + + +def _is_call_backend(func: ast.AST) -> bool: + """True for ``self.call_backend`` (Attribute ``call_backend`` on Name ``self``).""" + return (isinstance(func, ast.Attribute) and func.attr == 'call_backend' and isinstance(func.value, ast.Name) + and func.value.id == 'self') + + +def _is_mapped(value: ast.AST) -> bool: + """True when the ``adapter_name`` value is wrapped by ``self.resolve_model_adapter_name(...)``.""" + return (isinstance(value, ast.Call) and isinstance(value.func, ast.Attribute) + and value.func.attr == 'resolve_model_adapter_name') + + +def _unmapped_adapter_name_calls(tree: ast.AST) -> list[tuple[int, str]]: + offenders: list[tuple[int, str]] = [] + for node in ast.walk(tree): + if not isinstance(node, ast.Call) or not _is_call_backend(node.func): + continue + for kw in node.keywords: + if kw.arg != 'adapter_name': + continue + # Bare local ``adapter_name`` is the tenant-scoped name and must be mapped. + if isinstance(kw.value, ast.Name) and kw.value.id == 'adapter_name' and not _is_mapped(kw.value): + offenders.append((node.lineno, ast.unparse(kw.value))) + return offenders + + +def test_model_backend_calls_map_adapter_name(): + tree = ast.parse(_HANDLERS.read_text(), filename=str(_HANDLERS)) + offenders = _unmapped_adapter_name_calls(tree) + assert not offenders, ('These self.call_backend(...) sites pass the raw tenant adapter_name instead of ' + f'self.resolve_model_adapter_name(adapter_name): {offenders}') + + +def test_checker_detects_unmapped_adapter_name(): + source = ('async def route(self, body, adapter_name, token):\n' + ' await self.call_backend(self.model.forward, inputs=[], adapter_name=adapter_name)\n') + assert _unmapped_adapter_name_calls(ast.parse(source)) == [(2, 'adapter_name')] + + +def test_checker_allows_mapped_and_positional(): + source = ( + 'async def route(self, body, adapter_name, token):\n' + ' await self.call_backend(self.model.forward, adapter_name=self.resolve_model_adapter_name(adapter_name))\n' + ' await self.call_backend(self.model.add_adapter_to_model, adapter_name, config)\n') + assert _unmapped_adapter_name_calls(ast.parse(source)) == [] diff --git a/tests/server/static/test_client_architecture_imports.py b/tests/server/static/test_client_architecture_imports.py new file mode 100644 index 000000000..6b0e02421 --- /dev/null +++ b/tests/server/static/test_client_architecture_imports.py @@ -0,0 +1,43 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Executable client/server package-boundary contracts.""" +from __future__ import annotations + +import ast +from pathlib import Path + +_ROOT = Path(__file__).parents[3] / 'src' +_ALLOWED_SERVER_IMPORTS = ( + 'twinkle.protocol.types', + 'twinkle.protocol.headers', + 'twinkle.protocol.json_utils', + 'twinkle.protocol.serialize', +) + + +def _imports(path: Path): + tree = ast.parse(path.read_text(), filename=str(path)) + for node in ast.walk(tree): + if isinstance(node, ast.Import): + yield from (alias.name for alias in node.names) + elif isinstance(node, ast.ImportFrom) and node.module: + yield node.module + + +def test_client_does_not_import_server(): + offenders = [] + for path in (_ROOT / 'twinkle_client').rglob('*.py'): + for module in _imports(path): + if module == 'twinkle.server' or module.startswith('twinkle.server.'): + offenders.append(f'{path.relative_to(_ROOT)} -> {module}') + assert offenders == [] + + +def test_server_only_imports_shared_client_contracts(): + offenders = [] + for path in (_ROOT / 'twinkle' / 'server').rglob('*.py'): + for module in _imports(path): + if module == 'twinkle_client' or module.startswith('twinkle_client.'): + if not any(module == allowed or module.startswith(f'{allowed}.') + for allowed in _ALLOWED_SERVER_IMPORTS): + offenders.append(f'{path.relative_to(_ROOT)} -> {module}') + assert offenders == [] diff --git a/tests/server/static/test_no_degraded_path.py b/tests/server/static/test_no_degraded_path.py new file mode 100644 index 000000000..dcac5283e --- /dev/null +++ b/tests/server/static/test_no_degraded_path.py @@ -0,0 +1,68 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Static check: no silent-degradation symbols remain. + +One wildcard search covering eight symbols; each must occur zero times in its scope. +The symbols are matched as identifiers (word boundaries) so that ``nccl_safe_megatron`` +(the retained decorator), the ``twinkle.utils.nccl_safe`` module path, and unrelated +test names like ``test_zero_loss_...`` are not counted. +""" +from __future__ import annotations + +import pathlib +import re + +import twinkle + +_TWINKLE_SRC = pathlib.Path(twinkle.__file__).resolve().parent +_REPO_ROOT = _TWINKLE_SRC.parent.parent # .../src/twinkle -> repo root +_TESTS = _REPO_ROOT / 'tests' +_COOKBOOK = _REPO_ROOT / 'cookbook' +_SELF = pathlib.Path(__file__).resolve() + +# symbol -> compiled identifier pattern. +_IDENT = { + 'safe_loss': re.compile(r'(? str | None: + if not isinstance(node, ast.Attribute): + return None + owner = node.value + if isinstance(owner, ast.Attribute) and owner.attr in ('model', 'sampler'): + return owner.attr + return None + + +def _getattr_backend_method(node: ast.AST) -> str | None: + if not isinstance(node, ast.Call) or not isinstance(node.func, ast.Name) or node.func.id != 'getattr': + return None + if not node.args: + return None + owner = node.args[0] + if isinstance(owner, ast.Attribute) and owner.attr in ('model', 'sampler'): + return owner.attr + return None + + +def _is_backend_derived(node: ast.AST, backend_names: set[str]) -> bool: + """True if *node* is (transitively) the backend or a value bound from it. + + Matches ``self.model`` / ``self.sampler`` and any attribute/subscript chain + rooted at them, plus locals recorded in ``backend_names`` (e.g. from + ``actors = self.sampler._actors`` then ``actor = actors[0]``). This catches the + "bind the private actor list to a local, then call ``.remote()``" bypass that + plain attribute-name matching cannot see. + """ + while isinstance(node, (ast.Attribute, ast.Subscript)): + if isinstance(node, ast.Attribute): + if isinstance(node.value, ast.Name) and node.value.id == 'self' and node.attr in ('model', 'sampler'): + return True + node = node.value + else: + node = node.value + return isinstance(node, ast.Name) and node.id in backend_names + + +class _Collector(ast.NodeVisitor): + + def __init__(self, relpath: str) -> None: + self.relpath = relpath + self.func_stack: list[str] = [] + self.backend_aliases: set[str] = set() + self.backend_derived: set[str] = set() + self.offenders: list[tuple[str, str, int, str]] = [] + + def _visit_func(self, node: ast.AST) -> None: + self.func_stack.append(node.name) + self.generic_visit(node) + self.func_stack.pop() + + visit_FunctionDef = _visit_func + visit_AsyncFunctionDef = _visit_func + + def visit_Assign(self, node: ast.Assign) -> None: + if _getattr_backend_method(node.value) is not None: + self.backend_aliases.update(target.id for target in node.targets if isinstance(target, ast.Name)) + # Track locals bound (transitively) to the private backend actor list, so a + # later ``.remote()`` on them is still counted as a backend call. + if _is_backend_derived(node.value, self.backend_derived): + self.backend_derived.update(target.id for target in node.targets if isinstance(target, ast.Name)) + self.generic_visit(node) + + def visit_Call(self, node: ast.Call) -> None: + owner = _backend_method(node.func) + label = ast.unparse(node.func) if owner is not None else None + if isinstance(node.func, ast.Name) and node.func.id in self.backend_aliases: + owner = 'alias' + label = node.func.id + if isinstance(node.func, ast.Attribute) and node.func.attr in ('to_thread', 'run_in_executor') and node.args: + escaped_owner = _backend_method(node.args[0]) or _getattr_backend_method(node.args[0]) + if escaped_owner is not None: + owner = escaped_owner + label = f'{ast.unparse(node.func)}({ast.unparse(node.args[0])})' + # ``..remote(...)`` where traces back to the private + # backend actor list bypasses call_backend just like a direct call. + if (owner is None and isinstance(node.func, ast.Attribute) and node.func.attr == 'remote' + and _is_backend_derived(node.func.value, self.backend_derived)): + owner = 'remote' + label = ast.unparse(node.func) + if owner is not None: + enclosing = self.func_stack[-1] if self.func_stack else '' + if (self.relpath, enclosing) not in BACKEND_CALL_EXEMPTIONS: + self.offenders.append((self.relpath, enclosing, node.lineno, label or owner)) + self.generic_visit(node) + + +def test_no_direct_backend_call_in_server(): + offenders: list[tuple[str, str, int, str]] = [] + for path in _SERVER_ROOT.rglob('*.py'): + relpath = str(path.relative_to(_SERVER_ROOT)) + collector = _Collector(relpath) + collector.visit(ast.parse(path.read_text(), filename=str(path))) + offenders.extend(collector.offenders) + + assert not offenders, ('Direct backend calls must go through call_backend (or be listed in ' + f'backend_call_exemptions): {offenders}') + + +def test_exemptions_are_read_from_shared_file(): + assert ('sampler/twinkle_handlers.py', '_stream_queue') in BACKEND_CALL_EXEMPTIONS + + +def test_checker_detects_indirect_backend_calls(): + source = """ +async def route(self): + await asyncio.to_thread(self.model.save) + unload = getattr(self.sampler, 'unload_adapter_paths') + unload([]) +""" + collector = _Collector('example.py') + collector.visit(ast.parse(source)) + assert len(collector.offenders) == 2 + + +def test_checker_allows_call_backend(): + source = """ +async def route(self): + unload = getattr(self.sampler, 'unload_adapter_paths') + await self.call_backend(unload, []) +""" + collector = _Collector('example.py') + collector.visit(ast.parse(source)) + assert collector.offenders == [] + + +def test_checker_detects_aliased_remote_backend_call(): + # The sample_stream bypass shape: bind the private actor list to a local, then + # call .remote() on an element. The guard must see this through the aliasing. + source = """ +async def sample_stream(self): + actors = self.sampler._actors + actor = actors[0] + actor.sample_stream_to_queue.remote(q) +""" + collector = _Collector('example.py') + collector.visit(ast.parse(source)) + assert len(collector.offenders) == 1 + assert collector.offenders[0][1] == 'sample_stream' + + +def test_checker_ignores_remote_on_unrelated_object(): + # A .remote() on a handle not derived from self.model/self.sampler is not a + # backend-boundary bypass and must not be flagged. + source = """ +async def route(self): + handle = get_some_actor() + handle.do.remote(1) +""" + collector = _Collector('example.py') + collector.visit(ast.parse(source)) + assert collector.offenders == [] diff --git a/tests/server/static/test_no_package_root_imports.py b/tests/server/static/test_no_package_root_imports.py new file mode 100644 index 000000000..b14c0f78f --- /dev/null +++ b/tests/server/static/test_no_package_root_imports.py @@ -0,0 +1,61 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Architecture check (F007 / P006): no module imports its own package root. + +A module under ``twinkle.server`` must not import one of its own *ancestor* +sub-packages (e.g. ``checkpoint/tinker.py`` importing ``twinkle.server.checkpoint``). +Such an import is a package-initialisation-order dependency: it re-enters the +ancestor's ``__init__`` while that ``__init__`` is still importing the child, +which is the module-level cycle this check forbids. Importing a *sibling* module +directly (``from .paths import ...``) or the top-level ``twinkle`` public API is +fine and excluded. + +``import-linter``'s ``forbidden`` contract cannot express this (it treats the +source module as part of the forbidden package and reports the contract as kept), +so the check is written directly against grimp's static graph. grimp parses files +without importing the target package, so this stays safe to run in CI. +""" +from __future__ import annotations + +import pathlib +import pytest + +grimp = pytest.importorskip('grimp', reason='grimp is required for the package-root import architecture check') + +import twinkle # noqa: E402 + +_PACKAGE = 'twinkle.server' +_SRC = str(pathlib.Path(twinkle.__file__).resolve().parent.parent) + + +def _package_root_imports() -> list[dict]: + """Return every import where a module imports one of its own ancestor sub-packages.""" + import sys + if _SRC not in sys.path: + sys.path.insert(0, _SRC) + graph = grimp.build_graph('twinkle', include_external_packages=False) + offenders: list[dict] = [] + for module in sorted(graph.modules): + if not module.startswith(_PACKAGE + '.'): + continue + parts = module.split('.') + ancestors = {'.'.join(parts[:i]) for i in range(1, len(parts))} + ancestors = {a for a in ancestors if a.startswith(_PACKAGE + '.')} + for imported in graph.find_modules_directly_imported_by(module): + if imported in ancestors: + for detail in graph.get_import_details(importer=module, imported=imported): + offenders.append({ + 'importer': module, + 'imported': imported, + 'line_number': detail.get('line_number'), + 'line_contents': (detail.get('line_contents') or '').strip(), + }) + return offenders + + +def test_no_module_imports_its_own_package_root(): + offenders = _package_root_imports() + assert not offenders, ( + 'These modules import one of their own ancestor sub-packages (a package-init cycle); ' + 'import the sibling module directly instead:\n' + + '\n'.join(f" {o['importer']} -> {o['imported']} (L{o['line_number']}): {o['line_contents']}" + for o in offenders)) diff --git a/tests/server/static/test_no_twinkle_http_exception.py b/tests/server/static/test_no_twinkle_http_exception.py new file mode 100644 index 000000000..28a67a392 --- /dev/null +++ b/tests/server/static/test_no_twinkle_http_exception.py @@ -0,0 +1,42 @@ +from __future__ import annotations + +import ast +from pathlib import Path + +_SERVER = Path(__file__).resolve().parents[3] / 'src' / 'twinkle' / 'server' +_TWINKLE_HANDLERS = ( + _SERVER / 'gateway' / 'twinkle_handlers.py', + _SERVER / 'model' / 'twinkle_handlers.py', + _SERVER / 'sampler' / 'twinkle_handlers.py', + _SERVER / 'processor' / 'twinkle_handlers.py', +) + + +def _status_code(call: ast.Call) -> int | None: + if call.args and isinstance(call.args[0], ast.Constant) and isinstance(call.args[0].value, int): + return call.args[0].value + for keyword in call.keywords: + if keyword.arg == 'status_code' and isinstance(keyword.value, ast.Constant): + return keyword.value.value if isinstance(keyword.value.value, int) else None + return None + + +def test_twinkle_handlers_have_no_http_exception_bypass_except_iterator_410() -> None: + violations: list[str] = [] + allowed_410 = 0 + for path in _TWINKLE_HANDLERS: + tree = ast.parse(path.read_text(encoding='utf-8'), filename=str(path)) + for node in ast.walk(tree): + if not isinstance(node, ast.Raise) or not isinstance(node.exc, ast.Call): + continue + func = node.exc.func + if not isinstance(func, ast.Name) or func.id != 'HTTPException': + continue + status_code = _status_code(node.exc) + if path.parent.name == 'processor' and status_code == 410: + allowed_410 += 1 + else: + violations.append(f'{path.relative_to(_SERVER)}:{node.lineno} status={status_code}') + + assert allowed_410 == 1 + assert violations == [] diff --git a/tests/server/static/test_utils_bucket_is_light.py b/tests/server/static/test_utils_bucket_is_light.py new file mode 100644 index 000000000..6e970622b --- /dev/null +++ b/tests/server/static/test_utils_bucket_is_light.py @@ -0,0 +1,34 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""The ``twinkle.server.utils`` bucket must re-export only the two +dependency-light helpers, not the queue / session-resource machinery. + +The original intent was to assert ``import twinkle.server.utils`` pulls no +OpenTelemetry. That side-effect is dominated by the parent package's eager +``twinkle.server.__init__`` -> ``launcher`` -> ``application_spec`` -> ``task_queue.config`` +chain (which triggers ``task_queue/__init__`` -> ``mixin`` -> telemetry) and is out of +this Requirement's scope, so we assert the directly-controlled property instead: the +bucket's re-export surface. Re-exporting the mixins is what used to make every one of the +five light callers drag in the OpenTelemetry SDK. +""" +import twinkle.server.utils as bucket + +# The two genuinely dependency-free helpers the five call sites actually use. +_LIGHT_EXPORTS = ('get_template_for_model', 'wrap_builder_with_device_group_env') +# The heavy machinery that must no longer be re-exported through the bucket. +_HEAVY_EXPORTS = ( + 'TaskQueueMixin', + 'TaskQueueConfig', + 'RateLimiter', + 'QueueState', + 'TaskStatus', + 'SessionResourceMixin', + 'AdapterManagerMixin', + 'ProcessorManagerMixin', +) + + +def test_bucket_exports_only_light_helpers(): + for name in _LIGHT_EXPORTS: + assert hasattr(bucket, name), f'{name} should be re-exported by the utils bucket' + leaked = [name for name in _HEAVY_EXPORTS if hasattr(bucket, name)] + assert not leaked, f'utils bucket must not re-export heavy machinery: {leaked}' diff --git a/tests/server/telemetry/test_metrics_cache_invalidation.py b/tests/server/telemetry/test_metrics_cache_invalidation.py new file mode 100644 index 000000000..4f1c0fdf0 --- /dev/null +++ b/tests/server/telemetry/test_metrics_cache_invalidation.py @@ -0,0 +1,20 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Per-deployment metric adapters must not survive ``MetricsRegistry.reset()``. + +``ensure_telemetry_initialized`` resets the registry *in order to* rebind instruments to a +real MeterProvider; an adapter cached at module level survived that and kept recording into +NoOp instruments for the life of the process. This test fails before (module-level +cache) and passes once the caches live on the registry instance. +""" +from twinkle.server.telemetry.metrics import MetricsRegistry, get_task_metrics + + +def test_task_metrics_rebind_after_registry_reset(): + first = get_task_metrics('Model') + MetricsRegistry.reset() + second = get_task_metrics('Model') + assert first is not second + # The point is the *instrument*, not the wrapper identity: assert the bound + # instrument object differs, so a cheap "return a new wrapper around the same + # instrument" implementation cannot pass. + assert first.execution_seconds._instrument is not second.execution_seconds._instrument diff --git a/tests/server/test_deployment_exception_boundary.py b/tests/server/test_deployment_exception_boundary.py index 4f61b42ee..0ed0c63cf 100644 --- a/tests/server/test_deployment_exception_boundary.py +++ b/tests/server/test_deployment_exception_boundary.py @@ -34,9 +34,45 @@ async def boom(): response = client.get('/boom', headers={'x-request-id': 'boundary-test'}) assert response.status_code == 500 assert response.headers['X-Twinkle-Replica-Id'] == 'replica-test' - assert 'Traceback' in response.json()['detail'] - assert 'RuntimeError: boom with replica header' in response.json()['detail'] + # Unhandled exceptions now return the unified ErrorPayload (Server category + # keeps the traceback) instead of the legacy {'detail': } shape. + body = response.json() + assert body['category'] == 'server' + assert body['error_code'] == 500 + assert body['error'] == 'boom with replica header' + assert 'Traceback' in body['traceback'] + assert 'RuntimeError: boom with replica header' in body['traceback'] response = client.get('/healthz') assert response.status_code == 200 assert response.json() == {'ok': True, 'health_calls': 1} + + +def test_deployment_app_bounds_overlong_unhandled_error(monkeypatch): + + class _ReplicaId: + unique_id = 'replica-test' + + class _Context: + replica_id = _ReplicaId() + + from twinkle.server import deployment + + monkeypatch.setattr(deployment.serve, 'get_replica_context', lambda: _Context()) + + def register_routes(app: FastAPI, _get_self): + + @app.get('/boom') + async def boom(): + raise RuntimeError(f'first line {"X" * 2048}\nsecond line') + + client = TestClient(build_deployment_app('Test', register_routes)) + response = client.get('/boom', headers={'x-request-id': 'long-error'}) + + assert response.status_code == 500 + body = response.json() + assert len(body['error']) == 1024 + assert '\n' not in body['error'] + assert len(body['traceback']) <= 65536 + assert body['traceback'].endswith('second line\n') + assert body['request_id'] == 'long-error' diff --git a/tests/server/test_gateway_services.py b/tests/server/test_gateway_services.py new file mode 100644 index 000000000..5a5d032d8 --- /dev/null +++ b/tests/server/test_gateway_services.py @@ -0,0 +1,93 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +from __future__ import annotations + +import ast +import asyncio +from pathlib import Path + +from twinkle.server.gateway import use_cases as services + + +class _State: + + def __init__(self, records): + self.records = list(records) + + async def get_future(self, request_id): + return self.records.pop(0) + + +async def _no_sleep(_seconds): + return None + + +def test_poll_future_returns_canonical_terminal_record(monkeypatch): + monkeypatch.setattr(services, 'long_poll_window', lambda: 10) + monkeypatch.setattr(services, 'retrieve_poll_interval', lambda: 0) + monkeypatch.setattr(services.asyncio, 'sleep', _no_sleep) + state = _State([None, {'status': 'running'}, {'status': 'completed', 'result': None}]) + + outcome = asyncio.run(services.poll_future(state, 'request-1')) + + assert outcome.timed_out is False + assert outcome.record == {'status': 'completed', 'result': None} + + +def test_gateway_services_do_not_import_protocol_models(): + path = Path(services.__file__) + tree = ast.parse(path.read_text(), filename=str(path)) + imports = {node.module for node in ast.walk(tree) if isinstance(node, ast.ImportFrom) and node.module} + assert not any(module == 'tinker.types' or module.startswith('twinkle.protocol.types') for module in imports) + + +class _FakeGateway: + """Minimal ``GatewayServer`` stand-in; ``poll_future`` is patched so state is unused.""" + + state = None + + +def _parity_client(monkeypatch, canonical_record): + """Register both real wire adapters on one app, feeding both the same canonical record.""" + from fastapi import FastAPI + from fastapi.testclient import TestClient + + from twinkle.server.gateway import tinker_handlers, twinkle_handlers + + async def _fixed_poll(_state, _request_id): + return services.FuturePollResult(record=canonical_record, timed_out=False) + + app = FastAPI() + monkeypatch.setattr(tinker_handlers, 'poll_future', _fixed_poll) + monkeypatch.setattr(twinkle_handlers, 'poll_future', _fixed_poll) + tinker_handlers._register_gateway_tinker_routes(app, lambda: _FakeGateway()) + twinkle_handlers._register_gateway_twinkle_routes(app, lambda: _FakeGateway()) + return TestClient(app) + + +def test_same_completed_record_diverges_into_protocol_specific_wire_shapes(monkeypatch): + """One canonical ``completed`` record -> Tinker raw result vs Twinkle TaskEnvelope.""" + client = _parity_client(monkeypatch, {'status': 'completed', 'result': {'loss': 1.0}}) + + tinker_resp = client.post('/retrieve_future', json={'request_id': 'r1'}) + twinkle_resp = client.post('/twinkle/retrieve_future', json={'request_id': 'r1'}) + + assert tinker_resp.status_code == 200 + assert twinkle_resp.status_code == 200 + # Tinker returns the raw result payload; Twinkle wraps it in a canonical envelope. + assert tinker_resp.json() == {'loss': 1.0} + twinkle_body = twinkle_resp.json() + assert twinkle_body['request_id'] == 'r1' + assert twinkle_body['status'] == 'completed' + assert twinkle_body['result'] == {'loss': 1.0} + + +def test_completed_null_result_keeps_the_two_protocols_divergent(monkeypatch): + """The load-bearing difference: null result is a 500 for Tinker but a valid 200 envelope for Twinkle.""" + client = _parity_client(monkeypatch, {'status': 'completed', 'result': None}) + + tinker_resp = client.post('/retrieve_future', json={'request_id': 'r1'}) + twinkle_resp = client.post('/twinkle/retrieve_future', json={'request_id': 'r1'}) + + assert tinker_resp.status_code == 500 + assert twinkle_resp.status_code == 200 + assert twinkle_resp.json()['status'] == 'completed' diff --git a/tests/server/test_runtime.py b/tests/server/test_runtime.py new file mode 100644 index 000000000..84e4837db --- /dev/null +++ b/tests/server/test_runtime.py @@ -0,0 +1,42 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Directed tests for ``server/runtime.init_twinkle_runtime``. + +The processor deployment layer has no test under ``tests/server/`` and reuses this +function for it, so this pins the parameter trap it introduces: ``ncpu_proc_per_node`` is +forwarded only when provided (processor), never when unset (model/sampler), and the +DeviceMesh is built the same way as before (``mesh_dim_names`` -> ``DeviceMesh(**)`` else +``DeviceMesh.from_sizes(**)``). +""" +from __future__ import annotations + +from unittest import mock + +from twinkle.server.runtime import init_twinkle_runtime + + +def test_model_sampler_path_does_not_forward_ncpu_proc_per_node(): + with mock.patch('twinkle.initialize') as init, mock.patch('twinkle.DeviceMesh') as mesh: + mesh.from_sizes.return_value = 'MESH' + result = init_twinkle_runtime(False, 2, device_group='DG', device_mesh_dict={'sizes': [2]}) + assert result == 'MESH' + kwargs = init.call_args.kwargs + assert 'ncpu_proc_per_node' not in kwargs + assert kwargs['nproc_per_node'] == 2 and kwargs['groups'] == ['DG'] + mesh.from_sizes.assert_called_once_with(sizes=[2]) + + +def test_processor_path_forwards_ncpu_proc_per_node_and_named_mesh(): + with mock.patch('twinkle.initialize') as init, mock.patch('twinkle.DeviceMesh') as mesh: + mesh.return_value = 'NAMED_MESH' + result = init_twinkle_runtime( + False, 1, device_group='DG', device_mesh_dict={'mesh_dim_names': ['dp']}, ncpu_proc_per_node=8) + assert result == 'NAMED_MESH' + assert init.call_args.kwargs['ncpu_proc_per_node'] == 8 + mesh.assert_called_once_with(mesh_dim_names=['dp']) + + +def test_mock_backend_returns_none_and_uses_single_cpu_proc(): + with mock.patch('twinkle.initialize') as init, mock.patch('twinkle.DeviceMesh'): + result = init_twinkle_runtime(True, 1, device_group='DG', device_mesh_dict={}) + assert result is None + assert init.call_args.kwargs['ncpu_proc_per_node'] == 1 diff --git a/tests/server/utils/task_queue/test_config.py b/tests/server/utils/task_queue/test_config.py index 3126c507c..60787f567 100644 --- a/tests/server/utils/task_queue/test_config.py +++ b/tests/server/utils/task_queue/test_config.py @@ -13,7 +13,7 @@ from hypothesis import strategies as st from pydantic import ValidationError -from twinkle.server.utils.task_queue.config import TaskQueueConfig +from twinkle.server.task_queue.config import TaskQueueConfig # ---------- defaults snapshot used by the default-value test -------------- # @@ -22,6 +22,7 @@ 'tps_limit': 16000.0, 'window_seconds': 1.0, 'queue_timeout': 300.0, + 'execution_timeout': 1800.0, 'token_cleanup_interval': 60.0, 'max_input_tokens': 16000, } @@ -113,3 +114,15 @@ def test_extra_field_rejected() -> None: """``extra='forbid'`` rejects unknown keys.""" with pytest.raises(ValidationError): TaskQueueConfig(unknown_field=1) + + +def test_zero_execution_timeout_uses_finite_fallback() -> None: + assert TaskQueueConfig(execution_timeout=0).effective_execution_timeout == 3600 + + +def test_absolute_future_ttl_uses_conservative_backend_bound() -> None: + config = TaskQueueConfig(queue_timeout=10, execution_timeout=20) + assert config.absolute_future_ttl(collect_width=2) == 2 * (10 + 2 * 3600) + + config = TaskQueueConfig(queue_timeout=10, execution_timeout=5000) + assert config.absolute_future_ttl(collect_width=2) == 2 * (10 + 2 * 5000) diff --git a/tests/server/utils/test_rate_limiter.py b/tests/server/utils/test_rate_limiter.py new file mode 100644 index 000000000..435b43615 --- /dev/null +++ b/tests/server/utils/test_rate_limiter.py @@ -0,0 +1,58 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +from __future__ import annotations + +import pytest + +from twinkle.server.task_queue.rate_limiter import RateLimiter + + +class _RecordingGauge: + + def __init__(self) -> None: + self.calls: list[tuple[int, dict[str, str]]] = [] + + def set(self, value: int, *, tags: dict[str, str]) -> None: + self.calls.append((value, tags)) + + +@pytest.mark.asyncio +async def test_zero_rps_disables_request_limit() -> None: + limiter = RateLimiter(rps_limit=0, tps_limit=100) + + assert await limiter.check_and_record('token', 1) == (True, None) + assert await limiter.check_and_record('token', 1) == (True, None) + + +@pytest.mark.asyncio +async def test_zero_tps_disables_token_limit() -> None: + limiter = RateLimiter(rps_limit=100, tps_limit=0) + + assert await limiter.check_and_record('token', 1_000_000) == (True, None) + assert await limiter.check_and_record('token', 1_000_000) == (True, None) + + +@pytest.mark.asyncio +async def test_active_token_metric_distinguishes_replicas() -> None: + gauge = _RecordingGauge() + first = RateLimiter( + rps_limit=100, + tps_limit=100, + active_tokens_gauge=gauge, + deployment_name='model', + replica_id='replica-a', + ) + second = RateLimiter( + rps_limit=100, + tps_limit=100, + active_tokens_gauge=gauge, + deployment_name='model', + replica_id='replica-b', + ) + + assert await first.check_and_record('token-a', 1) == (True, None) + assert await second.check_and_record('token-b', 1) == (True, None) + + assert gauge.calls == [ + (1, {'deployment': 'model', 'replica': 'replica-a'}), + (1, {'deployment': 'model', 'replica': 'replica-b'}), + ] diff --git a/tests/server/utils/test_task_errors.py b/tests/server/utils/test_task_errors.py deleted file mode 100644 index 0c4986649..000000000 --- a/tests/server/utils/test_task_errors.py +++ /dev/null @@ -1,10 +0,0 @@ -from twinkle.server.utils.task_errors import task_error_payload - - -def test_task_error_payload_keeps_lora_traceback(): - error = 'Traceback...\nRuntimeError: No lora available for tenant session-default. Max loras: 3\n' - - assert task_error_payload(error) == { - 'error': error, - 'category': 'Server', - } diff --git a/tests/server/utils/test_task_queue_mixin.py b/tests/server/utils/test_task_queue_mixin.py index f0bdbf963..24743f750 100644 --- a/tests/server/utils/test_task_queue_mixin.py +++ b/tests/server/utils/test_task_queue_mixin.py @@ -1,19 +1,31 @@ import asyncio - import pytest -from twinkle.server.utils.task_queue.config import TaskQueueConfig -from twinkle.server.utils.task_queue.mixin import TaskQueueMixin -from twinkle.server.utils.task_queue.worker import ComputeWorker +from twinkle.server.task_queue.config import TaskQueueConfig +from twinkle.server.task_queue.mixin import TaskQueueMixin +from twinkle.server.task_queue.types import UserTaskError +from twinkle.server.task_queue.worker import ComputeWorker class _DummyState: def __init__(self): self.records = [] + self._latest = {} async def store_future_status(self, *args, **kwargs): self.records.append((args, kwargs)) + request_id, status = args[0], args[1] + self._latest[request_id] = { + 'status': status, + 'result': kwargs.get('result'), + 'failure': (kwargs['failure'].model_dump() if kwargs.get('failure') is not None else None), + 'queue_state': kwargs.get('queue_state'), + 'queue_state_reason': kwargs.get('queue_state_reason'), + } + + async def get_future(self, request_id): + return self._latest.get(request_id) class _AllowingRateLimiter: @@ -26,7 +38,9 @@ class _DummyQueue(TaskQueueMixin): def __init__(self): self.state = _DummyState() - self._task_queue_config = TaskQueueConfig() + # A generous Inline_Fast_Path window keeps "task settles inside submit" + # deterministic for the trivial in-process coroutines used here. + self._task_queue_config = TaskQueueConfig(inline_fast_path_timeout=5.0) self._rate_limiter = _AllowingRateLimiter() self._task_metrics = None self._deployment_name = 'test' @@ -44,21 +58,20 @@ def enable_compute_worker(self): @pytest.mark.asyncio async def test_preflight_rejects_batch_without_per_dp_multiple(): queue = _DummyQueue() + from twinkle.server.exceptions import BatchSizeError - result = await queue._perform_preflight_checks( - request_id='req1', - model_id='model1', - token='token1', - input_tokens=0, - batch_size=2, - data_world_size=2, - batch_size_multiple=2, - ) + with pytest.raises(BatchSizeError, match='must be divisible by 4'): + await queue._perform_preflight_checks( + model_id='model1', + token='token1', + input_tokens=0, + batch_size=2, + data_world_size=2, + batch_size_multiple=2, + ) - assert result == {'request_id': 'req1', 'model_id': 'model1'} - _, kwargs = queue.state.records[-1] - assert kwargs['result']['category'] == 'User' - assert 'Batch size 2 must be divisible by 4' in kwargs['result']['error'] + # A rejection writes no future record. + assert queue.state.records == [] @pytest.mark.asyncio @@ -66,7 +79,6 @@ async def test_preflight_accepts_batch_with_per_dp_multiple(): queue = _DummyQueue() result = await queue._perform_preflight_checks( - request_id='req1', model_id='model1', token='token1', input_tokens=0, @@ -93,12 +105,14 @@ async def work(): await asyncio.sleep(0) assert [args[1] for args, _ in queue.state.records] == ['running', 'completed'] + assert queue.state.records[0][1]['absolute_deadline'] > 0 assert queue.state.records[-1][1]['result'] == {'ok': True} @pytest.mark.asyncio -async def test_schedule_task_and_wait_returns_large_result_without_persisting_it(): +async def test_submit_and_peek_returns_completed_envelope_and_persists(): queue = _DummyQueue() + queue.replica_id = 'replica-1' queue.enable_compute_worker() result = {'logps': [[float(index) for index in range(128)]]} @@ -106,7 +120,7 @@ async def work(): return result try: - actual = await queue.schedule_task_and_wait( + env = await queue.submit_and_peek( work, model_id='model1', token='token1', @@ -115,13 +129,16 @@ async def work(): finally: await queue._compute_worker.stop() - assert actual is result - assert queue.state.records == [] + assert env.status == 'completed' + assert env.result == result + # The future record is now the single delivery channel: the result IS persisted. + assert any(args[1] == 'completed' for args, _ in queue.state.records) @pytest.mark.asyncio async def test_polling_schedule_task_still_persists_its_result(): queue = _DummyQueue() + queue.replica_id = 'replica-1' queue.enable_compute_worker() result = {'value': 42} @@ -131,52 +148,106 @@ async def work(): try: await queue.schedule_task(work, model_id='model1', token='token1') for _ in range(100): - completed = [ - kwargs - for args, kwargs in queue.state.records - if args[1] == 'completed' - ] + completed = [kwargs for args, kwargs in queue.state.records if args[1] == 'completed'] if completed: break await asyncio.sleep(0) finally: await queue._compute_worker.stop() + pending = next(kwargs for args, kwargs in queue.state.records if args[1] == 'pending') + assert pending['replica_id'] == 'replica-1' + assert pending['absolute_deadline'] > 0 assert completed[-1]['result'] is result @pytest.mark.asyncio -async def test_schedule_task_and_wait_propagates_failure_without_persisting_it(): +async def test_submit_and_peek_failure_returns_failed_envelope_and_persists(): queue = _DummyQueue() + queue.replica_id = 'replica-1' queue.enable_compute_worker() async def work(): raise ValueError('model failed') try: - with pytest.raises(RuntimeError, match='ValueError: model failed'): - await queue.schedule_task_and_wait( - work, - model_id='model1', - token='token1', - task_type='forward_backward', - ) + env = await queue.submit_and_peek( + work, + model_id='model1', + token='token1', + task_type='forward_backward', + ) finally: await queue._compute_worker.stop() - assert queue.state.records == [] + # The failure payload rides the envelope's `error` field. + assert env.status == 'failed' + assert env.error is not None and 'model failed' in env.error.error + assert any(args[1] == 'failed' for args, _ in queue.state.records) + + +@pytest.mark.asyncio +async def test_user_task_error_is_stored_as_user_failure(): + queue = _DummyQueue() + queue.enable_compute_worker() + + async def work(): + raise UserTaskError('invalid request') + + try: + await queue.schedule_task(work, model_id='model1', token='token1') + for _ in range(100): + failed = [kwargs for args, kwargs in queue.state.records if args[1] == 'failed'] + if failed: + break + await asyncio.sleep(0) + finally: + await queue._compute_worker.stop() + + failure = failed[-1]['failure'] + assert failure.reason_code == 'request_rejected' + assert failure.attribution == 'user' + assert failure.diagnostic is None + + +@pytest.mark.asyncio +async def test_typed_server_error_keeps_its_domain_reason_and_attribution(): + """A typed user failure must retain its domain meaning in persisted state.""" + from twinkle.server.exceptions import ResourceNotFoundError + + queue = _DummyQueue() + queue.enable_compute_worker() + + async def work(): + raise ResourceNotFoundError('adapter foo not found') + + try: + await queue.schedule_task(work, model_id='model1', token='token1') + for _ in range(100): + failed = [kwargs for args, kwargs in queue.state.records if args[1] == 'failed'] + if failed: + break + await asyncio.sleep(0) + finally: + await queue._compute_worker.stop() + + failure = failed[-1]['failure'] + assert failure.reason_code == 'resource_not_found' + assert failure.attribution == 'user' + assert failure.diagnostic is None @pytest.mark.asyncio -async def test_schedule_task_and_wait_reports_preflight_failure_without_persisting_it(): +async def test_submit_and_peek_preflight_rejection_raises_without_writing_or_queuing(): + from twinkle.server.exceptions import BatchSizeError queue = _DummyQueue() queue.enable_compute_worker() async def work(): raise AssertionError('preflight rejection must not execute the task') - with pytest.raises(RuntimeError, match='Batch size 2 must be divisible by 4'): - await queue.schedule_task_and_wait( + with pytest.raises(BatchSizeError, match='must be divisible by 4'): + await queue.submit_and_peek( work, model_id='model1', token='token1', @@ -187,3 +258,19 @@ async def work(): assert queue.state.records == [] assert queue._compute_worker._worker_task is None + + +@pytest.mark.asyncio +async def test_submit_and_peek_honors_explicit_request_id(): + # run_submit generates the request_id up front (to claim the seq dedup key + # atomically) and threads it through submit_and_peek; the future record and the + # returned envelope must use that exact id, not a freshly generated one. + queue = _DummyQueue() + queue.enable_compute_worker() + + async def _factory(): + return {'ok': True} + + env = await queue.submit_and_peek(_factory, task_type='step', request_id='req_fixed123') + assert env.request_id == 'req_fixed123' + assert env.status == 'completed' diff --git a/tests/server/validation/__init__.py b/tests/server/validation/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/server/validation/test_preflight.py b/tests/server/validation/test_preflight.py new file mode 100644 index 000000000..2912b1e8e --- /dev/null +++ b/tests/server/validation/test_preflight.py @@ -0,0 +1,243 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Preflight: everything decidable about a request before it is enqueued. + +The property under test is not "an error is returned" but *where* it is returned. A +rejection that happens inside the queued task has already written a future record and +fanned the call out to every data-parallel rank; a rejection in preflight has done +neither. Each test therefore asserts on the side effects (future records, backend calls) +as well as the status code. +""" +from __future__ import annotations + +import pytest +from types import SimpleNamespace + +from twinkle.server.exceptions import EndpointUnavailableError, RequestRejectedError +from twinkle.server.lifecycle.submit import backend_kwargs, run_submit +from twinkle.server.validation import BackendCapability, assert_request_supported +from twinkle.server.validation.backend_compat import resolve_backend +from twinkle.protocol.types import model as model_types + + +class _Deployment: + """A deployment stub that records what the request managed to reach.""" + + def __init__(self, backend: str = 'transformers'): + self.backend = backend + self.data_world_size = 1 + self.futures: dict[str, dict] = {} + self.backend_calls: list = [] + self.claimed: list = [] + self._task_queue_config = SimpleNamespace(effective_execution_timeout=60.0) + self.state = SimpleNamespace( + claim_seq=self._claim_seq, + get_future=self._get_future, + release_seq=self._release_seq, + ) + + async def _claim_seq(self, dedup_key, request_id, ttl): + self.claimed.append(dedup_key) + return None + + async def _get_future(self, request_id): + return self.futures.get(request_id) + + async def _release_seq(self, dedup_key): + self.claimed.remove(dedup_key) + + async def _on_request_start(self, request): + return 'token' + + async def submit_and_peek(self, task, *, request_id, **kwargs): + self.futures[request_id] = {'status': 'queued'} + return await task() + + +def _request(): + return SimpleNamespace(state=SimpleNamespace(session_id='sess', request_id='rq')) + + +async def _call(self, body, adapter_name, token): + self.backend_calls.append(type(body).__name__) + return {'ok': True} + + +# --------------------------------------------------------------------------- # +# Backend resolution +# --------------------------------------------------------------------------- # + + +def test_backend_is_read_from_the_deployment_not_guessed(): + assert resolve_backend(_Deployment('megatron')) == 'megatron' + # A sampler deployment has no backend concept at all, and must not be invented. + assert resolve_backend(SimpleNamespace()) is None + + +# --------------------------------------------------------------------------- # +# Endpoint capability +# --------------------------------------------------------------------------- # + + +def test_megatron_rejects_the_split_gradient_endpoints_with_501(): + for capability in (BackendCapability.Forward, BackendCapability.Backward, BackendCapability.CalculateLoss): + with pytest.raises(EndpointUnavailableError) as raised: + assert_request_supported( + _Deployment('megatron'), model_types.AdapterRequest(adapter_name='a'), capability=capability) + assert raised.value.error_code == 501 + assert 'forward_backward' in str(raised.value), 'the alternative endpoints must be named' + + +def test_transformers_serves_the_split_gradient_endpoints(): + assert_request_supported( + _Deployment('transformers'), + model_types.AdapterRequest(adapter_name='a'), + capability=BackendCapability.Forward) + + +def test_mock_backend_serves_every_endpoint(): + """The mock is a test double; restricting it would only break tests.""" + assert_request_supported( + _Deployment('mock'), model_types.AdapterRequest(adapter_name='a'), capability=BackendCapability.Forward) + + +@pytest.mark.asyncio +async def test_capability_rejection_writes_no_future_and_calls_no_backend(): + deployment = _Deployment('megatron') + with pytest.raises(EndpointUnavailableError): + await run_submit( + deployment, + _request(), + model_types.AdapterRequest(adapter_name='a'), + task_type='backward', + backend_call=_call, + capability=BackendCapability.Backward) + assert deployment.futures == {}, 'a rejected request must leave no future record' + assert deployment.backend_calls == [], 'the backend must not run for a rejected request' + assert deployment.claimed == [], 'no seq claim may be taken before preflight passes' + + +# --------------------------------------------------------------------------- # +# Backend-only parameters +# --------------------------------------------------------------------------- # + + +def test_a_megatron_only_parameter_is_rejected_on_transformers(): + with pytest.raises(RequestRejectedError) as raised: + assert_request_supported(_Deployment('transformers'), model_types.SaveRequest(adapter_name='a', merge_lora=True)) + assert raised.value.error_code == 422 + assert 'merge_lora' in str(raised.value) + assert 'megatron' in str(raised.value) + + +def test_a_transformers_only_parameter_is_rejected_on_megatron(): + with pytest.raises(RequestRejectedError): + assert_request_supported( + _Deployment('megatron'), model_types.LoadRequest(adapter_name='a', name='ckpt', strict=True)) + + +def test_an_unset_backend_only_parameter_is_not_rejected(): + """Only a *sent* value is checked. + + This is why every restricted field is ``Optional[...] = None``: had one carried its + backend's own default, it would look sent on every request and the other half of the + fleet would reject everything. + """ + assert_request_supported(_Deployment('transformers'), model_types.SaveRequest(adapter_name='a')) + assert_request_supported(_Deployment('megatron'), model_types.LoadRequest(adapter_name='a', name='ckpt')) + + +def test_every_restricted_field_is_optional_with_a_none_default(): + from twinkle.protocol.types.base import FieldRole, fields_with_role, read_backend_only + offenders = [] + for model_cls in vars(model_types).values(): + if not isinstance(model_cls, type) or not hasattr(model_cls, 'model_fields'): + continue + for name, info in fields_with_role(model_cls, FieldRole.BackendKwarg).items(): + if read_backend_only(info) and info.get_default() is not None: + offenders.append(f'{model_cls.__name__}.{name}') + assert not offenders, ('these backend-restricted fields carry a non-None default, so "non-None means wrongly ' + f'targeted" would reject every request on the other backend: {offenders}') + + +# --------------------------------------------------------------------------- # +# Passthrough keys are forwarded, not judged +# --------------------------------------------------------------------------- # + + +def test_a_real_parameter_read_from_kwargs_is_not_rejected(): + """The case that removed the passthrough spelling check. + + ``InputProcessor`` declares ``padding_free`` and reads ``padding_side`` via + ``kwargs.get`` -- both are real parameters, and ``inspect.signature`` only sees the + first. A similarity check scored them 0.75 and rejected + ``set_processor('InputProcessor', padding_side='right')``, a call used throughout the + cookbook and the E2E suite. No threshold fixes that: it also has to catch + ``bate`` -> ``beta`` at 0.5. Rejecting valid requests is worse than missing a typo, so + passthrough contents are forwarded unjudged. + """ + assert_request_supported( + _Deployment(), + model_types.SetProcessorRequest( + processor_cls='InputProcessor', adapter_name='a', init_kwargs={'padding_side': 'right'})) + + +def test_an_unrecognised_plugin_argument_is_forwarded(): + """Plugins accept ``**kwargs``, so "unknown" cannot mean "wrong".""" + assert_request_supported( + _Deployment(), + model_types.SetLossRequest(loss_cls='DPOLoss', adapter_name='a', init_kwargs={'my_custom_knob': 1})) + + +def test_no_plugin_download_happens_during_validation(monkeypatch): + """A validation path must have no side effects, and resolving a remote id downloads.""" + import twinkle.utils.loader as loader + monkeypatch.setattr(loader.Plugin, 'load_plugin', + lambda *a, **k: pytest.fail('validation must not download a plugin')) + assert_request_supported( + _Deployment(), + model_types.SetLossRequest(loss_cls='ms://someone/MyLoss', adapter_name='a', init_kwargs={'beta': 0.1})) + + +# --------------------------------------------------------------------------- # +# Forwarding +# --------------------------------------------------------------------------- # + + +def test_control_fields_are_never_forwarded_to_the_backend(): + """``inputs`` / ``adapter_name`` / ``seq_id`` are already passed explicitly. + + Forwarding them again would duplicate a keyword argument, or leak a protocol field + into a backend signature. + """ + body = model_types.ForwardBackwardTaskRequest( + inputs=[{'input_ids': [1, 2]}], adapter_name='a', seq_id=3, task='embedding') + assert backend_kwargs(body) == {'task': 'embedding'} + + +def test_unset_backend_parameters_are_not_forwarded(): + body = model_types.ForwardRequest(inputs=[{'input_ids': [1]}], adapter_name='a') + assert backend_kwargs(body) == {} + + +def test_passthrough_contents_are_flattened(): + body = model_types.ForwardBackwardTaskRequest( + inputs=[{'input_ids': [1]}], adapter_name='a', loss_kwargs={'advantages': [0.5]}) + assert backend_kwargs(body) == {'advantages': [0.5]} + + +def test_a_passthrough_key_shadowing_a_declared_parameter_is_an_error(): + """Silently letting one win would make the effective value depend on merge order.""" + body = model_types.ForwardRequest(inputs=[{'input_ids': [1]}], adapter_name='a', task='causal_lm', + loss_kwargs={'task': 'embedding'}) + with pytest.raises(ValueError, match='collides'): + backend_kwargs(body) + + +def test_plugin_identifiers_are_not_forwarded_as_kwargs(): + """The handler passes ``loss_cls`` positionally; forwarding it too would duplicate it.""" + body = model_types.SetLossRequest(loss_cls='DPOLoss', adapter_name='a', init_kwargs={'beta': 0.1}) + assert backend_kwargs(body) == {'beta': 0.1} + + +if __name__ == '__main__': + raise SystemExit(pytest.main([__file__, '-v'])) diff --git a/tests/server/validation/test_request_wire.py b/tests/server/validation/test_request_wire.py new file mode 100644 index 000000000..73759d25d --- /dev/null +++ b/tests/server/validation/test_request_wire.py @@ -0,0 +1,248 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Level 1: what a malformed request body looks like on the wire. + +Driven through FastAPI's ``TestClient`` against the real route table, so what is asserted +is the response a client actually receives -- not a schema call in isolation. A rejection +here happens during body parsing, before any handler runs, which is what makes it free of +queue and backend side effects. +""" +from __future__ import annotations + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient +from pydantic import ValidationError + +from fastapi.exceptions import RequestValidationError + +from twinkle.server.deployment import validation_error_handler +from twinkle.protocol.types import model as model_types +from twinkle.protocol.types.base import StrictRequest + + +@pytest.fixture(scope='module') +def client() -> TestClient: + """An app carrying the real request models and the shared error handler. + + Handlers are stubs on purpose: the point is that a bad body never reaches one, so a + stub that records nothing is the strongest possible witness -- if it is invoked, the + check did not happen. + """ + app = FastAPI() + app.add_exception_handler(RequestValidationError, validation_error_handler) + + @app.post('/forward') + async def forward(body: model_types.ForwardRequest): + return {'reached_handler': True} + + @app.post('/set_loss') + async def set_loss(body: model_types.SetLossRequest): + return {'reached_handler': True} + + @app.post('/save') + async def save(body: model_types.SaveRequest): + return {'reached_handler': True} + + @app.post('/forward_backward') + async def forward_backward(body: model_types.ForwardBackwardTaskRequest): + return {'reached_handler': True} + + return TestClient(app) + + +def _post(client: TestClient, path: str, body: dict): + return client.post(path, json=body) + + +def _valid_forward(**overrides) -> dict: + return {'inputs': [{'input_ids': [1, 2, 3]}], 'adapter_name': 'a', **overrides} + + +def test_a_valid_body_reaches_the_handler(client): + response = _post(client, '/forward', _valid_forward()) + assert response.status_code == 200 + assert response.json() == {'reached_handler': True} + + +def test_an_unknown_top_level_field_is_rejected(client): + response = _post(client, '/forward', _valid_forward(adapter_nmae='typo')) + assert response.status_code == 422 + body = response.json() + assert body['category'] == 'user' + assert body['error_code'] == 422 + assert 'adapter_nmae' in body['error'] + assert any(detail['field'] == 'adapter_nmae' for detail in body['details']) + assert 'reached_handler' not in body + + +def test_the_error_names_the_client_version_mismatch(client): + """An unknown top-level field is exactly what an outdated client looks like.""" + response = _post(client, '/forward', _valid_forward(advantages=[0.1])) + assert response.status_code == 422 + assert 'upgrade' in response.json()['error'].lower() + + +def test_overlong_validation_summary_is_bounded_without_losing_details(client): + field = 'unknown_' + ('X' * 2048) + response = _post(client, '/forward', _valid_forward(**{field: 1})) + + assert response.status_code == 422 + body = response.json() + assert len(body['error']) == 1024 + assert body['category'] == 'user' + assert body['details'][0]['field'] == field + assert 'traceback' not in body + + +def test_the_error_body_is_an_error_payload_not_fastapi_detail(client): + """One error shape on the wire, or a client has to learn two.""" + body = _post(client, '/forward', _valid_forward(unknown=1)).json() + assert set(body) >= {'error', 'category', 'error_code', 'request_id'} + assert 'detail' not in body + + +def test_no_traceback_is_returned(client): + """A rejected body is the caller's problem, not a crash to be dumped at them.""" + assert 'traceback' not in _post(client, '/forward', _valid_forward(unknown=1)).json() + + +def test_details_locate_the_field_inside_the_body(client): + body = _post(client, '/forward', {'inputs': [{'input_ids': [1.5]}], 'adapter_name': 'a'}).json() + assert body['error_code'] == 422 + assert any('inputs' in detail['path'] for detail in body['details']) + + +def test_a_token_in_the_body_is_rejected(client): + """``token`` comes from the Authorization header only. + + Declaring it as a field would make a body-supplied token *legal* under + ``extra='forbid'``, which is a credential-forgery path, not a convenience. + """ + assert _post(client, '/forward', _valid_forward(token='stolen')).status_code == 422 + assert 'token' not in model_types.ForwardRequest.model_fields + + +def test_seq_id_is_still_accepted(client): + """Strictness must not break the retry idempotency key. + + Were ``seq_id`` undeclared, ``extra='forbid'`` would reject every retried + gradient-mutating call -- disabling the dedup that prevents a double-apply. It is + declared on the gradient-mutating models only, which is where the client sends it. + """ + assert _post(client, '/forward_backward', _valid_forward(seq_id=4)).status_code == 200 + assert 'seq_id' in model_types.ForwardBackwardTaskRequest.model_fields + assert 'seq_id' in model_types.AdapterRequest.model_fields + assert 'seq_id' in model_types.DataPlaneForwardRequest.model_fields + + +def test_a_dynamic_plugin_argument_is_accepted_inside_its_region(client): + """Strictness at the top level, freedom inside the declared dict.""" + response = _post(client, '/set_loss', { + 'loss_cls': 'DPOLoss', + 'adapter_name': 'a', + 'init_kwargs': {'beta': 0.1, 'anything_at_all': [1, 2]}, + }) + assert response.status_code == 200 + + +def test_a_non_json_value_fails_in_the_client_before_any_request(): + """``JsonValue`` earns its keep at Level 0, not Level 1. + + Anything that arrived as JSON is by definition a JSON value, so this annotation can + only ever reject something in the caller's own process -- which is the useful place, + because the caller still has the offending object and a stack trace pointing at it. + """ + from pydantic import ValidationError + with pytest.raises(ValidationError) as raised: + model_types.SetLossRequest(loss_cls='DPOLoss', adapter_name='a', init_kwargs={'beta': object()}) + assert 'init_kwargs' in str(raised.value), 'the error must name the field that holds the bad value' + + +def test_a_checkpoint_dict_without_its_key_is_a_422_not_a_500(client): + """It used to surface as ``KeyError`` -> 500, blaming the server for a bad body.""" + response = _post(client, '/save', {'adapter_name': 'a', 'checkpoint_dir': {'wrong': 'shape'}}) + assert response.status_code == 422 + + +# --------------------------------------------------------------------------- # +# Coverage: no twinkle-native route may keep a lax body +# --------------------------------------------------------------------------- # + +# The one twinkle-native route with no request body at all. +_BODYLESS = {('model', 'GET /healthz')} + + +def test_every_twinkle_route_body_is_strict(): + """Enumerated from the live route table, not from a hardcoded count. + + A count would have to be updated by whoever adds a route -- exactly the person who + would also forget the base class. + """ + from fastapi.routing import APIRoute + from tests.server.contract.client_api_harness import build_model_app, build_processor_app, build_sampler_app + + offenders = [] + for app_name, builder in (('model', build_model_app), ('sampler', build_sampler_app), ('processor', + build_processor_app)): + for route in builder().routes: + if not isinstance(route, APIRoute): + continue + for method in sorted(route.methods & {'GET', 'POST', 'PUT', 'PATCH', 'DELETE'}): + key = f'{method} {route.path}' + if not route.path.startswith(('/twinkle', '/healthz')) or (app_name, key) in _BODYLESS: + continue + for field in route.dependant.body_params: + annotation = field.field_info.annotation + if not (isinstance(annotation, type) and issubclass(annotation, StrictRequest)): + offenders.append(f'{app_name} {key}: {getattr(annotation, "__name__", annotation)}') + assert not offenders, ('these twinkle-native routes accept a body that is not a StrictRequest, so an unknown ' + f'field reaches the backend instead of failing: {offenders}') + + +# --------------------------------------------------------------------------- # +# Regression: the sampler routes must bind the sampler-domain models, not model.py's +# +# ``sampler.py`` and ``model.py`` once both declared bare ``AddAdapterRequest`` / +# ``SetTemplateRequest``. Because the handler does ``import twinkle.protocol.types as +# types`` and the package ``__init__`` re-exported ``model.py`` first, +# ``types.AddAdapterRequest`` resolved to *model.py*'s model -- whose ``config`` is a +# ``str``, so it rejected the dict a real ``add_adapter_to_sampler`` call sends. The +# Sampler-prefixed names remove the collision; this pins the binding so it cannot +# silently regress. +# --------------------------------------------------------------------------- # + + +def _body_model(app, path: str): + from fastapi.routing import APIRoute + for route in app.routes: + if isinstance(route, APIRoute) and route.path == path: + (param, ) = route.dependant.body_params + return param.field_info.annotation + raise AssertionError(f'route {path} not found') + + +def test_sampler_routes_bind_sampler_domain_models(): + from tests.server.contract.client_api_harness import build_sampler_app + from twinkle.protocol.types import sampler as sampler_types + + app = build_sampler_app() + assert _body_model(app, '/twinkle/add_adapter_to_sampler') is sampler_types.SamplerAddAdapterRequest + assert _body_model(app, '/twinkle/set_template') is sampler_types.SamplerSetTemplateRequest + + +def test_sampler_add_adapter_accepts_a_dict_config_where_model_rejects_it(): + """The exact divergence the collision hid: the client sends ``config`` as a dict.""" + from twinkle.protocol.types import model as model_types + from twinkle.protocol.types import sampler as sampler_types + + # The sampler contract (``config: Any``) accepts the LoRA config dict the client sends. + ok = sampler_types.SamplerAddAdapterRequest(adapter_name='a', config={'r': 8}) + assert ok.config == {'r': 8} + # model.py's same-shaped-name model declares ``config: Optional[str]`` and would 422 it, + # which is why the two must not share a bare class name. + with pytest.raises(ValidationError): + model_types.AddAdapterRequest(adapter_name='a', config={'r': 8}) + + +if __name__ == '__main__': + raise SystemExit(pytest.main([__file__, '-v'])) diff --git a/tests/server/validation/test_wire_schema.py b/tests/server/validation/test_wire_schema.py new file mode 100644 index 000000000..09b18e7d7 --- /dev/null +++ b/tests/server/validation/test_wire_schema.py @@ -0,0 +1,209 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Wire schema behaviour for the inline ``inputs`` data plane. + +These assertions are made against the schema directly, not through a mock backend: the +mock accepts ``**kwargs`` without inspecting anything, so "the mock did not complain" +is evidence of nothing about validation. +""" +from __future__ import annotations + +import pytest +from pydantic import TypeAdapter, ValidationError +from typing import Union, get_args, get_origin + +from twinkle.data_format.encoding import ENCODED_INPUT_KEYS, is_encoded +from twinkle.processor.base import InputProcessor +from twinkle.protocol.types import data as wire + +_INPUTS = TypeAdapter(wire.WireInputBatch) + + +def _parse(payload): + return _INPUTS.validate_python(payload) + + +# --------------------------------------------------------------------------- # +# Classification and homogeneity +# --------------------------------------------------------------------------- # + + +def test_encoded_entry_parses_as_input_feature(): + (entry, ) = _parse([{'input_ids': [1, 2, 3]}]) + assert isinstance(entry, wire.WireInputFeature) + + +def test_embedding_only_entry_is_encoded_not_a_trajectory(): + """An embedding-only batch is already encoded. + + Reading it as a ``Trajectory`` would send it through ``template.batch_encode`` and + fail far from the cause, which is why ``input_embedding`` is part of the shared + predicate rather than only ``input_ids``. + """ + (entry, ) = _parse([{'input_embedding': [[0.1, 0.2]]}]) + assert isinstance(entry, wire.WireInputFeature) + + +def test_message_entry_parses_as_trajectory(): + (entry, ) = _parse([{'messages': [{'role': 'user', 'content': 'hi'}]}]) + assert isinstance(entry, wire.WireTrajectory) + + +def test_mixed_batch_is_rejected(): + with pytest.raises(ValidationError): + _parse([{'input_ids': [1]}, {'messages': [{'role': 'user', 'content': 'x'}]}]) + + +def test_entry_without_a_required_key_is_rejected(): + with pytest.raises(ValidationError): + _parse([{'labels': [1, 2]}]) + + +@pytest.mark.parametrize('payload', [[1, 2, 3], 'text', 42, True, None]) +def test_non_object_inputs_are_rejected(payload): + with pytest.raises(ValidationError): + _parse(payload) + + +def test_a_single_entry_is_accepted_as_a_one_element_batch(): + assert len(_parse({'input_ids': [1, 2]})) == 1 + + +def test_required_key_rule_matches_the_shared_predicate(): + """The schema's required-key rule and ``is_encoded`` must stay the same rule. + + If they drifted, an entry could be an ``InputFeature`` to one and a ``Trajectory`` + to the other -- the exact divergence the single shared predicate exists to prevent. + """ + for key in ENCODED_INPUT_KEYS: + entry = {key: [1] if key == 'input_ids' else [[0.5]]} + assert is_encoded(entry) + assert isinstance(_parse([entry])[0], wire.WireInputFeature) + assert not is_encoded({'messages': []}) + + +# --------------------------------------------------------------------------- # +# Strictness on declared numeric fields +# --------------------------------------------------------------------------- # + + +def test_bool_tokens_are_rejected(): + """Lax ``int`` would coerce ``[true, false]`` to ``[1, 0]`` and train on it.""" + with pytest.raises(ValidationError): + _parse([{'input_ids': [True, False]}]) + + +def test_float_tokens_are_rejected(): + """Values come from a tensor's ``tolist()``; a ``1.0`` there is an upstream defect.""" + with pytest.raises(ValidationError): + _parse([{'input_ids': [1.0]}]) + + +def test_negative_labels_are_accepted(): + (entry, ) = _parse([{'input_ids': [1, 2], 'labels': [-100, 5]}]) + assert entry.labels == [-100, 5] + + +def test_float_vlm_values_are_accepted(): + (entry, ) = _parse([{'input_ids': [1], 'pixel_values': [[0.5, 0.25]]}]) + assert entry.pixel_values == [[0.5, 0.25]] + + +@pytest.mark.parametrize('position_ids', [[0, 1], [[0, 1], [2, 3]], [[[0, 1]]]]) +def test_position_ids_accept_one_to_three_dimensions(position_ids): + (entry, ) = _parse([{'input_ids': [1, 2], 'position_ids': position_ids}]) + assert entry.position_ids == position_ids + + +def test_position_ids_reject_a_scalar(): + with pytest.raises(ValidationError): + _parse([{'input_ids': [1, 2], 'position_ids': 0}]) + + +def test_routed_experts_require_exactly_three_dimensions(): + with pytest.raises(ValidationError): + _parse([{'input_ids': [1, 2], 'routed_experts': [0, 1]}]) + + +# --------------------------------------------------------------------------- # +# Extension data and export +# --------------------------------------------------------------------------- # + + +def test_unknown_json_fields_survive_a_round_trip(): + """A preprocessor's leftover columns must reach the backend, not be dropped. + + ``extra='ignore'`` would accept the request and then silently strip these on export, + which loses data the caller sent -- a worse outcome than rejecting it. + """ + payload = {'input_ids': [1, 2], 'source_id': 'row-7', 'score': 0.5} + exported = wire.export(_parse([payload])[0]) + assert exported == payload + + +def test_export_omits_unset_optional_fields(): + """Twinkle_Core branches on key *presence*, so ``None`` must not be emitted.""" + exported = wire.export(_parse([{'input_ids': [1, 2]}])[0]) + assert exported == {'input_ids': [1, 2]} + + +def test_round_trip_is_idempotent(): + samples = [ + {'input_ids': [1, 2, 3]}, + {'input_ids': [[1, 2], [3, 4]]}, + {'input_ids': [1, 2], 'position_ids': [[0, 1]]}, + {'input_ids': [1, 2], 'position_ids': [[[0, 1]]]}, + {'input_ids': [1, 2], 'routed_experts': [[[0, 1]]]}, + {'input_ids': [1, 2], 'labels': [-100, 4]}, + {'input_embedding': [[0.5]]}, + {'messages': [{'role': 'user', 'content': 'x'}], 'user_data': [['k', '"v"']]}, + {'messages': []}, + ] + for sample in samples: + once = wire.export(_parse([sample])[0]) + twice = wire.export(_parse([once])[0]) + assert once == twice, sample + + +# --------------------------------------------------------------------------- # +# Structural invariants +# --------------------------------------------------------------------------- # + + +def _depth(annotation) -> int: + depth = 0 + while get_origin(annotation) is list: + depth += 1 + annotation = get_args(annotation)[0] + return depth + + +@pytest.mark.parametrize('alias', ['Ints1to2', 'Ints1to3', 'Numbers1to2', 'Numbers1to4']) +def test_union_members_are_declared_shallowest_first(alias): + """Deepest-first is ~41x slower on a 2-D input, so the order is load-bearing. + + Asserted structurally rather than by timing: a timing assertion would have to build + the anti-pattern to compare against, take seconds, and could fail on a busy runner. + """ + annotation = getattr(wire, alias) + union = get_args(annotation)[0] # unwrap Annotated + assert get_origin(union) is Union + depths = [_depth(member) for member in get_args(union)] + assert depths == sorted(depths) and len(set(depths)) == len(depths), depths + + +def test_schema_covers_every_key_core_reads(): + missing = wire.CORE_INPUT_KEYS - wire.declared_wire_keys() + assert not missing, f'Twinkle_Core reads these keys but the wire schema drops them: {sorted(missing)}' + + +def test_vlm_field_set_matches_the_processor(): + """Kept as a test, not an import, so this module stays free of Twinkle_Core's deps. + + A field added to ``VLM_CONCAT_FIELDS`` and not here would be silently absent from + the wire while the batching code still expects it. + """ + assert wire.VLM_TENSOR_FIELDS == frozenset(InputProcessor.VLM_CONCAT_FIELDS) + + +if __name__ == '__main__': + raise SystemExit(pytest.main([__file__, '-v'])) diff --git a/tests/twinkle_agentic/evaluator/test_client_sampler.py b/tests/twinkle_agentic/evaluator/test_client_sampler.py index a542f0182..91fec847f 100644 --- a/tests/twinkle_agentic/evaluator/test_client_sampler.py +++ b/tests/twinkle_agentic/evaluator/test_client_sampler.py @@ -1,26 +1,46 @@ from twinkle.data_format import SamplingParams +from twinkle_client.http import ClientContext, ClientTransport from twinkle_client.sampler.vllm_sampler import vLLMSampler -def test_http_sampler_serializes_sampling_params_dataclass_once(monkeypatch): - request = {} +class _Response: + + def __init__(self, payload): + self._payload = payload + self.ok = True + self.status_code = 200 + + def json(self): + return self._payload + - class Response: - def raise_for_status(self): - pass +class _Session: - def json(self): - return {'samples': []} + def __init__(self, request): + self.request = request - def post(*, url, json_data): - request['url'] = url - request['body'] = json_data - return Response() + def post(self, url, *, data=None, json=None, **_kwargs): + self.request['url'] = url + if data is not None: + import json as json_module + self.request['body'] = json_module.loads(data) + else: + self.request['body'] = json + if url.endswith('/create'): + return _Response({}) + return _Response({'request_id': 'test', 'status': 'completed', 'result': {'samples': []}}) - import twinkle_client.sampler.vllm_sampler as module - monkeypatch.setattr(module, 'http_post', post) - sampler = object.__new__(vLLMSampler) - sampler.server_url = 'http://example/sampler/model/twinkle' + def close(self): + pass + + +def test_http_sampler_serializes_sampling_params_dataclass_once(): + request = {} + transport = ClientTransport( + ClientContext(base_url='http://example', api_key='test'), + session=_Session(request), + ) + sampler = vLLMSampler('model', transport=transport) sampler.sample([{'messages': []}], SamplingParams(max_tokens=4, num_samples=2)) assert request['body']['sampling_params']['num_samples'] == 2 assert 'num_samples' not in request['body'] diff --git a/tests/twinkle_agentic/test_vllm_sampler_tq_generation.py b/tests/twinkle_agentic/test_vllm_sampler_tq_generation.py index e64669b59..3b81f387d 100644 --- a/tests/twinkle_agentic/test_vllm_sampler_tq_generation.py +++ b/tests/twinkle_agentic/test_vllm_sampler_tq_generation.py @@ -3,26 +3,22 @@ import asyncio import inspect import json +import pytest import time from concurrent.futures import Future -import pytest - from twinkle import DeviceMesh from twinkle.data_format import SampledSequence, SampleResponse, SamplingParams from twinkle.infra import _dispatch_args from twinkle.server.sampler.twinkle_handlers import _await_generation from twinkle_agentic.async_rl import LoraContext from twinkle_agentic.async_rl.types import PartitionAdmission, PromptGroup, RolloutPolicy -from twinkle_agentic.async_rl.vllm_sampler_tq import ( - VLLMSamplerTQ, - _GeneratedSample, - _PromptGroupRolloutStats, - _dispatch_generation, -) +from twinkle_agentic.async_rl.vllm_sampler_tq import (VLLMSamplerTQ, _dispatch_generation, _GeneratedSample, + _PromptGroupRolloutStats) class LocalActorHandle: + def __init__(self, target): self.target = target @@ -30,6 +26,7 @@ def __getattr__(self, name): method = getattr(self.target, name) class RemoteMethod: + async def remote(_, *args, **kwargs): result = method(*args, **kwargs) return await result if inspect.isawaitable(result) else result @@ -38,6 +35,7 @@ async def remote(_, *args, **kwargs): class PolicyProvider: + def __init__(self, policies): self.policies = iter(policies) self.released = [] @@ -105,10 +103,11 @@ def test_generation_dispatch_allows_one_prompt_with_multiple_dp_workers() -> Non _dispatch_generation( 3, worker_index, - ('submission', [{'input_ids': [1]}], 'params'), + ('submission', [{ + 'input_ids': [1] + }], 'params'), {}, - )[0][1] - for worker_index in range(3) + )[0][1] for worker_index in range(3) ] assert shards == [[{'input_ids': [1]}], [], []] @@ -128,7 +127,9 @@ def submit(coro): result = sampler.submit_generation( 'submission-1', - [{'input_ids': [1]}], + [{ + 'input_ids': [1] + }], SamplingParams(max_tokens=4), ) @@ -155,7 +156,11 @@ async def sample_single(feat, _params, **_kwargs): sampler._sample_single = sample_single responses = asyncio.run( sampler._generate_inputs( - [{'input_ids': [10]}, {'input_ids': [20]}], + [{ + 'input_ids': [10] + }, { + 'input_ids': [20] + }], SamplingParams(max_tokens=4), adapter_name='', adapter_path=None, @@ -214,6 +219,15 @@ def test_native_prompt_group_sampling_requires_context_manager() -> None: sampler.submit_prompt_groups([], SamplingParams(max_tokens=4)) +class _DirectGenerationService: + + def __init__(self, sampler): + self.sampler = sampler + + async def call_backend(self, fn, /, *args, **kwargs): + return await asyncio.to_thread(fn, *args, **kwargs) + + def test_server_waiter_admits_later_submission_before_first_finishes() -> None: class Sampler: @@ -237,14 +251,13 @@ def cancel_generation(self, submission_id): self.futures.pop(submission_id, None) sampler = Sampler() + service = _DirectGenerationService(sampler) async def run(): sampler.submit_generation('first') sampler.submit_generation('second') - first = asyncio.create_task( - _await_generation(sampler, 'first')) - second = asyncio.create_task( - _await_generation(sampler, 'second')) + first = asyncio.create_task(_await_generation(service, 'first', timeout=5)) + second = asyncio.create_task(_await_generation(service, 'second', timeout=5)) while len(sampler.submission_order) < 2: await asyncio.sleep(0) assert not first.done() @@ -284,8 +297,8 @@ def cancel_generation(self, _submission_id): sampler = Sampler() sampler.submit_generation('submission') - result = asyncio.run( - _await_generation(sampler, 'submission')) + service = _DirectGenerationService(sampler) + result = asyncio.run(_await_generation(service, 'submission', timeout=5)) assert result == ['completed'] assert sampler.status_calls == 2 @@ -314,11 +327,11 @@ def test_sampler_reports_submission_throughput_at_partition_or_shard_scope(dp_si context = _context() admission = PartitionAdmission(context, context.partition_id(0), 0, 2, 2, 0) groups = [ - PromptGroup(context, admission, f'{admission.partition_id}/group_{index}', {}, object()) - for index in range(2) + PromptGroup(context, admission, f'{admission.partition_id}/group_{index}', {}, object()) for index in range(2) ] class RolloutMetricsHarness: + def __init__(self): self.device_mesh = DeviceMesh.from_sizes(world_size=dp_size, dp_size=dp_size) self.events = [] @@ -375,28 +388,25 @@ def test_sampler_writes_one_atomic_rollout_file_per_prompt_group(tmp_path): sequences=[SampledSequence('stop', [20 + index], decoded=f'completion-{index}')], prompt_token_ids=[10, 11], ), - (policy,), + (policy, ), attempts=1, was_aborted=False, resumed_partial_output=False, - ) - for index in range(2) - ] - rows = [ - { - 'generation_idx': index, - 'rollout_policy_version': 7, - 'initial_policy_version': 7, - 'final_policy_version': 7, - 'rollout_policy_versions': [7], - 'rollout_adapter_path': '/tmp/adapter-v7', - 'stop_reason': 'stop', - 'logprobs': [-0.1], - } - for index in range(2) + ) for index in range(2) ] + rows = [{ + 'generation_idx': index, + 'rollout_policy_version': 7, + 'initial_policy_version': 7, + 'final_policy_version': 7, + 'rollout_policy_versions': [7], + 'rollout_adapter_path': '/tmp/adapter-v7', + 'stop_reason': 'stop', + 'logprobs': [-0.1], + } for index in range(2)] class Template: + @staticmethod def decode(token_ids, **_kwargs): return ' '.join(map(str, token_ids)) @@ -410,13 +420,8 @@ def decode(token_ids, **_kwargs): sampler._write_rollout_group('submission-2', group, generated, rows, [1.0, 0.0]) output_path = ( - tmp_path - / context.tenant_id - / context.training_run_id - / context.adapter_name - / 'policy_7' - / 'train_3-group_0.jsonl' - ) + tmp_path / context.tenant_id / context.training_run_id / context.adapter_name / 'policy_7' + / 'train_3-group_0.jsonl') records = [json.loads(line) for line in output_path.read_text().splitlines()] assert len(records) == 2 assert records[0]['submission_id'] == 'submission-2' @@ -446,7 +451,10 @@ def test_aborted_generation_restarts_from_original_prompt_when_partial_is_disabl VLLMSamplerTQ._generate_sample( sampler, context, - {'input_ids': [1, 2], 'labels': [-100, -100]}, + { + 'input_ids': [1, 2], + 'labels': [-100, -100] + }, SamplingParams(max_tokens=4, logprobs=1), multi_modal_data=None, logprobs_only=False, @@ -478,7 +486,10 @@ def test_aborted_generation_continues_from_partial_tokens_when_enabled(): VLLMSamplerTQ._generate_sample( sampler, context, - {'input_ids': [1, 2], 'labels': [-100, -100]}, + { + 'input_ids': [1, 2], + 'labels': [-100, -100] + }, SamplingParams(max_tokens=4, logprobs=1), multi_modal_data=None, logprobs_only=False, diff --git a/tests/twinkle_client/test_async_components.py b/tests/twinkle_client/test_async_components.py index 99b91fcc9..0a092627c 100644 --- a/tests/twinkle_client/test_async_components.py +++ b/tests/twinkle_client/test_async_components.py @@ -1,9 +1,18 @@ # Copyright (c) ModelScope Contributors. All rights reserved. +"""Client-side wire shape of the async / data-plane component calls. + +The recorded body is the *serialized request model*, not a hand-built dict: every +twinkle-native client method now instantiates its endpoint's model and posts one +``model_dump_json``. Asserting on the parsed JSON therefore checks the real wire shape, +including the fact that a caller's undeclared arguments land in the model's passthrough +region rather than at the top level. +""" from __future__ import annotations import asyncio +import json -from twinkle_client.types import DataRef +from twinkle.protocol.types import DataRef class _Response: @@ -11,6 +20,7 @@ class _Response: def __init__(self, payload, status_code: int = 200): self._payload = payload self.status_code = status_code + self.ok = status_code < 400 def raise_for_status(self) -> None: if self.status_code >= 400: @@ -20,20 +30,39 @@ def json(self): return self._payload -def test_model_forward_backward_sends_multiple_data_refs(monkeypatch) -> None: - import twinkle_client.http as http_module - from twinkle_client.model import multi_lora_transformers as module +def _completed(result): + """Wrap a business result in a completed Task_Envelope (the new wire shape).""" + return {'request_id': 'req-test', 'status': 'completed', 'result': result} + + +def _recorder(calls, result_factory): + """A Session.post stand-in that records the URL and decoded JSON body.""" + + def post(url, headers=None, data=None, timeout=None, **kwargs): + body = json.loads(data) if data else kwargs.get('json') or {} + calls.append((url, body)) + return _Response(result_factory(url)) + + return post + + +def _patch_transport(monkeypatch, calls, result_factory): + from twinkle_client.http import ClientContext, ClientTransport + from twinkle_client.http.context import set_default_transport + + transport = ClientTransport(ClientContext(base_url='http://server', api_key='test-key')) + monkeypatch.setattr(transport._session, 'post', _recorder(calls, result_factory)) + set_default_transport(transport) - calls = [] - def post(*, url, json_data=None, **_kwargs): - calls.append((url, json_data)) - if url.endswith('/create'): - return _Response({}) - return _Response({'result': {'loss': 1.0}}) +def test_model_forward_backward_sends_multiple_data_refs(monkeypatch) -> None: + from twinkle_client.model import multi_lora_transformers as module - monkeypatch.setattr(http_module, 'get_base_url', lambda: 'http://server/api/v1') - monkeypatch.setattr(module, 'http_post', post) + calls: list = [] + _patch_transport(monkeypatch, calls, lambda url: {} + if url.endswith('/create') else _completed({'result': { + 'loss': 1.0 + }})) model = module.MultiLoraTransformersModel('ms://base') model.adapter_name = 'adapter' @@ -56,17 +85,10 @@ def post(*, url, json_data=None, **_kwargs): def test_model_inline_forward_methods_keep_the_original_endpoints(monkeypatch) -> None: - import twinkle_client.http as http_module from twinkle_client.model import multi_lora_transformers as module - calls = [] - - def post(*, url, json_data=None, **_kwargs): - calls.append((url, json_data)) - return _Response({} if url.endswith('/create') else {'result': {}}) - - monkeypatch.setattr(http_module, 'get_base_url', lambda: 'http://server/api/v1') - monkeypatch.setattr(module, 'http_post', post) + calls: list = [] + _patch_transport(monkeypatch, calls, lambda url: {} if url.endswith('/create') else _completed({'result': {}})) model = module.MultiLoraTransformersModel('ms://base') model.adapter_name = 'adapter' @@ -82,22 +104,61 @@ def post(*, url, json_data=None, **_kwargs): ] assert all(body['inputs'] == inputs for _, body in calls[-3:]) assert all('input_refs' not in body for _, body in calls[-3:]) + # Declared backend parameters stay top level; ``exclude_none`` keeps the rest off. + assert calls[-3][1]['return_logits'] is True + assert calls[-2][1]['disable_lora'] is True + assert calls[-1][1]['micro_batch_size'] == 1 -def test_model_data_plane_forward_uses_a_separate_api(monkeypatch) -> None: - import twinkle_client.http as http_module +def test_model_add_adapter_serializes_dict_lora_config(monkeypatch) -> None: + """Dict input remains compatible with the strict serialized LoRA wire field.""" from twinkle_client.model import multi_lora_transformers as module - calls = [] + calls: list = [] + _patch_transport(monkeypatch, calls, + lambda url: {} if url.endswith('/create') else _completed({'status': 'ok'})) + + model = module.MultiLoraTransformersModel('ms://base') + model.add_adapter_to_model('adapter', {'r': 4, 'target_modules': 'all-linear'}) + + url, body = calls[-1] + assert url.endswith('/model/base/twinkle/add_adapter_to_model') + assert isinstance(body['config'], str) + assert 'LoraConfig' in body['config'] + assert model.adapter_name == 'adapter' - def post(url, json_data=None, **_kwargs): - calls.append((url, json_data)) - return _Response({} if url.endswith('/create') else {'result': {'value': 1}}) - monkeypatch.setattr(http_module, 'get_base_url', lambda: 'http://server/api/v1') - monkeypatch.setattr(module, 'http_post', post) +def test_undeclared_forward_arguments_are_routed_to_loss_kwargs(monkeypatch) -> None: + """A loss input is not a declared field, so it travels in the passthrough region. + + The public signature is unchanged -- callers still pass ``advantages=...`` -- which is + what lets the body be strict without breaking existing scripts. + """ + from twinkle_client.model import multi_lora_transformers as module + + calls: list = [] + _patch_transport(monkeypatch, calls, lambda url: {} if url.endswith('/create') else _completed({'result': {}})) model = module.MultiLoraTransformersModel('ms://base') + model.adapter_name = 'adapter' + model.forward_backward([{'input_ids': [1, 2]}], advantages=[0.5], old_logps=[[-1.0, -2.0]]) + + _, body = calls[-1] + assert body['loss_kwargs'] == {'advantages': [0.5], 'old_logps': [[-1.0, -2.0]]} + assert 'advantages' not in body + + +def test_model_data_plane_forward_uses_a_separate_api(monkeypatch) -> None: + from twinkle_client.model import multi_lora_transformers as module + + calls: list = [] + _patch_transport(monkeypatch, calls, lambda url: {} + if url.endswith('/create') else _completed({'result': { + 'value': 1 + }})) + + model = module.MultiLoraTransformersModel('ms://base') + model.adapter_name = 'adapter' ref = DataRef(ref_id='data-1', size=2, fields=['train_input']) model.forward_from_data_plane(ref, input_field='train_input') @@ -108,21 +169,18 @@ def post(url, json_data=None, **_kwargs): def test_model_data_plane_forward_only_can_append_selected_outputs(monkeypatch) -> None: - import twinkle_client.http as http_module from twinkle_client.model import multi_lora_transformers as module ref = DataRef(ref_id='data-1', size=2, fields=['input_ids']) updated_ref = ref.model_copy(update={'fields': ['input_ids', 'ref_logps']}) - calls = [] - - def post(*, url, json_data=None, **_kwargs): - calls.append((url, json_data)) - return _Response({} if url.endswith('/create') else {'result': updated_ref.model_dump()}) - - monkeypatch.setattr(http_module, 'get_base_url', lambda: 'http://server/api/v1') - monkeypatch.setattr(module, 'http_post', post) + calls: list = [] + # The handler wraps its payload as ``{'result': ...}``, so the stub must too -- + # otherwise the test asserts against a reply shape the server never sends. + _patch_transport(monkeypatch, calls, lambda url: {} + if url.endswith('/create') else _completed({'result': updated_ref.model_dump()})) model = module.MultiLoraTransformersModel('ms://base') + model.adapter_name = 'adapter' result = model.forward_only_from_data_plane( ref, output_ref=ref, @@ -138,7 +196,6 @@ def post(*, url, json_data=None, **_kwargs): def test_sampler_async_data_plane_path_returns_reference_without_materializing(monkeypatch) -> None: - import twinkle_client.http as http_module from twinkle_client.sampler import vllm_sampler as module output_ref = DataRef( @@ -147,20 +204,16 @@ def test_sampler_async_data_plane_path_returns_reference_without_materializing(m fields=['train_input', 'sampled_logprobs', 'decoded'], kind='rollout', ) - calls = [] - - def post(*, url, json_data=None, **_kwargs): - calls.append((url, json_data)) - if url.endswith('/create'): - return _Response({}) - return _Response(output_ref.model_dump()) + calls: list = [] + _patch_transport(monkeypatch, calls, lambda url: {} + if url.endswith('/create') else _completed(output_ref.model_dump())) - monkeypatch.setattr(http_module, 'get_base_url', lambda: 'http://server/api/v1') - monkeypatch.setattr(module, 'http_post', post) sampler = module.vLLMSampler('ms://base') result = asyncio.run(sampler.asample_to_data_plane( - [{'input_ids': [1]}], + [{ + 'input_ids': [1] + }], num_samples=4, group_ids=['group-1'], )) diff --git a/tests/twinkle_client/test_async_rl_workers.py b/tests/twinkle_client/test_async_rl_workers.py index 7a79cb8a5..332f20cda 100644 --- a/tests/twinkle_client/test_async_rl_workers.py +++ b/tests/twinkle_client/test_async_rl_workers.py @@ -1,7 +1,6 @@ from __future__ import annotations import asyncio - import pytest from twinkle_client.async_rl import Worker, WorkerPipeline @@ -59,6 +58,7 @@ async def failure(): def test_worker_pipeline_rejects_duplicate_role_names() -> None: + async def noop(): return None diff --git a/tests/twinkle_client/test_auto_agent_tools.py b/tests/twinkle_client/test_auto_agent_tools.py new file mode 100644 index 000000000..a62ebd731 --- /dev/null +++ b/tests/twinkle_client/test_auto_agent_tools.py @@ -0,0 +1,69 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Focused unit contracts for the auto-agent tool dispatcher.""" +from __future__ import annotations + +import json +import pytest + +from twinkle_client.auto.agent.tool_schemas import TOOL_SCHEMAS as DECLARED_TOOL_SCHEMAS +from twinkle_client.auto.agent.tools import TOOL_SCHEMAS, ToolExecutor + + +class _Connection: + current_run_id: str | None = None + + def list_training_runs(self): + return [{'run_id': 'run-1'}] + + +@pytest.mark.asyncio +async def test_execute_dispatches_and_serializes_result(): + result = json.loads(await ToolExecutor(_Connection()).execute('list_training_runs', {})) + assert result == [{'run_id': 'run-1'}] + + +@pytest.mark.asyncio +async def test_execute_reports_unknown_tool_without_raising(): + result = json.loads(await ToolExecutor(_Connection()).execute('missing', {})) + assert result == {'error': 'Unknown tool: missing'} + + +@pytest.mark.asyncio +async def test_execute_turns_handler_exception_into_error(monkeypatch): + executor = ToolExecutor(_Connection()) + + async def fail(): + raise RuntimeError('broken') + + monkeypatch.setattr(executor, '_tool_list_training_runs', fail) + result = json.loads(await executor.execute('list_training_runs', {})) + assert result == {'error': 'list_training_runs failed: broken'} + + +@pytest.mark.asyncio +async def test_search_dispatch_stays_behind_executor(monkeypatch): + executor = ToolExecutor(_Connection()) + monkeypatch.setattr(executor, '_search_datasets_impl', lambda query, limit: [{'id': query, 'limit': limit}]) + result = await executor._tool_search_datasets('demo', limit=2) + assert result == { + 'query': 'demo', + 'results': [{ + 'id': 'demo', + 'limit': 2 + }], + } + + +@pytest.mark.asyncio +async def test_server_health_helper_stays_behind_executor(monkeypatch): + import urllib.request + + monkeypatch.setattr(urllib.request, 'urlopen', lambda *args, **kwargs: object()) + assert await ToolExecutor(_Connection())._check_server_health('http://server') is True + + +def test_tool_schema_names_are_unique_and_dispatchable(): + assert TOOL_SCHEMAS is DECLARED_TOOL_SCHEMAS + names = [item['function']['name'] for item in TOOL_SCHEMAS] + assert len(names) == len(set(names)) + assert all(hasattr(ToolExecutor, f'_tool_{name}') for name in names) diff --git a/tests/twinkle_client/test_client_multi_turn_rollout.py b/tests/twinkle_client/test_client_multi_turn_rollout.py index f19bea041..cb1440419 100644 --- a/tests/twinkle_client/test_client_multi_turn_rollout.py +++ b/tests/twinkle_client/test_client_multi_turn_rollout.py @@ -7,7 +7,7 @@ ``tests/twinkle_agentic/test_multi_turn_rollout.py`` but adapt the fake sampler to the ``twinkle_client`` HTTP contract: ``FakeClientSampler.sample()`` mirrors ``vLLMSampler.sample()`` and returns ``List[SampleResponseModel]`` (pydantic, -from ``twinkle_client.types.sampler``) whose ``sequences[0]`` carries a populated +from ``twinkle.protocol.types.sampler``) whose ``sequences[0]`` carries a populated ``new_input_feature`` so the multi-turn loop can proceed round after round. Properties covered: @@ -21,19 +21,18 @@ import copy import json +import pytest import re from collections import defaultdict -from typing import Any, Dict, List, Optional - -import pytest from hypothesis import given, settings from hypothesis import strategies as st +from typing import Any, Dict, List, Optional from twinkle.data_format.sampling import SamplingParams from twinkle_agentic.tools.base import Tool from twinkle_agentic.tools.tool_manager import ToolManager from twinkle_client.rollout.multi_turn import ClientMultiTurnRollout -from twinkle_client.types.sampler import SampledSequenceModel, SampleResponseModel +from twinkle.protocol.types.sampler import SampledSequenceModel, SampleResponseModel # ============================================================================= @@ -114,8 +113,7 @@ def __init__(self, tokenizer: FakeTokenizer) -> None: def encode(self, trajectory: Dict[str, Any], add_generation_prompt: bool = False) -> Dict[str, Any]: messages = trajectory.get('messages', []) - s = self.tokenizer.apply_chat_template( - messages, tokenize=False, add_generation_prompt=add_generation_prompt) + s = self.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=add_generation_prompt) input_ids = self.tokenizer.encode(s, add_special_tokens=False) pif: Dict[str, Any] = dict(trajectory) # preserve top-level fields (incl. _tid) pif['input_ids'] = input_ids @@ -416,8 +414,7 @@ def test_logprobs_align_with_trainable_labels(scripts_spec, max_turns): logprobs = out.get('logprobs') if logprobs: trainable = _count_trainable(out.get('labels') or []) - assert len(logprobs) == trainable, ( - f'logprobs({len(logprobs)}) != trainable labels({trainable})') + assert len(logprobs) == trainable, (f'logprobs({len(logprobs)}) != trainable labels({trainable})') # ============================================================================= @@ -450,8 +447,7 @@ def test_max_turns_one_forces_truncation(logprobs_flags): # Every trajectory emits a tool_call on its first (and only allowed) turn. scripts_spec = [{'num_tools': 3, 'terminal': 'stop', 'logprobs': lp} for lp in logprobs_flags] trajectories, sampler, template = _build_from_scripts(scripts_spec) - rollout = ClientMultiTurnRollout( - sampler=sampler, template=template, tool_manager=_make_tool_manager(), max_turns=1) + rollout = ClientMultiTurnRollout(sampler=sampler, template=template, tool_manager=_make_tool_manager(), max_turns=1) outs = rollout(copy.deepcopy(trajectories)) @@ -475,8 +471,7 @@ def test_length_stop_marks_truncated(logprobs_flags): # very first generation is the one that gets cut. scripts_spec = [{'num_tools': 0, 'terminal': 'length', 'logprobs': lp} for lp in logprobs_flags] trajectories, sampler, template = _build_from_scripts(scripts_spec) - rollout = ClientMultiTurnRollout( - sampler=sampler, template=template, tool_manager=_make_tool_manager(), max_turns=4) + rollout = ClientMultiTurnRollout(sampler=sampler, template=template, tool_manager=_make_tool_manager(), max_turns=4) outs = rollout(copy.deepcopy(trajectories)) @@ -546,11 +541,13 @@ def sample(self, inputs, sampling_params=None, **kwargs): def test_missing_new_input_feature_raises_indexed_runtime_error(): """new_input_feature=None -> RuntimeError naming batch AND trajectory index.""" - trajectories, _script_sampler, template = _build_from_scripts( - [{'num_tools': 0, 'terminal': 'stop', 'logprobs': False}]) + trajectories, _script_sampler, template = _build_from_scripts([{ + 'num_tools': 0, + 'terminal': 'stop', + 'logprobs': False + }]) sampler = _NullFeatureSampler(template) - rollout = ClientMultiTurnRollout( - sampler=sampler, template=template, tool_manager=_make_tool_manager(), max_turns=3) + rollout = ClientMultiTurnRollout(sampler=sampler, template=template, tool_manager=_make_tool_manager(), max_turns=3) with pytest.raises(RuntimeError) as excinfo: rollout(copy.deepcopy(trajectories)) @@ -567,10 +564,8 @@ def test_tool_calls_without_tool_manager_raises_value_error(): """tool_calls produced but tool_manager missing -> ValueError.""" # One tool-call turn then a terminal turn; max_turns=2 so the tool-dispatch # site (not the max_turns truncation edge) is what fails. - trajectories, sampler, template = _build_from_scripts( - [{'num_tools': 1, 'terminal': 'stop', 'logprobs': False}]) - rollout = ClientMultiTurnRollout( - sampler=sampler, template=template, tool_manager=None, max_turns=2) + trajectories, sampler, template = _build_from_scripts([{'num_tools': 1, 'terminal': 'stop', 'logprobs': False}]) + rollout = ClientMultiTurnRollout(sampler=sampler, template=template, tool_manager=None, max_turns=2) with pytest.raises(ValueError) as excinfo: rollout(copy.deepcopy(trajectories)) @@ -582,11 +577,9 @@ def test_tool_calls_without_tool_manager_raises_value_error(): def test_tool_calls_without_tool_manager_via_per_call_kwarg_raises_value_error(): """Passing tool_manager=None as a per-call kwarg also raises at dispatch.""" - trajectories, sampler, template = _build_from_scripts( - [{'num_tools': 1, 'terminal': 'stop', 'logprobs': False}]) + trajectories, sampler, template = _build_from_scripts([{'num_tools': 1, 'terminal': 'stop', 'logprobs': False}]) # Constructed WITH a manager, but the per-call override nulls it out. - rollout = ClientMultiTurnRollout( - sampler=sampler, template=template, tool_manager=_make_tool_manager(), max_turns=2) + rollout = ClientMultiTurnRollout(sampler=sampler, template=template, tool_manager=_make_tool_manager(), max_turns=2) with pytest.raises(ValueError): rollout(copy.deepcopy(trajectories), tool_manager=None) @@ -594,12 +587,14 @@ def test_tool_calls_without_tool_manager_via_per_call_kwarg_raises_value_error() def test_sampler_network_error_propagates_unchanged(): """vLLMSampler.sample() network error propagates unchanged (not swallowed/wrapped).""" - trajectories, _script_sampler, template = _build_from_scripts( - [{'num_tools': 0, 'terminal': 'stop', 'logprobs': False}]) + trajectories, _script_sampler, template = _build_from_scripts([{ + 'num_tools': 0, + 'terminal': 'stop', + 'logprobs': False + }]) sentinel = NetworkError('simulated connection reset by peer') sampler = _NetworkFailingSampler(template, sentinel) - rollout = ClientMultiTurnRollout( - sampler=sampler, template=template, tool_manager=_make_tool_manager(), max_turns=3) + rollout = ClientMultiTurnRollout(sampler=sampler, template=template, tool_manager=_make_tool_manager(), max_turns=3) with pytest.raises(NetworkError) as excinfo: rollout(copy.deepcopy(trajectories)) diff --git a/tests/twinkle_client/test_client_orchestrated_grpo.py b/tests/twinkle_client/test_client_orchestrated_grpo.py index 2302e3bfc..796287418 100644 --- a/tests/twinkle_client/test_client_orchestrated_grpo.py +++ b/tests/twinkle_client/test_client_orchestrated_grpo.py @@ -5,12 +5,9 @@ import sys from pathlib import Path -from twinkle_client.types import DataRef +from twinkle.protocol.types import DataRef - -MODULE_PATH = ( - Path(__file__).parents[2] / 'cookbook' / 'client' / 'async_rl' / 'client_orchestrated_grpo.py' -) +MODULE_PATH = (Path(__file__).parents[2] / 'cookbook' / 'client' / 'async_rl' / 'client_orchestrated_grpo.py') def _load_module(): @@ -57,6 +54,7 @@ async def fake_rollout(_sampler, prompt, policy, _semaphore, _group_id): ) class FakeModel: + def __init__(self): self.saved = [] self.steps = 0 @@ -78,6 +76,7 @@ async def calculate_metric(self, **_kwargs): return {'result': {'loss': 1.0 / self.steps, 'grad_norm': 0.5}} class FakeDataPlane: + def __init__(self): self.released = [] @@ -95,9 +94,21 @@ async def run(): model = FakeModel() data_plane = FakeDataPlane() batches = [ - [{'name': 'p0-g0'}, {'name': 'p0-g1'}], - [{'name': 'p1-g0'}, {'name': 'p1-g1'}], - [{'name': 'p2-g0'}, {'name': 'p2-g1'}], + [{ + 'name': 'p0-g0' + }, { + 'name': 'p0-g1' + }], + [{ + 'name': 'p1-g0' + }, { + 'name': 'p1-g1' + }], + [{ + 'name': 'p2-g0' + }, { + 'name': 'p2-g1' + }], ] await module.run_grpo(batches, model, object(), data_plane) return model, data_plane @@ -152,6 +163,7 @@ async def fake_rollout(_sampler, prompt, _policy, _semaphore, _group_id): monkeypatch.setattr(module, 'GRPOAdvantage', lambda: lambda rewards, **_kwargs: [1.0]) class FakeModel: + def __init__(self): self.saved = [] @@ -169,6 +181,7 @@ async def calculate_metric(self, **_kwargs): return {'result': {'loss': 1.0}} class FakeDataPlane: + async def aget(self, ref, *, fields=None): assert fields == ['decoded'] return [{'decoded': ref.ref_id}] @@ -183,7 +196,13 @@ async def run(): model = FakeModel() try: await module.run_grpo( - [[{'name': 'p0'}], [{'name': 'p1'}], [{'name': 'p2'}]], + [[{ + 'name': 'p0' + }], [{ + 'name': 'p1' + }], [{ + 'name': 'p2' + }]], model, object(), FakeDataPlane(), diff --git a/tests/twinkle_client/test_component_rpc.py b/tests/twinkle_client/test_component_rpc.py new file mode 100644 index 000000000..f01b3bb99 --- /dev/null +++ b/tests/twinkle_client/test_component_rpc.py @@ -0,0 +1,56 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +from __future__ import annotations + +from twinkle.dataset import DatasetMeta +from twinkle.protocol.serialize import deserialize_object, serialize_object +from twinkle_client.common.component_rpc import call_remote_component, create_remote_component +from twinkle_client.http import ClientContext, ClientTransport + + +class _Response: + + def __init__(self, payload): + self._payload = payload + self.ok = True + self.status_code = 200 + self.url = 'http://server' + self.text = '' + + def json(self): + return self._payload + + +class _Session: + + def __init__(self): + self.calls = [] + + def post(self, url, **kwargs): + self.calls.append((url, kwargs)) + if url.endswith('/create'): + return _Response({'processor_id': 'pid:1'}) + return _Response({'result': 'ok'}) + + def close(self): + pass + + +def test_component_rpc_uses_transport_url_and_preserves_timeout_semantics() -> None: + session = _Session() + transport = ClientTransport(ClientContext(base_url='http://server', api_key='key'), session=session, timeout=90) + + assert create_remote_component('dataset', 'Dataset', transport=transport) == 'pid:1' + assert call_remote_component('pid:1', 'check', transport=transport) == 'ok' + assert call_remote_component('pid:1', 'check', None, transport=transport) == 'ok' + + assert session.calls[0][0] == 'http://server/api/v1/processor/twinkle/create' + assert session.calls[1][0] == 'http://server/api/v1/processor/twinkle/call' + assert session.calls[1][1]['timeout'] == 90 + assert session.calls[2][1]['timeout'] is None + + +def test_dataset_meta_data_slice_round_trips() -> None: + for data_slice in (range(1, 9, 2), [1, 3, 5]): + restored = deserialize_object(serialize_object(DatasetMeta(dataset_id='demo', data_slice=data_slice))) + assert restored.dataset_id == 'demo' + assert list(restored.data_slice) == list(data_slice) diff --git a/tests/twinkle_client/test_data_plane_async.py b/tests/twinkle_client/test_data_plane_async.py index fd616f347..5be9ebbdf 100644 --- a/tests/twinkle_client/test_data_plane_async.py +++ b/tests/twinkle_client/test_data_plane_async.py @@ -2,12 +2,11 @@ from __future__ import annotations import asyncio -import threading - import pytest +import threading from twinkle_client.data_plane import DataPlaneClient -from twinkle_client.types import DataRef, DataRowsResponse +from twinkle.protocol.types import DataRef, DataRowsResponse def test_async_convenience_methods_delegate_to_sync_operations(monkeypatch) -> None: @@ -46,9 +45,13 @@ async def run(): asyncio.run(run()) assert [call[:-1] for call in calls] == [ - ('put', [{'value': 1}], 'rollout'), + ('put', [{ + 'value': 1 + }], 'rollout'), ('get', original_ref, ['value']), - ('append', original_ref, [{'value': 2}]), + ('append', original_ref, [{ + 'value': 2 + }]), ('release', appended_ref), ] assert all(call[-1] != caller_thread for call in calls) @@ -96,7 +99,11 @@ async def run(): asyncio.run(run()) assert calls == [ - ('put', [{'value': 1}], 'data', tags), + ('put', [{ + 'value': 1 + }], 'data', tags), ('get_batch', ref, None), - ('append', ref, [{'reward': 1.0}], tags), + ('append', ref, [{ + 'reward': 1.0 + }], tags), ] diff --git a/tests/twinkle_client/test_error_parsing.py b/tests/twinkle_client/test_error_parsing.py new file mode 100644 index 000000000..037fd0db5 --- /dev/null +++ b/tests/twinkle_client/test_error_parsing.py @@ -0,0 +1,107 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Client error-response parsing tests.""" +from __future__ import annotations + +import pytest +import requests + +from twinkle_client.exceptions import TwinkleHTTPError +from twinkle_client.http.client import _handle_response + + +class _Resp: + """Minimal stand-in for requests.Response for _handle_response.""" + + def __init__(self, status_code, *, body=None, text='', url='http://x'): + self.status_code = status_code + self.ok = status_code < 400 + self._body = body + self.text = text + self.url = url + + def json(self): + if self._body is None: + raise ValueError('no json') + return self._body + + +def test_structured_error_reads_top_level_fields(): + """Top-level category/error_code/request_id are preferred.""" + resp = _Resp(422, body={'error': 'bad input', 'category': 'user', 'error_code': 422, 'request_id': 'req-7'}) + with pytest.raises(TwinkleHTTPError) as exc: + _handle_response(resp) + assert isinstance(exc.value, requests.HTTPError) # existing except clauses keep working + assert exc.value.status_code == 422 + assert exc.value.error_code == 422 + assert exc.value.category == 'user' + assert exc.value.request_id == 'req-7' + assert exc.value.details is None + assert exc.value.traceback is None + assert 'bad input' in str(exc.value) + + +def test_detail_only_error_falls_back_to_unknown_category(): + """A non-ErrorPayload JSON body uses the lowercase unknown category.""" + resp = _Resp(404, body={'detail': 'Not Found'}) + with pytest.raises(TwinkleHTTPError) as exc: + _handle_response(resp) + assert exc.value.status_code == 404 + assert exc.value.category == 'unknown' + assert exc.value.error_code is None + assert 'Not Found' in str(exc.value) + + +def test_non_json_body_falls_back_to_text(): + resp = _Resp(500, body=None, text='raw traceback text') + with pytest.raises(TwinkleHTTPError) as exc: + _handle_response(resp) + assert exc.value.category == 'unknown' + assert 'raw traceback text' in str(exc.value) + + +def test_validation_details_are_preserved(): + details = [{'loc': ['body', 'items', 0], 'msg': 'invalid', 'type': 'value_error'}] + resp = _Resp( + 422, + body={ + 'error': 'request validation failed', + 'category': 'user', + 'error_code': 422, + 'request_id': 'req-details', + 'details': details, + }, + ) + with pytest.raises(TwinkleHTTPError) as exc: + _handle_response(resp) + assert exc.value.details == details + assert exc.value.traceback is None + + +def test_server_traceback_is_preserved(): + traceback_text = 'Traceback (most recent call last):\n File "/srv/app.py", line 1\nRuntimeError: boom' + resp = _Resp( + 500, + body={ + 'error': 'RuntimeError: boom', + 'category': 'server', + 'error_code': 500, + 'request_id': 'req-trace', + 'traceback': traceback_text, + }, + ) + with pytest.raises(TwinkleHTTPError) as exc: + _handle_response(resp) + assert exc.value.traceback == traceback_text + assert exc.value.details is None + + +def test_410_raises_stop_iteration_not_http_error(): + """410 keeps raising StopIteration, not an HTTP error.""" + resp = _Resp(410, body={'detail': 'exhausted'}) + with pytest.raises(StopIteration): + _handle_response(resp) + + +def test_ok_response_passes_through(): + resp = _Resp(200, body={'status': 'ok'}) + assert _handle_response(resp) is resp diff --git a/tests/twinkle_client/test_future_layer.py b/tests/twinkle_client/test_future_layer.py new file mode 100644 index 000000000..e078b5388 --- /dev/null +++ b/tests/twinkle_client/test_future_layer.py @@ -0,0 +1,170 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Client future-layer unit tests. + +``resolve`` is exercised against fabricated envelopes and a monkeypatched +``_post_retrieve``; no server or network is involved. +""" +from __future__ import annotations + +import pytest +import requests + +from twinkle_client import _future +from twinkle_client.exceptions import TaskFailedError, TaskRecordLostError, TaskWaitTimeoutError +from twinkle.protocol.types.errors import ErrorPayload +from twinkle.protocol.types.lifecycle import TaskEnvelope + + +class _Model: + """A model_cls that records what it deserialized.""" + + def __init__(self, result): + self.result = result + + @classmethod + def model_validate(cls, value): + return _Model(value) + + +def _completed(result): + return TaskEnvelope(request_id='r', status='completed', result=result) + + +def _failed(**kw): + payload = ErrorPayload(error='boom', category='server', error_code=500, request_id='r', **kw) + return TaskEnvelope(request_id='r', status='failed', error=payload) + + +def _running(): + return TaskEnvelope(request_id='r', status='running', queue_state='active') + + +def test_terminal_submit_issues_no_retrieve(monkeypatch): + """A task terminal in the submit envelope makes zero retrieve calls.""" + + def _boom(_request_id, _transport): + raise AssertionError('retrieve must not be called for a terminal submit') + + monkeypatch.setattr(_future, '_post_retrieve', _boom) + out = _future.resolve(_completed({'loss': 1.0}), model_cls=_Model) + assert out.result == {'loss': 1.0} + + +def test_terminal_submit_failure_raises_taskfailed_with_payload(monkeypatch): + """A failure in the submit envelope raises TaskFailedError, payload intact.""" + monkeypatch.setattr(_future, '_post_retrieve', lambda _r, _transport: pytest.fail('no retrieve')) + with pytest.raises(TaskFailedError) as exc: + _future.resolve(_failed(), model_cls=_Model) + assert exc.value.error == 'boom' + assert exc.value.category == 'server' + assert exc.value.request_id == 'r' + assert exc.value.error_code == 500 + assert not isinstance(exc.value, requests.HTTPError) + + +def test_model_cls_none_returns_none_result(monkeypatch): + """A method that returned None before still returns None (not swallowed).""" + monkeypatch.setattr(_future, '_post_retrieve', lambda _r, _transport: pytest.fail('no retrieve')) + assert _future.resolve(_completed(None), model_cls=None) is None + + +def test_non_terminal_submit_polls_until_terminal(monkeypatch): + replies = [_running(), _running(), _completed({'ok': 1})] + monkeypatch.setattr(_future, '_post_retrieve', lambda _r, _transport: replies.pop(0)) + out = _future.resolve(_running(), model_cls=_Model) + assert out.result == {'ok': 1} + assert replies == [] + + +def test_404_is_bounded_then_raises_record_lost(monkeypatch): + + def _always_404(_request_id, _transport): + e = requests.HTTPError('404') + e.status_code = 404 + raise e + + monkeypatch.setattr(_future, '_post_retrieve', _always_404) + with pytest.raises(TaskRecordLostError): + _future.resolve(_running(), model_cls=_Model) + + +def test_transport_5xx_is_bounded_then_reraises(monkeypatch): + monkeypatch.setattr(_future.time, 'sleep', lambda _s: None) # no real backoff sleeps + + def _always_503(_request_id, _transport): + e = requests.HTTPError('503') + e.status_code = 503 + raise e + + monkeypatch.setattr(_future, '_post_retrieve', _always_503) + with pytest.raises(requests.HTTPError): + _future.resolve(_running(), model_cls=_Model) + + +def test_connection_error_is_retried_then_succeeds(monkeypatch): + monkeypatch.setattr(_future.time, 'sleep', lambda _s: None) + replies = [requests.ConnectionError('reset'), _completed({'ok': True})] + + def _next(_request_id, _transport): + value = replies.pop(0) + if isinstance(value, BaseException): + raise value + return value + + monkeypatch.setattr(_future, '_post_retrieve', _next) + out = _future.resolve(_running(), model_cls=_Model) + assert out.result == {'ok': True} + + +def test_connection_error_is_bounded_then_reraised(monkeypatch): + monkeypatch.setattr(_future.time, 'sleep', lambda _s: None) + calls = 0 + + def _always_fails(_request_id, _transport): + nonlocal calls + calls += 1 + raise requests.ConnectionError('reset') + + monkeypatch.setattr(_future, '_post_retrieve', _always_fails) + with pytest.raises(requests.ConnectionError, match='reset'): + _future.resolve(_running(), model_cls=_Model) + assert calls == _future._TRANSPORT_RETRY_MAX + 1 + + +def test_non_retryable_4xx_reraises_immediately(monkeypatch): + + def _400(_request_id, _transport): + e = requests.HTTPError('400') + e.status_code = 400 + raise e + + monkeypatch.setattr(_future, '_post_retrieve', _400) + with pytest.raises(requests.HTTPError): + _future.resolve(_running(), model_cls=_Model) + + +def test_total_timeout_raises_wait_timeout(monkeypatch): + monkeypatch.setattr(_future, '_post_retrieve', lambda _r, _transport: _running()) + with pytest.raises(TaskWaitTimeoutError) as exc: + _future.resolve(_running(), model_cls=_Model, total_timeout=0.0) + assert exc.value.request_id == 'r' + + +def test_success_resets_both_retry_counters(monkeypatch): + """A successful reply zeroes both counters, so intermittent 404s never sum up.""" + seq = [] + + def _mixed(_request_id, _transport): + seq.append(1) + n = len(seq) + if n in (1, 2, 4, 5): # 404s interleaved with a success at n==3 + e = requests.HTTPError('404') + e.status_code = 404 + raise e + if n == 3: + return _running() # success resets not_found_count + return _completed({'done': True}) + + monkeypatch.setattr(_future, '_post_retrieve', _mixed) + out = _future.resolve(_running(), model_cls=_Model) + assert out.result == {'done': True} diff --git a/tests/twinkle_client/test_import_surface.py b/tests/twinkle_client/test_import_surface.py new file mode 100644 index 000000000..fcce13399 --- /dev/null +++ b/tests/twinkle_client/test_import_surface.py @@ -0,0 +1,59 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +from __future__ import annotations + +import json +import subprocess +import sys + + +def _run_import_probe(source: str) -> dict[str, object]: + completed = subprocess.run( + [sys.executable, '-c', source], + check=True, + capture_output=True, + text=True, + ) + return json.loads(completed.stdout) + + +def test_import_twinkle_does_not_eagerly_import_client_or_torch() -> None: + result = _run_import_probe( + "import json, sys, twinkle; " + "print(json.dumps({'client': 'twinkle_client' in sys.modules, 'torch': 'torch' in sys.modules}))") + assert result == {'client': False, 'torch': False} + + +def test_twinkle_client_entry_points_remain_available_without_eager_import() -> None: + result = _run_import_probe( + "import json, sys, twinkle; " + "before = 'twinkle_client' in sys.modules; " + "from twinkle import init_tinker_client, init_twinkle_client; " + "after = 'twinkle_client' in sys.modules; " + "print(json.dumps({'before': before, 'after': after, " + "'tinker': callable(init_tinker_client), 'twinkle': callable(init_twinkle_client)}))") + assert result == {'before': False, 'after': False, 'tinker': True, 'twinkle': True} + + +def test_twinkle_client_entry_points_delegate(monkeypatch) -> None: + import twinkle + import twinkle_client + + calls = [] + monkeypatch.setattr(twinkle_client, 'init_tinker_client', lambda **kwargs: calls.append(('tinker', kwargs))) + monkeypatch.setattr(twinkle_client, 'init_twinkle_client', lambda **kwargs: ('twinkle', kwargs)) + + assert twinkle.init_tinker_client(feature=True) is None + assert calls == [('tinker', {'feature': True})] + assert twinkle.init_twinkle_client(base_url='http://server', api_key='key') == ( + 'twinkle', { + 'base_url': 'http://server', + 'api_key': 'key', + 'session_heartbeat_interval': 10, + }) + + +def test_import_twinkle_client_does_not_load_heavy_data_dependencies() -> None: + result = _run_import_probe( + "import json, sys, twinkle_client; " + "print(json.dumps({name: name in sys.modules for name in ('torch', 'datasets', 'pandas')}))") + assert result == {'torch': False, 'datasets': False, 'pandas': False} diff --git a/tests/twinkle_client/test_remote_components.py b/tests/twinkle_client/test_remote_components.py new file mode 100644 index 000000000..85657d272 --- /dev/null +++ b/tests/twinkle_client/test_remote_components.py @@ -0,0 +1,75 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +from __future__ import annotations + +import json +import subprocess +import sys + +from torch.utils.data import IterableDataset as TorchIterableDataset + +from twinkle_client.common import remote_component +from twinkle_client.dataloader import DataLoader +from twinkle_client.dataset import Dataset, IterableDataset, IterablePackingDataset, LazyDataset, PackingDataset +from twinkle_client.http import ClientContext, ClientTransport + + +class _Session: + + def close(self) -> None: + pass + + +class _CallableDataset: + + def __call__(self): + return None + + +def _transport() -> ClientTransport: + return ClientTransport(ClientContext(base_url='http://server', api_key='key'), session=_Session()) + + +def test_remote_component_binding_and_dispatch_are_shared(monkeypatch) -> None: + created: list[tuple[str, str, ClientTransport, dict]] = [] + called: list[tuple[str, str, tuple, ClientTransport, dict]] = [] + + def _create(processor_type, class_type, *, transport, **kwargs): + created.append((processor_type, class_type, transport, kwargs)) + return f'pid:{class_type}' + + def _call(processor_id, function, *args, transport, **kwargs): + called.append((processor_id, function, args, transport, kwargs)) + return 'result' + + monkeypatch.setattr(remote_component, 'create_remote_component', _create) + monkeypatch.setattr(remote_component, 'call_remote_component', _call) + transport = _transport() + + dataset = Dataset(transport=transport) + loader = DataLoader(_CallableDataset(), transport=transport) + + assert dataset.check(flag=True) == 'result' + assert loader.get_state() == 'result' + assert created[0][:3] == ('dataset', 'Dataset', transport) + assert created[1][:3] == ('dataloader', 'DataLoader', transport) + assert called[0] == ('pid:Dataset', 'check', (), transport, {'flag': True}) + assert called[1] == ('pid:DataLoader', 'get_state', (), transport, {}) + + +def test_dataset_method_surfaces_and_iterable_mro() -> None: + assert 'map' not in LazyDataset.__dict__ + assert hasattr(LazyDataset, 'map') + assert hasattr(PackingDataset, 'map') + assert '__len__' not in IterableDataset.__dict__ + assert '__getitem__' not in IterableDataset.__dict__ + assert IterableDataset.__mro__[1] is TorchIterableDataset + assert IterablePackingDataset.__mro__[1] is TorchIterableDataset + + +def test_dataloader_import_does_not_load_transformers() -> None: + source = ( + "import json, sys, twinkle_client.dataloader; " + "print(json.dumps({'transformers': 'transformers' in sys.modules}))" + ) + result = subprocess.run([sys.executable, '-c', source], check=True, capture_output=True, text=True) + assert json.loads(result.stdout) == {'transformers': False} diff --git a/tests/twinkle_client/test_request_builder.py b/tests/twinkle_client/test_request_builder.py new file mode 100644 index 000000000..0c7a15bce --- /dev/null +++ b/tests/twinkle_client/test_request_builder.py @@ -0,0 +1,175 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Level 0: the client builds the request from the schema, in-process. + +The property worth testing is that a bad call fails *without a network round trip*. Every +test here monkeypatches ``requests.post`` to fail loudly, so any test that passes has +proved no request was sent. +""" +from __future__ import annotations + +import pytest +from pydantic import JsonValue, ValidationError +from typing import Dict + +from twinkle_client._request_builder import build_request, request_json, to_wire_value +from twinkle.protocol.serialize import serialize_object +from twinkle_client.exceptions import TwinkleClientValidationError +from twinkle.protocol.types import model as model_types +from twinkle.protocol.types.base import StrictRequest, passthrough + + +@pytest.fixture(autouse=True) +def no_network(monkeypatch): + """Any HTTP call in this module is a bug in the code under test.""" + import twinkle_client.http.client as http_client + monkeypatch.setattr(http_client.requests, 'post', + lambda *a, **k: pytest.fail('a Level 0 failure must not produce a request')) + + +# --------------------------------------------------------------------------- # +# Routing +# --------------------------------------------------------------------------- # + + +def test_a_declared_name_goes_to_its_field(): + body = build_request(model_types.ForwardRequest, inputs=[{'input_ids': [1]}], adapter_name='a', task='embedding') + assert body.task == 'embedding' + assert body.loss_kwargs == {} + + +def test_an_undeclared_name_goes_to_the_single_passthrough_region(): + """Public signatures stay ``**kwargs``; only the wire shape changes.""" + body = build_request( + model_types.ForwardBackwardTaskRequest, + inputs=[{ + 'input_ids': [1] + }], + adapter_name='a', + advantages=[0.5], + old_logps=[[-1.0]]) + assert body.loss_kwargs == {'advantages': [0.5], 'old_logps': [[-1.0]]} + + +def test_an_ambiguous_target_is_an_error_rather_than_a_guess(): + """With two regions the builder refuses to choose instead of guessing. + + A model that grows a second region -- constructor arguments *and* invoked-method + arguments, say -- has no name-based rule that can tell them apart, and guessing wrong + sends a valid argument to the wrong callable: a wrong result rather than an error. No + shipped model has two today; the rule exists so that adding one fails loudly. + """ + + class _TwoRegions(StrictRequest): + target: str + init_kwargs: Dict[str, JsonValue] = passthrough() + call_kwargs: Dict[str, JsonValue] = passthrough() + + with pytest.raises(TwinkleClientValidationError) as raised: + build_request(_TwoRegions, target='t', unplaceable=1) + assert 'init_kwargs' in str(raised.value) and 'call_kwargs' in str(raised.value) + + +def test_a_model_without_a_region_rejects_an_unknown_name(): + with pytest.raises(TwinkleClientValidationError) as raised: + build_request(model_types.SaveRequest, adapter_name='a', name='ckpt', typo=1) + assert 'typo' in str(raised.value) + + +def test_an_explicit_region_and_a_routed_key_are_merged(): + body = build_request( + model_types.SetLossRequest, + loss_cls='DPOLoss', + adapter_name='a', + init_kwargs={'beta': 0.1}, + loss_type='sigmoid') + assert body.init_kwargs == {'beta': 0.1, 'loss_type': 'sigmoid'} + + +def test_a_key_passed_twice_is_an_error(): + """Silently letting one win would make the effective value depend on merge order.""" + with pytest.raises(TwinkleClientValidationError, match='both directly and inside'): + build_request( + model_types.SetLossRequest, loss_cls='DPOLoss', adapter_name='a', init_kwargs={'beta': 0.1}, beta=0.2) + + +def test_an_omitted_optional_argument_is_not_routed_into_the_region(): + """Client methods pass optionals unconditionally. + + Routing an explicit ``None`` would hand the backend a null argument it never received + before, changing behaviour for callers who simply did not pass anything. + """ + body = build_request(model_types.SetLossRequest, loss_cls='DPOLoss', adapter_name='a', unused=None) + assert body.init_kwargs == {} + + +# --------------------------------------------------------------------------- # +# Validation before the wire +# --------------------------------------------------------------------------- # + + +def test_a_wrongly_typed_field_fails_in_process(): + with pytest.raises(ValidationError): + build_request(model_types.ForwardRequest, inputs=[{'input_ids': [1]}], adapter_name='a', temperature='hot') + + +def test_an_out_of_range_value_fails_in_process(): + with pytest.raises(ValidationError): + build_request(model_types.ForwardRequest, inputs=[{'input_ids': [1]}], adapter_name='a', temperature=0) + + +def test_malformed_inputs_fail_in_process(): + """The client shares the server's schema, so it catches this without asking.""" + with pytest.raises(ValidationError): + build_request(model_types.ForwardRequest, inputs=[{'input_ids': [True]}], adapter_name='a') + + +def test_a_missing_required_field_fails_in_process(): + with pytest.raises(ValidationError): + build_request(model_types.ForwardRequest, inputs=[{'input_ids': [1]}]) + + +def test_forward_only_requires_an_adapter_context(): + with pytest.raises(ValidationError): + build_request(model_types.ForwardOnlyRequest, inputs=[{'input_ids': [1]}]) + + +# --------------------------------------------------------------------------- # +# Serialization +# --------------------------------------------------------------------------- # + + +def test_unset_optionals_stay_off_the_wire(): + """Absent and "not requested" must look the same, or the server needs its own defaults.""" + import json + body = build_request(model_types.ForwardRequest, inputs=[{'input_ids': [1]}], adapter_name='a') + payload = json.loads(request_json(body)) + assert payload == {'inputs': [{'input_ids': [1]}], 'adapter_name': 'a', 'loss_kwargs': {}} + + +def test_a_lora_config_is_serialized_to_the_form_the_server_decodes(): + from peft import LoraConfig + wire = to_wire_value(LoraConfig(target_modules='all-linear')) + assert isinstance(wire, str) and 'LoraConfig' in wire + + +def test_binary_values_are_rejected_by_legacy_serializer(): + for value in (b'raw', bytearray(b'raw'), memoryview(b'raw')): + with pytest.raises(TypeError, match='Unsupported binary object'): + serialize_object(value) + + +def test_a_component_handle_is_sent_as_its_id(): + + class _Handle: + processor_id = 'pid:abc' + + assert to_wire_value(_Handle()) == 'pid:abc' + + +def test_numpy_values_are_converted_to_lists(): + import numpy as np + assert to_wire_value(np.array([1, 2])) == [1, 2] + + +if __name__ == '__main__': + raise SystemExit(pytest.main([__file__, '-v'])) diff --git a/tests/twinkle_client/test_transport.py b/tests/twinkle_client/test_transport.py new file mode 100644 index 000000000..917c869b2 --- /dev/null +++ b/tests/twinkle_client/test_transport.py @@ -0,0 +1,202 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +from __future__ import annotations + +import logging +import pytest + +from twinkle_client.http import ClientContext, ClientTransport +from twinkle_client.http.context import capture_transport, clear_default_transport, set_default_transport +from twinkle_client.manager import TwinkleClient + + +class _Response: + ok = True + status_code = 200 + url = 'http://server' + text = '' + + def json(self): + return {} + + +class _Session: + + def __init__(self): + self.calls = [] + self.closed = False + + def post(self, url, **kwargs): + self.calls.append(('post', url, kwargs)) + return _Response() + + def get(self, url, **kwargs): + self.calls.append(('get', url, kwargs)) + return _Response() + + def delete(self, url, **kwargs): + self.calls.append(('delete', url, kwargs)) + return _Response() + + def close(self): + self.closed = True + + +def _transport(name: str) -> ClientTransport: + return ClientTransport( + ClientContext( + base_url=f'http://{name}', + api_key=f'{name}-key', + session_id=f'{name}-session', + routing_id=f'{name}-routing', + ), + session=_Session(), + ) + + +def test_context_normalizes_url_and_transport_builds_stable_headers(): + transport = _transport('alpha') + transport.post('/resource') + _, url, kwargs = transport._session.calls[-1] + + assert url == 'http://alpha/api/v1/resource' + assert kwargs['headers']['Authorization'] == 'Bearer alpha-key' + assert kwargs['headers']['X-Twinkle-Session-Id'] == 'alpha-session' + assert kwargs['headers']['x-request-id'] == 'alpha-routing' + + +def test_wrapper_captures_default_once_and_factory_is_explicit(): + from twinkle_client.model import MultiLoraTransformersModel + + transport_a = _transport('alpha') + transport_b = _transport('beta') + set_default_transport(transport_a) + legacy_model = MultiLoraTransformersModel('model') + set_default_transport(transport_b) + + client_a = TwinkleClient(transport=transport_a) + factory_model = client_a.model('factory-model') + + assert legacy_model._transport is transport_a + assert factory_model._transport is transport_a + assert legacy_model.server_url.startswith('http://alpha/api/v1/') + assert factory_model.server_url.startswith('http://alpha/api/v1/') + + +def test_close_is_idempotent_and_does_not_clear_another_default(): + transport_a = _transport('alpha') + transport_b = _transport('beta') + client_a = TwinkleClient(transport=transport_a) + set_default_transport(transport_b) + + client_a.close() + client_a.close() + + assert transport_a.closed + assert transport_a._session.closed + assert not transport_b.closed + + +def test_http_public_api_has_no_legacy_context_getters_or_setters(): + import twinkle_client.http as http + + assert not ({ + 'get_base_url', + 'get_api_key', + 'get_session_id', + 'get_request_id', + 'set_base_url', + 'set_api_key', + 'set_session_id', + 'set_request_id', + } & set(http.__all__)) + + +def test_default_transport_fallback_logs_info_once(monkeypatch, caplog): + import twinkle_client.http.context as context + + monkeypatch.setattr(context, '_default_transport', None) + with caplog.at_level(logging.INFO, logger='twinkle_client'): + transport = capture_transport() + assert capture_transport() is transport + records = [record for record in caplog.records if 'No explicit Twinkle client configured' in record.message] + assert len(records) == 1 + assert transport.context.base_url in records[0].message + transport.close() + clear_default_transport(transport) + + +def test_publish_blocks_rebind_and_replacement_does_not_close_old(caplog): + first = _transport('first') + first.bind_context(ClientContext(base_url='http://bound', api_key='key')) + set_default_transport(first) + + with pytest.raises(RuntimeError, match='published'): + first.bind_context(ClientContext(base_url='http://too-late', api_key='key')) + + second = _transport('second') + with caplog.at_level(logging.WARNING, logger='twinkle_client'): + set_default_transport(second) + assert any('Replacing default Twinkle transport' in record.message for record in caplog.records) + assert not first.closed + clear_default_transport(second) + first.close() + second.close() + + +def test_transport_adapter_configuration_and_post_retry_boundary(): + transport = ClientTransport(ClientContext(base_url='http://pool', api_key='key')) + adapter = transport._session.get_adapter('http://') + assert adapter._pool_maxsize == 32 + assert adapter._pool_block is True + assert adapter.max_retries.allowed_methods == frozenset({'GET', 'DELETE'}) + assert 'POST' not in adapter.max_retries.allowed_methods + transport.close() + + +def test_client_closes_distinct_heartbeat_transport(): + main = _transport('main') + heartbeat = _transport('heartbeat') + client = TwinkleClient(transport=main, heartbeat_transport=heartbeat) + client.close() + assert main.closed and heartbeat.closed + + +class _CapabilitiesResponse(_Response): + + def json(self): + return {'supported_models': [{'model_name': 'm'}]} + + +class _CapabilitiesSession(_Session): + + def get(self, url, **kwargs): + self.calls.append(('get', url, kwargs)) + return _CapabilitiesResponse() + + +def _capability_client(name: str) -> TwinkleClient: + transport = ClientTransport( + ClientContext(base_url=f'http://{name}', api_key=f'{name}-key', session_id=f'{name}-session'), + session=_CapabilitiesSession(), + ) + return TwinkleClient(transport=transport) + + +def test_capability_cache_is_per_transport_and_not_process_global(): + from twinkle.protocol.types.server import GetServerCapabilitiesResponse + + client_a = _capability_client('alpha') + client_b = _capability_client('beta') + + first = client_a.get_server_capabilities() + again = client_a.get_server_capabilities() + + # Cached on the transport instance: the second query reuses it, no second GET. + assert isinstance(first, GetServerCapabilitiesResponse) + assert again is first + assert sum(1 for call in client_a.transport._session.calls if call[0] == 'get') == 1 + + # Isolation: caching on A must never populate another transport's cache. + assert client_b.transport.cached_capabilities is None + client_b.get_server_capabilities() + assert client_b.transport.cached_capabilities is not client_a.transport.cached_capabilities diff --git a/tests/twinkle_client/test_types_contract.py b/tests/twinkle_client/test_types_contract.py new file mode 100644 index 000000000..7dbaa25d1 --- /dev/null +++ b/tests/twinkle_client/test_types_contract.py @@ -0,0 +1,152 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Contract-base consistency and naming-disambiguation tests. + +- ``QueueStateLiteral`` value set equals the server ``QueueState`` enum. +- naming disambiguation guard. + +The two SDKs already share public names. The contract freezes that legacy set and +rejects new collisions while requiring explicit aliases when both SDKs are imported +in one module. +""" +from __future__ import annotations + +import ast +import pathlib +import typing + +import twinkle +from twinkle.server.task_queue.types import QueueState +from twinkle.protocol.types import model as model_types +from twinkle.protocol.types.errors import QueueStateLiteral +from twinkle.protocol.types.server import GetServerCapabilitiesResponse + +_TWINKLE_SRC = pathlib.Path(twinkle.__file__).resolve().parent +_LEGACY_PUBLIC_NAME_OVERLAP = frozenset({ + 'Checkpoint', + 'CheckpointsListResponse', + 'CreateModelRequest', + 'CreateSessionRequest', + 'CreateSessionResponse', + 'Cursor', + 'ForwardRequest', + 'GetServerCapabilitiesResponse', + 'HealthResponse', + 'LoraConfig', + 'SampleRequest', + 'SessionHeartbeatRequest', + 'SessionHeartbeatResponse', + 'SupportedModel', + 'TrainingRun', + 'TrainingRunsResponse', + 'WeightsInfoResponse', + 'checkpoint', +}) + + +def test_queue_state_literal_matches_server_enum(): + literal_values = set(typing.get_args(QueueStateLiteral)) + enum_values = {state.value for state in QueueState} + assert literal_values == enum_values, (f'QueueStateLiteral {literal_values} != QueueState {enum_values}') + + +def test_old_capabilities_response_gets_conservative_defaults(): + response = GetServerCapabilitiesResponse.model_validate({'supported_models': []}) + assert response.protocol_version == 1 + assert response.features.task_envelope is True + assert response.features.cancel is False + assert response.features.batch_retrieve is False + + +def test_capabilities_response_ignores_future_fields(): + response = GetServerCapabilitiesResponse.model_validate({ + 'supported_models': [], + 'future_top_level': True, + 'features': { + 'cancel': True, + 'future_feature': True + }, + 'limits': { + 'max_batch_size': 8, + 'future_limit': 9 + }, + }) + assert response.features.cancel is True + assert response.limits.max_batch_size == 8 + + +def test_data_plane_forward_only_has_no_seq_id() -> None: + fields = model_types.DataPlaneForwardOnlyRequest.model_fields + assert 'seq_id' not in fields + assert 'seq_id' in model_types.DataPlaneForwardRequest.model_fields + + +def test_void_response_names_are_canonical_ok_response_aliases() -> None: + for name in ( + 'BackwardResponse', + 'StepResponse', + 'ZeroGradResponse', + 'LrStepResponse', + 'SetLossResponse', + 'SetOptimizerResponse', + 'SetLrSchedulerResponse', + 'LoadResponse', + 'SetTemplateResponse', + 'SetProcessorResponse', + 'ClipGradAndStepResponse', + 'ApplyPatchResponse', + 'AddMetricResponse', + ): + assert getattr(model_types, name) is model_types.OkResponse + + +def _origin(module: str | None) -> str | None: + """Classify an import's source module as 'tinker', 'twinkle_client', or None.""" + if not module: + return None + if module == 'tinker' or module.startswith('tinker.'): + return 'tinker' + if module == 'twinkle_client' or module.startswith('twinkle_client.'): + return 'twinkle_client' + return None + + +def _binding_collisions(tree: ast.AST) -> set[str]: + """Return local names bound to BOTH a tinker and a twinkle_client import.""" + tinker_names: set[str] = set() + twinkle_names: set[str] = set() + for node in ast.walk(tree): + if isinstance(node, ast.ImportFrom): + origin = _origin(node.module) + if origin is None: + continue + for alias in node.names: + bound = alias.asname or alias.name + (tinker_names if origin == 'tinker' else twinkle_names).add(bound) + elif isinstance(node, ast.Import): + for alias in node.names: + origin = _origin(alias.name) + if origin is None: + continue + bound = alias.asname or alias.name.split('.')[0] + (tinker_names if origin == 'tinker' else twinkle_names).add(bound) + return tinker_names & twinkle_names + + +def test_public_name_overlap_does_not_grow(): + import tinker.types + + import twinkle.protocol.types + + overlap = {name for name in set(dir(tinker.types)) & set(dir(twinkle.protocol.types)) if not name.startswith('_')} + assert overlap == _LEGACY_PUBLIC_NAME_OVERLAP + + +def test_no_tinker_twinkle_same_name_binding(): + offenders: dict[str, set[str]] = {} + for path in _TWINKLE_SRC.rglob('*.py'): + tree = ast.parse(path.read_text(), filename=str(path)) + collisions = _binding_collisions(tree) + if collisions: + offenders[str(path.relative_to(_TWINKLE_SRC))] = collisions + assert not offenders, ('tinker and twinkle_client types bound to the same local name (alias tinker ' + f'to disambiguate): {offenders}') diff --git a/tests/utils/test_nccl_safe.py b/tests/utils/test_nccl_safe.py new file mode 100644 index 000000000..62b0d8e15 --- /dev/null +++ b/tests/utils/test_nccl_safe.py @@ -0,0 +1,18 @@ +from unittest.mock import patch + +import pytest + +from twinkle.utils.nccl_safe import nccl_safe_megatron + + +def test_nccl_failure_preserves_type_and_adds_rank_context(): + + @nccl_safe_megatron + def fail(_self): + raise ValueError('bad shape') + + with patch('twinkle.utils.nccl_safe._global_rank', return_value=3): + with pytest.raises(ValueError) as caught: + fail(object()) + + assert 'global_rank=3' in ''.join(getattr(caught.value, '__notes__', caught.value.args))