Skip to content

Commit fbd0f8c

Browse files
xfgongclaude
andcommitted
refactor: introduce StreamOptions bag for provider stream parameters
Replace scattered thinking/signal/on_payload/on_response parameters with a single StreamOptions object that flows transparently from Agent through the loop to Provider.stream(). This mirrors pi's SimpleStreamOptions pattern — adding new stream options no longer requires changing signatures at every layer. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
1 parent 2be76f7 commit fbd0f8c

12 files changed

Lines changed: 92 additions & 86 deletions

File tree

cubepi/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
Model,
1717
Provider,
1818
StreamEvent,
19+
StreamOptions,
1920
TextContent,
2021
ThinkingBudgets,
2122
ThinkingLevel,
@@ -38,6 +39,7 @@
3839
"Model",
3940
"Provider",
4041
"StreamEvent",
42+
"StreamOptions",
4143
"TextContent",
4244
"ThinkingBudgets",
4345
"ThinkingLevel",

cubepi/agent/agent.py

Lines changed: 17 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,10 @@
1919
AssistantMessage,
2020
Message,
2121
Model,
22+
OnPayloadCallback,
23+
OnResponseCallback,
2224
Provider,
25+
StreamOptions,
2326
TextContent,
2427
ThinkingLevel,
2528
Usage,
@@ -116,6 +119,8 @@ def __init__(
116119
before_tool_call: Callable | None = None,
117120
after_tool_call: Callable | None = None,
118121
should_stop_after_turn: Callable | None = None,
122+
on_payload: OnPayloadCallback | None = None,
123+
on_response: OnResponseCallback | None = None,
119124
steering_mode: str = "one-at-a-time",
120125
follow_up_mode: str = "one-at-a-time",
121126
tool_execution: str = "parallel",
@@ -135,6 +140,8 @@ def __init__(
135140
self.before_tool_call = before_tool_call
136141
self.after_tool_call = after_tool_call
137142
self.should_stop_after_turn = should_stop_after_turn
143+
self.on_payload = on_payload
144+
self.on_response = on_response
138145
self.tool_execution = tool_execution
139146
self.checkpointer = checkpointer
140147
self.thread_id = thread_id
@@ -222,6 +229,14 @@ async def resume(self) -> None:
222229

223230
await self._run_continuation()
224231

232+
def _build_stream_options(self, signal: asyncio.Event) -> StreamOptions:
233+
return StreamOptions(
234+
thinking=self._state.thinking,
235+
signal=signal,
236+
on_payload=self.on_payload,
237+
on_response=self.on_response,
238+
)
239+
225240
async def _run_prompt(self, messages: list[Any]) -> None:
226241
await self._run_with_lifecycle(
227242
lambda signal: run_agent_loop(
@@ -236,9 +251,8 @@ async def _run_prompt(self, messages: list[Any]) -> None:
236251
should_stop_after_turn=self.should_stop_after_turn,
237252
get_steering_messages=self._make_async_drain(self._steering_queue),
238253
get_follow_up_messages=self._make_async_drain(self._follow_up_queue),
239-
thinking=self._state.thinking,
254+
stream_options=self._build_stream_options(signal),
240255
tool_execution=self.tool_execution,
241-
signal=signal,
242256
emit=lambda e: self._process_event(e),
243257
)
244258
)
@@ -256,9 +270,8 @@ async def _run_continuation(self) -> None:
256270
should_stop_after_turn=self.should_stop_after_turn,
257271
get_steering_messages=self._make_async_drain(self._steering_queue),
258272
get_follow_up_messages=self._make_async_drain(self._follow_up_queue),
259-
thinking=self._state.thinking,
273+
stream_options=self._build_stream_options(signal),
260274
tool_execution=self.tool_execution,
261-
signal=signal,
262275
emit=lambda e: self._process_event(e),
263276
)
264277
)

cubepi/agent/loop.py

