File size: 5,062 Bytes
c496840
f927995
c496840
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
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
"""
Endpoints for supply chain graph visualization and geospatial risk aggregation.
"""

from datetime import datetime, timedelta
from typing import List, Dict, Annotated, Any
from fastapi import APIRouter, Depends

from ..deps import get_graph_service, get_gnn, get_data_loader
from ..services.graph_service import GraphService
from ..services.shock_propagation import ShockPropagator
from ..services.criticality import compute_score
from ..models.graph import GraphResponse, StateRiskAggregate, GraphNode, GraphEdge


router = APIRouter(prefix="/api/v1", tags=["graph"])

# In-memory cache for graph data
_graph_cache: Dict[str, Any] = {
    "data": None,
    "expiry": datetime.min
}

# Mapping of Indian states to their representative drug subsets for risk aggregation
STATE_DRUG_MAP = {
    "MH": ("Maharashtra", ["paracetamol", "amoxicillin", "metformin", "atorvastatin"]),
    "DL": ("Delhi", ["azithromycin", "ceftriaxone", "amlodipine"]),
    "KA": ("Karnataka", ["losartan", "omeprazole", "ranitidine", "ibuprofen"]),
    "TN": ("Tamil Nadu", ["diclofenac", "aspirin", "penicillin_g"]),
    "WB": ("West Bengal", ["gentamicin", "fluconazole", "ciprofloxacin"]),
    "GJ": ("Gujarat", ["levofloxacin", "insulin", "salbutamol"])
}


def _get_current_graph_data(
    graph_service: GraphService,
    gnn: ShockPropagator,
    data_loader: Any
) -> GraphResponse:
    """Computes risks and aggregates state-level data."""
    # 1. Get raw graph structure
    graph_dict = graph_service.to_serializable_dict()
    
    # 2. Compute current risk for every drug
    risk_map: Dict[str, float] = {}
    drugs = data_loader.get_drugs()
    
    use_gnn = hasattr(gnn, "is_available") and gnn.is_available()

    if use_gnn:
        risk_map = gnn.compute_current_risk()
    else:
        for drug in drugs:
            hhi = graph_service.compute_concentration_hhi(drug.id)
            risk_map[drug.id] = compute_score(drug, hhi)
            
    # 3. Apply computed risks back to the graph service nodes
    graph_service.apply_risk_scores(risk_map)
    
    # Refresh nodes list after applying risks
    updated_graph = graph_service.to_serializable_dict()
    
    # 4. Compute state_risk_aggregates
    state_aggregates = []
    for state_id, (state_name, drug_ids) in STATE_DRUG_MAP.items():
        state_risks = [risk_map.get(d_id, 0.0) for d_id in drug_ids if d_id in risk_map]
        
        # Calculate mean risk for the state
        mean_risk = sum(state_risks) / len(state_risks) if state_risks else 0.0
        
        # Identify top 3 at-risk drugs for this state
        # Sort state drugs by risk descending
        sorted_drugs = sorted(
            [(d_id, risk_map.get(d_id, 0.0)) for d_id in drug_ids],
            key=lambda x: x[1],
            reverse=True
        )
        top_drugs = [d[0] for d in sorted_drugs[:3]]
        
        state_aggregates.append(StateRiskAggregate(
            state_id=state_id,
            state_name=state_name,
            risk_score=float(mean_risk),
            top_at_risk_drugs=top_drugs
        ))
        
    return GraphResponse(
        nodes=[GraphNode(**n) for n in updated_graph["nodes"]],
        edges=[GraphEdge(**e) for e in updated_graph["edges"]],
        state_risk_aggregates=state_aggregates,
        generated_at=datetime.utcnow()
    )


@router.get("/graph", response_model=GraphResponse)
async def get_full_graph(
    graph_service: Annotated[GraphService, Depends(get_graph_service)],
    gnn: Annotated[ShockPropagator, Depends(get_gnn)],
    data_loader: Annotated[Any, Depends(get_data_loader)]
) -> GraphResponse:
    """
    Returns the complete supply chain graph with computed risks and state-level aggregates.
    
    The result is cached for 1 hour to optimize performance.
    """
    global _graph_cache
    now = datetime.utcnow()
    
    if _graph_cache["data"] and _graph_cache["expiry"] > now:
        return _graph_cache["data"]
        
    response = _get_current_graph_data(graph_service, gnn, data_loader)
    
    _graph_cache["data"] = response
    _graph_cache["expiry"] = now + timedelta(hours=1)
    
    return response


@router.get("/graph/states", response_model=List[StateRiskAggregate])
async def get_state_risks(
    graph_service: Annotated[GraphService, Depends(get_graph_service)],
    gnn: Annotated[ShockPropagator, Depends(get_gnn)],
    data_loader: Annotated[Any, Depends(get_data_loader)]
) -> List[StateRiskAggregate]:
    """
    Returns just the geospatial risk aggregates for India.
    
    Useful for lightweight map visualizations.
    """
    # Reuse cached data if available
    global _graph_cache
    now = datetime.utcnow()
    
    if _graph_cache["data"] and _graph_cache["expiry"] > now:
        return _graph_cache["data"].state_risk_aggregates
        
    full_graph = _get_current_graph_data(graph_service, gnn, data_loader)
    
    # Update cache while we are at it
    _graph_cache["data"] = full_graph
    _graph_cache["expiry"] = now + timedelta(hours=1)
    
    return full_graph.state_risk_aggregates