tyrion-cli 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.
- tyrion_agent/__init__.py +0 -0
- tyrion_agent/events.py +89 -0
- tyrion_agent/harness.py +273 -0
- tyrion_agent/loop.py +273 -0
- tyrion_agent/messages.py +119 -0
- tyrion_agent/provider.py +32 -0
- tyrion_agent/provider_events.py +70 -0
- tyrion_agent/py.typed +0 -0
- tyrion_agent/sessions/__init__.py +28 -0
- tyrion_agent/sessions/entries.py +68 -0
- tyrion_agent/sessions/jsonl.py +67 -0
- tyrion_agent/sessions/tree.py +132 -0
- tyrion_agent/tool_history.py +59 -0
- tyrion_agent/tools.py +84 -0
- tyrion_agent/types.py +7 -0
- tyrion_ai/__init__.py +0 -0
- tyrion_ai/env.py +38 -0
- tyrion_ai/fake.py +50 -0
- tyrion_ai/http.py +21 -0
- tyrion_ai/model_limits.py +152 -0
- tyrion_ai/openai_compatible.py +432 -0
- tyrion_ai/py.typed +0 -0
- tyrion_cli-0.1.0.dist-info/METADATA +247 -0
- tyrion_cli-0.1.0.dist-info/RECORD +53 -0
- tyrion_cli-0.1.0.dist-info/WHEEL +4 -0
- tyrion_cli-0.1.0.dist-info/entry_points.txt +2 -0
- tyrion_cli-0.1.0.dist-info/licenses/LICENSE +21 -0
- tyrion_coding/__init__.py +0 -0
- tyrion_coding/cli.py +166 -0
- tyrion_coding/commands.py +173 -0
- tyrion_coding/context_window.py +104 -0
- tyrion_coding/data/__init__.py +0 -0
- tyrion_coding/data/system_prompt.md +19 -0
- tyrion_coding/display.py +106 -0
- tyrion_coding/prompt_templates.py +41 -0
- tyrion_coding/provider_catalog.py +76 -0
- tyrion_coding/provider_config.py +161 -0
- tyrion_coding/py.typed +0 -0
- tyrion_coding/rendering.py +86 -0
- tyrion_coding/resources.py +36 -0
- tyrion_coding/session.py +258 -0
- tyrion_coding/session_coding.py +107 -0
- tyrion_coding/skills.py +71 -0
- tyrion_coding/system_prompt.py +42 -0
- tyrion_coding/theme.py +128 -0
- tyrion_coding/tools.py +529 -0
- tyrion_coding/tui/__init__.py +6 -0
- tyrion_coding/tui/app.py +411 -0
- tyrion_coding/tui/connect_modal.py +158 -0
- tyrion_coding/tui/picker.py +150 -0
- tyrion_coding/tui/styles.py +222 -0
- tyrion_coding/tui/welcome.py +89 -0
- tyrion_coding/tui/widgets.py +474 -0
tyrion_agent/__init__.py
ADDED
|
File without changes
|
tyrion_agent/events.py
ADDED
|
@@ -0,0 +1,89 @@
|
|
|
1
|
+
"""Events emitted by the portable agent layer."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
from typing import Annotated, Literal
|
|
5
|
+
|
|
6
|
+
from pydantic import Field
|
|
7
|
+
|
|
8
|
+
from tyrion_agent.messages import AgentMessage, ToolResultMessage, WireModel
|
|
9
|
+
from tyrion_agent.provider_events import AssistantMessageEvent
|
|
10
|
+
from tyrion_agent.tools import AgentToolResult
|
|
11
|
+
from tyrion_agent.types import JSONValue
|
|
12
|
+
|
|
13
|
+
class AgentStartEvent(WireModel):
|
|
14
|
+
"""The agent has started processing."""
|
|
15
|
+
type: Literal["agent_start"] = "agent_start"
|
|
16
|
+
|
|
17
|
+
class AgentEndEvent(WireModel):
|
|
18
|
+
"""The agent has finished processing."""
|
|
19
|
+
|
|
20
|
+
type: Literal["agent_end"] = "agent_end"
|
|
21
|
+
messages: list[AgentMessage] = Field(default_factory=list)
|
|
22
|
+
|
|
23
|
+
class TurnStartEvent(WireModel):
|
|
24
|
+
"""A new model turn has started."""
|
|
25
|
+
|
|
26
|
+
type: Literal["turn_start"] = "turn_start"
|
|
27
|
+
|
|
28
|
+
class TurnEndEvent(WireModel):
|
|
29
|
+
"""A model turn has completed."""
|
|
30
|
+
|
|
31
|
+
type: Literal["turn_end"] = "turn_end"
|
|
32
|
+
message: AgentMessage
|
|
33
|
+
tool_results: list[ToolResultMessage] = Field(default_factory=list)
|
|
34
|
+
|
|
35
|
+
class MessageStartEvent(WireModel):
|
|
36
|
+
"""A message has started (user, assistant or tool result)."""
|
|
37
|
+
type: Literal["message_start"] = "message_start"
|
|
38
|
+
message: AgentMessage
|
|
39
|
+
|
|
40
|
+
class MessageUpdateEvent(WireModel):
|
|
41
|
+
"""A streaming delta has arrived for the current message."""
|
|
42
|
+
|
|
43
|
+
type: Literal["message_update"] = "message_update"
|
|
44
|
+
message: AgentMessage
|
|
45
|
+
assistant_message_event: AssistantMessageEvent
|
|
46
|
+
|
|
47
|
+
class MessageEndEvent(WireModel):
|
|
48
|
+
"""A message has been finalized."""
|
|
49
|
+
type: Literal["message_end"] = "message_end"
|
|
50
|
+
message: AgentMessage
|
|
51
|
+
|
|
52
|
+
class ToolExecutionStartEvent(WireModel):
|
|
53
|
+
"""A tool has started executing."""
|
|
54
|
+
|
|
55
|
+
type: Literal["tool_execution_start"] = "tool_execution_start"
|
|
56
|
+
tool_call_id: str
|
|
57
|
+
tool_name : str
|
|
58
|
+
args: dict[str, JSONValue] = Field(default_factory=dict)
|
|
59
|
+
|
|
60
|
+
class ToolExecutionUpdateEvent(WireModel):
|
|
61
|
+
"""A tool has started executing."""
|
|
62
|
+
|
|
63
|
+
type: Literal["tool_execution_update"] = "tool_execution_update"
|
|
64
|
+
tool_call_id: str
|
|
65
|
+
tool_name : str
|
|
66
|
+
args: dict[str, JSONValue] = Field(default_factory=dict)
|
|
67
|
+
partial_result: AgentToolResult
|
|
68
|
+
|
|
69
|
+
class ToolExecutionEndEvent(WireModel):
|
|
70
|
+
"""A tool has finished executing."""
|
|
71
|
+
type: Literal["tool_execution_end"] = "tool_execution_end"
|
|
72
|
+
tool_call_id: str
|
|
73
|
+
tool_name: str
|
|
74
|
+
result: AgentToolResult
|
|
75
|
+
is_error: bool
|
|
76
|
+
# Union of all agent events
|
|
77
|
+
type AgentEvent = Annotated[
|
|
78
|
+
AgentStartEvent
|
|
79
|
+
| AgentEndEvent
|
|
80
|
+
| TurnStartEvent
|
|
81
|
+
| TurnEndEvent
|
|
82
|
+
| MessageStartEvent
|
|
83
|
+
| MessageUpdateEvent
|
|
84
|
+
| MessageEndEvent
|
|
85
|
+
| ToolExecutionStartEvent
|
|
86
|
+
| ToolExecutionUpdateEvent
|
|
87
|
+
| ToolExecutionEndEvent,
|
|
88
|
+
Field(discriminator="type"),
|
|
89
|
+
]
|
tyrion_agent/harness.py
ADDED
|
@@ -0,0 +1,273 @@
|
|
|
1
|
+
"""Stateful reusable agent brain built on top of the pure agent loop."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import AsyncIterator, Awaitable, Callable, Sequence
|
|
6
|
+
from contextlib import suppress
|
|
7
|
+
from dataclasses import dataclass, field
|
|
8
|
+
from inspect import isawaitable
|
|
9
|
+
|
|
10
|
+
from tyrion_agent.events import AgentEvent, MessageEndEvent, MessageStartEvent
|
|
11
|
+
from tyrion_agent.loop import run_agent_loop
|
|
12
|
+
from tyrion_agent.messages import (
|
|
13
|
+
AgentMessage,
|
|
14
|
+
AssistantMessage,
|
|
15
|
+
TextContent,
|
|
16
|
+
ToolResultMessage,
|
|
17
|
+
UserMessage,
|
|
18
|
+
)
|
|
19
|
+
from tyrion_agent.provider import ModelProvider
|
|
20
|
+
from tyrion_agent.tools import AgentTool
|
|
21
|
+
|
|
22
|
+
# A listener is any function that receives an event (sync or async)
|
|
23
|
+
EventListener = Callable[[AgentEvent], Awaitable[None] | None]
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
# ---------------------------------------------------------------------------
|
|
27
|
+
# Configuration
|
|
28
|
+
# ---------------------------------------------------------------------------
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
@dataclass(slots=True)
|
|
32
|
+
class AgentHarnessConfig:
|
|
33
|
+
"""Everything the harness needs to run the loop."""
|
|
34
|
+
|
|
35
|
+
provider: ModelProvider
|
|
36
|
+
model: str
|
|
37
|
+
system: str
|
|
38
|
+
tools: list[AgentTool] = field(default_factory=list)
|
|
39
|
+
max_turns: int | None = None
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
# ---------------------------------------------------------------------------
|
|
43
|
+
# Cancellation
|
|
44
|
+
# ---------------------------------------------------------------------------
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
class SimpleCancellationToken:
|
|
48
|
+
"""A simple flag that can be checked to see if work should stop."""
|
|
49
|
+
|
|
50
|
+
def __init__(self) -> None:
|
|
51
|
+
self._cancelled = False
|
|
52
|
+
|
|
53
|
+
def cancel(self) -> None:
|
|
54
|
+
self._cancelled = True
|
|
55
|
+
|
|
56
|
+
def is_cancelled(self) -> bool:
|
|
57
|
+
return self._cancelled
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
# ---------------------------------------------------------------------------
|
|
61
|
+
# The Harness
|
|
62
|
+
# ---------------------------------------------------------------------------
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
class AgentHarness:
|
|
66
|
+
"""Reusable stateful agent brain.
|
|
67
|
+
|
|
68
|
+
Owns the transcript, delegates to the pure agent loop,
|
|
69
|
+
and provides a clean API for frontends.
|
|
70
|
+
|
|
71
|
+
Usage:
|
|
72
|
+
harness = AgentHarness(AgentHarnessConfig(
|
|
73
|
+
provider=provider,
|
|
74
|
+
model="gpt-4",
|
|
75
|
+
system="You are a coding agent.",
|
|
76
|
+
tools=[read_tool, bash_tool],
|
|
77
|
+
))
|
|
78
|
+
|
|
79
|
+
async for event in harness.prompt("Read README.md"):
|
|
80
|
+
print(event)
|
|
81
|
+
"""
|
|
82
|
+
|
|
83
|
+
def __init__(
|
|
84
|
+
self,
|
|
85
|
+
config: AgentHarnessConfig,
|
|
86
|
+
*,
|
|
87
|
+
messages: Sequence[AgentMessage] = (),
|
|
88
|
+
) -> None:
|
|
89
|
+
self._config = config
|
|
90
|
+
self._messages: list[AgentMessage] = list(messages)
|
|
91
|
+
self._listeners: list[EventListener] = []
|
|
92
|
+
self._current_signal: SimpleCancellationToken | None = None
|
|
93
|
+
self._running = False
|
|
94
|
+
|
|
95
|
+
# --- Public properties ---
|
|
96
|
+
|
|
97
|
+
@property
|
|
98
|
+
def messages(self) -> tuple[AgentMessage, ...]:
|
|
99
|
+
"""Immutable snapshot of the current transcript."""
|
|
100
|
+
return tuple(self._messages)
|
|
101
|
+
|
|
102
|
+
@property
|
|
103
|
+
def config(self) -> AgentHarnessConfig:
|
|
104
|
+
return self._config
|
|
105
|
+
|
|
106
|
+
@property
|
|
107
|
+
def is_running(self) -> bool:
|
|
108
|
+
return self._running
|
|
109
|
+
|
|
110
|
+
# --- Core API ---
|
|
111
|
+
|
|
112
|
+
def prompt(self, content: str) -> AsyncIterator[AgentEvent]:
|
|
113
|
+
"""Send a user message and run the agent loop.
|
|
114
|
+
|
|
115
|
+
This is the main entry point. It:
|
|
116
|
+
1. Appends a UserMessage to the transcript
|
|
117
|
+
2. Runs the agent loop
|
|
118
|
+
3. Yields events as they happen
|
|
119
|
+
"""
|
|
120
|
+
return self.prompt_message(UserMessage(content=content))
|
|
121
|
+
|
|
122
|
+
def prompt_message(self, message: AgentMessage) -> AsyncIterator[AgentEvent]:
|
|
123
|
+
"""Send an arbitrary message and run the agent loop."""
|
|
124
|
+
self._ensure_not_running()
|
|
125
|
+
self._running = True
|
|
126
|
+
return self._run(prompts=(message,))
|
|
127
|
+
|
|
128
|
+
def continue_(self) -> AsyncIterator[AgentEvent]:
|
|
129
|
+
"""Run the agent loop without appending a new message.
|
|
130
|
+
|
|
131
|
+
Useful for:
|
|
132
|
+
- Resuming after restoring a session
|
|
133
|
+
- Continuing after an interrupted run
|
|
134
|
+
"""
|
|
135
|
+
self._ensure_not_running()
|
|
136
|
+
self._running = True
|
|
137
|
+
return self._run()
|
|
138
|
+
|
|
139
|
+
# --- Cancellation ---
|
|
140
|
+
|
|
141
|
+
def cancel(self) -> None:
|
|
142
|
+
"""Cancel the current run.
|
|
143
|
+
|
|
144
|
+
The loop will stop after the current tool finishes.
|
|
145
|
+
Any unanswered tool calls get synthetic error results.
|
|
146
|
+
"""
|
|
147
|
+
if self._current_signal is not None:
|
|
148
|
+
self._current_signal.cancel()
|
|
149
|
+
|
|
150
|
+
# --- Event listeners ---
|
|
151
|
+
|
|
152
|
+
def subscribe(self, listener: EventListener) -> Callable[[], None]:
|
|
153
|
+
"""Subscribe to agent events. Returns an unsubscribe function.
|
|
154
|
+
|
|
155
|
+
Listeners receive the same events as the prompt()/continue_()
|
|
156
|
+
consumer. This lets persistence, logging, and other observers
|
|
157
|
+
watch runs without being the main consumer.
|
|
158
|
+
|
|
159
|
+
Usage:
|
|
160
|
+
def on_event(event):
|
|
161
|
+
print(f"Got: {event.type}")
|
|
162
|
+
|
|
163
|
+
unsubscribe = harness.subscribe(on_event)
|
|
164
|
+
# ... later ...
|
|
165
|
+
unsubscribe()
|
|
166
|
+
"""
|
|
167
|
+
self._listeners.append(listener)
|
|
168
|
+
|
|
169
|
+
def unsubscribe() -> None:
|
|
170
|
+
with suppress(ValueError):
|
|
171
|
+
self._listeners.remove(listener)
|
|
172
|
+
|
|
173
|
+
return unsubscribe
|
|
174
|
+
|
|
175
|
+
# --- Transcript management ---
|
|
176
|
+
|
|
177
|
+
def append_message(self, message: AgentMessage) -> None:
|
|
178
|
+
"""Manually append a message (used when restoring sessions)."""
|
|
179
|
+
self._messages.append(message)
|
|
180
|
+
|
|
181
|
+
def replace_messages(self, messages: Sequence[AgentMessage]) -> None:
|
|
182
|
+
"""Replace the entire transcript (used when restoring sessions)."""
|
|
183
|
+
self._messages = list(messages)
|
|
184
|
+
|
|
185
|
+
# --- Internal machinery ---
|
|
186
|
+
|
|
187
|
+
async def _run(
|
|
188
|
+
self,
|
|
189
|
+
*,
|
|
190
|
+
prompts: Sequence[AgentMessage] = (),
|
|
191
|
+
) -> AsyncIterator[AgentEvent]:
|
|
192
|
+
"""Run the agent loop and handle cleanup."""
|
|
193
|
+
signal = SimpleCancellationToken()
|
|
194
|
+
self._current_signal = signal
|
|
195
|
+
|
|
196
|
+
try:
|
|
197
|
+
# Repair any dangling tool calls from a previous interrupted run
|
|
198
|
+
self._repair_interrupted_tools()
|
|
199
|
+
|
|
200
|
+
# Run the loop
|
|
201
|
+
async for event in run_agent_loop(
|
|
202
|
+
provider=self._config.provider,
|
|
203
|
+
model=self._config.model,
|
|
204
|
+
system=self._config.system,
|
|
205
|
+
messages=self._messages,
|
|
206
|
+
prompts=prompts,
|
|
207
|
+
tools=self._config.tools,
|
|
208
|
+
max_turns=self._config.max_turns,
|
|
209
|
+
signal=signal,
|
|
210
|
+
):
|
|
211
|
+
# Notify all listeners
|
|
212
|
+
await self._notify(event)
|
|
213
|
+
# Yield to the consumer (CLI, TUI, etc.)
|
|
214
|
+
yield event
|
|
215
|
+
|
|
216
|
+
finally:
|
|
217
|
+
# If we were cancelled, repair any new dangling tool calls
|
|
218
|
+
if signal.is_cancelled():
|
|
219
|
+
before = len(self._messages)
|
|
220
|
+
self._repair_interrupted_tools()
|
|
221
|
+
# Push repair messages to listeners so persistence sees them
|
|
222
|
+
for message in self._messages[before:]:
|
|
223
|
+
with suppress(Exception):
|
|
224
|
+
await self._notify(MessageStartEvent(message=message))
|
|
225
|
+
await self._notify(MessageEndEvent(message=message))
|
|
226
|
+
|
|
227
|
+
if self._current_signal is signal:
|
|
228
|
+
self._current_signal = None
|
|
229
|
+
self._running = False
|
|
230
|
+
|
|
231
|
+
async def _notify(self, event: AgentEvent) -> None:
|
|
232
|
+
"""Send an event to all listeners."""
|
|
233
|
+
for listener in list(self._listeners):
|
|
234
|
+
result = listener(event)
|
|
235
|
+
if isawaitable(result):
|
|
236
|
+
await result
|
|
237
|
+
|
|
238
|
+
def _ensure_not_running(self) -> None:
|
|
239
|
+
"""Prevent overlapping runs."""
|
|
240
|
+
if self._running:
|
|
241
|
+
raise RuntimeError(
|
|
242
|
+
"AgentHarness is already running. "
|
|
243
|
+
"Wait for the current run to finish before starting a new one."
|
|
244
|
+
)
|
|
245
|
+
|
|
246
|
+
def _repair_interrupted_tools(self) -> None:
|
|
247
|
+
"""Append synthetic error results for any unanswered tool calls.
|
|
248
|
+
|
|
249
|
+
This happens when a run is cancelled mid-tool-execution.
|
|
250
|
+
The model asked to call a tool, but we stopped before getting
|
|
251
|
+
the result. We need to add a "Tool call interrupted" result
|
|
252
|
+
so the transcript stays valid.
|
|
253
|
+
"""
|
|
254
|
+
returned_ids: set[str] = {
|
|
255
|
+
msg.tool_call_id
|
|
256
|
+
for msg in self._messages
|
|
257
|
+
if isinstance(msg, ToolResultMessage)
|
|
258
|
+
}
|
|
259
|
+
|
|
260
|
+
for msg in tuple(self._messages):
|
|
261
|
+
if not isinstance(msg, AssistantMessage):
|
|
262
|
+
continue
|
|
263
|
+
for call in msg.tool_calls:
|
|
264
|
+
if call.id not in returned_ids:
|
|
265
|
+
returned_ids.add(call.id)
|
|
266
|
+
self._messages.append(
|
|
267
|
+
ToolResultMessage(
|
|
268
|
+
tool_call_id=call.id,
|
|
269
|
+
tool_name=call.name,
|
|
270
|
+
content=[TextContent(text="Tool call interrupted by user")],
|
|
271
|
+
is_error=True,
|
|
272
|
+
)
|
|
273
|
+
)
|
tyrion_agent/loop.py
ADDED
|
@@ -0,0 +1,273 @@
|
|
|
1
|
+
"""the pure agent loop - the engine that drives model <-> tool interaction."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
from collections.abc import AsyncIterator, Mapping, Sequence
|
|
7
|
+
|
|
8
|
+
from tyrion_agent.events import (
|
|
9
|
+
AgentEndEvent,
|
|
10
|
+
AgentEvent,
|
|
11
|
+
AgentStartEvent,
|
|
12
|
+
MessageEndEvent,
|
|
13
|
+
MessageStartEvent,
|
|
14
|
+
MessageUpdateEvent,
|
|
15
|
+
ToolExecutionEndEvent,
|
|
16
|
+
ToolExecutionStartEvent,
|
|
17
|
+
TurnEndEvent,
|
|
18
|
+
TurnStartEvent,
|
|
19
|
+
)
|
|
20
|
+
from tyrion_agent.messages import (
|
|
21
|
+
AgentMessage,
|
|
22
|
+
AssistantMessage,
|
|
23
|
+
TextContent,
|
|
24
|
+
ToolCall,
|
|
25
|
+
ToolResultMessage,
|
|
26
|
+
)
|
|
27
|
+
from tyrion_agent.provider import CancellationToken, ModelProvider
|
|
28
|
+
from tyrion_agent.provider_events import (
|
|
29
|
+
AssistantErrorEvent,
|
|
30
|
+
AssistantMessageEvent,
|
|
31
|
+
AssistantStartEvent,
|
|
32
|
+
AssistantDoneEvent
|
|
33
|
+
)
|
|
34
|
+
from tyrion_agent.tool_history import repair_tool_history
|
|
35
|
+
from tyrion_agent.tools import AgentTool, AgentToolResult
|
|
36
|
+
|
|
37
|
+
async def run_agent_loop(
|
|
38
|
+
*,
|
|
39
|
+
provider: ModelProvider,
|
|
40
|
+
model: str,
|
|
41
|
+
system: str,
|
|
42
|
+
messages :list[AgentMessage],
|
|
43
|
+
tools : list[AgentTool],
|
|
44
|
+
prompts: Sequence[AgentMessage] =(),
|
|
45
|
+
max_turns: int | None = None,
|
|
46
|
+
signal: CancellationToken | None = None,
|
|
47
|
+
) -> AsyncIterator[AgentEvent]:
|
|
48
|
+
"""Run the provider/tool loop and emit the agent events.
|
|
49
|
+
Args:
|
|
50
|
+
provider: The model provider to stream responses from.
|
|
51
|
+
model: Model name to use.
|
|
52
|
+
system: System prompt text.
|
|
53
|
+
messages: The mutable transcript. The loop APPENDS to this list.
|
|
54
|
+
tools: Available tools the model can call.
|
|
55
|
+
prompts: Initial messages to append (e.g. a new user message).
|
|
56
|
+
max_turns: Optional safety cap on model turns.
|
|
57
|
+
signal: Optional cancellation token.
|
|
58
|
+
"""
|
|
59
|
+
new_messages: list[AgentMessage] = list(prompts)
|
|
60
|
+
if prompts:
|
|
61
|
+
messages.extend(prompts)
|
|
62
|
+
|
|
63
|
+
tool_by_name = {tool.name: tool for tool in tools}
|
|
64
|
+
turn =1
|
|
65
|
+
|
|
66
|
+
## Emit start events
|
|
67
|
+
yield AgentStartEvent()
|
|
68
|
+
yield TurnStartEvent()
|
|
69
|
+
|
|
70
|
+
# Emit events for initial prompts messages
|
|
71
|
+
for prompt in prompts:
|
|
72
|
+
yield MessageStartEvent(message = prompt)
|
|
73
|
+
yield MessageEndEvent(message = prompt)
|
|
74
|
+
|
|
75
|
+
#Validate max_turns
|
|
76
|
+
if max_turns is not None and max_turns < 1:
|
|
77
|
+
error = _error_message(model, "max turns must be atleat 1")
|
|
78
|
+
messages.append(error)
|
|
79
|
+
new_messages.append(error)
|
|
80
|
+
yield MessageStartEvent(message=error)
|
|
81
|
+
yield MessageEndEvent(message=error)
|
|
82
|
+
yield TurnEndEvent(message=error)
|
|
83
|
+
yield AgentEndEvent(messages=new_messages)
|
|
84
|
+
return
|
|
85
|
+
|
|
86
|
+
## --- Main Loop ---
|
|
87
|
+
while True:
|
|
88
|
+
# check max turns
|
|
89
|
+
if max_turns is not None and turn > max_turns:
|
|
90
|
+
error = _error_message(model, f"Reached the max turns limit ({max_turns})")
|
|
91
|
+
messages.append(error)
|
|
92
|
+
new_messages.append(error)
|
|
93
|
+
yield MessageStartEvent(message=error)
|
|
94
|
+
yield MessageEndEvent(message=error)
|
|
95
|
+
yield TurnEndEvent(message=error)
|
|
96
|
+
yield AgentEndEvent(messages=new_messages)
|
|
97
|
+
return
|
|
98
|
+
assistant : AssistantMessage | None = None
|
|
99
|
+
|
|
100
|
+
async for event in _stream_assistant(
|
|
101
|
+
provider = provider,
|
|
102
|
+
model = model,
|
|
103
|
+
system = system,
|
|
104
|
+
messages = messages,
|
|
105
|
+
tools =tools,
|
|
106
|
+
signal= signal,
|
|
107
|
+
):
|
|
108
|
+
yield event
|
|
109
|
+
# capture the final assistant message
|
|
110
|
+
|
|
111
|
+
if isinstance(event, MessageEndEvent) and isinstance(event.message, AssistantMessage):
|
|
112
|
+
assistant = event.message
|
|
113
|
+
|
|
114
|
+
# The provider should always produce a message, but incase it doesn't
|
|
115
|
+
if assistant is None:
|
|
116
|
+
assistant = _error_message(model, "Provider produced no assistant message")
|
|
117
|
+
yield MessageStartEvent(message=assistant)
|
|
118
|
+
yield MessageEndEvent(message=assistant)
|
|
119
|
+
|
|
120
|
+
## Add assistant message to transcript
|
|
121
|
+
messages.append(assistant)
|
|
122
|
+
new_messages.append(assistant)
|
|
123
|
+
|
|
124
|
+
if assistant.stop_reason in {"error", "aborted"}:
|
|
125
|
+
yield TurnEndEvent(message=assistant)
|
|
126
|
+
yield AgentEndEvent(messages=new_messages)
|
|
127
|
+
return
|
|
128
|
+
## Execute tool calls (if any)
|
|
129
|
+
calls = list(assistant.tool_calls)
|
|
130
|
+
tool_results: list[ToolResultMessage] =[]
|
|
131
|
+
|
|
132
|
+
for call in calls:
|
|
133
|
+
async for event in _execute_tool_call(call, tool_by_name, signal):
|
|
134
|
+
yield event
|
|
135
|
+
|
|
136
|
+
if isinstance(event, MessageEndEvent) and isinstance(event.message, ToolResultMessage):
|
|
137
|
+
tool_results.append(event.message)
|
|
138
|
+
messages.append(event.message)
|
|
139
|
+
new_messages.append(event.message)
|
|
140
|
+
|
|
141
|
+
# Emit turn end
|
|
142
|
+
yield TurnEndEvent(message = assistant, tool_results=tool_results)
|
|
143
|
+
turn +=1
|
|
144
|
+
|
|
145
|
+
if not calls:
|
|
146
|
+
break
|
|
147
|
+
|
|
148
|
+
# Cancelled while tools were running: stop here instead of asking the
|
|
149
|
+
# model for another turn.
|
|
150
|
+
if signal is not None and signal.is_cancelled():
|
|
151
|
+
yield AgentEndEvent(messages=new_messages)
|
|
152
|
+
return
|
|
153
|
+
yield TurnStartEvent()
|
|
154
|
+
yield AgentEndEvent(messages=new_messages)
|
|
155
|
+
|
|
156
|
+
|
|
157
|
+
### Stream assistant response from porvider
|
|
158
|
+
|
|
159
|
+
async def _stream_assistant(
|
|
160
|
+
*,
|
|
161
|
+
provider: ModelProvider,
|
|
162
|
+
model: str,
|
|
163
|
+
system : str,
|
|
164
|
+
messages: list[AgentMessage],
|
|
165
|
+
tools: list[AgentTool],
|
|
166
|
+
signal: CancellationToken | None = None,
|
|
167
|
+
) -> AsyncIterator[AgentEvent]:
|
|
168
|
+
"""Strema one model's response and convert provider events to agent events."""
|
|
169
|
+
|
|
170
|
+
clean = repair_tool_history(messages)
|
|
171
|
+
|
|
172
|
+
source : AsyncIterator[AssistantMessageEvent] = provider.stream_response(
|
|
173
|
+
model = model,
|
|
174
|
+
system = system,
|
|
175
|
+
messages = clean.messages,
|
|
176
|
+
tools=tools,
|
|
177
|
+
signal = signal,
|
|
178
|
+
)
|
|
179
|
+
started = False
|
|
180
|
+
async for event in source:
|
|
181
|
+
if isinstance(event,AssistantStartEvent):
|
|
182
|
+
started = True
|
|
183
|
+
yield MessageStartEvent(message = event.partial)
|
|
184
|
+
|
|
185
|
+
elif isinstance(event, AssistantDoneEvent):
|
|
186
|
+
if not started:
|
|
187
|
+
yield MessageStartEvent(message = event.message)
|
|
188
|
+
yield MessageEndEvent(message = event.message)
|
|
189
|
+
|
|
190
|
+
elif isinstance(event, AssistantErrorEvent):
|
|
191
|
+
if not started:
|
|
192
|
+
yield MessageStartEvent(message=event.error)
|
|
193
|
+
yield MessageEndEvent(message=event.error)
|
|
194
|
+
else :
|
|
195
|
+
yield MessageUpdateEvent(
|
|
196
|
+
message = event.partial,
|
|
197
|
+
assistant_message_event=event,
|
|
198
|
+
)
|
|
199
|
+
|
|
200
|
+
async def _execute_tool_call(
|
|
201
|
+
call: ToolCall,
|
|
202
|
+
tools: Mapping[str, AgentTool],
|
|
203
|
+
signal: CancellationToken | None,
|
|
204
|
+
) -> AsyncIterator[AgentEvent]:
|
|
205
|
+
"""Execute one tool call and emit events."""
|
|
206
|
+
|
|
207
|
+
# Announce the tool execution
|
|
208
|
+
yield ToolExecutionStartEvent(
|
|
209
|
+
tool_call_id = call.id,
|
|
210
|
+
tool_name = call.name,
|
|
211
|
+
args = call.arguments,
|
|
212
|
+
)
|
|
213
|
+
|
|
214
|
+
#check if cancelled
|
|
215
|
+
if signal is not None and signal.is_cancelled():
|
|
216
|
+
result = _error_result("Operation aborted")
|
|
217
|
+
is_error= True
|
|
218
|
+
|
|
219
|
+
else:
|
|
220
|
+
## Find the tool
|
|
221
|
+
tool = tools.get(call.name)
|
|
222
|
+
if tool is None:
|
|
223
|
+
result = _error_result(f"Unknown Tool: {call.name}")
|
|
224
|
+
is_error = True
|
|
225
|
+
else:
|
|
226
|
+
result, is_error = await _run_tool(tool, call, signal)
|
|
227
|
+
|
|
228
|
+
## Announce the result
|
|
229
|
+
yield ToolExecutionEndEvent(
|
|
230
|
+
tool_call_id = call.id,
|
|
231
|
+
tool_name = call.name,
|
|
232
|
+
result = result,
|
|
233
|
+
is_error = is_error
|
|
234
|
+
)
|
|
235
|
+
|
|
236
|
+
# Emit tool result message
|
|
237
|
+
message = ToolResultMessage(
|
|
238
|
+
tool_call_id=call.id,
|
|
239
|
+
tool_name=call.name,
|
|
240
|
+
content=result.content,
|
|
241
|
+
is_error=is_error,
|
|
242
|
+
)
|
|
243
|
+
yield MessageStartEvent(message=message)
|
|
244
|
+
yield MessageEndEvent(message=message)
|
|
245
|
+
|
|
246
|
+
async def _run_tool(
|
|
247
|
+
tool: AgentTool,
|
|
248
|
+
call: ToolCall,
|
|
249
|
+
signal: CancellationToken | None,
|
|
250
|
+
) -> tuple[AgentToolResult, bool]:
|
|
251
|
+
"""Run a tool catching any exceptions"""
|
|
252
|
+
try:
|
|
253
|
+
result = await tool.execute(call.id, call.arguments, signal)
|
|
254
|
+
return result, False
|
|
255
|
+
except asyncio.CancelledError:
|
|
256
|
+
raise
|
|
257
|
+
except Exception as exc:
|
|
258
|
+
return _error_result(str(exc)), True
|
|
259
|
+
|
|
260
|
+
## Helpers
|
|
261
|
+
def _error_result(message: str) -> AgentToolResult:
|
|
262
|
+
"""Create an error tool result."""
|
|
263
|
+
return AgentToolResult(content=[TextContent(text=message)])
|
|
264
|
+
|
|
265
|
+
|
|
266
|
+
def _error_message(model: str, message: str) -> AssistantMessage:
|
|
267
|
+
"""Create an error assistant message."""
|
|
268
|
+
return AssistantMessage(
|
|
269
|
+
model=model,
|
|
270
|
+
content=[],
|
|
271
|
+
stop_reason="error",
|
|
272
|
+
error_message=message,
|
|
273
|
+
)
|