Improve the typing coverage

This commit is contained in:
Christoph Reiter
2019-08-27 23:44:18 +02:00
parent dd31279f1b
commit 236188fed2
5 changed files with 127 additions and 91 deletions

1
.gitignore vendored
View File

@@ -6,3 +6,4 @@ package-lock.json
*srcinfo.json
.mypy_cache
_cache
.vscode

175
main.py
View File

@@ -1,6 +1,5 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
# Copyright 2016 Christoph Reiter
# Copyright 2016-2019 Christoph Reiter
#
# Permission is hereby granted, free of charge, to any person obtaining
# a copy of this software and associated documentation files (the
@@ -21,6 +20,8 @@
# TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE
# SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
from __future__ import annotations
import argparse
import traceback
import contextlib
@@ -39,11 +40,11 @@ import hashlib
from itertools import zip_longest
from functools import cmp_to_key, wraps
from urllib.parse import quote_plus, quote
from typing import List, Set, Dict, Tuple, Optional, Generator, Any
from typing import List, Set, Dict, Tuple, Optional, Generator, Any, Type, Callable, Union
import requests
from flask import Flask, render_template, request, url_for, redirect, \
make_response, Blueprint, abort, jsonify
make_response, Blueprint, abort, jsonify, Request
from jinja2 import StrictUndefined
@@ -124,7 +125,7 @@ class AppState:
self._last_update = 0.0
self._sources: List[Source] = []
self._sourceinfos: Dict[str, SrcInfoPackage] = {}
self._versions: Dict[str, Tuple[str, str, str]] = {}
self._versions: Dict[str, Tuple[str, str, int]] = {}
self._update_etag()
def _update_etag(self) -> None:
@@ -140,29 +141,29 @@ class AppState:
return self._etag
@property
def sources(self) -> "List[Source]":
def sources(self) -> List[Source]:
return self._sources
@sources.setter
def sources(self, sources: "List[Source]") -> None:
def sources(self, sources: List[Source]) -> None:
self._sources = sources
self._update_etag()
@property
def sourceinfos(self) -> "Dict[str,SrcInfoPackage]":
def sourceinfos(self) -> Dict[str, SrcInfoPackage]:
return self._sourceinfos
@sourceinfos.setter
def sourceinfos(self, sourceinfos):
def sourceinfos(self, sourceinfos: Dict[str, SrcInfoPackage]) -> None:
self._sourceinfos = sourceinfos
self._update_etag()
@property
def versions(self):
def versions(self) -> Dict[str, Tuple[str, str, int]]:
return self._versions
@versions.setter
def versions(self, versions):
def versions(self, versions: Dict[str, Tuple[str, str, int]]) -> None:
self._versions = versions
self._update_etag()
@@ -175,10 +176,10 @@ state = AppState()
packages = Blueprint('packages', __name__, template_folder='templates')
def cache_route(f):
def cache_route(f: Callable) -> Callable:
@wraps(f)
def wrapper(*args, **kwargs):
def wrapper(*args: Any, **kwargs: Any) -> Any:
global state
response = make_response()
@@ -230,23 +231,26 @@ def cleanup_files(files: List[str]) -> List[str]:
return result[::-1]
PackageKey = Tuple[str, str, str, str, str]
class Package:
def __init__(self, builddate, csize, depends, filename, files, isize,
makedepends, md5sum, name, pgpsig, sha256sum, arch,
base_url, repo, repo_variant, provides, conflicts, replaces,
version, base, desc, groups, licenses, optdepends,
checkdepends):
def __init__(self, builddate: str, csize: str, depends: List[str], filename: str, files: List[str], isize: str,
makedepends: List[str], md5sum: str, name: str, pgpsig: str, sha256sum: str, arch: str,
base_url: str, repo: str, repo_variant: str, provides: List[str], conflicts: List[str], replaces: List[str],
version: str, base: str, desc: str, groups: List[str], licenses: List[str], optdepends: List[str],
checkdepends: List[str]) -> None:
self.builddate = int(builddate)
self.csize = csize
def split_depends(deps):
def split_depends(deps: List[str]) -> List[Tuple[str, str]]:
r = []
for d in deps:
parts = re.split("([<>=]+)", d, 1)
first = parts[0].strip()
second = "".join(parts[1:]).strip()
r.append([first, second])
r.append((first, second))
return r
self.depends = split_depends(depends)
@@ -271,24 +275,25 @@ class Package:
self.desc = desc
self.groups = groups
self.licenses = licenses
self.rdepends = []
self.rdepends: List[Tuple[Package, str]] = []
def split_opt(deps):
def split_opt(deps: List[str]) -> List[Tuple[str, str]]:
r = []
for d in deps:
if ":" in d:
r.append([p.strip() for p in d.split(":", 1)])
a, b = d.split(":", 1)
r.append((a.strip(), b.strip()))
else:
r.append([d.strip(), ""])
r.append((d.strip(), ""))
return r
self.optdepends = split_opt(optdepends)
def __repr__(self):
def __repr__(self) -> str:
return "Package(%s)" % self.fileurl
@property
def realprovides(self):
def realprovides(self) -> Dict[str, str]:
prov = {}
for key, info in self.provides.items():
if key.startswith("mingw"):
@@ -297,24 +302,24 @@ class Package:
return prov
@property
def realname(self):
def realname(self) -> str:
if self.repo.startswith("mingw"):
return self.name.split("-", 3)[-1]
return self.name
@property
def git_version(self):
def git_version(self) -> str:
if self.name in state.sourceinfos:
return state.sourceinfos[self.name].build_version
return ""
@property
def key(self):
def key(self) -> PackageKey:
return (self.repo, self.repo_variant,
self.name, self.arch, self.fileurl)
@classmethod
def from_desc(cls, d, base, base_url, repo, repo_variant):
def from_desc(cls: Type[Package], d: Dict[str, List[str]], base: str, base_url: str, repo: str, repo_variant: str) -> Package:
return cls(d["%BUILDDATE%"][0], d["%CSIZE%"][0],
d.get("%DEPENDS%", []), d["%FILENAME%"][0],
d.get("%FILES%", []), d["%ISIZE%"][0],
@@ -340,14 +345,14 @@ class Source:
self._repo = repo
self._repo_variant = repo_variant
self.packages: Dict[str, Package] = {}
self.packages: Dict[PackageKey, Package] = {}
@property
def repos(self):
def repos(self) -> List[str]:
return sorted(set([p.repo for p in self.packages.values()]))
@property
def arches(self):
def arches(self) -> List[str]:
return sorted(set([p.arch for p in self.packages.values()]))
@property
@@ -377,14 +382,14 @@ class Source:
return sorted(licenses)
@property
def arch_url(self):
def arch_url(self) -> str:
arch_info = get_arch_info_for_base(self)
if arch_info is not None:
return arch_info[1]
return ""
@property
def upstream_version(self):
def upstream_version(self) -> str:
arch_info = get_arch_info_for_base(self)
if arch_info is not None:
return extract_upstream_version(arch_info[0])
@@ -407,7 +412,7 @@ class Source:
return self.name
@property
def date(self):
def date(self) -> int:
"""The build date of the newest package"""
return sorted([p.builddate for p in self.packages.values()])[-1]
@@ -469,7 +474,7 @@ class Source:
self.packages[p.key] = p
def get_content_cached(url, *args, **kwargs):
def get_content_cached(url: str, *args: Any, **kwargs: Any) -> bytes:
if not CACHE_LOCAL:
r = requests.get(url, *args, **kwargs)
return r.content
@@ -550,14 +555,15 @@ def _jinja2_filter_filesize(d: int) -> str:
@packages.context_processor
def funcs():
def funcs() -> Dict[str, Callable]:
def is_endpoint(value):
def is_endpoint(value: str) -> bool:
if value.startswith(".") and request.blueprint is not None:
value = request.blueprint + value
return value == request.endpoint
def package_url(package, name=None):
def package_url(package: Package, name: str = None) -> str:
res: str = ""
if name is None:
res = url_for(".package", name=name or package.name)
res += "?repo=" + package.repo
@@ -570,17 +576,17 @@ def funcs():
res += "&variant=" + package.repo_variant
return res
def package_name(package, name=None):
def package_name(package: Package, name: str = None) -> str:
name = name or package.name
name = re.split("[<>=]+", name, 1)[0]
return (name or package.name) + (
"/" + package.repo_variant if package.repo_variant else "")
def package_restriction(package, name=None):
def package_restriction(package: Package, name: str = None) -> str:
name = name or package.name
return name[len(re.split("[<>=]+", name)[0]):].strip()
def update_timestamp():
def update_timestamp() -> float:
global state
return state.last_update
@@ -590,23 +596,26 @@ def funcs():
update_timestamp=update_timestamp, is_endpoint=is_endpoint)
RouteResponse = Any
@packages.route('/repos')
@cache_route
def repos():
def repos() -> RouteResponse:
global REPOSITORIES
return render_template('repos.html', repos=REPOSITORIES)
@packages.route('/')
def index():
def index() -> RouteResponse:
return redirect(url_for('.updates'))
@packages.route('/base')
@packages.route('/base/<name>')
@cache_route
def base(name=None):
def base(name: str = None) -> RouteResponse:
global state
if name is not None:
@@ -619,7 +628,7 @@ def base(name=None):
@packages.route('/group/')
@packages.route('/group/<name>')
@cache_route
def group(name=None):
def group(name: Optional[str] = None) -> RouteResponse:
global state
if name is not None:
@@ -631,7 +640,7 @@ def group(name=None):
return render_template('group.html', name=name, packages=res)
else:
groups = {}
groups: Dict[str, int] = {}
for s in state.sources:
for k, p in sorted(s.packages.items()):
for name in p.groups:
@@ -641,7 +650,7 @@ def group(name=None):
@packages.route('/package/<name>')
@cache_route
def package(name):
def package(name: str) -> RouteResponse:
global state
repo = request.args.get('repo')
@@ -659,10 +668,10 @@ def package(name):
@packages.route('/updates')
@cache_route
def updates():
def updates() -> RouteResponse:
global state
packages = []
packages: List[Package] = []
for s in state.sources:
packages.extend(s.packages.values())
packages.sort(key=lambda p: p.builddate, reverse=True)
@@ -767,7 +776,7 @@ def get_arch_names(name: str) -> List[str]:
names: List[str] = []
def add(n):
def add(n: str) -> None:
if n not in names:
names.append(n)
@@ -845,9 +854,18 @@ def vercmp(v1: str, v2: str) -> int:
def cmp(a: int, b: int) -> int:
return (a > b) - (a < b)
def split(v):
e, v = v.split("~", 1) if "~" in v else ("0", v)
v, r = v.rsplit("-", 1) if "-" in v else (v, None)
def split(v: str) -> Tuple[str, str, Optional[str]]:
if "~" in v:
e, v = v.split("~", 1)
else:
e, v = ("0", v)
r: Optional[str] = None
if "-" in v:
v, r = v.rsplit("-", 1)
else:
v, r = (v, None)
return (e, v, r)
digit, alpha, other = range(3)
@@ -861,8 +879,8 @@ def vercmp(v1: str, v2: str) -> int:
else:
return other
def parse(v):
parts = []
def parse(v: str) -> List[Tuple[int, Optional[str]]]:
parts: List[Tuple[int, Optional[str]]] = []
seps = 0
current = ""
for c in v:
@@ -948,7 +966,7 @@ def update_versions() -> None:
global VERSION_CONFIG, state
print("update versions")
arch_versions: Dict[str, Tuple[str, str, str]] = {}
arch_versions: Dict[str, Tuple[str, str, int]] = {}
for (url, repo, variant) in VERSION_CONFIG:
for source in parse_repo(repo, variant, url).values():
msys_ver = arch_version_to_msys(source.version)
@@ -1017,7 +1035,7 @@ def extract_upstream_version(version: str) -> str:
"-")[0].split("+", 1)[0].split("~", 1)[-1].split(":", 1)[-1]
def get_arch_info_for_base(s: Source) -> Optional[tuple]:
def get_arch_info_for_base(s: Source) -> Optional[Tuple[str, str, int]]:
"""tuple or None"""
global state
@@ -1033,13 +1051,13 @@ def get_arch_info_for_base(s: Source) -> Optional[tuple]:
for realname in variants:
for arch_name in get_arch_names(realname):
if arch_name in state.versions:
return tuple(state.versions[arch_name])
return state.versions[arch_name]
return None
@packages.route('/outofdate')
@cache_route
def outofdate():
def outofdate() -> RouteResponse:
global state
missing = []
@@ -1085,7 +1103,7 @@ def outofdate():
@packages.route('/queue')
@cache_route
def queue():
def queue() -> RouteResponse:
global state
# Create entries for all packages where the version doesn't match
@@ -1109,7 +1127,7 @@ def queue():
@packages.route('/new')
@cache_route
def new():
def new() -> RouteResponse:
global state
# Create dummy entries for all GIT only packages
@@ -1131,7 +1149,7 @@ def new():
@packages.route('/removals')
@cache_route
def removals():
def removals() -> RouteResponse:
global state
# get all packages in the pacman repo which are no in GIT
@@ -1147,9 +1165,9 @@ def removals():
@packages.route('/python2')
@cache_route
def test():
def test() -> RouteResponse:
def is_split_package(p):
def is_split_package(p: Package) -> bool:
c = 0
for name, type_ in p.makedepends:
if name == "mingw-w64-x86_64-python3":
@@ -1160,7 +1178,7 @@ def test():
return True
return False
def get_rdep_count(p):
def get_rdep_count(p: Package) -> int:
todo = {p.name: p}
done = set()
while todo:
@@ -1190,7 +1208,7 @@ def test():
@packages.route('/search')
@cache_route
def search():
def search() -> str:
global state
query = request.args.get('q', '')
@@ -1200,7 +1218,7 @@ def search():
qtype = "pkg"
parts = query.split()
res_pkg = []
res_pkg: List[Union[Package, Source]] = []
if not query:
pass
@@ -1245,14 +1263,14 @@ def trigger_appveyor_build(account: str, project: str, token: str) -> str:
account, project, build_id)
def check_github_signature(request, secret):
def check_github_signature(request: Request, secret: str) -> bool:
signature = request.headers.get('X-Hub-Signature', '')
mac = hmac.new(secret.encode("utf-8"), request.get_data(), hashlib.sha1)
return hmac.compare_digest("sha1=" + mac.hexdigest(), signature)
@packages.route("/webhook", methods=['POST'])
def github_payload():
def github_payload() -> RouteResponse:
secret = os.environ.get("GITHUB_WEBHOOK_SECRET")
if not secret:
abort(500, 'webhook secret config incomplete')
@@ -1267,7 +1285,7 @@ def github_payload():
account = os.environ.get("APPVEYOR_ACCOUNT")
project = os.environ.get("APPVEYOR_PROJECT")
token = os.environ.get("APPVEYOR_TOKEN")
if not all([account, project, token]):
if not account or not project or not token:
abort(500, 'appveyor config incomplete')
build_url = trigger_appveyor_build(account, project, token)
return jsonify({'msg': 'triggered a build: %s' % build_url})
@@ -1350,9 +1368,10 @@ def fill_rdepends(sources: List[Source]) -> None:
for s in sources:
for p in s.packages.values():
rdepends = list(deps.get(p.name, []))
rdepends = list(deps.get(p.name, set()))
for prov in p.provides:
rdepends += list(deps.get(prov, []))
rdepends += list(deps.get(prov, set()))
p.rdepends = sorted(rdepends, key=lambda e: (e[0].key, e[1]))
# filter out other arches for msys packages
@@ -1381,7 +1400,7 @@ def update_thread() -> None:
time.sleep(UPDATE_INTERVAL)
def start_update_thread():
def start_update_thread() -> None:
thread = threading.Thread(target=update_thread)
thread.daemon = True
thread.start()
@@ -1418,7 +1437,7 @@ class SrcInfoPackage(object):
version = "%s~%s" % (self.epoch, version)
return version
def __repr__(self):
def __repr__(self) -> str:
return "<%s %s %s>" % (
type(self).__name__, self.pkgname, self.build_version)
@@ -1465,7 +1484,7 @@ app.jinja_env.undefined = StrictUndefined
start_update_thread()
def main(argv: List[str]) -> Any:
def main(argv: List[str]) -> Optional[Union[int, str]]:
global CACHE_LOCAL
parser = argparse.ArgumentParser()
@@ -1480,6 +1499,8 @@ def main(argv: List[str]) -> Any:
print("http://localhost:%d" % args.port)
app.run(port=args.port, debug=args.debug)
return None
if __name__ == "__main__":
sys.exit(main(sys.argv))

