schemarouter 0.2.0a1__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.
- schemarouter/__init__.py +88 -0
- schemarouter/_version.py +8 -0
- schemarouter/adapters/__init__.py +29 -0
- schemarouter/adapters/base.py +71 -0
- schemarouter/adapters/mcp.py +179 -0
- schemarouter/adapters/openapi.py +418 -0
- schemarouter/adapters/optimade.py +656 -0
- schemarouter/adapters/python.py +188 -0
- schemarouter/analyzers/__init__.py +3 -0
- schemarouter/analyzers/model.py +185 -0
- schemarouter/errors.py +50 -0
- schemarouter/executor.py +207 -0
- schemarouter/ingestion.py +348 -0
- schemarouter/integrations/__init__.py +3 -0
- schemarouter/integrations/langchain.py +96 -0
- schemarouter/models.py +154 -0
- schemarouter/planner.py +253 -0
- schemarouter/policy.py +51 -0
- schemarouter/proposals.py +391 -0
- schemarouter/py.typed +0 -0
- schemarouter/registry.py +84 -0
- schemarouter/runs.py +77 -0
- schemarouter/runtime.py +686 -0
- schemarouter/validation.py +96 -0
- schemarouter-0.2.0a1.dist-info/METADATA +297 -0
- schemarouter-0.2.0a1.dist-info/RECORD +28 -0
- schemarouter-0.2.0a1.dist-info/WHEEL +4 -0
- schemarouter-0.2.0a1.dist-info/licenses/LICENSE +21 -0
schemarouter/runtime.py
ADDED
|
@@ -0,0 +1,686 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import asyncio
|
|
4
|
+
from collections.abc import AsyncIterator, Awaitable, Callable, Iterator, Sequence
|
|
5
|
+
from typing import Any, TypeVar
|
|
6
|
+
from urllib.parse import urlparse
|
|
7
|
+
from uuid import uuid4
|
|
8
|
+
|
|
9
|
+
import httpx
|
|
10
|
+
from pydantic import TypeAdapter
|
|
11
|
+
|
|
12
|
+
from .adapters.base import AdapterRegistry, SourceAdapter
|
|
13
|
+
from .adapters.openapi import OpenAPIRemoteInvoker
|
|
14
|
+
from .adapters.python import PythonCallableInvoker, callable_options, tool_from_callable
|
|
15
|
+
from .errors import ProposalApprovalError, RegistrationError
|
|
16
|
+
from .executor import RegistryExecutor
|
|
17
|
+
from .ingestion import SourceKind, URLSchemaLoader
|
|
18
|
+
from .models import ExecutionPlan, PlanRequest, ToolResult, ToolSpec
|
|
19
|
+
from .planner import QueryAnalyzer, SchemaPlanner
|
|
20
|
+
from .policy import ExecutionPolicy
|
|
21
|
+
from .proposals import DocumentationModelCallable, SchemaProposal, inspect_documentation_url
|
|
22
|
+
from .registry import InMemoryRegistry, ToolRegistry
|
|
23
|
+
from .runs import RunConfig, RunEvent
|
|
24
|
+
|
|
25
|
+
_T = TypeVar("_T")
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def _coerce_config(config: RunConfig | dict[str, Any] | None) -> RunConfig:
|
|
29
|
+
if config is None:
|
|
30
|
+
return RunConfig()
|
|
31
|
+
if isinstance(config, RunConfig):
|
|
32
|
+
return config
|
|
33
|
+
return RunConfig.model_validate(config)
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def _run_sync(factory: Callable[[], Awaitable[_T]]) -> _T:
|
|
37
|
+
try:
|
|
38
|
+
asyncio.get_running_loop()
|
|
39
|
+
except RuntimeError:
|
|
40
|
+
return asyncio.run(factory())
|
|
41
|
+
raise RuntimeError(
|
|
42
|
+
"synchronous SchemaRouter API cannot run inside an active event loop; "
|
|
43
|
+
"use the async API instead"
|
|
44
|
+
)
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def _stream_sync(factory: Callable[[], AsyncIterator[_T]]) -> Iterator[_T]:
|
|
48
|
+
try:
|
|
49
|
+
asyncio.get_running_loop()
|
|
50
|
+
except RuntimeError:
|
|
51
|
+
pass
|
|
52
|
+
else:
|
|
53
|
+
raise RuntimeError(
|
|
54
|
+
"synchronous SchemaRouter streaming cannot run inside an active event loop; "
|
|
55
|
+
"use the async streaming API instead"
|
|
56
|
+
)
|
|
57
|
+
|
|
58
|
+
loop = asyncio.new_event_loop()
|
|
59
|
+
iterator = factory()
|
|
60
|
+
try:
|
|
61
|
+
while True:
|
|
62
|
+
try:
|
|
63
|
+
yield loop.run_until_complete(iterator.__anext__())
|
|
64
|
+
except StopAsyncIteration:
|
|
65
|
+
break
|
|
66
|
+
finally:
|
|
67
|
+
loop.run_until_complete(iterator.aclose())
|
|
68
|
+
loop.close()
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
class SchemaRouter:
|
|
72
|
+
"""High-level facade for schema-aware planning and execution."""
|
|
73
|
+
|
|
74
|
+
def __init__(
|
|
75
|
+
self,
|
|
76
|
+
*,
|
|
77
|
+
analyzer: QueryAnalyzer | None = None,
|
|
78
|
+
http_client: httpx.AsyncClient | None = None,
|
|
79
|
+
policy: ExecutionPolicy | None = None,
|
|
80
|
+
registry: ToolRegistry | None = None,
|
|
81
|
+
adapter_registry: AdapterRegistry | None = None,
|
|
82
|
+
) -> None:
|
|
83
|
+
self.registry = registry if registry is not None else InMemoryRegistry()
|
|
84
|
+
self.planner = SchemaPlanner(self.registry, analyzer=analyzer)
|
|
85
|
+
self.executor = RegistryExecutor(self.registry, policy=policy)
|
|
86
|
+
self.loader = URLSchemaLoader(
|
|
87
|
+
self.registry,
|
|
88
|
+
self.executor,
|
|
89
|
+
http_client=http_client,
|
|
90
|
+
adapters=adapter_registry,
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
@property
|
|
94
|
+
def input_schema(self) -> dict[str, Any]:
|
|
95
|
+
return PlanRequest.model_json_schema()
|
|
96
|
+
|
|
97
|
+
@property
|
|
98
|
+
def output_schema(self) -> dict[str, Any]:
|
|
99
|
+
return TypeAdapter(list[ToolResult]).json_schema()
|
|
100
|
+
|
|
101
|
+
@property
|
|
102
|
+
def config_schema(self) -> dict[str, Any]:
|
|
103
|
+
return RunConfig.model_json_schema()
|
|
104
|
+
|
|
105
|
+
def with_config(
|
|
106
|
+
self,
|
|
107
|
+
config: RunConfig | dict[str, Any],
|
|
108
|
+
) -> ConfiguredSchemaRouter:
|
|
109
|
+
return ConfiguredSchemaRouter(self, _coerce_config(config))
|
|
110
|
+
|
|
111
|
+
@classmethod
|
|
112
|
+
async def from_url(
|
|
113
|
+
cls,
|
|
114
|
+
url: str,
|
|
115
|
+
*,
|
|
116
|
+
kind: SourceKind = "auto",
|
|
117
|
+
name: str | None = None,
|
|
118
|
+
namespace: str | None = None,
|
|
119
|
+
analyzer: QueryAnalyzer | None = None,
|
|
120
|
+
http_client: httpx.AsyncClient | None = None,
|
|
121
|
+
policy: ExecutionPolicy | None = None,
|
|
122
|
+
registry: ToolRegistry | None = None,
|
|
123
|
+
adapter_registry: AdapterRegistry | None = None,
|
|
124
|
+
base_url: str | None = None,
|
|
125
|
+
schema_headers: dict[str, str] | None = None,
|
|
126
|
+
trusted_headers: dict[str, str] | None = None,
|
|
127
|
+
) -> SchemaRouter:
|
|
128
|
+
router = cls(
|
|
129
|
+
analyzer=analyzer,
|
|
130
|
+
http_client=http_client,
|
|
131
|
+
policy=policy,
|
|
132
|
+
registry=registry,
|
|
133
|
+
adapter_registry=adapter_registry,
|
|
134
|
+
)
|
|
135
|
+
await router.add_url(
|
|
136
|
+
url,
|
|
137
|
+
kind=kind,
|
|
138
|
+
name=name,
|
|
139
|
+
namespace=namespace,
|
|
140
|
+
base_url=base_url,
|
|
141
|
+
schema_headers=schema_headers,
|
|
142
|
+
trusted_headers=trusted_headers,
|
|
143
|
+
)
|
|
144
|
+
return router
|
|
145
|
+
|
|
146
|
+
def add_tool(self, tool: ToolSpec, *, replace: bool = False) -> str:
|
|
147
|
+
return self.registry.register(tool, replace=replace)
|
|
148
|
+
|
|
149
|
+
def register_adapter(
|
|
150
|
+
self,
|
|
151
|
+
adapter: SourceAdapter,
|
|
152
|
+
*,
|
|
153
|
+
replace: bool = False,
|
|
154
|
+
) -> None:
|
|
155
|
+
self.loader.register_adapter(adapter, replace=replace)
|
|
156
|
+
|
|
157
|
+
@property
|
|
158
|
+
def adapter_registry(self) -> AdapterRegistry:
|
|
159
|
+
return self.loader.adapters
|
|
160
|
+
|
|
161
|
+
def add_callable(
|
|
162
|
+
self,
|
|
163
|
+
function: Callable[..., Any],
|
|
164
|
+
*,
|
|
165
|
+
name: str | None = None,
|
|
166
|
+
namespace: str | None = None,
|
|
167
|
+
description: str | None = None,
|
|
168
|
+
read_only: bool | None = None,
|
|
169
|
+
destructive: bool | None = None,
|
|
170
|
+
replace: bool = False,
|
|
171
|
+
) -> str:
|
|
172
|
+
decorated = callable_options(function)
|
|
173
|
+
tool = tool_from_callable(
|
|
174
|
+
function,
|
|
175
|
+
name=name if name is not None else decorated.get("name"),
|
|
176
|
+
namespace=(
|
|
177
|
+
namespace if namespace is not None else decorated.get("namespace")
|
|
178
|
+
),
|
|
179
|
+
description=(
|
|
180
|
+
description if description is not None else decorated.get("description")
|
|
181
|
+
),
|
|
182
|
+
read_only=(
|
|
183
|
+
read_only if read_only is not None else decorated.get("read_only")
|
|
184
|
+
),
|
|
185
|
+
destructive=(
|
|
186
|
+
destructive if destructive is not None else decorated.get("destructive")
|
|
187
|
+
),
|
|
188
|
+
)
|
|
189
|
+
invoker = PythonCallableInvoker(function)
|
|
190
|
+
key = self.registry.register(tool, replace=replace)
|
|
191
|
+
self.executor.bind(key, invoker)
|
|
192
|
+
return key
|
|
193
|
+
|
|
194
|
+
async def inspect_url(
|
|
195
|
+
self,
|
|
196
|
+
url: str,
|
|
197
|
+
*,
|
|
198
|
+
model: DocumentationModelCallable,
|
|
199
|
+
timeout: float = 20.0,
|
|
200
|
+
max_document_chars: int = 60_000,
|
|
201
|
+
) -> SchemaProposal:
|
|
202
|
+
return await inspect_documentation_url(
|
|
203
|
+
url,
|
|
204
|
+
model=model,
|
|
205
|
+
http_client=self.loader.http_client,
|
|
206
|
+
timeout=timeout,
|
|
207
|
+
max_document_chars=max_document_chars,
|
|
208
|
+
)
|
|
209
|
+
|
|
210
|
+
def bind_openapi(
|
|
211
|
+
self,
|
|
212
|
+
tool_key: str,
|
|
213
|
+
*,
|
|
214
|
+
base_url: str,
|
|
215
|
+
trusted_headers: dict[str, str] | None = None,
|
|
216
|
+
timeout: float = 20.0,
|
|
217
|
+
) -> None:
|
|
218
|
+
tool = self.registry.get(tool_key)
|
|
219
|
+
if tool.metadata.get("adapter") != "openapi":
|
|
220
|
+
raise RegistrationError(
|
|
221
|
+
f"tool {tool_key!r} was not imported from OpenAPI"
|
|
222
|
+
)
|
|
223
|
+
try:
|
|
224
|
+
invoker = OpenAPIRemoteInvoker(
|
|
225
|
+
tool,
|
|
226
|
+
base_url,
|
|
227
|
+
trusted_headers=trusted_headers,
|
|
228
|
+
timeout=timeout,
|
|
229
|
+
)
|
|
230
|
+
except ValueError as exc:
|
|
231
|
+
raise RegistrationError("invalid OpenAPI execution binding") from exc
|
|
232
|
+
updated = tool.model_copy(deep=True)
|
|
233
|
+
updated.metadata.update(
|
|
234
|
+
{
|
|
235
|
+
"execution_bound": True,
|
|
236
|
+
"approved_base_url": base_url,
|
|
237
|
+
"requires_explicit_base_url": False,
|
|
238
|
+
}
|
|
239
|
+
)
|
|
240
|
+
self.registry.register(updated, replace=True)
|
|
241
|
+
self.executor.bind(tool_key, invoker)
|
|
242
|
+
|
|
243
|
+
def approve_proposal(
|
|
244
|
+
self,
|
|
245
|
+
proposal: SchemaProposal,
|
|
246
|
+
*,
|
|
247
|
+
base_url: str,
|
|
248
|
+
min_grounding_score: float = 0.8,
|
|
249
|
+
allow_mutations: bool = False,
|
|
250
|
+
replace: bool = False,
|
|
251
|
+
trusted_headers: dict[str, str] | None = None,
|
|
252
|
+
timeout: float = 20.0,
|
|
253
|
+
) -> str:
|
|
254
|
+
if proposal.status != "grounded" or proposal.tool is None:
|
|
255
|
+
raise ProposalApprovalError("proposal has no grounded tool to approve")
|
|
256
|
+
if proposal.grounding_score < min_grounding_score:
|
|
257
|
+
raise ProposalApprovalError(
|
|
258
|
+
"proposal grounding score is below the required threshold"
|
|
259
|
+
)
|
|
260
|
+
|
|
261
|
+
parsed = urlparse(base_url)
|
|
262
|
+
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
|
|
263
|
+
raise ProposalApprovalError("base_url must be an absolute http(s) URL")
|
|
264
|
+
if parsed.username or parsed.password:
|
|
265
|
+
raise ProposalApprovalError(
|
|
266
|
+
"credentials must not be embedded in base_url; use trusted runtime auth"
|
|
267
|
+
)
|
|
268
|
+
|
|
269
|
+
mutating = [
|
|
270
|
+
endpoint.name
|
|
271
|
+
for endpoint in proposal.tool.endpoints
|
|
272
|
+
if endpoint.method not in {"GET", "HEAD", "OPTIONS"}
|
|
273
|
+
]
|
|
274
|
+
if mutating and not allow_mutations:
|
|
275
|
+
raise ProposalApprovalError(
|
|
276
|
+
"proposal contains mutating endpoints; set allow_mutations=True "
|
|
277
|
+
"after explicit review"
|
|
278
|
+
)
|
|
279
|
+
|
|
280
|
+
tool = proposal.tool.model_copy(deep=True)
|
|
281
|
+
tool.metadata.update(
|
|
282
|
+
{
|
|
283
|
+
"approved_from_proposal": True,
|
|
284
|
+
"executable": True,
|
|
285
|
+
"grounding_score": proposal.grounding_score,
|
|
286
|
+
"approved_base_url": base_url,
|
|
287
|
+
}
|
|
288
|
+
)
|
|
289
|
+
try:
|
|
290
|
+
invoker = OpenAPIRemoteInvoker(
|
|
291
|
+
tool,
|
|
292
|
+
base_url,
|
|
293
|
+
trusted_headers=trusted_headers,
|
|
294
|
+
timeout=timeout,
|
|
295
|
+
)
|
|
296
|
+
except ValueError as exc:
|
|
297
|
+
raise ProposalApprovalError("invalid proposal execution binding") from exc
|
|
298
|
+
key = self.registry.register(tool, replace=replace)
|
|
299
|
+
self.executor.bind(key, invoker)
|
|
300
|
+
return key
|
|
301
|
+
|
|
302
|
+
async def add_url(
|
|
303
|
+
self,
|
|
304
|
+
url: str,
|
|
305
|
+
*,
|
|
306
|
+
kind: SourceKind = "auto",
|
|
307
|
+
name: str | None = None,
|
|
308
|
+
namespace: str | None = None,
|
|
309
|
+
replace: bool = False,
|
|
310
|
+
base_url: str | None = None,
|
|
311
|
+
schema_headers: dict[str, str] | None = None,
|
|
312
|
+
trusted_headers: dict[str, str] | None = None,
|
|
313
|
+
timeout: float = 20.0,
|
|
314
|
+
) -> ToolSpec:
|
|
315
|
+
return await self.loader.load(
|
|
316
|
+
url,
|
|
317
|
+
kind=kind,
|
|
318
|
+
name=name,
|
|
319
|
+
namespace=namespace,
|
|
320
|
+
replace=replace,
|
|
321
|
+
base_url=base_url,
|
|
322
|
+
schema_headers=schema_headers,
|
|
323
|
+
trusted_headers=trusted_headers,
|
|
324
|
+
timeout=timeout,
|
|
325
|
+
)
|
|
326
|
+
|
|
327
|
+
def plan(self, request: PlanRequest | str) -> ExecutionPlan:
|
|
328
|
+
return self.planner.plan(request)
|
|
329
|
+
|
|
330
|
+
async def aplan(self, request: PlanRequest | str) -> ExecutionPlan:
|
|
331
|
+
return await self.planner.aplan(request)
|
|
332
|
+
|
|
333
|
+
async def execute(
|
|
334
|
+
self,
|
|
335
|
+
plan: ExecutionPlan,
|
|
336
|
+
*,
|
|
337
|
+
config: RunConfig | dict[str, Any] | None = None,
|
|
338
|
+
) -> list[ToolResult]:
|
|
339
|
+
run_config = _coerce_config(config)
|
|
340
|
+
return await self.executor.execute(plan, retry=run_config.retry)
|
|
341
|
+
|
|
342
|
+
async def ainvoke(
|
|
343
|
+
self,
|
|
344
|
+
request: PlanRequest | str,
|
|
345
|
+
*,
|
|
346
|
+
config: RunConfig | dict[str, Any] | None = None,
|
|
347
|
+
) -> list[ToolResult]:
|
|
348
|
+
run_config = _coerce_config(config)
|
|
349
|
+
plan = await self.aplan(request)
|
|
350
|
+
return await self.executor.execute(plan, retry=run_config.retry)
|
|
351
|
+
|
|
352
|
+
def invoke(
|
|
353
|
+
self,
|
|
354
|
+
request: PlanRequest | str,
|
|
355
|
+
*,
|
|
356
|
+
config: RunConfig | dict[str, Any] | None = None,
|
|
357
|
+
) -> list[ToolResult]:
|
|
358
|
+
return _run_sync(lambda: self.ainvoke(request, config=config))
|
|
359
|
+
|
|
360
|
+
async def abatch(
|
|
361
|
+
self,
|
|
362
|
+
requests: Sequence[PlanRequest | str],
|
|
363
|
+
*,
|
|
364
|
+
config: RunConfig | dict[str, Any] | None = None,
|
|
365
|
+
return_exceptions: bool = False,
|
|
366
|
+
) -> list[list[ToolResult] | Exception]:
|
|
367
|
+
run_config = _coerce_config(config)
|
|
368
|
+
semaphore = asyncio.Semaphore(run_config.max_concurrency)
|
|
369
|
+
|
|
370
|
+
async def invoke_one(request: PlanRequest | str) -> list[ToolResult]:
|
|
371
|
+
async with semaphore:
|
|
372
|
+
return await self.ainvoke(request, config=run_config)
|
|
373
|
+
|
|
374
|
+
return await asyncio.gather(
|
|
375
|
+
*(invoke_one(request) for request in requests),
|
|
376
|
+
return_exceptions=return_exceptions,
|
|
377
|
+
)
|
|
378
|
+
|
|
379
|
+
def batch(
|
|
380
|
+
self,
|
|
381
|
+
requests: Sequence[PlanRequest | str],
|
|
382
|
+
*,
|
|
383
|
+
config: RunConfig | dict[str, Any] | None = None,
|
|
384
|
+
return_exceptions: bool = False,
|
|
385
|
+
) -> list[list[ToolResult] | Exception]:
|
|
386
|
+
return _run_sync(
|
|
387
|
+
lambda: self.abatch(
|
|
388
|
+
requests,
|
|
389
|
+
config=config,
|
|
390
|
+
return_exceptions=return_exceptions,
|
|
391
|
+
)
|
|
392
|
+
)
|
|
393
|
+
|
|
394
|
+
async def abatch_as_completed(
|
|
395
|
+
self,
|
|
396
|
+
requests: Sequence[PlanRequest | str],
|
|
397
|
+
*,
|
|
398
|
+
config: RunConfig | dict[str, Any] | None = None,
|
|
399
|
+
return_exceptions: bool = False,
|
|
400
|
+
) -> AsyncIterator[tuple[int, list[ToolResult] | Exception]]:
|
|
401
|
+
run_config = _coerce_config(config)
|
|
402
|
+
semaphore = asyncio.Semaphore(run_config.max_concurrency)
|
|
403
|
+
|
|
404
|
+
async def invoke_indexed(
|
|
405
|
+
index: int,
|
|
406
|
+
request: PlanRequest | str,
|
|
407
|
+
) -> tuple[int, list[ToolResult] | Exception]:
|
|
408
|
+
try:
|
|
409
|
+
async with semaphore:
|
|
410
|
+
result = await self.ainvoke(request, config=run_config)
|
|
411
|
+
return index, result
|
|
412
|
+
except Exception as exc:
|
|
413
|
+
if return_exceptions:
|
|
414
|
+
return index, exc
|
|
415
|
+
raise
|
|
416
|
+
|
|
417
|
+
tasks = [
|
|
418
|
+
asyncio.create_task(invoke_indexed(index, request))
|
|
419
|
+
for index, request in enumerate(requests)
|
|
420
|
+
]
|
|
421
|
+
try:
|
|
422
|
+
for completed in asyncio.as_completed(tasks):
|
|
423
|
+
yield await completed
|
|
424
|
+
finally:
|
|
425
|
+
for task in tasks:
|
|
426
|
+
if not task.done():
|
|
427
|
+
task.cancel()
|
|
428
|
+
await asyncio.gather(*tasks, return_exceptions=True)
|
|
429
|
+
|
|
430
|
+
def batch_as_completed(
|
|
431
|
+
self,
|
|
432
|
+
requests: Sequence[PlanRequest | str],
|
|
433
|
+
*,
|
|
434
|
+
config: RunConfig | dict[str, Any] | None = None,
|
|
435
|
+
return_exceptions: bool = False,
|
|
436
|
+
) -> Iterator[tuple[int, list[ToolResult] | Exception]]:
|
|
437
|
+
return _stream_sync(
|
|
438
|
+
lambda: self.abatch_as_completed(
|
|
439
|
+
requests,
|
|
440
|
+
config=config,
|
|
441
|
+
return_exceptions=return_exceptions,
|
|
442
|
+
)
|
|
443
|
+
)
|
|
444
|
+
|
|
445
|
+
async def astream(
|
|
446
|
+
self,
|
|
447
|
+
request: PlanRequest | str,
|
|
448
|
+
*,
|
|
449
|
+
config: RunConfig | dict[str, Any] | None = None,
|
|
450
|
+
) -> AsyncIterator[ToolResult]:
|
|
451
|
+
run_config = _coerce_config(config)
|
|
452
|
+
plan = await self.aplan(request)
|
|
453
|
+
async for result in self.executor.execute_iter(
|
|
454
|
+
plan,
|
|
455
|
+
retry=run_config.retry,
|
|
456
|
+
):
|
|
457
|
+
yield result
|
|
458
|
+
|
|
459
|
+
def stream(
|
|
460
|
+
self,
|
|
461
|
+
request: PlanRequest | str,
|
|
462
|
+
*,
|
|
463
|
+
config: RunConfig | dict[str, Any] | None = None,
|
|
464
|
+
) -> Iterator[ToolResult]:
|
|
465
|
+
return _stream_sync(lambda: self.astream(request, config=config))
|
|
466
|
+
|
|
467
|
+
async def astream_events(
|
|
468
|
+
self,
|
|
469
|
+
request: PlanRequest | str,
|
|
470
|
+
*,
|
|
471
|
+
config: RunConfig | dict[str, Any] | None = None,
|
|
472
|
+
) -> AsyncIterator[RunEvent]:
|
|
473
|
+
run_config = _coerce_config(config)
|
|
474
|
+
run_id = uuid4().hex
|
|
475
|
+
sequence = 0
|
|
476
|
+
|
|
477
|
+
request_payload: dict[str, Any] = {
|
|
478
|
+
"input_type": type(request).__name__,
|
|
479
|
+
}
|
|
480
|
+
if run_config.include_payloads:
|
|
481
|
+
request_payload["input"] = (
|
|
482
|
+
request.model_dump(mode="json")
|
|
483
|
+
if isinstance(request, PlanRequest)
|
|
484
|
+
else request
|
|
485
|
+
)
|
|
486
|
+
|
|
487
|
+
yield RunEvent.create(
|
|
488
|
+
event="run.start",
|
|
489
|
+
run_id=run_id,
|
|
490
|
+
sequence=sequence,
|
|
491
|
+
config=run_config,
|
|
492
|
+
data=request_payload,
|
|
493
|
+
)
|
|
494
|
+
sequence += 1
|
|
495
|
+
|
|
496
|
+
try:
|
|
497
|
+
plan = await self.aplan(request)
|
|
498
|
+
except Exception as exc:
|
|
499
|
+
data = {"error_type": type(exc).__name__, "stage": "planning"}
|
|
500
|
+
if run_config.include_payloads:
|
|
501
|
+
data["message"] = str(exc)
|
|
502
|
+
yield RunEvent.create(
|
|
503
|
+
event="run.error",
|
|
504
|
+
run_id=run_id,
|
|
505
|
+
sequence=sequence,
|
|
506
|
+
config=run_config,
|
|
507
|
+
data=data,
|
|
508
|
+
)
|
|
509
|
+
raise
|
|
510
|
+
|
|
511
|
+
plan_data: dict[str, Any] = {
|
|
512
|
+
"call_count": len(plan.calls),
|
|
513
|
+
"warnings": list(plan.warnings),
|
|
514
|
+
}
|
|
515
|
+
if run_config.include_payloads:
|
|
516
|
+
plan_data["plan"] = plan.model_dump(mode="json")
|
|
517
|
+
yield RunEvent.create(
|
|
518
|
+
event="plan.end",
|
|
519
|
+
run_id=run_id,
|
|
520
|
+
sequence=sequence,
|
|
521
|
+
config=run_config,
|
|
522
|
+
data=plan_data,
|
|
523
|
+
)
|
|
524
|
+
sequence += 1
|
|
525
|
+
|
|
526
|
+
result_count = 0
|
|
527
|
+
for call in plan.calls:
|
|
528
|
+
start_data: dict[str, Any] = {
|
|
529
|
+
"argument_names": sorted(call.arguments),
|
|
530
|
+
"fields": list(call.fields),
|
|
531
|
+
}
|
|
532
|
+
if run_config.include_payloads:
|
|
533
|
+
start_data["arguments"] = dict(call.arguments)
|
|
534
|
+
yield RunEvent.create(
|
|
535
|
+
event="tool.start",
|
|
536
|
+
run_id=run_id,
|
|
537
|
+
sequence=sequence,
|
|
538
|
+
config=run_config,
|
|
539
|
+
tool=call.tool,
|
|
540
|
+
endpoint=call.endpoint,
|
|
541
|
+
data=start_data,
|
|
542
|
+
)
|
|
543
|
+
sequence += 1
|
|
544
|
+
|
|
545
|
+
try:
|
|
546
|
+
result = await self.executor.execute_call(
|
|
547
|
+
call,
|
|
548
|
+
retry=run_config.retry,
|
|
549
|
+
)
|
|
550
|
+
except Exception as exc:
|
|
551
|
+
error_data = {"error_type": type(exc).__name__}
|
|
552
|
+
if run_config.include_payloads:
|
|
553
|
+
error_data["message"] = str(exc)
|
|
554
|
+
yield RunEvent.create(
|
|
555
|
+
event="tool.error",
|
|
556
|
+
run_id=run_id,
|
|
557
|
+
sequence=sequence,
|
|
558
|
+
config=run_config,
|
|
559
|
+
tool=call.tool,
|
|
560
|
+
endpoint=call.endpoint,
|
|
561
|
+
data=error_data,
|
|
562
|
+
)
|
|
563
|
+
sequence += 1
|
|
564
|
+
yield RunEvent.create(
|
|
565
|
+
event="run.error",
|
|
566
|
+
run_id=run_id,
|
|
567
|
+
sequence=sequence,
|
|
568
|
+
config=run_config,
|
|
569
|
+
data={
|
|
570
|
+
"error_type": type(exc).__name__,
|
|
571
|
+
"stage": "execution",
|
|
572
|
+
},
|
|
573
|
+
)
|
|
574
|
+
raise
|
|
575
|
+
|
|
576
|
+
end_data: dict[str, Any] = {
|
|
577
|
+
"projected_fields": list(result.projected_fields),
|
|
578
|
+
}
|
|
579
|
+
if run_config.include_payloads:
|
|
580
|
+
end_data["result"] = result.model_dump(mode="json")
|
|
581
|
+
yield RunEvent.create(
|
|
582
|
+
event="tool.end",
|
|
583
|
+
run_id=run_id,
|
|
584
|
+
sequence=sequence,
|
|
585
|
+
config=run_config,
|
|
586
|
+
tool=result.tool,
|
|
587
|
+
endpoint=result.endpoint,
|
|
588
|
+
data=end_data,
|
|
589
|
+
)
|
|
590
|
+
sequence += 1
|
|
591
|
+
result_count += 1
|
|
592
|
+
|
|
593
|
+
yield RunEvent.create(
|
|
594
|
+
event="run.end",
|
|
595
|
+
run_id=run_id,
|
|
596
|
+
sequence=sequence,
|
|
597
|
+
config=run_config,
|
|
598
|
+
data={"result_count": result_count},
|
|
599
|
+
)
|
|
600
|
+
|
|
601
|
+
async def run(
|
|
602
|
+
self,
|
|
603
|
+
request: PlanRequest | str,
|
|
604
|
+
*,
|
|
605
|
+
config: RunConfig | dict[str, Any] | None = None,
|
|
606
|
+
) -> list[ToolResult]:
|
|
607
|
+
return await self.ainvoke(request, config=config)
|
|
608
|
+
|
|
609
|
+
async def arun(
|
|
610
|
+
self,
|
|
611
|
+
request: PlanRequest | str,
|
|
612
|
+
*,
|
|
613
|
+
config: RunConfig | dict[str, Any] | None = None,
|
|
614
|
+
) -> list[ToolResult]:
|
|
615
|
+
return await self.ainvoke(request, config=config)
|
|
616
|
+
|
|
617
|
+
|
|
618
|
+
class ConfiguredSchemaRouter:
|
|
619
|
+
"""SchemaRouter with an immutable default RunConfig."""
|
|
620
|
+
|
|
621
|
+
def __init__(self, router: SchemaRouter, config: RunConfig) -> None:
|
|
622
|
+
self.router = router
|
|
623
|
+
self.config = config
|
|
624
|
+
|
|
625
|
+
def invoke(self, request: PlanRequest | str) -> list[ToolResult]:
|
|
626
|
+
return self.router.invoke(request, config=self.config)
|
|
627
|
+
|
|
628
|
+
async def ainvoke(self, request: PlanRequest | str) -> list[ToolResult]:
|
|
629
|
+
return await self.router.ainvoke(request, config=self.config)
|
|
630
|
+
|
|
631
|
+
def batch(
|
|
632
|
+
self,
|
|
633
|
+
requests: Sequence[PlanRequest | str],
|
|
634
|
+
*,
|
|
635
|
+
return_exceptions: bool = False,
|
|
636
|
+
) -> list[list[ToolResult] | Exception]:
|
|
637
|
+
return self.router.batch(
|
|
638
|
+
requests,
|
|
639
|
+
config=self.config,
|
|
640
|
+
return_exceptions=return_exceptions,
|
|
641
|
+
)
|
|
642
|
+
|
|
643
|
+
async def abatch(
|
|
644
|
+
self,
|
|
645
|
+
requests: Sequence[PlanRequest | str],
|
|
646
|
+
*,
|
|
647
|
+
return_exceptions: bool = False,
|
|
648
|
+
) -> list[list[ToolResult] | Exception]:
|
|
649
|
+
return await self.router.abatch(
|
|
650
|
+
requests,
|
|
651
|
+
config=self.config,
|
|
652
|
+
return_exceptions=return_exceptions,
|
|
653
|
+
)
|
|
654
|
+
|
|
655
|
+
def batch_as_completed(
|
|
656
|
+
self,
|
|
657
|
+
requests: Sequence[PlanRequest | str],
|
|
658
|
+
*,
|
|
659
|
+
return_exceptions: bool = False,
|
|
660
|
+
) -> Iterator[tuple[int, list[ToolResult] | Exception]]:
|
|
661
|
+
return self.router.batch_as_completed(
|
|
662
|
+
requests,
|
|
663
|
+
config=self.config,
|
|
664
|
+
return_exceptions=return_exceptions,
|
|
665
|
+
)
|
|
666
|
+
|
|
667
|
+
def abatch_as_completed(
|
|
668
|
+
self,
|
|
669
|
+
requests: Sequence[PlanRequest | str],
|
|
670
|
+
*,
|
|
671
|
+
return_exceptions: bool = False,
|
|
672
|
+
) -> AsyncIterator[tuple[int, list[ToolResult] | Exception]]:
|
|
673
|
+
return self.router.abatch_as_completed(
|
|
674
|
+
requests,
|
|
675
|
+
config=self.config,
|
|
676
|
+
return_exceptions=return_exceptions,
|
|
677
|
+
)
|
|
678
|
+
|
|
679
|
+
def stream(self, request: PlanRequest | str) -> Iterator[ToolResult]:
|
|
680
|
+
return self.router.stream(request, config=self.config)
|
|
681
|
+
|
|
682
|
+
def astream(self, request: PlanRequest | str) -> AsyncIterator[ToolResult]:
|
|
683
|
+
return self.router.astream(request, config=self.config)
|
|
684
|
+
|
|
685
|
+
def astream_events(self, request: PlanRequest | str) -> AsyncIterator[RunEvent]:
|
|
686
|
+
return self.router.astream_events(request, config=self.config)
|