Skip to content

Commit ed6a0bc

Browse files
refactor: one is_tier, in user_tables, used everywhere
The tier test was a closure over `table_name` defined inside `_append_platform_attributes`, with a docstring narrating how job tables once acquired `_prov` -- an incident the code no longer has. `is_tier(table_name, tier)` now lives in user_tables.py beside the tier classes it tests, takes both arguments explicitly, and carries one line of docstring. Adopted at every boolean tier test in the library, replacing either a bare `re.fullmatch(X.tier_regexp, name)` or hand-rolled prefix arithmetic: - declare.py -- which tier receives which platform attribute - deploy.py -- which tables add_prov_column touches - schemas.py -- master lookup and the Part test - user_tables.py -- _get_tier itself - migrate.py -- _is_autopopulated_table, which had carried its own copy of the prefix arithmetic, including the `"__" not in table_name[2:]` part-exclusion that each tier_regexp already does by construction What remains using `tier_regexp` directly is the definition of `is_tier` and one `groupdict()` in schemas.py, which extracts the match rather than testing it. deploy.py no longer needs `re` at all. 681 passed, 14 skipped.
1 parent 750283e commit ed6a0bc

5 files changed

Lines changed: 18 additions & 30 deletions

File tree

‎src/datajoint/declare.py‎

Lines changed: 3 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -512,32 +512,22 @@ def _append_platform_attributes(definition, table_name: str, config) -> str:
512512
the exception: it *is* the primary key, so it goes into the key section of a
513513
table that declares none of its own.
514514
"""
515-
from .user_tables import Computed, Imported, Manual
515+
from .user_tables import Computed, Imported, Manual, is_tier
516516

517517
lines = list(definition) if not isinstance(definition, str) else definition.split("\n")
518518

519519
def is_attribute(line: str) -> bool:
520520
stripped = line.strip()
521521
return bool(stripped) and not stripped.startswith("#") and not stripped.startswith("---")
522522

523-
def is_tier(tier) -> bool:
524-
"""Match the tier's own definition rather than re-deriving it from prefixes.
525-
526-
Enumerating prefixes here is what let job tables (`~`) acquire `_prov`
527-
once already, and it silently miscategorises any tier added later. Each
528-
`tier_regexp` also excludes parts by construction, since a part's name
529-
carries its master's and fails the master's own pattern.
530-
"""
531-
return re.fullmatch(tier.tier_regexp, table_name) is not None
532-
533523
separator = next((i for i, line in enumerate(lines) if line.strip().startswith("---")), None)
534524
key_lines = lines[:separator] if separator is not None else lines
535525

536526
secondary = []
537-
if config.jobs.add_job_metadata and (is_tier(Computed) or is_tier(Imported)):
527+
if config.jobs.add_job_metadata and (is_tier(table_name, Computed) or is_tier(table_name, Imported)):
538528
secondary.extend(JOB_METADATA_DEFINITION)
539529

540-
if config.provenance.capture and is_tier(Manual):
530+
if config.provenance.capture and is_tier(table_name, Manual):
541531
secondary.append(PROV_DEFINITION)
542532

543533
# A table that declares no primary key of its own gets the sentinel.

‎src/datajoint/deploy.py‎

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -232,13 +232,12 @@ def add_prov_column(target: "TargetType", dry_run: bool = True) -> dict:
232232
nothing. This is a deploy-time operation: run it before the workers that
233233
will write through it, as with :func:`set_replica_identity`.
234234
"""
235-
import re
236235

237236
from . import provenance
238237
from .declare import PROV_DEFINITION, compile_attribute
239238
from .schemas import _Schema
240239
from .table import Table
241-
from .user_tables import Manual
240+
from .user_tables import Manual, is_tier
242241

243242
if isinstance(target, _Schema):
244243
connection = target.connection
@@ -281,7 +280,7 @@ def add_prov_column(target: "TargetType", dry_run: bool = True) -> dict:
281280
}
282281

283282
for table_name in table_names:
284-
if not re.fullmatch(Manual.tier_regexp, table_name):
283+
if not is_tier(table_name, Manual):
285284
continue
286285
result["tables_analyzed"] += 1
287286

‎src/datajoint/migrate.py‎

Lines changed: 4 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -617,16 +617,10 @@ def _get_existing_columns(connection, database: str, table_name: str) -> set[str
617617

618618

619619
def _is_autopopulated_table(table_name: str) -> bool:
620-
"""Check if a table name indicates a Computed or Imported table."""
621-
# Computed tables start with __ (but not part tables which have __ in middle)
622-
# Imported tables start with _ (but not __)
623-
if table_name.startswith("__"):
624-
# Computed table if no __ after the prefix
625-
return "__" not in table_name[2:]
626-
elif table_name.startswith("_"):
627-
# Imported table
628-
return True
629-
return False
620+
"""Whether a table name denotes a Computed or Imported table."""
621+
from .user_tables import Computed, Imported, is_tier
622+
623+
return is_tier(table_name, Computed) or is_tier(table_name, Imported)
630624

631625

632626
def add_job_metadata_columns(target, dry_run: bool = True) -> dict:

‎src/datajoint/schemas.py‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@
2222
from .heading import Heading
2323
from .jobs import Job
2424
from .table import FreeTable, lookup_class_name
25-
from .user_tables import Computed, Imported, Lookup, Manual, Part, _get_tier
25+
from .user_tables import Computed, Imported, Lookup, Manual, Part, _get_tier, is_tier
2626
from .utils import to_camel_case, user_choice
2727

2828
logger = logging.getLogger(__name__.split(".")[0])
@@ -377,9 +377,9 @@ def make_classes(self, into: dict[str, Any] | None = None) -> None:
377377
class_name = to_camel_case(table_name)
378378
if class_name not in into:
379379
try:
380-
cls = next(cls for cls in master_classes if re.fullmatch(cls.tier_regexp, table_name))
380+
cls = next(cls for cls in master_classes if is_tier(table_name, cls))
381381
except StopIteration:
382-
if re.fullmatch(Part.tier_regexp, table_name):
382+
if is_tier(table_name, Part):
383383
part_tables.append(table_name)
384384
else:
385385
# declare and decorate master table classes

‎src/datajoint/user_tables.py‎

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -278,6 +278,11 @@ def alter(self, prompt=True, context=None):
278278
user_table_classes = (Manual, Lookup, Computed, Imported, Part)
279279

280280

281+
def is_tier(table_name: str, tier) -> bool:
282+
"""Whether a stripped table name belongs to ``tier``."""
283+
return re.fullmatch(tier.tier_regexp, table_name) is not None
284+
285+
281286
def _get_tier(table_name):
282287
"""given the table name, return the user table class."""
283288
# Handle both MySQL backticks and PostgreSQL double quotes
@@ -290,6 +295,6 @@ def _get_tier(table_name):
290295
else:
291296
return None
292297
try:
293-
return next(tier for tier in user_table_classes if re.fullmatch(tier.tier_regexp, extracted_name))
298+
return next(tier for tier in user_table_classes if is_tier(extracted_name, tier))
294299
except StopIteration:
295300
return None

0 commit comments

Comments
 (0)