forked from github/copilot-sdk
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest_set_tools_e2e.py
More file actions
195 lines (153 loc) · 6.92 KB
/
Copy pathtest_set_tools_e2e.py
File metadata and controls
195 lines (153 loc) · 6.92 KB
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
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
"""Live tool replacement E2Es using shared set_tools provider captures."""
import json
import pytest
from pydantic import BaseModel, Field
from copilot import define_tool
from copilot.session import PermissionHandler
from copilot.tools import Tool, ToolInvocation, ToolResult
from .testharness import E2ETestContext
pytestmark = pytest.mark.asyncio(loop_scope="module")
FRUIT_PROMPT = "Use lookup_fruit to find the fruit for code 42."
FRUIT_AND_VEGETABLE_PROMPT = (
"Use lookup_fruit to find the fruit for code 42 again, and use lookup_vegetable to find the "
"vegetable for code 7."
)
VEGETABLE_PROMPT = "Use lookup_vegetable to find the vegetable for code 7."
class FruitParams(BaseModel):
code: int = Field(description="Fruit code")
class VegetableParams(BaseModel):
code: int = Field(description="Vegetable code")
def lookup_fruit(fruit: str, calls: list[int] | None = None) -> Tool:
@define_tool("lookup_fruit", description="Looks up the fruit for a numeric code")
def tool(params: FruitParams) -> str:
if calls is not None:
calls.append(params.code)
return fruit
return tool
def lookup_vegetable(calls: list[int] | None = None) -> Tool:
@define_tool("lookup_vegetable", description="Looks up the vegetable for a numeric code")
def tool(params: VegetableParams) -> str:
if calls is not None:
calls.append(params.code)
return "carrot"
return tool
def retired_lookup() -> Tool:
async def handler(_invocation: ToolInvocation) -> ToolResult:
return ToolResult(text_result_for_llm="retired")
return Tool("retired_lookup", "Looks up a retired value", handler=handler)
def invalid_tool(calls: list[int]) -> Tool:
async def handler(_invocation: ToolInvocation) -> ToolResult:
calls.append(1)
return ToolResult(text_result_for_llm="never")
return Tool("invalid.tool", "Has a name the runtime rejects", handler=handler)
def offered_tools(exchange: dict) -> set[str]:
"""Names of the tools a model request offered."""
return {
tool["function"]["name"]
for tool in exchange["request"].get("tools") or []
if "function" in tool
}
def index_of_prompt(exchanges: list[dict], prompt: str) -> int:
"""Index of the first model request that carries ``prompt`` as a user message."""
for index, exchange in enumerate(exchanges):
for message in exchange["request"]["messages"]:
content = message.get("content")
text = content if isinstance(content, str) else json.dumps(content)
if message.get("role") == "user" and prompt in text:
return index
return -1
class TestSetTools:
async def test_replaces_tools_on_a_created_session(self, ctx: E2ETestContext):
original_lookups: list[int] = []
replacement_lookups: list[int] = []
vegetable_lookups: list[int] = []
session = await ctx.client.create_session(
on_permission_request=PermissionHandler.approve_all,
tools=[lookup_fruit("apple", original_lookups), retired_lookup()],
)
first = await session.send_and_wait(FRUIT_PROMPT, timeout=30)
assert first is not None
assert "apple" in (first.data.content or "")
await session.set_tools(
[
lookup_fruit("dragonfruit", replacement_lookups),
lookup_vegetable(vegetable_lookups),
]
)
second = await session.send_and_wait(FRUIT_AND_VEGETABLE_PROMPT, timeout=30)
assert second is not None
assert "dragonfruit" in (second.data.content or "")
assert "carrot" in (second.data.content or "")
assert original_lookups == [42]
assert replacement_lookups == [42]
assert vegetable_lookups == [7]
# Model requests after the replacement offer exactly the new tool set.
exchanges = await ctx.get_exchanges()
replaced_from = index_of_prompt(exchanges, FRUIT_AND_VEGETABLE_PROMPT)
assert replaced_from > 0
for exchange in exchanges[:replaced_from]:
tools = offered_tools(exchange)
assert {"lookup_fruit", "retired_lookup"} <= tools
assert "lookup_vegetable" not in tools
for exchange in exchanges[replaced_from:]:
tools = offered_tools(exchange)
assert {"lookup_fruit", "lookup_vegetable"} <= tools
assert "retired_lookup" not in tools
await session.disconnect()
async def test_replaces_tools_on_a_resumed_session(self, ctx: E2ETestContext):
created_lookups: list[int] = []
created = await ctx.client.create_session(
on_permission_request=PermissionHandler.approve_all,
tools=[lookup_fruit("apple", created_lookups)],
)
session_id = created.session_id
first = await created.send_and_wait(FRUIT_PROMPT, timeout=30)
assert first is not None
assert "apple" in (first.data.content or "")
assert created_lookups == [42]
await created.disconnect()
fruit_lookups: list[int] = []
vegetable_lookups: list[int] = []
resumed = await ctx.client.resume_session(
session_id,
on_permission_request=PermissionHandler.approve_all,
tools=[lookup_fruit("apple", fruit_lookups)],
)
await resumed.set_tools([lookup_vegetable(vegetable_lookups)])
answer = await resumed.send_and_wait(VEGETABLE_PROMPT, timeout=30)
assert answer is not None
assert "carrot" in (answer.data.content or "")
assert vegetable_lookups == [7]
assert fruit_lookups == []
exchanges = await ctx.get_exchanges()
replaced_from = index_of_prompt(exchanges, VEGETABLE_PROMPT)
assert replaced_from > 0
for exchange in exchanges[replaced_from:]:
tools = offered_tools(exchange)
assert "lookup_vegetable" in tools
assert "lookup_fruit" not in tools
await resumed.disconnect()
async def test_keeps_the_previous_tools_when_a_replacement_is_rejected(
self, ctx: E2ETestContext
):
original_lookups: list[int] = []
replacement_lookups: list[int] = []
invalid_calls: list[int] = []
session = await ctx.client.create_session(
on_permission_request=PermissionHandler.approve_all,
tools=[lookup_fruit("apple", original_lookups)],
)
with pytest.raises(Exception):
await session.set_tools(
[
lookup_fruit("dragonfruit", replacement_lookups),
invalid_tool(invalid_calls),
]
)
answer = await session.send_and_wait(FRUIT_PROMPT, timeout=30)
assert answer is not None
assert "apple" in (answer.data.content or "")
assert original_lookups == [42]
assert replacement_lookups == []
assert invalid_calls == []
await session.disconnect()