Introduce naming convention in io_struct and base sglang io classes. (#10133)
This commit is contained in:
@@ -18,6 +18,7 @@ processes (TokenizerManager, DetokenizerManager, Scheduler).
|
||||
|
||||
import copy
|
||||
import uuid
|
||||
from abc import ABC
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
@@ -36,10 +37,32 @@ else:
|
||||
|
||||
|
||||
# Parameters for a session
|
||||
@dataclass
|
||||
class BaseReq(ABC):
|
||||
rid: Optional[Union[str, List[str]]] = field(default=None, kw_only=True)
|
||||
|
||||
def regenerate_rid(self):
|
||||
"""Generate a new request ID and return it."""
|
||||
if isinstance(self.rid, list):
|
||||
self.rid = [uuid.uuid4().hex for _ in range(len(self.rid))]
|
||||
else:
|
||||
self.rid = uuid.uuid4().hex
|
||||
return self.rid
|
||||
|
||||
|
||||
@dataclass
|
||||
class BaseBatchReq(ABC):
|
||||
rids: Optional[List[str]] = field(default=None, kw_only=True)
|
||||
|
||||
def regenerate_rids(self):
|
||||
"""Generate new request IDs and return them."""
|
||||
self.rids = [uuid.uuid4().hex for _ in range(len(self.rids))]
|
||||
return self.rids
|
||||
|
||||
|
||||
@dataclass
|
||||
class SessionParams:
|
||||
id: Optional[str] = None
|
||||
rid: Optional[str] = None
|
||||
offset: Optional[int] = None
|
||||
replace: Optional[bool] = None
|
||||
drop_previous_output: Optional[bool] = None
|
||||
@@ -63,7 +86,7 @@ MultimodalDataInputFormat = Union[
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenerateReqInput:
|
||||
class GenerateReqInput(BaseReq):
|
||||
# The input prompt. It can be a single prompt or a batch of prompts.
|
||||
text: Optional[Union[List[str], str]] = None
|
||||
# The token ids for text; one can specify either text or input_ids
|
||||
@@ -83,8 +106,6 @@ class GenerateReqInput:
|
||||
audio_data: Optional[MultimodalDataInputFormat] = None
|
||||
# The sampling_params. See descriptions below.
|
||||
sampling_params: Optional[Union[List[Dict], Dict]] = None
|
||||
# The request id.
|
||||
rid: Optional[Union[List[str], str]] = None
|
||||
# Whether to return logprobs.
|
||||
return_logprob: Optional[Union[List[bool], bool]] = None
|
||||
# If return logprobs, the start location in the prompt for returning logprobs.
|
||||
@@ -491,11 +512,6 @@ class GenerateReqInput:
|
||||
):
|
||||
raise ValueError("Session params must be a dict or a list of dicts.")
|
||||
|
||||
def regenerate_rid(self):
|
||||
"""Generate a new request ID and return it."""
|
||||
self.rid = uuid.uuid4().hex
|
||||
return self.rid
|
||||
|
||||
def __getitem__(self, i):
|
||||
return GenerateReqInput(
|
||||
text=self.text[i] if self.text is not None else None,
|
||||
@@ -558,9 +574,7 @@ class GenerateReqInput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class TokenizedGenerateReqInput:
|
||||
# The request id
|
||||
rid: str
|
||||
class TokenizedGenerateReqInput(BaseReq):
|
||||
# The input text
|
||||
input_text: str
|
||||
# The input token ids
|
||||
@@ -625,7 +639,7 @@ class TokenizedGenerateReqInput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class BatchTokenizedGenerateReqInput:
|
||||
class BatchTokenizedGenerateReqInput(BaseBatchReq):
|
||||
# The batch of tokenized requests
|
||||
batch: List[TokenizedGenerateReqInput]
|
||||
|
||||
@@ -640,7 +654,7 @@ class BatchTokenizedGenerateReqInput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class EmbeddingReqInput:
|
||||
class EmbeddingReqInput(BaseReq):
|
||||
# The input prompt. It can be a single prompt or a batch of prompts.
|
||||
text: Optional[Union[List[List[str]], List[str], str]] = None
|
||||
# The image input. It can be an image instance, file name, URL, or base64 encoded string.
|
||||
@@ -656,8 +670,6 @@ class EmbeddingReqInput:
|
||||
audio_data: Optional[MultimodalDataInputFormat] = None
|
||||
# The token ids for text; one can either specify text or input_ids.
|
||||
input_ids: Optional[Union[List[List[int]], List[int]]] = None
|
||||
# The request id.
|
||||
rid: Optional[Union[List[str], str]] = None
|
||||
# Dummy sampling params for compatibility
|
||||
sampling_params: Optional[Union[List[Dict], Dict]] = None
|
||||
# Dummy input embeds for compatibility
|
||||
@@ -728,10 +740,6 @@ class EmbeddingReqInput:
|
||||
for i in range(self.batch_size):
|
||||
self.sampling_params[i]["max_new_tokens"] = 0
|
||||
|
||||
def regenerate_rid(self):
|
||||
self.rid = uuid.uuid4().hex
|
||||
return self.rid
|
||||
|
||||
def contains_mm_input(self) -> bool:
|
||||
return (
|
||||
has_valid_data(self.image_data)
|
||||
@@ -760,9 +768,7 @@ class EmbeddingReqInput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class TokenizedEmbeddingReqInput:
|
||||
# The request id
|
||||
rid: str
|
||||
class TokenizedEmbeddingReqInput(BaseReq):
|
||||
# The input text
|
||||
input_text: str
|
||||
# The input token ids
|
||||
@@ -780,7 +786,7 @@ class TokenizedEmbeddingReqInput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class BatchTokenizedEmbeddingReqInput:
|
||||
class BatchTokenizedEmbeddingReqInput(BaseBatchReq):
|
||||
# The batch of tokenized embedding requests
|
||||
batch: List[TokenizedEmbeddingReqInput]
|
||||
|
||||
@@ -795,9 +801,7 @@ class BatchTokenizedEmbeddingReqInput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class BatchTokenIDOut:
|
||||
# The request id
|
||||
rids: List[str]
|
||||
class BatchTokenIDOutput(BaseBatchReq):
|
||||
# The finish reason
|
||||
finished_reasons: List[BaseFinishReason]
|
||||
# For incremental decoding
|
||||
@@ -842,7 +846,7 @@ class BatchTokenIDOut:
|
||||
|
||||
|
||||
@dataclass
|
||||
class BatchMultimodalDecodeReq:
|
||||
class BatchMultimodalDecodeReq(BaseBatchReq):
|
||||
decoded_ids: List[int]
|
||||
input_token_logprobs_val: List[float]
|
||||
input_token_logprobs_idx: List[int]
|
||||
@@ -854,8 +858,6 @@ class BatchMultimodalDecodeReq:
|
||||
image_resolutions: List[List[int]]
|
||||
resize_image_resolutions: List[List[int]]
|
||||
|
||||
# The request id
|
||||
rids: List[str]
|
||||
finished_reasons: List[BaseFinishReason]
|
||||
|
||||
# Token counts
|
||||
@@ -871,9 +873,7 @@ class BatchMultimodalDecodeReq:
|
||||
|
||||
|
||||
@dataclass
|
||||
class BatchStrOut:
|
||||
# The request id
|
||||
rids: List[str]
|
||||
class BatchStrOutput(BaseBatchReq):
|
||||
# The finish reason
|
||||
finished_reasons: List[dict]
|
||||
# The output decoded strings
|
||||
@@ -909,9 +909,7 @@ class BatchStrOut:
|
||||
|
||||
|
||||
@dataclass
|
||||
class BatchMultimodalOut:
|
||||
# The request id
|
||||
rids: List[str]
|
||||
class BatchMultimodalOutput(BaseBatchReq):
|
||||
# The finish reason
|
||||
finished_reasons: List[dict]
|
||||
decoded_ids: List[List[int]]
|
||||
@@ -936,9 +934,7 @@ class BatchMultimodalOut:
|
||||
|
||||
|
||||
@dataclass
|
||||
class BatchEmbeddingOut:
|
||||
# The request id
|
||||
rids: List[str]
|
||||
class BatchEmbeddingOutput(BaseBatchReq):
|
||||
# The finish reason
|
||||
finished_reasons: List[BaseFinishReason]
|
||||
# The output embedding
|
||||
@@ -952,27 +948,27 @@ class BatchEmbeddingOut:
|
||||
|
||||
|
||||
@dataclass
|
||||
class ClearHiCacheReqInput:
|
||||
class ClearHiCacheReqInput(BaseReq):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class ClearHiCacheReqOutput:
|
||||
class ClearHiCacheReqOutput(BaseReq):
|
||||
success: bool
|
||||
|
||||
|
||||
@dataclass
|
||||
class FlushCacheReqInput:
|
||||
class FlushCacheReqInput(BaseReq):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class FlushCacheReqOutput:
|
||||
class FlushCacheReqOutput(BaseReq):
|
||||
success: bool
|
||||
|
||||
|
||||
@dataclass
|
||||
class UpdateWeightFromDiskReqInput:
|
||||
class UpdateWeightFromDiskReqInput(BaseReq):
|
||||
# The model path with the new weights
|
||||
model_path: str
|
||||
# The format to load the weights
|
||||
@@ -990,7 +986,7 @@ class UpdateWeightFromDiskReqInput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class UpdateWeightFromDiskReqOutput:
|
||||
class UpdateWeightFromDiskReqOutput(BaseReq):
|
||||
success: bool
|
||||
message: str
|
||||
# Number of paused requests during weight sync.
|
||||
@@ -998,7 +994,7 @@ class UpdateWeightFromDiskReqOutput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class UpdateWeightsFromDistributedReqInput:
|
||||
class UpdateWeightsFromDistributedReqInput(BaseReq):
|
||||
names: List[str]
|
||||
dtypes: List[str]
|
||||
shapes: List[List[int]]
|
||||
@@ -1013,13 +1009,13 @@ class UpdateWeightsFromDistributedReqInput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class UpdateWeightsFromDistributedReqOutput:
|
||||
class UpdateWeightsFromDistributedReqOutput(BaseReq):
|
||||
success: bool
|
||||
message: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class UpdateWeightsFromTensorReqInput:
|
||||
class UpdateWeightsFromTensorReqInput(BaseReq):
|
||||
"""Update model weights from tensor input.
|
||||
|
||||
- Tensors are serialized for transmission
|
||||
@@ -1038,13 +1034,13 @@ class UpdateWeightsFromTensorReqInput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class UpdateWeightsFromTensorReqOutput:
|
||||
class UpdateWeightsFromTensorReqOutput(BaseReq):
|
||||
success: bool
|
||||
message: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class InitWeightsSendGroupForRemoteInstanceReqInput:
|
||||
class InitWeightsSendGroupForRemoteInstanceReqInput(BaseReq):
|
||||
# The master address
|
||||
master_address: str
|
||||
# The ports for each rank's communication group
|
||||
@@ -1060,13 +1056,13 @@ class InitWeightsSendGroupForRemoteInstanceReqInput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class InitWeightsSendGroupForRemoteInstanceReqOutput:
|
||||
class InitWeightsSendGroupForRemoteInstanceReqOutput(BaseReq):
|
||||
success: bool
|
||||
message: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class SendWeightsToRemoteInstanceReqInput:
|
||||
class SendWeightsToRemoteInstanceReqInput(BaseReq):
|
||||
# The master address
|
||||
master_address: str
|
||||
# The ports for each rank's communication group
|
||||
@@ -1076,13 +1072,13 @@ class SendWeightsToRemoteInstanceReqInput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class SendWeightsToRemoteInstanceReqOutput:
|
||||
class SendWeightsToRemoteInstanceReqOutput(BaseReq):
|
||||
success: bool
|
||||
message: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class InitWeightsUpdateGroupReqInput:
|
||||
class InitWeightsUpdateGroupReqInput(BaseReq):
|
||||
# The master address
|
||||
master_address: str
|
||||
# The master port
|
||||
@@ -1098,24 +1094,24 @@ class InitWeightsUpdateGroupReqInput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class InitWeightsUpdateGroupReqOutput:
|
||||
class InitWeightsUpdateGroupReqOutput(BaseReq):
|
||||
success: bool
|
||||
message: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class DestroyWeightsUpdateGroupReqInput:
|
||||
class DestroyWeightsUpdateGroupReqInput(BaseReq):
|
||||
group_name: str = "weight_update_group"
|
||||
|
||||
|
||||
@dataclass
|
||||
class DestroyWeightsUpdateGroupReqOutput:
|
||||
class DestroyWeightsUpdateGroupReqOutput(BaseReq):
|
||||
success: bool
|
||||
message: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class UpdateWeightVersionReqInput:
|
||||
class UpdateWeightVersionReqInput(BaseReq):
|
||||
# The new weight version
|
||||
new_version: str
|
||||
# Whether to abort all running requests before updating
|
||||
@@ -1123,89 +1119,87 @@ class UpdateWeightVersionReqInput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class GetWeightsByNameReqInput:
|
||||
class GetWeightsByNameReqInput(BaseReq):
|
||||
name: str
|
||||
truncate_size: int = 100
|
||||
|
||||
|
||||
@dataclass
|
||||
class GetWeightsByNameReqOutput:
|
||||
class GetWeightsByNameReqOutput(BaseReq):
|
||||
parameter: list
|
||||
|
||||
|
||||
@dataclass
|
||||
class ReleaseMemoryOccupationReqInput:
|
||||
class ReleaseMemoryOccupationReqInput(BaseReq):
|
||||
# Optional tags to identify the memory region, which is primarily used for RL
|
||||
# Currently we only support `weights` and `kv_cache`
|
||||
tags: Optional[List[str]] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ReleaseMemoryOccupationReqOutput:
|
||||
class ReleaseMemoryOccupationReqOutput(BaseReq):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class ResumeMemoryOccupationReqInput:
|
||||
class ResumeMemoryOccupationReqInput(BaseReq):
|
||||
# Optional tags to identify the memory region, which is primarily used for RL
|
||||
# Currently we only support `weights` and `kv_cache`
|
||||
tags: Optional[List[str]] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ResumeMemoryOccupationReqOutput:
|
||||
class ResumeMemoryOccupationReqOutput(BaseReq):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class SlowDownReqInput:
|
||||
class SlowDownReqInput(BaseReq):
|
||||
forward_sleep_time: Optional[float]
|
||||
|
||||
|
||||
@dataclass
|
||||
class SlowDownReqOutput:
|
||||
class SlowDownReqOutput(BaseReq):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class AbortReq:
|
||||
# The request id
|
||||
rid: str = ""
|
||||
class AbortReq(BaseReq):
|
||||
# Whether to abort all requests
|
||||
abort_all: bool = False
|
||||
# The finished reason data
|
||||
finished_reason: Optional[Dict[str, Any]] = None
|
||||
abort_reason: Optional[str] = None
|
||||
# used in MultiTokenzierManager mode
|
||||
rids: Optional[Union[List[str], str]] = None
|
||||
|
||||
def __post_init__(self):
|
||||
self.rids = self.rid
|
||||
# FIXME: This is a hack to keep the same with the old code
|
||||
if self.rid is None:
|
||||
self.rid = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class GetInternalStateReq:
|
||||
class GetInternalStateReq(BaseReq):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class GetInternalStateReqOutput:
|
||||
class GetInternalStateReqOutput(BaseReq):
|
||||
internal_state: Dict[Any, Any]
|
||||
|
||||
|
||||
@dataclass
|
||||
class SetInternalStateReq:
|
||||
class SetInternalStateReq(BaseReq):
|
||||
server_args: Dict[str, Any]
|
||||
|
||||
|
||||
@dataclass
|
||||
class SetInternalStateReqOutput:
|
||||
class SetInternalStateReqOutput(BaseReq):
|
||||
updated: bool
|
||||
server_args: Dict[str, Any]
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProfileReqInput:
|
||||
class ProfileReqInput(BaseReq):
|
||||
# The output directory
|
||||
output_dir: Optional[str] = None
|
||||
# If set, it profile as many as this number of steps.
|
||||
@@ -1225,7 +1219,7 @@ class ProfileReqType(Enum):
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProfileReq:
|
||||
class ProfileReq(BaseReq):
|
||||
type: ProfileReqType
|
||||
output_dir: Optional[str] = None
|
||||
start_step: Optional[int] = None
|
||||
@@ -1238,18 +1232,18 @@ class ProfileReq:
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProfileReqOutput:
|
||||
class ProfileReqOutput(BaseReq):
|
||||
success: bool
|
||||
message: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class FreezeGCReq:
|
||||
class FreezeGCReq(BaseReq):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConfigureLoggingReq:
|
||||
class ConfigureLoggingReq(BaseReq):
|
||||
log_requests: Optional[bool] = None
|
||||
log_requests_level: Optional[int] = None
|
||||
dump_requests_folder: Optional[str] = None
|
||||
@@ -1258,35 +1252,39 @@ class ConfigureLoggingReq:
|
||||
|
||||
|
||||
@dataclass
|
||||
class OpenSessionReqInput:
|
||||
class OpenSessionReqInput(BaseReq):
|
||||
capacity_of_str_len: int
|
||||
session_id: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class CloseSessionReqInput:
|
||||
class CloseSessionReqInput(BaseReq):
|
||||
session_id: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class OpenSessionReqOutput:
|
||||
class OpenSessionReqOutput(BaseReq):
|
||||
session_id: Optional[str]
|
||||
success: bool
|
||||
|
||||
|
||||
@dataclass
|
||||
class HealthCheckOutput:
|
||||
class HealthCheckOutput(BaseReq):
|
||||
pass
|
||||
|
||||
|
||||
class ExpertDistributionReq(Enum):
|
||||
class ExpertDistributionReqType(Enum):
|
||||
START_RECORD = 1
|
||||
STOP_RECORD = 2
|
||||
DUMP_RECORD = 3
|
||||
|
||||
|
||||
class ExpertDistributionReq(BaseReq):
|
||||
action: ExpertDistributionReqType
|
||||
|
||||
|
||||
@dataclass
|
||||
class ExpertDistributionReqOutput:
|
||||
class ExpertDistributionReqOutput(BaseReq):
|
||||
pass
|
||||
|
||||
|
||||
@@ -1304,7 +1302,7 @@ class Tool:
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParseFunctionCallReq:
|
||||
class ParseFunctionCallReq(BaseReq):
|
||||
text: str # The text to parse.
|
||||
tools: List[Tool] = field(
|
||||
default_factory=list
|
||||
@@ -1315,31 +1313,31 @@ class ParseFunctionCallReq:
|
||||
|
||||
|
||||
@dataclass
|
||||
class SeparateReasoningReqInput:
|
||||
class SeparateReasoningReqInput(BaseReq):
|
||||
text: str # The text to parse.
|
||||
reasoning_parser: str # Specify the parser type, e.g., "deepseek-r1".
|
||||
|
||||
|
||||
@dataclass
|
||||
class VertexGenerateReqInput:
|
||||
class VertexGenerateReqInput(BaseReq):
|
||||
instances: List[dict]
|
||||
parameters: Optional[dict] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class RpcReqInput:
|
||||
class RpcReqInput(BaseReq):
|
||||
method: str
|
||||
parameters: Optional[Dict] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class RpcReqOutput:
|
||||
class RpcReqOutput(BaseReq):
|
||||
success: bool
|
||||
message: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class LoadLoRAAdapterReqInput:
|
||||
class LoadLoRAAdapterReqInput(BaseReq):
|
||||
# The name of the lora module to newly loaded.
|
||||
lora_name: str
|
||||
# The path of loading.
|
||||
@@ -1359,7 +1357,7 @@ class LoadLoRAAdapterReqInput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class UnloadLoRAAdapterReqInput:
|
||||
class UnloadLoRAAdapterReqInput(BaseReq):
|
||||
# The name of lora module to unload.
|
||||
lora_name: str
|
||||
# The unique identifier for the LoRA adapter, which automatically generated in the `TokenizerManager`.
|
||||
@@ -1373,23 +1371,23 @@ class UnloadLoRAAdapterReqInput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class LoRAUpdateResult:
|
||||
class LoRAUpdateOutput(BaseReq):
|
||||
success: bool
|
||||
error_message: Optional[str] = None
|
||||
loaded_adapters: Optional[Dict[str, LoRARef]] = None
|
||||
|
||||
|
||||
LoadLoRAAdapterReqOutput = UnloadLoRAAdapterReqOutput = LoRAUpdateResult
|
||||
LoadLoRAAdapterReqOutput = UnloadLoRAAdapterReqOutput = LoRAUpdateOutput
|
||||
|
||||
|
||||
@dataclass
|
||||
class MultiTokenizerRegisterReq:
|
||||
rids: Optional[Union[List[str], str]] = None
|
||||
class MultiTokenizerRegisterReq(BaseBatchReq):
|
||||
ipc_name: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class MultiTokenizerWrapper:
|
||||
# FIXME(lsyin): remove this
|
||||
worker_id: int
|
||||
obj: Optional[Any] = None
|
||||
|
||||
@@ -1400,17 +1398,17 @@ class BlockReqType(Enum):
|
||||
|
||||
|
||||
@dataclass
|
||||
class BlockReqInput:
|
||||
class BlockReqInput(BaseReq):
|
||||
type: BlockReqType
|
||||
|
||||
|
||||
@dataclass
|
||||
class GetLoadReqInput:
|
||||
class GetLoadReqInput(BaseReq):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class GetLoadReqOutput:
|
||||
class GetLoadReqOutput(BaseReq):
|
||||
dp_rank: int
|
||||
num_reqs: int
|
||||
num_waiting_reqs: int
|
||||
@@ -1418,5 +1416,31 @@ class GetLoadReqOutput:
|
||||
|
||||
|
||||
@dataclass
|
||||
class WatchLoadUpdateReq:
|
||||
class WatchLoadUpdateReq(BaseReq):
|
||||
loads: List[GetLoadReqOutput]
|
||||
|
||||
|
||||
def _check_all_req_types():
|
||||
"""A helper function to check all request types are defined in this file."""
|
||||
import inspect
|
||||
import sys
|
||||
|
||||
all_classes = inspect.getmembers(sys.modules[__name__], inspect.isclass)
|
||||
for class_type in all_classes:
|
||||
# check its name
|
||||
name = class_type[0]
|
||||
is_io_struct = (
|
||||
name.endswith("Req") or name.endswith("Input") or name.endswith("Output")
|
||||
)
|
||||
is_base_req = issubclass(class_type[1], BaseReq) or issubclass(
|
||||
class_type[1], BaseBatchReq
|
||||
)
|
||||
if is_io_struct and not is_base_req:
|
||||
raise ValueError(f"{name} is not a subclass of BaseReq or BaseBatchReq.")
|
||||
if is_base_req and not is_io_struct:
|
||||
raise ValueError(
|
||||
f"{name} is a subclass of BaseReq but not follow the naming convention."
|
||||
)
|
||||
|
||||
|
||||
_check_all_req_types()
|
||||
|
||||
Reference in New Issue
Block a user