File size: 4,907 Bytes
5df4501
 
 
 
 
 
8b00c96
5df4501
 
 
91bef46
5df4501
91bef46
5df4501
 
f592868
 
 
5df4501
 
 
c721b39
 
 
 
 
 
5df4501
 
 
c721b39
f592868
 
 
 
 
c721b39
f592868
c721b39
f592868
c721b39
 
5df4501
c721b39
f592868
 
c721b39
 
 
 
 
 
 
5df4501
 
 
c721b39
 
 
 
 
5df4501
c721b39
f592868
 
c721b39
 
 
 
 
 
 
5df4501
 
 
c721b39
 
 
5df4501
 
 
 
f592868
 
5df4501
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8b00c96
 
 
 
f592868
 
8b00c96
 
 
 
 
 
 
 
 
 
 
5df4501
 
f592868
 
5df4501
 
 
 
 
 
 
 
 
f592868
 
5df4501
 
 
 
 
 
 
 
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
133
134
135
136
137
138
139
140
"""Robust action parsing for untrusted LLM outputs."""

from __future__ import annotations

import json
import re
import os
from typing import Any, Mapping, Optional

from utils.schemas import ParsedAction
from utils.constants import PARSER_DEFAULT_U, PARSER_DEFAULT_F, INVALID_ACTION_PENALTY

INVALID_OUTPUT_PENALTY = -INVALID_ACTION_PENALTY

_NUMBER_PATTERN = r"-?(?:\d+(?:\.\d*)?|\.\d+)"
_U_TARGET_PATTERN = re.compile(r'(?i)["\']?u_target["\']?(?:\s*[:=]\s*|\s+)(' + _NUMBER_PATTERN + r')')
_F_TARGET_PATTERN = re.compile(r'(?i)["\']?f_target["\']?(?:\s*[:=]\s*|\s+)(' + _NUMBER_PATTERN + r')')
_PAIR_PATTERN = re.compile(r'^\s*(' + _NUMBER_PATTERN + r')[\s,;]+(' + _NUMBER_PATTERN + r')\s*$')


def _clamp_unit_interval(value: Any, *, fallback: float) -> float:
    """Coerce to float and clamp to [0.0, 1.0]."""
    try:
        numeric_value = float(value)
    except (TypeError, ValueError):
        numeric_value = fallback
    return max(0.0, min(1.0, numeric_value))


def _from_mapping(payload: Mapping[str, Any], raw_text: str) -> ParsedAction:
    """Parse the expected JSON payload from a mapping."""
    # Find keys ignoring case
    u_key = next((k for k in payload.keys() if k.lower() == "u_target"), "U_target")
    f_key = next((k for k in payload.keys() if k.lower() == "f_target"), "F_target")
    
    if u_key not in payload or f_key not in payload:
        missing = []
        if u_key not in payload:
            missing.append("U_target")
        if f_key not in payload:
            missing.append("F_target")
        raise ValueError(f"missing required keys: {', '.join(missing)}")

    return ParsedAction(
        u_target=_clamp_unit_interval(payload[u_key], fallback=PARSER_DEFAULT_U),
        f_target=_clamp_unit_interval(payload[f_key], fallback=PARSER_DEFAULT_F),
        source="json",
        used_fallback=False,
        invalid_output=False,
        penalty_applied=0.0,
        raw_text=raw_text,
        parse_error=None,
    )


def _anchored_extract(raw_text: str) -> ParsedAction:
    """Extract targets from anchored key/value pairs embedded in free text."""
    u_match = _U_TARGET_PATTERN.search(raw_text)
    f_match = _F_TARGET_PATTERN.search(raw_text)
    if not u_match or not f_match:
        raise ValueError("anchored extraction failed")

    return ParsedAction(
        u_target=_clamp_unit_interval(u_match.group(1), fallback=PARSER_DEFAULT_U),
        f_target=_clamp_unit_interval(f_match.group(1), fallback=PARSER_DEFAULT_F),
        source="fallback",
        used_fallback=True,
        invalid_output=True,
        penalty_applied=INVALID_OUTPUT_PENALTY,
        raw_text=raw_text,
        parse_error="strict_json_failed",
    )


def parse_llm_action(
    raw_text: Optional[str],
    previous_valid_action: Optional[Mapping[str, float]] = None,
    default_action: Optional[Mapping[str, float]] = None,
) -> ParsedAction:
	"""Parse untrusted model output into a canonical action without raising."""
	text = (raw_text or "").strip()
	fallback_default = default_action or {
		"U_target": PARSER_DEFAULT_U,
		"F_target": PARSER_DEFAULT_F,
	}

	try:
		payload = json.loads(text)
		if not isinstance(payload, Mapping):
			raise ValueError("json payload is not an object")
		return _from_mapping(payload, text)
	except Exception as json_error:
		json_error_message = str(json_error)

	try:
		return _anchored_extract(text)
	except Exception as fallback_error:
		fallback_error_message = str(fallback_error)

	# Always try to match the compact "value value" pair format as a valid
	# alternative to JSON, since the system prompt may be tuned for pair output.
	pair_match = _PAIR_PATTERN.match(text)
	if pair_match:
		u_val = _clamp_unit_interval(pair_match.group(1), fallback=PARSER_DEFAULT_U)
		f_val = _clamp_unit_interval(pair_match.group(2), fallback=PARSER_DEFAULT_F)
		return ParsedAction(
			u_target=u_val,
			f_target=f_val,
			source="fallback",
			used_fallback=False,
			invalid_output=False,
			penalty_applied=0.0,
			raw_text=text,
			parse_error=None,
		)

	if previous_valid_action is not None:
		return ParsedAction(
			u_target=_clamp_unit_interval(previous_valid_action.get("U_target"), fallback=PARSER_DEFAULT_U),
			f_target=_clamp_unit_interval(previous_valid_action.get("F_target"), fallback=PARSER_DEFAULT_F),
			source="previous_valid",
			used_fallback=True,
			invalid_output=True,
			penalty_applied=INVALID_OUTPUT_PENALTY,
			raw_text=text,
			parse_error=f"{json_error_message}; {fallback_error_message}",
		)

	return ParsedAction(
		u_target=_clamp_unit_interval(fallback_default.get("U_target"), fallback=PARSER_DEFAULT_U),
		f_target=_clamp_unit_interval(fallback_default.get("F_target"), fallback=PARSER_DEFAULT_F),
		source="default",
		used_fallback=True,
		invalid_output=True,
		penalty_applied=INVALID_OUTPUT_PENALTY,
		raw_text=text,
		parse_error=f"{json_error_message}; {fallback_error_message}",
	)