DGXAI/driftcall
0
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 