diff --git a/src/datajoint/builtin_codecs/__init__.py b/src/datajoint/builtin_codecs/__init__.py index 1f2dd2ec7..d8d462dd9 100644 --- a/src/datajoint/builtin_codecs/__init__.py +++ b/src/datajoint/builtin_codecs/__init__.py @@ -29,14 +29,14 @@ class GraphCodec(dj.Codec): def get_dtype(self, is_store: bool) -> str: return "" # Compose with blob for serialization - def encode(self, graph, *, key=None, store_name=None): + def encode(self, graph, *, key=None, context=None, store_name=None): # Convert graph to a serializable format return { 'nodes': list(graph.nodes(data=True)), 'edges': list(graph.edges(data=True)), } - def decode(self, stored, *, key=None): + def decode(self, stored, *, key=None, context=None): # Reconstruct graph from stored format G = nx.Graph() G.add_nodes_from(stored['nodes']) diff --git a/src/datajoint/builtin_codecs/attach.py b/src/datajoint/builtin_codecs/attach.py index 9aff7bbde..8035810be 100644 --- a/src/datajoint/builtin_codecs/attach.py +++ b/src/datajoint/builtin_codecs/attach.py @@ -50,7 +50,9 @@ def get_dtype(self, is_store: bool) -> str: """Return bytes for in-table, for in-store storage.""" return "" if is_store else "bytes" - def encode(self, value: Any, *, key: dict | None = None, store_name: str | None = None) -> bytes: + def encode( + self, value: Any, *, key: dict | None = None, context: dict | None = None, store_name: str | None = None + ) -> bytes: """ Read file and encode as filename + contents. @@ -80,7 +82,7 @@ def encode(self, value: Any, *, key: dict | None = None, store_name: str | None contents = path.read_bytes() return filename.encode("utf-8") + b"\x00" + contents - def decode(self, stored: bytes, *, key: dict | None = None) -> str: + def decode(self, stored: bytes, *, key: dict | None = None, context: dict | None = None) -> str: """ Extract file to download path and return local path. @@ -104,7 +106,7 @@ def decode(self, stored: bytes, *, key: dict | None = None) -> str: contents = stored[null_pos + 1 :] # Write to download path - config = (key or {}).get("_config") + config = self._codec_config(key, context) if config is None: from ..settings import config # type: ignore[assignment] assert config is not None diff --git a/src/datajoint/builtin_codecs/filepath.py b/src/datajoint/builtin_codecs/filepath.py index 6be05b5cd..f20047963 100644 --- a/src/datajoint/builtin_codecs/filepath.py +++ b/src/datajoint/builtin_codecs/filepath.py @@ -74,7 +74,9 @@ def get_dtype(self, is_store: bool) -> str: ) return "json" - def encode(self, value: Any, *, key: dict | None = None, store_name: str | None = None) -> dict: + def encode( + self, value: Any, *, key: dict | None = None, context: dict | None = None, store_name: str | None = None + ) -> dict: """ Store path reference as JSON metadata. @@ -104,7 +106,7 @@ def encode(self, value: Any, *, key: dict | None = None, store_name: str | None from ..hash_registry import get_store_backend - config = (key or {}).get("_config") + config = self._codec_config(key, context) if config is None: from ..settings import config # type: ignore[assignment] assert config is not None @@ -168,7 +170,7 @@ def encode(self, value: Any, *, key: dict | None = None, store_name: str | None "timestamp": datetime.now(timezone.utc).isoformat(), } - def decode(self, stored: dict, *, key: dict | None = None) -> Any: + def decode(self, stored: dict, *, key: dict | None = None, context: dict | None = None) -> Any: """ Create ObjectRef handle for lazy access. @@ -187,7 +189,7 @@ def decode(self, stored: dict, *, key: dict | None = None) -> Any: from ..objectref import ObjectRef from ..hash_registry import get_store_backend - config = (key or {}).get("_config") + config = self._codec_config(key, context) store_name = stored.get("store") backend = get_store_backend(store_name, config=config) return ObjectRef.from_json(stored, backend=backend) diff --git a/src/datajoint/builtin_codecs/hash.py b/src/datajoint/builtin_codecs/hash.py index db9de7a3d..90a7ff7bc 100644 --- a/src/datajoint/builtin_codecs/hash.py +++ b/src/datajoint/builtin_codecs/hash.py @@ -57,7 +57,9 @@ def get_dtype(self, is_store: bool) -> str: raise DataJointError(" requires @ (in-store storage only)") return "json" - def encode(self, value: bytes, *, key: dict | None = None, store_name: str | None = None) -> dict: + def encode( + self, value: bytes, *, key: dict | None = None, context: dict | None = None, store_name: str | None = None + ) -> dict: """ Store content and return metadata. @@ -78,10 +80,10 @@ def encode(self, value: bytes, *, key: dict | None = None, store_name: str | Non from ..hash_registry import put_hash schema_name = (key or {}).get("_schema", "unknown") - config = (key or {}).get("_config") + config = self._codec_config(key, context) return put_hash(value, schema_name=schema_name, store_name=store_name, config=config) - def decode(self, stored: dict, *, key: dict | None = None) -> bytes: + def decode(self, stored: dict, *, key: dict | None = None, context: dict | None = None) -> bytes: """ Retrieve content using stored metadata. @@ -99,7 +101,7 @@ def decode(self, stored: dict, *, key: dict | None = None) -> bytes: """ from ..hash_registry import get_hash - config = (key or {}).get("_config") + config = self._codec_config(key, context) return get_hash(stored, config=config) def validate(self, value: Any) -> None: diff --git a/src/datajoint/builtin_codecs/npy.py b/src/datajoint/builtin_codecs/npy.py index d9315322b..a2bb7519a 100644 --- a/src/datajoint/builtin_codecs/npy.py +++ b/src/datajoint/builtin_codecs/npy.py @@ -312,6 +312,7 @@ def encode( value: Any, *, key: dict | None = None, + context: dict | None = None, store_name: str | None = None, ) -> dict: """ @@ -337,8 +338,8 @@ def encode( import numpy as np # Extract context using inherited helper - schema, table, field, primary_key = self._extract_context(key) - config = (key or {}).get("_config") + schema, table, field, primary_key = self._extract_context(key, context) + config = self._codec_config(key, context) # Build schema-addressed storage path path, _ = self._build_path(schema, table, field, primary_key, ext=".npy", store_name=store_name, config=config) @@ -360,7 +361,7 @@ def encode( "shape": list(value.shape), } - def decode(self, stored: dict, *, key: dict | None = None) -> NpyRef: + def decode(self, stored: dict, *, key: dict | None = None, context: dict | None = None) -> NpyRef: """ Create lazy NpyRef from stored metadata. @@ -376,6 +377,6 @@ def decode(self, stored: dict, *, key: dict | None = None) -> NpyRef: NpyRef Lazy array reference with metadata access and numpy integration. """ - config = (key or {}).get("_config") + config = self._codec_config(key, context) backend = self._get_backend(stored.get("store"), config=config) return NpyRef(stored, backend) diff --git a/src/datajoint/builtin_codecs/object.py b/src/datajoint/builtin_codecs/object.py index edff4bede..755688262 100644 --- a/src/datajoint/builtin_codecs/object.py +++ b/src/datajoint/builtin_codecs/object.py @@ -81,6 +81,7 @@ def encode( value: Any, *, key: dict | None = None, + context: dict | None = None, store_name: str | None = None, ) -> dict: """ @@ -105,8 +106,8 @@ def encode( from pathlib import Path # Extract context using inherited helper - schema, table, field, primary_key = self._extract_context(key) - config = (key or {}).get("_config") + schema, table, field, primary_key = self._extract_context(key, context) + config = self._codec_config(key, context) # Check for pre-computed metadata (from staged insert) if isinstance(value, dict) and "path" in value: @@ -177,7 +178,7 @@ def encode( return metadata - def decode(self, stored: dict, *, key: dict | None = None) -> Any: + def decode(self, stored: dict, *, key: dict | None = None, context: dict | None = None) -> Any: """ Create ObjectRef handle for lazy access. @@ -195,7 +196,7 @@ def decode(self, stored: dict, *, key: dict | None = None) -> Any: """ from ..objectref import ObjectRef - config = (key or {}).get("_config") + config = self._codec_config(key, context) backend = self._get_backend(stored.get("store"), config=config) return ObjectRef.from_json(stored, backend=backend) diff --git a/src/datajoint/builtin_codecs/schema.py b/src/datajoint/builtin_codecs/schema.py index 260560f4d..9fc215fe6 100644 --- a/src/datajoint/builtin_codecs/schema.py +++ b/src/datajoint/builtin_codecs/schema.py @@ -4,6 +4,8 @@ from __future__ import annotations +import warnings + from ..codecs import Codec from ..errors import DataJointError @@ -24,14 +26,25 @@ class SchemaCodec(Codec, register=False): - ``validate()``: Validate input values Helper Methods: - - ``_extract_context()``: Parse key dict into schema/table/field/pk + - ``_extract_context()``: Parse key/context into schema/table/field/pk + - ``_codec_config()``: Read the calling connection's config - ``_build_path()``: Construct storage path from context - ``_get_backend()``: Get storage backend by name - Both helpers take a ``config`` and fall back to the global ``dj.config`` - without one. Read it off ``key["_config"]`` and pass it through, as below: it - is the calling connection's config, and in a process holding connections for - several users the global one belongs to none of them. + ``_build_path`` and ``_get_backend`` take a ``config`` and fall back to the + global ``dj.config`` without one. Always pass the calling connection's + config: in a process holding connections for several users, the global one + belongs to none of them, and the fallback resolves a different store + silently rather than raising. + + Since 2.3.4 that config arrives in an explicit ``context`` argument rather + than hidden among the primary key values. **Declare ``context=None`` in + ``encode`` and ``decode`` and pass it to the helpers**, as the example below + does. + + The old underscore keys in ``key`` still work and are still populated, so a + codec written before 2.3.4 keeps running — but reading them raises a + ``DeprecationWarning`` and they are removed in 2.4. Comparison with Hash-addressed: - **Schema-addressed** (this): Path from schema structure, no dedup @@ -42,9 +55,9 @@ class SchemaCodec(Codec, register=False): class MyCodec(SchemaCodec): name = "my" - def encode(self, value, *, key=None, store_name=None): - schema, table, field, pk = self._extract_context(key) - config = (key or {}).get("_config") + def encode(self, value, *, key=None, context=None, store_name=None): + schema, table, field, pk = self._extract_context(key, context) + config = self._codec_config(key, context) path, _ = self._build_path( schema, table, field, pk, ext=".dat", store_name=store_name, config=config, @@ -53,8 +66,8 @@ def encode(self, value, *, key=None, store_name=None): backend.put_buffer(serialize(value), path) return {"path": path, "store": store_name, ...} - def decode(self, stored, *, key=None): - config = (key or {}).get("_config") + def decode(self, stored, *, key=None, context=None): + config = self._codec_config(key, context) backend = self._get_backend(stored.get("store"), config=config) return MyRef(stored, backend) @@ -88,15 +101,21 @@ def get_dtype(self, is_store: bool) -> str: raise DataJointError(f"<{self.name}> requires @ (store only)") return "json" - def _extract_context(self, key: dict | None) -> tuple[str, str, str, dict]: + def _extract_context(self, key: dict | None, context: dict | None = None) -> tuple[str, str, str, dict]: """ - Extract schema, table, field, and primary key from context dict. + Extract schema, table, field, and primary key. Parameters ---------- key : dict or None - Context dict with ``_schema``, ``_table``, ``_field``, - and primary key values. + Primary key values. Before 2.3.4 this also carried connection + context under ``_schema``, ``_table``, ``_field`` and ``_config``; + those keys are still populated and still read, with a + ``DeprecationWarning``, when ``context`` is not supplied. + context : dict or None + Connection context with ``schema``, ``table``, ``field`` and + ``config``. Pass the ``context`` argument your ``encode``/``decode`` + received. Returns ------- @@ -104,9 +123,19 @@ def _extract_context(self, key: dict | None) -> tuple[str, str, str, dict]: ``(schema, table, field, primary_key)`` """ key = dict(key) if key else {} - schema = key.pop("_schema", "unknown") - table = key.pop("_table", "unknown") - field = key.pop("_field", "data") + if context is None and any(k.startswith("_") for k in key): + warnings.warn( + "Reading connection context from the `key` dict is deprecated and will " + "be removed in DataJoint 2.4. Accept a `context` argument in encode()/decode() " + "and pass it to _extract_context(key, context). See " + "https://github.com/datajoint/datajoint-python/issues/1550", + DeprecationWarning, + stacklevel=2, + ) + context = context or {} + schema = context.get("schema", key.pop("_schema", "unknown")) + table = context.get("table", key.pop("_table", "unknown")) + field = context.get("field", key.pop("_field", "data")) primary_key = {k: v for k, v in key.items() if not k.startswith("_")} return schema, table, field, primary_key diff --git a/src/datajoint/codecs.py b/src/datajoint/codecs.py index c68e33413..5b323b154 100644 --- a/src/datajoint/codecs.py +++ b/src/datajoint/codecs.py @@ -16,10 +16,10 @@ class GraphCodec(dj.Codec): def get_dtype(self, is_store: bool) -> str: return "" - def encode(self, graph, *, key=None, store_name=None): + def encode(self, graph, *, key=None, context=None, store_name=None): return {'nodes': list(graph.nodes()), 'edges': list(graph.edges())} - def decode(self, stored, *, key=None): + def decode(self, stored, *, key=None, context=None): import networkx as nx G = nx.Graph() G.add_nodes_from(stored['nodes']) @@ -38,6 +38,8 @@ class MyTable(dj.Manual): from __future__ import annotations +import functools +import inspect import json import logging from abc import ABC, abstractmethod @@ -79,10 +81,10 @@ class Codec(ABC): ... def get_dtype(self, is_store: bool) -> str: ... return "" ... - ... def encode(self, graph, *, key=None, store_name=None): + ... def encode(self, graph, *, key=None, context=None, store_name=None): ... return {'nodes': list(graph.nodes()), 'edges': list(graph.edges())} ... - ... def decode(self, stored, *, key=None): + ... def decode(self, stored, *, key=None, context=None): ... import networkx as nx ... G = nx.Graph() ... G.add_nodes_from(stored['nodes']) @@ -176,6 +178,18 @@ def encode(self, value: Any, *, key: dict | None = None, store_name: str | None ------- any Value in the format expected by the dtype. + + Notes + ----- + **Declare ``context`` as well**: ``encode(self, value, *, key=None, + context=None, store_name=None)``. It carries ``schema``, ``table``, + ``field`` and ``config`` — the calling connection's configuration, which + :meth:`_codec_config` reads and which store resolution needs. + + Before 2.3.4 those four arrived inside ``key`` under underscore-prefixed + names. That path still works and DataJoint still populates it, so a codec + written against it keeps running — but it is deprecated and removed in + 2.4. Write new codecs against ``context``. """ ... @@ -195,9 +209,42 @@ def decode(self, stored: Any, *, key: dict | None = None) -> Any: ------- any The reconstructed Python object. + + Notes + ----- + **Declare ``context`` as well**: ``decode(self, stored, *, key=None, + context=None)``. See :meth:`encode`. """ ... + @staticmethod + def _codec_config(key: dict | None = None, context: dict | None = None): + """ + Return the calling connection's config, or None. + + Always thread the result into ``_build_path`` and ``_get_backend``. Those + helpers fall back to the global ``dj.config`` without it, which in a + process holding connections for several users belongs to none of them and + resolves a different store silently rather than raising. + + Prefers ``context["config"]``. Falls back to ``key["_config"]``, the + pre-2.3.4 location, which DataJoint still populates. + + Parameters + ---------- + key : dict, optional + The ``key`` argument the codec received. + context : dict, optional + The ``context`` argument the codec received, if it declares one. + + Returns + ------- + Config or None + """ + if context and context.get("config") is not None: + return context["config"] + return (key or {}).get("_config") + def validate(self, value: Any) -> None: """ Validate a value before encoding. @@ -562,6 +609,28 @@ def lookup_codec(codec_spec: str) -> tuple[Codec, str | None]: # ============================================================================= +@functools.lru_cache(maxsize=None) +def _accepts_kwarg(func, name: str) -> bool: + """ + Whether ``func`` declares a keyword parameter ``name``. + + Used to offer optional arguments -- ``store_name``, ``context`` -- only to + codecs whose signature declares them, so a codec written before either + existed is called exactly as it was. + + Cached on the underlying function: ``inspect.signature`` costs several + times more than a small ``encode`` call, and this runs per attribute per + row. Pass the unbound function (``type(codec).encode``), which is stable + per class, rather than a bound method, which is not. + + Introspection rather than calling and catching ``TypeError``: an ``encode`` + body serializes and uploads, so a ``TypeError`` raised inside it would be + indistinguishable from an unexpected-keyword error at the call site, and + retrying would both mask the real failure and repeat the upload. + """ + return name in inspect.signature(func).parameters + + def decode_attribute(attr, data, squeeze: bool = False, connection=None): """ Decode raw database value using attribute's codec or native type handling. @@ -617,14 +686,21 @@ def decode_attribute(attr, data, squeeze: bool = False, connection=None): elif final_dtype.lower() == "binary(16)": data = uuid_module.UUID(bytes=data) - # Build decode key with config if connection is available + # Build decode key with config if connection is available. The + # underscore key stays for codecs written against it; `context` carries + # the same config to codecs that declare one -- see #1550. decode_key = None + decode_context = None if connection is not None: decode_key = {"_config": connection._config} + decode_context = {"config": connection._config} # Apply decoders in reverse order: innermost first, then outermost for codec in reversed(type_chain): - data = codec.decode(data, key=decode_key) + if _accepts_kwarg(type(codec).decode, "context"): + data = codec.decode(data, key=decode_key, context=decode_context) + else: + data = codec.decode(data, key=decode_key) # Squeeze arrays if requested if squeeze and isinstance(data, np.ndarray): diff --git a/src/datajoint/heading.py b/src/datajoint/heading.py index 2816a22dc..2762a59a8 100644 --- a/src/datajoint/heading.py +++ b/src/datajoint/heading.py @@ -44,12 +44,12 @@ def get_dtype(self, is_store: bool) -> str: f"Codec <{self._codec_name}> is not registered. Define a Codec subclass with name='{self._codec_name}'." ) - def encode(self, value, *, key=None, store_name=None): + def encode(self, value, *, key=None, context=None, store_name=None): raise DataJointError( f"Codec <{self._codec_name}> is not registered. Define a Codec subclass with name='{self._codec_name}'." ) - def decode(self, stored, *, key=None): + def decode(self, stored, *, key=None, context=None): raise DataJointError( f"Codec <{self._codec_name}> is not registered. Define a Codec subclass with name='{self._codec_name}'." ) diff --git a/src/datajoint/spark.py b/src/datajoint/spark.py index 29397b64f..d3550037f 100644 --- a/src/datajoint/spark.py +++ b/src/datajoint/spark.py @@ -59,8 +59,8 @@ class SparkAdapter(Protocol): class FloatArrayCodec(dj.Codec): name = "float_array" - def encode(self, value, *, key=None, store_name=None): ... - def decode(self, stored, *, key=None) -> np.ndarray: ... + def encode(self, value, *, key=None, context=None, store_name=None): ... + def decode(self, stored, *, key=None, context=None) -> np.ndarray: ... def to_spark(self, decoded: np.ndarray, *, key=None) -> list[float]: return decoded.tolist() # → Spark ARRAY diff --git a/src/datajoint/table.py b/src/datajoint/table.py index 5c1c84598..821695f14 100644 --- a/src/datajoint/table.py +++ b/src/datajoint/table.py @@ -1421,6 +1421,15 @@ def __make_placeholder(self, name, value, ignore_extra_fields=False, row=None): "_field": name, "_config": self.connection._config, } + # The same four facts, named without underscores, for codecs that + # declare `context`. The underscore keys above stay in `key` so that + # codecs written against them keep working -- see #1550. + codec_context = { + "schema": self.database, + "table": self.table_name, + "field": name, + "config": self.connection._config, + } # Add primary key values from row if available if row is not None: for pk_name in self.primary_key: @@ -1429,14 +1438,18 @@ def __make_placeholder(self, name, value, ignore_extra_fields=False, row=None): # Apply encoders from outermost to innermost for attr_type in type_chain: - # Pass store_name to encoders that support it (check via introspection) - import inspect - - sig = inspect.signature(attr_type.encode) - if "store_name" in sig.parameters: - value = attr_type.encode(value, key=context, store_name=resolved_store) - else: - value = attr_type.encode(value, key=context) + # Offer store_name and context only to encoders that declare + # them, so a codec written before either existed is called + # exactly as it was. The check is cached per codec class. + from .codecs import _accepts_kwarg + + encode_fn = type(attr_type).encode + kwargs = {} + if _accepts_kwarg(encode_fn, "store_name"): + kwargs["store_name"] = resolved_store + if _accepts_kwarg(encode_fn, "context"): + kwargs["context"] = codec_context + value = attr_type.encode(value, key=context, **kwargs) # Handle NULL values if value is None or (attr.numeric and (value == "" or np.isnan(float(value)))): diff --git a/tests/integration/test_object.py b/tests/integration/test_object.py index ba5cb47ee..86a76cdbd 100644 --- a/tests/integration/test_object.py +++ b/tests/integration/test_object.py @@ -845,12 +845,12 @@ def capture_file(row, **kwargs): codec = ObjectCodec() encode_file_meta = codec.encode( ref_path, - key={ - "_schema": table.database, - "_table": table.class_name, - "_field": "data_file", - "_config": table.connection._config, - "file_id": 801, + key={"file_id": 801}, + context={ + "schema": table.database, + "table": table.class_name, + "field": "data_file", + "config": table.connection._config, }, store_name="local", ) @@ -888,12 +888,12 @@ def capture_folder(row, **kwargs): (ref_dir / "y.bin").write_bytes(b"yy") encode_dir_meta = codec.encode( ref_dir, - key={ - "_schema": table_folder.database, - "_table": table_folder.class_name, - "_field": "data_folder", - "_config": table_folder.connection._config, - "folder_id": 803, + key={"folder_id": 803}, + context={ + "schema": table_folder.database, + "table": table_folder.class_name, + "field": "data_folder", + "config": table_folder.connection._config, }, store_name="local", ) diff --git a/tests/unit/test_codec_context.py b/tests/unit/test_codec_context.py new file mode 100644 index 000000000..de1c02db8 --- /dev/null +++ b/tests/unit/test_codec_context.py @@ -0,0 +1,140 @@ +"""The `context` argument is additive: a codec written before it keeps working. + +DataJoint passes `context` only to codecs whose signature declares it, and keeps +populating the pre-2.3.4 underscore keys in `key`. These tests pin both halves of +that contract, since breaking either would break every third-party codec. +""" + +import inspect +import warnings + +import pytest + +from datajoint.builtin_codecs.schema import SchemaCodec +from datajoint.codecs import Codec + + +class LegacyCodec(SchemaCodec): + """A codec written against the pre-2.3.4 signature. Must not need changing.""" + + name = "legacy_ctx_test" + + def encode(self, value, *, key=None, store_name=None): + schema, table, field, pk = self._extract_context(key) + return {"schema": schema, "table": table, "field": field, "pk": pk, "config": self._codec_config(key)} + + def decode(self, stored, *, key=None): + return self._codec_config(key) + + def validate(self, value): + pass + + +class ModernCodec(SchemaCodec): + """A codec that declares `context`.""" + + name = "modern_ctx_test" + + def encode(self, value, *, key=None, context=None, store_name=None): + schema, table, field, pk = self._extract_context(key, context) + return {"schema": schema, "table": table, "field": field, "pk": pk, "config": self._codec_config(key, context)} + + def decode(self, stored, *, key=None, context=None): + return self._codec_config(key, context) + + def validate(self, value): + pass + + +LEGACY_KEY = {"_schema": "lab", "_table": "subject", "_field": "data", "_config": "CFG", "subject_id": 7} +CONTEXT = {"schema": "lab", "table": "subject", "field": "data", "config": "CFG"} + + +def test_legacy_codec_signature_is_not_asked_for_context(): + """The framework introspects before passing, so the old signature is safe.""" + assert "context" not in inspect.signature(LegacyCodec().encode).parameters + assert "context" in inspect.signature(ModernCodec().encode).parameters + + +def test_legacy_codec_still_resolves_everything_from_key(): + with warnings.catch_warnings(): + warnings.simplefilter("ignore", DeprecationWarning) + out = LegacyCodec().encode(b"x", key=dict(LEGACY_KEY)) + assert out == {"schema": "lab", "table": "subject", "field": "data", "pk": {"subject_id": 7}, "config": "CFG"} + + +def test_modern_codec_resolves_everything_from_context(): + out = ModernCodec().encode(b"x", key={"subject_id": 7}, context=CONTEXT) + assert out == {"schema": "lab", "table": "subject", "field": "data", "pk": {"subject_id": 7}, "config": "CFG"} + + +def test_context_takes_precedence_over_the_legacy_keys(): + """Both are populated during the deprecation window; context wins.""" + out = ModernCodec().encode(b"x", key=dict(LEGACY_KEY), context=CONTEXT) + assert out["config"] == "CFG" + assert out["pk"] == {"subject_id": 7} + + +def test_reading_context_from_key_warns(): + """The deprecated path works and says so.""" + with pytest.warns(DeprecationWarning, match="1550"): + LegacyCodec()._extract_context(dict(LEGACY_KEY)) + + +def test_no_warning_when_context_is_supplied(): + with warnings.catch_warnings(): + warnings.simplefilter("error", DeprecationWarning) + ModernCodec()._extract_context(dict(LEGACY_KEY), CONTEXT) + + +def test_no_warning_for_a_plain_primary_key(): + """A codec that never needed context must not be nagged.""" + with warnings.catch_warnings(): + warnings.simplefilter("error", DeprecationWarning) + assert LegacyCodec()._extract_context({"subject_id": 7}) == ("unknown", "unknown", "data", {"subject_id": 7}) + + +@pytest.mark.parametrize( + "key, context, expected", + [ + ({"_config": "FROM_KEY"}, None, "FROM_KEY"), + (None, {"config": "FROM_CONTEXT"}, "FROM_CONTEXT"), + ({"_config": "FROM_KEY"}, {"config": "FROM_CONTEXT"}, "FROM_CONTEXT"), + ({"_config": "FROM_KEY"}, {}, "FROM_KEY"), + (None, None, None), + ], +) +def test_codec_config_resolution(key, context, expected): + assert Codec._codec_config(key, context) == expected + + +def test_builtin_codecs_declare_context(): + """Every built-in accepts the new argument, so none falls back silently.""" + from datajoint.builtin_codecs import attach, filepath, hash as hash_codec, npy, object as object_codec + + for cls in ( + attach.AttachCodec, + filepath.FilepathCodec, + hash_codec.HashCodec, + npy.NpyCodec, + object_codec.ObjectCodec, + ): + for method in ("encode", "decode"): + params = inspect.signature(getattr(cls, method)).parameters + assert "context" in params, f"{cls.__name__}.{method} does not accept context" + + +def test_capability_check_is_cached_per_class(): + """inspect.signature costs more than a small encode; it must not run per row.""" + from datajoint.codecs import _accepts_kwarg + + _accepts_kwarg.cache_clear() + assert _accepts_kwarg(ModernCodec.encode, "context") is True + assert _accepts_kwarg(LegacyCodec.encode, "context") is False + before = _accepts_kwarg.cache_info() + for _ in range(1000): + _accepts_kwarg(ModernCodec.encode, "context") + _accepts_kwarg(LegacyCodec.encode, "context") + after = _accepts_kwarg.cache_info() + assert after.misses == before.misses, "signature was re-inspected after caching" + assert after.hits - before.hits == 2000