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.
@@ -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)