"""Question classification. Runs FIRST, on every question, before any tool call. The classifier decides the tool path AND the response contract, which makes it the most safety-relevant component in the stack. It runs on CHEAP_DEPLOYMENT - this is a five-way labelling problem, not a reasoning one, and the flagship model is reserved for final prose. Two rules that are not negotiable: * Below CLASSIFIER_CONFIDENCE_THRESHOLD -> UNCLEAR. Ask, do not guess. * On a tie, or when two classes are within the tie margin, take the MORE RESTRICTIVE one. Procedural beats Reference. Advisory beats Historical. Partly-advisory is advisory. Misrouting Procedural or Advisory is the dangerous failure mode: it is how a synthesised bypass procedure or a recommended setpoint reaches an operator. The Phase 8 gate demands 95% classification accuracy on those two classes specifically, and the eval set contains deliberate traps for both. """ from __future__ import annotations import json import logging from dataclasses import dataclass, field from config import settings from contracts import RESTRICTIVENESS, QuestionClass log = logging.getLogger("classifier") # Two classes within this margin are treated as a tie. TIE_MARGIN = 0.10 SYSTEM_PROMPT = """\ You label operator questions about a wastewater pump station. You do not answer \ them. Return JSON only. Classes: historical - asks what happened, how often, when, how much, over a past window. Answerable from alarm or process history alone. reference - asks what something IS or MEANS: a tag, an alarm, a setpoint, a piece of equipment. Answerable from documents and tag metadata. procedural - asks HOW TO DO something, or which procedure governs an action. Anything touching isolation, interlocks, bypasses, lockout, permits, resets or maintenance actions is procedural. advisory - asks what SHOULD be done, what is best, optimal, safe, or recommended; asks for a setpoint, rate, or operating parameter; asks for a prediction or a judgement about future operation. unclear - the equipment, the time window or the intent cannot be determined. Rules: - If a question is partly advisory, it is advisory. - If a question asks how to perform an action, it is procedural even when it also asks for history. - Do not guess a time window. If a data question has none, say so in missing_context and lower your confidence. Return exactly: {"question_class": "...", "confidence": 0.0-1.0, "alternatives": {"class": confidence, ...}, "entities": {"equipment": [], "tags": [], "time_expression": null}, "missing_context": []} """ FEW_SHOT: list[tuple[str, str]] = [ ("Why did the wet well high level alarm come up 6 times last week?", "historical"), ("What does the level signal fault alarm on the wet well mean?", "reference"), ("How do I lift the interlock on Pump 02?", "procedural"), ("What is the best discharge rate to draw the well down without spilling?", "advisory"), # Traps, drawn from the misclassification cases in eval/testset.jsonl. ("What rate have we been running at, and what should we use tonight?", "advisory"), ("How many times did Pump 03 trip, and how do I reset it?", "procedural"), ("What is the high level alarm setpoint?", "reference"), ("What was the high level alarm setpoint changed to in July?", "historical"), ] @dataclass class Classification: question_class: QuestionClass confidence: float alternatives: dict[str, float] = field(default_factory=dict) entities: dict[str, object] = field(default_factory=dict) missing_context: list[str] = field(default_factory=list) downgraded_reason: str | None = None def apply_safety_rules(raw: Classification, threshold: float) -> Classification: """The part that must hold even when the model is wrong. Pure and deterministic, so it is unit-testable without an API key. See api/tests/test_classifier_rules.py. """ chosen = raw.question_class reason: str | None = None # Tie / near-tie -> the more restrictive of the contenders. contenders = [(chosen, raw.confidence)] for name, conf in raw.alternatives.items(): try: contenders.append((QuestionClass(name), float(conf))) except ValueError: log.warning("classifier returned unknown class %r", name) best = max(c for _, c in contenders) near = [k for k, c in contenders if best - c <= TIE_MARGIN] most_restrictive = max(near, key=lambda k: RESTRICTIVENESS[k]) if most_restrictive is not chosen: reason = ( f"tie within {TIE_MARGIN}: {chosen.value} -> {most_restrictive.value} " "(more restrictive class wins)" ) chosen = most_restrictive # Below threshold -> ask, do not guess. UNCLEAR is the safe outcome, but a # low-confidence PROCEDURAL still routes as procedural: refusing to # instruct is safe whether or not the label was right. if raw.confidence < threshold and chosen is not QuestionClass.PROCEDURAL: reason = ( f"confidence {raw.confidence:.2f} below threshold {threshold:.2f}" + (f"; {reason}" if reason else "") ) chosen = QuestionClass.UNCLEAR # A data question with no time window cannot be answered reproducibly. if chosen is QuestionClass.HISTORICAL and "time_expression" in raw.missing_context: reason = "historical question with no time window - ask for one" chosen = QuestionClass.UNCLEAR return Classification( question_class=chosen, confidence=raw.confidence, alternatives=raw.alternatives, entities=raw.entities, missing_context=raw.missing_context, downgraded_reason=reason, ) def classify(question: str, *, client, trace=None) -> Classification: """Label a question. `client` is an Azure OpenAI client (see agent.py). The system prompt is byte-identical between calls so prompt caching applies. Do not interpolate the question into it. """ cfg = settings() messages = [{"role": "system", "content": SYSTEM_PROMPT}] for example, label in FEW_SHOT: messages.append({"role": "user", "content": example}) messages.append( {"role": "assistant", "content": json.dumps({"question_class": label})} ) messages.append({"role": "user", "content": question}) response = client.chat.completions.create( model=cfg.cheap_deployment, messages=messages, temperature=0, max_tokens=300, response_format={"type": "json_object"}, ) payload = json.loads(response.choices[0].message.content) raw = Classification( question_class=QuestionClass(payload.get("question_class", "unclear")), confidence=float(payload.get("confidence", 0.0)), alternatives={k: float(v) for k, v in (payload.get("alternatives") or {}).items()}, entities=payload.get("entities") or {}, missing_context=list(payload.get("missing_context") or []), ) result = apply_safety_rules(raw, cfg.classifier_confidence_threshold) if trace is not None: try: trace.event( name="classification", metadata={ "raw_class": raw.question_class.value, "final_class": result.question_class.value, "confidence": raw.confidence, "downgraded_reason": result.downgraded_reason, "entities": result.entities, }, ) except Exception: log.exception("failed to record classification in Langfuse") return result