evidencebound-core 0.4.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- evidencebound/__init__.py +105 -0
- evidencebound/adapters/__init__.py +3 -0
- evidencebound/adapters/google_adk.py +71 -0
- evidencebound/adapters/plain.py +37 -0
- evidencebound/api.py +138 -0
- evidencebound/canonical.py +49 -0
- evidencebound/events.py +28 -0
- evidencebound/graph.py +110 -0
- evidencebound/models.py +164 -0
- evidencebound/persistence.py +59 -0
- evidencebound/providers/__init__.py +1 -0
- evidencebound/providers/ed25519.py +91 -0
- evidencebound/providers/sqlite.py +636 -0
- evidencebound/py.typed +0 -0
- evidencebound/recovery.py +137 -0
- evidencebound/replay.py +31 -0
- evidencebound/signing.py +301 -0
- evidencebound/verification.py +268 -0
- evidencebound_core-0.4.0.dist-info/METADATA +253 -0
- evidencebound_core-0.4.0.dist-info/RECORD +23 -0
- evidencebound_core-0.4.0.dist-info/WHEEL +5 -0
- evidencebound_core-0.4.0.dist-info/licenses/LICENSE +201 -0
- evidencebound_core-0.4.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,105 @@
|
|
|
1
|
+
"""EvidenceBound public API."""
|
|
2
|
+
from .api import EvidenceBound
|
|
3
|
+
from .canonical import CANONICALIZATION_VERSION, CanonicalizationError, canonical_bytes, digest
|
|
4
|
+
from .events import Event, EventSink, ListEventSink, NullEventSink
|
|
5
|
+
from .graph import DependencyGraph, GraphError
|
|
6
|
+
from .models import (
|
|
7
|
+
ActionDecision,
|
|
8
|
+
ApplicabilityStatus,
|
|
9
|
+
Checkpoint,
|
|
10
|
+
EvidenceAssessment,
|
|
11
|
+
EvidenceRecord,
|
|
12
|
+
EvidenceState,
|
|
13
|
+
IntegrityStatus,
|
|
14
|
+
InvalidationReason,
|
|
15
|
+
PolicyBinding,
|
|
16
|
+
ProofReceipt,
|
|
17
|
+
ProvenanceRecord,
|
|
18
|
+
VerificationResult,
|
|
19
|
+
)
|
|
20
|
+
from .persistence import (
|
|
21
|
+
PERSISTED_RECORD_VERSION,
|
|
22
|
+
PERSISTENCE_SCHEMA_VERSION,
|
|
23
|
+
PersistenceError,
|
|
24
|
+
PersistenceIntegrityError,
|
|
25
|
+
PersistenceStore,
|
|
26
|
+
UnsupportedPersistenceSchema,
|
|
27
|
+
)
|
|
28
|
+
from .recovery import RecoveryPlan, action_after_recovery, plan_recovery
|
|
29
|
+
from .replay import ReplayConflict, ReplayGuard, ReplayStatus
|
|
30
|
+
from .signing import (
|
|
31
|
+
KeyState,
|
|
32
|
+
ReceiptSigner,
|
|
33
|
+
ReceiptVerifier,
|
|
34
|
+
SignatureCheck,
|
|
35
|
+
SignatureCheckStatus,
|
|
36
|
+
SignedProofReceipt,
|
|
37
|
+
SignedReceiptStatus,
|
|
38
|
+
SignedReceiptVerification,
|
|
39
|
+
sign_receipt,
|
|
40
|
+
signed_receipt_bytes,
|
|
41
|
+
verify_signed_receipt,
|
|
42
|
+
)
|
|
43
|
+
from .verification import (
|
|
44
|
+
checkpoint_digest,
|
|
45
|
+
evidence_digest,
|
|
46
|
+
verify_checkpoint,
|
|
47
|
+
verify_receipt,
|
|
48
|
+
verify_verification_receipt,
|
|
49
|
+
)
|
|
50
|
+
|
|
51
|
+
__version__ = "0.4.0"
|
|
52
|
+
|
|
53
|
+
__all__ = [
|
|
54
|
+
"ActionDecision",
|
|
55
|
+
"ApplicabilityStatus",
|
|
56
|
+
"CANONICALIZATION_VERSION",
|
|
57
|
+
"CanonicalizationError",
|
|
58
|
+
"Checkpoint",
|
|
59
|
+
"DependencyGraph",
|
|
60
|
+
"EvidenceAssessment",
|
|
61
|
+
"EvidenceBound",
|
|
62
|
+
"EvidenceRecord",
|
|
63
|
+
"EvidenceState",
|
|
64
|
+
"Event",
|
|
65
|
+
"EventSink",
|
|
66
|
+
"GraphError",
|
|
67
|
+
"IntegrityStatus",
|
|
68
|
+
"InvalidationReason",
|
|
69
|
+
"KeyState",
|
|
70
|
+
"ListEventSink",
|
|
71
|
+
"NullEventSink",
|
|
72
|
+
"PERSISTED_RECORD_VERSION",
|
|
73
|
+
"PERSISTENCE_SCHEMA_VERSION",
|
|
74
|
+
"PersistenceError",
|
|
75
|
+
"PersistenceIntegrityError",
|
|
76
|
+
"PersistenceStore",
|
|
77
|
+
"PolicyBinding",
|
|
78
|
+
"ProofReceipt",
|
|
79
|
+
"ProvenanceRecord",
|
|
80
|
+
"ReceiptSigner",
|
|
81
|
+
"ReceiptVerifier",
|
|
82
|
+
"RecoveryPlan",
|
|
83
|
+
"ReplayConflict",
|
|
84
|
+
"ReplayGuard",
|
|
85
|
+
"ReplayStatus",
|
|
86
|
+
"SignatureCheck",
|
|
87
|
+
"SignatureCheckStatus",
|
|
88
|
+
"SignedProofReceipt",
|
|
89
|
+
"SignedReceiptStatus",
|
|
90
|
+
"SignedReceiptVerification",
|
|
91
|
+
"UnsupportedPersistenceSchema",
|
|
92
|
+
"VerificationResult",
|
|
93
|
+
"action_after_recovery",
|
|
94
|
+
"canonical_bytes",
|
|
95
|
+
"checkpoint_digest",
|
|
96
|
+
"digest",
|
|
97
|
+
"evidence_digest",
|
|
98
|
+
"plan_recovery",
|
|
99
|
+
"sign_receipt",
|
|
100
|
+
"signed_receipt_bytes",
|
|
101
|
+
"verify_checkpoint",
|
|
102
|
+
"verify_receipt",
|
|
103
|
+
"verify_signed_receipt",
|
|
104
|
+
"verify_verification_receipt",
|
|
105
|
+
]
|
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
"""Dependency-free integration seam for Google ADK agent callbacks.
|
|
2
|
+
|
|
3
|
+
Google ADK's tested callback contract invokes ``after_agent_callback`` with a
|
|
4
|
+
single keyword argument named ``callback_context``. EvidenceBound does not import
|
|
5
|
+
ADK; callers provide functions that extract evidence and the protected output from
|
|
6
|
+
that context (commonly from ADK state populated through ``output_key``).
|
|
7
|
+
"""
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
from collections.abc import Callable, Mapping
|
|
11
|
+
from dataclasses import dataclass
|
|
12
|
+
from typing import Any
|
|
13
|
+
|
|
14
|
+
from ..api import EvidenceBound
|
|
15
|
+
from ..models import ActionDecision, EvidenceRecord, VerificationResult
|
|
16
|
+
|
|
17
|
+
EvidenceFactory = Callable[[Any], list[EvidenceRecord]]
|
|
18
|
+
OutputFactory = Callable[[Any], object]
|
|
19
|
+
DEFAULT_ADK_STATE_KEY = "evidencebound.verification"
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def adk_consequential_action_allowed(
|
|
23
|
+
state: object,
|
|
24
|
+
*,
|
|
25
|
+
state_key: str = DEFAULT_ADK_STATE_KEY,
|
|
26
|
+
) -> bool:
|
|
27
|
+
"""Return True only for an explicit serialized EvidenceBound ``ALLOW`` result.
|
|
28
|
+
|
|
29
|
+
The helper deliberately treats missing, malformed or incomplete callback state
|
|
30
|
+
as blocked. It does not infer approval from an ADK response event or from agent
|
|
31
|
+
completion alone.
|
|
32
|
+
"""
|
|
33
|
+
if not isinstance(state, Mapping):
|
|
34
|
+
return False
|
|
35
|
+
result = state.get(state_key)
|
|
36
|
+
if not isinstance(result, Mapping):
|
|
37
|
+
return False
|
|
38
|
+
return result.get("action") == ActionDecision.ALLOW.value
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
@dataclass(slots=True)
|
|
42
|
+
class AdkCallbackAdapter:
|
|
43
|
+
"""Callable-compatible ADK after-agent adapter without an ADK dependency.
|
|
44
|
+
|
|
45
|
+
Register ``adapter.after_agent`` as ``after_agent_callback``. The method
|
|
46
|
+
returns ``None`` so it does not replace the agent response. It stores the
|
|
47
|
+
structured EvidenceBound verification in ``callback_context.state`` when
|
|
48
|
+
that object exposes a mutable mapping-like ``state`` attribute.
|
|
49
|
+
"""
|
|
50
|
+
|
|
51
|
+
evidencebound: EvidenceBound
|
|
52
|
+
evidence_factory: EvidenceFactory
|
|
53
|
+
output_factory: OutputFactory
|
|
54
|
+
agent_name: str
|
|
55
|
+
state_key: str = DEFAULT_ADK_STATE_KEY
|
|
56
|
+
last_verification: VerificationResult | None = None
|
|
57
|
+
|
|
58
|
+
def after_agent(self, callback_context: Any) -> None:
|
|
59
|
+
evidence = self.evidence_factory(callback_context)
|
|
60
|
+
output = self.output_factory(callback_context)
|
|
61
|
+
checkpoint = self.evidencebound.checkpoint(
|
|
62
|
+
agent=self.agent_name,
|
|
63
|
+
evidence=evidence,
|
|
64
|
+
output=output,
|
|
65
|
+
)
|
|
66
|
+
result = self.evidencebound.verify(checkpoint)
|
|
67
|
+
self.last_verification = result
|
|
68
|
+
state = getattr(callback_context, "state", None)
|
|
69
|
+
if state is not None:
|
|
70
|
+
state[self.state_key] = result.to_dict()
|
|
71
|
+
return None
|
|
@@ -0,0 +1,37 @@
|
|
|
1
|
+
"""Thin framework-neutral hook adapter."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
from collections.abc import Callable
|
|
5
|
+
from dataclasses import dataclass
|
|
6
|
+
from typing import Generic, TypeVar
|
|
7
|
+
|
|
8
|
+
from ..api import EvidenceBound
|
|
9
|
+
from ..models import Checkpoint, EvidenceRecord, VerificationResult
|
|
10
|
+
|
|
11
|
+
T = TypeVar("T")
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
@dataclass(frozen=True, slots=True)
|
|
15
|
+
class GuardedRun(Generic[T]):
|
|
16
|
+
value: T
|
|
17
|
+
checkpoint: Checkpoint
|
|
18
|
+
verification: VerificationResult
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def run_guarded(
|
|
22
|
+
eb: EvidenceBound,
|
|
23
|
+
*,
|
|
24
|
+
agent: str,
|
|
25
|
+
evidence: list[EvidenceRecord],
|
|
26
|
+
worker: Callable[[], T],
|
|
27
|
+
depends_on: tuple[str, ...] = (),
|
|
28
|
+
) -> GuardedRun[T]:
|
|
29
|
+
value = worker()
|
|
30
|
+
checkpoint = eb.checkpoint(
|
|
31
|
+
agent=agent,
|
|
32
|
+
evidence=evidence,
|
|
33
|
+
output=value,
|
|
34
|
+
depends_on=depends_on,
|
|
35
|
+
)
|
|
36
|
+
verification = eb.verify(checkpoint)
|
|
37
|
+
return GuardedRun(value, checkpoint, verification)
|
evidencebound/api.py
ADDED
|
@@ -0,0 +1,138 @@
|
|
|
1
|
+
"""High-level EvidenceBound facade."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
from collections.abc import Mapping
|
|
5
|
+
from datetime import datetime
|
|
6
|
+
|
|
7
|
+
from .canonical import digest
|
|
8
|
+
from .events import Event, EventSink, NullEventSink
|
|
9
|
+
from .graph import DependencyGraph
|
|
10
|
+
from .models import (
|
|
11
|
+
ActionDecision,
|
|
12
|
+
Checkpoint,
|
|
13
|
+
EvidenceRecord,
|
|
14
|
+
InvalidationReason,
|
|
15
|
+
PolicyBinding,
|
|
16
|
+
ProofReceipt,
|
|
17
|
+
VerificationResult,
|
|
18
|
+
)
|
|
19
|
+
from .recovery import RecoveryPlan, plan_recovery
|
|
20
|
+
from .verification import verify_checkpoint
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class EvidenceBound:
|
|
24
|
+
def __init__(self, *, policy: PolicyBinding, event_sink: EventSink | None = None) -> None:
|
|
25
|
+
self.policy = policy
|
|
26
|
+
self.events = event_sink or NullEventSink()
|
|
27
|
+
self._checkpoints: dict[str, Checkpoint] = {}
|
|
28
|
+
self._verification: dict[str, VerificationResult] = {}
|
|
29
|
+
|
|
30
|
+
def checkpoint(
|
|
31
|
+
self,
|
|
32
|
+
*,
|
|
33
|
+
agent: str,
|
|
34
|
+
evidence: list[EvidenceRecord] | tuple[EvidenceRecord, ...],
|
|
35
|
+
output: object,
|
|
36
|
+
depends_on: list[str] | tuple[str, ...] = (),
|
|
37
|
+
checkpoint_id: str | None = None,
|
|
38
|
+
operation_id: str | None = None,
|
|
39
|
+
) -> Checkpoint:
|
|
40
|
+
if checkpoint_id is None:
|
|
41
|
+
checkpoint_id = digest(
|
|
42
|
+
{
|
|
43
|
+
"agent": agent,
|
|
44
|
+
"evidence": [record.to_dict() for record in evidence],
|
|
45
|
+
"output": output,
|
|
46
|
+
"depends_on": list(depends_on),
|
|
47
|
+
"policy": self.policy.to_dict(),
|
|
48
|
+
"operation_id": operation_id,
|
|
49
|
+
},
|
|
50
|
+
domain="checkpoint-id",
|
|
51
|
+
)[:24]
|
|
52
|
+
if checkpoint_id in self._checkpoints:
|
|
53
|
+
raise ValueError(f"duplicate checkpoint id: {checkpoint_id}")
|
|
54
|
+
evidence_ids = [record.evidence_id for record in evidence]
|
|
55
|
+
if len(evidence_ids) != len(set(evidence_ids)):
|
|
56
|
+
raise ValueError("duplicate evidence ids are not allowed")
|
|
57
|
+
missing = sorted(set(depends_on) - set(self._checkpoints))
|
|
58
|
+
if missing:
|
|
59
|
+
raise ValueError(f"unknown dependencies: {missing}")
|
|
60
|
+
checkpoint = Checkpoint(
|
|
61
|
+
checkpoint_id,
|
|
62
|
+
agent,
|
|
63
|
+
tuple(evidence),
|
|
64
|
+
output, # type: ignore[arg-type]
|
|
65
|
+
self.policy,
|
|
66
|
+
tuple(depends_on),
|
|
67
|
+
operation_id,
|
|
68
|
+
)
|
|
69
|
+
self._checkpoints[checkpoint_id] = checkpoint
|
|
70
|
+
self.events.emit(
|
|
71
|
+
Event("checkpoint_created", {"checkpoint_id": checkpoint_id, "agent": agent})
|
|
72
|
+
)
|
|
73
|
+
return checkpoint
|
|
74
|
+
|
|
75
|
+
def verify(
|
|
76
|
+
self,
|
|
77
|
+
checkpoint: Checkpoint,
|
|
78
|
+
*,
|
|
79
|
+
current_evidence: Mapping[str, EvidenceRecord] | None = None,
|
|
80
|
+
expected_policy: PolicyBinding | None = None,
|
|
81
|
+
prior_receipt: ProofReceipt | None = None,
|
|
82
|
+
now: datetime | None = None,
|
|
83
|
+
) -> VerificationResult:
|
|
84
|
+
result = verify_checkpoint(
|
|
85
|
+
checkpoint,
|
|
86
|
+
current_evidence=current_evidence,
|
|
87
|
+
expected_policy=expected_policy or self.policy,
|
|
88
|
+
prior_receipt=prior_receipt,
|
|
89
|
+
now=now,
|
|
90
|
+
)
|
|
91
|
+
self._verification[checkpoint.checkpoint_id] = result
|
|
92
|
+
self.events.emit(Event("verification_completed", {
|
|
93
|
+
"checkpoint_id": checkpoint.checkpoint_id,
|
|
94
|
+
"action": result.action.value,
|
|
95
|
+
"applicability": result.applicability.value,
|
|
96
|
+
}))
|
|
97
|
+
if result.integrity.value == "FAILED":
|
|
98
|
+
self.events.emit(Event("integrity_failed", {"checkpoint_id": checkpoint.checkpoint_id}))
|
|
99
|
+
if result.action is not ActionDecision.ALLOW:
|
|
100
|
+
self.events.emit(Event("action_blocked", {
|
|
101
|
+
"checkpoint_id": checkpoint.checkpoint_id,
|
|
102
|
+
"decision": result.action.value,
|
|
103
|
+
}))
|
|
104
|
+
return result
|
|
105
|
+
|
|
106
|
+
def graph(self) -> DependencyGraph:
|
|
107
|
+
return DependencyGraph(self._checkpoints.values())
|
|
108
|
+
|
|
109
|
+
def invalidate(
|
|
110
|
+
self,
|
|
111
|
+
*,
|
|
112
|
+
checkpoint_id: str,
|
|
113
|
+
reason: InvalidationReason,
|
|
114
|
+
consequential: set[str] | None = None,
|
|
115
|
+
) -> RecoveryPlan:
|
|
116
|
+
self.events.emit(Event("invalidation_started", {
|
|
117
|
+
"checkpoint_id": checkpoint_id,
|
|
118
|
+
"reason": reason.value,
|
|
119
|
+
}))
|
|
120
|
+
graph = self.graph()
|
|
121
|
+
affected = graph.blast_radius([checkpoint_id])
|
|
122
|
+
self.events.emit(Event("blast_radius_computed", {
|
|
123
|
+
"checkpoint_id": checkpoint_id,
|
|
124
|
+
"affected": list(affected),
|
|
125
|
+
}))
|
|
126
|
+
plan = plan_recovery(
|
|
127
|
+
graph,
|
|
128
|
+
{checkpoint_id: reason},
|
|
129
|
+
verification=self._verification,
|
|
130
|
+
consequential=consequential,
|
|
131
|
+
)
|
|
132
|
+
self.events.emit(Event("recovery_planned", {
|
|
133
|
+
"recompute": list(plan.recompute),
|
|
134
|
+
"reusable": list(plan.reusable),
|
|
135
|
+
"requires_verification": list(plan.requires_verification),
|
|
136
|
+
"blocked": list(plan.blocked),
|
|
137
|
+
}))
|
|
138
|
+
return plan
|
|
@@ -0,0 +1,49 @@
|
|
|
1
|
+
"""Versioned deterministic canonicalization used by EvidenceBound receipts."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
import hashlib
|
|
5
|
+
import json
|
|
6
|
+
from collections.abc import Mapping, Sequence
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
CANONICALIZATION_VERSION = "EBCJ-1"
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class CanonicalizationError(ValueError):
|
|
13
|
+
"""Raised when a value cannot be represented unambiguously by EBCJ-1."""
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def _validate(value: Any, path: str = "$") -> None:
|
|
17
|
+
if value is None or isinstance(value, (bool, int, str)):
|
|
18
|
+
return
|
|
19
|
+
if isinstance(value, float):
|
|
20
|
+
raise CanonicalizationError(f"floats are not supported by EBCJ-1 at {path}")
|
|
21
|
+
if isinstance(value, Mapping):
|
|
22
|
+
for key, item in value.items():
|
|
23
|
+
if not isinstance(key, str):
|
|
24
|
+
raise CanonicalizationError(f"mapping key is not a string at {path}")
|
|
25
|
+
_validate(item, f"{path}.{key}")
|
|
26
|
+
return
|
|
27
|
+
if isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)):
|
|
28
|
+
for index, item in enumerate(value):
|
|
29
|
+
_validate(item, f"{path}[{index}]")
|
|
30
|
+
return
|
|
31
|
+
raise CanonicalizationError(f"unsupported type {type(value).__name__} at {path}")
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def canonical_bytes(value: Any) -> bytes:
|
|
35
|
+
"""Return deterministic UTF-8 JSON bytes for the intentionally narrow EBCJ-1 domain."""
|
|
36
|
+
_validate(value)
|
|
37
|
+
return json.dumps(
|
|
38
|
+
value,
|
|
39
|
+
sort_keys=True,
|
|
40
|
+
separators=(",", ":"),
|
|
41
|
+
ensure_ascii=False,
|
|
42
|
+
allow_nan=False,
|
|
43
|
+
).encode("utf-8")
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def digest(value: Any, *, domain: str) -> str:
|
|
47
|
+
"""SHA-256 digest with versioned domain separation."""
|
|
48
|
+
prefix = f"evidencebound:{domain}:{CANONICALIZATION_VERSION}\0".encode()
|
|
49
|
+
return hashlib.sha256(prefix + canonical_bytes(value)).hexdigest()
|
evidencebound/events.py
ADDED
|
@@ -0,0 +1,28 @@
|
|
|
1
|
+
"""Provider-neutral structured event hooks."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
from dataclasses import dataclass
|
|
5
|
+
from typing import Any, Protocol
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
@dataclass(frozen=True, slots=True)
|
|
9
|
+
class Event:
|
|
10
|
+
name: str
|
|
11
|
+
attributes: dict[str, Any]
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class EventSink(Protocol):
|
|
15
|
+
def emit(self, event: Event) -> None: ...
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class NullEventSink:
|
|
19
|
+
def emit(self, event: Event) -> None:
|
|
20
|
+
return None
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class ListEventSink:
|
|
24
|
+
def __init__(self) -> None:
|
|
25
|
+
self.events: list[Event] = []
|
|
26
|
+
|
|
27
|
+
def emit(self, event: Event) -> None:
|
|
28
|
+
self.events.append(event)
|
evidencebound/graph.py
ADDED
|
@@ -0,0 +1,110 @@
|
|
|
1
|
+
"""Deterministic dependency graph with exact blast-radius traversal."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
from collections import defaultdict, deque
|
|
5
|
+
from collections.abc import Iterable
|
|
6
|
+
|
|
7
|
+
from .models import Checkpoint
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class GraphError(ValueError):
|
|
11
|
+
pass
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class DependencyGraph:
|
|
15
|
+
def __init__(self, checkpoints: Iterable[Checkpoint] = ()) -> None:
|
|
16
|
+
self._nodes: dict[str, Checkpoint] = {}
|
|
17
|
+
for checkpoint in checkpoints:
|
|
18
|
+
if checkpoint.checkpoint_id in self._nodes:
|
|
19
|
+
raise GraphError(f"duplicate checkpoint id: {checkpoint.checkpoint_id}")
|
|
20
|
+
self._nodes[checkpoint.checkpoint_id] = checkpoint
|
|
21
|
+
self._validate_dependencies()
|
|
22
|
+
self._validate_acyclic()
|
|
23
|
+
|
|
24
|
+
@property
|
|
25
|
+
def checkpoint_ids(self) -> tuple[str, ...]:
|
|
26
|
+
return tuple(sorted(self._nodes))
|
|
27
|
+
|
|
28
|
+
def _validate_dependencies(self) -> None:
|
|
29
|
+
for checkpoint in self._nodes.values():
|
|
30
|
+
for dependency in checkpoint.depends_on:
|
|
31
|
+
if dependency not in self._nodes:
|
|
32
|
+
raise GraphError(
|
|
33
|
+
f"checkpoint {checkpoint.checkpoint_id} depends on missing {dependency}"
|
|
34
|
+
)
|
|
35
|
+
|
|
36
|
+
def _children(self) -> dict[str, set[str]]:
|
|
37
|
+
children: dict[str, set[str]] = defaultdict(set)
|
|
38
|
+
for checkpoint in self._nodes.values():
|
|
39
|
+
for dependency in checkpoint.depends_on:
|
|
40
|
+
children[dependency].add(checkpoint.checkpoint_id)
|
|
41
|
+
return children
|
|
42
|
+
|
|
43
|
+
def _validate_acyclic(self) -> None:
|
|
44
|
+
self.topological_order()
|
|
45
|
+
|
|
46
|
+
def topological_order(self) -> tuple[str, ...]:
|
|
47
|
+
indegree = {node: 0 for node in self._nodes}
|
|
48
|
+
children = self._children()
|
|
49
|
+
for checkpoint in self._nodes.values():
|
|
50
|
+
indegree[checkpoint.checkpoint_id] = len(checkpoint.depends_on)
|
|
51
|
+
ready = deque(sorted(node for node, degree in indegree.items() if degree == 0))
|
|
52
|
+
result: list[str] = []
|
|
53
|
+
while ready:
|
|
54
|
+
node = ready.popleft()
|
|
55
|
+
result.append(node)
|
|
56
|
+
for child in sorted(children.get(node, ())):
|
|
57
|
+
indegree[child] -= 1
|
|
58
|
+
if indegree[child] == 0:
|
|
59
|
+
ready.append(child)
|
|
60
|
+
if len(ready) > 1:
|
|
61
|
+
ready = deque(sorted(ready))
|
|
62
|
+
if len(result) != len(self._nodes):
|
|
63
|
+
raise GraphError("dependency graph contains a cycle")
|
|
64
|
+
return tuple(result)
|
|
65
|
+
|
|
66
|
+
def descendants(self, checkpoint_id: str) -> tuple[str, ...]:
|
|
67
|
+
if checkpoint_id not in self._nodes:
|
|
68
|
+
raise GraphError(f"unknown checkpoint: {checkpoint_id}")
|
|
69
|
+
children = self._children()
|
|
70
|
+
seen: set[str] = set()
|
|
71
|
+
queue = deque(sorted(children.get(checkpoint_id, ())))
|
|
72
|
+
while queue:
|
|
73
|
+
current = queue.popleft()
|
|
74
|
+
if current in seen:
|
|
75
|
+
continue
|
|
76
|
+
seen.add(current)
|
|
77
|
+
queue.extend(sorted(children.get(current, ())))
|
|
78
|
+
order = self.topological_order()
|
|
79
|
+
return tuple(node for node in order if node in seen)
|
|
80
|
+
|
|
81
|
+
def ancestors(self, checkpoint_id: str) -> tuple[str, ...]:
|
|
82
|
+
if checkpoint_id not in self._nodes:
|
|
83
|
+
raise GraphError(f"unknown checkpoint: {checkpoint_id}")
|
|
84
|
+
seen: set[str] = set()
|
|
85
|
+
queue = deque(sorted(self._nodes[checkpoint_id].depends_on))
|
|
86
|
+
while queue:
|
|
87
|
+
current = queue.popleft()
|
|
88
|
+
if current in seen:
|
|
89
|
+
continue
|
|
90
|
+
seen.add(current)
|
|
91
|
+
queue.extend(sorted(self._nodes[current].depends_on))
|
|
92
|
+
order = self.topological_order()
|
|
93
|
+
return tuple(node for node in order if node in seen)
|
|
94
|
+
|
|
95
|
+
def blast_radius(self, checkpoint_ids: Iterable[str]) -> tuple[str, ...]:
|
|
96
|
+
roots = set(checkpoint_ids)
|
|
97
|
+
for root in roots:
|
|
98
|
+
if root not in self._nodes:
|
|
99
|
+
raise GraphError(f"unknown checkpoint: {root}")
|
|
100
|
+
affected = set(roots)
|
|
101
|
+
for root in roots:
|
|
102
|
+
affected.update(self.descendants(root))
|
|
103
|
+
order = self.topological_order()
|
|
104
|
+
return tuple(node for node in order if node in affected)
|
|
105
|
+
|
|
106
|
+
def get(self, checkpoint_id: str) -> Checkpoint:
|
|
107
|
+
try:
|
|
108
|
+
return self._nodes[checkpoint_id]
|
|
109
|
+
except KeyError as exc:
|
|
110
|
+
raise GraphError(f"unknown checkpoint: {checkpoint_id}") from exc
|