Lines changed: 12 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,7 @@
1919
AssistantMessage,
2020
Model,
2121
Provider,
22-
ThinkingLevel,
22+
StreamOptions,
2323
ToolCall,
2424
ToolResultMessage,
2525
)
@@ -45,9 +45,8 @@ async def run_agent_loop(
4545
should_stop_after_turn: Callable | None = None,
4646
get_steering_messages: Callable | None = None,
4747
get_follow_up_messages: Callable | None = None,
48-
thinking: ThinkingLevel = "off",
48+
stream_options: StreamOptions | None = None,
4949
tool_execution: str = "parallel",
50-
signal: asyncio.Event | None = None,
5150
system_prompt: str | None = None,
5251
) -> list[Any]:
5352
new_messages: list[Any] = list(prompts)
@@ -77,9 +76,8 @@ async def run_agent_loop(
7776
should_stop_after_turn=should_stop_after_turn,
7877
get_steering_messages=get_steering_messages,
7978
get_follow_up_messages=get_follow_up_messages,
80-
thinking=thinking,
79+
stream_options=stream_options,
8180
tool_execution=tool_execution,
82-
signal=signal,
8381
emit=emit,
8482
)
8583
return new_messages
@@ -98,9 +96,8 @@ async def run_agent_loop_continue(
9896
should_stop_after_turn: Callable | None = None,
9997
get_steering_messages: Callable | None = None,
10098
get_follow_up_messages: Callable | None = None,
101-
thinking: ThinkingLevel = "off",
99+
stream_options: StreamOptions | None = None,
102100
tool_execution: str = "parallel",
103-
signal: asyncio.Event | None = None,
104101
system_prompt: str | None = None,
105102
) -> list[Any]:
106103
if not context.messages:
@@ -130,9 +127,8 @@ async def run_agent_loop_continue(
130127
should_stop_after_turn=should_stop_after_turn,
131128
get_steering_messages=get_steering_messages,
132129
get_follow_up_messages=get_follow_up_messages,
133-
thinking=thinking,
130+
stream_options=stream_options,
134131
tool_execution=tool_execution,
135-
signal=signal,
136132
emit=emit,
137133
)
138134
return new_messages
@@ -151,11 +147,11 @@ async def _run_loop(
151147
should_stop_after_turn: Callable | None,
152148
get_steering_messages: Callable | None,
153149
get_follow_up_messages: Callable | None,
154-
thinking: ThinkingLevel,
150+
stream_options: StreamOptions | None,
155151
tool_execution: str,
156-
signal: asyncio.Event | None,
157152
emit: Callable,
158153
) -> None:
154+
opts = stream_options or StreamOptions()
159155
first_turn = True
160156

161157
while True:
@@ -173,8 +169,7 @@ async def _run_loop(
173169
model,
174170
convert_to_llm,
175171
transform_context,
176-
thinking,
177-
signal,
172+
opts,
178173
emit,
179174
)
180175
new_messages.append(message)
@@ -195,7 +190,7 @@ async def _run_loop(
195190
tool_execution=tool_execution,
196191
before_tool_call=before_tool_call,
197192
after_tool_call=after_tool_call,
198-
signal=signal,
193+
signal=opts.signal,
199194
emit=emit,
200195
)
201196
tool_results = batch.messages
@@ -251,13 +246,12 @@ async def _stream_assistant_response(
251246
model: Model,
252247
convert_to_llm: Callable,
253248
transform_context: Callable | None,
254-
thinking: ThinkingLevel,
255-
signal: asyncio.Event | None,
249+
options: StreamOptions,
256250
emit: Callable,
257251
) -> AssistantMessage:
258252
messages = context.messages
259253
if transform_context:
260-
messages = await transform_context(messages, signal=signal)
254+
messages = await transform_context(messages, signal=options.signal)
261255

262256
llm_messages = convert_to_llm(messages)
263257
if asyncio.iscoroutine(llm_messages):
@@ -272,8 +266,7 @@ async def _stream_assistant_response(
272266
llm_messages,
273267
system_prompt=context.system_prompt,
274268
tools=tools_defs,
275-
thinking=thinking,
276-
signal=signal,
269+
options=options,
277270
)
278271

279272
partial_message: AssistantMessage | None = None

cubepi/providers/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111
Provider,
1212
ProviderResponse,
1313
StreamEvent,
14+
StreamOptions,
1415
TextContent,
1516
ThinkingBudgets,
1617
ThinkingContent,
@@ -70,6 +71,7 @@ def get_openai_responses_provider():
7071
"Provider",
7172
"ProviderResponse",
7273
"StreamEvent",
74+
"StreamOptions",
7375
"TextContent",
7476
"ThinkingBudgets",
7577
"ThinkingContent",

cubepi/providers/anthropic.py

Lines changed: 8 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -10,14 +10,11 @@
1010
Message,
1111
MessageStream,
1212
Model,
13-
OnPayloadCallback,
14-
OnResponseCallback,
1513
ProviderResponse,
1614
StreamEvent,
15+
StreamOptions,
1716
TextContent,
18-
ThinkingBudgets,
1917
ThinkingContent,
20-
ThinkingLevel,
2118
ToolCall,
2219
ToolDefinition,
2320
ToolResultMessage,
@@ -51,14 +48,11 @@ async def stream(
5148
*,
5249
system_prompt: str = "",
5350
tools: list[ToolDefinition] | None = None,
54-
thinking: ThinkingLevel = "off",
55-
thinking_budgets: ThinkingBudgets | None = None,
56-
signal: asyncio.Event | None = None,
57-
on_payload: OnPayloadCallback | None = None,
58-
on_response: OnResponseCallback | None = None,
51+
options: StreamOptions | None = None,
5952
) -> MessageStream:
53+
opts = options or StreamOptions()
6054
ms = MessageStream()
61-
thinking = clamp_thinking_level(model, thinking)
55+
thinking = clamp_thinking_level(model, opts.thinking)
6256

