Skip to content

Commit 47cda65

Browse files
authored
Rework local openml directory (#987)
* Fix #883 #884 #906 #972 * Address Mitar's comments * rework for Windows/OSX, some mypy pleasing due to pre-commit * type fixes and removing unused code
1 parent ab793a6 commit 47cda65

6 files changed

Lines changed: 116 additions & 83 deletions

File tree

‎openml/_api_calls.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -155,7 +155,7 @@ def _read_url_files(url, data=None, file_elements=None):
155155

156156
def __read_url(url, request_method, data=None, md5_checksum=None):
157157
data = {} if data is None else data
158-
if config.apikey is not None:
158+
if config.apikey:
159159
data["api_key"] = config.apikey
160160
return _send_request(
161161
request_method=request_method, url=url, data=data, md5_checksum=md5_checksum

‎openml/config.py‎

Lines changed: 70 additions & 53 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,8 @@
77
import logging
88
import logging.handlers
99
import os
10+
from pathlib import Path
11+
import platform
1012
from typing import Tuple, cast
1113

1214
from io import StringIO
@@ -19,7 +21,7 @@
1921
file_handler = None
2022

2123

22-
def _create_log_handlers():
24+
def _create_log_handlers(create_file_handler=True):
2325
""" Creates but does not attach the log handlers. """
2426
global console_handler, file_handler
2527
if console_handler is not None or file_handler is not None:
@@ -32,12 +34,13 @@ def _create_log_handlers():
3234
console_handler = logging.StreamHandler()
3335
console_handler.setFormatter(output_formatter)
3436

35-
one_mb = 2 ** 20
36-
log_path = os.path.join(cache_directory, "openml_python.log")
37-
file_handler = logging.handlers.RotatingFileHandler(
38-
log_path, maxBytes=one_mb, backupCount=1, delay=True
39-
)
40-
file_handler.setFormatter(output_formatter)
37+
if create_file_handler:
38+
one_mb = 2 ** 20
39+
log_path = os.path.join(cache_directory, "openml_python.log")
40+
file_handler = logging.handlers.RotatingFileHandler(
41+
log_path, maxBytes=one_mb, backupCount=1, delay=True
42+
)
43+
file_handler.setFormatter(output_formatter)
4144

4245

4346
def _convert_log_levels(log_level: int) -> Tuple[int, int]:
@@ -83,15 +86,18 @@ def set_file_log_level(file_output_level: int):
8386

8487
# Default values (see also https://github.com/openml/OpenML/wiki/Client-API-Standards)
8588
_defaults = {
86-
"apikey": None,
89+
"apikey": "",
8790
"server": "https://www.openml.org/api/v1/xml",
88-
"cachedir": os.path.expanduser(os.path.join("~", ".openml", "cache")),
91+
"cachedir": (
92+
os.environ.get("XDG_CACHE_HOME", os.path.join("~", ".cache", "openml",))
93+
if platform.system() == "Linux"
94+
else os.path.join("~", ".openml")
95+
),
8996
"avoid_duplicate_runs": "True",
90-
"connection_n_retries": 10,
91-
"max_retries": 20,
97+
"connection_n_retries": "10",
98+
"max_retries": "20",
9299
}
93100

94-
config_file = os.path.expanduser(os.path.join("~", ".openml", "config"))
95101

96102
# Default values are actually added here in the _setup() function which is
97103
# called at the end of this module
@@ -116,8 +122,8 @@ def get_server_base_url() -> str:
116122
avoid_duplicate_runs = True if _defaults["avoid_duplicate_runs"] == "True" else False
117123

118124
# Number of retries if the connection breaks
119-
connection_n_retries = _defaults["connection_n_retries"]
120-
max_retries = _defaults["max_retries"]
125+
connection_n_retries = int(_defaults["connection_n_retries"])
126+
max_retries = int(_defaults["max_retries"])
121127

122128

123129
class ConfigurationForExamples:
@@ -187,62 +193,78 @@ def _setup():
187193
global connection_n_retries
188194
global max_retries
189195

190-
# read config file, create cache directory
191-
try:
192-
os.mkdir(os.path.expanduser(os.path.join("~", ".openml")))
193-
except FileExistsError:
194-
# For other errors, we want to propagate the error as openml does not work without cache
195-
pass
196+
if platform.system() == "Linux":
197+
config_dir = Path(os.environ.get("XDG_CONFIG_HOME", Path("~") / ".config" / "openml"))
198+
else:
199+
config_dir = Path("~") / ".openml"
200+
# Still use os.path.expanduser to trigger the mock in the unit test
201+
config_dir = Path(os.path.expanduser(config_dir))
202+
config_file = config_dir / "config"
203+
204+
# read config file, create directory for config file
205+
if not os.path.exists(config_dir):
206+
try:
207+
os.mkdir(config_dir)
208+
cache_exists = True
209+
except PermissionError:
210+
cache_exists = False
211+
else:
212+
cache_exists = True
213+
214+
if cache_exists:
215+
_create_log_handlers()
216+
else:
217+
_create_log_handlers(create_file_handler=False)
218+
openml_logger.warning(
219+
"No permission to create OpenML directory at %s! This can result in OpenML-Python "
220+
"not working properly." % config_dir
221+
)
196222

197-
config = _parse_config()
223+
config = _parse_config(config_file)
198224
apikey = config.get("FAKE_SECTION", "apikey")
199225
server = config.get("FAKE_SECTION", "server")
200226

201-
short_cache_dir = config.get("FAKE_SECTION", "cachedir")
202-
cache_directory = os.path.expanduser(short_cache_dir)
227+
cache_dir = config.get("FAKE_SECTION", "cachedir")
228+
cache_directory = os.path.expanduser(cache_dir)
203229

204230
# create the cache subdirectory
205-
try:
206-
os.mkdir(cache_directory)
207-
except FileExistsError:
208-
# For other errors, we want to propagate the error as openml does not work without cache
209-
pass
231+
if not os.path.exists(cache_directory):
232+
try:
233+
os.mkdir(cache_directory)
234+
except PermissionError:
235+
openml_logger.warning(
236+
"No permission to create openml cache directory at %s! This can result in "
237+
"OpenML-Python not working properly." % cache_directory
238+
)
210239

211240
avoid_duplicate_runs = config.getboolean("FAKE_SECTION", "avoid_duplicate_runs")
212-
connection_n_retries = config.get("FAKE_SECTION", "connection_n_retries")
213-
max_retries = config.get("FAKE_SECTION", "max_retries")
241+
connection_n_retries = int(config.get("FAKE_SECTION", "connection_n_retries"))
242+
max_retries = int(config.get("FAKE_SECTION", "max_retries"))
214243
if connection_n_retries > max_retries:
215244
raise ValueError(
216245
"A higher number of retries than {} is not allowed to keep the "
217246
"server load reasonable".format(max_retries)
218247
)
219248

220249

221-
def _parse_config():
250+
def _parse_config(config_file: str):
222251
""" Parse the config file, set up defaults. """
223252
config = configparser.RawConfigParser(defaults=_defaults)
224253

225-
if not os.path.exists(config_file):
226-
# Create an empty config file if there was none so far
227-
fh = open(config_file, "w")
228-
fh.close()
229-
logger.info(
230-
"Could not find a configuration file at %s. Going to "
231-
"create an empty file there." % config_file
232-
)
233-
254+
# The ConfigParser requires a [SECTION_HEADER], which we do not expect in our config file.
255+
# Cheat the ConfigParser module by adding a fake section header
256+
config_file_ = StringIO()
257+
config_file_.write("[FAKE_SECTION]\n")
234258
try:
235-
# The ConfigParser requires a [SECTION_HEADER], which we do not expect in our config file.
236-
# Cheat the ConfigParser module by adding a fake section header
237-
config_file_ = StringIO()
238-
config_file_.write("[FAKE_SECTION]\n")
239259
with open(config_file) as fh:
240260
for line in fh:
241261
config_file_.write(line)
242-
config_file_.seek(0)
243-
config.read_file(config_file_)
262+
except FileNotFoundError:
263+
logger.info("No config file found at %s, using default configuration.", config_file)
244264
except OSError as e:
245-
logger.info("Error opening file %s: %s", config_file, e.message)
265+
logger.info("Error opening file %s: %s", config_file, e.args[0])
266+
config_file_.seek(0)
267+
config.read_file(config_file_)
246268
return config
247269

248270

@@ -257,11 +279,7 @@ def get_cache_directory():
257279
"""
258280
url_suffix = urlparse(server).netloc
259281
reversed_url_suffix = os.sep.join(url_suffix.split(".")[::-1])
260-
if not cache_directory:
261-
_cachedir = _defaults(cache_directory)
262-
else:
263-
_cachedir = cache_directory
264-
_cachedir = os.path.join(_cachedir, reversed_url_suffix)
282+
_cachedir = os.path.join(cache_directory, reversed_url_suffix)
265283
return _cachedir
266284

267285

@@ -297,4 +315,3 @@ def set_cache_directory(cachedir):
297315
]
298316

299317
_setup()
300-
_create_log_handlers()

‎openml/testing.py‎

Lines changed: 0 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -8,15 +8,8 @@
88
import time
99
from typing import Dict, Union, cast
1010
import unittest
11-
import warnings
1211
import pandas as pd
1312

14-
# Currently, importing oslo raises a lot of warning that it will stop working
15-
# under python3.8; remove this once they disappear
16-
with warnings.catch_warnings():
17-
warnings.simplefilter("ignore")
18-
from oslo_concurrency import lockutils
19-
2013
import openml
2114
from openml.tasks import TaskType
2215
from openml.exceptions import OpenMLServerException
@@ -100,13 +93,6 @@ def setUp(self, n_levels: int = 1):
10093
openml.config.avoid_duplicate_runs = False
10194
openml.config.cache_directory = self.workdir
10295

103-
# If we're on travis, we save the api key in the config file to allow
104-
# the notebook tests to read them.
105-
if os.environ.get("TRAVIS") or os.environ.get("APPVEYOR"):
106-
with lockutils.external_lock("config", lock_path=self.workdir):
107-
with open(openml.config.config_file, "w") as fh:
108-
fh.write("apikey = %s" % openml.config.apikey)
109-
11096
# Increase the number of retries to avoid spurious server failures
11197
self.connection_n_retries = openml.config.connection_n_retries
11298
openml.config.connection_n_retries = 10

‎openml/utils.py‎

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -244,7 +244,7 @@ def _list_all(listing_call, output_format="dict", *args, **filters):
244244
limit=batch_size,
245245
offset=current_offset,
246246
output_format=output_format,
247-
**active_filters
247+
**active_filters,
248248
)
249249
except openml.exceptions.OpenMLServerNoResult:
250250
# we want to return an empty dict in this case
@@ -277,9 +277,11 @@ def _create_cache_directory(key):
277277
cache = config.get_cache_directory()
278278
cache_dir = os.path.join(cache, key)
279279
try:
280-
os.makedirs(cache_dir)
281-
except OSError:
282-
pass
280+
os.makedirs(cache_dir, exist_ok=True)
281+
except Exception as e:
282+
raise openml.exceptions.OpenMLCacheException(
283+
f"Cannot create cache directory {cache_dir}."
284+
) from e
283285
return cache_dir
284286

285287

‎tests/test_openml/test_config.py‎

Lines changed: 17 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,15 +1,29 @@
11
# License: BSD 3-Clause
22

3+
import tempfile
34
import os
5+
import unittest.mock
46

57
import openml.config
68
import openml.testing
79

810

911
class TestConfig(openml.testing.TestBase):
10-
def test_config_loading(self):
11-
self.assertTrue(os.path.exists(openml.config.config_file))
12-
self.assertTrue(os.path.isdir(os.path.expanduser("~/.openml")))
12+
@unittest.mock.patch("os.path.expanduser")
13+
@unittest.mock.patch("openml.config.openml_logger.warning")
14+
@unittest.mock.patch("openml.config._create_log_handlers")
15+
def test_non_writable_home(self, log_handler_mock, warnings_mock, expanduser_mock):
16+
with tempfile.TemporaryDirectory(dir=self.workdir) as td:
17+
expanduser_mock.side_effect = (
18+
os.path.join(td, "openmldir"),
19+
os.path.join(td, "cachedir"),
20+
)
21+
os.chmod(td, 0o444)
22+
openml.config._setup()
23+
24+
self.assertEqual(warnings_mock.call_count, 2)
25+
self.assertEqual(log_handler_mock.call_count, 1)
26+
self.assertFalse(log_handler_mock.call_args_list[0][1]["create_file_handler"])
1327

1428

1529
class TestConfigurationForExamples(openml.testing.TestBase):

‎tests/test_utils/test_utils.py‎

Lines changed: 22 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,11 @@
1-
from openml.testing import TestBase
1+
import os
2+
import tempfile
3+
import unittest.mock
4+
25
import numpy as np
3-
import openml
4-
import sys
56

6-
if sys.version_info[0] >= 3:
7-
from unittest import mock
8-
else:
9-
import mock
7+
import openml
8+
from openml.testing import TestBase
109

1110

1211
class OpenMLTaskTest(TestBase):
@@ -20,7 +19,7 @@ def mocked_perform_api_call(call, request_method):
2019
def test_list_all(self):
2120
openml.utils._list_all(listing_call=openml.tasks.functions._list_tasks)
2221

23-
@mock.patch("openml._api_calls._perform_api_call", side_effect=mocked_perform_api_call)
22+
@unittest.mock.patch("openml._api_calls._perform_api_call", side_effect=mocked_perform_api_call)
2423
def test_list_all_few_results_available(self, _perform_api_call):
2524
# we want to make sure that the number of api calls is only 1.
2625
# Although we have multiple versions of the iris dataset, there is only
@@ -86,3 +85,18 @@ def test_list_all_for_evaluations(self):
8685

8786
# might not be on test server after reset, please rerun test at least once if fails
8887
self.assertEqual(len(evaluations), required_size)
88+
89+
@unittest.mock.patch("openml.config.get_cache_directory")
90+
def test__create_cache_directory(self, config_mock):
91+
with tempfile.TemporaryDirectory(dir=self.workdir) as td:
92+
config_mock.return_value = td
93+
openml.utils._create_cache_directory("abc")
94+
self.assertTrue(os.path.exists(os.path.join(td, "abc")))
95+
subdir = os.path.join(td, "def")
96+
os.mkdir(subdir)
97+
os.chmod(subdir, 0o444)
98+
config_mock.return_value = subdir
99+
with self.assertRaisesRegex(
100+
openml.exceptions.OpenMLCacheException, r"Cannot create cache directory",
101+
):
102+
openml.utils._create_cache_directory("ghi")

0 commit comments

Comments
 (0)