ssebench-sdk 1.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- sse/__init__.py +18 -0
- sse/ai.py +471 -0
- sse/cheating.py +31 -0
- sse/daemon.py +144 -0
- sse/error.py +18 -0
- sse/grading.py +178 -0
- sse/helper.py +66 -0
- sse/metadata.py +74 -0
- sse/project.py +78 -0
- sse/prompt.py +148 -0
- sse/py.typed +0 -0
- sse/reference.py +29 -0
- sse/tools/__init__.py +2 -0
- sse/tools/bash.py +15 -0
- sse/tools/bencher.py +90 -0
- ssebench_sdk-1.0.0.dist-info/METADATA +95 -0
- ssebench_sdk-1.0.0.dist-info/RECORD +21 -0
- ssebench_sdk-1.0.0.dist-info/WHEEL +5 -0
- ssebench_sdk-1.0.0.dist-info/licenses/LICENSE +202 -0
- ssebench_sdk-1.0.0.dist-info/licenses/NOTICE +10 -0
- ssebench_sdk-1.0.0.dist-info/top_level.txt +1 -0
sse/__init__.py
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
1
|
+
"""SSEBench SDK: the task, and the build, test and grading actions of the daemon, for code in a task container.
|
|
2
|
+
|
|
3
|
+
Agents, plugins and the evaluator run inside the task container and import this package as
|
|
4
|
+
``sse``. It talks to ``ssebench-daemon`` over the Unix socket in ``SSE_DAEMON_SOCKET``; see
|
|
5
|
+
:mod:`sse.daemon`. ``import sse`` loads :mod:`sse.ai`, :mod:`sse.grading`, :mod:`sse.reference`
|
|
6
|
+
and :mod:`sse.tools` without contacting the daemon; :mod:`sse.project` and :mod:`sse.prompt` ask
|
|
7
|
+
the daemon for the task when they are imported.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from importlib.metadata import version
|
|
11
|
+
|
|
12
|
+
from . import ai as ai
|
|
13
|
+
from . import grading as grading
|
|
14
|
+
from . import reference as reference
|
|
15
|
+
from . import tools as tools
|
|
16
|
+
from .error import SDKError as SDKError
|
|
17
|
+
|
|
18
|
+
__version__ = version("ssebench.sdk")
|
sse/ai.py
ADDED
|
@@ -0,0 +1,471 @@
|
|
|
1
|
+
"""OpenCode agent wrapper.
|
|
2
|
+
|
|
3
|
+
Provides a reusable async client for starting an OpenCode server,
|
|
4
|
+
creating sessions, sending prompts, and collecting responses.
|
|
5
|
+
|
|
6
|
+
Usage::
|
|
7
|
+
|
|
8
|
+
from sse.ai import OpenCodeAgent, build_opencode_config
|
|
9
|
+
|
|
10
|
+
config = build_opencode_config("claude-sonnet-4-20250514")
|
|
11
|
+
async with OpenCodeAgent(directory="/path/to/project", config=config) as agent:
|
|
12
|
+
session_id = await agent.create_session()
|
|
13
|
+
response = await agent.send_prompt(session_id, "hello")
|
|
14
|
+
print(response.final_message)
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
from __future__ import annotations
|
|
18
|
+
|
|
19
|
+
import asyncio
|
|
20
|
+
import json
|
|
21
|
+
import logging
|
|
22
|
+
import os
|
|
23
|
+
import time
|
|
24
|
+
from dataclasses import dataclass
|
|
25
|
+
from typing import Any, Literal, Required
|
|
26
|
+
|
|
27
|
+
import httpx
|
|
28
|
+
from typing_extensions import TypedDict
|
|
29
|
+
|
|
30
|
+
log = logging.getLogger(__name__)
|
|
31
|
+
|
|
32
|
+
# ---------------------------------------------------------------------------
|
|
33
|
+
# API type definitions (matching OpenCode OpenAPI 3.1 spec)
|
|
34
|
+
# ---------------------------------------------------------------------------
|
|
35
|
+
|
|
36
|
+
# -- Request types --
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
class _TextPartInput(TypedDict):
|
|
40
|
+
"""A text part in a message request body."""
|
|
41
|
+
|
|
42
|
+
type: Literal["text"]
|
|
43
|
+
text: str
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class _SendMessageBody(TypedDict, total=False):
|
|
47
|
+
"""Request body for ``POST /session/:id/message``."""
|
|
48
|
+
|
|
49
|
+
parts: Required[list[_TextPartInput]]
|
|
50
|
+
messageID: str
|
|
51
|
+
agent: str
|
|
52
|
+
noReply: bool
|
|
53
|
+
system: str
|
|
54
|
+
tools: dict[str, bool]
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
class _CreateSessionBody(TypedDict, total=False):
|
|
58
|
+
"""Request body for ``POST /session``."""
|
|
59
|
+
|
|
60
|
+
title: str
|
|
61
|
+
parentID: str
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
# -- Response types --
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
class _SessionResponse(TypedDict):
|
|
68
|
+
"""Subset of the ``Session`` schema returned by ``POST /session``."""
|
|
69
|
+
|
|
70
|
+
id: str
|
|
71
|
+
title: str
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
class _PartResponse(TypedDict, total=False):
|
|
75
|
+
"""Union of all ``Part`` variants.
|
|
76
|
+
|
|
77
|
+
Only the common fields and the ``text`` field (for ``TextPart``) are
|
|
78
|
+
typed here; other variant-specific fields are accessed via ``dict``
|
|
79
|
+
when needed.
|
|
80
|
+
"""
|
|
81
|
+
|
|
82
|
+
id: Required[str]
|
|
83
|
+
sessionID: Required[str]
|
|
84
|
+
messageID: Required[str]
|
|
85
|
+
type: Required[str] # discriminator: "text", "tool", "reasoning", ...
|
|
86
|
+
text: str # present only when type == "text"
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
class _TokenCache(TypedDict):
|
|
90
|
+
"""Cache token counts inside ``_TokenUsage``."""
|
|
91
|
+
|
|
92
|
+
read: int
|
|
93
|
+
write: int
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
class _TokenUsage(TypedDict, total=False):
|
|
97
|
+
"""Token usage reported in an ``AssistantMessage``."""
|
|
98
|
+
|
|
99
|
+
input: Required[int]
|
|
100
|
+
output: Required[int]
|
|
101
|
+
reasoning: Required[int]
|
|
102
|
+
cache: Required[_TokenCache]
|
|
103
|
+
total: int
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
class _AssistantMessageResponse(TypedDict, total=False):
|
|
107
|
+
"""Subset of the ``AssistantMessage`` schema we inspect."""
|
|
108
|
+
|
|
109
|
+
id: Required[str]
|
|
110
|
+
sessionID: Required[str]
|
|
111
|
+
role: Required[Literal["assistant"]]
|
|
112
|
+
cost: Required[float]
|
|
113
|
+
tokens: Required[_TokenUsage]
|
|
114
|
+
modelID: str
|
|
115
|
+
providerID: str
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
class _SyncMessageResponse(TypedDict):
|
|
119
|
+
"""Response from ``POST /session/:id/message`` (synchronous)."""
|
|
120
|
+
|
|
121
|
+
info: _AssistantMessageResponse
|
|
122
|
+
parts: list[_PartResponse]
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
class _MessageEntry(TypedDict):
|
|
126
|
+
"""One element returned by ``GET /session/:id/message``.
|
|
127
|
+
|
|
128
|
+
The ``info`` field can be either a ``UserMessage`` or an
|
|
129
|
+
``AssistantMessage``; we only inspect the ``role`` field here.
|
|
130
|
+
"""
|
|
131
|
+
|
|
132
|
+
info: dict[str, Any]
|
|
133
|
+
parts: list[_PartResponse]
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
# ---------------------------------------------------------------------------
|
|
137
|
+
# Public data structures
|
|
138
|
+
# ---------------------------------------------------------------------------
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
@dataclass(frozen=True)
|
|
142
|
+
class AgentResponse:
|
|
143
|
+
"""Structured response from an OpenCode agent session."""
|
|
144
|
+
|
|
145
|
+
final_message: str
|
|
146
|
+
"""Text parts from the last assistant message only."""
|
|
147
|
+
full_log: str
|
|
148
|
+
"""All text parts from every message in the session."""
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
# ---------------------------------------------------------------------------
|
|
152
|
+
# Configuration helper
|
|
153
|
+
# ---------------------------------------------------------------------------
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
def build_opencode_config(
|
|
157
|
+
model_name: str | None = None,
|
|
158
|
+
mcp_url: str | None = None,
|
|
159
|
+
) -> dict[str, Any]:
|
|
160
|
+
"""Build a complete OpenCode config.
|
|
161
|
+
|
|
162
|
+
When ``SSE_MODEL_NAME``, ``SSE_BASE_URL`` and ``SSE_API_KEY`` are set, as they are in a task
|
|
163
|
+
container, the config uses the run's model through the LiteLLM proxy and ignores
|
|
164
|
+
``model_name``. Otherwise it uses ``model_name`` with OpenCode's built-in Anthropic provider,
|
|
165
|
+
which reads the ``ANTHROPIC_API_KEY`` environment variable at runtime -- the key is **not**
|
|
166
|
+
embedded in the config dict.
|
|
167
|
+
|
|
168
|
+
Args:
|
|
169
|
+
model_name: Anthropic model identifier (e.g. ``claude-sonnet-4-20250514``); required
|
|
170
|
+
when the ``SSE_*`` variables are not all set.
|
|
171
|
+
mcp_url: Optional MCP server URL to register.
|
|
172
|
+
|
|
173
|
+
Returns:
|
|
174
|
+
Config dict suitable for the ``OPENCODE_CONFIG_CONTENT`` env var.
|
|
175
|
+
|
|
176
|
+
Raises:
|
|
177
|
+
ValueError: If ``model_name`` is missing and the ``SSE_*`` variables are not all set.
|
|
178
|
+
"""
|
|
179
|
+
config: dict[str, Any] = {}
|
|
180
|
+
|
|
181
|
+
sse_model_name = os.environ.get("SSE_MODEL_NAME")
|
|
182
|
+
sse_base_url = os.environ.get("SSE_BASE_URL")
|
|
183
|
+
sse_api_key = os.environ.get("SSE_API_KEY")
|
|
184
|
+
|
|
185
|
+
if sse_model_name and sse_base_url and sse_api_key:
|
|
186
|
+
config["provider"] = {
|
|
187
|
+
"ssebench": {
|
|
188
|
+
"npm": "@ai-sdk/openai-compatible",
|
|
189
|
+
"name": "SSEBench built-in provider",
|
|
190
|
+
"options": {
|
|
191
|
+
"baseURL": sse_base_url,
|
|
192
|
+
"apiKey": sse_api_key,
|
|
193
|
+
},
|
|
194
|
+
"models": {sse_model_name: {"name": "SSEBench built-in model"}},
|
|
195
|
+
}
|
|
196
|
+
}
|
|
197
|
+
config["model"] = f"ssebench/{sse_model_name}"
|
|
198
|
+
else:
|
|
199
|
+
if model_name is None:
|
|
200
|
+
raise ValueError(
|
|
201
|
+
"model_name is required when SSE_MODEL_NAME, SSE_BASE_URL, "
|
|
202
|
+
"and SSE_API_KEY environment variables are not all set"
|
|
203
|
+
)
|
|
204
|
+
config["model"] = f"anthropic/{model_name}"
|
|
205
|
+
|
|
206
|
+
if mcp_url:
|
|
207
|
+
config["mcp"] = {
|
|
208
|
+
"ssebench": {
|
|
209
|
+
"type": "remote",
|
|
210
|
+
"url": mcp_url,
|
|
211
|
+
}
|
|
212
|
+
}
|
|
213
|
+
|
|
214
|
+
config["permission"] = "allow"
|
|
215
|
+
|
|
216
|
+
return config
|
|
217
|
+
|
|
218
|
+
|
|
219
|
+
# ---------------------------------------------------------------------------
|
|
220
|
+
# Text extraction helpers
|
|
221
|
+
# ---------------------------------------------------------------------------
|
|
222
|
+
|
|
223
|
+
|
|
224
|
+
def _extract_text(parts: list[_PartResponse]) -> str:
|
|
225
|
+
"""Join text content from a list of message parts.
|
|
226
|
+
|
|
227
|
+
Only parts with ``type == "text"`` are included; empty/whitespace-only
|
|
228
|
+
parts are skipped.
|
|
229
|
+
"""
|
|
230
|
+
return "\n".join(
|
|
231
|
+
text
|
|
232
|
+
for p in parts
|
|
233
|
+
if p["type"] == "text" and (text := p.get("text", "")).strip()
|
|
234
|
+
)
|
|
235
|
+
|
|
236
|
+
|
|
237
|
+
def _extract_all_text(messages: list[_MessageEntry]) -> str:
|
|
238
|
+
"""Join all text content across every message in a session."""
|
|
239
|
+
segments: list[str] = []
|
|
240
|
+
for entry in messages:
|
|
241
|
+
text = _extract_text(entry["parts"])
|
|
242
|
+
if text:
|
|
243
|
+
segments.append(text)
|
|
244
|
+
return "\n".join(segments)
|
|
245
|
+
|
|
246
|
+
|
|
247
|
+
# ---------------------------------------------------------------------------
|
|
248
|
+
# Agent
|
|
249
|
+
# ---------------------------------------------------------------------------
|
|
250
|
+
|
|
251
|
+
|
|
252
|
+
class OpenCodeAgent:
|
|
253
|
+
"""Async wrapper for the OpenCode agent server.
|
|
254
|
+
|
|
255
|
+
Manages the full lifecycle: start the server, create sessions,
|
|
256
|
+
send prompts, collect responses, and shut down cleanly.
|
|
257
|
+
|
|
258
|
+
With a config that uses OpenCode's built-in Anthropic provider (see
|
|
259
|
+
:func:`build_opencode_config`), the Anthropic API key must be available as
|
|
260
|
+
the ``ANTHROPIC_API_KEY`` environment variable in the process that runs
|
|
261
|
+
this agent.
|
|
262
|
+
|
|
263
|
+
Usage::
|
|
264
|
+
|
|
265
|
+
from sse.ai import OpenCodeAgent, build_opencode_config
|
|
266
|
+
|
|
267
|
+
config = build_opencode_config("claude-sonnet-4-20250514")
|
|
268
|
+
async with OpenCodeAgent(directory="/path/to/project", config=config) as agent:
|
|
269
|
+
session_id = await agent.create_session()
|
|
270
|
+
response = await agent.send_prompt(session_id, "hello")
|
|
271
|
+
print(response)
|
|
272
|
+
"""
|
|
273
|
+
|
|
274
|
+
def __init__(
|
|
275
|
+
self,
|
|
276
|
+
directory: str,
|
|
277
|
+
config: dict[str, Any],
|
|
278
|
+
port: int = 4097,
|
|
279
|
+
hostname: str = "127.0.0.1",
|
|
280
|
+
) -> None:
|
|
281
|
+
self.directory = directory
|
|
282
|
+
self.config = config
|
|
283
|
+
self.port = port
|
|
284
|
+
self.hostname = hostname
|
|
285
|
+
self.base_url = f"http://{hostname}:{port}"
|
|
286
|
+
self._process: asyncio.subprocess.Process | None = None
|
|
287
|
+
self._client: httpx.AsyncClient | None = None
|
|
288
|
+
|
|
289
|
+
# ------------------------------------------------------------------
|
|
290
|
+
# Lifecycle
|
|
291
|
+
# ------------------------------------------------------------------
|
|
292
|
+
|
|
293
|
+
async def start(self, timeout: int = 120) -> None:
|
|
294
|
+
"""Start the OpenCode server and block until it is ready.
|
|
295
|
+
|
|
296
|
+
Args:
|
|
297
|
+
timeout: Maximum seconds to wait for server readiness.
|
|
298
|
+
|
|
299
|
+
Raises:
|
|
300
|
+
RuntimeError: If the server process exits prematurely.
|
|
301
|
+
TimeoutError: If the server does not become ready in time.
|
|
302
|
+
"""
|
|
303
|
+
env = {**os.environ, "OPENCODE_CONFIG_CONTENT": json.dumps(self.config)}
|
|
304
|
+
self._process = await asyncio.create_subprocess_exec(
|
|
305
|
+
"opencode",
|
|
306
|
+
"serve",
|
|
307
|
+
"--port",
|
|
308
|
+
str(self.port),
|
|
309
|
+
"--hostname",
|
|
310
|
+
self.hostname,
|
|
311
|
+
cwd=self.directory,
|
|
312
|
+
env=env,
|
|
313
|
+
stdout=asyncio.subprocess.DEVNULL,
|
|
314
|
+
stderr=asyncio.subprocess.DEVNULL,
|
|
315
|
+
)
|
|
316
|
+
log.info("Server started (PID %s) on port %s", self._process.pid, self.port)
|
|
317
|
+
|
|
318
|
+
self._client = httpx.AsyncClient(
|
|
319
|
+
base_url=self.base_url,
|
|
320
|
+
headers={"x-opencode-directory": self.directory},
|
|
321
|
+
)
|
|
322
|
+
|
|
323
|
+
try:
|
|
324
|
+
start_time = time.monotonic()
|
|
325
|
+
while time.monotonic() - start_time < timeout:
|
|
326
|
+
try:
|
|
327
|
+
resp = await self._client.get("/session")
|
|
328
|
+
if resp.status_code < 400:
|
|
329
|
+
elapsed = int(time.monotonic() - start_time)
|
|
330
|
+
log.info("Server ready after %ss", elapsed)
|
|
331
|
+
return
|
|
332
|
+
except (httpx.ConnectError, httpx.ReadError, httpx.ConnectTimeout):
|
|
333
|
+
pass
|
|
334
|
+
if self._process.returncode is not None:
|
|
335
|
+
raise RuntimeError(
|
|
336
|
+
f"OpenCode server exited prematurely "
|
|
337
|
+
f"with code {self._process.returncode}"
|
|
338
|
+
)
|
|
339
|
+
await asyncio.sleep(1)
|
|
340
|
+
|
|
341
|
+
raise TimeoutError(
|
|
342
|
+
f"OpenCode server did not become ready within {timeout}s"
|
|
343
|
+
)
|
|
344
|
+
except BaseException:
|
|
345
|
+
await self.stop()
|
|
346
|
+
raise
|
|
347
|
+
|
|
348
|
+
async def stop(self) -> None:
|
|
349
|
+
"""Shut down the OpenCode server and close the HTTP client."""
|
|
350
|
+
try:
|
|
351
|
+
if self._client:
|
|
352
|
+
await self._client.aclose()
|
|
353
|
+
finally:
|
|
354
|
+
self._client = None
|
|
355
|
+
proc = self._process
|
|
356
|
+
self._process = None
|
|
357
|
+
if proc and proc.returncode is None:
|
|
358
|
+
proc.terminate()
|
|
359
|
+
try:
|
|
360
|
+
await asyncio.wait_for(proc.wait(), timeout=5)
|
|
361
|
+
except TimeoutError:
|
|
362
|
+
proc.kill()
|
|
363
|
+
await proc.wait()
|
|
364
|
+
log.info("Server stopped")
|
|
365
|
+
|
|
366
|
+
async def __aenter__(self) -> OpenCodeAgent:
|
|
367
|
+
await self.start()
|
|
368
|
+
return self
|
|
369
|
+
|
|
370
|
+
async def __aexit__(self, *args: Any) -> None:
|
|
371
|
+
await self.stop()
|
|
372
|
+
|
|
373
|
+
# ------------------------------------------------------------------
|
|
374
|
+
# HTTP client accessor
|
|
375
|
+
# ------------------------------------------------------------------
|
|
376
|
+
|
|
377
|
+
@property
|
|
378
|
+
def _http(self) -> httpx.AsyncClient:
|
|
379
|
+
"""Return the HTTP client, raising if the agent has not been started."""
|
|
380
|
+
if self._client is None:
|
|
381
|
+
raise RuntimeError("Agent not started -- call start() first")
|
|
382
|
+
return self._client
|
|
383
|
+
|
|
384
|
+
# ------------------------------------------------------------------
|
|
385
|
+
# Session management
|
|
386
|
+
# ------------------------------------------------------------------
|
|
387
|
+
|
|
388
|
+
async def create_session(self, title: str = "Session") -> str:
|
|
389
|
+
"""Create a new chat session.
|
|
390
|
+
|
|
391
|
+
Args:
|
|
392
|
+
title: Human-readable session title.
|
|
393
|
+
|
|
394
|
+
Returns:
|
|
395
|
+
The session ID string.
|
|
396
|
+
|
|
397
|
+
Raises:
|
|
398
|
+
httpx.HTTPStatusError: If session creation fails.
|
|
399
|
+
"""
|
|
400
|
+
body: _CreateSessionBody = {"title": title}
|
|
401
|
+
resp = await self._http.post("/session", json=body)
|
|
402
|
+
resp.raise_for_status()
|
|
403
|
+
session: _SessionResponse = resp.json()
|
|
404
|
+
log.info("Session created: %s", session["id"])
|
|
405
|
+
return session["id"]
|
|
406
|
+
|
|
407
|
+
# ------------------------------------------------------------------
|
|
408
|
+
# Prompt and response
|
|
409
|
+
# ------------------------------------------------------------------
|
|
410
|
+
|
|
411
|
+
async def send_prompt(
|
|
412
|
+
self,
|
|
413
|
+
session_id: str,
|
|
414
|
+
text: str,
|
|
415
|
+
timeout: float = 300.0,
|
|
416
|
+
) -> AgentResponse:
|
|
417
|
+
"""Send a prompt and wait for the complete response.
|
|
418
|
+
|
|
419
|
+
Uses the synchronous ``POST /session/:id/message`` endpoint which
|
|
420
|
+
blocks until the assistant finishes responding. After receiving the
|
|
421
|
+
response, all session messages are fetched to build the full
|
|
422
|
+
conversation log.
|
|
423
|
+
|
|
424
|
+
Args:
|
|
425
|
+
session_id: Target session ID (from :meth:`create_session`).
|
|
426
|
+
text: The prompt text.
|
|
427
|
+
timeout: Maximum seconds to wait for the assistant to finish.
|
|
428
|
+
If exceeded, ``final_message`` will contain a timeout
|
|
429
|
+
indicator and ``full_log`` will contain whatever messages
|
|
430
|
+
were available at that point.
|
|
431
|
+
|
|
432
|
+
Returns:
|
|
433
|
+
An :class:`AgentResponse` containing the final assistant
|
|
434
|
+
message text (or a timeout indicator) and the full
|
|
435
|
+
conversation log.
|
|
436
|
+
|
|
437
|
+
Raises:
|
|
438
|
+
httpx.HTTPStatusError: On HTTP-level failures from the server.
|
|
439
|
+
"""
|
|
440
|
+
body: _SendMessageBody = {"parts": [{"type": "text", "text": text}]}
|
|
441
|
+
|
|
442
|
+
timed_out = False
|
|
443
|
+
try:
|
|
444
|
+
resp = await self._http.post(
|
|
445
|
+
f"/session/{session_id}/message",
|
|
446
|
+
json=body,
|
|
447
|
+
timeout=httpx.Timeout(timeout, connect=5.0),
|
|
448
|
+
)
|
|
449
|
+
resp.raise_for_status()
|
|
450
|
+
result: _SyncMessageResponse = resp.json()
|
|
451
|
+
final_message = _extract_text(result["parts"])
|
|
452
|
+
except httpx.TimeoutException:
|
|
453
|
+
timed_out = True
|
|
454
|
+
final_message = f"[timeout] Response not completed within {timeout}s"
|
|
455
|
+
|
|
456
|
+
log.info("Prompt sent to session %s", session_id)
|
|
457
|
+
|
|
458
|
+
# Always fetch the full conversation log -- partial progress is
|
|
459
|
+
# still valuable when the response timed out.
|
|
460
|
+
all_resp = await self._http.get(
|
|
461
|
+
f"/session/{session_id}/message",
|
|
462
|
+
timeout=httpx.Timeout(30.0, connect=5.0),
|
|
463
|
+
)
|
|
464
|
+
all_resp.raise_for_status()
|
|
465
|
+
messages: list[_MessageEntry] = all_resp.json()
|
|
466
|
+
full_log = _extract_all_text(messages)
|
|
467
|
+
|
|
468
|
+
if timed_out:
|
|
469
|
+
log.warning("Timed out after %ss, returning partial log", timeout)
|
|
470
|
+
|
|
471
|
+
return AgentResponse(final_message=final_message, full_log=full_log)
|
sse/cheating.py
ADDED
|
@@ -0,0 +1,31 @@
|
|
|
1
|
+
"""Deprecated alias for :mod:`sse.reference`.
|
|
2
|
+
|
|
3
|
+
Kept for one release. Import :mod:`sse.reference` and call
|
|
4
|
+
:func:`sse.reference.get_reference_patch` instead.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import warnings
|
|
10
|
+
|
|
11
|
+
from sse.reference import get_reference_patch
|
|
12
|
+
|
|
13
|
+
warnings.warn(
|
|
14
|
+
"sse.cheating is deprecated; use sse.reference instead",
|
|
15
|
+
DeprecationWarning,
|
|
16
|
+
stacklevel=2,
|
|
17
|
+
)
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def get_ground_truth() -> str:
|
|
21
|
+
"""Deprecated alias for :func:`sse.reference.get_reference_patch`."""
|
|
22
|
+
warnings.warn(
|
|
23
|
+
"sse.cheating.get_ground_truth is deprecated; "
|
|
24
|
+
"use sse.reference.get_reference_patch instead",
|
|
25
|
+
DeprecationWarning,
|
|
26
|
+
stacklevel=2,
|
|
27
|
+
)
|
|
28
|
+
return get_reference_patch()
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
__all__ = ["get_ground_truth"]
|
sse/daemon.py
ADDED
|
@@ -0,0 +1,144 @@
|
|
|
1
|
+
"""HTTP client of ``ssebench-daemon``, which the rest of the SDK goes through.
|
|
2
|
+
|
|
3
|
+
The client connects to the Unix socket in ``SSE_DAEMON_SOCKET``, or, when that is unset, over
|
|
4
|
+
TCP to the ``host:port`` in ``SSE_AGENT_DOCKER``. The daemon's HTTP API is described in
|
|
5
|
+
docs/reference/daemon-api.md.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import logging
|
|
11
|
+
import os
|
|
12
|
+
import urllib.parse
|
|
13
|
+
from importlib.metadata import version
|
|
14
|
+
|
|
15
|
+
import requests
|
|
16
|
+
import requests_unixsocket
|
|
17
|
+
from packaging.version import InvalidVersion, Version
|
|
18
|
+
|
|
19
|
+
from sse.error import SDKError
|
|
20
|
+
|
|
21
|
+
__version__ = version("ssebench.sdk")
|
|
22
|
+
|
|
23
|
+
log = logging.getLogger(__name__)
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def same_version(sdk_version: str, daemon_version: object) -> bool:
|
|
27
|
+
"""Compare the SDK's PEP 440 version with the daemon's SemVer one.
|
|
28
|
+
|
|
29
|
+
Both come from the same release, spelled per ecosystem (1.0.0rc1 and
|
|
30
|
+
1.0.0-rc.1); PEP 440 parsing normalizes the SemVer spelling.
|
|
31
|
+
"""
|
|
32
|
+
try:
|
|
33
|
+
return Version(str(daemon_version)) == Version(sdk_version)
|
|
34
|
+
except InvalidVersion:
|
|
35
|
+
return False
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
class Daemon:
|
|
39
|
+
"""A connection to the daemon.
|
|
40
|
+
|
|
41
|
+
Creating one asks the daemon for its version and logs a warning when it differs from the
|
|
42
|
+
SDK's. It connects to ``socket`` when given, else to ``SSE_DAEMON_SOCKET``, else to the
|
|
43
|
+
daemon at ``SSE_AGENT_DOCKER``. Raises ``Exception`` when there is none of them, and
|
|
44
|
+
``RuntimeError`` when the daemon does not answer.
|
|
45
|
+
"""
|
|
46
|
+
|
|
47
|
+
def __init__(self, socket: str | None = None):
|
|
48
|
+
if (sock := socket or os.getenv("SSE_DAEMON_SOCKET")) is not None:
|
|
49
|
+
self.conn = requests_unixsocket.Session()
|
|
50
|
+
self.base_url = "http+unix://" + urllib.parse.quote_plus(sock)
|
|
51
|
+
elif os.getenv("SSE_AGENT_DOCKER") is not None:
|
|
52
|
+
docker_host = os.getenv("SSE_AGENT_DOCKER")
|
|
53
|
+
self.conn = requests.Session()
|
|
54
|
+
self.base_url = f"http://{docker_host}"
|
|
55
|
+
else:
|
|
56
|
+
raise Exception("Cannot establish connection to daemon")
|
|
57
|
+
|
|
58
|
+
# Version check: SDK and daemon should match
|
|
59
|
+
self._check_version()
|
|
60
|
+
|
|
61
|
+
def _check_version(self):
|
|
62
|
+
"""Check that SDK version matches daemon version."""
|
|
63
|
+
try:
|
|
64
|
+
response = self.conn.get(
|
|
65
|
+
f"{self.base_url}/version", headers={"Connection": "close"}
|
|
66
|
+
)
|
|
67
|
+
data = response.json()
|
|
68
|
+
daemon_version = data.get("version")
|
|
69
|
+
if not same_version(__version__, daemon_version):
|
|
70
|
+
log.warning(
|
|
71
|
+
"SDK version (%s) does not match daemon version (%s)",
|
|
72
|
+
__version__,
|
|
73
|
+
daemon_version,
|
|
74
|
+
)
|
|
75
|
+
except Exception as e:
|
|
76
|
+
raise RuntimeError(f"Failed to verify daemon version: {e}") from e
|
|
77
|
+
|
|
78
|
+
def get(self, path, headers=None):
|
|
79
|
+
"""Send ``GET path`` and return the JSON body.
|
|
80
|
+
|
|
81
|
+
Raises :class:`~sse.error.SDKError` with the daemon's message when it answers with an
|
|
82
|
+
error status, and on network errors.
|
|
83
|
+
"""
|
|
84
|
+
merged_headers = {"Connection": "close"}
|
|
85
|
+
if headers:
|
|
86
|
+
merged_headers.update(headers)
|
|
87
|
+
try:
|
|
88
|
+
response = self.conn.get(f"{self.base_url}{path}", headers=merged_headers)
|
|
89
|
+
data = response.json()
|
|
90
|
+
if not response.ok:
|
|
91
|
+
raise SDKError(data.get("error", "unknown error"))
|
|
92
|
+
return data
|
|
93
|
+
except SDKError:
|
|
94
|
+
raise
|
|
95
|
+
except Exception as e:
|
|
96
|
+
raise SDKError(f"network error: {e}") from e
|
|
97
|
+
|
|
98
|
+
def post(self, path, data, headers=None):
|
|
99
|
+
"""Send ``POST path`` with ``data`` as the JSON body and return the JSON body of the answer.
|
|
100
|
+
|
|
101
|
+
Raises :class:`~sse.error.SDKError` as :meth:`get` does.
|
|
102
|
+
"""
|
|
103
|
+
merged_headers = {"Connection": "close"}
|
|
104
|
+
if headers:
|
|
105
|
+
merged_headers.update(headers)
|
|
106
|
+
try:
|
|
107
|
+
response = self.conn.post(
|
|
108
|
+
f"{self.base_url}{path}", json=data, headers=merged_headers
|
|
109
|
+
)
|
|
110
|
+
result = response.json()
|
|
111
|
+
if not response.ok:
|
|
112
|
+
raise SDKError(result.get("error", "unknown error"))
|
|
113
|
+
return result
|
|
114
|
+
except SDKError:
|
|
115
|
+
raise
|
|
116
|
+
except Exception as e:
|
|
117
|
+
raise SDKError(f"network error: {e}") from e
|
|
118
|
+
|
|
119
|
+
def prepare_grading(self):
|
|
120
|
+
"""``POST /prepare_grading``: save the agent's diff and apply it to the copy that is graded.
|
|
121
|
+
|
|
122
|
+
Only the admin socket accepts it.
|
|
123
|
+
"""
|
|
124
|
+
return self.post("/prepare_grading", {})
|
|
125
|
+
|
|
126
|
+
def version(self):
|
|
127
|
+
"""``GET /version``: the daemon's version."""
|
|
128
|
+
return self.get("/version")
|
|
129
|
+
|
|
130
|
+
def project(self):
|
|
131
|
+
"""``GET /project``: the task's metadata, as :class:`sse.project.Metadata` holds it."""
|
|
132
|
+
return self.get("/project")
|
|
133
|
+
|
|
134
|
+
def capabilities(self):
|
|
135
|
+
"""``GET /capabilities``: the checks the task supports."""
|
|
136
|
+
return self.get("/capabilities")
|
|
137
|
+
|
|
138
|
+
def tool(self, name, action, data):
|
|
139
|
+
"""``POST /tool/{name}?action={action}`` with ``data`` as the JSON body."""
|
|
140
|
+
return self.post(
|
|
141
|
+
f"/tool/{name}?action={action}",
|
|
142
|
+
data,
|
|
143
|
+
{"Content-Type": "application/json"},
|
|
144
|
+
)
|
sse/error.py
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
1
|
+
"""Error types for SSEBench SDK."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class SDKError(Exception):
|
|
7
|
+
"""Raised when the SDK encounters an error from the daemon.
|
|
8
|
+
|
|
9
|
+
This exception is raised for daemon communication failures, config errors,
|
|
10
|
+
file not found, etc. - NOT for script execution failures (non-zero exit code).
|
|
11
|
+
|
|
12
|
+
Script execution results (success or failure) are returned as ScriptResult.
|
|
13
|
+
Use result.is_success() to check if the script succeeded.
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
def __init__(self, message: str):
|
|
17
|
+
self.message = message
|
|
18
|
+
super().__init__(message)
|