feat: keep RGB-D captures only on workstation

This commit is contained in:
2026-08-10 17:14:44 +08:00
parent d02b99d2aa
commit 67c9b65768
15 changed files with 2817 additions and 154 deletions

View File

@@ -32,6 +32,15 @@ import yaml
ToggleAction = Literal["start", "stop"]
def _valid_ros_topic(topic: Any) -> bool:
return (
isinstance(topic, str)
and topic.startswith("/")
and topic.strip() == topic
and not any(character.isspace() for character in topic)
)
def left_joystick_pressed(data: Mapping[str, Any]) -> bool | None:
"""Strictly parse the live-frame ``button_joystick.left`` value.
@@ -195,6 +204,55 @@ class RecordingToggleGate:
return "active" if self.active else "idle"
@dataclass(frozen=True)
class OptionalTopicGroupConfig:
"""Non-fatal completeness and rate checks for an optional sensor group."""
topics: Sequence[str]
minimum_topic_rates_hz: Mapping[str, float] = field(default_factory=dict)
def __post_init__(self) -> None:
topics = (
tuple(self.topics)
if isinstance(self.topics, Sequence)
and not isinstance(self.topics, str)
else ()
)
if not topics:
raise ValueError("optional topic group must contain at least one topic")
if any(not _valid_ros_topic(topic) for topic in topics):
raise ValueError(
"optional topic group topics must be absolute ROS topic names"
)
if len(set(topics)) != len(topics):
raise ValueError("optional topic group topics must not contain duplicates")
if not isinstance(self.minimum_topic_rates_hz, Mapping):
raise ValueError("optional minimum topic rates must be a mapping")
rates: dict[str, float] = {}
for topic, rate in self.minimum_topic_rates_hz.items():
if topic not in topics:
raise ValueError(
"optional minimum-rate topics must be a subset of group topics"
)
if (
isinstance(rate, bool)
or not isinstance(rate, (int, float))
or not math.isfinite(rate)
or rate <= 0.0
):
raise ValueError(
f"optional minimum topic rate for {topic!r} must be "
"positive and finite"
)
rates[topic] = float(rate)
object.__setattr__(self, "topics", topics)
object.__setattr__(
self,
"minimum_topic_rates_hz",
MappingProxyType(rates),
)
@dataclass(frozen=True)
class RecorderConfig:
"""Static configuration for :class:`DataRecorderManager`."""
@@ -203,6 +261,10 @@ class RecorderConfig:
topics: Sequence[str] = ()
required_topics: Sequence[str] = ()
minimum_topic_rates_hz: Mapping[str, float] = field(default_factory=dict)
optional_topic_groups: Mapping[str, OptionalTopicGroupConfig] = field(
default_factory=dict
)
retain_failed_episodes: bool = True
minimum_free_bytes: int = 5 * 1024**3
max_duration_seconds: float = 30 * 60.0
poll_interval_seconds: float = 0.1
@@ -219,25 +281,13 @@ class RecorderConfig:
minimum_topic_rates: dict[str, float] = {}
if not topics:
raise ValueError("at least one recording topic is required")
if any(
not isinstance(topic, str)
or not topic.startswith("/")
or topic.strip() != topic
or any(character.isspace() for character in topic)
for topic in topics
):
if any(not _valid_ros_topic(topic) for topic in topics):
raise ValueError("recording topics must be absolute ROS topic names")
if len(set(topics)) != len(topics):
raise ValueError("recording topics must not contain duplicates")
if len(set(required_topics)) != len(required_topics):
raise ValueError("required topics must not contain duplicates")
if any(
not isinstance(topic, str)
or not topic.startswith("/")
or topic.strip() != topic
or any(character.isspace() for character in topic)
for topic in required_topics
):
if any(not _valid_ros_topic(topic) for topic in required_topics):
raise ValueError("required topics must be absolute ROS topic names")
if not set(required_topics).issubset(topics):
raise ValueError("required topics must be a subset of recording topics")
@@ -258,6 +308,45 @@ class RecorderConfig:
f"minimum topic rate for {topic!r} must be positive and finite"
)
minimum_topic_rates[topic] = float(rate)
if not isinstance(self.optional_topic_groups, Mapping):
raise ValueError("optional topic groups must be a mapping")
optional_groups: dict[str, OptionalTopicGroupConfig] = {}
grouped_topics: set[str] = set()
for name, group in self.optional_topic_groups.items():
if (
not isinstance(name, str)
or not name
or len(name) > 64
or any(
character
not in "-_0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"
for character in name
)
):
raise ValueError(
"optional topic group names must contain 1-64 safe characters"
)
if not isinstance(group, OptionalTopicGroupConfig):
raise ValueError(
f"optional topic group {name!r} has an invalid configuration"
)
group_topics = set(group.topics)
if not group_topics.issubset(topics):
raise ValueError(
f"optional topic group {name!r} must be a subset of recording topics"
)
if group_topics.intersection(required_topics):
raise ValueError(
f"optional topic group {name!r} must be disjoint from required topics"
)
overlap = group_topics.intersection(grouped_topics)
if overlap:
raise ValueError(
"optional topic groups must be disjoint; repeated topics: "
+ ", ".join(sorted(overlap))
)
grouped_topics.update(group_topics)
optional_groups[name] = group
if (
isinstance(self.minimum_free_bytes, bool)
or not isinstance(self.minimum_free_bytes, int)
@@ -277,6 +366,8 @@ class RecorderConfig:
raise ValueError("ROS 2 executable must be a non-empty string")
if type(self.validate_bag_info) is not bool:
raise ValueError("validate bag info must be a bool")
if type(self.retain_failed_episodes) is not bool:
raise ValueError("retain failed episodes must be a bool")
object.__setattr__(self, "base_directory", base)
object.__setattr__(self, "topics", topics)
object.__setattr__(self, "required_topics", required_topics)
@@ -285,6 +376,11 @@ class RecorderConfig:
"minimum_topic_rates_hz",
MappingProxyType(minimum_topic_rates),
)
object.__setattr__(
self,
"optional_topic_groups",
MappingProxyType(optional_groups),
)
class RecorderProcess(Protocol):
@@ -325,6 +421,8 @@ class _Episode:
required_topic_message_counts: dict[str, int] | None = None
minimum_topic_message_counts: dict[str, int] | None = None
observed_topic_rates_hz: dict[str, float] | None = None
optional_topic_groups: dict[str, Any] | None = None
data_quality_warnings: list[str] | None = None
metadata_duration_nanoseconds: int | None = None
process: RecorderProcess | None = None
stdout_stream: Any = None
@@ -767,6 +865,8 @@ class DataRecorderManager:
episode.minimum_topic_message_counts
),
"observed_topic_rates_hz": episode.observed_topic_rates_hz,
"optional_topic_groups": episode.optional_topic_groups,
"data_quality_warnings": episode.data_quality_warnings,
"metadata_duration_nanoseconds": (
episode.metadata_duration_nanoseconds
),
@@ -859,6 +959,7 @@ class DataRecorderManager:
topic: all_counts.get(topic, 0) for topic in self.config.topics
}
episode.required_topic_message_counts = required_counts
self._observe_optional_topic_groups(episode, information, all_counts)
missing = [
topic for topic in self.config.required_topics if topic not in all_counts
]
@@ -877,6 +978,101 @@ class DataRecorderManager:
self._validate_minimum_topic_rates(episode, information, all_counts)
return required_counts
def _observe_optional_topic_groups(
self,
episode: _Episode,
information: Mapping[str, Any],
all_counts: Mapping[str, int],
) -> None:
"""Classify optional sensor quality without invalidating core data."""
duration_nanoseconds: int | None = None
duration = information.get("duration")
if isinstance(duration, Mapping):
candidate = duration.get("nanoseconds")
if (
type(candidate) is int
and candidate > 0
):
duration_nanoseconds = candidate
episode.metadata_duration_nanoseconds = candidate
duration_seconds = (
None
if duration_nanoseconds is None
else duration_nanoseconds / 1_000_000_000.0
)
observations: dict[str, Any] = {}
warnings: list[str] = []
for name, group in self.config.optional_topic_groups.items():
counts = {
topic: int(all_counts.get(topic, 0)) for topic in group.topics
}
empty_topics = [
topic for topic, count in counts.items() if count <= 0
]
all_absent = len(empty_topics) == len(group.topics)
minimum_counts: dict[str, int | None] = {}
observed_rates: dict[str, float | None] = {}
below_rate: list[str] = []
for topic, minimum_rate in group.minimum_topic_rates_hz.items():
if duration_seconds is None:
minimum_counts[topic] = None
observed_rates[topic] = None
if not all_absent:
below_rate.append(topic)
continue
minimum_count = math.floor(duration_seconds * minimum_rate)
observed_rate = counts[topic] / duration_seconds
minimum_counts[topic] = minimum_count
observed_rates[topic] = observed_rate
if not all_absent and counts[topic] < minimum_count:
below_rate.append(topic)
if all_absent:
state = "absent"
elif empty_topics:
state = "partial"
warnings.append(
f"optional topic group {name!r} is partial; zero-message "
"topics: " + ", ".join(empty_topics)
)
elif below_rate:
state = "low_rate"
if duration_seconds is None:
warnings.append(
f"optional topic group {name!r} rate could not be "
"validated because bag duration is unavailable"
)
else:
details = [
f"{topic}={observed_rates[topic]:.3f}Hz<"
f"{group.minimum_topic_rates_hz[topic]:.3f}Hz"
for topic in below_rate
]
warnings.append(
f"optional topic group {name!r} is below its observed "
"minimum rate: " + "; ".join(details)
)
else:
state = "healthy"
observations[name] = {
"state": state,
"topics": list(group.topics),
"topic_message_counts": counts,
"minimum_topic_rates_hz": dict(
group.minimum_topic_rates_hz
),
"minimum_topic_message_counts": minimum_counts,
"observed_topic_rates_hz": observed_rates,
"zero_message_topics": empty_topics,
"below_minimum_rate_topics": below_rate,
}
episode.optional_topic_groups = observations
episode.data_quality_warnings = warnings
def _validate_minimum_topic_rates(
self,
episode: _Episode,
@@ -1036,6 +1232,31 @@ class DataRecorderManager:
) -> tuple[Path | None, str]:
if not episode.active_directory.exists():
return None, "active episode directory is missing"
if not self.config.retain_failed_episodes:
try:
active_root = Path(self.config.base_directory) / "active"
target = episode.active_directory
if (
target.parent != active_root
or target.name != episode.request.episode_id
or target.is_symlink()
or active_root.is_symlink()
):
raise RecordingError(
"refusing to discard a failed episode outside its "
"fixed active root"
)
shutil.rmtree(target)
try:
self._sync_directory(active_root)
except OSError:
pass
return None, ""
except Exception as discard_error:
return None, (
"failed to discard project-owned failed episode: "
f"{type(discard_error).__name__}: {discard_error}"
)
try:
ready_marker = episode.active_directory / "READY"
if ready_marker.exists():
@@ -1063,6 +1284,8 @@ class DataRecorderManager:
episode.minimum_topic_message_counts
),
"observed_topic_rates_hz": episode.observed_topic_rates_hz,
"optional_topic_groups": episode.optional_topic_groups,
"data_quality_warnings": episode.data_quality_warnings,
"metadata_duration_nanoseconds": (
episode.metadata_duration_nanoseconds
),
@@ -1242,6 +1465,7 @@ class DataRecorderManager:
__all__ = [
"DataRecorderManager",
"OptionalTopicGroupConfig",
"RecorderConfig",
"RecordingError",
"RecordingToggleGate",