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

DRY get_unique

parent 96770872
......@@ -119,15 +119,22 @@ class SimStadtResults(ABC):
@property
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
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
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
def workflow_path(self) -> Path:
......@@ -136,7 +143,10 @@ class SimStadtResults(ABC):
@property
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
def project_path(self) -> Path:
......@@ -170,13 +180,7 @@ class SimStadtResults(ABC):
@property
def csv_path(self) -> Path:
csvs = self.get_all_by_extension(".csv")
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]
return self.get_unique_by_extension('.csv', self.csv_identifier)
def csv_content(self) -> bytes:
with open(self.csv_path, "rb") as f:
......@@ -185,6 +189,16 @@ class SimStadtResults(ABC):
def get_all_by_extension(self, ext: str) -> list[Path]:
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
def diagrams(self) -> dict[str, Path]:
return {}
......@@ -243,6 +257,10 @@ class SimStadtResults(ABC):
def create(cls, provider: str, **kwargs) -> "SimStadtResults":
if provider not in cls._registry:
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 cls._registry[provider](**kwargs)
......@@ -19,9 +19,7 @@ class EnergyGridResults(SimStadtResults):
csv_identifier = "json_only"
def _parse_results(self) -> pd.DataFrame:
jsons = self.get_all_by_extension(".json")
assert len(jsons) == 1, f"Exactly one JSON file should have been written: {len(jsons)} were found."
json_output = jsons[0]
json_output = self.get_unique_by_extension(".json")
with open(json_output, encoding="utf-8") as out:
content = json.load(out)
buildings = content["buildings"]
......
......@@ -20,11 +20,9 @@ class SolarPotentialResults(SimStadtResults):
csv_identifier = "_solar_potential" # no CSV output; csv_path will raise if called
def _parse_results(self) -> pd.DataFrame:
prn_paths = self.get_all_by_extension(".prn")
if len(prn_paths) != 1:
raise ValueError(f"Expected exactly one .prn file, found {len(prn_paths)}")
prn_path = self.get_unique_by_extension(".prn")
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
......
......@@ -619,11 +619,11 @@ def test_broken_simulation():
/ "CGSC.proj/103_PV.flow/04_PhotovoltaicPotential.step/StillNotATable.png",
]
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)
with pytest.raises(ValueError, match="any CSV file"):
with pytest.raises(ValueError, match="any .csv file"):
results.to_json()
with pytest.raises(ValueError, match="any CSV file"):
with pytest.raises(ValueError, match="any .csv file"):
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