17.12
This commit is contained in:
@@ -0,0 +1,172 @@
|
||||
from django.conf import settings
|
||||
from django.contrib.auth import (
|
||||
BACKEND_SESSION_KEY,
|
||||
HASH_SESSION_KEY,
|
||||
SESSION_KEY,
|
||||
_get_backends,
|
||||
get_user_model,
|
||||
load_backend,
|
||||
user_logged_in,
|
||||
user_logged_out,
|
||||
)
|
||||
from django.contrib.auth.models import AnonymousUser
|
||||
from django.utils.crypto import constant_time_compare
|
||||
from django.utils.functional import LazyObject
|
||||
from django.utils.translation import LANGUAGE_SESSION_KEY
|
||||
|
||||
from channels.db import database_sync_to_async
|
||||
from channels.middleware import BaseMiddleware
|
||||
from channels.sessions import CookieMiddleware, SessionMiddleware
|
||||
|
||||
|
||||
@database_sync_to_async
|
||||
def get_user(scope):
|
||||
"""
|
||||
Return the user model instance associated with the given scope.
|
||||
If no user is retrieved, return an instance of `AnonymousUser`.
|
||||
"""
|
||||
if "session" not in scope:
|
||||
raise ValueError(
|
||||
"Cannot find session in scope. You should wrap your consumer in SessionMiddleware."
|
||||
)
|
||||
session = scope["session"]
|
||||
user = None
|
||||
try:
|
||||
user_id = _get_user_session_key(session)
|
||||
backend_path = session[BACKEND_SESSION_KEY]
|
||||
except KeyError:
|
||||
pass
|
||||
else:
|
||||
if backend_path in settings.AUTHENTICATION_BACKENDS:
|
||||
backend = load_backend(backend_path)
|
||||
user = backend.get_user(user_id)
|
||||
# Verify the session
|
||||
if hasattr(user, "get_session_auth_hash"):
|
||||
session_hash = session.get(HASH_SESSION_KEY)
|
||||
session_hash_verified = session_hash and constant_time_compare(
|
||||
session_hash, user.get_session_auth_hash()
|
||||
)
|
||||
if not session_hash_verified:
|
||||
session.flush()
|
||||
user = None
|
||||
return user or AnonymousUser()
|
||||
|
||||
|
||||
@database_sync_to_async
|
||||
def login(scope, user, backend=None):
|
||||
"""
|
||||
Persist a user id and a backend in the request.
|
||||
This way a user doesn't have to re-authenticate on every request.
|
||||
Note that data set during the anonymous session is retained when the user logs in.
|
||||
"""
|
||||
if "session" not in scope:
|
||||
raise ValueError(
|
||||
"Cannot find session in scope. You should wrap your consumer in SessionMiddleware."
|
||||
)
|
||||
session = scope["session"]
|
||||
session_auth_hash = ""
|
||||
if user is None:
|
||||
user = scope.get("user", None)
|
||||
if user is None:
|
||||
raise ValueError(
|
||||
"User must be passed as an argument or must be present in the scope."
|
||||
)
|
||||
if hasattr(user, "get_session_auth_hash"):
|
||||
session_auth_hash = user.get_session_auth_hash()
|
||||
if SESSION_KEY in session:
|
||||
if _get_user_session_key(session) != user.pk or (
|
||||
session_auth_hash
|
||||
and not constant_time_compare(
|
||||
session.get(HASH_SESSION_KEY, ""), session_auth_hash
|
||||
)
|
||||
):
|
||||
# To avoid reusing another user's session, create a new, empty
|
||||
# session if the existing session corresponds to a different
|
||||
# authenticated user.
|
||||
session.flush()
|
||||
else:
|
||||
session.cycle_key()
|
||||
try:
|
||||
backend = backend or user.backend
|
||||
except AttributeError:
|
||||
backends = _get_backends(return_tuples=True)
|
||||
if len(backends) == 1:
|
||||
_, backend = backends[0]
|
||||
else:
|
||||
raise ValueError(
|
||||
"You have multiple authentication backends configured and therefore must provide the `backend` "
|
||||
"argument or set the `backend` attribute on the user."
|
||||
)
|
||||
session[SESSION_KEY] = user._meta.pk.value_to_string(user)
|
||||
session[BACKEND_SESSION_KEY] = backend
|
||||
session[HASH_SESSION_KEY] = session_auth_hash
|
||||
scope["user"] = user
|
||||
# note this does not reset the CSRF_COOKIE/Token
|
||||
user_logged_in.send(sender=user.__class__, request=None, user=user)
|
||||
|
||||
|
||||
@database_sync_to_async
|
||||
def logout(scope):
|
||||
"""
|
||||
Remove the authenticated user's ID from the request and flush their session data.
|
||||
"""
|
||||
if "session" not in scope:
|
||||
raise ValueError(
|
||||
"Login cannot find session in scope. You should wrap your consumer in SessionMiddleware."
|
||||
)
|
||||
session = scope["session"]
|
||||
# Dispatch the signal before the user is logged out so the receivers have a
|
||||
# chance to find out *who* logged out.
|
||||
user = scope.get("user", None)
|
||||
if hasattr(user, "is_authenticated") and not user.is_authenticated:
|
||||
user = None
|
||||
if user is not None:
|
||||
user_logged_out.send(sender=user.__class__, request=None, user=user)
|
||||
# remember language choice saved to session
|
||||
language = session.get(LANGUAGE_SESSION_KEY)
|
||||
session.flush()
|
||||
if language is not None:
|
||||
session[LANGUAGE_SESSION_KEY] = language
|
||||
if "user" in scope:
|
||||
scope["user"] = AnonymousUser()
|
||||
|
||||
|
||||
def _get_user_session_key(session):
|
||||
# This value in the session is always serialized to a string, so we need
|
||||
# to convert it back to Python whenever we access it.
|
||||
return get_user_model()._meta.pk.to_python(session[SESSION_KEY])
|
||||
|
||||
|
||||
class UserLazyObject(LazyObject):
|
||||
"""
|
||||
Throw a more useful error message when scope['user'] is accessed before it's resolved
|
||||
"""
|
||||
|
||||
def _setup(self):
|
||||
raise ValueError("Accessing scope user before it is ready.")
|
||||
|
||||
|
||||
class AuthMiddleware(BaseMiddleware):
|
||||
"""
|
||||
Middleware which populates scope["user"] from a Django session.
|
||||
Requires SessionMiddleware to function.
|
||||
"""
|
||||
|
||||
def populate_scope(self, scope):
|
||||
# Make sure we have a session
|
||||
if "session" not in scope:
|
||||
raise ValueError(
|
||||
"AuthMiddleware cannot find session in scope. SessionMiddleware must be above it."
|
||||
)
|
||||
# Add it to the scope if it's not there already
|
||||
if "user" not in scope:
|
||||
scope["user"] = UserLazyObject()
|
||||
|
||||
async def resolve_scope(self, scope):
|
||||
scope["user"]._wrapped = await get_user(scope)
|
||||
|
||||
|
||||
# Handy shortcut for applying all three layers at once
|
||||
AuthMiddlewareStack = lambda inner: CookieMiddleware(
|
||||
SessionMiddleware(AuthMiddleware(inner))
|
||||
)
|
||||
@@ -0,0 +1,113 @@
|
||||
import functools
|
||||
|
||||
from asgiref.sync import async_to_sync
|
||||
|
||||
from . import DEFAULT_CHANNEL_LAYER
|
||||
from .db import database_sync_to_async
|
||||
from .exceptions import StopConsumer
|
||||
from .layers import get_channel_layer
|
||||
from .utils import await_many_dispatch
|
||||
|
||||
|
||||
def get_handler_name(message):
|
||||
"""
|
||||
Looks at a message, checks it has a sensible type, and returns the
|
||||
handler name for that type.
|
||||
"""
|
||||
# Check message looks OK
|
||||
if "type" not in message:
|
||||
raise ValueError("Incoming message has no 'type' attribute")
|
||||
if message["type"].startswith("_"):
|
||||
raise ValueError("Malformed type in message (leading underscore)")
|
||||
# Extract type and replace . with _
|
||||
return message["type"].replace(".", "_")
|
||||
|
||||
|
||||
class AsyncConsumer:
|
||||
"""
|
||||
Base consumer class. Implements the ASGI application spec, and adds on
|
||||
channel layer management and routing of events to named methods based
|
||||
on their type.
|
||||
"""
|
||||
|
||||
_sync = False
|
||||
channel_layer_alias = DEFAULT_CHANNEL_LAYER
|
||||
|
||||
def __init__(self, scope):
|
||||
self.scope = scope
|
||||
|
||||
async def __call__(self, receive, send):
|
||||
"""
|
||||
Dispatches incoming messages to type-based handlers asynchronously.
|
||||
"""
|
||||
# Initialize channel layer
|
||||
self.channel_layer = get_channel_layer(self.channel_layer_alias)
|
||||
if self.channel_layer is not None:
|
||||
self.channel_name = await self.channel_layer.new_channel()
|
||||
self.channel_receive = functools.partial(
|
||||
self.channel_layer.receive, self.channel_name
|
||||
)
|
||||
# Store send function
|
||||
if self._sync:
|
||||
self.base_send = async_to_sync(send)
|
||||
else:
|
||||
self.base_send = send
|
||||
# Pass messages in from channel layer or client to dispatch method
|
||||
try:
|
||||
if self.channel_layer is not None:
|
||||
await await_many_dispatch(
|
||||
[receive, self.channel_receive], self.dispatch
|
||||
)
|
||||
else:
|
||||
await await_many_dispatch([receive], self.dispatch)
|
||||
except StopConsumer:
|
||||
# Exit cleanly
|
||||
pass
|
||||
|
||||
async def dispatch(self, message):
|
||||
"""
|
||||
Works out what to do with a message.
|
||||
"""
|
||||
handler = getattr(self, get_handler_name(message), None)
|
||||
if handler:
|
||||
await handler(message)
|
||||
else:
|
||||
raise ValueError("No handler for message type %s" % message["type"])
|
||||
|
||||
async def send(self, message):
|
||||
"""
|
||||
Overrideable/callable-by-subclasses send method.
|
||||
"""
|
||||
await self.base_send(message)
|
||||
|
||||
|
||||
class SyncConsumer(AsyncConsumer):
|
||||
"""
|
||||
Synchronous version of the consumer, which is what we write most of the
|
||||
generic consumers against (for now). Calls handlers in a threadpool and
|
||||
uses CallBouncer to get the send method out to the main event loop.
|
||||
|
||||
It would have been possible to have "mixed" consumers and auto-detect
|
||||
if a handler was awaitable or not, but that would have made the API
|
||||
for user-called methods very confusing as there'd be two types of each.
|
||||
"""
|
||||
|
||||
_sync = True
|
||||
|
||||
@database_sync_to_async
|
||||
def dispatch(self, message):
|
||||
"""
|
||||
Dispatches incoming messages to type-based handlers asynchronously.
|
||||
"""
|
||||
# Get and execute the handler
|
||||
handler = getattr(self, get_handler_name(message), None)
|
||||
if handler:
|
||||
handler(message)
|
||||
else:
|
||||
raise ValueError("No handler for message type %s" % message["type"])
|
||||
|
||||
def send(self, message):
|
||||
"""
|
||||
Overrideable/callable-by-subclasses send method.
|
||||
"""
|
||||
self.base_send(message)
|
||||
@@ -0,0 +1,12 @@
|
||||
def monkeypatch_django():
|
||||
"""
|
||||
Monkeypatches support for us into parts of Django.
|
||||
"""
|
||||
# Ensure that the staticfiles version of runserver bows down to us
|
||||
# This one is particularly horrible
|
||||
from django.contrib.staticfiles.management.commands.runserver import (
|
||||
Command as StaticRunserverCommand,
|
||||
)
|
||||
from .management.commands.runserver import Command as RunserverCommand
|
||||
|
||||
StaticRunserverCommand.__bases__ = (RunserverCommand,)
|
||||
@@ -0,0 +1,371 @@
|
||||
import cgi
|
||||
import codecs
|
||||
import logging
|
||||
import sys
|
||||
import tempfile
|
||||
import traceback
|
||||
|
||||
from django import http
|
||||
from django.conf import settings
|
||||
from django.core import signals
|
||||
from django.core.exceptions import RequestDataTooBig
|
||||
from django.core.handlers import base
|
||||
from django.http import FileResponse, HttpResponse, HttpResponseServerError
|
||||
from django.urls import set_script_prefix
|
||||
from django.utils.functional import cached_property
|
||||
|
||||
from asgiref.sync import async_to_sync, sync_to_async
|
||||
from channels.exceptions import RequestAborted, RequestTimeout
|
||||
|
||||
logger = logging.getLogger("django.request")
|
||||
|
||||
|
||||
class AsgiRequest(http.HttpRequest):
|
||||
"""
|
||||
Custom request subclass that decodes from an ASGI-standard request
|
||||
dict, and wraps request body handling.
|
||||
"""
|
||||
|
||||
# Number of seconds until a Request gives up on trying to read a request
|
||||
# body and aborts.
|
||||
body_receive_timeout = 60
|
||||
|
||||
def __init__(self, scope, stream):
|
||||
self.scope = scope
|
||||
self._content_length = 0
|
||||
self._post_parse_error = False
|
||||
self._read_started = False
|
||||
self.resolver_match = None
|
||||
self.script_name = self.scope.get("root_path", "")
|
||||
if self.script_name and scope["path"].startswith(self.script_name):
|
||||
# TODO: Better is-prefix checking, slash handling?
|
||||
self.path_info = scope["path"][len(self.script_name) :]
|
||||
else:
|
||||
self.path_info = scope["path"]
|
||||
|
||||
# django path is different from asgi scope path args, it should combine with script name
|
||||
if self.script_name:
|
||||
self.path = "%s/%s" % (
|
||||
self.script_name.rstrip("/"),
|
||||
self.path_info.replace("/", "", 1),
|
||||
)
|
||||
else:
|
||||
self.path = scope["path"]
|
||||
|
||||
# HTTP basics
|
||||
self.method = self.scope["method"].upper()
|
||||
# fix https://github.com/django/channels/issues/622
|
||||
query_string = self.scope.get("query_string", "")
|
||||
if isinstance(query_string, bytes):
|
||||
query_string = query_string.decode("utf-8")
|
||||
self.META = {
|
||||
"REQUEST_METHOD": self.method,
|
||||
"QUERY_STRING": query_string,
|
||||
"SCRIPT_NAME": self.script_name,
|
||||
"PATH_INFO": self.path_info,
|
||||
# Old code will need these for a while
|
||||
"wsgi.multithread": True,
|
||||
"wsgi.multiprocess": True,
|
||||
}
|
||||
if self.scope.get("client", None):
|
||||
self.META["REMOTE_ADDR"] = self.scope["client"][0]
|
||||
self.META["REMOTE_HOST"] = self.META["REMOTE_ADDR"]
|
||||
self.META["REMOTE_PORT"] = self.scope["client"][1]
|
||||
if self.scope.get("server", None):
|
||||
self.META["SERVER_NAME"] = self.scope["server"][0]
|
||||
self.META["SERVER_PORT"] = str(self.scope["server"][1])
|
||||
else:
|
||||
self.META["SERVER_NAME"] = "unknown"
|
||||
self.META["SERVER_PORT"] = "0"
|
||||
# Handle old style-headers for a transition period
|
||||
if "headers" in self.scope and isinstance(self.scope["headers"], dict):
|
||||
self.scope["headers"] = [
|
||||
(x.encode("latin1"), y) for x, y in self.scope["headers"].items()
|
||||
]
|
||||
# Headers go into META
|
||||
for name, value in self.scope.get("headers", []):
|
||||
name = name.decode("latin1")
|
||||
if name == "content-length":
|
||||
corrected_name = "CONTENT_LENGTH"
|
||||
elif name == "content-type":
|
||||
corrected_name = "CONTENT_TYPE"
|
||||
else:
|
||||
corrected_name = "HTTP_%s" % name.upper().replace("-", "_")
|
||||
# HTTPbis say only ASCII chars are allowed in headers, but we latin1 just in case
|
||||
value = value.decode("latin1")
|
||||
if corrected_name in self.META:
|
||||
value = self.META[corrected_name] + "," + value
|
||||
self.META[corrected_name] = value
|
||||
# Pull out request encoding if we find it
|
||||
if "CONTENT_TYPE" in self.META:
|
||||
self.content_type, self.content_params = cgi.parse_header(
|
||||
self.META["CONTENT_TYPE"]
|
||||
)
|
||||
if "charset" in self.content_params:
|
||||
try:
|
||||
codecs.lookup(self.content_params["charset"])
|
||||
except LookupError:
|
||||
pass
|
||||
else:
|
||||
self.encoding = self.content_params["charset"]
|
||||
else:
|
||||
self.content_type, self.content_params = "", {}
|
||||
# Pull out content length info
|
||||
if self.META.get("CONTENT_LENGTH", None):
|
||||
try:
|
||||
self._content_length = int(self.META["CONTENT_LENGTH"])
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
# Body handling
|
||||
self._stream = stream
|
||||
# Other bits
|
||||
self.resolver_match = None
|
||||
|
||||
@cached_property
|
||||
def GET(self):
|
||||
return http.QueryDict(self.scope.get("query_string", ""))
|
||||
|
||||
def _get_scheme(self):
|
||||
return self.scope.get("scheme", "http")
|
||||
|
||||
def _get_post(self):
|
||||
if not hasattr(self, "_post"):
|
||||
self._load_post_and_files()
|
||||
return self._post
|
||||
|
||||
def _set_post(self, post):
|
||||
self._post = post
|
||||
|
||||
def _get_files(self):
|
||||
if not hasattr(self, "_files"):
|
||||
self._load_post_and_files()
|
||||
return self._files
|
||||
|
||||
POST = property(_get_post, _set_post)
|
||||
FILES = property(_get_files)
|
||||
|
||||
@cached_property
|
||||
def COOKIES(self):
|
||||
return http.parse_cookie(self.META.get("HTTP_COOKIE", ""))
|
||||
|
||||
|
||||
class AsgiHandler(base.BaseHandler):
|
||||
"""
|
||||
Handler for ASGI requests for the view system only (it will have got here
|
||||
after traversing the dispatch-by-channel-name system, which decides it's
|
||||
a HTTP request)
|
||||
|
||||
You can also manually construct it with a get_response callback if you
|
||||
want to run a single Django view yourself. If you do this, though, it will
|
||||
not do any URL routing or middleware (Channels uses it for staticfiles'
|
||||
serving code)
|
||||
"""
|
||||
|
||||
request_class = AsgiRequest
|
||||
|
||||
# Size to chunk response bodies into for multiple response messages
|
||||
chunk_size = 512 * 1024
|
||||
|
||||
def __init__(self, scope):
|
||||
if scope["type"] != "http":
|
||||
raise ValueError(
|
||||
"The AsgiHandler can only handle HTTP connections, not %s"
|
||||
% scope["type"]
|
||||
)
|
||||
super(AsgiHandler, self).__init__()
|
||||
self.scope = scope
|
||||
self.load_middleware()
|
||||
|
||||
async def __call__(self, receive, send):
|
||||
"""
|
||||
Async entrypoint - uses the sync_to_async wrapper to run things in a
|
||||
threadpool.
|
||||
"""
|
||||
self.send = async_to_sync(send)
|
||||
|
||||
# Receive the HTTP request body as a stream object.
|
||||
try:
|
||||
body_stream = await self.read_body(receive)
|
||||
except RequestAborted:
|
||||
return
|
||||
# Launch into body handling (and a synchronous subthread).
|
||||
await self.handle(body_stream)
|
||||
|
||||
async def read_body(self, receive):
|
||||
"""Reads a HTTP body from an ASGI connection."""
|
||||
# Use the tempfile that auto rolls-over to a disk file as it fills up.
|
||||
body_file = tempfile.SpooledTemporaryFile(
|
||||
max_size=settings.FILE_UPLOAD_MAX_MEMORY_SIZE, mode="w+b"
|
||||
)
|
||||
while True:
|
||||
message = await receive()
|
||||
if message["type"] == "http.disconnect":
|
||||
# Early client disconnect.
|
||||
raise RequestAborted()
|
||||
# Add a body chunk from the message, if provided.
|
||||
if "body" in message:
|
||||
body_file.write(message["body"])
|
||||
# Quit out if that's the end.
|
||||
if not message.get("more_body", False):
|
||||
break
|
||||
body_file.seek(0)
|
||||
return body_file
|
||||
|
||||
@sync_to_async
|
||||
def handle(self, body):
|
||||
"""
|
||||
Synchronous message processing.
|
||||
"""
|
||||
# Set script prefix from message root_path, turning None into empty string
|
||||
script_prefix = self.scope.get("root_path", "") or ""
|
||||
if settings.FORCE_SCRIPT_NAME:
|
||||
script_prefix = settings.FORCE_SCRIPT_NAME
|
||||
set_script_prefix(script_prefix)
|
||||
signals.request_started.send(sender=self.__class__, scope=self.scope)
|
||||
# Run request through view system
|
||||
try:
|
||||
request = self.request_class(self.scope, body)
|
||||
except UnicodeDecodeError:
|
||||
logger.warning(
|
||||
"Bad Request (UnicodeDecodeError)",
|
||||
exc_info=sys.exc_info(),
|
||||
extra={"status_code": 400},
|
||||
)
|
||||
response = http.HttpResponseBadRequest()
|
||||
except RequestTimeout:
|
||||
# Parsing the request failed, so the response is a Request Timeout error
|
||||
response = HttpResponse("408 Request Timeout (upload too slow)", status=408)
|
||||
except RequestAborted:
|
||||
# Client closed connection on us mid request. Abort!
|
||||
return
|
||||
except RequestDataTooBig:
|
||||
response = HttpResponse("413 Payload too large", status=413)
|
||||
else:
|
||||
response = self.get_response(request)
|
||||
# Fix chunk size on file responses
|
||||
if isinstance(response, FileResponse):
|
||||
response.block_size = 1024 * 512
|
||||
# Transform response into messages, which we yield back to caller
|
||||
for response_message in self.encode_response(response):
|
||||
self.send(response_message)
|
||||
# Close the response now we're done with it
|
||||
response.close()
|
||||
|
||||
def handle_uncaught_exception(self, request, resolver, exc_info):
|
||||
"""
|
||||
Last-chance handler for exceptions.
|
||||
"""
|
||||
# There's no WSGI server to catch the exception further up if this fails,
|
||||
# so translate it into a plain text response.
|
||||
try:
|
||||
return super(AsgiHandler, self).handle_uncaught_exception(
|
||||
request, resolver, exc_info
|
||||
)
|
||||
except Exception:
|
||||
return HttpResponseServerError(
|
||||
traceback.format_exc() if settings.DEBUG else "Internal Server Error",
|
||||
content_type="text/plain",
|
||||
)
|
||||
|
||||
def load_middleware(self):
|
||||
"""
|
||||
Loads the Django middleware chain and caches it on the class.
|
||||
"""
|
||||
# Because we create an AsgiHandler on every HTTP request
|
||||
# we need to preserve the Django middleware chain once we load it.
|
||||
if (
|
||||
hasattr(self.__class__, "_middleware_chain")
|
||||
and self.__class__._middleware_chain
|
||||
):
|
||||
self._middleware_chain = self.__class__._middleware_chain
|
||||
self._view_middleware = self.__class__._view_middleware
|
||||
self._template_response_middleware = (
|
||||
self.__class__._template_response_middleware
|
||||
)
|
||||
self._exception_middleware = self.__class__._exception_middleware
|
||||
|
||||
# Support additional arguments for Django 1.11 and 2.0.
|
||||
if hasattr(self.__class__, "_request_middleware"):
|
||||
self._request_middleware = self.__class__._request_middleware
|
||||
self._response_middleware = self.__class__._response_middleware
|
||||
|
||||
else:
|
||||
super(AsgiHandler, self).load_middleware()
|
||||
self.__class__._middleware_chain = self._middleware_chain
|
||||
self.__class__._view_middleware = self._view_middleware
|
||||
self.__class__._template_response_middleware = (
|
||||
self._template_response_middleware
|
||||
)
|
||||
self.__class__._exception_middleware = self._exception_middleware
|
||||
|
||||
# Support additional arguments for Django 1.11 and 2.0.
|
||||
if hasattr(self, "_request_middleware"):
|
||||
self.__class__._request_middleware = self._request_middleware
|
||||
self.__class__._response_middleware = self._response_middleware
|
||||
|
||||
@classmethod
|
||||
def encode_response(cls, response):
|
||||
"""
|
||||
Encodes a Django HTTP response into ASGI http.response message(s).
|
||||
"""
|
||||
# Collect cookies into headers.
|
||||
# Note that we have to preserve header case as there are some non-RFC
|
||||
# compliant clients that want things like Content-Type correct. Ugh.
|
||||
response_headers = []
|
||||
for header, value in response.items():
|
||||
if isinstance(header, str):
|
||||
header = header.encode("ascii")
|
||||
if isinstance(value, str):
|
||||
value = value.encode("latin1")
|
||||
response_headers.append((bytes(header), bytes(value)))
|
||||
for c in response.cookies.values():
|
||||
response_headers.append(
|
||||
(b"Set-Cookie", c.output(header="").encode("ascii").strip())
|
||||
)
|
||||
# Make initial response message
|
||||
yield {
|
||||
"type": "http.response.start",
|
||||
"status": response.status_code,
|
||||
"headers": response_headers,
|
||||
}
|
||||
# Streaming responses need to be pinned to their iterator
|
||||
if response.streaming:
|
||||
# Access `__iter__` and not `streaming_content` directly in case
|
||||
# it has been overridden in a subclass.
|
||||
for part in response:
|
||||
for chunk, _ in cls.chunk_bytes(part):
|
||||
yield {
|
||||
"type": "http.response.body",
|
||||
"body": chunk,
|
||||
# We ignore "more" as there may be more parts; instead,
|
||||
# we use an empty final closing message with False.
|
||||
"more_body": True,
|
||||
}
|
||||
# Final closing message
|
||||
yield {"type": "http.response.body"}
|
||||
# Other responses just need chunking
|
||||
else:
|
||||
# Yield chunks of response
|
||||
for chunk, last in cls.chunk_bytes(response.content):
|
||||
yield {
|
||||
"type": "http.response.body",
|
||||
"body": chunk,
|
||||
"more_body": not last,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def chunk_bytes(cls, data):
|
||||
"""
|
||||
Chunks some data up so it can be sent in reasonable size messages.
|
||||
Yields (chunk, last_chunk) tuples.
|
||||
"""
|
||||
position = 0
|
||||
if not data:
|
||||
yield data, True
|
||||
return
|
||||
while position < len(data):
|
||||
yield (
|
||||
data[position : position + cls.chunk_size],
|
||||
(position + cls.chunk_size) >= len(data),
|
||||
)
|
||||
position += cls.chunk_size
|
||||
@@ -0,0 +1,41 @@
|
||||
from functools import partial
|
||||
|
||||
|
||||
class BaseMiddleware:
|
||||
"""
|
||||
Base class for implementing ASGI middleware. Inherit from this and
|
||||
override the setup() method if you want to do things before you
|
||||
get to.
|
||||
|
||||
Note that subclasses of this are not self-safe; don't store state on
|
||||
the instance, as it serves multiple application instances. Instead, use
|
||||
scope.
|
||||
"""
|
||||
|
||||
def __init__(self, inner):
|
||||
"""
|
||||
Middleware constructor - just takes inner application.
|
||||
"""
|
||||
self.inner = inner
|
||||
|
||||
def __call__(self, scope):
|
||||
"""
|
||||
ASGI constructor; can insert things into the scope, but not
|
||||
run asynchronous code.
|
||||
"""
|
||||
# Copy scope to stop changes going upstream
|
||||
scope = dict(scope)
|
||||
# Allow subclasses to change the scope
|
||||
self.populate_scope(scope)
|
||||
# Call the inner application's init
|
||||
inner_instance = self.inner(scope)
|
||||
# Partially bind it to our coroutine entrypoint along with the scope
|
||||
return partial(self.coroutine_call, inner_instance, scope)
|
||||
|
||||
async def coroutine_call(self, inner_instance, scope, receive, send):
|
||||
"""
|
||||
ASGI coroutine; where we can resolve items in the scope
|
||||
(but you can't modify it at the top level here!)
|
||||
"""
|
||||
await self.resolve_scope(scope)
|
||||
await inner_instance(receive, send)
|
||||
@@ -0,0 +1,253 @@
|
||||
import datetime
|
||||
import time
|
||||
from importlib import import_module
|
||||
|
||||
from django.conf import settings
|
||||
from django.contrib.sessions.backends.base import UpdateError
|
||||
from django.core.exceptions import SuspiciousOperation
|
||||
from django.http import parse_cookie
|
||||
from django.http.cookie import SimpleCookie
|
||||
from django.utils import timezone
|
||||
from django.utils.encoding import force_str
|
||||
from django.utils.functional import LazyObject
|
||||
|
||||
from channels.db import database_sync_to_async
|
||||
|
||||
try:
|
||||
from django.utils.http import http_date
|
||||
except ImportError:
|
||||
from django.utils.http import cookie_date as http_date
|
||||
|
||||
|
||||
class CookieMiddleware:
|
||||
"""
|
||||
Extracts cookies from HTTP or WebSocket-style scopes and adds them as a
|
||||
scope["cookies"] entry with the same format as Django's request.COOKIES.
|
||||
"""
|
||||
|
||||
def __init__(self, inner):
|
||||
self.inner = inner
|
||||
|
||||
def __call__(self, scope):
|
||||
# Check this actually has headers. They're a required scope key for HTTP and WS.
|
||||
if "headers" not in scope:
|
||||
raise ValueError(
|
||||
"CookieMiddleware was passed a scope that did not have a headers key "
|
||||
+ "(make sure it is only passed HTTP or WebSocket connections)"
|
||||
)
|
||||
# Go through headers to find the cookie one
|
||||
for name, value in scope.get("headers", []):
|
||||
if name == b"cookie":
|
||||
cookies = parse_cookie(value.decode("ascii"))
|
||||
break
|
||||
else:
|
||||
# No cookie header found - add an empty default.
|
||||
cookies = {}
|
||||
# Return inner application
|
||||
return self.inner(dict(scope, cookies=cookies))
|
||||
|
||||
@classmethod
|
||||
def set_cookie(
|
||||
cls,
|
||||
message,
|
||||
key,
|
||||
value="",
|
||||
max_age=None,
|
||||
expires=None,
|
||||
path="/",
|
||||
domain=None,
|
||||
secure=False,
|
||||
httponly=False,
|
||||
):
|
||||
"""
|
||||
Sets a cookie in the passed HTTP response message.
|
||||
|
||||
``expires`` can be:
|
||||
- a string in the correct format,
|
||||
- a naive ``datetime.datetime`` object in UTC,
|
||||
- an aware ``datetime.datetime`` object in any time zone.
|
||||
If it is a ``datetime.datetime`` object then ``max_age`` will be calculated.
|
||||
"""
|
||||
value = force_str(value)
|
||||
cookies = SimpleCookie()
|
||||
cookies[key] = value
|
||||
if expires is not None:
|
||||
if isinstance(expires, datetime.datetime):
|
||||
if timezone.is_aware(expires):
|
||||
expires = timezone.make_naive(expires, timezone.utc)
|
||||
delta = expires - expires.utcnow()
|
||||
# Add one second so the date matches exactly (a fraction of
|
||||
# time gets lost between converting to a timedelta and
|
||||
# then the date string).
|
||||
delta = delta + datetime.timedelta(seconds=1)
|
||||
# Just set max_age - the max_age logic will set expires.
|
||||
expires = None
|
||||
max_age = max(0, delta.days * 86400 + delta.seconds)
|
||||
else:
|
||||
cookies[key]["expires"] = expires
|
||||
else:
|
||||
cookies[key]["expires"] = ""
|
||||
if max_age is not None:
|
||||
cookies[key]["max-age"] = max_age
|
||||
# IE requires expires, so set it if hasn't been already.
|
||||
if not expires:
|
||||
cookies[key]["expires"] = http_date(time.time() + max_age)
|
||||
if path is not None:
|
||||
cookies[key]["path"] = path
|
||||
if domain is not None:
|
||||
cookies[key]["domain"] = domain
|
||||
if secure:
|
||||
cookies[key]["secure"] = True
|
||||
if httponly:
|
||||
cookies[key]["httponly"] = True
|
||||
# Write out the cookies to the response
|
||||
for c in cookies.values():
|
||||
message.setdefault("headers", []).append(
|
||||
(b"Set-Cookie", bytes(c.output(header=""), encoding="utf-8"))
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def delete_cookie(cls, message, key, path="/", domain=None):
|
||||
"""
|
||||
Deletes a cookie in a response.
|
||||
"""
|
||||
return cls.set_cookie(
|
||||
message,
|
||||
key,
|
||||
max_age=0,
|
||||
path=path,
|
||||
domain=domain,
|
||||
expires="Thu, 01-Jan-1970 00:00:00 GMT",
|
||||
)
|
||||
|
||||
|
||||
class SessionMiddleware:
|
||||
"""
|
||||
Class that adds Django sessions (from HTTP cookies) to the
|
||||
scope. Works with HTTP or WebSocket protocol types (or anything that
|
||||
provides a "headers" entry in the scope).
|
||||
|
||||
Requires the CookieMiddleware to be higher up in the stack.
|
||||
"""
|
||||
|
||||
# Message types that trigger a session save if it's modified
|
||||
save_message_types = ["http.response.start"]
|
||||
|
||||
# Message types that can carry session cookies back
|
||||
cookie_response_message_types = ["http.response.start"]
|
||||
|
||||
def __init__(self, inner):
|
||||
self.inner = inner
|
||||
self.cookie_name = settings.SESSION_COOKIE_NAME
|
||||
self.session_store = import_module(settings.SESSION_ENGINE).SessionStore
|
||||
|
||||
def __call__(self, scope):
|
||||
return SessionMiddlewareInstance(scope, self)
|
||||
|
||||
|
||||
class SessionMiddlewareInstance:
|
||||
"""
|
||||
Inner class that is instantiated once per scope.
|
||||
"""
|
||||
|
||||
def __init__(self, scope, middleware):
|
||||
self.middleware = middleware
|
||||
self.scope = dict(scope)
|
||||
if "session" in self.scope:
|
||||
# There's already session middleware of some kind above us, pass that through
|
||||
self.activated = False
|
||||
else:
|
||||
# Make sure there are cookies in the scope
|
||||
if "cookies" not in self.scope:
|
||||
raise ValueError(
|
||||
"No cookies in scope - SessionMiddleware needs to run inside of CookieMiddleware."
|
||||
)
|
||||
# Parse the headers in the scope into cookies
|
||||
self.scope["session"] = LazyObject()
|
||||
self.activated = True
|
||||
# Instantiate our inner application
|
||||
self.inner = self.middleware.inner(self.scope)
|
||||
|
||||
async def __call__(self, receive, send):
|
||||
"""
|
||||
We intercept the send() callable so we can do session saves and
|
||||
add session cookie overrides to send back.
|
||||
"""
|
||||
# Resolve the session now we can do it in a blocking way
|
||||
session_key = self.scope["cookies"].get(self.middleware.cookie_name)
|
||||
self.scope["session"]._wrapped = await database_sync_to_async(
|
||||
self.middleware.session_store
|
||||
)(session_key)
|
||||
# Override send
|
||||
self.real_send = send
|
||||
return await self.inner(receive, self.send)
|
||||
|
||||
async def send(self, message):
|
||||
"""
|
||||
Overridden send that also does session saves/cookies.
|
||||
"""
|
||||
# Only save session if we're the outermost session middleware
|
||||
if self.activated:
|
||||
modified = self.scope["session"].modified
|
||||
empty = self.scope["session"].is_empty()
|
||||
# If this is a message type that we want to save on, and there's
|
||||
# changed data, save it. We also save if it's empty as we might
|
||||
# not be able to send a cookie-delete along with this message.
|
||||
if (
|
||||
message["type"] in self.middleware.save_message_types
|
||||
and message.get("status", 200) != 500
|
||||
and (modified or settings.SESSION_SAVE_EVERY_REQUEST)
|
||||
):
|
||||
self.save_session()
|
||||
# If this is a message type that can transport cookies back to the
|
||||
# client, then do so.
|
||||
if message["type"] in self.middleware.cookie_response_message_types:
|
||||
if empty:
|
||||
# Delete cookie if it's set
|
||||
if settings.SESSION_COOKIE_NAME in self.scope["cookies"]:
|
||||
CookieMiddleware.delete_cookie(
|
||||
message,
|
||||
settings.SESSION_COOKIE_NAME,
|
||||
path=settings.SESSION_COOKIE_PATH,
|
||||
domain=settings.SESSION_COOKIE_DOMAIN,
|
||||
)
|
||||
else:
|
||||
# Get the expiry data
|
||||
if self.scope["session"].get_expire_at_browser_close():
|
||||
max_age = None
|
||||
expires = None
|
||||
else:
|
||||
max_age = self.scope["session"].get_expiry_age()
|
||||
expires_time = time.time() + max_age
|
||||
expires = http_date(expires_time)
|
||||
# Set the cookie
|
||||
CookieMiddleware.set_cookie(
|
||||
message,
|
||||
self.middleware.cookie_name,
|
||||
self.scope["session"].session_key,
|
||||
max_age=max_age,
|
||||
expires=expires,
|
||||
domain=settings.SESSION_COOKIE_DOMAIN,
|
||||
path=settings.SESSION_COOKIE_PATH,
|
||||
secure=settings.SESSION_COOKIE_SECURE or None,
|
||||
httponly=settings.SESSION_COOKIE_HTTPONLY or None,
|
||||
)
|
||||
# Pass up the send
|
||||
return await self.real_send(message)
|
||||
|
||||
def save_session(self):
|
||||
"""
|
||||
Saves the current session.
|
||||
"""
|
||||
try:
|
||||
self.scope["session"].save()
|
||||
except UpdateError:
|
||||
raise SuspiciousOperation(
|
||||
"The request's session was deleted before the "
|
||||
"request completed. The user may have logged "
|
||||
"out in a concurrent request, for example."
|
||||
)
|
||||
|
||||
|
||||
# Shortcut to include cookie middleware
|
||||
SessionMiddlewareStack = lambda inner: CookieMiddleware(SessionMiddleware(inner))
|
||||
@@ -0,0 +1,56 @@
|
||||
from urllib.parse import unquote, urlparse
|
||||
|
||||
from asgiref.testing import ApplicationCommunicator
|
||||
|
||||
|
||||
class HttpCommunicator(ApplicationCommunicator):
|
||||
"""
|
||||
ApplicationCommunicator subclass that has HTTP shortcut methods.
|
||||
|
||||
It will construct the scope for you, so you need to pass the application
|
||||
(uninstantiated) along with HTTP parameters.
|
||||
|
||||
This does not support full chunking - for that, just use ApplicationCommunicator
|
||||
directly.
|
||||
"""
|
||||
|
||||
def __init__(self, application, method, path, body=b"", headers=None):
|
||||
parsed = urlparse(path)
|
||||
self.scope = {
|
||||
"type": "http",
|
||||
"http_version": "1.1",
|
||||
"method": method.upper(),
|
||||
"path": unquote(parsed.path),
|
||||
"query_string": parsed.query.encode("utf-8"),
|
||||
"headers": headers or [],
|
||||
}
|
||||
assert isinstance(body, bytes)
|
||||
self.body = body
|
||||
self.sent_request = False
|
||||
super().__init__(application, self.scope)
|
||||
|
||||
async def get_response(self, timeout=1):
|
||||
"""
|
||||
Get the application's response. Returns a dict with keys of
|
||||
"body", "headers" and "status".
|
||||
"""
|
||||
# If we've not sent the request yet, do so
|
||||
if not self.sent_request:
|
||||
self.sent_request = True
|
||||
await self.send_input({"type": "http.request", "body": self.body})
|
||||
# Get the response start
|
||||
response_start = await self.receive_output(timeout)
|
||||
assert response_start["type"] == "http.response.start"
|
||||
# Get all body parts
|
||||
response_start["body"] = b""
|
||||
while True:
|
||||
chunk = await self.receive_output(timeout)
|
||||
assert chunk["type"] == "http.response.body"
|
||||
assert isinstance(chunk["body"], bytes)
|
||||
response_start["body"] += chunk["body"]
|
||||
if not chunk.get("more_body", False):
|
||||
break
|
||||
# Return structured info
|
||||
del response_start["type"]
|
||||
response_start.setdefault("headers", [])
|
||||
return response_start
|
||||
@@ -0,0 +1,67 @@
|
||||
from django.core.exceptions import ImproperlyConfigured
|
||||
from django.db import connections
|
||||
from django.test.testcases import TransactionTestCase
|
||||
from django.test.utils import modify_settings
|
||||
|
||||
from channels.routing import get_default_application
|
||||
from channels.staticfiles import StaticFilesWrapper
|
||||
from daphne.testing import DaphneProcess
|
||||
|
||||
|
||||
class ChannelsLiveServerTestCase(TransactionTestCase):
|
||||
"""
|
||||
Does basically the same as TransactionTestCase but also launches a
|
||||
live Daphne server in a separate process, so
|
||||
that the tests may use another test framework, such as Selenium,
|
||||
instead of the built-in dummy client.
|
||||
"""
|
||||
|
||||
host = "localhost"
|
||||
ProtocolServerProcess = DaphneProcess
|
||||
static_wrapper = StaticFilesWrapper
|
||||
serve_static = True
|
||||
|
||||
@property
|
||||
def live_server_url(self):
|
||||
return "http://%s:%s" % (self.host, self._port)
|
||||
|
||||
@property
|
||||
def live_server_ws_url(self):
|
||||
return "ws://%s:%s" % (self.host, self._port)
|
||||
|
||||
def _pre_setup(self):
|
||||
for connection in connections.all():
|
||||
if self._is_in_memory_db(connection):
|
||||
raise ImproperlyConfigured(
|
||||
"ChannelLiveServerTestCase can not be used with in memory databases"
|
||||
)
|
||||
|
||||
super(ChannelsLiveServerTestCase, self)._pre_setup()
|
||||
|
||||
self._live_server_modified_settings = modify_settings(
|
||||
ALLOWED_HOSTS={"append": self.host}
|
||||
)
|
||||
self._live_server_modified_settings.enable()
|
||||
|
||||
if self.serve_static:
|
||||
application = self.static_wrapper(get_default_application())
|
||||
else:
|
||||
application = get_default_application()
|
||||
|
||||
self._server_process = self.ProtocolServerProcess(self.host, application)
|
||||
self._server_process.start()
|
||||
self._server_process.ready.wait()
|
||||
self._port = self._server_process.port.value
|
||||
|
||||
def _post_teardown(self):
|
||||
self._server_process.terminate()
|
||||
self._server_process.join()
|
||||
self._live_server_modified_settings.disable()
|
||||
super(ChannelsLiveServerTestCase, self)._post_teardown()
|
||||
|
||||
def _is_in_memory_db(self, connection):
|
||||
"""
|
||||
Check if DatabaseWrapper holds in memory database.
|
||||
"""
|
||||
if connection.vendor == "sqlite":
|
||||
return connection.is_in_memory_db()
|
||||
@@ -0,0 +1,102 @@
|
||||
import json
|
||||
from urllib.parse import unquote, urlparse
|
||||
|
||||
from asgiref.testing import ApplicationCommunicator
|
||||
|
||||
|
||||
class WebsocketCommunicator(ApplicationCommunicator):
|
||||
"""
|
||||
ApplicationCommunicator subclass that has WebSocket shortcut methods.
|
||||
|
||||
It will construct the scope for you, so you need to pass the application
|
||||
(uninstantiated) along with the initial connection parameters.
|
||||
"""
|
||||
|
||||
def __init__(self, application, path, headers=None, subprotocols=None):
|
||||
if not isinstance(path, str):
|
||||
raise TypeError("Expected str, got {}".format(type(path)))
|
||||
parsed = urlparse(path)
|
||||
self.scope = {
|
||||
"type": "websocket",
|
||||
"path": unquote(parsed.path),
|
||||
"query_string": parsed.query.encode("utf-8"),
|
||||
"headers": headers or [],
|
||||
"subprotocols": subprotocols or [],
|
||||
}
|
||||
super().__init__(application, self.scope)
|
||||
|
||||
async def connect(self, timeout=1):
|
||||
"""
|
||||
Trigger the connection code.
|
||||
|
||||
On an accepted connection, returns (True, <chosen-subprotocol>)
|
||||
On a rejected connection, returns (False, <close-code>)
|
||||
"""
|
||||
await self.send_input({"type": "websocket.connect"})
|
||||
response = await self.receive_output(timeout)
|
||||
if response["type"] == "websocket.close":
|
||||
return (False, response.get("code", 1000))
|
||||
else:
|
||||
return (True, response.get("subprotocol", None))
|
||||
|
||||
async def send_to(self, text_data=None, bytes_data=None):
|
||||
"""
|
||||
Sends a WebSocket frame to the application.
|
||||
"""
|
||||
# Make sure we have exactly one of the arguments
|
||||
assert bool(text_data) != bool(
|
||||
bytes_data
|
||||
), "You must supply exactly one of text_data or bytes_data"
|
||||
# Send the right kind of event
|
||||
if text_data:
|
||||
assert isinstance(text_data, str), "The text_data argument must be a str"
|
||||
await self.send_input({"type": "websocket.receive", "text": text_data})
|
||||
else:
|
||||
assert isinstance(
|
||||
bytes_data, bytes
|
||||
), "The bytes_data argument must be bytes"
|
||||
await self.send_input({"type": "websocket.receive", "bytes": bytes_data})
|
||||
|
||||
async def send_json_to(self, data):
|
||||
"""
|
||||
Sends JSON data as a text frame
|
||||
"""
|
||||
await self.send_to(text_data=json.dumps(data))
|
||||
|
||||
async def receive_from(self, timeout=1):
|
||||
"""
|
||||
Receives a data frame from the view. Will fail if the connection
|
||||
closes instead. Returns either a bytestring or a unicode string
|
||||
depending on what sort of frame you got.
|
||||
"""
|
||||
response = await self.receive_output(timeout)
|
||||
# Make sure this is a send message
|
||||
assert response["type"] == "websocket.send"
|
||||
# Make sure there's exactly one key in the response
|
||||
assert ("text" in response) != (
|
||||
"bytes" in response
|
||||
), "The response needs exactly one of 'text' or 'bytes'"
|
||||
# Pull out the right key and typecheck it for our users
|
||||
if "text" in response:
|
||||
assert isinstance(response["text"], str), "Text frame payload is not str"
|
||||
return response["text"]
|
||||
else:
|
||||
assert isinstance(
|
||||
response["bytes"], bytes
|
||||
), "Binary frame payload is not bytes"
|
||||
return response["bytes"]
|
||||
|
||||
async def receive_json_from(self, timeout=1):
|
||||
"""
|
||||
Receives a JSON text frame payload and decodes it
|
||||
"""
|
||||
payload = await self.receive_from(timeout)
|
||||
assert isinstance(payload, str), "JSON data is not a text frame"
|
||||
return json.loads(payload)
|
||||
|
||||
async def disconnect(self, code=1000, timeout=1):
|
||||
"""
|
||||
Closes the socket
|
||||
"""
|
||||
await self.send_input({"type": "websocket.disconnect", "code": code})
|
||||
await self.wait(timeout)
|
||||
@@ -0,0 +1,44 @@
|
||||
import asyncio
|
||||
|
||||
from asgiref.server import StatelessServer
|
||||
|
||||
|
||||
class Worker(StatelessServer):
|
||||
"""
|
||||
ASGI protocol server that surfaces events sent to specific channels
|
||||
on the channel layer into a single application instance.
|
||||
"""
|
||||
|
||||
def __init__(self, application, channels, channel_layer, max_applications=1000):
|
||||
super().__init__(application, max_applications)
|
||||
self.channels = channels
|
||||
self.channel_layer = channel_layer
|
||||
if self.channel_layer is None:
|
||||
raise ValueError("Channel layer is not valid")
|
||||
|
||||
async def handle(self):
|
||||
"""
|
||||
Listens on all the provided channels and handles the messages.
|
||||
"""
|
||||
# For each channel, launch its own listening coroutine
|
||||
listeners = []
|
||||
for channel in self.channels:
|
||||
listeners.append(asyncio.ensure_future(self.listener(channel)))
|
||||
# Wait for them all to exit
|
||||
await asyncio.wait(listeners)
|
||||
# See if any of the listeners had an error (e.g. channel layer error)
|
||||
[listener.result() for listener in listeners]
|
||||
|
||||
async def listener(self, channel):
|
||||
"""
|
||||
Single-channel listener
|
||||
"""
|
||||
while True:
|
||||
message = await self.channel_layer.receive(channel)
|
||||
if not message.get("type", None):
|
||||
raise ValueError("Worker received message with no type.")
|
||||
# Make a scope and get an application instance for it
|
||||
scope = {"type": "channel", "channel": channel}
|
||||
instance_queue = self.get_or_create_application_instance(channel, scope)
|
||||
# Run the message into the app
|
||||
await instance_queue.put(message)
|
||||
Reference in New Issue
Block a user