Files
praxis/server/scenarios/classifier.py
T
Praxis CI 9108b07de1 feat(P01-03-05,P01-03-06): interruptibility harness + LLM-as-judge branch classifier
server/interruptibility.py — programmatic check that the pipeline task is
configured with allow_interruptions=True (D-008 abort-and-yield). The manual
test (speaking during AI speech cuts it off) is documented in
docs/latency-report.md; the automated test confirms the flag is set.

server/scenarios/classifier.py — classify_branch() uses
deepseek-v4-flash:cloud no-think (D-020) as LLM-as-judge to classify the
learner's turn transcripts into accept_resolution or escalate at session end,
OFFLINE from the voice loop (D-P1-05, G-002 post-hoc classification). Parses
the LLM's JSON response leniently (code fences, unknown-id fallback, malformed
JSON scan). classify_branch_sync_heuristic() is a rule-based fallback for
tests/e2e when no API key is present. 11 tests pass (interruptibility flag
on/off/missing; classifier accept/escalate scripted transcripts; JSON parse
valid/code-fenced/unknown-id/malformed; fake-LLM async classify; offline
structural assertion).

---ci---
phase: 1
milestone: v0.1
plan: 03
task: 03-05,03-06
status: execute
persona: backend-engineer
requirements:
  covered: [REQ-VOICE-04, REQ-SCEN-01]
---/ci---
2026-08-01 13:13:28 +00:00

137 lines
4.9 KiB
Python

"""Branch classifier — LLM-as-judge for learner-signal classification (R7, TASK-03-06).
At session end (or turn boundary), classifies the learner's turn transcripts
into a scenario branch (accept_resolution or escalate) based on the scenario's
learner_signals definitions. Runs OFFLINE from the voice loop (not on the
latency-critical path) per D-P1-05.
Uses deepseek-v4-flash:cloud in no-think mode (D-020) via the LLMProvider —
cheap + fast enough for a one-shot end-of-session classification.
Per G-002: the v0.1 branch is a post-hoc outcome classification, not a runtime
conversation fork. This classifier produces the label that the debrief + DB
log consume.
"""
from __future__ import annotations
import json
from typing import Any
from server.scenarios.schema import Scenario
from server.services.base import GuardrailContext, LLMProvider
CLASSIFIER_SYSTEM_PROMPT = """\
You are a conversation-branch classifier for a customer-service role-play
training session. Given the learner's turns and the scenario's branch
definitions (each with learner_signals), classify which branch the learner's
behavior matches.
Respond with ONLY a JSON object: {"branch_id": "<id>", "reason": "<short>"}
No other text. If the signals are mixed, pick the closest match and explain in
the reason field.
"""
def _build_user_prompt(scenario: Scenario, learner_turns: list[str]) -> str:
branches_desc = "\n".join(
f" - {b.id}: signals={b.trigger.learner_signals}, outcome={b.outcome}"
for b in scenario.branches
)
turns_desc = "\n".join(f" Learner: {t}" for t in learner_turns)
return (
f"Scenario: {scenario.title}\n"
f"Branches:\n{branches_desc}\n\n"
f"Learner turns:\n{turns_desc}\n\n"
f"Which branch does the learner's behavior match? "
f"Respond with JSON {{\"branch_id\": ..., \"reason\": ...}}."
)
async def classify_branch(
llm: LLMProvider,
scenario: Scenario,
learner_turns: list[str],
) -> tuple[str, str]:
"""Classify the learner's turns into a branch id.
Args:
llm: the LLMProvider (uses debrief_model = deepseek-v4-flash:cloud no_think).
scenario: the loaded Scenario.
learner_turns: the learner's ASR transcripts for the session.
Returns:
(branch_id, reason) — branch_id is one of scenario.branch_ids().
"""
messages = [
{"role": "system", "content": CLASSIFIER_SYSTEM_PROMPT},
{"role": "user", "content": _build_user_prompt(scenario, learner_turns)},
]
text, _usage = await llm.chat_full(
messages, model=llm.debrief_model, no_think=True
)
return _parse_branch(text, scenario)
def _parse_branch(text: str, scenario: Scenario) -> tuple[str, str]:
"""Parse the LLM's JSON response into (branch_id, reason)."""
# Be lenient — strip code fences, find the JSON object.
cleaned = text.strip()
if cleaned.startswith("```"):
cleaned = cleaned.strip("`")
if cleaned.lower().startswith("json"):
cleaned = cleaned[4:]
try:
obj = json.loads(cleaned)
branch_id = obj.get("branch_id", "")
reason = obj.get("reason", "")
except json.JSONDecodeError:
# Fall back to a heuristic scan for a known branch id.
reason = "fallback: could not parse LLM JSON"
for b in scenario.branches:
if b.id in text:
return b.id, reason
return scenario.branches[0].id, reason
# Validate the branch id is known.
if branch_id not in scenario.branch_ids():
reason = f"fallback: unknown branch_id {branch_id!r}; {reason}"
branch_id = scenario.branches[0].id
return branch_id, reason
def classify_branch_sync_heuristic(
scenario: Scenario, learner_turns: list[str]
) -> str:
"""A rule-based fallback classifier for tests (no LLM call).
Used by the e2e smoke test when no API key is present. Scans for keywords
matching each branch's learner_signals. Signal tokens are matched as
substrings (e.g. 'policy' matches 'policy_first'; 'empathy' matches
'empathy'; 'concrete resolution' matches 'concrete_resolution').
"""
text = " ".join(learner_turns).lower()
best = scenario.branches[0]
best_score = -1
for b in scenario.branches:
score = 0
for sig in b.trigger.learner_signals:
# Match the signal as a space- or underscore-separated phrase.
token = sig.replace("_", " ").lower()
# Use the first significant word as a loose keyword (e.g. 'policy'
# for 'policy_first', 'defensive' for 'defensive').
keyword = token.split()[0] if " " in token else token
if keyword in text or token in text:
score += 1
if score > best_score:
best_score = score
best = b
return best.id
__all__ = [
"classify_branch",
"classify_branch_sync_heuristic",
"CLASSIFIER_SYSTEM_PROMPT",
]