aboutsummaryrefslogtreecommitdiff
path: root/server/szurubooru/rest
diff options
context:
space:
mode:
authorrr-2017-02-04 01:08:12 +0100
committerrr-2017-02-05 16:34:45 +0100
commitad842ee8a54c57463b8e28b52970f173ea1d64ea (patch)
tree1238e8270878801a2ddb10bc4db469115c026cab /server/szurubooru/rest
parentabf1fc2b2d135299fe7e8a700d4e7a18966d7bfe (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__.py2
-rw-r--r--server/szurubooru/rest/app.py19
-rw-r--r--server/szurubooru/rest/context.py187
-rw-r--r--server/szurubooru/rest/errors.py18
-rw-r--r--server/szurubooru/rest/middleware.py12
-rw-r--r--server/szurubooru/rest/routes.py22
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

© 2015 - 2026 Jakob L. Kreuze