17.12
This commit is contained in:
@@ -0,0 +1,16 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
# Make the twisted module executable with the default behaviour of
|
||||
# running twist.
|
||||
# This is not a docstring to avoid changing the string output of twist.
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
import sys
|
||||
from pkg_resources import load_entry_point
|
||||
|
||||
if __name__ == '__main__':
|
||||
sys.exit(
|
||||
load_entry_point('Twisted', 'console_scripts', 'twist')()
|
||||
)
|
||||
@@ -0,0 +1,61 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Test cases for convenience functionality in L{twisted._threads._convenience}.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division, print_function
|
||||
|
||||
from twisted.trial.unittest import SynchronousTestCase
|
||||
|
||||
from .._convenience import Quit
|
||||
from .._ithreads import AlreadyQuit
|
||||
|
||||
|
||||
class QuitTests(SynchronousTestCase):
|
||||
"""
|
||||
Tests for L{Quit}
|
||||
"""
|
||||
|
||||
def test_isInitiallySet(self):
|
||||
"""
|
||||
L{Quit.isSet} starts as L{False}.
|
||||
"""
|
||||
quit = Quit()
|
||||
self.assertEqual(quit.isSet, False)
|
||||
|
||||
|
||||
def test_setSetsSet(self):
|
||||
"""
|
||||
L{Quit.set} sets L{Quit.isSet} to L{True}.
|
||||
"""
|
||||
quit = Quit()
|
||||
quit.set()
|
||||
self.assertEqual(quit.isSet, True)
|
||||
|
||||
|
||||
def test_checkDoesNothing(self):
|
||||
"""
|
||||
L{Quit.check} initially does nothing and returns L{None}.
|
||||
"""
|
||||
quit = Quit()
|
||||
self.assertIs(quit.check(), None)
|
||||
|
||||
|
||||
def test_checkAfterSetRaises(self):
|
||||
"""
|
||||
L{Quit.check} raises L{AlreadyQuit} if L{Quit.set} has been called.
|
||||
"""
|
||||
quit = Quit()
|
||||
quit.set()
|
||||
self.assertRaises(AlreadyQuit, quit.check)
|
||||
|
||||
|
||||
def test_setTwiceRaises(self):
|
||||
"""
|
||||
L{Quit.set} raises L{AlreadyQuit} if it has been called previously.
|
||||
"""
|
||||
quit = Quit()
|
||||
quit.set()
|
||||
self.assertRaises(AlreadyQuit, quit.set)
|
||||
@@ -0,0 +1,65 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted._threads._memory}.
|
||||
"""
|
||||
from __future__ import absolute_import, division, print_function
|
||||
|
||||
from zope.interface.verify import verifyObject
|
||||
|
||||
from twisted.trial.unittest import SynchronousTestCase
|
||||
from .. import AlreadyQuit, IWorker, createMemoryWorker
|
||||
|
||||
|
||||
class MemoryWorkerTests(SynchronousTestCase):
|
||||
"""
|
||||
Tests for L{MemoryWorker}.
|
||||
"""
|
||||
|
||||
def test_createWorkerAndPerform(self):
|
||||
"""
|
||||
L{createMemoryWorker} creates an L{IWorker} and a callable that can
|
||||
perform work on it. The performer returns C{True} if it accomplished
|
||||
useful work.
|
||||
"""
|
||||
worker, performer = createMemoryWorker()
|
||||
verifyObject(IWorker, worker)
|
||||
done = []
|
||||
worker.do(lambda: done.append(3))
|
||||
worker.do(lambda: done.append(4))
|
||||
self.assertEqual(done, [])
|
||||
self.assertEqual(performer(), True)
|
||||
self.assertEqual(done, [3])
|
||||
self.assertEqual(performer(), True)
|
||||
self.assertEqual(done, [3, 4])
|
||||
|
||||
|
||||
def test_quitQuits(self):
|
||||
"""
|
||||
Calling C{quit} on the worker returned by L{createMemoryWorker} causes
|
||||
its C{do} and C{quit} methods to raise L{AlreadyQuit}; its C{perform}
|
||||
callable will start raising L{AlreadyQuit} when the work already
|
||||
provided to C{do} has been exhausted.
|
||||
"""
|
||||
worker, performer = createMemoryWorker()
|
||||
done = []
|
||||
def moreWork():
|
||||
done.append(7)
|
||||
worker.do(moreWork)
|
||||
worker.quit()
|
||||
self.assertRaises(AlreadyQuit, worker.do, moreWork)
|
||||
self.assertRaises(AlreadyQuit, worker.quit)
|
||||
performer()
|
||||
self.assertEqual(done, [7])
|
||||
self.assertEqual(performer(), False)
|
||||
|
||||
|
||||
def test_performWhenNothingToDoYet(self):
|
||||
"""
|
||||
The C{perform} callable returned by L{createMemoryWorker} will return
|
||||
no result when there's no work to do yet. Since there is no work to
|
||||
do, the performer returns C{False}.
|
||||
"""
|
||||
worker, performer = createMemoryWorker()
|
||||
self.assertEqual(performer(), False)
|
||||
@@ -0,0 +1,11 @@
|
||||
"""
|
||||
Provides Twisted version information.
|
||||
"""
|
||||
|
||||
# This file is auto-generated! Do not edit!
|
||||
# Use `python -m incremental.update Twisted` to change this file.
|
||||
|
||||
from incremental import Version
|
||||
|
||||
__version__ = Version('Twisted', 19, 10, 0)
|
||||
__all__ = ["__version__"]
|
||||
@@ -0,0 +1,6 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Configuration objects for Twisted Applications.
|
||||
"""
|
||||
@@ -0,0 +1,138 @@
|
||||
# -*- test-case-name: twisted.application.runner.test.test_exit -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
System exit support.
|
||||
"""
|
||||
|
||||
from sys import stdout, stderr, exit as sysexit
|
||||
|
||||
from constantly import Values, ValueConstant
|
||||
|
||||
|
||||
|
||||
def exit(status, message=None):
|
||||
"""
|
||||
Exit the python interpreter with the given status and an optional message.
|
||||
|
||||
@param status: An exit status.
|
||||
@type status: L{int} or L{ValueConstant} from L{ExitStatus}.
|
||||
|
||||
@param message: An options message to print.
|
||||
@type status: L{str}
|
||||
"""
|
||||
if isinstance(status, ValueConstant):
|
||||
code = status.value
|
||||
else:
|
||||
code = int(status)
|
||||
|
||||
if message:
|
||||
if code == 0:
|
||||
out = stdout
|
||||
else:
|
||||
out = stderr
|
||||
out.write(message)
|
||||
out.write("\n")
|
||||
|
||||
sysexit(code)
|
||||
|
||||
|
||||
|
||||
try:
|
||||
import posix as Status
|
||||
except ImportError:
|
||||
class Status(object):
|
||||
"""
|
||||
Object to hang C{EX_*} values off of as a substitute for L{posix}.
|
||||
"""
|
||||
EX__BASE = 64
|
||||
|
||||
EX_OK = 0
|
||||
EX_USAGE = EX__BASE
|
||||
EX_DATAERR = EX__BASE + 1
|
||||
EX_NOINPUT = EX__BASE + 2
|
||||
EX_NOUSER = EX__BASE + 3
|
||||
EX_NOHOST = EX__BASE + 4
|
||||
EX_UNAVAILABLE = EX__BASE + 5
|
||||
EX_SOFTWARE = EX__BASE + 6
|
||||
EX_OSERR = EX__BASE + 7
|
||||
EX_OSFILE = EX__BASE + 8
|
||||
EX_CANTCREAT = EX__BASE + 9
|
||||
EX_IOERR = EX__BASE + 10
|
||||
EX_TEMPFAIL = EX__BASE + 11
|
||||
EX_PROTOCOL = EX__BASE + 12
|
||||
EX_NOPERM = EX__BASE + 13
|
||||
EX_CONFIG = EX__BASE + 14
|
||||
|
||||
|
||||
|
||||
class ExitStatus(Values):
|
||||
"""
|
||||
Standard exit status codes for system programs.
|
||||
|
||||
@cvar EX_OK: Successful termination.
|
||||
@type EX_OK: L{ValueConstant}
|
||||
|
||||
@cvar EX_USAGE: Command line usage error.
|
||||
@type EX_USAGE: L{ValueConstant}
|
||||
|
||||
@cvar EX_DATAERR: Data format error.
|
||||
@type EX_DATAERR: L{ValueConstant}
|
||||
|
||||
@cvar EX_NOINPUT: Cannot open input.
|
||||
@type EX_NOINPUT: L{ValueConstant}
|
||||
|
||||
@cvar EX_NOUSER: Addressee unknown.
|
||||
@type EX_NOUSER: L{ValueConstant}
|
||||
|
||||
@cvar EX_NOHOST: Host name unknown.
|
||||
@type EX_NOHOST: L{ValueConstant}
|
||||
|
||||
@cvar EX_UNAVAILABLE: Service unavailable.
|
||||
@type EX_UNAVAILABLE: L{ValueConstant}
|
||||
|
||||
@cvar EX_SOFTWARE: Internal software error.
|
||||
@type EX_SOFTWARE: L{ValueConstant}
|
||||
|
||||
@cvar EX_OSERR: System error (e.g., can't fork).
|
||||
@type EX_OSERR: L{ValueConstant}
|
||||
|
||||
@cvar EX_OSFILE: Critical OS file missing.
|
||||
@type EX_OSFILE: L{ValueConstant}
|
||||
|
||||
@cvar EX_CANTCREAT: Can't create (user) output file.
|
||||
@type EX_CANTCREAT: L{ValueConstant}
|
||||
|
||||
@cvar EX_IOERR: Input/output error.
|
||||
@type EX_IOERR: L{ValueConstant}
|
||||
|
||||
@cvar EX_TEMPFAIL: Temporary failure; the user is invited to retry.
|
||||
@type EX_TEMPFAIL: L{ValueConstant}
|
||||
|
||||
@cvar EX_PROTOCOL: Remote error in protocol.
|
||||
@type EX_PROTOCOL: L{ValueConstant}
|
||||
|
||||
@cvar EX_NOPERM: Permission denied.
|
||||
@type EX_NOPERM: L{ValueConstant}
|
||||
|
||||
@cvar EX_CONFIG: Configuration error.
|
||||
@type EX_CONFIG: L{ValueConstant}
|
||||
"""
|
||||
|
||||
EX_OK = ValueConstant(Status.EX_OK)
|
||||
EX_USAGE = ValueConstant(Status.EX_USAGE)
|
||||
EX_DATAERR = ValueConstant(Status.EX_DATAERR)
|
||||
EX_NOINPUT = ValueConstant(Status.EX_NOINPUT)
|
||||
EX_NOUSER = ValueConstant(Status.EX_NOUSER)
|
||||
EX_NOHOST = ValueConstant(Status.EX_NOHOST)
|
||||
EX_UNAVAILABLE = ValueConstant(Status.EX_UNAVAILABLE)
|
||||
EX_SOFTWARE = ValueConstant(Status.EX_SOFTWARE)
|
||||
EX_OSERR = ValueConstant(Status.EX_OSERR)
|
||||
EX_OSFILE = ValueConstant(Status.EX_OSFILE)
|
||||
EX_CANTCREAT = ValueConstant(Status.EX_CANTCREAT)
|
||||
EX_IOERR = ValueConstant(Status.EX_IOERR)
|
||||
EX_TEMPFAIL = ValueConstant(Status.EX_TEMPFAIL)
|
||||
EX_PROTOCOL = ValueConstant(Status.EX_PROTOCOL)
|
||||
EX_NOPERM = ValueConstant(Status.EX_NOPERM)
|
||||
EX_CONFIG = ValueConstant(Status.EX_CONFIG)
|
||||
@@ -0,0 +1,7 @@
|
||||
# -*- test-case-name: twisted.application.runner.test -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.application.runner}.
|
||||
"""
|
||||
@@ -0,0 +1,476 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.application.runner._pidfile}.
|
||||
"""
|
||||
|
||||
from functools import wraps
|
||||
import errno
|
||||
from os import getpid, name as SYSTEM_NAME
|
||||
from io import BytesIO
|
||||
|
||||
from zope.interface import implementer
|
||||
from zope.interface.verify import verifyObject
|
||||
|
||||
from twisted.python.filepath import IFilePath
|
||||
from twisted.python.runtime import platform
|
||||
|
||||
from ...runner import _pidfile
|
||||
from .._pidfile import (
|
||||
IPIDFile, PIDFile, NonePIDFile,
|
||||
AlreadyRunningError, InvalidPIDFileError, StalePIDFileError,
|
||||
NoPIDFound,
|
||||
)
|
||||
|
||||
import twisted.trial.unittest
|
||||
from twisted.trial.unittest import SkipTest
|
||||
|
||||
|
||||
def ifPlatformSupported(f):
|
||||
"""
|
||||
Decorator for tests that are not expected to work on all platforms.
|
||||
|
||||
Calling L{PIDFile.isRunning} currently raises L{NotImplementedError} on
|
||||
non-POSIX platforms.
|
||||
|
||||
On an unsupported platform, we expect to see any test that calls
|
||||
L{PIDFile.isRunning} to raise either L{NotImplementedError}, L{SkipTest},
|
||||
or C{self.failureException}.
|
||||
(C{self.failureException} may occur in a test that checks for a specific
|
||||
exception but it gets NotImplementedError instead.)
|
||||
|
||||
@param f: The test method to decorate.
|
||||
@type f: method
|
||||
|
||||
@return: The wrapped callable.
|
||||
"""
|
||||
@wraps(f)
|
||||
def wrapper(self, *args, **kwargs):
|
||||
supported = platform.getType() == "posix"
|
||||
|
||||
if supported:
|
||||
return f(self, *args, **kwargs)
|
||||
else:
|
||||
e = self.assertRaises(
|
||||
(NotImplementedError, SkipTest, self.failureException),
|
||||
f, self, *args, **kwargs
|
||||
)
|
||||
if isinstance(e, NotImplementedError):
|
||||
self.assertTrue(
|
||||
str(e).startswith("isRunning is not implemented on ")
|
||||
)
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
|
||||
class PIDFileTests(twisted.trial.unittest.TestCase):
|
||||
"""
|
||||
Tests for L{PIDFile}.
|
||||
"""
|
||||
|
||||
def test_interface(self):
|
||||
"""
|
||||
L{PIDFile} conforms to L{IPIDFile}.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
verifyObject(IPIDFile, pidFile)
|
||||
|
||||
|
||||
def test_formatWithPID(self):
|
||||
"""
|
||||
L{PIDFile._format} returns the expected format when given a PID.
|
||||
"""
|
||||
self.assertEqual(PIDFile._format(pid=1337), b"1337\n")
|
||||
|
||||
|
||||
def test_readWithPID(self):
|
||||
"""
|
||||
L{PIDFile.read} returns the PID from the given file path.
|
||||
"""
|
||||
pid = 1337
|
||||
|
||||
pidFile = PIDFile(DummyFilePath(PIDFile._format(pid=pid)))
|
||||
|
||||
self.assertEqual(pid, pidFile.read())
|
||||
|
||||
|
||||
def test_readEmptyPID(self):
|
||||
"""
|
||||
L{PIDFile.read} raises L{InvalidPIDFileError} when given an empty file
|
||||
path.
|
||||
"""
|
||||
pidValue = b""
|
||||
pidFile = PIDFile(DummyFilePath(b""))
|
||||
|
||||
e = self.assertRaises(InvalidPIDFileError, pidFile.read)
|
||||
self.assertEqual(
|
||||
str(e),
|
||||
"non-integer PID value in PID file: {!r}".format(pidValue)
|
||||
)
|
||||
|
||||
|
||||
def test_readWithBogusPID(self):
|
||||
"""
|
||||
L{PIDFile.read} raises L{InvalidPIDFileError} when given an empty file
|
||||
path.
|
||||
"""
|
||||
pidValue = b"$foo!"
|
||||
pidFile = PIDFile(DummyFilePath(pidValue))
|
||||
|
||||
e = self.assertRaises(InvalidPIDFileError, pidFile.read)
|
||||
self.assertEqual(
|
||||
str(e),
|
||||
"non-integer PID value in PID file: {!r}".format(pidValue)
|
||||
)
|
||||
|
||||
|
||||
def test_readDoesntExist(self):
|
||||
"""
|
||||
L{PIDFile.read} raises L{NoPIDFound} when given a non-existing file
|
||||
path.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
|
||||
e = self.assertRaises(NoPIDFound, pidFile.read)
|
||||
self.assertEqual(str(e), "PID file does not exist")
|
||||
|
||||
|
||||
def test_readOpenRaisesOSErrorNotENOENT(self):
|
||||
"""
|
||||
L{PIDFile.read} re-raises L{OSError} if the associated C{errno} is
|
||||
anything other than L{errno.ENOENT}.
|
||||
"""
|
||||
def oops(mode="r"):
|
||||
raise OSError(errno.EIO, "I/O error")
|
||||
|
||||
self.patch(DummyFilePath, "open", oops)
|
||||
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
|
||||
error = self.assertRaises(OSError, pidFile.read)
|
||||
self.assertEqual(error.errno, errno.EIO)
|
||||
|
||||
|
||||
def test_writePID(self):
|
||||
"""
|
||||
L{PIDFile._write} stores the given PID.
|
||||
"""
|
||||
pid = 1995
|
||||
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
pidFile._write(pid)
|
||||
|
||||
self.assertEqual(pidFile.read(), pid)
|
||||
|
||||
|
||||
def test_writePIDInvalid(self):
|
||||
"""
|
||||
L{PIDFile._write} raises L{ValueError} when given an invalid PID.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
|
||||
self.assertRaises(ValueError, pidFile._write, u"burp")
|
||||
|
||||
|
||||
def test_writeRunningPID(self):
|
||||
"""
|
||||
L{PIDFile.writeRunningPID} stores the PID for the current process.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
pidFile.writeRunningPID()
|
||||
|
||||
self.assertEqual(pidFile.read(), getpid())
|
||||
|
||||
|
||||
def test_remove(self):
|
||||
"""
|
||||
L{PIDFile.remove} removes the PID file.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath(b""))
|
||||
self.assertTrue(pidFile.filePath.exists())
|
||||
|
||||
pidFile.remove()
|
||||
self.assertFalse(pidFile.filePath.exists())
|
||||
|
||||
|
||||
@ifPlatformSupported
|
||||
def test_isRunningDoesExist(self):
|
||||
"""
|
||||
L{PIDFile.isRunning} returns true for a process that does exist.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
pidFile._write(1337)
|
||||
|
||||
def kill(pid, signal):
|
||||
return # Don't actually kill anything
|
||||
|
||||
self.patch(_pidfile, "kill", kill)
|
||||
|
||||
self.assertTrue(pidFile.isRunning())
|
||||
|
||||
|
||||
@ifPlatformSupported
|
||||
def test_isRunningThis(self):
|
||||
"""
|
||||
L{PIDFile.isRunning} returns true for this process (which is running).
|
||||
|
||||
@note: This differs from L{PIDFileTests.test_isRunningDoesExist} in
|
||||
that it actually invokes the C{kill} system call, which is useful for
|
||||
testing of our chosen method for probing the existence of a process.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
pidFile.writeRunningPID()
|
||||
|
||||
self.assertTrue(pidFile.isRunning())
|
||||
|
||||
|
||||
@ifPlatformSupported
|
||||
def test_isRunningDoesNotExist(self):
|
||||
"""
|
||||
L{PIDFile.isRunning} raises L{StalePIDFileError} for a process that
|
||||
does not exist (errno=ESRCH).
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
pidFile._write(1337)
|
||||
|
||||
def kill(pid, signal):
|
||||
raise OSError(errno.ESRCH, "No such process")
|
||||
|
||||
self.patch(_pidfile, "kill", kill)
|
||||
|
||||
self.assertRaises(StalePIDFileError, pidFile.isRunning)
|
||||
|
||||
|
||||
@ifPlatformSupported
|
||||
def test_isRunningNotAllowed(self):
|
||||
"""
|
||||
L{PIDFile.isRunning} returns true for a process that we are not allowed
|
||||
to kill (errno=EPERM).
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
pidFile._write(1337)
|
||||
|
||||
def kill(pid, signal):
|
||||
raise OSError(errno.EPERM, "Operation not permitted")
|
||||
|
||||
self.patch(_pidfile, "kill", kill)
|
||||
|
||||
self.assertTrue(pidFile.isRunning())
|
||||
|
||||
|
||||
@ifPlatformSupported
|
||||
def test_isRunningInit(self):
|
||||
"""
|
||||
L{PIDFile.isRunning} returns true for a process that we are not allowed
|
||||
to kill (errno=EPERM).
|
||||
|
||||
@note: This differs from L{PIDFileTests.test_isRunningNotAllowed} in
|
||||
that it actually invokes the C{kill} system call, which is useful for
|
||||
testing of our chosen method for probing the existence of a process
|
||||
that we are not allowed to kill.
|
||||
|
||||
@note: In this case, we try killing C{init}, which is process #1 on
|
||||
POSIX systems, so this test is not portable. C{init} should always be
|
||||
running and should not be killable by non-root users.
|
||||
"""
|
||||
if SYSTEM_NAME != "posix":
|
||||
raise SkipTest("This test assumes POSIX")
|
||||
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
pidFile._write(1) # PID 1 is init on POSIX systems
|
||||
|
||||
self.assertTrue(pidFile.isRunning())
|
||||
|
||||
|
||||
@ifPlatformSupported
|
||||
def test_isRunningUnknownErrno(self):
|
||||
"""
|
||||
L{PIDFile.isRunning} re-raises L{OSError} if the attached C{errno}
|
||||
value from L{os.kill} is not an expected one.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
pidFile.writeRunningPID()
|
||||
|
||||
def kill(pid, signal):
|
||||
raise OSError(errno.EEXIST, "File exists")
|
||||
|
||||
self.patch(_pidfile, "kill", kill)
|
||||
|
||||
self.assertRaises(OSError, pidFile.isRunning)
|
||||
|
||||
|
||||
def test_isRunningNoPIDFile(self):
|
||||
"""
|
||||
L{PIDFile.isRunning} returns false if the PID file doesn't exist.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
|
||||
self.assertFalse(pidFile.isRunning())
|
||||
|
||||
|
||||
def test_contextManager(self):
|
||||
"""
|
||||
When used as a context manager, a L{PIDFile} will store the current pid
|
||||
on entry, then removes the PID file on exit.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
self.assertFalse(pidFile.filePath.exists())
|
||||
|
||||
with pidFile:
|
||||
self.assertTrue(pidFile.filePath.exists())
|
||||
self.assertEqual(pidFile.read(), getpid())
|
||||
|
||||
self.assertFalse(pidFile.filePath.exists())
|
||||
|
||||
|
||||
@ifPlatformSupported
|
||||
def test_contextManagerDoesntExist(self):
|
||||
"""
|
||||
When used as a context manager, a L{PIDFile} will replace the
|
||||
underlying PIDFile rather than raising L{AlreadyRunningError} if the
|
||||
contained PID file exists but refers to a non-running PID.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
pidFile._write(1337)
|
||||
|
||||
def kill(pid, signal):
|
||||
raise OSError(errno.ESRCH, "No such process")
|
||||
|
||||
self.patch(_pidfile, "kill", kill)
|
||||
|
||||
e = self.assertRaises(StalePIDFileError, pidFile.isRunning)
|
||||
self.assertEqual(str(e), "PID file refers to non-existing process")
|
||||
|
||||
with pidFile:
|
||||
self.assertEqual(pidFile.read(), getpid())
|
||||
|
||||
|
||||
@ifPlatformSupported
|
||||
def test_contextManagerAlreadyRunning(self):
|
||||
"""
|
||||
When used as a context manager, a L{PIDFile} will raise
|
||||
L{AlreadyRunningError} if the there is already a running process with
|
||||
the contained PID.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath())
|
||||
pidFile._write(1337)
|
||||
|
||||
def kill(pid, signal):
|
||||
return # Don't actually kill anything
|
||||
|
||||
self.patch(_pidfile, "kill", kill)
|
||||
|
||||
self.assertTrue(pidFile.isRunning())
|
||||
|
||||
self.assertRaises(AlreadyRunningError, pidFile.__enter__)
|
||||
|
||||
|
||||
|
||||
class NonePIDFileTests(twisted.trial.unittest.TestCase):
|
||||
"""
|
||||
Tests for L{NonePIDFile}.
|
||||
"""
|
||||
|
||||
def test_interface(self):
|
||||
"""
|
||||
L{NonePIDFile} conforms to L{IPIDFile}.
|
||||
"""
|
||||
pidFile = NonePIDFile()
|
||||
verifyObject(IPIDFile, pidFile)
|
||||
|
||||
|
||||
def test_read(self):
|
||||
"""
|
||||
L{NonePIDFile.read} raises L{NoPIDFound}.
|
||||
"""
|
||||
pidFile = NonePIDFile()
|
||||
|
||||
e = self.assertRaises(NoPIDFound, pidFile.read)
|
||||
self.assertEqual(str(e), "PID file does not exist")
|
||||
|
||||
|
||||
def test_write(self):
|
||||
"""
|
||||
L{NonePIDFile._write} raises L{OSError} with an errno of L{errno.EPERM}.
|
||||
"""
|
||||
pidFile = NonePIDFile()
|
||||
|
||||
error = self.assertRaises(OSError, pidFile._write, 0)
|
||||
self.assertEqual(error.errno, errno.EPERM)
|
||||
|
||||
|
||||
def test_writeRunningPID(self):
|
||||
"""
|
||||
L{NonePIDFile.writeRunningPID} raises L{OSError} with an errno of
|
||||
L{errno.EPERM}.
|
||||
"""
|
||||
pidFile = NonePIDFile()
|
||||
|
||||
error = self.assertRaises(OSError, pidFile.writeRunningPID)
|
||||
self.assertEqual(error.errno, errno.EPERM)
|
||||
|
||||
|
||||
def test_remove(self):
|
||||
"""
|
||||
L{NonePIDFile.remove} raises L{OSError} with an errno of L{errno.EPERM}.
|
||||
"""
|
||||
pidFile = NonePIDFile()
|
||||
|
||||
error = self.assertRaises(OSError, pidFile.remove)
|
||||
self.assertEqual(error.errno, errno.ENOENT)
|
||||
|
||||
|
||||
def test_isRunning(self):
|
||||
"""
|
||||
L{NonePIDFile.isRunning} returns L{False}.
|
||||
"""
|
||||
pidFile = NonePIDFile()
|
||||
|
||||
self.assertEqual(pidFile.isRunning(), False)
|
||||
|
||||
|
||||
def test_contextManager(self):
|
||||
"""
|
||||
When used as a context manager, a L{NonePIDFile} doesn't raise, despite
|
||||
not existing.
|
||||
"""
|
||||
pidFile = NonePIDFile()
|
||||
|
||||
with pidFile:
|
||||
pass
|
||||
|
||||
|
||||
|
||||
@implementer(IFilePath)
|
||||
class DummyFilePath(object):
|
||||
"""
|
||||
In-memory L{IFilePath}.
|
||||
"""
|
||||
|
||||
def __init__(self, content=None):
|
||||
self.setContent(content)
|
||||
|
||||
|
||||
def open(self, mode="r"):
|
||||
if not self._exists:
|
||||
raise OSError(errno.ENOENT, "No such file or directory")
|
||||
return BytesIO(self.getContent())
|
||||
|
||||
|
||||
def setContent(self, content):
|
||||
self._exists = content is not None
|
||||
self._content = content
|
||||
|
||||
|
||||
def getContent(self):
|
||||
return self._content
|
||||
|
||||
|
||||
def remove(self):
|
||||
self.setContent(None)
|
||||
|
||||
|
||||
def exists(self):
|
||||
return self._exists
|
||||
@@ -0,0 +1,460 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.application.runner._runner}.
|
||||
"""
|
||||
|
||||
from signal import SIGTERM
|
||||
from io import BytesIO
|
||||
import errno
|
||||
|
||||
from attr import attrib, attrs, Factory
|
||||
|
||||
from twisted.logger import (
|
||||
LogLevel, LogPublisher, LogBeginner,
|
||||
FileLogObserver, FilteringLogObserver, LogLevelFilterPredicate,
|
||||
)
|
||||
from twisted.test.proto_helpers import MemoryReactor
|
||||
|
||||
from ...runner import _runner
|
||||
from .._exit import ExitStatus
|
||||
from .._pidfile import PIDFile, NonePIDFile
|
||||
from .._runner import Runner
|
||||
from .test_pidfile import DummyFilePath
|
||||
|
||||
import twisted.trial.unittest
|
||||
|
||||
|
||||
|
||||
class RunnerTests(twisted.trial.unittest.TestCase):
|
||||
"""
|
||||
Tests for L{Runner}.
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
# Patch exit and kill so we can capture usage and prevent actual exits
|
||||
# and kills.
|
||||
|
||||
self.exit = DummyExit()
|
||||
self.kill = DummyKill()
|
||||
|
||||
self.patch(_runner, "exit", self.exit)
|
||||
self.patch(_runner, "kill", self.kill)
|
||||
|
||||
# Patch getpid so we get a known result
|
||||
|
||||
self.pid = 1337
|
||||
self.pidFileContent = u"{}\n".format(self.pid).encode("utf-8")
|
||||
|
||||
# Patch globalLogBeginner so that we aren't trying to install multiple
|
||||
# global log observers.
|
||||
|
||||
self.stdout = BytesIO()
|
||||
self.stderr = BytesIO()
|
||||
self.stdio = DummyStandardIO(self.stdout, self.stderr)
|
||||
self.warnings = DummyWarningsModule()
|
||||
|
||||
self.globalLogPublisher = LogPublisher()
|
||||
self.globalLogBeginner = LogBeginner(
|
||||
self.globalLogPublisher,
|
||||
self.stdio.stderr, self.stdio,
|
||||
self.warnings,
|
||||
)
|
||||
|
||||
self.patch(_runner, "stderr", self.stderr)
|
||||
self.patch(_runner, "globalLogBeginner", self.globalLogBeginner)
|
||||
|
||||
|
||||
def test_runInOrder(self):
|
||||
"""
|
||||
L{Runner.run} calls the expected methods in order.
|
||||
"""
|
||||
runner = DummyRunner(reactor=MemoryReactor())
|
||||
runner.run()
|
||||
|
||||
self.assertEqual(
|
||||
runner.calledMethods,
|
||||
[
|
||||
"killIfRequested",
|
||||
"startLogging",
|
||||
"startReactor",
|
||||
"reactorExited",
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def test_runUsesPIDFile(self):
|
||||
"""
|
||||
L{Runner.run} uses the provided PID file.
|
||||
"""
|
||||
pidFile = DummyPIDFile()
|
||||
|
||||
runner = Runner(reactor=MemoryReactor(), pidFile=pidFile)
|
||||
|
||||
self.assertFalse(pidFile.entered)
|
||||
self.assertFalse(pidFile.exited)
|
||||
|
||||
runner.run()
|
||||
|
||||
self.assertTrue(pidFile.entered)
|
||||
self.assertTrue(pidFile.exited)
|
||||
|
||||
|
||||
def test_runAlreadyRunning(self):
|
||||
"""
|
||||
L{Runner.run} exits with L{ExitStatus.EX_USAGE} and the expected
|
||||
message if a process is already running that corresponds to the given
|
||||
PID file.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath(self.pidFileContent))
|
||||
pidFile.isRunning = lambda: True
|
||||
|
||||
runner = Runner(reactor=MemoryReactor(), pidFile=pidFile)
|
||||
runner.run()
|
||||
|
||||
self.assertEqual(self.exit.status, ExitStatus.EX_CONFIG)
|
||||
self.assertEqual(self.exit.message, "Already running.")
|
||||
|
||||
|
||||
def test_killNotRequested(self):
|
||||
"""
|
||||
L{Runner.killIfRequested} when C{kill} is false doesn't exit and
|
||||
doesn't indiscriminately murder anyone.
|
||||
"""
|
||||
runner = Runner(reactor=MemoryReactor())
|
||||
runner.killIfRequested()
|
||||
|
||||
self.assertEqual(self.kill.calls, [])
|
||||
self.assertFalse(self.exit.exited)
|
||||
|
||||
|
||||
def test_killRequestedWithoutPIDFile(self):
|
||||
"""
|
||||
L{Runner.killIfRequested} when C{kill} is true but C{pidFile} is
|
||||
L{nonePIDFile} exits with L{ExitStatus.EX_USAGE} and the expected
|
||||
message; and also doesn't indiscriminately murder anyone.
|
||||
"""
|
||||
runner = Runner(reactor=MemoryReactor(), kill=True)
|
||||
runner.killIfRequested()
|
||||
|
||||
self.assertEqual(self.kill.calls, [])
|
||||
self.assertEqual(self.exit.status, ExitStatus.EX_USAGE)
|
||||
self.assertEqual(self.exit.message, "No PID file specified.")
|
||||
|
||||
|
||||
def test_killRequestedWithPIDFile(self):
|
||||
"""
|
||||
L{Runner.killIfRequested} when C{kill} is true and given a C{pidFile}
|
||||
performs a targeted killing of the appropriate process.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath(self.pidFileContent))
|
||||
runner = Runner(reactor=MemoryReactor(), kill=True, pidFile=pidFile)
|
||||
runner.killIfRequested()
|
||||
|
||||
self.assertEqual(self.kill.calls, [(self.pid, SIGTERM)])
|
||||
self.assertEqual(self.exit.status, ExitStatus.EX_OK)
|
||||
self.assertIdentical(self.exit.message, None)
|
||||
|
||||
|
||||
def test_killRequestedWithPIDFileCantRead(self):
|
||||
"""
|
||||
L{Runner.killIfRequested} when C{kill} is true and given a C{pidFile}
|
||||
that it can't read exits with L{ExitStatus.EX_IOERR}.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath(None))
|
||||
|
||||
def read():
|
||||
raise OSError(errno.EACCES, "Permission denied")
|
||||
|
||||
pidFile.read = read
|
||||
|
||||
runner = Runner(reactor=MemoryReactor(), kill=True, pidFile=pidFile)
|
||||
runner.killIfRequested()
|
||||
|
||||
self.assertEqual(self.exit.status, ExitStatus.EX_IOERR)
|
||||
self.assertEqual(self.exit.message, "Unable to read PID file.")
|
||||
|
||||
|
||||
def test_killRequestedWithPIDFileEmpty(self):
|
||||
"""
|
||||
L{Runner.killIfRequested} when C{kill} is true and given a C{pidFile}
|
||||
containing no value exits with L{ExitStatus.EX_DATAERR}.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath(b""))
|
||||
runner = Runner(reactor=MemoryReactor(), kill=True, pidFile=pidFile)
|
||||
runner.killIfRequested()
|
||||
|
||||
self.assertEqual(self.exit.status, ExitStatus.EX_DATAERR)
|
||||
self.assertEqual(self.exit.message, "Invalid PID file.")
|
||||
|
||||
|
||||
def test_killRequestedWithPIDFileNotAnInt(self):
|
||||
"""
|
||||
L{Runner.killIfRequested} when C{kill} is true and given a C{pidFile}
|
||||
containing a non-integer value exits with L{ExitStatus.EX_DATAERR}.
|
||||
"""
|
||||
pidFile = PIDFile(DummyFilePath(b"** totally not a number, dude **"))
|
||||
runner = Runner(reactor=MemoryReactor(), kill=True, pidFile=pidFile)
|
||||
runner.killIfRequested()
|
||||
|
||||
self.assertEqual(self.exit.status, ExitStatus.EX_DATAERR)
|
||||
self.assertEqual(self.exit.message, "Invalid PID file.")
|
||||
|
||||
|
||||
def test_startLogging(self):
|
||||
"""
|
||||
L{Runner.startLogging} sets up a filtering observer with a log level
|
||||
predicate set to the given log level that contains a file observer of
|
||||
the given type which writes to the given file.
|
||||
"""
|
||||
logFile = BytesIO()
|
||||
|
||||
# Patch the log beginner so that we don't try to start the already
|
||||
# running (started by trial) logging system.
|
||||
|
||||
class LogBeginner(object):
|
||||
def beginLoggingTo(self, observers):
|
||||
LogBeginner.observers = observers
|
||||
|
||||
self.patch(_runner, "globalLogBeginner", LogBeginner())
|
||||
|
||||
# Patch FilteringLogObserver so we can capture its arguments
|
||||
|
||||
class MockFilteringLogObserver(FilteringLogObserver):
|
||||
def __init__(
|
||||
self, observer, predicates,
|
||||
negativeObserver=lambda event: None
|
||||
):
|
||||
MockFilteringLogObserver.observer = observer
|
||||
MockFilteringLogObserver.predicates = predicates
|
||||
FilteringLogObserver.__init__(
|
||||
self, observer, predicates, negativeObserver
|
||||
)
|
||||
|
||||
self.patch(_runner, "FilteringLogObserver", MockFilteringLogObserver)
|
||||
|
||||
# Patch FileLogObserver so we can capture its arguments
|
||||
|
||||
class MockFileLogObserver(FileLogObserver):
|
||||
def __init__(self, outFile):
|
||||
MockFileLogObserver.outFile = outFile
|
||||
FileLogObserver.__init__(self, outFile, str)
|
||||
|
||||
# Start logging
|
||||
runner = Runner(
|
||||
reactor=MemoryReactor(),
|
||||
defaultLogLevel=LogLevel.critical,
|
||||
logFile=logFile,
|
||||
fileLogObserverFactory=MockFileLogObserver,
|
||||
)
|
||||
runner.startLogging()
|
||||
|
||||
# Check for a filtering observer
|
||||
self.assertEqual(len(LogBeginner.observers), 1)
|
||||
self.assertIsInstance(LogBeginner.observers[0], FilteringLogObserver)
|
||||
|
||||
# Check log level predicate with the correct default log level
|
||||
self.assertEqual(len(MockFilteringLogObserver.predicates), 1)
|
||||
self.assertIsInstance(
|
||||
MockFilteringLogObserver.predicates[0],
|
||||
LogLevelFilterPredicate
|
||||
)
|
||||
self.assertIdentical(
|
||||
MockFilteringLogObserver.predicates[0].defaultLogLevel,
|
||||
LogLevel.critical
|
||||
)
|
||||
|
||||
# Check for a file observer attached to the filtering observer
|
||||
self.assertIsInstance(
|
||||
MockFilteringLogObserver.observer, MockFileLogObserver
|
||||
)
|
||||
|
||||
# Check for the file we gave it
|
||||
self.assertIdentical(
|
||||
MockFilteringLogObserver.observer.outFile, logFile
|
||||
)
|
||||
|
||||
|
||||
def test_startReactorWithReactor(self):
|
||||
"""
|
||||
L{Runner.startReactor} with the C{reactor} argument runs the given
|
||||
reactor.
|
||||
"""
|
||||
reactor = MemoryReactor()
|
||||
runner = Runner(reactor=reactor)
|
||||
runner.startReactor()
|
||||
|
||||
self.assertTrue(reactor.hasRun)
|
||||
|
||||
|
||||
def test_startReactorWhenRunning(self):
|
||||
"""
|
||||
L{Runner.startReactor} ensures that C{whenRunning} is called with
|
||||
C{whenRunningArguments} when the reactor is running.
|
||||
"""
|
||||
self._testHook("whenRunning", "startReactor")
|
||||
|
||||
|
||||
def test_whenRunningWithArguments(self):
|
||||
"""
|
||||
L{Runner.whenRunning} calls C{whenRunning} with
|
||||
C{whenRunningArguments}.
|
||||
"""
|
||||
self._testHook("whenRunning")
|
||||
|
||||
|
||||
def test_reactorExitedWithArguments(self):
|
||||
"""
|
||||
L{Runner.whenRunning} calls C{reactorExited} with
|
||||
C{reactorExitedArguments}.
|
||||
"""
|
||||
self._testHook("reactorExited")
|
||||
|
||||
|
||||
def _testHook(self, methodName, callerName=None):
|
||||
"""
|
||||
Verify that the named hook is run with the expected arguments as
|
||||
specified by the arguments used to create the L{Runner}, when the
|
||||
specified caller is invoked.
|
||||
|
||||
@param methodName: The name of the hook to verify.
|
||||
@type methodName: L{str}
|
||||
|
||||
@param callerName: The name of the method that is expected to cause the
|
||||
hook to be called.
|
||||
If C{None}, use the L{Runner} method with the same name as the
|
||||
hook.
|
||||
@type callerName: L{str}
|
||||
"""
|
||||
if callerName is None:
|
||||
callerName = methodName
|
||||
|
||||
arguments = dict(a=object(), b=object(), c=object())
|
||||
argumentsSeen = []
|
||||
|
||||
def hook(**arguments):
|
||||
argumentsSeen.append(arguments)
|
||||
|
||||
runnerArguments = {
|
||||
methodName: hook,
|
||||
"{}Arguments".format(methodName): arguments.copy(),
|
||||
}
|
||||
runner = Runner(reactor=MemoryReactor(), **runnerArguments)
|
||||
|
||||
hookCaller = getattr(runner, callerName)
|
||||
hookCaller()
|
||||
|
||||
self.assertEqual(len(argumentsSeen), 1)
|
||||
self.assertEqual(argumentsSeen[0], arguments)
|
||||
|
||||
|
||||
|
||||
@attrs(frozen=True)
|
||||
class DummyRunner(Runner):
|
||||
"""
|
||||
Stub for L{Runner}.
|
||||
|
||||
Keep track of calls to some methods without actually doing anything.
|
||||
"""
|
||||
|
||||
calledMethods = attrib(default=Factory(list))
|
||||
|
||||
|
||||
def killIfRequested(self):
|
||||
self.calledMethods.append("killIfRequested")
|
||||
|
||||
|
||||
def startLogging(self):
|
||||
self.calledMethods.append("startLogging")
|
||||
|
||||
|
||||
def startReactor(self):
|
||||
self.calledMethods.append("startReactor")
|
||||
|
||||
|
||||
def reactorExited(self):
|
||||
self.calledMethods.append("reactorExited")
|
||||
|
||||
|
||||
|
||||
class DummyPIDFile(NonePIDFile):
|
||||
"""
|
||||
Stub for L{PIDFile}.
|
||||
|
||||
Tracks context manager entry/exit without doing anything.
|
||||
"""
|
||||
def __init__(self):
|
||||
NonePIDFile.__init__(self)
|
||||
|
||||
self.entered = False
|
||||
self.exited = False
|
||||
|
||||
|
||||
def __enter__(self):
|
||||
self.entered = True
|
||||
return self
|
||||
|
||||
|
||||
def __exit__(self, excType, excValue, traceback):
|
||||
self.exited = True
|
||||
|
||||
|
||||
|
||||
class DummyExit(object):
|
||||
"""
|
||||
Stub for L{exit} that remembers whether it's been called and, if it has,
|
||||
what arguments it was given.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.exited = False
|
||||
|
||||
|
||||
def __call__(self, status, message=None):
|
||||
assert not self.exited
|
||||
|
||||
self.status = status
|
||||
self.message = message
|
||||
self.exited = True
|
||||
|
||||
|
||||
|
||||
class DummyKill(object):
|
||||
"""
|
||||
Stub for L{os.kill} that remembers whether it's been called and, if it has,
|
||||
what arguments it was given.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
|
||||
|
||||
def __call__(self, pid, sig):
|
||||
self.calls.append((pid, sig))
|
||||
|
||||
|
||||
|
||||
class DummyStandardIO(object):
|
||||
"""
|
||||
Stub for L{sys} which provides L{BytesIO} streams as stdout and stderr.
|
||||
"""
|
||||
|
||||
def __init__(self, stdout, stderr):
|
||||
self.stdout = stdout
|
||||
self.stderr = stderr
|
||||
|
||||
|
||||
|
||||
class DummyWarningsModule(object):
|
||||
"""
|
||||
Stub for L{warnings} which provides a C{showwarning} method that is a no-op.
|
||||
"""
|
||||
|
||||
def showwarning(*args, **kwargs):
|
||||
"""
|
||||
Do nothing.
|
||||
|
||||
@param args: ignored.
|
||||
@param kwargs: ignored.
|
||||
"""
|
||||
@@ -0,0 +1,424 @@
|
||||
# -*- test-case-name: twisted.application.test.test_service -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Service architecture for Twisted.
|
||||
|
||||
Services are arranged in a hierarchy. At the leafs of the hierarchy,
|
||||
the services which actually interact with the outside world are started.
|
||||
Services can be named or anonymous -- usually, they will be named if
|
||||
there is need to access them through the hierarchy (from a parent or
|
||||
a sibling).
|
||||
|
||||
Maintainer: Moshe Zadka
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
from zope.interface import implementer, Interface, Attribute
|
||||
|
||||
from twisted.persisted import sob
|
||||
from twisted.python.reflect import namedAny
|
||||
from twisted.python import components
|
||||
from twisted.python._oldstyle import _oldStyle
|
||||
from twisted.internet import defer
|
||||
from twisted.plugin import IPlugin
|
||||
|
||||
|
||||
class IServiceMaker(Interface):
|
||||
"""
|
||||
An object which can be used to construct services in a flexible
|
||||
way.
|
||||
|
||||
This interface should most often be implemented along with
|
||||
L{twisted.plugin.IPlugin}, and will most often be used by the
|
||||
'twistd' command.
|
||||
"""
|
||||
tapname = Attribute(
|
||||
"A short string naming this Twisted plugin, for example 'web' or "
|
||||
"'pencil'. This name will be used as the subcommand of 'twistd'.")
|
||||
|
||||
description = Attribute(
|
||||
"A brief summary of the features provided by this "
|
||||
"Twisted application plugin.")
|
||||
|
||||
options = Attribute(
|
||||
"A C{twisted.python.usage.Options} subclass defining the "
|
||||
"configuration options for this application.")
|
||||
|
||||
|
||||
def makeService(options):
|
||||
"""
|
||||
Create and return an object providing
|
||||
L{twisted.application.service.IService}.
|
||||
|
||||
@param options: A mapping (typically a C{dict} or
|
||||
L{twisted.python.usage.Options} instance) of configuration
|
||||
options to desired configuration values.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
@implementer(IPlugin, IServiceMaker)
|
||||
class ServiceMaker(object):
|
||||
"""
|
||||
Utility class to simplify the definition of L{IServiceMaker} plugins.
|
||||
"""
|
||||
def __init__(self, name, module, description, tapname):
|
||||
self.name = name
|
||||
self.module = module
|
||||
self.description = description
|
||||
self.tapname = tapname
|
||||
|
||||
|
||||
def options():
|
||||
def get(self):
|
||||
return namedAny(self.module).Options
|
||||
return get,
|
||||
options = property(*options())
|
||||
|
||||
|
||||
def makeService():
|
||||
def get(self):
|
||||
return namedAny(self.module).makeService
|
||||
return get,
|
||||
makeService = property(*makeService())
|
||||
|
||||
|
||||
|
||||
class IService(Interface):
|
||||
"""
|
||||
A service.
|
||||
|
||||
Run start-up and shut-down code at the appropriate times.
|
||||
"""
|
||||
|
||||
name = Attribute(
|
||||
"A C{str} which is the name of the service or C{None}.")
|
||||
|
||||
running = Attribute(
|
||||
"A C{boolean} which indicates whether the service is running.")
|
||||
|
||||
parent = Attribute(
|
||||
"An C{IServiceCollection} which is the parent or C{None}.")
|
||||
|
||||
def setName(name):
|
||||
"""
|
||||
Set the name of the service.
|
||||
|
||||
@type name: C{str}
|
||||
@raise RuntimeError: Raised if the service already has a parent.
|
||||
"""
|
||||
|
||||
def setServiceParent(parent):
|
||||
"""
|
||||
Set the parent of the service. This method is responsible for setting
|
||||
the C{parent} attribute on this service (the child service).
|
||||
|
||||
@type parent: L{IServiceCollection}
|
||||
@raise RuntimeError: Raised if the service already has a parent
|
||||
or if the service has a name and the parent already has a child
|
||||
by that name.
|
||||
"""
|
||||
|
||||
def disownServiceParent():
|
||||
"""
|
||||
Use this API to remove an L{IService} from an L{IServiceCollection}.
|
||||
|
||||
This method is used symmetrically with L{setServiceParent} in that it
|
||||
sets the C{parent} attribute on the child.
|
||||
|
||||
@rtype: L{Deferred<defer.Deferred>}
|
||||
@return: a L{Deferred<defer.Deferred>} which is triggered when the
|
||||
service has finished shutting down. If shutting down is immediate,
|
||||
a value can be returned (usually, L{None}).
|
||||
"""
|
||||
|
||||
def startService():
|
||||
"""
|
||||
Start the service.
|
||||
"""
|
||||
|
||||
def stopService():
|
||||
"""
|
||||
Stop the service.
|
||||
|
||||
@rtype: L{Deferred<defer.Deferred>}
|
||||
@return: a L{Deferred<defer.Deferred>} which is triggered when the
|
||||
service has finished shutting down. If shutting down is immediate,
|
||||
a value can be returned (usually, L{None}).
|
||||
"""
|
||||
|
||||
def privilegedStartService():
|
||||
"""
|
||||
Do preparation work for starting the service.
|
||||
|
||||
Here things which should be done before changing directory,
|
||||
root or shedding privileges are done.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
@implementer(IService)
|
||||
class Service(object):
|
||||
"""
|
||||
Base class for services.
|
||||
|
||||
Most services should inherit from this class. It handles the
|
||||
book-keeping responsibilities of starting and stopping, as well
|
||||
as not serializing this book-keeping information.
|
||||
"""
|
||||
|
||||
running = 0
|
||||
name = None
|
||||
parent = None
|
||||
|
||||
def __getstate__(self):
|
||||
dict = self.__dict__.copy()
|
||||
if "running" in dict:
|
||||
del dict['running']
|
||||
return dict
|
||||
|
||||
def setName(self, name):
|
||||
if self.parent is not None:
|
||||
raise RuntimeError("cannot change name when parent exists")
|
||||
self.name = name
|
||||
|
||||
def setServiceParent(self, parent):
|
||||
if self.parent is not None:
|
||||
self.disownServiceParent()
|
||||
parent = IServiceCollection(parent, parent)
|
||||
self.parent = parent
|
||||
self.parent.addService(self)
|
||||
|
||||
def disownServiceParent(self):
|
||||
d = self.parent.removeService(self)
|
||||
self.parent = None
|
||||
return d
|
||||
|
||||
def privilegedStartService(self):
|
||||
pass
|
||||
|
||||
def startService(self):
|
||||
self.running = 1
|
||||
|
||||
def stopService(self):
|
||||
self.running = 0
|
||||
|
||||
|
||||
|
||||
class IServiceCollection(Interface):
|
||||
"""
|
||||
Collection of services.
|
||||
|
||||
Contain several services, and manage their start-up/shut-down.
|
||||
Services can be accessed by name if they have a name, and it
|
||||
is always possible to iterate over them.
|
||||
"""
|
||||
|
||||
def getServiceNamed(name):
|
||||
"""
|
||||
Get the child service with a given name.
|
||||
|
||||
@type name: C{str}
|
||||
@rtype: L{IService}
|
||||
@raise KeyError: Raised if the service has no child with the
|
||||
given name.
|
||||
"""
|
||||
|
||||
def __iter__():
|
||||
"""
|
||||
Get an iterator over all child services.
|
||||
"""
|
||||
|
||||
def addService(service):
|
||||
"""
|
||||
Add a child service.
|
||||
|
||||
Only implementations of L{IService.setServiceParent} should use this
|
||||
method.
|
||||
|
||||
@type service: L{IService}
|
||||
@raise RuntimeError: Raised if the service has a child with
|
||||
the given name.
|
||||
"""
|
||||
|
||||
def removeService(service):
|
||||
"""
|
||||
Remove a child service.
|
||||
|
||||
Only implementations of L{IService.disownServiceParent} should
|
||||
use this method.
|
||||
|
||||
@type service: L{IService}
|
||||
@raise ValueError: Raised if the given service is not a child.
|
||||
@rtype: L{Deferred<defer.Deferred>}
|
||||
@return: a L{Deferred<defer.Deferred>} which is triggered when the
|
||||
service has finished shutting down. If shutting down is immediate,
|
||||
a value can be returned (usually, L{None}).
|
||||
"""
|
||||
|
||||
|
||||
|
||||
@implementer(IServiceCollection)
|
||||
class MultiService(Service):
|
||||
"""
|
||||
Straightforward Service Container.
|
||||
|
||||
Hold a collection of services, and manage them in a simplistic
|
||||
way. No service will wait for another, but this object itself
|
||||
will not finish shutting down until all of its child services
|
||||
will finish.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.services = []
|
||||
self.namedServices = {}
|
||||
self.parent = None
|
||||
|
||||
def privilegedStartService(self):
|
||||
Service.privilegedStartService(self)
|
||||
for service in self:
|
||||
service.privilegedStartService()
|
||||
|
||||
def startService(self):
|
||||
Service.startService(self)
|
||||
for service in self:
|
||||
service.startService()
|
||||
|
||||
def stopService(self):
|
||||
Service.stopService(self)
|
||||
l = []
|
||||
services = list(self)
|
||||
services.reverse()
|
||||
for service in services:
|
||||
l.append(defer.maybeDeferred(service.stopService))
|
||||
return defer.DeferredList(l)
|
||||
|
||||
def getServiceNamed(self, name):
|
||||
return self.namedServices[name]
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self.services)
|
||||
|
||||
def addService(self, service):
|
||||
if service.name is not None:
|
||||
if service.name in self.namedServices:
|
||||
raise RuntimeError("cannot have two services with same name"
|
||||
" '%s'" % service.name)
|
||||
self.namedServices[service.name] = service
|
||||
self.services.append(service)
|
||||
if self.running:
|
||||
# It may be too late for that, but we will do our best
|
||||
service.privilegedStartService()
|
||||
service.startService()
|
||||
|
||||
def removeService(self, service):
|
||||
if service.name:
|
||||
del self.namedServices[service.name]
|
||||
self.services.remove(service)
|
||||
if self.running:
|
||||
# Returning this so as not to lose information from the
|
||||
# MultiService.stopService deferred.
|
||||
return service.stopService()
|
||||
else:
|
||||
return None
|
||||
|
||||
|
||||
|
||||
class IProcess(Interface):
|
||||
"""
|
||||
Process running parameters.
|
||||
|
||||
Represents parameters for how processes should be run.
|
||||
"""
|
||||
processName = Attribute(
|
||||
"""
|
||||
A C{str} giving the name the process should have in ps (or L{None}
|
||||
to leave the name alone).
|
||||
""")
|
||||
|
||||
uid = Attribute(
|
||||
"""
|
||||
An C{int} giving the user id as which the process should run (or
|
||||
L{None} to leave the UID alone).
|
||||
""")
|
||||
|
||||
gid = Attribute(
|
||||
"""
|
||||
An C{int} giving the group id as which the process should run (or
|
||||
L{None} to leave the GID alone).
|
||||
""")
|
||||
|
||||
|
||||
|
||||
@implementer(IProcess)
|
||||
@_oldStyle
|
||||
class Process:
|
||||
"""
|
||||
Process running parameters.
|
||||
|
||||
Sets up uid/gid in the constructor, and has a default
|
||||
of L{None} as C{processName}.
|
||||
"""
|
||||
processName = None
|
||||
|
||||
def __init__(self, uid=None, gid=None):
|
||||
"""
|
||||
Set uid and gid.
|
||||
|
||||
@param uid: The user ID as whom to execute the process. If
|
||||
this is L{None}, no attempt will be made to change the UID.
|
||||
|
||||
@param gid: The group ID as whom to execute the process. If
|
||||
this is L{None}, no attempt will be made to change the GID.
|
||||
"""
|
||||
self.uid = uid
|
||||
self.gid = gid
|
||||
|
||||
|
||||
|
||||
def Application(name, uid=None, gid=None):
|
||||
"""
|
||||
Return a compound class.
|
||||
|
||||
Return an object supporting the L{IService}, L{IServiceCollection},
|
||||
L{IProcess} and L{sob.IPersistable} interfaces, with the given
|
||||
parameters. Always access the return value by explicit casting to
|
||||
one of the interfaces.
|
||||
"""
|
||||
ret = components.Componentized()
|
||||
availableComponents = [MultiService(), Process(uid, gid),
|
||||
sob.Persistent(ret, name)]
|
||||
|
||||
for comp in availableComponents:
|
||||
ret.addComponent(comp, ignoreClass=1)
|
||||
IService(ret).setName(name)
|
||||
return ret
|
||||
|
||||
|
||||
|
||||
def loadApplication(filename, kind, passphrase=None):
|
||||
"""
|
||||
Load Application from a given file.
|
||||
|
||||
The serialization format it was saved in should be given as
|
||||
C{kind}, and is one of C{pickle}, C{source}, C{xml} or C{python}. If
|
||||
C{passphrase} is given, the application was encrypted with the
|
||||
given passphrase.
|
||||
|
||||
@type filename: C{str}
|
||||
@type kind: C{str}
|
||||
@type passphrase: C{str}
|
||||
"""
|
||||
if kind == 'python':
|
||||
application = sob.loadValueFromFile(filename, 'application')
|
||||
else:
|
||||
application = sob.load(filename, kind)
|
||||
return application
|
||||
|
||||
|
||||
__all__ = ['IServiceMaker', 'IService', 'Service',
|
||||
'IServiceCollection', 'MultiService',
|
||||
'IProcess', 'Process', 'Application', 'loadApplication']
|
||||
@@ -0,0 +1,6 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.internet.application}.
|
||||
"""
|
||||
@@ -0,0 +1,128 @@
|
||||
# -*- test-case-name: twisted.application.twist.test.test_twist -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Run a Twisted application.
|
||||
"""
|
||||
|
||||
import sys
|
||||
|
||||
from twisted.python.usage import UsageError
|
||||
from ..service import Application, IService
|
||||
from ..runner._exit import exit, ExitStatus
|
||||
from ..runner._runner import Runner
|
||||
from ._options import TwistOptions
|
||||
from twisted.application.app import _exitWithSignal
|
||||
from twisted.internet.interfaces import _ISupportsExitSignalCapturing
|
||||
|
||||
|
||||
|
||||
class Twist(object):
|
||||
"""
|
||||
Run a Twisted application.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def options(argv):
|
||||
"""
|
||||
Parse command line options.
|
||||
|
||||
@param argv: Command line arguments.
|
||||
@type argv: L{list}
|
||||
|
||||
@return: The parsed options.
|
||||
@rtype: L{TwistOptions}
|
||||
"""
|
||||
options = TwistOptions()
|
||||
|
||||
try:
|
||||
options.parseOptions(argv[1:])
|
||||
except UsageError as e:
|
||||
exit(ExitStatus.EX_USAGE, "Error: {}\n\n{}".format(e, options))
|
||||
|
||||
return options
|
||||
|
||||
|
||||
@staticmethod
|
||||
def service(plugin, options):
|
||||
"""
|
||||
Create the application service.
|
||||
|
||||
@param plugin: The name of the plugin that implements the service
|
||||
application to run.
|
||||
@type plugin: L{str}
|
||||
|
||||
@param options: Options to pass to the application.
|
||||
@type options: L{twisted.python.usage.Options}
|
||||
|
||||
@return: The created application service.
|
||||
@rtype: L{IService}
|
||||
"""
|
||||
service = plugin.makeService(options)
|
||||
application = Application(plugin.tapname)
|
||||
service.setServiceParent(application)
|
||||
|
||||
return IService(application)
|
||||
|
||||
|
||||
@staticmethod
|
||||
def startService(reactor, service):
|
||||
"""
|
||||
Start the application service.
|
||||
|
||||
@param reactor: The reactor to run the service with.
|
||||
@type reactor: L{twisted.internet.interfaces.IReactorCore}
|
||||
|
||||
@param service: The application service to run.
|
||||
@type service: L{IService}
|
||||
"""
|
||||
service.startService()
|
||||
|
||||
# Ask the reactor to stop the service before shutting down
|
||||
reactor.addSystemEventTrigger(
|
||||
"before", "shutdown", service.stopService
|
||||
)
|
||||
|
||||
|
||||
@staticmethod
|
||||
def run(twistOptions):
|
||||
"""
|
||||
Run the application service.
|
||||
|
||||
@param twistOptions: Command line options to convert to runner
|
||||
arguments.
|
||||
@type twistOptions: L{TwistOptions}
|
||||
"""
|
||||
runner = Runner(
|
||||
reactor=twistOptions["reactor"],
|
||||
defaultLogLevel=twistOptions["logLevel"],
|
||||
logFile=twistOptions["logFile"],
|
||||
fileLogObserverFactory=twistOptions["fileLogObserverFactory"],
|
||||
)
|
||||
runner.run()
|
||||
reactor = twistOptions["reactor"]
|
||||
if _ISupportsExitSignalCapturing.providedBy(reactor):
|
||||
if reactor._exitSignal is not None:
|
||||
_exitWithSignal(reactor._exitSignal)
|
||||
|
||||
|
||||
@classmethod
|
||||
def main(cls, argv=sys.argv):
|
||||
"""
|
||||
Executable entry point for L{Twist}.
|
||||
Processes options and run a twisted reactor with a service.
|
||||
|
||||
@param argv: Command line arguments.
|
||||
@type argv: L{list}
|
||||
"""
|
||||
options = cls.options(argv)
|
||||
|
||||
reactor = options["reactor"]
|
||||
service = cls.service(
|
||||
plugin=options.plugins[options.subCommand],
|
||||
options=options.subOptions,
|
||||
)
|
||||
|
||||
cls.startService(reactor, service)
|
||||
cls.run(options)
|
||||
@@ -0,0 +1,45 @@
|
||||
# -*- test-case-name: twisted.conch.test.test_conch -*-
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
from zope.interface import implementer
|
||||
|
||||
from twisted.conch.error import ConchError
|
||||
from twisted.conch.interfaces import IConchUser
|
||||
from twisted.conch.ssh.connection import OPEN_UNKNOWN_CHANNEL_TYPE
|
||||
from twisted.python import log
|
||||
from twisted.python.compat import nativeString
|
||||
|
||||
|
||||
@implementer(IConchUser)
|
||||
class ConchUser:
|
||||
def __init__(self):
|
||||
self.channelLookup = {}
|
||||
self.subsystemLookup = {}
|
||||
|
||||
|
||||
def lookupChannel(self, channelType, windowSize, maxPacket, data):
|
||||
klass = self.channelLookup.get(channelType, None)
|
||||
if not klass:
|
||||
raise ConchError(OPEN_UNKNOWN_CHANNEL_TYPE, "unknown channel")
|
||||
else:
|
||||
return klass(remoteWindow=windowSize,
|
||||
remoteMaxPacket=maxPacket,
|
||||
data=data, avatar=self)
|
||||
|
||||
|
||||
def lookupSubsystem(self, subsystem, data):
|
||||
log.msg(repr(self.subsystemLookup))
|
||||
klass = self.subsystemLookup.get(subsystem, None)
|
||||
if not klass:
|
||||
return False
|
||||
return klass(data, avatar=self)
|
||||
|
||||
|
||||
def gotGlobalRequest(self, requestType, data):
|
||||
# XXX should this use method dispatch?
|
||||
requestType = nativeString(requestType.replace(b'-', b'_'))
|
||||
f = getattr(self, "global_%s" % requestType, None)
|
||||
if not f:
|
||||
return 0
|
||||
return f(data)
|
||||
@@ -0,0 +1,103 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
#
|
||||
from twisted.conch.ssh.transport import SSHClientTransport, SSHCiphers
|
||||
from twisted.python import usage
|
||||
from twisted.python.compat import unicode
|
||||
|
||||
import sys
|
||||
|
||||
class ConchOptions(usage.Options):
|
||||
|
||||
optParameters = [['user', 'l', None, 'Log in using this user name.'],
|
||||
['identity', 'i', None],
|
||||
['ciphers', 'c', None],
|
||||
['macs', 'm', None],
|
||||
['port', 'p', None, 'Connect to this port. Server must be on the same port.'],
|
||||
['option', 'o', None, 'Ignored OpenSSH options'],
|
||||
['host-key-algorithms', '', None],
|
||||
['known-hosts', '', None, 'File to check for host keys'],
|
||||
['user-authentications', '', None, 'Types of user authentications to use.'],
|
||||
['logfile', '', None, 'File to log to, or - for stdout'],
|
||||
]
|
||||
|
||||
optFlags = [['version', 'V', 'Display version number only.'],
|
||||
['compress', 'C', 'Enable compression.'],
|
||||
['log', 'v', 'Enable logging (defaults to stderr)'],
|
||||
['nox11', 'x', 'Disable X11 connection forwarding (default)'],
|
||||
['agent', 'A', 'Enable authentication agent forwarding'],
|
||||
['noagent', 'a', 'Disable authentication agent forwarding (default)'],
|
||||
['reconnect', 'r', 'Reconnect to the server if the connection is lost.'],
|
||||
]
|
||||
|
||||
compData = usage.Completions(
|
||||
mutuallyExclusive=[("agent", "noagent")],
|
||||
optActions={
|
||||
"user": usage.CompleteUsernames(),
|
||||
"ciphers": usage.CompleteMultiList(
|
||||
SSHCiphers.cipherMap.keys(),
|
||||
descr='ciphers to choose from'),
|
||||
"macs": usage.CompleteMultiList(
|
||||
SSHCiphers.macMap.keys(),
|
||||
descr='macs to choose from'),
|
||||
"host-key-algorithms": usage.CompleteMultiList(
|
||||
SSHClientTransport.supportedPublicKeys,
|
||||
descr='host key algorithms to choose from'),
|
||||
#"user-authentications": usage.CompleteMultiList(?
|
||||
# descr='user authentication types' ),
|
||||
},
|
||||
extraActions=[usage.CompleteUserAtHost(),
|
||||
usage.Completer(descr="command"),
|
||||
usage.Completer(descr='argument',
|
||||
repeat=True)]
|
||||
)
|
||||
|
||||
def __init__(self, *args, **kw):
|
||||
usage.Options.__init__(self, *args, **kw)
|
||||
self.identitys = []
|
||||
self.conns = None
|
||||
|
||||
def opt_identity(self, i):
|
||||
"""Identity for public-key authentication"""
|
||||
self.identitys.append(i)
|
||||
|
||||
def opt_ciphers(self, ciphers):
|
||||
"Select encryption algorithms"
|
||||
ciphers = ciphers.split(',')
|
||||
for cipher in ciphers:
|
||||
if cipher not in SSHCiphers.cipherMap:
|
||||
sys.exit("Unknown cipher type '%s'" % cipher)
|
||||
self['ciphers'] = ciphers
|
||||
|
||||
|
||||
def opt_macs(self, macs):
|
||||
"Specify MAC algorithms"
|
||||
if isinstance(macs, unicode):
|
||||
macs = macs.encode("utf-8")
|
||||
macs = macs.split(b',')
|
||||
for mac in macs:
|
||||
if mac not in SSHCiphers.macMap:
|
||||
sys.exit("Unknown mac type '%r'" % mac)
|
||||
self['macs'] = macs
|
||||
|
||||
def opt_host_key_algorithms(self, hkas):
|
||||
"Select host key algorithms"
|
||||
if isinstance(hkas, unicode):
|
||||
hkas = hkas.encode("utf-8")
|
||||
hkas = hkas.split(b',')
|
||||
for hka in hkas:
|
||||
if hka not in SSHClientTransport.supportedPublicKeys:
|
||||
sys.exit("Unknown host key type '%r'" % hka)
|
||||
self['host-key-algorithms'] = hkas
|
||||
|
||||
def opt_user_authentications(self, uas):
|
||||
"Choose how to authenticate to the remote server"
|
||||
if isinstance(uas, unicode):
|
||||
uas = uas.encode("utf-8")
|
||||
self['user-authentications'] = uas.split(b',')
|
||||
|
||||
# def opt_compress(self):
|
||||
# "Enable compression"
|
||||
# self.enableCompression = 1
|
||||
# SSHClientTransport.supportedCompressions[0:1] = ['zlib']
|
||||
@@ -0,0 +1,103 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
An error to represent bad things happening in Conch.
|
||||
|
||||
Maintainer: Paul Swartz
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
from twisted.cred.error import UnauthorizedLogin
|
||||
|
||||
|
||||
class ConchError(Exception):
|
||||
def __init__(self, value, data = None):
|
||||
Exception.__init__(self, value, data)
|
||||
self.value = value
|
||||
self.data = data
|
||||
|
||||
|
||||
|
||||
class NotEnoughAuthentication(Exception):
|
||||
"""
|
||||
This is thrown if the authentication is valid, but is not enough to
|
||||
successfully verify the user. i.e. don't retry this type of
|
||||
authentication, try another one.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class ValidPublicKey(UnauthorizedLogin):
|
||||
"""
|
||||
Raised by public key checkers when they receive public key credentials
|
||||
that don't contain a signature at all, but are valid in every other way.
|
||||
(e.g. the public key matches one in the user's authorized_keys file).
|
||||
|
||||
Protocol code (eg
|
||||
L{SSHUserAuthServer<twisted.conch.ssh.userauth.SSHUserAuthServer>}) which
|
||||
attempts to log in using
|
||||
L{ISSHPrivateKey<twisted.cred.credentials.ISSHPrivateKey>} credentials
|
||||
should be prepared to handle a failure of this type by telling the user to
|
||||
re-authenticate using the same key and to include a signature with the new
|
||||
attempt.
|
||||
|
||||
See U{http://www.ietf.org/rfc/rfc4252.txt} section 7 for more details.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class IgnoreAuthentication(Exception):
|
||||
"""
|
||||
This is thrown to let the UserAuthServer know it doesn't need to handle the
|
||||
authentication anymore.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class MissingKeyStoreError(Exception):
|
||||
"""
|
||||
Raised if an SSHAgentServer starts receiving data without its factory
|
||||
providing a keys dict on which to read/write key data.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class UserRejectedKey(Exception):
|
||||
"""
|
||||
The user interactively rejected a key.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class InvalidEntry(Exception):
|
||||
"""
|
||||
An entry in a known_hosts file could not be interpreted as a valid entry.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class HostKeyChanged(Exception):
|
||||
"""
|
||||
The host key of a remote host has changed.
|
||||
|
||||
@ivar offendingEntry: The entry which contains the persistent host key that
|
||||
disagrees with the given host key.
|
||||
|
||||
@type offendingEntry: L{twisted.conch.interfaces.IKnownHostEntry}
|
||||
|
||||
@ivar path: a reference to the known_hosts file that the offending entry
|
||||
was loaded from
|
||||
|
||||
@type path: L{twisted.python.filepath.FilePath}
|
||||
|
||||
@ivar lineno: The line number of the offending entry in the given path.
|
||||
|
||||
@type lineno: L{int}
|
||||
"""
|
||||
def __init__(self, offendingEntry, path, lineno):
|
||||
Exception.__init__(self)
|
||||
self.offendingEntry = offendingEntry
|
||||
self.path = path
|
||||
self.lineno = lineno
|
||||
@@ -0,0 +1,4 @@
|
||||
"""
|
||||
Insults: a replacement for Curses/S-Lang.
|
||||
|
||||
Very basic at the moment."""
|
||||
@@ -0,0 +1,165 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
TAP plugin for creating telnet- and ssh-accessible manhole servers.
|
||||
|
||||
@author: Jp Calderone
|
||||
"""
|
||||
|
||||
from zope.interface import implementer
|
||||
|
||||
from twisted.internet import protocol
|
||||
from twisted.application import service, strports
|
||||
from twisted.cred import portal, checkers
|
||||
from twisted.python import usage, filepath
|
||||
|
||||
from twisted.conch import manhole, manhole_ssh, telnet
|
||||
from twisted.conch.insults import insults
|
||||
from twisted.conch.ssh import keys
|
||||
|
||||
|
||||
|
||||
class makeTelnetProtocol:
|
||||
def __init__(self, portal):
|
||||
self.portal = portal
|
||||
|
||||
def __call__(self):
|
||||
auth = telnet.AuthenticatingTelnetProtocol
|
||||
args = (self.portal,)
|
||||
return telnet.TelnetTransport(auth, *args)
|
||||
|
||||
|
||||
|
||||
class chainedProtocolFactory:
|
||||
def __init__(self, namespace):
|
||||
self.namespace = namespace
|
||||
|
||||
def __call__(self):
|
||||
return insults.ServerProtocol(manhole.ColoredManhole, self.namespace)
|
||||
|
||||
|
||||
|
||||
@implementer(portal.IRealm)
|
||||
class _StupidRealm:
|
||||
def __init__(self, proto, *a, **kw):
|
||||
self.protocolFactory = proto
|
||||
self.protocolArgs = a
|
||||
self.protocolKwArgs = kw
|
||||
|
||||
def requestAvatar(self, avatarId, *interfaces):
|
||||
if telnet.ITelnetProtocol in interfaces:
|
||||
return (telnet.ITelnetProtocol,
|
||||
self.protocolFactory(*self.protocolArgs,
|
||||
**self.protocolKwArgs),
|
||||
lambda: None)
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
|
||||
class Options(usage.Options):
|
||||
optParameters = [
|
||||
["telnetPort", "t", None,
|
||||
("strports description of the address on which to listen for telnet "
|
||||
"connections")],
|
||||
["sshPort", "s", None,
|
||||
("strports description of the address on which to listen for ssh "
|
||||
"connections")],
|
||||
["passwd", "p", "/etc/passwd",
|
||||
"name of a passwd(5)-format username/password file"],
|
||||
["sshKeyDir", None, "<USER DATA DIR>",
|
||||
"Directory where the autogenerated SSH key is kept."],
|
||||
["sshKeyName", None, "server.key",
|
||||
"Filename of the autogenerated SSH key."],
|
||||
["sshKeySize", None, 4096,
|
||||
"Size of the automatically generated SSH key."],
|
||||
]
|
||||
|
||||
def __init__(self):
|
||||
usage.Options.__init__(self)
|
||||
self['namespace'] = None
|
||||
|
||||
def postOptions(self):
|
||||
if self['telnetPort'] is None and self['sshPort'] is None:
|
||||
raise usage.UsageError(
|
||||
"At least one of --telnetPort and --sshPort must be specified")
|
||||
|
||||
|
||||
|
||||
def makeService(options):
|
||||
"""
|
||||
Create a manhole server service.
|
||||
|
||||
@type options: L{dict}
|
||||
@param options: A mapping describing the configuration of
|
||||
the desired service. Recognized key/value pairs are::
|
||||
|
||||
"telnetPort": strports description of the address on which
|
||||
to listen for telnet connections. If None,
|
||||
no telnet service will be started.
|
||||
|
||||
"sshPort": strports description of the address on which to
|
||||
listen for ssh connections. If None, no ssh
|
||||
service will be started.
|
||||
|
||||
"namespace": dictionary containing desired initial locals
|
||||
for manhole connections. If None, an empty
|
||||
dictionary will be used.
|
||||
|
||||
"passwd": Name of a passwd(5)-format username/password file.
|
||||
|
||||
"sshKeyDir": The folder that the SSH server key will be kept in.
|
||||
|
||||
"sshKeyName": The filename of the key.
|
||||
|
||||
"sshKeySize": The size of the key, in bits. Default is 4096.
|
||||
|
||||
@rtype: L{twisted.application.service.IService}
|
||||
@return: A manhole service.
|
||||
"""
|
||||
svc = service.MultiService()
|
||||
|
||||
namespace = options['namespace']
|
||||
if namespace is None:
|
||||
namespace = {}
|
||||
|
||||
checker = checkers.FilePasswordDB(options['passwd'])
|
||||
|
||||
if options['telnetPort']:
|
||||
telnetRealm = _StupidRealm(telnet.TelnetBootstrapProtocol,
|
||||
insults.ServerProtocol,
|
||||
manhole.ColoredManhole,
|
||||
namespace)
|
||||
|
||||
telnetPortal = portal.Portal(telnetRealm, [checker])
|
||||
|
||||
telnetFactory = protocol.ServerFactory()
|
||||
telnetFactory.protocol = makeTelnetProtocol(telnetPortal)
|
||||
telnetService = strports.service(options['telnetPort'],
|
||||
telnetFactory)
|
||||
telnetService.setServiceParent(svc)
|
||||
|
||||
if options['sshPort']:
|
||||
sshRealm = manhole_ssh.TerminalRealm()
|
||||
sshRealm.chainedProtocolFactory = chainedProtocolFactory(namespace)
|
||||
|
||||
sshPortal = portal.Portal(sshRealm, [checker])
|
||||
sshFactory = manhole_ssh.ConchFactory(sshPortal)
|
||||
|
||||
if options['sshKeyDir'] != "<USER DATA DIR>":
|
||||
keyDir = options['sshKeyDir']
|
||||
else:
|
||||
from twisted.python._appdirs import getDataDirectory
|
||||
keyDir = getDataDirectory()
|
||||
|
||||
keyLocation = filepath.FilePath(keyDir).child(options['sshKeyName'])
|
||||
|
||||
sshKey = keys._getPersistentRSAKey(keyLocation,
|
||||
int(options['sshKeySize']))
|
||||
sshFactory.publicKeys[b"ssh-rsa"] = sshKey
|
||||
sshFactory.privateKeys[b"ssh-rsa"] = sshKey
|
||||
|
||||
sshService = strports.service(options['sshPort'], sshFactory)
|
||||
sshService.setServiceParent(svc)
|
||||
|
||||
return svc
|
||||
@@ -0,0 +1,585 @@
|
||||
# -*- test-case-name: twisted.conch.test.test_conch -*-
|
||||
#
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
#
|
||||
# $Id: conch.py,v 1.65 2004/03/11 00:29:14 z3p Exp $
|
||||
|
||||
#""" Implementation module for the `conch` command.
|
||||
#"""
|
||||
from __future__ import print_function
|
||||
|
||||
from twisted.conch.client import connect, default, options
|
||||
from twisted.conch.error import ConchError
|
||||
from twisted.conch.ssh import connection, common
|
||||
from twisted.conch.ssh import session, forwarding, channel
|
||||
from twisted.internet import reactor, stdio, task
|
||||
from twisted.python import log, usage
|
||||
from twisted.python.compat import ioType, networkString, unicode
|
||||
|
||||
import os
|
||||
import sys
|
||||
import getpass
|
||||
import struct
|
||||
import tty
|
||||
import fcntl
|
||||
import signal
|
||||
|
||||
|
||||
|
||||
class ClientOptions(options.ConchOptions):
|
||||
|
||||
synopsis = """Usage: conch [options] host [command]
|
||||
"""
|
||||
longdesc = ("conch is a SSHv2 client that allows logging into a remote "
|
||||
"machine and executing commands.")
|
||||
|
||||
optParameters = [['escape', 'e', '~'],
|
||||
['localforward', 'L', None, 'listen-port:host:port Forward local port to remote address'],
|
||||
['remoteforward', 'R', None, 'listen-port:host:port Forward remote port to local address'],
|
||||
]
|
||||
|
||||
optFlags = [['null', 'n', 'Redirect input from /dev/null.'],
|
||||
['fork', 'f', 'Fork to background after authentication.'],
|
||||
['tty', 't', 'Tty; allocate a tty even if command is given.'],
|
||||
['notty', 'T', 'Do not allocate a tty.'],
|
||||
['noshell', 'N', 'Do not execute a shell or command.'],
|
||||
['subsystem', 's', 'Invoke command (mandatory) as SSH2 subsystem.'],
|
||||
]
|
||||
|
||||
compData = usage.Completions(
|
||||
mutuallyExclusive=[("tty", "notty")],
|
||||
optActions={
|
||||
"localforward": usage.Completer(descr="listen-port:host:port"),
|
||||
"remoteforward": usage.Completer(descr="listen-port:host:port")},
|
||||
extraActions=[usage.CompleteUserAtHost(),
|
||||
usage.Completer(descr="command"),
|
||||
usage.Completer(descr="argument", repeat=True)]
|
||||
)
|
||||
|
||||
localForwards = []
|
||||
remoteForwards = []
|
||||
|
||||
def opt_escape(self, esc):
|
||||
"""
|
||||
Set escape character; ``none'' = disable
|
||||
"""
|
||||
if esc == 'none':
|
||||
self['escape'] = None
|
||||
elif esc[0] == '^' and len(esc) == 2:
|
||||
self['escape'] = chr(ord(esc[1])-64)
|
||||
elif len(esc) == 1:
|
||||
self['escape'] = esc
|
||||
else:
|
||||
sys.exit("Bad escape character '{}'.".format(esc))
|
||||
|
||||
|
||||
def opt_localforward(self, f):
|
||||
"""
|
||||
Forward local port to remote address (lport:host:port)
|
||||
"""
|
||||
localPort, remoteHost, remotePort = f.split(':') # Doesn't do v6 yet
|
||||
localPort = int(localPort)
|
||||
remotePort = int(remotePort)
|
||||
self.localForwards.append((localPort, (remoteHost, remotePort)))
|
||||
|
||||
|
||||
def opt_remoteforward(self, f):
|
||||
"""
|
||||
Forward remote port to local address (rport:host:port)
|
||||
"""
|
||||
remotePort, connHost, connPort = f.split(':') # Doesn't do v6 yet
|
||||
remotePort = int(remotePort)
|
||||
connPort = int(connPort)
|
||||
self.remoteForwards.append((remotePort, (connHost, connPort)))
|
||||
|
||||
|
||||
def parseArgs(self, host, *command):
|
||||
self['host'] = host
|
||||
self['command'] = ' '.join(command)
|
||||
|
||||
|
||||
|
||||
# Rest of code in "run"
|
||||
options = None
|
||||
conn = None
|
||||
exitStatus = 0
|
||||
old = None
|
||||
_inRawMode = 0
|
||||
_savedRawMode = None
|
||||
|
||||
|
||||
|
||||
def run():
|
||||
global options, old
|
||||
args = sys.argv[1:]
|
||||
if '-l' in args: # CVS is an idiot
|
||||
i = args.index('-l')
|
||||
args = args[i:i+2]+args
|
||||
del args[i+2:i+4]
|
||||
for arg in args[:]:
|
||||
try:
|
||||
i = args.index(arg)
|
||||
if arg[:2] == '-o' and args[i+1][0] != '-':
|
||||
args[i:i+2] = [] # Suck on it scp
|
||||
except ValueError:
|
||||
pass
|
||||
options = ClientOptions()
|
||||
try:
|
||||
options.parseOptions(args)
|
||||
except usage.UsageError as u:
|
||||
print('ERROR: {}'.format(u))
|
||||
options.opt_help()
|
||||
sys.exit(1)
|
||||
if options['log']:
|
||||
if options['logfile']:
|
||||
if options['logfile'] == '-':
|
||||
f = sys.stdout
|
||||
else:
|
||||
f = open(options['logfile'], 'a+')
|
||||
else:
|
||||
f = sys.stderr
|
||||
realout = sys.stdout
|
||||
log.startLogging(f)
|
||||
sys.stdout = realout
|
||||
else:
|
||||
log.discardLogs()
|
||||
doConnect()
|
||||
fd = sys.stdin.fileno()
|
||||
try:
|
||||
old = tty.tcgetattr(fd)
|
||||
except:
|
||||
old = None
|
||||
try:
|
||||
oldUSR1 = signal.signal(signal.SIGUSR1, lambda *a: reactor.callLater(0, reConnect))
|
||||
except:
|
||||
oldUSR1 = None
|
||||
try:
|
||||
reactor.run()
|
||||
finally:
|
||||
if old:
|
||||
tty.tcsetattr(fd, tty.TCSANOW, old)
|
||||
if oldUSR1:
|
||||
signal.signal(signal.SIGUSR1, oldUSR1)
|
||||
if (options['command'] and options['tty']) or not options['notty']:
|
||||
signal.signal(signal.SIGWINCH, signal.SIG_DFL)
|
||||
if sys.stdout.isatty() and not options['command']:
|
||||
print('Connection to {} closed.'.format(options['host']))
|
||||
sys.exit(exitStatus)
|
||||
|
||||
|
||||
|
||||
def handleError():
|
||||
from twisted.python import failure
|
||||
global exitStatus
|
||||
exitStatus = 2
|
||||
reactor.callLater(0.01, _stopReactor)
|
||||
log.err(failure.Failure())
|
||||
raise
|
||||
|
||||
|
||||
|
||||
def _stopReactor():
|
||||
try:
|
||||
reactor.stop()
|
||||
except: pass
|
||||
|
||||
|
||||
|
||||
def doConnect():
|
||||
if '@' in options['host']:
|
||||
options['user'], options['host'] = options['host'].split('@', 1)
|
||||
if not options.identitys:
|
||||
options.identitys = ['~/.ssh/id_rsa', '~/.ssh/id_dsa']
|
||||
host = options['host']
|
||||
if not options['user']:
|
||||
options['user'] = getpass.getuser()
|
||||
if not options['port']:
|
||||
options['port'] = 22
|
||||
else:
|
||||
options['port'] = int(options['port'])
|
||||
host = options['host']
|
||||
port = options['port']
|
||||
vhk = default.verifyHostKey
|
||||
if not options['host-key-algorithms']:
|
||||
options['host-key-algorithms'] = default.getHostKeyAlgorithms(
|
||||
host, options)
|
||||
uao = default.SSHUserAuthClient(options['user'], options, SSHConnection())
|
||||
connect.connect(host, port, options, vhk, uao).addErrback(_ebExit)
|
||||
|
||||
|
||||
|
||||
def _ebExit(f):
|
||||
global exitStatus
|
||||
exitStatus = "conch: exiting with error {}".format(f)
|
||||
reactor.callLater(0.1, _stopReactor)
|
||||
|
||||
|
||||
|
||||
def onConnect():
|
||||
# if keyAgent and options['agent']:
|
||||
# cc = protocol.ClientCreator(reactor, SSHAgentForwardingLocal, conn)
|
||||
# cc.connectUNIX(os.environ['SSH_AUTH_SOCK'])
|
||||
if hasattr(conn.transport, 'sendIgnore'):
|
||||
_KeepAlive(conn)
|
||||
if options.localForwards:
|
||||
for localPort, hostport in options.localForwards:
|
||||
s = reactor.listenTCP(localPort,
|
||||
forwarding.SSHListenForwardingFactory(conn,
|
||||
hostport,
|
||||
SSHListenClientForwardingChannel))
|
||||
conn.localForwards.append(s)
|
||||
if options.remoteForwards:
|
||||
for remotePort, hostport in options.remoteForwards:
|
||||
log.msg('asking for remote forwarding for {}:{}'.format(
|
||||
remotePort, hostport))
|
||||
conn.requestRemoteForwarding(remotePort, hostport)
|
||||
reactor.addSystemEventTrigger('before', 'shutdown', beforeShutdown)
|
||||
if not options['noshell'] or options['agent']:
|
||||
conn.openChannel(SSHSession())
|
||||
if options['fork']:
|
||||
if os.fork():
|
||||
os._exit(0)
|
||||
os.setsid()
|
||||
for i in range(3):
|
||||
try:
|
||||
os.close(i)
|
||||
except OSError as e:
|
||||
import errno
|
||||
if e.errno != errno.EBADF:
|
||||
raise
|
||||
|
||||
|
||||
|
||||
def reConnect():
|
||||
beforeShutdown()
|
||||
conn.transport.transport.loseConnection()
|
||||
|
||||
|
||||
|
||||
def beforeShutdown():
|
||||
remoteForwards = options.remoteForwards
|
||||
for remotePort, hostport in remoteForwards:
|
||||
log.msg('cancelling {}:{}'.format(remotePort, hostport))
|
||||
conn.cancelRemoteForwarding(remotePort)
|
||||
|
||||
|
||||
|
||||
def stopConnection():
|
||||
if not options['reconnect']:
|
||||
reactor.callLater(0.1, _stopReactor)
|
||||
|
||||
|
||||
|
||||
class _KeepAlive:
|
||||
|
||||
def __init__(self, conn):
|
||||
self.conn = conn
|
||||
self.globalTimeout = None
|
||||
self.lc = task.LoopingCall(self.sendGlobal)
|
||||
self.lc.start(300)
|
||||
|
||||
|
||||
def sendGlobal(self):
|
||||
d = self.conn.sendGlobalRequest(b"conch-keep-alive@twistedmatrix.com",
|
||||
b"", wantReply=1)
|
||||
d.addBoth(self._cbGlobal)
|
||||
self.globalTimeout = reactor.callLater(30, self._ebGlobal)
|
||||
|
||||
|
||||
def _cbGlobal(self, res):
|
||||
if self.globalTimeout:
|
||||
self.globalTimeout.cancel()
|
||||
self.globalTimeout = None
|
||||
|
||||
|
||||
def _ebGlobal(self):
|
||||
if self.globalTimeout:
|
||||
self.globalTimeout = None
|
||||
self.conn.transport.loseConnection()
|
||||
|
||||
|
||||
|
||||
class SSHConnection(connection.SSHConnection):
|
||||
def serviceStarted(self):
|
||||
global conn
|
||||
conn = self
|
||||
self.localForwards = []
|
||||
self.remoteForwards = {}
|
||||
if not isinstance(self, connection.SSHConnection):
|
||||
# make these fall through
|
||||
del self.__class__.requestRemoteForwarding
|
||||
del self.__class__.cancelRemoteForwarding
|
||||
onConnect()
|
||||
|
||||
|
||||
def serviceStopped(self):
|
||||
lf = self.localForwards
|
||||
self.localForwards = []
|
||||
for s in lf:
|
||||
s.loseConnection()
|
||||
stopConnection()
|
||||
|
||||
|
||||
def requestRemoteForwarding(self, remotePort, hostport):
|
||||
data = forwarding.packGlobal_tcpip_forward(('0.0.0.0', remotePort))
|
||||
d = self.sendGlobalRequest(b'tcpip-forward', data,
|
||||
wantReply=1)
|
||||
log.msg('requesting remote forwarding {}:{}'.format(
|
||||
remotePort, hostport))
|
||||
d.addCallback(self._cbRemoteForwarding, remotePort, hostport)
|
||||
d.addErrback(self._ebRemoteForwarding, remotePort, hostport)
|
||||
|
||||
|
||||
def _cbRemoteForwarding(self, result, remotePort, hostport):
|
||||
log.msg('accepted remote forwarding {}:{}'.format(
|
||||
remotePort, hostport))
|
||||
self.remoteForwards[remotePort] = hostport
|
||||
log.msg(repr(self.remoteForwards))
|
||||
|
||||
|
||||
def _ebRemoteForwarding(self, f, remotePort, hostport):
|
||||
log.msg('remote forwarding {}:{} failed'.format(
|
||||
remotePort, hostport))
|
||||
log.msg(f)
|
||||
|
||||
|
||||
def cancelRemoteForwarding(self, remotePort):
|
||||
data = forwarding.packGlobal_tcpip_forward(('0.0.0.0', remotePort))
|
||||
self.sendGlobalRequest(b'cancel-tcpip-forward', data)
|
||||
log.msg('cancelling remote forwarding {}'.format(remotePort))
|
||||
try:
|
||||
del self.remoteForwards[remotePort]
|
||||
except Exception:
|
||||
pass
|
||||
log.msg(repr(self.remoteForwards))
|
||||
|
||||
|
||||
def channel_forwarded_tcpip(self, windowSize, maxPacket, data):
|
||||
log.msg('FTCP {!r}'.format(data))
|
||||
remoteHP, origHP = forwarding.unpackOpen_forwarded_tcpip(data)
|
||||
log.msg(self.remoteForwards)
|
||||
log.msg(remoteHP)
|
||||
if remoteHP[1] in self.remoteForwards:
|
||||
connectHP = self.remoteForwards[remoteHP[1]]
|
||||
log.msg('connect forwarding {}'.format(connectHP))
|
||||
return SSHConnectForwardingChannel(connectHP,
|
||||
remoteWindow=windowSize,
|
||||
remoteMaxPacket=maxPacket,
|
||||
conn=self)
|
||||
else:
|
||||
raise ConchError(connection.OPEN_CONNECT_FAILED,
|
||||
"don't know about that port")
|
||||
|
||||
|
||||
def channelClosed(self, channel):
|
||||
log.msg('connection closing {}'.format(channel))
|
||||
log.msg(self.channels)
|
||||
if len(self.channels) == 1: # Just us left
|
||||
log.msg('stopping connection')
|
||||
stopConnection()
|
||||
else:
|
||||
# Because of the unix thing
|
||||
self.__class__.__bases__[0].channelClosed(self, channel)
|
||||
|
||||
|
||||
|
||||
class SSHSession(channel.SSHChannel):
|
||||
|
||||
name = b'session'
|
||||
|
||||
def channelOpen(self, foo):
|
||||
log.msg('session {} open'.format(self.id))
|
||||
if options['agent']:
|
||||
d = self.conn.sendRequest(self, b'auth-agent-req@openssh.com',
|
||||
b'', wantReply=1)
|
||||
d.addBoth(lambda x: log.msg(x))
|
||||
if options['noshell']:
|
||||
return
|
||||
if (options['command'] and options['tty']) or not options['notty']:
|
||||
_enterRawMode()
|
||||
c = session.SSHSessionClient()
|
||||
if options['escape'] and not options['notty']:
|
||||
self.escapeMode = 1
|
||||
c.dataReceived = self.handleInput
|
||||
else:
|
||||
c.dataReceived = self.write
|
||||
c.connectionLost = lambda x: self.sendEOF()
|
||||
self.stdio = stdio.StandardIO(c)
|
||||
fd = 0
|
||||
if options['subsystem']:
|
||||
self.conn.sendRequest(self, b'subsystem',
|
||||
common.NS(options['command']))
|
||||
elif options['command']:
|
||||
if options['tty']:
|
||||
term = os.environ['TERM']
|
||||
winsz = fcntl.ioctl(fd, tty.TIOCGWINSZ, '12345678')
|
||||
winSize = struct.unpack('4H', winsz)
|
||||
ptyReqData = session.packRequest_pty_req(term, winSize, '')
|
||||
self.conn.sendRequest(self, b'pty-req', ptyReqData)
|
||||
signal.signal(signal.SIGWINCH, self._windowResized)
|
||||
self.conn.sendRequest(self, b'exec', common.NS(options['command']))
|
||||
else:
|
||||
if not options['notty']:
|
||||
term = os.environ['TERM']
|
||||
winsz = fcntl.ioctl(fd, tty.TIOCGWINSZ, '12345678')
|
||||
winSize = struct.unpack('4H', winsz)
|
||||
ptyReqData = session.packRequest_pty_req(term, winSize, '')
|
||||
self.conn.sendRequest(self, b'pty-req', ptyReqData)
|
||||
signal.signal(signal.SIGWINCH, self._windowResized)
|
||||
self.conn.sendRequest(self, b'shell', b'')
|
||||
#if hasattr(conn.transport, 'transport'):
|
||||
# conn.transport.transport.setTcpNoDelay(1)
|
||||
|
||||
|
||||
def handleInput(self, char):
|
||||
if char in (b'\n', b'\r'):
|
||||
self.escapeMode = 1
|
||||
self.write(char)
|
||||
elif self.escapeMode == 1 and char == options['escape']:
|
||||
self.escapeMode = 2
|
||||
elif self.escapeMode == 2:
|
||||
self.escapeMode = 1 # So we can chain escapes together
|
||||
if char == b'.': # Disconnect
|
||||
log.msg('disconnecting from escape')
|
||||
stopConnection()
|
||||
return
|
||||
elif char == b'\x1a': # ^Z, suspend
|
||||
def _():
|
||||
_leaveRawMode()
|
||||
sys.stdout.flush()
|
||||
sys.stdin.flush()
|
||||
os.kill(os.getpid(), signal.SIGTSTP)
|
||||
_enterRawMode()
|
||||
reactor.callLater(0, _)
|
||||
return
|
||||
elif char == b'R': # Rekey connection
|
||||
log.msg('rekeying connection')
|
||||
self.conn.transport.sendKexInit()
|
||||
return
|
||||
elif char == b'#': # Display connections
|
||||
self.stdio.write(
|
||||
b'\r\nThe following connections are open:\r\n')
|
||||
channels = self.conn.channels.keys()
|
||||
channels.sort()
|
||||
for channelId in channels:
|
||||
self.stdio.write(networkString(' #{} {}\r\n'.format(
|
||||
channelId,
|
||||
self.conn.channels[channelId])))
|
||||
return
|
||||
self.write(b'~' + char)
|
||||
else:
|
||||
self.escapeMode = 0
|
||||
self.write(char)
|
||||
|
||||
|
||||
def dataReceived(self, data):
|
||||
self.stdio.write(data)
|
||||
|
||||
|
||||
def extReceived(self, t, data):
|
||||
if t == connection.EXTENDED_DATA_STDERR:
|
||||
log.msg('got {} stderr data'.format(len(data)))
|
||||
if ioType(sys.stderr) == unicode:
|
||||
sys.stderr.buffer.write(data)
|
||||
else:
|
||||
sys.stderr.write(data)
|
||||
|
||||
|
||||
def eofReceived(self):
|
||||
log.msg('got eof')
|
||||
self.stdio.loseWriteConnection()
|
||||
|
||||
|
||||
def closeReceived(self):
|
||||
log.msg('remote side closed {}'.format(self))
|
||||
self.conn.sendClose(self)
|
||||
|
||||
|
||||
def closed(self):
|
||||
global old
|
||||
log.msg('closed {}'.format(self))
|
||||
log.msg(repr(self.conn.channels))
|
||||
|
||||
|
||||
def request_exit_status(self, data):
|
||||
global exitStatus
|
||||
exitStatus = int(struct.unpack('>L', data)[0])
|
||||
log.msg('exit status: {}'.format(exitStatus))
|
||||
|
||||
|
||||
def sendEOF(self):
|
||||
self.conn.sendEOF(self)
|
||||
|
||||
|
||||
def stopWriting(self):
|
||||
self.stdio.pauseProducing()
|
||||
|
||||
|
||||
def startWriting(self):
|
||||
self.stdio.resumeProducing()
|
||||
|
||||
|
||||
def _windowResized(self, *args):
|
||||
winsz = fcntl.ioctl(0, tty.TIOCGWINSZ, '12345678')
|
||||
winSize = struct.unpack('4H', winsz)
|
||||
newSize = winSize[1], winSize[0], winSize[2], winSize[3]
|
||||
self.conn.sendRequest(self, b'window-change', struct.pack('!4L', *newSize))
|
||||
|
||||
|
||||
|
||||
class SSHListenClientForwardingChannel(forwarding.SSHListenClientForwardingChannel): pass
|
||||
class SSHConnectForwardingChannel(forwarding.SSHConnectForwardingChannel): pass
|
||||
|
||||
|
||||
|
||||
def _leaveRawMode():
|
||||
global _inRawMode
|
||||
if not _inRawMode:
|
||||
return
|
||||
fd = sys.stdin.fileno()
|
||||
tty.tcsetattr(fd, tty.TCSANOW, _savedRawMode)
|
||||
_inRawMode = 0
|
||||
|
||||
|
||||
|
||||
def _enterRawMode():
|
||||
global _inRawMode, _savedRawMode
|
||||
if _inRawMode:
|
||||
return
|
||||
fd = sys.stdin.fileno()
|
||||
try:
|
||||
old = tty.tcgetattr(fd)
|
||||
new = old[:]
|
||||
except:
|
||||
log.msg('not a typewriter!')
|
||||
else:
|
||||
# iflage
|
||||
new[0] = new[0] | tty.IGNPAR
|
||||
new[0] = new[0] & ~(tty.ISTRIP | tty.INLCR | tty.IGNCR | tty.ICRNL |
|
||||
tty.IXON | tty.IXANY | tty.IXOFF)
|
||||
if hasattr(tty, 'IUCLC'):
|
||||
new[0] = new[0] & ~tty.IUCLC
|
||||
|
||||
# lflag
|
||||
new[3] = new[3] & ~(tty.ISIG | tty.ICANON | tty.ECHO | tty.ECHO |
|
||||
tty.ECHOE | tty.ECHOK | tty.ECHONL)
|
||||
if hasattr(tty, 'IEXTEN'):
|
||||
new[3] = new[3] & ~tty.IEXTEN
|
||||
|
||||
#oflag
|
||||
new[1] = new[1] & ~tty.OPOST
|
||||
|
||||
new[6][tty.VMIN] = 1
|
||||
new[6][tty.VTIME] = 0
|
||||
|
||||
_savedRawMode = old
|
||||
tty.tcsetattr(fd, tty.TCSANOW, new)
|
||||
#tty.setraw(fd)
|
||||
_inRawMode = 1
|
||||
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
run()
|
||||
@@ -0,0 +1,10 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
#
|
||||
|
||||
"""
|
||||
An SSHv2 implementation for Twisted. Part of the Twisted.Conch package.
|
||||
|
||||
Maintainer: Paul Swartz
|
||||
"""
|
||||
@@ -0,0 +1,47 @@
|
||||
# -*- test-case-name: twisted.conch.test.test_address -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Address object for SSH network connections.
|
||||
|
||||
Maintainer: Paul Swartz
|
||||
|
||||
@since: 12.1
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
from zope.interface import implementer
|
||||
|
||||
from twisted.internet.interfaces import IAddress
|
||||
from twisted.python import util
|
||||
|
||||
|
||||
|
||||
@implementer(IAddress)
|
||||
class SSHTransportAddress(util.FancyEqMixin, object):
|
||||
"""
|
||||
Object representing an SSH Transport endpoint.
|
||||
|
||||
This is used to ensure that any code inspecting this address and
|
||||
attempting to construct a similar connection based upon it is not
|
||||
mislead into creating a transport which is not similar to the one it is
|
||||
indicating.
|
||||
|
||||
@ivar address: An instance of an object which implements I{IAddress} to
|
||||
which this transport address is connected.
|
||||
"""
|
||||
|
||||
compareAttributes = ('address',)
|
||||
|
||||
def __init__(self, address):
|
||||
self.address = address
|
||||
|
||||
|
||||
def __repr__(self):
|
||||
return 'SSHTransportAddress(%r)' % (self.address,)
|
||||
|
||||
|
||||
def __hash__(self):
|
||||
return hash(('SSH', self.address))
|
||||
@@ -0,0 +1,93 @@
|
||||
# -*- test-case-name: twisted.conch.test.test_ssh -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Common functions for the SSH classes.
|
||||
|
||||
Maintainer: Paul Swartz
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
import struct
|
||||
|
||||
from cryptography.utils import int_from_bytes, int_to_bytes
|
||||
|
||||
from twisted.python.compat import unicode
|
||||
from twisted.python.deprecate import deprecated
|
||||
from twisted.python.versions import Version
|
||||
|
||||
__all__ = ["NS", "getNS", "MP", "getMP", "ffs"]
|
||||
|
||||
|
||||
|
||||
def NS(t):
|
||||
"""
|
||||
net string
|
||||
"""
|
||||
if isinstance(t, unicode):
|
||||
t = t.encode("utf-8")
|
||||
return struct.pack('!L', len(t)) + t
|
||||
|
||||
|
||||
|
||||
def getNS(s, count=1):
|
||||
"""
|
||||
get net string
|
||||
"""
|
||||
ns = []
|
||||
c = 0
|
||||
for i in range(count):
|
||||
l, = struct.unpack('!L', s[c:c + 4])
|
||||
ns.append(s[c + 4:4 + l + c])
|
||||
c += 4 + l
|
||||
return tuple(ns) + (s[c:],)
|
||||
|
||||
|
||||
|
||||
def MP(number):
|
||||
if number == 0:
|
||||
return b'\000' * 4
|
||||
assert number > 0
|
||||
bn = int_to_bytes(number)
|
||||
if ord(bn[0:1]) & 128:
|
||||
bn = b'\000' + bn
|
||||
return struct.pack('>L', len(bn)) + bn
|
||||
|
||||
|
||||
|
||||
def getMP(data, count=1):
|
||||
"""
|
||||
Get multiple precision integer out of the string. A multiple precision
|
||||
integer is stored as a 4-byte length followed by length bytes of the
|
||||
integer. If count is specified, get count integers out of the string.
|
||||
The return value is a tuple of count integers followed by the rest of
|
||||
the data.
|
||||
"""
|
||||
mp = []
|
||||
c = 0
|
||||
for i in range(count):
|
||||
length, = struct.unpack('>L', data[c:c + 4])
|
||||
mp.append(int_from_bytes(data[c + 4:c + 4 + length], 'big'))
|
||||
c += 4 + length
|
||||
return tuple(mp) + (data[c:],)
|
||||
|
||||
|
||||
|
||||
def ffs(c, s):
|
||||
"""
|
||||
first from second
|
||||
goes through the first list, looking for items in the second, returns the first one
|
||||
"""
|
||||
for i in c:
|
||||
if i in s:
|
||||
return i
|
||||
|
||||
|
||||
|
||||
@deprecated(Version("Twisted", 16, 5, 0))
|
||||
def install():
|
||||
# This used to install gmpy, but is technically public API, so just do
|
||||
# nothing.
|
||||
pass
|
||||
@@ -0,0 +1,48 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
The parent class for all the SSH services. Currently implemented services
|
||||
are ssh-userauth and ssh-connection.
|
||||
|
||||
Maintainer: Paul Swartz
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
from twisted.python import log
|
||||
|
||||
class SSHService(log.Logger):
|
||||
name = None # this is the ssh name for the service
|
||||
protocolMessages = {} # these map #'s -> protocol names
|
||||
transport = None # gets set later
|
||||
|
||||
def serviceStarted(self):
|
||||
"""
|
||||
called when the service is active on the transport.
|
||||
"""
|
||||
|
||||
def serviceStopped(self):
|
||||
"""
|
||||
called when the service is stopped, either by the connection ending
|
||||
or by another service being started
|
||||
"""
|
||||
|
||||
def logPrefix(self):
|
||||
return "SSHService %r on %s" % (self.name,
|
||||
self.transport.transport.logPrefix())
|
||||
|
||||
def packetReceived(self, messageNum, packet):
|
||||
"""
|
||||
called when we receive a packet on the transport
|
||||
"""
|
||||
#print self.protocolMessages
|
||||
if messageNum in self.protocolMessages:
|
||||
messageType = self.protocolMessages[messageNum]
|
||||
f = getattr(self,'ssh_%s' % messageType[4:],
|
||||
None)
|
||||
if f is not None:
|
||||
return f(packet)
|
||||
log.msg("couldn't handle %r" % messageNum)
|
||||
log.msg(repr(packet))
|
||||
self.transport.sendUnimplemented()
|
||||
@@ -0,0 +1,45 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
from twisted.python.compat import intToBytes
|
||||
|
||||
|
||||
def parse(s):
|
||||
s = s.strip()
|
||||
expr = []
|
||||
while s:
|
||||
if s[0:1] == b'(':
|
||||
newSexp = []
|
||||
if expr:
|
||||
expr[-1].append(newSexp)
|
||||
expr.append(newSexp)
|
||||
s = s[1:]
|
||||
continue
|
||||
if s[0:1] == b')':
|
||||
aList = expr.pop()
|
||||
s=s[1:]
|
||||
if not expr:
|
||||
assert not s
|
||||
return aList
|
||||
continue
|
||||
i = 0
|
||||
while s[i:i+1].isdigit(): i+=1
|
||||
assert i
|
||||
length = int(s[:i])
|
||||
data = s[i+1:i+1+length]
|
||||
expr[-1].append(data)
|
||||
s=s[i+1+length:]
|
||||
assert 0, "this should not happen"
|
||||
|
||||
def pack(sexp):
|
||||
s = b""
|
||||
for o in sexp:
|
||||
if type(o) in (type(()), type([])):
|
||||
s+=b'('
|
||||
s+=pack(o)
|
||||
s+=b')'
|
||||
else:
|
||||
s+=intToBytes(len(o)) + b":" + o
|
||||
return s
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,86 @@
|
||||
# -*- test-case-name: twisted.conch.test.test_tap -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Support module for making SSH servers with twistd.
|
||||
"""
|
||||
|
||||
from twisted.conch import unix
|
||||
from twisted.conch import checkers as conch_checkers
|
||||
from twisted.conch.openssh_compat import factory
|
||||
from twisted.cred import portal, strcred
|
||||
from twisted.python import usage
|
||||
from twisted.application import strports
|
||||
|
||||
|
||||
class Options(usage.Options, strcred.AuthOptionMixin):
|
||||
synopsis = "[-i <interface>] [-p <port>] [-d <dir>] "
|
||||
longdesc = ("Makes a Conch SSH server. If no authentication methods are "
|
||||
"specified, the default authentication methods are UNIX passwords "
|
||||
"and SSH public keys. If --auth options are "
|
||||
"passed, only the measures specified will be used.")
|
||||
optParameters = [
|
||||
["interface", "i", "", "local interface to which we listen"],
|
||||
["port", "p", "tcp:22", "Port on which to listen"],
|
||||
["data", "d", "/etc", "directory to look for host keys in"],
|
||||
["moduli", "", None, "directory to look for moduli in "
|
||||
"(if different from --data)"]
|
||||
]
|
||||
compData = usage.Completions(
|
||||
optActions={"data": usage.CompleteDirs(descr="data directory"),
|
||||
"moduli": usage.CompleteDirs(descr="moduli directory"),
|
||||
"interface": usage.CompleteNetInterfaces()}
|
||||
)
|
||||
|
||||
|
||||
def __init__(self, *a, **kw):
|
||||
usage.Options.__init__(self, *a, **kw)
|
||||
|
||||
# Call the default addCheckers (for backwards compatibility) that will
|
||||
# be used if no --auth option is provided - note that conch's
|
||||
# UNIXPasswordDatabase is used, instead of twisted.plugins.cred_unix's
|
||||
# checker
|
||||
super(Options, self).addChecker(conch_checkers.UNIXPasswordDatabase())
|
||||
super(Options, self).addChecker(conch_checkers.SSHPublicKeyChecker(
|
||||
conch_checkers.UNIXAuthorizedKeysFiles()))
|
||||
self._usingDefaultAuth = True
|
||||
|
||||
|
||||
def addChecker(self, checker):
|
||||
"""
|
||||
Add the checker specified. If any checkers are added, the default
|
||||
checkers are automatically cleared and the only checkers will be the
|
||||
specified one(s).
|
||||
"""
|
||||
if self._usingDefaultAuth:
|
||||
self['credCheckers'] = []
|
||||
self['credInterfaces'] = {}
|
||||
self._usingDefaultAuth = False
|
||||
super(Options, self).addChecker(checker)
|
||||
|
||||
|
||||
|
||||
def makeService(config):
|
||||
"""
|
||||
Construct a service for operating a SSH server.
|
||||
|
||||
@param config: An L{Options} instance specifying server options, including
|
||||
where server keys are stored and what authentication methods to use.
|
||||
|
||||
@return: A L{twisted.application.service.IService} provider which contains
|
||||
the requested SSH server.
|
||||
"""
|
||||
|
||||
t = factory.OpenSSHFactory()
|
||||
|
||||
r = unix.UnixSSHRealm()
|
||||
t.portal = portal.Portal(r, config.get('credCheckers', []))
|
||||
t.dataRoot = config['data']
|
||||
t.moduliRoot = config['moduli'] or config['data']
|
||||
|
||||
port = config['port']
|
||||
if config['interface']:
|
||||
# Add warning here
|
||||
port += ':interface=' + config['interface']
|
||||
return strports.service(port, t)
|
||||
@@ -0,0 +1,571 @@
|
||||
# -*- test-case-name: twisted.conch.test.test_keys -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
# pylint: disable=I0011,C0103,W9401,W9402
|
||||
|
||||
"""
|
||||
Data used by test_keys as well as others.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
from twisted.python.compat import long, _b64decodebytes as decodebytes
|
||||
|
||||
RSAData = {
|
||||
'n': long('269413617238113438198661010376758399219880277968382122687862697'
|
||||
'296942471209955603071120391975773283844560230371884389952067978'
|
||||
'789684135947515341209478065209455427327369102356204259106807047'
|
||||
'964139525310539133073743116175821417513079706301100600025815509'
|
||||
'786721808719302671068052414466483676821987505720384645561708425'
|
||||
'794379383191274856941628512616355437197560712892001107828247792'
|
||||
'561858327085521991407807015047750218508971611590850575870321007'
|
||||
'991909043252470730134547038841839367764074379439843108550888709'
|
||||
'430958143271417044750314742880542002948053835745429446485015316'
|
||||
'60749404403945254975473896534482849256068133525751'),
|
||||
'e': long(65537),
|
||||
'd': long('420335724286999695680502438485489819800002417295071059780489811'
|
||||
'840828351636754206234982682752076205397047218449504537476523960'
|
||||
'987613148307573487322720481066677105211155388802079519869249746'
|
||||
'774085882219244493290663802569201213676433159425782937159766786'
|
||||
'329742053214957933941260042101377175565683849732354700525628975'
|
||||
'239000548651346620826136200952740446562751690924335365940810658'
|
||||
'931238410612521441739702170503547025018016868116037053013935451'
|
||||
'477930426013703886193016416453215950072147440344656137718959053'
|
||||
'897268663969428680144841987624962928576808352739627262941675617'
|
||||
'7724661940425316604626522633351193810751757014073'),
|
||||
'p': long('152689878451107675391723141129365667732639179427453246378763774'
|
||||
'448531436802867910180261906924087589684175595016060014593521649'
|
||||
'964959248408388984465569934780790357826811592229318702991401054'
|
||||
'226302790395714901636384511513449977061729214247279176398290513'
|
||||
'085108930550446985490864812445551198848562639933888780317'),
|
||||
'q': long('176444974592327996338888725079951900172097062203378367409936859'
|
||||
'072670162290963119826394224277287608693818012745872307600855894'
|
||||
'647300295516866118620024751601329775653542084052616260193174546'
|
||||
'400544176890518564317596334518015173606460860373958663673307503'
|
||||
'231977779632583864454001476729233959405710696795574874403'),
|
||||
'u': long('936018002388095842969518498561007090965136403384715613439364803'
|
||||
'229386793506402222847415019772053080458257034241832795210460612'
|
||||
'924445085372678524176842007912276654532773301546269997020970818'
|
||||
'155956828553418266110329867222673040098885651348225673298948529'
|
||||
'93885224775891490070400861134282266967852120152546563278')
|
||||
}
|
||||
|
||||
DSAData = {
|
||||
'g': long("10253261326864117157640690761723586967382334319435778695"
|
||||
"29171533815411392477819921538350732400350395446211982054"
|
||||
"96512489289702949127531056893725702005035043292195216541"
|
||||
"11525058911428414042792836395195432445511200566318251789"
|
||||
"10575695836669396181746841141924498545494149998282951407"
|
||||
"18645344764026044855941864175"),
|
||||
'p': long("10292031726231756443208850082191198787792966516790381991"
|
||||
"77502076899763751166291092085666022362525614129374702633"
|
||||
"26262930887668422949051881895212412718444016917144560705"
|
||||
"45675251775747156453237145919794089496168502517202869160"
|
||||
"78674893099371444940800865897607102159386345313384716752"
|
||||
"18590012064772045092956919481"),
|
||||
'q': long(1393384845225358996250882900535419012502712821577),
|
||||
'x': long(1220877188542930584999385210465204342686893855021),
|
||||
'y': long("14604423062661947579790240720337570315008549983452208015"
|
||||
"39426429789435409684914513123700756086453120500041882809"
|
||||
"10283610277194188071619191739512379408443695946763554493"
|
||||
"86398594314468629823767964702559709430618263927529765769"
|
||||
"10270265745700231533660131769648708944711006508965764877"
|
||||
"684264272082256183140297951")
|
||||
}
|
||||
|
||||
ECDatanistp256 = {
|
||||
'x': long('762825130203920963171185031449647317742997734817505505433829043'
|
||||
'45687059013883'),
|
||||
'y': long('815431978646028526322656647694416475343443758943143196810611371'
|
||||
'59310646683104'),
|
||||
'privateValue': long('3463874347721034170096400845565569825355565567882605'
|
||||
'9678074967909361042656500'),
|
||||
'curve': b'ecdsa-sha2-nistp256'
|
||||
}
|
||||
|
||||
ECDatanistp384 = {
|
||||
'privateValue': long('280814107134858470598753916394807521398239633534281633982576099083'
|
||||
'35787109896602102090002196616273211495718603965098'),
|
||||
'x': long('10036914308591746758780165503819213553101287571902957054148542'
|
||||
'504671046744460374996612408381962208627004841444205030'),
|
||||
'y': long('17337335659928075994560513699823544906448896792102247714689323'
|
||||
'575406618073069185107088229463828921069465902299522926'),
|
||||
'curve': b'ecdsa-sha2-nistp384'
|
||||
}
|
||||
|
||||
ECDatanistp521 = {
|
||||
'x': long('12944742826257420846659527752683763193401384271391513286022917'
|
||||
'29910013082920512632908350502247952686156279140016049549948975'
|
||||
'670668730618745449113644014505462'),
|
||||
'y': long('10784108810271976186737587749436295782985563640368689081052886'
|
||||
'16296815984553198866894145509329328086635278430266482551941240'
|
||||
'591605833440825557820439734509311'),
|
||||
'privateValue': long('662751235215460886290293902658128847495347691199214706697089140769'
|
||||
'672273950767961331442265530524063943548846724348048614239791498442'
|
||||
'5997823106818915698960565'),
|
||||
'curve': b'ecdsa-sha2-nistp521'
|
||||
}
|
||||
|
||||
privateECDSA_openssh521 = b"""-----BEGIN EC PRIVATE KEY-----
|
||||
MIHcAgEBBEIAjn0lSVF6QweS4bjOGP9RHwqxUiTastSE0MVuLtFvkxygZqQ712oZ
|
||||
ewMvqKkxthMQgxzSpGtRBcmkL7RqZ94+18qgBwYFK4EEACOhgYkDgYYABAFpX/6B
|
||||
mxxglwD+VpEvw0hcyxVzLxNnMGzxZGF7xmNj8nlF7M+TQctdlR2Xv/J+AgIeVGmB
|
||||
j2p84bkV9jBzrUNJEACsJjttZw8NbUrhxjkLT/3rMNtuwjE4vLja0P7DMTE0EV8X
|
||||
f09ETdku/z/1tOSSrSvRwmUcM9nQUJtHHAZlr5Q0fw==
|
||||
-----END EC PRIVATE KEY-----"""
|
||||
|
||||
# New format introduced in OpenSSH 6.5
|
||||
privateECDSA_openssh521_new = b"""-----BEGIN OPENSSH PRIVATE KEY-----
|
||||
b3BlbnNzaC1rZXktdjEAAAAABG5vbmUAAAAEbm9uZQAAAAAAAAABAAAArAAAABNlY2RzYS
|
||||
1zaGEyLW5pc3RwNTIxAAAACG5pc3RwNTIxAAAAhQQBaV/+gZscYJcA/laRL8NIXMsVcy8T
|
||||
ZzBs8WRhe8ZjY/J5RezPk0HLXZUdl7/yfgICHlRpgY9qfOG5FfYwc61DSRAArCY7bWcPDW
|
||||
1K4cY5C0/96zDbbsIxOLy42tD+wzExNBFfF39PRE3ZLv8/9bTkkq0r0cJlHDPZ0FCbRxwG
|
||||
Za+UNH8AAAEAeRISlnkSEpYAAAATZWNkc2Etc2hhMi1uaXN0cDUyMQAAAAhuaXN0cDUyMQ
|
||||
AAAIUEAWlf/oGbHGCXAP5WkS/DSFzLFXMvE2cwbPFkYXvGY2PyeUXsz5NBy12VHZe/8n4C
|
||||
Ah5UaYGPanzhuRX2MHOtQ0kQAKwmO21nDw1tSuHGOQtP/esw227CMTi8uNrQ/sMxMTQRXx
|
||||
d/T0RN2S7/P/W05JKtK9HCZRwz2dBQm0ccBmWvlDR/AAAAQgCOfSVJUXpDB5LhuM4Y/1Ef
|
||||
CrFSJNqy1ITQxW4u0W+THKBmpDvXahl7Ay+oqTG2ExCDHNKka1EFyaQvtGpn3j7XygAAAA
|
||||
ABAg==
|
||||
-----END OPENSSH PRIVATE KEY-----"""
|
||||
|
||||
publicECDSA_openssh521 = (
|
||||
b"ecdsa-sha2-nistp521 AAAAE2VjZHNhLXNoYTItbmlzdHA1MjEAAAAIbmlzdHA1MjEAAACF"
|
||||
b"BAFpX/6BmxxglwD+VpEvw0hcyxVzLxNnMGzxZGF7xmNj8nlF7M+TQctdlR2Xv/J+AgIeVGmB"
|
||||
b"j2p84bkV9jBzrUNJEACsJjttZw8NbUrhxjkLT/3rMNtuwjE4vLja0P7DMTE0EV8Xf09ETdku"
|
||||
b"/z/1tOSSrSvRwmUcM9nQUJtHHAZlr5Q0fw== comment"
|
||||
)
|
||||
|
||||
privateECDSA_openssh384 = b"""-----BEGIN EC PRIVATE KEY-----
|
||||
MIGkAgEBBDAtAi7I8j73WCX20qUM5hhHwHuFzYWYYILs2Sh8UZ+awNkARZ/Fu2LU
|
||||
LLl5RtOQpbWgBwYFK4EEACKhZANiAATU17sA9P5FRwSknKcFsjjsk0+E3CeXPYX0
|
||||
Tk/M0HK3PpWQWgrO8JdRHP9eFE9O/23P8BumwFt7F/AvPlCzVd35VfraFT0o4cCW
|
||||
G0RqpQ+np31aKmeJshkcYALEchnU+tQ=
|
||||
-----END EC PRIVATE KEY-----"""
|
||||
|
||||
# New format introduced in OpenSSH 6.5
|
||||
privateECDSA_openssh384_new = b"""-----BEGIN OPENSSH PRIVATE KEY-----
|
||||
b3BlbnNzaC1rZXktdjEAAAAABG5vbmUAAAAEbm9uZQAAAAAAAAABAAAAiAAAABNlY2RzYS
|
||||
1zaGEyLW5pc3RwMzg0AAAACG5pc3RwMzg0AAAAYQTU17sA9P5FRwSknKcFsjjsk0+E3CeX
|
||||
PYX0Tk/M0HK3PpWQWgrO8JdRHP9eFE9O/23P8BumwFt7F/AvPlCzVd35VfraFT0o4cCWG0
|
||||
RqpQ+np31aKmeJshkcYALEchnU+tQAAADIiktpWIpLaVgAAAATZWNkc2Etc2hhMi1uaXN0
|
||||
cDM4NAAAAAhuaXN0cDM4NAAAAGEE1Ne7APT+RUcEpJynBbI47JNPhNwnlz2F9E5PzNBytz
|
||||
6VkFoKzvCXURz/XhRPTv9tz/AbpsBbexfwLz5Qs1Xd+VX62hU9KOHAlhtEaqUPp6d9Wipn
|
||||
ibIZHGACxHIZ1PrUAAAAMC0CLsjyPvdYJfbSpQzmGEfAe4XNhZhgguzZKHxRn5rA2QBFn8
|
||||
W7YtQsuXlG05CltQAAAAA=
|
||||
-----END OPENSSH PRIVATE KEY-----"""
|
||||
|
||||
publicECDSA_openssh384 = (
|
||||
b"ecdsa-sha2-nistp384 AAAAE2VjZHNhLXNoYTItbmlzdHAzODQAAAAIbmlzdHAzODQAAABh"
|
||||
b"BNTXuwD0/kVHBKScpwWyOOyTT4TcJ5c9hfROT8zQcrc+lZBaCs7wl1Ec/14UT07/bc/wG6bA"
|
||||
b"W3sX8C8+ULNV3flV+toVPSjhwJYbRGqlD6enfVoqZ4myGRxgAsRyGdT61A== comment"
|
||||
)
|
||||
|
||||
publicECDSA_openssh = (
|
||||
b"ecdsa-sha2-nistp256 AAAAE2VjZHNhLXNoYTItbmlzdHAyNTYAAAAIbmlzdHAyNTYAAABB"
|
||||
b"BKimX1DZ7+Qj0SpfePMbo1pb6yGkAb5l7duC1l855yD7tEfQfqk7bc7v46We1hLMyz6ObUBY"
|
||||
b"gkN/34n42F4vpeA= comment"
|
||||
)
|
||||
|
||||
privateECDSA_openssh = b"""-----BEGIN EC PRIVATE KEY-----
|
||||
MHcCAQEEIEyU1YOT2JxxofwbJXIjGftdNcJK55aQdNrhIt2xYQz0oAoGCCqGSM49
|
||||
AwEHoUQDQgAEqKZfUNnv5CPRKl948xujWlvrIaQBvmXt24LWXznnIPu0R9B+qTtt
|
||||
zu/jpZ7WEszLPo5tQFiCQ3/fifjYXi+l4A==
|
||||
-----END EC PRIVATE KEY-----"""
|
||||
|
||||
# New format introduced in OpenSSH 6.5
|
||||
privateECDSA_openssh_new = b"""-----BEGIN OPENSSH PRIVATE KEY-----
|
||||
b3BlbnNzaC1rZXktdjEAAAAABG5vbmUAAAAEbm9uZQAAAAAAAAABAAAAaAAAABNlY2RzYS
|
||||
1zaGEyLW5pc3RwMjU2AAAACG5pc3RwMjU2AAAAQQSopl9Q2e/kI9EqX3jzG6NaW+shpAG+
|
||||
Ze3bgtZfOecg+7RH0H6pO23O7+OlntYSzMs+jm1AWIJDf9+J+NheL6XgAAAAmCKU4hcilO
|
||||
IXAAAAE2VjZHNhLXNoYTItbmlzdHAyNTYAAAAIbmlzdHAyNTYAAABBBKimX1DZ7+Qj0Spf
|
||||
ePMbo1pb6yGkAb5l7duC1l855yD7tEfQfqk7bc7v46We1hLMyz6ObUBYgkN/34n42F4vpe
|
||||
AAAAAgTJTVg5PYnHGh/BslciMZ+101wkrnlpB02uEi3bFhDPQAAAAA
|
||||
-----END OPENSSH PRIVATE KEY-----"""
|
||||
|
||||
publicRSA_openssh = (
|
||||
b"ssh-rsa AAAAB3NzaC1yc2EAAAADAQABAAABAQDVaqx4I9bWG+wloVDEd2NQhEUBVUIUKirg"
|
||||
b"0GDu1OmjrUr6OQZehFV1XwA2v2+qKj+DJjfBaS5b/fDz0n3WmM06QHjVyqgYwBGTJAkMgUyP"
|
||||
b"95ztExZqpATpSXfD5FVks3loniwI66zoBC0hdwWnju9TMA2l5bs9auIJNm/9NNN9b0b/h9qp"
|
||||
b"KSeq/631heY+Grh6HUqx6sBa9zDfH8Kk5O8/kUmWQNUZdy03w17snaY6RKXCpCnd1bqcPUWz"
|
||||
b"xiwYZNW6Pd+rf81CrKfxGAugWBViC6QqbkPD5ASfNaNHjkbtM6Vlvbw7KW4CC1ffdOgTtDc1"
|
||||
b"foNfICZgptyti8ZseZj3 comment"
|
||||
)
|
||||
|
||||
privateRSA_openssh = b'''-----BEGIN RSA PRIVATE KEY-----
|
||||
MIIEogIBAAKCAQEA1WqseCPW1hvsJaFQxHdjUIRFAVVCFCoq4NBg7tTpo61K+jkG
|
||||
XoRVdV8ANr9vqio/gyY3wWkuW/3w89J91pjNOkB41cqoGMARkyQJDIFMj/ec7RMW
|
||||
aqQE6Ul3w+RVZLN5aJ4sCOus6AQtIXcFp47vUzANpeW7PWriCTZv/TTTfW9G/4fa
|
||||
qSknqv+t9YXmPhq4eh1KserAWvcw3x/CpOTvP5FJlkDVGXctN8Ne7J2mOkSlwqQp
|
||||
3dW6nD1Fs8YsGGTVuj3fq3/NQqyn8RgLoFgVYgukKm5Dw+QEnzWjR45G7TOlZb28
|
||||
OyluAgtX33ToE7Q3NX6DXyAmYKbcrYvGbHmY9wIDAQABAoIBACFMCGaiKNW0+44P
|
||||
chuFCQC58k438BxXS+NRf54jp+Q6mFUb6ot6mB682Lqx+YkSGGCs6MwLTglaQGq6
|
||||
L5n4syRghLnOaZWa+eL8H1FNJxXbKyet77RprL59EOuGR3BztACHlRU7N/nnFOeA
|
||||
u2geG+bdu3NjuWfmsid/z88wm8KY/dkYNi82LvE9gXqf4QMtR9s0UWI53U/prKiL
|
||||
2dbzhMQXuXGdBghCeE27xSr0w1jNVSvtvjNfBOp75gQkY/It1z0bbNWcY0MvkoiN
|
||||
Pm7aGDfYDyVniR25RjReyc7Ei+2SWjMHD9+GCPmS6dvrOAg2yc3NCgFIWzk+esrG
|
||||
gKnc1DkCgYEA2XAG2OK81HiRUJTUwRuJOGxGZFpRoJoHPUiPA1HMaxKOfRqxZedx
|
||||
dTngMgV1jRhMr5OxSbFmX3hietEMyuZNQ7Oc9Gt95gyY3M8hYo7VLhLeBK7XJG6D
|
||||
MaIVokQ9IqliJiK5su1UCp0Ig6cHDf8ZGI7Yqx3aSJwxaBGhZm3j2B0CgYEA+0QX
|
||||
i6Q2vh43Haf2YWwExKrdeD4HjB4zAq4DFIeDeuWefQhnqPKqvxJwz3Kpp8cLHYjV
|
||||
IP2cY8pHMFVOi8TP9H8WpJISdKEJwsRunIwz76Xl9+ArrU9cEaoahDdb/Xrqw818
|
||||
sMjkH1Rjtcev3/QJp/zHJfxc6ZHXksWYHlbTsSMCgYBRr+mSn5QLSoRlPpSzO5IQ
|
||||
tXS4jMnvyQ4BMvovaBKhAyauz1FoFEwmmyikAjMIX+GncJgBNHleUo7Ezza8H0tV
|
||||
rOvBU4TH4WGoStSi/0ANgB8SqVDAKhh1lAwGmxZQqEvsQc177/dLyXUCaMSYuIaI
|
||||
GFpD5wIzlyJkk4MMRSp87QKBgGlmN8ZA3SHFBPOwuD5HlHx2/C3rPzk8lcNDAVHE
|
||||
Qpfz6Bakxu7s1EkQUDgE7jvN19DMzDJpkAegG1qf/jHNHjp+cR4ZlBpOTwzfX1LV
|
||||
0Rdu7NectlWd244hX7wkiLb8r6vw76QssNyfhrADEriL4t0PwO4jPUpQ/i+4KUZY
|
||||
v7YnAoGAZhb5IDTQVCW8YTGsgvvvnDUefkpVAmiVDQqTvh6/4UD6kKdUcDHpePzg
|
||||
Zrcid5rr3dXSMEbK4tdeQZvPtUg1Uaol3N7bNClIIdvWdPx+5S9T95wJcLnkoHam
|
||||
rXp0IjScTxfLP+Cq5V6lJ94/pX8Ppoj1FdZfNxeS4NYFSRA7kvY=
|
||||
-----END RSA PRIVATE KEY-----'''
|
||||
|
||||
# Some versions of OpenSSH generate these (slightly different keys): the PKCS#1
|
||||
# structure is wrapped in an extra ASN.1 SEQUENCE and there's an empty SEQUENCE
|
||||
# following it. It is not any standard key format and was probably a bug in
|
||||
# OpenSSH at some point.
|
||||
privateRSA_openssh_alternate = b"""-----BEGIN RSA PRIVATE KEY-----
|
||||
MIIEqTCCBKMCAQACggEBANVqrHgj1tYb7CWhUMR3Y1CERQFVQhQqKuDQYO7U6aOtSvo5Bl6EVXVf
|
||||
ADa/b6oqP4MmN8FpLlv98PPSfdaYzTpAeNXKqBjAEZMkCQyBTI/3nO0TFmqkBOlJd8PkVWSzeWie
|
||||
LAjrrOgELSF3BaeO71MwDaXluz1q4gk2b/00031vRv+H2qkpJ6r/rfWF5j4auHodSrHqwFr3MN8f
|
||||
wqTk7z+RSZZA1Rl3LTfDXuydpjpEpcKkKd3Vupw9RbPGLBhk1bo936t/zUKsp/EYC6BYFWILpCpu
|
||||
Q8PkBJ81o0eORu0zpWW9vDspbgILV9906BO0NzV+g18gJmCm3K2Lxmx5mPcCAwEAAQKCAQAhTAhm
|
||||
oijVtPuOD3IbhQkAufJON/AcV0vjUX+eI6fkOphVG+qLepgevNi6sfmJEhhgrOjMC04JWkBqui+Z
|
||||
+LMkYIS5zmmVmvni/B9RTScV2ysnre+0aay+fRDrhkdwc7QAh5UVOzf55xTngLtoHhvm3btzY7ln
|
||||
5rInf8/PMJvCmP3ZGDYvNi7xPYF6n+EDLUfbNFFiOd1P6ayoi9nW84TEF7lxnQYIQnhNu8Uq9MNY
|
||||
zVUr7b4zXwTqe+YEJGPyLdc9G2zVnGNDL5KIjT5u2hg32A8lZ4kduUY0XsnOxIvtklozBw/fhgj5
|
||||
kunb6zgINsnNzQoBSFs5PnrKxoCp3NQ5AoGBANlwBtjivNR4kVCU1MEbiThsRmRaUaCaBz1IjwNR
|
||||
zGsSjn0asWXncXU54DIFdY0YTK+TsUmxZl94YnrRDMrmTUOznPRrfeYMmNzPIWKO1S4S3gSu1yRu
|
||||
gzGiFaJEPSKpYiYiubLtVAqdCIOnBw3/GRiO2Ksd2kicMWgRoWZt49gdAoGBAPtEF4ukNr4eNx2n
|
||||
9mFsBMSq3Xg+B4weMwKuAxSHg3rlnn0IZ6jyqr8ScM9yqafHCx2I1SD9nGPKRzBVTovEz/R/FqSS
|
||||
EnShCcLEbpyMM++l5ffgK61PXBGqGoQ3W/166sPNfLDI5B9UY7XHr9/0Caf8xyX8XOmR15LFmB5W
|
||||
07EjAoGAUa/pkp+UC0qEZT6UszuSELV0uIzJ78kOATL6L2gSoQMmrs9RaBRMJpsopAIzCF/hp3CY
|
||||
ATR5XlKOxM82vB9LVazrwVOEx+FhqErUov9ADYAfEqlQwCoYdZQMBpsWUKhL7EHNe+/3S8l1AmjE
|
||||
mLiGiBhaQ+cCM5ciZJODDEUqfO0CgYBpZjfGQN0hxQTzsLg+R5R8dvwt6z85PJXDQwFRxEKX8+gW
|
||||
pMbu7NRJEFA4BO47zdfQzMwyaZAHoBtan/4xzR46fnEeGZQaTk8M319S1dEXbuzXnLZVnduOIV+8
|
||||
JIi2/K+r8O+kLLDcn4awAxK4i+LdD8DuIz1KUP4vuClGWL+2JwKBgQCFSxt6mxIQN54frV7a/saW
|
||||
/t81a7k04haXkiYJvb1wIAOnNb0tG6DSB0cr1N6oqAcHG7gEIKcnQTxsOTnpQc7nFx3RTFy8PdIm
|
||||
Jv5q1v1Icq5G+nvD0xlgRB2lE6eA9WMp1HpdBgcWXfaLPctkOuKEWk2MBi0tnRzrg0x4PXlUzjAA
|
||||
-----END RSA PRIVATE KEY-----"""
|
||||
|
||||
# New format introduced in OpenSSH 6.5
|
||||
privateRSA_openssh_new = b'''-----BEGIN OPENSSH PRIVATE KEY-----
|
||||
b3BlbnNzaC1rZXktdjEAAAAABG5vbmUAAAAEbm9uZQAAAAAAAAABAAABFwAAAAdzc2gtcn
|
||||
NhAAAAAwEAAQAAAQEA1WqseCPW1hvsJaFQxHdjUIRFAVVCFCoq4NBg7tTpo61K+jkGXoRV
|
||||
dV8ANr9vqio/gyY3wWkuW/3w89J91pjNOkB41cqoGMARkyQJDIFMj/ec7RMWaqQE6Ul3w+
|
||||
RVZLN5aJ4sCOus6AQtIXcFp47vUzANpeW7PWriCTZv/TTTfW9G/4faqSknqv+t9YXmPhq4
|
||||
eh1KserAWvcw3x/CpOTvP5FJlkDVGXctN8Ne7J2mOkSlwqQp3dW6nD1Fs8YsGGTVuj3fq3
|
||||
/NQqyn8RgLoFgVYgukKm5Dw+QEnzWjR45G7TOlZb28OyluAgtX33ToE7Q3NX6DXyAmYKbc
|
||||
rYvGbHmY9wAAA7gXkBoMF5AaDAAAAAdzc2gtcnNhAAABAQDVaqx4I9bWG+wloVDEd2NQhE
|
||||
UBVUIUKirg0GDu1OmjrUr6OQZehFV1XwA2v2+qKj+DJjfBaS5b/fDz0n3WmM06QHjVyqgY
|
||||
wBGTJAkMgUyP95ztExZqpATpSXfD5FVks3loniwI66zoBC0hdwWnju9TMA2l5bs9auIJNm
|
||||
/9NNN9b0b/h9qpKSeq/631heY+Grh6HUqx6sBa9zDfH8Kk5O8/kUmWQNUZdy03w17snaY6
|
||||
RKXCpCnd1bqcPUWzxiwYZNW6Pd+rf81CrKfxGAugWBViC6QqbkPD5ASfNaNHjkbtM6Vlvb
|
||||
w7KW4CC1ffdOgTtDc1foNfICZgptyti8ZseZj3AAAAAwEAAQAAAQAhTAhmoijVtPuOD3Ib
|
||||
hQkAufJON/AcV0vjUX+eI6fkOphVG+qLepgevNi6sfmJEhhgrOjMC04JWkBqui+Z+LMkYI
|
||||
S5zmmVmvni/B9RTScV2ysnre+0aay+fRDrhkdwc7QAh5UVOzf55xTngLtoHhvm3btzY7ln
|
||||
5rInf8/PMJvCmP3ZGDYvNi7xPYF6n+EDLUfbNFFiOd1P6ayoi9nW84TEF7lxnQYIQnhNu8
|
||||
Uq9MNYzVUr7b4zXwTqe+YEJGPyLdc9G2zVnGNDL5KIjT5u2hg32A8lZ4kduUY0XsnOxIvt
|
||||
klozBw/fhgj5kunb6zgINsnNzQoBSFs5PnrKxoCp3NQ5AAAAgQCFSxt6mxIQN54frV7a/s
|
||||
aW/t81a7k04haXkiYJvb1wIAOnNb0tG6DSB0cr1N6oqAcHG7gEIKcnQTxsOTnpQc7nFx3R
|
||||
TFy8PdImJv5q1v1Icq5G+nvD0xlgRB2lE6eA9WMp1HpdBgcWXfaLPctkOuKEWk2MBi0tnR
|
||||
zrg0x4PXlUzgAAAIEA2XAG2OK81HiRUJTUwRuJOGxGZFpRoJoHPUiPA1HMaxKOfRqxZedx
|
||||
dTngMgV1jRhMr5OxSbFmX3hietEMyuZNQ7Oc9Gt95gyY3M8hYo7VLhLeBK7XJG6DMaIVok
|
||||
Q9IqliJiK5su1UCp0Ig6cHDf8ZGI7Yqx3aSJwxaBGhZm3j2B0AAACBAPtEF4ukNr4eNx2n
|
||||
9mFsBMSq3Xg+B4weMwKuAxSHg3rlnn0IZ6jyqr8ScM9yqafHCx2I1SD9nGPKRzBVTovEz/
|
||||
R/FqSSEnShCcLEbpyMM++l5ffgK61PXBGqGoQ3W/166sPNfLDI5B9UY7XHr9/0Caf8xyX8
|
||||
XOmR15LFmB5W07EjAAAAAAEC
|
||||
-----END OPENSSH PRIVATE KEY-----'''
|
||||
|
||||
# Encrypted with the passphrase 'encrypted'
|
||||
privateRSA_openssh_encrypted = b"""-----BEGIN RSA PRIVATE KEY-----
|
||||
Proc-Type: 4,ENCRYPTED
|
||||
DEK-Info: DES-EDE3-CBC,FFFFFFFFFFFFFFFF
|
||||
|
||||
p2A1YsHLXkpMVcsEqhh/nCYb5AqL0uMzfEIqc8hpZ/Ub8PtLsypilMkqzYTnZIGS
|
||||
ouyPjU/WgtR4VaDnutPWdgYaKdixSEmGhKghCtXFySZqCTJ4O8NCczsktYjUK3D4
|
||||
Jtl90zL6O81WBY6xP76PBQo9lrI/heAetATeyqutc18bwQIGU+gKk32qvfo15DfS
|
||||
VYiY0Ds4D7F7fd9pz+f5+UbFUCgU+tfDvBrqodYrUgmH7jKoW/CRDCHHyeEIZDbF
|
||||
mcMwdcKOyw1sRLaPdihRSVx3kOMvIotHKVTkIDMp+0RTNeXzQnp5U2qzsxzTcG/M
|
||||
UyJN38XXkuvq5VMj2zmmjHzx34w3NK3ZxpZcoaFUqUBlNp2C8hkCLrAa/DWobKqN
|
||||
5xA1ElrQvli9XXkT/RIuy4Gc10bbGEoJjuxNRibtSxxWd5Bd1E40ocOd4l1ebI8+
|
||||
w69XvMTnsmHvkBEADGF2zfRszKnMelg+W5NER1UDuNT03i+1cuhp+2AZg8z7niTO
|
||||
M17XP3ScGVxrQAEYgtxPrPeIpFJvOx2j5Yt78U9Y2WlaAG6DrubbYv2RsMIibhOG
|
||||
yk139vMdD8FwCey6yMkkhFAJwnBtC22MAWgjmC5c6AF3SRQSjjQXepPsJcLgpOjy
|
||||
YwjhnL8w56x9kVDUNPw9A9Cqgxo2sty34ATnKrh4h59PsP83LOL6OC5WjbASgZRd
|
||||
OIBD8RloQPISo+RUF7X0i4kdaHVNPlR0KyapR+3M5BwhQuvEO99IArDV2LNKGzfc
|
||||
W4ssugm8iyAJlmwmb2yRXIDHXabInWY7XCdGk8J2qPFbDTvnPbiagJBimjVjgpWw
|
||||
tV3sVlJYqmOqmCDP78J6he04l0vaHtiOWTDEmNCrK7oFMXIIp3XWjOZGPSOJFdPs
|
||||
6Go3YB+EGWfOQxqkFM28gcqmYfVPF2sa1FbZLz0ffO11Ma/rliZxZu7WdrAXe/tc
|
||||
BgIQ8etp2PwAK4jCwwVwjIO8FzqQGpS23Y9NY3rfi97ckgYXKESFtXPsMMA+drZd
|
||||
ThbXvccfh4EPmaqQXKf4WghHiVJ+/yuY1kUIDEl/O0jRZWT7STgBim/Aha1m6qRs
|
||||
zl1H7hkDbU4solb1GM5oPzbgGTzyBc+z0XxM9iFRM+fMzPB8+yYHTr4kPbVmKBjy
|
||||
SCovjQQVsHE4YeUGTq6k/NF5cVIRKTW/RlHvzxsky1Zj31MC736jrxGw4KG7VSLZ
|
||||
fP6F5jj+mXwS7m0v5to42JBZmRJdKUD88QaGE3ncyQ4yleW5bn9Lf9SuzQg1Dhao
|
||||
3rSA1RuexsHlIAHvGxx/17X+pyygl8DJbt6TBfbLQk9wc707DJTfh5M/bnk9wwIX
|
||||
l/Hsa1WtylAMW/2MzgiVy83MbYz4+Ss6GQ5W66okWji+NxrnrYEy6q+WgVQanp7X
|
||||
D+D7oKykqE1Cdvvulvtfl5fh8wlAs8mrUnKPBBUru348u++2lfacLkxRXyT1ooqY
|
||||
uSNE5nlwFt08N2Ou/bl7yq6QNRMYrRkn+UEfHWCNYDoGMHln2/i6Z1RapQzNarik
|
||||
tJf7radBz5nBwBjP08YAEACNSQvpsUgdqiuYjLwX7efFXQva2RzqaQ==
|
||||
-----END RSA PRIVATE KEY-----"""
|
||||
|
||||
# Encrypted with the passphrase 'encrypted', and using the new format
|
||||
# introduced in OpenSSH 6.5
|
||||
privateRSA_openssh_encrypted_new = b"""-----BEGIN OPENSSH PRIVATE KEY-----
|
||||
b3BlbnNzaC1rZXktdjEAAAAACmFlczI1Ni1jdHIAAAAGYmNyeXB0AAAAGAAAABD0f9WAof
|
||||
DTbmwztb8pdrSeAAAAEAAAAAEAAAEXAAAAB3NzaC1yc2EAAAADAQABAAABAQDVaqx4I9bW
|
||||
G+wloVDEd2NQhEUBVUIUKirg0GDu1OmjrUr6OQZehFV1XwA2v2+qKj+DJjfBaS5b/fDz0n
|
||||
3WmM06QHjVyqgYwBGTJAkMgUyP95ztExZqpATpSXfD5FVks3loniwI66zoBC0hdwWnju9T
|
||||
MA2l5bs9auIJNm/9NNN9b0b/h9qpKSeq/631heY+Grh6HUqx6sBa9zDfH8Kk5O8/kUmWQN
|
||||
UZdy03w17snaY6RKXCpCnd1bqcPUWzxiwYZNW6Pd+rf81CrKfxGAugWBViC6QqbkPD5ASf
|
||||
NaNHjkbtM6Vlvbw7KW4CC1ffdOgTtDc1foNfICZgptyti8ZseZj3AAADwPQaac8s1xX3af
|
||||
hQTQexj0vEAWDQsLYzDHN9G7W+UP5WHUu7igeu2GqAC/TOnjUXDP73I+EN3n7T3JFeDRfs
|
||||
U1Z6Zqb0NKHSRVYwDIdIi8qVohFv85g6+xQ01OpaoOzz+vI34OUvCRHQGTgR6L9fQShZyC
|
||||
McopYMYfbIse6KcqkfxX3KSdG1Pao6Njx/ShFRbgvmALpR/z0EaGCzHCDxpfUyAdnxm621
|
||||
Jzaf+LverWdN7sfrfMptaS9//9iJb70sL67K+YIB64qhDnA/w9UOQfXGQFL+AEtdM0BPv8
|
||||
thP1bs7T0yucBl+ZXdrDKVLZfaS3S/w85Jlgfu+a1DG73pOBOuag435iEJ9EnspjXiiydx
|
||||
GrfSRk2C+/c4fBDZVGFscK5bfQuUUZyU1qOagekxX7WLHFKk9xajnud+nrAN070SeNwlX8
|
||||
FZ2CI4KGlQfDvVUpKanYn8Kkj3fZ+YBGyx4M+19clF65FKSM0x1Rrh5tAmNT/SNDbSc28m
|
||||
ASxrBhztzxUFTrIn3tp+uqkJniFLmFsUtiAUmj8fNyE9blykU7dqq+CqpLA872nQ9bOHHA
|
||||
JsS1oBYmQ0n6AJz8WrYMdcepqWVld6Q8QSD1zdrY/sAWUovuBA1s4oIEXZhpXSS4ZJiMfh
|
||||
PVktKBwj5bmoG/mmwYLbo0JHntK8N3TGTzTGLq5TpSBBdVvWSWo7tnfEkrFObmhi1uJSrQ
|
||||
3zfPVP6BguboxBv+oxhaUBK8UOANe6ZwM4vfiu+QN+sZqWymHIfAktz7eWzwlToe4cKpdG
|
||||
Uv+e3/7Lo2dyMl3nke5HsSUrlsMGPREuGkBih8+o85ii6D+cuCiVtus3f5c78Cir80zLIr
|
||||
Z0wWvEAjciEvml00DWaA+JIaOrWwvXySaOzFGpCqC9SQjao379bvn9P3b7kVZsy6zBfHqm
|
||||
bNEJUOuhBZaY8Okz36chh1xqh4sz7m3nsZ3GYGcvM+3mvRY72QnqsQEG0Sp1XYIn2bHa29
|
||||
tqp7CG9X8J6dqMcPeoPRDWIX9gw7EPl/M0LP6xgewGJ9bgxwle6Mnr9kNITIswjAJqrLec
|
||||
zx7dfixjAPc42ADqrw/tEdFQcSqxigcfJNKO1LbDBjh+Hk/cSBou2PoxbIcl0qfQfbGcqI
|
||||
Dbpd695IEuiW9pYR22txNoIi+7cbMsuFHxQ/OqbrX/jCsprGNNJLAjgGsVEI1JnHWDH0db
|
||||
3UbqbOHAeY3ufoYXNY1utVOIACpW3r9wBw3FjRi04d70VcKr16OXvOAHGN2G++Y+kMya84
|
||||
Hl/Kt/gA==
|
||||
-----END OPENSSH PRIVATE KEY-----"""
|
||||
|
||||
# Encrypted with the passphrase 'testxp'. NB: this key was generated by
|
||||
# OpenSSH, so it doesn't use the same key data as the other keys here.
|
||||
privateRSA_openssh_encrypted_aes = b"""-----BEGIN RSA PRIVATE KEY-----
|
||||
Proc-Type: 4,ENCRYPTED
|
||||
DEK-Info: AES-128-CBC,0673309A6ACCAB4B77DEE1C1E536AC26
|
||||
|
||||
4Ed/a9OgJWHJsne7yOGWeWMzHYKsxuP9w1v0aYcp+puS75wvhHLiUnNwxz0KDi6n
|
||||
T3YkKLBsoCWS68ApR2J9yeQ6R+EyS+UQDrO9nwqo3DB5BT3Ggt8S1wE7vjNLQD0H
|
||||
g/SJnlqwsECNhh8aAx+Ag0m3ZKOZiRD5mCkcDQsZET7URSmFytDKOjhFn3u6ZFVB
|
||||
sXrfpYc6TJtOQlHd/52JB6aAbjt6afSv955Z7enIi+5yEJ5y7oYQTaE5zrFMP7N5
|
||||
9LbfJFlKXxEddy/DErRLxEjmC+t4svHesoJKc2jjjyNPiOoGGF3kJXea62vsjdNV
|
||||
gMK5Eged3TBVIk2dv8rtJUvyFeCUtjQ1UJZIebScRR47KrbsIpCmU8I4/uHWm5hW
|
||||
0mOwvdx1L/mqx/BHqVU9Dw2COhOdLbFxlFI92chkovkmNk4P48ziyVnpm7ME22sE
|
||||
vfCMsyirdqB1mrL4CSM7FXONv+CgfBfeYVkYW8RfJac9U1L/O+JNn7yee414O/rS
|
||||
hRYw4UdWnH6Gg6niklVKWNY0ZwUZC8zgm2iqy8YCYuneS37jC+OEKP+/s6HSKuqk
|
||||
2bzcl3/TcZXNSM815hnFRpz0anuyAsvwPNRyvxG2/DacJHL1f6luV4B0o6W410yf
|
||||
qXQx01DLo7nuyhJqoH3UGCyyXB+/QUs0mbG2PAEn3f5dVs31JMdbt+PrxURXXjKk
|
||||
4cexpUcIpqqlfpIRe3RD0sDVbH4OXsGhi2kiTfPZu7mgyFxKopRbn1KwU1qKinfY
|
||||
EU9O4PoTak/tPT+5jFNhaP+HrURoi/pU8EAUNSktl7xAkHYwkN/9Cm7DeBghgf3n
|
||||
8+tyCGYDsB5utPD0/Xe9yx0Qhc/kMm4xIyQDyA937dk3mUvLC9vulnAP8I+Izim0
|
||||
fZ182+D1bWwykoD0997mUHG/AUChWR01V1OLwRyPv2wUtiS8VNG76Y2aqKlgqP1P
|
||||
V+IvIEqR4ERvSBVFzXNF8Y6j/sVxo8+aZw+d0L1Ns/R55deErGg3B8i/2EqGd3r+
|
||||
0jps9BqFHHWW87n3VyEB3jWCMj8Vi2EJIfa/7pSaViFIQn8LiBLf+zxG5LTOToK5
|
||||
xkN42fReDcqi3UNfKNGnv4dsplyTR2hyx65lsj4bRKDGLKOuB1y7iB0AGb0LtcAI
|
||||
dcsVlcCeUquDXtqKvRnwfIMg+ZunyjqHBhj3qgRgbXbT6zjaSdNnih569aTg0Vup
|
||||
VykzZ7+n/KVcGLmvX0NesdoI7TKbq4TnEIOynuG5Sf+2GpARO5bjcWKSZeN/Ybgk
|
||||
gccf8Cqf6XWqiwlWd0B7BR3SymeHIaSymC45wmbgdstrbk7Ppa2Tp9AZku8M2Y7c
|
||||
8mY9b+onK075/ypiwBm4L4GRNTFLnoNQJXx0OSl4FNRWsn6ztbD+jZhu8Seu10Jw
|
||||
SEJVJ+gmTKdRLYORJKyqhDet6g7kAxs4EoJ25WsOnX5nNr00rit+NkMPA7xbJT+7
|
||||
CfI51GQLw7pUPeO2WNt6yZO/YkzZrqvTj5FEwybkUyBv7L0gkqu9wjfDdUw0fVHE
|
||||
xEm4DxjEoaIp8dW/JOzXQ2EF+WaSOgdYsw3Ac+rnnjnNptCdOEDGP6QBkt+oXj4P
|
||||
-----END RSA PRIVATE KEY-----"""
|
||||
|
||||
publicRSA_lsh = (
|
||||
b'{KDEwOnB1YmxpYy1rZXkoMTQ6cnNhLXBrY3MxLXNoYTEoMTpuMjU3OgDVaqx4I9bWG+wloVD'
|
||||
b'Ed2NQhEUBVUIUKirg0GDu1OmjrUr6OQZehFV1XwA2v2+qKj+DJjfBaS5b/fDz0n3WmM06QHj'
|
||||
b'VyqgYwBGTJAkMgUyP95ztExZqpATpSXfD5FVks3loniwI66zoBC0hdwWnju9TMA2l5bs9auI'
|
||||
b'JNm/9NNN9b0b/h9qpKSeq/631heY+Grh6HUqx6sBa9zDfH8Kk5O8/kUmWQNUZdy03w17snaY'
|
||||
b'6RKXCpCnd1bqcPUWzxiwYZNW6Pd+rf81CrKfxGAugWBViC6QqbkPD5ASfNaNHjkbtM6Vlvbw'
|
||||
b'7KW4CC1ffdOgTtDc1foNfICZgptyti8ZseZj3KSgxOmUzOgEAASkpKQ==}'
|
||||
)
|
||||
|
||||
privateRSA_lsh = (
|
||||
b"(11:private-key(9:rsa-pkcs1(1:n257:\x00\xd5j\xacx#\xd6\xd6\x1b\xec%\xa1P"
|
||||
b"\xc4wcP\x84E\x01UB\x14**\xe0\xd0`\xee\xd4\xe9\xa3\xadJ\xfa9\x06^\x84Uu_"
|
||||
b"\x006\xbfo\xaa*?\x83&7\xc1i.[\xfd\xf0\xf3\xd2}\xd6\x98\xcd:@x\xd5\xca"
|
||||
b"\xa8\x18\xc0\x11\x93$\t\x0c\x81L\x8f\xf7\x9c\xed\x13\x16j\xa4\x04\xe9Iw"
|
||||
b"\xc3\xe4Ud\xb3yh\x9e,\x08\xeb\xac\xe8\x04-!w\x05\xa7\x8e\xefS0\r\xa5\xe5"
|
||||
b"\xbb=j\xe2\t6o\xfd4\xd3}oF\xff\x87\xda\xa9)'\xaa\xff\xad\xf5\x85\xe6>"
|
||||
b"\x1a\xb8z\x1dJ\xb1\xea\xc0Z\xf70\xdf\x1f\xc2\xa4\xe4\xef?\x91I\x96@\xd5"
|
||||
b"\x19w-7\xc3^\xec\x9d\xa6:D\xa5\xc2\xa4)\xdd\xd5\xba\x9c=E\xb3\xc6,\x18d"
|
||||
b"\xd5\xba=\xdf\xab\x7f\xcdB\xac\xa7\xf1\x18\x0b\xa0X\x15b\x0b\xa4*nC\xc3"
|
||||
b"\xe4\x04\x9f5\xa3G\x8eF\xed3\xa5e\xbd\xbc;)n\x02\x0bW\xdft\xe8\x13\xb475"
|
||||
b"~\x83_ &`\xa6\xdc\xad\x8b\xc6ly\x98\xf7)(1:e3:\x01\x00\x01)(1:d256:!L"
|
||||
b"\x08f\xa2(\xd5\xb4\xfb\x8e\x0fr\x1b\x85\t\x00\xb9\xf2N7\xf0\x1cWK\xe3Q"
|
||||
b"\x7f\x9e#\xa7\xe4:\x98U\x1b\xea\x8bz\x98\x1e\xbc\xd8\xba\xb1\xf9\x89\x12"
|
||||
b"\x18`\xac\xe8\xcc\x0bN\tZ@j\xba/\x99\xf8\xb3$`\x84\xb9\xcei\x95\x9a\xf9"
|
||||
b"\xe2\xfc\x1fQM'\x15\xdb+'\xad\xef\xb4i\xac\xbe}\x10\xeb\x86Gps\xb4\x00"
|
||||
b"\x87\x95\x15;7\xf9\xe7\x14\xe7\x80\xbbh\x1e\x1b\xe6\xdd\xbbsc\xb9g\xe6"
|
||||
b"\xb2'\x7f\xcf\xcf0\x9b\xc2\x98\xfd\xd9\x186/6.\xf1=\x81z\x9f\xe1\x03-G"
|
||||
b"\xdb4Qb9\xddO\xe9\xac\xa8\x8b\xd9\xd6\xf3\x84\xc4\x17\xb9q\x9d\x06\x08Bx"
|
||||
b"M\xbb\xc5*\xf4\xc3X\xcdU+\xed\xbe3_\x04\xea{\xe6\x04$c\xf2-\xd7=\x1bl"
|
||||
b"\xd5\x9ccC/\x92\x88\x8d>n\xda\x187\xd8\x0f%g\x89\x1d\xb9F4^\xc9\xce\xc4"
|
||||
b"\x8b\xed\x92Z3\x07\x0f\xdf\x86\x08\xf9\x92\xe9\xdb\xeb8\x086\xc9\xcd\xcd"
|
||||
b"\n\x01H[9>z\xca\xc6\x80\xa9\xdc\xd49)(1:p129:\x00\xfbD\x17\x8b\xa46\xbe"
|
||||
b"\x1e7\x1d\xa7\xf6al\x04\xc4\xaa\xddx>\x07\x8c\x1e3\x02\xae\x03\x14\x87"
|
||||
b"\x83z\xe5\x9e}\x08g\xa8\xf2\xaa\xbf\x12p\xcfr\xa9\xa7\xc7\x0b\x1d\x88"
|
||||
b"\xd5 \xfd\x9cc\xcaG0UN\x8b\xc4\xcf\xf4\x7f\x16\xa4\x92\x12t\xa1\t\xc2"
|
||||
b"\xc4n\x9c\x8c3\xef\xa5\xe5\xf7\xe0+\xadO\\\x11\xaa\x1a\x847[\xfdz\xea"
|
||||
b"\xc3\xcd|\xb0\xc8\xe4\x1fTc\xb5\xc7\xaf\xdf\xf4\t\xa7\xfc\xc7%\xfc\\\xe9"
|
||||
b"\x91\xd7\x92\xc5\x98\x1eV\xd3\xb1#)(1:q129:\x00\xd9p\x06\xd8\xe2\xbc\xd4"
|
||||
b"x\x91P\x94\xd4\xc1\x1b\x898lFdZQ\xa0\x9a\x07=H\x8f\x03Q\xcck\x12\x8e}"
|
||||
b"\x1a\xb1e\xe7qu9\xe02\x05u\x8d\x18L\xaf\x93\xb1I\xb1f_xbz\xd1\x0c\xca"
|
||||
b"\xe6MC\xb3\x9c\xf4k}\xe6\x0c\x98\xdc\xcf!b\x8e\xd5.\x12\xde\x04\xae\xd7$"
|
||||
b"n\x831\xa2\x15\xa2D=\"\xa9b&\"\xb9\xb2\xedT\n\x9d\x08\x83\xa7\x07\r\xff"
|
||||
b"\x19\x18\x8e\xd8\xab\x1d\xdaH\x9c1h\x11\xa1fm\xe3\xd8\x1d)(1:a128:if7"
|
||||
b"\xc6@\xdd!\xc5\x04\xf3\xb0\xb8>G\x94|v\xfc-\xeb?9<\x95\xc3C\x01Q\xc4B"
|
||||
b"\x97\xf3\xe8\x16\xa4\xc6\xee\xec\xd4I\x10P8\x04\xee;\xcd\xd7\xd0\xcc\xcc"
|
||||
b"2i\x90\x07\xa0\x1bZ\x9f\xfe1\xcd\x1e:~q\x1e\x19\x94\x1aNO\x0c\xdf_R\xd5"
|
||||
b"\xd1\x17n\xec\xd7\x9c\xb6U\x9d\xdb\x8e!_\xbc$\x88\xb6\xfc\xaf\xab\xf0"
|
||||
b"\xef\xa4,\xb0\xdc\x9f\x86\xb0\x03\x12\xb8\x8b\xe2\xdd\x0f\xc0\xee#=JP"
|
||||
b"\xfe/\xb8)FX\xbf\xb6')(1:b128:Q\xaf\xe9\x92\x9f\x94\x0bJ\x84e>\x94\xb3;"
|
||||
b"\x92\x10\xb5t\xb8\x8c\xc9\xef\xc9\x0e\x012\xfa/h\x12\xa1\x03&\xae\xcfQh"
|
||||
b"\x14L&\x9b(\xa4\x023\x08_\xe1\xa7p\x98\x014y^R\x8e\xc4\xcf6\xbc\x1fKU"
|
||||
b"\xac\xeb\xc1S\x84\xc7\xe1a\xa8J\xd4\xa2\xff@\r\x80\x1f\x12\xa9P\xc0*\x18"
|
||||
b"u\x94\x0c\x06\x9b\x16P\xa8K\xecA\xcd{\xef\xf7K\xc9u\x02h\xc4\x98\xb8\x86"
|
||||
b"\x88\x18ZC\xe7\x023\x97\"d\x93\x83\x0cE*|\xed)(1:c128:f\x16\xf9 4\xd0T%"
|
||||
b"\xbca1\xac\x82\xfb\xef\x9c5\x1e~JU\x02h\x95\r\n\x93\xbe\x1e\xbf\xe1@\xfa"
|
||||
b"\x90\xa7Tp1\xe9x\xfc\xe0f\xb7\"w\x9a\xeb\xdd\xd5\xd20F\xca\xe2\xd7^A\x9b"
|
||||
b"\xcf\xb5H5Q\xaa%\xdc\xde\xdb4)H!\xdb\xd6t\xfc~\xe5/S\xf7\x9c\tp\xb9\xe4"
|
||||
b"\xa0v\xa6\xadzt\"4\x9cO\x17\xcb?\xe0\xaa\xe5^\xa5\'\xde?\xa5\x7f\x0f\xa6"
|
||||
b"\x88\xf5\x15\xd6_7\x17\x92\xe0\xd6\x05I\x10;\x92\xf6)))"
|
||||
)
|
||||
|
||||
privateRSA_agentv3 = (
|
||||
b"\x00\x00\x00\x07ssh-rsa\x00\x00\x00\x03\x01\x00\x01\x00\x00\x01\x00!L"
|
||||
b"\x08f\xa2(\xd5\xb4\xfb\x8e\x0fr\x1b\x85\t\x00\xb9\xf2N7\xf0\x1cWK\xe3Q"
|
||||
b"\x7f\x9e#\xa7\xe4:\x98U\x1b\xea\x8bz\x98\x1e\xbc\xd8\xba\xb1\xf9\x89\x12"
|
||||
b"\x18`\xac\xe8\xcc\x0bN\tZ@j\xba/\x99\xf8\xb3$`\x84\xb9\xcei\x95\x9a\xf9"
|
||||
b"\xe2\xfc\x1fQM'\x15\xdb+'\xad\xef\xb4i\xac\xbe}\x10\xeb\x86Gps\xb4\x00"
|
||||
b"\x87\x95\x15;7\xf9\xe7\x14\xe7\x80\xbbh\x1e\x1b\xe6\xdd\xbbsc\xb9g\xe6"
|
||||
b"\xb2'\x7f\xcf\xcf0\x9b\xc2\x98\xfd\xd9\x186/6.\xf1=\x81z\x9f\xe1\x03-G"
|
||||
b"\xdb4Qb9\xddO\xe9\xac\xa8\x8b\xd9\xd6\xf3\x84\xc4\x17\xb9q\x9d\x06\x08Bx"
|
||||
b"M\xbb\xc5*\xf4\xc3X\xcdU+\xed\xbe3_\x04\xea{\xe6\x04$c\xf2-\xd7=\x1bl"
|
||||
b"\xd5\x9ccC/\x92\x88\x8d>n\xda\x187\xd8\x0f%g\x89\x1d\xb9F4^\xc9\xce\xc4"
|
||||
b"\x8b\xed\x92Z3\x07\x0f\xdf\x86\x08\xf9\x92\xe9\xdb\xeb8\x086\xc9\xcd\xcd"
|
||||
b"\n\x01H[9>z\xca\xc6\x80\xa9\xdc\xd49\x00\x00\x01\x01\x00\xd5j\xacx#\xd6"
|
||||
b"\xd6\x1b\xec%\xa1P\xc4wcP\x84E\x01UB\x14**\xe0\xd0`\xee\xd4\xe9\xa3\xadJ"
|
||||
b"\xfa9\x06^\x84Uu_\x006\xbfo\xaa*?\x83&7\xc1i.[\xfd\xf0\xf3\xd2}\xd6\x98"
|
||||
b"\xcd:@x\xd5\xca\xa8\x18\xc0\x11\x93$\t\x0c\x81L\x8f\xf7\x9c\xed\x13\x16j"
|
||||
b"\xa4\x04\xe9Iw\xc3\xe4Ud\xb3yh\x9e,\x08\xeb\xac\xe8\x04-!w\x05\xa7\x8e"
|
||||
b"\xefS0\r\xa5\xe5\xbb=j\xe2\t6o\xfd4\xd3}oF\xff\x87\xda\xa9)'\xaa\xff\xad"
|
||||
b"\xf5\x85\xe6>\x1a\xb8z\x1dJ\xb1\xea\xc0Z\xf70\xdf\x1f\xc2\xa4\xe4\xef?"
|
||||
b"\x91I\x96@\xd5\x19w-7\xc3^\xec\x9d\xa6:D\xa5\xc2\xa4)\xdd\xd5\xba\x9c=E"
|
||||
b"\xb3\xc6,\x18d\xd5\xba=\xdf\xab\x7f\xcdB\xac\xa7\xf1\x18\x0b\xa0X\x15b"
|
||||
b"\x0b\xa4*nC\xc3\xe4\x04\x9f5\xa3G\x8eF\xed3\xa5e\xbd\xbc;)n\x02\x0bW\xdf"
|
||||
b"t\xe8\x13\xb475~\x83_ &`\xa6\xdc\xad\x8b\xc6ly\x98\xf7\x00\x00\x00\x81"
|
||||
b"\x00\x85K\x1bz\x9b\x12\x107\x9e\x1f\xad^\xda\xfe\xc6\x96\xfe\xdf5k\xb94"
|
||||
b"\xe2\x16\x97\x92&\t\xbd\xbdp \x03\xa75\xbd-\x1b\xa0\xd2\x07G+\xd4\xde"
|
||||
b"\xa8\xa8\x07\x07\x1b\xb8\x04 \xa7'A<l99\xe9A\xce\xe7\x17\x1d\xd1L\\\xbc="
|
||||
b"\xd2&&\xfej\xd6\xfdHr\xaeF\xfa{\xc3\xd3\x19`D\x1d\xa5\x13\xa7\x80\xf5c)"
|
||||
b"\xd4z]\x06\x07\x16]\xf6\x8b=\xcbd:\xe2\x84ZM\x8c\x06--\x9d\x1c\xeb\x83Lx"
|
||||
b"=yT\xce\x00\x00\x00\x81\x00\xd9p\x06\xd8\xe2\xbc\xd4x\x91P\x94\xd4\xc1"
|
||||
b"\x1b\x898lFdZQ\xa0\x9a\x07=H\x8f\x03Q\xcck\x12\x8e}\x1a\xb1e\xe7qu9\xe02"
|
||||
b"\x05u\x8d\x18L\xaf\x93\xb1I\xb1f_xbz\xd1\x0c\xca\xe6MC\xb3\x9c\xf4k}\xe6"
|
||||
b"\x0c\x98\xdc\xcf!b\x8e\xd5.\x12\xde\x04\xae\xd7$n\x831\xa2\x15\xa2D=\""
|
||||
b"\xa9b&\"\xb9\xb2\xedT\n\x9d\x08\x83\xa7\x07\r\xff\x19\x18\x8e\xd8\xab"
|
||||
b"\x1d\xdaH\x9c1h\x11\xa1fm\xe3\xd8\x1d\x00\x00\x00\x81\x00\xfbD\x17\x8b"
|
||||
b"\xa46\xbe\x1e7\x1d\xa7\xf6al\x04\xc4\xaa\xddx>\x07\x8c\x1e3\x02\xae\x03"
|
||||
b"\x14\x87\x83z\xe5\x9e}\x08g\xa8\xf2\xaa\xbf\x12p\xcfr\xa9\xa7\xc7\x0b"
|
||||
b"\x1d\x88\xd5 \xfd\x9cc\xcaG0UN\x8b\xc4\xcf\xf4\x7f\x16\xa4\x92\x12t\xa1"
|
||||
b"\t\xc2\xc4n\x9c\x8c3\xef\xa5\xe5\xf7\xe0+\xadO\\\x11\xaa\x1a\x847[\xfdz"
|
||||
b"\xea\xc3\xcd|\xb0\xc8\xe4\x1fTc\xb5\xc7\xaf\xdf\xf4\t\xa7\xfc\xc7%\xfc\\"
|
||||
b"\xe9\x91\xd7\x92\xc5\x98\x1eV\xd3\xb1#"
|
||||
)
|
||||
|
||||
publicDSA_openssh = b"""\
|
||||
ssh-dss AAAAB3NzaC1kc3MAAACBAJKQOsVERVDQIpANHH+JAAylo9\
|
||||
LvFYmFFVMIuHFGlZpIL7sh3IMkqy+cssINM/lnHD3fmsAyLlUXZtt6PD9LgZRazsPOgptuH+Gu48G\
|
||||
+yFuE8l0fVVUivos/MmYVJ66qT99htcZKatrTWZnpVW7gFABoqw+he2LZ0gkeU0+Sx9a5AAAAFQD0\
|
||||
EYmTNaFJ8CS0+vFSF4nYcyEnSQAAAIEAkgLjxHJAE7qFWdTqf7EZngu7jAGmdB9k3YzMHe1ldMxEB\
|
||||
7zNw5aOnxjhoYLtiHeoEcOk2XOyvnE+VfhIWwWAdOiKRTEZlmizkvhGbq0DCe2EPMXirjqWACI5nD\
|
||||
ioQX1oEMonR8N3AEO5v9SfBqS2Q9R6OBr6lf04RvwpHZ0UGu8AAACAAhRpxGMIWEyaEh8YnjiazQT\
|
||||
NEpklRZqeBGo1gotJggNmVaIQNIClGlLyCi359efEUuQcZ9SXxM59P+hecc/GU/GHakW5YWE4dP2G\
|
||||
gdgMQWC7S6WFIXePGGXqNQDdWxlX8umhenvQqa1PnKrFRhDrJw8Z7GjdHxflsxCEmXPoLN8= \
|
||||
comment\
|
||||
"""
|
||||
|
||||
privateDSA_openssh = b"""\
|
||||
-----BEGIN DSA PRIVATE KEY-----
|
||||
MIIBvAIBAAKBgQCSkDrFREVQ0CKQDRx/iQAMpaPS7xWJhRVTCLhxRpWaSC+7IdyD
|
||||
JKsvnLLCDTP5Zxw935rAMi5VF2bbejw/S4GUWs7DzoKbbh/hruPBvshbhPJdH1VV
|
||||
Ir6LPzJmFSeuqk/fYbXGSmra01mZ6VVu4BQAaKsPoXti2dIJHlNPksfWuQIVAPQR
|
||||
iZM1oUnwJLT68VIXidhzISdJAoGBAJIC48RyQBO6hVnU6n+xGZ4Lu4wBpnQfZN2M
|
||||
zB3tZXTMRAe8zcOWjp8Y4aGC7Yh3qBHDpNlzsr5xPlX4SFsFgHToikUxGZZos5L4
|
||||
Rm6tAwnthDzF4q46lgAiOZw4qEF9aBDKJ0fDdwBDub/UnwaktkPUejga+pX9OEb8
|
||||
KR2dFBrvAoGAAhRpxGMIWEyaEh8YnjiazQTNEpklRZqeBGo1gotJggNmVaIQNICl
|
||||
GlLyCi359efEUuQcZ9SXxM59P+hecc/GU/GHakW5YWE4dP2GgdgMQWC7S6WFIXeP
|
||||
GGXqNQDdWxlX8umhenvQqa1PnKrFRhDrJw8Z7GjdHxflsxCEmXPoLN8CFQDV2gbL
|
||||
czUdxCus0pfEP1bddaXRLQ==
|
||||
-----END DSA PRIVATE KEY-----\
|
||||
"""
|
||||
|
||||
privateDSA_openssh_new = b"""\
|
||||
-----BEGIN OPENSSH PRIVATE KEY-----
|
||||
b3BlbnNzaC1rZXktdjEAAAAABG5vbmUAAAAEbm9uZQAAAAAAAAABAAABsgAAAAdzc2gtZH
|
||||
NzAAAAgQCSkDrFREVQ0CKQDRx/iQAMpaPS7xWJhRVTCLhxRpWaSC+7IdyDJKsvnLLCDTP5
|
||||
Zxw935rAMi5VF2bbejw/S4GUWs7DzoKbbh/hruPBvshbhPJdH1VVIr6LPzJmFSeuqk/fYb
|
||||
XGSmra01mZ6VVu4BQAaKsPoXti2dIJHlNPksfWuQAAABUA9BGJkzWhSfAktPrxUheJ2HMh
|
||||
J0kAAACBAJIC48RyQBO6hVnU6n+xGZ4Lu4wBpnQfZN2MzB3tZXTMRAe8zcOWjp8Y4aGC7Y
|
||||
h3qBHDpNlzsr5xPlX4SFsFgHToikUxGZZos5L4Rm6tAwnthDzF4q46lgAiOZw4qEF9aBDK
|
||||
J0fDdwBDub/UnwaktkPUejga+pX9OEb8KR2dFBrvAAAAgAIUacRjCFhMmhIfGJ44ms0EzR
|
||||
KZJUWangRqNYKLSYIDZlWiEDSApRpS8got+fXnxFLkHGfUl8TOfT/oXnHPxlPxh2pFuWFh
|
||||
OHT9hoHYDEFgu0ulhSF3jxhl6jUA3VsZV/LpoXp70KmtT5yqxUYQ6ycPGexo3R8X5bMQhJ
|
||||
lz6CzfAAAB2MVcBjzFXAY8AAAAB3NzaC1kc3MAAACBAJKQOsVERVDQIpANHH+JAAylo9Lv
|
||||
FYmFFVMIuHFGlZpIL7sh3IMkqy+cssINM/lnHD3fmsAyLlUXZtt6PD9LgZRazsPOgptuH+
|
||||
Gu48G+yFuE8l0fVVUivos/MmYVJ66qT99htcZKatrTWZnpVW7gFABoqw+he2LZ0gkeU0+S
|
||||
x9a5AAAAFQD0EYmTNaFJ8CS0+vFSF4nYcyEnSQAAAIEAkgLjxHJAE7qFWdTqf7EZngu7jA
|
||||
GmdB9k3YzMHe1ldMxEB7zNw5aOnxjhoYLtiHeoEcOk2XOyvnE+VfhIWwWAdOiKRTEZlmiz
|
||||
kvhGbq0DCe2EPMXirjqWACI5nDioQX1oEMonR8N3AEO5v9SfBqS2Q9R6OBr6lf04RvwpHZ
|
||||
0UGu8AAACAAhRpxGMIWEyaEh8YnjiazQTNEpklRZqeBGo1gotJggNmVaIQNIClGlLyCi35
|
||||
9efEUuQcZ9SXxM59P+hecc/GU/GHakW5YWE4dP2GgdgMQWC7S6WFIXePGGXqNQDdWxlX8u
|
||||
mhenvQqa1PnKrFRhDrJw8Z7GjdHxflsxCEmXPoLN8AAAAVANXaBstzNR3EK6zSl8Q/Vt11
|
||||
pdEtAAAAAAE=
|
||||
-----END OPENSSH PRIVATE KEY-----\
|
||||
"""
|
||||
|
||||
publicDSA_lsh = decodebytes(b"""\
|
||||
e0tERXdPbkIxWW14cFl5MXJaWGtvTXpwa2MyRW9NVHB3TVRJNU9nQ1NrRHJGUkVWUTBDS1FEUngv
|
||||
aVFBTXBhUFM3eFdKaFJWVENMaHhScFdhU0MrN0lkeURKS3N2bkxMQ0RUUDVaeHc5MzVyQU1pNVZG
|
||||
MmJiZWp3L1M0R1VXczdEem9LYmJoL2hydVBCdnNoYmhQSmRIMVZWSXI2TFB6Sm1GU2V1cWsvZlli
|
||||
WEdTbXJhMDFtWjZWVnU0QlFBYUtzUG9YdGkyZElKSGxOUGtzZld1U2tvTVRweE1qRTZBUFFSaVpN
|
||||
MW9VbndKTFQ2OFZJWGlkaHpJU2RKS1NneE9tY3hNams2QUpJQzQ4UnlRQk82aFZuVTZuK3hHWjRM
|
||||
dTR3QnBuUWZaTjJNekIzdFpYVE1SQWU4emNPV2pwOFk0YUdDN1loM3FCSERwTmx6c3I1eFBsWDRT
|
||||
RnNGZ0hUb2lrVXhHWlpvczVMNFJtNnRBd250aER6RjRxNDZsZ0FpT1p3NHFFRjlhQkRLSjBmRGR3
|
||||
QkR1Yi9Vbndha3RrUFVlamdhK3BYOU9FYjhLUjJkRkJydktTZ3hPbmt4TWpnNkFoUnB4R01JV0V5
|
||||
YUVoOFluamlhelFUTkVwa2xSWnFlQkdvMWdvdEpnZ05tVmFJUU5JQ2xHbEx5Q2kzNTllZkVVdVFj
|
||||
WjlTWHhNNTlQK2hlY2MvR1UvR0hha1c1WVdFNGRQMkdnZGdNUVdDN1M2V0ZJWGVQR0dYcU5RRGRX
|
||||
eGxYOHVtaGVudlFxYTFQbktyRlJoRHJKdzhaN0dqZEh4ZmxzeENFbVhQb0xOOHBLU2s9fQ==
|
||||
""")
|
||||
|
||||
privateDSA_lsh = decodebytes(b"""\
|
||||
KDExOnByaXZhdGUta2V5KDM6ZHNhKDE6cDEyOToAkpA6xURFUNAikA0cf4kADKWj0u8ViYUVUwi4
|
||||
cUaVmkgvuyHcgySrL5yywg0z+WccPd+awDIuVRdm23o8P0uBlFrOw86Cm24f4a7jwb7IW4TyXR9V
|
||||
VSK+iz8yZhUnrqpP32G1xkpq2tNZmelVbuAUAGirD6F7YtnSCR5TT5LH1rkpKDE6cTIxOgD0EYmT
|
||||
NaFJ8CS0+vFSF4nYcyEnSSkoMTpnMTI5OgCSAuPEckATuoVZ1Op/sRmeC7uMAaZ0H2TdjMwd7WV0
|
||||
zEQHvM3Dlo6fGOGhgu2Id6gRw6TZc7K+cT5V+EhbBYB06IpFMRmWaLOS+EZurQMJ7YQ8xeKuOpYA
|
||||
IjmcOKhBfWgQyidHw3cAQ7m/1J8GpLZD1Ho4GvqV/ThG/CkdnRQa7ykoMTp5MTI4OgIUacRjCFhM
|
||||
mhIfGJ44ms0EzRKZJUWangRqNYKLSYIDZlWiEDSApRpS8got+fXnxFLkHGfUl8TOfT/oXnHPxlPx
|
||||
h2pFuWFhOHT9hoHYDEFgu0ulhSF3jxhl6jUA3VsZV/LpoXp70KmtT5yqxUYQ6ycPGexo3R8X5bMQ
|
||||
hJlz6CzfKSgxOngyMToA1doGy3M1HcQrrNKXxD9W3XWl0S0pKSk=
|
||||
""")
|
||||
|
||||
privateDSA_agentv3 = decodebytes(b"""\
|
||||
AAAAB3NzaC1kc3MAAACBAJKQOsVERVDQIpANHH+JAAylo9LvFYmFFVMIuHFGlZpIL7sh3IMkqy+c
|
||||
ssINM/lnHD3fmsAyLlUXZtt6PD9LgZRazsPOgptuH+Gu48G+yFuE8l0fVVUivos/MmYVJ66qT99h
|
||||
tcZKatrTWZnpVW7gFABoqw+he2LZ0gkeU0+Sx9a5AAAAFQD0EYmTNaFJ8CS0+vFSF4nYcyEnSQAA
|
||||
AIEAkgLjxHJAE7qFWdTqf7EZngu7jAGmdB9k3YzMHe1ldMxEB7zNw5aOnxjhoYLtiHeoEcOk2XOy
|
||||
vnE+VfhIWwWAdOiKRTEZlmizkvhGbq0DCe2EPMXirjqWACI5nDioQX1oEMonR8N3AEO5v9SfBqS2
|
||||
Q9R6OBr6lf04RvwpHZ0UGu8AAACAAhRpxGMIWEyaEh8YnjiazQTNEpklRZqeBGo1gotJggNmVaIQ
|
||||
NIClGlLyCi359efEUuQcZ9SXxM59P+hecc/GU/GHakW5YWE4dP2GgdgMQWC7S6WFIXePGGXqNQDd
|
||||
WxlX8umhenvQqa1PnKrFRhDrJw8Z7GjdHxflsxCEmXPoLN8AAAAVANXaBstzNR3EK6zSl8Q/Vt11
|
||||
pdEt
|
||||
""")
|
||||
|
||||
__all__ = ['DSAData', 'RSAData', 'privateDSA_agentv3', 'privateDSA_lsh',
|
||||
'privateDSA_openssh', 'privateRSA_agentv3', 'privateRSA_lsh',
|
||||
'privateRSA_openssh', 'publicDSA_lsh', 'publicDSA_openssh',
|
||||
'publicRSA_lsh', 'publicRSA_openssh', 'privateRSA_openssh_alternate']
|
||||
@@ -0,0 +1,28 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
"""
|
||||
Loopback helper used in test_ssh and test_recvline
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
from twisted.protocols import loopback
|
||||
class LoopbackRelay(loopback.LoopbackRelay):
|
||||
clearCall = None
|
||||
|
||||
def logPrefix(self):
|
||||
return "LoopbackRelay(%r)" % (self.target.__class__.__name__,)
|
||||
|
||||
|
||||
def write(self, data):
|
||||
loopback.LoopbackRelay.write(self, data)
|
||||
if self.clearCall is not None:
|
||||
self.clearCall.cancel()
|
||||
|
||||
from twisted.internet import reactor
|
||||
self.clearCall = reactor.callLater(0, self._clearBuffer)
|
||||
|
||||
|
||||
def _clearBuffer(self):
|
||||
self.clearCall = None
|
||||
loopback.LoopbackRelay.clearBuffer(self)
|
||||
@@ -0,0 +1,50 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{SSHTransportAddrress} in ssh/address.py
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
from twisted.trial import unittest
|
||||
from twisted.internet.address import IPv4Address
|
||||
from twisted.internet.test.test_address import AddressTestCaseMixin
|
||||
|
||||
from twisted.conch.ssh.address import SSHTransportAddress
|
||||
|
||||
|
||||
|
||||
class SSHTransportAddressTests(unittest.TestCase, AddressTestCaseMixin):
|
||||
"""
|
||||
L{twisted.conch.ssh.address.SSHTransportAddress} is what Conch transports
|
||||
use to represent the other side of the SSH connection. This tests the
|
||||
basic functionality of that class (string representation, comparison, &c).
|
||||
"""
|
||||
|
||||
|
||||
def _stringRepresentation(self, stringFunction):
|
||||
"""
|
||||
The string representation of C{SSHTransportAddress} should be
|
||||
"SSHTransportAddress(<stringFunction on address>)".
|
||||
"""
|
||||
addr = self.buildAddress()
|
||||
stringValue = stringFunction(addr)
|
||||
addressValue = stringFunction(addr.address)
|
||||
self.assertEqual(stringValue,
|
||||
"SSHTransportAddress(%s)" % addressValue)
|
||||
|
||||
|
||||
def buildAddress(self):
|
||||
"""
|
||||
Create an arbitrary new C{SSHTransportAddress}. A new instance is
|
||||
created for each call, but always for the same address.
|
||||
"""
|
||||
return SSHTransportAddress(IPv4Address("TCP", "127.0.0.1", 22))
|
||||
|
||||
|
||||
def buildDifferentAddress(self):
|
||||
"""
|
||||
Like C{buildAddress}, but with a different fixed address.
|
||||
"""
|
||||
return SSHTransportAddress(IPv4Address("TCP", "127.0.0.2", 22))
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,567 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.conch.scripts.ckeygen}.
|
||||
"""
|
||||
|
||||
import getpass
|
||||
import sys
|
||||
import subprocess
|
||||
|
||||
from io import BytesIO, StringIO
|
||||
|
||||
from twisted.python.compat import unicode, _PY3
|
||||
from twisted.python.reflect import requireModule
|
||||
|
||||
if requireModule('cryptography') and requireModule('pyasn1'):
|
||||
from twisted.conch.ssh.keys import (Key, BadKeyError,
|
||||
BadFingerPrintFormat, FingerprintFormats)
|
||||
from twisted.conch.scripts.ckeygen import (
|
||||
changePassPhrase, displayPublicKey, printFingerprint,
|
||||
_saveKey, enumrepresentation)
|
||||
else:
|
||||
skip = "cryptography and pyasn1 required for twisted.conch.scripts.ckeygen"
|
||||
|
||||
from twisted.python.filepath import FilePath
|
||||
from twisted.trial.unittest import TestCase
|
||||
from twisted.conch.test.keydata import (
|
||||
publicRSA_openssh, privateRSA_openssh, privateRSA_openssh_encrypted, privateECDSA_openssh)
|
||||
|
||||
|
||||
|
||||
def makeGetpass(*passphrases):
|
||||
"""
|
||||
Return a callable to patch C{getpass.getpass}. Yields a passphrase each
|
||||
time called. Use case is to provide an old, then new passphrase(s) as if
|
||||
requested interactively.
|
||||
|
||||
@param passphrases: The list of passphrases returned, one per each call.
|
||||
|
||||
@return: A callable to patch C{getpass.getpass}.
|
||||
"""
|
||||
passphrases = iter(passphrases)
|
||||
|
||||
def fakeGetpass(_):
|
||||
return next(passphrases)
|
||||
|
||||
return fakeGetpass
|
||||
|
||||
|
||||
|
||||
class KeyGenTests(TestCase):
|
||||
"""
|
||||
Tests for various functions used to implement the I{ckeygen} script.
|
||||
"""
|
||||
def setUp(self):
|
||||
"""
|
||||
Patch C{sys.stdout} so tests can make assertions about what's printed.
|
||||
"""
|
||||
if _PY3:
|
||||
self.stdout = StringIO()
|
||||
else:
|
||||
self.stdout = BytesIO()
|
||||
self.patch(sys, 'stdout', self.stdout)
|
||||
|
||||
|
||||
|
||||
def _testrun(self, keyType, keySize=None):
|
||||
filename = self.mktemp()
|
||||
if keySize is None:
|
||||
subprocess.call(['ckeygen', '-t', keyType, '-f', filename, '--no-passphrase'])
|
||||
else:
|
||||
subprocess.call(['ckeygen', '-t', keyType, '-f', filename, '--no-passphrase',
|
||||
'-b', keySize])
|
||||
privKey = Key.fromFile(filename)
|
||||
pubKey = Key.fromFile(filename + '.pub')
|
||||
if keyType == 'ecdsa':
|
||||
self.assertEqual(privKey.type(), 'EC')
|
||||
else:
|
||||
self.assertEqual(privKey.type(), keyType.upper())
|
||||
self.assertTrue(pubKey.isPublic())
|
||||
|
||||
|
||||
def test_keygeneration(self):
|
||||
self._testrun('ecdsa', '384')
|
||||
self._testrun('ecdsa')
|
||||
self._testrun('dsa', '2048')
|
||||
self._testrun('dsa')
|
||||
self._testrun('rsa', '2048')
|
||||
self._testrun('rsa')
|
||||
|
||||
|
||||
|
||||
def test_runBadKeytype(self):
|
||||
filename = self.mktemp()
|
||||
with self.assertRaises(subprocess.CalledProcessError):
|
||||
subprocess.check_call(['ckeygen', '-t', 'foo', '-f', filename])
|
||||
|
||||
|
||||
|
||||
def test_enumrepresentation(self):
|
||||
"""
|
||||
L{enumrepresentation} takes a dictionary as input and returns a
|
||||
dictionary with its attributes changed to enum representation.
|
||||
"""
|
||||
options = enumrepresentation({'format': 'md5-hex'})
|
||||
self.assertIs(options['format'],
|
||||
FingerprintFormats.MD5_HEX)
|
||||
|
||||
|
||||
def test_enumrepresentationsha256(self):
|
||||
"""
|
||||
Test for format L{FingerprintFormats.SHA256-BASE64}.
|
||||
"""
|
||||
options = enumrepresentation({'format': 'sha256-base64'})
|
||||
self.assertIs(options['format'],
|
||||
FingerprintFormats.SHA256_BASE64)
|
||||
|
||||
|
||||
|
||||
def test_enumrepresentationBadFormat(self):
|
||||
"""
|
||||
Test for unsupported fingerprint format
|
||||
"""
|
||||
with self.assertRaises(BadFingerPrintFormat) as em:
|
||||
enumrepresentation({'format': 'sha-base64'})
|
||||
self.assertEqual('Unsupported fingerprint format: sha-base64',
|
||||
em.exception.args[0])
|
||||
|
||||
|
||||
|
||||
def test_printFingerprint(self):
|
||||
"""
|
||||
L{printFingerprint} writes a line to standard out giving the number of
|
||||
bits of the key, its fingerprint, and the basename of the file from it
|
||||
was read.
|
||||
"""
|
||||
filename = self.mktemp()
|
||||
FilePath(filename).setContent(publicRSA_openssh)
|
||||
printFingerprint({'filename': filename,
|
||||
'format': 'md5-hex'})
|
||||
self.assertEqual(
|
||||
self.stdout.getvalue(),
|
||||
'2048 85:25:04:32:58:55:96:9f:57:ee:fb:a8:1a:ea:69:da temp\n')
|
||||
|
||||
|
||||
def test_printFingerprintsha256(self):
|
||||
"""
|
||||
L{printFigerprint} will print key fingerprint in
|
||||
L{FingerprintFormats.SHA256-BASE64} format if explicitly specified.
|
||||
"""
|
||||
filename = self.mktemp()
|
||||
FilePath(filename).setContent(publicRSA_openssh)
|
||||
printFingerprint({'filename': filename,
|
||||
'format': 'sha256-base64'})
|
||||
self.assertEqual(
|
||||
self.stdout.getvalue(),
|
||||
'2048 FBTCOoknq0mHy+kpfnY9tDdcAJuWtCpuQMaV3EsvbUI= temp\n')
|
||||
|
||||
|
||||
def test_printFingerprintBadFingerPrintFormat(self):
|
||||
"""
|
||||
L{printFigerprint} raises C{keys.BadFingerprintFormat} when unsupported
|
||||
formats are requested.
|
||||
"""
|
||||
filename = self.mktemp()
|
||||
FilePath(filename).setContent(publicRSA_openssh)
|
||||
with self.assertRaises(BadFingerPrintFormat) as em:
|
||||
printFingerprint({'filename': filename, 'format':'sha-base64'})
|
||||
self.assertEqual('Unsupported fingerprint format: sha-base64',
|
||||
em.exception.args[0])
|
||||
|
||||
|
||||
|
||||
def test_saveKey(self):
|
||||
"""
|
||||
L{_saveKey} writes the private and public parts of a key to two
|
||||
different files and writes a report of this to standard out.
|
||||
"""
|
||||
base = FilePath(self.mktemp())
|
||||
base.makedirs()
|
||||
filename = base.child('id_rsa').path
|
||||
key = Key.fromString(privateRSA_openssh)
|
||||
_saveKey(key, {'filename': filename, 'pass': 'passphrase',
|
||||
'format': 'md5-hex'})
|
||||
self.assertEqual(
|
||||
self.stdout.getvalue(),
|
||||
"Your identification has been saved in %s\n"
|
||||
"Your public key has been saved in %s.pub\n"
|
||||
"The key fingerprint in <FingerprintFormats=MD5_HEX> is:\n"
|
||||
"85:25:04:32:58:55:96:9f:57:ee:fb:a8:1a:ea:69:da\n" % (
|
||||
filename,
|
||||
filename))
|
||||
self.assertEqual(
|
||||
key.fromString(
|
||||
base.child('id_rsa').getContent(), None, 'passphrase'),
|
||||
key)
|
||||
self.assertEqual(
|
||||
Key.fromString(base.child('id_rsa.pub').getContent()),
|
||||
key.public())
|
||||
|
||||
|
||||
def test_saveKeyECDSA(self):
|
||||
"""
|
||||
L{_saveKey} writes the private and public parts of a key to two
|
||||
different files and writes a report of this to standard out.
|
||||
Test with ECDSA key.
|
||||
"""
|
||||
base = FilePath(self.mktemp())
|
||||
base.makedirs()
|
||||
filename = base.child('id_ecdsa').path
|
||||
key = Key.fromString(privateECDSA_openssh)
|
||||
_saveKey(key, {'filename': filename, 'pass': 'passphrase',
|
||||
'format': 'md5-hex'})
|
||||
self.assertEqual(
|
||||
self.stdout.getvalue(),
|
||||
"Your identification has been saved in %s\n"
|
||||
"Your public key has been saved in %s.pub\n"
|
||||
"The key fingerprint in <FingerprintFormats=MD5_HEX> is:\n"
|
||||
"1e:ab:83:a6:f2:04:22:99:7c:64:14:d2:ab:fa:f5:16\n" % (
|
||||
filename,
|
||||
filename))
|
||||
self.assertEqual(
|
||||
key.fromString(
|
||||
base.child('id_ecdsa').getContent(), None, 'passphrase'),
|
||||
key)
|
||||
self.assertEqual(
|
||||
Key.fromString(base.child('id_ecdsa.pub').getContent()),
|
||||
key.public())
|
||||
|
||||
|
||||
def test_saveKeysha256(self):
|
||||
"""
|
||||
L{_saveKey} will generate key fingerprint in
|
||||
L{FingerprintFormats.SHA256-BASE64} format if explicitly specified.
|
||||
"""
|
||||
base = FilePath(self.mktemp())
|
||||
base.makedirs()
|
||||
filename = base.child('id_rsa').path
|
||||
key = Key.fromString(privateRSA_openssh)
|
||||
_saveKey(key, {'filename': filename, 'pass': 'passphrase',
|
||||
'format': 'sha256-base64'})
|
||||
self.assertEqual(
|
||||
self.stdout.getvalue(),
|
||||
"Your identification has been saved in %s\n"
|
||||
"Your public key has been saved in %s.pub\n"
|
||||
"The key fingerprint in <FingerprintFormats=SHA256_BASE64> is:\n"
|
||||
"FBTCOoknq0mHy+kpfnY9tDdcAJuWtCpuQMaV3EsvbUI=\n" % (
|
||||
filename,
|
||||
filename))
|
||||
self.assertEqual(
|
||||
key.fromString(
|
||||
base.child('id_rsa').getContent(), None, 'passphrase'),
|
||||
key)
|
||||
self.assertEqual(
|
||||
Key.fromString(base.child('id_rsa.pub').getContent()),
|
||||
key.public())
|
||||
|
||||
|
||||
def test_saveKeyBadFingerPrintformat(self):
|
||||
"""
|
||||
L{_saveKey} raises C{keys.BadFingerprintFormat} when unsupported
|
||||
formats are requested.
|
||||
"""
|
||||
base = FilePath(self.mktemp())
|
||||
base.makedirs()
|
||||
filename = base.child('id_rsa').path
|
||||
key = Key.fromString(privateRSA_openssh)
|
||||
with self.assertRaises(BadFingerPrintFormat) as em:
|
||||
_saveKey(key, {'filename': filename, 'pass': 'passphrase',
|
||||
'format': 'sha-base64'})
|
||||
self.assertEqual('Unsupported fingerprint format: sha-base64',
|
||||
em.exception.args[0])
|
||||
|
||||
|
||||
def test_saveKeyEmptyPassphrase(self):
|
||||
"""
|
||||
L{_saveKey} will choose an empty string for the passphrase if
|
||||
no-passphrase is C{True}.
|
||||
"""
|
||||
base = FilePath(self.mktemp())
|
||||
base.makedirs()
|
||||
filename = base.child('id_rsa').path
|
||||
key = Key.fromString(privateRSA_openssh)
|
||||
_saveKey(key, {'filename': filename, 'no-passphrase': True,
|
||||
'format': 'md5-hex'})
|
||||
self.assertEqual(
|
||||
key.fromString(
|
||||
base.child('id_rsa').getContent(), None, b''),
|
||||
key)
|
||||
|
||||
|
||||
def test_saveKeyECDSAEmptyPassphrase(self):
|
||||
"""
|
||||
L{_saveKey} will choose an empty string for the passphrase if
|
||||
no-passphrase is C{True}.
|
||||
"""
|
||||
base = FilePath(self.mktemp())
|
||||
base.makedirs()
|
||||
filename = base.child('id_ecdsa').path
|
||||
key = Key.fromString(privateECDSA_openssh)
|
||||
_saveKey(key, {'filename': filename, 'no-passphrase': True,
|
||||
'format': 'md5-hex'})
|
||||
self.assertEqual(
|
||||
key.fromString(
|
||||
base.child('id_ecdsa').getContent(), None),
|
||||
key)
|
||||
|
||||
|
||||
|
||||
def test_saveKeyNoFilename(self):
|
||||
"""
|
||||
When no path is specified, it will ask for the path used to store the
|
||||
key.
|
||||
"""
|
||||
base = FilePath(self.mktemp())
|
||||
base.makedirs()
|
||||
keyPath = base.child('custom_key').path
|
||||
|
||||
import twisted.conch.scripts.ckeygen
|
||||
self.patch(twisted.conch.scripts.ckeygen, 'raw_input', lambda _: keyPath)
|
||||
key = Key.fromString(privateRSA_openssh)
|
||||
_saveKey(key, {'filename': None, 'no-passphrase': True,
|
||||
'format': 'md5-hex'})
|
||||
|
||||
persistedKeyContent = base.child('custom_key').getContent()
|
||||
persistedKey = key.fromString(persistedKeyContent, None, b'')
|
||||
self.assertEqual(key, persistedKey)
|
||||
|
||||
|
||||
def test_displayPublicKey(self):
|
||||
"""
|
||||
L{displayPublicKey} prints out the public key associated with a given
|
||||
private key.
|
||||
"""
|
||||
filename = self.mktemp()
|
||||
pubKey = Key.fromString(publicRSA_openssh)
|
||||
FilePath(filename).setContent(privateRSA_openssh)
|
||||
displayPublicKey({'filename': filename})
|
||||
displayed = self.stdout.getvalue().strip('\n')
|
||||
if isinstance(displayed, unicode):
|
||||
displayed = displayed.encode("ascii")
|
||||
self.assertEqual(
|
||||
displayed,
|
||||
pubKey.toString('openssh'))
|
||||
|
||||
|
||||
def test_displayPublicKeyEncrypted(self):
|
||||
"""
|
||||
L{displayPublicKey} prints out the public key associated with a given
|
||||
private key using the given passphrase when it's encrypted.
|
||||
"""
|
||||
filename = self.mktemp()
|
||||
pubKey = Key.fromString(publicRSA_openssh)
|
||||
FilePath(filename).setContent(privateRSA_openssh_encrypted)
|
||||
displayPublicKey({'filename': filename, 'pass': 'encrypted'})
|
||||
displayed = self.stdout.getvalue().strip('\n')
|
||||
if isinstance(displayed, unicode):
|
||||
displayed = displayed.encode("ascii")
|
||||
self.assertEqual(
|
||||
displayed,
|
||||
pubKey.toString('openssh'))
|
||||
|
||||
|
||||
def test_displayPublicKeyEncryptedPassphrasePrompt(self):
|
||||
"""
|
||||
L{displayPublicKey} prints out the public key associated with a given
|
||||
private key, asking for the passphrase when it's encrypted.
|
||||
"""
|
||||
filename = self.mktemp()
|
||||
pubKey = Key.fromString(publicRSA_openssh)
|
||||
FilePath(filename).setContent(privateRSA_openssh_encrypted)
|
||||
self.patch(getpass, 'getpass', lambda x: 'encrypted')
|
||||
displayPublicKey({'filename': filename})
|
||||
displayed = self.stdout.getvalue().strip('\n')
|
||||
if isinstance(displayed, unicode):
|
||||
displayed = displayed.encode("ascii")
|
||||
self.assertEqual(
|
||||
displayed,
|
||||
pubKey.toString('openssh'))
|
||||
|
||||
|
||||
def test_displayPublicKeyWrongPassphrase(self):
|
||||
"""
|
||||
L{displayPublicKey} fails with a L{BadKeyError} when trying to decrypt
|
||||
an encrypted key with the wrong password.
|
||||
"""
|
||||
filename = self.mktemp()
|
||||
FilePath(filename).setContent(privateRSA_openssh_encrypted)
|
||||
self.assertRaises(
|
||||
BadKeyError, displayPublicKey,
|
||||
{'filename': filename, 'pass': 'wrong'})
|
||||
|
||||
|
||||
def test_changePassphrase(self):
|
||||
"""
|
||||
L{changePassPhrase} allows a user to change the passphrase of a
|
||||
private key interactively.
|
||||
"""
|
||||
oldNewConfirm = makeGetpass('encrypted', 'newpass', 'newpass')
|
||||
self.patch(getpass, 'getpass', oldNewConfirm)
|
||||
|
||||
filename = self.mktemp()
|
||||
FilePath(filename).setContent(privateRSA_openssh_encrypted)
|
||||
|
||||
changePassPhrase({'filename': filename})
|
||||
self.assertEqual(
|
||||
self.stdout.getvalue().strip('\n'),
|
||||
'Your identification has been saved with the new passphrase.')
|
||||
self.assertNotEqual(privateRSA_openssh_encrypted,
|
||||
FilePath(filename).getContent())
|
||||
|
||||
|
||||
def test_changePassphraseWithOld(self):
|
||||
"""
|
||||
L{changePassPhrase} allows a user to change the passphrase of a
|
||||
private key, providing the old passphrase and prompting for new one.
|
||||
"""
|
||||
newConfirm = makeGetpass('newpass', 'newpass')
|
||||
self.patch(getpass, 'getpass', newConfirm)
|
||||
|
||||
filename = self.mktemp()
|
||||
FilePath(filename).setContent(privateRSA_openssh_encrypted)
|
||||
|
||||
changePassPhrase({'filename': filename, 'pass': 'encrypted'})
|
||||
self.assertEqual(
|
||||
self.stdout.getvalue().strip('\n'),
|
||||
'Your identification has been saved with the new passphrase.')
|
||||
self.assertNotEqual(privateRSA_openssh_encrypted,
|
||||
FilePath(filename).getContent())
|
||||
|
||||
|
||||
def test_changePassphraseWithBoth(self):
|
||||
"""
|
||||
L{changePassPhrase} allows a user to change the passphrase of a private
|
||||
key by providing both old and new passphrases without prompting.
|
||||
"""
|
||||
filename = self.mktemp()
|
||||
FilePath(filename).setContent(privateRSA_openssh_encrypted)
|
||||
|
||||
changePassPhrase(
|
||||
{'filename': filename, 'pass': 'encrypted',
|
||||
'newpass': 'newencrypt'})
|
||||
self.assertEqual(
|
||||
self.stdout.getvalue().strip('\n'),
|
||||
'Your identification has been saved with the new passphrase.')
|
||||
self.assertNotEqual(privateRSA_openssh_encrypted,
|
||||
FilePath(filename).getContent())
|
||||
|
||||
|
||||
def test_changePassphraseWrongPassphrase(self):
|
||||
"""
|
||||
L{changePassPhrase} exits if passed an invalid old passphrase when
|
||||
trying to change the passphrase of a private key.
|
||||
"""
|
||||
filename = self.mktemp()
|
||||
FilePath(filename).setContent(privateRSA_openssh_encrypted)
|
||||
error = self.assertRaises(
|
||||
SystemExit, changePassPhrase,
|
||||
{'filename': filename, 'pass': 'wrong'})
|
||||
self.assertEqual('Could not change passphrase: old passphrase error',
|
||||
str(error))
|
||||
self.assertEqual(privateRSA_openssh_encrypted,
|
||||
FilePath(filename).getContent())
|
||||
|
||||
|
||||
def test_changePassphraseEmptyGetPass(self):
|
||||
"""
|
||||
L{changePassPhrase} exits if no passphrase is specified for the
|
||||
C{getpass} call and the key is encrypted.
|
||||
"""
|
||||
self.patch(getpass, 'getpass', makeGetpass(''))
|
||||
filename = self.mktemp()
|
||||
FilePath(filename).setContent(privateRSA_openssh_encrypted)
|
||||
error = self.assertRaises(
|
||||
SystemExit, changePassPhrase, {'filename': filename})
|
||||
self.assertEqual(
|
||||
'Could not change passphrase: Passphrase must be provided '
|
||||
'for an encrypted key',
|
||||
str(error))
|
||||
self.assertEqual(privateRSA_openssh_encrypted,
|
||||
FilePath(filename).getContent())
|
||||
|
||||
|
||||
def test_changePassphraseBadKey(self):
|
||||
"""
|
||||
L{changePassPhrase} exits if the file specified points to an invalid
|
||||
key.
|
||||
"""
|
||||
filename = self.mktemp()
|
||||
FilePath(filename).setContent(b'foobar')
|
||||
error = self.assertRaises(
|
||||
SystemExit, changePassPhrase, {'filename': filename})
|
||||
|
||||
if _PY3:
|
||||
expected = "Could not change passphrase: cannot guess the type of b'foobar'"
|
||||
else:
|
||||
expected = "Could not change passphrase: cannot guess the type of 'foobar'"
|
||||
self.assertEqual(expected, str(error))
|
||||
self.assertEqual(b'foobar', FilePath(filename).getContent())
|
||||
|
||||
|
||||
def test_changePassphraseCreateError(self):
|
||||
"""
|
||||
L{changePassPhrase} doesn't modify the key file if an unexpected error
|
||||
happens when trying to create the key with the new passphrase.
|
||||
"""
|
||||
filename = self.mktemp()
|
||||
FilePath(filename).setContent(privateRSA_openssh)
|
||||
|
||||
def toString(*args, **kwargs):
|
||||
raise RuntimeError('oops')
|
||||
|
||||
self.patch(Key, 'toString', toString)
|
||||
|
||||
error = self.assertRaises(
|
||||
SystemExit, changePassPhrase,
|
||||
{'filename': filename,
|
||||
'newpass': 'newencrypt'})
|
||||
|
||||
self.assertEqual(
|
||||
'Could not change passphrase: oops', str(error))
|
||||
|
||||
self.assertEqual(privateRSA_openssh, FilePath(filename).getContent())
|
||||
|
||||
|
||||
def test_changePassphraseEmptyStringError(self):
|
||||
"""
|
||||
L{changePassPhrase} doesn't modify the key file if C{toString} returns
|
||||
an empty string.
|
||||
"""
|
||||
filename = self.mktemp()
|
||||
FilePath(filename).setContent(privateRSA_openssh)
|
||||
|
||||
def toString(*args, **kwargs):
|
||||
return ''
|
||||
|
||||
self.patch(Key, 'toString', toString)
|
||||
|
||||
error = self.assertRaises(
|
||||
SystemExit, changePassPhrase,
|
||||
{'filename': filename, 'newpass': 'newencrypt'})
|
||||
|
||||
if _PY3:
|
||||
expected = (
|
||||
"Could not change passphrase: cannot guess the type of b''")
|
||||
else:
|
||||
expected = (
|
||||
"Could not change passphrase: cannot guess the type of ''")
|
||||
self.assertEqual(expected, str(error))
|
||||
|
||||
self.assertEqual(privateRSA_openssh, FilePath(filename).getContent())
|
||||
|
||||
|
||||
def test_changePassphrasePublicKey(self):
|
||||
"""
|
||||
L{changePassPhrase} exits when trying to change the passphrase on a
|
||||
public key, and doesn't change the file.
|
||||
"""
|
||||
filename = self.mktemp()
|
||||
FilePath(filename).setContent(publicRSA_openssh)
|
||||
error = self.assertRaises(
|
||||
SystemExit, changePassPhrase,
|
||||
{'filename': filename, 'newpass': 'pass'})
|
||||
self.assertEqual(
|
||||
'Could not change passphrase: key not encrypted', str(error))
|
||||
self.assertEqual(publicRSA_openssh, FilePath(filename).getContent())
|
||||
@@ -0,0 +1,813 @@
|
||||
# -*- test-case-name: twisted.conch.test.test_conch -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
import os, sys, socket
|
||||
import subprocess
|
||||
from itertools import count
|
||||
|
||||
from zope.interface import implementer
|
||||
from twisted.python.reflect import requireModule
|
||||
from twisted.conch.error import ConchError
|
||||
from twisted.cred import portal
|
||||
from twisted.internet import reactor, defer, protocol
|
||||
from twisted.internet.error import ProcessExitedAlready
|
||||
from twisted.internet.task import LoopingCall
|
||||
from twisted.internet.utils import getProcessValue
|
||||
from twisted.python import filepath, log, runtime
|
||||
from twisted.python.compat import unicode, _PYPY
|
||||
from twisted.trial import unittest
|
||||
from twisted.conch.test.test_ssh import ConchTestRealm
|
||||
from twisted.python.procutils import which
|
||||
|
||||
from twisted.conch.test.keydata import publicRSA_openssh, privateRSA_openssh
|
||||
from twisted.conch.test.keydata import publicDSA_openssh, privateDSA_openssh
|
||||
|
||||
try:
|
||||
from twisted.conch.test.test_ssh import ConchTestServerFactory, \
|
||||
conchTestPublicKeyChecker
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
try:
|
||||
import pyasn1
|
||||
except ImportError:
|
||||
pyasn1 = None
|
||||
|
||||
cryptography = requireModule("cryptography")
|
||||
if cryptography:
|
||||
from twisted.conch.avatar import ConchUser
|
||||
from twisted.conch.ssh.session import ISession, SSHSession, wrapProtocol
|
||||
else:
|
||||
from twisted.conch.interfaces import ISession
|
||||
|
||||
class ConchUser:
|
||||
pass
|
||||
try:
|
||||
from twisted.conch.scripts.conch import (
|
||||
SSHSession as StdioInteractingSession
|
||||
)
|
||||
except ImportError as e:
|
||||
StdioInteractingSession = None
|
||||
_reason = str(e)
|
||||
del e
|
||||
|
||||
|
||||
|
||||
def _has_ipv6():
|
||||
""" Returns True if the system can bind an IPv6 address."""
|
||||
sock = None
|
||||
has_ipv6 = False
|
||||
|
||||
try:
|
||||
sock = socket.socket(socket.AF_INET6)
|
||||
sock.bind(('::1', 0))
|
||||
has_ipv6 = True
|
||||
except socket.error:
|
||||
pass
|
||||
|
||||
if sock:
|
||||
sock.close()
|
||||
return has_ipv6
|
||||
|
||||
|
||||
HAS_IPV6 = _has_ipv6()
|
||||
|
||||
|
||||
class FakeStdio(object):
|
||||
"""
|
||||
A fake for testing L{twisted.conch.scripts.conch.SSHSession.eofReceived} and
|
||||
L{twisted.conch.scripts.cftp.SSHSession.eofReceived}.
|
||||
|
||||
@ivar writeConnLost: A flag which records whether L{loserWriteConnection}
|
||||
has been called.
|
||||
"""
|
||||
writeConnLost = False
|
||||
|
||||
def loseWriteConnection(self):
|
||||
"""
|
||||
Record the call to loseWriteConnection.
|
||||
"""
|
||||
self.writeConnLost = True
|
||||
|
||||
|
||||
|
||||
class StdioInteractingSessionTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for L{twisted.conch.scripts.conch.SSHSession}.
|
||||
"""
|
||||
if StdioInteractingSession is None:
|
||||
skip = _reason
|
||||
|
||||
|
||||
def test_eofReceived(self):
|
||||
"""
|
||||
L{twisted.conch.scripts.conch.SSHSession.eofReceived} loses the
|
||||
write half of its stdio connection.
|
||||
"""
|
||||
stdio = FakeStdio()
|
||||
channel = StdioInteractingSession()
|
||||
channel.stdio = stdio
|
||||
channel.eofReceived()
|
||||
self.assertTrue(stdio.writeConnLost)
|
||||
|
||||
|
||||
|
||||
class Echo(protocol.Protocol):
|
||||
def connectionMade(self):
|
||||
log.msg('ECHO CONNECTION MADE')
|
||||
|
||||
|
||||
def connectionLost(self, reason):
|
||||
log.msg('ECHO CONNECTION DONE')
|
||||
|
||||
|
||||
def dataReceived(self, data):
|
||||
self.transport.write(data)
|
||||
if b'\n' in data:
|
||||
self.transport.loseConnection()
|
||||
|
||||
|
||||
|
||||
class EchoFactory(protocol.Factory):
|
||||
protocol = Echo
|
||||
|
||||
|
||||
|
||||
class ConchTestOpenSSHProcess(protocol.ProcessProtocol):
|
||||
"""
|
||||
Test protocol for launching an OpenSSH client process.
|
||||
|
||||
@ivar deferred: Set by whatever uses this object. Accessed using
|
||||
L{_getDeferred}, which destroys the value so the Deferred is not
|
||||
fired twice. Fires when the process is terminated.
|
||||
"""
|
||||
|
||||
deferred = None
|
||||
buf = b''
|
||||
|
||||
def _getDeferred(self):
|
||||
d, self.deferred = self.deferred, None
|
||||
return d
|
||||
|
||||
|
||||
def outReceived(self, data):
|
||||
self.buf += data
|
||||
|
||||
|
||||
def processEnded(self, reason):
|
||||
"""
|
||||
Called when the process has ended.
|
||||
|
||||
@param reason: a Failure giving the reason for the process' end.
|
||||
"""
|
||||
if reason.value.exitCode != 0:
|
||||
self._getDeferred().errback(
|
||||
ConchError("exit code was not 0: {}".format(
|
||||
reason.value.exitCode)))
|
||||
else:
|
||||
buf = self.buf.replace(b'\r\n', b'\n')
|
||||
self._getDeferred().callback(buf)
|
||||
|
||||
|
||||
|
||||
class ConchTestForwardingProcess(protocol.ProcessProtocol):
|
||||
"""
|
||||
Manages a third-party process which launches a server.
|
||||
|
||||
Uses L{ConchTestForwardingPort} to connect to the third-party server.
|
||||
Once L{ConchTestForwardingPort} has disconnected, kill the process and fire
|
||||
a Deferred with the data received by the L{ConchTestForwardingPort}.
|
||||
|
||||
@ivar deferred: Set by whatever uses this object. Accessed using
|
||||
L{_getDeferred}, which destroys the value so the Deferred is not
|
||||
fired twice. Fires when the process is terminated.
|
||||
"""
|
||||
|
||||
deferred = None
|
||||
|
||||
def __init__(self, port, data):
|
||||
"""
|
||||
@type port: L{int}
|
||||
@param port: The port on which the third-party server is listening.
|
||||
(it is assumed that the server is running on localhost).
|
||||
|
||||
@type data: L{str}
|
||||
@param data: This is sent to the third-party server. Must end with '\n'
|
||||
in order to trigger a disconnect.
|
||||
"""
|
||||
self.port = port
|
||||
self.buffer = None
|
||||
self.data = data
|
||||
|
||||
|
||||
def _getDeferred(self):
|
||||
d, self.deferred = self.deferred, None
|
||||
return d
|
||||
|
||||
|
||||
def connectionMade(self):
|
||||
self._connect()
|
||||
|
||||
|
||||
def _connect(self):
|
||||
"""
|
||||
Connect to the server, which is often a third-party process.
|
||||
Tries to reconnect if it fails because we have no way of determining
|
||||
exactly when the port becomes available for listening -- we can only
|
||||
know when the process starts.
|
||||
"""
|
||||
cc = protocol.ClientCreator(reactor, ConchTestForwardingPort, self,
|
||||
self.data)
|
||||
d = cc.connectTCP('127.0.0.1', self.port)
|
||||
d.addErrback(self._ebConnect)
|
||||
return d
|
||||
|
||||
|
||||
def _ebConnect(self, f):
|
||||
reactor.callLater(.1, self._connect)
|
||||
|
||||
|
||||
def forwardingPortDisconnected(self, buffer):
|
||||
"""
|
||||
The network connection has died; save the buffer of output
|
||||
from the network and attempt to quit the process gracefully,
|
||||
and then (after the reactor has spun) send it a KILL signal.
|
||||
"""
|
||||
self.buffer = buffer
|
||||
self.transport.write(b'\x03')
|
||||
self.transport.loseConnection()
|
||||
reactor.callLater(0, self._reallyDie)
|
||||
|
||||
|
||||
def _reallyDie(self):
|
||||
try:
|
||||
self.transport.signalProcess('KILL')
|
||||
except ProcessExitedAlready:
|
||||
pass
|
||||
|
||||
|
||||
def processEnded(self, reason):
|
||||
"""
|
||||
Fire the Deferred at self.deferred with the data collected
|
||||
from the L{ConchTestForwardingPort} connection, if any.
|
||||
"""
|
||||
self._getDeferred().callback(self.buffer)
|
||||
|
||||
|
||||
|
||||
class ConchTestForwardingPort(protocol.Protocol):
|
||||
"""
|
||||
Connects to server launched by a third-party process (managed by
|
||||
L{ConchTestForwardingProcess}) sends data, then reports whatever it
|
||||
received back to the L{ConchTestForwardingProcess} once the connection
|
||||
is ended.
|
||||
"""
|
||||
|
||||
def __init__(self, protocol, data):
|
||||
"""
|
||||
@type protocol: L{ConchTestForwardingProcess}
|
||||
@param protocol: The L{ProcessProtocol} which made this connection.
|
||||
|
||||
@type data: str
|
||||
@param data: The data to be sent to the third-party server.
|
||||
"""
|
||||
self.protocol = protocol
|
||||
self.data = data
|
||||
|
||||
|
||||
def connectionMade(self):
|
||||
self.buffer = b''
|
||||
self.transport.write(self.data)
|
||||
|
||||
|
||||
def dataReceived(self, data):
|
||||
self.buffer += data
|
||||
|
||||
|
||||
def connectionLost(self, reason):
|
||||
self.protocol.forwardingPortDisconnected(self.buffer)
|
||||
|
||||
|
||||
|
||||
def _makeArgs(args, mod="conch"):
|
||||
start = [sys.executable, '-c'
|
||||
"""
|
||||
### Twisted Preamble
|
||||
import sys, os
|
||||
path = os.path.abspath(sys.argv[0])
|
||||
while os.path.dirname(path) != path:
|
||||
if os.path.basename(path).startswith('Twisted'):
|
||||
sys.path.insert(0, path)
|
||||
break
|
||||
path = os.path.dirname(path)
|
||||
|
||||
from twisted.conch.scripts.%s import run
|
||||
run()""" % mod]
|
||||
madeArgs = []
|
||||
for arg in start + list(args):
|
||||
if isinstance(arg, unicode):
|
||||
arg = arg.encode("utf-8")
|
||||
madeArgs.append(arg)
|
||||
return madeArgs
|
||||
|
||||
|
||||
|
||||
class ConchServerSetupMixin:
|
||||
if not cryptography:
|
||||
skip = "can't run without cryptography"
|
||||
|
||||
if not pyasn1:
|
||||
skip = "Cannot run without PyASN1"
|
||||
|
||||
# FIXME: https://twistedmatrix.com/trac/ticket/8506
|
||||
|
||||
# This should be un-skipped on Travis after the ticket is fixed. For now
|
||||
# is enabled so that we can continue with fixing other stuff using Travis.
|
||||
if _PYPY:
|
||||
skip = 'PyPy known_host not working yet on Travis.'
|
||||
|
||||
realmFactory = staticmethod(lambda: ConchTestRealm(b'testuser'))
|
||||
|
||||
def _createFiles(self):
|
||||
for f in ['rsa_test','rsa_test.pub','dsa_test','dsa_test.pub',
|
||||
'kh_test']:
|
||||
if os.path.exists(f):
|
||||
os.remove(f)
|
||||
with open('rsa_test','wb') as f:
|
||||
f.write(privateRSA_openssh)
|
||||
with open('rsa_test.pub','wb') as f:
|
||||
f.write(publicRSA_openssh)
|
||||
with open('dsa_test.pub','wb') as f:
|
||||
f.write(publicDSA_openssh)
|
||||
with open('dsa_test','wb') as f:
|
||||
f.write(privateDSA_openssh)
|
||||
os.chmod('dsa_test', 33152)
|
||||
os.chmod('rsa_test', 33152)
|
||||
with open('kh_test','wb') as f:
|
||||
f.write(b'127.0.0.1 '+publicRSA_openssh)
|
||||
|
||||
|
||||
def _getFreePort(self):
|
||||
s = socket.socket()
|
||||
s.bind(('', 0))
|
||||
port = s.getsockname()[1]
|
||||
s.close()
|
||||
return port
|
||||
|
||||
|
||||
def _makeConchFactory(self):
|
||||
"""
|
||||
Make a L{ConchTestServerFactory}, which allows us to start a
|
||||
L{ConchTestServer} -- i.e. an actually listening conch.
|
||||
"""
|
||||
realm = self.realmFactory()
|
||||
p = portal.Portal(realm)
|
||||
p.registerChecker(conchTestPublicKeyChecker())
|
||||
factory = ConchTestServerFactory()
|
||||
factory.portal = p
|
||||
return factory
|
||||
|
||||
|
||||
def setUp(self):
|
||||
self._createFiles()
|
||||
self.conchFactory = self._makeConchFactory()
|
||||
self.conchFactory.expectedLoseConnection = 1
|
||||
self.conchServer = reactor.listenTCP(0, self.conchFactory,
|
||||
interface="127.0.0.1")
|
||||
self.echoServer = reactor.listenTCP(0, EchoFactory())
|
||||
self.echoPort = self.echoServer.getHost().port
|
||||
if HAS_IPV6:
|
||||
self.echoServerV6 = reactor.listenTCP(0, EchoFactory(), interface="::1")
|
||||
self.echoPortV6 = self.echoServerV6.getHost().port
|
||||
|
||||
|
||||
def tearDown(self):
|
||||
try:
|
||||
self.conchFactory.proto.done = 1
|
||||
except AttributeError:
|
||||
pass
|
||||
else:
|
||||
self.conchFactory.proto.transport.loseConnection()
|
||||
deferreds = [
|
||||
defer.maybeDeferred(self.conchServer.stopListening),
|
||||
defer.maybeDeferred(self.echoServer.stopListening),
|
||||
]
|
||||
if HAS_IPV6:
|
||||
deferreds.append(defer.maybeDeferred(self.echoServerV6.stopListening))
|
||||
return defer.gatherResults(deferreds)
|
||||
|
||||
|
||||
|
||||
class ForwardingMixin(ConchServerSetupMixin):
|
||||
"""
|
||||
Template class for tests of the Conch server's ability to forward arbitrary
|
||||
protocols over SSH.
|
||||
|
||||
These tests are integration tests, not unit tests. They launch a Conch
|
||||
server, a custom TCP server (just an L{EchoProtocol}) and then call
|
||||
L{execute}.
|
||||
|
||||
L{execute} is implemented by subclasses of L{ForwardingMixin}. It should
|
||||
cause an SSH client to connect to the Conch server, asking it to forward
|
||||
data to the custom TCP server.
|
||||
"""
|
||||
|
||||
def test_exec(self):
|
||||
"""
|
||||
Test that we can use whatever client to send the command "echo goodbye"
|
||||
to the Conch server. Make sure we receive "goodbye" back from the
|
||||
server.
|
||||
"""
|
||||
d = self.execute('echo goodbye', ConchTestOpenSSHProcess())
|
||||
return d.addCallback(self.assertEqual, b'goodbye\n')
|
||||
|
||||
|
||||
def test_localToRemoteForwarding(self):
|
||||
"""
|
||||
Test that we can use whatever client to forward a local port to a
|
||||
specified port on the server.
|
||||
"""
|
||||
localPort = self._getFreePort()
|
||||
process = ConchTestForwardingProcess(localPort, b'test\n')
|
||||
d = self.execute('', process,
|
||||
sshArgs='-N -L%i:127.0.0.1:%i'
|
||||
% (localPort, self.echoPort))
|
||||
d.addCallback(self.assertEqual, b'test\n')
|
||||
return d
|
||||
|
||||
|
||||
def test_remoteToLocalForwarding(self):
|
||||
"""
|
||||
Test that we can use whatever client to forward a port from the server
|
||||
to a port locally.
|
||||
"""
|
||||
localPort = self._getFreePort()
|
||||
process = ConchTestForwardingProcess(localPort, b'test\n')
|
||||
d = self.execute('', process,
|
||||
sshArgs='-N -R %i:127.0.0.1:%i'
|
||||
% (localPort, self.echoPort))
|
||||
d.addCallback(self.assertEqual, b'test\n')
|
||||
return d
|
||||
|
||||
|
||||
|
||||
# Conventionally there is a separate adapter object which provides ISession for
|
||||
# the user, but making the user provide ISession directly works too. This isn't
|
||||
# a full implementation of ISession though, just enough to make these tests
|
||||
# pass.
|
||||
@implementer(ISession)
|
||||
class RekeyAvatar(ConchUser):
|
||||
"""
|
||||
This avatar implements a shell which sends 60 numbered lines to whatever
|
||||
connects to it, then closes the session with a 0 exit status.
|
||||
|
||||
60 lines is selected as being enough to send more than 2kB of traffic, the
|
||||
amount the client is configured to initiate a rekey after.
|
||||
"""
|
||||
def __init__(self):
|
||||
ConchUser.__init__(self)
|
||||
self.channelLookup[b'session'] = SSHSession
|
||||
|
||||
|
||||
def openShell(self, transport):
|
||||
"""
|
||||
Write 60 lines of data to the transport, then exit.
|
||||
"""
|
||||
proto = protocol.Protocol()
|
||||
proto.makeConnection(transport)
|
||||
transport.makeConnection(wrapProtocol(proto))
|
||||
|
||||
# Send enough bytes to the connection so that a rekey is triggered in
|
||||
# the client.
|
||||
def write(counter):
|
||||
i = next(counter)
|
||||
if i == 60:
|
||||
call.stop()
|
||||
transport.session.conn.sendRequest(
|
||||
transport.session, b'exit-status', b'\x00\x00\x00\x00')
|
||||
transport.loseConnection()
|
||||
else:
|
||||
line = "line #%02d\n" % (i,)
|
||||
line = line.encode("utf-8")
|
||||
transport.write(line)
|
||||
|
||||
# The timing for this loop is an educated guess (and/or the result of
|
||||
# experimentation) to exercise the case where a packet is generated
|
||||
# mid-rekey. Since the other side of the connection is (so far) the
|
||||
# OpenSSH command line client, there's no easy way to determine when the
|
||||
# rekey has been initiated. If there were, then generating a packet
|
||||
# immediately at that time would be a better way to test the
|
||||
# functionality being tested here.
|
||||
call = LoopingCall(write, count())
|
||||
call.start(0.01)
|
||||
|
||||
|
||||
def closed(self):
|
||||
"""
|
||||
Ignore the close of the session.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class RekeyRealm:
|
||||
"""
|
||||
This realm gives out new L{RekeyAvatar} instances for any avatar request.
|
||||
"""
|
||||
def requestAvatar(self, avatarID, mind, *interfaces):
|
||||
return interfaces[0], RekeyAvatar(), lambda: None
|
||||
|
||||
|
||||
|
||||
class RekeyTestsMixin(ConchServerSetupMixin):
|
||||
"""
|
||||
TestCase mixin which defines tests exercising L{SSHTransportBase}'s handling
|
||||
of rekeying messages.
|
||||
"""
|
||||
realmFactory = RekeyRealm
|
||||
|
||||
def test_clientRekey(self):
|
||||
"""
|
||||
After a client-initiated rekey is completed, application data continues
|
||||
to be passed over the SSH connection.
|
||||
"""
|
||||
process = ConchTestOpenSSHProcess()
|
||||
d = self.execute("", process, '-o RekeyLimit=2K')
|
||||
def finished(result):
|
||||
expectedResult = '\n'.join(['line #%02d' % (i,) for i in range(60)]) + '\n'
|
||||
expectedResult = expectedResult.encode("utf-8")
|
||||
self.assertEqual(result, expectedResult)
|
||||
d.addCallback(finished)
|
||||
return d
|
||||
|
||||
|
||||
|
||||
class OpenSSHClientMixin:
|
||||
if not which('ssh'):
|
||||
skip = "no ssh command-line client available"
|
||||
|
||||
|
||||
def execute(self, remoteCommand, process, sshArgs=''):
|
||||
"""
|
||||
Connects to the SSH server started in L{ConchServerSetupMixin.setUp} by
|
||||
running the 'ssh' command line tool.
|
||||
|
||||
@type remoteCommand: str
|
||||
@param remoteCommand: The command (with arguments) to run on the
|
||||
remote end.
|
||||
|
||||
@type process: L{ConchTestOpenSSHProcess}
|
||||
|
||||
@type sshArgs: str
|
||||
@param sshArgs: Arguments to pass to the 'ssh' process.
|
||||
|
||||
@return: L{defer.Deferred}
|
||||
"""
|
||||
# PubkeyAcceptedKeyTypes does not exist prior to OpenSSH 7.0 so we
|
||||
# first need to check if we can set it. If we can, -V will just print
|
||||
# the version without doing anything else; if we can't, we will get a
|
||||
# configuration error.
|
||||
d = getProcessValue(
|
||||
which('ssh')[0], ('-o', 'PubkeyAcceptedKeyTypes=ssh-dss', '-V'))
|
||||
def hasPAKT(status):
|
||||
if status == 0:
|
||||
opts = '-oPubkeyAcceptedKeyTypes=ssh-dss '
|
||||
else:
|
||||
opts = ''
|
||||
|
||||
process.deferred = defer.Deferred()
|
||||
# Pass -F /dev/null to avoid the user's configuration file from
|
||||
# being loaded, as it may contain settings that cause our tests to
|
||||
# fail or hang.
|
||||
cmdline = ('ssh -2 -l testuser -p %i '
|
||||
'-F /dev/null '
|
||||
'-oUserKnownHostsFile=kh_test '
|
||||
'-oPasswordAuthentication=no '
|
||||
# Always use the RSA key, since that's the one in kh_test.
|
||||
'-oHostKeyAlgorithms=ssh-rsa '
|
||||
'-a '
|
||||
'-i dsa_test ') + opts + sshArgs + \
|
||||
' 127.0.0.1 ' + remoteCommand
|
||||
port = self.conchServer.getHost().port
|
||||
cmds = (cmdline % port).split()
|
||||
encodedCmds = []
|
||||
for cmd in cmds:
|
||||
if isinstance(cmd, unicode):
|
||||
cmd = cmd.encode("utf-8")
|
||||
encodedCmds.append(cmd)
|
||||
reactor.spawnProcess(process, which('ssh')[0], encodedCmds)
|
||||
return process.deferred
|
||||
return d.addCallback(hasPAKT)
|
||||
|
||||
|
||||
|
||||
class OpenSSHKeyExchangeTests(ConchServerSetupMixin, OpenSSHClientMixin,
|
||||
unittest.TestCase):
|
||||
"""
|
||||
Tests L{SSHTransportBase}'s key exchange algorithm compatibility with
|
||||
OpenSSH.
|
||||
"""
|
||||
|
||||
def assertExecuteWithKexAlgorithm(self, keyExchangeAlgo):
|
||||
"""
|
||||
Call execute() method of L{OpenSSHClientMixin} with an ssh option that
|
||||
forces the exclusive use of the key exchange algorithm specified by
|
||||
keyExchangeAlgo
|
||||
|
||||
@type keyExchangeAlgo: L{str}
|
||||
@param keyExchangeAlgo: The key exchange algorithm to use
|
||||
|
||||
@return: L{defer.Deferred}
|
||||
"""
|
||||
kexAlgorithms = []
|
||||
try:
|
||||
output = subprocess.check_output([which('ssh')[0], '-Q', 'kex'],
|
||||
stderr=subprocess.STDOUT)
|
||||
if not isinstance(output, str):
|
||||
output = output.decode("utf-8")
|
||||
kexAlgorithms = output.split()
|
||||
except:
|
||||
pass
|
||||
|
||||
if keyExchangeAlgo not in kexAlgorithms:
|
||||
raise unittest.SkipTest(
|
||||
"{} not supported by ssh client".format(
|
||||
keyExchangeAlgo))
|
||||
|
||||
d = self.execute('echo hello', ConchTestOpenSSHProcess(),
|
||||
'-oKexAlgorithms=' + keyExchangeAlgo)
|
||||
return d.addCallback(self.assertEqual, b'hello\n')
|
||||
|
||||
|
||||
def test_ECDHSHA256(self):
|
||||
"""
|
||||
The ecdh-sha2-nistp256 key exchange algorithm is compatible with
|
||||
OpenSSH
|
||||
"""
|
||||
return self.assertExecuteWithKexAlgorithm(
|
||||
'ecdh-sha2-nistp256')
|
||||
|
||||
|
||||
def test_ECDHSHA384(self):
|
||||
"""
|
||||
The ecdh-sha2-nistp384 key exchange algorithm is compatible with
|
||||
OpenSSH
|
||||
"""
|
||||
return self.assertExecuteWithKexAlgorithm(
|
||||
'ecdh-sha2-nistp384')
|
||||
|
||||
|
||||
def test_ECDHSHA521(self):
|
||||
"""
|
||||
The ecdh-sha2-nistp521 key exchange algorithm is compatible with
|
||||
OpenSSH
|
||||
"""
|
||||
return self.assertExecuteWithKexAlgorithm(
|
||||
'ecdh-sha2-nistp521')
|
||||
|
||||
|
||||
def test_DH_GROUP14(self):
|
||||
"""
|
||||
The diffie-hellman-group14-sha1 key exchange algorithm is compatible
|
||||
with OpenSSH.
|
||||
"""
|
||||
return self.assertExecuteWithKexAlgorithm(
|
||||
'diffie-hellman-group14-sha1')
|
||||
|
||||
|
||||
def test_DH_GROUP_EXCHANGE_SHA1(self):
|
||||
"""
|
||||
The diffie-hellman-group-exchange-sha1 key exchange algorithm is
|
||||
compatible with OpenSSH.
|
||||
"""
|
||||
return self.assertExecuteWithKexAlgorithm(
|
||||
'diffie-hellman-group-exchange-sha1')
|
||||
|
||||
|
||||
def test_DH_GROUP_EXCHANGE_SHA256(self):
|
||||
"""
|
||||
The diffie-hellman-group-exchange-sha256 key exchange algorithm is
|
||||
compatible with OpenSSH.
|
||||
"""
|
||||
return self.assertExecuteWithKexAlgorithm(
|
||||
'diffie-hellman-group-exchange-sha256')
|
||||
|
||||
|
||||
def test_unsupported_algorithm(self):
|
||||
"""
|
||||
The list of key exchange algorithms supported
|
||||
by OpenSSH client is obtained with C{ssh -Q kex}.
|
||||
"""
|
||||
self.assertRaises(unittest.SkipTest,
|
||||
self.assertExecuteWithKexAlgorithm,
|
||||
'unsupported-algorithm')
|
||||
|
||||
|
||||
|
||||
class OpenSSHClientForwardingTests(ForwardingMixin, OpenSSHClientMixin,
|
||||
unittest.TestCase):
|
||||
"""
|
||||
Connection forwarding tests run against the OpenSSL command line client.
|
||||
"""
|
||||
def test_localToRemoteForwardingV6(self):
|
||||
"""
|
||||
Forwarding of arbitrary IPv6 TCP connections via SSH.
|
||||
"""
|
||||
localPort = self._getFreePort()
|
||||
process = ConchTestForwardingProcess(localPort, b'test\n')
|
||||
d = self.execute('', process,
|
||||
sshArgs='-N -L%i:[::1]:%i'
|
||||
% (localPort, self.echoPortV6))
|
||||
d.addCallback(self.assertEqual, b'test\n')
|
||||
return d
|
||||
if not HAS_IPV6:
|
||||
test_localToRemoteForwardingV6.skip = "Requires IPv6 support"
|
||||
|
||||
|
||||
|
||||
class OpenSSHClientRekeyTests(RekeyTestsMixin, OpenSSHClientMixin,
|
||||
unittest.TestCase):
|
||||
"""
|
||||
Rekeying tests run against the OpenSSL command line client.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class CmdLineClientTests(ForwardingMixin, unittest.TestCase):
|
||||
"""
|
||||
Connection forwarding tests run against the Conch command line client.
|
||||
"""
|
||||
if runtime.platformType == 'win32':
|
||||
skip = "can't run cmdline client on win32"
|
||||
|
||||
|
||||
def execute(self, remoteCommand, process, sshArgs='', conchArgs=None):
|
||||
"""
|
||||
As for L{OpenSSHClientTestCase.execute}, except it runs the 'conch'
|
||||
command line tool, not 'ssh'.
|
||||
"""
|
||||
if conchArgs is None:
|
||||
conchArgs = []
|
||||
|
||||
process.deferred = defer.Deferred()
|
||||
port = self.conchServer.getHost().port
|
||||
cmd = ('-p {} -l testuser '
|
||||
'--known-hosts kh_test '
|
||||
'--user-authentications publickey '
|
||||
'-a '
|
||||
'-i dsa_test '
|
||||
'-v '.format(port) + sshArgs +
|
||||
' 127.0.0.1 ' + remoteCommand)
|
||||
cmds = _makeArgs(conchArgs + cmd.split())
|
||||
env = os.environ.copy()
|
||||
env['PYTHONPATH'] = os.pathsep.join(sys.path)
|
||||
encodedCmds = []
|
||||
encodedEnv = {}
|
||||
for cmd in cmds:
|
||||
if isinstance(cmd, unicode):
|
||||
cmd = cmd.encode("utf-8")
|
||||
encodedCmds.append(cmd)
|
||||
for var in env:
|
||||
val = env[var]
|
||||
if isinstance(var, unicode):
|
||||
var = var.encode("utf-8")
|
||||
if isinstance(val, unicode):
|
||||
val = val.encode("utf-8")
|
||||
encodedEnv[var] = val
|
||||
reactor.spawnProcess(process, sys.executable, encodedCmds, env=encodedEnv)
|
||||
return process.deferred
|
||||
|
||||
|
||||
def test_runWithLogFile(self):
|
||||
"""
|
||||
It can store logs to a local file.
|
||||
"""
|
||||
def cb_check_log(result):
|
||||
logContent = logPath.getContent()
|
||||
self.assertIn(b'Log opened.', logContent)
|
||||
|
||||
logPath = filepath.FilePath(self.mktemp())
|
||||
|
||||
d = self.execute(
|
||||
remoteCommand='echo goodbye',
|
||||
process=ConchTestOpenSSHProcess(),
|
||||
conchArgs=['--log', '--logfile', logPath.path,
|
||||
'--host-key-algorithms', 'ssh-rsa']
|
||||
)
|
||||
|
||||
d.addCallback(self.assertEqual, b'goodbye\n')
|
||||
d.addCallback(cb_check_log)
|
||||
return d
|
||||
|
||||
|
||||
def test_runWithNoHostAlgorithmsSpecified(self):
|
||||
"""
|
||||
Do not use --host-key-algorithms flag on command line.
|
||||
"""
|
||||
d = self.execute(
|
||||
remoteCommand='echo goodbye',
|
||||
process=ConchTestOpenSSHProcess()
|
||||
)
|
||||
|
||||
d.addCallback(self.assertEqual, b'goodbye\n')
|
||||
return d
|
||||
@@ -0,0 +1,761 @@
|
||||
# Copyright (c) 2007-2010 Twisted Matrix Laboratories.
|
||||
# See LICENSE for details
|
||||
|
||||
"""
|
||||
This module tests twisted.conch.ssh.connection.
|
||||
"""
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
import struct
|
||||
|
||||
from twisted.python.reflect import requireModule
|
||||
|
||||
cryptography = requireModule("cryptography")
|
||||
|
||||
from twisted.conch import error
|
||||
if cryptography:
|
||||
from twisted.conch.ssh import common, connection
|
||||
else:
|
||||
class connection:
|
||||
class SSHConnection: pass
|
||||
|
||||
from twisted.conch.ssh import channel
|
||||
from twisted.python.compat import long
|
||||
from twisted.trial import unittest
|
||||
from twisted.conch.test import test_userauth
|
||||
|
||||
|
||||
class TestChannel(channel.SSHChannel):
|
||||
"""
|
||||
A mocked-up version of twisted.conch.ssh.channel.SSHChannel.
|
||||
|
||||
@ivar gotOpen: True if channelOpen has been called.
|
||||
@type gotOpen: L{bool}
|
||||
@ivar specificData: the specific channel open data passed to channelOpen.
|
||||
@type specificData: L{bytes}
|
||||
@ivar openFailureReason: the reason passed to openFailed.
|
||||
@type openFailed: C{error.ConchError}
|
||||
@ivar inBuffer: a C{list} of strings received by the channel.
|
||||
@type inBuffer: C{list}
|
||||
@ivar extBuffer: a C{list} of 2-tuples (type, extended data) of received by
|
||||
the channel.
|
||||
@type extBuffer: C{list}
|
||||
@ivar numberRequests: the number of requests that have been made to this
|
||||
channel.
|
||||
@type numberRequests: L{int}
|
||||
@ivar gotEOF: True if the other side sent EOF.
|
||||
@type gotEOF: L{bool}
|
||||
@ivar gotOneClose: True if the other side closed the connection.
|
||||
@type gotOneClose: L{bool}
|
||||
@ivar gotClosed: True if the channel is closed.
|
||||
@type gotClosed: L{bool}
|
||||
"""
|
||||
name = b"TestChannel"
|
||||
gotOpen = False
|
||||
gotClosed = False
|
||||
|
||||
def logPrefix(self):
|
||||
return "TestChannel %i" % self.id
|
||||
|
||||
def channelOpen(self, specificData):
|
||||
"""
|
||||
The channel is open. Set up the instance variables.
|
||||
"""
|
||||
self.gotOpen = True
|
||||
self.specificData = specificData
|
||||
self.inBuffer = []
|
||||
self.extBuffer = []
|
||||
self.numberRequests = 0
|
||||
self.gotEOF = False
|
||||
self.gotOneClose = False
|
||||
self.gotClosed = False
|
||||
|
||||
def openFailed(self, reason):
|
||||
"""
|
||||
Opening the channel failed. Store the reason why.
|
||||
"""
|
||||
self.openFailureReason = reason
|
||||
|
||||
def request_test(self, data):
|
||||
"""
|
||||
A test request. Return True if data is 'data'.
|
||||
|
||||
@type data: L{bytes}
|
||||
"""
|
||||
self.numberRequests += 1
|
||||
return data == b'data'
|
||||
|
||||
def dataReceived(self, data):
|
||||
"""
|
||||
Data was received. Store it in the buffer.
|
||||
"""
|
||||
self.inBuffer.append(data)
|
||||
|
||||
def extReceived(self, code, data):
|
||||
"""
|
||||
Extended data was received. Store it in the buffer.
|
||||
"""
|
||||
self.extBuffer.append((code, data))
|
||||
|
||||
def eofReceived(self):
|
||||
"""
|
||||
EOF was received. Remember it.
|
||||
"""
|
||||
self.gotEOF = True
|
||||
|
||||
def closeReceived(self):
|
||||
"""
|
||||
Close was received. Remember it.
|
||||
"""
|
||||
self.gotOneClose = True
|
||||
|
||||
def closed(self):
|
||||
"""
|
||||
The channel is closed. Rembember it.
|
||||
"""
|
||||
self.gotClosed = True
|
||||
|
||||
class TestAvatar:
|
||||
"""
|
||||
A mocked-up version of twisted.conch.avatar.ConchUser
|
||||
"""
|
||||
_ARGS_ERROR_CODE = 123
|
||||
|
||||
def lookupChannel(self, channelType, windowSize, maxPacket, data):
|
||||
"""
|
||||
The server wants us to return a channel. If the requested channel is
|
||||
our TestChannel, return it, otherwise return None.
|
||||
"""
|
||||
if channelType == TestChannel.name:
|
||||
return TestChannel(remoteWindow=windowSize,
|
||||
remoteMaxPacket=maxPacket,
|
||||
data=data, avatar=self)
|
||||
elif channelType == b"conch-error-args":
|
||||
# Raise a ConchError with backwards arguments to make sure the
|
||||
# connection fixes it for us. This case should be deprecated and
|
||||
# deleted eventually, but only after all of Conch gets the argument
|
||||
# order right.
|
||||
raise error.ConchError(
|
||||
self._ARGS_ERROR_CODE, "error args in wrong order")
|
||||
|
||||
|
||||
def gotGlobalRequest(self, requestType, data):
|
||||
"""
|
||||
The client has made a global request. If the global request is
|
||||
'TestGlobal', return True. If the global request is 'TestData',
|
||||
return True and the request-specific data we received. Otherwise,
|
||||
return False.
|
||||
"""
|
||||
if requestType == b'TestGlobal':
|
||||
return True
|
||||
elif requestType == b'TestData':
|
||||
return True, data
|
||||
else:
|
||||
return False
|
||||
|
||||
|
||||
|
||||
class TestConnection(connection.SSHConnection):
|
||||
"""
|
||||
A subclass of SSHConnection for testing.
|
||||
|
||||
@ivar channel: the current channel.
|
||||
@type channel. C{TestChannel}
|
||||
"""
|
||||
|
||||
if not cryptography:
|
||||
skip = "Cannot run without cryptography"
|
||||
|
||||
def logPrefix(self):
|
||||
return "TestConnection"
|
||||
|
||||
def global_TestGlobal(self, data):
|
||||
"""
|
||||
The other side made the 'TestGlobal' global request. Return True.
|
||||
"""
|
||||
return True
|
||||
|
||||
def global_Test_Data(self, data):
|
||||
"""
|
||||
The other side made the 'Test-Data' global request. Return True and
|
||||
the data we received.
|
||||
"""
|
||||
return True, data
|
||||
|
||||
def channel_TestChannel(self, windowSize, maxPacket, data):
|
||||
"""
|
||||
The other side is requesting the TestChannel. Create a C{TestChannel}
|
||||
instance, store it, and return it.
|
||||
"""
|
||||
self.channel = TestChannel(remoteWindow=windowSize,
|
||||
remoteMaxPacket=maxPacket, data=data)
|
||||
return self.channel
|
||||
|
||||
def channel_ErrorChannel(self, windowSize, maxPacket, data):
|
||||
"""
|
||||
The other side is requesting the ErrorChannel. Raise an exception.
|
||||
"""
|
||||
raise AssertionError('no such thing')
|
||||
|
||||
|
||||
|
||||
class ConnectionTests(unittest.TestCase):
|
||||
|
||||
if not cryptography:
|
||||
skip = "Cannot run without cryptography"
|
||||
if test_userauth.transport is None:
|
||||
skip = "Cannot run without both cryptography and pyasn1"
|
||||
|
||||
def setUp(self):
|
||||
self.transport = test_userauth.FakeTransport(None)
|
||||
self.transport.avatar = TestAvatar()
|
||||
self.conn = TestConnection()
|
||||
self.conn.transport = self.transport
|
||||
self.conn.serviceStarted()
|
||||
|
||||
def _openChannel(self, channel):
|
||||
"""
|
||||
Open the channel with the default connection.
|
||||
"""
|
||||
self.conn.openChannel(channel)
|
||||
self.transport.packets = self.transport.packets[:-1]
|
||||
self.conn.ssh_CHANNEL_OPEN_CONFIRMATION(struct.pack('>2L',
|
||||
channel.id, 255) + b'\x00\x02\x00\x00\x00\x00\x80\x00')
|
||||
|
||||
def tearDown(self):
|
||||
self.conn.serviceStopped()
|
||||
|
||||
def test_linkAvatar(self):
|
||||
"""
|
||||
Test that the connection links itself to the avatar in the
|
||||
transport.
|
||||
"""
|
||||
self.assertIs(self.transport.avatar.conn, self.conn)
|
||||
|
||||
def test_serviceStopped(self):
|
||||
"""
|
||||
Test that serviceStopped() closes any open channels.
|
||||
"""
|
||||
channel1 = TestChannel()
|
||||
channel2 = TestChannel()
|
||||
self.conn.openChannel(channel1)
|
||||
self.conn.openChannel(channel2)
|
||||
self.conn.ssh_CHANNEL_OPEN_CONFIRMATION(b'\x00\x00\x00\x00' * 4)
|
||||
self.assertTrue(channel1.gotOpen)
|
||||
self.assertFalse(channel1.gotClosed)
|
||||
self.assertFalse(channel2.gotOpen)
|
||||
self.assertFalse(channel2.gotClosed)
|
||||
self.conn.serviceStopped()
|
||||
self.assertTrue(channel1.gotClosed)
|
||||
self.assertFalse(channel2.gotOpen)
|
||||
self.assertFalse(channel2.gotClosed)
|
||||
from twisted.internet.error import ConnectionLost
|
||||
self.assertIsInstance(channel2.openFailureReason,
|
||||
ConnectionLost)
|
||||
|
||||
def test_GLOBAL_REQUEST(self):
|
||||
"""
|
||||
Test that global request packets are dispatched to the global_*
|
||||
methods and the return values are translated into success or failure
|
||||
messages.
|
||||
"""
|
||||
self.conn.ssh_GLOBAL_REQUEST(common.NS(b'TestGlobal') + b'\xff')
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_REQUEST_SUCCESS, b'')])
|
||||
self.transport.packets = []
|
||||
self.conn.ssh_GLOBAL_REQUEST(common.NS(b'TestData') + b'\xff' +
|
||||
b'test data')
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_REQUEST_SUCCESS, b'test data')])
|
||||
self.transport.packets = []
|
||||
self.conn.ssh_GLOBAL_REQUEST(common.NS(b'TestBad') + b'\xff')
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_REQUEST_FAILURE, b'')])
|
||||
self.transport.packets = []
|
||||
self.conn.ssh_GLOBAL_REQUEST(common.NS(b'TestGlobal') + b'\x00')
|
||||
self.assertEqual(self.transport.packets, [])
|
||||
|
||||
def test_REQUEST_SUCCESS(self):
|
||||
"""
|
||||
Test that global request success packets cause the Deferred to be
|
||||
called back.
|
||||
"""
|
||||
d = self.conn.sendGlobalRequest(b'request', b'data', True)
|
||||
self.conn.ssh_REQUEST_SUCCESS(b'data')
|
||||
def check(data):
|
||||
self.assertEqual(data, b'data')
|
||||
d.addCallback(check)
|
||||
d.addErrback(self.fail)
|
||||
return d
|
||||
|
||||
def test_REQUEST_FAILURE(self):
|
||||
"""
|
||||
Test that global request failure packets cause the Deferred to be
|
||||
erred back.
|
||||
"""
|
||||
d = self.conn.sendGlobalRequest(b'request', b'data', True)
|
||||
self.conn.ssh_REQUEST_FAILURE(b'data')
|
||||
def check(f):
|
||||
self.assertEqual(f.value.data, b'data')
|
||||
d.addCallback(self.fail)
|
||||
d.addErrback(check)
|
||||
return d
|
||||
|
||||
def test_CHANNEL_OPEN(self):
|
||||
"""
|
||||
Test that open channel packets cause a channel to be created and
|
||||
opened or a failure message to be returned.
|
||||
"""
|
||||
del self.transport.avatar
|
||||
self.conn.ssh_CHANNEL_OPEN(common.NS(b'TestChannel') +
|
||||
b'\x00\x00\x00\x01' * 4)
|
||||
self.assertTrue(self.conn.channel.gotOpen)
|
||||
self.assertEqual(self.conn.channel.conn, self.conn)
|
||||
self.assertEqual(self.conn.channel.data, b'\x00\x00\x00\x01')
|
||||
self.assertEqual(self.conn.channel.specificData, b'\x00\x00\x00\x01')
|
||||
self.assertEqual(self.conn.channel.remoteWindowLeft, 1)
|
||||
self.assertEqual(self.conn.channel.remoteMaxPacket, 1)
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_OPEN_CONFIRMATION,
|
||||
b'\x00\x00\x00\x01\x00\x00\x00\x00\x00\x02\x00\x00'
|
||||
b'\x00\x00\x80\x00')])
|
||||
self.transport.packets = []
|
||||
self.conn.ssh_CHANNEL_OPEN(common.NS(b'BadChannel') +
|
||||
b'\x00\x00\x00\x02' * 4)
|
||||
self.flushLoggedErrors()
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_OPEN_FAILURE,
|
||||
b'\x00\x00\x00\x02\x00\x00\x00\x03' + common.NS(
|
||||
b'unknown channel') + common.NS(b''))])
|
||||
self.transport.packets = []
|
||||
self.conn.ssh_CHANNEL_OPEN(common.NS(b'ErrorChannel') +
|
||||
b'\x00\x00\x00\x02' * 4)
|
||||
self.flushLoggedErrors()
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_OPEN_FAILURE,
|
||||
b'\x00\x00\x00\x02\x00\x00\x00\x02' + common.NS(
|
||||
b'unknown failure') + common.NS(b''))])
|
||||
|
||||
|
||||
def _lookupChannelErrorTest(self, code):
|
||||
"""
|
||||
Deliver a request for a channel open which will result in an exception
|
||||
being raised during channel lookup. Assert that an error response is
|
||||
delivered as a result.
|
||||
"""
|
||||
self.transport.avatar._ARGS_ERROR_CODE = code
|
||||
self.conn.ssh_CHANNEL_OPEN(
|
||||
common.NS(b'conch-error-args') + b'\x00\x00\x00\x01' * 4)
|
||||
errors = self.flushLoggedErrors(error.ConchError)
|
||||
self.assertEqual(
|
||||
len(errors), 1, "Expected one error, got: %r" % (errors,))
|
||||
self.assertEqual(errors[0].value.args, (long(123), "error args in wrong order"))
|
||||
self.assertEqual(
|
||||
self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_OPEN_FAILURE,
|
||||
# The response includes some bytes which identifying the
|
||||
# associated request, as well as the error code (7b in hex) and
|
||||
# the error message.
|
||||
b'\x00\x00\x00\x01\x00\x00\x00\x7b' + common.NS(
|
||||
b'error args in wrong order') + common.NS(b''))])
|
||||
|
||||
|
||||
def test_lookupChannelError(self):
|
||||
"""
|
||||
If a C{lookupChannel} implementation raises L{error.ConchError} with the
|
||||
arguments in the wrong order, a C{MSG_CHANNEL_OPEN} failure is still
|
||||
sent in response to the message.
|
||||
|
||||
This is a temporary work-around until L{error.ConchError} is given
|
||||
better attributes and all of the Conch code starts constructing
|
||||
instances of it properly. Eventually this functionality should be
|
||||
deprecated and then removed.
|
||||
"""
|
||||
self._lookupChannelErrorTest(123)
|
||||
|
||||
|
||||
def test_lookupChannelErrorLongCode(self):
|
||||
"""
|
||||
Like L{test_lookupChannelError}, but for the case where the failure code
|
||||
is represented as a L{long} instead of a L{int}.
|
||||
"""
|
||||
self._lookupChannelErrorTest(long(123))
|
||||
|
||||
|
||||
def test_CHANNEL_OPEN_CONFIRMATION(self):
|
||||
"""
|
||||
Test that channel open confirmation packets cause the channel to be
|
||||
notified that it's open.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self.conn.openChannel(channel)
|
||||
self.conn.ssh_CHANNEL_OPEN_CONFIRMATION(b'\x00\x00\x00\x00'*5)
|
||||
self.assertEqual(channel.remoteWindowLeft, 0)
|
||||
self.assertEqual(channel.remoteMaxPacket, 0)
|
||||
self.assertEqual(channel.specificData, b'\x00\x00\x00\x00')
|
||||
self.assertEqual(self.conn.channelsToRemoteChannel[channel],
|
||||
0)
|
||||
self.assertEqual(self.conn.localToRemoteChannel[0], 0)
|
||||
|
||||
def test_CHANNEL_OPEN_FAILURE(self):
|
||||
"""
|
||||
Test that channel open failure packets cause the channel to be
|
||||
notified that its opening failed.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self.conn.openChannel(channel)
|
||||
self.conn.ssh_CHANNEL_OPEN_FAILURE(b'\x00\x00\x00\x00\x00\x00\x00'
|
||||
b'\x01' + common.NS(b'failure!'))
|
||||
self.assertEqual(channel.openFailureReason.args, (b'failure!', 1))
|
||||
self.assertIsNone(self.conn.channels.get(channel))
|
||||
|
||||
|
||||
def test_CHANNEL_WINDOW_ADJUST(self):
|
||||
"""
|
||||
Test that channel window adjust messages add bytes to the channel
|
||||
window.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self._openChannel(channel)
|
||||
oldWindowSize = channel.remoteWindowLeft
|
||||
self.conn.ssh_CHANNEL_WINDOW_ADJUST(b'\x00\x00\x00\x00\x00\x00\x00'
|
||||
b'\x01')
|
||||
self.assertEqual(channel.remoteWindowLeft, oldWindowSize + 1)
|
||||
|
||||
def test_CHANNEL_DATA(self):
|
||||
"""
|
||||
Test that channel data messages are passed up to the channel, or
|
||||
cause the channel to be closed if the data is too large.
|
||||
"""
|
||||
channel = TestChannel(localWindow=6, localMaxPacket=5)
|
||||
self._openChannel(channel)
|
||||
self.conn.ssh_CHANNEL_DATA(b'\x00\x00\x00\x00' + common.NS(b'data'))
|
||||
self.assertEqual(channel.inBuffer, [b'data'])
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_WINDOW_ADJUST, b'\x00\x00\x00\xff'
|
||||
b'\x00\x00\x00\x04')])
|
||||
self.transport.packets = []
|
||||
longData = b'a' * (channel.localWindowLeft + 1)
|
||||
self.conn.ssh_CHANNEL_DATA(b'\x00\x00\x00\x00' + common.NS(longData))
|
||||
self.assertEqual(channel.inBuffer, [b'data'])
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_CLOSE, b'\x00\x00\x00\xff')])
|
||||
channel = TestChannel()
|
||||
self._openChannel(channel)
|
||||
bigData = b'a' * (channel.localMaxPacket + 1)
|
||||
self.transport.packets = []
|
||||
self.conn.ssh_CHANNEL_DATA(b'\x00\x00\x00\x01' + common.NS(bigData))
|
||||
self.assertEqual(channel.inBuffer, [])
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_CLOSE, b'\x00\x00\x00\xff')])
|
||||
|
||||
def test_CHANNEL_EXTENDED_DATA(self):
|
||||
"""
|
||||
Test that channel extended data messages are passed up to the channel,
|
||||
or cause the channel to be closed if they're too big.
|
||||
"""
|
||||
channel = TestChannel(localWindow=6, localMaxPacket=5)
|
||||
self._openChannel(channel)
|
||||
self.conn.ssh_CHANNEL_EXTENDED_DATA(b'\x00\x00\x00\x00\x00\x00\x00'
|
||||
b'\x00' + common.NS(b'data'))
|
||||
self.assertEqual(channel.extBuffer, [(0, b'data')])
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_WINDOW_ADJUST, b'\x00\x00\x00\xff'
|
||||
b'\x00\x00\x00\x04')])
|
||||
self.transport.packets = []
|
||||
longData = b'a' * (channel.localWindowLeft + 1)
|
||||
self.conn.ssh_CHANNEL_EXTENDED_DATA(b'\x00\x00\x00\x00\x00\x00\x00'
|
||||
b'\x00' + common.NS(longData))
|
||||
self.assertEqual(channel.extBuffer, [(0, b'data')])
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_CLOSE, b'\x00\x00\x00\xff')])
|
||||
channel = TestChannel()
|
||||
self._openChannel(channel)
|
||||
bigData = b'a' * (channel.localMaxPacket + 1)
|
||||
self.transport.packets = []
|
||||
self.conn.ssh_CHANNEL_EXTENDED_DATA(b'\x00\x00\x00\x01\x00\x00\x00'
|
||||
b'\x00' + common.NS(bigData))
|
||||
self.assertEqual(channel.extBuffer, [])
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_CLOSE, b'\x00\x00\x00\xff')])
|
||||
|
||||
def test_CHANNEL_EOF(self):
|
||||
"""
|
||||
Test that channel eof messages are passed up to the channel.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self._openChannel(channel)
|
||||
self.conn.ssh_CHANNEL_EOF(b'\x00\x00\x00\x00')
|
||||
self.assertTrue(channel.gotEOF)
|
||||
|
||||
def test_CHANNEL_CLOSE(self):
|
||||
"""
|
||||
Test that channel close messages are passed up to the channel. Also,
|
||||
test that channel.close() is called if both sides are closed when this
|
||||
message is received.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self._openChannel(channel)
|
||||
self.assertTrue(channel.gotOpen)
|
||||
self.assertFalse(channel.gotOneClose)
|
||||
self.assertFalse(channel.gotClosed)
|
||||
self.conn.sendClose(channel)
|
||||
self.conn.ssh_CHANNEL_CLOSE(b'\x00\x00\x00\x00')
|
||||
self.assertTrue(channel.gotOneClose)
|
||||
self.assertTrue(channel.gotClosed)
|
||||
|
||||
def test_CHANNEL_REQUEST_success(self):
|
||||
"""
|
||||
Test that channel requests that succeed send MSG_CHANNEL_SUCCESS.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self._openChannel(channel)
|
||||
self.conn.ssh_CHANNEL_REQUEST(b'\x00\x00\x00\x00' + common.NS(b'test')
|
||||
+ b'\x00')
|
||||
self.assertEqual(channel.numberRequests, 1)
|
||||
d = self.conn.ssh_CHANNEL_REQUEST(b'\x00\x00\x00\x00' + common.NS(
|
||||
b'test') + b'\xff' + b'data')
|
||||
def check(result):
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_SUCCESS, b'\x00\x00\x00\xff')])
|
||||
d.addCallback(check)
|
||||
return d
|
||||
|
||||
def test_CHANNEL_REQUEST_failure(self):
|
||||
"""
|
||||
Test that channel requests that fail send MSG_CHANNEL_FAILURE.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self._openChannel(channel)
|
||||
d = self.conn.ssh_CHANNEL_REQUEST(b'\x00\x00\x00\x00' + common.NS(
|
||||
b'test') + b'\xff')
|
||||
def check(result):
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_FAILURE, b'\x00\x00\x00\xff'
|
||||
)])
|
||||
d.addCallback(self.fail)
|
||||
d.addErrback(check)
|
||||
return d
|
||||
|
||||
def test_CHANNEL_REQUEST_SUCCESS(self):
|
||||
"""
|
||||
Test that channel request success messages cause the Deferred to be
|
||||
called back.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self._openChannel(channel)
|
||||
d = self.conn.sendRequest(channel, b'test', b'data', True)
|
||||
self.conn.ssh_CHANNEL_SUCCESS(b'\x00\x00\x00\x00')
|
||||
def check(result):
|
||||
self.assertTrue(result)
|
||||
return d
|
||||
|
||||
def test_CHANNEL_REQUEST_FAILURE(self):
|
||||
"""
|
||||
Test that channel request failure messages cause the Deferred to be
|
||||
erred back.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self._openChannel(channel)
|
||||
d = self.conn.sendRequest(channel, b'test', b'', True)
|
||||
self.conn.ssh_CHANNEL_FAILURE(b'\x00\x00\x00\x00')
|
||||
def check(result):
|
||||
self.assertEqual(result.value.value, 'channel request failed')
|
||||
d.addCallback(self.fail)
|
||||
d.addErrback(check)
|
||||
return d
|
||||
|
||||
def test_sendGlobalRequest(self):
|
||||
"""
|
||||
Test that global request messages are sent in the right format.
|
||||
"""
|
||||
d = self.conn.sendGlobalRequest(b'wantReply', b'data', True)
|
||||
# must be added to prevent errbacking during teardown
|
||||
d.addErrback(lambda failure: None)
|
||||
self.conn.sendGlobalRequest(b'noReply', b'', False)
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_GLOBAL_REQUEST, common.NS(b'wantReply') +
|
||||
b'\xffdata'),
|
||||
(connection.MSG_GLOBAL_REQUEST, common.NS(b'noReply') +
|
||||
b'\x00')])
|
||||
self.assertEqual(self.conn.deferreds, {'global':[d]})
|
||||
|
||||
def test_openChannel(self):
|
||||
"""
|
||||
Test that open channel messages are sent in the right format.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self.conn.openChannel(channel, b'aaaa')
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_OPEN, common.NS(b'TestChannel') +
|
||||
b'\x00\x00\x00\x00\x00\x02\x00\x00\x00\x00\x80\x00aaaa')])
|
||||
self.assertEqual(channel.id, 0)
|
||||
self.assertEqual(self.conn.localChannelID, 1)
|
||||
|
||||
def test_sendRequest(self):
|
||||
"""
|
||||
Test that channel request messages are sent in the right format.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self._openChannel(channel)
|
||||
d = self.conn.sendRequest(channel, b'test', b'test', True)
|
||||
# needed to prevent errbacks during teardown.
|
||||
d.addErrback(lambda failure: None)
|
||||
self.conn.sendRequest(channel, b'test2', b'', False)
|
||||
channel.localClosed = True # emulate sending a close message
|
||||
self.conn.sendRequest(channel, b'test3', b'', True)
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_REQUEST, b'\x00\x00\x00\xff' +
|
||||
common.NS(b'test') + b'\x01test'),
|
||||
(connection.MSG_CHANNEL_REQUEST, b'\x00\x00\x00\xff' +
|
||||
common.NS(b'test2') + b'\x00')])
|
||||
self.assertEqual(self.conn.deferreds[0], [d])
|
||||
|
||||
def test_adjustWindow(self):
|
||||
"""
|
||||
Test that channel window adjust messages cause bytes to be added
|
||||
to the window.
|
||||
"""
|
||||
channel = TestChannel(localWindow=5)
|
||||
self._openChannel(channel)
|
||||
channel.localWindowLeft = 0
|
||||
self.conn.adjustWindow(channel, 1)
|
||||
self.assertEqual(channel.localWindowLeft, 1)
|
||||
channel.localClosed = True
|
||||
self.conn.adjustWindow(channel, 2)
|
||||
self.assertEqual(channel.localWindowLeft, 1)
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_WINDOW_ADJUST, b'\x00\x00\x00\xff'
|
||||
b'\x00\x00\x00\x01')])
|
||||
|
||||
def test_sendData(self):
|
||||
"""
|
||||
Test that channel data messages are sent in the right format.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self._openChannel(channel)
|
||||
self.conn.sendData(channel, b'a')
|
||||
channel.localClosed = True
|
||||
self.conn.sendData(channel, b'b')
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_DATA, b'\x00\x00\x00\xff' +
|
||||
common.NS(b'a'))])
|
||||
|
||||
def test_sendExtendedData(self):
|
||||
"""
|
||||
Test that channel extended data messages are sent in the right format.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self._openChannel(channel)
|
||||
self.conn.sendExtendedData(channel, 1, b'test')
|
||||
channel.localClosed = True
|
||||
self.conn.sendExtendedData(channel, 2, b'test2')
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_EXTENDED_DATA, b'\x00\x00\x00\xff' +
|
||||
b'\x00\x00\x00\x01' + common.NS(b'test'))])
|
||||
|
||||
def test_sendEOF(self):
|
||||
"""
|
||||
Test that channel EOF messages are sent in the right format.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self._openChannel(channel)
|
||||
self.conn.sendEOF(channel)
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_EOF, b'\x00\x00\x00\xff')])
|
||||
channel.localClosed = True
|
||||
self.conn.sendEOF(channel)
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_EOF, b'\x00\x00\x00\xff')])
|
||||
|
||||
def test_sendClose(self):
|
||||
"""
|
||||
Test that channel close messages are sent in the right format.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self._openChannel(channel)
|
||||
self.conn.sendClose(channel)
|
||||
self.assertTrue(channel.localClosed)
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_CLOSE, b'\x00\x00\x00\xff')])
|
||||
self.conn.sendClose(channel)
|
||||
self.assertEqual(self.transport.packets,
|
||||
[(connection.MSG_CHANNEL_CLOSE, b'\x00\x00\x00\xff')])
|
||||
|
||||
channel2 = TestChannel()
|
||||
self._openChannel(channel2)
|
||||
self.assertTrue(channel2.gotOpen)
|
||||
self.assertFalse(channel2.gotClosed)
|
||||
channel2.remoteClosed = True
|
||||
self.conn.sendClose(channel2)
|
||||
self.assertTrue(channel2.gotClosed)
|
||||
|
||||
def test_getChannelWithAvatar(self):
|
||||
"""
|
||||
Test that getChannel dispatches to the avatar when an avatar is
|
||||
present. Correct functioning without the avatar is verified in
|
||||
test_CHANNEL_OPEN.
|
||||
"""
|
||||
channel = self.conn.getChannel(b'TestChannel', 50, 30, b'data')
|
||||
self.assertEqual(channel.data, b'data')
|
||||
self.assertEqual(channel.remoteWindowLeft, 50)
|
||||
self.assertEqual(channel.remoteMaxPacket, 30)
|
||||
self.assertRaises(error.ConchError, self.conn.getChannel,
|
||||
b'BadChannel', 50, 30, b'data')
|
||||
|
||||
def test_gotGlobalRequestWithoutAvatar(self):
|
||||
"""
|
||||
Test that gotGlobalRequests dispatches to global_* without an avatar.
|
||||
"""
|
||||
del self.transport.avatar
|
||||
self.assertTrue(self.conn.gotGlobalRequest(b'TestGlobal', b'data'))
|
||||
self.assertEqual(self.conn.gotGlobalRequest(b'Test-Data', b'data'),
|
||||
(True, b'data'))
|
||||
self.assertFalse(self.conn.gotGlobalRequest(b'BadGlobal', b'data'))
|
||||
|
||||
|
||||
def test_channelClosedCausesLeftoverChannelDeferredsToErrback(self):
|
||||
"""
|
||||
Whenever an SSH channel gets closed any Deferred that was returned by a
|
||||
sendRequest() on its parent connection must be errbacked.
|
||||
"""
|
||||
channel = TestChannel()
|
||||
self._openChannel(channel)
|
||||
|
||||
d = self.conn.sendRequest(
|
||||
channel, b"dummyrequest", b"dummydata", wantReply=1)
|
||||
d = self.assertFailure(d, error.ConchError)
|
||||
self.conn.channelClosed(channel)
|
||||
return d
|
||||
|
||||
|
||||
|
||||
class CleanConnectionShutdownTests(unittest.TestCase):
|
||||
"""
|
||||
Check whether correct cleanup is performed on connection shutdown.
|
||||
"""
|
||||
if not cryptography:
|
||||
skip = "Cannot run without cryptography"
|
||||
|
||||
if test_userauth.transport is None:
|
||||
skip = "Cannot run without both cryptography and pyasn1"
|
||||
|
||||
def setUp(self):
|
||||
self.transport = test_userauth.FakeTransport(None)
|
||||
self.transport.avatar = TestAvatar()
|
||||
self.conn = TestConnection()
|
||||
self.conn.transport = self.transport
|
||||
|
||||
|
||||
def test_serviceStoppedCausesLeftoverGlobalDeferredsToErrback(self):
|
||||
"""
|
||||
Once the service is stopped any leftover global deferred returned by
|
||||
a sendGlobalRequest() call must be errbacked.
|
||||
"""
|
||||
self.conn.serviceStarted()
|
||||
|
||||
d = self.conn.sendGlobalRequest(
|
||||
b"dummyrequest", b"dummydata", wantReply=1)
|
||||
d = self.assertFailure(d, error.ConchError)
|
||||
self.conn.serviceStopped()
|
||||
return d
|
||||
@@ -0,0 +1,333 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.conch.client.default}.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
import sys
|
||||
|
||||
from twisted.python.reflect import requireModule
|
||||
|
||||
if requireModule('cryptography') and requireModule('pyasn1'):
|
||||
from twisted.conch.client.agent import SSHAgentClient
|
||||
from twisted.conch.client.default import SSHUserAuthClient
|
||||
from twisted.conch.client.options import ConchOptions
|
||||
from twisted.conch.client import default
|
||||
from twisted.conch.ssh.keys import Key
|
||||
skip = None
|
||||
else:
|
||||
skip = "cryptography and PyASN1 required for twisted.conch.client.default."
|
||||
|
||||
from twisted.trial.unittest import TestCase
|
||||
from twisted.python.filepath import FilePath
|
||||
from twisted.conch.error import ConchError
|
||||
from twisted.conch.test import keydata
|
||||
from twisted.test.proto_helpers import StringTransport
|
||||
from twisted.python.compat import nativeString
|
||||
from twisted.python.runtime import platform
|
||||
|
||||
if platform.isWindows():
|
||||
windowsSkip = (
|
||||
"genericAnswers and getPassword does not work on Windows."
|
||||
" Should be fixed as part of fixing bug 6409 and 6410")
|
||||
else:
|
||||
windowsSkip = skip
|
||||
|
||||
ttySkip = None
|
||||
if not sys.stdin.isatty():
|
||||
ttySkip = "sys.stdin is not an interactive tty"
|
||||
if not sys.stdout.isatty():
|
||||
ttySkip = "sys.stdout is not an interactive tty"
|
||||
|
||||
|
||||
|
||||
class SSHUserAuthClientTests(TestCase):
|
||||
"""
|
||||
Tests for L{SSHUserAuthClient}.
|
||||
|
||||
@type rsaPublic: L{Key}
|
||||
@ivar rsaPublic: A public RSA key.
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
self.rsaPublic = Key.fromString(keydata.publicRSA_openssh)
|
||||
self.tmpdir = FilePath(self.mktemp())
|
||||
self.tmpdir.makedirs()
|
||||
self.rsaFile = self.tmpdir.child('id_rsa')
|
||||
self.rsaFile.setContent(keydata.privateRSA_openssh)
|
||||
self.tmpdir.child('id_rsa.pub').setContent(keydata.publicRSA_openssh)
|
||||
|
||||
|
||||
def test_signDataWithAgent(self):
|
||||
"""
|
||||
When connected to an agent, L{SSHUserAuthClient} can use it to
|
||||
request signatures of particular data with a particular L{Key}.
|
||||
"""
|
||||
client = SSHUserAuthClient(b"user", ConchOptions(), None)
|
||||
agent = SSHAgentClient()
|
||||
transport = StringTransport()
|
||||
agent.makeConnection(transport)
|
||||
client.keyAgent = agent
|
||||
cleartext = b"Sign here"
|
||||
client.signData(self.rsaPublic, cleartext)
|
||||
self.assertEqual(
|
||||
transport.value(),
|
||||
b"\x00\x00\x01\x2d\r\x00\x00\x01\x17" + self.rsaPublic.blob() +
|
||||
b"\x00\x00\x00\t" + cleartext +
|
||||
b"\x00\x00\x00\x00")
|
||||
|
||||
|
||||
def test_agentGetPublicKey(self):
|
||||
"""
|
||||
L{SSHUserAuthClient} looks up public keys from the agent using the
|
||||
L{SSHAgentClient} class. That L{SSHAgentClient.getPublicKey} returns a
|
||||
L{Key} object with one of the public keys in the agent. If no more
|
||||
keys are present, it returns L{None}.
|
||||
"""
|
||||
agent = SSHAgentClient()
|
||||
agent.blobs = [self.rsaPublic.blob()]
|
||||
key = agent.getPublicKey()
|
||||
self.assertTrue(key.isPublic())
|
||||
self.assertEqual(key, self.rsaPublic)
|
||||
self.assertIsNone(agent.getPublicKey())
|
||||
|
||||
|
||||
def test_getPublicKeyFromFile(self):
|
||||
"""
|
||||
L{SSHUserAuthClient.getPublicKey()} is able to get a public key from
|
||||
the first file described by its options' C{identitys} list, and return
|
||||
the corresponding public L{Key} object.
|
||||
"""
|
||||
options = ConchOptions()
|
||||
options.identitys = [self.rsaFile.path]
|
||||
client = SSHUserAuthClient(b"user", options, None)
|
||||
key = client.getPublicKey()
|
||||
self.assertTrue(key.isPublic())
|
||||
self.assertEqual(key, self.rsaPublic)
|
||||
|
||||
|
||||
def test_getPublicKeyAgentFallback(self):
|
||||
"""
|
||||
If an agent is present, but doesn't return a key,
|
||||
L{SSHUserAuthClient.getPublicKey} continue with the normal key lookup.
|
||||
"""
|
||||
options = ConchOptions()
|
||||
options.identitys = [self.rsaFile.path]
|
||||
agent = SSHAgentClient()
|
||||
client = SSHUserAuthClient(b"user", options, None)
|
||||
client.keyAgent = agent
|
||||
key = client.getPublicKey()
|
||||
self.assertTrue(key.isPublic())
|
||||
self.assertEqual(key, self.rsaPublic)
|
||||
|
||||
|
||||
def test_getPublicKeyBadKeyError(self):
|
||||
"""
|
||||
If L{keys.Key.fromFile} raises a L{keys.BadKeyError}, the
|
||||
L{SSHUserAuthClient.getPublicKey} tries again to get a public key by
|
||||
calling itself recursively.
|
||||
"""
|
||||
options = ConchOptions()
|
||||
self.tmpdir.child('id_dsa.pub').setContent(keydata.publicDSA_openssh)
|
||||
dsaFile = self.tmpdir.child('id_dsa')
|
||||
dsaFile.setContent(keydata.privateDSA_openssh)
|
||||
options.identitys = [self.rsaFile.path, dsaFile.path]
|
||||
self.tmpdir.child('id_rsa.pub').setContent(b'not a key!')
|
||||
client = SSHUserAuthClient(b"user", options, None)
|
||||
key = client.getPublicKey()
|
||||
self.assertTrue(key.isPublic())
|
||||
self.assertEqual(key, Key.fromString(keydata.publicDSA_openssh))
|
||||
self.assertEqual(client.usedFiles, [self.rsaFile.path, dsaFile.path])
|
||||
|
||||
|
||||
def test_getPrivateKey(self):
|
||||
"""
|
||||
L{SSHUserAuthClient.getPrivateKey} will load a private key from the
|
||||
last used file populated by L{SSHUserAuthClient.getPublicKey}, and
|
||||
return a L{Deferred} which fires with the corresponding private L{Key}.
|
||||
"""
|
||||
rsaPrivate = Key.fromString(keydata.privateRSA_openssh)
|
||||
options = ConchOptions()
|
||||
options.identitys = [self.rsaFile.path]
|
||||
client = SSHUserAuthClient(b"user", options, None)
|
||||
# Populate the list of used files
|
||||
client.getPublicKey()
|
||||
|
||||
def _cbGetPrivateKey(key):
|
||||
self.assertFalse(key.isPublic())
|
||||
self.assertEqual(key, rsaPrivate)
|
||||
|
||||
return client.getPrivateKey().addCallback(_cbGetPrivateKey)
|
||||
|
||||
|
||||
def test_getPrivateKeyPassphrase(self):
|
||||
"""
|
||||
L{SSHUserAuthClient} can get a private key from a file, and return a
|
||||
Deferred called back with a private L{Key} object, even if the key is
|
||||
encrypted.
|
||||
"""
|
||||
rsaPrivate = Key.fromString(keydata.privateRSA_openssh)
|
||||
passphrase = b'this is the passphrase'
|
||||
self.rsaFile.setContent(rsaPrivate.toString('openssh', passphrase))
|
||||
options = ConchOptions()
|
||||
options.identitys = [self.rsaFile.path]
|
||||
client = SSHUserAuthClient(b"user", options, None)
|
||||
# Populate the list of used files
|
||||
client.getPublicKey()
|
||||
|
||||
def _getPassword(prompt):
|
||||
self.assertEqual(
|
||||
prompt,
|
||||
"Enter passphrase for key '%s': " % (self.rsaFile.path,))
|
||||
return nativeString(passphrase)
|
||||
|
||||
def _cbGetPrivateKey(key):
|
||||
self.assertFalse(key.isPublic())
|
||||
self.assertEqual(key, rsaPrivate)
|
||||
|
||||
self.patch(client, '_getPassword', _getPassword)
|
||||
return client.getPrivateKey().addCallback(_cbGetPrivateKey)
|
||||
|
||||
|
||||
def test_getPassword(self):
|
||||
"""
|
||||
Get the password using
|
||||
L{twisted.conch.client.default.SSHUserAuthClient.getPassword}
|
||||
"""
|
||||
class FakeTransport:
|
||||
def __init__(self, host):
|
||||
self.transport = self
|
||||
self.host = host
|
||||
def getPeer(self):
|
||||
return self
|
||||
|
||||
options = ConchOptions()
|
||||
client = SSHUserAuthClient(b"user", options, None)
|
||||
client.transport = FakeTransport("127.0.0.1")
|
||||
|
||||
def getpass(prompt):
|
||||
self.assertEqual(prompt, "user@127.0.0.1's password: ")
|
||||
return 'bad password'
|
||||
|
||||
self.patch(default.getpass, 'getpass', getpass)
|
||||
d = client.getPassword()
|
||||
d.addCallback(self.assertEqual, b'bad password')
|
||||
return d
|
||||
|
||||
test_getPassword.skip = windowsSkip or ttySkip
|
||||
|
||||
|
||||
def test_getPasswordPrompt(self):
|
||||
"""
|
||||
Get the password using
|
||||
L{twisted.conch.client.default.SSHUserAuthClient.getPassword}
|
||||
using a different prompt.
|
||||
"""
|
||||
options = ConchOptions()
|
||||
client = SSHUserAuthClient(b"user", options, None)
|
||||
prompt = b"Give up your password"
|
||||
|
||||
def getpass(p):
|
||||
self.assertEqual(p, nativeString(prompt))
|
||||
return 'bad password'
|
||||
|
||||
self.patch(default.getpass, 'getpass', getpass)
|
||||
d = client.getPassword(prompt)
|
||||
d.addCallback(self.assertEqual, b'bad password')
|
||||
return d
|
||||
|
||||
test_getPasswordPrompt.skip = windowsSkip or ttySkip
|
||||
|
||||
|
||||
def test_getPasswordConchError(self):
|
||||
"""
|
||||
Get the password using
|
||||
L{twisted.conch.client.default.SSHUserAuthClient.getPassword}
|
||||
and trigger a {twisted.conch.error import ConchError}.
|
||||
"""
|
||||
options = ConchOptions()
|
||||
client = SSHUserAuthClient(b"user", options, None)
|
||||
|
||||
def getpass(prompt):
|
||||
raise KeyboardInterrupt("User pressed CTRL-C")
|
||||
|
||||
self.patch(default.getpass, 'getpass', getpass)
|
||||
stdout, stdin = sys.stdout, sys.stdin
|
||||
d = client.getPassword(b'?')
|
||||
@d.addErrback
|
||||
def check_sys(fail):
|
||||
self.assertEqual(
|
||||
[stdout, stdin], [sys.stdout, sys.stdin])
|
||||
return fail
|
||||
self.assertFailure(d, ConchError)
|
||||
|
||||
test_getPasswordConchError.skip = windowsSkip or ttySkip
|
||||
|
||||
|
||||
def test_getGenericAnswers(self):
|
||||
"""
|
||||
L{twisted.conch.client.default.SSHUserAuthClient.getGenericAnswers}
|
||||
"""
|
||||
options = ConchOptions()
|
||||
client = SSHUserAuthClient(b"user", options, None)
|
||||
|
||||
def getpass(prompt):
|
||||
self.assertEqual(prompt, "pass prompt")
|
||||
return "getpass"
|
||||
|
||||
self.patch(default.getpass, 'getpass', getpass)
|
||||
|
||||
def raw_input(prompt):
|
||||
self.assertEqual(prompt, "raw_input prompt")
|
||||
return "raw_input"
|
||||
|
||||
self.patch(default, 'raw_input', raw_input)
|
||||
d = client.getGenericAnswers(
|
||||
b"Name", b"Instruction", [
|
||||
(b"pass prompt", False), (b"raw_input prompt", True)])
|
||||
d.addCallback(
|
||||
self.assertListEqual, ["getpass", "raw_input"])
|
||||
return d
|
||||
|
||||
test_getGenericAnswers.skip = windowsSkip or ttySkip
|
||||
|
||||
|
||||
|
||||
|
||||
class ConchOptionsParsing(TestCase):
|
||||
"""
|
||||
Options parsing.
|
||||
"""
|
||||
def test_macs(self):
|
||||
"""
|
||||
Specify MAC algorithms.
|
||||
"""
|
||||
opts = ConchOptions()
|
||||
e = self.assertRaises(SystemExit, opts.opt_macs, "invalid-mac")
|
||||
self.assertIn("Unknown mac type", e.code)
|
||||
opts = ConchOptions()
|
||||
opts.opt_macs("hmac-sha2-512")
|
||||
self.assertEqual(opts['macs'], [b"hmac-sha2-512"])
|
||||
opts.opt_macs(b"hmac-sha2-512")
|
||||
self.assertEqual(opts['macs'], [b"hmac-sha2-512"])
|
||||
opts.opt_macs("hmac-sha2-256,hmac-sha1,hmac-md5")
|
||||
self.assertEqual(opts['macs'], [b"hmac-sha2-256", b"hmac-sha1", b"hmac-md5"])
|
||||
|
||||
|
||||
def test_host_key_algorithms(self):
|
||||
"""
|
||||
Specify host key algorithms.
|
||||
"""
|
||||
opts = ConchOptions()
|
||||
e = self.assertRaises(SystemExit, opts.opt_host_key_algorithms, "invalid-key")
|
||||
self.assertIn("Unknown host key type", e.code)
|
||||
opts = ConchOptions()
|
||||
opts.opt_host_key_algorithms("ssh-rsa")
|
||||
self.assertEqual(opts['host-key-algorithms'], [b"ssh-rsa"])
|
||||
opts.opt_host_key_algorithms(b"ssh-dss")
|
||||
self.assertEqual(opts['host-key-algorithms'], [b"ssh-dss"])
|
||||
opts.opt_host_key_algorithms("ssh-rsa,ssh-dss")
|
||||
self.assertEqual(opts['host-key-algorithms'], [b"ssh-rsa", b"ssh-dss"])
|
||||
@@ -0,0 +1,63 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.conch.ssh.forwarding}.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
from twisted.python.reflect import requireModule
|
||||
|
||||
cryptography = requireModule("cryptography")
|
||||
if cryptography:
|
||||
from twisted.conch.ssh import forwarding
|
||||
|
||||
from twisted.internet.address import IPv6Address
|
||||
from twisted.trial import unittest
|
||||
from twisted.internet.test.test_endpoints import deterministicResolvingReactor
|
||||
from twisted.test.proto_helpers import MemoryReactorClock, StringTransport
|
||||
|
||||
|
||||
class TestSSHConnectForwardingChannel(unittest.TestCase):
|
||||
"""
|
||||
Unit and integration tests for L{SSHConnectForwardingChannel}.
|
||||
"""
|
||||
|
||||
if not cryptography:
|
||||
skip = "Cannot run without cryptography"
|
||||
|
||||
def makeTCPConnection(self, reactor):
|
||||
"""
|
||||
Fake that connection was established for first connectTCP request made
|
||||
on C{reactor}.
|
||||
|
||||
@param reactor: Reactor on which to fake the connection.
|
||||
@type reactor: A reactor.
|
||||
"""
|
||||
factory = reactor.tcpClients[0][2]
|
||||
connector = reactor.connectors[0]
|
||||
protocol = factory.buildProtocol(None)
|
||||
transport = StringTransport(peerAddress=connector.getDestination())
|
||||
protocol.makeConnection(transport)
|
||||
|
||||
|
||||
def test_channelOpenHostnameRequests(self):
|
||||
"""
|
||||
When a hostname is sent as part of forwarding requests, it
|
||||
is resolved using HostnameEndpoint's resolver.
|
||||
"""
|
||||
sut = forwarding.SSHConnectForwardingChannel(
|
||||
hostport=('fwd.example.org', 1234))
|
||||
# Patch channel and resolver to not touch the network.
|
||||
memoryReactor = MemoryReactorClock()
|
||||
sut._reactor = deterministicResolvingReactor(memoryReactor, ['::1'])
|
||||
sut.channelOpen(None)
|
||||
|
||||
self.makeTCPConnection(memoryReactor)
|
||||
self.successResultOf(sut._channelOpenDeferred)
|
||||
# Channel is connected using a forwarding client to the resolved
|
||||
# address of the requested host.
|
||||
self.assertIsInstance(sut.client, forwarding.SSHForwardingClient)
|
||||
self.assertEqual(
|
||||
IPv6Address('TCP', '::1', 1234), sut.client.transport.getPeer())
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,77 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for the command-line interfaces to conch.
|
||||
"""
|
||||
from twisted.python.reflect import requireModule
|
||||
|
||||
if requireModule('pyasn1'):
|
||||
pyasn1Skip = None
|
||||
else:
|
||||
pyasn1Skip = "Cannot run without PyASN1"
|
||||
|
||||
if requireModule('cryptography'):
|
||||
cryptoSkip = None
|
||||
else:
|
||||
cryptoSkip = "can't run w/o cryptography"
|
||||
|
||||
if requireModule('tty'):
|
||||
ttySkip = None
|
||||
else:
|
||||
ttySkip = "can't run w/o tty"
|
||||
|
||||
try:
|
||||
import Tkinter
|
||||
except ImportError:
|
||||
tkskip = "can't run w/o Tkinter"
|
||||
else:
|
||||
try:
|
||||
Tkinter.Tk().destroy()
|
||||
except Tkinter.TclError as e:
|
||||
tkskip = "Can't test Tkinter: " + str(e)
|
||||
else:
|
||||
tkskip = None
|
||||
|
||||
from twisted.trial.unittest import TestCase
|
||||
from twisted.scripts.test.test_scripts import ScriptTestsMixin
|
||||
from twisted.python.test.test_shellcomp import ZshScriptTestMixin
|
||||
|
||||
|
||||
|
||||
class ScriptTests(TestCase, ScriptTestsMixin):
|
||||
"""
|
||||
Tests for the Conch scripts.
|
||||
"""
|
||||
skip = pyasn1Skip or cryptoSkip
|
||||
|
||||
|
||||
def test_conch(self):
|
||||
self.scriptTest("conch/conch")
|
||||
test_conch.skip = ttySkip or skip
|
||||
|
||||
|
||||
def test_cftp(self):
|
||||
self.scriptTest("conch/cftp")
|
||||
test_cftp.skip = ttySkip or skip
|
||||
|
||||
|
||||
def test_ckeygen(self):
|
||||
self.scriptTest("conch/ckeygen")
|
||||
|
||||
|
||||
def test_tkconch(self):
|
||||
self.scriptTest("conch/tkconch")
|
||||
test_tkconch.skip = tkskip or skip
|
||||
|
||||
|
||||
|
||||
class ZshIntegrationTests(TestCase, ZshScriptTestMixin):
|
||||
"""
|
||||
Test that zsh completion functions are generated without error
|
||||
"""
|
||||
generateFor = [('conch', 'twisted.conch.scripts.conch.ClientOptions'),
|
||||
('cftp', 'twisted.conch.scripts.cftp.ClientOptions'),
|
||||
('ckeygen', 'twisted.conch.scripts.ckeygen.GeneralOptions'),
|
||||
('tkconch', 'twisted.conch.scripts.tkconch.GeneralOptions'),
|
||||
]
|
||||
@@ -0,0 +1,997 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.conch.ssh}.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
import struct
|
||||
|
||||
from twisted.python.reflect import requireModule
|
||||
|
||||
cryptography = requireModule("cryptography")
|
||||
pyasn1 = requireModule("pyasn1")
|
||||
|
||||
if cryptography:
|
||||
from twisted.conch.ssh import common, forwarding, session, _kex
|
||||
from twisted.conch import avatar, error
|
||||
else:
|
||||
class avatar:
|
||||
class ConchUser: pass
|
||||
|
||||
from twisted.conch.test.keydata import publicRSA_openssh, privateRSA_openssh
|
||||
from twisted.conch.test.keydata import publicDSA_openssh, privateDSA_openssh
|
||||
from twisted.cred import portal
|
||||
from twisted.cred.error import UnauthorizedLogin
|
||||
from twisted.internet import defer, protocol, reactor
|
||||
from twisted.internet.error import ProcessTerminated
|
||||
from twisted.python import failure, log
|
||||
from twisted.trial import unittest
|
||||
|
||||
from twisted.conch.test.loopback import LoopbackRelay
|
||||
|
||||
|
||||
|
||||
class ConchTestRealm(object):
|
||||
"""
|
||||
A realm which expects a particular avatarId to log in once and creates a
|
||||
L{ConchTestAvatar} for that request.
|
||||
|
||||
@ivar expectedAvatarID: The only avatarID that this realm will produce an
|
||||
avatar for.
|
||||
|
||||
@ivar avatar: A reference to the avatar after it is requested.
|
||||
"""
|
||||
avatar = None
|
||||
|
||||
def __init__(self, expectedAvatarID):
|
||||
self.expectedAvatarID = expectedAvatarID
|
||||
|
||||
|
||||
def requestAvatar(self, avatarID, mind, *interfaces):
|
||||
"""
|
||||
Return a new L{ConchTestAvatar} if the avatarID matches the expected one
|
||||
and this is the first avatar request.
|
||||
"""
|
||||
if avatarID == self.expectedAvatarID:
|
||||
if self.avatar is not None:
|
||||
raise UnauthorizedLogin("Only one login allowed")
|
||||
self.avatar = ConchTestAvatar()
|
||||
return interfaces[0], self.avatar, self.avatar.logout
|
||||
raise UnauthorizedLogin(
|
||||
"Only %r may log in, not %r" % (self.expectedAvatarID, avatarID))
|
||||
|
||||
|
||||
|
||||
class ConchTestAvatar(avatar.ConchUser):
|
||||
"""
|
||||
An avatar against which various SSH features can be tested.
|
||||
|
||||
@ivar loggedOut: A flag indicating whether the avatar logout method has been
|
||||
called.
|
||||
"""
|
||||
if not cryptography:
|
||||
skip = "cannot run without cryptography"
|
||||
|
||||
loggedOut = False
|
||||
|
||||
def __init__(self):
|
||||
avatar.ConchUser.__init__(self)
|
||||
self.listeners = {}
|
||||
self.globalRequests = {}
|
||||
self.channelLookup.update(
|
||||
{b'session': session.SSHSession,
|
||||
b'direct-tcpip':forwarding.openConnectForwardingClient})
|
||||
self.subsystemLookup.update({b'crazy': CrazySubsystem})
|
||||
|
||||
|
||||
def global_foo(self, data):
|
||||
self.globalRequests['foo'] = data
|
||||
return 1
|
||||
|
||||
|
||||
def global_foo_2(self, data):
|
||||
self.globalRequests['foo_2'] = data
|
||||
return 1, b'data'
|
||||
|
||||
|
||||
def global_tcpip_forward(self, data):
|
||||
host, port = forwarding.unpackGlobal_tcpip_forward(data)
|
||||
try:
|
||||
listener = reactor.listenTCP(
|
||||
port, forwarding.SSHListenForwardingFactory(
|
||||
self.conn, (host, port),
|
||||
forwarding.SSHListenServerForwardingChannel),
|
||||
interface=host)
|
||||
except:
|
||||
log.err(None, "something went wrong with remote->local forwarding")
|
||||
return 0
|
||||
else:
|
||||
self.listeners[(host, port)] = listener
|
||||
return 1
|
||||
|
||||
|
||||
def global_cancel_tcpip_forward(self, data):
|
||||
host, port = forwarding.unpackGlobal_tcpip_forward(data)
|
||||
listener = self.listeners.get((host, port), None)
|
||||
if not listener:
|
||||
return 0
|
||||
del self.listeners[(host, port)]
|
||||
listener.stopListening()
|
||||
return 1
|
||||
|
||||
|
||||
def logout(self):
|
||||
self.loggedOut = True
|
||||
for listener in self.listeners.values():
|
||||
log.msg('stopListening %s' % listener)
|
||||
listener.stopListening()
|
||||
|
||||
|
||||
|
||||
class ConchSessionForTestAvatar(object):
|
||||
"""
|
||||
An ISession adapter for ConchTestAvatar.
|
||||
"""
|
||||
def __init__(self, avatar):
|
||||
"""
|
||||
Initialize the session and create a reference to it on the avatar for
|
||||
later inspection.
|
||||
"""
|
||||
self.avatar = avatar
|
||||
self.avatar._testSession = self
|
||||
self.cmd = None
|
||||
self.proto = None
|
||||
self.ptyReq = False
|
||||
self.eof = 0
|
||||
self.onClose = defer.Deferred()
|
||||
|
||||
|
||||
def getPty(self, term, windowSize, attrs):
|
||||
log.msg('pty req')
|
||||
self._terminalType = term
|
||||
self._windowSize = windowSize
|
||||
self.ptyReq = True
|
||||
|
||||
|
||||
def openShell(self, proto):
|
||||
log.msg('opening shell')
|
||||
self.proto = proto
|
||||
EchoTransport(proto)
|
||||
self.cmd = b'shell'
|
||||
|
||||
|
||||
def execCommand(self, proto, cmd):
|
||||
self.cmd = cmd
|
||||
self.proto = proto
|
||||
f = cmd.split()[0]
|
||||
if f == b'false':
|
||||
t = FalseTransport(proto)
|
||||
# Avoid disconnecting this immediately. If the channel is closed
|
||||
# before execCommand even returns the caller gets confused.
|
||||
reactor.callLater(0, t.loseConnection)
|
||||
elif f == b'echo':
|
||||
t = EchoTransport(proto)
|
||||
t.write(cmd[5:])
|
||||
t.loseConnection()
|
||||
elif f == b'secho':
|
||||
t = SuperEchoTransport(proto)
|
||||
t.write(cmd[6:])
|
||||
t.loseConnection()
|
||||
elif f == b'eecho':
|
||||
t = ErrEchoTransport(proto)
|
||||
t.write(cmd[6:])
|
||||
t.loseConnection()
|
||||
else:
|
||||
raise error.ConchError('bad exec')
|
||||
self.avatar.conn.transport.expectedLoseConnection = 1
|
||||
|
||||
|
||||
def eofReceived(self):
|
||||
self.eof = 1
|
||||
|
||||
|
||||
def closed(self):
|
||||
log.msg('closed cmd "%s"' % self.cmd)
|
||||
self.remoteWindowLeftAtClose = self.proto.session.remoteWindowLeft
|
||||
self.onClose.callback(None)
|
||||
|
||||
from twisted.python import components
|
||||
|
||||
if cryptography:
|
||||
components.registerAdapter(ConchSessionForTestAvatar, ConchTestAvatar,
|
||||
session.ISession)
|
||||
|
||||
class CrazySubsystem(protocol.Protocol):
|
||||
|
||||
def __init__(self, *args, **kw):
|
||||
pass
|
||||
|
||||
def connectionMade(self):
|
||||
"""
|
||||
good ... good
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class FalseTransport:
|
||||
"""
|
||||
False transport should act like a /bin/false execution, i.e. just exit with
|
||||
nonzero status, writing nothing to the terminal.
|
||||
|
||||
@ivar proto: The protocol associated with this transport.
|
||||
@ivar closed: A flag tracking whether C{loseConnection} has been called yet.
|
||||
"""
|
||||
|
||||
def __init__(self, p):
|
||||
"""
|
||||
@type p L{twisted.conch.ssh.session.SSHSessionProcessProtocol} instance
|
||||
"""
|
||||
self.proto = p
|
||||
p.makeConnection(self)
|
||||
self.closed = 0
|
||||
|
||||
|
||||
def loseConnection(self):
|
||||
"""
|
||||
Disconnect the protocol associated with this transport.
|
||||
"""
|
||||
if self.closed:
|
||||
return
|
||||
self.closed = 1
|
||||
self.proto.inConnectionLost()
|
||||
self.proto.outConnectionLost()
|
||||
self.proto.errConnectionLost()
|
||||
self.proto.processEnded(failure.Failure(ProcessTerminated(255, None, None)))
|
||||
|
||||
|
||||
|
||||
class EchoTransport:
|
||||
|
||||
def __init__(self, p):
|
||||
self.proto = p
|
||||
p.makeConnection(self)
|
||||
self.closed = 0
|
||||
|
||||
def write(self, data):
|
||||
log.msg(repr(data))
|
||||
self.proto.outReceived(data)
|
||||
self.proto.outReceived(b'\r\n')
|
||||
if b'\x00' in data: # mimic 'exit' for the shell test
|
||||
self.loseConnection()
|
||||
|
||||
def loseConnection(self):
|
||||
if self.closed: return
|
||||
self.closed = 1
|
||||
self.proto.inConnectionLost()
|
||||
self.proto.outConnectionLost()
|
||||
self.proto.errConnectionLost()
|
||||
self.proto.processEnded(failure.Failure(ProcessTerminated(0, None, None)))
|
||||
|
||||
class ErrEchoTransport:
|
||||
|
||||
def __init__(self, p):
|
||||
self.proto = p
|
||||
p.makeConnection(self)
|
||||
self.closed = 0
|
||||
|
||||
def write(self, data):
|
||||
self.proto.errReceived(data)
|
||||
self.proto.errReceived(b'\r\n')
|
||||
|
||||
def loseConnection(self):
|
||||
if self.closed: return
|
||||
self.closed = 1
|
||||
self.proto.inConnectionLost()
|
||||
self.proto.outConnectionLost()
|
||||
self.proto.errConnectionLost()
|
||||
self.proto.processEnded(failure.Failure(ProcessTerminated(0, None, None)))
|
||||
|
||||
class SuperEchoTransport:
|
||||
|
||||
def __init__(self, p):
|
||||
self.proto = p
|
||||
p.makeConnection(self)
|
||||
self.closed = 0
|
||||
|
||||
def write(self, data):
|
||||
self.proto.outReceived(data)
|
||||
self.proto.outReceived(b'\r\n')
|
||||
self.proto.errReceived(data)
|
||||
self.proto.errReceived(b'\r\n')
|
||||
|
||||
def loseConnection(self):
|
||||
if self.closed: return
|
||||
self.closed = 1
|
||||
self.proto.inConnectionLost()
|
||||
self.proto.outConnectionLost()
|
||||
self.proto.errConnectionLost()
|
||||
self.proto.processEnded(failure.Failure(ProcessTerminated(0, None, None)))
|
||||
|
||||
|
||||
if cryptography is not None and pyasn1 is not None:
|
||||
from twisted.conch import checkers
|
||||
from twisted.conch.ssh import channel, connection, factory, keys
|
||||
from twisted.conch.ssh import transport, userauth
|
||||
|
||||
class ConchTestPasswordChecker:
|
||||
credentialInterfaces = checkers.IUsernamePassword,
|
||||
|
||||
def requestAvatarId(self, credentials):
|
||||
if credentials.username == b'testuser' and credentials.password == b'testpass':
|
||||
return defer.succeed(credentials.username)
|
||||
return defer.fail(Exception("Bad credentials"))
|
||||
|
||||
|
||||
class ConchTestSSHChecker(checkers.SSHProtocolChecker):
|
||||
|
||||
def areDone(self, avatarId):
|
||||
if avatarId != b'testuser' or len(self.successfulCredentials[avatarId]) < 2:
|
||||
return False
|
||||
return True
|
||||
|
||||
class ConchTestServerFactory(factory.SSHFactory):
|
||||
noisy = 0
|
||||
|
||||
services = {
|
||||
b'ssh-userauth':userauth.SSHUserAuthServer,
|
||||
b'ssh-connection':connection.SSHConnection
|
||||
}
|
||||
|
||||
def buildProtocol(self, addr):
|
||||
proto = ConchTestServer()
|
||||
proto.supportedPublicKeys = self.privateKeys.keys()
|
||||
proto.factory = self
|
||||
|
||||
if hasattr(self, 'expectedLoseConnection'):
|
||||
proto.expectedLoseConnection = self.expectedLoseConnection
|
||||
|
||||
self.proto = proto
|
||||
return proto
|
||||
|
||||
def getPublicKeys(self):
|
||||
return {
|
||||
b'ssh-rsa': keys.Key.fromString(publicRSA_openssh),
|
||||
b'ssh-dss': keys.Key.fromString(publicDSA_openssh)
|
||||
}
|
||||
|
||||
def getPrivateKeys(self):
|
||||
return {
|
||||
b'ssh-rsa': keys.Key.fromString(privateRSA_openssh),
|
||||
b'ssh-dss': keys.Key.fromString(privateDSA_openssh)
|
||||
}
|
||||
|
||||
def getPrimes(self):
|
||||
"""
|
||||
Diffie-Hellman primes that can be used for the
|
||||
diffie-hellman-group-exchange-sha1 key exchange.
|
||||
|
||||
@return: The primes and generators.
|
||||
@rtype: L{dict} mapping the key size to a C{list} of
|
||||
C{(generator, prime)} tupple.
|
||||
"""
|
||||
# In these tests, we hardwire the prime values to those defined by
|
||||
# the diffie-hellman-group14-sha1 key exchange algorithm, to avoid
|
||||
# requiring a moduli file when running tests.
|
||||
# See OpenSSHFactory.getPrimes.
|
||||
return {
|
||||
2048: [
|
||||
_kex.getDHGeneratorAndPrime(
|
||||
b'diffie-hellman-group14-sha1')]
|
||||
}
|
||||
|
||||
def getService(self, trans, name):
|
||||
return factory.SSHFactory.getService(self, trans, name)
|
||||
|
||||
class ConchTestBase:
|
||||
|
||||
done = 0
|
||||
|
||||
def connectionLost(self, reason):
|
||||
if self.done:
|
||||
return
|
||||
if not hasattr(self, 'expectedLoseConnection'):
|
||||
raise unittest.FailTest(
|
||||
'unexpectedly lost connection %s\n%s' % (self, reason))
|
||||
self.done = 1
|
||||
|
||||
def receiveError(self, reasonCode, desc):
|
||||
self.expectedLoseConnection = 1
|
||||
# Some versions of OpenSSH (for example, OpenSSH_5.3p1) will
|
||||
# send a DISCONNECT_BY_APPLICATION error before closing the
|
||||
# connection. Other, older versions (for example,
|
||||
# OpenSSH_5.1p1), won't. So accept this particular error here,
|
||||
# but no others.
|
||||
if reasonCode != transport.DISCONNECT_BY_APPLICATION:
|
||||
log.err(
|
||||
Exception(
|
||||
'got disconnect for %s: reason %s, desc: %s' % (
|
||||
self, reasonCode, desc)))
|
||||
self.loseConnection()
|
||||
|
||||
def receiveUnimplemented(self, seqID):
|
||||
raise unittest.FailTest('got unimplemented: seqid %s' % (seqID,))
|
||||
self.expectedLoseConnection = 1
|
||||
self.loseConnection()
|
||||
|
||||
class ConchTestServer(ConchTestBase, transport.SSHServerTransport):
|
||||
|
||||
def connectionLost(self, reason):
|
||||
ConchTestBase.connectionLost(self, reason)
|
||||
transport.SSHServerTransport.connectionLost(self, reason)
|
||||
|
||||
|
||||
class ConchTestClient(ConchTestBase, transport.SSHClientTransport):
|
||||
"""
|
||||
@ivar _channelFactory: A callable which accepts an SSH connection and
|
||||
returns a channel which will be attached to a new channel on that
|
||||
connection.
|
||||
"""
|
||||
def __init__(self, channelFactory):
|
||||
self._channelFactory = channelFactory
|
||||
|
||||
def connectionLost(self, reason):
|
||||
ConchTestBase.connectionLost(self, reason)
|
||||
transport.SSHClientTransport.connectionLost(self, reason)
|
||||
|
||||
def verifyHostKey(self, key, fp):
|
||||
keyMatch = key == keys.Key.fromString(publicRSA_openssh).blob()
|
||||
fingerprintMatch = (
|
||||
fp == b'85:25:04:32:58:55:96:9f:57:ee:fb:a8:1a:ea:69:da')
|
||||
if keyMatch and fingerprintMatch:
|
||||
return defer.succeed(1)
|
||||
return defer.fail(Exception("Key or fingerprint mismatch"))
|
||||
|
||||
def connectionSecure(self):
|
||||
self.requestService(ConchTestClientAuth(b'testuser',
|
||||
ConchTestClientConnection(self._channelFactory)))
|
||||
|
||||
|
||||
class ConchTestClientAuth(userauth.SSHUserAuthClient):
|
||||
|
||||
hasTriedNone = 0 # have we tried the 'none' auth yet?
|
||||
canSucceedPublicKey = 0 # can we succeed with this yet?
|
||||
canSucceedPassword = 0
|
||||
|
||||
def ssh_USERAUTH_SUCCESS(self, packet):
|
||||
if not self.canSucceedPassword and self.canSucceedPublicKey:
|
||||
raise unittest.FailTest(
|
||||
'got USERAUTH_SUCCESS before password and publickey')
|
||||
userauth.SSHUserAuthClient.ssh_USERAUTH_SUCCESS(self, packet)
|
||||
|
||||
def getPassword(self):
|
||||
self.canSucceedPassword = 1
|
||||
return defer.succeed(b'testpass')
|
||||
|
||||
def getPrivateKey(self):
|
||||
self.canSucceedPublicKey = 1
|
||||
return defer.succeed(keys.Key.fromString(privateDSA_openssh))
|
||||
|
||||
def getPublicKey(self):
|
||||
return keys.Key.fromString(publicDSA_openssh)
|
||||
|
||||
|
||||
class ConchTestClientConnection(connection.SSHConnection):
|
||||
"""
|
||||
@ivar _completed: A L{Deferred} which will be fired when the number of
|
||||
results collected reaches C{totalResults}.
|
||||
"""
|
||||
name = b'ssh-connection'
|
||||
results = 0
|
||||
totalResults = 8
|
||||
|
||||
def __init__(self, channelFactory):
|
||||
connection.SSHConnection.__init__(self)
|
||||
self._channelFactory = channelFactory
|
||||
|
||||
def serviceStarted(self):
|
||||
self.openChannel(self._channelFactory(conn=self))
|
||||
|
||||
|
||||
class SSHTestChannel(channel.SSHChannel):
|
||||
|
||||
def __init__(self, name, opened, *args, **kwargs):
|
||||
self.name = name
|
||||
self._opened = opened
|
||||
self.received = []
|
||||
self.receivedExt = []
|
||||
self.onClose = defer.Deferred()
|
||||
channel.SSHChannel.__init__(self, *args, **kwargs)
|
||||
|
||||
|
||||
def openFailed(self, reason):
|
||||
self._opened.errback(reason)
|
||||
|
||||
|
||||
def channelOpen(self, ignore):
|
||||
self._opened.callback(self)
|
||||
|
||||
|
||||
def dataReceived(self, data):
|
||||
self.received.append(data)
|
||||
|
||||
|
||||
def extReceived(self, dataType, data):
|
||||
if dataType == connection.EXTENDED_DATA_STDERR:
|
||||
self.receivedExt.append(data)
|
||||
else:
|
||||
log.msg("Unrecognized extended data: %r" % (dataType,))
|
||||
|
||||
|
||||
def request_exit_status(self, status):
|
||||
[self.status] = struct.unpack('>L', status)
|
||||
|
||||
|
||||
def eofReceived(self):
|
||||
self.eofCalled = True
|
||||
|
||||
|
||||
def closed(self):
|
||||
self.onClose.callback(None)
|
||||
|
||||
|
||||
def conchTestPublicKeyChecker():
|
||||
"""
|
||||
Produces a SSHPublicKeyChecker with an in-memory key mapping with
|
||||
a single use: 'testuser'
|
||||
|
||||
@return: L{twisted.conch.checkers.SSHPublicKeyChecker}
|
||||
"""
|
||||
conchTestPublicKeyDB = checkers.InMemorySSHKeyDB(
|
||||
{b'testuser': [keys.Key.fromString(publicDSA_openssh)]})
|
||||
return checkers.SSHPublicKeyChecker(conchTestPublicKeyDB)
|
||||
|
||||
|
||||
|
||||
class SSHProtocolTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for communication between L{SSHServerTransport} and
|
||||
L{SSHClientTransport}.
|
||||
"""
|
||||
|
||||
if not cryptography:
|
||||
skip = "can't run without cryptography"
|
||||
|
||||
if not pyasn1:
|
||||
skip = "Cannot run without PyASN1"
|
||||
|
||||
def _ourServerOurClientTest(self, name=b'session', **kwargs):
|
||||
"""
|
||||
Create a connected SSH client and server protocol pair and return a
|
||||
L{Deferred} which fires with an L{SSHTestChannel} instance connected to
|
||||
a channel on that SSH connection.
|
||||
"""
|
||||
result = defer.Deferred()
|
||||
self.realm = ConchTestRealm(b'testuser')
|
||||
p = portal.Portal(self.realm)
|
||||
sshpc = ConchTestSSHChecker()
|
||||
sshpc.registerChecker(ConchTestPasswordChecker())
|
||||
sshpc.registerChecker(conchTestPublicKeyChecker())
|
||||
p.registerChecker(sshpc)
|
||||
fac = ConchTestServerFactory()
|
||||
fac.portal = p
|
||||
fac.startFactory()
|
||||
self.server = fac.buildProtocol(None)
|
||||
self.clientTransport = LoopbackRelay(self.server)
|
||||
self.client = ConchTestClient(
|
||||
lambda conn: SSHTestChannel(name, result, conn=conn, **kwargs))
|
||||
|
||||
self.serverTransport = LoopbackRelay(self.client)
|
||||
|
||||
self.server.makeConnection(self.serverTransport)
|
||||
self.client.makeConnection(self.clientTransport)
|
||||
return result
|
||||
|
||||
|
||||
def test_subsystemsAndGlobalRequests(self):
|
||||
"""
|
||||
Run the Conch server against the Conch client. Set up several different
|
||||
channels which exercise different behaviors and wait for them to
|
||||
complete. Verify that the channels with errors log them.
|
||||
"""
|
||||
channel = self._ourServerOurClientTest()
|
||||
|
||||
def cbSubsystem(channel):
|
||||
self.channel = channel
|
||||
return self.assertFailure(
|
||||
channel.conn.sendRequest(
|
||||
channel, b'subsystem', common.NS(b'not-crazy'), 1),
|
||||
Exception)
|
||||
channel.addCallback(cbSubsystem)
|
||||
|
||||
def cbNotCrazyFailed(ignored):
|
||||
channel = self.channel
|
||||
return channel.conn.sendRequest(
|
||||
channel, b'subsystem', common.NS(b'crazy'), 1)
|
||||
channel.addCallback(cbNotCrazyFailed)
|
||||
|
||||
def cbGlobalRequests(ignored):
|
||||
channel = self.channel
|
||||
d1 = channel.conn.sendGlobalRequest(b'foo', b'bar', 1)
|
||||
|
||||
d2 = channel.conn.sendGlobalRequest(b'foo-2', b'bar2', 1)
|
||||
d2.addCallback(self.assertEqual, b'data')
|
||||
|
||||
d3 = self.assertFailure(
|
||||
channel.conn.sendGlobalRequest(b'bar', b'foo', 1),
|
||||
Exception)
|
||||
|
||||
return defer.gatherResults([d1, d2, d3])
|
||||
channel.addCallback(cbGlobalRequests)
|
||||
|
||||
def disconnect(ignored):
|
||||
self.assertEqual(
|
||||
self.realm.avatar.globalRequests,
|
||||
{"foo": b"bar", "foo_2": b"bar2"})
|
||||
channel = self.channel
|
||||
channel.conn.transport.expectedLoseConnection = True
|
||||
channel.conn.serviceStopped()
|
||||
channel.loseConnection()
|
||||
channel.addCallback(disconnect)
|
||||
|
||||
return channel
|
||||
|
||||
|
||||
def test_shell(self):
|
||||
"""
|
||||
L{SSHChannel.sendRequest} can open a shell with a I{pty-req} request,
|
||||
specifying a terminal type and window size.
|
||||
"""
|
||||
channel = self._ourServerOurClientTest()
|
||||
|
||||
data = session.packRequest_pty_req(
|
||||
b'conch-test-term', (24, 80, 0, 0), b'')
|
||||
def cbChannel(channel):
|
||||
self.channel = channel
|
||||
return channel.conn.sendRequest(channel, b'pty-req', data, 1)
|
||||
channel.addCallback(cbChannel)
|
||||
|
||||
def cbPty(ignored):
|
||||
# The server-side object corresponding to our client side channel.
|
||||
session = self.realm.avatar.conn.channels[0].session
|
||||
self.assertIs(session.avatar, self.realm.avatar)
|
||||
self.assertEqual(session._terminalType, b'conch-test-term')
|
||||
self.assertEqual(session._windowSize, (24, 80, 0, 0))
|
||||
self.assertTrue(session.ptyReq)
|
||||
channel = self.channel
|
||||
return channel.conn.sendRequest(channel, b'shell', b'', 1)
|
||||
channel.addCallback(cbPty)
|
||||
|
||||
def cbShell(ignored):
|
||||
self.channel.write(b'testing the shell!\x00')
|
||||
self.channel.conn.sendEOF(self.channel)
|
||||
return defer.gatherResults([
|
||||
self.channel.onClose,
|
||||
self.realm.avatar._testSession.onClose])
|
||||
channel.addCallback(cbShell)
|
||||
|
||||
def cbExited(ignored):
|
||||
if self.channel.status != 0:
|
||||
log.msg(
|
||||
'shell exit status was not 0: %i' % (self.channel.status,))
|
||||
self.assertEqual(
|
||||
b"".join(self.channel.received),
|
||||
b'testing the shell!\x00\r\n')
|
||||
self.assertTrue(self.channel.eofCalled)
|
||||
self.assertTrue(
|
||||
self.realm.avatar._testSession.eof)
|
||||
channel.addCallback(cbExited)
|
||||
return channel
|
||||
|
||||
|
||||
def test_failedExec(self):
|
||||
"""
|
||||
If L{SSHChannel.sendRequest} issues an exec which the server responds to
|
||||
with an error, the L{Deferred} it returns fires its errback.
|
||||
"""
|
||||
channel = self._ourServerOurClientTest()
|
||||
|
||||
def cbChannel(channel):
|
||||
self.channel = channel
|
||||
return self.assertFailure(
|
||||
channel.conn.sendRequest(
|
||||
channel, b'exec', common.NS(b'jumboliah'), 1),
|
||||
Exception)
|
||||
channel.addCallback(cbChannel)
|
||||
|
||||
def cbFailed(ignored):
|
||||
# The server logs this exception when it cannot perform the
|
||||
# requested exec.
|
||||
errors = self.flushLoggedErrors(error.ConchError)
|
||||
self.assertEqual(errors[0].value.args, ('bad exec', None))
|
||||
channel.addCallback(cbFailed)
|
||||
return channel
|
||||
|
||||
|
||||
def test_falseChannel(self):
|
||||
"""
|
||||
When the process started by a L{SSHChannel.sendRequest} exec request
|
||||
exits, the exit status is reported to the channel.
|
||||
"""
|
||||
channel = self._ourServerOurClientTest()
|
||||
|
||||
def cbChannel(channel):
|
||||
self.channel = channel
|
||||
return channel.conn.sendRequest(
|
||||
channel, b'exec', common.NS(b'false'), 1)
|
||||
channel.addCallback(cbChannel)
|
||||
|
||||
def cbExec(ignored):
|
||||
return self.channel.onClose
|
||||
channel.addCallback(cbExec)
|
||||
|
||||
def cbClosed(ignored):
|
||||
# No data is expected
|
||||
self.assertEqual(self.channel.received, [])
|
||||
self.assertNotEqual(self.channel.status, 0)
|
||||
channel.addCallback(cbClosed)
|
||||
return channel
|
||||
|
||||
|
||||
def test_errorChannel(self):
|
||||
"""
|
||||
Bytes sent over the extended channel for stderr data are delivered to
|
||||
the channel's C{extReceived} method.
|
||||
"""
|
||||
channel = self._ourServerOurClientTest(localWindow=4, localMaxPacket=5)
|
||||
|
||||
def cbChannel(channel):
|
||||
self.channel = channel
|
||||
return channel.conn.sendRequest(
|
||||
channel, b'exec', common.NS(b'eecho hello'), 1)
|
||||
channel.addCallback(cbChannel)
|
||||
|
||||
def cbExec(ignored):
|
||||
return defer.gatherResults([
|
||||
self.channel.onClose,
|
||||
self.realm.avatar._testSession.onClose])
|
||||
channel.addCallback(cbExec)
|
||||
|
||||
def cbClosed(ignored):
|
||||
self.assertEqual(self.channel.received, [])
|
||||
self.assertEqual(b"".join(self.channel.receivedExt), b"hello\r\n")
|
||||
self.assertEqual(self.channel.status, 0)
|
||||
self.assertTrue(self.channel.eofCalled)
|
||||
self.assertEqual(self.channel.localWindowLeft, 4)
|
||||
self.assertEqual(
|
||||
self.channel.localWindowLeft,
|
||||
self.realm.avatar._testSession.remoteWindowLeftAtClose)
|
||||
channel.addCallback(cbClosed)
|
||||
return channel
|
||||
|
||||
|
||||
def test_unknownChannel(self):
|
||||
"""
|
||||
When an attempt is made to open an unknown channel type, the L{Deferred}
|
||||
returned by L{SSHChannel.sendRequest} fires its errback.
|
||||
"""
|
||||
d = self.assertFailure(
|
||||
self._ourServerOurClientTest(b'crazy-unknown-channel'), Exception)
|
||||
def cbFailed(ignored):
|
||||
errors = self.flushLoggedErrors(error.ConchError)
|
||||
self.assertEqual(errors[0].value.args, (3, 'unknown channel'))
|
||||
self.assertEqual(len(errors), 1)
|
||||
d.addCallback(cbFailed)
|
||||
return d
|
||||
|
||||
|
||||
def test_maxPacket(self):
|
||||
"""
|
||||
An L{SSHChannel} can be configured with a maximum packet size to
|
||||
receive.
|
||||
"""
|
||||
# localWindow needs to be at least 11 otherwise the assertion about it
|
||||
# in cbClosed is invalid.
|
||||
channel = self._ourServerOurClientTest(
|
||||
localWindow=11, localMaxPacket=1)
|
||||
|
||||
def cbChannel(channel):
|
||||
self.channel = channel
|
||||
return channel.conn.sendRequest(
|
||||
channel, b'exec', common.NS(b'secho hello'), 1)
|
||||
channel.addCallback(cbChannel)
|
||||
|
||||
def cbExec(ignored):
|
||||
return self.channel.onClose
|
||||
channel.addCallback(cbExec)
|
||||
|
||||
def cbClosed(ignored):
|
||||
self.assertEqual(self.channel.status, 0)
|
||||
self.assertEqual(b"".join(self.channel.received), b"hello\r\n")
|
||||
self.assertEqual(b"".join(self.channel.receivedExt), b"hello\r\n")
|
||||
self.assertEqual(self.channel.localWindowLeft, 11)
|
||||
self.assertTrue(self.channel.eofCalled)
|
||||
channel.addCallback(cbClosed)
|
||||
return channel
|
||||
|
||||
|
||||
def test_echo(self):
|
||||
"""
|
||||
Normal standard out bytes are sent to the channel's C{dataReceived}
|
||||
method.
|
||||
"""
|
||||
channel = self._ourServerOurClientTest(localWindow=4, localMaxPacket=5)
|
||||
|
||||
def cbChannel(channel):
|
||||
self.channel = channel
|
||||
return channel.conn.sendRequest(
|
||||
channel, b'exec', common.NS(b'echo hello'), 1)
|
||||
channel.addCallback(cbChannel)
|
||||
|
||||
def cbEcho(ignored):
|
||||
return defer.gatherResults([
|
||||
self.channel.onClose,
|
||||
self.realm.avatar._testSession.onClose])
|
||||
channel.addCallback(cbEcho)
|
||||
|
||||
def cbClosed(ignored):
|
||||
self.assertEqual(self.channel.status, 0)
|
||||
self.assertEqual(b"".join(self.channel.received), b"hello\r\n")
|
||||
self.assertEqual(self.channel.localWindowLeft, 4)
|
||||
self.assertTrue(self.channel.eofCalled)
|
||||
self.assertEqual(
|
||||
self.channel.localWindowLeft,
|
||||
self.realm.avatar._testSession.remoteWindowLeftAtClose)
|
||||
channel.addCallback(cbClosed)
|
||||
return channel
|
||||
|
||||
|
||||
|
||||
class SSHFactoryTests(unittest.TestCase):
|
||||
|
||||
if not cryptography:
|
||||
skip = "can't run without cryptography"
|
||||
|
||||
if not pyasn1:
|
||||
skip = "Cannot run without PyASN1"
|
||||
|
||||
def makeSSHFactory(self, primes=None):
|
||||
sshFactory = factory.SSHFactory()
|
||||
gpk = lambda: {'ssh-rsa' : keys.Key(None)}
|
||||
sshFactory.getPrimes = lambda: primes
|
||||
sshFactory.getPublicKeys = sshFactory.getPrivateKeys = gpk
|
||||
sshFactory.startFactory()
|
||||
return sshFactory
|
||||
|
||||
|
||||
def test_buildProtocol(self):
|
||||
"""
|
||||
By default, buildProtocol() constructs an instance of
|
||||
SSHServerTransport.
|
||||
"""
|
||||
factory = self.makeSSHFactory()
|
||||
protocol = factory.buildProtocol(None)
|
||||
self.assertIsInstance(protocol, transport.SSHServerTransport)
|
||||
|
||||
|
||||
def test_buildProtocolRespectsProtocol(self):
|
||||
"""
|
||||
buildProtocol() calls 'self.protocol()' to construct a protocol
|
||||
instance.
|
||||
"""
|
||||
calls = []
|
||||
def makeProtocol(*args):
|
||||
calls.append(args)
|
||||
return transport.SSHServerTransport()
|
||||
factory = self.makeSSHFactory()
|
||||
factory.protocol = makeProtocol
|
||||
factory.buildProtocol(None)
|
||||
self.assertEqual([()], calls)
|
||||
|
||||
|
||||
def test_buildProtocolNoPrimes(self):
|
||||
"""
|
||||
Group key exchanges are not supported when we don't have the primes
|
||||
database.
|
||||
"""
|
||||
f1 = self.makeSSHFactory(primes=None)
|
||||
|
||||
p1 = f1.buildProtocol(None)
|
||||
|
||||
self.assertNotIn(
|
||||
b'diffie-hellman-group-exchange-sha1', p1.supportedKeyExchanges)
|
||||
self.assertNotIn(
|
||||
b'diffie-hellman-group-exchange-sha256', p1.supportedKeyExchanges)
|
||||
|
||||
|
||||
def test_buildProtocolWithPrimes(self):
|
||||
"""
|
||||
Group key exchanges are supported when we have the primes database.
|
||||
"""
|
||||
f2 = self.makeSSHFactory(primes={1:(2,3)})
|
||||
|
||||
p2 = f2.buildProtocol(None)
|
||||
|
||||
self.assertIn(
|
||||
b'diffie-hellman-group-exchange-sha1', p2.supportedKeyExchanges)
|
||||
self.assertIn(
|
||||
b'diffie-hellman-group-exchange-sha256', p2.supportedKeyExchanges)
|
||||
|
||||
|
||||
|
||||
class MPTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for L{common.getMP}.
|
||||
|
||||
@cvar getMP: a method providing a MP parser.
|
||||
@type getMP: C{callable}
|
||||
"""
|
||||
if not cryptography:
|
||||
skip = "can't run without cryptography"
|
||||
|
||||
if not pyasn1:
|
||||
skip = "Cannot run without PyASN1"
|
||||
|
||||
if cryptography:
|
||||
getMP = staticmethod(common.getMP)
|
||||
|
||||
def test_getMP(self):
|
||||
"""
|
||||
L{common.getMP} should parse the a multiple precision integer from a
|
||||
string: a 4-byte length followed by length bytes of the integer.
|
||||
"""
|
||||
self.assertEqual(
|
||||
self.getMP(b'\x00\x00\x00\x04\x00\x00\x00\x01'),
|
||||
(1, b''))
|
||||
|
||||
|
||||
def test_getMPBigInteger(self):
|
||||
"""
|
||||
L{common.getMP} should be able to parse a big enough integer
|
||||
(that doesn't fit on one byte).
|
||||
"""
|
||||
self.assertEqual(
|
||||
self.getMP(b'\x00\x00\x00\x04\x01\x02\x03\x04'),
|
||||
(16909060, b''))
|
||||
|
||||
|
||||
def test_multipleGetMP(self):
|
||||
"""
|
||||
L{common.getMP} has the ability to parse multiple integer in the same
|
||||
string.
|
||||
"""
|
||||
self.assertEqual(
|
||||
self.getMP(b'\x00\x00\x00\x04\x00\x00\x00\x01'
|
||||
b'\x00\x00\x00\x04\x00\x00\x00\x02', 2),
|
||||
(1, 2, b''))
|
||||
|
||||
|
||||
def test_getMPRemainingData(self):
|
||||
"""
|
||||
When more data than needed is sent to L{common.getMP}, it should return
|
||||
the remaining data.
|
||||
"""
|
||||
self.assertEqual(
|
||||
self.getMP(b'\x00\x00\x00\x04\x00\x00\x00\x01foo'),
|
||||
(1, b'foo'))
|
||||
|
||||
|
||||
def test_notEnoughData(self):
|
||||
"""
|
||||
When the string passed to L{common.getMP} doesn't even make 5 bytes,
|
||||
it should raise a L{struct.error}.
|
||||
"""
|
||||
self.assertRaises(struct.error, self.getMP, b'\x02\x00')
|
||||
|
||||
|
||||
class GMPYInstallDeprecationTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for the deprecation of former GMPY accidental public API.
|
||||
"""
|
||||
|
||||
if not cryptography:
|
||||
skip = "cannot run without cryptography"
|
||||
|
||||
def test_deprecated(self):
|
||||
"""
|
||||
L{twisted.conch.ssh.common.install} is deprecated.
|
||||
"""
|
||||
common.install()
|
||||
warnings = self.flushWarnings([self.test_deprecated])
|
||||
self.assertEqual(len(warnings), 1)
|
||||
self.assertEqual(
|
||||
warnings[0]["message"],
|
||||
"twisted.conch.ssh.common.install was deprecated in Twisted 16.5.0"
|
||||
)
|
||||
@@ -0,0 +1,121 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
#
|
||||
|
||||
import tty
|
||||
# this module was autogenerated.
|
||||
|
||||
VINTR = 1
|
||||
VQUIT = 2
|
||||
VERASE = 3
|
||||
VKILL = 4
|
||||
VEOF = 5
|
||||
VEOL = 6
|
||||
VEOL2 = 7
|
||||
VSTART = 8
|
||||
VSTOP = 9
|
||||
VSUSP = 10
|
||||
VDSUSP = 11
|
||||
VREPRINT = 12
|
||||
VWERASE = 13
|
||||
VLNEXT = 14
|
||||
VFLUSH = 15
|
||||
VSWTCH = 16
|
||||
VSTATUS = 17
|
||||
VDISCARD = 18
|
||||
IGNPAR = 30
|
||||
PARMRK = 31
|
||||
INPCK = 32
|
||||
ISTRIP = 33
|
||||
INLCR = 34
|
||||
IGNCR = 35
|
||||
ICRNL = 36
|
||||
IUCLC = 37
|
||||
IXON = 38
|
||||
IXANY = 39
|
||||
IXOFF = 40
|
||||
IMAXBEL = 41
|
||||
ISIG = 50
|
||||
ICANON = 51
|
||||
XCASE = 52
|
||||
ECHO = 53
|
||||
ECHOE = 54
|
||||
ECHOK = 55
|
||||
ECHONL = 56
|
||||
NOFLSH = 57
|
||||
TOSTOP = 58
|
||||
IEXTEN = 59
|
||||
ECHOCTL = 60
|
||||
ECHOKE = 61
|
||||
PENDIN = 62
|
||||
OPOST = 70
|
||||
OLCUC = 71
|
||||
ONLCR = 72
|
||||
OCRNL = 73
|
||||
ONOCR = 74
|
||||
ONLRET = 75
|
||||
CS7 = 90
|
||||
CS8 = 91
|
||||
PARENB = 92
|
||||
PARODD = 93
|
||||
TTY_OP_ISPEED = 128
|
||||
TTY_OP_OSPEED = 129
|
||||
|
||||
TTYMODES = {
|
||||
1 : 'VINTR',
|
||||
2 : 'VQUIT',
|
||||
3 : 'VERASE',
|
||||
4 : 'VKILL',
|
||||
5 : 'VEOF',
|
||||
6 : 'VEOL',
|
||||
7 : 'VEOL2',
|
||||
8 : 'VSTART',
|
||||
9 : 'VSTOP',
|
||||
10 : 'VSUSP',
|
||||
11 : 'VDSUSP',
|
||||
12 : 'VREPRINT',
|
||||
13 : 'VWERASE',
|
||||
14 : 'VLNEXT',
|
||||
15 : 'VFLUSH',
|
||||
16 : 'VSWTCH',
|
||||
17 : 'VSTATUS',
|
||||
18 : 'VDISCARD',
|
||||
30 : (tty.IFLAG, 'IGNPAR'),
|
||||
31 : (tty.IFLAG, 'PARMRK'),
|
||||
32 : (tty.IFLAG, 'INPCK'),
|
||||
33 : (tty.IFLAG, 'ISTRIP'),
|
||||
34 : (tty.IFLAG, 'INLCR'),
|
||||
35 : (tty.IFLAG, 'IGNCR'),
|
||||
36 : (tty.IFLAG, 'ICRNL'),
|
||||
37 : (tty.IFLAG, 'IUCLC'),
|
||||
38 : (tty.IFLAG, 'IXON'),
|
||||
39 : (tty.IFLAG, 'IXANY'),
|
||||
40 : (tty.IFLAG, 'IXOFF'),
|
||||
41 : (tty.IFLAG, 'IMAXBEL'),
|
||||
50 : (tty.LFLAG, 'ISIG'),
|
||||
51 : (tty.LFLAG, 'ICANON'),
|
||||
52 : (tty.LFLAG, 'XCASE'),
|
||||
53 : (tty.LFLAG, 'ECHO'),
|
||||
54 : (tty.LFLAG, 'ECHOE'),
|
||||
55 : (tty.LFLAG, 'ECHOK'),
|
||||
56 : (tty.LFLAG, 'ECHONL'),
|
||||
57 : (tty.LFLAG, 'NOFLSH'),
|
||||
58 : (tty.LFLAG, 'TOSTOP'),
|
||||
59 : (tty.LFLAG, 'IEXTEN'),
|
||||
60 : (tty.LFLAG, 'ECHOCTL'),
|
||||
61 : (tty.LFLAG, 'ECHOKE'),
|
||||
62 : (tty.LFLAG, 'PENDIN'),
|
||||
70 : (tty.OFLAG, 'OPOST'),
|
||||
71 : (tty.OFLAG, 'OLCUC'),
|
||||
72 : (tty.OFLAG, 'ONLCR'),
|
||||
73 : (tty.OFLAG, 'OCRNL'),
|
||||
74 : (tty.OFLAG, 'ONOCR'),
|
||||
75 : (tty.OFLAG, 'ONLRET'),
|
||||
# 90 : (tty.CFLAG, 'CS7'),
|
||||
# 91 : (tty.CFLAG, 'CS8'),
|
||||
92 : (tty.CFLAG, 'PARENB'),
|
||||
93 : (tty.CFLAG, 'PARODD'),
|
||||
128 : 'ISPEED',
|
||||
129 : 'OSPEED'
|
||||
}
|
||||
@@ -0,0 +1,240 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
#
|
||||
"""Module to parse ANSI escape sequences
|
||||
|
||||
Maintainer: Jean-Paul Calderone
|
||||
"""
|
||||
|
||||
import string
|
||||
|
||||
# Twisted imports
|
||||
from twisted.python import log
|
||||
|
||||
class ColorText:
|
||||
"""
|
||||
Represents an element of text along with the texts colors and
|
||||
additional attributes.
|
||||
"""
|
||||
|
||||
# The colors to use
|
||||
COLORS = ('b', 'r', 'g', 'y', 'l', 'm', 'c', 'w')
|
||||
BOLD_COLORS = tuple([x.upper() for x in COLORS])
|
||||
BLACK, RED, GREEN, YELLOW, BLUE, MAGENTA, CYAN, WHITE = range(len(COLORS))
|
||||
|
||||
# Color names
|
||||
COLOR_NAMES = (
|
||||
'Black', 'Red', 'Green', 'Yellow', 'Blue', 'Magenta', 'Cyan', 'White'
|
||||
)
|
||||
|
||||
def __init__(self, text, fg, bg, display, bold, underline, flash, reverse):
|
||||
self.text, self.fg, self.bg = text, fg, bg
|
||||
self.display = display
|
||||
self.bold = bold
|
||||
self.underline = underline
|
||||
self.flash = flash
|
||||
self.reverse = reverse
|
||||
if self.reverse:
|
||||
self.fg, self.bg = self.bg, self.fg
|
||||
|
||||
|
||||
class AnsiParser:
|
||||
"""
|
||||
Parser class for ANSI codes.
|
||||
"""
|
||||
|
||||
# Terminators for cursor movement ansi controls - unsupported
|
||||
CURSOR_SET = ('H', 'f', 'A', 'B', 'C', 'D', 'R', 's', 'u', 'd','G')
|
||||
|
||||
# Terminators for erasure ansi controls - unsupported
|
||||
ERASE_SET = ('J', 'K', 'P')
|
||||
|
||||
# Terminators for mode change ansi controls - unsupported
|
||||
MODE_SET = ('h', 'l')
|
||||
|
||||
# Terminators for keyboard assignment ansi controls - unsupported
|
||||
ASSIGN_SET = ('p',)
|
||||
|
||||
# Terminators for color change ansi controls - supported
|
||||
COLOR_SET = ('m',)
|
||||
|
||||
SETS = (CURSOR_SET, ERASE_SET, MODE_SET, ASSIGN_SET, COLOR_SET)
|
||||
|
||||
def __init__(self, defaultFG, defaultBG):
|
||||
self.defaultFG, self.defaultBG = defaultFG, defaultBG
|
||||
self.currentFG, self.currentBG = self.defaultFG, self.defaultBG
|
||||
self.bold, self.flash, self.underline, self.reverse = 0, 0, 0, 0
|
||||
self.display = 1
|
||||
self.prepend = ''
|
||||
|
||||
|
||||
def stripEscapes(self, string):
|
||||
"""
|
||||
Remove all ANSI color escapes from the given string.
|
||||
"""
|
||||
result = ''
|
||||
show = 1
|
||||
i = 0
|
||||
L = len(string)
|
||||
while i < L:
|
||||
if show == 0 and string[i] in _sets:
|
||||
show = 1
|
||||
elif show:
|
||||
n = string.find('\x1B', i)
|
||||
if n == -1:
|
||||
return result + string[i:]
|
||||
else:
|
||||
result = result + string[i:n]
|
||||
i = n
|
||||
show = 0
|
||||
i = i + 1
|
||||
return result
|
||||
|
||||
def writeString(self, colorstr):
|
||||
pass
|
||||
|
||||
def parseString(self, str):
|
||||
"""
|
||||
Turn a string input into a list of L{ColorText} elements.
|
||||
"""
|
||||
|
||||
if self.prepend:
|
||||
str = self.prepend + str
|
||||
self.prepend = ''
|
||||
parts = str.split('\x1B')
|
||||
|
||||
if len(parts) == 1:
|
||||
self.writeString(self.formatText(parts[0]))
|
||||
else:
|
||||
self.writeString(self.formatText(parts[0]))
|
||||
for s in parts[1:]:
|
||||
L = len(s)
|
||||
i = 0
|
||||
type = None
|
||||
while i < L:
|
||||
if s[i] not in string.digits+'[;?':
|
||||
break
|
||||
i+=1
|
||||
if not s:
|
||||
self.prepend = '\x1b'
|
||||
return
|
||||
if s[0]!='[':
|
||||
self.writeString(self.formatText(s[i+1:]))
|
||||
continue
|
||||
else:
|
||||
s=s[1:]
|
||||
i-=1
|
||||
if i==L-1:
|
||||
self.prepend = '\x1b['
|
||||
return
|
||||
type = _setmap.get(s[i], None)
|
||||
if type is None:
|
||||
continue
|
||||
|
||||
if type == AnsiParser.COLOR_SET:
|
||||
self.parseColor(s[:i + 1])
|
||||
s = s[i + 1:]
|
||||
self.writeString(self.formatText(s))
|
||||
elif type == AnsiParser.CURSOR_SET:
|
||||
cursor, s = s[:i+1], s[i+1:]
|
||||
self.parseCursor(cursor)
|
||||
self.writeString(self.formatText(s))
|
||||
elif type == AnsiParser.ERASE_SET:
|
||||
erase, s = s[:i+1], s[i+1:]
|
||||
self.parseErase(erase)
|
||||
self.writeString(self.formatText(s))
|
||||
elif type == AnsiParser.MODE_SET:
|
||||
s = s[i+1:]
|
||||
#self.parseErase('2J')
|
||||
self.writeString(self.formatText(s))
|
||||
elif i == L:
|
||||
self.prepend = '\x1B[' + s
|
||||
else:
|
||||
log.msg('Unhandled ANSI control type: %c' % (s[i],))
|
||||
s = s[i + 1:]
|
||||
self.writeString(self.formatText(s))
|
||||
|
||||
def parseColor(self, str):
|
||||
"""
|
||||
Handle a single ANSI color sequence
|
||||
"""
|
||||
# Drop the trailing 'm'
|
||||
str = str[:-1]
|
||||
|
||||
if not str:
|
||||
str = '0'
|
||||
|
||||
try:
|
||||
parts = map(int, str.split(';'))
|
||||
except ValueError:
|
||||
log.msg('Invalid ANSI color sequence (%d): %s' % (len(str), str))
|
||||
self.currentFG, self.currentBG = self.defaultFG, self.defaultBG
|
||||
return
|
||||
|
||||
for x in parts:
|
||||
if x == 0:
|
||||
self.currentFG, self.currentBG = self.defaultFG, self.defaultBG
|
||||
self.bold, self.flash, self.underline, self.reverse = 0, 0, 0, 0
|
||||
self.display = 1
|
||||
elif x == 1:
|
||||
self.bold = 1
|
||||
elif 30 <= x <= 37:
|
||||
self.currentFG = x - 30
|
||||
elif 40 <= x <= 47:
|
||||
self.currentBG = x - 40
|
||||
elif x == 39:
|
||||
self.currentFG = self.defaultFG
|
||||
elif x == 49:
|
||||
self.currentBG = self.defaultBG
|
||||
elif x == 4:
|
||||
self.underline = 1
|
||||
elif x == 5:
|
||||
self.flash = 1
|
||||
elif x == 7:
|
||||
self.reverse = 1
|
||||
elif x == 8:
|
||||
self.display = 0
|
||||
elif x == 22:
|
||||
self.bold = 0
|
||||
elif x == 24:
|
||||
self.underline = 0
|
||||
elif x == 25:
|
||||
self.blink = 0
|
||||
elif x == 27:
|
||||
self.reverse = 0
|
||||
elif x == 28:
|
||||
self.display = 1
|
||||
else:
|
||||
log.msg('Unrecognised ANSI color command: %d' % (x,))
|
||||
|
||||
def parseCursor(self, cursor):
|
||||
pass
|
||||
|
||||
def parseErase(self, erase):
|
||||
pass
|
||||
|
||||
|
||||
def pickColor(self, value, mode, BOLD = ColorText.BOLD_COLORS):
|
||||
if mode:
|
||||
return ColorText.COLORS[value]
|
||||
else:
|
||||
return self.bold and BOLD[value] or ColorText.COLORS[value]
|
||||
|
||||
|
||||
def formatText(self, text):
|
||||
return ColorText(
|
||||
text,
|
||||
self.pickColor(self.currentFG, 0),
|
||||
self.pickColor(self.currentBG, 1),
|
||||
self.display, self.bold, self.underline, self.flash, self.reverse
|
||||
)
|
||||
|
||||
|
||||
_sets = ''.join(map(''.join, AnsiParser.SETS))
|
||||
|
||||
_setmap = {}
|
||||
for s in AnsiParser.SETS:
|
||||
for r in s:
|
||||
_setmap[r] = s
|
||||
del s
|
||||
@@ -0,0 +1,43 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Copyright information for Twisted.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
__all__ = ['copyright', 'disclaimer', 'longversion' ,'version']
|
||||
|
||||
from twisted import __version__ as version, version as longversion
|
||||
|
||||
longversion = str(longversion)
|
||||
|
||||
copyright="""\
|
||||
Copyright (c) 2001-2019 Twisted Matrix Laboratories.
|
||||
See LICENSE for details."""
|
||||
|
||||
disclaimer='''
|
||||
Twisted, the Framework of Your Internet
|
||||
%s
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining
|
||||
a copy of this software and associated documentation files (the
|
||||
"Software"), to deal in the Software without restriction, including
|
||||
without limitation the rights to use, copy, modify, merge, publish,
|
||||
distribute, sublicense, and/or sell copies of the Software, and to
|
||||
permit persons to whom the Software is furnished to do so, subject to
|
||||
the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be
|
||||
included in all copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND,
|
||||
EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF
|
||||
MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND
|
||||
NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE
|
||||
LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION
|
||||
OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION
|
||||
WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||
|
||||
''' % (copyright,)
|
||||
@@ -0,0 +1,132 @@
|
||||
# -*- test-case-name: twisted.cred.test.test_digestauth -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Calculations for HTTP Digest authentication.
|
||||
|
||||
@see: U{http://www.faqs.org/rfcs/rfc2617.html}
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
from binascii import hexlify
|
||||
from hashlib import md5, sha1
|
||||
|
||||
|
||||
|
||||
# The digest math
|
||||
|
||||
algorithms = {
|
||||
b'md5': md5,
|
||||
|
||||
# md5-sess is more complicated than just another algorithm. It requires
|
||||
# H(A1) state to be remembered from the first WWW-Authenticate challenge
|
||||
# issued and re-used to process any Authorization header in response to
|
||||
# that WWW-Authenticate challenge. It is *not* correct to simply
|
||||
# recalculate H(A1) each time an Authorization header is received. Read
|
||||
# RFC 2617, section 3.2.2.2 and do not try to make DigestCredentialFactory
|
||||
# support this unless you completely understand it. -exarkun
|
||||
b'md5-sess': md5,
|
||||
|
||||
b'sha': sha1,
|
||||
}
|
||||
|
||||
# DigestCalcHA1
|
||||
def calcHA1(pszAlg, pszUserName, pszRealm, pszPassword, pszNonce, pszCNonce,
|
||||
preHA1=None):
|
||||
"""
|
||||
Compute H(A1) from RFC 2617.
|
||||
|
||||
@param pszAlg: The name of the algorithm to use to calculate the digest.
|
||||
Currently supported are md5, md5-sess, and sha.
|
||||
@param pszUserName: The username
|
||||
@param pszRealm: The realm
|
||||
@param pszPassword: The password
|
||||
@param pszNonce: The nonce
|
||||
@param pszCNonce: The cnonce
|
||||
|
||||
@param preHA1: If available this is a str containing a previously
|
||||
calculated H(A1) as a hex string. If this is given then the values for
|
||||
pszUserName, pszRealm, and pszPassword must be L{None} and are ignored.
|
||||
"""
|
||||
|
||||
if (preHA1 and (pszUserName or pszRealm or pszPassword)):
|
||||
raise TypeError(("preHA1 is incompatible with the pszUserName, "
|
||||
"pszRealm, and pszPassword arguments"))
|
||||
|
||||
if preHA1 is None:
|
||||
# We need to calculate the HA1 from the username:realm:password
|
||||
m = algorithms[pszAlg]()
|
||||
m.update(pszUserName)
|
||||
m.update(b":")
|
||||
m.update(pszRealm)
|
||||
m.update(b":")
|
||||
m.update(pszPassword)
|
||||
HA1 = hexlify(m.digest())
|
||||
else:
|
||||
# We were given a username:realm:password
|
||||
HA1 = preHA1
|
||||
|
||||
if pszAlg == b"md5-sess":
|
||||
m = algorithms[pszAlg]()
|
||||
m.update(HA1)
|
||||
m.update(b":")
|
||||
m.update(pszNonce)
|
||||
m.update(b":")
|
||||
m.update(pszCNonce)
|
||||
HA1 = hexlify(m.digest())
|
||||
|
||||
return HA1
|
||||
|
||||
|
||||
def calcHA2(algo, pszMethod, pszDigestUri, pszQop, pszHEntity):
|
||||
"""
|
||||
Compute H(A2) from RFC 2617.
|
||||
|
||||
@param pszAlg: The name of the algorithm to use to calculate the digest.
|
||||
Currently supported are md5, md5-sess, and sha.
|
||||
@param pszMethod: The request method.
|
||||
@param pszDigestUri: The request URI.
|
||||
@param pszQop: The Quality-of-Protection value.
|
||||
@param pszHEntity: The hash of the entity body or L{None} if C{pszQop} is
|
||||
not C{'auth-int'}.
|
||||
@return: The hash of the A2 value for the calculation of the response
|
||||
digest.
|
||||
"""
|
||||
m = algorithms[algo]()
|
||||
m.update(pszMethod)
|
||||
m.update(b":")
|
||||
m.update(pszDigestUri)
|
||||
if pszQop == b"auth-int":
|
||||
m.update(b":")
|
||||
m.update(pszHEntity)
|
||||
return hexlify(m.digest())
|
||||
|
||||
|
||||
def calcResponse(HA1, HA2, algo, pszNonce, pszNonceCount, pszCNonce, pszQop):
|
||||
"""
|
||||
Compute the digest for the given parameters.
|
||||
|
||||
@param HA1: The H(A1) value, as computed by L{calcHA1}.
|
||||
@param HA2: The H(A2) value, as computed by L{calcHA2}.
|
||||
@param pszNonce: The challenge nonce.
|
||||
@param pszNonceCount: The (client) nonce count value for this response.
|
||||
@param pszCNonce: The client nonce.
|
||||
@param pszQop: The Quality-of-Protection value.
|
||||
"""
|
||||
m = algorithms[algo]()
|
||||
m.update(HA1)
|
||||
m.update(b":")
|
||||
m.update(pszNonce)
|
||||
m.update(b":")
|
||||
if pszNonceCount and pszCNonce:
|
||||
m.update(pszNonceCount)
|
||||
m.update(b":")
|
||||
m.update(pszCNonce)
|
||||
m.update(b":")
|
||||
m.update(pszQop)
|
||||
m.update(b":")
|
||||
m.update(HA2)
|
||||
respHash = hexlify(m.digest())
|
||||
return respHash
|
||||
@@ -0,0 +1,264 @@
|
||||
# -*- test-case-name: twisted.cred.test.test_cred -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
import os
|
||||
|
||||
from zope.interface import implementer, Interface, Attribute
|
||||
|
||||
from twisted.logger import Logger
|
||||
from twisted.internet import defer
|
||||
from twisted.python import failure
|
||||
from twisted.cred import error, credentials
|
||||
|
||||
|
||||
class ICredentialsChecker(Interface):
|
||||
"""
|
||||
An object that can check sub-interfaces of ICredentials.
|
||||
"""
|
||||
|
||||
credentialInterfaces = Attribute(
|
||||
'A list of sub-interfaces of ICredentials which specifies which I may check.')
|
||||
|
||||
|
||||
def requestAvatarId(credentials):
|
||||
"""
|
||||
@param credentials: something which implements one of the interfaces in
|
||||
self.credentialInterfaces.
|
||||
|
||||
@return: a Deferred which will fire a string which identifies an
|
||||
avatar, an empty tuple to specify an authenticated anonymous user
|
||||
(provided as checkers.ANONYMOUS) or fire a Failure(UnauthorizedLogin).
|
||||
Alternatively, return the result itself.
|
||||
|
||||
@see: L{twisted.cred.credentials}
|
||||
"""
|
||||
|
||||
|
||||
|
||||
# A note on anonymity - We do not want None as the value for anonymous
|
||||
# because it is too easy to accidentally return it. We do not want the
|
||||
# empty string, because it is too easy to mistype a password file. For
|
||||
# example, an .htpasswd file may contain the lines: ['hello:asdf',
|
||||
# 'world:asdf', 'goodbye', ':world']. This misconfiguration will have an
|
||||
# ill effect in any case, but accidentally granting anonymous access is a
|
||||
# worse failure mode than simply granting access to an untypeable
|
||||
# username. We do not want an instance of 'object', because that would
|
||||
# create potential problems with persistence.
|
||||
|
||||
ANONYMOUS = ()
|
||||
|
||||
|
||||
@implementer(ICredentialsChecker)
|
||||
class AllowAnonymousAccess:
|
||||
credentialInterfaces = credentials.IAnonymous,
|
||||
|
||||
def requestAvatarId(self, credentials):
|
||||
return defer.succeed(ANONYMOUS)
|
||||
|
||||
|
||||
|
||||
@implementer(ICredentialsChecker)
|
||||
class InMemoryUsernamePasswordDatabaseDontUse(object):
|
||||
"""
|
||||
An extremely simple credentials checker.
|
||||
|
||||
This is only of use in one-off test programs or examples which don't
|
||||
want to focus too much on how credentials are verified.
|
||||
|
||||
You really don't want to use this for anything else. It is, at best, a
|
||||
toy. If you need a simple credentials checker for a real application,
|
||||
see L{FilePasswordDB}.
|
||||
"""
|
||||
credentialInterfaces = (credentials.IUsernamePassword,
|
||||
credentials.IUsernameHashedPassword)
|
||||
|
||||
def __init__(self, **users):
|
||||
self.users = {x.encode('ascii'):y for x, y in users.items()}
|
||||
|
||||
|
||||
def addUser(self, username, password):
|
||||
self.users[username] = password
|
||||
|
||||
|
||||
def _cbPasswordMatch(self, matched, username):
|
||||
if matched:
|
||||
return username
|
||||
else:
|
||||
return failure.Failure(error.UnauthorizedLogin())
|
||||
|
||||
|
||||
def requestAvatarId(self, credentials):
|
||||
if credentials.username in self.users:
|
||||
return defer.maybeDeferred(
|
||||
credentials.checkPassword,
|
||||
self.users[credentials.username]).addCallback(
|
||||
self._cbPasswordMatch, credentials.username)
|
||||
else:
|
||||
return defer.fail(error.UnauthorizedLogin())
|
||||
|
||||
|
||||
|
||||
@implementer(ICredentialsChecker)
|
||||
class FilePasswordDB:
|
||||
"""
|
||||
A file-based, text-based username/password database.
|
||||
|
||||
Records in the datafile for this class are delimited by a particular
|
||||
string. The username appears in a fixed field of the columns delimited
|
||||
by this string, as does the password. Both fields are specifiable. If
|
||||
the passwords are not stored plaintext, a hash function must be supplied
|
||||
to convert plaintext passwords to the form stored on disk and this
|
||||
CredentialsChecker will only be able to check IUsernamePassword
|
||||
credentials. If the passwords are stored plaintext,
|
||||
IUsernameHashedPassword credentials will be checkable as well.
|
||||
"""
|
||||
|
||||
cache = False
|
||||
_credCache = None
|
||||
_cacheTimestamp = 0
|
||||
_log = Logger()
|
||||
|
||||
def __init__(self, filename, delim=b':', usernameField=0, passwordField=1,
|
||||
caseSensitive=True, hash=None, cache=False):
|
||||
"""
|
||||
@type filename: C{str}
|
||||
@param filename: The name of the file from which to read username and
|
||||
password information.
|
||||
|
||||
@type delim: C{str}
|
||||
@param delim: The field delimiter used in the file.
|
||||
|
||||
@type usernameField: C{int}
|
||||
@param usernameField: The index of the username after splitting a
|
||||
line on the delimiter.
|
||||
|
||||
@type passwordField: C{int}
|
||||
@param passwordField: The index of the password after splitting a
|
||||
line on the delimiter.
|
||||
|
||||
@type caseSensitive: C{bool}
|
||||
@param caseSensitive: If true, consider the case of the username when
|
||||
performing a lookup. Ignore it otherwise.
|
||||
|
||||
@type hash: Three-argument callable or L{None}
|
||||
@param hash: A function used to transform the plaintext password
|
||||
received over the network to a format suitable for comparison
|
||||
against the version stored on disk. The arguments to the callable
|
||||
are the username, the network-supplied password, and the in-file
|
||||
version of the password. If the return value compares equal to the
|
||||
version stored on disk, the credentials are accepted.
|
||||
|
||||
@type cache: C{bool}
|
||||
@param cache: If true, maintain an in-memory cache of the
|
||||
contents of the password file. On lookups, the mtime of the
|
||||
file will be checked, and the file will only be re-parsed if
|
||||
the mtime is newer than when the cache was generated.
|
||||
"""
|
||||
self.filename = filename
|
||||
self.delim = delim
|
||||
self.ufield = usernameField
|
||||
self.pfield = passwordField
|
||||
self.caseSensitive = caseSensitive
|
||||
self.hash = hash
|
||||
self.cache = cache
|
||||
|
||||
if self.hash is None:
|
||||
# The passwords are stored plaintext. We can support both
|
||||
# plaintext and hashed passwords received over the network.
|
||||
self.credentialInterfaces = (
|
||||
credentials.IUsernamePassword,
|
||||
credentials.IUsernameHashedPassword
|
||||
)
|
||||
else:
|
||||
# The passwords are hashed on disk. We can support only
|
||||
# plaintext passwords received over the network.
|
||||
self.credentialInterfaces = (
|
||||
credentials.IUsernamePassword,
|
||||
)
|
||||
|
||||
|
||||
def __getstate__(self):
|
||||
d = dict(vars(self))
|
||||
for k in '_credCache', '_cacheTimestamp':
|
||||
try:
|
||||
del d[k]
|
||||
except KeyError:
|
||||
pass
|
||||
return d
|
||||
|
||||
|
||||
def _cbPasswordMatch(self, matched, username):
|
||||
if matched:
|
||||
return username
|
||||
else:
|
||||
return failure.Failure(error.UnauthorizedLogin())
|
||||
|
||||
|
||||
def _loadCredentials(self):
|
||||
"""
|
||||
Loads the credentials from the configured file.
|
||||
|
||||
@return: An iterable of C{username, password} couples.
|
||||
@rtype: C{iterable}
|
||||
|
||||
@raise UnauthorizedLogin: when failing to read the credentials from the
|
||||
file.
|
||||
"""
|
||||
try:
|
||||
with open(self.filename, "rb") as f:
|
||||
for line in f:
|
||||
line = line.rstrip()
|
||||
parts = line.split(self.delim)
|
||||
|
||||
if self.ufield >= len(parts) or self.pfield >= len(parts):
|
||||
continue
|
||||
if self.caseSensitive:
|
||||
yield parts[self.ufield], parts[self.pfield]
|
||||
else:
|
||||
yield parts[self.ufield].lower(), parts[self.pfield]
|
||||
except IOError as e:
|
||||
self._log.error("Unable to load credentials db: {e!r}", e=e)
|
||||
raise error.UnauthorizedLogin()
|
||||
|
||||
|
||||
def getUser(self, username):
|
||||
if not self.caseSensitive:
|
||||
username = username.lower()
|
||||
|
||||
if self.cache:
|
||||
if self._credCache is None or os.path.getmtime(self.filename) > self._cacheTimestamp:
|
||||
self._cacheTimestamp = os.path.getmtime(self.filename)
|
||||
self._credCache = dict(self._loadCredentials())
|
||||
return username, self._credCache[username]
|
||||
else:
|
||||
for u, p in self._loadCredentials():
|
||||
if u == username:
|
||||
return u, p
|
||||
raise KeyError(username)
|
||||
|
||||
|
||||
def requestAvatarId(self, c):
|
||||
try:
|
||||
u, p = self.getUser(c.username)
|
||||
except KeyError:
|
||||
return defer.fail(error.UnauthorizedLogin())
|
||||
else:
|
||||
up = credentials.IUsernamePassword(c, None)
|
||||
if self.hash:
|
||||
if up is not None:
|
||||
h = self.hash(up.username, up.password, p)
|
||||
if h == p:
|
||||
return defer.succeed(u)
|
||||
return defer.fail(error.UnauthorizedLogin())
|
||||
else:
|
||||
return defer.maybeDeferred(c.checkPassword, p
|
||||
).addCallback(self._cbPasswordMatch, u)
|
||||
|
||||
|
||||
|
||||
# For backwards compatibility
|
||||
# Allow access as the old name.
|
||||
OnDiskUsernamePasswordDatabase = FilePasswordDB
|
||||
@@ -0,0 +1,508 @@
|
||||
# -*- test-case-name: twisted.cred.test.test_cred-*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
This module defines L{ICredentials}, an interface for objects that represent
|
||||
authentication credentials to provide, and also includes a number of useful
|
||||
implementations of that interface.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
from zope.interface import implementer, Interface
|
||||
|
||||
import base64
|
||||
import hmac
|
||||
import random
|
||||
import re
|
||||
import time
|
||||
|
||||
from binascii import hexlify
|
||||
from hashlib import md5
|
||||
|
||||
from twisted.python.randbytes import secureRandom
|
||||
from twisted.python.compat import networkString, nativeString
|
||||
from twisted.python.compat import intToBytes, unicode
|
||||
from twisted.cred._digest import calcResponse, calcHA1, calcHA2
|
||||
from twisted.cred import error
|
||||
|
||||
|
||||
|
||||
class ICredentials(Interface):
|
||||
"""
|
||||
I check credentials.
|
||||
|
||||
Implementors _must_ specify which sub-interfaces of ICredentials
|
||||
to which it conforms, using L{zope.interface.declarations.implementer}.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class IUsernameDigestHash(ICredentials):
|
||||
"""
|
||||
This credential is used when a CredentialChecker has access to the hash
|
||||
of the username:realm:password as in an Apache .htdigest file.
|
||||
"""
|
||||
def checkHash(digestHash):
|
||||
"""
|
||||
@param digestHash: The hashed username:realm:password to check against.
|
||||
|
||||
@return: C{True} if the credentials represented by this object match
|
||||
the given hash, C{False} if they do not, or a L{Deferred} which
|
||||
will be called back with one of these values.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class IUsernameHashedPassword(ICredentials):
|
||||
"""
|
||||
I encapsulate a username and a hashed password.
|
||||
|
||||
This credential is used when a hashed password is received from the
|
||||
party requesting authentication. CredentialCheckers which check this
|
||||
kind of credential must store the passwords in plaintext (or as
|
||||
password-equivalent hashes) form so that they can be hashed in a manner
|
||||
appropriate for the particular credentials class.
|
||||
|
||||
@type username: L{bytes}
|
||||
@ivar username: The username associated with these credentials.
|
||||
"""
|
||||
|
||||
def checkPassword(password):
|
||||
"""
|
||||
Validate these credentials against the correct password.
|
||||
|
||||
@type password: L{bytes}
|
||||
@param password: The correct, plaintext password against which to
|
||||
check.
|
||||
|
||||
@rtype: C{bool} or L{Deferred}
|
||||
@return: C{True} if the credentials represented by this object match the
|
||||
given password, C{False} if they do not, or a L{Deferred} which will
|
||||
be called back with one of these values.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class IUsernamePassword(ICredentials):
|
||||
"""
|
||||
I encapsulate a username and a plaintext password.
|
||||
|
||||
This encapsulates the case where the password received over the network
|
||||
has been hashed with the identity function (That is, not at all). The
|
||||
CredentialsChecker may store the password in whatever format it desires,
|
||||
it need only transform the stored password in a similar way before
|
||||
performing the comparison.
|
||||
|
||||
@type username: L{bytes}
|
||||
@ivar username: The username associated with these credentials.
|
||||
|
||||
@type password: L{bytes}
|
||||
@ivar password: The password associated with these credentials.
|
||||
"""
|
||||
|
||||
def checkPassword(password):
|
||||
"""
|
||||
Validate these credentials against the correct password.
|
||||
|
||||
@type password: L{bytes}
|
||||
@param password: The correct, plaintext password against which to
|
||||
check.
|
||||
|
||||
@rtype: C{bool} or L{Deferred}
|
||||
@return: C{True} if the credentials represented by this object match the
|
||||
given password, C{False} if they do not, or a L{Deferred} which will
|
||||
be called back with one of these values.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class IAnonymous(ICredentials):
|
||||
"""
|
||||
I am an explicitly anonymous request for access.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
@implementer(IUsernameHashedPassword, IUsernameDigestHash)
|
||||
class DigestedCredentials(object):
|
||||
"""
|
||||
Yet Another Simple HTTP Digest authentication scheme.
|
||||
"""
|
||||
|
||||
def __init__(self, username, method, realm, fields):
|
||||
self.username = username
|
||||
self.method = method
|
||||
self.realm = realm
|
||||
self.fields = fields
|
||||
|
||||
|
||||
def checkPassword(self, password):
|
||||
"""
|
||||
Verify that the credentials represented by this object agree with the
|
||||
given plaintext C{password} by hashing C{password} in the same way the
|
||||
response hash represented by this object was generated and comparing
|
||||
the results.
|
||||
"""
|
||||
response = self.fields.get('response')
|
||||
uri = self.fields.get('uri')
|
||||
nonce = self.fields.get('nonce')
|
||||
cnonce = self.fields.get('cnonce')
|
||||
nc = self.fields.get('nc')
|
||||
algo = self.fields.get('algorithm', b'md5').lower()
|
||||
qop = self.fields.get('qop', b'auth')
|
||||
|
||||
expected = calcResponse(
|
||||
calcHA1(algo, self.username, self.realm, password, nonce, cnonce),
|
||||
calcHA2(algo, self.method, uri, qop, None),
|
||||
algo, nonce, nc, cnonce, qop)
|
||||
|
||||
return expected == response
|
||||
|
||||
|
||||
def checkHash(self, digestHash):
|
||||
"""
|
||||
Verify that the credentials represented by this object agree with the
|
||||
credentials represented by the I{H(A1)} given in C{digestHash}.
|
||||
|
||||
@param digestHash: A precomputed H(A1) value based on the username,
|
||||
realm, and password associate with this credentials object.
|
||||
"""
|
||||
response = self.fields.get('response')
|
||||
uri = self.fields.get('uri')
|
||||
nonce = self.fields.get('nonce')
|
||||
cnonce = self.fields.get('cnonce')
|
||||
nc = self.fields.get('nc')
|
||||
algo = self.fields.get('algorithm', b'md5').lower()
|
||||
qop = self.fields.get('qop', b'auth')
|
||||
|
||||
expected = calcResponse(
|
||||
calcHA1(algo, None, None, None, nonce, cnonce, preHA1=digestHash),
|
||||
calcHA2(algo, self.method, uri, qop, None),
|
||||
algo, nonce, nc, cnonce, qop)
|
||||
|
||||
return expected == response
|
||||
|
||||
|
||||
|
||||
class DigestCredentialFactory(object):
|
||||
"""
|
||||
Support for RFC2617 HTTP Digest Authentication
|
||||
|
||||
@cvar CHALLENGE_LIFETIME_SECS: The number of seconds for which an
|
||||
opaque should be valid.
|
||||
|
||||
@type privateKey: L{bytes}
|
||||
@ivar privateKey: A random string used for generating the secure opaque.
|
||||
|
||||
@type algorithm: L{bytes}
|
||||
@param algorithm: Case insensitive string specifying the hash algorithm to
|
||||
use. Must be either C{'md5'} or C{'sha'}. C{'md5-sess'} is B{not}
|
||||
supported.
|
||||
|
||||
@type authenticationRealm: L{bytes}
|
||||
@param authenticationRealm: case sensitive string that specifies the realm
|
||||
portion of the challenge
|
||||
"""
|
||||
|
||||
_parseparts = re.compile(
|
||||
b'([^= ]+)' # The key
|
||||
b'=' # Conventional key/value separator (literal)
|
||||
b'(?:' # Group together a couple options
|
||||
b'"([^"]*)"' # A quoted string of length 0 or more
|
||||
b'|' # The other option in the group is coming
|
||||
b'([^,]+)' # An unquoted string of length 1 or more, up to a comma
|
||||
b')' # That non-matching group ends
|
||||
b',?') # There might be a comma at the end (none on last pair)
|
||||
|
||||
CHALLENGE_LIFETIME_SECS = 15 * 60 # 15 minutes
|
||||
|
||||
scheme = b"digest"
|
||||
|
||||
def __init__(self, algorithm, authenticationRealm):
|
||||
self.algorithm = algorithm
|
||||
self.authenticationRealm = authenticationRealm
|
||||
self.privateKey = secureRandom(12)
|
||||
|
||||
|
||||
def getChallenge(self, address):
|
||||
"""
|
||||
Generate the challenge for use in the WWW-Authenticate header.
|
||||
|
||||
@param address: The client address to which this challenge is being
|
||||
sent.
|
||||
|
||||
@return: The L{dict} that can be used to generate a WWW-Authenticate
|
||||
header.
|
||||
"""
|
||||
c = self._generateNonce()
|
||||
o = self._generateOpaque(c, address)
|
||||
|
||||
return {'nonce': c,
|
||||
'opaque': o,
|
||||
'qop': b'auth',
|
||||
'algorithm': self.algorithm,
|
||||
'realm': self.authenticationRealm}
|
||||
|
||||
|
||||
def _generateNonce(self):
|
||||
"""
|
||||
Create a random value suitable for use as the nonce parameter of a
|
||||
WWW-Authenticate challenge.
|
||||
|
||||
@rtype: L{bytes}
|
||||
"""
|
||||
return hexlify(secureRandom(12))
|
||||
|
||||
|
||||
def _getTime(self):
|
||||
"""
|
||||
Parameterize the time based seed used in C{_generateOpaque}
|
||||
so we can deterministically unittest it's behavior.
|
||||
"""
|
||||
return time.time()
|
||||
|
||||
|
||||
def _generateOpaque(self, nonce, clientip):
|
||||
"""
|
||||
Generate an opaque to be returned to the client. This is a unique
|
||||
string that can be returned to us and verified.
|
||||
"""
|
||||
# Now, what we do is encode the nonce, client ip and a timestamp in the
|
||||
# opaque value with a suitable digest.
|
||||
now = intToBytes(int(self._getTime()))
|
||||
|
||||
if not clientip:
|
||||
clientip = b''
|
||||
elif isinstance(clientip, unicode):
|
||||
clientip = clientip.encode('ascii')
|
||||
|
||||
key = b",".join((nonce, clientip, now))
|
||||
digest = hexlify(md5(key + self.privateKey).digest())
|
||||
ekey = base64.b64encode(key)
|
||||
return b"-".join((digest, ekey.replace(b'\n', b'')))
|
||||
|
||||
|
||||
def _verifyOpaque(self, opaque, nonce, clientip):
|
||||
"""
|
||||
Given the opaque and nonce from the request, as well as the client IP
|
||||
that made the request, verify that the opaque was generated by us.
|
||||
And that it's not too old.
|
||||
|
||||
@param opaque: The opaque value from the Digest response
|
||||
@param nonce: The nonce value from the Digest response
|
||||
@param clientip: The remote IP address of the client making the request
|
||||
or L{None} if the request was submitted over a channel where this
|
||||
does not make sense.
|
||||
|
||||
@return: C{True} if the opaque was successfully verified.
|
||||
|
||||
@raise error.LoginFailed: if C{opaque} could not be parsed or
|
||||
contained the wrong values.
|
||||
"""
|
||||
# First split the digest from the key
|
||||
opaqueParts = opaque.split(b'-')
|
||||
if len(opaqueParts) != 2:
|
||||
raise error.LoginFailed('Invalid response, invalid opaque value')
|
||||
|
||||
if not clientip:
|
||||
clientip = b''
|
||||
elif isinstance(clientip, unicode):
|
||||
clientip = clientip.encode('ascii')
|
||||
|
||||
# Verify the key
|
||||
key = base64.b64decode(opaqueParts[1])
|
||||
keyParts = key.split(b',')
|
||||
|
||||
if len(keyParts) != 3:
|
||||
raise error.LoginFailed('Invalid response, invalid opaque value')
|
||||
|
||||
if keyParts[0] != nonce:
|
||||
raise error.LoginFailed(
|
||||
'Invalid response, incompatible opaque/nonce values')
|
||||
|
||||
if keyParts[1] != clientip:
|
||||
raise error.LoginFailed(
|
||||
'Invalid response, incompatible opaque/client values')
|
||||
|
||||
try:
|
||||
when = int(keyParts[2])
|
||||
except ValueError:
|
||||
raise error.LoginFailed(
|
||||
'Invalid response, invalid opaque/time values')
|
||||
|
||||
if (int(self._getTime()) - when >
|
||||
DigestCredentialFactory.CHALLENGE_LIFETIME_SECS):
|
||||
|
||||
raise error.LoginFailed(
|
||||
'Invalid response, incompatible opaque/nonce too old')
|
||||
|
||||
# Verify the digest
|
||||
digest = hexlify(md5(key + self.privateKey).digest())
|
||||
if digest != opaqueParts[0]:
|
||||
raise error.LoginFailed('Invalid response, invalid opaque value')
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def decode(self, response, method, host):
|
||||
"""
|
||||
Decode the given response and attempt to generate a
|
||||
L{DigestedCredentials} from it.
|
||||
|
||||
@type response: L{bytes}
|
||||
@param response: A string of comma separated key=value pairs
|
||||
|
||||
@type method: L{bytes}
|
||||
@param method: The action requested to which this response is addressed
|
||||
(GET, POST, INVITE, OPTIONS, etc).
|
||||
|
||||
@type host: L{bytes}
|
||||
@param host: The address the request was sent from.
|
||||
|
||||
@raise error.LoginFailed: If the response does not contain a username,
|
||||
a nonce, an opaque, or if the opaque is invalid.
|
||||
|
||||
@return: L{DigestedCredentials}
|
||||
"""
|
||||
response = b' '.join(response.splitlines())
|
||||
parts = self._parseparts.findall(response)
|
||||
auth = {}
|
||||
for (key, bare, quoted) in parts:
|
||||
value = (quoted or bare).strip()
|
||||
auth[nativeString(key.strip())] = value
|
||||
|
||||
username = auth.get('username')
|
||||
if not username:
|
||||
raise error.LoginFailed('Invalid response, no username given.')
|
||||
|
||||
if 'opaque' not in auth:
|
||||
raise error.LoginFailed('Invalid response, no opaque given.')
|
||||
|
||||
if 'nonce' not in auth:
|
||||
raise error.LoginFailed('Invalid response, no nonce given.')
|
||||
|
||||
# Now verify the nonce/opaque values for this client
|
||||
if self._verifyOpaque(auth.get('opaque'), auth.get('nonce'), host):
|
||||
return DigestedCredentials(username,
|
||||
method,
|
||||
self.authenticationRealm,
|
||||
auth)
|
||||
|
||||
|
||||
|
||||
@implementer(IUsernameHashedPassword)
|
||||
class CramMD5Credentials(object):
|
||||
"""
|
||||
An encapsulation of some CramMD5 hashed credentials.
|
||||
|
||||
@ivar challenge: The challenge to be sent to the client.
|
||||
@type challenge: L{bytes}
|
||||
|
||||
@ivar response: The hashed response from the client.
|
||||
@type response: L{bytes}
|
||||
|
||||
@ivar username: The username from the response from the client.
|
||||
@type username: L{bytes} or L{None} if not yet provided.
|
||||
"""
|
||||
username = None
|
||||
challenge = b''
|
||||
response = b''
|
||||
|
||||
def __init__(self, host=None):
|
||||
self.host = host
|
||||
|
||||
|
||||
def getChallenge(self):
|
||||
if self.challenge:
|
||||
return self.challenge
|
||||
# The data encoded in the first ready response contains an
|
||||
# presumptively arbitrary string of random digits, a timestamp, and
|
||||
# the fully-qualified primary host name of the server. The syntax of
|
||||
# the unencoded form must correspond to that of an RFC 822 'msg-id'
|
||||
# [RFC822] as described in [POP3].
|
||||
# -- RFC 2195
|
||||
r = random.randrange(0x7fffffff)
|
||||
t = time.time()
|
||||
self.challenge = networkString('<%d.%d@%s>' % (
|
||||
r, t, nativeString(self.host) if self.host else None))
|
||||
return self.challenge
|
||||
|
||||
|
||||
def setResponse(self, response):
|
||||
self.username, self.response = response.split(None, 1)
|
||||
|
||||
|
||||
def moreChallenges(self):
|
||||
return False
|
||||
|
||||
|
||||
def checkPassword(self, password):
|
||||
verify = hexlify(hmac.HMAC(password, self.challenge).digest())
|
||||
return verify == self.response
|
||||
|
||||
|
||||
|
||||
@implementer(IUsernameHashedPassword)
|
||||
class UsernameHashedPassword:
|
||||
|
||||
def __init__(self, username, hashed):
|
||||
self.username = username
|
||||
self.hashed = hashed
|
||||
|
||||
def checkPassword(self, password):
|
||||
return self.hashed == password
|
||||
|
||||
|
||||
|
||||
@implementer(IUsernamePassword)
|
||||
class UsernamePassword:
|
||||
|
||||
def __init__(self, username, password):
|
||||
self.username = username
|
||||
self.password = password
|
||||
|
||||
def checkPassword(self, password):
|
||||
return self.password == password
|
||||
|
||||
|
||||
|
||||
@implementer(IAnonymous)
|
||||
class Anonymous:
|
||||
pass
|
||||
|
||||
|
||||
|
||||
class ISSHPrivateKey(ICredentials):
|
||||
"""
|
||||
L{ISSHPrivateKey} credentials encapsulate an SSH public key to be checked
|
||||
against a user's private key.
|
||||
|
||||
@ivar username: The username associated with these credentials.
|
||||
@type username: L{bytes}
|
||||
|
||||
@ivar algName: The algorithm name for the blob.
|
||||
@type algName: L{bytes}
|
||||
|
||||
@ivar blob: The public key blob as sent by the client.
|
||||
@type blob: L{bytes}
|
||||
|
||||
@ivar sigData: The data the signature was made from.
|
||||
@type sigData: L{bytes}
|
||||
|
||||
@ivar signature: The signed data. This is checked to verify that the user
|
||||
owns the private key.
|
||||
@type signature: L{bytes} or L{None}
|
||||
"""
|
||||
|
||||
|
||||
|
||||
@implementer(ISSHPrivateKey)
|
||||
class SSHPrivateKey:
|
||||
def __init__(self, username, algName, blob, sigData, signature):
|
||||
self.username = username
|
||||
self.algName = algName
|
||||
self.blob = blob
|
||||
self.sigData = sigData
|
||||
self.signature = signature
|
||||
@@ -0,0 +1,124 @@
|
||||
# -*- test-case-name: twisted.cred.test.test_cred -*-
|
||||
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
The point of integration of application and authentication.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
from twisted.internet import defer
|
||||
from twisted.internet.defer import maybeDeferred
|
||||
from twisted.python import failure, reflect
|
||||
from twisted.cred import error
|
||||
from zope.interface import providedBy, Interface
|
||||
|
||||
|
||||
class IRealm(Interface):
|
||||
"""
|
||||
The realm connects application-specific objects to the
|
||||
authentication system.
|
||||
"""
|
||||
def requestAvatar(avatarId, mind, *interfaces):
|
||||
"""
|
||||
Return avatar which provides one of the given interfaces.
|
||||
|
||||
@param avatarId: a string that identifies an avatar, as returned by
|
||||
L{ICredentialsChecker.requestAvatarId<twisted.cred.checkers.ICredentialsChecker.requestAvatarId>}
|
||||
(via a Deferred). Alternatively, it may be
|
||||
C{twisted.cred.checkers.ANONYMOUS}.
|
||||
@param mind: usually None. See the description of mind in
|
||||
L{Portal.login}.
|
||||
@param interfaces: the interface(s) the returned avatar should
|
||||
implement, e.g. C{IMailAccount}. See the description of
|
||||
L{Portal.login}.
|
||||
|
||||
@returns: a deferred which will fire a tuple of (interface,
|
||||
avatarAspect, logout), or the tuple itself. The interface will be
|
||||
one of the interfaces passed in the 'interfaces' argument. The
|
||||
'avatarAspect' will implement that interface. The 'logout' object
|
||||
is a callable which will detach the mind from the avatar.
|
||||
"""
|
||||
|
||||
|
||||
class Portal(object):
|
||||
"""
|
||||
A mediator between clients and a realm.
|
||||
|
||||
A portal is associated with one Realm and zero or more credentials checkers.
|
||||
When a login is attempted, the portal finds the appropriate credentials
|
||||
checker for the credentials given, invokes it, and if the credentials are
|
||||
valid, retrieves the appropriate avatar from the Realm.
|
||||
|
||||
This class is not intended to be subclassed. Customization should be done
|
||||
in the realm object and in the credentials checker objects.
|
||||
"""
|
||||
def __init__(self, realm, checkers=()):
|
||||
"""
|
||||
Create a Portal to a L{IRealm}.
|
||||
"""
|
||||
self.realm = realm
|
||||
self.checkers = {}
|
||||
for checker in checkers:
|
||||
self.registerChecker(checker)
|
||||
|
||||
|
||||
def listCredentialsInterfaces(self):
|
||||
"""
|
||||
Return list of credentials interfaces that can be used to login.
|
||||
"""
|
||||
return list(self.checkers.keys())
|
||||
|
||||
|
||||
def registerChecker(self, checker, *credentialInterfaces):
|
||||
if not credentialInterfaces:
|
||||
credentialInterfaces = checker.credentialInterfaces
|
||||
for credentialInterface in credentialInterfaces:
|
||||
self.checkers[credentialInterface] = checker
|
||||
|
||||
|
||||
def login(self, credentials, mind, *interfaces):
|
||||
"""
|
||||
@param credentials: an implementor of
|
||||
L{twisted.cred.credentials.ICredentials}
|
||||
|
||||
@param mind: an object which implements a client-side interface for
|
||||
your particular realm. In many cases, this may be None, so if the
|
||||
word 'mind' confuses you, just ignore it.
|
||||
|
||||
@param interfaces: list of interfaces for the perspective that the mind
|
||||
wishes to attach to. Usually, this will be only one interface, for
|
||||
example IMailAccount. For highly dynamic protocols, however, this
|
||||
may be a list like (IMailAccount, IUserChooser, IServiceInfo). To
|
||||
expand: if we are speaking to the system over IMAP, any information
|
||||
that will be relayed to the user MUST be returned as an
|
||||
IMailAccount implementor; IMAP clients would not be able to
|
||||
understand anything else. Any information about unusual status
|
||||
would have to be relayed as a single mail message in an
|
||||
otherwise-empty mailbox. However, in a web-based mail system, or a
|
||||
PB-based client, the ``mind'' object inside the web server
|
||||
(implemented with a dynamic page-viewing mechanism such as a
|
||||
Twisted Web Resource) or on the user's client program may be
|
||||
intelligent enough to respond to several ``server''-side
|
||||
interfaces.
|
||||
|
||||
@return: A deferred which will fire a tuple of (interface,
|
||||
avatarAspect, logout). The interface will be one of the interfaces
|
||||
passed in the 'interfaces' argument. The 'avatarAspect' will
|
||||
implement that interface. The 'logout' object is a callable which
|
||||
will detach the mind from the avatar. It must be called when the
|
||||
user has conceptually disconnected from the service. Although in
|
||||
some cases this will not be in connectionLost (such as in a
|
||||
web-based session), it will always be at the end of a user's
|
||||
interactive session.
|
||||
"""
|
||||
for i in self.checkers:
|
||||
if i.providedBy(credentials):
|
||||
return maybeDeferred(self.checkers[i].requestAvatarId, credentials
|
||||
).addCallback(self.realm.requestAvatar, mind, *interfaces
|
||||
)
|
||||
ifac = providedBy(credentials)
|
||||
return defer.fail(failure.Failure(error.UnhandledCredentials(
|
||||
"No checker for %s" % ', '.join(map(reflect.qual, ifac)))))
|
||||
@@ -0,0 +1,441 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.cred}, now with 30% more starch.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
from zope.interface import implementer, Interface
|
||||
|
||||
from binascii import hexlify, unhexlify
|
||||
|
||||
from twisted.trial import unittest
|
||||
from twisted.python.compat import nativeString, networkString
|
||||
from twisted.python import components
|
||||
from twisted.internet import defer
|
||||
from twisted.cred import checkers, credentials, portal, error
|
||||
|
||||
try:
|
||||
from crypt import crypt
|
||||
except ImportError:
|
||||
crypt = None
|
||||
|
||||
|
||||
|
||||
class ITestable(Interface):
|
||||
"""
|
||||
An interface for a theoretical protocol.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
|
||||
class TestAvatar(object):
|
||||
"""
|
||||
A test avatar.
|
||||
"""
|
||||
def __init__(self, name):
|
||||
self.name = name
|
||||
self.loggedIn = False
|
||||
self.loggedOut = False
|
||||
|
||||
|
||||
def login(self):
|
||||
assert not self.loggedIn
|
||||
self.loggedIn = True
|
||||
|
||||
|
||||
def logout(self):
|
||||
self.loggedOut = True
|
||||
|
||||
|
||||
|
||||
@implementer(ITestable)
|
||||
class Testable(components.Adapter):
|
||||
"""
|
||||
A theoretical protocol for testing.
|
||||
"""
|
||||
pass
|
||||
|
||||
components.registerAdapter(Testable, TestAvatar, ITestable)
|
||||
|
||||
|
||||
|
||||
class IDerivedCredentials(credentials.IUsernamePassword):
|
||||
pass
|
||||
|
||||
|
||||
|
||||
@implementer(IDerivedCredentials, ITestable)
|
||||
class DerivedCredentials(object):
|
||||
|
||||
def __init__(self, username, password):
|
||||
self.username = username
|
||||
self.password = password
|
||||
|
||||
|
||||
def checkPassword(self, password):
|
||||
return password == self.password
|
||||
|
||||
|
||||
|
||||
@implementer(portal.IRealm)
|
||||
class TestRealm(object):
|
||||
"""
|
||||
A basic test realm.
|
||||
"""
|
||||
def __init__(self):
|
||||
self.avatars = {}
|
||||
|
||||
|
||||
def requestAvatar(self, avatarId, mind, *interfaces):
|
||||
if avatarId in self.avatars:
|
||||
avatar = self.avatars[avatarId]
|
||||
else:
|
||||
avatar = TestAvatar(avatarId)
|
||||
self.avatars[avatarId] = avatar
|
||||
avatar.login()
|
||||
return (interfaces[0], interfaces[0](avatar),
|
||||
avatar.logout)
|
||||
|
||||
|
||||
|
||||
class CredTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for the meat of L{twisted.cred} -- realms, portals, avatars, and
|
||||
checkers.
|
||||
"""
|
||||
def setUp(self):
|
||||
self.realm = TestRealm()
|
||||
self.portal = portal.Portal(self.realm)
|
||||
self.checker = checkers.InMemoryUsernamePasswordDatabaseDontUse()
|
||||
self.checker.addUser(b"bob", b"hello")
|
||||
self.portal.registerChecker(self.checker)
|
||||
|
||||
|
||||
def test_listCheckers(self):
|
||||
"""
|
||||
The checkers in a portal can check only certain types of credentials.
|
||||
Since this portal has
|
||||
L{checkers.InMemoryUsernamePasswordDatabaseDontUse} registered, it
|
||||
"""
|
||||
expected = [credentials.IUsernamePassword,
|
||||
credentials.IUsernameHashedPassword]
|
||||
got = self.portal.listCredentialsInterfaces()
|
||||
self.assertEqual(sorted(got), sorted(expected))
|
||||
|
||||
|
||||
def test_basicLogin(self):
|
||||
"""
|
||||
Calling C{login} on a portal with correct credentials and an interface
|
||||
that the portal's realm supports works.
|
||||
"""
|
||||
login = self.successResultOf(self.portal.login(
|
||||
credentials.UsernamePassword(b"bob", b"hello"), self, ITestable))
|
||||
iface, impl, logout = login
|
||||
|
||||
# whitebox
|
||||
self.assertEqual(iface, ITestable)
|
||||
self.assertTrue(iface.providedBy(impl),
|
||||
"%s does not implement %s" % (impl, iface))
|
||||
|
||||
# greybox
|
||||
self.assertTrue(impl.original.loggedIn)
|
||||
self.assertTrue(not impl.original.loggedOut)
|
||||
logout()
|
||||
self.assertTrue(impl.original.loggedOut)
|
||||
|
||||
|
||||
def test_derivedInterface(self):
|
||||
"""
|
||||
Logging in with correct derived credentials and an interface
|
||||
that the portal's realm supports works.
|
||||
"""
|
||||
login = self.successResultOf(self.portal.login(
|
||||
DerivedCredentials(b"bob", b"hello"), self, ITestable))
|
||||
iface, impl, logout = login
|
||||
|
||||
# whitebox
|
||||
self.assertEqual(iface, ITestable)
|
||||
self.assertTrue(iface.providedBy(impl),
|
||||
"%s does not implement %s" % (impl, iface))
|
||||
|
||||
# greybox
|
||||
self.assertTrue(impl.original.loggedIn)
|
||||
self.assertTrue(not impl.original.loggedOut)
|
||||
logout()
|
||||
self.assertTrue(impl.original.loggedOut)
|
||||
|
||||
|
||||
def test_failedLoginPassword(self):
|
||||
"""
|
||||
Calling C{login} with incorrect credentials (in this case a wrong
|
||||
password) causes L{error.UnauthorizedLogin} to be raised.
|
||||
"""
|
||||
login = self.failureResultOf(self.portal.login(
|
||||
credentials.UsernamePassword(b"bob", b"h3llo"), self, ITestable))
|
||||
self.assertTrue(login)
|
||||
self.assertEqual(error.UnauthorizedLogin, login.type)
|
||||
|
||||
|
||||
def test_failedLoginName(self):
|
||||
"""
|
||||
Calling C{login} with incorrect credentials (in this case no known
|
||||
user) causes L{error.UnauthorizedLogin} to be raised.
|
||||
"""
|
||||
login = self.failureResultOf(self.portal.login(
|
||||
credentials.UsernamePassword(b"jay", b"hello"), self, ITestable))
|
||||
self.assertTrue(login)
|
||||
self.assertEqual(error.UnauthorizedLogin, login.type)
|
||||
|
||||
|
||||
|
||||
class OnDiskDatabaseTests(unittest.TestCase):
|
||||
users = [
|
||||
(b'user1', b'pass1'),
|
||||
(b'user2', b'pass2'),
|
||||
(b'user3', b'pass3'),
|
||||
]
|
||||
|
||||
def setUp(self):
|
||||
self.dbfile = self.mktemp()
|
||||
with open(self.dbfile, 'wb') as f:
|
||||
for (u, p) in self.users:
|
||||
f.write(u + b":" + p + b"\n")
|
||||
|
||||
|
||||
def test_getUserNonexistentDatabase(self):
|
||||
"""
|
||||
A missing db file will cause a permanent rejection of authorization
|
||||
attempts.
|
||||
"""
|
||||
self.db = checkers.FilePasswordDB('test_thisbetternoteverexist.db')
|
||||
|
||||
self.assertRaises(error.UnauthorizedLogin, self.db.getUser, 'user')
|
||||
|
||||
|
||||
def testUserLookup(self):
|
||||
self.db = checkers.FilePasswordDB(self.dbfile)
|
||||
for (u, p) in self.users:
|
||||
self.assertRaises(KeyError, self.db.getUser, u.upper())
|
||||
self.assertEqual(self.db.getUser(u), (u, p))
|
||||
|
||||
|
||||
def testCaseInSensitivity(self):
|
||||
self.db = checkers.FilePasswordDB(self.dbfile, caseSensitive=False)
|
||||
for (u, p) in self.users:
|
||||
self.assertEqual(self.db.getUser(u.upper()), (u, p))
|
||||
|
||||
|
||||
def testRequestAvatarId(self):
|
||||
self.db = checkers.FilePasswordDB(self.dbfile)
|
||||
creds = [credentials.UsernamePassword(u, p) for u, p in self.users]
|
||||
d = defer.gatherResults(
|
||||
[defer.maybeDeferred(self.db.requestAvatarId, c) for c in creds])
|
||||
d.addCallback(self.assertEqual, [u for u, p in self.users])
|
||||
return d
|
||||
|
||||
|
||||
def testRequestAvatarId_hashed(self):
|
||||
self.db = checkers.FilePasswordDB(self.dbfile)
|
||||
creds = [credentials.UsernameHashedPassword(u, p)
|
||||
for u, p in self.users]
|
||||
d = defer.gatherResults(
|
||||
[defer.maybeDeferred(self.db.requestAvatarId, c) for c in creds])
|
||||
d.addCallback(self.assertEqual, [u for u, p in self.users])
|
||||
return d
|
||||
|
||||
|
||||
|
||||
class HashedPasswordOnDiskDatabaseTests(unittest.TestCase):
|
||||
users = [
|
||||
(b'user1', b'pass1'),
|
||||
(b'user2', b'pass2'),
|
||||
(b'user3', b'pass3'),
|
||||
]
|
||||
|
||||
def setUp(self):
|
||||
dbfile = self.mktemp()
|
||||
self.db = checkers.FilePasswordDB(dbfile, hash=self.hash)
|
||||
with open(dbfile, 'wb') as f:
|
||||
for (u, p) in self.users:
|
||||
f.write(u + b":" + self.hash(u, p, u[:2]) + b"\n")
|
||||
|
||||
r = TestRealm()
|
||||
self.port = portal.Portal(r)
|
||||
self.port.registerChecker(self.db)
|
||||
|
||||
|
||||
def hash(self, u, p, s):
|
||||
return networkString(crypt(nativeString(p), nativeString(s)))
|
||||
|
||||
|
||||
def testGoodCredentials(self):
|
||||
goodCreds = [credentials.UsernamePassword(u, p) for u, p in self.users]
|
||||
d = defer.gatherResults([self.db.requestAvatarId(c)
|
||||
for c in goodCreds])
|
||||
d.addCallback(self.assertEqual, [u for u, p in self.users])
|
||||
return d
|
||||
|
||||
|
||||
def testGoodCredentials_login(self):
|
||||
goodCreds = [credentials.UsernamePassword(u, p) for u, p in self.users]
|
||||
d = defer.gatherResults([self.port.login(c, None, ITestable)
|
||||
for c in goodCreds])
|
||||
d.addCallback(lambda x: [a.original.name for i, a, l in x])
|
||||
d.addCallback(self.assertEqual, [u for u, p in self.users])
|
||||
return d
|
||||
|
||||
|
||||
def testBadCredentials(self):
|
||||
badCreds = [credentials.UsernamePassword(u, 'wrong password')
|
||||
for u, p in self.users]
|
||||
d = defer.DeferredList([self.port.login(c, None, ITestable)
|
||||
for c in badCreds], consumeErrors=True)
|
||||
d.addCallback(self._assertFailures, error.UnauthorizedLogin)
|
||||
return d
|
||||
|
||||
|
||||
def testHashedCredentials(self):
|
||||
hashedCreds = [credentials.UsernameHashedPassword(
|
||||
u, self.hash(None, p, u[:2])) for u, p in self.users]
|
||||
d = defer.DeferredList([self.port.login(c, None, ITestable)
|
||||
for c in hashedCreds], consumeErrors=True)
|
||||
d.addCallback(self._assertFailures, error.UnhandledCredentials)
|
||||
return d
|
||||
|
||||
|
||||
def _assertFailures(self, failures, *expectedFailures):
|
||||
for flag, failure in failures:
|
||||
self.assertEqual(flag, defer.FAILURE)
|
||||
failure.trap(*expectedFailures)
|
||||
return None
|
||||
|
||||
if crypt is None:
|
||||
skip = "crypt module not available"
|
||||
|
||||
|
||||
|
||||
class CheckersMixin(object):
|
||||
"""
|
||||
L{unittest.TestCase} mixin for testing that some checkers accept
|
||||
and deny specified credentials.
|
||||
|
||||
Subclasses must provide
|
||||
- C{getCheckers} which returns a sequence of
|
||||
L{checkers.ICredentialChecker}
|
||||
- C{getGoodCredentials} which returns a list of 2-tuples of
|
||||
credential to check and avaterId to expect.
|
||||
- C{getBadCredentials} which returns a list of credentials
|
||||
which are expected to be unauthorized.
|
||||
"""
|
||||
|
||||
@defer.inlineCallbacks
|
||||
def test_positive(self):
|
||||
"""
|
||||
The given credentials are accepted by all the checkers, and give
|
||||
the expected C{avatarID}s
|
||||
"""
|
||||
for chk in self.getCheckers():
|
||||
for (cred, avatarId) in self.getGoodCredentials():
|
||||
r = yield chk.requestAvatarId(cred)
|
||||
self.assertEqual(r, avatarId)
|
||||
|
||||
|
||||
@defer.inlineCallbacks
|
||||
def test_negative(self):
|
||||
"""
|
||||
The given credentials are rejected by all the checkers.
|
||||
"""
|
||||
for chk in self.getCheckers():
|
||||
for cred in self.getBadCredentials():
|
||||
d = chk.requestAvatarId(cred)
|
||||
yield self.assertFailure(d, error.UnauthorizedLogin)
|
||||
|
||||
|
||||
|
||||
class HashlessFilePasswordDBMixin(object):
|
||||
credClass = credentials.UsernamePassword
|
||||
diskHash = None
|
||||
networkHash = staticmethod(lambda x: x)
|
||||
|
||||
_validCredentials = [
|
||||
(b'user1', b'password1'),
|
||||
(b'user2', b'password2'),
|
||||
(b'user3', b'password3')]
|
||||
|
||||
|
||||
def getGoodCredentials(self):
|
||||
for u, p in self._validCredentials:
|
||||
yield self.credClass(u, self.networkHash(p)), u
|
||||
|
||||
|
||||
def getBadCredentials(self):
|
||||
for u, p in [(b'user1', b'password3'),
|
||||
(b'user2', b'password1'),
|
||||
(b'bloof', b'blarf')]:
|
||||
yield self.credClass(u, self.networkHash(p))
|
||||
|
||||
|
||||
def getCheckers(self):
|
||||
diskHash = self.diskHash or (lambda x: x)
|
||||
hashCheck = self.diskHash and (lambda username, password,
|
||||
stored: self.diskHash(password))
|
||||
|
||||
for cache in True, False:
|
||||
fn = self.mktemp()
|
||||
with open(fn, 'wb') as fObj:
|
||||
for u, p in self._validCredentials:
|
||||
fObj.write(u + b":" + diskHash(p) + b"\n")
|
||||
yield checkers.FilePasswordDB(fn, cache=cache, hash=hashCheck)
|
||||
|
||||
fn = self.mktemp()
|
||||
with open(fn, 'wb') as fObj:
|
||||
for u, p in self._validCredentials:
|
||||
fObj.write(diskHash(p) + b' dingle dongle ' + u + b'\n')
|
||||
yield checkers.FilePasswordDB(fn, b' ', 3, 0,
|
||||
cache=cache, hash=hashCheck)
|
||||
|
||||
fn = self.mktemp()
|
||||
with open(fn, 'wb') as fObj:
|
||||
for u, p in self._validCredentials:
|
||||
fObj.write(b'zip,zap,' + u.title() + b',zup,'\
|
||||
+ diskHash(p) + b'\n',)
|
||||
yield checkers.FilePasswordDB(fn, b',', 2, 4, False,
|
||||
cache=cache, hash=hashCheck)
|
||||
|
||||
|
||||
|
||||
class LocallyHashedFilePasswordDBMixin(HashlessFilePasswordDBMixin):
|
||||
diskHash = staticmethod(lambda x: hexlify(x))
|
||||
|
||||
|
||||
|
||||
class NetworkHashedFilePasswordDBMixin(HashlessFilePasswordDBMixin):
|
||||
networkHash = staticmethod(lambda x: hexlify(x))
|
||||
|
||||
class credClass(credentials.UsernameHashedPassword):
|
||||
def checkPassword(self, password):
|
||||
return unhexlify(self.hashed) == password
|
||||
|
||||
|
||||
|
||||
class HashlessFilePasswordDBCheckerTests(HashlessFilePasswordDBMixin,
|
||||
CheckersMixin, unittest.TestCase):
|
||||
pass
|
||||
|
||||
|
||||
|
||||
class LocallyHashedFilePasswordDBCheckerTests(LocallyHashedFilePasswordDBMixin,
|
||||
CheckersMixin,
|
||||
unittest.TestCase):
|
||||
pass
|
||||
|
||||
|
||||
|
||||
class NetworkHashedFilePasswordDBCheckerTests(NetworkHashedFilePasswordDBMixin,
|
||||
CheckersMixin,
|
||||
unittest.TestCase):
|
||||
pass
|
||||
@@ -0,0 +1,698 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.cred._digest} and the associated bits in
|
||||
L{twisted.cred.credentials}.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
import base64
|
||||
|
||||
from binascii import hexlify
|
||||
from hashlib import md5, sha1
|
||||
|
||||
from zope.interface.verify import verifyObject
|
||||
from twisted.trial.unittest import TestCase
|
||||
from twisted.internet.address import IPv4Address
|
||||
from twisted.cred.error import LoginFailed
|
||||
from twisted.cred.credentials import calcHA1, calcHA2, IUsernameDigestHash
|
||||
from twisted.cred.credentials import calcResponse, DigestCredentialFactory
|
||||
from twisted.python.compat import networkString
|
||||
|
||||
def b64encode(s):
|
||||
return base64.b64encode(s).strip()
|
||||
|
||||
|
||||
|
||||
class FakeDigestCredentialFactory(DigestCredentialFactory):
|
||||
"""
|
||||
A Fake Digest Credential Factory that generates a predictable
|
||||
nonce and opaque
|
||||
"""
|
||||
def __init__(self, *args, **kwargs):
|
||||
super(FakeDigestCredentialFactory, self).__init__(*args, **kwargs)
|
||||
self.privateKey = b"0"
|
||||
|
||||
|
||||
def _generateNonce(self):
|
||||
"""
|
||||
Generate a static nonce
|
||||
"""
|
||||
return b'178288758716122392881254770685'
|
||||
|
||||
|
||||
def _getTime(self):
|
||||
"""
|
||||
Return a stable time
|
||||
"""
|
||||
return 0
|
||||
|
||||
|
||||
|
||||
class DigestAuthTests(TestCase):
|
||||
"""
|
||||
L{TestCase} mixin class which defines a number of tests for
|
||||
L{DigestCredentialFactory}. Because this mixin defines C{setUp}, it
|
||||
must be inherited before L{TestCase}.
|
||||
"""
|
||||
def setUp(self):
|
||||
"""
|
||||
Create a DigestCredentialFactory for testing
|
||||
"""
|
||||
self.username = b"foobar"
|
||||
self.password = b"bazquux"
|
||||
self.realm = b"test realm"
|
||||
self.algorithm = b"md5"
|
||||
self.cnonce = b"29fc54aa1641c6fa0e151419361c8f23"
|
||||
self.qop = b"auth"
|
||||
self.uri = b"/write/"
|
||||
self.clientAddress = IPv4Address('TCP', '10.2.3.4', 43125)
|
||||
self.method = b'GET'
|
||||
self.credentialFactory = DigestCredentialFactory(
|
||||
self.algorithm, self.realm)
|
||||
|
||||
|
||||
def test_MD5HashA1(self, _algorithm=b'md5', _hash=md5):
|
||||
"""
|
||||
L{calcHA1} accepts the C{'md5'} algorithm and returns an MD5 hash of
|
||||
its parameters, excluding the nonce and cnonce.
|
||||
"""
|
||||
nonce = b'abc123xyz'
|
||||
hashA1 = calcHA1(_algorithm, self.username, self.realm, self.password,
|
||||
nonce, self.cnonce)
|
||||
a1 = b":".join((self.username, self.realm, self.password))
|
||||
expected = hexlify(_hash(a1).digest())
|
||||
self.assertEqual(hashA1, expected)
|
||||
|
||||
|
||||
def test_MD5SessionHashA1(self):
|
||||
"""
|
||||
L{calcHA1} accepts the C{'md5-sess'} algorithm and returns an MD5 hash
|
||||
of its parameters, including the nonce and cnonce.
|
||||
"""
|
||||
nonce = b'xyz321abc'
|
||||
hashA1 = calcHA1(b'md5-sess', self.username, self.realm, self.password,
|
||||
nonce, self.cnonce)
|
||||
a1 = self.username + b':' + self.realm + b':' + self.password
|
||||
ha1 = hexlify(md5(a1).digest())
|
||||
a1 = ha1 + b':' + nonce + b':' + self.cnonce
|
||||
expected = hexlify(md5(a1).digest())
|
||||
self.assertEqual(hashA1, expected)
|
||||
|
||||
|
||||
def test_SHAHashA1(self):
|
||||
"""
|
||||
L{calcHA1} accepts the C{'sha'} algorithm and returns a SHA hash of its
|
||||
parameters, excluding the nonce and cnonce.
|
||||
"""
|
||||
self.test_MD5HashA1(b'sha', sha1)
|
||||
|
||||
|
||||
def test_MD5HashA2Auth(self, _algorithm=b'md5', _hash=md5):
|
||||
"""
|
||||
L{calcHA2} accepts the C{'md5'} algorithm and returns an MD5 hash of
|
||||
its arguments, excluding the entity hash for QOP other than
|
||||
C{'auth-int'}.
|
||||
"""
|
||||
method = b'GET'
|
||||
hashA2 = calcHA2(_algorithm, method, self.uri, b'auth', None)
|
||||
a2 = method + b':' + self.uri
|
||||
expected = hexlify(_hash(a2).digest())
|
||||
self.assertEqual(hashA2, expected)
|
||||
|
||||
|
||||
def test_MD5HashA2AuthInt(self, _algorithm=b'md5', _hash=md5):
|
||||
"""
|
||||
L{calcHA2} accepts the C{'md5'} algorithm and returns an MD5 hash of
|
||||
its arguments, including the entity hash for QOP of C{'auth-int'}.
|
||||
"""
|
||||
method = b'GET'
|
||||
hentity = b'foobarbaz'
|
||||
hashA2 = calcHA2(_algorithm, method, self.uri, b'auth-int', hentity)
|
||||
a2 = method + b':' + self.uri + b':' + hentity
|
||||
expected = hexlify(_hash(a2).digest())
|
||||
self.assertEqual(hashA2, expected)
|
||||
|
||||
|
||||
def test_MD5SessHashA2Auth(self):
|
||||
"""
|
||||
L{calcHA2} accepts the C{'md5-sess'} algorithm and QOP of C{'auth'} and
|
||||
returns the same value as it does for the C{'md5'} algorithm.
|
||||
"""
|
||||
self.test_MD5HashA2Auth(b'md5-sess')
|
||||
|
||||
|
||||
def test_MD5SessHashA2AuthInt(self):
|
||||
"""
|
||||
L{calcHA2} accepts the C{'md5-sess'} algorithm and QOP of C{'auth-int'}
|
||||
and returns the same value as it does for the C{'md5'} algorithm.
|
||||
"""
|
||||
self.test_MD5HashA2AuthInt(b'md5-sess')
|
||||
|
||||
|
||||
def test_SHAHashA2Auth(self):
|
||||
"""
|
||||
L{calcHA2} accepts the C{'sha'} algorithm and returns a SHA hash of
|
||||
its arguments, excluding the entity hash for QOP other than
|
||||
C{'auth-int'}.
|
||||
"""
|
||||
self.test_MD5HashA2Auth(b'sha', sha1)
|
||||
|
||||
|
||||
def test_SHAHashA2AuthInt(self):
|
||||
"""
|
||||
L{calcHA2} accepts the C{'sha'} algorithm and returns a SHA hash of
|
||||
its arguments, including the entity hash for QOP of C{'auth-int'}.
|
||||
"""
|
||||
self.test_MD5HashA2AuthInt(b'sha', sha1)
|
||||
|
||||
|
||||
def test_MD5HashResponse(self, _algorithm=b'md5', _hash=md5):
|
||||
"""
|
||||
L{calcResponse} accepts the C{'md5'} algorithm and returns an MD5 hash
|
||||
of its parameters, excluding the nonce count, client nonce, and QoP
|
||||
value if the nonce count and client nonce are L{None}
|
||||
"""
|
||||
hashA1 = b'abc123'
|
||||
hashA2 = b'789xyz'
|
||||
nonce = b'lmnopq'
|
||||
|
||||
response = hashA1 + b':' + nonce + b':' + hashA2
|
||||
expected = hexlify(_hash(response).digest())
|
||||
|
||||
digest = calcResponse(hashA1, hashA2, _algorithm, nonce, None, None,
|
||||
None)
|
||||
self.assertEqual(expected, digest)
|
||||
|
||||
|
||||
def test_MD5SessionHashResponse(self):
|
||||
"""
|
||||
L{calcResponse} accepts the C{'md5-sess'} algorithm and returns an MD5
|
||||
hash of its parameters, excluding the nonce count, client nonce, and
|
||||
QoP value if the nonce count and client nonce are L{None}
|
||||
"""
|
||||
self.test_MD5HashResponse(b'md5-sess')
|
||||
|
||||
|
||||
def test_SHAHashResponse(self):
|
||||
"""
|
||||
L{calcResponse} accepts the C{'sha'} algorithm and returns a SHA hash
|
||||
of its parameters, excluding the nonce count, client nonce, and QoP
|
||||
value if the nonce count and client nonce are L{None}
|
||||
"""
|
||||
self.test_MD5HashResponse(b'sha', sha1)
|
||||
|
||||
|
||||
def test_MD5HashResponseExtra(self, _algorithm=b'md5', _hash=md5):
|
||||
"""
|
||||
L{calcResponse} accepts the C{'md5'} algorithm and returns an MD5 hash
|
||||
of its parameters, including the nonce count, client nonce, and QoP
|
||||
value if they are specified.
|
||||
"""
|
||||
hashA1 = b'abc123'
|
||||
hashA2 = b'789xyz'
|
||||
nonce = b'lmnopq'
|
||||
nonceCount = b'00000004'
|
||||
clientNonce = b'abcxyz123'
|
||||
qop = b'auth'
|
||||
|
||||
response = hashA1 + b':' + nonce + b':' + nonceCount + b':' +\
|
||||
clientNonce + b':' + qop + b':' + hashA2
|
||||
expected = hexlify(_hash(response).digest())
|
||||
|
||||
digest = calcResponse(
|
||||
hashA1, hashA2, _algorithm, nonce, nonceCount, clientNonce, qop)
|
||||
self.assertEqual(expected, digest)
|
||||
|
||||
|
||||
def test_MD5SessionHashResponseExtra(self):
|
||||
"""
|
||||
L{calcResponse} accepts the C{'md5-sess'} algorithm and returns an MD5
|
||||
hash of its parameters, including the nonce count, client nonce, and
|
||||
QoP value if they are specified.
|
||||
"""
|
||||
self.test_MD5HashResponseExtra(b'md5-sess')
|
||||
|
||||
|
||||
def test_SHAHashResponseExtra(self):
|
||||
"""
|
||||
L{calcResponse} accepts the C{'sha'} algorithm and returns a SHA hash
|
||||
of its parameters, including the nonce count, client nonce, and QoP
|
||||
value if they are specified.
|
||||
"""
|
||||
self.test_MD5HashResponseExtra(b'sha', sha1)
|
||||
|
||||
|
||||
def formatResponse(self, quotes=True, **kw):
|
||||
"""
|
||||
Format all given keyword arguments and their values suitably for use as
|
||||
the value of an HTTP header.
|
||||
|
||||
@types quotes: C{bool}
|
||||
@param quotes: A flag indicating whether to quote the values of each
|
||||
field in the response.
|
||||
|
||||
@param **kw: Keywords and C{bytes} values which will be treated as field
|
||||
name/value pairs to include in the result.
|
||||
|
||||
@rtype: C{bytes}
|
||||
@return: The given fields formatted for use as an HTTP header value.
|
||||
"""
|
||||
if 'username' not in kw:
|
||||
kw['username'] = self.username
|
||||
if 'realm' not in kw:
|
||||
kw['realm'] = self.realm
|
||||
if 'algorithm' not in kw:
|
||||
kw['algorithm'] = self.algorithm
|
||||
if 'qop' not in kw:
|
||||
kw['qop'] = self.qop
|
||||
if 'cnonce' not in kw:
|
||||
kw['cnonce'] = self.cnonce
|
||||
if 'uri' not in kw:
|
||||
kw['uri'] = self.uri
|
||||
if quotes:
|
||||
quote = b'"'
|
||||
else:
|
||||
quote = b''
|
||||
|
||||
return b', '.join([
|
||||
b"".join((networkString(k), b"=", quote, v, quote))
|
||||
for (k, v)
|
||||
in kw.items()
|
||||
if v is not None])
|
||||
|
||||
|
||||
def getDigestResponse(self, challenge, ncount):
|
||||
"""
|
||||
Calculate the response for the given challenge
|
||||
"""
|
||||
nonce = challenge.get('nonce')
|
||||
algo = challenge.get('algorithm').lower()
|
||||
qop = challenge.get('qop')
|
||||
|
||||
ha1 = calcHA1(
|
||||
algo, self.username, self.realm, self.password, nonce, self.cnonce)
|
||||
ha2 = calcHA2(algo, b"GET", self.uri, qop, None)
|
||||
expected = calcResponse(
|
||||
ha1, ha2, algo, nonce, ncount, self.cnonce, qop)
|
||||
return expected
|
||||
|
||||
|
||||
def test_response(self, quotes=True):
|
||||
"""
|
||||
L{DigestCredentialFactory.decode} accepts a digest challenge response
|
||||
and parses it into an L{IUsernameHashedPassword} provider.
|
||||
"""
|
||||
challenge = self.credentialFactory.getChallenge(
|
||||
self.clientAddress.host)
|
||||
|
||||
nc = b"00000001"
|
||||
clientResponse = self.formatResponse(
|
||||
quotes=quotes,
|
||||
nonce=challenge['nonce'],
|
||||
response=self.getDigestResponse(challenge, nc),
|
||||
nc=nc,
|
||||
opaque=challenge['opaque'])
|
||||
creds = self.credentialFactory.decode(
|
||||
clientResponse, self.method, self.clientAddress.host)
|
||||
self.assertTrue(creds.checkPassword(self.password))
|
||||
self.assertFalse(creds.checkPassword(self.password + b'wrong'))
|
||||
|
||||
|
||||
def test_responseWithoutQuotes(self):
|
||||
"""
|
||||
L{DigestCredentialFactory.decode} accepts a digest challenge response
|
||||
which does not quote the values of its fields and parses it into an
|
||||
L{IUsernameHashedPassword} provider in the same way it would a
|
||||
response which included quoted field values.
|
||||
"""
|
||||
self.test_response(False)
|
||||
|
||||
|
||||
def test_responseWithCommaURI(self):
|
||||
"""
|
||||
L{DigestCredentialFactory.decode} accepts a digest challenge response
|
||||
which quotes the values of its fields and includes a C{b","} in the URI
|
||||
field.
|
||||
"""
|
||||
self.uri = b"/some,path/"
|
||||
self.test_response(True)
|
||||
|
||||
|
||||
def test_caseInsensitiveAlgorithm(self):
|
||||
"""
|
||||
The case of the algorithm value in the response is ignored when
|
||||
checking the credentials.
|
||||
"""
|
||||
self.algorithm = b'MD5'
|
||||
self.test_response()
|
||||
|
||||
|
||||
def test_md5DefaultAlgorithm(self):
|
||||
"""
|
||||
The algorithm defaults to MD5 if it is not supplied in the response.
|
||||
"""
|
||||
self.algorithm = None
|
||||
self.test_response()
|
||||
|
||||
|
||||
def test_responseWithoutClientIP(self):
|
||||
"""
|
||||
L{DigestCredentialFactory.decode} accepts a digest challenge response
|
||||
even if the client address it is passed is L{None}.
|
||||
"""
|
||||
challenge = self.credentialFactory.getChallenge(None)
|
||||
|
||||
nc = b"00000001"
|
||||
clientResponse = self.formatResponse(
|
||||
nonce=challenge['nonce'],
|
||||
response=self.getDigestResponse(challenge, nc),
|
||||
nc=nc,
|
||||
opaque=challenge['opaque'])
|
||||
creds = self.credentialFactory.decode(clientResponse, self.method,
|
||||
None)
|
||||
self.assertTrue(creds.checkPassword(self.password))
|
||||
self.assertFalse(creds.checkPassword(self.password + b'wrong'))
|
||||
|
||||
|
||||
def test_multiResponse(self):
|
||||
"""
|
||||
L{DigestCredentialFactory.decode} handles multiple responses to a
|
||||
single challenge.
|
||||
"""
|
||||
challenge = self.credentialFactory.getChallenge(
|
||||
self.clientAddress.host)
|
||||
|
||||
nc = b"00000001"
|
||||
clientResponse = self.formatResponse(
|
||||
nonce=challenge['nonce'],
|
||||
response=self.getDigestResponse(challenge, nc),
|
||||
nc=nc,
|
||||
opaque=challenge['opaque'])
|
||||
|
||||
creds = self.credentialFactory.decode(clientResponse, self.method,
|
||||
self.clientAddress.host)
|
||||
self.assertTrue(creds.checkPassword(self.password))
|
||||
self.assertFalse(creds.checkPassword(self.password + b'wrong'))
|
||||
|
||||
nc = b"00000002"
|
||||
clientResponse = self.formatResponse(
|
||||
nonce=challenge['nonce'],
|
||||
response=self.getDigestResponse(challenge, nc),
|
||||
nc=nc,
|
||||
opaque=challenge['opaque'])
|
||||
|
||||
creds = self.credentialFactory.decode(clientResponse, self.method,
|
||||
self.clientAddress.host)
|
||||
self.assertTrue(creds.checkPassword(self.password))
|
||||
self.assertFalse(creds.checkPassword(self.password + b'wrong'))
|
||||
|
||||
|
||||
def test_failsWithDifferentMethod(self):
|
||||
"""
|
||||
L{DigestCredentialFactory.decode} returns an L{IUsernameHashedPassword}
|
||||
provider which rejects a correct password for the given user if the
|
||||
challenge response request is made using a different HTTP method than
|
||||
was used to request the initial challenge.
|
||||
"""
|
||||
challenge = self.credentialFactory.getChallenge(
|
||||
self.clientAddress.host)
|
||||
|
||||
nc = b"00000001"
|
||||
clientResponse = self.formatResponse(
|
||||
nonce=challenge['nonce'],
|
||||
response=self.getDigestResponse(challenge, nc),
|
||||
nc=nc,
|
||||
opaque=challenge['opaque'])
|
||||
creds = self.credentialFactory.decode(clientResponse, b'POST',
|
||||
self.clientAddress.host)
|
||||
self.assertFalse(creds.checkPassword(self.password))
|
||||
self.assertFalse(creds.checkPassword(self.password + b'wrong'))
|
||||
|
||||
|
||||
def test_noUsername(self):
|
||||
"""
|
||||
L{DigestCredentialFactory.decode} raises L{LoginFailed} if the response
|
||||
has no username field or if the username field is empty.
|
||||
"""
|
||||
# Check for no username
|
||||
e = self.assertRaises(
|
||||
LoginFailed,
|
||||
self.credentialFactory.decode,
|
||||
self.formatResponse(username=None),
|
||||
self.method, self.clientAddress.host)
|
||||
self.assertEqual(str(e), "Invalid response, no username given.")
|
||||
|
||||
# Check for an empty username
|
||||
e = self.assertRaises(
|
||||
LoginFailed,
|
||||
self.credentialFactory.decode,
|
||||
self.formatResponse(username=b""),
|
||||
self.method, self.clientAddress.host)
|
||||
self.assertEqual(str(e), "Invalid response, no username given.")
|
||||
|
||||
|
||||
def test_noNonce(self):
|
||||
"""
|
||||
L{DigestCredentialFactory.decode} raises L{LoginFailed} if the response
|
||||
has no nonce.
|
||||
"""
|
||||
e = self.assertRaises(
|
||||
LoginFailed,
|
||||
self.credentialFactory.decode,
|
||||
self.formatResponse(opaque=b"abc123"),
|
||||
self.method, self.clientAddress.host)
|
||||
self.assertEqual(str(e), "Invalid response, no nonce given.")
|
||||
|
||||
|
||||
def test_noOpaque(self):
|
||||
"""
|
||||
L{DigestCredentialFactory.decode} raises L{LoginFailed} if the response
|
||||
has no opaque.
|
||||
"""
|
||||
e = self.assertRaises(
|
||||
LoginFailed,
|
||||
self.credentialFactory.decode,
|
||||
self.formatResponse(),
|
||||
self.method, self.clientAddress.host)
|
||||
self.assertEqual(str(e), "Invalid response, no opaque given.")
|
||||
|
||||
|
||||
def test_checkHash(self):
|
||||
"""
|
||||
L{DigestCredentialFactory.decode} returns an L{IUsernameDigestHash}
|
||||
provider which can verify a hash of the form 'username:realm:password'.
|
||||
"""
|
||||
challenge = self.credentialFactory.getChallenge(
|
||||
self.clientAddress.host)
|
||||
|
||||
nc = b"00000001"
|
||||
clientResponse = self.formatResponse(
|
||||
nonce=challenge['nonce'],
|
||||
response=self.getDigestResponse(challenge, nc),
|
||||
nc=nc,
|
||||
opaque=challenge['opaque'])
|
||||
|
||||
creds = self.credentialFactory.decode(clientResponse, self.method,
|
||||
self.clientAddress.host)
|
||||
self.assertTrue(verifyObject(IUsernameDigestHash, creds))
|
||||
|
||||
cleartext = self.username + b":" + self.realm + b":" + self.password
|
||||
hash = md5(cleartext)
|
||||
self.assertTrue(creds.checkHash(hexlify(hash.digest())))
|
||||
hash.update(b'wrong')
|
||||
self.assertFalse(creds.checkHash(hexlify(hash.digest())))
|
||||
|
||||
|
||||
def test_invalidOpaque(self):
|
||||
"""
|
||||
L{DigestCredentialFactory.decode} raises L{LoginFailed} when the opaque
|
||||
value does not contain all the required parts.
|
||||
"""
|
||||
credentialFactory = FakeDigestCredentialFactory(self.algorithm,
|
||||
self.realm)
|
||||
challenge = credentialFactory.getChallenge(self.clientAddress.host)
|
||||
|
||||
exc = self.assertRaises(
|
||||
LoginFailed,
|
||||
credentialFactory._verifyOpaque,
|
||||
b'badOpaque',
|
||||
challenge['nonce'],
|
||||
self.clientAddress.host)
|
||||
self.assertEqual(str(exc), 'Invalid response, invalid opaque value')
|
||||
|
||||
badOpaque = b'foo-' + b64encode(b'nonce,clientip')
|
||||
|
||||
exc = self.assertRaises(
|
||||
LoginFailed,
|
||||
credentialFactory._verifyOpaque,
|
||||
badOpaque,
|
||||
challenge['nonce'],
|
||||
self.clientAddress.host)
|
||||
self.assertEqual(str(exc), 'Invalid response, invalid opaque value')
|
||||
|
||||
exc = self.assertRaises(
|
||||
LoginFailed,
|
||||
credentialFactory._verifyOpaque,
|
||||
b'',
|
||||
challenge['nonce'],
|
||||
self.clientAddress.host)
|
||||
self.assertEqual(str(exc), 'Invalid response, invalid opaque value')
|
||||
|
||||
badOpaque = b'foo-' + b64encode(
|
||||
b",".join((challenge['nonce'],
|
||||
networkString(self.clientAddress.host),
|
||||
b"foobar")))
|
||||
exc = self.assertRaises(
|
||||
LoginFailed,
|
||||
credentialFactory._verifyOpaque,
|
||||
badOpaque,
|
||||
challenge['nonce'],
|
||||
self.clientAddress.host)
|
||||
self.assertEqual(
|
||||
str(exc), 'Invalid response, invalid opaque/time values')
|
||||
|
||||
|
||||
def test_incompatibleNonce(self):
|
||||
"""
|
||||
L{DigestCredentialFactory.decode} raises L{LoginFailed} when the given
|
||||
nonce from the response does not match the nonce encoded in the opaque.
|
||||
"""
|
||||
credentialFactory = FakeDigestCredentialFactory(self.algorithm,
|
||||
self.realm)
|
||||
challenge = credentialFactory.getChallenge(self.clientAddress.host)
|
||||
|
||||
badNonceOpaque = credentialFactory._generateOpaque(
|
||||
b'1234567890',
|
||||
self.clientAddress.host)
|
||||
|
||||
exc = self.assertRaises(
|
||||
LoginFailed,
|
||||
credentialFactory._verifyOpaque,
|
||||
badNonceOpaque,
|
||||
challenge['nonce'],
|
||||
self.clientAddress.host)
|
||||
self.assertEqual(
|
||||
str(exc),
|
||||
'Invalid response, incompatible opaque/nonce values')
|
||||
|
||||
exc = self.assertRaises(
|
||||
LoginFailed,
|
||||
credentialFactory._verifyOpaque,
|
||||
badNonceOpaque,
|
||||
b'',
|
||||
self.clientAddress.host)
|
||||
self.assertEqual(
|
||||
str(exc),
|
||||
'Invalid response, incompatible opaque/nonce values')
|
||||
|
||||
|
||||
def test_incompatibleClientIP(self):
|
||||
"""
|
||||
L{DigestCredentialFactory.decode} raises L{LoginFailed} when the
|
||||
request comes from a client IP other than what is encoded in the
|
||||
opaque.
|
||||
"""
|
||||
credentialFactory = FakeDigestCredentialFactory(self.algorithm,
|
||||
self.realm)
|
||||
challenge = credentialFactory.getChallenge(self.clientAddress.host)
|
||||
|
||||
badAddress = '10.0.0.1'
|
||||
# Sanity check
|
||||
self.assertNotEqual(self.clientAddress.host, badAddress)
|
||||
|
||||
badNonceOpaque = credentialFactory._generateOpaque(
|
||||
challenge['nonce'], badAddress)
|
||||
|
||||
self.assertRaises(
|
||||
LoginFailed,
|
||||
credentialFactory._verifyOpaque,
|
||||
badNonceOpaque,
|
||||
challenge['nonce'],
|
||||
self.clientAddress.host)
|
||||
|
||||
|
||||
def test_oldNonce(self):
|
||||
"""
|
||||
L{DigestCredentialFactory.decode} raises L{LoginFailed} when the given
|
||||
opaque is older than C{DigestCredentialFactory.CHALLENGE_LIFETIME_SECS}
|
||||
"""
|
||||
credentialFactory = FakeDigestCredentialFactory(self.algorithm,
|
||||
self.realm)
|
||||
challenge = credentialFactory.getChallenge(self.clientAddress.host)
|
||||
|
||||
key = b",".join((challenge['nonce'],
|
||||
networkString(self.clientAddress.host),
|
||||
b'-137876876'))
|
||||
digest = hexlify(md5(key + credentialFactory.privateKey).digest())
|
||||
ekey = b64encode(key)
|
||||
|
||||
oldNonceOpaque = b"-".join((digest, ekey.strip(b'\n')))
|
||||
|
||||
self.assertRaises(
|
||||
LoginFailed,
|
||||
credentialFactory._verifyOpaque,
|
||||
oldNonceOpaque,
|
||||
challenge['nonce'],
|
||||
self.clientAddress.host)
|
||||
|
||||
|
||||
def test_mismatchedOpaqueChecksum(self):
|
||||
"""
|
||||
L{DigestCredentialFactory.decode} raises L{LoginFailed} when the opaque
|
||||
checksum fails verification.
|
||||
"""
|
||||
credentialFactory = FakeDigestCredentialFactory(self.algorithm,
|
||||
self.realm)
|
||||
challenge = credentialFactory.getChallenge(self.clientAddress.host)
|
||||
|
||||
key = b",".join((challenge['nonce'],
|
||||
networkString(self.clientAddress.host),
|
||||
b'0'))
|
||||
|
||||
digest = hexlify(md5(key + b'this is not the right pkey').digest())
|
||||
badChecksum = b"-".join((digest, b64encode(key)))
|
||||
|
||||
self.assertRaises(
|
||||
LoginFailed,
|
||||
credentialFactory._verifyOpaque,
|
||||
badChecksum,
|
||||
challenge['nonce'],
|
||||
self.clientAddress.host)
|
||||
|
||||
|
||||
def test_incompatibleCalcHA1Options(self):
|
||||
"""
|
||||
L{calcHA1} raises L{TypeError} when any of the pszUsername, pszRealm,
|
||||
or pszPassword arguments are specified with the preHA1 keyword
|
||||
argument.
|
||||
"""
|
||||
arguments = (
|
||||
(b"user", b"realm", b"password", b"preHA1"),
|
||||
(None, b"realm", None, b"preHA1"),
|
||||
(None, None, b"password", b"preHA1"),
|
||||
)
|
||||
|
||||
for pszUsername, pszRealm, pszPassword, preHA1 in arguments:
|
||||
self.assertRaises(
|
||||
TypeError,
|
||||
calcHA1,
|
||||
b"md5",
|
||||
pszUsername,
|
||||
pszRealm,
|
||||
pszPassword,
|
||||
b"nonce",
|
||||
b"cnonce",
|
||||
preHA1=preHA1)
|
||||
|
||||
|
||||
def test_noNewlineOpaque(self):
|
||||
"""
|
||||
L{DigestCredentialFactory._generateOpaque} returns a value without
|
||||
newlines, regardless of the length of the nonce.
|
||||
"""
|
||||
opaque = self.credentialFactory._generateOpaque(
|
||||
b"long nonce " * 10, None)
|
||||
self.assertNotIn(b'\n', opaque)
|
||||
@@ -0,0 +1,748 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
L{twisted.cred.strcred}.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
import os
|
||||
|
||||
from twisted import plugin
|
||||
from twisted.trial import unittest
|
||||
from twisted.cred import credentials, checkers, error, strcred
|
||||
from twisted.plugins import cred_file, cred_anonymous, cred_unix
|
||||
from twisted.python import usage
|
||||
from twisted.python.compat import NativeStringIO
|
||||
from twisted.python.filepath import FilePath
|
||||
from twisted.python.fakepwd import UserDatabase
|
||||
from twisted.python.reflect import requireModule
|
||||
|
||||
try:
|
||||
import crypt
|
||||
except ImportError:
|
||||
crypt = None
|
||||
|
||||
try:
|
||||
import pwd
|
||||
except ImportError:
|
||||
pwd = None
|
||||
|
||||
try:
|
||||
import spwd
|
||||
except ImportError:
|
||||
spwd = None
|
||||
|
||||
|
||||
|
||||
def getInvalidAuthType():
|
||||
"""
|
||||
Helper method to produce an auth type that doesn't exist.
|
||||
"""
|
||||
invalidAuthType = 'ThisPluginDoesNotExist'
|
||||
while (invalidAuthType in
|
||||
[factory.authType for factory in strcred.findCheckerFactories()]):
|
||||
invalidAuthType += '_'
|
||||
return invalidAuthType
|
||||
|
||||
|
||||
|
||||
class PublicAPITests(unittest.TestCase):
|
||||
|
||||
def test_emptyDescription(self):
|
||||
"""
|
||||
The description string cannot be empty.
|
||||
"""
|
||||
iat = getInvalidAuthType()
|
||||
self.assertRaises(strcred.InvalidAuthType, strcred.makeChecker, iat)
|
||||
self.assertRaises(
|
||||
strcred.InvalidAuthType, strcred.findCheckerFactory, iat)
|
||||
|
||||
|
||||
def test_invalidAuthType(self):
|
||||
"""
|
||||
An unrecognized auth type raises an exception.
|
||||
"""
|
||||
iat = getInvalidAuthType()
|
||||
self.assertRaises(strcred.InvalidAuthType, strcred.makeChecker, iat)
|
||||
self.assertRaises(
|
||||
strcred.InvalidAuthType, strcred.findCheckerFactory, iat)
|
||||
|
||||
|
||||
|
||||
class StrcredFunctionsTests(unittest.TestCase):
|
||||
|
||||
def test_findCheckerFactories(self):
|
||||
"""
|
||||
L{strcred.findCheckerFactories} returns all available plugins.
|
||||
"""
|
||||
availablePlugins = list(strcred.findCheckerFactories())
|
||||
for plg in plugin.getPlugins(strcred.ICheckerFactory):
|
||||
self.assertIn(plg, availablePlugins)
|
||||
|
||||
|
||||
def test_findCheckerFactory(self):
|
||||
"""
|
||||
L{strcred.findCheckerFactory} returns the first plugin
|
||||
available for a given authentication type.
|
||||
"""
|
||||
self.assertIdentical(strcred.findCheckerFactory('file'),
|
||||
cred_file.theFileCheckerFactory)
|
||||
|
||||
|
||||
|
||||
class MemoryCheckerTests(unittest.TestCase):
|
||||
|
||||
def setUp(self):
|
||||
self.admin = credentials.UsernamePassword('admin', 'asdf')
|
||||
self.alice = credentials.UsernamePassword('alice', 'foo')
|
||||
self.badPass = credentials.UsernamePassword('alice', 'foobar')
|
||||
self.badUser = credentials.UsernamePassword('x', 'yz')
|
||||
self.checker = strcred.makeChecker('memory:admin:asdf:alice:foo')
|
||||
|
||||
|
||||
def test_isChecker(self):
|
||||
"""
|
||||
Verifies that strcred.makeChecker('memory') returns an object
|
||||
that implements the L{ICredentialsChecker} interface.
|
||||
"""
|
||||
self.assertTrue(checkers.ICredentialsChecker.providedBy(self.checker))
|
||||
self.assertIn(credentials.IUsernamePassword,
|
||||
self.checker.credentialInterfaces)
|
||||
|
||||
|
||||
def test_badFormatArgString(self):
|
||||
"""
|
||||
An argument string which does not contain user:pass pairs
|
||||
(i.e., an odd number of ':' characters) raises an exception.
|
||||
"""
|
||||
self.assertRaises(strcred.InvalidAuthArgumentString,
|
||||
strcred.makeChecker, 'memory:a:b:c')
|
||||
|
||||
|
||||
def test_memoryCheckerSucceeds(self):
|
||||
"""
|
||||
The checker works with valid credentials.
|
||||
"""
|
||||
def _gotAvatar(username):
|
||||
self.assertEqual(username, self.admin.username)
|
||||
return (self.checker
|
||||
.requestAvatarId(self.admin)
|
||||
.addCallback(_gotAvatar))
|
||||
|
||||
|
||||
def test_memoryCheckerFailsUsername(self):
|
||||
"""
|
||||
The checker fails with an invalid username.
|
||||
"""
|
||||
return self.assertFailure(self.checker.requestAvatarId(self.badUser),
|
||||
error.UnauthorizedLogin)
|
||||
|
||||
|
||||
def test_memoryCheckerFailsPassword(self):
|
||||
"""
|
||||
The checker fails with an invalid password.
|
||||
"""
|
||||
return self.assertFailure(self.checker.requestAvatarId(self.badPass),
|
||||
error.UnauthorizedLogin)
|
||||
|
||||
|
||||
|
||||
class AnonymousCheckerTests(unittest.TestCase):
|
||||
|
||||
def test_isChecker(self):
|
||||
"""
|
||||
Verifies that strcred.makeChecker('anonymous') returns an object
|
||||
that implements the L{ICredentialsChecker} interface.
|
||||
"""
|
||||
checker = strcred.makeChecker('anonymous')
|
||||
self.assertTrue(checkers.ICredentialsChecker.providedBy(checker))
|
||||
self.assertIn(credentials.IAnonymous, checker.credentialInterfaces)
|
||||
|
||||
|
||||
def testAnonymousAccessSucceeds(self):
|
||||
"""
|
||||
We can log in anonymously using this checker.
|
||||
"""
|
||||
checker = strcred.makeChecker('anonymous')
|
||||
request = checker.requestAvatarId(credentials.Anonymous())
|
||||
def _gotAvatar(avatar):
|
||||
self.assertIdentical(checkers.ANONYMOUS, avatar)
|
||||
return request.addCallback(_gotAvatar)
|
||||
|
||||
|
||||
|
||||
class UnixCheckerTests(unittest.TestCase):
|
||||
users = {
|
||||
'admin': 'asdf',
|
||||
'alice': 'foo',
|
||||
}
|
||||
|
||||
|
||||
def _spwd_getspnam(self, username):
|
||||
return spwd.struct_spwd((username,
|
||||
crypt.crypt(self.users[username], 'F/'),
|
||||
0, 0, 99999, 7, -1, -1, -1))
|
||||
|
||||
|
||||
def setUp(self):
|
||||
self.admin = credentials.UsernamePassword('admin', 'asdf')
|
||||
self.alice = credentials.UsernamePassword('alice', 'foo')
|
||||
self.badPass = credentials.UsernamePassword('alice', 'foobar')
|
||||
self.badUser = credentials.UsernamePassword('x', 'yz')
|
||||
self.checker = strcred.makeChecker('unix')
|
||||
self.adminBytes = credentials.UsernamePassword(b'admin', b'asdf')
|
||||
self.aliceBytes = credentials.UsernamePassword(b'alice', b'foo')
|
||||
self.badPassBytes = credentials.UsernamePassword(b'alice', b'foobar')
|
||||
self.badUserBytes = credentials.UsernamePassword(b'x', b'yz')
|
||||
self.checkerBytes = strcred.makeChecker('unix')
|
||||
|
||||
# Hack around the pwd and spwd modules, since we can't really
|
||||
# go about reading your /etc/passwd or /etc/shadow files
|
||||
if pwd:
|
||||
database = UserDatabase()
|
||||
for username, password in self.users.items():
|
||||
database.addUser(
|
||||
username, crypt.crypt(password, 'F/'),
|
||||
1000, 1000, username, '/home/' + username, '/bin/sh')
|
||||
self.patch(pwd, 'getpwnam', database.getpwnam)
|
||||
if spwd:
|
||||
self.patch(spwd, 'getspnam', self._spwd_getspnam)
|
||||
|
||||
|
||||
def test_isChecker(self):
|
||||
"""
|
||||
Verifies that strcred.makeChecker('unix') returns an object
|
||||
that implements the L{ICredentialsChecker} interface.
|
||||
"""
|
||||
self.assertTrue(checkers.ICredentialsChecker.providedBy(self.checker))
|
||||
self.assertIn(credentials.IUsernamePassword,
|
||||
self.checker.credentialInterfaces)
|
||||
self.assertTrue(checkers.ICredentialsChecker.providedBy(
|
||||
self.checkerBytes))
|
||||
self.assertIn(credentials.IUsernamePassword,
|
||||
self.checkerBytes.credentialInterfaces)
|
||||
|
||||
|
||||
def test_unixCheckerSucceeds(self):
|
||||
"""
|
||||
The checker works with valid credentials.
|
||||
"""
|
||||
def _gotAvatar(username):
|
||||
self.assertEqual(username, self.admin.username)
|
||||
return (self.checker
|
||||
.requestAvatarId(self.admin)
|
||||
.addCallback(_gotAvatar))
|
||||
|
||||
|
||||
def test_unixCheckerSucceedsBytes(self):
|
||||
"""
|
||||
The checker works with valid L{bytes} credentials.
|
||||
"""
|
||||
def _gotAvatar(username):
|
||||
self.assertEqual(username,
|
||||
self.adminBytes.username.decode("utf-8"))
|
||||
return (self.checkerBytes
|
||||
.requestAvatarId(self.adminBytes)
|
||||
.addCallback(_gotAvatar))
|
||||
|
||||
|
||||
def test_unixCheckerFailsUsername(self):
|
||||
"""
|
||||
The checker fails with an invalid username.
|
||||
"""
|
||||
return self.assertFailure(self.checker.requestAvatarId(self.badUser),
|
||||
error.UnauthorizedLogin)
|
||||
|
||||
|
||||
def test_unixCheckerFailsUsernameBytes(self):
|
||||
"""
|
||||
The checker fails with an invalid L{bytes} username.
|
||||
"""
|
||||
return self.assertFailure(self.checkerBytes.requestAvatarId(
|
||||
self.badUserBytes), error.UnauthorizedLogin)
|
||||
|
||||
|
||||
def test_unixCheckerFailsPassword(self):
|
||||
"""
|
||||
The checker fails with an invalid password.
|
||||
"""
|
||||
return self.assertFailure(self.checker.requestAvatarId(self.badPass),
|
||||
error.UnauthorizedLogin)
|
||||
|
||||
|
||||
def test_unixCheckerFailsPasswordBytes(self):
|
||||
"""
|
||||
The checker fails with an invalid L{bytes} password.
|
||||
"""
|
||||
return self.assertFailure(self.checkerBytes.requestAvatarId(
|
||||
self.badPassBytes), error.UnauthorizedLogin)
|
||||
|
||||
|
||||
if None in (pwd, spwd, crypt):
|
||||
availability = []
|
||||
for module, name in ((pwd, "pwd"), (spwd, "spwd"), (crypt, "crypt")):
|
||||
if module is None:
|
||||
availability += [name]
|
||||
for method in (test_unixCheckerSucceeds,
|
||||
test_unixCheckerSucceedsBytes,
|
||||
test_unixCheckerFailsUsername,
|
||||
test_unixCheckerFailsUsernameBytes,
|
||||
test_unixCheckerFailsPassword,
|
||||
test_unixCheckerFailsPasswordBytes):
|
||||
method.skip = ("Required module(s) are unavailable: " +
|
||||
", ".join(availability))
|
||||
|
||||
|
||||
class CryptTests(unittest.TestCase):
|
||||
"""
|
||||
L{crypt} has functions for encrypting password.
|
||||
"""
|
||||
if not crypt:
|
||||
skip = "Required module is unavailable: crypt"
|
||||
|
||||
def test_verifyCryptedPassword(self):
|
||||
"""
|
||||
L{cred_unix.verifyCryptedPassword}
|
||||
"""
|
||||
password = "sample password ^%$"
|
||||
|
||||
for salt in (None, "ab"):
|
||||
try:
|
||||
cryptedCorrect = crypt.crypt(password, salt)
|
||||
except TypeError:
|
||||
# Older Python versions would throw a TypeError if
|
||||
# a value of None was is used for the salt.
|
||||
# Newer Python versions allow it.
|
||||
continue
|
||||
cryptedIncorrect = "$1x1234"
|
||||
self.assertTrue(cred_unix.verifyCryptedPassword(cryptedCorrect,
|
||||
password))
|
||||
self.assertFalse(cred_unix.verifyCryptedPassword(cryptedIncorrect,
|
||||
password))
|
||||
|
||||
|
||||
# Python 3.3+ has crypt.METHOD_*, but not all
|
||||
# platforms implement all methods.
|
||||
for method in ("METHOD_SHA512", "METHOD_SHA256", "METHOD_MD5",
|
||||
"METHOD_CRYPT"):
|
||||
cryptMethod = getattr(crypt, method, None)
|
||||
if not cryptMethod:
|
||||
continue
|
||||
password = "interesting password xyz"
|
||||
crypted = crypt.crypt(password, cryptMethod)
|
||||
incorrectCrypted = crypted + "blahfooincorrect"
|
||||
result = cred_unix.verifyCryptedPassword(crypted, password)
|
||||
self.assertTrue(result)
|
||||
# Try to pass in bytes
|
||||
result = cred_unix.verifyCryptedPassword(crypted.encode("utf-8"),
|
||||
password.encode("utf-8"))
|
||||
self.assertTrue(result)
|
||||
result = cred_unix.verifyCryptedPassword(incorrectCrypted, password)
|
||||
self.assertFalse(result)
|
||||
# Try to pass in bytes
|
||||
result = cred_unix.verifyCryptedPassword(incorrectCrypted.encode("utf-8"),
|
||||
password.encode("utf-8"))
|
||||
self.assertFalse(result)
|
||||
|
||||
|
||||
|
||||
class FileDBCheckerTests(unittest.TestCase):
|
||||
"""
|
||||
C{--auth=file:...} file checker.
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
self.admin = credentials.UsernamePassword(b'admin', b'asdf')
|
||||
self.alice = credentials.UsernamePassword(b'alice', b'foo')
|
||||
self.badPass = credentials.UsernamePassword(b'alice', b'foobar')
|
||||
self.badUser = credentials.UsernamePassword(b'x', b'yz')
|
||||
self.filename = self.mktemp()
|
||||
FilePath(self.filename).setContent(b'admin:asdf\nalice:foo\n')
|
||||
self.checker = strcred.makeChecker('file:' + self.filename)
|
||||
|
||||
|
||||
def _fakeFilename(self):
|
||||
filename = '/DoesNotExist'
|
||||
while os.path.exists(filename):
|
||||
filename += '_'
|
||||
return filename
|
||||
|
||||
|
||||
def test_isChecker(self):
|
||||
"""
|
||||
Verifies that strcred.makeChecker('memory') returns an object
|
||||
that implements the L{ICredentialsChecker} interface.
|
||||
"""
|
||||
self.assertTrue(checkers.ICredentialsChecker.providedBy(self.checker))
|
||||
self.assertIn(credentials.IUsernamePassword,
|
||||
self.checker.credentialInterfaces)
|
||||
|
||||
|
||||
def test_fileCheckerSucceeds(self):
|
||||
"""
|
||||
The checker works with valid credentials.
|
||||
"""
|
||||
def _gotAvatar(username):
|
||||
self.assertEqual(username, self.admin.username)
|
||||
return (self.checker
|
||||
.requestAvatarId(self.admin)
|
||||
.addCallback(_gotAvatar))
|
||||
|
||||
|
||||
def test_fileCheckerFailsUsername(self):
|
||||
"""
|
||||
The checker fails with an invalid username.
|
||||
"""
|
||||
return self.assertFailure(self.checker.requestAvatarId(self.badUser),
|
||||
error.UnauthorizedLogin)
|
||||
|
||||
|
||||
def test_fileCheckerFailsPassword(self):
|
||||
"""
|
||||
The checker fails with an invalid password.
|
||||
"""
|
||||
return self.assertFailure(self.checker.requestAvatarId(self.badPass),
|
||||
error.UnauthorizedLogin)
|
||||
|
||||
|
||||
def test_failsWithEmptyFilename(self):
|
||||
"""
|
||||
An empty filename raises an error.
|
||||
"""
|
||||
self.assertRaises(ValueError, strcred.makeChecker, 'file')
|
||||
self.assertRaises(ValueError, strcred.makeChecker, 'file:')
|
||||
|
||||
|
||||
def test_warnWithBadFilename(self):
|
||||
"""
|
||||
When the file auth plugin is given a file that doesn't exist, it
|
||||
should produce a warning.
|
||||
"""
|
||||
oldOutput = cred_file.theFileCheckerFactory.errorOutput
|
||||
newOutput = NativeStringIO()
|
||||
cred_file.theFileCheckerFactory.errorOutput = newOutput
|
||||
strcred.makeChecker('file:' + self._fakeFilename())
|
||||
cred_file.theFileCheckerFactory.errorOutput = oldOutput
|
||||
self.assertIn(cred_file.invalidFileWarning, newOutput.getvalue())
|
||||
|
||||
|
||||
|
||||
class SSHCheckerTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for the C{--auth=sshkey:...} checker. The majority of the tests for the
|
||||
ssh public key database checker are in
|
||||
L{twisted.conch.test.test_checkers.SSHPublicKeyCheckerTestCase}.
|
||||
"""
|
||||
|
||||
skip = None
|
||||
|
||||
if requireModule('cryptography') is None:
|
||||
skip = 'cryptography is not available'
|
||||
|
||||
if requireModule('pyasn1') is None:
|
||||
skip = 'pyasn1 is not available'
|
||||
|
||||
|
||||
def test_isChecker(self):
|
||||
"""
|
||||
Verifies that strcred.makeChecker('sshkey') returns an object
|
||||
that implements the L{ICredentialsChecker} interface.
|
||||
"""
|
||||
sshChecker = strcred.makeChecker('sshkey')
|
||||
self.assertTrue(checkers.ICredentialsChecker.providedBy(sshChecker))
|
||||
self.assertIn(
|
||||
credentials.ISSHPrivateKey, sshChecker.credentialInterfaces)
|
||||
|
||||
|
||||
|
||||
class DummyOptions(usage.Options, strcred.AuthOptionMixin):
|
||||
"""
|
||||
Simple options for testing L{strcred.AuthOptionMixin}.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class CheckerOptionsTests(unittest.TestCase):
|
||||
|
||||
def test_createsList(self):
|
||||
"""
|
||||
The C{--auth} command line creates a list in the
|
||||
Options instance and appends values to it.
|
||||
"""
|
||||
options = DummyOptions()
|
||||
options.parseOptions(['--auth', 'memory'])
|
||||
self.assertEqual(len(options['credCheckers']), 1)
|
||||
options = DummyOptions()
|
||||
options.parseOptions(['--auth', 'memory', '--auth', 'memory'])
|
||||
self.assertEqual(len(options['credCheckers']), 2)
|
||||
|
||||
|
||||
def test_invalidAuthError(self):
|
||||
"""
|
||||
The C{--auth} command line raises an exception when it
|
||||
gets a parameter it doesn't understand.
|
||||
"""
|
||||
options = DummyOptions()
|
||||
# If someone adds a 'ThisPluginDoesNotExist' then this unit
|
||||
# test should still run.
|
||||
invalidParameter = getInvalidAuthType()
|
||||
self.assertRaises(
|
||||
usage.UsageError,
|
||||
options.parseOptions, ['--auth', invalidParameter])
|
||||
self.assertRaises(
|
||||
usage.UsageError,
|
||||
options.parseOptions, ['--help-auth-type', invalidParameter])
|
||||
|
||||
|
||||
def test_createsDictionary(self):
|
||||
"""
|
||||
The C{--auth} command line creates a dictionary mapping supported
|
||||
interfaces to the list of credentials checkers that support it.
|
||||
"""
|
||||
options = DummyOptions()
|
||||
options.parseOptions(['--auth', 'memory', '--auth', 'anonymous'])
|
||||
chd = options['credInterfaces']
|
||||
self.assertEqual(len(chd[credentials.IAnonymous]), 1)
|
||||
self.assertEqual(len(chd[credentials.IUsernamePassword]), 1)
|
||||
chdAnonymous = chd[credentials.IAnonymous][0]
|
||||
chdUserPass = chd[credentials.IUsernamePassword][0]
|
||||
self.assertTrue(checkers.ICredentialsChecker.providedBy(chdAnonymous))
|
||||
self.assertTrue(checkers.ICredentialsChecker.providedBy(chdUserPass))
|
||||
self.assertIn(credentials.IAnonymous,
|
||||
chdAnonymous.credentialInterfaces)
|
||||
self.assertIn(credentials.IUsernamePassword,
|
||||
chdUserPass.credentialInterfaces)
|
||||
|
||||
|
||||
def test_credInterfacesProvidesLists(self):
|
||||
"""
|
||||
When two C{--auth} arguments are passed along which support the same
|
||||
interface, a list with both is created.
|
||||
"""
|
||||
options = DummyOptions()
|
||||
options.parseOptions(['--auth', 'memory', '--auth', 'unix'])
|
||||
self.assertEqual(
|
||||
options['credCheckers'],
|
||||
options['credInterfaces'][credentials.IUsernamePassword])
|
||||
|
||||
|
||||
def test_listDoesNotDisplayDuplicates(self):
|
||||
"""
|
||||
The list for C{--help-auth} does not duplicate items.
|
||||
"""
|
||||
authTypes = []
|
||||
options = DummyOptions()
|
||||
for cf in options._checkerFactoriesForOptHelpAuth():
|
||||
self.assertNotIn(cf.authType, authTypes)
|
||||
authTypes.append(cf.authType)
|
||||
|
||||
|
||||
def test_displaysListCorrectly(self):
|
||||
"""
|
||||
The C{--help-auth} argument correctly displays all
|
||||
available authentication plugins, then exits.
|
||||
"""
|
||||
newStdout = NativeStringIO()
|
||||
options = DummyOptions()
|
||||
options.authOutput = newStdout
|
||||
self.assertRaises(SystemExit, options.parseOptions, ['--help-auth'])
|
||||
for checkerFactory in strcred.findCheckerFactories():
|
||||
self.assertIn(checkerFactory.authType, newStdout.getvalue())
|
||||
|
||||
|
||||
def test_displaysHelpCorrectly(self):
|
||||
"""
|
||||
The C{--help-auth-for} argument will correctly display the help file for a
|
||||
particular authentication plugin.
|
||||
"""
|
||||
newStdout = NativeStringIO()
|
||||
options = DummyOptions()
|
||||
options.authOutput = newStdout
|
||||
self.assertRaises(
|
||||
SystemExit, options.parseOptions, ['--help-auth-type', 'file'])
|
||||
for line in cred_file.theFileCheckerFactory.authHelp:
|
||||
if line.strip():
|
||||
self.assertIn(line.strip(), newStdout.getvalue())
|
||||
|
||||
|
||||
def test_unexpectedException(self):
|
||||
"""
|
||||
When the checker specified by C{--auth} raises an unexpected error, it
|
||||
should be caught and re-raised within a L{usage.UsageError}.
|
||||
"""
|
||||
options = DummyOptions()
|
||||
err = self.assertRaises(usage.UsageError, options.parseOptions,
|
||||
['--auth', 'file'])
|
||||
self.assertEqual(str(err),
|
||||
"Unexpected error: 'file' requires a filename")
|
||||
|
||||
|
||||
|
||||
class OptionsForUsernamePassword(usage.Options, strcred.AuthOptionMixin):
|
||||
supportedInterfaces = (credentials.IUsernamePassword,)
|
||||
|
||||
|
||||
|
||||
class OptionsForUsernameHashedPassword(usage.Options, strcred.AuthOptionMixin):
|
||||
supportedInterfaces = (credentials.IUsernameHashedPassword,)
|
||||
|
||||
|
||||
|
||||
class OptionsSupportsAllInterfaces(usage.Options, strcred.AuthOptionMixin):
|
||||
supportedInterfaces = None
|
||||
|
||||
|
||||
|
||||
class OptionsSupportsNoInterfaces(usage.Options, strcred.AuthOptionMixin):
|
||||
supportedInterfaces = []
|
||||
|
||||
|
||||
|
||||
class LimitingInterfacesTests(unittest.TestCase):
|
||||
"""
|
||||
Tests functionality that allows an application to limit the
|
||||
credential interfaces it can support. For the purposes of this
|
||||
test, we use IUsernameHashedPassword, although this will never
|
||||
really be used by the command line.
|
||||
|
||||
(I have, to date, not thought of a half-decent way for a user to
|
||||
specify a hash algorithm via the command-line. Nor do I think it's
|
||||
very useful.)
|
||||
|
||||
I should note that, at first, this test is counter-intuitive,
|
||||
because we're using the checker with a pre-defined hash function
|
||||
as the 'bad' checker. See the documentation for
|
||||
L{twisted.cred.checkers.FilePasswordDB.hash} for more details.
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
self.filename = self.mktemp()
|
||||
with open(self.filename, 'wb') as f:
|
||||
f.write(b'admin:asdf\nalice:foo\n')
|
||||
self.goodChecker = checkers.FilePasswordDB(self.filename)
|
||||
self.badChecker = checkers.FilePasswordDB(
|
||||
self.filename, hash=self._hash)
|
||||
self.anonChecker = checkers.AllowAnonymousAccess()
|
||||
|
||||
|
||||
def _hash(self, networkUsername, networkPassword, storedPassword):
|
||||
"""
|
||||
A dumb hash that doesn't really do anything.
|
||||
"""
|
||||
return networkPassword
|
||||
|
||||
|
||||
def test_supportsInterface(self):
|
||||
"""
|
||||
The supportsInterface method behaves appropriately.
|
||||
"""
|
||||
options = OptionsForUsernamePassword()
|
||||
self.assertTrue(
|
||||
options.supportsInterface(credentials.IUsernamePassword))
|
||||
self.assertFalse(
|
||||
options.supportsInterface(credentials.IAnonymous))
|
||||
self.assertRaises(
|
||||
strcred.UnsupportedInterfaces, options.addChecker,
|
||||
self.anonChecker)
|
||||
|
||||
|
||||
def test_supportsAllInterfaces(self):
|
||||
"""
|
||||
The supportsInterface method behaves appropriately
|
||||
when the supportedInterfaces attribute is None.
|
||||
"""
|
||||
options = OptionsSupportsAllInterfaces()
|
||||
self.assertTrue(
|
||||
options.supportsInterface(credentials.IUsernamePassword))
|
||||
self.assertTrue(
|
||||
options.supportsInterface(credentials.IAnonymous))
|
||||
|
||||
|
||||
def test_supportsCheckerFactory(self):
|
||||
"""
|
||||
The supportsCheckerFactory method behaves appropriately.
|
||||
"""
|
||||
options = OptionsForUsernamePassword()
|
||||
fileCF = cred_file.theFileCheckerFactory
|
||||
anonCF = cred_anonymous.theAnonymousCheckerFactory
|
||||
self.assertTrue(options.supportsCheckerFactory(fileCF))
|
||||
self.assertFalse(options.supportsCheckerFactory(anonCF))
|
||||
|
||||
|
||||
def test_canAddSupportedChecker(self):
|
||||
"""
|
||||
When addChecker is called with a checker that implements at least one
|
||||
of the interfaces our application supports, it is successful.
|
||||
"""
|
||||
options = OptionsForUsernamePassword()
|
||||
options.addChecker(self.goodChecker)
|
||||
iface = options.supportedInterfaces[0]
|
||||
# Test that we did get IUsernamePassword
|
||||
self.assertIdentical(
|
||||
options['credInterfaces'][iface][0], self.goodChecker)
|
||||
self.assertIdentical(options['credCheckers'][0], self.goodChecker)
|
||||
# Test that we didn't get IUsernameHashedPassword
|
||||
self.assertEqual(len(options['credInterfaces'][iface]), 1)
|
||||
self.assertEqual(len(options['credCheckers']), 1)
|
||||
|
||||
|
||||
def test_failOnAddingUnsupportedChecker(self):
|
||||
"""
|
||||
When addChecker is called with a checker that does not implement any
|
||||
supported interfaces, it fails.
|
||||
"""
|
||||
options = OptionsForUsernameHashedPassword()
|
||||
self.assertRaises(strcred.UnsupportedInterfaces,
|
||||
options.addChecker, self.badChecker)
|
||||
|
||||
|
||||
def test_unsupportedInterfaceError(self):
|
||||
"""
|
||||
The C{--auth} command line raises an exception when it
|
||||
gets a checker we don't support.
|
||||
"""
|
||||
options = OptionsSupportsNoInterfaces()
|
||||
authType = cred_anonymous.theAnonymousCheckerFactory.authType
|
||||
self.assertRaises(
|
||||
usage.UsageError,
|
||||
options.parseOptions, ['--auth', authType])
|
||||
|
||||
|
||||
def test_helpAuthLimitsOutput(self):
|
||||
"""
|
||||
C{--help-auth} will only list checkers that purport to
|
||||
supply at least one of the credential interfaces our
|
||||
application can use.
|
||||
"""
|
||||
options = OptionsForUsernamePassword()
|
||||
for factory in options._checkerFactoriesForOptHelpAuth():
|
||||
invalid = True
|
||||
for interface in factory.credentialInterfaces:
|
||||
if options.supportsInterface(interface):
|
||||
invalid = False
|
||||
if invalid:
|
||||
raise strcred.UnsupportedInterfaces()
|
||||
|
||||
|
||||
def test_helpAuthTypeLimitsOutput(self):
|
||||
"""
|
||||
C{--help-auth-type} will display a warning if you get
|
||||
help for an authType that does not supply at least one of the
|
||||
credential interfaces our application can use.
|
||||
"""
|
||||
options = OptionsForUsernamePassword()
|
||||
# Find an interface that we can use for our test
|
||||
invalidFactory = None
|
||||
for factory in strcred.findCheckerFactories():
|
||||
if not options.supportsCheckerFactory(factory):
|
||||
invalidFactory = factory
|
||||
break
|
||||
self.assertNotIdentical(invalidFactory, None)
|
||||
# Capture output and make sure the warning is there
|
||||
newStdout = NativeStringIO()
|
||||
options.authOutput = newStdout
|
||||
self.assertRaises(SystemExit, options.parseOptions,
|
||||
['--help-auth-type', 'anonymous'])
|
||||
self.assertIn(strcred.notSupportedWarning, newStdout.getvalue())
|
||||
@@ -0,0 +1,8 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Twisted Enterprise: Database support for Twisted services.
|
||||
"""
|
||||
|
||||
__all__ = ['adbapi']
|
||||
@@ -0,0 +1,168 @@
|
||||
# -*- test-case-name: twisted.test.test_stdio -*-
|
||||
|
||||
"""Standard input/out/err support.
|
||||
|
||||
Future Plans::
|
||||
|
||||
support for stderr, perhaps
|
||||
Rewrite to use the reactor instead of an ad-hoc mechanism for connecting
|
||||
protocols to transport.
|
||||
|
||||
Maintainer: James Y Knight
|
||||
"""
|
||||
|
||||
from zope.interface import implementer
|
||||
|
||||
from twisted.internet import process, error, interfaces
|
||||
from twisted.python import log, failure
|
||||
|
||||
|
||||
@implementer(interfaces.IAddress)
|
||||
class PipeAddress(object):
|
||||
pass
|
||||
|
||||
|
||||
@implementer(interfaces.ITransport, interfaces.IProducer,
|
||||
interfaces.IConsumer, interfaces.IHalfCloseableDescriptor)
|
||||
class StandardIO(object):
|
||||
|
||||
_reader = None
|
||||
_writer = None
|
||||
disconnected = False
|
||||
disconnecting = False
|
||||
|
||||
def __init__(self, proto, stdin=0, stdout=1, reactor=None):
|
||||
if reactor is None:
|
||||
from twisted.internet import reactor
|
||||
self.protocol = proto
|
||||
|
||||
self._writer = process.ProcessWriter(reactor, self, 'write', stdout)
|
||||
self._reader = process.ProcessReader(reactor, self, 'read', stdin)
|
||||
self._reader.startReading()
|
||||
self.protocol.makeConnection(self)
|
||||
|
||||
# ITransport
|
||||
|
||||
# XXX Actually, see #3597.
|
||||
def loseWriteConnection(self):
|
||||
if self._writer is not None:
|
||||
self._writer.loseConnection()
|
||||
|
||||
def write(self, data):
|
||||
if self._writer is not None:
|
||||
self._writer.write(data)
|
||||
|
||||
def writeSequence(self, data):
|
||||
if self._writer is not None:
|
||||
self._writer.writeSequence(data)
|
||||
|
||||
def loseConnection(self):
|
||||
self.disconnecting = True
|
||||
|
||||
if self._writer is not None:
|
||||
self._writer.loseConnection()
|
||||
if self._reader is not None:
|
||||
# Don't loseConnection, because we don't want to SIGPIPE it.
|
||||
self._reader.stopReading()
|
||||
|
||||
def getPeer(self):
|
||||
return PipeAddress()
|
||||
|
||||
def getHost(self):
|
||||
return PipeAddress()
|
||||
|
||||
|
||||
# Callbacks from process.ProcessReader/ProcessWriter
|
||||
def childDataReceived(self, fd, data):
|
||||
self.protocol.dataReceived(data)
|
||||
|
||||
def childConnectionLost(self, fd, reason):
|
||||
if self.disconnected:
|
||||
return
|
||||
|
||||
if reason.value.__class__ == error.ConnectionDone:
|
||||
# Normal close
|
||||
if fd == 'read':
|
||||
self._readConnectionLost(reason)
|
||||
else:
|
||||
self._writeConnectionLost(reason)
|
||||
else:
|
||||
self.connectionLost(reason)
|
||||
|
||||
def connectionLost(self, reason):
|
||||
self.disconnected = True
|
||||
|
||||
# Make sure to cleanup the other half
|
||||
_reader = self._reader
|
||||
_writer = self._writer
|
||||
protocol = self.protocol
|
||||
self._reader = self._writer = None
|
||||
self.protocol = None
|
||||
|
||||
if _writer is not None and not _writer.disconnected:
|
||||
_writer.connectionLost(reason)
|
||||
|
||||
if _reader is not None and not _reader.disconnected:
|
||||
_reader.connectionLost(reason)
|
||||
|
||||
try:
|
||||
protocol.connectionLost(reason)
|
||||
except:
|
||||
log.err()
|
||||
|
||||
def _writeConnectionLost(self, reason):
|
||||
self._writer=None
|
||||
if self.disconnecting:
|
||||
self.connectionLost(reason)
|
||||
return
|
||||
|
||||
p = interfaces.IHalfCloseableProtocol(self.protocol, None)
|
||||
if p:
|
||||
try:
|
||||
p.writeConnectionLost()
|
||||
except:
|
||||
log.err()
|
||||
self.connectionLost(failure.Failure())
|
||||
|
||||
def _readConnectionLost(self, reason):
|
||||
self._reader=None
|
||||
p = interfaces.IHalfCloseableProtocol(self.protocol, None)
|
||||
if p:
|
||||
try:
|
||||
p.readConnectionLost()
|
||||
except:
|
||||
log.err()
|
||||
self.connectionLost(failure.Failure())
|
||||
else:
|
||||
self.connectionLost(reason)
|
||||
|
||||
# IConsumer
|
||||
def registerProducer(self, producer, streaming):
|
||||
if self._writer is None:
|
||||
producer.stopProducing()
|
||||
else:
|
||||
self._writer.registerProducer(producer, streaming)
|
||||
|
||||
def unregisterProducer(self):
|
||||
if self._writer is not None:
|
||||
self._writer.unregisterProducer()
|
||||
|
||||
# IProducer
|
||||
def stopProducing(self):
|
||||
self.loseConnection()
|
||||
|
||||
def pauseProducing(self):
|
||||
if self._reader is not None:
|
||||
self._reader.pauseProducing()
|
||||
|
||||
def resumeProducing(self):
|
||||
if self._reader is not None:
|
||||
self._reader.resumeProducing()
|
||||
|
||||
def stopReading(self):
|
||||
"""Compatibility only, don't use. Call pauseProducing."""
|
||||
self.pauseProducing()
|
||||
|
||||
def startReading(self):
|
||||
"""Compatibility only, don't use. Call resumeProducing."""
|
||||
self.resumeProducing()
|
||||
@@ -0,0 +1,133 @@
|
||||
# -*- test-case-name: twisted.internet.test.test_win32serialport -*-
|
||||
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Serial port support for Windows.
|
||||
|
||||
Requires PySerial and pywin32.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
# system imports
|
||||
from serial import PARITY_NONE
|
||||
from serial import STOPBITS_ONE
|
||||
from serial import EIGHTBITS
|
||||
from serial.serialutil import to_bytes
|
||||
import win32file, win32event
|
||||
|
||||
# twisted imports
|
||||
from twisted.internet import abstract
|
||||
|
||||
# sibling imports
|
||||
from twisted.internet.serialport import BaseSerialPort
|
||||
|
||||
|
||||
class SerialPort(BaseSerialPort, abstract.FileDescriptor):
|
||||
"""A serial device, acting as a transport, that uses a win32 event."""
|
||||
|
||||
connected = 1
|
||||
|
||||
def __init__(self, protocol, deviceNameOrPortNumber, reactor,
|
||||
baudrate = 9600, bytesize = EIGHTBITS, parity = PARITY_NONE,
|
||||
stopbits = STOPBITS_ONE, xonxoff = 0, rtscts = 0):
|
||||
self._serial = self._serialFactory(
|
||||
deviceNameOrPortNumber, baudrate=baudrate, bytesize=bytesize,
|
||||
parity=parity, stopbits=stopbits, timeout=None,
|
||||
xonxoff=xonxoff, rtscts=rtscts)
|
||||
self.flushInput()
|
||||
self.flushOutput()
|
||||
self.reactor = reactor
|
||||
self.protocol = protocol
|
||||
self.outQueue = []
|
||||
self.closed = 0
|
||||
self.closedNotifies = 0
|
||||
self.writeInProgress = 0
|
||||
|
||||
self.protocol = protocol
|
||||
self._overlappedRead = win32file.OVERLAPPED()
|
||||
self._overlappedRead.hEvent = win32event.CreateEvent(None, 1, 0, None)
|
||||
self._overlappedWrite = win32file.OVERLAPPED()
|
||||
self._overlappedWrite.hEvent = win32event.CreateEvent(None, 0, 0, None)
|
||||
|
||||
self.reactor.addEvent(self._overlappedRead.hEvent, self, 'serialReadEvent')
|
||||
self.reactor.addEvent(self._overlappedWrite.hEvent, self, 'serialWriteEvent')
|
||||
|
||||
self.protocol.makeConnection(self)
|
||||
self._finishPortSetup()
|
||||
|
||||
|
||||
def _finishPortSetup(self):
|
||||
"""
|
||||
Finish setting up the serial port.
|
||||
|
||||
This is a separate method to facilitate testing.
|
||||
"""
|
||||
flags, comstat = self._clearCommError()
|
||||
rc, self.read_buf = win32file.ReadFile(self._serial._port_handle,
|
||||
win32file.AllocateReadBuffer(1),
|
||||
self._overlappedRead)
|
||||
|
||||
|
||||
def _clearCommError(self):
|
||||
return win32file.ClearCommError(self._serial._port_handle)
|
||||
|
||||
|
||||
def serialReadEvent(self):
|
||||
#get that character we set up
|
||||
n = win32file.GetOverlappedResult(self._serial._port_handle, self._overlappedRead, 0)
|
||||
first = to_bytes(self.read_buf[:n])
|
||||
#now we should get everything that is already in the buffer
|
||||
flags, comstat = self._clearCommError()
|
||||
if comstat.cbInQue:
|
||||
win32event.ResetEvent(self._overlappedRead.hEvent)
|
||||
rc, buf = win32file.ReadFile(self._serial._port_handle,
|
||||
win32file.AllocateReadBuffer(comstat.cbInQue),
|
||||
self._overlappedRead)
|
||||
n = win32file.GetOverlappedResult(self._serial._port_handle, self._overlappedRead, 1)
|
||||
#handle all the received data:
|
||||
self.protocol.dataReceived(first + to_bytes(buf[:n]))
|
||||
else:
|
||||
#handle all the received data:
|
||||
self.protocol.dataReceived(first)
|
||||
|
||||
#set up next one
|
||||
win32event.ResetEvent(self._overlappedRead.hEvent)
|
||||
rc, self.read_buf = win32file.ReadFile(self._serial._port_handle,
|
||||
win32file.AllocateReadBuffer(1),
|
||||
self._overlappedRead)
|
||||
|
||||
|
||||
def write(self, data):
|
||||
if data:
|
||||
if self.writeInProgress:
|
||||
self.outQueue.append(data)
|
||||
else:
|
||||
self.writeInProgress = 1
|
||||
win32file.WriteFile(self._serial._port_handle, data, self._overlappedWrite)
|
||||
|
||||
|
||||
def serialWriteEvent(self):
|
||||
try:
|
||||
dataToWrite = self.outQueue.pop(0)
|
||||
except IndexError:
|
||||
self.writeInProgress = 0
|
||||
return
|
||||
else:
|
||||
win32file.WriteFile(self._serial._port_handle, dataToWrite, self._overlappedWrite)
|
||||
|
||||
|
||||
def connectionLost(self, reason):
|
||||
"""
|
||||
Called when the serial port disconnects.
|
||||
|
||||
Will call C{connectionLost} on the protocol that is handling the
|
||||
serial data.
|
||||
"""
|
||||
self.reactor.removeEvent(self._overlappedRead.hEvent)
|
||||
self.reactor.removeEvent(self._overlappedWrite.hEvent)
|
||||
abstract.FileDescriptor.connectionLost(self, reason)
|
||||
self._serial.close()
|
||||
self.protocol.connectionLost(reason)
|
||||
@@ -0,0 +1,133 @@
|
||||
# -*- test-case-name: twisted.test.test_stdio -*-
|
||||
|
||||
"""
|
||||
Windows-specific implementation of the L{twisted.internet.stdio} interface.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
import win32api
|
||||
import os
|
||||
import msvcrt
|
||||
|
||||
from zope.interface import implementer
|
||||
|
||||
from twisted.internet.interfaces import (IHalfCloseableProtocol, ITransport,
|
||||
IConsumer, IPushProducer, IAddress)
|
||||
|
||||
from twisted.internet import _pollingfile, main
|
||||
from twisted.python.failure import Failure
|
||||
|
||||
@implementer(IAddress)
|
||||
class Win32PipeAddress(object):
|
||||
pass
|
||||
|
||||
|
||||
@implementer(ITransport, IConsumer, IPushProducer)
|
||||
class StandardIO(_pollingfile._PollingTimer):
|
||||
|
||||
disconnecting = False
|
||||
disconnected = False
|
||||
|
||||
def __init__(self, proto, reactor=None):
|
||||
"""
|
||||
Start talking to standard IO with the given protocol.
|
||||
|
||||
Also, put it stdin/stdout/stderr into binary mode.
|
||||
"""
|
||||
if reactor is None:
|
||||
from twisted.internet import reactor
|
||||
|
||||
for stdfd in range(0, 1, 2):
|
||||
msvcrt.setmode(stdfd, os.O_BINARY)
|
||||
|
||||
_pollingfile._PollingTimer.__init__(self, reactor)
|
||||
self.proto = proto
|
||||
|
||||
hstdin = win32api.GetStdHandle(win32api.STD_INPUT_HANDLE)
|
||||
hstdout = win32api.GetStdHandle(win32api.STD_OUTPUT_HANDLE)
|
||||
|
||||
self.stdin = _pollingfile._PollableReadPipe(
|
||||
hstdin, self.dataReceived, self.readConnectionLost)
|
||||
|
||||
self.stdout = _pollingfile._PollableWritePipe(
|
||||
hstdout, self.writeConnectionLost)
|
||||
|
||||
self._addPollableResource(self.stdin)
|
||||
self._addPollableResource(self.stdout)
|
||||
|
||||
self.proto.makeConnection(self)
|
||||
|
||||
|
||||
def dataReceived(self, data):
|
||||
self.proto.dataReceived(data)
|
||||
|
||||
|
||||
def readConnectionLost(self):
|
||||
if IHalfCloseableProtocol.providedBy(self.proto):
|
||||
self.proto.readConnectionLost()
|
||||
self.checkConnLost()
|
||||
|
||||
|
||||
def writeConnectionLost(self):
|
||||
if IHalfCloseableProtocol.providedBy(self.proto):
|
||||
self.proto.writeConnectionLost()
|
||||
self.checkConnLost()
|
||||
|
||||
connsLost = 0
|
||||
|
||||
|
||||
def checkConnLost(self):
|
||||
self.connsLost += 1
|
||||
if self.connsLost >= 2:
|
||||
self.disconnecting = True
|
||||
self.disconnected = True
|
||||
self.proto.connectionLost(Failure(main.CONNECTION_DONE))
|
||||
|
||||
# ITransport
|
||||
|
||||
def write(self, data):
|
||||
self.stdout.write(data)
|
||||
|
||||
|
||||
def writeSequence(self, seq):
|
||||
self.stdout.write(b''.join(seq))
|
||||
|
||||
|
||||
def loseConnection(self):
|
||||
self.disconnecting = True
|
||||
self.stdin.close()
|
||||
self.stdout.close()
|
||||
|
||||
|
||||
def getPeer(self):
|
||||
return Win32PipeAddress()
|
||||
|
||||
|
||||
def getHost(self):
|
||||
return Win32PipeAddress()
|
||||
|
||||
# IConsumer
|
||||
|
||||
def registerProducer(self, producer, streaming):
|
||||
return self.stdout.registerProducer(producer, streaming)
|
||||
|
||||
|
||||
def unregisterProducer(self):
|
||||
return self.stdout.unregisterProducer()
|
||||
|
||||
# def write() above
|
||||
|
||||
# IProducer
|
||||
|
||||
def stopProducing(self):
|
||||
self.stdin.stopProducing()
|
||||
|
||||
# IPushProducer
|
||||
|
||||
def pauseProducing(self):
|
||||
self.stdin.pauseProducing()
|
||||
|
||||
|
||||
def resumeProducing(self):
|
||||
self.stdin.resumeProducing()
|
||||
@@ -0,0 +1,546 @@
|
||||
# -*- test-case-name: twisted.test.test_abstract -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Support for generic select()able objects.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
from socket import AF_INET, AF_INET6, inet_pton, error
|
||||
|
||||
from zope.interface import implementer
|
||||
|
||||
# Twisted Imports
|
||||
from twisted.python.compat import unicode, lazyByteSlice, _PY3
|
||||
from twisted.python import reflect, failure
|
||||
from twisted.internet import interfaces, main
|
||||
|
||||
if _PY3:
|
||||
# Python 3.4+ can join bytes and memoryviews; using a
|
||||
# memoryview prevents the slice from copying
|
||||
def _concatenate(bObj, offset, bArray):
|
||||
return b''.join([memoryview(bObj)[offset:]] + bArray)
|
||||
else:
|
||||
from __builtin__ import buffer
|
||||
|
||||
def _concatenate(bObj, offset, bArray):
|
||||
# Avoid one extra string copy by using a buffer to limit what
|
||||
# we include in the result.
|
||||
return buffer(bObj, offset) + b"".join(bArray)
|
||||
|
||||
|
||||
|
||||
class _ConsumerMixin(object):
|
||||
"""
|
||||
L{IConsumer} implementations can mix this in to get C{registerProducer} and
|
||||
C{unregisterProducer} methods which take care of keeping track of a
|
||||
producer's state.
|
||||
|
||||
Subclasses must provide three attributes which L{_ConsumerMixin} will read
|
||||
but not write:
|
||||
|
||||
- connected: A C{bool} which is C{True} as long as the consumer has
|
||||
someplace to send bytes (for example, a TCP connection), and then
|
||||
C{False} when it no longer does.
|
||||
|
||||
- disconnecting: A C{bool} which is C{False} until something like
|
||||
L{ITransport.loseConnection} is called, indicating that the send buffer
|
||||
should be flushed and the connection lost afterwards. Afterwards,
|
||||
C{True}.
|
||||
|
||||
- disconnected: A C{bool} which is C{False} until the consumer no longer
|
||||
has a place to send bytes, then C{True}.
|
||||
|
||||
Subclasses must also override the C{startWriting} method.
|
||||
|
||||
@ivar producer: L{None} if no producer is registered, otherwise the
|
||||
registered producer.
|
||||
|
||||
@ivar producerPaused: A flag indicating whether the producer is currently
|
||||
paused.
|
||||
@type producerPaused: L{bool}
|
||||
|
||||
@ivar streamingProducer: A flag indicating whether the producer was
|
||||
registered as a streaming (ie push) producer or not (ie a pull
|
||||
producer). This will determine whether the consumer may ever need to
|
||||
pause and resume it, or if it can merely call C{resumeProducing} on it
|
||||
when buffer space is available.
|
||||
@ivar streamingProducer: C{bool} or C{int}
|
||||
|
||||
"""
|
||||
producer = None
|
||||
producerPaused = False
|
||||
streamingProducer = False
|
||||
|
||||
def startWriting(self):
|
||||
"""
|
||||
Override in a subclass to cause the reactor to monitor this selectable
|
||||
for write events. This will be called once in C{unregisterProducer} if
|
||||
C{loseConnection} has previously been called, so that the connection can
|
||||
actually close.
|
||||
"""
|
||||
raise NotImplementedError("%r did not implement startWriting")
|
||||
|
||||
|
||||
def registerProducer(self, producer, streaming):
|
||||
"""
|
||||
Register to receive data from a producer.
|
||||
|
||||
This sets this selectable to be a consumer for a producer. When this
|
||||
selectable runs out of data on a write() call, it will ask the producer
|
||||
to resumeProducing(). When the FileDescriptor's internal data buffer is
|
||||
filled, it will ask the producer to pauseProducing(). If the connection
|
||||
is lost, FileDescriptor calls producer's stopProducing() method.
|
||||
|
||||
If streaming is true, the producer should provide the IPushProducer
|
||||
interface. Otherwise, it is assumed that producer provides the
|
||||
IPullProducer interface. In this case, the producer won't be asked to
|
||||
pauseProducing(), but it has to be careful to write() data only when its
|
||||
resumeProducing() method is called.
|
||||
"""
|
||||
if self.producer is not None:
|
||||
raise RuntimeError(
|
||||
"Cannot register producer %s, because producer %s was never "
|
||||
"unregistered." % (producer, self.producer))
|
||||
if self.disconnected:
|
||||
producer.stopProducing()
|
||||
else:
|
||||
self.producer = producer
|
||||
self.streamingProducer = streaming
|
||||
if not streaming:
|
||||
producer.resumeProducing()
|
||||
|
||||
|
||||
def unregisterProducer(self):
|
||||
"""
|
||||
Stop consuming data from a producer, without disconnecting.
|
||||
"""
|
||||
self.producer = None
|
||||
if self.connected and self.disconnecting:
|
||||
self.startWriting()
|
||||
|
||||
|
||||
|
||||
@implementer(interfaces.ILoggingContext)
|
||||
class _LogOwner(object):
|
||||
"""
|
||||
Mixin to help implement L{interfaces.ILoggingContext} for transports which
|
||||
have a protocol, the log prefix of which should also appear in the
|
||||
transport's log prefix.
|
||||
"""
|
||||
|
||||
def _getLogPrefix(self, applicationObject):
|
||||
"""
|
||||
Determine the log prefix to use for messages related to
|
||||
C{applicationObject}, which may or may not be an
|
||||
L{interfaces.ILoggingContext} provider.
|
||||
|
||||
@return: A C{str} giving the log prefix to use.
|
||||
"""
|
||||
if interfaces.ILoggingContext.providedBy(applicationObject):
|
||||
return applicationObject.logPrefix()
|
||||
return applicationObject.__class__.__name__
|
||||
|
||||
|
||||
def logPrefix(self):
|
||||
"""
|
||||
Override this method to insert custom logging behavior. Its
|
||||
return value will be inserted in front of every line. It may
|
||||
be called more times than the number of output lines.
|
||||
"""
|
||||
return "-"
|
||||
|
||||
|
||||
|
||||
@implementer(
|
||||
interfaces.IPushProducer, interfaces.IReadWriteDescriptor,
|
||||
interfaces.IConsumer, interfaces.ITransport,
|
||||
interfaces.IHalfCloseableDescriptor)
|
||||
class FileDescriptor(_ConsumerMixin, _LogOwner):
|
||||
"""
|
||||
An object which can be operated on by select().
|
||||
|
||||
This is an abstract superclass of all objects which may be notified when
|
||||
they are readable or writable; e.g. they have a file-descriptor that is
|
||||
valid to be passed to select(2).
|
||||
"""
|
||||
connected = 0
|
||||
disconnected = 0
|
||||
disconnecting = 0
|
||||
_writeDisconnecting = False
|
||||
_writeDisconnected = False
|
||||
dataBuffer = b""
|
||||
offset = 0
|
||||
|
||||
SEND_LIMIT = 128*1024
|
||||
|
||||
def __init__(self, reactor=None):
|
||||
"""
|
||||
@param reactor: An L{IReactorFDSet} provider which this descriptor will
|
||||
use to get readable and writeable event notifications. If no value
|
||||
is given, the global reactor will be used.
|
||||
"""
|
||||
if not reactor:
|
||||
from twisted.internet import reactor
|
||||
self.reactor = reactor
|
||||
self._tempDataBuffer = [] # will be added to dataBuffer in doWrite
|
||||
self._tempDataLen = 0
|
||||
|
||||
|
||||
def connectionLost(self, reason):
|
||||
"""The connection was lost.
|
||||
|
||||
This is called when the connection on a selectable object has been
|
||||
lost. It will be called whether the connection was closed explicitly,
|
||||
an exception occurred in an event handler, or the other end of the
|
||||
connection closed it first.
|
||||
|
||||
Clean up state here, but make sure to call back up to FileDescriptor.
|
||||
"""
|
||||
self.disconnected = 1
|
||||
self.connected = 0
|
||||
if self.producer is not None:
|
||||
self.producer.stopProducing()
|
||||
self.producer = None
|
||||
self.stopReading()
|
||||
self.stopWriting()
|
||||
|
||||
|
||||
def writeSomeData(self, data):
|
||||
"""
|
||||
Write as much as possible of the given data, immediately.
|
||||
|
||||
This is called to invoke the lower-level writing functionality, such
|
||||
as a socket's send() method, or a file's write(); this method
|
||||
returns an integer or an exception. If an integer, it is the number
|
||||
of bytes written (possibly zero); if an exception, it indicates the
|
||||
connection was lost.
|
||||
"""
|
||||
raise NotImplementedError("%s does not implement writeSomeData" %
|
||||
reflect.qual(self.__class__))
|
||||
|
||||
|
||||
def doRead(self):
|
||||
"""
|
||||
Called when data is available for reading.
|
||||
|
||||
Subclasses must override this method. The result will be interpreted
|
||||
in the same way as a result of doWrite().
|
||||
"""
|
||||
raise NotImplementedError("%s does not implement doRead" %
|
||||
reflect.qual(self.__class__))
|
||||
|
||||
def doWrite(self):
|
||||
"""
|
||||
Called when data can be written.
|
||||
|
||||
@return: L{None} on success, an exception or a negative integer on
|
||||
failure.
|
||||
|
||||
@see: L{twisted.internet.interfaces.IWriteDescriptor.doWrite}.
|
||||
"""
|
||||
if len(self.dataBuffer) - self.offset < self.SEND_LIMIT:
|
||||
# If there is currently less than SEND_LIMIT bytes left to send
|
||||
# in the string, extend it with the array data.
|
||||
self.dataBuffer = _concatenate(
|
||||
self.dataBuffer, self.offset, self._tempDataBuffer)
|
||||
self.offset = 0
|
||||
self._tempDataBuffer = []
|
||||
self._tempDataLen = 0
|
||||
|
||||
# Send as much data as you can.
|
||||
if self.offset:
|
||||
l = self.writeSomeData(lazyByteSlice(self.dataBuffer, self.offset))
|
||||
else:
|
||||
l = self.writeSomeData(self.dataBuffer)
|
||||
|
||||
# There is no writeSomeData implementation in Twisted which returns
|
||||
# < 0, but the documentation for writeSomeData used to claim negative
|
||||
# integers meant connection lost. Keep supporting this here,
|
||||
# although it may be worth deprecating and removing at some point.
|
||||
if isinstance(l, Exception) or l < 0:
|
||||
return l
|
||||
self.offset += l
|
||||
# If there is nothing left to send,
|
||||
if self.offset == len(self.dataBuffer) and not self._tempDataLen:
|
||||
self.dataBuffer = b""
|
||||
self.offset = 0
|
||||
# stop writing.
|
||||
self.stopWriting()
|
||||
# If I've got a producer who is supposed to supply me with data,
|
||||
if self.producer is not None and ((not self.streamingProducer)
|
||||
or self.producerPaused):
|
||||
# tell them to supply some more.
|
||||
self.producerPaused = False
|
||||
self.producer.resumeProducing()
|
||||
elif self.disconnecting:
|
||||
# But if I was previously asked to let the connection die, do
|
||||
# so.
|
||||
return self._postLoseConnection()
|
||||
elif self._writeDisconnecting:
|
||||
# I was previously asked to half-close the connection. We
|
||||
# set _writeDisconnected before calling handler, in case the
|
||||
# handler calls loseConnection(), which will want to check for
|
||||
# this attribute.
|
||||
self._writeDisconnected = True
|
||||
result = self._closeWriteConnection()
|
||||
return result
|
||||
return None
|
||||
|
||||
def _postLoseConnection(self):
|
||||
"""Called after a loseConnection(), when all data has been written.
|
||||
|
||||
Whatever this returns is then returned by doWrite.
|
||||
"""
|
||||
# default implementation, telling reactor we're finished
|
||||
return main.CONNECTION_DONE
|
||||
|
||||
def _closeWriteConnection(self):
|
||||
# override in subclasses
|
||||
pass
|
||||
|
||||
def writeConnectionLost(self, reason):
|
||||
# in current code should never be called
|
||||
self.connectionLost(reason)
|
||||
|
||||
def readConnectionLost(self, reason):
|
||||
# override in subclasses
|
||||
self.connectionLost(reason)
|
||||
|
||||
|
||||
def _isSendBufferFull(self):
|
||||
"""
|
||||
Determine whether the user-space send buffer for this transport is full
|
||||
or not.
|
||||
|
||||
When the buffer contains more than C{self.bufferSize} bytes, it is
|
||||
considered full. This might be improved by considering the size of the
|
||||
kernel send buffer and how much of it is free.
|
||||
|
||||
@return: C{True} if it is full, C{False} otherwise.
|
||||
"""
|
||||
return len(self.dataBuffer) + self._tempDataLen > self.bufferSize
|
||||
|
||||
|
||||
def _maybePauseProducer(self):
|
||||
"""
|
||||
Possibly pause a producer, if there is one and the send buffer is full.
|
||||
"""
|
||||
# If we are responsible for pausing our producer,
|
||||
if self.producer is not None and self.streamingProducer:
|
||||
# and our buffer is full,
|
||||
if self._isSendBufferFull():
|
||||
# pause it.
|
||||
self.producerPaused = True
|
||||
self.producer.pauseProducing()
|
||||
|
||||
|
||||
def write(self, data):
|
||||
"""Reliably write some data.
|
||||
|
||||
The data is buffered until the underlying file descriptor is ready
|
||||
for writing. If there is more than C{self.bufferSize} data in the
|
||||
buffer and this descriptor has a registered streaming producer, its
|
||||
C{pauseProducing()} method will be called.
|
||||
"""
|
||||
if isinstance(data, unicode): # no, really, I mean it
|
||||
raise TypeError("Data must not be unicode")
|
||||
if not self.connected or self._writeDisconnected:
|
||||
return
|
||||
if data:
|
||||
self._tempDataBuffer.append(data)
|
||||
self._tempDataLen += len(data)
|
||||
self._maybePauseProducer()
|
||||
self.startWriting()
|
||||
|
||||
|
||||
def writeSequence(self, iovec):
|
||||
"""
|
||||
Reliably write a sequence of data.
|
||||
|
||||
Currently, this is a convenience method roughly equivalent to::
|
||||
|
||||
for chunk in iovec:
|
||||
fd.write(chunk)
|
||||
|
||||
It may have a more efficient implementation at a later time or in a
|
||||
different reactor.
|
||||
|
||||
As with the C{write()} method, if a buffer size limit is reached and a
|
||||
streaming producer is registered, it will be paused until the buffered
|
||||
data is written to the underlying file descriptor.
|
||||
"""
|
||||
for i in iovec:
|
||||
if isinstance(i, unicode): # no, really, I mean it
|
||||
raise TypeError("Data must not be unicode")
|
||||
if not self.connected or not iovec or self._writeDisconnected:
|
||||
return
|
||||
self._tempDataBuffer.extend(iovec)
|
||||
for i in iovec:
|
||||
self._tempDataLen += len(i)
|
||||
self._maybePauseProducer()
|
||||
self.startWriting()
|
||||
|
||||
|
||||
def loseConnection(self, _connDone=failure.Failure(main.CONNECTION_DONE)):
|
||||
"""Close the connection at the next available opportunity.
|
||||
|
||||
Call this to cause this FileDescriptor to lose its connection. It will
|
||||
first write any data that it has buffered.
|
||||
|
||||
If there is data buffered yet to be written, this method will cause the
|
||||
transport to lose its connection as soon as it's done flushing its
|
||||
write buffer. If you have a producer registered, the connection won't
|
||||
be closed until the producer is finished. Therefore, make sure you
|
||||
unregister your producer when it's finished, or the connection will
|
||||
never close.
|
||||
"""
|
||||
|
||||
if self.connected and not self.disconnecting:
|
||||
if self._writeDisconnected:
|
||||
# doWrite won't trigger the connection close anymore
|
||||
self.stopReading()
|
||||
self.stopWriting()
|
||||
self.connectionLost(_connDone)
|
||||
else:
|
||||
self.stopReading()
|
||||
self.startWriting()
|
||||
self.disconnecting = 1
|
||||
|
||||
def loseWriteConnection(self):
|
||||
self._writeDisconnecting = True
|
||||
self.startWriting()
|
||||
|
||||
def stopReading(self):
|
||||
"""Stop waiting for read availability.
|
||||
|
||||
Call this to remove this selectable from being notified when it is
|
||||
ready for reading.
|
||||
"""
|
||||
self.reactor.removeReader(self)
|
||||
|
||||
def stopWriting(self):
|
||||
"""Stop waiting for write availability.
|
||||
|
||||
Call this to remove this selectable from being notified when it is ready
|
||||
for writing.
|
||||
"""
|
||||
self.reactor.removeWriter(self)
|
||||
|
||||
def startReading(self):
|
||||
"""Start waiting for read availability.
|
||||
"""
|
||||
self.reactor.addReader(self)
|
||||
|
||||
def startWriting(self):
|
||||
"""Start waiting for write availability.
|
||||
|
||||
Call this to have this FileDescriptor be notified whenever it is ready for
|
||||
writing.
|
||||
"""
|
||||
self.reactor.addWriter(self)
|
||||
|
||||
# Producer/consumer implementation
|
||||
|
||||
# first, the consumer stuff. This requires no additional work, as
|
||||
# any object you can write to can be a consumer, really.
|
||||
|
||||
producer = None
|
||||
bufferSize = 2**2**2**2
|
||||
|
||||
def stopConsuming(self):
|
||||
"""Stop consuming data.
|
||||
|
||||
This is called when a producer has lost its connection, to tell the
|
||||
consumer to go lose its connection (and break potential circular
|
||||
references).
|
||||
"""
|
||||
self.unregisterProducer()
|
||||
self.loseConnection()
|
||||
|
||||
# producer interface implementation
|
||||
|
||||
def resumeProducing(self):
|
||||
if self.connected and not self.disconnecting:
|
||||
self.startReading()
|
||||
|
||||
def pauseProducing(self):
|
||||
self.stopReading()
|
||||
|
||||
def stopProducing(self):
|
||||
self.loseConnection()
|
||||
|
||||
|
||||
def fileno(self):
|
||||
"""File Descriptor number for select().
|
||||
|
||||
This method must be overridden or assigned in subclasses to
|
||||
indicate a valid file descriptor for the operating system.
|
||||
"""
|
||||
return -1
|
||||
|
||||
|
||||
|
||||
def isIPAddress(addr, family=AF_INET):
|
||||
"""
|
||||
Determine whether the given string represents an IP address of the given
|
||||
family; by default, an IPv4 address.
|
||||
|
||||
@type addr: C{str}
|
||||
@param addr: A string which may or may not be the decimal dotted
|
||||
representation of an IPv4 address.
|
||||
|
||||
@param family: The address family to test for; one of the C{AF_*} constants
|
||||
from the L{socket} module. (This parameter has only been available
|
||||
since Twisted 17.1.0; previously L{isIPAddress} could only test for IPv4
|
||||
addresses.)
|
||||
@type family: C{int}
|
||||
|
||||
@rtype: C{bool}
|
||||
@return: C{True} if C{addr} represents an IPv4 address, C{False} otherwise.
|
||||
"""
|
||||
if isinstance(addr, bytes):
|
||||
try:
|
||||
addr = addr.decode("ascii")
|
||||
except UnicodeDecodeError:
|
||||
return False
|
||||
if family == AF_INET6:
|
||||
# On some platforms, inet_ntop fails unless the scope ID is valid; this
|
||||
# is a test for whether the given string *is* an IP address, so strip
|
||||
# any potential scope ID before checking.
|
||||
addr = addr.split(u"%", 1)[0]
|
||||
elif family == AF_INET:
|
||||
# On Windows, where 3.5+ implement inet_pton, "0" is considered a valid
|
||||
# IPv4 address, but we want to ensure we have all 4 segments.
|
||||
if addr.count(u".") != 3:
|
||||
return False
|
||||
else:
|
||||
raise ValueError("unknown address family {!r}".format(family))
|
||||
try:
|
||||
# This might be a native implementation or the one from
|
||||
# twisted.python.compat.
|
||||
inet_pton(family, addr)
|
||||
except (ValueError, error):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
|
||||
def isIPv6Address(addr):
|
||||
"""
|
||||
Determine whether the given string represents an IPv6 address.
|
||||
|
||||
@param addr: A string which may or may not be the hex
|
||||
representation of an IPv6 address.
|
||||
@type addr: C{str}
|
||||
|
||||
@return: C{True} if C{addr} represents an IPv6 address, C{False}
|
||||
otherwise.
|
||||
@rtype: C{bool}
|
||||
"""
|
||||
return isIPAddress(addr, AF_INET6)
|
||||
|
||||
|
||||
__all__ = ["FileDescriptor", "isIPAddress", "isIPv6Address"]
|
||||
@@ -0,0 +1,249 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
An epoll() based implementation of the twisted main loop.
|
||||
|
||||
To install the event loop (and you should do this before any connections,
|
||||
listeners or connectors are added)::
|
||||
|
||||
from twisted.internet import epollreactor
|
||||
epollreactor.install()
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
from select import epoll, EPOLLHUP, EPOLLERR, EPOLLIN, EPOLLOUT
|
||||
import errno
|
||||
|
||||
from zope.interface import implementer
|
||||
|
||||
from twisted.internet.interfaces import IReactorFDSet
|
||||
|
||||
from twisted.python import log
|
||||
from twisted.internet import posixbase
|
||||
|
||||
|
||||
@implementer(IReactorFDSet)
|
||||
class EPollReactor(posixbase.PosixReactorBase, posixbase._PollLikeMixin):
|
||||
"""
|
||||
A reactor that uses epoll(7).
|
||||
|
||||
@ivar _poller: A C{epoll} which will be used to check for I/O
|
||||
readiness.
|
||||
|
||||
@ivar _selectables: A dictionary mapping integer file descriptors to
|
||||
instances of C{FileDescriptor} which have been registered with the
|
||||
reactor. All C{FileDescriptors} which are currently receiving read or
|
||||
write readiness notifications will be present as values in this
|
||||
dictionary.
|
||||
|
||||
@ivar _reads: A set containing integer file descriptors. Values in this
|
||||
set will be registered with C{_poller} for read readiness notifications
|
||||
which will be dispatched to the corresponding C{FileDescriptor}
|
||||
instances in C{_selectables}.
|
||||
|
||||
@ivar _writes: A set containing integer file descriptors. Values in this
|
||||
set will be registered with C{_poller} for write readiness
|
||||
notifications which will be dispatched to the corresponding
|
||||
C{FileDescriptor} instances in C{_selectables}.
|
||||
|
||||
@ivar _continuousPolling: A L{_ContinuousPolling} instance, used to handle
|
||||
file descriptors (e.g. filesystem files) that are not supported by
|
||||
C{epoll(7)}.
|
||||
"""
|
||||
|
||||
# Attributes for _PollLikeMixin
|
||||
_POLL_DISCONNECTED = (EPOLLHUP | EPOLLERR)
|
||||
_POLL_IN = EPOLLIN
|
||||
_POLL_OUT = EPOLLOUT
|
||||
|
||||
def __init__(self):
|
||||
"""
|
||||
Initialize epoll object, file descriptor tracking dictionaries, and the
|
||||
base class.
|
||||
"""
|
||||
# Create the poller we're going to use. The 1024 here is just a hint
|
||||
# to the kernel, it is not a hard maximum. After Linux 2.6.8, the size
|
||||
# argument is completely ignored.
|
||||
self._poller = epoll(1024)
|
||||
self._reads = set()
|
||||
self._writes = set()
|
||||
self._selectables = {}
|
||||
self._continuousPolling = posixbase._ContinuousPolling(self)
|
||||
posixbase.PosixReactorBase.__init__(self)
|
||||
|
||||
|
||||
def _add(self, xer, primary, other, selectables, event, antievent):
|
||||
"""
|
||||
Private method for adding a descriptor from the event loop.
|
||||
|
||||
It takes care of adding it if new or modifying it if already added
|
||||
for another state (read -> read/write for example).
|
||||
"""
|
||||
fd = xer.fileno()
|
||||
if fd not in primary:
|
||||
flags = event
|
||||
# epoll_ctl can raise all kinds of IOErrors, and every one
|
||||
# indicates a bug either in the reactor or application-code.
|
||||
# Let them all through so someone sees a traceback and fixes
|
||||
# something. We'll do the same thing for every other call to
|
||||
# this method in this file.
|
||||
if fd in other:
|
||||
flags |= antievent
|
||||
self._poller.modify(fd, flags)
|
||||
else:
|
||||
self._poller.register(fd, flags)
|
||||
|
||||
# Update our own tracking state *only* after the epoll call has
|
||||
# succeeded. Otherwise we may get out of sync.
|
||||
primary.add(fd)
|
||||
selectables[fd] = xer
|
||||
|
||||
|
||||
def addReader(self, reader):
|
||||
"""
|
||||
Add a FileDescriptor for notification of data available to read.
|
||||
"""
|
||||
try:
|
||||
self._add(reader, self._reads, self._writes, self._selectables,
|
||||
EPOLLIN, EPOLLOUT)
|
||||
except IOError as e:
|
||||
if e.errno == errno.EPERM:
|
||||
# epoll(7) doesn't support certain file descriptors,
|
||||
# e.g. filesystem files, so for those we just poll
|
||||
# continuously:
|
||||
self._continuousPolling.addReader(reader)
|
||||
else:
|
||||
raise
|
||||
|
||||
|
||||
def addWriter(self, writer):
|
||||
"""
|
||||
Add a FileDescriptor for notification of data available to write.
|
||||
"""
|
||||
try:
|
||||
self._add(writer, self._writes, self._reads, self._selectables,
|
||||
EPOLLOUT, EPOLLIN)
|
||||
except IOError as e:
|
||||
if e.errno == errno.EPERM:
|
||||
# epoll(7) doesn't support certain file descriptors,
|
||||
# e.g. filesystem files, so for those we just poll
|
||||
# continuously:
|
||||
self._continuousPolling.addWriter(writer)
|
||||
else:
|
||||
raise
|
||||
|
||||
|
||||
def _remove(self, xer, primary, other, selectables, event, antievent):
|
||||
"""
|
||||
Private method for removing a descriptor from the event loop.
|
||||
|
||||
It does the inverse job of _add, and also add a check in case of the fd
|
||||
has gone away.
|
||||
"""
|
||||
fd = xer.fileno()
|
||||
if fd == -1:
|
||||
for fd, fdes in selectables.items():
|
||||
if xer is fdes:
|
||||
break
|
||||
else:
|
||||
return
|
||||
if fd in primary:
|
||||
if fd in other:
|
||||
flags = antievent
|
||||
# See comment above modify call in _add.
|
||||
self._poller.modify(fd, flags)
|
||||
else:
|
||||
del selectables[fd]
|
||||
# See comment above _control call in _add.
|
||||
self._poller.unregister(fd)
|
||||
primary.remove(fd)
|
||||
|
||||
|
||||
def removeReader(self, reader):
|
||||
"""
|
||||
Remove a Selectable for notification of data available to read.
|
||||
"""
|
||||
if self._continuousPolling.isReading(reader):
|
||||
self._continuousPolling.removeReader(reader)
|
||||
return
|
||||
self._remove(reader, self._reads, self._writes, self._selectables,
|
||||
EPOLLIN, EPOLLOUT)
|
||||
|
||||
|
||||
def removeWriter(self, writer):
|
||||
"""
|
||||
Remove a Selectable for notification of data available to write.
|
||||
"""
|
||||
if self._continuousPolling.isWriting(writer):
|
||||
self._continuousPolling.removeWriter(writer)
|
||||
return
|
||||
self._remove(writer, self._writes, self._reads, self._selectables,
|
||||
EPOLLOUT, EPOLLIN)
|
||||
|
||||
|
||||
def removeAll(self):
|
||||
"""
|
||||
Remove all selectables, and return a list of them.
|
||||
"""
|
||||
return (self._removeAll(
|
||||
[self._selectables[fd] for fd in self._reads],
|
||||
[self._selectables[fd] for fd in self._writes]) +
|
||||
self._continuousPolling.removeAll())
|
||||
|
||||
|
||||
def getReaders(self):
|
||||
return ([self._selectables[fd] for fd in self._reads] +
|
||||
self._continuousPolling.getReaders())
|
||||
|
||||
|
||||
def getWriters(self):
|
||||
return ([self._selectables[fd] for fd in self._writes] +
|
||||
self._continuousPolling.getWriters())
|
||||
|
||||
|
||||
def doPoll(self, timeout):
|
||||
"""
|
||||
Poll the poller for new events.
|
||||
"""
|
||||
if timeout is None:
|
||||
timeout = -1 # Wait indefinitely.
|
||||
|
||||
try:
|
||||
# Limit the number of events to the number of io objects we're
|
||||
# currently tracking (because that's maybe a good heuristic) and
|
||||
# the amount of time we block to the value specified by our
|
||||
# caller.
|
||||
l = self._poller.poll(timeout, len(self._selectables))
|
||||
except IOError as err:
|
||||
if err.errno == errno.EINTR:
|
||||
return
|
||||
# See epoll_wait(2) for documentation on the other conditions
|
||||
# under which this can fail. They can only be due to a serious
|
||||
# programming error on our part, so let's just announce them
|
||||
# loudly.
|
||||
raise
|
||||
|
||||
_drdw = self._doReadOrWrite
|
||||
for fd, event in l:
|
||||
try:
|
||||
selectable = self._selectables[fd]
|
||||
except KeyError:
|
||||
pass
|
||||
else:
|
||||
log.callWithLogger(selectable, _drdw, selectable, fd, event)
|
||||
|
||||
doIteration = doPoll
|
||||
|
||||
|
||||
def install():
|
||||
"""
|
||||
Install the epoll() reactor.
|
||||
"""
|
||||
p = EPollReactor()
|
||||
from twisted.internet.main import installReactor
|
||||
installReactor(p)
|
||||
|
||||
|
||||
__all__ = ["EPollReactor", "install"]
|
||||
@@ -0,0 +1,517 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Exceptions and errors for use in twisted.internet modules.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
import socket
|
||||
|
||||
from twisted.python import deprecate
|
||||
from incremental import Version
|
||||
|
||||
|
||||
|
||||
class BindError(Exception):
|
||||
"""An error occurred binding to an interface"""
|
||||
|
||||
def __str__(self):
|
||||
s = self.__doc__
|
||||
if self.args:
|
||||
s = '%s: %s' % (s, ' '.join(self.args))
|
||||
s = '%s.' % s
|
||||
return s
|
||||
|
||||
|
||||
|
||||
class CannotListenError(BindError):
|
||||
"""
|
||||
This gets raised by a call to startListening, when the object cannotstart
|
||||
listening.
|
||||
|
||||
@ivar interface: the interface I tried to listen on
|
||||
@ivar port: the port I tried to listen on
|
||||
@ivar socketError: the exception I got when I tried to listen
|
||||
@type socketError: L{socket.error}
|
||||
"""
|
||||
def __init__(self, interface, port, socketError):
|
||||
BindError.__init__(self, interface, port, socketError)
|
||||
self.interface = interface
|
||||
self.port = port
|
||||
self.socketError = socketError
|
||||
|
||||
def __str__(self):
|
||||
iface = self.interface or 'any'
|
||||
return "Couldn't listen on %s:%s: %s." % (iface, self.port,
|
||||
self.socketError)
|
||||
|
||||
|
||||
|
||||
class MulticastJoinError(Exception):
|
||||
"""
|
||||
An attempt to join a multicast group failed.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class MessageLengthError(Exception):
|
||||
"""Message is too long to send"""
|
||||
|
||||
def __str__(self):
|
||||
s = self.__doc__
|
||||
if self.args:
|
||||
s = '%s: %s' % (s, ' '.join(self.args))
|
||||
s = '%s.' % s
|
||||
return s
|
||||
|
||||
|
||||
|
||||
class DNSLookupError(IOError):
|
||||
"""DNS lookup failed"""
|
||||
|
||||
def __str__(self):
|
||||
s = self.__doc__
|
||||
if self.args:
|
||||
s = '%s: %s' % (s, ' '.join(self.args))
|
||||
s = '%s.' % s
|
||||
return s
|
||||
|
||||
|
||||
|
||||
class ConnectInProgressError(Exception):
|
||||
"""A connect operation was started and isn't done yet."""
|
||||
|
||||
|
||||
# connection errors
|
||||
|
||||
class ConnectError(Exception):
|
||||
"""An error occurred while connecting"""
|
||||
|
||||
def __init__(self, osError=None, string=""):
|
||||
self.osError = osError
|
||||
Exception.__init__(self, string)
|
||||
|
||||
def __str__(self):
|
||||
s = self.__doc__ or self.__class__.__name__
|
||||
if self.osError:
|
||||
s = '%s: %s' % (s, self.osError)
|
||||
if self.args[0]:
|
||||
s = '%s: %s' % (s, self.args[0])
|
||||
s = '%s.' % s
|
||||
return s
|
||||
|
||||
|
||||
|
||||
class ConnectBindError(ConnectError):
|
||||
"""Couldn't bind"""
|
||||
|
||||
|
||||
|
||||
class UnknownHostError(ConnectError):
|
||||
"""Hostname couldn't be looked up"""
|
||||
|
||||
|
||||
|
||||
class NoRouteError(ConnectError):
|
||||
"""No route to host"""
|
||||
|
||||
|
||||
|
||||
class ConnectionRefusedError(ConnectError):
|
||||
"""Connection was refused by other side"""
|
||||
|
||||
|
||||
|
||||
class TCPTimedOutError(ConnectError):
|
||||
"""TCP connection timed out"""
|
||||
|
||||
|
||||
|
||||
class BadFileError(ConnectError):
|
||||
"""File used for UNIX socket is no good"""
|
||||
|
||||
|
||||
|
||||
class ServiceNameUnknownError(ConnectError):
|
||||
"""Service name given as port is unknown"""
|
||||
|
||||
|
||||
|
||||
class UserError(ConnectError):
|
||||
"""User aborted connection"""
|
||||
|
||||
|
||||
|
||||
class TimeoutError(UserError):
|
||||
"""User timeout caused connection failure"""
|
||||
|
||||
|
||||
|
||||
class SSLError(ConnectError):
|
||||
"""An SSL error occurred"""
|
||||
|
||||
|
||||
|
||||
class VerifyError(Exception):
|
||||
"""Could not verify something that was supposed to be signed.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class PeerVerifyError(VerifyError):
|
||||
"""The peer rejected our verify error.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class CertificateError(Exception):
|
||||
"""
|
||||
We did not find a certificate where we expected to find one.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
try:
|
||||
import errno
|
||||
errnoMapping = {
|
||||
errno.ENETUNREACH: NoRouteError,
|
||||
errno.ECONNREFUSED: ConnectionRefusedError,
|
||||
errno.ETIMEDOUT: TCPTimedOutError,
|
||||
}
|
||||
if hasattr(errno, "WSAECONNREFUSED"):
|
||||
errnoMapping[errno.WSAECONNREFUSED] = ConnectionRefusedError
|
||||
errnoMapping[errno.WSAENETUNREACH] = NoRouteError
|
||||
except ImportError:
|
||||
errnoMapping = {}
|
||||
|
||||
|
||||
|
||||
def getConnectError(e):
|
||||
"""Given a socket exception, return connection error."""
|
||||
if isinstance(e, Exception):
|
||||
args = e.args
|
||||
else:
|
||||
args = e
|
||||
try:
|
||||
number, string = args
|
||||
except ValueError:
|
||||
return ConnectError(string=e)
|
||||
|
||||
if hasattr(socket, 'gaierror') and isinstance(e, socket.gaierror):
|
||||
# Only works in 2.2 in newer. Really that means always; #5978 covers
|
||||
# this and other weirdnesses in this function.
|
||||
klass = UnknownHostError
|
||||
else:
|
||||
klass = errnoMapping.get(number, ConnectError)
|
||||
return klass(number, string)
|
||||
|
||||
|
||||
|
||||
class ConnectionClosed(Exception):
|
||||
"""
|
||||
Connection was closed, whether cleanly or non-cleanly.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class ConnectionLost(ConnectionClosed):
|
||||
"""Connection to the other side was lost in a non-clean fashion"""
|
||||
|
||||
def __str__(self):
|
||||
s = self.__doc__.strip().splitlines()[0]
|
||||
if self.args:
|
||||
s = '%s: %s' % (s, ' '.join(self.args))
|
||||
s = '%s.' % s
|
||||
return s
|
||||
|
||||
|
||||
|
||||
class ConnectionAborted(ConnectionLost):
|
||||
"""
|
||||
Connection was aborted locally, using
|
||||
L{twisted.internet.interfaces.ITCPTransport.abortConnection}.
|
||||
|
||||
@since: 11.1
|
||||
"""
|
||||
|
||||
def __str__(self):
|
||||
s = [(
|
||||
"Connection was aborted locally using"
|
||||
" ITCPTransport.abortConnection"
|
||||
)]
|
||||
if self.args:
|
||||
s.append(': ')
|
||||
s.append(' '.join(self.args))
|
||||
s.append('.')
|
||||
return ''.join(s)
|
||||
|
||||
|
||||
|
||||
class ConnectionDone(ConnectionClosed):
|
||||
"""Connection was closed cleanly"""
|
||||
|
||||
def __str__(self):
|
||||
s = self.__doc__
|
||||
if self.args:
|
||||
s = '%s: %s' % (s, ' '.join(self.args))
|
||||
s = '%s.' % s
|
||||
return s
|
||||
|
||||
|
||||
|
||||
class FileDescriptorOverrun(ConnectionLost):
|
||||
"""
|
||||
A mis-use of L{IUNIXTransport.sendFileDescriptor} caused the connection to
|
||||
be closed.
|
||||
|
||||
Each file descriptor sent using C{sendFileDescriptor} must be associated
|
||||
with at least one byte sent using L{ITransport.write}. If at any point
|
||||
fewer bytes have been written than file descriptors have been sent, the
|
||||
connection is closed with this exception.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class ConnectionFdescWentAway(ConnectionLost):
|
||||
"""Uh""" #TODO
|
||||
|
||||
|
||||
|
||||
class AlreadyCalled(ValueError):
|
||||
"""Tried to cancel an already-called event"""
|
||||
|
||||
def __str__(self):
|
||||
s = self.__doc__
|
||||
if self.args:
|
||||
s = '%s: %s' % (s, ' '.join(self.args))
|
||||
s = '%s.' % s
|
||||
return s
|
||||
|
||||
|
||||
|
||||
class AlreadyCancelled(ValueError):
|
||||
"""Tried to cancel an already-cancelled event"""
|
||||
|
||||
def __str__(self):
|
||||
s = self.__doc__
|
||||
if self.args:
|
||||
s = '%s: %s' % (s, ' '.join(self.args))
|
||||
s = '%s.' % s
|
||||
return s
|
||||
|
||||
|
||||
|
||||
class PotentialZombieWarning(Warning):
|
||||
"""
|
||||
Emitted when L{IReactorProcess.spawnProcess} is called in a way which may
|
||||
result in termination of the created child process not being reported.
|
||||
|
||||
Deprecated in Twisted 10.0.
|
||||
"""
|
||||
MESSAGE = (
|
||||
"spawnProcess called, but the SIGCHLD handler is not "
|
||||
"installed. This probably means you have not yet "
|
||||
"called reactor.run, or called "
|
||||
"reactor.run(installSignalHandler=0). You will probably "
|
||||
"never see this process finish, and it may become a "
|
||||
"zombie process.")
|
||||
|
||||
deprecate.deprecatedModuleAttribute(
|
||||
Version("Twisted", 10, 0, 0),
|
||||
"There is no longer any potential for zombie process.",
|
||||
__name__,
|
||||
"PotentialZombieWarning")
|
||||
|
||||
|
||||
|
||||
class ProcessDone(ConnectionDone):
|
||||
"""A process has ended without apparent errors"""
|
||||
|
||||
def __init__(self, status):
|
||||
Exception.__init__(self, "process finished with exit code 0")
|
||||
self.exitCode = 0
|
||||
self.signal = None
|
||||
self.status = status
|
||||
|
||||
|
||||
|
||||
class ProcessTerminated(ConnectionLost):
|
||||
"""
|
||||
A process has ended with a probable error condition
|
||||
|
||||
@ivar exitCode: See L{__init__}
|
||||
@ivar signal: See L{__init__}
|
||||
@ivar status: See L{__init__}
|
||||
"""
|
||||
def __init__(self, exitCode=None, signal=None, status=None):
|
||||
"""
|
||||
@param exitCode: The exit status of the process. This is roughly like
|
||||
the value you might pass to L{os.exit}. This is L{None} if the
|
||||
process exited due to a signal.
|
||||
@type exitCode: L{int} or L{None}
|
||||
|
||||
@param signal: The exit signal of the process. This is L{None} if the
|
||||
process did not exit due to a signal.
|
||||
@type signal: L{int} or L{None}
|
||||
|
||||
@param status: The exit code of the process. This is a platform
|
||||
specific combination of the exit code and the exit signal. See
|
||||
L{os.WIFEXITED} and related functions.
|
||||
@type status: L{int}
|
||||
"""
|
||||
self.exitCode = exitCode
|
||||
self.signal = signal
|
||||
self.status = status
|
||||
s = "process ended"
|
||||
if exitCode is not None: s = s + " with exit code %s" % exitCode
|
||||
if signal is not None: s = s + " by signal %s" % signal
|
||||
Exception.__init__(self, s)
|
||||
|
||||
|
||||
|
||||
class ProcessExitedAlready(Exception):
|
||||
"""
|
||||
The process has already exited and the operation requested can no longer
|
||||
be performed.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class NotConnectingError(RuntimeError):
|
||||
"""The Connector was not connecting when it was asked to stop connecting"""
|
||||
|
||||
def __str__(self):
|
||||
s = self.__doc__
|
||||
if self.args:
|
||||
s = '%s: %s' % (s, ' '.join(self.args))
|
||||
s = '%s.' % s
|
||||
return s
|
||||
|
||||
|
||||
|
||||
class NotListeningError(RuntimeError):
|
||||
"""The Port was not listening when it was asked to stop listening"""
|
||||
|
||||
def __str__(self):
|
||||
s = self.__doc__
|
||||
if self.args:
|
||||
s = '%s: %s' % (s, ' '.join(self.args))
|
||||
s = '%s.' % s
|
||||
return s
|
||||
|
||||
|
||||
|
||||
class ReactorNotRunning(RuntimeError):
|
||||
"""
|
||||
Error raised when trying to stop a reactor which is not running.
|
||||
"""
|
||||
|
||||
|
||||
class ReactorNotRestartable(RuntimeError):
|
||||
"""
|
||||
Error raised when trying to run a reactor which was stopped.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class ReactorAlreadyRunning(RuntimeError):
|
||||
"""
|
||||
Error raised when trying to start the reactor multiple times.
|
||||
"""
|
||||
|
||||
|
||||
class ReactorAlreadyInstalledError(AssertionError):
|
||||
"""
|
||||
Could not install reactor because one is already installed.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class ConnectingCancelledError(Exception):
|
||||
"""
|
||||
An C{Exception} that will be raised when an L{IStreamClientEndpoint} is
|
||||
cancelled before it connects.
|
||||
|
||||
@ivar address: The L{IAddress} that is the destination of the
|
||||
cancelled L{IStreamClientEndpoint}.
|
||||
"""
|
||||
|
||||
def __init__(self, address):
|
||||
"""
|
||||
@param address: The L{IAddress} that is the destination of the
|
||||
L{IStreamClientEndpoint} that was cancelled.
|
||||
"""
|
||||
Exception.__init__(self, address)
|
||||
self.address = address
|
||||
|
||||
|
||||
|
||||
class NoProtocol(Exception):
|
||||
"""
|
||||
An C{Exception} that will be raised when the factory given to a
|
||||
L{IStreamClientEndpoint} returns L{None} from C{buildProtocol}.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class UnsupportedAddressFamily(Exception):
|
||||
"""
|
||||
An attempt was made to use a socket with an address family (eg I{AF_INET},
|
||||
I{AF_INET6}, etc) which is not supported by the reactor.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class UnsupportedSocketType(Exception):
|
||||
"""
|
||||
An attempt was made to use a socket of a type (eg I{SOCK_STREAM},
|
||||
I{SOCK_DGRAM}, etc) which is not supported by the reactor.
|
||||
"""
|
||||
|
||||
|
||||
class AlreadyListened(Exception):
|
||||
"""
|
||||
An attempt was made to listen on a file descriptor which can only be
|
||||
listened on once.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class InvalidAddressError(ValueError):
|
||||
"""
|
||||
An invalid address was specified (i.e. neither IPv4 or IPv6, or expected
|
||||
one and got the other).
|
||||
|
||||
@ivar address: See L{__init__}
|
||||
@ivar message: See L{__init__}
|
||||
"""
|
||||
|
||||
def __init__(self, address, message):
|
||||
"""
|
||||
@param address: The address that was provided.
|
||||
@type address: L{bytes}
|
||||
@param message: A native string of additional information provided by
|
||||
the calling context.
|
||||
@type address: L{str}
|
||||
"""
|
||||
self.address = address
|
||||
self.message = message
|
||||
|
||||
|
||||
|
||||
__all__ = [
|
||||
'BindError', 'CannotListenError', 'MulticastJoinError',
|
||||
'MessageLengthError', 'DNSLookupError', 'ConnectInProgressError',
|
||||
'ConnectError', 'ConnectBindError', 'UnknownHostError', 'NoRouteError',
|
||||
'ConnectionRefusedError', 'TCPTimedOutError', 'BadFileError',
|
||||
'ServiceNameUnknownError', 'UserError', 'TimeoutError', 'SSLError',
|
||||
'VerifyError', 'PeerVerifyError', 'CertificateError',
|
||||
'getConnectError', 'ConnectionClosed', 'ConnectionLost',
|
||||
'ConnectionDone', 'ConnectionFdescWentAway', 'AlreadyCalled',
|
||||
'AlreadyCancelled', 'PotentialZombieWarning', 'ProcessDone',
|
||||
'ProcessTerminated', 'ProcessExitedAlready', 'NotConnectingError',
|
||||
'NotListeningError', 'ReactorNotRunning', 'ReactorAlreadyRunning',
|
||||
'ReactorAlreadyInstalledError', 'ConnectingCancelledError',
|
||||
'UnsupportedAddressFamily', 'UnsupportedSocketType', 'InvalidAddressError']
|
||||
@@ -0,0 +1,118 @@
|
||||
# -*- test-case-name: twisted.test.test_fdesc -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
|
||||
"""
|
||||
Utility functions for dealing with POSIX file descriptors.
|
||||
"""
|
||||
|
||||
import os
|
||||
import errno
|
||||
try:
|
||||
import fcntl
|
||||
except ImportError:
|
||||
fcntl = None
|
||||
|
||||
# twisted imports
|
||||
from twisted.internet.main import CONNECTION_LOST, CONNECTION_DONE
|
||||
|
||||
|
||||
def setNonBlocking(fd):
|
||||
"""
|
||||
Set the file description of the given file descriptor to non-blocking.
|
||||
"""
|
||||
flags = fcntl.fcntl(fd, fcntl.F_GETFL)
|
||||
flags = flags | os.O_NONBLOCK
|
||||
fcntl.fcntl(fd, fcntl.F_SETFL, flags)
|
||||
|
||||
|
||||
def setBlocking(fd):
|
||||
"""
|
||||
Set the file description of the given file descriptor to blocking.
|
||||
"""
|
||||
flags = fcntl.fcntl(fd, fcntl.F_GETFL)
|
||||
flags = flags & ~os.O_NONBLOCK
|
||||
fcntl.fcntl(fd, fcntl.F_SETFL, flags)
|
||||
|
||||
|
||||
if fcntl is None:
|
||||
# fcntl isn't available on Windows. By default, handles aren't
|
||||
# inherited on Windows, so we can do nothing here.
|
||||
_setCloseOnExec = _unsetCloseOnExec = lambda fd: None
|
||||
else:
|
||||
def _setCloseOnExec(fd):
|
||||
"""
|
||||
Make a file descriptor close-on-exec.
|
||||
"""
|
||||
flags = fcntl.fcntl(fd, fcntl.F_GETFD)
|
||||
flags = flags | fcntl.FD_CLOEXEC
|
||||
fcntl.fcntl(fd, fcntl.F_SETFD, flags)
|
||||
|
||||
|
||||
def _unsetCloseOnExec(fd):
|
||||
"""
|
||||
Make a file descriptor close-on-exec.
|
||||
"""
|
||||
flags = fcntl.fcntl(fd, fcntl.F_GETFD)
|
||||
flags = flags & ~fcntl.FD_CLOEXEC
|
||||
fcntl.fcntl(fd, fcntl.F_SETFD, flags)
|
||||
|
||||
|
||||
def readFromFD(fd, callback):
|
||||
"""
|
||||
Read from file descriptor, calling callback with resulting data.
|
||||
|
||||
If successful, call 'callback' with a single argument: the
|
||||
resulting data.
|
||||
|
||||
Returns same thing FileDescriptor.doRead would: CONNECTION_LOST,
|
||||
CONNECTION_DONE, or None.
|
||||
|
||||
@type fd: C{int}
|
||||
@param fd: non-blocking file descriptor to be read from.
|
||||
@param callback: a callable which accepts a single argument. If
|
||||
data is read from the file descriptor it will be called with this
|
||||
data. Handling exceptions from calling the callback is up to the
|
||||
caller.
|
||||
|
||||
Note that if the descriptor is still connected but no data is read,
|
||||
None will be returned but callback will not be called.
|
||||
|
||||
@return: CONNECTION_LOST on error, CONNECTION_DONE when fd is
|
||||
closed, otherwise None.
|
||||
"""
|
||||
try:
|
||||
output = os.read(fd, 8192)
|
||||
except (OSError, IOError) as ioe:
|
||||
if ioe.args[0] in (errno.EAGAIN, errno.EINTR):
|
||||
return
|
||||
else:
|
||||
return CONNECTION_LOST
|
||||
if not output:
|
||||
return CONNECTION_DONE
|
||||
callback(output)
|
||||
|
||||
|
||||
def writeToFD(fd, data):
|
||||
"""
|
||||
Write data to file descriptor.
|
||||
|
||||
Returns same thing FileDescriptor.writeSomeData would.
|
||||
|
||||
@type fd: C{int}
|
||||
@param fd: non-blocking file descriptor to be written to.
|
||||
@type data: C{str} or C{buffer}
|
||||
@param data: bytes to write to fd.
|
||||
|
||||
@return: number of bytes written, or CONNECTION_LOST.
|
||||
"""
|
||||
try:
|
||||
return os.write(fd, data)
|
||||
except (OSError, IOError) as io:
|
||||
if io.errno in (errno.EAGAIN, errno.EINTR):
|
||||
return 0
|
||||
return CONNECTION_LOST
|
||||
|
||||
|
||||
__all__ = ["setNonBlocking", "setBlocking", "readFromFD", "writeToFD"]
|
||||
@@ -0,0 +1,188 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
This module provides support for Twisted to interact with the glib
|
||||
mainloop via GObject Introspection.
|
||||
|
||||
In order to use this support, simply do the following::
|
||||
|
||||
from twisted.internet import gireactor
|
||||
gireactor.install()
|
||||
|
||||
If you wish to use a GApplication, register it with the reactor::
|
||||
|
||||
from twisted.internet import reactor
|
||||
reactor.registerGApplication(app)
|
||||
|
||||
Then use twisted.internet APIs as usual.
|
||||
|
||||
On Python 3, pygobject v3.4 or later is required.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
from twisted.python.compat import _PY3
|
||||
from twisted.internet.error import ReactorAlreadyRunning
|
||||
from twisted.internet import _glibbase
|
||||
from twisted.python import runtime
|
||||
|
||||
if _PY3:
|
||||
# We require a sufficiently new version of pygobject, so always exists:
|
||||
_pygtkcompatPresent = True
|
||||
else:
|
||||
# We can't just try to import gi.pygtkcompat, because that would import
|
||||
# gi, and the goal here is to not import gi in cases where that would
|
||||
# cause segfault.
|
||||
from twisted.python.modules import theSystemPath
|
||||
_pygtkcompatPresent = True
|
||||
try:
|
||||
theSystemPath["gi.pygtkcompat"]
|
||||
except KeyError:
|
||||
_pygtkcompatPresent = False
|
||||
|
||||
|
||||
# Modules that we want to ensure aren't imported if we're on older versions of
|
||||
# GI:
|
||||
_PYGTK_MODULES = ['gobject', 'glib', 'gio', 'gtk']
|
||||
|
||||
def _oldGiInit():
|
||||
"""
|
||||
Make sure pygtk and gi aren't loaded at the same time, and import Glib if
|
||||
possible.
|
||||
"""
|
||||
# We can't immediately prevent imports, because that confuses some buggy
|
||||
# code in gi:
|
||||
_glibbase.ensureNotImported(
|
||||
_PYGTK_MODULES,
|
||||
"Introspected and static glib/gtk bindings must not be mixed; can't "
|
||||
"import gireactor since pygtk2 module is already imported.")
|
||||
|
||||
global GLib
|
||||
from gi.repository import GLib
|
||||
if getattr(GLib, "threads_init", None) is not None:
|
||||
GLib.threads_init()
|
||||
|
||||
_glibbase.ensureNotImported([], "",
|
||||
preventImports=_PYGTK_MODULES)
|
||||
|
||||
|
||||
if not _pygtkcompatPresent:
|
||||
# Older versions of gi don't have compatibility layer, so just enforce no
|
||||
# imports of pygtk and gi at same time:
|
||||
_oldGiInit()
|
||||
else:
|
||||
# Newer version of gi, so we can try to initialize compatibility layer; if
|
||||
# real pygtk was already imported we'll get ImportError at this point
|
||||
# rather than segfault, so unconditional import is fine.
|
||||
import gi.pygtkcompat
|
||||
gi.pygtkcompat.enable()
|
||||
# At this point importing gobject will get you gi version, and importing
|
||||
# e.g. gtk will either fail in non-segfaulty way or use gi version if user
|
||||
# does gi.pygtkcompat.enable_gtk(). So, no need to prevent imports of
|
||||
# old school pygtk modules.
|
||||
from gi.repository import GLib
|
||||
if getattr(GLib, "threads_init", None) is not None:
|
||||
GLib.threads_init()
|
||||
|
||||
|
||||
|
||||
class GIReactor(_glibbase.GlibReactorBase):
|
||||
"""
|
||||
GObject-introspection event loop reactor.
|
||||
|
||||
@ivar _gapplication: A C{Gio.Application} instance that was registered
|
||||
with C{registerGApplication}.
|
||||
"""
|
||||
_POLL_DISCONNECTED = (GLib.IOCondition.HUP | GLib.IOCondition.ERR |
|
||||
GLib.IOCondition.NVAL)
|
||||
_POLL_IN = GLib.IOCondition.IN
|
||||
_POLL_OUT = GLib.IOCondition.OUT
|
||||
|
||||
# glib's iochannel sources won't tell us about any events that we haven't
|
||||
# asked for, even if those events aren't sensible inputs to the poll()
|
||||
# call.
|
||||
INFLAGS = _POLL_IN | _POLL_DISCONNECTED
|
||||
OUTFLAGS = _POLL_OUT | _POLL_DISCONNECTED
|
||||
|
||||
# By default no Application is registered:
|
||||
_gapplication = None
|
||||
|
||||
|
||||
def __init__(self, useGtk=False):
|
||||
_gtk = None
|
||||
if useGtk is True:
|
||||
from gi.repository import Gtk as _gtk
|
||||
|
||||
_glibbase.GlibReactorBase.__init__(self, GLib, _gtk, useGtk=useGtk)
|
||||
|
||||
|
||||
def registerGApplication(self, app):
|
||||
"""
|
||||
Register a C{Gio.Application} or C{Gtk.Application}, whose main loop
|
||||
will be used instead of the default one.
|
||||
|
||||
We will C{hold} the application so it doesn't exit on its own. In
|
||||
versions of C{python-gi} 3.2 and later, we exit the event loop using
|
||||
the C{app.quit} method which overrides any holds. Older versions are
|
||||
not supported.
|
||||
"""
|
||||
if self._gapplication is not None:
|
||||
raise RuntimeError(
|
||||
"Can't register more than one application instance.")
|
||||
if self._started:
|
||||
raise ReactorAlreadyRunning(
|
||||
"Can't register application after reactor was started.")
|
||||
if not hasattr(app, "quit"):
|
||||
raise RuntimeError("Application registration is not supported in"
|
||||
" versions of PyGObject prior to 3.2.")
|
||||
self._gapplication = app
|
||||
def run():
|
||||
app.hold()
|
||||
app.run(None)
|
||||
self._run = run
|
||||
|
||||
self._crash = app.quit
|
||||
|
||||
|
||||
|
||||
class PortableGIReactor(_glibbase.PortableGlibReactorBase):
|
||||
"""
|
||||
Portable GObject Introspection event loop reactor.
|
||||
"""
|
||||
def __init__(self, useGtk=False):
|
||||
_gtk = None
|
||||
if useGtk is True:
|
||||
from gi.repository import Gtk as _gtk
|
||||
|
||||
_glibbase.PortableGlibReactorBase.__init__(self, GLib, _gtk,
|
||||
useGtk=useGtk)
|
||||
|
||||
|
||||
def registerGApplication(self, app):
|
||||
"""
|
||||
Register a C{Gio.Application} or C{Gtk.Application}, whose main loop
|
||||
will be used instead of the default one.
|
||||
"""
|
||||
raise NotImplementedError("GApplication is not currently supported on Windows.")
|
||||
|
||||
|
||||
|
||||
def install(useGtk=False):
|
||||
"""
|
||||
Configure the twisted mainloop to be run inside the glib mainloop.
|
||||
|
||||
@param useGtk: should GTK+ rather than glib event loop be
|
||||
used (this will be slightly slower but does support GUI).
|
||||
"""
|
||||
if runtime.platform.getType() == 'posix':
|
||||
reactor = GIReactor(useGtk=useGtk)
|
||||
else:
|
||||
reactor = PortableGIReactor(useGtk=useGtk)
|
||||
|
||||
from twisted.internet.main import installReactor
|
||||
installReactor(reactor)
|
||||
return reactor
|
||||
|
||||
|
||||
__all__ = ['install']
|
||||
@@ -0,0 +1,10 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
I/O Completion Ports reactor
|
||||
"""
|
||||
|
||||
from twisted.internet.iocpreactor.reactor import install
|
||||
|
||||
__all__ = ['install']
|
||||
@@ -0,0 +1,399 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Abstract file handle class
|
||||
"""
|
||||
|
||||
from twisted.internet import main, error, interfaces
|
||||
from twisted.internet.abstract import _ConsumerMixin, _LogOwner
|
||||
from twisted.python import failure
|
||||
from twisted.python.compat import unicode
|
||||
|
||||
from zope.interface import implementer
|
||||
import errno
|
||||
|
||||
from twisted.internet.iocpreactor.const import ERROR_HANDLE_EOF
|
||||
from twisted.internet.iocpreactor.const import ERROR_IO_PENDING
|
||||
from twisted.internet.iocpreactor import iocpsupport as _iocp
|
||||
|
||||
|
||||
@implementer(interfaces.IPushProducer, interfaces.IConsumer,
|
||||
interfaces.ITransport, interfaces.IHalfCloseableDescriptor)
|
||||
class FileHandle(_ConsumerMixin, _LogOwner):
|
||||
"""
|
||||
File handle that can read and write asynchronously
|
||||
"""
|
||||
# read stuff
|
||||
maxReadBuffers = 16
|
||||
readBufferSize = 4096
|
||||
reading = False
|
||||
dynamicReadBuffers = True # set this to false if subclass doesn't do iovecs
|
||||
_readNextBuffer = 0
|
||||
_readSize = 0 # how much data we have in the read buffer
|
||||
_readScheduled = None
|
||||
_readScheduledInOS = False
|
||||
|
||||
|
||||
def startReading(self):
|
||||
self.reactor.addActiveHandle(self)
|
||||
if not self._readScheduled and not self.reading:
|
||||
self.reading = True
|
||||
self._readScheduled = self.reactor.callLater(0,
|
||||
self._resumeReading)
|
||||
|
||||
|
||||
def stopReading(self):
|
||||
if self._readScheduled:
|
||||
self._readScheduled.cancel()
|
||||
self._readScheduled = None
|
||||
self.reading = False
|
||||
|
||||
|
||||
def _resumeReading(self):
|
||||
self._readScheduled = None
|
||||
if self._dispatchData() and not self._readScheduledInOS:
|
||||
self.doRead()
|
||||
|
||||
|
||||
def _dispatchData(self):
|
||||
"""
|
||||
Dispatch previously read data. Return True if self.reading and we don't
|
||||
have any more data
|
||||
"""
|
||||
if not self._readSize:
|
||||
return self.reading
|
||||
size = self._readSize
|
||||
full_buffers = size // self.readBufferSize
|
||||
while self._readNextBuffer < full_buffers:
|
||||
self.dataReceived(self._readBuffers[self._readNextBuffer])
|
||||
self._readNextBuffer += 1
|
||||
if not self.reading:
|
||||
return False
|
||||
remainder = size % self.readBufferSize
|
||||
if remainder:
|
||||
self.dataReceived(self._readBuffers[full_buffers][0:remainder])
|
||||
if self.dynamicReadBuffers:
|
||||
total_buffer_size = self.readBufferSize * len(self._readBuffers)
|
||||
# we have one buffer too many
|
||||
if size < total_buffer_size - self.readBufferSize:
|
||||
del self._readBuffers[-1]
|
||||
# we filled all buffers, so allocate one more
|
||||
elif (size == total_buffer_size and
|
||||
len(self._readBuffers) < self.maxReadBuffers):
|
||||
self._readBuffers.append(bytearray(self.readBufferSize))
|
||||
self._readNextBuffer = 0
|
||||
self._readSize = 0
|
||||
return self.reading
|
||||
|
||||
|
||||
def _cbRead(self, rc, data, evt):
|
||||
self._readScheduledInOS = False
|
||||
if self._handleRead(rc, data, evt):
|
||||
self.doRead()
|
||||
|
||||
|
||||
def _handleRead(self, rc, data, evt):
|
||||
"""
|
||||
Returns False if we should stop reading for now
|
||||
"""
|
||||
if self.disconnected:
|
||||
return False
|
||||
# graceful disconnection
|
||||
if (not (rc or data)) or rc in (errno.WSAEDISCON, ERROR_HANDLE_EOF):
|
||||
self.reactor.removeActiveHandle(self)
|
||||
self.readConnectionLost(failure.Failure(main.CONNECTION_DONE))
|
||||
return False
|
||||
# XXX: not handling WSAEWOULDBLOCK
|
||||
# ("too many outstanding overlapped I/O requests")
|
||||
elif rc:
|
||||
self.connectionLost(failure.Failure(
|
||||
error.ConnectionLost("read error -- %s (%s)" %
|
||||
(errno.errorcode.get(rc, 'unknown'), rc))))
|
||||
return False
|
||||
else:
|
||||
assert self._readSize == 0
|
||||
assert self._readNextBuffer == 0
|
||||
self._readSize = data
|
||||
return self._dispatchData()
|
||||
|
||||
|
||||
def doRead(self):
|
||||
evt = _iocp.Event(self._cbRead, self)
|
||||
|
||||
evt.buff = buff = self._readBuffers
|
||||
rc, numBytesRead = self.readFromHandle(buff, evt)
|
||||
|
||||
if not rc or rc == ERROR_IO_PENDING:
|
||||
self._readScheduledInOS = True
|
||||
else:
|
||||
self._handleRead(rc, numBytesRead, evt)
|
||||
|
||||
|
||||
def readFromHandle(self, bufflist, evt):
|
||||
raise NotImplementedError() # TODO: this should default to ReadFile
|
||||
|
||||
|
||||
def dataReceived(self, data):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
def readConnectionLost(self, reason):
|
||||
self.connectionLost(reason)
|
||||
|
||||
|
||||
# write stuff
|
||||
dataBuffer = b''
|
||||
offset = 0
|
||||
writing = False
|
||||
_writeScheduled = None
|
||||
_writeDisconnecting = False
|
||||
_writeDisconnected = False
|
||||
writeBufferSize = 2**2**2**2
|
||||
|
||||
|
||||
def loseWriteConnection(self):
|
||||
self._writeDisconnecting = True
|
||||
self.startWriting()
|
||||
|
||||
|
||||
def _closeWriteConnection(self):
|
||||
# override in subclasses
|
||||
pass
|
||||
|
||||
|
||||
def writeConnectionLost(self, reason):
|
||||
# in current code should never be called
|
||||
self.connectionLost(reason)
|
||||
|
||||
|
||||
def startWriting(self):
|
||||
self.reactor.addActiveHandle(self)
|
||||
|
||||
if not self._writeScheduled and not self.writing:
|
||||
self.writing = True
|
||||
self._writeScheduled = self.reactor.callLater(0,
|
||||
self._resumeWriting)
|
||||
|
||||
|
||||
def stopWriting(self):
|
||||
if self._writeScheduled:
|
||||
self._writeScheduled.cancel()
|
||||
self._writeScheduled = None
|
||||
self.writing = False
|
||||
|
||||
|
||||
def _resumeWriting(self):
|
||||
self._writeScheduled = None
|
||||
self.doWrite()
|
||||
|
||||
|
||||
def _cbWrite(self, rc, numBytesWritten, evt):
|
||||
if self._handleWrite(rc, numBytesWritten, evt):
|
||||
self.doWrite()
|
||||
|
||||
|
||||
def _handleWrite(self, rc, numBytesWritten, evt):
|
||||
"""
|
||||
Returns false if we should stop writing for now
|
||||
"""
|
||||
if self.disconnected or self._writeDisconnected:
|
||||
return False
|
||||
# XXX: not handling WSAEWOULDBLOCK
|
||||
# ("too many outstanding overlapped I/O requests")
|
||||
if rc:
|
||||
self.connectionLost(failure.Failure(
|
||||
error.ConnectionLost("write error -- %s (%s)" %
|
||||
(errno.errorcode.get(rc, 'unknown'), rc))))
|
||||
return False
|
||||
else:
|
||||
self.offset += numBytesWritten
|
||||
# If there is nothing left to send,
|
||||
if self.offset == len(self.dataBuffer) and not self._tempDataLen:
|
||||
self.dataBuffer = b""
|
||||
self.offset = 0
|
||||
# stop writing
|
||||
self.stopWriting()
|
||||
# If I've got a producer who is supposed to supply me with data
|
||||
if self.producer is not None and ((not self.streamingProducer)
|
||||
or self.producerPaused):
|
||||
# tell them to supply some more.
|
||||
self.producerPaused = True
|
||||
self.producer.resumeProducing()
|
||||
elif self.disconnecting:
|
||||
# But if I was previously asked to let the connection die,
|
||||
# do so.
|
||||
self.connectionLost(failure.Failure(main.CONNECTION_DONE))
|
||||
elif self._writeDisconnecting:
|
||||
# I was previously asked to half-close the connection.
|
||||
self._writeDisconnected = True
|
||||
self._closeWriteConnection()
|
||||
return False
|
||||
else:
|
||||
return True
|
||||
|
||||
|
||||
def doWrite(self):
|
||||
if len(self.dataBuffer) - self.offset < self.SEND_LIMIT:
|
||||
# If there is currently less than SEND_LIMIT bytes left to send
|
||||
# in the string, extend it with the array data.
|
||||
self.dataBuffer = (self.dataBuffer[self.offset:] +
|
||||
b"".join(self._tempDataBuffer))
|
||||
self.offset = 0
|
||||
self._tempDataBuffer = []
|
||||
self._tempDataLen = 0
|
||||
|
||||
evt = _iocp.Event(self._cbWrite, self)
|
||||
|
||||
# Send as much data as you can.
|
||||
if self.offset:
|
||||
sendView = memoryview(self.dataBuffer)
|
||||
evt.buff = buff = sendView[self.offset:]
|
||||
else:
|
||||
evt.buff = buff = self.dataBuffer
|
||||
rc, data = self.writeToHandle(buff, evt)
|
||||
if rc and rc != ERROR_IO_PENDING:
|
||||
self._handleWrite(rc, data, evt)
|
||||
|
||||
|
||||
def writeToHandle(self, buff, evt):
|
||||
raise NotImplementedError() # TODO: this should default to WriteFile
|
||||
|
||||
|
||||
def write(self, data):
|
||||
"""Reliably write some data.
|
||||
|
||||
The data is buffered until his file descriptor is ready for writing.
|
||||
"""
|
||||
if isinstance(data, unicode): # no, really, I mean it
|
||||
raise TypeError("Data must not be unicode")
|
||||
if not self.connected or self._writeDisconnected:
|
||||
return
|
||||
if data:
|
||||
self._tempDataBuffer.append(data)
|
||||
self._tempDataLen += len(data)
|
||||
if self.producer is not None and self.streamingProducer:
|
||||
if (len(self.dataBuffer) + self._tempDataLen
|
||||
> self.writeBufferSize):
|
||||
self.producerPaused = True
|
||||
self.producer.pauseProducing()
|
||||
self.startWriting()
|
||||
|
||||
|
||||
def writeSequence(self, iovec):
|
||||
for i in iovec:
|
||||
if isinstance(i, unicode): # no, really, I mean it
|
||||
raise TypeError("Data must not be unicode")
|
||||
if not self.connected or not iovec or self._writeDisconnected:
|
||||
return
|
||||
self._tempDataBuffer.extend(iovec)
|
||||
for i in iovec:
|
||||
self._tempDataLen += len(i)
|
||||
if self.producer is not None and self.streamingProducer:
|
||||
if len(self.dataBuffer) + self._tempDataLen > self.writeBufferSize:
|
||||
self.producerPaused = True
|
||||
self.producer.pauseProducing()
|
||||
self.startWriting()
|
||||
|
||||
|
||||
# general stuff
|
||||
connected = False
|
||||
disconnected = False
|
||||
disconnecting = False
|
||||
logstr = "Uninitialized"
|
||||
|
||||
SEND_LIMIT = 128*1024
|
||||
|
||||
|
||||
def __init__(self, reactor = None):
|
||||
if not reactor:
|
||||
from twisted.internet import reactor
|
||||
self.reactor = reactor
|
||||
self._tempDataBuffer = [] # will be added to dataBuffer in doWrite
|
||||
self._tempDataLen = 0
|
||||
self._readBuffers = [bytearray(self.readBufferSize)]
|
||||
|
||||
|
||||
def connectionLost(self, reason):
|
||||
"""
|
||||
The connection was lost.
|
||||
|
||||
This is called when the connection on a selectable object has been
|
||||
lost. It will be called whether the connection was closed explicitly,
|
||||
an exception occurred in an event handler, or the other end of the
|
||||
connection closed it first.
|
||||
|
||||
Clean up state here, but make sure to call back up to FileDescriptor.
|
||||
"""
|
||||
|
||||
self.disconnected = True
|
||||
self.connected = False
|
||||
if self.producer is not None:
|
||||
self.producer.stopProducing()
|
||||
self.producer = None
|
||||
self.stopReading()
|
||||
self.stopWriting()
|
||||
self.reactor.removeActiveHandle(self)
|
||||
|
||||
|
||||
def getFileHandle(self):
|
||||
return -1
|
||||
|
||||
|
||||
def loseConnection(self, _connDone=failure.Failure(main.CONNECTION_DONE)):
|
||||
"""
|
||||
Close the connection at the next available opportunity.
|
||||
|
||||
Call this to cause this FileDescriptor to lose its connection. It will
|
||||
first write any data that it has buffered.
|
||||
|
||||
If there is data buffered yet to be written, this method will cause the
|
||||
transport to lose its connection as soon as it's done flushing its
|
||||
write buffer. If you have a producer registered, the connection won't
|
||||
be closed until the producer is finished. Therefore, make sure you
|
||||
unregister your producer when it's finished, or the connection will
|
||||
never close.
|
||||
"""
|
||||
|
||||
if self.connected and not self.disconnecting:
|
||||
if self._writeDisconnected:
|
||||
# doWrite won't trigger the connection close anymore
|
||||
self.stopReading()
|
||||
self.stopWriting
|
||||
self.connectionLost(_connDone)
|
||||
else:
|
||||
self.stopReading()
|
||||
self.startWriting()
|
||||
self.disconnecting = 1
|
||||
|
||||
|
||||
# Producer/consumer implementation
|
||||
|
||||
def stopConsuming(self):
|
||||
"""
|
||||
Stop consuming data.
|
||||
|
||||
This is called when a producer has lost its connection, to tell the
|
||||
consumer to go lose its connection (and break potential circular
|
||||
references).
|
||||
"""
|
||||
self.unregisterProducer()
|
||||
self.loseConnection()
|
||||
|
||||
|
||||
# producer interface implementation
|
||||
|
||||
def resumeProducing(self):
|
||||
if self.connected and not self.disconnecting:
|
||||
self.startReading()
|
||||
|
||||
|
||||
def pauseProducing(self):
|
||||
self.stopReading()
|
||||
|
||||
|
||||
def stopProducing(self):
|
||||
self.loseConnection()
|
||||
|
||||
|
||||
__all__ = ['FileHandle']
|
||||
@@ -0,0 +1,23 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
|
||||
"""
|
||||
Distutils file for building low-level IOCP bindings from their Pyrex source
|
||||
"""
|
||||
|
||||
|
||||
from distutils.core import setup
|
||||
from distutils.extension import Extension
|
||||
from Cython.Distutils import build_ext
|
||||
|
||||
setup(name='iocpsupport',
|
||||
ext_modules=[Extension('iocpsupport',
|
||||
['iocpsupport/iocpsupport.pyx',
|
||||
'iocpsupport/winsock_pointers.c'],
|
||||
libraries = ['ws2_32'],
|
||||
)
|
||||
],
|
||||
cmdclass = {'build_ext': build_ext},
|
||||
)
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
# -*- test-case-name: twisted.internet.test.test_main -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
|
||||
"""
|
||||
Backwards compatibility, and utility functions.
|
||||
|
||||
In general, this module should not be used, other than by reactor authors
|
||||
who need to use the 'installReactor' method.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
from twisted.internet import error
|
||||
|
||||
CONNECTION_DONE = error.ConnectionDone('Connection done')
|
||||
CONNECTION_LOST = error.ConnectionLost('Connection lost')
|
||||
|
||||
|
||||
|
||||
def installReactor(reactor):
|
||||
"""
|
||||
Install reactor C{reactor}.
|
||||
|
||||
@param reactor: An object that provides one or more IReactor* interfaces.
|
||||
"""
|
||||
# this stuff should be common to all reactors.
|
||||
import twisted.internet
|
||||
import sys
|
||||
if 'twisted.internet.reactor' in sys.modules:
|
||||
raise error.ReactorAlreadyInstalledError("reactor already installed")
|
||||
twisted.internet.reactor = reactor
|
||||
sys.modules['twisted.internet.reactor'] = reactor
|
||||
|
||||
|
||||
__all__ = ["CONNECTION_LOST", "CONNECTION_DONE", "installReactor"]
|
||||
@@ -0,0 +1,37 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
|
||||
"""
|
||||
This module integrates PyUI with twisted.internet's mainloop.
|
||||
|
||||
Maintainer: Jp Calderone
|
||||
|
||||
See doc/examples/pyuidemo.py for example usage.
|
||||
"""
|
||||
|
||||
# System imports
|
||||
import pyui
|
||||
|
||||
def _guiUpdate(reactor, delay):
|
||||
pyui.draw()
|
||||
if pyui.update() == 0:
|
||||
pyui.quit()
|
||||
reactor.stop()
|
||||
else:
|
||||
reactor.callLater(delay, _guiUpdate, reactor, delay)
|
||||
|
||||
|
||||
def install(ms=10, reactor=None, args=(), kw={}):
|
||||
"""
|
||||
Schedule PyUI's display to be updated approximately every C{ms}
|
||||
milliseconds, and initialize PyUI with the specified arguments.
|
||||
"""
|
||||
d = pyui.init(*args, **kw)
|
||||
|
||||
if reactor is None:
|
||||
from twisted.internet import reactor
|
||||
_guiUpdate(reactor, ms / 1000.0)
|
||||
return d
|
||||
|
||||
__all__ = ["install"]
|
||||
@@ -0,0 +1,39 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
The reactor is the Twisted event loop within Twisted, the loop which drives
|
||||
applications using Twisted. The reactor provides APIs for networking,
|
||||
threading, dispatching events, and more.
|
||||
|
||||
The default reactor depends on the platform and will be installed if this
|
||||
module is imported without another reactor being explicitly installed
|
||||
beforehand. Regardless of which reactor is installed, importing this module is
|
||||
the correct way to get a reference to it.
|
||||
|
||||
New application code should prefer to pass and accept the reactor as a
|
||||
parameter where it is needed, rather than relying on being able to import this
|
||||
module to get a reference. This simplifies unit testing and may make it easier
|
||||
to one day support multiple reactors (as a performance enhancement), though
|
||||
this is not currently possible.
|
||||
|
||||
@see: L{IReactorCore<twisted.internet.interfaces.IReactorCore>}
|
||||
@see: L{IReactorTime<twisted.internet.interfaces.IReactorTime>}
|
||||
@see: L{IReactorProcess<twisted.internet.interfaces.IReactorProcess>}
|
||||
@see: L{IReactorTCP<twisted.internet.interfaces.IReactorTCP>}
|
||||
@see: L{IReactorSSL<twisted.internet.interfaces.IReactorSSL>}
|
||||
@see: L{IReactorUDP<twisted.internet.interfaces.IReactorUDP>}
|
||||
@see: L{IReactorMulticast<twisted.internet.interfaces.IReactorMulticast>}
|
||||
@see: L{IReactorUNIX<twisted.internet.interfaces.IReactorUNIX>}
|
||||
@see: L{IReactorUNIXDatagram<twisted.internet.interfaces.IReactorUNIXDatagram>}
|
||||
@see: L{IReactorFDSet<twisted.internet.interfaces.IReactorFDSet>}
|
||||
@see: L{IReactorThreads<twisted.internet.interfaces.IReactorThreads>}
|
||||
@see: L{IReactorPluggableResolver<twisted.internet.interfaces.IReactorPluggableResolver>}
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
import sys
|
||||
del sys.modules['twisted.internet.reactor']
|
||||
from twisted.internet import default
|
||||
default.install()
|
||||
@@ -0,0 +1,948 @@
|
||||
# -*- test-case-name: twisted.test.test_task,twisted.test.test_cooperator -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Scheduling utility methods and classes.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
__metaclass__ = type
|
||||
|
||||
import sys
|
||||
import time
|
||||
import warnings
|
||||
|
||||
from zope.interface import implementer
|
||||
|
||||
from twisted.python import log
|
||||
from twisted.python import reflect
|
||||
from twisted.python.deprecate import _getDeprecationWarningString
|
||||
from twisted.python.failure import Failure
|
||||
from incremental import Version
|
||||
|
||||
from twisted.internet import base, defer
|
||||
from twisted.internet.interfaces import IReactorTime
|
||||
from twisted.internet.error import ReactorNotRunning
|
||||
|
||||
|
||||
class LoopingCall:
|
||||
"""Call a function repeatedly.
|
||||
|
||||
If C{f} returns a deferred, rescheduling will not take place until the
|
||||
deferred has fired. The result value is ignored.
|
||||
|
||||
@ivar f: The function to call.
|
||||
@ivar a: A tuple of arguments to pass the function.
|
||||
@ivar kw: A dictionary of keyword arguments to pass to the function.
|
||||
@ivar clock: A provider of
|
||||
L{twisted.internet.interfaces.IReactorTime}. The default is
|
||||
L{twisted.internet.reactor}. Feel free to set this to
|
||||
something else, but it probably ought to be set *before*
|
||||
calling L{start}.
|
||||
|
||||
@type running: C{bool}
|
||||
@ivar running: A flag which is C{True} while C{f} is scheduled to be called
|
||||
(or is currently being called). It is set to C{True} when L{start} is
|
||||
called and set to C{False} when L{stop} is called or if C{f} raises an
|
||||
exception. In either case, it will be C{False} by the time the
|
||||
C{Deferred} returned by L{start} fires its callback or errback.
|
||||
|
||||
@type _realLastTime: C{float}
|
||||
@ivar _realLastTime: When counting skips, the time at which the skip
|
||||
counter was last invoked.
|
||||
|
||||
@type _runAtStart: C{bool}
|
||||
@ivar _runAtStart: A flag indicating whether the 'now' argument was passed
|
||||
to L{LoopingCall.start}.
|
||||
"""
|
||||
|
||||
call = None
|
||||
running = False
|
||||
_deferred = None
|
||||
interval = None
|
||||
_runAtStart = False
|
||||
starttime = None
|
||||
|
||||
def __init__(self, f, *a, **kw):
|
||||
self.f = f
|
||||
self.a = a
|
||||
self.kw = kw
|
||||
from twisted.internet import reactor
|
||||
self.clock = reactor
|
||||
|
||||
@property
|
||||
def deferred(self):
|
||||
"""
|
||||
DEPRECATED. L{Deferred} fired when loop stops or fails.
|
||||
|
||||
Use the L{Deferred} returned by L{LoopingCall.start}.
|
||||
"""
|
||||
warningString = _getDeprecationWarningString(
|
||||
"twisted.internet.task.LoopingCall.deferred",
|
||||
Version("Twisted", 16, 0, 0),
|
||||
replacement='the deferred returned by start()')
|
||||
warnings.warn(warningString, DeprecationWarning, stacklevel=2)
|
||||
|
||||
return self._deferred
|
||||
|
||||
def withCount(cls, countCallable):
|
||||
"""
|
||||
An alternate constructor for L{LoopingCall} that makes available the
|
||||
number of calls which should have occurred since it was last invoked.
|
||||
|
||||
Note that this number is an C{int} value; It represents the discrete
|
||||
number of calls that should have been made. For example, if you are
|
||||
using a looping call to display an animation with discrete frames, this
|
||||
number would be the number of frames to advance.
|
||||
|
||||
The count is normally 1, but can be higher. For example, if the reactor
|
||||
is blocked and takes too long to invoke the L{LoopingCall}, a Deferred
|
||||
returned from a previous call is not fired before an interval has
|
||||
elapsed, or if the callable itself blocks for longer than an interval,
|
||||
preventing I{itself} from being called.
|
||||
|
||||
When running with an interval if 0, count will be always 1.
|
||||
|
||||
@param countCallable: A callable that will be invoked each time the
|
||||
resulting LoopingCall is run, with an integer specifying the number
|
||||
of calls that should have been invoked.
|
||||
|
||||
@type countCallable: 1-argument callable which takes an C{int}
|
||||
|
||||
@return: An instance of L{LoopingCall} with call counting enabled,
|
||||
which provides the count as the first positional argument.
|
||||
|
||||
@rtype: L{LoopingCall}
|
||||
|
||||
@since: 9.0
|
||||
"""
|
||||
|
||||
def counter():
|
||||
now = self.clock.seconds()
|
||||
|
||||
if self.interval == 0:
|
||||
self._realLastTime = now
|
||||
return countCallable(1)
|
||||
|
||||
lastTime = self._realLastTime
|
||||
if lastTime is None:
|
||||
lastTime = self.starttime
|
||||
if self._runAtStart:
|
||||
lastTime -= self.interval
|
||||
lastInterval = self._intervalOf(lastTime)
|
||||
thisInterval = self._intervalOf(now)
|
||||
count = thisInterval - lastInterval
|
||||
if count > 0:
|
||||
self._realLastTime = now
|
||||
return countCallable(count)
|
||||
|
||||
self = cls(counter)
|
||||
|
||||
self._realLastTime = None
|
||||
|
||||
return self
|
||||
|
||||
withCount = classmethod(withCount)
|
||||
|
||||
|
||||
def _intervalOf(self, t):
|
||||
"""
|
||||
Determine the number of intervals passed as of the given point in
|
||||
time.
|
||||
|
||||
@param t: The specified time (from the start of the L{LoopingCall}) to
|
||||
be measured in intervals
|
||||
|
||||
@return: The C{int} number of intervals which have passed as of the
|
||||
given point in time.
|
||||
"""
|
||||
elapsedTime = t - self.starttime
|
||||
intervalNum = int(elapsedTime / self.interval)
|
||||
return intervalNum
|
||||
|
||||
|
||||
def start(self, interval, now=True):
|
||||
"""
|
||||
Start running function every interval seconds.
|
||||
|
||||
@param interval: The number of seconds between calls. May be
|
||||
less than one. Precision will depend on the underlying
|
||||
platform, the available hardware, and the load on the system.
|
||||
|
||||
@param now: If True, run this call right now. Otherwise, wait
|
||||
until the interval has elapsed before beginning.
|
||||
|
||||
@return: A Deferred whose callback will be invoked with
|
||||
C{self} when C{self.stop} is called, or whose errback will be
|
||||
invoked when the function raises an exception or returned a
|
||||
deferred that has its errback invoked.
|
||||
"""
|
||||
assert not self.running, ("Tried to start an already running "
|
||||
"LoopingCall.")
|
||||
if interval < 0:
|
||||
raise ValueError("interval must be >= 0")
|
||||
self.running = True
|
||||
# Loop might fail to start and then self._deferred will be cleared.
|
||||
# This why the local C{deferred} variable is used.
|
||||
deferred = self._deferred = defer.Deferred()
|
||||
self.starttime = self.clock.seconds()
|
||||
self.interval = interval
|
||||
self._runAtStart = now
|
||||
if now:
|
||||
self()
|
||||
else:
|
||||
self._scheduleFrom(self.starttime)
|
||||
return deferred
|
||||
|
||||
def stop(self):
|
||||
"""Stop running function.
|
||||
"""
|
||||
assert self.running, ("Tried to stop a LoopingCall that was "
|
||||
"not running.")
|
||||
self.running = False
|
||||
if self.call is not None:
|
||||
self.call.cancel()
|
||||
self.call = None
|
||||
d, self._deferred = self._deferred, None
|
||||
d.callback(self)
|
||||
|
||||
def reset(self):
|
||||
"""
|
||||
Skip the next iteration and reset the timer.
|
||||
|
||||
@since: 11.1
|
||||
"""
|
||||
assert self.running, ("Tried to reset a LoopingCall that was "
|
||||
"not running.")
|
||||
if self.call is not None:
|
||||
self.call.cancel()
|
||||
self.call = None
|
||||
self.starttime = self.clock.seconds()
|
||||
self._scheduleFrom(self.starttime)
|
||||
|
||||
def __call__(self):
|
||||
def cb(result):
|
||||
if self.running:
|
||||
self._scheduleFrom(self.clock.seconds())
|
||||
else:
|
||||
d, self._deferred = self._deferred, None
|
||||
d.callback(self)
|
||||
|
||||
def eb(failure):
|
||||
self.running = False
|
||||
d, self._deferred = self._deferred, None
|
||||
d.errback(failure)
|
||||
|
||||
self.call = None
|
||||
d = defer.maybeDeferred(self.f, *self.a, **self.kw)
|
||||
d.addCallback(cb)
|
||||
d.addErrback(eb)
|
||||
|
||||
|
||||
def _scheduleFrom(self, when):
|
||||
"""
|
||||
Schedule the next iteration of this looping call.
|
||||
|
||||
@param when: The present time from whence the call is scheduled.
|
||||
"""
|
||||
def howLong():
|
||||
# How long should it take until the next invocation of our
|
||||
# callable? Split out into a function because there are multiple
|
||||
# places we want to 'return' out of this.
|
||||
if self.interval == 0:
|
||||
# If the interval is 0, just go as fast as possible, always
|
||||
# return zero, call ourselves ASAP.
|
||||
return 0
|
||||
# Compute the time until the next interval; how long has this call
|
||||
# been running for?
|
||||
runningFor = when - self.starttime
|
||||
# And based on that start time, when does the current interval end?
|
||||
untilNextInterval = self.interval - (runningFor % self.interval)
|
||||
# Now that we know how long it would be, we have to tell if the
|
||||
# number is effectively zero. However, we can't just test against
|
||||
# zero. If a number with a small exponent is added to a number
|
||||
# with a large exponent, it may be so small that the digits just
|
||||
# fall off the end, which means that adding the increment makes no
|
||||
# difference; it's time to tick over into the next interval.
|
||||
if when == when + untilNextInterval:
|
||||
# If it's effectively zero, then we need to add another
|
||||
# interval.
|
||||
return self.interval
|
||||
# Finally, if everything else is normal, we just return the
|
||||
# computed delay.
|
||||
return untilNextInterval
|
||||
self.call = self.clock.callLater(howLong(), self)
|
||||
|
||||
|
||||
def __repr__(self):
|
||||
if hasattr(self.f, '__qualname__'):
|
||||
func = self.f.__qualname__
|
||||
elif hasattr(self.f, '__name__'):
|
||||
func = self.f.__name__
|
||||
if hasattr(self.f, 'im_class'):
|
||||
func = self.f.im_class.__name__ + '.' + func
|
||||
else:
|
||||
func = reflect.safe_repr(self.f)
|
||||
|
||||
return 'LoopingCall<%r>(%s, *%s, **%s)' % (
|
||||
self.interval, func, reflect.safe_repr(self.a),
|
||||
reflect.safe_repr(self.kw))
|
||||
|
||||
|
||||
|
||||
class SchedulerError(Exception):
|
||||
"""
|
||||
The operation could not be completed because the scheduler or one of its
|
||||
tasks was in an invalid state. This exception should not be raised
|
||||
directly, but is a superclass of various scheduler-state-related
|
||||
exceptions.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class SchedulerStopped(SchedulerError):
|
||||
"""
|
||||
The operation could not complete because the scheduler was stopped in
|
||||
progress or was already stopped.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class TaskFinished(SchedulerError):
|
||||
"""
|
||||
The operation could not complete because the task was already completed,
|
||||
stopped, encountered an error or otherwise permanently stopped running.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class TaskDone(TaskFinished):
|
||||
"""
|
||||
The operation could not complete because the task was already completed.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class TaskStopped(TaskFinished):
|
||||
"""
|
||||
The operation could not complete because the task was stopped.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class TaskFailed(TaskFinished):
|
||||
"""
|
||||
The operation could not complete because the task died with an unhandled
|
||||
error.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class NotPaused(SchedulerError):
|
||||
"""
|
||||
This exception is raised when a task is resumed which was not previously
|
||||
paused.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class _Timer(object):
|
||||
MAX_SLICE = 0.01
|
||||
def __init__(self):
|
||||
self.end = time.time() + self.MAX_SLICE
|
||||
|
||||
|
||||
def __call__(self):
|
||||
return time.time() >= self.end
|
||||
|
||||
|
||||
|
||||
_EPSILON = 0.00000001
|
||||
def _defaultScheduler(x):
|
||||
from twisted.internet import reactor
|
||||
return reactor.callLater(_EPSILON, x)
|
||||
|
||||
|
||||
class CooperativeTask(object):
|
||||
"""
|
||||
A L{CooperativeTask} is a task object inside a L{Cooperator}, which can be
|
||||
paused, resumed, and stopped. It can also have its completion (or
|
||||
termination) monitored.
|
||||
|
||||
@see: L{Cooperator.cooperate}
|
||||
|
||||
@ivar _iterator: the iterator to iterate when this L{CooperativeTask} is
|
||||
asked to do work.
|
||||
|
||||
@ivar _cooperator: the L{Cooperator} that this L{CooperativeTask}
|
||||
participates in, which is used to re-insert it upon resume.
|
||||
|
||||
@ivar _deferreds: the list of L{defer.Deferred}s to fire when this task
|
||||
completes, fails, or finishes.
|
||||
|
||||
@type _deferreds: C{list}
|
||||
|
||||
@type _cooperator: L{Cooperator}
|
||||
|
||||
@ivar _pauseCount: the number of times that this L{CooperativeTask} has
|
||||
been paused; if 0, it is running.
|
||||
|
||||
@type _pauseCount: C{int}
|
||||
|
||||
@ivar _completionState: The completion-state of this L{CooperativeTask}.
|
||||
L{None} if the task is not yet completed, an instance of L{TaskStopped}
|
||||
if C{stop} was called to stop this task early, of L{TaskFailed} if the
|
||||
application code in the iterator raised an exception which caused it to
|
||||
terminate, and of L{TaskDone} if it terminated normally via raising
|
||||
C{StopIteration}.
|
||||
|
||||
@type _completionState: L{TaskFinished}
|
||||
"""
|
||||
|
||||
def __init__(self, iterator, cooperator):
|
||||
"""
|
||||
A private constructor: to create a new L{CooperativeTask}, see
|
||||
L{Cooperator.cooperate}.
|
||||
"""
|
||||
self._iterator = iterator
|
||||
self._cooperator = cooperator
|
||||
self._deferreds = []
|
||||
self._pauseCount = 0
|
||||
self._completionState = None
|
||||
self._completionResult = None
|
||||
cooperator._addTask(self)
|
||||
|
||||
|
||||
def whenDone(self):
|
||||
"""
|
||||
Get a L{defer.Deferred} notification of when this task is complete.
|
||||
|
||||
@return: a L{defer.Deferred} that fires with the C{iterator} that this
|
||||
L{CooperativeTask} was created with when the iterator has been
|
||||
exhausted (i.e. its C{next} method has raised C{StopIteration}), or
|
||||
fails with the exception raised by C{next} if it raises some other
|
||||
exception.
|
||||
|
||||
@rtype: L{defer.Deferred}
|
||||
"""
|
||||
d = defer.Deferred()
|
||||
if self._completionState is None:
|
||||
self._deferreds.append(d)
|
||||
else:
|
||||
d.callback(self._completionResult)
|
||||
return d
|
||||
|
||||
|
||||
def pause(self):
|
||||
"""
|
||||
Pause this L{CooperativeTask}. Stop doing work until
|
||||
L{CooperativeTask.resume} is called. If C{pause} is called more than
|
||||
once, C{resume} must be called an equal number of times to resume this
|
||||
task.
|
||||
|
||||
@raise TaskFinished: if this task has already finished or completed.
|
||||
"""
|
||||
self._checkFinish()
|
||||
self._pauseCount += 1
|
||||
if self._pauseCount == 1:
|
||||
self._cooperator._removeTask(self)
|
||||
|
||||
|
||||
def resume(self):
|
||||
"""
|
||||
Resume processing of a paused L{CooperativeTask}.
|
||||
|
||||
@raise NotPaused: if this L{CooperativeTask} is not paused.
|
||||
"""
|
||||
if self._pauseCount == 0:
|
||||
raise NotPaused()
|
||||
self._pauseCount -= 1
|
||||
if self._pauseCount == 0 and self._completionState is None:
|
||||
self._cooperator._addTask(self)
|
||||
|
||||
|
||||
def _completeWith(self, completionState, deferredResult):
|
||||
"""
|
||||
@param completionState: a L{TaskFinished} exception or a subclass
|
||||
thereof, indicating what exception should be raised when subsequent
|
||||
operations are performed.
|
||||
|
||||
@param deferredResult: the result to fire all the deferreds with.
|
||||
"""
|
||||
self._completionState = completionState
|
||||
self._completionResult = deferredResult
|
||||
if not self._pauseCount:
|
||||
self._cooperator._removeTask(self)
|
||||
|
||||
# The Deferreds need to be invoked after all this is completed, because
|
||||
# a Deferred may want to manipulate other tasks in a Cooperator. For
|
||||
# example, if you call "stop()" on a cooperator in a callback on a
|
||||
# Deferred returned from whenDone(), this CooperativeTask must be gone
|
||||
# from the Cooperator by that point so that _completeWith is not
|
||||
# invoked reentrantly; that would cause these Deferreds to blow up with
|
||||
# an AlreadyCalledError, or the _removeTask to fail with a ValueError.
|
||||
for d in self._deferreds:
|
||||
d.callback(deferredResult)
|
||||
|
||||
|
||||
def stop(self):
|
||||
"""
|
||||
Stop further processing of this task.
|
||||
|
||||
@raise TaskFinished: if this L{CooperativeTask} has previously
|
||||
completed, via C{stop}, completion, or failure.
|
||||
"""
|
||||
self._checkFinish()
|
||||
self._completeWith(TaskStopped(), Failure(TaskStopped()))
|
||||
|
||||
|
||||
def _checkFinish(self):
|
||||
"""
|
||||
If this task has been stopped, raise the appropriate subclass of
|
||||
L{TaskFinished}.
|
||||
"""
|
||||
if self._completionState is not None:
|
||||
raise self._completionState
|
||||
|
||||
|
||||
def _oneWorkUnit(self):
|
||||
"""
|
||||
Perform one unit of work for this task, retrieving one item from its
|
||||
iterator, stopping if there are no further items in the iterator, and
|
||||
pausing if the result was a L{defer.Deferred}.
|
||||
"""
|
||||
try:
|
||||
result = next(self._iterator)
|
||||
except StopIteration:
|
||||
self._completeWith(TaskDone(), self._iterator)
|
||||
except:
|
||||
self._completeWith(TaskFailed(), Failure())
|
||||
else:
|
||||
if isinstance(result, defer.Deferred):
|
||||
self.pause()
|
||||
def failLater(f):
|
||||
self._completeWith(TaskFailed(), f)
|
||||
result.addCallbacks(lambda result: self.resume(),
|
||||
failLater)
|
||||
|
||||
|
||||
|
||||
class Cooperator(object):
|
||||
"""
|
||||
Cooperative task scheduler.
|
||||
|
||||
A cooperative task is an iterator where each iteration represents an
|
||||
atomic unit of work. When the iterator yields, it allows the
|
||||
L{Cooperator} to decide which of its tasks to execute next. If the
|
||||
iterator yields a L{defer.Deferred} then work will pause until the
|
||||
L{defer.Deferred} fires and completes its callback chain.
|
||||
|
||||
When a L{Cooperator} has more than one task, it distributes work between
|
||||
all tasks.
|
||||
|
||||
There are two ways to add tasks to a L{Cooperator}, L{cooperate} and
|
||||
L{coiterate}. L{cooperate} is the more useful of the two, as it returns a
|
||||
L{CooperativeTask}, which can be L{paused<CooperativeTask.pause>},
|
||||
L{resumed<CooperativeTask.resume>} and L{waited
|
||||
on<CooperativeTask.whenDone>}. L{coiterate} has the same effect, but
|
||||
returns only a L{defer.Deferred} that fires when the task is done.
|
||||
|
||||
L{Cooperator} can be used for many things, including but not limited to:
|
||||
|
||||
- running one or more computationally intensive tasks without blocking
|
||||
- limiting parallelism by running a subset of the total tasks
|
||||
simultaneously
|
||||
- doing one thing, waiting for a L{Deferred<defer.Deferred>} to fire,
|
||||
doing the next thing, repeat (i.e. serializing a sequence of
|
||||
asynchronous tasks)
|
||||
|
||||
Multiple L{Cooperator}s do not cooperate with each other, so for most
|
||||
cases you should use the L{global cooperator<task.cooperate>}.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
terminationPredicateFactory=_Timer,
|
||||
scheduler=_defaultScheduler,
|
||||
started=True):
|
||||
"""
|
||||
Create a scheduler-like object to which iterators may be added.
|
||||
|
||||
@param terminationPredicateFactory: A no-argument callable which will
|
||||
be invoked at the beginning of each step and should return a
|
||||
no-argument callable which will return True when the step should be
|
||||
terminated. The default factory is time-based and allows iterators to
|
||||
run for 1/100th of a second at a time.
|
||||
|
||||
@param scheduler: A one-argument callable which takes a no-argument
|
||||
callable and should invoke it at some future point. This will be used
|
||||
to schedule each step of this Cooperator.
|
||||
|
||||
@param started: A boolean which indicates whether iterators should be
|
||||
stepped as soon as they are added, or if they will be queued up until
|
||||
L{Cooperator.start} is called.
|
||||
"""
|
||||
self._tasks = []
|
||||
self._metarator = iter(())
|
||||
self._terminationPredicateFactory = terminationPredicateFactory
|
||||
self._scheduler = scheduler
|
||||
self._delayedCall = None
|
||||
self._stopped = False
|
||||
self._started = started
|
||||
|
||||
|
||||
def coiterate(self, iterator, doneDeferred=None):
|
||||
"""
|
||||
Add an iterator to the list of iterators this L{Cooperator} is
|
||||
currently running.
|
||||
|
||||
Equivalent to L{cooperate}, but returns a L{defer.Deferred} that will
|
||||
be fired when the task is done.
|
||||
|
||||
@param doneDeferred: If specified, this will be the Deferred used as
|
||||
the completion deferred. It is suggested that you use the default,
|
||||
which creates a new Deferred for you.
|
||||
|
||||
@return: a Deferred that will fire when the iterator finishes.
|
||||
"""
|
||||
if doneDeferred is None:
|
||||
doneDeferred = defer.Deferred()
|
||||
CooperativeTask(iterator, self).whenDone().chainDeferred(doneDeferred)
|
||||
return doneDeferred
|
||||
|
||||
|
||||
def cooperate(self, iterator):
|
||||
"""
|
||||
Start running the given iterator as a long-running cooperative task, by
|
||||
calling next() on it as a periodic timed event.
|
||||
|
||||
@param iterator: the iterator to invoke.
|
||||
|
||||
@return: a L{CooperativeTask} object representing this task.
|
||||
"""
|
||||
return CooperativeTask(iterator, self)
|
||||
|
||||
|
||||
def _addTask(self, task):
|
||||
"""
|
||||
Add a L{CooperativeTask} object to this L{Cooperator}.
|
||||
"""
|
||||
if self._stopped:
|
||||
self._tasks.append(task) # XXX silly, I know, but _completeWith
|
||||
# does the inverse
|
||||
task._completeWith(SchedulerStopped(), Failure(SchedulerStopped()))
|
||||
else:
|
||||
self._tasks.append(task)
|
||||
self._reschedule()
|
||||
|
||||
|
||||
def _removeTask(self, task):
|
||||
"""
|
||||
Remove a L{CooperativeTask} from this L{Cooperator}.
|
||||
"""
|
||||
self._tasks.remove(task)
|
||||
# If no work left to do, cancel the delayed call:
|
||||
if not self._tasks and self._delayedCall:
|
||||
self._delayedCall.cancel()
|
||||
self._delayedCall = None
|
||||
|
||||
|
||||
def _tasksWhileNotStopped(self):
|
||||
"""
|
||||
Yield all L{CooperativeTask} objects in a loop as long as this
|
||||
L{Cooperator}'s termination condition has not been met.
|
||||
"""
|
||||
terminator = self._terminationPredicateFactory()
|
||||
while self._tasks:
|
||||
for t in self._metarator:
|
||||
yield t
|
||||
if terminator():
|
||||
return
|
||||
self._metarator = iter(self._tasks)
|
||||
|
||||
|
||||
def _tick(self):
|
||||
"""
|
||||
Run one scheduler tick.
|
||||
"""
|
||||
self._delayedCall = None
|
||||
for taskObj in self._tasksWhileNotStopped():
|
||||
taskObj._oneWorkUnit()
|
||||
self._reschedule()
|
||||
|
||||
|
||||
_mustScheduleOnStart = False
|
||||
def _reschedule(self):
|
||||
if not self._started:
|
||||
self._mustScheduleOnStart = True
|
||||
return
|
||||
if self._delayedCall is None and self._tasks:
|
||||
self._delayedCall = self._scheduler(self._tick)
|
||||
|
||||
|
||||
def start(self):
|
||||
"""
|
||||
Begin scheduling steps.
|
||||
"""
|
||||
self._stopped = False
|
||||
self._started = True
|
||||
if self._mustScheduleOnStart:
|
||||
del self._mustScheduleOnStart
|
||||
self._reschedule()
|
||||
|
||||
|
||||
def stop(self):
|
||||
"""
|
||||
Stop scheduling steps. Errback the completion Deferreds of all
|
||||
iterators which have been added and forget about them.
|
||||
"""
|
||||
self._stopped = True
|
||||
for taskObj in self._tasks:
|
||||
taskObj._completeWith(SchedulerStopped(),
|
||||
Failure(SchedulerStopped()))
|
||||
self._tasks = []
|
||||
if self._delayedCall is not None:
|
||||
self._delayedCall.cancel()
|
||||
self._delayedCall = None
|
||||
|
||||
|
||||
@property
|
||||
def running(self):
|
||||
"""
|
||||
Is this L{Cooperator} is currently running?
|
||||
|
||||
@return: C{True} if the L{Cooperator} is running, C{False} otherwise.
|
||||
@rtype: C{bool}
|
||||
"""
|
||||
return (self._started and not self._stopped)
|
||||
|
||||
|
||||
|
||||
_theCooperator = Cooperator()
|
||||
|
||||
def coiterate(iterator):
|
||||
"""
|
||||
Cooperatively iterate over the given iterator, dividing runtime between it
|
||||
and all other iterators which have been passed to this function and not yet
|
||||
exhausted.
|
||||
|
||||
@param iterator: the iterator to invoke.
|
||||
|
||||
@return: a Deferred that will fire when the iterator finishes.
|
||||
"""
|
||||
return _theCooperator.coiterate(iterator)
|
||||
|
||||
|
||||
|
||||
def cooperate(iterator):
|
||||
"""
|
||||
Start running the given iterator as a long-running cooperative task, by
|
||||
calling next() on it as a periodic timed event.
|
||||
|
||||
This is very useful if you have computationally expensive tasks that you
|
||||
want to run without blocking the reactor. Just break each task up so that
|
||||
it yields frequently, pass it in here and the global L{Cooperator} will
|
||||
make sure work is distributed between them without blocking longer than a
|
||||
single iteration of a single task.
|
||||
|
||||
@param iterator: the iterator to invoke.
|
||||
|
||||
@return: a L{CooperativeTask} object representing this task.
|
||||
"""
|
||||
return _theCooperator.cooperate(iterator)
|
||||
|
||||
|
||||
|
||||
@implementer(IReactorTime)
|
||||
class Clock:
|
||||
"""
|
||||
Provide a deterministic, easily-controlled implementation of
|
||||
L{IReactorTime.callLater}. This is commonly useful for writing
|
||||
deterministic unit tests for code which schedules events using this API.
|
||||
"""
|
||||
|
||||
rightNow = 0.0
|
||||
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
|
||||
|
||||
def seconds(self):
|
||||
"""
|
||||
Pretend to be time.time(). This is used internally when an operation
|
||||
such as L{IDelayedCall.reset} needs to determine a time value
|
||||
relative to the current time.
|
||||
|
||||
@rtype: C{float}
|
||||
@return: The time which should be considered the current time.
|
||||
"""
|
||||
return self.rightNow
|
||||
|
||||
|
||||
def _sortCalls(self):
|
||||
"""
|
||||
Sort the pending calls according to the time they are scheduled.
|
||||
"""
|
||||
self.calls.sort(key=lambda a: a.getTime())
|
||||
|
||||
|
||||
def callLater(self, when, what, *a, **kw):
|
||||
"""
|
||||
See L{twisted.internet.interfaces.IReactorTime.callLater}.
|
||||
"""
|
||||
dc = base.DelayedCall(self.seconds() + when,
|
||||
what, a, kw,
|
||||
self.calls.remove,
|
||||
lambda c: None,
|
||||
self.seconds)
|
||||
self.calls.append(dc)
|
||||
self._sortCalls()
|
||||
return dc
|
||||
|
||||
|
||||
def getDelayedCalls(self):
|
||||
"""
|
||||
See L{twisted.internet.interfaces.IReactorTime.getDelayedCalls}
|
||||
"""
|
||||
return self.calls
|
||||
|
||||
|
||||
def advance(self, amount):
|
||||
"""
|
||||
Move time on this clock forward by the given amount and run whatever
|
||||
pending calls should be run.
|
||||
|
||||
@type amount: C{float}
|
||||
@param amount: The number of seconds which to advance this clock's
|
||||
time.
|
||||
"""
|
||||
self.rightNow += amount
|
||||
self._sortCalls()
|
||||
while self.calls and self.calls[0].getTime() <= self.seconds():
|
||||
call = self.calls.pop(0)
|
||||
call.called = 1
|
||||
call.func(*call.args, **call.kw)
|
||||
self._sortCalls()
|
||||
|
||||
|
||||
def pump(self, timings):
|
||||
"""
|
||||
Advance incrementally by the given set of times.
|
||||
|
||||
@type timings: iterable of C{float}
|
||||
"""
|
||||
for amount in timings:
|
||||
self.advance(amount)
|
||||
|
||||
|
||||
|
||||
def deferLater(clock, delay, callable=None, *args, **kw):
|
||||
"""
|
||||
Call the given function after a certain period of time has passed.
|
||||
|
||||
@type clock: L{IReactorTime} provider
|
||||
@param clock: The object which will be used to schedule the delayed
|
||||
call.
|
||||
|
||||
@type delay: C{float} or C{int}
|
||||
@param delay: The number of seconds to wait before calling the function.
|
||||
|
||||
@param callable: The object to call after the delay.
|
||||
|
||||
@param *args: The positional arguments to pass to C{callable}.
|
||||
|
||||
@param **kw: The keyword arguments to pass to C{callable}.
|
||||
|
||||
@rtype: L{defer.Deferred}
|
||||
|
||||
@return: A deferred that fires with the result of the callable when the
|
||||
specified time has elapsed.
|
||||
"""
|
||||
def deferLaterCancel(deferred):
|
||||
delayedCall.cancel()
|
||||
d = defer.Deferred(deferLaterCancel)
|
||||
if callable is not None:
|
||||
d.addCallback(lambda ignored: callable(*args, **kw))
|
||||
delayedCall = clock.callLater(delay, d.callback, None)
|
||||
return d
|
||||
|
||||
|
||||
|
||||
def react(main, argv=(), _reactor=None):
|
||||
"""
|
||||
Call C{main} and run the reactor until the L{Deferred} it returns fires.
|
||||
|
||||
This is intended as the way to start up an application with a well-defined
|
||||
completion condition. Use it to write clients or one-off asynchronous
|
||||
operations. Prefer this to calling C{reactor.run} directly, as this
|
||||
function will also:
|
||||
|
||||
- Take care to call C{reactor.stop} once and only once, and at the right
|
||||
time.
|
||||
- Log any failures from the C{Deferred} returned by C{main}.
|
||||
- Exit the application when done, with exit code 0 in case of success and
|
||||
1 in case of failure. If C{main} fails with a C{SystemExit} error, the
|
||||
code returned is used.
|
||||
|
||||
The following demonstrates the signature of a C{main} function which can be
|
||||
used with L{react}::
|
||||
def main(reactor, username, password):
|
||||
return defer.succeed('ok')
|
||||
|
||||
task.react(main, ('alice', 'secret'))
|
||||
|
||||
@param main: A callable which returns a L{Deferred}. It should
|
||||
take the reactor as its first parameter, followed by the elements of
|
||||
C{argv}.
|
||||
|
||||
@param argv: A list of arguments to pass to C{main}. If omitted the
|
||||
callable will be invoked with no additional arguments.
|
||||
|
||||
@param _reactor: An implementation detail to allow easier unit testing. Do
|
||||
not supply this parameter.
|
||||
|
||||
@since: 12.3
|
||||
"""
|
||||
if _reactor is None:
|
||||
from twisted.internet import reactor as _reactor
|
||||
finished = main(_reactor, *argv)
|
||||
codes = [0]
|
||||
|
||||
stopping = []
|
||||
_reactor.addSystemEventTrigger('before', 'shutdown', stopping.append, True)
|
||||
|
||||
def stop(result, stopReactor):
|
||||
if stopReactor:
|
||||
try:
|
||||
_reactor.stop()
|
||||
except ReactorNotRunning:
|
||||
pass
|
||||
|
||||
if isinstance(result, Failure):
|
||||
if result.check(SystemExit) is not None:
|
||||
code = result.value.code
|
||||
else:
|
||||
log.err(result, "main function encountered error")
|
||||
code = 1
|
||||
codes[0] = code
|
||||
|
||||
def cbFinish(result):
|
||||
if stopping:
|
||||
stop(result, False)
|
||||
else:
|
||||
_reactor.callWhenRunning(stop, result, True)
|
||||
|
||||
finished.addBoth(cbFinish)
|
||||
_reactor.run()
|
||||
sys.exit(codes[0])
|
||||
|
||||
|
||||
__all__ = [
|
||||
'LoopingCall',
|
||||
|
||||
'Clock',
|
||||
|
||||
'SchedulerStopped', 'Cooperator', 'coiterate',
|
||||
|
||||
'deferLater', 'react']
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,221 @@
|
||||
# -*- test-case-name: twisted.internet.test.test_coroutines -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for C{await} support in Deferreds.
|
||||
|
||||
These tests can only work and be imported on Python 3.5+!
|
||||
"""
|
||||
|
||||
import types
|
||||
|
||||
from twisted.python.failure import Failure
|
||||
from twisted.internet.defer import (
|
||||
Deferred, maybeDeferred, ensureDeferred, fail
|
||||
)
|
||||
from twisted.trial.unittest import TestCase
|
||||
from twisted.internet.task import Clock
|
||||
|
||||
class SampleException(Exception):
|
||||
"""
|
||||
A specific sample exception for testing.
|
||||
"""
|
||||
|
||||
|
||||
class AwaitTests(TestCase):
|
||||
"""
|
||||
Tests for using Deferreds in conjunction with PEP-492.
|
||||
"""
|
||||
def test_awaitReturnsIterable(self):
|
||||
"""
|
||||
C{Deferred.__await__} returns an iterable.
|
||||
"""
|
||||
d = Deferred()
|
||||
awaitedDeferred = d.__await__()
|
||||
self.assertEqual(awaitedDeferred, iter(awaitedDeferred))
|
||||
|
||||
|
||||
def test_ensureDeferred(self):
|
||||
"""
|
||||
L{ensureDeferred} will turn a coroutine into a L{Deferred}.
|
||||
"""
|
||||
async def run():
|
||||
d = Deferred()
|
||||
d.callback("bar")
|
||||
await d
|
||||
res = await run2()
|
||||
return res
|
||||
|
||||
async def run2():
|
||||
d = Deferred()
|
||||
d.callback("foo")
|
||||
res = await d
|
||||
return res
|
||||
|
||||
# It's a coroutine...
|
||||
r = run()
|
||||
self.assertIsInstance(r, types.CoroutineType)
|
||||
|
||||
# Now it's a Deferred.
|
||||
d = ensureDeferred(r)
|
||||
self.assertIsInstance(d, Deferred)
|
||||
|
||||
# The Deferred has the result we want.
|
||||
res = self.successResultOf(d)
|
||||
self.assertEqual(res, "foo")
|
||||
|
||||
|
||||
def test_basic(self):
|
||||
"""
|
||||
L{ensureDeferred} allows a function to C{await} on a L{Deferred}.
|
||||
"""
|
||||
async def run():
|
||||
d = Deferred()
|
||||
d.callback("foo")
|
||||
res = await d
|
||||
return res
|
||||
|
||||
d = ensureDeferred(run())
|
||||
res = self.successResultOf(d)
|
||||
self.assertEqual(res, "foo")
|
||||
|
||||
|
||||
def test_exception(self):
|
||||
"""
|
||||
An exception in a coroutine wrapped with L{ensureDeferred} will cause
|
||||
the returned L{Deferred} to fire with a failure.
|
||||
"""
|
||||
async def run():
|
||||
d = Deferred()
|
||||
d.callback("foo")
|
||||
await d
|
||||
raise ValueError("Oh no!")
|
||||
|
||||
d = ensureDeferred(run())
|
||||
res = self.failureResultOf(d)
|
||||
self.assertEqual(type(res.value), ValueError)
|
||||
self.assertEqual(res.value.args, ("Oh no!",))
|
||||
|
||||
|
||||
def test_synchronousDeferredFailureTraceback(self):
|
||||
"""
|
||||
When a Deferred is awaited upon that has already failed with a Failure
|
||||
that has a traceback, both the place that the synchronous traceback
|
||||
comes from and the awaiting line are shown in the traceback.
|
||||
"""
|
||||
def raises():
|
||||
raise SampleException()
|
||||
it = maybeDeferred(raises)
|
||||
async def doomed():
|
||||
return await it
|
||||
failure = self.failureResultOf(ensureDeferred(doomed()))
|
||||
|
||||
self.assertIn(", in doomed\n", failure.getTraceback())
|
||||
self.assertIn(", in raises\n", failure.getTraceback())
|
||||
|
||||
|
||||
def test_asyncDeferredFailureTraceback(self):
|
||||
"""
|
||||
When a Deferred is awaited upon that later fails with a Failure that
|
||||
has a traceback, both the place that the synchronous traceback comes
|
||||
from and the awaiting line are shown in the traceback.
|
||||
"""
|
||||
def returnsFailure():
|
||||
try:
|
||||
raise SampleException()
|
||||
except SampleException:
|
||||
return Failure()
|
||||
it = Deferred()
|
||||
async def doomed():
|
||||
return await it
|
||||
started = ensureDeferred(doomed())
|
||||
self.assertNoResult(started)
|
||||
it.errback(returnsFailure())
|
||||
failure = self.failureResultOf(started)
|
||||
self.assertIn(", in doomed\n", failure.getTraceback())
|
||||
self.assertIn(", in returnsFailure\n", failure.getTraceback())
|
||||
|
||||
|
||||
def test_twoDeep(self):
|
||||
"""
|
||||
A coroutine wrapped with L{ensureDeferred} that awaits a L{Deferred}
|
||||
suspends its execution until the inner L{Deferred} fires.
|
||||
"""
|
||||
reactor = Clock()
|
||||
sections = []
|
||||
|
||||
async def runone():
|
||||
sections.append(2)
|
||||
d = Deferred()
|
||||
reactor.callLater(1, d.callback, 2)
|
||||
await d
|
||||
sections.append(3)
|
||||
return "Yay!"
|
||||
|
||||
|
||||
async def run():
|
||||
sections.append(1)
|
||||
result = await runone()
|
||||
sections.append(4)
|
||||
d = Deferred()
|
||||
reactor.callLater(1, d.callback, 1)
|
||||
await d
|
||||
sections.append(5)
|
||||
return result
|
||||
|
||||
d = ensureDeferred(run())
|
||||
|
||||
reactor.advance(0.9)
|
||||
self.assertEqual(sections, [1, 2])
|
||||
|
||||
reactor.advance(0.1)
|
||||
self.assertEqual(sections, [1, 2, 3, 4])
|
||||
|
||||
reactor.advance(0.9)
|
||||
self.assertEqual(sections, [1, 2, 3, 4])
|
||||
|
||||
reactor.advance(0.1)
|
||||
self.assertEqual(sections, [1, 2, 3, 4, 5])
|
||||
|
||||
res = self.successResultOf(d)
|
||||
self.assertEqual(res, "Yay!")
|
||||
|
||||
|
||||
def test_reraise(self):
|
||||
"""
|
||||
Awaiting an already failed Deferred will raise the exception.
|
||||
"""
|
||||
async def test():
|
||||
try:
|
||||
await fail(ValueError("Boom"))
|
||||
except ValueError as e:
|
||||
self.assertEqual(e.args, ("Boom",))
|
||||
return 1
|
||||
return 0
|
||||
|
||||
res = self.successResultOf(ensureDeferred(test()))
|
||||
self.assertEqual(res, 1)
|
||||
|
||||
|
||||
def test_chained(self):
|
||||
"""
|
||||
Awaiting a paused & chained Deferred will give the result when it has
|
||||
one.
|
||||
"""
|
||||
reactor = Clock()
|
||||
|
||||
async def test():
|
||||
d = Deferred()
|
||||
d2 = Deferred()
|
||||
d.addCallback(lambda ignored: d2)
|
||||
|
||||
d.callback(None)
|
||||
reactor.callLater(0, d2.callback, "bye")
|
||||
return await d
|
||||
|
||||
d = ensureDeferred(test())
|
||||
reactor.advance(0.1)
|
||||
|
||||
res = self.successResultOf(d)
|
||||
self.assertEqual(res, "bye")
|
||||
@@ -0,0 +1,177 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
POSIX implementation of local network interface enumeration.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
import sys, socket
|
||||
|
||||
from socket import AF_INET, AF_INET6, inet_ntop
|
||||
from ctypes import (
|
||||
CDLL, POINTER, Structure, c_char_p, c_ushort, c_int,
|
||||
c_uint32, c_uint8, c_void_p, c_ubyte, pointer, cast)
|
||||
from ctypes.util import find_library
|
||||
|
||||
from twisted.python.compat import _PY3, nativeString
|
||||
|
||||
if _PY3:
|
||||
# Once #6070 is implemented, this can be replaced with the implementation
|
||||
# from that ticket:
|
||||
def chr(i):
|
||||
"""
|
||||
Python 3 implementation of Python 2 chr(), i.e. convert an integer to
|
||||
corresponding byte.
|
||||
"""
|
||||
return bytes([i])
|
||||
|
||||
|
||||
libc = CDLL(find_library("c"))
|
||||
|
||||
if sys.platform.startswith('freebsd') or sys.platform == 'darwin':
|
||||
_sockaddrCommon = [
|
||||
("sin_len", c_uint8),
|
||||
("sin_family", c_uint8),
|
||||
]
|
||||
else:
|
||||
_sockaddrCommon = [
|
||||
("sin_family", c_ushort),
|
||||
]
|
||||
|
||||
|
||||
|
||||
class in_addr(Structure):
|
||||
_fields_ = [
|
||||
("in_addr", c_ubyte * 4),
|
||||
]
|
||||
|
||||
|
||||
|
||||
class in6_addr(Structure):
|
||||
_fields_ = [
|
||||
("in_addr", c_ubyte * 16),
|
||||
]
|
||||
|
||||
|
||||
|
||||
class sockaddr(Structure):
|
||||
_fields_ = _sockaddrCommon + [
|
||||
("sin_port", c_ushort),
|
||||
]
|
||||
|
||||
|
||||
|
||||
class sockaddr_in(Structure):
|
||||
_fields_ = _sockaddrCommon + [
|
||||
("sin_port", c_ushort),
|
||||
("sin_addr", in_addr),
|
||||
]
|
||||
|
||||
|
||||
|
||||
class sockaddr_in6(Structure):
|
||||
_fields_ = _sockaddrCommon + [
|
||||
("sin_port", c_ushort),
|
||||
("sin_flowinfo", c_uint32),
|
||||
("sin_addr", in6_addr),
|
||||
]
|
||||
|
||||
|
||||
|
||||
class ifaddrs(Structure):
|
||||
pass
|
||||
|
||||
ifaddrs_p = POINTER(ifaddrs)
|
||||
ifaddrs._fields_ = [
|
||||
('ifa_next', ifaddrs_p),
|
||||
('ifa_name', c_char_p),
|
||||
('ifa_flags', c_uint32),
|
||||
('ifa_addr', POINTER(sockaddr)),
|
||||
('ifa_netmask', POINTER(sockaddr)),
|
||||
('ifa_dstaddr', POINTER(sockaddr)),
|
||||
('ifa_data', c_void_p)]
|
||||
|
||||
getifaddrs = libc.getifaddrs
|
||||
getifaddrs.argtypes = [POINTER(ifaddrs_p)]
|
||||
getifaddrs.restype = c_int
|
||||
|
||||
freeifaddrs = libc.freeifaddrs
|
||||
freeifaddrs.argtypes = [ifaddrs_p]
|
||||
|
||||
|
||||
|
||||
def _maybeCleanupScopeIndex(family, packed):
|
||||
"""
|
||||
On FreeBSD, kill the embedded interface indices in link-local scoped
|
||||
addresses.
|
||||
|
||||
@param family: The address family of the packed address - one of the
|
||||
I{socket.AF_*} constants.
|
||||
|
||||
@param packed: The packed representation of the address (ie, the bytes of a
|
||||
I{in_addr} field).
|
||||
@type packed: L{bytes}
|
||||
|
||||
@return: The packed address with any FreeBSD-specific extra bits cleared.
|
||||
@rtype: L{bytes}
|
||||
|
||||
@see: U{https://twistedmatrix.com/trac/ticket/6843}
|
||||
@see: U{http://www.freebsd.org/doc/en/books/developers-handbook/ipv6.html#ipv6-scope-index}
|
||||
|
||||
@note: Indications are that the need for this will be gone in FreeBSD >=10.
|
||||
"""
|
||||
if sys.platform.startswith('freebsd') and packed[:2] == b"\xfe\x80":
|
||||
return packed[:2] + b"\x00\x00" + packed[4:]
|
||||
return packed
|
||||
|
||||
|
||||
|
||||
def _interfaces():
|
||||
"""
|
||||
Call C{getifaddrs(3)} and return a list of tuples of interface name, address
|
||||
family, and human-readable address representing its results.
|
||||
"""
|
||||
ifaddrs = ifaddrs_p()
|
||||
if getifaddrs(pointer(ifaddrs)) < 0:
|
||||
raise OSError()
|
||||
results = []
|
||||
try:
|
||||
while ifaddrs:
|
||||
if ifaddrs[0].ifa_addr:
|
||||
family = ifaddrs[0].ifa_addr[0].sin_family
|
||||
if family == AF_INET:
|
||||
addr = cast(ifaddrs[0].ifa_addr, POINTER(sockaddr_in))
|
||||
elif family == AF_INET6:
|
||||
addr = cast(ifaddrs[0].ifa_addr, POINTER(sockaddr_in6))
|
||||
else:
|
||||
addr = None
|
||||
|
||||
if addr:
|
||||
packed = b''.join(map(chr, addr[0].sin_addr.in_addr[:]))
|
||||
packed = _maybeCleanupScopeIndex(family, packed)
|
||||
results.append((
|
||||
ifaddrs[0].ifa_name,
|
||||
family,
|
||||
inet_ntop(family, packed)))
|
||||
|
||||
ifaddrs = ifaddrs[0].ifa_next
|
||||
finally:
|
||||
freeifaddrs(ifaddrs)
|
||||
return results
|
||||
|
||||
|
||||
|
||||
def posixGetLinkLocalIPv6Addresses():
|
||||
"""
|
||||
Return a list of strings in colon-hex format representing all the link local
|
||||
IPv6 addresses available on the system, as reported by I{getifaddrs(3)}.
|
||||
"""
|
||||
retList = []
|
||||
for (interface, family, address) in _interfaces():
|
||||
interface = nativeString(interface)
|
||||
address = nativeString(address)
|
||||
if family == socket.AF_INET6 and address.startswith('fe80:'):
|
||||
retList.append('%s%%%s' % (address, interface))
|
||||
return retList
|
||||
@@ -0,0 +1,26 @@
|
||||
|
||||
This is a self-signed certificate authority certificate to be used in tests.
|
||||
|
||||
It was created with the following command:
|
||||
certcreate -f thing1.pem -h fake-ca-1.example.com -e noreply@example.com \
|
||||
-S 1234 -o 'Twisted Matrix Labs'
|
||||
|
||||
'certcreate' may be obtained from <http://divmod.org/trac/wiki/DivmodEpsilon>
|
||||
|
||||
-----BEGIN CERTIFICATE-----
|
||||
MIICwjCCAisCAgTSMA0GCSqGSIb3DQEBBAUAMIGoMREwDwYDVQQLEwhTZWN1cml0
|
||||
eTEcMBoGA1UEChMTVHdpc3RlZCBNYXRyaXggTGFiczEeMBwGA1UEAxMVZmFrZS1j
|
||||
YS0xLmV4YW1wbGUuY29tMREwDwYDVQQIEwhOZXcgWW9yazELMAkGA1UEBhMCVVMx
|
||||
IjAgBgkqhkiG9w0BCQEWE25vcmVwbHlAZXhhbXBsZS5jb20xETAPBgNVBAcTCE5l
|
||||
dyBZb3JrMB4XDTEwMDkyMTAxMjUxNFoXDTExMDkyMTAxMjUxNFowgagxETAPBgNV
|
||||
BAsTCFNlY3VyaXR5MRwwGgYDVQQKExNUd2lzdGVkIE1hdHJpeCBMYWJzMR4wHAYD
|
||||
VQQDExVmYWtlLWNhLTEuZXhhbXBsZS5jb20xETAPBgNVBAgTCE5ldyBZb3JrMQsw
|
||||
CQYDVQQGEwJVUzEiMCAGCSqGSIb3DQEJARYTbm9yZXBseUBleGFtcGxlLmNvbTER
|
||||
MA8GA1UEBxMITmV3IFlvcmswgZ8wDQYJKoZIhvcNAQEBBQADgY0AMIGJAoGBALRb
|
||||
VqC0CsaFgq1vbwPfs8zoP3ZYC/0sPMv0RJN+f3Dc7Q6YgNHS7o7TM3uAy/McADeW
|
||||
rwVuNJGe9k+4ZBHysmBH1sG64fHT5TlK9saPcUQqkubSWj4cKSDtVbQERWqC5Dy+
|
||||
qTQeZGYoPEMlnRXgMpST04DG//Dgzi4PYqUOjwxTAgMBAAEwDQYJKoZIhvcNAQEE
|
||||
BQADgYEAqNEdMXWEs8Co76wxL3/cSV3MjiAroVxJdI/3EzlnfPi1JeibbdWw31fC
|
||||
bn6428KTjjfhS31zo1yHG3YNXFEJXRscwLAH7ogz5kJwZMy/oS/96EFM10bkNwkK
|
||||
v+nWKN8i3t/E5TEIl3BPN8tchtWmH0rycVuzs5LwaewwR1AnUE4=
|
||||
-----END CERTIFICATE-----
|
||||
@@ -0,0 +1,27 @@
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
import sys
|
||||
import os
|
||||
|
||||
try:
|
||||
# On Windows, stdout is not opened in binary mode by default,
|
||||
# so newline characters are munged on writing, interfering with
|
||||
# the tests.
|
||||
import msvcrt
|
||||
msvcrt.setmode(sys.stdout.fileno(), os.O_BINARY)
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
|
||||
# Loop over each of the arguments given and print it to stdout
|
||||
for arg in sys.argv[1:]:
|
||||
res = arg + chr(0)
|
||||
|
||||
if sys.version_info < (3, 0):
|
||||
stdout = sys.stdout
|
||||
else:
|
||||
stdout = sys.stdout.buffer
|
||||
res = res.encode(sys.getfilesystemencoding(), "surrogateescape")
|
||||
|
||||
stdout.write(res)
|
||||
stdout.flush()
|
||||
@@ -0,0 +1,350 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Utilities for unit testing reactor implementations.
|
||||
|
||||
The main feature of this module is L{ReactorBuilder}, a base class for use when
|
||||
writing interface/blackbox tests for reactor implementations. Test case classes
|
||||
for reactor features should subclass L{ReactorBuilder} instead of
|
||||
L{SynchronousTestCase}. All of the features of L{SynchronousTestCase} will be
|
||||
available. Additionally, the tests will automatically be applied to all
|
||||
available reactor implementations.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
__metaclass__ = type
|
||||
|
||||
__all__ = ['TestTimeoutError', 'ReactorBuilder', 'needsRunningReactor']
|
||||
|
||||
import os, signal, time
|
||||
|
||||
from twisted.trial.unittest import SynchronousTestCase, SkipTest
|
||||
from twisted.trial.util import DEFAULT_TIMEOUT_DURATION, acquireAttribute
|
||||
from twisted.python.runtime import platform
|
||||
from twisted.python.reflect import namedAny
|
||||
from twisted.python.deprecate import _fullyQualifiedName as fullyQualifiedName
|
||||
|
||||
from twisted.python import log
|
||||
from twisted.python.failure import Failure
|
||||
from twisted.python.compat import _PY3
|
||||
|
||||
|
||||
# Access private APIs.
|
||||
if platform.isWindows():
|
||||
process = None
|
||||
else:
|
||||
from twisted.internet import process
|
||||
|
||||
|
||||
|
||||
class TestTimeoutError(Exception):
|
||||
"""
|
||||
The reactor was still running after the timeout period elapsed in
|
||||
L{ReactorBuilder.runReactor}.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
def needsRunningReactor(reactor, thunk):
|
||||
"""
|
||||
Various functions within these tests need an already-running reactor at
|
||||
some point. They need to stop the reactor when the test has completed, and
|
||||
that means calling reactor.stop(). However, reactor.stop() raises an
|
||||
exception if the reactor isn't already running, so if the L{Deferred} that
|
||||
a particular API under test returns fires synchronously (as especially an
|
||||
endpoint's C{connect()} method may do, if the connect is to a local
|
||||
interface address) then the test won't be able to stop the reactor being
|
||||
tested and finish. So this calls C{thunk} only once C{reactor} is running.
|
||||
|
||||
(This is just an alias for
|
||||
L{twisted.internet.interfaces.IReactorCore.callWhenRunning} on the given
|
||||
reactor parameter, in order to centrally reference the above paragraph and
|
||||
repeating it everywhere as a comment.)
|
||||
|
||||
@param reactor: the L{twisted.internet.interfaces.IReactorCore} under test
|
||||
|
||||
@param thunk: a 0-argument callable, which eventually finishes the test in
|
||||
question, probably in a L{Deferred} callback.
|
||||
"""
|
||||
reactor.callWhenRunning(thunk)
|
||||
|
||||
|
||||
|
||||
def stopOnError(case, reactor, publisher=None):
|
||||
"""
|
||||
Stop the reactor as soon as any error is logged on the given publisher.
|
||||
|
||||
This is beneficial for tests which will wait for a L{Deferred} to fire
|
||||
before completing (by passing or failing). Certain implementation bugs may
|
||||
prevent the L{Deferred} from firing with any result at all (consider a
|
||||
protocol's {dataReceived} method that raises an exception: this exception
|
||||
is logged but it won't ever cause a L{Deferred} to fire). In that case the
|
||||
test would have to complete by timing out which is a much less desirable
|
||||
outcome than completing as soon as the unexpected error is encountered.
|
||||
|
||||
@param case: A L{SynchronousTestCase} to use to clean up the necessary log
|
||||
observer when the test is over.
|
||||
@param reactor: The reactor to stop.
|
||||
@param publisher: A L{LogPublisher} to watch for errors. If L{None}, the
|
||||
global log publisher will be watched.
|
||||
"""
|
||||
if publisher is None:
|
||||
from twisted.python import log as publisher
|
||||
running = [None]
|
||||
def stopIfError(event):
|
||||
if running and event.get('isError'):
|
||||
running.pop()
|
||||
reactor.stop()
|
||||
publisher.addObserver(stopIfError)
|
||||
case.addCleanup(publisher.removeObserver, stopIfError)
|
||||
|
||||
|
||||
|
||||
class ReactorBuilder:
|
||||
"""
|
||||
L{SynchronousTestCase} mixin which provides a reactor-creation API. This
|
||||
mixin defines C{setUp} and C{tearDown}, so mix it in before
|
||||
L{SynchronousTestCase} or call its methods from the overridden ones in the
|
||||
subclass.
|
||||
|
||||
@cvar skippedReactors: A dict mapping FQPN strings of reactors for
|
||||
which the tests defined by this class will be skipped to strings
|
||||
giving the skip message.
|
||||
@cvar requiredInterfaces: A C{list} of interfaces which the reactor must
|
||||
provide or these tests will be skipped. The default, L{None}, means
|
||||
that no interfaces are required.
|
||||
@ivar reactorFactory: A no-argument callable which returns the reactor to
|
||||
use for testing.
|
||||
@ivar originalHandler: The SIGCHLD handler which was installed when setUp
|
||||
ran and which will be re-installed when tearDown runs.
|
||||
@ivar _reactors: A list of FQPN strings giving the reactors for which
|
||||
L{SynchronousTestCase}s will be created.
|
||||
"""
|
||||
|
||||
_reactors = [
|
||||
# Select works everywhere
|
||||
"twisted.internet.selectreactor.SelectReactor",
|
||||
]
|
||||
|
||||
if platform.isWindows():
|
||||
# PortableGtkReactor is only really interesting on Windows,
|
||||
# but not really Windows specific; if you want you can
|
||||
# temporarily move this up to the all-platforms list to test
|
||||
# it on other platforms. It's not there in general because
|
||||
# it's not _really_ worth it to support on other platforms,
|
||||
# since no one really wants to use it on other platforms.
|
||||
_reactors.extend([
|
||||
"twisted.internet.gtk2reactor.PortableGtkReactor",
|
||||
"twisted.internet.gireactor.PortableGIReactor",
|
||||
"twisted.internet.gtk3reactor.PortableGtk3Reactor",
|
||||
"twisted.internet.win32eventreactor.Win32Reactor",
|
||||
"twisted.internet.iocpreactor.reactor.IOCPReactor"])
|
||||
else:
|
||||
_reactors.extend([
|
||||
"twisted.internet.glib2reactor.Glib2Reactor",
|
||||
"twisted.internet.gtk2reactor.Gtk2Reactor",
|
||||
"twisted.internet.gireactor.GIReactor",
|
||||
"twisted.internet.gtk3reactor.Gtk3Reactor"])
|
||||
|
||||
if _PY3:
|
||||
_reactors.append(
|
||||
"twisted.internet.asyncioreactor.AsyncioSelectorReactor")
|
||||
|
||||
if platform.isMacOSX():
|
||||
_reactors.append("twisted.internet.cfreactor.CFReactor")
|
||||
else:
|
||||
_reactors.extend([
|
||||
"twisted.internet.pollreactor.PollReactor",
|
||||
"twisted.internet.epollreactor.EPollReactor"])
|
||||
if not platform.isLinux():
|
||||
# Presumably Linux is not going to start supporting kqueue, so
|
||||
# skip even trying this configuration.
|
||||
_reactors.extend([
|
||||
# Support KQueue on non-OS-X POSIX platforms for now.
|
||||
"twisted.internet.kqreactor.KQueueReactor",
|
||||
])
|
||||
|
||||
reactorFactory = None
|
||||
originalHandler = None
|
||||
requiredInterfaces = None
|
||||
skippedReactors = {}
|
||||
|
||||
def setUp(self):
|
||||
"""
|
||||
Clear the SIGCHLD handler, if there is one, to ensure an environment
|
||||
like the one which exists prior to a call to L{reactor.run}.
|
||||
"""
|
||||
if not platform.isWindows():
|
||||
self.originalHandler = signal.signal(signal.SIGCHLD, signal.SIG_DFL)
|
||||
|
||||
|
||||
def tearDown(self):
|
||||
"""
|
||||
Restore the original SIGCHLD handler and reap processes as long as
|
||||
there seem to be any remaining.
|
||||
"""
|
||||
if self.originalHandler is not None:
|
||||
signal.signal(signal.SIGCHLD, self.originalHandler)
|
||||
if process is not None:
|
||||
begin = time.time()
|
||||
while process.reapProcessHandlers:
|
||||
log.msg(
|
||||
"ReactorBuilder.tearDown reaping some processes %r" % (
|
||||
process.reapProcessHandlers,))
|
||||
process.reapAllProcesses()
|
||||
|
||||
# The process should exit on its own. However, if it
|
||||
# doesn't, we're stuck in this loop forever. To avoid
|
||||
# hanging the test suite, eventually give the process some
|
||||
# help exiting and move on.
|
||||
time.sleep(0.001)
|
||||
if time.time() - begin > 60:
|
||||
for pid in process.reapProcessHandlers:
|
||||
os.kill(pid, signal.SIGKILL)
|
||||
raise Exception(
|
||||
"Timeout waiting for child processes to exit: %r" % (
|
||||
process.reapProcessHandlers,))
|
||||
|
||||
|
||||
def unbuildReactor(self, reactor):
|
||||
"""
|
||||
Clean up any resources which may have been allocated for the given
|
||||
reactor by its creation or by a test which used it.
|
||||
"""
|
||||
# Chris says:
|
||||
#
|
||||
# XXX These explicit calls to clean up the waker (and any other
|
||||
# internal readers) should become obsolete when bug #3063 is
|
||||
# fixed. -radix, 2008-02-29. Fortunately it should probably cause an
|
||||
# error when bug #3063 is fixed, so it should be removed in the same
|
||||
# branch that fixes it.
|
||||
#
|
||||
# -exarkun
|
||||
reactor._uninstallHandler()
|
||||
if getattr(reactor, '_internalReaders', None) is not None:
|
||||
for reader in reactor._internalReaders:
|
||||
reactor.removeReader(reader)
|
||||
reader.connectionLost(None)
|
||||
reactor._internalReaders.clear()
|
||||
|
||||
# Here's an extra thing unrelated to wakers but necessary for
|
||||
# cleaning up after the reactors we make. -exarkun
|
||||
reactor.disconnectAll()
|
||||
|
||||
# It would also be bad if any timed calls left over were allowed to
|
||||
# run.
|
||||
calls = reactor.getDelayedCalls()
|
||||
for c in calls:
|
||||
c.cancel()
|
||||
|
||||
|
||||
def buildReactor(self):
|
||||
"""
|
||||
Create and return a reactor using C{self.reactorFactory}.
|
||||
"""
|
||||
try:
|
||||
from twisted.internet.cfreactor import CFReactor
|
||||
from twisted.internet import reactor as globalReactor
|
||||
except ImportError:
|
||||
pass
|
||||
else:
|
||||
if (isinstance(globalReactor, CFReactor)
|
||||
and self.reactorFactory is CFReactor):
|
||||
raise SkipTest(
|
||||
"CFReactor uses APIs which manipulate global state, "
|
||||
"so it's not safe to run its own reactor-builder tests "
|
||||
"under itself")
|
||||
try:
|
||||
reactor = self.reactorFactory()
|
||||
except:
|
||||
# Unfortunately, not all errors which result in a reactor
|
||||
# being unusable are detectable without actually
|
||||
# instantiating the reactor. So we catch some more here
|
||||
# and skip the test if necessary. We also log it to aid
|
||||
# with debugging, but flush the logged error so the test
|
||||
# doesn't fail.
|
||||
log.err(None, "Failed to install reactor")
|
||||
self.flushLoggedErrors()
|
||||
raise SkipTest(Failure().getErrorMessage())
|
||||
else:
|
||||
if self.requiredInterfaces is not None:
|
||||
missing = [
|
||||
required for required in self.requiredInterfaces
|
||||
if not required.providedBy(reactor)]
|
||||
if missing:
|
||||
self.unbuildReactor(reactor)
|
||||
raise SkipTest("%s does not provide %s" % (
|
||||
fullyQualifiedName(reactor.__class__),
|
||||
",".join([fullyQualifiedName(x) for x in missing])))
|
||||
self.addCleanup(self.unbuildReactor, reactor)
|
||||
return reactor
|
||||
|
||||
|
||||
def getTimeout(self):
|
||||
"""
|
||||
Determine how long to run the test before considering it failed.
|
||||
|
||||
@return: A C{int} or C{float} giving a number of seconds.
|
||||
"""
|
||||
return acquireAttribute(self._parents, 'timeout', DEFAULT_TIMEOUT_DURATION)
|
||||
|
||||
|
||||
def runReactor(self, reactor, timeout=None):
|
||||
"""
|
||||
Run the reactor for at most the given amount of time.
|
||||
|
||||
@param reactor: The reactor to run.
|
||||
|
||||
@type timeout: C{int} or C{float}
|
||||
@param timeout: The maximum amount of time, specified in seconds, to
|
||||
allow the reactor to run. If the reactor is still running after
|
||||
this much time has elapsed, it will be stopped and an exception
|
||||
raised. If L{None}, the default test method timeout imposed by
|
||||
Trial will be used. This depends on the L{IReactorTime}
|
||||
implementation of C{reactor} for correct operation.
|
||||
|
||||
@raise TestTimeoutError: If the reactor is still running after
|
||||
C{timeout} seconds.
|
||||
"""
|
||||
if timeout is None:
|
||||
timeout = self.getTimeout()
|
||||
|
||||
timedOut = []
|
||||
def stop():
|
||||
timedOut.append(None)
|
||||
reactor.stop()
|
||||
|
||||
timedOutCall = reactor.callLater(timeout, stop)
|
||||
reactor.run()
|
||||
if timedOut:
|
||||
raise TestTimeoutError(
|
||||
"reactor still running after %s seconds" % (timeout,))
|
||||
else:
|
||||
timedOutCall.cancel()
|
||||
|
||||
|
||||
def makeTestCaseClasses(cls):
|
||||
"""
|
||||
Create a L{SynchronousTestCase} subclass which mixes in C{cls} for each
|
||||
known reactor and return a dict mapping their names to them.
|
||||
"""
|
||||
classes = {}
|
||||
for reactor in cls._reactors:
|
||||
shortReactorName = reactor.split(".")[-1]
|
||||
name = (cls.__name__ + "." + shortReactorName + "Tests").replace(".", "_")
|
||||
class testcase(cls, SynchronousTestCase):
|
||||
__module__ = cls.__module__
|
||||
if reactor in cls.skippedReactors:
|
||||
skip = cls.skippedReactors[reactor]
|
||||
try:
|
||||
reactorFactory = namedAny(reactor)
|
||||
except:
|
||||
skip = Failure().getErrorMessage()
|
||||
testcase.__name__ = name
|
||||
if hasattr(cls, "__qualname__"):
|
||||
testcase.__qualname__ = ".".join(cls.__qualname__.split()[0:-1] + [name])
|
||||
classes[testcase.__name__] = testcase
|
||||
return classes
|
||||
makeTestCaseClasses = classmethod(makeTestCaseClasses)
|
||||
@@ -0,0 +1,450 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.internet.base}.
|
||||
"""
|
||||
|
||||
import socket
|
||||
try:
|
||||
from Queue import Queue
|
||||
except ImportError:
|
||||
from queue import Queue
|
||||
|
||||
from zope.interface import implementer
|
||||
|
||||
from twisted.python.threadpool import ThreadPool
|
||||
from twisted.internet.interfaces import (IReactorTime, IReactorThreads,
|
||||
IResolverSimple)
|
||||
from twisted.internet.error import DNSLookupError
|
||||
from twisted.internet._resolver import FirstOneWins
|
||||
from twisted.internet.defer import Deferred
|
||||
from twisted.internet.base import ThreadedResolver, DelayedCall, ReactorBase
|
||||
from twisted.internet.task import Clock
|
||||
from twisted.trial.unittest import TestCase, SkipTest
|
||||
|
||||
|
||||
@implementer(IReactorTime, IReactorThreads)
|
||||
class FakeReactor(object):
|
||||
"""
|
||||
A fake reactor implementation which just supports enough reactor APIs for
|
||||
L{ThreadedResolver}.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._clock = Clock()
|
||||
self.callLater = self._clock.callLater
|
||||
|
||||
self._threadpool = ThreadPool()
|
||||
self._threadpool.start()
|
||||
self.getThreadPool = lambda: self._threadpool
|
||||
|
||||
self._threadCalls = Queue()
|
||||
|
||||
|
||||
def callFromThread(self, f, *args, **kwargs):
|
||||
self._threadCalls.put((f, args, kwargs))
|
||||
|
||||
|
||||
def _runThreadCalls(self):
|
||||
f, args, kwargs = self._threadCalls.get()
|
||||
f(*args, **kwargs)
|
||||
|
||||
|
||||
def _stop(self):
|
||||
self._threadpool.stop()
|
||||
|
||||
|
||||
|
||||
class ThreadedResolverTests(TestCase):
|
||||
"""
|
||||
Tests for L{ThreadedResolver}.
|
||||
"""
|
||||
def test_success(self):
|
||||
"""
|
||||
L{ThreadedResolver.getHostByName} returns a L{Deferred} which fires
|
||||
with the value returned by the call to L{socket.gethostbyname} in the
|
||||
threadpool of the reactor passed to L{ThreadedResolver.__init__}.
|
||||
"""
|
||||
ip = "10.0.0.17"
|
||||
name = "foo.bar.example.com"
|
||||
timeout = 30
|
||||
|
||||
reactor = FakeReactor()
|
||||
self.addCleanup(reactor._stop)
|
||||
|
||||
lookedUp = []
|
||||
resolvedTo = []
|
||||
def fakeGetHostByName(name):
|
||||
lookedUp.append(name)
|
||||
return ip
|
||||
|
||||
self.patch(socket, 'gethostbyname', fakeGetHostByName)
|
||||
|
||||
resolver = ThreadedResolver(reactor)
|
||||
d = resolver.getHostByName(name, (timeout,))
|
||||
d.addCallback(resolvedTo.append)
|
||||
|
||||
reactor._runThreadCalls()
|
||||
|
||||
self.assertEqual(lookedUp, [name])
|
||||
self.assertEqual(resolvedTo, [ip])
|
||||
|
||||
# Make sure that any timeout-related stuff gets cleaned up.
|
||||
reactor._clock.advance(timeout + 1)
|
||||
self.assertEqual(reactor._clock.calls, [])
|
||||
|
||||
|
||||
def test_failure(self):
|
||||
"""
|
||||
L{ThreadedResolver.getHostByName} returns a L{Deferred} which fires a
|
||||
L{Failure} if the call to L{socket.gethostbyname} raises an exception.
|
||||
"""
|
||||
timeout = 30
|
||||
|
||||
reactor = FakeReactor()
|
||||
self.addCleanup(reactor._stop)
|
||||
|
||||
def fakeGetHostByName(name):
|
||||
raise IOError("ENOBUFS (this is a funny joke)")
|
||||
|
||||
self.patch(socket, 'gethostbyname', fakeGetHostByName)
|
||||
|
||||
failedWith = []
|
||||
resolver = ThreadedResolver(reactor)
|
||||
d = resolver.getHostByName("some.name", (timeout,))
|
||||
self.assertFailure(d, DNSLookupError)
|
||||
d.addCallback(failedWith.append)
|
||||
|
||||
reactor._runThreadCalls()
|
||||
|
||||
self.assertEqual(len(failedWith), 1)
|
||||
|
||||
# Make sure that any timeout-related stuff gets cleaned up.
|
||||
reactor._clock.advance(timeout + 1)
|
||||
self.assertEqual(reactor._clock.calls, [])
|
||||
|
||||
|
||||
def test_timeout(self):
|
||||
"""
|
||||
If L{socket.gethostbyname} does not complete before the specified
|
||||
timeout elapsed, the L{Deferred} returned by
|
||||
L{ThreadedResolver.getHostByName} fails with L{DNSLookupError}.
|
||||
"""
|
||||
timeout = 10
|
||||
|
||||
reactor = FakeReactor()
|
||||
self.addCleanup(reactor._stop)
|
||||
|
||||
result = Queue()
|
||||
def fakeGetHostByName(name):
|
||||
raise result.get()
|
||||
|
||||
self.patch(socket, 'gethostbyname', fakeGetHostByName)
|
||||
|
||||
failedWith = []
|
||||
resolver = ThreadedResolver(reactor)
|
||||
d = resolver.getHostByName("some.name", (timeout,))
|
||||
self.assertFailure(d, DNSLookupError)
|
||||
d.addCallback(failedWith.append)
|
||||
|
||||
reactor._clock.advance(timeout - 1)
|
||||
self.assertEqual(failedWith, [])
|
||||
reactor._clock.advance(1)
|
||||
self.assertEqual(len(failedWith), 1)
|
||||
|
||||
# Eventually the socket.gethostbyname does finish - in this case, with
|
||||
# an exception. Nobody cares, though.
|
||||
result.put(IOError("The I/O was errorful"))
|
||||
|
||||
|
||||
def test_resolverGivenStr(self):
|
||||
"""
|
||||
L{ThreadedResolver.getHostByName} is passed L{str}, encoded using IDNA
|
||||
if required.
|
||||
"""
|
||||
calls = []
|
||||
|
||||
@implementer(IResolverSimple)
|
||||
class FakeResolver(object):
|
||||
def getHostByName(self, name, timeouts=()):
|
||||
calls.append(name)
|
||||
return Deferred()
|
||||
|
||||
class JustEnoughReactor(ReactorBase):
|
||||
def installWaker(self):
|
||||
pass
|
||||
|
||||
fake = FakeResolver()
|
||||
reactor = JustEnoughReactor()
|
||||
reactor.installResolver(fake)
|
||||
rec = FirstOneWins(Deferred())
|
||||
reactor.nameResolver.resolveHostName(
|
||||
rec, u"example.example")
|
||||
reactor.nameResolver.resolveHostName(
|
||||
rec, "example.example")
|
||||
reactor.nameResolver.resolveHostName(
|
||||
rec, u"v\xe4\xe4ntynyt.example")
|
||||
reactor.nameResolver.resolveHostName(
|
||||
rec, u"\u0440\u0444.example")
|
||||
reactor.nameResolver.resolveHostName(
|
||||
rec, "xn----7sbb4ac0ad0be6cf.xn--p1ai")
|
||||
|
||||
self.assertEqual(len(calls), 5)
|
||||
self.assertEqual(list(map(type, calls)), [str]*5)
|
||||
self.assertEqual("example.example", calls[0])
|
||||
self.assertEqual("example.example", calls[1])
|
||||
self.assertEqual("xn--vntynyt-5waa.example", calls[2])
|
||||
self.assertEqual("xn--p1ai.example", calls[3])
|
||||
self.assertEqual("xn----7sbb4ac0ad0be6cf.xn--p1ai", calls[4])
|
||||
|
||||
|
||||
|
||||
def nothing():
|
||||
"""
|
||||
Function used by L{DelayedCallTests.test_str}.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class DelayedCallMixin(object):
|
||||
"""
|
||||
L{DelayedCall}
|
||||
"""
|
||||
def _getDelayedCallAt(self, time):
|
||||
"""
|
||||
Get a L{DelayedCall} instance at a given C{time}.
|
||||
|
||||
@param time: The absolute time at which the returned L{DelayedCall}
|
||||
will be scheduled.
|
||||
"""
|
||||
def noop(call):
|
||||
pass
|
||||
return DelayedCall(time, lambda: None, (), {}, noop, noop, None)
|
||||
|
||||
|
||||
def setUp(self):
|
||||
"""
|
||||
Create two L{DelayedCall} instanced scheduled to run at different
|
||||
times.
|
||||
"""
|
||||
self.zero = self._getDelayedCallAt(0)
|
||||
self.one = self._getDelayedCallAt(1)
|
||||
|
||||
|
||||
def test_str(self):
|
||||
"""
|
||||
The string representation of a L{DelayedCall} instance, as returned by
|
||||
L{str}, includes the unsigned id of the instance, as well as its state,
|
||||
the function to be called, and the function arguments.
|
||||
"""
|
||||
dc = DelayedCall(12, nothing, (3, ), {"A": 5}, None, None, lambda: 1.5)
|
||||
self.assertEqual(
|
||||
str(dc),
|
||||
"<DelayedCall 0x%x [10.5s] called=0 cancelled=0 nothing(3, A=5)>"
|
||||
% (id(dc),),
|
||||
)
|
||||
|
||||
|
||||
def test_repr(self):
|
||||
"""
|
||||
The string representation of a L{DelayedCall} instance, as returned by
|
||||
{repr}, is identical to that returned by L{str}.
|
||||
"""
|
||||
dc = DelayedCall(13, nothing, (6, ), {"A": 9}, None, None, lambda: 1.6)
|
||||
self.assertEqual(str(dc), repr(dc))
|
||||
|
||||
|
||||
def test_lt(self):
|
||||
"""
|
||||
For two instances of L{DelayedCall} C{a} and C{b}, C{a < b} is true
|
||||
if and only if C{a} is scheduled to run before C{b}.
|
||||
"""
|
||||
zero, one = self.zero, self.one
|
||||
self.assertTrue(zero < one)
|
||||
self.assertFalse(one < zero)
|
||||
self.assertFalse(zero < zero)
|
||||
self.assertFalse(one < one)
|
||||
|
||||
|
||||
def test_le(self):
|
||||
"""
|
||||
For two instances of L{DelayedCall} C{a} and C{b}, C{a <= b} is true
|
||||
if and only if C{a} is scheduled to run before C{b} or at the same
|
||||
time as C{b}.
|
||||
"""
|
||||
zero, one = self.zero, self.one
|
||||
self.assertTrue(zero <= one)
|
||||
self.assertFalse(one <= zero)
|
||||
self.assertTrue(zero <= zero)
|
||||
self.assertTrue(one <= one)
|
||||
|
||||
|
||||
def test_gt(self):
|
||||
"""
|
||||
For two instances of L{DelayedCall} C{a} and C{b}, C{a > b} is true
|
||||
if and only if C{a} is scheduled to run after C{b}.
|
||||
"""
|
||||
zero, one = self.zero, self.one
|
||||
self.assertTrue(one > zero)
|
||||
self.assertFalse(zero > one)
|
||||
self.assertFalse(zero > zero)
|
||||
self.assertFalse(one > one)
|
||||
|
||||
|
||||
def test_ge(self):
|
||||
"""
|
||||
For two instances of L{DelayedCall} C{a} and C{b}, C{a > b} is true
|
||||
if and only if C{a} is scheduled to run after C{b} or at the same
|
||||
time as C{b}.
|
||||
"""
|
||||
zero, one = self.zero, self.one
|
||||
self.assertTrue(one >= zero)
|
||||
self.assertFalse(zero >= one)
|
||||
self.assertTrue(zero >= zero)
|
||||
self.assertTrue(one >= one)
|
||||
|
||||
|
||||
def test_eq(self):
|
||||
"""
|
||||
A L{DelayedCall} instance is only equal to itself.
|
||||
"""
|
||||
# Explicitly use == here, instead of assertEqual, to be more
|
||||
# confident __eq__ is being tested.
|
||||
self.assertFalse(self.zero == self.one)
|
||||
self.assertTrue(self.zero == self.zero)
|
||||
self.assertTrue(self.one == self.one)
|
||||
|
||||
|
||||
def test_ne(self):
|
||||
"""
|
||||
A L{DelayedCall} instance is not equal to any other object.
|
||||
"""
|
||||
# Explicitly use != here, instead of assertEqual, to be more
|
||||
# confident __ne__ is being tested.
|
||||
self.assertTrue(self.zero != self.one)
|
||||
self.assertFalse(self.zero != self.zero)
|
||||
self.assertFalse(self.one != self.one)
|
||||
|
||||
|
||||
|
||||
class DelayedCallNoDebugTests(DelayedCallMixin, TestCase):
|
||||
"""
|
||||
L{DelayedCall}
|
||||
"""
|
||||
def setUp(self):
|
||||
"""
|
||||
Turn debug off.
|
||||
"""
|
||||
self.patch(DelayedCall, 'debug', False)
|
||||
DelayedCallMixin.setUp(self)
|
||||
|
||||
|
||||
def test_str(self):
|
||||
"""
|
||||
The string representation of a L{DelayedCall} instance, as returned by
|
||||
L{str}, includes the unsigned id of the instance, as well as its state,
|
||||
the function to be called, and the function arguments.
|
||||
"""
|
||||
dc = DelayedCall(12, nothing, (3, ), {"A": 5}, None, None, lambda: 1.5)
|
||||
expected = (
|
||||
"<DelayedCall 0x{:x} [10.5s] called=0 cancelled=0 "
|
||||
"nothing(3, A=5)>".format(id(dc)))
|
||||
self.assertEqual(str(dc), expected)
|
||||
|
||||
|
||||
|
||||
class DelayedCallDebugTests(DelayedCallMixin, TestCase):
|
||||
"""
|
||||
L{DelayedCall}
|
||||
"""
|
||||
def setUp(self):
|
||||
"""
|
||||
Turn debug on.
|
||||
"""
|
||||
self.patch(DelayedCall, 'debug', True)
|
||||
DelayedCallMixin.setUp(self)
|
||||
|
||||
|
||||
def test_str(self):
|
||||
"""
|
||||
The string representation of a L{DelayedCall} instance, as returned by
|
||||
L{str}, includes the unsigned id of the instance, as well as its state,
|
||||
the function to be called, and the function arguments.
|
||||
"""
|
||||
dc = DelayedCall(12, nothing, (3, ), {"A": 5}, None, None, lambda: 1.5)
|
||||
expectedRegexp = (
|
||||
"<DelayedCall 0x{:x} \\[10.5s\\] called=0 cancelled=0 "
|
||||
"nothing\\(3, A=5\\)\n\n"
|
||||
"traceback at creation:".format(id(dc)))
|
||||
self.assertRegex(
|
||||
str(dc), expectedRegexp)
|
||||
|
||||
|
||||
|
||||
class TestSpySignalCapturingReactor(ReactorBase):
|
||||
|
||||
"""
|
||||
Subclass of ReactorBase to capture signals delivered to the
|
||||
reactor for inspection.
|
||||
"""
|
||||
|
||||
def installWaker(self):
|
||||
"""
|
||||
Required method, unused.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class ReactorBaseSignalTests(TestCase):
|
||||
|
||||
"""
|
||||
Tests to exercise ReactorBase's signal exit reporting path.
|
||||
"""
|
||||
|
||||
def test_exitSignalDefaultsToNone(self):
|
||||
"""
|
||||
The default value of the _exitSignal attribute is None.
|
||||
"""
|
||||
reactor = TestSpySignalCapturingReactor()
|
||||
self.assertIs(None, reactor._exitSignal)
|
||||
|
||||
|
||||
def test_captureSIGINT(self):
|
||||
"""
|
||||
ReactorBase's SIGINT handler saves the value of SIGINT to the
|
||||
_exitSignal attribute.
|
||||
"""
|
||||
reactor = TestSpySignalCapturingReactor()
|
||||
reactor.sigInt(signal.SIGINT, None)
|
||||
self.assertEquals(signal.SIGINT, reactor._exitSignal)
|
||||
|
||||
|
||||
def test_captureSIGTERM(self):
|
||||
"""
|
||||
ReactorBase's SIGTERM handler saves the value of SIGTERM to the
|
||||
_exitSignal attribute.
|
||||
"""
|
||||
reactor = TestSpySignalCapturingReactor()
|
||||
reactor.sigTerm(signal.SIGTERM, None)
|
||||
self.assertEquals(signal.SIGTERM, reactor._exitSignal)
|
||||
|
||||
|
||||
def test_captureSIGBREAK(self):
|
||||
"""
|
||||
ReactorBase's SIGBREAK handler saves the value of SIGBREAK to the
|
||||
_exitSignal attribute.
|
||||
"""
|
||||
if not hasattr(signal, "SIGBREAK"):
|
||||
raise SkipTest("signal module does not have SIGBREAK")
|
||||
|
||||
reactor = TestSpySignalCapturingReactor()
|
||||
reactor.sigBreak(signal.SIGBREAK, None)
|
||||
self.assertEquals(signal.SIGBREAK, reactor._exitSignal)
|
||||
|
||||
|
||||
|
||||
try:
|
||||
import signal
|
||||
except ImportError:
|
||||
ReactorBaseSignalTests.skip = "signal module not available"
|
||||
@@ -0,0 +1,73 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.internet._baseprocess} which implements process-related
|
||||
functionality that is useful in all platforms supporting L{IReactorProcess}.
|
||||
"""
|
||||
|
||||
__metaclass__ = type
|
||||
|
||||
from twisted.python.deprecate import getWarningMethod, setWarningMethod
|
||||
from twisted.trial.unittest import TestCase
|
||||
from twisted.internet._baseprocess import BaseProcess
|
||||
|
||||
|
||||
class BaseProcessTests(TestCase):
|
||||
"""
|
||||
Tests for L{BaseProcess}, a parent class for other classes which represent
|
||||
processes which implements functionality common to many different process
|
||||
implementations.
|
||||
"""
|
||||
def test_callProcessExited(self):
|
||||
"""
|
||||
L{BaseProcess._callProcessExited} calls the C{processExited} method of
|
||||
its C{proto} attribute and passes it a L{Failure} wrapping the given
|
||||
exception.
|
||||
"""
|
||||
class FakeProto:
|
||||
reason = None
|
||||
|
||||
def processExited(self, reason):
|
||||
self.reason = reason
|
||||
|
||||
reason = RuntimeError("fake reason")
|
||||
process = BaseProcess(FakeProto())
|
||||
process._callProcessExited(reason)
|
||||
process.proto.reason.trap(RuntimeError)
|
||||
self.assertIs(reason, process.proto.reason.value)
|
||||
|
||||
|
||||
def test_callProcessExitedMissing(self):
|
||||
"""
|
||||
L{BaseProcess._callProcessExited} emits a L{DeprecationWarning} if the
|
||||
object referred to by its C{proto} attribute has no C{processExited}
|
||||
method.
|
||||
"""
|
||||
class FakeProto:
|
||||
pass
|
||||
|
||||
reason = object()
|
||||
process = BaseProcess(FakeProto())
|
||||
|
||||
self.addCleanup(setWarningMethod, getWarningMethod())
|
||||
warnings = []
|
||||
def collect(message, category, stacklevel):
|
||||
warnings.append((message, category, stacklevel))
|
||||
setWarningMethod(collect)
|
||||
|
||||
process._callProcessExited(reason)
|
||||
|
||||
[(message, category, stacklevel)] = warnings
|
||||
self.assertEqual(
|
||||
message,
|
||||
"Since Twisted 8.2, IProcessProtocol.processExited is required. "
|
||||
"%s.%s must implement it." % (
|
||||
FakeProto.__module__, FakeProto.__name__))
|
||||
self.assertIs(category, DeprecationWarning)
|
||||
# The stacklevel doesn't really make sense for this kind of
|
||||
# deprecation. Requiring it to be 0 will at least avoid pointing to
|
||||
# any part of Twisted or a random part of the application's code, which
|
||||
# I think would be more misleading than having it point inside the
|
||||
# warning system itself. -exarkun
|
||||
self.assertEqual(stacklevel, 0)
|
||||
@@ -0,0 +1,58 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
A wrapper for L{twisted.internet.test._awaittests}, as that test module
|
||||
includes keywords not valid in Pythons before 3.5.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
from twisted.python.compat import _PY35PLUS, _PY3, execfile
|
||||
from twisted.python.filepath import FilePath
|
||||
from twisted.trial import unittest
|
||||
|
||||
|
||||
if _PY35PLUS:
|
||||
_path = FilePath(__file__).parent().child("_awaittests.py.3only")
|
||||
|
||||
_g = {"__name__": __name__ + ".3-only.awaittests"}
|
||||
execfile(_path.path, _g)
|
||||
AwaitTests = _g["AwaitTests"]
|
||||
else:
|
||||
class AwaitTests(unittest.SynchronousTestCase):
|
||||
"""
|
||||
A dummy class to show that this test file was discovered but the tests
|
||||
are unable to be run in this version of Python.
|
||||
"""
|
||||
skip = "async/await is not available before Python 3.5"
|
||||
|
||||
def test_notAvailable(self):
|
||||
"""
|
||||
A skipped test to show that this was not run because the Python is
|
||||
too old.
|
||||
"""
|
||||
|
||||
|
||||
if _PY3:
|
||||
_path = FilePath(__file__).parent().child("_yieldfromtests.py.3only")
|
||||
|
||||
_g = {"__name__": __name__ + ".3-only.yieldfromtests"}
|
||||
execfile(_path.path, _g)
|
||||
YieldFromTests = _g["YieldFromTests"]
|
||||
else:
|
||||
class YieldFromTests(unittest.SynchronousTestCase):
|
||||
"""
|
||||
A dummy class to show that this test file was discovered but the tests
|
||||
are unable to be run in this version of Python.
|
||||
"""
|
||||
skip = "yield from is not available before Python 3"
|
||||
|
||||
def test_notAvailable(self):
|
||||
"""
|
||||
A skipped test to show that this was not run because the Python is
|
||||
too old.
|
||||
"""
|
||||
|
||||
|
||||
__all__ = ["AwaitTests", "YieldFromTests"]
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,248 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.internet.epollreactor}.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
from twisted.trial.unittest import TestCase
|
||||
try:
|
||||
from twisted.internet.epollreactor import _ContinuousPolling
|
||||
except ImportError:
|
||||
_ContinuousPolling = None
|
||||
from twisted.internet.task import Clock
|
||||
from twisted.internet.error import ConnectionDone
|
||||
|
||||
|
||||
|
||||
class Descriptor(object):
|
||||
"""
|
||||
Records reads and writes, as if it were a C{FileDescriptor}.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.events = []
|
||||
|
||||
|
||||
def fileno(self):
|
||||
return 1
|
||||
|
||||
|
||||
def doRead(self):
|
||||
self.events.append("read")
|
||||
|
||||
|
||||
def doWrite(self):
|
||||
self.events.append("write")
|
||||
|
||||
|
||||
def connectionLost(self, reason):
|
||||
reason.trap(ConnectionDone)
|
||||
self.events.append("lost")
|
||||
|
||||
|
||||
|
||||
class ContinuousPollingTests(TestCase):
|
||||
"""
|
||||
L{_ContinuousPolling} can be used to read and write from C{FileDescriptor}
|
||||
objects.
|
||||
"""
|
||||
|
||||
def test_addReader(self):
|
||||
"""
|
||||
Adding a reader when there was previously no reader starts up a
|
||||
C{LoopingCall}.
|
||||
"""
|
||||
poller = _ContinuousPolling(Clock())
|
||||
self.assertIsNone(poller._loop)
|
||||
reader = object()
|
||||
self.assertFalse(poller.isReading(reader))
|
||||
poller.addReader(reader)
|
||||
self.assertIsNotNone(poller._loop)
|
||||
self.assertTrue(poller._loop.running)
|
||||
self.assertIs(poller._loop.clock, poller._reactor)
|
||||
self.assertTrue(poller.isReading(reader))
|
||||
|
||||
|
||||
def test_addWriter(self):
|
||||
"""
|
||||
Adding a writer when there was previously no writer starts up a
|
||||
C{LoopingCall}.
|
||||
"""
|
||||
poller = _ContinuousPolling(Clock())
|
||||
self.assertIsNone(poller._loop)
|
||||
writer = object()
|
||||
self.assertFalse(poller.isWriting(writer))
|
||||
poller.addWriter(writer)
|
||||
self.assertIsNotNone(poller._loop)
|
||||
self.assertTrue(poller._loop.running)
|
||||
self.assertIs(poller._loop.clock, poller._reactor)
|
||||
self.assertTrue(poller.isWriting(writer))
|
||||
|
||||
|
||||
def test_removeReader(self):
|
||||
"""
|
||||
Removing a reader stops the C{LoopingCall}.
|
||||
"""
|
||||
poller = _ContinuousPolling(Clock())
|
||||
reader = object()
|
||||
poller.addReader(reader)
|
||||
poller.removeReader(reader)
|
||||
self.assertIsNone(poller._loop)
|
||||
self.assertEqual(poller._reactor.getDelayedCalls(), [])
|
||||
self.assertFalse(poller.isReading(reader))
|
||||
|
||||
|
||||
def test_removeWriter(self):
|
||||
"""
|
||||
Removing a writer stops the C{LoopingCall}.
|
||||
"""
|
||||
poller = _ContinuousPolling(Clock())
|
||||
writer = object()
|
||||
poller.addWriter(writer)
|
||||
poller.removeWriter(writer)
|
||||
self.assertIsNone(poller._loop)
|
||||
self.assertEqual(poller._reactor.getDelayedCalls(), [])
|
||||
self.assertFalse(poller.isWriting(writer))
|
||||
|
||||
|
||||
def test_removeUnknown(self):
|
||||
"""
|
||||
Removing unknown readers and writers silently does nothing.
|
||||
"""
|
||||
poller = _ContinuousPolling(Clock())
|
||||
poller.removeWriter(object())
|
||||
poller.removeReader(object())
|
||||
|
||||
|
||||
def test_multipleReadersAndWriters(self):
|
||||
"""
|
||||
Adding multiple readers and writers results in a single
|
||||
C{LoopingCall}.
|
||||
"""
|
||||
poller = _ContinuousPolling(Clock())
|
||||
writer = object()
|
||||
poller.addWriter(writer)
|
||||
self.assertIsNotNone(poller._loop)
|
||||
poller.addWriter(object())
|
||||
self.assertIsNotNone(poller._loop)
|
||||
poller.addReader(object())
|
||||
self.assertIsNotNone(poller._loop)
|
||||
poller.addReader(object())
|
||||
poller.removeWriter(writer)
|
||||
self.assertIsNotNone(poller._loop)
|
||||
self.assertTrue(poller._loop.running)
|
||||
self.assertEqual(len(poller._reactor.getDelayedCalls()), 1)
|
||||
|
||||
|
||||
def test_readerPolling(self):
|
||||
"""
|
||||
Adding a reader causes its C{doRead} to be called every 1
|
||||
milliseconds.
|
||||
"""
|
||||
reactor = Clock()
|
||||
poller = _ContinuousPolling(reactor)
|
||||
desc = Descriptor()
|
||||
poller.addReader(desc)
|
||||
self.assertEqual(desc.events, [])
|
||||
reactor.advance(0.00001)
|
||||
self.assertEqual(desc.events, ["read"])
|
||||
reactor.advance(0.00001)
|
||||
self.assertEqual(desc.events, ["read", "read"])
|
||||
reactor.advance(0.00001)
|
||||
self.assertEqual(desc.events, ["read", "read", "read"])
|
||||
|
||||
|
||||
def test_writerPolling(self):
|
||||
"""
|
||||
Adding a writer causes its C{doWrite} to be called every 1
|
||||
milliseconds.
|
||||
"""
|
||||
reactor = Clock()
|
||||
poller = _ContinuousPolling(reactor)
|
||||
desc = Descriptor()
|
||||
poller.addWriter(desc)
|
||||
self.assertEqual(desc.events, [])
|
||||
reactor.advance(0.001)
|
||||
self.assertEqual(desc.events, ["write"])
|
||||
reactor.advance(0.001)
|
||||
self.assertEqual(desc.events, ["write", "write"])
|
||||
reactor.advance(0.001)
|
||||
self.assertEqual(desc.events, ["write", "write", "write"])
|
||||
|
||||
|
||||
def test_connectionLostOnRead(self):
|
||||
"""
|
||||
If a C{doRead} returns a value indicating disconnection,
|
||||
C{connectionLost} is called on it.
|
||||
"""
|
||||
reactor = Clock()
|
||||
poller = _ContinuousPolling(reactor)
|
||||
desc = Descriptor()
|
||||
desc.doRead = lambda: ConnectionDone()
|
||||
poller.addReader(desc)
|
||||
self.assertEqual(desc.events, [])
|
||||
reactor.advance(0.001)
|
||||
self.assertEqual(desc.events, ["lost"])
|
||||
|
||||
|
||||
def test_connectionLostOnWrite(self):
|
||||
"""
|
||||
If a C{doWrite} returns a value indicating disconnection,
|
||||
C{connectionLost} is called on it.
|
||||
"""
|
||||
reactor = Clock()
|
||||
poller = _ContinuousPolling(reactor)
|
||||
desc = Descriptor()
|
||||
desc.doWrite = lambda: ConnectionDone()
|
||||
poller.addWriter(desc)
|
||||
self.assertEqual(desc.events, [])
|
||||
reactor.advance(0.001)
|
||||
self.assertEqual(desc.events, ["lost"])
|
||||
|
||||
|
||||
def test_removeAll(self):
|
||||
"""
|
||||
L{_ContinuousPolling.removeAll} removes all descriptors and returns
|
||||
the readers and writers.
|
||||
"""
|
||||
poller = _ContinuousPolling(Clock())
|
||||
reader = object()
|
||||
writer = object()
|
||||
both = object()
|
||||
poller.addReader(reader)
|
||||
poller.addReader(both)
|
||||
poller.addWriter(writer)
|
||||
poller.addWriter(both)
|
||||
removed = poller.removeAll()
|
||||
self.assertEqual(poller.getReaders(), [])
|
||||
self.assertEqual(poller.getWriters(), [])
|
||||
self.assertEqual(len(removed), 3)
|
||||
self.assertEqual(set(removed), set([reader, writer, both]))
|
||||
|
||||
|
||||
def test_getReaders(self):
|
||||
"""
|
||||
L{_ContinuousPolling.getReaders} returns a list of the read
|
||||
descriptors.
|
||||
"""
|
||||
poller = _ContinuousPolling(Clock())
|
||||
reader = object()
|
||||
poller.addReader(reader)
|
||||
self.assertIn(reader, poller.getReaders())
|
||||
|
||||
|
||||
def test_getWriters(self):
|
||||
"""
|
||||
L{_ContinuousPolling.getWriters} returns a list of the write
|
||||
descriptors.
|
||||
"""
|
||||
poller = _ContinuousPolling(Clock())
|
||||
writer = object()
|
||||
poller.addWriter(writer)
|
||||
self.assertIn(writer, poller.getWriters())
|
||||
|
||||
if _ContinuousPolling is None:
|
||||
skip = "epoll not supported in this environment."
|
||||
@@ -0,0 +1,99 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Whitebox tests for L{twisted.internet.abstract.FileDescriptor}.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
from zope.interface.verify import verifyClass
|
||||
|
||||
from twisted.internet.abstract import FileDescriptor
|
||||
from twisted.internet.interfaces import IPushProducer
|
||||
from twisted.trial.unittest import SynchronousTestCase
|
||||
|
||||
|
||||
|
||||
class MemoryFile(FileDescriptor):
|
||||
"""
|
||||
A L{FileDescriptor} customization which writes to a Python list in memory
|
||||
with certain limitations.
|
||||
|
||||
@ivar _written: A C{list} of C{bytes} which have been accepted as written.
|
||||
|
||||
@ivar _freeSpace: A C{int} giving the number of bytes which will be accepted
|
||||
by future writes.
|
||||
"""
|
||||
connected = True
|
||||
|
||||
def __init__(self):
|
||||
FileDescriptor.__init__(self, reactor=object())
|
||||
self._written = []
|
||||
self._freeSpace = 0
|
||||
|
||||
|
||||
def startWriting(self):
|
||||
pass
|
||||
|
||||
|
||||
def stopWriting(self):
|
||||
pass
|
||||
|
||||
|
||||
def writeSomeData(self, data):
|
||||
"""
|
||||
Copy at most C{self._freeSpace} bytes from C{data} into C{self._written}.
|
||||
|
||||
@return: A C{int} indicating how many bytes were copied from C{data}.
|
||||
"""
|
||||
acceptLength = min(self._freeSpace, len(data))
|
||||
if acceptLength:
|
||||
self._freeSpace -= acceptLength
|
||||
self._written.append(data[:acceptLength])
|
||||
return acceptLength
|
||||
|
||||
|
||||
|
||||
class FileDescriptorTests(SynchronousTestCase):
|
||||
"""
|
||||
Tests for L{FileDescriptor}.
|
||||
"""
|
||||
def test_writeWithUnicodeRaisesException(self):
|
||||
"""
|
||||
L{FileDescriptor.write} doesn't accept unicode data.
|
||||
"""
|
||||
fileDescriptor = FileDescriptor(reactor=object())
|
||||
self.assertRaises(TypeError, fileDescriptor.write, u'foo')
|
||||
|
||||
|
||||
def test_writeSequenceWithUnicodeRaisesException(self):
|
||||
"""
|
||||
L{FileDescriptor.writeSequence} doesn't accept unicode data.
|
||||
"""
|
||||
fileDescriptor = FileDescriptor(reactor=object())
|
||||
self.assertRaises(
|
||||
TypeError, fileDescriptor.writeSequence, [b'foo', u'bar', b'baz'])
|
||||
|
||||
|
||||
def test_implementInterfaceIPushProducer(self):
|
||||
"""
|
||||
L{FileDescriptor} should implement L{IPushProducer}.
|
||||
"""
|
||||
self.assertTrue(verifyClass(IPushProducer, FileDescriptor))
|
||||
|
||||
|
||||
|
||||
class WriteDescriptorTests(SynchronousTestCase):
|
||||
"""
|
||||
Tests for L{FileDescriptor}'s implementation of L{IWriteDescriptor}.
|
||||
"""
|
||||
def test_kernelBufferFull(self):
|
||||
"""
|
||||
When L{FileDescriptor.writeSomeData} returns C{0} to indicate no more
|
||||
data can be written immediately, L{FileDescriptor.doWrite} returns
|
||||
L{None}.
|
||||
"""
|
||||
descriptor = MemoryFile()
|
||||
descriptor.write(b"hello, world")
|
||||
self.assertIsNone(descriptor.doWrite())
|
||||
@@ -0,0 +1,257 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
GI/GTK3 reactor tests.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import, print_function
|
||||
|
||||
import sys, os
|
||||
try:
|
||||
from twisted.internet import gireactor
|
||||
from gi.repository import Gio
|
||||
except ImportError:
|
||||
gireactor = None
|
||||
gtk3reactor = None
|
||||
else:
|
||||
# gtk3reactor may be unavailable even if gireactor is available; in
|
||||
# particular in pygobject 3.4/gtk 3.6, when no X11 DISPLAY is found.
|
||||
try:
|
||||
from twisted.internet import gtk3reactor
|
||||
except ImportError:
|
||||
gtk3reactor = None
|
||||
else:
|
||||
from gi.repository import Gtk
|
||||
|
||||
from twisted.python.filepath import FilePath
|
||||
from twisted.python.runtime import platform
|
||||
from twisted.internet.defer import Deferred
|
||||
from twisted.internet.error import ReactorAlreadyRunning
|
||||
from twisted.internet.protocol import ProcessProtocol
|
||||
from twisted.trial.unittest import TestCase, SkipTest
|
||||
from twisted.internet.test.reactormixins import ReactorBuilder
|
||||
from twisted.test.test_twisted import SetAsideModule
|
||||
from twisted.internet.interfaces import IReactorProcess
|
||||
from twisted.python.compat import _PY3
|
||||
|
||||
# Skip all tests if gi is unavailable:
|
||||
if gireactor is None:
|
||||
skip = "gtk3/gi not importable"
|
||||
|
||||
|
||||
|
||||
class GApplicationRegistrationTests(ReactorBuilder, TestCase):
|
||||
"""
|
||||
GtkApplication and GApplication are supported by
|
||||
L{twisted.internet.gtk3reactor} and L{twisted.internet.gireactor}.
|
||||
|
||||
We inherit from L{ReactorBuilder} in order to use some of its
|
||||
reactor-running infrastructure, but don't need its test-creation
|
||||
functionality.
|
||||
"""
|
||||
def runReactor(self, app, reactor):
|
||||
"""
|
||||
Register the app, run the reactor, make sure app was activated, and
|
||||
that reactor was running, and that reactor can be stopped.
|
||||
"""
|
||||
if not hasattr(app, "quit"):
|
||||
raise SkipTest("Version of PyGObject is too old.")
|
||||
|
||||
result = []
|
||||
def stop():
|
||||
result.append("stopped")
|
||||
reactor.stop()
|
||||
def activate(widget):
|
||||
result.append("activated")
|
||||
reactor.callLater(0, stop)
|
||||
app.connect('activate', activate)
|
||||
|
||||
# We want reactor.stop() to *always* stop the event loop, even if
|
||||
# someone has called hold() on the application and never done the
|
||||
# corresponding release() -- for more details see
|
||||
# http://developer.gnome.org/gio/unstable/GApplication.html.
|
||||
app.hold()
|
||||
|
||||
reactor.registerGApplication(app)
|
||||
ReactorBuilder.runReactor(self, reactor)
|
||||
self.assertEqual(result, ["activated", "stopped"])
|
||||
|
||||
|
||||
def test_gApplicationActivate(self):
|
||||
"""
|
||||
L{Gio.Application} instances can be registered with a gireactor.
|
||||
"""
|
||||
reactor = gireactor.GIReactor(useGtk=False)
|
||||
self.addCleanup(self.unbuildReactor, reactor)
|
||||
app = Gio.Application(
|
||||
application_id='com.twistedmatrix.trial.gireactor',
|
||||
flags=Gio.ApplicationFlags.FLAGS_NONE)
|
||||
|
||||
self.runReactor(app, reactor)
|
||||
|
||||
|
||||
def test_gtkApplicationActivate(self):
|
||||
"""
|
||||
L{Gtk.Application} instances can be registered with a gtk3reactor.
|
||||
"""
|
||||
reactor = gtk3reactor.Gtk3Reactor()
|
||||
self.addCleanup(self.unbuildReactor, reactor)
|
||||
app = Gtk.Application(
|
||||
application_id='com.twistedmatrix.trial.gtk3reactor',
|
||||
flags=Gio.ApplicationFlags.FLAGS_NONE)
|
||||
|
||||
self.runReactor(app, reactor)
|
||||
|
||||
if gtk3reactor is None:
|
||||
test_gtkApplicationActivate.skip = (
|
||||
"Gtk unavailable (may require running with X11 DISPLAY env set)")
|
||||
|
||||
|
||||
def test_portable(self):
|
||||
"""
|
||||
L{gireactor.PortableGIReactor} doesn't support application
|
||||
registration at this time.
|
||||
"""
|
||||
reactor = gireactor.PortableGIReactor()
|
||||
self.addCleanup(self.unbuildReactor, reactor)
|
||||
app = Gio.Application(
|
||||
application_id='com.twistedmatrix.trial.gireactor',
|
||||
flags=Gio.ApplicationFlags.FLAGS_NONE)
|
||||
self.assertRaises(NotImplementedError,
|
||||
reactor.registerGApplication, app)
|
||||
|
||||
|
||||
def test_noQuit(self):
|
||||
"""
|
||||
Older versions of PyGObject lack C{Application.quit}, and so won't
|
||||
allow registration.
|
||||
"""
|
||||
reactor = gireactor.GIReactor(useGtk=False)
|
||||
self.addCleanup(self.unbuildReactor, reactor)
|
||||
# An app with no "quit" method:
|
||||
app = object()
|
||||
exc = self.assertRaises(RuntimeError, reactor.registerGApplication, app)
|
||||
self.assertTrue(exc.args[0].startswith(
|
||||
"Application registration is not"))
|
||||
|
||||
|
||||
def test_cantRegisterAfterRun(self):
|
||||
"""
|
||||
It is not possible to register a C{Application} after the reactor has
|
||||
already started.
|
||||
"""
|
||||
reactor = gireactor.GIReactor(useGtk=False)
|
||||
self.addCleanup(self.unbuildReactor, reactor)
|
||||
app = Gio.Application(
|
||||
application_id='com.twistedmatrix.trial.gireactor',
|
||||
flags=Gio.ApplicationFlags.FLAGS_NONE)
|
||||
|
||||
def tryRegister():
|
||||
exc = self.assertRaises(ReactorAlreadyRunning,
|
||||
reactor.registerGApplication, app)
|
||||
self.assertEqual(exc.args[0],
|
||||
"Can't register application after reactor was started.")
|
||||
reactor.stop()
|
||||
reactor.callLater(0, tryRegister)
|
||||
ReactorBuilder.runReactor(self, reactor)
|
||||
|
||||
|
||||
def test_cantRegisterTwice(self):
|
||||
"""
|
||||
It is not possible to register more than one C{Application}.
|
||||
"""
|
||||
reactor = gireactor.GIReactor(useGtk=False)
|
||||
self.addCleanup(self.unbuildReactor, reactor)
|
||||
app = Gio.Application(
|
||||
application_id='com.twistedmatrix.trial.gireactor',
|
||||
flags=Gio.ApplicationFlags.FLAGS_NONE)
|
||||
reactor.registerGApplication(app)
|
||||
app2 = Gio.Application(
|
||||
application_id='com.twistedmatrix.trial.gireactor2',
|
||||
flags=Gio.ApplicationFlags.FLAGS_NONE)
|
||||
exc = self.assertRaises(RuntimeError,
|
||||
reactor.registerGApplication, app2)
|
||||
self.assertEqual(exc.args[0],
|
||||
"Can't register more than one application instance.")
|
||||
|
||||
|
||||
|
||||
class PygtkCompatibilityTests(TestCase):
|
||||
"""
|
||||
pygtk imports are either prevented, or a compatibility layer is used if
|
||||
possible.
|
||||
"""
|
||||
def test_noCompatibilityLayer(self):
|
||||
"""
|
||||
If no compatibility layer is present, imports of gobject and friends
|
||||
are disallowed.
|
||||
|
||||
We do this by running a process where we make sure gi.pygtkcompat
|
||||
isn't present.
|
||||
"""
|
||||
if _PY3:
|
||||
raise SkipTest("Python3 always has the compatibility layer.")
|
||||
|
||||
from twisted.internet import reactor
|
||||
if not IReactorProcess.providedBy(reactor):
|
||||
raise SkipTest("No process support available in this reactor.")
|
||||
|
||||
result = Deferred()
|
||||
class Stdout(ProcessProtocol):
|
||||
data = b""
|
||||
|
||||
def errReceived(self, err):
|
||||
print(err)
|
||||
|
||||
def outReceived(self, data):
|
||||
self.data += data
|
||||
|
||||
def processExited(self, reason):
|
||||
result.callback(self.data)
|
||||
|
||||
path = FilePath(__file__).sibling(b"process_gireactornocompat.py").path
|
||||
pyExe = FilePath(sys.executable)._asBytesPath()
|
||||
# Pass in a PYTHONPATH that is the test runner's os.path, to make sure
|
||||
# we're running from a checkout
|
||||
reactor.spawnProcess(Stdout(), pyExe, [pyExe, path],
|
||||
env={"PYTHONPATH": ":".join(sys.path)})
|
||||
result.addCallback(self.assertEqual, b"success")
|
||||
return result
|
||||
|
||||
|
||||
def test_compatibilityLayer(self):
|
||||
"""
|
||||
If compatibility layer is present, importing gobject uses the gi
|
||||
compatibility layer.
|
||||
"""
|
||||
if "gi.pygtkcompat" not in sys.modules:
|
||||
raise SkipTest("This version of gi doesn't include pygtkcompat.")
|
||||
import gobject
|
||||
self.assertTrue(gobject.__name__.startswith("gi."))
|
||||
|
||||
|
||||
|
||||
class Gtk3ReactorTests(TestCase):
|
||||
"""
|
||||
Tests for L{gtk3reactor}.
|
||||
"""
|
||||
|
||||
def test_requiresDISPLAY(self):
|
||||
"""
|
||||
On X11, L{gtk3reactor} is unimportable if the C{DISPLAY} environment
|
||||
variable is not set.
|
||||
"""
|
||||
display = os.environ.get("DISPLAY", None)
|
||||
if display is not None:
|
||||
self.addCleanup(os.environ.__setitem__, "DISPLAY", display)
|
||||
del os.environ["DISPLAY"]
|
||||
with SetAsideModule("twisted.internet.gtk3reactor"):
|
||||
exc = self.assertRaises(ImportError,
|
||||
__import__, "twisted.internet.gtk3reactor")
|
||||
self.assertEqual(
|
||||
exc.args[0],
|
||||
"Gtk3 requires X11, and no DISPLAY environment variable is set")
|
||||
|
||||
if platform.getType() != "posix" or platform.isMacOSX():
|
||||
test_requiresDISPLAY.skip = "This test is only relevant when using X11"
|
||||
@@ -0,0 +1,68 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for twisted.internet.glibbase.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
import sys
|
||||
from twisted.trial.unittest import TestCase
|
||||
from twisted.internet._glibbase import ensureNotImported
|
||||
|
||||
|
||||
|
||||
class EnsureNotImportedTests(TestCase):
|
||||
"""
|
||||
L{ensureNotImported} protects against unwanted past and future imports.
|
||||
"""
|
||||
|
||||
def test_ensureWhenNotImported(self):
|
||||
"""
|
||||
If the specified modules have never been imported, and import
|
||||
prevention is requested, L{ensureNotImported} makes sure they will not
|
||||
be imported in the future.
|
||||
"""
|
||||
modules = {}
|
||||
self.patch(sys, "modules", modules)
|
||||
ensureNotImported(["m1", "m2"], "A message.",
|
||||
preventImports=["m1", "m2", "m3"])
|
||||
self.assertEqual(modules, {"m1": None, "m2": None, "m3": None})
|
||||
|
||||
|
||||
def test_ensureWhenNotImportedDontPrevent(self):
|
||||
"""
|
||||
If the specified modules have never been imported, and import
|
||||
prevention is not requested, L{ensureNotImported} has no effect.
|
||||
"""
|
||||
modules = {}
|
||||
self.patch(sys, "modules", modules)
|
||||
ensureNotImported(["m1", "m2"], "A message.")
|
||||
self.assertEqual(modules, {})
|
||||
|
||||
|
||||
def test_ensureWhenFailedToImport(self):
|
||||
"""
|
||||
If the specified modules have been set to L{None} in C{sys.modules},
|
||||
L{ensureNotImported} does not complain.
|
||||
"""
|
||||
modules = {"m2": None}
|
||||
self.patch(sys, "modules", modules)
|
||||
ensureNotImported(["m1", "m2"], "A message.", preventImports=["m1", "m2"])
|
||||
self.assertEqual(modules, {"m1": None, "m2": None})
|
||||
|
||||
|
||||
def test_ensureFailsWhenImported(self):
|
||||
"""
|
||||
If one of the specified modules has been previously imported,
|
||||
L{ensureNotImported} raises an exception.
|
||||
"""
|
||||
module = object()
|
||||
modules = {"m2": module}
|
||||
self.patch(sys, "modules", modules)
|
||||
e = self.assertRaises(ImportError, ensureNotImported,
|
||||
["m1", "m2"], "A message.",
|
||||
preventImports=["m1", "m2"])
|
||||
self.assertEqual(modules, {"m2": module})
|
||||
self.assertEqual(e.args, ("A message.",))
|
||||
@@ -0,0 +1,384 @@
|
||||
# -*- test-case-name: twisted.internet.test.test_inlinecb -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.internet.defer.inlineCallbacks}.
|
||||
|
||||
Some tests for inlineCallbacks are defined in L{twisted.test.test_defgen} as
|
||||
well.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
import sys
|
||||
|
||||
from twisted.trial.unittest import TestCase, SynchronousTestCase
|
||||
from twisted.internet.defer import (
|
||||
Deferred, returnValue, inlineCallbacks, CancelledError)
|
||||
|
||||
|
||||
class StopIterationReturnTests(TestCase):
|
||||
"""
|
||||
On Python 3.4 and newer generator functions may use the C{return} statement
|
||||
with a value, which is attached to the L{StopIteration} exception that is
|
||||
raised.
|
||||
|
||||
L{inlineCallbacks} will use this value when it fires the C{callback}.
|
||||
"""
|
||||
|
||||
def test_returnWithValue(self):
|
||||
"""
|
||||
If the C{return} statement has a value it is propagated back to the
|
||||
L{Deferred} that the C{inlineCallbacks} function returned.
|
||||
"""
|
||||
environ = {"inlineCallbacks": inlineCallbacks}
|
||||
exec("""
|
||||
@inlineCallbacks
|
||||
def f(d):
|
||||
yield d
|
||||
return 14
|
||||
""", environ)
|
||||
d1 = Deferred()
|
||||
d2 = environ["f"](d1)
|
||||
d1.callback(None)
|
||||
self.assertEqual(self.successResultOf(d2), 14)
|
||||
|
||||
|
||||
|
||||
if sys.version_info < (3, 4):
|
||||
StopIterationReturnTests.skip = "Test requires Python 3.4 or greater"
|
||||
|
||||
|
||||
|
||||
class NonLocalExitTests(TestCase):
|
||||
"""
|
||||
It's possible for L{returnValue} to be (accidentally) invoked at a stack
|
||||
level below the L{inlineCallbacks}-decorated function which it is exiting.
|
||||
If this happens, L{returnValue} should report useful errors.
|
||||
|
||||
If L{returnValue} is invoked from a function not decorated by
|
||||
L{inlineCallbacks}, it will emit a warning if it causes an
|
||||
L{inlineCallbacks} function further up the stack to exit.
|
||||
"""
|
||||
|
||||
def mistakenMethod(self):
|
||||
"""
|
||||
This method mistakenly invokes L{returnValue}, despite the fact that it
|
||||
is not decorated with L{inlineCallbacks}.
|
||||
"""
|
||||
returnValue(1)
|
||||
|
||||
|
||||
def assertMistakenMethodWarning(self, resultList):
|
||||
"""
|
||||
Flush the current warnings and assert that we have been told that
|
||||
C{mistakenMethod} was invoked, and that the result from the Deferred
|
||||
that was fired (appended to the given list) is C{mistakenMethod}'s
|
||||
result. The warning should indicate that an inlineCallbacks function
|
||||
called 'inline' was made to exit.
|
||||
"""
|
||||
self.assertEqual(resultList, [1])
|
||||
warnings = self.flushWarnings(offendingFunctions=[self.mistakenMethod])
|
||||
self.assertEqual(len(warnings), 1)
|
||||
self.assertEqual(warnings[0]['category'], DeprecationWarning)
|
||||
self.assertEqual(
|
||||
warnings[0]['message'],
|
||||
"returnValue() in 'mistakenMethod' causing 'inline' to exit: "
|
||||
"returnValue should only be invoked by functions decorated with "
|
||||
"inlineCallbacks")
|
||||
|
||||
|
||||
def test_returnValueNonLocalWarning(self):
|
||||
"""
|
||||
L{returnValue} will emit a non-local exit warning in the simplest case,
|
||||
where the offending function is invoked immediately.
|
||||
"""
|
||||
@inlineCallbacks
|
||||
def inline():
|
||||
self.mistakenMethod()
|
||||
returnValue(2)
|
||||
yield 0
|
||||
d = inline()
|
||||
results = []
|
||||
d.addCallback(results.append)
|
||||
self.assertMistakenMethodWarning(results)
|
||||
|
||||
|
||||
def test_returnValueNonLocalDeferred(self):
|
||||
"""
|
||||
L{returnValue} will emit a non-local warning in the case where the
|
||||
L{inlineCallbacks}-decorated function has already yielded a Deferred
|
||||
and therefore moved its generator function along.
|
||||
"""
|
||||
cause = Deferred()
|
||||
@inlineCallbacks
|
||||
def inline():
|
||||
yield cause
|
||||
self.mistakenMethod()
|
||||
returnValue(2)
|
||||
effect = inline()
|
||||
results = []
|
||||
effect.addCallback(results.append)
|
||||
self.assertEqual(results, [])
|
||||
cause.callback(1)
|
||||
self.assertMistakenMethodWarning(results)
|
||||
|
||||
|
||||
|
||||
class ForwardTraceBackTests(SynchronousTestCase):
|
||||
|
||||
def test_forwardTracebacks(self):
|
||||
"""
|
||||
Chained inlineCallbacks are forwarding the traceback information
|
||||
from generator to generator.
|
||||
|
||||
A first simple test with a couple of inline callbacks.
|
||||
"""
|
||||
|
||||
@inlineCallbacks
|
||||
def erroring():
|
||||
yield "forcing generator"
|
||||
raise Exception('Error Marker')
|
||||
|
||||
@inlineCallbacks
|
||||
def calling():
|
||||
yield erroring()
|
||||
|
||||
d = calling()
|
||||
f = self.failureResultOf(d)
|
||||
tb = f.getTraceback()
|
||||
self.assertIn("in erroring", tb)
|
||||
self.assertIn("in calling", tb)
|
||||
self.assertIn("Error Marker", tb)
|
||||
|
||||
|
||||
def test_forwardLotsOfTracebacks(self):
|
||||
"""
|
||||
Several Chained inlineCallbacks gives information about all generators.
|
||||
|
||||
A wider test with a 4 chained inline callbacks.
|
||||
|
||||
Application stack-trace should be reported, and implementation details
|
||||
like "throwExceptionIntoGenerator" symbols are omitted from the stack.
|
||||
|
||||
Note that the previous test is testing the simple case, and this one is
|
||||
testing the deep recursion case.
|
||||
|
||||
That case needs specific code in failure.py to accomodate to stack
|
||||
breakage introduced by throwExceptionIntoGenerator.
|
||||
|
||||
Hence we keep the two tests in order to sort out which code we
|
||||
might have regression in.
|
||||
"""
|
||||
|
||||
@inlineCallbacks
|
||||
def erroring():
|
||||
yield "forcing generator"
|
||||
raise Exception('Error Marker')
|
||||
|
||||
@inlineCallbacks
|
||||
def calling3():
|
||||
yield erroring()
|
||||
|
||||
@inlineCallbacks
|
||||
def calling2():
|
||||
yield calling3()
|
||||
|
||||
@inlineCallbacks
|
||||
def calling():
|
||||
yield calling2()
|
||||
|
||||
d = calling()
|
||||
f = self.failureResultOf(d)
|
||||
tb = f.getTraceback()
|
||||
self.assertIn("in erroring", tb)
|
||||
self.assertIn("in calling", tb)
|
||||
self.assertIn("in calling2", tb)
|
||||
self.assertIn("in calling3", tb)
|
||||
self.assertNotIn("throwExceptionIntoGenerator", tb)
|
||||
self.assertIn("Error Marker", tb)
|
||||
self.assertIn("in erroring", f.getTraceback())
|
||||
|
||||
|
||||
|
||||
class UntranslatedError(Exception):
|
||||
"""
|
||||
Untranslated exception type when testing an exception translation.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class TranslatedError(Exception):
|
||||
"""
|
||||
Translated exception type when testing an exception translation.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
class DontFail(Exception):
|
||||
"""
|
||||
Sample exception type.
|
||||
"""
|
||||
|
||||
def __init__(self, actual):
|
||||
Exception.__init__(self)
|
||||
self.actualValue = actual
|
||||
|
||||
|
||||
|
||||
class CancellationTests(SynchronousTestCase):
|
||||
"""
|
||||
Tests for cancellation of L{Deferred}s returned by L{inlineCallbacks}.
|
||||
For each of these tests, let:
|
||||
- C{G} be a generator decorated with C{inlineCallbacks}
|
||||
- C{D} be a L{Deferred} returned by C{G}
|
||||
- C{C} be a L{Deferred} awaited by C{G} with C{yield}
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
"""
|
||||
Set up the list of outstanding L{Deferred}s.
|
||||
"""
|
||||
self.deferredsOutstanding = []
|
||||
|
||||
|
||||
def tearDown(self):
|
||||
"""
|
||||
If any L{Deferred}s are still outstanding, fire them.
|
||||
"""
|
||||
while self.deferredsOutstanding:
|
||||
self.deferredGotten()
|
||||
|
||||
|
||||
@inlineCallbacks
|
||||
def sampleInlineCB(self, getChildDeferred=None):
|
||||
"""
|
||||
Generator for testing cascade cancelling cases.
|
||||
|
||||
@param getChildDeferred: Some callable returning L{Deferred} that we
|
||||
awaiting (with C{yield})
|
||||
"""
|
||||
if getChildDeferred is None:
|
||||
getChildDeferred = self.getDeferred
|
||||
try:
|
||||
x = yield getChildDeferred()
|
||||
except UntranslatedError:
|
||||
raise TranslatedError()
|
||||
except DontFail as df:
|
||||
x = df.actualValue - 2
|
||||
returnValue(x + 1)
|
||||
|
||||
|
||||
def getDeferred(self):
|
||||
"""
|
||||
A sample function that returns a L{Deferred} that can be fired on
|
||||
demand, by L{CancellationTests.deferredGotten}.
|
||||
|
||||
@return: L{Deferred} that can be fired on demand.
|
||||
"""
|
||||
self.deferredsOutstanding.append(Deferred())
|
||||
return self.deferredsOutstanding[-1]
|
||||
|
||||
|
||||
def deferredGotten(self, result=None):
|
||||
"""
|
||||
Fire the L{Deferred} returned from the least-recent call to
|
||||
L{CancellationTests.getDeferred}.
|
||||
|
||||
@param result: result object to be used when firing the L{Deferred}.
|
||||
"""
|
||||
self.deferredsOutstanding.pop(0).callback(result)
|
||||
|
||||
|
||||
def test_cascadeCancellingOnCancel(self):
|
||||
"""
|
||||
When C{D} cancelled, C{C} will be immediately cancelled too.
|
||||
"""
|
||||
childResultHolder = ['FAILURE']
|
||||
def getChildDeferred():
|
||||
d = Deferred()
|
||||
def _eb(result):
|
||||
childResultHolder[0] = result.check(CancelledError)
|
||||
return result
|
||||
d.addErrback(_eb)
|
||||
return d
|
||||
d = self.sampleInlineCB(getChildDeferred=getChildDeferred)
|
||||
d.addErrback(lambda result: None)
|
||||
d.cancel()
|
||||
self.assertEqual(
|
||||
childResultHolder[0],
|
||||
CancelledError,
|
||||
"no cascade cancelling occurs",
|
||||
)
|
||||
|
||||
|
||||
def test_errbackCancelledErrorOnCancel(self):
|
||||
"""
|
||||
When C{D} cancelled, CancelledError from C{C} will be errbacked
|
||||
through C{D}.
|
||||
"""
|
||||
d = self.sampleInlineCB()
|
||||
d.cancel()
|
||||
self.assertRaises(
|
||||
CancelledError,
|
||||
self.failureResultOf(d).raiseException,
|
||||
)
|
||||
|
||||
|
||||
def test_errorToErrorTranslation(self):
|
||||
"""
|
||||
When C{D} is cancelled, and C raises a particular type of error, C{G}
|
||||
may catch that error at the point of yielding and translate it into
|
||||
a different error which may be received by application code.
|
||||
"""
|
||||
def cancel(it):
|
||||
it.errback(UntranslatedError())
|
||||
a = Deferred(cancel)
|
||||
d = self.sampleInlineCB(lambda: a)
|
||||
d.cancel()
|
||||
self.assertRaises(
|
||||
TranslatedError,
|
||||
self.failureResultOf(d).raiseException,
|
||||
)
|
||||
|
||||
|
||||
def test_errorToSuccessTranslation(self):
|
||||
"""
|
||||
When C{D} is cancelled, and C{C} raises a particular type of error,
|
||||
C{G} may catch that error at the point of yielding and translate it
|
||||
into a result value which may be received by application code.
|
||||
"""
|
||||
def cancel(it):
|
||||
it.errback(DontFail(4321))
|
||||
a = Deferred(cancel)
|
||||
d = self.sampleInlineCB(lambda: a)
|
||||
results = []
|
||||
d.addCallback(results.append)
|
||||
d.cancel()
|
||||
self.assertEquals(results, [4320])
|
||||
|
||||
|
||||
def test_asynchronousCancellation(self):
|
||||
"""
|
||||
When C{D} is cancelled, it won't reach the callbacks added to it by
|
||||
application code until C{C} reaches the point in its callback chain
|
||||
where C{G} awaits it. Otherwise, application code won't be able to
|
||||
track resource usage that C{D} may be using.
|
||||
"""
|
||||
moreDeferred = Deferred()
|
||||
|
||||
def deferMeMore(result):
|
||||
result.trap(CancelledError)
|
||||
return moreDeferred
|
||||
|
||||
def deferMe():
|
||||
d = Deferred()
|
||||
d.addErrback(deferMeMore)
|
||||
return d
|
||||
|
||||
d = self.sampleInlineCB(getChildDeferred=deferMe)
|
||||
d.cancel()
|
||||
self.assertNoResult(d)
|
||||
moreDeferred.callback(6543)
|
||||
self.assertEqual(self.successResultOf(d), 6544)
|
||||
@@ -0,0 +1,71 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.internet.kqueuereactor}.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
import errno
|
||||
|
||||
from zope.interface import implementer
|
||||
|
||||
from twisted.trial.unittest import TestCase
|
||||
|
||||
try:
|
||||
from twisted.internet.kqreactor import KQueueReactor, _IKQueue
|
||||
kqueueSkip = None
|
||||
except ImportError:
|
||||
kqueueSkip = "KQueue not available."
|
||||
|
||||
|
||||
def _fakeKEvent(*args, **kwargs):
|
||||
"""
|
||||
Do nothing.
|
||||
"""
|
||||
|
||||
|
||||
|
||||
def makeFakeKQueue(testKQueue, testKEvent):
|
||||
"""
|
||||
Create a fake that implements L{_IKQueue}.
|
||||
|
||||
@param testKQueue: Something that acts like L{select.kqueue}.
|
||||
@param testKEvent: Something that acts like L{select.kevent}.
|
||||
@return: An implementation of L{_IKQueue} that includes C{testKQueue} and
|
||||
C{testKEvent}.
|
||||
"""
|
||||
@implementer(_IKQueue)
|
||||
class FakeKQueue(object):
|
||||
kqueue = testKQueue
|
||||
kevent = testKEvent
|
||||
|
||||
return FakeKQueue()
|
||||
|
||||
|
||||
|
||||
class KQueueTests(TestCase):
|
||||
"""
|
||||
These are tests for L{KQueueReactor}'s implementation, not its real world
|
||||
behaviour. For that, look at
|
||||
L{twisted.internet.test.reactormixins.ReactorBuilder}.
|
||||
"""
|
||||
skip = kqueueSkip
|
||||
|
||||
def test_EINTR(self):
|
||||
"""
|
||||
L{KQueueReactor} handles L{errno.EINTR} in C{doKEvent} by returning.
|
||||
"""
|
||||
class FakeKQueue(object):
|
||||
"""
|
||||
A fake KQueue that raises L{errno.EINTR} when C{control} is called,
|
||||
like a real KQueue would if it was interrupted.
|
||||
"""
|
||||
def control(self, *args, **kwargs):
|
||||
raise OSError(errno.EINTR, "Interrupted")
|
||||
|
||||
reactor = KQueueReactor(makeFakeKQueue(FakeKQueue, _fakeKEvent))
|
||||
# This should return cleanly -- should not raise the OSError we're
|
||||
# spawning, nor get upset and raise about the incomplete KQueue fake.
|
||||
reactor.doKEvent(0)
|
||||
@@ -0,0 +1,908 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for implementations of L{IReactorProcess}.
|
||||
|
||||
@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 division, absolute_import, print_function
|
||||
|
||||
import io
|
||||
import os
|
||||
import signal
|
||||
import sys
|
||||
import threading
|
||||
import twisted
|
||||
import subprocess
|
||||
|
||||
from twisted.trial.unittest import TestCase
|
||||
from twisted.internet.test.reactormixins import ReactorBuilder
|
||||
from twisted.python.log import msg, err
|
||||
from twisted.python.runtime import platform
|
||||
from twisted.python.filepath import FilePath, _asFilesystemBytes
|
||||
from twisted.python.compat import (networkString, range, items,
|
||||
bytesEnviron, unicode)
|
||||
from twisted.internet import utils
|
||||
from twisted.internet.interfaces import IReactorProcess, IProcessTransport
|
||||
from twisted.internet.defer import Deferred, succeed
|
||||
from twisted.internet.protocol import ProcessProtocol
|
||||
from twisted.internet.error import ProcessDone, ProcessTerminated
|
||||
|
||||
|
||||
# Get the current Python executable as a bytestring.
|
||||
pyExe = FilePath(sys.executable)._asBytesPath()
|
||||
twistedRoot = FilePath(twisted.__file__).parent().parent()
|
||||
|
||||
_uidgidSkip = None
|
||||
if platform.isWindows():
|
||||
resource = None
|
||||
process = None
|
||||
_uidgidSkip = "Cannot change UID/GID on Windows"
|
||||
|
||||
properEnv = dict(os.environ)
|
||||
properEnv["PYTHONPATH"] = os.pathsep.join(sys.path)
|
||||
else:
|
||||
import resource
|
||||
from twisted.internet import process
|
||||
if os.getuid() != 0:
|
||||
_uidgidSkip = "Cannot change UID/GID except as root"
|
||||
|
||||
properEnv = bytesEnviron()
|
||||
properEnv[b"PYTHONPATH"] = os.pathsep.join(sys.path).encode(
|
||||
sys.getfilesystemencoding())
|
||||
|
||||
|
||||
|
||||
def onlyOnPOSIX(testMethod):
|
||||
"""
|
||||
Only run this test on POSIX platforms.
|
||||
|
||||
@param testMethod: A test function, being decorated.
|
||||
|
||||
@return: the C{testMethod} argument.
|
||||
"""
|
||||
if resource is None:
|
||||
testMethod.skip = "Test only applies to POSIX platforms."
|
||||
return testMethod
|
||||
|
||||
|
||||
|
||||
class _ShutdownCallbackProcessProtocol(ProcessProtocol):
|
||||
"""
|
||||
An L{IProcessProtocol} which fires a Deferred when the process it is
|
||||
associated with ends.
|
||||
|
||||
@ivar received: A C{dict} mapping file descriptors to lists of bytes
|
||||
received from the child process on those file descriptors.
|
||||
"""
|
||||
def __init__(self, whenFinished):
|
||||
self.whenFinished = whenFinished
|
||||
self.received = {}
|
||||
|
||||
|
||||
def childDataReceived(self, fd, bytes):
|
||||
self.received.setdefault(fd, []).append(bytes)
|
||||
|
||||
|
||||
def processEnded(self, reason):
|
||||
self.whenFinished.callback(None)
|
||||
|
||||
|
||||
|
||||
class ProcessTestsBuilderBase(ReactorBuilder):
|
||||
"""
|
||||
Base class for L{IReactorProcess} tests which defines some tests which
|
||||
can be applied to PTY or non-PTY uses of C{spawnProcess}.
|
||||
|
||||
Subclasses are expected to set the C{usePTY} attribute to C{True} or
|
||||
C{False}.
|
||||
"""
|
||||
requiredInterfaces = [IReactorProcess]
|
||||
|
||||
|
||||
def test_processTransportInterface(self):
|
||||
"""
|
||||
L{IReactorProcess.spawnProcess} connects the protocol passed to it
|
||||
to a transport which provides L{IProcessTransport}.
|
||||
"""
|
||||
ended = Deferred()
|
||||
protocol = _ShutdownCallbackProcessProtocol(ended)
|
||||
|
||||
reactor = self.buildReactor()
|
||||
transport = reactor.spawnProcess(
|
||||
protocol, pyExe, [pyExe, b"-c", b""],
|
||||
usePTY=self.usePTY)
|
||||
|
||||
# The transport is available synchronously, so we can check it right
|
||||
# away (unlike many transport-based tests). This is convenient even
|
||||
# though it's probably not how the spawnProcess interface should really
|
||||
# work.
|
||||
# We're not using verifyObject here because part of
|
||||
# IProcessTransport is a lie - there are no getHost or getPeer
|
||||
# methods. See #1124.
|
||||
self.assertTrue(IProcessTransport.providedBy(transport))
|
||||
|
||||
# Let the process run and exit so we don't leave a zombie around.
|
||||
ended.addCallback(lambda ignored: reactor.stop())
|
||||
self.runReactor(reactor)
|
||||
|
||||
|
||||
def _writeTest(self, write):
|
||||
"""
|
||||
Helper for testing L{IProcessTransport} write functionality. This
|
||||
method spawns a child process and gives C{write} a chance to write some
|
||||
bytes to it. It then verifies that the bytes were actually written to
|
||||
it (by relying on the child process to echo them back).
|
||||
|
||||
@param write: A two-argument callable. This is invoked with a process
|
||||
transport and some bytes to write to it.
|
||||
"""
|
||||
reactor = self.buildReactor()
|
||||
|
||||
ended = Deferred()
|
||||
protocol = _ShutdownCallbackProcessProtocol(ended)
|
||||
|
||||
bytesToSend = b"hello, world" + networkString(os.linesep)
|
||||
program = (
|
||||
b"import sys\n"
|
||||
b"sys.stdout.write(sys.stdin.readline())\n"
|
||||
)
|
||||
|
||||
def startup():
|
||||
transport = reactor.spawnProcess(
|
||||
protocol, pyExe, [pyExe, b"-c", program])
|
||||
try:
|
||||
write(transport, bytesToSend)
|
||||
except:
|
||||
err(None, "Unhandled exception while writing")
|
||||
transport.signalProcess('KILL')
|
||||
reactor.callWhenRunning(startup)
|
||||
|
||||
ended.addCallback(lambda ignored: reactor.stop())
|
||||
|
||||
self.runReactor(reactor)
|
||||
self.assertEqual(bytesToSend, b"".join(protocol.received[1]))
|
||||
|
||||
|
||||
def test_write(self):
|
||||
"""
|
||||
L{IProcessTransport.write} writes the specified C{bytes} to the standard
|
||||
input of the child process.
|
||||
"""
|
||||
def write(transport, bytesToSend):
|
||||
transport.write(bytesToSend)
|
||||
self._writeTest(write)
|
||||
|
||||
|
||||
def test_writeSequence(self):
|
||||
"""
|
||||
L{IProcessTransport.writeSequence} writes the specified C{list} of
|
||||
C{bytes} to the standard input of the child process.
|
||||
"""
|
||||
def write(transport, bytesToSend):
|
||||
transport.writeSequence([bytesToSend])
|
||||
self._writeTest(write)
|
||||
|
||||
|
||||
def test_writeToChild(self):
|
||||
"""
|
||||
L{IProcessTransport.writeToChild} writes the specified C{bytes} to the
|
||||
specified file descriptor of the child process.
|
||||
"""
|
||||
def write(transport, bytesToSend):
|
||||
transport.writeToChild(0, bytesToSend)
|
||||
self._writeTest(write)
|
||||
|
||||
|
||||
def test_writeToChildBadFileDescriptor(self):
|
||||
"""
|
||||
L{IProcessTransport.writeToChild} raises L{KeyError} if passed a file
|
||||
descriptor which is was not set up by L{IReactorProcess.spawnProcess}.
|
||||
"""
|
||||
def write(transport, bytesToSend):
|
||||
try:
|
||||
self.assertRaises(KeyError, transport.writeToChild, 13, bytesToSend)
|
||||
finally:
|
||||
# Just get the process to exit so the test can complete
|
||||
transport.write(bytesToSend)
|
||||
self._writeTest(write)
|
||||
|
||||
|
||||
def test_spawnProcessEarlyIsReaped(self):
|
||||
"""
|
||||
If, before the reactor is started with L{IReactorCore.run}, a
|
||||
process is started with L{IReactorProcess.spawnProcess} and
|
||||
terminates, the process is reaped once the reactor is started.
|
||||
"""
|
||||
reactor = self.buildReactor()
|
||||
|
||||
# Create the process with no shared file descriptors, so that there
|
||||
# are no other events for the reactor to notice and "cheat" with.
|
||||
# We want to be sure it's really dealing with the process exiting,
|
||||
# not some associated event.
|
||||
if self.usePTY:
|
||||
childFDs = None
|
||||
else:
|
||||
childFDs = {}
|
||||
|
||||
# Arrange to notice the SIGCHLD.
|
||||
signaled = threading.Event()
|
||||
def handler(*args):
|
||||
signaled.set()
|
||||
signal.signal(signal.SIGCHLD, handler)
|
||||
|
||||
# Start a process - before starting the reactor!
|
||||
ended = Deferred()
|
||||
reactor.spawnProcess(
|
||||
_ShutdownCallbackProcessProtocol(ended), pyExe,
|
||||
[pyExe, b"-c", b""], usePTY=self.usePTY, childFDs=childFDs)
|
||||
|
||||
# Wait for the SIGCHLD (which might have been delivered before we got
|
||||
# here, but that's okay because the signal handler was installed above,
|
||||
# before we could have gotten it).
|
||||
signaled.wait(120)
|
||||
if not signaled.isSet():
|
||||
self.fail("Timed out waiting for child process to exit.")
|
||||
|
||||
# Capture the processEnded callback.
|
||||
result = []
|
||||
ended.addCallback(result.append)
|
||||
|
||||
if result:
|
||||
# The synchronous path through spawnProcess / Process.__init__ /
|
||||
# registerReapProcessHandler was encountered. There's no reason to
|
||||
# start the reactor, because everything is done already.
|
||||
return
|
||||
|
||||
# Otherwise, though, start the reactor so it can tell us the process
|
||||
# exited.
|
||||
ended.addCallback(lambda ignored: reactor.stop())
|
||||
self.runReactor(reactor)
|
||||
|
||||
# Make sure the reactor stopped because the Deferred fired.
|
||||
self.assertTrue(result)
|
||||
|
||||
if getattr(signal, 'SIGCHLD', None) is None:
|
||||
test_spawnProcessEarlyIsReaped.skip = (
|
||||
"Platform lacks SIGCHLD, early-spawnProcess test can't work.")
|
||||
|
||||
|
||||
def test_processExitedWithSignal(self):
|
||||
"""
|
||||
The C{reason} argument passed to L{IProcessProtocol.processExited} is a
|
||||
L{ProcessTerminated} instance if the child process exits with a signal.
|
||||
"""
|
||||
sigName = 'TERM'
|
||||
sigNum = getattr(signal, 'SIG' + sigName)
|
||||
exited = Deferred()
|
||||
source = (
|
||||
b"import sys\n"
|
||||
# Talk so the parent process knows the process is running. This is
|
||||
# necessary because ProcessProtocol.makeConnection may be called
|
||||
# before this process is exec'd. It would be unfortunate if we
|
||||
# SIGTERM'd the Twisted process while it was on its way to doing
|
||||
# the exec.
|
||||
b"sys.stdout.write('x')\n"
|
||||
b"sys.stdout.flush()\n"
|
||||
b"sys.stdin.read()\n")
|
||||
|
||||
class Exiter(ProcessProtocol):
|
||||
def childDataReceived(self, fd, data):
|
||||
msg('childDataReceived(%d, %r)' % (fd, data))
|
||||
self.transport.signalProcess(sigName)
|
||||
|
||||
def childConnectionLost(self, fd):
|
||||
msg('childConnectionLost(%d)' % (fd,))
|
||||
|
||||
def processExited(self, reason):
|
||||
msg('processExited(%r)' % (reason,))
|
||||
# Protect the Deferred from the failure so that it follows
|
||||
# the callback chain. This doesn't use the errback chain
|
||||
# because it wants to make sure reason is a Failure. An
|
||||
# Exception would also make an errback-based test pass, and
|
||||
# that would be wrong.
|
||||
exited.callback([reason])
|
||||
|
||||
def processEnded(self, reason):
|
||||
msg('processEnded(%r)' % (reason,))
|
||||
|
||||
reactor = self.buildReactor()
|
||||
reactor.callWhenRunning(
|
||||
reactor.spawnProcess, Exiter(), pyExe,
|
||||
[pyExe, b"-c", source], usePTY=self.usePTY)
|
||||
|
||||
def cbExited(args):
|
||||
failure, = args
|
||||
# Trapping implicitly verifies that it's a Failure (rather than
|
||||
# an exception) and explicitly makes sure it's the right type.
|
||||
failure.trap(ProcessTerminated)
|
||||
err = failure.value
|
||||
if platform.isWindows():
|
||||
# Windows can't really /have/ signals, so it certainly can't
|
||||
# report them as the reason for termination. Maybe there's
|
||||
# something better we could be doing here, anyway? Hard to
|
||||
# say. Anyway, this inconsistency between different platforms
|
||||
# is extremely unfortunate and I would remove it if I
|
||||
# could. -exarkun
|
||||
self.assertIsNone(err.signal)
|
||||
self.assertEqual(err.exitCode, 1)
|
||||
else:
|
||||
self.assertEqual(err.signal, sigNum)
|
||||
self.assertIsNone(err.exitCode)
|
||||
|
||||
exited.addCallback(cbExited)
|
||||
exited.addErrback(err)
|
||||
exited.addCallback(lambda ign: reactor.stop())
|
||||
|
||||
self.runReactor(reactor)
|
||||
|
||||
|
||||
def test_systemCallUninterruptedByChildExit(self):
|
||||
"""
|
||||
If a child process exits while a system call is in progress, the system
|
||||
call should not be interfered with. In particular, it should not fail
|
||||
with EINTR.
|
||||
|
||||
Older versions of Twisted installed a SIGCHLD handler on POSIX without
|
||||
using the feature exposed by the SA_RESTART flag to sigaction(2). The
|
||||
most noticeable problem this caused was for blocking reads and writes to
|
||||
sometimes fail with EINTR.
|
||||
"""
|
||||
reactor = self.buildReactor()
|
||||
result = []
|
||||
|
||||
def f():
|
||||
try:
|
||||
exe = pyExe.decode(sys.getfilesystemencoding())
|
||||
|
||||
subprocess.Popen([exe, "-c", "import time; time.sleep(0.1)"])
|
||||
f2 = subprocess.Popen([exe, "-c",
|
||||
("import time; time.sleep(0.5);"
|
||||
"print(\'Foo\')")],
|
||||
stdout=subprocess.PIPE)
|
||||
# The read call below will blow up with an EINTR from the
|
||||
# SIGCHLD from the first process exiting if we install a
|
||||
# SIGCHLD handler without SA_RESTART. (which we used to do)
|
||||
with f2.stdout:
|
||||
result.append(f2.stdout.read())
|
||||
finally:
|
||||
reactor.stop()
|
||||
|
||||
reactor.callWhenRunning(f)
|
||||
self.runReactor(reactor)
|
||||
self.assertEqual(result, [b"Foo" + os.linesep.encode('ascii')])
|
||||
|
||||
|
||||
@onlyOnPOSIX
|
||||
def test_openFileDescriptors(self):
|
||||
"""
|
||||
Processes spawned with spawnProcess() close all extraneous file
|
||||
descriptors in the parent. They do have a stdin, stdout, and stderr
|
||||
open.
|
||||
"""
|
||||
|
||||
# To test this, we are going to open a file descriptor in the parent
|
||||
# that is unlikely to be opened in the child, then verify that it's not
|
||||
# open in the child.
|
||||
source = networkString("""
|
||||
import sys
|
||||
sys.path.insert(0, '{0}')
|
||||
from twisted.internet import process
|
||||
sys.stdout.write(repr(process._listOpenFDs()))
|
||||
sys.stdout.flush()""".format(twistedRoot.path))
|
||||
|
||||
r, w = os.pipe()
|
||||
self.addCleanup(os.close, r)
|
||||
self.addCleanup(os.close, w)
|
||||
|
||||
# The call to "os.listdir()" (in _listOpenFDs's implementation) opens a
|
||||
# file descriptor (with "opendir"), which shows up in _listOpenFDs's
|
||||
# result. And speaking of "random" file descriptors, the code required
|
||||
# for _listOpenFDs itself imports logger, which imports random, which
|
||||
# (depending on your Python version) might leave /dev/urandom open.
|
||||
|
||||
# More generally though, even if we were to use an extremely minimal C
|
||||
# program, the operating system would be within its rights to open file
|
||||
# descriptors we might not know about in the C library's
|
||||
# initialization; things like debuggers, profilers, or nsswitch plugins
|
||||
# might open some and this test should pass in those environments.
|
||||
|
||||
# Although some of these file descriptors aren't predictable, we should
|
||||
# at least be able to select a very large file descriptor which is very
|
||||
# unlikely to be opened automatically in the subprocess. (Apply a
|
||||
# fudge factor to avoid hard-coding something too near a limit
|
||||
# condition like the maximum possible file descriptor, which a library
|
||||
# might at least hypothetically select.)
|
||||
|
||||
fudgeFactor = 17
|
||||
unlikelyFD = (resource.getrlimit(resource.RLIMIT_NOFILE)[0]
|
||||
- fudgeFactor)
|
||||
|
||||
os.dup2(w, unlikelyFD)
|
||||
self.addCleanup(os.close, unlikelyFD)
|
||||
|
||||
output = io.BytesIO()
|
||||
class GatheringProtocol(ProcessProtocol):
|
||||
outReceived = output.write
|
||||
def processEnded(self, reason):
|
||||
reactor.stop()
|
||||
|
||||
reactor = self.buildReactor()
|
||||
|
||||
reactor.callWhenRunning(
|
||||
reactor.spawnProcess, GatheringProtocol(), pyExe,
|
||||
[pyExe, b"-Wignore", b"-c", source], usePTY=self.usePTY)
|
||||
|
||||
self.runReactor(reactor)
|
||||
reportedChildFDs = set(eval(output.getvalue()))
|
||||
|
||||
stdFDs = [0, 1, 2]
|
||||
|
||||
# Unfortunately this assertion is still not *entirely* deterministic,
|
||||
# since hypothetically, any library could open any file descriptor at
|
||||
# any time. See comment above.
|
||||
self.assertEqual(
|
||||
reportedChildFDs.intersection(set(stdFDs + [unlikelyFD])),
|
||||
set(stdFDs)
|
||||
)
|
||||
|
||||
|
||||
@onlyOnPOSIX
|
||||
def test_errorDuringExec(self):
|
||||
"""
|
||||
When L{os.execvpe} raises an exception, it will format that exception
|
||||
on stderr as UTF-8, regardless of system encoding information.
|
||||
"""
|
||||
|
||||
def execvpe(*args, **kw):
|
||||
# Ensure that real traceback formatting has some non-ASCII in it,
|
||||
# by forcing the filename of the last frame to contain non-ASCII.
|
||||
filename = u"<\N{SNOWMAN}>"
|
||||
if not isinstance(filename, str):
|
||||
filename = filename.encode("utf-8")
|
||||
codeobj = compile("1/0", filename, "single")
|
||||
eval(codeobj)
|
||||
|
||||
self.patch(os, "execvpe", execvpe)
|
||||
self.patch(sys, "getfilesystemencoding", lambda: "ascii")
|
||||
|
||||
reactor = self.buildReactor()
|
||||
output = io.BytesIO()
|
||||
|
||||
@reactor.callWhenRunning
|
||||
def whenRunning():
|
||||
class TracebackCatcher(ProcessProtocol, object):
|
||||
errReceived = output.write
|
||||
def processEnded(self, reason):
|
||||
reactor.stop()
|
||||
reactor.spawnProcess(TracebackCatcher(), pyExe,
|
||||
[pyExe, b"-c", b""])
|
||||
|
||||
self.runReactor(reactor, timeout=30)
|
||||
self.assertIn(u"\N{SNOWMAN}".encode("utf-8"), output.getvalue())
|
||||
|
||||
|
||||
def test_timelyProcessExited(self):
|
||||
"""
|
||||
If a spawned process exits, C{processExited} will be called in a
|
||||
timely manner.
|
||||
"""
|
||||
reactor = self.buildReactor()
|
||||
|
||||
class ExitingProtocol(ProcessProtocol):
|
||||
exited = False
|
||||
|
||||
def processExited(protoSelf, reason):
|
||||
protoSelf.exited = True
|
||||
reactor.stop()
|
||||
self.assertEqual(reason.value.exitCode, 0)
|
||||
|
||||
protocol = ExitingProtocol()
|
||||
reactor.callWhenRunning(
|
||||
reactor.spawnProcess, protocol, pyExe,
|
||||
[pyExe, b"-c", b"raise SystemExit(0)"],
|
||||
usePTY=self.usePTY)
|
||||
|
||||
# This will timeout if processExited isn't called:
|
||||
self.runReactor(reactor, timeout=30)
|
||||
self.assertTrue(protocol.exited)
|
||||
|
||||
|
||||
def _changeIDTest(self, which):
|
||||
"""
|
||||
Launch a child process, using either the C{uid} or C{gid} argument to
|
||||
L{IReactorProcess.spawnProcess} to change either its UID or GID to a
|
||||
different value. If the child process reports this hasn't happened,
|
||||
raise an exception to fail the test.
|
||||
|
||||
@param which: Either C{b"uid"} or C{b"gid"}.
|
||||
"""
|
||||
program = [
|
||||
"import os",
|
||||
"raise SystemExit(os.get%s() != 1)" % (which,)]
|
||||
|
||||
container = []
|
||||
class CaptureExitStatus(ProcessProtocol):
|
||||
def processEnded(self, reason):
|
||||
container.append(reason)
|
||||
reactor.stop()
|
||||
|
||||
reactor = self.buildReactor()
|
||||
protocol = CaptureExitStatus()
|
||||
reactor.callWhenRunning(
|
||||
reactor.spawnProcess, protocol, pyExe,
|
||||
[pyExe, "-c", "\n".join(program)],
|
||||
**{which: 1})
|
||||
|
||||
self.runReactor(reactor)
|
||||
|
||||
self.assertEqual(0, container[0].value.exitCode)
|
||||
|
||||
|
||||
def test_changeUID(self):
|
||||
"""
|
||||
If a value is passed for L{IReactorProcess.spawnProcess}'s C{uid}, the
|
||||
child process is run with that UID.
|
||||
"""
|
||||
self._changeIDTest("uid")
|
||||
if _uidgidSkip is not None:
|
||||
test_changeUID.skip = _uidgidSkip
|
||||
|
||||
|
||||
def test_changeGID(self):
|
||||
"""
|
||||
If a value is passed for L{IReactorProcess.spawnProcess}'s C{gid}, the
|
||||
child process is run with that GID.
|
||||
"""
|
||||
self._changeIDTest("gid")
|
||||
if _uidgidSkip is not None:
|
||||
test_changeGID.skip = _uidgidSkip
|
||||
|
||||
|
||||
def test_processExitedRaises(self):
|
||||
"""
|
||||
If L{IProcessProtocol.processExited} raises an exception, it is logged.
|
||||
"""
|
||||
# Ideally we wouldn't need to poke the process module; see
|
||||
# https://twistedmatrix.com/trac/ticket/6889
|
||||
reactor = self.buildReactor()
|
||||
|
||||
class TestException(Exception):
|
||||
pass
|
||||
|
||||
class Protocol(ProcessProtocol):
|
||||
def processExited(self, reason):
|
||||
reactor.stop()
|
||||
raise TestException("processedExited raised")
|
||||
|
||||
protocol = Protocol()
|
||||
transport = reactor.spawnProcess(
|
||||
protocol, pyExe, [pyExe, b"-c", b""],
|
||||
usePTY=self.usePTY)
|
||||
self.runReactor(reactor)
|
||||
|
||||
# Manually clean-up broken process handler.
|
||||
# Only required if the test fails on systems that support
|
||||
# the process module.
|
||||
if process is not None:
|
||||
for pid, handler in items(process.reapProcessHandlers):
|
||||
if handler is not transport:
|
||||
continue
|
||||
process.unregisterReapProcessHandler(pid, handler)
|
||||
self.fail("After processExited raised, transport was left in"
|
||||
" reapProcessHandlers")
|
||||
|
||||
self.assertEqual(1, len(self.flushLoggedErrors(TestException)))
|
||||
|
||||
|
||||
|
||||
class ProcessTestsBuilder(ProcessTestsBuilderBase):
|
||||
"""
|
||||
Builder defining tests relating to L{IReactorProcess} for child processes
|
||||
which do not have a PTY.
|
||||
"""
|
||||
usePTY = False
|
||||
|
||||
keepStdioOpenProgram = b'twisted.internet.test.process_helper'
|
||||
if platform.isWindows():
|
||||
keepStdioOpenArg = b"windows"
|
||||
else:
|
||||
# Just a value that doesn't equal "windows"
|
||||
keepStdioOpenArg = b""
|
||||
|
||||
|
||||
# Define this test here because PTY-using processes only have stdin and
|
||||
# stdout and the test would need to be different for that to work.
|
||||
def test_childConnectionLost(self):
|
||||
"""
|
||||
L{IProcessProtocol.childConnectionLost} is called each time a file
|
||||
descriptor associated with a child process is closed.
|
||||
"""
|
||||
connected = Deferred()
|
||||
lost = {0: Deferred(), 1: Deferred(), 2: Deferred()}
|
||||
|
||||
class Closer(ProcessProtocol):
|
||||
def makeConnection(self, transport):
|
||||
connected.callback(transport)
|
||||
|
||||
def childConnectionLost(self, childFD):
|
||||
lost[childFD].callback(None)
|
||||
|
||||
target = b"twisted.internet.test.process_loseconnection"
|
||||
|
||||
reactor = self.buildReactor()
|
||||
reactor.callWhenRunning(
|
||||
reactor.spawnProcess, Closer(), pyExe,
|
||||
[pyExe, b"-m", target], env=properEnv, usePTY=self.usePTY)
|
||||
|
||||
def cbConnected(transport):
|
||||
transport.write(b'2\n')
|
||||
return lost[2].addCallback(lambda ign: transport)
|
||||
connected.addCallback(cbConnected)
|
||||
|
||||
def lostSecond(transport):
|
||||
transport.write(b'1\n')
|
||||
return lost[1].addCallback(lambda ign: transport)
|
||||
connected.addCallback(lostSecond)
|
||||
|
||||
def lostFirst(transport):
|
||||
transport.write(b'\n')
|
||||
connected.addCallback(lostFirst)
|
||||
connected.addErrback(err)
|
||||
|
||||
def cbEnded(ignored):
|
||||
reactor.stop()
|
||||
connected.addCallback(cbEnded)
|
||||
|
||||
self.runReactor(reactor)
|
||||
|
||||
|
||||
# This test is here because PTYProcess never delivers childConnectionLost.
|
||||
def test_processEnded(self):
|
||||
"""
|
||||
L{IProcessProtocol.processEnded} is called after the child process
|
||||
exits and L{IProcessProtocol.childConnectionLost} is called for each of
|
||||
its file descriptors.
|
||||
"""
|
||||
ended = Deferred()
|
||||
lost = []
|
||||
|
||||
class Ender(ProcessProtocol):
|
||||
def childDataReceived(self, fd, data):
|
||||
msg('childDataReceived(%d, %r)' % (fd, data))
|
||||
self.transport.loseConnection()
|
||||
|
||||
def childConnectionLost(self, childFD):
|
||||
msg('childConnectionLost(%d)' % (childFD,))
|
||||
lost.append(childFD)
|
||||
|
||||
def processExited(self, reason):
|
||||
msg('processExited(%r)' % (reason,))
|
||||
|
||||
def processEnded(self, reason):
|
||||
msg('processEnded(%r)' % (reason,))
|
||||
ended.callback([reason])
|
||||
|
||||
reactor = self.buildReactor()
|
||||
reactor.callWhenRunning(
|
||||
reactor.spawnProcess, Ender(), pyExe,
|
||||
[pyExe, b"-m", self.keepStdioOpenProgram, b"child",
|
||||
self.keepStdioOpenArg],
|
||||
env=properEnv, usePTY=self.usePTY)
|
||||
|
||||
def cbEnded(args):
|
||||
failure, = args
|
||||
failure.trap(ProcessDone)
|
||||
self.assertEqual(set(lost), set([0, 1, 2]))
|
||||
ended.addCallback(cbEnded)
|
||||
|
||||
ended.addErrback(err)
|
||||
ended.addCallback(lambda ign: reactor.stop())
|
||||
|
||||
self.runReactor(reactor)
|
||||
|
||||
|
||||
# This test is here because PTYProcess.loseConnection does not actually
|
||||
# close the file descriptors to the child process. This test needs to be
|
||||
# written fairly differently for PTYProcess.
|
||||
def test_processExited(self):
|
||||
"""
|
||||
L{IProcessProtocol.processExited} is called when the child process
|
||||
exits, even if file descriptors associated with the child are still
|
||||
open.
|
||||
"""
|
||||
exited = Deferred()
|
||||
allLost = Deferred()
|
||||
lost = []
|
||||
|
||||
class Waiter(ProcessProtocol):
|
||||
def childDataReceived(self, fd, data):
|
||||
msg('childDataReceived(%d, %r)' % (fd, data))
|
||||
|
||||
def childConnectionLost(self, childFD):
|
||||
msg('childConnectionLost(%d)' % (childFD,))
|
||||
lost.append(childFD)
|
||||
if len(lost) == 3:
|
||||
allLost.callback(None)
|
||||
|
||||
def processExited(self, reason):
|
||||
msg('processExited(%r)' % (reason,))
|
||||
# See test_processExitedWithSignal
|
||||
exited.callback([reason])
|
||||
self.transport.loseConnection()
|
||||
|
||||
reactor = self.buildReactor()
|
||||
reactor.callWhenRunning(
|
||||
reactor.spawnProcess, Waiter(), pyExe,
|
||||
[pyExe, b"-u", b"-m", self.keepStdioOpenProgram, b"child",
|
||||
self.keepStdioOpenArg],
|
||||
env=properEnv, usePTY=self.usePTY)
|
||||
|
||||
def cbExited(args):
|
||||
failure, = args
|
||||
failure.trap(ProcessDone)
|
||||
msg('cbExited; lost = %s' % (lost,))
|
||||
self.assertEqual(lost, [])
|
||||
return allLost
|
||||
exited.addCallback(cbExited)
|
||||
|
||||
def cbAllLost(ignored):
|
||||
self.assertEqual(set(lost), set([0, 1, 2]))
|
||||
exited.addCallback(cbAllLost)
|
||||
|
||||
exited.addErrback(err)
|
||||
exited.addCallback(lambda ign: reactor.stop())
|
||||
|
||||
self.runReactor(reactor)
|
||||
|
||||
|
||||
def makeSourceFile(self, sourceLines):
|
||||
"""
|
||||
Write the given list of lines to a text file and return the absolute
|
||||
path to it.
|
||||
"""
|
||||
script = _asFilesystemBytes(self.mktemp())
|
||||
with open(script, 'wt') as scriptFile:
|
||||
scriptFile.write(os.linesep.join(sourceLines) + os.linesep)
|
||||
return os.path.abspath(script)
|
||||
|
||||
|
||||
def test_shebang(self):
|
||||
"""
|
||||
Spawning a process with an executable which is a script starting
|
||||
with an interpreter definition line (#!) uses that interpreter to
|
||||
evaluate the script.
|
||||
"""
|
||||
shebangOutput = b'this is the shebang output'
|
||||
|
||||
scriptFile = self.makeSourceFile([
|
||||
"#!%s" % (pyExe.decode('ascii'),),
|
||||
"import sys",
|
||||
"sys.stdout.write('%s')" % (shebangOutput.decode('ascii'),),
|
||||
"sys.stdout.flush()"])
|
||||
os.chmod(scriptFile, 0o700)
|
||||
|
||||
reactor = self.buildReactor()
|
||||
|
||||
def cbProcessExited(args):
|
||||
out, err, code = args
|
||||
msg("cbProcessExited((%r, %r, %d))" % (out, err, code))
|
||||
self.assertEqual(out, shebangOutput)
|
||||
self.assertEqual(err, b"")
|
||||
self.assertEqual(code, 0)
|
||||
|
||||
def shutdown(passthrough):
|
||||
reactor.stop()
|
||||
return passthrough
|
||||
|
||||
def start():
|
||||
d = utils.getProcessOutputAndValue(scriptFile, reactor=reactor)
|
||||
d.addBoth(shutdown)
|
||||
d.addCallback(cbProcessExited)
|
||||
d.addErrback(err)
|
||||
|
||||
reactor.callWhenRunning(start)
|
||||
self.runReactor(reactor)
|
||||
|
||||
|
||||
def test_processCommandLineArguments(self):
|
||||
"""
|
||||
Arguments given to spawnProcess are passed to the child process as
|
||||
originally intended.
|
||||
"""
|
||||
us = b"twisted.internet.test.process_cli"
|
||||
|
||||
args = [b'hello', b'"', b' \t|<>^&', br'"\\"hello\\"', br'"foo\ bar baz\""']
|
||||
# Ensure that all non-NUL characters can be passed too.
|
||||
allChars = "".join(map(chr, range(1, 255)))
|
||||
if isinstance(allChars, unicode):
|
||||
allChars.encode("utf-8")
|
||||
|
||||
reactor = self.buildReactor()
|
||||
|
||||
def processFinished(finishedArgs):
|
||||
output, err, code = finishedArgs
|
||||
output = output.split(b'\0')
|
||||
# Drop the trailing \0.
|
||||
output.pop()
|
||||
self.assertEqual(args, output)
|
||||
|
||||
def shutdown(result):
|
||||
reactor.stop()
|
||||
return result
|
||||
|
||||
def spawnChild():
|
||||
d = succeed(None)
|
||||
d.addCallback(lambda dummy: utils.getProcessOutputAndValue(
|
||||
pyExe, [b"-m", us] + args, env=properEnv,
|
||||
reactor=reactor))
|
||||
d.addCallback(processFinished)
|
||||
d.addBoth(shutdown)
|
||||
|
||||
reactor.callWhenRunning(spawnChild)
|
||||
self.runReactor(reactor)
|
||||
globals().update(ProcessTestsBuilder.makeTestCaseClasses())
|
||||
|
||||
|
||||
|
||||
class PTYProcessTestsBuilder(ProcessTestsBuilderBase):
|
||||
"""
|
||||
Builder defining tests relating to L{IReactorProcess} for child processes
|
||||
which have a PTY.
|
||||
"""
|
||||
usePTY = True
|
||||
|
||||
if platform.isWindows():
|
||||
skip = "PTYs are not supported on Windows."
|
||||
elif platform.isMacOSX():
|
||||
skip = "PTYs are flaky from a Darwin bug. See #8840."
|
||||
|
||||
skippedReactors = {
|
||||
"twisted.internet.pollreactor.PollReactor":
|
||||
"macOS's poll() does not support PTYs"}
|
||||
globals().update(PTYProcessTestsBuilder.makeTestCaseClasses())
|
||||
|
||||
|
||||
|
||||
class PotentialZombieWarningTests(TestCase):
|
||||
"""
|
||||
Tests for L{twisted.internet.error.PotentialZombieWarning}.
|
||||
"""
|
||||
def test_deprecated(self):
|
||||
"""
|
||||
Accessing L{PotentialZombieWarning} via the
|
||||
I{PotentialZombieWarning} attribute of L{twisted.internet.error}
|
||||
results in a deprecation warning being emitted.
|
||||
"""
|
||||
from twisted.internet import error
|
||||
error.PotentialZombieWarning
|
||||
|
||||
warnings = self.flushWarnings([self.test_deprecated])
|
||||
self.assertEqual(warnings[0]['category'], DeprecationWarning)
|
||||
self.assertEqual(
|
||||
warnings[0]['message'],
|
||||
"twisted.internet.error.PotentialZombieWarning was deprecated in "
|
||||
"Twisted 10.0.0: There is no longer any potential for zombie "
|
||||
"process.")
|
||||
self.assertEqual(len(warnings), 1)
|
||||
|
||||
|
||||
|
||||
class ProcessIsUnimportableOnUnsupportedPlatormsTests(TestCase):
|
||||
"""
|
||||
Tests to ensure that L{twisted.internet.process} is unimportable on
|
||||
platforms where it does not work (namely Windows).
|
||||
"""
|
||||
def test_unimportableOnWindows(self):
|
||||
"""
|
||||
L{twisted.internet.process} is unimportable on Windows.
|
||||
"""
|
||||
with self.assertRaises(ImportError):
|
||||
import twisted.internet.process
|
||||
twisted.internet.process # shh pyflakes
|
||||
|
||||
if not platform.isWindows():
|
||||
test_unimportableOnWindows.skip = "Only relevant on Windows."
|
||||
@@ -0,0 +1,125 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.internet._sigchld}, an alternate, superior SIGCHLD
|
||||
monitoring API.
|
||||
"""
|
||||
|
||||
from __future__ import division, absolute_import
|
||||
|
||||
import os, signal, errno
|
||||
|
||||
from twisted.python.runtime import platformType
|
||||
from twisted.python.log import msg
|
||||
from twisted.trial.unittest import SynchronousTestCase
|
||||
if platformType == "posix":
|
||||
from twisted.internet.fdesc import setNonBlocking
|
||||
from twisted.internet._signals import installHandler, isDefaultHandler
|
||||
else:
|
||||
skip = "These tests can only run on POSIX platforms."
|
||||
|
||||
|
||||
class SetWakeupSIGCHLDTests(SynchronousTestCase):
|
||||
"""
|
||||
Tests for the L{signal.set_wakeup_fd} implementation of the
|
||||
L{installHandler} and L{isDefaultHandler} APIs.
|
||||
"""
|
||||
|
||||
def pipe(self):
|
||||
"""
|
||||
Create a non-blocking pipe which will be closed after the currently
|
||||
running test.
|
||||
"""
|
||||
read, write = os.pipe()
|
||||
self.addCleanup(os.close, read)
|
||||
self.addCleanup(os.close, write)
|
||||
setNonBlocking(read)
|
||||
setNonBlocking(write)
|
||||
return read, write
|
||||
|
||||
|
||||
def setUp(self):
|
||||
"""
|
||||
Save the current SIGCHLD handler as reported by L{signal.signal} and
|
||||
the current file descriptor registered with L{installHandler}.
|
||||
"""
|
||||
handler = signal.getsignal(signal.SIGCHLD)
|
||||
if handler != signal.SIG_DFL:
|
||||
self.signalModuleHandler = handler
|
||||
signal.signal(signal.SIGCHLD, signal.SIG_DFL)
|
||||
else:
|
||||
self.signalModuleHandler = None
|
||||
|
||||
self.oldFD = installHandler(-1)
|
||||
|
||||
if self.signalModuleHandler is not None and self.oldFD != -1:
|
||||
msg("Previous test didn't clean up after its SIGCHLD setup: %r %r"
|
||||
% (self.signalModuleHandler, self.oldFD))
|
||||
|
||||
|
||||
def tearDown(self):
|
||||
"""
|
||||
Restore whatever signal handler was present when setUp ran.
|
||||
"""
|
||||
# If tests set up any kind of handlers, clear them out.
|
||||
installHandler(-1)
|
||||
signal.signal(signal.SIGCHLD, signal.SIG_DFL)
|
||||
|
||||
# Now restore whatever the setup was before the test ran.
|
||||
if self.signalModuleHandler is not None:
|
||||
signal.signal(signal.SIGCHLD, self.signalModuleHandler)
|
||||
elif self.oldFD != -1:
|
||||
installHandler(self.oldFD)
|
||||
|
||||
|
||||
def test_isDefaultHandler(self):
|
||||
"""
|
||||
L{isDefaultHandler} returns true if the SIGCHLD handler is SIG_DFL,
|
||||
false otherwise.
|
||||
"""
|
||||
self.assertTrue(isDefaultHandler())
|
||||
signal.signal(signal.SIGCHLD, signal.SIG_IGN)
|
||||
self.assertFalse(isDefaultHandler())
|
||||
signal.signal(signal.SIGCHLD, signal.SIG_DFL)
|
||||
self.assertTrue(isDefaultHandler())
|
||||
signal.signal(signal.SIGCHLD, lambda *args: None)
|
||||
self.assertFalse(isDefaultHandler())
|
||||
|
||||
|
||||
def test_returnOldFD(self):
|
||||
"""
|
||||
L{installHandler} returns the previously registered file descriptor.
|
||||
"""
|
||||
read, write = self.pipe()
|
||||
oldFD = installHandler(write)
|
||||
self.assertEqual(installHandler(oldFD), write)
|
||||
|
||||
|
||||
def test_uninstallHandler(self):
|
||||
"""
|
||||
C{installHandler(-1)} removes the SIGCHLD handler completely.
|
||||
"""
|
||||
read, write = self.pipe()
|
||||
self.assertTrue(isDefaultHandler())
|
||||
installHandler(write)
|
||||
self.assertFalse(isDefaultHandler())
|
||||
installHandler(-1)
|
||||
self.assertTrue(isDefaultHandler())
|
||||
|
||||
|
||||
def test_installHandler(self):
|
||||
"""
|
||||
The file descriptor passed to L{installHandler} has a byte written to
|
||||
it when SIGCHLD is delivered to the process.
|
||||
"""
|
||||
read, write = self.pipe()
|
||||
installHandler(write)
|
||||
|
||||
exc = self.assertRaises(OSError, os.read, read, 1)
|
||||
self.assertEqual(exc.errno, errno.EAGAIN)
|
||||
|
||||
os.kill(os.getpid(), signal.SIGCHLD)
|
||||
|
||||
self.assertEqual(len(os.read(read, 5)), 1)
|
||||
|
||||
@@ -0,0 +1,199 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.internet.stdio}.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
from twisted.python.runtime import platform
|
||||
from twisted.internet.test.reactormixins import ReactorBuilder
|
||||
from twisted.internet.protocol import Protocol
|
||||
|
||||
if not platform.isWindows():
|
||||
from twisted.internet.stdio import StandardIO
|
||||
|
||||
|
||||
|
||||
class StdioFilesTests(ReactorBuilder):
|
||||
"""
|
||||
L{StandardIO} supports reading and writing to filesystem files.
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
path = self.mktemp()
|
||||
open(path, "wb").close()
|
||||
self.extraFile = open(path, "rb+")
|
||||
self.addCleanup(self.extraFile.close)
|
||||
|
||||
|
||||
def test_addReader(self):
|
||||
"""
|
||||
Adding a filesystem file reader to a reactor will make sure it is
|
||||
polled.
|
||||
"""
|
||||
reactor = self.buildReactor()
|
||||
|
||||
class DataProtocol(Protocol):
|
||||
data = b""
|
||||
def dataReceived(self, data):
|
||||
self.data += data
|
||||
# It'd be better to stop reactor on connectionLost, but that
|
||||
# fails on FreeBSD, probably due to
|
||||
# http://bugs.python.org/issue9591:
|
||||
if self.data == b"hello!":
|
||||
reactor.stop()
|
||||
|
||||
path = self.mktemp()
|
||||
|
||||
with open(path, "wb") as f:
|
||||
f.write(b"hello!")
|
||||
|
||||
with open(path, "rb") as f:
|
||||
# Read bytes from a file, deliver them to a protocol instance:
|
||||
protocol = DataProtocol()
|
||||
StandardIO(protocol, stdin=f.fileno(),
|
||||
stdout=self.extraFile.fileno(),
|
||||
reactor=reactor)
|
||||
self.runReactor(reactor)
|
||||
|
||||
self.assertEqual(protocol.data, b"hello!")
|
||||
|
||||
|
||||
def test_addWriter(self):
|
||||
"""
|
||||
Adding a filesystem file writer to a reactor will make sure it is
|
||||
polled.
|
||||
"""
|
||||
reactor = self.buildReactor()
|
||||
|
||||
class DisconnectProtocol(Protocol):
|
||||
def connectionLost(self, reason):
|
||||
reactor.stop()
|
||||
|
||||
path = self.mktemp()
|
||||
|
||||
with open(path, "wb") as f:
|
||||
# Write bytes to a transport, hopefully have them written to a
|
||||
# file:
|
||||
protocol = DisconnectProtocol()
|
||||
StandardIO(protocol, stdout=f.fileno(),
|
||||
stdin=self.extraFile.fileno(), reactor=reactor)
|
||||
protocol.transport.write(b"hello")
|
||||
protocol.transport.write(b", world")
|
||||
protocol.transport.loseConnection()
|
||||
|
||||
self.runReactor(reactor)
|
||||
|
||||
with open(path, "rb") as f:
|
||||
self.assertEqual(f.read(), b"hello, world")
|
||||
|
||||
|
||||
def test_removeReader(self):
|
||||
"""
|
||||
Removing a filesystem file reader from a reactor will make sure it is
|
||||
no longer polled.
|
||||
"""
|
||||
reactor = self.buildReactor()
|
||||
self.addCleanup(self.unbuildReactor, reactor)
|
||||
|
||||
path = self.mktemp()
|
||||
open(path, "wb").close()
|
||||
|
||||
with open(path, "rb") as f:
|
||||
# Have the reader added:
|
||||
stdio = StandardIO(Protocol(), stdin=f.fileno(),
|
||||
stdout=self.extraFile.fileno(),
|
||||
reactor=reactor)
|
||||
self.assertIn(stdio._reader, reactor.getReaders())
|
||||
stdio._reader.stopReading()
|
||||
self.assertNotIn(stdio._reader, reactor.getReaders())
|
||||
|
||||
|
||||
def test_removeWriter(self):
|
||||
"""
|
||||
Removing a filesystem file writer from a reactor will make sure it is
|
||||
no longer polled.
|
||||
"""
|
||||
reactor = self.buildReactor()
|
||||
self.addCleanup(self.unbuildReactor, reactor)
|
||||
|
||||
# Cleanup might fail if file is GCed too soon:
|
||||
self.f = f = open(self.mktemp(), "wb")
|
||||
|
||||
# Have the reader added:
|
||||
protocol = Protocol()
|
||||
stdio = StandardIO(protocol, stdout=f.fileno(),
|
||||
stdin=self.extraFile.fileno(),
|
||||
reactor=reactor)
|
||||
protocol.transport.write(b"hello")
|
||||
self.assertIn(stdio._writer, reactor.getWriters())
|
||||
stdio._writer.stopWriting()
|
||||
self.assertNotIn(stdio._writer, reactor.getWriters())
|
||||
|
||||
|
||||
def test_removeAll(self):
|
||||
"""
|
||||
Calling C{removeAll} on a reactor includes descriptors that are
|
||||
filesystem files.
|
||||
"""
|
||||
reactor = self.buildReactor()
|
||||
self.addCleanup(self.unbuildReactor, reactor)
|
||||
|
||||
path = self.mktemp()
|
||||
open(path, "wb").close()
|
||||
|
||||
# Cleanup might fail if file is GCed too soon:
|
||||
self.f = f = open(path, "rb")
|
||||
|
||||
# Have the reader added:
|
||||
stdio = StandardIO(Protocol(), stdin=f.fileno(),
|
||||
stdout=self.extraFile.fileno(), reactor=reactor)
|
||||
# And then removed:
|
||||
removed = reactor.removeAll()
|
||||
self.assertIn(stdio._reader, removed)
|
||||
self.assertNotIn(stdio._reader, reactor.getReaders())
|
||||
|
||||
|
||||
def test_getReaders(self):
|
||||
"""
|
||||
C{reactor.getReaders} includes descriptors that are filesystem files.
|
||||
"""
|
||||
reactor = self.buildReactor()
|
||||
self.addCleanup(self.unbuildReactor, reactor)
|
||||
|
||||
path = self.mktemp()
|
||||
open(path, "wb").close()
|
||||
|
||||
# Cleanup might fail if file is GCed too soon:
|
||||
with open(path, "rb") as f:
|
||||
# Have the reader added:
|
||||
stdio = StandardIO(Protocol(), stdin=f.fileno(),
|
||||
stdout=self.extraFile.fileno(), reactor=reactor)
|
||||
self.assertIn(stdio._reader, reactor.getReaders())
|
||||
|
||||
|
||||
def test_getWriters(self):
|
||||
"""
|
||||
C{reactor.getWriters} includes descriptors that are filesystem files.
|
||||
"""
|
||||
reactor = self.buildReactor()
|
||||
self.addCleanup(self.unbuildReactor, reactor)
|
||||
|
||||
# Cleanup might fail if file is GCed too soon:
|
||||
self.f = f = open(self.mktemp(), "wb")
|
||||
|
||||
# Have the reader added:
|
||||
stdio = StandardIO(Protocol(), stdout=f.fileno(),
|
||||
stdin=self.extraFile.fileno(), reactor=reactor)
|
||||
self.assertNotIn(stdio._writer, reactor.getWriters())
|
||||
stdio._writer.startWriting()
|
||||
self.assertIn(stdio._writer, reactor.getWriters())
|
||||
|
||||
if platform.isWindows():
|
||||
skip = ("StandardIO does not accept stdout as an argument to Windows. "
|
||||
"Testing redirection to a file is therefore harder.")
|
||||
|
||||
|
||||
globals().update(StdioFilesTests.makeTestCaseClasses())
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,515 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.internet.testing}.
|
||||
"""
|
||||
|
||||
from zope.interface.verify import verifyObject
|
||||
|
||||
from twisted.internet.interfaces import (
|
||||
ITransport,
|
||||
IPushProducer,
|
||||
IConsumer,
|
||||
IReactorTCP,
|
||||
IReactorSSL,
|
||||
IReactorUNIX,
|
||||
IAddress,
|
||||
IListeningPort,
|
||||
IConnector
|
||||
)
|
||||
from twisted.internet.address import IPv4Address
|
||||
from twisted.trial.unittest import TestCase
|
||||
from twisted.internet.testing import (
|
||||
StringTransport,
|
||||
MemoryReactor,
|
||||
RaisingMemoryReactor,
|
||||
NonStreamingProducer
|
||||
)
|
||||
from twisted.internet.protocol import ClientFactory, Factory
|
||||
from twisted.python.reflect import namedAny
|
||||
|
||||
|
||||
|
||||
class StringTransportTests(TestCase):
|
||||
"""
|
||||
Tests for L{twisted.internet.testing.StringTransport}.
|
||||
"""
|
||||
def setUp(self):
|
||||
self.transport = StringTransport()
|
||||
|
||||
|
||||
def test_interfaces(self):
|
||||
"""
|
||||
L{StringTransport} instances provide L{ITransport}, L{IPushProducer},
|
||||
and L{IConsumer}.
|
||||
"""
|
||||
self.assertTrue(verifyObject(ITransport, self.transport))
|
||||
self.assertTrue(verifyObject(IPushProducer, self.transport))
|
||||
self.assertTrue(verifyObject(IConsumer, self.transport))
|
||||
|
||||
|
||||
def test_registerProducer(self):
|
||||
"""
|
||||
L{StringTransport.registerProducer} records the arguments supplied to
|
||||
it as instance attributes.
|
||||
"""
|
||||
producer = object()
|
||||
streaming = object()
|
||||
self.transport.registerProducer(producer, streaming)
|
||||
self.assertIs(self.transport.producer, producer)
|
||||
self.assertIs(self.transport.streaming, streaming)
|
||||
|
||||
|
||||
def test_disallowedRegisterProducer(self):
|
||||
"""
|
||||
L{StringTransport.registerProducer} raises L{RuntimeError} if a
|
||||
producer is already registered.
|
||||
"""
|
||||
producer = object()
|
||||
self.transport.registerProducer(producer, True)
|
||||
self.assertRaises(
|
||||
RuntimeError, self.transport.registerProducer, object(), False)
|
||||
self.assertIs(self.transport.producer, producer)
|
||||
self.assertTrue(self.transport.streaming)
|
||||
|
||||
|
||||
def test_unregisterProducer(self):
|
||||
"""
|
||||
L{StringTransport.unregisterProducer} causes the transport to forget
|
||||
about the registered producer and makes it possible to register a new
|
||||
one.
|
||||
"""
|
||||
oldProducer = object()
|
||||
newProducer = object()
|
||||
self.transport.registerProducer(oldProducer, False)
|
||||
self.transport.unregisterProducer()
|
||||
self.assertIsNone(self.transport.producer)
|
||||
self.transport.registerProducer(newProducer, True)
|
||||
self.assertIs(self.transport.producer, newProducer)
|
||||
self.assertTrue(self.transport.streaming)
|
||||
|
||||
|
||||
def test_invalidUnregisterProducer(self):
|
||||
"""
|
||||
L{StringTransport.unregisterProducer} raises L{RuntimeError} if called
|
||||
when no producer is registered.
|
||||
"""
|
||||
self.assertRaises(RuntimeError, self.transport.unregisterProducer)
|
||||
|
||||
|
||||
def test_initialProducerState(self):
|
||||
"""
|
||||
L{StringTransport.producerState} is initially C{'producing'}.
|
||||
"""
|
||||
self.assertEqual(self.transport.producerState, 'producing')
|
||||
|
||||
|
||||
def test_pauseProducing(self):
|
||||
"""
|
||||
L{StringTransport.pauseProducing} changes the C{producerState} of the
|
||||
transport to C{'paused'}.
|
||||
"""
|
||||
self.transport.pauseProducing()
|
||||
self.assertEqual(self.transport.producerState, 'paused')
|
||||
|
||||
|
||||
def test_resumeProducing(self):
|
||||
"""
|
||||
L{StringTransport.resumeProducing} changes the C{producerState} of the
|
||||
transport to C{'producing'}.
|
||||
"""
|
||||
self.transport.pauseProducing()
|
||||
self.transport.resumeProducing()
|
||||
self.assertEqual(self.transport.producerState, 'producing')
|
||||
|
||||
|
||||
def test_stopProducing(self):
|
||||
"""
|
||||
L{StringTransport.stopProducing} changes the C{'producerState'} of the
|
||||
transport to C{'stopped'}.
|
||||
"""
|
||||
self.transport.stopProducing()
|
||||
self.assertEqual(self.transport.producerState, 'stopped')
|
||||
|
||||
|
||||
def test_stoppedTransportCannotPause(self):
|
||||
"""
|
||||
L{StringTransport.pauseProducing} raises L{RuntimeError} if the
|
||||
transport has been stopped.
|
||||
"""
|
||||
self.transport.stopProducing()
|
||||
self.assertRaises(RuntimeError, self.transport.pauseProducing)
|
||||
|
||||
|
||||
def test_stoppedTransportCannotResume(self):
|
||||
"""
|
||||
L{StringTransport.resumeProducing} raises L{RuntimeError} if the
|
||||
transport has been stopped.
|
||||
"""
|
||||
self.transport.stopProducing()
|
||||
self.assertRaises(RuntimeError, self.transport.resumeProducing)
|
||||
|
||||
|
||||
def test_disconnectingTransportCannotPause(self):
|
||||
"""
|
||||
L{StringTransport.pauseProducing} raises L{RuntimeError} if the
|
||||
transport is being disconnected.
|
||||
"""
|
||||
self.transport.loseConnection()
|
||||
self.assertRaises(RuntimeError, self.transport.pauseProducing)
|
||||
|
||||
|
||||
def test_disconnectingTransportCannotResume(self):
|
||||
"""
|
||||
L{StringTransport.resumeProducing} raises L{RuntimeError} if the
|
||||
transport is being disconnected.
|
||||
"""
|
||||
self.transport.loseConnection()
|
||||
self.assertRaises(RuntimeError, self.transport.resumeProducing)
|
||||
|
||||
|
||||
def test_loseConnectionSetsDisconnecting(self):
|
||||
"""
|
||||
L{StringTransport.loseConnection} toggles the C{disconnecting} instance
|
||||
variable to C{True}.
|
||||
"""
|
||||
self.assertFalse(self.transport.disconnecting)
|
||||
self.transport.loseConnection()
|
||||
self.assertTrue(self.transport.disconnecting)
|
||||
|
||||
|
||||
def test_specifiedHostAddress(self):
|
||||
"""
|
||||
If a host address is passed to L{StringTransport.__init__}, that
|
||||
value is returned from L{StringTransport.getHost}.
|
||||
"""
|
||||
address = object()
|
||||
self.assertIs(StringTransport(address).getHost(), address)
|
||||
|
||||
|
||||
def test_specifiedPeerAddress(self):
|
||||
"""
|
||||
If a peer address is passed to L{StringTransport.__init__}, that
|
||||
value is returned from L{StringTransport.getPeer}.
|
||||
"""
|
||||
address = object()
|
||||
self.assertIs(
|
||||
StringTransport(peerAddress=address).getPeer(), address)
|
||||
|
||||
|
||||
def test_defaultHostAddress(self):
|
||||
"""
|
||||
If no host address is passed to L{StringTransport.__init__}, an
|
||||
L{IPv4Address} is returned from L{StringTransport.getHost}.
|
||||
"""
|
||||
address = StringTransport().getHost()
|
||||
self.assertIsInstance(address, IPv4Address)
|
||||
|
||||
|
||||
def test_defaultPeerAddress(self):
|
||||
"""
|
||||
If no peer address is passed to L{StringTransport.__init__}, an
|
||||
L{IPv4Address} is returned from L{StringTransport.getPeer}.
|
||||
"""
|
||||
address = StringTransport().getPeer()
|
||||
self.assertIsInstance(address, IPv4Address)
|
||||
|
||||
|
||||
|
||||
class ReactorTests(TestCase):
|
||||
"""
|
||||
Tests for L{MemoryReactor} and L{RaisingMemoryReactor}.
|
||||
"""
|
||||
|
||||
def test_memoryReactorProvides(self):
|
||||
"""
|
||||
L{MemoryReactor} provides all of the attributes described by the
|
||||
interfaces it advertises.
|
||||
"""
|
||||
memoryReactor = MemoryReactor()
|
||||
verifyObject(IReactorTCP, memoryReactor)
|
||||
verifyObject(IReactorSSL, memoryReactor)
|
||||
verifyObject(IReactorUNIX, memoryReactor)
|
||||
|
||||
|
||||
def test_raisingReactorProvides(self):
|
||||
"""
|
||||
L{RaisingMemoryReactor} provides all of the attributes described by the
|
||||
interfaces it advertises.
|
||||
"""
|
||||
raisingReactor = RaisingMemoryReactor()
|
||||
verifyObject(IReactorTCP, raisingReactor)
|
||||
verifyObject(IReactorSSL, raisingReactor)
|
||||
verifyObject(IReactorUNIX, raisingReactor)
|
||||
|
||||
|
||||
def test_connectDestination(self):
|
||||
"""
|
||||
L{MemoryReactor.connectTCP}, L{MemoryReactor.connectSSL}, and
|
||||
L{MemoryReactor.connectUNIX} will return an L{IConnector} whose
|
||||
C{getDestination} method returns an L{IAddress} with attributes which
|
||||
reflect the values passed.
|
||||
"""
|
||||
memoryReactor = MemoryReactor()
|
||||
for connector in [memoryReactor.connectTCP(
|
||||
"test.example.com", 8321, ClientFactory()),
|
||||
memoryReactor.connectSSL(
|
||||
"test.example.com", 8321, ClientFactory(),
|
||||
None)]:
|
||||
verifyObject(IConnector, connector)
|
||||
address = connector.getDestination()
|
||||
verifyObject(IAddress, address)
|
||||
self.assertEqual(address.host, "test.example.com")
|
||||
self.assertEqual(address.port, 8321)
|
||||
connector = memoryReactor.connectUNIX(b"/fake/path", ClientFactory())
|
||||
verifyObject(IConnector, connector)
|
||||
address = connector.getDestination()
|
||||
verifyObject(IAddress, address)
|
||||
self.assertEqual(address.name, b"/fake/path")
|
||||
|
||||
|
||||
def test_listenDefaultHost(self):
|
||||
"""
|
||||
L{MemoryReactor.listenTCP}, L{MemoryReactor.listenSSL} and
|
||||
L{MemoryReactor.listenUNIX} will return an L{IListeningPort} whose
|
||||
C{getHost} method returns an L{IAddress}; C{listenTCP} and C{listenSSL}
|
||||
will have a default host of C{'0.0.0.0'}, and a port that reflects the
|
||||
value passed, and C{listenUNIX} will have a name that reflects the path
|
||||
passed.
|
||||
"""
|
||||
memoryReactor = MemoryReactor()
|
||||
for port in [memoryReactor.listenTCP(8242, Factory()),
|
||||
memoryReactor.listenSSL(8242, Factory(), None)]:
|
||||
verifyObject(IListeningPort, port)
|
||||
address = port.getHost()
|
||||
verifyObject(IAddress, address)
|
||||
self.assertEqual(address.host, '0.0.0.0')
|
||||
self.assertEqual(address.port, 8242)
|
||||
port = memoryReactor.listenUNIX(b"/path/to/socket", Factory())
|
||||
verifyObject(IListeningPort, port)
|
||||
address = port.getHost()
|
||||
verifyObject(IAddress, address)
|
||||
self.assertEqual(address.name, b"/path/to/socket")
|
||||
|
||||
|
||||
def test_readers(self):
|
||||
"""
|
||||
Adding, removing, and listing readers works.
|
||||
"""
|
||||
reader = object()
|
||||
reactor = MemoryReactor()
|
||||
|
||||
reactor.addReader(reader)
|
||||
reactor.addReader(reader)
|
||||
|
||||
self.assertEqual(reactor.getReaders(), [reader])
|
||||
|
||||
reactor.removeReader(reader)
|
||||
|
||||
self.assertEqual(reactor.getReaders(), [])
|
||||
|
||||
|
||||
def test_writers(self):
|
||||
"""
|
||||
Adding, removing, and listing writers works.
|
||||
"""
|
||||
writer = object()
|
||||
reactor = MemoryReactor()
|
||||
|
||||
reactor.addWriter(writer)
|
||||
reactor.addWriter(writer)
|
||||
|
||||
self.assertEqual(reactor.getWriters(), [writer])
|
||||
|
||||
reactor.removeWriter(writer)
|
||||
|
||||
self.assertEqual(reactor.getWriters(), [])
|
||||
|
||||
|
||||
|
||||
class TestConsumer(object):
|
||||
"""
|
||||
A very basic test consumer for use with the NonStreamingProducerTests.
|
||||
"""
|
||||
def __init__(self):
|
||||
self.writes = []
|
||||
self.producer = None
|
||||
self.producerStreaming = None
|
||||
|
||||
|
||||
def registerProducer(self, producer, streaming):
|
||||
"""
|
||||
Registers a single producer with this consumer. Just keeps track of it.
|
||||
|
||||
@param producer: The producer to register.
|
||||
@param streaming: Whether the producer is a streaming one or not.
|
||||
"""
|
||||
self.producer = producer
|
||||
self.producerStreaming = streaming
|
||||
|
||||
|
||||
def unregisterProducer(self):
|
||||
"""
|
||||
Forget the producer we had previously registered.
|
||||
"""
|
||||
self.producer = None
|
||||
self.producerStreaming = None
|
||||
|
||||
|
||||
def write(self, data):
|
||||
"""
|
||||
Some data was written to the consumer: stores it for later use.
|
||||
|
||||
@param data: The data to write.
|
||||
"""
|
||||
self.writes.append(data)
|
||||
|
||||
|
||||
|
||||
class NonStreamingProducerTests(TestCase):
|
||||
"""
|
||||
Tests for the L{NonStreamingProducer} to validate behaviour.
|
||||
"""
|
||||
def test_producesOnly10Times(self):
|
||||
"""
|
||||
When the L{NonStreamingProducer} has resumeProducing called 10 times,
|
||||
it writes the counter each time and then fails.
|
||||
"""
|
||||
consumer = TestConsumer()
|
||||
producer = NonStreamingProducer(consumer)
|
||||
consumer.registerProducer(producer, False)
|
||||
|
||||
self.assertIs(consumer.producer, producer)
|
||||
self.assertIs(producer.consumer, consumer)
|
||||
self.assertFalse(consumer.producerStreaming)
|
||||
|
||||
for _ in range(10):
|
||||
producer.resumeProducing()
|
||||
|
||||
# We should have unregistered the producer and printed the 10 results.
|
||||
expectedWrites = [
|
||||
b'0', b'1', b'2', b'3', b'4', b'5', b'6', b'7', b'8', b'9'
|
||||
]
|
||||
self.assertIsNone(consumer.producer)
|
||||
self.assertIsNone(consumer.producerStreaming)
|
||||
self.assertIsNone(producer.consumer)
|
||||
self.assertEqual(consumer.writes, expectedWrites)
|
||||
|
||||
# Another attempt to produce fails.
|
||||
self.assertRaises(RuntimeError, producer.resumeProducing)
|
||||
|
||||
|
||||
def test_cannotPauseProduction(self):
|
||||
"""
|
||||
When the L{NonStreamingProducer} is paused, it raises a
|
||||
L{RuntimeError}.
|
||||
"""
|
||||
consumer = TestConsumer()
|
||||
producer = NonStreamingProducer(consumer)
|
||||
consumer.registerProducer(producer, False)
|
||||
|
||||
# Produce once, just to be safe.
|
||||
producer.resumeProducing()
|
||||
|
||||
self.assertRaises(RuntimeError, producer.pauseProducing)
|
||||
|
||||
|
||||
|
||||
class DeprecationTests(TestCase):
|
||||
"""
|
||||
Deprecations in L{twisted.test.proto_helpers}.
|
||||
"""
|
||||
def helper(self, test, obj):
|
||||
new_path = 'twisted.internet.testing.{}'.format(obj.__name__)
|
||||
warnings = self.flushWarnings(
|
||||
[test])
|
||||
self.assertEqual(DeprecationWarning, warnings[0]['category'])
|
||||
self.assertEqual(1, len(warnings))
|
||||
self.assertIn(new_path, warnings[0]['message'])
|
||||
self.assertIs(obj, namedAny(new_path))
|
||||
|
||||
def test_accumulatingProtocol(self):
|
||||
from twisted.test.proto_helpers import AccumulatingProtocol
|
||||
self.helper(self.test_accumulatingProtocol,
|
||||
AccumulatingProtocol)
|
||||
|
||||
|
||||
def test_lineSendingProtocol(self):
|
||||
from twisted.test.proto_helpers import LineSendingProtocol
|
||||
self.helper(self.test_lineSendingProtocol,
|
||||
LineSendingProtocol)
|
||||
|
||||
|
||||
def test_fakeDatagramTransport(self):
|
||||
from twisted.test.proto_helpers import FakeDatagramTransport
|
||||
self.helper(self.test_fakeDatagramTransport,
|
||||
FakeDatagramTransport)
|
||||
|
||||
|
||||
def test_stringTransport(self):
|
||||
from twisted.test.proto_helpers import StringTransport
|
||||
self.helper(self.test_stringTransport,
|
||||
StringTransport)
|
||||
|
||||
|
||||
def test_stringTransportWithDisconnection(self):
|
||||
from twisted.test.proto_helpers import (
|
||||
StringTransportWithDisconnection)
|
||||
self.helper(self.test_stringTransportWithDisconnection,
|
||||
StringTransportWithDisconnection)
|
||||
|
||||
|
||||
def test_stringIOWithoutClosing(self):
|
||||
from twisted.test.proto_helpers import StringIOWithoutClosing
|
||||
self.helper(self.test_stringIOWithoutClosing,
|
||||
StringIOWithoutClosing)
|
||||
|
||||
|
||||
def test__fakeConnector(self):
|
||||
from twisted.test.proto_helpers import _FakeConnector
|
||||
self.helper(self.test__fakeConnector,
|
||||
_FakeConnector)
|
||||
|
||||
|
||||
def test__fakePort(self):
|
||||
from twisted.test.proto_helpers import _FakePort
|
||||
self.helper(self.test__fakePort,
|
||||
_FakePort)
|
||||
|
||||
|
||||
def test_memoryReactor(self):
|
||||
from twisted.test.proto_helpers import MemoryReactor
|
||||
self.helper(self.test_memoryReactor,
|
||||
MemoryReactor)
|
||||
|
||||
|
||||
def test_memoryReactorClock(self):
|
||||
from twisted.test.proto_helpers import MemoryReactorClock
|
||||
self.helper(self.test_memoryReactorClock,
|
||||
MemoryReactorClock)
|
||||
|
||||
|
||||
def test_raisingMemoryReactor(self):
|
||||
from twisted.test.proto_helpers import RaisingMemoryReactor
|
||||
self.helper(self.test_raisingMemoryReactor,
|
||||
RaisingMemoryReactor)
|
||||
|
||||
|
||||
def test_nonStreamingProducer(self):
|
||||
from twisted.test.proto_helpers import NonStreamingProducer
|
||||
self.helper(self.test_nonStreamingProducer,
|
||||
NonStreamingProducer)
|
||||
|
||||
|
||||
def test_waitUntilAllDisconnected(self):
|
||||
from twisted.test.proto_helpers import (
|
||||
waitUntilAllDisconnected)
|
||||
self.helper(self.test_waitUntilAllDisconnected,
|
||||
waitUntilAllDisconnected)
|
||||
|
||||
|
||||
def test_eventLoggingObserver(self):
|
||||
from twisted.test.proto_helpers import EventLoggingObserver
|
||||
self.helper(self.test_eventLoggingObserver,
|
||||
EventLoggingObserver)
|
||||
@@ -0,0 +1,175 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Tests for L{twisted.internet.serialport}.
|
||||
"""
|
||||
|
||||
import os
|
||||
import shutil
|
||||
import tempfile
|
||||
|
||||
from twisted.trial import unittest
|
||||
from twisted.internet.protocol import Protocol
|
||||
from twisted.python.failure import Failure
|
||||
from twisted.python.runtime import platform
|
||||
from twisted.internet.test.test_serialport import DoNothing
|
||||
|
||||
|
||||
testingForced = 'TWISTED_FORCE_SERIAL_TESTS' in os.environ
|
||||
|
||||
|
||||
try:
|
||||
from twisted.internet import serialport
|
||||
import serial
|
||||
except ImportError:
|
||||
if testingForced:
|
||||
raise
|
||||
|
||||
serialport = None
|
||||
serial = None
|
||||
|
||||
|
||||
if serialport is not None:
|
||||
class RegularFileSerial(serial.Serial):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super(RegularFileSerial, self).__init__(*args, **kwargs)
|
||||
self.captured_args = args
|
||||
self.captured_kwargs = kwargs
|
||||
|
||||
def _reconfigurePort(self):
|
||||
pass
|
||||
|
||||
def _reconfigure_port(self):
|
||||
pass
|
||||
|
||||
|
||||
class RegularFileSerialPort(serialport.SerialPort):
|
||||
_serialFactory = RegularFileSerial
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
cbInQue = kwargs.get('cbInQue')
|
||||
|
||||
if 'cbInQue' in kwargs:
|
||||
del kwargs['cbInQue']
|
||||
|
||||
self.comstat = serial.win32.COMSTAT
|
||||
self.comstat.cbInQue = cbInQue
|
||||
|
||||
super(RegularFileSerialPort, self).__init__(*args, **kwargs)
|
||||
|
||||
def _clearCommError(self):
|
||||
return True, self.comstat
|
||||
|
||||
|
||||
class CollectReceivedProtocol(Protocol):
|
||||
def __init__(self):
|
||||
self.received_data = []
|
||||
|
||||
def dataReceived(self, data):
|
||||
self.received_data.append(data)
|
||||
|
||||
|
||||
class Win32SerialPortTests(unittest.TestCase):
|
||||
"""
|
||||
Minimal testing for Twisted's Win32 serial port support.
|
||||
"""
|
||||
|
||||
if not testingForced:
|
||||
if not platform.isWindows():
|
||||
skip = "This test must run on Windows."
|
||||
|
||||
elif not serialport:
|
||||
skip = "Windows serial port support is not available."
|
||||
|
||||
def setUp(self):
|
||||
# Re-usable protocol and reactor
|
||||
self.protocol = Protocol()
|
||||
self.reactor = DoNothing()
|
||||
|
||||
self.directory = tempfile.mkdtemp()
|
||||
self.path = os.path.join(self.directory, 'fake_serial')
|
||||
|
||||
data = b'1234'
|
||||
with open(self.path, 'wb') as f:
|
||||
f.write(data)
|
||||
|
||||
def tearDown(self):
|
||||
shutil.rmtree(self.directory)
|
||||
|
||||
def test_serialPortDefaultArgs(self):
|
||||
"""
|
||||
Test correct positional and keyword arguments have been
|
||||
passed to the C{serial.Serial} object.
|
||||
"""
|
||||
port = RegularFileSerialPort(self.protocol, self.path, self.reactor)
|
||||
# Validate args
|
||||
self.assertEqual((self.path,), port._serial.captured_args)
|
||||
# Validate kwargs
|
||||
kwargs = port._serial.captured_kwargs
|
||||
self.assertEqual(9600, kwargs["baudrate"])
|
||||
self.assertEqual(serial.EIGHTBITS, kwargs["bytesize"])
|
||||
self.assertEqual(serial.PARITY_NONE, kwargs["parity"])
|
||||
self.assertEqual(serial.STOPBITS_ONE, kwargs["stopbits"])
|
||||
self.assertEqual(0, kwargs["xonxoff"])
|
||||
self.assertEqual(0, kwargs["rtscts"])
|
||||
self.assertEqual(None, kwargs["timeout"])
|
||||
port.connectionLost(Failure(Exception("Cleanup")))
|
||||
|
||||
def test_serialPortInitiallyConnected(self):
|
||||
"""
|
||||
Test the port is connected at initialization time, and
|
||||
C{Protocol.makeConnection} has been called on the desired protocol.
|
||||
"""
|
||||
self.assertEqual(0, self.protocol.connected)
|
||||
|
||||
port = RegularFileSerialPort(self.protocol, self.path, self.reactor)
|
||||
self.assertEqual(1, port.connected)
|
||||
self.assertEqual(1, self.protocol.connected)
|
||||
self.assertEqual(port, self.protocol.transport)
|
||||
port.connectionLost(Failure(Exception("Cleanup")))
|
||||
|
||||
def common_exerciseHandleAccess(self, cbInQue):
|
||||
port = RegularFileSerialPort(
|
||||
protocol=self.protocol,
|
||||
deviceNameOrPortNumber=self.path,
|
||||
reactor=self.reactor,
|
||||
cbInQue=cbInQue,
|
||||
)
|
||||
port.serialReadEvent()
|
||||
port.write(b'')
|
||||
port.write(b'abcd')
|
||||
port.write(b'ABCD')
|
||||
port.serialWriteEvent()
|
||||
port.serialWriteEvent()
|
||||
port.connectionLost(Failure(Exception("Cleanup")))
|
||||
|
||||
# No assertion since the point is simply to make sure that in all cases
|
||||
# the port handle resolves instead of raising an exception.
|
||||
|
||||
def test_exerciseHandleAccess_1(self):
|
||||
self.common_exerciseHandleAccess(cbInQue=False)
|
||||
|
||||
def test_exerciseHandleAccess_2(self):
|
||||
self.common_exerciseHandleAccess(cbInQue=True)
|
||||
|
||||
def common_serialPortReturnsBytes(self, cbInQue):
|
||||
protocol = CollectReceivedProtocol()
|
||||
|
||||
port = RegularFileSerialPort(
|
||||
protocol=protocol,
|
||||
deviceNameOrPortNumber=self.path,
|
||||
reactor=self.reactor,
|
||||
cbInQue=cbInQue,
|
||||
)
|
||||
port.serialReadEvent()
|
||||
self.assertTrue(all(
|
||||
isinstance(d, bytes) for d in protocol.received_data
|
||||
))
|
||||
port.connectionLost(Failure(Exception("Cleanup")))
|
||||
|
||||
def test_serialPortReturnsBytes_1(self):
|
||||
self.common_serialPortReturnsBytes(cbInQue=False)
|
||||
|
||||
def test_serialPortReturnsBytes_2(self):
|
||||
self.common_serialPortReturnsBytes(cbInQue=True)
|
||||
@@ -0,0 +1,78 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
|
||||
"""
|
||||
This module integrates Tkinter with twisted.internet's mainloop.
|
||||
|
||||
Maintainer: Itamar Shtull-Trauring
|
||||
|
||||
To use, do::
|
||||
|
||||
| tksupport.install(rootWidget)
|
||||
|
||||
and then run your reactor as usual - do *not* call Tk's mainloop(),
|
||||
use Twisted's regular mechanism for running the event loop.
|
||||
|
||||
Likewise, to stop your program you will need to stop Twisted's
|
||||
event loop. For example, if you want closing your root widget to
|
||||
stop Twisted::
|
||||
|
||||
| root.protocol('WM_DELETE_WINDOW', reactor.stop)
|
||||
|
||||
When using Aqua Tcl/Tk on macOS the standard Quit menu item in
|
||||
your application might become unresponsive without the additional
|
||||
fix::
|
||||
|
||||
| root.createcommand("::tk::mac::Quit", reactor.stop)
|
||||
|
||||
@see: U{Tcl/TkAqua FAQ for more info<http://wiki.tcl.tk/12987>}
|
||||
"""
|
||||
|
||||
from twisted.internet import task
|
||||
from twisted.python.compat import _PY3
|
||||
|
||||
if _PY3:
|
||||
import tkinter.simpledialog as tkSimpleDialog
|
||||
import tkinter.messagebox as tkMessageBox
|
||||
else:
|
||||
import tkSimpleDialog, tkMessageBox
|
||||
|
||||
|
||||
|
||||
_task = None
|
||||
|
||||
def install(widget, ms=10, reactor=None):
|
||||
"""Install a Tkinter.Tk() object into the reactor."""
|
||||
installTkFunctions()
|
||||
global _task
|
||||
_task = task.LoopingCall(widget.update)
|
||||
_task.start(ms / 1000.0, False)
|
||||
|
||||
def uninstall():
|
||||
"""Remove the root Tk widget from the reactor.
|
||||
|
||||
Call this before destroy()ing the root widget.
|
||||
"""
|
||||
global _task
|
||||
_task.stop()
|
||||
_task = None
|
||||
|
||||
|
||||
def installTkFunctions():
|
||||
import twisted.python.util
|
||||
twisted.python.util.getPassword = getPassword
|
||||
|
||||
|
||||
def getPassword(prompt = '', confirm = 0):
|
||||
while 1:
|
||||
try1 = tkSimpleDialog.askstring('Password Dialog', prompt, show='*')
|
||||
if not confirm:
|
||||
return try1
|
||||
try2 = tkSimpleDialog.askstring('Password Dialog', 'Confirm Password', show='*')
|
||||
if try1 == try2:
|
||||
return try1
|
||||
else:
|
||||
tkMessageBox.showerror('Password Mismatch', 'Passwords did not match, starting over')
|
||||
|
||||
__all__ = ["install", "uninstall"]
|
||||
@@ -0,0 +1,188 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
This module provides wxPython event loop support for Twisted.
|
||||
|
||||
In order to use this support, simply do the following::
|
||||
|
||||
| from twisted.internet import wxreactor
|
||||
| wxreactor.install()
|
||||
|
||||
Then, when your root wxApp has been created::
|
||||
|
||||
| from twisted.internet import reactor
|
||||
| reactor.registerWxApp(yourApp)
|
||||
| reactor.run()
|
||||
|
||||
Then use twisted.internet APIs as usual. Stop the event loop using
|
||||
reactor.stop(), not yourApp.ExitMainLoop().
|
||||
|
||||
IMPORTANT: tests will fail when run under this reactor. This is
|
||||
expected and probably does not reflect on the reactor's ability to run
|
||||
real applications.
|
||||
"""
|
||||
|
||||
try:
|
||||
from queue import Empty, Queue
|
||||
except ImportError:
|
||||
from Queue import Empty, Queue
|
||||
|
||||
try:
|
||||
from wx import PySimpleApp as wxPySimpleApp, CallAfter as wxCallAfter, \
|
||||
Timer as wxTimer
|
||||
except ImportError:
|
||||
# older version of wxPython:
|
||||
from wxPython.wx import wxPySimpleApp, wxCallAfter, wxTimer
|
||||
|
||||
from twisted.python import log, runtime
|
||||
from twisted.internet import _threadedselect
|
||||
|
||||
|
||||
class ProcessEventsTimer(wxTimer):
|
||||
"""
|
||||
Timer that tells wx to process pending events.
|
||||
|
||||
This is necessary on macOS, probably due to a bug in wx, if we want
|
||||
wxCallAfters to be handled when modal dialogs, menus, etc. are open.
|
||||
"""
|
||||
def __init__(self, wxapp):
|
||||
wxTimer.__init__(self)
|
||||
self.wxapp = wxapp
|
||||
|
||||
|
||||
def Notify(self):
|
||||
"""
|
||||
Called repeatedly by wx event loop.
|
||||
"""
|
||||
self.wxapp.ProcessPendingEvents()
|
||||
|
||||
|
||||
|
||||
class WxReactor(_threadedselect.ThreadedSelectReactor):
|
||||
"""
|
||||
wxPython reactor.
|
||||
|
||||
wxPython drives the event loop, select() runs in a thread.
|
||||
"""
|
||||
|
||||
_stopping = False
|
||||
|
||||
def registerWxApp(self, wxapp):
|
||||
"""
|
||||
Register wxApp instance with the reactor.
|
||||
"""
|
||||
self.wxapp = wxapp
|
||||
|
||||
|
||||
def _installSignalHandlersAgain(self):
|
||||
"""
|
||||
wx sometimes removes our own signal handlers, so re-add them.
|
||||
"""
|
||||
try:
|
||||
# make _handleSignals happy:
|
||||
import signal
|
||||
signal.signal(signal.SIGINT, signal.default_int_handler)
|
||||
except ImportError:
|
||||
return
|
||||
self._handleSignals()
|
||||
|
||||
|
||||
def stop(self):
|
||||
"""
|
||||
Stop the reactor.
|
||||
"""
|
||||
if self._stopping:
|
||||
return
|
||||
self._stopping = True
|
||||
_threadedselect.ThreadedSelectReactor.stop(self)
|
||||
|
||||
|
||||
def _runInMainThread(self, f):
|
||||
"""
|
||||
Schedule function to run in main wx/Twisted thread.
|
||||
|
||||
Called by the select() thread.
|
||||
"""
|
||||
if hasattr(self, "wxapp"):
|
||||
wxCallAfter(f)
|
||||
else:
|
||||
# wx shutdown but twisted hasn't
|
||||
self._postQueue.put(f)
|
||||
|
||||
|
||||
def _stopWx(self):
|
||||
"""
|
||||
Stop the wx event loop if it hasn't already been stopped.
|
||||
|
||||
Called during Twisted event loop shutdown.
|
||||
"""
|
||||
if hasattr(self, "wxapp"):
|
||||
self.wxapp.ExitMainLoop()
|
||||
|
||||
|
||||
def run(self, installSignalHandlers=True):
|
||||
"""
|
||||
Start the reactor.
|
||||
"""
|
||||
self._postQueue = Queue()
|
||||
if not hasattr(self, "wxapp"):
|
||||
log.msg("registerWxApp() was not called on reactor, "
|
||||
"registering my own wxApp instance.")
|
||||
self.registerWxApp(wxPySimpleApp())
|
||||
|
||||
# start select() thread:
|
||||
self.interleave(self._runInMainThread,
|
||||
installSignalHandlers=installSignalHandlers)
|
||||
if installSignalHandlers:
|
||||
self.callLater(0, self._installSignalHandlersAgain)
|
||||
|
||||
# add cleanup events:
|
||||
self.addSystemEventTrigger("after", "shutdown", self._stopWx)
|
||||
self.addSystemEventTrigger("after", "shutdown",
|
||||
lambda: self._postQueue.put(None))
|
||||
|
||||
# On macOS, work around wx bug by starting timer to ensure
|
||||
# wxCallAfter calls are always processed. We don't wake up as
|
||||
# often as we could since that uses too much CPU.
|
||||
if runtime.platform.isMacOSX():
|
||||
t = ProcessEventsTimer(self.wxapp)
|
||||
t.Start(2) # wake up every 2ms
|
||||
|
||||
self.wxapp.MainLoop()
|
||||
wxapp = self.wxapp
|
||||
del self.wxapp
|
||||
|
||||
if not self._stopping:
|
||||
# wx event loop exited without reactor.stop() being
|
||||
# called. At this point events from select() thread will
|
||||
# be added to _postQueue, but some may still be waiting
|
||||
# unprocessed in wx, thus the ProcessPendingEvents()
|
||||
# below.
|
||||
self.stop()
|
||||
wxapp.ProcessPendingEvents() # deal with any queued wxCallAfters
|
||||
while 1:
|
||||
try:
|
||||
f = self._postQueue.get(timeout=0.01)
|
||||
except Empty:
|
||||
continue
|
||||
else:
|
||||
if f is None:
|
||||
break
|
||||
try:
|
||||
f()
|
||||
except:
|
||||
log.err()
|
||||
|
||||
|
||||
def install():
|
||||
"""
|
||||
Configure the twisted mainloop to be run inside the wxPython mainloop.
|
||||
"""
|
||||
reactor = WxReactor()
|
||||
from twisted.internet.main import installReactor
|
||||
installReactor(reactor)
|
||||
return reactor
|
||||
|
||||
|
||||
__all__ = ['install']
|
||||
@@ -0,0 +1,59 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
#
|
||||
"""Old method of wxPython support for Twisted.
|
||||
|
||||
twisted.internet.wxreactor is probably a better choice.
|
||||
|
||||
To use::
|
||||
|
||||
| # given a wxApp instance called myWxAppInstance:
|
||||
| from twisted.internet import wxsupport
|
||||
| wxsupport.install(myWxAppInstance)
|
||||
|
||||
Use Twisted's APIs for running and stopping the event loop, don't use
|
||||
wxPython's methods.
|
||||
|
||||
On Windows the Twisted event loop might block when dialogs are open
|
||||
or menus are selected.
|
||||
|
||||
Maintainer: Itamar Shtull-Trauring
|
||||
"""
|
||||
|
||||
import warnings
|
||||
warnings.warn("wxsupport is not fully functional on Windows, wxreactor is better.")
|
||||
|
||||
from twisted.python._oldstyle import _oldStyle
|
||||
from twisted.internet import reactor
|
||||
|
||||
|
||||
|
||||
@_oldStyle
|
||||
class wxRunner:
|
||||
"""Make sure GUI events are handled."""
|
||||
|
||||
def __init__(self, app):
|
||||
self.app = app
|
||||
|
||||
def run(self):
|
||||
"""
|
||||
Execute pending WX events followed by WX idle events and
|
||||
reschedule.
|
||||
"""
|
||||
# run wx events
|
||||
while self.app.Pending():
|
||||
self.app.Dispatch()
|
||||
|
||||
# run wx idle events
|
||||
self.app.ProcessIdle()
|
||||
reactor.callLater(0.02, self.run)
|
||||
|
||||
|
||||
def install(app):
|
||||
"""Install the wxPython support, given a wxApp instance"""
|
||||
runner = wxRunner(app)
|
||||
reactor.callLater(0.02, runner.run)
|
||||
|
||||
|
||||
__all__ = ["install"]
|
||||
@@ -0,0 +1,59 @@
|
||||
# -*- test-case-name: twisted.logger.test.test_buffer -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Log observer that maintains a buffer.
|
||||
"""
|
||||
|
||||
from collections import deque
|
||||
|
||||
from zope.interface import implementer
|
||||
|
||||
from ._observer import ILogObserver
|
||||
|
||||
|
||||
_DEFAULT_BUFFER_MAXIMUM = 64 * 1024
|
||||
|
||||
|
||||
|
||||
@implementer(ILogObserver)
|
||||
class LimitedHistoryLogObserver(object):
|
||||
"""
|
||||
L{ILogObserver} that stores events in a buffer of a fixed size::
|
||||
|
||||
>>> from twisted.logger import LimitedHistoryLogObserver
|
||||
>>> history = LimitedHistoryLogObserver(5)
|
||||
>>> for n in range(10): history({'n': n})
|
||||
...
|
||||
>>> repeats = []
|
||||
>>> history.replayTo(repeats.append)
|
||||
>>> len(repeats)
|
||||
5
|
||||
>>> repeats
|
||||
[{'n': 5}, {'n': 6}, {'n': 7}, {'n': 8}, {'n': 9}]
|
||||
>>>
|
||||
"""
|
||||
|
||||
def __init__(self, size=_DEFAULT_BUFFER_MAXIMUM):
|
||||
"""
|
||||
@param size: The maximum number of events to buffer. If L{None}, the
|
||||
buffer is unbounded.
|
||||
@type size: L{int}
|
||||
"""
|
||||
self._buffer = deque(maxlen=size)
|
||||
|
||||
|
||||
def __call__(self, event):
|
||||
self._buffer.append(event)
|
||||
|
||||
|
||||
def replayTo(self, otherObserver):
|
||||
"""
|
||||
Re-play the buffered events to another log observer.
|
||||
|
||||
@param otherObserver: An observer to replay events to.
|
||||
@type otherObserver: L{ILogObserver}
|
||||
"""
|
||||
for event in self._buffer:
|
||||
otherObserver(event)
|
||||
@@ -0,0 +1,86 @@
|
||||
# -*- test-case-name: twisted.logger.test.test_file -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
File log observer.
|
||||
"""
|
||||
|
||||
from zope.interface import implementer
|
||||
|
||||
from twisted.python.compat import ioType, unicode
|
||||
from ._observer import ILogObserver
|
||||
from ._format import formatTime
|
||||
from ._format import timeFormatRFC3339
|
||||
from ._format import formatEventAsClassicLogText
|
||||
|
||||
|
||||
|
||||
@implementer(ILogObserver)
|
||||
class FileLogObserver(object):
|
||||
"""
|
||||
Log observer that writes to a file-like object.
|
||||
"""
|
||||
def __init__(self, outFile, formatEvent):
|
||||
"""
|
||||
@param outFile: A file-like object. Ideally one should be passed which
|
||||
accepts L{unicode} data. Otherwise, UTF-8 L{bytes} will be used.
|
||||
@type outFile: L{io.IOBase}
|
||||
|
||||
@param formatEvent: A callable that formats an event.
|
||||
@type formatEvent: L{callable} that takes an C{event} argument and
|
||||
returns a formatted event as L{unicode}.
|
||||
"""
|
||||
if ioType(outFile) is not unicode:
|
||||
self._encoding = "utf-8"
|
||||
else:
|
||||
self._encoding = None
|
||||
|
||||
self._outFile = outFile
|
||||
self.formatEvent = formatEvent
|
||||
|
||||
|
||||
def __call__(self, event):
|
||||
"""
|
||||
Write event to file.
|
||||
|
||||
@param event: An event.
|
||||
@type event: L{dict}
|
||||
"""
|
||||
text = self.formatEvent(event)
|
||||
|
||||
if text is None:
|
||||
text = u""
|
||||
|
||||
if self._encoding is not None:
|
||||
text = text.encode(self._encoding)
|
||||
|
||||
if text:
|
||||
self._outFile.write(text)
|
||||
self._outFile.flush()
|
||||
|
||||
|
||||
|
||||
def textFileLogObserver(outFile, timeFormat=timeFormatRFC3339):
|
||||
"""
|
||||
Create a L{FileLogObserver} that emits text to a specified (writable)
|
||||
file-like object.
|
||||
|
||||
@param outFile: A file-like object. Ideally one should be passed which
|
||||
accepts L{unicode} data. Otherwise, UTF-8 L{bytes} will be used.
|
||||
@type outFile: L{io.IOBase}
|
||||
|
||||
@param timeFormat: The format to use when adding timestamp prefixes to
|
||||
logged events. If L{None}, or for events with no C{"log_timestamp"}
|
||||
key, the default timestamp prefix of C{u"-"} is used.
|
||||
@type timeFormat: L{unicode} or L{None}
|
||||
|
||||
@return: A file log observer.
|
||||
@rtype: L{FileLogObserver}
|
||||
"""
|
||||
def formatEvent(event):
|
||||
return formatEventAsClassicLogText(
|
||||
event, formatTime=lambda e: formatTime(e, timeFormat)
|
||||
)
|
||||
|
||||
return FileLogObserver(outFile, formatEvent)
|
||||
@@ -0,0 +1,110 @@
|
||||
# -*- test-case-name: twisted.logger.test.test_levels -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Log levels.
|
||||
"""
|
||||
|
||||
from constantly import NamedConstant, Names
|
||||
|
||||
|
||||
|
||||
class InvalidLogLevelError(Exception):
|
||||
"""
|
||||
Someone tried to use a L{LogLevel} that is unknown to the logging system.
|
||||
"""
|
||||
def __init__(self, level):
|
||||
"""
|
||||
@param level: A log level.
|
||||
@type level: L{LogLevel}
|
||||
"""
|
||||
super(InvalidLogLevelError, self).__init__(str(level))
|
||||
self.level = level
|
||||
|
||||
|
||||
|
||||
class LogLevel(Names):
|
||||
"""
|
||||
Constants describing log levels.
|
||||
|
||||
@cvar debug: Debugging events: Information of use to a developer of the
|
||||
software, not generally of interest to someone running the software
|
||||
unless they are attempting to diagnose a software issue.
|
||||
|
||||
@cvar info: Informational events: Routine information about the status of
|
||||
an application, such as incoming connections, startup of a subsystem,
|
||||
etc.
|
||||
|
||||
@cvar warn: Warning events: Events that may require greater attention than
|
||||
informational events but are not a systemic failure condition, such as
|
||||
authorization failures, bad data from a network client, etc. Such
|
||||
events are of potential interest to system administrators, and should
|
||||
ideally be phrased in such a way, or documented, so as to indicate an
|
||||
action that an administrator might take to mitigate the warning.
|
||||
|
||||
@cvar error: Error conditions: Events indicating a systemic failure, such
|
||||
as programming errors in the form of unhandled exceptions, loss of
|
||||
connectivity to an external system without which no useful work can
|
||||
proceed, such as a database or API endpoint, or resource exhaustion.
|
||||
Similarly to warnings, errors that are related to operational
|
||||
parameters may be actionable to system administrators and should
|
||||
provide references to resources which an administrator might use to
|
||||
resolve them.
|
||||
|
||||
@cvar critical: Critical failures: Errors indicating systemic failure (ie.
|
||||
service outage), data corruption, imminent data loss, etc. which must
|
||||
be handled immediately. This includes errors unanticipated by the
|
||||
software, such as unhandled exceptions, wherein the cause and
|
||||
consequences are unknown.
|
||||
"""
|
||||
|
||||
debug = NamedConstant()
|
||||
info = NamedConstant()
|
||||
warn = NamedConstant()
|
||||
error = NamedConstant()
|
||||
critical = NamedConstant()
|
||||
|
||||
|
||||
@classmethod
|
||||
def levelWithName(cls, name):
|
||||
"""
|
||||
Get the log level with the given name.
|
||||
|
||||
@param name: The name of a log level.
|
||||
@type name: L{str} (native string)
|
||||
|
||||
@return: The L{LogLevel} with the specified C{name}.
|
||||
@rtype: L{LogLevel}
|
||||
|
||||
@raise InvalidLogLevelError: if the C{name} does not name a valid log
|
||||
level.
|
||||
"""
|
||||
try:
|
||||
return cls.lookupByName(name)
|
||||
except ValueError:
|
||||
raise InvalidLogLevelError(name)
|
||||
|
||||
|
||||
@classmethod
|
||||
def _priorityForLevel(cls, level):
|
||||
"""
|
||||
We want log levels to have defined ordering - the order of definition -
|
||||
but they aren't value constants (the only value is the name). This is
|
||||
arguably a bug in Twisted, so this is just a workaround for U{until
|
||||
this is fixed in some way
|
||||
<https://twistedmatrix.com/trac/ticket/6523>}.
|
||||
|
||||
@param level: A log level.
|
||||
@type level: L{LogLevel}
|
||||
|
||||
@return: A numeric index indicating priority (lower is higher level).
|
||||
@rtype: L{int}
|
||||
"""
|
||||
return cls._levelPriorities[level]
|
||||
|
||||
|
||||
LogLevel._levelPriorities = dict(
|
||||
(level, index) for (index, level) in
|
||||
(enumerate(LogLevel.iterconstants()))
|
||||
)
|
||||
@@ -0,0 +1,275 @@
|
||||
# -*- test-case-name: twisted.logger.test.test_logger -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Logger class.
|
||||
"""
|
||||
|
||||
from time import time
|
||||
|
||||
from twisted.python.compat import currentframe
|
||||
from twisted.python.failure import Failure
|
||||
from ._levels import InvalidLogLevelError, LogLevel
|
||||
|
||||
|
||||
|
||||
class Logger(object):
|
||||
"""
|
||||
A L{Logger} emits log messages to an observer. You should instantiate it
|
||||
as a class or module attribute, as documented in L{this module's
|
||||
documentation <twisted.logger>}.
|
||||
|
||||
@type namespace: L{str}
|
||||
@ivar namespace: the namespace for this logger
|
||||
|
||||
@type source: L{object}
|
||||
@ivar source: The object which is emitting events via this logger
|
||||
|
||||
@type: L{ILogObserver}
|
||||
@ivar observer: The observer that this logger will send events to.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _namespaceFromCallingContext():
|
||||
"""
|
||||
Derive a namespace from the module containing the caller's caller.
|
||||
|
||||
@return: the fully qualified python name of a module.
|
||||
@rtype: L{str} (native string)
|
||||
"""
|
||||
try:
|
||||
return currentframe(2).f_globals["__name__"]
|
||||
except KeyError:
|
||||
return "<unknown>"
|
||||
|
||||
|
||||
def __init__(self, namespace=None, source=None, observer=None):
|
||||
"""
|
||||
@param namespace: The namespace for this logger. Uses a dotted
|
||||
notation, as used by python modules. If not L{None}, then the name
|
||||
of the module of the caller is used.
|
||||
@type namespace: L{str} (native string)
|
||||
|
||||
@param source: The object which is emitting events via this
|
||||
logger; this is automatically set on instances of a class
|
||||
if this L{Logger} is an attribute of that class.
|
||||
@type source: L{object}
|
||||
|
||||
@param observer: The observer that this logger will send events to.
|
||||
If L{None}, use the L{global log publisher <globalLogPublisher>}.
|
||||
@type observer: L{ILogObserver}
|
||||
"""
|
||||
if namespace is None:
|
||||
namespace = self._namespaceFromCallingContext()
|
||||
|
||||
self.namespace = namespace
|
||||
self.source = source
|
||||
|
||||
if observer is None:
|
||||
from ._global import globalLogPublisher
|
||||
self.observer = globalLogPublisher
|
||||
else:
|
||||
self.observer = observer
|
||||
|
||||
|
||||
def __get__(self, oself, type=None):
|
||||
"""
|
||||
When used as a descriptor, i.e.::
|
||||
|
||||
# File: athing.py
|
||||
class Something(object):
|
||||
log = Logger()
|
||||
def hello(self):
|
||||
self.log.info("Hello")
|
||||
|
||||
a L{Logger}'s namespace will be set to the name of the class it is
|
||||
declared on. In the above example, the namespace would be
|
||||
C{athing.Something}.
|
||||
|
||||
Additionally, its source will be set to the actual object referring to
|
||||
the L{Logger}. In the above example, C{Something.log.source} would be
|
||||
C{Something}, and C{Something().log.source} would be an instance of
|
||||
C{Something}.
|
||||
"""
|
||||
if oself is None:
|
||||
source = type
|
||||
else:
|
||||
source = oself
|
||||
|
||||
return self.__class__(
|
||||
".".join([type.__module__, type.__name__]),
|
||||
source,
|
||||
observer=self.observer,
|
||||
)
|
||||
|
||||
|
||||
def __repr__(self):
|
||||
return "<%s %r>" % (self.__class__.__name__, self.namespace)
|
||||
|
||||
|
||||
def emit(self, level, format=None, **kwargs):
|
||||
"""
|
||||
Emit a log event to all log observers at the given level.
|
||||
|
||||
@param level: a L{LogLevel}
|
||||
|
||||
@param format: a message format using new-style (PEP 3101)
|
||||
formatting. The logging event (which is a L{dict}) is
|
||||
used to render this format string.
|
||||
|
||||
@param kwargs: additional key/value pairs to include in the event.
|
||||
Note that values which are later mutated may result in
|
||||
non-deterministic behavior from observers that schedule work for
|
||||
later execution.
|
||||
"""
|
||||
if level not in LogLevel.iterconstants():
|
||||
self.failure(
|
||||
"Got invalid log level {invalidLevel!r} in {logger}.emit().",
|
||||
Failure(InvalidLogLevelError(level)),
|
||||
invalidLevel=level,
|
||||
logger=self,
|
||||
)
|
||||
return
|
||||
|
||||
event = kwargs
|
||||
event.update(
|
||||
log_logger=self, log_level=level, log_namespace=self.namespace,
|
||||
log_source=self.source, log_format=format, log_time=time(),
|
||||
)
|
||||
|
||||
if "log_trace" in event:
|
||||
event["log_trace"].append((self, self.observer))
|
||||
|
||||
self.observer(event)
|
||||
|
||||
|
||||
def failure(self, format, failure=None, level=LogLevel.critical, **kwargs):
|
||||
"""
|
||||
Log a failure and emit a traceback.
|
||||
|
||||
For example::
|
||||
|
||||
try:
|
||||
frob(knob)
|
||||
except Exception:
|
||||
log.failure("While frobbing {knob}", knob=knob)
|
||||
|
||||
or::
|
||||
|
||||
d = deferredFrob(knob)
|
||||
d.addErrback(lambda f: log.failure("While frobbing {knob}",
|
||||
f, knob=knob))
|
||||
|
||||
This method is generally meant to capture unexpected exceptions in
|
||||
code; an exception that is caught and handled somehow should be logged,
|
||||
if appropriate, via L{Logger.error} instead. If some unknown exception
|
||||
occurs and your code doesn't know how to handle it, as in the above
|
||||
example, then this method provides a means to describe the failure in
|
||||
nerd-speak. This is done at L{LogLevel.critical} by default, since no
|
||||
corrective guidance can be offered to an user/administrator, and the
|
||||
impact of the condition is unknown.
|
||||
|
||||
@param format: a message format using new-style (PEP 3101) formatting.
|
||||
The logging event (which is a L{dict}) is used to render this
|
||||
format string.
|
||||
|
||||
@param failure: a L{Failure} to log. If L{None}, a L{Failure} is
|
||||
created from the exception in flight.
|
||||
|
||||
@param level: a L{LogLevel} to use.
|
||||
|
||||
@param kwargs: additional key/value pairs to include in the event.
|
||||
Note that values which are later mutated may result in
|
||||
non-deterministic behavior from observers that schedule work for
|
||||
later execution.
|
||||
"""
|
||||
if failure is None:
|
||||
failure = Failure()
|
||||
|
||||
self.emit(level, format, log_failure=failure, **kwargs)
|
||||
|
||||
|
||||
def debug(self, format=None, **kwargs):
|
||||
"""
|
||||
Emit a log event at log level L{LogLevel.debug}.
|
||||
|
||||
@param format: a message format using new-style (PEP 3101) formatting.
|
||||
The logging event (which is a L{dict}) is used to render this
|
||||
format string.
|
||||
|
||||
@param kwargs: additional key/value pairs to include in the event.
|
||||
Note that values which are later mutated may result in
|
||||
non-deterministic behavior from observers that schedule work for
|
||||
later execution.
|
||||
"""
|
||||
self.emit(LogLevel.debug, format, **kwargs)
|
||||
|
||||
|
||||
def info(self, format=None, **kwargs):
|
||||
"""
|
||||
Emit a log event at log level L{LogLevel.info}.
|
||||
|
||||
@param format: a message format using new-style (PEP 3101) formatting.
|
||||
The logging event (which is a L{dict}) is used to render this
|
||||
format string.
|
||||
|
||||
@param kwargs: additional key/value pairs to include in the event.
|
||||
Note that values which are later mutated may result in
|
||||
non-deterministic behavior from observers that schedule work for
|
||||
later execution.
|
||||
"""
|
||||
self.emit(LogLevel.info, format, **kwargs)
|
||||
|
||||
|
||||
def warn(self, format=None, **kwargs):
|
||||
"""
|
||||
Emit a log event at log level L{LogLevel.warn}.
|
||||
|
||||
@param format: a message format using new-style (PEP 3101) formatting.
|
||||
The logging event (which is a L{dict}) is used to render this
|
||||
format string.
|
||||
|
||||
@param kwargs: additional key/value pairs to include in the event.
|
||||
Note that values which are later mutated may result in
|
||||
non-deterministic behavior from observers that schedule work for
|
||||
later execution.
|
||||
"""
|
||||
self.emit(LogLevel.warn, format, **kwargs)
|
||||
|
||||
|
||||
def error(self, format=None, **kwargs):
|
||||
"""
|
||||
Emit a log event at log level L{LogLevel.error}.
|
||||
|
||||
@param format: a message format using new-style (PEP 3101) formatting.
|
||||
The logging event (which is a L{dict}) is used to render this
|
||||
format string.
|
||||
|
||||
@param kwargs: additional key/value pairs to include in the event.
|
||||
Note that values which are later mutated may result in
|
||||
non-deterministic behavior from observers that schedule work for
|
||||
later execution.
|
||||
"""
|
||||
self.emit(LogLevel.error, format, **kwargs)
|
||||
|
||||
|
||||
def critical(self, format=None, **kwargs):
|
||||
"""
|
||||
Emit a log event at log level L{LogLevel.critical}.
|
||||
|
||||
@param format: a message format using new-style (PEP 3101) formatting.
|
||||
The logging event (which is a L{dict}) is used to render this
|
||||
format string.
|
||||
|
||||
@param kwargs: additional key/value pairs to include in the event.
|
||||
Note that values which are later mutated may result in
|
||||
non-deterministic behavior from observers that schedule work for
|
||||
later execution.
|
||||
"""
|
||||
self.emit(LogLevel.critical, format, **kwargs)
|
||||
|
||||
|
||||
|
||||
_log = Logger()
|
||||
_loggerFor = lambda obj:_log.__get__(obj, obj.__class__)
|
||||
@@ -0,0 +1,62 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Test cases for L{twisted.logger._buffer}.
|
||||
"""
|
||||
|
||||
from zope.interface.verify import verifyObject, BrokenMethodImplementation
|
||||
|
||||
from twisted.trial import unittest
|
||||
|
||||
from .._observer import ILogObserver
|
||||
from .._buffer import LimitedHistoryLogObserver
|
||||
|
||||
|
||||
|
||||
class LimitedHistoryLogObserverTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for L{LimitedHistoryLogObserver}.
|
||||
"""
|
||||
|
||||
def test_interface(self):
|
||||
"""
|
||||
L{LimitedHistoryLogObserver} provides L{ILogObserver}.
|
||||
"""
|
||||
observer = LimitedHistoryLogObserver(0)
|
||||
try:
|
||||
verifyObject(ILogObserver, observer)
|
||||
except BrokenMethodImplementation as e:
|
||||
self.fail(e)
|
||||
|
||||
|
||||
def test_order(self):
|
||||
"""
|
||||
L{LimitedHistoryLogObserver} saves history in the order it is received.
|
||||
"""
|
||||
size = 4
|
||||
events = [dict(n=n) for n in range(size//2)]
|
||||
observer = LimitedHistoryLogObserver(size)
|
||||
|
||||
for event in events:
|
||||
observer(event)
|
||||
|
||||
outEvents = []
|
||||
observer.replayTo(outEvents.append)
|
||||
self.assertEqual(events, outEvents)
|
||||
|
||||
|
||||
def test_limit(self):
|
||||
"""
|
||||
When more events than a L{LimitedHistoryLogObserver}'s maximum size are
|
||||
buffered, older events will be dropped.
|
||||
"""
|
||||
size = 4
|
||||
events = [dict(n=n) for n in range(size*2)]
|
||||
observer = LimitedHistoryLogObserver(size)
|
||||
|
||||
for event in events:
|
||||
observer(event)
|
||||
outEvents = []
|
||||
observer.replayTo(outEvents.append)
|
||||
self.assertEqual(events[-size:], outEvents)
|
||||
@@ -0,0 +1,39 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Test cases for L{twisted.logger._capture}.
|
||||
"""
|
||||
|
||||
from twisted.logger import Logger, LogLevel
|
||||
from twisted.trial.unittest import TestCase
|
||||
|
||||
from .._capture import capturedLogs
|
||||
|
||||
|
||||
|
||||
class LogCaptureTests(TestCase):
|
||||
"""
|
||||
Tests for L{LogCaptureTests}.
|
||||
"""
|
||||
|
||||
log = Logger()
|
||||
|
||||
|
||||
def test_capture(self):
|
||||
"""
|
||||
Events logged within context are captured.
|
||||
"""
|
||||
foo = object()
|
||||
|
||||
with capturedLogs() as captured:
|
||||
self.log.debug("Capture this, please", foo=foo)
|
||||
self.log.info("Capture this too, please", foo=foo)
|
||||
|
||||
self.assertTrue(len(captured) == 2)
|
||||
self.assertEqual(captured[0]["log_format"], "Capture this, please")
|
||||
self.assertEqual(captured[0]["log_level"], LogLevel.debug)
|
||||
self.assertEqual(captured[0]["foo"], foo)
|
||||
self.assertEqual(captured[1]["log_format"], "Capture this too, please")
|
||||
self.assertEqual(captured[1]["log_level"], LogLevel.info)
|
||||
self.assertEqual(captured[1]["foo"], foo)
|
||||
@@ -0,0 +1,199 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Test cases for L{twisted.logger._file}.
|
||||
"""
|
||||
|
||||
from io import StringIO
|
||||
|
||||
from zope.interface.verify import verifyObject, BrokenMethodImplementation
|
||||
|
||||
from twisted.trial.unittest import TestCase
|
||||
|
||||
from twisted.python.failure import Failure
|
||||
from twisted.python.compat import unicode
|
||||
from .._observer import ILogObserver
|
||||
from .._file import FileLogObserver
|
||||
from .._file import textFileLogObserver
|
||||
|
||||
|
||||
|
||||
class FileLogObserverTests(TestCase):
|
||||
"""
|
||||
Tests for L{FileLogObserver}.
|
||||
"""
|
||||
|
||||
def test_interface(self):
|
||||
"""
|
||||
L{FileLogObserver} is an L{ILogObserver}.
|
||||
"""
|
||||
with StringIO() as fileHandle:
|
||||
observer = FileLogObserver(fileHandle, lambda e: unicode(e))
|
||||
try:
|
||||
verifyObject(ILogObserver, observer)
|
||||
except BrokenMethodImplementation as e:
|
||||
self.fail(e)
|
||||
|
||||
|
||||
def test_observeWrites(self):
|
||||
"""
|
||||
L{FileLogObserver} writes to the given file when it observes events.
|
||||
"""
|
||||
with StringIO() as fileHandle:
|
||||
observer = FileLogObserver(fileHandle, lambda e: unicode(e))
|
||||
event = dict(x=1)
|
||||
observer(event)
|
||||
self.assertEqual(fileHandle.getvalue(), unicode(event))
|
||||
|
||||
|
||||
def _test_observeWrites(self, what, count):
|
||||
"""
|
||||
Verify that observer performs an expected number of writes when the
|
||||
formatter returns a given value.
|
||||
|
||||
@param what: the value for the formatter to return.
|
||||
@type what: L{unicode}
|
||||
|
||||
@param count: the expected number of writes.
|
||||
@type count: L{int}
|
||||
"""
|
||||
with DummyFile() as fileHandle:
|
||||
observer = FileLogObserver(fileHandle, lambda e: what)
|
||||
event = dict(x=1)
|
||||
observer(event)
|
||||
self.assertEqual(fileHandle.writes, count)
|
||||
|
||||
|
||||
def test_observeWritesNone(self):
|
||||
"""
|
||||
L{FileLogObserver} does not write to the given file when it observes
|
||||
events and C{formatEvent} returns L{None}.
|
||||
"""
|
||||
self._test_observeWrites(None, 0)
|
||||
|
||||
|
||||
def test_observeWritesEmpty(self):
|
||||
"""
|
||||
L{FileLogObserver} does not write to the given file when it observes
|
||||
events and C{formatEvent} returns C{u""}.
|
||||
"""
|
||||
self._test_observeWrites(u"", 0)
|
||||
|
||||
|
||||
def test_observeFlushes(self):
|
||||
"""
|
||||
L{FileLogObserver} calles C{flush()} on the output file when it
|
||||
observes an event.
|
||||
"""
|
||||
with DummyFile() as fileHandle:
|
||||
observer = FileLogObserver(fileHandle, lambda e: unicode(e))
|
||||
event = dict(x=1)
|
||||
observer(event)
|
||||
self.assertEqual(fileHandle.flushes, 1)
|
||||
|
||||
|
||||
class TextFileLogObserverTests(TestCase):
|
||||
"""
|
||||
Tests for L{textFileLogObserver}.
|
||||
"""
|
||||
|
||||
def test_returnsFileLogObserver(self):
|
||||
"""
|
||||
L{textFileLogObserver} returns a L{FileLogObserver}.
|
||||
"""
|
||||
with StringIO() as fileHandle:
|
||||
observer = textFileLogObserver(fileHandle)
|
||||
self.assertIsInstance(observer, FileLogObserver)
|
||||
|
||||
|
||||
def test_outFile(self):
|
||||
"""
|
||||
Returned L{FileLogObserver} has the correct outFile.
|
||||
"""
|
||||
with StringIO() as fileHandle:
|
||||
observer = textFileLogObserver(fileHandle)
|
||||
self.assertIs(observer._outFile, fileHandle)
|
||||
|
||||
|
||||
def test_timeFormat(self):
|
||||
"""
|
||||
Returned L{FileLogObserver} has the correct outFile.
|
||||
"""
|
||||
with StringIO() as fileHandle:
|
||||
observer = textFileLogObserver(fileHandle, timeFormat=u"%f")
|
||||
observer(dict(log_format=u"XYZZY", log_time=112345.6))
|
||||
self.assertEqual(fileHandle.getvalue(), u"600000 [-#-] XYZZY\n")
|
||||
|
||||
|
||||
def test_observeFailure(self):
|
||||
"""
|
||||
If the C{"log_failure"} key exists in an event, the observer appends
|
||||
the failure's traceback to the output.
|
||||
"""
|
||||
with StringIO() as fileHandle:
|
||||
observer = textFileLogObserver(fileHandle)
|
||||
|
||||
try:
|
||||
1 / 0
|
||||
except ZeroDivisionError:
|
||||
failure = Failure()
|
||||
|
||||
event = dict(log_failure=failure)
|
||||
observer(event)
|
||||
output = fileHandle.getvalue()
|
||||
self.assertTrue(output.split("\n")[1].startswith("\tTraceback "),
|
||||
msg=repr(output))
|
||||
|
||||
|
||||
def test_observeFailureThatRaisesInGetTraceback(self):
|
||||
"""
|
||||
If the C{"log_failure"} key exists in an event, and contains an object
|
||||
that raises when you call its C{getTraceback()}, then the observer
|
||||
appends a message noting the problem, instead of raising.
|
||||
"""
|
||||
with StringIO() as fileHandle:
|
||||
observer = textFileLogObserver(fileHandle)
|
||||
event = dict(log_failure=object()) # object has no getTraceback()
|
||||
observer(event)
|
||||
output = fileHandle.getvalue()
|
||||
expected = (
|
||||
"(UNABLE TO OBTAIN TRACEBACK FROM EVENT)"
|
||||
)
|
||||
self.assertIn(expected, output)
|
||||
|
||||
|
||||
|
||||
class DummyFile(object):
|
||||
"""
|
||||
File that counts writes and flushes.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.writes = 0
|
||||
self.flushes = 0
|
||||
|
||||
|
||||
def write(self, data):
|
||||
"""
|
||||
Write data.
|
||||
|
||||
@param data: data
|
||||
@type data: L{unicode} or L{bytes}
|
||||
"""
|
||||
self.writes += 1
|
||||
|
||||
|
||||
def flush(self):
|
||||
"""
|
||||
Flush buffers.
|
||||
"""
|
||||
self.flushes += 1
|
||||
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
|
||||
def __exit__(self, exc_type, exc_value, traceback):
|
||||
pass
|
||||
@@ -0,0 +1,372 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Test cases for L{twisted.logger._global}.
|
||||
"""
|
||||
|
||||
from __future__ import print_function
|
||||
|
||||
import io
|
||||
|
||||
from twisted.trial import unittest
|
||||
|
||||
from .._file import textFileLogObserver
|
||||
from .._observer import LogPublisher
|
||||
from .._logger import Logger
|
||||
from .._global import LogBeginner
|
||||
from .._global import MORE_THAN_ONCE_WARNING
|
||||
from .._levels import LogLevel
|
||||
from ..test.test_stdlib import nextLine
|
||||
from twisted.python.failure import Failure
|
||||
|
||||
|
||||
|
||||
def compareEvents(test, actualEvents, expectedEvents):
|
||||
"""
|
||||
Compare two sequences of log events, examining only the the keys which are
|
||||
present in both.
|
||||
|
||||
@param test: a test case doing the comparison
|
||||
@type test: L{unittest.TestCase}
|
||||
|
||||
@param actualEvents: A list of log events that were emitted by a logger.
|
||||
@type actualEvents: L{list} of L{dict}
|
||||
|
||||
@param expectedEvents: A list of log events that were expected by a test.
|
||||
@type expected: L{list} of L{dict}
|
||||
"""
|
||||
if len(actualEvents) != len(expectedEvents):
|
||||
test.assertEqual(actualEvents, expectedEvents)
|
||||
allMergedKeys = set()
|
||||
|
||||
for event in expectedEvents:
|
||||
allMergedKeys |= set(event.keys())
|
||||
|
||||
def simplify(event):
|
||||
copy = event.copy()
|
||||
for key in event.keys():
|
||||
if key not in allMergedKeys:
|
||||
copy.pop(key)
|
||||
return copy
|
||||
|
||||
simplifiedActual = [simplify(event) for event in actualEvents]
|
||||
test.assertEqual(simplifiedActual, expectedEvents)
|
||||
|
||||
|
||||
|
||||
class LogBeginnerTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for L{LogBeginner}.
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
self.publisher = LogPublisher()
|
||||
self.errorStream = io.StringIO()
|
||||
|
||||
class NotSys(object):
|
||||
stdout = object()
|
||||
stderr = object()
|
||||
|
||||
class NotWarnings(object):
|
||||
def __init__(self):
|
||||
self.warnings = []
|
||||
|
||||
def showwarning(
|
||||
self, message, category, filename, lineno,
|
||||
file=None, line=None
|
||||
):
|
||||
"""
|
||||
Emulate warnings.showwarning.
|
||||
|
||||
@param message: A warning message to emit.
|
||||
@type message: L{str}
|
||||
|
||||
@param category: A warning category to associate with
|
||||
C{message}.
|
||||
@type category: L{warnings.Warning}
|
||||
|
||||
@param filename: A file name for the source code file issuing
|
||||
the warning.
|
||||
@type warning: L{str}
|
||||
|
||||
@param lineno: A line number in the source file where the
|
||||
warning was issued.
|
||||
@type lineno: L{int}
|
||||
|
||||
@param file: A file to write the warning message to. If
|
||||
L{None}, write to L{sys.stderr}.
|
||||
@type file: file-like object
|
||||
|
||||
@param line: A line of source code to include with the warning
|
||||
message. If L{None}, attempt to read the line from
|
||||
C{filename} and C{lineno}.
|
||||
@type line: L{str}
|
||||
"""
|
||||
self.warnings.append(
|
||||
(message, category, filename, lineno, file, line)
|
||||
)
|
||||
|
||||
self.sysModule = NotSys()
|
||||
self.warningsModule = NotWarnings()
|
||||
self.beginner = LogBeginner(
|
||||
self.publisher, self.errorStream, self.sysModule,
|
||||
self.warningsModule
|
||||
)
|
||||
|
||||
|
||||
def test_beginLoggingToAddObservers(self):
|
||||
"""
|
||||
Test that C{beginLoggingTo()} adds observers.
|
||||
"""
|
||||
event = dict(foo=1, bar=2)
|
||||
|
||||
events1 = []
|
||||
events2 = []
|
||||
|
||||
o1 = lambda e: events1.append(e)
|
||||
o2 = lambda e: events2.append(e)
|
||||
|
||||
self.beginner.beginLoggingTo((o1, o2))
|
||||
self.publisher(event)
|
||||
|
||||
self.assertEqual([event], events1)
|
||||
self.assertEqual([event], events2)
|
||||
|
||||
|
||||
def test_beginLoggingToBufferedEvents(self):
|
||||
"""
|
||||
Test that events are buffered until C{beginLoggingTo()} is
|
||||
called.
|
||||
"""
|
||||
event = dict(foo=1, bar=2)
|
||||
|
||||
events1 = []
|
||||
events2 = []
|
||||
|
||||
o1 = lambda e: events1.append(e)
|
||||
o2 = lambda e: events2.append(e)
|
||||
|
||||
self.publisher(event) # Before beginLoggingTo; this is buffered
|
||||
self.beginner.beginLoggingTo((o1, o2))
|
||||
|
||||
self.assertEqual([event], events1)
|
||||
self.assertEqual([event], events2)
|
||||
|
||||
|
||||
def _bufferLimitTest(self, limit, beginner):
|
||||
"""
|
||||
Verify that when more than C{limit} events are logged to L{LogBeginner},
|
||||
only the last C{limit} are replayed by L{LogBeginner.beginLoggingTo}.
|
||||
|
||||
@param limit: The maximum number of events the log beginner should
|
||||
buffer.
|
||||
@type limit: L{int}
|
||||
|
||||
@param beginner: The L{LogBeginner} against which to verify.
|
||||
@type beginner: L{LogBeginner}
|
||||
|
||||
@raise: C{self.failureException} if the wrong events are replayed by
|
||||
C{beginner}.
|
||||
|
||||
@return: L{None}
|
||||
"""
|
||||
for count in range(limit + 1):
|
||||
self.publisher(dict(count=count))
|
||||
events = []
|
||||
beginner.beginLoggingTo([events.append])
|
||||
self.assertEqual(
|
||||
list(range(1, limit + 1)),
|
||||
list(event["count"] for event in events),
|
||||
)
|
||||
|
||||
|
||||
def test_defaultBufferLimit(self):
|
||||
"""
|
||||
Up to C{LogBeginner._DEFAULT_BUFFER_SIZE} log events are buffered for
|
||||
replay by L{LogBeginner.beginLoggingTo}.
|
||||
"""
|
||||
limit = LogBeginner._DEFAULT_BUFFER_SIZE
|
||||
self._bufferLimitTest(limit, self.beginner)
|
||||
|
||||
|
||||
def test_overrideBufferLimit(self):
|
||||
"""
|
||||
The size of the L{LogBeginner} event buffer can be overridden with the
|
||||
C{initialBufferSize} initilizer argument.
|
||||
"""
|
||||
limit = 3
|
||||
beginner = LogBeginner(
|
||||
self.publisher, self.errorStream, self.sysModule,
|
||||
self.warningsModule, initialBufferSize=limit,
|
||||
)
|
||||
self._bufferLimitTest(limit, beginner)
|
||||
|
||||
|
||||
def test_beginLoggingToTwice(self):
|
||||
"""
|
||||
When invoked twice, L{LogBeginner.beginLoggingTo} will emit a log
|
||||
message warning the user that they previously began logging, and add
|
||||
the new log observers.
|
||||
"""
|
||||
events1 = []
|
||||
events2 = []
|
||||
fileHandle = io.StringIO()
|
||||
textObserver = textFileLogObserver(fileHandle)
|
||||
self.publisher(dict(event="prebuffer"))
|
||||
firstFilename, firstLine = nextLine()
|
||||
self.beginner.beginLoggingTo([events1.append, textObserver])
|
||||
self.publisher(dict(event="postbuffer"))
|
||||
secondFilename, secondLine = nextLine()
|
||||
self.beginner.beginLoggingTo([events2.append, textObserver])
|
||||
self.publisher(dict(event="postwarn"))
|
||||
warning = dict(
|
||||
log_format=MORE_THAN_ONCE_WARNING,
|
||||
log_level=LogLevel.warn,
|
||||
fileNow=secondFilename, lineNow=secondLine,
|
||||
fileThen=firstFilename, lineThen=firstLine
|
||||
)
|
||||
|
||||
compareEvents(
|
||||
self, events1,
|
||||
[
|
||||
dict(event="prebuffer"),
|
||||
dict(event="postbuffer"),
|
||||
warning,
|
||||
dict(event="postwarn")
|
||||
]
|
||||
)
|
||||
compareEvents(self, events2, [warning, dict(event="postwarn")])
|
||||
|
||||
output = fileHandle.getvalue()
|
||||
self.assertIn('<{0}:{1}>'.format(firstFilename, firstLine),
|
||||
output)
|
||||
self.assertIn('<{0}:{1}>'.format(secondFilename, secondLine),
|
||||
output)
|
||||
|
||||
|
||||
def test_criticalLogging(self):
|
||||
"""
|
||||
Critical messages will be written as text to the error stream.
|
||||
"""
|
||||
log = Logger(observer=self.publisher)
|
||||
log.info("ignore this")
|
||||
log.critical("a critical {message}", message="message")
|
||||
self.assertEqual(self.errorStream.getvalue(), u"a critical message\n")
|
||||
|
||||
|
||||
def test_criticalLoggingStops(self):
|
||||
"""
|
||||
Once logging has begun with C{beginLoggingTo}, critical messages are no
|
||||
longer written to the output stream.
|
||||
"""
|
||||
log = Logger(observer=self.publisher)
|
||||
self.beginner.beginLoggingTo(())
|
||||
log.critical("another critical message")
|
||||
self.assertEqual(self.errorStream.getvalue(), u"")
|
||||
|
||||
|
||||
def test_beginLoggingToRedirectStandardIO(self):
|
||||
"""
|
||||
L{LogBeginner.beginLoggingTo} will re-direct the standard output and
|
||||
error streams by setting the C{stdio} and C{stderr} attributes on its
|
||||
sys module object.
|
||||
"""
|
||||
x = []
|
||||
self.beginner.beginLoggingTo([x.append])
|
||||
print("Hello, world.", file=self.sysModule.stdout)
|
||||
compareEvents(
|
||||
self, x, [dict(log_namespace="stdout", log_io="Hello, world.")]
|
||||
)
|
||||
del x[:]
|
||||
print("Error, world.", file=self.sysModule.stderr)
|
||||
compareEvents(
|
||||
self, x, [dict(log_namespace="stderr", log_io="Error, world.")]
|
||||
)
|
||||
|
||||
|
||||
def test_beginLoggingToDontRedirect(self):
|
||||
"""
|
||||
L{LogBeginner.beginLoggingTo} will leave the existing stdout/stderr in
|
||||
place if it has been told not to replace them.
|
||||
"""
|
||||
oldOut = self.sysModule.stdout
|
||||
oldErr = self.sysModule.stderr
|
||||
self.beginner.beginLoggingTo((), redirectStandardIO=False)
|
||||
self.assertIs(self.sysModule.stdout, oldOut)
|
||||
self.assertIs(self.sysModule.stderr, oldErr)
|
||||
|
||||
|
||||
def test_beginLoggingToPreservesEncoding(self):
|
||||
"""
|
||||
When L{LogBeginner.beginLoggingTo} redirects stdout/stderr streams, the
|
||||
replacement streams will preserve the encoding of the replaced streams,
|
||||
to minimally disrupt any application relying on a specific encoding.
|
||||
"""
|
||||
|
||||
weird = io.TextIOWrapper(io.BytesIO(), "shift-JIS")
|
||||
weirderr = io.TextIOWrapper(io.BytesIO(), "big5")
|
||||
|
||||
self.sysModule.stdout = weird
|
||||
self.sysModule.stderr = weirderr
|
||||
|
||||
x = []
|
||||
self.beginner.beginLoggingTo([x.append])
|
||||
self.assertEqual(self.sysModule.stdout.encoding, "shift-JIS")
|
||||
self.assertEqual(self.sysModule.stderr.encoding, "big5")
|
||||
|
||||
self.sysModule.stdout.write(b"\x97\x9B\n")
|
||||
self.sysModule.stderr.write(b"\xBC\xFC\n")
|
||||
compareEvents(
|
||||
self, x, [dict(log_io=u"\u674e"), dict(log_io=u"\u7469")]
|
||||
)
|
||||
|
||||
|
||||
def test_warningsModule(self):
|
||||
"""
|
||||
L{LogBeginner.beginLoggingTo} will redirect the warnings of its
|
||||
warnings module into the logging system.
|
||||
"""
|
||||
self.warningsModule.showwarning(
|
||||
"a message", DeprecationWarning, __file__, 1
|
||||
)
|
||||
x = []
|
||||
self.beginner.beginLoggingTo([x.append])
|
||||
self.warningsModule.showwarning(
|
||||
"another message", DeprecationWarning, __file__, 2
|
||||
)
|
||||
f = io.StringIO()
|
||||
self.warningsModule.showwarning(
|
||||
"yet another", DeprecationWarning, __file__, 3, file=f
|
||||
)
|
||||
self.assertEqual(
|
||||
self.warningsModule.warnings,
|
||||
[
|
||||
("a message", DeprecationWarning, __file__, 1, None, None),
|
||||
("yet another", DeprecationWarning, __file__, 3, f, None),
|
||||
]
|
||||
)
|
||||
compareEvents(
|
||||
self, x,
|
||||
[dict(
|
||||
warning="another message",
|
||||
category=(
|
||||
DeprecationWarning.__module__ + "." +
|
||||
DeprecationWarning.__name__
|
||||
),
|
||||
filename=__file__, lineno=2,
|
||||
)]
|
||||
)
|
||||
|
||||
|
||||
def test_failuresAppendTracebacks(self):
|
||||
"""
|
||||
The string resulting from a logged failure contains a traceback.
|
||||
"""
|
||||
f = Failure(Exception("this is not the behavior you are looking for"))
|
||||
log = Logger(observer=self.publisher)
|
||||
log.failure('a failure', failure=f)
|
||||
msg = self.errorStream.getvalue()
|
||||
self.assertIn('a failure', msg)
|
||||
self.assertIn('this is not the behavior you are looking for', msg)
|
||||
self.assertIn('Traceback', msg)
|
||||
@@ -0,0 +1,38 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Test cases for L{twisted.logger._levels}.
|
||||
"""
|
||||
|
||||
from twisted.trial import unittest
|
||||
|
||||
from .._levels import InvalidLogLevelError
|
||||
from .._levels import LogLevel
|
||||
|
||||
|
||||
|
||||
class LogLevelTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for L{LogLevel}.
|
||||
"""
|
||||
|
||||
def test_levelWithName(self):
|
||||
"""
|
||||
Look up log level by name.
|
||||
"""
|
||||
for level in LogLevel.iterconstants():
|
||||
self.assertIs(LogLevel.levelWithName(level.name), level)
|
||||
|
||||
|
||||
def test_levelWithInvalidName(self):
|
||||
"""
|
||||
You can't make up log level names.
|
||||
"""
|
||||
bogus = "*bogus*"
|
||||
try:
|
||||
LogLevel.levelWithName(bogus)
|
||||
except InvalidLogLevelError as e:
|
||||
self.assertIs(e.level, bogus)
|
||||
else:
|
||||
self.fail("Expected InvalidLogLevelError.")
|
||||
@@ -0,0 +1,304 @@
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Test cases for L{twisted.logger._format}.
|
||||
"""
|
||||
|
||||
import sys
|
||||
from io import BytesIO, TextIOWrapper
|
||||
import logging as py_logging
|
||||
from inspect import getsourcefile
|
||||
|
||||
from zope.interface.verify import verifyObject, BrokenMethodImplementation
|
||||
|
||||
from twisted.python.compat import _PY3, currentframe
|
||||
from twisted.python.failure import Failure
|
||||
from twisted.trial import unittest
|
||||
|
||||
from .._levels import LogLevel
|
||||
from .._observer import ILogObserver
|
||||
from .._stdlib import STDLibLogObserver
|
||||
|
||||
|
||||
def nextLine():
|
||||
"""
|
||||
Retrive the file name and line number immediately after where this function
|
||||
is called.
|
||||
|
||||
@return: the file name and line number
|
||||
@rtype: 2-L{tuple} of L{str}, L{int}
|
||||
"""
|
||||
caller = currentframe(1)
|
||||
return (getsourcefile(sys.modules[caller.f_globals['__name__']]),
|
||||
caller.f_lineno + 1)
|
||||
|
||||
|
||||
|
||||
class STDLibLogObserverTests(unittest.TestCase):
|
||||
"""
|
||||
Tests for L{STDLibLogObserver}.
|
||||
"""
|
||||
|
||||
def test_interface(self):
|
||||
"""
|
||||
L{STDLibLogObserver} is an L{ILogObserver}.
|
||||
"""
|
||||
observer = STDLibLogObserver()
|
||||
try:
|
||||
verifyObject(ILogObserver, observer)
|
||||
except BrokenMethodImplementation as e:
|
||||
self.fail(e)
|
||||
|
||||
|
||||
def py_logger(self):
|
||||
"""
|
||||
Create a logging object we can use to test with.
|
||||
|
||||
@return: a stdlib-style logger
|
||||
@rtype: L{StdlibLoggingContainer}
|
||||
"""
|
||||
logger = StdlibLoggingContainer()
|
||||
self.addCleanup(logger.close)
|
||||
return logger
|
||||
|
||||
|
||||
def logEvent(self, *events):
|
||||
"""
|
||||
Send one or more events to Python's logging module, and
|
||||
capture the emitted L{logging.LogRecord}s and output stream as
|
||||
a string.
|
||||
|
||||
@param events: events
|
||||
@type events: L{tuple} of L{dict}
|
||||
|
||||
@return: a tuple: (records, output)
|
||||
@rtype: 2-tuple of (L{list} of L{logging.LogRecord}, L{bytes}.)
|
||||
"""
|
||||
pl = self.py_logger()
|
||||
observer = STDLibLogObserver(
|
||||
# Add 1 to default stack depth to skip *this* frame, since
|
||||
# tests will want to know about their own frames.
|
||||
stackDepth=STDLibLogObserver.defaultStackDepth + 1
|
||||
)
|
||||
for event in events:
|
||||
observer(event)
|
||||
return pl.bufferedHandler.records, pl.outputAsText()
|
||||
|
||||
|
||||
def test_name(self):
|
||||
"""
|
||||
Logger name.
|
||||
"""
|
||||
records, output = self.logEvent({})
|
||||
|
||||
self.assertEqual(len(records), 1)
|
||||
self.assertEqual(records[0].name, "twisted")
|
||||
|
||||
|
||||
def test_levels(self):
|
||||
"""
|
||||
Log levels.
|
||||
"""
|
||||
levelMapping = {
|
||||
None: py_logging.INFO, # Default
|
||||
LogLevel.debug: py_logging.DEBUG,
|
||||
LogLevel.info: py_logging.INFO,
|
||||
LogLevel.warn: py_logging.WARNING,
|
||||
LogLevel.error: py_logging.ERROR,
|
||||
LogLevel.critical: py_logging.CRITICAL,
|
||||
}
|
||||
|
||||
# Build a set of events for each log level
|
||||
events = []
|
||||
for level, pyLevel in levelMapping.items():
|
||||
event = {}
|
||||
|
||||
# Set the log level on the event, except for default
|
||||
if level is not None:
|
||||
event["log_level"] = level
|
||||
|
||||
# Remember the Python log level we expect to see for this
|
||||
# event (as an int)
|
||||
event["py_levelno"] = int(pyLevel)
|
||||
|
||||
events.append(event)
|
||||
|
||||
records, output = self.logEvent(*events)
|
||||
self.assertEqual(len(records), len(levelMapping))
|
||||
|
||||
# Check that each event has the correct level
|
||||
for i in range(len(records)):
|
||||
self.assertEqual(records[i].levelno, events[i]["py_levelno"])
|
||||
|
||||
|
||||
def test_callerInfo(self):
|
||||
"""
|
||||
C{pathname}, C{lineno}, C{exc_info}, C{func} is set properly on
|
||||
records.
|
||||
"""
|
||||
filename, logLine = nextLine()
|
||||
records, output = self.logEvent({})
|
||||
|
||||
self.assertEqual(len(records), 1)
|
||||
self.assertEqual(records[0].pathname, filename)
|
||||
self.assertEqual(records[0].lineno, logLine)
|
||||
self.assertIsNone(records[0].exc_info)
|
||||
|
||||
# Attribute "func" is missing from record, which is weird because it's
|
||||
# documented.
|
||||
# self.assertEqual(records[0].func, "test_callerInfo")
|
||||
|
||||
|
||||
def test_basicFormat(self):
|
||||
"""
|
||||
Basic formattable event passes the format along correctly.
|
||||
"""
|
||||
event = dict(log_format="Hello, {who}!", who="dude")
|
||||
records, output = self.logEvent(event)
|
||||
|
||||
self.assertEqual(len(records), 1)
|
||||
self.assertEqual(str(records[0].msg), u"Hello, dude!")
|
||||
self.assertEqual(records[0].args, ())
|
||||
|
||||
|
||||
def test_basicFormatRendered(self):
|
||||
"""
|
||||
Basic formattable event renders correctly.
|
||||
"""
|
||||
event = dict(log_format="Hello, {who}!", who="dude")
|
||||
records, output = self.logEvent(event)
|
||||
|
||||
self.assertEqual(len(records), 1)
|
||||
self.assertTrue(output.endswith(u":Hello, dude!\n"),
|
||||
repr(output))
|
||||
|
||||
|
||||
def test_noFormat(self):
|
||||
"""
|
||||
Event with no format.
|
||||
"""
|
||||
records, output = self.logEvent({})
|
||||
|
||||
self.assertEqual(len(records), 1)
|
||||
self.assertEqual(str(records[0].msg), "")
|
||||
|
||||
|
||||
def test_failure(self):
|
||||
"""
|
||||
An event with a failure logs the failure details as well.
|
||||
"""
|
||||
def failing_func():
|
||||
1 / 0
|
||||
try:
|
||||
failing_func()
|
||||
except ZeroDivisionError:
|
||||
failure = Failure()
|
||||
event = dict(log_format='Hi mom', who='me', log_failure=failure)
|
||||
records, output = self.logEvent(event)
|
||||
self.assertEqual(len(records), 1)
|
||||
self.assertIn(u'Hi mom', output)
|
||||
self.assertIn(u'in failing_func', output)
|
||||
self.assertIn(u'ZeroDivisionError', output)
|
||||
|
||||
|
||||
def test_cleanedFailure(self):
|
||||
"""
|
||||
A cleaned Failure object has a fake traceback object; make sure that
|
||||
logging such a failure still results in the exception details being
|
||||
logged.
|
||||
"""
|
||||
def failing_func():
|
||||
1 / 0
|
||||
try:
|
||||
failing_func()
|
||||
except ZeroDivisionError:
|
||||
failure = Failure()
|
||||
failure.cleanFailure()
|
||||
event = dict(log_format='Hi mom', who='me', log_failure=failure)
|
||||
records, output = self.logEvent(event)
|
||||
self.assertEqual(len(records), 1)
|
||||
self.assertIn(u'Hi mom', output)
|
||||
self.assertIn(u'in failing_func', output)
|
||||
self.assertIn(u'ZeroDivisionError', output)
|
||||
|
||||
|
||||
|
||||
class StdlibLoggingContainer(object):
|
||||
"""
|
||||
Continer for a test configuration of stdlib logging objects.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.rootLogger = py_logging.getLogger("")
|
||||
|
||||
self.originalLevel = self.rootLogger.getEffectiveLevel()
|
||||
self.rootLogger.setLevel(py_logging.DEBUG)
|
||||
|
||||
self.bufferedHandler = BufferedHandler()
|
||||
self.rootLogger.addHandler(self.bufferedHandler)
|
||||
|
||||
self.streamHandler, self.output = handlerAndBytesIO()
|
||||
self.rootLogger.addHandler(self.streamHandler)
|
||||
|
||||
|
||||
def close(self):
|
||||
"""
|
||||
Close the logger.
|
||||
"""
|
||||
self.rootLogger.setLevel(self.originalLevel)
|
||||
self.rootLogger.removeHandler(self.bufferedHandler)
|
||||
self.rootLogger.removeHandler(self.streamHandler)
|
||||
self.streamHandler.close()
|
||||
self.output.close()
|
||||
|
||||
|
||||
def outputAsText(self):
|
||||
"""
|
||||
Get the output to the underlying stream as text.
|
||||
|
||||
@return: the output text
|
||||
@rtype: L{unicode}
|
||||
"""
|
||||
return self.output.getvalue().decode("utf-8")
|
||||
|
||||
|
||||
|
||||
def handlerAndBytesIO():
|
||||
"""
|
||||
Construct a 2-tuple of C{(StreamHandler, BytesIO)} for testing interaction
|
||||
with the 'logging' module.
|
||||
|
||||
@return: handler and io object
|
||||
@rtype: tuple of L{StreamHandler} and L{io.BytesIO}
|
||||
"""
|
||||
output = BytesIO()
|
||||
stream = output
|
||||
template = py_logging.BASIC_FORMAT
|
||||
if _PY3:
|
||||
stream = TextIOWrapper(output, encoding="utf-8", newline="\n")
|
||||
formatter = py_logging.Formatter(template)
|
||||
handler = py_logging.StreamHandler(stream)
|
||||
handler.setFormatter(formatter)
|
||||
return handler, output
|
||||
|
||||
|
||||
|
||||
class BufferedHandler(py_logging.Handler):
|
||||
"""
|
||||
A L{py_logging.Handler} that remembers all logged records in a list.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
"""
|
||||
Initialize this L{BufferedHandler}.
|
||||
"""
|
||||
py_logging.Handler.__init__(self)
|
||||
self.records = []
|
||||
|
||||
|
||||
def emit(self, record):
|
||||
"""
|
||||
Remember the record.
|
||||
"""
|
||||
self.records.append(record)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,404 @@
|
||||
# -*- test-case-name: twisted.mail.test.test_mail -*-
|
||||
# Copyright (c) Twisted Matrix Laboratories.
|
||||
# See LICENSE for details.
|
||||
|
||||
"""
|
||||
Mail protocol support.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
from twisted.mail import pop3
|
||||
from twisted.mail import smtp
|
||||
from twisted.internet import protocol
|
||||
from twisted.internet import defer
|
||||
from twisted.copyright import longversion
|
||||
from twisted.python import log
|
||||
|
||||
from twisted.cred.credentials import CramMD5Credentials, UsernamePassword
|
||||
from twisted.cred.error import UnauthorizedLogin
|
||||
|
||||
from twisted.mail import relay
|
||||
|
||||
from zope.interface import implementer
|
||||
|
||||
|
||||
|
||||
@implementer(smtp.IMessageDelivery)
|
||||
class DomainDeliveryBase:
|
||||
"""
|
||||
A base class for message delivery using the domains of a mail service.
|
||||
|
||||
@ivar service: See L{__init__}
|
||||
@ivar user: See L{__init__}
|
||||
@ivar host: See L{__init__}
|
||||
|
||||
@type protocolName: L{bytes}
|
||||
@ivar protocolName: The protocol being used to deliver the mail.
|
||||
Sub-classes should set this appropriately.
|
||||
"""
|
||||
service = None
|
||||
protocolName = None
|
||||
|
||||
def __init__(self, service, user, host=smtp.DNSNAME):
|
||||
"""
|
||||
@type service: L{MailService}
|
||||
@param service: A mail service.
|
||||
|
||||
@type user: L{bytes} or L{None}
|
||||
@param user: The authenticated SMTP user.
|
||||
|
||||
@type host: L{bytes}
|
||||
@param host: The hostname.
|
||||
"""
|
||||
self.service = service
|
||||
self.user = user
|
||||
self.host = host
|
||||
|
||||
|
||||
def receivedHeader(self, helo, origin, recipients):
|
||||
"""
|
||||
Generate a received header string for a message.
|
||||
|
||||
@type helo: 2-L{tuple} of (L{bytes}, L{bytes})
|
||||
@param helo: The client's identity as sent in the HELO command and its
|
||||
IP address.
|
||||
|
||||
@type origin: L{Address}
|
||||
@param origin: The origination address of the message.
|
||||
|
||||
@type recipients: L{list} of L{User}
|
||||
@param recipients: The destination addresses for the message.
|
||||
|
||||
@rtype: L{bytes}
|
||||
@return: A received header string.
|
||||
"""
|
||||
authStr = heloStr = b""
|
||||
if self.user:
|
||||
authStr = b" auth=" + self.user.encode('xtext')
|
||||
if helo[0]:
|
||||
heloStr = b" helo=" + helo[0]
|
||||
fromUser = (b"from " + helo[0] + b" ([" + helo[1] + b"]" +
|
||||
heloStr + authStr)
|
||||
by = (b"by " + self.host + b" with " + self.protocolName +
|
||||
b" (" + longversion.encode("ascii") + b")")
|
||||
forUser = (b"for <" + b' '.join(map(bytes, recipients)) + b"> " +
|
||||
smtp.rfc822date())
|
||||
return (b"Received: " + fromUser + b"\n\t" + by +
|
||||
b"\n\t" + forUser)
|
||||
|
||||
|
||||
def validateTo(self, user):
|
||||
"""
|
||||
Validate the address for which a message is destined.
|
||||
|
||||
@type user: L{User}
|
||||
@param user: The destination address.
|
||||
|
||||
@rtype: L{Deferred <defer.Deferred>} which successfully fires with
|
||||
no-argument callable which returns L{IMessage <smtp.IMessage>}
|
||||
provider.
|
||||
@return: A deferred which successfully fires with a no-argument
|
||||
callable which returns a message receiver for the destination.
|
||||
|
||||
@raise SMTPBadRcpt: When messages cannot be accepted for the
|
||||
destination address.
|
||||
"""
|
||||
# XXX - Yick. This needs cleaning up.
|
||||
if self.user and self.service.queue:
|
||||
d = self.service.domains.get(user.dest.domain, None)
|
||||
if d is None:
|
||||
d = relay.DomainQueuer(self.service, True)
|
||||
else:
|
||||
d = self.service.domains[user.dest.domain]
|
||||
return defer.maybeDeferred(d.exists, user)
|
||||
|
||||
|
||||
def validateFrom(self, helo, origin):
|
||||
"""
|
||||
Validate the address from which a message originates.
|
||||
|
||||
@type helo: 2-L{tuple} of (L{bytes}, L{bytes})
|
||||
@param helo: The client's identity as sent in the HELO command and its
|
||||
IP address.
|
||||
|
||||
@type origin: L{Address}
|
||||
@param origin: The origination address of the message.
|
||||
|
||||
@rtype: L{Address}
|
||||
@return: The origination address.
|
||||
|
||||
@raise SMTPBadSender: When messages cannot be accepted from the
|
||||
origination address.
|
||||
"""
|
||||
if not helo:
|
||||
raise smtp.SMTPBadSender(origin, 503,
|
||||
"Who are you? Say HELO first.")
|
||||
if origin.local != b'' and origin.domain == b'':
|
||||
raise smtp.SMTPBadSender(origin, 501,
|
||||
"Sender address must contain domain.")
|
||||
return origin
|
||||
|
||||
|
||||
|
||||
class SMTPDomainDelivery(DomainDeliveryBase):
|
||||
"""
|
||||
A domain delivery base class for use in an SMTP server.
|
||||
"""
|
||||
protocolName = b'smtp'
|
||||
|
||||
|
||||
|
||||
class ESMTPDomainDelivery(DomainDeliveryBase):
|
||||
"""
|
||||
A domain delivery base class for use in an ESMTP server.
|
||||
"""
|
||||
protocolName = b'esmtp'
|
||||
|
||||
|
||||
|
||||
class SMTPFactory(smtp.SMTPFactory):
|
||||
"""
|
||||
An SMTP server protocol factory.
|
||||
|
||||
@ivar service: See L{__init__}
|
||||
@ivar portal: See L{__init__}
|
||||
|
||||
@type protocol: no-argument callable which returns a L{Protocol
|
||||
<protocol.Protocol>} subclass
|
||||
@ivar protocol: A callable which creates a protocol. The default value is
|
||||
L{SMTP}.
|
||||
"""
|
||||
protocol = smtp.SMTP
|
||||
portal = None
|
||||
|
||||
def __init__(self, service, portal = None):
|
||||
"""
|
||||
@type service: L{MailService}
|
||||
@param service: An email service.
|
||||
|
||||
@type portal: L{Portal <twisted.cred.portal.Portal>} or
|
||||
L{None}
|
||||
@param portal: A portal to use for authentication.
|
||||
"""
|
||||
smtp.SMTPFactory.__init__(self)
|
||||
self.service = service
|
||||
self.portal = portal
|
||||
|
||||
|
||||
def buildProtocol(self, addr):
|
||||
"""
|
||||
Create an instance of an SMTP server protocol.
|
||||
|
||||
@type addr: L{IAddress <twisted.internet.interfaces.IAddress>} provider
|
||||
@param addr: The address of the SMTP client.
|
||||
|
||||
@rtype: L{SMTP}
|
||||
@return: An SMTP protocol.
|
||||
"""
|
||||
log.msg('Connection from %s' % (addr,))
|
||||
p = smtp.SMTPFactory.buildProtocol(self, addr)
|
||||
p.service = self.service
|
||||
p.portal = self.portal
|
||||
return p
|
||||
|
||||
|
||||
|
||||
class ESMTPFactory(SMTPFactory):
|
||||
"""
|
||||
An ESMTP server protocol factory.
|
||||
|
||||
@type protocol: no-argument callable which returns a L{Protocol
|
||||
<protocol.Protocol>} subclass
|
||||
@ivar protocol: A callable which creates a protocol. The default value is
|
||||
L{ESMTP}.
|
||||
|
||||
@type context: L{IOpenSSLContextFactory
|
||||
<twisted.internet.interfaces.IOpenSSLContextFactory>} or L{None}
|
||||
@ivar context: A factory to generate contexts to be used in negotiating
|
||||
encrypted communication.
|
||||
|
||||
@type challengers: L{dict} mapping L{bytes} to no-argument callable which
|
||||
returns L{ICredentials <twisted.cred.credentials.ICredentials>}
|
||||
subclass provider.
|
||||
@ivar challengers: A mapping of acceptable authorization mechanism to
|
||||
callable which creates credentials to use for authentication.
|
||||
"""
|
||||
protocol = smtp.ESMTP
|
||||
context = None
|
||||
|
||||
def __init__(self, *args):
|
||||
"""
|
||||
@param args: Arguments for L{SMTPFactory.__init__}
|
||||
|
||||
@see: L{SMTPFactory.__init__}
|
||||
"""
|
||||
SMTPFactory.__init__(self, *args)
|
||||
self.challengers = {
|
||||
b'CRAM-MD5': CramMD5Credentials
|
||||
}
|
||||
|
||||
|
||||
def buildProtocol(self, addr):
|
||||
"""
|
||||
Create an instance of an ESMTP server protocol.
|
||||
|
||||
@type addr: L{IAddress <twisted.internet.interfaces.IAddress>} provider
|
||||
@param addr: The address of the ESMTP client.
|
||||
|
||||
@rtype: L{ESMTP}
|
||||
@return: An ESMTP protocol.
|
||||
"""
|
||||
p = SMTPFactory.buildProtocol(self, addr)
|
||||
p.challengers = self.challengers
|
||||
p.ctx = self.context
|
||||
return p
|
||||
|
||||
|
||||
|
||||
class VirtualPOP3(pop3.POP3):
|
||||
"""
|
||||
A virtual hosting POP3 server.
|
||||
|
||||
@type service: L{MailService}
|
||||
@ivar service: The email service that created this server. This must be
|
||||
set by the service.
|
||||
|
||||
@type domainSpecifier: L{bytes}
|
||||
@ivar domainSpecifier: The character to use to split an email address into
|
||||
local-part and domain. The default is '@'.
|
||||
"""
|
||||
service = None
|
||||
|
||||
domainSpecifier = b'@' # Gaagh! I hate POP3. No standardized way
|
||||
# to indicate user@host. '@' doesn't work
|
||||
# with NS, e.g.
|
||||
|
||||
def authenticateUserAPOP(self, user, digest):
|
||||
"""
|
||||
Perform APOP authentication.
|
||||
|
||||
Override the default lookup scheme to allow virtual domains.
|
||||
|
||||
@type user: L{bytes}
|
||||
@param user: The name of the user attempting to log in.
|
||||
|
||||
@type digest: L{bytes}
|
||||
@param digest: The challenge response.
|
||||
|
||||
@rtype: L{Deferred} which successfully results in 3-L{tuple} of
|
||||
(L{IMailbox <pop3.IMailbox>}, L{IMailbox <pop3.IMailbox>}
|
||||
provider, no-argument callable)
|
||||
@return: A deferred which fires when authentication is complete.
|
||||
If successful, it returns an L{IMailbox <pop3.IMailbox>} interface,
|
||||
a mailbox and a logout function. If authentication fails, the
|
||||
deferred fails with an L{UnauthorizedLogin
|
||||
<twisted.cred.error.UnauthorizedLogin>} error.
|
||||
"""
|
||||
user, domain = self.lookupDomain(user)
|
||||
try:
|
||||
portal = self.service.lookupPortal(domain)
|
||||
except KeyError:
|
||||
return defer.fail(UnauthorizedLogin())
|
||||
else:
|
||||
return portal.login(
|
||||
pop3.APOPCredentials(self.magic, user, digest),
|
||||
None,
|
||||
pop3.IMailbox
|
||||
)
|
||||
|
||||
|
||||
def authenticateUserPASS(self, user, password):
|
||||
"""
|
||||
Perform authentication for a username/password login.
|
||||
|
||||
Override the default lookup scheme to allow virtual domains.
|
||||
|
||||
@type user: L{bytes}
|
||||
@param user: The name of the user attempting to log in.
|
||||
|
||||
@type password: L{bytes}
|
||||
@param password: The password to authenticate with.
|
||||
|
||||
@rtype: L{Deferred} which successfully results in 3-L{tuple} of
|
||||
(L{IMailbox <pop3.IMailbox>}, L{IMailbox <pop3.IMailbox>}
|
||||
provider, no-argument callable)
|
||||
@return: A deferred which fires when authentication is complete.
|
||||
If successful, it returns an L{IMailbox <pop3.IMailbox>} interface,
|
||||
a mailbox and a logout function. If authentication fails, the
|
||||
deferred fails with an L{UnauthorizedLogin
|
||||
<twisted.cred.error.UnauthorizedLogin>} error.
|
||||
"""
|
||||
user, domain = self.lookupDomain(user)
|
||||
try:
|
||||
portal = self.service.lookupPortal(domain)
|
||||
except KeyError:
|
||||
return defer.fail(UnauthorizedLogin())
|
||||
else:
|
||||
return portal.login(
|
||||
UsernamePassword(user, password),
|
||||
None,
|
||||
pop3.IMailbox
|
||||
)
|
||||
|
||||
|
||||
def lookupDomain(self, user):
|
||||
"""
|
||||
Check whether a domain is among the virtual domains supported by the
|
||||
mail service.
|
||||
|
||||
@type user: L{bytes}
|
||||
@param user: An email address.
|
||||
|
||||
@rtype: 2-L{tuple} of (L{bytes}, L{bytes})
|
||||
@return: The local part and the domain part of the email address if the
|
||||
domain is supported.
|
||||
|
||||
@raise POP3Error: When the domain is not supported by the mail service.
|
||||
"""
|
||||
try:
|
||||
user, domain = user.split(self.domainSpecifier, 1)
|
||||
except ValueError:
|
||||
domain = b''
|
||||
if domain not in self.service.domains:
|
||||
raise pop3.POP3Error(
|
||||
"no such domain {}".format(domain.decode("utf-8")))
|
||||
return user, domain
|
||||
|
||||
|
||||
|
||||
class POP3Factory(protocol.ServerFactory):
|
||||
"""
|
||||
A POP3 server protocol factory.
|
||||
|
||||
@ivar service: See L{__init__}
|
||||
|
||||
@type protocol: no-argument callable which returns a L{Protocol
|
||||
<protocol.Protocol>} subclass
|
||||
@ivar protocol: A callable which creates a protocol. The default value is
|
||||
L{VirtualPOP3}.
|
||||
"""
|
||||
protocol = VirtualPOP3
|
||||
service = None
|
||||
|
||||
def __init__(self, service):
|
||||
"""
|
||||
@type service: L{MailService}
|
||||
@param service: An email service.
|
||||
"""
|
||||
self.service = service
|
||||
|
||||
|
||||
def buildProtocol(self, addr):
|
||||
"""
|
||||
Create an instance of a POP3 server protocol.
|
||||
|
||||
@type addr: L{IAddress <twisted.internet.interfaces.IAddress>} provider
|
||||
@param addr: The address of the POP3 client.
|
||||
|
||||
@rtype: L{POP3}
|
||||
@return: A POP3 protocol.
|
||||
"""
|
||||
p = protocol.ServerFactory.buildProtocol(self, addr)
|
||||
p.service = self.service
|
||||
return p
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user