diff --git a/openml/tasks/task.py b/openml/tasks/task.py index ab3cb3da48..c03809dfa2 100644 --- a/openml/tasks/task.py +++ b/openml/tasks/task.py @@ -240,6 +240,11 @@ def _to_dict(self) -> dict[str, dict[str, int | str | list[dict[str, Any]]]]: def _parse_publish_response(self, xml_response: dict) -> None: """Parse the id from the xml_response and assign it to self.""" self.task_id = int(xml_response["oml:upload_task"]["oml:id"]) + from openml.tasks.functions import _get_task_description + + completed = _get_task_description(self.task_id) + self.estimation_procedure = completed.estimation_procedure + self.estimation_procedure_id = completed.estimation_procedure_id class OpenMLSupervisedTask(OpenMLTask, ABC): diff --git a/tests/files/mock_responses/tasks/task_upload_successful.xml b/tests/files/mock_responses/tasks/task_upload_successful.xml new file mode 100644 index 0000000000..e8d3a73e87 --- /dev/null +++ b/tests/files/mock_responses/tasks/task_upload_successful.xml @@ -0,0 +1,3 @@ + + 999 + \ No newline at end of file diff --git a/tests/test_tasks/test_task_functions.py b/tests/test_tasks/test_task_functions.py index bf2fcfeae8..7e2cc73c3c 100644 --- a/tests/test_tasks/test_task_functions.py +++ b/tests/test_tasks/test_task_functions.py @@ -315,3 +315,50 @@ def test_delete_unknown_task(mock_delete, test_files_directory, test_server_v1, task_url = test_server_v1 + "task/9999999" assert task_url == mock_delete.call_args.args[0] assert test_apikey_v1 == mock_delete.call_args.kwargs.get("params", {}).get("api_key") + + +@mock.patch("openml.tasks.functions._get_task_description") +@mock.patch.object(requests.Session, "post") +def test_create_task_publish_populates_estimation_procedure( + mock_post, + mock_get_task_description, + test_files_directory, +): + """After publish(), estimation_procedure must be populated from the server.""" + # Mock the POST (publish) response + content_file = ( + test_files_directory / "mock_responses" / "tasks" / "task_upload_successful.xml" + ) + mock_post.return_value = create_request_response( + status_code=200, + content_filepath=content_file, + ) + + # Mock the follow-up GET (what _get_task_description returns) + from openml.tasks import OpenMLClassificationTask, TaskType + fake_task = OpenMLClassificationTask( + task_id=999, + task_type_id=TaskType.SUPERVISED_CLASSIFICATION, + task_type="Supervised Classification", + data_set_id=128, + target_name="class", + estimation_procedure_id=1, + estimation_procedure_type="crossvalidation", + estimation_parameters={"number_folds": "10", "number_repeats": "1"}, + data_splits_url="https://www.openml.org/api_splits/get/999/Task_999_splits.arff", + ) + mock_get_task_description.return_value = fake_task + + task = openml.tasks.create_task( + task_type=TaskType.SUPERVISED_CLASSIFICATION, + dataset_id=128, + target_name="class", + evaluation_measure="predictive_accuracy", + estimation_procedure_id=1, + ) + task.publish() + + assert task.task_id == 999 + assert task.estimation_procedure["type"] == "crossvalidation" + assert task.estimation_procedure["data_splits_url"] is not None + assert task.estimation_procedure["parameters"] is not None \ No newline at end of file