diff --git a/create_trace_mapping.py b/create_trace_mapping.py index 2b8ef57..1dc449a 100644 --- a/create_trace_mapping.py +++ b/create_trace_mapping.py @@ -1,3 +1,5 @@ +from pathlib import Path + import yaml from generator_to_trace_draft_mapper import ( draft_solar_generator_to_trace_mapping, @@ -19,14 +21,14 @@ solar_generator_mapping = draft_solar_generator_to_trace_mapping( solar_gens, solar_traces ) -with open("draft_solar_generator_mapping.yaml", "w") as file: +with Path("draft_solar_generator_mapping.yaml").open("w") as file: yaml.dump(solar_generator_mapping, file, default_flow_style=False) solar_traces = "/media/nick/Samsung_T5/isp_2024_data/trace_data/solar/solar_2023" rezs = gets_rezs(workbook) solar_rez_mapping = draft_solar_rez_mapping(rezs, solar_traces) -with open("solar_area_mapping.yaml", "w") as file: +with Path("solar_area_mapping.yaml").open("w") as file: yaml.dump(solar_rez_mapping, file, default_flow_style=False) duids_and_station_names = static_table( @@ -47,12 +49,12 @@ wind_generator_mapping = draft_wind_generator_to_trace_mapping( wind_gens, wind_duids_and_station_names, wind_traces ) -with open("draft_wind_generator_mapping.yaml", "w") as file: +with Path("draft_wind_generator_mapping.yaml").open("w") as file: yaml.dump(wind_generator_mapping, file, default_flow_style=False, sort_keys=False) wind_traces = "D:/isp_2024_data/trace_data/wind/wind_2023" rezs = gets_rezs(workbook) wind_rez_mapping = draft_wind_rez_mapping(rezs, wind_traces) -with open("draft_wind_rez_mapping.yaml", "w") as file: +with Path("draft_wind_rez_mapping.yaml").open("w") as file: yaml.dump(wind_rez_mapping, file, default_flow_style=False) diff --git a/src/isp_trace_parser/__init__.py b/src/isp_trace_parser/__init__.py index 648baeb..0101b45 100644 --- a/src/isp_trace_parser/__init__.py +++ b/src/isp_trace_parser/__init__.py @@ -15,13 +15,13 @@ from isp_trace_parser.wind_traces import WindMetadataFilter, parse_wind_traces __all__ = [ - "trace_formatter", + "DemandMetadataFilter", + "SolarMetadataFilter", + "WindMetadataFilter", + "construct_reference_year_mapping", "get_data", - "parse_wind_traces", "parse_demand_traces", "parse_solar_traces", - "construct_reference_year_mapping", - "WindMetadataFilter", - "SolarMetadataFilter", - "DemandMetadataFilter", + "parse_wind_traces", + "trace_formatter", ] diff --git a/src/isp_trace_parser/construct_reference_year_mapping.py b/src/isp_trace_parser/construct_reference_year_mapping.py index 95b9299..b79ba76 100644 --- a/src/isp_trace_parser/construct_reference_year_mapping.py +++ b/src/isp_trace_parser/construct_reference_year_mapping.py @@ -13,7 +13,7 @@ @validate_call def construct_reference_year_mapping( start_year: int, end_year: int, reference_years: list[int] -): +) -> dict: """Constructs a dictionary mapping a sequence of modeling years to a cycle of reference years. Examples: @@ -42,4 +42,4 @@ def construct_reference_year_mapping( reference_years = ( reference_years * full_reference_year_cycles ) + reference_years[:partial_cycle_length] - return dict(zip(years, reference_years)) + return dict(zip(years, reference_years, strict=True)) diff --git a/src/isp_trace_parser/demand_trace_metadata.py b/src/isp_trace_parser/demand_trace_metadata.py index 7a1eff9..2b77896 100644 --- a/src/isp_trace_parser/demand_trace_metadata.py +++ b/src/isp_trace_parser/demand_trace_metadata.py @@ -29,7 +29,8 @@ def build( refyear, _, dimensions_suffix = after.partition("_") key = (location_prefix, dimensions_suffix) if not refyear.isdigit() or key not in lookup: - raise ValueError(f"Unexpected trace filename: {path.name}") + msg = f"Unexpected trace filename: {path.name}" + raise ValueError(msg) file_metadata[path] = {**lookup[key], "reference_year": int(refyear)} return file_metadata diff --git a/src/isp_trace_parser/demand_traces.py b/src/isp_trace_parser/demand_traces.py index 56b71d8..4a0a690 100644 --- a/src/isp_trace_parser/demand_traces.py +++ b/src/isp_trace_parser/demand_traces.py @@ -8,7 +8,7 @@ import functools import os from pathlib import Path -from typing import Literal, Optional +from typing import Literal import polars as pl from joblib import Parallel, delayed @@ -52,15 +52,16 @@ class DemandMetadataFilter(BaseModel): reference_year: list of ints specifying reference_years """ - subregion: Optional[list[str]] = None - scenario: Optional[ + subregion: list[str] | None = None + scenario: ( list[Literal["Step Change", "Progressive Change", "Green Energy Exports"]] - ] = None - poe: Optional[list[Literal["POE50", "POE10"]]] = None - demand_type: Optional[ - list[Literal["OPSO_MODELLING", "OPSO_MODELLING_PVLITE", "PV_TOT"]] - ] = None - reference_year: Optional[list[int]] = None + | None + ) = None + poe: list[Literal["POE50", "POE10"]] | None = None + demand_type: ( + list[Literal["OPSO_MODELLING", "OPSO_MODELLING_PVLITE", "PV_TOT"]] | None + ) = None + reference_year: list[int] | None = None @validate_call @@ -69,7 +70,7 @@ def parse_demand_traces( parsed_directory: str | Path, use_concurrency: bool = True, filters: DemandMetadataFilter | None = None, -): +) -> None: """Takes a directory with AEMO demand trace data and reformats the data, saving it to a new directory. AEMO demand trace data comes in CSVs with columns specifying the year, day, and month, and data columns diff --git a/src/isp_trace_parser/get_data.py b/src/isp_trace_parser/get_data.py index e058aa4..1cf4f91 100644 --- a/src/isp_trace_parser/get_data.py +++ b/src/isp_trace_parser/get_data.py @@ -7,7 +7,7 @@ import datetime from pathlib import Path -from typing import List, Literal +from typing import Literal import pandas as pd import polars as pl @@ -16,7 +16,7 @@ def _year_range_to_dt_range( start_year: int, end_year: int, year_type: Literal["fy", "calendar"] = "fy" -): +) -> datetime.datetime: """ Convert year range to datetime boundaries for efficient time filtering. @@ -44,10 +44,11 @@ def _year_range_to_dt_range( end_year, 7, 1 ) - elif year_type == "calendar": + if year_type == "calendar": return datetime.datetime(start_year, 1, 1), datetime.datetime( end_year + 1, 1, 1 ) + raise ValueError(year_type) def _query_parquet_single_reference_year( @@ -55,8 +56,8 @@ def _query_parquet_single_reference_year( end_year: int, reference_year: int, directory: str | Path, - filters: dict[str, any] = None, - select_columns: list[str] = None, + filters: dict[str, any] | None = None, + select_columns: list[str] | None = None, year_type: Literal["fy", "calendar"] = "fy", ) -> pd.DataFrame: """ @@ -143,8 +144,7 @@ def _query_parquet_multiple_reference_years( start_year=year, end_year=year, reference_year=reference_year, **kwargs ) ) - data = pd.concat(data).reset_index(drop=True) - return data + return pd.concat(data).reset_index(drop=True) @validate_call @@ -152,11 +152,11 @@ def get_project_single_reference_year( start_year: int, end_year: int, reference_year: int, - project: str | List, + project: str | list, directory: str | Path, year_type: Literal["fy", "calendar"] = "fy", - select_columns: list[str] = None, -): + select_columns: list[str] | None = None, +) -> pd.DataFrame: """ Query project trace data for a single reference year. @@ -244,12 +244,12 @@ def get_zone_single_reference_year( start_year: int, end_year: int, reference_year: int, - zone: str | List, - resource_type: str | List, + zone: str | list, + resource_type: str | list, directory: str | Path, year_type: Literal["fy", "calendar"] = "fy", - select_columns: list[str] = None, -): + select_columns: list[str] | None = None, +) -> pd.DataFrame: """ Query zone trace data for a single reference year. @@ -340,14 +340,14 @@ def get_demand_single_reference_year( start_year: int, end_year: int, reference_year: int, - scenario: str | List, - subregion: str | List, - demand_type: str | List, - poe: str | List, + scenario: str | list, + subregion: str | list, + demand_type: str | list, + poe: str | list, directory: str | Path, year_type: Literal["fy", "calendar"] = "fy", - select_columns: list[str] = None, -): + select_columns: list[str] | None = None, +) -> pd.DataFrame: """ Query demand trace data for a single reference year. @@ -448,11 +448,11 @@ def get_demand_single_reference_year( @validate_call def get_project_multiple_reference_years( reference_year_mapping: dict[int, int], - project: str | List, + project: str | list, directory: str | Path, year_type: Literal["fy", "calendar"] = "fy", - select_columns: list[str] = None, -): + select_columns: list[str] | None = None, +) -> pd.DataFrame: """ Query project trace data across multiple reference years. @@ -537,12 +537,12 @@ def get_project_multiple_reference_years( @validate_call def get_zone_multiple_reference_years( reference_year_mapping: dict[int, int], - zone: str | List, - resource_type: str | List, + zone: str | list, + resource_type: str | list, directory: str | Path, year_type: Literal["fy", "calendar"] = "fy", - select_columns: list[str] = None, -): + select_columns: list[str] | None = None, +) -> pd.DataFrame: """ Query zone trace data across multiple reference years. @@ -630,14 +630,14 @@ def get_zone_multiple_reference_years( @validate_call def get_demand_multiple_reference_years( reference_year_mapping: dict[int, int], - scenario: str | List, - subregion: str | List, - demand_type: str | List, - poe: str | List, + scenario: str | list, + subregion: str | list, + demand_type: str | list, + poe: str | list, directory: str | Path, year_type: Literal["fy", "calendar"] = "fy", - select_columns: list[str] = None, -): + select_columns: list[str] | None = None, +) -> pd.DataFrame: """ Query demand trace data across multiple reference years. diff --git a/src/isp_trace_parser/input_validation.py b/src/isp_trace_parser/input_validation.py index 6c696a3..a56cc44 100644 --- a/src/isp_trace_parser/input_validation.py +++ b/src/isp_trace_parser/input_validation.py @@ -11,7 +11,7 @@ def input_directory(path: Path | str) -> Path: path = is_valid_path(path) if not path.is_dir(): - raise ValueError(f"Directory {path} does not exist") + raise FileNotFoundError(path) return path @@ -22,10 +22,12 @@ def parsed_directory(path: str | Path) -> Path: def is_valid_path(path: str | Path) -> Path: try: return Path(path) - except (TypeError, ValueError): - raise ValueError(f"Invalid parsed directory path: {path}") + except (TypeError, ValueError) as exc: + msg = f"Invalid parsed directory path: {path}" + raise ValueError(msg) from exc -def start_year_before_end_year(start_year, end_year): +def start_year_before_end_year(start_year: int, end_year: int) -> None: if end_year < start_year: - raise ValueError(f"Start year {end_year} < end year {start_year}") + msg = f"Start year {end_year} < end year {start_year}" + raise ValueError(msg) diff --git a/src/isp_trace_parser/optimise_parquet.py b/src/isp_trace_parser/optimise_parquet.py index 9ced190..a549377 100644 --- a/src/isp_trace_parser/optimise_parquet.py +++ b/src/isp_trace_parser/optimise_parquet.py @@ -7,7 +7,6 @@ from itertools import product from pathlib import Path -from typing import Optional import duckdb from pydantic import validate_call @@ -30,7 +29,7 @@ def partition_traces_by_columns( input_directory: str | Path, output_directory: str | Path, partition_cols: list[str], - sort_by: Optional[list[str]] = ["datetime"], + sort_by: list[str] | None = None, ) -> None: """Partition parquet traces by specified columns with optional sorting. @@ -60,6 +59,12 @@ def partition_traces_by_columns( ... partition_cols=["scenario", "reference_year"] ... ) # doctest: +SKIP """ + + if sort_by is None: + # Avoid use of mutable data structure for argument defaults + # (see Ruff rule B006). + sort_by = ["datetime"] + output_path = Path(output_directory) output_path.mkdir(parents=True, exist_ok=True) @@ -71,23 +76,21 @@ def partition_traces_by_columns( values = con.execute(f""" SELECT DISTINCT {col} FROM read_parquet('{input_directory}') - """).fetchall() + """).fetchall() # noqa: S608 distinct_values.append(values) partitions = [tuple(val[0] for val in vals) for vals in product(*distinct_values)] for partition_values in partitions: - # print(*partition_values) - conditions = [] - for col, val in zip(partition_cols, partition_values): + for col, val in zip(partition_cols, partition_values, strict=True): if isinstance(val, str): conditions.append(f"{col}='{val}'") else: conditions.append(f"{col}={val}") where_clause = " AND ".join(conditions) - query = f"SELECT * FROM read_parquet('{input_directory}') WHERE {where_clause}" + query = f"SELECT * FROM read_parquet('{input_directory}') WHERE {where_clause}" # noqa: S608 if sort_by: query += f" ORDER BY {', '.join(sort_by)}" diff --git a/src/isp_trace_parser/remote/download.py b/src/isp_trace_parser/remote/download.py index 17c751f..d16bfbd 100644 --- a/src/isp_trace_parser/remote/download.py +++ b/src/isp_trace_parser/remote/download.py @@ -61,14 +61,15 @@ def _download_from_manifest( manifest_path = files("isp_trace_parser.remote.manifests") / f"{manifest_name}.txt" if not manifest_path.exists(): - raise FileNotFoundError(f"Manifest file not found: {manifest_path}") + raise FileNotFoundError(manifest_path) # Read URLs from manifest - with open(manifest_path) as f: + with Path(manifest_path).open("r") as f: urls = [line.strip() for line in f if line.strip()] if not urls: - raise ValueError(f"No URLs found in manifest: {manifest_path}") + msg = f"No URLs found in manifest: {manifest_path}" + raise ValueError(msg) save_directory = Path(save_directory) @@ -88,12 +89,13 @@ def _download_with_retry( for attempt in range(max_retries): try: _download_file(url, save_directory, strip_levels, unquote_path) - return - except requests.exceptions.RequestException: + except requests.exceptions.RequestException: # noqa: PERF203 if attempt < max_retries - 1: time.sleep(2**attempt) else: raise + else: + return def _download_file( @@ -131,10 +133,11 @@ def _download_file( # Strip specified number of directory levels path_parts = url_path.split("/") if strip_levels >= len(path_parts): - raise ValueError( + msg = ( f"Cannot strip {strip_levels} levels from path with only " f"{len(path_parts)} parts: {url_path}" ) + raise ValueError(msg) stripped_path = "/".join(path_parts[strip_levels:]) destination = save_directory / stripped_path @@ -151,7 +154,7 @@ def _download_file( # Write file with progress bar with ( - open(destination, "wb") as f, + Path(destination).open("wb") as f, tqdm( total=total_size, unit="B", @@ -218,17 +221,16 @@ def fetch_trace_data( # Validate inputs if dataset_type not in ["full", "example"]: - raise ValueError( - f"dataset_type must be 'full' or 'example', got: {dataset_type}" - ) + msg = f"dataset_type must be 'full' or 'example', got: {dataset_type}" + raise ValueError(msg) if dataset_src != "isp_2024": - raise ValueError(f"Only isp_2024 is currently supported, got: {dataset_src}") + msg = f"Only isp_2024 is currently supported, got: {dataset_src}" + raise ValueError(msg) if data_format not in ["processed", "archive"]: - raise ValueError( - f"data_format must be 'processed' or 'archive', got: {data_format}" - ) + msg = f"data_format must be 'processed' or 'archive', got: {data_format}" + raise ValueError(msg) # Construct manifest name and download manifest_name = f"{data_format}/{dataset_type}_{dataset_src}" diff --git a/src/isp_trace_parser/resource_trace_metadata.py b/src/isp_trace_parser/resource_trace_metadata.py index f5ffa95..8cedc7a 100644 --- a/src/isp_trace_parser/resource_trace_metadata.py +++ b/src/isp_trace_parser/resource_trace_metadata.py @@ -40,7 +40,8 @@ def build( for path in files: stem, sep, ref = path.stem.rpartition("_RefYear") if not sep or not ref.isdigit() or stem not in resource_mapping: - raise ValueError(f"Unexpected trace filename: {path.name}") + msg = f"Unexpected trace filename: {path.name}" + raise ValueError(msg) entry = resource_mapping[stem] file_metadata[path] = { "name": entry["location"], diff --git a/src/isp_trace_parser/solar_traces.py b/src/isp_trace_parser/solar_traces.py index aa8dbde..f858fdb 100644 --- a/src/isp_trace_parser/solar_traces.py +++ b/src/isp_trace_parser/solar_traces.py @@ -8,7 +8,7 @@ import functools import os from pathlib import Path -from typing import Literal, Optional +from typing import Literal from joblib import Parallel, delayed from pydantic import BaseModel, validate_call @@ -56,10 +56,10 @@ class SolarMetadataFilter(BaseModel): reference_year: list of ints specifying reference_years """ - name: Optional[list[str]] = None - file_type: Optional[list[Literal["zone", "project"]]] = None - resource_type: Optional[list[Literal["SAT", "FFP", "CST"]]] = None - reference_year: Optional[list[int]] = None + name: list[str] | None = None + file_type: list[Literal["zone", "project"]] | None = None + resource_type: list[Literal["SAT", "FFP", "CST"]] | None = None + reference_year: list[int] | None = None @validate_call @@ -68,7 +68,7 @@ def parse_solar_traces( parsed_directory: str | Path, use_concurrency: bool = True, filters: SolarMetadataFilter | None = None, -): +) -> None: """Takes a directory with AEMO solar trace data and reformats the data, saving it to a new directory. AEMO solar trace data comes in CSVs with columns specifying the year, day, and month, and data columns @@ -164,7 +164,7 @@ def parse_solar_traces( } project_and_zone_output_names, project_and_zone_input_names = zip( - *name_mappings.items() + *name_mappings.items(), strict=True ) partial_func = functools.partial( @@ -179,12 +179,12 @@ def parse_solar_traces( Parallel(n_jobs=max_workers)( delayed(partial_func)(save_name, old_trace_name) for save_name, old_trace_name in zip( - project_and_zone_output_names, project_and_zone_input_names + project_and_zone_output_names, project_and_zone_input_names, strict=True ) ) else: for save_name, old_trace_name in zip( - project_and_zone_output_names, project_and_zone_input_names + project_and_zone_output_names, project_and_zone_input_names, strict=True ): partial_func(save_name, old_trace_name) @@ -283,7 +283,7 @@ def get_unique_resource_types_in_metadata( A list of unique resource types. """ return list( - set(metadata["resource_type"] for metadata in metadata_for_trace_files.values()) + {metadata["resource_type"] for metadata in metadata_for_trace_files.values()} ) diff --git a/src/isp_trace_parser/trace_formatter.py b/src/isp_trace_parser/trace_formatter.py index c889520..c9ec713 100644 --- a/src/isp_trace_parser/trace_formatter.py +++ b/src/isp_trace_parser/trace_formatter.py @@ -72,10 +72,10 @@ def trace_formatter(trace_data: pl.DataFrame) -> pl.DataFrame: value_name="value", ) - def get_hour(time_label): + def get_hour(time_label: str) -> timedelta: return timedelta(hours=int(time_label) // 2) - def get_minute(time_label): + def get_minute(time_label: str) -> timedelta: return timedelta(minutes=int(time_label) % 2 * 30) trace_data = trace_data.with_columns( diff --git a/src/isp_trace_parser/trace_restructure_helper_functions.py b/src/isp_trace_parser/trace_restructure_helper_functions.py index 8d9b43b..0a74dbb 100644 --- a/src/isp_trace_parser/trace_restructure_helper_functions.py +++ b/src/isp_trace_parser/trace_restructure_helper_functions.py @@ -5,7 +5,6 @@ # the Free Software Foundation; either version 3 of the License, or # (at your option) any later version. -from datetime import timedelta from pathlib import Path import polars as pl @@ -17,14 +16,12 @@ def get_all_filepaths(directory: Path) -> list[Path]: if directory.is_dir(): return [path for path in Path(directory).rglob("*.csv") if path.is_file()] - else: - raise ValueError(f"{directory} not found.") + raise FileNotFoundError(directory) def read_trace_csv(file: Path) -> pl.DataFrame: pl_types = [pl.Int64] * 3 + [pl.Float64] * 48 - data = pl.read_csv(file, schema_overrides=pl_types) - return data + return pl.read_csv(file, schema_overrides=pl_types) def read_and_format_traces(files: list[Path]) -> list[pl.DataFrame]: @@ -38,10 +35,9 @@ def read_and_format_traces(files: list[Path]) -> list[pl.DataFrame]: def calculate_average_trace(traces: list[pl.DataFrame]) -> pl.DataFrame: combined_traces = pl.concat(traces) - average_trace = combined_traces.group_by("datetime").agg( + return combined_traces.group_by("datetime").agg( [pl.col("value").mean().alias("value")] ) - return average_trace def _frame_with_metadata(trace: pl.DataFrame, file_metadata: dict) -> pl.DataFrame: @@ -79,11 +75,7 @@ def process_and_save_files( ) -> None: traces = read_and_format_traces(files) - if len(traces) > 1: - trace = calculate_average_trace(traces) - else: - trace = traces[0] - + trace = calculate_average_trace(traces) if len(traces) > 1 else traces[0] trace = _frame_with_metadata(trace, file_metadata) save_trace(trace, file_metadata, output_directory, write_output_filepath) @@ -94,21 +86,18 @@ def get_metadata_that_matches_trace_names( ) -> dict[Path, dict[str, str]]: if isinstance(trace_names, str): trace_names = [trace_names] - matching_meta_data = { + return { f: metadata.copy() for f, metadata in all_input_file_metadata.items() if metadata["name"] in trace_names } - return matching_meta_data def get_unique_reference_years_in_metadata( metadata_for_trace_files: dict[Path, dict[str, str]], ) -> list[str]: return list( - set( - metadata["reference_year"] for metadata in metadata_for_trace_files.values() - ) + {metadata["reference_year"] for metadata in metadata_for_trace_files.values()} ) @@ -142,9 +131,12 @@ def check_filter_by_metadata( return True for field, allowed_values in filters.model_dump(exclude_unset=True).items(): - if field in metadata and allowed_values is not None: - if metadata[field] not in allowed_values: - return False + if ( + field in metadata + and allowed_values is not None + and metadata[field] not in allowed_values + ): + return False return True @@ -152,9 +144,7 @@ def check_filter_by_metadata( def get_unique_project_and_zone_names_in_input_files( metadata_for_trace_files: dict[Path, dict[str, str]], ) -> list[str]: - names = [] - for filepath, meta_data in metadata_for_trace_files.items(): - names.append(meta_data["name"]) + names = [meta_data["name"] for meta_data in metadata_for_trace_files.values()] return list(set(names)) @@ -172,5 +162,5 @@ def filter_mapping_by_names_in_input_files( return filtered_mapping -def get_just_filepaths(metadata_for_files): +def get_just_filepaths(metadata_for_files: dict) -> list: return [file for file, metadata in metadata_for_files.items()] diff --git a/src/isp_trace_parser/wind_traces.py b/src/isp_trace_parser/wind_traces.py index 8c3d2b7..2ff7467 100644 --- a/src/isp_trace_parser/wind_traces.py +++ b/src/isp_trace_parser/wind_traces.py @@ -8,7 +8,7 @@ import functools import os from pathlib import Path -from typing import Literal, Optional +from typing import Literal from joblib import Parallel, delayed from pydantic import BaseModel, validate_call @@ -56,10 +56,10 @@ class WindMetadataFilter(BaseModel): reference_year: list of ints specifying reference_years """ - name: Optional[list[str]] = None - file_type: Optional[list[Literal["zone", "project"]]] = None - resource_type: Optional[list[Literal["WH", "WM", "WL", "WX", "wind"]]] = None - reference_year: Optional[list[int]] = None + name: list[str] | None = None + file_type: list[Literal["zone", "project"]] | None = None + resource_type: list[Literal["WH", "WM", "WL", "WX", "wind"]] | None = None + reference_year: list[int] | None = None @validate_call @@ -68,7 +68,7 @@ def parse_wind_traces( parsed_directory: str | Path, use_concurrency: bool = True, filters: WindMetadataFilter | None = None, -): +) -> None: """Takes a directory with AEMO wind trace data and reformats the data, saving it to a new directory. AEMO wind trace data comes in CSVs with columns specifying the year, day, and month, and data columns @@ -163,12 +163,14 @@ def parse_wind_traces( zone_name_mappings = filter_mapping_by_names_in_input_files( zone_name_mappings, project_and_zone_input_names ) - zone_output_names, zone_input_names = zip(*zone_name_mappings.items()) + zone_output_names, zone_input_names = zip(*zone_name_mappings.items(), strict=True) project_name_mappings = filter_mapping_by_names_in_input_files( project_name_mappings, project_and_zone_input_names ) - project_output_names, project_input_names = zip(*project_name_mappings.items()) + project_output_names, project_input_names = zip( + *project_name_mappings.items(), strict=True + ) zone_partial_func = functools.partial( restructure_wind_zone_files, @@ -189,21 +191,27 @@ def parse_wind_traces( Parallel(n_jobs=max_workers)( delayed(zone_partial_func)(save_name, old_trace_name) - for save_name, old_trace_name in zip(zone_output_names, zone_input_names) + for save_name, old_trace_name in zip( + zone_output_names, zone_input_names, strict=True + ) ) Parallel(n_jobs=max_workers)( delayed(project_partial_func)(save_name, old_trace_name) for save_name, old_trace_name in zip( - project_output_names, project_input_names + project_output_names, project_input_names, strict=True ) ) else: - for save_name, old_trace_name in zip(zone_output_names, zone_input_names): + for save_name, old_trace_name in zip( + zone_output_names, zone_input_names, strict=True + ): zone_partial_func(save_name, old_trace_name) - for save_name, old_trace_name in zip(project_output_names, project_input_names): + for save_name, old_trace_name in zip( + project_output_names, project_input_names, strict=True + ): project_partial_func(save_name, old_trace_name) @@ -330,7 +338,7 @@ def get_unique_resource_types_in_metadata( metadata_for_trace_files: dict[str:str], ) -> list: return list( - set(metadata["resource_type"] for metadata in metadata_for_trace_files.values()) + {metadata["resource_type"] for metadata in metadata_for_trace_files.values()} ) diff --git a/tests/conftest.py b/tests/conftest.py index 8275ad8..04285b3 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -16,7 +16,7 @@ @pytest.fixture(params=[True, False], ids=["concurrent", "sequential"], scope="module") -def parsed_trace_trace_directory(request): +def parsed_trace_trace_directory(request) -> Path: """Fixture that performs parsing of wind and solar trace directory once, providing the output directory to multiple test cases that validate different files. @@ -59,7 +59,7 @@ def parsed_trace_trace_directory(request): optimise_parquet.partition_traces_by_columns( input_directory=tmp_parsed_directory / "demand", - output_directory=tmp_parsed_directory / f"demand_optimised", + output_directory=tmp_parsed_directory / "demand_optimised", partition_cols=["scenario", "reference_year"], ) yield tmp_parsed_directory diff --git a/tests/create_end_to_end_test_data.py b/tests/create_end_to_end_test_data.py index 5bf3ed6..e45e95f 100644 --- a/tests/create_end_to_end_test_data.py +++ b/tests/create_end_to_end_test_data.py @@ -13,7 +13,7 @@ import pandas as pd -def generate_random_data(start_year, end_year): +def generate_random_data(start_year: int, end_year: int) -> pd.DataFrame: # Generate date range from July 1st of the start year to July 1st of the end year (excluding end) date_range = pd.date_range( start=f"{start_year}-01-01", end=f"{end_year}-01-01", freq="D", inclusive="left" @@ -31,14 +31,13 @@ def generate_random_data(start_year, end_year): half_hour_columns = [f"{i:02d}" for i in range(1, 49)] # Combine the date components with the random data - df = pd.concat([df, pd.DataFrame(random_data, columns=half_hour_columns)], axis=1) - return df + return pd.concat([df, pd.DataFrame(random_data, columns=half_hour_columns)], axis=1) data = generate_random_data(start_year=config.start, end_year=config.end) -def simple_flatten(nested_list): +def simple_flatten(nested_list: list) -> list: flattened = [] for item in nested_list: if isinstance(item, list): @@ -48,7 +47,7 @@ def simple_flatten(nested_list): return flattened -def create_solar_csvs(directory): +def create_solar_csvs(directory: Path) -> None: combos = itertools.product(config.reference_years, config.solar_projects) for y, project in combos: data.to_csv(directory / Path(f"{project}_FFP_RefYear{y}.csv"), index=False) @@ -60,7 +59,7 @@ def create_solar_csvs(directory): ) -def create_wind_csvs(directory): +def create_wind_csvs(directory: Path) -> None: combos = itertools.product( config.reference_years, simple_flatten(config.wind_projects.values()) ) @@ -78,7 +77,7 @@ def create_wind_csvs(directory): ) -def create_demand_csvs(directory): +def create_demand_csvs(directory: Path) -> None: combos = itertools.product( config.reference_years, config.sub_regions, diff --git a/tests/test_demand_trace_metadata.py b/tests/test_demand_trace_metadata.py index 22414e2..8c48ae4 100644 --- a/tests/test_demand_trace_metadata.py +++ b/tests/test_demand_trace_metadata.py @@ -12,7 +12,7 @@ from isp_trace_parser import demand_trace_metadata -def test_build(): +def test_build() -> None: """Two examples spanning different scenario / poe / demand_type / subregion values. Every combination resolves through the same single dict lookup, so two are enough for testing. @@ -47,6 +47,6 @@ def test_build(): "VIC_RefYear_2011_MYSTERY_POE10_OPSO_MODELLING.csv", # lookup miss ], ) -def test_build_rejects_unexpected_filename(filename): +def test_build_rejects_unexpected_filename(filename: str) -> None: with pytest.raises(ValueError, match="Unexpected trace filename"): demand_trace_metadata.build([Path(filename)], version="2024") diff --git a/tests/test_download.py b/tests/test_download.py index 7c8bed0..4edbd89 100644 --- a/tests/test_download.py +++ b/tests/test_download.py @@ -16,7 +16,7 @@ TEST_EXPECTED_CONTENT = b"ISP Trace Parser Test File\n" -def test_download_test_file(): +def test_download_test_file() -> None: """Test download with actual server file.""" with TemporaryDirectory() as tmp_path: @@ -28,7 +28,7 @@ def test_download_test_file(): assert downloaded.read_bytes() == TEST_EXPECTED_CONTENT -def test_download_with_retry(): +def test_download_with_retry() -> None: """Test retry logic with real server.""" with TemporaryDirectory() as tmp_path: @@ -38,7 +38,7 @@ def test_download_with_retry(): assert (tmp_path / "test" / "test" / "test_file.txt").exists() -def test_fetch_trace_data_with_test_manifest(monkeypatch): +def test_fetch_trace_data_with_test_manifest(monkeypatch: pytest.MonkeyPatch) -> None: """Test downloading from a small, test manifest. The testing manifest, while still named "full_isp_2024" here, is just a test manifest with containing a single url ("https://data.openisp.au/test/test/test_file.txt") @@ -48,7 +48,7 @@ def test_fetch_trace_data_with_test_manifest(monkeypatch): tmp_path = Path(tmp_path) # Point to test fixtures instead of production manifests - def mock_files(package): + def mock_files(package) -> Path: return Path(__file__).parent / "fixtures" / "manifests" monkeypatch.setattr("isp_trace_parser.remote.download.files", mock_files) @@ -65,7 +65,7 @@ def mock_files(package): assert downloaded.read_bytes() == TEST_EXPECTED_CONTENT -def test_manifest_not_found(): +def test_manifest_not_found() -> None: """Test downloading from a small, test manifest.""" with pytest.raises(FileNotFoundError): @@ -75,7 +75,7 @@ def test_manifest_not_found(): @pytest.mark.parametrize("unquote", [True, False]) -def test_fetch_trace_data(unquote: bool, monkeypatch): +def test_fetch_trace_data(unquote: bool, monkeypatch: pytest.MonkeyPatch) -> None: """Test downloading via fetch_trace_data with test fixtures. This, while still download a dataset name "full", is just a pointing to a test manifest manifest with containing a single url ("https://data.openisp.au/test/test/test_file.txt") @@ -85,7 +85,7 @@ def test_fetch_trace_data(unquote: bool, monkeypatch): tmp_path = Path(tmp_path) # Point to test manifests instead of production manifests - def mock_files(package): + def mock_files(package) -> Path: return Path(__file__).parent / "fixtures" / "manifests" monkeypatch.setattr("isp_trace_parser.remote.download.files", mock_files) @@ -102,13 +102,13 @@ def mock_files(package): assert downloaded.read_bytes() == TEST_EXPECTED_CONTENT -def test_wrong_source(): +def test_wrong_source() -> None: # no ISP 2025 data with pytest.raises(ValueError, match="Only isp_2024 is currently supported"): download.fetch_trace_data("example", "isp_2025", "/", "archive") -def test_wrong_format(): +def test_wrong_format() -> None: # only archive or processed data (not other) with pytest.raises( ValueError, match="data_format must be 'processed' or 'archive'" @@ -116,21 +116,19 @@ def test_wrong_format(): download.fetch_trace_data("example", "isp_2024", "/", "other") -def test_wrong_type(): +def test_wrong_type() -> None: # only full or example type with pytest.raises(ValueError): download.fetch_trace_data("other", "isp_2024", "/", "archive") -def test_empty_manifest(monkeypatch): +def test_empty_manifest(monkeypatch: pytest.MonkeyPatch) -> None: """Test that empty manifest raises ValueError.""" - from importlib.resources import files - with TemporaryDirectory() as tmp_path: tmp_path = Path(tmp_path) # Point to test manifest instead of production manifests - def mock_files(package): + def mock_files(package) -> Path: return Path(__file__).parent / "fixtures" / "manifests" monkeypatch.setattr("isp_trace_parser.remote.download.files", mock_files) @@ -139,7 +137,7 @@ def mock_files(package): download._download_from_manifest("empty_manifest", tmp_path, strip_levels=0) -def test_strip_levels_too_high(): +def test_strip_levels_too_high() -> None: """Test that strip_levels >= path parts raises ValueError.""" with TemporaryDirectory() as tmp_path: tmp_path = Path(tmp_path) diff --git a/tests/test_get_data.py b/tests/test_get_data.py index 3f454d3..44637d4 100644 --- a/tests/test_get_data.py +++ b/tests/test_get_data.py @@ -32,7 +32,7 @@ TEST_DATA = Path(__file__).parent / "test_data" -def test_year_range_to_dt_range_fy(): +def test_year_range_to_dt_range_fy() -> None: """Test financial year conversion.""" start_dt, end_dt = _year_range_to_dt_range(2022, 2024, year_type="fy") @@ -40,7 +40,7 @@ def test_year_range_to_dt_range_fy(): assert end_dt == datetime.datetime(2024, 7, 1, 0, 0) -def test_year_range_to_dt_range_calendar(): +def test_year_range_to_dt_range_calendar() -> None: """Test calendar year conversion.""" start_dt, end_dt = _year_range_to_dt_range(2022, 2024, year_type="calendar") @@ -49,7 +49,9 @@ def test_year_range_to_dt_range_calendar(): @pytest.mark.parametrize("year_type", ["fy", "calendar"]) -def test_get_zone_single_reference_year(parsed_trace_trace_directory: Path, year_type): +def test_get_zone_single_reference_year( + parsed_trace_trace_directory: Path, year_type: str +) -> None: test_df_lazy = pl.scan_parquet(TEST_DATA / "output" / "RefYear2022_N2_CST.parquet") start_dt, end_dt = _year_range_to_dt_range(2023, 2024, year_type=year_type) @@ -76,7 +78,7 @@ def test_get_zone_single_reference_year(parsed_trace_trace_directory: Path, year pd.testing.assert_frame_equal(test_df, df) -def test_get_zone_multiple_reference_year(parsed_trace_trace_directory: Path): +def test_get_zone_multiple_reference_year(parsed_trace_trace_directory: Path) -> None: test_df_lazy = pl.scan_parquet(TEST_DATA / "output" / "RefYear2022_N1_WM.parquet") test_df = ( @@ -100,7 +102,7 @@ def test_get_zone_multiple_reference_year(parsed_trace_trace_directory: Path): pd.testing.assert_frame_equal(test_df, df) -def test_get_project_single_reference_year(parsed_trace_trace_directory: Path): +def test_get_project_single_reference_year(parsed_trace_trace_directory: Path) -> None: test_df_lazy = pl.scan_parquet( TEST_DATA / "output" / "RefYear2022_Bodangora_Wind_Farm.parquet" ) @@ -127,7 +129,9 @@ def test_get_project_single_reference_year(parsed_trace_trace_directory: Path): pd.testing.assert_frame_equal(test_df, df) -def test_get_project_multiple_reference_year(parsed_trace_trace_directory: Path): +def test_get_project_multiple_reference_year( + parsed_trace_trace_directory: Path, +) -> None: test_df_lazy = pl.scan_parquet( TEST_DATA / "output" / "RefYear2022_Broken_Hill_Solar_Farm_FFP.parquet" ) @@ -152,7 +156,7 @@ def test_get_project_multiple_reference_year(parsed_trace_trace_directory: Path) pd.testing.assert_frame_equal(test_df, df) -def test_get_demand_single_reference_year(parsed_trace_trace_directory: Path): +def test_get_demand_single_reference_year(parsed_trace_trace_directory: Path) -> None: test_df_lazy = pl.scan_parquet( TEST_DATA / "output" @@ -185,7 +189,7 @@ def test_get_demand_single_reference_year(parsed_trace_trace_directory: Path): pd.testing.assert_frame_equal(test_df, df) -def test_get_demand_multiple_reference_year(parsed_trace_trace_directory: Path): +def test_get_demand_multiple_reference_year(parsed_trace_trace_directory: Path) -> None: test_df_lazy = pl.scan_parquet( TEST_DATA / "output" @@ -215,7 +219,7 @@ def test_get_demand_multiple_reference_year(parsed_trace_trace_directory: Path): pd.testing.assert_frame_equal(test_df, df) -def test_explicit_select_columns(parsed_trace_trace_directory): +def test_explicit_select_columns(parsed_trace_trace_directory: Path) -> None: df = get_zone_single_reference_year( start_year=2023, end_year=2024, @@ -228,7 +232,7 @@ def test_explicit_select_columns(parsed_trace_trace_directory): assert list(df.columns) == ["datetime", "value", "zone"] -def test_multi_value_filter(parsed_trace_trace_directory): +def test_multi_value_filter(parsed_trace_trace_directory: Path) -> None: df = get_zone_single_reference_year( start_year=2023, end_year=2024, @@ -241,7 +245,7 @@ def test_multi_value_filter(parsed_trace_trace_directory): assert "zone" in df.columns -def test_wind_project_single_reference_year(parsed_trace_trace_directory): +def test_wind_project_single_reference_year(parsed_trace_trace_directory: Path) -> None: test_df_lazy = pl.scan_parquet( TEST_DATA / "output" / "RefYear2022_Bodangora_Wind_Farm.parquet" ) @@ -267,7 +271,9 @@ def test_wind_project_single_reference_year(parsed_trace_trace_directory): pd.testing.assert_frame_equal(test_df, df) -def test_solar_project_single_reference_year(parsed_trace_trace_directory): +def test_solar_project_single_reference_year( + parsed_trace_trace_directory: Path, +) -> None: test_df_lazy = pl.scan_parquet( TEST_DATA / "output" / "RefYear2022_Broken_Hill_Solar_Farm_FFP.parquet" ) @@ -293,7 +299,9 @@ def test_solar_project_single_reference_year(parsed_trace_trace_directory): pd.testing.assert_frame_equal(test_df, df) -def test_solar_project_multiple_reference_years(parsed_trace_trace_directory: Path): +def test_solar_project_multiple_reference_years( + parsed_trace_trace_directory: Path, +) -> None: test_df_lazy = pl.scan_parquet( TEST_DATA / "output" / "RefYear2022_Broken_Hill_Solar_Farm_FFP.parquet" ) @@ -318,7 +326,9 @@ def test_solar_project_multiple_reference_years(parsed_trace_trace_directory: Pa pd.testing.assert_frame_equal(test_df, df) -def test_wind_project_multiple_reference_years(parsed_trace_trace_directory: Path): +def test_wind_project_multiple_reference_years( + parsed_trace_trace_directory: Path, +) -> None: test_df_lazy = pl.scan_parquet( TEST_DATA / "output" / "RefYear2022_Bodangora_Wind_Farm.parquet" ) @@ -343,7 +353,7 @@ def test_wind_project_multiple_reference_years(parsed_trace_trace_directory: Pat pd.testing.assert_frame_equal(test_df, df) -def test_solar_area_single_reference_year(parsed_trace_trace_directory: Path): +def test_solar_area_single_reference_year(parsed_trace_trace_directory: Path) -> None: test_df_lazy = pl.scan_parquet(TEST_DATA / "output" / "RefYear2022_N2_CST.parquet") start_dt, end_dt = _year_range_to_dt_range(2023, 2024, year_type="fy") @@ -369,7 +379,7 @@ def test_solar_area_single_reference_year(parsed_trace_trace_directory: Path): pd.testing.assert_frame_equal(test_df, df) -def test_demand_single_reference_year(parsed_trace_trace_directory: Path): +def test_demand_single_reference_year(parsed_trace_trace_directory: Path) -> None: test_df_lazy = pl.scan_parquet( TEST_DATA / "output" @@ -402,7 +412,7 @@ def test_demand_single_reference_year(parsed_trace_trace_directory: Path): pd.testing.assert_frame_equal(test_df, df) -def test_demand_multiple_reference_years(parsed_trace_trace_directory: Path): +def test_demand_multiple_reference_years(parsed_trace_trace_directory: Path) -> None: test_df_lazy = pl.scan_parquet( TEST_DATA / "output" diff --git a/tests/test_input_validation.py b/tests/test_input_validation.py index a631127..f0d5a3c 100644 --- a/tests/test_input_validation.py +++ b/tests/test_input_validation.py @@ -38,7 +38,7 @@ }, ], ) -def test_solar_metadata_filter_valid(valid_input): +def test_solar_metadata_filter_valid(valid_input: dict[str, list]) -> None: assert SolarMetadataFilter(**valid_input) @@ -51,7 +51,9 @@ def test_solar_metadata_filter_valid(valid_input): ({"name": 123}, "Input should be a valid list"), ], ) -def test_solar_metadata_filter_invalid(invalid_input, expected_error): +def test_solar_metadata_filter_invalid( + invalid_input: dict[str, list], expected_error: str +) -> None: with pytest.raises(ValidationError, match=expected_error): SolarMetadataFilter(**invalid_input) @@ -71,7 +73,7 @@ def test_solar_metadata_filter_invalid(invalid_input, expected_error): }, ], ) -def test_wind_metadata_filter_valid(valid_input): +def test_wind_metadata_filter_valid(valid_input: dict[str, list]) -> None: assert WindMetadataFilter(**valid_input) @@ -87,7 +89,9 @@ def test_wind_metadata_filter_valid(valid_input): ({"name": 123}, "Input should be a valid list"), ], ) -def test_wind_metadata_filter_invalid(invalid_input, expected_error): +def test_wind_metadata_filter_invalid( + invalid_input: dict[str, list], expected_error: str +) -> None: with pytest.raises(ValidationError, match=expected_error): WindMetadataFilter(**invalid_input) @@ -109,7 +113,7 @@ def test_wind_metadata_filter_invalid(invalid_input, expected_error): }, ], ) -def test_demand_metadata_filter_valid(valid_input): +def test_demand_metadata_filter_valid(valid_input: dict[str, list]) -> None: assert DemandMetadataFilter(**valid_input) @@ -129,7 +133,9 @@ def test_demand_metadata_filter_valid(valid_input): ({"subregion": 123}, "Input should be a valid list"), ], ) -def test_demand_metadata_filter_invalid(invalid_input, expected_error): +def test_demand_metadata_filter_invalid( + invalid_input: dict[str, str | int | list], expected_error: str +) -> None: with pytest.raises(ValidationError, match=expected_error): DemandMetadataFilter(**invalid_input) @@ -152,7 +158,7 @@ def test_demand_metadata_filter_invalid(invalid_input, expected_error): }, ], ) -def test_parse_traces_validation(invalid_input): +def test_parse_traces_validation(invalid_input: dict[str, str | int | list]) -> None: with pytest.raises(ValidationError): parse_solar_traces(**invalid_input) with pytest.raises(ValidationError): @@ -170,12 +176,14 @@ def test_parse_traces_validation(invalid_input): {"start_year": 2030, "end_year": 2035, "reference_years": [2011, "x", 2018]}, ], ) -def test_construct_reference_year_mapping_validation_invalid(invalid_input): +def test_construct_reference_year_mapping_validation_invalid( + invalid_input: dict, +) -> None: with pytest.raises(ValidationError): construct_reference_year_mapping(**invalid_input) -def test_construct_reference_year_mapping_validation_valid(): +def test_construct_reference_year_mapping_validation_valid() -> None: result = construct_reference_year_mapping( start_year=2030, end_year=2035, reference_years=[2011, 2013, 2018] ) @@ -185,12 +193,11 @@ def test_construct_reference_year_mapping_validation_valid(): # Tests for custom input validation functions -def test_input_directory(tmp_path): +def test_input_directory(tmp_path: Path) -> None: valid_dir = tmp_path / "valid_dir" valid_dir.mkdir() assert input_validation.input_directory(valid_dir) == valid_dir - - with pytest.raises(ValueError, match="Directory .* does not exist"): + with pytest.raises(FileNotFoundError): input_validation.input_directory(tmp_path / "non_existent_dir") @@ -201,7 +208,7 @@ def test_input_directory(tmp_path): Path("/valid/path"), ], ) -def test_parsed_directory_valid(valid_path): +def test_parsed_directory_valid(valid_path: Path | str) -> None: result = input_validation.parsed_directory(valid_path) assert isinstance(result, Path) @@ -214,7 +221,7 @@ def test_parsed_directory_valid(valid_path): [], ], ) -def test_parsed_directory_invalid(invalid_path): +def test_parsed_directory_invalid(invalid_path: Path | str) -> None: with pytest.raises(ValueError, match="Invalid parsed directory path"): input_validation.parsed_directory(invalid_path) @@ -226,7 +233,7 @@ def test_parsed_directory_invalid(invalid_path): Path("/valid/path"), ], ) -def test_is_valid_path_valid(valid_path): +def test_is_valid_path_valid(valid_path: Path | str) -> None: result = input_validation.is_valid_path(valid_path) assert isinstance(result, Path) @@ -239,7 +246,7 @@ def test_is_valid_path_valid(valid_path): [], ], ) -def test_is_valid_path_invalid(invalid_path): +def test_is_valid_path_invalid(invalid_path: Path | str) -> None: with pytest.raises(ValueError, match="Invalid parsed directory path"): input_validation.is_valid_path(invalid_path) @@ -252,7 +259,7 @@ def test_is_valid_path_invalid(invalid_path): (-10, 0), ], ) -def test_start_year_before_end_year_valid(start, end): +def test_start_year_before_end_year_valid(start: int, end: int) -> None: assert input_validation.start_year_before_end_year(start, end) is None @@ -264,6 +271,6 @@ def test_start_year_before_end_year_valid(start, end): (2020, 2019), ], ) -def test_start_year_before_end_year_invalid(start, end): +def test_start_year_before_end_year_invalid(start: int, end: int) -> None: with pytest.raises(ValueError, match="Start year .* < end year"): input_validation.start_year_before_end_year(start, end) diff --git a/tests/test_optimise_parquet.py b/tests/test_optimise_parquet.py index beec176..01c9b59 100644 --- a/tests/test_optimise_parquet.py +++ b/tests/test_optimise_parquet.py @@ -20,7 +20,9 @@ "expected_data, file_type", [("zone_data_0.parquet", "zone"), ("project_data_0.parquet", "project")], ) -def test_optimisation(parsed_trace_trace_directory, expected_data, file_type): +def test_optimisation( + parsed_trace_trace_directory: Path, expected_data: str, file_type: str +) -> None: """Test wind trace parsing produces expected parquet outputs (both for a sample wind project and wind zone)""" test_output_parquet = TEST_DATA / "output" / expected_data diff --git a/tests/test_resource_trace_metadata.py b/tests/test_resource_trace_metadata.py index 61861ae..2f5c063 100644 --- a/tests/test_resource_trace_metadata.py +++ b/tests/test_resource_trace_metadata.py @@ -12,7 +12,7 @@ from isp_trace_parser import resource_trace_metadata -def test_build(): +def test_build() -> None: """One test covers function logic compared with regex approach Solar zones / wind zones / extra reference years add no new code-path @@ -42,6 +42,6 @@ def test_build(): "Mystery_Plant_RefYear2011.csv", # stem not in mapping ], ) -def test_build_rejects_unexpected_filename(filename): +def test_build_rejects_unexpected_filename(filename: str) -> None: with pytest.raises(ValueError, match="Unexpected trace filename"): resource_trace_metadata.build([Path(filename)], version="2024") diff --git a/tests/test_trace_formatter.py b/tests/test_trace_formatter.py index 90be4a6..9cb93c7 100644 --- a/tests/test_trace_formatter.py +++ b/tests/test_trace_formatter.py @@ -11,7 +11,7 @@ from isp_trace_parser import trace_formatter, trace_restructure_helper_functions -def test_trace_formatter(): +def test_trace_formatter() -> None: # Test trace formatting works by using formatting function works by performing formatting and then # reversing the formatting changes and checking the result matches the original data. filepath = ( diff --git a/tests/test_trace_parsers.py b/tests/test_trace_parsers.py index a49954a..64ae3a3 100644 --- a/tests/test_trace_parsers.py +++ b/tests/test_trace_parsers.py @@ -12,13 +12,13 @@ import pytest from polars.testing import assert_frame_equal -from isp_trace_parser import demand_traces, solar_traces, wind_traces +from isp_trace_parser import demand_traces TEST_DATA = Path(__file__).parent / "test_data" @pytest.mark.parametrize("use_concurrency", [True, False]) -def test_demand_trace_parsing(use_concurrency: bool): +def test_demand_trace_parsing(use_concurrency: bool) -> None: """Test demand trace parsing produces expected parquet output.""" test_demand_csv_directory = TEST_DATA / "demand" expected_filename = "CNSW_RefYear_2011_HYDROGEN_EXPORT_POE10_OPSO_MODELLING.parquet" @@ -50,7 +50,9 @@ def test_demand_trace_parsing(use_concurrency: bool): ("RefYear2022_N1_WM.parquet", "zone"), ], ) -def test_wind_trace_parsing(parsed_trace_trace_directory, expected_filename, file_type): +def test_wind_trace_parsing( + parsed_trace_trace_directory: Path, expected_filename: str, file_type: str +) -> None: """Test wind trace parsing produces expected parquet outputs (both for a sample wind project and wind zone)""" test_output_parquet = TEST_DATA / "output" / expected_filename @@ -70,8 +72,8 @@ def test_wind_trace_parsing(parsed_trace_trace_directory, expected_filename, fil ], ) def test_solar_trace_parsing( - parsed_trace_trace_directory, expected_filename, file_type -): + parsed_trace_trace_directory: Path, expected_filename: str, file_type: str +) -> None: """Test solar trace parsing produces expected parquet output (both for a sample solar project and solar zone)""" test_output_parquet = TEST_DATA / "output" / expected_filename diff --git a/tests/test_writing_save_names.py b/tests/test_writing_save_names.py index 36e3c30..ffd3be3 100644 --- a/tests/test_writing_save_names.py +++ b/tests/test_writing_save_names.py @@ -8,7 +8,7 @@ import isp_trace_parser -def test_write_solar_save_names(): +def test_write_solar_save_names() -> None: meta_data = { "name": "a", "reference_year": "1", @@ -32,7 +32,7 @@ def test_write_solar_save_names(): assert str(save_filepath) == "RefYear1_a_x.parquet" -def test_write_wind_save_names(): +def test_write_wind_save_names() -> None: meta_data = { "name": "a", "reference_year": "1",