This commit is contained in:
2024-12-17 14:36:15 -08:00
parent b2dbf46d28
commit 06d106de53
17731 changed files with 3037186 additions and 144 deletions

View File

@@ -0,0 +1,6 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Distributed trial test runner tests.
"""

View File

@@ -0,0 +1,192 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Hamcrest matchers useful throughout the test suite.
"""
__all__ = [
"matches_result",
"HasSum",
"IsSequenceOf",
]
from typing import Any, List, Sequence, Tuple, TypeVar
from hamcrest import (
contains_exactly,
contains_string,
equal_to,
has_length,
has_properties,
instance_of,
)
from hamcrest.core.base_matcher import BaseMatcher
from hamcrest.core.core.allof import AllOf
from hamcrest.core.description import Description
from hamcrest.core.matcher import Matcher
from typing_extensions import Protocol
from twisted.python.failure import Failure
T = TypeVar("T")
class Semigroup(Protocol[T]):
"""
A type with an associative binary operator.
Common examples of a semigroup are integers with addition and strings with
concatenation.
"""
def __add__(self, other: T) -> T:
"""
This must be associative: a + (b + c) == (a + b) + c
"""
S = TypeVar("S", bound=Semigroup[Any])
def matches_result(
successes: Matcher[Any] = equal_to(0),
errors: Matcher[Any] = has_length(0),
failures: Matcher[Any] = has_length(0),
skips: Matcher[Any] = has_length(0),
expectedFailures: Matcher[Any] = has_length(0),
unexpectedSuccesses: Matcher[Any] = has_length(0),
) -> Matcher[Any]:
"""
Match a L{TestCase} instances with matching attributes.
"""
return has_properties(
{
"successes": successes,
"errors": errors,
"failures": failures,
"skips": skips,
"expectedFailures": expectedFailures,
"unexpectedSuccesses": unexpectedSuccesses,
}
)
class HasSum(BaseMatcher[Sequence[S]]):
"""
Match a sequence the elements of which sum to a value matched by
another matcher.
:ivar sumMatcher: The matcher which must match the sum.
:ivar zero: The zero value for the matched type.
"""
def __init__(self, sumMatcher: Matcher[S], zero: S) -> None:
self.sumMatcher = sumMatcher
self.zero = zero
def _sum(self, sequence: Sequence[S]) -> S:
if not sequence:
return self.zero
result = self.zero
for elem in sequence:
result = result + elem
return result
def _matches(self, item: Sequence[S]) -> bool:
"""
Determine whether the sum of the sequence is matched.
"""
s = self._sum(item)
return self.sumMatcher.matches(s)
def describe_mismatch(self, item: Sequence[S], description: Description) -> None:
"""
Describe the mismatch.
"""
s = self._sum(item)
description.append_description_of(self)
self.sumMatcher.describe_mismatch(s, description)
return None
def describe_to(self, description: Description) -> None:
"""
Describe this matcher for error messages.
"""
description.append_text("a sequence with sum ")
description.append_description_of(self.sumMatcher)
description.append_text(", ")
class IsSequenceOf(BaseMatcher[Sequence[T]]):
"""
Match a sequence where every element is matched by another matcher.
:ivar elementMatcher: The matcher which must match every element of the
sequence.
"""
def __init__(self, elementMatcher: Matcher[T]) -> None:
self.elementMatcher = elementMatcher
def _matches(self, item: Sequence[T]) -> bool:
"""
Determine whether every element of the sequence is matched.
"""
for elem in item:
if not self.elementMatcher.matches(elem):
return False
return True
def describe_mismatch(self, item: Sequence[T], description: Description) -> None:
"""
Describe the mismatch.
"""
for idx, elem in enumerate(item):
if not self.elementMatcher.matches(elem):
description.append_description_of(self)
description.append_text(f"not sequence with element #{idx} {elem!r}")
def describe_to(self, description: Description) -> None:
"""
Describe this matcher for error messages.
"""
description.append_text("a sequence containing only ")
description.append_description_of(self.elementMatcher)
description.append_text(", ")
def isFailure(**properties: Matcher[object]) -> Matcher[object]:
"""
Match an instance of L{Failure} with matching attributes.
"""
return AllOf(
instance_of(Failure),
has_properties(**properties),
)
def similarFrame(
functionName: str, fileName: str
) -> Matcher[Sequence[Tuple[str, str, int, List[object], List[object]]]]:
"""
Match a tuple representation of a frame like those used by
L{twisted.python.failure.Failure}.
"""
# The frames depend on exact layout of the source
# code in files and on the filesystem so we won't
# bother being very precise here. Just verify we
# see some distinctive fragments.
#
# In particular, the last frame should be a tuple like
#
# (functionName, fileName, someint, [], [])
return contains_exactly(
equal_to(functionName),
contains_string(fileName), # type: ignore[arg-type]
instance_of(int), # type: ignore[arg-type]
# Unfortunately Failure makes them sometimes tuples, sometimes
# dict_items.
has_length(0), # type: ignore[arg-type]
has_length(0), # type: ignore[arg-type]
)

View File

