pi-python-core 0.8.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.
- pi_python/__init__.py +160 -0
- pi_python/_version.py +1 -0
- pi_python/agent.py +396 -0
- pi_python/cancellation.py +24 -0
- pi_python/data/models.json +3315 -0
- pi_python/errors.py +49 -0
- pi_python/estimate.py +144 -0
- pi_python/events.py +138 -0
- pi_python/function_tools.py +438 -0
- pi_python/hooks.py +44 -0
- pi_python/limits.py +28 -0
- pi_python/loop.py +431 -0
- pi_python/lowlevel.py +179 -0
- pi_python/mcp.py +187 -0
- pi_python/messages.py +405 -0
- pi_python/models.py +155 -0
- pi_python/provider.py +123 -0
- pi_python/providers/__init__.py +21 -0
- pi_python/providers/anthropic.py +673 -0
- pi_python/providers/common.py +201 -0
- pi_python/providers/completions.py +1149 -0
- pi_python/providers/oauth.py +542 -0
- pi_python/providers/openai.py +681 -0
- pi_python/providers/transport.py +574 -0
- pi_python/proxy.py +304 -0
- pi_python/py.typed +0 -0
- pi_python/queues.py +76 -0
- pi_python/recovery.py +209 -0
- pi_python/run.py +419 -0
- pi_python/stream.py +251 -0
- pi_python/sync.py +78 -0
- pi_python/testing.py +25 -0
- pi_python/tools.py +546 -0
- pi_python/transcript.py +167 -0
- pi_python_core-0.8.1.dist-info/METADATA +119 -0
- pi_python_core-0.8.1.dist-info/RECORD +39 -0
- pi_python_core-0.8.1.dist-info/WHEEL +4 -0
- pi_python_core-0.8.1.dist-info/licenses/LICENSE +21 -0
- pi_python_core-0.8.1.dist-info/licenses/NOTICE +8 -0
pi_python/mcp.py
ADDED
|
@@ -0,0 +1,187 @@
|
|
|
1
|
+
"""Use the tools of an MCP server as Agent tools.
|
|
2
|
+
|
|
3
|
+
`mcp_tools(session)` wraps every tool of a connected session from the official `mcp` SDK
|
|
4
|
+
(1.x or 2.x); `connect_stdio(...)` also starts a stdio server and closes it afterwards.
|
|
5
|
+
The conversion follows Pi's MCP adapter: text and images pass through, embedded text and
|
|
6
|
+
image resources are unwrapped, other blocks become short text placeholders, and MCP's
|
|
7
|
+
`isError` marks a failed call. Install with ``pip install 'pi-python-core[mcp]'``.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
import asyncio
|
|
13
|
+
import hashlib
|
|
14
|
+
import inspect
|
|
15
|
+
import json
|
|
16
|
+
import re
|
|
17
|
+
import warnings
|
|
18
|
+
from collections.abc import AsyncIterator, Iterable, Mapping, Sequence
|
|
19
|
+
from contextlib import asynccontextmanager
|
|
20
|
+
from typing import Any
|
|
21
|
+
|
|
22
|
+
from .errors import ConfigurationError
|
|
23
|
+
from .messages import ImageContent, TextContent
|
|
24
|
+
from .tools import Tool, ToolContext, ToolResult
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def _get(value: Any, snake: str, camel: str | None = None) -> Any:
|
|
28
|
+
"""Read a field in either SDK's naming (2.x snake_case, 1.x camelCase) or from a dict."""
|
|
29
|
+
for name in (snake, camel) if camel else (snake,):
|
|
30
|
+
if isinstance(value, Mapping):
|
|
31
|
+
if name in value:
|
|
32
|
+
return value[name]
|
|
33
|
+
elif hasattr(value, name):
|
|
34
|
+
return getattr(value, name)
|
|
35
|
+
return None
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def _block(block: Any) -> TextContent | ImageContent:
|
|
39
|
+
kind = _get(block, "type")
|
|
40
|
+
if kind == "text":
|
|
41
|
+
return TextContent(_get(block, "text") or "")
|
|
42
|
+
if kind == "image":
|
|
43
|
+
return ImageContent(_get(block, "data"), _get(block, "mime_type", "mimeType"))
|
|
44
|
+
if kind == "audio":
|
|
45
|
+
return TextContent(f"[audio {_get(block, 'mime_type', 'mimeType')} omitted]")
|
|
46
|
+
if kind == "resource_link":
|
|
47
|
+
return TextContent(f"{_get(block, 'name')}: {_get(block, 'uri')}")
|
|
48
|
+
if kind == "resource":
|
|
49
|
+
resource = _get(block, "resource")
|
|
50
|
+
mime = _get(resource, "mime_type", "mimeType")
|
|
51
|
+
if _get(resource, "text") is not None:
|
|
52
|
+
return TextContent(_get(resource, "text"))
|
|
53
|
+
if mime and str(mime).startswith("image/"):
|
|
54
|
+
return ImageContent(_get(resource, "blob"), mime)
|
|
55
|
+
return TextContent(
|
|
56
|
+
f"[binary resource {_get(resource, 'uri')} ({mime or 'unknown type'}) omitted]"
|
|
57
|
+
)
|
|
58
|
+
return TextContent(f"[unsupported MCP content {kind}]")
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def _tool_name(prefix: str | None, name: str, taken: set[str]) -> str:
|
|
62
|
+
# Model APIs accept letters, digits, "_" and "-", up to 64 characters.
|
|
63
|
+
full = f"{prefix}_{name}" if prefix else name
|
|
64
|
+
result = re.sub(r"[^A-Za-z0-9_-]", "_", full)[:64]
|
|
65
|
+
if result in taken: # "get.file" and "get_file", or two long names with one prefix
|
|
66
|
+
result = f"{result[:55]}_{hashlib.sha1(full.encode()).hexdigest()[:8]}"
|
|
67
|
+
warnings.warn(f"MCP tool {name!r} renamed to {result!r} to keep names unique", stacklevel=3)
|
|
68
|
+
taken.add(result)
|
|
69
|
+
return result
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def _wrap(session: Any, spec: Any, prefix: str | None, taken: set[str]) -> Tool:
|
|
73
|
+
name = _get(spec, "name")
|
|
74
|
+
schema = dict(_get(spec, "input_schema", "inputSchema") or {})
|
|
75
|
+
# Model APIs require an object schema, and some reject one without `properties`.
|
|
76
|
+
schema.update(type="object", properties=schema.get("properties") or {})
|
|
77
|
+
|
|
78
|
+
async def execute(args: dict[str, Any], context: ToolContext) -> ToolResult:
|
|
79
|
+
async def progress(
|
|
80
|
+
done: float, total: float | None = None, message: str | None = None
|
|
81
|
+
) -> None:
|
|
82
|
+
await context.emit_update({"progress": done, "total": total, "message": message})
|
|
83
|
+
|
|
84
|
+
result = await session.call_tool(name, args, progress_callback=progress)
|
|
85
|
+
blocks = _get(result, "content")
|
|
86
|
+
if blocks is None:
|
|
87
|
+
# MCP 2.x can ask the client for input mid-call; this adapter cannot answer.
|
|
88
|
+
return ToolResult.text(
|
|
89
|
+
f"MCP tool {name} asked for input this client cannot provide", is_error=True
|
|
90
|
+
)
|
|
91
|
+
content = [_block(b) for b in blocks]
|
|
92
|
+
structured = _get(result, "structured_content", "structuredContent")
|
|
93
|
+
is_error = _get(result, "is_error", "isError") is True
|
|
94
|
+
if not content and structured is not None:
|
|
95
|
+
content = [TextContent(json.dumps(structured, indent=2, ensure_ascii=False))]
|
|
96
|
+
if is_error and not content:
|
|
97
|
+
content = [TextContent(f"MCP tool {name} failed")]
|
|
98
|
+
return ToolResult(content, structured_content=structured, is_error=is_error)
|
|
99
|
+
|
|
100
|
+
description = _get(spec, "description") or _get(spec, "title") or f"MCP tool {name}"
|
|
101
|
+
tool = Tool(_tool_name(prefix, name, set(taken)), description, schema, execute)
|
|
102
|
+
taken.add(tool.name)
|
|
103
|
+
return tool
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
async def mcp_tools(
|
|
107
|
+
session: Any, *, prefix: str | None = None, names: Iterable[str] | None = None
|
|
108
|
+
) -> list[Tool]:
|
|
109
|
+
"""Wrap the tools of an initialized MCP client session.
|
|
110
|
+
|
|
111
|
+
`prefix` namespaces the tool names (``prefix_tool``); `names` keeps only those MCP
|
|
112
|
+
tools. A tool whose input schema cannot be used is skipped with a warning, so one
|
|
113
|
+
unusual tool does not make the whole server unusable.
|
|
114
|
+
"""
|
|
115
|
+
wanted = set(names) if names is not None else None
|
|
116
|
+
specs: list[Any] = []
|
|
117
|
+
cursor = None
|
|
118
|
+
while True:
|
|
119
|
+
page = await _list_page(session, cursor)
|
|
120
|
+
specs.extend(_get(page, "tools") or [])
|
|
121
|
+
cursor = _get(page, "next_cursor", "nextCursor")
|
|
122
|
+
if not cursor:
|
|
123
|
+
break
|
|
124
|
+
tools: list[Tool] = []
|
|
125
|
+
taken: set[str] = set()
|
|
126
|
+
for spec in specs:
|
|
127
|
+
if wanted is not None and _get(spec, "name") not in wanted:
|
|
128
|
+
continue
|
|
129
|
+
try:
|
|
130
|
+
tools.append(_wrap(session, spec, prefix, taken))
|
|
131
|
+
except (ConfigurationError, RecursionError) as exc:
|
|
132
|
+
warnings.warn(f"Skipping MCP tool {_get(spec, 'name')}: {exc}", stacklevel=2)
|
|
133
|
+
return tools
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
async def _list_page(session: Any, cursor: Any) -> Any:
|
|
137
|
+
if cursor is None:
|
|
138
|
+
return await session.list_tools()
|
|
139
|
+
if "params" in inspect.signature(session.list_tools).parameters: # MCP SDK 2.x
|
|
140
|
+
from mcp.types import PaginatedRequestParams # type: ignore[import-not-found,unused-ignore]
|
|
141
|
+
|
|
142
|
+
return await session.list_tools(params=PaginatedRequestParams(cursor=cursor))
|
|
143
|
+
return await session.list_tools(cursor=cursor)
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
@asynccontextmanager
|
|
147
|
+
async def connect_stdio(
|
|
148
|
+
command: str,
|
|
149
|
+
args: Sequence[str] = (),
|
|
150
|
+
*,
|
|
151
|
+
env: dict[str, str] | None = None,
|
|
152
|
+
cwd: str | None = None,
|
|
153
|
+
prefix: str | None = None,
|
|
154
|
+
names: Iterable[str] | None = None,
|
|
155
|
+
init_timeout: float = 30.0,
|
|
156
|
+
) -> AsyncIterator[list[Tool]]:
|
|
157
|
+
"""Start a stdio MCP server, yield its tools, and stop it when the block ends.
|
|
158
|
+
|
|
159
|
+
``async with connect_stdio("uvx", ["mcp-server-fetch"]) as tools: ...``
|
|
160
|
+
|
|
161
|
+
A server that does not complete the MCP handshake within `init_timeout` seconds
|
|
162
|
+
raises TimeoutError instead of hanging.
|
|
163
|
+
"""
|
|
164
|
+
try:
|
|
165
|
+
from mcp import ClientSession, StdioServerParameters # type: ignore[import-not-found,unused-ignore]
|
|
166
|
+
from mcp.client.stdio import stdio_client # type: ignore[import-not-found,unused-ignore]
|
|
167
|
+
except ImportError as exc:
|
|
168
|
+
raise ImportError(
|
|
169
|
+
"connect_stdio needs the MCP SDK: pip install 'pi-python-core[mcp]'"
|
|
170
|
+
) from exc
|
|
171
|
+
params = StdioServerParameters(command=command, args=list(args), env=env, cwd=cwd)
|
|
172
|
+
try:
|
|
173
|
+
async with stdio_client(params) as (read, write):
|
|
174
|
+
async with ClientSession(read, write) as session:
|
|
175
|
+
async with asyncio.timeout(init_timeout):
|
|
176
|
+
await session.initialize()
|
|
177
|
+
tools = await mcp_tools(session, prefix=prefix, names=names)
|
|
178
|
+
yield tools
|
|
179
|
+
except BaseExceptionGroup as group:
|
|
180
|
+
# The SDK's task groups wrap a single failure ("Connection closed", a timeout,
|
|
181
|
+
# or an error from the caller's block), sometimes twice; raise that failure itself.
|
|
182
|
+
error: BaseException = group
|
|
183
|
+
while isinstance(error, BaseExceptionGroup) and len(error.exceptions) == 1:
|
|
184
|
+
error = error.exceptions[0]
|
|
185
|
+
if error is group:
|
|
186
|
+
raise
|
|
187
|
+
raise error from None
|
pi_python/messages.py
ADDED
|
@@ -0,0 +1,405 @@
|
|
|
1
|
+
"""Versioned, JSON-only messages. All boundaries copy caller-owned values."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import json
|
|
6
|
+
import math
|
|
7
|
+
import time
|
|
8
|
+
from dataclasses import asdict, dataclass, field
|
|
9
|
+
from typing import Any, Literal, TypeAlias
|
|
10
|
+
|
|
11
|
+
from .errors import MessageValidationError, UnsupportedCapabilityError
|
|
12
|
+
import base64
|
|
13
|
+
|
|
14
|
+
JsonValue: TypeAlias = "None | bool | int | float | str | list[JsonValue] | dict[str, JsonValue]"
|
|
15
|
+
SCHEMA_VERSION = 3
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def now() -> float:
|
|
19
|
+
return time.time()
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def validate_json(value: Any) -> None:
|
|
23
|
+
if value is None or type(value) in (str, bool, int):
|
|
24
|
+
return
|
|
25
|
+
if type(value) is float and math.isfinite(value):
|
|
26
|
+
return
|
|
27
|
+
if type(value) is list:
|
|
28
|
+
for item in value:
|
|
29
|
+
validate_json(item)
|
|
30
|
+
return
|
|
31
|
+
if type(value) is dict and all(type(k) is str for k in value):
|
|
32
|
+
for item in value.values():
|
|
33
|
+
validate_json(item)
|
|
34
|
+
return
|
|
35
|
+
raise MessageValidationError(f"Not a finite JSON value: {type(value).__name__}")
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
@dataclass
|
|
39
|
+
class TextContent:
|
|
40
|
+
text: str
|
|
41
|
+
text_signature: str | None = None
|
|
42
|
+
type: Literal["text"] = field(default="text", init=False)
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
@dataclass
|
|
46
|
+
class ImageContent:
|
|
47
|
+
data: str
|
|
48
|
+
mime_type: str
|
|
49
|
+
type: Literal["image"] = field(default="image", init=False)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
@dataclass
|
|
53
|
+
class ThinkingContent:
|
|
54
|
+
thinking: str
|
|
55
|
+
thinking_signature: str | None = None
|
|
56
|
+
redacted: bool = False
|
|
57
|
+
type: Literal["thinking"] = field(default="thinking", init=False)
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
@dataclass
|
|
61
|
+
class ToolCall:
|
|
62
|
+
id: str
|
|
63
|
+
name: str
|
|
64
|
+
arguments: dict[str, Any]
|
|
65
|
+
thought_signature: str | None = None
|
|
66
|
+
namespace: str | None = None
|
|
67
|
+
type: Literal["tool_call"] = field(default="tool_call", init=False)
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
@dataclass
|
|
71
|
+
class ToolDeclaration:
|
|
72
|
+
name: str
|
|
73
|
+
description: str
|
|
74
|
+
input_schema: dict[str, Any]
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
@dataclass
|
|
78
|
+
class SystemMessage:
|
|
79
|
+
content: str | list[TextContent] = ""
|
|
80
|
+
sections: dict[str, str | None] = field(default_factory=dict)
|
|
81
|
+
tools_added: list[ToolDeclaration] = field(default_factory=list)
|
|
82
|
+
tools_removed: list[str] = field(default_factory=list)
|
|
83
|
+
timestamp: float = field(default_factory=now)
|
|
84
|
+
role: Literal["system"] = field(default="system", init=False)
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
@dataclass
|
|
88
|
+
class UserMessage:
|
|
89
|
+
content: str | list[TextContent | ImageContent]
|
|
90
|
+
timestamp: float = field(default_factory=now)
|
|
91
|
+
role: Literal["user"] = field(default="user", init=False)
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
@dataclass
|
|
95
|
+
class AssistantMessage:
|
|
96
|
+
content: list[TextContent | ThinkingContent | ToolCall]
|
|
97
|
+
stop_reason: str = "stop"
|
|
98
|
+
provider: str = "mock"
|
|
99
|
+
model: str = "mock"
|
|
100
|
+
usage: dict[str, Any] = field(default_factory=dict)
|
|
101
|
+
timestamp: float = field(default_factory=now)
|
|
102
|
+
error: str | None = None
|
|
103
|
+
api: str | None = None
|
|
104
|
+
provider_thinking_level: str | None = None
|
|
105
|
+
response_model: str | None = None
|
|
106
|
+
response_id: str | None = None
|
|
107
|
+
thinking_level: str | None = None
|
|
108
|
+
diagnostics: list[dict[str, Any]] | None = None
|
|
109
|
+
raw_stop_reason: str | None = None
|
|
110
|
+
end_turn: bool | None = None
|
|
111
|
+
deferred: dict[str, Any] | None = None
|
|
112
|
+
role: Literal["assistant"] = field(default="assistant", init=False)
|
|
113
|
+
|
|
114
|
+
@classmethod
|
|
115
|
+
def text(cls, text: str, **kwargs: Any) -> AssistantMessage:
|
|
116
|
+
return cls([TextContent(text)], **kwargs)
|
|
117
|
+
|
|
118
|
+
@property
|
|
119
|
+
def tool_calls(self) -> list[ToolCall]:
|
|
120
|
+
return [b for b in self.content if isinstance(b, ToolCall)]
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
@dataclass
|
|
124
|
+
class ToolResultMessage:
|
|
125
|
+
call_id: str
|
|
126
|
+
name: str
|
|
127
|
+
content: list[TextContent | ImageContent]
|
|
128
|
+
is_error: bool = False
|
|
129
|
+
timestamp: float = field(default_factory=now)
|
|
130
|
+
details: Any = None
|
|
131
|
+
usage: dict[str, Any] | None = None
|
|
132
|
+
nested_calls: dict[str, Any] | None = None
|
|
133
|
+
role: Literal["tool_result"] = field(default="tool_result", init=False)
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
@dataclass
|
|
137
|
+
class CustomMessage:
|
|
138
|
+
custom_type: str
|
|
139
|
+
data: Any
|
|
140
|
+
timestamp: float = field(default_factory=now)
|
|
141
|
+
role: Literal["custom"] = field(default="custom", init=False)
|
|
142
|
+
|
|
143
|
+
|
|
144
|
+
Message: TypeAlias = (
|
|
145
|
+
SystemMessage | UserMessage | AssistantMessage | ToolResultMessage | CustomMessage
|
|
146
|
+
)
|
|
147
|
+
|
|
148
|
+
|
|
149
|
+
def message_to_dict(message: Message) -> dict[str, Any]:
|
|
150
|
+
if not isinstance(
|
|
151
|
+
message, (SystemMessage, UserMessage, AssistantMessage, ToolResultMessage, CustomMessage)
|
|
152
|
+
):
|
|
153
|
+
raise MessageValidationError("Expected a supported Message")
|
|
154
|
+
optional_fields = {
|
|
155
|
+
"text_signature",
|
|
156
|
+
"thought_signature",
|
|
157
|
+
"namespace",
|
|
158
|
+
"api",
|
|
159
|
+
"provider_thinking_level",
|
|
160
|
+
"response_model",
|
|
161
|
+
"response_id",
|
|
162
|
+
"thinking_level",
|
|
163
|
+
"diagnostics",
|
|
164
|
+
"raw_stop_reason",
|
|
165
|
+
"end_turn",
|
|
166
|
+
"deferred",
|
|
167
|
+
"details",
|
|
168
|
+
"usage",
|
|
169
|
+
"nested_calls",
|
|
170
|
+
}
|
|
171
|
+
# dict_factory visits dataclass fields, not arbitrary user dictionaries.
|
|
172
|
+
data = asdict(
|
|
173
|
+
message,
|
|
174
|
+
dict_factory=lambda pairs: {
|
|
175
|
+
k: v for k, v in pairs if not (k in optional_fields and v is None)
|
|
176
|
+
},
|
|
177
|
+
)
|
|
178
|
+
validate_json(data)
|
|
179
|
+
# Parsing validates semantic fields as well as JSON representability.
|
|
180
|
+
message_from_dict(data)
|
|
181
|
+
return data
|
|
182
|
+
|
|
183
|
+
|
|
184
|
+
def _text(value: Any, label: str) -> str:
|
|
185
|
+
if not isinstance(value, str):
|
|
186
|
+
raise MessageValidationError(f"{label} must be a string")
|
|
187
|
+
return value
|
|
188
|
+
|
|
189
|
+
|
|
190
|
+
def _blocks(values: Any, tools: bool = False, images: bool = False) -> list:
|
|
191
|
+
if not isinstance(values, list):
|
|
192
|
+
raise MessageValidationError("content must be a list")
|
|
193
|
+
result: list = []
|
|
194
|
+
for b in values:
|
|
195
|
+
if not isinstance(b, dict):
|
|
196
|
+
raise MessageValidationError("content block must be an object")
|
|
197
|
+
if b.get("type") == "text":
|
|
198
|
+
if set(b) - {"type", "text", "text_signature"} or "text" not in b:
|
|
199
|
+
raise MessageValidationError("Invalid text block fields")
|
|
200
|
+
signature = b.get("text_signature")
|
|
201
|
+
if signature is not None:
|
|
202
|
+
_text(signature, "text_signature")
|
|
203
|
+
result.append(TextContent(_text(b["text"], "text"), signature))
|
|
204
|
+
elif b.get("type") == "image" and images:
|
|
205
|
+
if set(b) != {"type", "data", "mime_type"}:
|
|
206
|
+
raise MessageValidationError("Invalid image block")
|
|
207
|
+
data, mime = _text(b["data"], "image data"), _text(b["mime_type"], "mime_type")
|
|
208
|
+
if mime not in {"image/png", "image/jpeg", "image/gif", "image/webp"}:
|
|
209
|
+
raise UnsupportedCapabilityError("Unsupported image MIME type")
|
|
210
|
+
try:
|
|
211
|
+
base64.b64decode(data, validate=True)
|
|
212
|
+
except ValueError as exc:
|
|
213
|
+
raise MessageValidationError("Invalid image base64") from exc
|
|
214
|
+
result.append(ImageContent(data, mime))
|
|
215
|
+
elif b.get("type") == "thinking" and tools:
|
|
216
|
+
if set(b) - {"type", "thinking", "thinking_signature", "redacted"}:
|
|
217
|
+
raise MessageValidationError("Invalid thinking block")
|
|
218
|
+
signature = b.get("thinking_signature")
|
|
219
|
+
if signature is not None:
|
|
220
|
+
_text(signature, "thinking_signature")
|
|
221
|
+
if type(b.get("redacted", False)) is not bool:
|
|
222
|
+
raise MessageValidationError("redacted must be boolean")
|
|
223
|
+
result.append(
|
|
224
|
+
ThinkingContent(
|
|
225
|
+
_text(b["thinking"], "thinking"), signature, b.get("redacted", False)
|
|
226
|
+
)
|
|
227
|
+
)
|
|
228
|
+
elif b.get("type") == "tool_call" and tools:
|
|
229
|
+
if (
|
|
230
|
+
set(b) - {"type", "id", "name", "arguments", "thought_signature", "namespace"}
|
|
231
|
+
or not {"id", "name", "arguments"} <= b.keys()
|
|
232
|
+
):
|
|
233
|
+
raise MessageValidationError("Invalid tool call fields")
|
|
234
|
+
if not b["id"] or not b["name"] or not isinstance(b["arguments"], dict):
|
|
235
|
+
raise MessageValidationError("Invalid tool call ID, name or arguments")
|
|
236
|
+
for key in ("thought_signature", "namespace"):
|
|
237
|
+
if b.get(key) is not None:
|
|
238
|
+
_text(b[key], key)
|
|
239
|
+
result.append(
|
|
240
|
+
ToolCall(
|
|
241
|
+
_text(b["id"], "id"),
|
|
242
|
+
_text(b["name"], "name"),
|
|
243
|
+
b["arguments"],
|
|
244
|
+
b.get("thought_signature"),
|
|
245
|
+
b.get("namespace"),
|
|
246
|
+
)
|
|
247
|
+
)
|
|
248
|
+
else:
|
|
249
|
+
raise UnsupportedCapabilityError(f"Unsupported content: {b.get('type')}")
|
|
250
|
+
return result
|
|
251
|
+
|
|
252
|
+
|
|
253
|
+
def message_from_dict(data: dict[str, Any]) -> Message:
|
|
254
|
+
validate_json(data)
|
|
255
|
+
if not isinstance(data, dict):
|
|
256
|
+
raise MessageValidationError("Message must be an object")
|
|
257
|
+
d = json.loads(json.dumps(data, allow_nan=False))
|
|
258
|
+
role = d.pop("role", None)
|
|
259
|
+
try:
|
|
260
|
+
if "timestamp" in d and (type(d["timestamp"]) not in (float, int)):
|
|
261
|
+
raise MessageValidationError("timestamp must be a number")
|
|
262
|
+
if role == "user":
|
|
263
|
+
if not isinstance(d["content"], str):
|
|
264
|
+
d["content"] = _blocks(d["content"], images=True)
|
|
265
|
+
return UserMessage(**d)
|
|
266
|
+
if role == "system":
|
|
267
|
+
if isinstance(d.get("content", ""), str):
|
|
268
|
+
_text(d.get("content", ""), "content")
|
|
269
|
+
else:
|
|
270
|
+
d["content"] = _blocks(d["content"])
|
|
271
|
+
sections = d.get("sections", {})
|
|
272
|
+
if not isinstance(sections, dict) or any(
|
|
273
|
+
v is not None and not isinstance(v, str) for v in sections.values()
|
|
274
|
+
):
|
|
275
|
+
raise MessageValidationError("Invalid sections")
|
|
276
|
+
d["tools_added"] = [ToolDeclaration(**t) for t in d.get("tools_added", [])]
|
|
277
|
+
for t in d["tools_added"]:
|
|
278
|
+
_text(t.name, "tool name")
|
|
279
|
+
_text(t.description, "description")
|
|
280
|
+
if not isinstance(t.input_schema, dict):
|
|
281
|
+
raise MessageValidationError("Invalid tool schema")
|
|
282
|
+
if not isinstance(d.get("tools_removed", []), list):
|
|
283
|
+
raise MessageValidationError("Invalid tool removals")
|
|
284
|
+
for name in d.get("tools_removed", []):
|
|
285
|
+
_text(name, "tool removal")
|
|
286
|
+
return SystemMessage(**d)
|
|
287
|
+
if role == "assistant":
|
|
288
|
+
d["content"] = _blocks(d["content"], tools=True)
|
|
289
|
+
if d.get("stop_reason", "stop") not in {
|
|
290
|
+
"stop",
|
|
291
|
+
"tool_use",
|
|
292
|
+
"length",
|
|
293
|
+
"error",
|
|
294
|
+
"aborted",
|
|
295
|
+
"pending",
|
|
296
|
+
"deferred",
|
|
297
|
+
}:
|
|
298
|
+
raise UnsupportedCapabilityError("Unsupported stop reason")
|
|
299
|
+
calls = [b.id for b in d["content"] if isinstance(b, ToolCall)]
|
|
300
|
+
if len(set(calls)) != len(calls):
|
|
301
|
+
raise MessageValidationError("Duplicate tool call ID")
|
|
302
|
+
for key in ("provider", "model"):
|
|
303
|
+
_text(d.get(key, "mock"), key)
|
|
304
|
+
for key in (
|
|
305
|
+
"api",
|
|
306
|
+
"provider_thinking_level",
|
|
307
|
+
"response_model",
|
|
308
|
+
"response_id",
|
|
309
|
+
"thinking_level",
|
|
310
|
+
"raw_stop_reason",
|
|
311
|
+
):
|
|
312
|
+
if d.get(key) is not None:
|
|
313
|
+
_text(d[key], key)
|
|
314
|
+
if not isinstance(d.get("usage", {}), dict):
|
|
315
|
+
raise MessageValidationError("Invalid usage")
|
|
316
|
+
if d.get("error") is not None:
|
|
317
|
+
_text(d["error"], "error")
|
|
318
|
+
if d.get("end_turn") is not None and type(d["end_turn"]) is not bool:
|
|
319
|
+
raise MessageValidationError("end_turn must be boolean")
|
|
320
|
+
if d.get("diagnostics") is not None and (
|
|
321
|
+
not isinstance(d["diagnostics"], list)
|
|
322
|
+
or any(not isinstance(v, dict) for v in d["diagnostics"])
|
|
323
|
+
):
|
|
324
|
+
raise MessageValidationError("diagnostics must be a list of objects")
|
|
325
|
+
if d.get("deferred") is not None:
|
|
326
|
+
handle = d["deferred"]
|
|
327
|
+
if not isinstance(handle, dict) or not all(
|
|
328
|
+
isinstance(handle.get(k), str) for k in ("provider", "model_id", "api", "id")
|
|
329
|
+
):
|
|
330
|
+
raise MessageValidationError("Invalid deferred handle")
|
|
331
|
+
return AssistantMessage(**d)
|
|
332
|
+
if role == "tool_result":
|
|
333
|
+
d["content"] = _blocks(d["content"], images=True)
|
|
334
|
+
_text(d["call_id"], "call_id")
|
|
335
|
+
_text(d["name"], "name")
|
|
336
|
+
if type(d.get("is_error", False)) is not bool:
|
|
337
|
+
raise MessageValidationError("is_error must be boolean")
|
|
338
|
+
if d.get("usage") is not None and not isinstance(d["usage"], dict):
|
|
339
|
+
raise MessageValidationError("Invalid tool usage")
|
|
340
|
+
if d.get("nested_calls") is not None:
|
|
341
|
+
nested = d["nested_calls"]
|
|
342
|
+
if (
|
|
343
|
+
not isinstance(nested, dict)
|
|
344
|
+
or type(nested.get("complete")) is not bool
|
|
345
|
+
or not isinstance(nested.get("calls"), list)
|
|
346
|
+
):
|
|
347
|
+
raise MessageValidationError("Invalid nested calls")
|
|
348
|
+
for call in nested["calls"]:
|
|
349
|
+
if (
|
|
350
|
+
not isinstance(call, dict)
|
|
351
|
+
or not all(isinstance(call.get(k), str) for k in ("id", "name"))
|
|
352
|
+
or call.get("status") not in {"ok", "error", "unfinished"}
|
|
353
|
+
):
|
|
354
|
+
raise MessageValidationError("Invalid nested call")
|
|
355
|
+
return ToolResultMessage(**d)
|
|
356
|
+
if role == "custom":
|
|
357
|
+
_text(d["custom_type"], "custom_type")
|
|
358
|
+
return CustomMessage(**d)
|
|
359
|
+
raise UnsupportedCapabilityError(f"Unsupported role: {role}")
|
|
360
|
+
except (TypeError, KeyError, ValueError) as exc:
|
|
361
|
+
raise MessageValidationError(str(exc)) from exc
|
|
362
|
+
|
|
363
|
+
|
|
364
|
+
def validate_history(messages: list[Message]) -> None:
|
|
365
|
+
pending: dict[str, str] = {}
|
|
366
|
+
for message in messages:
|
|
367
|
+
message_to_dict(message)
|
|
368
|
+
if isinstance(message, ToolResultMessage):
|
|
369
|
+
if pending.pop(message.call_id, None) != message.name:
|
|
370
|
+
raise MessageValidationError("Unmatched or duplicate tool result")
|
|
371
|
+
else:
|
|
372
|
+
if pending:
|
|
373
|
+
raise MessageValidationError("Dangling tool calls")
|
|
374
|
+
if isinstance(message, AssistantMessage):
|
|
375
|
+
if message.stop_reason == "pending":
|
|
376
|
+
raise MessageValidationError("Pending response cannot be committed to history")
|
|
377
|
+
if message.tool_calls and message.stop_reason in {"error", "aborted"}:
|
|
378
|
+
raise MessageValidationError("Failed assistant cannot declare tool calls")
|
|
379
|
+
pending = {c.id: c.name for c in message.tool_calls}
|
|
380
|
+
if pending:
|
|
381
|
+
raise MessageValidationError("Dangling tool calls")
|
|
382
|
+
|
|
383
|
+
|
|
384
|
+
def encode_messages(messages: list[Message]) -> str:
|
|
385
|
+
validate_history(messages)
|
|
386
|
+
return json.dumps(
|
|
387
|
+
{"schema_version": SCHEMA_VERSION, "messages": [message_to_dict(m) for m in messages]},
|
|
388
|
+
ensure_ascii=False,
|
|
389
|
+
allow_nan=False,
|
|
390
|
+
)
|
|
391
|
+
|
|
392
|
+
|
|
393
|
+
def decode_messages(value: str) -> list[Message]:
|
|
394
|
+
try:
|
|
395
|
+
data = json.loads(value)
|
|
396
|
+
validate_json(data)
|
|
397
|
+
if type(data.get("schema_version")) is not int or data["schema_version"] != SCHEMA_VERSION:
|
|
398
|
+
raise MessageValidationError("Unsupported message schema version")
|
|
399
|
+
if set(data) != {"schema_version", "messages"} or not isinstance(data["messages"], list):
|
|
400
|
+
raise MessageValidationError("Invalid message envelope")
|
|
401
|
+
messages = [message_from_dict(m) for m in data["messages"]]
|
|
402
|
+
validate_history(messages)
|
|
403
|
+
return messages
|
|
404
|
+
except (ValueError, TypeError, AttributeError, KeyError) as exc:
|
|
405
|
+
raise MessageValidationError(str(exc)) from exc
|