File size: 3,652 Bytes
d04144c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
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