agenthub-python 0.3.0__py3-none-any.whl → 0.3.2__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.
- agenthub/abort_signal.py +135 -0
- agenthub/auto_client.py +34 -15
- agenthub/base_client.py +96 -13
- agenthub/claude4_6/client.py +28 -14
- agenthub/claude4_8/__init__.py +18 -0
- agenthub/claude4_8/client.py +429 -0
- agenthub/deepseek_v4/__init__.py +18 -0
- agenthub/deepseek_v4/client.py +337 -0
- agenthub/gemini3/client.py +117 -2
- agenthub/{glm5 → glm5_1}/__init__.py +2 -2
- agenthub/{glm5 → glm5_1}/client.py +13 -12
- agenthub/{gpt5_4 → gpt5_5}/__init__.py +2 -2
- agenthub/{gpt5_4 → gpt5_5}/client.py +18 -18
- agenthub/integration/playground.py +617 -71
- agenthub/integration/tracer.py +360 -102
- agenthub/{kimi_k2_5 → kimi_k2_6}/__init__.py +2 -2
- agenthub/{kimi_k2_5 → kimi_k2_6}/client.py +14 -13
- agenthub/{qwen3 → openai}/__init__.py +2 -2
- agenthub/{qwen3 → openai}/client.py +92 -95
- agenthub/types.py +57 -1
- agenthub_python-0.3.2.dist-info/METADATA +351 -0
- agenthub_python-0.3.2.dist-info/RECORD +28 -0
- {agenthub_python-0.3.0.dist-info → agenthub_python-0.3.2.dist-info}/WHEEL +1 -1
- agenthub_python-0.3.0.dist-info/METADATA +0 -10
- agenthub_python-0.3.0.dist-info/RECORD +0 -23
agenthub/abort_signal.py
ADDED
|
@@ -0,0 +1,135 @@
|
|
|
1
|
+
# Copyright 2025 Prism Shadow. and/or its affiliates
|
|
2
|
+
#
|
|
3
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
+
# you may not use this file except in compliance with the License.
|
|
5
|
+
# You may obtain a copy of the License at
|
|
6
|
+
#
|
|
7
|
+
# http://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
+
#
|
|
9
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
+
# See the License for the specific language governing permissions and
|
|
13
|
+
# limitations under the License.
|
|
14
|
+
|
|
15
|
+
import asyncio
|
|
16
|
+
import threading
|
|
17
|
+
from contextlib import suppress
|
|
18
|
+
from typing import Any, Awaitable, TypeVar
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
T = TypeVar("T")
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class AbortSignal:
|
|
25
|
+
"""Abort signal that can also trigger its own aborted state."""
|
|
26
|
+
|
|
27
|
+
def __init__(self) -> None:
|
|
28
|
+
self._lock = threading.Lock()
|
|
29
|
+
self._aborted = False
|
|
30
|
+
self._reason: Any = None
|
|
31
|
+
self._waiters: set[asyncio.Future[None]] = set()
|
|
32
|
+
|
|
33
|
+
@property
|
|
34
|
+
def aborted(self) -> bool:
|
|
35
|
+
with self._lock:
|
|
36
|
+
return self._aborted
|
|
37
|
+
|
|
38
|
+
@property
|
|
39
|
+
def reason(self) -> Any:
|
|
40
|
+
with self._lock:
|
|
41
|
+
return self._reason
|
|
42
|
+
|
|
43
|
+
def abort(self, reason: Any = None) -> None:
|
|
44
|
+
with self._lock:
|
|
45
|
+
if self._aborted:
|
|
46
|
+
return
|
|
47
|
+
|
|
48
|
+
self._aborted = True
|
|
49
|
+
self._reason = reason
|
|
50
|
+
waiters = tuple(self._waiters)
|
|
51
|
+
self._waiters.clear()
|
|
52
|
+
|
|
53
|
+
for waiter in waiters:
|
|
54
|
+
_notify_waiter(waiter)
|
|
55
|
+
|
|
56
|
+
async def wait(self) -> None:
|
|
57
|
+
loop = asyncio.get_running_loop()
|
|
58
|
+
waiter = loop.create_future()
|
|
59
|
+
with self._lock:
|
|
60
|
+
if self._aborted:
|
|
61
|
+
return
|
|
62
|
+
|
|
63
|
+
self._waiters.add(waiter)
|
|
64
|
+
|
|
65
|
+
try:
|
|
66
|
+
await waiter
|
|
67
|
+
finally:
|
|
68
|
+
with self._lock:
|
|
69
|
+
self._waiters.discard(waiter)
|
|
70
|
+
|
|
71
|
+
def throw_if_aborted(self) -> None:
|
|
72
|
+
with self._lock:
|
|
73
|
+
aborted = self._aborted
|
|
74
|
+
reason = self._reason
|
|
75
|
+
|
|
76
|
+
if aborted:
|
|
77
|
+
raise _cancelled_error(reason)
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
async def run_with_abort(awaitable: Awaitable[T], signal: AbortSignal) -> T:
|
|
81
|
+
"""Run an awaitable and cancel it when the signal is aborted."""
|
|
82
|
+
|
|
83
|
+
task = asyncio.ensure_future(awaitable)
|
|
84
|
+
|
|
85
|
+
if signal.aborted:
|
|
86
|
+
task.cancel(signal.reason)
|
|
87
|
+
with suppress(asyncio.CancelledError):
|
|
88
|
+
await task
|
|
89
|
+
raise _cancelled_error(signal.reason)
|
|
90
|
+
|
|
91
|
+
abort_task = asyncio.create_task(signal.wait())
|
|
92
|
+
|
|
93
|
+
try:
|
|
94
|
+
done, _ = await asyncio.wait((task, abort_task), return_when=asyncio.FIRST_COMPLETED)
|
|
95
|
+
if task in done:
|
|
96
|
+
return await task
|
|
97
|
+
|
|
98
|
+
task.cancel(signal.reason)
|
|
99
|
+
with suppress(asyncio.CancelledError):
|
|
100
|
+
await task
|
|
101
|
+
raise _cancelled_error(signal.reason)
|
|
102
|
+
except asyncio.CancelledError:
|
|
103
|
+
if not task.done():
|
|
104
|
+
task.cancel()
|
|
105
|
+
with suppress(asyncio.CancelledError):
|
|
106
|
+
await task
|
|
107
|
+
raise
|
|
108
|
+
finally:
|
|
109
|
+
if not abort_task.done():
|
|
110
|
+
abort_task.cancel()
|
|
111
|
+
with suppress(asyncio.CancelledError):
|
|
112
|
+
await abort_task
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
def _set_waiter_result(waiter: asyncio.Future[None]) -> None:
|
|
116
|
+
if not waiter.done():
|
|
117
|
+
waiter.set_result(None)
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
def _notify_waiter(waiter: asyncio.Future[None]) -> None:
|
|
121
|
+
loop = waiter.get_loop()
|
|
122
|
+
if loop.is_closed():
|
|
123
|
+
return
|
|
124
|
+
|
|
125
|
+
if loop.is_running():
|
|
126
|
+
loop.call_soon_threadsafe(_set_waiter_result, waiter)
|
|
127
|
+
else:
|
|
128
|
+
_set_waiter_result(waiter)
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
def _cancelled_error(reason: Any) -> asyncio.CancelledError:
|
|
132
|
+
if reason is None:
|
|
133
|
+
return asyncio.CancelledError()
|
|
134
|
+
|
|
135
|
+
return asyncio.CancelledError(reason)
|
agenthub/auto_client.py
CHANGED
|
@@ -15,6 +15,7 @@
|
|
|
15
15
|
import os
|
|
16
16
|
from typing import Any, AsyncIterator
|
|
17
17
|
|
|
18
|
+
from .abort_signal import AbortSignal
|
|
18
19
|
from .base_client import LLMClient
|
|
19
20
|
from .types import UniConfig, UniEvent, UniMessage
|
|
20
21
|
|
|
@@ -45,35 +46,45 @@ class AutoLLMClient(LLMClient):
|
|
|
45
46
|
self, model: str, api_key: str | None = None, base_url: str | None = None, client_type: str | None = None
|
|
46
47
|
) -> LLMClient:
|
|
47
48
|
"""Create the appropriate client for the given model."""
|
|
48
|
-
client_type = client_type or os.getenv("CLIENT_TYPE", model.lower()
|
|
49
|
-
if
|
|
49
|
+
client_type = (client_type or os.getenv("CLIENT_TYPE", model)).lower()
|
|
50
|
+
if any(
|
|
51
|
+
prefix in client_type for prefix in ("gemini-3", "gemini-embedding")
|
|
52
|
+
): # e.g., gemini-3-flash-preview, gemini-embedding-2
|
|
50
53
|
from .gemini3 import Gemini3Client
|
|
51
54
|
|
|
52
55
|
return Gemini3Client(model=model, api_key=api_key, base_url=base_url)
|
|
56
|
+
elif "claude" in client_type and ("4-7" in client_type or "4-8" in client_type): # e.g., claude-opus-4-7
|
|
57
|
+
from .claude4_8 import Claude4_8Client
|
|
58
|
+
|
|
59
|
+
return Claude4_8Client(model=model, api_key=api_key, base_url=base_url)
|
|
53
60
|
elif "claude" in client_type and "4-6" in client_type: # e.g., claude-sonnet-4-6
|
|
54
61
|
from .claude4_6 import Claude4_6Client
|
|
55
62
|
|
|
56
63
|
return Claude4_6Client(model=model, api_key=api_key, base_url=base_url)
|
|
57
|
-
elif "gpt-5.4" in client_type: # e.g., gpt-5.
|
|
58
|
-
from .
|
|
64
|
+
elif "gpt-5.4" in client_type or "gpt-5.5" in client_type: # e.g., gpt-5.5
|
|
65
|
+
from .gpt5_5 import GPT5_5Client
|
|
66
|
+
|
|
67
|
+
return GPT5_5Client(model=model, api_key=api_key, base_url=base_url)
|
|
68
|
+
elif "glm-5" in client_type or "glm-5.1" in client_type:
|
|
69
|
+
from .glm5_1 import GLM5_1Client
|
|
59
70
|
|
|
60
|
-
return
|
|
61
|
-
elif "
|
|
62
|
-
from .
|
|
71
|
+
return GLM5_1Client(model=model, api_key=api_key, base_url=base_url)
|
|
72
|
+
elif "kimi-k2.5" in client_type or "kimi-k2.6" in client_type:
|
|
73
|
+
from .kimi_k2_6 import KimiK2_6Client
|
|
63
74
|
|
|
64
|
-
return
|
|
65
|
-
elif "
|
|
66
|
-
from .
|
|
75
|
+
return KimiK2_6Client(model=model, api_key=api_key, base_url=base_url)
|
|
76
|
+
elif "deepseek-v4" in client_type:
|
|
77
|
+
from .deepseek_v4 import DeepSeekV4Client
|
|
67
78
|
|
|
68
|
-
return
|
|
69
|
-
elif "
|
|
70
|
-
from .
|
|
79
|
+
return DeepSeekV4Client(model=model, api_key=api_key, base_url=base_url)
|
|
80
|
+
elif "openai" in client_type:
|
|
81
|
+
from .openai import OpenaiClient
|
|
71
82
|
|
|
72
|
-
return
|
|
83
|
+
return OpenaiClient(model=model, api_key=api_key, base_url=base_url)
|
|
73
84
|
else:
|
|
74
85
|
raise ValueError(
|
|
75
86
|
f"{client_type} is not supported. "
|
|
76
|
-
"Supported client types: gemini-3, claude-4-6, gpt-5.4, glm-5, kimi-k2.5,
|
|
87
|
+
"Supported client types: gemini-3, claude-4-8, claude-4-7, claude-4-6, gpt-5.5, gpt-5.4, glm-5.1, kimi-k2.6, kimi-k2.5, deepseek-v4, openai."
|
|
77
88
|
)
|
|
78
89
|
|
|
79
90
|
def transform_uni_config_to_model_config(self, config: UniConfig) -> Any:
|
|
@@ -99,11 +110,13 @@ class AutoLLMClient(LLMClient):
|
|
|
99
110
|
self,
|
|
100
111
|
messages: list[UniMessage],
|
|
101
112
|
config: UniConfig,
|
|
113
|
+
signal: AbortSignal | None = None,
|
|
102
114
|
) -> AsyncIterator[UniEvent]:
|
|
103
115
|
"""Route to underlying client's streaming_response."""
|
|
104
116
|
async for event in self._client.streaming_response(
|
|
105
117
|
messages=messages,
|
|
106
118
|
config=config,
|
|
119
|
+
signal=signal,
|
|
107
120
|
):
|
|
108
121
|
yield event
|
|
109
122
|
|
|
@@ -111,11 +124,13 @@ class AutoLLMClient(LLMClient):
|
|
|
111
124
|
self,
|
|
112
125
|
message: UniMessage,
|
|
113
126
|
config: UniConfig,
|
|
127
|
+
signal: AbortSignal | None = None,
|
|
114
128
|
) -> AsyncIterator[UniEvent]:
|
|
115
129
|
"""Route to underlying client's streaming_response_stateful."""
|
|
116
130
|
async for event in self._client.streaming_response_stateful(
|
|
117
131
|
message=message,
|
|
118
132
|
config=config,
|
|
133
|
+
signal=signal,
|
|
119
134
|
):
|
|
120
135
|
yield event
|
|
121
136
|
|
|
@@ -126,3 +141,7 @@ class AutoLLMClient(LLMClient):
|
|
|
126
141
|
def get_history(self) -> list[UniMessage]:
|
|
127
142
|
"""Get history from the underlying client."""
|
|
128
143
|
return self._client.get_history()
|
|
144
|
+
|
|
145
|
+
def set_history(self, history: list[UniMessage]) -> None:
|
|
146
|
+
"""Set history in the underlying client."""
|
|
147
|
+
self._client.set_history(history)
|
agenthub/base_client.py
CHANGED
|
@@ -12,10 +12,21 @@
|
|
|
12
12
|
# See the License for the specific language governing permissions and
|
|
13
13
|
# limitations under the License.
|
|
14
14
|
|
|
15
|
+
import asyncio
|
|
16
|
+
import time
|
|
15
17
|
from abc import ABC, abstractmethod
|
|
18
|
+
from contextlib import suppress
|
|
16
19
|
from typing import Any, AsyncIterator
|
|
17
20
|
|
|
18
|
-
from .
|
|
21
|
+
from .abort_signal import AbortSignal
|
|
22
|
+
from .types import (
|
|
23
|
+
ContentItem,
|
|
24
|
+
FinishReason,
|
|
25
|
+
UniConfig,
|
|
26
|
+
UniEvent,
|
|
27
|
+
UniMessage,
|
|
28
|
+
UsageMetadata,
|
|
29
|
+
)
|
|
19
30
|
|
|
20
31
|
|
|
21
32
|
class LLMClient(ABC):
|
|
@@ -84,6 +95,7 @@ class LLMClient(ABC):
|
|
|
84
95
|
content_items: list[ContentItem] = []
|
|
85
96
|
usage_metadata: UsageMetadata | None = None
|
|
86
97
|
finish_reason: FinishReason | None = None
|
|
98
|
+
created_at: int | None = None
|
|
87
99
|
|
|
88
100
|
for event in events:
|
|
89
101
|
# Merge content_items from all events
|
|
@@ -119,12 +131,14 @@ class LLMClient(ABC):
|
|
|
119
131
|
|
|
120
132
|
usage_metadata = event.get("usage_metadata") # usage_metadata is taken from the last event
|
|
121
133
|
finish_reason = event.get("finish_reason") # finish_reason is taken from the last event
|
|
134
|
+
created_at = event.get("created_at") # created_at is taken from the last event
|
|
122
135
|
|
|
123
136
|
return {
|
|
124
137
|
"role": "assistant",
|
|
125
138
|
"content_items": content_items,
|
|
126
139
|
"usage_metadata": usage_metadata,
|
|
127
140
|
"finish_reason": finish_reason,
|
|
141
|
+
"created_at": created_at,
|
|
128
142
|
}
|
|
129
143
|
|
|
130
144
|
@abstractmethod
|
|
@@ -152,6 +166,7 @@ class LLMClient(ABC):
|
|
|
152
166
|
self,
|
|
153
167
|
messages: list[UniMessage],
|
|
154
168
|
config: UniConfig,
|
|
169
|
+
signal: AbortSignal | None = None,
|
|
155
170
|
) -> AsyncIterator[UniEvent]:
|
|
156
171
|
"""
|
|
157
172
|
Generate content in streaming mode (stateless).
|
|
@@ -163,21 +178,86 @@ class LLMClient(ABC):
|
|
|
163
178
|
Args:
|
|
164
179
|
messages: List of universal message dictionaries containing conversation history
|
|
165
180
|
config: Universal configuration dict
|
|
181
|
+
signal: Optional abort signal used to cancel the active request
|
|
166
182
|
|
|
167
183
|
Yields:
|
|
168
184
|
Universal events from the streaming response
|
|
169
185
|
"""
|
|
186
|
+
# Stamp any messages that don't yet have a created_at timestamp
|
|
187
|
+
for msg in messages:
|
|
188
|
+
if "created_at" not in msg:
|
|
189
|
+
msg["created_at"] = int(time.time() * 1000)
|
|
190
|
+
|
|
170
191
|
last_event: UniEvent | None = None
|
|
171
|
-
|
|
172
|
-
|
|
173
|
-
|
|
192
|
+
events = []
|
|
193
|
+
if signal is not None:
|
|
194
|
+
signal.throw_if_aborted()
|
|
195
|
+
|
|
196
|
+
stream = self._streaming_response_internal(messages, config)
|
|
197
|
+
abort_task: asyncio.Task[None] | None = None
|
|
198
|
+
waiting_for_stream = False
|
|
199
|
+
if signal is not None:
|
|
200
|
+
streaming_task = asyncio.current_task()
|
|
201
|
+
abort_task = asyncio.create_task(signal.wait())
|
|
202
|
+
|
|
203
|
+
def cancel_streaming_task(task: asyncio.Task[None]) -> None:
|
|
204
|
+
if (
|
|
205
|
+
task.cancelled()
|
|
206
|
+
or not signal.aborted
|
|
207
|
+
or not waiting_for_stream
|
|
208
|
+
or streaming_task is None
|
|
209
|
+
or streaming_task.done()
|
|
210
|
+
):
|
|
211
|
+
return
|
|
212
|
+
|
|
213
|
+
streaming_task.cancel(signal.reason)
|
|
214
|
+
|
|
215
|
+
abort_task.add_done_callback(cancel_streaming_task)
|
|
216
|
+
|
|
217
|
+
try:
|
|
218
|
+
while True:
|
|
219
|
+
try:
|
|
220
|
+
if signal is not None:
|
|
221
|
+
signal.throw_if_aborted()
|
|
222
|
+
waiting_for_stream = True
|
|
223
|
+
signal.throw_if_aborted()
|
|
224
|
+
|
|
225
|
+
event = await anext(stream)
|
|
226
|
+
except StopAsyncIteration:
|
|
227
|
+
break
|
|
228
|
+
except asyncio.CancelledError:
|
|
229
|
+
if signal is not None and signal.aborted:
|
|
230
|
+
signal.throw_if_aborted()
|
|
231
|
+
raise
|
|
232
|
+
finally:
|
|
233
|
+
waiting_for_stream = False
|
|
234
|
+
|
|
235
|
+
event["created_at"] = int(time.time() * 1000)
|
|
236
|
+
last_event = event
|
|
237
|
+
events.append(event)
|
|
238
|
+
yield event
|
|
239
|
+
finally:
|
|
240
|
+
if abort_task is not None and not abort_task.done():
|
|
241
|
+
abort_task.cancel()
|
|
242
|
+
with suppress(asyncio.CancelledError):
|
|
243
|
+
await abort_task
|
|
244
|
+
await stream.aclose()
|
|
174
245
|
|
|
175
246
|
self._validate_last_event(last_event)
|
|
176
247
|
|
|
248
|
+
# Save history to file if trace_id is specified
|
|
249
|
+
if config.get("trace_id") and events:
|
|
250
|
+
from .integration.tracer import Tracer
|
|
251
|
+
|
|
252
|
+
assistant_message = self.concat_uni_events_to_uni_message(events)
|
|
253
|
+
tracer = Tracer()
|
|
254
|
+
tracer.save_history(self._model, messages + [assistant_message], config["trace_id"], config)
|
|
255
|
+
|
|
177
256
|
async def streaming_response_stateful(
|
|
178
257
|
self,
|
|
179
258
|
message: UniMessage,
|
|
180
259
|
config: UniConfig,
|
|
260
|
+
signal: AbortSignal | None = None,
|
|
181
261
|
) -> AsyncIterator[UniEvent]:
|
|
182
262
|
"""
|
|
183
263
|
Generate content in streaming mode (stateful).
|
|
@@ -189,6 +269,7 @@ class LLMClient(ABC):
|
|
|
189
269
|
Args:
|
|
190
270
|
message: Latest universal message dictionary to add to conversation
|
|
191
271
|
config: Universal configuration dict
|
|
272
|
+
signal: Optional abort signal used to cancel the active request
|
|
192
273
|
|
|
193
274
|
Yields:
|
|
194
275
|
Universal events from the streaming response
|
|
@@ -198,23 +279,17 @@ class LLMClient(ABC):
|
|
|
198
279
|
|
|
199
280
|
# Collect all events for history
|
|
200
281
|
events = []
|
|
201
|
-
async for event in self.streaming_response(messages=temp_messages, config=config):
|
|
282
|
+
async for event in self.streaming_response(messages=temp_messages, config=config, signal=signal):
|
|
202
283
|
events.append(event)
|
|
203
284
|
yield event
|
|
204
285
|
|
|
205
286
|
# Only update history after successful inference
|
|
287
|
+
# temp_messages[-1] is the user message, now stamped with created_at by streaming_response
|
|
206
288
|
if events:
|
|
207
289
|
assistant_message = self.concat_uni_events_to_uni_message(events)
|
|
208
|
-
self._history.append(
|
|
290
|
+
self._history.append(temp_messages[-1])
|
|
209
291
|
self._history.append(assistant_message)
|
|
210
292
|
|
|
211
|
-
# Save history to file if trace_id is specified
|
|
212
|
-
if config.get("trace_id"):
|
|
213
|
-
from .integration.tracer import Tracer
|
|
214
|
-
|
|
215
|
-
tracer = Tracer()
|
|
216
|
-
tracer.save_history(self._model, self._history, config["trace_id"], config)
|
|
217
|
-
|
|
218
293
|
@staticmethod
|
|
219
294
|
def _validate_last_event(last_event: UniEvent | None) -> None:
|
|
220
295
|
"""Validate that the last event has usage_metadata and finish_reason.
|
|
@@ -244,3 +319,11 @@ class LLMClient(ABC):
|
|
|
244
319
|
def get_history(self) -> list[UniMessage]:
|
|
245
320
|
"""Get the current message history."""
|
|
246
321
|
return self._history.copy()
|
|
322
|
+
|
|
323
|
+
def set_history(self, history: list[UniMessage]) -> None:
|
|
324
|
+
"""Replace the message history with a copy of the provided history.
|
|
325
|
+
|
|
326
|
+
Args:
|
|
327
|
+
history: List of universal message dictionaries to set as the new history
|
|
328
|
+
"""
|
|
329
|
+
self._history = list(history)
|
agenthub/claude4_6/client.py
CHANGED
|
@@ -110,6 +110,7 @@ class Claude4_6Client(LLMClient):
|
|
|
110
110
|
ThinkingLevel.LOW: {"thinking": {"type": "adaptive"}, "output_config": {"effort": "low"}},
|
|
111
111
|
ThinkingLevel.MEDIUM: {"thinking": {"type": "adaptive"}, "output_config": {"effort": "medium"}},
|
|
112
112
|
ThinkingLevel.HIGH: {"thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}},
|
|
113
|
+
ThinkingLevel.XHIGH: {"thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}},
|
|
113
114
|
}
|
|
114
115
|
return mapping.get(thinking_level)
|
|
115
116
|
|
|
@@ -145,7 +146,7 @@ class Claude4_6Client(LLMClient):
|
|
|
145
146
|
if config.get("max_tokens") is not None:
|
|
146
147
|
claude_config["max_tokens"] = config["max_tokens"]
|
|
147
148
|
else:
|
|
148
|
-
claude_config["max_tokens"] =
|
|
149
|
+
claude_config["max_tokens"] = 64000 # Claude requires max_tokens to be specified
|
|
149
150
|
|
|
150
151
|
if config.get("temperature") is not None:
|
|
151
152
|
claude_config["temperature"] = config["temperature"]
|
|
@@ -171,6 +172,15 @@ class Claude4_6Client(LLMClient):
|
|
|
171
172
|
if config.get("tool_choice") is not None:
|
|
172
173
|
claude_config["tool_choice"] = self._convert_tool_choice(config["tool_choice"])
|
|
173
174
|
|
|
175
|
+
# Add cache_control if prompt caching is enabled
|
|
176
|
+
# TODO: wait for bedrock to support cache_control in config
|
|
177
|
+
if not self._use_bedrock:
|
|
178
|
+
prompt_caching = config.get("prompt_caching", PromptCaching.ENABLE)
|
|
179
|
+
if prompt_caching == PromptCaching.ENABLE:
|
|
180
|
+
claude_config["cache_control"] = {"type": "ephemeral"}
|
|
181
|
+
elif prompt_caching == PromptCaching.ENHANCE:
|
|
182
|
+
claude_config["cache_control"] = {"type": "ephemeral", "ttl": "1h"}
|
|
183
|
+
|
|
174
184
|
return claude_config
|
|
175
185
|
|
|
176
186
|
async def transform_uni_message_to_model_input(self, messages: list[UniMessage]) -> list[BetaMessageParam]:
|
|
@@ -334,18 +344,20 @@ class Claude4_6Client(LLMClient):
|
|
|
334
344
|
# Use unified message conversion
|
|
335
345
|
claude_messages = await self.transform_uni_message_to_model_input(messages)
|
|
336
346
|
|
|
337
|
-
# Add cache_control to last user message's last item if enabled
|
|
338
|
-
|
|
339
|
-
if
|
|
340
|
-
|
|
341
|
-
|
|
342
|
-
|
|
343
|
-
|
|
344
|
-
|
|
345
|
-
"
|
|
346
|
-
|
|
347
|
-
|
|
348
|
-
|
|
347
|
+
# Add cache_control to last user message's last item if using bedrock and enabled prompt caching
|
|
348
|
+
# TODO: remove after bedrock supports cache_control in config
|
|
349
|
+
if self._use_bedrock:
|
|
350
|
+
prompt_caching = config.get("prompt_caching", PromptCaching.ENABLE)
|
|
351
|
+
if prompt_caching != PromptCaching.DISABLE and claude_messages:
|
|
352
|
+
try:
|
|
353
|
+
last_user_message = next(filter(lambda x: x["role"] == "user", claude_messages[::-1]))
|
|
354
|
+
last_content_item = last_user_message["content"][-1]
|
|
355
|
+
last_content_item["cache_control"] = {
|
|
356
|
+
"type": "ephemeral",
|
|
357
|
+
"ttl": "1h" if prompt_caching == PromptCaching.ENHANCE else "5m",
|
|
358
|
+
}
|
|
359
|
+
except StopIteration:
|
|
360
|
+
pass
|
|
349
361
|
|
|
350
362
|
# Stream generate
|
|
351
363
|
partial_tool_call = {}
|
|
@@ -362,7 +374,9 @@ class Claude4_6Client(LLMClient):
|
|
|
362
374
|
"arguments": "",
|
|
363
375
|
"tool_call_id": item["tool_call_id"],
|
|
364
376
|
}
|
|
365
|
-
|
|
377
|
+
|
|
378
|
+
if event["content_items"]:
|
|
379
|
+
yield event
|
|
366
380
|
|
|
367
381
|
if event["usage_metadata"] is not None:
|
|
368
382
|
# initialize partial_usage
|
|
@@ -0,0 +1,18 @@
|
|
|
1
|
+
# Copyright 2025 Prism Shadow. and/or its affiliates
|
|
2
|
+
#
|
|
3
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
+
# you may not use this file except in compliance with the License.
|
|
5
|
+
# You may obtain a copy of the License at
|
|
6
|
+
#
|
|
7
|
+
# http://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
+
#
|
|
9
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
+
# See the License for the specific language governing permissions and
|
|
13
|
+
# limitations under the License.
|
|
14
|
+
|
|
15
|
+
from .client import Claude4_8Client
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
__all__ = ["Claude4_8Client"]
|