View File

@@ -10,6 +10,8 @@ warn_return_any=True
warn_unused_ignores=True
strict_equality=True
disallow_incomplete_defs=True
disallow_untyped_calls=True
disallow_untyped_defs=True
[tool:pytest]
filterwarnings =

View File

@@ -1,3 +1,5 @@
# type: ignore
import pytest
from main import app

View File

@@ -30,8 +30,15 @@ import subprocess
from multiprocessing.pool import ThreadPool
from multiprocessing import cpu_count
from typing import List, Iterator, Tuple, Dict, Optional, Union
def get_cache_key(pkgbuild_path):
CacheEntry = Dict[str, str]
CacheTuple = Tuple[str, CacheEntry]
Cache = Dict[str, CacheEntry]
def get_cache_key(pkgbuild_path: str) -> str:
pkgbuild_path = os.path.abspath(pkgbuild_path)
git_cwd = os.path.dirname(pkgbuild_path)
git_path = os.path.relpath(pkgbuild_path, git_cwd)
@@ -53,7 +60,7 @@ def get_cache_key(pkgbuild_path):
return h.hexdigest()
def fixup_makepkg_output(text):
def fixup_makepkg_output(text: str) -> str:
# makepkg-mingw runs makepkg twice for mingw32/64. In case of msys
# packages this results in the output geting duplicated.
# Dedup the output so we can use makepkg-mingw for all packages and still
@@ -63,7 +70,7 @@ def fixup_makepkg_output(text):
return text
def get_srcinfo_for_pkgbuild(pkgbuild_path):
def get_srcinfo_for_pkgbuild(pkgbuild_path: str) -> Optional[CacheTuple]:
pkgbuild_path = os.path.abspath(pkgbuild_path)
git_cwd = os.path.dirname(pkgbuild_path)
git_path = os.path.relpath(pkgbuild_path, git_cwd)
@@ -97,12 +104,12 @@ def get_srcinfo_for_pkgbuild(pkgbuild_path):
meta = {"repo": repo, "path": relpath, "date": date, "srcinfo": text}
except subprocess.CalledProcessError as e:
print("ERROR: %s %s" % (pkgbuild_path, e.output.splitlines()))
return
return None
return (key, meta)
def iter_pkgbuild_paths(repo_path):
def iter_pkgbuild_paths(repo_path: str) -> Iterator[str]:
repo_path = os.path.abspath(repo_path)
print("Searching for PKGBUILD files in %s" % repo_path)
for base, dirs, files in os.walk(repo_path):
@@ -114,7 +121,7 @@ def iter_pkgbuild_paths(repo_path):
yield path
def get_srcinfo_from_cache(args):
def get_srcinfo_from_cache(args: Tuple[str, Cache]) -> Tuple[str, Optional[CacheTuple]]:
pkgbuild_path, cache = args
key = get_cache_key(pkgbuild_path)
if key in cache:
@@ -123,14 +130,14 @@ def get_srcinfo_from_cache(args):
return (pkgbuild_path, None)
def iter_srcinfo(repo_paths, cache):
def iter_srcinfo(repo_paths: List[str], cache: Cache) -> Iterator[Optional[CacheTuple]]:
to_check = []
for repo_path in repo_paths:
for pkgbuild_path in iter_pkgbuild_paths(repo_path):
to_check.append(pkgbuild_path)
pool = ThreadPool(cpu_count() * 3)
to_parse = []
to_parse: List[str] = []
pool_iter = pool.imap_unordered(
get_srcinfo_from_cache, ((p, cache) for p in to_check))
for pkgbuild_path, srcinfo in pool_iter:
@@ -139,20 +146,21 @@ def iter_srcinfo(repo_paths, cache):
else:
to_parse.append(pkgbuild_path)
pool_iter = pool.imap_unordered(get_srcinfo_for_pkgbuild, to_parse)
print("Parsing PKGBUILD files...")
for srcinfo in pool_iter:
for srcinfo in pool.imap_unordered(get_srcinfo_for_pkgbuild, to_parse):
yield srcinfo
pool.close()
def main(argv):
def main(argv: List[str]) -> Optional[Union[int, str]]:
srcinfo_path = os.path.abspath(argv[1])
cache: Cache = {}
try:
with open(srcinfo_path, "rb") as h:
cache = json.loads(h.read())
except FileNotFoundError:
cache = {}
pass
t = time.monotonic()
srcinfos = []
@@ -165,9 +173,11 @@ def main(argv):
if time.monotonic() - t > 60 * 30:
break
srcinfos = OrderedDict(sorted(srcinfos))
srcinfos_dict = OrderedDict(sorted(srcinfos))
with open(srcinfo_path, "wb") as h:
h.write(json.dumps(srcinfos, indent=2).encode("utf-8"))
h.write(json.dumps(srcinfos_dict, indent=2).encode("utf-8"))
return None
if __name__ == "__main__":