Commit cbe708e5 authored by Eric Duminil's avatar Eric Duminil
Browse files

DRY get_unique

parent 96770872
...@@ -119,15 +119,22 @@ class SimStadtResults(ABC): ...@@ -119,15 +119,22 @@ class SimStadtResults(ABC):
@property @property
def name(self) -> str: def name(self) -> str:
return self._params_element(".//void[@property='name']/string").text or "Unknown" return (
self._params_element(".//void[@property='name']/string").text or "Unknown"
)
@property @property
def short_name(self) -> str: def short_name(self) -> str:
return self._params_element(".//void[@property='shortName']/string").text or "Unknown" return (
self._params_element(".//void[@property='shortName']/string").text
or "Unknown"
)
@property @property
def workflow_provider(self) -> str: def workflow_provider(self) -> str:
return self._params_element(".//void[@property='workflowProvider']/object").get("class", "Unknown provider") return self._params_element(".//void[@property='workflowProvider']/object").get(
"class", "Unknown provider"
)
@property @property
def workflow_path(self) -> Path: def workflow_path(self) -> Path:
...@@ -136,7 +143,10 @@ class SimStadtResults(ABC): ...@@ -136,7 +143,10 @@ class SimStadtResults(ABC):
@property @property
def citygml(self) -> str: def citygml(self) -> str:
return self._params_element(".//void[@property='cityGmlFileNames']//string").text or "Unknown CityGML" return (
self._params_element(".//void[@property='cityGmlFileNames']//string").text
or "Unknown CityGML"
)
@property @property
def project_path(self) -> Path: def project_path(self) -> Path:
...@@ -170,13 +180,7 @@ class SimStadtResults(ABC): ...@@ -170,13 +180,7 @@ class SimStadtResults(ABC):
@property @property
def csv_path(self) -> Path: def csv_path(self) -> Path:
csvs = self.get_all_by_extension(".csv") return self.get_unique_by_extension('.csv', self.csv_identifier)
relevant_csvs = [csv for csv in csvs if self.csv_identifier in csv.name]
if len(relevant_csvs) == 0:
raise ValueError("Workflow didn't seem to have returned any CSV file!")
if len(relevant_csvs) > 1:
logger.warning("Too many CSV results found for %s", self)
return relevant_csvs[-1]
def csv_content(self) -> bytes: def csv_content(self) -> bytes:
with open(self.csv_path, "rb") as f: with open(self.csv_path, "rb") as f:
...@@ -185,6 +189,16 @@ class SimStadtResults(ABC): ...@@ -185,6 +189,16 @@ class SimStadtResults(ABC):
def get_all_by_extension(self, ext: str) -> list[Path]: def get_all_by_extension(self, ext: str) -> list[Path]:
return [f for f in sorted(self.output_files) if f.suffix == ext] return [f for f in sorted(self.output_files) if f.suffix == ext]
def get_unique_by_extension(self, ext: str, filter: str | None = None) -> Path:
files_by_extension = self.get_all_by_extension(ext)
if filter:
files_by_extension = [f for f in files_by_extension if filter in f.name]
if len(files_by_extension) == 0:
raise ValueError(f"Workflow didn't seem to have returned any {ext} file!")
if len(files_by_extension) > 1:
raise ValueError(f"Too many {ext} results found for {self}!")
return files_by_extension[0]
@property @property
def diagrams(self) -> dict[str, Path]: def diagrams(self) -> dict[str, Path]:
return {} return {}
...@@ -243,6 +257,10 @@ class SimStadtResults(ABC): ...@@ -243,6 +257,10 @@ class SimStadtResults(ABC):
def create(cls, provider: str, **kwargs) -> "SimStadtResults": def create(cls, provider: str, **kwargs) -> "SimStadtResults":
if provider not in cls._registry: if provider not in cls._registry:
from .unknown_workflow import UnknownWorkflowResults from .unknown_workflow import UnknownWorkflowResults
logger.warning("Unknown workflow provider: %s. Returning UnknownWorkflowResults.", provider)
logger.warning(
"Unknown workflow provider: %s. Returning UnknownWorkflowResults.",
provider,
)
return UnknownWorkflowResults(**kwargs) return UnknownWorkflowResults(**kwargs)
return cls._registry[provider](**kwargs) return cls._registry[provider](**kwargs)
...@@ -19,9 +19,7 @@ class EnergyGridResults(SimStadtResults): ...@@ -19,9 +19,7 @@ class EnergyGridResults(SimStadtResults):
csv_identifier = "json_only" csv_identifier = "json_only"
def _parse_results(self) -> pd.DataFrame: def _parse_results(self) -> pd.DataFrame:
jsons = self.get_all_by_extension(".json") json_output = self.get_unique_by_extension(".json")
assert len(jsons) == 1, f"Exactly one JSON file should have been written: {len(jsons)} were found."
json_output = jsons[0]
with open(json_output, encoding="utf-8") as out: with open(json_output, encoding="utf-8") as out:
content = json.load(out) content = json.load(out)
buildings = content["buildings"] buildings = content["buildings"]
......
...@@ -20,11 +20,9 @@ class SolarPotentialResults(SimStadtResults): ...@@ -20,11 +20,9 @@ class SolarPotentialResults(SimStadtResults):
csv_identifier = "_solar_potential" # no CSV output; csv_path will raise if called csv_identifier = "_solar_potential" # no CSV output; csv_path will raise if called
def _parse_results(self) -> pd.DataFrame: def _parse_results(self) -> pd.DataFrame:
prn_paths = self.get_all_by_extension(".prn") prn_path = self.get_unique_by_extension(".prn")
if len(prn_paths) != 1:
raise ValueError(f"Expected exactly one .prn file, found {len(prn_paths)}")
return pd.read_csv( return pd.read_csv(
prn_paths[0], sep=r"\s+", names=["GHI", "DHI", "Ta"], header=None prn_path, sep=r"\s+", names=["GHI", "DHI", "Ta"], header=None
) )
@property @property
......
...@@ -619,11 +619,11 @@ def test_broken_simulation(): ...@@ -619,11 +619,11 @@ def test_broken_simulation():
/ "CGSC.proj/103_PV.flow/04_PhotovoltaicPotential.step/StillNotATable.png", / "CGSC.proj/103_PV.flow/04_PhotovoltaicPotential.step/StillNotATable.png",
] ]
results = create_simstadt_results("FakeTests", output_files) results = create_simstadt_results("FakeTests", output_files)
with pytest.raises(ValueError, match="any CSV file"): with pytest.raises(ValueError, match="any .csv file"):
repr(results) repr(results)
with pytest.raises(ValueError, match="any CSV file"): with pytest.raises(ValueError, match="any .csv file"):
results.to_json() results.to_json()
with pytest.raises(ValueError, match="any CSV file"): with pytest.raises(ValueError, match="any .csv file"):
results.kpis results.kpis
......
Supports Markdown
0% or .
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment