lucashudsn commited on
Commit
d36bfb4
Β·
1 Parent(s): bed0f9c

wip: coordinate-check utility (pre-wow snapshot)

Browse files
Files changed (1) hide show
  1. app/check_coords.py +466 -0
app/check_coords.py ADDED
@@ -0,0 +1,466 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ check_coords.py β€” verify (and optionally fix) the coordinates of every
3
+ enriched surf break in data/australia-surf-breaks-enriched.json.
4
+
5
+ Two phases:
6
+
7
+ 1. CHECK (free, no API calls) β€” deterministic sanity checks per break:
8
+ - invalid: coords missing or not finite numbers
9
+ - out_of_australia: lat/lng outside Australia's bounding box
10
+ - wrong_state: lat/lng outside the break's own state bounding box
11
+ - duplicate: same point (within 0.005 deg) as a *different* break β€”
12
+ a common LLM failure mode is copying the nearest town
13
+ centre or another break's coordinates
14
+ - region_outlier: > REGION_OUTLIER_KM from the median of its
15
+ (state, region) cluster
16
+
17
+ 2. FIX (--fix) β€” for each flagged break, look the spot up in
18
+ OpenStreetMap via Nominatim (the same data the UI map renders, so the
19
+ marker lands exactly on the feature). The best candidate must pass the
20
+ deterministic checks above (in Australia, in the right state, not a
21
+ duplicate) before it is accepted. Spots OSM doesn't know can fall back
22
+ to a focused LLM estimate with --source llm or --source both.
23
+ A timestamped backup of the file is written before the first change;
24
+ the file is saved after every accepted fix.
25
+
26
+ Usage:
27
+ uv run python app/check_coords.py # report only (no network)
28
+ uv run python app/check_coords.py --fix # fix hard failures via OSM (no API key needed)
29
+ uv run python app/check_coords.py --fix --all # also fix region outliers
30
+ uv run python app/check_coords.py --fix --source both --max-fixes 10
31
+ # --source both additionally needs HF_TOKEN for the LLM fallback.
32
+ """
33
+
34
+ import json
35
+ import math
36
+ import os
37
+ import re
38
+ import sys
39
+ import time
40
+ import urllib.parse
41
+ import urllib.request
42
+ from collections import defaultdict
43
+ from datetime import datetime
44
+ from pathlib import Path
45
+
46
+ from huggingface_hub import InferenceClient
47
+
48
+ try: # package import from repo root
49
+ from app.generate_surf_break import MODEL_ID, PROVIDER, extract_json
50
+ except ImportError: # pragma: no cover β€” standalone fallback
51
+ MODEL_ID = "nvidia/NVIDIA-Nemotron-3-Ultra-550B-A55B-BF16"
52
+ PROVIDER = "deepinfra"
53
+
54
+ def extract_json(text: str) -> dict: # type: ignore[no-redef]
55
+ text = text.strip()
56
+ fence = re.search(r"```(?:json)?\s*(\{.*\})\s*```", text, re.E | re.S)
57
+ if fence:
58
+ text = fence.group(1)
59
+ else:
60
+ brace = re.search(r"\{.*\}", text, re.E | re.S)
61
+ text = brace.group(0)
62
+ return json.loads(text)
63
+
64
+ DATA_DIR = Path(__file__).parent.parent / "data"
65
+ ENRICHED_FILE = DATA_DIR / "australia-surf-breaks-enriched.json"
66
+
67
+ REQUEST_DELAY = 1.0 # seconds between inference calls
68
+ MAX_RETRIES = 2
69
+
70
+ # Rough bounding boxes (lat_min, lat_max, lng_min, lng_max) per state.
71
+ # Good enough to catch a break dropped in the wrong state/territory.
72
+ AUSTRALIA_BOX = (-44.5, -8.5, 109.5, 156.5)
73
+ STATE_BOXES = {
74
+ "Western Australia": (-50.5, -10.7, 112.5, 129.3),
75
+ "Northern Territory": (-26.6, -10.6, 129.0, 138.2),
76
+ "Queensland": (-43.8, -10.6, 135.5, 154.0),
77
+ "South Australia": (-39.1, -25.8, 129.0, 141.1),
78
+ # NSW/QLD border runs along 28.5S out to the coast at ~153.63E (Cape Byron);
79
+ # NSW/SA border is the 141E meridian.
80
+ "New South Wales": (-39.1, -28.0, 140.9, 153.65),
81
+ "Victoria": (-39.3, -33.9, 140.9, 150.2),
82
+ "Tasmania": (-43.8, -40.4, 143.4, 149.6),
83
+ "ACT": (-36.0, -35.2, 148.6, 149.5),
84
+ }
85
+
86
+ DUP_RADIUS_DEG = 0.005 # ~0.55 km β€” effectively "same point"
87
+ REGION_OUTLIER_KM = 100.0
88
+
89
+
90
+ # ---------------------------------------------------------------- checks
91
+
92
+ def _coords(break_: dict) -> tuple[float, float] | None:
93
+ try:
94
+ lat = float(break_["location"]["coordinates"]["lat"])
95
+ lng = float(break_["location"]["coordinates"]["lng"])
96
+ if not (math.isfinite(lat) and math.isfinite(lng)):
97
+ return None
98
+ return lat, lng
99
+ except (KeyError, TypeError, ValueError):
100
+ return None
101
+
102
+
103
+ def _in_box(lat: float, lng: float, box: tuple) -> bool:
104
+ lat_min, lat_max, lng_min, lng_max = box
105
+ return lat_min <= lat <= lat_max and lng_min <= lng <= lng_max
106
+
107
+
108
+ def _haversine_km(a: tuple[float, float], b: tuple[float, float]) -> float:
109
+ lat1, lon1, lat2, lon2 = map(math.radians, (*a, *b))
110
+ h = (
111
+ math.sin((lat2 - lat1) / 2) ** 2
112
+ + math.cos(lat1) * math.cos(lat2) * math.sin((lon2 - lon1) / 2) ** 2
113
+ )
114
+ return 2 * 6371.0 * math.asin(math.sqrt(h))
115
+
116
+
117
+ def check_point(lat: float, lng: float, state: str, other_points: list[tuple[float, float]]) -> list[str]:
118
+ """Run the deterministic checks for one (lat, lng) against one state.
119
+
120
+ ``other_points`` are the coordinates of the *other* breaks (already
121
+ placed), used for the duplicate check. Returns a list of reason tags
122
+ (empty == passes).
123
+ """
124
+ reasons: list[str] = []
125
+ if not _in_box(lat, lng, AUSTRALIA_BOX):
126
+ reasons.append("out_of_australia")
127
+ box = STATE_BOXES.get(state)
128
+ if box and not _in_box(lat, lng, box):
129
+ reasons.append("wrong_state")
130
+ for o_lat, o_lng in other_points:
131
+ if abs(lat - o_lat) < DUP_RADIUS_DEG and abs(lng - o_lng) < DUP_RADIUS_DEG:
132
+ reasons.append("duplicate")
133
+ break
134
+ return reasons
135
+
136
+
137
+ def check_all(breaks: list[dict]) -> list[dict]:
138
+ """Check every break; returns [{break, coords, reasons}] for flagged ones."""
139
+ points = {id(b): _coords(b) for b in breaks}
140
+ flagged = []
141
+
142
+ # Region medians (per state+region cluster) for the outlier check.
143
+ clusters: dict[tuple, list[tuple[float, float]]] = defaultdict(list)
144
+ for b in breaks:
145
+ p = points[id(b)]
146
+ if p:
147
+ clusters[(b.get("state", ""), b.get("region", ""))].append(p)
148
+ medians = {}
149
+ for key, pts in clusters.items():
150
+ medians[key] = (
151
+ sorted(p[0] for p in pts)[len(pts) // 2],
152
+ sorted(p[1] for p in pts)[len(pts) // 2],
153
+ )
154
+
155
+ for b in breaks:
156
+ p = points[id(b)]
157
+ reasons: list[str] = []
158
+ if p is None:
159
+ reasons.append("invalid")
160
+ else:
161
+ others = [
162
+ points[id(o)]
163
+ for o in breaks
164
+ if o is not b and (points[id(o)] is not None)
165
+ ]
166
+ reasons.extend(check_point(p[0], p[1], b.get("state", ""), others))
167
+ med = medians.get((b.get("state", ""), b.get("region", "")))
168
+ if med and _haversine_km(p, med) > REGION_OUTLIER_KM:
169
+ reasons.append("region_outlier")
170
+ if reasons:
171
+ flagged.append({"break": b, "coords": p, "reasons": reasons})
172
+ return flagged
173
+
174
+
175
+ # ---------------------------------------------------------------- fix
176
+
177
+ # OpenStreetMap (Nominatim) is the ground-truth source: the UI map renders
178
+ # OSM tiles, so a point taken from OSM sits exactly on the feature shown.
179
+ # Usage policy: <=1 request/second, descriptive User-Agent, no caching needed
180
+ # for a one-off maintenance run.
181
+ NOMINATIM_URL = "https://nominatim.openstreetmap.org/search"
182
+ NOMINATIM_UA = "wavereader-coord-check/1.0 (one-off surf-break data maintenance)"
183
+ _nominatim_last_call = 0.0
184
+
185
+ # Score bonuses for candidate feature kinds: we want the physical spot
186
+ # (beach/reef/point), not the suburb it happens to be inside.
187
+ _TYPE_SCORE = {
188
+ "beach": 6, "reef": 6, "bay": 5, "water": 5, "shoreline": 5,
189
+ "coastline": 4, "headland": 4, "cape": 4, "point": 4, "rock": 3,
190
+ "bare_rock": 3, "strait": 3, "harbour": 2, "island": 2, "river": 1,
191
+ "viewpoint": 2, "attraction": 1, "place": 0, "administrative": -6,
192
+ }
193
+ _CLASS_SCORE = {"natural": 3, "waterway": 3, "man_made": 1, "tourism": 2,
194
+ "highway": -5, "boundary": -10, "administrative": -10}
195
+
196
+
197
+ def _pace_nominatim() -> None:
198
+ global _nominatim_last_call
199
+ wait = 1.05 - (time.time() - _nominatim_last_call)
200
+ if wait > 0:
201
+ time.sleep(wait)
202
+ _nominatim_last_call = time.time()
203
+
204
+
205
+ def nominatim_search(query: str, limit: int = 5) -> list[dict]:
206
+ """One Nominatim free-form search, Australia-only, paced to 1 req/s."""
207
+ _pace_nominatim()
208
+ url = NOMINATIM_URL + "?" + urllib.parse.urlencode(
209
+ {"q": query, "format": "jsonv2", "limit": limit, "countrycodes": "au"}
210
+ )
211
+ req = urllib.request.Request(url, headers={"User-Agent": NOMINATIM_UA})
212
+ with urllib.request.urlopen(req, timeout=20) as resp:
213
+ return json.load(resp)
214
+
215
+
216
+ def _norm(s: str) -> str:
217
+ return re.sub(r"[^a-z0-9]+", " ", s.lower()).strip()
218
+
219
+
220
+ def _candidate_score(cand: dict, core: str, anchor: tuple[float, float] | None) -> float | None:
221
+ """Score one Nominatim hit; None means it is disqualified."""
222
+ dn = _norm(cand.get("display_name", ""))
223
+ if _norm(core) not in dn:
224
+ return None # not actually this spot
225
+ try:
226
+ lat, lng = float(cand["lat"]), float(cand["lon"])
227
+ except (KeyError, ValueError):
228
+ return None
229
+ score = _TYPE_SCORE.get(cand.get("type", ""), 1) + _CLASS_SCORE.get(cand.get("class", ""), 0)
230
+ if anchor is not None: # prefer the candidate nearest the region's trusted cluster
231
+ score -= 0.05 * _haversine_km((lat, lng), anchor)
232
+ return score
233
+
234
+
235
+ def osm_match(break_: dict, state: str, anchor: tuple[float, float] | None) -> tuple[float, float, str] | None:
236
+ """Geocode one break against OSM. Returns (lat, lng, display_name) of the
237
+ best candidate that lies inside the break's state, else None.
238
+
239
+ ``anchor`` is the median of the *other, unflagged* breaks in the same
240
+ region β€” used only to disambiguate, never to disqualify.
241
+ """
242
+ name = str(break_.get("name", ""))
243
+ base, inner = name, ""
244
+ m = re.match(r"^(.*?)\s*\((.*?)\)\s*$", name) # "Kelp Beds (Esperance)" -> core + locality
245
+ if m:
246
+ base, inner = m.group(1).strip(), m.group(2).strip()
247
+
248
+ variants: list[str] = []
249
+ for core in dict.fromkeys((name, base)):
250
+ variants.append(f"{core}, {state}, Australia")
251
+ if inner:
252
+ variants.insert(0, f"{base}, {inner}, {state}, Australia")
253
+ variants.append(f"{base}, Australia")
254
+
255
+ box = STATE_BOXES.get(state)
256
+ for query in dict.fromkeys(variants):
257
+ try:
258
+ hits = nominatim_search(query)
259
+ except Exception as e: # noqa: BLE001
260
+ print(f" ! Nominatim error for {query!r}: {e}")
261
+ continue
262
+ best, best_score = None, None
263
+ for cand in hits:
264
+ try:
265
+ lat, lng = float(cand["lat"]), float(cand["lon"])
266
+ except (KeyError, ValueError):
267
+ continue
268
+ if box and not _in_box(lat, lng, box):
269
+ continue # wrong corner of the country β€” never accept
270
+ score = _candidate_score(cand, base if inner else name, anchor)
271
+ if score is None:
272
+ continue
273
+ if best_score is None or score > best_score:
274
+ best, best_score = (lat, lng, cand.get("display_name", "")), score
275
+ if best:
276
+ return best
277
+ return None
278
+
279
+
280
+ def build_coord_prompt(break_: dict, current) -> str:
281
+ cur = (
282
+ f"The coordinates currently on file are lat {current[0]}, lng {current[1]} β€” "
283
+ f"they are suspected to be wrong, so re-estimate from your real-world knowledge "
284
+ f"of the spot rather than trusting them."
285
+ if current
286
+ else "The break currently has no usable coordinates on file."
287
+ )
288
+ return (
289
+ "You are a surf-break geocoding assistant. You are given one Australian surf break. "
290
+ "Return the real-world geographic coordinates of the break's primary takeoff zone "
291
+ "(the actual point on the coast, not the nearest town centre) as JSON ONLY β€” "
292
+ "no markdown fences, no commentary β€” in exactly this shape:\n"
293
+ '{"lat": <decimal degrees, south is negative>, "lng": <decimal degrees, east is positive>}\n\n'
294
+ f"Break: {break_.get('name')} β€” state: {break_.get('state')}, region: {break_.get('region')}\n"
295
+ f"{cur}\n"
296
+ f"Known description: {str(break_.get('description', ''))[:300]}"
297
+ )
298
+
299
+
300
+ def fetch_corrected_coords(client: InferenceClient, break_: dict, current) -> tuple[float, float] | None:
301
+ """One chat-completion call (with retries) asking for lat/lng only."""
302
+ prompt = build_coord_prompt(break_, current)
303
+ for attempt in range(1, MAX_RETRIES + 1):
304
+ try:
305
+ completion = client.chat.completions.create(
306
+ model=MODEL_ID,
307
+ messages=[{"role": "user", "content": prompt}],
308
+ max_tokens=100,
309
+ temperature=0.0,
310
+ )
311
+ data = extract_json(completion.choices[0].message.content or "")
312
+ lat, lng = float(data["lat"]), float(data["lng"])
313
+ if math.isfinite(lat) and math.isfinite(lng):
314
+ return lat, lng
315
+ except Exception as e: # noqa: BLE001
316
+ print(f" ! Attempt {attempt}/{MAX_RETRIES} failed: {e}")
317
+ if attempt < MAX_RETRIES:
318
+ time.sleep(2 * attempt)
319
+ return None
320
+
321
+
322
+ def fix_flagged(
323
+ client: InferenceClient | None,
324
+ breaks: list[dict],
325
+ flagged: list[dict],
326
+ max_fixes: int,
327
+ source: str = "osm",
328
+ ) -> None:
329
+ """Replace coordinates on flagged breaks using a ground-truth source
330
+ (OSM by default, LLM fallback for spots OSM doesn't know); save after
331
+ each accepted fix."""
332
+ fixed, kept = [], []
333
+ backup = ENRICHED_FILE.with_name(f"{ENRICHED_FILE.stem}-backup-{datetime.now():%Y%m%d-%H%M%S}{ENRICHED_FILE.suffix}")
334
+ backup_written = False
335
+
336
+ # Anchor per (state, region): median of the unflagged breaks in the cluster.
337
+ flagged_ids = {id(item["break"]) for item in flagged}
338
+ clusters: dict[tuple, list[tuple[float, float]]] = defaultdict(list)
339
+ for b in breaks:
340
+ if id(b) in flagged_ids:
341
+ continue
342
+ p = _coords(b)
343
+ if p:
344
+ clusters[(b.get("state", ""), b.get("region", ""))].append(p)
345
+ anchors = {
346
+ k: (sorted(p[0] for p in pts)[len(pts) // 2], sorted(p[1] for p in pts)[len(pts) // 2])
347
+ for k, pts in clusters.items()
348
+ }
349
+
350
+ for item in flagged:
351
+ if len(fixed) + len(kept) >= max_fixes:
352
+ print(f"Reached --max-fixes {max_fixes}, stopping.")
353
+ break
354
+ b = item["break"]
355
+ label = f"{b.get('name')} ({b.get('state')} / {b.get('region')})"
356
+ print(f"Fixing: {label} β€” flagged: {', '.join(item['reasons'])}")
357
+
358
+ anchor = anchors.get((b.get("state", ""), b.get("region", "")))
359
+ new, note = None, ""
360
+ if source in ("osm", "both"):
361
+ match = osm_match(b, b.get("state", ""), anchor)
362
+ if match:
363
+ new, note = (match[0], match[1]), match[2]
364
+ else:
365
+ print(" -> not found in OpenStreetMap")
366
+ if new is None and source in ("llm", "both") and client is not None:
367
+ llm = fetch_corrected_coords(client, b, item["coords"])
368
+ if llm is not None:
369
+ new, note = llm, "LLM estimate"
370
+ if new is None:
371
+ kept.append(label)
372
+ print(" -> kept old coords (no ground-truth match found)")
373
+ continue
374
+
375
+ # Validate the proposal against the deterministic checks before accepting.
376
+ others = [
377
+ p
378
+ for p in (_coords(o) for o in breaks if o is not b)
379
+ if p is not None
380
+ ]
381
+ problems = check_point(new[0], new[1], b.get("state", ""), others)
382
+ problems = [r for r in problems if r != "region_outlier"] # outlier vs old cluster is expected
383
+ if "duplicate" in item["reasons"]: # old point was already a duplicate;
384
+ problems = [r for r in problems if r != "duplicate"] # adjacent spots may share a point
385
+ if problems:
386
+ kept.append(label)
387
+ print(f" -> kept old coords (new point fails: {', '.join(problems)})")
388
+ continue
389
+
390
+ if not backup_written:
391
+ ENRICHED_FILE.replace(backup)
392
+ backup_written = True
393
+ print(f"Backup written to {backup}")
394
+ b["location"]["coordinates"]["lat"] = round(new[0], 6)
395
+ b["location"]["coordinates"]["lng"] = round(new[1], 6)
396
+ fixed.append(label)
397
+ print(f" -> set to lat {new[0]:.4f}, lng {new[1]:.4f} [{note}]")
398
+ ENRICHED_FILE.write_text(json.dumps(breaks, indent=2))
399
+
400
+ print(f"\nFixed {len(fixed)} break(s), kept old coords on {len(kept)}.")
401
+ for label in kept:
402
+ print(f" - unresolved: {label}")
403
+
404
+
405
+ # ---------------------------------------------------------------- main
406
+
407
+ def main() -> None:
408
+ args = [a for a in sys.argv[1:]]
409
+ do_fix = "--fix" in args
410
+ fix_all = "--all" in args
411
+ source = "osm"
412
+ if "--source" in args:
413
+ source = args[args.index("--source") + 1]
414
+ if source not in ("osm", "llm", "both"):
415
+ raise SystemExit("--source must be osm, llm, or both")
416
+ max_fixes = float("inf")
417
+ if "--max-fixes" in args:
418
+ max_fixes = int(args[args.index("--max-fixes") + 1])
419
+
420
+ breaks = json.loads(ENRICHED_FILE.read_text())
421
+ if not isinstance(breaks, list):
422
+ raise SystemExit(f"{ENRICHED_FILE} is not a JSON list")
423
+ print(f"Loaded {len(breaks)} breaks from {ENRICHED_FILE}")
424
+
425
+ flagged = check_all(breaks)
426
+ if not flagged:
427
+ print("All coordinates passed every check. Nothing to do.")
428
+ return
429
+
430
+ print(f"\n{len(flagged)} flagged:\n")
431
+ for item in flagged:
432
+ b = item["break"]
433
+ c = item["coords"]
434
+ coords_str = f"lat {c[0]}, lng {c[1]}" if c else "missing/invalid"
435
+ print(f" [{', '.join(item['reasons'])}] {b.get('name')} ({b.get('state')} / {b.get('region')}) β€” {coords_str}")
436
+
437
+ if not do_fix:
438
+ print("\nRun with --fix to correct these against OpenStreetMap.")
439
+ return
440
+
441
+ targets = flagged if fix_all else [
442
+ f for f in flagged if "invalid" in f["reasons"]
443
+ or "out_of_australia" in f["reasons"]
444
+ or "wrong_state" in f["reasons"]
445
+ or "duplicate" in f["reasons"]
446
+ ]
447
+ skipped = len(flagged) - len(targets)
448
+ if skipped > 0:
449
+ print(f"\n--fix targets the {len(targets)} hard failures; {skipped} region outlier(s) "
450
+ "skipped (use --all to include them).")
451
+
452
+ client = None
453
+ if source in ("llm", "both"):
454
+ if not os.environ.get("HF_TOKEN"):
455
+ raise SystemExit("HF_TOKEN is not set β€” required for --source llm/both.")
456
+ client = InferenceClient(provider=PROVIDER, api_key=os.environ["HF_TOKEN"])
457
+
458
+ fix_flagged(client, breaks, targets, max_fixes, source=source)
459
+
460
+ # Re-run the checks on the in-memory list to show the final state.
461
+ remaining = check_all(breaks)
462
+ print(f"{len(remaining)} break(s) still flagged after fix attempt.")
463
+
464
+
465
+ if __name__ == "__main__":
466
+ main()