diff --git a/README.md b/README.md index edc91f7..f1524ac 100644 --- a/README.md +++ b/README.md @@ -40,7 +40,7 @@ With microbootstrap, you receive an application with lightweight built-in - `opentelemetry` - `logging` - `cors` -- `swagger` - with additional offline version support +- `swagger` - offline UI assets, OpenAPI security definitions, and optional Accept-version documentation - `health-checks` Those instruments can be bootstrapped for: @@ -594,6 +594,64 @@ Parameter descriptions: - `swagger_offline_docs` - A boolean value that, when set to True, allows the Swagger JS bundles to be accessed offline. This is because the service starts to host via static. - `swagger_extra_params` - Additional parameters to pass into the OpenAPI configuration. +#### OpenAPI security schemes + +Security schemes are disabled by default. They add reusable OpenAPI definitions under +`components.securitySchemes`; they do not authenticate requests or add global or operation-level security requirements. +Keep requirements and authentication in your application routes and dependencies. + +```python +from microbootstrap import ( + LitestarSettings, + OpenApiHttpSecurityScheme, +) + + +class Settings(LitestarSettings): + security_schemes: dict[str, OpenApiHttpSecurityScheme] = { + "serviceAuth": OpenApiHttpSecurityScheme(scheme="bearer", bearer_format="JWT"), + } +``` + +HTTP, API key, OAuth 2.0, and OpenID Connect definitions are supported by `SwaggerConfig`. Annotating a consumer +setting with a concrete scheme class intentionally rejects other kinds for that consumer. Python field names and +OpenAPI aliases are accepted; output uses canonical names such as `bearerFormat`, `in`, `tokenUrl`, and +`openIdConnectUrl`. A same-named definition must be identical to the service-owned definition or schema generation +raises `ValueError`. + +The definitions are added to the framework's normal schema. Repeated schema reads remain stable; after correcting +a configuration conflict, rebuild the application. + +#### API-version documentation + +Version documentation is disabled by default. It describes supported versions in OpenAPI without implementing +runtime version negotiation. + +```python +from microbootstrap import LitestarSettings, OpenApiOperationVersionOverride, OpenApiVersionDocsConfig + + +class Settings(LitestarSettings): + openapi_version_docs: OpenApiVersionDocsConfig | None = OpenApiVersionDocsConfig( + vendor_media_type="application/vnd.example+json", + supported_versions=("1.0",), + operation_versions=( + OpenApiOperationVersionOverride(path="/widgets", method="post", supported_versions=("2.0",)), + OpenApiOperationVersionOverride(path="/internal/widgets", method="get", supported_versions=()), + ), + ) +``` + +Set `openapi_version_docs` to `None` to disable version documentation. A configured non-empty global +`supported_versions` list adds an `x-accept-versioning` extension and matching description text to each operation. +`operation_versions` replaces that list for one exact path and lower-case HTTP method; an explicit empty tuple skips +microbootstrap's additions for that operation without asserting that no service-owned version metadata exists. This does +not negotiate requests, add an `Accept` parameter, change response media types, or provide a Swagger UI version selector. + +Repeated schema reads remain stable. A conflicting service-owned `x-accept-versioning` extension raises `ValueError`; +correct the configuration and rebuild the application. Litestar supports its standard `Operation` type for version +documentation and rejects other custom operation subclasses. + #### FastStream AsyncAPI documentation AsyncAPI documentation is available by default under `/asyncapi` route. You can change that by setting `asyncapi_path`: diff --git a/microbootstrap/__init__.py b/microbootstrap/__init__.py index 784b5e2..51fe81e 100644 --- a/microbootstrap/__init__.py +++ b/microbootstrap/__init__.py @@ -1,6 +1,18 @@ from microbootstrap.instruments.cors_instrument import CorsConfig from microbootstrap.instruments.health_checks_instrument import HealthChecksConfig from microbootstrap.instruments.logging_instrument import LoggingConfig +from microbootstrap.instruments.openapi_security_schemes import ( + OpenApiApiKeySecurityScheme, + OpenApiHttpSecurityScheme, + OpenApiOAuth2SecurityScheme, + OpenApiOAuthFlow, + OpenApiOAuthFlows, + OpenApiOpenIdConnectSecurityScheme, +) +from microbootstrap.instruments.openapi_version_docs import ( + OpenApiOperationVersionOverride, + OpenApiVersionDocsConfig, +) from microbootstrap.instruments.opentelemetry_instrument import ( FastStreamOpentelemetryConfig, FastStreamTelemetryMiddlewareProtocol, @@ -42,6 +54,14 @@ "LitestarPrometheusConfig", "LitestarSettings", "LoggingConfig", + "OpenApiApiKeySecurityScheme", + "OpenApiHttpSecurityScheme", + "OpenApiOAuth2SecurityScheme", + "OpenApiOAuthFlow", + "OpenApiOAuthFlows", + "OpenApiOpenIdConnectSecurityScheme", + "OpenApiOperationVersionOverride", + "OpenApiVersionDocsConfig", "OpentelemetryConfig", "PyroscopeConfig", "SentryConfig", diff --git a/microbootstrap/bootstrappers/fastapi.py b/microbootstrap/bootstrappers/fastapi.py index 628a953..07aab1b 100644 --- a/microbootstrap/bootstrappers/fastapi.py +++ b/microbootstrap/bootstrappers/fastapi.py @@ -12,6 +12,8 @@ from microbootstrap.instruments.cors_instrument import CorsInstrument from microbootstrap.instruments.health_checks_instrument import HealthChecksInstrument, HealthCheckTypedDict from microbootstrap.instruments.logging_instrument import LoggingInstrument +from microbootstrap.instruments.openapi_security_schemes import serialize_security_schemes +from microbootstrap.instruments.openapi_version_docs import SUPPORTED_HTTP_METHODS from microbootstrap.instruments.opentelemetry_instrument import OpentelemetryInstrument from microbootstrap.instruments.prometheus_instrument import FastApiPrometheusConfig, PrometheusInstrument from microbootstrap.instruments.pyroscope_instrument import PyroscopeInstrument @@ -67,8 +69,48 @@ def bootstrap_before(self) -> dict[str, typing.Any]: def bootstrap_after(self, application: ApplicationT) -> ApplicationT: if self.instrument_config.swagger_offline_docs: enable_offline_docs(application, static_files_handler=self.instrument_config.service_static_path) + if self.instrument_config.openapi_version_docs is None and not self.instrument_config.security_schemes: + return application + + original_openapi: typing.Final = application.openapi + + def documented_openapi() -> dict[str, typing.Any]: + openapi_schema: typing.Final = original_openapi() + if self.instrument_config.security_schemes: + self._merge_security_schemes(openapi_schema) + if self.instrument_config.openapi_version_docs is not None: + self._document_operations(openapi_schema) + return openapi_schema + + application.openapi = documented_openapi # type: ignore[method-assign] # FastAPI's public custom OpenAPI hook. return application + def _merge_security_schemes(self, openapi_schema: dict[str, typing.Any]) -> None: + expected_schemes: typing.Final = serialize_security_schemes(self.instrument_config.security_schemes) + components = openapi_schema.setdefault("components", {}) + security_schemes = components.setdefault("securitySchemes", {}) + self._validate_security_scheme_conflicts(security_schemes, expected_schemes) + security_schemes.update( + {name: scheme for name, scheme in expected_schemes.items() if name not in security_schemes} + ) + + def _document_operations(self, openapi_schema: dict[str, typing.Any]) -> None: + for path, path_item in openapi_schema.get("paths", {}).items(): + if path.startswith("x-"): + continue + for method, operation in path_item.items(): + if method not in SUPPORTED_HTTP_METHODS: + continue + documentation = self._build_version_documentation( + path, + method, + operation.get("description"), + operation.get("x-accept-versioning"), + has_existing_extension="x-accept-versioning" in operation, + ) + if documentation is not None: + operation["x-accept-versioning"], operation["description"] = documentation + @FastApiBootstrapper.use_instrument() class FastApiCorsInstrument(CorsInstrument): diff --git a/microbootstrap/bootstrappers/faststream.py b/microbootstrap/bootstrappers/faststream.py index ecef4c9..9b533b9 100644 --- a/microbootstrap/bootstrappers/faststream.py +++ b/microbootstrap/bootstrappers/faststream.py @@ -57,7 +57,11 @@ def _isolate_faststream_subscribers(application: AsgiFastStream) -> None: class KwargsAsgiFastStream(AsgiFastStream): def __init__(self, **kwargs: typing.Any) -> None: # noqa: ANN401 # `broker` argument is positional-only - super().__init__(kwargs.pop("broker", None), **kwargs) + broker = kwargs.pop("broker", None) + if broker is None: + super().__init__(**kwargs) + else: + super().__init__(broker, **kwargs) class FastStreamBootstrapper(ApplicationBootstrapper[FastStreamSettings, AsgiFastStream, FastStreamConfig]): diff --git a/microbootstrap/bootstrappers/litestar.py b/microbootstrap/bootstrappers/litestar.py index bf827a8..e0b8fdc 100644 --- a/microbootstrap/bootstrappers/litestar.py +++ b/microbootstrap/bootstrappers/litestar.py @@ -1,4 +1,5 @@ from __future__ import annotations +import dataclasses import typing import litestar @@ -25,6 +26,15 @@ HealthCheckTypedDict, ) from microbootstrap.instruments.logging_instrument import LoggingInstrument +from microbootstrap.instruments.openapi_security_schemes import ( + OpenApiApiKeySecurityScheme, + OpenApiHttpSecurityScheme, + OpenApiOAuth2SecurityScheme, + OpenApiOpenIdConnectSecurityScheme, + _OpenApiSecurityScheme, + serialize_security_schemes, +) +from microbootstrap.instruments.openapi_version_docs import SUPPORTED_HTTP_METHODS from microbootstrap.instruments.opentelemetry_instrument import OpentelemetryInstrument from microbootstrap.instruments.prometheus_instrument import ( LitestarPrometheusConfig, @@ -37,6 +47,17 @@ from microbootstrap.settings import LitestarSettings +ApplicationT = typing.TypeVar("ApplicationT", bound=litestar.Litestar) + + +@dataclasses.dataclass +class AcceptVersionedOperation(openapi.spec.Operation): + accept_versioning: dict[str, str | list[str]] | None = dataclasses.field( + default=None, + metadata={"alias": "x-accept-versioning"}, + ) + + if typing.TYPE_CHECKING: from litestar.contrib.opentelemetry import OpenTelemetryConfig from litestar.types import ASGIApp, Scope @@ -102,6 +123,102 @@ def bootstrap_before(self) -> dict[str, typing.Any]: ] return bootstrap_result + def bootstrap_after(self, application: ApplicationT) -> ApplicationT: + if (self.instrument_config.openapi_version_docs is None and not self.instrument_config.security_schemes) or ( + application.openapi_config is None + ): + return application + if self.instrument_config.security_schemes: + self._merge_security_schemes(application.openapi_schema) + if self.instrument_config.openapi_version_docs is not None: + self._document_operations(application.openapi_schema) + return application + + def _merge_security_schemes(self, openapi_schema: openapi.spec.OpenAPI) -> None: + expected_schemes: typing.Final = serialize_security_schemes(self.instrument_config.security_schemes) + security_schemes = openapi_schema.components.security_schemes + if security_schemes is not None: + canonical_schemes: typing.Final = { + name: scheme.to_schema() if isinstance(scheme, openapi.spec.SecurityScheme) else scheme + for name, scheme in security_schemes.items() + } + self._validate_security_scheme_conflicts(canonical_schemes, expected_schemes) + configured_schemes = { + scheme_name: self._build_litestar_security_scheme(security_scheme) + for scheme_name, security_scheme in self.instrument_config.security_schemes.items() + } + if security_schemes is None: + openapi_schema.components.security_schemes = typing.cast( + "dict[str, openapi.spec.SecurityScheme | openapi.spec.Reference]", + configured_schemes, + ) + return + security_schemes.update( + {name: scheme for name, scheme in configured_schemes.items() if name not in security_schemes} + ) + + def _document_operations(self, openapi_schema: openapi.spec.OpenAPI) -> None: + if openapi_schema.paths is None: + return + for path, path_item in openapi_schema.paths.items(): + for method in SUPPORTED_HTTP_METHODS: + operation = getattr(path_item, method) + if operation is None: + continue + existing_extension = ( + operation.accept_versioning if isinstance(operation, AcceptVersionedOperation) else None + ) + documentation = self._build_version_documentation( + path, + method, + operation.description, + existing_extension, + has_existing_extension=existing_extension is not None, + ) + if documentation is None: + continue + extension, description = documentation + if type(operation) is openapi.spec.Operation: + init_fields = { + field.name: getattr(operation, field.name) + for field in dataclasses.fields(openapi.spec.Operation) + if field.init + } + operation = AcceptVersionedOperation(**init_fields, accept_versioning=extension) + setattr(path_item, method, operation) + elif type(operation) is not AcceptVersionedOperation: + message = ( + f"OpenAPI operation {type(operation).__name__} is not supported " + "for Accept version documentation." + ) + raise TypeError(message) + operation.accept_versioning = extension + operation.description = description + + @classmethod + def _build_litestar_security_scheme(cls, security_scheme: _OpenApiSecurityScheme) -> openapi.spec.SecurityScheme: + if not isinstance( + security_scheme, + ( + OpenApiHttpSecurityScheme, + OpenApiApiKeySecurityScheme, + OpenApiOAuth2SecurityScheme, + OpenApiOpenIdConnectSecurityScheme, + ), + ): + raise AssertionError("Unsupported OpenAPI security scheme.") # noqa: TRY004 + scheme_data = security_scheme.model_dump(by_alias=False, exclude_none=True) + if "location" in scheme_data: + scheme_data["security_scheme_in"] = scheme_data.pop("location") + if "flows" in scheme_data: + flow_data = scheme_data["flows"] + if "resource_owner" in flow_data: + flow_data["password"] = flow_data.pop("resource_owner") + scheme_data["flows"] = openapi.spec.OAuthFlows( + **{flow_name: openapi.spec.OAuthFlow(**flow) for flow_name, flow in flow_data.items()} + ) + return openapi.spec.SecurityScheme(**scheme_data) + @LitestarBootstrapper.use_instrument() class LitestarCorsInstrument(CorsInstrument): diff --git a/microbootstrap/helpers.py b/microbootstrap/helpers.py index e0dffa9..1cae380 100644 --- a/microbootstrap/helpers.py +++ b/microbootstrap/helpers.py @@ -26,7 +26,7 @@ def dataclass_to_dict_no_defaults(dataclass_to_convert: "_DataclassT") -> dict[s if dataclass_field.default != value and isinstance(dataclass_field.default_factory, _MISSING_TYPE): conversion_result[dataclass_field.name] = value continue - if value != dataclass_field.default and value != dataclass_field.default_factory(): # type: ignore[misc] + if value != dataclass_field.default and value != dataclass_field.default_factory(): # type: ignore[operator] conversion_result[dataclass_field.name] = value return conversion_result diff --git a/microbootstrap/instruments/openapi_security_schemes.py b/microbootstrap/instruments/openapi_security_schemes.py new file mode 100644 index 0000000..398703e --- /dev/null +++ b/microbootstrap/instruments/openapi_security_schemes.py @@ -0,0 +1,127 @@ +from __future__ import annotations +import typing +import unicodedata + +import pydantic + + +OpenApiApiKeyLocation: typing.TypeAlias = typing.Literal["header", "query", "cookie"] + + +class OpenApiSecuritySchemeModel(pydantic.BaseModel): + model_config = pydantic.ConfigDict(extra="forbid", populate_by_name=True) + + +class OpenApiHttpSecurityScheme(OpenApiSecuritySchemeModel): + type: typing.Literal["http"] = "http" + scheme: str + bearer_format: str | None = pydantic.Field(default=None, alias="bearerFormat") + description: str | None = None + + +class OpenApiApiKeySecurityScheme(OpenApiSecuritySchemeModel): + type: typing.Literal["apiKey"] = "apiKey" + name: str + location: OpenApiApiKeyLocation = pydantic.Field(alias="in") + description: str | None = None + + +class OpenApiOAuthFlow(OpenApiSecuritySchemeModel): + authorization_url: str | None = pydantic.Field(default=None, alias="authorizationUrl") + token_url: str | None = pydantic.Field(default=None, alias="tokenUrl") + refresh_url: str | None = pydantic.Field(default=None, alias="refreshUrl") + scopes: dict[str, str] = pydantic.Field(default_factory=dict) + + @pydantic.field_validator("authorization_url", "token_url", "refresh_url") + @classmethod + def validate_url(cls, value: str | None) -> str | None: + return validate_openapi_url(value) + + +class OpenApiOAuthFlows(OpenApiSecuritySchemeModel): + implicit: OpenApiOAuthFlow | None = None + resource_owner: OpenApiOAuthFlow | None = pydantic.Field(default=None, alias="password") + client_credentials: OpenApiOAuthFlow | None = pydantic.Field(default=None, alias="clientCredentials") + authorization_code: OpenApiOAuthFlow | None = pydantic.Field(default=None, alias="authorizationCode") + + @pydantic.model_validator(mode="after") + def validate_required_urls(self) -> OpenApiOAuthFlows: + if all( + flow is None + for flow in (self.implicit, self.resource_owner, self.client_credentials, self.authorization_code) + ): + message = "OAuth2 flows must configure at least one grant type." + raise ValueError(message) + self._validate_flow_urls("implicit", self.implicit, requires_authorization_url=True) + self._validate_flow_urls("password", self.resource_owner, requires_token_url=True) + self._validate_flow_urls("client credentials", self.client_credentials, requires_token_url=True) + self._validate_flow_urls( + "authorization code", + self.authorization_code, + requires_authorization_url=True, + requires_token_url=True, + ) + return self + + @staticmethod + def _validate_flow_urls( + flow_name: str, + flow: OpenApiOAuthFlow | None, + *, + requires_authorization_url: bool = False, + requires_token_url: bool = False, + ) -> None: + if flow is None: + return + if requires_authorization_url and flow.authorization_url is None: + message = f"OAuth2 {flow_name} flow requires authorizationUrl." + raise ValueError(message) + if requires_token_url and flow.token_url is None: + message = f"OAuth2 {flow_name} flow requires tokenUrl." + raise ValueError(message) + + +class OpenApiOAuth2SecurityScheme(OpenApiSecuritySchemeModel): + type: typing.Literal["oauth2"] = "oauth2" + flows: OpenApiOAuthFlows + description: str | None = None + + +class OpenApiOpenIdConnectSecurityScheme(OpenApiSecuritySchemeModel): + type: typing.Literal["openIdConnect"] = "openIdConnect" + open_id_connect_url: str = pydantic.Field(alias="openIdConnectUrl") + description: str | None = None + + @pydantic.field_validator("open_id_connect_url") + @classmethod + def validate_url(cls, value: str) -> str: + validated_value = validate_openapi_url(value) + assert validated_value is not None # noqa: S101 - this field is required. + return validated_value + + +_OpenApiSecurityScheme: typing.TypeAlias = typing.Annotated[ + OpenApiHttpSecurityScheme + | OpenApiApiKeySecurityScheme + | OpenApiOAuth2SecurityScheme + | OpenApiOpenIdConnectSecurityScheme, + pydantic.Field(discriminator="type"), +] + + +def validate_openapi_url(value: str | None) -> str | None: + if value is None: + return value + if not value or any(character.isspace() or unicodedata.category(character) == "Cc" for character in value): + message = "OpenAPI URL values must be non-empty and contain no whitespace or control characters." + raise ValueError(message) + return value + + +def serialize_security_schemes( + security_schemes: typing.Mapping[str, _OpenApiSecurityScheme], +) -> dict[str, dict[str, typing.Any]]: + return { + scheme_name: security_scheme.model_dump(by_alias=True, exclude_none=True) + for scheme_name, security_scheme in security_schemes.items() + } diff --git a/microbootstrap/instruments/openapi_version_docs.py b/microbootstrap/instruments/openapi_version_docs.py new file mode 100644 index 0000000..fb03e5d --- /dev/null +++ b/microbootstrap/instruments/openapi_version_docs.py @@ -0,0 +1,86 @@ +from __future__ import annotations +import re +import typing + +import pydantic + + +SUPPORTED_HTTP_METHODS: typing.Final = frozenset({"delete", "get", "head", "options", "patch", "post", "put", "trace"}) +SAFE_MEDIA_TYPE_TOKEN: typing.Final[re.Pattern[str]] = re.compile(r"[!#$%&'*+\-.^_|~0-9A-Za-z]+\Z") +VENDOR_MEDIA_TYPE: typing.Final = re.compile(r"application/vnd\.([!#$%&'*+\-.^_|~0-9A-Za-z]+)\+json\Z") + + +class OpenApiOperationVersionOverride(pydantic.BaseModel): + model_config = pydantic.ConfigDict(extra="forbid") + + path: str + method: str + supported_versions: tuple[str, ...] + + @pydantic.field_validator("path") + @classmethod + def validate_path(cls, value: str) -> str: + if not value.startswith("/") or "?" in value or "#" in value: + message = "Operation path must be an absolute path without a query string or fragment." + raise ValueError(message) + return value + + @pydantic.field_validator("method") + @classmethod + def validate_method(cls, value: str) -> str: + if value not in SUPPORTED_HTTP_METHODS: + message = f"Operation method must be one of: {', '.join(sorted(SUPPORTED_HTTP_METHODS))}." + raise ValueError(message) + return value + + @pydantic.field_validator("supported_versions") + @classmethod + def validate_supported_versions(cls, value: tuple[str, ...]) -> tuple[str, ...]: + return validate_versions(value) + + +class OpenApiVersionDocsConfig(pydantic.BaseModel): + model_config = pydantic.ConfigDict(extra="forbid") + + vendor_media_type: str + supported_versions: tuple[str, ...] + operation_versions: tuple[OpenApiOperationVersionOverride, ...] = () + + @pydantic.field_validator("vendor_media_type") + @classmethod + def validate_vendor_media_type(cls, value: str) -> str: + if VENDOR_MEDIA_TYPE.fullmatch(value) is None: + message = "Vendor media type must use the application/vnd.+json form." + raise ValueError(message) + return value + + @pydantic.field_validator("supported_versions") + @classmethod + def validate_supported_versions(cls, value: tuple[str, ...]) -> tuple[str, ...]: + validated_versions = validate_versions(value) + if not validated_versions: + message = "OpenAPI version documentation requires at least one supported API version." + raise ValueError(message) + return validated_versions + + @pydantic.field_validator("operation_versions") + @classmethod + def validate_operation_versions( + cls, + value: tuple[OpenApiOperationVersionOverride, ...], + ) -> tuple[OpenApiOperationVersionOverride, ...]: + selectors = {(override.path, override.method) for override in value} + if len(value) != len(selectors): + message = "Operation version overrides must not contain duplicate path and method pairs." + raise ValueError(message) + return value + + +def validate_versions(value: tuple[str, ...]) -> tuple[str, ...]: + if len(value) != len(set(value)): + message = "Supported API versions must not contain duplicates." + raise ValueError(message) + if any(SAFE_MEDIA_TYPE_TOKEN.fullmatch(version) is None for version in value): + message = "Each supported API version must be a non-empty safe media-type token." + raise ValueError(message) + return value diff --git a/microbootstrap/instruments/swagger_instrument.py b/microbootstrap/instruments/swagger_instrument.py index aca3765..00efbda 100644 --- a/microbootstrap/instruments/swagger_instrument.py +++ b/microbootstrap/instruments/swagger_instrument.py @@ -1,10 +1,19 @@ from __future__ import annotations +import re import typing import pydantic from microbootstrap.helpers import is_valid_path from microbootstrap.instruments.base import BaseInstrumentConfig, Instrument +from microbootstrap.instruments.openapi_security_schemes import _OpenApiSecurityScheme # noqa: TC001 +from microbootstrap.instruments.openapi_version_docs import ( + SUPPORTED_HTTP_METHODS, + OpenApiVersionDocsConfig, +) + + +SECURITY_SCHEME_NAME_PATTERN: typing.Final = re.compile(r"^[a-zA-Z0-9._-]+$") class SwaggerConfig(BaseInstrumentConfig): @@ -16,6 +25,20 @@ class SwaggerConfig(BaseInstrumentConfig): swagger_path: str = "/docs" swagger_offline_docs: bool = False swagger_extra_params: dict[str, typing.Any] = pydantic.Field(default_factory=dict) + security_schemes: typing.Mapping[str, _OpenApiSecurityScheme] = pydantic.Field(default_factory=dict) + openapi_version_docs: OpenApiVersionDocsConfig | None = None + + @pydantic.field_validator("security_schemes") + @classmethod + def validate_security_scheme_names( + cls, + security_schemes: typing.Mapping[str, _OpenApiSecurityScheme], + ) -> typing.Mapping[str, _OpenApiSecurityScheme]: + for scheme_name in security_schemes: + if SECURITY_SCHEME_NAME_PATTERN.fullmatch(scheme_name) is None: + message = "OpenAPI security scheme names must match ^[a-zA-Z0-9._-]+$." + raise ValueError(message) + return security_schemes class SwaggerInstrument(Instrument[SwaggerConfig]): @@ -25,6 +48,74 @@ class SwaggerInstrument(Instrument[SwaggerConfig]): def is_ready(self) -> bool: return bool(self.instrument_config.swagger_path) and is_valid_path(self.instrument_config.swagger_path) + def _build_version_documentation( + self, + path: str, + method: str, + description: object, + existing_extension: object, + *, + has_existing_extension: bool, + ) -> tuple[dict[str, str | list[str]], str] | None: + configuration = self.instrument_config.openapi_version_docs + if configuration is None or method not in SUPPORTED_HTTP_METHODS: + return None + + supported_versions = next( + ( + override.supported_versions + for override in configuration.operation_versions + if override.path == path and override.method == method + ), + configuration.supported_versions, + ) + if not supported_versions: + return None + if description is not None and not isinstance(description, str): + message = f"OpenAPI operation {method.upper()} {path} has a non-string description." + raise ValueError(message) + + expected_extension: typing.Final[dict[str, str | list[str]]] = { + "header": "Accept", + "mediaType": configuration.vendor_media_type, + "parameter": "version", + "supportedVersions": list(supported_versions), + } + if has_existing_extension and existing_extension != expected_extension: + message = "OpenAPI operation x-accept-versioning conflicts with configured Accept version documentation." + raise ValueError(message) + return expected_extension, self._format_version_documentation(description, configuration, supported_versions) + + @staticmethod + def _format_version_documentation( + description: str | None, + configuration: OpenApiVersionDocsConfig, + supported_versions: tuple[str, ...], + ) -> str: + media_types: typing.Final = tuple( + f"{configuration.vendor_media_type}; version={version}" for version in supported_versions + ) + version_documentation: typing.Final = ( + f"Supported API version: `{media_types[0]}`." + if len(media_types) == 1 + else "Supported API versions:\n" + "\n".join(f"- `{media_type}`." for media_type in media_types) + ) + if not description: + return version_documentation + if version_documentation in description: + return description + return f"{description}\n\n{version_documentation}" + + @staticmethod + def _validate_security_scheme_conflicts( + existing_schemes: typing.Mapping[str, object], + expected_schemes: typing.Mapping[str, object], + ) -> None: + for scheme_name, expected_scheme in expected_schemes.items(): + if scheme_name in existing_schemes and existing_schemes[scheme_name] != expected_scheme: + message = f"OpenAPI security scheme '{scheme_name}' conflicts with the configured security scheme." + raise ValueError(message) + @classmethod def get_config_type(cls) -> type[SwaggerConfig]: return SwaggerConfig diff --git a/pyproject.toml b/pyproject.toml index f1621b3..2434a86 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -63,18 +63,18 @@ authors = [{ name = "community-of-python" }] [project.optional-dependencies] fastapi = [ - "fastapi>=0.100", + "fastapi>=0.110.1", "fastapi-offline-docs>=1", "opentelemetry-instrumentation-fastapi>=0.46b0", - "prometheus-fastapi-instrumentator>=6.1", + "prometheus-fastapi-instrumentator>=7.1", ] litestar = [ - "litestar>=2.9", + "litestar>=2.21.1", "litestar-offline-docs>=1", "prometheus-client>=0.20", ] granian = ["granian[reload]>=1"] -faststream = ["faststream~=0.6.2", "prometheus-client>=0.20"] +faststream = ["faststream>=0.6.7,<0.8", "prometheus-client>=0.20"] fastmcp = ["fastmcp>=2,<4", "prometheus-client>=0.20"] [dependency-groups] diff --git a/tests/bootstrappers/test_faststream.py b/tests/bootstrappers/test_faststream.py index 35f14fe..feddd24 100644 --- a/tests/bootstrappers/test_faststream.py +++ b/tests/bootstrappers/test_faststream.py @@ -193,17 +193,19 @@ async def handler(conversation_id: str) -> None: FastStreamConfig(broker=broker) ).configure_instruments(minimal_sentry_config).bootstrap() - original_log = broker.config.logger.log + logger_state: typing.Final = broker.config.logger + logger_type: typing.Final = type(logger_state) + original_log: typing.Final = logger_type.log - def record_error_tag(*args: typing.Any, **kwargs: typing.Any) -> None: # noqa: ANN401 - if kwargs.get("log_level") == logging.ERROR: + def record_error_tag(state: typing.Any, *args: typing.Any, **kwargs: typing.Any) -> None: # noqa: ANN401 + if state is logger_state and kwargs.get("log_level") == logging.ERROR: error_message = typing.cast("str", kwargs["message"]) captured_tags[error_message] = sentry_sdk.get_isolation_scope()._tags.get(conversation_id_tag) # noqa: SLF001 if error_message.endswith("second"): second_logged.set() - original_log(*args, **kwargs) + original_log(state, *args, **kwargs) - monkeypatch.setattr(broker.config.logger, "log", record_error_tag) + monkeypatch.setattr(logger_type, "log", record_error_tag) event_loop = asyncio.get_running_loop() previous_exception_handler = event_loop.get_exception_handler() event_loop.set_exception_handler(lambda *_: None) @@ -267,10 +269,12 @@ async def handler(conversation_id: str) -> None: client.get_integration.return_value = baggage_integration monkeypatch.setattr(sentry_sdk, "get_client", mock.Mock(return_value=client)) before_send: typing.Final = init.call_args.kwargs["before_send"] - original_log = broker.config.logger.log + logger_state: typing.Final = broker.config.logger + logger_type: typing.Final = type(logger_state) + original_log: typing.Final = logger_type.log - def record_automatic_error(*args: typing.Any, **kwargs: typing.Any) -> None: # noqa: ANN401 - if kwargs.get("log_level") == logging.ERROR: + def record_automatic_error(state: typing.Any, *args: typing.Any, **kwargs: typing.Any) -> None: # noqa: ANN401 + if state is logger_state and kwargs.get("log_level") == logging.ERROR: exception = typing.cast("ValueError", kwargs["exc_info"]) event = before_send( {}, @@ -279,9 +283,9 @@ def record_automatic_error(*args: typing.Any, **kwargs: typing.Any) -> None: # captured_tags[str(exception)] = event.get("tags", {}).get(conversation_id_tag) if str(exception) == "second": second_captured.set() - original_log(*args, **kwargs) + original_log(state, *args, **kwargs) - monkeypatch.setattr(broker.config.logger, "log", record_automatic_error) + monkeypatch.setattr(logger_type, "log", record_automatic_error) event_loop = asyncio.get_running_loop() previous_exception_handler = event_loop.get_exception_handler() diff --git a/tests/bootstrappers/test_openapi_version_docs.py b/tests/bootstrappers/test_openapi_version_docs.py new file mode 100644 index 0000000..cd1d2fa --- /dev/null +++ b/tests/bootstrappers/test_openapi_version_docs.py @@ -0,0 +1,671 @@ +import copy +import dataclasses +import typing + +import fastapi +import litestar +import pytest +from fastapi.security import HTTPBearer +from fastapi.testclient import TestClient as FastAPITestClient +from litestar import get, openapi, post, status_codes +from litestar.config.app import AppConfig +from litestar.openapi import spec as litestar_openapi +from litestar.openapi.plugins import SwaggerRenderPlugin +from litestar.testing import TestClient as LitestarTestClient +from pydantic import BaseModel + +from microbootstrap import ( + OpenApiApiKeySecurityScheme, + OpenApiHttpSecurityScheme, + OpenApiOAuth2SecurityScheme, + OpenApiOAuthFlow, + OpenApiOAuthFlows, + OpenApiOpenIdConnectSecurityScheme, + OpenApiOperationVersionOverride, + OpenApiVersionDocsConfig, + SwaggerConfig, +) +from microbootstrap.bootstrappers.fastapi import FastApiBootstrapper, FastApiSwaggerInstrument +from microbootstrap.bootstrappers.litestar import LitestarBootstrapper, LitestarSwaggerInstrument +from microbootstrap.config.litestar import LitestarConfig +from microbootstrap.instruments.openapi_security_schemes import _OpenApiSecurityScheme +from microbootstrap.settings import FastApiSettings, LitestarSettings + + +TARGET_PATH: typing.Final = "/widgets" +HEALTH_PATH: typing.Final = "/service-health" +NO_DESCRIPTION_PATH: typing.Final = "/without-description" +EXTENSION: typing.Final = { + "header": "Accept", + "mediaType": "application/vnd.real-api+json", + "parameter": "version", + "supportedVersions": ["2026-01", "release-candidate"], +} +VERSION_TEXT: typing.Final = ( + "Supported API versions:\n" + "- `application/vnd.real-api+json; version=2026-01`.\n" + "- `application/vnd.real-api+json; version=release-candidate`." +) +EXPECTED_GENERATOR_CALLS: typing.Final = 2 +VERSIONED_OPERATIONS: typing.Final = ( + (TARGET_PATH, "get"), + (TARGET_PATH, "post"), + (HEALTH_PATH, "get"), + (NO_DESCRIPTION_PATH, "get"), +) + + +class ServiceResponse(BaseModel): + status: str + + +@dataclasses.dataclass +class BuiltApplication: + application: fastapi.FastAPI | litestar.Litestar + renderer: SwaggerRenderPlugin | None = None + + def schema(self) -> dict[str, typing.Any]: + if isinstance(self.application, litestar.Litestar): + assert self.application.openapi_schema is not None + return self.application.openapi_schema.to_schema() + return self.application.openapi() + + def served_schema(self) -> dict[str, typing.Any]: + if isinstance(self.application, litestar.Litestar): + with LitestarTestClient(app=self.application) as client: + response = client.get("/schema/openapi.json") + else: + with FastAPITestClient(app=self.application) as client: + response = client.get("/openapi.json") + assert response.status_code == status_codes.HTTP_200_OK + return typing.cast("dict[str, typing.Any]", response.json()) + + def create_widget(self, headers: dict[str, str] | None = None) -> None: + if isinstance(self.application, litestar.Litestar): + with LitestarTestClient(app=self.application) as client: + response = client.post(TARGET_PATH, headers=headers) + expected_status = status_codes.HTTP_201_CREATED + else: + with FastAPITestClient(app=self.application) as client: + response = client.post(TARGET_PATH, headers=headers) + expected_status = status_codes.HTTP_200_OK + assert response.status_code == expected_status + assert response.headers["content-type"] == "application/json" + assert response.json() == {"status": "created"} + + +@pytest.fixture(params=("fastapi", "litestar")) +def framework(request: pytest.FixtureRequest) -> str: + return typing.cast("str", request.param) + + +def version_docs( + *, + overrides: tuple[OpenApiOperationVersionOverride, ...] = (), +) -> OpenApiVersionDocsConfig: + return OpenApiVersionDocsConfig( + vendor_media_type="application/vnd.real-api+json", + supported_versions=("2026-01", "release-candidate"), + operation_versions=overrides, + ) + + +def build_application( + framework: str, + config: OpenApiVersionDocsConfig | None, + security_schemes: dict[str, _OpenApiSecurityScheme] | None = None, +) -> BuiltApplication: + if framework == "fastapi": + fastapi_application = FastApiBootstrapper( + FastApiSettings( + service_debug=False, + security_schemes=security_schemes or {}, + openapi_version_docs=config, + ) + ).bootstrap() + service_auth = HTTPBearer() + + @fastapi_application.get(TARGET_PATH, description="List widgets", dependencies=[fastapi.Security(service_auth)]) + async def fastapi_list_widgets() -> dict[str, str]: + return {"status": "ok"} + + @fastapi_application.post(TARGET_PATH, description="Create widget") + async def fastapi_create_widget() -> dict[str, str]: + return {"status": "created"} + + @fastapi_application.get(HEALTH_PATH, description="Service health") + async def fastapi_service_health() -> dict[str, str]: + return {"status": "ok"} + + @fastapi_application.get(NO_DESCRIPTION_PATH) + async def fastapi_without_description() -> dict[str, str]: + return {"status": "ok"} + + return BuiltApplication(fastapi_application) + + @get(TARGET_PATH, description="List widgets", security=[{"ServiceAuth": []}]) + async def litestar_list_widgets() -> ServiceResponse: + return ServiceResponse(status="ok") + + @post(TARGET_PATH, description="Create widget") + async def litestar_create_widget() -> ServiceResponse: + return ServiceResponse(status="created") + + @get(HEALTH_PATH, description="Service health") + async def litestar_service_health() -> dict[str, str]: + return {"status": "ok"} + + @get(NO_DESCRIPTION_PATH) + async def litestar_without_description() -> dict[str, str]: + return {"status": "ok"} + + renderer = SwaggerRenderPlugin() + litestar_application = ( + LitestarBootstrapper( + LitestarSettings( + service_debug=False, + security_schemes=security_schemes or {}, + openapi_version_docs=config, + ) + ) + .configure_application( + LitestarConfig( + route_handlers=[ + litestar_list_widgets, + litestar_create_widget, + litestar_service_health, + litestar_without_description, + ], + openapi_config=openapi.OpenAPIConfig( + title="Service API", + version="1.0.0", + components=litestar_openapi.Components( + security_schemes={ + "GlobalAuth": litestar_openapi.SecurityScheme(type="http", scheme="bearer"), + "ServiceAuth": litestar_openapi.SecurityScheme(type="http", scheme="bearer"), + } + ), + security=[{"GlobalAuth": []}], + render_plugins=[renderer], + ), + ) + ) + .bootstrap() + ) + return BuiltApplication(litestar_application, renderer) + + +def test_version_docs_change_only_selected_operations_and_match_served_schema(framework: str) -> None: + baseline = build_application(framework, None).schema() + application = build_application( + framework, + version_docs( + overrides=(OpenApiOperationVersionOverride(path=TARGET_PATH, method="get", supported_versions=()),), + ), + ) + + schema = application.schema() + comparable_schema = copy.deepcopy(schema) + comparable_baseline = copy.deepcopy(baseline) + for path, method in ((TARGET_PATH, "post"), (HEALTH_PATH, "get"), (NO_DESCRIPTION_PATH, "get")): + comparable_schema["paths"][path][method].pop("description") + comparable_schema["paths"][path][method].pop("x-accept-versioning") + comparable_baseline["paths"][path][method].pop("description", None) + + assert application.schema() == schema + assert application.served_schema() == schema + assert comparable_schema == comparable_baseline + assert schema["components"] == baseline["components"] + assert schema.get("security") == baseline.get("security") + assert schema["paths"][TARGET_PATH]["get"] == baseline["paths"][TARGET_PATH]["get"] + assert schema["paths"][TARGET_PATH]["post"]["description"] == f"Create widget\n\n{VERSION_TEXT}" + assert schema["paths"][HEALTH_PATH]["get"]["description"] == f"Service health\n\n{VERSION_TEXT}" + assert schema["paths"][NO_DESCRIPTION_PATH]["get"]["description"] == VERSION_TEXT + assert schema["paths"][TARGET_PATH]["post"]["x-accept-versioning"] == EXTENSION + application.create_widget() + + +def test_absent_version_docs_leave_schema_and_runtime_unchanged( + framework: str, +) -> None: + baseline = build_application(framework, None) + application = build_application(framework, None) + + assert application.schema() == baseline.schema() + assert application.served_schema() == baseline.schema() + application.create_widget() + + +def test_empty_operation_override_skips_malformed_service_owned_metadata(framework: str) -> None: + configuration = version_docs( + overrides=(OpenApiOperationVersionOverride(path=TARGET_PATH, method="get", supported_versions=()),), + ) + if framework == "fastapi": + application, schema = custom_fastapi_application() + operation = schema["paths"][TARGET_PATH]["get"] + operation.update({"description": 1, "x-accept-versioning": {"header": "X-Service-Version"}}) + snapshot = copy.deepcopy(operation) + + FastApiSwaggerInstrument(SwaggerConfig(openapi_version_docs=configuration)).bootstrap_after(application) + + assert application.openapi() is schema + assert schema["paths"][TARGET_PATH]["get"] is operation + assert operation == snapshot + return + + built = build_application("litestar", None) + assert isinstance(built.application, litestar.Litestar) + assert built.application.openapi_schema is not None + assert built.application.openapi_schema.paths is not None + path_item = built.application.openapi_schema.paths[TARGET_PATH] + assert path_item.get is not None + + @dataclasses.dataclass + class ServiceOperation(litestar_openapi.Operation): + accept_versioning: dict[str, str | list[str]] | None = dataclasses.field( + default=None, + metadata={"alias": "x-accept-versioning"}, + ) + + fields = {field.name: getattr(path_item.get, field.name) for field in dataclasses.fields(path_item.get)} + operation = ServiceOperation(**fields, accept_versioning={"header": "X-Service-Version"}) + operation.description = 1 # type: ignore[assignment] # Deliberately invalid service-owned schema. + path_item.get = operation + snapshot = copy.deepcopy(operation) + + LitestarSwaggerInstrument(SwaggerConfig(openapi_version_docs=configuration)).bootstrap_after(built.application) + + assert path_item.get is operation + assert operation == snapshot + + +def test_version_docs_do_not_create_absent_security_scheme_containers(framework: str) -> None: + if framework == "fastapi": + application, schema = custom_fastapi_application() + schema.pop("components") + FastApiSwaggerInstrument(SwaggerConfig(openapi_version_docs=version_docs())).bootstrap_after(application) + + assert application.openapi() is schema + assert "components" not in schema + return + + built = build_application("litestar", None) + assert isinstance(built.application, litestar.Litestar) + assert built.application.openapi_schema is not None + built.application.openapi_schema.components.security_schemes = None + LitestarSwaggerInstrument(SwaggerConfig(openapi_version_docs=version_docs())).bootstrap_after(built.application) + + assert built.application.openapi_schema.components.security_schemes is None + + +@pytest.mark.parametrize( + ("configuration", "expected_description", "expected_extension"), + [ + ( + OpenApiVersionDocsConfig( + vendor_media_type="application/vnd.real-api+json", + supported_versions=("2027-01",), + ), + "Create widget\n\nSupported API version: `application/vnd.real-api+json; version=2027-01`.", + { + "header": "Accept", + "mediaType": "application/vnd.real-api+json", + "parameter": "version", + "supportedVersions": ["2027-01"], + }, + ), + ( + version_docs( + overrides=( + OpenApiOperationVersionOverride( + path=TARGET_PATH, + method="post", + supported_versions=("2027-01",), + ), + ) + ), + "Create widget\n\nSupported API version: `application/vnd.real-api+json; version=2027-01`.", + { + "header": "Accept", + "mediaType": "application/vnd.real-api+json", + "parameter": "version", + "supportedVersions": ["2027-01"], + }, + ), + ], +) +def test_single_global_or_operation_override_version_and_accept_header_do_not_change_runtime( + framework: str, + configuration: OpenApiVersionDocsConfig, + expected_description: str, + expected_extension: dict[str, str | list[str]], +) -> None: + baseline = build_application(framework, None).schema() + application = build_application( + framework, + configuration, + ) + + operation = application.schema()["paths"][TARGET_PATH]["post"] + assert operation["description"] == expected_description + assert operation["x-accept-versioning"] == expected_extension + assert operation.get("parameters") == baseline["paths"][TARGET_PATH]["post"].get("parameters") + assert operation.get("responses") == baseline["paths"][TARGET_PATH]["post"].get("responses") + application.create_widget() + application.create_widget({"Accept": "application/vnd.real-api+json; version=2027-01"}) + + +def configured_schemes() -> dict[str, _OpenApiSecurityScheme]: + return { + "httpAuth": OpenApiHttpSecurityScheme(scheme="bearer", bearer_format="JWT"), + "apiKeyAuth": OpenApiApiKeySecurityScheme(name="X-API-Key", location="header"), + "oauth": OpenApiOAuth2SecurityScheme( + flows=OpenApiOAuthFlows( + implicit=OpenApiOAuthFlow( + authorization_url="/implicit-authorize", + refresh_url="/implicit-refresh", + scopes={"read": "Read widgets"}, + ), + resource_owner=OpenApiOAuthFlow( + token_url="/password-token", # noqa: S106 + refresh_url="/password-refresh", + scopes={"write": "Write widgets"}, + ), + client_credentials=OpenApiOAuthFlow( + token_url="/client-token", # noqa: S106 + refresh_url="/client-refresh", + scopes={"service": "Service access"}, + ), + authorization_code=OpenApiOAuthFlow( + authorization_url="/code-authorize", + token_url="/code-token", # noqa: S106 + refresh_url="/code-refresh", + scopes={"admin": "Administer widgets"}, + ), + ) + ), + "oidc": OpenApiOpenIdConnectSecurityScheme(open_id_connect_url="/.well-known/openid-configuration"), + } + + +def test_security_schemes_preserve_service_owned_security_and_combine_with_version_docs(framework: str) -> None: + schemes = configured_schemes() + baseline = build_application(framework, None).schema() + application = build_application(framework, version_docs(), schemes) + + schema = application.schema() + comparable_schema = copy.deepcopy(schema) + for scheme_name in schemes: + comparable_schema["components"]["securitySchemes"].pop(scheme_name) + for path, method in VERSIONED_OPERATIONS: + comparable_schema["paths"][path][method].pop("description") + comparable_schema["paths"][path][method].pop("x-accept-versioning") + comparable_baseline = copy.deepcopy(baseline) + for path, method in VERSIONED_OPERATIONS: + comparable_baseline["paths"][path][method].pop("description", None) + + assert comparable_schema == comparable_baseline + assert schema["components"]["securitySchemes"] == { + **( + {"HTTPBearer": {"type": "http", "scheme": "bearer"}} + if framework == "fastapi" + else { + "GlobalAuth": {"type": "http", "scheme": "bearer"}, + "ServiceAuth": {"type": "http", "scheme": "bearer"}, + } + ), + "httpAuth": {"type": "http", "scheme": "bearer", "bearerFormat": "JWT"}, + "apiKeyAuth": {"type": "apiKey", "name": "X-API-Key", "in": "header"}, + "oauth": { + "type": "oauth2", + "flows": { + "implicit": { + "authorizationUrl": "/implicit-authorize", + "refreshUrl": "/implicit-refresh", + "scopes": {"read": "Read widgets"}, + }, + "password": { + "tokenUrl": "/password-token", + "refreshUrl": "/password-refresh", + "scopes": {"write": "Write widgets"}, + }, + "clientCredentials": { + "tokenUrl": "/client-token", + "refreshUrl": "/client-refresh", + "scopes": {"service": "Service access"}, + }, + "authorizationCode": { + "authorizationUrl": "/code-authorize", + "tokenUrl": "/code-token", + "refreshUrl": "/code-refresh", + "scopes": {"admin": "Administer widgets"}, + }, + }, + }, + "oidc": {"type": "openIdConnect", "openIdConnectUrl": "/.well-known/openid-configuration"}, + } + assert schema.get("security") == baseline.get("security") + assert schema["paths"][TARGET_PATH]["get"].get("security") == baseline["paths"][TARGET_PATH]["get"].get("security") + assert schema["paths"][TARGET_PATH]["post"]["description"] == f"Create widget\n\n{VERSION_TEXT}" + assert schema["paths"][TARGET_PATH]["post"]["x-accept-versioning"] == EXTENSION + assert application.served_schema() == schema + + +def custom_fastapi_application() -> tuple[fastapi.FastAPI, dict[str, typing.Any]]: + schema: dict[str, typing.Any] = { + "openapi": "3.1.0", + "info": {"title": "Service API", "version": "1.0.0"}, + "paths": {TARGET_PATH: {"get": {"description": "List widgets"}, "post": {"description": "Create widget"}}}, + "components": {"securitySchemes": {}}, + } + application = fastapi.FastAPI() + application.openapi = lambda: schema # type: ignore[method-assign] # Public custom OpenAPI hook. + return application, schema + + +def test_security_scheme_conflicts_leave_the_definition_batch_unchanged(framework: str) -> None: + if framework == "fastapi": + application, schema = custom_fastapi_application() + schema["components"]["securitySchemes"] = { + "matching": {"type": "http", "scheme": "bearer"}, + "conflict": {"type": "http", "scheme": "basic"}, + } + FastApiSwaggerInstrument( + SwaggerConfig(security_schemes={"matching": OpenApiHttpSecurityScheme(scheme="bearer")}) + ).bootstrap_after(application) + assert application.openapi() is schema + instrument = FastApiSwaggerInstrument( + SwaggerConfig( + security_schemes={ + "insert": OpenApiApiKeySecurityScheme(name="X-API-Key", location="header"), + "conflict": OpenApiHttpSecurityScheme(scheme="bearer"), + } + ) + ) + instrument.bootstrap_after(application) + with pytest.raises(ValueError, match="security scheme 'conflict' conflicts"): + application.openapi() + assert "insert" not in schema["components"]["securitySchemes"] + return + + litestar_application = build_application("litestar", None).application + assert isinstance(litestar_application, litestar.Litestar) + assert litestar_application.openapi_schema is not None + assert litestar_application.openapi_schema.components.security_schemes is not None + matching_scheme = litestar_openapi.SecurityScheme(type="http", scheme="bearer") + litestar_application.openapi_schema.components.security_schemes["matching"] = matching_scheme + LitestarSwaggerInstrument( + SwaggerConfig(security_schemes={"matching": OpenApiHttpSecurityScheme(scheme="bearer")}) + ).bootstrap_after(litestar_application) + assert litestar_application.openapi_schema.components.security_schemes["matching"] is matching_scheme + litestar_application.openapi_schema.components.security_schemes["conflict"] = litestar_openapi.SecurityScheme( + type="http", + scheme="basic", + ) + with pytest.raises(ValueError, match="security scheme 'conflict' conflicts"): + LitestarSwaggerInstrument( + SwaggerConfig( + security_schemes={ + "insert": OpenApiApiKeySecurityScheme(name="X-API-Key", location="header"), + "conflict": OpenApiHttpSecurityScheme(scheme="bearer"), + } + ) + ).bootstrap_after(litestar_application) + assert "insert" not in litestar_application.openapi_schema.components.security_schemes + + +def test_fastapi_preserves_custom_generator_calls_and_errors() -> None: + application, schema = custom_fastapi_application() + calls = 0 + + def generator() -> dict[str, typing.Any]: + nonlocal calls + calls += 1 + return schema + + application.openapi = generator # type: ignore[method-assign] # Public custom OpenAPI hook. + FastApiSwaggerInstrument(SwaggerConfig(openapi_version_docs=version_docs())).bootstrap_after(application) + assert application.openapi() is schema + assert application.openapi() is schema + assert calls == EXPECTED_GENERATOR_CALLS + + failing_application = fastapi.FastAPI() + failing_application.openapi = lambda: (_ for _ in ()).throw(RuntimeError("service-owned generator failed")) # type: ignore[method-assign] + FastApiSwaggerInstrument(SwaggerConfig(openapi_version_docs=version_docs())).bootstrap_after(failing_application) + with pytest.raises(RuntimeError, match="service-owned generator failed"): + failing_application.openapi() + + +def test_fastapi_documents_http_operations_without_altering_path_item_metadata() -> None: + application = fastapi.FastAPI() + schema: dict[str, typing.Any] = { + "openapi": "3.1.0", + "info": {"title": "Service API", "version": "1.0.0"}, + "paths": { + TARGET_PATH: { + "summary": "Widgets", + "parameters": [{"$ref": "#/components/parameters/RequestedBy"}], + "servers": [{"url": "https://widgets.example.test"}], + "x-service-metadata": {"owner": "widgets"}, + "get": {"description": "List widgets", "responses": {"200": {"description": "OK"}}}, + } + }, + "components": { + "parameters": {"RequestedBy": {"name": "X-Requested-By", "in": "header", "schema": {"type": "string"}}} + }, + } + application.openapi = lambda: schema # type: ignore[method-assign] # Public custom OpenAPI hook. + FastApiSwaggerInstrument(SwaggerConfig(openapi_version_docs=version_docs())).bootstrap_after(application) + + canonical_schema = application.openapi() + repeated_schema = application.openapi() + with FastAPITestClient(app=application) as client: + served_schema = client.get("/openapi.json") + + assert served_schema.status_code == status_codes.HTTP_200_OK + assert canonical_schema is schema + assert repeated_schema is schema + assert served_schema.json() == schema + assert schema["paths"][TARGET_PATH] == { + "summary": "Widgets", + "parameters": [{"$ref": "#/components/parameters/RequestedBy"}], + "servers": [{"url": "https://widgets.example.test"}], + "x-service-metadata": {"owner": "widgets"}, + "get": { + "description": f"List widgets\n\n{VERSION_TEXT}", + "responses": {"200": {"description": "OK"}}, + "x-accept-versioning": EXTENSION, + }, + } + + +def test_litestar_served_schema_matches_canonical_schema_on_repeated_reads() -> None: + built = build_application("litestar", version_docs()) + assert isinstance(built.application, litestar.Litestar) + assert built.application.openapi_config is not None + assert built.application.openapi_config.render_plugins[0] is built.renderer + + expected_schema = built.schema() + assert built.served_schema() == expected_schema + assert built.served_schema() == expected_schema + + +@pytest.mark.parametrize( + ("has_security_schemes", "has_version_docs"), + [(False, False), (True, False), (False, True), (True, True)], +) +def test_litestar_bootstraps_without_openapi_config( + has_security_schemes: bool, + has_version_docs: bool, +) -> None: + @get(TARGET_PATH) + async def handler() -> dict[str, str]: + return {"status": "ok"} + + def disable_openapi(configuration: AppConfig) -> AppConfig: + configuration.openapi_config = None + return configuration + + application = ( + LitestarBootstrapper( + LitestarSettings( + service_debug=False, + security_schemes={"ServiceAuth": OpenApiHttpSecurityScheme(scheme="bearer")} + if has_security_schemes + else {}, + openapi_version_docs=version_docs() if has_version_docs else None, + ) + ) + .configure_application(LitestarConfig(route_handlers=[handler], on_app_init=[disable_openapi])) + .bootstrap() + ) + + assert application.openapi_config is None + with LitestarTestClient(app=application) as client: + response = client.get(TARGET_PATH) + openapi_response = client.get("/docs/openapi.json") + assert response.status_code == status_codes.HTTP_200_OK + assert response.json() == {"status": "ok"} + assert openapi_response.status_code == status_codes.HTTP_404_NOT_FOUND + + +def test_litestar_standard_operations_preserve_fields_and_are_stable() -> None: + built = build_application("litestar", None) + assert isinstance(built.application, litestar.Litestar) + assert built.application.openapi_schema is not None + assert built.application.openapi_schema.paths is not None + path_item = built.application.openapi_schema.paths[TARGET_PATH] + standard_post = path_item.post + assert standard_post is not None + standard_fields = { + field.name: getattr(standard_post, field.name) + for field in dataclasses.fields(litestar_openapi.Operation) + if field.name != "description" + } + instrument = LitestarSwaggerInstrument(SwaggerConfig(openapi_version_docs=version_docs())) + instrument.bootstrap_after(built.application) + + assert all(getattr(path_item.post, name) == value for name, value in standard_fields.items()) + assert built.schema()["paths"][TARGET_PATH]["post"]["x-accept-versioning"] == EXTENSION + assert built.served_schema()["paths"][TARGET_PATH]["post"]["x-accept-versioning"] == EXTENSION + documented_post = path_item.post + instrument.bootstrap_after(built.application) + assert path_item.post is documented_post + + +def test_litestar_rejects_unsupported_custom_operation() -> None: + built = build_application("litestar", None) + assert isinstance(built.application, litestar.Litestar) + assert built.application.openapi_schema is not None + assert built.application.openapi_schema.paths is not None + path_item = built.application.openapi_schema.paths[TARGET_PATH] + assert path_item.post is not None + + @dataclasses.dataclass + class UnsupportedOperation(litestar_openapi.Operation): + metadata: dict[str, str] | None = dataclasses.field(default=None, metadata={"alias": "x-service-metadata"}) + + fields = {field.name: getattr(path_item.post, field.name) for field in dataclasses.fields(path_item.post)} + path_item.post = UnsupportedOperation(**fields, metadata={"owner": "widgets"}) + with pytest.raises(TypeError, match="is not supported for Accept version documentation"): + LitestarSwaggerInstrument(SwaggerConfig(openapi_version_docs=version_docs())).bootstrap_after(built.application) diff --git a/tests/instruments/test_openapi_security_schemes.py b/tests/instruments/test_openapi_security_schemes.py new file mode 100644 index 0000000..46d557c --- /dev/null +++ b/tests/instruments/test_openapi_security_schemes.py @@ -0,0 +1,253 @@ +import typing + +import pytest +from pydantic import Field, ValidationError + +from microbootstrap import OpenApiApiKeySecurityScheme as ApiKeySecurityScheme +from microbootstrap import ( + OpenApiHttpSecurityScheme, + OpenApiOAuth2SecurityScheme, + OpenApiOAuthFlow, + OpenApiOAuthFlows, + OpenApiOpenIdConnectSecurityScheme, +) +from microbootstrap.instruments.instrument_box import InstrumentBox +from microbootstrap.instruments.openapi_security_schemes import _OpenApiSecurityScheme, serialize_security_schemes +from microbootstrap.instruments.swagger_instrument import SwaggerConfig, SwaggerInstrument +from microbootstrap.settings import LitestarSettings + + +class HttpOnlySettings(LitestarSettings): + security_schemes: dict[str, OpenApiHttpSecurityScheme] = Field(default_factory=dict) + + +class ApiKeyOnlySettings(LitestarSettings): + security_schemes: dict[str, ApiKeySecurityScheme] = Field(default_factory=dict) + + +class MixedSecuritySettings(LitestarSettings): + security_schemes: dict[str, OpenApiHttpSecurityScheme | ApiKeySecurityScheme] = Field(default_factory=dict) + + +def test_security_schemes_accept_python_names_and_serialize_openapi_aliases() -> None: + security_schemes: typing.Final[dict[str, _OpenApiSecurityScheme]] = { + "http.auth": OpenApiHttpSecurityScheme(scheme="bearer", bearer_format="JWT"), + "api-key": ApiKeySecurityScheme(name="X-API-Key", location="header"), + "oauth2": OpenApiOAuth2SecurityScheme( + flows=OpenApiOAuthFlows( + implicit=OpenApiOAuthFlow(authorization_url="/authorize", scopes={"read": "Read access"}), + resource_owner=OpenApiOAuthFlow(token_url="/token", scopes={"write": "Write access"}), # noqa: S106 + client_credentials=OpenApiOAuthFlow(token_url="/client-token"), # noqa: S106 + authorization_code=OpenApiOAuthFlow( + authorization_url="/code-authorize", + token_url="/code-token", # noqa: S106 + refresh_url="/refresh", + ), + ) + ), + "oidc": OpenApiOpenIdConnectSecurityScheme(open_id_connect_url="/.well-known/openid-configuration"), + } + + assert serialize_security_schemes(security_schemes) == { + "http.auth": {"type": "http", "scheme": "bearer", "bearerFormat": "JWT"}, + "api-key": {"type": "apiKey", "name": "X-API-Key", "in": "header"}, + "oauth2": { + "type": "oauth2", + "flows": { + "implicit": {"authorizationUrl": "/authorize", "scopes": {"read": "Read access"}}, + "password": {"tokenUrl": "/token", "scopes": {"write": "Write access"}}, + "clientCredentials": {"tokenUrl": "/client-token", "scopes": {}}, + "authorizationCode": { + "authorizationUrl": "/code-authorize", + "tokenUrl": "/code-token", + "refreshUrl": "/refresh", + "scopes": {}, + }, + }, + }, + "oidc": {"type": "openIdConnect", "openIdConnectUrl": "/.well-known/openid-configuration"}, + } + + +def test_swagger_config_round_trips_all_security_scheme_types() -> None: + configuration: typing.Final = SwaggerConfig( + security_schemes={ + "http": OpenApiHttpSecurityScheme(scheme="bearer", bearer_format="JWT"), + "api": ApiKeySecurityScheme(name="X-API-Key", location="header"), + "oauth": OpenApiOAuth2SecurityScheme( + flows=OpenApiOAuthFlows(client_credentials=OpenApiOAuthFlow(token_url="/token")) # noqa: S106 + ), + "oidc": OpenApiOpenIdConnectSecurityScheme(open_id_connect_url="/.well-known/openid-configuration"), + } + ) + + from_dump: typing.Final = SwaggerConfig.model_validate(configuration.model_dump(by_alias=True)) + from_json: typing.Final = SwaggerConfig.model_validate_json(configuration.model_dump_json(by_alias=True)) + + for round_tripped in (from_dump, from_json): + assert type(round_tripped.security_schemes) is dict + assert round_tripped.security_schemes == configuration.security_schemes + assert isinstance(round_tripped.security_schemes["http"], OpenApiHttpSecurityScheme) + assert isinstance(round_tripped.security_schemes["api"], ApiKeySecurityScheme) + assert isinstance(round_tripped.security_schemes["oauth"], OpenApiOAuth2SecurityScheme) + assert isinstance(round_tripped.security_schemes["oidc"], OpenApiOpenIdConnectSecurityScheme) + + +def test_instrument_box_reconstructs_concrete_security_scheme_settings() -> None: + configurations: list[SwaggerConfig] = [] + for settings in ( + HttpOnlySettings(security_schemes={"http": OpenApiHttpSecurityScheme(scheme="bearer", bearer_format="JWT")}), + ApiKeyOnlySettings(security_schemes={"api": ApiKeySecurityScheme(name="X-API-Key", location="header")}), + MixedSecuritySettings( + security_schemes={ + "http": OpenApiHttpSecurityScheme(scheme="bearer"), + "api": ApiKeySecurityScheme(name="X-API-Key", location="query"), + } + ), + ): + instrument_box = InstrumentBox(__instruments__=[SwaggerInstrument]) + instrument_box.initialize(settings) + + configuration = instrument_box.instruments[0].instrument_config + + assert isinstance(configuration, SwaggerConfig) + assert type(configuration.security_schemes) is dict + configurations.append(configuration) + + http_configuration, api_key_configuration, mixed_configuration = configurations + assert isinstance(http_configuration.security_schemes["http"], OpenApiHttpSecurityScheme) + assert isinstance(api_key_configuration.security_schemes["api"], ApiKeySecurityScheme) + assert isinstance(mixed_configuration.security_schemes["http"], OpenApiHttpSecurityScheme) + assert isinstance(mixed_configuration.security_schemes["api"], ApiKeySecurityScheme) + + +def test_instrument_box_merges_security_scheme_dicts() -> None: + instrument_box: typing.Final = InstrumentBox(__instruments__=[SwaggerInstrument]) + instrument_box.initialize(HttpOnlySettings(security_schemes={"http": OpenApiHttpSecurityScheme(scheme="bearer")})) + instrument_box.configure_instrument( + SwaggerConfig(security_schemes={"api": ApiKeySecurityScheme(name="X-API-Key", location="header")}) + ) + + configuration = instrument_box.instruments[0].instrument_config + + assert isinstance(configuration, SwaggerConfig) + assert type(configuration.security_schemes) is dict + assert set(configuration.security_schemes) == {"http", "api"} + assert isinstance(configuration.security_schemes["http"], OpenApiHttpSecurityScheme) + assert isinstance(configuration.security_schemes["api"], ApiKeySecurityScheme) + + +def test_http_only_settings_reject_other_security_scheme_kinds() -> None: + with pytest.raises(ValidationError): + HttpOnlySettings(security_schemes={"api": ApiKeySecurityScheme(name="X-API-Key", location="header")}) + + +def test_security_schemes_accept_openapi_aliases() -> None: + configuration: typing.Final = SwaggerConfig( + security_schemes={ + "api": {"type": "apiKey", "name": "X-API-Key", "in": "query"}, + "oauth": { + "type": "oauth2", + "flows": {"clientCredentials": {"tokenUrl": "relative-token", "scopes": {}}}, + }, + "oidc": {"type": "openIdConnect", "openIdConnectUrl": "relative-discovery"}, + } + ) + + assert serialize_security_schemes(configuration.security_schemes) == { + "api": {"type": "apiKey", "name": "X-API-Key", "in": "query"}, + "oauth": { + "type": "oauth2", + "flows": {"clientCredentials": {"tokenUrl": "relative-token", "scopes": {}}}, + }, + "oidc": {"type": "openIdConnect", "openIdConnectUrl": "relative-discovery"}, + } + + +def test_oauth_flows_accept_python_and_openapi_password_names() -> None: + flow: typing.Final = OpenApiOAuthFlow(token_url="/token") # noqa: S106 + flows: typing.Final = OpenApiOAuthFlows(resource_owner=flow) + flow_arguments: typing.Final = {"password": flow} + alias_flows: typing.Final = OpenApiOAuthFlows(**flow_arguments) + + assert flows.resource_owner == alias_flows.resource_owner == flow + + +@pytest.mark.parametrize( + ("flows", "error"), + [ + ({"implicit": {"scopes": {}}}, "implicit flow requires authorizationUrl"), + ({"password": {"scopes": {}}}, "password flow requires tokenUrl"), + ({"clientCredentials": {"scopes": {}}}, "client credentials flow requires tokenUrl"), + ( + {"authorizationCode": {"tokenUrl": "/token", "scopes": {}}}, + "authorization code flow requires authorizationUrl", + ), + ( + {"authorizationCode": {"authorizationUrl": "/authorize", "scopes": {}}}, + "authorization code flow requires tokenUrl", + ), + ], +) +def test_oauth_flows_require_urls_for_configured_grant_types(flows: dict[str, object], error: str) -> None: + with pytest.raises(ValidationError, match=error): + OpenApiOAuthFlows.model_validate(flows) + + +def test_oauth_flows_require_at_least_one_grant_type() -> None: + with pytest.raises(ValidationError, match="must configure at least one grant type"): + OpenApiOAuthFlows() + + +@pytest.mark.parametrize("url", ["", " ", "/oauth token", "/oauth\ttoken", "/oauth\ntoken", "/oauth\x00token"]) +def test_oauth_and_openid_urls_reject_empty_whitespace_and_control_characters(url: str) -> None: + with pytest.raises(ValidationError, match="must be non-empty"): + OpenApiOAuthFlow(token_url=url) + with pytest.raises(ValidationError, match="must be non-empty"): + OpenApiOpenIdConnectSecurityScheme(open_id_connect_url=url) + + +def test_oauth_and_openid_urls_allow_relative_references() -> None: + flow: typing.Final = OpenApiOAuthFlow( + authorization_url="authorize", + token_url="../token", # noqa: S106 - a relative URL is under test. + refresh_url="./refresh", + ) + oidc: typing.Final = OpenApiOpenIdConnectSecurityScheme(open_id_connect_url=".well-known/openid-configuration") + + assert flow.authorization_url == "authorize" + assert flow.token_url == "../token" # noqa: S105 - a relative URL is under test. + assert flow.refresh_url == "./refresh" + assert oidc.open_id_connect_url == ".well-known/openid-configuration" + + +@pytest.mark.parametrize( + "configuration", + [ + lambda: OpenApiHttpSecurityScheme.model_validate({"scheme": "bearer", "unexpected": "value"}), + lambda: ApiKeySecurityScheme.model_validate({"name": "X-API-Key", "location": "header", "unexpected": "value"}), + lambda: OpenApiOAuthFlow.model_validate({"tokenUrl": "/token", "unexpected": "value"}), + lambda: OpenApiOAuthFlows.model_validate({"clientCredentials": {"tokenUrl": "/token"}, "unexpected": "value"}), + lambda: OpenApiOAuth2SecurityScheme.model_validate({"flows": {}, "unexpected": "value"}), + lambda: OpenApiOpenIdConnectSecurityScheme.model_validate( + {"openIdConnectUrl": "/discovery", "unexpected": "value"} + ), + ], +) +def test_security_scheme_models_forbid_unknown_fields(configuration: typing.Callable[[], object]) -> None: + with pytest.raises(ValidationError, match="Extra inputs are not permitted"): + configuration() + + +@pytest.mark.parametrize("scheme_name", ["", "service auth", "service/auth", "служебная"]) +def test_security_scheme_component_names_must_match_openapi_pattern(scheme_name: str) -> None: + with pytest.raises(ValidationError, match=r"\^\[a-zA-Z0-9\._-\]\+\$"): + SwaggerConfig(security_schemes={scheme_name: OpenApiHttpSecurityScheme(scheme="bearer")}) + + +def test_security_scheme_component_names_allow_openapi_non_identifiers() -> None: + configuration: typing.Final = SwaggerConfig( + security_schemes={"service.auth-1": OpenApiHttpSecurityScheme(scheme="bearer")} + ) + + assert set(configuration.security_schemes) == {"service.auth-1"} diff --git a/tests/instruments/test_swagger.py b/tests/instruments/test_swagger.py index 619d8f5..a76e6da 100644 --- a/tests/instruments/test_swagger.py +++ b/tests/instruments/test_swagger.py @@ -1,17 +1,34 @@ +import copy import typing import fastapi import litestar +import pytest from fastapi.testclient import TestClient as FastAPITestClient from litestar import openapi, status_codes from litestar.openapi import spec as litestar_openapi from litestar.openapi.plugins import ScalarRenderPlugin from litestar.static_files import StaticFilesConfig from litestar.testing import TestClient as LitestarTestClient +from pydantic import ValidationError from microbootstrap.bootstrappers.fastapi import FastApiSwaggerInstrument from microbootstrap.bootstrappers.litestar import LitestarSwaggerInstrument +from microbootstrap.instruments.openapi_security_schemes import ( + OpenApiApiKeySecurityScheme, + OpenApiHttpSecurityScheme, + OpenApiOAuth2SecurityScheme, + OpenApiOAuthFlow, + OpenApiOAuthFlows, + OpenApiSecuritySchemeModel, + _OpenApiSecurityScheme, +) +from microbootstrap.instruments.openapi_version_docs import OpenApiOperationVersionOverride, OpenApiVersionDocsConfig from microbootstrap.instruments.swagger_instrument import SwaggerConfig, SwaggerInstrument +from microbootstrap.settings import FastApiSettings, LitestarSettings + + +TARGET_PATH: typing.Final = "/widgets" def test_swagger_is_ready(minimal_swagger_config: SwaggerConfig) -> None: @@ -37,7 +54,7 @@ def test_swagger_teardown( minimal_swagger_config: SwaggerConfig, ) -> None: swagger_instrument: typing.Final = SwaggerInstrument(minimal_swagger_config) - assert swagger_instrument.teardown() is None # type: ignore[func-returns-value] + swagger_instrument.teardown() def test_litestar_swagger_bootstrap_online_docs(minimal_swagger_config: SwaggerConfig) -> None: @@ -80,6 +97,84 @@ def test_litestar_swagger_bootstrap_extra_params_have_correct_types(minimal_swag assert type(bootstrap_result["openapi_config"].components) is litestar_openapi.Components +def test_litestar_swagger_builds_native_security_schemes_with_optional_fields() -> None: + swagger_instrument: typing.Final = LitestarSwaggerInstrument( + SwaggerConfig( + security_schemes={ + "http": OpenApiHttpSecurityScheme(scheme="bearer"), + "api-key": OpenApiApiKeySecurityScheme(name="X-API-Key", location="header"), + "oauth": OpenApiOAuth2SecurityScheme( + flows=OpenApiOAuthFlows(resource_owner=OpenApiOAuthFlow(token_url="/token")) # noqa: S106 + ), + } + ) + ) + application: typing.Final = litestar.Litestar(**swagger_instrument.bootstrap_before()) + + swagger_instrument.bootstrap_after(application) + + security_schemes = application.openapi_schema.components.security_schemes + assert security_schemes is not None + http_scheme, api_key_scheme, oauth_scheme = ( + security_schemes["http"], + security_schemes["api-key"], + security_schemes["oauth"], + ) + assert isinstance(http_scheme, litestar_openapi.SecurityScheme) + assert isinstance(api_key_scheme, litestar_openapi.SecurityScheme) + assert isinstance(oauth_scheme, litestar_openapi.SecurityScheme) + assert http_scheme.bearer_format is None + assert http_scheme.description is None + assert api_key_scheme.security_scheme_in == "header" + assert api_key_scheme.description is None + assert oauth_scheme.flows is not None + assert oauth_scheme.flows.password is not None + assert oauth_scheme.flows.password.token_url == "/token" # noqa: S105 + + +def test_litestar_version_docs_preserves_components_only_schema() -> None: + swagger_instrument: typing.Final = LitestarSwaggerInstrument( + SwaggerConfig( + openapi_version_docs=OpenApiVersionDocsConfig( + vendor_media_type="application/vnd.example+json", + supported_versions=("2026-01",), + ) + ) + ) + application: typing.Final = litestar.Litestar(**swagger_instrument.bootstrap_before()) + schema: typing.Final = application.openapi_schema + schema.paths = None + expected_schema: typing.Final = copy.deepcopy(schema.to_schema()) + + swagger_instrument.bootstrap_after(application) + + assert application.openapi_schema is schema + assert schema.paths is None + assert schema.to_schema() == expected_schema + + +def test_litestar_rejects_unsupported_security_scheme_model() -> None: + unsupported_scheme: typing.Final = typing.cast("_OpenApiSecurityScheme", OpenApiSecuritySchemeModel()) + + with pytest.raises(AssertionError, match=r"^Unsupported OpenAPI security scheme\.$"): + LitestarSwaggerInstrument._build_litestar_security_scheme(unsupported_scheme) # noqa: SLF001 + + +def test_swagger_version_documentation_returns_none_when_disabled() -> None: + swagger_instrument: typing.Final = SwaggerInstrument(SwaggerConfig()) + + assert ( + swagger_instrument._build_version_documentation( # noqa: SLF001 + TARGET_PATH, + "get", + "Service-owned description", + {"service": "owned"}, + has_existing_extension=True, + ) + is None + ) + + def test_litestar_swagger_bootstrap_offline_docs(minimal_swagger_config: SwaggerConfig) -> None: minimal_swagger_config.swagger_offline_docs = True swagger_instrument: typing.Final = LitestarSwaggerInstrument(minimal_swagger_config) @@ -130,6 +225,8 @@ def test_litestar_swagger_bootstrap_working_offline_docs( assert response.status_code == status_codes.HTTP_200_OK response = test_client.get(f"{minimal_swagger_config.service_static_path}/swagger-ui.css") assert response.status_code == status_codes.HTTP_200_OK + response = test_client.get(f"{minimal_swagger_config.service_static_path}/swagger-ui-bundle.js") + assert response.status_code == status_codes.HTTP_200_OK def test_fastapi_swagger_bootstrap_online_docs(minimal_swagger_config: SwaggerConfig) -> None: @@ -172,3 +269,332 @@ def test_fastapi_swagger_bootstrap_working_offline_docs( assert response.status_code == status_codes.HTTP_200_OK response = test_client.get(f"{minimal_swagger_config.service_static_path}/swagger-ui.css") assert response.status_code == status_codes.HTTP_200_OK + response = test_client.get(f"{minimal_swagger_config.service_static_path}/swagger-ui-bundle.js") + assert response.status_code == status_codes.HTTP_200_OK + + +@pytest.mark.parametrize( + ("operation", "error"), + [ + ({"description": "List widgets"}, None), + ( + {"description": "List widgets", "x-accept-versioning": None}, + "x-accept-versioning conflicts", + ), + ({"description": 42}, "non-string description"), + ], +) +def test_fastapi_version_docs_distinguishes_missing_and_invalid_operation_values( + operation: dict[str, object], error: str | None +) -> None: + configuration: typing.Final = SwaggerConfig( + openapi_version_docs=OpenApiVersionDocsConfig( + vendor_media_type="application/vnd.example+json", + supported_versions=("2026-01",), + ) + ) + swagger_instrument: typing.Final = FastApiSwaggerInstrument(configuration) + application: typing.Final = fastapi.FastAPI() + schema: dict[str, typing.Any] = {"paths": {TARGET_PATH: {"get": operation}}} + application.openapi = lambda: schema # type: ignore[method-assign] # Exercise FastAPI's public OpenAPI hook. + + swagger_instrument.bootstrap_after(application) + + if error is not None: + with pytest.raises(ValueError, match=error): + application.openapi() + else: + documented_schema = application.openapi() + assert documented_schema["paths"][TARGET_PATH]["get"]["x-accept-versioning"] == { + "header": "Accept", + "mediaType": "application/vnd.example+json", + "parameter": "version", + "supportedVersions": ["2026-01"], + } + + +def test_fastapi_version_docs_empty_operation_override_leaves_operation_unchanged() -> None: + swagger_instrument: typing.Final = FastApiSwaggerInstrument( + SwaggerConfig( + openapi_version_docs=OpenApiVersionDocsConfig( + vendor_media_type="application/vnd.example+json", + supported_versions=("2026-01",), + operation_versions=( + OpenApiOperationVersionOverride(path=TARGET_PATH, method="get", supported_versions=()), + ), + ) + ) + ) + application: typing.Final = fastapi.FastAPI() + operation: dict[str, typing.Any] = {"description": "List widgets"} + application.openapi = lambda: {"paths": {TARGET_PATH: {"get": operation}}} # type: ignore[method-assign] # Exercise FastAPI's public OpenAPI hook. + + swagger_instrument.bootstrap_after(application) + + assert application.openapi()["paths"][TARGET_PATH]["get"] == {"description": "List widgets"} + + +@pytest.mark.parametrize( + ("schema", "expected_schema"), + [ + ( + { + "openapi": "3.1.0", + "info": {"title": "Service API", "version": "1.0.0"}, + "components": {"schemas": {"Widget": {"type": "object"}}}, + }, + { + "openapi": "3.1.0", + "info": {"title": "Service API", "version": "1.0.0"}, + "components": {"schemas": {"Widget": {"type": "object"}}}, + }, + ), + ( + { + "openapi": "3.1.0", + "info": {"title": "Service API", "version": "1.0.0"}, + "paths": {"x-owner": "widgets", TARGET_PATH: {"get": {"description": "List widgets"}}}, + }, + { + "openapi": "3.1.0", + "info": {"title": "Service API", "version": "1.0.0"}, + "paths": { + "x-owner": "widgets", + TARGET_PATH: { + "get": { + "description": "List widgets\n\nSupported API version: " + "`application/vnd.example+json; version=2026-01`.", + "x-accept-versioning": { + "header": "Accept", + "mediaType": "application/vnd.example+json", + "parameter": "version", + "supportedVersions": ["2026-01"], + }, + } + }, + }, + }, + ), + ], +) +def test_fastapi_version_docs_handles_optional_paths_and_extensions( + schema: dict[str, typing.Any], expected_schema: dict[str, typing.Any] +) -> None: + swagger_instrument: typing.Final = FastApiSwaggerInstrument( + SwaggerConfig( + openapi_version_docs=OpenApiVersionDocsConfig( + vendor_media_type="application/vnd.example+json", + supported_versions=("2026-01",), + ) + ) + ) + application: typing.Final = fastapi.FastAPI() + application.openapi = lambda: schema # type: ignore[method-assign] # Exercise FastAPI's public OpenAPI hook. + + swagger_instrument.bootstrap_after(application) + + assert application.openapi() == expected_schema + assert application.openapi() == expected_schema + + +@pytest.mark.parametrize( + ("configuration", "error"), + [ + ( + lambda: OpenApiVersionDocsConfig.model_validate({"vendor_media_type": "application/vnd.example+json"}), + "Field required", + ), + ( + lambda: OpenApiVersionDocsConfig.model_validate({"supported_versions": ("2026-01",)}), + "Field required", + ), + ( + lambda: OpenApiVersionDocsConfig( + vendor_media_type=None, + supported_versions=("2026-01",), + ), + "Input should be a valid string", + ), + ( + lambda: OpenApiVersionDocsConfig( + vendor_media_type="application/vnd.example+json", + supported_versions=None, + ), + "Input should be a valid tuple", + ), + ( + lambda: OpenApiVersionDocsConfig( + vendor_media_type="application/json", + supported_versions=("1.0",), + ), + "Vendor media type must use", + ), + ( + lambda: OpenApiVersionDocsConfig( + vendor_media_type="application/vnd.example+json", + supported_versions=("1.0", "1.0"), + ), + "must not contain duplicates", + ), + ( + lambda: OpenApiOperationVersionOverride(path="widgets", method="get", supported_versions=()), + "must be an absolute path", + ), + ( + lambda: OpenApiOperationVersionOverride(path="/widgets", method="GET", supported_versions=()), + "Operation method must be one of", + ), + ], +) +def test_openapi_version_docs_configuration_rejects_invalid_values( + configuration: typing.Callable[[], object], + error: str, +) -> None: + with pytest.raises(ValidationError, match=error): + configuration() + + +@pytest.mark.parametrize("enabled", [False, True]) +def test_openapi_version_docs_rejects_removed_enabled_field(enabled: bool) -> None: + with pytest.raises(ValidationError, match="Extra inputs are not permitted"): + OpenApiVersionDocsConfig.model_validate( + {"enabled": enabled, "vendor_media_type": "application/vnd.example+json", "supported_versions": ("1.0",)} + ) + + +def test_openapi_version_docs_rejects_removed_suppression_and_unknown_override_fields() -> None: + with pytest.raises(ValidationError, match="Extra inputs are not permitted"): + OpenApiVersionDocsConfig.model_validate( + { + "vendor_media_type": "application/vnd.example+json", + "supported_versions": ("1.0",), + "suppressed_operations": (), + } + ) + with pytest.raises(ValidationError, match="Extra inputs are not permitted"): + OpenApiOperationVersionOverride.model_validate( + {"path": "/widgets", "method": "get", "supported_versions": (), "enabled": False} + ) + + +@pytest.mark.parametrize( + "vendor_media_type", + [ + "", + "version1.0", + "application/json", + "application/vnd.+json", + "application/vnd.example api+json", + "application/vnd.bad name+json", + "application/vnd.example\tapi+json", + "application/vnd.example\napi+json", + "application/vnd.example\x00api+json", + "application/vnd.example`api+json", + "application/vnd.example,api+json", + "application/vnd.example;api+json", + ], +) +def test_openapi_version_docs_rejects_unsafe_vendor_media_types(vendor_media_type: str) -> None: + with pytest.raises(ValidationError, match="Vendor media type must use"): + OpenApiVersionDocsConfig( + vendor_media_type=vendor_media_type, + supported_versions=("release-2026",), + ) + + +@pytest.mark.parametrize( + "version", + [ + "", + "1.0 ", + "1.0\tnext", + "1.0\nnext", + "1.0\x00next", + "1.0`next", + "1.0,next", + "1.0;next", + "version1.0, application/json", + ], +) +def test_openapi_version_docs_rejects_unsafe_versions(version: str) -> None: + with pytest.raises(ValidationError, match="safe media-type token"): + OpenApiVersionDocsConfig( + vendor_media_type="application/vnd.real-api+json", + supported_versions=(version,), + ) + + +def test_openapi_version_docs_accepts_real_vendor_and_non_semver_versions() -> None: + configuration: typing.Final = OpenApiVersionDocsConfig( + vendor_media_type="application/vnd.real-api_v2+json", + supported_versions=("2026-01", "release-candidate", "v1.0+beta"), + ) + + assert configuration.vendor_media_type == "application/vnd.real-api_v2+json" + assert configuration.supported_versions == ("2026-01", "release-candidate", "v1.0+beta") + + +def test_openapi_version_docs_accepts_explicit_empty_operation_override_versions() -> None: + configuration: typing.Final = OpenApiOperationVersionOverride( + path="/widgets", + method="get", + supported_versions=(), + ) + + assert configuration.supported_versions == () + + +@pytest.mark.parametrize("settings_type", [FastApiSettings, LitestarSettings]) +def test_settings_validate_security_schemes_and_operation_versions( + settings_type: type[FastApiSettings] | type[LitestarSettings], +) -> None: + settings = settings_type( + security_schemes={"serviceBearer": {"type": "http", "scheme": "bearer"}}, + openapi_version_docs={ + "vendor_media_type": "application/vnd.real-api+json", + "supported_versions": ("2026-01",), + "operation_versions": ({"path": TARGET_PATH, "method": "post", "supported_versions": ("2027-01",)},), + }, + ) + assert settings.security_schemes == {"serviceBearer": OpenApiHttpSecurityScheme(scheme="bearer")} + assert settings.openapi_version_docs is not None + assert settings.openapi_version_docs.operation_versions[0].supported_versions == ("2027-01",) + + +@pytest.mark.parametrize( + ("configuration", "message"), + [ + ( + lambda: OpenApiVersionDocsConfig( + vendor_media_type="application/vnd.real-api+json", + supported_versions=(), + ), + "requires at least one supported API version", + ), + ( + lambda: OpenApiVersionDocsConfig( + vendor_media_type="application/vnd.real-api+json", + supported_versions=("2026-01",), + operation_versions=( + OpenApiOperationVersionOverride( + path=TARGET_PATH, + method="post", + supported_versions=("2027-01",), + ), + OpenApiOperationVersionOverride( + path=TARGET_PATH, + method="post", + supported_versions=("2028-01",), + ), + ), + ), + "duplicate path and method pairs", + ), + ], +) +def test_version_docs_settings_reject_empty_global_versions_and_duplicate_override_selectors( + configuration: typing.Callable[[], object], + message: str, +) -> None: + with pytest.raises(ValueError, match=message): + configuration()