File size: 3,727 Bytes
f927995
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
"""
Builds the MVP import graph:
- 20 pharma API nodes
- 8 rare earth mineral nodes
- edges weighted by India's monthly import volume (USD millions)
"""

from __future__ import annotations

import json
from pathlib import Path
from typing import Any

import networkx as nx

BASE_DIR = Path(__file__).resolve().parent.parent
APIS_FILE = BASE_DIR / "data" / "seed" / "apis.json"
RARE_EARTHS_FILE = BASE_DIR / "data" / "seed" / "rare_earths.json"
OUTPUT_FILE = BASE_DIR / "data" / "graph_imports.json"

INDIA_NODE_ID = "india"
PHARMA_NODE_LIMIT = 20
RARE_EARTH_NODE_LIMIT = 8


def _read_json(path: Path) -> list[dict[str, Any]]:
    with open(path, "r", encoding="utf-8") as file:
        return json.load(file)


def _top_by_import_value(records: list[dict[str, Any]], limit: int) -> list[dict[str, Any]]:
    sorted_rows = sorted(
        records,
        key=lambda row: float(row.get("monthly_import_value_usd_millions", 0) or 0),
        reverse=True,
    )
    return sorted_rows[:limit]


def build_import_graph() -> nx.DiGraph:
    apis = _top_by_import_value(_read_json(APIS_FILE), PHARMA_NODE_LIMIT)
    rare_earths = _top_by_import_value(_read_json(RARE_EARTHS_FILE), RARE_EARTH_NODE_LIMIT)

    graph = nx.DiGraph(name="pharmashield_mvp_import_graph")
    graph.add_node(INDIA_NODE_ID, type="country", name="India")

    max_value = max(
        [float(item.get("monthly_import_value_usd_millions", 0) or 0) for item in apis + rare_earths] or [1.0]
    )

    for api in apis:
        node_id = api["id"]
        value = float(api.get("monthly_import_value_usd_millions", 0) or 0)
        graph.add_node(
            node_id,
            type="pharma_api",
            name=api.get("name", node_id),
            sector="pharma",
            china_share=api.get("china_share"),
            primary_provinces=api.get("primary_provinces", []),
            monthly_import_value_usd_millions=value,
        )
        graph.add_edge(
            node_id,
            INDIA_NODE_ID,
            edge_type="india_imports_api",
            import_volume_usd_millions=value,
            weight=round(value / max_value, 6),
        )

    for mineral in rare_earths:
        node_id = mineral["id"]
        value = float(mineral.get("monthly_import_value_usd_millions", 0) or 0)
        graph.add_node(
            node_id,
            type="rare_earth",
            name=mineral.get("name", node_id),
            sector="rare_earth",
            china_share=mineral.get("china_share"),
            primary_provinces=mineral.get("primary_provinces", []),
            monthly_import_value_usd_millions=value,
        )
        graph.add_edge(
            node_id,
            INDIA_NODE_ID,
            edge_type="india_imports_rare_earth",
            import_volume_usd_millions=value,
            weight=round(value / max_value, 6),
        )

    return graph


def graph_to_dict(graph: nx.DiGraph) -> dict[str, Any]:
    return {
        "nodes": [
            {"id": node_id, **attrs}
            for node_id, attrs in graph.nodes(data=True)
        ],
        "edges": [
            {"source": source, "target": target, **attrs}
            for source, target, attrs in graph.edges(data=True)
        ],
    }


def save_graph(path: Path = OUTPUT_FILE) -> dict[str, Any]:
    graph = build_import_graph()
    payload = graph_to_dict(graph)
    path.parent.mkdir(parents=True, exist_ok=True)
    with open(path, "w", encoding="utf-8") as file:
        json.dump(payload, file, indent=2, ensure_ascii=False)
    return payload


if __name__ == "__main__":
    result = save_graph()
    print(
        f"Built import graph with {len(result['nodes'])} nodes "
        f"and {len(result['edges'])} edges -> {OUTPUT_FILE}"
    )