95 lines
2.7 KiB
Python
95 lines
2.7 KiB
Python
from __future__ import annotations
|
|
|
|
import csv
|
|
import io
|
|
import unittest
|
|
|
|
from fastapi import HTTPException
|
|
|
|
from app.main import (
|
|
SimulationResultCsvPayload,
|
|
SimulationResultVariablePayload,
|
|
export_simulation_results_csv,
|
|
)
|
|
|
|
|
|
def result_variable(
|
|
key: str,
|
|
component_id: str,
|
|
name: str,
|
|
label: str,
|
|
unit: str,
|
|
) -> SimulationResultVariablePayload:
|
|
return SimulationResultVariablePayload(
|
|
key=key,
|
|
componentId=component_id,
|
|
componentType="tank",
|
|
scope="component",
|
|
name=name,
|
|
label=label,
|
|
quantity="pressure",
|
|
unit=unit,
|
|
)
|
|
|
|
|
|
class ResultCsvExportTests(unittest.TestCase):
|
|
def valid_payload(self) -> SimulationResultCsvPayload:
|
|
return SimulationResultCsvPayload(
|
|
projectName="储气系统",
|
|
variables=[
|
|
result_variable(
|
|
"cylinder_1.p",
|
|
"cylinder_1",
|
|
"p",
|
|
"压力",
|
|
"Pa",
|
|
),
|
|
result_variable("tank_1.p", "tank_1", "p", "压力", "Pa"),
|
|
],
|
|
series={
|
|
"time": [0.0, 0.1],
|
|
"cylinder_1.p": [35000000.0, 34900000.0],
|
|
"tank_1.p": [100000.0, 101000.0],
|
|
},
|
|
)
|
|
|
|
def test_csv_export_preserves_result_keys_and_rows(self) -> None:
|
|
response = export_simulation_results_csv(self.valid_payload())
|
|
|
|
text = response.body.decode("utf-8-sig")
|
|
rows = list(csv.reader(io.StringIO(text)))
|
|
self.assertEqual(
|
|
rows[0],
|
|
["time", "cylinder_1.p", "tank_1.p"],
|
|
)
|
|
self.assertEqual(rows[1], ["0.0", "35000000.0", "100000.0"])
|
|
self.assertEqual(rows[2], ["0.1", "34900000.0", "101000.0"])
|
|
self.assertIn(
|
|
"filename*=UTF-8''",
|
|
response.headers["content-disposition"],
|
|
)
|
|
|
|
def test_csv_export_rejects_inconsistent_column_lengths(self) -> None:
|
|
payload = self.valid_payload()
|
|
payload.series["tank_1.p"] = [100000.0]
|
|
|
|
with self.assertRaises(HTTPException) as caught:
|
|
export_simulation_results_csv(payload)
|
|
|
|
self.assertEqual(caught.exception.status_code, 422)
|
|
self.assertIn("inconsistent length", str(caught.exception.detail))
|
|
|
|
def test_csv_export_requires_metadata_for_every_result_column(self) -> None:
|
|
payload = self.valid_payload()
|
|
payload.series["orphan.value"] = [1.0, 2.0]
|
|
|
|
with self.assertRaises(HTTPException) as caught:
|
|
export_simulation_results_csv(payload)
|
|
|
|
self.assertEqual(caught.exception.status_code, 422)
|
|
self.assertIn("unmapped orphan.value", str(caught.exception.detail))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|