astera-customerAI / runtime /answer_quality.py
G-ACE's picture
Deploy current Customer AI runtime 96b5a02fe7bbc525787d1a4359b7358d63d4d89a
d04144c verified
Raw
History Blame Contribute Delete
3.65 kB
from __future__ import annotations
from dataclasses import dataclass
from .schemas import NeedTask, ResolutionMode, TaskResolution
@dataclass(frozen=True)
class IntegratedAnswerPlan:
needs: tuple[NeedTask, ...]
resolutions: tuple[TaskResolution, ...]
blocked_task_ids: tuple[str, ...] = ()
missing_evidence_task_ids: tuple[str, ...] = ()
missing_user_inputs: tuple[str, ...] = ()
safety_blocked: bool = False
runtime_failure: bool = False
@dataclass(frozen=True)
class ComposedAnswer:
mode: ResolutionMode
answer: str | None
resolved_task_ids: tuple[str, ...]
unresolved_task_ids: tuple[str, ...]
clarification_questions: tuple[str, ...] = ()
class FinalAnswerComposer:
def compose(self, plan: IntegratedAnswerPlan) -> ComposedAnswer:
all_task_ids = tuple(n.task_id for n in plan.needs)
if plan.runtime_failure:
return ComposedAnswer(ResolutionMode.RUNTIME_FAILURE, None, (), all_task_ids)
if plan.safety_blocked:
return ComposedAnswer(ResolutionMode.SAFETY_BLOCKED, None, (), all_task_ids)
by_task = {r.task_id: r for r in plan.resolutions}
blocked = set(plan.blocked_task_ids) | set(plan.missing_evidence_task_ids)
resolved: list[TaskResolution] = []
unresolved: list[str] = []
for need in plan.needs:
item = by_task.get(need.task_id)
if need.task_id in blocked or item is None or not item.resolved:
unresolved.append(need.task_id)
else:
resolved.append(item)
useful = "\n\n".join(r.public_text.strip() for r in resolved) or None
if plan.missing_user_inputs:
question = f"{plan.missing_user_inputs[0]}γ‚’η’Ίθͺγ—てください。"
return ComposedAnswer(ResolutionMode.NEEDS_USER_INPUT, useful, tuple(r.task_id for r in resolved), tuple(unresolved), (question,))
if unresolved:
return ComposedAnswer(ResolutionMode.SAFE_PARTIAL, useful, tuple(r.task_id for r in resolved), tuple(unresolved))
return ComposedAnswer(ResolutionMode.RESOLVED, useful, tuple(r.task_id for r in resolved), ())
@dataclass(frozen=True)
class RuntimeSatisfactionSignals:
all_major_needs_covered: bool
evidence_complete: bool
context_consistent: bool
required_actionability_present: bool
false_premise_corrected: bool
unnecessary_clarification_count: int = 0
unsupported_claim_count: int = 0
stale_grounding_count: int = 0
terminology_violation_count: int = 0
class RuntimeSatisfactionGate:
def evaluate(self, mode: ResolutionMode, signals: RuntimeSatisfactionSignals) -> tuple[bool, list[str]]:
failures: list[str] = []
if mode != ResolutionMode.RESOLVED: failures.append("conversation_not_resolved")
if not signals.all_major_needs_covered: failures.append("major_need_missing")
if not signals.evidence_complete: failures.append("evidence_incomplete")
if not signals.context_consistent: failures.append("context_inconsistent")
if not signals.required_actionability_present: failures.append("required_actionability_missing")
if not signals.false_premise_corrected: failures.append("false_premise_uncorrected")
if signals.unnecessary_clarification_count: failures.append("unnecessary_clarification")
if signals.unsupported_claim_count: failures.append("unsupported_claim")
if signals.stale_grounding_count: failures.append("stale_grounding")
if signals.terminology_violation_count: failures.append("terminology_violation")
return not failures, failures