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/loop.py
ADDED
|
@@ -0,0 +1,431 @@
|
|
|
1
|
+
"""Model/tool loop. Shared observable behavior derives from pinned Pi (see NOTICE)."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
from typing import Any
|
|
5
|
+
from .stream import checked_events
|
|
6
|
+
import asyncio
|
|
7
|
+
from contextlib import nullcontext
|
|
8
|
+
from copy import deepcopy
|
|
9
|
+
from typing import TYPE_CHECKING
|
|
10
|
+
|
|
11
|
+
from .errors import ConfigurationError, ProviderProtocolError, UnsupportedCapabilityError
|
|
12
|
+
from .messages import (
|
|
13
|
+
Message,
|
|
14
|
+
AssistantMessage,
|
|
15
|
+
CustomMessage,
|
|
16
|
+
TextContent,
|
|
17
|
+
ToolCall,
|
|
18
|
+
ToolResultMessage,
|
|
19
|
+
message_from_dict,
|
|
20
|
+
message_to_dict,
|
|
21
|
+
validate_history,
|
|
22
|
+
)
|
|
23
|
+
from .models import ModelInfo
|
|
24
|
+
from .provider import ModelRequest
|
|
25
|
+
from .tools import (
|
|
26
|
+
ToolContext,
|
|
27
|
+
ToolOutcome,
|
|
28
|
+
aborted_result,
|
|
29
|
+
error_result,
|
|
30
|
+
prepare_tool_call,
|
|
31
|
+
run_tool_call,
|
|
32
|
+
)
|
|
33
|
+
from .transcript import current_tools
|
|
34
|
+
|
|
35
|
+
if TYPE_CHECKING:
|
|
36
|
+
from .run import Run
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def abort_error(run: Run) -> str:
|
|
40
|
+
return f"Request aborted ({run.token.reason})" if run.token.reason else "Request aborted"
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def without_tool_calls(message: AssistantMessage) -> AssistantMessage:
|
|
44
|
+
"""A failed response cannot declare calls (D7); keep its other content and say what was removed."""
|
|
45
|
+
calls = [b.id for b in message.content if isinstance(b, ToolCall)]
|
|
46
|
+
if calls and message.stop_reason in {"error", "aborted"}:
|
|
47
|
+
message.content = [b for b in message.content if not isinstance(b, ToolCall)]
|
|
48
|
+
message.diagnostics = [
|
|
49
|
+
*(message.diagnostics or []),
|
|
50
|
+
{"type": "removed_tool_calls", "call_ids": calls},
|
|
51
|
+
]
|
|
52
|
+
return message
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def failure_message(
|
|
56
|
+
run: Run, stop_reason: str, error: str | None, partial: dict[str, Any] | None = None
|
|
57
|
+
) -> AssistantMessage:
|
|
58
|
+
"""Pi's record of a failed or aborted response: what streamed so far, plus the error."""
|
|
59
|
+
model = run.context.model
|
|
60
|
+
if partial is not None:
|
|
61
|
+
message = message_from_dict(partial)
|
|
62
|
+
assert isinstance(message, AssistantMessage)
|
|
63
|
+
else:
|
|
64
|
+
message = AssistantMessage(
|
|
65
|
+
[TextContent("")],
|
|
66
|
+
provider=getattr(run.provider, "name", "custom"),
|
|
67
|
+
model=model.id if isinstance(model, ModelInfo) else model,
|
|
68
|
+
api=getattr(run.provider, "api", None),
|
|
69
|
+
)
|
|
70
|
+
message.stop_reason, message.error = stop_reason, error
|
|
71
|
+
message.thinking_level = run.context.options.get("reasoning", "off")
|
|
72
|
+
return without_tool_calls(message)
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
async def stream_response(run: Run) -> tuple[AssistantMessage, bool]:
|
|
76
|
+
context = run.context
|
|
77
|
+
assert context is not None
|
|
78
|
+
run.partial, run.started = None, False
|
|
79
|
+
messages = deepcopy(context.messages)
|
|
80
|
+
if run.hooks.transform_context:
|
|
81
|
+
messages = await run.hook("transform_context", messages, run.token)
|
|
82
|
+
if run.hooks.convert_to_llm:
|
|
83
|
+
messages = await run.hook("convert_to_llm", deepcopy(messages))
|
|
84
|
+
if any(isinstance(m, CustomMessage) for m in messages):
|
|
85
|
+
raise UnsupportedCapabilityError("Custom messages require convert_to_llm")
|
|
86
|
+
validate_history(messages)
|
|
87
|
+
for message in messages:
|
|
88
|
+
if isinstance(message, ToolResultMessage):
|
|
89
|
+
message.details = None
|
|
90
|
+
message.usage = None
|
|
91
|
+
message.nested_calls = None
|
|
92
|
+
model = context.model
|
|
93
|
+
request = ModelRequest(
|
|
94
|
+
deepcopy(messages),
|
|
95
|
+
current_tools(messages),
|
|
96
|
+
model.id if isinstance(model, ModelInfo) else model,
|
|
97
|
+
deepcopy(context.options),
|
|
98
|
+
model_info=deepcopy(model) if isinstance(model, ModelInfo) else None,
|
|
99
|
+
)
|
|
100
|
+
if run.hooks.get_api_key:
|
|
101
|
+
request.api_key = await run.hook(
|
|
102
|
+
"get_api_key", getattr(run.provider, "name", request.model)
|
|
103
|
+
)
|
|
104
|
+
request.on_payload = run.hooks.on_payload
|
|
105
|
+
request.on_response = run.hooks.on_response
|
|
106
|
+
request.on_provider_stream_event = run.hooks.on_provider_stream_event
|
|
107
|
+
if run.token.cancelled:
|
|
108
|
+
# A Pi stream function given an aborted signal answers at once; no request is sent.
|
|
109
|
+
return failure_message(run, "aborted", abort_error(run)), False
|
|
110
|
+
source = run.provider.stream(request, run.token)
|
|
111
|
+
stream_method = getattr(run.provider, "stream", None)
|
|
112
|
+
# Remote providers already check their own events; check every other source here.
|
|
113
|
+
iterator = (
|
|
114
|
+
source
|
|
115
|
+
if getattr(stream_method, "checked", False)
|
|
116
|
+
else checked_events(
|
|
117
|
+
source,
|
|
118
|
+
AssistantMessage(
|
|
119
|
+
[],
|
|
120
|
+
stop_reason="pending",
|
|
121
|
+
provider=getattr(run.provider, "name", "custom"),
|
|
122
|
+
model=request.model,
|
|
123
|
+
api=getattr(run.provider, "api", None),
|
|
124
|
+
),
|
|
125
|
+
)
|
|
126
|
+
)
|
|
127
|
+
final = None
|
|
128
|
+
try:
|
|
129
|
+
async for event in iterator:
|
|
130
|
+
run.token.raise_if_cancelled()
|
|
131
|
+
if event.type == "error":
|
|
132
|
+
# Pi commits the provider's own failed message: partial content, usage, error.
|
|
133
|
+
stop_reason = "aborted" if event.reason == "aborted" else "error"
|
|
134
|
+
failed = event.message
|
|
135
|
+
if isinstance(failed, AssistantMessage) and failed.stop_reason == stop_reason:
|
|
136
|
+
final = without_tool_calls(deepcopy(failed))
|
|
137
|
+
final.thinking_level = context.options.get("reasoning", "off")
|
|
138
|
+
else:
|
|
139
|
+
error = getattr(failed, "error", None) or "Provider failed"
|
|
140
|
+
final = failure_message(run, stop_reason, error, run.partial)
|
|
141
|
+
break
|
|
142
|
+
if event.type == "done":
|
|
143
|
+
assert event.message is not None
|
|
144
|
+
final = without_tool_calls(deepcopy(event.message))
|
|
145
|
+
final.thinking_level = context.options.get("reasoning", "off")
|
|
146
|
+
continue
|
|
147
|
+
assert event.partial is not None # checked events always carry a snapshot
|
|
148
|
+
run.partial = message_to_dict(event.partial)
|
|
149
|
+
if not run.started:
|
|
150
|
+
run.started = True
|
|
151
|
+
await run.emit("message_start", message=run.partial)
|
|
152
|
+
if event.type != "start":
|
|
153
|
+
await run.emit(
|
|
154
|
+
"message_update",
|
|
155
|
+
delta_type=event.type,
|
|
156
|
+
block_index=event.index,
|
|
157
|
+
delta=event.delta,
|
|
158
|
+
content=event.content,
|
|
159
|
+
tool_call_id=event.call_id,
|
|
160
|
+
partial=run.partial,
|
|
161
|
+
)
|
|
162
|
+
finally:
|
|
163
|
+
close = getattr(iterator, "aclose", None)
|
|
164
|
+
if close is not None:
|
|
165
|
+
await close()
|
|
166
|
+
if final is None:
|
|
167
|
+
raise ProviderProtocolError("Stream ended without final message")
|
|
168
|
+
run.partial = None
|
|
169
|
+
return final, run.started
|
|
170
|
+
|
|
171
|
+
|
|
172
|
+
async def execute_batch(run: Run, outcomes: list[ToolOutcome]) -> None:
|
|
173
|
+
context = run.context
|
|
174
|
+
assert context is not None
|
|
175
|
+
tools = {t.name: t for t in context.tools}
|
|
176
|
+
sequential = run.execution_mode == "sequential" or any(
|
|
177
|
+
tools[o.call.name].execution_mode == "sequential" for o in outcomes if o.call.name in tools
|
|
178
|
+
)
|
|
179
|
+
limit = run.limits.max_concurrency
|
|
180
|
+
# Unlimited by default, like Pi's Promise.all over the batch.
|
|
181
|
+
semaphore = asyncio.Semaphore(limit) if limit is not None else nullcontext()
|
|
182
|
+
|
|
183
|
+
# One snapshot per batch, shared read-only by its calls: copying the whole context
|
|
184
|
+
# for every call made large batches over long histories slow.
|
|
185
|
+
snapshot: dict[str, Any] = {}
|
|
186
|
+
contexts: dict[int, ToolContext] = {}
|
|
187
|
+
|
|
188
|
+
def tool_context(outcome: ToolOutcome) -> ToolContext:
|
|
189
|
+
if id(outcome) in contexts:
|
|
190
|
+
return contexts[id(outcome)]
|
|
191
|
+
if not snapshot:
|
|
192
|
+
snapshot.update(message=deepcopy(context.message), context=deepcopy(context))
|
|
193
|
+
|
|
194
|
+
async def emit(value: Any) -> None:
|
|
195
|
+
await run.emit("tool_execution_update", outcome.call.id, update=value)
|
|
196
|
+
|
|
197
|
+
contexts[id(outcome)] = ToolContext(
|
|
198
|
+
run.id,
|
|
199
|
+
outcome.call.id,
|
|
200
|
+
run.token,
|
|
201
|
+
emit,
|
|
202
|
+
assistant_message=snapshot["message"],
|
|
203
|
+
agent_context=snapshot["context"],
|
|
204
|
+
)
|
|
205
|
+
return contexts[id(outcome)]
|
|
206
|
+
|
|
207
|
+
def abort_before_start(outcome: ToolOutcome) -> bool:
|
|
208
|
+
"""Pi checks the abort signal before preparing and before executing each call."""
|
|
209
|
+
run.checkpoint()
|
|
210
|
+
if run.token.cancelled and outcome.result is None:
|
|
211
|
+
outcome.result = aborted_result()
|
|
212
|
+
return outcome.result is not None
|
|
213
|
+
|
|
214
|
+
async def prepare(outcome: ToolOutcome) -> None:
|
|
215
|
+
if abort_before_start(outcome):
|
|
216
|
+
return
|
|
217
|
+
await run.emit(
|
|
218
|
+
"tool_execution_start",
|
|
219
|
+
outcome.call.id,
|
|
220
|
+
name=outcome.call.name,
|
|
221
|
+
arguments=outcome.original_arguments,
|
|
222
|
+
)
|
|
223
|
+
try:
|
|
224
|
+
await run.await_owned(
|
|
225
|
+
prepare_tool_call(
|
|
226
|
+
tools.get(outcome.call.name),
|
|
227
|
+
outcome,
|
|
228
|
+
tool_context(outcome),
|
|
229
|
+
run.hooks.before_tool_call,
|
|
230
|
+
)
|
|
231
|
+
)
|
|
232
|
+
except asyncio.CancelledError:
|
|
233
|
+
if not run.aborted():
|
|
234
|
+
raise
|
|
235
|
+
abort_before_start(outcome)
|
|
236
|
+
|
|
237
|
+
async def execute(outcome: ToolOutcome) -> None:
|
|
238
|
+
def timed_out() -> None:
|
|
239
|
+
# Freeze the fact known at the deadline before delivering cancellation.
|
|
240
|
+
# A coroutine which suppresses cancellation cannot rewrite this decision.
|
|
241
|
+
if outcome.execution_status == "running":
|
|
242
|
+
outcome.execution_status = "cancelled"
|
|
243
|
+
outcome.result = error_result("tool_timeout", "Tool timed out")
|
|
244
|
+
else:
|
|
245
|
+
outcome.result = error_result(
|
|
246
|
+
"finalization_timeout", "Execution finished; finalization timed out"
|
|
247
|
+
)
|
|
248
|
+
outcome._settled = True
|
|
249
|
+
|
|
250
|
+
async with semaphore:
|
|
251
|
+
if not abort_before_start(outcome):
|
|
252
|
+
try:
|
|
253
|
+
await run.await_owned(
|
|
254
|
+
run_tool_call(
|
|
255
|
+
tools.get(outcome.call.name),
|
|
256
|
+
outcome.call,
|
|
257
|
+
tool_context(outcome),
|
|
258
|
+
after_tool_call=run.hooks.after_tool_call,
|
|
259
|
+
outcome=outcome,
|
|
260
|
+
prepared=True,
|
|
261
|
+
),
|
|
262
|
+
run.limits.tool_timeout,
|
|
263
|
+
on_timeout=timed_out,
|
|
264
|
+
)
|
|
265
|
+
except TimeoutError:
|
|
266
|
+
pass # timed_out already recorded the deadline result
|
|
267
|
+
except asyncio.CancelledError:
|
|
268
|
+
# Abort reached the tool. Its own result, or Pi's "Operation aborted",
|
|
269
|
+
# was recorded by run_tool_call unless it is still running.
|
|
270
|
+
if not run.aborted():
|
|
271
|
+
raise
|
|
272
|
+
await run.end_tool(outcome)
|
|
273
|
+
|
|
274
|
+
if sequential:
|
|
275
|
+
for outcome in outcomes:
|
|
276
|
+
await prepare(outcome)
|
|
277
|
+
await execute(outcome)
|
|
278
|
+
await run.commit_outcome(outcome)
|
|
279
|
+
if outcome.execution_status == "unknown":
|
|
280
|
+
break
|
|
281
|
+
else:
|
|
282
|
+
for outcome in outcomes:
|
|
283
|
+
await prepare(outcome)
|
|
284
|
+
if outcome.result is not None:
|
|
285
|
+
await run.end_tool(outcome)
|
|
286
|
+
# Every await inside a worker is owned and bounded, so an abort reaches each tool
|
|
287
|
+
# directly and the batch settles; caller cancellation cancels the workers below.
|
|
288
|
+
workers = [asyncio.create_task(execute(o)) for o in outcomes]
|
|
289
|
+
try:
|
|
290
|
+
await asyncio.gather(*workers)
|
|
291
|
+
finally:
|
|
292
|
+
for worker in workers:
|
|
293
|
+
if not worker.done():
|
|
294
|
+
worker.cancel()
|
|
295
|
+
await asyncio.gather(*workers, return_exceptions=True)
|
|
296
|
+
for outcome in outcomes:
|
|
297
|
+
await run.commit_outcome(outcome)
|
|
298
|
+
|
|
299
|
+
|
|
300
|
+
async def run_loop(run: Run, initial: list[Message]) -> tuple[str, str]:
|
|
301
|
+
context = run.context
|
|
302
|
+
assert context is not None
|
|
303
|
+
requests = tool_count = 0
|
|
304
|
+
run.checkpoint()
|
|
305
|
+
await run.emit("agent_start")
|
|
306
|
+
run.turn = 1
|
|
307
|
+
await run.emit("turn_start")
|
|
308
|
+
await run.pending(initial)
|
|
309
|
+
pending = [] if run.skip_initial_steering else await run.poll_queue(True)
|
|
310
|
+
while True:
|
|
311
|
+
# After an abort the loop still runs to Pi's end: the next response is an
|
|
312
|
+
# aborted message, recorded like any other, without a provider request.
|
|
313
|
+
run.checkpoint()
|
|
314
|
+
max_requests = run.limits.max_model_requests
|
|
315
|
+
if max_requests is not None and requests >= max_requests:
|
|
316
|
+
return "limit_reached", "max_model_requests"
|
|
317
|
+
prepared = []
|
|
318
|
+
if requests:
|
|
319
|
+
prepared = await run.prepare("prepare_next_turn")
|
|
320
|
+
if not pending:
|
|
321
|
+
pending = await run.poll_queue(True)
|
|
322
|
+
run.turn += 1
|
|
323
|
+
await run.emit("turn_start")
|
|
324
|
+
await run.apply_updates()
|
|
325
|
+
await run.pending(prepared + pending)
|
|
326
|
+
extra = await run.prepare("prepare_request")
|
|
327
|
+
# Run hook updates persist, but never rewrite the committed transcript.
|
|
328
|
+
await run.pending(extra)
|
|
329
|
+
run.checkpoint()
|
|
330
|
+
requests += 1
|
|
331
|
+
failed_reason = None
|
|
332
|
+
try:
|
|
333
|
+
message, started = await run.await_owned(stream_response(run))
|
|
334
|
+
except asyncio.CancelledError:
|
|
335
|
+
if not run.aborted():
|
|
336
|
+
raise
|
|
337
|
+
message = failure_message(run, "aborted", abort_error(run), run.partial)
|
|
338
|
+
started = run.started
|
|
339
|
+
except Exception as exc:
|
|
340
|
+
failed_reason = f"{type(exc).__name__}: {exc}"
|
|
341
|
+
message = failure_message(run, "error", failed_reason, run.partial)
|
|
342
|
+
started = run.started
|
|
343
|
+
run.partial = None
|
|
344
|
+
outcomes = [
|
|
345
|
+
ToolOutcome(deepcopy(c), original_arguments=deepcopy(c.arguments))
|
|
346
|
+
for c in message.tool_calls
|
|
347
|
+
]
|
|
348
|
+
run.active_batch = outcomes
|
|
349
|
+
run.outcomes.extend(outcomes)
|
|
350
|
+
await run.append(message, started=started)
|
|
351
|
+
context.message = deepcopy(message)
|
|
352
|
+
status = "completed"
|
|
353
|
+
reason = message.stop_reason
|
|
354
|
+
if message.stop_reason == "error":
|
|
355
|
+
status = "failed"
|
|
356
|
+
failed_reason = failed_reason or message.error
|
|
357
|
+
run.fail(message.error or "Provider returned an error response")
|
|
358
|
+
elif message.stop_reason == "aborted":
|
|
359
|
+
# Pi exposes the aborted message's error as the agent's error message.
|
|
360
|
+
run.fail(message.error or "Request aborted")
|
|
361
|
+
if run.token.reason == "run_timeout":
|
|
362
|
+
status, reason = "limit_reached", "run_timeout"
|
|
363
|
+
else:
|
|
364
|
+
status, reason = "cancelled", run.token.reason or "aborted"
|
|
365
|
+
elif outcomes:
|
|
366
|
+
max_calls = run.limits.max_tool_calls
|
|
367
|
+
if max_calls is not None and tool_count + len(outcomes) > max_calls:
|
|
368
|
+
status, reason = "limit_reached", "max_tool_calls"
|
|
369
|
+
for outcome in outcomes:
|
|
370
|
+
outcome.result = error_result(
|
|
371
|
+
"limit", "Entire tool batch exceeds remaining budget"
|
|
372
|
+
)
|
|
373
|
+
elif message.stop_reason == "length":
|
|
374
|
+
for outcome in outcomes:
|
|
375
|
+
await run.emit(
|
|
376
|
+
"tool_execution_start",
|
|
377
|
+
outcome.call.id,
|
|
378
|
+
name=outcome.call.name,
|
|
379
|
+
arguments=outcome.original_arguments,
|
|
380
|
+
)
|
|
381
|
+
outcome.result = error_result(
|
|
382
|
+
"truncated", "Tool call not executed: output length limit"
|
|
383
|
+
)
|
|
384
|
+
await run.end_tool(outcome)
|
|
385
|
+
await run.commit_outcome(outcome)
|
|
386
|
+
else:
|
|
387
|
+
tool_count += len(outcomes)
|
|
388
|
+
await execute_batch(run, outcomes)
|
|
389
|
+
for outcome in outcomes:
|
|
390
|
+
await run.end_tool(outcome)
|
|
391
|
+
await run.commit_outcome(outcome)
|
|
392
|
+
if run.agent._unknown:
|
|
393
|
+
status, reason = "failed", "outcome_unknown"
|
|
394
|
+
elif run.stuck():
|
|
395
|
+
# Never start another request while a cancelled tool keeps running.
|
|
396
|
+
run.fail("Tool did not stop after cancellation")
|
|
397
|
+
status = "cancelled" if run.token.cancelled else "failed"
|
|
398
|
+
reason = "tool_not_stopped"
|
|
399
|
+
context.tool_results = [
|
|
400
|
+
ToolResultMessage(
|
|
401
|
+
o.call.id,
|
|
402
|
+
o.call.name,
|
|
403
|
+
deepcopy(o.result.content),
|
|
404
|
+
o.result.is_error,
|
|
405
|
+
details=deepcopy(o.result.details),
|
|
406
|
+
usage=deepcopy(o.result.usage),
|
|
407
|
+
nested_calls=deepcopy(o.result.nested_calls),
|
|
408
|
+
)
|
|
409
|
+
for o in outcomes
|
|
410
|
+
if o.result
|
|
411
|
+
]
|
|
412
|
+
decision = await run.hook("finish_turn", deepcopy(context), run.token)
|
|
413
|
+
if decision not in {None, "continue", "end"}:
|
|
414
|
+
raise ConfigurationError("finish_turn must return continue, end or None")
|
|
415
|
+
await run.emit(
|
|
416
|
+
"turn_end",
|
|
417
|
+
message=message_to_dict(message),
|
|
418
|
+
tool_results=[message_to_dict(m) for m in context.tool_results],
|
|
419
|
+
)
|
|
420
|
+
if status != "completed":
|
|
421
|
+
return status, failed_reason or reason
|
|
422
|
+
if decision == "end":
|
|
423
|
+
return "completed", "finish_turn"
|
|
424
|
+
pending = await run.poll_queue(True)
|
|
425
|
+
natural = bool(outcomes) and not all(o.result and o.result.terminate for o in outcomes)
|
|
426
|
+
if natural or pending:
|
|
427
|
+
continue
|
|
428
|
+
pending = await run.poll_queue(False)
|
|
429
|
+
if pending or decision == "continue":
|
|
430
|
+
continue
|
|
431
|
+
return "completed", reason
|
pi_python/lowlevel.py
ADDED
|
@@ -0,0 +1,179 @@
|
|
|
1
|
+
"""Public low-level loop variants backed by the same Agent execution engine."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
import asyncio
|
|
5
|
+
from collections.abc import Awaitable, Callable
|
|
6
|
+
from copy import deepcopy
|
|
7
|
+
from dataclasses import dataclass, field
|
|
8
|
+
from typing import Any
|
|
9
|
+
from .agent import Agent
|
|
10
|
+
from .cancellation import CancelToken
|
|
11
|
+
from .events import Event, EventListener
|
|
12
|
+
from .hooks import Hooks
|
|
13
|
+
from .limits import RunLimits
|
|
14
|
+
from .messages import Message
|
|
15
|
+
from .models import ModelInfo
|
|
16
|
+
from .provider import Provider
|
|
17
|
+
from .tools import Tool
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
@dataclass
|
|
21
|
+
class AgentContext:
|
|
22
|
+
messages: list[Message] = field(default_factory=list)
|
|
23
|
+
tools: list[Tool] = field(default_factory=list)
|
|
24
|
+
system_prompt: str = ""
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@dataclass
|
|
28
|
+
class AgentLoopConfig:
|
|
29
|
+
provider: Provider | None = None
|
|
30
|
+
stream_fn: Any = None
|
|
31
|
+
model: str | ModelInfo = "mock"
|
|
32
|
+
options: dict[str, Any] = field(default_factory=dict)
|
|
33
|
+
hooks: Hooks = field(default_factory=Hooks)
|
|
34
|
+
limits: RunLimits = field(default_factory=RunLimits)
|
|
35
|
+
tool_execution: str = "parallel"
|
|
36
|
+
get_steering_messages: Any = None
|
|
37
|
+
get_follow_up_messages: Any = None
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
async def _run(
|
|
41
|
+
prompts: list[Message] | None,
|
|
42
|
+
context: AgentContext,
|
|
43
|
+
config: AgentLoopConfig,
|
|
44
|
+
emit: EventListener | None,
|
|
45
|
+
cancel: CancelToken | None,
|
|
46
|
+
continuing: bool,
|
|
47
|
+
) -> list[Message]:
|
|
48
|
+
agent = Agent(
|
|
49
|
+
provider=config.provider,
|
|
50
|
+
stream_fn=config.stream_fn,
|
|
51
|
+
model=config.model,
|
|
52
|
+
options=config.options,
|
|
53
|
+
messages=context.messages,
|
|
54
|
+
tools=context.tools,
|
|
55
|
+
system_prompt=context.system_prompt,
|
|
56
|
+
hooks=config.hooks,
|
|
57
|
+
limits=config.limits,
|
|
58
|
+
execution_mode=config.tool_execution,
|
|
59
|
+
)
|
|
60
|
+
agent._get_steering_messages = config.get_steering_messages
|
|
61
|
+
agent._get_follow_up_messages = config.get_follow_up_messages
|
|
62
|
+
if emit is not None:
|
|
63
|
+
agent.subscribe(emit)
|
|
64
|
+
|
|
65
|
+
async def watch(token: CancelToken) -> None:
|
|
66
|
+
await token.wait()
|
|
67
|
+
agent.abort(token.reason or "requested")
|
|
68
|
+
|
|
69
|
+
watcher = asyncio.create_task(watch(cancel)) if cancel else None
|
|
70
|
+
try:
|
|
71
|
+
if cancel and cancel.cancelled:
|
|
72
|
+
return []
|
|
73
|
+
result = await (agent.continue_run() if continuing else agent.prompt(prompts or []))
|
|
74
|
+
if continuing:
|
|
75
|
+
context.messages[:] = list(agent.state.messages)
|
|
76
|
+
return result.messages
|
|
77
|
+
finally:
|
|
78
|
+
if watcher:
|
|
79
|
+
watcher.cancel()
|
|
80
|
+
await asyncio.gather(watcher, return_exceptions=True)
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
async def run_agent_loop(
|
|
84
|
+
prompts: list[Message],
|
|
85
|
+
context: AgentContext,
|
|
86
|
+
config: AgentLoopConfig,
|
|
87
|
+
emit: EventListener | None = None,
|
|
88
|
+
cancel: CancelToken | None = None,
|
|
89
|
+
) -> list[Message]:
|
|
90
|
+
return await _run(prompts, context, config, emit, cancel, False)
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
async def run_agent_loop_continue(
|
|
94
|
+
context: AgentContext,
|
|
95
|
+
config: AgentLoopConfig,
|
|
96
|
+
emit: EventListener | None = None,
|
|
97
|
+
cancel: CancelToken | None = None,
|
|
98
|
+
) -> list[Message]:
|
|
99
|
+
return await _run(None, context, config, emit, cancel, True)
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
class AgentEventStream:
|
|
103
|
+
"""Consume via async for, or await result() to drain events and obtain messages.
|
|
104
|
+
|
|
105
|
+
The bounded queue applies backpressure. Close an abandoned stream explicitly.
|
|
106
|
+
"""
|
|
107
|
+
|
|
108
|
+
def __init__(
|
|
109
|
+
self, runner: Callable[[Callable], Awaitable[list[Message]]], *, maxsize: int = 128
|
|
110
|
+
) -> None:
|
|
111
|
+
self.queue: asyncio.Queue = asyncio.Queue(maxsize)
|
|
112
|
+
self._iterating = False
|
|
113
|
+
self._ended = False
|
|
114
|
+
|
|
115
|
+
async def emit(event: Event) -> None:
|
|
116
|
+
await self.queue.put(deepcopy(event))
|
|
117
|
+
|
|
118
|
+
async def drive() -> list[Message]:
|
|
119
|
+
try:
|
|
120
|
+
return await runner(emit)
|
|
121
|
+
finally:
|
|
122
|
+
await self.queue.put(None)
|
|
123
|
+
|
|
124
|
+
self.task = asyncio.create_task(drive())
|
|
125
|
+
|
|
126
|
+
def __aiter__(self) -> AgentEventStream:
|
|
127
|
+
self._iterating = True
|
|
128
|
+
return self
|
|
129
|
+
|
|
130
|
+
async def __anext__(self) -> Event:
|
|
131
|
+
if self._ended:
|
|
132
|
+
raise StopAsyncIteration
|
|
133
|
+
if self.task.done() and self.queue.empty():
|
|
134
|
+
self._ended = True
|
|
135
|
+
await self.task
|
|
136
|
+
raise StopAsyncIteration
|
|
137
|
+
event = await self.queue.get()
|
|
138
|
+
if event is None:
|
|
139
|
+
self._ended = True
|
|
140
|
+
await self.task
|
|
141
|
+
raise StopAsyncIteration
|
|
142
|
+
return event
|
|
143
|
+
|
|
144
|
+
async def result(self) -> list[Message]:
|
|
145
|
+
if self.task.done():
|
|
146
|
+
return await self.task
|
|
147
|
+
if not self._iterating:
|
|
148
|
+
async for _ in self:
|
|
149
|
+
pass
|
|
150
|
+
return await asyncio.shield(self.task)
|
|
151
|
+
|
|
152
|
+
async def aclose(self) -> None:
|
|
153
|
+
self.task.cancel()
|
|
154
|
+
|
|
155
|
+
# Drain so a cancelled producer cannot block on its terminal marker.
|
|
156
|
+
async def drain() -> None:
|
|
157
|
+
while not self.task.done():
|
|
158
|
+
try:
|
|
159
|
+
await asyncio.wait_for(self.queue.get(), 0.05)
|
|
160
|
+
except TimeoutError:
|
|
161
|
+
pass
|
|
162
|
+
|
|
163
|
+
await drain()
|
|
164
|
+
await asyncio.gather(self.task, return_exceptions=True)
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
def agent_loop(
|
|
168
|
+
prompts: list[Message],
|
|
169
|
+
context: AgentContext,
|
|
170
|
+
config: AgentLoopConfig,
|
|
171
|
+
cancel: CancelToken | None = None,
|
|
172
|
+
) -> AgentEventStream:
|
|
173
|
+
return AgentEventStream(lambda emit: run_agent_loop(prompts, context, config, emit, cancel))
|
|
174
|
+
|
|
175
|
+
|
|
176
|
+
def agent_loop_continue(
|
|
177
|
+
context: AgentContext, config: AgentLoopConfig, cancel: CancelToken | None = None
|
|
178
|
+
) -> AgentEventStream:
|
|
179
|
+
return AgentEventStream(lambda emit: run_agent_loop_continue(context, config, emit, cancel))
|