Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 18 additions & 5 deletions engine/src/agent_control_engine/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
EvaluationRequest,
EvaluationResponse,
EvaluatorResult,
JSONObject,
)

from .selectors import select_data
Expand Down Expand Up @@ -295,6 +296,7 @@ async def _evaluate_leaf(
node: ConditionNode,
request: EvaluationRequest,
semaphore: asyncio.Semaphore,
extensions: JSONObject | None,
) -> _ConditionEvaluation:
"""Evaluate a leaf selector/evaluator pair.

Expand All @@ -320,7 +322,7 @@ async def _evaluate_leaf(
timeout = DEFAULT_EVALUATOR_TIMEOUT

result = await asyncio.wait_for(
evaluator.evaluate_with_context(data, request.step),
evaluator.evaluate_with_extensions(data, request.step, extensions),
timeout=timeout,
)
except TimeoutError:
Expand Down Expand Up @@ -440,17 +442,20 @@ async def _evaluate_condition(
node: ConditionNode,
request: EvaluationRequest,
semaphore: asyncio.Semaphore,
extensions: JSONObject | None,
) -> _ConditionEvaluation:
"""Evaluate a recursive condition tree."""
if node.is_leaf():
return await self._evaluate_leaf(item, node, request, semaphore)
return await self._evaluate_leaf(item, node, request, semaphore, extensions)

kind = node.kind()
children = node.children_in_order()
child_evaluations: list[_ConditionEvaluation] = []

if kind == "not":
child_eval = await self._evaluate_condition(item, children[0], request, semaphore)
child_eval = await self._evaluate_condition(
item, children[0], request, semaphore, extensions
)
trace = {
"type": "not",
"evaluated": True,
Expand Down Expand Up @@ -478,7 +483,9 @@ async def _evaluate_condition(
return _ConditionEvaluation(result=result, trace=trace)

for index, child in enumerate(children):
child_eval = await self._evaluate_condition(item, child, request, semaphore)
child_eval = await self._evaluate_condition(
item, child, request, semaphore, extensions
)
child_evaluations.append(child_eval)

if child_eval.result.error:
Expand Down Expand Up @@ -624,14 +631,19 @@ def get_applicable_controls(

return applicable

async def process(self, request: EvaluationRequest) -> EvaluationResponse:
async def process(
self,
request: EvaluationRequest,
extensions: JSONObject | None = None,
) -> EvaluationResponse:
"""Process controls in parallel with cancel-on-deny.

All applicable controls are evaluated concurrently. If any control
matches with action=deny, remaining evaluations are cancelled.

Args:
request: The evaluation request containing step and context
extensions: Opaque authenticated metadata supplied by the caller.

Returns:
EvaluationResponse with is_safe status and any matches
Expand Down Expand Up @@ -667,6 +679,7 @@ async def evaluate_control(eval_task: _EvalTask) -> None:
eval_task.item.control.condition,
request,
semaphore,
extensions,
)
eval_task.result = evaluation.result

Expand Down
54 changes: 53 additions & 1 deletion engine/tests/test_core.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
EvaluationRequest,
EvaluatorResult,
EvaluatorSpec,
JSONObject,
SteeringContext,
Step,
)
Expand All @@ -40,14 +41,16 @@ class SimpleConfig(BaseModel):
_execution_log: list[str] = []
_blocker_event: asyncio.Event | None = None
_context_calls: list[tuple[Any, Step]] = []
_extension_calls: list[JSONObject | None] = []


def reset_test_state() -> None:
"""Reset shared test state."""
global _execution_log, _blocker_event, _context_calls
global _execution_log, _blocker_event, _context_calls, _extension_calls
_execution_log = []
_blocker_event = asyncio.Event()
_context_calls = []
_extension_calls = []


class AllowEvaluator(Evaluator[SimpleConfig]):
Expand Down Expand Up @@ -182,6 +185,12 @@ async def evaluate_with_context(self, data: Any, step: Step) -> EvaluatorResult:
_context_calls.append((data, step))
return EvaluatorResult(matched=False, confidence=1.0, message="context received")

async def evaluate_with_extensions(
self, data: Any, step: Step, extensions: JSONObject | None
) -> EvaluatorResult:
_extension_calls.append(extensions)
return await self.evaluate_with_context(data, step)


