diff --git a/.gitignore b/.gitignore index cb1a262..6198ea5 100644 --- a/.gitignore +++ b/.gitignore @@ -218,3 +218,7 @@ __marimo__/ .streamlit/secrets.toml docs/ + +# Geocoder response caches: +*.sqlite +*.sqlite-journal diff --git a/README.md b/README.md index b70eb13..99cf52a 100644 --- a/README.md +++ b/README.md @@ -91,6 +91,15 @@ poetry run geocoder input.xlsx output.xlsx --api census --cacheRead archive.sqli poetry run geocoder input.xlsx output.xlsx --api census --noCache ``` +`--noCache` opens no cache file at all, so `--cache` and `--cacheRead` are +ignored when it is given; repeated addresses within a run still cost one call. +A `--cacheRead` file must hold entries for the provider being run, and is +rejected if it does not. + +Each provider's responses are stored in a table of their own, tagged with the +version of the API that produced them: entries written by an earlier version are +never served, and are replaced as their addresses are looked up again. + ## Providers | Provider | `--api` value | Coverage | API key | diff --git a/src/api.py b/src/api.py index 7ab9bc0..b114597 100644 --- a/src/api.py +++ b/src/api.py @@ -11,10 +11,12 @@ from abc import ABC, abstractmethod from dataclasses import dataclass, field from enum import IntEnum -from typing import Any, Callable, Dict, List, Optional, Type, TypeVar +from typing import Any, Callable, Dict, Iterator, List, Optional, Tuple, Type, TypeVar import requests +from .cache import Cache + KEY_ENV_VARS = { "geocodio": "GEOCODIO_API_KEY", "google": "GOOGLE_GEOCODING_API_KEY", @@ -273,6 +275,32 @@ def format_street_address(country: str, street: str, number: str = "", subpremis return f"{address}, {sublocality}" if sublocality else address +def cache_key(record: SourceRecord) -> str: + """ + Canonicalizes a record into the key its cached response is stored under. + + Every provider is asked the same question — the record's address string — so + the key is the same for all of them, and what differs between them is only + the table it is stored in and the version tag it carries. + + Letter case and runs of whitespace are folded so equivalent rows share one + cached response; nothing else is altered, since punctuation and + abbreviations can change what a geocoder returns. A row left with nothing + to ask about yields an empty key. + + Parameters + ---------- + record : SourceRecord + The record whose query is being composed. + + Return + ---------- + str + The canonical key for that record's response. + """ + return " ".join(record.address_string().split()).casefold() + + PROVIDERS: Dict[str, Type["Provider"]] = {} T = TypeVar("T", bound="Provider") @@ -289,25 +317,110 @@ def decorator(cls: Type[T]) -> Type[T]: class Provider(ABC): - """Base class for geocoding providers; subclasses self-register via @register.""" + """ + Base class for geocoding providers; subclasses self-register via @register. + + Subclasses supply ``_fetch`` and ``parse``; deduplication, caching, and the + limits a source row imposes are handled once here so no provider repeats + them. + + A subclass also declares CACHE_VERSION, which names everything about how it + calls its API that a cached response depends on — the API version, and any + dataset the call is pinned to. It tags the entries a provider writes, so a + later change to the call retires them instead of serving answers the + provider would no longer give. + """ name: str = "" requires_key: bool = False + CACHE_VERSION: str = "" MAX_ATTEMPTS = 3 RETRY_BACKOFF = 5 MAX_ACCURACY = AccuracyLevel.ROOFTOP - def __init__(self, api_key: Optional[str] = None): + def __init__(self, api_key: Optional[str] = None, cache: Optional[Cache] = None): self.api_key = api_key + self.cache = cache if cache is not None else Cache(self.name, self.CACHE_VERSION) @abstractmethod + def _fetch(self, records: List[SourceRecord]) -> Iterator[Tuple[SourceRecord, Dict[str, Any]]]: + """ + Yields each record paired with its raw provider response, as replies arrive. + + Everything yielded is cached, so only responses validated against the + provider's expected shape may be yielded; a failed response must raise. + """ + + @abstractmethod + def parse(self, raw: Dict[str, Any]) -> GeocodeResult: + """ + Converts one raw provider response into a normalized, graded GeocodeResult. + + Cached and freshly fetched responses both come through here, so a stored + response is always scored by the current grading rules. + + Parameters + ---------- + raw : Dict[str, Any] + One response as the provider returned it. + + Return + ---------- + GeocodeResult + The normalized result that response describes. + """ + def geocode(self, records: List[SourceRecord]) -> List[GeocodeResult]: - """Returns one GeocodeResult per input record, in the same order.""" + """ + Returns one GeocodeResult per input record, in the same order. - @staticmethod - def unqueryable_result() -> GeocodeResult: - """Returns the no-match a record with nothing left to query resolves to.""" - return GeocodeResult(match_notes=NO_ADDRESS_NOTE) + Records sharing a query are collapsed to a single API call and cached + queries skip the API entirely, with or without a cache file configured. + A row left with nothing to ask about is never sent, and every result is + bounded by the precision its source row could support. + + Parameters + ---------- + records : List[SourceRecord] + The source rows to geocode. + + Return + ---------- + List[GeocodeResult] + One result per input record, aligned by position. + """ + keys = [cache_key(record) for record in records] + responses = self.cache.lookup(key for key in keys if key) + + pending: Dict[str, SourceRecord] = {} + for key, record in zip(keys, records): + if key and key not in responses and key not in pending: + pending[key] = record + + if records: + print(f"{len(records)} rows, {len(responses) + len(pending)} distinct queries: {len(responses)} from cache, {len(pending)} to fetch") + + if pending: + for record, raw in self._fetch(list(pending.values())): + key = cache_key(record) + responses[key] = raw + self.cache.store(key, raw) + self.cache.commit() + + return [apply_legal_land_limit(record, self._result(key, responses)) for record, key in zip(records, keys)] + + def _result(self, key: str, responses: Dict[str, Any]) -> GeocodeResult: + """ + Reads one record's result out of the responses gathered for the run. + + A record with no key had nothing left to query and was never sent; one + whose key no response answers is a no match. + """ + if not key: + return GeocodeResult(match_notes=NO_ADDRESS_NOTE) + if key in responses: + return self.parse(responses[key]) + return GeocodeResult(match_notes="No match") def _request_with_retry(self, send: Callable[[], requests.Response]) -> requests.Response: """ diff --git a/src/cache.py b/src/cache.py new file mode 100644 index 0000000..377f9b5 --- /dev/null +++ b/src/cache.py @@ -0,0 +1,234 @@ +#!/usr/bin/python3 +# -.- coding: utf-8 -.- +# -.- dependencies: Python 3.8+ -.- + +""" +Geocoder — provider response cache + +Copyright (c) 2026 Pangaea Information Technologies, Ltd. +""" + +import json +import sqlite3 +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Dict, Iterable, Iterator, List, Optional, Sequence + +DEFAULT_CACHE_FILE = "geocoder-cache.sqlite" + +REQUIRED_COLUMNS = ("query", "version", "response") + +SCHEMA = """ +CREATE TABLE IF NOT EXISTS {table} ( + query TEXT PRIMARY KEY, + version TEXT NOT NULL, + response TEXT NOT NULL, + fetched_at TEXT NOT NULL +) +""" + + +class Cache: + """ + A SQLite-backed store of raw provider responses keyed by the query that produced them. + + The key is the address a run asks about, and nothing else: how the provider + was called is held apart from it, as the version tag every row carries. + Lookups are filtered to the version the run is calling now, so entries + written under an earlier one are never served, and re-fetching a query + overwrites its row rather than leaving an obsolete one behind for good. + + Each provider owns a table named for it, so a lookup reads only the entries + that could answer it and the tables of providers never called cost nothing. + + A cache is never a source of truth: deleting a cache file changes nothing + but how many API calls a run costs. + """ + + COMMIT_INTERVAL = 250 + LOOKUP_CHUNK = 500 + + def __init__(self, api: str, version: str, path: Optional[str] = None, read_paths: Sequence[str] = ()): + self._api = api + self._version = version + self._writer = self._open_writer(path) if path else None + self._readers = [self._open_reader(read_path) for read_path in read_paths] + self._uncommitted = 0 + + @property + def _table(self) -> str: + """Returns this provider's table name, quoted for use in a statement.""" + escaped = self._api.replace('"', '""') + return f'"{escaped}"' + + def _open_writer(self, path: str) -> sqlite3.Connection: + """ + Opens the read-write cache, creating the file and this provider's table when absent. + + Parameters + ---------- + path : str + The path of the cache file to open or create. + + Return + ---------- + sqlite3.Connection + A connection with this provider's table in place. + + Raises + ---------- + ValueError + If the path names an existing file that is not a SQLite database. + """ + connection = sqlite3.connect(path) + try: + connection.execute(SCHEMA.format(table=self._table)) + connection.commit() + except sqlite3.DatabaseError as error: + connection.close() + raise ValueError(f"'{path}' cannot be used as a cache: {error}") from error + return connection + + def _open_reader(self, path: str) -> sqlite3.Connection: + """ + Opens a cache file read-only and confirms it can answer this provider's lookups. + + A file named as a read-only cache is one the run expects to save calls, + so a file holding nothing for this provider is refused here rather than + quietly costing the calls it was meant to spare. The columns lookups + read are checked as well as the table itself, so a foreign database that + happens to hold a table of that name is refused too. + + Parameters + ---------- + path : str + The path of an existing cache file. + + Return + ---------- + sqlite3.Connection + A connection that cannot modify the file. + + Raises + ---------- + ValueError + If the file is not a SQLite database, or holds no usable table for + this provider. + """ + connection = sqlite3.connect(f"{Path(path).resolve().as_uri()}?mode=ro", uri=True) + try: + columns = {row[1] for row in connection.execute(f"PRAGMA table_info({self._table})")} + except sqlite3.DatabaseError as error: + connection.close() + raise ValueError(f"'{path}' is not a readable SQLite database: {error}") from error + + if not columns: + connection.close() + raise ValueError(f"'{path}' holds no cached '{self._api}' responses") + + missing = [column for column in REQUIRED_COLUMNS if column not in columns] + if missing: + connection.close() + raise ValueError(f"'{path}' is not a geocoder cache (its '{self._api}' table is missing {', '.join(missing)})") + return connection + + def _connections(self) -> Iterator[sqlite3.Connection]: + """Yields the writable cache first, then each read-only cache in the order given.""" + if self._writer is not None: + yield self._writer + yield from self._readers + + def lookup(self, keys: Iterable[str]) -> Dict[str, Any]: + """ + Fetches the stored responses for the given keys, writable cache first. + + The first file holding a key under the current version wins; keys with + no such response are absent from the result, whether they were never + stored or were stored under a version this run no longer calls. + + Parameters + ---------- + keys : Iterable[str] + The query keys to look for. + + Return + ---------- + Dict[str, Any] + Each found key mapped to its decoded response. + """ + found: Dict[str, Any] = {} + outstanding = list(dict.fromkeys(keys)) + + for connection in self._connections(): + if not outstanding: + break + for start in range(0, len(outstanding), self.LOOKUP_CHUNK): + chunk = outstanding[start : start + self.LOOKUP_CHUNK] + placeholders = ",".join("?" * len(chunk)) + rows = connection.execute( + f"SELECT query, response FROM {self._table} WHERE version = ? AND query IN ({placeholders})", + [self._version, *chunk], + ) + for key, response in rows: + found[key] = json.loads(response) + outstanding = [key for key in outstanding if key not in found] + + return found + + def store(self, key: str, response: Any) -> None: + """ + Records one provider response, committing once a batch has accumulated. + + The response is tagged with the version that produced it, replacing any + entry the query already had, so an obsolete response is retired by the + call that supersedes it. + + Committing as the run proceeds means a crash costs only the calls made + since the last commit. + + Parameters + ---------- + key : str + The query key to store the response under. + response : Any + The raw provider response, stored as JSON. + """ + if self._writer is None: + return + + fetched_at = datetime.now(timezone.utc).isoformat(timespec="seconds") + self._writer.execute( + f"INSERT OR REPLACE INTO {self._table} (query, version, response, fetched_at) VALUES (?, ?, ?, ?)", + (key, self._version, json.dumps(response, ensure_ascii=False), fetched_at), + ) + + self._uncommitted += 1 + if self._uncommitted >= self.COMMIT_INTERVAL: + self.commit() + + def commit(self) -> None: + """Flushes any stored responses that have not yet been committed.""" + if self._writer is not None and self._uncommitted: + self._writer.commit() + self._uncommitted = 0 + + def close(self) -> None: + """Commits outstanding writes and closes every open cache file.""" + self.commit() + for connection in self._connections(): + connection.close() + self._writer = None + self._readers = [] + + def __enter__(self) -> "Cache": + """Returns the cache so it can be used as a context manager.""" + return self + + def __exit__(self, *_exception: object) -> None: + """Closes the cache when the context exits.""" + self.close() + + +def missing_files(paths: Sequence[str]) -> List[str]: + """Returns the subset of the given paths that are not existing files.""" + return [path for path in paths if not Path(path).is_file()] diff --git a/src/census.py b/src/census.py index 9d50fbf..d71cda8 100644 --- a/src/census.py +++ b/src/census.py @@ -7,7 +7,7 @@ import csv import io -from typing import Dict, List +from typing import Any, Dict, Iterator, List, Tuple import requests @@ -24,7 +24,9 @@ class CensusProvider(Provider): The benchmark and vintage are pinned to the frozen Census2020 dataset so the positional response layout stays stable; the addressbatch endpoint returns - headerless CSV, so columns can only be read by their fixed position. + headerless CSV, so columns can only be read by their fixed position. That + pinned pair is also the version cached responses are tagged with, since + moving to another dataset asks a different question of the service. """ requires_key = False @@ -33,6 +35,7 @@ class CensusProvider(Provider): ENDPOINT = "https://geocoding.geo.census.gov/geocoder/geographies/addressbatch" BENCHMARK = "Public_AR_Census2020" VINTAGE = "Census2020_Census2020" + CACHE_VERSION = f"{BENCHMARK}/{VINTAGE}" BATCH_SIZE = 10000 TIMEOUT = 300 @@ -51,9 +54,12 @@ class CensusProvider(Provider): "block", ] - def geocode(self, records: List[SourceRecord]) -> List[GeocodeResult]: + def _fetch(self, records: List[SourceRecord]) -> Iterator[Tuple[SourceRecord, Dict[str, Any]]]: """ - Geocodes records in batches and returns one result per record, in order. + Posts records in batches, yielding each record with its response row. + + Yielding per batch rather than per run means an interrupted job keeps + every batch that already came back. Parameters ---------- @@ -62,23 +68,24 @@ def geocode(self, records: List[SourceRecord]) -> List[GeocodeResult]: Return ---------- - List[GeocodeResult] - One result per input record, aligned by position. + Iterator[Tuple[SourceRecord, Dict[str, Any]]] + Each record paired with its raw Census response row. """ - results_by_key: Dict[int, GeocodeResult] = {} for start in range(0, len(records), self.BATCH_SIZE): - self._geocode_batch(records[start : start + self.BATCH_SIZE], results_by_key) - - return [results_by_key.get(record.internal_key, GeocodeResult(match_notes="No match")) for record in records] + yield from self._fetch_batch(records[start : start + self.BATCH_SIZE]) - def _geocode_batch(self, batch: List[SourceRecord], results_by_key: Dict[int, GeocodeResult]) -> None: - """Posts one CSV batch, verifies its row count, and stores each result by internal key.""" + def _fetch_batch(self, batch: List[SourceRecord]) -> Iterator[Tuple[SourceRecord, Dict[str, Any]]]: + """Posts one CSV batch, verifies its row count, and matches rows back by internal key.""" response = self._post_batch(batch) rows = [row for row in csv.reader(io.StringIO(response.text)) if row] if len(rows) != len(batch): raise ValueError(f"Census returned {len(rows)} rows for {len(batch)} submitted records; the service or benchmark may have changed") + + records_by_key = {record.internal_key: record for record in batch} for row in rows: - results_by_key[int(row[0])] = self._parse_row(row) + record = records_by_key.get(int(row[0])) + if record is not None: + yield record, self._build_raw(row) def _post_batch(self, batch: List[SourceRecord]) -> requests.Response: """Posts one CSV batch through the retrying request helper.""" @@ -109,27 +116,50 @@ def _build_csv(batch: List[SourceRecord]) -> str: ) return buffer.getvalue() - def _parse_row(self, row: List[str]) -> GeocodeResult: + def _build_raw(self, row: List[str]) -> Dict[str, Any]: + """ + Names the positional fields of one response row, validating a match first. + + The pinned layout is checked here, where the row arrives from the network, + and the echoed id is dropped as it is only meaningful within this run. + """ + status = row[2] if len(row) > 2 else "No_Match" + if status == "Match": + self._check_layout(row) + + raw = dict(zip(self.RESPONSE_FIELDS, row)) + raw.pop("id", None) + return raw + + def parse(self, raw: Dict[str, Any]) -> GeocodeResult: """ Converts one Census response row into a normalized GeocodeResult. Accuracy is graded from the populated result fields rather than the Census match type, so it reflects how specific the returned location is. + + Parameters + ---------- + raw : Dict[str, Any] + One named Census response row. + + Return + ---------- + GeocodeResult + The normalized, graded result that row describes. """ - result = self._build_result(row) + result = self._build_result(raw) result.accuracy = grade_accuracy(result, self.MAX_ACCURACY) return result - def _build_result(self, row: List[str]) -> GeocodeResult: + def _build_result(self, raw: Dict[str, Any]) -> GeocodeResult: """Maps a Census response row to a result without scoring its accuracy.""" - raw = dict(zip(self.RESPONSE_FIELDS, row)) - status = row[2] if len(row) > 2 else "No_Match" + status = raw.get("match_status", "No_Match") if status == "Match": - self._check_layout(row) - exact = row[3].strip().lower() == "exact" - address, city, stateprov, postalcode = self._split_address(row[4]) - longitude, latitude = self._split_coordinates(row[5]) + exact = raw.get("match_type", "").strip().lower() == "exact" + address, city, stateprov, postalcode = self._split_address(raw.get("matched_address", "")) + longitude, latitude = self._split_coordinates(raw.get("coordinates", "")) return GeocodeResult( result_id=raw.get("tigerline_id", ""), result_address=address, diff --git a/src/geocoder.py b/src/geocoder.py index 2bf410e..1a6b9a7 100644 --- a/src/geocoder.py +++ b/src/geocoder.py @@ -12,7 +12,7 @@ import json import os from dataclasses import dataclass, replace -from typing import Any, Dict, Iterable, Iterator, List, Optional, Sequence, Tuple +from typing import Any, Dict, Iterable, Iterator, List, Optional, Sequence, Tuple, Type from openpyxl import Workbook, load_workbook @@ -25,6 +25,7 @@ SourceRecord, resolve_api_key, ) +from .cache import DEFAULT_CACHE_FILE, Cache, missing_files from .flags import format_flags from .postprocess import MATCH_HEADERS, compare_record from .preprocess import BLANK_CHECKED_FIELDS, PRE_HEADERS, QUERY_FIELDS, check_records @@ -760,6 +761,26 @@ def parse_args(argv: Optional[List[str]] = None) -> argparse.Namespace: action="store_true", help="use the worksheet name as COUNTRY when a row's country is blank", ) + parser.add_argument( + "--cache", + default=None, + metavar="FILE", + help=f"SQLite cache of prior API responses, read and written " f"(default: {DEFAULT_CACHE_FILE} in the working directory)", + ) + parser.add_argument( + "--cacheRead", + dest="cache_read", + action="append", + default=[], + metavar="FILE", + help="existing SQLite cache to read but never write; repeatable", + ) + parser.add_argument( + "--noCache", + dest="no_cache", + action="store_true", + help="open no cache file, ignoring --cache and --cacheRead; repeated addresses within the run are still collapsed", + ) parser.add_argument( "--preProcess", dest="preprocess", @@ -780,6 +801,49 @@ def parse_args(argv: Optional[List[str]] = None) -> argparse.Namespace: return parser.parse_args(argv) +def open_cache(args: argparse.Namespace, provider_cls: Type[Provider]) -> Cache: + """ + Opens the cache files the given provider will read and write. + + A cache answers one provider calling one version of its API, so the + provider selects both the table read and the version tag written. + + Parameters + ---------- + args + The parsed command-line arguments. + provider_cls : Type[Provider] + The provider class whose responses the cache will hold. + + Return + ---------- + Cache + The cache to hand the provider; one backed by no file when --noCache + was given. + + Raises + ---------- + SystemExit + If a read-only cache file is missing, or a named file cannot be used + as a cache. + """ + if args.no_cache: + return Cache(provider_cls.name, provider_cls.CACHE_VERSION) + + absent = missing_files(args.cache_read) + if absent: + raise SystemExit(f"error: read-only cache file(s) do not exist: {', '.join(absent)}") + + path = args.cache or DEFAULT_CACHE_FILE + try: + cache = Cache(provider_cls.name, provider_cls.CACHE_VERSION, path, args.cache_read) + except ValueError as error: + raise SystemExit(f"error: {error}") from error + + print(f"using cache {os.path.abspath(path)}") + return cache + + def select_provider(args: argparse.Namespace) -> Optional[Provider]: """ Builds the provider named on the command line, or None when it is 'none'. @@ -797,7 +861,8 @@ def select_provider(args: argparse.Namespace) -> Optional[Provider]: Raises ---------- SystemExit - If the api is unknown, or a required API key is not configured. + If the api is unknown, a required API key is not configured, or a named + cache file cannot be used. """ if args.api == NO_API: return None @@ -813,7 +878,7 @@ def select_provider(args: argparse.Namespace) -> Optional[Provider]: if not api_key: raise SystemExit(f"error: api '{args.api}' requires an API key " f"(pass --apiKey or set {KEY_ENV_VARS[args.api]})") - return provider_cls(api_key) + return provider_cls(api_key, open_cache(args, provider_cls)) def main(argv: Optional[List[str]] = None) -> None: @@ -832,9 +897,9 @@ def main(argv: Optional[List[str]] = None) -> None: Raises ---------- SystemExit - If the input file is missing, the output file already exists, the api is - unknown, a required API key is not configured, or no provider and no - check was asked for. + If the input file is missing, the output file already exists, a read-only + cache is missing or unusable, the api is unknown, a required API key is + not configured, or no provider and no check was asked for. """ load_dotenv() args = parse_args(argv) @@ -854,10 +919,14 @@ def main(argv: Optional[List[str]] = None) -> None: compare=args.compare, debug=args.debug, ) - if provider is None and holds_output(args.infile, args.worksheet): - recheck_workbook(args.infile, args.outfile, worksheet=args.worksheet, options=options) - else: - process_workbook(args.infile, args.outfile, provider, worksheet=args.worksheet, options=options) + try: + if provider is None and holds_output(args.infile, args.worksheet): + recheck_workbook(args.infile, args.outfile, worksheet=args.worksheet, options=options) + else: + process_workbook(args.infile, args.outfile, provider, worksheet=args.worksheet, options=options) + finally: + if provider is not None: + provider.cache.close() print(f"wrote {args.outfile}") diff --git a/src/geocodio.py b/src/geocodio.py index 3094c3c..291eb4c 100644 --- a/src/geocodio.py +++ b/src/geocodio.py @@ -5,11 +5,11 @@ Geocoder — Geocodio provider """ -from typing import Dict, List +from typing import Any, Dict, Iterator, List, Tuple import requests -from .api import AccuracyLevel, GeocodeResult, Provider, SourceRecord, apply_legal_land_limit, format_street_address, grade_accuracy, register +from .api import AccuracyLevel, GeocodeResult, Provider, SourceRecord, format_street_address, grade_accuracy, register @register("geocodio") @@ -47,8 +47,9 @@ class GeocodioProvider(Provider): """ requires_key = True + CACHE_VERSION = "v1.7" - ENDPOINT = "https://api.geocod.io/v1.7/geocode" + ENDPOINT = f"https://api.geocod.io/{CACHE_VERSION}/geocode" BATCH_SIZE = 10000 TIMEOUT = 600 @@ -66,9 +67,12 @@ class GeocodioProvider(Provider): EXACT_TYPES = {"rooftop", "point"} - def geocode(self, records: List[SourceRecord]) -> List[GeocodeResult]: + def _fetch(self, records: List[SourceRecord]) -> Iterator[Tuple[SourceRecord, Dict[str, Any]]]: """ - Geocodes records in batches and returns one result per record, in order. + Posts records in batches, yielding each record with its response entry. + + Yielding per batch rather than per run means an interrupted job keeps + every batch that already came back. Parameters ---------- @@ -77,24 +81,18 @@ def geocode(self, records: List[SourceRecord]) -> List[GeocodeResult]: Return ---------- - List[GeocodeResult] - One result per input record, aligned by position. + Iterator[Tuple[SourceRecord, Dict[str, Any]]] + Each record paired with its raw Geocodio response entry. """ - results_by_key: Dict[int, GeocodeResult] = {} - queryable = [record for record in records if record.address_string()] - for start in range(0, len(queryable), self.BATCH_SIZE): - self._geocode_batch(queryable[start : start + self.BATCH_SIZE], results_by_key) - - results = [results_by_key.get(record.internal_key, self.unqueryable_result()) for record in records] - return [apply_legal_land_limit(record, result) for record, result in zip(records, results)] + for start in range(0, len(records), self.BATCH_SIZE): + yield from self._fetch_batch(records[start : start + self.BATCH_SIZE]) - def _geocode_batch(self, batch: List[SourceRecord], results_by_key: Dict[int, GeocodeResult]) -> None: - """Posts one batch, verifies its entry count, and stores each result by internal key.""" + def _fetch_batch(self, batch: List[SourceRecord]) -> Iterator[Tuple[SourceRecord, Dict[str, Any]]]: + """Posts one batch, verifies its entry count, and matches entries back by position.""" entries = self._post_batch(batch) if len(entries) != len(batch): raise ValueError(f"Geocodio returned {len(entries)} entries for {len(batch)} submitted records; the batch response may have changed") - for record, entry in zip(batch, entries): - results_by_key[record.internal_key] = self._parse_entry(entry) + yield from zip(batch, entries) def _post_batch(self, batch: List[SourceRecord]) -> List[Dict]: """Posts one batch through the retrying request helper and returns the ordered result entries.""" @@ -109,23 +107,30 @@ def _post_batch(self, batch: List[SourceRecord]) -> List[Dict]: ) return response.json().get("results", []) - def _parse_entry(self, entry: Dict) -> GeocodeResult: - """Grades the best candidate for one input, or returns a no-match when none was found.""" - matches = entry.get("response", {}).get("results", []) - if not matches: - return GeocodeResult(match_notes="No match", raw=entry) - return self._parse_result(matches[0]) - - def _parse_result(self, match: Dict) -> GeocodeResult: + def parse(self, raw: Dict[str, Any]) -> GeocodeResult: """ - Converts one Geocodio candidate into a normalized GeocodeResult. + Converts one Geocodio response entry into a normalized GeocodeResult. The graded accuracy is capped at the tier implied by ``accuracy_type`` so an interpolated or centroid match cannot report rooftop precision on the strength of the echoed address fields. + + Parameters + ---------- + raw : Dict[str, Any] + One Geocodio response entry. + + Return + ---------- + GeocodeResult + The normalized, graded result that entry describes. """ - components = match.get("address_components", {}) - result = self._build_result(match, components) + matches = raw.get("response", {}).get("results", []) + if not matches: + return GeocodeResult(match_notes="No match", raw=raw) + + components = matches[0].get("address_components", {}) + result = self._build_result(raw, components) result.accuracy = grade_accuracy(result, self._accuracy_cap(result, components)) return result @@ -142,8 +147,15 @@ def _accuracy_cap(self, result: GeocodeResult, components: Dict[str, str]) -> Ac return cap return min(cap, AccuracyLevel.STREET) - def _build_result(self, match: Dict, components: Dict[str, str]) -> GeocodeResult: - """Maps a Geocodio candidate to a GeocodeResult without scoring its accuracy.""" + def _build_result(self, raw: Dict[str, Any], components: Dict[str, str]) -> GeocodeResult: + """ + Maps the best Geocodio candidate to a GeocodeResult without scoring its accuracy. + + The whole response entry is kept as the raw value rather than the single + candidate it was read from, so what is cached and reported is what + Geocodio actually said. + """ + match = raw["response"]["results"][0] location = match.get("location", {}) accuracy_type = match.get("accuracy_type", "") @@ -157,7 +169,7 @@ def _build_result(self, match: Dict, components: Dict[str, str]) -> GeocodeResul longitude=str(location.get("lng", "")), match_type="exact" if accuracy_type in self.EXACT_TYPES else "non-exact", location_type=accuracy_type, - raw=match, + raw=raw, ) @staticmethod diff --git a/src/google.py b/src/google.py index 103b3d1..d856420 100644 --- a/src/google.py +++ b/src/google.py @@ -5,11 +5,11 @@ Geocoder — Google Geocoding API provider """ -from typing import Dict, List +from typing import Any, Dict, Iterator, List, Tuple import requests -from .api import AccuracyLevel, GeocodeResult, Provider, SourceRecord, apply_legal_land_limit, format_street_address, grade_accuracy, register +from .api import AccuracyLevel, GeocodeResult, Provider, SourceRecord, format_street_address, grade_accuracy, register @register("google") @@ -39,9 +39,14 @@ class GoogleProvider(Provider): and matches it to an unrelated road. Whatever ordinary place name the row also carries is still sent, but the result is capped at the province the parcel sits in. + + The endpoint carries no version of its own, so cached responses are tagged + with one of ours, to be raised whenever a change here would make Google + answer differently. """ requires_key = True + CACHE_VERSION = "v1" ENDPOINT = "https://maps.googleapis.com/maps/api/geocode/json" TIMEOUT = 30 @@ -64,9 +69,13 @@ class GoogleProvider(Provider): "country": "result_country", } - def geocode(self, records: List[SourceRecord]) -> List[GeocodeResult]: + def _fetch(self, records: List[SourceRecord]) -> Iterator[Tuple[SourceRecord, Dict[str, Any]]]: """ - Geocodes each record with a single request and returns results in order. + Queries one address per request, yielding each response as it arrives. + + A status other than a match or an empty result means the request itself + failed — a rejected key or an exhausted quota — so it is raised here rather + than yielded, keeping a failure out of the cache and off the output. Parameters ---------- @@ -75,25 +84,20 @@ def geocode(self, records: List[SourceRecord]) -> List[GeocodeResult]: Return ---------- - List[GeocodeResult] - One result per input record, aligned by position. - """ - return [apply_legal_land_limit(record, self._geocode_one(record)) for record in records] - - def _geocode_one(self, record: SourceRecord) -> GeocodeResult: - """Queries one address and grades the parsed response by its location_type.""" - query = record.address_string() - if not query: - return self.unqueryable_result() - - payload = self._request(query) - status = payload.get("status", "UNKNOWN") + Iterator[Tuple[SourceRecord, Dict[str, Any]]] + Each record paired with its raw Google response. - if status == "OK": - return self._parse_result(payload["results"][0]) - if status == "ZERO_RESULTS": - return GeocodeResult(match_notes="No match", raw=payload) - raise ValueError(f"Google geocoding failed with status {status!r}: {payload.get('error_message', '')}".strip()) + Raises + ---------- + ValueError + If Google reports a status other than OK or ZERO_RESULTS. + """ + for record in records: + payload = self._request(record.address_string()) + status = payload.get("status", "UNKNOWN") + if status not in ("OK", "ZERO_RESULTS"): + raise ValueError(f"Google geocoding failed with status {status!r}: {payload.get('error_message', '')}".strip()) + yield record, payload def _request(self, address: str) -> Dict: """Sends one geocoding request through the retrying request helper.""" @@ -106,16 +110,29 @@ def _request(self, address: str) -> Dict: ) return response.json() - def _parse_result(self, match: Dict) -> GeocodeResult: + def parse(self, raw: Dict[str, Any]) -> GeocodeResult: """ - Converts one Google result into a normalized GeocodeResult. + Converts one Google response into a normalized GeocodeResult. - The graded accuracy is capped at the tier implied by ``location_type`` so - an interpolated or centroid match cannot report rooftop precision on the - strength of the echoed address fields. + The graded accuracy is bounded by _accuracy_cap so an interpolated match, + an area centroid, or a route with no street number cannot report rooftop + precision on the strength of the echoed address fields. + + Parameters + ---------- + raw : Dict[str, Any] + One Google response envelope. + + Return + ---------- + GeocodeResult + The normalized, graded result that response describes. """ - components = self._extract_components(match) - result = self._build_result(match, components) + if raw.get("status") != "OK": + return GeocodeResult(match_notes="No match", raw=raw) + + components = self._extract_components(raw["results"][0]) + result = self._build_result(raw, components) result.accuracy = grade_accuracy(result, self._accuracy_cap(result, components)) return result @@ -132,8 +149,15 @@ def _accuracy_cap(self, result: GeocodeResult, components: Dict[str, str]) -> Ac return cap return min(cap, AccuracyLevel.STREET) - def _build_result(self, match: Dict, components: Dict[str, str]) -> GeocodeResult: - """Maps a Google result to a GeocodeResult without scoring its accuracy.""" + def _build_result(self, raw: Dict[str, Any], components: Dict[str, str]) -> GeocodeResult: + """ + Maps the best Google match to a GeocodeResult without scoring its accuracy. + + The whole response envelope is kept as the raw value rather than the single + match it was read from, so what is cached and reported is what Google + actually said. + """ + match = raw["results"][0] geometry = match.get("geometry", {}) location = geometry.get("location", {}) @@ -148,7 +172,7 @@ def _build_result(self, match: Dict, components: Dict[str, str]) -> GeocodeResul longitude=str(location.get("lng", "")), match_type="partial" if match.get("partial_match") else "exact", location_type=geometry.get("location_type", ""), - raw=match, + raw=raw, ) def _extract_components(self, match: Dict) -> Dict[str, str]: diff --git a/test/api_test.py b/test/api_test.py index d103611..8c1058c 100644 --- a/test/api_test.py +++ b/test/api_test.py @@ -27,9 +27,13 @@ def test_register_adds_to_registry(): @register("temp_provider") class _Temp(Provider): - def geocode(self, records): - """Returns no results; the class only exercises registration.""" - return [] + def _fetch(self, records): + """Yields nothing; the class only exercises registration.""" + return iter(()) + + def parse(self, raw): + """Returns an empty result; the class only exercises registration.""" + return GeocodeResult(raw=raw) try: assert api.PROVIDERS["temp_provider"] is _Temp @@ -240,14 +244,22 @@ def raise_for_status(self): class _RetryProvider(Provider): - """Routes a single retried request through geocode for the retry tests.""" + """Routes a single retried request through the provider for the retry tests.""" def __init__(self, send): super().__init__() self.send = send - def geocode(self, records): - """Returns the response from one retried request, ignoring records.""" + def _fetch(self, records): + """Yields nothing; the class only exercises the retry helper.""" + return iter(()) + + def parse(self, raw): + """Returns an empty result; the class only exercises the retry helper.""" + return GeocodeResult(raw=raw) + + def send_once(self): + """Returns the response from one retried request.""" return self._request_with_retry(self.send) @@ -262,7 +274,7 @@ def send(): raise api.requests.ConnectionError("connection reset") return _FakeResponse("ok") - response = _RetryProvider(send).geocode([]) + response = _RetryProvider(send).send_once() assert attempts["count"] == 3 assert response.text == "ok" @@ -278,5 +290,5 @@ def send(): raise api.requests.ConnectionError("connection reset") with pytest.raises(api.requests.ConnectionError): - _RetryProvider(send).geocode([]) + _RetryProvider(send).send_once() assert attempts["count"] == Provider.MAX_ATTEMPTS diff --git a/test/cache_test.py b/test/cache_test.py new file mode 100644 index 0000000..10eadfd --- /dev/null +++ b/test/cache_test.py @@ -0,0 +1,250 @@ +#!/usr/bin/python3 +# -.- coding: utf-8 -.- +# -.- dependencies: Python 3.8+ -.- + +""" +Geocoder Cache Tests + +Copyright (c) 2026 Pangaea Information Technologies, Ltd. +""" + +import sqlite3 + +import pytest + +from src.api import GeocodeResult, Provider, SourceRecord, cache_key +from src.cache import Cache, missing_files + + +class _CountingProvider(Provider): + """Counts the records it is asked to fetch so cache hits can be observed.""" + + name = "counting" + CACHE_VERSION = "v1" + + def __init__(self, cache=None): + super().__init__(cache=cache) + self.fetched = [] + + def _fetch(self, records): + """Yields a response per record and records which records were requested.""" + for record in records: + self.fetched.append(record.address_string()) + yield record, {"echo": record.address_string()} + + def parse(self, raw): + """Returns a result carrying the echoed query so hits can be identified.""" + return GeocodeResult(result_address=raw["echo"], raw=raw) + + +def _record(key, address, city="Town", stateprov="CA"): + """Builds a SourceRecord with the given internal key and address.""" + return SourceRecord(internal_key=key, address=address, city=city, stateprov=stateprov) + + +def _cache(path=None, read_paths=()): + """Opens a cache scoped to the counting provider under test.""" + return Cache(_CountingProvider.name, _CountingProvider.CACHE_VERSION, path, read_paths) + + +def test_cache_key_folds_case_and_whitespace(): + """Case and runs of whitespace collapse so equivalent rows share one key.""" + assert cache_key(_record(0, "123 Main St")) == cache_key(_record(1, "123 MAIN ST")) + assert cache_key(SourceRecord(internal_key=0, address=" 1 A St\t")) == "1 a st" + + +def test_cache_key_keeps_distinct_addresses_distinct(): + """Normalization does not merge addresses that differ in substance.""" + assert cache_key(_record(0, "1 Main St")) != cache_key(_record(1, "1 Main Ave")) + + +def test_cache_key_is_empty_for_a_row_with_nothing_to_query(): + """A row left with no address to send has no query, and so no key.""" + assert cache_key(SourceRecord(internal_key=0, address="01-17-040-06w4")) == "" + + +def test_duplicates_collapse_to_one_call_without_a_cache_file(): + """Repeated addresses in one run cost a single call even with no cache configured.""" + provider = _CountingProvider() + records = [ + _record(0, "1 Main St"), + _record(1, "1 MAIN ST"), + _record(2, "2 Oak St"), + _record(3, "1 Main St"), + ] + + results = provider.geocode(records) + + assert len(provider.fetched) == 2 + assert len(results) == 4 + assert results[0].result_address == results[1].result_address == results[3].result_address + assert results[2].result_address == "2 Oak St, Town, CA" + + +def test_cache_hit_across_runs_skips_the_api(tmp_path): + """A second run over an overlapping list only calls the API for the new addresses.""" + path = str(tmp_path / "cache.sqlite") + + with _cache(path) as cache: + first = _CountingProvider(cache) + first.geocode([_record(0, "1 Main St"), _record(1, "2 Oak St")]) + assert len(first.fetched) == 2 + + with _cache(path) as cache: + second = _CountingProvider(cache) + results = second.geocode([_record(0, "1 Main St"), _record(1, "3 Elm St")]) + + assert second.fetched == ["3 Elm St, Town, CA"] + assert results[0].result_address == "1 Main St, Town, CA" + assert results[1].result_address == "3 Elm St, Town, CA" + + +def test_cache_is_scoped_per_api(tmp_path): + """One provider's stored response is never served to another provider.""" + path = str(tmp_path / "cache.sqlite") + + with _cache(path) as cache: + _CountingProvider(cache).geocode([_record(0, "1 Main St")]) + + with Cache("different", "v1", path) as cache: + other = _CountingProvider(cache) + other.geocode([_record(0, "1 Main St")]) + + assert other.fetched == ["1 Main St, Town, CA"] + + +def test_cache_is_scoped_per_api_version(tmp_path): + """An entry stored under an earlier version of the API is never served.""" + path = str(tmp_path / "cache.sqlite") + + with _cache(path) as cache: + _CountingProvider(cache).geocode([_record(0, "1 Main St")]) + + with Cache("counting", "v2", path) as cache: + later = _CountingProvider(cache) + later.geocode([_record(0, "1 Main St")]) + + assert later.fetched == ["1 Main St, Town, CA"] + + +def test_refetching_replaces_the_entry_of_an_earlier_version(tmp_path): + """A query re-asked under a new version overwrites its row instead of accumulating one.""" + path = str(tmp_path / "cache.sqlite") + + with _cache(path) as cache: + _CountingProvider(cache).geocode([_record(0, "1 Main St")]) + with Cache("counting", "v2", path) as cache: + _CountingProvider(cache).geocode([_record(0, "1 Main St")]) + + assert sqlite3.connect(path).execute("SELECT version FROM counting").fetchall() == [("v2",)] + + +def test_read_only_cache_is_used_but_not_written(tmp_path): + """A --cacheRead file answers lookups and is left untouched by the run.""" + shared = tmp_path / "shared.sqlite" + with _cache(str(shared)) as cache: + _CountingProvider(cache).geocode([_record(0, "1 Main St")]) + + before = shared.read_bytes() + + with _cache(str(tmp_path / "own.sqlite"), [str(shared)]) as cache: + provider = _CountingProvider(cache) + provider.geocode([_record(0, "1 Main St"), _record(1, "2 Oak St")]) + + assert provider.fetched == ["2 Oak St, Town, CA"] + assert shared.read_bytes() == before + + +def test_writable_cache_wins_over_read_only(tmp_path): + """The writable cache answers first when both files hold the same key.""" + stale = str(tmp_path / "stale.sqlite") + fresh = str(tmp_path / "fresh.sqlite") + key = cache_key(_record(0, "1 Main St")) + + with _cache(stale) as cache: + cache.store(key, {"echo": "from stale"}) + with _cache(fresh) as cache: + cache.store(key, {"echo": "from fresh"}) + + with _cache(fresh, [stale]) as cache: + result = _CountingProvider(cache).geocode([_record(0, "1 Main St")])[0] + + assert result.result_address == "from fresh" + + +def test_store_commits_in_batches_for_crash_safety(tmp_path): + """Responses reach disk before the run ends so a crash keeps completed calls.""" + path = str(tmp_path / "cache.sqlite") + cache = _cache(path) + for index in range(Cache.COMMIT_INTERVAL): + cache.store(f"key {index}", {"echo": index}) + + committed = sqlite3.connect(path).execute("SELECT COUNT(*) FROM counting").fetchone()[0] + cache.close() + + assert committed == Cache.COMMIT_INTERVAL + + +def test_stored_row_records_the_version_and_the_fetch_time(tmp_path): + """Each entry carries the key it answers, the version behind it, and when it arrived.""" + path = str(tmp_path / "cache.sqlite") + with _cache(path) as cache: + _CountingProvider(cache).geocode([_record(0, "1 MAIN St")]) + + row = sqlite3.connect(path).execute("SELECT query, version, fetched_at FROM counting").fetchone() + assert row[0] == "1 main st, town, ca" + assert row[1] == "v1" + assert row[2] + + +def test_read_only_cache_rejects_a_database_holding_nothing_for_the_provider(tmp_path): + """A file that cannot answer this provider is refused instead of quietly saving nothing.""" + path = str(tmp_path / "other.sqlite") + connection = sqlite3.connect(path) + connection.execute("CREATE TABLE unrelated (id INTEGER)") + connection.commit() + connection.close() + + with pytest.raises(ValueError): + _cache(None, [path]) + + +def test_read_only_cache_rejects_a_differently_shaped_table(tmp_path): + """A table named for the provider without the columns lookups read is refused up front.""" + path = str(tmp_path / "other.sqlite") + connection = sqlite3.connect(path) + connection.execute("CREATE TABLE counting (key TEXT, value TEXT)") + connection.commit() + connection.close() + + with pytest.raises(ValueError): + _cache(None, [path]) + + +def test_read_only_cache_rejects_a_non_database(tmp_path): + """A file that is not a SQLite database is refused with a clear error.""" + path = tmp_path / "notes.txt" + path.write_text("not a database") + + with pytest.raises(ValueError): + _cache(None, [str(path)]) + + +def test_writable_cache_rejects_a_non_database(tmp_path): + """Pointing --cache at an existing non-database file is refused, not overwritten.""" + path = tmp_path / "addresses.xlsx" + path.write_text("not a database") + + with pytest.raises(ValueError): + _cache(str(path)) + + assert path.read_text() == "not a database" + + +def test_missing_files_reports_only_absent_paths(tmp_path): + """Existing paths are dropped and absent ones are reported in order.""" + present = tmp_path / "here.sqlite" + present.write_text("") + absent = str(tmp_path / "gone.sqlite") + + assert missing_files([str(present), absent]) == [absent] diff --git a/test/census_test.py b/test/census_test.py index 0dfa443..f38ba82 100644 --- a/test/census_test.py +++ b/test/census_test.py @@ -159,3 +159,46 @@ def fake_post(*_args, **_kwargs): ] with pytest.raises(ValueError): census.CensusProvider().geocode(records) + + +def test_census_cache_version_pins_the_dataset(): + """Cached entries are tagged with the dataset that answered them.""" + assert census.CensusProvider.BENCHMARK in census.CensusProvider.CACHE_VERSION + assert census.CensusProvider.VINTAGE in census.CensusProvider.CACHE_VERSION + + +def test_census_posts_each_distinct_address_once(monkeypatch): + """Repeated addresses are collapsed so the posted CSV carries one row each.""" + posted = [] + + def fake_post(_url, files=None, **_kwargs): + posted.append(files["addressFile"][1]) + keys = [line.split(",")[0] for line in files["addressFile"][1].splitlines() if line] + return _FakeResponse("".join(f'"{key}","q","No_Match"\r\n' for key in keys)) + + monkeypatch.setattr(census.requests, "post", fake_post) + + records = [ + SourceRecord(internal_key=0, address="1 Main St", city="Town", stateprov="CA"), + SourceRecord(internal_key=1, address="1 MAIN ST", city="Town", stateprov="CA"), + SourceRecord(internal_key=2, address="2 Oak St", city="Town", stateprov="CA"), + ] + results = census.CensusProvider().geocode(records) + + assert len(posted) == 1 + assert len([line for line in posted[0].splitlines() if line]) == 2 + assert len(results) == 3 + + +def test_census_raw_drops_the_run_local_id(monkeypatch): + """The echoed internal key is not kept, since it means nothing outside its run.""" + + def fake_post(*_args, **_kwargs): + return _FakeResponse('"0","1 Main St, Town, CA","Match","Exact","1 MAIN ST, TOWN, CA, 90210","-118.0,34.0","1","L","06","037","1","1"\r\n') + + monkeypatch.setattr(census.requests, "post", fake_post) + + result = census.CensusProvider().geocode([SourceRecord(internal_key=0, address="1 Main St", city="Town", stateprov="CA")])[0] + + assert "id" not in result.raw + assert result.raw["match_status"] == "Match" diff --git a/test/geocoder_test.py b/test/geocoder_test.py index f86cea1..3fd6144 100644 --- a/test/geocoder_test.py +++ b/test/geocoder_test.py @@ -16,6 +16,7 @@ from src import postprocess from src import preprocess from src.api import GeocodeResult, Provider, SourceRecord +from src.cache import DEFAULT_CACHE_FILE from src.geocoder import Options, detect_columns, main, process_workbook, write_output_sheet @@ -25,26 +26,33 @@ class MockProvider(Provider): name = "mock" requires_key = False - def geocode(self, records): - """Returns a fixed result per record, echoing the source address fields.""" - results = [] + def _fetch(self, records): + """Yields a canned response per record, echoing the source address fields.""" for record in records: - results.append( - GeocodeResult( - result_address=record.address, - result_city=record.city, - result_stateprov=record.stateprov, - result_country=record.country, - latitude="40.0", - longitude="-75.0", - match_type="exact", - accuracy=100, - location_type="rooftop", - match_notes="", - raw={"key": record.internal_key, "q": record.address_string()}, - ) - ) - return results + yield record, { + "key": record.internal_key, + "q": record.address_string(), + "address": record.address, + "city": record.city, + "stateprov": record.stateprov, + "country": record.country, + } + + def parse(self, raw): + """Rebuilds the echoed result from a raw response.""" + return GeocodeResult( + result_address=raw["address"], + result_city=raw["city"], + result_stateprov=raw["stateprov"], + result_country=raw["country"], + latitude="40.0", + longitude="-75.0", + match_type="exact", + accuracy=100, + location_type="rooftop", + match_notes="", + raw=raw, + ) class CoarseProvider(MockProvider): @@ -319,13 +327,59 @@ def test_main_runs_registered_provider(tmp_path): api.PROVIDERS["mock"] = MockProvider try: - main([str(infile), str(outfile), "--api", "mock"]) + main([str(infile), str(outfile), "--api", "mock", "--noCache"]) finally: del api.PROVIDERS["mock"] assert outfile.exists() +def test_main_defaults_to_a_cache_in_the_working_directory(tmp_path, monkeypatch): + """With no cache flag the default file is created beside the working directory.""" + infile = tmp_path / "in.xlsx" + _make_workbook(infile, {"S": [["Address", "City", "State"], ["1 A St", "Town", "CA"]]}) + monkeypatch.chdir(tmp_path) + + api.PROVIDERS["mock"] = MockProvider + try: + main([str(infile), "out.xlsx", "--api", "mock"]) + finally: + del api.PROVIDERS["mock"] + + assert (tmp_path / DEFAULT_CACHE_FILE).is_file() + + +def test_main_no_cache_writes_nothing(tmp_path, monkeypatch): + """--noCache leaves no cache file behind, including the default one.""" + infile = tmp_path / "in.xlsx" + _make_workbook(infile, {"S": [["Address", "City", "State"], ["1 A St", "Town", "CA"]]}) + monkeypatch.chdir(tmp_path) + + api.PROVIDERS["mock"] = MockProvider + try: + main([str(infile), "out.xlsx", "--api", "mock", "--noCache"]) + finally: + del api.PROVIDERS["mock"] + + assert not (tmp_path / DEFAULT_CACHE_FILE).exists() + + +def test_main_no_cache_ignores_the_cache_files_named(tmp_path, monkeypatch): + """--noCache wins over --cache and --cacheRead rather than being an error.""" + infile = tmp_path / "in.xlsx" + _make_workbook(infile, {"S": [["Address", "City", "State"], ["1 A St", "Town", "CA"]]}) + monkeypatch.chdir(tmp_path) + + api.PROVIDERS["mock"] = MockProvider + try: + main([str(infile), "out.xlsx", "--api", "mock", "--cache", "c.sqlite", "--cacheRead", "gone.sqlite", "--noCache"]) + finally: + del api.PROVIDERS["mock"] + + assert not (tmp_path / "c.sqlite").exists() + assert not (tmp_path / DEFAULT_CACHE_FILE).exists() + + def test_preprocess_adds_flag_columns(tmp_path): """--preProcess inserts PRE_ columns between the source and result columns.""" infile = tmp_path / "in.xlsx" @@ -588,6 +642,59 @@ def test_country_per_sheet_reports_no_missing_country(tmp_path, capsys): assert not dict(zip(header, [cell.value for cell in sheet[2]]))["PRE_FLAGS"] +def test_main_cache_flag_persists_across_runs(tmp_path): + """--cache creates the file, and a rerun of the same input calls nothing.""" + cache_path = tmp_path / "cache.sqlite" + infile = tmp_path / "in.xlsx" + _make_workbook(infile, {"S": [["Address", "City", "State"], ["1 A St", "Town", "CA"], ["1 A St", "Town", "CA"]]}) + + calls = [] + + class _CountingMock(MockProvider): + """Records every record handed to the API so calls can be counted.""" + + def _fetch(self, records): + """Counts the fetched records before delegating to the mock response.""" + calls.extend(record.address_string() for record in records) + yield from super()._fetch(records) + + api.PROVIDERS["mock"] = _CountingMock + try: + main([str(infile), str(tmp_path / "a.xlsx"), "--api", "mock", "--cache", str(cache_path)]) + assert calls == ["1 A St, Town, CA"] + + main([str(infile), str(tmp_path / "b.xlsx"), "--api", "mock", "--cache", str(cache_path)]) + finally: + del api.PROVIDERS["mock"] + + assert calls == ["1 A St, Town, CA"] + assert cache_path.is_file() + + +def test_main_missing_read_cache_exits(tmp_path): + """A --cacheRead file that does not exist exits before any provider work.""" + infile = tmp_path / "in.xlsx" + _make_workbook(infile, {"S": [["Address", "City", "State"], ["1 A St", "Town", "CA"]]}) + + with pytest.raises(SystemExit): + main([str(infile), str(tmp_path / "out.xlsx"), "--api", "mock", "--cacheRead", str(tmp_path / "gone.sqlite")]) + + +def test_main_unusable_read_cache_exits(tmp_path): + """A --cacheRead file that is not a geocoder cache exits with a message.""" + infile = tmp_path / "in.xlsx" + _make_workbook(infile, {"S": [["Address", "City", "State"], ["1 A St", "Town", "CA"]]}) + foreign = tmp_path / "foreign.sqlite" + foreign.write_text("not a database") + + api.PROVIDERS["mock"] = MockProvider + try: + with pytest.raises(SystemExit): + main([str(infile), str(tmp_path / "out.xlsx"), "--api", "mock", "--cache", str(tmp_path / "c.sqlite"), "--cacheRead", str(foreign)]) + finally: + del api.PROVIDERS["mock"] + + def test_main_runs_census_provider(tmp_path, monkeypatch): """--api census runs end to end and writes the census result columns.""" infile = tmp_path / "in.xlsx" @@ -607,7 +714,7 @@ def fake_post(*_args, **_kwargs): monkeypatch.setattr(census.requests, "post", fake_post) - main([str(infile), str(outfile), "--api", "census"]) + main([str(infile), str(outfile), "--api", "census", "--noCache"]) sheet = openpyxl.load_workbook(outfile)["S"] header = [cell.value for cell in sheet[1]] diff --git a/test/geocodio_test.py b/test/geocodio_test.py index 959a3c0..7ad5972 100644 --- a/test/geocodio_test.py +++ b/test/geocodio_test.py @@ -237,6 +237,29 @@ def test_geocodio_aligns_entries_by_position(monkeypatch): assert results[2].accuracy == 40 +def test_geocodio_posts_each_distinct_address_once(monkeypatch): + """Repeated addresses are collapsed so the posted batch carries one entry each.""" + posted = [] + + def fake_post(_url, json=None, **_kwargs): + posted.append(json) + return FakeResponse(_batch(*[[_candidate("rooftop", ROOFTOP_COMPONENTS)] for _ in json])) + + monkeypatch.setattr(geocodio.requests, "post", fake_post) + + records = [ + SourceRecord(internal_key=0, address="1 Main St", city="Town", stateprov="CA"), + SourceRecord(internal_key=1, address="1 MAIN ST", city="Town", stateprov="CA"), + SourceRecord(internal_key=2, address="2 Oak St", city="Town", stateprov="CA"), + ] + results = geocodio.GeocodioProvider("key").geocode(records) + + assert len(posted) == 1 + assert len(posted[0]) == 2 + assert len(results) == 3 + assert all(result.accuracy == 100 for result in results) + + def test_geocodio_raises_on_entry_count_mismatch(monkeypatch): """A response short of one entry per record fails loudly instead of misaligning.""" _patch_response(monkeypatch, _batch([_candidate("rooftop", ROOFTOP_COMPONENTS)])) diff --git a/test/google_test.py b/test/google_test.py index 4275802..4941671 100644 --- a/test/google_test.py +++ b/test/google_test.py @@ -10,7 +10,8 @@ import pytest from src import google -from src.api import LEGAL_LAND_NOTE, NO_ADDRESS_NOTE, PROVIDERS, SourceRecord +from src.api import LEGAL_LAND_NOTE, NO_ADDRESS_NOTE, PROVIDERS, SourceRecord, cache_key +from src.cache import Cache def _result(location_type, components, **overrides): @@ -234,6 +235,53 @@ def test_google_raises_on_error_status(monkeypatch): google.GoogleProvider("key").geocode([SourceRecord(internal_key=0, address="1 Main St")]) +def test_google_raw_is_the_whole_response(monkeypatch): + """The stored raw value is the full envelope, not just the match read from it.""" + _patch_response(monkeypatch, _result("ROOFTOP", ROOFTOP_COMPONENTS)) + + result = google.GoogleProvider("key").geocode([SourceRecord(internal_key=0, address="1600 Pennsylvania Ave NW")])[0] + + assert result.raw["status"] == "OK" + assert result.raw["results"][0]["place_id"] == "PLACE" + + +def test_google_requests_each_distinct_address_once(monkeypatch): + """Records sharing an address cost one request and all receive the result.""" + requested = [] + + def fake_get(_url, params=None, **_kwargs): + requested.append(params["address"]) + return FakeResponse(_result("ROOFTOP", ROOFTOP_COMPONENTS)) + + monkeypatch.setattr(google.requests, "get", fake_get) + + records = [ + SourceRecord(internal_key=0, address="1 Main St", city="Town", stateprov="CA"), + SourceRecord(internal_key=1, address="1 MAIN ST", city="Town", stateprov="CA"), + SourceRecord(internal_key=2, address="2 Oak St", city="Town", stateprov="CA"), + ] + results = google.GoogleProvider("key").geocode(records) + + assert len(requested) == 2 + assert len(results) == 3 + assert all(result.accuracy == 100 for result in results) + + +def test_google_error_status_is_never_cached(monkeypatch, tmp_path): + """A failed request leaves nothing behind, so a later run retries it.""" + path = str(tmp_path / "cache.sqlite") + version = google.GoogleProvider.CACHE_VERSION + _patch_response(monkeypatch, {"status": "OVER_QUERY_LIMIT", "error_message": "quota exceeded"}) + + record = SourceRecord(internal_key=0, address="1 Main St") + with Cache("google", version, path) as cache: + with pytest.raises(ValueError): + google.GoogleProvider("key", cache).geocode([record]) + + with Cache("google", version, path) as cache: + assert not cache.lookup([cache_key(record)]) + + def test_google_withholds_legal_land_description(monkeypatch): """A rig row asks Google only for its province and is capped there.""" captured = {}