forked from github/copilot-sdk
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest_runtime_host_e2e.py
More file actions
310 lines (280 loc) · 12.4 KB
/
Copy pathtest_runtime_host_e2e.py
File metadata and controls
310 lines (280 loc) · 12.4 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
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
"""All SDKs replay the same application-owned AHP hosting conversations."""
import asyncio
import json
import os
from contextlib import asynccontextmanager
import pytest
import pytest_asyncio
from pydantic import BaseModel, Field
from copilot import (
AhpHost,
AhpHostOptions,
AhpSessionCreateRequest,
AhpSessionResumeRequest,
CopilotClient,
CopilotSession,
define_tool,
)
from copilot.session import PermissionHandler
from copilot.tools import ToolInvocation
from .testharness import E2ETestContext
from .testharness.ahp import AhpTestClient
from .testharness.context import DEFAULT_GITHUB_TOKEN
pytestmark = [
pytest.mark.asyncio(loop_scope="module"),
pytest.mark.skipif(
os.environ.get("COPILOT_RUNTIME_HOST_E2E") != "1",
reason="Requires an integrated runtime; set COPILOT_RUNTIME_HOST_E2E=1",
),
]
TOOL_PROMPT = "Use the magic_number tool with seed 'hello' and tell me the result"
COMPOSED_PROMPT = (
"Call magic_number with seed 'hello' and client_echo with text 'ping', then report both results"
)
MARKER = "APPLICATION_OWNED_AHP_PROMPT"
@pytest_asyncio.fixture(loop_scope="module")
async def ahp():
async with AhpTestClient() as client:
yield client
class Seed(BaseModel):
seed: str = Field(description="A seed value")
class Application:
def __init__(self, client: CopilotClient, work_dir: str):
self.client = client
self.work_dir = work_dir
self.sessions: list[CopilotSession] = []
self.releases: list[CopilotSession] = []
self.release_event = asyncio.Event()
self.exit_event = asyncio.Event()
self.exits = []
self.tool_calls: list[str] = []
self.hook_calls: list[str] = []
self.create_calls = 0
self.resume_calls = 0
@define_tool("magic_number", description="Returns a magic number")
def magic_number(params: Seed, invocation: ToolInvocation) -> str:
assert params.seed == "hello"
self.tool_calls.append(invocation.session_id)
return f"MAGIC_{params.seed}_42"
self.tool = magic_number
async def pre_tool(self, _input, invocation):
self.hook_calls.append(invocation["session_id"])
return None
def config(self, requested: dict) -> dict:
return {
**requested,
"on_permission_request": PermissionHandler.approve_all,
"system_message": {"mode": "append", "content": MARKER},
"tools": [self.tool],
"hooks": {"on_pre_tool_use": self.pre_tool},
}
async def create(self, request: AhpSessionCreateRequest) -> CopilotSession:
self.create_calls += 1
assert not request.cancellation_event.is_set()
assert request.config["working_directory"] == self.work_dir
session = await self.client.create_session(**self.config(request.config))
self.sessions.append(session)
return session
async def resume(self, request: AhpSessionResumeRequest) -> CopilotSession:
self.resume_calls += 1
assert not request.cancellation_event.is_set()
assert request.config["continue_pending_work"] is False
assert request.config["working_directory"] == self.work_dir
session = await self.client.resume_session(
request.session_id, **self.config(request.config)
)
self.sessions.append(session)
return session
def release(self, session: CopilotSession) -> None:
self.releases.append(session)
self.release_event.set()
def exited(self, event) -> None:
self.exits.append(event)
self.exit_event.set()
def options(self) -> AhpHostOptions:
from copilot.rpc import HostLocalServerOptions
return AhpHostOptions(
local_server=HostLocalServerOptions(),
create_session=self.create,
resume_session=self.resume,
on_session_released=self.release,
on_exit=self.exited,
)
@asynccontextmanager
async def connect(ahp: AhpTestClient, host: AhpHost, client_id: str | None = None):
assert host.url is not None
assert host.environment_id is None
command = {
"op": "connect",
"url": host.url,
"token": host.token,
"githubToken": DEFAULT_GITHUB_TOKEN,
}
if client_id is not None:
command["clientId"] = client_id
connection = await ahp.request(command)
try:
yield connection["clientId"]
finally:
await ahp.request({"op": "close", "clientId": connection["clientId"]})
async def assert_tools(ctx: E2ETestContext, app: Application, session_id: str):
assert app.tool_calls == [session_id]
assert session_id in app.hook_calls
exchanges = await ctx.get_exchanges()
assert any(MARKER in json.dumps(exchange["request"]["messages"]) for exchange in exchanges)
assert any(
tool["function"]["name"] == "magic_number"
for exchange in exchanges
for tool in exchange["request"]["tools"]
)
class TestRuntimeHost:
async def test_creates_application_session_and_preserves_callbacks(
self, ctx: E2ETestContext, ahp: AhpTestClient
):
await ctx.configure_for_test(
"multi_client", "both_clients_see_tool_request_and_completion_events"
)
app = Application(ctx.client, ctx.work_dir)
try:
async with await ctx.client.start_ahp_host(app.options()) as host:
assert host.pid is None
async with connect(ahp, host) as client_id:
created = await ahp.request(
{"op": "create", "clientId": client_id, "workDir": ctx.work_dir}
)
session_id = created["sessionId"]
assert app.create_calls == 1
assert app.sessions[0].session_id == session_id
response = await ahp.request(
{
"op": "turn",
"clientId": client_id,
"sessionId": session_id,
"prompt": TOOL_PROMPT,
}
)
assert "MAGIC_hello_42" in response["text"]
await assert_tools(ctx, app, session_id)
await asyncio.gather(host.dispose(), host.dispose())
await ahp.request({"op": "stopped", "clientId": client_id, "url": host.url})
await asyncio.wait_for(app.release_event.wait(), 10)
await asyncio.wait_for(app.exit_event.wait(), 10)
assert len(app.releases) == 1
assert app.releases[0] is app.sessions[0]
assert len(app.exits) == 1
assert await app.sessions[0].get_events()
assert (await ctx.client.ping("still alive")).message == "pong: still alive"
finally:
await ctx.client.stop()
async def test_publishes_exact_resident_session_without_factory(
self, ctx: E2ETestContext, ahp: AhpTestClient
):
await ctx.configure_for_test(
"multi_client", "both_clients_see_tool_request_and_completion_events"
)
app = Application(ctx.client, ctx.work_dir)
try:
original = await ctx.client.create_session(
**app.config({"working_directory": ctx.work_dir})
)
async with await ctx.client.start_ahp_host(app.options()) as host:
published = await host.publish_session(original.session_id)
assert published.session_id == original.session_id
assert published.session_uri == f"ahp-session:/{original.session_id}"
async with connect(ahp, host) as client_id:
await ahp.request(
{"op": "attach", "clientId": client_id, "sessionId": original.session_id}
)
response = await ahp.request(
{
"op": "turn",
"clientId": client_id,
"sessionId": original.session_id,
"prompt": TOOL_PROMPT,
}
)
assert "MAGIC_hello_42" in response["text"]
await assert_tools(ctx, app, original.session_id)
await host.dispose()
await ahp.request({"op": "stopped", "clientId": client_id, "url": host.url})
assert app.create_calls == app.resume_calls == 0
assert not app.releases
assert await original.get_events()
assert (await ctx.client.ping("still alive")).message == "pong: still alive"
finally:
await ctx.client.stop()
async def test_resumes_after_runtime_restart_and_composes_tools(
self, ctx: E2ETestContext, ahp: AhpTestClient
):
await ctx.configure_for_test(
"runtime_host", "app_resume_callback_composes_tools_after_history"
)
app = Application(ctx.client, ctx.work_dir)
try:
async with await ctx.client.start_ahp_host(app.options()) as host:
async with connect(ahp, host) as client_id:
created = await ahp.request(
{
"op": "create",
"clientId": client_id,
"workDir": ctx.work_dir,
"clientTools": True,
}
)
session_id = created["sessionId"]
first = await ahp.request(
{
"op": "turn",
"clientId": client_id,
"sessionId": session_id,
"prompt": "What is 2+2?",
}
)
assert "4" in first["text"]
await host.dispose()
await ahp.request({"op": "stopped", "clientId": client_id, "url": host.url})
await asyncio.wait_for(app.release_event.wait(), 10)
assert app.releases[0] is app.sessions[0]
await ctx.client.stop()
resumed = Application(ctx.client, ctx.work_dir)
async with await ctx.client.start_ahp_host(resumed.options()) as replacement:
assert replacement.pid is None
async with connect(ahp, replacement, client_id) as resumed_client_id:
attached = await ahp.request(
{
"op": "attach",
"clientId": resumed_client_id,
"sessionId": session_id,
"clientTools": True,
}
)
assert [turn["message"]["text"] for turn in attached["history"]] == [
"What is 2+2?"
]
assert resumed.create_calls == 0
assert resumed.resume_calls == 1
assert resumed.sessions[0].session_id == session_id
assert resumed.sessions[0] is not app.sessions[0]
response = await ahp.request(
{
"op": "turn",
"clientId": resumed_client_id,
"sessionId": session_id,
"prompt": COMPOSED_PROMPT,
"clientTools": True,
}
)
assert "MAGIC_hello_42" in response["text"]
assert "CLIENT_ECHO_ping" in response["text"]
assert response["clientToolCalls"] == 1
await assert_tools(ctx, resumed, session_id)
await replacement.dispose()
await ahp.request(
{"op": "stopped", "clientId": resumed_client_id, "url": replacement.url}
)
await asyncio.wait_for(resumed.release_event.wait(), 10)
assert len(resumed.releases) == 1
assert resumed.releases[0] is resumed.sessions[0]
assert await resumed.sessions[0].get_events()
finally:
await ctx.client.stop()