Skip to content

Commit 821feaf

Browse files
committed
feat(uwm): add Gemma4 scenario parsing and multistage planning
1 parent 1421e00 commit 821feaf

45 files changed

Lines changed: 8798 additions & 256 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.
Lines changed: 133 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,133 @@
1+
"""API routes for demand-7 UWM livability intervention planning."""
2+
3+
from __future__ import annotations
4+
5+
import asyncio
6+
import os
7+
from pathlib import Path
8+
9+
from starlette.requests import Request
10+
from starlette.responses import JSONResponse
11+
from starlette.routing import Route
12+
13+
from .helpers import _get_user_from_request, _set_user_context
14+
from ..uwm.livability_demand7.service import Demand7ProductInvalid, Demand7Service
15+
16+
17+
ROOT = Path(__file__).resolve().parents[2]
18+
DATA_ROOT = ROOT / "data/uwm_public_proxy/chongqing_central"
19+
DEFAULT_PANEL = DATA_ROOT / "admin_livability_target_full_admin_graph_2024_07_2026_07_08/uwm_admin_livability_target_full_admin_graph_panel.json"
20+
DEFAULT_PLANNER = DATA_ROOT / "data_calibrated_planner_replay_full_admin_graph_2026_07_08/uwm_full_admin_graph_model_based_graph_search.json"
21+
DEFAULT_GEOMETRY = DATA_ROOT / "admin_units/chongqing_township_admin_units.geojson"
22+
_SERVICE_CACHE: tuple[Path, Path, Path, Demand7Service] | None = None
23+
24+
25+
def _path(name: str, default: Path) -> Path:
26+
configured = os.environ.get(name, "").strip()
27+
return Path(configured).expanduser() if configured else default
28+
29+
30+
def _service() -> Demand7Service:
31+
global _SERVICE_CACHE
32+
paths = (
33+
_path("UWM_LIVABILITY_DEMAND7_PANEL_PATH", DEFAULT_PANEL),
34+
_path("UWM_LIVABILITY_DEMAND7_PLANNER_PATH", DEFAULT_PLANNER),
35+
_path("UWM_LIVABILITY_DEMAND7_GEOMETRY_PATH", DEFAULT_GEOMETRY),
36+
)
37+
if _SERVICE_CACHE is None or _SERVICE_CACHE[:3] != paths:
38+
_SERVICE_CACHE = (*paths, Demand7Service(*paths))
39+
return _SERVICE_CACHE[3]
40+
41+
42+
def _reset_service_cache() -> None:
43+
global _SERVICE_CACHE
44+
_SERVICE_CACHE = None
45+
46+
47+
def _authorized(request: Request) -> JSONResponse | None:
48+
user = _get_user_from_request(request)
49+
if not user:
50+
return JSONResponse({"error": "Unauthorized"}, status_code=401)
51+
_set_user_context(user)
52+
return None
53+
54+
55+
def _unavailable(error: Exception) -> JSONResponse:
56+
return JSONResponse(
57+
{
58+
"schema": "uwm.livability.demand7.unavailable.v1",
59+
"ready": False,
60+
"blockers": [str(error)],
61+
"claim_boundary": "fail_closed",
62+
},
63+
status_code=503,
64+
)
65+
66+
67+
async def demand7_overview(request: Request):
68+
if unauthorized := _authorized(request):
69+
return unauthorized
70+
try:
71+
return JSONResponse(await asyncio.to_thread(_service().overview))
72+
except Demand7ProductInvalid as error:
73+
return _unavailable(error)
74+
75+
76+
async def demand7_units(request: Request):
77+
if unauthorized := _authorized(request):
78+
return unauthorized
79+
try:
80+
query = request.query_params
81+
result = await asyncio.to_thread(
82+
_service().list_units,
83+
query.get("search", ""),
84+
query.get("county", ""),
85+
int(query.get("limit", "100")),
86+
)
87+
return JSONResponse(result)
88+
except (Demand7ProductInvalid, ValueError) as error:
89+
return _unavailable(error) if isinstance(error, Demand7ProductInvalid) else JSONResponse({"error": str(error)}, status_code=400)
90+
91+
92+
async def demand7_unit(request: Request):
93+
if unauthorized := _authorized(request):
94+
return unauthorized
95+
try:
96+
return JSONResponse(await asyncio.to_thread(_service().unit_detail, str(request.path_params.get("unit_id") or "")))
97+
except Demand7ProductInvalid as error:
98+
return _unavailable(error)
99+
except ValueError as error:
100+
return JSONResponse({"error": str(error)}, status_code=404)
101+
102+
103+
async def demand7_plan(request: Request):
104+
if unauthorized := _authorized(request):
105+
return unauthorized
106+
try:
107+
payload = await request.json()
108+
except Exception:
109+
return JSONResponse({"error": "Invalid JSON payload"}, status_code=400)
110+
if not isinstance(payload, dict):
111+
return JSONResponse({"error": "Request object required"}, status_code=400)
112+
try:
113+
result = await asyncio.to_thread(
114+
_service().plan,
115+
str(payload.get("unit_id") or ""),
116+
str(payload.get("target_profile") or "balanced"),
117+
str(payload.get("horizon") or "simulator_step"),
118+
)
119+
return JSONResponse(result)
120+
except Demand7ProductInvalid as error:
121+
return _unavailable(error)
122+
except ValueError as error:
123+
status = 404 if str(error) == "unit_not_found" else 400
124+
return JSONResponse({"error": str(error)}, status_code=status)
125+
126+
127+
def get_uwm_livability_demand7_routes() -> list:
128+
return [
129+
Route("/api/uwm/livability/demand7/overview", demand7_overview, methods=["GET"]),
130+
Route("/api/uwm/livability/demand7/units", demand7_units, methods=["GET"]),
131+
Route("/api/uwm/livability/demand7/units/{unit_id}", demand7_unit, methods=["GET"]),
132+
Route("/api/uwm/livability/demand7/plan", demand7_plan, methods=["POST"]),
133+
]

