Coverage for backend/django/PinchAnalysis/views/SegmentViewSet.py: 65%

148 statements  

« prev     ^ index     » next       coverage.py v7.10.7, created at 2026-07-22 05:22 +0000

1from PinchAnalysis.models.HenNode import HenNode 

2from core.auxiliary.views.ExtractSegmentDataFromFS import _calc_area 

3from core.auxiliary.enums.generalEnums import AbstractionType 

4from core.viewset import ModelViewSet 

5from PinchAnalysis.models.InputModels import PinchInputs, Segment, StreamDataEntry 

6from drf_spectacular.utils import extend_schema, OpenApiParameter, OpenApiTypes 

7from rest_framework.response import Response 

8from rest_framework import serializers 

9from rest_framework.decorators import action 

10import traceback 

11from PinchAnalysis.models.StreamDataProject import StreamDataProject 

12from flowsheetInternals.graphicData.models.groupingModel import Grouping 

13from PinchAnalysis.serializers.PinchInputSerializers import SegmentSerializer 

14from core.auxiliary.models.Flowsheet import Flowsheet 

15from flowsheetInternals.unitops.services.edit_operations.recorder import ( 

16 tracked_bulk_create, 

17) 

18 

19class BulkCreateStreamsSerializer(serializers.Serializer): 

20 projectID = serializers.IntegerField(required=True) 

21 streams = serializers.ListField( 

22 child=serializers.DictField(), 

23 required=True, 

24 ) 

25 

26class CustomSegmentSerializer(serializers.Serializer): 

27 id = serializers.IntegerField(allow_null=True) 

28 name = serializers.CharField() 

29 custom = serializers.BooleanField() 

30 type = serializers.CharField() 

31 t_supply = serializers.FloatField(allow_null=True) 

32 t_target = serializers.FloatField(allow_null=True) 

33 p_supply = serializers.FloatField(allow_null=True) 

34 p_target = serializers.FloatField(allow_null=True) 

35 h_supply = serializers.FloatField(allow_null=True) 

36 h_target = serializers.FloatField(allow_null=True) 

37 heat_flow = serializers.FloatField(allow_null=True) 

38 dt_cont = serializers.FloatField(allow_null=True) 

39 htc = serializers.FloatField(allow_null=True) 

40 area = serializers.FloatField(allow_null=True) 

41 children = serializers.ListField(child=serializers.DictField()) 

42 

43class CreateSegmentSerializer(serializers.Serializer): 

44 name = serializers.CharField(required=True) 

45 t_supply = serializers.FloatField(required=True) 

46 t_target = serializers.FloatField(required=True) 

47 heat_flow = serializers.FloatField(required=True) 

48 dt_cont = serializers.FloatField(required=True) 

49 htc = serializers.FloatField(required=True) 

50 parentZone = serializers.CharField(required=True) 

51 

52class SegmentViewSet(ModelViewSet): 

53 serializer_class = SegmentSerializer 

54 

55 def get_queryset(self): 

56 queryset = Segment.objects.all() 

57 return queryset 

58 

59 @extend_schema( 

60 responses=CustomSegmentSerializer 

61 ) 

62 def list(self, request): 

63 stream_data_projects = StreamDataProject.objects.all() 

64 all_groups = {} 

65 

66 # Step 1: Build all_groups with zones and child Zones/Segments 

67 for stream_data_project in stream_data_projects: 

68 for stream_data_entry in stream_data_project.StreamDataEntries.all(): 

69 zone = stream_data_entry.zone 

70 group = stream_data_entry.group 

71 

72 

73 # Prepare Segment nodes as child nodes 

74 segments = Segment.objects.filter( 

75 stream_data_entry=stream_data_entry) 

76 segment_children = [{ 

77 "id": segment.id, 

78 "name": segment.name, 

79 "type": "Segment", 

80 "custom": segment.stream_data_entry.custom, 

81 "t_supply": segment.t_supply, 

82 "t_target": segment.t_target, 

83 "p_supply": segment.p_supply, 

84 "p_target": segment.p_target, 

85 "h_supply": segment.h_supply, 

86 "h_target": segment.h_target, 

87 "heat_flow": segment.heat_flow, 

88 "dt_cont": segment.dt_cont, 

89 "htc": segment.htc, 

90 "area": segment.area, 

91 "stream_data_entry": segment.stream_data_entry.id, 

92 "hen_node": segment.hen_node.id if segment.hen_node else None, 

93 "children": [] # Segments don't have children 

94 } for segment in segments] 

95 

96 if zone not in all_groups: 

97 all_groups[zone] = { 

98 "group": group, 

99 "stream_data_entry": stream_data_entry, 

100 "children": [], 

101 } 

102 

103 # Add segments to the current zone's children 

104 all_groups[zone]["children"].extend(segment_children) 

105 

106 # Step 2: Organize into a tree 

107 # Need to review the parent groups creation here, it's a messy temporary fix but likely to cause problems longterm. 

108 root_node = None 

109 parent_groups = {} 

110 for zone_name, zone_data in all_groups.items(): 

111 group = zone_data["group"] 

112 parent_group = group.get_parent_group() 

113 if parent_group: 

114 parent_zone = parent_group.simulationObject.componentName 

115 if parent_zone in all_groups: 115 ↛ 117line 115 didn't jump to line 117 because the condition on line 115 was always true

116 all_groups[parent_zone]["children"].append(zone_data) 

117 elif parent_zone in parent_groups: 

118 parent_groups[parent_zone]["children"].append(zone_data) 

119 else: 

120 parent_groups[parent_zone] = { 

121 "group": parent_group, 

122 "children":[], 

123 } 

