kcs-agent 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- kcs_agent/__init__.py +175 -0
- kcs_agent/agent.py +409 -0
- kcs_agent/events.py +79 -0
- kcs_agent/exceptions.py +26 -0
- kcs_agent/extensions.py +170 -0
- kcs_agent/messages.py +240 -0
- kcs_agent/model.py +127 -0
- kcs_agent/providers.py +999 -0
- kcs_agent/py.typed +0 -0
- kcs_agent/schema.py +225 -0
- kcs_agent/state.py +210 -0
- kcs_agent/tools.py +331 -0
- kcs_agent-0.1.0.dist-info/METADATA +227 -0
- kcs_agent-0.1.0.dist-info/RECORD +16 -0
- kcs_agent-0.1.0.dist-info/WHEEL +4 -0
- kcs_agent-0.1.0.dist-info/licenses/LICENSE +21 -0
kcs_agent/__init__.py
ADDED
|
@@ -0,0 +1,175 @@
|
|
|
1
|
+
from .agent import Agent
|
|
2
|
+
from .events import AgentEvent, AgentEventType, AgentResult
|
|
3
|
+
from .exceptions import (
|
|
4
|
+
AgentCancelledError,
|
|
5
|
+
AgentError,
|
|
6
|
+
AgentIterationLimitError,
|
|
7
|
+
AgentProtocolError,
|
|
8
|
+
SessionVersionConflictError,
|
|
9
|
+
ToolDefinitionError,
|
|
10
|
+
ToolInvocationError,
|
|
11
|
+
)
|
|
12
|
+
from .extensions import (
|
|
13
|
+
AgentContext,
|
|
14
|
+
AgentExtension,
|
|
15
|
+
ExtensionFailurePolicy,
|
|
16
|
+
HistoryExtension,
|
|
17
|
+
StateExtension,
|
|
18
|
+
StaticContextExtension,
|
|
19
|
+
ToolExtension,
|
|
20
|
+
ToolPromptExtension,
|
|
21
|
+
)
|
|
22
|
+
from .messages import (
|
|
23
|
+
AnyMessage,
|
|
24
|
+
AssistantMessage,
|
|
25
|
+
AssistantMessageData,
|
|
26
|
+
AssistantMessageMetadata,
|
|
27
|
+
ImageAssetSource,
|
|
28
|
+
ImageBytesSource,
|
|
29
|
+
ImageContent,
|
|
30
|
+
ImageDetail,
|
|
31
|
+
ImageSource,
|
|
32
|
+
ImageUrlSource,
|
|
33
|
+
Message,
|
|
34
|
+
MessageRole,
|
|
35
|
+
SystemMessage,
|
|
36
|
+
SystemMessageData,
|
|
37
|
+
SystemMessageMetadata,
|
|
38
|
+
TextContent,
|
|
39
|
+
ToolCall,
|
|
40
|
+
ToolMessage,
|
|
41
|
+
ToolMessageData,
|
|
42
|
+
ToolMessageMetadata,
|
|
43
|
+
UserContent,
|
|
44
|
+
UserContentPart,
|
|
45
|
+
UserMessage,
|
|
46
|
+
UserMessageData,
|
|
47
|
+
UserMessageMetadata,
|
|
48
|
+
)
|
|
49
|
+
from .model import (
|
|
50
|
+
AgentModel,
|
|
51
|
+
ModelEvent,
|
|
52
|
+
ModelEventType,
|
|
53
|
+
ModelRequest,
|
|
54
|
+
ModelResponse,
|
|
55
|
+
ModelUsage,
|
|
56
|
+
ReasoningEffort,
|
|
57
|
+
ToolCallDelta,
|
|
58
|
+
ToolDefinition,
|
|
59
|
+
)
|
|
60
|
+
from .providers import (
|
|
61
|
+
AnthropicProvider,
|
|
62
|
+
DeepSeekProvider,
|
|
63
|
+
GoogleProvider,
|
|
64
|
+
OllamaProvider,
|
|
65
|
+
OpenAIProvider,
|
|
66
|
+
ProviderAuthError,
|
|
67
|
+
ProviderError,
|
|
68
|
+
ProviderResponseError,
|
|
69
|
+
)
|
|
70
|
+
from .schema import Parameter, annotation_schema, callable_schema
|
|
71
|
+
from .state import (
|
|
72
|
+
AgentErrorInfo,
|
|
73
|
+
AgentState,
|
|
74
|
+
CancellationToken,
|
|
75
|
+
CheckpointReason,
|
|
76
|
+
ExtensionStateRegistry,
|
|
77
|
+
RunState,
|
|
78
|
+
RunStatus,
|
|
79
|
+
SessionState,
|
|
80
|
+
ToolCallStreamState,
|
|
81
|
+
UsageState,
|
|
82
|
+
)
|
|
83
|
+
from .tools import (
|
|
84
|
+
AgentTool,
|
|
85
|
+
ToolDocstring,
|
|
86
|
+
ToolParameterDocumentation,
|
|
87
|
+
parse_tool,
|
|
88
|
+
parse_tool_docstring,
|
|
89
|
+
render_tool_guidance,
|
|
90
|
+
tool,
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
__all__ = [
|
|
94
|
+
"Agent",
|
|
95
|
+
"AgentContext",
|
|
96
|
+
"AgentCancelledError",
|
|
97
|
+
"AgentError",
|
|
98
|
+
"AgentErrorInfo",
|
|
99
|
+
"AgentEvent",
|
|
100
|
+
"AgentEventType",
|
|
101
|
+
"AgentExtension",
|
|
102
|
+
"AgentIterationLimitError",
|
|
103
|
+
"AgentModel",
|
|
104
|
+
"AgentProtocolError",
|
|
105
|
+
"AgentResult",
|
|
106
|
+
"AgentState",
|
|
107
|
+
"AgentTool",
|
|
108
|
+
"AnyMessage",
|
|
109
|
+
"AssistantMessage",
|
|
110
|
+
"AssistantMessageData",
|
|
111
|
+
"AssistantMessageMetadata",
|
|
112
|
+
"CancellationToken",
|
|
113
|
+
"CheckpointReason",
|
|
114
|
+
"ExtensionFailurePolicy",
|
|
115
|
+
"ExtensionStateRegistry",
|
|
116
|
+
"HistoryExtension",
|
|
117
|
+
"ImageAssetSource",
|
|
118
|
+
"ImageBytesSource",
|
|
119
|
+
"ImageContent",
|
|
120
|
+
"ImageDetail",
|
|
121
|
+
"ImageSource",
|
|
122
|
+
"ImageUrlSource",
|
|
123
|
+
"Message",
|
|
124
|
+
"MessageRole",
|
|
125
|
+
"ModelEvent",
|
|
126
|
+
"ModelEventType",
|
|
127
|
+
"ModelRequest",
|
|
128
|
+
"ModelResponse",
|
|
129
|
+
"ModelUsage",
|
|
130
|
+
"Parameter",
|
|
131
|
+
"ReasoningEffort",
|
|
132
|
+
"RunState",
|
|
133
|
+
"RunStatus",
|
|
134
|
+
"SessionState",
|
|
135
|
+
"SessionVersionConflictError",
|
|
136
|
+
"StateExtension",
|
|
137
|
+
"StaticContextExtension",
|
|
138
|
+
"SystemMessage",
|
|
139
|
+
"SystemMessageData",
|
|
140
|
+
"SystemMessageMetadata",
|
|
141
|
+
"ToolCall",
|
|
142
|
+
"ToolCallDelta",
|
|
143
|
+
"ToolCallStreamState",
|
|
144
|
+
"ToolDefinition",
|
|
145
|
+
"ToolDefinitionError",
|
|
146
|
+
"ToolDocstring",
|
|
147
|
+
"ToolExtension",
|
|
148
|
+
"ToolPromptExtension",
|
|
149
|
+
"ToolInvocationError",
|
|
150
|
+
"ToolMessage",
|
|
151
|
+
"ToolMessageData",
|
|
152
|
+
"ToolMessageMetadata",
|
|
153
|
+
"ToolParameterDocumentation",
|
|
154
|
+
"TextContent",
|
|
155
|
+
"UsageState",
|
|
156
|
+
"annotation_schema",
|
|
157
|
+
"callable_schema",
|
|
158
|
+
"parse_tool",
|
|
159
|
+
"parse_tool_docstring",
|
|
160
|
+
"render_tool_guidance",
|
|
161
|
+
"AnthropicProvider",
|
|
162
|
+
"DeepSeekProvider",
|
|
163
|
+
"GoogleProvider",
|
|
164
|
+
"OllamaProvider",
|
|
165
|
+
"OpenAIProvider",
|
|
166
|
+
"ProviderAuthError",
|
|
167
|
+
"ProviderError",
|
|
168
|
+
"ProviderResponseError",
|
|
169
|
+
"tool",
|
|
170
|
+
"UserMessage",
|
|
171
|
+
"UserContent",
|
|
172
|
+
"UserContentPart",
|
|
173
|
+
"UserMessageData",
|
|
174
|
+
"UserMessageMetadata",
|
|
175
|
+
]
|
kcs_agent/agent.py
ADDED
|
@@ -0,0 +1,409 @@
|
|
|
1
|
+
import asyncio
|
|
2
|
+
import json
|
|
3
|
+
from collections.abc import AsyncIterator, Callable, Sequence
|
|
4
|
+
from contextlib import aclosing
|
|
5
|
+
from typing import Any, cast
|
|
6
|
+
|
|
7
|
+
from .events import AgentEvent, AgentEventType, AgentResult
|
|
8
|
+
from .exceptions import AgentCancelledError, AgentIterationLimitError, AgentProtocolError, ToolDefinitionError
|
|
9
|
+
from .extensions import AgentContext, AgentExtension, ExtensionFailurePolicy
|
|
10
|
+
from .messages import (
|
|
11
|
+
AnyMessage,
|
|
12
|
+
Message,
|
|
13
|
+
SystemMessage,
|
|
14
|
+
SystemMessageData,
|
|
15
|
+
ToolCall,
|
|
16
|
+
ToolMessage,
|
|
17
|
+
ToolMessageData,
|
|
18
|
+
ToolMessageMetadata,
|
|
19
|
+
)
|
|
20
|
+
from .model import AgentModel, ModelEventType, ModelRequest, ModelResponse, ReasoningEffort
|
|
21
|
+
from .state import (
|
|
22
|
+
AgentErrorInfo,
|
|
23
|
+
AgentState,
|
|
24
|
+
CancellationToken,
|
|
25
|
+
CheckpointReason,
|
|
26
|
+
RunState,
|
|
27
|
+
RunStatus,
|
|
28
|
+
SessionState,
|
|
29
|
+
utc_now,
|
|
30
|
+
)
|
|
31
|
+
from .tools import AgentTool, parse_tool
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
class Agent[SessionDataT]:
|
|
35
|
+
"""Minimal provider-neutral agent loop with typed state and extension-owned capabilities."""
|
|
36
|
+
|
|
37
|
+
def __init__(
|
|
38
|
+
self,
|
|
39
|
+
model: AgentModel,
|
|
40
|
+
*,
|
|
41
|
+
system_prompt: str = "You are a helpful assistant.",
|
|
42
|
+
extensions: Sequence[AgentExtension[SessionDataT]] = (),
|
|
43
|
+
max_iterations: int = 12,
|
|
44
|
+
default_reasoning_effort: ReasoningEffort = ReasoningEffort.MEDIUM,
|
|
45
|
+
state_factory: Callable[[], SessionDataT] | None = None,
|
|
46
|
+
) -> None:
|
|
47
|
+
if max_iterations < 1:
|
|
48
|
+
raise ValueError("max_iterations must be at least one")
|
|
49
|
+
self._model = model
|
|
50
|
+
self._system_prompt = system_prompt.strip()
|
|
51
|
+
self._extensions = tuple(extensions)
|
|
52
|
+
self._max_iterations = max_iterations
|
|
53
|
+
self._default_reasoning_effort = default_reasoning_effort
|
|
54
|
+
self._state_factory = state_factory
|
|
55
|
+
|
|
56
|
+
async def run(
|
|
57
|
+
self,
|
|
58
|
+
messages: AnyMessage | Sequence[AnyMessage],
|
|
59
|
+
*,
|
|
60
|
+
session: SessionState[SessionDataT] | None = None,
|
|
61
|
+
metadata: dict[str, Any] | None = None,
|
|
62
|
+
reasoning_effort: ReasoningEffort | None = None,
|
|
63
|
+
cancellation_token: CancellationToken | None = None,
|
|
64
|
+
) -> AgentResult[SessionDataT]:
|
|
65
|
+
"""Run to completion and return the final assistant message with typed state."""
|
|
66
|
+
result: AgentResult[SessionDataT] | None = None
|
|
67
|
+
cancelled: AgentCancelledError | None = None
|
|
68
|
+
async for event in self.stream(
|
|
69
|
+
messages,
|
|
70
|
+
session=session,
|
|
71
|
+
metadata=metadata,
|
|
72
|
+
reasoning_effort=reasoning_effort,
|
|
73
|
+
cancellation_token=cancellation_token,
|
|
74
|
+
):
|
|
75
|
+
match event.type:
|
|
76
|
+
case AgentEventType.RUN_COMPLETED:
|
|
77
|
+
result = cast(AgentResult[SessionDataT], event.result)
|
|
78
|
+
case AgentEventType.RUN_CANCELLED:
|
|
79
|
+
cancelled = cast(AgentCancelledError, event.error)
|
|
80
|
+
case _:
|
|
81
|
+
pass
|
|
82
|
+
if cancelled is not None:
|
|
83
|
+
raise cancelled
|
|
84
|
+
if result is None:
|
|
85
|
+
raise AgentProtocolError("Agent stream ended without a final result")
|
|
86
|
+
return result
|
|
87
|
+
|
|
88
|
+
async def stream(
|
|
89
|
+
self,
|
|
90
|
+
messages: AnyMessage | Sequence[AnyMessage],
|
|
91
|
+
*,
|
|
92
|
+
session: SessionState[SessionDataT] | None = None,
|
|
93
|
+
metadata: dict[str, Any] | None = None,
|
|
94
|
+
reasoning_effort: ReasoningEffort | None = None,
|
|
95
|
+
cancellation_token: CancellationToken | None = None,
|
|
96
|
+
) -> AsyncIterator[AgentEvent]:
|
|
97
|
+
"""Execute model and tool steps while yielding typed lifecycle events in exact order."""
|
|
98
|
+
initial_messages = self._normalize_messages(messages)
|
|
99
|
+
context = AgentContext(
|
|
100
|
+
state=AgentState(
|
|
101
|
+
run=RunState(
|
|
102
|
+
messages=initial_messages,
|
|
103
|
+
reasoning_effort=(
|
|
104
|
+
reasoning_effort if reasoning_effort is not None else self._default_reasoning_effort
|
|
105
|
+
),
|
|
106
|
+
),
|
|
107
|
+
session=session or self._create_session(),
|
|
108
|
+
),
|
|
109
|
+
metadata=dict(metadata or {}),
|
|
110
|
+
)
|
|
111
|
+
run = context.state.run
|
|
112
|
+
token = cancellation_token or CancellationToken()
|
|
113
|
+
try:
|
|
114
|
+
await self._call_extensions("load_state", context)
|
|
115
|
+
run.status = RunStatus.RUNNING
|
|
116
|
+
run.started_at = utc_now()
|
|
117
|
+
await self._call_extensions("on_run_start", context)
|
|
118
|
+
tools = await self._collect_tools(context)
|
|
119
|
+
injected = await self._collect_context(context)
|
|
120
|
+
system = [SystemMessage(data=SystemMessageData(self._system_prompt))] if self._system_prompt else []
|
|
121
|
+
run.messages = [*system, *injected, *run.messages]
|
|
122
|
+
yield self._event(context, AgentEventType.RUN_STARTED)
|
|
123
|
+
await self._checkpoint(context, CheckpointReason.RUN_STARTED)
|
|
124
|
+
|
|
125
|
+
for iteration in range(1, self._max_iterations + 1):
|
|
126
|
+
self._ensure_not_cancelled(run, token)
|
|
127
|
+
run.iteration = iteration
|
|
128
|
+
run.text_buffer = ""
|
|
129
|
+
run.reasoning_buffer = ""
|
|
130
|
+
run.tool_call_streams = {}
|
|
131
|
+
request = ModelRequest(
|
|
132
|
+
messages=tuple(run.messages),
|
|
133
|
+
tools=tuple(item.definition for item in tools.values()),
|
|
134
|
+
reasoning_effort=run.reasoning_effort,
|
|
135
|
+
metadata=context.metadata,
|
|
136
|
+
)
|
|
137
|
+
for extension in self._extensions:
|
|
138
|
+
replacement = await self._invoke_extension(extension, "before_model", context, request)
|
|
139
|
+
if replacement is not None:
|
|
140
|
+
if not isinstance(replacement, ModelRequest):
|
|
141
|
+
raise AgentProtocolError("before_model must return ModelRequest or None")
|
|
142
|
+
request = replacement
|
|
143
|
+
run.status = RunStatus.WAITING_FOR_MODEL
|
|
144
|
+
yield self._event(context, AgentEventType.MODEL_STARTED)
|
|
145
|
+
response: ModelResponse | None = None
|
|
146
|
+
response_seen = False
|
|
147
|
+
async with aclosing(self._model.stream(request)) as model_events:
|
|
148
|
+
async for model_event in model_events:
|
|
149
|
+
self._ensure_not_cancelled(run, token)
|
|
150
|
+
if response_seen:
|
|
151
|
+
raise AgentProtocolError("A model adapter emitted events after its final response")
|
|
152
|
+
match model_event.type:
|
|
153
|
+
case ModelEventType.TEXT_DELTA:
|
|
154
|
+
run.text_buffer += model_event.delta
|
|
155
|
+
yield self._event(
|
|
156
|
+
context,
|
|
157
|
+
AgentEventType.TEXT_DELTA,
|
|
158
|
+
delta=model_event.delta,
|
|
159
|
+
)
|
|
160
|
+
case ModelEventType.REASONING_DELTA:
|
|
161
|
+
run.reasoning_buffer += model_event.delta
|
|
162
|
+
yield self._event(
|
|
163
|
+
context,
|
|
164
|
+
AgentEventType.REASONING_DELTA,
|
|
165
|
+
delta=model_event.delta,
|
|
166
|
+
)
|
|
167
|
+
case ModelEventType.TOOL_CALL_DELTA:
|
|
168
|
+
if model_event.tool_call_delta is None:
|
|
169
|
+
raise AgentProtocolError("A tool-call event must contain a ToolCallDelta")
|
|
170
|
+
run.append_tool_call_delta(model_event.tool_call_delta)
|
|
171
|
+
yield self._event(
|
|
172
|
+
context,
|
|
173
|
+
AgentEventType.TOOL_CALL_DELTA,
|
|
174
|
+
tool_call_delta=model_event.tool_call_delta,
|
|
175
|
+
)
|
|
176
|
+
case ModelEventType.RESPONSE:
|
|
177
|
+
if model_event.response is None:
|
|
178
|
+
raise AgentProtocolError("A response event must contain a ModelResponse")
|
|
179
|
+
response = model_event.response
|
|
180
|
+
response_seen = True
|
|
181
|
+
case _:
|
|
182
|
+
raise AgentProtocolError(f"Unsupported model event: {model_event.type}")
|
|
183
|
+
if response is None:
|
|
184
|
+
raise AgentProtocolError("A model stream must finish with exactly one response event")
|
|
185
|
+
|
|
186
|
+
run.status = RunStatus.RUNNING
|
|
187
|
+
run.usage.add_model_usage(response.usage)
|
|
188
|
+
await self._call_extensions("after_model", context, response)
|
|
189
|
+
await self._append_message(context, response.message)
|
|
190
|
+
yield self._event(context, AgentEventType.MODEL_COMPLETED, response=response)
|
|
191
|
+
await self._checkpoint(context, CheckpointReason.MODEL_COMPLETED)
|
|
192
|
+
|
|
193
|
+
match response.message.data.tool_calls:
|
|
194
|
+
case ():
|
|
195
|
+
run.status = RunStatus.COMPLETED
|
|
196
|
+
run.finished_at = utc_now()
|
|
197
|
+
await self._call_extensions("on_run_end", context, response.message)
|
|
198
|
+
await self._checkpoint(context, CheckpointReason.RUN_COMPLETED)
|
|
199
|
+
result = AgentResult(response.message, context.state, dict(context.metadata))
|
|
200
|
+
yield self._event(context, AgentEventType.RUN_COMPLETED, result=result)
|
|
201
|
+
return
|
|
202
|
+
case calls:
|
|
203
|
+
for call in calls:
|
|
204
|
+
self._ensure_not_cancelled(run, token)
|
|
205
|
+
async for event in self._execute_tool(context, tools, call):
|
|
206
|
+
yield event
|
|
207
|
+
|
|
208
|
+
raise AgentIterationLimitError(f"Agent exceeded {self._max_iterations} model iterations")
|
|
209
|
+
except AgentCancelledError as error:
|
|
210
|
+
run.status = RunStatus.CANCELLED
|
|
211
|
+
run.finished_at = utc_now()
|
|
212
|
+
run.error = AgentErrorInfo.from_exception(error)
|
|
213
|
+
await self._terminal_checkpoint(context, CheckpointReason.RUN_CANCELLED)
|
|
214
|
+
yield self._event(context, AgentEventType.RUN_CANCELLED, error=error)
|
|
215
|
+
except (asyncio.CancelledError, GeneratorExit) as error:
|
|
216
|
+
if run.status not in (RunStatus.COMPLETED, RunStatus.FAILED, RunStatus.CANCELLED):
|
|
217
|
+
run.status = RunStatus.CANCELLED
|
|
218
|
+
run.finished_at = utc_now()
|
|
219
|
+
run.error = AgentErrorInfo.from_exception(error)
|
|
220
|
+
await self._terminal_checkpoint(context, CheckpointReason.RUN_CANCELLED)
|
|
221
|
+
raise
|
|
222
|
+
except Exception as error:
|
|
223
|
+
run.status = RunStatus.FAILED
|
|
224
|
+
run.finished_at = utc_now()
|
|
225
|
+
run.error = AgentErrorInfo.from_exception(error)
|
|
226
|
+
await self._notify_error(context, error)
|
|
227
|
+
await self._terminal_checkpoint(context, CheckpointReason.RUN_FAILED)
|
|
228
|
+
yield self._event(context, AgentEventType.RUN_FAILED, error=error)
|
|
229
|
+
raise
|
|
230
|
+
finally:
|
|
231
|
+
await self._release_state(context)
|
|
232
|
+
|
|
233
|
+
def _normalize_messages(self, messages: AnyMessage | Sequence[AnyMessage]) -> list[AnyMessage]:
|
|
234
|
+
match messages:
|
|
235
|
+
case Message():
|
|
236
|
+
return [messages]
|
|
237
|
+
case _:
|
|
238
|
+
return list(messages)
|
|
239
|
+
|
|
240
|
+
def _create_session(self) -> SessionState[SessionDataT]:
|
|
241
|
+
data = self._state_factory() if self._state_factory is not None else cast(SessionDataT, {})
|
|
242
|
+
return SessionState(data=data)
|
|
243
|
+
|
|
244
|
+
async def _collect_tools(self, context: AgentContext[SessionDataT]) -> dict[str, AgentTool]:
|
|
245
|
+
result: dict[str, AgentTool] = {}
|
|
246
|
+
for extension in self._extensions:
|
|
247
|
+
candidates = await self._invoke_extension(extension, "tools", context)
|
|
248
|
+
for candidate in candidates or ():
|
|
249
|
+
registered = candidate if isinstance(candidate, AgentTool) else parse_tool(candidate)
|
|
250
|
+
if registered.name in result:
|
|
251
|
+
raise ToolDefinitionError(f"Duplicate tool name: {registered.name}")
|
|
252
|
+
result[registered.name] = registered
|
|
253
|
+
return result
|
|
254
|
+
|
|
255
|
+
async def _collect_context(self, context: AgentContext[SessionDataT]) -> list[AnyMessage]:
|
|
256
|
+
messages: list[AnyMessage] = []
|
|
257
|
+
for extension in self._extensions:
|
|
258
|
+
injected = await self._invoke_extension(extension, "context_messages", context)
|
|
259
|
+
messages.extend(injected or ())
|
|
260
|
+
return messages
|
|
261
|
+
|
|
262
|
+
async def _execute_tool(
|
|
263
|
+
self,
|
|
264
|
+
context: AgentContext[SessionDataT],
|
|
265
|
+
tools: dict[str, AgentTool],
|
|
266
|
+
call: ToolCall,
|
|
267
|
+
) -> AsyncIterator[AgentEvent]:
|
|
268
|
+
run = context.state.run
|
|
269
|
+
run.status = RunStatus.WAITING_FOR_TOOL
|
|
270
|
+
run.current_tool_call = call
|
|
271
|
+
run.usage.tool_calls += 1
|
|
272
|
+
yield self._event(context, AgentEventType.TOOL_STARTED, call=call)
|
|
273
|
+
try:
|
|
274
|
+
match tools.get(call.name):
|
|
275
|
+
case AgentTool() as registered:
|
|
276
|
+
pass
|
|
277
|
+
case None:
|
|
278
|
+
raise ToolDefinitionError(f"Unknown tool requested by model: {call.name}")
|
|
279
|
+
case invalid:
|
|
280
|
+
raise ToolDefinitionError(f"Invalid tool registration for {call.name}: {invalid!r}")
|
|
281
|
+
await self._call_extensions("before_tool", context, call)
|
|
282
|
+
value = await registered(call.arguments)
|
|
283
|
+
tool_message = ToolMessage(
|
|
284
|
+
data=ToolMessageData(
|
|
285
|
+
content=registered.serialize_result(value),
|
|
286
|
+
tool_call_id=call.id,
|
|
287
|
+
name=call.name,
|
|
288
|
+
),
|
|
289
|
+
metadata=ToolMessageMetadata(success=True),
|
|
290
|
+
)
|
|
291
|
+
except Exception as error:
|
|
292
|
+
await self._notify_error(context, error)
|
|
293
|
+
tool_message = ToolMessage(
|
|
294
|
+
data=ToolMessageData(
|
|
295
|
+
content=json.dumps({"error": str(error)}, ensure_ascii=False, separators=(",", ":")),
|
|
296
|
+
tool_call_id=call.id,
|
|
297
|
+
name=call.name,
|
|
298
|
+
),
|
|
299
|
+
metadata=ToolMessageMetadata(success=False),
|
|
300
|
+
)
|
|
301
|
+
await self._append_message(context, tool_message)
|
|
302
|
+
run.status = RunStatus.RUNNING
|
|
303
|
+
run.current_tool_call = None
|
|
304
|
+
yield self._event(
|
|
305
|
+
context,
|
|
306
|
+
AgentEventType.TOOL_FAILED,
|
|
307
|
+
call=call,
|
|
308
|
+
message=tool_message,
|
|
309
|
+
error=error,
|
|
310
|
+
)
|
|
311
|
+
await self._checkpoint(context, CheckpointReason.TOOL_FAILED)
|
|
312
|
+
return
|
|
313
|
+
|
|
314
|
+
await self._append_message(context, tool_message)
|
|
315
|
+
await self._call_extensions("after_tool", context, call, tool_message)
|
|
316
|
+
run.status = RunStatus.RUNNING
|
|
317
|
+
run.current_tool_call = None
|
|
318
|
+
yield self._event(
|
|
319
|
+
context,
|
|
320
|
+
AgentEventType.TOOL_COMPLETED,
|
|
321
|
+
call=call,
|
|
322
|
+
message=tool_message,
|
|
323
|
+
)
|
|
324
|
+
await self._checkpoint(context, CheckpointReason.TOOL_COMPLETED)
|
|
325
|
+
|
|
326
|
+
async def _append_message(self, context: AgentContext[SessionDataT], message: AnyMessage) -> None:
|
|
327
|
+
context.state.run.messages.append(message)
|
|
328
|
+
await self._call_extensions("on_message", context, message)
|
|
329
|
+
|
|
330
|
+
async def _checkpoint(self, context: AgentContext[SessionDataT], reason: CheckpointReason) -> None:
|
|
331
|
+
await self._call_extensions("on_checkpoint", context, reason)
|
|
332
|
+
|
|
333
|
+
async def _terminal_checkpoint(
|
|
334
|
+
self,
|
|
335
|
+
context: AgentContext[SessionDataT],
|
|
336
|
+
reason: CheckpointReason,
|
|
337
|
+
) -> None:
|
|
338
|
+
"""Attempt terminal persistence without replacing the original terminal outcome."""
|
|
339
|
+
try:
|
|
340
|
+
await self._checkpoint(context, reason)
|
|
341
|
+
except Exception:
|
|
342
|
+
pass
|
|
343
|
+
|
|
344
|
+
async def _release_state(self, context: AgentContext[SessionDataT]) -> None:
|
|
345
|
+
"""Best-effort cleanup that cannot invalidate an already emitted terminal event."""
|
|
346
|
+
for extension in self._extensions:
|
|
347
|
+
try:
|
|
348
|
+
await self._invoke_extension(extension, "release_state", context)
|
|
349
|
+
except Exception:
|
|
350
|
+
continue
|
|
351
|
+
|
|
352
|
+
async def _call_extensions(
|
|
353
|
+
self,
|
|
354
|
+
hook: str,
|
|
355
|
+
context: AgentContext[SessionDataT],
|
|
356
|
+
*arguments: Any,
|
|
357
|
+
) -> None:
|
|
358
|
+
for extension in self._extensions:
|
|
359
|
+
await self._invoke_extension(extension, hook, context, *arguments)
|
|
360
|
+
|
|
361
|
+
async def _invoke_extension(
|
|
362
|
+
self,
|
|
363
|
+
extension: AgentExtension[SessionDataT],
|
|
364
|
+
hook: str,
|
|
365
|
+
context: AgentContext[SessionDataT],
|
|
366
|
+
*arguments: Any,
|
|
367
|
+
) -> Any:
|
|
368
|
+
callback = getattr(extension, hook)
|
|
369
|
+
attempts = max(extension.max_retries, 0) + 1 if extension.failure_policy == ExtensionFailurePolicy.RETRY else 1
|
|
370
|
+
for attempt in range(attempts):
|
|
371
|
+
try:
|
|
372
|
+
return await callback(context, *arguments)
|
|
373
|
+
except Exception:
|
|
374
|
+
if extension.failure_policy == ExtensionFailurePolicy.CONTINUE:
|
|
375
|
+
return None
|
|
376
|
+
if extension.failure_policy != ExtensionFailurePolicy.RETRY or attempt == attempts - 1:
|
|
377
|
+
raise
|
|
378
|
+
return None
|
|
379
|
+
|
|
380
|
+
async def _notify_error(self, context: AgentContext[SessionDataT], error: Exception) -> None:
|
|
381
|
+
for extension in self._extensions:
|
|
382
|
+
try:
|
|
383
|
+
await self._invoke_extension(extension, "on_error", context, error)
|
|
384
|
+
except Exception:
|
|
385
|
+
continue
|
|
386
|
+
|
|
387
|
+
def _ensure_not_cancelled(self, run: RunState, token: CancellationToken) -> None:
|
|
388
|
+
if not token.cancelled:
|
|
389
|
+
return
|
|
390
|
+
run.cancellation_requested = True
|
|
391
|
+
run.status = RunStatus.CANCELLING
|
|
392
|
+
raise AgentCancelledError("Agent run was cancelled")
|
|
393
|
+
|
|
394
|
+
def _event(
|
|
395
|
+
self,
|
|
396
|
+
context: AgentContext[SessionDataT],
|
|
397
|
+
event_type: AgentEventType,
|
|
398
|
+
**values: Any,
|
|
399
|
+
) -> AgentEvent:
|
|
400
|
+
run = context.state.run
|
|
401
|
+
return AgentEvent(
|
|
402
|
+
type=event_type,
|
|
403
|
+
run_id=run.run_id,
|
|
404
|
+
sequence=run.next_event_sequence(),
|
|
405
|
+
timestamp=utc_now(),
|
|
406
|
+
status=run.status,
|
|
407
|
+
iteration=run.iteration,
|
|
408
|
+
**values,
|
|
409
|
+
)
|
kcs_agent/events.py
ADDED
|
@@ -0,0 +1,79 @@
|
|
|
1
|
+
from collections.abc import Mapping
|
|
2
|
+
from dataclasses import dataclass, field
|
|
3
|
+
from datetime import datetime
|
|
4
|
+
from enum import StrEnum
|
|
5
|
+
from typing import Any
|
|
6
|
+
|
|
7
|
+
from .messages import AssistantMessage, ToolCall, ToolMessage
|
|
8
|
+
from .model import ModelResponse, ToolCallDelta
|
|
9
|
+
from .state import AgentState, RunStatus
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class AgentEventType(StrEnum):
|
|
13
|
+
"""Observable lifecycle events emitted by the agent loop."""
|
|
14
|
+
|
|
15
|
+
RUN_STARTED = "run_started"
|
|
16
|
+
MODEL_STARTED = "model_started"
|
|
17
|
+
TEXT_DELTA = "text_delta"
|
|
18
|
+
REASONING_DELTA = "reasoning_delta"
|
|
19
|
+
TOOL_CALL_DELTA = "tool_call_delta"
|
|
20
|
+
MODEL_COMPLETED = "model_completed"
|
|
21
|
+
TOOL_STARTED = "tool_started"
|
|
22
|
+
TOOL_COMPLETED = "tool_completed"
|
|
23
|
+
TOOL_FAILED = "tool_failed"
|
|
24
|
+
RUN_COMPLETED = "run_completed"
|
|
25
|
+
RUN_FAILED = "run_failed"
|
|
26
|
+
RUN_CANCELLED = "run_cancelled"
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
@dataclass(frozen=True, slots=True)
|
|
30
|
+
class AgentResult[SessionDataT]:
|
|
31
|
+
"""Final state returned after an agent run completes."""
|
|
32
|
+
|
|
33
|
+
message: AssistantMessage
|
|
34
|
+
state: AgentState[SessionDataT]
|
|
35
|
+
metadata: Mapping[str, Any] = field(default_factory=dict)
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
@dataclass(frozen=True, slots=True)
|
|
39
|
+
class AgentEvent:
|
|
40
|
+
"""One ordered and observable event emitted by the agent loop."""
|
|
41
|
+
|
|
42
|
+
# Lifecycle discriminator that determines which optional payload fields are populated.
|
|
43
|
+
type: AgentEventType
|
|
44
|
+
|
|
45
|
+
# Stable identifier shared by every event produced by one agent invocation.
|
|
46
|
+
run_id: str
|
|
47
|
+
|
|
48
|
+
# One-based event number within a run, used to order events, detect gaps, and discard duplicates.
|
|
49
|
+
sequence: int
|
|
50
|
+
|
|
51
|
+
# Timezone-aware UTC timestamp captured when the core creates the event.
|
|
52
|
+
timestamp: datetime
|
|
53
|
+
|
|
54
|
+
# Run lifecycle state at emission time, such as waiting for a model, completed, or cancelled.
|
|
55
|
+
status: RunStatus
|
|
56
|
+
|
|
57
|
+
# One-based model-loop iteration, or zero before the first model step starts.
|
|
58
|
+
iteration: int = 0
|
|
59
|
+
|
|
60
|
+
# Incremental text for TEXT_DELTA or REASONING_DELTA, appended by consumers in sequence order.
|
|
61
|
+
delta: str = ""
|
|
62
|
+
|
|
63
|
+
# Partial model-generated tool ID, name, or JSON arguments; this fragment is not safe to execute.
|
|
64
|
+
tool_call_delta: ToolCallDelta | None = None
|
|
65
|
+
|
|
66
|
+
# Complete normalized invocation used by TOOL_STARTED, TOOL_COMPLETED, and TOOL_FAILED events.
|
|
67
|
+
call: ToolCall | None = None
|
|
68
|
+
|
|
69
|
+
# Authoritative normalized model response and usage, populated by MODEL_COMPLETED.
|
|
70
|
+
response: ModelResponse | None = None
|
|
71
|
+
|
|
72
|
+
# Tool result populated by TOOL_COMPLETED or TOOL_FAILED; model output is available through response.
|
|
73
|
+
message: AssistantMessage | ToolMessage | None = None
|
|
74
|
+
|
|
75
|
+
# Successful terminal output and typed state, populated only by RUN_COMPLETED.
|
|
76
|
+
result: AgentResult[Any] | None = None
|
|
77
|
+
|
|
78
|
+
# In-process failure or cancellation cause; transport adapters must serialize it safely.
|
|
79
|
+
error: Exception | None = None
|
kcs_agent/exceptions.py
ADDED
|
@@ -0,0 +1,26 @@
|
|
|
1
|
+
class AgentError(Exception):
|
|
2
|
+
"""Base exception for agent runtime failures."""
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
class AgentProtocolError(AgentError):
|
|
6
|
+
"""Raised when a model adapter violates the streaming protocol."""
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class AgentIterationLimitError(AgentError):
|
|
10
|
+
"""Raised when the model does not finish within the configured loop bound."""
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class AgentCancelledError(AgentError):
|
|
14
|
+
"""Raised when a caller cooperatively cancels an active run."""
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class SessionVersionConflictError(AgentError):
|
|
18
|
+
"""Raised when durable session state changed since it was loaded."""
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class ToolDefinitionError(AgentError):
|
|
22
|
+
"""Raised when a decorated function cannot become a model tool."""
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class ToolInvocationError(AgentError):
|
|
26
|
+
"""Raised when tool arguments are invalid or execution fails."""
|