from __future__ import annotations import argparse import base64 import html import importlib.util import json import os import re from dataclasses import asdict, dataclass from pathlib import Path from typing import Any, Iterable from PIL import Image, ImageDraw from modules.asset_regenerator.asset_regenerator import _absolute_bbox, _element_transform, _load_svg from modules.asset_regenerator.asset_regenerator import render_svg_to_png from modules.slot_layout_planner.chart_sanitizer import ( is_inside_chart_internal_slot, sanitize_chart_svg, ) @dataclass class SlotPlanItem: id: str kind: str bbox_px: tuple[int, int, int, int] may_overlap_chart: bool element: str anchor: str collision_rule: str asset_source_policy: str slot_requirements: list[str] @dataclass class PlannedSlotPackage: slot_plan: Path planned_svg: Path planned_png: Path reference_map: Path sanitized_svg: Path sanitized_png: Path sanitizer_report: Path @dataclass class ProtectedRegion: id: str reason: str bbox_px: tuple[int, int, int, int] source_slot: str anchor: str = "" collision_rule: str = "" semantic: str = "" GPT_LAYOUT_INSTRUCTIONS = """You are designing editable generation slots for a data infographic. Return valid JSON only. Do not include markdown fences. The chart image is already rendered and will be placed at chart_bbox_px on the final canvas. Design 3 to 6 editable slots for GPT image editing around or near that chart. The slots should make the final piece feel intentionally designed, not like a fixed template. Hard constraints: - Use the provided canvas coordinate system in pixels. - Return only editable slots; do not return chart_internal_slots. They are collected separately. - Avoid every chart_internal_slot bbox completely. - Do not cover axis labels, numeric values, category labels, legends, or data marks unless the slot is a small callout and may_overlap_chart is true. - If allow_chart_overlap is false, every returned slot must avoid the chart_bbox_px. - Keep text slots large enough for legible typography and image slots large enough for clear imagery. - Do not invent exact numeric values in slot text. Required JSON shape: { "background": "#F7F8F4", "design_rationale": "short explanation", "slots": [ { "id": "short_snake_case_id", "kind": "title|subtitle|image|callout|context|annotation|badge", "bbox_px": [x, y, width, height], "may_overlap_chart": false, "element": "what the image model should generate inside this slot", "anchor": "where it sits in the composition", "collision_rule": "explicit preservation rule", "asset_source_policy": "llm_generated_text|image2_generated|llm_generated_text_and_shape|llm_generated_text_and_image", "slot_requirements": ["short requirement", "short requirement"] } ] } """ ALLOWED_SLOT_POLICIES = { "llm_generated_text", "image2_generated", "llm_generated_text_and_shape", "llm_generated_text_and_image", } def _read_json(path: Path | None) -> dict[str, Any]: if not path or not path.is_file(): return {} try: return json.loads(path.read_text(encoding="utf-8")) except Exception: return {} def _load_openai_config( *, api_key: str | None = None, api_key_env: str = "OPENAI_API_KEY", base_url: str | None = None, ) -> tuple[str | None, str | None]: resolved_key = api_key or os.environ.get(api_key_env) resolved_base_url = base_url or os.environ.get("OPENAI_BASE_URL") config_path = Path(__file__).resolve().parents[2] / "config.py" if config_path.exists(): try: spec = importlib.util.spec_from_file_location("_slot_layout_planner_config", config_path) if spec is not None and spec.loader is not None: module = importlib.util.module_from_spec(spec) spec.loader.exec_module(module) if not resolved_key: for name in ("openai_api_key", "api_key", "client_key"): value = getattr(module, name, None) if isinstance(value, str) and value.strip(): resolved_key = value.strip() break if not resolved_base_url: for name in ("openai_base_url", "base_url"): value = getattr(module, name, None) if isinstance(value, str) and value.strip(): resolved_base_url = value.strip() break except Exception: pass return resolved_key, resolved_base_url def _summarize_data_for_layout(data: dict[str, Any]) -> dict[str, Any]: if not data: return {} summary: dict[str, Any] = {} for key in ("titles", "metadata", "description", "insights", "unit", "source"): value = data.get(key) if value: summary[key] = value rows = data.get("data") or data.get("rows") or data.get("values") if isinstance(rows, list): summary["row_count"] = len(rows) summary["sample_rows"] = rows[:8] elif isinstance(rows, dict): summary["data_keys"] = list(rows.keys())[:20] for key, value in data.items(): if key in summary or key in {"data", "rows", "values"}: continue if len(summary) >= 10: break if isinstance(value, (str, int, float, bool)) or value is None: summary[key] = value elif isinstance(value, list): summary[key] = value[:5] elif isinstance(value, dict): summary[key] = {k: value[k] for k in list(value.keys())[:8]} return summary def _title_hint(data: dict[str, Any]) -> str: titles = data.get("titles") if isinstance(data.get("titles"), dict) else {} metadata = data.get("metadata") if isinstance(data.get("metadata"), dict) else {} for value in ( titles.get("main_title"), titles.get("title"), metadata.get("title"), metadata.get("name"), ): if value: return str(value) return "Generate a concise data-aware headline" def _extract_response_text(response: Any) -> str: output_text = getattr(response, "output_text", None) if isinstance(output_text, str) and output_text.strip(): return output_text.strip() parts: list[str] = [] for item in getattr(response, "output", []) or []: for content in getattr(item, "content", []) or []: text = getattr(content, "text", None) if isinstance(text, str): parts.append(text) elif isinstance(text, dict): value = text.get("value") if isinstance(value, str): parts.append(value) text = "\n".join(parts).strip() if not text: raise RuntimeError("OpenAI layout response did not include text output.") return text def _parse_json_object(text: str) -> dict[str, Any]: cleaned = text.strip() if cleaned.startswith("```"): cleaned = cleaned.strip("`") if cleaned.lower().startswith("json"): cleaned = cleaned[4:].strip() try: data = json.loads(cleaned) except json.JSONDecodeError: start = cleaned.find("{") end = cleaned.rfind("}") if start < 0 or end <= start: raise data = json.loads(cleaned[start : end + 1]) if not isinstance(data, dict): raise RuntimeError("OpenAI layout response must be a JSON object.") return data def _clamp_slot( x: float, y: float, w: float, h: float, canvas_w: int, canvas_h: int, ) -> tuple[int, int, int, int]: x0 = max(0, min(int(round(x)), canvas_w - 1)) y0 = max(0, min(int(round(y)), canvas_h - 1)) ww = max(8, min(int(round(w)), canvas_w - x0)) hh = max(8, min(int(round(h)), canvas_h - y0)) return x0, y0, ww, hh def _intersection_area(a: tuple[int, int, int, int], b: tuple[int, int, int, int]) -> int: ax, ay, aw, ah = a bx, by, bw, bh = b ix0 = max(ax, bx) iy0 = max(ay, by) ix1 = min(ax + aw, bx + bw) iy1 = min(ay + ah, by + bh) return max(0, ix1 - ix0) * max(0, iy1 - iy0) def _choose_overlay_slot( chart_bbox: tuple[int, int, int, int], canvas_w: int, canvas_h: int, protected_regions: list[ProtectedRegion], ) -> tuple[int, int, int, int]: chart_x, chart_y, chart_w, chart_h = chart_bbox candidate_w = chart_w * 0.34 candidate_h = min(190, chart_h * 0.22) candidates = [ (chart_x + chart_w * 0.56, chart_y + chart_h * 0.16), (chart_x + chart_w * 0.08, chart_y + chart_h * 0.16), (chart_x + chart_w * 0.56, chart_y + chart_h * 0.54), (chart_x + chart_w * 0.08, chart_y + chart_h * 0.54), (chart_x + chart_w * 0.33, chart_y + chart_h * 0.34), ] slots = [ _clamp_slot(x, y, candidate_w, candidate_h, canvas_w, canvas_h) for x, y in candidates ] if not protected_regions: return slots[0] return min( slots, key=lambda bbox: ( sum(_intersection_area(bbox, region.bbox_px) for region in protected_regions), abs((bbox[0] + bbox[2] / 2) - (chart_x + chart_w * 0.73)), ), ) def _build_slots( canvas_w: int, canvas_h: int, chart_bbox: tuple[int, int, int, int], data: dict[str, Any], allow_chart_overlap: bool, protected_regions: list[ProtectedRegion] | None = None, ) -> list[SlotPlanItem]: margin = max(56, int(canvas_w * 0.047)) chart_x, chart_y, chart_w, chart_h = chart_bbox title_hint = _title_hint(data) protected_regions = protected_regions or [] def slot( slot_id: str, kind: str, bbox: tuple[int, int, int, int], element: str, anchor: str, overlap: bool = False, policy: str = "image2_generated", requirements: list[str] | None = None, ) -> SlotPlanItem: return SlotPlanItem( id=slot_id, kind=kind, bbox_px=bbox, may_overlap_chart=bool(overlap and allow_chart_overlap), element=element, anchor=anchor, collision_rule=( "may overlap and visually integrate with chart content" if overlap and allow_chart_overlap else "stay visually coherent with the chart without requiring pixel preservation" ), asset_source_policy=policy, slot_requirements=requirements or [], ) overlay = _choose_overlay_slot(chart_bbox, canvas_w, canvas_h, protected_regions) visual = _clamp_slot(canvas_w - margin - 360, margin + 26, 360, 260, canvas_w, canvas_h) title = _clamp_slot(margin, margin, canvas_w - 2 * margin - 250, 210, canvas_w, canvas_h) subtitle = _clamp_slot(margin, margin + 232, canvas_w * 0.60, 100, canvas_w, canvas_h) bottom_y = min(canvas_h - margin - 180, chart_y + chart_h + 42) bottom = _clamp_slot(margin, bottom_y, canvas_w - 2 * margin, 170, canvas_w, canvas_h) return [ slot( "hero_title", "title", title, f"Generate a polished editorial headline. Title hint: {title_hint}", "top-left headline area", policy="llm_generated_text", requirements=[ "text may be rewritten by the image model", "keep the headline legible and data-aware", ], ), slot( "subtitle_context", "subtitle", subtitle, "Generate a short explanatory subtitle or deck based on the chart data.", "under the main headline", policy="llm_generated_text", requirements=["one or two concise lines", "avoid invented numeric values"], ), slot( "topic_visual", "image", visual, "Generate a thematic editorial image or icon cluster that matches the data topic.", "upper-right visual support area", policy="image2_generated", requirements=["no logos", "no watermarks", "avoid small unreadable text"], ), slot( "chart_overlay_callout", "callout", overlay, "Generate a concise visual callout overlay highlighting the most important visible trend.", "inside or near the chart plot area", overlap=True, policy="llm_generated_text_and_shape", requirements=[ "may overlap chart if useful", "do not fabricate exact values unless they are visible in the chart", ], ), slot( "bottom_context_band", "context", bottom, "Generate supporting context, iconography, or a short narrative footer for the infographic.", "below the chart", policy="llm_generated_text_and_image", requirements=["keep text brief", "do not add source logos or watermarks"], ), ] def _safe_slot_id(raw_id: Any, kind: str, index: int) -> str: text = str(raw_id or "").strip().lower() text = re.sub(r"[^a-z0-9]+", "_", text).strip("_") if not text: text = f"{kind}_{index:02d}" if not re.match(r"^[a-z]", text): text = f"slot_{text}" return text[:64] def _slot_kind(raw_kind: Any) -> str: kind = re.sub(r"[^a-z0-9_]+", "_", str(raw_kind or "context").strip().lower()).strip("_") return kind or "context" def _coerce_requirements(value: Any) -> list[str]: if not isinstance(value, list): return [] requirements = [] for item in value[:8]: text = str(item).strip() if text: requirements.append(text[:220]) return requirements def _slot_overlap_area( bbox: tuple[int, int, int, int], regions: list[ProtectedRegion], ) -> int: return sum(_intersection_area(bbox, region.bbox_px) for region in regions) def _nudge_slot_clear_of_regions( bbox: tuple[int, int, int, int], regions: list[ProtectedRegion], canvas_w: int, canvas_h: int, ) -> tuple[int, int, int, int] | None: if not regions or _slot_overlap_area(bbox, regions) == 0: return bbox x, y, w, h = bbox candidates = [bbox] for region in regions: rx, ry, rw, rh = region.bbox_px candidates.extend( [ _clamp_slot(x, ry - h - 12, w, h, canvas_w, canvas_h), _clamp_slot(x, ry + rh + 12, w, h, canvas_w, canvas_h), _clamp_slot(rx - w - 12, y, w, h, canvas_w, canvas_h), _clamp_slot(rx + rw + 12, y, w, h, canvas_w, canvas_h), ] ) clear_candidates = [candidate for candidate in candidates if _slot_overlap_area(candidate, regions) == 0] if clear_candidates: return min( clear_candidates, key=lambda candidate: abs(candidate[0] - x) + abs(candidate[1] - y), ) return None def _slots_from_layout_payload( data: dict[str, Any], *, canvas_w: int, canvas_h: int, chart_bbox: tuple[int, int, int, int], allow_chart_overlap: bool, chart_internal_slots: list[ProtectedRegion], ) -> tuple[list[SlotPlanItem], str, dict[str, Any]]: raw_slots = data.get("slots") if not isinstance(raw_slots, list): raise RuntimeError("GPT layout response is missing a slots array.") slots: list[SlotPlanItem] = [] seen_ids: set[str] = set() validation_notes: list[dict[str, Any]] = [] for index, raw in enumerate(raw_slots, 1): if not isinstance(raw, dict): validation_notes.append({"index": index, "status": "skipped", "reason": "slot is not an object"}) continue raw_bbox = raw.get("bbox_px") if not isinstance(raw_bbox, (list, tuple)) or len(raw_bbox) != 4: validation_notes.append({"index": index, "status": "skipped", "reason": "invalid bbox_px"}) continue try: bbox = _clamp_slot( float(raw_bbox[0]), float(raw_bbox[1]), float(raw_bbox[2]), float(raw_bbox[3]), canvas_w, canvas_h, ) except (TypeError, ValueError): validation_notes.append({"index": index, "status": "skipped", "reason": "non-numeric bbox_px"}) continue if bbox[2] < 40 or bbox[3] < 32: validation_notes.append({"index": index, "status": "skipped", "reason": "bbox too small"}) continue nudged_bbox = _nudge_slot_clear_of_regions(bbox, chart_internal_slots, canvas_w, canvas_h) if nudged_bbox is None: validation_notes.append( { "index": index, "status": "skipped", "reason": "overlaps chart_internal_slots", "bbox_px": list(bbox), } ) continue if nudged_bbox != bbox: validation_notes.append( { "index": index, "status": "adjusted", "reason": "nudged clear of chart_internal_slots", "from_bbox_px": list(bbox), "to_bbox_px": list(nudged_bbox), } ) bbox = nudged_bbox requested_overlap = bool(raw.get("may_overlap_chart")) may_overlap_chart = requested_overlap and allow_chart_overlap if not may_overlap_chart and _intersection_area(bbox, chart_bbox) > 0: validation_notes.append( { "index": index, "status": "skipped", "reason": "overlaps chart while may_overlap_chart is false", "bbox_px": list(bbox), } ) continue kind = _slot_kind(raw.get("kind")) slot_id = _safe_slot_id(raw.get("id"), kind, index) if slot_id in seen_ids: slot_id = f"{slot_id}_{index:02d}" seen_ids.add(slot_id) policy = str(raw.get("asset_source_policy") or "").strip() if policy not in ALLOWED_SLOT_POLICIES: policy = "llm_generated_text_and_image" if kind in {"context", "callout", "annotation"} else "image2_generated" collision_rule = str(raw.get("collision_rule") or "").strip() if not collision_rule: collision_rule = ( "may overlap chart but must preserve all chart labels, marks, icons, and numeric values" if may_overlap_chart else "avoid chart content and chart-internal icon slots" ) slots.append( SlotPlanItem( id=slot_id, kind=kind, bbox_px=bbox, may_overlap_chart=may_overlap_chart, element=str(raw.get("element") or f"Generate {kind} content for this infographic.").strip(), anchor=str(raw.get("anchor") or "").strip(), collision_rule=collision_rule, asset_source_policy=policy, slot_requirements=_coerce_requirements(raw.get("slot_requirements")), ) ) if not slots: raise RuntimeError("GPT layout response did not produce any valid editable slots.") background = str(data.get("background") or "#F7F8F4").strip() if not re.match(r"^#[0-9a-fA-F]{6}$", background): background = "#F7F8F4" meta = { "design_rationale": str(data.get("design_rationale") or "").strip(), "validation_notes": validation_notes, } return slots, background, meta def _build_gpt_slots( *, output_dir: Path, chart_png: Path, canvas_w: int, canvas_h: int, chart_bbox: tuple[int, int, int, int], data: dict[str, Any], allow_chart_overlap: bool, chart_internal_slots: list[ProtectedRegion], model: str, api_key: str | None, api_key_env: str, base_url: str | None, timeout_seconds: float, max_retries: int, max_output_tokens: int, ) -> tuple[list[SlotPlanItem], str, dict[str, Any]]: request_path = output_dir / "layout_request.json" response_path = output_dir / "layout_response.json" payload = { "canvas": {"width": canvas_w, "height": canvas_h}, "chart_bbox_px": list(chart_bbox), "allow_chart_overlap": allow_chart_overlap, "title_hint": _title_hint(data), "data_summary": _summarize_data_for_layout(data), "chart_internal_slots": [ { "id": region.id, "source_slot": region.source_slot, "bbox_px": list(region.bbox_px), "anchor": region.anchor, "collision_rule": region.collision_rule, } for region in chart_internal_slots ], "required_output_shape": { "background": "#RRGGBB", "design_rationale": "short string", "slots": [ { "id": "short_snake_case_id", "kind": "title|subtitle|image|callout|context|annotation|badge", "bbox_px": [0, 0, 100, 100], "may_overlap_chart": False, "element": "generation instruction", "anchor": "layout anchor", "collision_rule": "preservation rule", "asset_source_policy": "llm_generated_text|image2_generated|llm_generated_text_and_shape|llm_generated_text_and_image", "slot_requirements": ["requirement"], } ], }, } request_log = { "agent": "slot_layout_planner", "model": model, "instructions": GPT_LAYOUT_INSTRUCTIONS, "payload": payload, "images": [{"label": "sanitized_chart", "path": str(chart_png)}], } request_path.write_text(json.dumps(request_log, indent=2, ensure_ascii=False), encoding="utf-8") resolved_key, resolved_base_url = _load_openai_config( api_key=api_key, api_key_env=api_key_env, base_url=base_url, ) if not resolved_key: raise RuntimeError( f"{api_key_env} is not set and config.py does not define api_key/client_key; GPT layout cannot run." ) try: from openai import OpenAI except ImportError as exc: raise RuntimeError("The openai package is required for GPT layout planning.") from exc client_kwargs: dict[str, Any] = { "api_key": resolved_key, "timeout": timeout_seconds, "max_retries": max_retries, } if resolved_base_url: client_kwargs["base_url"] = resolved_base_url client = OpenAI(**client_kwargs) content: list[dict[str, Any]] = [ { "type": "input_text", "text": ( "Return valid JSON only. No markdown fences.\n\n" f"Payload:\n{json.dumps(payload, ensure_ascii=False, indent=2)}" ), }, {"type": "input_text", "text": "Image: sanitized_chart"}, {"type": "input_image", "image_url": _image_data_uri(chart_png)}, ] kwargs: dict[str, Any] = { "model": model, "instructions": GPT_LAYOUT_INSTRUCTIONS, "input": [{"role": "user", "content": content}], "max_output_tokens": max_output_tokens, "text": {"format": {"type": "json_object"}}, } response = client.responses.create(**kwargs) response_meta = { "model": getattr(response, "model", model), "response_id": getattr(response, "id", None), } text = _extract_response_text(response) response_log: dict[str, Any] = { "metadata": response_meta, "output_text": text, } try: parsed = _parse_json_object(text) slots, background, validation_meta = _slots_from_layout_payload( parsed, canvas_w=canvas_w, canvas_h=canvas_h, chart_bbox=chart_bbox, allow_chart_overlap=allow_chart_overlap, chart_internal_slots=chart_internal_slots, ) except Exception as exc: response_log["parse_or_validation_error"] = f"{type(exc).__name__}: {exc}" response_path.write_text(json.dumps(response_log, indent=2, ensure_ascii=False), encoding="utf-8") raise response_log["parsed"] = parsed response_log["validated"] = { "background": background, "reserved_slots": [asdict(slot) for slot in slots], **validation_meta, } response_path.write_text(json.dumps(response_log, indent=2, ensure_ascii=False), encoding="utf-8") meta = { "engine": "gpt", "model": response_meta["model"], "request_path": str(request_path), "response_path": str(response_path), **validation_meta, } return slots, background, meta def _image_data_uri(path: Path) -> str: payload = base64.b64encode(path.read_bytes()).decode("ascii") return f"data:image/png;base64,{payload}" def _slot_svg(slot: SlotPlanItem) -> str: x, y, w, h = slot.bbox_px cls = f"planned-slot planned-{slot.kind}-slot" marker = html.escape(slot.id, quote=True) element = html.escape(slot.element, quote=True) anchor = html.escape(slot.anchor, quote=True) collision = html.escape(slot.collision_rule, quote=True) policy = html.escape(slot.asset_source_policy, quote=True) return ( f'' f'' "" ) def _write_planned_svg( output_svg: Path, chart_png: Path, canvas_size: tuple[int, int], chart_bbox: tuple[int, int, int, int], slots: list[SlotPlanItem], background: str, ) -> None: canvas_w, canvas_h = canvas_size chart_x, chart_y, chart_w, chart_h = chart_bbox chart_href = _image_data_uri(chart_png) slot_markup = "\n".join(_slot_svg(slot) for slot in slots) svg = f''' {slot_markup} ''' output_svg.parent.mkdir(parents=True, exist_ok=True) output_svg.write_text(svg, encoding="utf-8") def _render_planned_png( output_png: Path, chart_png: Path, canvas_size: tuple[int, int], chart_bbox: tuple[int, int, int, int], slots: list[SlotPlanItem], background_rgb: tuple[int, int, int] = (247, 248, 244), ) -> None: canvas_w, canvas_h = canvas_size output_png.parent.mkdir(parents=True, exist_ok=True) canvas = Image.new("RGB", (canvas_w, canvas_h), background_rgb) with Image.open(chart_png) as chart: chart = chart.convert("RGBA") x, y, w, h = chart_bbox fitted = chart.copy() fitted.thumbnail((w, h), Image.Resampling.LANCZOS) paste_x = x + (w - fitted.width) // 2 paste_y = y + (h - fitted.height) // 2 canvas.paste(fitted.convert("RGB"), (paste_x, paste_y), fitted.getchannel("A")) overlay = Image.new("RGBA", (canvas_w, canvas_h), (0, 0, 0, 0)) draw = ImageDraw.Draw(overlay) colors = [ (93, 117, 150, 72), (35, 150, 120, 68), (208, 132, 48, 70), (135, 88, 178, 70), (60, 128, 190, 68), ] for index, slot in enumerate(slots): x, y, w, h = slot.bbox_px fill = colors[index % len(colors)] outline = (54, 66, 82, 155) draw.rounded_rectangle((x, y, x + w, y + h), radius=10, fill=fill, outline=outline, width=2) canvas = Image.alpha_composite(canvas.convert("RGBA"), overlay).convert("RGB") canvas.save(output_png) def _write_reference_map( path: Path, slots: list[SlotPlanItem], chart_internal_slots: list[ProtectedRegion], ) -> None: path.parent.mkdir(parents=True, exist_ok=True) payload = { "reserved_slots": [], "chart_internal_slots": [ { "id": region.id, "reason": region.reason, "bbox_px": list(region.bbox_px), "source_slot": region.source_slot, "anchor": region.anchor, "collision_rule": region.collision_rule, "semantic": region.semantic, } for region in chart_internal_slots ], "protected_regions": [], } for slot in slots: payload["reserved_slots"].append( { "id": slot.id, "element": slot.element, "reason": f"planned {slot.kind} slot for image-edit generation", "slot_requirements": slot.slot_requirements, "semantic_slot": { "class_or_data_marker": f'data-asset-slot="{slot.id}"', "anchor": slot.anchor, "collision_rule": slot.collision_rule, "asset_source_policy": slot.asset_source_policy, "may_overlap_chart": slot.may_overlap_chart, }, } ) path.write_text(json.dumps(payload, indent=2, ensure_ascii=False), encoding="utf-8") def _parse_bbox_text(text: str | None) -> tuple[float, float, float, float] | None: if not text: return None try: parts = [float(part) for part in str(text).replace(",", " ").split() if part] except ValueError: return None if len(parts) != 4: return None x, y, w, h = parts if w <= 0 or h <= 0: return None return x, y, w, h def _bbox_with_matrix( bbox: tuple[float, float, float, float], matrix: tuple[float, float, float, float, float, float], ) -> tuple[float, float, float, float]: x, y, w, h = bbox a, b, c, d, e, f = matrix points = [ (a * x + c * y + e, b * x + d * y + f), (a * (x + w) + c * y + e, b * (x + w) + d * y + f), (a * x + c * (y + h) + e, b * x + d * (y + h) + f), (a * (x + w) + c * (y + h) + e, b * (x + w) + d * (y + h) + f), ] xs = [point[0] for point in points] ys = [point[1] for point in points] return min(xs), min(ys), max(xs) - min(xs), max(ys) - min(ys) def _float_attr(elem: Any, name: str, default: float | None = None) -> float | None: value = elem.get(name) if value is None: return default try: return float(str(value).strip()) except ValueError: return default def _element_local_bbox(elem: Any) -> tuple[float, float, float, float] | None: tag = etree_local_name(elem) if tag in {"rect", "image"}: x = _float_attr(elem, "x", 0.0) y = _float_attr(elem, "y", 0.0) w = _float_attr(elem, "width") h = _float_attr(elem, "height") if x is not None and y is not None and w and h and w > 0 and h > 0: return x, y, w, h if tag == "circle": cx = _float_attr(elem, "cx", 0.0) cy = _float_attr(elem, "cy", 0.0) r = _float_attr(elem, "r") if cx is not None and cy is not None and r and r > 0: return cx - r, cy - r, r * 2, r * 2 if tag == "ellipse": cx = _float_attr(elem, "cx", 0.0) cy = _float_attr(elem, "cy", 0.0) rx = _float_attr(elem, "rx") ry = _float_attr(elem, "ry") if cx is not None and cy is not None and rx and ry and rx > 0 and ry > 0: return cx - rx, cy - ry, rx * 2, ry * 2 return None def etree_local_name(elem: Any) -> str: try: return elem.tag.rsplit("}", 1)[-1] if isinstance(elem.tag, str) else "" except Exception: return "" def _collect_protected_regions( chart_svg: Path, source_chart_png: Path, chart_bbox: tuple[int, int, int, int], padding_px: int = 10, ) -> list[ProtectedRegion]: try: tree, root, svg_width, svg_height = _load_svg(chart_svg) except Exception: return [] del tree with Image.open(source_chart_png) as chart_image: png_w, png_h = chart_image.size chart_x, chart_y, chart_w, chart_h = chart_bbox svg_to_png_x = png_w / svg_width svg_to_png_y = png_h / svg_height png_to_canvas_x = chart_w / png_w png_to_canvas_y = chart_h / png_h nodes = root.xpath( ".//*[@data-asset-slot or @data-reserved-slot or @data-slot-id or @data-slot-kind]" ) regions: list[ProtectedRegion] = [] for index, elem in enumerate(nodes, 1): if not is_inside_chart_internal_slot(elem): continue raw_bbox = ( _parse_bbox_text(elem.get("data-bbox")) or _parse_bbox_text(elem.get("data-asset-bbox")) or _parse_bbox_text(elem.get("data-slot-bbox")) ) if raw_bbox: bbox_svg = _bbox_with_matrix(raw_bbox, _element_transform(elem, root)) else: local_bbox = _element_local_bbox(elem) if local_bbox is not None: bbox_svg = _bbox_with_matrix(local_bbox, _element_transform(elem, root)) else: abs_bbox = _absolute_bbox(elem, root) if abs_bbox is None: continue bbox_svg = (abs_bbox.x, abs_bbox.y, abs_bbox.width, abs_bbox.height) if bbox_svg is None: continue x, y, w, h = bbox_svg if w <= 0 or h <= 0: continue px0 = chart_x + int(round(x * svg_to_png_x * png_to_canvas_x)) - padding_px py0 = chart_y + int(round(y * svg_to_png_y * png_to_canvas_y)) - padding_px px1 = chart_x + int(round((x + w) * svg_to_png_x * png_to_canvas_x)) + padding_px py1 = chart_y + int(round((y + h) * svg_to_png_y * png_to_canvas_y)) + padding_px px0 = max(0, px0) py0 = max(0, py0) px1 = min(chart_x + chart_w, px1) py1 = min(chart_y + chart_h, py1) if px1 <= px0 or py1 <= py0: continue slot_name = ( elem.get("data-asset-slot") or elem.get("data-reserved-slot") or elem.get("data-slot-id") or elem.get("data-slot-kind") or f"chart_internal_slot_{index}" ) anchor = elem.get("data-asset-anchor") or elem.get("data-anchor") or elem.get("data-slot-anchor") or "" collision_rule = elem.get("data-asset-collision") or elem.get("data-collision-rule") or "" regions.append( ProtectedRegion( id=f"CHART_ICON_{len(regions) + 1:02d}", reason="chart-internal icon/asset slot", bbox_px=(px0, py0, px1 - px0, py1 - py0), source_slot=str(slot_name), anchor=str(anchor), collision_rule=str(collision_rule), semantic=( "Generate or polish only the icon inside this chart slot; " "preserve nearby chart labels, ranks, lines, points, and numeric values exactly." ), ) ) return regions def build_planned_slot_package( chart_svg: Path, output_dir: Path, *, chart_png: Path | None = None, data_json: Path | None = None, canvas_width: int = 1536, canvas_height: int = 2048, allow_chart_overlap: bool = True, render_longest_side: int | None = None, layout_model: str | None = "gpt-5.5", layout_api_key: str | None = None, layout_api_key_env: str = "OPENAI_API_KEY", layout_base_url: str | None = None, layout_timeout_seconds: float = 120.0, layout_max_retries: int = 0, layout_max_output_tokens: int = 2500, deterministic_layout: bool = False, ) -> PlannedSlotPackage: output_dir.mkdir(parents=True, exist_ok=True) sanitized_svg = output_dir / "sanitized_chart.svg" sanitized_png = output_dir / "sanitized_chart.png" sanitizer_report = output_dir / "sanitized_chart.report.json" sanitize_chart_svg(chart_svg, sanitized_svg, report_path=sanitizer_report) if render_longest_side is not None: import os previous = os.environ.get("RENDER_LONGEST_SIDE") os.environ["RENDER_LONGEST_SIDE"] = str(render_longest_side) try: render_svg_to_png(sanitized_svg, sanitized_png) finally: if previous is None: os.environ.pop("RENDER_LONGEST_SIDE", None) else: os.environ["RENDER_LONGEST_SIDE"] = previous else: render_svg_to_png(sanitized_svg, sanitized_png) source_chart_png = sanitized_png if sanitized_png.is_file() else chart_png if source_chart_png is None or not source_chart_png.is_file(): raise FileNotFoundError("planned slot package requires a chart PNG") with Image.open(source_chart_png) as chart_image: chart_w, chart_h = chart_image.size max_w = int(canvas_width * 0.84) max_h = int(canvas_height * 0.58) scale = min(max_w / chart_w, max_h / chart_h, 1.0) placed_w = max(1, int(chart_w * scale)) placed_h = max(1, int(chart_h * scale)) chart_x = (canvas_width - placed_w) // 2 chart_y = int(canvas_height * 0.30) chart_bbox = (chart_x, chart_y, placed_w, placed_h) chart_internal_slots = _collect_protected_regions(chart_svg, source_chart_png, chart_bbox) data = _read_json(data_json) background = "#F7F8F4" if deterministic_layout or not layout_model: slots = _build_slots( canvas_width, canvas_height, chart_bbox, data, allow_chart_overlap, chart_internal_slots, ) layout_engine: dict[str, Any] = { "engine": "deterministic", "model": None, "reason": "deterministic_layout option enabled" if deterministic_layout else "layout_model is empty", } else: slots, background, layout_engine = _build_gpt_slots( output_dir=output_dir, chart_png=source_chart_png, canvas_w=canvas_width, canvas_h=canvas_height, chart_bbox=chart_bbox, data=data, allow_chart_overlap=allow_chart_overlap, chart_internal_slots=chart_internal_slots, model=layout_model, api_key=layout_api_key, api_key_env=layout_api_key_env, base_url=layout_base_url, timeout_seconds=layout_timeout_seconds, max_retries=layout_max_retries, max_output_tokens=layout_max_output_tokens, ) slot_plan_path = output_dir / "slot_plan.json" planned_svg = output_dir / "planned_slots.svg" planned_png = output_dir / "planned_slots.png" reference_map = output_dir / "reference_element_map.json" slot_plan = { "canvas": {"width": canvas_width, "height": canvas_height, "background": background}, "chart": { "sanitized_svg": str(sanitized_svg), "sanitized_png": str(source_chart_png), "bbox_px": list(chart_bbox), "overlap_policy": "editable_overlap_allowed" if allow_chart_overlap else "avoid_chart_overlap", }, "layout_engine": layout_engine, "text_policy": "llm_generated", "chart_internal_slots": [asdict(region) for region in chart_internal_slots], "protected_regions": [], "reserved_slots": [asdict(slot) for slot in slots], } slot_plan_path.write_text(json.dumps(slot_plan, indent=2, ensure_ascii=False), encoding="utf-8") _write_reference_map(reference_map, slots, chart_internal_slots) _write_planned_svg(planned_svg, source_chart_png, (canvas_width, canvas_height), chart_bbox, slots, background) _render_planned_png(planned_png, source_chart_png, (canvas_width, canvas_height), chart_bbox, slots) return PlannedSlotPackage( slot_plan=slot_plan_path, planned_svg=planned_svg, planned_png=planned_png, reference_map=reference_map, sanitized_svg=sanitized_svg, sanitized_png=sanitized_png, sanitizer_report=sanitizer_report, ) def parse_args(argv: Iterable[str] | None = None) -> argparse.Namespace: parser = argparse.ArgumentParser(description="Build a planned slot canvas from a chart-only SVG.") parser.add_argument("--chart-svg", type=Path, required=True) parser.add_argument("--chart-png", type=Path, default=None) parser.add_argument("--data-json", type=Path, default=None) parser.add_argument("--output-dir", type=Path, required=True) parser.add_argument("--canvas-width", type=int, default=1536) parser.add_argument("--canvas-height", type=int, default=2048) parser.add_argument("--disallow-chart-overlap", action="store_true") parser.add_argument("--render-longest-side", type=int, default=None) parser.add_argument("--layout-model", default="gpt-5.5") parser.add_argument("--layout-api-key-env", default="OPENAI_API_KEY") parser.add_argument("--layout-base-url", default=None) parser.add_argument("--layout-timeout", type=float, default=120.0) parser.add_argument("--layout-max-retries", type=int, default=0) parser.add_argument("--layout-max-output-tokens", type=int, default=2500) parser.add_argument( "--deterministic-layout", action="store_true", help="Use the old deterministic slot planner instead of GPT layout design.", ) return parser.parse_args(argv) def main(argv: Iterable[str] | None = None) -> int: args = parse_args(argv) package = build_planned_slot_package( chart_svg=args.chart_svg, chart_png=args.chart_png, data_json=args.data_json, output_dir=args.output_dir, canvas_width=args.canvas_width, canvas_height=args.canvas_height, allow_chart_overlap=not args.disallow_chart_overlap, render_longest_side=args.render_longest_side, layout_model=args.layout_model, layout_api_key_env=args.layout_api_key_env, layout_base_url=args.layout_base_url, layout_timeout_seconds=args.layout_timeout, layout_max_retries=args.layout_max_retries, layout_max_output_tokens=args.layout_max_output_tokens, deterministic_layout=args.deterministic_layout, ) print(json.dumps({k: str(v) for k, v in asdict(package).items()}, ensure_ascii=False)) return 0 if __name__ == "__main__": raise SystemExit(main())