data_agent/api/uwm_livability_s2_routes.py

Lines changed: 45 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@
1414
from .helpers import _get_user_from_request, _set_user_context
1515
from ..uwm.livability_s2.scenario_service import (
1616
S2ProductInvalid,
17+
S2RunInvalid,
1718
S2RunNotFound,
1819
S2ScenarioService,
1920
)
@@ -23,7 +24,7 @@
2324
DEFAULT_PRODUCT_DIR = (
2425
ROOT / "data/uwm_public_proxy/chongqing_central/uwm_livability_s2_fulu"
2526
)
26-
_SERVICE_CACHE: tuple[Path, S2ScenarioService] | None = None
27+
_SERVICE_CACHE: tuple[Path, Path | None, S2ScenarioService] | None = None
2728

2829

2930
def _product_dir() -> Path:
@@ -34,9 +35,11 @@ def _product_dir() -> Path:
3435
def _service() -> S2ScenarioService:
3536
global _SERVICE_CACHE
3637
path = _product_dir()
37-
if _SERVICE_CACHE is None or _SERVICE_CACHE[0] != path:
38-
_SERVICE_CACHE = (path, S2ScenarioService(path))
39-
return _SERVICE_CACHE[1]
38+
configured_store = os.environ.get("UWM_LIVABILITY_S2_RUN_STORE", "").strip()
39+
store_path = Path(configured_store).expanduser() if configured_store else None
40+
if _SERVICE_CACHE is None or _SERVICE_CACHE[:2] != (path, store_path):
41+
_SERVICE_CACHE = (path, store_path, S2ScenarioService(path, store_path))
42+
return _SERVICE_CACHE[2]
4043

4144

4245
def _reset_service_cache() -> None:
@@ -96,6 +99,26 @@ async def uwm_livability_s2_parcels(request: Request):
9699
return _product_error(error)
97100

98101

