• Home
  • Features
  • Pricing
  • Docs
  • Announcements
  • Sign In

pyta-uoft / pyta / 36646375051

29 Sep 2026 11:39PM UTC coverage: 90.867% (-0.06%) from 90.925%
36646375051

push

github

web-flow
Added type annotations and fixed mypy errors in python-ta package (#1397)

256 of 269 new or added lines in 33 files covered. (95.17%)

2 existing lines in 2 files now uncovered.

3761 of 4139 relevant lines covered (90.87%)

17.67 hits per line

Source File
Press 'n' to go to next uncovered line, 'b' for previous

91.57
/packages/python-ta/src/python_ta/contracts/__init__.py
1
"""This module provides the functionality for PythonTA contracts.
2

3
Representation invariants, preconditions, and postconditions are parsed, compiled, and stored.
4
Below are some notes on how they are stored.
5
    - Representation invariants are stored in a class attribute __representation_invariants__
6
    as a list [(assertion, compiled)].
7
    - Preconditions are stored in an attribute __preconditions__ of the function as a list
8
    [(assertion, compiled)].
9
    - Postconditions are stored in an attribute __postconditions__ of the function as a list
10
    [(assertion, compiled, return_val_var_name)].
11
"""
12

13
from __future__ import annotations
20 ✔
14

15
import inspect
20 ✔
16
import logging
20 ✔
17
import re
20 ✔
18
import sys
20 ✔
19
import typing
20 ✔
20
from types import CodeType, FunctionType, ModuleType
20 ✔
21
from typing import (
20 ✔
22
    Any,
23
    Callable,
24
    Optional,
25
    TypeVar,
26
    Union,
27
    get_args,
28
    get_origin,
29
    overload,
30
)
31

32
import wrapt
20 ✔
33
from typeguard import CollectionCheckStrategy, TypeCheckError, check_type
20 ✔
34

35
# Configuration options
36

37
ENABLE_CONTRACT_CHECKING = True
20 ✔
38
"""
20 ✔
39
Set to True to enable contract checking.
40
"""
41

42
DEBUG_CONTRACTS = False
20 ✔
43
"""
20 ✔
44
Set to True to display debugging messages when checking contracts.
45
"""
46

47
RENAME_MAIN_TO_PYDEV_UMD = True
20 ✔
48
"""
20 ✔
49
Set to False to disable workaround for PyCharm's "Run File in Python Console" action.
50
In most cases you should not need to change this!
51
"""
52

53
STRICT_NUMERIC_TYPES = True
20 ✔
54
"""
20 ✔
55
Set to False to allow more specific numeric types to be accepted by more general type annotations.
56
"""
57

58
_PYDEV_UMD_NAME = "pydev_umd"
20 ✔
59

60

61
_DEFAULT_MAX_VALUE_LENGTH = 30
20 ✔
62
FUNCTION_RETURN_VALUE = "$return_value"
20 ✔
63

64

65
class PyTAContractError(Exception):
20 ✔
66
    """Error raised when a PyTA contract assertion is violated."""
67

68

69
def check_all_contracts(*mod_names: str, decorate_main: bool = True) -> None:
20 ✔
70
    """Automatically check contracts for all functions and classes in the given modules.
71

72
    By default (when called with no arguments), the current module is used.
73

74
    Args:
75
        *mod_names: The names of modules to check contracts for. These modules must have been
76
            previously imported.
77
        decorate_main: True if the module being run (where __name__ == '__main__') should
78
            have contracts checked.
79
    """
80
    if not ENABLE_CONTRACT_CHECKING:
20 ✔
81
        return
×
82

83
    modules: list[ModuleType | None] = []
20 ✔
84
    if decorate_main:
20 ✔
85
        mod_names = mod_names + ("__main__",)
×
86

87
        # Also add _PYDEV_UMD_NAME, handling when the file is being run in PyCharm
88
        # with the "Run in Python Console" action.
89
        if RENAME_MAIN_TO_PYDEV_UMD:
×
90
            mod_names = mod_names + (_PYDEV_UMD_NAME,)
×
91

92
    for module_name in mod_names:
20 ✔
93
        modules.append(sys.modules.get(module_name, None))
20 ✔
94

95
    for module in modules:
20 ✔
96
        if not module:
20 ✔
97
            # Module name was passed in incorrectly.
98
            continue
×
99
        for name, value in inspect.getmembers(module):
20 ✔
100
            if inspect.isfunction(value) or inspect.isclass(value):
20 ✔
101
                module.__dict__[name] = check_contracts(value, module_names=set(mod_names))
20 ✔
102

103

104
# Wildcard Type Variable
105
Class = TypeVar("Class", bound=type)
20 ✔
106

107

108
@overload
109
def check_contracts(
110
    func: FunctionType,
111
    module_names: Optional[set[str]] = None,
112
    argument_types: bool = True,
113
    return_type: bool = True,
114
    preconditions: bool = True,
115
    postconditions: bool = True,
116
) -> FunctionType: ...
117

118

119
@overload
120
def check_contracts(
121
    func: Class,
122
    module_names: Optional[set[str]] = None,
123
    argument_types: bool = True,
124
    return_type: bool = True,
125
    preconditions: bool = True,
126
    postconditions: bool = True,
127
) -> Class: ...
128

129

130
def check_contracts(  # type: ignore[misc]
20 ✔
131
    func_or_class: Optional[Union[Class, FunctionType]] = None,
132
    *,
133
    module_names: Optional[set[str]] = None,
134
    argument_types: bool = True,
135
    return_type: bool = True,
136
    preconditions: bool = True,
137
    postconditions: bool = True,
138
) -> Union[Class, FunctionType]:
139
    """A decorator to enable contract checking for a function or class.
140

141
    When used with a class, all methods defined within the class have contract checking enabled.
142
    If module_names is not None, only functions or classes defined in a module whose name is in module_names are checked.
143

144
    When used with functions, `check_contracts` accepts four optional boolean keyword arguments to selectively disable checks when set to `False`:
145

146
    - `argument_types`: check parameter type annotations
147
    - `return_type`: check the return type annotation
148
    - `preconditions`: check preconditions
149
    - `postconditions`: check postconditions
150

151
    By default, all four checks are enabled. These arguments only affect functions, and are ignored when `check_contracts` is applied to a class.
152

153
    Example:
154
        >>> from python_ta.contracts import check_contracts
155
        >>> @check_contracts
156
        ... def divide(x: int, y: int) -> int:
157
        ...     \"\"\"Return x // y.
158
        ...
159
        ...     Preconditions:
160
        ...        - y != 0
161
        ...     \"\"\"
162
        ...     return x // y
163
    """
164

165
    @wrapt.decorator
20 ✔
166
    def _enable_function_contracts(wrapped, instance, args, kwargs):
20 ✔
167
        """A decorator that enables checking contracts for a function."""
168
        try:
20 ✔
169
            if instance is not None and inspect.isclass(instance):
20 ✔
170
                # This is a class method, so there is no instance.
171
                return _check_function_contracts(
20 ✔
172
                    wrapped,
173
                    None,
174
                    args,
175
                    kwargs,
176
                    argument_types_enabled=argument_types,
177
                    return_type_enabled=return_type,
178
                    preconditions_enabled=preconditions,
179
                    postconditions_enabled=postconditions,
180
                )
181
            else:
182
                return _check_function_contracts(
20 ✔
183
                    wrapped,
184
                    instance,
185
                    args,
186
                    kwargs,
187
                    argument_types_enabled=argument_types,
188
                    return_type_enabled=return_type,
189
                    preconditions_enabled=preconditions,
190
                    postconditions_enabled=postconditions,
191
                )
192
        except PyTAContractError as e:
20 ✔
193
            raise AssertionError(str(e)) from None
20 ✔
194

195
    # Optional Arguments passed to the decorator
196
    if func_or_class is None:
20 ✔
197
        return wrapt.PartialCallableObjectProxy(
20 ✔
198
            check_contracts,
199
            module_names=module_names,
200
            argument_types=argument_types,
201
            return_type=return_type,
202
            preconditions=preconditions,
203
            postconditions=postconditions,
204
        )
205

206
    if not ENABLE_CONTRACT_CHECKING:
20 ✔
207
        return func_or_class
20 ✔
208

209
    if module_names is not None and func_or_class.__module__ not in module_names:
20 ✔
210
        _debug(
20 ✔
211
            f"Warning: skipping contract check for {func_or_class.__name__} defined in {func_or_class.__module__} because module is not included as an argument."
212
        )
213
        return func_or_class
20 ✔
214
    elif inspect.isroutine(func_or_class):
20 ✔
215
        return _enable_function_contracts(func_or_class)
20 ✔
216
    elif inspect.isclass(func_or_class):
20 ✔
217
        add_class_invariants(func_or_class)
20 ✔
218
        return func_or_class  # type: ignore[return-value]
20 ✔
219
    else:
220
        # Default action
221
        return func_or_class
×
222

223

224
def add_class_invariants(klass: type[Class]) -> None:
20 ✔
225
    """Modify the given class to check representation invariants and method contracts."""
226
    if not ENABLE_CONTRACT_CHECKING or "__representation_invariants__" in vars(klass):
20 ✔
227
        # This means the class has already been decorated
228
        return
×
229

230
    _set_invariants(klass)
20 ✔
231

232
    klass_mod = _get_module(klass)
20 ✔
233
    cls_annotations: Optional[dict[str, Any]] = (
20 ✔
234
        None  # This is a cached value set the first time new_setattr is called
235
    )
236

237
    def new_setattr(self: Class, name: str, value: Any) -> None:
20 ✔
238
        """Set the value of the given attribute on self to the given value.
239

240
        Check representation invariants for this class when not within an instance method of the class.
241
        """
242
        if not ENABLE_CONTRACT_CHECKING:
20 ✔
243
            super(klass, self).__setattr__(name, value)
20 ✔
244
            return
20 ✔
245

246
        nonlocal cls_annotations
247
        if cls_annotations is None:
20 ✔
248
            cls_annotations = typing.get_type_hints(klass, localns=klass_mod.__dict__)
20 ✔
249

250
        if name in cls_annotations:
20 ✔
251
            try:
20 ✔
252
                _debug(f"Checking type of attribute {name} for {klass.__qualname__} instance")
20 ✔
253
                check_type(
20 ✔
254
                    value,
255
                    cls_annotations[name],
256
                    collection_check_strategy=CollectionCheckStrategy.ALL_ITEMS,
257
                )
258
            except TypeCheckError:
20 ✔
259
                raise AssertionError(
20 ✔
260
                    f"Value {_display_value(value)} for attribute {name} did not match expected type "
261
                    f"{_display_annotation(cls_annotations[name])}"
262
                ) from None
263
        original_attr_value_exists = False
20 ✔
264
        original_attr_value = None
20 ✔
265
        if hasattr(self, name):
20 ✔
266
            original_attr_value_exists = True
20 ✔
267
            original_attr_value = super(klass, self).__getattribute__(name)  # type: ignore[arg-type]
20 ✔
268
        super(klass, self).__setattr__(name, value)  # type: ignore[arg-type]
20 ✔
269
        current_frame = inspect.currentframe()
20 ✔
270
        if current_frame is None or current_frame.f_back is None:
20 ✔
NEW
271
            return
×
272
        frame_locals = current_frame.f_back.f_locals
20 ✔
273
        caller_self = frame_locals.get("self")
20 ✔
274
        if not isinstance(caller_self, type(self)):
20 ✔
275
            # Only validating if the attribute is not being set in a instance/class method
276
            # AND caller_self is an instance of self's type
277
            if klass_mod is not None:
20 ✔
278
                try:
20 ✔
279
                    _check_invariants(self, klass, klass_mod.__dict__)
20 ✔
280
                except PyTAContractError as e:
20 ✔
281
                    if original_attr_value_exists:
20 ✔
282
                        super(klass, self).__setattr__(name, original_attr_value)  # type: ignore[arg-type]
20 ✔
283
                    else:
284
                        super(klass, self).__delattr__(name)  # type: ignore[arg-type]
20 ✔
285
                    raise AssertionError(str(e)) from None
20 ✔
286
        elif caller_self is not self:
20 ✔
287
            # Keep track of mutations to instances that are of the same type as caller_self (and are also not `self`)
288
            # to enforce RIs on them only after the caller function returns.
289
            caller_klass = type(caller_self)
20 ✔
290
            if hasattr(caller_klass, "__mutated_instances__"):
20 ✔
291
                mutated_instances = getattr(caller_klass, "__mutated_instances__")
20 ✔
292
                if self not in mutated_instances:
20 ✔
293
                    mutated_instances.append(self)
20 ✔
294

295
    for attr, value in vars(klass).items():
20 ✔
296
        # Skip built-in __annotate_func__, which was introduced in Python 3.14
297
        if attr == "__annotate_func__":
20 ✔
298
            continue
4 ✔
299
        if inspect.isroutine(value):
20 ✔
300
            if isinstance(value, (staticmethod, classmethod)):
20 ✔
301
                # Don't check rep invariants for staticmethod and classmethod
302
                setattr(klass, attr, check_contracts(value))
20 ✔
303
            else:
304
                setattr(klass, attr, _instance_method_wrapper(value, klass))
20 ✔
305

306
    klass.__setattr__ = new_setattr  # type: ignore[assignment, method-assign]
20 ✔
307

308

309
def _check_function_contracts(
20 ✔
310
    wrapped,
311
    instance,
312
    args,
313
    kwargs,
314
    argument_types_enabled: bool = True,
315
    return_type_enabled: bool = True,
316
    preconditions_enabled: bool = True,
317
    postconditions_enabled: bool = True,
318
) -> Any:
319
    params = wrapped.__code__.co_varnames[: wrapped.__code__.co_argcount]
20 ✔
320
    if instance is not None:
20 ✔
321
        klass_mod = _get_module(type(instance))
20 ✔
322
        annotations = typing.get_type_hints(wrapped, globalns=klass_mod.__dict__)
20 ✔
323
    else:
324
        annotations = typing.get_type_hints(wrapped)
20 ✔
325
    args_with_self = args if instance is None else (instance,) + args
20 ✔
326

327
    if argument_types_enabled:
20 ✔
328
        # Check function parameter types
329
        for arg, param in zip(args_with_self, params):
20 ✔
330
            if param in annotations:
20 ✔
331
                try:
20 ✔
332
                    _debug(f"Checking type of parameter {param} in call to {wrapped.__qualname__}")
20 ✔
333
                    if STRICT_NUMERIC_TYPES:
20 ✔
334
                        check_type_strict(param, arg, annotations[param])
20 ✔
335
                    else:
336
                        check_type(arg, annotations[param])
20 ✔
337
                except (TypeError, TypeCheckError):
20 ✔
338
                    additional_suggestions = _get_argument_suggestions(arg, annotations[param])
20 ✔
339

340
                    raise PyTAContractError(
20 ✔
341
                        f"Argument value {_display_value(arg)} for {wrapped.__name__} parameter {param} "
342
                        f"did not match expected type {_display_annotation(annotations[param])}"
343
                        + (f"\n{additional_suggestions}" if additional_suggestions else "")
344
                    )
345

346
    function_locals = dict(zip(params, args_with_self))
20 ✔
347

348
    # Check bounded function
349
    if hasattr(wrapped, "__self__"):
20 ✔
350
        target = wrapped.__func__
20 ✔
351
    else:
352
        target = wrapped
20 ✔
353

354
    # Check function preconditions
355
    if not hasattr(target, "__preconditions__") and preconditions_enabled:
20 ✔
356
        target_preconditions: list[tuple[str, CodeType]] = []
20 ✔
357
        preconditions = parse_assertions(wrapped)
20 ✔
358
        for precondition in preconditions:
20 ✔
359
            try:
20 ✔
360
                compiled = compile(precondition, "<string>", "eval")
20 ✔
361
            except:
20 ✔
362
                _debug(
20 ✔
363
                    f"Warning: precondition {precondition} could not be parsed as a valid Python expression"
364
                )
365
                continue
20 ✔
366
            target_preconditions.append((precondition, compiled))
20 ✔
367
        target.__preconditions__ = target_preconditions
20 ✔
368

369
    if ENABLE_CONTRACT_CHECKING and preconditions_enabled:
20 ✔
370
        _check_assertions(wrapped, function_locals)
20 ✔
371

372
    # Check return type
373
    r = wrapped(*args, **kwargs)
20 ✔
374
    if return_type_enabled and "return" in annotations:
20 ✔
375
        return_type = annotations["return"]
20 ✔
376
        try:
20 ✔
377
            _debug(f"Checking return type from call to {wrapped.__qualname__}")
20 ✔
378
            if STRICT_NUMERIC_TYPES:
20 ✔
379
                check_type_strict("return", r, return_type)
20 ✔
380
            else:
381
                check_type(r, return_type)
20 ✔
382
        except (TypeError, TypeCheckError):
20 ✔
383
            raise PyTAContractError(
20 ✔
384
                f"Return value {_display_value(r)} for {wrapped.__name__} did not match "
385
                f"expected type {_display_annotation(return_type)}"
386
            )
387

388
    # Check function postconditions
389
    if postconditions_enabled and not hasattr(target, "__postconditions__"):
20 ✔
390
        target_postconditions: list[tuple[str, CodeType, str]] = []
20 ✔
391
        return_val_var_name = _get_legal_return_val_var_name(
20 ✔
392
            {**wrapped.__globals__, **function_locals}
393
        )
394
        postconditions = parse_assertions(wrapped, parse_token="Postcondition")
20 ✔
395
        for postcondition in postconditions:
20 ✔
396
            assertion = _replace_return_val_assertion(postcondition, return_val_var_name)
20 ✔
397
            try:
20 ✔
398
                compiled = compile(assertion, "<string>", "eval")
20 ✔
399
            except:
×
400
                _debug(
×
401
                    f"Warning: postcondition {postcondition} could not be parsed as a valid Python expression"
402
                )
403
                continue
×
404
            target_postconditions.append((postcondition, compiled, return_val_var_name))
20 ✔
405
        target.__postconditions__ = target_postconditions
20 ✔
406

407
    if ENABLE_CONTRACT_CHECKING and postconditions_enabled:
20 ✔
408
        _check_assertions(
20 ✔
409
            wrapped,
410
            function_locals,
411
            function_return_val=r,
412
            condition_type="postcondition",
413
        )
414

415
    return r
20 ✔
416

417

418
def check_type_strict(argname: str, value: Any, expected_type: type) -> None:
20 ✔
419
    """Ensure that `value` matches ``expected_type`` with strict type checking.
420

421
    This function enforces strict type distinctions within the numeric hierarchy (bool, int, float,
422
    complex), ensuring that the type of value is exactly the same as expected_type.
423
    """
424
    if not ENABLE_CONTRACT_CHECKING:
20 ✔
425
        return
20 ✔
426
    try:
20 ✔
427
        _check_inner_type(argname, value, expected_type)
20 ✔
428
    except (TypeError, TypeCheckError):
20 ✔
429
        raise TypeError(f"type of {argname} must be {expected_type}; got {value} instead")
20 ✔
430

431

432
def _check_inner_type(argname: str, value: Any, expected_type: type) -> None:
20 ✔
433
    """Recursively checks if `value` matches `expected_type` for strict type validation, specifically supports checking
434
    collections (list[int], dicts[float]) and Union types (bool | int).
435
    """
436
    inner_types = get_args(expected_type)
20 ✔
437
    outer_type = get_origin(expected_type)
20 ✔
438
    if outer_type is None:
20 ✔
439
        if (
20 ✔
440
            (type(value) is bool and expected_type in {int, float, complex})
441
            or (type(value) is int and expected_type in {float, complex})
442
            or (type(value) is float and expected_type is complex)
443
        ):
444
            raise TypeError(
20 ✔
445
                f"type of {argname} must be {expected_type}; got {type(value).__name__} instead"
446
            )
447
        else:
448
            check_type(
20 ✔
449
                value, expected_type, collection_check_strategy=CollectionCheckStrategy.ALL_ITEMS
450
            )
451
    elif outer_type is typing.Union:
20 ✔
452
        for inner_type in inner_types:
20 ✔
453
            try:
20 ✔
454
                _check_inner_type(argname, value, inner_type)
20 ✔
455
                return
20 ✔
456
            except (TypeError, TypeCheckError):
20 ✔
457
                pass
20 ✔
458
        raise TypeError(f"type of {argname} must be {expected_type}; got {value} instead")
20 ✔
459
    elif outer_type in {list, set}:
20 ✔
460
        if isinstance(value, outer_type):
20 ✔
461
            for item in value:
20 ✔
462
                _check_inner_type(argname, item, inner_types[0])
20 ✔
463
        else:
464
            raise TypeError(f"type of {argname} must be {expected_type}; got {value} instead")
20 ✔
465
    elif outer_type is dict:
20 ✔
466
        if isinstance(value, dict):
20 ✔
467
            for key, item in value.items():
20 ✔
468
                _check_inner_type(argname, key, inner_types[0])
20 ✔
469
                _check_inner_type(argname, item, inner_types[1])
20 ✔
470
        else:
471
            raise TypeError(f"type of {argname} must be {expected_type}; got {value} instead")
20 ✔
472
    elif outer_type is tuple:
20 ✔
473
        if isinstance(value, tuple) and len(inner_types) == 2 and inner_types[1] is Ellipsis:
20 ✔
474
            for item in value:
20 ✔
475
                _check_inner_type(argname, item, inner_types[0])
20 ✔
476
        elif isinstance(value, tuple) and len(value) == len(inner_types):
20 ✔
477
            for item, inner_type in zip(value, inner_types):
20 ✔
478
                _check_inner_type(argname, item, inner_type)
20 ✔
479
        else:
480
            raise TypeError(f"type of {argname} must be {expected_type}; got {value} instead")
20 ✔
481
    else:
482
        check_type(
20 ✔
483
            value, expected_type, collection_check_strategy=CollectionCheckStrategy.ALL_ITEMS
484
        )
485

486

487
def _get_argument_suggestions(arg: Any, annotation: type) -> str:
20 ✔
488
    """Returns potential suggestions for the given arg and its annotation"""
489
    try:
20 ✔
490
        if isinstance(arg, type) and issubclass(arg, annotation):
20 ✔
491
            return "Did you mean {cls}(...) instead of {cls}?".format(cls=arg.__name__)
20 ✔
492
    except TypeError:
20 ✔
493
        pass
20 ✔
494

495
    return ""
20 ✔
496

497

498
def _instance_method_wrapper(wrapped: Callable, klass: type) -> Callable:
20 ✔
499
    @wrapt.decorator
20 ✔
500
    def wrapper(wrapped, instance, args, kwargs):
20 ✔
501
        # Create an accumulator to store the instances mutated across this function call.
502
        # Store and restore existing mutated instance lists in case the instance method
503
        # executes another instance method.
504
        instance_klass = type(instance)
20 ✔
505
        mutated_instances_to_restore = None
20 ✔
506
        if hasattr(instance_klass, "__mutated_instances__"):
20 ✔
507
            mutated_instances_to_restore = getattr(instance_klass, "__mutated_instances__")
20 ✔
508
        setattr(instance_klass, "__mutated_instances__", [])
20 ✔
509

510
        try:
20 ✔
511
            r = _check_function_contracts(wrapped, instance, args, kwargs)
20 ✔
512
            if _instance_init_in_callstack(instance):
20 ✔
513
                return r
20 ✔
514
            _check_class_type_annotations(klass, instance)
20 ✔
515
            klass_mod = _get_module(klass)
20 ✔
516
            if klass_mod is not None and ENABLE_CONTRACT_CHECKING:
20 ✔
517
                _check_invariants(instance, klass, klass_mod.__dict__)
20 ✔
518

519
                # Additionally check RI violations on PyTA-decorated instances that were mutated
520
                # across the function call.
521
                mutated_instances = getattr(instance_klass, "__mutated_instances__", [])
20 ✔
522
                for mutated_instance in mutated_instances:
20 ✔
523
                    # Mutated instances may be of parent class types so the invariants to check should also be
524
                    # for the parent class and not the child class.
525
                    mutated_instance_klass = type(mutated_instance)
20 ✔
526
                    mutated_instance_klass_mod = _get_module(mutated_instance_klass)
20 ✔
527
                    _check_invariants(
20 ✔
528
                        mutated_instance,
529
                        mutated_instance_klass,
530
                        mutated_instance_klass_mod.__dict__,
531
                    )
532
        except PyTAContractError as e:
20 ✔
533
            raise AssertionError(str(e)) from None
20 ✔
534
        else:
535
            return r
20 ✔
536
        finally:
537
            if mutated_instances_to_restore is None:
20 ✔
538
                delattr(instance_klass, "__mutated_instances__")
20 ✔
539
            else:
540
                setattr(instance_klass, "__mutated_instances__", mutated_instances_to_restore)
20 ✔
541

542
    return wrapper(wrapped)
20 ✔
543

544

545
def _instance_init_in_callstack(instance: Any) -> bool:
20 ✔
546
    """Return whether instance's init is part of the current callstack
547

548
    Note: due to the nature of the check, externally defined __init__ functions with
549
    'self' defined as the first parameter may pass this check.
550
    """
551
    current_frame = inspect.currentframe()
20 ✔
552
    if current_frame is None:
20 ✔
NEW
553
        return False
×
554
    frame = current_frame.f_back
20 ✔
555
    while frame is not None:
20 ✔
556
        frame_context_name = inspect.getframeinfo(frame).function
20 ✔
557
        frame_context_self = frame.f_locals.get("self")
20 ✔
558
        frame_context_vars = frame.f_code.co_varnames
20 ✔
559
        if (
20 ✔
560
            frame_context_name == "__init__"
561
            and frame_context_self is instance
562
            and frame_context_vars[0] == "self"
563
        ):
564
            return True
20 ✔
565
        frame = frame.f_back
20 ✔
566
    return False
20 ✔
567

568

569
def _check_class_type_annotations(klass: type, instance: Any) -> None:
20 ✔
570
    """Check that the type annotations for the class still hold.
571

572
    Precondition:
573
        - isinstance(instance, klass)
574
    """
575
    klass_mod = _get_module(klass)
20 ✔
576
    cls_annotations = typing.get_type_hints(klass, localns=klass_mod.__dict__)
20 ✔
577

578
    for attr, annotation in cls_annotations.items():
20 ✔
579
        _debug(f"Checking type of attribute {attr} for {klass.__qualname__} instance")
20 ✔
580
        if not hasattr(instance, attr):
20 ✔
581
            raise AssertionError(
20 ✔
582
                f"Attribute {attr} is not defined for this {klass.__qualname__} instance, but "
583
                f"is expected to have type {_display_annotation(annotation)}"
584
            )
585
        value = getattr(instance, attr)
20 ✔
586
        try:
20 ✔
587
            check_type(
20 ✔
588
                value, annotation, collection_check_strategy=CollectionCheckStrategy.ALL_ITEMS
589
            )
590
        except TypeCheckError:
20 ✔
591
            raise AssertionError(
20 ✔
592
                f"Value {_display_value(value)} for attribute {attr} did not match expected type "
593
                f"{_display_annotation(annotation)}"
594
            )
595

596

597
def _check_invariants(instance, klass: type, global_scope: dict) -> None:
20 ✔
598
    """Check that the representation invariants for the instance are satisfied."""
599
    if hasattr(instance, "__pyta_currently_checking"):
20 ✔
600
        # If already checking invariants for this instance, skip to avoid infinite recursion
601
        return
20 ✔
602

603
    super(type(instance), instance).__setattr__("__pyta_currently_checking", True)
20 ✔
604

605
    rep_invariants: set[tuple[str, CodeType]] = getattr(
20 ✔
606
        klass, "__representation_invariants__", set()
607
    )
608

609
    try:
20 ✔
610
        for invariant, compiled in rep_invariants:
20 ✔
611
            try:
20 ✔
612
                _debug(
20 ✔
613
                    "Checking representation invariant for "
614
                    f"{instance.__class__.__qualname__}: {invariant}"
615
                )
616
                check = eval(compiled, {**global_scope, "self": instance})
20 ✔
617
            except AssertionError as e:
20 ✔
618
                raise AssertionError(str(e)) from None
20 ✔
619
            except NameError as e:
20 ✔
620
                # Get the missing name
621
                missing = getattr(e, "name", None)
20 ✔
622
                if missing is None:
20 ✔
623
                    # Failsafe for version 3.9
624
                    message = re.search(r"name '(.+?)' is not defined", str(e))
×
625
                    if message:
×
626
                        missing = message.group(1)
×
627

628
                # Check if missing name is an attribute
629
                if missing is not None and hasattr(instance, missing):
20 ✔
630
                    print(
20 ✔
631
                        f"[WARNING] Could not find variable `{missing}` when evaluating representation invariant. Did you mean `self.{missing}`?",
632
                        file=sys.stderr,
633
                    )
634
                else:
635
                    _debug(f"Warning: could not evaluate representation invariant: {invariant}")
20 ✔
636
            except:
×
637
                _debug(f"Warning: could not evaluate representation invariant: {invariant}")
×
638
            else:
639
                if not check:
20 ✔
640
                    curr_attributes = ", ".join(
20 ✔
641
                        f"{k}: {_display_value(v)}"
642
                        for k, v in vars(instance).items()
643
                        if k != "__pyta_currently_checking"
644
                    )
645

646
                    curr_attributes = "{" + curr_attributes + "}"
20 ✔
647

648
                    raise PyTAContractError(
20 ✔
649
                        f'{instance.__class__.__name__} representation invariant "{invariant}" was violated for'
650
                        f" instance attributes {curr_attributes}"
651
                    )
652

653
    finally:
654
        delattr(instance, "__pyta_currently_checking")
20 ✔
655

656

657
def _get_legal_return_val_var_name(var_dict: dict) -> str:
20 ✔
658
    """
659
    Add '_' to the end of __function_return_value__ until a variable name that has not been used for any other
660
    variable in the function's scope is created. This is used to refer to the function's return value when evaluating
661
    postconditions.
662
    """
663
    legal_var_name = "__function_return_value__"
20 ✔
664

665
    while legal_var_name in var_dict:
20 ✔
666
        legal_var_name += "_"
×
667

668
    return legal_var_name
20 ✔
669

670

671
def _replace_return_val_assertion(assertion: str, return_val_var_name: str) -> str:
20 ✔
672
    """
673
    Replace FUNCTION_RETURN_VALUE in the assertion with the legal python variable name generated and return the new
674
    assertion. If FUNCTION_RETURN_VALUE does not appear in assertion, then simply return the original assertion.
675

676
    Precondition: If FUNCTION_RETURN_VALUE is in assertion, then return_val_var_name is not None
677
    """
678

679
    if FUNCTION_RETURN_VALUE in assertion:
20 ✔
680
        return assertion.replace(FUNCTION_RETURN_VALUE, return_val_var_name)
20 ✔
681
    return assertion
20 ✔
682

683

684
def _check_assertions(
20 ✔
685
    wrapped: Callable[..., Any],
686
    function_locals: dict,
687
    condition_type: str = "precondition",
688
    function_return_val: Any = None,
689
) -> None:
690
    """Check that the given assertions are still satisfied."""
691
    # Check bounded function
692
    if hasattr(wrapped, "__self__"):
20 ✔
693
        target = wrapped.__func__  # type: ignore[attr-defined]
20 ✔
694
    else:
695
        target = wrapped
20 ✔
696
    assertions: list[Any] = []
20 ✔
697
    if condition_type == "precondition":
20 ✔
698
        assertions = target.__preconditions__
20 ✔
699
    elif condition_type == "postcondition":
20 ✔
700
        assertions = target.__postconditions__
20 ✔
701
    for assertion_str, compiled, *return_val_var_name in assertions:
20 ✔
702
        return_val_dict = {}
20 ✔
703
        if condition_type == "postcondition":
20 ✔
704
            return_val_dict = {return_val_var_name[0]: function_return_val}
20 ✔
705
        try:
20 ✔
706
            _debug(f"Checking {condition_type} for {wrapped.__qualname__}: {assertion_str}")
20 ✔
707
            check = eval(compiled, {**wrapped.__globals__, **function_locals, **return_val_dict})
20 ✔
708
        except AssertionError as e:
20 ✔
709
            raise AssertionError(str(e)) from None
20 ✔
710
        except:
×
711
            _debug(f"Warning: could not evaluate {condition_type}: {assertion_str}")
×
712
        else:
713
            if not check:
20 ✔
714
                arg_string = ", ".join(
20 ✔
715
                    f"{k}: {_display_value(v)}" for k, v in function_locals.items()
716
                )
717
                arg_string = "{" + arg_string + "}"
20 ✔
718

719
                return_val_string = ""
20 ✔
720

721
                if condition_type == "postcondition":
20 ✔
722
                    return_val_string = f" and return value {function_return_val}"
20 ✔
723
                raise PyTAContractError(
20 ✔
724
                    f'{wrapped.__name__} {condition_type} "{assertion_str}" was '
725
                    f"violated for arguments {arg_string}{return_val_string}"
726
                )
727

728

729
def parse_assertions(obj: Any, parse_token: str = "Precondition") -> list[str]:
20 ✔
730
    """Return a list of preconditions/postconditions/representation invariants parsed from the given entity's docstring.
731

732
    Uses parse_token to determine what to look for. parse_token defaults to Precondition.
733

734
    Currently only supports two forms:
735

736
    1. A single line of the form "<parse_token>: <cond>"
737
    2. A group of lines starting with "<parse_token>s:", where each subsequent
738
       line is of the form "- <cond>". Each line is considered a separate condition.
739
       The lines can be separated by blank lines, but no other text.
740
    """
741
    if hasattr(obj, "doc_node") and obj.doc_node is not None:
20 ✔
742
        # Check if obj is an astroid node
743
        docstring = obj.doc_node.value
20 ✔
744
    else:
745
        docstring = getattr(obj, "__doc__") or ""
20 ✔
746
    lines = [line.strip() for line in docstring.split("\n")]
20 ✔
747
    assertion_lines = [
20 ✔
748
        i for i, line in enumerate(lines) if line.lower().startswith(parse_token.lower())
749
    ]
750

751
    if assertion_lines == []:
20 ✔
752
        return []
20 ✔
753

754
    first = assertion_lines[0]
20 ✔
755

756
    if lines[first].startswith(parse_token + ":"):
20 ✔
757
        return [lines[first][len(parse_token + ":") :].strip()]
20 ✔
758
    elif lines[first].startswith(parse_token + "s:"):
20 ✔
759
        assertions: list[str] = []
20 ✔
760
        for line in lines[first + 1 :]:
20 ✔
761
            if line.startswith("-"):
20 ✔
762
                assertion = line[1:].strip()
20 ✔
763
                if hasattr(obj, "__qualname__"):
20 ✔
764
                    _debug(f"Adding assertion to {obj.__qualname__}: {assertion}")
20 ✔
765
                assertions.append(assertion)
20 ✔
766
            elif line != "":
20 ✔
767
                break
×
768
        return assertions
16 ✔
769
    else:
770
        return []
×
771

772

773
def _display_value(value: Any, max_length: int = _DEFAULT_MAX_VALUE_LENGTH) -> str:
20 ✔
774
    """Return a human-friendly representation of the given value.
775

776
    If DEBUG_CONTRACTS is False, truncate long strings to max_length characters.
777

778
    Preconditions:
779
        - max_length >= 5
780
    """
781
    s = repr(value)
20 ✔
782
    if not DEBUG_CONTRACTS and len(s) > max_length:
20 ✔
783
        i = (max_length - 3) // 2
×
784
        return s[:i] + "..." + s[-i:]
×
785
    else:
786
        return s
20 ✔
787

788

789
def _display_annotation(annotation: Any) -> str:
20 ✔
790
    """Return a human-friendly representation of the given type annotation.
791

792
    >>> _display_annotation(int)
793
    'int'
794
    >>> _display_annotation(list[int])
795
    'list[int]'
796
    >>> from typing import List
797
    >>> _display_annotation(List[int])
798
    'typing.List[int]'
799
    """
800
    if annotation is type(None):  # Use 'None' instead of 'NoneType'
20 ✔
801
        return "None"
×
802
    if hasattr(annotation, "__origin__"):  # Generic type annotations
20 ✔
803
        return repr(annotation)
20 ✔
804
    elif hasattr(annotation, "__name__"):
20 ✔
805
        return annotation.__name__
20 ✔
806
    else:
807
        return repr(annotation)
×
808

809

810
def _get_module(obj: Any) -> ModuleType:
20 ✔
811
    """Return the module where obj was defined (normally obj.__module__).
812

813
    NOTE: this function defines a special case when using PyCharm and the file
814
    defining the object is "Run in Python Console". In this case, the pydevd runner
815
    renames the '__main__' module to 'pydev_umd', and so we need to access that
816
    module instead. This behaviour can be disabled by setting RENAME_MAIN_TO_PYDEV_UMD
817
    to False.
818
    """
819
    module_name = obj.__module__
20 ✔
820
    module = sys.modules[module_name]
20 ✔
821

822
    if (
20 ✔
823
        module_name != "__main__"
824
        or not RENAME_MAIN_TO_PYDEV_UMD
825
        or _PYDEV_UMD_NAME not in sys.modules
826
    ):
827
        return module
20 ✔
828

829
    # Get a function/class name to check whether it is defined in the module
830
    if isinstance(obj, (FunctionType, type)):
×
831
        name = obj.__name__
×
832
    else:
833
        # For any other type of object, be conservative and just return the module
834
        return module
×
835

836
    if name in vars(module):
×
837
        return module
×
838
    else:
839
        return sys.modules[_PYDEV_UMD_NAME]
×
840

841

842
def _debug(msg: str) -> None:
20 ✔
843
    """Display a debugging message.
844

845
    Do nothing if DEBUG_CONTRACTS is False.
846
    """
847
    if not DEBUG_CONTRACTS:
20 ✔
848
        return
20 ✔
849
    logging.basicConfig(format="[%(levelname)s] %(message)s", level=logging.DEBUG)
20 ✔
850
    logging.debug(msg)
20 ✔
851

852

853
def _set_invariants(klass: type) -> None:
20 ✔
854
    """Retrieve and set the representation invariants of this class"""
855
    # Update representation invariants from this class' docstring and those of its superclasses.
856
    rep_invariants: list[tuple[str, CodeType]] = []
20 ✔
857

858
    # Iterate over all inherited classes except builtins
859
    for cls in reversed(klass.__mro__):
20 ✔
860
        if "__representation_invariants__" in cls.__dict__:
20 ✔
861
            rep_invariants.extend(cls.__representation_invariants__)  # type: ignore[attr-defined]
20 ✔
862
        elif cls.__module__ != "builtins":
20 ✔
863
            assertions = parse_assertions(cls, parse_token="Representation Invariant")
20 ✔
864
            # Try compiling assertions
865
            for assertion in assertions:
20 ✔
866
                try:
20 ✔
867
                    compiled = compile(assertion, "<string>", "eval")
20 ✔
868
                except:
×
869
                    _debug(
×
870
                        f"Warning: representation invariant {assertion} could not be parsed as a valid Python expression"
871
                    )
872
                    continue
×
873
                rep_invariants.append((assertion, compiled))
20 ✔
874

875
    setattr(klass, "__representation_invariants__", rep_invariants)
20 ✔
876

877

878
def validate_invariants(obj: object) -> None:
20 ✔
879
    """Check that the representation invariants of obj are satisfied."""
880
    klass = obj.__class__
20 ✔
881
    klass_mod = _get_module(klass)
20 ✔
882

883
    try:
20 ✔
884
        _check_invariants(obj, klass, klass_mod.__dict__)
20 ✔
885
    except PyTAContractError as e:
20 ✔
886
        raise AssertionError(str(e)) from None
20 ✔
STATUS · Troubleshooting · Open an Issue · Sales · Support · CAREERS · ENTERPRISE · START FREE TRIAL · SCHEDULE DEMO
ANNOUNCEMENTS · TWITTER · TOS & SLA · Supported CI Services · What's a CI service? · Automated Testing

© 2026 Coveralls, Inc