From f59ce6732aad0a325eb21fdea8c2ffb6864170f4 Mon Sep 17 00:00:00 2001 From: Namrata Ghadi Date: Wed, 30 Sep 2026 12:55:20 -0700 Subject: [PATCH 1/7] add extensions/metadata for additional fields that we want to pass as context to execution of evaluator --- engine/src/agent_control_engine/core.py | 23 ++++-- engine/tests/test_core.py | 38 ++++++++- .../src/agent_control_evaluators/_base.py | 15 +++- .../luna/client.py | 13 ++++ .../luna/evaluator.py | 74 +++++++++++++++++- .../galileo/tests/test_luna_evaluator.py | 78 +++++++++++++++++++ .../auth_framework/core.py | 4 + .../auth_framework/providers/http_upstream.py | 21 ++++- .../auth_framework/providers/local_jwt.py | 1 + .../auth_framework/runtime_token.py | 23 ++++++ .../agent_control_server/endpoints/auth.py | 1 + .../endpoints/evaluation.py | 19 ++++- server/tests/test_auth_framework.py | 20 +++-- .../test_runtime_token_exchange_endpoint.py | 75 +++++++++++++++++- 14 files changed, 383 insertions(+), 22 deletions(-) diff --git a/engine/src/agent_control_engine/core.py b/engine/src/agent_control_engine/core.py index 7dbba3c72..ab170e822 100644 --- a/engine/src/agent_control_engine/core.py +++ b/engine/src/agent_control_engine/core.py @@ -21,6 +21,7 @@ EvaluationRequest, EvaluationResponse, EvaluatorResult, + JSONObject, ) from .selectors import select_data @@ -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. @@ -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: @@ -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, @@ -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: @@ -624,7 +631,11 @@ 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 @@ -632,6 +643,7 @@ async def process(self, request: EvaluationRequest) -> EvaluationResponse: 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 @@ -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 diff --git a/engine/tests/test_core.py b/engine/tests/test_core.py index e16d437f6..fe15f5ed6 100644 --- a/engine/tests/test_core.py +++ b/engine/tests/test_core.py @@ -20,6 +20,7 @@ EvaluationRequest, EvaluatorResult, EvaluatorSpec, + JSONObject, SteeringContext, Step, ) @@ -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]): @@ -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: @@ -303,6 +312,33 @@ 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_cached_context_evaluator_handles_concurrent_steps_without_retaining_state() -> None: # Given: one cached evaluator configuration and two independent requests diff --git a/evaluators/builtin/src/agent_control_evaluators/_base.py b/evaluators/builtin/src/agent_control_evaluators/_base.py index d83c0c3ab..100ebd80b 100644 --- a/evaluators/builtin/src/agent_control_evaluators/_base.py +++ b/evaluators/builtin/src/agent_control_evaluators/_base.py @@ -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: @@ -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) diff --git a/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/luna/client.py b/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/luna/client.py index d00d14571..cca27c665 100644 --- a/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/luna/client.py +++ b/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/luna/client.py @@ -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. @@ -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") @@ -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, @@ -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. @@ -531,6 +543,7 @@ async def invoke( selected_input=input, selected_output=output, ), + execution_context=execution_context, config=invoke_config, ).to_dict() diff --git a/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/luna/evaluator.py b/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/luna/evaluator.py index e4147a9ad..99085e26e 100644 --- a/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/luna/evaluator.py +++ b/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/luna/evaluator.py @@ -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__) @@ -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)): @@ -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, diff --git a/evaluators/contrib/galileo/tests/test_luna_evaluator.py b/evaluators/contrib/galileo/tests/test_luna_evaluator.py index 4334ac5e8..b368f24da 100644 --- a/evaluators/contrib/galileo/tests/test_luna_evaluator.py +++ b/evaluators/contrib/galileo/tests/test_luna_evaluator.py @@ -627,6 +627,7 @@ def test_client_reports_invalid_connection_tuning_env( @pytest.mark.asyncio async def test_client_posts_to_luna_invoke_scorer_invoke(self) -> None: from agent_control_evaluator_galileo.luna import GalileoLunaClient + from agent_control_evaluator_galileo.luna.client import GalileoExecutionContext captured: dict[str, object] = {} @@ -654,6 +655,12 @@ def handler(request: httpx.Request) -> httpx.Response: scorer_id=SCORER_ID, input="user prompt", output="model answer", + execution_context=GalileoExecutionContext( + organization_id="org-1", + user_id="user-2", + project_id="project-3", + run_id="run-4", + ), config={"request_timeout_seconds": 7}, headers={"Galileo-API-Key": "blocked", "X-Request-ID": "safe-id"}, ) @@ -665,6 +672,12 @@ def handler(request: httpx.Request) -> httpx.Response: assert captured["url"] == "http://luna-invoke:8090/api/v1/scorers/invoke" expected_body = { "scorer_id": SCORER_ID, + "execution_context": { + "organization_id": "org-1", + "user_id": "user-2", + "project_id": "project-3", + "run_id": "run-4", + }, "inputs": {"query": "user prompt", "response": "model answer"}, "config": {"request_timeout_seconds": 7.0}, } @@ -1049,6 +1062,71 @@ async def test_evaluator_contextual_hook_forwards_complete_step(self) -> None: timeout=10.0, ) + @patch.dict(os.environ, LUNA_ENV) + @pytest.mark.asyncio + async def test_evaluator_consumes_verified_opaque_extensions(self) -> None: + from agent_control_evaluator_galileo.luna import LunaEvaluator, ScorerInvokeResponse + from agent_control_evaluator_galileo.luna.client import ( + GalileoExecutionContext, + GalileoLunaClient, + ) + + evaluator = LunaEvaluator.from_dict( + {"scorer_id": "scorer-123", "threshold": 0.5, "operator": "gte"} + ) + extensions = { + "namespace_key": "org-1", + "caller_id": "verified-user-2", + "target_type": "log_stream", + "target_id": "run-4", + "metadata": {"project_id": "project-3"}, + } + + with patch.object(GalileoLunaClient, "invoke", new_callable=AsyncMock) as mock_invoke: + mock_invoke.return_value = ScorerInvokeResponse(score=0.8, status="success") + result = await evaluator.evaluate_with_extensions( + "selected input", Step(type="llm", name="answer", input="prompt"), extensions + ) + + assert result.matched is True + mock_invoke.assert_awaited_once_with( + scorer_id="scorer-123", + step=Step(type="llm", name="answer", input="prompt"), + execution_context=GalileoExecutionContext( + organization_id="org-1", + user_id="verified-user-2", + project_id="project-3", + run_id="run-4", + ), + input="selected input", + output=None, + config=None, + timeout=10.0, + ) + + @patch.dict(os.environ, LUNA_ENV) + @pytest.mark.asyncio + async def test_evaluator_requires_authenticated_caller_id(self) -> None: + from agent_control_evaluator_galileo.luna import LunaEvaluator + from agent_control_evaluator_galileo.luna.client import GalileoLunaClient + + evaluator = LunaEvaluator.from_dict({"scorer_id": "scorer-123"}) + extensions = { + "namespace_key": "org-1", + "target_type": "log_stream", + "target_id": "run-4", + "metadata": {"project_id": "project-3"}, + } + + with patch.object(GalileoLunaClient, "invoke", new_callable=AsyncMock) as mock_invoke: + result = await evaluator.evaluate_with_extensions( + "selected input", Step(type="llm", name="answer", input="prompt"), extensions + ) + + assert result.error is not None + assert "user_id" in result.error + mock_invoke.assert_not_called() + @patch.dict(os.environ, LUNA_ENV) @pytest.mark.asyncio async def test_evaluator_labels_forwarded_scorer_version_id_as_requested(self) -> None: diff --git a/server/src/agent_control_server/auth_framework/core.py b/server/src/agent_control_server/auth_framework/core.py index 011c62de2..fde0f524b 100644 --- a/server/src/agent_control_server/auth_framework/core.py +++ b/server/src/agent_control_server/auth_framework/core.py @@ -28,6 +28,7 @@ from enum import StrEnum from typing import Any, Protocol +from agent_control_models import JSONObject from fastapi import Request @@ -83,6 +84,8 @@ class Principal: grant_expires_at: When the upstream grant expires. Used by the runtime-token exchange endpoint to bound the local token's lifetime. + extensions: Opaque trusted metadata returned by the authorizer. + Generic authorization code preserves this without interpreting it. """ namespace_key: str @@ -92,6 +95,7 @@ class Principal: target_id: str | None = None scopes: tuple[str, ...] = () grant_expires_at: datetime | None = None + extensions: JSONObject | None = None ContextBuilder = Callable[[Request], dict[str, Any] | Awaitable[dict[str, Any]]] diff --git a/server/src/agent_control_server/auth_framework/providers/http_upstream.py b/server/src/agent_control_server/auth_framework/providers/http_upstream.py index b61705b98..451d9ae07 100644 --- a/server/src/agent_control_server/auth_framework/providers/http_upstream.py +++ b/server/src/agent_control_server/auth_framework/providers/http_upstream.py @@ -48,6 +48,7 @@ from typing import Any import httpx +from agent_control_models import JSONObject from agent_control_models.errors import ErrorCode, ErrorReason from fastapi import Request from prometheus_client import Counter, Histogram @@ -55,6 +56,7 @@ BaseModel, ConfigDict, Field, + TypeAdapter, ValidationError, field_validator, model_validator, @@ -78,6 +80,7 @@ "Duration of auth upstream HTTP attempts made by Agent Control.", ("operation", "outcome"), ) +_JSON_OBJECT_ADAPTER: TypeAdapter[JSONObject] = TypeAdapter(JSONObject) class _UpstreamGrant(BaseModel): @@ -89,7 +92,7 @@ class _UpstreamGrant(BaseModel): with a 502. """ - model_config = ConfigDict(extra="ignore", strict=True) + model_config = ConfigDict(extra="allow", strict=True) namespace_key: str = Field(min_length=1) is_admin: bool = False @@ -408,6 +411,21 @@ def _parse_principal(self, response: httpx.Response) -> Principal: hint="Contact the operator.", ) from exc + try: + extensions = _JSON_OBJECT_ADAPTER.validate_python(grant.model_extra or {}) + except ValidationError as exc: + _logger.error( + "Auth upstream returned malformed extension metadata: %s", + exc.errors(), + ) + raise APIError( + status_code=502, + error_code=ErrorCode.AUTH_MISCONFIGURED, + reason=ErrorReason.INTERNAL_ERROR, + detail="Authorization service returned malformed extension metadata.", + hint="Contact the operator.", + ) from exc + return Principal( namespace_key=grant.namespace_key, is_admin=grant.is_admin, @@ -416,6 +434,7 @@ def _parse_principal(self, response: httpx.Response) -> Principal: target_id=grant.target_id, scopes=grant.scopes, grant_expires_at=grant.expires_at, + extensions=extensions or None, ) diff --git a/server/src/agent_control_server/auth_framework/providers/local_jwt.py b/server/src/agent_control_server/auth_framework/providers/local_jwt.py index 7cab77f46..182f013ff 100644 --- a/server/src/agent_control_server/auth_framework/providers/local_jwt.py +++ b/server/src/agent_control_server/auth_framework/providers/local_jwt.py @@ -118,6 +118,7 @@ async def authorize( target_id=claims.target_id, scopes=claims.scopes, grant_expires_at=claims.expires_at, + extensions=claims.extensions, ) def _extract_bearer_token(self, request: Request) -> str: diff --git a/server/src/agent_control_server/auth_framework/runtime_token.py b/server/src/agent_control_server/auth_framework/runtime_token.py index 54c59fbb2..e40e008e9 100644 --- a/server/src/agent_control_server/auth_framework/runtime_token.py +++ b/server/src/agent_control_server/auth_framework/runtime_token.py @@ -33,10 +33,13 @@ from typing import Any import jwt +from agent_control_models import JSONObject +from pydantic import TypeAdapter, ValidationError _ALGORITHM = "HS256" _ISSUER = "agent-control/server" _DOMAIN = "runtime" +_JSON_OBJECT_ADAPTER: TypeAdapter[JSONObject] = TypeAdapter(JSONObject) class RuntimeTokenError(Exception): @@ -65,6 +68,7 @@ class RuntimeTokenClaims: expires_at: datetime issued_at: datetime jti: str + extensions: JSONObject | None = None def mint_runtime_token( @@ -77,6 +81,7 @@ def mint_runtime_token( secret: str, ttl_seconds: int, upstream_expires_at: datetime | None = None, + extensions: JSONObject | None = None, now: datetime | None = None, ) -> tuple[str, RuntimeTokenClaims]: """Mint a runtime token. Returns ``(token, claims)``. @@ -116,6 +121,12 @@ def mint_runtime_token( "(e.g., tz=UTC); naive datetimes are not supported." ) + if extensions is not None: + try: + extensions = _JSON_OBJECT_ADAPTER.validate_python(extensions) + except ValidationError as exc: + raise RuntimeTokenError("Runtime token extensions must be a JSON object.") from exc + issued_at = now or datetime.now(UTC) if upstream_expires_at is not None and upstream_expires_at <= issued_at: # Minting with an already-expired ``exp`` would return a 200 with @@ -142,6 +153,8 @@ def mint_runtime_token( "exp": int(expires_at.timestamp()), "jti": jti, } + if extensions: + payload["extensions"] = extensions token = jwt.encode(payload, secret, algorithm=_ALGORITHM) claims = RuntimeTokenClaims( namespace_key=namespace_key, @@ -152,6 +165,7 @@ def mint_runtime_token( expires_at=expires_at, issued_at=issued_at, jti=jti, + extensions=extensions or None, ) return token, claims @@ -187,6 +201,7 @@ def verify_runtime_token(token: str, secret: str) -> RuntimeTokenClaims: actor_id = payload.get("actor_id") target_type = payload.get("target_type") target_id = payload.get("target_id") + raw_extensions = payload.get("extensions") if not isinstance(namespace_key, str) or not namespace_key: raise RuntimeTokenError("Runtime token missing namespace_key.") if not isinstance(actor_id, str) or not actor_id: @@ -196,6 +211,13 @@ def verify_runtime_token(token: str, secret: str) -> RuntimeTokenClaims: if not isinstance(target_id, str) or not target_id: raise RuntimeTokenError("Runtime token missing target_id.") + extensions: JSONObject | None = None + if raw_extensions is not None: + try: + extensions = _JSON_OBJECT_ADAPTER.validate_python(raw_extensions) + except ValidationError as exc: + raise RuntimeTokenError("Runtime token has malformed extensions.") from exc + raw_scopes = payload.get("scopes", []) if not isinstance(raw_scopes, list) or not all(isinstance(s, str) for s in raw_scopes): raise RuntimeTokenError("Runtime token has malformed scopes.") @@ -214,4 +236,5 @@ def verify_runtime_token(token: str, secret: str) -> RuntimeTokenClaims: expires_at=datetime.fromtimestamp(payload["exp"], tz=UTC), issued_at=datetime.fromtimestamp(payload["iat"], tz=UTC), jti=jti, + extensions=extensions, ) diff --git a/server/src/agent_control_server/endpoints/auth.py b/server/src/agent_control_server/endpoints/auth.py index 2372eb6f8..f2797bee2 100644 --- a/server/src/agent_control_server/endpoints/auth.py +++ b/server/src/agent_control_server/endpoints/auth.py @@ -211,6 +211,7 @@ async def runtime_token_exchange( secret=config.secret, ttl_seconds=config.ttl_seconds, upstream_expires_at=principal.grant_expires_at, + extensions=principal.extensions, ) except UpstreamGrantExpiredError as exc: # Upstream returned a grant whose ``expires_at`` is already in diff --git a/server/src/agent_control_server/endpoints/evaluation.py b/server/src/agent_control_server/endpoints/evaluation.py index a31d757dd..7a46829e0 100644 --- a/server/src/agent_control_server/endpoints/evaluation.py +++ b/server/src/agent_control_server/endpoints/evaluation.py @@ -9,6 +9,7 @@ ControlMatch, EvaluationRequest, EvaluationResponse, + JSONObject, ) from agent_control_models.errors import ErrorCode, ValidationErrorItem from fastapi import APIRouter, Depends, Request @@ -117,6 +118,22 @@ def _sanitize_evaluation_response(response: EvaluationResponse) -> EvaluationRes ) +def _principal_extensions(principal: Principal) -> JSONObject | None: + """Build the opaque evaluator extension from the authenticated principal.""" + if principal.target_type is None or principal.target_id is None: + return None + + extensions: JSONObject = { + "namespace_key": principal.namespace_key, + "target_type": principal.target_type, + "target_id": principal.target_id, + "metadata": principal.extensions or {}, + } + if principal.caller_id is not None: + extensions["caller_id"] = principal.caller_id + return extensions + + async def _evaluation_context(request: Request) -> dict[str, object]: """Surface target identifiers to the runtime authorizer.""" try: @@ -199,7 +216,7 @@ async def evaluate( engine_controls = await _load_engine_controls(request, principal) engine = ControlEngine(engine_controls) try: - raw_response = await engine.process(request) + raw_response = await engine.process(request, extensions=_principal_extensions(principal)) except ValueError: _logger.exception("Evaluation failed due to invalid configuration or input") raise APIValidationError( diff --git a/server/tests/test_auth_framework.py b/server/tests/test_auth_framework.py index 289a18899..8f0e32a04 100644 --- a/server/tests/test_auth_framework.py +++ b/server/tests/test_auth_framework.py @@ -1187,11 +1187,13 @@ async def test_http_upstream_accepts_iso_datetime_and_array_scopes(): lambda req: httpx.Response( 200, json={ - "namespace_key": "org-1", + "namespace_key": "namespace-1", "is_admin": False, "scopes": ["runtime.use", "runtime.read_only"], - "target_type": "log_stream", - "target_id": "ls-1", + "target_type": "custom_target", + "target_id": "target-1", + "provider_field": "opaque_value", + "nested": {"opaque": "metadata"}, "expires_at": iso_expiry, }, ) @@ -1199,12 +1201,16 @@ async def test_http_upstream_accepts_iso_datetime_and_array_scopes(): principal = await provider.authorize( _build_request(), Operation.RUNTIME_TOKEN_EXCHANGE, - context={"target_type": "log_stream", "target_id": "ls-1"}, + context={"target_type": "custom_target", "target_id": "target-1"}, ) - assert principal.namespace_key == "org-1" + assert principal.namespace_key == "namespace-1" assert principal.scopes == ("runtime.use", "runtime.read_only") - assert principal.target_type == "log_stream" - assert principal.target_id == "ls-1" + assert principal.target_type == "custom_target" + assert principal.target_id == "target-1" + assert principal.extensions == { + "provider_field": "opaque_value", + "nested": {"opaque": "metadata"}, + } assert principal.grant_expires_at is not None assert principal.grant_expires_at.isoformat() == iso_expiry diff --git a/server/tests/test_runtime_token_exchange_endpoint.py b/server/tests/test_runtime_token_exchange_endpoint.py index a4c5d4405..6a3a86451 100644 --- a/server/tests/test_runtime_token_exchange_endpoint.py +++ b/server/tests/test_runtime_token_exchange_endpoint.py @@ -14,8 +14,11 @@ from datetime import UTC, datetime, timedelta from unittest.mock import patch +import httpx +import jwt import pytest from fastapi.testclient import TestClient +from starlette.requests import Request from agent_control_server.auth_framework import Operation, Principal from agent_control_server.auth_framework.config import ( @@ -26,9 +29,8 @@ clear_authorizers, set_authorizer, ) -from agent_control_server.auth_framework.providers import ( - LocalJwtVerifyProvider, -) +from agent_control_server.auth_framework.providers import HttpUpstreamAuthProvider, LocalJwtVerifyProvider +from agent_control_server.auth_framework.providers.http_upstream import HttpUpstreamConfig from agent_control_server.auth_framework.runtime_token import RuntimeTokenError _TEST_SECRET = "test-runtime-secret-12345678901234567890" @@ -125,6 +127,73 @@ def test_exchange_endpoint_mints_token_when_configured( assert body["expires_at"] +@pytest.mark.asyncio +async def test_upstream_extensions_round_trip_through_signed_runtime_token( + runtime_config_enabled, +): + from agent_control_server.endpoints.auth import ( + RuntimeTokenExchangeRequest, + runtime_token_exchange, + ) + + trusted_metadata = {"provider_field": "opaque_value", "nested": {"opaque": "value"}} + + def authz_response(_request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + json={ + "namespace_key": "namespace-1", + "caller_id": "opaque-caller", + "target_type": "custom_target", + "target_id": "target-9", + "scopes": ["runtime.use"], + **trusted_metadata, + }, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(authz_response)) as client: + provider = HttpUpstreamAuthProvider( + HttpUpstreamConfig(url="https://authz.test/check"), client=client + ) + upstream_request = Request( + {"type": "http", "headers": [], "method": "POST", "path": "/"} + ) + principal = await provider.authorize( + upstream_request, + Operation.RUNTIME_TOKEN_EXCHANGE, + context={"target_type": "custom_target", "target_id": "target-9"}, + ) + + assert principal.extensions == trusted_metadata + response = await runtime_token_exchange( + RuntimeTokenExchangeRequest(target_type="custom_target", target_id="target-9"), + principal=principal, + ) + token_payload = jwt.decode( + response.token, + _TEST_SECRET, + algorithms=["HS256"], + issuer="agent-control/server", + ) + assert "provider_field" not in token_payload + assert token_payload["extensions"] == trusted_metadata + + verify_request = Request( + { + "type": "http", + "headers": [(b"authorization", f"Bearer {response.token}".encode())], + "method": "POST", + "path": "/", + } + ) + verified_principal = await LocalJwtVerifyProvider(secret=_TEST_SECRET).authorize( + verify_request, + Operation.RUNTIME_USE, + context={"target_type": "custom_target", "target_id": "target-9"}, + ) + assert verified_principal.extensions == trusted_metadata + + @pytest.mark.parametrize( ("actor_id", "expected_log_actor_id"), [ From 48e9edaa92b973993494328cf686cbab892c5d69 Mon Sep 17 00:00:00 2001 From: Namrata Ghadi Date: Wed, 30 Sep 2026 13:11:18 -0700 Subject: [PATCH 2/7] test: cover evaluator auth extension fallbacks --- engine/tests/test_core.py | 16 +++++++ evaluators/builtin/tests/test_base.py | 13 ++++++ server/tests/test_auth_framework.py | 67 +++++++++++++++++++++++++++ 3 files changed, 96 insertions(+) diff --git a/engine/tests/test_core.py b/engine/tests/test_core.py index fe15f5ed6..34c620d5d 100644 --- a/engine/tests/test_core.py +++ b/engine/tests/test_core.py @@ -339,6 +339,22 @@ async def test_engine_passes_opaque_extensions_to_evaluator() -> None: 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 diff --git a/evaluators/builtin/tests/test_base.py b/evaluators/builtin/tests/test_base.py index d81cc3f5b..5dfe6915e 100644 --- a/evaluators/builtin/tests/test_base.py +++ b/evaluators/builtin/tests/test_base.py @@ -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}) diff --git a/server/tests/test_auth_framework.py b/server/tests/test_auth_framework.py index 8f0e32a04..48645d6c8 100644 --- a/server/tests/test_auth_framework.py +++ b/server/tests/test_auth_framework.py @@ -696,6 +696,51 @@ def test_runtime_token_round_trips(): assert decoded.scopes == ("runtime.use",) +def test_runtime_token_rejects_non_json_extensions_when_minting(): + from agent_control_server.auth_framework.runtime_token import ( + RuntimeTokenError, + mint_runtime_token, + ) + + with pytest.raises(RuntimeTokenError, match="extensions must be a JSON object"): + mint_runtime_token( + namespace_key="default", + actor_id="actor", + target_type="target", + target_id="target-id", + scopes=("runtime.use",), + secret=_TEST_SECRET, + ttl_seconds=60, + extensions={"provider_field": object()}, + ) + + +def test_runtime_token_rejects_malformed_signed_extensions(): + import jwt + + from agent_control_server.auth_framework.runtime_token import ( + RuntimeTokenError, + mint_runtime_token, + verify_runtime_token, + ) + + token, _ = mint_runtime_token( + namespace_key="default", + actor_id="actor", + target_type="target", + target_id="target-id", + scopes=("runtime.use",), + secret=_TEST_SECRET, + ttl_seconds=60, + ) + payload = jwt.decode(token, _TEST_SECRET, algorithms=["HS256"]) + payload["extensions"] = [] + malformed_token = jwt.encode(payload, _TEST_SECRET, algorithm="HS256") + + with pytest.raises(RuntimeTokenError, match="malformed extensions"): + verify_runtime_token(malformed_token, _TEST_SECRET) + + def test_runtime_token_rejects_wrong_secret(): from agent_control_server.auth_framework.runtime_token import ( RuntimeTokenError, @@ -1215,6 +1260,28 @@ async def test_http_upstream_accepts_iso_datetime_and_array_scopes(): assert principal.grant_expires_at.isoformat() == iso_expiry +@pytest.mark.asyncio +async def test_http_upstream_rejects_invalid_extension_metadata(monkeypatch): + from pydantic import TypeAdapter + + from agent_control_server.auth_framework.providers import http_upstream + + provider = _build_upstream( + lambda req: httpx.Response( + 200, + json={"namespace_key": "namespace-1", "provider_field": "opaque_value"}, + ) + ) + # Force the defensive conversion error path for invalid extension payloads. + monkeypatch.setattr(http_upstream, "_JSON_OBJECT_ADAPTER", TypeAdapter(str)) + + with pytest.raises(APIError) as exc_info: + await provider.authorize(_build_request(), Operation.RUNTIME_TOKEN_EXCHANGE) + + assert exc_info.value.status_code == 502 + assert "malformed extension metadata" in exc_info.value.detail + + @pytest.mark.asyncio async def test_http_upstream_rejects_target_grant_mismatch(): provider = _build_upstream( From f48ea25de9646c53c1612f1e355c093bebe1c0df Mon Sep 17 00:00:00 2001 From: Namrata Ghadi Date: Wed, 30 Sep 2026 13:39:26 -0700 Subject: [PATCH 3/7] refactor(galileo): normalize records with galileo-core --- evaluators/contrib/galileo/pyproject.toml | 2 +- .../records/__init__.py | 2 + .../records/factory.py | 100 +-- .../records/normalization.py | 570 ++++++++++++++---- .../galileo/tests/test_records_factory.py | 404 +++++++++++-- 5 files changed, 871 insertions(+), 207 deletions(-) diff --git a/evaluators/contrib/galileo/pyproject.toml b/evaluators/contrib/galileo/pyproject.toml index 953ab04cd..6f0a19c47 100644 --- a/evaluators/contrib/galileo/pyproject.toml +++ b/evaluators/contrib/galileo/pyproject.toml @@ -9,8 +9,8 @@ authors = [{ name = "Agent Control Team" }] dependencies = [ "agent-control-evaluators>=8.6.0", "agent-control-models>=8.6.0", + "galileo-core>=4.4.0,<5.0.0", "httpx>=0.24.0", - "splunk-ao==0.4.0", "pydantic>=2.12.4", ] diff --git a/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/__init__.py b/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/__init__.py index 6cd089edf..7265e8a19 100644 --- a/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/__init__.py +++ b/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/__init__.py @@ -9,9 +9,11 @@ record_from_scorer_invoke_record, record_from_step, ) +from .normalization import GalileoRecordNormalizer __all__ = [ "GalileoRecord", + "GalileoRecordNormalizer", "RecordFactoryError", "UnsupportedStepTypeError", "build_galileo_record", diff --git a/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/factory.py b/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/factory.py index 64de8dca3..4b7959a24 100644 --- a/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/factory.py +++ b/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/factory.py @@ -6,25 +6,12 @@ from typing import Any, cast from agent_control_models import Step +from galileo_core.schemas.logging.session import Session +from galileo_core.schemas.logging.span import LlmSpan, RetrieverSpan, ToolSpan +from galileo_core.schemas.logging.trace import Trace from pydantic import BaseModel -from splunk_ao import ( # type: ignore[import-untyped] - LlmSpan, - RetrieverSpan, - Session, - ToolSpan, - Trace, -) - -from .normalization import ( - documents, - message_value, - session_value, - string_metadata, - text_value, - tool_value, - trace_input_value, - trace_output_value, -) + +from .normalization import GalileoRecordNormalizer type GalileoRecord = LlmSpan | ToolSpan | RetrieverSpan | Trace | Session type GalileoSpan = LlmSpan | ToolSpan | RetrieverSpan @@ -111,7 +98,9 @@ def _record_from_mapping(raw: Mapping[str, Any], record_type: str) -> GalileoRec context = _mapping(raw.get("context")) or {} common = { "name": "" if raw.get("name") is None else str(raw["name"]), - "user_metadata": string_metadata(raw.get("user_metadata", raw.get("metadata"))), + "user_metadata": GalileoRecordNormalizer.metadata( + raw.get("user_metadata", raw.get("metadata")) + ), } for field in ( "tags", @@ -126,9 +115,9 @@ def _record_from_mapping(raw: Mapping[str, Any], record_type: str) -> GalileoRec ): if raw.get(field) is not None: common[field] = ( - text_value(raw[field]) + GalileoRecordNormalizer.json_text(raw[field]) if field in {"dataset_input", "dataset_output"} - else string_metadata(raw[field]) + else GalileoRecordNormalizer.metadata(raw[field]) if field == "dataset_metadata" else raw[field] ) @@ -138,39 +127,55 @@ def _record_from_mapping(raw: Mapping[str, Any], record_type: str) -> GalileoRec if record_type == "llm": kwargs = { **common, - "input": message_value(input_value), - "output": message_value(output_value, output=True), + "input": GalileoRecordNormalizer.llm_input(input_value), + "output": GalileoRecordNormalizer.llm_output(output_value), } if raw.get("tools") is not None: - kwargs["tools"] = raw["tools"] + kwargs["tools"] = GalileoRecordNormalizer.llm_tools(raw["tools"]) for field in ("events", "model", "temperature", "finish_reason"): if raw.get(field) is not None: kwargs[field] = raw[field] if raw.get("redacted_input") is not None: - kwargs["redacted_input"] = message_value(raw["redacted_input"]) + kwargs["redacted_input"] = GalileoRecordNormalizer.llm_redacted_input( + raw["redacted_input"] + ) if raw.get("redacted_output") is not None: - kwargs["redacted_output"] = message_value(raw["redacted_output"], output=True) + kwargs["redacted_output"] = GalileoRecordNormalizer.llm_redacted_output( + raw["redacted_output"] + ) return LlmSpan(**kwargs) if record_type == "tool": kwargs = { **common, - "input": tool_value(input_value) or "", - "output": tool_value(output_value), + "input": GalileoRecordNormalizer.tool_input(input_value), + "output": GalileoRecordNormalizer.tool_output(output_value), } if raw.get("tool_call_id") is not None: kwargs["tool_call_id"] = str(raw["tool_call_id"]) if raw.get("redacted_input") is not None: - kwargs["redacted_input"] = tool_value(raw["redacted_input"]) + kwargs["redacted_input"] = GalileoRecordNormalizer.tool_redacted_input( + raw["redacted_input"] + ) if raw.get("redacted_output") is not None: - kwargs["redacted_output"] = tool_value(raw["redacted_output"]) + kwargs["redacted_output"] = GalileoRecordNormalizer.tool_redacted_output( + raw["redacted_output"] + ) kwargs["spans"] = [_nested_span(item) for item in _span_payloads(raw, context)] return ToolSpan(**kwargs) if record_type == "retriever": - kwargs = {**common, "input": text_value(input_value), "output": documents(output_value)} + kwargs = { + **common, + "input": GalileoRecordNormalizer.retriever_input(input_value), + "output": GalileoRecordNormalizer.retriever_output(output_value), + } if raw.get("redacted_input") is not None: - kwargs["redacted_input"] = text_value(raw["redacted_input"]) + kwargs["redacted_input"] = GalileoRecordNormalizer.retriever_redacted_input( + raw["redacted_input"] + ) if raw.get("redacted_output") is not None: - kwargs["redacted_output"] = documents(raw["redacted_output"]) + kwargs["redacted_output"] = GalileoRecordNormalizer.retriever_redacted_output( + raw["redacted_output"] + ) kwargs["spans"] = [_nested_span(item) for item in _span_payloads(raw, context)] return RetrieverSpan(**kwargs) if record_type == "trace": @@ -179,34 +184,45 @@ def _record_from_mapping(raw: Mapping[str, Any], record_type: str) -> GalileoRec raise RecordFactoryError("Galileo trace context must include a 'spans' list.") trace_kwargs: dict[str, Any] = { **common, - "input": trace_input_value(input_value), - "output": trace_output_value(output_value), + "input": GalileoRecordNormalizer.trace_input(input_value), + "output": GalileoRecordNormalizer.trace_output(output_value), "spans": [_nested_span(item) for item in spans], } if raw.get("redacted_input") is not None: - trace_kwargs["redacted_input"] = trace_input_value(raw["redacted_input"]) + trace_kwargs["redacted_input"] = GalileoRecordNormalizer.trace_redacted_input( + raw["redacted_input"] + ) if raw.get("redacted_output") is not None: - trace_kwargs["redacted_output"] = trace_output_value(raw["redacted_output"]) + trace_kwargs["redacted_output"] = GalileoRecordNormalizer.trace_redacted_output( + raw["redacted_output"] + ) return Trace( **trace_kwargs, ) traces = raw.get("traces", context.get("traces")) if not isinstance(traces, Sequence) or isinstance(traces, str | bytes | bytearray): raise RecordFactoryError("Galileo session context must include a 'traces' list.") - nested_records = [_nested_record(item) for item in traces] + nested_records = GalileoRecordNormalizer.session_traces( + traces, + normalize_record=_nested_record, + ) nested_traces = [item for item in nested_records if isinstance(item, Trace)] if len(nested_traces) != len(nested_records): raise RecordFactoryError("Galileo session 'traces' must contain trace records.") session_kwargs: dict[str, Any] = { **common, - "input": session_value(input_value), - "output": session_value(output_value), + "input": GalileoRecordNormalizer.session_input(input_value), + "output": GalileoRecordNormalizer.session_output(output_value), "traces": cast(list[Trace], nested_traces), } if raw.get("redacted_input") is not None: - session_kwargs["redacted_input"] = session_value(raw["redacted_input"]) + session_kwargs["redacted_input"] = GalileoRecordNormalizer.session_redacted_input( + raw["redacted_input"] + ) if raw.get("redacted_output") is not None: - session_kwargs["redacted_output"] = session_value(raw["redacted_output"]) + session_kwargs["redacted_output"] = GalileoRecordNormalizer.session_redacted_output( + raw["redacted_output"] + ) return Session( **session_kwargs, ) diff --git a/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/normalization.py b/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/normalization.py index cfa44b6ef..d9fb41ab1 100644 --- a/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/normalization.py +++ b/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/normalization.py @@ -1,150 +1,474 @@ -"""Value normalization shared by the Galileo record factory. - -The conversions in this module deliberately follow the published ``splunk-ao`` -helpers. The factory owns the Agent Control-to-record mapping, while the SDK -continues to own serialization and document coercion. - -Orbit's scorer-input normalizer retains list-of-message dictionaries as an -intermediate LLM output, but the published ``LlmSpan`` model accepts one output -message. The factory therefore serializes the complete sequence before public -model validation; this is an intentional Orbit/SDK boundary difference, not a -last-item selection rule. Likewise, retriever ``None`` follows the SDK logger -helper (one empty document), rather than the direct model default (an empty -list). -""" +"""Normalize Agent Control values for canonical Galileo record fields.""" from __future__ import annotations -from collections.abc import Mapping, Sequence -from typing import Any, cast +import json +from collections.abc import Callable, Mapping, Sequence +from dataclasses import asdict, is_dataclass +from datetime import UTC, date, datetime +from enum import Enum +from pathlib import Path +from typing import Any from uuid import UUID -from splunk_ao import Document, Message, SplunkAOLogger # type: ignore[import-untyped] -from splunk_ao.utils.retrievers import convert_to_documents # type: ignore[import-untyped] -from splunk_ao.utils.serialization import serialize_to_str # type: ignore[import-untyped] +from galileo_core.schemas.logging.llm import Message +from galileo_core.schemas.shared.content_parts import ContentPart, FileContentPart, TextContentPart +from galileo_core.schemas.shared.document import Document +from pydantic import BaseModel, TypeAdapter, ValidationError +_CONTENT_PART_ADAPTER: TypeAdapter[ContentPart] = TypeAdapter(ContentPart) +_MAX_SAFE_INTEGER = 2**53 - 1 -def text_value(value: Any, *, default: str = "") -> str: - """Use the SDK serializer for a value that a Galileo field stores as text.""" - if value is None: - return default - return cast(str, serialize_to_str(value)) +class _GalileoJSONEncoder(json.JSONEncoder): + """Serialize values using the stable subset of Splunk AO's JSON behavior.""" -def tool_value(value: Any) -> str | None: - """Serialize tool arguments/results using the SDK's canonical serializer.""" - if value is None: - return None - return cast(str, serialize_to_str(value)) + def default(self, value: Any) -> Any: + if isinstance(value, BaseModel): + return value.model_dump( + mode="json", + exclude_none=True, + exclude_unset=True, + exclude_defaults=True, + ) + if isinstance(value, datetime): + if value.tzinfo is None: + value = value.replace(tzinfo=datetime.now().astimezone().tzinfo) + serialized = value.isoformat() + if value.tzinfo is not None and value.tzinfo.tzname(None) == UTC.tzname(None): + return serialized.replace("+00:00", "Z") + return serialized + if isinstance(value, date): + return value.isoformat() + if isinstance(value, UUID | Path): + return str(value) + if isinstance(value, Enum): + return value.value + if isinstance(value, bytes): + try: + return value.decode("utf-8") + except UnicodeDecodeError: + return "" + if is_dataclass(value) and not isinstance(value, type): + return asdict(value) + if isinstance(value, int) and not isinstance(value, bool): + return value if -_MAX_SAFE_INTEGER <= value <= _MAX_SAFE_INTEGER else str(value) + if isinstance(value, set | frozenset): + return list(value) + return f"<{type(value).__name__}>" -def message_value(value: Any, *, output: bool = False) -> Any: - """Prepare values for the public ``LlmSpan`` model. +def _sdk_json_value(value: Any, *, active: set[int] | None = None) -> Any: + """Apply SDK-compatible conversion recursively, including mapping keys.""" + if value is None or isinstance(value, str | bool | float): + return value + if isinstance(value, int): + return value if -_MAX_SAFE_INTEGER <= value <= _MAX_SAFE_INTEGER else str(value) + if isinstance(value, UUID | Path): + return str(value) + if isinstance(value, datetime): + return _GalileoJSONEncoder().default(value) + if isinstance(value, date): + return value.isoformat() + if isinstance(value, Enum): + return _sdk_json_value(value.value, active=active) + if isinstance(value, bytes): + try: + return value.decode("utf-8") + except UnicodeDecodeError: + return "" - LLM inputs are passed through as message sequences so the canonical model - can validate each message. The public model has a single-message output - contract, so a list or tuple output is serialized as one complete JSON - string. This preserves every item and lets the canonical model construct - the assistant message. - """ - if value is None or isinstance(value, (str, Mapping, Message)): + seen = active if active is not None else set() + value_id = id(value) + if value_id in seen: + return type(value).__name__ + seen.add(value_id) + try: + if isinstance(value, BaseModel): + return _sdk_json_value( + value.model_dump( + mode="json", + exclude_none=True, + exclude_unset=True, + exclude_defaults=True, + ), + active=seen, + ) + if is_dataclass(value) and not isinstance(value, type): + return _sdk_json_value(asdict(value), active=seen) + if isinstance(value, Mapping): + return { + _sdk_json_value(key, active=seen): _sdk_json_value(item, active=seen) + for key, item in value.items() + } + if isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray): + return [_sdk_json_value(item, active=seen) for item in value] + if isinstance(value, set | frozenset): + return [_sdk_json_value(item, active=seen) for item in value] + if hasattr(value, "__slots__") and value.__slots__: + return _sdk_json_value( + {slot: getattr(value, slot, None) for slot in value.__slots__}, active=seen + ) + if hasattr(value, "__dict__"): + return _sdk_json_value( + {key: item for key, item in vars(value).items() if not key.startswith("_")}, + active=seen, + ) + return f"<{type(value).__name__}>" + finally: + seen.discard(value_id) + + +def _json_text(value: Any) -> str: + """Return the Galileo text representation of a structured value.""" + if isinstance(value, str): return value - if output and isinstance(value, Sequence) and not isinstance(value, bytes | bytearray): - return serialize_to_str(value) - return value + if value is None or isinstance(value, bool | int | float): + return json.dumps(value) + return json.dumps(_sdk_json_value(value)) -def documents(value: Any) -> Any: - """Use the SDK retriever helper, including its fallback/error behavior. +def _json_compatible(value: Any) -> Any: + """Convert models and containers into values accepted by Galileo Pydantic models.""" + if isinstance(value, BaseModel): + return value.model_dump(mode="python", exclude_none=True) + if isinstance(value, Mapping): + return {key: _json_compatible(item) for key, item in value.items()} + if isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray): + return [_json_compatible(item) for item in value] + if isinstance(value, UUID | datetime | date | Enum): + return _GalileoJSONEncoder().default(value) + return value - ``splunk_ao`` currently exposes an attrs-based ``Document`` while its - retriever helper consumes the logging document representation. Convert - that public resource model to its public dictionary form before delegating - to the SDK helper; no Galileo core type is imported here. - """ - if isinstance(value, Document): - value = _document_dict(value) - elif isinstance(value, list) and any(isinstance(item, Document) for item in value): - value = [_document_dict(item) if isinstance(item, Document) else item for item in value] - return convert_to_documents(value) - - -def _document_dict(value: Document) -> dict[str, Any]: - """Convert either public SDK Document metadata representation to a dict.""" - metadata = getattr(value, "metadata", None) - metadata_to_dict = getattr(metadata, "to_dict", None) - if callable(metadata_to_dict): - metadata = cast(Any, metadata_to_dict)() - elif isinstance(metadata, Mapping): - metadata = dict(metadata) - else: - metadata = None - - result: dict[str, Any] = {"content": value.content} - if metadata is not None: - result["metadata"] = metadata - return result - - -def string_metadata(value: Any) -> dict[str, str]: - """Apply the SDK's flat metadata conversion.""" + +def _content_part(value: Any) -> dict[str, Any] | None: + """Validate and serialize a public Galileo text or file content part.""" + if isinstance(value, (TextContentPart, FileContentPart)): + return value.model_dump(mode="json", exclude_none=True) if not isinstance(value, Mapping): - return {} - return { - str(key): "None" if item is None else item if isinstance(item, str) else str(item) - for key, item in value.items() - } + return None + try: + part = _CONTENT_PART_ADAPTER.validate_python(dict(value)) + except ValidationError: + return None + return part.model_dump(mode="json", exclude_none=True) -def trace_input_value(value: Any) -> Any: - """Delegate trace-input coercion to the SDK logger implementation.""" - if _is_public_trace_content_parts(value): - return value - return _public_content_blocks(SplunkAOLogger._coerce_trace_input("input", value)) +def _message_content(value: Any) -> Any: + if isinstance(value, Message): + return value.content + if isinstance(value, Mapping) and ("role" in value or "content" in value): + return value.get("content", "") + return None -def trace_output_value(value: Any) -> Any: - """Delegate trace-output coercion to the SDK logger implementation.""" - if value is None: +def _normalize_content_items(value: Any) -> list[dict[str, Any]] | None: + if not isinstance(value, Sequence) or isinstance(value, str | bytes | bytearray): return None - if _is_public_trace_content_parts(value): - return value - return _public_content_blocks(SplunkAOLogger._coerce_output(value)) + normalized: list[dict[str, Any]] = [] + for item in value: + block = _content_part(item) + if block is None: + return None + normalized.append(block) + return normalized -def _is_public_trace_content_parts(value: Any) -> bool: - """Recognize content parts already emitted by a public ``Trace`` dump.""" - if not isinstance(value, list): - return False - for item in value: - if not isinstance(item, Mapping): - return False - if item.get("type") == "text" and isinstance(item.get("text"), str): - continue - if item.get("type") == "file": - try: - if UUID(str(item["file_id"])).version == 4: - continue - except (KeyError, ValueError, AttributeError): - pass - return False - return True - - -def _public_content_blocks(value: Any) -> Any: - """Convert SDK ingestion blocks to the public model's content-part dicts.""" - if isinstance(value, list): - return [ - item.model_dump(exclude_none=True) if hasattr(item, "model_dump") else item - for item in value - ] - return value +def _message_sequence(value: Any) -> bool: + return isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray) and all( + isinstance(item, Message) + or (isinstance(item, Mapping) and ("role" in item or "content" in item)) + for item in value + ) -def session_value(value: Any) -> Any: - """Preserve sequences accepted by the public ``Session`` validator.""" +def _flatten_message_sequence(value: Sequence[Any]) -> list[dict[str, Any]]: + blocks: list[dict[str, Any]] = [] + for message in value: + content = _message_content(message) + if isinstance(content, str): + if content: + blocks.append(TextContentPart(text=content).model_dump(mode="json")) + elif isinstance(content, Sequence) and not isinstance(content, str | bytes | bytearray): + for item in content: + block = _content_part(item) + if block is not None: + blocks.append(block) + elif isinstance(item, str): + blocks.append(TextContentPart(text=item).model_dump(mode="json")) + else: + # Preserve unsupported multimodal payloads instead of dropping them. + blocks.append(TextContentPart(text=_json_text(item)).model_dump(mode="json")) + elif content is not None: + blocks.append(TextContentPart(text=_json_text(content)).model_dump(mode="json")) + return blocks + + +def _normalize_document(value: Any) -> Document: if isinstance(value, Document): - return _document_dict(value) - if isinstance(value, list) and any(isinstance(item, Document) for item in value): - return [_document_dict(item) if isinstance(item, Document) else item for item in value] - return value + return value + if isinstance(value, BaseModel): + value = value.model_dump(mode="python", exclude_none=True) + if isinstance(value, Mapping): + try: + return Document.model_validate(dict(value)) + except ValidationError: + return Document(content=_json_text(value)) + if isinstance(value, str): + return Document(content=value) + return Document(content=_json_text(value)) + + +class GalileoRecordNormalizer: + """Normalize destination fields for public Galileo record models. + + The normalizer owns conversions that differ by record destination. It uses + only public ``galileo-core`` models and preserves values that cannot be + represented natively as serialized text rather than discarding them. + """ + + @staticmethod + def json_text(value: Any) -> str: + """Serialize structured data while preserving strings and SDK JSON semantics.""" + return _json_text(value) + + @staticmethod + def metadata(value: Any) -> dict[str, str]: + """Convert metadata to the flat string mapping required by Galileo records.""" + if not isinstance(value, Mapping): + return {} + return { + str(key): "None" if item is None else item if isinstance(item, str) else str(item) + for key, item in value.items() + } + + @classmethod + def llm_input(cls, value: Any) -> Any: + """Normalize an LLM input without losing any message in a sequence.""" + if value is None: + return "" + if isinstance(value, str | Message): + return value + if isinstance(value, Mapping): + return _json_compatible(value) + if isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray): + if all(isinstance(item, Message | Mapping) for item in value): + return [_json_compatible(item) for item in value] + return cls.json_text(value) + + @classmethod + def llm_output(cls, value: Any) -> Any: + """Normalize an LLM output, serializing full message sequences as one message.""" + if value is None: + return "" + if isinstance(value, str | Message): + return value + if isinstance(value, Mapping): + return _json_compatible(value) + if _message_sequence(value): + return cls.json_text(value) + return cls.json_text(value) + + @classmethod + def llm_redacted_input(cls, value: Any) -> Any: + return None if value is None else cls.llm_input(value) + + @classmethod + def llm_redacted_output(cls, value: Any) -> Any: + return None if value is None else cls.llm_output(value) + + @staticmethod + def llm_tools(value: Any) -> Any: + """Normalize tool definitions nested in an LLM span.""" + return None if value is None else _sdk_json_value(value) + + @classmethod + def tool_input(cls, value: Any) -> str: + """Serialize tool input as text, using an empty string for missing input.""" + return "" if value is None else cls.json_text(value) + + @classmethod + def tool_output(cls, value: Any) -> str | None: + """Serialize a present tool output as text and preserve missing output.""" + return None if value is None else cls.json_text(value) + + @classmethod + def tool_redacted_input(cls, value: Any) -> str | None: + return None if value is None else cls.json_text(value) + + @classmethod + def tool_redacted_output(cls, value: Any) -> str | None: + return None if value is None else cls.json_text(value) + + @classmethod + def retriever_input(cls, value: Any) -> str: + return "" if value is None else cls.json_text(value) + + @staticmethod + def retriever_output(value: Any) -> list[Document]: + """Convert supported retrieval values into canonical Galileo Documents.""" + if value is None: + return [Document(content="")] + if isinstance(value, Document): + return [value] + if isinstance(value, BaseModel): + return [_normalize_document(value)] + if isinstance(value, str): + return [Document(content=value)] + if isinstance(value, Mapping): + return [_normalize_document(value)] + if isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray): + items = list(value) + if not items: + return [] + if all(isinstance(item, Document) for item in items): + return list(items) + if all(isinstance(item, str) for item in items): + return [Document(content=item) for item in items] + if all(isinstance(item, Mapping) for item in items): + return [_normalize_document(item) for item in items] + if all(isinstance(item, BaseModel) for item in items): + return [_normalize_document(item) for item in items] + raise ValueError( + "Invalid document output. Expected a list of strings, dictionaries, or Documents." + ) + return [Document(content="")] + + @classmethod + def retriever_redacted_input(cls, value: Any) -> str | None: + return None if value is None else cls.json_text(value) + + @classmethod + def retriever_redacted_output(cls, value: Any) -> list[Document] | None: + return None if value is None else cls.retriever_output(value) + + @classmethod + def trace_input(cls, value: Any) -> Any: + """Normalize trace input to a string or validated public content parts.""" + if value is None: + return "" + if isinstance(value, str): + return value + if isinstance(value, Mapping): + return cls.json_text(value) + if isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray): + normalized = _normalize_content_items(value) + if normalized is not None: + return normalized + if all(isinstance(item, Mapping) for item in value): + return cls.json_text(value) + raise TypeError("Trace input must be text, a mapping, or a list of content blocks.") + raise TypeError(f"Trace input does not support {type(value).__name__} values.") + + @classmethod + def trace_output(cls, value: Any) -> Any: + """Normalize trace output and flatten message sequences to content parts.""" + if value is None or isinstance(value, str): + return value + if _message_sequence(value): + return _flatten_message_sequence(value) + normalized = _normalize_content_items(value) + if normalized is not None: + return normalized + return cls.json_text(value) + + @classmethod + def trace_redacted_input(cls, value: Any) -> Any: + return None if value is None else cls.trace_input(value) + + @classmethod + def trace_redacted_output(cls, value: Any) -> Any: + return None if value is None else cls.trace_output(value) + + @classmethod + def session_input(cls, value: Any) -> Any: + """Normalize the session input variants accepted by the public model.""" + if value is None or isinstance(value, str | Message): + return value + if isinstance(value, TextContentPart | FileContentPart): + return [value] + if isinstance(value, Document): + return cls.json_text(value) + if isinstance(value, BaseModel): + return value + if isinstance(value, Mapping): + if "role" in value: + return _json_compatible(value) + part = _content_part(value) + if part is not None: + return [part] + return cls.json_text(value) + if isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray): + if not value: + return [] + if all(isinstance(item, Message) for item in value): + return list(value) + if all(isinstance(item, Mapping) for item in value): + if all("role" in item for item in value): + return [_json_compatible(item) for item in value] + content_parts = _normalize_content_items(value) + if content_parts is not None: + return content_parts + if all(isinstance(item, BaseModel) for item in value): + return cls.json_text(value) + return [_json_compatible(item) for item in value] + return cls.json_text(value) + + @classmethod + def session_output(cls, value: Any) -> Any: + """Normalize session output while preserving messages, documents and parts.""" + if value is None or isinstance(value, str | Message): + return value + if isinstance(value, Document): + return [value] + if isinstance(value, TextContentPart | FileContentPart): + return [value] + if isinstance(value, BaseModel): + return value + if isinstance(value, Mapping): + if "role" in value: + return _json_compatible(value) + if "content" in value or "page_content" in value: + return _normalize_document(value) + part = _content_part(value) + if part is not None: + return [part] + return cls.json_text(value) + if isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray): + if not value: + return [] + if all(isinstance(item, Document) for item in value): + return list(value) + if all(isinstance(item, Mapping) and "role" in item for item in value): + return _flatten_message_sequence(value) + if all( + isinstance(item, Mapping) and ("content" in item or "page_content" in item) + for item in value + ): + return [_normalize_document(item) for item in value] + if _message_sequence(value): + return _flatten_message_sequence(value) + content_parts = _normalize_content_items(value) + if content_parts is not None: + return content_parts + if all(isinstance(item, Mapping) for item in value): + return [_json_compatible(item) for item in value] + return [_json_compatible(item) for item in value] + return cls.json_text(value) + + @classmethod + def session_redacted_input(cls, value: Any) -> Any: + return None if value is None else cls.session_input(value) + + @classmethod + def session_redacted_output(cls, value: Any) -> Any: + return None if value is None else cls.session_output(value) + + @staticmethod + def session_traces( + values: Sequence[Any], + *, + normalize_record: Callable[[Any], Any], + ) -> list[Any]: + """Normalize every nested session trace through the shared record path.""" + return [normalize_record(value) for value in values] diff --git a/evaluators/contrib/galileo/tests/test_records_factory.py b/evaluators/contrib/galileo/tests/test_records_factory.py index dd489fdf5..aae63f2b2 100644 --- a/evaluators/contrib/galileo/tests/test_records_factory.py +++ b/evaluators/contrib/galileo/tests/test_records_factory.py @@ -3,11 +3,15 @@ from __future__ import annotations import json -from datetime import UTC, datetime +from dataclasses import dataclass +from datetime import UTC, date, datetime, timedelta, timezone +from enum import Enum +from pathlib import Path from uuid import uuid4 import pytest from agent_control_evaluator_galileo.records import ( + GalileoRecordNormalizer, RecordFactoryError, UnsupportedStepTypeError, build_galileo_record, @@ -16,10 +20,13 @@ record_from_step, ) from agent_control_models import Step -from pydantic import BaseModel, ConfigDict -from splunk_ao import Document, LlmSpan, Message, RetrieverSpan, Session, ToolSpan, Trace -from splunk_ao.utils.retrievers import convert_to_documents -from splunk_ao.utils.serialization import serialize_to_str +from galileo_core.schemas.logging.llm import Message +from galileo_core.schemas.logging.session import Session +from galileo_core.schemas.logging.span import LlmSpan, RetrieverSpan, ToolSpan +from galileo_core.schemas.logging.trace import Trace +from galileo_core.schemas.shared.content_parts import FileContentPart, TextContentPart +from galileo_core.schemas.shared.document import Document +from pydantic import BaseModel, ConfigDict, Field class _RecordPayload(BaseModel): @@ -37,7 +44,7 @@ class _RecordPayload(BaseModel): tool_call_id: object | None = None -def test_required_splunk_ao_public_exports_are_importable() -> None: +def test_required_galileo_core_public_exports_are_importable() -> None: assert all((LlmSpan, ToolSpan, RetrieverSpan, Trace, Session, Document, Message)) @@ -55,7 +62,7 @@ def test_llm_messages_use_public_canonical_validation() -> None: assert isinstance(record, LlmSpan) assert record.input[0].role.value == "user" assert record.output.role.value == "assistant" - assert record.dataset_output == serialize_to_str({"expected": "answer"}) + assert record.dataset_output == json.dumps({"expected": "answer"}) @pytest.mark.parametrize( @@ -101,7 +108,7 @@ def test_llm_tuple_output_uses_the_same_sdk_serializer() -> None: record = record_from_step(Step(type="llm", name="answer", input="question", output=output)) assert isinstance(record, LlmSpan) - assert record.output.content == serialize_to_str(output) + assert record.output.content == json.dumps(list(output)) assert json.loads(record.output.content) == list(output) @@ -133,8 +140,8 @@ def test_tool_values_are_json_strings_and_missing_output_stays_missing() -> None expected = ToolSpan( name="search", - input=serialize_to_str({"query": "q"}), - output=serialize_to_str({"hits": [1, 2]}), + input=json.dumps({"query": "q"}), + output=json.dumps({"hits": [1, 2]}), ) assert type(record) is type(expected) assert record.input == expected.input @@ -146,21 +153,20 @@ def test_tool_values_are_json_strings_and_missing_output_stays_missing() -> None [None, "text", {"content": "document"}, {"invalid": True}, 42, ["one", "two"]], ) def test_retriever_matches_sdk_document_coercion(output: object) -> None: - expected = convert_to_documents(output) - - try: - record = record_from_step( - Step(type="retriever", name="retrieve", input="question", output=output) - ) - except (TypeError, ValueError): - with pytest.raises((TypeError, ValueError)): - RetrieverSpan(input="question", output=expected) - return - + record = record_from_step( + Step(type="retriever", name="retrieve", input="question", output=output) + ) assert isinstance(record, RetrieverSpan) - assert [document.model_dump() for document in record.output] == [ - document.model_dump() for document in expected - ] + if output is None or isinstance(output, int): + assert [document.content for document in record.output] == [""] + elif isinstance(output, str): + assert [document.content for document in record.output] == [output] + elif isinstance(output, dict) and "content" in output: + assert [document.content for document in record.output] == [output["content"]] + elif output == ["one", "two"]: + assert [document.content for document in record.output] == output + else: + assert record.output[0].content == json.dumps(output) def test_retriever_rejects_mixed_lists_like_the_sdk_helper() -> None: @@ -190,7 +196,7 @@ def test_retriever_results_become_documents_with_scalar_metadata_only() -> None: assert record.output[0].metadata == {"score": 0.9, "source": "kb"} -def test_retriever_accepts_the_public_splunk_ao_document_model() -> None: +def test_retriever_accepts_the_public_galileo_core_document_model() -> None: from pydantic import BaseModel document = Document(content="context", metadata={"source": "kb"}) @@ -210,7 +216,7 @@ class RecordInput(BaseModel): def test_public_document_metadata_variants_are_supported() -> None: document_without_metadata = Document(content="plain") - document_with_model_metadata = Document.from_dict( + document_with_model_metadata = Document.model_validate( {"content": "model metadata", "metadata": {"source": "kb"}} ) @@ -406,8 +412,8 @@ def test_malformed_trace_content_parts_follow_sdk_serialization() -> None: assert isinstance(scalar_item, Trace) assert isinstance(invalid_file, Trace) - assert scalar_item.output == serialize_to_str([1]) - assert invalid_file.output == serialize_to_str([{"type": "file", "file_id": "not-a-uuid"}]) + assert scalar_item.output == json.dumps([1]) + assert invalid_file.output == json.dumps([{"type": "file", "file_id": "not-a-uuid"}]) def test_nested_record_errors_and_existing_models_are_explicit() -> None: @@ -469,14 +475,9 @@ def test_invalid_nested_context_shapes_are_rejected() -> None: ) -def test_trace_structured_values_match_sdk_logger_coercion() -> None: - from splunk_ao import SplunkAOLogger - +def test_trace_structured_values_are_normalized_without_losing_messages() -> None: trace_input = {"question": "hello"} trace_output = [{"role": "assistant", "content": "answer"}] - expected_input = SplunkAOLogger._coerce_trace_input("input", trace_input) - expected_output = SplunkAOLogger._coerce_output(trace_output) - record = record_from_step( Step( type="trace", @@ -488,8 +489,325 @@ def test_trace_structured_values_match_sdk_logger_coercion() -> None: ) assert isinstance(record, Trace) - assert record.input == expected_input - assert [block.text for block in record.output] == [block.text for block in expected_output] + assert record.input == json.dumps(trace_input) + assert [block.text for block in record.output] == ["answer"] + + +def test_trace_content_blocks_preserve_text_files_and_unrepresentable_data() -> None: + file_id = uuid4() + blocks = [ + {"type": "text", "text": "look at this"}, + {"type": "file", "file_id": str(file_id)}, + ] + trace = record_from_step( + Step( + type="trace", + name="multimodal", + input=blocks, + output=blocks, + context={"spans": []}, + ) + ) + serialized_data_block = { + "type": "data", + "modality": "image", + "url": "https://example/image.png", + } + data_trace = record_from_step( + Step( + type="trace", + name="inline-image", + input=[serialized_data_block], + output=[serialized_data_block], + context={"spans": []}, + ) + ) + + assert isinstance(trace, Trace) + assert isinstance(trace.input[0], TextContentPart) + assert isinstance(trace.input[1], FileContentPart) + assert trace.input[1].file_id == file_id + assert isinstance(trace.output[0], TextContentPart) + assert isinstance(trace.output[1], FileContentPart) + assert isinstance(data_trace, Trace) + assert json.loads(data_trace.input) == [serialized_data_block] + assert json.loads(data_trace.output or "") == [serialized_data_block] + + +def test_trace_message_sequences_flatten_every_text_and_content_part() -> None: + file_id = uuid4() + output = [ + {"role": "assistant", "content": "first message"}, + { + "role": "assistant", + "content": [ + {"type": "text", "text": "second message"}, + {"type": "file", "file_id": str(file_id)}, + ], + }, + ] + + trace = record_from_step( + Step(type="trace", name="messages", input="question", output=output, context={"spans": []}) + ) + + assert isinstance(trace, Trace) + assert [part.text for part in trace.output if isinstance(part, TextContentPart)] == [ + "first message", + "second message", + ] + file_parts = [part for part in trace.output if isinstance(part, FileContentPart)] + assert len(file_parts) == 1 + assert file_parts[0].file_id == file_id + + +def test_record_normalizer_defines_none_and_empty_value_behavior() -> None: + assert GalileoRecordNormalizer.llm_input(None) == "" + assert GalileoRecordNormalizer.llm_output(None) == "" + assert GalileoRecordNormalizer.tool_input(None) == "" + assert GalileoRecordNormalizer.tool_output(None) is None + assert GalileoRecordNormalizer.retriever_output(None)[0].content == "" + assert GalileoRecordNormalizer.retriever_output([]) == [] + assert GalileoRecordNormalizer.trace_input(None) == "" + assert GalileoRecordNormalizer.trace_output(None) is None + assert GalileoRecordNormalizer.trace_output([]) == [] + assert GalileoRecordNormalizer.session_input(None) is None + assert GalileoRecordNormalizer.session_output(None) is None + assert GalileoRecordNormalizer.session_input([]) == [] + assert GalileoRecordNormalizer.session_output([]) == [] + + +def test_session_content_parts_are_preserved_by_input_and_output_normalizers() -> None: + parts = [TextContentPart(text="text"), FileContentPart(file_id=uuid4())] + + assert GalileoRecordNormalizer.session_input(parts) == [ + part.model_dump(mode="json") for part in parts + ] + assert GalileoRecordNormalizer.session_output(parts) == [ + part.model_dump(mode="json") for part in parts + ] + + +def test_record_normalizer_serializes_nested_models_uuids_and_timestamps() -> None: + identifier = uuid4() + timestamp = datetime(2025, 1, 2, 3, 4, 5, tzinfo=UTC) + + class NestedPayload(BaseModel): + identifier: object + timestamp: datetime + + payload = { + "nested": {"answer": "42"}, + "payload": NestedPayload(identifier=identifier, timestamp=timestamp), + "uuid": identifier, + "timestamp": timestamp, + } + serialized = GalileoRecordNormalizer.tool_output(payload) + + assert serialized == ( + '{"nested": {"answer": "42"}, "payload": {"identifier": "' + + str(identifier) + + '", "timestamp": "2025-01-02T03:04:05Z"}, "uuid": "' + + str(identifier) + + '", "timestamp": "2025-01-02T03:04:05Z"}' + ) + + +def test_json_normalization_supports_python_value_types_and_fallbacks() -> None: + from agent_control_evaluator_galileo.records.normalization import ( + _GalileoJSONEncoder, + _json_compatible, + _sdk_json_value, + ) + + class State(Enum): + READY = "ready" + + class NestedModel(BaseModel): + value: int + optional: str | None = None + + @dataclass + class NestedDataclass: + identifier: object + + class SlottedValue: + __slots__ = ("value", "missing") + + def __init__(self) -> None: + self.value = "slot" + + class PlainValue: + def __init__(self) -> None: + self.visible = "public" + self._hidden = "private" + + encoder = _GalileoJSONEncoder() + identifier = uuid4() + utc_timestamp = datetime(2025, 1, 2, tzinfo=UTC) + naive_timestamp = datetime(2025, 1, 2) + offset_timestamp = datetime(2025, 1, 2, tzinfo=timezone(timedelta(hours=2))) + + assert encoder.default(NestedModel(value=3)) == {"value": 3} + assert encoder.default(utc_timestamp) == "2025-01-02T00:00:00Z" + assert isinstance(encoder.default(naive_timestamp), str) + assert encoder.default(offset_timestamp) == "2025-01-02T00:00:00+02:00" + assert encoder.default(date(2025, 1, 2)) == "2025-01-02" + assert encoder.default(identifier) == str(identifier) + assert encoder.default(Path("folder/file.txt")) == "folder/file.txt" + assert encoder.default(State.READY) == "ready" + assert encoder.default(b"text") == "text" + assert encoder.default(b"\xff") == "" + assert encoder.default(NestedDataclass(identifier)) == {"identifier": identifier} + assert encoder.default(2**53) == str(2**53) + assert encoder.default({"item"}) == ["item"] + assert encoder.default(object()) == "" + + recursive: list[object] = [] + recursive.append(recursive) + normalized = _sdk_json_value( + { + identifier: NestedModel(value=4), + "date": date(2025, 1, 2), + "enum": State.READY, + "bytes": b"text", + "invalid_bytes": b"\xff", + "large_integer": 2**53, + "set": frozenset({"item"}), + "dataclass": NestedDataclass(identifier), + "slotted": SlottedValue(), + "plain": PlainValue(), + "unsupported": object(), + "recursive": recursive, + } + ) + assert normalized[str(identifier)] == {"value": 4} + assert normalized["date"] == "2025-01-02" + assert normalized["enum"] == "ready" + assert normalized["bytes"] == "text" + assert normalized["invalid_bytes"] == "" + assert normalized["large_integer"] == str(2**53) + assert normalized["set"] == ["item"] + assert normalized["dataclass"] == {"identifier": str(identifier)} + assert normalized["slotted"] == {"value": "slot", "missing": None} + assert normalized["plain"] == {"visible": "public"} + assert normalized["unsupported"] == "" + assert normalized["recursive"] == ["list"] + assert _json_compatible(NestedModel(value=5)) == {"value": 5} + assert _json_compatible({"values": [State.READY, date(2025, 1, 2)]}) == { + "values": ["ready", "2025-01-02"] + } + + +def test_normalizer_preserves_message_content_and_reports_invalid_trace_inputs() -> None: + from agent_control_evaluator_galileo.records.normalization import ( + _flatten_message_sequence, + _message_sequence, + _normalize_content_items, + ) + + messages = [ + {"role": "user", "content": ""}, + {"role": "assistant", "content": ["plain", {"unsupported": True}]}, + {"content": 7}, + {}, + ] + flattened = _flatten_message_sequence(messages) + + assert [part["text"] for part in flattened] == [ + "plain", + '{"unsupported": true}', + "7", + ] + assert _message_sequence([{"role": "user", "content": "question"}]) + assert not _message_sequence([{"other": "field"}]) + assert _normalize_content_items("not a sequence") is None + assert _normalize_content_items([{"invalid": True}]) is None + assert GalileoRecordNormalizer.llm_input(42) == "42" + assert GalileoRecordNormalizer.llm_output([1, 2]) == "[1, 2]" + with pytest.raises(TypeError, match="Trace input must be"): + GalileoRecordNormalizer.trace_input([1]) + with pytest.raises(TypeError, match="does not support int"): + GalileoRecordNormalizer.trace_input(1) + + +def test_retriever_and_session_normalizers_accept_public_models() -> None: + class DocumentModel(BaseModel): + content: str + metadata: dict[str, str] = Field(default_factory=dict) + + pydantic_document = DocumentModel(content="model document") + document = Document(content="canonical document") + part = TextContentPart(text="content part") + message = Message(role="user", content="question") + + assert GalileoRecordNormalizer.retriever_output(pydantic_document) == [ + Document(content="model document", metadata={}) + ] + assert GalileoRecordNormalizer.retriever_output([pydantic_document]) == [ + Document(content="model document", metadata={}) + ] + assert GalileoRecordNormalizer.retriever_output([document]) == [document] + assert GalileoRecordNormalizer.session_input(part) == [part] + assert GalileoRecordNormalizer.session_input(document) == '{"content": "canonical document"}' + assert GalileoRecordNormalizer.session_input(pydantic_document) is pydantic_document + assert GalileoRecordNormalizer.session_input({"role": "user", "content": "question"}) == { + "role": "user", + "content": "question", + } + assert GalileoRecordNormalizer.session_input({"type": "text", "text": "hello"}) == [ + {"type": "text", "text": "hello"} + ] + assert GalileoRecordNormalizer.session_input([message]) == [message] + assert GalileoRecordNormalizer.session_input([{"role": "user", "content": "question"}]) == [ + {"role": "user", "content": "question"} + ] + assert json.loads(GalileoRecordNormalizer.session_input([pydantic_document]))[0][ + "content" + ] == "model document" + assert GalileoRecordNormalizer.session_output(document) == [document] + assert GalileoRecordNormalizer.session_output(part) == [part] + assert GalileoRecordNormalizer.session_output(pydantic_document) is pydantic_document + assert GalileoRecordNormalizer.session_output({"role": "assistant", "content": "answer"}) == { + "role": "assistant", + "content": "answer", + } + assert GalileoRecordNormalizer.session_output({"page_content": "page"}).content == "page" + assert GalileoRecordNormalizer.session_output({"type": "text", "text": "hello"}) == [ + {"type": "text", "text": "hello"} + ] + assert GalileoRecordNormalizer.session_output([document]) == [document] + assert GalileoRecordNormalizer.session_output([{"role": "assistant", "content": "answer"}]) == [ + {"type": "text", "text": "answer"} + ] + + +def test_llm_tool_definitions_normalize_nested_pydantic_uuid_and_datetime_values() -> None: + identifier = uuid4() + timestamp = datetime(2025, 1, 2, tzinfo=UTC) + source = _RecordPayload( + type="llm", + input="question", + tools=[{"id": identifier, "created_at": timestamp}], + ) + + record = record_from_scorer_invoke_record(source) + + assert isinstance(record, LlmSpan) + assert record.tools == [ + {"id": str(identifier), "created_at": timestamp.isoformat().replace("+00:00", "Z")} + ] + + +def test_retriever_dictionaries_become_canonical_documents() -> None: + documents = GalileoRecordNormalizer.retriever_output( + [{"content": "document", "metadata": {"source": "kb"}}] + ) + + assert len(documents) == 1 + assert isinstance(documents[0], Document) + assert documents[0].content == "document" + assert documents[0].metadata == {"source": "kb"} def test_trace_and_session_envelopes_are_supported() -> None: @@ -660,7 +978,7 @@ def test_record_serialization_is_json_safe() -> None: payload = record.model_dump(mode="json", exclude_none=True) assert payload["type"] == "tool" - assert payload["input"] == serialize_to_str({"q": "x"}) + assert payload["input"] == json.dumps({"q": "x"}) def test_factory_accepts_the_existing_pydantic_luna_record() -> None: @@ -675,12 +993,16 @@ def test_factory_accepts_the_existing_pydantic_luna_record() -> None: def test_factory_rejects_invalid_boundaries_and_normalizes_missing_text() -> None: - from agent_control_evaluator_galileo.records.normalization import session_value, text_value + from agent_control_evaluator_galileo.records import GalileoRecordNormalizer with pytest.raises(RecordFactoryError, match="complete Agent Control Step"): record_from_step(object()) # type: ignore[arg-type] with pytest.raises(UnsupportedStepTypeError, match="unsupported"): record_from_scorer_invoke_record(_RecordPayload(type="unsupported")) - assert text_value(None) == "" - assert session_value(Document(content="plain")) == {"content": "plain"} - assert session_value([Document(content="plain")]) == [{"content": "plain"}] + assert GalileoRecordNormalizer.retriever_input(None) == "" + assert GalileoRecordNormalizer.session_input(Document(content="plain")) == json.dumps( + {"content": "plain"} + ) + assert GalileoRecordNormalizer.session_input([Document(content="plain")]) == json.dumps( + [{"content": "plain"}] + ) From 9e25a3b9022912eeb138ca2061c422540a46f4a3 Mon Sep 17 00:00:00 2001 From: Namrata Ghadi Date: Wed, 30 Sep 2026 14:00:12 -0700 Subject: [PATCH 4/7] test(galileo): verify factory BaseStep compatibility --- .../records/factory.py | 5 +- .../galileo/tests/test_records_factory.py | 59 +++++++++++++++++++ 2 files changed, 63 insertions(+), 1 deletion(-) diff --git a/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/factory.py b/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/factory.py index 4b7959a24..d34047f5f 100644 --- a/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/factory.py +++ b/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/factory.py @@ -242,7 +242,10 @@ def record_from_step( values override only the corresponding input/output fields, preserving the selected evaluator payload while the unselected side comes from ``step``. Trace and session records require structured child context; scalar selected - values are never promoted into those record types. + values are never promoted into those record types. The result is a concrete + Galileo Core step model (a ``BaseStep`` subclass) that can be JSON-serialized + and revalidated using that subtype's schema. The factory doesn't assign + execution IDs or storage ownership fields. """ if not isinstance(step, Step): raise RecordFactoryError("A complete Agent Control Step is required.") diff --git a/evaluators/contrib/galileo/tests/test_records_factory.py b/evaluators/contrib/galileo/tests/test_records_factory.py index aae63f2b2..66ac9c9e5 100644 --- a/evaluators/contrib/galileo/tests/test_records_factory.py +++ b/evaluators/contrib/galileo/tests/test_records_factory.py @@ -23,6 +23,7 @@ from galileo_core.schemas.logging.llm import Message from galileo_core.schemas.logging.session import Session from galileo_core.schemas.logging.span import LlmSpan, RetrieverSpan, ToolSpan +from galileo_core.schemas.logging.step import BaseStep from galileo_core.schemas.logging.trace import Trace from galileo_core.schemas.shared.content_parts import FileContentPart, TextContentPart from galileo_core.schemas.shared.document import Document @@ -48,6 +49,64 @@ def test_required_galileo_core_public_exports_are_importable() -> None: assert all((LlmSpan, ToolSpan, RetrieverSpan, Trace, Session, Document, Message)) +def test_factory_records_are_serializable_concrete_base_steps() -> None: + # Given: one record of each type accepted by the Galileo factory + records = [ + record_from_step(Step(type="llm", name="answer", input="question", output="answer")), + record_from_step(Step(type="tool", name="search", input={"query": "q"}, output="result")), + record_from_step( + Step(type="retriever", name="retrieve", input="question", output=["document"]) + ), + record_from_step( + Step( + type="trace", + name="request", + input="question", + output="answer", + context={"spans": [{"type": "llm", "name": "answer", "input": "question"}]}, + ) + ), + record_from_step( + Step( + type="session", + name="conversation", + input="question", + context={ + "traces": [ + { + "type": "trace", + "name": "request", + "input": "question", + "spans": [{"type": "llm", "name": "answer", "input": "question"}], + } + ] + }, + ) + ), + ] + + # When: each concrete Galileo step is serialized and parsed back through its own schema + rebuilt = [ + type(record).model_validate(record.model_dump(mode="json", exclude_none=True)) + for record in records + ] + + # Then: every concrete class remains a BaseStep and keeps its discriminator and children + assert all(isinstance(record, BaseStep) for record in records) + assert all(isinstance(record, BaseStep) for record in rebuilt) + assert [record.type.value for record in rebuilt] == [ + "llm", + "tool", + "retriever", + "trace", + "session", + ] + assert isinstance(rebuilt[3], Trace) + assert isinstance(rebuilt[4], Session) + assert isinstance(rebuilt[3].spans[0], LlmSpan) + assert isinstance(rebuilt[4].traces[0].spans[0], LlmSpan) + + def test_llm_messages_use_public_canonical_validation() -> None: record = record_from_step( Step( From a0dc0000a8c12e8e9513eaced1418f5e88810ab3 Mon Sep 17 00:00:00 2001 From: Namrata Ghadi Date: Wed, 30 Sep 2026 14:07:32 -0700 Subject: [PATCH 5/7] test(galileo): verify record conversion compatibility --- .../galileo/tests/test_records_factory.py | 49 +++++++++++++++++++ 1 file changed, 49 insertions(+) diff --git a/evaluators/contrib/galileo/tests/test_records_factory.py b/evaluators/contrib/galileo/tests/test_records_factory.py index 66ac9c9e5..814eca283 100644 --- a/evaluators/contrib/galileo/tests/test_records_factory.py +++ b/evaluators/contrib/galileo/tests/test_records_factory.py @@ -27,6 +27,7 @@ from galileo_core.schemas.logging.trace import Trace from galileo_core.schemas.shared.content_parts import FileContentPart, TextContentPart from galileo_core.schemas.shared.document import Document +from galileo_core.schemas.shared.records import BaseRecord, RecordTypeAdapter from pydantic import BaseModel, ConfigDict, Field @@ -107,6 +108,54 @@ def test_factory_records_are_serializable_concrete_base_steps() -> None: assert isinstance(rebuilt[4].traces[0].spans[0], LlmSpan) +def test_factory_steps_convert_to_core_records_after_execution_ids_are_added() -> None: + # Given: each root step subtype emitted by the factory, plus the IDs the execution caller supplies + steps = [ + record_from_step(Step(type="llm", name="answer", input="question", output="answer")), + record_from_step(Step(type="tool", name="search", input={"query": "q"}, output="result")), + record_from_step( + Step(type="retriever", name="retrieve", input="question", output=["document"]) + ), + record_from_step( + Step(type="trace", name="request", input="question", context={"spans": []}) + ), + record_from_step( + Step(type="session", name="conversation", input="question", context={"traces": []}) + ), + ] + project_id = uuid4() + run_id = uuid4() + session_id = uuid4() + trace_id = uuid4() + rebuilt_records: list[BaseRecord] = [] + + # When: add execution IDs and relationships, then apply Galileo Core's record discriminator + for step in steps: + step_id = uuid4() + values = step.model_dump( + mode="python", + exclude={"id", "project_id", "run_id", "session_id", "trace_id", "parent_id", "spans", "traces"}, + ) + values.update(id=step_id, project_id=project_id, run_id=run_id, type=step.type) + if isinstance(step, Session): + values["session_id"] = step_id + elif isinstance(step, Trace): + values.update(session_id=session_id, trace_id=step_id) + else: + values.update(session_id=session_id, trace_id=trace_id, parent_id=trace_id) + rebuilt_records.append(RecordTypeAdapter.validate_python(values)) + + # Then: the concrete step payloads satisfy Galileo Core's stored-record schemas + assert all(isinstance(record, BaseRecord) for record in rebuilt_records) + assert [record.type.value for record in rebuilt_records] == [ + "llm", + "tool", + "retriever", + "trace", + "session", + ] + + def test_llm_messages_use_public_canonical_validation() -> None: record = record_from_step( Step( From 6c84192b3376e1da948eb1193ae612d710d25e2b Mon Sep 17 00:00:00 2001 From: Namrata Ghadi Date: Thu, 1 Oct 2026 10:04:52 -0700 Subject: [PATCH 6/7] address comments --- .../records/normalization.py | 22 +++---- .../galileo/tests/test_records_factory.py | 57 +++++++++++++++++-- 2 files changed, 61 insertions(+), 18 deletions(-) diff --git a/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/normalization.py b/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/normalization.py index d9fb41ab1..5144cbbb7 100644 --- a/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/normalization.py +++ b/evaluators/contrib/galileo/src/agent_control_evaluator_galileo/records/normalization.py @@ -33,7 +33,7 @@ def default(self, value: Any) -> Any: ) if isinstance(value, datetime): if value.tzinfo is None: - value = value.replace(tzinfo=datetime.now().astimezone().tzinfo) + value = value.astimezone() serialized = value.isoformat() if value.tzinfo is not None and value.tzinfo.tzname(None) == UTC.tzname(None): return serialized.replace("+00:00", "Z") @@ -43,7 +43,7 @@ def default(self, value: Any) -> Any: if isinstance(value, UUID | Path): return str(value) if isinstance(value, Enum): - return value.value + return _sdk_json_value(value.value) if isinstance(value, bytes): try: return value.decode("utf-8") @@ -51,8 +51,6 @@ def default(self, value: Any) -> Any: return "" if is_dataclass(value) and not isinstance(value, type): return asdict(value) - if isinstance(value, int) and not isinstance(value, bool): - return value if -_MAX_SAFE_INTEGER <= value <= _MAX_SAFE_INTEGER else str(value) if isinstance(value, set | frozenset): return list(value) return f"<{type(value).__name__}>" @@ -123,20 +121,20 @@ def _json_text(value: Any) -> str: """Return the Galileo text representation of a structured value.""" if isinstance(value, str): return value - if value is None or isinstance(value, bool | int | float): - return json.dumps(value) return json.dumps(_sdk_json_value(value)) def _json_compatible(value: Any) -> Any: """Convert models and containers into values accepted by Galileo Pydantic models.""" if isinstance(value, BaseModel): - return value.model_dump(mode="python", exclude_none=True) + return _json_compatible(value.model_dump(mode="python", exclude_none=True)) if isinstance(value, Mapping): return {key: _json_compatible(item) for key, item in value.items()} if isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray): return [_json_compatible(item) for item in value] - if isinstance(value, UUID | datetime | date | Enum): + if isinstance(value, set | frozenset): + return [_json_compatible(item) for item in value] + if isinstance(value, UUID | datetime | date | Enum | Path | bytes): return _GalileoJSONEncoder().default(value) return value @@ -157,7 +155,7 @@ def _content_part(value: Any) -> dict[str, Any] | None: def _message_content(value: Any) -> Any: if isinstance(value, Message): return value.content - if isinstance(value, Mapping) and ("role" in value or "content" in value): + if isinstance(value, Mapping) and "role" in value: return value.get("content", "") return None @@ -177,7 +175,7 @@ def _normalize_content_items(value: Any) -> list[dict[str, Any]] | None: def _message_sequence(value: Any) -> bool: return isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray) and all( isinstance(item, Message) - or (isinstance(item, Mapping) and ("role" in item or "content" in item)) + or (isinstance(item, Mapping) and "role" in item) for item in value ) @@ -451,9 +449,7 @@ def session_output(cls, value: Any) -> Any: content_parts = _normalize_content_items(value) if content_parts is not None: return content_parts - if all(isinstance(item, Mapping) for item in value): - return [_json_compatible(item) for item in value] - return [_json_compatible(item) for item in value] + return cls.json_text(value) return cls.json_text(value) @classmethod diff --git a/evaluators/contrib/galileo/tests/test_records_factory.py b/evaluators/contrib/galileo/tests/test_records_factory.py index 814eca283..482688f83 100644 --- a/evaluators/contrib/galileo/tests/test_records_factory.py +++ b/evaluators/contrib/galileo/tests/test_records_factory.py @@ -696,6 +696,16 @@ def test_session_content_parts_are_preserved_by_input_and_output_normalizers() - ] +def test_session_output_serializes_unsupported_sequences_for_core_model() -> None: + output = GalileoRecordNormalizer.session_output([1, 2]) + + assert isinstance(output, str) + assert output == json.dumps([1, 2]) + assert json.loads(output) == [1, 2] + session = Session(input=[], output=output, traces=[]) + assert session.output == output + + def test_record_normalizer_serializes_nested_models_uuids_and_timestamps() -> None: identifier = uuid4() timestamp = datetime(2025, 1, 2, 3, 4, 5, tzinfo=UTC) @@ -735,6 +745,9 @@ class NestedModel(BaseModel): value: int optional: str | None = None + class NestedValuesModel(BaseModel): + payload: object + @dataclass class NestedDataclass: identifier: object @@ -756,21 +769,32 @@ def __init__(self) -> None: naive_timestamp = datetime(2025, 1, 2) offset_timestamp = datetime(2025, 1, 2, tzinfo=timezone(timedelta(hours=2))) + class StructuredState(Enum): + VALUE = {"identifier": identifier, "timestamp": utc_timestamp} + assert encoder.default(NestedModel(value=3)) == {"value": 3} assert encoder.default(utc_timestamp) == "2025-01-02T00:00:00Z" - assert isinstance(encoder.default(naive_timestamp), str) + assert encoder.default(naive_timestamp) == naive_timestamp.astimezone().isoformat() assert encoder.default(offset_timestamp) == "2025-01-02T00:00:00+02:00" assert encoder.default(date(2025, 1, 2)) == "2025-01-02" assert encoder.default(identifier) == str(identifier) assert encoder.default(Path("folder/file.txt")) == "folder/file.txt" assert encoder.default(State.READY) == "ready" + assert json.dumps(StructuredState.VALUE, cls=_GalileoJSONEncoder) == json.dumps( + {"identifier": str(identifier), "timestamp": "2025-01-02T00:00:00Z"} + ) assert encoder.default(b"text") == "text" assert encoder.default(b"\xff") == "" assert encoder.default(NestedDataclass(identifier)) == {"identifier": identifier} - assert encoder.default(2**53) == str(2**53) assert encoder.default({"item"}) == ["item"] assert encoder.default(object()) == "" + assert GalileoRecordNormalizer.json_text(42) == "42" + assert GalileoRecordNormalizer.json_text(2**53) == json.dumps(str(2**53)) + assert GalileoRecordNormalizer.json_text({"large_integer": 2**53}) == json.dumps( + {"large_integer": str(2**53)} + ) + recursive: list[object] = [] recursive.append(recursive) normalized = _sdk_json_value( @@ -805,11 +829,33 @@ def __init__(self) -> None: assert _json_compatible({"values": [State.READY, date(2025, 1, 2)]}) == { "values": ["ready", "2025-01-02"] } + nested_values = _json_compatible( + { + "nested": { + "set": {b"bytes"}, + "frozenset": frozenset({Path("nested/path")}), + "invalid_bytes": {b"\xff"}, + }, + "model": NestedValuesModel( + payload={"set": {b"model bytes"}, "path": Path("model/path")} + ), + } + ) + assert nested_values == { + "nested": { + "set": ["bytes"], + "frozenset": ["nested/path"], + "invalid_bytes": [""], + }, + "model": {"payload": {"set": ["model bytes"], "path": "model/path"}}, + } + assert json.loads(json.dumps(nested_values)) == nested_values def test_normalizer_preserves_message_content_and_reports_invalid_trace_inputs() -> None: from agent_control_evaluator_galileo.records.normalization import ( _flatten_message_sequence, + _message_content, _message_sequence, _normalize_content_items, ) @@ -817,17 +863,18 @@ def test_normalizer_preserves_message_content_and_reports_invalid_trace_inputs() messages = [ {"role": "user", "content": ""}, {"role": "assistant", "content": ["plain", {"unsupported": True}]}, - {"content": 7}, - {}, ] flattened = _flatten_message_sequence(messages) assert [part["text"] for part in flattened] == [ "plain", '{"unsupported": true}', - "7", ] assert _message_sequence([{"role": "user", "content": "question"}]) + assert not _message_sequence([{"content": "question"}]) + retriever_document = {"content": "retrieved text", "metadata": {"source": "kb"}} + assert not _message_sequence([retriever_document]) + assert _message_content(retriever_document) is None assert not _message_sequence([{"other": "field"}]) assert _normalize_content_items("not a sequence") is None assert _normalize_content_items([{"invalid": True}]) is None From a47ece4298eeec7bfc7511102b269325d77cea86 Mon Sep 17 00:00:00 2001 From: Namrata Ghadi Date: Thu, 1 Oct 2026 10:11:42 -0700 Subject: [PATCH 7/7] fix the CI --- .../galileo/tests/test_records_factory.py | 19 ++++++++++++++++--- 1 file changed, 16 insertions(+), 3 deletions(-) diff --git a/evaluators/contrib/galileo/tests/test_records_factory.py b/evaluators/contrib/galileo/tests/test_records_factory.py index 482688f83..f47f4c917 100644 --- a/evaluators/contrib/galileo/tests/test_records_factory.py +++ b/evaluators/contrib/galileo/tests/test_records_factory.py @@ -109,7 +109,7 @@ def test_factory_records_are_serializable_concrete_base_steps() -> None: def test_factory_steps_convert_to_core_records_after_execution_ids_are_added() -> None: - # Given: each root step subtype emitted by the factory, plus the IDs the execution caller supplies + # Given: each root subtype emitted by the factory, plus caller-supplied IDs steps = [ record_from_step(Step(type="llm", name="answer", input="question", output="answer")), record_from_step(Step(type="tool", name="search", input={"query": "q"}, output="result")), @@ -134,7 +134,16 @@ def test_factory_steps_convert_to_core_records_after_execution_ids_are_added() - step_id = uuid4() values = step.model_dump( mode="python", - exclude={"id", "project_id", "run_id", "session_id", "trace_id", "parent_id", "spans", "traces"}, + exclude={ + "id", + "project_id", + "run_id", + "session_id", + "trace_id", + "parent_id", + "spans", + "traces", + }, ) values.update(id=step_id, project_id=project_id, run_id=run_id, type=step.type) if isinstance(step, Session): @@ -774,7 +783,11 @@ class StructuredState(Enum): assert encoder.default(NestedModel(value=3)) == {"value": 3} assert encoder.default(utc_timestamp) == "2025-01-02T00:00:00Z" - assert encoder.default(naive_timestamp) == naive_timestamp.astimezone().isoformat() + expected_naive_timestamp = naive_timestamp.astimezone() + expected_naive_serialized = expected_naive_timestamp.isoformat() + if expected_naive_timestamp.tzname() == UTC.tzname(None): + expected_naive_serialized = expected_naive_serialized.replace("+00:00", "Z") + assert encoder.default(naive_timestamp) == expected_naive_serialized assert encoder.default(offset_timestamp) == "2025-01-02T00:00:00+02:00" assert encoder.default(date(2025, 1, 2)) == "2025-01-02" assert encoder.default(identifier) == str(identifier)