ai 0.2.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.
- ai/__init__.py +149 -0
- ai/_modelsdev.py +89 -0
- ai/agents/__init__.py +66 -0
- ai/agents/_middleware.py +363 -0
- ai/agents/agent.py +1225 -0
- ai/agents/hooks.py +252 -0
- ai/agents/mcp/__init__.py +7 -0
- ai/agents/mcp/client.py +293 -0
- ai/agents/runtime.py +97 -0
- ai/agents/ui/__init__.py +5 -0
- ai/agents/ui/ai_sdk/__init__.py +23 -0
- ai/agents/ui/ai_sdk/_approvals.py +33 -0
- ai/agents/ui/ai_sdk/_parts.py +144 -0
- ai/agents/ui/ai_sdk/inbound.py +510 -0
- ai/agents/ui/ai_sdk/outbound/__init__.py +7 -0
- ai/agents/ui/ai_sdk/outbound/_state.py +340 -0
- ai/agents/ui/ai_sdk/outbound/history.py +66 -0
- ai/agents/ui/ai_sdk/outbound/sse.py +54 -0
- ai/agents/ui/ai_sdk/outbound/stream.py +38 -0
- ai/agents/ui/ai_sdk/protocol.py +331 -0
- ai/agents/ui/ai_sdk/ui_message.py +236 -0
- ai/errors.py +337 -0
- ai/models/__init__.py +68 -0
- ai/models/core/__init__.py +36 -0
- ai/models/core/api.py +475 -0
- ai/models/core/helpers/__init__.py +3 -0
- ai/models/core/helpers/files.py +101 -0
- ai/models/core/model.py +92 -0
- ai/models/core/params.py +42 -0
- ai/providers/__init__.py +14 -0
- ai/providers/_optional.py +23 -0
- ai/providers/ai_gateway/__init__.py +42 -0
- ai/providers/ai_gateway/client/__init__.py +6 -0
- ai/providers/ai_gateway/client/_client.py +257 -0
- ai/providers/ai_gateway/client/errors.py +322 -0
- ai/providers/ai_gateway/errors.py +55 -0
- ai/providers/ai_gateway/protocol.py +637 -0
- ai/providers/ai_gateway/provider.py +162 -0
- ai/providers/ai_gateway/tools.py +128 -0
- ai/providers/anthropic/__init__.py +27 -0
- ai/providers/anthropic/_sdk.py +31 -0
- ai/providers/anthropic/errors.py +153 -0
- ai/providers/anthropic/protocol.py +585 -0
- ai/providers/anthropic/provider.py +217 -0
- ai/providers/anthropic/tools.py +240 -0
- ai/providers/base.py +289 -0
- ai/providers/openai/__init__.py +19 -0
- ai/providers/openai/_sdk.py +43 -0
- ai/providers/openai/errors.py +151 -0
- ai/providers/openai/protocol.py +378 -0
- ai/providers/openai/provider.py +215 -0
- ai/providers/openai/tools.py +341 -0
- ai/py.typed +0 -0
- ai/types/__init__.py +10 -0
- ai/types/builders.py +214 -0
- ai/types/events.py +346 -0
- ai/types/integrity.py +263 -0
- ai/types/media.py +278 -0
- ai/types/messages.py +292 -0
- ai/types/proto.py +3 -0
- ai/types/tools.py +39 -0
- ai/types/usage.py +57 -0
- ai/util.py +204 -0
- ai-0.2.0.dist-info/METADATA +197 -0
- ai-0.2.0.dist-info/RECORD +67 -0
- ai-0.2.0.dist-info/WHEEL +4 -0
- ai-0.2.0.dist-info/licenses/LICENSE +13 -0
ai/__init__.py
ADDED
|
@@ -0,0 +1,149 @@
|
|
|
1
|
+
from . import errors, models, providers, util
|
|
2
|
+
from .agents import (
|
|
3
|
+
Agent,
|
|
4
|
+
AgentTool,
|
|
5
|
+
Context,
|
|
6
|
+
StreamingStatusTool,
|
|
7
|
+
StreamingTextTool,
|
|
8
|
+
SubAgentTool,
|
|
9
|
+
Tool,
|
|
10
|
+
ToolCall,
|
|
11
|
+
ToolRunner,
|
|
12
|
+
abort_pending_hook,
|
|
13
|
+
agent,
|
|
14
|
+
cancel_hook,
|
|
15
|
+
hook,
|
|
16
|
+
mcp,
|
|
17
|
+
pending_tool_result,
|
|
18
|
+
resolve_hook,
|
|
19
|
+
tool,
|
|
20
|
+
tool_result,
|
|
21
|
+
yield_from,
|
|
22
|
+
)
|
|
23
|
+
from .errors import (
|
|
24
|
+
AIError,
|
|
25
|
+
ConfigurationError,
|
|
26
|
+
HTTPErrorContext,
|
|
27
|
+
InstallationError,
|
|
28
|
+
ProviderAPIError,
|
|
29
|
+
ProviderAuthenticationError,
|
|
30
|
+
ProviderBadRequestError,
|
|
31
|
+
ProviderConflictError,
|
|
32
|
+
ProviderConnectionError,
|
|
33
|
+
ProviderDeadlineExceededError,
|
|
34
|
+
ProviderError,
|
|
35
|
+
ProviderInternalServerError,
|
|
36
|
+
ProviderModelNotFoundError,
|
|
37
|
+
ProviderNotConfiguredError,
|
|
38
|
+
ProviderNotFoundError,
|
|
39
|
+
ProviderOverloadedError,
|
|
40
|
+
ProviderPermissionDeniedError,
|
|
41
|
+
ProviderRateLimitError,
|
|
42
|
+
ProviderRequestTooLargeError,
|
|
43
|
+
ProviderResponseError,
|
|
44
|
+
ProviderServiceUnavailableError,
|
|
45
|
+
ProviderStatusError,
|
|
46
|
+
ProviderTimeoutError,
|
|
47
|
+
ProviderUnprocessableEntityError,
|
|
48
|
+
UnsupportedProviderError,
|
|
49
|
+
)
|
|
50
|
+
from .models import (
|
|
51
|
+
ImageParams,
|
|
52
|
+
Model,
|
|
53
|
+
Provider,
|
|
54
|
+
Stream,
|
|
55
|
+
VideoParams,
|
|
56
|
+
generate,
|
|
57
|
+
get_model,
|
|
58
|
+
probe,
|
|
59
|
+
stream,
|
|
60
|
+
)
|
|
61
|
+
from .providers import get_provider
|
|
62
|
+
from .types import events, messages, tools
|
|
63
|
+
from .types.builders import (
|
|
64
|
+
assistant_message,
|
|
65
|
+
file_part,
|
|
66
|
+
system_message,
|
|
67
|
+
thinking,
|
|
68
|
+
tool_message,
|
|
69
|
+
tool_result_part,
|
|
70
|
+
user_message,
|
|
71
|
+
)
|
|
72
|
+
|
|
73
|
+
__all__ = [
|
|
74
|
+
# Builders (from types/builders)
|
|
75
|
+
"user_message",
|
|
76
|
+
"assistant_message",
|
|
77
|
+
"system_message",
|
|
78
|
+
"tool_message",
|
|
79
|
+
"tool_result",
|
|
80
|
+
"tool_result_part",
|
|
81
|
+
"pending_tool_result",
|
|
82
|
+
"file_part",
|
|
83
|
+
"thinking",
|
|
84
|
+
# Models (from models/)
|
|
85
|
+
"AIError",
|
|
86
|
+
"ConfigurationError",
|
|
87
|
+
"HTTPErrorContext",
|
|
88
|
+
"InstallationError",
|
|
89
|
+
"ProviderAPIError",
|
|
90
|
+
"ProviderAuthenticationError",
|
|
91
|
+
"ProviderBadRequestError",
|
|
92
|
+
"ProviderConflictError",
|
|
93
|
+
"ProviderConnectionError",
|
|
94
|
+
"ProviderDeadlineExceededError",
|
|
95
|
+
"ProviderError",
|
|
96
|
+
"ProviderInternalServerError",
|
|
97
|
+
"ProviderModelNotFoundError",
|
|
98
|
+
"ProviderNotConfiguredError",
|
|
99
|
+
"ProviderNotFoundError",
|
|
100
|
+
"ProviderOverloadedError",
|
|
101
|
+
"ProviderPermissionDeniedError",
|
|
102
|
+
"ProviderRateLimitError",
|
|
103
|
+
"ProviderRequestTooLargeError",
|
|
104
|
+
"ProviderResponseError",
|
|
105
|
+
"ProviderServiceUnavailableError",
|
|
106
|
+
"ProviderStatusError",
|
|
107
|
+
"ProviderTimeoutError",
|
|
108
|
+
"ProviderUnprocessableEntityError",
|
|
109
|
+
"UnsupportedProviderError",
|
|
110
|
+
"Model",
|
|
111
|
+
"Provider",
|
|
112
|
+
"ImageParams",
|
|
113
|
+
"VideoParams",
|
|
114
|
+
"Stream",
|
|
115
|
+
"stream",
|
|
116
|
+
"generate",
|
|
117
|
+
"get_model",
|
|
118
|
+
"probe",
|
|
119
|
+
"get_provider",
|
|
120
|
+
"models",
|
|
121
|
+
"providers",
|
|
122
|
+
# Agents — primary API
|
|
123
|
+
"Agent",
|
|
124
|
+
"agent",
|
|
125
|
+
"Context",
|
|
126
|
+
# Agents — tools
|
|
127
|
+
"AgentTool",
|
|
128
|
+
"Tool",
|
|
129
|
+
"ToolCall",
|
|
130
|
+
"ToolRunner",
|
|
131
|
+
"tool",
|
|
132
|
+
"StreamingTextTool",
|
|
133
|
+
"SubAgentTool",
|
|
134
|
+
"StreamingStatusTool",
|
|
135
|
+
# Agents — composition
|
|
136
|
+
"yield_from",
|
|
137
|
+
# Agents — hooks
|
|
138
|
+
"hook",
|
|
139
|
+
"resolve_hook",
|
|
140
|
+
"cancel_hook",
|
|
141
|
+
"abort_pending_hook",
|
|
142
|
+
# Submodules
|
|
143
|
+
"events",
|
|
144
|
+
"errors",
|
|
145
|
+
"messages",
|
|
146
|
+
"mcp",
|
|
147
|
+
"tools",
|
|
148
|
+
"util",
|
|
149
|
+
]
|
ai/_modelsdev.py
ADDED
|
@@ -0,0 +1,89 @@
|
|
|
1
|
+
"""Helpers for models.dev metadata."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import re
|
|
6
|
+
|
|
7
|
+
import modelsdotdev
|
|
8
|
+
|
|
9
|
+
_ENV_REFERENCE_RE = re.compile(r"\$\{?([A-Z_][A-Z0-9_]*)\}?")
|
|
10
|
+
_SECRET_ENV_MARKERS = ("API_KEY", "TOKEN", "SECRET", "BEARER")
|
|
11
|
+
_PROVIDER_ID_ALIASES = {"ai-gateway": "vercel", "gateway": "vercel"}
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def parse_model_id(model_id: str) -> modelsdotdev.ModelRef:
|
|
15
|
+
return modelsdotdev.parse_model_id(_canonical_model_id(model_id))
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def get_provider_by_id(provider_id: str) -> modelsdotdev.Provider | None:
|
|
19
|
+
return modelsdotdev.get_provider_by_id(_canonical_provider_id(provider_id))
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def get_model_by_id(model_id: str) -> modelsdotdev.Model | None:
|
|
23
|
+
return modelsdotdev.get_model_by_id(_canonical_model_id(model_id))
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def _canonical_provider_id(provider_id: str) -> str:
|
|
27
|
+
return _PROVIDER_ID_ALIASES.get(provider_id, provider_id)
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def _canonical_model_id(model_id: str) -> str:
|
|
31
|
+
for separator in (":", "/"):
|
|
32
|
+
prefix, sep, rest = model_id.partition(separator)
|
|
33
|
+
if sep and prefix in _PROVIDER_ID_ALIASES:
|
|
34
|
+
return f"{_PROVIDER_ID_ALIASES[prefix]}{separator}{rest}"
|
|
35
|
+
return model_id
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def provider_base_url(
|
|
39
|
+
provider: modelsdotdev.Provider,
|
|
40
|
+
model_provider_config: modelsdotdev.ModelProviderConfig | None = None,
|
|
41
|
+
) -> str | None:
|
|
42
|
+
if model_provider_config is not None and model_provider_config.api is not None:
|
|
43
|
+
return model_provider_config.api
|
|
44
|
+
return provider.api
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def provider_config(
|
|
48
|
+
provider: modelsdotdev.Provider,
|
|
49
|
+
model_provider_config: modelsdotdev.ModelProviderConfig | None = None,
|
|
50
|
+
) -> tuple[str | None, tuple[str, ...]]:
|
|
51
|
+
"""Return ``api_key_env`` and non-secret config envs from models.dev data."""
|
|
52
|
+
api = provider_base_url(provider, model_provider_config)
|
|
53
|
+
envs = _provider_envs(provider, api)
|
|
54
|
+
api_key_env = _api_key_env(envs, api)
|
|
55
|
+
config_envs = tuple(env for env in envs if env != api_key_env)
|
|
56
|
+
return api_key_env, config_envs
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def provider_npm(
|
|
60
|
+
provider: modelsdotdev.Provider,
|
|
61
|
+
model_provider_config: modelsdotdev.ModelProviderConfig | None = None,
|
|
62
|
+
) -> str:
|
|
63
|
+
if model_provider_config is not None and model_provider_config.npm is not None:
|
|
64
|
+
return model_provider_config.npm
|
|
65
|
+
return provider.npm
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def _provider_envs(provider: modelsdotdev.Provider, api: str | None) -> tuple[str, ...]:
|
|
69
|
+
envs = list(provider.env)
|
|
70
|
+
for env in _ENV_REFERENCE_RE.findall(api or ""):
|
|
71
|
+
if env not in envs:
|
|
72
|
+
envs.append(env)
|
|
73
|
+
return tuple(envs)
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def _api_key_env(envs: tuple[str, ...], api: str | None) -> str | None:
|
|
77
|
+
if not envs:
|
|
78
|
+
return None
|
|
79
|
+
|
|
80
|
+
referenced_envs = set(_ENV_REFERENCE_RE.findall(api or ""))
|
|
81
|
+
candidates = [env for env in envs if env not in referenced_envs]
|
|
82
|
+
if not candidates:
|
|
83
|
+
candidates = list(envs)
|
|
84
|
+
|
|
85
|
+
for marker in _SECRET_ENV_MARKERS:
|
|
86
|
+
for env in candidates:
|
|
87
|
+
if marker in env:
|
|
88
|
+
return env
|
|
89
|
+
return candidates[0]
|
ai/agents/__init__.py
ADDED
|
@@ -0,0 +1,66 @@
|
|
|
1
|
+
from . import mcp, ui
|
|
2
|
+
from .agent import (
|
|
3
|
+
Agent,
|
|
4
|
+
AgentTool,
|
|
5
|
+
Aggregate,
|
|
6
|
+
BoundToolCall,
|
|
7
|
+
ConcatAggregator,
|
|
8
|
+
Context,
|
|
9
|
+
GatedToolCall,
|
|
10
|
+
LastAggregator,
|
|
11
|
+
MessageAggregator,
|
|
12
|
+
MessageBundle,
|
|
13
|
+
SimpleAggregator,
|
|
14
|
+
StreamingStatusTool,
|
|
15
|
+
StreamingTextTool,
|
|
16
|
+
SubAgentTool,
|
|
17
|
+
Tool,
|
|
18
|
+
ToolCall,
|
|
19
|
+
ToolCallCallable,
|
|
20
|
+
ToolRunner,
|
|
21
|
+
agent,
|
|
22
|
+
pending_tool_result,
|
|
23
|
+
tool,
|
|
24
|
+
tool_result,
|
|
25
|
+
yield_from,
|
|
26
|
+
)
|
|
27
|
+
from .hooks import (
|
|
28
|
+
TOOL_APPROVAL_HOOK_TYPE,
|
|
29
|
+
abort_pending_hook,
|
|
30
|
+
cancel_hook,
|
|
31
|
+
hook,
|
|
32
|
+
resolve_hook,
|
|
33
|
+
)
|
|
34
|
+
|
|
35
|
+
__all__ = [
|
|
36
|
+
"Agent",
|
|
37
|
+
"AgentTool",
|
|
38
|
+
"Aggregate",
|
|
39
|
+
"ConcatAggregator",
|
|
40
|
+
"Context",
|
|
41
|
+
"LastAggregator",
|
|
42
|
+
"MessageAggregator",
|
|
43
|
+
"MessageBundle",
|
|
44
|
+
"SimpleAggregator",
|
|
45
|
+
"StreamingTextTool",
|
|
46
|
+
"SubAgentTool",
|
|
47
|
+
"BoundToolCall",
|
|
48
|
+
"GatedToolCall",
|
|
49
|
+
"Tool",
|
|
50
|
+
"ToolCall",
|
|
51
|
+
"ToolCallCallable",
|
|
52
|
+
"ToolRunner",
|
|
53
|
+
"StreamingStatusTool",
|
|
54
|
+
"TOOL_APPROVAL_HOOK_TYPE",
|
|
55
|
+
"abort_pending_hook",
|
|
56
|
+
"agent",
|
|
57
|
+
"cancel_hook",
|
|
58
|
+
"hook",
|
|
59
|
+
"mcp",
|
|
60
|
+
"pending_tool_result",
|
|
61
|
+
"resolve_hook",
|
|
62
|
+
"tool",
|
|
63
|
+
"tool_result",
|
|
64
|
+
"ui",
|
|
65
|
+
"yield_from",
|
|
66
|
+
]
|
ai/agents/_middleware.py
ADDED
|
@@ -0,0 +1,363 @@
|
|
|
1
|
+
"""Middleware: composable wrappers around all execution surfaces.
|
|
2
|
+
|
|
3
|
+
Middleware is run-scoped — pass it to :meth:`Agent.run`::
|
|
4
|
+
|
|
5
|
+
agent.run(model, messages, middleware=[LoggingMiddleware()])
|
|
6
|
+
|
|
7
|
+
Middleware wraps agent runs, model calls, generate calls, tool calls, and
|
|
8
|
+
hook calls. Subclass :class:`Middleware` and override the methods you care
|
|
9
|
+
about — unimplemented methods pass through to the next middleware (or the
|
|
10
|
+
real implementation).
|
|
11
|
+
|
|
12
|
+
Ordering: first in the list = outermost. ``[A(), B()]`` means A wraps B
|
|
13
|
+
wraps the real call. A sees the call first and the result last.
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
from __future__ import annotations
|
|
17
|
+
|
|
18
|
+
import contextvars
|
|
19
|
+
import dataclasses
|
|
20
|
+
from collections.abc import AsyncGenerator, Awaitable, Callable, Sequence
|
|
21
|
+
from typing import TYPE_CHECKING, Any
|
|
22
|
+
|
|
23
|
+
import pydantic
|
|
24
|
+
|
|
25
|
+
from ..types import messages as messages_
|
|
26
|
+
from ..types.tools import Tool
|
|
27
|
+
|
|
28
|
+
# Compat shim: ``StreamResultLike`` was removed from ``ai.types.proto`` when
|
|
29
|
+
# the model layer was reworked. Middleware is dead code under the new
|
|
30
|
+
# ``Executor``-based ``api.py`` and is kept around only so the agents
|
|
31
|
+
# rewrite can land separately; ``Any`` is enough to keep the existing
|
|
32
|
+
# annotations type-checking.
|
|
33
|
+
type StreamResultLike = Any
|
|
34
|
+
|
|
35
|
+
# ---------------------------------------------------------------------------
|
|
36
|
+
# Call context objects — frozen dataclasses with isolated mutable fields.
|
|
37
|
+
#
|
|
38
|
+
# Mutable container fields (``list``, ``dict``) are shallow-copied at
|
|
39
|
+
# construction via ``__post_init__`` so that middleware sees its own copy
|
|
40
|
+
# and cannot accidentally mutate the caller's data. To modify fields,
|
|
41
|
+
# use ``dataclasses.replace(call, messages=new_msgs)`` before passing
|
|
42
|
+
# to ``next``.
|
|
43
|
+
# ---------------------------------------------------------------------------
|
|
44
|
+
|
|
45
|
+
if TYPE_CHECKING:
|
|
46
|
+
from ..models.core.model import Model
|
|
47
|
+
from ..types import events as events_
|
|
48
|
+
from .agent import Context
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
@dataclasses.dataclass(frozen=True)
|
|
52
|
+
class ModelContext:
|
|
53
|
+
"""Context for a model streaming call."""
|
|
54
|
+
|
|
55
|
+
model: Model
|
|
56
|
+
messages: list[messages_.Message]
|
|
57
|
+
tools: Sequence[Tool] | None
|
|
58
|
+
output_type: type[pydantic.BaseModel] | None
|
|
59
|
+
kwargs: dict[str, Any]
|
|
60
|
+
|
|
61
|
+
def __post_init__(self) -> None:
|
|
62
|
+
object.__setattr__(self, "messages", list(self.messages))
|
|
63
|
+
if self.tools is not None:
|
|
64
|
+
object.__setattr__(self, "tools", list(self.tools))
|
|
65
|
+
object.__setattr__(self, "kwargs", dict(self.kwargs))
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
@dataclasses.dataclass(frozen=True)
|
|
69
|
+
class GenerateContext:
|
|
70
|
+
"""Context for a model generate call (images, video, etc.)."""
|
|
71
|
+
|
|
72
|
+
model: Model
|
|
73
|
+
messages: list[messages_.Message]
|
|
74
|
+
params: Any
|
|
75
|
+
|
|
76
|
+
def __post_init__(self) -> None:
|
|
77
|
+
object.__setattr__(self, "messages", list(self.messages))
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
@dataclasses.dataclass(frozen=True)
|
|
81
|
+
class ToolContext:
|
|
82
|
+
"""Context for a tool execution."""
|
|
83
|
+
|
|
84
|
+
tool_call_id: str
|
|
85
|
+
tool_name: str
|
|
86
|
+
kwargs: dict[str, Any]
|
|
87
|
+
|
|
88
|
+
def __post_init__(self) -> None:
|
|
89
|
+
object.__setattr__(self, "kwargs", dict(self.kwargs))
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
@dataclasses.dataclass(frozen=True)
|
|
93
|
+
class HookContext:
|
|
94
|
+
"""Context for a hook suspension point."""
|
|
95
|
+
|
|
96
|
+
label: str
|
|
97
|
+
payload: type[pydantic.BaseModel]
|
|
98
|
+
metadata: dict[str, Any]
|
|
99
|
+
|
|
100
|
+
def __post_init__(self) -> None:
|
|
101
|
+
object.__setattr__(self, "metadata", dict(self.metadata))
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
# ---------------------------------------------------------------------------
|
|
105
|
+
# Middleware base class — override the methods you care about.
|
|
106
|
+
# ---------------------------------------------------------------------------
|
|
107
|
+
|
|
108
|
+
# Event/message aliases for brevity in signatures. ``_Event`` is intentionally
|
|
109
|
+
# typed as ``Any`` so the agent-run chain accepts the wider ``AgentEvent``
|
|
110
|
+
# union (which includes ``ToolCallResult``/``HookEvent``) without a circular
|
|
111
|
+
# import from ``ai.agents``.
|
|
112
|
+
_Event = Any
|
|
113
|
+
_Message = messages_.Message
|
|
114
|
+
|
|
115
|
+
# Agent run next-function type: call -> async generator of events.
|
|
116
|
+
_AgentRunNext = Callable[["Context"], AsyncGenerator[_Event]]
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
class _Middleware:
|
|
120
|
+
"""Base middleware class. Override the methods you need.
|
|
121
|
+
|
|
122
|
+
Default implementations call ``next(call)`` — a transparent pass-through.
|
|
123
|
+
"""
|
|
124
|
+
|
|
125
|
+
async def wrap_agent_run(
|
|
126
|
+
self,
|
|
127
|
+
call: Context,
|
|
128
|
+
next: _AgentRunNext,
|
|
129
|
+
) -> AsyncGenerator[_Event]:
|
|
130
|
+
"""Wrap an agent run.
|
|
131
|
+
|
|
132
|
+
``next(call)`` returns an async generator of ``Event`` objects.
|
|
133
|
+
Override to add tracing, durability checkpoints, or other
|
|
134
|
+
run-scoped behavior::
|
|
135
|
+
|
|
136
|
+
async def wrap_agent_run(self, call, next):
|
|
137
|
+
span = start_span("agent.run")
|
|
138
|
+
async for event in next(call):
|
|
139
|
+
yield event
|
|
140
|
+
span.end()
|
|
141
|
+
"""
|
|
142
|
+
async for event in next(call):
|
|
143
|
+
yield event
|
|
144
|
+
|
|
145
|
+
async def wrap_model(
|
|
146
|
+
self,
|
|
147
|
+
call: ModelContext,
|
|
148
|
+
next: Callable[[ModelContext], Awaitable[StreamResultLike]],
|
|
149
|
+
) -> StreamResultLike:
|
|
150
|
+
"""Wrap a model streaming call.
|
|
151
|
+
|
|
152
|
+
``next(call)`` returns a :class:`~ai.types.StreamResultLike` that
|
|
153
|
+
is async-iterable over ``Event`` objects. You can do work
|
|
154
|
+
before, iterate / transform the stream, or do cleanup after.
|
|
155
|
+
|
|
156
|
+
To transform the stream, use
|
|
157
|
+
:meth:`~ai.models.StreamResult.from_generator`::
|
|
158
|
+
|
|
159
|
+
async def wrap_model(self, call, next):
|
|
160
|
+
stream = await next(call)
|
|
161
|
+
async def _add_suffix():
|
|
162
|
+
async for event in stream:
|
|
163
|
+
yield event
|
|
164
|
+
from ai.models import StreamResult
|
|
165
|
+
return StreamResult.from_generator(_add_suffix())
|
|
166
|
+
"""
|
|
167
|
+
return await next(call)
|
|
168
|
+
|
|
169
|
+
async def wrap_generate(
|
|
170
|
+
self,
|
|
171
|
+
call: GenerateContext,
|
|
172
|
+
next: Callable[[GenerateContext], Awaitable[_Message]],
|
|
173
|
+
) -> _Message:
|
|
174
|
+
"""Wrap a model generate call (images, video, etc.)."""
|
|
175
|
+
return await next(call)
|
|
176
|
+
|
|
177
|
+
async def wrap_tool(
|
|
178
|
+
self,
|
|
179
|
+
call: ToolContext,
|
|
180
|
+
next: Callable[[ToolContext], Awaitable[events_.ToolCallResult]],
|
|
181
|
+
) -> events_.ToolCallResult:
|
|
182
|
+
"""Wrap a tool execution.
|
|
183
|
+
|
|
184
|
+
``next(call)`` returns a :class:`ToolCallResult`.
|
|
185
|
+
"""
|
|
186
|
+
return await next(call)
|
|
187
|
+
|
|
188
|
+
async def wrap_hook(
|
|
189
|
+
self,
|
|
190
|
+
call: HookContext,
|
|
191
|
+
next: Callable[[HookContext], Awaitable[pydantic.BaseModel]],
|
|
192
|
+
) -> pydantic.BaseModel:
|
|
193
|
+
"""Wrap a hook suspension point.
|
|
194
|
+
|
|
195
|
+
``next(call)`` blocks until the hook is resolved and returns the
|
|
196
|
+
validated payload instance.
|
|
197
|
+
"""
|
|
198
|
+
return await next(call)
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
# ---------------------------------------------------------------------------
|
|
202
|
+
# Run-scoped middleware via ContextVar
|
|
203
|
+
# ---------------------------------------------------------------------------
|
|
204
|
+
|
|
205
|
+
_active: contextvars.ContextVar[list[_Middleware]] = contextvars.ContextVar(
|
|
206
|
+
"middleware",
|
|
207
|
+
)
|
|
208
|
+
|
|
209
|
+
_EMPTY: list[_Middleware] = []
|
|
210
|
+
|
|
211
|
+
|
|
212
|
+
def get() -> list[_Middleware]:
|
|
213
|
+
"""Return the middleware stack for the current run (empty if none)."""
|
|
214
|
+
return _active.get(_EMPTY)
|
|
215
|
+
|
|
216
|
+
|
|
217
|
+
Token = contextvars.Token[list[_Middleware]]
|
|
218
|
+
|
|
219
|
+
|
|
220
|
+
def activate(mw: list[_Middleware]) -> Token:
|
|
221
|
+
"""Set the middleware stack for the current run. Returns a token for reset."""
|
|
222
|
+
return _active.set(mw)
|
|
223
|
+
|
|
224
|
+
|
|
225
|
+
def deactivate(token: Token) -> None:
|
|
226
|
+
"""Restore the previous middleware stack."""
|
|
227
|
+
_active.reset(token)
|
|
228
|
+
|
|
229
|
+
|
|
230
|
+
# ---------------------------------------------------------------------------
|
|
231
|
+
# Chain builders — compose the middleware stack for each surface.
|
|
232
|
+
#
|
|
233
|
+
# Each builder takes the *real* implementation as a callable and returns
|
|
234
|
+
# a callable with the same signature that routes through middleware.
|
|
235
|
+
#
|
|
236
|
+
# When no middleware is active, the real implementation is returned
|
|
237
|
+
# directly — zero overhead.
|
|
238
|
+
# ---------------------------------------------------------------------------
|
|
239
|
+
|
|
240
|
+
|
|
241
|
+
def _build_model_chain(
|
|
242
|
+
real: Callable[[ModelContext], Awaitable[StreamResultLike]],
|
|
243
|
+
) -> Callable[[ModelContext], Awaitable[StreamResultLike]]:
|
|
244
|
+
mw = get()
|
|
245
|
+
if not mw:
|
|
246
|
+
return real
|
|
247
|
+
|
|
248
|
+
chain = real
|
|
249
|
+
for m in reversed(mw):
|
|
250
|
+
|
|
251
|
+
def _make(
|
|
252
|
+
m: _Middleware,
|
|
253
|
+
nxt: Callable[[ModelContext], Awaitable[StreamResultLike]],
|
|
254
|
+
) -> Callable[[ModelContext], Awaitable[StreamResultLike]]:
|
|
255
|
+
async def _wrapped(call: ModelContext) -> StreamResultLike:
|
|
256
|
+
return await m.wrap_model(call, nxt)
|
|
257
|
+
|
|
258
|
+
return _wrapped
|
|
259
|
+
|
|
260
|
+
chain = _make(m, chain)
|
|
261
|
+
return chain
|
|
262
|
+
|
|
263
|
+
|
|
264
|
+
def _build_generate_chain(
|
|
265
|
+
real: Callable[[GenerateContext], Awaitable[_Message]],
|
|
266
|
+
) -> Callable[[GenerateContext], Awaitable[_Message]]:
|
|
267
|
+
mw = get()
|
|
268
|
+
if not mw:
|
|
269
|
+
return real
|
|
270
|
+
|
|
271
|
+
chain = real
|
|
272
|
+
for m in reversed(mw):
|
|
273
|
+
|
|
274
|
+
def _make(
|
|
275
|
+
m: _Middleware, nxt: Callable[[GenerateContext], Awaitable[_Message]]
|
|
276
|
+
) -> Callable[[GenerateContext], Awaitable[_Message]]:
|
|
277
|
+
async def _wrapped(call: GenerateContext) -> _Message:
|
|
278
|
+
return await m.wrap_generate(call, nxt)
|
|
279
|
+
|
|
280
|
+
return _wrapped
|
|
281
|
+
|
|
282
|
+
chain = _make(m, chain)
|
|
283
|
+
return chain
|
|
284
|
+
|
|
285
|
+
|
|
286
|
+
def _build_tool_chain(
|
|
287
|
+
real: Callable[[ToolContext], Awaitable[events_.ToolCallResult]],
|
|
288
|
+
) -> Callable[[ToolContext], Awaitable[events_.ToolCallResult]]:
|
|
289
|
+
mw = get()
|
|
290
|
+
if not mw:
|
|
291
|
+
return real
|
|
292
|
+
|
|
293
|
+
chain = real
|
|
294
|
+
for m in reversed(mw):
|
|
295
|
+
|
|
296
|
+
def _make(
|
|
297
|
+
m: _Middleware,
|
|
298
|
+
nxt: Callable[[ToolContext], Awaitable[events_.ToolCallResult]],
|
|
299
|
+
) -> Callable[[ToolContext], Awaitable[events_.ToolCallResult]]:
|
|
300
|
+
async def _wrapped(call: ToolContext) -> events_.ToolCallResult:
|
|
301
|
+
return await m.wrap_tool(call, nxt)
|
|
302
|
+
|
|
303
|
+
return _wrapped
|
|
304
|
+
|
|
305
|
+
chain = _make(m, chain)
|
|
306
|
+
return chain
|
|
307
|
+
|
|
308
|
+
|
|
309
|
+
def _build_hook_chain(
|
|
310
|
+
real: Callable[[HookContext], Awaitable[pydantic.BaseModel]],
|
|
311
|
+
) -> Callable[[HookContext], Awaitable[pydantic.BaseModel]]:
|
|
312
|
+
mw = get()
|
|
313
|
+
if not mw:
|
|
314
|
+
return real
|
|
315
|
+
|
|
316
|
+
chain = real
|
|
317
|
+
for m in reversed(mw):
|
|
318
|
+
|
|
319
|
+
def _make(
|
|
320
|
+
m: _Middleware,
|
|
321
|
+
nxt: Callable[[HookContext], Awaitable[pydantic.BaseModel]],
|
|
322
|
+
) -> Callable[[HookContext], Awaitable[pydantic.BaseModel]]:
|
|
323
|
+
async def _wrapped(call: HookContext) -> pydantic.BaseModel:
|
|
324
|
+
return await m.wrap_hook(call, nxt)
|
|
325
|
+
|
|
326
|
+
return _wrapped
|
|
327
|
+
|
|
328
|
+
chain = _make(m, chain)
|
|
329
|
+
return chain
|
|
330
|
+
|
|
331
|
+
|
|
332
|
+
def _build_agent_run_chain(
|
|
333
|
+
real: _AgentRunNext,
|
|
334
|
+
) -> _AgentRunNext:
|
|
335
|
+
mw = get()
|
|
336
|
+
if not mw:
|
|
337
|
+
return real
|
|
338
|
+
|
|
339
|
+
chain = real
|
|
340
|
+
for m in reversed(mw):
|
|
341
|
+
|
|
342
|
+
def _make(m: _Middleware, nxt: _AgentRunNext) -> _AgentRunNext:
|
|
343
|
+
async def _wrapped(call: Context) -> AsyncGenerator[_Event]:
|
|
344
|
+
async for event in m.wrap_agent_run(call, nxt):
|
|
345
|
+
yield event
|
|
346
|
+
|
|
347
|
+
return _wrapped
|
|
348
|
+
|
|
349
|
+
chain = _make(m, chain)
|
|
350
|
+
return chain
|
|
351
|
+
|
|
352
|
+
|
|
353
|
+
__all__ = [
|
|
354
|
+
"GenerateContext",
|
|
355
|
+
"HookContext",
|
|
356
|
+
"ModelContext",
|
|
357
|
+
"StreamResultLike",
|
|
358
|
+
"ToolContext",
|
|
359
|
+
"_Middleware",
|
|
360
|
+
"activate",
|
|
361
|
+
"deactivate",
|
|
362
|
+
"get",
|
|
363
|
+
]
|