cursor-cloud-mcp 0.3.1__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.
- cursor_cloud_mcp/__init__.py +3 -0
- cursor_cloud_mcp/__main__.py +11 -0
- cursor_cloud_mcp/activity.py +158 -0
- cursor_cloud_mcp/artifacts.py +156 -0
- cursor_cloud_mcp/budget.py +34 -0
- cursor_cloud_mcp/catalog.py +155 -0
- cursor_cloud_mcp/client.py +884 -0
- cursor_cloud_mcp/compat.py +74 -0
- cursor_cloud_mcp/config.py +111 -0
- cursor_cloud_mcp/errors.py +140 -0
- cursor_cloud_mcp/fixture.py +393 -0
- cursor_cloud_mcp/models.py +701 -0
- cursor_cloud_mcp/present.py +342 -0
- cursor_cloud_mcp/redaction.py +89 -0
- cursor_cloud_mcp/server.py +776 -0
- cursor_cloud_mcp/sessions.py +410 -0
- cursor_cloud_mcp/shapes.py +35 -0
- cursor_cloud_mcp/slicing.py +22 -0
- cursor_cloud_mcp/stream.py +513 -0
- cursor_cloud_mcp/supervision.py +321 -0
- cursor_cloud_mcp/validation.py +181 -0
- cursor_cloud_mcp-0.3.1.dist-info/METADATA +213 -0
- cursor_cloud_mcp-0.3.1.dist-info/RECORD +26 -0
- cursor_cloud_mcp-0.3.1.dist-info/WHEEL +4 -0
- cursor_cloud_mcp-0.3.1.dist-info/entry_points.txt +2 -0
- cursor_cloud_mcp-0.3.1.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,158 @@
|
|
|
1
|
+
"""Activity summary of a run, built from its event stream.
|
|
2
|
+
|
|
3
|
+
The run object says little: a RUNNING run keeps the ``updatedAt`` of its creation, an ERROR run
|
|
4
|
+
carries neither ``error`` nor ``result``, and FINISHED only means the agent ended its turn. The
|
|
5
|
+
stream says more: event ids are millisecond timestamps (liveness), and tool calls show the last
|
|
6
|
+
command and the background tasks the agent started and awaited. Everything here is the last
|
|
7
|
+
*observed* state, never a guarantee about the VM.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
import datetime as dt
|
|
11
|
+
import re
|
|
12
|
+
from typing import Any
|
|
13
|
+
|
|
14
|
+
from cursor_cloud_mcp.config import ACTIVITY_TEXT_MAX_CHARS
|
|
15
|
+
from cursor_cloud_mcp.models import ActivityView, BackgroundTaskView, RunEventView, ToolCallSummaryView
|
|
16
|
+
from cursor_cloud_mcp.shapes import clip, tool_payload
|
|
17
|
+
|
|
18
|
+
# Stream ids look like Redis stream ids: "<milliseconds>-<sequence>" (checked 2026-10-06).
|
|
19
|
+
_EVENT_ID = re.compile(r"^(\d{12,14})-\d+$")
|
|
20
|
+
# Bounds a parsed timestamp must fall within to be trusted: 2020-01-01 to 2100-01-01.
|
|
21
|
+
_MIN_MS = 1_577_836_800_000
|
|
22
|
+
_MAX_MS = 4_102_444_800_000
|
|
23
|
+
_MAX_TASKS = 20
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def event_time(event_id: str | None) -> dt.datetime | None:
|
|
27
|
+
"""UTC time encoded in an event id, or None if the id does not have the expected shape."""
|
|
28
|
+
if event_id is None:
|
|
29
|
+
return None
|
|
30
|
+
match = _EVENT_ID.match(event_id)
|
|
31
|
+
if match is None:
|
|
32
|
+
return None
|
|
33
|
+
ms = int(match.group(1))
|
|
34
|
+
if not _MIN_MS <= ms <= _MAX_MS:
|
|
35
|
+
return None
|
|
36
|
+
return dt.datetime.fromtimestamp(ms / 1000, tz=dt.UTC)
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def iso(moment: dt.datetime | None) -> str | None:
|
|
40
|
+
if moment is None:
|
|
41
|
+
return None
|
|
42
|
+
return moment.isoformat(timespec="milliseconds").replace("+00:00", "Z")
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
class ActivityTracker:
|
|
46
|
+
"""Fed with every stream event, in order; keeps only what the summary needs."""
|
|
47
|
+
|
|
48
|
+
def __init__(self) -> None:
|
|
49
|
+
self.last_event_id: str | None = None
|
|
50
|
+
self.scanned = 0
|
|
51
|
+
self._assistant: list[str] = []
|
|
52
|
+
self._assistant_open = False
|
|
53
|
+
self._last_assistant: str | None = None
|
|
54
|
+
self._last_tool: RunEventView | None = None
|
|
55
|
+
self._tasks: dict[str, dict[str, Any]] = {}
|
|
56
|
+
|
|
57
|
+
def feed(self, event_id: str | None, kind: str, payload: object, view: RunEventView | None) -> None:
|
|
58
|
+
if event_id:
|
|
59
|
+
self.last_event_id = event_id
|
|
60
|
+
if kind in {"heartbeat", "interaction_update"}:
|
|
61
|
+
return
|
|
62
|
+
self.scanned += 1
|
|
63
|
+
if kind == "assistant":
|
|
64
|
+
text = view.text if view is not None else None
|
|
65
|
+
if not self._assistant_open:
|
|
66
|
+
self._assistant = []
|
|
67
|
+
self._assistant_open = True
|
|
68
|
+
if text:
|
|
69
|
+
self._assistant.append(text)
|
|
70
|
+
return
|
|
71
|
+
if self._assistant_open:
|
|
72
|
+
self._close_assistant()
|
|
73
|
+
if kind == "tool_call" and view is not None:
|
|
74
|
+
self._last_tool = view
|
|
75
|
+
self._track_task(event_id, tool_payload(payload))
|
|
76
|
+
|
|
77
|
+
def summary(self, *, run_terminal: bool | None, now: dt.datetime | None = None) -> ActivityView:
|
|
78
|
+
if self._assistant_open:
|
|
79
|
+
self._close_assistant()
|
|
80
|
+
now = now or dt.datetime.now(dt.UTC)
|
|
81
|
+
last_at = event_time(self.last_event_id)
|
|
82
|
+
# Counted over every task: the oldest one may be the job still running.
|
|
83
|
+
running = sum(1 for task in self._tasks.values() if task.get("state") == "running")
|
|
84
|
+
tasks = [
|
|
85
|
+
BackgroundTaskView(
|
|
86
|
+
task_id=task_id,
|
|
87
|
+
command=task.get("command"),
|
|
88
|
+
last_state=task.get("state", "unknown"),
|
|
89
|
+
runtime_ms=task.get("runtime_ms"),
|
|
90
|
+
observed_at=iso(event_time(task.get("event_id"))),
|
|
91
|
+
)
|
|
92
|
+
for task_id, task in list(self._tasks.items())[-_MAX_TASKS:]
|
|
93
|
+
]
|
|
94
|
+
tool = self._last_tool
|
|
95
|
+
return ActivityView(
|
|
96
|
+
last_event_id=self.last_event_id,
|
|
97
|
+
last_event_at=iso(last_at),
|
|
98
|
+
idle_seconds=None if last_at is None else max(0, int((now - last_at).total_seconds())),
|
|
99
|
+
last_assistant_text=self._last_assistant,
|
|
100
|
+
last_tool_call=None
|
|
101
|
+
if tool is None
|
|
102
|
+
else ToolCallSummaryView(
|
|
103
|
+
name=tool.tool_name,
|
|
104
|
+
status=tool.tool_status,
|
|
105
|
+
args=tool.tool_args,
|
|
106
|
+
result=tool.tool_result,
|
|
107
|
+
),
|
|
108
|
+
background_tasks=tasks or None,
|
|
109
|
+
background_tasks_total=len(self._tasks) or None,
|
|
110
|
+
unfinished_background_tasks=running if self._tasks else None,
|
|
111
|
+
run_terminal=run_terminal,
|
|
112
|
+
scanned_events=self.scanned,
|
|
113
|
+
)
|
|
114
|
+
|
|
115
|
+
def _close_assistant(self) -> None:
|
|
116
|
+
text = "".join(self._assistant).strip()
|
|
117
|
+
if text:
|
|
118
|
+
# The end of a long message is what tells where the agent stands.
|
|
119
|
+
self._last_assistant = text if len(text) <= ACTIVITY_TEXT_MAX_CHARS else "…" + text[-ACTIVITY_TEXT_MAX_CHARS:]
|
|
120
|
+
self._assistant_open = False
|
|
121
|
+
|
|
122
|
+
def _track_task(self, event_id: str | None, payload: dict[str, Any]) -> None:
|
|
123
|
+
"""Background launches (run_terminal_cmd) and their observed states (await)."""
|
|
124
|
+
if payload.get("status") != "completed":
|
|
125
|
+
return
|
|
126
|
+
name = payload.get("name")
|
|
127
|
+
args = payload.get("args") if isinstance(payload.get("args"), dict) else {}
|
|
128
|
+
result = payload.get("result") if isinstance(payload.get("result"), dict) else {}
|
|
129
|
+
success = result.get("success") if isinstance(result.get("success"), dict) else {}
|
|
130
|
+
if name == "run_terminal_cmd" and (args.get("isBackground") is True or result.get("isBackground") is True):
|
|
131
|
+
task_id = success.get("shellId")
|
|
132
|
+
if task_id is None:
|
|
133
|
+
return
|
|
134
|
+
command = success.get("command") or args.get("command")
|
|
135
|
+
self._tasks[str(task_id)] = {
|
|
136
|
+
"command": clip(command)[0],
|
|
137
|
+
"state": "running",
|
|
138
|
+
"runtime_ms": None,
|
|
139
|
+
"event_id": event_id,
|
|
140
|
+
}
|
|
141
|
+
return
|
|
142
|
+
if name == "await":
|
|
143
|
+
for state, key in (("running", "stillRunning"), ("complete", "complete")):
|
|
144
|
+
observed = success.get(key)
|
|
145
|
+
if isinstance(observed, dict) and observed.get("taskId") is not None:
|
|
146
|
+
task_id = str(observed["taskId"])
|
|
147
|
+
task = self._tasks.setdefault(task_id, {"command": None})
|
|
148
|
+
task.update(state=state, runtime_ms=_int(observed.get("runtimeMs")), event_id=event_id)
|
|
149
|
+
# Keep insertion order meaningful: the most recently observed task comes last.
|
|
150
|
+
self._tasks[task_id] = self._tasks.pop(task_id)
|
|
151
|
+
return
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def _int(value: object) -> int | None:
|
|
155
|
+
try:
|
|
156
|
+
return int(value) # type: ignore[arg-type]
|
|
157
|
+
except (TypeError, ValueError):
|
|
158
|
+
return None
|
|
@@ -0,0 +1,156 @@
|
|
|
1
|
+
"""List, URL and text reading of artifacts. The download does not send the Cursor key."""
|
|
2
|
+
|
|
3
|
+
import asyncio
|
|
4
|
+
import logging
|
|
5
|
+
from typing import Literal
|
|
6
|
+
|
|
7
|
+
import httpx
|
|
8
|
+
|
|
9
|
+
from cursor_cloud_mcp import budget
|
|
10
|
+
from cursor_cloud_mcp.client import CursorCloudClient, ResponseTooLarge, close_quietly, read_bounded
|
|
11
|
+
from cursor_cloud_mcp.config import ARTIFACT_MAX_BYTES, DEFAULT_DEADLINE_SECONDS
|
|
12
|
+
from cursor_cloud_mcp.errors import ErrorCode, failure
|
|
13
|
+
from cursor_cloud_mcp.models import (
|
|
14
|
+
ArtifactItemView,
|
|
15
|
+
ArtifactListView,
|
|
16
|
+
ArtifactReadView,
|
|
17
|
+
ArtifactUrlView,
|
|
18
|
+
)
|
|
19
|
+
from cursor_cloud_mcp.slicing import slice_text
|
|
20
|
+
from cursor_cloud_mcp.validation import require_artifact_path
|
|
21
|
+
|
|
22
|
+
logger = logging.getLogger(__name__)
|
|
23
|
+
_REDIRECTS = {301, 302, 303, 307, 308}
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
async def list_artifacts(client: CursorCloudClient, agent_id: str) -> ArtifactListView:
|
|
27
|
+
remote = await client.list_artifacts(agent_id)
|
|
28
|
+
return ArtifactListView(
|
|
29
|
+
items=[
|
|
30
|
+
ArtifactItemView(path=item.path, size_bytes=item.sizeBytes, updated_at=item.updatedAt)
|
|
31
|
+
for item in remote.items
|
|
32
|
+
]
|
|
33
|
+
)
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
async def artifact_url(client: CursorCloudClient, agent_id: str, path: str) -> ArtifactUrlView:
|
|
37
|
+
checked = require_artifact_path(path)
|
|
38
|
+
remote = await client.artifact_download(agent_id, checked)
|
|
39
|
+
_require_presigned(remote.url)
|
|
40
|
+
return ArtifactUrlView(path=checked, url=remote.url, expires_at=remote.expiresAt)
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
async def read_artifact(
|
|
44
|
+
client: CursorCloudClient,
|
|
45
|
+
agent_id: str,
|
|
46
|
+
path: str,
|
|
47
|
+
*,
|
|
48
|
+
offset: int,
|
|
49
|
+
limit: int,
|
|
50
|
+
url_only: bool = False,
|
|
51
|
+
) -> ArtifactReadView:
|
|
52
|
+
"""UTF-8 text of at most 5 MB. Otherwise, or on request, the presigned URL to download elsewhere."""
|
|
53
|
+
located = await artifact_url(client, agent_id, path)
|
|
54
|
+
if url_only:
|
|
55
|
+
return ArtifactReadView(path=located.path, expires_at=located.expires_at, url=located.url)
|
|
56
|
+
try:
|
|
57
|
+
raw = await fetch_presigned(
|
|
58
|
+
located.url,
|
|
59
|
+
transport=client.download_transport,
|
|
60
|
+
max_bytes=ARTIFACT_MAX_BYTES,
|
|
61
|
+
deadline=DEFAULT_DEADLINE_SECONDS,
|
|
62
|
+
)
|
|
63
|
+
except _TooLarge:
|
|
64
|
+
return _unreadable(located, "too_large")
|
|
65
|
+
try:
|
|
66
|
+
text = raw.decode("utf-8")
|
|
67
|
+
except UnicodeError:
|
|
68
|
+
return _unreadable(located, "not_utf8")
|
|
69
|
+
chunk, truncated, next_offset = slice_text(text, offset, limit)
|
|
70
|
+
return ArtifactReadView(
|
|
71
|
+
path=located.path,
|
|
72
|
+
text=chunk,
|
|
73
|
+
offset=offset,
|
|
74
|
+
limit=limit,
|
|
75
|
+
total_chars=len(text),
|
|
76
|
+
truncated=truncated,
|
|
77
|
+
next_offset=next_offset,
|
|
78
|
+
expires_at=located.expires_at,
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
class _TooLarge(Exception):
|
|
83
|
+
"""Artifact beyond the text-reading limit: the URL is returned instead."""
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def _unreadable(located: ArtifactUrlView, reason: Literal["not_utf8", "too_large"]) -> ArtifactReadView:
|
|
87
|
+
return ArtifactReadView(
|
|
88
|
+
path=located.path,
|
|
89
|
+
expires_at=located.expires_at,
|
|
90
|
+
url=located.url,
|
|
91
|
+
text_unavailable=reason,
|
|
92
|
+
)
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
async def fetch_presigned(
|
|
96
|
+
url: str,
|
|
97
|
+
*,
|
|
98
|
+
transport: httpx.AsyncBaseTransport | None,
|
|
99
|
+
max_bytes: int,
|
|
100
|
+
deadline: float,
|
|
101
|
+
) -> bytes:
|
|
102
|
+
"""Download bounded end to end: connection, reading and closing."""
|
|
103
|
+
_require_presigned(url)
|
|
104
|
+
deadline_at = budget.deadline_at(deadline)
|
|
105
|
+
remaining = deadline_at - asyncio.get_running_loop().time()
|
|
106
|
+
if remaining <= 0:
|
|
107
|
+
raise failure(ErrorCode.TIMEOUT, "Tool budget exhausted before the artifact download.")
|
|
108
|
+
try:
|
|
109
|
+
async with asyncio.timeout_at(deadline_at):
|
|
110
|
+
async with httpx.AsyncClient(
|
|
111
|
+
transport=transport,
|
|
112
|
+
follow_redirects=False,
|
|
113
|
+
timeout=httpx.Timeout(remaining),
|
|
114
|
+
) as http:
|
|
115
|
+
request = http.build_request("GET", url)
|
|
116
|
+
if "authorization" in {name.lower() for name in request.headers}:
|
|
117
|
+
raise failure(ErrorCode.VALIDATION, "The download must not carry the Cursor key.")
|
|
118
|
+
response = await http.send(request, stream=True)
|
|
119
|
+
try:
|
|
120
|
+
return await _read_download(response, max_bytes)
|
|
121
|
+
finally:
|
|
122
|
+
await close_quietly(response)
|
|
123
|
+
except (TimeoutError, httpx.TimeoutException):
|
|
124
|
+
raise failure(ErrorCode.TIMEOUT, "Timed out during the artifact download.") from None
|
|
125
|
+
except httpx.RequestError:
|
|
126
|
+
raise failure(ErrorCode.TIMEOUT, "Artifact download interrupted.") from None
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
async def _read_download(response: httpx.Response, max_bytes: int) -> bytes:
|
|
130
|
+
logger.info("artifact_download status=%s", response.status_code)
|
|
131
|
+
if response.status_code in _REDIRECTS:
|
|
132
|
+
raise failure(ErrorCode.INCOMPATIBLE_RESPONSE, "Download redirect refused.")
|
|
133
|
+
if response.status_code != 200:
|
|
134
|
+
raise failure(
|
|
135
|
+
ErrorCode.UPSTREAM,
|
|
136
|
+
f"Artifact download refused ({response.status_code}).",
|
|
137
|
+
)
|
|
138
|
+
try:
|
|
139
|
+
return await read_bounded(response, max_bytes)
|
|
140
|
+
except ResponseTooLarge:
|
|
141
|
+
raise _TooLarge from None
|
|
142
|
+
|
|
143
|
+
|
|
144
|
+
def _require_presigned(url: str) -> None:
|
|
145
|
+
try:
|
|
146
|
+
parsed = httpx.URL(url)
|
|
147
|
+
except httpx.InvalidURL:
|
|
148
|
+
raise failure(ErrorCode.VALIDATION, "Invalid artifact URL.") from None
|
|
149
|
+
host = parsed.host
|
|
150
|
+
if parsed.scheme != "https" or not host.endswith(".amazonaws.com"):
|
|
151
|
+
raise failure(
|
|
152
|
+
ErrorCode.VALIDATION,
|
|
153
|
+
"The download only accepts an HTTPS URL whose host ends with .amazonaws.com.",
|
|
154
|
+
)
|
|
155
|
+
if parsed.username or parsed.password:
|
|
156
|
+
raise failure(ErrorCode.VALIDATION, "The artifact URL must not contain credentials.")
|
|
@@ -0,0 +1,34 @@
|
|
|
1
|
+
"""Absolute budget of a tool call. Sub-operations receive the remaining time."""
|
|
2
|
+
|
|
3
|
+
import asyncio
|
|
4
|
+
from collections.abc import Iterator
|
|
5
|
+
from contextlib import contextmanager
|
|
6
|
+
from contextvars import ContextVar
|
|
7
|
+
|
|
8
|
+
_DEADLINE: ContextVar[float | None] = ContextVar("cursor_mcp_tool_deadline", default=None)
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
@contextmanager
|
|
12
|
+
def tool_budget(seconds: float) -> Iterator[None]:
|
|
13
|
+
"""Set the deadline of the current call. A nested budget can only shorten it."""
|
|
14
|
+
now = asyncio.get_running_loop().time()
|
|
15
|
+
current = _DEADLINE.get()
|
|
16
|
+
wanted = now + seconds
|
|
17
|
+
token = _DEADLINE.set(wanted if current is None else min(current, wanted))
|
|
18
|
+
try:
|
|
19
|
+
yield
|
|
20
|
+
finally:
|
|
21
|
+
_DEADLINE.reset(token)
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def deadline_at(cap_seconds: float) -> float:
|
|
25
|
+
"""Absolute deadline: the nearer of ``cap_seconds`` and the tool budget."""
|
|
26
|
+
now = asyncio.get_running_loop().time()
|
|
27
|
+
local = now + cap_seconds
|
|
28
|
+
budget = _DEADLINE.get()
|
|
29
|
+
return local if budget is None else min(local, budget)
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def remaining(cap_seconds: float) -> float:
|
|
33
|
+
"""Remaining time, capped by ``cap_seconds``. May be negative or zero."""
|
|
34
|
+
return deadline_at(cap_seconds) - asyncio.get_running_loop().time()
|
|
@@ -0,0 +1,155 @@
|
|
|
1
|
+
"""Resolution of the model and reasoning level from the Cursor catalog."""
|
|
2
|
+
|
|
3
|
+
from cursor_cloud_mcp.errors import ErrorCode, failure
|
|
4
|
+
from cursor_cloud_mcp.models import (
|
|
5
|
+
ModelParam,
|
|
6
|
+
RemoteModel,
|
|
7
|
+
RemoteModelList,
|
|
8
|
+
RemoteModelParameter,
|
|
9
|
+
)
|
|
10
|
+
|
|
11
|
+
_REASONING_IDS = ("effort", "reasoning_effort", "reasoning")
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def reasoning_parameter(model: RemoteModel) -> RemoteModelParameter | None:
|
|
15
|
+
"""First exposed reasoning parameter, in the order effort, reasoning_effort, reasoning."""
|
|
16
|
+
by_id = {parameter.id: parameter for parameter in model.parameters or []}
|
|
17
|
+
for name in _REASONING_IDS:
|
|
18
|
+
found = by_id.get(name)
|
|
19
|
+
if found is not None:
|
|
20
|
+
return found
|
|
21
|
+
return None
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def find_model(catalog: RemoteModelList, model_id: str) -> RemoteModel:
|
|
25
|
+
"""A catalog id, or an alias that designates only one model."""
|
|
26
|
+
exact = next((item for item in catalog.items if item.id == model_id), None)
|
|
27
|
+
if exact is not None:
|
|
28
|
+
return exact
|
|
29
|
+
matches = [item for item in catalog.items if model_id in (item.aliases or [])]
|
|
30
|
+
if len(matches) == 1:
|
|
31
|
+
return matches[0]
|
|
32
|
+
if matches:
|
|
33
|
+
raise failure(
|
|
34
|
+
ErrorCode.VALIDATION,
|
|
35
|
+
f"The alias {model_id} designates several models: {', '.join(item.id for item in matches)}. "
|
|
36
|
+
"Choose an id.",
|
|
37
|
+
)
|
|
38
|
+
raise failure(
|
|
39
|
+
ErrorCode.VALIDATION,
|
|
40
|
+
f"{model_id} is neither an id nor an alias in the Cursor catalog. Call cursor_list_models.",
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def default_selection(model: RemoteModel) -> dict[str, str] | None:
|
|
45
|
+
"""Values of the default variant, limited to the published parameters."""
|
|
46
|
+
published = {parameter.id for parameter in model.parameters or []}
|
|
47
|
+
for variant in model.variants or []:
|
|
48
|
+
if variant.isDefault:
|
|
49
|
+
values = {item.id: item.value for item in variant.params if item.id in published}
|
|
50
|
+
return values or None
|
|
51
|
+
return None
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def restricted(model: RemoteModel) -> bool:
|
|
55
|
+
"""True if the catalog does not offer every combination of published values."""
|
|
56
|
+
if not model.variants or not model.parameters:
|
|
57
|
+
return False
|
|
58
|
+
return len(_published_combinations(model)) < _product(model)
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def resolve_model_selection(
|
|
62
|
+
catalog: RemoteModelList,
|
|
63
|
+
*,
|
|
64
|
+
model_id: str,
|
|
65
|
+
model_params: list[ModelParam] | None,
|
|
66
|
+
reasoning_level: str | None,
|
|
67
|
+
) -> dict[str, object]:
|
|
68
|
+
"""Build ``model`` for POST /v1/agents. Refuses an id, a value or a combination outside the catalog."""
|
|
69
|
+
model = find_model(catalog, model_id)
|
|
70
|
+
ordered: list[tuple[str, str]] = []
|
|
71
|
+
seen: set[str] = set()
|
|
72
|
+
for item in model_params or []:
|
|
73
|
+
_require_param(model, item.id, item.value)
|
|
74
|
+
if item.id in seen:
|
|
75
|
+
raise failure(ErrorCode.VALIDATION, f"The parameter {item.id} is repeated.")
|
|
76
|
+
seen.add(item.id)
|
|
77
|
+
ordered.append((item.id, item.value))
|
|
78
|
+
if reasoning_level is not None:
|
|
79
|
+
parameter = reasoning_parameter(model)
|
|
80
|
+
if parameter is None:
|
|
81
|
+
raise failure(
|
|
82
|
+
ErrorCode.VALIDATION,
|
|
83
|
+
f"{model.id} does not expose a reasoning level (effort, reasoning_effort or reasoning).",
|
|
84
|
+
)
|
|
85
|
+
allowed = _values(parameter)
|
|
86
|
+
if reasoning_level not in allowed:
|
|
87
|
+
raise failure(
|
|
88
|
+
ErrorCode.VALIDATION,
|
|
89
|
+
f"{parameter.id} for {model.id} accepts: {', '.join(allowed)}. "
|
|
90
|
+
"No translation is done: xhigh and extra-high stay distinct.",
|
|
91
|
+
)
|
|
92
|
+
if parameter.id in seen:
|
|
93
|
+
current = next(value for key, value in ordered if key == parameter.id)
|
|
94
|
+
if current != reasoning_level:
|
|
95
|
+
raise failure(ErrorCode.VALIDATION, "reasoning_level contradicts model_params.")
|
|
96
|
+
else:
|
|
97
|
+
ordered.append((parameter.id, reasoning_level))
|
|
98
|
+
seen.add(parameter.id)
|
|
99
|
+
_require_combination(model, ordered)
|
|
100
|
+
body: dict[str, object] = {"id": model.id}
|
|
101
|
+
if ordered:
|
|
102
|
+
body["params"] = [{"id": key, "value": value} for key, value in ordered]
|
|
103
|
+
return body
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
def _require_combination(model: RemoteModel, chosen: list[tuple[str, str]]) -> None:
|
|
107
|
+
"""A partial selection must fit within at least one published variant."""
|
|
108
|
+
if not chosen or not model.variants:
|
|
109
|
+
return
|
|
110
|
+
wanted = set(chosen)
|
|
111
|
+
if any(wanted <= combination for combination in _published_combinations(model)):
|
|
112
|
+
return
|
|
113
|
+
text = ", ".join(f"{key}={value}" for key, value in chosen)
|
|
114
|
+
raise failure(
|
|
115
|
+
ErrorCode.VALIDATION,
|
|
116
|
+
f"Combination refused by the catalog of {model.id}: {text}. "
|
|
117
|
+
f"cursor_list_models with model_id={model.id} lists the valid variants.",
|
|
118
|
+
)
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
def _published_combinations(model: RemoteModel) -> list[set[tuple[str, str]]]:
|
|
122
|
+
published = {parameter.id for parameter in model.parameters or []}
|
|
123
|
+
found: list[set[tuple[str, str]]] = []
|
|
124
|
+
for variant in model.variants or []:
|
|
125
|
+
combination = {(item.id, item.value) for item in variant.params if item.id in published}
|
|
126
|
+
if combination not in found:
|
|
127
|
+
found.append(combination)
|
|
128
|
+
return found
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
def _product(model: RemoteModel) -> int:
|
|
132
|
+
total = 1
|
|
133
|
+
for parameter in model.parameters or []:
|
|
134
|
+
total *= len(parameter.values)
|
|
135
|
+
return total
|
|
136
|
+
|
|
137
|
+
|
|
138
|
+
def _values(parameter: RemoteModelParameter) -> list[str]:
|
|
139
|
+
return [item.value for item in parameter.values]
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
def _require_param(model: RemoteModel, param_id: str, value: str) -> None:
|
|
143
|
+
found = next((parameter for parameter in model.parameters or [] if parameter.id == param_id), None)
|
|
144
|
+
if found is None:
|
|
145
|
+
names = ", ".join(parameter.id for parameter in model.parameters or []) or "none"
|
|
146
|
+
raise failure(
|
|
147
|
+
ErrorCode.VALIDATION,
|
|
148
|
+
f"Unknown parameter {param_id} for {model.id}. Parameters: {names}.",
|
|
149
|
+
)
|
|
150
|
+
allowed = _values(found)
|
|
151
|
+
if value not in allowed:
|
|
152
|
+
raise failure(
|
|
153
|
+
ErrorCode.VALIDATION,
|
|
154
|
+
f"{param_id} for {model.id} accepts: {', '.join(allowed)}.",
|
|
155
|
+
)
|