Coverage for backend/django/core/auxiliary/models/MLModel.py: 96%

48 statements  

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

1from django.db import models 

2from django.contrib.postgres.fields import ArrayField 

3from flowsheetInternals.unitops.models import SimulationObject 

4 

5from core.managers import AccessControlManager, AllFlowsheetStatesManager 

6 

7 

8class MLModel(models.Model): 

9 class ModelType(models.TextChoices): 

10 POLYNOMIAL_REGRESSION = "polynomial_regression", "Polynomial Regression" 

11 RBF_REGRESSION = "rbf_regression", "RBF Regression" 

12 

13 class ResultState(models.TextChoices): 

14 PENDING = "pending", "Pending" 

15 TRAINING = "training", "Training" 

16 READY = "ready", "Ready" 

17 READY_WITHOUT_DIAGNOSTICS = ( 

18 "ready_without_diagnostics", 

19 "Ready without diagnostics", 

20 ) 

21 FAILED = "failed", "Failed" 

22 

23 flowsheet_state = models.ForeignKey( 

24 "FlowsheetState", 

25 on_delete=models.CASCADE, 

26 related_name="MLModels", 

27 ) 

28 simulationObject = models.ForeignKey(SimulationObject, on_delete=models.CASCADE, related_name="MLModels", null=True) 

29 surrogate_model = models.JSONField(default=dict, blank=True, null=True) 

30 model_type = models.CharField(default="rbf_regression", max_length=40, choices=ModelType.choices) 

31 csv_file_name = models.CharField(max_length=255, blank=True, default="") 

32 csv_bucket = models.CharField(max_length=255, blank=True, default="") 

33 csv_object_key = models.CharField(max_length=1024, blank=True, default="") 

34 csv_headers = models.JSONField(default=list, blank=True) 

35 csv_delimiter = models.CharField(max_length=1, blank=True, default="") 

36 csv_upload_session = models.ForeignKey( 

37 "UploadSession", 

38 on_delete=models.SET_NULL, 

39 related_name="ml_models", 

40 null=True, 

41 blank=True, 

42 ) 

43 

44 active_step = models.IntegerField(default=0) 

45 return_step = models.IntegerField(null=True, blank=True) 

46 completed_steps = ArrayField( 

47 models.IntegerField(), 

48 size=None, 

49 blank=True, 

50 null=True, 

51 default=list 

52 ) 

53 

54 charts = models.JSONField(default=list, blank=True, null=True) 

55 metrics = models.JSONField(default=list, blank=True, null=True) 

56 result_state = models.CharField( 

57 max_length=32, 

58 choices=ResultState.choices, 

59 default=ResultState.PENDING, 

60 ) 

61 mapping_update_snapshot = models.JSONField(default=None, blank=True, null=True) 

62 test_results_bucket = models.CharField(max_length=255, blank=True, default="") 

63 test_results_key = models.CharField(max_length=1024, blank=True, default="") 

64 displayName = models.CharField(default="ML Model", max_length=64) 

65 created_at = models.DateTimeField(auto_now_add=True) 

66 is_updating = models.BooleanField(default=False) 

67 is_resetting = models.BooleanField(default=False) 

68 

69 objects = AccessControlManager() 

70 all_states = AllFlowsheetStatesManager() 

71 

72 def save(self, *args, **kwargs): 

73 # Track the previous active_step if active_step has changed 

74 if self.pk: # Only if the instance already exists in the database 

75 try: 

76 db_instance = MLModel.objects.get(pk=self.pk) 

77 if db_instance.active_step != self.active_step: 

78 # active_step has changed, so save the previous value as return_step 

79 self.return_step = db_instance.active_step 

80 except MLModel.DoesNotExist: 

81 pass 

82 super().save(*args, **kwargs) 

83 

84