agentenv-framework-protocol 0.1.269__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.
- agentenv_framework_protocol-0.1.269.dist-info/METADATA +599 -0
- agentenv_framework_protocol-0.1.269.dist-info/RECORD +20 -0
- agentenv_framework_protocol-0.1.269.dist-info/WHEEL +4 -0
- agentenv_framework_protocol-0.1.269.dist-info/licenses/LICENSE +202 -0
- agentenv_framework_protocol-0.1.269.dist-info/licenses/NOTICE +4 -0
- agentenv_framework_protocol-0.1.269.dist-info/licenses/THIRD_PARTY_NOTICES.md +1701 -0
- agentenv_protocol/__init__.py +121 -0
- agentenv_protocol/a2a_agent/__init__.py +204 -0
- agentenv_protocol/a2a_agent/_triggers.py +489 -0
- agentenv_protocol/a2a_agent/extensions.py +1151 -0
- agentenv_protocol/a2a_agent/framework.py +1283 -0
- agentenv_protocol/a2a_agent/registry.py +449 -0
- agentenv_protocol/a2a_agent/tasks/__init__.py +39 -0
- agentenv_protocol/a2a_agent/tasks/v1.py +408 -0
- agentenv_protocol/agent_env_environment.py +653 -0
- agentenv_protocol/client.py +185 -0
- agentenv_protocol/manifest.py +203 -0
- agentenv_protocol/preflight.py +81 -0
- agentenv_protocol/transfers.py +554 -0
- agentenv_protocol/types.py +165 -0
|
@@ -0,0 +1,1283 @@
|
|
|
1
|
+
"""A2A agent application assembly built from declarations and handlers."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import inspect
|
|
7
|
+
import json
|
|
8
|
+
import logging
|
|
9
|
+
import os
|
|
10
|
+
import uuid
|
|
11
|
+
from collections import OrderedDict
|
|
12
|
+
from collections.abc import Callable, Iterable, Mapping
|
|
13
|
+
from contextlib import AbstractAsyncContextManager, asynccontextmanager
|
|
14
|
+
from dataclasses import dataclass, field
|
|
15
|
+
from pathlib import Path
|
|
16
|
+
from typing import Any, AsyncIterator, TypeVar, get_args, get_origin, get_type_hints
|
|
17
|
+
from urllib.parse import urlparse
|
|
18
|
+
|
|
19
|
+
from a2a.types import AgentCapabilities
|
|
20
|
+
from pydantic import BaseModel, ValidationError
|
|
21
|
+
from pydantic.json_schema import GenerateJsonSchema, JsonSchemaValue
|
|
22
|
+
from pydantic_core import core_schema
|
|
23
|
+
from starlette.applications import Starlette
|
|
24
|
+
from starlette.exceptions import HTTPException
|
|
25
|
+
from starlette.requests import Request
|
|
26
|
+
from starlette.responses import JSONResponse
|
|
27
|
+
from starlette.routing import Route
|
|
28
|
+
|
|
29
|
+
from ._triggers import TriggerEngine, TriggerError
|
|
30
|
+
from .extensions import (
|
|
31
|
+
AGENT_CONFIG_V1,
|
|
32
|
+
MCP_CONFIG_V1,
|
|
33
|
+
SKILL_CONFIG_V1,
|
|
34
|
+
TRAJECTORY_V1,
|
|
35
|
+
TRIGGERS_V1,
|
|
36
|
+
ExtensionActivation,
|
|
37
|
+
ExtensionDefinition,
|
|
38
|
+
ImplementationOwner,
|
|
39
|
+
McpAddRequest,
|
|
40
|
+
OperationReference,
|
|
41
|
+
TaskObjectTrajectoryRequest,
|
|
42
|
+
TaskTrajectoryRequest,
|
|
43
|
+
TriggerDecideRequest,
|
|
44
|
+
TriggerRegisterRequest,
|
|
45
|
+
enable,
|
|
46
|
+
)
|
|
47
|
+
from .registry import RegisteredOperation, build_registry
|
|
48
|
+
from .tasks.v1 import (
|
|
49
|
+
AgentConfig,
|
|
50
|
+
AgentRunResult,
|
|
51
|
+
TaskOutcome,
|
|
52
|
+
TaskProgress,
|
|
53
|
+
TaskRequest,
|
|
54
|
+
TaskResult,
|
|
55
|
+
)
|
|
56
|
+
from .tasks.v1 import (
|
|
57
|
+
DataPart as TaskDataPart,
|
|
58
|
+
)
|
|
59
|
+
from .tasks.v1 import (
|
|
60
|
+
FilePart as TaskFilePart,
|
|
61
|
+
)
|
|
62
|
+
from .tasks.v1 import (
|
|
63
|
+
TextPart as TaskTextPart,
|
|
64
|
+
)
|
|
65
|
+
from ..transfers import TransferError, upload
|
|
66
|
+
|
|
67
|
+
logger = logging.getLogger(__name__)
|
|
68
|
+
|
|
69
|
+
_AGENT_DEFINITION = "_agentenv_a2a_definition"
|
|
70
|
+
_AgentT = TypeVar("_AgentT", bound="AgentEnvAgent")
|
|
71
|
+
_MAX_CACHED_CONTEXTS = 1_024
|
|
72
|
+
_MAX_CACHED_TRAJECTORIES = 1_024
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def _config_wire_names(config: type[AgentConfig]) -> dict[str, str]:
|
|
76
|
+
"""Map model field names to their single top-level validation name."""
|
|
77
|
+
result: dict[str, str] = {}
|
|
78
|
+
for field_name, field_info in config.model_fields.items():
|
|
79
|
+
validation_alias = field_info.validation_alias
|
|
80
|
+
if validation_alias is None:
|
|
81
|
+
wire_name = field_name
|
|
82
|
+
elif isinstance(validation_alias, str):
|
|
83
|
+
wire_name = validation_alias
|
|
84
|
+
else:
|
|
85
|
+
raise TypeError(
|
|
86
|
+
f"AgentConfig field {field_name!r} must use a simple string alias"
|
|
87
|
+
)
|
|
88
|
+
if wire_name in result.values():
|
|
89
|
+
raise TypeError(f"AgentConfig fields have duplicate alias {wire_name!r}")
|
|
90
|
+
result[field_name] = wire_name
|
|
91
|
+
return result
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def _config_write_only_wire_names(config: type[AgentConfig]) -> set[str]:
|
|
95
|
+
"""Return wire names explicitly marked as write-only in the config schema."""
|
|
96
|
+
wire_names = _config_wire_names(config)
|
|
97
|
+
return {
|
|
98
|
+
wire_names[field_name]
|
|
99
|
+
for field_name, field_info in config.model_fields.items()
|
|
100
|
+
if isinstance(field_info.json_schema_extra, Mapping)
|
|
101
|
+
and field_info.json_schema_extra.get("writeOnly") is True
|
|
102
|
+
}
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
class _ConfigSchemaGenerator(GenerateJsonSchema):
|
|
106
|
+
"""Generate config schemas without exposing their runtime defaults."""
|
|
107
|
+
|
|
108
|
+
def default_schema(
|
|
109
|
+
self, schema: core_schema.WithDefaultSchema
|
|
110
|
+
) -> JsonSchemaValue:
|
|
111
|
+
return self.generate_inner(schema["schema"])
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
class _BoundedContextSessions:
|
|
115
|
+
"""Least-recently-used in-process cache of native runtime sessions."""
|
|
116
|
+
|
|
117
|
+
def __init__(self, max_entries: int = _MAX_CACHED_CONTEXTS) -> None:
|
|
118
|
+
if max_entries < 1:
|
|
119
|
+
raise ValueError("max_entries must be positive")
|
|
120
|
+
self._max_entries = max_entries
|
|
121
|
+
self._entries: OrderedDict[str, str] = OrderedDict()
|
|
122
|
+
|
|
123
|
+
def get(self, context_id: str) -> str | None:
|
|
124
|
+
try:
|
|
125
|
+
value = self._entries.pop(context_id)
|
|
126
|
+
except KeyError:
|
|
127
|
+
return None
|
|
128
|
+
self._entries[context_id] = value
|
|
129
|
+
return value
|
|
130
|
+
|
|
131
|
+
def set(self, context_id: str, session_ref: str) -> None:
|
|
132
|
+
self._entries.pop(context_id, None)
|
|
133
|
+
self._entries[context_id] = session_ref
|
|
134
|
+
while len(self._entries) > self._max_entries:
|
|
135
|
+
self._entries.popitem(last=False)
|
|
136
|
+
|
|
137
|
+
|
|
138
|
+
class _BoundedTaskTrajectories:
|
|
139
|
+
"""Least-recently-used cache of completed native task trajectories."""
|
|
140
|
+
|
|
141
|
+
def __init__(self, max_entries: int = _MAX_CACHED_TRAJECTORIES) -> None:
|
|
142
|
+
if max_entries < 1:
|
|
143
|
+
raise ValueError("max_entries must be positive")
|
|
144
|
+
self._max_entries = max_entries
|
|
145
|
+
self._entries: OrderedDict[str, Any] = OrderedDict()
|
|
146
|
+
|
|
147
|
+
def __setitem__(self, task_id: str, trajectory: Any) -> None:
|
|
148
|
+
self._entries.pop(task_id, None)
|
|
149
|
+
self._entries[task_id] = trajectory
|
|
150
|
+
while len(self._entries) > self._max_entries:
|
|
151
|
+
self._entries.popitem(last=False)
|
|
152
|
+
|
|
153
|
+
def get(self, task_id: str) -> Any | None:
|
|
154
|
+
try:
|
|
155
|
+
trajectory = self._entries.pop(task_id)
|
|
156
|
+
except KeyError:
|
|
157
|
+
return None
|
|
158
|
+
self._entries[task_id] = trajectory
|
|
159
|
+
return trajectory
|
|
160
|
+
|
|
161
|
+
def pop(self, task_id: str, default: Any = None) -> Any:
|
|
162
|
+
return self._entries.pop(task_id, default)
|
|
163
|
+
|
|
164
|
+
def __len__(self) -> int:
|
|
165
|
+
return len(self._entries)
|
|
166
|
+
|
|
167
|
+
|
|
168
|
+
@dataclass(slots=True)
|
|
169
|
+
class _ContextLockState:
|
|
170
|
+
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
|
|
171
|
+
users: int = 0
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
class _BoundedContextLocks:
|
|
175
|
+
"""Bound idle per-context locks without disrupting active or waiting tasks."""
|
|
176
|
+
|
|
177
|
+
def __init__(self, max_entries: int = _MAX_CACHED_CONTEXTS) -> None:
|
|
178
|
+
if max_entries < 1:
|
|
179
|
+
raise ValueError("max_entries must be positive")
|
|
180
|
+
self._max_entries = max_entries
|
|
181
|
+
self._entries: OrderedDict[str, _ContextLockState] = OrderedDict()
|
|
182
|
+
|
|
183
|
+
@asynccontextmanager
|
|
184
|
+
async def acquire(self, context_id: str) -> AsyncIterator[None]:
|
|
185
|
+
state = self._entries.get(context_id)
|
|
186
|
+
if state is None:
|
|
187
|
+
state = _ContextLockState()
|
|
188
|
+
self._entries[context_id] = state
|
|
189
|
+
else:
|
|
190
|
+
self._entries.move_to_end(context_id)
|
|
191
|
+
state.users += 1
|
|
192
|
+
try:
|
|
193
|
+
async with state.lock:
|
|
194
|
+
yield
|
|
195
|
+
finally:
|
|
196
|
+
state.users -= 1
|
|
197
|
+
self._entries.move_to_end(context_id)
|
|
198
|
+
self._trim_idle()
|
|
199
|
+
|
|
200
|
+
def _trim_idle(self) -> None:
|
|
201
|
+
while len(self._entries) > self._max_entries:
|
|
202
|
+
for context_id, state in tuple(self._entries.items()):
|
|
203
|
+
if state.users == 0:
|
|
204
|
+
del self._entries[context_id]
|
|
205
|
+
break
|
|
206
|
+
else:
|
|
207
|
+
# The temporary excess is proportional only to live work; it
|
|
208
|
+
# is removed as soon as one of those contexts becomes idle.
|
|
209
|
+
return
|
|
210
|
+
|
|
211
|
+
|
|
212
|
+
class AgentEnvAgent:
|
|
213
|
+
"""Typed base class for AgentEnv-compatible A2A agents."""
|
|
214
|
+
|
|
215
|
+
def run(self, request: TaskRequest[Any]) -> AgentRunResult:
|
|
216
|
+
"""Execute one normalized A2A task."""
|
|
217
|
+
raise NotImplementedError
|
|
218
|
+
|
|
219
|
+
def create_app(self) -> Starlette:
|
|
220
|
+
"""Build this agent's unserved Starlette application."""
|
|
221
|
+
return create_app(self)
|
|
222
|
+
|
|
223
|
+
@property
|
|
224
|
+
def default_handlers(self) -> DefaultExtensionHandlers:
|
|
225
|
+
"""Access default SDK handlers for composing extension overrides."""
|
|
226
|
+
application = getattr(self, "_agentenv_a2a_application", None)
|
|
227
|
+
if application is None:
|
|
228
|
+
raise RuntimeError("default handlers are available after create_app()")
|
|
229
|
+
return application.default_handlers
|
|
230
|
+
|
|
231
|
+
def session_ref_for_context(self, context_id: str) -> str | None:
|
|
232
|
+
"""Return the framework-managed native session reference for a context."""
|
|
233
|
+
application = getattr(self, "_agentenv_a2a_application", None)
|
|
234
|
+
if application is None:
|
|
235
|
+
raise RuntimeError("session references are available after create_app()")
|
|
236
|
+
return application.services.session_ref_for_context(context_id)
|
|
237
|
+
|
|
238
|
+
def set_session_ref_for_context(self, context_id: str, session_ref: str) -> None:
|
|
239
|
+
"""Associate a native session reference with an A2A context."""
|
|
240
|
+
if not context_id:
|
|
241
|
+
raise ValueError("context_id must not be empty")
|
|
242
|
+
if not session_ref:
|
|
243
|
+
raise ValueError("session_ref must not be empty")
|
|
244
|
+
application = getattr(self, "_agentenv_a2a_application", None)
|
|
245
|
+
if application is None:
|
|
246
|
+
raise RuntimeError("session references are available after create_app()")
|
|
247
|
+
application.services.set_session_ref_for_context(context_id, session_ref)
|
|
248
|
+
|
|
249
|
+
def serve(
|
|
250
|
+
self,
|
|
251
|
+
*,
|
|
252
|
+
host: str | None = None,
|
|
253
|
+
port: int | None = None,
|
|
254
|
+
**uvicorn_kwargs: Any,
|
|
255
|
+
) -> None:
|
|
256
|
+
"""Build and serve this agent with uvicorn."""
|
|
257
|
+
serve(self, host=host, port=port, **uvicorn_kwargs)
|
|
258
|
+
|
|
259
|
+
|
|
260
|
+
@dataclass(frozen=True, slots=True, kw_only=True)
|
|
261
|
+
class AgentIdentity:
|
|
262
|
+
name: str
|
|
263
|
+
description: str
|
|
264
|
+
version: str
|
|
265
|
+
input_modes: tuple[str, ...] = ("text",)
|
|
266
|
+
output_modes: tuple[str, ...] = ("text",)
|
|
267
|
+
skills: tuple[Any, ...] = ()
|
|
268
|
+
url: str = "/a2a"
|
|
269
|
+
|
|
270
|
+
def __post_init__(self) -> None:
|
|
271
|
+
object.__setattr__(self, "input_modes", tuple(self.input_modes))
|
|
272
|
+
object.__setattr__(self, "output_modes", tuple(self.output_modes))
|
|
273
|
+
object.__setattr__(self, "skills", tuple(self.skills))
|
|
274
|
+
|
|
275
|
+
|
|
276
|
+
def _rpc_path(card_url: str) -> str:
|
|
277
|
+
"""Resolve the local RPC route while preserving the card's public URL."""
|
|
278
|
+
parsed = urlparse(card_url)
|
|
279
|
+
if parsed.query or parsed.fragment:
|
|
280
|
+
raise ValueError("AgentIdentity.url must not contain a query or fragment")
|
|
281
|
+
if parsed.scheme or parsed.netloc:
|
|
282
|
+
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
|
|
283
|
+
raise ValueError("AgentIdentity.url must be an HTTP(S) URL or absolute path")
|
|
284
|
+
return parsed.path or "/"
|
|
285
|
+
if not parsed.path.startswith("/"):
|
|
286
|
+
raise ValueError("AgentIdentity.url path must start with '/'")
|
|
287
|
+
return parsed.path
|
|
288
|
+
|
|
289
|
+
|
|
290
|
+
@dataclass(frozen=True, slots=True)
|
|
291
|
+
class _AgentDefinition:
|
|
292
|
+
identity: AgentIdentity
|
|
293
|
+
extensions: tuple[ExtensionActivation, ...]
|
|
294
|
+
config: type[AgentConfig] | None = None
|
|
295
|
+
lifespan: Callable[[Starlette], AbstractAsyncContextManager] | None = None
|
|
296
|
+
workspace: Path | None = None
|
|
297
|
+
|
|
298
|
+
|
|
299
|
+
def a2a_agent(
|
|
300
|
+
*,
|
|
301
|
+
identity: AgentIdentity,
|
|
302
|
+
extensions: Iterable[ExtensionDefinition | ExtensionActivation] = (),
|
|
303
|
+
config: type[AgentConfig] | None = None,
|
|
304
|
+
config_description: str | None = None,
|
|
305
|
+
config_readback: bool = True,
|
|
306
|
+
lifespan: Callable[[Starlette], AbstractAsyncContextManager] | None = None,
|
|
307
|
+
workspace: str | Path | None = None,
|
|
308
|
+
) -> Callable[[type[_AgentT]], type[_AgentT]]:
|
|
309
|
+
"""Attach an A2A definition to an ``AgentEnvAgent`` subclass."""
|
|
310
|
+
|
|
311
|
+
declared_extensions = tuple(
|
|
312
|
+
ExtensionActivation(item) if isinstance(item, ExtensionDefinition) else item
|
|
313
|
+
for item in extensions
|
|
314
|
+
)
|
|
315
|
+
if not all(isinstance(item, ExtensionActivation) for item in declared_extensions):
|
|
316
|
+
raise TypeError(
|
|
317
|
+
"extensions must contain ExtensionDefinition or ExtensionActivation"
|
|
318
|
+
)
|
|
319
|
+
if not isinstance(config_readback, bool):
|
|
320
|
+
raise TypeError("config_readback must be a boolean")
|
|
321
|
+
has_explicit_agent_config = any(
|
|
322
|
+
activation.definition.uri == AGENT_CONFIG_V1.uri
|
|
323
|
+
for activation in declared_extensions
|
|
324
|
+
)
|
|
325
|
+
if config is None and config_description is not None:
|
|
326
|
+
raise TypeError("config_description requires config=")
|
|
327
|
+
if config is None and not config_readback:
|
|
328
|
+
raise TypeError("config_readback=False requires config=")
|
|
329
|
+
if config is None and has_explicit_agent_config:
|
|
330
|
+
raise TypeError(
|
|
331
|
+
"agent-config/v1 is derived from config=; replace the explicit "
|
|
332
|
+
"AGENT_CONFIG_V1 declaration with an AgentConfig subclass"
|
|
333
|
+
)
|
|
334
|
+
if config is not None:
|
|
335
|
+
if not isinstance(config, type) or not issubclass(config, AgentConfig):
|
|
336
|
+
raise TypeError("config must be an AgentConfig subclass")
|
|
337
|
+
if config.model_config.get("extra") != "forbid" or not config.model_config.get(
|
|
338
|
+
"frozen"
|
|
339
|
+
):
|
|
340
|
+
raise TypeError(
|
|
341
|
+
"AgentConfig subclasses must retain extra='forbid' and frozen=True"
|
|
342
|
+
)
|
|
343
|
+
if has_explicit_agent_config:
|
|
344
|
+
raise TypeError(
|
|
345
|
+
"config= automatically enables agent-config/v1; remove the explicit "
|
|
346
|
+
"AGENT_CONFIG_V1 declaration"
|
|
347
|
+
)
|
|
348
|
+
try:
|
|
349
|
+
default_config = config()
|
|
350
|
+
except ValidationError as exc:
|
|
351
|
+
raise TypeError(
|
|
352
|
+
"AgentConfig fields must have defaults so partial deployment-time "
|
|
353
|
+
"configuration updates can be validated"
|
|
354
|
+
) from exc
|
|
355
|
+
wire_names = _config_wire_names(config)
|
|
356
|
+
internal_defaults = default_config.model_dump(mode="python")
|
|
357
|
+
default_values = {
|
|
358
|
+
wire_names[name]: value for name, value in internal_defaults.items()
|
|
359
|
+
}
|
|
360
|
+
try:
|
|
361
|
+
config.model_validate(default_values, by_name=True)
|
|
362
|
+
except ValidationError as exc:
|
|
363
|
+
raise TypeError(
|
|
364
|
+
"AgentConfig serializers must return values accepted by their "
|
|
365
|
+
"declared field types"
|
|
366
|
+
) from exc
|
|
367
|
+
config_activation = enable(
|
|
368
|
+
AGENT_CONFIG_V1,
|
|
369
|
+
description=config_description,
|
|
370
|
+
fields=tuple(wire_names.values()),
|
|
371
|
+
defaults=default_values,
|
|
372
|
+
schema=config.model_json_schema(
|
|
373
|
+
mode="validation", schema_generator=_ConfigSchemaGenerator
|
|
374
|
+
),
|
|
375
|
+
readback=config_readback,
|
|
376
|
+
)
|
|
377
|
+
declared_extensions = (config_activation, *declared_extensions)
|
|
378
|
+
|
|
379
|
+
definition = _AgentDefinition(
|
|
380
|
+
identity=identity,
|
|
381
|
+
extensions=declared_extensions,
|
|
382
|
+
config=config,
|
|
383
|
+
lifespan=lifespan,
|
|
384
|
+
workspace=Path(workspace) if workspace is not None else None,
|
|
385
|
+
)
|
|
386
|
+
|
|
387
|
+
def decorate(cls: type[_AgentT]) -> type[_AgentT]:
|
|
388
|
+
if not issubclass(cls, AgentEnvAgent):
|
|
389
|
+
raise TypeError("@a2a_agent requires an AgentEnvAgent subclass")
|
|
390
|
+
setattr(cls, _AGENT_DEFINITION, definition)
|
|
391
|
+
return cls
|
|
392
|
+
|
|
393
|
+
return decorate
|
|
394
|
+
|
|
395
|
+
|
|
396
|
+
def create_app(agent: AgentEnvAgent) -> Starlette:
|
|
397
|
+
"""Build an unserved Starlette application for a declared agent."""
|
|
398
|
+
application = getattr(agent, "_agentenv_a2a_application", None)
|
|
399
|
+
if application is not None:
|
|
400
|
+
return application.app
|
|
401
|
+
return A2AAgentApplication(agent).app
|
|
402
|
+
|
|
403
|
+
|
|
404
|
+
def serve(
|
|
405
|
+
agent: AgentEnvAgent,
|
|
406
|
+
*,
|
|
407
|
+
host: str | None = None,
|
|
408
|
+
port: int | None = None,
|
|
409
|
+
**uvicorn_kwargs: Any,
|
|
410
|
+
) -> None:
|
|
411
|
+
"""Build and serve a declared agent with uvicorn."""
|
|
412
|
+
try:
|
|
413
|
+
import uvicorn
|
|
414
|
+
except ImportError as exc: # pragma: no cover - installation error
|
|
415
|
+
raise ImportError(
|
|
416
|
+
"Serving A2A agents requires: pip install 'agentenv-framework-protocol[agent]'"
|
|
417
|
+
) from exc
|
|
418
|
+
|
|
419
|
+
resolved_host = host or os.environ.get("A2A_HOST", "0.0.0.0")
|
|
420
|
+
resolved_port = (
|
|
421
|
+
port if port is not None else int(os.environ.get("A2A_PORT", "8000"))
|
|
422
|
+
)
|
|
423
|
+
uvicorn.run(
|
|
424
|
+
create_app(agent),
|
|
425
|
+
host=resolved_host,
|
|
426
|
+
port=resolved_port,
|
|
427
|
+
**uvicorn_kwargs,
|
|
428
|
+
)
|
|
429
|
+
|
|
430
|
+
|
|
431
|
+
def _thaw(value: Any) -> Any:
|
|
432
|
+
if isinstance(value, Mapping):
|
|
433
|
+
return {key: _thaw(item) for key, item in value.items()}
|
|
434
|
+
if isinstance(value, (tuple, list, set, frozenset)):
|
|
435
|
+
return [_thaw(item) for item in value]
|
|
436
|
+
return value
|
|
437
|
+
|
|
438
|
+
|
|
439
|
+
def _validation_error_detail(exc: ValidationError) -> str:
|
|
440
|
+
"""Return client-safe Pydantic errors without rejected values or docs URLs."""
|
|
441
|
+
return exc.json(include_input=False, include_url=False)
|
|
442
|
+
|
|
443
|
+
|
|
444
|
+
async def _invoke(handler: Callable, request: Any | None) -> Any:
|
|
445
|
+
arguments = () if request is None else (request,)
|
|
446
|
+
return await handler(*arguments)
|
|
447
|
+
|
|
448
|
+
|
|
449
|
+
def _validate_run_handler_signature(
|
|
450
|
+
handler: Callable, config_model: type[AgentConfig]
|
|
451
|
+
) -> None:
|
|
452
|
+
parameters = list(inspect.signature(handler).parameters.values())
|
|
453
|
+
if len(parameters) != 1 or parameters[0].kind not in (
|
|
454
|
+
inspect.Parameter.POSITIONAL_ONLY,
|
|
455
|
+
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
|
456
|
+
):
|
|
457
|
+
raise TypeError("agent.run must accept exactly one request argument")
|
|
458
|
+
if not (
|
|
459
|
+
inspect.iscoroutinefunction(handler) or inspect.isasyncgenfunction(handler)
|
|
460
|
+
):
|
|
461
|
+
raise TypeError("agent.run must be an async function or async generator")
|
|
462
|
+
parameter = parameters[0]
|
|
463
|
+
# Include parents so postponed local annotations like TaskRequest[Parent] resolve.
|
|
464
|
+
config_types = {
|
|
465
|
+
candidate.__name__: candidate
|
|
466
|
+
for candidate in config_model.__mro__
|
|
467
|
+
if isinstance(candidate, type) and issubclass(candidate, AgentConfig)
|
|
468
|
+
}
|
|
469
|
+
try:
|
|
470
|
+
annotation = get_type_hints(handler, localns=config_types).get(
|
|
471
|
+
parameter.name, parameter.annotation
|
|
472
|
+
)
|
|
473
|
+
except (NameError, TypeError):
|
|
474
|
+
annotation = parameter.annotation
|
|
475
|
+
if annotation is TaskRequest:
|
|
476
|
+
return
|
|
477
|
+
if get_origin(annotation) is not TaskRequest:
|
|
478
|
+
raise TypeError(
|
|
479
|
+
"agent.run request argument must be annotated as TaskRequest or "
|
|
480
|
+
f"TaskRequest[{config_model.__name__}]"
|
|
481
|
+
)
|
|
482
|
+
config_arguments = get_args(annotation)
|
|
483
|
+
accepted_config = config_arguments[0] if len(config_arguments) == 1 else None
|
|
484
|
+
if accepted_config is Any:
|
|
485
|
+
return
|
|
486
|
+
if not (
|
|
487
|
+
isinstance(accepted_config, type)
|
|
488
|
+
and issubclass(config_model, accepted_config)
|
|
489
|
+
and issubclass(accepted_config, AgentConfig)
|
|
490
|
+
):
|
|
491
|
+
raise TypeError(
|
|
492
|
+
"agent.run request argument must accept the configured "
|
|
493
|
+
f"{config_model.__name__} type"
|
|
494
|
+
)
|
|
495
|
+
|
|
496
|
+
|
|
497
|
+
def _unwrap_bound_handler(handler: Callable) -> Callable:
|
|
498
|
+
"""Unwrap a decorated method while preserving its original binding."""
|
|
499
|
+
resolved = inspect.unwrap(handler)
|
|
500
|
+
if inspect.ismethod(handler) and not inspect.ismethod(resolved):
|
|
501
|
+
owner = handler.__self__
|
|
502
|
+
owner_type = owner if isinstance(owner, type) else type(owner)
|
|
503
|
+
resolved = resolved.__get__(owner, owner_type)
|
|
504
|
+
return resolved
|
|
505
|
+
|
|
506
|
+
|
|
507
|
+
def _opaque_extension_error(stage: str) -> HTTPException:
|
|
508
|
+
"""Log an extension failure while keeping implementation details off the wire."""
|
|
509
|
+
correlation_id = uuid.uuid4().hex
|
|
510
|
+
logger.exception("extension %s failed (correlation_id=%s)", stage, correlation_id)
|
|
511
|
+
return HTTPException(
|
|
512
|
+
status_code=500,
|
|
513
|
+
detail=(
|
|
514
|
+
f"An unexpected framework error occurred. Correlation ID: {correlation_id}"
|
|
515
|
+
),
|
|
516
|
+
)
|
|
517
|
+
|
|
518
|
+
|
|
519
|
+
class _SdkServices:
|
|
520
|
+
def __init__(
|
|
521
|
+
self,
|
|
522
|
+
activations: Iterable[ExtensionActivation],
|
|
523
|
+
config_model: type[AgentConfig] | None,
|
|
524
|
+
) -> None:
|
|
525
|
+
self._activations = {
|
|
526
|
+
activation.definition.uri: activation for activation in activations
|
|
527
|
+
}
|
|
528
|
+
agent_config = self._activations.get(AGENT_CONFIG_V1.uri)
|
|
529
|
+
self.config_defaults: dict[str, Any] = (
|
|
530
|
+
dict(agent_config.options.get("defaults", {}))
|
|
531
|
+
if agent_config is not None
|
|
532
|
+
else {}
|
|
533
|
+
)
|
|
534
|
+
self.config: dict[str, Any] = {}
|
|
535
|
+
self.config_model = config_model or AgentConfig
|
|
536
|
+
self._write_only_config_fields = _config_write_only_wire_names(
|
|
537
|
+
self.config_model
|
|
538
|
+
)
|
|
539
|
+
self.mcp_servers: dict[str, dict[str, Any]] = {}
|
|
540
|
+
self.task_trajectories = _BoundedTaskTrajectories()
|
|
541
|
+
self._context_sessions = _BoundedContextSessions()
|
|
542
|
+
self.skills: list[Mapping[str, Any]] = []
|
|
543
|
+
self.identity_skills: dict[str, str] = {}
|
|
544
|
+
self._skill_registration_lock = asyncio.Lock()
|
|
545
|
+
self.card: Any = None
|
|
546
|
+
self.trigger_engine = TriggerEngine()
|
|
547
|
+
|
|
548
|
+
def task_config(self) -> AgentConfig:
|
|
549
|
+
return self.config_model.model_validate(
|
|
550
|
+
{**self.config_defaults, **self.config}, by_name=True
|
|
551
|
+
)
|
|
552
|
+
|
|
553
|
+
def session_ref_for_context(self, context_id: str) -> str | None:
|
|
554
|
+
return self._context_sessions.get(context_id)
|
|
555
|
+
|
|
556
|
+
def set_session_ref_for_context(self, context_id: str, session_ref: str) -> None:
|
|
557
|
+
self._context_sessions.set(context_id, session_ref)
|
|
558
|
+
|
|
559
|
+
def attach_card(self, card: Any) -> None:
|
|
560
|
+
self.card = card
|
|
561
|
+
self.identity_skills = {skill.name: skill.description for skill in card.skills}
|
|
562
|
+
|
|
563
|
+
def record_skill(self, registration: Mapping[str, Any]) -> None:
|
|
564
|
+
"""Record a successfully installed skill for subsequent task executions.
|
|
565
|
+
|
|
566
|
+
A bundle's read grants are secret and short-lived, so they are not kept."""
|
|
567
|
+
name = str(registration["name"])
|
|
568
|
+
self.skills.append(
|
|
569
|
+
{
|
|
570
|
+
key: registration[key]
|
|
571
|
+
for key in ("name", "description", "skill_md")
|
|
572
|
+
if key in registration
|
|
573
|
+
}
|
|
574
|
+
)
|
|
575
|
+
if self.card is not None:
|
|
576
|
+
from a2a.types import AgentSkill
|
|
577
|
+
|
|
578
|
+
self.card.skills.append(
|
|
579
|
+
AgentSkill(
|
|
580
|
+
id=f"skill-{name}",
|
|
581
|
+
name=name,
|
|
582
|
+
description=str(registration["description"]),
|
|
583
|
+
tags=["skill"],
|
|
584
|
+
)
|
|
585
|
+
)
|
|
586
|
+
|
|
587
|
+
def ensure_skill_is_new(self, registration: Mapping[str, Any]) -> None:
|
|
588
|
+
name = str(registration["name"])
|
|
589
|
+
if name in self.identity_skills or any(
|
|
590
|
+
skill["name"] == name for skill in self.skills
|
|
591
|
+
):
|
|
592
|
+
raise HTTPException(
|
|
593
|
+
status_code=409,
|
|
594
|
+
detail=f"Skill '{name}' is already registered",
|
|
595
|
+
)
|
|
596
|
+
|
|
597
|
+
@asynccontextmanager
|
|
598
|
+
async def skill_registration(
|
|
599
|
+
self, registration: Mapping[str, Any]
|
|
600
|
+
) -> AsyncIterator[None]:
|
|
601
|
+
"""Serialize duplicate checking, installation, and registration."""
|
|
602
|
+
async with self._skill_registration_lock:
|
|
603
|
+
self.ensure_skill_is_new(registration)
|
|
604
|
+
yield
|
|
605
|
+
self.record_skill(registration)
|
|
606
|
+
|
|
607
|
+
def handlers(self) -> dict[tuple[str, str], Callable]:
|
|
608
|
+
return {
|
|
609
|
+
(AGENT_CONFIG_V1.uri, "set"): self.agent_config_set,
|
|
610
|
+
(AGENT_CONFIG_V1.uri, "get"): self.agent_config_get,
|
|
611
|
+
(MCP_CONFIG_V1.uri, "add"): self.mcp_add,
|
|
612
|
+
(MCP_CONFIG_V1.uri, "list"): self.mcp_list,
|
|
613
|
+
(SKILL_CONFIG_V1.uri, "list"): self.skill_list,
|
|
614
|
+
(TRAJECTORY_V1.uri, "get"): self.trajectory_get,
|
|
615
|
+
(TRIGGERS_V1.uri, "register"): self.triggers_register,
|
|
616
|
+
(TRIGGERS_V1.uri, "decide"): self.triggers_decide,
|
|
617
|
+
(TRIGGERS_V1.uri, "state"): self.triggers_state,
|
|
618
|
+
}
|
|
619
|
+
|
|
620
|
+
async def agent_config_set(self, request: AgentConfig) -> dict[str, Any]:
|
|
621
|
+
normalized_request = request.model_dump(mode="python")
|
|
622
|
+
wire_names = _config_wire_names(type(request))
|
|
623
|
+
payload = {
|
|
624
|
+
wire_names[field_name]: normalized_request[field_name]
|
|
625
|
+
for field_name in request.model_fields_set
|
|
626
|
+
}
|
|
627
|
+
try:
|
|
628
|
+
validated = self.config_model.model_validate(
|
|
629
|
+
{**self.config_defaults, **self.config, **payload}, by_name=True
|
|
630
|
+
)
|
|
631
|
+
except ValidationError as exc:
|
|
632
|
+
raise HTTPException(
|
|
633
|
+
status_code=400, detail=_validation_error_detail(exc)
|
|
634
|
+
) from exc
|
|
635
|
+
normalized = validated.model_dump(mode="python")
|
|
636
|
+
field_for_wire = {
|
|
637
|
+
wire_name: field_name
|
|
638
|
+
for field_name, wire_name in _config_wire_names(self.config_model).items()
|
|
639
|
+
}
|
|
640
|
+
pending = {
|
|
641
|
+
**self.config,
|
|
642
|
+
**{
|
|
643
|
+
wire_name: normalized[field_for_wire[wire_name]]
|
|
644
|
+
for wire_name in payload
|
|
645
|
+
},
|
|
646
|
+
}
|
|
647
|
+
try:
|
|
648
|
+
self.config_model.model_validate(
|
|
649
|
+
{**self.config_defaults, **pending}, by_name=True
|
|
650
|
+
)
|
|
651
|
+
except ValidationError as exc:
|
|
652
|
+
raise HTTPException(
|
|
653
|
+
status_code=400,
|
|
654
|
+
detail=(
|
|
655
|
+
"AgentConfig serializers must return values accepted by their "
|
|
656
|
+
"declared field types: "
|
|
657
|
+
f"{_validation_error_detail(exc)}"
|
|
658
|
+
),
|
|
659
|
+
) from exc
|
|
660
|
+
provided_fields = {field_for_wire[name] for name in payload}
|
|
661
|
+
for identity_field in ("name", "description"):
|
|
662
|
+
if identity_field in provided_fields:
|
|
663
|
+
identity_value = getattr(validated, identity_field)
|
|
664
|
+
if not isinstance(identity_value, str) or not identity_value.strip():
|
|
665
|
+
raise HTTPException(
|
|
666
|
+
status_code=400,
|
|
667
|
+
detail=f"{identity_field} must be a non-empty string",
|
|
668
|
+
)
|
|
669
|
+
self.config.update(pending)
|
|
670
|
+
if self.card is not None:
|
|
671
|
+
if "name" in provided_fields:
|
|
672
|
+
self.card.name = validated.name
|
|
673
|
+
if "description" in provided_fields:
|
|
674
|
+
self.card.description = validated.description
|
|
675
|
+
return {"status": "updated"}
|
|
676
|
+
|
|
677
|
+
async def agent_config_get(self) -> dict[str, Any]:
|
|
678
|
+
config = _thaw(self.config)
|
|
679
|
+
for wire_name in self._write_only_config_fields:
|
|
680
|
+
if wire_name in config:
|
|
681
|
+
config[wire_name] = "***"
|
|
682
|
+
return {"config": config}
|
|
683
|
+
|
|
684
|
+
async def mcp_add(self, request: McpAddRequest) -> dict[str, Any]:
|
|
685
|
+
url = request.url.strip()
|
|
686
|
+
if not url:
|
|
687
|
+
raise HTTPException(status_code=400, detail="url must not be empty")
|
|
688
|
+
for existing_name, existing in self.mcp_servers.items():
|
|
689
|
+
if existing["url"] == url:
|
|
690
|
+
raise HTTPException(
|
|
691
|
+
status_code=409,
|
|
692
|
+
detail=f"MCP URL '{url}' already registered as '{existing_name}'",
|
|
693
|
+
)
|
|
694
|
+
name = request.name or f"mcp_{uuid.uuid4().hex[:8]}"
|
|
695
|
+
if name in self.mcp_servers:
|
|
696
|
+
raise HTTPException(status_code=409, detail=f"MCP name '{name}' is already registered")
|
|
697
|
+
self.mcp_servers[name] = {
|
|
698
|
+
"url": url,
|
|
699
|
+
"headers": dict(request.headers) if request.headers else None,
|
|
700
|
+
}
|
|
701
|
+
return {"status": "added", "name": name, "url": url}
|
|
702
|
+
|
|
703
|
+
async def mcp_list(self) -> dict[str, Any]:
|
|
704
|
+
return {
|
|
705
|
+
"mcp_servers": {
|
|
706
|
+
name: {
|
|
707
|
+
"url": registration["url"],
|
|
708
|
+
"has_headers": bool(registration.get("headers")),
|
|
709
|
+
}
|
|
710
|
+
for name, registration in self.mcp_servers.items()
|
|
711
|
+
}
|
|
712
|
+
}
|
|
713
|
+
|
|
714
|
+
async def skill_list(self) -> dict[str, Any]:
|
|
715
|
+
skills = {
|
|
716
|
+
name: {"description": description}
|
|
717
|
+
for name, description in self.identity_skills.items()
|
|
718
|
+
}
|
|
719
|
+
skills.update(
|
|
720
|
+
{
|
|
721
|
+
str(registration["name"]): {
|
|
722
|
+
"description": str(registration["description"])
|
|
723
|
+
}
|
|
724
|
+
for registration in self.skills
|
|
725
|
+
}
|
|
726
|
+
)
|
|
727
|
+
return {"skills": skills}
|
|
728
|
+
|
|
729
|
+
async def trajectory_get(
|
|
730
|
+
self, request: TaskTrajectoryRequest | TaskObjectTrajectoryRequest
|
|
731
|
+
) -> dict[str, Any]:
|
|
732
|
+
task_id = request.task_id
|
|
733
|
+
trajectory = self.task_trajectories.get(task_id)
|
|
734
|
+
if trajectory is None:
|
|
735
|
+
raise HTTPException(
|
|
736
|
+
status_code=404,
|
|
737
|
+
detail=f"No completed trajectory available for task '{task_id}'",
|
|
738
|
+
)
|
|
739
|
+
if isinstance(request, TaskTrajectoryRequest):
|
|
740
|
+
return {"trajectory": _thaw(trajectory.payload)}
|
|
741
|
+
body = json.dumps(
|
|
742
|
+
_thaw(trajectory.payload), separators=(",", ":"), default=str
|
|
743
|
+
).encode()
|
|
744
|
+
uploaded = await upload(request.objects.trajectory, body)
|
|
745
|
+
return {
|
|
746
|
+
"objects": {
|
|
747
|
+
"trajectory": uploaded.model_dump(mode="json"),
|
|
748
|
+
}
|
|
749
|
+
}
|
|
750
|
+
|
|
751
|
+
async def triggers_register(
|
|
752
|
+
self, request: TriggerRegisterRequest
|
|
753
|
+
) -> dict[str, Any]:
|
|
754
|
+
return self.trigger_engine.register(request.model_dump(mode="python"))
|
|
755
|
+
|
|
756
|
+
async def triggers_decide(self, request: TriggerDecideRequest) -> dict[str, Any]:
|
|
757
|
+
return self.trigger_engine.decide(
|
|
758
|
+
turn=request.turn,
|
|
759
|
+
solver_message=request.solver_message,
|
|
760
|
+
context_id=request.context_id,
|
|
761
|
+
env_triggers=request.env_triggers,
|
|
762
|
+
)
|
|
763
|
+
|
|
764
|
+
async def triggers_state(self) -> dict[str, Any]:
|
|
765
|
+
return self.trigger_engine.state()
|
|
766
|
+
|
|
767
|
+
|
|
768
|
+
class DefaultExtensionHandlers:
|
|
769
|
+
"""Stable access to default implementations of SDK-owned operations."""
|
|
770
|
+
|
|
771
|
+
def __init__(self, services: _SdkServices) -> None:
|
|
772
|
+
self._services = services
|
|
773
|
+
|
|
774
|
+
async def call(
|
|
775
|
+
self,
|
|
776
|
+
operation: OperationReference,
|
|
777
|
+
request: BaseModel | None = None,
|
|
778
|
+
) -> Any:
|
|
779
|
+
if not isinstance(operation, OperationReference):
|
|
780
|
+
raise TypeError("operation must be an OperationReference")
|
|
781
|
+
if operation.request_variant is not None:
|
|
782
|
+
raise ValueError("SDK delegation does not accept request variants")
|
|
783
|
+
definition = operation.extension_definition.operation(operation.operation)
|
|
784
|
+
if definition.implementation is not ImplementationOwner.SDK:
|
|
785
|
+
raise ValueError(
|
|
786
|
+
f"{operation.extension_definition.uri}.{operation.operation} "
|
|
787
|
+
"is not SDK-owned"
|
|
788
|
+
)
|
|
789
|
+
handler = self._services.handlers().get(
|
|
790
|
+
(operation.extension_definition.uri, operation.operation)
|
|
791
|
+
)
|
|
792
|
+
if handler is None:
|
|
793
|
+
raise ValueError(
|
|
794
|
+
f"no SDK implementation for "
|
|
795
|
+
f"{operation.extension_definition.uri}.{operation.operation}"
|
|
796
|
+
)
|
|
797
|
+
expected = None
|
|
798
|
+
if definition.request is not None:
|
|
799
|
+
if definition.request.model is not None:
|
|
800
|
+
expected = definition.request.model
|
|
801
|
+
else:
|
|
802
|
+
sdk_variant_models = [
|
|
803
|
+
variant.model
|
|
804
|
+
for variant in definition.request.variants
|
|
805
|
+
if (variant.implementation or definition.implementation)
|
|
806
|
+
is ImplementationOwner.SDK
|
|
807
|
+
]
|
|
808
|
+
matching_models = [
|
|
809
|
+
model for model in sdk_variant_models if isinstance(request, model)
|
|
810
|
+
]
|
|
811
|
+
if len(matching_models) == 1:
|
|
812
|
+
expected = matching_models[0]
|
|
813
|
+
if definition.request is None:
|
|
814
|
+
if request is not None:
|
|
815
|
+
raise TypeError(f"{definition.name} does not accept a request")
|
|
816
|
+
elif expected is None or not isinstance(request, expected):
|
|
817
|
+
expected_names = (
|
|
818
|
+
[variant.model.__name__ for variant in definition.request.variants]
|
|
819
|
+
if definition.request.variants
|
|
820
|
+
else [definition.request.model.__name__]
|
|
821
|
+
)
|
|
822
|
+
raise TypeError(
|
|
823
|
+
f"{definition.name} requires one of {expected_names}, got "
|
|
824
|
+
f"{type(request).__name__}"
|
|
825
|
+
)
|
|
826
|
+
return await _invoke(handler, request)
|
|
827
|
+
|
|
828
|
+
|
|
829
|
+
class A2AAgentApplication:
|
|
830
|
+
"""Resolved Agent Card, extension routes, and A2A task executor."""
|
|
831
|
+
|
|
832
|
+
def __init__(self, agent: Any, definition: _AgentDefinition | None = None) -> None:
|
|
833
|
+
definition = definition or getattr(agent, _AGENT_DEFINITION, None)
|
|
834
|
+
if definition is None:
|
|
835
|
+
raise TypeError("agent class must be decorated with @a2a_agent")
|
|
836
|
+
if getattr(agent, "_agentenv_a2a_application", None) is not None:
|
|
837
|
+
raise RuntimeError(
|
|
838
|
+
"agent instance is already bound to an application; create a new "
|
|
839
|
+
"agent instance for each application"
|
|
840
|
+
)
|
|
841
|
+
self.agent = agent
|
|
842
|
+
self.definition = definition
|
|
843
|
+
self.rpc_path = _rpc_path(definition.identity.url)
|
|
844
|
+
handler = getattr(agent, "run", None)
|
|
845
|
+
if (
|
|
846
|
+
not callable(handler)
|
|
847
|
+
or getattr(type(agent), "run", None) is AgentEnvAgent.run
|
|
848
|
+
):
|
|
849
|
+
raise TypeError("an agent must define run(request)")
|
|
850
|
+
resolved_handler = _unwrap_bound_handler(handler)
|
|
851
|
+
_validate_run_handler_signature(
|
|
852
|
+
resolved_handler, definition.config or AgentConfig
|
|
853
|
+
)
|
|
854
|
+
self.run_handler = handler
|
|
855
|
+
self.streaming = inspect.isasyncgenfunction(resolved_handler)
|
|
856
|
+
self.services = _SdkServices(definition.extensions, definition.config)
|
|
857
|
+
self.default_handlers = DefaultExtensionHandlers(self.services)
|
|
858
|
+
self.registry = build_registry(
|
|
859
|
+
agent,
|
|
860
|
+
definition.extensions,
|
|
861
|
+
sdk_handlers=self.services.handlers(),
|
|
862
|
+
request_model_overrides=(
|
|
863
|
+
{(AGENT_CONFIG_V1.uri, "set"): definition.config}
|
|
864
|
+
if definition.config is not None
|
|
865
|
+
else None
|
|
866
|
+
),
|
|
867
|
+
)
|
|
868
|
+
self.registry.reject_framework_route_collisions(
|
|
869
|
+
{
|
|
870
|
+
("/health", "GET"): "health",
|
|
871
|
+
(self.rpc_path, "POST"): "A2A JSON-RPC",
|
|
872
|
+
("/.well-known/agent.json", "GET"): "Agent Card",
|
|
873
|
+
("/.well-known/agent-card.json", "GET"): "Agent Card",
|
|
874
|
+
}
|
|
875
|
+
)
|
|
876
|
+
conformance = self.registry.conformance()
|
|
877
|
+
for override in conformance["standard_operation_overrides"]:
|
|
878
|
+
logger.warning(
|
|
879
|
+
"Agent overrides SDK operation %s.%s",
|
|
880
|
+
override["uri"],
|
|
881
|
+
override["operation"],
|
|
882
|
+
)
|
|
883
|
+
self.card = self._build_card()
|
|
884
|
+
self.services.attach_card(self.card)
|
|
885
|
+
try:
|
|
886
|
+
setattr(agent, "_agentenv_a2a_application", self)
|
|
887
|
+
except AttributeError:
|
|
888
|
+
# Slotted agents can still use the framework; only advanced
|
|
889
|
+
# runtime integrations that need application state lose this hook.
|
|
890
|
+
pass
|
|
891
|
+
|
|
892
|
+
self.executor = _StandardExecutor(
|
|
893
|
+
self.run_handler,
|
|
894
|
+
self.services,
|
|
895
|
+
workspace=definition.workspace,
|
|
896
|
+
streaming=self.streaming,
|
|
897
|
+
)
|
|
898
|
+
|
|
899
|
+
routes = [Route("/health", self._health, methods=["GET"])]
|
|
900
|
+
for extension in self.registry.extensions:
|
|
901
|
+
for operation in extension.operations.values():
|
|
902
|
+
routes.append(
|
|
903
|
+
Route(
|
|
904
|
+
operation.definition.path,
|
|
905
|
+
self._route_handler(extension.definition.uri, operation),
|
|
906
|
+
methods=[operation.definition.method],
|
|
907
|
+
)
|
|
908
|
+
)
|
|
909
|
+
kwargs = {"routes": routes}
|
|
910
|
+
if definition.lifespan is not None:
|
|
911
|
+
kwargs["lifespan"] = definition.lifespan
|
|
912
|
+
self.app = Starlette(**kwargs)
|
|
913
|
+
|
|
914
|
+
try:
|
|
915
|
+
from a2a.server.apps import A2AStarletteApplication
|
|
916
|
+
from a2a.server.request_handlers import DefaultRequestHandler
|
|
917
|
+
from a2a.server.tasks import InMemoryTaskStore
|
|
918
|
+
except ImportError as exc: # pragma: no cover - depends on installation extra
|
|
919
|
+
raise ImportError(
|
|
920
|
+
"A2A agent applications require: pip install 'agentenv-framework-protocol[agent]'"
|
|
921
|
+
) from exc
|
|
922
|
+
|
|
923
|
+
request_handler = DefaultRequestHandler(
|
|
924
|
+
agent_executor=self.executor,
|
|
925
|
+
task_store=InMemoryTaskStore(),
|
|
926
|
+
)
|
|
927
|
+
A2AStarletteApplication(
|
|
928
|
+
agent_card=self.card,
|
|
929
|
+
http_handler=request_handler,
|
|
930
|
+
).add_routes_to_app(
|
|
931
|
+
self.app,
|
|
932
|
+
rpc_url=self.rpc_path,
|
|
933
|
+
)
|
|
934
|
+
self.app.state.agentenv_a2a = self
|
|
935
|
+
|
|
936
|
+
async def _health(self, _request: Request) -> JSONResponse:
|
|
937
|
+
return JSONResponse({"status": "ok"})
|
|
938
|
+
|
|
939
|
+
def _build_card(self) -> Any:
|
|
940
|
+
try:
|
|
941
|
+
from a2a.types import AgentCard, AgentExtension, AgentSkill
|
|
942
|
+
except ImportError as exc: # pragma: no cover - depends on installation extra
|
|
943
|
+
raise ImportError(
|
|
944
|
+
"A2A agent applications require: pip install 'agentenv-framework-protocol[agent]'"
|
|
945
|
+
) from exc
|
|
946
|
+
|
|
947
|
+
identity = self.definition.identity
|
|
948
|
+
skills = [
|
|
949
|
+
skill if isinstance(skill, AgentSkill) else AgentSkill.model_validate(skill)
|
|
950
|
+
for skill in identity.skills
|
|
951
|
+
]
|
|
952
|
+
extensions = [
|
|
953
|
+
AgentExtension.model_validate(item)
|
|
954
|
+
for item in self.registry.card_extensions()
|
|
955
|
+
]
|
|
956
|
+
resolved_capabilities = AgentCapabilities(
|
|
957
|
+
streaming=self.streaming,
|
|
958
|
+
push_notifications=False,
|
|
959
|
+
state_transition_history=False,
|
|
960
|
+
extensions=extensions,
|
|
961
|
+
)
|
|
962
|
+
return AgentCard(
|
|
963
|
+
name=identity.name,
|
|
964
|
+
description=identity.description,
|
|
965
|
+
version=identity.version,
|
|
966
|
+
url=identity.url,
|
|
967
|
+
default_input_modes=list(identity.input_modes),
|
|
968
|
+
default_output_modes=list(identity.output_modes),
|
|
969
|
+
skills=skills,
|
|
970
|
+
capabilities=resolved_capabilities,
|
|
971
|
+
)
|
|
972
|
+
|
|
973
|
+
def _route_handler(
|
|
974
|
+
self, extension_uri: str, operation: RegisteredOperation
|
|
975
|
+
) -> Callable:
|
|
976
|
+
async def invoke(request: Request) -> JSONResponse:
|
|
977
|
+
try:
|
|
978
|
+
if operation.definition.method == "GET":
|
|
979
|
+
payload = dict(request.query_params)
|
|
980
|
+
else:
|
|
981
|
+
raw = await request.body()
|
|
982
|
+
payload = json.loads(raw) if raw else {}
|
|
983
|
+
if not isinstance(payload, dict):
|
|
984
|
+
raise ValueError("request body must be a JSON object")
|
|
985
|
+
input_payload = dict(payload)
|
|
986
|
+
handler, variant = operation.select_handler(payload)
|
|
987
|
+
if handler is None:
|
|
988
|
+
raise RuntimeError("operation has no implementation")
|
|
989
|
+
request_model = operation.request_model(variant)
|
|
990
|
+
if request_model is None:
|
|
991
|
+
if payload:
|
|
992
|
+
raise ValueError(
|
|
993
|
+
f"{operation.definition.name} does not accept request fields"
|
|
994
|
+
)
|
|
995
|
+
extension_request = None
|
|
996
|
+
else:
|
|
997
|
+
extension_request = request_model.model_validate(payload)
|
|
998
|
+
except HTTPException:
|
|
999
|
+
raise
|
|
1000
|
+
except ValidationError as exc:
|
|
1001
|
+
raise HTTPException(
|
|
1002
|
+
status_code=400, detail=_validation_error_detail(exc)
|
|
1003
|
+
) from exc
|
|
1004
|
+
except (TypeError, ValueError) as exc:
|
|
1005
|
+
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
|
1006
|
+
except Exception as exc:
|
|
1007
|
+
raise _opaque_extension_error("request resolution") from exc
|
|
1008
|
+
|
|
1009
|
+
try:
|
|
1010
|
+
async def invoke_and_validate() -> Any:
|
|
1011
|
+
result = await _invoke(handler, extension_request)
|
|
1012
|
+
result = (
|
|
1013
|
+
_thaw(result.model_dump(mode="json", by_alias=True))
|
|
1014
|
+
if isinstance(result, BaseModel)
|
|
1015
|
+
else _thaw(result)
|
|
1016
|
+
)
|
|
1017
|
+
if result is None:
|
|
1018
|
+
result = {}
|
|
1019
|
+
if operation.definition.response is not None:
|
|
1020
|
+
if not isinstance(result, Mapping):
|
|
1021
|
+
raise RuntimeError(
|
|
1022
|
+
"operation response must be a JSON object"
|
|
1023
|
+
)
|
|
1024
|
+
response_model = operation.definition.response_model
|
|
1025
|
+
if response_model is not None:
|
|
1026
|
+
return response_model.model_validate(result).model_dump(
|
|
1027
|
+
mode="json", by_alias=True, exclude_none=True
|
|
1028
|
+
)
|
|
1029
|
+
missing = set(operation.definition.response.required) - set(
|
|
1030
|
+
result
|
|
1031
|
+
)
|
|
1032
|
+
if missing:
|
|
1033
|
+
raise RuntimeError(
|
|
1034
|
+
"operation response is missing fields: "
|
|
1035
|
+
f"{sorted(missing)}"
|
|
1036
|
+
)
|
|
1037
|
+
return result
|
|
1038
|
+
|
|
1039
|
+
if (
|
|
1040
|
+
extension_uri == SKILL_CONFIG_V1.uri
|
|
1041
|
+
and operation.definition.name == "add"
|
|
1042
|
+
):
|
|
1043
|
+
async with self.services.skill_registration(input_payload):
|
|
1044
|
+
result = await invoke_and_validate()
|
|
1045
|
+
else:
|
|
1046
|
+
result = await invoke_and_validate()
|
|
1047
|
+
return JSONResponse(result)
|
|
1048
|
+
except HTTPException:
|
|
1049
|
+
raise
|
|
1050
|
+
except TriggerError as exc:
|
|
1051
|
+
raise HTTPException(
|
|
1052
|
+
status_code=exc.status_code, detail=str(exc)
|
|
1053
|
+
) from exc
|
|
1054
|
+
except TransferError as exc:
|
|
1055
|
+
return JSONResponse(exc.body(), status_code=exc.status_code)
|
|
1056
|
+
except Exception as exc:
|
|
1057
|
+
raise _opaque_extension_error("invocation") from exc
|
|
1058
|
+
|
|
1059
|
+
return invoke
|
|
1060
|
+
|
|
1061
|
+
|
|
1062
|
+
class _StandardExecutor:
|
|
1063
|
+
"""A2A SDK executor that delegates one normalized request to ``agent.run``."""
|
|
1064
|
+
|
|
1065
|
+
def __init__(
|
|
1066
|
+
self,
|
|
1067
|
+
handler: Callable,
|
|
1068
|
+
services: _SdkServices,
|
|
1069
|
+
*,
|
|
1070
|
+
workspace: Path | None,
|
|
1071
|
+
streaming: bool,
|
|
1072
|
+
) -> None:
|
|
1073
|
+
self._handler = handler
|
|
1074
|
+
self._services = services
|
|
1075
|
+
self._workspace = workspace
|
|
1076
|
+
self._streaming = streaming
|
|
1077
|
+
self._context_locks = _BoundedContextLocks()
|
|
1078
|
+
|
|
1079
|
+
async def execute(self, context: Any, event_queue: Any) -> None:
|
|
1080
|
+
from a2a.server.tasks import TaskUpdater
|
|
1081
|
+
from a2a.types import InvalidParamsError
|
|
1082
|
+
from a2a.utils import new_task
|
|
1083
|
+
from a2a.utils.errors import ServerError
|
|
1084
|
+
|
|
1085
|
+
try:
|
|
1086
|
+
task = context.current_task or new_task(context.message)
|
|
1087
|
+
except Exception as exc:
|
|
1088
|
+
raise ServerError(error=InvalidParamsError(message=str(exc))) from exc
|
|
1089
|
+
|
|
1090
|
+
updater = TaskUpdater(event_queue, task.id, task.context_id)
|
|
1091
|
+
try:
|
|
1092
|
+
if context.current_task is None:
|
|
1093
|
+
await event_queue.enqueue_event(task)
|
|
1094
|
+
await updater.start_work()
|
|
1095
|
+
|
|
1096
|
+
task_config = self._services.task_config()
|
|
1097
|
+
metadata = {
|
|
1098
|
+
**_thaw(getattr(context.message, "metadata", None) or {}),
|
|
1099
|
+
**({"role": task_config.role} if task_config.role is not None else {}),
|
|
1100
|
+
}
|
|
1101
|
+
async with self._context_locks.acquire(task.context_id):
|
|
1102
|
+
request = TaskRequest(
|
|
1103
|
+
task_id=task.id,
|
|
1104
|
+
context_id=task.context_id,
|
|
1105
|
+
parts=tuple(
|
|
1106
|
+
_from_a2a_part(part) for part in (context.message.parts or [])
|
|
1107
|
+
),
|
|
1108
|
+
config=task_config,
|
|
1109
|
+
mcp_servers=self._services.mcp_servers,
|
|
1110
|
+
skills=tuple(self._services.skills),
|
|
1111
|
+
metadata=metadata,
|
|
1112
|
+
session_ref=self._services.session_ref_for_context(task.context_id),
|
|
1113
|
+
workspace=self._workspace,
|
|
1114
|
+
)
|
|
1115
|
+
if self._streaming:
|
|
1116
|
+
result = await self._run_streaming(request, updater, event_queue)
|
|
1117
|
+
else:
|
|
1118
|
+
result = await self._handler(request)
|
|
1119
|
+
if not isinstance(result, TaskResult):
|
|
1120
|
+
raise TypeError("agent.run must return TaskResult")
|
|
1121
|
+
if result.session_ref is not None:
|
|
1122
|
+
self._services.set_session_ref_for_context(
|
|
1123
|
+
task.context_id, result.session_ref
|
|
1124
|
+
)
|
|
1125
|
+
if result.native_trajectory is not None:
|
|
1126
|
+
self._services.task_trajectories[task.id] = result.native_trajectory
|
|
1127
|
+
|
|
1128
|
+
parts = _to_a2a_parts(result)
|
|
1129
|
+
message = updater.new_agent_message(parts=parts)
|
|
1130
|
+
if result.outcome is TaskOutcome.FAILED:
|
|
1131
|
+
await updater.failed(message=message)
|
|
1132
|
+
else:
|
|
1133
|
+
await updater.complete(message=message)
|
|
1134
|
+
except Exception:
|
|
1135
|
+
correlation_id = uuid.uuid4().hex
|
|
1136
|
+
logger.exception(
|
|
1137
|
+
"unhandled task execution error (correlation_id=%s)",
|
|
1138
|
+
correlation_id,
|
|
1139
|
+
)
|
|
1140
|
+
failure = TaskResult.failure(
|
|
1141
|
+
"framework.unhandled_exception",
|
|
1142
|
+
"An unexpected framework error occurred. "
|
|
1143
|
+
f"Correlation ID: {correlation_id}",
|
|
1144
|
+
error_type="infra_error",
|
|
1145
|
+
)
|
|
1146
|
+
await updater.failed(
|
|
1147
|
+
message=updater.new_agent_message(parts=_to_a2a_parts(failure))
|
|
1148
|
+
)
|
|
1149
|
+
|
|
1150
|
+
async def _run_streaming(
|
|
1151
|
+
self,
|
|
1152
|
+
request: TaskRequest[Any],
|
|
1153
|
+
updater: Any,
|
|
1154
|
+
event_queue: Any,
|
|
1155
|
+
) -> TaskResult:
|
|
1156
|
+
from a2a.types import (
|
|
1157
|
+
TaskState,
|
|
1158
|
+
TaskStatusUpdateEvent,
|
|
1159
|
+
)
|
|
1160
|
+
|
|
1161
|
+
stream = self._handler(request)
|
|
1162
|
+
if inspect.isawaitable(stream):
|
|
1163
|
+
stream = await stream
|
|
1164
|
+
try:
|
|
1165
|
+
async for item in stream:
|
|
1166
|
+
if isinstance(item, TaskResult):
|
|
1167
|
+
return item
|
|
1168
|
+
if isinstance(item, TaskProgress):
|
|
1169
|
+
await updater.update_status(
|
|
1170
|
+
TaskState.working,
|
|
1171
|
+
message=updater.new_agent_message(
|
|
1172
|
+
parts=_to_a2a_parts_collection(item.parts)
|
|
1173
|
+
),
|
|
1174
|
+
metadata=_thaw(item.metadata) or None,
|
|
1175
|
+
)
|
|
1176
|
+
continue
|
|
1177
|
+
if isinstance(item, TaskStatusUpdateEvent):
|
|
1178
|
+
self._validate_status_event(item, request)
|
|
1179
|
+
await event_queue.enqueue_event(item)
|
|
1180
|
+
continue
|
|
1181
|
+
raise TypeError(
|
|
1182
|
+
"agent.run (streaming) must yield TaskProgress, "
|
|
1183
|
+
"a non-terminal A2A status update event, or a terminal "
|
|
1184
|
+
"TaskResult"
|
|
1185
|
+
)
|
|
1186
|
+
finally:
|
|
1187
|
+
await stream.aclose()
|
|
1188
|
+
raise TypeError(
|
|
1189
|
+
"agent.run (streaming) must yield a TaskResult before completing"
|
|
1190
|
+
)
|
|
1191
|
+
|
|
1192
|
+
@staticmethod
|
|
1193
|
+
def _validate_event_identity(item: Any, request: TaskRequest[Any]) -> None:
|
|
1194
|
+
if item.task_id != request.task_id or item.context_id != request.context_id:
|
|
1195
|
+
raise ValueError(
|
|
1196
|
+
"streaming A2A event task_id and context_id must match the request"
|
|
1197
|
+
)
|
|
1198
|
+
|
|
1199
|
+
@classmethod
|
|
1200
|
+
def _validate_status_event(cls, item: Any, request: TaskRequest[Any]) -> None:
|
|
1201
|
+
from a2a.types import TaskState
|
|
1202
|
+
|
|
1203
|
+
cls._validate_event_identity(item, request)
|
|
1204
|
+
terminal_states = {
|
|
1205
|
+
TaskState.completed,
|
|
1206
|
+
TaskState.canceled,
|
|
1207
|
+
TaskState.failed,
|
|
1208
|
+
TaskState.rejected,
|
|
1209
|
+
}
|
|
1210
|
+
if item.final or item.status.state in terminal_states:
|
|
1211
|
+
raise ValueError(
|
|
1212
|
+
"streaming TaskStatusUpdateEvent must be non-terminal and final=False"
|
|
1213
|
+
)
|
|
1214
|
+
|
|
1215
|
+
async def cancel(self, _context: Any, _event_queue: Any) -> None:
|
|
1216
|
+
"""Satisfy the upstream executor interface without offering cancellation."""
|
|
1217
|
+
from a2a.types import UnsupportedOperationError
|
|
1218
|
+
from a2a.utils.errors import ServerError
|
|
1219
|
+
|
|
1220
|
+
raise ServerError(error=UnsupportedOperationError())
|
|
1221
|
+
|
|
1222
|
+
|
|
1223
|
+
def _from_a2a_part(part: Any) -> Any:
|
|
1224
|
+
root = part.root
|
|
1225
|
+
kind = getattr(root, "kind", None)
|
|
1226
|
+
metadata = getattr(root, "metadata", None) or {}
|
|
1227
|
+
if kind == "text":
|
|
1228
|
+
return TaskTextPart(text=root.text, metadata=metadata)
|
|
1229
|
+
if kind == "data":
|
|
1230
|
+
return TaskDataPart(data=root.data, metadata=metadata)
|
|
1231
|
+
if kind == "file":
|
|
1232
|
+
file = root.file
|
|
1233
|
+
return TaskFilePart(
|
|
1234
|
+
name=getattr(file, "name", None),
|
|
1235
|
+
mime_type=getattr(file, "mime_type", None),
|
|
1236
|
+
bytes=getattr(file, "bytes", None),
|
|
1237
|
+
uri=getattr(file, "uri", None),
|
|
1238
|
+
metadata=metadata,
|
|
1239
|
+
)
|
|
1240
|
+
raise ValueError(f"unsupported A2A part kind: {kind!r}")
|
|
1241
|
+
|
|
1242
|
+
|
|
1243
|
+
def _to_a2a_parts(result: TaskResult) -> list[Any]:
|
|
1244
|
+
converted = _to_a2a_parts_collection(result.parts)
|
|
1245
|
+
|
|
1246
|
+
from a2a.types import DataPart, Part
|
|
1247
|
+
|
|
1248
|
+
metadata: dict[str, Any] = {}
|
|
1249
|
+
if not result.usage.is_empty:
|
|
1250
|
+
metadata["usage"] = result.usage.to_dict()
|
|
1251
|
+
if result.error is not None:
|
|
1252
|
+
metadata["error_type"] = result.error.error_type
|
|
1253
|
+
metadata["error_code"] = result.error.code
|
|
1254
|
+
metadata["error_message"] = result.error.message
|
|
1255
|
+
if metadata:
|
|
1256
|
+
converted.append(Part(root=DataPart(data=metadata)))
|
|
1257
|
+
return converted
|
|
1258
|
+
|
|
1259
|
+
|
|
1260
|
+
def _to_a2a_parts_collection(parts: Iterable[Any]) -> list[Any]:
|
|
1261
|
+
from a2a.types import DataPart, FilePart, FileWithBytes, FileWithUri, Part, TextPart
|
|
1262
|
+
|
|
1263
|
+
converted = []
|
|
1264
|
+
for part in parts:
|
|
1265
|
+
metadata = _thaw(part.metadata) or None
|
|
1266
|
+
if isinstance(part, TaskTextPart):
|
|
1267
|
+
root = TextPart(text=part.text, metadata=metadata)
|
|
1268
|
+
elif isinstance(part, TaskDataPart):
|
|
1269
|
+
root = DataPart(data=_thaw(part.data), metadata=metadata)
|
|
1270
|
+
elif isinstance(part, TaskFilePart):
|
|
1271
|
+
file = (
|
|
1272
|
+
FileWithBytes(
|
|
1273
|
+
bytes=part.bytes, name=part.name, mime_type=part.mime_type
|
|
1274
|
+
)
|
|
1275
|
+
if part.bytes is not None
|
|
1276
|
+
else FileWithUri(uri=part.uri, name=part.name, mime_type=part.mime_type)
|
|
1277
|
+
)
|
|
1278
|
+
root = FilePart(file=file, metadata=metadata)
|
|
1279
|
+
else: # pragma: no cover - TaskPart is closed
|
|
1280
|
+
raise TypeError(f"unsupported task result part: {type(part).__name__}")
|
|
1281
|
+
converted.append(Part(root=root))
|
|
1282
|
+
|
|
1283
|
+
return converted
|