diff --git a/deepaas/cmd/cli.py b/deepaas/cmd/cli.py index a41fc351..d232e07e 100644 --- a/deepaas/cmd/cli.py +++ b/deepaas/cmd/cli.py @@ -203,8 +203,8 @@ def _get_model_name(): model_name, model_obj = _get_model_name() # use deepaas.model.v2.wrapper.ModelWrapper(). deepaas>1.2.1dev4 -# model_obj = v2_wrapper.ModelWrapper(name=model_name, -# model_obj=model_obj) +model_obj = v2_wrapper.ModelWrapper(name=model_name, + model_obj=model_obj) # Once we know the model name, # we get arguments for predict and train as dictionaries diff --git a/deepaas/tests/test_cli_fix.py b/deepaas/tests/test_cli_fix.py new file mode 100644 index 00000000..875756a8 --- /dev/null +++ b/deepaas/tests/test_cli_fix.py @@ -0,0 +1,117 @@ +# -*- coding: utf-8 -*- + +# Copyright 2024 Spanish National Research Council (CSIC) +# +# Licensed under the Apache License, Version 2.0 (the "License"); you may +# not use this file except in compliance with the License. You may obtain +# a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, WITHOUT +# WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the +# License for the specific language governing permissions and limitations +# under the License. + +"""Tests for CLI fix handling models with missing get_predict_args or get_train_args""" + +import unittest + +from webargs import fields +from deepaas.model.v2 import wrapper as v2_wrapper + + +class TestCLIModelHandling(unittest.TestCase): + """Test CLI functionality with models that have missing methods""" + + def test_model_wrapper_handles_missing_predict_args(self): + """Test that ModelWrapper handles models without get_predict_args""" + + class TrainOnlyModel: + def get_metadata(self): + return {"name": "train-only", "version": "1.0"} + + def train(self, **kwargs): + return {"status": "trained"} + + def get_train_args(self): + return {"epochs": fields.Int(required=True)} + + # NOTE: Missing get_predict_args method + + model = TrainOnlyModel() + wrapper = v2_wrapper.ModelWrapper("train-only", model, None) + + # Should not raise AttributeError + predict_args = wrapper.get_predict_args() + train_args = wrapper.get_train_args() + + self.assertEqual(predict_args, {}) + self.assertIn("epochs", train_args) + + def test_model_wrapper_handles_missing_train_args(self): + """Test that ModelWrapper handles models without get_train_args""" + + class PredictOnlyModel: + def get_metadata(self): + return {"name": "predict-only", "version": "1.0"} + + def predict(self, **kwargs): + return {"prediction": "result"} + + def get_predict_args(self): + return {"data": fields.Str(required=True)} + + # NOTE: Missing get_train_args method + + model = PredictOnlyModel() + wrapper = v2_wrapper.ModelWrapper("predict-only", model, None) + + # Should not raise AttributeError + predict_args = wrapper.get_predict_args() + train_args = wrapper.get_train_args() + + self.assertIn("data", predict_args) + self.assertEqual(train_args, {}) + + def test_model_wrapper_handles_missing_both_args(self): + """Test that ModelWrapper handles models without both args methods""" + + class MinimalModel: + def get_metadata(self): + return {"name": "minimal", "version": "1.0"} + + # NOTE: Missing both get_predict_args and get_train_args methods + + model = MinimalModel() + wrapper = v2_wrapper.ModelWrapper("minimal", model, None) + + # Should not raise AttributeError for either method + predict_args = wrapper.get_predict_args() + train_args = wrapper.get_train_args() + + self.assertEqual(predict_args, {}) + self.assertEqual(train_args, {}) + + def test_direct_model_fails_as_expected(self): + """Test that direct model access fails as expected (validating the problem)""" + + class PredictOnlyModel: + def get_predict_args(self): + return {"data": fields.Str(required=True)} + # NOTE: Missing get_train_args method + + model = PredictOnlyModel() + + # Direct access should work for existing method + predict_args = model.get_predict_args() + self.assertIn("data", predict_args) + + # Direct access should fail for missing method + with self.assertRaises(AttributeError): + model.get_train_args() + + +if __name__ == '__main__': + unittest.main() \ No newline at end of file diff --git a/deepaas/tests/test_cmd.py b/deepaas/tests/test_cmd.py index 10029284..85a4b017 100644 --- a/deepaas/tests/test_cmd.py +++ b/deepaas/tests/test_cmd.py @@ -24,6 +24,8 @@ from deepaas.cmd import execute from deepaas.cmd import run from deepaas.tests import base +from deepaas.model.v2 import wrapper as v2_wrapper +from webargs import fields class TestRun(base.TestCase): @@ -112,3 +114,76 @@ def test_execute_ct(self, m_out_pred): execute.main() shutil.rmtree(output_dir) os.remove(output_dir + ".zip") + + +class TestCLIModels(base.TestCase): + """Test CLI functionality with models that have missing methods""" + + def test_model_wrapper_handles_missing_predict_args(self): + """Test that ModelWrapper handles models without get_predict_args""" + + class TrainOnlyModel: + def get_metadata(self): + return {"name": "train-only", "version": "1.0"} + + def train(self, **kwargs): + return {"status": "trained"} + + def get_train_args(self): + return {"epochs": fields.Int(required=True)} + + # NOTE: Missing get_predict_args method + + model = TrainOnlyModel() + wrapper = v2_wrapper.ModelWrapper("train-only", model, None) + + # Should not raise AttributeError + predict_args = wrapper.get_predict_args() + train_args = wrapper.get_train_args() + + self.assertEqual(predict_args, {}) + self.assertIn("epochs", train_args) + + def test_model_wrapper_handles_missing_train_args(self): + """Test that ModelWrapper handles models without get_train_args""" + + class PredictOnlyModel: + def get_metadata(self): + return {"name": "predict-only", "version": "1.0"} + + def predict(self, **kwargs): + return {"prediction": "result"} + + def get_predict_args(self): + return {"data": fields.Str(required=True)} + + # NOTE: Missing get_train_args method + + model = PredictOnlyModel() + wrapper = v2_wrapper.ModelWrapper("predict-only", model, None) + + # Should not raise AttributeError + predict_args = wrapper.get_predict_args() + train_args = wrapper.get_train_args() + + self.assertIn("data", predict_args) + self.assertEqual(train_args, {}) + + def test_model_wrapper_handles_missing_both_args(self): + """Test that ModelWrapper handles models without both args methods""" + + class MinimalModel: + def get_metadata(self): + return {"name": "minimal", "version": "1.0"} + + # NOTE: Missing both get_predict_args and get_train_args methods + + model = MinimalModel() + wrapper = v2_wrapper.ModelWrapper("minimal", model, None) + + # Should not raise AttributeError for either method + predict_args = wrapper.get_predict_args() + train_args = wrapper.get_train_args() + + self.assertEqual(predict_args, {}) + self.assertEqual(train_args, {})