DCE / decision /engine /flow_engine.py
That guy James Bond :)
Deploy Medical Intent Escalation API
af61b34
Raw
History Blame Contribute Delete
27.7 kB
"""
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})"
)