@dataclass
class MockControlWithIdentity:
Expand Down Expand Up @@ -303,6 +312,49 @@ async def test_context_evaluator_receives_selected_data_and_complete_step() -> N
assert _context_calls == [("answer", step)]


@pytest.mark.asyncio
async def test_engine_passes_opaque_extensions_to_evaluator() -> None:
# Given: an authenticated extension mapping and a context-aware evaluator
engine = ControlEngine(
[make_control(1, "context", "test-context", action="observe", path="input")]
)
extensions: JSONObject = {
"namespace_key": "namespace-1",
"target_type": "custom_target",
"target_id": "target-9",
"metadata": {"provider_field": "opaque_value"},
}

# When: evaluating a request with the trusted extensions
await engine.process(
EvaluationRequest(
agent_name="00000000-0000-0000-0000-000000000001",
step=Step(type="llm", name="test-step", input="question"),
stage="pre",
),
extensions=extensions,
)

# Then: the engine passes the mapping unchanged without interpreting it
assert _extension_calls == [extensions]


@pytest.mark.asyncio
async def test_default_extension_hook_preserves_existing_evaluator_behavior() -> None:
# Given: an evaluator that implements only the original evaluate method
evaluator = AllowEvaluator(SimpleConfig())
step = Step(type="llm", name="test-step", input="question")

# When: the new hook is called with opaque extensions
result = await evaluator.evaluate_with_extensions(
"question", step, {"provider_field": "opaque_value"}
)

# Then: the existing evaluation method still handles the request
assert result.message == "Allowed"
assert _execution_log == ["allow:default:start", "allow:default:end"]


@pytest.mark.asyncio
async def test_cached_context_evaluator_handles_concurrent_steps_without_retaining_state() -> None:
# Given: one cached evaluator configuration and two independent requests
Expand Down
15 changes: 14 additions & 1 deletion evaluators/builtin/src/agent_control_evaluators/_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, ClassVar, Generic, TypeVar

from agent_control_models import EvaluatorResult, Step
from agent_control_models import EvaluatorResult, JSONObject, Step
from agent_control_models.base import BaseModel

if TYPE_CHECKING:
Expand Down Expand Up @@ -178,6 +178,19 @@ async def evaluate_with_context(self, data: Any, step: Step) -> EvaluatorResult:
"""
return await self.evaluate(data)

async def evaluate_with_extensions(
self,
data: Any,
step: Step,
extensions: JSONObject | None,
) -> EvaluatorResult:
"""Evaluate with opaque trusted request extensions.

Existing evaluators retain their behavior because the default
implementation delegates to :meth:`evaluate_with_context`.
"""
return await self.evaluate_with_context(data, step)

def get_timeout_seconds(self) -> float:
"""Get timeout in seconds from config or metadata default."""
timeout_ms: int = getattr(self.config, "timeout_ms", self.metadata.timeout_ms)
Expand Down
13 changes: 13 additions & 0 deletions evaluators/builtin/tests/test_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,6 +121,19 @@ async def test_contextual_evaluation_delegates_to_existing_evaluate(self):
assert result.matched is True
assert result.metadata == {"data": "selected data"}

@pytest.mark.asyncio
async def test_extension_evaluation_delegates_to_existing_context_hook(self):
"""The extension hook preserves behavior for legacy evaluators."""
evaluator = MockEvaluator.from_dict({"should_match": True})
step = Step(type="llm", name="answer", input="full input")

result = await evaluator.evaluate_with_extensions(
"selected data", step, {"provider_field": "opaque_value"}
)

assert result.matched is True
assert result.metadata == {"data": "selected data"}

def test_evaluator_config_stored(self):
"""Test that evaluator stores config."""
evaluator = MockEvaluator.from_dict({"should_match": True})
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -228,6 +228,15 @@ class ScorerInvokeRecord(BaseModel):
dataset_output: JSONValue = None


class GalileoExecutionContext(BaseModel):
"""Authenticated identity fields required by Galileo scorer invocation."""

organization_id: str = Field(min_length=1)
user_id: str = Field(min_length=1)
project_id: str = Field(min_length=1)
run_id: str = Field(min_length=1)


class ScorerInvokeRequest(BaseModel):
"""Request payload for Luna scorer invocation.

Expand All @@ -246,6 +255,7 @@ class ScorerInvokeRequest(BaseModel):
scorer_label: str | None = Field(default=None, min_length=1)
inputs: ScorerInvokeInputs
record: ScorerInvokeRecord | None = None
execution_context: GalileoExecutionContext | None = None
config: ScorerInvokeConfig = Field(default_factory=ScorerInvokeConfig)