102+
async def uwm_livability_s2_facilities(request: Request):
103+
_, unauthorized = _authorized(request)
104+
if unauthorized:
105+
return unauthorized
106+
try:
107+
return JSONResponse(await asyncio.to_thread(_service().list_facilities))
108+
except S2ProductInvalid as error:
109+
return _product_error(error)
110+
111+
112+
async def uwm_livability_s2_planning_projects(request: Request):
113+
_, unauthorized = _authorized(request)
114+
if unauthorized:
115+
return unauthorized
116+
try:
117+
return JSONResponse(await asyncio.to_thread(_service().list_planning_projects))
118+
except S2ProductInvalid as error:
119+
return _product_error(error)
120+
121+
99122
async def uwm_livability_s2_parcel(request: Request):
100123
_, unauthorized = _authorized(request)
101124
if unauthorized:
@@ -128,6 +151,13 @@ async def uwm_livability_s2_validate_action(request: Request):
128151
rationale=str(payload.get("rationale") or ""),
129152
requested_at=str(payload.get("requested_at") or ""),
130153
actor_id=str(username),
154+
action_type=str(payload.get("action_type") or "change_land_use"),
155+
facility_class=payload.get("facility_class"),
156+
facility_id=payload.get("facility_id"),
157+
service_radius_m=payload.get("service_radius_m"),
158+
radius_evidence_source=payload.get("radius_evidence_source"),
159+
critical_facility=bool(payload.get("critical_facility")),
160+
planning_project_id=payload.get("planning_project_id"),
131161
)
132162
status = 200 if result["validation"]["valid"] else 400
133163
if "snapshot_digest_mismatch" in result["validation"]["errors"]:
@@ -157,6 +187,13 @@ async def uwm_livability_s2_rollout(request: Request):
157187
requested_at=str(payload.get("requested_at") or ""),
158188
actor_id=str(username),
159189
alternative_land_use_class=payload.get("alternative_land_use_class"),
190+
action_type=str(payload.get("action_type") or "change_land_use"),
191+
facility_class=payload.get("facility_class"),
192+
facility_id=payload.get("facility_id"),
193+
service_radius_m=payload.get("service_radius_m"),
194+
radius_evidence_source=payload.get("radius_evidence_source"),
195+
critical_facility=bool(payload.get("critical_facility")),
196+
planning_project_id=payload.get("planning_project_id"),
160197
)
161198
return JSONResponse(result)
162199
except S2ProductInvalid as error:
@@ -184,12 +221,16 @@ async def uwm_livability_s2_run(request: Request):
184221
return _product_error(error)
185222
except S2RunNotFound as error:
186223
return JSONResponse({"error": str(error).strip("'\"")}, status_code=404)
224+
except S2RunInvalid as error:
225+
return JSONResponse({"error": str(error)}, status_code=409)
187226

188227

