r47309 - something approaching a test

hawkowl-TA+aISz0psMTMxyoc4vAAJOcrHinNvQL0E9HWUfgJXw@public.gmane.org Sat, 23 Apr 2016 09:29:19 -0600 (MDT)
Newsgroups gmane.comp.python.twisted.commits
Message-ID <[email protected]>
Author: hawkowl
Date: Sat Apr 23 09:29:15 2016
New Revision: 47309

Modified:
   branches/proper-upgrade-8301-2/twisted/web/http.py
   branches/proper-upgrade-8301-2/twisted/web/server.py
   branches/proper-upgrade-8301-2/twisted/web/test/test_http.py

Log:
something approaching a test

Modified: branches/proper-upgrade-8301-2/twisted/web/http.py
==============================================================================
--- branches/proper-upgrade-8301-2/twisted/web/http.py	(original)
+++ branches/proper-upgrade-8301-2/twisted/web/http.py	Sat Apr 23 09:29:15 2016
@@ -89,6 +89,7 @@
 from zope.interface import implementer, provider
 
 # twisted imports
+from twisted import copyright
 from twisted.python.compat import (
     _PY3, unicode, intToBytes, networkString, nativeString)
 from twisted.python.deprecate import deprecated
@@ -132,6 +133,8 @@
 else:
     _intTypes = (int, long)
 
+_version = networkString("TwistedWeb/%s" % (copyright.version,))
+
 protocol_version = "HTTP/1.1"
 
 CACHED = """Magic constant returned by http.Request methods to set cache
@@ -1987,6 +1990,22 @@
 
 
 
+def _respondToUpgrade(transport, newProtocol, headers):
+
+    transport.write(b"HTTP/1.1 101 Switching Protocols\r\n")
+
+    transport.write(b"Server: " + _version + b"\r\n")
+    transport.write(b"Upgrade: " + newProtocol + b"\r\n")
+    transport.write(b"Connection: Upgrade\r\n")
+
+    for k, v in headers.items():
+
+        transport.write(k + b": " + v + b"\r\n")
+
+    transport.write(b"\r\n")
+
+
+
 class _GenericHTTPChannelProtocol(proxyForInterface(IProtocol, "_channel")):
     """
     A proxy object that wraps one of the HTTP protocol objects, and switches
@@ -2063,7 +2082,7 @@
         """
         Look at the headers and determine if they want us to upgrade.
         """
-        content = b"".join(self._buffer).split(b"\r\n\r\n", 1)[0].split("\r\n")
+        content = b"".join(self._buffer).split(b"\r\n\r\n", 1)[0].split(b"\r\n")
         headers = {}
 
         for line in content[1:]:
@@ -2081,11 +2100,16 @@
 
             for upgrade in headers[b"upgrade"].split(b", "):
 
-                upgrader = self.site.upgradeables.get(upgrade.lower())
+                upgrader = self.factory.upgradeables.get(upgrade.lower())
 
                 if upgrader:
                     try:
-                        self._channel, self._replay = upgrader(self, headers)
+                        res = upgrader(self, headers)
+                        transport = self._channel.transport
+                        self._channel, self._replay, headersToSend = res
+                        _respondToUpgrade(transport, upgrade, headersToSend)
+                        self._channel.makeConnection(transport)
+                        return upgrade
                     except CannotUpgrade:
                         pass
 
@@ -2197,6 +2221,7 @@
         if logFormatter is None:
             logFormatter = combinedLogFormatter
         self._logFormatter = logFormatter
+        self.upgradeables = {}
 
         # For storing the cached log datetime and the callback to update it
         self._logDateTime = None

Modified: branches/proper-upgrade-8301-2/twisted/web/server.py
==============================================================================
--- branches/proper-upgrade-8301-2/twisted/web/server.py	(original)
+++ branches/proper-upgrade-8301-2/twisted/web/server.py	Sat Apr 23 09:29:15 2016
@@ -34,9 +34,8 @@
     from twisted.spread.pb import Copyable, ViewPoint
 from twisted.internet import address
 from twisted.web import iweb, http, util
-from twisted.web.http import unquote
+from twisted.web.http import unquote, _version as version
 from twisted.python import log, reflect, failure, components
