Skip to content

Commit ed75db2

Browse files
authored
Merge pull request #477 from hud-evals/fix-template-annotation-resolution
feat: enhance parameter annotation resolution in task signatures
2 parents 4089df0 + a6ac044 commit ed75db2

3 files changed

Lines changed: 68 additions & 1 deletion

File tree

hud/environment/env.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -98,7 +98,7 @@ def __init__(
9898
#: Type the agent must produce (``None`` = plain text). Drives answer
9999
#: deserialization into ``Answer[T]``.
100100
self.return_type = returns
101-
self.sig = inspect.signature(func)
101+
self.sig = inspect.signature(func, eval_str=True)
102102
functools.update_wrapper(self, func)
103103

104104
def manifest_entry(self) -> dict[str, Any]:

hud/environment/tests/test_manifest.py

Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,8 @@
77

88
from __future__ import annotations
99

10+
from typing import Literal
11+
1012
from pydantic import BaseModel
1113

1214
from hud.environment import Environment
@@ -17,6 +19,14 @@ class _Point(BaseModel):
1719
y: int
1820

1921

22+
class _Payload(BaseModel):
23+
text: str
24+
25+
26+
_Mode = Literal["easy", "hard"]
27+
_Retries = int | None
28+
29+
2030
def test_args_schema_captures_params_defaults_and_required() -> None:
2131
env = Environment("manifests")
2232

@@ -86,3 +96,19 @@ async def typed():
8696
assert entry["input"]["properties"]["x"]["type"] == "integer"
8797
assert entry["returns"]["properties"]["y"]["type"] == "integer"
8898
assert entry["args"]["properties"] == {}
99+
100+
101+
def test_args_schema_resolves_postponed_rich_annotations() -> None:
102+
env = Environment("manifests")
103+
104+
@env.template()
105+
async def rich(mode: _Mode, payload: _Payload, retries: _Retries = None):
106+
yield "go"
107+
yield 1.0
108+
109+
assert callable(rich)
110+
schema = env.tasks["rich"].manifest_entry()["args"]
111+
assert schema["properties"]["mode"]["enum"] == ["easy", "hard"]
112+
assert schema["properties"]["retries"]["default"] is None
113+
assert "$ref" in schema["properties"]["payload"]
114+
assert set(schema["required"]) == {"mode", "payload"}

hud/environment/tests/test_server.py

Lines changed: 41 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,10 @@
77

88
from __future__ import annotations
99

10+
from typing import Literal
11+
1012
import pytest
13+
from pydantic import BaseModel
1114

1215
from hud.clients import HudProtocolError
1316
from hud.environment import Answer, Environment
@@ -17,6 +20,13 @@
1720
from .conftest import served
1821

1922

23+
class _Payload(BaseModel):
24+
text: str
25+
26+
27+
_Mode = Literal["upper", "lower"]
28+
29+
2030
async def test_dict_grade_without_numeric_score_errors_loudly() -> None:
2131
env = Environment("badgrade")
2232

@@ -79,3 +89,34 @@ def test_answer_holds_parsed_content_and_raw_string() -> None:
7989
answer = Answer(content={"final": "42"}, raw='{"final": "42"}')
8090
assert answer.content == {"final": "42"}
8191
assert answer.raw == '{"final": "42"}'
92+
93+
94+
async def test_start_coerces_postponed_rich_annotations() -> None:
95+
env = Environment("coerce")
96+
97+
@env.template()
98+
async def typed(mode: _Mode, payload: _Payload, retries: int | None = None):
99+
if mode == "upper":
100+
prompt = payload.text.upper()
101+
elif mode == "lower":
102+
prompt = payload.text.lower()
103+
else:
104+
raise ValueError(f"unexpected mode: {mode!r}")
105+
if retries is not None:
106+
prompt += "!" * retries
107+
yield prompt
108+
yield 1.0
109+
110+
assert callable(typed)
111+
async with served(env) as client:
112+
async with Run(
113+
client,
114+
"typed",
115+
{
116+
"mode": '"upper"',
117+
"payload": '{"text":"hello"}',
118+
"retries": "3",
119+
},
120+
) as run:
121+
run.trace.content = "x"
122+
assert run.prompt == "HELLO!!!"

0 commit comments

Comments
 (0)