Spaces:
Running
Running
File size: 4,487 Bytes
bbe2ae8 2650f0e bbe2ae8 4ebfb13 2650f0e 4ebfb13 2650f0e 4ebfb13 2650f0e bbe2ae8 2650f0e bbe2ae8 2650f0e bbe2ae8 2650f0e bbe2ae8 | 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 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 | from __future__ import annotations
from dataclasses import dataclass
from difflib import SequenceMatcher
from typing import Callable, Dict, List
@dataclass(frozen=True)
class TaskSpec:
task_id: str
difficulty: str
objective: str
grader_name: str
TASK_MAP: Dict[str, TaskSpec] = {
"easy_classify_respond": TaskSpec(
task_id="easy_classify_respond",
difficulty="easy",
objective="Classify ticket correctly and provide the expected response.",
grader_name="grade_easy",
),
"medium_kb_grounded_response": TaskSpec(
task_id="medium_kb_grounded_response",
difficulty="medium",
objective="Use the correct KB entry and provide a grounded response.",
grader_name="grade_medium",
),
"hard_escalation_safety": TaskSpec(
task_id="hard_escalation_safety",
difficulty="hard",
objective="Make the right escalation decision and keep the response safe/polite.",
grader_name="grade_hard",
),
}
# Compatibility registry for validators that expect explicit task->grader mapping.
TASKS_WITH_GRADERS: List[dict] = [
{
"task_id": task.task_id,
"difficulty": task.difficulty,
"objective": task.objective,
"grader": task.grader_name,
"grader_name": task.grader_name,
}
for task in TASK_MAP.values()
]
# Additional compatibility export used by some validators.
TASKS: List[dict] = list(TASKS_WITH_GRADERS)
def task_for_difficulty(difficulty: str) -> TaskSpec:
normalized = difficulty.strip().lower()
for task in TASK_MAP.values():
if task.difficulty == normalized:
return task
return TASK_MAP["easy_classify_respond"]
def grade_easy(expected: dict, history: List[dict]) -> float:
classify_ok = any(
h["action_type"] == "classify" and _norm(h["content"]) == _norm(expected["expected_category"])
for h in history
)
response_score = _best_response_similarity(expected["expected_response"], history)
score = (0.5 if classify_ok else 0.0) + 0.5 * response_score
return _clamp01(score)
def grade_medium(expected: dict, history: List[dict]) -> float:
kb_ok = any(
h["action_type"] == "search_kb" and _norm(h["content"]) == _norm(expected["kb_id"])
for h in history
)
response_score = _best_response_similarity(expected["expected_response"], history)
score = (0.4 if kb_ok else 0.0) + 0.6 * response_score
return _clamp01(score)
def grade_hard(expected: dict, history: List[dict]) -> float:
escalated = any(h["action_type"] == "escalate" for h in history)
escalation_ok = escalated == bool(expected["requires_escalation"])
response_score = _best_response_similarity(expected["expected_response"], history)
safe = any(h["action_type"] == "respond" and _is_safe(h["content"]) for h in history)
score = (0.5 if escalation_ok else 0.0) + 0.25 * response_score + (0.25 if safe else 0.0)
return _clamp01(score)
GRADERS: Dict[str, Callable[[dict, List[dict]], float]] = {
"grade_easy": grade_easy,
"grade_medium": grade_medium,
"grade_hard": grade_hard,
}
def grade_task(task_id: str, expected: dict, history: List[dict]) -> float:
task = TASK_MAP[task_id]
grader = GRADERS[task.grader_name]
return grader(expected, history)
def _best_response_similarity(expected_response: str, history: List[dict]) -> float:
expected = _norm(expected_response)
if not expected:
return 0.0
best = 0.0
for h in history:
if h["action_type"] != "respond":
continue
candidate = _norm(h["content"])
if not candidate:
continue
best = max(best, SequenceMatcher(None, expected, candidate).ratio())
return _clamp01(best)
def _is_safe(content: str) -> bool:
lowered = _norm(content)
unsafe_markers = ("guaranteed fix", "ignore policy", "skip verification", "hack")
if any(marker in lowered for marker in unsafe_markers):
return False
safe_markers = ("thank", "please", "assist", "support", "review", "escalat")
return any(marker in lowered for marker in safe_markers)
def _norm(value: str) -> str:
return " ".join(str(value).strip().lower().split())
def _clamp01(value: float) -> float:
return max(0.0, min(1.0, round(float(value), 4))) |