@@ -0,0 +1,61 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Tests for L{twisted.trial._dist.distreporter}.
"""
from io import StringIO
from twisted.python.failure import Failure
from twisted.trial._dist.distreporter import DistReporter
from twisted.trial.reporter import TreeReporter
from twisted.trial.unittest import TestCase
class DistReporterTests(TestCase):
"""
Tests for L{DistReporter}.
"""
def setUp(self) -> None:
self.stream = StringIO()
self.distReporter = DistReporter(TreeReporter(self.stream))
self.test = TestCase()
def test_startSuccessStop(self) -> None:
"""
Success output only gets sent to the stream after the test has stopped.
"""
self.distReporter.startTest(self.test)
self.assertEqual(self.stream.getvalue(), "")
self.distReporter.addSuccess(self.test)
self.assertEqual(self.stream.getvalue(), "")
self.distReporter.stopTest(self.test)
self.assertNotEqual(self.stream.getvalue(), "")
def test_startErrorStop(self) -> None:
"""
Error output only gets sent to the stream after the test has stopped.
"""
self.distReporter.startTest(self.test)
self.assertEqual(self.stream.getvalue(), "")
self.distReporter.addError(self.test, Failure(Exception("error")))
self.assertEqual(self.stream.getvalue(), "")
self.distReporter.stopTest(self.test)
self.assertNotEqual(self.stream.getvalue(), "")
def test_forwardedMethods(self) -> None:
"""
Calling methods of L{DistReporter} add calls to the running queue of
the test.
"""
self.distReporter.startTest(self.test)
self.distReporter.addFailure(self.test, Failure(Exception("foo")))
self.distReporter.addError(self.test, Failure(Exception("bar")))
self.distReporter.addSkip(self.test, "egg")
self.distReporter.addUnexpectedSuccess(self.test, "spam")
self.distReporter.addExpectedFailure(
self.test, Failure(Exception("err")), "foo"
)
self.assertEqual(len(self.distReporter.running[self.test.id()]), 6)

View File

@@ -0,0 +1,861 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Tests for L{twisted.trial._dist.disttrial}.
"""
import os
import sys
from functools import partial
from io import StringIO
from os.path import sep
from typing import Callable, List, Set
from unittest import TestCase as PyUnitTestCase
from zope.interface import implementer, verify
from attrs import Factory, assoc, define, field
from hamcrest import (
assert_that,
contains,
ends_with,
equal_to,
has_length,
none,
starts_with,
)
from hamcrest.core.core.allof import AllOf
from hypothesis import given
from hypothesis.strategies import booleans, sampled_from
from twisted.internet import interfaces
from twisted.internet.base import ReactorBase
from twisted.internet.defer import CancelledError, Deferred, succeed
from twisted.internet.error import ProcessDone
from twisted.internet.protocol import ProcessProtocol, Protocol
from twisted.internet.test.modulehelpers import AlternateReactor
from twisted.internet.testing import MemoryReactorClock
from twisted.python.failure import Failure
from twisted.python.filepath import FilePath
from twisted.python.lockfile import FilesystemLock
from twisted.trial._dist import _WORKER_AMP_STDIN
from twisted.trial._dist.distreporter import DistReporter
from twisted.trial._dist.disttrial import DistTrialRunner, WorkerPool, WorkerPoolConfig
from twisted.trial._dist.functional import (
countingCalls,
discardResult,
fromOptional,
iterateWhile,
sequence,
)
from twisted.trial._dist.worker import LocalWorker, RunResult, Worker, WorkerAction
from twisted.trial.reporter import (
Reporter,
TestResult,
TreeReporter,
UncleanWarningsReporterWrapper,
)
from twisted.trial.runner import ErrorHolder, TrialSuite
from twisted.trial.unittest import SynchronousTestCase, TestCase
from ...test import erroneous, sample
from .matchers import matches_result
@define
class FakeTransport:
"""
A simple fake process transport.
"""
_closed: Set[int] = field(default=Factory(set))
def writeToChild(self, fd, data):
"""
Ignore write calls.
"""
def closeChildFD(self, fd):
"""
Mark one of the child descriptors as closed.
"""
self._closed.add(fd)
@implementer(interfaces.IReactorProcess)
class CountingReactor(MemoryReactorClock):
"""
A fake reactor that counts the calls to L{IReactorCore.run},
L{IReactorCore.stop}, and L{IReactorProcess.spawnProcess}.
"""
spawnCount = 0
stopCount = 0
runCount = 0
def __init__(self, workers):
MemoryReactorClock.__init__(self)
self._workers = workers
def spawnProcess(
self,
workerProto,
executable,
args=(),
env={},
path=None,
uid=None,
gid=None,
usePTY=0,
childFDs=None,
):
"""
See L{IReactorProcess.spawnProcess}.
@param workerProto: See L{IReactorProcess.spawnProcess}.
@param args: See L{IReactorProcess.spawnProcess}.
@param kwargs: See L{IReactorProcess.spawnProcess}.
"""
self._workers.append(workerProto)
workerProto.makeConnection(FakeTransport())
self.spawnCount += 1
def stop(self):
"""
See L{IReactorCore.stop}.
"""
MemoryReactorClock.stop(self)
# TODO: implementing this more comprehensively in MemoryReactor would
# be nice, this is rather hard-coded to disttrial's current
# implementation.
if "before" in self.triggers:
self.triggers["before"]["shutdown"][0][0]()
self.stopCount += 1
def run(self):
"""
See L{IReactorCore.run}.
"""
self.runCount += 1
# The same as IReactorCore.run, except no stop.
self.running = True
self.hasRun = True
for f, args, kwargs in self.whenRunningHooks:
f(*args, **kwargs)
self.stop()
# do not count internal 'stop' against trial-initiated .stop() count
self.stopCount -= 1
class CountingReactorTests(SynchronousTestCase):
"""
Tests for L{CountingReactor}.
"""
def setUp(self):
self.workers = []
self.reactor = CountingReactor(self.workers)
def test_providesIReactorProcess(self):
"""
L{CountingReactor} instances provide L{IReactorProcess}.
"""
verify.verifyObject(interfaces.IReactorProcess, self.reactor)
def test_spawnProcess(self):
"""
The process protocol for a spawned process is connected to a
transport and appended onto the provided C{workers} list, and
the reactor's C{spawnCount} increased.
"""
self.assertFalse(self.reactor.spawnCount)
proto = Protocol()
for count in [1, 2]:
self.reactor.spawnProcess(proto, sys.executable, args=[sys.executable])
self.assertTrue(proto.transport)
self.assertEqual(self.workers, [proto] * count)
self.assertEqual(self.reactor.spawnCount, count)
def test_stop(self):
"""
Stopping the reactor increments its C{stopCount}
"""
self.assertFalse(self.reactor.stopCount)
for count in [1, 2]:
self.reactor.stop()
self.assertEqual(self.reactor.stopCount, count)
def test_run(self):
"""
Running the reactor increments its C{runCount}, does not imply
C{stop}, and calls L{IReactorCore.callWhenRunning} hooks.
"""
self.assertFalse(self.reactor.runCount)
whenRunningCalls = []
self.reactor.callWhenRunning(whenRunningCalls.append, None)
for count in [1, 2]:
self.reactor.run()
self.assertEqual(self.reactor.runCount, count)
self.assertEqual(self.reactor.stopCount, 0)
self.assertEqual(len(whenRunningCalls), count)
class WorkerPoolTests(TestCase):
"""
Tests for L{WorkerPool}.
"""
def setUp(self):
self.parent = FilePath(self.mktemp())
self.workingDirectory = self.parent.child("_trial_temp")
self.config = WorkerPoolConfig(
numWorkers=4,
workingDirectory=self.workingDirectory,
workerArguments=[],
logFile="out.log",
)
self.pool = WorkerPool(self.config)
def test_createLocalWorkers(self):
"""
C{_createLocalWorkers} iterates the list of protocols and create one
L{LocalWorker} for each.
"""
protocols = [object() for x in range(4)]
workers = self.pool._createLocalWorkers(protocols, FilePath("path"), StringIO())
for s in workers:
self.assertIsInstance(s, LocalWorker)
self.assertEqual(4, len(workers))
def test_launchWorkerProcesses(self):
"""
Given a C{spawnProcess} function, C{_launchWorkerProcess} launches a
python process with an existing path as its argument.
"""
protocols = [ProcessProtocol() for i in range(4)]
arguments = []
environment = {}
def fakeSpawnProcess(
processProtocol,
executable,
args=(),
env={},
path=None,
uid=None,
gid=None,
usePTY=0,
childFDs=None,
):
arguments.append(executable)
arguments.extend(args)
environment.update(env)
self.pool._launchWorkerProcesses(fakeSpawnProcess, protocols, ["foo"])
self.assertEqual(arguments[0], arguments[1])
self.assertTrue(os.path.exists(arguments[2]))
self.assertEqual("foo", arguments[3])
# The child process runs with PYTHONPATH set to exactly the parent's
# import search path so that the child has a good chance of finding
# the same source files the parent would have found.
self.assertEqual(os.pathsep.join(sys.path), environment["PYTHONPATH"])
def test_run(self):
"""
C{run} dispatches the given action to each of its workers exactly once.
"""
# Make sure the parent of the working directory exists so
# manage a lock in it.
self.parent.makedirs()
workers = []
starting = self.pool.start(CountingReactor([]))
started = self.successResultOf(starting)
running = started.run(lambda w: succeed(workers.append(w)))
self.successResultOf(running)
assert_that(workers, has_length(self.config.numWorkers))
def test_runUsedDirectory(self):
"""
L{WorkerPool.start} checks if the test directory is already locked, and if
it is generates a name based on it.
"""
# Make sure the parent of the working directory exists so we can
# manage a lock in it.
self.parent.makedirs()
# Lock the directory the runner will expect to use.
lock = FilesystemLock(self.workingDirectory.path + ".lock")
self.assertTrue(lock.lock())
self.addCleanup(lock.unlock)
# Start up the pool
fakeReactor = CountingReactor([])
started = self.successResultOf(self.pool.start(fakeReactor))
# Verify it took a nearby directory instead.
self.assertEqual(
started.workingDirectory,
self.workingDirectory.sibling("_trial_temp-1"),
)
def test_join(self):
"""
L{StartedWorkerPool.join} causes all of the workers to exit, closes the
log file, and unlocks the test directory.
"""
self.parent.makedirs()
reactor = CountingReactor([])
started = self.successResultOf(self.pool.start(reactor))
joining = Deferred.fromCoroutine(started.join())
self.assertNoResult(joining)
for w in reactor._workers:
assert_that(w.transport._closed, contains(_WORKER_AMP_STDIN))
for fd in w.transport._closed:
w.childConnectionLost(fd)
for f in [w.processExited, w.processEnded]:
f(Failure(ProcessDone(0)))
assert_that(self.successResultOf(joining), none())
assert_that(started.testLog.closed, equal_to(True))
assert_that(started.testDirLock.locked, equal_to(False))
@given(
booleans(),
sampled_from(
[
"out.log",
f"subdir{sep}out.log",
]
),
)
def test_logFile(self, absolute: bool, logFile: str) -> None:
"""
L{WorkerPool.start} creates a L{StartedWorkerPool} configured with a
log file based on the L{WorkerPoolConfig.logFile}.
"""
if absolute:
logFile = self.parent.path + sep + logFile
config = assoc(self.config, logFile=logFile)
if absolute:
matches = equal_to(logFile)
else:
matches = AllOf(
# This might have a suffix if the configured workingDirectory
# was found to be in-use already so we don't add a sep suffix.
starts_with(config.workingDirectory.path),
# This should be exactly the suffix so we add a sep prefix.
ends_with(sep + logFile),
)
pool = WorkerPool(config)
started = self.successResultOf(pool.start(CountingReactor([])))
assert_that(started.testLog.name, matches)
class DistTrialRunnerTests(TestCase):
"""
Tests for L{DistTrialRunner}.
"""
suite = TrialSuite([sample.FooTest("test_foo")])
def getRunner(self, **overrides):
"""
Create a runner for testing.
"""
args = dict(
reporterFactory=TreeReporter,
workingDirectory=self.mktemp(),
stream=StringIO(),
maxWorkers=4,
workerArguments=[],
workerPoolFactory=partial(LocalWorkerPool, autostop=True),
reactor=CountingReactor([]),
)
args.update(overrides)
return DistTrialRunner(**args)
def test_writeResults(self):
"""
L{DistTrialRunner.writeResults} writes to the stream specified in the
init.
"""
stringIO = StringIO()
result = DistReporter(Reporter(stringIO))
runner = self.getRunner()
runner.writeResults(result)
self.assertTrue(stringIO.tell() > 0)
def test_minimalWorker(self):
"""
L{DistTrialRunner.runAsync} doesn't try to start more workers than the
number of tests.
"""
pool = None
def recordingFactory(*a, **kw):
nonlocal pool
pool = LocalWorkerPool(*a, autostop=True, **kw)
return pool
maxWorkers = 7
numTests = 3
runner = self.getRunner(
maxWorkers=maxWorkers, workerPoolFactory=recordingFactory
)
suite = TrialSuite([TestCase() for n in range(numTests)])
self.successResultOf(runner.runAsync(suite))
assert_that(pool._started[0].workers, has_length(numTests))
def test_runUncleanWarnings(self) -> None:
"""
Running with the C{unclean-warnings} option makes L{DistTrialRunner} uses
the L{UncleanWarningsReporterWrapper}.
"""
runner = self.getRunner(uncleanWarnings=True)
d = runner.runAsync(self.suite)
result = self.successResultOf(d)
self.assertIsInstance(result, DistReporter)
self.assertIsInstance(result.original, UncleanWarningsReporterWrapper)
def test_runWithoutTest(self):
"""
L{DistTrialRunner} can run an empty test suite.
"""
stream = StringIO()
runner = self.getRunner(stream=stream)
result = self.successResultOf(runner.runAsync(TrialSuite()))
self.assertIsInstance(result, DistReporter)
output = stream.getvalue()
self.assertIn("Running 0 test", output)
self.assertIn("PASSED", output)
def test_runWithoutTestButWithAnError(self):
"""
Even if there is no test, the suite can contain an error (most likely,
an import error): this should make the run fail, and the error should
be printed.
"""
err = ErrorHolder("an error", Failure(RuntimeError("foo bar")))
stream = StringIO()
runner = self.getRunner(stream=stream)
result = self.successResultOf(runner.runAsync(err))
self.assertIsInstance(result, DistReporter)
output = stream.getvalue()
self.assertIn("Running 0 test", output)
self.assertIn("foo bar", output)
self.assertIn("an error", output)
self.assertIn("errors=1", output)
self.assertIn("FAILED", output)
def test_runUnexpectedError(self) -> None:
"""
If for some reasons we can't connect to the worker process, the error is
recorded in the result object.
"""
runner = self.getRunner(workerPoolFactory=BrokenWorkerPool)
result = self.successResultOf(runner.runAsync(self.suite))
errors = result.original.errors
assert_that(errors, has_length(1))
assert_that(errors[0][1].type, equal_to(WorkerPoolBroken))
def test_runUnexpectedErrorCtrlC(self) -> None:
"""
If the reactor is stopped by C-c (i.e. `run` returns before the test
case's Deferred has been fired) we should cancel the pending test run.
"""
runner = self.getRunner(workerPoolFactory=LocalWorkerPool)
with self.assertRaises(CancelledError):
runner.run(self.suite)
def test_runUnexpectedWorkerError(self) -> None:
"""
If for some reason the worker process cannot run a test, the error is
recorded in the result object.
"""
runner = self.getRunner(
workerPoolFactory=partial(
LocalWorkerPool, workerFactory=_BrokenLocalWorker, autostop=True
)
)
result = self.successResultOf(runner.runAsync(self.suite))
errors = result.original.errors
assert_that(errors, has_length(1))
assert_that(errors[0][1].type, equal_to(WorkerBroken))
def test_runWaitForProcessesDeferreds(self) -> None:
"""
L{DistTrialRunner} waits for the worker pool to stop.
"""
pool = None
def recordingFactory(*a, **kw):
nonlocal pool
pool = LocalWorkerPool(*a, autostop=False, **kw)
return pool
runner = self.getRunner(
workerPoolFactory=recordingFactory,
)
d = Deferred.fromCoroutine(runner.runAsync(self.suite))
if pool is None:
self.fail("worker pool was never created")
assert pool is not None
stopped = pool._started[0]._stopped
self.assertNoResult(d)
stopped.callback(None)
result = self.successResultOf(d)
self.assertIsInstance(result, DistReporter)
def test_exitFirst(self):
"""
L{DistTrialRunner} can run in C{exitFirst} mode where it will run until a
test fails and then abandon the rest of the suite.
"""
stream = StringIO()
# Construct a suite with a failing test in the middle.
suite = TrialSuite(
[
sample.FooTest("test_foo"),
erroneous.TestRegularFail("test_fail"),
sample.FooTest("test_bar"),
]
)
runner = self.getRunner(stream=stream, exitFirst=True, maxWorkers=2)
d = runner.runAsync(suite)
result = self.successResultOf(d)
assert_that(
result.original,
matches_result(
successes=1,
failures=has_length(1),
),
)
def test_runUntilFailure(self):
"""
L{DistTrialRunner} can run in C{untilFailure} mode where it will run
the given tests until they fail.
"""
stream = StringIO()
case = erroneous.EventuallyFailingTestCase("test_it")
runner = self.getRunner(stream=stream)
d = runner.runAsync(case, untilFailure=True)
result = self.successResultOf(d)
# The case is hard-coded to fail on its 5th run.
self.assertEqual(5, case.n)
self.assertFalse(result.wasSuccessful())
output = stream.getvalue()
# It passes each time except the last.
self.assertEqual(
output.count("PASSED"),
case.n - 1,
"expected to see PASSED in output",
)
# It also fails at the end.
self.assertIn("FAIL", output)
# It also reports its progress.
for i in range(1, 6):
self.assertIn(f"Test Pass {i}", output)
# It also reports the number of tests run as part of each iteration.
self.assertEqual(
output.count("Ran 1 tests in"),
case.n,
"expected to see per-iteration test count in output",
)
def test_run(self) -> None:
"""
L{DistTrialRunner.run} returns a L{DistReporter} containing the result of
the test suite run.
"""
runner = self.getRunner()
result = runner.run(self.suite)
assert_that(result.wasSuccessful(), equal_to(True))
assert_that(result.successes, equal_to(1))
def test_installedReactor(self) -> None:
"""
L{DistTrialRunner.run} uses the installed reactor L{DistTrialRunner} was
constructed without a reactor.
"""
reactor = CountingReactor([])
with AlternateReactor(reactor):
runner = self.getRunner(reactor=None)
result = runner.run(self.suite)
assert_that(result.errors, equal_to([]))
assert_that(result.failures, equal_to([]))
assert_that(result.wasSuccessful(), equal_to(True))
assert_that(result.successes, equal_to(1))
assert_that(reactor.runCount, equal_to(1))
assert_that(reactor.stopCount, equal_to(1))
def test_wrongInstalledReactor(self) -> None:
"""
L{DistTrialRunner} raises L{TypeError} if the installed reactor provides
neither L{IReactorCore} nor L{IReactorProcess} and no other reactor is
given.
"""
class Core(ReactorBase):
def installWaker(self):
pass
@implementer(interfaces.IReactorProcess)
class Process:
def spawnProcess(
self,
processProtocol,
executable,
args,
env=None,
path=None,
uid=None,
gid=None,
usePTY=False,
childFDs=None,
):
pass
class Neither:
pass
# It provides neither
with AlternateReactor(Neither()):
with self.assertRaises(TypeError):
self.getRunner(reactor=None)
# It is missing IReactorProcess
with AlternateReactor(Core()):
with self.assertRaises(TypeError):
self.getRunner(reactor=None)
# It is missing IReactorCore
with AlternateReactor(Process()):
with self.assertRaises(TypeError):
self.getRunner(reactor=None)
def test_runFailure(self):
"""
If there is an unexpected exception running the test suite then it is
re-raised by L{DistTrialRunner.run}.
"""
# Give it a broken worker pool factory. There's no exception handling
# for such an error in the implementation..
class BrokenFactory(Exception):
pass
def brokenFactory(*args, **kwargs):
raise BrokenFactory()
runner = self.getRunner(workerPoolFactory=brokenFactory)
with self.assertRaises(BrokenFactory):
runner.run(self.suite)
class FunctionalTests(TestCase):
"""
Tests for the functional helpers that need it.
"""
def test_fromOptional(self) -> None:
"""
``fromOptional`` accepts a default value and an ``Optional`` value of the
same type and returns the default value if the optional value is
``None`` or the optional value otherwise.
"""
assert_that(fromOptional(1, None), equal_to(1))
assert_that(fromOptional(2, 2), equal_to(2))
def test_discardResult(self) -> None:
"""
``discardResult`` accepts an awaitable and returns a ``Deferred`` that
fires with ``None`` after the awaitable completes.
"""
a: Deferred[str] = Deferred()
d = discardResult(a)
self.assertNoResult(d)
a.callback("result")
assert_that(self.successResultOf(d), none())
def test_sequence(self) -> None:
"""
``sequence`` accepts two awaitables and returns an awaitable that waits
for the first one to complete and then completes with the result of
the second one.
"""
a: Deferred[str] = Deferred()
b: Deferred[int] = Deferred()
c = Deferred.fromCoroutine(sequence(a, b))
b.callback(42)
self.assertNoResult(c)
a.callback("hello")
assert_that(self.successResultOf(c), equal_to(42))
def test_iterateWhile(self) -> None:
"""
``iterateWhile`` executes the actions from its factory until the predicate
does not match an action result.
"""
actions: List[Deferred[int]] = [Deferred(), Deferred(), Deferred()]
def predicate(value):
return value != 42
d: Deferred[int] = Deferred.fromCoroutine(
iterateWhile(predicate, list(actions).pop)
)
# Let the action it is waiting on complete
actions.pop().callback(7)
# It does not match the predicate so it is not done yet.
self.assertNoResult(d)
# Let the action it is waiting on now complete - with the result it
# wants.
actions.pop().callback(42)
assert_that(self.successResultOf(d), equal_to(42))
def test_countingCalls(self) -> None:
"""
``countingCalls`` decorates a function so that it is called with an
increasing counter and passes the return value through.
"""
@countingCalls
def target(n: int) -> int:
return n + 1
for expected in range(1, 10):
assert_that(target(), equal_to(expected))
class WorkerPoolBroken(Exception):
"""
An exception for ``StartedWorkerPoolBroken`` to fail with to allow tests
to exercise exception code paths.
"""
class StartedWorkerPoolBroken:
"""
A broken, started worker pool. Its workers cannot run actions. They
always raise an exception.
"""
async def run(self, workerAction: WorkerAction[None]) -> None:
raise WorkerPoolBroken()
async def join(self) -> None:
return None
@define
class BrokenWorkerPool:
"""
A worker pool that has workers with a broken ``run`` method.
"""
_config: WorkerPoolConfig
async def start(
self, reactor: interfaces.IReactorProcess
) -> StartedWorkerPoolBroken:
return StartedWorkerPoolBroken()
class _LocalWorker:
"""
A L{Worker} that runs tests in this process in the usual way.
This is a test double for L{LocalWorkerAMP} which allows testing worker
pool logic without sending tests over an AMP connection to be run
somewhere else..
"""
async def run(self, case: PyUnitTestCase, result: TestResult) -> RunResult:
"""
Directly run C{case} in the usual way.
"""
TrialSuite([case]).run(result)
return {"success": True}
class WorkerBroken(Exception):
"""
A worker tried to run a test case but the worker is broken.
"""
class _BrokenLocalWorker:
"""
A L{Worker} that always fails to run test cases.
"""
async def run(self, case: PyUnitTestCase, result: TestResult) -> None:
"""
Raise an exception instead of running C{case}.
"""
raise WorkerBroken()
@define
class StartedLocalWorkerPool:
"""
A started L{LocalWorkerPool}.
"""
workingDirectory: FilePath[str]
workers: List[Worker]
_stopped: Deferred[None]
async def run(self, workerAction: WorkerAction[None]) -> None:
"""
Run the action with each local worker.
"""
for worker in self.workers:
await workerAction(worker)
async def join(self):
await self._stopped
@define
class LocalWorkerPool:
"""
Implement a worker pool that runs tests in-process instead of in child
processes.
"""
_config: WorkerPoolConfig
_started: List[StartedLocalWorkerPool] = field(default=Factory(list))
_autostop: bool = False
_workerFactory: Callable[[], Worker] = _LocalWorker
async def start(
self, reactor: interfaces.IReactorProcess
) -> StartedLocalWorkerPool:
workers = [self._workerFactory() for i in range(self._config.numWorkers)]
started = StartedLocalWorkerPool(
self._config.workingDirectory,
workers,
(succeed(None) if self._autostop else Deferred()),
)
self._started.append(started)
return started

