# ============================================================ # 项目 -> 会话 -> 分支 三级树(moduleId: state-branches, 可重生 ✅) # plan.md §4.x / 矩阵 51 行:BranchNode 生命周期(创建/命名/切换/丢弃/合并) # 每会话一棵分支树(tree_root / active_node),父子关系持久化且可审计; # 分支锚定 checkpoint pairId(世界状态点),切换分支即切换成对状态。 # ============================================================ from __future__ import annotations import json import os import tempfile import threading import uuid from typing import Any from server.timeutil import fmt_dt # 分支状态 _ACTIVE = "active" _MERGED = "merged" _DISCARDED = "discarded" def _now() -> str: from datetime import datetime return fmt_dt(datetime.now()) def _uid() -> str: return f"branch_{uuid.uuid4().hex[:10]}" def _safe_scope(value: str) -> str: cleaned = "".join(c for c in (value or "default") if c.isalnum() or c in "-_") return cleaned[:64] or "default" class BranchError(RuntimeError): code = "BRANCH_ERROR" class BranchNotFoundError(BranchError): code = "BRANCH_NOT_FOUND" class BranchConflictError(BranchError): code = "BRANCH_CONFLICT" class BranchStore: """分支树仓:每会话一棵树,JSON 持久化(原子写)+ 并发锁。 结构: {sessions: {session_id: {tree_root, active_node, branches: {id: {...}}}}} """ def __init__(self, path: str | None = None) -> None: self.path = path or os.environ.get( "APS_BRANCH_PATH", "server/data/branches.json") self._lock = threading.RLock() self._data = self._load() def _load(self) -> dict[str, Any]: try: with open(self.path, "r", encoding="utf-8") as f: data = json.load(f) if isinstance(data, dict) and isinstance(data.get("sessions"), dict): return data except (FileNotFoundError, json.JSONDecodeError, TypeError): pass return {"sessions": {}} def _write(self) -> None: os.makedirs(os.path.dirname(self.path) or ".", exist_ok=True) fd, tmp = tempfile.mkstemp(dir=os.path.dirname(self.path) or ".", suffix=".tmp") try: with os.fdopen(fd, "w", encoding="utf-8") as f: json.dump(self._data, f, ensure_ascii=False, sort_keys=True) os.replace(tmp, self.path) except BaseException: if os.path.exists(tmp): os.unlink(tmp) raise # ---------------- 会话树 ---------------- def ensure_session(self, session_id: str) -> None: """确保会话有根分支(惰性创建;幂等)。""" with self._lock: if session_id in self._data["sessions"]: return root = { "id": _uid(), "name": "主干", "parentId": None, "status": _ACTIVE, "createdAt": _now(), "checkpointId": None, } self._data["sessions"][session_id] = { "tree_root": root["id"], "active_node": root["id"], "branches": {root["id"]: root}, } self._write() def _tree(self, session_id: str) -> dict[str, Any]: tree = self._data["sessions"].get(session_id) if tree is None: raise BranchNotFoundError(f"session {session_id!r} has no branch tree") return tree def _branch(self, session_id: str, branch_id: str) -> dict[str, Any]: tree = self._tree(session_id) branch = tree["branches"].get(branch_id) if branch is None: raise BranchNotFoundError(f"branch {branch_id!r} not found in session {session_id!r}") return branch # ---------------- 查询 ---------------- def tree(self, session_id: str) -> dict[str, Any]: """返回会话分支树(只读快照)。""" with self._lock: self.ensure_session(session_id) tree = self._tree(session_id) return { "treeRoot": tree["tree_root"], "activeNode": tree["active_node"], "branches": list(tree["branches"].values()), } def active(self, session_id: str) -> str: """当前活动分支 id。""" with self._lock: self.ensure_session(session_id) return self._tree(session_id)["active_node"] # ---------------- 生命周期 ---------------- def fork(self, session_id: str, name: str, checkpoint_id: str | None = None) -> dict[str, Any]: """从当前活动分支派生新分支(继承活动分支的 checkpoint 锚点)。""" with self._lock: self.ensure_session(session_id) tree = self._tree(session_id) parent_id = tree["active_node"] parent = tree["branches"][parent_id] branch = { "id": _uid(), "name": name.strip() or "新分支", "parentId": parent_id, "status": _ACTIVE, "createdAt": _now(), "checkpointId": checkpoint_id or parent.get("checkpointId"), } tree["branches"][branch["id"]] = branch tree["active_node"] = branch["id"] self._write() return branch def rename(self, session_id: str, branch_id: str, name: str) -> dict[str, Any]: with self._lock: branch = self._branch(session_id, branch_id) if branch["status"] != _ACTIVE: raise BranchConflictError("只能重命名活动状态的分支") branch["name"] = name.strip() or branch["name"] self._write() return branch def switch(self, session_id: str, branch_id: str) -> dict[str, Any]: """切换活动分支;返回目标分支(含 checkpointId 供成对状态恢复)。""" with self._lock: tree = self._tree(session_id) branch = self._branch(session_id, branch_id) if branch["status"] == _DISCARDED: raise BranchConflictError("已丢弃的分支不可切换") tree["active_node"] = branch_id self._write() return branch def discard(self, session_id: str, branch_id: str) -> dict[str, Any]: """丢弃分支(仅允许丢弃非主干、非当前活动分支;标记 DISCARDED 保留审计)。""" with self._lock: tree = self._tree(session_id) branch = self._branch(session_id, branch_id) if branch_id == tree["tree_root"]: raise BranchConflictError("主干分支不可丢弃") if branch_id == tree["active_node"]: raise BranchConflictError("当前活动分支不可丢弃,请先切换") if branch["status"] == _MERGED: raise BranchConflictError("已合并的分支不可再次丢弃") branch["status"] = _DISCARDED self._write() return branch def merge(self, session_id: str, source_id: str, target_id: str) -> dict[str, Any]: """合并:source 并入 target(target 保持活动);source 标记 MERGED。""" with self._lock: tree = self._tree(session_id) source = self._branch(session_id, source_id) target = self._branch(session_id, target_id) if source["status"] == _MERGED or source["status"] == _DISCARDED: raise BranchConflictError("源分支不可合并(已合并/已丢弃)") if target["status"] == _DISCARDED: raise BranchConflictError("目标分支已丢弃") if source_id == target_id: raise BranchConflictError("不能合并到自身") source["status"] = _MERGED source["mergedInto"] = target_id target["checkpointId"] = source.get("checkpointId") or target.get("checkpointId") tree["active_node"] = target_id self._write() return {"source": source, "target": target} # ---------------- 检查点世界 diff(矩阵 51 方向:merge diff 视图) ---------------- def _safe_projection(world: dict[str, Any]) -> dict[str, Any]: """世界状态的安全只读投影:任一视图异常都降级为空结构,不阻断合并""" from server.aps_domain.flex import capacity_analysis, flex_gantt_view from server.aps_domain.views import gantt_view, world_summary empty = { "summary": {"hasVersion": False}, "gantt": {"lines": [], "workstations": [], "workOrders": []}, "flexGantt": {"versionNo": None, "sortMode": None, "zones": [], "equipment": [], "workOrders": []}, "flexCapacity": {"pools": [], "bottleneckPool": None}, } try: proj = { "summary": world_summary(world), "gantt": gantt_view(world), "flexGantt": flex_gantt_view(world), "flexCapacity": capacity_analysis(world), } for key, fallback in empty.items(): if not isinstance(proj.get(key), dict): proj[key] = fallback return proj except Exception: # noqa: BLE001 - 投影降级边界:任何视图异常都不阻断合并 return empty def _num(value: Any) -> int | float | None: """JSON 数字规范化:NaN/Infinity 归为 None,保证 diff 结果稳定""" if isinstance(value, bool) or value is None: return None if isinstance(value, int): return value if isinstance(value, float): import math return value if math.isfinite(value) else None return None def _pair(src: Any, dst: Any) -> dict[str, Any]: """{source, target, delta} 三栏结构;任一侧非数字时 delta=None""" s, t = _num(src), _num(dst) delta = (t - s) if (s is not None and t is not None) else None return {"source": s, "target": t, "delta": delta} def compute_world_diff(source_world: dict[str, Any] | None, target_world: dict[str, Any] | None, *, source_label: str = "source", target_label: str = "target") -> dict[str, Any]: """两个 checkpoint 世界的结构化 diff(矩阵 51 方向:merge diff) - 投影仅保留 world_summary / gantt_view / flex_gantt_view / capacity_analysis 四个只读视图(矩阵 54 方向:投影降级不阻断) - 同一 checkpoint / 同一视图逐项对比,diff 结果均为可序列化 JSON - 变化字段含 version / kpi / operationTimeDiffs / flexVersion / capacityPools / changed,供前端并排展示 """ import copy sa = _safe_projection(copy.deepcopy(source_world or {})) ta = _safe_projection(copy.deepcopy(target_world or {})) # ---- 版本元信息(versionNo / status / KPI) ---- version = { "sourceVersionNo": sa["summary"].get("versionNo"), "targetVersionNo": ta["summary"].get("versionNo"), "sourceStatus": sa["summary"].get("status"), "targetStatus": ta["summary"].get("status"), "changed": (sa["summary"].get("versionNo") != ta["summary"].get("versionNo") or sa["summary"].get("status") != ta["summary"].get("status")), } kpi = { "workOrderCount": _pair(sa["summary"].get("woCount"), ta["summary"].get("woCount")), "conflictCount": _pair(sa["summary"].get("conflictCount"), ta["summary"].get("conflictCount")), "totalTardiness": _pair(sa["summary"].get("totalTardiness"), ta["summary"].get("totalTardiness")), "avgUtilization": _pair(sa["summary"].get("avgUtilization"), ta["summary"].get("avgUtilization")), } # ---- 工单时间对比:固定轨 + 柔性轨逐活动计算 start/end 偏移 ---- def _fixed_key(row: dict[str, Any]) -> tuple[str, Any]: return ("fixed", row.get("orderNo"), row.get("operationName"), row.get("id")) def _flex_key(row: dict[str, Any]) -> tuple[str, Any]: return ("flex", row.get("orderNo"), row.get("woNo"), row.get("operationName"), row.get("seq")) def _collect(proj: dict[str, Any], key_fn) -> dict[tuple[str, Any], dict[str, Any]]: out: dict[tuple[str, Any], dict[str, Any]] = {} for track_rows, track in ((proj["gantt"].get("workOrders") or [], "fixed"), (proj["flexGantt"].get("workOrders") or [], "flex")): for row in track_rows: out[key_fn(row)] = dict(row) return out src_rows = _collect(sa, _fixed_key) src_rows.update(_collect(sa, _flex_key)) dst_rows = _collect(ta, _fixed_key) dst_rows.update(_collect(ta, _flex_key)) op_diffs: list[dict[str, Any]] = [] for key in sorted(set(src_rows) | set(dst_rows)): a, b = src_rows.get(key), dst_rows.get(key) if a is None or b is None: changed = True # 新增/删除活动直接标记变化 flags = {"added": a is None, "deleted": b is None} else: changed = (a.get("start") != b.get("start") or a.get("end") != b.get("end") or a.get("equipmentCode") != b.get("equipmentCode") or a.get("moldCode") != b.get("moldCode")) flags = {} if changed: op_diffs.append({ "track": key[0], "orderNo": b.get("orderNo") if b else a.get("orderNo"), "operationName": b.get("operationName") if b else a.get("operationName"), "source": {"start": (a or {}).get("start"), "end": (a or {}).get("end"), "equipment": (a or {}).get("equipmentCode"), "mold": (a or {}).get("moldCode")}, "target": {"start": (b or {}).get("start"), "end": (b or {}).get("end"), "equipment": (b or {}).get("equipmentCode"), "mold": (b or {}).get("moldCode")}, **flags, }) op_diffs.sort(key=lambda d: (d["track"], str(d["orderNo"]), str(d["operationName"]))) truncated = len(op_diffs) > 20 operation_time_diffs = { "count": len(op_diffs), "orders": op_diffs[:20], "truncated": truncated, } # ---- 柔性版本对比 ---- flex_version = { "sourceVersionNo": sa["flexGantt"].get("versionNo"), "targetVersionNo": ta["flexGantt"].get("versionNo"), "sourceSortMode": sa["flexGantt"].get("sortMode"), "targetSortMode": ta["flexGantt"].get("sortMode"), "changed": (sa["flexGantt"].get("versionNo") != ta["flexGantt"].get("versionNo") or sa["flexGantt"].get("sortMode") != ta["flexGantt"].get("sortMode")), } # ---- 产能池对比 ---- src_pools = {p.get("operationCode"): p for p in (sa["flexCapacity"].get("pools") or [])} dst_pools = {p.get("operationCode"): p for p in (ta["flexCapacity"].get("pools") or [])} pool_rows: list[dict[str, Any]] = [] for code in sorted(set(src_pools) | set(dst_pools)): a, b = src_pools.get(code), dst_pools.get(code) changed = (a is None or b is None or a.get("equipmentCount") != b.get("equipmentCount") or a.get("dailyCapacity") != b.get("dailyCapacity")) if changed: pool_rows.append({ "operationCode": code, "operationName": (b or a or {}).get("operationName"), "source": {"equipmentCount": (a or {}).get("equipmentCount"), "dailyCapacity": (a or {}).get("dailyCapacity")}, "target": {"equipmentCount": (b or {}).get("equipmentCount"), "dailyCapacity": (b or {}).get("dailyCapacity")}, }) capacity_pools = { "sourceCount": len(src_pools), "targetCount": len(dst_pools), "changed": len(pool_rows) > 0 or len(src_pools) != len(dst_pools), "pools": pool_rows[:20], } changed = (version["changed"] or flex_version["changed"] or capacity_pools["changed"] or operation_time_diffs["count"] > 0) return { "computedAt": _now(), "source": {"label": source_label, "checkpointId": None}, "target": {"label": target_label, "checkpointId": None}, "version": version, "kpi": kpi, "operationTimeDiffs": operation_time_diffs, "flexVersion": flex_version, "capacityPools": capacity_pools, "changed": changed, } _branches: dict[tuple[str, str], BranchStore] = {} _branches_lock = threading.Lock() def get_branches(tenant_uuid: str = "platform", world_key: str = "default") -> BranchStore: """按租户/项目隔离的分支树仓实例。""" key = (_safe_scope(tenant_uuid), _safe_scope(world_key)) with _branches_lock: store = _branches.get(key) if store is None: base = os.environ.get("APS_BRANCH_DIR", "server/data") store = BranchStore(os.path.join(base, f"branches-{key[0]}-{key[1]}.json")) _branches[key] = store return store def reset_branch_stores() -> None: """清空分支树单例(测试/重配置用)。""" with _branches_lock: _branches.clear()