chore: update version to 0.0.3 and refactor ZaiSampler error handling (#24)

Co-authored-by: zhengweijun <weijun.zheng@aminer.cn>
This commit is contained in:
wellenzheng 2025-08-05 17:55:38 +08:00 committed by GitHub
parent 47dc7ca5c0
commit 49ca5af5cb
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 53 additions and 50 deletions

View file

@ -6,7 +6,7 @@ import traceback
from typing import Optional from typing import Optional
from example_types import MessageList, SamplerBase from example_types import MessageList, SamplerBase
from zai import ZhipuAiClient from zai import ZhipuAiClient, ZaiClient
class ZaiSampler(SamplerBase): class ZaiSampler(SamplerBase):
@ -27,57 +27,49 @@ class ZaiSampler(SamplerBase):
self.temperature = temperature self.temperature = temperature
self.max_tokens = max_tokens self.max_tokens = max_tokens
self.model = model self.model = model
self.client = ZhipuAiClient(api_key=api_key) self.client = ZaiClient(api_key=api_key)
self.stream = stream self.stream = stream
def get_resp(self, message_list): def get_resp(self, message_list):
for _ in range(3): try:
try: chat_completion = self.client.chat.completions.create(
chat_completion = self.client.chat.completions.create( messages=message_list,
messages=message_list, model=self.model,
model=self.model, temperature=self.temperature,
temperature=self.temperature, top_p=self.top_p,
top_p=self.top_p, max_tokens=self.max_tokens
max_tokens=self.max_tokens )
) output = chat_completion.choices[0].message.content
output = chat_completion.choices[0].message.content return output
return output except Exception as e:
except Exception as e: print(f"Exception: {e}\nTraceback: {traceback.format_exc()}")
print(f"Exception: {e}\nTraceback: {traceback.format_exc()}") raise
time.sleep(1)
continue
print(f"failed, last exception: {e if 'e' in locals() else ''}")
return ''
def get_resp_stream(self, message_list, top_p=-1, temperature=-1): def get_resp_stream(self, message_list, top_p=-1, temperature=-1):
temperature = temperature if temperature > 0 else self.temperature temperature = temperature if temperature > 0 else self.temperature
top_p = top_p if top_p > 0 else 0.95 top_p = top_p if top_p > 0 else 0.95
final = '' final = ''
for _ in range(200): try:
try: chat_completion_res = self.client.chat.completions.create(
chat_completion_res = self.client.chat.completions.create( model=self.model,
model=self.model, messages=message_list,
messages=message_list, thinking={
thinking={ "type": "enabled",
"type": "enabled", },
}, stream=True,
stream=True, max_tokens=self.max_tokens,
max_tokens=self.max_tokens, temperature=temperature
temperature=temperature )
) for chunk in chat_completion_res:
for chunk in chat_completion_res: if chunk.choices[0].delta.content:
if chunk.choices[0].delta.content: final += chunk.choices[0].delta.content
final += chunk.choices[0].delta.content except Exception as e:
break print(f"Exception: {e}\nTraceback: {traceback.format_exc()}")
except Exception as e: raise
final = ""
print(f"Exception: {e}\nTraceback: {traceback.format_exc()}")
time.sleep(5)
continue
if final == '': if final == '':
print(f"failed in get_resp for 50 times, last exception: {e if 'e' in locals() else ''}") print(f"failed in get_resp, no content received")
return '' return ''
content = '' content = ''
@ -105,9 +97,13 @@ class ZaiSampler(SamplerBase):
if __name__ == "__main__": if __name__ == "__main__":
client = ZaiSampler(model="glm-4.5", api_key=os.getenv("ZAI_API_KEY"), stream=True) try:
messages = [ client = ZaiSampler(model="glm-4.5", api_key=os.getenv("ZAI_API_KEY"), stream=True)
{"role": "user", "content": "Hi?"}, messages = [
] {"role": "user", "content": "Hi? Tell me a joke."},
response = client(messages) ]
print(response) response = client(messages)
print(response)
except Exception as e:
print(f"Fatal error: {e}\nTraceback: {traceback.format_exc()}")
sys.exit(1)

View file

@ -1,6 +1,6 @@
[tool.poetry] [tool.poetry]
name = "zai-sdk" name = "zai-sdk"
version = "0.0.2" version = "0.0.3"
description = "A SDK library for accessing big model apis from Z.ai" description = "A SDK library for accessing big model apis from Z.ai"
authors = ["Z.ai"] authors = ["Z.ai"]
readme = "README.md" readme = "README.md"

View file

@ -218,6 +218,13 @@ class ZaiClient(BaseClient):
def default_base_url(self): def default_base_url(self):
return 'https://api.z.ai/api/paas/v4' return 'https://api.z.ai/api/paas/v4'
@property
@override
def auth_headers(self) -> dict[str, str]:
headers = super().auth_headers
headers['Accept-Language'] = 'en-US,en'
return headers
class ZhipuAiClient(BaseClient): class ZhipuAiClient(BaseClient):
@property @property

View file

@ -1,2 +1,2 @@
__title__ = 'Z.ai' __title__ = 'Z.ai'
__version__ = '0.0.2' __version__ = '0.0.3'