-from twisted import copyright
 from twisted.web import resource
 from twisted.web.error import UnsupportedMethod, CannotUpgrade
 
@@ -612,17 +611,6 @@
             self._expireCall.reset(self.sessionTimeout)
 
 
-version = networkString("TwistedWeb/%s" % (copyright.version,))
-
-
-
-def H2C(self, headers):
-    print("Negotiating H2C!")
-
-    raise CannotUpgrade()
-
-    return None, True
-
 
 class Site(http.HTTPFactory):
     """
@@ -656,7 +644,6 @@
         http.HTTPFactory.__init__(self, *args, **kwargs)
         self.sessions = {}
         self.resource = resource
-        self.upgradeables = {}
         if requestFactory is not None:
             self.requestFactory = requestFactory
 

Modified: branches/proper-upgrade-8301-2/twisted/web/test/test_http.py
==============================================================================
--- branches/proper-upgrade-8301-2/twisted/web/test/test_http.py	(original)
+++ branches/proper-upgrade-8301-2/twisted/web/test/test_http.py	Sat Apr 23 09:29:15 2016
@@ -8,6 +8,7 @@
 from __future__ import absolute_import, division
 
 import random, cgi, base64
+import math
 
 try:
     from urlparse import urlparse, urlunsplit, clear_cache
@@ -20,8 +21,9 @@
 from twisted.trial import unittest
 from twisted.trial.unittest import TestCase
 from twisted.web import http, http_headers
-from twisted.web.http import PotentialDataLoss, _DataLoss
+from twisted.web.http import PotentialDataLoss, _DataLoss, _version
 from twisted.web.http import _IdentityTransferDecoder
+from twisted.internet.protocol import Protocol, Factory
 from twisted.internet.task import Clock
 from twisted.internet.error import ConnectionLost
 from twisted.protocols import loopback
@@ -2533,3 +2535,72 @@
                     "in Twisted 15.0.0; please use Twisted Names to "
                     "resolve hostnames instead")},
                          sub(["category", "message"], warnings[0]))
+
+
+class HTTPUpgradeTests(unittest.TestCase):
+    """
+    Tests for HTTP/1.1 protocol upgrade.
+    """
+
+    def test_basic(self):
+
+        piTimes = 10
+
+        class Pitocol(Protocol):
+            """
+            A protocol that writes pi to the transport.
+            """
+            def dataReceived(protoself, data):
+                """
+                A C{dataReceived} that expects "GO" and will then write out
+                "3.14" * C{piTimes}. If there's any other data, that won't
+                """
+                if not protoself.connected:
+                    self.fail("dataReceived called when disconnected!")
+                if data == b"GO":
+                    for i in range(piTimes):
+                        protoself.transport.write(b"3.14")
+                protoself.transport.loseConnection()
+
+        piFactory = Factory()
+        piFactory.protocol = Pitocol
+
+        def piNegotiate(channel, headers):
+            pi = piFactory.buildProtocol(None)
+            return pi, False, {}
+
+
+        from twisted.web.http import _respondToUpgrade, HTTPFactory
+
+        factory = HTTPFactory()
+        factory._logDateTime = "sometime"
+        factory._logDateTimeCall = True
+        factory.startFactory()
+
+        factory.upgradeables[b"pitocol"] = piNegotiate
+
+        protocol = factory.buildProtocol(None)
+
+        trans = StringTransport()
+        protocol.makeConnection(trans)
+
+        val = [
+            b"GET / HTTP/1.1\r\n"
+            b"Connection: keep-alive, Upgrade\r\n",
+            b"Upgrade: pitocol\r\n\r\n",
+        ]
+
+        for x in val:
+            protocol.dataReceived(x)
+
+        expectedValue = b"".join([
+            b"HTTP/1.1 101 Switching Protocols\r\nServer: ",
+            _version, b"\r\nUpgrade: pitocol\r\nConnection: Upgrade\r\n\r\n"])
+
+        self.assertEqual(trans.value(), expectedValue)
+        trans.clear()
+
+        protocol.dataReceived(b"GO")
+
+        self.assertEqual(trans.value(), b"3.14" * 10)
+        self.assertTrue(trans.disconnecting)