local_output_sink.py 6.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175
  1. from __future__ import annotations
  2. import json
  3. from datetime import datetime
  4. from pathlib import Path
  5. from typing import Any
  6. from zoneinfo import ZoneInfo
  7. from examples.demand.data_query_tools import (
  8. build_dwd_multi_demand_pool_di_rows,
  9. build_feature_point_data_rows,
  10. )
  11. OUTPUT_FILENAMES = {
  12. "run_manifest": "run_manifest.json",
  13. "demand_task": "demand_task.json",
  14. "demand_items": "demand_items.json",
  15. "rejected_demand_items": "rejected_demand_items.json",
  16. "demand_content": "demand_content.json",
  17. "dwd_multi_demand_pool_di": "dwd_multi_demand_pool_di.json",
  18. "feature_point_data": "feature_point_data.json",
  19. }
  20. _CHINA_TZ = ZoneInfo("Asia/Shanghai")
  21. _DT_FMT = "%Y%m%d"
  22. def _today_dt() -> str:
  23. return datetime.now(_CHINA_TZ).strftime(_DT_FMT)
  24. def _as_list(value: Any) -> list[Any]:
  25. if value is None:
  26. return []
  27. if isinstance(value, list):
  28. return value
  29. if isinstance(value, tuple):
  30. return list(value)
  31. return [value]
  32. def _normalize_demand_content_row(row: Any) -> Any:
  33. if not isinstance(row, dict):
  34. return row
  35. normalized = dict(row)
  36. ext_data = normalized.get("ext_data")
  37. if isinstance(ext_data, str) and ext_data.strip():
  38. try:
  39. normalized["ext_data"] = json.loads(ext_data)
  40. except json.JSONDecodeError:
  41. pass
  42. return normalized
  43. def _write_json(path: Path, payload: Any) -> Path:
  44. path.parent.mkdir(parents=True, exist_ok=True)
  45. tmp_path = path.with_name(f".{path.name}.tmp")
  46. tmp_path.write_text(
  47. json.dumps(payload, ensure_ascii=False, indent=2, default=str) + "\n",
  48. encoding="utf-8",
  49. )
  50. tmp_path.replace(path)
  51. return path
  52. class LocalOutputSink:
  53. """Writes DemandAgent local-json output mirrors for one run."""
  54. def __init__(self, output_dir: str | Path):
  55. self.output_dir = Path(output_dir)
  56. @classmethod
  57. def for_run(cls, base_dir: str | Path, run_id: str | int) -> "LocalOutputSink":
  58. return cls(Path(base_dir) / str(run_id))
  59. def path_for(self, name: str) -> Path:
  60. try:
  61. filename = OUTPUT_FILENAMES[name]
  62. except KeyError as exc:
  63. raise ValueError(f"unknown local output name: {name}") from exc
  64. return self.output_dir / filename
  65. def default_run_manifest(self) -> dict[str, Any]:
  66. return {
  67. "output_mode": "local_json",
  68. "output_dir": str(self.output_dir),
  69. "created_at": datetime.now(_CHINA_TZ).isoformat(),
  70. }
  71. def initialize(self, run_manifest: dict[str, Any] | None = None) -> dict[str, Path]:
  72. return self.write_all(run_manifest=run_manifest)
  73. def write_run_manifest(self, payload: dict[str, Any] | None = None) -> Path:
  74. if payload is None:
  75. payload = self.default_run_manifest()
  76. return _write_json(self.path_for("run_manifest"), payload)
  77. def write_demand_task(self, payload: Any = None) -> Path:
  78. if payload is None:
  79. payload = {}
  80. return _write_json(self.path_for("demand_task"), payload)
  81. def write_demand_items(self, items: Any = None) -> Path:
  82. return _write_json(self.path_for("demand_items"), _as_list(items))
  83. def write_rejected_demand_items(self, items: Any = None) -> Path:
  84. return _write_json(self.path_for("rejected_demand_items"), _as_list(items))
  85. def write_demand_content(self, rows: Any = None) -> Path:
  86. normalized_rows = [_normalize_demand_content_row(row) for row in _as_list(rows)]
  87. return _write_json(self.path_for("demand_content"), normalized_rows)
  88. def write_dwd_multi_demand_pool_di(self, rows: Any = None) -> Path:
  89. return _write_json(self.path_for("dwd_multi_demand_pool_di"), _as_list(rows))
  90. def write_dwd_multi_demand_pool_di_from_demand_content(
  91. self,
  92. rows: list[dict],
  93. partition_dt: str | None = None,
  94. ) -> Path:
  95. output_rows = build_dwd_multi_demand_pool_di_rows(rows=rows, partition_dt=partition_dt or _today_dt())
  96. return self.write_dwd_multi_demand_pool_di(output_rows)
  97. def write_feature_point_data(self, rows: Any = None) -> Path:
  98. return _write_json(self.path_for("feature_point_data"), _as_list(rows))
  99. def write_feature_point_data_from_names(self, names: list[str], dt: str | None = None) -> Path:
  100. output_rows = build_feature_point_data_rows(names=names, dt=dt or _today_dt())
  101. return self.write_feature_point_data(output_rows)
  102. def write_all(
  103. self,
  104. *,
  105. run_manifest: dict[str, Any] | None = None,
  106. demand_task: Any = None,
  107. demand_items: Any = None,
  108. rejected_demand_items: Any = None,
  109. demand_content_rows: Any = None,
  110. dwd_multi_demand_pool_di_rows: Any = None,
  111. feature_point_data_rows: Any = None,
  112. ) -> dict[str, Path]:
  113. paths = {
  114. "run_manifest": self.write_run_manifest(run_manifest),
  115. "demand_task": self.write_demand_task(demand_task),
  116. "demand_items": self.write_demand_items(demand_items),
  117. "rejected_demand_items": self.write_rejected_demand_items(rejected_demand_items),
  118. "demand_content": self.write_demand_content(demand_content_rows),
  119. "dwd_multi_demand_pool_di": self.write_dwd_multi_demand_pool_di(dwd_multi_demand_pool_di_rows),
  120. "feature_point_data": self.write_feature_point_data(feature_point_data_rows),
  121. }
  122. return paths
  123. def write_local_outputs(
  124. output_dir: str | Path,
  125. *,
  126. run_manifest: dict[str, Any] | None = None,
  127. demand_task: Any = None,
  128. demand_items: Any = None,
  129. rejected_demand_items: Any = None,
  130. demand_content_rows: Any = None,
  131. dwd_multi_demand_pool_di_rows: Any = None,
  132. feature_point_data_rows: Any = None,
  133. ) -> dict[str, Path]:
  134. return LocalOutputSink(output_dir).write_all(
  135. run_manifest=run_manifest,
  136. demand_task=demand_task,
  137. demand_items=demand_items,
  138. rejected_demand_items=rejected_demand_items,
  139. demand_content_rows=demand_content_rows,
  140. dwd_multi_demand_pool_di_rows=dwd_multi_demand_pool_di_rows,
  141. feature_point_data_rows=feature_point_data_rows,
  142. )