From 62182a1ecfc02e5934b327d661b6814541379261 Mon Sep 17 00:00:00 2001 From: Ananya Date: Mon, 23 Mar 2026 23:14:49 +0530 Subject: [PATCH 1/2] feat(core): improve risk aggregator with bounded signals, weighting, and extensibility guards --- sdk/python/acf/sdk_integration/__init__.py | 5 + .../acf/sdk_integration/risk_context.py | 133 +++++++++++++ sdk/python/tests/test_risk_context.py | 177 ++++++++++++++++++ 3 files changed, 315 insertions(+) create mode 100644 sdk/python/acf/sdk_integration/__init__.py create mode 100644 sdk/python/acf/sdk_integration/risk_context.py create mode 100644 sdk/python/tests/test_risk_context.py diff --git a/sdk/python/acf/sdk_integration/__init__.py b/sdk/python/acf/sdk_integration/__init__.py new file mode 100644 index 0000000..2a47b72 --- /dev/null +++ b/sdk/python/acf/sdk_integration/__init__.py @@ -0,0 +1,5 @@ +"""SDK integration helpers for sidecar risk context contracts.""" + +from .risk_context import RiskContext, aggregate_risk + +__all__ = ["RiskContext", "aggregate_risk"] diff --git a/sdk/python/acf/sdk_integration/risk_context.py b/sdk/python/acf/sdk_integration/risk_context.py new file mode 100644 index 0000000..7630664 --- /dev/null +++ b/sdk/python/acf/sdk_integration/risk_context.py @@ -0,0 +1,133 @@ +"""Risk context contract between aggregate and policy stages. + +This module defines a fixed-size, O(1) aggregation surface for policy input. +""" +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + + +_SIGNAL_OBFUSCATION = "obfuscation" +_SIGNAL_LEXICAL = "lexical" +_SIGNAL_SEMANTIC = "semantic" +_SIGNAL_PROVENANCE = "provenance" + +ALLOWED_SIGNALS = { + _SIGNAL_OBFUSCATION, + _SIGNAL_LEXICAL, + _SIGNAL_SEMANTIC, + _SIGNAL_PROVENANCE, +} + +WEIGHTS = { + _SIGNAL_OBFUSCATION: 0.3, + _SIGNAL_LEXICAL: 0.3, + _SIGNAL_SEMANTIC: 0.2, + _SIGNAL_PROVENANCE: 0.2, +} + +# TODO(v2): evolve hook multipliers into policy-configured profiles. +HOOK_MULTIPLIERS = { + "on_prompt": 1.0, + "on_context": 1.0, + "on_tool_call": 1.1, + "on_memory": 1.0, +} + + +def _clamp01(value: float) -> float: + if value < 0.0: + return 0.0 + if value > 1.0: + return 1.0 + return value + + +def _normalize_signals(signals: dict[str, float]) -> dict[str, float]: + """Return a fixed-size signal map, ignoring non-allowed keys.""" + + # Explicit fixed-key extraction keeps aggregation O(1), even if callers + # pass additional keys. + return { + _SIGNAL_OBFUSCATION: _clamp01( + float(signals.get(_SIGNAL_OBFUSCATION, 0.0)) + ), + _SIGNAL_LEXICAL: _clamp01(float(signals.get(_SIGNAL_LEXICAL, 0.0))), + _SIGNAL_SEMANTIC: _clamp01(float(signals.get(_SIGNAL_SEMANTIC, 0.0))), + _SIGNAL_PROVENANCE: _clamp01( + float(signals.get(_SIGNAL_PROVENANCE, 0.0)) + ), + } + + +@dataclass(frozen=True) +class RiskContext: + """Normalized policy input object produced by the aggregator.""" + + score: float + signals: dict[str, float] + provenance: dict[str, Any] + metadata: dict[str, Any] + + def to_dict(self) -> dict[str, Any]: + return { + "score": self.score, + "signals": self.signals, + "provenance": self.provenance, + "metadata": self.metadata, + } + + +def aggregate_risk( + *, + signals: dict[str, float], + provenance: dict[str, Any], + metadata: dict[str, Any], +) -> RiskContext: + """Build a fixed-shape `RiskContext` with O(1) weighted scoring. + + Expected signals keys: obfuscation, lexical, semantic, provenance. + Missing keys default to 0.0. All values are clamped to [0.0, 1.0]. + """ + + normalized_signals = _normalize_signals(signals) + + normalized_provenance = { + "execution_id": str(provenance.get("execution_id", "")), + "trusted": bool(provenance.get("trusted", False)), + "nonce_valid": bool(provenance.get("nonce_valid", False)), + } + + normalized_metadata = { + "hook": str(metadata.get("hook", "")), + "timestamp": int(metadata.get("timestamp", 0)), + } + + score = _clamp01( + ( + WEIGHTS[_SIGNAL_OBFUSCATION] + * normalized_signals[_SIGNAL_OBFUSCATION] + ) + + (WEIGHTS[_SIGNAL_LEXICAL] * normalized_signals[_SIGNAL_LEXICAL]) + + (WEIGHTS[_SIGNAL_SEMANTIC] * normalized_signals[_SIGNAL_SEMANTIC]) + + ( + WEIGHTS[_SIGNAL_PROVENANCE] + * normalized_signals[_SIGNAL_PROVENANCE] + ) + ) + + trust_penalty = 0.2 if not normalized_provenance["trusted"] else 0.0 + nonce_penalty = 0.1 if not normalized_provenance["nonce_valid"] else 0.0 + score = _clamp01(score + trust_penalty + nonce_penalty) + + hook = normalized_metadata["hook"] + hook_multiplier = HOOK_MULTIPLIERS.get(hook, 1.0) + score = _clamp01(score * hook_multiplier) + + return RiskContext( + score=score, + signals=normalized_signals, + provenance=normalized_provenance, + metadata=normalized_metadata, + ) diff --git a/sdk/python/tests/test_risk_context.py b/sdk/python/tests/test_risk_context.py new file mode 100644 index 0000000..36be405 --- /dev/null +++ b/sdk/python/tests/test_risk_context.py @@ -0,0 +1,177 @@ +"""Tests for acf.sdk_integration.risk_context.""" + +from acf.sdk_integration.risk_context import ( + RiskContext, + WEIGHTS, + aggregate_risk, +) + + +def test_risk_context_structure(): + ctx = aggregate_risk( + signals={ + "obfuscation": 0.8, + "lexical": 0.7, + "semantic": 0.5, + "provenance": 1.0, + }, + provenance={ + "execution_id": "exec-123", + "trusted": True, + "nonce_valid": True, + }, + metadata={ + "hook": "on_prompt", + "timestamp": 1710000000, + }, + ) + + assert isinstance(ctx, RiskContext) + + payload = ctx.to_dict() + assert set(payload.keys()) == { + "score", + "signals", + "provenance", + "metadata", + } + assert set(payload["signals"].keys()) == { + "obfuscation", + "lexical", + "semantic", + "provenance", + } + assert set(payload["provenance"].keys()) == { + "execution_id", + "trusted", + "nonce_valid", + } + assert set(payload["metadata"].keys()) == {"hook", "timestamp"} + + opa_input = {"input": payload} + assert "input" in opa_input + assert opa_input["input"]["metadata"]["hook"] == "on_prompt" + + +def test_score_bounds(): + ctx = aggregate_risk( + signals={ + "obfuscation": 99.0, + "lexical": 50.0, + "semantic": 22.0, + "provenance": 9.0, + }, + provenance={ + "execution_id": "exec-bounds", + "trusted": False, + "nonce_valid": False, + }, + metadata={ + "hook": "on_context", + "timestamp": 1710000001, + }, + ) + + assert 0.0 <= ctx.score <= 1.0 + + +def test_weighted_scoring(): + ctx = aggregate_risk( + signals={ + "obfuscation": 0.5, + "lexical": 0.5, + "semantic": 0.5, + "provenance": 0.5, + }, + provenance={ + "execution_id": "exec-weighted", + "trusted": True, + "nonce_valid": True, + }, + metadata={ + "hook": "on_prompt", + "timestamp": 1710000100, + }, + ) + + expected = ( + (WEIGHTS["obfuscation"] * 0.5) + + (WEIGHTS["lexical"] * 0.5) + + (WEIGHTS["semantic"] * 0.5) + + (WEIGHTS["provenance"] * 0.5) + ) + assert round(ctx.score, 2) == round(expected, 2) + + +def test_deterministic_output(): + kwargs = { + "signals": { + "obfuscation": 0.21, + "lexical": 0.49, + "semantic": 0.13, + "provenance": 0.77, + }, + "provenance": { + "execution_id": "exec-deterministic", + "trusted": True, + "nonce_valid": True, + }, + "metadata": { + "hook": "on_tool_call", + "timestamp": 1710000002, + }, + } + + a = aggregate_risk(**kwargs) + b = aggregate_risk(**kwargs) + + assert a.to_dict() == b.to_dict() + + +def test_signal_normalization(): + ctx = aggregate_risk( + signals={ + "obfuscation": -10.0, + "lexical": 0.25, + "semantic": 100.0, + "provenance": -1.0, + }, + provenance={ + "execution_id": "exec-normalize", + "trusted": 1, + "nonce_valid": 0, + }, + metadata={ + "hook": "on_memory", + "timestamp": "1710000003", + }, + ) + + assert ctx.signals["obfuscation"] == 0.0 + assert ctx.signals["lexical"] == 0.25 + assert ctx.signals["semantic"] == 1.0 + assert ctx.signals["provenance"] == 0.0 + + assert ctx.provenance["trusted"] is True + assert ctx.provenance["nonce_valid"] is False + assert ctx.metadata["timestamp"] == 1710000003 + + +def test_missing_signals_defaults(): + ctx = aggregate_risk( + signals={}, + provenance={ + "execution_id": "exec-defaults", + "trusted": True, + "nonce_valid": True, + }, + metadata={ + "hook": "on_prompt", + "timestamp": 1710000200, + }, + ) + + assert ctx.signals["obfuscation"] == 0.0 + assert ctx.signals["lexical"] == 0.0 + assert ctx.signals["semantic"] == 0.0 + assert ctx.signals["provenance"] == 0.0 From af55b4a7d6a00d012c987617e99e19c988867ab7 Mon Sep 17 00:00:00 2001 From: Ananya Date: Mon, 23 Mar 2026 23:38:45 +0530 Subject: [PATCH 2/2] fixes --- .../acf/sdk_integration/risk_context.py | 81 ++++++++++-- sdk/python/tests/test_risk_context.py | 124 ++++++++++++++++++ 2 files changed, 193 insertions(+), 12 deletions(-) diff --git a/sdk/python/acf/sdk_integration/risk_context.py b/sdk/python/acf/sdk_integration/risk_context.py index 7630664..782307c 100644 --- a/sdk/python/acf/sdk_integration/risk_context.py +++ b/sdk/python/acf/sdk_integration/risk_context.py @@ -4,6 +4,7 @@ """ from __future__ import annotations +import math from dataclasses import dataclass from typing import Any @@ -37,6 +38,8 @@ def _clamp01(value: float) -> float: + if not math.isfinite(value): + return 0.0 if value < 0.0: return 0.0 if value > 1.0: @@ -44,19 +47,68 @@ def _clamp01(value: float) -> float: return value -def _normalize_signals(signals: dict[str, float]) -> dict[str, float]: +def _safe_float(value: Any, default: float = 0.0) -> float: + """Safely convert untyped input to finite float.""" + + try: + parsed = float(value) + except (TypeError, ValueError): + return default + if not math.isfinite(parsed): + return default + return parsed + + +def _safe_bool(value: Any, default: bool = False) -> bool: + """Parse bool-like values without truthiness pitfalls.""" + + if isinstance(value, bool): + return value + + if isinstance(value, (int, float)): + if value == 1: + return True + if value == 0: + return False + return default + + if isinstance(value, str): + normalized = value.strip().lower() + if normalized in {"true", "1", "yes", "y"}: + return True + if normalized in {"false", "0", "no", "n"}: + return False + return default + + return default + + +def _safe_int(value: Any, default: int = 0) -> int: + """Safely convert untyped input to int.""" + + try: + return int(value) + except (TypeError, ValueError): + return default + + +def _normalize_signals(signals: dict[str, Any]) -> dict[str, float]: """Return a fixed-size signal map, ignoring non-allowed keys.""" # Explicit fixed-key extraction keeps aggregation O(1), even if callers # pass additional keys. return { _SIGNAL_OBFUSCATION: _clamp01( - float(signals.get(_SIGNAL_OBFUSCATION, 0.0)) + _safe_float(signals.get(_SIGNAL_OBFUSCATION, 0.0)) + ), + _SIGNAL_LEXICAL: _clamp01( + _safe_float(signals.get(_SIGNAL_LEXICAL, 0.0)) + ), + _SIGNAL_SEMANTIC: _clamp01( + _safe_float(signals.get(_SIGNAL_SEMANTIC, 0.0)) ), - _SIGNAL_LEXICAL: _clamp01(float(signals.get(_SIGNAL_LEXICAL, 0.0))), - _SIGNAL_SEMANTIC: _clamp01(float(signals.get(_SIGNAL_SEMANTIC, 0.0))), _SIGNAL_PROVENANCE: _clamp01( - float(signals.get(_SIGNAL_PROVENANCE, 0.0)) + _safe_float(signals.get(_SIGNAL_PROVENANCE, 0.0)) ), } @@ -70,18 +122,23 @@ class RiskContext: provenance: dict[str, Any] metadata: dict[str, Any] + def __post_init__(self) -> None: + object.__setattr__(self, "signals", dict(self.signals)) + object.__setattr__(self, "provenance", dict(self.provenance)) + object.__setattr__(self, "metadata", dict(self.metadata)) + def to_dict(self) -> dict[str, Any]: return { "score": self.score, - "signals": self.signals, - "provenance": self.provenance, - "metadata": self.metadata, + "signals": dict(self.signals), + "provenance": dict(self.provenance), + "metadata": dict(self.metadata), } def aggregate_risk( *, - signals: dict[str, float], + signals: dict[str, Any], provenance: dict[str, Any], metadata: dict[str, Any], ) -> RiskContext: @@ -95,13 +152,13 @@ def aggregate_risk( normalized_provenance = { "execution_id": str(provenance.get("execution_id", "")), - "trusted": bool(provenance.get("trusted", False)), - "nonce_valid": bool(provenance.get("nonce_valid", False)), + "trusted": _safe_bool(provenance.get("trusted", False)), + "nonce_valid": _safe_bool(provenance.get("nonce_valid", False)), } normalized_metadata = { "hook": str(metadata.get("hook", "")), - "timestamp": int(metadata.get("timestamp", 0)), + "timestamp": _safe_int(metadata.get("timestamp", 0)), } score = _clamp01( diff --git a/sdk/python/tests/test_risk_context.py b/sdk/python/tests/test_risk_context.py index 36be405..ae83221 100644 --- a/sdk/python/tests/test_risk_context.py +++ b/sdk/python/tests/test_risk_context.py @@ -1,5 +1,7 @@ """Tests for acf.sdk_integration.risk_context.""" +import math + from acf.sdk_integration.risk_context import ( RiskContext, WEIGHTS, @@ -175,3 +177,125 @@ def test_missing_signals_defaults(): assert ctx.signals["lexical"] == 0.0 assert ctx.signals["semantic"] == 0.0 assert ctx.signals["provenance"] == 0.0 + + +def test_nan_signal_normalization(): + ctx = aggregate_risk( + signals={ + "obfuscation": math.nan, + "lexical": 0.25, + "semantic": math.inf, + "provenance": -math.inf, + }, + provenance={ + "execution_id": "exec-nan", + "trusted": True, + "nonce_valid": True, + }, + metadata={ + "hook": "on_prompt", + "timestamp": 1710000300, + }, + ) + + assert ctx.signals["obfuscation"] == 0.0 + assert ctx.signals["semantic"] == 0.0 + assert ctx.signals["provenance"] == 0.0 + assert 0.0 <= ctx.score <= 1.0 + + +def test_non_numeric_signal_values_default_to_zero(): + ctx = aggregate_risk( + signals={ + "obfuscation": "high", + "lexical": None, + "semantic": "0.4", + "provenance": {}, + }, + provenance={ + "execution_id": "exec-bad-signal", + "trusted": True, + "nonce_valid": True, + }, + metadata={ + "hook": "on_prompt", + "timestamp": 1710000400, + }, + ) + + assert ctx.signals["obfuscation"] == 0.0 + assert ctx.signals["lexical"] == 0.0 + assert ctx.signals["semantic"] == 0.4 + assert ctx.signals["provenance"] == 0.0 + + +def test_provenance_bool_parsing_from_strings(): + ctx = aggregate_risk( + signals={ + "obfuscation": 0.0, + "lexical": 0.0, + "semantic": 0.0, + "provenance": 0.0, + }, + provenance={ + "execution_id": "exec-bool-str", + "trusted": "false", + "nonce_valid": "true", + }, + metadata={ + "hook": "on_prompt", + "timestamp": 1710000500, + }, + ) + + assert ctx.provenance["trusted"] is False + assert ctx.provenance["nonce_valid"] is True + + +def test_invalid_metadata_timestamp_defaults_zero(): + ctx = aggregate_risk( + signals={ + "obfuscation": 0.0, + "lexical": 0.0, + "semantic": 0.0, + "provenance": 0.0, + }, + provenance={ + "execution_id": "exec-ts", + "trusted": True, + "nonce_valid": True, + }, + metadata={ + "hook": "on_prompt", + "timestamp": "not-a-number", + }, + ) + + assert ctx.metadata["timestamp"] == 0 + + +def test_to_dict_returns_defensive_copies(): + ctx = aggregate_risk( + signals={ + "obfuscation": 0.1, + "lexical": 0.2, + "semantic": 0.3, + "provenance": 0.4, + }, + provenance={ + "execution_id": "exec-copy", + "trusted": True, + "nonce_valid": True, + }, + metadata={ + "hook": "on_prompt", + "timestamp": 1710000600, + }, + ) + + payload = ctx.to_dict() + payload["signals"]["lexical"] = 0.99 + payload["metadata"]["hook"] = "on_tool_call" + + assert ctx.signals["lexical"] == 0.2 + assert ctx.metadata["hook"] == "on_prompt"