View File

@@ -0,0 +1,188 @@
"""
Tests for L{twisted.trial._dist.test.matchers}.
"""
from typing import Callable, Sequence, Tuple, Type
from hamcrest import anything, assert_that, contains, contains_string, equal_to, not_
from hamcrest.core.matcher import Matcher
from hamcrest.core.string_description import StringDescription
from hypothesis import given
from hypothesis.strategies import (
binary,
booleans,
integers,
just,
lists,
one_of,
sampled_from,
text,
tuples,
)
from twisted.python.failure import Failure
from twisted.trial.unittest import SynchronousTestCase
from .matchers import HasSum, IsSequenceOf, S, isFailure, similarFrame
Summer = Callable[[Sequence[S]], S]
concatInt = sum
concatStr = "".join
concatBytes = b"".join
class HasSumTests(SynchronousTestCase):
"""
Tests for L{HasSum}.
"""
summables = one_of(
tuples(lists(integers()), just(concatInt)),
tuples(lists(text()), just(concatStr)),
tuples(lists(binary()), just(concatBytes)),
)
@given(summables)
def test_matches(self, summable: Tuple[Sequence[S], Summer[S]]) -> None:
"""
L{HasSum} matches a sequence if the elements sum to a value matched by
the parameterized matcher.
:param summable: A tuple of a sequence of values to try to match and a
function which can compute the correct sum for that sequence.
"""
seq, sumFunc = summable
expected = sumFunc(seq)
zero = sumFunc([])
matcher = HasSum(equal_to(expected), zero)
description = StringDescription()
assert_that(matcher.matches(seq, description), equal_to(True))
assert_that(str(description), equal_to(""))
@given(summables)
def test_mismatches(
self,
summable: Tuple[
Sequence[S],
Summer[S],
],
) -> None:
"""
L{HasSum} does not match a sequence if the elements do not sum to a
value matched by the parameterized matcher.
:param summable: See L{test_matches}.
"""
seq, sumFunc = summable
zero = sumFunc([])
# A matcher that never matches.
sumMatcher: Matcher[S] = not_(anything())
matcher = HasSum(sumMatcher, zero)
actualDescription = StringDescription()
assert_that(matcher.matches(seq, actualDescription), equal_to(False))
sumMatcherDescription = StringDescription()
sumMatcherDescription.append_description_of(sumMatcher)
actualStr = str(actualDescription)
assert_that(actualStr, contains_string("a sequence with sum"))
assert_that(actualStr, contains_string(str(sumMatcherDescription)))
class IsSequenceOfTests(SynchronousTestCase):
"""
Tests for L{IsSequenceOf}.
"""
sequences = lists(booleans())
@given(integers(min_value=0, max_value=1000))
def test_matches(self, numItems: int) -> None:
"""
L{IsSequenceOf} matches a sequence if all of the elements are
matched by the parameterized matcher.
:param numItems: The length of a sequence to try to match.
"""
seq = [True] * numItems
matcher = IsSequenceOf(equal_to(True))
actualDescription = StringDescription()
assert_that(matcher.matches(seq, actualDescription), equal_to(True))
assert_that(str(actualDescription), equal_to(""))
@given(integers(min_value=0, max_value=1000), integers(min_value=0, max_value=1000))
def test_mismatches(self, numBefore: int, numAfter: int) -> None:
"""
L{IsSequenceOf} does not match a sequence if any of the elements
are not matched by the parameterized matcher.
:param numBefore: In the sequence to try to match, the number of
elements expected to match before an expected mismatch.
:param numAfter: In the sequence to try to match, the number of
elements expected expected to match after an expected mismatch.
"""
# Hide the non-matching value somewhere in the sequence.
seq = [True] * numBefore + [False] + [True] * numAfter
matcher = IsSequenceOf(equal_to(True))
actualDescription = StringDescription()
assert_that(matcher.matches(seq, actualDescription), equal_to(False))
actualStr = str(actualDescription)
assert_that(actualStr, contains_string("a sequence containing only"))
assert_that(
actualStr, contains_string(f"not sequence with element #{numBefore}")
)
class IsFailureTests(SynchronousTestCase):
"""
Tests for L{isFailure}.
"""
@given(sampled_from([ValueError, ZeroDivisionError, RuntimeError]))
def test_matches(self, excType: Type[BaseException]) -> None:
"""
L{isFailure} matches instances of L{Failure} with matching
attributes.
:param excType: An exception type to wrap in a L{Failure} to be
matched against.
"""
matcher = isFailure(type=equal_to(excType))
failure = Failure(excType())
assert_that(matcher.matches(failure), equal_to(True))
@given(sampled_from([ValueError, ZeroDivisionError, RuntimeError]))
def test_mismatches(self, excType: Type[BaseException]) -> None:
"""
L{isFailure} does not match instances of L{Failure} with
attributes that don't match.
:param excType: An exception type to wrap in a L{Failure} to be
matched against.
"""
matcher = isFailure(type=equal_to(excType), other=not_(anything()))
failure = Failure(excType())
assert_that(matcher.matches(failure), equal_to(False))
def test_frames(self):
"""
The L{similarFrame} matcher matches elements of the C{frames} list
of a L{Failure}.
"""
try:
raise ValueError("Oh no")
except BaseException:
f = Failure()
actualDescription = StringDescription()
matcher = isFailure(
frames=contains(similarFrame("test_frames", "test_matchers"))
)
assert_that(
matcher.matches(f, actualDescription),
equal_to(True),
actualDescription,
)

