feat(perception): add bounded reference graph
This commit is contained in:
@@ -6,6 +6,7 @@ import re
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
from threading import Event
|
||||
from typing import Final, Protocol
|
||||
|
||||
from .contracts import (
|
||||
@@ -18,6 +19,7 @@ from .contracts import (
|
||||
)
|
||||
|
||||
REFERENCE_GRAPH_CONFIG_SCHEMA: Final = "missioncore.reference-perception-graph-config/v1"
|
||||
REFERENCE_GRAPH_STAGE_IDS: Final = frozenset({"detector", "geometry", "temporal", "threat"})
|
||||
_SHA256 = re.compile(r"^[a-f0-9]{64}$")
|
||||
_IDENTIFIER = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:/-]{0,159}$")
|
||||
|
||||
@@ -35,6 +37,30 @@ class ProviderRole(StrEnum):
|
||||
THREAT = "threat"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SourcePacket:
|
||||
"""Execution-only carrier; the graph never interprets opaque sensor payloads."""
|
||||
|
||||
envelope: SourceEnvelope
|
||||
image_payload: object | None
|
||||
registered_point_increment_payload: object | None
|
||||
pose_payload: object | None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
bindings = (
|
||||
(self.envelope.image.available, self.image_payload, "image"),
|
||||
(
|
||||
self.envelope.registered_point_increment.available,
|
||||
self.registered_point_increment_payload,
|
||||
"registered point increment",
|
||||
),
|
||||
(self.envelope.pose.available, self.pose_payload, "pose"),
|
||||
)
|
||||
for available, payload, label in bindings:
|
||||
if available is not (payload is not None):
|
||||
raise ProviderContractError(f"{label} availability and payload disagree")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ProviderPin:
|
||||
role: ProviderRole
|
||||
@@ -188,8 +214,10 @@ class ReferencePerceptionGraphConfig:
|
||||
if len(set(roles)) != len(roles) or set(roles) != set(ProviderRole):
|
||||
raise ProviderContractError("graph must pin each provider role exactly once")
|
||||
stage_ids = [queue.stage_id for queue in self.queues]
|
||||
if not stage_ids or len(set(stage_ids)) != len(stage_ids):
|
||||
raise ProviderContractError("graph queue policies must be nonempty and unique")
|
||||
if len(set(stage_ids)) != len(stage_ids):
|
||||
raise ProviderContractError("graph queue policies must be unique")
|
||||
if set(stage_ids) != REFERENCE_GRAPH_STAGE_IDS:
|
||||
raise ProviderContractError("graph must bound each reference stage exactly once")
|
||||
|
||||
def to_dict(self) -> dict[str, object]:
|
||||
return {
|
||||
@@ -223,13 +251,13 @@ class ReferencePerceptionGraphConfig:
|
||||
class SourceProvider(Protocol):
|
||||
provider_id: str
|
||||
|
||||
def envelopes(self) -> Iterator[SourceEnvelope]: ...
|
||||
def packets(self, stop_event: Event) -> Iterator[SourcePacket]: ...
|
||||
|
||||
|
||||
class DetectorProvider(Protocol):
|
||||
provider_id: str
|
||||
|
||||
def detect(self, envelope: SourceEnvelope) -> tuple[ObjectProposal2D, ...]: ...
|
||||
def detect(self, packet: SourcePacket) -> tuple[ObjectProposal2D, ...]: ...
|
||||
|
||||
|
||||
class GeometryAssociationProvider(Protocol):
|
||||
@@ -237,7 +265,7 @@ class GeometryAssociationProvider(Protocol):
|
||||
|
||||
def associate(
|
||||
self,
|
||||
envelope: SourceEnvelope,
|
||||
packet: SourcePacket,
|
||||
proposals: tuple[ObjectProposal2D, ...],
|
||||
) -> tuple[ObstacleObservation, ...]: ...
|
||||
|
||||
@@ -247,7 +275,7 @@ class TemporalStateProvider(Protocol):
|
||||
|
||||
def update(
|
||||
self,
|
||||
envelope: SourceEnvelope,
|
||||
packet: SourcePacket,
|
||||
observations: tuple[ObstacleObservation, ...],
|
||||
) -> tuple[TemporalObstacle, ...]: ...
|
||||
|
||||
@@ -257,7 +285,7 @@ class MotionProvider(Protocol):
|
||||
|
||||
def estimate(
|
||||
self,
|
||||
envelope: SourceEnvelope,
|
||||
packet: SourcePacket,
|
||||
obstacles: tuple[TemporalObstacle, ...],
|
||||
) -> tuple[TemporalObstacle, ...]: ...
|
||||
|
||||
|
||||
Reference in New Issue
Block a user