diff --git a/roughpy/src/streams/tick_stream.cpp b/roughpy/src/streams/tick_stream.cpp index 9fefae296..827c9e401 100644 --- a/roughpy/src/streams/tick_stream.cpp +++ b/roughpy/src/streams/tick_stream.cpp @@ -172,41 +172,41 @@ static py::object construct(const py::object& data, py::kwargs kwargs) lie_elt[key] += python::py_to_scalar(pmd.scalar_type, tick.data); break; - case streams::ChannelType::Value: + case streams::ChannelType::Value: { + const auto tick_scalar + = python::py_to_scalar(pmd.scalar_type, tick.data); if (channel->is_lead_lag()) { auto lag_key = key + 1; const auto& prev_lead = previous_values[idx]; const auto& prev_lag = previous_values[idx + 1]; - auto lead = pmd.ctx->zero_lie(meta.cached_vector_type); - auto lag = pmd.ctx->zero_lie(meta.cached_vector_type); + auto next_lead = pmd.ctx->zero_lie(meta.cached_vector_type); + auto next_lag = pmd.ctx->zero_lie(meta.cached_vector_type); - lead[key] - += python::py_to_scalar(pmd.scalar_type, tick.data); - lead[lag_key] - += python::py_to_scalar(pmd.scalar_type, tick.data); + next_lead[key] += tick_scalar; + next_lag[lag_key] += tick_scalar; lie_elt = pmd.ctx->cbh( - lead.sub(prev_lead), lag.sub(prev_lag), + next_lead.sub(prev_lead), next_lag.sub(prev_lag), meta.cached_vector_type ); - previous_values[idx] = std::move(lead); - previous_values[idx + 1] = std::move(lag); + previous_values[idx] = std::move(next_lead); + previous_values[idx + 1] = std::move(next_lag); break; } else { const auto& prev_val = previous_values[idx]; auto new_val = pmd.ctx->zero_lie(meta.cached_vector_type); - new_val[key] - += python::py_to_scalar(pmd.scalar_type, tick.data); + new_val[key] += tick_scalar; lie_elt = new_val.sub(prev_val); previous_values[idx] = std::move(new_val); break; } + } case streams::ChannelType::Categorical: { lie_elt[key] += scalars::Scalar(1); break; diff --git a/tests/streams/test_tick_stream.py b/tests/streams/test_tick_stream.py index 4e094aa10..9737d916f 100644 --- a/tests/streams/test_tick_stream.py +++ b/tests/streams/test_tick_stream.py @@ -30,6 +30,7 @@ import pytest +import roughpy as rp from roughpy import DPReal, Lie, RealInterval, TickStream DATA_FORMATS = [ @@ -135,3 +136,58 @@ def test_zero_data_at_zero_timestamp(): sig = stream.signature(RealInterval(0.0, 2.0)) assert sig is not None + + +def test_tick_stream_lead_lag(): + # Data from the bug report + data = [ + (0.0, "s1", "value", 1.0), + (1.0, "s1", "value", 2.0), + (2.0, "s1", "value", 3.0), + (3.0, "s1", "value", 2.0), + ] + + # Manual implementation (known correct) + ll_data = [] + for i in range(len(data)): + t_curr, _, _, v_curr = data[i] + + if i == 0: + ll_data.append((t_curr, "lead", "value", v_curr)) + ll_data.append((t_curr, "lag", "value", v_curr)) + else: + t_prev, _, _, v_prev = data[i - 1] + + ll_data.append((t_curr, "lead", "value", v_curr)) + ll_data.append((t_curr, "lag", "value", v_prev)) + + ll_data.append((t_curr, "lead", "value", v_curr)) + ll_data.append((t_curr, "lag", "value", v_curr)) + + ctx = rp.get_context(2, 2, rp.DPReal) + schema_manual = [("lead", "value", {}), ("lag", "value", {})] + + s_manual = rp.TickStream.from_data( + ll_data, + schema=schema_manual, + ctx=ctx, + support=rp.RealInterval(0, 4), + ) + + sig_manual = s_manual.signature( + rp.RealInterval(1.0, 3.0001), ctx=ctx, resolution=10 + ) + + # Native implementation + s_native = rp.TickStream.from_data( + data, + schema=[("s1", "value", {"lead_lag": True})], + ctx=ctx, + support=rp.RealInterval(0, 4), + ) + sig_native = s_native.signature( + rp.RealInterval(1.0, 3.0001), ctx=ctx, resolution=10 + ) + + # They should match + assert str(sig_native) == str(sig_manual)