This commit is contained in:
Mai Gillmann
2019-12-17 14:09:10 +01:00
parent 66e908fc8a
commit 4791d00a43
2122 changed files with 423791 additions and 0 deletions
@@ -0,0 +1,33 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
from twisted.python import components
from zope.interface import implementer, Interface
def foo():
return 2
class X:
def __init__(self, x):
self.x = x
def do(self):
#print 'X',self.x,'doing!'
pass
class XComponent(components.Componentized):
pass
class IX(Interface):
pass
@implementer(IX)
class XA(components.Adapter):
def method(self):
# Kick start :(
pass
components.registerAdapter(XA, X, IX)
@@ -0,0 +1,557 @@
# -*- test-case-name: twisted.test.test_amp,twisted.test.test_iosim -*-
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Utilities and helpers for simulating a network
"""
from __future__ import absolute_import, division, print_function
import itertools
try:
from OpenSSL.SSL import Error as NativeOpenSSLError
except ImportError:
pass
from zope.interface import implementer, directlyProvides
from twisted.internet.endpoints import TCP4ClientEndpoint, TCP4ServerEndpoint
from twisted.internet.protocol import Factory, Protocol
from twisted.internet.error import ConnectionRefusedError
from twisted.python.failure import Failure
from twisted.internet import error
from twisted.internet import interfaces
from twisted.internet.testing import MemoryReactorClock
class TLSNegotiation:
def __init__(self, obj, connectState):
self.obj = obj
self.connectState = connectState
self.sent = False
self.readyToSend = connectState
def __repr__(self):
return 'TLSNegotiation(%r)' % (self.obj,)
def pretendToVerify(self, other, tpt):
# Set the transport problems list here? disconnections?
# hmmmmm... need some negative path tests.
if not self.obj.iosimVerify(other.obj):
tpt.disconnectReason = NativeOpenSSLError()
tpt.loseConnection()
@implementer(interfaces.IAddress)
class FakeAddress(object):
"""
The default address type for the host and peer of L{FakeTransport}
connections.
"""
@implementer(interfaces.ITransport,
interfaces.ITLSTransport)
class FakeTransport:
"""
A wrapper around a file-like object to make it behave as a Transport.
This doesn't actually stream the file to the attached protocol,
and is thus useful mainly as a utility for debugging protocols.
"""
_nextserial = staticmethod(lambda counter=itertools.count(): next(counter))
closed = 0
disconnecting = 0
disconnected = 0
disconnectReason = error.ConnectionDone("Connection done")
producer = None
streamingProducer = 0
tls = None
def __init__(self, protocol, isServer, hostAddress=None, peerAddress=None):
"""
@param protocol: This transport will deliver bytes to this protocol.
@type protocol: L{IProtocol} provider
@param isServer: C{True} if this is the accepting side of the
connection, C{False} if it is the connecting side.
@type isServer: L{bool}
@param hostAddress: The value to return from C{getHost}. L{None}
results in a new L{FakeAddress} being created to use as the value.
@type hostAddress: L{IAddress} provider or L{None}
@param peerAddress: The value to return from C{getPeer}. L{None}
results in a new L{FakeAddress} being created to use as the value.
@type peerAddress: L{IAddress} provider or L{None}
"""
self.protocol = protocol
self.isServer = isServer
self.stream = []
self.serial = self._nextserial()
if hostAddress is None:
hostAddress = FakeAddress()
self.hostAddress = hostAddress
if peerAddress is None:
peerAddress = FakeAddress()
self.peerAddress = peerAddress
def __repr__(self):
return 'FakeTransport<%s,%s,%s>' % (
self.isServer and 'S' or 'C', self.serial,
self.protocol.__class__.__name__)
def write(self, data):
# If transport is closed, we should accept writes but drop the data.
if self.disconnecting:
return
if self.tls is not None:
self.tlsbuf.append(data)
else:
self.stream.append(data)
def _checkProducer(self):
# Cheating; this is called at "idle" times to allow producers to be
# found and dealt with
if self.producer and not self.streamingProducer:
self.producer.resumeProducing()
def registerProducer(self, producer, streaming):
"""
From abstract.FileDescriptor
"""
self.producer = producer
self.streamingProducer = streaming
if not streaming:
producer.resumeProducing()
def unregisterProducer(self):
self.producer = None
def stopConsuming(self):
self.unregisterProducer()
self.loseConnection()
def writeSequence(self, iovec):
self.write(b"".join(iovec))
def loseConnection(self):
self.disconnecting = True
def abortConnection(self):
"""
For the time being, this is the same as loseConnection; no buffered
data will be lost.
"""
self.disconnecting = True
def reportDisconnect(self):
if self.tls is not None:
# We were in the middle of negotiating! Must have been a TLS
# problem.
err = NativeOpenSSLError()
else:
err = self.disconnectReason
self.protocol.connectionLost(Failure(err))
def logPrefix(self):
"""
Identify this transport/event source to the logging system.
"""
return "iosim"
def getPeer(self):
return self.peerAddress
def getHost(self):
return self.hostAddress
def resumeProducing(self):
# Never sends data anyways
pass
def pauseProducing(self):
# Never sends data anyways
pass
def stopProducing(self):
self.loseConnection()
def startTLS(self, contextFactory, beNormal=True):
# Nothing's using this feature yet, but startTLS has an undocumented
# second argument which defaults to true; if set to False, servers will
# behave like clients and clients will behave like servers.
connectState = self.isServer ^ beNormal
self.tls = TLSNegotiation(contextFactory, connectState)
self.tlsbuf = []
def getOutBuffer(self):
"""
Get the pending writes from this transport, clearing them from the
pending buffer.
@return: the bytes written with C{transport.write}
@rtype: L{bytes}
"""
S = self.stream
if S:
self.stream = []
return b''.join(S)
elif self.tls is not None:
if self.tls.readyToSend:
# Only _send_ the TLS negotiation "packet" if I'm ready to.
self.tls.sent = True
return self.tls
else:
return None
else:
return None
def bufferReceived(self, buf):
if isinstance(buf, TLSNegotiation):
assert self.tls is not None # By the time you're receiving a
# negotiation, you have to have called
# startTLS already.
if self.tls.sent:
self.tls.pretendToVerify(buf, self)
self.tls = None # We're done with the handshake if we've gotten
# this far... although maybe it failed...?
# TLS started! Unbuffer...
b, self.tlsbuf = self.tlsbuf, None
self.writeSequence(b)
directlyProvides(self, interfaces.ISSLTransport)
else:
# We haven't sent our own TLS negotiation: time to do that!
self.tls.readyToSend = True
else:
self.protocol.dataReceived(buf)
def makeFakeClient(clientProtocol):
"""
Create and return a new in-memory transport hooked up to the given protocol.
@param clientProtocol: The client protocol to use.
@type clientProtocol: L{IProtocol} provider
@return: The transport.
@rtype: L{FakeTransport}
"""
return FakeTransport(clientProtocol, isServer=False)
def makeFakeServer(serverProtocol):
"""
Create and return a new in-memory transport hooked up to the given protocol.
@param serverProtocol: The server protocol to use.
@type serverProtocol: L{IProtocol} provider
@return: The transport.
@rtype: L{FakeTransport}
"""
return FakeTransport(serverProtocol, isServer=True)
class IOPump:
"""
Utility to pump data between clients and servers for protocol testing.
Perhaps this is a utility worthy of being in protocol.py?
"""
def __init__(self, client, server, clientIO, serverIO, debug):
self.client = client
self.server = server
self.clientIO = clientIO
self.serverIO = serverIO
self.debug = debug
def flush(self, debug=False):
"""
Pump until there is no more input or output.
Returns whether any data was moved.
"""
result = False
for x in range(1000):
if self.pump(debug):
result = True
else:
break
else:
assert 0, "Too long"
return result
def pump(self, debug=False):
"""
Move data back and forth.
Returns whether any data was moved.
"""
if self.debug or debug:
print('-- GLUG --')
sData = self.serverIO.getOutBuffer()
cData = self.clientIO.getOutBuffer()
self.clientIO._checkProducer()
self.serverIO._checkProducer()
if self.debug or debug:
print('.')
# XXX slightly buggy in the face of incremental output
if cData:
print('C: ' + repr(cData))
if sData:
print('S: ' + repr(sData))
if cData:
self.serverIO.bufferReceived(cData)
if sData:
self.clientIO.bufferReceived(sData)
if cData or sData:
return True
if (self.serverIO.disconnecting and
not self.serverIO.disconnected):
if self.debug or debug:
print('* C')
self.serverIO.disconnected = True
self.clientIO.disconnecting = True
self.clientIO.reportDisconnect()
return True
if self.clientIO.disconnecting and not self.clientIO.disconnected:
if self.debug or debug:
print('* S')
self.clientIO.disconnected = True
self.serverIO.disconnecting = True
self.serverIO.reportDisconnect()
return True
return False
def connect(serverProtocol, serverTransport, clientProtocol, clientTransport,
debug=False, greet=True):
"""
Create a new L{IOPump} connecting two protocols.
@param serverProtocol: The protocol to use on the accepting side of the
connection.
@type serverProtocol: L{IProtocol} provider
@param serverTransport: The transport to associate with C{serverProtocol}.
@type serverTransport: L{FakeTransport}
@param clientProtocol: The protocol to use on the initiating side of the
connection.
@type clientProtocol: L{IProtocol} provider
@param clientTransport: The transport to associate with C{clientProtocol}.
@type clientTransport: L{FakeTransport}
@param debug: A flag indicating whether to log information about what the
L{IOPump} is doing.
@type debug: L{bool}
@param greet: Should the L{IOPump} be L{flushed <IOPump.flush>} once before
returning to put the protocols into their post-handshake or
post-server-greeting state?
@type greet: L{bool}
@return: An L{IOPump} which connects C{serverProtocol} and
C{clientProtocol} and delivers bytes between them when it is pumped.
@rtype: L{IOPump}
"""
serverProtocol.makeConnection(serverTransport)
clientProtocol.makeConnection(clientTransport)
pump = IOPump(
clientProtocol, serverProtocol, clientTransport, serverTransport, debug
)
if greet:
# Kick off server greeting, etc
pump.flush()
return pump
def connectedServerAndClient(ServerClass, ClientClass,
clientTransportFactory=makeFakeClient,
serverTransportFactory=makeFakeServer,
debug=False, greet=True):
"""
Connect a given server and client class to each other.
@param ServerClass: a callable that produces the server-side protocol.
@type ServerClass: 0-argument callable returning L{IProtocol} provider.
@param ClientClass: like C{ServerClass} but for the other side of the
connection.
@type ClientClass: 0-argument callable returning L{IProtocol} provider.
@param clientTransportFactory: a callable that produces the transport which
will be attached to the protocol returned from C{ClientClass}.
@type clientTransportFactory: callable taking (L{IProtocol}) and returning
L{FakeTransport}
@param serverTransportFactory: a callable that produces the transport which
will be attached to the protocol returned from C{ServerClass}.
@type serverTransportFactory: callable taking (L{IProtocol}) and returning
L{FakeTransport}
@param debug: Should this dump an escaped version of all traffic on this
connection to stdout for inspection?
@type debug: L{bool}
@param greet: Should the L{IOPump} be L{flushed <IOPump.flush>} once before
returning to put the protocols into their post-handshake or
post-server-greeting state?
@type greet: L{bool}
@return: the client protocol, the server protocol, and an L{IOPump} which,
when its C{pump} and C{flush} methods are called, will move data
between the created client and server protocol instances.
@rtype: 3-L{tuple} of L{IProtocol}, L{IProtocol}, L{IOPump}
"""
c = ClientClass()
s = ServerClass()
cio = clientTransportFactory(c)
sio = serverTransportFactory(s)
return c, s, connect(s, sio, c, cio, debug, greet)
def _factoriesShouldConnect(clientInfo, serverInfo):
"""
Should the client and server described by the arguments be connected to
each other, i.e. do their port numbers match?
@param clientInfo: the args for connectTCP
@type clientInfo: L{tuple}
@param serverInfo: the args for listenTCP
@type serverInfo: L{tuple}
@return: If they do match, return factories for the client and server that
should connect; otherwise return L{None}, indicating they shouldn't be
connected.
@rtype: L{None} or 2-L{tuple} of (L{ClientFactory},
L{IProtocolFactory})
"""
(clientHost, clientPort, clientFactory, clientTimeout,
clientBindAddress) = clientInfo
(serverPort, serverFactory, serverBacklog,
serverInterface) = serverInfo
if serverPort == clientPort:
return clientFactory, serverFactory
else:
return None
class ConnectionCompleter(object):
"""
A L{ConnectionCompleter} can cause synthetic TCP connections established by
L{MemoryReactor.connectTCP} and L{MemoryReactor.listenTCP} to succeed or
fail.
"""
def __init__(self, memoryReactor):
"""
Create a L{ConnectionCompleter} from a L{MemoryReactor}.
@param memoryReactor: The reactor to attach to.
@type memoryReactor: L{MemoryReactor}
"""
self._reactor = memoryReactor
def succeedOnce(self, debug=False):
"""
Complete a single TCP connection established on this
L{ConnectionCompleter}'s L{MemoryReactor}.
@param debug: A flag; whether to dump output from the established
connection to stdout.
@type debug: L{bool}
@return: a pump for the connection, or L{None} if no connection could
be established.
@rtype: L{IOPump} or L{None}
"""
memoryReactor = self._reactor
for clientIdx, clientInfo in enumerate(memoryReactor.tcpClients):
for serverInfo in memoryReactor.tcpServers:
factories = _factoriesShouldConnect(clientInfo, serverInfo)
if factories:
memoryReactor.tcpClients.remove(clientInfo)
memoryReactor.connectors.pop(clientIdx)
clientFactory, serverFactory = factories
clientProtocol = clientFactory.buildProtocol(None)
serverProtocol = serverFactory.buildProtocol(None)
serverTransport = makeFakeServer(serverProtocol)
clientTransport = makeFakeClient(clientProtocol)
return connect(serverProtocol, serverTransport,
clientProtocol, clientTransport,
debug)
def failOnce(self, reason=Failure(ConnectionRefusedError())):
"""
Fail a single TCP connection established on this
L{ConnectionCompleter}'s L{MemoryReactor}.
@param reason: the reason to provide that the connection failed.
@type reason: L{Failure}
"""
self._reactor.tcpClients.pop(0)[2].clientConnectionFailed(
self._reactor.connectors.pop(0), reason
)
def connectableEndpoint(debug=False):
"""
Create an endpoint that can be fired on demand.
@param debug: A flag; whether to dump output from the established
connection to stdout.
@type debug: L{bool}
@return: A client endpoint, and an object that will cause one of the
L{Deferred}s returned by that client endpoint.
@rtype: 2-L{tuple} of (L{IStreamClientEndpoint}, L{ConnectionCompleter})
"""
reactor = MemoryReactorClock()
clientEndpoint = TCP4ClientEndpoint(reactor, "0.0.0.0", 4321)
serverEndpoint = TCP4ServerEndpoint(reactor, 4321)
serverEndpoint.listen(Factory.forProtocol(Protocol))
return clientEndpoint, ConnectionCompleter(reactor)
@@ -0,0 +1,12 @@
class A:
def a(self):
return 'b'
class B(A, object):
def b(self):
return 'c'
class Inherit(A):
def a(self):
return 'd'
@@ -0,0 +1,31 @@
# Copyright (c) 2005 Divmod, Inc.
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Test plugin used in L{twisted.test.test_plugin}.
"""
from zope.interface import provider
from twisted.plugin import IPlugin
from twisted.test.test_plugin import ITestPlugin
@provider(ITestPlugin, IPlugin)
class FourthTestPlugin:
def test1():
pass
test1 = staticmethod(test1)
@provider(ITestPlugin, IPlugin)
class FifthTestPlugin:
"""
More documentation: I hate you.
"""
def test1():
pass
test1 = staticmethod(test1)
@@ -0,0 +1,11 @@
"""
Write to stdout the command line args it received, one per line.
"""
from __future__ import print_function
import sys
for x in sys.argv[1:]:
print(x)
@@ -0,0 +1,11 @@
"""Write back all data it receives."""
import sys
data = sys.stdin.read(1)
while data:
sys.stdout.write(data)
sys.stdout.flush()
data = sys.stdin.read(1)
sys.stderr.write("byebye")
sys.stderr.flush()
@@ -0,0 +1,43 @@
"""Write to a handful of file descriptors, to test the childFDs= argument of
reactor.spawnProcess()
"""
from __future__ import print_function
import os, sys
if __name__ == "__main__":
debug = 0
if debug: stderr = os.fdopen(2, "w")
if debug: print("this is stderr", file=stderr)
abcd = os.read(0, 4)
if debug: print("read(0):", abcd, file=stderr)
if abcd != b"abcd":
sys.exit(1)
if debug: print("os.write(1, righto)", file=stderr)
os.write(1, b"righto")
efgh = os.read(3, 4)
if debug: print("read(3):", file=stderr)
if efgh != b"efgh":
sys.exit(2)
if debug: print("os.close(4)", file=stderr)
os.close(4)
eof = os.read(5, 4)
if debug: print("read(5):", eof, file=stderr)
if eof != b"":
sys.exit(3)
if debug: print("os.write(1, closed)", file=stderr)
os.write(1, b"closed")
if debug: print("sys.exit(0)", file=stderr)
sys.exit(0)
@@ -0,0 +1,53 @@
"""Test program for processes."""
import sys, os
# Twisted is unimportable from this file, so just do the PY3 check manually
if sys.version_info < (3, 0):
_PY3 = False
else:
_PY3 = True
test_file_match = "process_test.log.*"
test_file = "process_test.log.%d" % os.getpid()
def main():
f = open(test_file, 'wb')
if _PY3:
stdin = sys.stdin.buffer
stderr = sys.stderr.buffer
stdout = sys.stdout.buffer
else:
stdin = sys.stdin
stdout = sys.stdout
stderr = sys.stderr
# stage 1
b = stdin.read(4)
f.write(b"one: " + b + b"\n")
# stage 2
stdout.write(b)
stdout.flush()
os.close(sys.stdout.fileno())
# and a one, and a two, and a...
b = stdin.read(4)
f.write(b"two: " + b + b"\n")
# stage 3
stderr.write(b)
stderr.flush()
os.close(stderr.fileno())
# stage 4
b = stdin.read(4)
f.write(b"three: " + b + b"\n")
# exit with status code 23
sys.exit(23)
if __name__ == '__main__':
main()
@@ -0,0 +1,45 @@
"""A process that reads from stdin and out using Twisted."""
from __future__ import division, absolute_import, print_function
### Twisted Preamble
# This makes sure that users don't have to set up their environment
# specially in order to run these programs from bin/.
import sys, os
pos = os.path.abspath(sys.argv[0]).find(os.sep+'Twisted')
if pos != -1:
sys.path.insert(0, os.path.abspath(sys.argv[0])[:pos+8])
sys.path.insert(0, os.curdir)
### end of preamble
from twisted.python import log
from zope.interface import implementer
from twisted.internet import interfaces
log.startLogging(sys.stderr)
from twisted.internet import protocol, reactor, stdio
@implementer(interfaces.IHalfCloseableProtocol)
class Echo(protocol.Protocol):
def connectionMade(self):
print("connection made")
def dataReceived(self, data):
self.transport.write(data)
def readConnectionLost(self):
print("readConnectionLost")
self.transport.loseConnection()
def writeConnectionLost(self):
print("writeConnectionLost")
def connectionLost(self, reason):
print("connectionLost", reason)
reactor.stop()
stdio.StandardIO(Echo())
reactor.run()
@@ -0,0 +1,4 @@
# Helper module for a test_reflect test
1//0
@@ -0,0 +1,50 @@
# -*- test-case-name: twisted.test.test_stdio.StandardInputOutputTests.test_loseConnection -*-
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Main program for the child process run by
L{twisted.test.test_stdio.StandardInputOutputTests.test_loseConnection} to
test that ITransport.loseConnection() works for process transports.
"""
from __future__ import absolute_import, division
import sys
from twisted.internet.error import ConnectionDone
from twisted.internet import stdio, protocol
from twisted.python import reflect, log
class LoseConnChild(protocol.Protocol):
exitCode = 0
def connectionMade(self):
self.transport.loseConnection()
def connectionLost(self, reason):
"""
Check that C{reason} is a L{Failure} wrapping a L{ConnectionDone}
instance and stop the reactor. If C{reason} is wrong for some reason,
log something about that in C{self.errorLogFile} and make sure the
process exits with a non-zero status.
"""
try:
try:
reason.trap(ConnectionDone)
except:
log.err(None, "Problem with reason passed to connectionLost")
self.exitCode = 1
finally:
reactor.stop()
if __name__ == '__main__':
reflect.namedAny(sys.argv[1]).install()
log.startLogging(open(sys.argv[2], 'wb'))
from twisted.internet import reactor
protocol = LoseConnChild()
stdio.StandardIO(protocol)
reactor.run()
sys.exit(protocol.exitCode)
@@ -0,0 +1,59 @@
# -*- test-case-name: twisted.test.test_stdio.StandardInputOutputTests.test_producer -*-
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Main program for the child process run by
L{twisted.test.test_stdio.StandardInputOutputTests.test_producer} to test
that process transports implement IProducer properly.
"""
from __future__ import absolute_import, division
import sys
from twisted.internet import stdio, protocol
from twisted.python import log, reflect
class ProducerChild(protocol.Protocol):
_paused = False
buf = b''
def connectionLost(self, reason):
log.msg("*****OVER*****")
reactor.callLater(1, reactor.stop)
def dataReceived(self, data):
self.buf += data
if self._paused:
log.startLogging(sys.stderr)
log.msg("dataReceived while transport paused!")
self.transport.loseConnection()
else:
self.transport.write(data)
if self.buf.endswith(b'\n0\n'):
self.transport.loseConnection()
else:
self.pause()
def pause(self):
self._paused = True
self.transport.pauseProducing()
reactor.callLater(0.01, self.unpause)
def unpause(self):
self._paused = False
self.transport.resumeProducing()
if __name__ == '__main__':
reflect.namedAny(sys.argv[1]).install()
from twisted.internet import reactor
stdio.StandardIO(ProducerChild())
reactor.run()
@@ -0,0 +1,935 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Tests for L{twisted.python.compat}.
"""
from __future__ import division, absolute_import
import socket, sys, traceback, io, codecs
from twisted.trial import unittest
from twisted.python.compat import (
reduce, execfile, _PY3, _PYPY, comparable, cmp, nativeString,
networkString, unicode as unicodeCompat, lazyByteSlice, reraise,
NativeStringIO, iterbytes, intToBytes, ioType, bytesEnviron, iteritems,
_coercedUnicode, unichr, raw_input, _bytesRepr, _get_async_param,
)
from twisted.python.filepath import FilePath
from twisted.python.runtime import platform
class IOTypeTests(unittest.SynchronousTestCase):
"""
Test cases for determining a file-like object's type.
"""
def test_3StringIO(self):
"""
An L{io.StringIO} accepts and returns text.
"""
self.assertEqual(ioType(io.StringIO()), unicodeCompat)
def test_3BytesIO(self):
"""
An L{io.BytesIO} accepts and returns bytes.
"""
self.assertEqual(ioType(io.BytesIO()), bytes)
def test_3openTextMode(self):
"""
A file opened via 'io.open' in text mode accepts and returns text.
"""
with io.open(self.mktemp(), "w") as f:
self.assertEqual(ioType(f), unicodeCompat)
def test_3openBinaryMode(self):
"""
A file opened via 'io.open' in binary mode accepts and returns bytes.
"""
with io.open(self.mktemp(), "wb") as f:
self.assertEqual(ioType(f), bytes)
def test_2openTextMode(self):
"""
The special built-in console file in Python 2 which has an 'encoding'
attribute should qualify as a special type, since it accepts both bytes
and text faithfully.
"""
class VerySpecificLie(file):
"""
In their infinite wisdom, the CPython developers saw fit not to
allow us a writable 'encoding' attribute on the built-in 'file'
type in Python 2, despite making it writable in C with
PyFile_SetEncoding.
Pretend they did not do that.
"""
encoding = 'utf-8'
self.assertEqual(ioType(VerySpecificLie(self.mktemp(), "wb")),
basestring)
def test_2StringIO(self):
"""
Python 2's L{StringIO} and L{cStringIO} modules are both binary I/O.
"""
from cStringIO import StringIO as cStringIO
from StringIO import StringIO
self.assertEqual(ioType(StringIO()), bytes)
self.assertEqual(ioType(cStringIO()), bytes)
def test_2openBinaryMode(self):
"""
The normal 'open' builtin in Python 2 will always result in bytes I/O.
"""
with open(self.mktemp(), "w") as f:
self.assertEqual(ioType(f), bytes)
if _PY3:
test_2openTextMode.skip = "The 'file' type is no longer available."
test_2openBinaryMode.skip = "'io.open' is now the same as 'open'."
test_2StringIO.skip = ("The 'StringIO' and 'cStringIO' modules were "
"subsumed by the 'io' module.")
def test_codecsOpenBytes(self):
"""
The L{codecs} module, oddly, returns a file-like object which returns
bytes when not passed an 'encoding' argument.
"""
with codecs.open(self.mktemp(), 'wb') as f:
self.assertEqual(ioType(f), bytes)
def test_codecsOpenText(self):
"""
When passed an encoding, however, the L{codecs} module returns unicode.
"""
with codecs.open(self.mktemp(), 'wb', encoding='utf-8') as f:
self.assertEqual(ioType(f), unicodeCompat)
def test_defaultToText(self):
"""
When passed an object about which no sensible decision can be made, err
on the side of unicode.
"""
self.assertEqual(ioType(object()), unicodeCompat)
class CompatTests(unittest.SynchronousTestCase):
"""
Various utility functions in C{twisted.python.compat} provide same
functionality as modern Python variants.
"""
def test_set(self):
"""
L{set} should behave like the expected set interface.
"""
a = set()
a.add('b')
a.add('c')
a.add('a')
b = list(a)
b.sort()
self.assertEqual(b, ['a', 'b', 'c'])
a.remove('b')
b = list(a)
b.sort()
self.assertEqual(b, ['a', 'c'])
a.discard('d')
b = set(['r', 's'])
d = a.union(b)
b = list(d)
b.sort()
self.assertEqual(b, ['a', 'c', 'r', 's'])
def test_frozenset(self):
"""
L{frozenset} should behave like the expected frozenset interface.
"""
a = frozenset(['a', 'b'])
self.assertRaises(AttributeError, getattr, a, "add")
self.assertEqual(sorted(a), ['a', 'b'])
b = frozenset(['r', 's'])
d = a.union(b)
b = list(d)
b.sort()
self.assertEqual(b, ['a', 'b', 'r', 's'])
def test_reduce(self):
"""
L{reduce} should behave like the builtin reduce.
"""
self.assertEqual(15, reduce(lambda x, y: x + y, [1, 2, 3, 4, 5]))
self.assertEqual(16, reduce(lambda x, y: x + y, [1, 2, 3, 4, 5], 1))
class IPv6Tests(unittest.SynchronousTestCase):
"""
C{inet_pton} and C{inet_ntop} implementations support IPv6.
"""
def testNToP(self):
from twisted.python.compat import inet_ntop
f = lambda a: inet_ntop(socket.AF_INET6, a)
g = lambda a: inet_ntop(socket.AF_INET, a)
self.assertEqual('::', f('\x00' * 16))
self.assertEqual('::1', f('\x00' * 15 + '\x01'))
self.assertEqual(
'aef:b01:506:1001:ffff:9997:55:170',
f('\x0a\xef\x0b\x01\x05\x06\x10\x01\xff\xff\x99\x97\x00\x55\x01'
'\x70'))
self.assertEqual('1.0.1.0', g('\x01\x00\x01\x00'))
self.assertEqual('170.85.170.85', g('\xaa\x55\xaa\x55'))
self.assertEqual('255.255.255.255', g('\xff\xff\xff\xff'))
self.assertEqual('100::', f('\x01' + '\x00' * 15))
self.assertEqual('100::1', f('\x01' + '\x00' * 14 + '\x01'))
def testPToN(self):
"""
L{twisted.python.compat.inet_pton} parses IPv4 and IPv6 addresses in a
manner similar to that of L{socket.inet_pton}.
"""
from twisted.python.compat import inet_pton
f = lambda a: inet_pton(socket.AF_INET6, a)
g = lambda a: inet_pton(socket.AF_INET, a)
self.assertEqual('\x00\x00\x00\x00', g('0.0.0.0'))
self.assertEqual('\xff\x00\xff\x00', g('255.0.255.0'))
self.assertEqual('\xaa\xaa\xaa\xaa', g('170.170.170.170'))
self.assertEqual('\x00' * 16, f('::'))
self.assertEqual('\x00' * 16, f('0::0'))
self.assertEqual('\x00\x01' + '\x00' * 14, f('1::'))
self.assertEqual(
'\x45\xef\x76\xcb\x00\x1a\x56\xef\xaf\xeb\x0b\xac\x19\x24\xae\xae',
f('45ef:76cb:1a:56ef:afeb:bac:1924:aeae'))
# Scope ID doesn't affect the binary representation.
self.assertEqual(
'\x45\xef\x76\xcb\x00\x1a\x56\xef\xaf\xeb\x0b\xac\x19\x24\xae\xae',
f('45ef:76cb:1a:56ef:afeb:bac:1924:aeae%en0'))
self.assertEqual('\x00' * 14 + '\x00\x01', f('::1'))
self.assertEqual('\x00' * 12 + '\x01\x02\x03\x04', f('::1.2.3.4'))
self.assertEqual(
'\x00\x01\x00\x02\x00\x03\x00\x04\x00\x05\x00\x06\x01\x02\x03\xff',
f('1:2:3:4:5:6:1.2.3.255'))
for badaddr in ['1:2:3:4:5:6:7:8:', ':1:2:3:4:5:6:7:8', '1::2::3',
'1:::3', ':::', '1:2', '::1.2', '1.2.3.4::',
'abcd:1.2.3.4:abcd:abcd:abcd:abcd:abcd',
'1234:1.2.3.4:1234:1234:1234:1234:1234:1234',
'1.2.3.4', '', '%eth0']:
self.assertRaises(ValueError, f, badaddr)
if _PY3:
IPv6Tests.skip = "These tests are only relevant to old versions of Python"
class ExecfileCompatTests(unittest.SynchronousTestCase):
"""
Tests for the Python 3-friendly L{execfile} implementation.
"""
def writeScript(self, content):
"""
Write L{content} to a new temporary file, returning the L{FilePath}
for the new file.
"""
path = self.mktemp()
with open(path, "wb") as f:
f.write(content.encode("ascii"))
return FilePath(path.encode("utf-8"))
def test_execfileGlobals(self):
"""
L{execfile} executes the specified file in the given global namespace.
"""
script = self.writeScript(u"foo += 1\n")
globalNamespace = {"foo": 1}
execfile(script.path, globalNamespace)
self.assertEqual(2, globalNamespace["foo"])
def test_execfileGlobalsAndLocals(self):
"""
L{execfile} executes the specified file in the given global and local
namespaces.
"""
script = self.writeScript(u"foo += 1\n")
globalNamespace = {"foo": 10}
localNamespace = {"foo": 20}
execfile(script.path, globalNamespace, localNamespace)
self.assertEqual(10, globalNamespace["foo"])
self.assertEqual(21, localNamespace["foo"])
def test_execfileUniversalNewlines(self):
"""
L{execfile} reads in the specified file using universal newlines so
that scripts written on one platform will work on another.
"""
for lineEnding in u"\n", u"\r", u"\r\n":
script = self.writeScript(u"foo = 'okay'" + lineEnding)
globalNamespace = {"foo": None}
execfile(script.path, globalNamespace)
self.assertEqual("okay", globalNamespace["foo"])
class PY3Tests(unittest.SynchronousTestCase):
"""
Identification of Python 2 vs. Python 3.
"""
def test_python2(self):
"""
On Python 2, C{_PY3} is False.
"""
if sys.version.startswith("2."):
self.assertFalse(_PY3)
def test_python3(self):
"""
On Python 3, C{_PY3} is True.
"""
if sys.version.startswith("3."):
self.assertTrue(_PY3)
class PYPYTest(unittest.SynchronousTestCase):
"""
Identification of PyPy.
"""
def test_PYPY(self):
"""
On PyPy, L{_PYPY} is True.
"""
if 'PyPy' in sys.version:
self.assertTrue(_PYPY)
else:
self.assertFalse(_PYPY)
@comparable
class Comparable(object):
"""
Objects that can be compared to each other, but not others.
"""
def __init__(self, value):
self.value = value
def __cmp__(self, other):
if not isinstance(other, Comparable):
return NotImplemented
return cmp(self.value, other.value)
class ComparableTests(unittest.SynchronousTestCase):
"""
L{comparable} decorated classes emulate Python 2's C{__cmp__} semantics.
"""
def test_equality(self):
"""
Instances of a class that is decorated by C{comparable} support
equality comparisons.
"""
# Make explicitly sure we're using ==:
self.assertTrue(Comparable(1) == Comparable(1))
self.assertFalse(Comparable(2) == Comparable(1))
def test_nonEquality(self):
"""
Instances of a class that is decorated by C{comparable} support
inequality comparisons.
"""
# Make explicitly sure we're using !=:
self.assertFalse(Comparable(1) != Comparable(1))
self.assertTrue(Comparable(2) != Comparable(1))
def test_greaterThan(self):
"""
Instances of a class that is decorated by C{comparable} support
greater-than comparisons.
"""
self.assertTrue(Comparable(2) > Comparable(1))
self.assertFalse(Comparable(0) > Comparable(3))
def test_greaterThanOrEqual(self):
"""
Instances of a class that is decorated by C{comparable} support
greater-than-or-equal comparisons.
"""
self.assertTrue(Comparable(1) >= Comparable(1))
self.assertTrue(Comparable(2) >= Comparable(1))
self.assertFalse(Comparable(0) >= Comparable(3))
def test_lessThan(self):
"""
Instances of a class that is decorated by C{comparable} support
less-than comparisons.
"""
self.assertTrue(Comparable(0) < Comparable(3))
self.assertFalse(Comparable(2) < Comparable(0))
def test_lessThanOrEqual(self):
"""
Instances of a class that is decorated by C{comparable} support
less-than-or-equal comparisons.
"""
self.assertTrue(Comparable(3) <= Comparable(3))
self.assertTrue(Comparable(0) <= Comparable(3))
self.assertFalse(Comparable(2) <= Comparable(0))
class Python3ComparableTests(unittest.SynchronousTestCase):
"""
Python 3-specific functionality of C{comparable}.
"""
def test_notImplementedEquals(self):
"""
Instances of a class that is decorated by C{comparable} support
returning C{NotImplemented} from C{__eq__} if it is returned by the
underlying C{__cmp__} call.
"""
self.assertEqual(Comparable(1).__eq__(object()), NotImplemented)
def test_notImplementedNotEquals(self):
"""
Instances of a class that is decorated by C{comparable} support
returning C{NotImplemented} from C{__ne__} if it is returned by the
underlying C{__cmp__} call.
"""
self.assertEqual(Comparable(1).__ne__(object()), NotImplemented)
def test_notImplementedGreaterThan(self):
"""
Instances of a class that is decorated by C{comparable} support
returning C{NotImplemented} from C{__gt__} if it is returned by the
underlying C{__cmp__} call.
"""
self.assertEqual(Comparable(1).__gt__(object()), NotImplemented)
def test_notImplementedLessThan(self):
"""
Instances of a class that is decorated by C{comparable} support
returning C{NotImplemented} from C{__lt__} if it is returned by the
underlying C{__cmp__} call.
"""
self.assertEqual(Comparable(1).__lt__(object()), NotImplemented)
def test_notImplementedGreaterThanEquals(self):
"""
Instances of a class that is decorated by C{comparable} support
returning C{NotImplemented} from C{__ge__} if it is returned by the
underlying C{__cmp__} call.
"""
self.assertEqual(Comparable(1).__ge__(object()), NotImplemented)
def test_notImplementedLessThanEquals(self):
"""
Instances of a class that is decorated by C{comparable} support
returning C{NotImplemented} from C{__le__} if it is returned by the
underlying C{__cmp__} call.
"""
self.assertEqual(Comparable(1).__le__(object()), NotImplemented)
if not _PY3:
# On Python 2, we just use __cmp__ directly, so checking detailed
# comparison methods doesn't makes sense.
Python3ComparableTests.skip = "Python 3 only."
class CmpTests(unittest.SynchronousTestCase):
"""
L{cmp} should behave like the built-in Python 2 C{cmp}.
"""
def test_equals(self):
"""
L{cmp} returns 0 for equal objects.
"""
self.assertEqual(cmp(u"a", u"a"), 0)
self.assertEqual(cmp(1, 1), 0)
self.assertEqual(cmp([1], [1]), 0)
def test_greaterThan(self):
"""
L{cmp} returns 1 if its first argument is bigger than its second.
"""
self.assertEqual(cmp(4, 0), 1)
self.assertEqual(cmp(b"z", b"a"), 1)
def test_lessThan(self):
"""
L{cmp} returns -1 if its first argument is smaller than its second.
"""
self.assertEqual(cmp(0.1, 2.3), -1)
self.assertEqual(cmp(b"a", b"d"), -1)
class StringTests(unittest.SynchronousTestCase):
"""
Compatibility functions and types for strings.
"""
def assertNativeString(self, original, expected):
"""
Raise an exception indicating a failed test if the output of
C{nativeString(original)} is unequal to the expected string, or is not
a native string.
"""
self.assertEqual(nativeString(original), expected)
self.assertIsInstance(nativeString(original), str)
def test_nonASCIIBytesToString(self):
"""
C{nativeString} raises a C{UnicodeError} if input bytes are not ASCII
decodable.
"""
self.assertRaises(UnicodeError, nativeString, b"\xFF")
def test_nonASCIIUnicodeToString(self):
"""
C{nativeString} raises a C{UnicodeError} if input Unicode is not ASCII
encodable.
"""
self.assertRaises(UnicodeError, nativeString, u"\u1234")
def test_bytesToString(self):
"""
C{nativeString} converts bytes to the native string format, assuming
an ASCII encoding if applicable.
"""
self.assertNativeString(b"hello", "hello")
def test_unicodeToString(self):
"""
C{nativeString} converts unicode to the native string format, assuming
an ASCII encoding if applicable.
"""
self.assertNativeString(u"Good day", "Good day")
def test_stringToString(self):
"""
C{nativeString} leaves native strings as native strings.
"""
self.assertNativeString("Hello!", "Hello!")
def test_unexpectedType(self):
"""
C{nativeString} raises a C{TypeError} if given an object that is not a
string of some sort.
"""
self.assertRaises(TypeError, nativeString, 1)
def test_unicode(self):
"""
C{compat.unicode} is C{str} on Python 3, C{unicode} on Python 2.
"""
if _PY3:
expected = str
else:
expected = unicode
self.assertIs(unicodeCompat, expected)
def test_nativeStringIO(self):
"""
L{NativeStringIO} is a file-like object that stores native strings in
memory.
"""
f = NativeStringIO()
f.write("hello")
f.write(" there")
self.assertEqual(f.getvalue(), "hello there")
class NetworkStringTests(unittest.SynchronousTestCase):
"""
Tests for L{networkString}.
"""
def test_bytes(self):
"""
L{networkString} returns a C{bytes} object passed to it unmodified.
"""
self.assertEqual(b"foo", networkString(b"foo"))
def test_bytesOutOfRange(self):
"""
L{networkString} raises C{UnicodeError} if passed a C{bytes} instance
containing bytes not used by ASCII.
"""
self.assertRaises(
UnicodeError, networkString, u"\N{SNOWMAN}".encode('utf-8'))
if _PY3:
test_bytes.skip = test_bytesOutOfRange.skip = (
"Bytes behavior of networkString only provided on Python 2.")
def test_unicode(self):
"""
L{networkString} returns a C{unicode} object passed to it encoded into
a C{bytes} instance.
"""
self.assertEqual(b"foo", networkString(u"foo"))
def test_unicodeOutOfRange(self):
"""
L{networkString} raises L{UnicodeError} if passed a C{unicode} instance
containing characters not encodable in ASCII.
"""
self.assertRaises(
UnicodeError, networkString, u"\N{SNOWMAN}")
if not _PY3:
test_unicode.skip = test_unicodeOutOfRange.skip = (
"Unicode behavior of networkString only provided on Python 3.")
def test_nonString(self):
"""
L{networkString} raises L{TypeError} if passed a non-string object or
the wrong type of string object.
"""
self.assertRaises(TypeError, networkString, object())
if _PY3:
self.assertRaises(TypeError, networkString, b"bytes")
else:
self.assertRaises(TypeError, networkString, u"text")
class ReraiseTests(unittest.SynchronousTestCase):
"""
L{reraise} re-raises exceptions on both Python 2 and Python 3.
"""
def test_reraiseWithNone(self):
"""
Calling L{reraise} with an exception instance and a traceback of
L{None} re-raises it with a new traceback.
"""
try:
1/0
except:
typ, value, tb = sys.exc_info()
try:
reraise(value, None)
except:
typ2, value2, tb2 = sys.exc_info()
self.assertEqual(typ2, ZeroDivisionError)
self.assertIs(value, value2)
self.assertNotEqual(traceback.format_tb(tb)[-1],
traceback.format_tb(tb2)[-1])
else:
self.fail("The exception was not raised.")
def test_reraiseWithTraceback(self):
"""
Calling L{reraise} with an exception instance and a traceback
re-raises the exception with the given traceback.
"""
try:
1/0
except:
typ, value, tb = sys.exc_info()
try:
reraise(value, tb)
except:
typ2, value2, tb2 = sys.exc_info()
self.assertEqual(typ2, ZeroDivisionError)
self.assertIs(value, value2)
self.assertEqual(traceback.format_tb(tb)[-1],
traceback.format_tb(tb2)[-1])
else:
self.fail("The exception was not raised.")
class Python3BytesTests(unittest.SynchronousTestCase):
"""
Tests for L{iterbytes}, L{intToBytes}, L{lazyByteSlice}.
"""
def test_iteration(self):
"""
When L{iterbytes} is called with a bytestring, the returned object
can be iterated over, resulting in the individual bytes of the
bytestring.
"""
input = b"abcd"
result = list(iterbytes(input))
self.assertEqual(result, [b'a', b'b', b'c', b'd'])
def test_intToBytes(self):
"""
When L{intToBytes} is called with an integer, the result is an
ASCII-encoded string representation of the number.
"""
self.assertEqual(intToBytes(213), b"213")
def test_lazyByteSliceNoOffset(self):
"""
L{lazyByteSlice} called with some bytes returns a semantically equal
version of these bytes.
"""
data = b'123XYZ'
self.assertEqual(bytes(lazyByteSlice(data)), data)
def test_lazyByteSliceOffset(self):
"""
L{lazyByteSlice} called with some bytes and an offset returns a
semantically equal version of these bytes starting at the given offset.
"""
data = b'123XYZ'
self.assertEqual(bytes(lazyByteSlice(data, 2)), data[2:])
def test_lazyByteSliceOffsetAndLength(self):
"""
L{lazyByteSlice} called with some bytes, an offset and a length returns
a semantically equal version of these bytes starting at the given
offset, up to the given length.
"""
data = b'123XYZ'
self.assertEqual(bytes(lazyByteSlice(data, 2, 3)), data[2:5])
class BytesEnvironTests(unittest.TestCase):
"""
Tests for L{BytesEnviron}.
"""
def test_alwaysBytes(self):
"""
The output of L{BytesEnviron} should always be a L{dict} with L{bytes}
values and L{bytes} keys.
"""
result = bytesEnviron()
types = set()
for key, val in iteritems(result):
types.add(type(key))
types.add(type(val))
self.assertEqual(list(types), [bytes])
if platform.isWindows():
test_alwaysBytes.skip = "Environment vars are always str on Windows."
class CoercedUnicodeTests(unittest.TestCase):
"""
Tests for L{twisted.python.compat._coercedUnicode}.
"""
def test_unicodeASCII(self):
"""
Unicode strings with ASCII code points are unchanged.
"""
result = _coercedUnicode(u'text')
self.assertEqual(result, u'text')
self.assertIsInstance(result, unicodeCompat)
def test_unicodeNonASCII(self):
"""
Unicode strings with non-ASCII code points are unchanged.
"""
result = _coercedUnicode(u'\N{SNOWMAN}')
self.assertEqual(result, u'\N{SNOWMAN}')
self.assertIsInstance(result, unicodeCompat)
def test_nativeASCII(self):
"""
Native strings with ASCII code points are unchanged.
On Python 2, this verifies that ASCII-only byte strings are accepted,
whereas for Python 3 it is identical to L{test_unicodeASCII}.
"""
result = _coercedUnicode('text')
self.assertEqual(result, u'text')
self.assertIsInstance(result, unicodeCompat)
def test_bytesPy3(self):
"""
Byte strings are not accceptable in Python 3.
"""
exc = self.assertRaises(TypeError, _coercedUnicode, b'bytes')
self.assertEqual(str(exc), "Expected str not b'bytes' (bytes)")
if not _PY3:
test_bytesPy3.skip = (
"Bytes behavior of _coercedUnicode only provided on Python 2.")
def test_bytesNonASCII(self):
"""
Byte strings with non-ASCII code points raise an exception.
"""
self.assertRaises(UnicodeError, _coercedUnicode, b'\xe2\x98\x83')
if _PY3:
test_bytesNonASCII.skip = (
"Bytes behavior of _coercedUnicode only provided on Python 2.")
class UnichrTests(unittest.TestCase):
"""
Tests for L{unichr}.
"""
def test_unichr(self):
"""
unichar exists and returns a unicode string with the given code point.
"""
self.assertEqual(unichr(0x2603), u"\N{SNOWMAN}")
class RawInputTests(unittest.TestCase):
"""
Tests for L{raw_input}
"""
def test_raw_input(self):
"""
L{twisted.python.compat.raw_input}
"""
class FakeStdin:
def readline(self):
return "User input\n"
class FakeStdout:
data = ""
def write(self, data):
self.data += data
self.patch(sys, "stdin", FakeStdin())
stdout = FakeStdout()
self.patch(sys, "stdout", stdout)
self.assertEqual(raw_input("Prompt"), "User input")
self.assertEqual(stdout.data, "Prompt")
class FutureBytesReprTests(unittest.TestCase):
"""
Tests for L{twisted.python.compat._bytesRepr}.
"""
def test_bytesReprNotBytes(self):
"""
L{twisted.python.compat._bytesRepr} raises a
L{TypeError} when called any object that is not an instance of
L{bytes}.
"""
exc = self.assertRaises(TypeError, _bytesRepr, ["not bytes"])
self.assertEquals(str(exc), "Expected bytes not ['not bytes']")
def test_bytesReprPrefix(self):
"""
L{twisted.python.compat._bytesRepr} always prepends
``b`` to the returned repr on both Python 2 and 3.
"""
self.assertEqual(_bytesRepr(b'\x00'), "b'\\x00'")
class GetAsyncParamTests(unittest.SynchronousTestCase):
"""
Tests for L{twisted.python.compat._get_async_param}
"""
def test_get_async_param(self):
"""
L{twisted.python.compat._get_async_param} uses isAsync by default,
or deprecated async keyword argument if isAsync is None.
"""
self.assertEqual(_get_async_param(isAsync=False), False)
self.assertEqual(_get_async_param(isAsync=True), True)
self.assertEqual(
_get_async_param(isAsync=None, **{'async': False}), False)
self.assertEqual(
_get_async_param(isAsync=None, **{'async': True}), True)
self.assertRaises(TypeError, _get_async_param, False, {'async': False})
def test_get_async_param_deprecation(self):
"""
L{twisted.python.compat._get_async_param} raises a deprecation
warning if async keyword argument is passed.
"""
self.assertEqual(
_get_async_param(isAsync=None, **{'async': False}), False)
currentWarnings = self.flushWarnings(
offendingFunctions=[self.test_get_async_param_deprecation])
self.assertEqual(
currentWarnings[0]['message'],
"'async' keyword argument is deprecated, please use isAsync")
@@ -0,0 +1,711 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
This module contains tests for L{twisted.internet.task.Cooperator} and
related functionality.
"""
from __future__ import division, absolute_import
from twisted.internet import reactor, defer, task
from twisted.trial import unittest
class FakeDelayedCall(object):
"""
Fake delayed call which lets us simulate the scheduler.
"""
def __init__(self, func):
"""
A function to run, later.
"""
self.func = func
self.cancelled = False
def cancel(self):
"""
Don't run my function later.
"""
self.cancelled = True
class FakeScheduler(object):
"""
A fake scheduler for testing against.
"""
def __init__(self):
"""
Create a fake scheduler with a list of work to do.
"""
self.work = []
def __call__(self, thunk):
"""
Schedule a unit of work to be done later.
"""
unit = FakeDelayedCall(thunk)
self.work.append(unit)
return unit
def pump(self):
"""
Do all of the work that is currently available to be done.
"""
work, self.work = self.work, []
for unit in work:
if not unit.cancelled:
unit.func()
class CooperatorTests(unittest.TestCase):
RESULT = 'done'
def ebIter(self, err):
err.trap(task.SchedulerStopped)
return self.RESULT
def cbIter(self, ign):
self.fail()
def testStoppedRejectsNewTasks(self):
"""
Test that Cooperators refuse new tasks when they have been stopped.
"""
def testwith(stuff):
c = task.Cooperator()
c.stop()
d = c.coiterate(iter(()), stuff)
d.addCallback(self.cbIter)
d.addErrback(self.ebIter)
return d.addCallback(lambda result:
self.assertEqual(result, self.RESULT))
return testwith(None).addCallback(lambda ign: testwith(defer.Deferred()))
def testStopRunning(self):
"""
Test that a running iterator will not run to completion when the
cooperator is stopped.
"""
c = task.Cooperator()
def myiter():
for myiter.value in range(3):
yield myiter.value
myiter.value = -1
d = c.coiterate(myiter())
d.addCallback(self.cbIter)
d.addErrback(self.ebIter)
c.stop()
def doasserts(result):
self.assertEqual(result, self.RESULT)
self.assertEqual(myiter.value, -1)
d.addCallback(doasserts)
return d
def testStopOutstanding(self):
"""
An iterator run with L{Cooperator.coiterate} paused on a L{Deferred}
yielded by that iterator will fire its own L{Deferred} (the one
returned by C{coiterate}) when L{Cooperator.stop} is called.
"""
testControlD = defer.Deferred()
outstandingD = defer.Deferred()
def myiter():
reactor.callLater(0, testControlD.callback, None)
yield outstandingD
self.fail()
c = task.Cooperator()
d = c.coiterate(myiter())
def stopAndGo(ign):
c.stop()
outstandingD.callback('arglebargle')
testControlD.addCallback(stopAndGo)
d.addCallback(self.cbIter)
d.addErrback(self.ebIter)
return d.addCallback(
lambda result: self.assertEqual(result, self.RESULT))
def testUnexpectedError(self):
c = task.Cooperator()
def myiter():
if 0:
yield None
else:
raise RuntimeError()
d = c.coiterate(myiter())
return self.assertFailure(d, RuntimeError)
def testUnexpectedErrorActuallyLater(self):
def myiter():
D = defer.Deferred()
reactor.callLater(0, D.errback, RuntimeError())
yield D
c = task.Cooperator()
d = c.coiterate(myiter())
return self.assertFailure(d, RuntimeError)
def testUnexpectedErrorNotActuallyLater(self):
def myiter():
yield defer.fail(RuntimeError())
c = task.Cooperator()
d = c.coiterate(myiter())
return self.assertFailure(d, RuntimeError)
def testCooperation(self):
L = []
def myiter(things):
for th in things:
L.append(th)
yield None
groupsOfThings = ['abc', (1, 2, 3), 'def', (4, 5, 6)]
c = task.Cooperator()
tasks = []
for stuff in groupsOfThings:
tasks.append(c.coiterate(myiter(stuff)))
return defer.DeferredList(tasks).addCallback(
lambda ign: self.assertEqual(tuple(L), sum(zip(*groupsOfThings), ())))
def testResourceExhaustion(self):
output = []
def myiter():
for i in range(100):
output.append(i)
if i == 9:
_TPF.stopped = True
yield i
class _TPF:
stopped = False
def __call__(self):
return self.stopped
c = task.Cooperator(terminationPredicateFactory=_TPF)
c.coiterate(myiter()).addErrback(self.ebIter)
c._delayedCall.cancel()
# testing a private method because only the test case will ever care
# about this, so we have to carefully clean up after ourselves.
c._tick()
c.stop()
self.assertTrue(_TPF.stopped)
self.assertEqual(output, list(range(10)))
def testCallbackReCoiterate(self):
"""
If a callback to a deferred returned by coiterate calls coiterate on
the same Cooperator, we should make sure to only do the minimal amount
of scheduling work. (This test was added to demonstrate a specific bug
that was found while writing the scheduler.)
"""
calls = []
class FakeCall:
def __init__(self, func):
self.func = func
def __repr__(self):
return '<FakeCall %r>' % (self.func,)
def sched(f):
self.assertFalse(calls, repr(calls))
calls.append(FakeCall(f))
return calls[-1]
c = task.Cooperator(scheduler=sched, terminationPredicateFactory=lambda: lambda: True)
d = c.coiterate(iter(()))
done = []
def anotherTask(ign):
c.coiterate(iter(())).addBoth(done.append)
d.addCallback(anotherTask)
work = 0
while not done:
work += 1
while calls:
calls.pop(0).func()
work += 1
if work > 50:
self.fail("Cooperator took too long")
def test_removingLastTaskStopsScheduledCall(self):
"""
If the last task in a Cooperator is removed, the scheduled call for
the next tick is cancelled, since it is no longer necessary.
This behavior is useful for tests that want to assert they have left
no reactor state behind when they're done.
"""
calls = [None]
def sched(f):
calls[0] = FakeDelayedCall(f)
return calls[0]
coop = task.Cooperator(scheduler=sched)
# Add two task; this should schedule the tick:
task1 = coop.cooperate(iter([1, 2]))
task2 = coop.cooperate(iter([1, 2]))
self.assertEqual(calls[0].func, coop._tick)
# Remove first task; scheduled call should still be going:
task1.stop()
self.assertFalse(calls[0].cancelled)
self.assertEqual(coop._delayedCall, calls[0])
# Remove second task; scheduled call should be cancelled:
task2.stop()
self.assertTrue(calls[0].cancelled)
self.assertIsNone(coop._delayedCall)
# Add another task; scheduled call will be recreated:
coop.cooperate(iter([1, 2]))
self.assertFalse(calls[0].cancelled)
self.assertEqual(coop._delayedCall, calls[0])
def test_runningWhenStarted(self):
"""
L{Cooperator.running} reports C{True} if the L{Cooperator}
was started on creation.
"""
c = task.Cooperator()
self.assertTrue(c.running)
def test_runningWhenNotStarted(self):
"""
L{Cooperator.running} reports C{False} if the L{Cooperator}
has not been started.
"""
c = task.Cooperator(started=False)
self.assertFalse(c.running)
def test_runningWhenRunning(self):
"""
L{Cooperator.running} reports C{True} when the L{Cooperator}
is running.
"""
c = task.Cooperator(started=False)
c.start()
self.addCleanup(c.stop)
self.assertTrue(c.running)
def test_runningWhenStopped(self):
"""
L{Cooperator.running} reports C{False} after the L{Cooperator}
has been stopped.
"""
c = task.Cooperator(started=False)
c.start()
c.stop()
self.assertFalse(c.running)
class UnhandledException(Exception):
"""
An exception that should go unhandled.
"""
class AliasTests(unittest.TestCase):
"""
Integration test to verify that the global singleton aliases do what
they're supposed to.
"""
def test_cooperate(self):
"""
L{twisted.internet.task.cooperate} ought to run the generator that it is
"""
d = defer.Deferred()
def doit():
yield 1
yield 2
yield 3
d.callback("yay")
it = doit()
theTask = task.cooperate(it)
self.assertIn(theTask, task._theCooperator._tasks)
return d
class RunStateTests(unittest.TestCase):
"""
Tests to verify the behavior of L{CooperativeTask.pause},
L{CooperativeTask.resume}, L{CooperativeTask.stop}, exhausting the
underlying iterator, and their interactions with each other.
"""
def setUp(self):
"""
Create a cooperator with a fake scheduler and a termination predicate
that ensures only one unit of work will take place per tick.
"""
self._doDeferNext = False
self._doStopNext = False
self._doDieNext = False
self.work = []
self.scheduler = FakeScheduler()
self.cooperator = task.Cooperator(
scheduler=self.scheduler,
# Always stop after one iteration of work (return a function which
# returns a function which always returns True)
terminationPredicateFactory=lambda: lambda: True)
self.task = self.cooperator.cooperate(self.worker())
self.cooperator.start()
def worker(self):
"""
This is a sample generator which yields Deferreds when we are testing
deferral and an ascending integer count otherwise.
"""
i = 0
while True:
i += 1
if self._doDeferNext:
self._doDeferNext = False
d = defer.Deferred()
self.work.append(d)
yield d
elif self._doStopNext:
return
elif self._doDieNext:
raise UnhandledException()
else:
self.work.append(i)
yield i
def tearDown(self):
"""
Drop references to interesting parts of the fixture to allow Deferred
errors to be noticed when things start failing.
"""
del self.task
del self.scheduler
def deferNext(self):
"""
Defer the next result from my worker iterator.
"""
self._doDeferNext = True
def stopNext(self):
"""
Make the next result from my worker iterator be completion (raising
StopIteration).
"""
self._doStopNext = True
def dieNext(self):
"""
Make the next result from my worker iterator be raising an
L{UnhandledException}.
"""
def ignoreUnhandled(failure):
failure.trap(UnhandledException)
return None
self._doDieNext = True
def test_pauseResume(self):
"""
Cooperators should stop running their tasks when they're paused, and
start again when they're resumed.
"""
# first, sanity check
self.scheduler.pump()
self.assertEqual(self.work, [1])
self.scheduler.pump()
self.assertEqual(self.work, [1, 2])
# OK, now for real
self.task.pause()
self.scheduler.pump()
self.assertEqual(self.work, [1, 2])
self.task.resume()
# Resuming itself shoult not do any work
self.assertEqual(self.work, [1, 2])
self.scheduler.pump()
# But when the scheduler rolls around again...
self.assertEqual(self.work, [1, 2, 3])
def test_resumeNotPaused(self):
"""
L{CooperativeTask.resume} should raise a L{TaskNotPaused} exception if
it was not paused; e.g. if L{CooperativeTask.pause} was not invoked
more times than L{CooperativeTask.resume} on that object.
"""
self.assertRaises(task.NotPaused, self.task.resume)
self.task.pause()
self.task.resume()
self.assertRaises(task.NotPaused, self.task.resume)
def test_pauseTwice(self):
"""
Pauses on tasks should behave like a stack. If a task is paused twice,
it needs to be resumed twice.
"""
# pause once
self.task.pause()
self.scheduler.pump()
self.assertEqual(self.work, [])
# pause twice
self.task.pause()
self.scheduler.pump()
self.assertEqual(self.work, [])
# resume once (it shouldn't)
self.task.resume()
self.scheduler.pump()
self.assertEqual(self.work, [])
# resume twice (now it should go)
self.task.resume()
self.scheduler.pump()
self.assertEqual(self.work, [1])
def test_pauseWhileDeferred(self):
"""
C{pause()}ing a task while it is waiting on an outstanding
L{defer.Deferred} should put the task into a state where the
outstanding L{defer.Deferred} must be called back I{and} the task is
C{resume}d before it will continue processing.
"""
self.deferNext()
self.scheduler.pump()
self.assertEqual(len(self.work), 1)
self.assertIsInstance(self.work[0], defer.Deferred)
self.scheduler.pump()
self.assertEqual(len(self.work), 1)
self.task.pause()
self.scheduler.pump()
self.assertEqual(len(self.work), 1)
self.task.resume()
self.scheduler.pump()
self.assertEqual(len(self.work), 1)
self.work[0].callback("STUFF!")
self.scheduler.pump()
self.assertEqual(len(self.work), 2)
self.assertEqual(self.work[1], 2)
def test_whenDone(self):
"""
L{CooperativeTask.whenDone} returns a Deferred which fires when the
Cooperator's iterator is exhausted. It returns a new Deferred each
time it is called; callbacks added to other invocations will not modify
the value that subsequent invocations will fire with.
"""
deferred1 = self.task.whenDone()
deferred2 = self.task.whenDone()
results1 = []
results2 = []
final1 = []
final2 = []
def callbackOne(result):
results1.append(result)
return 1
def callbackTwo(result):
results2.append(result)
return 2
deferred1.addCallback(callbackOne)
deferred2.addCallback(callbackTwo)
deferred1.addCallback(final1.append)
deferred2.addCallback(final2.append)
# exhaust the task iterator
# callbacks fire
self.stopNext()
self.scheduler.pump()
self.assertEqual(len(results1), 1)
self.assertEqual(len(results2), 1)
self.assertIs(results1[0], self.task._iterator)
self.assertIs(results2[0], self.task._iterator)
self.assertEqual(final1, [1])
self.assertEqual(final2, [2])
def test_whenDoneError(self):
"""
L{CooperativeTask.whenDone} returns a L{defer.Deferred} that will fail
when the iterable's C{next} method raises an exception, with that
exception.
"""
deferred1 = self.task.whenDone()
results = []
deferred1.addErrback(results.append)
self.dieNext()
self.scheduler.pump()
self.assertEqual(len(results), 1)
self.assertEqual(results[0].check(UnhandledException), UnhandledException)
def test_whenDoneStop(self):
"""
L{CooperativeTask.whenDone} returns a L{defer.Deferred} that fails with
L{TaskStopped} when the C{stop} method is called on that
L{CooperativeTask}.
"""
deferred1 = self.task.whenDone()
errors = []
deferred1.addErrback(errors.append)
self.task.stop()
self.assertEqual(len(errors), 1)
self.assertEqual(errors[0].check(task.TaskStopped), task.TaskStopped)
def test_whenDoneAlreadyDone(self):
"""
L{CooperativeTask.whenDone} will return a L{defer.Deferred} that will
succeed immediately if its iterator has already completed.
"""
self.stopNext()
self.scheduler.pump()
results = []
self.task.whenDone().addCallback(results.append)
self.assertEqual(results, [self.task._iterator])
def test_stopStops(self):
"""
C{stop()}ping a task should cause it to be removed from the run just as
C{pause()}ing, with the distinction that C{resume()} will raise a
L{TaskStopped} exception.
"""
self.task.stop()
self.scheduler.pump()
self.assertEqual(len(self.work), 0)
self.assertRaises(task.TaskStopped, self.task.stop)
self.assertRaises(task.TaskStopped, self.task.pause)
# Sanity check - it's still not scheduled, is it?
self.scheduler.pump()
self.assertEqual(self.work, [])
def test_pauseStopResume(self):
"""
C{resume()}ing a paused, stopped task should be a no-op; it should not
raise an exception, because it's paused, but neither should it actually
do more work from the task.
"""
self.task.pause()
self.task.stop()
self.task.resume()
self.scheduler.pump()
self.assertEqual(self.work, [])
def test_stopDeferred(self):
"""
As a corrolary of the interaction of C{pause()} and C{unpause()},
C{stop()}ping a task which is waiting on a L{Deferred} should cause the
task to gracefully shut down, meaning that it should not be unpaused
when the deferred fires.
"""
self.deferNext()
self.scheduler.pump()
d = self.work.pop()
self.assertEqual(self.task._pauseCount, 1)
results = []
d.addBoth(results.append)
self.scheduler.pump()
self.task.stop()
self.scheduler.pump()
d.callback(7)
self.scheduler.pump()
# Let's make sure that Deferred doesn't come out fried with an
# unhandled error that will be logged. The value is None, rather than
# our test value, 7, because this Deferred is returned to and consumed
# by the cooperator code. Its callback therefore has no contract.
self.assertEqual(results, [None])
# But more importantly, no further work should have happened.
self.assertEqual(self.work, [])
def test_stopExhausted(self):
"""
C{stop()}ping a L{CooperativeTask} whose iterator has been exhausted
should raise L{TaskDone}.
"""
self.stopNext()
self.scheduler.pump()
self.assertRaises(task.TaskDone, self.task.stop)
def test_stopErrored(self):
"""
C{stop()}ping a L{CooperativeTask} whose iterator has encountered an
error should raise L{TaskFailed}.
"""
self.dieNext()
self.scheduler.pump()
self.assertRaises(task.TaskFailed, self.task.stop)
def test_stopCooperatorReentrancy(self):
"""
If a callback of a L{Deferred} from L{CooperativeTask.whenDone} calls
C{Cooperator.stop} on its L{CooperativeTask._cooperator}, the
L{Cooperator} will stop, but the L{CooperativeTask} whose callback is
calling C{stop} should already be considered 'stopped' by the time the
callback is running, and therefore removed from the
L{CoooperativeTask}.
"""
callbackPhases = []
def stopit(result):
callbackPhases.append(result)
self.cooperator.stop()
# "done" here is a sanity check to make sure that we get all the
# way through the callback; i.e. stop() shouldn't be raising an
# exception due to the stopped-ness of our main task.
callbackPhases.append("done")
self.task.whenDone().addCallback(stopit)
self.stopNext()
self.scheduler.pump()
self.assertEqual(callbackPhases, [self.task._iterator, "done"])
@@ -0,0 +1,43 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
from twisted.trial import unittest
from twisted.trial.unittest import SynchronousTestCase
from twisted.protocols import dict
paramString = b"\"This is a dqstring \\w\\i\\t\\h boring stuff like: \\\"\" and t\\hes\\\"e are a\\to\\ms"
goodparams = [b"This is a dqstring with boring stuff like: \"", b"and", b"thes\"e", b"are", b"atoms"]
class ParamTests(unittest.TestCase):
def testParseParam(self):
"""Testing command response handling"""
params = []
rest = paramString
while 1:
(param, rest) = dict.parseParam(rest)
if param == None:
break
params.append(param)
self.assertEqual(params, goodparams)#, "DictClient.parseParam returns unexpected results")
class DictDeprecationTests(SynchronousTestCase):
"""
L{twisted.protocols.dict} is deprecated.
"""
def test_dictDeprecation(self):
"""
L{twisted.protocols.dict} is deprecated since Twisted 17.9.0.
"""
from twisted.protocols import dict
dict
warningsShown = self.flushWarnings([self.test_dictDeprecation])
self.assertEqual(1, len(warningsShown))
self.assertEqual(
("twisted.protocols.dict was deprecated in Twisted 17.9.0:"
" There is no replacement for this module."),
warningsShown[0]['message'])
@@ -0,0 +1,233 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Test cases for twisted.protocols.ident module.
"""
import struct
from twisted.protocols import ident
from twisted.python import failure
from twisted.internet import error
from twisted.internet import defer
from twisted.python.compat import NativeStringIO
from twisted.trial import unittest
from twisted.test.proto_helpers import StringTransport
try:
import builtins
except ImportError:
import __builtin__ as builtins
class ClassParserTests(unittest.TestCase):
"""
Test parsing of ident responses.
"""
def setUp(self):
"""
Create an ident client used in tests.
"""
self.client = ident.IdentClient()
def test_indentError(self):
"""
'UNKNOWN-ERROR' error should map to the L{ident.IdentError} exception.
"""
d = defer.Deferred()
self.client.queries.append((d, 123, 456))
self.client.lineReceived('123, 456 : ERROR : UNKNOWN-ERROR')
return self.assertFailure(d, ident.IdentError)
def test_noUSerError(self):
"""
'NO-USER' error should map to the L{ident.NoUser} exception.
"""
d = defer.Deferred()
self.client.queries.append((d, 234, 456))
self.client.lineReceived('234, 456 : ERROR : NO-USER')
return self.assertFailure(d, ident.NoUser)
def test_invalidPortError(self):
"""
'INVALID-PORT' error should map to the L{ident.InvalidPort} exception.
"""
d = defer.Deferred()
self.client.queries.append((d, 345, 567))
self.client.lineReceived('345, 567 : ERROR : INVALID-PORT')
return self.assertFailure(d, ident.InvalidPort)
def test_hiddenUserError(self):
"""
'HIDDEN-USER' error should map to the L{ident.HiddenUser} exception.
"""
d = defer.Deferred()
self.client.queries.append((d, 567, 789))
self.client.lineReceived('567, 789 : ERROR : HIDDEN-USER')
return self.assertFailure(d, ident.HiddenUser)
def test_lostConnection(self):
"""
A pending query which failed because of a ConnectionLost should
receive an L{ident.IdentError}.
"""
d = defer.Deferred()
self.client.queries.append((d, 765, 432))
self.client.connectionLost(failure.Failure(error.ConnectionLost()))
return self.assertFailure(d, ident.IdentError)
class TestIdentServer(ident.IdentServer):
def lookup(self, serverAddress, clientAddress):
return self.resultValue
class TestErrorIdentServer(ident.IdentServer):
def lookup(self, serverAddress, clientAddress):
raise self.exceptionType()
class NewException(RuntimeError):
pass
class ServerParserTests(unittest.TestCase):
def testErrors(self):
p = TestErrorIdentServer()
p.makeConnection(StringTransport())
L = []
p.sendLine = L.append
p.exceptionType = ident.IdentError
p.lineReceived('123, 345')
self.assertEqual(L[0], '123, 345 : ERROR : UNKNOWN-ERROR')
p.exceptionType = ident.NoUser
p.lineReceived('432, 210')
self.assertEqual(L[1], '432, 210 : ERROR : NO-USER')
p.exceptionType = ident.InvalidPort
p.lineReceived('987, 654')
self.assertEqual(L[2], '987, 654 : ERROR : INVALID-PORT')
p.exceptionType = ident.HiddenUser
p.lineReceived('756, 827')
self.assertEqual(L[3], '756, 827 : ERROR : HIDDEN-USER')
p.exceptionType = NewException
p.lineReceived('987, 789')
self.assertEqual(L[4], '987, 789 : ERROR : UNKNOWN-ERROR')
errs = self.flushLoggedErrors(NewException)
self.assertEqual(len(errs), 1)
for port in -1, 0, 65536, 65537:
del L[:]
p.lineReceived('%d, 5' % (port,))
p.lineReceived('5, %d' % (port,))
self.assertEqual(
L, ['%d, 5 : ERROR : INVALID-PORT' % (port,),
'5, %d : ERROR : INVALID-PORT' % (port,)])
def testSuccess(self):
p = TestIdentServer()
p.makeConnection(StringTransport())
L = []
p.sendLine = L.append
p.resultValue = ('SYS', 'USER')
p.lineReceived('123, 456')
self.assertEqual(L[0], '123, 456 : USERID : SYS : USER')
if struct.pack('=L', 1)[0:1] == b'\x01':
_addr1 = '0100007F'
_addr2 = '04030201'
else:
_addr1 = '7F000001'
_addr2 = '01020304'
class ProcMixinTests(unittest.TestCase):
line = ('4: %s:0019 %s:02FA 0A 00000000:00000000 '
'00:00000000 00000000 0 0 10927 1 f72a5b80 '
'3000 0 0 2 -1') % (_addr1, _addr2)
sampleFile = (' sl local_address rem_address st tx_queue rx_queue tr '
'tm->when retrnsmt uid timeout inode\n ' + line)
def testDottedQuadFromHexString(self):
p = ident.ProcServerMixin()
self.assertEqual(p.dottedQuadFromHexString(_addr1), '127.0.0.1')
def testUnpackAddress(self):
p = ident.ProcServerMixin()
self.assertEqual(p.unpackAddress(_addr1 + ':0277'),
('127.0.0.1', 631))
def testLineParser(self):
p = ident.ProcServerMixin()
self.assertEqual(
p.parseLine(self.line),
(('127.0.0.1', 25), ('1.2.3.4', 762), 0))
def testExistingAddress(self):
username = []
p = ident.ProcServerMixin()
p.entries = lambda: iter([self.line])
p.getUsername = lambda uid: (username.append(uid), 'root')[1]
self.assertEqual(
p.lookup(('127.0.0.1', 25), ('1.2.3.4', 762)),
(p.SYSTEM_NAME, 'root'))
self.assertEqual(username, [0])
def testNonExistingAddress(self):
p = ident.ProcServerMixin()
p.entries = lambda: iter([self.line])
self.assertRaises(ident.NoUser, p.lookup, ('127.0.0.1', 26),
('1.2.3.4', 762))
self.assertRaises(ident.NoUser, p.lookup, ('127.0.0.1', 25),
('1.2.3.5', 762))
self.assertRaises(ident.NoUser, p.lookup, ('127.0.0.1', 25),
('1.2.3.4', 763))
def testLookupProcNetTcp(self):
"""
L{ident.ProcServerMixin.lookup} uses the Linux TCP process table.
"""
open_calls = []
def mocked_open(*args, **kwargs):
"""
Mock for the open call to prevent actually opening /proc/net/tcp.
"""
open_calls.append((args, kwargs))
return NativeStringIO(self.sampleFile)
self.patch(builtins, 'open', mocked_open)
p = ident.ProcServerMixin()
self.assertRaises(ident.NoUser, p.lookup, ('127.0.0.1', 26),
('1.2.3.4', 762))
self.assertEqual([(('/proc/net/tcp',), {})], open_calls)
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,562 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
from __future__ import division, absolute_import
import contextlib
import errno
import os
import stat
import time
from twisted.trial import unittest
from twisted.python import logfile, runtime
class LogFileTests(unittest.TestCase):
"""
Test the rotating log file.
"""
def setUp(self):
self.dir = self.mktemp()
os.makedirs(self.dir)
self.name = "test.log"
self.path = os.path.join(self.dir, self.name)
def tearDown(self):
"""
Restore back write rights on created paths: if tests modified the
rights, that will allow the paths to be removed easily afterwards.
"""
os.chmod(self.dir, 0o777)
if os.path.exists(self.path):
os.chmod(self.path, 0o777)
def test_abstractShouldRotate(self):
"""
L{BaseLogFile.shouldRotate} is abstract and must be implemented by
subclass.
"""
log = logfile.BaseLogFile(self.name, self.dir)
self.addCleanup(log.close)
self.assertRaises(NotImplementedError, log.shouldRotate)
def test_writing(self):
"""
Log files can be written to, flushed and closed. Closing a log file
also flushes it.
"""
with contextlib.closing(logfile.LogFile(self.name, self.dir)) as log:
log.write("123")
log.write("456")
log.flush()
log.write("7890")
with open(self.path) as f:
self.assertEqual(f.read(), "1234567890")
def test_rotation(self):
"""
Rotating log files autorotate after a period of time, and can also be
manually rotated.
"""
# this logfile should rotate every 10 bytes
with contextlib.closing(
logfile.LogFile(self.name, self.dir, rotateLength=10)) as log:
# test automatic rotation
log.write("123")
log.write("4567890")
log.write("1" * 11)
self.assertTrue(os.path.exists("{0}.1".format(self.path)))
self.assertFalse(os.path.exists("{0}.2".format(self.path)))
log.write('')
self.assertTrue(os.path.exists("{0}.1".format(self.path)))
self.assertTrue(os.path.exists("{0}.2".format(self.path)))
self.assertFalse(os.path.exists("{0}.3".format(self.path)))
log.write("3")
self.assertFalse(os.path.exists("{0}.3".format(self.path)))
# test manual rotation
log.rotate()
self.assertTrue(os.path.exists("{0}.3".format(self.path)))
self.assertFalse(os.path.exists("{0}.4".format(self.path)))
self.assertEqual(log.listLogs(), [1, 2, 3])
def test_append(self):
"""
Log files can be written to, closed. Their size is the number of
bytes written to them. Everything that was written to them can
be read, even if the writing happened on separate occasions,
and even if the log file was closed in between.
"""
with contextlib.closing(logfile.LogFile(self.name, self.dir)) as log:
log.write("0123456789")
log = logfile.LogFile(self.name, self.dir)
self.addCleanup(log.close)
self.assertEqual(log.size, 10)
self.assertEqual(log._file.tell(), log.size)
log.write("abc")
self.assertEqual(log.size, 13)
self.assertEqual(log._file.tell(), log.size)
f = log._file
f.seek(0, 0)
self.assertEqual(f.read(), b"0123456789abc")
def test_logReader(self):
"""
Various tests for log readers.
First of all, log readers can get logs by number and read what
was written to those log files. Getting nonexistent log files
raises C{ValueError}. Using anything other than an integer
index raises C{TypeError}. As logs get older, their log
numbers increase.
"""
log = logfile.LogFile(self.name, self.dir)
self.addCleanup(log.close)
log.write("abc\n")
log.write("def\n")
log.rotate()
log.write("ghi\n")
log.flush()
# check reading logs
self.assertEqual(log.listLogs(), [1])
with contextlib.closing(log.getCurrentLog()) as reader:
reader._file.seek(0)
self.assertEqual(reader.readLines(), ["ghi\n"])
self.assertEqual(reader.readLines(), [])
with contextlib.closing(log.getLog(1)) as reader:
self.assertEqual(reader.readLines(), ["abc\n", "def\n"])
self.assertEqual(reader.readLines(), [])
# check getting illegal log readers
self.assertRaises(ValueError, log.getLog, 2)
self.assertRaises(TypeError, log.getLog, "1")
# check that log numbers are higher for older logs
log.rotate()
self.assertEqual(log.listLogs(), [1, 2])
with contextlib.closing(log.getLog(1)) as reader:
reader._file.seek(0)
self.assertEqual(reader.readLines(), ["ghi\n"])
self.assertEqual(reader.readLines(), [])
with contextlib.closing(log.getLog(2)) as reader:
self.assertEqual(reader.readLines(), ["abc\n", "def\n"])
self.assertEqual(reader.readLines(), [])
def test_LogReaderReadsZeroLine(self):
"""
L{LogReader.readLines} supports reading no line.
"""
# We don't need any content, just a file path that can be opened.
with open(self.path, "w"):
pass
reader = logfile.LogReader(self.path)
self.addCleanup(reader.close)
self.assertEqual([], reader.readLines(0))
def test_modePreservation(self):
"""
Check rotated files have same permissions as original.
"""
open(self.path, "w").close()
os.chmod(self.path, 0o707)
mode = os.stat(self.path)[stat.ST_MODE]
log = logfile.LogFile(self.name, self.dir)
self.addCleanup(log.close)
log.write("abc")
log.rotate()
self.assertEqual(mode, os.stat(self.path)[stat.ST_MODE])
def test_noPermission(self):
"""
Check it keeps working when permission on dir changes.
"""
log = logfile.LogFile(self.name, self.dir)
self.addCleanup(log.close)
log.write("abc")
# change permissions so rotation would fail
os.chmod(self.dir, 0o555)
# if this succeeds, chmod doesn't restrict us, so we can't
# do the test
try:
f = open(os.path.join(self.dir,"xxx"), "w")
except (OSError, IOError):
pass
else:
f.close()
return
log.rotate() # this should not fail
log.write("def")
log.flush()
f = log._file
self.assertEqual(f.tell(), 6)
f.seek(0, 0)
self.assertEqual(f.read(), b"abcdef")
def test_maxNumberOfLog(self):
"""
Test it respect the limit on the number of files when maxRotatedFiles
is not None.
"""
log = logfile.LogFile(self.name, self.dir, rotateLength=10,
maxRotatedFiles=3)
self.addCleanup(log.close)
log.write("1" * 11)
log.write("2" * 11)
self.assertTrue(os.path.exists("{0}.1".format(self.path)))
log.write("3" * 11)
self.assertTrue(os.path.exists("{0}.2".format(self.path)))
log.write("4" * 11)
self.assertTrue(os.path.exists("{0}.3".format(self.path)))
with open("{0}.3".format(self.path)) as fp:
self.assertEqual(fp.read(), "1" * 11)
log.write("5" * 11)
with open("{0}.3".format(self.path)) as fp:
self.assertEqual(fp.read(), "2" * 11)
self.assertFalse(os.path.exists("{0}.4".format(self.path)))
def test_fromFullPath(self):
"""
Test the fromFullPath method.
"""
log1 = logfile.LogFile(self.name, self.dir, 10, defaultMode=0o777)
self.addCleanup(log1.close)
log2 = logfile.LogFile.fromFullPath(self.path, 10, defaultMode=0o777)
self.addCleanup(log2.close)
self.assertEqual(log1.name, log2.name)
self.assertEqual(os.path.abspath(log1.path), log2.path)
self.assertEqual(log1.rotateLength, log2.rotateLength)
self.assertEqual(log1.defaultMode, log2.defaultMode)
def test_defaultPermissions(self):
"""
Test the default permission of the log file: if the file exist, it
should keep the permission.
"""
with open(self.path, "wb"):
os.chmod(self.path, 0o707)
currentMode = stat.S_IMODE(os.stat(self.path)[stat.ST_MODE])
log1 = logfile.LogFile(self.name, self.dir)
self.assertEqual(stat.S_IMODE(os.stat(self.path)[stat.ST_MODE]),
currentMode)
self.addCleanup(log1.close)
def test_specifiedPermissions(self):
"""
Test specifying the permissions used on the log file.
"""
log1 = logfile.LogFile(self.name, self.dir, defaultMode=0o066)
self.addCleanup(log1.close)
mode = stat.S_IMODE(os.stat(self.path)[stat.ST_MODE])
if runtime.platform.isWindows():
# The only thing we can get here is global read-only
self.assertEqual(mode, 0o444)
else:
self.assertEqual(mode, 0o066)
def test_reopen(self):
"""
L{logfile.LogFile.reopen} allows to rename the currently used file and
make L{logfile.LogFile} create a new file.
"""
with contextlib.closing(logfile.LogFile(self.name, self.dir)) as log1:
log1.write("hello1")
savePath = os.path.join(self.dir, "save.log")
os.rename(self.path, savePath)
log1.reopen()
log1.write("hello2")
with open(self.path) as f:
self.assertEqual(f.read(), "hello2")
with open(savePath) as f:
self.assertEqual(f.read(), "hello1")
if runtime.platform.isWindows():
test_reopen.skip = "Can't test reopen on Windows"
def test_nonExistentDir(self):
"""
Specifying an invalid directory to L{LogFile} raises C{IOError}.
"""
e = self.assertRaises(
IOError, logfile.LogFile, self.name, 'this_dir_does_not_exist')
self.assertEqual(e.errno, errno.ENOENT)
def test_cantChangeFileMode(self):
"""
Opening a L{LogFile} which can be read and write but whose mode can't
be changed doesn't trigger an error.
"""
if runtime.platform.isWindows():
name, directory = "NUL", ""
expectedPath = "NUL"
else:
name, directory = "null", "/dev"
expectedPath = "/dev/null"
log = logfile.LogFile(name, directory, defaultMode=0o555)
self.addCleanup(log.close)
self.assertEqual(log.path, expectedPath)
self.assertEqual(log.defaultMode, 0o555)
def test_listLogsWithBadlyNamedFiles(self):
"""
L{LogFile.listLogs} doesn't choke if it encounters a file with an
unexpected name.
"""
log = logfile.LogFile(self.name, self.dir)
self.addCleanup(log.close)
with open("{0}.1".format(log.path), "w") as fp:
fp.write("123")
with open("{0}.bad-file".format(log.path), "w") as fp:
fp.write("123")
self.assertEqual([1], log.listLogs())
def test_listLogsIgnoresZeroSuffixedFiles(self):
"""
L{LogFile.listLogs} ignores log files which rotated suffix is 0.
"""
log = logfile.LogFile(self.name, self.dir)
self.addCleanup(log.close)
for i in range(0, 3):
with open("{0}.{1}".format(log.path, i), "w") as fp:
fp.write("123")
self.assertEqual([1, 2], log.listLogs())
class RiggedDailyLogFile(logfile.DailyLogFile):
_clock = 0.0
def _openFile(self):
logfile.DailyLogFile._openFile(self)
# rig the date to match _clock, not mtime
self.lastDate = self.toDate()
def toDate(self, *args):
if args:
return time.gmtime(*args)[:3]
return time.gmtime(self._clock)[:3]
class DailyLogFileTests(unittest.TestCase):
"""
Test rotating log file.
"""
def setUp(self):
self.dir = self.mktemp()
os.makedirs(self.dir)
self.name = "testdaily.log"
self.path = os.path.join(self.dir, self.name)
def test_writing(self):
"""
A daily log file can be written to like an ordinary log file.
"""
with contextlib.closing(RiggedDailyLogFile(self.name, self.dir)) as log:
log.write("123")
log.write("456")
log.flush()
log.write("7890")
with open(self.path) as f:
self.assertEqual(f.read(), "1234567890")
def test_rotation(self):
"""
Daily log files rotate daily.
"""
log = RiggedDailyLogFile(self.name, self.dir)
self.addCleanup(log.close)
days = [(self.path + '.' + log.suffix(day * 86400)) for day in range(3)]
# test automatic rotation
log._clock = 0.0 # 1970/01/01 00:00.00
log.write("123")
log._clock = 43200 # 1970/01/01 12:00.00
log.write("4567890")
log._clock = 86400 # 1970/01/02 00:00.00
log.write("1" * 11)
self.assertTrue(os.path.exists(days[0]))
self.assertFalse(os.path.exists(days[1]))
log._clock = 172800 # 1970/01/03 00:00.00
log.write('')
self.assertTrue(os.path.exists(days[0]))
self.assertTrue(os.path.exists(days[1]))
self.assertFalse(os.path.exists(days[2]))
log._clock = 259199 # 1970/01/03 23:59.59
log.write("3")
self.assertFalse(os.path.exists(days[2]))
def test_getLog(self):
"""
Test retrieving log files with L{DailyLogFile.getLog}.
"""
data = ["1\n", "2\n", "3\n"]
log = RiggedDailyLogFile(self.name, self.dir)
self.addCleanup(log.close)
for d in data:
log.write(d)
log.flush()
# This returns the current log file.
r = log.getLog(0.0)
self.addCleanup(r.close)
self.assertEqual(data, r.readLines())
# We can't get this log, it doesn't exist yet.
self.assertRaises(ValueError, log.getLog, 86400)
log._clock = 86401 # New day
r.close()
log.rotate()
r = log.getLog(0) # We get the previous log
self.addCleanup(r.close)
self.assertEqual(data, r.readLines())
def test_rotateAlreadyExists(self):
"""
L{DailyLogFile.rotate} doesn't do anything if they new log file already
exists on the disk.
"""
log = RiggedDailyLogFile(self.name, self.dir)
self.addCleanup(log.close)
# Build a new file with the same name as the file which would be created
# if the log file is to be rotated.
newFilePath = "{0}.{1}".format(log.path, log.suffix(log.lastDate))
with open(newFilePath, "w") as fp:
fp.write("123")
previousFile = log._file
log.rotate()
self.assertEqual(previousFile, log._file)
def test_rotatePermissionDirectoryNotOk(self):
"""
L{DailyLogFile.rotate} doesn't do anything if the directory containing
the log files can't be written to.
"""
log = logfile.DailyLogFile(self.name, self.dir)
self.addCleanup(log.close)
os.chmod(log.directory, 0o444)
# Restore permissions so tests can be cleaned up.
self.addCleanup(os.chmod, log.directory, 0o755)
previousFile = log._file
log.rotate()
self.assertEqual(previousFile, log._file)
if runtime.platform.isWindows():
test_rotatePermissionDirectoryNotOk.skip = (
"Making read-only directories on Windows is too complex for this "
"test to reasonably do.")
def test_rotatePermissionFileNotOk(self):
"""
L{DailyLogFile.rotate} doesn't do anything if the log file can't be
written to.
"""
log = logfile.DailyLogFile(self.name, self.dir)
self.addCleanup(log.close)
os.chmod(log.path, 0o444)
previousFile = log._file
log.rotate()
self.assertEqual(previousFile, log._file)
def test_toDate(self):
"""
Test that L{DailyLogFile.toDate} converts its timestamp argument to a
time tuple (year, month, day).
"""
log = logfile.DailyLogFile(self.name, self.dir)
self.addCleanup(log.close)
timestamp = time.mktime((2000, 1, 1, 0, 0, 0, 0, 0, 0))
self.assertEqual((2000, 1, 1), log.toDate(timestamp))
def test_toDateDefaultToday(self):
"""
Test that L{DailyLogFile.toDate} returns today's date by default.
By mocking L{time.localtime}, we ensure that L{DailyLogFile.toDate}
returns the first 3 values of L{time.localtime} which is the current
date.
Note that we don't compare the *real* result of L{DailyLogFile.toDate}
to the *real* current date, as there's a slight possibility that the
date changes between the 2 function calls.
"""
def mock_localtime(*args):
self.assertEqual((), args)
return list(range(0, 9))
log = logfile.DailyLogFile(self.name, self.dir)
self.addCleanup(log.close)
self.patch(time, "localtime", mock_localtime)
logDate = log.toDate()
self.assertEqual([0, 1, 2], logDate)
def test_toDateUsesArgumentsToMakeADate(self):
"""
Test that L{DailyLogFile.toDate} uses its arguments to create a new
date.
"""
log = logfile.DailyLogFile(self.name, self.dir)
self.addCleanup(log.close)
date = (2014, 10, 22)
seconds = time.mktime(date + (0,)*6)
logDate = log.toDate(seconds)
self.assertEqual(date, logDate)
@@ -0,0 +1,474 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Test case for L{twisted.protocols.loopback}.
"""
from __future__ import division, absolute_import
from zope.interface import implementer
from twisted.python.compat import intToBytes
from twisted.trial import unittest
from twisted.protocols import basic, loopback
from twisted.internet import defer
from twisted.internet.protocol import Protocol
from twisted.internet.defer import Deferred
from twisted.internet.interfaces import IAddress, IPushProducer, IPullProducer
from twisted.internet import reactor, interfaces
class SimpleProtocol(basic.LineReceiver):
def __init__(self):
self.conn = defer.Deferred()
self.lines = []
self.connLost = []
def connectionMade(self):
self.conn.callback(None)
def lineReceived(self, line):
self.lines.append(line)
def connectionLost(self, reason):
self.connLost.append(reason)
class DoomProtocol(SimpleProtocol):
i = 0
def lineReceived(self, line):
self.i += 1
if self.i < 4:
# by this point we should have connection closed,
# but just in case we didn't we won't ever send 'Hello 4'
self.sendLine(b"Hello " + intToBytes(self.i))
SimpleProtocol.lineReceived(self, line)
if self.lines[-1] == b"Hello 3":
self.transport.loseConnection()
class LoopbackTestCaseMixin:
def testRegularFunction(self):
s = SimpleProtocol()
c = SimpleProtocol()
def sendALine(result):
s.sendLine(b"THIS IS LINE ONE!")
s.transport.loseConnection()
s.conn.addCallback(sendALine)
def check(ignored):
self.assertEqual(c.lines, [b"THIS IS LINE ONE!"])
self.assertEqual(len(s.connLost), 1)
self.assertEqual(len(c.connLost), 1)
d = defer.maybeDeferred(self.loopbackFunc, s, c)
d.addCallback(check)
return d
def testSneakyHiddenDoom(self):
s = DoomProtocol()
c = DoomProtocol()
def sendALine(result):
s.sendLine(b"DOOM LINE")
s.conn.addCallback(sendALine)
def check(ignored):
self.assertEqual(s.lines, [b'Hello 1', b'Hello 2', b'Hello 3'])
self.assertEqual(
c.lines, [b'DOOM LINE', b'Hello 1', b'Hello 2', b'Hello 3'])
self.assertEqual(len(s.connLost), 1)
self.assertEqual(len(c.connLost), 1)
d = defer.maybeDeferred(self.loopbackFunc, s, c)
d.addCallback(check)
return d
class LoopbackAsyncTests(LoopbackTestCaseMixin, unittest.TestCase):
loopbackFunc = staticmethod(loopback.loopbackAsync)
def test_makeConnection(self):
"""
Test that the client and server protocol both have makeConnection
invoked on them by loopbackAsync.
"""
class TestProtocol(Protocol):
transport = None
def makeConnection(self, transport):
self.transport = transport
server = TestProtocol()
client = TestProtocol()
loopback.loopbackAsync(server, client)
self.assertIsNotNone(client.transport)
self.assertIsNotNone(server.transport)
def _hostpeertest(self, get, testServer):
"""
Test one of the permutations of client/server host/peer.
"""
class TestProtocol(Protocol):
def makeConnection(self, transport):
Protocol.makeConnection(self, transport)
self.onConnection.callback(transport)
if testServer:
server = TestProtocol()
d = server.onConnection = Deferred()
client = Protocol()
else:
server = Protocol()
client = TestProtocol()
d = client.onConnection = Deferred()
loopback.loopbackAsync(server, client)
def connected(transport):
host = getattr(transport, get)()
self.assertTrue(IAddress.providedBy(host))
return d.addCallback(connected)
def test_serverHost(self):
"""
Test that the server gets a transport with a properly functioning
implementation of L{ITransport.getHost}.
"""
return self._hostpeertest("getHost", True)
def test_serverPeer(self):
"""
Like C{test_serverHost} but for L{ITransport.getPeer}
"""
return self._hostpeertest("getPeer", True)
def test_clientHost(self, get="getHost"):
"""
Test that the client gets a transport with a properly functioning
implementation of L{ITransport.getHost}.
"""
return self._hostpeertest("getHost", False)
def test_clientPeer(self):
"""
Like C{test_clientHost} but for L{ITransport.getPeer}.
"""
return self._hostpeertest("getPeer", False)
def _greetingtest(self, write, testServer):
"""
Test one of the permutations of write/writeSequence client/server.
@param write: The name of the method to test, C{"write"} or
C{"writeSequence"}.
"""
class GreeteeProtocol(Protocol):
bytes = b""
def dataReceived(self, bytes):
self.bytes += bytes
if self.bytes == b"bytes":
self.received.callback(None)
class GreeterProtocol(Protocol):
def connectionMade(self):
if write == "write":
self.transport.write(b"bytes")
else:
self.transport.writeSequence([b"byt", b"es"])
if testServer:
server = GreeterProtocol()
client = GreeteeProtocol()
d = client.received = Deferred()
else:
server = GreeteeProtocol()
d = server.received = Deferred()
client = GreeterProtocol()
loopback.loopbackAsync(server, client)
return d
def test_clientGreeting(self):
"""
Test that on a connection where the client speaks first, the server
receives the bytes sent by the client.
"""
return self._greetingtest("write", False)
def test_clientGreetingSequence(self):
"""
Like C{test_clientGreeting}, but use C{writeSequence} instead of
C{write} to issue the greeting.
"""
return self._greetingtest("writeSequence", False)
def test_serverGreeting(self, write="write"):
"""
Test that on a connection where the server speaks first, the client
receives the bytes sent by the server.
"""
return self._greetingtest("write", True)
def test_serverGreetingSequence(self):
"""
Like C{test_serverGreeting}, but use C{writeSequence} instead of
C{write} to issue the greeting.
"""
return self._greetingtest("writeSequence", True)
def _producertest(self, producerClass):
toProduce = list(map(intToBytes, range(0, 10)))
class ProducingProtocol(Protocol):
def connectionMade(self):
self.producer = producerClass(list(toProduce))
self.producer.start(self.transport)
class ReceivingProtocol(Protocol):
bytes = b""
def dataReceived(self, data):
self.bytes += data
if self.bytes == b''.join(toProduce):
self.received.callback((client, server))
server = ProducingProtocol()
client = ReceivingProtocol()
client.received = Deferred()
loopback.loopbackAsync(server, client)
return client.received
def test_pushProducer(self):
"""
Test a push producer registered against a loopback transport.
"""
@implementer(IPushProducer)
class PushProducer(object):
resumed = False
def __init__(self, toProduce):
self.toProduce = toProduce
def resumeProducing(self):
self.resumed = True
def start(self, consumer):
self.consumer = consumer
consumer.registerProducer(self, True)
self._produceAndSchedule()
def _produceAndSchedule(self):
if self.toProduce:
self.consumer.write(self.toProduce.pop(0))
reactor.callLater(0, self._produceAndSchedule)
else:
self.consumer.unregisterProducer()
d = self._producertest(PushProducer)
def finished(results):
(client, server) = results
self.assertFalse(
server.producer.resumed,
"Streaming producer should not have been resumed.")
d.addCallback(finished)
return d
def test_pullProducer(self):
"""
Test a pull producer registered against a loopback transport.
"""
@implementer(IPullProducer)
class PullProducer(object):
def __init__(self, toProduce):
self.toProduce = toProduce
def start(self, consumer):
self.consumer = consumer
self.consumer.registerProducer(self, False)
def resumeProducing(self):
self.consumer.write(self.toProduce.pop(0))
if not self.toProduce:
self.consumer.unregisterProducer()
return self._producertest(PullProducer)
def test_writeNotReentrant(self):
"""
L{loopback.loopbackAsync} does not call a protocol's C{dataReceived}
method while that protocol's transport's C{write} method is higher up
on the stack.
"""
class Server(Protocol):
def dataReceived(self, bytes):
self.transport.write(b"bytes")
class Client(Protocol):
ready = False
def connectionMade(self):
reactor.callLater(0, self.go)
def go(self):
self.transport.write(b"foo")
self.ready = True
def dataReceived(self, bytes):
self.wasReady = self.ready
self.transport.loseConnection()
server = Server()
client = Client()
d = loopback.loopbackAsync(client, server)
def cbFinished(ignored):
self.assertTrue(client.wasReady)
d.addCallback(cbFinished)
return d
def test_pumpPolicy(self):
"""
The callable passed as the value for the C{pumpPolicy} parameter to
L{loopbackAsync} is called with a L{_LoopbackQueue} of pending bytes
and a protocol to which they should be delivered.
"""
pumpCalls = []
def dummyPolicy(queue, target):
bytes = []
while queue:
bytes.append(queue.get())
pumpCalls.append((target, bytes))
client = Protocol()
server = Protocol()
finished = loopback.loopbackAsync(server, client, dummyPolicy)
self.assertEqual(pumpCalls, [])
client.transport.write(b"foo")
client.transport.write(b"bar")
server.transport.write(b"baz")
server.transport.write(b"quux")
server.transport.loseConnection()
def cbComplete(ignored):
self.assertEqual(
pumpCalls,
# The order here is somewhat arbitrary. The implementation
# happens to always deliver data to the client first.
[(client, [b"baz", b"quux", None]),
(server, [b"foo", b"bar"])])
finished.addCallback(cbComplete)
return finished
def test_identityPumpPolicy(self):
"""
L{identityPumpPolicy} is a pump policy which calls the target's
C{dataReceived} method one for each string in the queue passed to it.
"""
bytes = []
client = Protocol()
client.dataReceived = bytes.append
queue = loopback._LoopbackQueue()
queue.put(b"foo")
queue.put(b"bar")
queue.put(None)
loopback.identityPumpPolicy(queue, client)
self.assertEqual(bytes, [b"foo", b"bar"])
def test_collapsingPumpPolicy(self):
"""
L{collapsingPumpPolicy} is a pump policy which calls the target's
C{dataReceived} only once with all of the strings in the queue passed
to it joined together.
"""
bytes = []
client = Protocol()
client.dataReceived = bytes.append
queue = loopback._LoopbackQueue()
queue.put(b"foo")
queue.put(b"bar")
queue.put(None)
loopback.collapsingPumpPolicy(queue, client)
self.assertEqual(bytes, [b"foobar"])
class LoopbackTCPTests(LoopbackTestCaseMixin, unittest.TestCase):
loopbackFunc = staticmethod(loopback.loopbackTCP)
class LoopbackUNIXTests(LoopbackTestCaseMixin, unittest.TestCase):
loopbackFunc = staticmethod(loopback.loopbackUNIX)
if interfaces.IReactorUNIX(reactor, None) is None:
skip = "Current reactor does not support UNIX sockets"
class LoopbackRelayTest(unittest.TestCase):
"""
Test for L{twisted.protocols.loopback.LoopbackRelay}
"""
class Receiver(Protocol):
"""
Simple Receiver class used for testing LoopbackRelay
"""
data = b''
def dataReceived(self, data):
"Accumulate received data for verification"
self.data += data
def test_write(self):
"Test to verify that the write function works as expected"
receiver = self.Receiver()
relay = loopback.LoopbackRelay(receiver)
relay.write(b'abc')
relay.write(b'def')
self.assertEqual(receiver.data, b'')
relay.clearBuffer()
self.assertEqual(receiver.data, b'abcdef')
def test_writeSequence(self):
"Test to verify that the writeSequence function works as expected"
receiver = self.Receiver()
relay = loopback.LoopbackRelay(receiver)
relay.writeSequence(
[b'The ', b'quick ', b'brown ', b'fox '])
relay.writeSequence(
[b'jumps ', b'over ', b'the lazy dog'])
self.assertEqual(receiver.data, b'')
relay.clearBuffer()
self.assertEqual(
receiver.data, b'The quick brown fox jumps over the lazy dog')
@@ -0,0 +1,73 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Test that twisted scripts can be invoked as modules.
"""
from __future__ import division, absolute_import
import sys
from twisted.application.twist._options import TwistOptions
from twisted.scripts import trial
from twisted.internet import defer, reactor
from twisted.python.compat import NativeStringIO as StringIO
from twisted.test.test_process import Accumulator
from twisted.trial.unittest import TestCase
class MainTests(TestCase):
"""Test that twisted scripts can be invoked as modules."""
def test_twisted(self):
"""Invoking python -m twisted should execute twist."""
cmd = sys.executable
p = Accumulator()
d = p.endedDeferred = defer.Deferred()
reactor.spawnProcess(p, cmd, [cmd, '-m', 'twisted', '--help'], env=None)
p.transport.closeStdin()
# Fix up our sys args to match the command we issued
from twisted import __main__
self.patch(sys, 'argv', [__main__.__file__, '--help'])
def processEnded(ign):
f = p.outF
output = f.getvalue().replace(b'\r\n', b'\n')
options = TwistOptions()
message = '{}\n'.format(options).encode('utf-8')
self.assertEqual(output, message)
return d.addCallback(processEnded)
def test_trial(self):
"""Invoking python -m twisted.trial should execute trial."""
cmd = sys.executable
p = Accumulator()
d = p.endedDeferred = defer.Deferred()
reactor.spawnProcess(p, cmd, [cmd, '-m', 'twisted.trial', '--help'], env=None)
p.transport.closeStdin()
# Fix up our sys args to match the command we issued
from twisted.trial import __main__
self.patch(sys, 'argv', [__main__.__file__, '--help'])
def processEnded(ign):
f = p.outF
output = f.getvalue().replace(b'\r\n', b'\n')
options = trial.Options()
message = '{}\n'.format(options).encode('utf-8')
self.assertEqual(output, message)
return d.addCallback(processEnded)
def test_twisted_import(self):
"""Importing twisted.__main__ does not execute twist."""
output = StringIO()
monkey = self.patch(sys, 'stdout', output)
import twisted.__main__
self.assertTrue(twisted.__main__) # Appease pyflakes
monkey.restore()
self.assertEqual(output.getvalue(), "")
@@ -0,0 +1,237 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Test cases for twisted.protocols package.
"""
from twisted.trial import unittest
from twisted.protocols import wire, portforward
from twisted.python.compat import iterbytes
from twisted.internet import reactor, defer, address, protocol
from twisted.test import proto_helpers
class WireTests(unittest.TestCase):
"""
Test wire protocols.
"""
def test_echo(self):
"""
Test wire.Echo protocol: send some data and check it send it back.
"""
t = proto_helpers.StringTransport()
a = wire.Echo()
a.makeConnection(t)
a.dataReceived(b"hello")
a.dataReceived(b"world")
a.dataReceived(b"how")
a.dataReceived(b"are")
a.dataReceived(b"you")
self.assertEqual(t.value(), b"helloworldhowareyou")
def test_who(self):
"""
Test wire.Who protocol.
"""
t = proto_helpers.StringTransport()
a = wire.Who()
a.makeConnection(t)
self.assertEqual(t.value(), b"root\r\n")
def test_QOTD(self):
"""
Test wire.QOTD protocol.
"""
t = proto_helpers.StringTransport()
a = wire.QOTD()
a.makeConnection(t)
self.assertEqual(t.value(),
b"An apple a day keeps the doctor away.\r\n")
def test_discard(self):
"""
Test wire.Discard protocol.
"""
t = proto_helpers.StringTransport()
a = wire.Discard()
a.makeConnection(t)
a.dataReceived(b"hello")
a.dataReceived(b"world")
a.dataReceived(b"how")
a.dataReceived(b"are")
a.dataReceived(b"you")
self.assertEqual(t.value(), b"")
class TestableProxyClientFactory(portforward.ProxyClientFactory):
"""
Test proxy client factory that keeps the last created protocol instance.
@ivar protoInstance: the last instance of the protocol.
@type protoInstance: L{portforward.ProxyClient}
"""
def buildProtocol(self, addr):
"""
Create the protocol instance and keeps track of it.
"""
proto = portforward.ProxyClientFactory.buildProtocol(self, addr)
self.protoInstance = proto
return proto
class TestableProxyFactory(portforward.ProxyFactory):
"""
Test proxy factory that keeps the last created protocol instance.
@ivar protoInstance: the last instance of the protocol.
@type protoInstance: L{portforward.ProxyServer}
@ivar clientFactoryInstance: client factory used by C{protoInstance} to
create forward connections.
@type clientFactoryInstance: L{TestableProxyClientFactory}
"""
def buildProtocol(self, addr):
"""
Create the protocol instance, keeps track of it, and makes it use
C{clientFactoryInstance} as client factory.
"""
proto = portforward.ProxyFactory.buildProtocol(self, addr)
self.clientFactoryInstance = TestableProxyClientFactory()
# Force the use of this specific instance
proto.clientProtocolFactory = lambda: self.clientFactoryInstance
self.protoInstance = proto
return proto
class PortforwardingTests(unittest.TestCase):
"""
Test port forwarding.
"""
def setUp(self):
self.serverProtocol = wire.Echo()
self.clientProtocol = protocol.Protocol()
self.openPorts = []
def tearDown(self):
try:
self.proxyServerFactory.protoInstance.transport.loseConnection()
except AttributeError:
pass
try:
pi = self.proxyServerFactory.clientFactoryInstance.protoInstance
pi.transport.loseConnection()
except AttributeError:
pass
try:
self.clientProtocol.transport.loseConnection()
except AttributeError:
pass
try:
self.serverProtocol.transport.loseConnection()
except AttributeError:
pass
return defer.gatherResults(
[defer.maybeDeferred(p.stopListening) for p in self.openPorts])
def test_portforward(self):
"""
Test port forwarding through Echo protocol.
"""
realServerFactory = protocol.ServerFactory()
realServerFactory.protocol = lambda: self.serverProtocol
realServerPort = reactor.listenTCP(0, realServerFactory,
interface='127.0.0.1')
self.openPorts.append(realServerPort)
self.proxyServerFactory = TestableProxyFactory('127.0.0.1',
realServerPort.getHost().port)
proxyServerPort = reactor.listenTCP(0, self.proxyServerFactory,
interface='127.0.0.1')
self.openPorts.append(proxyServerPort)
nBytes = 1000
received = []
d = defer.Deferred()
def testDataReceived(data):
received.extend(iterbytes(data))
if len(received) >= nBytes:
self.assertEqual(b''.join(received), b'x' * nBytes)
d.callback(None)
self.clientProtocol.dataReceived = testDataReceived
def testConnectionMade():
self.clientProtocol.transport.write(b'x' * nBytes)
self.clientProtocol.connectionMade = testConnectionMade
clientFactory = protocol.ClientFactory()
clientFactory.protocol = lambda: self.clientProtocol
reactor.connectTCP(
'127.0.0.1', proxyServerPort.getHost().port, clientFactory)
return d
def test_registerProducers(self):
"""
The proxy client registers itself as a producer of the proxy server and
vice versa.
"""
# create a ProxyServer instance
addr = address.IPv4Address('TCP', '127.0.0.1', 0)
server = portforward.ProxyFactory('127.0.0.1', 0).buildProtocol(addr)
# set the reactor for this test
reactor = proto_helpers.MemoryReactor()
server.reactor = reactor
# make the connection
serverTransport = proto_helpers.StringTransport()
server.makeConnection(serverTransport)
# check that the ProxyClientFactory is connecting to the backend
self.assertEqual(len(reactor.tcpClients), 1)
# get the factory instance and check it's the one we expect
host, port, clientFactory, timeout, _ = reactor.tcpClients[0]
self.assertIsInstance(clientFactory, portforward.ProxyClientFactory)
# Connect it
client = clientFactory.buildProtocol(addr)
clientTransport = proto_helpers.StringTransport()
client.makeConnection(clientTransport)
# check that the producers are registered
self.assertIs(clientTransport.producer, serverTransport)
self.assertIs(serverTransport.producer, clientTransport)
# check the streaming attribute in both transports
self.assertTrue(clientTransport.streaming)
self.assertTrue(serverTransport.streaming)
class StringTransportTests(unittest.TestCase):
"""
Test L{proto_helpers.StringTransport} helper behaviour.
"""
def test_noUnicode(self):
"""
Test that L{proto_helpers.StringTransport} doesn't accept unicode data.
"""
s = proto_helpers.StringTransport()
self.assertRaises(TypeError, s.write, u'foo')
@@ -0,0 +1,892 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Test cases for the L{twisted.python.reflect} module.
"""
from __future__ import division, absolute_import
import os
import weakref
from collections import deque
from twisted.python.compat import _PY3
from twisted.trial import unittest
from twisted.trial.unittest import SynchronousTestCase as TestCase
from twisted.python import reflect
from twisted.python.reflect import (
accumulateMethods, prefixedMethods, prefixedMethodNames,
addMethodNamesToDict, fullyQualifiedName)
class Base(object):
"""
A no-op class which can be used to verify the behavior of
method-discovering APIs.
"""
def method(self):
"""
A no-op method which can be discovered.
"""
class Sub(Base):
"""
A subclass of a class with a method which can be discovered.
"""
class Separate(object):
"""
A no-op class with methods with differing prefixes.
"""
def good_method(self):
"""
A no-op method which a matching prefix to be discovered.
"""
def bad_method(self):
"""
A no-op method with a mismatched prefix to not be discovered.
"""
class AccumulateMethodsTests(TestCase):
"""
Tests for L{accumulateMethods} which finds methods on a class hierarchy and
adds them to a dictionary.
"""
def test_ownClass(self):
"""
If x is and instance of Base and Base defines a method named method,
L{accumulateMethods} adds an item to the given dictionary with
C{"method"} as the key and a bound method object for Base.method value.
"""
x = Base()
output = {}
accumulateMethods(x, output)
self.assertEqual({"method": x.method}, output)
def test_baseClass(self):
"""
If x is an instance of Sub and Sub is a subclass of Base and Base
defines a method named method, L{accumulateMethods} adds an item to the
given dictionary with C{"method"} as the key and a bound method object
for Base.method as the value.
"""
x = Sub()
output = {}
accumulateMethods(x, output)
self.assertEqual({"method": x.method}, output)
def test_prefix(self):
"""
If a prefix is given, L{accumulateMethods} limits its results to
methods beginning with that prefix. Keys in the resulting dictionary
also have the prefix removed from them.
"""
x = Separate()
output = {}
accumulateMethods(x, output, 'good_')
self.assertEqual({'method': x.good_method}, output)
class PrefixedMethodsTests(TestCase):
"""
Tests for L{prefixedMethods} which finds methods on a class hierarchy and
adds them to a dictionary.
"""
def test_onlyObject(self):
"""
L{prefixedMethods} returns a list of the methods discovered on an
object.
"""
x = Base()
output = prefixedMethods(x)
self.assertEqual([x.method], output)
def test_prefix(self):
"""
If a prefix is given, L{prefixedMethods} returns only methods named
with that prefix.
"""
x = Separate()
output = prefixedMethods(x, 'good_')
self.assertEqual([x.good_method], output)
class PrefixedMethodNamesTests(TestCase):
"""
Tests for L{prefixedMethodNames}.
"""
def test_method(self):
"""
L{prefixedMethodNames} returns a list including methods with the given
prefix defined on the class passed to it.
"""
self.assertEqual(["method"], prefixedMethodNames(Separate, "good_"))
def test_inheritedMethod(self):
"""
L{prefixedMethodNames} returns a list included methods with the given
prefix defined on base classes of the class passed to it.
"""
class Child(Separate):
pass
self.assertEqual(["method"], prefixedMethodNames(Child, "good_"))
class AddMethodNamesToDictTests(TestCase):
"""
Tests for L{addMethodNamesToDict}.
"""
def test_baseClass(self):
"""
If C{baseClass} is passed to L{addMethodNamesToDict}, only methods which
are a subclass of C{baseClass} are added to the result dictionary.
"""
class Alternate(object):
pass
class Child(Separate, Alternate):
def good_alternate(self):
pass
result = {}
addMethodNamesToDict(Child, result, 'good_', Alternate)
self.assertEqual({'alternate': 1}, result)
class Summer(object):
"""
A class we look up as part of the LookupsTests.
"""
def reallySet(self):
"""
Do something.
"""
class LookupsTests(TestCase):
"""
Tests for L{namedClass}, L{namedModule}, and L{namedAny}.
"""
def test_namedClassLookup(self):
"""
L{namedClass} should return the class object for the name it is passed.
"""
self.assertIs(
reflect.namedClass("twisted.test.test_reflect.Summer"),
Summer)
def test_namedModuleLookup(self):
"""
L{namedModule} should return the module object for the name it is
passed.
"""
from twisted.python import monkey
self.assertIs(
reflect.namedModule("twisted.python.monkey"), monkey)
def test_namedAnyPackageLookup(self):
"""
L{namedAny} should return the package object for the name it is passed.
"""
import twisted.python
self.assertIs(
reflect.namedAny("twisted.python"), twisted.python)
def test_namedAnyModuleLookup(self):
"""
L{namedAny} should return the module object for the name it is passed.
"""
from twisted.python import monkey
self.assertIs(
reflect.namedAny("twisted.python.monkey"), monkey)
def test_namedAnyClassLookup(self):
"""
L{namedAny} should return the class object for the name it is passed.
"""
self.assertIs(
reflect.namedAny("twisted.test.test_reflect.Summer"),
Summer)
def test_namedAnyAttributeLookup(self):
"""
L{namedAny} should return the object an attribute of a non-module,
non-package object is bound to for the name it is passed.
"""
# Note - not assertIs because unbound method lookup creates a new
# object every time. This is a foolishness of Python's object
# implementation, not a bug in Twisted.
self.assertEqual(
reflect.namedAny(
"twisted.test.test_reflect.Summer.reallySet"),
Summer.reallySet)
def test_namedAnySecondAttributeLookup(self):
"""
L{namedAny} should return the object an attribute of an object which
itself was an attribute of a non-module, non-package object is bound to
for the name it is passed.
"""
self.assertIs(
reflect.namedAny(
"twisted.test.test_reflect."
"Summer.reallySet.__doc__"),
Summer.reallySet.__doc__)
def test_importExceptions(self):
"""
Exceptions raised by modules which L{namedAny} causes to be imported
should pass through L{namedAny} to the caller.
"""
self.assertRaises(
ZeroDivisionError,
reflect.namedAny, "twisted.test.reflect_helper_ZDE")
# Make sure that there is post-failed-import cleanup
self.assertRaises(
ZeroDivisionError,
reflect.namedAny, "twisted.test.reflect_helper_ZDE")
self.assertRaises(
ValueError,
reflect.namedAny, "twisted.test.reflect_helper_VE")
# Modules which themselves raise ImportError when imported should
# result in an ImportError
self.assertRaises(
ImportError,
reflect.namedAny, "twisted.test.reflect_helper_IE")
def test_attributeExceptions(self):
"""
If segments on the end of a fully-qualified Python name represents
attributes which aren't actually present on the object represented by
the earlier segments, L{namedAny} should raise an L{AttributeError}.
"""
self.assertRaises(
AttributeError,
reflect.namedAny, "twisted.nosuchmoduleintheworld")
# ImportError behaves somewhat differently between "import
# extant.nonextant" and "import extant.nonextant.nonextant", so test
# the latter as well.
self.assertRaises(
AttributeError,
reflect.namedAny, "twisted.nosuch.modulein.theworld")
self.assertRaises(
AttributeError,
reflect.namedAny,
"twisted.test.test_reflect.Summer.nosuchattribute")
def test_invalidNames(self):
"""
Passing a name which isn't a fully-qualified Python name to L{namedAny}
should result in one of the following exceptions:
- L{InvalidName}: the name is not a dot-separated list of Python
objects
- L{ObjectNotFound}: the object doesn't exist
- L{ModuleNotFound}: the object doesn't exist and there is only one
component in the name
"""
err = self.assertRaises(reflect.ModuleNotFound, reflect.namedAny,
'nosuchmoduleintheworld')
self.assertEqual(str(err), "No module named 'nosuchmoduleintheworld'")
# This is a dot-separated list, but it isn't valid!
err = self.assertRaises(reflect.ObjectNotFound, reflect.namedAny,
"@#$@(#.!@(#!@#")
self.assertEqual(str(err), "'@#$@(#.!@(#!@#' does not name an object")
err = self.assertRaises(reflect.ObjectNotFound, reflect.namedAny,
"tcelfer.nohtyp.detsiwt")
self.assertEqual(
str(err),
"'tcelfer.nohtyp.detsiwt' does not name an object")
err = self.assertRaises(reflect.InvalidName, reflect.namedAny, '')
self.assertEqual(str(err), 'Empty module name')
for invalidName in ['.twisted', 'twisted.', 'twisted..python']:
err = self.assertRaises(
reflect.InvalidName, reflect.namedAny, invalidName)
self.assertEqual(
str(err),
"name must be a string giving a '.'-separated list of Python "
"identifiers, not %r" % (invalidName,))
def test_requireModuleImportError(self):
"""
When module import fails with ImportError it returns the specified
default value.
"""
for name in ['nosuchmtopodule', 'no.such.module']:
default = object()
result = reflect.requireModule(name, default=default)
self.assertIs(result, default)
def test_requireModuleDefaultNone(self):
"""
When module import fails it returns L{None} by default.
"""
result = reflect.requireModule('no.such.module')
self.assertIsNone(result)
def test_requireModuleRequestedImport(self):
"""
When module import succeed it returns the module and not the default
value.
"""
from twisted.python import monkey
default = object()
self.assertIs(
reflect.requireModule('twisted.python.monkey', default=default),
monkey,
)
class Breakable(object):
breakRepr = False
breakStr = False
def __str__(self):
if self.breakStr:
raise RuntimeError("str!")
else:
return '<Breakable>'
def __repr__(self):
if self.breakRepr:
raise RuntimeError("repr!")
else:
return 'Breakable()'
class BrokenType(Breakable, type):
breakName = False
def get___name__(self):
if self.breakName:
raise RuntimeError("no name")
return 'BrokenType'
__name__ = property(get___name__)
BTBase = BrokenType('BTBase', (Breakable,),
{"breakRepr": True,
"breakStr": True})
class NoClassAttr(Breakable):
__class__ = property(lambda x: x.not_class)
class SafeReprTests(TestCase):
"""
Tests for L{reflect.safe_repr} function.
"""
def test_workingRepr(self):
"""
L{reflect.safe_repr} produces the same output as C{repr} on a working
object.
"""
xs = ([1, 2, 3], b'a')
self.assertEqual(list(map(reflect.safe_repr, xs)), list(map(repr, xs)))
def test_brokenRepr(self):
"""
L{reflect.safe_repr} returns a string with class name, address, and
traceback when the repr call failed.
"""
b = Breakable()
b.breakRepr = True
bRepr = reflect.safe_repr(b)
self.assertIn("Breakable instance at 0x", bRepr)
# Check that the file is in the repr, but without the extension as it
# can be .py/.pyc
self.assertIn(os.path.splitext(__file__)[0], bRepr)
self.assertIn("RuntimeError: repr!", bRepr)
def test_brokenStr(self):
"""
L{reflect.safe_repr} isn't affected by a broken C{__str__} method.
"""
b = Breakable()
b.breakStr = True
self.assertEqual(reflect.safe_repr(b), repr(b))
def test_brokenClassRepr(self):
class X(BTBase):
breakRepr = True
reflect.safe_repr(X)
reflect.safe_repr(X())
def test_brokenReprIncludesID(self):
"""
C{id} is used to print the ID of the object in case of an error.
L{safe_repr} includes a traceback after a newline, so we only check
against the first line of the repr.
"""
class X(BTBase):
breakRepr = True
xRepr = reflect.safe_repr(X)
xReprExpected = ('<BrokenType instance at 0x%x with repr error:'
% (id(X),))
self.assertEqual(xReprExpected, xRepr.split('\n')[0])
def test_brokenClassStr(self):
class X(BTBase):
breakStr = True
reflect.safe_repr(X)
reflect.safe_repr(X())
def test_brokenClassAttribute(self):
"""
If an object raises an exception when accessing its C{__class__}
attribute, L{reflect.safe_repr} uses C{type} to retrieve the class
object.
"""
b = NoClassAttr()
b.breakRepr = True
bRepr = reflect.safe_repr(b)
self.assertIn("NoClassAttr instance at 0x", bRepr)
self.assertIn(os.path.splitext(__file__)[0], bRepr)
self.assertIn("RuntimeError: repr!", bRepr)
def test_brokenClassNameAttribute(self):
"""
If a class raises an exception when accessing its C{__name__} attribute
B{and} when calling its C{__str__} implementation, L{reflect.safe_repr}
returns 'BROKEN CLASS' instead of the class name.
"""
class X(BTBase):
breakName = True
xRepr = reflect.safe_repr(X())
self.assertIn("<BROKEN CLASS AT 0x", xRepr)
self.assertIn(os.path.splitext(__file__)[0], xRepr)
self.assertIn("RuntimeError: repr!", xRepr)
class SafeStrTests(TestCase):
"""
Tests for L{reflect.safe_str} function.
"""
def test_workingStr(self):
x = [1, 2, 3]
self.assertEqual(reflect.safe_str(x), str(x))
def test_brokenStr(self):
b = Breakable()
b.breakStr = True
reflect.safe_str(b)
def test_workingAscii(self):
"""
L{safe_str} for C{str} with ascii-only data should return the
value unchanged.
"""
x = 'a'
self.assertEqual(reflect.safe_str(x), 'a')
def test_workingUtf8_2(self):
"""
L{safe_str} for C{str} with utf-8 encoded data should return the
value unchanged.
"""
x = b't\xc3\xbcst'
self.assertEqual(reflect.safe_str(x), x)
def test_workingUtf8_3(self):
"""
L{safe_str} for C{bytes} with utf-8 encoded data should return
the value decoded into C{str}.
"""
x = b't\xc3\xbcst'
self.assertEqual(reflect.safe_str(x), x.decode('utf-8'))
if _PY3:
# TODO: after something like python.compat.nativeUtf8String is
# introduced, use that one for assertEqual. Then we can combine
# test_workingUtf8_* tests into one without needing _PY3.
# nativeUtf8String is needed for Python 3 anyway.
test_workingUtf8_2.skip = ("Skip Python 2 specific test for utf-8 str")
else:
test_workingUtf8_3.skip = (
"Skip Python 3 specific test for utf-8 bytes")
def test_brokenUtf8(self):
"""
Use str() for non-utf8 bytes: "b'non-utf8'"
"""
x = b'\xff'
xStr = reflect.safe_str(x)
self.assertEqual(xStr, str(x))
def test_brokenRepr(self):
b = Breakable()
b.breakRepr = True
reflect.safe_str(b)
def test_brokenClassStr(self):
class X(BTBase):
breakStr = True
reflect.safe_str(X)
reflect.safe_str(X())
def test_brokenClassRepr(self):
class X(BTBase):
breakRepr = True
reflect.safe_str(X)
reflect.safe_str(X())
def test_brokenClassAttribute(self):
"""
If an object raises an exception when accessing its C{__class__}
attribute, L{reflect.safe_str} uses C{type} to retrieve the class
object.
"""
b = NoClassAttr()
b.breakStr = True
bStr = reflect.safe_str(b)
self.assertIn("NoClassAttr instance at 0x", bStr)
self.assertIn(os.path.splitext(__file__)[0], bStr)
self.assertIn("RuntimeError: str!", bStr)
def test_brokenClassNameAttribute(self):
"""
If a class raises an exception when accessing its C{__name__} attribute
B{and} when calling its C{__str__} implementation, L{reflect.safe_str}
returns 'BROKEN CLASS' instead of the class name.
"""
class X(BTBase):
breakName = True
xStr = reflect.safe_str(X())
self.assertIn("<BROKEN CLASS AT 0x", xStr)
self.assertIn(os.path.splitext(__file__)[0], xStr)
self.assertIn("RuntimeError: str!", xStr)
class FilenameToModuleTests(TestCase):
"""
Test L{filenameToModuleName} detection.
"""
def setUp(self):
self.path = os.path.join(self.mktemp(), "fakepackage", "test")
os.makedirs(self.path)
with open(os.path.join(self.path, "__init__.py"), "w") as f:
f.write("")
with open(os.path.join(os.path.dirname(self.path), "__init__.py"),
"w") as f:
f.write("")
def test_directory(self):
"""
L{filenameToModuleName} returns the correct module (a package) given a
directory.
"""
module = reflect.filenameToModuleName(self.path)
self.assertEqual(module, 'fakepackage.test')
module = reflect.filenameToModuleName(self.path + os.path.sep)
self.assertEqual(module, 'fakepackage.test')
def test_file(self):
"""
L{filenameToModuleName} returns the correct module given the path to
its file.
"""
module = reflect.filenameToModuleName(
os.path.join(self.path, 'test_reflect.py'))
self.assertEqual(module, 'fakepackage.test.test_reflect')
def test_bytes(self):
"""
L{filenameToModuleName} returns the correct module given a C{bytes}
path to its file.
"""
module = reflect.filenameToModuleName(
os.path.join(self.path.encode("utf-8"), b'test_reflect.py'))
# Module names are always native string:
self.assertEqual(module, 'fakepackage.test.test_reflect')
class FullyQualifiedNameTests(TestCase):
"""
Test for L{fullyQualifiedName}.
"""
def _checkFullyQualifiedName(self, obj, expected):
"""
Helper to check that fully qualified name of C{obj} results to
C{expected}.
"""
self.assertEqual(fullyQualifiedName(obj), expected)
def test_package(self):
"""
L{fullyQualifiedName} returns the full name of a package and a
subpackage.
"""
import twisted
self._checkFullyQualifiedName(twisted, 'twisted')
import twisted.python
self._checkFullyQualifiedName(twisted.python, 'twisted.python')
def test_module(self):
"""
L{fullyQualifiedName} returns the name of a module inside a package.
"""
import twisted.python.compat
self._checkFullyQualifiedName(
twisted.python.compat, 'twisted.python.compat')
def test_class(self):
"""
L{fullyQualifiedName} returns the name of a class and its module.
"""
self._checkFullyQualifiedName(
FullyQualifiedNameTests,
'%s.FullyQualifiedNameTests' % (__name__,))
def test_function(self):
"""
L{fullyQualifiedName} returns the name of a function inside its module.
"""
self._checkFullyQualifiedName(
fullyQualifiedName, "twisted.python.reflect.fullyQualifiedName")
def test_boundMethod(self):
"""
L{fullyQualifiedName} returns the name of a bound method inside its
class and its module.
"""
self._checkFullyQualifiedName(
self.test_boundMethod,
"%s.%s.test_boundMethod" % (__name__, self.__class__.__name__))
def test_unboundMethod(self):
"""
L{fullyQualifiedName} returns the name of an unbound method inside its
class and its module.
"""
self._checkFullyQualifiedName(
self.__class__.test_unboundMethod,
"%s.%s.test_unboundMethod" % (__name__, self.__class__.__name__))
class ObjectGrepTests(unittest.TestCase):
if _PY3:
# This is to be removed when fixing #6986
skip = "twisted.python.reflect.objgrep hasn't been ported to Python 3"
def test_dictionary(self):
"""
Test references search through a dictionary, as a key or as a value.
"""
o = object()
d1 = {None: o}
d2 = {o: None}
self.assertIn("[None]", reflect.objgrep(d1, o, reflect.isSame))
self.assertIn("{None}", reflect.objgrep(d2, o, reflect.isSame))
def test_list(self):
"""
Test references search through a list.
"""
o = object()
L = [None, o]
self.assertIn("[1]", reflect.objgrep(L, o, reflect.isSame))
def test_tuple(self):
"""
Test references search through a tuple.
"""
o = object()
T = (o, None)
self.assertIn("[0]", reflect.objgrep(T, o, reflect.isSame))
def test_instance(self):
"""
Test references search through an object attribute.
"""
class Dummy:
pass
o = object()
d = Dummy()
d.o = o
self.assertIn(".o", reflect.objgrep(d, o, reflect.isSame))
def test_weakref(self):
"""
Test references search through a weakref object.
"""
class Dummy:
pass
o = Dummy()
w1 = weakref.ref(o)
self.assertIn("()", reflect.objgrep(w1, o, reflect.isSame))
def test_boundMethod(self):
"""
Test references search through method special attributes.
"""
class Dummy:
def dummy(self):
pass
o = Dummy()
m = o.dummy
self.assertIn(".__self__",
reflect.objgrep(m, m.__self__, reflect.isSame))
self.assertIn(".__self__.__class__",
reflect.objgrep(m, m.__self__.__class__, reflect.isSame))
self.assertIn(".__func__",
reflect.objgrep(m, m.__func__, reflect.isSame))
def test_everything(self):
"""
Test references search using complex set of objects.
"""
class Dummy:
def method(self):
pass
o = Dummy()
D1 = {(): "baz", None: "Quux", o: "Foosh"}
L = [None, (), D1, 3]
T = (L, {}, Dummy())
D2 = {0: "foo", 1: "bar", 2: T}
i = Dummy()
i.attr = D2
m = i.method
w = weakref.ref(m)
self.assertIn("().__self__.attr[2][0][2]{'Foosh'}",
reflect.objgrep(w, o, reflect.isSame))
def test_depthLimit(self):
"""
Test the depth of references search.
"""
a = []
b = [a]
c = [a, b]
d = [a, c]
self.assertEqual(['[0]'], reflect.objgrep(d, a, reflect.isSame, maxDepth=1))
self.assertEqual(['[0]', '[1][0]'], reflect.objgrep(d, a, reflect.isSame, maxDepth=2))
self.assertEqual(['[0]', '[1][0]', '[1][1][0]'], reflect.objgrep(d, a, reflect.isSame, maxDepth=3))
def test_deque(self):
"""
Test references search through a deque object.
"""
o = object()
D = deque()
D.append(None)
D.append(o)
self.assertIn("[1]", reflect.objgrep(D, o, reflect.isSame))
class GetClassTests(unittest.TestCase):
if _PY3:
oldClassNames = ['type']
else:
oldClassNames = ['class', 'classobj']
def test_old(self):
class OldClass:
pass
old = OldClass()
self.assertIn(reflect.getClass(OldClass).__name__, self.oldClassNames)
self.assertEqual(reflect.getClass(old).__name__, 'OldClass')
def test_new(self):
class NewClass(object):
pass
new = NewClass()
self.assertEqual(reflect.getClass(NewClass).__name__, 'type')
self.assertEqual(reflect.getClass(new).__name__, 'NewClass')
@@ -0,0 +1,67 @@
"""
Test win32 shortcut script
"""
import os.path
import sys
import tempfile
from twisted.trial import unittest
skipReason = None
try:
from win32com.shell import shell
from twisted.python import shortcut
except ImportError:
skipReason = "Only runs on Windows with win32com"
if sys.version_info[0:2] >= (3, 7):
skipReason = "Broken on Python 3.7+."
class ShortcutTests(unittest.TestCase):
skip = skipReason
def test_create(self):
"""
Create a simple shortcut.
"""
testFilename = __file__
baseFileName = os.path.basename(testFilename)
s1 = shortcut.Shortcut(testFilename)
tempname = self.mktemp() + '.lnk'
s1.save(tempname)
self.assertTrue(os.path.exists(tempname))
sc = shortcut.open(tempname)
scPath = sc.GetPath(shell.SLGP_RAWPATH)[0]
self.assertEqual(scPath[-len(baseFileName):].lower(),
baseFileName.lower())
def test_createPythonShortcut(self):
"""
Create a shortcut to the Python executable,
and set some values.
"""
testFilename = sys.executable
baseFileName = os.path.basename(testFilename)
tempDir = tempfile.gettempdir()
s1 = shortcut.Shortcut(
path=testFilename,
arguments="-V",
description="The Python executable",
workingdir=tempDir,
iconpath=tempDir,
iconidx=1,
)
tempname = self.mktemp() + '.lnk'
s1.save(tempname)
self.assertTrue(os.path.exists(tempname))
sc = shortcut.open(tempname)
scPath = sc.GetPath(shell.SLGP_RAWPATH)[0]
self.assertEqual(scPath[-len(baseFileName):].lower(),
baseFileName.lower())
self.assertEqual(sc.GetDescription(), "The Python executable")
self.assertEqual(sc.GetWorkingDirectory(), tempDir)
self.assertEqual(sc.GetIconLocation(), (tempDir, 1))
@@ -0,0 +1,178 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
from __future__ import division, absolute_import
import os
import sys
from textwrap import dedent
from twisted.trial import unittest
from twisted.persisted import sob
from twisted.python import components
from twisted.persisted.styles import Ephemeral
class Dummy(components.Componentized):
pass
objects = [
1,
"hello",
(1, "hello"),
[1, "hello"],
{1:"hello"},
]
class FakeModule(object):
pass
class PersistTests(unittest.TestCase):
def testStyles(self):
for o in objects:
p = sob.Persistent(o, '')
for style in 'source pickle'.split():
p.setStyle(style)
p.save(filename='persisttest.'+style)
o1 = sob.load('persisttest.'+style, style)
self.assertEqual(o, o1)
def testStylesBeingSet(self):
o = Dummy()
o.foo = 5
o.setComponent(sob.IPersistable, sob.Persistent(o, 'lala'))
for style in 'source pickle'.split():
sob.IPersistable(o).setStyle(style)
sob.IPersistable(o).save(filename='lala.'+style)
o1 = sob.load('lala.'+style, style)
self.assertEqual(o.foo, o1.foo)
self.assertEqual(sob.IPersistable(o1).style, style)
def testPassphraseError(self):
"""
Calling save() with a passphrase is an error.
"""
p = sob.Persistant(None, 'object')
self.assertRaises(
TypeError, p.save, 'filename.pickle', passphrase='abc')
def testNames(self):
o = [1,2,3]
p = sob.Persistent(o, 'object')
for style in 'source pickle'.split():
p.setStyle(style)
p.save()
o1 = sob.load('object.ta'+style[0], style)
self.assertEqual(o, o1)
for tag in 'lala lolo'.split():
p.save(tag)
o1 = sob.load('object-'+tag+'.ta'+style[0], style)
self.assertEqual(o, o1)
def testPython(self):
with open("persisttest.python", 'w') as f:
f.write('foo=[1,2,3] ')
o = sob.loadValueFromFile('persisttest.python', 'foo')
self.assertEqual(o, [1,2,3])
def testTypeGuesser(self):
self.assertRaises(KeyError, sob.guessType, "file.blah")
self.assertEqual('python', sob.guessType("file.py"))
self.assertEqual('python', sob.guessType("file.tac"))
self.assertEqual('python', sob.guessType("file.etac"))
self.assertEqual('pickle', sob.guessType("file.tap"))
self.assertEqual('pickle', sob.guessType("file.etap"))
self.assertEqual('source', sob.guessType("file.tas"))
self.assertEqual('source', sob.guessType("file.etas"))
def testEverythingEphemeralGetattr(self):
"""
L{_EverythingEphermal.__getattr__} will proxy the __main__ module as an
L{Ephemeral} object, and during load will be transparent, but after
load will return L{Ephemeral} objects from any accessed attributes.
"""
self.fakeMain.testMainModGetattr = 1
dirname = self.mktemp()
os.mkdir(dirname)
filename = os.path.join(dirname, 'persisttest.ee_getattr')
global mainWhileLoading
mainWhileLoading = None
with open(filename, "w") as f:
f.write(dedent("""
app = []
import __main__
app.append(__main__.testMainModGetattr == 1)
try:
__main__.somethingElse
except AttributeError:
app.append(True)
else:
app.append(False)
from twisted.test import test_sob
test_sob.mainWhileLoading = __main__
"""))
loaded = sob.load(filename, 'source')
self.assertIsInstance(loaded, list)
self.assertTrue(loaded[0], "Expected attribute not set.")
self.assertTrue(loaded[1], "Unexpected attribute set.")
self.assertIsInstance(mainWhileLoading, Ephemeral)
self.assertIsInstance(mainWhileLoading.somethingElse, Ephemeral)
del mainWhileLoading
def testEverythingEphemeralSetattr(self):
"""
Verify that _EverythingEphemeral.__setattr__ won't affect __main__.
"""
self.fakeMain.testMainModSetattr = 1
dirname = self.mktemp()
os.mkdir(dirname)
filename = os.path.join(dirname, 'persisttest.ee_setattr')
with open(filename, 'w') as f:
f.write('import __main__\n')
f.write('__main__.testMainModSetattr = 2\n')
f.write('app = None\n')
sob.load(filename, 'source')
self.assertEqual(self.fakeMain.testMainModSetattr, 1)
def testEverythingEphemeralException(self):
"""
Test that an exception during load() won't cause _EE to mask __main__
"""
dirname = self.mktemp()
os.mkdir(dirname)
filename = os.path.join(dirname, 'persisttest.ee_exception')
with open(filename, 'w') as f:
f.write('raise ValueError\n')
self.assertRaises(ValueError, sob.load, filename, 'source')
self.assertEqual(type(sys.modules['__main__']), FakeModule)
def setUp(self):
"""
Replace the __main__ module with a fake one, so that it can be mutated
in tests
"""
self.realMain = sys.modules['__main__']
self.fakeMain = sys.modules['__main__'] = FakeModule()
def tearDown(self):
"""
Restore __main__ to its original value
"""
sys.modules['__main__'] = self.realMain
@@ -0,0 +1,380 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Tests for L{twisted.internet.stdio}.
@var properEnv: A copy of L{os.environ} which has L{bytes} keys/values on POSIX
platforms and native L{str} keys/values on Windows.
"""
from __future__ import absolute_import, division
import os
import sys
import itertools
from twisted.trial import unittest
from twisted.python import filepath, log
from twisted.python.reflect import requireModule
from twisted.python.runtime import platform
from twisted.python.compat import range, intToBytes, bytesEnviron
from twisted.internet import error, defer, protocol, stdio, reactor
from twisted.test.test_tcp import ConnectionLostNotifyingProtocol
# A short string which is intended to appear here and nowhere else,
# particularly not in any random garbage output CPython unavoidable
# generates (such as in warning text and so forth). This is searched
# for in the output from stdio_test_lastwrite and if it is found at
# the end, the functionality works.
UNIQUE_LAST_WRITE_STRING = b'xyz123abc Twisted is great!'
skipWindowsNopywin32 = None
if platform.isWindows():
if requireModule('win32process') is None:
skipWindowsNopywin32 = ("On windows, spawnProcess is not available "
"in the absence of win32process.")
properEnv = dict(os.environ)
properEnv["PYTHONPATH"] = os.pathsep.join(sys.path)
else:
properEnv = bytesEnviron()
properEnv[b"PYTHONPATH"] = os.pathsep.join(sys.path).encode(
sys.getfilesystemencoding())
class StandardIOTestProcessProtocol(protocol.ProcessProtocol):
"""
Test helper for collecting output from a child process and notifying
something when it exits.
@ivar onConnection: A L{defer.Deferred} which will be called back with
L{None} when the connection to the child process is established.
@ivar onCompletion: A L{defer.Deferred} which will be errbacked with the
failure associated with the child process exiting when it exits.
@ivar onDataReceived: A L{defer.Deferred} which will be called back with
this instance whenever C{childDataReceived} is called, or L{None} to
suppress these callbacks.
@ivar data: A C{dict} mapping file descriptors to strings containing all
bytes received from the child process on each file descriptor.
"""
onDataReceived = None
def __init__(self):
self.onConnection = defer.Deferred()
self.onCompletion = defer.Deferred()
self.data = {}
def connectionMade(self):
self.onConnection.callback(None)
def childDataReceived(self, name, bytes):
"""
Record all bytes received from the child process in the C{data}
dictionary. Fire C{onDataReceived} if it is not L{None}.
"""
self.data[name] = self.data.get(name, b'') + bytes
if self.onDataReceived is not None:
d, self.onDataReceived = self.onDataReceived, None
d.callback(self)
def processEnded(self, reason):
self.onCompletion.callback(reason)
class StandardInputOutputTests(unittest.TestCase):
skip = skipWindowsNopywin32
def _spawnProcess(self, proto, sibling, *args, **kw):
"""
Launch a child Python process and communicate with it using the
given ProcessProtocol.
@param proto: A L{ProcessProtocol} instance which will be connected
to the child process.
@param sibling: The basename of a file containing the Python program
to run in the child process.
@param *args: strings which will be passed to the child process on
the command line as C{argv[2:]}.
@param **kw: additional arguments to pass to L{reactor.spawnProcess}.
@return: The L{IProcessTransport} provider for the spawned process.
"""
args = [sys.executable,
b"-m", b"twisted.test." + sibling,
reactor.__class__.__module__] + list(args)
return reactor.spawnProcess(
proto,
sys.executable,
args,
env=properEnv,
**kw)
def _requireFailure(self, d, callback):
def cb(result):
self.fail("Process terminated with non-Failure: %r" % (result,))
def eb(err):
return callback(err)
return d.addCallbacks(cb, eb)
def test_loseConnection(self):
"""
Verify that a protocol connected to L{StandardIO} can disconnect
itself using C{transport.loseConnection}.
"""
errorLogFile = self.mktemp()
log.msg("Child process logging to " + errorLogFile)
p = StandardIOTestProcessProtocol()
d = p.onCompletion
self._spawnProcess(p, b'stdio_test_loseconn', errorLogFile)
def processEnded(reason):
# Copy the child's log to ours so it's more visible.
with open(errorLogFile, 'r') as f:
for line in f:
log.msg("Child logged: " + line.rstrip())
self.failIfIn(1, p.data)
reason.trap(error.ProcessDone)
return self._requireFailure(d, processEnded)
def test_readConnectionLost(self):
"""
When stdin is closed and the protocol connected to it implements
L{IHalfCloseableProtocol}, the protocol's C{readConnectionLost} method
is called.
"""
errorLogFile = self.mktemp()
log.msg("Child process logging to " + errorLogFile)
p = StandardIOTestProcessProtocol()
p.onDataReceived = defer.Deferred()
def cbBytes(ignored):
d = p.onCompletion
p.transport.closeStdin()
return d
p.onDataReceived.addCallback(cbBytes)
def processEnded(reason):
reason.trap(error.ProcessDone)
d = self._requireFailure(p.onDataReceived, processEnded)
self._spawnProcess(
p, b'stdio_test_halfclose', errorLogFile)
return d
def test_lastWriteReceived(self):
"""
Verify that a write made directly to stdout using L{os.write}
after StandardIO has finished is reliably received by the
process reading that stdout.
"""
p = StandardIOTestProcessProtocol()
# Note: the macOS bug which prompted the addition of this test
# is an apparent race condition involving non-blocking PTYs.
# Delaying the parent process significantly increases the
# likelihood of the race going the wrong way. If you need to
# fiddle with this code at all, uncommenting the next line
# will likely make your life much easier. It is commented out
# because it makes the test quite slow.
# p.onConnection.addCallback(lambda ign: __import__('time').sleep(5))
try:
self._spawnProcess(
p, b'stdio_test_lastwrite', UNIQUE_LAST_WRITE_STRING,
usePTY=True)
except ValueError as e:
# Some platforms don't work with usePTY=True
raise unittest.SkipTest(str(e))
def processEnded(reason):
"""
Asserts that the parent received the bytes written by the child
immediately after the child starts.
"""
self.assertTrue(
p.data[1].endswith(UNIQUE_LAST_WRITE_STRING),
"Received %r from child, did not find expected bytes." % (
p.data,))
reason.trap(error.ProcessDone)
return self._requireFailure(p.onCompletion, processEnded)
def test_hostAndPeer(self):
"""
Verify that the transport of a protocol connected to L{StandardIO}
has C{getHost} and C{getPeer} methods.
"""
p = StandardIOTestProcessProtocol()
d = p.onCompletion
self._spawnProcess(p, b'stdio_test_hostpeer')
def processEnded(reason):
host, peer = p.data[1].splitlines()
self.assertTrue(host)
self.assertTrue(peer)
reason.trap(error.ProcessDone)
return self._requireFailure(d, processEnded)
def test_write(self):
"""
Verify that the C{write} method of the transport of a protocol
connected to L{StandardIO} sends bytes to standard out.
"""
p = StandardIOTestProcessProtocol()
d = p.onCompletion
self._spawnProcess(p, b'stdio_test_write')
def processEnded(reason):
self.assertEqual(p.data[1], b'ok!')
reason.trap(error.ProcessDone)
return self._requireFailure(d, processEnded)
def test_writeSequence(self):
"""
Verify that the C{writeSequence} method of the transport of a
protocol connected to L{StandardIO} sends bytes to standard out.
"""
p = StandardIOTestProcessProtocol()
d = p.onCompletion
self._spawnProcess(p, b'stdio_test_writeseq')
def processEnded(reason):
self.assertEqual(p.data[1], b'ok!')
reason.trap(error.ProcessDone)
return self._requireFailure(d, processEnded)
def _junkPath(self):
junkPath = self.mktemp()
with open(junkPath, 'wb') as junkFile:
for i in range(1024):
junkFile.write(intToBytes(i) + b'\n')
return junkPath
def test_producer(self):
"""
Verify that the transport of a protocol connected to L{StandardIO}
is a working L{IProducer} provider.
"""
p = StandardIOTestProcessProtocol()
d = p.onCompletion
written = []
toWrite = list(range(100))
def connectionMade(ign):
if toWrite:
written.append(intToBytes(toWrite.pop()) + b"\n")
proc.write(written[-1])
reactor.callLater(0.01, connectionMade, None)
proc = self._spawnProcess(p, b'stdio_test_producer')
p.onConnection.addCallback(connectionMade)
def processEnded(reason):
self.assertEqual(p.data[1], b''.join(written))
self.assertFalse(
toWrite,
"Connection lost with %d writes left to go." % (len(toWrite),))
reason.trap(error.ProcessDone)
return self._requireFailure(d, processEnded)
def test_consumer(self):
"""
Verify that the transport of a protocol connected to L{StandardIO}
is a working L{IConsumer} provider.
"""
p = StandardIOTestProcessProtocol()
d = p.onCompletion
junkPath = self._junkPath()
self._spawnProcess(p, b'stdio_test_consumer', junkPath)
def processEnded(reason):
with open(junkPath, 'rb') as f:
self.assertEqual(p.data[1], f.read())
reason.trap(error.ProcessDone)
return self._requireFailure(d, processEnded)
def test_normalFileStandardOut(self):
"""
If L{StandardIO} is created with a file descriptor which refers to a
normal file (ie, a file from the filesystem), L{StandardIO.write}
writes bytes to that file. In particular, it does not immediately
consider the file closed or call its protocol's C{connectionLost}
method.
"""
onConnLost = defer.Deferred()
proto = ConnectionLostNotifyingProtocol(onConnLost)
path = filepath.FilePath(self.mktemp())
self.normal = normal = path.open('wb')
self.addCleanup(normal.close)
kwargs = dict(stdout=normal.fileno())
if not platform.isWindows():
# Make a fake stdin so that StandardIO doesn't mess with the *real*
# stdin.
r, w = os.pipe()
self.addCleanup(os.close, r)
self.addCleanup(os.close, w)
kwargs['stdin'] = r
connection = stdio.StandardIO(proto, **kwargs)
# The reactor needs to spin a bit before it might have incorrectly
# decided stdout is closed. Use this counter to keep track of how
# much we've let it spin. If it closes before we expected, this
# counter will have a value that's too small and we'll know.
howMany = 5
count = itertools.count()
def spin():
for value in count:
if value == howMany:
connection.loseConnection()
return
connection.write(intToBytes(value))
break
reactor.callLater(0, spin)
reactor.callLater(0, spin)
# Once the connection is lost, make sure the counter is at the
# appropriate value.
def cbLost(reason):
self.assertEqual(next(count), howMany + 1)
self.assertEqual(
path.getContent(),
b''.join(map(intToBytes, range(howMany))))
onConnLost.addCallback(cbLost)
return onConnLost
if platform.isWindows():
test_normalFileStandardOut.skip = (
"StandardIO does not accept stdout as an argument to Windows. "
"Testing redirection to a file is therefore harder.")
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,132 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Tests for L{twisted.python.threadable}.
"""
from __future__ import division, absolute_import
import sys, pickle
try:
import threading
except ImportError:
threadingSkip = "Platform lacks thread support"
else:
threadingSkip = None
from twisted.python.compat import _PY3
from twisted.trial import unittest
from twisted.python import threadable
class TestObject:
synchronized = ['aMethod']
x = -1
y = 1
def aMethod(self):
for i in range(10):
self.x, self.y = self.y, self.x
self.z = self.x + self.y
assert self.z == 0, "z == %d, not 0 as expected" % (self.z,)
threadable.synchronize(TestObject)
class SynchronizationTests(unittest.SynchronousTestCase):
def setUp(self):
"""
Reduce the CPython check interval so that thread switches happen much
more often, hopefully exercising more possible race conditions. Also,
delay actual test startup until the reactor has been started.
"""
if _PY3:
if getattr(sys, 'getswitchinterval', None) is not None:
self.addCleanup(sys.setswitchinterval, sys.getswitchinterval())
sys.setswitchinterval(0.0000001)
else:
if getattr(sys, 'getcheckinterval', None) is not None:
self.addCleanup(sys.setcheckinterval, sys.getcheckinterval())
sys.setcheckinterval(7)
def test_synchronizedName(self):
"""
The name of a synchronized method is inaffected by the synchronization
decorator.
"""
self.assertEqual("aMethod", TestObject.aMethod.__name__)
def test_isInIOThread(self):
"""
L{threadable.isInIOThread} returns C{True} if and only if it is called
in the same thread as L{threadable.registerAsIOThread}.
"""
threadable.registerAsIOThread()
foreignResult = []
t = threading.Thread(
target=lambda: foreignResult.append(threadable.isInIOThread()))
t.start()
t.join()
self.assertFalse(
foreignResult[0], "Non-IO thread reported as IO thread")
self.assertTrue(
threadable.isInIOThread(), "IO thread reported as not IO thread")
def testThreadedSynchronization(self):
o = TestObject()
errors = []
def callMethodLots():
try:
for i in range(1000):
o.aMethod()
except AssertionError as e:
errors.append(str(e))
threads = []
for x in range(5):
t = threading.Thread(target=callMethodLots)
threads.append(t)
t.start()
for t in threads:
t.join()
if errors:
raise unittest.FailTest(errors)
if threadingSkip is not None:
testThreadedSynchronization.skip = threadingSkip
test_isInIOThread.skip = threadingSkip
def testUnthreadedSynchronization(self):
o = TestObject()
for i in range(1000):
o.aMethod()
class SerializationTests(unittest.SynchronousTestCase):
def testPickling(self):
lock = threadable.XLock()
lockType = type(lock)
lockPickle = pickle.dumps(lock)
newLock = pickle.loads(lockPickle)
self.assertIsInstance(newLock, lockType)
if threadingSkip is not None:
testPickling.skip = threadingSkip
def testUnpickling(self):
lockPickle = b'ctwisted.python.threadable\nunpickle_lock\np0\n(tp1\nRp2\n.'
lock = pickle.loads(lockPickle)
newPickle = pickle.dumps(lock, 2)
pickle.loads(newPickle)
@@ -0,0 +1,734 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Tests for L{twisted.python.threadpool}
"""
from __future__ import division, absolute_import
import pickle
import time
import weakref
import gc
import threading
from twisted.python.compat import range
from twisted.trial import unittest
from twisted.python import threadpool, threadable, failure, context
from twisted._threads import Team, createMemoryWorker
class Synchronization(object):
failures = 0
def __init__(self, N, waiting):
self.N = N
self.waiting = waiting
self.lock = threading.Lock()
self.runs = []
def run(self):
# This is the testy part: this is supposed to be invoked
# serially from multiple threads. If that is actually the
# case, we will never fail to acquire this lock. If it is
# *not* the case, we might get here while someone else is
# holding the lock.
if self.lock.acquire(False):
if not len(self.runs) % 5:
# Constant selected based on empirical data to maximize the
# chance of a quick failure if this code is broken.
time.sleep(0.0002)
self.lock.release()
else:
self.failures += 1
# This is just the only way I can think of to wake up the test
# method. It doesn't actually have anything to do with the
# test.
self.lock.acquire()
self.runs.append(None)
if len(self.runs) == self.N:
self.waiting.release()
self.lock.release()
synchronized = ["run"]
threadable.synchronize(Synchronization)
class ThreadPoolTests(unittest.SynchronousTestCase):
"""
Test threadpools.
"""
def getTimeout(self):
"""
Return number of seconds to wait before giving up.
"""
return 5 # Really should be order of magnitude less
def _waitForLock(self, lock):
items = range(1000000)
for i in items:
if lock.acquire(False):
break
time.sleep(1e-5)
else:
self.fail("A long time passed without succeeding")
def test_attributes(self):
"""
L{ThreadPool.min} and L{ThreadPool.max} are set to the values passed to
L{ThreadPool.__init__}.
"""
pool = threadpool.ThreadPool(12, 22)
self.assertEqual(pool.min, 12)
self.assertEqual(pool.max, 22)
def test_start(self):
"""
L{ThreadPool.start} creates the minimum number of threads specified.
"""
pool = threadpool.ThreadPool(0, 5)
pool.start()
self.addCleanup(pool.stop)
self.assertEqual(len(pool.threads), 0)
pool = threadpool.ThreadPool(3, 10)
self.assertEqual(len(pool.threads), 0)
pool.start()
self.addCleanup(pool.stop)
self.assertEqual(len(pool.threads), 3)
def test_adjustingWhenPoolStopped(self):
"""
L{ThreadPool.adjustPoolsize} only modifies the pool size and does not
start new workers while the pool is not running.
"""
pool = threadpool.ThreadPool(0, 5)
pool.start()
pool.stop()
pool.adjustPoolsize(2)
self.assertEqual(len(pool.threads), 0)
def test_threadCreationArguments(self):
"""
Test that creating threads in the threadpool with application-level
objects as arguments doesn't results in those objects never being
freed, with the thread maintaining a reference to them as long as it
exists.
"""
tp = threadpool.ThreadPool(0, 1)
tp.start()
self.addCleanup(tp.stop)
# Sanity check - no threads should have been started yet.
self.assertEqual(tp.threads, [])
# Here's our function
def worker(arg):
pass
# weakref needs an object subclass
class Dumb(object):
pass
# And here's the unique object
unique = Dumb()
workerRef = weakref.ref(worker)
uniqueRef = weakref.ref(unique)
# Put some work in
tp.callInThread(worker, unique)
# Add an event to wait completion
event = threading.Event()
tp.callInThread(event.set)
event.wait(self.getTimeout())
del worker
del unique
gc.collect()
self.assertIsNone(uniqueRef())
self.assertIsNone(workerRef())
def test_threadCreationArgumentsCallInThreadWithCallback(self):
"""
As C{test_threadCreationArguments} above, but for
callInThreadWithCallback.
"""
tp = threadpool.ThreadPool(0, 1)
tp.start()
self.addCleanup(tp.stop)
# Sanity check - no threads should have been started yet.
self.assertEqual(tp.threads, [])
# this holds references obtained in onResult
refdict = {} # name -> ref value
onResultWait = threading.Event()
onResultDone = threading.Event()
resultRef = []
# result callback
def onResult(success, result):
# Spin the GC, which should now delete worker and unique if it's
# not held on to by callInThreadWithCallback after it is complete
gc.collect()
onResultWait.wait(self.getTimeout())
refdict['workerRef'] = workerRef()
refdict['uniqueRef'] = uniqueRef()
onResultDone.set()
resultRef.append(weakref.ref(result))
# Here's our function
def worker(arg, test):
return Dumb()
# weakref needs an object subclass
class Dumb(object):
pass
# And here's the unique object
unique = Dumb()
onResultRef = weakref.ref(onResult)
workerRef = weakref.ref(worker)
uniqueRef = weakref.ref(unique)
# Put some work in
tp.callInThreadWithCallback(onResult, worker, unique, test=unique)
del worker
del unique
# let onResult collect the refs
onResultWait.set()
# wait for onResult
onResultDone.wait(self.getTimeout())
gc.collect()
self.assertIsNone(uniqueRef())
self.assertIsNone(workerRef())
# XXX There's a race right here - has onResult in the worker thread
# returned and the locals in _worker holding it and the result been
# deleted yet?
del onResult
gc.collect()
self.assertIsNone(onResultRef())
self.assertIsNone(resultRef[0]())
# The callback shouldn't have been able to resolve the references.
self.assertEqual(list(refdict.values()), [None, None])
def test_persistence(self):
"""
Threadpools can be pickled and unpickled, which should preserve the
number of threads and other parameters.
"""
pool = threadpool.ThreadPool(7, 20)
self.assertEqual(pool.min, 7)
self.assertEqual(pool.max, 20)
# check that unpickled threadpool has same number of threads
copy = pickle.loads(pickle.dumps(pool))
self.assertEqual(copy.min, 7)
self.assertEqual(copy.max, 20)
def _threadpoolTest(self, method):
"""
Test synchronization of calls made with C{method}, which should be
one of the mechanisms of the threadpool to execute work in threads.
"""
# This is a schizophrenic test: it seems to be trying to test
# both the callInThread()/dispatch() behavior of the ThreadPool as well
# as the serialization behavior of threadable.synchronize(). It
# would probably make more sense as two much simpler tests.
N = 10
tp = threadpool.ThreadPool()
tp.start()
self.addCleanup(tp.stop)
waiting = threading.Lock()
waiting.acquire()
actor = Synchronization(N, waiting)
for i in range(N):
method(tp, actor)
self._waitForLock(waiting)
self.assertFalse(
actor.failures,
"run() re-entered {} times".format(actor.failures))
def test_callInThread(self):
"""
Call C{_threadpoolTest} with C{callInThread}.
"""
return self._threadpoolTest(
lambda tp, actor: tp.callInThread(actor.run))
def test_callInThreadException(self):
"""
L{ThreadPool.callInThread} logs exceptions raised by the callable it
is passed.
"""
class NewError(Exception):
pass
def raiseError():
raise NewError()
tp = threadpool.ThreadPool(0, 1)
tp.callInThread(raiseError)
tp.start()
tp.stop()
errors = self.flushLoggedErrors(NewError)
self.assertEqual(len(errors), 1)
def test_callInThreadWithCallback(self):
"""
L{ThreadPool.callInThreadWithCallback} calls C{onResult} with a
two-tuple of C{(True, result)} where C{result} is the value returned
by the callable supplied.
"""
waiter = threading.Lock()
waiter.acquire()
results = []
def onResult(success, result):
waiter.release()
results.append(success)
results.append(result)
tp = threadpool.ThreadPool(0, 1)
tp.callInThreadWithCallback(onResult, lambda: "test")
tp.start()
try:
self._waitForLock(waiter)
finally:
tp.stop()
self.assertTrue(results[0])
self.assertEqual(results[1], "test")
def test_callInThreadWithCallbackExceptionInCallback(self):
"""
L{ThreadPool.callInThreadWithCallback} calls C{onResult} with a
two-tuple of C{(False, failure)} where C{failure} represents the
exception raised by the callable supplied.
"""
class NewError(Exception):
pass
def raiseError():
raise NewError()
waiter = threading.Lock()
waiter.acquire()
results = []
def onResult(success, result):
waiter.release()
results.append(success)
results.append(result)
tp = threadpool.ThreadPool(0, 1)
tp.callInThreadWithCallback(onResult, raiseError)
tp.start()
try:
self._waitForLock(waiter)
finally:
tp.stop()
self.assertFalse(results[0])
self.assertIsInstance(results[1], failure.Failure)
self.assertTrue(issubclass(results[1].type, NewError))
def test_callInThreadWithCallbackExceptionInOnResult(self):
"""
L{ThreadPool.callInThreadWithCallback} logs the exception raised by
C{onResult}.
"""
class NewError(Exception):
pass
waiter = threading.Lock()
waiter.acquire()
results = []
def onResult(success, result):
results.append(success)
results.append(result)
raise NewError()
tp = threadpool.ThreadPool(0, 1)
tp.callInThreadWithCallback(onResult, lambda: None)
tp.callInThread(waiter.release)
tp.start()
try:
self._waitForLock(waiter)
finally:
tp.stop()
errors = self.flushLoggedErrors(NewError)
self.assertEqual(len(errors), 1)
self.assertTrue(results[0])
self.assertIsNone(results[1])
def test_callbackThread(self):
"""
L{ThreadPool.callInThreadWithCallback} calls the function it is
given and the C{onResult} callback in the same thread.
"""
threadIds = []
event = threading.Event()
def onResult(success, result):
threadIds.append(threading.currentThread().ident)
event.set()
def func():
threadIds.append(threading.currentThread().ident)
tp = threadpool.ThreadPool(0, 1)
tp.callInThreadWithCallback(onResult, func)
tp.start()
self.addCleanup(tp.stop)
event.wait(self.getTimeout())
self.assertEqual(len(threadIds), 2)
self.assertEqual(threadIds[0], threadIds[1])
def test_callbackContext(self):
"""
The context L{ThreadPool.callInThreadWithCallback} is invoked in is
shared by the context the callable and C{onResult} callback are
invoked in.
"""
myctx = context.theContextTracker.currentContext().contexts[-1]
myctx['testing'] = 'this must be present'
contexts = []
event = threading.Event()
def onResult(success, result):
ctx = context.theContextTracker.currentContext().contexts[-1]
contexts.append(ctx)
event.set()
def func():
ctx = context.theContextTracker.currentContext().contexts[-1]
contexts.append(ctx)
tp = threadpool.ThreadPool(0, 1)
tp.callInThreadWithCallback(onResult, func)
tp.start()
self.addCleanup(tp.stop)
event.wait(self.getTimeout())
self.assertEqual(len(contexts), 2)
self.assertEqual(myctx, contexts[0])
self.assertEqual(myctx, contexts[1])
def test_existingWork(self):
"""
Work added to the threadpool before its start should be executed once
the threadpool is started: this is ensured by trying to release a lock
previously acquired.
"""
waiter = threading.Lock()
waiter.acquire()
tp = threadpool.ThreadPool(0, 1)
tp.callInThread(waiter.release) # Before start()
tp.start()
try:
self._waitForLock(waiter)
finally:
tp.stop()
def test_workerStateTransition(self):
"""
As the worker receives and completes work, it transitions between
the working and waiting states.
"""
pool = threadpool.ThreadPool(0, 1)
pool.start()
self.addCleanup(pool.stop)
# Sanity check
self.assertEqual(pool.workers, 0)
self.assertEqual(len(pool.waiters), 0)
self.assertEqual(len(pool.working), 0)
# Fire up a worker and give it some 'work'
threadWorking = threading.Event()
threadFinish = threading.Event()
def _thread():
threadWorking.set()
threadFinish.wait(10)
pool.callInThread(_thread)
threadWorking.wait(10)
self.assertEqual(pool.workers, 1)
self.assertEqual(len(pool.waiters), 0)
self.assertEqual(len(pool.working), 1)
# Finish work, and spin until state changes
threadFinish.set()
while not len(pool.waiters):
time.sleep(0.0005)
# Make sure state changed correctly
self.assertEqual(len(pool.waiters), 1)
self.assertEqual(len(pool.working), 0)
class RaceConditionTests(unittest.SynchronousTestCase):
def setUp(self):
self.threadpool = threadpool.ThreadPool(0, 10)
self.event = threading.Event()
self.threadpool.start()
def done():
self.threadpool.stop()
del self.threadpool
self.addCleanup(done)
def getTimeout(self):
"""
A reasonable number of seconds to time out.
"""
return 5
def test_synchronization(self):
"""
If multiple threads are waiting on an event (via blocking on something
in a callable passed to L{threadpool.ThreadPool.callInThread}), and
there is spare capacity in the threadpool, sending another callable
which will cause those to un-block to
L{threadpool.ThreadPool.callInThread} will reliably run that callable
and un-block the blocked threads promptly.
@note: This is not really a unit test, it is a stress-test. You may
need to run it with C{trial -u} to fail reliably if there is a
problem. It is very hard to regression-test for this particular
bug - one where the thread pool may consider itself as having
"enough capacity" when it really needs to spin up a new thread if
it possibly can - in a deterministic way, since the bug can only be
provoked by subtle race conditions.
"""
timeout = self.getTimeout()
self.threadpool.callInThread(self.event.set)
self.event.wait(timeout)
self.event.clear()
for i in range(3):
self.threadpool.callInThread(self.event.wait)
self.threadpool.callInThread(self.event.set)
self.event.wait(timeout)
if not self.event.isSet():
self.event.set()
self.fail(
"'set' did not run in thread; timed out waiting on 'wait'."
)
class MemoryPool(threadpool.ThreadPool):
"""
A deterministic threadpool that uses in-memory data structures to queue
work rather than threads to execute work.
"""
def __init__(self, coordinator, failTest, newWorker, *args, **kwargs):
"""
Initialize this L{MemoryPool} with a test case.
@param coordinator: a worker used to coordinate work in the L{Team}
underlying this threadpool.
@type coordinator: L{twisted._threads.IExclusiveWorker}
@param failTest: A 1-argument callable taking an exception and raising
a test-failure exception.
@type failTest: 1-argument callable taking (L{Failure}) and raising
L{unittest.FailTest}.
@param newWorker: a 0-argument callable that produces a new
L{twisted._threads.IWorker} provider on each invocation.
@type newWorker: 0-argument callable returning
L{twisted._threads.IWorker}.
"""
self._coordinator = coordinator
self._failTest = failTest
self._newWorker = newWorker
threadpool.ThreadPool.__init__(self, *args, **kwargs)
def _pool(self, currentLimit, threadFactory):
"""
Override testing hook to create a deterministic threadpool.
@param currentLimit: A 1-argument callable which returns the current
threadpool size limit.
@param threadFactory: ignored in this invocation; a 0-argument callable
that would produce a thread.
@return: a L{Team} backed by the coordinator and worker passed to
L{MemoryPool.__init__}.
"""
def respectLimit():
# The expression in this method copied and pasted from
# twisted.threads._pool, which is unfortunately bound up
# with lots of actual-threading stuff.
stats = team.statistics()
if ((stats.busyWorkerCount +
stats.idleWorkerCount) >= currentLimit()):
return None
return self._newWorker()
team = Team(coordinator=self._coordinator,
createWorker=respectLimit,
logException=self._failTest)
return team
class PoolHelper(object):
"""
A L{PoolHelper} constructs a L{threadpool.ThreadPool} that doesn't actually
use threads, by using the internal interfaces in L{twisted._threads}.
@ivar performCoordination: a 0-argument callable that will perform one unit
of "coordination" - work involved in delegating work to other threads -
and return L{True} if it did any work, L{False} otherwise.
@ivar workers: the workers which represent the threads within the pool -
the workers other than the coordinator.
@type workers: L{list} of 2-tuple of (L{IWorker}, C{workPerformer}) where
C{workPerformer} is a 0-argument callable like C{performCoordination}.
@ivar threadpool: a modified L{threadpool.ThreadPool} to test.
@type threadpool: L{MemoryPool}
"""
def __init__(self, testCase, *args, **kwargs):
"""
Create a L{PoolHelper}.
@param testCase: a test case attached to this helper.
@type args: The arguments passed to a L{threadpool.ThreadPool}.
@type kwargs: The arguments passed to a L{threadpool.ThreadPool}
"""
coordinator, self.performCoordination = createMemoryWorker()
self.workers = []
def newWorker():
self.workers.append(createMemoryWorker())
return self.workers[-1][0]
self.threadpool = MemoryPool(coordinator, testCase.fail, newWorker,
*args, **kwargs)
def performAllCoordination(self):
"""
Perform all currently scheduled "coordination", which is the work
involved in delegating work to other threads.
"""
while self.performCoordination():
pass
class MemoryBackedTests(unittest.SynchronousTestCase):
"""
Tests using L{PoolHelper} to deterministically test properties of the
threadpool implementation.
"""
def test_workBeforeStarting(self):
"""
If a threadpool is told to do work before starting, then upon starting
up, it will start enough workers to handle all of the enqueued work
that it's been given.
"""
helper = PoolHelper(self, 0, 10)
n = 5
for x in range(n):
helper.threadpool.callInThread(lambda: None)
helper.performAllCoordination()
self.assertEqual(helper.workers, [])
helper.threadpool.start()
helper.performAllCoordination()
self.assertEqual(len(helper.workers), n)
def test_tooMuchWorkBeforeStarting(self):
"""
If the amount of work before starting exceeds the maximum number of
threads allowed to the threadpool, only the maximum count will be
started.
"""
helper = PoolHelper(self, 0, 10)
n = 50
for x in range(n):
helper.threadpool.callInThread(lambda: None)
helper.performAllCoordination()
self.assertEqual(helper.workers, [])
helper.threadpool.start()
helper.performAllCoordination()
self.assertEqual(len(helper.workers), helper.threadpool.max)
@@ -0,0 +1,420 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Test methods in twisted.internet.threads and reactor thread APIs.
"""
from __future__ import division, absolute_import
import sys, os, time
from twisted.trial import unittest
from twisted.python.compat import range
from twisted.internet import reactor, defer, interfaces, threads, protocol, error
from twisted.python import failure, threadable, log, threadpool
class ReactorThreadsTests(unittest.TestCase):
"""
Tests for the reactor threading API.
"""
def test_suggestThreadPoolSize(self):
"""
Try to change maximum number of threads.
"""
reactor.suggestThreadPoolSize(34)
self.assertEqual(reactor.threadpool.max, 34)
reactor.suggestThreadPoolSize(4)
self.assertEqual(reactor.threadpool.max, 4)
def _waitForThread(self):
"""
The reactor's threadpool is only available when the reactor is running,
so to have a sane behavior during the tests we make a dummy
L{threads.deferToThread} call.
"""
return threads.deferToThread(time.sleep, 0)
def test_callInThread(self):
"""
Test callInThread functionality: set a C{threading.Event}, and check
that it's not in the main thread.
"""
def cb(ign):
waiter = threading.Event()
result = []
def threadedFunc():
result.append(threadable.isInIOThread())
waiter.set()
reactor.callInThread(threadedFunc)
waiter.wait(120)
if not waiter.isSet():
self.fail("Timed out waiting for event.")
else:
self.assertEqual(result, [False])
return self._waitForThread().addCallback(cb)
def test_callFromThread(self):
"""
Test callFromThread functionality: from the main thread, and from
another thread.
"""
def cb(ign):
firedByReactorThread = defer.Deferred()
firedByOtherThread = defer.Deferred()
def threadedFunc():
reactor.callFromThread(firedByOtherThread.callback, None)
reactor.callInThread(threadedFunc)
reactor.callFromThread(firedByReactorThread.callback, None)
return defer.DeferredList(
[firedByReactorThread, firedByOtherThread],
fireOnOneErrback=True)
return self._waitForThread().addCallback(cb)
def test_wakerOverflow(self):
"""
Try to make an overflow on the reactor waker using callFromThread.
"""
def cb(ign):
self.failure = None
waiter = threading.Event()
def threadedFunction():
# Hopefully a hundred thousand queued calls is enough to
# trigger the error condition
for i in range(100000):
try:
reactor.callFromThread(lambda: None)
except:
self.failure = failure.Failure()
break
waiter.set()
reactor.callInThread(threadedFunction)
waiter.wait(120)
if not waiter.isSet():
self.fail("Timed out waiting for event")
if self.failure is not None:
return defer.fail(self.failure)
return self._waitForThread().addCallback(cb)
def _testBlockingCallFromThread(self, reactorFunc):
"""
Utility method to test L{threads.blockingCallFromThread}.
"""
waiter = threading.Event()
results = []
errors = []
def cb1(ign):
def threadedFunc():
try:
r = threads.blockingCallFromThread(reactor, reactorFunc)
except Exception as e:
errors.append(e)
else:
results.append(r)
waiter.set()
reactor.callInThread(threadedFunc)
return threads.deferToThread(waiter.wait, self.getTimeout())
def cb2(ign):
if not waiter.isSet():
self.fail("Timed out waiting for event")
return results, errors
return self._waitForThread().addCallback(cb1).addBoth(cb2)
def test_blockingCallFromThread(self):
"""
Test blockingCallFromThread facility: create a thread, call a function
in the reactor using L{threads.blockingCallFromThread}, and verify the
result returned.
"""
def reactorFunc():
return defer.succeed("foo")
def cb(res):
self.assertEqual(res[0][0], "foo")
return self._testBlockingCallFromThread(reactorFunc).addCallback(cb)
def test_asyncBlockingCallFromThread(self):
"""
Test blockingCallFromThread as above, but be sure the resulting
Deferred is not already fired.
"""
def reactorFunc():
d = defer.Deferred()
reactor.callLater(0.1, d.callback, "egg")
return d
def cb(res):
self.assertEqual(res[0][0], "egg")
return self._testBlockingCallFromThread(reactorFunc).addCallback(cb)
def test_errorBlockingCallFromThread(self):
"""
Test error report for blockingCallFromThread.
"""
def reactorFunc():
return defer.fail(RuntimeError("bar"))
def cb(res):
self.assertIsInstance(res[1][0], RuntimeError)
self.assertEqual(res[1][0].args[0], "bar")
return self._testBlockingCallFromThread(reactorFunc).addCallback(cb)
def test_asyncErrorBlockingCallFromThread(self):
"""
Test error report for blockingCallFromThread as above, but be sure the
resulting Deferred is not already fired.
"""
def reactorFunc():
d = defer.Deferred()
reactor.callLater(0.1, d.errback, RuntimeError("spam"))
return d
def cb(res):
self.assertIsInstance(res[1][0], RuntimeError)
self.assertEqual(res[1][0].args[0], "spam")
return self._testBlockingCallFromThread(reactorFunc).addCallback(cb)
class Counter:
index = 0
problem = 0
def add(self):
"""A non thread-safe method."""
next = self.index + 1
# another thread could jump in here and increment self.index on us
if next != self.index + 1:
self.problem = 1
raise ValueError
# or here, same issue but we wouldn't catch it. We'd overwrite
# their results, and the index will have lost a count. If
# several threads get in here, we will actually make the count
# go backwards when we overwrite it.
self.index = next
class DeferredResultTests(unittest.TestCase):
"""
Test twisted.internet.threads.
"""
def setUp(self):
reactor.suggestThreadPoolSize(8)
def tearDown(self):
reactor.suggestThreadPoolSize(0)
def test_callMultiple(self):
"""
L{threads.callMultipleInThread} calls multiple functions in a thread.
"""
L = []
N = 10
d = defer.Deferred()
def finished():
self.assertEqual(L, list(range(N)))
d.callback(None)
threads.callMultipleInThread([
(L.append, (i,), {}) for i in range(N)
] + [(reactor.callFromThread, (finished,), {})])
return d
def test_deferredResult(self):
"""
L{threads.deferToThread} executes the function passed, and correctly
handles the positional and keyword arguments given.
"""
d = threads.deferToThread(lambda x, y=5: x + y, 3, y=4)
d.addCallback(self.assertEqual, 7)
return d
def test_deferredFailure(self):
"""
Check that L{threads.deferToThread} return a failure object
with an appropriate exception instance when the called
function raises an exception.
"""
class NewError(Exception):
pass
def raiseError():
raise NewError()
d = threads.deferToThread(raiseError)
return self.assertFailure(d, NewError)
def test_deferredFailureAfterSuccess(self):
"""
Check that a successful L{threads.deferToThread} followed by a one
that raises an exception correctly result as a failure.
"""
# set up a condition that causes cReactor to hang. These conditions
# can also be set by other tests when the full test suite is run in
# alphabetical order (test_flow.FlowTest.testThreaded followed by
# test_internet.ReactorCoreTestCase.testStop, to be precise). By
# setting them up explicitly here, we can reproduce the hang in a
# single precise test case instead of depending upon side effects of
# other tests.
#
# alas, this test appears to flunk the default reactor too
d = threads.deferToThread(lambda: None)
d.addCallback(lambda ign: threads.deferToThread(lambda: 1//0))
return self.assertFailure(d, ZeroDivisionError)
class DeferToThreadPoolTests(unittest.TestCase):
"""
Test L{twisted.internet.threads.deferToThreadPool}.
"""
def setUp(self):
self.tp = threadpool.ThreadPool(0, 8)
self.tp.start()
def tearDown(self):
self.tp.stop()
def test_deferredResult(self):
"""
L{threads.deferToThreadPool} executes the function passed, and
correctly handles the positional and keyword arguments given.
"""
d = threads.deferToThreadPool(reactor, self.tp,
lambda x, y=5: x + y, 3, y=4)
d.addCallback(self.assertEqual, 7)
return d
def test_deferredFailure(self):
"""
Check that L{threads.deferToThreadPool} return a failure object with an
appropriate exception instance when the called function raises an
exception.
"""
class NewError(Exception):
pass
def raiseError():
raise NewError()
d = threads.deferToThreadPool(reactor, self.tp, raiseError)
return self.assertFailure(d, NewError)
_callBeforeStartupProgram = """
import time
import %(reactor)s
%(reactor)s.install()
from twisted.internet import reactor
def threadedCall():
print('threaded call')
reactor.callInThread(threadedCall)
# Spin very briefly to try to give the thread a chance to run, if it
# is going to. Is there a better way to achieve this behavior?
for i in range(100):
time.sleep(0.0)
"""
class ThreadStartupProcessProtocol(protocol.ProcessProtocol):
def __init__(self, finished):
self.finished = finished
self.out = []
self.err = []
def outReceived(self, out):
self.out.append(out)
def errReceived(self, err):
self.err.append(err)
def processEnded(self, reason):
self.finished.callback((self.out, self.err, reason))
class StartupBehaviorTests(unittest.TestCase):
"""
Test cases for the behavior of the reactor threadpool near startup
boundary conditions.
In particular, this asserts that no threaded calls are attempted
until the reactor starts up, that calls attempted before it starts
are in fact executed once it has started, and that in both cases,
the reactor properly cleans itself up (which is tested for
somewhat implicitly, by requiring a child process be able to exit,
something it cannot do unless the threadpool has been properly
torn down).
"""
def testCallBeforeStartupUnexecuted(self):
progname = self.mktemp()
with open(progname, 'w') as progfile:
progfile.write(_callBeforeStartupProgram % {'reactor': reactor.__module__})
def programFinished(result):
(out, err, reason) = result
if reason.check(error.ProcessTerminated):
self.fail("Process did not exit cleanly (out: %s err: %s)" % (out, err))
if err:
log.msg("Unexpected output on standard error: %s" % (err,))
self.assertFalse(
out,
"Expected no output, instead received:\n%s" % (out,))
def programTimeout(err):
err.trap(error.TimeoutError)
proto.signalProcess('KILL')
return err
env = os.environ.copy()
env['PYTHONPATH'] = os.pathsep.join(sys.path)
d = defer.Deferred().addCallbacks(programFinished, programTimeout)
proto = ThreadStartupProcessProtocol(d)
reactor.spawnProcess(proto, sys.executable, ('python', progname), env)
return d
if interfaces.IReactorThreads(reactor, None) is None:
for cls in (ReactorThreadsTests,
DeferredResultTests,
StartupBehaviorTests):
cls.skip = "No thread support, nothing to test here."
else:
import threading
if interfaces.IReactorProcess(reactor, None) is None:
for cls in (StartupBehaviorTests,):
cls.skip = "No process support, cannot run subprocess thread tests."
@@ -0,0 +1,55 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
from twisted.trial import unittest
from twisted.protocols import loopback
from twisted.protocols import basic
from twisted.internet import protocol, abstract
from io import BytesIO
class BufferingServer(protocol.Protocol):
buffer = b''
def dataReceived(self, data):
self.buffer += data
class FileSendingClient(protocol.Protocol):
def __init__(self, f):
self.f = f
def connectionMade(self):
s = basic.FileSender()
d = s.beginFileTransfer(self.f, self.transport, lambda x: x)
d.addCallback(lambda r: self.transport.loseConnection())
class FileSenderTests(unittest.TestCase):
def testSendingFile(self):
testStr = b'xyz' * 100 + b'abc' * 100 + b'123' * 100
s = BufferingServer()
c = FileSendingClient(BytesIO(testStr))
d = loopback.loopbackTCP(s, c)
d.addCallback(lambda x : self.assertEqual(s.buffer, testStr))
return d
def testSendingEmptyFile(self):
fileSender = basic.FileSender()
consumer = abstract.FileDescriptor()
consumer.connected = 1
emptyFile = BytesIO(b'')
d = fileSender.beginFileTransfer(emptyFile, consumer, lambda x: x)
# The producer will be immediately exhausted, and so immediately
# unregistered
self.assertIsNone(consumer.producer)
# Which means the Deferred from FileSender should have been called
self.assertTrue(d.called,
'producer unregistered with deferred being called')
@@ -0,0 +1,715 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Tests for L{twisted.python.usage}, a command line option parsing library.
"""
from __future__ import division, absolute_import
from twisted.trial import unittest
from twisted.python import usage
class WellBehaved(usage.Options):
optParameters = [['long', 'w', 'default', 'and a docstring'],
['another', 'n', 'no docstring'],
['longonly', None, 'noshort'],
['shortless', None, 'except',
'this one got docstring'],
]
optFlags = [['aflag', 'f',
"""
flagallicious docstringness for this here
"""],
['flout', 'o'],
]
def opt_myflag(self):
self.opts['myflag'] = "PONY!"
def opt_myparam(self, value):
self.opts['myparam'] = "%s WITH A PONY!" % (value,)
class ParseCorrectnessTests(unittest.TestCase):
"""
Test L{usage.Options.parseOptions} for correct values under
good conditions.
"""
def setUp(self):
"""
Instantiate and parseOptions a well-behaved Options class.
"""
self.niceArgV = ("--long Alpha -n Beta "
"--shortless Gamma -f --myflag "
"--myparam Tofu").split()
self.nice = WellBehaved()
self.nice.parseOptions(self.niceArgV)
def test_checkParameters(self):
"""
Parameters have correct values.
"""
self.assertEqual(self.nice.opts['long'], "Alpha")
self.assertEqual(self.nice.opts['another'], "Beta")
self.assertEqual(self.nice.opts['longonly'], "noshort")
self.assertEqual(self.nice.opts['shortless'], "Gamma")
def test_checkFlags(self):
"""
Flags have correct values.
"""
self.assertEqual(self.nice.opts['aflag'], 1)
self.assertEqual(self.nice.opts['flout'], 0)
def test_checkCustoms(self):
"""
Custom flags and parameters have correct values.
"""
self.assertEqual(self.nice.opts['myflag'], "PONY!")
self.assertEqual(self.nice.opts['myparam'], "Tofu WITH A PONY!")
class TypedOptions(usage.Options):
optParameters = [
['fooint', None, 392, 'Foo int', int],
['foofloat', None, 4.23, 'Foo float', float],
['eggint', None, None, 'Egg int without default', int],
['eggfloat', None, None, 'Egg float without default', float],
]
def opt_under_score(self, value):
"""
This option has an underscore in its name to exercise the _ to -
translation.
"""
self.underscoreValue = value
opt_u = opt_under_score
class TypedTests(unittest.TestCase):
"""
Test L{usage.Options.parseOptions} for options with forced types.
"""
def setUp(self):
self.usage = TypedOptions()
def test_defaultValues(self):
"""
Default values are parsed.
"""
argV = []
self.usage.parseOptions(argV)
self.assertEqual(self.usage.opts['fooint'], 392)
self.assertIsInstance(self.usage.opts['fooint'], int)
self.assertEqual(self.usage.opts['foofloat'], 4.23)
self.assertIsInstance(self.usage.opts['foofloat'], float)
self.assertIsNone(self.usage.opts['eggint'])
self.assertIsNone(self.usage.opts['eggfloat'])
def test_parsingValues(self):
"""
int and float values are parsed.
"""
argV = ("--fooint 912 --foofloat -823.1 "
"--eggint 32 --eggfloat 21").split()
self.usage.parseOptions(argV)
self.assertEqual(self.usage.opts['fooint'], 912)
self.assertIsInstance(self.usage.opts['fooint'], int)
self.assertEqual(self.usage.opts['foofloat'], -823.1)
self.assertIsInstance(self.usage.opts['foofloat'], float)
self.assertEqual(self.usage.opts['eggint'], 32)
self.assertIsInstance(self.usage.opts['eggint'], int)
self.assertEqual(self.usage.opts['eggfloat'], 21.)
self.assertIsInstance(self.usage.opts['eggfloat'], float)
def test_underscoreOption(self):
"""
A dash in an option name is translated to an underscore before being
dispatched to a handler.
"""
self.usage.parseOptions(['--under-score', 'foo'])
self.assertEqual(self.usage.underscoreValue, 'foo')
def test_underscoreOptionAlias(self):
"""
An option name with a dash in it can have an alias.
"""
self.usage.parseOptions(['-u', 'bar'])
self.assertEqual(self.usage.underscoreValue, 'bar')
def test_invalidValues(self):
"""
Passing wrong values raises an error.
"""
argV = "--fooint egg".split()
self.assertRaises(usage.UsageError, self.usage.parseOptions, argV)
class WrongTypedOptions(usage.Options):
optParameters = [
['barwrong', None, None, 'Bar with wrong coerce', 'he']
]
class WeirdCallableOptions(usage.Options):
def _bar(value):
raise RuntimeError("Ouch")
def _foo(value):
raise ValueError("Yay")
optParameters = [
['barwrong', None, None, 'Bar with strange callable', _bar],
['foowrong', None, None, 'Foo with strange callable', _foo]
]
class WrongTypedTests(unittest.TestCase):
"""
Test L{usage.Options.parseOptions} for wrong coerce options.
"""
def test_nonCallable(self):
"""
Using a non-callable type fails.
"""
us = WrongTypedOptions()
argV = "--barwrong egg".split()
self.assertRaises(TypeError, us.parseOptions, argV)
def test_notCalledInDefault(self):
"""
The coerce functions are not called if no values are provided.
"""
us = WeirdCallableOptions()
argV = []
us.parseOptions(argV)
def test_weirdCallable(self):
"""
Errors raised by coerce functions are handled properly.
"""
us = WeirdCallableOptions()
argV = "--foowrong blah".split()
# ValueError is swallowed as UsageError
e = self.assertRaises(usage.UsageError, us.parseOptions, argV)
self.assertEqual(str(e), "Parameter type enforcement failed: Yay")
us = WeirdCallableOptions()
argV = "--barwrong blah".split()
# RuntimeError is not swallowed
self.assertRaises(RuntimeError, us.parseOptions, argV)
class OutputTests(unittest.TestCase):
def test_uppercasing(self):
"""
Error output case adjustment does not mangle options
"""
opt = WellBehaved()
e = self.assertRaises(usage.UsageError,
opt.parseOptions, ['-Z'])
self.assertEqual(str(e), 'option -Z not recognized')
class InquisitionOptions(usage.Options):
optFlags = [
('expect', 'e'),
]
optParameters = [
('torture-device', 't',
'comfy-chair',
'set preferred torture device'),
]
class HolyQuestOptions(usage.Options):
optFlags = [('horseback', 'h',
'use a horse'),
('for-grail', 'g'),
]
class SubCommandOptions(usage.Options):
optFlags = [('europian-swallow', None,
'set default swallow type to Europian'),
]
subCommands = [
('inquisition', 'inquest', InquisitionOptions,
'Perform an inquisition'),
('holyquest', 'quest', HolyQuestOptions,
'Embark upon a holy quest'),
]
class SubCommandTests(unittest.TestCase):
"""
Test L{usage.Options.parseOptions} for options with subcommands.
"""
def test_simpleSubcommand(self):
"""
A subcommand is recognized.
"""
o = SubCommandOptions()
o.parseOptions(['--europian-swallow', 'inquisition'])
self.assertTrue(o['europian-swallow'])
self.assertEqual(o.subCommand, 'inquisition')
self.assertIsInstance(o.subOptions, InquisitionOptions)
self.assertFalse(o.subOptions['expect'])
self.assertEqual(o.subOptions['torture-device'], 'comfy-chair')
def test_subcommandWithFlagsAndOptions(self):
"""
Flags and options of a subcommand are assigned.
"""
o = SubCommandOptions()
o.parseOptions(['inquisition', '--expect', '--torture-device=feather'])
self.assertFalse(o['europian-swallow'])
self.assertEqual(o.subCommand, 'inquisition')
self.assertIsInstance(o.subOptions, InquisitionOptions)
self.assertTrue(o.subOptions['expect'])
self.assertEqual(o.subOptions['torture-device'], 'feather')
def test_subcommandAliasWithFlagsAndOptions(self):
"""
Flags and options of a subcommand alias are assigned.
"""
o = SubCommandOptions()
o.parseOptions(['inquest', '--expect', '--torture-device=feather'])
self.assertFalse(o['europian-swallow'])
self.assertEqual(o.subCommand, 'inquisition')
self.assertIsInstance(o.subOptions, InquisitionOptions)
self.assertTrue(o.subOptions['expect'])
self.assertEqual(o.subOptions['torture-device'], 'feather')
def test_anotherSubcommandWithFlagsAndOptions(self):
"""
Flags and options of another subcommand are assigned.
"""
o = SubCommandOptions()
o.parseOptions(['holyquest', '--for-grail'])
self.assertFalse(o['europian-swallow'])
self.assertEqual(o.subCommand, 'holyquest')
self.assertIsInstance(o.subOptions, HolyQuestOptions)
self.assertFalse(o.subOptions['horseback'])
self.assertTrue(o.subOptions['for-grail'])
def test_noSubcommand(self):
"""
If no subcommand is specified and no default subcommand is assigned,
a subcommand will not be implied.
"""
o = SubCommandOptions()
o.parseOptions(['--europian-swallow'])
self.assertTrue(o['europian-swallow'])
self.assertIsNone(o.subCommand)
self.assertFalse(hasattr(o, 'subOptions'))
def test_defaultSubcommand(self):
"""
Flags and options in the default subcommand are assigned.
"""
o = SubCommandOptions()
o.defaultSubCommand = 'inquest'
o.parseOptions(['--europian-swallow'])
self.assertTrue(o['europian-swallow'])
self.assertEqual(o.subCommand, 'inquisition')
self.assertIsInstance(o.subOptions, InquisitionOptions)
self.assertFalse(o.subOptions['expect'])
self.assertEqual(o.subOptions['torture-device'], 'comfy-chair')
def test_subCommandParseOptionsHasParent(self):
"""
The parseOptions method from the Options object specified for the
given subcommand is called.
"""
class SubOpt(usage.Options):
def parseOptions(self, *a, **kw):
self.sawParent = self.parent
usage.Options.parseOptions(self, *a, **kw)
class Opt(usage.Options):
subCommands = [
('foo', 'f', SubOpt, 'bar'),
]
o = Opt()
o.parseOptions(['foo'])
self.assertTrue(hasattr(o.subOptions, 'sawParent'))
self.assertEqual(o.subOptions.sawParent , o)
def test_subCommandInTwoPlaces(self):
"""
The .parent pointer is correct even when the same Options class is
used twice.
"""
class SubOpt(usage.Options):
pass
class OptFoo(usage.Options):
subCommands = [
('foo', 'f', SubOpt, 'quux'),
]
class OptBar(usage.Options):
subCommands = [
('bar', 'b', SubOpt, 'quux'),
]
oFoo = OptFoo()
oFoo.parseOptions(['foo'])
oBar=OptBar()
oBar.parseOptions(['bar'])
self.assertTrue(hasattr(oFoo.subOptions, 'parent'))
self.assertTrue(hasattr(oBar.subOptions, 'parent'))
self.failUnlessIdentical(oFoo.subOptions.parent, oFoo)
self.failUnlessIdentical(oBar.subOptions.parent, oBar)
class HelpStringTests(unittest.TestCase):
"""
Test generated help strings.
"""
def setUp(self):
"""
Instantiate a well-behaved Options class.
"""
self.niceArgV = ("--long Alpha -n Beta "
"--shortless Gamma -f --myflag "
"--myparam Tofu").split()
self.nice = WellBehaved()
def test_noGoBoom(self):
"""
__str__ shouldn't go boom.
"""
try:
self.nice.__str__()
except Exception as e:
self.fail(e)
def test_whitespaceStripFlagsAndParameters(self):
"""
Extra whitespace in flag and parameters docs is stripped.
"""
# We test this by making sure aflag and it's help string are on the
# same line.
lines = [s for s in str(self.nice).splitlines() if s.find("aflag")>=0]
self.assertTrue(len(lines) > 0)
self.assertTrue(lines[0].find("flagallicious") >= 0)
class PortCoerceTests(unittest.TestCase):
"""
Test the behavior of L{usage.portCoerce}.
"""
def test_validCoerce(self):
"""
Test the answers with valid input.
"""
self.assertEqual(0, usage.portCoerce("0"))
self.assertEqual(3210, usage.portCoerce("3210"))
self.assertEqual(65535, usage.portCoerce("65535"))
def test_errorCoerce(self):
"""
Test error path.
"""
self.assertRaises(ValueError, usage.portCoerce, "")
self.assertRaises(ValueError, usage.portCoerce, "-21")
self.assertRaises(ValueError, usage.portCoerce, "212189")
self.assertRaises(ValueError, usage.portCoerce, "foo")
class ZshCompleterTests(unittest.TestCase):
"""
Test the behavior of the various L{twisted.usage.Completer} classes
for producing output usable by zsh tab-completion system.
"""
def test_completer(self):
"""
Completer produces zsh shell-code that produces no completion matches.
"""
c = usage.Completer()
got = c._shellCode('some-option', usage._ZSH)
self.assertEqual(got, ':some-option:')
c = usage.Completer(descr='some action', repeat=True)
got = c._shellCode('some-option', usage._ZSH)
self.assertEqual(got, '*:some action:')
def test_files(self):
"""
CompleteFiles produces zsh shell-code that completes file names
according to a glob.
"""
c = usage.CompleteFiles()
got = c._shellCode('some-option', usage._ZSH)
self.assertEqual(got, ':some-option (*):_files -g "*"')
c = usage.CompleteFiles('*.py')
got = c._shellCode('some-option', usage._ZSH)
self.assertEqual(got, ':some-option (*.py):_files -g "*.py"')
c = usage.CompleteFiles('*.py', descr="some action", repeat=True)
got = c._shellCode('some-option', usage._ZSH)
self.assertEqual(got, '*:some action (*.py):_files -g "*.py"')
def test_dirs(self):
"""
CompleteDirs produces zsh shell-code that completes directory names.
"""
c = usage.CompleteDirs()
got = c._shellCode('some-option', usage._ZSH)
self.assertEqual(got, ':some-option:_directories')
c = usage.CompleteDirs(descr="some action", repeat=True)
got = c._shellCode('some-option', usage._ZSH)
self.assertEqual(got, '*:some action:_directories')
def test_list(self):
"""
CompleteList produces zsh shell-code that completes words from a fixed
list of possibilities.
"""
c = usage.CompleteList('ABC')
got = c._shellCode('some-option', usage._ZSH)
self.assertEqual(got, ':some-option:(A B C)')
c = usage.CompleteList(['1', '2', '3'])
got = c._shellCode('some-option', usage._ZSH)
self.assertEqual(got, ':some-option:(1 2 3)')
c = usage.CompleteList(['1', '2', '3'], descr='some action',
repeat=True)
got = c._shellCode('some-option', usage._ZSH)
self.assertEqual(got, '*:some action:(1 2 3)')
def test_multiList(self):
"""
CompleteMultiList produces zsh shell-code that completes multiple
comma-separated words from a fixed list of possibilities.
"""
c = usage.CompleteMultiList('ABC')
got = c._shellCode('some-option', usage._ZSH)
self.assertEqual(got, ':some-option:_values -s , \'some-option\' A B C')
c = usage.CompleteMultiList(['1','2','3'])
got = c._shellCode('some-option', usage._ZSH)
self.assertEqual(got, ':some-option:_values -s , \'some-option\' 1 2 3')
c = usage.CompleteMultiList(['1','2','3'], descr='some action',
repeat=True)
got = c._shellCode('some-option', usage._ZSH)
expected = '*:some action:_values -s , \'some action\' 1 2 3'
self.assertEqual(got, expected)
def test_usernames(self):
"""
CompleteUsernames produces zsh shell-code that completes system
usernames.
"""
c = usage.CompleteUsernames()
out = c._shellCode('some-option', usage._ZSH)
self.assertEqual(out, ':some-option:_users')
c = usage.CompleteUsernames(descr='some action', repeat=True)
out = c._shellCode('some-option', usage._ZSH)
self.assertEqual(out, '*:some action:_users')
def test_groups(self):
"""
CompleteGroups produces zsh shell-code that completes system group
names.
"""
c = usage.CompleteGroups()
out = c._shellCode('some-option', usage._ZSH)
self.assertEqual(out, ':group:_groups')
c = usage.CompleteGroups(descr='some action', repeat=True)
out = c._shellCode('some-option', usage._ZSH)
self.assertEqual(out, '*:some action:_groups')
def test_hostnames(self):
"""
CompleteHostnames produces zsh shell-code that completes hostnames.
"""
c = usage.CompleteHostnames()
out = c._shellCode('some-option', usage._ZSH)
self.assertEqual(out, ':some-option:_hosts')
c = usage.CompleteHostnames(descr='some action', repeat=True)
out = c._shellCode('some-option', usage._ZSH)
self.assertEqual(out, '*:some action:_hosts')
def test_userAtHost(self):
"""
CompleteUserAtHost produces zsh shell-code that completes hostnames or
a word of the form <username>@<hostname>.
"""
c = usage.CompleteUserAtHost()
out = c._shellCode('some-option', usage._ZSH)
self.assertTrue(out.startswith(':host | user@host:'))
c = usage.CompleteUserAtHost(descr='some action', repeat=True)
out = c._shellCode('some-option', usage._ZSH)
self.assertTrue(out.startswith('*:some action:'))
def test_netInterfaces(self):
"""
CompleteNetInterfaces produces zsh shell-code that completes system
network interface names.
"""
c = usage.CompleteNetInterfaces()
out = c._shellCode('some-option', usage._ZSH)
self.assertEqual(out, ':some-option:_net_interfaces')
c = usage.CompleteNetInterfaces(descr='some action', repeat=True)
out = c._shellCode('some-option', usage._ZSH)
self.assertEqual(out, '*:some action:_net_interfaces')
class CompleterNotImplementedTests(unittest.TestCase):
"""
Using an unknown shell constant with the various Completer() classes
should raise NotImplementedError
"""
def test_unknownShell(self):
"""
Using an unknown shellType should raise NotImplementedError
"""
classes = [usage.Completer, usage.CompleteFiles,
usage.CompleteDirs, usage.CompleteList,
usage.CompleteMultiList, usage.CompleteUsernames,
usage.CompleteGroups, usage.CompleteHostnames,
usage.CompleteUserAtHost, usage.CompleteNetInterfaces]
for cls in classes:
try:
action = cls()
except:
action = cls(None)
self.assertRaises(NotImplementedError, action._shellCode,
None, "bad_shell_type")
class FlagFunctionTests(unittest.TestCase):
"""
Tests for L{usage.flagFunction}.
"""
class SomeClass(object):
"""
Dummy class for L{usage.flagFunction} tests.
"""
def oneArg(self, a):
"""
A one argument method to be tested by L{usage.flagFunction}.
@param a: a useless argument to satisfy the function's signature.
"""
def noArg(self):
"""
A no argument method to be tested by L{usage.flagFunction}.
"""
def manyArgs(self, a, b, c):
"""
A multiple arguments method to be tested by L{usage.flagFunction}.
@param a: a useless argument to satisfy the function's signature.
@param b: a useless argument to satisfy the function's signature.
@param c: a useless argument to satisfy the function's signature.
"""
def test_hasArg(self):
"""
L{usage.flagFunction} returns C{False} if the method checked allows
exactly one argument.
"""
self.assertIs(False, usage.flagFunction(self.SomeClass().oneArg))
def test_noArg(self):
"""
L{usage.flagFunction} returns C{True} if the method checked allows
exactly no argument.
"""
self.assertIs(True, usage.flagFunction(self.SomeClass().noArg))
def test_tooManyArguments(self):
"""
L{usage.flagFunction} raises L{usage.UsageError} if the method checked
allows more than one argument.
"""
exc = self.assertRaises(
usage.UsageError, usage.flagFunction, self.SomeClass().manyArgs)
self.assertEqual("Invalid Option function for manyArgs", str(exc))
def test_tooManyArgumentsAndSpecificErrorMessage(self):
"""
L{usage.flagFunction} uses the given method name in the error message
raised when the method allows too many arguments.
"""
exc = self.assertRaises(
usage.UsageError,
usage.flagFunction, self.SomeClass().manyArgs, "flubuduf")
self.assertEqual("Invalid Option function for flubuduf", str(exc))
class OptionsInternalTests(unittest.TestCase):
"""
Tests internal behavior of C{usage.Options}.
"""
def test_optionsAliasesOrder(self):
"""
Options which are synonyms to another option are aliases towards the
longest option name.
"""
class Opts(usage.Options):
def opt_very_very_long(self):
"""
This is an option method with a very long name, that is going to
be aliased.
"""
opt_short = opt_very_very_long
opt_s = opt_very_very_long
opts = Opts()
self.assertEqual(
dict.fromkeys(
["s", "short", "very-very-long"], "very-very-long"), {
"s": opts.synonyms["s"],
"short": opts.synonyms["short"],
"very-very-long": opts.synonyms["very-very-long"],
})