CoolFace
Apppublic

DGXAI/driftcall

sourceHugging Faceapache-2.0updated 5mo agoView on Hugging Face
0likes
step_10_env.py1020 linesDownload Raw Back to cells
1"""Cell 10 — DriftCallEnv integration class.2 3Implements ``docs/modules/env.md`` and DESIGN.md §4. ``DriftCallEnv`` is the4single public surface that composes models, vendors, drift_injector,5task_generator, rewards, and the optional audio boundary into an6OpenEnv-compliant RL environment.7 8Hard rules (env.md §3.8, CLAUDE.md §0):9- All public dataclasses are frozen.10- State transitions go through ``dataclasses.replace``; no in-place mutation.11- Validation is pure: ``InvalidActionError`` raises BEFORE any state mutation.12- Rewards are computed exactly once at termination and memoized.13- No LLM judge anywhere; no network/disk I/O at ``__init__``.14"""15 16from __future__ import annotations17 18import os19import struct20import uuid21from dataclasses import dataclass, field, replace22from datetime import datetime, timedelta, timezone23from typing import TYPE_CHECKING, Any, Literal, Protocol, cast24 25from cells.step_04_models import (26    ActionType,27    DriftCallAction,28    DriftCallObservation,29    DriftCallState,30    DriftEvent,31    GoalSpec,32    ToolResult,33)34from cells.step_05_vendors import TOOLS as VENDOR_TOOLS35from cells.step_05_vendors import VENDOR_REGISTRY36from cells.step_06_drift_injector import (37    DriftCatalogueError,38    DriftDomainMismatchError,39    DriftReapplicationError,40    DriftScheduleConflictError,41    UnknownDriftPatternError,42    apply_drift,43    build_schedule,44    list_patterns,45)46from cells.step_07_task_generator import (47    InvalidLanguageWeightError,48    InvalidStageError,49)50from cells.step_07_task_generator import (51    generate as task_generate,52)53 54if TYPE_CHECKING:55    from collections.abc import Mapping56 57# rewards is imported lazily inside _compute_rewards to keep the env importable58# even before step_08_rewards.py lands; failures surface as RewardComputationError.59 60_DEFAULT_LANGUAGE_WEIGHTS: dict[str, float] = {61    "en": 0.4,62    "hinglish": 0.4,63    "hi": 0.1,64    "ta": 0.05,65    "kn": 0.05,66}67 68_LANGUAGE_CODES: frozenset[str] = frozenset({"hi", "ta", "kn", "en", "hinglish"})69 70_STAGE_MAX_TURNS: dict[int, int] = {1: 8, 2: 12, 3: 16}71 72_VENDOR_DOMAINS: tuple[str, ...] = ("airline", "cab", "restaurant", "hotel", "payment")73 74_TERMINATED_VALUES: frozenset[str] = frozenset({"SUBMIT", "ABORT", "TIMEOUT", "ANTI_HACK"})75 76_NOW_IST: datetime = datetime(2026, 4, 25, 10, 0, tzinfo=timezone(timedelta(hours=5, minutes=30)))77 78 79# ---------------------------------------------------------------------------80# Error taxonomy (env.md §5)81# ---------------------------------------------------------------------------82 83 84class DriftCallEnvError(Exception):85    """Root for every typed env error (env.md §5)."""86 87 88class InvalidConfigError(DriftCallEnvError):89    """E1 — malformed config dict."""90 91 92class EnvNotReadyError(DriftCallEnvError):93    """E2 — operation issued before reset()."""94 95 96class EnvClosedError(DriftCallEnvError):97    """E3 — operation issued after close()."""98 99 100class InvalidActionError(DriftCallEnvError):101    """E4 — action fails the per-ActionType field matrix."""102 103 104class EpisodeAlreadyTerminalError(DriftCallEnvError):105    """E5 — step() called after termination."""106 107 108class EpisodeNotTerminalError(DriftCallEnvError):109    """E6 — episode()/rewards() called before termination."""110 111 112class ConcurrentStepError(DriftCallEnvError):113    """E7 — reentrant step() detected."""114 115 116class UnknownDomainError(DriftCallEnvError):117    """E8 — PROBE_SCHEMA on a domain that is not registered."""118 119 120class UnknownToolError(DriftCallEnvError):121    """E9 — TOOL_CALL with a tool_name not in available_tools()."""122 123 124class DriftInjectionError(DriftCallEnvError):125    """E10 — drift fold raised; surfaced as-is."""126 127 128class RewardComputationError(DriftCallEnvError):129    """E11 — compute_rewards raised; surfaced as-is."""130 131 132class AudioPipelineError(DriftCallEnvError):133    """E12 — TTS/ASR engine raised on a step()/reset() boundary."""134 135 136_ALL_ERROR_CLASSES: tuple[type[DriftCallEnvError], ...] = (137    InvalidConfigError,138    EnvNotReadyError,139    EnvClosedError,140    InvalidActionError,141    EpisodeAlreadyTerminalError,142    EpisodeNotTerminalError,143    ConcurrentStepError,144    UnknownDomainError,145    UnknownToolError,146    DriftInjectionError,147    RewardComputationError,148    AudioPipelineError,149)150 151 152# ---------------------------------------------------------------------------153# Protocols (env.md §2.1)154# ---------------------------------------------------------------------------155 156 157class DriftScheduler(Protocol):158    def __call__(159        self, stage: int, episode_seed: int, goal: GoalSpec160    ) -> tuple[DriftEvent, ...]: ...161 162 163class TTSEngine(Protocol):164    def synthesize(165        self,166        text: str,167        language_code: str,168        voice_pack: Any | None = None,169        *,170        seed: int = 0,171        sample_rate_hz: int = 16000,172    ) -> bytes: ...173 174 175class ASREngine(Protocol):176    def transcribe(177        self,178        audio_bytes: bytes,179        language_hint: str | None,180        *,181        beam_size: int = 1,182        vad_filter: bool = True,183        max_duration_s: float = 30.0,184    ) -> Any: ...185 186 187def _default_scheduler(188    stage: int, episode_seed: int, goal: GoalSpec189) -> tuple[DriftEvent, ...]:190    return build_schedule(stage, episode_seed, goal)191 192 193# ---------------------------------------------------------------------------194# Episode (env.md §4.3) — built at termination, fed to rewards.compute_rewards.195# Matches the Episode shape consumed by step_08_rewards (kw fields).196# ---------------------------------------------------------------------------197 198 199@dataclass(frozen=True)200class Episode:201    episode_id: str202    goal: GoalSpec203    actions: tuple[DriftCallAction, ...]204    action_turns: tuple[int, ...]205    tool_results: tuple[ToolResult, ...]206    tool_result_turns: tuple[int, ...]207    drift_log: tuple[DriftEvent, ...]208    vendor_states_final: dict[str, dict[str, Any]]209    schema_versions_final: dict[str, str]210    max_turns: int211    turns_used: int212    terminated_by: Literal["SUBMIT", "ABORT", "TIMEOUT", "ANTI_HACK"]213    stage: Literal[1, 2, 3]214    drift_pattern_overrides: dict[str, Any] = field(default_factory=dict)215 216 217# ---------------------------------------------------------------------------218# EnvConfig (env.md §4.1)219# ---------------------------------------------------------------------------220 221 222@dataclass(frozen=True)223class EnvConfig:224    curriculum_stage: Literal[1, 2, 3]225    language_weights: dict[str, float]226    audio_boundary_enabled: bool227    max_turns_override: int | None228    scheduler: DriftScheduler229    tts_engine: TTSEngine | None230    asr_engine: ASREngine | None231 232    @classmethod233    def from_mapping(cls, raw: Mapping[str, Any] | None) -> EnvConfig:234        allowed = {235            "curriculum_stage",236            "language_weights",237            "audio_boundary_enabled",238            "max_turns_override",239            "scheduler",240            "tts_engine",241            "asr_engine",242        }243        if raw is None:244            raw = {}245        if not isinstance(raw, dict):246            raise InvalidConfigError(247                f"config must be a dict or None, got {type(raw).__name__}"248            )249 250        unknown = set(raw.keys()) - allowed251        if unknown:252            raise InvalidConfigError(253                f"unknown config key(s): {sorted(unknown)}; "254                f"allowed: {sorted(allowed)}"255            )256 257        stage_raw = raw.get("curriculum_stage", 1)258        if isinstance(stage_raw, bool) or not isinstance(stage_raw, int):259            raise InvalidConfigError(260                f"curriculum_stage must be int in {{1,2,3}}, got "261                f"{type(stage_raw).__name__}"262            )263        if stage_raw not in (1, 2, 3):264            raise InvalidConfigError(265                f"curriculum_stage must be 1, 2, or 3; got {stage_raw!r}"266            )267        stage = cast("Literal[1, 2, 3]", stage_raw)268 269        weights_raw = raw.get("language_weights", _DEFAULT_LANGUAGE_WEIGHTS)270        if not isinstance(weights_raw, dict) or not weights_raw:271            raise InvalidConfigError(272                "language_weights must be a non-empty dict"273            )274        for k, v in weights_raw.items():275            if k not in _LANGUAGE_CODES:276                raise InvalidConfigError(277                    f"language_weights: unknown language {k!r}; "278                    f"allowed: {sorted(_LANGUAGE_CODES)}"279                )280            if isinstance(v, bool) or not isinstance(v, (int, float)):281                raise InvalidConfigError(282                    f"language_weights[{k!r}] must be numeric, got "283                    f"{type(v).__name__}"284                )285            if v < 0:286                raise InvalidConfigError(287                    f"language_weights[{k!r}]={v} is negative"288                )289        total = sum(float(v) for v in weights_raw.values())290        if abs(total - 1.0) > 1e-6:291            raise InvalidConfigError(292                f"language_weights sum {total!r} not within 1.0 ± 1e-6"293            )294        # Frozen copy.295        weights: dict[str, float] = {k: float(v) for k, v in weights_raw.items()}296 297        audio_enabled_raw = raw.get("audio_boundary_enabled", False)298        if not isinstance(audio_enabled_raw, bool):299            raise InvalidConfigError(300                f"audio_boundary_enabled must be bool, got "301                f"{type(audio_enabled_raw).__name__}"302            )303        audio_enabled = audio_enabled_raw304 305        max_turns_override = raw.get("max_turns_override")306        if max_turns_override is not None:307            if isinstance(max_turns_override, bool) or not isinstance(308                max_turns_override, int309            ):310                raise InvalidConfigError(311                    f"max_turns_override must be int or None, got "312                    f"{type(max_turns_override).__name__}"313                )314            if max_turns_override < 1:315                raise InvalidConfigError(316                    f"max_turns_override must be >= 1, got {max_turns_override}"317                )318 319        scheduler = raw.get("scheduler", _default_scheduler)320        if not callable(scheduler):321            raise InvalidConfigError("scheduler must be callable")322 323        tts_engine = raw.get("tts_engine")324        asr_engine = raw.get("asr_engine")325 326        if audio_enabled:327            if tts_engine is None:328                raise InvalidConfigError(329                    "tts_engine is required when audio_boundary_enabled is True"330                )331            if asr_engine is None:332                raise InvalidConfigError(333                    "asr_engine is required when audio_boundary_enabled is True"334                )335        else:336            if tts_engine is not None:337                raise InvalidConfigError(338                    "tts_engine must be None when audio_boundary_enabled is False"339                )340            if asr_engine is not None:341                raise InvalidConfigError(342                    "asr_engine must be None when audio_boundary_enabled is False"343                )344 345        return cls(346            curriculum_stage=stage,347            language_weights=weights,348            audio_boundary_enabled=audio_enabled,349            max_turns_override=max_turns_override,350            scheduler=cast("DriftScheduler", scheduler),351            tts_engine=cast("TTSEngine | None", tts_engine),352            asr_engine=cast("ASREngine | None", asr_engine),353        )354 355 356# ---------------------------------------------------------------------------357# DriftCallEnv358# ---------------------------------------------------------------------------359 360 361def _make_seed_from_urandom() -> int:362    raw = os.urandom(8)363    (value,) = struct.unpack("<Q", raw)364    return int(value)365 366 367def _vendor_state_to_dict(state: Any) -> dict[str, Any]:368    """Coerce a frozen vendor dataclass (or already-dict) into a plain dict."""369    if isinstance(state, dict):370        return dict(state)371    # All vendor states are frozen dataclasses.372    import dataclasses as _dc373 374    if _dc.is_dataclass(state) and not isinstance(state, type):375        return _dc.asdict(state)376    # Defensive: best-effort fallback.377    return {"_raw": repr(state)}378 379 380class DriftCallEnv:381    """OpenEnv-compliant RL environment for DriftCall (env.md §1)."""382 383    # -- construction --------------------------------------------------------384 385    def __init__(self, config: dict[str, Any] | None = None) -> None:386        self._config: EnvConfig = EnvConfig.from_mapping(config)387        self._state: DriftCallState | None = None388        self._rewards: Any | None = None389        self._episode: Episode | None = None390        self._closed: bool = False391        self._seed: int | None = None392        self._episode_id: str | None = None393        # Pending side-channel notices keyed by domain (env.md §3.3).394        self._side_channel_pending: dict[str, str] = {}395        # Per-vendor-state cache (frozen dataclass or dict). Kept on the env396        # because DriftCallState.vendor_states is a dict[str, dict] for397        # compatibility with the design dataclass.398        self._vendor_state_objects: dict[str, Any] = {}399        # Re-entrancy guard (E7).400        self._step_in_progress: bool = False401 402    # -- internal helpers ----------------------------------------------------403 404    @property405    def _max_turns(self) -> int:406        if self._config.max_turns_override is not None:407            return int(self._config.max_turns_override)408        return _STAGE_MAX_TURNS[self._config.curriculum_stage]409 410    def _available_tools(self) -> tuple[str, ...]:411        return VENDOR_TOOLS412 413    def _ensure_ready_for_step(self) -> None:414        if self._closed:415            raise EnvClosedError("env is closed")416        if self._state is None:417            raise EnvNotReadyError("reset() must be called before step()")418        if self._state.done:419            raise EpisodeAlreadyTerminalError(420                f"episode already terminated (terminated_by={self._terminated_by()})"421            )422 423    def _terminated_by(self) -> str | None:424        return self._episode.terminated_by if self._episode is not None else None425 426    # -- OpenEnv primitives --------------------------------------------------427 428    def reset(self, seed: int | None = None) -> DriftCallObservation:429        if self._closed:430            raise EnvClosedError("env is closed")431 432        if seed is None:433            seed = _make_seed_from_urandom()434        if isinstance(seed, bool) or not isinstance(seed, int):435            raise InvalidActionError(436                f"seed must be int or None, got {type(seed).__name__}"437            )438 439        self._seed = int(seed)440        # Reset memoization; legacy state is dropped before any propagatable441        # exception can leak (env.md §2.2 docstring).442        self._state = None443        self._rewards = None444        self._episode = None445        self._side_channel_pending = {}446        self._vendor_state_objects = {}447        self._episode_id = None448 449        try:450            goal = task_generate(451                self._seed,452                self._config.curriculum_stage,453                cast("dict[Any, float]", self._config.language_weights),454            )455        except (InvalidLanguageWeightError, InvalidStageError) as exc:456            # E1-class reset failure (env.md §2.2 raises clause).457            raise InvalidConfigError(str(exc)) from exc458 459        # Initial per-domain vendor state objects (frozen dataclasses).460        vendor_state_objects: dict[str, Any] = {}461        vendor_states_dict: dict[str, dict[str, Any]] = {}462        for domain in _VENDOR_DOMAINS:463            ns = VENDOR_REGISTRY[domain]464            vs = ns.initial_state(self._seed, goal)465            vendor_state_objects[domain] = vs466            vendor_states_dict[domain] = _vendor_state_to_dict(vs)467 468        schema_versions = {d: "v1" for d in _VENDOR_DOMAINS}469 470        try:471            schedule = self._config.scheduler(472                self._config.curriculum_stage, self._seed, goal473            )474        except (475            DriftScheduleConflictError,476            DriftCatalogueError,477            UnknownDriftPatternError,478            DriftDomainMismatchError,479        ) as exc:480            # Bad scheduler at reset is an E1 (env.md §7.4).481            raise InvalidConfigError(f"scheduler failure: {exc}") from exc482 483        self._episode_id = uuid.uuid4().hex484 485        max_turns = self._max_turns486        new_state = DriftCallState(487            episode_id=self._episode_id,488            goal=goal,489            vendor_states=vendor_states_dict,490            schema_versions=schema_versions,491            drift_schedule=tuple(schedule),492            drift_fired=(),493            turn=0,494            max_turns=max_turns,495            actions=(),496            done=False,497        )498        self._state = new_state499        self._vendor_state_objects = vendor_state_objects500 501        if self._config.audio_boundary_enabled:502            tts = self._config.tts_engine503            assert tts is not None  # validated in EnvConfig504            try:505                tts.synthesize(goal.seed_utterance, goal.language)506            except Exception as exc:  # noqa: BLE001 — surface as E12-class507                # Audio failure on reset leaves env unready (env.md §5 E12).508                self._state = None509                self._vendor_state_objects = {}510                self._episode_id = None511                raise AudioPipelineError(f"TTS reset failure: {exc}") from exc512 513        return self._build_observation()514 515    def step(516        self,517        action: DriftCallAction,518        *,519        force_drift_pattern: str | None = None,520    ) -> DriftCallObservation:521        # 1a. Pure validation — must raise before any state mutation.522        self._ensure_ready_for_step()523        self._validate_action(action)524        if force_drift_pattern is not None:525            valid_ids = {p.id for p in list_patterns()}526            if force_drift_pattern not in valid_ids:527                raise InvalidActionError(528                    f"force_drift_pattern {force_drift_pattern!r} not a known "529                    f"pattern_id"530                )531 532        if self._step_in_progress:533            raise ConcurrentStepError("reentrant step() detected")534        self._step_in_progress = True535        try:536            return self._step_inner(action, force_drift_pattern)537        finally:538            self._step_in_progress = False539 540    def _step_inner(541        self,542        action: DriftCallAction,543        force_drift_pattern: str | None,544    ) -> DriftCallObservation:545        assert self._state is not None  # ensured above546        # 2. Increment turn counter.547        turn_current = self._state.turn + 1548        self._state = replace(self._state, turn=turn_current)549 550        # 3. Fire drifts for this turn.551        self._fire_drifts(turn_current, force_drift_pattern)552 553        # 4. Side-channel emit pass — refresh pending notices for any vendor554        # whose state just mutated.555        self._emit_side_channel()556 557        # 5. Dispatch action.558        new_tool_result, terminate, terminated_by = self._dispatch(action)559 560        # 6. Record action (and ToolResult, if any) via dataclasses.replace.561        new_actions = self._state.actions + (action,)562        if new_tool_result is not None:563            # Tool result history lives on the state's vendor history; here we564            # rely on the running observation history we will rebuild in §3.4.565            self._tool_results = self._tool_results + (new_tool_result,)566            self._tool_result_turns = self._tool_result_turns + (turn_current,)567        self._action_turns = self._action_turns + (turn_current,)568        self._state = replace(self._state, actions=new_actions)569 570        # 7. Budget check — only if action did not already terminate.571        if not terminate and turn_current >= self._state.max_turns:572            terminate = True573            terminated_by = "TIMEOUT"574 575        # 8. If terminal, build Episode + compute rewards.576        if terminate:577            assert terminated_by is not None578            self._terminate(terminated_by)579 580        # 9. Build observation.581        return self._build_observation()582 583    def state(self) -> DriftCallState:584        if self._state is None:585            raise EnvNotReadyError("reset() must be called before state()")586        return self._state587 588    def close(self) -> None:589        # Idempotent.590        self._closed = True591        # Per env.md §9 Q7: never invoke close on shared audio engines.592        # Only drop per-env state.593        self._side_channel_pending = {}594        self._vendor_state_objects = {}595        # Note: we keep self._state, self._rewards, self._episode so post-close596        # audits still work (env.md §7.11).597 598    def episode(self) -> Episode:599        if self._episode is None:600            raise EpisodeNotTerminalError("episode is not terminal")601        return self._episode602 603    def rewards(self) -> Any:604        if self._rewards is None:605            raise EpisodeNotTerminalError("episode is not terminal")606        return self._rewards607 608    def done(self) -> bool:609        if self._state is None:610            return False611        return bool(self._state.done)612 613    # -- validation ----------------------------------------------------------614 615    def _validate_action(self, action: DriftCallAction) -> None:616        if not isinstance(action, DriftCallAction):617            raise InvalidActionError(618                f"action must be DriftCallAction, got {type(action).__name__}"619            )620        atype = action.action_type621        if not isinstance(atype, ActionType):622            raise InvalidActionError(623                f"action_type must be ActionType, got {type(atype).__name__}"624            )625 626        # rationale length cap (env.md §3.1).627        if action.rationale is not None and len(action.rationale) > 200:628            raise InvalidActionError(629                f"rationale length {len(action.rationale)} exceeds 200"630            )631 632        if atype == ActionType.TOOL_CALL:633            if not action.tool_name or not isinstance(action.tool_name, str):634                raise InvalidActionError("TOOL_CALL requires non-empty tool_name")635            if action.tool_args is None or not isinstance(action.tool_args, dict):636                raise InvalidActionError(637                    "TOOL_CALL requires tool_args dict (may be empty)"638                )639            if action.message is not None or action.confidence is not None:640                raise InvalidActionError(641                    "TOOL_CALL forbids message/confidence"642                )643            if action.tool_name not in self._available_tools():644                raise UnknownToolError(645                    f"tool_name {action.tool_name!r} not in available_tools()"646                )647            # JSON-serializability (shallow check: must be dict; values arbitrary).648            return649 650        if atype == ActionType.SPEAK or atype == ActionType.CLARIFY:651            if not isinstance(action.message, str):652                raise InvalidActionError(653                    f"{atype.value} requires str message"654                )655            if not (1 <= len(action.message) <= 2000):656                raise InvalidActionError(657                    f"{atype.value} message length must be in [1, 2000], "658                    f"got {len(action.message)}"659                )660            if "\x00" in action.message:661                raise InvalidActionError(662                    f"{atype.value} message contains NUL byte"663                )664            if (665                action.tool_name is not None666                or action.tool_args is not None667                or action.confidence is not None668            ):669                raise InvalidActionError(670                    f"{atype.value} forbids tool_name/tool_args/confidence"671                )672            return673 674        if atype == ActionType.PROBE_SCHEMA:675            if not action.tool_name or not isinstance(action.tool_name, str):676                raise InvalidActionError(677                    "PROBE_SCHEMA requires tool_name (domain string)"678                )679            if (680                action.tool_args is not None681                or action.message is not None682                or action.confidence is not None683            ):684                raise InvalidActionError(685                    "PROBE_SCHEMA forbids tool_args/message/confidence"686                )687            assert self._state is not None688            if action.tool_name not in self._state.vendor_states:689                raise UnknownDomainError(690                    f"PROBE_SCHEMA: domain {action.tool_name!r} not registered"691                )692            return693 694        if atype == ActionType.SUBMIT:695            if action.confidence is None or not isinstance(696                action.confidence, (int, float)697            ) or isinstance(action.confidence, bool):698                raise InvalidActionError("SUBMIT requires float confidence")699            conf = float(action.confidence)700            if not (0.0 <= conf <= 1.0):701                raise InvalidActionError(702                    f"SUBMIT confidence {conf!r} outside [0.0, 1.0]"703                )704            if action.tool_name is not None or action.tool_args is not None:705                raise InvalidActionError(706                    "SUBMIT forbids tool_name/tool_args"707                )708            if action.message is not None and not isinstance(action.message, str):709                raise InvalidActionError("SUBMIT message must be str if present")710            return711 712        if atype == ActionType.ABORT:713            if (714                action.tool_name is not None715                or action.tool_args is not None716                or action.confidence is not None717            ):718                raise InvalidActionError(719                    "ABORT forbids tool_name/tool_args/confidence"720                )721            return722 723        # Unreachable — all six ActionType members handled above.724        raise InvalidActionError(f"unhandled action_type {atype!r}")725 726    # -- drift firing --------------------------------------------------------727 728    def _fire_drifts(self, turn_current: int, force_pattern: str | None) -> None:729        assert self._state is not None730        if force_pattern is not None:731            patterns_by_id = {p.id: p for p in list_patterns()}732            pattern = patterns_by_id[force_pattern]733            if pattern.domain not in self._state.vendor_states:734                raise DriftInjectionError(735                    f"force_drift_pattern {force_pattern!r}: domain "736                    f"{pattern.domain!r} not registered"737                )738            event = DriftEvent(739                turn=turn_current,740                drift_type=pattern.drift_type,741                domain=pattern.domain,742                description=pattern.description,743                from_version=pattern.from_version,744                to_version=pattern.to_version,745                pattern_id=pattern.id,746            )747            try:748                self._state = apply_drift(self._state, event)749            except (750                UnknownDriftPatternError,751                DriftDomainMismatchError,752                DriftReapplicationError,753            ) as exc:754                raise DriftInjectionError(str(exc)) from exc755            return756 757        # Schedule-driven fold.758        pending = tuple(759            e for e in self._state.drift_schedule760            if e.turn == turn_current and e not in self._state.drift_fired761        )762        if not pending:763            return764        ordered = tuple(sorted(pending, key=lambda e: (e.turn, e.pattern_id)))765        for event in ordered:766            try:767                self._state = apply_drift(self._state, event)768            except (769                UnknownDriftPatternError,770                DriftDomainMismatchError,771                DriftReapplicationError,772            ) as exc:773                raise DriftInjectionError(str(exc)) from exc774 775    def _emit_side_channel(self) -> None:776        """Refresh pending side-channel notices per env.md §3.3 clause 3."""777        assert self._state is not None778        new_pending = dict(self._side_channel_pending)779        for domain in _VENDOR_DOMAINS:780            ns = VENDOR_REGISTRY[domain]781            vs_obj = self._vendor_state_objects.get(domain)782            if vs_obj is None:783                continue784            try:785                notice, new_state = ns.emit_side_channel_if_pending(vs_obj)786            except Exception as exc:  # noqa: BLE001 — defensive787                raise DriftInjectionError(788                    f"side-channel emit failed for {domain}: {exc}"789                ) from exc790            if notice is not None:791                existing = new_pending.get(domain)792                merged = (793                    f"{existing}\n---\n{notice}" if existing else notice794                )795                new_pending[domain] = merged796            self._vendor_state_objects[domain] = new_state797        self._side_channel_pending = new_pending798 799    # -- dispatch ------------------------------------------------------------800 801    @property802    def _tool_results(self) -> tuple[ToolResult, ...]:803        return getattr(self, "_tool_results_internal", ())804 805    @_tool_results.setter806    def _tool_results(self, value: tuple[ToolResult, ...]) -> None:807        self._tool_results_internal = value808 809    @property810    def _tool_result_turns(self) -> tuple[int, ...]:811        return getattr(self, "_tool_result_turns_internal", ())812 813    @_tool_result_turns.setter814    def _tool_result_turns(self, value: tuple[int, ...]) -> None:815        self._tool_result_turns_internal = value816 817    @property818    def _action_turns(self) -> tuple[int, ...]:819        return getattr(self, "_action_turns_internal", ())820 821    @_action_turns.setter822    def _action_turns(self, value: tuple[int, ...]) -> None:823        self._action_turns_internal = value824 825    def _dispatch(826        self, action: DriftCallAction827    ) -> tuple[ToolResult | None, bool, str | None]:828        """Return (tool_result, terminate?, terminated_by?)."""829        assert self._state is not None830        atype = action.action_type831 832        if atype == ActionType.SUBMIT:833            return None, True, "SUBMIT"834        if atype == ActionType.ABORT:835            return None, True, "ABORT"836        if atype == ActionType.SPEAK or atype == ActionType.CLARIFY:837            return None, False, None838 839        if atype == ActionType.PROBE_SCHEMA:840            assert action.tool_name is not None841            domain = action.tool_name842            ns = VENDOR_REGISTRY[domain]843            vs_obj = self._vendor_state_objects[domain]844            schema_version = self._state.schema_versions[domain]845            schema = ns.describe_schema(vs_obj, schema_version)846            tr = ToolResult(847                tool_name=f"probe:{domain}",848                status="ok",849                response=dict(schema),850                schema_version=schema_version,851                latency_ms=0,852            )853            return tr, False, None854 855        if atype == ActionType.TOOL_CALL:856            assert action.tool_name is not None and action.tool_args is not None857            tool_name = action.tool_name858            domain = tool_name.split(".", 1)[0]859            if domain not in self._state.vendor_states:860                raise UnknownDomainError(861                    f"tool {tool_name!r} targets unknown domain {domain!r}"862                )863            ns = VENDOR_REGISTRY[domain]864            vs_obj = self._vendor_state_objects[domain]865            schema_version = self._state.schema_versions[domain]866            try:867                if domain == "payment":868                    tr, new_vs = ns.dispatch(869                        tool_name,870                        action.tool_args,871                        vs_obj,872                        schema_version,873                        self._seed,874                        _NOW_IST,875                    )876                    payment_state = new_vs877                else:878                    payment_state = self._vendor_state_objects.get("payment")879                    tr, new_vs, payment_state = ns.dispatch(880                        tool_name,881                        action.tool_args,882                        vs_obj,883                        schema_version,884                        self._seed,885                        _NOW_IST,886                        payment_state,887                    )888            except ValueError as exc:889                # Unknown tool inside a known domain → treat as anti-hack.890                raise UnknownToolError(str(exc)) from exc891 892            self._vendor_state_objects[domain] = new_vs893            if payment_state is not None:894                self._vendor_state_objects["payment"] = payment_state895 896            # Refresh state.vendor_states snapshot.897            new_vendor_states = dict(self._state.vendor_states)898            new_vendor_states[domain] = _vendor_state_to_dict(new_vs)899            if domain != "payment" and payment_state is not None:900                new_vendor_states["payment"] = _vendor_state_to_dict(payment_state)901            self._state = replace(self._state, vendor_states=new_vendor_states)902 903            # Attach pending side-channel notice (one-shot per domain).904            notice = self._side_channel_pending.pop(domain, None)905            if notice is not None:906                merged_response = dict(tr.response)907                merged_response["_notice"] = notice908                tr = ToolResult(909                    tool_name=tr.tool_name,910                    status=tr.status,911                    response=merged_response,912                    schema_version=tr.schema_version,913                    latency_ms=tr.latency_ms,914                )915            return tr, False, None916 917        # Unreachable.918        raise InvalidActionError(f"unhandled action_type {atype!r}")919 920    # -- termination ---------------------------------------------------------921 922    def _terminate(self, terminated_by: str) -> None:923        assert self._state is not None924        if terminated_by not in _TERMINATED_VALUES:925            raise RewardComputationError(926                f"unknown terminated_by sentinel {terminated_by!r}"927            )928        self._state = replace(self._state, done=True)929        episode = Episode(930            episode_id=self._state.episode_id,931            goal=self._state.goal,932            actions=self._state.actions,933            action_turns=self._action_turns,934            tool_results=self._tool_results,935            tool_result_turns=self._tool_result_turns,936            drift_log=self._state.drift_fired,937            vendor_states_final={938                d: _vendor_state_to_dict(self._vendor_state_objects[d])939                for d in _VENDOR_DOMAINS940            },941            schema_versions_final=dict(self._state.schema_versions),942            max_turns=self._state.max_turns,943            turns_used=len(self._state.actions),944            terminated_by=cast(945                "Literal['SUBMIT','ABORT','TIMEOUT','ANTI_HACK']", terminated_by946            ),947            stage=self._config.curriculum_stage,948        )949        self._episode = episode950        self._rewards = self._compute_rewards(episode)951 952    @staticmethod953    def _compute_rewards(episode: Episode) -> Any:954        import importlib955 956        try:957            mod = importlib.import_module("cells.step_08_rewards")958        except ImportError as exc:959            raise RewardComputationError(960                f"rewards module unavailable: {exc}"961            ) from exc962        compute = getattr(mod, "compute_rewards", None)963        if compute is None:964            raise RewardComputationError(965                "cells.step_08_rewards has no compute_rewards"966            )967        try:968            return compute(episode)969        except Exception as exc:970            raise RewardComputationError(str(exc)) from exc971 972    # -- observation builder -------------------------------------------------973 974    def _build_observation(self) -> DriftCallObservation:975        assert self._state is not None976        st = self._state977        if st.turn == 0:978            last_transcript = st.goal.seed_utterance979            last_lang = st.goal.language980            last_confidence = 1.0981        else:982            last_transcript = st.goal.seed_utterance983            last_lang = st.goal.language984            last_confidence = 1.0985 986        return DriftCallObservation(987            turn=st.turn,988            goal=st.goal,989            last_transcript=last_transcript,990            last_lang=last_lang,991            last_confidence=last_confidence,992            tool_results=self._tool_results,993            drift_log=st.drift_fired,994            budget_remaining=max(0, st.max_turns - st.turn),995            available_tools=self._available_tools(),996        )997 998 999__all__ = [1000    "ASREngine",1001    "AudioPipelineError",1002    "ConcurrentStepError",1003    "DriftCallEnv",1004    "DriftCallEnvError",1005    "DriftInjectionError",1006    "DriftScheduler",1007    "EnvClosedError",1008    "EnvConfig",1009    "EnvNotReadyError",1010    "Episode",1011    "EpisodeAlreadyTerminalError",1012    "EpisodeNotTerminalError",1013    "InvalidActionError",1014    "InvalidConfigError",1015    "RewardComputationError",1016    "TTSEngine",1017    "UnknownDomainError",1018    "UnknownToolError",1019]1020