mcptoolforge 0.1.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.
- mcptoolforge/__init__.py +64 -0
- mcptoolforge/cli/__init__.py +1 -0
- mcptoolforge/cli/commands.py +278 -0
- mcptoolforge/cli/main.py +101 -0
- mcptoolforge/config.py +45 -0
- mcptoolforge/decorators.py +49 -0
- mcptoolforge/errors.py +94 -0
- mcptoolforge/execution.py +25 -0
- mcptoolforge/mcp/__init__.py +4 -0
- mcptoolforge/mcp/adapter.py +405 -0
- mcptoolforge/mcp/server.py +111 -0
- mcptoolforge/middleware/__init__.py +11 -0
- mcptoolforge/middleware/context.py +15 -0
- mcptoolforge/middleware/logging.py +30 -0
- mcptoolforge/middleware/manager.py +68 -0
- mcptoolforge/middleware/timing.py +37 -0
- mcptoolforge/project.py +134 -0
- mcptoolforge/prompts/__init__.py +5 -0
- mcptoolforge/prompts/decorator.py +32 -0
- mcptoolforge/prompts/prompt.py +129 -0
- mcptoolforge/prompts/registry.py +37 -0
- mcptoolforge/registry.py +241 -0
- mcptoolforge/resources/__init__.py +5 -0
- mcptoolforge/resources/decorator.py +33 -0
- mcptoolforge/resources/registry.py +39 -0
- mcptoolforge/resources/resource.py +97 -0
- mcptoolforge/schema.py +146 -0
- mcptoolforge/server.py +91 -0
- mcptoolforge/testing/__init__.py +15 -0
- mcptoolforge/testing/client.py +364 -0
- mcptoolforge/transports/__init__.py +1 -0
- mcptoolforge/validation.py +229 -0
- mcptoolforge-0.1.0.dist-info/METADATA +655 -0
- mcptoolforge-0.1.0.dist-info/RECORD +37 -0
- mcptoolforge-0.1.0.dist-info/WHEEL +4 -0
- mcptoolforge-0.1.0.dist-info/entry_points.txt +2 -0
- mcptoolforge-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,364 @@
|
|
|
1
|
+
import asyncio
|
|
2
|
+
import base64
|
|
3
|
+
import inspect
|
|
4
|
+
import json
|
|
5
|
+
import time
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
from mcptoolforge.errors import (
|
|
9
|
+
MCPToolForgeError,
|
|
10
|
+
ResourceExecutionError,
|
|
11
|
+
ResourceNotFoundError,
|
|
12
|
+
ToolExecutionError,
|
|
13
|
+
ToolNotFoundError,
|
|
14
|
+
ToolValidationError,
|
|
15
|
+
)
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class MCPToolForgeTestingError(MCPToolForgeError):
|
|
19
|
+
"""Exception raised for general MCPToolForge test client errors."""
|
|
20
|
+
|
|
21
|
+
pass
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class ResourceReadResult:
|
|
25
|
+
"""Convenient result container for reading resources in tests."""
|
|
26
|
+
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
uri: str,
|
|
30
|
+
text: str | None = None,
|
|
31
|
+
blob: str | None = None,
|
|
32
|
+
mime_type: str | None = None,
|
|
33
|
+
):
|
|
34
|
+
self.uri = uri
|
|
35
|
+
self.text = text
|
|
36
|
+
self.blob = blob
|
|
37
|
+
self.mime_type = mime_type
|
|
38
|
+
|
|
39
|
+
@property
|
|
40
|
+
def is_blob(self) -> bool:
|
|
41
|
+
return self.blob is not None
|
|
42
|
+
|
|
43
|
+
def __repr__(self) -> str:
|
|
44
|
+
return (
|
|
45
|
+
f"ResourceReadResult(uri={self.uri!r}, text={self.text!r}, "
|
|
46
|
+
f"blob={self.blob!r}, mime_type={self.mime_type!r})"
|
|
47
|
+
)
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
class PromptMessageResult:
|
|
51
|
+
"""Represents a message result of a prompt retrieval in tests."""
|
|
52
|
+
|
|
53
|
+
def __init__(self, role: str, content: str):
|
|
54
|
+
self.role = role
|
|
55
|
+
self.content = content
|
|
56
|
+
|
|
57
|
+
def __repr__(self) -> str:
|
|
58
|
+
return f"PromptMessageResult(role={self.role!r}, content={self.content!r})"
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
class PromptGetResult:
|
|
62
|
+
"""Convenient result container for retrieved prompts in tests."""
|
|
63
|
+
|
|
64
|
+
def __init__(self, messages: list[PromptMessageResult], description: str | None = None):
|
|
65
|
+
self.messages = messages
|
|
66
|
+
self.description = description
|
|
67
|
+
|
|
68
|
+
def __repr__(self) -> str:
|
|
69
|
+
return f"PromptGetResult(messages={self.messages!r}, description={self.description!r})"
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
class MCPTestClient:
|
|
73
|
+
"""First-class in-process testing client for MCPToolForge MCPServers."""
|
|
74
|
+
|
|
75
|
+
def __init__(self, server: Any) -> None:
|
|
76
|
+
self._server = server
|
|
77
|
+
|
|
78
|
+
def _run_sync(self, coro_fn: Any) -> Any:
|
|
79
|
+
import anyio
|
|
80
|
+
|
|
81
|
+
try:
|
|
82
|
+
asyncio.get_running_loop()
|
|
83
|
+
# If there's already a running loop, running anyio.run raises a RuntimeError.
|
|
84
|
+
# We must await or warn. To be safe:
|
|
85
|
+
raise MCPToolForgeTestingError(
|
|
86
|
+
"Cannot execute synchronous test method inside a running event loop. "
|
|
87
|
+
"Use the async version (e.g. `await client.call_tool_async(...)`) instead."
|
|
88
|
+
)
|
|
89
|
+
except RuntimeError:
|
|
90
|
+
pass
|
|
91
|
+
|
|
92
|
+
return anyio.run(coro_fn)
|
|
93
|
+
|
|
94
|
+
# --- Startup/Shutdown Lifecycle ---
|
|
95
|
+
|
|
96
|
+
def start(self) -> None:
|
|
97
|
+
"""Run startup hooks synchronously."""
|
|
98
|
+
for hook in self._server.startup_hooks:
|
|
99
|
+
if inspect.iscoroutinefunction(hook):
|
|
100
|
+
self._run_sync(lambda h=hook: h())
|
|
101
|
+
else:
|
|
102
|
+
hook()
|
|
103
|
+
|
|
104
|
+
def stop(self) -> None:
|
|
105
|
+
"""Run shutdown hooks synchronously."""
|
|
106
|
+
for hook in self._server.shutdown_hooks:
|
|
107
|
+
if inspect.iscoroutinefunction(hook):
|
|
108
|
+
self._run_sync(lambda h=hook: h())
|
|
109
|
+
else:
|
|
110
|
+
hook()
|
|
111
|
+
|
|
112
|
+
async def start_async(self) -> None:
|
|
113
|
+
"""Run startup hooks asynchronously."""
|
|
114
|
+
for hook in self._server.startup_hooks:
|
|
115
|
+
if inspect.iscoroutinefunction(hook):
|
|
116
|
+
await hook()
|
|
117
|
+
else:
|
|
118
|
+
hook()
|
|
119
|
+
|
|
120
|
+
async def stop_async(self) -> None:
|
|
121
|
+
"""Run shutdown hooks asynchronously."""
|
|
122
|
+
for hook in self._server.shutdown_hooks:
|
|
123
|
+
if inspect.iscoroutinefunction(hook):
|
|
124
|
+
await hook()
|
|
125
|
+
else:
|
|
126
|
+
hook()
|
|
127
|
+
|
|
128
|
+
def __enter__(self) -> "MCPTestClient":
|
|
129
|
+
self.start()
|
|
130
|
+
return self
|
|
131
|
+
|
|
132
|
+
def __exit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None:
|
|
133
|
+
self.stop()
|
|
134
|
+
|
|
135
|
+
async def __aenter__(self) -> "MCPTestClient":
|
|
136
|
+
await self.start_async()
|
|
137
|
+
return self
|
|
138
|
+
|
|
139
|
+
async def __aexit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None:
|
|
140
|
+
await self.stop_async()
|
|
141
|
+
|
|
142
|
+
# --- Tools ---
|
|
143
|
+
|
|
144
|
+
def list_tools(self) -> list[Any]:
|
|
145
|
+
"""List registered tools from the MCPServer."""
|
|
146
|
+
return self._server.list_tools()
|
|
147
|
+
|
|
148
|
+
async def call_tool_async(self, name: str, arguments: dict[str, Any] | None = None) -> Any:
|
|
149
|
+
"""Execute a tool asynchronously through the server's validation and middleware pipeline."""
|
|
150
|
+
if not self._server.has_tool(name):
|
|
151
|
+
raise ToolNotFoundError(f"Tool '{name}' is not registered.")
|
|
152
|
+
|
|
153
|
+
tf_tool = self._server.get_tool(name)
|
|
154
|
+
raw_args = arguments or {}
|
|
155
|
+
|
|
156
|
+
if self._server.middlewares:
|
|
157
|
+
from mcptoolforge.middleware.context import MiddlewareContext
|
|
158
|
+
from mcptoolforge.middleware.manager import build_chain, is_async_callable
|
|
159
|
+
|
|
160
|
+
context = MiddlewareContext(
|
|
161
|
+
tool_name=tf_tool.name,
|
|
162
|
+
tool=tf_tool,
|
|
163
|
+
arguments=raw_args,
|
|
164
|
+
server=self._server,
|
|
165
|
+
)
|
|
166
|
+
|
|
167
|
+
called_next = False
|
|
168
|
+
|
|
169
|
+
if inspect.iscoroutinefunction(tf_tool.fn):
|
|
170
|
+
|
|
171
|
+
async def final_call() -> Any:
|
|
172
|
+
nonlocal called_next
|
|
173
|
+
called_next = True
|
|
174
|
+
from mcptoolforge.validation import validate_tool_arguments
|
|
175
|
+
|
|
176
|
+
validated_args = validate_tool_arguments(tf_tool, context.arguments)
|
|
177
|
+
try:
|
|
178
|
+
return await tf_tool.fn(**validated_args)
|
|
179
|
+
except Exception as inner_e:
|
|
180
|
+
if not isinstance(inner_e, (ToolValidationError, ToolExecutionError)):
|
|
181
|
+
raise ToolExecutionError(
|
|
182
|
+
f"Error executing tool '{tf_tool.name}': {inner_e}"
|
|
183
|
+
) from inner_e
|
|
184
|
+
raise
|
|
185
|
+
|
|
186
|
+
else:
|
|
187
|
+
|
|
188
|
+
def final_call() -> Any:
|
|
189
|
+
nonlocal called_next
|
|
190
|
+
called_next = True
|
|
191
|
+
from mcptoolforge.validation import validate_tool_arguments
|
|
192
|
+
|
|
193
|
+
validated_args = validate_tool_arguments(tf_tool, context.arguments)
|
|
194
|
+
try:
|
|
195
|
+
return tf_tool.fn(**validated_args)
|
|
196
|
+
except Exception as inner_e:
|
|
197
|
+
if not isinstance(inner_e, (ToolValidationError, ToolExecutionError)):
|
|
198
|
+
raise ToolExecutionError(
|
|
199
|
+
f"Error executing tool '{tf_tool.name}': {inner_e}"
|
|
200
|
+
) from inner_e
|
|
201
|
+
raise
|
|
202
|
+
|
|
203
|
+
start_time = time.perf_counter()
|
|
204
|
+
chain_callable = build_chain(self._server.middlewares, context, final_call)
|
|
205
|
+
|
|
206
|
+
if inspect.iscoroutinefunction(chain_callable) or is_async_callable(chain_callable):
|
|
207
|
+
result = await chain_callable()
|
|
208
|
+
else:
|
|
209
|
+
result = chain_callable()
|
|
210
|
+
|
|
211
|
+
context.duration = time.perf_counter() - start_time
|
|
212
|
+
return result
|
|
213
|
+
else:
|
|
214
|
+
from mcptoolforge.execution import execute_tool
|
|
215
|
+
|
|
216
|
+
return await execute_tool(tf_tool, raw_args)
|
|
217
|
+
|
|
218
|
+
def call_tool(self, name: str, arguments: dict[str, Any] | None = None) -> Any:
|
|
219
|
+
"""Execute a tool synchronously through the server's validation and middleware pipeline."""
|
|
220
|
+
return self._run_sync(lambda: self.call_tool_async(name, arguments))
|
|
221
|
+
|
|
222
|
+
# --- Resources ---
|
|
223
|
+
|
|
224
|
+
def list_resources(self) -> list[Any]:
|
|
225
|
+
"""List registered resources from the MCPServer."""
|
|
226
|
+
return self._server.list_resources()
|
|
227
|
+
|
|
228
|
+
async def read_resource_async(self, uri: str) -> ResourceReadResult:
|
|
229
|
+
"""Read a resource asynchronously through the server's resource handler."""
|
|
230
|
+
if not self._server.has_resource(uri):
|
|
231
|
+
raise ResourceNotFoundError(f"Resource '{uri}' is not registered.")
|
|
232
|
+
|
|
233
|
+
tf_resource = self._server.get_resource(uri)
|
|
234
|
+
|
|
235
|
+
try:
|
|
236
|
+
if inspect.iscoroutinefunction(tf_resource.fn):
|
|
237
|
+
raw_result = await tf_resource.fn()
|
|
238
|
+
else:
|
|
239
|
+
raw_result = tf_resource.fn()
|
|
240
|
+
except Exception as e:
|
|
241
|
+
raise ResourceExecutionError(f"Resource execution failed: {e}") from e
|
|
242
|
+
|
|
243
|
+
mime_type = tf_resource.mime_type
|
|
244
|
+
if isinstance(raw_result, (dict, list)):
|
|
245
|
+
if mime_type is None:
|
|
246
|
+
mime_type = "application/json"
|
|
247
|
+
try:
|
|
248
|
+
serialized_text = json.dumps(raw_result)
|
|
249
|
+
except Exception as json_e:
|
|
250
|
+
raise ResourceExecutionError(
|
|
251
|
+
f"Failed to serialize resource content to JSON: {json_e}"
|
|
252
|
+
) from json_e
|
|
253
|
+
return ResourceReadResult(
|
|
254
|
+
uri=tf_resource.uri,
|
|
255
|
+
text=serialized_text,
|
|
256
|
+
mime_type=mime_type,
|
|
257
|
+
)
|
|
258
|
+
elif isinstance(raw_result, str):
|
|
259
|
+
return ResourceReadResult(
|
|
260
|
+
uri=tf_resource.uri,
|
|
261
|
+
text=raw_result,
|
|
262
|
+
mime_type=mime_type,
|
|
263
|
+
)
|
|
264
|
+
elif isinstance(raw_result, bytes):
|
|
265
|
+
b64_str = base64.b64encode(raw_result).decode("utf-8")
|
|
266
|
+
return ResourceReadResult(
|
|
267
|
+
uri=tf_resource.uri,
|
|
268
|
+
blob=b64_str,
|
|
269
|
+
mime_type=mime_type,
|
|
270
|
+
)
|
|
271
|
+
else:
|
|
272
|
+
raise ResourceExecutionError(
|
|
273
|
+
f"Unsupported resource return type: {type(raw_result).__name__}"
|
|
274
|
+
)
|
|
275
|
+
|
|
276
|
+
def read_resource(self, uri: str) -> ResourceReadResult:
|
|
277
|
+
"""Read a resource synchronously through the server's resource handler."""
|
|
278
|
+
return self._run_sync(lambda: self.read_resource_async(uri))
|
|
279
|
+
|
|
280
|
+
# --- Prompts ---
|
|
281
|
+
|
|
282
|
+
def list_prompts(self) -> list[Any]:
|
|
283
|
+
"""List registered prompts from the MCPServer."""
|
|
284
|
+
return self._server.list_prompts()
|
|
285
|
+
|
|
286
|
+
async def get_prompt_async(
|
|
287
|
+
self, name: str, arguments: dict[str, Any] | None = None
|
|
288
|
+
) -> PromptGetResult:
|
|
289
|
+
"""Retrieve a prompt template asynchronously with validation and mapping."""
|
|
290
|
+
if not self._server.has_prompt(name):
|
|
291
|
+
raise ToolNotFoundError(f"Prompt '{name}' is not registered.")
|
|
292
|
+
|
|
293
|
+
tf_prompt = self._server.get_prompt(name)
|
|
294
|
+
raw_args = arguments or {}
|
|
295
|
+
|
|
296
|
+
from mcptoolforge.validation import validate_prompt_arguments
|
|
297
|
+
|
|
298
|
+
try:
|
|
299
|
+
validated_args = validate_prompt_arguments(tf_prompt, raw_args)
|
|
300
|
+
except Exception as val_e:
|
|
301
|
+
raise ToolValidationError(f"Prompt argument validation failed: {val_e}") from val_e
|
|
302
|
+
|
|
303
|
+
try:
|
|
304
|
+
if inspect.iscoroutinefunction(tf_prompt.fn):
|
|
305
|
+
raw_result = await tf_prompt.fn(**validated_args)
|
|
306
|
+
else:
|
|
307
|
+
raw_result = tf_prompt.fn(**validated_args)
|
|
308
|
+
except Exception as e:
|
|
309
|
+
raise ToolExecutionError(f"Prompt execution failed: {e}") from e
|
|
310
|
+
|
|
311
|
+
messages = []
|
|
312
|
+
|
|
313
|
+
def make_message(text_val: str, role_val: str = "user") -> PromptMessageResult:
|
|
314
|
+
if role_val not in ("user", "assistant"):
|
|
315
|
+
raise ToolExecutionError(
|
|
316
|
+
f"Role '{role_val}' is not supported. Only 'user' or 'assistant' are allowed."
|
|
317
|
+
)
|
|
318
|
+
return PromptMessageResult(role=role_val, content=text_val)
|
|
319
|
+
|
|
320
|
+
if isinstance(raw_result, str):
|
|
321
|
+
messages.append(make_message(raw_result))
|
|
322
|
+
elif isinstance(raw_result, dict):
|
|
323
|
+
role = raw_result.get("role", "user")
|
|
324
|
+
content = raw_result.get("content", "")
|
|
325
|
+
if not isinstance(content, str):
|
|
326
|
+
tname = type(content).__name__
|
|
327
|
+
raise ToolExecutionError(f"Prompt dict content must be a string, received {tname}")
|
|
328
|
+
messages.append(make_message(content, role))
|
|
329
|
+
elif isinstance(raw_result, list):
|
|
330
|
+
for item in raw_result:
|
|
331
|
+
if isinstance(item, str):
|
|
332
|
+
messages.append(make_message(item))
|
|
333
|
+
elif isinstance(item, dict):
|
|
334
|
+
role = item.get("role", "user")
|
|
335
|
+
content = item.get("content", "")
|
|
336
|
+
if not isinstance(content, str):
|
|
337
|
+
tname = type(content).__name__
|
|
338
|
+
raise ToolExecutionError(
|
|
339
|
+
f"Prompt dict content must be a string, received {tname}"
|
|
340
|
+
)
|
|
341
|
+
messages.append(make_message(content, role))
|
|
342
|
+
elif hasattr(item, "role") and hasattr(item, "content"):
|
|
343
|
+
content_text = item.content
|
|
344
|
+
if hasattr(content_text, "text"):
|
|
345
|
+
content_text = content_text.text
|
|
346
|
+
messages.append(make_message(str(content_text), str(item.role)))
|
|
347
|
+
else:
|
|
348
|
+
tname = type(item).__name__
|
|
349
|
+
raise ToolExecutionError(
|
|
350
|
+
f"Unsupported item type in prompt result list: {tname}"
|
|
351
|
+
)
|
|
352
|
+
elif hasattr(raw_result, "role") and hasattr(raw_result, "content"):
|
|
353
|
+
content_text = raw_result.content
|
|
354
|
+
if hasattr(content_text, "text"):
|
|
355
|
+
content_text = content_text.text
|
|
356
|
+
messages.append(make_message(str(content_text), str(raw_result.role)))
|
|
357
|
+
else:
|
|
358
|
+
raise ToolExecutionError(f"Unsupported prompt return type: {type(raw_result).__name__}")
|
|
359
|
+
|
|
360
|
+
return PromptGetResult(messages=messages, description=tf_prompt.description)
|
|
361
|
+
|
|
362
|
+
def get_prompt(self, name: str, arguments: dict[str, Any] | None = None) -> PromptGetResult:
|
|
363
|
+
"""Retrieve a prompt template synchronously with validation and mapping."""
|
|
364
|
+
return self._run_sync(lambda: self.get_prompt_async(name, arguments))
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
# Transport package placeholder for MCPToolForge
|
|
@@ -0,0 +1,229 @@
|
|
|
1
|
+
import enum
|
|
2
|
+
import inspect
|
|
3
|
+
import types
|
|
4
|
+
import typing
|
|
5
|
+
from typing import Any
|
|
6
|
+
|
|
7
|
+
from mcptoolforge.errors import ToolValidationError
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def _validate_type(val: Any, ann: Any, param_name: str, tool_name: str) -> Any:
|
|
11
|
+
"""Validate and normalize a value against a type annotation."""
|
|
12
|
+
if ann is inspect.Parameter.empty or ann is Any:
|
|
13
|
+
return val
|
|
14
|
+
|
|
15
|
+
# None / NoneType
|
|
16
|
+
if ann is None or ann is type(None):
|
|
17
|
+
if val is None:
|
|
18
|
+
return None
|
|
19
|
+
raise ToolValidationError(
|
|
20
|
+
f"Tool '{tool_name}': parameter '{param_name}' expected None, "
|
|
21
|
+
f"received {type(val).__name__}."
|
|
22
|
+
)
|
|
23
|
+
|
|
24
|
+
# Union Types / Optional
|
|
25
|
+
origin = getattr(ann, "__origin__", None)
|
|
26
|
+
if origin is typing.Union or (hasattr(types, "UnionType") and isinstance(ann, types.UnionType)):
|
|
27
|
+
args = typing.get_args(ann)
|
|
28
|
+
# Try to validate against each union type option
|
|
29
|
+
for arg in args:
|
|
30
|
+
try:
|
|
31
|
+
return _validate_type(val, arg, param_name, tool_name)
|
|
32
|
+
except ToolValidationError:
|
|
33
|
+
continue
|
|
34
|
+
|
|
35
|
+
# If all options failed, build a clean error message
|
|
36
|
+
expected_types = " | ".join(
|
|
37
|
+
getattr(arg, "__name__", str(arg)) for arg in args if arg is not type(None)
|
|
38
|
+
)
|
|
39
|
+
if type(None) in args:
|
|
40
|
+
expected_types += " | None"
|
|
41
|
+
raise ToolValidationError(
|
|
42
|
+
f"Tool '{tool_name}': parameter '{param_name}' expected {expected_types}, "
|
|
43
|
+
f"received {type(val).__name__}."
|
|
44
|
+
)
|
|
45
|
+
|
|
46
|
+
# Lists (list, List)
|
|
47
|
+
if ann is list or origin is list or origin is list:
|
|
48
|
+
if not isinstance(val, list):
|
|
49
|
+
raise ToolValidationError(
|
|
50
|
+
f"Tool '{tool_name}': parameter '{param_name}' expected list, "
|
|
51
|
+
f"received {type(val).__name__}."
|
|
52
|
+
)
|
|
53
|
+
args = typing.get_args(ann)
|
|
54
|
+
if not args:
|
|
55
|
+
return val
|
|
56
|
+
item_type = args[0]
|
|
57
|
+
validated_list = []
|
|
58
|
+
for i, item in enumerate(val):
|
|
59
|
+
try:
|
|
60
|
+
validated_list.append(
|
|
61
|
+
_validate_type(item, item_type, f"{param_name}[{i}]", tool_name)
|
|
62
|
+
)
|
|
63
|
+
except ToolValidationError as e:
|
|
64
|
+
raise ToolValidationError(
|
|
65
|
+
f"Tool '{tool_name}': parameter '{param_name}' "
|
|
66
|
+
f"element at index {i} invalid: {e}"
|
|
67
|
+
) from e
|
|
68
|
+
return validated_list
|
|
69
|
+
|
|
70
|
+
# Dictionaries (dict, Dict)
|
|
71
|
+
if ann is dict or origin is dict or origin is dict:
|
|
72
|
+
if not isinstance(val, dict):
|
|
73
|
+
raise ToolValidationError(
|
|
74
|
+
f"Tool '{tool_name}': parameter '{param_name}' expected dict, "
|
|
75
|
+
f"received {type(val).__name__}."
|
|
76
|
+
)
|
|
77
|
+
args = typing.get_args(ann)
|
|
78
|
+
if not args:
|
|
79
|
+
return val
|
|
80
|
+
key_type, val_type = args[0], args[1]
|
|
81
|
+
validated_dict = {}
|
|
82
|
+
for k, v in val.items():
|
|
83
|
+
try:
|
|
84
|
+
validated_k = _validate_type(k, key_type, f"{param_name}.key", tool_name)
|
|
85
|
+
except ToolValidationError as e:
|
|
86
|
+
raise ToolValidationError(
|
|
87
|
+
f"Tool '{tool_name}': parameter '{param_name}' "
|
|
88
|
+
f"dictionary key '{k}' invalid: {e}"
|
|
89
|
+
) from e
|
|
90
|
+
try:
|
|
91
|
+
validated_v = _validate_type(v, val_type, f"{param_name}['{k}']", tool_name)
|
|
92
|
+
except ToolValidationError as e:
|
|
93
|
+
raise ToolValidationError(
|
|
94
|
+
f"Tool '{tool_name}': parameter '{param_name}' "
|
|
95
|
+
f"dictionary value for key '{k}' invalid: {e}"
|
|
96
|
+
) from e
|
|
97
|
+
validated_dict[validated_k] = validated_v
|
|
98
|
+
return validated_dict
|
|
99
|
+
|
|
100
|
+
# Enums
|
|
101
|
+
if isinstance(ann, type) and issubclass(ann, enum.Enum):
|
|
102
|
+
if isinstance(val, ann):
|
|
103
|
+
return val
|
|
104
|
+
try:
|
|
105
|
+
return ann(val)
|
|
106
|
+
except ValueError:
|
|
107
|
+
for member in ann:
|
|
108
|
+
if member.name == val:
|
|
109
|
+
return member
|
|
110
|
+
valid_choices = [m.value for m in ann]
|
|
111
|
+
raise ToolValidationError(
|
|
112
|
+
f"Tool '{tool_name}': parameter '{param_name}' expected one of {valid_choices}, "
|
|
113
|
+
f"received '{val}'."
|
|
114
|
+
) from None
|
|
115
|
+
|
|
116
|
+
# Primitive types (strictly rejecting booleans from matching integer/string/float)
|
|
117
|
+
if ann is str:
|
|
118
|
+
if isinstance(val, str) and not isinstance(val, bool):
|
|
119
|
+
return val
|
|
120
|
+
raise ToolValidationError(
|
|
121
|
+
f"Tool '{tool_name}': parameter '{param_name}' expected string, "
|
|
122
|
+
f"received {type(val).__name__}."
|
|
123
|
+
)
|
|
124
|
+
elif ann is int:
|
|
125
|
+
if isinstance(val, int) and not isinstance(val, bool):
|
|
126
|
+
return val
|
|
127
|
+
raise ToolValidationError(
|
|
128
|
+
f"Tool '{tool_name}': parameter '{param_name}' expected integer, "
|
|
129
|
+
f"received {type(val).__name__}."
|
|
130
|
+
)
|
|
131
|
+
elif ann is float:
|
|
132
|
+
if isinstance(val, (int, float)) and not isinstance(val, bool):
|
|
133
|
+
return float(val)
|
|
134
|
+
raise ToolValidationError(
|
|
135
|
+
f"Tool '{tool_name}': parameter '{param_name}' expected float, "
|
|
136
|
+
f"received {type(val).__name__}."
|
|
137
|
+
)
|
|
138
|
+
elif ann is bool:
|
|
139
|
+
if isinstance(val, bool):
|
|
140
|
+
return val
|
|
141
|
+
raise ToolValidationError(
|
|
142
|
+
f"Tool '{tool_name}': parameter '{param_name}' expected boolean, "
|
|
143
|
+
f"received {type(val).__name__}."
|
|
144
|
+
)
|
|
145
|
+
|
|
146
|
+
# Fallback for unhandled type annotations
|
|
147
|
+
return val
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
def validate_tool_arguments(tool: Any, arguments: dict[str, Any]) -> dict[str, Any]:
|
|
151
|
+
"""Validate and normalize incoming arguments against a Tool's parameters."""
|
|
152
|
+
validated: dict[str, Any] = {}
|
|
153
|
+
|
|
154
|
+
# 1. Reject unexpected parameters
|
|
155
|
+
for k in arguments:
|
|
156
|
+
if k not in tool.parameters:
|
|
157
|
+
raise ToolValidationError(f"Tool '{tool.name}': unexpected parameter '{k}'.")
|
|
158
|
+
|
|
159
|
+
# 2. Validate expected parameters
|
|
160
|
+
for name, param in tool.parameters.items():
|
|
161
|
+
if name not in arguments:
|
|
162
|
+
if param.required:
|
|
163
|
+
raise ToolValidationError(
|
|
164
|
+
f"Tool '{tool.name}': missing required parameter '{name}'."
|
|
165
|
+
)
|
|
166
|
+
else:
|
|
167
|
+
validated[name] = param.default
|
|
168
|
+
else:
|
|
169
|
+
val = arguments[name]
|
|
170
|
+
validated[name] = _validate_type(val, param.annotation, name, tool.name)
|
|
171
|
+
|
|
172
|
+
return validated
|
|
173
|
+
|
|
174
|
+
|
|
175
|
+
def _coerce_prompt_type(val: Any, ann: Any) -> Any:
|
|
176
|
+
"""Attempt to coerce string values into correct types for prompt parameters."""
|
|
177
|
+
if isinstance(val, str):
|
|
178
|
+
if ann is int:
|
|
179
|
+
try:
|
|
180
|
+
return int(val)
|
|
181
|
+
except ValueError:
|
|
182
|
+
pass
|
|
183
|
+
elif ann is float:
|
|
184
|
+
try:
|
|
185
|
+
return float(val)
|
|
186
|
+
except ValueError:
|
|
187
|
+
pass
|
|
188
|
+
elif ann is bool:
|
|
189
|
+
if val.lower() == "true":
|
|
190
|
+
return True
|
|
191
|
+
if val.lower() == "false":
|
|
192
|
+
return False
|
|
193
|
+
|
|
194
|
+
# Support unions (e.g. Union[int, str], int | None)
|
|
195
|
+
origin = getattr(ann, "__origin__", None)
|
|
196
|
+
if origin is typing.Union or (
|
|
197
|
+
hasattr(types, "UnionType") and isinstance(ann, types.UnionType)
|
|
198
|
+
):
|
|
199
|
+
for arg in typing.get_args(ann):
|
|
200
|
+
coerced = _coerce_prompt_type(val, arg)
|
|
201
|
+
if coerced is not val:
|
|
202
|
+
return coerced
|
|
203
|
+
return val
|
|
204
|
+
|
|
205
|
+
|
|
206
|
+
def validate_prompt_arguments(prompt: Any, arguments: dict[str, Any]) -> dict[str, Any]:
|
|
207
|
+
"""Validate and normalize incoming arguments against a Prompt's parameters."""
|
|
208
|
+
validated: dict[str, Any] = {}
|
|
209
|
+
|
|
210
|
+
# 1. Reject unexpected parameters
|
|
211
|
+
for k in arguments:
|
|
212
|
+
if k not in prompt.parameters:
|
|
213
|
+
raise ToolValidationError(f"Prompt '{prompt.name}': unexpected parameter '{k}'.")
|
|
214
|
+
|
|
215
|
+
# 2. Validate expected parameters
|
|
216
|
+
for name, param in prompt.parameters.items():
|
|
217
|
+
if name not in arguments:
|
|
218
|
+
if param.required:
|
|
219
|
+
raise ToolValidationError(
|
|
220
|
+
f"Prompt '{prompt.name}': missing required parameter '{name}'."
|
|
221
|
+
)
|
|
222
|
+
else:
|
|
223
|
+
validated[name] = param.default
|
|
224
|
+
else:
|
|
225
|
+
val = arguments[name]
|
|
226
|
+
coerced_val = _coerce_prompt_type(val, param.annotation)
|
|
227
|
+
validated[name] = _validate_type(coerced_val, param.annotation, name, prompt.name)
|
|
228
|
+
|
|
229
|
+
return validated
|