128 lines
5.0 KiB
Python
128 lines
5.0 KiB
Python
import copy
|
||
|
||
from pydantic import BaseModel
|
||
|
||
from ...sync_system.strategy import DefaultSyncStrategy
|
||
from ...sync_system.config import StrategyConfig, OrphanAction, UpdateDirection
|
||
from schemas.project.project_base import ProjectResponseBase
|
||
|
||
|
||
class ProjectDomainOption(BaseModel):
|
||
pull_group_plan_production_date: bool = False
|
||
|
||
|
||
class ProjectSyncStrategy(DefaultSyncStrategy[ProjectResponseBase]):
|
||
"""
|
||
Project 同步策略
|
||
|
||
业务规则:
|
||
1. 手动绑定(根节点)
|
||
2. 允许拉取更新(不自动创建)
|
||
3. 依赖公司(company)
|
||
|
||
schema 已在类定义中设置,禁止在初始化时传入其他 schema。
|
||
"""
|
||
|
||
# 类变量:schema 固定为 ProjectResponseBase
|
||
schema = ProjectResponseBase
|
||
domain_option_model = ProjectDomainOption
|
||
|
||
_PULL_ONLY_FIELDS = [
|
||
"ic_plan_first_production_date",
|
||
"ic_plan_full_production_date",
|
||
]
|
||
|
||
# 类变量:默认配置
|
||
default_config = StrategyConfig(
|
||
auto_bind=False,
|
||
auto_bind_fields=[],
|
||
depend_fields={"company_id": "company"},
|
||
local_orphan_action=OrphanAction.CREATE_REMOTE,
|
||
remote_orphan_action=OrphanAction.NONE,
|
||
update_direction=UpdateDirection.PUSH,
|
||
)
|
||
|
||
@property
|
||
def domain_option(self) -> ProjectDomainOption:
|
||
return ProjectDomainOption.model_validate(self.config.domain_option)
|
||
|
||
def get_node_update_payload(self, node):
|
||
payload = dict(node.get_data() or {})
|
||
context = getattr(node, "context", None)
|
||
if not isinstance(context, dict):
|
||
return payload
|
||
|
||
all_info = context.get("all_info")
|
||
if not isinstance(all_info, dict):
|
||
return payload
|
||
|
||
for section_name, section_data in all_info.items():
|
||
if not section_name.endswith("_extension") or not isinstance(section_data, dict):
|
||
continue
|
||
for field_name in self._PULL_ONLY_FIELDS:
|
||
if payload.get(field_name) not in (None, ""):
|
||
continue
|
||
if field_name in section_data:
|
||
payload[field_name] = section_data.get(field_name)
|
||
|
||
return payload
|
||
|
||
@staticmethod
|
||
def _snapshot_node(node):
|
||
# project 是一个很特殊的双向更新场景:
|
||
# 1. 先做一次 PULL,把远端 ic_plan_* 回填到本地;
|
||
# 2. 再做一次常规 PUSH,把本地其余字段推到远端。
|
||
#
|
||
# 但 prepare_update_for_direction -> e30_update_prepare 会校验
|
||
# source_node 当前仍然处于可进入更新准备的稳定态(S01)。
|
||
# 如果直接复用同一个真实节点作为两次 update 的 source,第一次
|
||
# 准备完成后真实节点的 action/status 已经变化,第二次就不再满足
|
||
# “source 仍在 S01”的前置条件,导致第二次更新被状态机拒绝。
|
||
#
|
||
# 这里复制一个只用于“提供源数据/源状态视图”的快照,让两次 update
|
||
# 都从各自的原始稳定态出发做判断;真正被推进到 S07 / 写入 payload
|
||
# 的仍然是目标侧真实节点,而不是这个 snapshot。
|
||
snapshot = copy.deepcopy(node)
|
||
if hasattr(snapshot, "sync_log"):
|
||
snapshot.sync_log = None
|
||
return snapshot
|
||
|
||
async def update_pair(self, local_node, remote_node, data_id_map=None):
|
||
updated_by_id = {}
|
||
# 注意:这两个 snapshot 只会在各自方向上充当 source_node。
|
||
# - PULL: source=remote_snapshot, target=local_node
|
||
# - PUSH: source=local_snapshot, target=remote_node
|
||
#
|
||
# 因此真实 project 节点的状态变更仍只发生在 target 节点上,
|
||
# 不会因为传入 snapshot 而把本体节点替换掉,也不会把本体节点的
|
||
# action/status/error/sync_log 写到 snapshot 上后丢失掉。
|
||
local_snapshot = self._snapshot_node(local_node)
|
||
remote_snapshot = self._snapshot_node(remote_node)
|
||
|
||
if self.domain_option.pull_group_plan_production_date:
|
||
pull_update = await self.prepare_update_for_direction(
|
||
local_node,
|
||
remote_snapshot,
|
||
UpdateDirection.PULL,
|
||
data_id_map,
|
||
include_fields=self._PULL_ONLY_FIELDS,
|
||
)
|
||
if pull_update is not None:
|
||
updated_by_id[pull_update.node_id] = pull_update
|
||
|
||
generic_exclude_fields = self._PULL_ONLY_FIELDS
|
||
else:
|
||
generic_exclude_fields = None
|
||
|
||
if self.config.update_direction != UpdateDirection.NONE:
|
||
generic_update = await self.prepare_update_for_direction(
|
||
local_snapshot,
|
||
remote_node,
|
||
self.config.update_direction,
|
||
data_id_map,
|
||
exclude_fields=generic_exclude_fields,
|
||
)
|
||
if generic_update is not None:
|
||
updated_by_id[generic_update.node_id] = generic_update
|
||
|
||
return list(updated_by_id.values()) |