pi-python-core 0.8.1__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.
- pi_python/__init__.py +160 -0
- pi_python/_version.py +1 -0
- pi_python/agent.py +396 -0
- pi_python/cancellation.py +24 -0
- pi_python/data/models.json +3315 -0
- pi_python/errors.py +49 -0
- pi_python/estimate.py +144 -0
- pi_python/events.py +138 -0
- pi_python/function_tools.py +438 -0
- pi_python/hooks.py +44 -0
- pi_python/limits.py +28 -0
- pi_python/loop.py +431 -0
- pi_python/lowlevel.py +179 -0
- pi_python/mcp.py +187 -0
- pi_python/messages.py +405 -0
- pi_python/models.py +155 -0
- pi_python/provider.py +123 -0
- pi_python/providers/__init__.py +21 -0
- pi_python/providers/anthropic.py +673 -0
- pi_python/providers/common.py +201 -0
- pi_python/providers/completions.py +1149 -0
- pi_python/providers/oauth.py +542 -0
- pi_python/providers/openai.py +681 -0
- pi_python/providers/transport.py +574 -0
- pi_python/proxy.py +304 -0
- pi_python/py.typed +0 -0
- pi_python/queues.py +76 -0
- pi_python/recovery.py +209 -0
- pi_python/run.py +419 -0
- pi_python/stream.py +251 -0
- pi_python/sync.py +78 -0
- pi_python/testing.py +25 -0
- pi_python/tools.py +546 -0
- pi_python/transcript.py +167 -0
- pi_python_core-0.8.1.dist-info/METADATA +119 -0
- pi_python_core-0.8.1.dist-info/RECORD +39 -0
- pi_python_core-0.8.1.dist-info/WHEEL +4 -0
- pi_python_core-0.8.1.dist-info/licenses/LICENSE +21 -0
- pi_python_core-0.8.1.dist-info/licenses/NOTICE +8 -0
pi_python/__init__.py
ADDED
|
@@ -0,0 +1,160 @@
|
|
|
1
|
+
"""Embeddable asyncio agent core with optional model providers."""
|
|
2
|
+
|
|
3
|
+
from .models import ModelInfo, ModelCatalog
|
|
4
|
+
from ._version import __version__ as __version__
|
|
5
|
+
from .agent import Agent, AgentStateView, RunResult
|
|
6
|
+
from .cancellation import CancelToken
|
|
7
|
+
from .errors import (
|
|
8
|
+
PiError,
|
|
9
|
+
ConfigurationError,
|
|
10
|
+
MessageValidationError,
|
|
11
|
+
UnsupportedCapabilityError,
|
|
12
|
+
ProviderProtocolError,
|
|
13
|
+
AgentBusyError,
|
|
14
|
+
AgentClosedError,
|
|
15
|
+
InvalidContinuationError,
|
|
16
|
+
ToolOutcomeUnknownError,
|
|
17
|
+
CleanupTimeoutError,
|
|
18
|
+
SubscriptionError,
|
|
19
|
+
CandidateValidationError,
|
|
20
|
+
)
|
|
21
|
+
from .events import Event, EventListener, Unsubscribe, EventQueue, encode_event, decode_event
|
|
22
|
+
from .hooks import AgentConfigUpdate, Hooks, RunContext, TurnUpdate
|
|
23
|
+
from .limits import RunLimits
|
|
24
|
+
from .messages import (
|
|
25
|
+
JsonValue,
|
|
26
|
+
Message,
|
|
27
|
+
TextContent,
|
|
28
|
+
ImageContent,
|
|
29
|
+
ThinkingContent,
|
|
30
|
+
ToolCall,
|
|
31
|
+
ToolDeclaration,
|
|
32
|
+
SystemMessage,
|
|
33
|
+
UserMessage,
|
|
34
|
+
AssistantMessage,
|
|
35
|
+
ToolResultMessage,
|
|
36
|
+
CustomMessage,
|
|
37
|
+
encode_messages,
|
|
38
|
+
decode_messages,
|
|
39
|
+
message_to_dict,
|
|
40
|
+
message_from_dict,
|
|
41
|
+
validate_history,
|
|
42
|
+
)
|
|
43
|
+
from .provider import ModelRequest, ModelEvent, Provider, set_default_stream_fn
|
|
44
|
+
from .lowlevel import (
|
|
45
|
+
AgentContext,
|
|
46
|
+
AgentLoopConfig,
|
|
47
|
+
AgentEventStream,
|
|
48
|
+
agent_loop,
|
|
49
|
+
agent_loop_continue,
|
|
50
|
+
run_agent_loop,
|
|
51
|
+
run_agent_loop_continue,
|
|
52
|
+
)
|
|
53
|
+
from .proxy import ProxyProvider, stream_proxy
|
|
54
|
+
from .testing import ScriptedProvider
|
|
55
|
+
from .function_tools import tool
|
|
56
|
+
from .sync import run_sync
|
|
57
|
+
from .recovery import (
|
|
58
|
+
is_context_overflow,
|
|
59
|
+
is_recoverable_length,
|
|
60
|
+
is_retryable_error,
|
|
61
|
+
retry_delay,
|
|
62
|
+
)
|
|
63
|
+
from .tools import (
|
|
64
|
+
Tool,
|
|
65
|
+
ToolContext,
|
|
66
|
+
ToolExecutor,
|
|
67
|
+
ToolResult,
|
|
68
|
+
ToolResultUpdate,
|
|
69
|
+
ToolOutcome,
|
|
70
|
+
run_tool_call,
|
|
71
|
+
)
|
|
72
|
+
from .transcript import (
|
|
73
|
+
current_tools,
|
|
74
|
+
current_system_message,
|
|
75
|
+
current_system_prompt,
|
|
76
|
+
render_system_update,
|
|
77
|
+
)
|
|
78
|
+
from .estimate import clamp_max_tokens_to_context, estimate_context_tokens
|
|
79
|
+
|
|
80
|
+
__all__ = [
|
|
81
|
+
"ModelInfo",
|
|
82
|
+
"ModelCatalog",
|
|
83
|
+
"Agent",
|
|
84
|
+
"ImageContent",
|
|
85
|
+
"ThinkingContent",
|
|
86
|
+
"set_default_stream_fn",
|
|
87
|
+
"AgentContext",
|
|
88
|
+
"AgentLoopConfig",
|
|
89
|
+
"AgentEventStream",
|
|
90
|
+
"agent_loop",
|
|
91
|
+
"agent_loop_continue",
|
|
92
|
+
"run_agent_loop",
|
|
93
|
+
"run_agent_loop_continue",
|
|
94
|
+
"ProxyProvider",
|
|
95
|
+
"stream_proxy",
|
|
96
|
+
"AgentStateView",
|
|
97
|
+
"RunResult",
|
|
98
|
+
"CancelToken",
|
|
99
|
+
"PiError",
|
|
100
|
+
"ConfigurationError",
|
|
101
|
+
"MessageValidationError",
|
|
102
|
+
"UnsupportedCapabilityError",
|
|
103
|
+
"ProviderProtocolError",
|
|
104
|
+
"AgentBusyError",
|
|
105
|
+
"AgentClosedError",
|
|
106
|
+
"InvalidContinuationError",
|
|
107
|
+
"ToolOutcomeUnknownError",
|
|
108
|
+
"CleanupTimeoutError",
|
|
109
|
+
"SubscriptionError",
|
|
110
|
+
"CandidateValidationError",
|
|
111
|
+
"Event",
|
|
112
|
+
"EventListener",
|
|
113
|
+
"Unsubscribe",
|
|
114
|
+
"EventQueue",
|
|
115
|
+
"encode_event",
|
|
116
|
+
"decode_event",
|
|
117
|
+
"AgentConfigUpdate",
|
|
118
|
+
"Hooks",
|
|
119
|
+
"RunContext",
|
|
120
|
+
"TurnUpdate",
|
|
121
|
+
"RunLimits",
|
|
122
|
+
"JsonValue",
|
|
123
|
+
"Message",
|
|
124
|
+
"TextContent",
|
|
125
|
+
"ToolCall",
|
|
126
|
+
"ToolDeclaration",
|
|
127
|
+
"SystemMessage",
|
|
128
|
+
"UserMessage",
|
|
129
|
+
"AssistantMessage",
|
|
130
|
+
"ToolResultMessage",
|
|
131
|
+
"CustomMessage",
|
|
132
|
+
"encode_messages",
|
|
133
|
+
"decode_messages",
|
|
134
|
+
"message_to_dict",
|
|
135
|
+
"message_from_dict",
|
|
136
|
+
"validate_history",
|
|
137
|
+
"ModelRequest",
|
|
138
|
+
"ModelEvent",
|
|
139
|
+
"Provider",
|
|
140
|
+
"ScriptedProvider",
|
|
141
|
+
"Tool",
|
|
142
|
+
"tool",
|
|
143
|
+
"run_sync",
|
|
144
|
+
"is_context_overflow",
|
|
145
|
+
"is_recoverable_length",
|
|
146
|
+
"is_retryable_error",
|
|
147
|
+
"retry_delay",
|
|
148
|
+
"ToolContext",
|
|
149
|
+
"ToolExecutor",
|
|
150
|
+
"ToolResult",
|
|
151
|
+
"ToolResultUpdate",
|
|
152
|
+
"ToolOutcome",
|
|
153
|
+
"run_tool_call",
|
|
154
|
+
"current_tools",
|
|
155
|
+
"current_system_message",
|
|
156
|
+
"current_system_prompt",
|
|
157
|
+
"render_system_update",
|
|
158
|
+
"estimate_context_tokens",
|
|
159
|
+
"clamp_max_tokens_to_context",
|
|
160
|
+
]
|
pi_python/_version.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "0.8.1"
|
pi_python/agent.py
ADDED
|
@@ -0,0 +1,396 @@
|
|
|
1
|
+
"""The Agent session: history, default configuration, queues and subscribers.
|
|
2
|
+
|
|
3
|
+
Each prompt or continuation executes as a Run (run.py) through the loop (loop.py).
|
|
4
|
+
Execution behavior is adapted from Pi; see NOTICE.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
import asyncio
|
|
9
|
+
from copy import copy, deepcopy
|
|
10
|
+
from dataclasses import dataclass
|
|
11
|
+
from typing import Any, Callable
|
|
12
|
+
|
|
13
|
+
from .cancellation import CancelToken
|
|
14
|
+
from .errors import (
|
|
15
|
+
AgentBusyError,
|
|
16
|
+
AgentClosedError,
|
|
17
|
+
CleanupTimeoutError,
|
|
18
|
+
ConfigurationError,
|
|
19
|
+
InvalidContinuationError,
|
|
20
|
+
ToolOutcomeUnknownError,
|
|
21
|
+
)
|
|
22
|
+
from .events import EventDispatcher, EventListener, Unsubscribe
|
|
23
|
+
from .hooks import AgentConfigUpdate, Hooks
|
|
24
|
+
from .limits import RunLimits
|
|
25
|
+
from .models import ModelInfo
|
|
26
|
+
from .messages import (
|
|
27
|
+
AssistantMessage,
|
|
28
|
+
ImageContent,
|
|
29
|
+
TextContent,
|
|
30
|
+
Message,
|
|
31
|
+
SystemMessage,
|
|
32
|
+
UserMessage,
|
|
33
|
+
validate_history,
|
|
34
|
+
validate_json,
|
|
35
|
+
)
|
|
36
|
+
from .provider import Provider, DefaultProvider, FunctionProvider, has_default_stream
|
|
37
|
+
from .queues import MessageQueues
|
|
38
|
+
from .run import CALLER_CANCELLED, Run, RunResult, merge
|
|
39
|
+
from .sync import run_sync
|
|
40
|
+
from .tools import Tool
|
|
41
|
+
from .transcript import current_system_message
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
@dataclass(frozen=True)
|
|
45
|
+
class AgentStateView:
|
|
46
|
+
messages: tuple[Message, ...]
|
|
47
|
+
is_running: bool
|
|
48
|
+
partial_response: Any
|
|
49
|
+
pending_calls: tuple[str, ...]
|
|
50
|
+
last_error: str | None
|
|
51
|
+
reconciliation_required: bool
|
|
52
|
+
cleanup_complete: bool
|
|
53
|
+
closed: bool
|
|
54
|
+
diagnostics: tuple[dict[str, Any], ...]
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
class Agent:
|
|
58
|
+
def __init__(
|
|
59
|
+
self,
|
|
60
|
+
*,
|
|
61
|
+
provider: Provider | None = None,
|
|
62
|
+
stream_fn: Callable | None = None,
|
|
63
|
+
model: str | ModelInfo = "mock",
|
|
64
|
+
options: dict[str, Any] | None = None,
|
|
65
|
+
system_prompt: str = "",
|
|
66
|
+
thinking_level: str | None = None,
|
|
67
|
+
thinking_budgets: dict[str, int] | None = None,
|
|
68
|
+
transport: str | None = None,
|
|
69
|
+
session_id: str | None = None,
|
|
70
|
+
get_api_key: Callable | None = None,
|
|
71
|
+
on_payload: Callable | None = None,
|
|
72
|
+
on_response: Callable | None = None,
|
|
73
|
+
on_provider_stream_event: Callable | None = None,
|
|
74
|
+
tools: list[Tool] | None = None,
|
|
75
|
+
messages: list[Message] | None = None,
|
|
76
|
+
hooks: Hooks | None = None,
|
|
77
|
+
limits: RunLimits | None = None,
|
|
78
|
+
execution_mode: str = "parallel",
|
|
79
|
+
steering_mode: str = "one_at_a_time",
|
|
80
|
+
follow_up_mode: str = "one_at_a_time",
|
|
81
|
+
):
|
|
82
|
+
if execution_mode not in {"parallel", "sequential"}:
|
|
83
|
+
raise ConfigurationError("Invalid execution mode")
|
|
84
|
+
self._queues = MessageQueues(steering_mode, follow_up_mode)
|
|
85
|
+
if provider is not None and stream_fn is not None:
|
|
86
|
+
raise ConfigurationError("Pass provider or stream_fn, not both")
|
|
87
|
+
self.provider = provider or (
|
|
88
|
+
FunctionProvider(stream_fn) if stream_fn else DefaultProvider()
|
|
89
|
+
)
|
|
90
|
+
self.hooks = copy(hooks) if hooks else Hooks()
|
|
91
|
+
for name, callback in (
|
|
92
|
+
("get_api_key", get_api_key),
|
|
93
|
+
("on_payload", on_payload),
|
|
94
|
+
("on_response", on_response),
|
|
95
|
+
("on_provider_stream_event", on_provider_stream_event),
|
|
96
|
+
):
|
|
97
|
+
if callback is not None:
|
|
98
|
+
setattr(self.hooks, name, callback)
|
|
99
|
+
options = deepcopy(options or {})
|
|
100
|
+
for name, value in (
|
|
101
|
+
("reasoning", thinking_level),
|
|
102
|
+
("thinking_budgets", thinking_budgets),
|
|
103
|
+
("transport", transport),
|
|
104
|
+
("session_id", session_id),
|
|
105
|
+
):
|
|
106
|
+
if value is not None:
|
|
107
|
+
options[name] = value
|
|
108
|
+
self.limits = limits or RunLimits()
|
|
109
|
+
self.execution_mode = execution_mode
|
|
110
|
+
self._get_steering_messages: Callable | None = None
|
|
111
|
+
self._get_follow_up_messages: Callable | None = None
|
|
112
|
+
self._defaults = AgentConfigUpdate(tools or [], model, options or {})
|
|
113
|
+
self._validate_update(self._defaults)
|
|
114
|
+
self._defaults = deepcopy(self._defaults)
|
|
115
|
+
self._messages = deepcopy(messages or [])
|
|
116
|
+
validate_history(self._messages)
|
|
117
|
+
if (system_prompt or tools) and not (
|
|
118
|
+
self._messages and isinstance(self._messages[0], SystemMessage)
|
|
119
|
+
):
|
|
120
|
+
self._messages.insert(
|
|
121
|
+
0,
|
|
122
|
+
SystemMessage(
|
|
123
|
+
system_prompt, tools_added=[t.declaration() for t in tools or []], timestamp=0
|
|
124
|
+
),
|
|
125
|
+
)
|
|
126
|
+
validate_history(self._messages)
|
|
127
|
+
self._updates: list[AgentConfigUpdate] = []
|
|
128
|
+
self._events = EventDispatcher()
|
|
129
|
+
self._run: Run | None = None
|
|
130
|
+
self._running = False
|
|
131
|
+
self._closed = False
|
|
132
|
+
# These outlive a run: an outcome a tool explicitly reported as unknown, or an
|
|
133
|
+
# unfinished cleanup, makes the instance unusable until the application
|
|
134
|
+
# reconciles or discards it. Cancelling a tool never does.
|
|
135
|
+
self._unknown = False
|
|
136
|
+
self._cleanup_complete = True
|
|
137
|
+
self._last_error: str | None = None
|
|
138
|
+
self._idle = asyncio.Event()
|
|
139
|
+
self._idle.set()
|
|
140
|
+
|
|
141
|
+
@property
|
|
142
|
+
def steering_mode(self) -> str:
|
|
143
|
+
return self._queues.steering_mode
|
|
144
|
+
|
|
145
|
+
@property
|
|
146
|
+
def follow_up_mode(self) -> str:
|
|
147
|
+
return self._queues.follow_up_mode
|
|
148
|
+
|
|
149
|
+
@property
|
|
150
|
+
def state(self) -> AgentStateView:
|
|
151
|
+
run = self._run if self._running else None
|
|
152
|
+
return AgentStateView(
|
|
153
|
+
tuple(deepcopy(self._messages)),
|
|
154
|
+
self._running,
|
|
155
|
+
deepcopy(run.partial) if run else None,
|
|
156
|
+
run.pending_calls if run else (),
|
|
157
|
+
self._last_error,
|
|
158
|
+
self._unknown,
|
|
159
|
+
self._cleanup_complete,
|
|
160
|
+
self._closed,
|
|
161
|
+
tuple(deepcopy(self._events.diagnostics)),
|
|
162
|
+
)
|
|
163
|
+
|
|
164
|
+
def subscribe(self, listener: EventListener) -> Unsubscribe:
|
|
165
|
+
return self._events.subscribe(listener)
|
|
166
|
+
|
|
167
|
+
def _usable(self) -> None:
|
|
168
|
+
if self._closed:
|
|
169
|
+
raise AgentClosedError("Agent is closed")
|
|
170
|
+
if not self._cleanup_complete:
|
|
171
|
+
raise CleanupTimeoutError(
|
|
172
|
+
"Managed operations did not finish cleanup; discard this Agent"
|
|
173
|
+
)
|
|
174
|
+
if self._unknown:
|
|
175
|
+
raise ToolOutcomeUnknownError("External outcome unknown; reconcile outside this Agent")
|
|
176
|
+
|
|
177
|
+
def _validate_update(self, update: AgentConfigUpdate) -> None:
|
|
178
|
+
if not isinstance(update, AgentConfigUpdate):
|
|
179
|
+
raise ConfigurationError("Expected AgentConfigUpdate")
|
|
180
|
+
if update.model is not None and not isinstance(update.model, (str, ModelInfo)):
|
|
181
|
+
raise ConfigurationError("model must be a model name or a ModelInfo")
|
|
182
|
+
if update.options is not None:
|
|
183
|
+
validate_json(update.options)
|
|
184
|
+
if not isinstance(update.options, dict):
|
|
185
|
+
raise ConfigurationError("options must be an object")
|
|
186
|
+
if update.tools is not None:
|
|
187
|
+
names = []
|
|
188
|
+
for tool in update.tools:
|
|
189
|
+
if not isinstance(tool, Tool):
|
|
190
|
+
raise ConfigurationError("Expected Tool")
|
|
191
|
+
tool.__post_init__()
|
|
192
|
+
names.append(tool.name)
|
|
193
|
+
if len(set(names)) != len(names):
|
|
194
|
+
raise ConfigurationError("Duplicate tool names")
|
|
195
|
+
|
|
196
|
+
def update_config(self, update: AgentConfigUpdate) -> None:
|
|
197
|
+
self._usable()
|
|
198
|
+
self._validate_update(update)
|
|
199
|
+
update = deepcopy(update)
|
|
200
|
+
if self._running:
|
|
201
|
+
self._updates.append(update)
|
|
202
|
+
else:
|
|
203
|
+
merge(self._defaults, update)
|
|
204
|
+
|
|
205
|
+
@staticmethod
|
|
206
|
+
def _input(message: str | Message | list[Message]) -> list[Message]:
|
|
207
|
+
result: list[Message] = (
|
|
208
|
+
[UserMessage(message)]
|
|
209
|
+
if isinstance(message, str)
|
|
210
|
+
else (message if isinstance(message, list) else [message])
|
|
211
|
+
)
|
|
212
|
+
if not result:
|
|
213
|
+
raise InvalidContinuationError("Empty prompt")
|
|
214
|
+
result = deepcopy(result)
|
|
215
|
+
validate_history(result)
|
|
216
|
+
return result
|
|
217
|
+
|
|
218
|
+
def steer(self, message: str | Message) -> None:
|
|
219
|
+
self._usable()
|
|
220
|
+
self._queues.steering.extend(self._input(message))
|
|
221
|
+
|
|
222
|
+
def follow_up(self, message: str | Message) -> None:
|
|
223
|
+
self._usable()
|
|
224
|
+
self._queues.follow_up.extend(self._input(message))
|
|
225
|
+
|
|
226
|
+
def clear_queues(self, *, steering: bool, follow_up: bool) -> None:
|
|
227
|
+
self._queues.clear(steering=steering, follow_up=follow_up)
|
|
228
|
+
|
|
229
|
+
def has_queued_messages(self) -> bool:
|
|
230
|
+
return bool(self._queues)
|
|
231
|
+
|
|
232
|
+
def peek_queued_messages(self) -> list[Message]:
|
|
233
|
+
return self._queues.peek()
|
|
234
|
+
|
|
235
|
+
def clear_steering_queue(self) -> None:
|
|
236
|
+
self.clear_queues(steering=True, follow_up=False)
|
|
237
|
+
|
|
238
|
+
def clear_follow_up_queue(self) -> None:
|
|
239
|
+
self.clear_queues(steering=False, follow_up=True)
|
|
240
|
+
|
|
241
|
+
def clear_all_queues(self) -> None:
|
|
242
|
+
self.clear_queues(steering=True, follow_up=True)
|
|
243
|
+
|
|
244
|
+
@property
|
|
245
|
+
def signal(self) -> CancelToken | None:
|
|
246
|
+
return self._run.token if self._running and self._run else None
|
|
247
|
+
|
|
248
|
+
def reset(self) -> None:
|
|
249
|
+
self._usable()
|
|
250
|
+
if self._running:
|
|
251
|
+
raise AgentBusyError("Cannot reset while running")
|
|
252
|
+
baseline = current_system_message(self._messages)
|
|
253
|
+
self._messages = [baseline] if baseline else []
|
|
254
|
+
self._last_error = None
|
|
255
|
+
self._run = None
|
|
256
|
+
self.clear_all_queues()
|
|
257
|
+
|
|
258
|
+
def abort(self, reason: str = "requested") -> None:
|
|
259
|
+
"""Signal the run, as Pi does. Running operations receive task cancellation and
|
|
260
|
+
settle within the cleanup deadline; the run then ends with an aborted response.
|
|
261
|
+
|
|
262
|
+
Safe to call from any thread, for example a GUI or a watchdog.
|
|
263
|
+
"""
|
|
264
|
+
run = self._run
|
|
265
|
+
if not (self._running and run) or run.token.cancelled:
|
|
266
|
+
return
|
|
267
|
+
try:
|
|
268
|
+
current = asyncio.get_running_loop()
|
|
269
|
+
except RuntimeError:
|
|
270
|
+
current = None
|
|
271
|
+
if run.loop is not None and current is not run.loop:
|
|
272
|
+
run.loop.call_soon_threadsafe(self._abort_on_loop, run, reason)
|
|
273
|
+
else:
|
|
274
|
+
self._abort_on_loop(run, reason)
|
|
275
|
+
|
|
276
|
+
def _abort_on_loop(self, run: Run, reason: str) -> None:
|
|
277
|
+
if self._run is run and self._running and not run.token.cancelled:
|
|
278
|
+
run.token.cancel(reason)
|
|
279
|
+
|
|
280
|
+
def _cleanup_finished(self) -> None:
|
|
281
|
+
"""Operations left running by a cleanup timeout have all ended."""
|
|
282
|
+
self._cleanup_complete = True
|
|
283
|
+
run = self._run
|
|
284
|
+
if run is None or run.driver is None or run.driver.done():
|
|
285
|
+
self._running = False
|
|
286
|
+
self._idle.set()
|
|
287
|
+
|
|
288
|
+
async def wait_for_idle(self) -> None:
|
|
289
|
+
if not self._cleanup_complete:
|
|
290
|
+
raise CleanupTimeoutError("Cleanup deadline exceeded")
|
|
291
|
+
await self._idle.wait()
|
|
292
|
+
if not self._cleanup_complete:
|
|
293
|
+
raise CleanupTimeoutError("Cleanup deadline exceeded")
|
|
294
|
+
|
|
295
|
+
async def aclose(self) -> None:
|
|
296
|
+
self._closed = True
|
|
297
|
+
self.abort("closed")
|
|
298
|
+
await self.wait_for_idle()
|
|
299
|
+
|
|
300
|
+
async def __aenter__(self) -> Agent:
|
|
301
|
+
self._usable()
|
|
302
|
+
return self
|
|
303
|
+
|
|
304
|
+
async def __aexit__(self, *args: Any) -> None:
|
|
305
|
+
await self.aclose()
|
|
306
|
+
|
|
307
|
+
async def prompt(
|
|
308
|
+
self, message: str | Message | list[Message], images: list[ImageContent] | None = None
|
|
309
|
+
) -> RunResult:
|
|
310
|
+
if images is not None:
|
|
311
|
+
if not isinstance(message, str):
|
|
312
|
+
raise ConfigurationError("images requires a string prompt")
|
|
313
|
+
message = UserMessage([TextContent(message), *images])
|
|
314
|
+
self._usable()
|
|
315
|
+
if self._running:
|
|
316
|
+
raise AgentBusyError("Agent already running")
|
|
317
|
+
return await self._start(self._input(message))
|
|
318
|
+
|
|
319
|
+
def prompt_sync(
|
|
320
|
+
self, message: str | Message | list[Message], images: list[ImageContent] | None = None
|
|
321
|
+
) -> RunResult:
|
|
322
|
+
"""Blocking `prompt` for plain scripts; Ctrl+C aborts the run.
|
|
323
|
+
|
|
324
|
+
Runs on a shared background event loop, so use one style per Agent: either
|
|
325
|
+
these blocking calls or the async API inside your own event loop.
|
|
326
|
+
"""
|
|
327
|
+
return run_sync(self.prompt(message, images), on_interrupt=self.abort)
|
|
328
|
+
|
|
329
|
+
def continue_run_sync(self) -> RunResult:
|
|
330
|
+
"""Blocking `continue_run`; see `prompt_sync`."""
|
|
331
|
+
return run_sync(self.continue_run(), on_interrupt=self.abort)
|
|
332
|
+
|
|
333
|
+
async def continue_run(self) -> RunResult:
|
|
334
|
+
self._usable()
|
|
335
|
+
if self._running:
|
|
336
|
+
raise AgentBusyError("Agent already running")
|
|
337
|
+
validate_history(self._messages)
|
|
338
|
+
queues = self._queues
|
|
339
|
+
tail = self._messages[-1] if self._messages else None
|
|
340
|
+
# A failed or aborted last response can be retried: Provider replay skips it.
|
|
341
|
+
# (Pi's coding agent deletes it first; here it stays in the record.)
|
|
342
|
+
retry = isinstance(tail, AssistantMessage) and tail.stop_reason in {"error", "aborted"}
|
|
343
|
+
if (
|
|
344
|
+
tail is None
|
|
345
|
+
or all(isinstance(m, SystemMessage) for m in self._messages)
|
|
346
|
+
or (isinstance(tail, AssistantMessage) and not queues and not retry)
|
|
347
|
+
):
|
|
348
|
+
raise InvalidContinuationError("No unfinished interaction or queued messages")
|
|
349
|
+
if isinstance(tail, AssistantMessage) and queues:
|
|
350
|
+
if queues.steering:
|
|
351
|
+
return await self._start(queues.take(True), skip_initial_steering=True)
|
|
352
|
+
return await self._start(queues.take(False))
|
|
353
|
+
return await self._start([])
|
|
354
|
+
|
|
355
|
+
async def _start(
|
|
356
|
+
self, pending: list[Message], skip_initial_steering: bool = False
|
|
357
|
+
) -> RunResult:
|
|
358
|
+
if isinstance(self.provider, DefaultProvider) and not has_default_stream():
|
|
359
|
+
raise ConfigurationError(
|
|
360
|
+
"No model provider: pass Agent(provider=...) or call set_default_stream_fn(...)"
|
|
361
|
+
)
|
|
362
|
+
self._running = True
|
|
363
|
+
self._idle.clear()
|
|
364
|
+
self._last_error = None
|
|
365
|
+
self._events.failed = False
|
|
366
|
+
run = self._run = Run(self, skip_initial_steering)
|
|
367
|
+
run.loop = asyncio.get_running_loop()
|
|
368
|
+
# Driver starts with a checkpoint so cancellation before scheduling still finalizes.
|
|
369
|
+
run.driver = asyncio.create_task(self._drive(run, pending))
|
|
370
|
+
timer = None
|
|
371
|
+
if self.limits.run_timeout is not None:
|
|
372
|
+
timer = asyncio.get_running_loop().call_later(
|
|
373
|
+
self.limits.run_timeout, self.abort, "run_timeout"
|
|
374
|
+
)
|
|
375
|
+
try:
|
|
376
|
+
return await asyncio.shield(run.driver)
|
|
377
|
+
except asyncio.CancelledError:
|
|
378
|
+
run.token.cancel(CALLER_CANCELLED)
|
|
379
|
+
if run.driver_started and not run.finalizing and not run.driver.done():
|
|
380
|
+
run.driver.cancel()
|
|
381
|
+
# A second caller cancellation must not orphan the cleanup owner.
|
|
382
|
+
while not run.driver.done():
|
|
383
|
+
try:
|
|
384
|
+
await asyncio.shield(run.driver)
|
|
385
|
+
except asyncio.CancelledError:
|
|
386
|
+
continue
|
|
387
|
+
raise
|
|
388
|
+
finally:
|
|
389
|
+
if timer is not None:
|
|
390
|
+
timer.cancel()
|
|
391
|
+
|
|
392
|
+
async def _drive(self, run: Run, pending: list[Message]) -> RunResult:
|
|
393
|
+
result = await run.drive(pending)
|
|
394
|
+
self._running = not self._cleanup_complete
|
|
395
|
+
self._idle.set() # Wakes waiters; they check cleanup_complete before returning.
|
|
396
|
+
return result
|
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
import asyncio
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
class CancelToken:
|
|
6
|
+
def __init__(self) -> None:
|
|
7
|
+
self._event = asyncio.Event()
|
|
8
|
+
self.reason: str | None = None
|
|
9
|
+
|
|
10
|
+
@property
|
|
11
|
+
def cancelled(self) -> bool:
|
|
12
|
+
return self._event.is_set()
|
|
13
|
+
|
|
14
|
+
def cancel(self, reason: str = "requested") -> None:
|
|
15
|
+
if not self.cancelled:
|
|
16
|
+
self.reason = reason
|
|
17
|
+
self._event.set()
|
|
18
|
+
|
|
19
|
+
async def wait(self) -> None:
|
|
20
|
+
await self._event.wait()
|
|
21
|
+
|
|
22
|
+
def raise_if_cancelled(self) -> None:
|
|
23
|
+
if self.cancelled:
|
|
24
|
+
raise asyncio.CancelledError(self.reason)
|