Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import matplotlib.pyplot as plt

from rivretrieve import UKFetcher, constants
from rivretrieve import UKEAFetcher, constants

gauge_ids = [
"http://environment.data.gov.uk/hydrology/id/stations/3c5cba29-2321-4289-a1fd-c355e135f4cb",
Expand All @@ -11,7 +11,7 @@

plt.figure(figsize=(12, 6))

fetcher = UKFetcher()
fetcher = UKEAFetcher()
for gauge_id in gauge_ids:
print(f"Fetching data for {gauge_id}...")
data = fetcher.get_data(gauge_id=gauge_id, variable=variable, start_date=start_date, end_date=end_date)
Expand Down
2 changes: 1 addition & 1 deletion rivretrieve/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from .poland import PolandFetcher
from .slovenia import SloveniaFetcher
from .southafrica import SouthAfricaFetcher
from .uk import UKFetcher
from .uk_ea import UKEAFetcher
from .uk_nrfa import UKNRFAFetcher
from .usa import USAFetcher

Expand Down
9,195 changes: 9,195 additions & 0 deletions rivretrieve/cached_site_data/uk_ea_sites.csv

Large diffs are not rendered by default.

905 changes: 0 additions & 905 deletions rivretrieve/cached_site_data/uk_sites.csv

This file was deleted.

45 changes: 41 additions & 4 deletions rivretrieve/uk.py → rivretrieve/uk_ea.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,20 +12,58 @@
logger = logging.getLogger(__name__)


class UKFetcher(base.RiverDataFetcher):
class UKEAFetcher(base.RiverDataFetcher):
"""Fetches river gauge data from the UK Environment Agency."""

BASE_URL = "http://environment.data.gov.uk"

METADATA_TRANSLATION_MAPPING = {
"notation": constants.GAUGE_ID,
"stationReference": "stationReference",
"label": constants.STATION_NAME,
"lat": constants.LATITUDE,
"long": constants.LONGITUDE,
"riverName": constants.RIVER,
"catchmentArea": constants.AREA,
}

@staticmethod
def get_gauge_ids() -> pd.DataFrame:
"""Retrieves a DataFrame of available UK gauge IDs."""
return utils.load_sites_csv("uk")
return utils.load_sites_csv("uk_ea")

@staticmethod
def get_available_variables() -> tuple[str, ...]:
return (constants.DISCHARGE, constants.STAGE)

def get_metadata(self) -> pd.DataFrame:
"""Fetches site metadata for all stations from the EA API.

Returns:
A pandas DataFrame indexed by gauge_id, containing site metadata.
"""
params = {"_limit": 10000}
url = f"{self.BASE_URL}/hydrology/id/stations.json"
s = utils.requests_retry_session()
try:
response = s.get(url, params=params)
response.raise_for_status()
data = response.json()
df = pd.DataFrame(data.get("items", []))

if df.empty:
return pd.DataFrame().set_index(constants.GAUGE_ID)

df = df.rename(columns=self.METADATA_TRANSLATION_MAPPING)

return df.set_index(constants.GAUGE_ID)
except requests.exceptions.RequestException as e:
logger.error(f"Error fetching EA stations list: {e}")
raise
except Exception as e:
logger.error(f"Error processing EA stations list: {e}")
raise

def _get_measure_notation(self, variable: str) -> str:
"""Gets the notation for the given variable."""
if variable == constants.STAGE:
Expand All @@ -38,10 +76,9 @@ def _get_measure_notation(self, variable: str) -> str:
def _download_data(self, gauge_id: str, variable: str, start_date: str, end_date: str) -> List[Dict[str, Any]]:
"""Downloads the raw data from the UK Environment Agency API."""
notation = self._get_measure_notation(variable)
site = gauge_id.split("/")[-1]

# Check if the station has data for the given variable
measure_url = f"{self.BASE_URL}/hydrology/id/measures?station={site}"
measure_url = f"{self.BASE_URL}/hydrology/id/measures?station={gauge_id}"
try:
r = utils.requests_retry_session().get(measure_url)
r.raise_for_status()
Expand Down
71 changes: 0 additions & 71 deletions tests/test_uk.py

This file was deleted.

141 changes: 141 additions & 0 deletions tests/test_uk_ea.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,141 @@
import json
import os
import unittest
from pathlib import Path
from unittest.mock import MagicMock, patch

import pandas as pd
from pandas.testing import assert_frame_equal

from rivretrieve import UKEAFetcher, constants


class TestUKEAFetcher(unittest.TestCase):
def setUp(self):
self.fetcher = UKEAFetcher()
self.test_data_dir = Path(os.path.dirname(__file__)) / "test_data"
self.measures_file = self.test_data_dir / "uk_measures.json"
self.readings_file = self.test_data_dir / "uk_readings_discharge.json"

if not self.measures_file.exists():
raise FileNotFoundError(f"Measures file not found at {self.measures_file}")
if not self.readings_file.exists():
raise FileNotFoundError(f"Readings file not found at {self.readings_file}")

def load_sample_json(self, filename):
with open(filename, "r", encoding="utf-8") as f:
return json.load(f)

@patch("rivretrieve.utils.requests_retry_session")
def test_get_data_discharge(self, mock_requests_session):
mock_session = MagicMock()
mock_requests_session.return_value = mock_session

mock_measures_response = MagicMock()
mock_measures_response.json.return_value = self.load_sample_json(self.measures_file)
mock_measures_response.raise_for_status = MagicMock()

mock_readings_response = MagicMock()
mock_readings_response.json.return_value = self.load_sample_json(self.readings_file)
mock_readings_response.raise_for_status = MagicMock()

def mock_get_side_effect(url, *args, **kwargs):
if "measures?station" in url:
return mock_measures_response
elif "readings" in url:
return mock_readings_response
return MagicMock()

mock_session.get.side_effect = mock_get_side_effect

gauge_id = "http://environment.data.gov.uk/hydrology/id/stations/3c5cba29-2321-4289-a1fd-c355e135f4cb"
variable = constants.DISCHARGE
start_date = "2024-01-01"
end_date = "2024-01-03"

result_df = self.fetcher.get_data(gauge_id, variable, start_date, end_date)

expected_dates = pd.to_datetime(["2024-01-01", "2024-01-02", "2024-01-03"])
expected_values = [72.777, 99.138, 68.020] # Values from sample file
expected_data = {
constants.TIME_INDEX: expected_dates,
constants.DISCHARGE: expected_values,
}
expected_df = pd.DataFrame(expected_data).set_index(constants.TIME_INDEX)

assert_frame_equal(result_df, expected_df, check_dtype=False)
self.assertEqual(mock_session.get.call_count, 2)

@patch("rivretrieve.utils.requests_retry_session")
def test_get_metadata(self, mock_requests_session):
mock_session = MagicMock()
mock_requests_session.return_value = mock_session

sample_metadata = {
"items": [
{
"@id": "http://environment.data.gov.uk/hydrology/id/stations/12345",
"notation": "12345",
"stationReference": "STREF001",
"label": "Test Station 1",
"lat": 51.5,
"long": -0.1,
"riverName": "Test River",
"catchmentArea": 1234,
"easting": 500000,
"northing": 180000,
"other_col": "some_value",
},
{
"@id": "http://environment.data.gov.uk/hydrology/id/stations/67890",
"notation": "67890",
"stationReference": "STREF002",
"label": "Test Station 2",
"lat": 52.1,
"long": -0.5,
"riverName": "Another River",
"catchmentArea": 4567,
"easting": 480000,
"northing": 200000,
"other_col": None, # Match the None for check_like
},
]
}

mock_response = MagicMock()
mock_response.json.return_value = sample_metadata
mock_response.raise_for_status = MagicMock()
mock_session.get.return_value = mock_response

result_df = self.fetcher.get_metadata()

expected_data = {
constants.GAUGE_ID: [
"12345",
"67890",
],
"@id": [
"http://environment.data.gov.uk/hydrology/id/stations/12345",
"http://environment.data.gov.uk/hydrology/id/stations/67890",
],
"stationReference": ["STREF001", "STREF002"],
constants.STATION_NAME: ["Test Station 1", "Test Station 2"],
constants.LATITUDE: [51.5, 52.1],
constants.LONGITUDE: [-0.1, -0.5],
constants.RIVER: ["Test River", "Another River"],
constants.AREA: [1234, 4567],
"easting": [500000, 480000],
"northing": [180000, 200000],
"other_col": ["some_value", None],
}
expected_df = pd.DataFrame(expected_data).set_index(constants.GAUGE_ID)

assert_frame_equal(result_df, expected_df, check_like=True)
mock_session.get.assert_called_once()
mock_args, mock_kwargs = mock_session.get.call_args
self.assertIn("/hydrology/id/stations.json", mock_args[0])
self.assertEqual(mock_kwargs["params"], {"_limit": 10000})


if __name__ == "__main__":
unittest.main()