ParthKulshreshtha's picture
Deploy personalized ad environment
526cf2e verified
Raw
History Blame Contribute Delete
10.7 kB
"""Core contracts for the phase-one personalized product-ad environment."""
from __future__ import annotations
from dataclasses import dataclass
from enum import Enum
from pathlib import Path
from typing import Any
HEADLINE_MAX_CHARS = 60
BODY_MAX_CHARS = 180
CTA_ALLOWLIST = frozenset({"Shop now", "Explore the style", "View details"})
class ValidationError(ValueError):
"""A domain value cannot be accepted safely."""
class DataProvenance(str, Enum):
HM_LOCAL = "hm_local"
HM_PILOT = "hm_pilot"
SYNTHETIC_DEMO = "synthetic_demo"
class ExecutionMode(str, Enum):
LIVE = "live"
RECORDED_REPLAY = "recorded_replay"
DETERMINISTIC_TEST = "deterministic_test"
class ActionSource(str, Enum):
EXTERNAL_AGENT = "external_agent"
LIVE_GENERATOR = "live_generator"
RECORDED_REPLAY = "recorded_replay"
DETERMINISTIC_TEST = "deterministic_test"
class ReviewStatus(str, Enum):
ACCEPTED = "accepted"
NEEDS_REVIEW = "needs_review"
FAILED = "failed"
class StageState(str, Enum):
STARTED = "started"
COMPLETED = "completed"
def _clean_required(value: str, field: str) -> str:
if not isinstance(value, str):
raise ValidationError(f"{field} must be a string")
cleaned = " ".join(value.split())
if not cleaned:
raise ValidationError(f"{field} must not be empty")
return cleaned
def _clean_tuple(values: tuple[str, ...], field: str) -> tuple[str, ...]:
if not isinstance(values, tuple):
raise ValidationError(f"{field} must be a tuple")
cleaned = tuple(_clean_required(value, field) for value in values)
if not cleaned:
raise ValidationError(f"{field} must not be empty")
return cleaned
@dataclass(frozen=True, slots=True)
class CustomerContext:
top_product_groups: tuple[str, ...]
top_product_types: tuple[str, ...]
top_colours: tuple[str, ...]
recent_purchase_summary: tuple[str, ...]
def __post_init__(self) -> None:
for field in (
"top_product_groups",
"top_product_types",
"top_colours",
"recent_purchase_summary",
):
object.__setattr__(self, field, _clean_tuple(getattr(self, field), field))
@dataclass(frozen=True, slots=True)
class Product:
name: str
product_type: str
colour: str
description: str
image_id: str
def __post_init__(self) -> None:
for field in ("name", "product_type", "colour", "description", "image_id"):
object.__setattr__(self, field, _clean_required(getattr(self, field), field))
@dataclass(frozen=True, slots=True)
class GraderTarget:
expected_intent: str
personalization_target: str
acceptable_realizations: tuple[str, ...]
version: str = "draft-v1"
def __post_init__(self) -> None:
object.__setattr__(
self, "expected_intent", _clean_required(self.expected_intent, "expected_intent")
)
object.__setattr__(
self,
"personalization_target",
_clean_required(self.personalization_target, "personalization_target"),
)
object.__setattr__(
self,
"acceptable_realizations",
_clean_tuple(self.acceptable_realizations, "acceptable_realizations"),
)
object.__setattr__(self, "version", _clean_required(self.version, "version"))
@dataclass(frozen=True, slots=True)
class CheckResult:
check_id: str
passed: bool
explanation: str
def __post_init__(self) -> None:
object.__setattr__(self, "check_id", _clean_required(self.check_id, "check_id"))
if not isinstance(self.passed, bool):
raise ValidationError("passed must be a bool")
object.__setattr__(
self, "explanation", _clean_required(self.explanation, "explanation")
)
@dataclass(frozen=True, slots=True)
class Observation:
customer_context: CustomerContext
query: str
product: Product
def __post_init__(self) -> None:
if not isinstance(self.customer_context, CustomerContext):
raise ValidationError("customer_context must be a CustomerContext")
if not isinstance(self.product, Product):
raise ValidationError("product must be a Product")
object.__setattr__(self, "query", _clean_required(self.query, "query"))
@dataclass(frozen=True, slots=True)
class Scenario:
"""Internal scenario record. `image_path` and `grader_target` are never observations."""
scenario_id: str
customer_context: CustomerContext
query: str
product: Product
image_path: Path
provenance: DataProvenance
grader_target: GraderTarget | None = None
def __post_init__(self) -> None:
object.__setattr__(self, "scenario_id", _clean_required(self.scenario_id, "scenario_id"))
object.__setattr__(self, "query", _clean_required(self.query, "query"))
if not isinstance(self.customer_context, CustomerContext):
raise ValidationError("customer_context must be a CustomerContext")
if not isinstance(self.product, Product):
raise ValidationError("product must be a Product")
if not isinstance(self.image_path, Path):
raise ValidationError("image_path must be a Path")
if not isinstance(self.provenance, DataProvenance):
raise ValidationError("provenance must be a DataProvenance")
if self.grader_target is not None and not isinstance(self.grader_target, GraderTarget):
raise ValidationError("grader_target must be a GraderTarget or None")
def to_observation(self) -> Observation:
return Observation(
customer_context=self.customer_context,
query=self.query,
product=self.product,
)
@dataclass(frozen=True, slots=True)
class AdCopy:
headline: str
body: str
cta: str | None = None
def __post_init__(self) -> None:
headline = _clean_required(self.headline, "headline")
body = _clean_required(self.body, "body")
if len(headline) > HEADLINE_MAX_CHARS:
raise ValidationError(f"headline must be at most {HEADLINE_MAX_CHARS} characters")
if len(body) > BODY_MAX_CHARS:
raise ValidationError(f"body must be at most {BODY_MAX_CHARS} characters")
cta = self.cta
if cta == "":
cta = None
elif cta is not None:
cta = _clean_required(cta, "cta")
if cta not in CTA_ALLOWLIST:
raise ValidationError(f"cta must be omitted or one of {sorted(CTA_ALLOWLIST)}")
object.__setattr__(self, "headline", headline)
object.__setattr__(self, "body", body)
object.__setattr__(self, "cta", cta)
@dataclass(frozen=True, slots=True)
class ActionProvenance:
"""Public-safe description of where a submitted copy action came from."""
source: ActionSource
execution_mode: ExecutionMode
identity: dict[str, str] | None = None
def __post_init__(self) -> None:
if not isinstance(self.source, ActionSource):
raise ValidationError("action provenance source must be an ActionSource")
if not isinstance(self.execution_mode, ExecutionMode):
raise ValidationError("action provenance execution_mode must be an ExecutionMode")
if self.identity is None:
return
if not isinstance(self.identity, dict) or not self.identity:
raise ValidationError("action provenance identity must be a non-empty mapping")
if len(self.identity) > 8:
raise ValidationError("action provenance identity has too many fields")
cleaned: dict[str, str] = {}
for key, value in self.identity.items():
clean_key = _clean_required(key, "action provenance identity key")
clean_value = _clean_required(value, "action provenance identity value")
if len(clean_key) > 64 or len(clean_value) > 200:
raise ValidationError("action provenance identity fields are too long")
cleaned[clean_key] = clean_value
object.__setattr__(self, "identity", cleaned)
@dataclass(frozen=True, slots=True)
class SafeStageEvent:
"""Public progress notification for an actual environment boundary."""
stage_id: str
label: str
state: StageState
sequence: int
public_data: dict[str, Any] | None = None
def __post_init__(self) -> None:
object.__setattr__(self, "stage_id", _clean_required(self.stage_id, "stage_id"))
object.__setattr__(self, "label", _clean_required(self.label, "label"))
if not isinstance(self.state, StageState):
raise ValidationError("stage state must be a StageState")
if not isinstance(self.sequence, int) or self.sequence < 1:
raise ValidationError("stage sequence must be a positive integer")
if self.public_data is not None and not isinstance(self.public_data, dict):
raise ValidationError("stage public_data must be a mapping or None")
@dataclass(frozen=True, slots=True)
class EpisodeResult:
"""Public-safe result envelope for one complete episode."""
scenario_id: str
action: AdCopy
data_provenance: DataProvenance
action_provenance: ActionProvenance
judge_execution_mode: ExecutionMode
safe_stages: tuple[str, ...] = ()
reward_policy_version: str | None = None
reward: float | None = None
review_status: ReviewStatus | None = None
card_artifact: str | None = None
checks: tuple[CheckResult, ...] = ()
judge_scores: dict[str, float] | None = None
judge_explanations: dict[str, str] | None = None
weighted_components: dict[str, float] | None = None
base_score: float | None = None
failed_checks: tuple[str, ...] = ()
applied_cap: float | None = None
judge_identity: dict[str, str] | None = None
def __post_init__(self) -> None:
object.__setattr__(self, "scenario_id", _clean_required(self.scenario_id, "scenario_id"))
if not isinstance(self.action, AdCopy):
raise ValidationError("action must be an AdCopy")
if not isinstance(self.data_provenance, DataProvenance):
raise ValidationError("data_provenance must be a DataProvenance")
if not isinstance(self.action_provenance, ActionProvenance):
raise ValidationError("action_provenance must be an ActionProvenance")
if not isinstance(self.judge_execution_mode, ExecutionMode):
raise ValidationError("judge_execution_mode must be an ExecutionMode")
if any(not isinstance(check, CheckResult) for check in self.checks):
raise ValidationError("checks must contain CheckResult values")