View File

@@ -0,0 +1,48 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Tests for distributed trial's options management.
"""
import gc
import os
import sys
from twisted.trial._dist.options import WorkerOptions
from twisted.trial.unittest import TestCase
class WorkerOptionsTests(TestCase):
"""
Tests for L{WorkerOptions}.
"""
def setUp(self) -> None:
"""
Build an L{WorkerOptions} object to be used in the tests.
"""
self.options = WorkerOptions()
def test_standardOptions(self) -> None:
"""
L{WorkerOptions} supports a subset of standard options supported by
trial.
"""
self.addCleanup(sys.setrecursionlimit, sys.getrecursionlimit())
if gc.isenabled():
self.addCleanup(gc.enable)
gc.enable()
self.options.parseOptions(["--recursionlimit", "2000", "--disablegc"])
self.assertEqual(2000, sys.getrecursionlimit())
self.assertFalse(gc.isenabled())
def test_coverage(self) -> None:
"""
L{WorkerOptions.coverdir} returns the C{coverage} child directory of
the current directory to be used for storing coverage data.
"""
self.assertEqual(
os.path.realpath(os.path.join(os.getcwd(), "coverage")),
self.options.coverdir().path,
)

View File

@@ -0,0 +1,208 @@
"""
Tests for L{twisted.trial._dist.stream}.
"""
from random import Random
from typing import Awaitable, Dict, List, TypeVar, Union
from hamcrest import (
all_of,
assert_that,
calling,
equal_to,
has_length,
is_,
less_than_or_equal_to,
raises,
)
from hypothesis import given
from hypothesis.strategies import binary, integers, just, lists, randoms, text
from twisted.internet.defer import Deferred, fail
from twisted.internet.interfaces import IProtocol
from twisted.internet.protocol import Protocol
from twisted.protocols.amp import AMP
from twisted.python.failure import Failure
from twisted.test.iosim import FakeTransport, connect
from twisted.trial.unittest import SynchronousTestCase
from ..stream import StreamOpen, StreamReceiver, StreamWrite, chunk, stream
from .matchers import HasSum, IsSequenceOf
T = TypeVar("T")
class StreamReceiverTests(SynchronousTestCase):
"""
Tests for L{StreamReceiver}
"""
@given(lists(lists(binary())), randoms())
def test_streamReceived(self, streams: List[List[bytes]], random: Random) -> None:
"""
All data passed to L{StreamReceiver.write} is returned by a call to
L{StreamReceiver.finish} with a matching C{streamId}.
"""
receiver = StreamReceiver()
streamIds = [receiver.open() for _ in streams]
# uncorrelate the results with open() order
random.shuffle(streamIds)
expectedData = dict(zip(streamIds, streams))
for streamId, strings in expectedData.items():
for s in strings:
receiver.write(streamId, s)
# uncorrelate the results with write() order
random.shuffle(streamIds)
actualData = {streamId: receiver.finish(streamId) for streamId in streamIds}
assert_that(actualData, is_(equal_to(expectedData)))
@given(integers(), just("data"))
def test_writeBadStreamId(self, streamId: int, data: str) -> None:
"""
L{StreamReceiver.write} raises L{KeyError} if called with a
streamId not associated with an open stream.
"""
receiver = StreamReceiver()
assert_that(calling(receiver.write).with_args(streamId, data), raises(KeyError))
@given(integers())
def test_badFinishStreamId(self, streamId: int) -> None:
"""
L{StreamReceiver.finish} raises L{KeyError} if called with a
streamId not associated with an open stream.
"""
receiver = StreamReceiver()
assert_that(calling(receiver.finish).with_args(streamId), raises(KeyError))
def test_finishRemovesStream(self) -> None:
"""
L{StreamReceiver.finish} removes the identified stream.
"""
receiver = StreamReceiver()
streamId = receiver.open()
receiver.finish(streamId)
assert_that(calling(receiver.finish).with_args(streamId), raises(KeyError))
class ChunkTests(SynchronousTestCase):
"""
Tests for ``chunk``.
"""
@given(data=text(), chunkSize=integers(min_value=1))
def test_chunk(self, data, chunkSize):
"""
L{chunk} returns an iterable of L{str} where each element is no
longer than the given limit. The concatenation of the strings is also
equal to the original input string.
"""
chunks = list(chunk(data, chunkSize))
assert_that(
chunks,
all_of(
IsSequenceOf(
has_length(less_than_or_equal_to(chunkSize)),
),
HasSum(equal_to(data), data[:0]),
),
)
class AMPStreamReceiver(AMP):
"""
A simple AMP interface to L{StreamReceiver}.
"""
def __init__(self, streams: StreamReceiver) -> None:
self.streams = streams
@StreamOpen.responder
def streamOpen(self) -> Dict[str, object]:
return {"streamId": self.streams.open()}
@StreamWrite.responder
def streamWrite(self, streamId: int, data: bytes) -> Dict[str, object]:
self.streams.write(streamId, data)
return {}
def interact(server: IProtocol, client: IProtocol, interaction: Awaitable[T]) -> T:
"""
Let C{server} and C{client} exchange bytes while C{interaction} runs.
"""
finished = False
result: Union[Failure, T]
async def to_coroutine() -> T:
return await interaction
def collect_result(r: Union[Failure, T]) -> None:
nonlocal result, finished
finished = True
result = r
pump = connect(
server,
FakeTransport(server, isServer=True),
client,
FakeTransport(client, isServer=False),
)
interacting = Deferred.fromCoroutine(to_coroutine())
interacting.addBoth(collect_result)
pump.flush()
if finished:
if isinstance(result, Failure):
result.raiseException()
return result
raise Exception("Interaction failed to produce a result.")
class InteractTests(SynchronousTestCase):
"""
Tests for the test helper L{interact}.
"""
def test_failure(self):
"""
If the interaction results in a failure then L{interact} raises an
exception.
"""
class ArbitraryException(Exception):
pass
with self.assertRaises(ArbitraryException):
interact(Protocol(), Protocol(), fail(ArbitraryException()))
def test_incomplete(self):
"""
If the interaction fails to produce a result then L{interact} raises
an exception.
"""
with self.assertRaises(Exception):
interact(Protocol(), Protocol(), Deferred())
class StreamTests(SynchronousTestCase):
"""
Tests for L{stream}.
"""
@given(lists(binary()))
def test_stream(self, chunks: List[bytes]) -> None:
"""
All of the chunks passed to L{stream} are sent in order over a
stream using the given AMP connection.
"""
sender = AMP()
streams = StreamReceiver()
streamId = interact(
AMPStreamReceiver(streams), sender, stream(sender, iter(chunks))
)
assert_that(streams.finish(streamId), is_(equal_to(chunks)))

View File

@@ -0,0 +1,532 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Test for distributed trial worker side.
"""
import os
from io import BytesIO, StringIO
from typing import Type
from unittest import TestCase as PyUnitTestCase
from zope.interface.verify import verifyObject
from hamcrest import assert_that, equal_to, has_item, has_length
from twisted.internet.defer import Deferred, fail
from twisted.internet.error import ConnectionLost, ProcessDone
from twisted.internet.interfaces import IAddress, ITransport
from twisted.python.failure import Failure
from twisted.python.filepath import FilePath
from twisted.test.iosim import connectedServerAndClient
from twisted.trial._dist import managercommands
from twisted.trial._dist.worker import (
LocalWorker,
LocalWorkerAMP,
LocalWorkerTransport,
NotRunning,
WorkerException,
WorkerProtocol,
)
from twisted.trial.reporter import TestResult
from twisted.trial.test import pyunitcases, skipping
from twisted.trial.unittest import TestCase, makeTodo
from .matchers import isFailure, matches_result, similarFrame
class WorkerProtocolTests(TestCase):
"""
Tests for L{WorkerProtocol}.
"""
worker: WorkerProtocol
server: LocalWorkerAMP
def setUp(self) -> None:
"""
Set up a transport, a result stream and a protocol instance.
"""
self.worker, self.server, pump = connectedServerAndClient(
LocalWorkerAMP, WorkerProtocol, greet=False
)
self.flush = pump.flush
def test_run(self) -> None:
"""
Sending the L{workercommands.Run} command to the worker returns a
response with C{success} sets to C{True}.
"""
d = Deferred.fromCoroutine(
self.server.run(pyunitcases.PyUnitTest("test_pass"), TestResult())
)
self.flush()
self.assertEqual({"success": True}, self.successResultOf(d))
def test_start(self) -> None:
"""
The C{start} command changes the current path.
"""
curdir = os.path.realpath(os.path.curdir)
self.addCleanup(os.chdir, curdir)
self.worker.start("..")
self.assertNotEqual(os.path.realpath(os.path.curdir), curdir)
class WorkerProtocolErrorTests(TestCase):
"""
Tests for L{WorkerProtocol}'s handling of certain errors related to
running the tests themselves (i.e., not test errors but test
infrastructure/runner errors).
"""
def _runErrorTest(
self, brokenTestName: str, loggedExceptionType: Type[BaseException]
) -> None:
worker, server, pump = connectedServerAndClient(
LocalWorkerAMP, WorkerProtocol, greet=False
)
expectedCase = pyunitcases.BrokenRunInfrastructure(brokenTestName)
result = TestResult()
Deferred.fromCoroutine(server.run(expectedCase, result))
pump.flush()
assert_that(result, matches_result(errors=has_length(1)))
[(actualCase, errors)] = result.errors
assert_that(actualCase, equal_to(expectedCase))
# Additionally, we expect that the worker protocol logged the failure
# once so that it is visible somewhere, even if it cannot deliver it
# back to the parent process (which it can in this case). Since the
# worker runs in process with us, that failure is in our log so we can
# easily make an assertion about it. Also, if we don't flush it, the
# test fails. As far as the type goes, we just have to be aware of
# the implementation details of `BrokenRunInfrastructure`.
assert_that(self.flushLoggedErrors(loggedExceptionType), has_length(1))
def test_addSuccessError(self) -> None:
"""
If there is an error reporting success then the test run is marked as
an error.
"""
self._runErrorTest("test_addSuccess", AttributeError)
def test_addErrorError(self) -> None:
"""
If there is an error reporting an error then the test run is marked as
an error.
"""
self._runErrorTest("test_addError", AttributeError)
def test_addFailureError(self) -> None:
"""
If there is an error reporting a failure then the test run is marked
as an error.
"""
self._runErrorTest("test_addFailure", AttributeError)
def test_addSkipError(self) -> None:
"""
If there is an error reporting a skip then the test run is marked
as an error.
"""
self._runErrorTest("test_addSkip", AttributeError)
def test_addExpectedFailure(self) -> None:
"""
If there is an error reporting an expected failure then the test
run is marked as an error.
"""
self._runErrorTest("test_addExpectedFailure", AttributeError)
def test_addUnexpectedSuccess(self) -> None:
"""
If there is an error reporting an unexpected ccess then the test
run is marked as an error.
"""
self._runErrorTest("test_addUnexpectedSuccess", AttributeError)
def test_failedFailureReport(self) -> None:
"""
A failure encountered while reporting a reporting failure is logged.
"""
worker, server, pump = connectedServerAndClient(
LocalWorkerAMP, WorkerProtocol, greet=False
)
# We can easily break everything by eliminating the worker protocol's
# transport. This prevents it from ever sending anything to the
# manager protocol.
worker.transport = None
expectedCase = pyunitcases.PyUnitTest("test_pass")
result = TestResult()
Deferred.fromCoroutine(server.run(expectedCase, result))
pump.flush()
# There should be two exceptions logged here. The first is from the
# attempt to report the success result. The second is a report that
# the first failed.
assert_that(self.flushLoggedErrors(ConnectionLost), has_length(2))
class LocalWorkerAMPTests(TestCase):
"""
Test case for distributed trial's manager-side local worker AMP protocol
"""
def setUp(self) -> None:
self.worker, self.managerAMP, pump = connectedServerAndClient(
LocalWorkerAMP, WorkerProtocol, greet=False
)
self.flush = pump.flush
def workerRunTest(
self, testCase: PyUnitTestCase, makeResult: Type[TestResult] = TestResult
) -> TestResult:
result = makeResult()
d = Deferred.fromCoroutine(self.managerAMP.run(testCase, result))
self.flush()
self.assertEqual({"success": True}, self.successResultOf(d))
return result
def test_runSuccess(self) -> None:
"""
Run a test, and succeed.
"""
result = self.workerRunTest(pyunitcases.PyUnitTest("test_pass"))
assert_that(result, matches_result(successes=equal_to(1)))
def test_runExpectedFailure(self) -> None:
"""
Run a test, and fail expectedly.
"""
expectedCase = skipping.SynchronousStrictTodo("test_todo1")
result = self.workerRunTest(expectedCase)
assert_that(result, matches_result(expectedFailures=has_length(1)))
[(actualCase, exceptionMessage, todoReason)] = result.expectedFailures
assert_that(actualCase, equal_to(expectedCase))
# Match the strings used in the test we ran.
assert_that(exceptionMessage, equal_to("expected failure"))
assert_that(todoReason, equal_to(makeTodo("todo1")))
def test_runError(self) -> None:
"""
Run a test, and encounter an error.
"""
expectedCase = pyunitcases.PyUnitTest("test_error")
result = self.workerRunTest(expectedCase)
assert_that(result, matches_result(errors=has_length(1)))
[(actualCase, failure)] = result.errors
assert_that(expectedCase, equal_to(actualCase))
assert_that(
failure,
isFailure(
type=equal_to(Exception),
value=equal_to(WorkerException("pyunit error")),
frames=has_item(similarFrame("test_error", "pyunitcases.py")), # type: ignore[arg-type]
),
)
def test_runFailure(self) -> None:
"""
Run a test, and fail.
"""
expectedCase = pyunitcases.PyUnitTest("test_fail")
result = self.workerRunTest(expectedCase)
assert_that(result, matches_result(failures=has_length(1)))
[(actualCase, failure)] = result.failures
assert_that(expectedCase, equal_to(actualCase))
assert_that(
failure,
isFailure(
# AssertionError is the type raised by TestCase.fail
type=equal_to(AssertionError),
value=equal_to(WorkerException("pyunit failure")),
),
)
def test_runSkip(self) -> None:
"""
Run a test, but skip it.
"""
expectedCase = pyunitcases.PyUnitTest("test_skip")
result = self.workerRunTest(expectedCase)
assert_that(result, matches_result(skips=has_length(1)))
[(actualCase, skip)] = result.skips
assert_that(expectedCase, equal_to(actualCase))
assert_that(skip, equal_to("pyunit skip"))
def test_runUnexpectedSuccesses(self) -> None:
"""
Run a test, and succeed unexpectedly.
"""
expectedCase = skipping.SynchronousStrictTodo("test_todo7")
result = self.workerRunTest(expectedCase)
assert_that(result, matches_result(unexpectedSuccesses=has_length(1)))
[(actualCase, unexpectedSuccess)] = result.unexpectedSuccesses
assert_that(expectedCase, equal_to(actualCase))
assert_that(unexpectedSuccess, equal_to("todo7"))
def test_testWrite(self) -> None:
"""
L{LocalWorkerAMP.testWrite} writes the data received to its test
stream.
"""
stream = StringIO()
self.managerAMP.setTestStream(stream)
d = self.worker.callRemote(managercommands.TestWrite, out="Some output")
self.flush()
self.assertEqual({"success": True}, self.successResultOf(d))
self.assertEqual("Some output\n", stream.getvalue())
def test_stopAfterRun(self) -> None:
"""
L{LocalWorkerAMP.run} calls C{stopTest} on its test result once the
C{Run} commands has succeeded.
"""
stopped = []
class StopTestResult(TestResult):
def stopTest(self, test: PyUnitTestCase) -> None:
stopped.append(test)
case = pyunitcases.PyUnitTest("test_pass")
self.workerRunTest(case, StopTestResult)
assert_that(stopped, equal_to([case]))
class SpyDataLocalWorkerAMP(LocalWorkerAMP):
"""
A fake implementation of L{LocalWorkerAMP} that records the received
data and doesn't automatically dispatch any command..
"""
id = 0
dataString = b""
def dataReceived(self, data):
self.dataString += data
class FakeTransport:
"""
A fake process transport implementation for testing.
"""
dataString = b""
calls = 0
def writeToChild(self, fd, data):
self.dataString += data
def loseConnection(self):
self.calls += 1
class LocalWorkerTests(TestCase):
"""
Tests for L{LocalWorker} and L{LocalWorkerTransport}.
"""
def tidyLocalWorker(self, *args, **kwargs):
"""
Create a L{LocalWorker}, connect it to a transport, and ensure
its log files are closed.
@param args: See L{LocalWorker}
@param kwargs: See L{LocalWorker}
@return: a L{LocalWorker} instance
"""
worker = LocalWorker(*args, **kwargs)
worker.makeConnection(FakeTransport())
self.addCleanup(worker._outLog.close)
self.addCleanup(worker._errLog.close)
return worker
def test_exitBeforeConnected(self):
"""
L{LocalWorker.exit} fails with L{NotRunning} if it is called before the
protocol is connected to a transport.
"""
worker = LocalWorker(
SpyDataLocalWorkerAMP(), FilePath(self.mktemp()), StringIO()
)
self.failureResultOf(worker.exit(), NotRunning)
def test_exitAfterDisconnected(self):
"""
L{LocalWorker.exit} fails with L{NotRunning} if it is called after the the
protocol is disconnected from its transport.
"""
worker = self.tidyLocalWorker(
SpyDataLocalWorkerAMP(), FilePath(self.mktemp()), StringIO()
)
worker.processEnded(Failure(ProcessDone(0)))
# Since we're not calling exit until after the process has ended, it
# won't consume the ProcessDone failure on the internal `endDeferred`.
# Swallow it here.
self.failureResultOf(worker.endDeferred, ProcessDone)
# Now assert that exit behaves.
self.failureResultOf(worker.exit(), NotRunning)
def test_childDataReceived(self):
"""
L{LocalWorker.childDataReceived} forwards the received data to linked
L{AMP} protocol if the right file descriptor, otherwise forwards to
C{ProcessProtocol.childDataReceived}.
"""
localWorker = self.tidyLocalWorker(
SpyDataLocalWorkerAMP(), FilePath(self.mktemp()), "test.log"
)
localWorker._outLog = BytesIO()
localWorker.childDataReceived(4, b"foo")
localWorker.childDataReceived(1, b"bar")
self.assertEqual(b"foo", localWorker._ampProtocol.dataString)
self.assertEqual(b"bar", localWorker._outLog.getvalue())
def test_newlineStyle(self):
"""
L{LocalWorker} writes the log data with local newlines.
"""
amp = SpyDataLocalWorkerAMP()
tempDir = FilePath(self.mktemp())
tempDir.makedirs()
logPath = tempDir.child("test.log")
with open(logPath.path, "wt", encoding="utf-8") as logFile:
worker = LocalWorker(amp, tempDir, logFile)
worker.makeConnection(FakeTransport())
self.addCleanup(worker._outLog.close)
self.addCleanup(worker._errLog.close)
expected = "Here comes the \N{sun}!"
amp.testWrite(expected)
self.assertEqual(
# os.linesep is the local newline.
(expected + os.linesep),
# getContent reads in binary mode so we'll see the bytes that
# actually ended up in the file.
logPath.getContent().decode("utf-8"),
)
def test_outReceived(self):
"""
L{LocalWorker.outReceived} logs the output into its C{_outLog} log
file.
"""
localWorker = self.tidyLocalWorker(
SpyDataLocalWorkerAMP(), FilePath(self.mktemp()), "test.log"
)
localWorker._outLog = BytesIO()
data = b"The quick brown fox jumps over the lazy dog"
localWorker.outReceived(data)
self.assertEqual(data, localWorker._outLog.getvalue())
def test_errReceived(self):
"""
L{LocalWorker.errReceived} logs the errors into its C{_errLog} log
file.
"""
localWorker = self.tidyLocalWorker(
SpyDataLocalWorkerAMP(), FilePath(self.mktemp()), "test.log"
)
localWorker._errLog = BytesIO()
data = b"The quick brown fox jumps over the lazy dog"
localWorker.errReceived(data)
self.assertEqual(data, localWorker._errLog.getvalue())
def test_write(self):
"""
L{LocalWorkerTransport.write} forwards the written data to the given
transport.
"""
transport = FakeTransport()
localTransport = LocalWorkerTransport(transport)
data = b"The quick brown fox jumps over the lazy dog"
localTransport.write(data)
self.assertEqual(data, transport.dataString)
def test_writeSequence(self):
"""
L{LocalWorkerTransport.writeSequence} forwards the written data to the
given transport.
"""
transport = FakeTransport()
localTransport = LocalWorkerTransport(transport)
data = (b"The quick ", b"brown fox jumps ", b"over the lazy dog")
localTransport.writeSequence(data)
self.assertEqual(b"".join(data), transport.dataString)
def test_loseConnection(self):
"""
L{LocalWorkerTransport.loseConnection} forwards the call to the given
transport.
"""
transport = FakeTransport()
localTransport = LocalWorkerTransport(transport)
localTransport.loseConnection()
self.assertEqual(transport.calls, 1)
def test_connectionLost(self):
"""
L{LocalWorker.connectionLost} closes the per-worker log streams.
"""
localWorker = self.tidyLocalWorker(
SpyDataLocalWorkerAMP(), FilePath(self.mktemp()), "test.log"
)
localWorker.connectionLost(None)
self.assertTrue(localWorker._outLog.closed)
self.assertTrue(localWorker._errLog.closed)
def test_processEnded(self):
"""
L{LocalWorker.processEnded} calls C{connectionLost} on itself and on
the L{AMP} protocol.
"""
transport = FakeTransport()
protocol = SpyDataLocalWorkerAMP()
localWorker = LocalWorker(protocol, FilePath(self.mktemp()), "test.log")
localWorker.makeConnection(transport)
localWorker.processEnded(Failure(ProcessDone(0)))
self.assertTrue(localWorker._outLog.closed)
self.assertTrue(localWorker._errLog.closed)
self.assertIdentical(None, protocol.transport)
return self.assertFailure(localWorker.endDeferred, ProcessDone)
def test_addresses(self):
"""
L{LocalWorkerTransport.getPeer} and L{LocalWorkerTransport.getHost}
return L{IAddress} objects.
"""
localTransport = LocalWorkerTransport(None)
self.assertTrue(verifyObject(IAddress, localTransport.getPeer()))
self.assertTrue(verifyObject(IAddress, localTransport.getHost()))
def test_transport(self):
"""
L{LocalWorkerTransport} implements L{ITransport} to be able to be used
by L{AMP}.
"""
localTransport = LocalWorkerTransport(None)
self.assertTrue(verifyObject(ITransport, localTransport))
def test_startError(self):
"""
L{LocalWorker} swallows the exceptions returned by the L{AMP} protocol
start method, as it generates unnecessary errors.
"""
def failCallRemote(command, directory):
return fail(RuntimeError("oops"))
protocol = SpyDataLocalWorkerAMP()
protocol.callRemote = failCallRemote
self.tidyLocalWorker(protocol, FilePath(self.mktemp()), "test.log")
self.assertEqual([], self.flushLoggedErrors(RuntimeError))