189228
def get_uwm_livability_s2_routes() -> list:
190229
return [
191230
Route("/api/uwm/livability/s2/catalog", uwm_livability_s2_catalog, methods=["GET"]),
192231
Route("/api/uwm/livability/s2/parcels", uwm_livability_s2_parcels, methods=["GET"]),
232+
Route("/api/uwm/livability/s2/facilities", uwm_livability_s2_facilities, methods=["GET"]),
233+
Route("/api/uwm/livability/s2/planning-projects", uwm_livability_s2_planning_projects, methods=["GET"]),
193234
Route("/api/uwm/livability/s2/parcels/{parcel_id}", uwm_livability_s2_parcel, methods=["GET"]),
194235
Route("/api/uwm/livability/s2/validate-action", uwm_livability_s2_validate_action, methods=["POST"]),
195236
Route("/api/uwm/livability/s2/rollout", uwm_livability_s2_rollout, methods=["POST"]),
Lines changed: 112 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,112 @@
1+
"""API routes for real-data UWM multi-stage intervention planning."""
2+
3+
from __future__ import annotations
4+
5+
import asyncio
6+
from typing import Any
7+
8+
from starlette.requests import Request
9+
from starlette.responses import JSONResponse
10+
from starlette.routing import Route
11+
12+
from data_agent.uwm.multistage_intervention_planner import (
13+
MultiStageInterventionPlannerService,
14+
)
15+
16+
from .helpers import _get_user_from_request, _set_user_context
17+
18+
19+
_SERVICE = MultiStageInterventionPlannerService()
20+
21+
22+
def _authorize(request: Request) -> JSONResponse | None:
23+
user = _get_user_from_request(request)
24+
if not user:
25+
return JSONResponse({"error": "Unauthorized"}, status_code=401)
26+
_set_user_context(user)
27+
return None
28+
29+
30+
async def overview(request: Request):
31+
denied = _authorize(request)
32+
if denied:
33+
return denied
34+
try:
35+
return JSONResponse(await asyncio.to_thread(_SERVICE.overview))
36+
except Exception as exc:
37+
return JSONResponse({"error": str(exc)}, status_code=500)
38+
39+
40+
async def actions(request: Request):
41+
denied = _authorize(request)
42+
if denied:
43+
return denied
44+
query = request.query_params
45+
action_types = [value for value in query.get("action_types", "").split(",") if value]
46+
try:
47+
payload = await asyncio.to_thread(
48+
_SERVICE.actions,
49+
county=query.get("county", ""),
50+
action_types=action_types or None,
51+
limit=int(query.get("limit", "100")),
52+
)
53+
return JSONResponse(payload)
54+
except ValueError as exc:
55+
return JSONResponse({"error": str(exc)}, status_code=400)
56+
except Exception as exc:
57+
return JSONResponse({"error": str(exc)}, status_code=500)
58+
59+
60+
async def plan(request: Request):
61+
denied = _authorize(request)
62+
if denied:
63+
return denied
64+
try:
65+
body: dict[str, Any] = await request.json()
66+
return JSONResponse(await asyncio.to_thread(_SERVICE.plan, body))
67+
except ValueError as exc:
68+
return JSONResponse({"error": str(exc)}, status_code=400)
69+
except Exception as exc:
70+
return JSONResponse({"error": str(exc)}, status_code=500)
71+
72+
73+
async def get_run(request: Request):
74+
denied = _authorize(request)
75+
if denied:
76+
return denied
77+
try:
78+
return JSONResponse(
79+
await asyncio.to_thread(_SERVICE.get_run, request.path_params["run_id"])
80+
)
81+
except FileNotFoundError as exc:
82+
return JSONResponse({"error": str(exc)}, status_code=404)
83+
except ValueError as exc:
84+
return JSONResponse({"error": str(exc)}, status_code=400)
85+
except Exception as exc:
86+
return JSONResponse({"error": str(exc)}, status_code=500)
87+
88+
89+
async def get_map(request: Request):
90+
denied = _authorize(request)
91+
if denied:
92+
return denied
93+
try:
94+
return JSONResponse(
95+
await asyncio.to_thread(_SERVICE.get_map, request.path_params["run_id"])
96+
)
97+
except FileNotFoundError as exc:
98+
return JSONResponse({"error": str(exc)}, status_code=404)
99+
except ValueError as exc:
100+
return JSONResponse({"error": str(exc)}, status_code=400)
101+
except Exception as exc:
102+
return JSONResponse({"error": str(exc)}, status_code=500)
103+
104+
105+
def get_uwm_multistage_intervention_routes() -> list:
106+
return [
107+
Route("/api/uwm/multistage-intervention/overview", overview, methods=["GET"]),
108+
Route("/api/uwm/multistage-intervention/actions", actions, methods=["GET"]),
109+
Route("/api/uwm/multistage-intervention/plan", plan, methods=["POST"]),
110+
Route("/api/uwm/multistage-intervention/runs/{run_id}", get_run, methods=["GET"]),
111+
Route("/api/uwm/multistage-intervention/runs/{run_id}/map", get_map, methods=["GET"]),
112+
]

0 commit comments

Comments
 (0)