6357
cache_control = self._get_cache_control()
6458
api_messages = [self._convert_message(m) for m in messages]
@@ -69,7 +63,7 @@ async def stream(
6963
base_max_tokens=model.max_tokens,
7064
model_max_tokens=model.context_window,
7165
reasoning_level=thinking,
72-
custom_budgets=thinking_budgets,
66+
custom_budgets=opts.thinking_budgets,
7367
)
7468

7569
kwargs: dict[str, Any] = {
@@ -99,14 +93,14 @@ async def stream(
9993
async def _produce() -> None:
10094
try:
10195
nonlocal kwargs
102-
kwargs = await _invoke_on_payload(on_payload, kwargs, model)
96+
kwargs = await _invoke_on_payload(opts.on_payload, kwargs, model)
10397

10498
async with self._client.messages.stream(**kwargs) as stream:
10599
# Invoke on_response with HTTP metadata if available
106100
http_response = getattr(stream, "response", None)
107101
if http_response is not None:
108102
await _invoke_on_response(
109-
on_response,
103+
opts.on_response,
110104
ProviderResponse(
111105
status=http_response.status_code,
112106
headers=dict(http_response.headers),
@@ -122,7 +116,7 @@ async def _produce() -> None:
122116
)
123117

124118
async for event in stream:
125-
if signal and signal.is_set():
119+
if opts.signal and opts.signal.is_set():
126120
aborted = partial.model_copy(
127121
update={
128122
"stop_reason": "aborted",

cubepi/providers/base.py

Lines changed: 14 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,7 @@
1313
runtime_checkable,
1414
)
1515

16-
from pydantic import BaseModel
16+
from pydantic import BaseModel, ConfigDict
1717

1818
ThinkingLevel = Literal["off", "minimal", "low", "medium", "high", "xhigh"]
1919

@@ -210,6 +210,18 @@ class ProviderResponse:
210210
"""Optional callback invoked after an HTTP response is received."""
211211

212212

213+
class StreamOptions(BaseModel):
214+
"""Options bag for Provider.stream(), transparent to the agent loop."""
215+
216+
model_config = ConfigDict(arbitrary_types_allowed=True)
217+
218+
thinking: ThinkingLevel = "off"
219+
thinking_budgets: ThinkingBudgets | None = None
220+
signal: asyncio.Event | None = None
221+
on_payload: OnPayloadCallback | None = None
222+
on_response: OnResponseCallback | None = None
223+
224+
213225
async def _invoke_on_payload(
214226
callback: OnPayloadCallback | None,
215227
payload: dict,
@@ -246,9 +258,5 @@ async def stream(
246258
*,
247259
system_prompt: str = "",
248260
tools: list[ToolDefinition] | None = None,
249-
thinking: ThinkingLevel = "off",
250-
thinking_budgets: ThinkingBudgets | None = None,
251-
signal: asyncio.Event | None = None,
252-
on_payload: OnPayloadCallback | None = None,
253-
on_response: OnResponseCallback | None = None,
261+
options: StreamOptions | None = None,
254262
) -> MessageStream: ...

cubepi/providers/faux.py

Lines changed: 4 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -13,10 +13,9 @@
1313
MessageStream,
1414
Model,
1515
StreamEvent,
16+
StreamOptions,
1617
TextContent,
17-
ThinkingBudgets,
1818
ThinkingContent,
19-
ThinkingLevel,
2019
ToolCall,
2120
ToolDefinition,
2221
Usage,
@@ -223,12 +222,9 @@ async def stream(
223222
*,
224223
system_prompt: str = "",
225224
tools: list[ToolDefinition] | None = None,
226-
thinking: ThinkingLevel = "off",
227-
thinking_budgets: ThinkingBudgets | None = None,
228-
signal: asyncio.Event | None = None,
229-
on_payload: Any = None,
230-
on_response: Any = None,
225+
options: StreamOptions | None = None,
231226
) -> MessageStream:
227+
opts = options or StreamOptions()
232228
ms = MessageStream()
233229
self.call_count += 1
234230

@@ -273,7 +269,7 @@ async def _produce() -> None:
273269
)
274270
resolved = resolved.model_copy(update={"usage": cache_usage})
275271

276-
await self._stream_with_deltas(ms, resolved, signal)
272+
await self._stream_with_deltas(ms, resolved, opts.signal)
277273
except BaseException as exc:
278274
error_msg = AssistantMessage(
279275
content=[],

0 commit comments

Comments
 (0)