diff --git a/examples/glm4_5_thinking.py b/examples/glm4_5_thinking.py index abe4b2f..e1df7ae 100644 --- a/examples/glm4_5_thinking.py +++ b/examples/glm4_5_thinking.py @@ -6,7 +6,7 @@ import traceback from typing import Optional from example_types import MessageList, SamplerBase -from zai import ZhipuAiClient +from zai import ZhipuAiClient, ZaiClient class ZaiSampler(SamplerBase): @@ -27,57 +27,49 @@ class ZaiSampler(SamplerBase): self.temperature = temperature self.max_tokens = max_tokens self.model = model - self.client = ZhipuAiClient(api_key=api_key) + self.client = ZaiClient(api_key=api_key) self.stream = stream def get_resp(self, message_list): - for _ in range(3): - try: - chat_completion = self.client.chat.completions.create( - messages=message_list, - model=self.model, - temperature=self.temperature, - top_p=self.top_p, - max_tokens=self.max_tokens - ) - output = chat_completion.choices[0].message.content - return output - except Exception as e: - print(f"Exception: {e}\nTraceback: {traceback.format_exc()}") - time.sleep(1) - continue - print(f"failed, last exception: {e if 'e' in locals() else ''}") - return '' + try: + chat_completion = self.client.chat.completions.create( + messages=message_list, + model=self.model, + temperature=self.temperature, + top_p=self.top_p, + max_tokens=self.max_tokens + ) + output = chat_completion.choices[0].message.content + return output + except Exception as e: + print(f"Exception: {e}\nTraceback: {traceback.format_exc()}") + raise def get_resp_stream(self, message_list, top_p=-1, temperature=-1): temperature = temperature if temperature > 0 else self.temperature top_p = top_p if top_p > 0 else 0.95 final = '' - for _ in range(200): - try: - chat_completion_res = self.client.chat.completions.create( - model=self.model, - messages=message_list, - thinking={ - "type": "enabled", - }, - stream=True, - max_tokens=self.max_tokens, - temperature=temperature - ) - for chunk in chat_completion_res: - if chunk.choices[0].delta.content: - final += chunk.choices[0].delta.content - break - except Exception as e: - final = "" - print(f"Exception: {e}\nTraceback: {traceback.format_exc()}") - time.sleep(5) - continue - + try: + chat_completion_res = self.client.chat.completions.create( + model=self.model, + messages=message_list, + thinking={ + "type": "enabled", + }, + stream=True, + max_tokens=self.max_tokens, + temperature=temperature + ) + for chunk in chat_completion_res: + if chunk.choices[0].delta.content: + final += chunk.choices[0].delta.content + except Exception as e: + print(f"Exception: {e}\nTraceback: {traceback.format_exc()}") + raise + 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 '' content = '' @@ -105,9 +97,13 @@ class ZaiSampler(SamplerBase): if __name__ == "__main__": - client = ZaiSampler(model="glm-4.5", api_key=os.getenv("ZAI_API_KEY"), stream=True) - messages = [ - {"role": "user", "content": "Hi?"}, - ] - response = client(messages) - print(response) \ No newline at end of file + try: + client = ZaiSampler(model="glm-4.5", api_key=os.getenv("ZAI_API_KEY"), stream=True) + messages = [ + {"role": "user", "content": "Hi? Tell me a joke."}, + ] + response = client(messages) + print(response) + except Exception as e: + print(f"Fatal error: {e}\nTraceback: {traceback.format_exc()}") + sys.exit(1) \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml index 4f68d63..469c3bb 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "zai-sdk" -version = "0.0.2" +version = "0.0.3" description = "A SDK library for accessing big model apis from Z.ai" authors = ["Z.ai"] readme = "README.md" diff --git a/src/zai/_client.py b/src/zai/_client.py index 7945be3..f7fedf4 100644 --- a/src/zai/_client.py +++ b/src/zai/_client.py @@ -218,6 +218,13 @@ class ZaiClient(BaseClient): def default_base_url(self): 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): @property diff --git a/src/zai/_version.py b/src/zai/_version.py index da36501..9098f97 100644 --- a/src/zai/_version.py +++ b/src/zai/_version.py @@ -1,2 +1,2 @@ __title__ = 'Z.ai' -__version__ = '0.0.2' +__version__ = '0.0.3'