Skip to content

Commit b6f9e31

Browse files
rwgkcursoragent
andcommitted
pathfinder: localize CTK coherence and driver policy
Require exact CTK matching only for authored same-component or companion relationships, so independent artifacts can coexist across minors. Add a Linux-only driver-compatibility override for forward-compatibility deployments without relaxing CTK-coherence checks. Co-authored-by: Cursor <cursoragent@cursor.com>
1 parent 7fe89a2 commit b6f9e31

4 files changed

Lines changed: 351 additions & 37 deletions

File tree

cuda_pathfinder/cuda/pathfinder/_compatibility_guard_rails.py

Lines changed: 85 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -59,7 +59,7 @@
5959
ItemKind: TypeAlias = str
6060
PackagedWith: TypeAlias = str
6161
CtkVersionConstraintArg: TypeAlias = str | SpecifierSet | None
62-
PairwiseItemRelation: TypeAlias = str
62+
PairwiseItemRelationKind: TypeAlias = str
6363

6464
_CTK_VERSION_RE = re.compile(r"^(?P<major>\d+)\.(?P<minor>\d+)")
6565
_CTK_VERSION_CONSTRAINT_ERROR = (
@@ -69,6 +69,13 @@
6969
_PAIRWISE_ITEM_RELATION_NONE = "none"
7070
_PAIRWISE_ITEM_RELATION_EXACT_CTK_MATCH_REQUIRED = "exact-ctk-match-required"
7171

72+
73+
@dataclass(frozen=True, slots=True)
74+
class PairwiseItemRelation:
75+
kind: PairwiseItemRelationKind
76+
reason: str | None = None
77+
78+
7279
_STATIC_LIBS_PACKAGED_WITH: dict[str, PackagedWith] = {
7380
"cudadevrt": "ctk",
7481
}
@@ -87,6 +94,10 @@ class CompatibilityInsufficientMetadataError(CompatibilityCheckError):
8794
"""Raised when v1 compatibility checks cannot reach a definitive answer."""
8895

8996

97+
class DriverCtkCompatibilityError(CompatibilityCheckError):
98+
"""Raised when driver-vs-CTK policy rejects a resolved item."""
99+
100+
90101
@dataclass(frozen=True, slots=True)
91102
class CtkMetadata:
92103
ctk_version: CtkVersion
@@ -140,13 +151,14 @@ def describe(self) -> str:
140151
class CompatibilityResult:
141152
status: str
142153
message: str
154+
error_type: type[CompatibilityCheckError] = CompatibilityCheckError
143155

144156
def require_compatible(self) -> None:
145157
if self.status == "compatible":
146158
return
147159
if self.status == "insufficient_metadata":
148160
raise CompatibilityInsufficientMetadataError(self.message)
149-
raise CompatibilityCheckError(self.message)
161+
raise self.error_type(self.message)
150162

151163

152164
def _parse_ctk_version(cuda_version: str) -> CtkVersion | None:
@@ -434,13 +446,19 @@ def _ctk_constraint_failure_message(item: ResolvedItem, constraint: CtkVersionCo
434446
return f"{item.describe()} resolves to CTK {item.ctk_version}, which does not satisfy ctk_version{constraint}."
435447

436448

437-
def _ctk_pair_mismatch_message(item1: ResolvedItem, item2: ResolvedItem) -> str:
449+
def _ctk_pair_mismatch_message(
450+
item1: ResolvedItem,
451+
item2: ResolvedItem,
452+
relation: PairwiseItemRelation,
453+
) -> str:
438454
assert item1.ctk_version is not None
439455
assert item2.ctk_version is not None
456+
assert relation.reason is not None
457+
requirement_reason = relation.reason[:1].upper() + relation.reason[1:]
440458
return (
441459
f"{item1.describe()} resolves to CTK {item1.ctk_version}, while "
442460
f"{item2.describe()} resolves to CTK {item2.ctk_version}. "
443-
"v1 requires an exact CTK major.minor match."
461+
f"{requirement_reason}, so v1 requires an exact CTK major.minor match."
444462
)
445463

446464

@@ -453,11 +471,26 @@ def _driver_major_mismatch_message(driver_cuda_version: DriverCudaVersion, item:
453471
)
454472

455473

456-
def _compatible_pair_message(driver_cuda_version: DriverCudaVersion, item1: ResolvedItem, item2: ResolvedItem) -> str:
474+
def _compatible_pair_message(
475+
driver_cuda_version: DriverCudaVersion,
476+
item1: ResolvedItem,
477+
item2: ResolvedItem,
478+
relation: PairwiseItemRelation,
479+
) -> str:
457480
assert item1.ctk_version is not None
481+
assert item2.ctk_version is not None
482+
if relation.kind == _PAIRWISE_ITEM_RELATION_NONE:
483+
return (
484+
f"{item1.describe()} resolves to CTK {item1.ctk_version}, "
485+
f"{item2.describe()} resolves to CTK {item2.ctk_version}, "
486+
"and v1 does not require exact CTK lockstep for this pair. "
487+
f"Driver version {driver_cuda_version.encoded} satisfies the v1 driver guard rail."
488+
)
489+
assert relation.reason is not None
458490
return (
459-
f"{item1.describe()} and {item2.describe()} both resolve to CTK {item1.ctk_version}, "
460-
f"and driver version {driver_cuda_version.encoded} satisfies the v1 driver guard rail."
491+
f"{item1.describe()} and {item2.describe()} both resolve to CTK {item1.ctk_version}. "
492+
f"{relation.reason[:1].upper() + relation.reason[1:]}, and driver version "
493+
f"{driver_cuda_version.encoded} satisfies the v1 driver guard rail."
461494
)
462495

463496

@@ -473,18 +506,40 @@ def _ctk_metadata_result(item: ResolvedItem) -> CompatibilityResult | None:
473506
return CompatibilityResult(status="insufficient_metadata", message=_missing_ctk_metadata_message(item))
474507

475508

476-
def _classify_pairwise_item_relation(item1: ResolvedItem, item2: ResolvedItem) -> PairwiseItemRelation:
477-
if item1.packaged_with == "driver" or item2.packaged_with == "driver":
478-
return _PAIRWISE_ITEM_RELATION_NONE
479-
return _PAIRWISE_ITEM_RELATION_EXACT_CTK_MATCH_REQUIRED
509+
def _shared_ctk_companion_tags(item1: ResolvedItem, item2: ResolvedItem) -> tuple[str, ...]:
510+
return tuple(sorted(set(item1.ctk_companion_tags).intersection(item2.ctk_companion_tags)))
480511

481512

482-
def _ctk_coherence_result(item1: ResolvedItem, item2: ResolvedItem) -> CompatibilityResult | None:
513+
def _classify_pairwise_item_relation(item1: ResolvedItem, item2: ResolvedItem) -> PairwiseItemRelation:
514+
if item1.packaged_with == "driver" or item2.packaged_with == "driver":
515+
return PairwiseItemRelation(_PAIRWISE_ITEM_RELATION_NONE)
516+
if item1.dynamic_link_component is not None and item1.dynamic_link_component == item2.dynamic_link_component:
517+
return PairwiseItemRelation(
518+
_PAIRWISE_ITEM_RELATION_EXACT_CTK_MATCH_REQUIRED,
519+
reason=f"they are in the same authored dynamic-link component {item1.dynamic_link_component!r}",
520+
)
521+
shared_companion_tags = _shared_ctk_companion_tags(item1, item2)
522+
if shared_companion_tags:
523+
if len(shared_companion_tags) == 1:
524+
tag_description = repr(shared_companion_tags[0])
525+
reason = f"they share the authored companion tag {tag_description}"
526+
else:
527+
tags_description = ", ".join(repr(tag) for tag in shared_companion_tags)
528+
reason = f"they share the authored companion tags {tags_description}"
529+
return PairwiseItemRelation(_PAIRWISE_ITEM_RELATION_EXACT_CTK_MATCH_REQUIRED, reason=reason)
530+
return PairwiseItemRelation(_PAIRWISE_ITEM_RELATION_NONE)
531+
532+
533+
def _ctk_coherence_result(
534+
item1: ResolvedItem,
535+
item2: ResolvedItem,
536+
relation: PairwiseItemRelation,
537+
) -> CompatibilityResult | None:
483538
assert item1.ctk_version is not None
484539
assert item2.ctk_version is not None
485540
if item1.ctk_version == item2.ctk_version:
486541
return None
487-
return CompatibilityResult(status="incompatible", message=_ctk_pair_mismatch_message(item1, item2))
542+
return CompatibilityResult(status="incompatible", message=_ctk_pair_mismatch_message(item1, item2, relation))
488543

489544

490545
def _pipeline_compatibility_result(_item1: ResolvedItem, _item2: ResolvedItem) -> CompatibilityResult | None:
@@ -493,16 +548,21 @@ def _pipeline_compatibility_result(_item1: ResolvedItem, _item2: ResolvedItem) -
493548
return None
494549

495550

496-
def _pairwise_policy_result(item1: ResolvedItem, item2: ResolvedItem) -> CompatibilityResult | None:
497-
relation = _classify_pairwise_item_relation(item1, item2)
498-
if relation == _PAIRWISE_ITEM_RELATION_NONE:
551+
def _pairwise_policy_result(
552+
item1: ResolvedItem,
553+
item2: ResolvedItem,
554+
relation: PairwiseItemRelation | None = None,
555+
) -> CompatibilityResult | None:
556+
if relation is None:
557+
relation = _classify_pairwise_item_relation(item1, item2)
558+
if relation.kind == _PAIRWISE_ITEM_RELATION_NONE:
499559
return None
500-
if relation == _PAIRWISE_ITEM_RELATION_EXACT_CTK_MATCH_REQUIRED:
501-
result = _ctk_coherence_result(item1, item2)
560+
if relation.kind == _PAIRWISE_ITEM_RELATION_EXACT_CTK_MATCH_REQUIRED:
561+
result = _ctk_coherence_result(item1, item2, relation)
502562
if result is not None:
503563
return result
504564
return _pipeline_compatibility_result(item1, item2)
505-
raise AssertionError(f"Unhandled pairwise item relation: {relation!r}")
565+
raise AssertionError(f"Unhandled pairwise item relation: {relation.kind!r}")
506566

507567

508568
def _driver_compatibility_result(
@@ -514,6 +574,7 @@ def _driver_compatibility_result(
514574
return CompatibilityResult(
515575
status="incompatible",
516576
message=_driver_major_mismatch_message(driver_cuda_version, item),
577+
error_type=DriverCtkCompatibilityError,
517578
)
518579

519580

@@ -528,7 +589,8 @@ def compatibility_check(
528589
if result is not None:
529590
return result
530591

531-
result = _pairwise_policy_result(item1, item2)
592+
relation = _classify_pairwise_item_relation(item1, item2)
593+
result = _pairwise_policy_result(item1, item2, relation)
532594
if result is not None:
533595
return result
534596

@@ -538,7 +600,7 @@ def compatibility_check(
538600

539601
return CompatibilityResult(
540602
status="compatible",
541-
message=_compatible_pair_message(driver_cuda_version, item1, item2),
603+
message=_compatible_pair_message(driver_cuda_version, item1, item2, relation),
542604
)
543605

544606

@@ -606,8 +668,8 @@ def _reset_for_testing(self) -> None:
606668

607669
def _register_and_check(self, item: ResolvedItem) -> None:
608670
# Driver libraries come from the installed display driver rather than a
609-
# CUDA Toolkit line, so they do not need CTK metadata and must not lock
610-
# the process-wide CTK anchor.
671+
# CUDA Toolkit line, so they do not need CTK metadata and must not
672+
# create CTK coherence relations by themselves.
611673
if item.packaged_with == "driver":
612674
self._remember(item)
613675
return

cuda_pathfinder/cuda/pathfinder/_process_wide_compatibility_guard_rails.py

Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@
1414
from cuda.pathfinder._compatibility_guard_rails import (
1515
CompatibilityGuardRails,
1616
CompatibilityInsufficientMetadataError,
17+
DriverCtkCompatibilityError,
1718
)
1819
from cuda.pathfinder._dynamic_libs.load_dl_common import LoadedDL
1920
from cuda.pathfinder._dynamic_libs.load_nvidia_dynamic_lib import (
@@ -52,6 +53,10 @@
5253
_COMPATIBILITY_GUARD_RAILS_MODES = ("off", "best_effort", "strict")
5354
_COMPATIBILITY_GUARD_RAILS_DEFAULT_MODE = "strict"
5455
assert _COMPATIBILITY_GUARD_RAILS_DEFAULT_MODE in _COMPATIBILITY_GUARD_RAILS_MODES
56+
_DRIVER_COMPATIBILITY_ENV_VAR = "CUDA_PATHFINDER_DRIVER_COMPATIBILITY"
57+
_DRIVER_COMPATIBILITY_MODES = ("default", "assume_forward_compatibility")
58+
_DRIVER_COMPATIBILITY_DEFAULT_MODE = "default"
59+
assert _DRIVER_COMPATIBILITY_DEFAULT_MODE in _DRIVER_COMPATIBILITY_MODES
5560

5661

5762
class _ProcessWideGuardRailsApi(Protocol):
@@ -93,6 +98,37 @@ def _compatibility_guard_rails_mode() -> str:
9398
)
9499

95100

101+
def _driver_compatibility_mode() -> str:
102+
value = os.environ.get(_DRIVER_COMPATIBILITY_ENV_VAR)
103+
if not value:
104+
return _DRIVER_COMPATIBILITY_DEFAULT_MODE
105+
if value not in _DRIVER_COMPATIBILITY_MODES:
106+
allowed_values = ", ".join(repr(mode) for mode in _DRIVER_COMPATIBILITY_MODES)
107+
raise RuntimeError(
108+
f"Invalid {_DRIVER_COMPATIBILITY_ENV_VAR}={value!r}. "
109+
f"Allowed values: {allowed_values}. "
110+
f"Unset or empty defaults to {_DRIVER_COMPATIBILITY_DEFAULT_MODE!r}."
111+
)
112+
if value == "assume_forward_compatibility" and not sys.platform.startswith("linux"):
113+
raise RuntimeError(f"{_DRIVER_COMPATIBILITY_ENV_VAR}={value!r} is only supported on Linux.")
114+
return value
115+
116+
117+
def _driver_compatibility_override_hint() -> str:
118+
return (
119+
"On supported Linux systems that intentionally rely on NVIDIA forward compatibility "
120+
f"(`cuda-compat-*`), set {_DRIVER_COMPATIBILITY_ENV_VAR}=assume_forward_compatibility "
121+
"to bypass this driver-vs-CTK check. This does not relax CTK-coherence checks "
122+
"between headers, libraries, and compiler/JIT components."
123+
)
124+
125+
126+
def _with_driver_compatibility_hint(message: str) -> str:
127+
if _DRIVER_COMPATIBILITY_ENV_VAR in message:
128+
return message
129+
return f"{message} {_driver_compatibility_override_hint()}"
130+
131+
96132
def _public_module() -> _PublicPathfinderModule | None:
97133
public_module = sys.modules.get("cuda.pathfinder")
98134
if public_module is None:
@@ -121,6 +157,7 @@ def _reset_process_wide_compatibility_guard_rails() -> None:
121157

122158

123159
def _try_process_wide_guard_rails_then_fallback(guard_rails_call: Callable[[], _T], raw_call: Callable[[], _T]) -> _T:
160+
driver_compatibility_mode = _driver_compatibility_mode()
124161
mode = _compatibility_guard_rails_mode()
125162
if mode == "off":
126163
return raw_call()
@@ -130,6 +167,12 @@ def _try_process_wide_guard_rails_then_fallback(guard_rails_call: Callable[[], _
130167
if mode == "best_effort":
131168
return raw_call()
132169
raise
170+
except DriverCtkCompatibilityError as exc:
171+
if driver_compatibility_mode == "assume_forward_compatibility":
172+
return raw_call()
173+
if sys.platform.startswith("linux"):
174+
raise DriverCtkCompatibilityError(_with_driver_compatibility_hint(str(exc))) from exc
175+
raise
133176

134177

135178
def _cache_clear_with_process_state_reset(cache_clear: Callable[[], object]) -> Callable[[], None]:

0 commit comments

Comments
 (0)