Skip to content

Commit d26c10e

Browse files
authored
conn bar and no barr (#510)
* conn bar and no barr * move valid vars in connectivity
1 parent 27fcdda commit d26c10e

6 files changed

Lines changed: 160 additions & 18 deletions

File tree

rapida/cli/assess.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -147,7 +147,7 @@ def build_variable_help():
147147
if vars_list:
148148
vars_str = ", ".join(vars_list)
149149
parts.append(f"{comp} ({vars_str})")
150-
if comp == 'population': print(len(vars_list))
150+
151151
return "The variable/s to be assessed. Will be filtered by selected components. Available variables per component:\n\n" + "\n\n".join(parts)
152152

153153

@@ -211,7 +211,7 @@ def assess(ctx, all=False, components=None, variables=None, year=None, datetime
211211
rapida assess -c landuse -dt 2025-02-01/2025-05-31 -cc 10: Search Sentinel 2 item which is less than 10% of cloud cover from February to May 2025.
212212
213213
"""
214-
214+
progress = ctx.obj.get('progress')
215215
if not is_rapida_initialized():
216216
return
217217

@@ -231,7 +231,7 @@ def assess(ctx, all=False, components=None, variables=None, year=None, datetime
231231
sys.exit(0)
232232

233233
logger.info(f'Current project/folder: {prj.path}')
234-
with Progress(disable=False, console=None) as progress:
234+
with progress:
235235
with Session() as session:
236236
all_components = session.get_components()
237237
target_components = components

rapida/cli/connectivity.py

Lines changed: 40 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
from typing import Union
1+
from typing import Union, Iterable
22

33
import click
44
import logging
@@ -7,10 +7,21 @@
77
from rapida.connectivity import run_connectivity_analysis
88
from rapida.cli.aclick import AsyncCommand
99
from rapida.connectivity.isochrone import MODE_MAP
10-
10+
from rapida.cli.assess import get_variables_by_components
1111
logger = logging.getLogger(__name__)
1212

1313

14+
15+
def validate_variables(ctx, param, value):
16+
"""
17+
click callback function to validate polulation
18+
"""
19+
valid_vars = get_variables_by_components(['population'])['population']
20+
invalid = [v for v in value if v not in valid_vars]
21+
if invalid:
22+
raise click.BadParameter(f"Invalid variable{'s' if len(invalid) > 1 else ''}: {', '.join(invalid)} for population . Valid options: {', '.join(valid_vars)}")
23+
return value
24+
1425
def parse_intervals(ctx, param, value):
1526
"""Parses a comma-separated string of numbers into a list of integers."""
1627
if not value:
@@ -46,6 +57,20 @@ def parse_intervals(ctx, param, value):
4657
help="Comma-separated time intervals in minutes for the catchment areas."
4758
)
4859

60+
61+
@click.option(
62+
'-sd', '--sites-dataset',
63+
type=click.Path(exists=True, file_okay=True, dir_okay=False, readable=True),
64+
default=None,
65+
help="Path to an OGR-supported vector data source (e.g., GPKG, Shapefile) containing sites."
66+
)
67+
@click.option(
68+
'-sl', '--sites-layer',
69+
type=str,
70+
default="0",
71+
help="Name or index of the layer to use from sites dataset. Defaults to the first layer (layer 0)."
72+
)
73+
4974
@click.option(
5075
'-bd', '--barriers-dataset',
5176
type=click.Path(exists=True, file_okay=True, dir_okay=False, readable=True),
@@ -59,12 +84,19 @@ def parse_intervals(ctx, param, value):
5984
help="Name or index of the layer to use from barriers dataset. Defaults to the first layer (layer 0)."
6085
)
6186

87+
6288
@click.option('-bb', "--barriers-buffer",
6389
type=int,
6490
default=5,
6591
required=False,
6692
help="The value in meters to used to buffer the geometries in barriers/dataset/layer in case the barriers are lines"
6793
)
94+
95+
@click.option('--popvar', required=False, multiple=True,
96+
type=str, callback=validate_variables,
97+
help=f"Open or more RAPIDA population variable to compute zonal stats for withing the connectivity zones"
98+
)
99+
68100
@click.option(
69101
"--dst-dir",
70102
"-d", # Short option
@@ -84,12 +116,15 @@ def parse_intervals(ctx, param, value):
84116
@click.pass_context
85117
async def connectivity(ctx, bbox:tuple[float, float, float, float]=None, travel_mode:str=None,
86118
time_intervals:list[int] =None, dst_dir:str=None,
87-
barriers_dataset:str=None, barriers_layer:str=None, barriers_buffer:int=None
119+
barriers_dataset:str=None, barriers_layer:str=None, barriers_buffer:int=None,
120+
sites_dataset:str=None, sites_layer:str=None,popvar:str|tuple[str]=None
88121
):
89-
logger.info(f'Running connectivity analysis ')
122+
logger.info(f'Running connectivity analysis')
90123
progress = ctx.obj.get('progress')
91124
with progress:
92125
return await run_connectivity_analysis(
93126
bbox=bbox, dst_dir=dst_dir, travel_mode=travel_mode, time_intervals=time_intervals,
94-
barriers_dataset=barriers_dataset, barriers_layer=barriers_layer, barriers_buffer=barriers_buffer, progress=progress
127+
barriers_dataset=barriers_dataset, barriers_layer=barriers_layer, barriers_buffer=barriers_buffer,
128+
sites_dataset=sites_dataset, sites_layer=sites_layer, pop_vars=popvar,
129+
progress=progress
95130
)

rapida/connectivity/__init__.py

Lines changed: 46 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -2,27 +2,66 @@
22
import os.path
33
from rapida.util.bbox_param_type import get_best_semantic_label
44
from rich.progress import Progress
5-
from rapida.connectivity.io import prepare_osm_pbf,extract_health_sites, extract_origins_from_geojson
5+
from rapida.connectivity.io import prepare_osm_pbf,extract_health_sites, extract_origins_from_geojson, extract_origins
66
from rapida.connectivity.graph import compile_valhalla_graph
77
from rapida.connectivity.isochrone import connectivity_areas
8+
# from rapida.cli.assess import assess
9+
# import click
10+
# from rapida.project.project import Project
11+
# from tempfile import TemporaryDirectory
812

913

1014

1115
async def run_connectivity_analysis(
1216
bbox:tuple[float, float, float, float]=None, travel_mode:str=None, time_intervals:list[int] =None,
13-
dst_dir:str=None, barriers_dataset:str=None, barriers_layer:str=None, barriers_buffer:int=None, progress:Progress=None
17+
dst_dir:str=None, barriers_dataset:str=None, barriers_layer:str=None, barriers_buffer:int=None,
18+
sites_dataset:str=None, sites_layer:str=None,pop_vars:str|tuple[str]=None,
19+
progress:Progress=None
1420
):
1521
bbox_label = get_best_semantic_label(bbox=bbox)
1622
dest_dir = os.path.join(dst_dir, bbox_label)
1723
bbox_pbf = await prepare_osm_pbf(bbox=bbox, dst_dir=dest_dir, progress=progress)
18-
health_sites = await extract_health_sites(pbf_path=bbox_pbf, dst_dir=dest_dir, progress=progress)
24+
if sites_dataset is None:
25+
sites = await extract_health_sites(pbf_path=bbox_pbf, dst_dir=dest_dir, progress=progress)
26+
else:
27+
sites = sites_dataset
28+
1929
dag_tar_path = await compile_valhalla_graph(pbf_path=bbox_pbf,dst_dir=dest_dir, progress=progress)
20-
origins = extract_origins_from_geojson(geojson_path=health_sites)
30+
origins = extract_origins(sites_dataset=sites, src_layer=sites_layer)
31+
32+
2133
results = await connectivity_areas(
34+
tar_path=dag_tar_path, origins=origins, travel_mode=travel_mode, intervals_minutes=time_intervals)
35+
isochrones_path = os.path.join(dest_dir, 'isochrones.geojson')
36+
with open(isochrones_path, "w") as f:
37+
json.dump(results, f, indent=2)
38+
39+
# with TemporaryDirectory(dir=dest_dir, delete=False) as project_folder:
40+
# project = Project(path=project_folder, polygons=isochrones_path, comment='temp project for conn isochrones')
41+
#
42+
#
43+
# with click.Context(assess) as ctx:
44+
# ctx.ensure_object(dict)
45+
# ctx.obj['progress'] = progress
46+
# # 2. Use invoke. Do NOT pass 'ctx' manually here.
47+
# # Click intercepts this and injects it as the first argument automatically.
48+
# await ctx.invoke(
49+
# assess,
50+
# components=('population',),
51+
# variables=pop_vars,
52+
# year=2026,
53+
# project=project.path,
54+
# force=False
55+
# )
56+
57+
58+
if barriers_dataset is not None:
59+
barrier_results = await connectivity_areas(
2260
tar_path=dag_tar_path, origins=origins, travel_mode=travel_mode, intervals_minutes=time_intervals,
2361
barriers_dataset=barriers_dataset, barriers_layer=barriers_layer, barriers_buffer=barriers_buffer
24-
)
25-
with open(os.path.join(dest_dir, 'isochrones.geojson'), "w") as f:
26-
json.dump(results, f, indent=2)
62+
)
63+
with open(os.path.join(dest_dir, 'isochrones_with_barriers.geojson'), "w") as f:
64+
json.dump(barrier_results, f, indent=2)
65+
2766

2867
return

rapida/connectivity/io.py

Lines changed: 63 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@
1111
from rich.progress import Progress
1212
import json
1313
from shapely.geometry import shape, mapping
14-
from osgeo import gdal
14+
from osgeo import gdal, ogr, osr
1515
from shapely.wkb import loads as load_wkb
1616
from shapely.ops import orient, transform
1717
import numpy as np
@@ -58,6 +58,7 @@ async def prepare_osm_pbf(bbox: tuple[float, float, float, float], dst_dir: str
5858

5959
# Perform quick spatial overlay check
6060
gdf = gpd.GeoDataFrame.from_features(features, crs="EPSG:4326")
61+
6162
intersecting = gdf[gdf.intersects(bbox_geom)]
6263

6364
if intersecting.empty:
@@ -214,6 +215,67 @@ def extract_origins_from_geojson(geojson_path: str) -> list[tuple[float, float]]
214215
return origins
215216

216217

218+
def extract_origins(sites_dataset: str=None, src_layer: str = None) -> list[tuple[float, float]]:
219+
"""
220+
Extracts a list of (longitude, latitude) tuples from a spatial file using OGR.
221+
Handles reprojection to WGS84 (EPSG:4326) if the source is not in lat/lon.
222+
"""
223+
# Open the dataset
224+
with gdal.OpenEx(sites_dataset, gdal.OF_VECTOR) as src_ds:
225+
226+
try:
227+
layer = src_ds.GetLayer(int(src_layer))
228+
except ValueError:
229+
layer = src_ds.GetLayerByName(str(src_layer))
230+
231+
if layer is None:
232+
raise ValueError(f"Layer '{src_layer}' could not be found in the dataset {sites_dataset}.")
233+
234+
235+
236+
# Set up coordinate transformation to WGS84 (EPSG:4326)
237+
source_srs = layer.GetSpatialRef()
238+
target_srs = osr.SpatialReference()
239+
target_srs.ImportFromEPSG(4326)
240+
241+
# OGR 3+ strict axis mapping strategy (ensures Longitude/Latitude order)
242+
target_srs.SetAxisMappingStrategy(osr.OAMS_TRADITIONAL_GIS_ORDER)
243+
244+
transform = None
245+
if source_srs and not source_srs.IsSame(target_srs):
246+
transform = osr.CoordinateTransformation(source_srs, target_srs)
247+
248+
origins = []
249+
layer.ResetReading()
250+
# Iterate through features
251+
for feature in layer:
252+
# If you still need the filter: if feature.GetField("osm_id") != 80: continue
253+
254+
geom = feature.GetGeometryRef()
255+
if geom is not None:
256+
# Clone geometry to avoid modifying original layer data during transform
257+
geom_clone = geom.Clone()
258+
259+
# Reproject if necessary
260+
if transform:
261+
geom_clone.Transform(transform)
262+
263+
# Helper function to recursively extract points from nested collections
264+
def extract_points(g):
265+
name = g.GetGeometryName()
266+
if name == "POINT":
267+
origins.append((g.GetX(), g.GetY()))
268+
elif name in ("MULTIPOINT", "GEOMETRYCOLLECTION"):
269+
for i in range(g.GetGeometryCount()):
270+
sub_geom = g.GetGeometryRef(i)
271+
if sub_geom is not None:
272+
extract_points(sub_geom)
273+
274+
extract_points(geom_clone)
275+
276+
return origins
277+
278+
217279
def read_barriers_grid(src_path: str, src_layer: str = None, barriers_buffer:float=None) -> list:
218280
"""Reads a vector source and cuts features into micro-tiles to stay under Valhalla's limit."""
219281
if not src_path:

rapida/project/project.py

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,7 @@
1616
from geopandas import GeoDataFrame
1717
from osgeo import gdal, ogr, osr
1818
from azure.storage.fileshare import ShareClient
19-
19+
import reverse_geocoder as rg
2020
from rapida import constants
2121
from rapida.az.blobstorage import check_blob_exists, delete_blob
2222
from rapida.session import Session
@@ -159,12 +159,17 @@ def __init__(self, path: str,polygons: str = None,
159159
iso3_codes = []
160160
for lat, lon in zip(lats, lons):
161161
try:
162-
iso3_codes.append(fetch_ccode(lat=lat, lon=lon))
162+
result = rg.search((lat, lon))[0]
163+
iso2_cc = result.get('cc', '')
164+
country = coco.convert(names=iso2_cc, to='ISO3')
165+
iso3_codes.append(country)
163166
except Exception as e:
164167
logger.warning(f"Failed to fetch ISO3 for point ({lat}, {lon}): {e}")
165168
iso3_codes.append(None)
166169
gdf["iso3"] = iso3_codes
170+
167171
self.countries = tuple(sorted(set(filter(lambda x: x in COUNTRY_CODES, gdf["iso3"]))))
172+
168173
gdf.to_file(
169174
filename=self.geopackage_file_path,
170175
driver="GPKG",

rapida/stats/raster_zonal_stats.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -257,4 +257,5 @@ def progress_callback(completed, message, progress=progress, task=task):
257257
if vname in egdf.columns.tolist():
258258
egdf.drop(columns=[vname], inplace=True)
259259
combined = egdf.merge(combined, on='geometry', how='inner')
260+
combined = combined.sort_values(by='geometry', key=lambda geom: geom.to_geo_index().area, ascending=False)
260261
return combined

0 commit comments

Comments
 (0)