Compare commits

..

20 Commits

Author SHA1 Message Date
Cristi Vîjdea 904895ba3c Add 1.9.2 changelog 2018-08-03 23:43:52 +03:00
Cristi Vîjdea 79adfc19be Update swagger-ui to 3.17.6 and ReDoc to 2.0.0-alpha.32 2018-08-03 20:27:41 +03:00
Cristi Vîjdea 2965e08e39 Force descriptions to be str objects
Fixes #159
2018-08-03 20:27:31 +03:00
Cristi Vîjdea f2e05ee4c0 Fix ReDoc configuration docs link
Closes #171
2018-08-03 16:50:28 +03:00
Amoki bbc70a7e3d Allow specific version generation in command
* Add --api-version parameter
* Fix request mocking
* Add tests
2018-08-03 16:43:26 +03:00
Bang Dao ca43a7de0c Fix readme wrong caching time expected behavior (#172) 2018-08-03 16:38:18 +03:00
Антон Вахмин (Anton Vakhmin) cc43bdf5cb Fix SHOW_COMMON_EXTENSIONS docs typos (#162)
Fix copy-paste errors.
2018-08-03 16:09:35 +03:00
Étienne Noss db86981dc1 make generate_swagger work for projects without authentication (#161)
* make generate_swagger work for projects without authentication
* use get_user_model instead of importing User
2018-07-06 16:13:19 +03:00
Paul Wayper b37ce3227a Fixing ` single-quote around @list-route (#160) 2018-07-06 15:26:34 +03:00
Cristi Vîjdea 8fbab88bda Fix changelog 2018-06-30 01:55:19 +03:00
Cristi Vîjdea 58b979f4ea Fix 1.9.1 tag fuckup
Accidentally disappeared 178390a9a0
2018-06-30 01:42:02 +03:00
Cristi Vîjdea d9b6fbc1b8 Add 1.9.1 changelog 2018-06-30 00:32:07 +03:00
Cristi Vîjdea 20370d912e Update swagger-ui to 3.17.1 and ReDoc to 2.0.0-alpha.25 2018-06-30 00:32:07 +03:00
Xiaohan Zhang 178390a9a0 Add get_default_response_serializer extension point (#153)
Enables easier request/response differentiation in SwaggerAutoSchema.
2018-06-30 00:26:29 +03:00
Cristi Vîjdea 7982f62f38 Add LICENSE file to wheel 2018-06-30 00:18:19 +03:00
Cristi Vîjdea 9fa4765121 Ignore None when passed in as response override
Closes #148
2018-06-29 23:13:36 +03:00
Jack Cushman 544d72db0a Handle duplicate urls in urlconf (#155)
Django resolves urls in order from top to bottom, and only uses the first matching URL found.
2018-06-29 22:57:37 +03:00
Cristi Vîjdea 696ec3a94a Add swagger_fake_view marker to help detect fake views in get_serializer
Cleaner fix for #154
2018-06-29 18:00:54 +03:00
Cristi Vîjdea e0aec3ff45 Test that get_serializer is not called when overriden
Views' ``get_serializer`` methods should not be called when the serializer is provided by ``request_body`` or ``responses``.

Closes #154
2018-06-29 17:41:41 +03:00
John Carter ee7b9a0734 Fix 1.9.0 changelog typo (#147) 2018-06-17 13:31:21 +03:00
27 changed files with 529 additions and 826 deletions
+3 -3
View File
@@ -134,9 +134,9 @@ In ``urls.py``:
) )
urlpatterns = [ urlpatterns = [
url(r'^swagger(?P<format>\.json|\.yaml)$', schema_view.without_ui(cache_timeout=None), name='schema-json'), url(r'^swagger(?P<format>\.json|\.yaml)$', schema_view.without_ui(cache_timeout=0), name='schema-json'),
url(r'^swagger/$', schema_view.with_ui('swagger', cache_timeout=None), name='schema-swagger-ui'), url(r'^swagger/$', schema_view.with_ui('swagger', cache_timeout=0), name='schema-swagger-ui'),
url(r'^redoc/$', schema_view.with_ui('redoc', cache_timeout=None), name='schema-redoc'), url(r'^redoc/$', schema_view.with_ui('redoc', cache_timeout=0), name='schema-redoc'),
... ...
] ]
+28 -2
View File
@@ -3,19 +3,45 @@ Changelog
######### #########
**********
**1.9.2**
**********
*Release date: Aug 03, 2018*
- **IMPROVED:** updated ``swagger-ui`` to version 3.17.6
- **IMPROVED:** updated ``ReDoc`` to version 2.0.0-alpha.32
- **IMPROVED:** added ``--api-version`` argument to the ``generate_swagger`` management command (:pr:`170`)
- **FIXED:** corrected various documentation typos (:pr:`160`, :pr:`162`, :issue:`171`, :pr:`172`)
- **FIXED:** made ``generate_swagger`` work for projects without authentication (:pr:`161`)
- **FIXED:** fixed ``SafeText`` interaction with YAML codec (:issue:`159`)
*********
**1.9.1**
*********
*Release date: Jun 30, 2018*
- **IMPROVED:** added a ``swagger_fake_view`` marker to more easily detect mock views in view methods;
``getattr(self, 'swagger_fake_view', False)`` inside a view method like ``get_serializer_class`` will tell you if the
view instnace is being used for swagger schema introspection (:issue:`154`)
- **IMPROVED:** updated ``swagger-ui`` to version 3.17.1
- **IMPROVED:** updated ``ReDoc`` to version 2.0.0-alpha.25
- **FIXED:** fixed wrong handling of duplicate urls in urlconf (:pr:`155`)
- **FIXED:** fixed crash when passing ``None`` as a response override (:issue:`148`)
********* *********
**1.9.0** **1.9.0**
********* *********
*Release date: Jun 16, 2018* *Release date: Jun 16, 2018*
- **ADDED:** added ``DEFAULT_GENERATOR_CLASS`` setting and ``--generator-clas`` argument to the ``generate_swagger`` - **ADDED:** added ``DEFAULT_GENERATOR_CLASS`` setting and ``--generator-class`` argument to the ``generate_swagger``
management command (:issue:`140`) management command (:issue:`140`)
- **FIXED:** fixed wrongly required ``'count'`` response field on ``CursorPagination`` (:issue:`141`) - **FIXED:** fixed wrongly required ``'count'`` response field on ``CursorPagination`` (:issue:`141`)
- **FIXED:** fixed some cases where ``swagger_extra_fields`` would not be handlded (:pr:`142`) - **FIXED:** fixed some cases where ``swagger_extra_fields`` would not be handlded (:pr:`142`)
- **FIXED:** fixed crash when encountering ``coreapi.Fields``\ s without a ``schema`` (:issue:`143`) - **FIXED:** fixed crash when encountering ``coreapi.Fields``\ s without a ``schema`` (:issue:`143`)
********* *********
**1.8.0** **1.8.0**
********* *********
+1 -1
View File
@@ -87,7 +87,7 @@ Where you can use the :func:`@swagger_auto_schema <.swagger_auto_schema>` decora
* for ``ViewSet``, ``GenericViewSet``, ``ModelViewSet``, because each viewset corresponds to multiple **paths**, you have * for ``ViewSet``, ``GenericViewSet``, ``ModelViewSet``, because each viewset corresponds to multiple **paths**, you have
to decorate the *action methods*, i.e. ``list``, ``create``, ``retrieve``, etc. |br| to decorate the *action methods*, i.e. ``list``, ``create``, ``retrieve``, etc. |br|
Additionally, ``@action``\ s, `@list_route``\ s or ``@detail_route``\ s defined on the viewset, like function based Additionally, ``@action``\ s, ``@list_route``\ s or ``@detail_route``\ s defined on the viewset, like function based
api views, can respond to multiple HTTP methods and thus have multiple operations that must be decorated separately: api views, can respond to multiple HTTP methods and thus have multiple operations that must be decorated separately:
+9
View File
@@ -174,6 +174,15 @@ in a few ways:
:ref:`@swagger_auto_schema <custom-spec-swagger-auto-schema>`, to prevent ``get_serializer`` from being called on :ref:`@swagger_auto_schema <custom-spec-swagger-auto-schema>`, to prevent ``get_serializer`` from being called on
the view the view
* :ref:`exclude your endpoint from introspection <custom-spec-excluding-endpoints>` * :ref:`exclude your endpoint from introspection <custom-spec-excluding-endpoints>`
* use the ``swagger_fake_view`` marker to detect requests generated by ``drf-yasg``:
.. code-block:: python
def get_serializer_class(self):
if getattr(self, 'swagger_fake_view', False):
return TodoTreeSerializer
raise NotImplementedError("must not call this")
.. _SCRIPT_NAME: https://www.python.org/dev/peps/pep-0333/#environ-variables .. _SCRIPT_NAME: https://www.python.org/dev/peps/pep-0333/#environ-variables
.. _FORCE_SCRIPT_NAME: https://docs.djangoproject.com/en/2.0/ref/settings/#force-script-name .. _FORCE_SCRIPT_NAME: https://docs.djangoproject.com/en/2.0/ref/settings/#force-script-name
+4 -4
View File
@@ -263,10 +263,10 @@ Controls how many levels are expaned by default when showing nested models.
**Default**: :python:`3` |br| **Default**: :python:`3` |br|
*Maps to parameter*: ``defaultModelExpandDepth`` *Maps to parameter*: ``defaultModelExpandDepth``
DEFAULT_MODEL_DEPTH SHOW_COMMON_EXTENSIONS
------------------- ----------------------
Controls the display of extensions (``pattern``, ``maxLength``, ``minLength``, ``maximum``, ```minimum``) fields and Controls the display of extensions (``pattern``, ``maxLength``, ``minLength``, ``maximum``, ``minimum``) fields and
values for Parameters. values for Parameters.
**Default**: :python:`True` |br| **Default**: :python:`True` |br|
@@ -313,7 +313,7 @@ ReDoc UI settings
================= =================
ReDoc UI configuration settings. |br| ReDoc UI configuration settings. |br|
See https://github.com/Rebilly/ReDoc#redoc-tag-attributes. See https://github.com/Rebilly/ReDoc#configuration.
LAZY_RENDERING LAZY_RENDERING
-------------- --------------
+175 -662
View File
File diff suppressed because it is too large Load Diff
+2 -2
View File
@@ -1,8 +1,8 @@
{ {
"name": "drf-yasg", "name": "drf-yasg",
"dependencies": { "dependencies": {
"redoc": "^2.0.0-alpha.22", "redoc": "^2.0.0-alpha.32",
"swagger-ui-dist": "^3.16.0" "swagger-ui-dist": "^3.17.6"
}, },
"repository": { "repository": {
"type": "git", "type": "git",
+3
View File
@@ -1,2 +1,5 @@
[bdist_wheel] [bdist_wheel]
universal = 1 universal = 1
[metadata]
license_file = LICENSE.rst
+1
View File
@@ -5,6 +5,7 @@ import json
from collections import OrderedDict from collections import OrderedDict
from coreapi.compat import force_bytes from coreapi.compat import force_bytes
from django.utils.safestring import SafeData, SafeText
from ruamel import yaml from ruamel import yaml
from . import openapi from . import openapi
+17 -6
View File
@@ -5,7 +5,6 @@ from collections import OrderedDict, defaultdict
import uritemplate import uritemplate
from coreapi.compat import urlparse from coreapi.compat import urlparse
from django.utils.encoding import force_text
from rest_framework import versioning from rest_framework import versioning
from rest_framework.compat import URLPattern, URLResolver, get_original_route from rest_framework.compat import URLPattern, URLResolver, get_original_route
from rest_framework.schemas.generators import EndpointEnumerator as _EndpointEnumerator from rest_framework.schemas.generators import EndpointEnumerator as _EndpointEnumerator
@@ -18,7 +17,7 @@ from .app_settings import swagger_settings
from .errors import SwaggerGenerationError from .errors import SwaggerGenerationError
from .inspectors.field import get_basic_type_info, get_queryset_field from .inspectors.field import get_basic_type_info, get_queryset_field
from .openapi import ReferenceResolver from .openapi import ReferenceResolver
from .utils import get_consumes, get_produces from .utils import force_real_str, get_consumes, get_produces
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -72,7 +71,7 @@ class EndpointEnumerator(_EndpointEnumerator):
return path return path
def get_api_endpoints(self, patterns=None, prefix='', app_name=None, namespace=None): def get_api_endpoints(self, patterns=None, prefix='', app_name=None, namespace=None, ignored_endpoints=None):
""" """
Return a list of all available API endpoints by inspecting the URL conf. Return a list of all available API endpoints by inspecting the URL conf.
@@ -82,6 +81,8 @@ class EndpointEnumerator(_EndpointEnumerator):
patterns = self.patterns patterns = self.patterns
api_endpoints = [] api_endpoints = []
if ignored_endpoints is None:
ignored_endpoints = set()
for pattern in patterns: for pattern in patterns:
path_regex = prefix + get_original_route(pattern) path_regex = prefix + get_original_route(pattern)
@@ -92,6 +93,13 @@ class EndpointEnumerator(_EndpointEnumerator):
url_name = pattern.name url_name = pattern.name
if self.should_include_endpoint(path, callback, app_name or '', namespace or '', url_name): if self.should_include_endpoint(path, callback, app_name or '', namespace or '', url_name):
path = self.replace_version(path, callback) path = self.replace_version(path, callback)
# avoid adding endpoints that have already been seen,
# as Django resolves urls in top-down order
if path in ignored_endpoints:
continue
ignored_endpoints.add(path)
for method in self.get_allowed_methods(callback): for method in self.get_allowed_methods(callback):
endpoint = (path, method, callback) endpoint = (path, method, callback)
api_endpoints.append(endpoint) api_endpoints.append(endpoint)
@@ -103,7 +111,8 @@ class EndpointEnumerator(_EndpointEnumerator):
patterns=pattern.url_patterns, patterns=pattern.url_patterns,
prefix=path_regex, prefix=path_regex,
app_name="%s:%s" % (app_name, pattern.app_name) if app_name else pattern.app_name, app_name="%s:%s" % (app_name, pattern.app_name) if app_name else pattern.app_name,
namespace="%s:%s" % (namespace, pattern.namespace) if namespace else pattern.namespace namespace="%s:%s" % (namespace, pattern.namespace) if namespace else pattern.namespace,
ignored_endpoints=ignored_endpoints
) )
api_endpoints.extend(nested_endpoints) api_endpoints.extend(nested_endpoints)
else: else:
@@ -243,6 +252,8 @@ class OpenAPISchemaGenerator(object):
view_method = getattr(view, method, None) view_method = getattr(view, method, None)
if view_method is not None: # pragma: no cover if view_method is not None: # pragma: no cover
setattr(view_method.__func__, '_swagger_auto_schema', overrides) setattr(view_method.__func__, '_swagger_auto_schema', overrides)
setattr(view, 'swagger_fake_view', True)
return view return view
def get_endpoints(self, request): def get_endpoints(self, request):
@@ -423,7 +434,7 @@ class OpenAPISchemaGenerator(object):
attrs['pattern'] = getattr(view_cls, 'lookup_value_regex', attrs.get('pattern', None)) attrs['pattern'] = getattr(view_cls, 'lookup_value_regex', attrs.get('pattern', None))
if model_field and getattr(model_field, 'help_text', False): if model_field and getattr(model_field, 'help_text', False):
description = force_text(model_field.help_text) description = model_field.help_text
elif model_field and getattr(model_field, 'primary_key', False): elif model_field and getattr(model_field, 'primary_key', False):
description = get_pk_description(model, model_field) description = get_pk_description(model, model_field)
else: else:
@@ -431,7 +442,7 @@ class OpenAPISchemaGenerator(object):
field = openapi.Parameter( field = openapi.Parameter(
name=variable, name=variable,
description=description, description=force_real_str(description),
required=True, required=True,
in_=openapi.IN_PATH, in_=openapi.IN_PATH,
**attrs **attrs
+4 -5
View File
@@ -1,12 +1,11 @@
import inspect import inspect
import logging import logging
from django.utils.encoding import force_text
from rest_framework import serializers from rest_framework import serializers
from rest_framework.utils import encoders, json from rest_framework.utils import encoders, json
from .. import openapi from .. import openapi
from ..utils import decimal_as_float, is_list_view from ..utils import decimal_as_float, force_real_str, is_list_view
#: Sentinel value that inspectors must return to signal that they do not know how to handle an object #: Sentinel value that inspectors must return to signal that they do not know how to handle an object
NotHandled = object() NotHandled = object()
@@ -196,9 +195,9 @@ class FieldInspector(BaseInspector):
""" """
assert swagger_object_type in (openapi.Schema, openapi.Parameter, openapi.Items) assert swagger_object_type in (openapi.Schema, openapi.Parameter, openapi.Items)
assert not isinstance(field, openapi.SwaggerDict), "passed field is already a SwaggerDict object" assert not isinstance(field, openapi.SwaggerDict), "passed field is already a SwaggerDict object"
title = force_text(field.label) if field.label else None title = force_real_str(field.label) if field.label else None
title = title if swagger_object_type == openapi.Schema else None # only Schema has title title = title if swagger_object_type == openapi.Schema else None # only Schema has title
description = force_text(field.help_text) if field.help_text else None description = force_real_str(field.help_text) if field.help_text else None
description = description if swagger_object_type != openapi.Items else None # Items has no description either description = description if swagger_object_type != openapi.Items else None # Items has no description either
def SwaggerType(existing_object=None, **instance_kwargs): def SwaggerType(existing_object=None, **instance_kwargs):
@@ -371,7 +370,7 @@ class ViewInspector(BaseInspector):
:param serializers.BaseSerializer serializer: the ``Serializer`` instance :param serializers.BaseSerializer serializer: the ``Serializer`` instance
:returns: the converted :class:`.Schema`, or ``None`` in case of an unknown serializer :returns: the converted :class:`.Schema`, or ``None`` in case of an unknown serializer
:rtype: openapi.Schema,openapi.SchemaRef,None :rtype: openapi.Schema,openapi.SchemaRef
""" """
return self.probe_inspectors( return self.probe_inspectors(
self.field_inspectors, 'get_schema', serializer, {'field_inspectors': self.field_inspectors} self.field_inspectors, 'get_schema', serializer, {'field_inspectors': self.field_inspectors}
+3 -1
View File
@@ -3,6 +3,8 @@ from collections import OrderedDict
import coreschema import coreschema
from rest_framework.pagination import CursorPagination, LimitOffsetPagination, PageNumberPagination from rest_framework.pagination import CursorPagination, LimitOffsetPagination, PageNumberPagination
from drf_yasg.utils import force_real_str
from .. import openapi from .. import openapi
from .base import FilterInspector, PaginatorInspector from .base import FilterInspector, PaginatorInspector
@@ -48,7 +50,7 @@ class CoreAPICompatInspector(PaginatorInspector, FilterInspector):
in_=location_to_in[field.location], in_=location_to_in[field.location],
type=coreapi_types.get(type(field.schema), openapi.TYPE_STRING), type=coreapi_types.get(type(field.schema), openapi.TYPE_STRING),
required=field.required, required=field.required,
description=field.schema.description if field.schema else None, description=force_real_str(field.schema.description) if field.schema else None,
) )
+31 -10
View File
@@ -8,7 +8,7 @@ from rest_framework.status import is_success
from .. import openapi from .. import openapi
from ..errors import SwaggerGenerationError from ..errors import SwaggerGenerationError
from ..utils import ( from ..utils import (
force_serializer_instance, get_consumes, get_produces, guess_response_status, is_list_view, no_body, force_real_str, force_serializer_instance, get_consumes, get_produces, guess_response_status, is_list_view, no_body,
param_list_to_odict param_list_to_odict
) )
from .base import ViewInspector from .base import ViewInspector
@@ -42,7 +42,7 @@ class SwaggerAutoSchema(ViewInspector):
return openapi.Operation( return openapi.Operation(
operation_id=operation_id, operation_id=operation_id,
description=description, description=force_real_str(description),
responses=responses, responses=responses,
parameters=parameters, parameters=parameters,
consumes=consumes, consumes=consumes,
@@ -89,13 +89,14 @@ class SwaggerAutoSchema(ViewInspector):
try: try:
return self.view.get_serializer() return self.view.get_serializer()
except Exception: except Exception:
log.warning("view's get_serializer raised exception (%s)", type(self.view).__name__, exc_info=True) log.warning("view's get_serializer raised exception (%s %s %s)",
self.method, self.path, type(self.view).__name__, exc_info=True)
return None return None
def get_request_serializer(self): def _get_request_body_override(self):
"""Return the request serializer (used for parsing the request payload) for this endpoint. """Parse the request_body key in the override dict. This method is not public API.
:return: the request serializer, or one of :class:`.Schema`, :class:`.SchemaRef`, ``None`` :return:
""" """
body_override = self.overrides.get('request_body', None) body_override = self.overrides.get('request_body', None)
@@ -108,10 +109,20 @@ class SwaggerAutoSchema(ViewInspector):
if isinstance(body_override, openapi.Schema.OR_REF): if isinstance(body_override, openapi.Schema.OR_REF):
return body_override return body_override
return force_serializer_instance(body_override) return force_serializer_instance(body_override)
elif self.method in self.implicit_body_methods:
return body_override
def get_request_serializer(self):
"""Return the request serializer (used for parsing the request payload) for this endpoint.
:return: the request serializer, or one of :class:`.Schema`, :class:`.SchemaRef`, ``None``
"""
body_override = self._get_request_body_override()
if body_override is None and self.method in self.implicit_body_methods:
return self.get_view_serializer() return self.get_view_serializer()
return None return body_override
def get_request_form_parameters(self, serializer): def get_request_form_parameters(self, serializer):
"""Given a Serializer, return a list of ``in: formData`` :class:`.Parameter`\ s. """Given a Serializer, return a list of ``in: formData`` :class:`.Parameter`\ s.
@@ -171,6 +182,14 @@ class SwaggerAutoSchema(ViewInspector):
responses=self.get_response_schemas(response_serializers) responses=self.get_response_schemas(response_serializers)
) )
def get_default_response_serializer(self):
"""Return the default response serializer for this endpoint. This is derived from either the ``request_body``
override or the request serializer (:meth:`.get_view_serializer`).
:return: response serializer, :class:`.Schema`, :class:`.SchemaRef`, ``None``
"""
return self._get_request_body_override() or self.get_view_serializer()
def get_default_responses(self): def get_default_responses(self):
"""Get the default responses determined for this view from the request serializer and request method. """Get the default responses determined for this view from the request serializer and request method.
@@ -181,7 +200,7 @@ class SwaggerAutoSchema(ViewInspector):
default_status = guess_response_status(method) default_status = guess_response_status(method)
default_schema = '' default_schema = ''
if method in ('get', 'post', 'put', 'patch'): if method in ('get', 'post', 'put', 'patch'):
default_schema = self.get_request_serializer() or self.get_view_serializer() default_schema = self.get_default_response_serializer()
default_schema = default_schema or '' default_schema = default_schema or ''
if any(is_form_media_type(encoding) for encoding in self.get_consumes()): if any(is_form_media_type(encoding) for encoding in self.get_consumes()):
@@ -227,8 +246,10 @@ class SwaggerAutoSchema(ViewInspector):
for sc, serializer in response_serializers.items(): for sc, serializer in response_serializers.items():
if isinstance(serializer, str): if isinstance(serializer, str):
response = openapi.Response( response = openapi.Response(
description=serializer description=force_real_str(serializer)
) )
elif not serializer:
continue
elif isinstance(serializer, openapi.Response): elif isinstance(serializer, openapi.Response):
response = serializer response = serializer
if hasattr(response, 'schema') and not isinstance(response.schema, openapi.Schema.OR_REF): if hasattr(response, 'schema') and not isinstance(response.schema, openapi.Schema.OR_REF):
@@ -4,7 +4,7 @@ import os
from collections import OrderedDict from collections import OrderedDict
from importlib import import_module from importlib import import_module
from django.contrib.auth.models import User from django.contrib.auth import get_user_model
from django.core.exceptions import ImproperlyConfigured from django.core.exceptions import ImproperlyConfigured
from django.core.management.base import BaseCommand from django.core.management.base import BaseCommand
from rest_framework.test import APIRequestFactory, force_authenticate from rest_framework.test import APIRequestFactory, force_authenticate
@@ -59,9 +59,13 @@ class Command(BaseCommand):
help='Use a mock request when generating the swagger schema. This is useful if your views or serializers' help='Use a mock request when generating the swagger schema. This is useful if your views or serializers'
'depend on context from a request in order to function.' 'depend on context from a request in order to function.'
) )
parser.add_argument(
'--api-version', dest='api_version',
type=str,
help='Version to use to generate schema. This option implies --mock-request.'
)
parser.add_argument( parser.add_argument(
'--user', dest='user', '--user', dest='user',
default='',
help='Username of an existing user to use for mocked authentication. This option implies --mock-request.' help='Username of an existing user to use for mocked authentication. This option implies --mock-request.'
) )
parser.add_argument( parser.add_argument(
@@ -103,7 +107,7 @@ class Command(BaseCommand):
request = APIView().initialize_request(request) request = APIView().initialize_request(request)
return request return request
def handle(self, output_file, overwrite, format, api_url, mock, user, private, generator_class_name, def handle(self, output_file, overwrite, format, api_url, mock, api_version, user, private, generator_class_name,
*args, **kwargs): *args, **kwargs):
# disable logs of WARNING and below # disable logs of WARNING and below
logging.disable(logging.WARNING) logging.disable(logging.WARNING)
@@ -122,20 +126,30 @@ class Command(BaseCommand):
api_url = api_url or swagger_settings.DEFAULT_API_URL api_url = api_url or swagger_settings.DEFAULT_API_URL
user = User.objects.get(username=user) if user else None if user:
mock = mock or private or (user is not None) # Only call get_user_model if --user was passed in order to
# avoid crashing if auth is not configured in the project
user = get_user_model().objects.get(username=user)
mock = mock or private or (user is not None) or (api_version is not None)
if mock and not api_url: if mock and not api_url:
raise ImproperlyConfigured( raise ImproperlyConfigured(
'--mock-request requires an API url; either provide ' '--mock-request requires an API url; either provide '
'the --url argument or set the DEFAULT_API_URL setting' 'the --url argument or set the DEFAULT_API_URL setting'
) )
request = self.get_mock_request(api_url, format, user) if mock else None request = None
if mock:
request = self.get_mock_request(api_url, format, user)
if request and api_version:
request.version = api_version
generator_class = import_class(generator_class_name) or swagger_settings.DEFAULT_GENERATOR_CLASS generator_class = import_class(generator_class_name) or swagger_settings.DEFAULT_GENERATOR_CLASS
generator = generator_class( generator = generator_class(
info=info, info=info,
url=api_url version=api_version,
url=api_url,
) )
schema = generator.get_schema(request=request, public=not private) schema = generator.get_schema(request=request, public=not private)
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+19 -2
View File
@@ -3,6 +3,7 @@ import logging
from collections import OrderedDict from collections import OrderedDict
from django.db import models from django.db import models
from django.utils.encoding import force_text
from rest_framework import serializers, status from rest_framework import serializers, status
from rest_framework.mixins import DestroyModelMixin, RetrieveModelMixin, UpdateModelMixin from rest_framework.mixins import DestroyModelMixin, RetrieveModelMixin, UpdateModelMixin
from rest_framework.request import is_form_media_type from rest_framework.request import is_form_media_type
@@ -76,6 +77,8 @@ def swagger_auto_schema(method=None, methods=None, auto_schema=unset, request_bo
* if a plain string is given as value, a :class:`.Response` with no body and that string as its description * if a plain string is given as value, a :class:`.Response` with no body and that string as its description
will be generated will be generated
* if ``None`` is given as a value, the response is ignored; this is mainly useful for disabling default
2xx responses, i.e. ``responses={200: None, 302: 'something'}``
* if a :class:`.Schema`, :class:`.SchemaRef` is given, a :class:`.Response` with the schema as its body and * if a :class:`.Schema`, :class:`.SchemaRef` is given, a :class:`.Response` with the schema as its body and
an empty description will be generated an empty description will be generated
* a ``Serializer`` class or instance will be converted into a :class:`.Schema` and treated as above * a ``Serializer`` class or instance will be converted into a :class:`.Schema` and treated as above
@@ -302,8 +305,8 @@ def get_serializer_ref_name(serializer):
Get serializer's ref_name (or None for ModelSerializer if it is named 'NestedSerializer') Get serializer's ref_name (or None for ModelSerializer if it is named 'NestedSerializer')
:param serializer: Serializer instance :param serializer: Serializer instance
:return: Serializer's ref_name or None for inline serializer :return: Serializer's ``ref_name`` or ``None`` for inline serializer
:rtype: str or None :rtype: str
""" """
serializer_meta = getattr(serializer, 'Meta', None) serializer_meta = getattr(serializer, 'Meta', None)
serializer_name = type(serializer).__name__ serializer_name = type(serializer).__name__
@@ -317,3 +320,17 @@ def get_serializer_ref_name(serializer):
if ref_name.endswith('Serializer'): if ref_name.endswith('Serializer'):
ref_name = ref_name[:-len('Serializer')] ref_name = ref_name[:-len('Serializer')]
return ref_name return ref_name
def force_real_str(s, encoding='utf-8', strings_only=False, errors='strict'):
"""
Force `s` into a ``str`` instance.
Fix for https://github.com/axnsan12/drf-yasg/issues/159
"""
if s is not None:
s = force_text(s, encoding, strings_only, errors)
if type(s) != str:
s = '' + s
return s
+2 -1
View File
@@ -1,9 +1,10 @@
from django.db import models from django.db import models
from django.utils.safestring import mark_safe
class Identity(models.Model): class Identity(models.Model):
firstName = models.CharField(max_length=30, null=True) firstName = models.CharField(max_length=30, null=True)
lastName = models.CharField(max_length=30, null=True) lastName = models.CharField(max_length=30, null=True, help_text=mark_safe("<strong>Here's some HTML!</strong>"))
class Person(models.Model): class Person(models.Model):
+40 -2
View File
@@ -1,6 +1,8 @@
from rest_framework import viewsets from rest_framework import viewsets
from rest_framework.generics import RetrieveAPIView from rest_framework.generics import RetrieveAPIView
from drf_yasg.utils import swagger_auto_schema
from .models import Todo, TodoAnother, TodoTree, TodoYetAnother from .models import Todo, TodoAnother, TodoTree, TodoYetAnother
from .serializer import ( from .serializer import (
TodoAnotherSerializer, TodoRecursiveSerializer, TodoSerializer, TodoTreeSerializer, TodoYetAnotherSerializer TodoAnotherSerializer, TodoRecursiveSerializer, TodoSerializer, TodoTreeSerializer, TodoYetAnotherSerializer
@@ -31,9 +33,45 @@ class NestedTodoView(RetrieveAPIView):
class TodoTreeView(viewsets.ReadOnlyModelViewSet): class TodoTreeView(viewsets.ReadOnlyModelViewSet):
queryset = TodoTree.objects.all() queryset = TodoTree.objects.all()
serializer_class = TodoTreeSerializer
def get_serializer_class(self):
if getattr(self, 'swagger_fake_view', False):
return TodoTreeSerializer
raise NotImplementedError("must not call this")
class TodoRecursiveView(viewsets.ModelViewSet): class TodoRecursiveView(viewsets.ModelViewSet):
queryset = TodoTree.objects.all() queryset = TodoTree.objects.all()
serializer_class = TodoRecursiveSerializer
def get_serializer(self, *args, **kwargs):
raise NotImplementedError("must not call this")
def get_serializer_class(self):
raise NotImplementedError("must not call this")
def get_serializer_context(self):
raise NotImplementedError("must not call this")
@swagger_auto_schema(request_body=TodoRecursiveSerializer)
def create(self, request, *args, **kwargs):
return super(TodoRecursiveView, self).create(request, *args, **kwargs)
@swagger_auto_schema(responses={200: None, 302: 'Redirect somewhere'})
def retrieve(self, request, *args, **kwargs):
return super(TodoRecursiveView, self).retrieve(request, *args, **kwargs)
@swagger_auto_schema(request_body=TodoRecursiveSerializer)
def update(self, request, *args, **kwargs):
return super(TodoRecursiveView, self).update(request, *args, **kwargs)
@swagger_auto_schema(request_body=TodoRecursiveSerializer)
def partial_update(self, request, *args, **kwargs):
return super(TodoRecursiveView, self).update(request, *args, **kwargs)
def destroy(self, request, *args, **kwargs):
return super(TodoRecursiveView, self).destroy(request, *args, **kwargs)
@swagger_auto_schema(responses={200: TodoRecursiveSerializer(many=True)})
def list(self, request, *args, **kwargs):
return super(TodoRecursiveView, self).list(request, *args, **kwargs)
+18
View File
@@ -1,3 +1,5 @@
from six import StringIO
import copy import copy
import json import json
import os import os
@@ -6,6 +8,7 @@ from collections import OrderedDict
import pytest import pytest
from datadiff.tools import assert_equal from datadiff.tools import assert_equal
from django.contrib.auth.models import User from django.contrib.auth.models import User
from django.core.management import call_command
from rest_framework.test import APIRequestFactory from rest_framework.test import APIRequestFactory
from rest_framework.views import APIView from rest_framework.views import APIView
@@ -64,6 +67,21 @@ def validate_schema(db):
return validate_schema return validate_schema
@pytest.fixture
def call_generate_swagger():
def call_generate_swagger(output_file='-', overwrite=False, format='', api_url='',
mock=False, user=None, private=False, generator_class_name='', **kwargs):
out = StringIO()
call_command(
'generate_swagger', stdout=out,
output_file=output_file, overwrite=overwrite, format=format, api_url=api_url, mock=mock, user=user,
private=private, generator_class_name=generator_class_name, **kwargs
)
return out.getvalue()
return call_generate_swagger
@pytest.fixture @pytest.fixture
def compare_schemas(): def compare_schemas():
def compare_schemas(schema1, schema2): def compare_schemas(schema1, schema2):
+3 -4
View File
@@ -560,10 +560,8 @@ paths:
description: '' description: ''
parameters: [] parameters: []
responses: responses:
'200': '302':
description: '' description: Redirect somewhere
schema:
$ref: '#/definitions/TodoRecursive'
tags: tags:
- todo - todo
put: put:
@@ -907,6 +905,7 @@ definitions:
minLength: 1 minLength: 1
lastName: lastName:
title: LastName title: LastName
description: <strong>Here's some HTML!</strong>
type: string type: string
maxLength: 30 maxLength: 30
minLength: 1 minLength: 1
+5 -17
View File
@@ -9,25 +9,13 @@ from collections import OrderedDict
import pytest import pytest
from django.contrib.auth.models import User from django.contrib.auth.models import User
from django.core.management import call_command
from drf_yasg import openapi from drf_yasg import openapi
from drf_yasg.codecs import yaml_sane_load from drf_yasg.codecs import yaml_sane_load
from drf_yasg.generators import OpenAPISchemaGenerator from drf_yasg.generators import OpenAPISchemaGenerator
def call_generate_swagger(output_file='-', overwrite=False, format='', api_url='', def test_reference_schema(call_generate_swagger, db, reference_schema):
mock=False, user='', private=False, generator_class_name='', **kwargs):
out = StringIO()
call_command(
'generate_swagger', stdout=out,
output_file=output_file, overwrite=overwrite, format=format, api_url=api_url, mock=mock, user=user,
private=private, generator_class_name=generator_class_name, **kwargs
)
return out.getvalue()
def test_reference_schema(db, reference_schema):
User.objects.create_superuser('admin', 'admin@admin.admin', 'blabla') User.objects.create_superuser('admin', 'admin@admin.admin', 'blabla')
output = call_generate_swagger(format='yaml', api_url='http://test.local:8002/', user='admin') output = call_generate_swagger(format='yaml', api_url='http://test.local:8002/', user='admin')
@@ -35,13 +23,13 @@ def test_reference_schema(db, reference_schema):
assert output_schema == reference_schema assert output_schema == reference_schema
def test_non_public(db): def test_non_public(call_generate_swagger, db):
output = call_generate_swagger(format='yaml', api_url='http://test.local:8002/', private=True) output = call_generate_swagger(format='yaml', api_url='http://test.local:8002/', private=True)
output_schema = yaml_sane_load(output) output_schema = yaml_sane_load(output)
assert len(output_schema['paths']) == 0 assert len(output_schema['paths']) == 0
def test_no_mock(db): def test_no_mock(call_generate_swagger, db):
output = call_generate_swagger() output = call_generate_swagger()
output_schema = json.loads(output, object_pairs_hook=OrderedDict) output_schema = json.loads(output, object_pairs_hook=OrderedDict)
assert len(output_schema['paths']) > 0 assert len(output_schema['paths']) > 0
@@ -52,7 +40,7 @@ class EmptySchemaGenerator(OpenAPISchemaGenerator):
return openapi.Paths(paths={}), '' return openapi.Paths(paths={}), ''
def test_generator_class(db): def test_generator_class(call_generate_swagger, db):
output = call_generate_swagger(generator_class_name='test_management.EmptySchemaGenerator') output = call_generate_swagger(generator_class_name='test_management.EmptySchemaGenerator')
output_schema = json.loads(output, object_pairs_hook=OrderedDict) output_schema = json.loads(output, object_pairs_hook=OrderedDict)
assert len(output_schema['paths']) == 0 assert len(output_schema['paths']) == 0
@@ -65,7 +53,7 @@ def silentremove(filename):
pass pass
def test_file_output(db): def test_file_output(call_generate_swagger, db):
prefix = os.path.join(tempfile.gettempdir(), tempfile.gettempprefix()) prefix = os.path.join(tempfile.gettempdir(), tempfile.gettempprefix())
name = ''.join(random.choice(string.ascii_lowercase + string.digits) for _ in range(8)) name = ''.join(random.choice(string.ascii_lowercase + string.digits) for _ in range(8))
yaml_file = prefix + name + '.yaml' yaml_file = prefix + name + '.yaml'
+34
View File
@@ -2,7 +2,9 @@ import json
from collections import OrderedDict from collections import OrderedDict
import pytest import pytest
from django.conf.urls import url
from rest_framework import routers, serializers, viewsets from rest_framework import routers, serializers, viewsets
from rest_framework.decorators import api_view
from rest_framework.response import Response from rest_framework.response import Response
from drf_yasg import codecs, openapi from drf_yasg import codecs, openapi
@@ -113,3 +115,35 @@ def test_replaced_serializer():
responses = swagger['paths']['/details/{id}/']['get']['responses'] responses = swagger['paths']['/details/{id}/']['get']['responses']
assert '404' in responses assert '404' in responses
assert responses['404']['schema']['$ref'] == "#/definitions/Detail" assert responses['404']['schema']['$ref'] == "#/definitions/Detail"
def test_url_order():
# this view with description override should show up in the schema ...
@swagger_auto_schema(method='get', operation_description="description override")
@api_view()
def test_override(request, pk=None):
return Response({"message": "Hello, world!"})
# ... instead of this view which appears later in the url patterns
@api_view()
def test_view(request, pk=None):
return Response({"message": "Hello, world!"})
patterns = [
url(r'^/test/$', test_override),
url(r'^/test/$', test_view),
]
generator = OpenAPISchemaGenerator(
info=openapi.Info(title="Test generator", default_version="v1"),
version="v2",
url='',
patterns=patterns
)
# description override is successful
swagger = generator.get_schema(None, True)
assert swagger['paths']['/test/']['get']['description'] == 'description override'
# get_endpoints only includes one endpoint
assert len(generator.get_endpoints(None)['/test/'][1]) == 1
+26
View File
@@ -7,6 +7,18 @@ def _get_versioned_schema(prefix, client, validate_schema):
response = client.get(prefix + '/swagger.yaml') response = client.get(prefix + '/swagger.yaml')
assert response.status_code == 200 assert response.status_code == 200
swagger = yaml_sane_load(response.content.decode('utf-8')) swagger = yaml_sane_load(response.content.decode('utf-8'))
_check_base(swagger, prefix, validate_schema)
return swagger
def _get_versioned_schema_management(prefix, call_generate_swagger, validate_schema, kwargs):
output = call_generate_swagger(format='yaml', api_url='http://localhost' + prefix + '/swagger.yaml', **kwargs)
swagger = yaml_sane_load(output)
_check_base(swagger, prefix, validate_schema)
return swagger
def _check_base(swagger, prefix, validate_schema):
assert swagger['basePath'] == prefix assert swagger['basePath'] == prefix
validate_schema(swagger) validate_schema(swagger)
assert '/snippets/' in swagger['paths'] assert '/snippets/' in swagger['paths']
@@ -51,3 +63,17 @@ def test_ns_v1(client, validate_schema):
def test_ns_v2(client, validate_schema): def test_ns_v2(client, validate_schema):
swagger = _get_versioned_schema('/versioned/ns/v2.0', client, validate_schema) swagger = _get_versioned_schema('/versioned/ns/v2.0', client, validate_schema)
_check_v2(swagger) _check_v2(swagger)
@pytest.mark.urls('urlconfs.url_versioning')
def test_url_v2_management(call_generate_swagger, validate_schema):
kwargs = {'api_version': '2.0'}
swagger = _get_versioned_schema_management('/versioned/url/v2.0', call_generate_swagger, validate_schema, kwargs)
_check_v2(swagger)
@pytest.mark.urls('urlconfs.ns_versioning')
def test_ns_v2_management(call_generate_swagger, validate_schema):
kwargs = {'api_version': '2.0'}
swagger = _get_versioned_schema_management('/versioned/ns/v2.0', call_generate_swagger, validate_schema, kwargs)
_check_v2(swagger)