diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index 83fdce8..cff6eaf 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -34,6 +34,8 @@ jobs: run: python -m pip install ".[dev]" - name: Run Python tests run: python -m pytest + - name: Type-check public contract sample + run: python -m mypy --strict examples/typed_consumer.py linux-distributions: name: Validate (${{ matrix.name }}) diff --git a/CHANGELOG.md b/CHANGELOG.md index 8672954..e84fb84 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -20,6 +20,9 @@ and versions are tracked in the repo-root `VERSION` file. Click parameters with domain-specific secret names. - Add `ConfigurationError` so consumer profiles can explicitly mark user-correctable configuration messages as safe usage errors. +- Add public typed runtime/profile/context contracts, generic attachment + factories, isolated command-schema registries/codecs, and a strict consumer + typing example. ### Changed @@ -32,6 +35,8 @@ and versions are tracked in the repo-root `VERSION` file. - Route `base_cli.testing.invoke()` through the production `run_app()` boundary and add a keyword-only `reraise_unexpected` opt-in for tests that need the original exception. +- Reject native async callbacks explicitly and preserve Click command subtypes + through typed `attach()` decorators and adapters. ### Fixed diff --git a/README.md b/README.md index 360ff3a..b818307 100644 --- a/README.md +++ b/README.md @@ -76,6 +76,41 @@ consumer-owned adapters should supply any product-specific policies. See [`docs/consumer-profiles.md`](docs/consumer-profiles.md) for the boundary and migration guidance. +### Typed extension contracts + +The public profile contract includes `ProjectDiscovery`, `ConfigLoader`, +`RuntimeResolver`, `HistoryWriter`, and the other resolver protocols exported +from `base_cli`. `RuntimeBinding.layout` uses the public immutable +`RuntimeLayout` type; consumers do not need to import private runtime modules. + +`Context` is generic over the validated configuration, application state, and +service payloads owned by a consumer: + +```python +Config = dict[str, object] +context: base_cli.Context[Config, ApplicationState, Services] +``` + +`App.command()`, `App.subcommand()`, `@base_cli.command()`, `@base_cli.option()`, +and `@base_cli.argument()` preserve the decorated callable's `ParamSpec` +signature. `base_cli.attach()` and `App.attach()` preserve the concrete Click +command subtype in their return type. + +`AttachmentAdapter`, `AttachmentContract`, and the typed context/service +factories define the Click attachment boundary for adapters that compose or +wrap an attached command. + +Command protocol schemas can be isolated per consumer with +`CommandSchemaRegistry` and `CommandCodec`. The module-level registration and +codec helpers remain compatible defaults backed by `RECORD_SCHEMAS`, but new +integrations should prefer an instance-owned registry when multiple protocol +boundaries share a process. + +Native async callbacks are intentionally rejected with an actionable error. +The core lifecycle is synchronous so cleanup, Click resource unwinding, and +outcome finalization remain deterministic; an adapter may provide an explicit +async runner without changing the core contract. + ## Public API The supported facade is `import base_cli`. It exports the command lifecycle @@ -101,6 +136,10 @@ Low-level implementation helpers are intentionally not included in the module `__all__` surfaces. Downstream code should use the documented facade or the explicitly supported symbols from those modules. +The repository includes [`examples/typed_consumer.py`](examples/typed_consumer.py), +a strict-typechecked consumer showing the public profile, runtime, and generic +context contracts. CI runs `mypy --strict` against that sample. + ## Minimal Command ```python diff --git a/docs/consumer-profiles.md b/docs/consumer-profiles.md index 8936b10..68802a0 100644 --- a/docs/consumer-profiles.md +++ b/docs/consumer-profiles.md @@ -123,3 +123,30 @@ whose compatibility names still reflect one historical consumer. The package rename is deliberately separate from this refactor. Names can be changed after the dependency boundary is stable. + +## Typed extension contracts + +The supported callback contracts are exported from `base_cli` as typed protocols +for static analyzers: `ProjectDiscovery`, `UserConfigLoader`, +`ConfigLoader`, `RuntimeResolver`, `WorkspaceRootResolver`, `HistoryWriter`, +`DisplayCommandResolver`, and `HistoryDisplayResolver`. A custom runtime +resolver returns `RuntimeBinding`, whose immutable `layout` is the public +`RuntimeLayout` dataclass. No consumer needs to import `_runtime`. + +`Context` accepts three consumer payload types: + +```python +Context[ConfigT, ApplicationStateT, ServicesT] +``` + +`config` is the validated configuration payload; `application_context` and +`services` are optional state and service payloads initialized by an attached +consumer. `AttachmentAdapter` and `AttachmentContract` describe the typed +boundary used by `App.attach()`. Attachment returns the same concrete Click +command object, so aliases, lazy groups, and custom Click subclasses remain +owned by the consumer. + +The core lifecycle is synchronous by design. Native `async def` callbacks and +callbacks that return awaitables are rejected with an actionable error. An +adapter that owns an event loop may run asynchronous work explicitly at its +boundary and return a normal synchronous callback result to base-cli. diff --git a/examples/typed_consumer.py b/examples/typed_consumer.py new file mode 100644 index 0000000..f6e6c2d --- /dev/null +++ b/examples/typed_consumer.py @@ -0,0 +1,77 @@ +"""Small strict-typing example for the public base-cli extension contracts.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import base_cli + + +@dataclass(frozen=True) +class ApplicationState: + invocation_count: int = 0 + + +@dataclass(frozen=True) +class Services: + workspace: Path + + +Config = dict[str, Any] +TypedContext = base_cli.Context[Config, ApplicationState, Services] + + +def profile() -> base_cli.CliProfile: + """Return a profile whose extension points use only public contracts.""" + + def resolve_runtime( + cli_name: str, + project: base_cli.ProjectInfo | None, + ) -> base_cli.RuntimeBinding: + root = Path(".base-cli-cache").resolve() + run_id = "sample-run" + layout = base_cli.RuntimeLayout( + owner_root=root / cli_name, + run_root=root / cli_name / "runs" / run_id, + state_dir=root / cli_name, + log_dir=root / cli_name / "runs" / run_id / "logs", + cache_dir=root / cli_name / "cache", + temp_dir=root / cli_name / "runs" / run_id / "tmp", + ) + return base_cli.RuntimeBinding( + cache_root=root, + layout=layout, + application_home=None, + runtime_owner="typed-sample", + project_root=project.root if project is not None else None, + project_name=project.name if project is not None else None, + inherited_path=None, + history_parent_run_id=None, + run_id=run_id, + ) + + return base_cli.CliProfile.generic(resolve_runtime=resolve_runtime) + + +app = base_cli.App( + name="typed-sample", + profile=profile(), + log_to_file=False, +) + + +@app.command() +@base_cli.option("--verbose", is_flag=True) +def main(ctx: TypedContext, verbose: bool) -> None: + """Use a consumer-owned context payload without private imports.""" + + del verbose + assert isinstance(ctx.config, dict) + _ = ctx.application_context + _ = ctx.services + + +if __name__ == "__main__": + raise SystemExit(base_cli.run_app(app)) diff --git a/lib/python/base_cli/__init__.py b/lib/python/base_cli/__init__.py index 6e6be43..f1d5fd7 100644 --- a/lib/python/base_cli/__init__.py +++ b/lib/python/base_cli/__init__.py @@ -31,6 +31,12 @@ def _resolve_version() -> str: __version__ = _resolve_version() from . import command_filters, command_protocol, history, testing +from .attachment import ( + AttachmentAdapter, + AttachmentContextFactory, + AttachmentContract, + AttachmentServiceFactory, +) from .app import ( App, argument, @@ -44,16 +50,26 @@ def _resolve_version() -> str: from .command_filters import CommandFilterNormalizer, command_matches, normalize_command_filter, normalize_command_filters from .command_protocol import ( BOOLEAN, + DEFAULT_SCHEMA_REGISTRY, + CommandCodec, NULLABLE_STRING, STRING, CommandProtocolError, + CommandSchemaRegistry, FieldSpec, + RECORD_SCHEMAS, dumps_record, dumps_records, loads_records, register_record_schema, ) -from .context import Context, get_current_context +from .context import ( + ApplicationStateT, + ConfigT, + Context, + ServicesT, + get_current_context, +) from .errors import ConfigurationError from .exit_codes import ExitCode from .inspection import inspection_envelope, render_inspection_json @@ -74,17 +90,41 @@ def _resolve_version() -> str: render_records, resolve_output_format, ) -from .profile import CliProfile, ProjectInfo, RuntimeBinding +from .profile import ( + CliProfile, + ConfigLoader, + DisplayCommandResolver, + HistoryDisplayResolver, + HistoryWriter, + ProjectDiscovery, + ProjectInfo, + RuntimeBinding, + RuntimeResolver, + UserConfigLoader, + WorkspaceRootResolver, +) +from .runtime import RuntimeLayout __all__ = [ "App", "__version__", + "AttachmentAdapter", + "AttachmentContextFactory", + "AttachmentContract", + "AttachmentServiceFactory", "BOOLEAN", + "ApplicationStateT", "CliProfile", + "ConfigLoader", "CommandFilterNormalizer", + "CommandCodec", "CommandProtocolError", + "CommandSchemaRegistry", "ConfigurationError", "Context", + "ConfigT", + "DEFAULT_SCHEMA_REGISTRY", + "DisplayCommandResolver", "ExitCode", "FieldSpec", "LIFECYCLE_META_KEY", @@ -121,6 +161,10 @@ def _resolve_version() -> str: "OutputFormatError", "PUBLIC_OUTPUT_FORMATS", "ProjectInfo", + "ProjectDiscovery", + "RECORD_SCHEMAS", + "RuntimeLayout", + "RuntimeResolver", "is_terminal", "output_format_choices", "option", @@ -130,4 +174,9 @@ def _resolve_version() -> str: "resolve_output_format", "run_app", "RuntimeBinding", + "ServicesT", + "HistoryWriter", + "HistoryDisplayResolver", + "UserConfigLoader", + "WorkspaceRootResolver", ] diff --git a/lib/python/base_cli/_lifecycle.py b/lib/python/base_cli/_lifecycle.py index e6a6ffe..2abbcbc 100644 --- a/lib/python/base_cli/_lifecycle.py +++ b/lib/python/base_cli/_lifecycle.py @@ -24,7 +24,7 @@ class InvocationOutcome: class RunRecorder: """Write core-owned lifecycle snapshots for one Context.""" - context: Context + context: Context[Any, Any, Any] started_at: datetime started_monotonic_ns: int diff --git a/lib/python/base_cli/_runtime.py b/lib/python/base_cli/_runtime.py index 3fe1537..d117543 100644 --- a/lib/python/base_cli/_runtime.py +++ b/lib/python/base_cli/_runtime.py @@ -4,21 +4,11 @@ import logging import os import stat -from dataclasses import dataclass from pathlib import Path from ._private_files import PRIVATE_DIRECTORY_MODE, restrict_directory, write_private_json from .paths import runtime_run_directory_name, runtime_slug - - -@dataclass(frozen=True) -class RuntimeLayout: - owner_root: Path - run_root: Path - state_dir: Path - log_dir: Path - cache_dir: Path - temp_dir: Path +from .runtime import RuntimeLayout _LOG_INDEX_NAME = ".base-cli-log-index.json" diff --git a/lib/python/base_cli/app.py b/lib/python/base_cli/app.py index 74eeec6..d990ea8 100644 --- a/lib/python/base_cli/app.py +++ b/lib/python/base_cli/app.py @@ -1,6 +1,7 @@ from __future__ import annotations import functools +import inspect import logging import os import stat @@ -13,7 +14,7 @@ from datetime import datetime from pathlib import Path from threading import RLock -from typing import Any, Callable +from typing import Any, Callable, ParamSpec, TypeVar from ._lifecycle import ( InvocationOutcome, @@ -29,6 +30,7 @@ create_runtime_directory, prune_log_files, ) +from .attachment import AttachmentContract from .context import Context, recover_current_context, reset_current_context, set_current_context from .errors import ConfigurationError from .exit_codes import ExitCode @@ -92,11 +94,20 @@ _CLICK_ORIGINAL_MAIN_ATTRIBUTE = "__base_cli_original_main__" _CLICK_APP_OWNER_ATTRIBUTE = "__base_cli_app_owner__" _CLICK_LIFECYCLE_BINDINGS_ATTRIBUTE = "__base_cli_lifecycle_bindings__" +_CLICK_INSTRUMENTED_SENTINEL = object() +_CLICK_MAIN_INSTRUMENTED_SENTINEL = object() _CLICK_ATTACHMENT_LOCK = RLock() _REGISTRATION_OPEN = "open" _REGISTRATION_MATERIALIZING = "materializing" _REGISTRATION_FROZEN = "frozen" _COMMAND_NAME_SUFFIXES = frozenset({"command", "cmd", "group", "grp"}) +_P = ParamSpec("_P") +_R = TypeVar("_R") +_ClickCommandT = TypeVar("_ClickCommandT") +_ASYNC_CALLBACK_ERROR = ( + "Native async Click callbacks are not supported by base-cli. " + "Use a synchronous callback or an adapter with an explicit async runner." +) @dataclass @@ -139,15 +150,7 @@ class _LifecycleResolution: raw: dict[str, _RawLifecycleValue] -@dataclass(frozen=True) -class _ClickAttachment: - app: Any - command: Any - context_factory: Callable[[Context], Any] | None - service_factory: Callable[[Context], Any] | None - sensitive_parameters: frozenset[str] - lifecycle_options: LifecycleOptions - standard_bindings: dict[str, _LifecycleBinding] +_ClickAttachment = AttachmentContract class _AttachedInvocation: @@ -155,9 +158,9 @@ class _AttachedInvocation: def __init__( self, - attachment: _ClickAttachment, + attachment: _ClickAttachment[Any], root_click_context: Any, - context: Context, + context: Context[Any, Any, Any], recorder: RunRecorder, ) -> None: self.attachment = attachment @@ -269,7 +272,7 @@ def _default_log_file(layout: Any, configured_log_file: Path | None) -> Path: return configured_log_file or layout.log_dir / "primary.log" -def _warn_lifecycle_failure(context: Context, message: str, exc: BaseException) -> None: +def _warn_lifecycle_failure(context: Context[Any, Any, Any], message: str, exc: BaseException) -> None: """Report a secondary lifecycle failure without breaking teardown.""" try: detail = str(exc) or type(exc).__name__ @@ -278,7 +281,7 @@ def _warn_lifecycle_failure(context: Context, message: str, exc: BaseException) pass -def _capture_invocation_context(context: Context, owner_app: App) -> None: +def _capture_invocation_context(context: Context[Any, Any, Any], owner_app: App) -> None: state = _INVOCATION_STATE.get() if state is None or state.owner_app is not owner_app: return @@ -310,7 +313,7 @@ def _capture_effective_output_options( state.quiet = quiet -def _record_unexpected_traceback(context: Context, outcome: InvocationOutcome) -> None: +def _record_unexpected_traceback(context: Context[Any, Any, Any], outcome: InvocationOutcome) -> None: if outcome.kind != "unexpected_error": return try: @@ -360,7 +363,7 @@ def _discard_owned_run_record(recorder: RunRecorder) -> None: ) -def _reset_active_context(context: Context, token: Any) -> None: +def _reset_active_context(context: Context[Any, Any, Any], token: Any) -> None: try: reset_current_context(token) except BaseException as exc: # pylint: disable=broad-exception-caught @@ -540,12 +543,17 @@ def _validate_single_command_name( f"the registered command cannot use '{explicit_name}'." ) - def command(self, *command_args: Any, **command_kwargs: Any): + def command( + self, + *command_args: Any, + **command_kwargs: Any, + ) -> Callable[[Callable[_P, _R]], Callable[_P, _R]]: with self._registration_lock: self._ensure_registration_open() self._validate_single_command_name(command_args, command_kwargs) - def decorator(func: Callable[..., Any]): + def decorator(func: Callable[_P, _R]) -> Callable[_P, _R]: + _reject_async_callback(func) with self._registration_lock: self._ensure_registration_open() self._validate_single_command_name(command_args, command_kwargs) @@ -566,12 +574,17 @@ def decorator(func: Callable[..., Any]): return decorator - def subcommand(self, *command_args: Any, **command_kwargs: Any): + def subcommand( + self, + *command_args: Any, + **command_kwargs: Any, + ) -> Callable[[Callable[_P, _R]], Callable[_P, _R]]: with self._registration_lock: self._ensure_registration_open() _explicit_command_name(command_args, command_kwargs) - def decorator(func: Callable[..., Any]): + def decorator(func: Callable[_P, _R]) -> Callable[_P, _R]: + _reject_async_callback(func) with self._registration_lock: self._ensure_registration_open() if self._command_func is not None: @@ -599,12 +612,12 @@ def decorator(func: Callable[..., Any]): def attach( self, - command: Any, + command: _ClickCommandT, *, - context_factory: Callable[[Context], Any] | None = None, - service_factory: Callable[[Context], Any] | None = None, + context_factory: Callable[[Context[Any, Any, Any]], Any] | None = None, + service_factory: Callable[[Context[Any, Any, Any]], Any] | None = None, sensitive_parameters: Iterable[str] = (), - ) -> Any: + ) -> _ClickCommandT: """Attach this app's lifecycle to an existing Click command tree. The same command object is returned rather than copied. Click continues @@ -615,6 +628,7 @@ def attach( click = _require_click() if not isinstance(command, click.Command): raise TypeError("App.attach() requires a click.Command instance.") + _reject_async_callback(getattr(command, "callback", None)) if context_factory is not None and not callable(context_factory): raise TypeError("context_factory must be callable or None.") if service_factory is not None and not callable(service_factory): @@ -668,18 +682,46 @@ def attach( ) added_parameters: list[Any] = [] - command_was_instrumented = bool( - getattr(command, _CLICK_INSTRUMENTED_ATTRIBUTE, False) - ) - main_was_instrumented = bool( - getattr(command, _CLICK_MAIN_INSTRUMENTED_ATTRIBUTE, False) - ) missing_marker = object() previous_marker = getattr( command, _CLICK_ATTACHMENT_ATTRIBUTE, missing_marker, ) + if previous_marker is not missing_marker and not isinstance( + previous_marker, + _ClickAttachment, + ): + raise RuntimeError( + f"Click command '{command_name}' uses base-cli's reserved " + "attachment marker. Remove that attribute before attaching." + ) + for marker_name, sentinel, description in ( + ( + _CLICK_INSTRUMENTED_ATTRIBUTE, + _CLICK_INSTRUMENTED_SENTINEL, + "command instrumentation", + ), + ( + _CLICK_MAIN_INSTRUMENTED_ATTRIBUTE, + _CLICK_MAIN_INSTRUMENTED_SENTINEL, + "main instrumentation", + ), + ): + marker = getattr(command, marker_name, missing_marker) + if marker is not missing_marker and marker is not sentinel: + raise RuntimeError( + f"Click command '{command_name}' uses base-cli's reserved " + f"{description} marker. Remove that attribute before attaching." + ) + command_was_instrumented = ( + getattr(command, _CLICK_INSTRUMENTED_ATTRIBUTE, None) + is _CLICK_INSTRUMENTED_SENTINEL + ) + main_was_instrumented = ( + getattr(command, _CLICK_MAIN_INSTRUMENTED_ATTRIBUTE, None) + is _CLICK_MAIN_INSTRUMENTED_SENTINEL + ) previous_redaction_plan = self._redaction_plan previous_attached_command = self._attached_command previous_click_command = self._click_command @@ -905,7 +947,7 @@ def wrapper(**kwargs: Any): _capture_standard_options(standard, self) started_at = utc_now() started_monotonic_ns = time.monotonic_ns() - context: Context | None = None + context: Context[Any, Any, Any] | None = None recorder: RunRecorder | None = None outcome = outcome_from_exit_code(ExitCode.SUCCESS) invocation_argv: list[str] = [] @@ -934,7 +976,7 @@ def wrapper(**kwargs: Any): context.log.debug("project_root=%s", context.project_root) if context.manifest_path is not None: context.log.debug("manifest_path=%s", context.manifest_path) - result = func(context, **kwargs) + result = _reject_async_result(func(context, **kwargs)) try: exit_code = _normalize_command_result(result) except TypeError as exc: @@ -1002,7 +1044,11 @@ def wrapper(**kwargs: Any): click_parameters[-1]._base_cli_sensitive = True return wrapper - def _create_context(self, standard: dict[str, Any], dry_run: bool = False) -> Context: + def _create_context( + self, + standard: dict[str, Any], + dry_run: bool = False, + ) -> Context[dict[str, Any], Any, Any]: project = self.profile.discover_project(current_working_dir()) manifest_path = project.manifest if project is not None else None explicit_config = Path(standard["config"]).expanduser() if standard.get("config") else None @@ -1125,7 +1171,7 @@ def _create_context(self, standard: dict[str, Any], dry_run: bool = False) -> Co def _rollback_context_creation( - context: Context, + context: Context[Any, Any, Any], *, logger_activation_started: bool, ) -> None: @@ -1152,7 +1198,7 @@ class _AttachedLifecycleResource: def __init__( self, click: Any, - attachment: _ClickAttachment, + attachment: _ClickAttachment[Any], click_context: Any, lifecycle_values: LifecycleValues, ) -> None: @@ -1163,7 +1209,7 @@ def __init__( self.standard = _standard_options_from_values(lifecycle_values) self.started_at = utc_now() self.started_monotonic_ns = time.monotonic_ns() - self.context: Context | None = None + self.context: Context[Any, Any, Any] | None = None self.invocation: _AttachedInvocation | None = None self.context_token: Any = None self.invocation_token: Any = None @@ -2336,8 +2382,14 @@ def _with_attached_lifecycle_resource( def _instrument_attached_click_command(click: Any, command: Any) -> None: with _CLICK_ATTACHMENT_LOCK: - if getattr(command, _CLICK_INSTRUMENTED_ATTRIBUTE, False): + marker = getattr(command, _CLICK_INSTRUMENTED_ATTRIBUTE, None) + if marker is _CLICK_INSTRUMENTED_SENTINEL: return + if marker is not None: + raise RuntimeError( + "Click command uses base-cli's reserved command instrumentation marker." + ) + _reject_async_callback(getattr(command, "callback", None)) original_invoke = command.invoke original_resolve = getattr(command, "resolve_command", None) @@ -2380,7 +2432,7 @@ def invoke(click_context: Any) -> Any: if resource.invocation is not None: resource.invocation.start(click_context) try: - result = original_invoke(click_context) + result = _reject_async_result(original_invoke(click_context)) except BaseException as exc: resource.record_exception(exc) raise @@ -2391,7 +2443,7 @@ def invoke(click_context: Any) -> Any: active.note_child_context(click_context) if not _click_command_has_pending_children(click_context, command): active.start(click_context) - return original_invoke(click_context) + return _reject_async_result(original_invoke(click_context)) try: setattr(command, _CLICK_ORIGINAL_INVOKE_ATTRIBUTE, original_invoke) @@ -2440,7 +2492,7 @@ def resolve_command(click_context: Any, args: list[str]) -> Any: command.resolve_command = resolve_command - setattr(command, _CLICK_INSTRUMENTED_ATTRIBUTE, True) + setattr(command, _CLICK_INSTRUMENTED_ATTRIBUTE, _CLICK_INSTRUMENTED_SENTINEL) except BaseException: _restore_attached_click_command(command) raise @@ -2460,7 +2512,6 @@ def _restore_attached_click_command(command: Any) -> None: except (AttributeError, TypeError): pass for attribute in ( - _CLICK_INSTRUMENTED_ATTRIBUTE, _CLICK_ORIGINAL_INVOKE_ATTRIBUTE, _CLICK_ORIGINAL_RESOLVE_ATTRIBUTE, ): @@ -2468,11 +2519,21 @@ def _restore_attached_click_command(command: Any) -> None: delattr(command, attribute) except (AttributeError, TypeError): pass + if getattr(command, _CLICK_INSTRUMENTED_ATTRIBUTE, None) is _CLICK_INSTRUMENTED_SENTINEL: + try: + delattr(command, _CLICK_INSTRUMENTED_ATTRIBUTE) + except (AttributeError, TypeError): + pass def _instrument_attached_click_main(command: Any) -> None: - if getattr(command, _CLICK_MAIN_INSTRUMENTED_ATTRIBUTE, False): + marker = getattr(command, _CLICK_MAIN_INSTRUMENTED_ATTRIBUTE, None) + if marker is _CLICK_MAIN_INSTRUMENTED_SENTINEL: return + if marker is not None: + raise RuntimeError( + "Click command uses base-cli's reserved main instrumentation marker." + ) original_main = command.main @functools.wraps(original_main) @@ -2506,7 +2567,7 @@ def main(*args: Any, **kwargs: Any) -> Any: try: setattr(command, _CLICK_ORIGINAL_MAIN_ATTRIBUTE, original_main) command.main = main - setattr(command, _CLICK_MAIN_INSTRUMENTED_ATTRIBUTE, True) + setattr(command, _CLICK_MAIN_INSTRUMENTED_ATTRIBUTE, _CLICK_MAIN_INSTRUMENTED_SENTINEL) except BaseException: _restore_attached_click_main(command) raise @@ -2519,10 +2580,12 @@ def _restore_attached_click_main(command: Any) -> None: command.main = original_main except (AttributeError, TypeError): pass - for attribute in ( - _CLICK_MAIN_INSTRUMENTED_ATTRIBUTE, - _CLICK_ORIGINAL_MAIN_ATTRIBUTE, - ): + if getattr(command, _CLICK_MAIN_INSTRUMENTED_ATTRIBUTE, None) is _CLICK_MAIN_INSTRUMENTED_SENTINEL: + try: + delattr(command, _CLICK_MAIN_INSTRUMENTED_ATTRIBUTE) + except (AttributeError, TypeError): + pass + for attribute in (_CLICK_ORIGINAL_MAIN_ATTRIBUTE,): try: delattr(command, attribute) except (AttributeError, TypeError): @@ -2560,14 +2623,14 @@ def get_command_app(command_func: Callable[..., Any]) -> App: def attach( - command: Any, + command: _ClickCommandT, *, app: App | None = None, - context_factory: Callable[[Context], Any] | None = None, - service_factory: Callable[[Context], Any] | None = None, + context_factory: Callable[[Context[Any, Any, Any]], Any] | None = None, + service_factory: Callable[[Context[Any, Any, Any]], Any] | None = None, sensitive_parameters: Iterable[str] = (), **app_kwargs: Any, -) -> Any: +) -> _ClickCommandT: """Attach lifecycle middleware and return the same Click command object. Attachment ownership, factories, and sensitivity policy are immutable; @@ -2741,6 +2804,20 @@ def _normalize_command_result(result: Any) -> int: ) +def _reject_async_callback(callback: Any) -> None: + if callback is not None and inspect.iscoroutinefunction(callback): + raise RuntimeError(_ASYNC_CALLBACK_ERROR) + + +def _reject_async_result(result: Any) -> Any: + if inspect.isawaitable(result): + close = getattr(result, "close", None) + if callable(close): + close() + raise RuntimeError(_ASYNC_CALLBACK_ERROR) + return result + + def _lifecycle_flag_declarations( option: LifecycleOption | None, ) -> tuple[tuple[str, ...], tuple[str, ...]]: @@ -2826,10 +2903,14 @@ def delegated_display_command(default: str | None = None) -> str | None: return default -def command(*args: Any, **kwargs: Any): +def command( + *args: Any, + **kwargs: Any, +) -> Callable[[Callable[_P, _R]], Callable[_P, _R]]: explicit_name = _explicit_command_name(args, kwargs) - def decorator(func: Callable[..., Any]): + def decorator(func: Callable[_P, _R]) -> Callable[_P, _R]: + _reject_async_callback(func) with _COMMAND_APP_LOCK: if getattr(func, _COMMAND_APP_ATTRIBUTE, None) is not None: raise RuntimeError( @@ -2849,8 +2930,13 @@ def decorator(func: Callable[..., Any]): return decorator -def option(*param_decls: str, sensitive: bool = False, dry_run: bool = False, **attrs: Any): - def decorator(func: Callable[..., Any]): +def option( + *param_decls: str, + sensitive: bool = False, + dry_run: bool = False, + **attrs: Any, +) -> Callable[[Callable[_P, _R]], Callable[_P, _R]]: + def decorator(func: Callable[_P, _R]) -> Callable[_P, _R]: specs = list(getattr(func, "__base_cli_param_specs__", [])) specs.append(("option", param_decls, attrs, sensitive)) func.__base_cli_param_specs__ = specs @@ -2868,8 +2954,12 @@ def decorator(func: Callable[..., Any]): return decorator -def argument(*param_decls: str, sensitive: bool = False, **attrs: Any): - def decorator(func: Callable[..., Any]): +def argument( + *param_decls: str, + sensitive: bool = False, + **attrs: Any, +) -> Callable[[Callable[_P, _R]], Callable[_P, _R]]: + def decorator(func: Callable[_P, _R]) -> Callable[_P, _R]: specs = list(getattr(func, "__base_cli_param_specs__", [])) specs.append(("argument", param_decls, attrs, sensitive)) func.__base_cli_param_specs__ = specs diff --git a/lib/python/base_cli/attachment.py b/lib/python/base_cli/attachment.py new file mode 100644 index 0000000..215d364 --- /dev/null +++ b/lib/python/base_cli/attachment.py @@ -0,0 +1,69 @@ +"""Typed contracts for attaching base-cli lifecycle behavior to Click trees. + +The Click adapter implementation remains private to :mod:`base_cli.app`, but +its boundary is deliberately small and typed. Consumers can use the adapter +protocol when wrapping or composing an attached command without importing the +framework's private Click machinery. +""" + +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass +from typing import Any, Generic, Protocol, TypeVar + +from .context import Context + + +__all__ = [ + "AttachmentAdapter", + "AttachmentContextFactory", + "AttachmentContract", + "AttachmentServiceFactory", +] + + +CommandT = TypeVar("CommandT") + + +class AttachmentContextFactory(Protocol): + """Create consumer-owned application state for an active Context.""" + + def __call__(self, context: Context[Any, Any, Any]) -> Any: ... + + +class AttachmentServiceFactory(Protocol): + """Create consumer-owned services for an active Context.""" + + def __call__(self, context: Context[Any, Any, Any]) -> Any: ... + + +class AttachmentAdapter(Protocol[CommandT]): + """Lifecycle adapter contract implemented by :class:`base_cli.App`. + + ``attach`` must return the exact command object it receives. This preserves + Click's concrete command subtype, aliases, lazy resolution, and callback + identity for consumers that compose multiple adapters. + """ + + def attach( + self, + command: CommandT, + *, + context_factory: AttachmentContextFactory | None = None, + service_factory: AttachmentServiceFactory | None = None, + sensitive_parameters: set[str] | frozenset[str] = frozenset(), + ) -> CommandT: ... + + +@dataclass(frozen=True) +class AttachmentContract(Generic[CommandT]): + """Immutable attachment state shared by the private Click adapter.""" + + app: Any + command: CommandT + context_factory: Callable[[Context[Any, Any, Any]], Any] | None + service_factory: Callable[[Context[Any, Any, Any]], Any] | None + sensitive_parameters: frozenset[str] + lifecycle_options: Any + standard_bindings: dict[str, Any] diff --git a/lib/python/base_cli/command_protocol.py b/lib/python/base_cli/command_protocol.py index 24376ef..ffca1de 100644 --- a/lib/python/base_cli/command_protocol.py +++ b/lib/python/base_cli/command_protocol.py @@ -7,9 +7,13 @@ __all__ = [ "BOOLEAN", + "CommandCodec", "CommandProtocolError", + "CommandSchemaRegistry", + "DEFAULT_SCHEMA_REGISTRY", "FieldSpec", "NULLABLE_STRING", + "RECORD_SCHEMAS", "STRING", "dumps_record", "dumps_records", @@ -36,13 +40,86 @@ class FieldSpec: NULLABLE_STRING = FieldSpec("string", nullable=True) BOOLEAN = FieldSpec("boolean") -# Consumers register their record schemas at their integration boundary. -RECORD_SCHEMAS: dict[str, dict[str, FieldSpec]] = {} - RecordValue = str | bool | None Record = Mapping[str, RecordValue] +class CommandSchemaRegistry: + """Own an isolated set of command record schemas. + + The module-level helpers below continue to use the default registry for + compatibility, while consumers that host more than one protocol boundary + can construct independent registries and codecs. + """ + + def __init__(self) -> None: + self.schemas: dict[str, dict[str, FieldSpec]] = {} + + def register(self, record_type: str, fields: Mapping[str, FieldSpec]) -> None: + _validate_and_store_schema(self.schemas, record_type, fields) + + def schema(self, record_type: str) -> dict[str, FieldSpec]: + return _lookup_schema(self.schemas, record_type) + + +class CommandCodec: + """Encode and decode records using one isolated schema registry.""" + + def __init__(self, registry: CommandSchemaRegistry | None = None) -> None: + self.registry = registry or CommandSchemaRegistry() + + def register_schema(self, record_type: str, fields: Mapping[str, FieldSpec]) -> None: + self.registry.register(record_type, fields) + + def dumps_record( + self, + record_type: str, + record: Record, + *, + protocol_header: str = PROTOCOL_HEADER, + ) -> str: + return dumps_record( + record_type, + record, + protocol_header=protocol_header, + registry=self.registry, + ) + + def dumps_records( + self, + record_type: str, + records: tuple[Record, ...] | list[Record], + *, + protocol_header: str = PROTOCOL_HEADER, + ) -> str: + return dumps_records( + record_type, + records, + protocol_header=protocol_header, + registry=self.registry, + ) + + def loads_records( + self, + payload: str, + expected_record_type: str | None = None, + *, + protocol_header: str = PROTOCOL_HEADER, + ) -> tuple[str, tuple[dict[str, RecordValue], ...]]: + return loads_records( + payload, + expected_record_type, + protocol_header=protocol_header, + registry=self.registry, + ) + + +DEFAULT_SCHEMA_REGISTRY = CommandSchemaRegistry() +# Preserve the existing mutable compatibility surface. New consumers should +# use CommandSchemaRegistry or CommandCodec instead of process-global state. +RECORD_SCHEMAS = DEFAULT_SCHEMA_REGISTRY.schemas + + def register_record_schema(record_type: str, fields: Mapping[str, FieldSpec]) -> None: """Register an application-specific record schema for the wire protocol. @@ -52,27 +129,7 @@ def register_record_schema(record_type: str, fields: Mapping[str, FieldSpec]) -> cannot silently change the meaning of an established record type. """ - if not isinstance(record_type, str) or re.fullmatch(r"[A-Za-z][A-Za-z0-9-]*", record_type) is None: - raise CommandProtocolError( - "record_type must start with a letter and contain only letters, digits, and hyphens" - ) - if record_type in RECORD_SCHEMAS: - raise CommandProtocolError(f"record_type '{record_type}' is already registered") - if not isinstance(fields, Mapping) or not fields: - raise CommandProtocolError("record schema fields must be a non-empty mapping") - - normalized: dict[str, FieldSpec] = {} - for field_name, spec in fields.items(): - if not isinstance(field_name, str) or re.fullmatch(r"[A-Za-z][A-Za-z0-9_]*", field_name) is None: - raise CommandProtocolError( - f"field name '{field_name}' must start with a letter and contain only letters, digits, and underscores" - ) - if not isinstance(spec, FieldSpec) or spec.value_type not in {"string", "boolean"}: - raise CommandProtocolError( - f"field '{field_name}' must use a FieldSpec with value_type 'string' or 'boolean'" - ) - normalized[field_name] = spec - RECORD_SCHEMAS[record_type] = normalized + DEFAULT_SCHEMA_REGISTRY.register(record_type, fields) def dumps_record( @@ -80,8 +137,14 @@ def dumps_record( record: Record, *, protocol_header: str = PROTOCOL_HEADER, + registry: CommandSchemaRegistry | None = None, ) -> str: - return dumps_records(record_type, (record,), protocol_header=protocol_header) + return dumps_records( + record_type, + (record,), + protocol_header=protocol_header, + registry=registry, + ) def dumps_records( @@ -89,8 +152,10 @@ def dumps_records( records: tuple[Record, ...] | list[Record], *, protocol_header: str = PROTOCOL_HEADER, + registry: CommandSchemaRegistry | None = None, ) -> str: - schema = _schema(record_type) + active_registry = registry or DEFAULT_SCHEMA_REGISTRY + schema = active_registry.schema(record_type) if len(records) > MAX_RECORD_COUNT: raise CommandProtocolError(f"record_count exceeds protocol maximum of {MAX_RECORD_COUNT}") lines = [ @@ -114,7 +179,9 @@ def loads_records( expected_record_type: str | None = None, *, protocol_header: str = PROTOCOL_HEADER, + registry: CommandSchemaRegistry | None = None, ) -> tuple[str, tuple[dict[str, RecordValue], ...]]: + active_registry = registry or DEFAULT_SCHEMA_REGISTRY # The wire framing is LF-delimited. `str.splitlines()` also accepts CR, # vertical tab, form feed, and Unicode separators, which would make the # Python decoder more permissive than the Bash and Zsh readers. @@ -136,7 +203,7 @@ def take(label: str) -> str: raise CommandProtocolError(f"unsupported protocol header; expected {protocol_header}") record_type = _metadata_value(take("record_type"), "record_type") - schema = _schema(record_type) + schema = active_registry.schema(record_type) if expected_record_type is not None and record_type != expected_record_type: raise CommandProtocolError(f"expected record_type '{expected_record_type}', got '{record_type}'") @@ -175,14 +242,45 @@ def take(label: str) -> str: return record_type, tuple(records) -def _schema(record_type: str) -> dict[str, FieldSpec]: +def _lookup_schema( + schemas: Mapping[str, dict[str, FieldSpec]], + record_type: str, +) -> dict[str, FieldSpec]: try: - return RECORD_SCHEMAS[record_type] + return schemas[record_type] except KeyError as exc: - supported = ", ".join(sorted(RECORD_SCHEMAS)) + supported = ", ".join(sorted(schemas)) raise CommandProtocolError(f"unsupported record_type '{record_type}'; expected one of: {supported}") from exc +def _validate_and_store_schema( + schemas: dict[str, dict[str, FieldSpec]], + record_type: str, + fields: Mapping[str, FieldSpec], +) -> None: + if not isinstance(record_type, str) or re.fullmatch(r"[A-Za-z][A-Za-z0-9-]*", record_type) is None: + raise CommandProtocolError( + "record_type must start with a letter and contain only letters, digits, and hyphens" + ) + if record_type in schemas: + raise CommandProtocolError(f"record_type '{record_type}' is already registered") + if not isinstance(fields, Mapping) or not fields: + raise CommandProtocolError("record schema fields must be a non-empty mapping") + + normalized: dict[str, FieldSpec] = {} + for field_name, spec in fields.items(): + if not isinstance(field_name, str) or re.fullmatch(r"[A-Za-z][A-Za-z0-9_]*", field_name) is None: + raise CommandProtocolError( + f"field name '{field_name}' must start with a letter and contain only letters, digits, and underscores" + ) + if not isinstance(spec, FieldSpec) or spec.value_type not in {"string", "boolean"}: + raise CommandProtocolError( + f"field '{field_name}' must use a FieldSpec with value_type 'string' or 'boolean'" + ) + normalized[field_name] = spec + schemas[record_type] = normalized + + def _validate_record(schema: Mapping[str, FieldSpec], record_type: str, record: Record) -> None: missing = sorted(set(schema) - set(record)) unknown = sorted(set(record) - set(schema)) diff --git a/lib/python/base_cli/context.py b/lib/python/base_cli/context.py index 039a010..ff95c4c 100644 --- a/lib/python/base_cli/context.py +++ b/lib/python/base_cli/context.py @@ -5,23 +5,38 @@ import os from dataclasses import dataclass, field from pathlib import Path -from typing import Any, Callable +from typing import Any, Callable, Generic, TypeVar from ._cleanup import remove_owned_temp_directory -_current_context: contextvars.ContextVar[Context | None] = contextvars.ContextVar( +_current_context: contextvars.ContextVar[Context[Any, Any, Any] | None] = contextvars.ContextVar( "base_cli_current_context", default=None, ) +ConfigT = TypeVar("ConfigT") +ApplicationStateT = TypeVar("ApplicationStateT") +ServicesT = TypeVar("ServicesT") + +__all__ = [ + "ApplicationStateT", + "ConfigT", + "Context", + "ServicesT", + "get_current_context", + "recover_current_context", + "reset_current_context", + "set_current_context", +] + def _default_history_display_command(cli_name: str, _argv: list[str]) -> str: return cli_name.replace("_", "-") @dataclass -class Context: +class Context(Generic[ConfigT, ApplicationStateT, ServicesT]): """Runtime state and cleanup hooks available to an active CLI command.""" cli_name: str @@ -31,7 +46,7 @@ class Context: cache_dir: Path temp_dir: Path log_file: Path | None - config: dict[str, Any] + config: ConfigT environment: str debug: bool keep_temp: bool @@ -51,8 +66,8 @@ class Context: runtime_owner: str = "default" owner_root: Path | None = None run_root: Path | None = None - application_context: Any = field(default=None, repr=False, compare=False) - services: Any = field(default=None, repr=False, compare=False) + application_context: ApplicationStateT | None = field(default=None, repr=False, compare=False) + services: ServicesT | None = field(default=None, repr=False, compare=False) _run_metadata_path: Path | None = field(default=None, init=False, repr=False, compare=False) _owns_temp_dir: bool = field(default=False, init=False, repr=False, compare=False) _owned_temp_identity: tuple[int, int] | None = field(default=None, init=False, repr=False, compare=False) @@ -147,23 +162,25 @@ def _cleanup_resources(self, *, preserve_temp_ownership: bool) -> None: pass -def set_current_context(context: Context | None) -> contextvars.Token[Context | None]: +def set_current_context( + context: Context[Any, Any, Any] | None, +) -> contextvars.Token[Context[Any, Any, Any] | None]: return _current_context.set(context) -def reset_current_context(token: contextvars.Token[Context | None]) -> None: +def reset_current_context(token: contextvars.Token[Context[Any, Any, Any] | None]) -> None: try: _current_context.reset(token) except BaseException: # pylint: disable=broad-exception-caught recover_current_context(token) -def recover_current_context(token: contextvars.Token[Context | None]) -> None: +def recover_current_context(token: contextvars.Token[Context[Any, Any, Any] | None]) -> None: previous = token.old_value _current_context.set(None if previous is contextvars.Token.MISSING else previous) -def get_current_context() -> Context: +def get_current_context() -> Context[Any, Any, Any]: context = _current_context.get() if context is None: raise RuntimeError("base_cli context is not active. Run inside a base_cli.App command.") diff --git a/lib/python/base_cli/history.py b/lib/python/base_cli/history.py index bf146e0..99b2a44 100644 --- a/lib/python/base_cli/history.py +++ b/lib/python/base_cli/history.py @@ -55,7 +55,7 @@ def utc_now() -> datetime: def build_finished_record( - context: Context, + context: Context[Any, Any, Any], argv: list[str], sensitive_options: set[str], started_at: datetime, diff --git a/lib/python/base_cli/profile.py b/lib/python/base_cli/profile.py index 6349637..c3f961b 100644 --- a/lib/python/base_cli/profile.py +++ b/lib/python/base_cli/profile.py @@ -1,14 +1,29 @@ from __future__ import annotations -from collections.abc import Callable from dataclasses import dataclass from datetime import datetime from pathlib import Path -from typing import Any +from typing import Any, Protocol, cast -from ._runtime import RuntimeLayout, runtime_layout +from ._runtime import runtime_layout from .config import load_yaml_file +from .context import Context from .paths import default_cache_root, make_run_id +from .runtime import RuntimeLayout + +__all__ = [ + "CliProfile", + "ConfigLoader", + "DisplayCommandResolver", + "HistoryDisplayResolver", + "HistoryWriter", + "ProjectDiscovery", + "ProjectInfo", + "RuntimeBinding", + "RuntimeResolver", + "UserConfigLoader", + "WorkspaceRootResolver", +] @dataclass(frozen=True) @@ -38,14 +53,63 @@ class RuntimeBinding: write_identity: bool = False -ProjectDiscovery = Callable[[Path], ProjectInfo | None] -UserConfigLoader = Callable[[], object | None] -ConfigLoader = Callable[[ProjectInfo | None, Path | None], dict[str, Any]] -RuntimeResolver = Callable[[str, ProjectInfo | None], RuntimeBinding] -WorkspaceRootResolver = Callable[[object | None], Path | None] -HistoryWriter = Callable[[Any, list[str], set[str], datetime, int], None] -DisplayCommandResolver = Callable[[], str | None] -HistoryDisplayResolver = Callable[[str, list[str]], str] +class ProjectDiscovery(Protocol): + """Discover consumer-owned project information for the current directory.""" + + def __call__(self, cwd: Path) -> ProjectInfo | None: ... + + +class UserConfigLoader(Protocol): + """Load opaque consumer-owned user configuration.""" + + def __call__(self) -> object | None: ... + + +class ConfigLoader(Protocol): + """Load validated framework configuration and opaque consumer settings.""" + + def __call__( + self, + project: ProjectInfo | None, + explicit_path: Path | None, + ) -> dict[str, Any]: ... + + +class RuntimeResolver(Protocol): + """Resolve the runtime directories and ownership for one invocation.""" + + def __call__(self, cli_name: str, project: ProjectInfo | None) -> RuntimeBinding: ... + + +class WorkspaceRootResolver(Protocol): + """Project a consumer-owned user configuration into a workspace root.""" + + def __call__(self, user_config: object | None) -> Path | None: ... + + +class HistoryWriter(Protocol): + """Persist one completed invocation using the active typed Context.""" + + def __call__( + self, + context: Context[Any, Any, Any], + argv: list[str], + sensitive_parameters: set[str], + started_at: datetime, + exit_code: int, + ) -> None: ... + + +class DisplayCommandResolver(Protocol): + """Resolve the process-facing command label used in diagnostics.""" + + def __call__(self) -> str | None: ... + + +class HistoryDisplayResolver(Protocol): + """Resolve the command label persisted in consumer history.""" + + def __call__(self, cli_name: str, argv: list[str]) -> str: ... def _no_display_command() -> str | None: @@ -75,8 +139,14 @@ class CliProfile: resolve_runtime: RuntimeResolver history_writer: HistoryWriter | None = None display_command: DisplayCommandResolver = _no_display_command - history_display_command: HistoryDisplayResolver = _generic_history_display_command - resolve_workspace_root: WorkspaceRootResolver = _no_workspace_root + history_display_command: HistoryDisplayResolver = cast( + HistoryDisplayResolver, + _generic_history_display_command, + ) + resolve_workspace_root: WorkspaceRootResolver = cast( + WorkspaceRootResolver, + _no_workspace_root, + ) @classmethod def generic( @@ -87,6 +157,7 @@ def generic( discover_project: ProjectDiscovery | None = None, load_user_config: UserConfigLoader | None = None, load_config: ConfigLoader | None = None, + resolve_runtime: RuntimeResolver | None = None, history_display_command: HistoryDisplayResolver | None = None, resolve_workspace_root: WorkspaceRootResolver | None = None, ) -> CliProfile: @@ -96,13 +167,15 @@ def generic( write command history unless the caller supplies those policies. """ return cls( - discover_project=discover_project or _discover_no_project, + discover_project=discover_project or cast(ProjectDiscovery, _discover_no_project), load_user_config=load_user_config or _empty_user_config, - load_config=load_config or _load_explicit_config, - resolve_runtime=_generic_runtime_resolver(cache_root, application_home), + load_config=load_config or cast(ConfigLoader, _load_explicit_config), + resolve_runtime=resolve_runtime or _generic_runtime_resolver(cache_root, application_home), display_command=_no_display_command, - history_display_command=history_display_command or _generic_history_display_command, - resolve_workspace_root=resolve_workspace_root or _no_workspace_root, + history_display_command=history_display_command + or cast(HistoryDisplayResolver, _generic_history_display_command), + resolve_workspace_root=resolve_workspace_root + or cast(WorkspaceRootResolver, _no_workspace_root), ) def _discover_no_project(_cwd: Path) -> ProjectInfo | None: diff --git a/lib/python/base_cli/runtime.py b/lib/python/base_cli/runtime.py new file mode 100644 index 0000000..0fda5b8 --- /dev/null +++ b/lib/python/base_cli/runtime.py @@ -0,0 +1,25 @@ +"""Public runtime layout contract for consumer profiles.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +__all__ = ["RuntimeLayout"] + + +@dataclass(frozen=True) +class RuntimeLayout: + """Filesystem locations owned by one base-cli runtime binding. + + Profiles may construct this value themselves or return it from a custom + runtime resolver. The layout deliberately contains paths only; ownership, + retention, and persistence policy remain profile decisions. + """ + + owner_root: Path + run_root: Path + state_dir: Path + log_dir: Path + cache_dir: Path + temp_dir: Path diff --git a/pyproject.toml b/pyproject.toml index 33f060c..ddcc57f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -38,6 +38,7 @@ dependencies = [ [project.optional-dependencies] dev = [ "build>=1.2", + "mypy>=1.17,<2", "pytest>=8.0", ] diff --git a/tests/test_command_protocol.py b/tests/test_command_protocol.py index 6b81f7c..0d6a681 100644 --- a/tests/test_command_protocol.py +++ b/tests/test_command_protocol.py @@ -4,7 +4,10 @@ from unittest.mock import patch from base_cli.command_protocol import BOOLEAN +from base_cli.command_protocol import CommandCodec from base_cli.command_protocol import CommandProtocolError +from base_cli.command_protocol import CommandSchemaRegistry +from base_cli.command_protocol import DEFAULT_SCHEMA_REGISTRY from base_cli.command_protocol import NULLABLE_STRING from base_cli.command_protocol import RECORD_SCHEMAS from base_cli.command_protocol import STRING @@ -58,6 +61,25 @@ def test_downstream_code_can_register_a_framing_safe_record_schema(self) -> None (record_type, ({"name": "demo", "enabled": True, "note": None},)), ) + def test_consumers_can_isolate_schema_registries_and_codecs(self) -> None: + first = CommandCodec() + second = CommandCodec(CommandSchemaRegistry()) + first.register_schema("isolated", {"name": STRING}) + second.register_schema("isolated", {"enabled": BOOLEAN}) + + first_payload = first.dumps_record("isolated", {"name": "first"}) + second_payload = second.dumps_record("isolated", {"enabled": True}) + + self.assertEqual(first.loads_records(first_payload), ("isolated", ({"name": "first"},))) + self.assertEqual(second.loads_records(second_payload), ("isolated", ({"enabled": True},))) + with self.assertRaisesRegex(CommandProtocolError, "unknown field 'enabled'"): + first.loads_records(second_payload) + with self.assertRaisesRegex(CommandProtocolError, "unknown field 'name'"): + second.loads_records(first_payload) + + def test_default_helpers_remain_backwards_compatible_with_registry_alias(self) -> None: + self.assertIs(RECORD_SCHEMAS, DEFAULT_SCHEMA_REGISTRY.schemas) + def test_record_schema_registration_rejects_invalid_or_duplicate_schemas(self) -> None: with self.assertRaisesRegex(CommandProtocolError, "already registered"): register_record_schema(RECORD_TYPE, {"name": STRING}) diff --git a/tests/test_profile.py b/tests/test_profile.py index 9a7cc0e..6a7145f 100644 --- a/tests/test_profile.py +++ b/tests/test_profile.py @@ -117,3 +117,31 @@ def test_generic_profile_accepts_consumer_history_display_policy(self) -> None: profile = base_cli.CliProfile.generic(history_display_command=formatter) self.assertIs(profile.history_display_command, formatter) + + def test_generic_profile_accepts_a_public_runtime_resolver(self) -> None: + def resolve_runtime( + _cli_name: str, + _project: base_cli.ProjectInfo | None, + ) -> base_cli.RuntimeBinding: + layout = base_cli.RuntimeLayout( + owner_root=Path("owner"), + run_root=Path("run"), + state_dir=Path("state"), + log_dir=Path("logs"), + cache_dir=Path("cache"), + temp_dir=Path("temp"), + ) + return base_cli.RuntimeBinding( + cache_root=Path("cache"), + layout=layout, + application_home=None, + runtime_owner="test", + project_root=None, + project_name=None, + inherited_path=None, + history_parent_run_id=None, + run_id="run", + ) + + profile = base_cli.CliProfile.generic(resolve_runtime=resolve_runtime) + self.assertIs(profile.resolve_runtime, resolve_runtime) diff --git a/tests/test_public_api.py b/tests/test_public_api.py index 971e911..bb4af91 100644 --- a/tests/test_public_api.py +++ b/tests/test_public_api.py @@ -6,7 +6,7 @@ from unittest import mock import base_cli -from base_cli import command_filters, command_protocol, history, lifecycle_options +from base_cli import attachment, command_filters, command_protocol, history, lifecycle_options class PublicApiTests(unittest.TestCase): @@ -32,7 +32,12 @@ def test_version_resolution_ignores_unrelated_ancestor_version_files(self) -> No def test_facade_exports_supported_modules_functions_and_types(self) -> None: expected = { "CommandProtocolError", + "CommandCodec", + "CommandSchemaRegistry", "ConfigurationError", + "AttachmentAdapter", + "AttachmentContract", + "RuntimeLayout", "attach", "command_filters", "command_matches", @@ -62,13 +67,26 @@ def test_module_all_surfaces_are_explicit(self) -> None: set(command_filters.__all__), {"CommandFilterNormalizer", "command_matches", "normalize_command_filter", "normalize_command_filters"}, ) + self.assertEqual( + set(attachment.__all__), + { + "AttachmentAdapter", + "AttachmentContextFactory", + "AttachmentContract", + "AttachmentServiceFactory", + }, + ) self.assertEqual( set(command_protocol.__all__), { "BOOLEAN", + "CommandCodec", "CommandProtocolError", + "CommandSchemaRegistry", + "DEFAULT_SCHEMA_REGISTRY", "FieldSpec", "NULLABLE_STRING", + "RECORD_SCHEMAS", "STRING", "dumps_record", "dumps_records", diff --git a/tests/test_typed_contracts.py b/tests/test_typed_contracts.py new file mode 100644 index 0000000..81f92c5 --- /dev/null +++ b/tests/test_typed_contracts.py @@ -0,0 +1,104 @@ +from __future__ import annotations + +import importlib.util +import tempfile +import unittest +from pathlib import Path +from typing import Any + +import base_cli +from base_cli.testing import invoke + + +@unittest.skipUnless(importlib.util.find_spec("click"), "Click is not installed") +class TypedContractTests(unittest.TestCase): + def test_attachment_contract_is_public_and_preserves_command_identity(self) -> None: + import click + + @click.command(name="contract") + def command() -> None: + pass + + app = base_cli.App(name="contract", log_to_file=False) + adapter: base_cli.AttachmentAdapter[Any] = app + self.assertIs(adapter.attach(command), command) + self.assertIsInstance(base_cli.AttachmentContract, type) + + def test_native_async_callbacks_are_rejected_at_registration(self) -> None: + app = base_cli.App(name="async-registration", log_to_file=False) + + with self.assertRaisesRegex(RuntimeError, "Native async Click callbacks"): + + @app.command() + async def command(_context: base_cli.Context[Any, Any, Any]) -> None: + pass + + def test_attached_async_callbacks_are_rejected_before_mutation(self) -> None: + import click + + @click.command(name="async-attached") + async def command() -> None: + pass + + original_params = tuple(command.params) + app = base_cli.App(name="async-attached", log_to_file=False) + with self.assertRaisesRegex(RuntimeError, "Native async Click callbacks"): + app.attach(command) + self.assertEqual(tuple(command.params), original_params) + self.assertFalse(hasattr(command, "__base_cli_attachment__")) + + def test_sync_callback_returning_awaitable_is_rejected_and_closed(self) -> None: + import click + + closed: list[bool] = [] + + class Awaitable: + def __await__(self) -> Any: + yield + return None + + def close(self) -> None: + closed.append(True) + + @click.command(name="awaitable-result") + def command() -> Any: + return Awaitable() + + app = base_cli.App(name="awaitable-result", log_to_file=False) + app.attach(command) + with tempfile.TemporaryDirectory() as tmpdir: + result = invoke( + app, + [], + home=Path(tmpdir), + reraise_unexpected=True, + ) + self.assertEqual(result.exit_code, 1) + self.assertIsInstance(result.exception, RuntimeError) + self.assertIn("Native async Click callbacks", str(result.exception)) + self.assertEqual(closed, [True]) + + def test_foreign_reserved_markers_are_rejected_transactionally(self) -> None: + import click + + for marker_name in ( + "__base_cli_attachment__", + "__base_cli_lifecycle_instrumented__", + "__base_cli_main_instrumented__", + ): + with self.subTest(marker_name=marker_name): + @click.command(name=f"marker-{marker_name[-5:]}") + def command() -> None: + pass + + original_params = tuple(command.params) + setattr(command, marker_name, True) + app = base_cli.App(name=command.name or "marker", log_to_file=False) + with self.assertRaisesRegex(RuntimeError, "reserved"): + app.attach(command) + self.assertEqual(tuple(command.params), original_params) + self.assertFalse(hasattr(app, "_attached_command") and app._attached_command is command) + + +if __name__ == "__main__": + unittest.main()