Wrap os.path.join to handle LOCALE issues

Closes gh-81.
This commit is contained in:
gfyoung
2017-05-19 16:27:12 -04:00
parent 65f04210c1
commit 2ff5dc212a
+89 -48
View File
@@ -11,6 +11,7 @@ from __future__ import (absolute_import, division,
from glob import glob from glob import glob
import os import os
import locale
import platform import platform
import re import re
import shutil import shutil
@@ -51,38 +52,51 @@ def write_data(f, data):
def list_dir_no_hidden(path): def list_dir_no_hidden(path):
# This function doesn't list hidden files # This function doesn't list hidden files
return glob(os.path.join(path, "*")) return glob(path_join_robust(path, "*"))
# Project Settings # Project Settings
BASEDIR_PATH = os.path.dirname(os.path.realpath(__file__)) BASEDIR_PATH = os.path.dirname(os.path.realpath(__file__))
defaults = {
"numberofrules": 0, def get_defaults():
"datapath": os.path.join(BASEDIR_PATH, "data"), """
"freshen": True, Helper method for getting the default settings.
"replace": False,
"backup": False, Returns
"skipstatichosts": False, -------
"keepdomaincomments": False, default_settings : dict
"extensionspath": os.path.join(BASEDIR_PATH, "extensions"), A dictionary of the default settings when updating host information.
"extensions": [], """
"outputsubfolder": "",
"hostfilename": "hosts", return {
"targetip": "0.0.0.0", "numberofrules": 0,
"ziphosts": False, "datapath": path_join_robust(BASEDIR_PATH, "data"),
"sourcedatafilename": "update.json", "freshen": True,
"sourcesdata": [], "replace": False,
"readmefilename": "readme.md", "backup": False,
"readmetemplate": os.path.join(BASEDIR_PATH, "readme_template.md"), "skipstatichosts": False,
"readmedata": {}, "keepdomaincomments": False,
"readmedatafilename": os.path.join(BASEDIR_PATH, "readmeData.json"), "extensionspath": path_join_robust(BASEDIR_PATH, "extensions"),
"exclusionpattern": "([a-zA-Z\d-]+\.){0,}", "extensions": [],
"exclusionregexs": [], "outputsubfolder": "",
"exclusions": [], "hostfilename": "hosts",
"commonexclusions": ["hulu.com"], "targetip": "0.0.0.0",
"blacklistfile": os.path.join(BASEDIR_PATH, "blacklist"), "ziphosts": False,
"whitelistfile": os.path.join(BASEDIR_PATH, "whitelist")} "sourcedatafilename": "update.json",
"sourcesdata": [],
"readmefilename": "readme.md",
"readmetemplate": path_join_robust(BASEDIR_PATH,
"readme_template.md"),
"readmedata": {},
"readmedatafilename": path_join_robust(BASEDIR_PATH,
"readmeData.json"),
"exclusionpattern": "([a-zA-Z\d-]+\.){0,}",
"exclusionregexs": [],
"exclusions": [],
"commonexclusions": ["hulu.com"],
"blacklistfile": path_join_robust(BASEDIR_PATH, "blacklist"),
"whitelistfile": path_join_robust(BASEDIR_PATH, "whitelist")}
def main(): def main():
@@ -129,12 +143,11 @@ def main():
options = vars(parser.parse_args()) options = vars(parser.parse_args())
options["outputpath"] = os.path.join(BASEDIR_PATH, options["outputpath"] = path_join_robust(BASEDIR_PATH,
options["outputsubfolder"]) options["outputsubfolder"])
options["freshen"] = not options["noupdate"] options["freshen"] = not options["noupdate"]
settings = {} settings = get_defaults()
settings.update(defaults)
settings.update(options) settings.update(options)
settings["sources"] = list_dir_no_hidden(settings["datapath"]) settings["sources"] = list_dir_no_hidden(settings["datapath"])
@@ -161,9 +174,9 @@ def main():
finalize_file(final_file) finalize_file(final_file)
if settings["ziphosts"]: if settings["ziphosts"]:
zf = zipfile.ZipFile(os.path.join(settings["outputsubfolder"], zf = zipfile.ZipFile(path_join_robust(settings["outputsubfolder"],
"hosts.zip"), mode='w') "hosts.zip"), mode='w')
zf.write(os.path.join(settings["outputsubfolder"], "hosts"), zf.write(path_join_robust(settings["outputsubfolder"], "hosts"),
compress_type=zipfile.ZIP_DEFLATED, arcname='hosts') compress_type=zipfile.ZIP_DEFLATED, arcname='hosts')
zf.close() zf.close()
@@ -179,9 +192,9 @@ def main():
# Prompt the User # Prompt the User
def prompt_for_update(): def prompt_for_update():
# Create hosts file if it doesn't exists # Create hosts file if it doesn't exists
if not os.path.isfile(os.path.join(BASEDIR_PATH, "hosts")): if not os.path.isfile(path_join_robust(BASEDIR_PATH, "hosts")):
try: try:
open(os.path.join(BASEDIR_PATH, "hosts"), "w+").close() open(path_join_robust(BASEDIR_PATH, "hosts"), "w+").close()
except: except:
print_failure("ERROR: No 'hosts' file in the folder," print_failure("ERROR: No 'hosts' file in the folder,"
"try creating one manually") "try creating one manually")
@@ -303,9 +316,9 @@ def update_all_sources():
# get rid of carriage-return symbols # get rid of carriage-return symbols
updated_file = updated_file.replace("\r", "") updated_file = updated_file.replace("\r", "")
hosts_file = open(os.path.join(BASEDIR_PATH, hosts_file = open(path_join_robust(BASEDIR_PATH,
os.path.dirname(source), os.path.dirname(source),
settings["hostfilename"]), "wb") settings["hostfilename"]), "wb")
write_data(hosts_file, updated_file) write_data(hosts_file, updated_file)
hosts_file.close() hosts_file.close()
except: except:
@@ -332,12 +345,12 @@ def create_initial_file():
# spin the sources for extensions to the base file # spin the sources for extensions to the base file
for source in settings["extensions"]: for source in settings["extensions"]:
for filename in recursive_glob(os.path.join( for filename in recursive_glob(path_join_robust(
settings["extensionspath"], source), settings["hostfilename"]): settings["extensionspath"], source), settings["hostfilename"]):
with open(filename, "r") as curFile: with open(filename, "r") as curFile:
write_data(merge_file, curFile.read()) write_data(merge_file, curFile.read())
for update_file_path in recursive_glob(os.path.join( for update_file_path in recursive_glob(path_join_robust(
settings["extensionspath"], source), settings["extensionspath"], source),
settings["sourcedatafilename"]): settings["sourcedatafilename"]):
update_file = open(update_file_path, "r") update_file = open(update_file_path, "r")
@@ -366,7 +379,7 @@ def remove_dups_and_excl(merge_file):
os.makedirs(settings["outputpath"]) os.makedirs(settings["outputpath"])
# Another mode is required to read and write the file in Python 3 # Another mode is required to read and write the file in Python 3
final_file = open(os.path.join(settings["outputpath"], "hosts"), final_file = open(path_join_robust(settings["outputpath"], "hosts"),
"w+b" if PY3 else "w+") "w+b" if PY3 else "w+")
merge_file.seek(0) # reset file pointer merge_file.seek(0) # reset file pointer
@@ -466,7 +479,7 @@ def write_opening_header(final_file):
write_data(final_file, "# Fetch the latest version of this file: " write_data(final_file, "# Fetch the latest version of this file: "
"https://raw.githubusercontent.com/" "https://raw.githubusercontent.com/"
"StevenBlack/hosts/master/" + "StevenBlack/hosts/master/" +
os.path.join(settings["outputsubfolder"], "") + "hosts\n") path_join_robust(settings["outputsubfolder"], "") + "hosts\n")
write_data(final_file, "# Project home page: https://github.com/" write_data(final_file, "# Project home page: https://github.com/"
"StevenBlack/hosts\n#\n") "StevenBlack/hosts\n#\n")
write_data(final_file, "# ===============================" write_data(final_file, "# ==============================="
@@ -486,7 +499,7 @@ def write_opening_header(final_file):
write_data(final_file, "127.0.0.53 " + socket.gethostname() + "\n") write_data(final_file, "127.0.0.53 " + socket.gethostname() + "\n")
write_data(final_file, "\n") write_data(final_file, "\n")
preamble = os.path.join(BASEDIR_PATH, "myhosts") preamble = path_join_robust(BASEDIR_PATH, "myhosts")
if os.path.isfile(preamble): if os.path.isfile(preamble):
with open(preamble, "r") as f: with open(preamble, "r") as f:
write_data(final_file, f.read()) write_data(final_file, f.read())
@@ -499,7 +512,7 @@ def update_readme_data():
if settings["extensions"]: if settings["extensions"]:
extensions_key = "-".join(settings["extensions"]) extensions_key = "-".join(settings["extensions"])
generation_data = {"location": os.path.join( generation_data = {"location": path_join_robust(
settings["outputsubfolder"], ""), settings["outputsubfolder"], ""),
"entries": settings["numberofrules"], "entries": settings["numberofrules"],
"sourcesdata": settings["sourcesdata"]} "sourcesdata": settings["sourcesdata"]}
@@ -626,12 +639,12 @@ def flush_dns_cache():
# Hotfix since merging with an already existing # Hotfix since merging with an already existing
# hosts file leads to artifacts and duplicates # hosts file leads to artifacts and duplicates
def remove_old_hosts_file(): def remove_old_hosts_file():
old_file_path = os.path.join(BASEDIR_PATH, "hosts") old_file_path = path_join_robust(BASEDIR_PATH, "hosts")
# create if already removed, so remove wont raise an error # create if already removed, so remove wont raise an error
open(old_file_path, "a").close() open(old_file_path, "a").close()
if settings["backup"]: if settings["backup"]:
backup_file_path = os.path.join(BASEDIR_PATH, "hosts-{}".format( backup_file_path = path_join_robust(BASEDIR_PATH, "hosts-{}".format(
time.strftime("%Y-%m-%d-%H-%M-%S"))) time.strftime("%Y-%m-%d-%H-%M-%S")))
# Make a backup copy, marking the date in which the list was updated # Make a backup copy, marking the date in which the list was updated
@@ -720,10 +733,38 @@ def recursive_glob(stem, file_pattern):
matches = [] matches = []
for root, dirnames, filenames in os.walk(stem): for root, dirnames, filenames in os.walk(stem):
for filename in fnmatch.filter(filenames, file_pattern): for filename in fnmatch.filter(filenames, file_pattern):
matches.append(os.path.join(root, filename)) matches.append(path_join_robust(root, filename))
return matches return matches
def path_join_robust(path_one, path_two):
"""
Wrapper around `os.path.join` with handling for locale issues.
Parameters
----------
path_one : str
The first path to join.
path_two : str
The second path to join.
Returns
-------
joined_path : str
The joined path string of the two path inputs.
Raises
------
locale.Error : A locale issue was detected that prevents path joining.
"""
try:
return os.path.join(path_one, path_two)
except UnicodeDecodeError as e:
raise locale.Error("Unable to construct path. This is "
"likely a LOCALE issue:\n\n" + str(e))
# Colors # Colors
class Colors(object): class Colors(object):
PROMPT = "\033[94m" PROMPT = "\033[94m"