File size: 5,063 Bytes
f019486
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from types import SimpleNamespace

from agent import graph


def _part(text=None, thought=False, function_call=None):
    return SimpleNamespace(text=text, thought=thought, function_call=function_call)


def _fc(name, args):
    return SimpleNamespace(name=name, args=args)


def _response(parts):
    content = SimpleNamespace(role="model", parts=parts)
    return SimpleNamespace(candidates=[SimpleNamespace(content=content)])


def test_split_response_separates_thought_answer_calls():
    resp = _response([
        _part(text="thinking hard", thought=True),
        _part(function_call=_fc("calculator", {"expression": "2+2"})),
    ])
    thought, answer, calls, content = graph.split_response(resp)
    assert thought == "thinking hard"
    assert answer == ""
    assert len(calls) == 1 and calls[0].name == "calculator"
    assert content.role == "model"


def test_split_response_final_answer():
    resp = _response([
        _part(text="done thinking", thought=True),
        _part(text="The answer is 42.", thought=False),
    ])
    thought, answer, calls, _ = graph.split_response(resp)
    assert thought == "done thinking"
    assert answer == "The answer is 42."
    assert calls == []


def test_classify_weak():
    assert graph.classify_weak("wikipedia_search", "")
    assert graph.classify_weak("wikipedia_search", "tiny")
    assert graph.classify_weak("wikipedia_search", "France may refer to: the country")
    assert graph.classify_weak("wikipedia_search", "ERROR: boom")
    assert not graph.classify_weak("wikipedia_search", "Paris is the capital of France " * 5)
    assert not graph.classify_weak("calculator", "42")
    assert graph.classify_weak("calculator", "ERROR: bad expression")


def test_signature_is_stable_and_distinguishes_args():
    a = graph.signature(_fc("wikipedia_search", {"query": "France"}))
    b = graph.signature(_fc("wikipedia_search", {"query": "France"}))
    c = graph.signature(_fc("wikipedia_search", {"query": "Paris"}))
    assert a == b
    assert a != c


class ScriptedClient:
    """Returns queued fake responses; mimics client.models.generate_content."""

    def __init__(self, responses):
        self._responses = list(responses)
        self.models = self

    def generate_content(self, *, model, contents, config):
        return self._responses.pop(0)


class LoopingClient:
    """Always returns the same response (to exercise the step cap)."""

    def __init__(self, response):
        self._response = response
        self.models = self

    def generate_content(self, *, model, contents, config):
        return self._response


def _events(task, client, tool_fns, **kw):
    return list(graph.stream_run(
        task, client=client, tool_fns=tool_fns,
        declarations=[], system_prompt="sys", sleep=lambda s: None, **kw,
    ))


def test_full_loop_shows_revision_and_final():
    wiki_calls = {"n": 0}

    def fake_wiki(query):
        wiki_calls["n"] += 1
        if wiki_calls["n"] == 1:
            return "France may refer to: the country, a film, ..."
        return "Paris is the capital of France. " * 6

    tool_fns = {"wikipedia_search": fake_wiki, "calculator": lambda expression: "2140"}

    responses = [
        _response([_part("Find the capital.", thought=True),
                   _part(function_call=_fc("wikipedia_search", {"query": "France"}))]),
        _response([_part("That was ambiguous; searching more specifically.", thought=True),
                   _part(function_call=_fc("wikipedia_search", {"query": "Paris capital population"}))]),
        _response([_part("Now compute.", thought=True),
                   _part(function_call=_fc("calculator", {"expression": "2140000/1000"}))]),
        _response([_part("Done.", thought=True), _part("The answer is 2140.", thought=False)]),
    ]

    events = _events("q", ScriptedClient(responses), tool_fns)
    kinds = [e["kind"] for e in events]
    assert "thought" in kinds and "tool_call" in kinds and "observation" in kinds

    thoughts = [e for e in events if e["kind"] == "thought"]
    assert thoughts[0]["revision"] is False
    assert any(t["revision"] for t in thoughts)  # step 2 thought is a revision

    obs = [e for e in events if e["kind"] == "observation"]
    assert obs[0]["weak"] is True
    assert obs[1]["weak"] is False

    finals = [e for e in events if e["kind"] == "final"]
    assert finals and finals[-1]["text"] == "The answer is 2140."
    assert len([e for e in events if e["kind"] == "tool_call"]) == 3


def test_step_cap_emits_limit():
    looping = LoopingClient(_response([
        _part("loop", thought=True),
        _part(function_call=_fc("calculator", {"expression": "1+1"})),
    ]))
    tool_fns = {"calculator": lambda expression: "2"}
    events = _events("q", looping, tool_fns, max_steps=3)
    assert any(e["kind"] == "limit" for e in events)
    assert len([e for e in events if e["kind"] == "tool_call"]) == 3