251 lines
10 KiB
Python
251 lines
10 KiB
Python
"""User-input adapter shared by HTTP and CLI; the numerical layer stays SI-only.
|
|
|
|
JSON v1 numbers (including decimal strings) were SI, but expressions used the
|
|
selected unit. JSON v2 consistently uses the selected unit for both. Missing
|
|
parameters use catalog defaults, which are always SI. Never infer a format from
|
|
magnitudes or relabel a legacy project without converting its values.
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import math
|
|
from pathlib import Path
|
|
import re
|
|
|
|
UNIT_TABLE = json.loads((Path(__file__).resolve().parent.parent / "schemas" / "parameter-units.json").read_text(encoding="utf-8"))
|
|
DECIMAL = re.compile(r"[+-]?(?:\d+(?:\.\d*)?|\.\d+)(?:[eE][+-]?\d+)?\Z", re.ASCII)
|
|
TOKEN = re.compile(r"(?:\d+(?:\.\d*)?|\.\d+)(?:[eE][+-]?\d+)?|[A-Za-z_][A-Za-z_0-9]*|\*\*|[+*/^(),-]", re.ASCII)
|
|
|
|
|
|
def finite(value: float) -> float:
|
|
if not math.isfinite(value):
|
|
raise ValueError("Parameter expression must produce a finite real number.")
|
|
return value
|
|
|
|
|
|
def numeric_literal(value: object) -> float | None:
|
|
if type(value) in (int, float):
|
|
try:
|
|
return finite(float(value))
|
|
except OverflowError as exc:
|
|
raise ValueError("Parameter magnitude exceeds finite float range.") from exc
|
|
if isinstance(value, str) and DECIMAL.fullmatch(value.strip()):
|
|
return finite(float(value))
|
|
return None
|
|
|
|
|
|
def expression_value(source: str) -> float:
|
|
"""Same bounded recursive-descent grammar as parameterExpression.ts; no eval."""
|
|
source = source.strip().removeprefix("=").strip()
|
|
if not source or len(source) > 512:
|
|
raise ValueError("Parameter expression must contain 1..512 characters.")
|
|
tokens: list[str] = []
|
|
position = 0
|
|
while position < len(source):
|
|
if source[position].isspace():
|
|
position += 1
|
|
continue
|
|
match = TOKEN.match(source, position)
|
|
if match is None:
|
|
raise ValueError(f"Unsupported expression character at {position + 1}.")
|
|
tokens.append(match[0])
|
|
position = match.end()
|
|
if len(tokens) > 256:
|
|
raise ValueError("Parameter expression exceeds 256 tokens.")
|
|
tokens.append("")
|
|
index = 0
|
|
operations = 0
|
|
|
|
def current():
|
|
return tokens[index]
|
|
|
|
def take():
|
|
nonlocal index
|
|
token = current()
|
|
if token:
|
|
index += 1
|
|
return token
|
|
|
|
def operation():
|
|
nonlocal operations
|
|
operations += 1
|
|
if operations > 256:
|
|
raise ValueError("Parameter expression exceeds 256 operations.")
|
|
|
|
def depth_check(depth):
|
|
if depth > 32:
|
|
raise ValueError("Parameter expression exceeds 32 nesting levels.")
|
|
|
|
def additive(depth):
|
|
value = multiplicative(depth)
|
|
while current() in ("+", "-"):
|
|
op = take()
|
|
right = multiplicative(depth)
|
|
operation()
|
|
value = finite(value + right if op == "+" else value - right)
|
|
return value
|
|
|
|
def multiplicative(depth):
|
|
value = unary(depth)
|
|
while current() in ("*", "/"):
|
|
op = take()
|
|
right = unary(depth)
|
|
operation()
|
|
value = finite(value * right if op == "*" else value / right)
|
|
return value
|
|
|
|
def unary(depth):
|
|
depth_check(depth)
|
|
if current() in ("+", "-"):
|
|
op = take()
|
|
operation()
|
|
value = unary(depth + 1)
|
|
return value if op == "+" else -value
|
|
return power(depth)
|
|
|
|
def power(depth):
|
|
depth_check(depth)
|
|
value = primary(depth)
|
|
if current() in ("^", "**"):
|
|
take()
|
|
exponent = unary(depth + 1)
|
|
operation()
|
|
value = finite(math.pow(value, exponent))
|
|
return value
|
|
|
|
def primary(depth):
|
|
depth_check(depth)
|
|
token = take()
|
|
if token == "(":
|
|
value = additive(depth + 1)
|
|
if take() != ")":
|
|
raise ValueError("Missing closing parenthesis.")
|
|
return value
|
|
if token and (token[0].isdigit() or token[0] == "."):
|
|
return finite(float(token))
|
|
name = token.lower()
|
|
if token and (token[0].isalpha() or token[0] == "_"):
|
|
if current() != "(":
|
|
if name in ("pi", "e"):
|
|
return math.pi if name == "pi" else math.e
|
|
raise ValueError(f"Unknown identifier: {token}.")
|
|
depth_check(depth + 1)
|
|
take()
|
|
args = []
|
|
if current() != ")":
|
|
while True:
|
|
if len(args) >= 16:
|
|
raise ValueError("Functions accept at most 16 arguments.")
|
|
args.append(additive(depth + 1))
|
|
if current() != ",":
|
|
break
|
|
take()
|
|
if take() != ")":
|
|
raise ValueError("Missing function closing parenthesis.")
|
|
operation()
|
|
functions = {"sqrt": math.sqrt, "abs": abs, "sin": math.sin,
|
|
"cos": math.cos, "tan": math.tan, "asin": math.asin,
|
|
"acos": math.acos, "atan": math.atan, "exp": math.exp,
|
|
"ln": math.log, "log": math.log, "log10": math.log10,
|
|
"pow": math.pow}
|
|
if name in ("min", "max") and args:
|
|
return finite((min if name == "min" else max)(args))
|
|
if name not in functions or len(args) != (2 if name == "pow" else 1):
|
|
raise ValueError(f"Unsupported function or argument count: {token}.")
|
|
return finite(functions[name](*args))
|
|
raise ValueError("Expected a number, constant or function.")
|
|
|
|
try:
|
|
result = additive(0)
|
|
if current():
|
|
raise ValueError("Unexpected trailing expression content.")
|
|
return finite(result)
|
|
except (ArithmeticError, RecursionError) as exc:
|
|
raise ValueError("Invalid arithmetic or expression domain.") from exc
|
|
|
|
|
|
def unit_conversion(definition, unit: str) -> tuple[float, float]:
|
|
options = UNIT_TABLE.get(definition.quantity, {}) if definition.unit else {}
|
|
if unit in options:
|
|
scale, offset, _ = options[unit]
|
|
return scale, offset
|
|
if unit == definition.unit:
|
|
return 1.0, 0.0
|
|
raise ValueError(f"Unsupported unit '{unit}' for {definition.name} ({definition.quantity}).")
|
|
|
|
|
|
def prepare_project(project):
|
|
"""Copy external input to a current-version, numeric SI execution project.
|
|
|
|
Returns consolidated version notices to the caller. Does not mutate saved
|
|
data and does not weaken the strict XML/native model-version checks.
|
|
"""
|
|
from app.simulation.registry import get_component_model_spec
|
|
|
|
normalized = project.model_copy(deep=True)
|
|
notices = []
|
|
specs = {}
|
|
for node in normalized.nodes:
|
|
model = node.data
|
|
spec = specs.get(model.modelType)
|
|
if spec is None:
|
|
spec = get_component_model_spec(model.modelType)
|
|
specs[model.modelType] = spec
|
|
if model.componentType != spec.model_type:
|
|
raise ValueError(f"COMPONENT_MODEL_TYPE_MISMATCH: {node.id}.")
|
|
if model.modelVersion != spec.model_version:
|
|
notices.append({"componentId": node.id, "label": model.label or node.id,
|
|
"storedVersion": model.modelVersion,
|
|
"currentVersion": spec.model_version})
|
|
for name, value in model.parameters.items():
|
|
definition = spec.parameter_by_name.get(name)
|
|
if definition is None:
|
|
raise ValueError(f"Component '{node.id}' contains unsupported parameters: {name}.")
|
|
try:
|
|
scale, offset = unit_conversion(definition, model.parameterUnits.get(name, definition.unit))
|
|
number = numeric_literal(value)
|
|
is_expression = number is None
|
|
if is_expression:
|
|
if not isinstance(value, str) or definition.editor:
|
|
raise ValueError("Expected a numeric value; discrete parameters cannot use expressions.")
|
|
number = expression_value(value)
|
|
if project.projectSchemaVersion == 2 or is_expression:
|
|
number = finite(number * scale + offset)
|
|
model.parameters[name] = number
|
|
except ValueError as exc:
|
|
raise ValueError(f"{node.id}.{name}: {exc}") from exc
|
|
# Explicit legacy migrations also used by the browser.
|
|
if model.modelVersion == "0.1.0" and model.modelType == "amesim_forc":
|
|
model.parameters.setdefault("direction", 1.0)
|
|
if model.modelVersion == "0.1.0" and model.modelType == "amesim_lmechn1":
|
|
count = model.parameters.get("v1")
|
|
if count in range(1, 9):
|
|
for edge in normalized.edges:
|
|
if edge.source == node.id and edge.sourceHandle == "port_9":
|
|
edge.sourceHandle = f"port_{int(count) + 1}"
|
|
if edge.target == node.id and edge.targetHandle == "port_9":
|
|
edge.targetHandle = f"port_{int(count) + 1}"
|
|
model.parameters["sum"] = 1.0
|
|
# Historical LMECHN1 exposed only nine ports; use its migrated contract.
|
|
from app.main import ReactFlowPortDefinition
|
|
model.ports = [ReactFlowPortDefinition(name=p.name, kind=p.kind, domain=p.domain,
|
|
nominalRole=p.nominal_role, positiveFlowDirection=p.positive_flow_direction)
|
|
for p in spec.ports]
|
|
model.modelVersion = spec.model_version
|
|
model.parameterUnits = {p.name: p.unit for p in spec.parameters}
|
|
model.parameterScientificNotation = {}
|
|
for name in ("t_start", "t_stop", "step", "max_step"):
|
|
value = getattr(normalized.simulation, name)
|
|
number = numeric_literal(value)
|
|
if number is None:
|
|
number = expression_value(value)
|
|
setattr(normalized.simulation, name, number)
|
|
normalized.projectSchemaVersion = 1 # Internal numeric SI contract, never a v2 wire payload.
|
|
return normalized, notices
|
|
|
|
|
|
def version_warning(notices):
|
|
return {"code": "COMPONENT_MODEL_VERSION_WARNING",
|
|
"message": "旧版或版本未知的组件将使用当前模型执行,可能仿真失败或结果与实际不符。",
|
|
"components": notices}
|