""" Flow Execution Engine (Phase 2A). Manages multi-step clinical conversation flows with: - State machine for flow step progression - Slot tracking and validation - Response spec enforcement (must_include, must_not, max_questions) - Primitive sequence execution (greeting → identity_verify → consent) - Domain-specific flow routing from risk assessment - Exit condition evaluation Safety invariants: - Escalation flows CANNOT be interrupted or rolled back - Hard-escalate domains skip screening and go directly to escalate_r3 - If unclear at any step, fail open (escalate, not suppress) - Response specs enforce must_not constraints — these are NEVER relaxed """ from __future__ import annotations import logging from typing import Any, Dict, List, Optional, Set, Tuple from decision.engine.config_loader import DecisionConfigLoader from decision.engine.models import ( CallPhase, ConfidenceTier, ConversationFlow, FlowStep, FlowType, ResponseSpec, RiskAssessment, RiskClass, SessionState, SlotState, TurnOutcome, ) logger = logging.getLogger("decision.flow_engine") # --------------------------------------------------------------------------- # Flow Resolution # --------------------------------------------------------------------------- class FlowResolver: """ Resolves which conversation flow to execute based on the current session state, risk assessment, and domain rules. """ def __init__(self, config: DecisionConfigLoader): self._config = config self._primitives = config.primitives self._taxonomy_flows = config.taxonomy_flows self._taxonomy_rules = config.taxonomy_rules self._primitive_sequence = config.global_rules.get("primitive_sequence", []) self._fallback_flow = config.global_rules.get("fallback_flow", "primitives/clarify") def resolve_initial_flow(self, session: SessionState) -> Optional[FlowDefinition]: """Resolve the first flow to execute when a call starts.""" if self._primitive_sequence: first = self._primitive_sequence[0] return self._load_primitive(first) return None def resolve_next_flow( self, session: SessionState, assessment: Optional[RiskAssessment] = None, current_exit: Optional[str] = None, ) -> Optional[FlowDefinition]: """ Resolve the next flow based on current state and exit condition. Priority order: 1. Hard-escalate → escalate_r3 (immediate, cannot be overridden) 2. R3 assessment → domain escalate_r3 flow 3. R2 assessment → domain escalate_r2 flow 4. Explicit exit rule (on_complete, on_refused, etc.) 5. Primitive sequence (if still in opening) 6. Domain-specific recommended_flow from assessment 7. Fallback → clarify """ # 1. Hard escalate overrides everything if assessment and assessment.hard_escalate: domain = assessment.primary_domain or "suicidal_ideation" flow = self._load_domain_flow(domain, "escalate_r3") if flow: logger.warning("HARD ESCALATE: routing to %s/escalate_r3", domain) return flow # Fallback to generic handoff return self._load_primitive("handoff") # 2-3. Risk-based routing if assessment and assessment.risk_class == RiskClass.R3: domain = assessment.primary_domain if domain: flow = self._load_domain_flow(domain, "escalate_r3") if flow: return flow return self._load_primitive("handoff") if assessment and assessment.risk_class == RiskClass.R2: domain = assessment.primary_domain if domain: flow = self._load_domain_flow(domain, "escalate_r2") if flow: return flow return self._load_primitive("handoff") # 4. Explicit exit rule if current_exit: return self._resolve_exit(current_exit, session) # 5. Check if still in primitive sequence prim_flow = self._advance_primitive_sequence(session) if prim_flow: return prim_flow # 6. Domain recommended flow if assessment and assessment.recommended_flow: domain = assessment.primary_domain if domain: flow = self._load_domain_flow(domain, assessment.recommended_flow) if flow: return flow # 7. Fallback return self._load_primitive("clarify") def _resolve_exit( self, exit_key: str, session: SessionState ) -> Optional[FlowDefinition]: """Resolve an exit rule like 'on_complete', 'on_refused', etc.""" # Special exit targets if exit_key == "end": return None # Session over if exit_key == "re_evaluate": return None # Signal to re-run assessment if exit_key == "evaluate_risk": return None # Signal to run risk evaluation if exit_key == "evaluate_concern": return None # Signal to evaluate a patient concern # Try as primitive name prim = self._load_primitive(exit_key) if prim: return prim # Try as domain flow (format: "domain/flow_id" or just "flow_id") if "/" in exit_key: domain, flow_id = exit_key.split("/", 1) return self._load_domain_flow(domain, flow_id) return None def _advance_primitive_sequence( self, session: SessionState ) -> Optional[FlowDefinition]: """Check if the session still needs to complete primitive sequence.""" if session.call_phase != CallPhase.SESSION_START: return None completed = set(session.domain_history) for prim_name in self._primitive_sequence: if prim_name not in completed: return self._load_primitive(prim_name) # All primitives complete — advance to active call phase return None def _load_primitive(self, name: str) -> Optional[FlowDefinition]: """Load a primitive flow definition by name.""" data = self._primitives.get(name) if not data: return None return FlowDefinition.from_yaml(data, source=f"primitives/{name}") def _load_domain_flow( self, domain: str, flow_id: str ) -> Optional[FlowDefinition]: """Load a domain-specific flow definition.""" domain_flows = self._taxonomy_flows.get(domain, {}) data = domain_flows.get(flow_id) if not data: return None return FlowDefinition.from_yaml(data, source=f"taxonomy/{domain}/{flow_id}") # --------------------------------------------------------------------------- # Flow Definition (parsed from YAML) # --------------------------------------------------------------------------- class FlowDefinition: """A parsed conversation flow with steps and exit rules.""" def __init__( self, flow_id: str, flow_type: str, domain: Optional[str], required_slots: List[str], steps: List[FlowStepDef], exit_rules: Dict[str, str], source: str = "", ): self.flow_id = flow_id self.flow_type = flow_type self.domain = domain self.required_slots = required_slots self.steps = steps self.exit_rules = exit_rules self.source = source @classmethod def from_yaml(cls, data: Dict[str, Any], source: str = "") -> FlowDefinition: """Parse a flow definition from YAML data.""" flow_id = data.get("flow_id", "unknown") flow_type = data.get("type", "primitive") domain = data.get("domain") required_slots = data.get("required_slots", []) exit_rules = data.get("exit", {}) steps = [] for action in data.get("actions", []): step = FlowStepDef.from_yaml(action) steps.append(step) return cls( flow_id=flow_id, flow_type=flow_type, domain=domain, required_slots=required_slots, steps=steps, exit_rules=exit_rules, source=source, ) @property def step_count(self) -> int: return len(self.steps) def get_step(self, index: int) -> Optional[FlowStepDef]: if 0 <= index < len(self.steps): return self.steps[index] return None def __repr__(self) -> str: return f"FlowDefinition({self.flow_id}, type={self.flow_type}, steps={self.step_count})" class FlowStepDef: """A single step in a flow, parsed from YAML.""" def __init__( self, step_number: int, step_type: str, response_spec: ResponseSpecDef, collects: List[str], handoff_target: Optional[str] = None, handoff_urgency: Optional[str] = None, ): self.step_number = step_number self.step_type = step_type self.response_spec = response_spec self.collects = collects self.handoff_target = handoff_target self.handoff_urgency = handoff_urgency @classmethod def from_yaml(cls, data: Dict[str, Any]) -> FlowStepDef: step_number = data.get("step", 0) step_type = data.get("type", "respond") spec_data = data.get("response_spec", {}) response_spec = ResponseSpecDef.from_yaml(spec_data) collects = data.get("collects", []) handoff_target = data.get("target") handoff_urgency = data.get("urgency") return cls( step_number=step_number, step_type=step_type, response_spec=response_spec, collects=collects, handoff_target=handoff_target, handoff_urgency=handoff_urgency, ) class ResponseSpecDef: """LLM response constraints parsed from YAML.""" def __init__( self, spec_id: str, goal: str, tone: str, must_include: List[str], must_ask: List[str], must_not: List[str], max_questions: int, ): self.spec_id = spec_id self.goal = goal self.tone = tone self.must_include = must_include self.must_ask = must_ask self.must_not = must_not self.max_questions = max_questions @classmethod def from_yaml(cls, data: Dict[str, Any]) -> ResponseSpecDef: return cls( spec_id=data.get("id", ""), goal=data.get("goal", ""), tone=data.get("tone", "professional"), must_include=data.get("must_include", []), must_ask=data.get("must_ask", []), must_not=data.get("must_not", []), max_questions=data.get("max_questions", 0), ) def validate_response(self, response_text: str) -> List[str]: """ Validate an LLM-generated response against this spec. Returns list of violations (empty = valid). """ violations = [] response_lower = response_text.lower() # Check must_include for phrase in self.must_include: if phrase.lower() not in response_lower: violations.append(f"MISSING required phrase: '{phrase}'") # Check must_not — SAFETY CRITICAL, never relaxed for phrase in self.must_not: if phrase.lower() in response_lower: violations.append(f"FORBIDDEN phrase found: '{phrase}'") # Check max_questions (count question marks) question_count = response_text.count("?") if question_count > self.max_questions and self.max_questions >= 0: violations.append( f"Too many questions: {question_count} > max {self.max_questions}" ) return violations def to_prompt_constraints(self) -> str: """Generate constraint text for LLM prompt injection.""" lines = [f"Goal: {self.goal}", f"Tone: {self.tone}"] if self.must_include: lines.append(f"Must include: {', '.join(self.must_include)}") if self.must_ask: lines.append(f"Must ask about: {', '.join(self.must_ask)}") if self.must_not: lines.append(f"NEVER say: {', '.join(self.must_not)}") if self.max_questions >= 0: lines.append(f"Maximum questions: {self.max_questions}") return "\n".join(lines) # --------------------------------------------------------------------------- # Flow Execution Engine # --------------------------------------------------------------------------- class FlowEngine: """ Executes conversation flows and manages session state. Usage: engine = FlowEngine(config) session = engine.create_session("session-123") # Start call result = engine.start_call(session) # result.response_spec has what Avery should say # Process patient response result = engine.process_turn(session, patient_text, assessment) # result tells you: next response_spec, slots collected, flow status """ def __init__(self, config: DecisionConfigLoader): self._config = config self._resolver = FlowResolver(config) self._sessions: Dict[str, SessionState] = {} def create_session( self, session_id: str, patient_id: Optional[str] = None, tenant_id: Optional[str] = None, ) -> SessionState: """Create a new call session.""" session = SessionState( session_id=session_id, patient_id=patient_id, tenant_id=tenant_id, call_phase=CallPhase.SESSION_START, ) self._sessions[session_id] = session logger.info("Session created: %s (patient=%s, tenant=%s)", session_id, patient_id, tenant_id) return session def get_session(self, session_id: str) -> Optional[SessionState]: """Retrieve an existing session.""" return self._sessions.get(session_id) def destroy_session(self, session_id: str) -> None: """Destroy a session.""" self._sessions.pop(session_id, None) def start_call(self, session: SessionState) -> TurnResult: """ Start a new call — returns the first flow step (greeting). """ flow = self._resolver.resolve_initial_flow(session) if not flow: return TurnResult( outcome=TurnOutcome.END, message="No initial flow available", ) session.current_flow = flow.flow_id session.current_step = 0 step = flow.get_step(0) if not step: return TurnResult(outcome=TurnOutcome.END, message="Flow has no steps") return TurnResult( outcome=TurnOutcome.PROCEED, flow_id=flow.flow_id, flow_type=flow.flow_type, step_number=0, step_type=step.step_type, response_spec=step.response_spec, collects=step.collects, handoff_target=step.handoff_target, handoff_urgency=step.handoff_urgency, ) def process_turn( self, session: SessionState, patient_text: str, assessment: Optional[RiskAssessment] = None, extracted_slots: Optional[Dict[str, Any]] = None, ) -> TurnResult: """ Process a patient turn and determine the next action. Args: session: Current session state patient_text: What the patient said assessment: Risk assessment from Phase 1 engine extracted_slots: Slots extracted from patient text (by NLU) Returns: TurnResult with next response_spec or escalation action """ session.turn_count += 1 # Update slots from extracted values if extracted_slots: for k, v in extracted_slots.items(): session.slots.set_slot(k, v) # SAFETY CHECK: If assessment triggers R3/hard_escalate, interrupt immediately if assessment and (assessment.hard_escalate or assessment.risk_class == RiskClass.R3): return self._handle_emergency_escalation(session, assessment) if assessment and assessment.risk_class == RiskClass.R2: return self._handle_nurse_transfer(session, assessment) # Get current flow current_flow = self._get_current_flow(session) if not current_flow: # No active flow — resolve from assessment flow = self._resolver.resolve_next_flow(session, assessment) if not flow: return TurnResult(outcome=TurnOutcome.END, message="No applicable flow") return self._enter_flow(session, flow) # Advance to next step in current flow return self._advance_flow(session, current_flow, assessment) def _handle_emergency_escalation( self, session: SessionState, assessment: RiskAssessment ) -> TurnResult: """Handle R3 / hard-escalate — interrupt everything, route to escalation.""" domain = assessment.primary_domain or "unknown" session.call_phase = CallPhase.ENDED # Try to load domain-specific escalation flow flow = self._resolver.resolve_next_flow(session, assessment) if flow and flow.steps: step = flow.get_step(0) logger.warning( "EMERGENCY ESCALATION: session=%s domain=%s risk=%s", session.session_id, domain, assessment.risk_class.value, ) return TurnResult( outcome=TurnOutcome.ESCALATE, flow_id=flow.flow_id, flow_type=flow.flow_type, step_number=0, step_type=step.step_type if step else "respond", response_spec=step.response_spec if step else None, handoff_target=step.handoff_target if step else "clinical_nurse", handoff_urgency="immediate", risk_class=assessment.risk_class, domain=domain, message=f"R3 EMERGENCY: {domain}", ) # Fallback: generic escalation return TurnResult( outcome=TurnOutcome.ESCALATE, handoff_target="clinical_nurse", handoff_urgency="immediate", risk_class=assessment.risk_class, domain=domain, message=f"R3 EMERGENCY: {domain} — transfer to nurse immediately", ) def _handle_nurse_transfer( self, session: SessionState, assessment: RiskAssessment ) -> TurnResult: """Handle R2 — warm transfer to nurse.""" domain = assessment.primary_domain or "unknown" flow = self._resolver.resolve_next_flow(session, assessment) if flow and flow.steps: step = flow.get_step(0) return TurnResult( outcome=TurnOutcome.HANDOFF, flow_id=flow.flow_id, flow_type=flow.flow_type, step_number=0, step_type=step.step_type if step else "respond", response_spec=step.response_spec if step else None, handoff_target=step.handoff_target if step else "clinical_nurse", handoff_urgency="soon", risk_class=assessment.risk_class, domain=domain, message=f"R2: {domain} — schedule nurse callback", ) return TurnResult( outcome=TurnOutcome.HANDOFF, handoff_target="clinical_nurse", handoff_urgency="soon", risk_class=assessment.risk_class, domain=domain, message=f"R2: {domain} — schedule nurse callback", ) def _enter_flow(self, session: SessionState, flow: FlowDefinition) -> TurnResult: """Enter a new flow and return the first step.""" session.current_flow = flow.flow_id session.current_step = 0 session.domain_history.append(flow.flow_id) step = flow.get_step(0) if not step: return TurnResult(outcome=TurnOutcome.PROCEED, message="Flow has no steps") return TurnResult( outcome=TurnOutcome.PROCEED, flow_id=flow.flow_id, flow_type=flow.flow_type, step_number=0, step_type=step.step_type, response_spec=step.response_spec, collects=step.collects, handoff_target=step.handoff_target, handoff_urgency=step.handoff_urgency, ) def _advance_flow( self, session: SessionState, flow: FlowDefinition, assessment: Optional[RiskAssessment], ) -> TurnResult: """Advance to the next step in the current flow.""" next_step_idx = session.current_step + 1 if next_step_idx < flow.step_count: # Move to next step session.current_step = next_step_idx step = flow.get_step(next_step_idx) return TurnResult( outcome=TurnOutcome.PROCEED, flow_id=flow.flow_id, flow_type=flow.flow_type, step_number=next_step_idx, step_type=step.step_type, response_spec=step.response_spec, collects=step.collects, handoff_target=step.handoff_target, handoff_urgency=step.handoff_urgency, ) # Flow complete — evaluate exit rules exit_rules = flow.exit_rules exit_target = exit_rules.get("on_complete", "end") # Check for special exits based on slot values if "on_refused" in exit_rules and session.slots.get_slot("consent_given") == "no": exit_target = exit_rules["on_refused"] if "on_unclear" in exit_rules and session.slots.get_slot("clarification_text") is None: exit_target = exit_rules.get("on_unclear", exit_target) if "on_concern" in exit_rules: # Check if any concern was raised during the flow concern_slots = ["pain_level", "medication_adherence"] for slot in concern_slots: val = session.slots.get_slot(slot) if val and isinstance(val, str) and any(w in val.lower() for w in ["severe", "bad", "not taking", "stopped"]): exit_target = exit_rules["on_concern"] break # Reset current flow session.current_flow = None session.current_step = 0 if exit_target == "end": session.call_phase = CallPhase.ENDED return TurnResult(outcome=TurnOutcome.END, message="Flow complete") if exit_target == "re_evaluate" or exit_target == "evaluate_risk": # Signal to re-run the risk assessment and determine next flow return TurnResult( outcome=TurnOutcome.PROCEED, message=f"Flow complete — re-evaluate ({exit_target})", requires_reassessment=True, ) # Resolve exit target as next flow next_flow = self._resolver.resolve_next_flow(session, assessment, current_exit=exit_target) if next_flow: return self._enter_flow(session, next_flow) return TurnResult(outcome=TurnOutcome.END, message="No next flow resolved") def _get_current_flow(self, session: SessionState) -> Optional[FlowDefinition]: """Get the current flow definition from session state.""" if not session.current_flow: return None # Check primitives first prim = self._config.primitives.get(session.current_flow) if prim: return FlowDefinition.from_yaml(prim, source=f"primitives/{session.current_flow}") # Check taxonomy flows for domain, flows in self._config.taxonomy_flows.items(): if session.current_flow in flows: return FlowDefinition.from_yaml( flows[session.current_flow], source=f"taxonomy/{domain}/{session.current_flow}", ) return None # --------------------------------------------------------------------------- # Turn Result # --------------------------------------------------------------------------- class TurnResult: """Result of processing a conversational turn.""" def __init__( self, outcome: TurnOutcome, flow_id: Optional[str] = None, flow_type: Optional[str] = None, step_number: int = 0, step_type: Optional[str] = None, response_spec: Optional[ResponseSpecDef] = None, collects: Optional[List[str]] = None, handoff_target: Optional[str] = None, handoff_urgency: Optional[str] = None, risk_class: Optional[RiskClass] = None, domain: Optional[str] = None, message: str = "", requires_reassessment: bool = False, ): self.outcome = outcome self.flow_id = flow_id self.flow_type = flow_type self.step_number = step_number self.step_type = step_type self.response_spec = response_spec self.collects = collects or [] self.handoff_target = handoff_target self.handoff_urgency = handoff_urgency self.risk_class = risk_class self.domain = domain self.message = message self.requires_reassessment = requires_reassessment @property def is_escalation(self) -> bool: return self.outcome == TurnOutcome.ESCALATE @property def is_handoff(self) -> bool: return self.outcome in (TurnOutcome.ESCALATE, TurnOutcome.HANDOFF) @property def is_end(self) -> bool: return self.outcome == TurnOutcome.END @property def prompt_constraints(self) -> Optional[str]: """Get LLM prompt constraints from response spec.""" if self.response_spec: return self.response_spec.to_prompt_constraints() return None def to_dict(self) -> Dict[str, Any]: """Serialize for API response.""" result = { "outcome": self.outcome.value, "flow_id": self.flow_id, "flow_type": self.flow_type, "step_number": self.step_number, "step_type": self.step_type, "collects": self.collects, "handoff_target": self.handoff_target, "handoff_urgency": self.handoff_urgency, "message": self.message, "requires_reassessment": self.requires_reassessment, } if self.response_spec: result["response_spec"] = { "id": self.response_spec.spec_id, "goal": self.response_spec.goal, "tone": self.response_spec.tone, "must_include": self.response_spec.must_include, "must_ask": self.response_spec.must_ask, "must_not": self.response_spec.must_not, "max_questions": self.response_spec.max_questions, } if self.risk_class: result["risk_class"] = self.risk_class.value if self.domain: result["domain"] = self.domain return result def __repr__(self) -> str: return ( f"TurnResult(outcome={self.outcome.value}, flow={self.flow_id}, " f"step={self.step_number}, type={self.step_type})" )