diff --git a/openapi_spec_validator/validation/validators.py b/openapi_spec_validator/validation/validators.py index 796db96b..b9d99654 100644 --- a/openapi_spec_validator/validation/validators.py +++ b/openapi_spec_validator/validation/validators.py @@ -2,9 +2,9 @@ import logging import warnings +from collections.abc import Iterable from collections.abc import Iterator from collections.abc import Mapping -from functools import lru_cache from typing import cast from jsonschema.exceptions import ValidationError @@ -19,6 +19,7 @@ from openapi_spec_validator.schemas.types import AnySchema from openapi_spec_validator.settings import OpenAPISpecValidatorSettings from openapi_spec_validator.validation import keywords +from openapi_spec_validator.validation.caches import CachedIterable from openapi_spec_validator.validation.decorators import unwraps_iter from openapi_spec_validator.validation.decorators import wraps_cached_iter from openapi_spec_validator.validation.decorators import wraps_errors @@ -67,6 +68,7 @@ def __init__( self.keyword_validators_registry = KeywordValidatorRegistry( self.keyword_validators ) + self._cached_errors: CachedIterable[ValidationError] | None = None def validate(self) -> None: for err in self.iter_errors(): @@ -83,15 +85,19 @@ def root_validator(self) -> keywords.RootValidator: self.keyword_validators_registry["__root__"], ) - @unwraps_iter - @lru_cache(maxsize=None) @wraps_cached_iter @wraps_errors - def iter_errors(self) -> Iterator[ValidationError]: + def _iter_errors(self) -> Iterator[ValidationError]: yield from self.schema_validator.iter_errors(self.schema) yield from self.root_validator(self.schema_path) + @unwraps_iter + def iter_errors(self) -> Iterable[ValidationError]: + if getattr(self, "_cached_errors", None) is None: + self._cached_errors = self._iter_errors() + return self._cached_errors + class OpenAPIV2SpecValidator(SpecValidator): schema_validator = openapi_v2_schema_validator diff --git a/tests/integration/validation/test_validators.py b/tests/integration/validation/test_validators.py index 17a347d2..e3d44f10 100644 --- a/tests/integration/validation/test_validators.py +++ b/tests/integration/validation/test_validators.py @@ -665,3 +665,33 @@ def test_failed(self, factory, spec_file): with pytest.raises(OpenAPIValidationError): OpenAPIV31SpecValidator(spec, base_uri=spec_url).validate() + + +def test_validator_iter_errors_instance_caching(): + spec = { + "openapi": "3.0.0", + "info": {"title": "Sample", "version": "1.0.0"}, + "paths": {}, + } + validator = OpenAPIV30SpecValidator(spec) + errors_1 = list(validator.iter_errors()) + errors_2 = list(validator.iter_errors()) + assert errors_1 == errors_2 == [] + + +def test_validator_instance_garbage_collected(): + import gc + import weakref + + spec = { + "openapi": "3.0.0", + "info": {"title": "Sample", "version": "1.0.0"}, + "paths": {}, + } + validator = OpenAPIV30SpecValidator(spec) + list(validator.iter_errors()) + ref = weakref.ref(validator) + del validator + gc.collect() + + assert ref() is None