View File

@@ -0,0 +1,165 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Tests for L{twisted.trial._dist.workerreporter}.
"""
from __future__ import annotations
from typing import Sized
from unittest import TestCase
from hamcrest import assert_that, equal_to, has_length
from hamcrest.core.matcher import Matcher
from twisted.internet.defer import Deferred
from twisted.test.iosim import connectedServerAndClient
from twisted.trial._dist.worker import LocalWorkerAMP, WorkerProtocol
from twisted.trial.reporter import TestResult
from twisted.trial.test import erroneous, pyunitcases, sample, skipping
from twisted.trial.unittest import SynchronousTestCase
from .matchers import matches_result
def run(case: SynchronousTestCase, target: TestCase) -> TestResult:
"""
Run C{target} and return a test result as populated by a worker reporter.
@param case: A test case to use to help run the target.
"""
result = TestResult()
worker, local, pump = connectedServerAndClient(LocalWorkerAMP, WorkerProtocol)
d = Deferred.fromCoroutine(local.run(target, result))
pump.flush()
assert_that(case.successResultOf(d), equal_to({"success": True}))
return result
class WorkerReporterTests(SynchronousTestCase):
"""
Tests for L{WorkerReporter}.
"""
def assertTestRun(self, target: TestCase, **expectations: Matcher[Sized]) -> None:
"""
Run the given test and assert that the result matches the given
expectations.
"""
assert_that(run(self, target), matches_result(**expectations))
def test_outsideReportingContext(self) -> None:
"""
L{WorkerReporter}'s implementation of test result methods raise
L{ValueError} when called outside of the
L{WorkerReporter.gatherReportingResults} context manager.
"""
worker, local, pump = connectedServerAndClient(LocalWorkerAMP, WorkerProtocol)
case = sample.FooTest("test_foo")
with self.assertRaises(ValueError):
worker._result.addSuccess(case)
def test_addSuccess(self) -> None:
"""
L{WorkerReporter} propagates successes.
"""
self.assertTestRun(sample.FooTest("test_foo"), successes=equal_to(1))
def test_addError(self) -> None:
"""
L{WorkerReporter} propagates errors from trial's TestCases.
"""
self.assertTestRun(
erroneous.TestAsynchronousFail("test_exception"), errors=has_length(1)
)
def test_addErrorGreaterThan64k(self) -> None:
"""
L{WorkerReporter} propagates errors with large string representations.
"""
self.assertTestRun(
erroneous.TestAsynchronousFail("test_exceptionGreaterThan64k"),
errors=has_length(1),
)
def test_addErrorGreaterThan64kEncoded(self) -> None:
"""
L{WorkerReporter} propagates errors with a string representation that
is smaller than an implementation-specific limit but which encode to a
byte representation that exceeds this limit.
"""
self.assertTestRun(
erroneous.TestAsynchronousFail("test_exceptionGreaterThan64kEncoded"),
errors=has_length(1),
)
def test_addErrorTuple(self) -> None:
"""
L{WorkerReporter} propagates errors from pyunit's TestCases.
"""
self.assertTestRun(pyunitcases.PyUnitTest("test_error"), errors=has_length(1))
def test_addFailure(self) -> None:
"""
L{WorkerReporter} propagates test failures from trial's TestCases.
"""
self.assertTestRun(
erroneous.TestRegularFail("test_fail"), failures=has_length(1)
)
def test_addFailureGreaterThan64k(self) -> None:
"""
L{WorkerReporter} propagates test failures with large string representations.
"""
self.assertTestRun(
erroneous.TestAsynchronousFail("test_failGreaterThan64k"),
failures=has_length(1),
)
def test_addFailureTuple(self) -> None:
"""
L{WorkerReporter} propagates test failures from pyunit's TestCases.
"""
self.assertTestRun(pyunitcases.PyUnitTest("test_fail"), failures=has_length(1))
def test_addSkip(self) -> None:
"""
L{WorkerReporter} propagates skips.
"""
self.assertTestRun(
skipping.SynchronousSkipping("test_skip1"), skips=has_length(1)
)
def test_addSkipPyunit(self) -> None:
"""
L{WorkerReporter} propagates skips from L{unittest.TestCase} cases.
"""
self.assertTestRun(
pyunitcases.PyUnitTest("test_skip"),
skips=has_length(1),
)
def test_addExpectedFailure(self) -> None:
"""
L{WorkerReporter} propagates expected failures.
"""
self.assertTestRun(
skipping.SynchronousStrictTodo("test_todo1"), expectedFailures=has_length(1)
)
def test_addExpectedFailureGreaterThan64k(self) -> None:
"""
WorkerReporter propagates expected failures with large string representations.
"""
self.assertTestRun(
skipping.ExpectedFailure("test_expectedFailureGreaterThan64k"),
expectedFailures=has_length(1),
)
def test_addUnexpectedSuccess(self) -> None:
"""
L{WorkerReporter} propagates unexpected successes.
"""
self.assertTestRun(
skipping.SynchronousTodo("test_todo3"), unexpectedSuccesses=has_length(1)
)

View File

@@ -0,0 +1,147 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Tests for L{twisted.trial._dist.workertrial}.
"""
import errno
import sys
from io import BytesIO
from twisted.internet.testing import StringTransport
from twisted.protocols.amp import AMP
from twisted.trial._dist import (
_WORKER_AMP_STDIN,
_WORKER_AMP_STDOUT,
managercommands,
workercommands,
workertrial,
)
from twisted.trial._dist.workertrial import WorkerLogObserver, main
from twisted.trial.unittest import TestCase
class FakeAMP(AMP):
"""
A fake amp protocol.
"""
class WorkerLogObserverTests(TestCase):
"""
Tests for L{WorkerLogObserver}.
"""
def test_emit(self):
"""
L{WorkerLogObserver} forwards data to L{managercommands.TestWrite}.
"""
calls = []
class FakeClient:
def callRemote(self, method, **kwargs):
calls.append((method, kwargs))
observer = WorkerLogObserver(FakeClient())
observer.emit({"message": ["Some log"]})
self.assertEqual(calls, [(managercommands.TestWrite, {"out": "Some log"})])
class MainTests(TestCase):
"""
Tests for L{main}.
"""
def setUp(self):
self.readStream = BytesIO()
self.writeStream = BytesIO()
self.patch(
workertrial, "startLoggingWithObserver", self.startLoggingWithObserver
)
self.addCleanup(setattr, sys, "argv", sys.argv)
sys.argv = ["trial"]
def fdopen(self, fd, mode=None):
"""
Fake C{os.fdopen} implementation which returns C{self.readStream} for
the stdin fd and C{self.writeStream} for the stdout fd.
"""
if fd == _WORKER_AMP_STDIN:
self.assertEqual("rb", mode)
return self.readStream
elif fd == _WORKER_AMP_STDOUT:
self.assertEqual("wb", mode)
return self.writeStream
else:
raise AssertionError(f"Unexpected fd {fd!r}")
def startLoggingWithObserver(self, emit, setStdout):
"""
Override C{startLoggingWithObserver} for not starting logging.
"""
self.assertFalse(setStdout)
def test_empty(self):
"""
If no data is ever written, L{main} exits without writing data out.
"""
main(self.fdopen)
self.assertEqual(b"", self.writeStream.getvalue())
def test_forwardCommand(self):
"""
L{main} forwards data from its input stream to a L{WorkerProtocol}
instance which writes data to the output stream.
"""
client = FakeAMP()
clientTransport = StringTransport()
client.makeConnection(clientTransport)
client.callRemote(workercommands.Run, testCase="doesntexist")
self.readStream = clientTransport.io
self.readStream.seek(0, 0)
main(self.fdopen)
# Just brazenly encode irrelevant implementation details here, why
# not.
self.assertIn(b"StreamOpen", self.writeStream.getvalue())
def test_readInterrupted(self):
"""
If reading the input stream fails with a C{IOError} with errno
C{EINTR}, L{main} ignores it and continues reading.
"""
excInfos = []
class FakeStream:
count = 0
def read(oself, size):
oself.count += 1
if oself.count == 1:
raise OSError(errno.EINTR)
else:
excInfos.append(sys.exc_info())
return b""
self.readStream = FakeStream()
main(self.fdopen)
self.assertEqual(b"", self.writeStream.getvalue())
self.assertEqual([(None, None, None)], excInfos)
def test_otherReadError(self):
"""
L{main} only ignores C{IOError} with C{EINTR} errno: otherwise, the
error pops out.
"""
class FakeStream:
count = 0
def read(oself, size):
oself.count += 1
if oself.count == 1:
raise OSError("Something else")
return ""
self.readStream = FakeStream()
self.assertRaises(IOError, main, self.fdopen)