diff --git a/src/us_equity_snapshot_pipelines/lifecycle/soxl_adjusted_last_acquisition.py b/src/us_equity_snapshot_pipelines/lifecycle/soxl_adjusted_last_acquisition.py index 7407857..e5969e2 100644 --- a/src/us_equity_snapshot_pipelines/lifecycle/soxl_adjusted_last_acquisition.py +++ b/src/us_equity_snapshot_pipelines/lifecycle/soxl_adjusted_last_acquisition.py @@ -3,10 +3,13 @@ from __future__ import annotations from collections.abc import Callable, Sequence +from dataclasses import dataclass from datetime import date, datetime +import hashlib import json import os from pathlib import Path +import re from typing import Any from quant_platform_kit.ibkr import ( @@ -44,12 +47,87 @@ "duplicate_sessions_sha256", } ) +_REQUEST_VALIDATION_321_CAUSES = frozenset( + { + "request_id_mismatch", + "invalid_end_datetime", + "invalid_duration", + "invalid_what_to_show", + "unknown_321", + } +) +_REQUEST_VALIDATION_321_COUNTS = ( + "matching_error_count", + "mismatching_error_count", + "matching_completion_count", + "mismatching_completion_count", + "expected_session_count", + "observed_session_count", +) +_PROVIDER_MESSAGE_CAUSE_PATTERNS = ( + ( + re.compile( + r"(? dict[str, Any]: + return { + "cause": self.cause, + "numeric_code": self.numeric_code, + "request_envelope_sha256": self.request_envelope_sha256, + "provider_message_sha256": self.provider_message_sha256, + "request_completion_observed": self.request_completion_observed, + "counts": dict(self.counts), + } + + def acquire_strict_adjusted_last( ib: Any, symbol: str, @@ -80,6 +158,113 @@ def _is_sha256(value: object) -> bool: ) +def _domain_separated_sha256(domain: bytes, value: bytes) -> str: + return hashlib.sha256(domain + b"\x00" + value).hexdigest() + + +def _nonnegative_count(value: object) -> bool: + return not isinstance(value, bool) and isinstance(value, int) and value >= 0 + + +def classify_request_validation_321( + *, + active_request_id: int, + error_request_id: int, + error_code: int, + provider_message: str | None, + request_envelope: bytes, + completion_request_ids: Sequence[int], + expected_session_count: int, + observed_session_count: int, +) -> RequestValidation321Diagnostic: + """Correlate one synthetic 321 event and retain only closed aggregates.""" + if ( + isinstance(active_request_id, bool) + or not isinstance(active_request_id, int) + or active_request_id < 0 + or isinstance(error_request_id, bool) + or not isinstance(error_request_id, int) + or isinstance(error_code, bool) + or error_code != 321 + or (provider_message is not None and not isinstance(provider_message, str)) + or not isinstance(request_envelope, bytes) + or not request_envelope + or not _nonnegative_count(expected_session_count) + or not _nonnegative_count(observed_session_count) + ): + raise SoxlAdjustedLastDiagnosticError( + "invalid request-validation 321 diagnostic input" + ) + try: + completion_ids = tuple(completion_request_ids) + except TypeError: + raise SoxlAdjustedLastDiagnosticError( + "invalid request-validation 321 diagnostic input" + ) from None + if any( + isinstance(request_id, bool) or not isinstance(request_id, int) + for request_id in completion_ids + ): + raise SoxlAdjustedLastDiagnosticError( + "invalid request-validation 321 diagnostic input" + ) + + error_matches = error_request_id == active_request_id + matching_completion_count = sum( + request_id == active_request_id for request_id in completion_ids + ) + if error_matches: + normalized_message = ( + " ".join(provider_message.split()).casefold() + if provider_message is not None + else "" + ) + matched_causes = { + candidate + for pattern, candidate in _PROVIDER_MESSAGE_CAUSE_PATTERNS + if pattern.search(normalized_message) is not None + } + mentioned_causes = { + candidate + for pattern, candidate in _PROVIDER_MESSAGE_FIELD_PATTERNS + if pattern.search(normalized_message) is not None + } + cause = ( + matched_causes.pop() + if len(matched_causes) == 1 and matched_causes == mentioned_causes + else "unknown_321" + ) + else: + cause = "request_id_mismatch" + + message_bytes = ( + provider_message.encode("utf-8") if provider_message is not None else b"" + ) + counts = { + "matching_error_count": int(error_matches), + "mismatching_error_count": int(not error_matches), + "matching_completion_count": matching_completion_count, + "mismatching_completion_count": len(completion_ids) + - matching_completion_count, + "expected_session_count": expected_session_count, + "observed_session_count": observed_session_count, + } + return RequestValidation321Diagnostic( + cause=cause, + numeric_code=error_code, + request_envelope_sha256=_domain_separated_sha256( + _REQUEST_ENVELOPE_COMMITMENT_DOMAIN, + request_envelope, + ), + provider_message_sha256=_domain_separated_sha256( + _PROVIDER_MESSAGE_COMMITMENT_DOMAIN, + message_bytes, + ), + request_completion_observed=matching_completion_count > 0, + counts=tuple((key, counts[key]) for key in _REQUEST_VALIDATION_321_COUNTS), + ) + + def _sanitized_payload(error: StrictAdjustedHistoryError) -> dict[str, Any]: diagnostic = error.diagnostic if diagnostic is None: @@ -143,14 +328,58 @@ def _sanitized_payload(error: StrictAdjustedHistoryError) -> dict[str, Any]: return payload -def write_sanitized_adjusted_last_diagnostic( +def _request_validation_321_payload( + diagnostic: RequestValidation321Diagnostic, +) -> dict[str, Any]: + if not isinstance(diagnostic, RequestValidation321Diagnostic): + raise SoxlAdjustedLastDiagnosticError( + "invalid request-validation 321 diagnostic" + ) + payload = diagnostic.to_dict() + if ( + set(payload) + != { + "cause", + "numeric_code", + "request_envelope_sha256", + "provider_message_sha256", + "request_completion_observed", + "counts", + } + or payload["cause"] not in _REQUEST_VALIDATION_321_CAUSES + or payload["numeric_code"] != 321 + or not _is_sha256(payload["request_envelope_sha256"]) + or not _is_sha256(payload["provider_message_sha256"]) + or not isinstance(payload["request_completion_observed"], bool) + ): + raise SoxlAdjustedLastDiagnosticError( + "invalid request-validation 321 diagnostic" + ) + counts = payload["counts"] + if ( + not isinstance(counts, dict) + or tuple(counts) != _REQUEST_VALIDATION_321_COUNTS + or any(not _nonnegative_count(value) for value in counts.values()) + or counts["matching_error_count"] + counts["mismatching_error_count"] + != 1 + or (payload["cause"] == "request_id_mismatch") + != (counts["mismatching_error_count"] == 1) + or payload["request_completion_observed"] + != (counts["matching_completion_count"] > 0) + ): + raise SoxlAdjustedLastDiagnosticError( + "invalid request-validation 321 diagnostic" + ) + return payload + + +def _write_exclusive_mode_0600_json( destination: str | Path, - error: StrictAdjustedHistoryError, + payload: dict[str, Any], ) -> None: - """Create one exclusive mode-0600 JSON diagnostic without raw market data.""" path = Path(destination) - payload = json.dumps( - _sanitized_payload(error), + serialized = json.dumps( + payload, sort_keys=True, separators=(",", ":"), ensure_ascii=True, @@ -163,7 +392,7 @@ def write_sanitized_adjusted_last_diagnostic( try: os.fchmod(descriptor, 0o600) with os.fdopen(descriptor, "wb", closefd=False) as handle: - handle.write(payload) + handle.write(serialized) handle.flush() os.fsync(handle.fileno()) except BaseException: @@ -171,3 +400,22 @@ def write_sanitized_adjusted_last_diagnostic( raise finally: os.close(descriptor) + + +def write_sanitized_adjusted_last_diagnostic( + destination: str | Path, + error: StrictAdjustedHistoryError, +) -> None: + """Create one exclusive mode-0600 JSON diagnostic without raw market data.""" + _write_exclusive_mode_0600_json(destination, _sanitized_payload(error)) + + +def write_sanitized_request_validation_321_diagnostic( + destination: str | Path, + diagnostic: RequestValidation321Diagnostic, +) -> None: + """Create one exclusive mode-0600 request-bound 321 diagnostic.""" + _write_exclusive_mode_0600_json( + destination, + _request_validation_321_payload(diagnostic), + ) diff --git a/tests/test_soxl_adjusted_last_acquisition.py b/tests/test_soxl_adjusted_last_acquisition.py index 95c8350..111891f 100644 --- a/tests/test_soxl_adjusted_last_acquisition.py +++ b/tests/test_soxl_adjusted_last_acquisition.py @@ -1,6 +1,8 @@ from __future__ import annotations +from dataclasses import replace from datetime import date, datetime, timezone +import hashlib import json from pathlib import Path import stat @@ -14,12 +16,23 @@ ) from us_equity_snapshot_pipelines.lifecycle.soxl_adjusted_last_acquisition import ( acquire_strict_adjusted_last, + classify_request_validation_321, write_sanitized_adjusted_last_diagnostic, + write_sanitized_request_validation_321_diagnostic, ) EXPECTED = (date(2026, 8, 1), date(2026, 8, 2)) CUTOFF = datetime(2026, 8, 5, 3, 59, 59, tzinfo=timezone.utc) +REQUEST_ENVELOPE = ( + b'{"barSizeSetting":"1 day","durationStr":"9 Y",' + b'"endDateTime":"20260805 03:59:59 UTC","formatDate":1,' + b'"keepUpToDate":false,"useRTH":true,"whatToShow":"ADJUSTED_LAST"}' +) + + +def _commitment(domain: bytes, value: bytes) -> str: + return hashlib.sha256(domain + b"\x00" + value).hexdigest() def _bar(session: date, close: float = 98_765.4321) -> SimpleNamespace: @@ -156,3 +169,268 @@ def test_exact_match_preserves_strict_request_and_returns_no_failure_artifact( assert tuple(candle.session for candle in result.candles) == EXPECTED assert list(tmp_path.iterdir()) == [] assert ib.history_calls == 0 + + +@pytest.mark.parametrize( + ("provider_message", "expected_cause"), + [ + ("invalid endDateTime", "invalid_end_datetime"), + ( + "Error validating request.-'bS' : cause - invalid end date/time", + "invalid_end_datetime", + ), + ("invalid duration", "invalid_duration"), + ( + "Error validating request.-'bS' : cause - invalid duration", + "invalid_duration", + ), + ("invalid whatToShow", "invalid_what_to_show"), + ( + "Error validating request.-'bS' : cause - invalid what to show", + "invalid_what_to_show", + ), + ("invalid endDateTime and duration", "unknown_321"), + ( + "Error validating request: invalid duration; invalid endDateTime", + "unknown_321", + ), + (None, "unknown_321"), + ], +) +def test_matching_request_id_uses_closed_321_cause_allowlist( + provider_message: str | None, + expected_cause: str, +) -> None: + diagnostic = classify_request_validation_321( + active_request_id=41, + error_request_id=41, + error_code=321, + provider_message=provider_message, + request_envelope=REQUEST_ENVELOPE, + completion_request_ids=(41,), + expected_session_count=2_264, + observed_session_count=0, + ) + + assert diagnostic.to_dict()["cause"] == expected_cause + assert diagnostic.to_dict()["request_completion_observed"] is True + + +def test_mismatching_ids_are_counted_without_persisting_raw_ids() -> None: + diagnostic = classify_request_validation_321( + active_request_id=41, + error_request_id=99, + error_code=321, + provider_message="invalid duration", + request_envelope=REQUEST_ENVELOPE, + completion_request_ids=(99, 41, 100), + expected_session_count=2_264, + observed_session_count=0, + ) + + payload = diagnostic.to_dict() + assert payload["cause"] == "request_id_mismatch" + assert payload["request_completion_observed"] is True + assert payload["counts"] == { + "matching_error_count": 0, + "mismatching_error_count": 1, + "matching_completion_count": 1, + "mismatching_completion_count": 2, + "expected_session_count": 2_264, + "observed_session_count": 0, + } + assert { + "active_request_id", + "error_request_id", + "completion_request_ids", + }.isdisjoint(payload) + + +def test_321_error_precedes_completion_and_session_state() -> None: + diagnostic = classify_request_validation_321( + active_request_id=41, + error_request_id=41, + error_code=321, + provider_message="invalid duration", + request_envelope=REQUEST_ENVELOPE, + completion_request_ids=(99,), + expected_session_count=2_264, + observed_session_count=0, + ) + + payload = diagnostic.to_dict() + assert payload["cause"] == "invalid_duration" + assert payload["request_completion_observed"] is False + assert payload["counts"]["mismatching_completion_count"] == 1 + assert payload["counts"]["observed_session_count"] == 0 + + +def test_request_validation_writer_is_exclusive_mode_0600_and_sanitized( + tmp_path: Path, +) -> None: + provider_message = "invalid duration" + diagnostic = classify_request_validation_321( + active_request_id=41, + error_request_id=41, + error_code=321, + provider_message=provider_message, + request_envelope=REQUEST_ENVELOPE, + completion_request_ids=(), + expected_session_count=2_264, + observed_session_count=0, + ) + destination = tmp_path / "request-validation-321.json" + + write_sanitized_request_validation_321_diagnostic(destination, diagnostic) + payload = json.loads(destination.read_bytes()) + + assert set(payload) == { + "cause", + "numeric_code", + "request_envelope_sha256", + "provider_message_sha256", + "request_completion_observed", + "counts", + } + assert payload["cause"] == "invalid_duration" + assert payload["numeric_code"] == 321 + assert payload["request_envelope_sha256"] == _commitment( + b"qsl.soxl.request-envelope.v1", + REQUEST_ENVELOPE, + ) + assert payload["provider_message_sha256"] == _commitment( + b"qsl.soxl.provider-message.v1", + provider_message.encode("utf-8"), + ) + assert stat.S_IMODE(destination.stat().st_mode) == 0o600 + + serialized = destination.read_text(encoding="utf-8") + for forbidden in ( + provider_message, + "20260805", + "ADJUSTED_LAST", + "SOXL", + "SMART", + "USD", + "clientId", + "account", + "open", + "close", + "volume", + ): + assert forbidden not in serialized + + with pytest.raises(FileExistsError): + write_sanitized_request_validation_321_diagnostic(destination, diagnostic) + assert json.loads(destination.read_bytes()) == payload + + +def test_request_validation_exception_and_logs_do_not_expose_sensitive_inputs( + caplog: pytest.LogCaptureFixture, +) -> None: + sensitive_message = "private provider detail token=never-persist" + sensitive_envelope = b'{"clientId":48291,"symbol":"SOXL"}' + + with pytest.raises(ValueError) as caught: + classify_request_validation_321( + active_request_id=41, + error_request_id=41, + error_code=10089, + provider_message=sensitive_message, + request_envelope=sensitive_envelope, + completion_request_ids=(), + expected_session_count=0, + observed_session_count=0, + ) + + exposed = repr(caught.value) + caplog.text + assert sensitive_message not in exposed + assert sensitive_envelope.decode("utf-8") not in exposed + assert "41" not in exposed + + +@pytest.mark.parametrize( + "contradictory_fields", + [ + { + "counts": ( + ("matching_error_count", 0), + ("mismatching_error_count", 1), + ("matching_completion_count", 0), + ("mismatching_completion_count", 0), + ("expected_session_count", 2_264), + ("observed_session_count", 0), + ) + }, + {"request_completion_observed": True}, + ], +) +def test_writer_rejects_contradictory_request_correlation_aggregates( + tmp_path: Path, + contradictory_fields: dict[str, object], +) -> None: + diagnostic = classify_request_validation_321( + active_request_id=41, + error_request_id=41, + error_code=321, + provider_message="invalid duration", + request_envelope=REQUEST_ENVELOPE, + completion_request_ids=(), + expected_session_count=2_264, + observed_session_count=0, + ) + + with pytest.raises(ValueError): + write_sanitized_request_validation_321_diagnostic( + tmp_path / "contradictory.json", + replace(diagnostic, **contradictory_fields), + ) + assert list(tmp_path.iterdir()) == [] + + +@pytest.mark.parametrize( + "provider_end_time", + ["20260805 03:59:59 UTC", "20260805-03:59:59"], +) +def test_frozen_request_fields_and_aware_or_legacy_end_time_representations( + provider_end_time: str, +) -> None: + captured: dict[str, object] = {} + + def requester(contract, **kwargs): + captured["contract"] = contract + captured.update(kwargs) + return StrictAdjustedHistoryRequestOutcome( + bars=[_bar(EXPECTED[0]), _bar(EXPECTED[1])], + completion_observed=True, + ) + + acquire_strict_adjusted_last( + OfflineIB(), + "SOXL", + end_datetime=CUTOFF, + duration="9 Y", + expected_sessions=EXPECTED, + stock_factory=_stock, + requester=requester, + ) + + contract = captured.pop("contract") + assert vars(contract) == { + "symbol": "SOXL", + "exchange": "SMART", + "currency": "USD", + } + assert captured == { + "endDateTime": CUTOFF, + "durationStr": "9 Y", + "barSizeSetting": "1 day", + "whatToShow": "ADJUSTED_LAST", + "useRTH": True, + "formatDate": 1, + "keepUpToDate": False, + } + assert provider_end_time in { + CUTOFF.strftime("%Y%m%d %H:%M:%S UTC"), + CUTOFF.strftime("%Y%m%d-%H:%M:%S"), + }