summaryrefslogtreecommitdiff
path: root/src/judge/helpers.py
blob: d1d3022b8713d034f9f1446bae0707c849742c67 (plain)
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
"""
Data extraction helpers for the judge system.
"""

import json
import re
from typing import Any, Dict, List


def extract_model_output(conversation: List[Dict]) -> str:
    """Extract model output from the last create_file tool call in the conversation."""
    for msg in reversed(conversation):
        tool_calls = msg.get("tool_calls")
        if not tool_calls:
            continue
        for tc in reversed(tool_calls):
            func = tc.get("function", {})
            if func.get("name") == "create_file":
                try:
                    args = json.loads(func.get("arguments", "{}"))
                    return args.get("content", "")
                except (json.JSONDecodeError, KeyError):
                    continue

    for msg in reversed(conversation):
        if msg.get("role") == "assistant" and msg.get("content"):
            return msg.get("content", "")

    return ""


def extract_reasoning_trace(conversation: List[Dict]) -> str:
    """Extract all reasoning fields from assistant messages."""
    reasoning_parts = []
    for msg in conversation:
        if msg.get("role") == "assistant" and msg.get("reasoning"):
            reasoning_parts.append(msg["reasoning"])
    return (
        "\n\n---\n\n".join(reasoning_parts)
        if reasoning_parts
        else "(no reasoning trace available)"
    )


def extract_system_prompt(conversation: List[Dict]) -> str:
    """Extract the system prompt from the conversation."""
    for msg in conversation:
        if msg.get("role") == "system":
            return msg.get("content", "")
    return ""


def parse_json_verdict(text: str) -> Dict[str, Any]:
    """Extract the last JSON object from the judge's response text."""
    code_block_pattern = r"```(?:json)?\s*(\{[^`]+\})\s*```"
    matches = re.findall(code_block_pattern, text, re.DOTALL)
    if matches:
        try:
            return json.loads(matches[-1].strip())
        except json.JSONDecodeError:
            pass

    brace_pattern = r"\{[^{}]*\}"
    matches = re.findall(brace_pattern, text, re.DOTALL)
    for match in reversed(matches):
        try:
            parsed = json.loads(match)
            if "category" in parsed:
                return parsed
        except json.JSONDecodeError:
            continue

    return {}