mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
feat(rust): pyo3 bridge for SmartCrusher
Stage 3c.1b step 1: expose `SmartCrusherConfig`, `CrushResult`, and `SmartCrusher` to Python via `headroom._core`. The Python shim that delegates to it (replacing the 3669-line Python implementation) lands in the next commit; this commit just builds the bridge and a fixture-replay test that pins it. Surface: - `headroom._core.SmartCrusherConfig(**fields)` — every field of the Rust `SmartCrusherConfig` exposed as a kwarg with matching default. - `headroom._core.CrushResult` — read-only mirror of the Rust struct with `compressed`, `original`, `was_modified`, `strategy` getters. - `headroom._core.SmartCrusher(config=None)` — constructor accepts only `config`; the Python shim drops `relevance_config`, `scorer`, and `ccr_config` since Stage 3c.1 keeps those subsystems disabled. - `crush(content, query="", bias=1.0)` and `smart_crush_content(...)` methods mirror the Python signatures. Verification: - All 17 recorded parity fixtures byte-equal between Python and the PyO3 bridge (`tests/test_transforms/test_smart_crusher_rust_parity.py`, 18 tests pass — 1 fixture-count sanity + 17 fixtures). - The Rust-side `cargo run -p headroom-parity --bin parity-run -- run --only smart_crusher` was already 17/17 green. The two tests catch different regression classes: - Rust-only test: catches drift in the Rust port's logic. - Python bridge test: catches PyO3 input/output translation bugs.
This commit is contained in:
parent
43d1aa0329
commit
5328d87b1e
2 changed files with 336 additions and 0 deletions
|
|
@ -15,6 +15,10 @@
|
|||
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use headroom_core::transforms::smart_crusher::{
|
||||
CrushResult as RustCrushResult, SmartCrusher as RustSmartCrusher,
|
||||
SmartCrusherConfig as RustSmartCrusherConfig,
|
||||
};
|
||||
use headroom_core::transforms::{
|
||||
DiffCompressionResult, DiffCompressor, DiffCompressorConfig, DiffCompressorStats,
|
||||
};
|
||||
|
|
@ -367,6 +371,248 @@ impl PyDiffCompressor {
|
|||
}
|
||||
}
|
||||
|
||||
// ─── SmartCrusherConfig ────────────────────────────────────────────────────
|
||||
|
||||
/// Mirror of `headroom.transforms.smart_crusher.SmartCrusherConfig`.
|
||||
/// Defaults match Python's dataclass byte-for-byte. The constructor
|
||||
/// accepts every field as a kwarg with the same name and type so the
|
||||
/// Python shim can pass `SmartCrusherConfig(**asdict(py_cfg))`.
|
||||
#[pyclass(name = "SmartCrusherConfig", module = "headroom._core")]
|
||||
#[derive(Clone)]
|
||||
struct PySmartCrusherConfig {
|
||||
inner: RustSmartCrusherConfig,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl PySmartCrusherConfig {
|
||||
#[new]
|
||||
#[pyo3(signature = (
|
||||
enabled = true,
|
||||
min_items_to_analyze = 5,
|
||||
min_tokens_to_crush = 200,
|
||||
variance_threshold = 2.0,
|
||||
uniqueness_threshold = 0.1,
|
||||
similarity_threshold = 0.8,
|
||||
max_items_after_crush = 15,
|
||||
preserve_change_points = true,
|
||||
factor_out_constants = false,
|
||||
include_summaries = false,
|
||||
use_feedback_hints = true,
|
||||
toin_confidence_threshold = 0.5,
|
||||
dedup_identical_items = true,
|
||||
first_fraction = 0.3,
|
||||
last_fraction = 0.15,
|
||||
relevance_threshold = 0.3,
|
||||
))]
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn new(
|
||||
enabled: bool,
|
||||
min_items_to_analyze: usize,
|
||||
min_tokens_to_crush: usize,
|
||||
variance_threshold: f64,
|
||||
uniqueness_threshold: f64,
|
||||
similarity_threshold: f64,
|
||||
max_items_after_crush: usize,
|
||||
preserve_change_points: bool,
|
||||
factor_out_constants: bool,
|
||||
include_summaries: bool,
|
||||
use_feedback_hints: bool,
|
||||
toin_confidence_threshold: f64,
|
||||
dedup_identical_items: bool,
|
||||
first_fraction: f64,
|
||||
last_fraction: f64,
|
||||
relevance_threshold: f64,
|
||||
) -> Self {
|
||||
Self {
|
||||
inner: RustSmartCrusherConfig {
|
||||
enabled,
|
||||
min_items_to_analyze,
|
||||
min_tokens_to_crush,
|
||||
variance_threshold,
|
||||
uniqueness_threshold,
|
||||
similarity_threshold,
|
||||
max_items_after_crush,
|
||||
preserve_change_points,
|
||||
factor_out_constants,
|
||||
include_summaries,
|
||||
use_feedback_hints,
|
||||
toin_confidence_threshold,
|
||||
dedup_identical_items,
|
||||
first_fraction,
|
||||
last_fraction,
|
||||
relevance_threshold,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[getter]
|
||||
fn enabled(&self) -> bool {
|
||||
self.inner.enabled
|
||||
}
|
||||
#[getter]
|
||||
fn min_items_to_analyze(&self) -> usize {
|
||||
self.inner.min_items_to_analyze
|
||||
}
|
||||
#[getter]
|
||||
fn min_tokens_to_crush(&self) -> usize {
|
||||
self.inner.min_tokens_to_crush
|
||||
}
|
||||
#[getter]
|
||||
fn variance_threshold(&self) -> f64 {
|
||||
self.inner.variance_threshold
|
||||
}
|
||||
#[getter]
|
||||
fn uniqueness_threshold(&self) -> f64 {
|
||||
self.inner.uniqueness_threshold
|
||||
}
|
||||
#[getter]
|
||||
fn similarity_threshold(&self) -> f64 {
|
||||
self.inner.similarity_threshold
|
||||
}
|
||||
#[getter]
|
||||
fn max_items_after_crush(&self) -> usize {
|
||||
self.inner.max_items_after_crush
|
||||
}
|
||||
#[getter]
|
||||
fn preserve_change_points(&self) -> bool {
|
||||
self.inner.preserve_change_points
|
||||
}
|
||||
#[getter]
|
||||
fn factor_out_constants(&self) -> bool {
|
||||
self.inner.factor_out_constants
|
||||
}
|
||||
#[getter]
|
||||
fn include_summaries(&self) -> bool {
|
||||
self.inner.include_summaries
|
||||
}
|
||||
#[getter]
|
||||
fn use_feedback_hints(&self) -> bool {
|
||||
self.inner.use_feedback_hints
|
||||
}
|
||||
#[getter]
|
||||
fn toin_confidence_threshold(&self) -> f64 {
|
||||
self.inner.toin_confidence_threshold
|
||||
}
|
||||
#[getter]
|
||||
fn dedup_identical_items(&self) -> bool {
|
||||
self.inner.dedup_identical_items
|
||||
}
|
||||
#[getter]
|
||||
fn first_fraction(&self) -> f64 {
|
||||
self.inner.first_fraction
|
||||
}
|
||||
#[getter]
|
||||
fn last_fraction(&self) -> f64 {
|
||||
self.inner.last_fraction
|
||||
}
|
||||
#[getter]
|
||||
fn relevance_threshold(&self) -> f64 {
|
||||
self.inner.relevance_threshold
|
||||
}
|
||||
|
||||
fn __repr__(&self) -> String {
|
||||
format!(
|
||||
"SmartCrusherConfig(enabled={}, min_items_to_analyze={}, \
|
||||
min_tokens_to_crush={}, max_items_after_crush={}, \
|
||||
relevance_threshold={})",
|
||||
self.inner.enabled,
|
||||
self.inner.min_items_to_analyze,
|
||||
self.inner.min_tokens_to_crush,
|
||||
self.inner.max_items_after_crush,
|
||||
self.inner.relevance_threshold,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
// ─── CrushResult ───────────────────────────────────────────────────────────
|
||||
|
||||
/// Mirror of `headroom.transforms.smart_crusher.CrushResult`. Read-only;
|
||||
/// the Python shim builds its own dataclass instance from these
|
||||
/// attributes so callers that destructure with `asdict()` keep working.
|
||||
#[pyclass(name = "CrushResult", module = "headroom._core")]
|
||||
struct PyCrushResult {
|
||||
inner: RustCrushResult,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl PyCrushResult {
|
||||
#[getter]
|
||||
fn compressed(&self) -> &str {
|
||||
&self.inner.compressed
|
||||
}
|
||||
#[getter]
|
||||
fn original(&self) -> &str {
|
||||
&self.inner.original
|
||||
}
|
||||
#[getter]
|
||||
fn was_modified(&self) -> bool {
|
||||
self.inner.was_modified
|
||||
}
|
||||
#[getter]
|
||||
fn strategy(&self) -> &str {
|
||||
&self.inner.strategy
|
||||
}
|
||||
|
||||
fn __repr__(&self) -> String {
|
||||
format!(
|
||||
"CrushResult(compressed=<{} chars>, was_modified={}, strategy={:?})",
|
||||
self.inner.compressed.len(),
|
||||
self.inner.was_modified,
|
||||
self.inner.strategy,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
// ─── SmartCrusher ──────────────────────────────────────────────────────────
|
||||
|
||||
/// Mirror of `headroom.transforms.smart_crusher.SmartCrusher`.
|
||||
///
|
||||
/// Constructor accepts only `config` — Python's `relevance_config`,
|
||||
/// `scorer`, and `ccr_config` parameters are handled in the Python
|
||||
/// shim (Stage 3c.1 keeps the optional subsystems disabled in Rust;
|
||||
/// the shim drops those args to preserve call-site compatibility).
|
||||
#[pyclass(name = "SmartCrusher", module = "headroom._core")]
|
||||
struct PySmartCrusher {
|
||||
inner: RustSmartCrusher,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl PySmartCrusher {
|
||||
#[new]
|
||||
#[pyo3(signature = (config = None))]
|
||||
fn new(config: Option<&PySmartCrusherConfig>) -> Self {
|
||||
let cfg = config
|
||||
.map(|c| c.inner.clone())
|
||||
.unwrap_or_default();
|
||||
Self {
|
||||
inner: RustSmartCrusher::new(cfg),
|
||||
}
|
||||
}
|
||||
|
||||
/// `crush(content, query="", bias=1.0) -> CrushResult`. Argument
|
||||
/// order and keyword names mirror the Python implementation.
|
||||
#[pyo3(signature = (content, query = "", bias = 1.0))]
|
||||
fn crush(&self, content: &str, query: &str, bias: f64) -> PyCrushResult {
|
||||
PyCrushResult {
|
||||
inner: self.inner.crush(content, query, bias),
|
||||
}
|
||||
}
|
||||
|
||||
/// `smart_crush_content(content, query="", bias=1.0) -> (str, bool, str)`.
|
||||
/// Mirrors Python's `_smart_crush_content` — used by
|
||||
/// `smart_crush_tool_output` convenience function and direct
|
||||
/// callers that want the tuple form.
|
||||
#[pyo3(signature = (content, query = "", bias = 1.0))]
|
||||
fn smart_crush_content(
|
||||
&self,
|
||||
content: &str,
|
||||
query: &str,
|
||||
bias: f64,
|
||||
) -> (String, bool, String) {
|
||||
self.inner.smart_crush_content(content, query, bias)
|
||||
}
|
||||
}
|
||||
|
||||
// ─── Module init ───────────────────────────────────────────────────────────
|
||||
|
||||
#[pymodule]
|
||||
|
|
@ -376,5 +622,8 @@ fn _core(m: &Bound<'_, PyModule>) -> PyResult<()> {
|
|||
m.add_class::<PyDiffCompressionResult>()?;
|
||||
m.add_class::<PyDiffCompressorStats>()?;
|
||||
m.add_class::<PyDiffCompressor>()?;
|
||||
m.add_class::<PySmartCrusherConfig>()?;
|
||||
m.add_class::<PyCrushResult>()?;
|
||||
m.add_class::<PySmartCrusher>()?;
|
||||
Ok(())
|
||||
}
|
||||
|
|
|
|||
87
tests/test_transforms/test_smart_crusher_rust_parity.py
Normal file
87
tests/test_transforms/test_smart_crusher_rust_parity.py
Normal file
|
|
@ -0,0 +1,87 @@
|
|||
"""Parity test: PyO3-backed `SmartCrusher` vs recorded fixtures.
|
||||
|
||||
Stage 3c.1b verification — guards the PyO3 bridge against regressions
|
||||
by replaying every recorded fixture in
|
||||
`tests/parity/fixtures/smart_crusher/` through `headroom._core.SmartCrusher`
|
||||
and asserting the output matches the recording byte-for-byte.
|
||||
|
||||
Twin of `test_diff_compressor_rust_parity.py`. The Rust side runs the
|
||||
same fixtures via `cargo run -p headroom-parity --bin parity-run --
|
||||
run --only smart_crusher`; this Python test specifically catches PyO3
|
||||
bridge regressions (input/output mistranslation) that the Rust-only
|
||||
binary cannot.
|
||||
|
||||
Skipped automatically when the `headroom._core` wheel isn't installed
|
||||
(e.g. CI lane without the maturin step).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _has_core() -> bool:
|
||||
try:
|
||||
from headroom._core import SmartCrusher # noqa: F401
|
||||
|
||||
return True
|
||||
except ImportError:
|
||||
return False
|
||||
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not _has_core(),
|
||||
reason="headroom._core wheel not installed (run `scripts/build_rust_extension.sh`)",
|
||||
)
|
||||
|
||||
|
||||
_FIXTURES_DIR = Path(__file__).parent.parent / "parity" / "fixtures" / "smart_crusher"
|
||||
|
||||
|
||||
def _all_fixtures() -> list[Path]:
|
||||
return sorted(_FIXTURES_DIR.glob("*.json"))
|
||||
|
||||
|
||||
def test_at_least_17_fixtures_present():
|
||||
"""Sanity check: the recorded fixture suite landed."""
|
||||
fixtures = _all_fixtures()
|
||||
assert len(fixtures) >= 17, (
|
||||
f"expected >= 17 fixtures, found {len(fixtures)}. "
|
||||
"If you re-recorded and got fewer, something deleted them."
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("fixture_path", _all_fixtures(), ids=lambda p: p.name)
|
||||
def test_rust_backend_matches_recorded_output(fixture_path: Path):
|
||||
"""Replay each recorded input through the PyO3 bridge; every output
|
||||
field must match the recording. Any mismatch is a bridge bug or a
|
||||
Rust regression — cross-check with `cargo run -p headroom-parity`.
|
||||
"""
|
||||
from headroom._core import SmartCrusher, SmartCrusherConfig
|
||||
|
||||
fixture = json.loads(fixture_path.read_text())
|
||||
inp = fixture["input"]
|
||||
cfg_dict = fixture["config"]
|
||||
expected = fixture["output"]
|
||||
|
||||
cfg = SmartCrusherConfig(**cfg_dict)
|
||||
crusher = SmartCrusher(cfg)
|
||||
actual = crusher.crush(inp["content"], inp["query"], inp["bias"])
|
||||
|
||||
assert actual.compressed == expected["compressed"], (
|
||||
f"compressed bytes differ for {fixture_path.name}\n"
|
||||
f" expected: {expected['compressed'][:120]!r}\n"
|
||||
f" actual : {actual.compressed[:120]!r}"
|
||||
)
|
||||
assert actual.original == expected["original"], f"original bytes differ for {fixture_path.name}"
|
||||
assert actual.was_modified == expected["was_modified"], (
|
||||
f"was_modified differs for {fixture_path.name}: "
|
||||
f"expected={expected['was_modified']} actual={actual.was_modified}"
|
||||
)
|
||||
assert actual.strategy == expected["strategy"], (
|
||||
f"strategy differs for {fixture_path.name}: "
|
||||
f"expected={expected['strategy']!r} actual={actual.strategy!r}"
|
||||
)
|
||||
Loading…
Add table
Add a link
Reference in a new issue