diff options
| author | rr- | 2017-02-04 01:08:12 +0100 |
|---|---|---|
| committer | rr- | 2017-02-05 16:34:45 +0100 |
| commit | ad842ee8a54c57463b8e28b52970f173ea1d64ea (patch) | |
| tree | 1238e8270878801a2ddb10bc4db469115c026cab /server/szurubooru/rest | |
| parent | abf1fc2b2d135299fe7e8a700d4e7a18966d7bfe (diff) | |
server: refactor + add type hinting
- Added type hinting (for now, 3.5-compatible)
- Split `db` namespace into `db` module and `model` namespace
- Changed elastic search to be created lazily for each operation
- Changed to class based approach in entity serialization to allow
stronger typing
- Removed `required` argument from `context.get_*` family of functions;
now it's implied if `default` argument is omitted
- Changed `unalias_dict` implementation to use less magic inputs
Diffstat (limited to 'server/szurubooru/rest')
| -rw-r--r-- | server/szurubooru/rest/__init__.py | 2 | ||||
| -rw-r--r-- | server/szurubooru/rest/app.py | 19 | ||||
| -rw-r--r-- | server/szurubooru/rest/context.py | 187 | ||||
| -rw-r--r-- | server/szurubooru/rest/errors.py | 18 | ||||
| -rw-r--r-- | server/szurubooru/rest/middleware.py | 12 | ||||
| -rw-r--r-- | server/szurubooru/rest/routes.py | 22 |
6 files changed, 164 insertions, 96 deletions
diff --git a/server/szurubooru/rest/__init__.py b/server/szurubooru/rest/__init__.py index ac9958a..14a3e30 100644 --- a/server/szurubooru/rest/__init__.py +++ b/server/szurubooru/rest/__init__.py @@ -1,2 +1,2 @@ from szurubooru.rest.app import application -from szurubooru.rest.context import Context +from szurubooru.rest.context import Context, Response diff --git a/server/szurubooru/rest/app.py b/server/szurubooru/rest/app.py index 1bbf8dc..b29110e 100644 --- a/server/szurubooru/rest/app.py +++ b/server/szurubooru/rest/app.py @@ -2,13 +2,14 @@ import urllib.parse import cgi import json import re +from typing import Dict, Any, Callable, Tuple from datetime import datetime from szurubooru import db from szurubooru.func import util from szurubooru.rest import errors, middleware, routes, context -def _json_serializer(obj): +def _json_serializer(obj: Any) -> str: ''' JSON serializer for objects not serializable by default JSON code ''' if isinstance(obj, datetime): serial = obj.isoformat('T') + 'Z' @@ -16,12 +17,12 @@ def _json_serializer(obj): raise TypeError('Type not serializable') -def _dump_json(obj): +def _dump_json(obj: Any) -> str: return json.dumps(obj, default=_json_serializer, indent=2) -def _get_headers(env): - headers = {} +def _get_headers(env: Dict[str, Any]) -> Dict[str, str]: + headers = {} # type: Dict[str, str] for key, value in env.items(): if key.startswith('HTTP_'): key = util.snake_case_to_upper_train_case(key[5:]) @@ -29,7 +30,7 @@ def _get_headers(env): return headers -def _create_context(env): +def _create_context(env: Dict[str, Any]) -> context.Context: method = env['REQUEST_METHOD'] path = '/' + env['PATH_INFO'].lstrip('/') headers = _get_headers(env) @@ -64,7 +65,9 @@ def _create_context(env): return context.Context(method, path, headers, params, files) -def application(env, start_response): +def application( + env: Dict[str, Any], + start_response: Callable[[str, Any], Any]) -> Tuple[bytes]: try: ctx = _create_context(env) if 'application/json' not in ctx.get_header('Accept'): @@ -106,9 +109,9 @@ def application(env, start_response): return (_dump_json(response).encode('utf-8'),) except Exception as ex: - for exception_type, handler in errors.error_handlers.items(): + for exception_type, ex_handler in errors.error_handlers.items(): if isinstance(ex, exception_type): - handler(ex) + ex_handler(ex) raise except errors.BaseHttpError as ex: diff --git a/server/szurubooru/rest/context.py b/server/szurubooru/rest/context.py index ae26f38..bb33bfa 100644 --- a/server/szurubooru/rest/context.py +++ b/server/szurubooru/rest/context.py @@ -1,111 +1,158 @@ -from szurubooru import errors +from typing import Any, Union, List, Dict, Optional, cast +from szurubooru import model, errors from szurubooru.func import net, file_uploads -def _lower_first(source): - return source[0].lower() + source[1:] - - -def _param_wrapper(func): - def wrapper(self, name, required=False, default=None, **kwargs): - # pylint: disable=protected-access - if name in self._params: - value = self._params[name] - try: - value = func(self, value, **kwargs) - except errors.InvalidParameterError as ex: - raise errors.InvalidParameterError( - 'Parameter %r is invalid: %s' % ( - name, _lower_first(str(ex)))) - return value - if not required: - return default - raise errors.MissingRequiredParameterError( - 'Required parameter %r is missing.' % name) - return wrapper +MISSING = object() +Request = Dict[str, Any] +Response = Optional[Dict[str, Any]] class Context: - def __init__(self, method, url, headers=None, params=None, files=None): + def __init__( + self, + method: str, + url: str, + headers: Dict[str, str]=None, + params: Request=None, + files: Dict[str, bytes]=None) -> None: self.method = method self.url = url self._headers = headers or {} self._params = params or {} self._files = files or {} - # provided by middleware - # self.session = None - # self.user = None + self.user = model.User() + self.user.name = None + self.user.rank = 'anonymous' + + self.session = None # type: Any - def has_header(self, name): + def has_header(self, name: str) -> bool: return name in self._headers - def get_header(self, name): - return self._headers.get(name, None) + def get_header(self, name: str) -> str: + return self._headers.get(name, '') - def has_file(self, name, allow_tokens=True): + def has_file(self, name: str, allow_tokens: bool=True) -> bool: return ( name in self._files or name + 'Url' in self._params or (allow_tokens and name + 'Token' in self._params)) - def get_file(self, name, required=False, allow_tokens=True): - ret = None - if name in self._files: - ret = self._files[name] - elif name + 'Url' in self._params: - ret = net.download(self._params[name + 'Url']) - elif allow_tokens and name + 'Token' in self._params: + def get_file( + self, + name: str, + default: Union[object, bytes]=MISSING, + allow_tokens: bool=True) -> bytes: + if name in self._files and self._files[name]: + return self._files[name] + + if name + 'Url' in self._params: + return net.download(self._params[name + 'Url']) + + if allow_tokens and name + 'Token' in self._params: ret = file_uploads.get(self._params[name + 'Token']) - if required and not ret: + if ret: + return ret + elif default is not MISSING: raise errors.MissingOrExpiredRequiredFileError( 'Required file %r is missing or has expired.' % name) - if required and not ret: - raise errors.MissingRequiredFileError( - 'Required file %r is missing.' % name) - return ret - def has_param(self, name): + if default is not MISSING: + return cast(bytes, default) + raise errors.MissingRequiredFileError( + 'Required file %r is missing.' % name) + + def has_param(self, name: str) -> bool: return name in self._params - @_param_wrapper - def get_param_as_list(self, value): - if not isinstance(value, list): + def get_param_as_list( + self, + name: str, + default: Union[object, List[Any]]=MISSING) -> List[Any]: + if name not in self._params: + if default is not MISSING: + return cast(List[Any], default) + raise errors.MissingRequiredParameterError( + 'Required parameter %r is missing.' % name) + value = self._params[name] + if type(value) is str: if ',' in value: return value.split(',') return [value] - return value + if type(value) is list: + return value + raise errors.InvalidParameterError( + 'Parameter %r must be a list.' % name) - @_param_wrapper - def get_param_as_string(self, value): - if isinstance(value, list): - try: - value = ','.join(value) - except TypeError: - raise errors.InvalidParameterError('Expected simple string.') - return value + def get_param_as_string( + self, + name: str, + default: Union[object, str]=MISSING) -> str: + if name not in self._params: + if default is not MISSING: + return cast(str, default) + raise errors.MissingRequiredParameterError( + 'Required parameter %r is missing.' % name) + value = self._params[name] + try: + if value is None: + return '' + if type(value) is list: + return ','.join(value) + if type(value) is int or type(value) is float: + return str(value) + if type(value) is str: + return value + except TypeError: + pass + raise errors.InvalidParameterError( + 'Parameter %r must be a string value.' % name) - @_param_wrapper - def get_param_as_int(self, value, min=None, max=None): + def get_param_as_int( + self, + name: str, + default: Union[object, int]=MISSING, + min: Optional[int]=None, + max: Optional[int]=None) -> int: + if name not in self._params: + if default is not MISSING: + return cast(int, default) + raise errors.MissingRequiredParameterError( + 'Required parameter %r is missing.' % name) + value = self._params[name] try: value = int(value) + if min is not None and value < min: + raise errors.InvalidParameterError( + 'Parameter %r must be at least %r.' % (name, min)) + if max is not None and value > max: + raise errors.InvalidParameterError( + 'Parameter %r may not exceed %r.' % (name, max)) + return value except (ValueError, TypeError): - raise errors.InvalidParameterError( - 'The value must be an integer.') - if min is not None and value < min: - raise errors.InvalidParameterError( - 'The value must be at least %r.' % min) - if max is not None and value > max: - raise errors.InvalidParameterError( - 'The value may not exceed %r.' % max) - return value + pass + raise errors.InvalidParameterError( + 'Parameter %r must be an integer value.' % name) - @_param_wrapper - def get_param_as_bool(self, value): - value = str(value).lower() + def get_param_as_bool( + self, + name: str, + default: Union[object, bool]=MISSING) -> bool: + if name not in self._params: + if default is not MISSING: + return cast(bool, default) + raise errors.MissingRequiredParameterError( + 'Required parameter %r is missing.' % name) + value = self._params[name] + try: + value = str(value).lower() + except TypeError: + pass if value in ['1', 'y', 'yes', 'yeah', 'yep', 'yup', 't', 'true']: return True if value in ['0', 'n', 'no', 'nope', 'f', 'false']: return False raise errors.InvalidParameterError( - 'The value must be a boolean value.') + 'Parameter %r must be a boolean value.' % name) diff --git a/server/szurubooru/rest/errors.py b/server/szurubooru/rest/errors.py index b0f5b88..6854e7d 100644 --- a/server/szurubooru/rest/errors.py +++ b/server/szurubooru/rest/errors.py @@ -1,11 +1,19 @@ +from typing import Callable, Type, Dict + + error_handlers = {} # pylint: disable=invalid-name class BaseHttpError(RuntimeError): - code = None - reason = None + code = -1 + reason = '' - def __init__(self, name, description, title=None, extra_fields=None): + def __init__( + self, + name: str, + description: str, + title: str=None, + extra_fields: Dict[str, str]=None) -> None: super().__init__() # error name for programmers self.name = name @@ -52,5 +60,7 @@ class HttpInternalServerError(BaseHttpError): reason = 'Internal Server Error' -def handle(exception_type, handler): +def handle( + exception_type: Type[Exception], + handler: Callable[[Exception], None]) -> None: error_handlers[exception_type] = handler diff --git a/server/szurubooru/rest/middleware.py b/server/szurubooru/rest/middleware.py index 7cf0729..05d9495 100644 --- a/server/szurubooru/rest/middleware.py +++ b/server/szurubooru/rest/middleware.py @@ -1,11 +1,15 @@ +from typing import Callable +from szurubooru.rest.context import Context + + # pylint: disable=invalid-name -pre_hooks = [] -post_hooks = [] +pre_hooks = [] # type: List[Callable[[Context], None]] +post_hooks = [] # type: List[Callable[[Context], None]] -def pre_hook(handler): +def pre_hook(handler: Callable) -> None: pre_hooks.append(handler) -def post_hook(handler): +def post_hook(handler: Callable) -> None: post_hooks.insert(0, handler) diff --git a/server/szurubooru/rest/routes.py b/server/szurubooru/rest/routes.py index ffa95f5..c0b6bea 100644 --- a/server/szurubooru/rest/routes.py +++ b/server/szurubooru/rest/routes.py @@ -1,32 +1,36 @@ +from typing import Callable, Dict, Any from collections import defaultdict +from szurubooru.rest.context import Context, Response -routes = defaultdict(dict) # pylint: disable=invalid-name +# pylint: disable=invalid-name +RouteHandler = Callable[[Context, Dict[str, str]], Response] +routes = defaultdict(dict) # type: Dict[str, Dict[str, RouteHandler]] -def get(url): - def wrapper(handler): +def get(url: str) -> Callable[[RouteHandler], RouteHandler]: + def wrapper(handler: RouteHandler) -> RouteHandler: routes[url]['GET'] = handler return handler return wrapper -def put(url): - def wrapper(handler): +def put(url: str) -> Callable[[RouteHandler], RouteHandler]: + def wrapper(handler: RouteHandler) -> RouteHandler: routes[url]['PUT'] = handler return handler return wrapper -def post(url): - def wrapper(handler): +def post(url: str) -> Callable[[RouteHandler], RouteHandler]: + def wrapper(handler: RouteHandler) -> RouteHandler: routes[url]['POST'] = handler return handler return wrapper -def delete(url): - def wrapper(handler): +def delete(url: str) -> Callable[[RouteHandler], RouteHandler]: + def wrapper(handler: RouteHandler) -> RouteHandler: routes[url]['DELETE'] = handler return handler return wrapper |