@model_validator(mode="after")
Expand Down Expand Up @@ -474,6 +484,7 @@ async def invoke(
input: JSONValue = None,
output: JSONValue = None,
step: Step | None = None,
execution_context: GalileoExecutionContext | None = None,
config: ScorerInvokeConfig | JSONObject | None = None,
timeout: float = DEFAULT_TIMEOUT_SECS,
headers: dict[str, str] | None = None,
Expand All @@ -488,6 +499,7 @@ async def invoke(
input: Optional user/system prompt text.
output: Optional model response text.
step: Optional complete runtime step used for structured dual-write.
execution_context: Optional authenticated Galileo scorer context.
config: Optional Orbit-supported scorer invocation configuration.
timeout: Request timeout in seconds.
headers: Additional request headers.
Expand Down Expand Up @@ -531,6 +543,7 @@ async def invoke(
selected_input=input,
selected_output=output,
),
execution_context=execution_context,
config=invoke_config,
).to_dict()

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,9 +10,9 @@

import httpx
from agent_control_evaluators import Evaluator, EvaluatorMetadata, register_evaluator
from agent_control_models import EvaluatorResult, JSONValue, Step
from agent_control_models import EvaluatorResult, JSONObject, JSONValue, Step

from .client import GalileoLunaClient, ScorerInvokeResponse
from .client import GalileoExecutionContext, GalileoLunaClient, ScorerInvokeResponse
from .config import LunaEvaluatorConfig, coerce_number

logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -212,7 +212,72 @@ async def evaluate_with_context(self, data: Any, step: Step) -> EvaluatorResult:
"""
return await self._evaluate(data, step=step)

async def _evaluate(self, data: Any, *, step: Step | None) -> EvaluatorResult:
async def evaluate_with_extensions(
self,
data: Any,
step: Step,
extensions: JSONObject | None,
) -> EvaluatorResult:
"""Evaluate using opaque authenticated metadata supplied by Agent Control."""
return await self._evaluate(
data,
step=step,
extensions=extensions,
)

@staticmethod
def _execution_context_from_extensions(
extensions: JSONObject | None,
) -> GalileoExecutionContext | None:
"""Translate trusted opaque auth metadata into Galileo's request context."""
if extensions is None:
return None

metadata = extensions.get("metadata")
metadata_obj = metadata if isinstance(metadata, dict) else {}
organization_id = extensions.get("namespace_key")
# Orbit's runtime-auth contract defines caller_id as the authenticated
# user ID for this flow. Keep that contract mapping inside Galileo.
user_id = extensions.get("caller_id")
project_id = metadata_obj.get("project_id")
target_type = extensions.get("target_type")
target_id = extensions.get("target_id")
run_id = target_id if target_type == "log_stream" else None

missing = [
field
for field, value in (
("organization_id", organization_id),
("user_id", user_id),
("project_id", project_id),
("run_id", run_id),
)
if not isinstance(value, str) or not value
]
if missing:
raise ValueError(
"Authenticated execution metadata is missing required fields: "
+ ", ".join(missing)
)

assert isinstance(organization_id, str)
assert isinstance(user_id, str)
assert isinstance(project_id, str)
assert isinstance(run_id, str)
return GalileoExecutionContext(
organization_id=organization_id,
user_id=user_id,
project_id=project_id,
run_id=run_id,
)

async def _evaluate(
self,
data: Any,
*,
step: Step | None,
extensions: JSONObject | None = None,
) -> EvaluatorResult:
"""Run a Luna evaluation with optional structured runtime context."""
input_text, output_text = self._prepare_payload(data)
if not (_has_text(input_text) or _has_text(output_text)):
Expand All @@ -224,9 +289,12 @@ async def _evaluate(self, data: Any, *, step: Step | None) -> EvaluatorResult:
)

try:
execution_context = self._execution_context_from_extensions(extensions)
scorer_kwargs = self._scorer_kwargs()
if step is not None:
scorer_kwargs["step"] = step
if execution_context is not None:
scorer_kwargs["execution_context"] = execution_context
response = await self._get_client().invoke(
**scorer_kwargs,
input=input_text if _has_text(input_text) else None,
Expand Down
Loading
Loading