From dc7cdc42861156fb338a9d29da92ea398b99017e Mon Sep 17 00:00:00 2001 From: Stefan Jansen Date: Mon, 7 Sep 2026 16:39:09 -0400 Subject: [PATCH] save_component carries the component's imports with it A component is validated by re-executing its source in an empty namespace, so anything it reads from another cell has to travel with it. `also` carries a function and `include` carries a value; an import is the third case and neither fits it. `also=[Ridge]` tries to inline sklearn's own source and dies on MultiOutputMixin; `include={'Ridge': Ridge}` writes the class's repr into the file and produces a SyntaxError. So a student who writes `from sklearn.linear_model import Ridge` in an import cell - the normal thing - had their correct model rejected, and was handed two remedies that both fail. That is the one place the conformance machinery was brittle in a way a student would experience as us being wrong. Now the imports are worked out from the names the component actually mentions and written into the saved file as imports. Name resolution walks up to the shallowest ancestor module that still exports the object, so the file records `from sklearn.linear_model import Ridge` rather than the private `sklearn.linear_model._ridge` path the class reports as its own. The genuine missing-symbol case is unchanged and still refused: a value or a helper defined in another cell does not resolve to an importable object, so it falls through to the existing message naming the symbol. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01T1QS8AeUNLQTvUhxZiVSbT --- src/ml4t_coursework/components.py | 108 +++++++++++++++++++++++++++++- tests/test_helper.py | 69 +++++++++++++++++++ 2 files changed, 174 insertions(+), 3 deletions(-) diff --git a/src/ml4t_coursework/components.py b/src/ml4t_coursework/components.py index 703c78a..a9f4dc3 100644 --- a/src/ml4t_coursework/components.py +++ b/src/ml4t_coursework/components.py @@ -8,10 +8,14 @@ from __future__ import annotations +import ast import datetime as dt +import importlib import inspect import json +import sys import textwrap +import types from collections.abc import Callable, Sequence from typing import Any @@ -49,6 +53,102 @@ def _source_of(obj: Any, name: str) -> str: ) from exc +# --- carrying the component's imports with it ------------------------------------------------- +# +# A component is validated on its source re-executed in an empty namespace, so anything it reads +# from another cell has to travel with it. `also` and `include` carry a function and a value. An +# import is the third case and neither of those fits it: `also` would try to inline the library's +# own source and `include` would write its repr into the file. So the import travels as what it +# is, an import statement, worked out from the names the component actually mentions. + + +def _mentioned(source: str) -> set[str]: + """Every bare name and attribute root the source refers to. + + Deliberately not scope-aware. It over-collects - a local variable's name lands here too - and + the caller-namespace lookup below is what filters, because a local name does not resolve to an + importable object. Precise scoping would be more code for the same result. + """ + try: + tree = ast.parse(textwrap.dedent(source)) + except SyntaxError: + return set() + found: set[str] = set() + for node in ast.walk(tree): + if isinstance(node, ast.Name) and isinstance(node.ctx, ast.Load): + found.add(node.id) + elif isinstance(node, ast.Attribute): + root = node + while isinstance(root, ast.Attribute): + root = root.value + if isinstance(root, ast.Name): + found.add(root.id) + return found + + +def _import_line(name: str, obj: Any) -> str | None: + """The import that brings `obj` back as `name` on a cold session, or None if there is none.""" + if isinstance(obj, types.ModuleType): + module = obj.__name__ + return f"import {module}" if module == name else f"import {module} as {name}" + module = getattr(obj, "__module__", None) + symbol = getattr(obj, "__qualname__", None) or getattr(obj, "__name__", None) + if not module or not symbol or "." in symbol or module in {"__main__", "builtins"}: + return None + module = _public_home(module, symbol, obj) + if module is None: + return None + return f"from {module} import {symbol}" if symbol == name else ( + f"from {module} import {symbol} as {name}") + + +def _public_home(module: str, symbol: str, obj: Any) -> str | None: + """The shortest importable path that re-exports `obj`, which is the one the student typed. + + `Ridge.__module__` is `sklearn.linear_model._ridge`, and writing that into a student's file + records a private path that the library is free to rename. Walking up to the shallowest + ancestor that still exports the same object recovers `sklearn.linear_model`. + """ + parts = module.split(".") + best = None + for depth in range(1, len(parts) + 1): + candidate = ".".join(parts[:depth]) + try: + found = importlib.import_module(candidate) + except Exception: + continue + if getattr(found, symbol, None) is obj: + best = candidate + break + if best is None: + # Not reachable from its own package - a class defined in a notebook cell, most often. + return None + return best + + +def _caller_namespace(depth: int) -> dict[str, Any]: + """The notebook cell's names, as seen from `depth` frames above this one.""" + try: + frame = sys._getframe(depth) + except ValueError: # pragma: no cover - only if the stack is shallower than the call + return {} + return {**frame.f_globals, **frame.f_locals} + + +def _carried_imports(source: str, namespace: dict[str, Any]) -> list[str]: + """The import statements the saved file needs so it stands on its own.""" + provided = {"np", "pd", "numpy", "pandas"} + lines = [] + for name in sorted(_mentioned(source) - provided): + obj = namespace.get(name) + if obj is None: + continue + line = _import_line(name, obj) + if line: + lines.append(line) + return lines + + def _rebuild(source: str, symbol: str, name: str) -> Any: """Execute the saved source in a fresh namespace and return the object. @@ -114,9 +214,11 @@ def save_component( raise ValueError( f"{name}: pass the function or class itself, by name, not an instance or a lambda." ) - parts = [f"{key} = {value!r}" for key, value in (include or {}).items()] - parts += [_source_of(helper, name) for helper in also] - parts.append(_source_of(obj, name)) + written = [f"{key} = {value!r}" for key, value in (include or {}).items()] + written += [_source_of(helper, name) for helper in also] + written.append(_source_of(obj, name)) + carried = _carried_imports("\n\n".join(written), _caller_namespace(2)) + parts = (["\n".join(carried)] if carried else []) + written body = "\n\n".join(parts) checked = _rebuild(body, symbol, name) target = folder / f"{name}.py" diff --git a/tests/test_helper.py b/tests/test_helper.py index 3752d28..65655ff 100644 --- a/tests/test_helper.py +++ b/tests/test_helper.py @@ -67,6 +67,75 @@ def fold_splitter(index, n_folds=4): assert "include=" in detail or "also=" in detail, "and say how to fix it" +def test_an_imported_class_travels_with_the_component(capsys): + """The normal thing a student does: import in one cell, use it in the component. + + Neither `also` nor `include` can carry an import - one inlines the library's source, the + other writes its repr - so the import has to be worked out and written as an import. + """ + from sklearn.linear_model import Ridge + + class LinearModel: + def __init__(self): + self.model = Ridge(alpha=1.0) + self.columns_ = None + + def fit(self, X, y): + self.columns_ = list(X.columns) + self.model.fit(X, y) + return self + + def predict(self, X): + if self.columns_ is None: + raise RuntimeError("fit the model before predicting") + return pd.Series(self.model.predict(X[self.columns_]), index=X.index, + name="prediction") + + @property + def coef_(self): + return self.model.coef_ + + result = save_component("model_linear", LinearModel, quiet=True) + assert result.conformant, [c.detail for c in result.checks if not c.passed] + + saved = (project.components_dir() / "model_linear.py").read_text() + assert "from sklearn.linear_model import Ridge" in saved + assert "class Ridge" not in saved, "the import travels, not the library's source" + + assert load_component("model_linear", quiet=True) is not None + assert source_of("model_linear") == "yours" + + +def test_a_module_imported_under_an_alias_travels_too(): + import numpy.linalg as la + + def fold_splitter(index, n_folds=4): + index = pd.Index(index).sort_values() + assert la.norm([1.0]) == 1.0 + block = len(index) // (n_folds + 1) + return [(index[: block * (k + 1) - 21], index[block * (k + 1): block * (k + 2)]) + for k in range(n_folds)] + + assert save_component("fold_splitter", fold_splitter, quiet=True).conformant + saved = (project.components_dir() / "fold_splitter.py").read_text() + assert "import numpy.linalg as la" in saved + + +def test_a_value_from_another_cell_is_still_refused(capsys): + """Carrying imports must not start swallowing the genuine missing-symbol case.""" + outside = 21 + + def fold_splitter(index, n_folds=4): + index = pd.Index(index).sort_values() + block = len(index) // (n_folds + 1) + return [(index[: block * (k + 1) - outside], index[block * (k + 1): block * (k + 2)]) + for k in range(n_folds)] + + result = save_component("fold_splitter", fold_splitter) + assert not result.conformant + assert "outside" in [c.detail for c in result.checks if not c.passed][0] + + def test_include_carries_a_value_the_component_reads(): threshold = 21