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/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))