124 parent_groups[parent_zone]["children"].append(zone_data) 

125 root_node = parent_groups[parent_zone] 

126 else: 

127 root_node = zone_data 

128 

129 # Step 3: Format the output recursively 

130 def clean_node(node): 

131 # If it's a Segment node, it's already cleaned 

132 if node.get("type") == "Segment" or node.get("type") == "CustomSegment": 

133 return node 

134 return { 

135 "name": f'{node["group"].simulationObject.componentName} ({node["group"].abstractionType})', 

136 "type": node["group"].abstractionType, 

137 "custom": False, 

138 "t_supply": None, 

139 "t_target": None, 

140 "p_supply": None, 

141 "p_target": None, 

142 "h_supply": None, 

143 "h_target": None, 

144 "heat_flow": None, 

145 "dt_cont": None, 

146 "htc": None, 

147 "children": [clean_node(child) for child in node["children"]], 

148 } 

149 

150 if root_node is None: 

151 return Response(data={}) 

152 

153 cleaned = clean_node(root_node) 

154 return Response(CustomSegmentSerializer(cleaned).data) 

155 

156 @extend_schema(request=CreateSegmentSerializer, responses=None) 

157 @action(methods=['post'], detail=False, url_path='create-new-segment') 

158 def create_segment(self, request): 

159 serializer = CreateSegmentSerializer(data=request.data) 

160 serializer.is_valid(raise_exception=True) 

161 

162 flowsheet = request.query_params.get("flowsheet") 

163 flowsheet = Flowsheet.objects.get(pk=flowsheet) 

164 flowsheet_state = flowsheet.current_state 

165 parentZone = serializer.validated_data.pop("parentZone") 

166 parentGroup = Grouping.objects.get( 

167 simulationObject__componentName=parentZone, 

168 flowsheet_state=flowsheet_state, 

169 ) 

170 

171 group = parentGroup 

172 

173 stream_data_entry = StreamDataEntry.objects.create( 

174 flowsheet_state=flowsheet_state, 

175 custom=True, 

176 group=group, 

177 streamDataProject=flowsheet_state.StreamDataProject, 

178 ) 

179 

180 segment = Segment.objects.create( 

181 stream_data_entry=stream_data_entry, 

182 flowsheet_state=flowsheet_state, 

183 **serializer.validated_data 

184 ) 

185 segment.area = segment._calc_area() 

186 segment.save(update_fields=["area"]) 

187 

188 return Response(SegmentSerializer(segment).data, status=201) 

189 

190 

191 @extend_schema(request=BulkCreateStreamsSerializer, responses=None) 

192 @action(methods=['post'], detail=False, url_path='bulk-create') 

193 def bulk_create(self, request): 

194 try: 

195 serializer = BulkCreateStreamsSerializer(data=request.data) 

196 serializer.is_valid(raise_exception=True) 

197 validated_data = serializer.validated_data 

198 streamZones = [] 

199 flowsheetID = request.query_params.get("flowsheet") 

200 flowsheet = Flowsheet.objects.get(pk=flowsheetID) 

201 flowsheet_state = flowsheet.current_state 

202 stream_objects = [] 

203 streams = validated_data.get("streams") 

204 for stream in streams: 

205 if stream["parentZone"] not in streamZones: 

206 streamZones.append(stream["parentZone"]) 

207 for zone in streamZones: 

208 if not Grouping.objects.filter(simulationObject__componentName=zone, flowsheet_state=flowsheet_state).exists(): 

209 #This should probably be handled by a bulk create within the upload process for open pinch. 

210 parentGroup = Grouping.create( 

211 flowsheet_state=flowsheet_state, 

212 group=flowsheet_state.root_grouping, 

213 componentName=zone, 

214 visible=False 

215 ) 

216 parentGroup.abstractionType = AbstractionType.Zone 

217 parentGroup.save() 

218 else: 

219 parentGroup = Grouping.objects.get( 

220 simulationObject__componentName=zone, 

221 flowsheet_state=flowsheet_state, 

222 ) 

223 stream_data_entry = StreamDataEntry.objects.create( 

224 flowsheet_state=flowsheet_state, 

225 custom=True, 

226 group=parentGroup, 

227 streamDataProject=flowsheet_state.StreamDataProject, 

228 ) 

229 zoneStreams = [stream for stream in streams if stream["parentZone"]==zone] 

230 for stream in zoneStreams: 

231 del stream["parentZone"] 

232 segment = Segment( 

233 flowsheet_state=flowsheet_state, 

234 stream_data_entry=stream_data_entry, 

235 **stream, 

236 ) 

237 segment.area = segment._calc_area() 

238 stream_objects.append(segment) 

239 streams.remove(stream) 

240 

241 tracked_bulk_create(Segment.objects, stream_objects) 

242 

243 return Response({'status': 'success'}, status=201) 

244 except Exception as e: 

245 return self.error_response(e) 

246 

247 @action(methods=['post'], detail=False, url_path='delete-all') 

248 def delete_all(self, request): 

249 try: 

250 flowsheet = request.query_params.get("flowsheet") 

251 flowsheet = Flowsheet.objects.get(pk=flowsheet) 

252 project = flowsheet.current_state.StreamDataProject 

253 

254 project.Inputs.PinchUtilities.all().delete() 

255 

256 return Response({'status': 'success'}, status=204) 

257 except Exception as e: 

258 return self.error_response(e) 

259 

260 def error_response(self, e): 

261 tb_info = traceback.format_exc() 

262 error_message = str(e) 

263 response_data = {'status': 'error', 

264 'message': error_message, 'traceback': tb_info} 

265 return Response(response_data, status=400)