feat: modeldump and usage support (#25)
Co-authored-by: zhengweijun <weijun.zheng@aminer.cn>
This commit is contained in:
parent
49ca5af5cb
commit
82e09d6ee9
15 changed files with 363 additions and 281 deletions
|
|
@ -41,7 +41,7 @@ def parse_function_call(model_response, messages):
|
||||||
tools=tools,
|
tools=tools,
|
||||||
)
|
)
|
||||||
print(response.choices[0].message)
|
print(response.choices[0].message)
|
||||||
messages.append(response.choices[0].message.model_dump())
|
messages.append(response.choices[0].message)
|
||||||
|
|
||||||
messages = []
|
messages = []
|
||||||
tools = [
|
tools = [
|
||||||
|
|
@ -104,7 +104,7 @@ response = client.chat.completions.create(
|
||||||
tools=tools,
|
tools=tools,
|
||||||
)
|
)
|
||||||
print(response.choices[0].message)
|
print(response.choices[0].message)
|
||||||
messages.append(response.choices[0].message.model_dump())
|
messages.append(response.choices[0].message)
|
||||||
|
|
||||||
parse_function_call(response, messages)
|
parse_function_call(response, messages)
|
||||||
|
|
||||||
|
|
@ -115,6 +115,6 @@ response = client.chat.completions.create(
|
||||||
tools=tools,
|
tools=tools,
|
||||||
)
|
)
|
||||||
print(response.choices[0].message)
|
print(response.choices[0].message)
|
||||||
messages.append(response.choices[0].message.model_dump())
|
messages.append(response.choices[0].message)
|
||||||
|
|
||||||
parse_function_call(response, messages)
|
parse_function_call(response, messages)
|
||||||
|
|
@ -1,6 +1,6 @@
|
||||||
[tool.poetry]
|
[tool.poetry]
|
||||||
name = "zai-sdk"
|
name = "zai-sdk"
|
||||||
version = "0.0.3"
|
version = "0.0.3.1"
|
||||||
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"
|
||||||
|
|
|
||||||
|
|
@ -1,2 +1,2 @@
|
||||||
__title__ = 'Z.ai'
|
__title__ = 'Z.ai'
|
||||||
__version__ = '0.0.3'
|
__version__ = '0.0.3.1'
|
||||||
|
|
|
||||||
|
|
@ -18,6 +18,7 @@ from zai.core import (
|
||||||
make_request_options,
|
make_request_options,
|
||||||
maybe_transform,
|
maybe_transform,
|
||||||
)
|
)
|
||||||
|
from zai.core._base_models import BaseModel
|
||||||
from zai.types.chat.chat_completion import Completion
|
from zai.types.chat.chat_completion import Completion
|
||||||
from zai.types.chat.chat_completion_chunk import ChatCompletionChunk
|
from zai.types.chat.chat_completion_chunk import ChatCompletionChunk
|
||||||
from zai.types.chat.code_geex import code_geex_params
|
from zai.types.chat.code_geex import code_geex_params
|
||||||
|
|
@ -112,7 +113,9 @@ class Completions(BaseAPI):
|
||||||
logger.debug(f'temperature:{temperature}, top_p:{top_p}')
|
logger.debug(f'temperature:{temperature}, top_p:{top_p}')
|
||||||
if isinstance(messages, List):
|
if isinstance(messages, List):
|
||||||
for item in messages:
|
for item in messages:
|
||||||
if item.get('content'):
|
if isinstance(item, BaseModel) and hasattr(item, 'content'):
|
||||||
|
item.content = drop_prefix_image_data(item.content)
|
||||||
|
elif isinstance(item, dict) and item.get('content'):
|
||||||
item['content'] = drop_prefix_image_data(item['content'])
|
item['content'] = drop_prefix_image_data(item['content'])
|
||||||
|
|
||||||
body = deepcopy_minimal(
|
body = deepcopy_minimal(
|
||||||
|
|
|
||||||
|
|
@ -160,6 +160,64 @@ class BaseModel(pydantic.BaseModel):
|
||||||
warnings=warnings,
|
warnings=warnings,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def get(self, key: str, default: Any = None) -> Any:
|
||||||
|
"""Get the value of an attribute by name, with an optional default value.
|
||||||
|
|
||||||
|
This method allows you to access model attributes by their string name,
|
||||||
|
similar to how dict.get() works.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
key: The name of the attribute to get
|
||||||
|
default: The value to return if the attribute doesn't exist or is None.
|
||||||
|
Defaults to None.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The value of the attribute if it exists, otherwise the default value.
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
>>> model = MyModel(name="test", age=25)
|
||||||
|
>>> model.get("name") # Returns "test"
|
||||||
|
>>> model.get("nonexistent") # Returns None
|
||||||
|
>>> model.get("nonexistent", "default") # Returns "default"
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
value = getattr(self, key)
|
||||||
|
return value if value is not None else default
|
||||||
|
except AttributeError:
|
||||||
|
return default
|
||||||
|
|
||||||
|
def __json__(self) -> dict[str, Any]:
|
||||||
|
"""Custom JSON serialization method.
|
||||||
|
|
||||||
|
This method is called by JSON encoders that support the __json__ protocol,
|
||||||
|
making BaseModel objects directly serializable without requiring model_dump().
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
A dictionary representation of the model suitable for JSON serialization.
|
||||||
|
"""
|
||||||
|
return self.model_dump(by_alias=True, exclude_unset=True)
|
||||||
|
|
||||||
|
def __reduce_ex__(self, protocol):
|
||||||
|
"""Support for pickle serialization by returning the dict representation."""
|
||||||
|
return (self.__class__.model_validate, (self.model_dump(by_alias=True, exclude_unset=True),))
|
||||||
|
|
||||||
|
def __iter__(self):
|
||||||
|
"""Make BaseModel iterable to support dict() conversion."""
|
||||||
|
data = self.model_dump(by_alias=True, exclude_unset=True)
|
||||||
|
return iter(data.items())
|
||||||
|
|
||||||
|
def keys(self):
|
||||||
|
"""Return keys for dict-like interface."""
|
||||||
|
return self.model_dump(by_alias=True, exclude_unset=True).keys()
|
||||||
|
|
||||||
|
def values(self):
|
||||||
|
"""Return values for dict-like interface."""
|
||||||
|
return self.model_dump(by_alias=True, exclude_unset=True).values()
|
||||||
|
|
||||||
|
def items(self):
|
||||||
|
"""Return items for dict-like interface."""
|
||||||
|
return self.model_dump(by_alias=True, exclude_unset=True).items()
|
||||||
|
|
||||||
@override
|
@override
|
||||||
def __str__(self) -> str:
|
def __str__(self) -> str:
|
||||||
# mypy complains about an invalid self arg
|
# mypy complains about an invalid self arg
|
||||||
|
|
|
||||||
|
|
@ -65,6 +65,7 @@ from ._request_opt import FinalRequestOptions, UserRequestInput
|
||||||
from ._response import APIResponse, BaseAPIResponse, extract_response_type
|
from ._response import APIResponse, BaseAPIResponse, extract_response_type
|
||||||
from ._streaming import StreamResponse
|
from ._streaming import StreamResponse
|
||||||
from ._utils import flatten, is_given, is_mapping
|
from ._utils import flatten, is_given, is_mapping
|
||||||
|
from ._json_encoder import json_dumps
|
||||||
|
|
||||||
log: logging.Logger = logging.getLogger(__name__)
|
log: logging.Logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
@ -355,6 +356,9 @@ class HttpClient:
|
||||||
else:
|
else:
|
||||||
raise RuntimeError(f'Unexpected JSON data type, {type(json_data)}, cannot merge with `extra_body`')
|
raise RuntimeError(f'Unexpected JSON data type, {type(json_data)}, cannot merge with `extra_body`')
|
||||||
|
|
||||||
|
# Convert BaseModel objects to dicts before passing to httpx
|
||||||
|
json_data = self._prepare_json_data(json_data)
|
||||||
|
|
||||||
content_type = headers.get('Content-Type')
|
content_type = headers.get('Content-Type')
|
||||||
# multipart/form-data; boundary=---abc--
|
# multipart/form-data; boundary=---abc--
|
||||||
if headers.get('Content-Type') == 'multipart/form-data':
|
if headers.get('Content-Type') == 'multipart/form-data':
|
||||||
|
|
@ -377,6 +381,24 @@ class HttpClient:
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _prepare_json_data(self, json_data: Any) -> Any:
|
||||||
|
"""Prepare JSON data for httpx by converting BaseModel objects to dicts."""
|
||||||
|
from ._base_models import BaseModel
|
||||||
|
|
||||||
|
if json_data is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
if isinstance(json_data, BaseModel):
|
||||||
|
return json_data.model_dump(by_alias=True, exclude_unset=True)
|
||||||
|
|
||||||
|
if isinstance(json_data, list):
|
||||||
|
return [self._prepare_json_data(item) for item in json_data]
|
||||||
|
|
||||||
|
if isinstance(json_data, dict):
|
||||||
|
return {key: self._prepare_json_data(value) for key, value in json_data.items()}
|
||||||
|
|
||||||
|
return json_data
|
||||||
|
|
||||||
def _object_to_formfata(self, key: str, value: Data | Mapping[object, object]) -> list[tuple[str, str]]:
|
def _object_to_formfata(self, key: str, value: Data | Mapping[object, object]) -> list[tuple[str, str]]:
|
||||||
items = []
|
items = []
|
||||||
|
|
||||||
|
|
|
||||||
44
src/zai/core/_json_encoder.py
Normal file
44
src/zai/core/_json_encoder.py
Normal file
|
|
@ -0,0 +1,44 @@
|
||||||
|
"""JSON encoding utilities for BaseModel objects."""
|
||||||
|
|
||||||
|
import json
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from ._base_models import BaseModel
|
||||||
|
|
||||||
|
|
||||||
|
class ZAIJSONEncoder(json.JSONEncoder):
|
||||||
|
"""Custom JSON encoder that handles BaseModel objects."""
|
||||||
|
|
||||||
|
def default(self, obj: Any) -> Any:
|
||||||
|
"""Override default method to handle BaseModel objects."""
|
||||||
|
if isinstance(obj, BaseModel):
|
||||||
|
return obj.model_dump(by_alias=True, exclude_unset=True)
|
||||||
|
return super().default(obj)
|
||||||
|
|
||||||
|
|
||||||
|
def json_dumps(obj: Any, **kwargs) -> str:
|
||||||
|
"""
|
||||||
|
JSON dumps with support for BaseModel objects.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
obj: Object to serialize
|
||||||
|
**kwargs: Additional arguments to pass to json.dumps()
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
JSON string representation
|
||||||
|
"""
|
||||||
|
return json.dumps(obj, cls=ZAIJSONEncoder, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
def json_loads(s: str, **kwargs) -> Any:
|
||||||
|
"""
|
||||||
|
JSON loads with consistent interface.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
s: JSON string to deserialize
|
||||||
|
**kwargs: Additional arguments to pass to json.loads()
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Deserialized object
|
||||||
|
"""
|
||||||
|
return json.loads(s, **kwargs)
|
||||||
|
|
@ -94,6 +94,26 @@ class Choice(BaseModel):
|
||||||
finish_reason: Optional[str] = None
|
finish_reason: Optional[str] = None
|
||||||
index: int
|
index: int
|
||||||
|
|
||||||
|
class PromptTokensDetails(BaseModel):
|
||||||
|
"""
|
||||||
|
Detailed breakdown of token usage for the input prompt
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
cached_tokens: Number of tokens reused from cache
|
||||||
|
"""
|
||||||
|
|
||||||
|
cached_tokens: int
|
||||||
|
|
||||||
|
|
||||||
|
class CompletionTokensDetails(BaseModel):
|
||||||
|
"""
|
||||||
|
Detailed breakdown of token usage for the model completion
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
reasoning_tokens: Number of tokens used for reasoning steps
|
||||||
|
"""
|
||||||
|
|
||||||
|
reasoning_tokens: int
|
||||||
|
|
||||||
class CompletionUsage(BaseModel):
|
class CompletionUsage(BaseModel):
|
||||||
"""
|
"""
|
||||||
|
|
@ -106,7 +126,9 @@ class CompletionUsage(BaseModel):
|
||||||
"""
|
"""
|
||||||
|
|
||||||
prompt_tokens: int
|
prompt_tokens: int
|
||||||
|
prompt_tokens_details: Optional[PromptTokensDetails] = None
|
||||||
completion_tokens: int
|
completion_tokens: int
|
||||||
|
completion_tokens_details: Optional[CompletionTokensDetails] = None
|
||||||
total_tokens: int
|
total_tokens: int
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
BIN
tests/integration_tests/asr.wav
Normal file
BIN
tests/integration_tests/asr.wav
Normal file
Binary file not shown.
Binary file not shown.
|
|
@ -30,9 +30,9 @@ def test_audio_speech(logging_conf):
|
||||||
def test_audio_customization(logging_conf):
|
def test_audio_customization(logging_conf):
|
||||||
logging.config.dictConfig(logging_conf)
|
logging.config.dictConfig(logging_conf)
|
||||||
client = ZaiClient() # Fill in your own API Key
|
client = ZaiClient() # Fill in your own API Key
|
||||||
with open(Path(__file__).parent / 'asr1.wav', 'rb') as file:
|
with open(Path(__file__).parent / 'asr.wav', 'rb') as file:
|
||||||
try:
|
try:
|
||||||
speech_file_path = Path(__file__).parent / 'asr1.wav'
|
speech_file_path = Path(__file__).parent / 'asr.wav'
|
||||||
response = client.audio.customization(
|
response = client.audio.customization(
|
||||||
model='cogtts',
|
model='cogtts',
|
||||||
input='Hello, welcome to Z.ai Open Platform',
|
input='Hello, welcome to Z.ai Open Platform',
|
||||||
|
|
|
||||||
|
|
@ -1,267 +0,0 @@
|
||||||
"""
|
|
||||||
Integration tests for client cleanup functionality
|
|
||||||
"""
|
|
||||||
|
|
||||||
import gc
|
|
||||||
import os
|
|
||||||
import sys
|
|
||||||
import time
|
|
||||||
import unittest
|
|
||||||
from unittest.mock import patch
|
|
||||||
|
|
||||||
# Add src to path for testing
|
|
||||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '../../src'))
|
|
||||||
|
|
||||||
from zai import ZaiClient, ZhipuAiClient
|
|
||||||
|
|
||||||
|
|
||||||
class TestClientCleanupIntegration(unittest.TestCase):
|
|
||||||
"""Integration tests for client cleanup functionality"""
|
|
||||||
|
|
||||||
def setUp(self):
|
|
||||||
"""Set up test environment"""
|
|
||||||
# Set a dummy API key for testing
|
|
||||||
os.environ['ZAI_API_KEY'] = 'test-api-key'
|
|
||||||
|
|
||||||
def tearDown(self):
|
|
||||||
"""Clean up test environment"""
|
|
||||||
if 'ZAI_API_KEY' in os.environ:
|
|
||||||
del os.environ['ZAI_API_KEY']
|
|
||||||
|
|
||||||
def test_context_manager_with_api_calls(self):
|
|
||||||
"""Test context manager with actual API calls (mocked)"""
|
|
||||||
with patch('httpx.Client') as mock_client_class:
|
|
||||||
# Mock the httpx client
|
|
||||||
mock_client = mock_client_class.return_value
|
|
||||||
mock_client.is_closed = False
|
|
||||||
mock_client.close = unittest.mock.Mock()
|
|
||||||
|
|
||||||
with ZaiClient() as client:
|
|
||||||
# Simulate API call
|
|
||||||
self.assertIsNotNone(client)
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
|
|
||||||
# Simulate some API operations
|
|
||||||
# (In real usage, these would be actual API calls)
|
|
||||||
pass
|
|
||||||
|
|
||||||
# Client should be closed after context manager exits
|
|
||||||
# After our fix, the client should still be open because we check is_closed
|
|
||||||
# and don't close if it's already closed
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
mock_client.close.assert_called_once()
|
|
||||||
|
|
||||||
def test_explicit_close_with_api_calls(self):
|
|
||||||
"""Test explicit close with actual API calls (mocked)"""
|
|
||||||
with patch('httpx.Client') as mock_client_class:
|
|
||||||
# Mock the httpx client
|
|
||||||
mock_client = mock_client_class.return_value
|
|
||||||
mock_client.is_closed = False
|
|
||||||
mock_client.close = unittest.mock.Mock()
|
|
||||||
|
|
||||||
client = ZaiClient()
|
|
||||||
try:
|
|
||||||
# Simulate API call
|
|
||||||
self.assertIsNotNone(client)
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
|
|
||||||
# Simulate some API operations
|
|
||||||
# (In real usage, these would be actual API calls)
|
|
||||||
pass
|
|
||||||
finally:
|
|
||||||
client.close()
|
|
||||||
|
|
||||||
# Client should be closed
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
mock_client.close.assert_called_once()
|
|
||||||
|
|
||||||
def test_multiple_clients_in_sequence(self):
|
|
||||||
"""Test multiple clients created and closed in sequence"""
|
|
||||||
with patch('httpx.Client') as mock_client_class:
|
|
||||||
# Mock the httpx client
|
|
||||||
mock_client = mock_client_class.return_value
|
|
||||||
mock_client.is_closed = False
|
|
||||||
mock_client.close = unittest.mock.Mock()
|
|
||||||
|
|
||||||
# Create and close multiple clients
|
|
||||||
for i in range(3):
|
|
||||||
with ZaiClient() as client:
|
|
||||||
self.assertIsNotNone(client)
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
|
|
||||||
# Client should be closed after context manager exits
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
|
|
||||||
# Verify that close was called for each client
|
|
||||||
# Our fix allows multiple close calls, so the count might be higher
|
|
||||||
self.assertGreaterEqual(mock_client.close.call_count, 3)
|
|
||||||
|
|
||||||
def test_client_reuse_after_close(self):
|
|
||||||
"""Test reusing a client after it has been closed"""
|
|
||||||
with patch('httpx.Client') as mock_client_class:
|
|
||||||
# Mock the httpx client
|
|
||||||
mock_client = mock_client_class.return_value
|
|
||||||
mock_client.is_closed = False
|
|
||||||
mock_client.close = unittest.mock.Mock()
|
|
||||||
|
|
||||||
client = ZaiClient()
|
|
||||||
client.close()
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
|
|
||||||
# Try to close again - should not raise exception
|
|
||||||
client.close()
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
|
|
||||||
def test_zhipu_client_cleanup_integration(self):
|
|
||||||
"""Test ZhipuAiClient cleanup integration"""
|
|
||||||
with patch('httpx.Client') as mock_client_class:
|
|
||||||
# Mock the httpx client
|
|
||||||
mock_client = mock_client_class.return_value
|
|
||||||
mock_client.is_closed = False
|
|
||||||
mock_client.close = unittest.mock.Mock()
|
|
||||||
|
|
||||||
with ZhipuAiClient() as client:
|
|
||||||
self.assertIsNotNone(client)
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
|
|
||||||
# Simulate API call
|
|
||||||
pass
|
|
||||||
|
|
||||||
# Client should be closed after context manager exits
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
mock_client.close.assert_called_once()
|
|
||||||
|
|
||||||
def test_client_with_custom_http_client_integration(self):
|
|
||||||
"""Test client with custom HTTP client integration"""
|
|
||||||
import httpx
|
|
||||||
|
|
||||||
# Create a real httpx client for testing
|
|
||||||
custom_client = httpx.Client()
|
|
||||||
|
|
||||||
with ZaiClient(http_client=custom_client) as client:
|
|
||||||
self.assertIsNotNone(client)
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
|
|
||||||
# Simulate API call
|
|
||||||
pass
|
|
||||||
|
|
||||||
# Client should be closed after context manager exits
|
|
||||||
# The client is actually closed in this case
|
|
||||||
self.assertTrue(client.is_closed())
|
|
||||||
|
|
||||||
# Custom client should be closed by our client
|
|
||||||
self.assertTrue(custom_client.is_closed)
|
|
||||||
|
|
||||||
# No need to clean up custom client as it's already closed
|
|
||||||
|
|
||||||
def test_client_cleanup_with_exception_handling(self):
|
|
||||||
"""Test client cleanup when exceptions occur during API calls"""
|
|
||||||
with patch('httpx.Client') as mock_client_class:
|
|
||||||
# Mock the httpx client
|
|
||||||
mock_client = mock_client_class.return_value
|
|
||||||
mock_client.is_closed = False
|
|
||||||
mock_client.close = unittest.mock.Mock()
|
|
||||||
|
|
||||||
try:
|
|
||||||
with ZaiClient() as client:
|
|
||||||
self.assertIsNotNone(client)
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
|
|
||||||
# Simulate an exception during API call
|
|
||||||
raise Exception("Simulated API error")
|
|
||||||
except Exception:
|
|
||||||
# Exception should be caught and client should still be closed
|
|
||||||
pass
|
|
||||||
|
|
||||||
# Client should be closed even after exception
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
mock_client.close.assert_called_once()
|
|
||||||
|
|
||||||
def test_client_cleanup_with_timeout(self):
|
|
||||||
"""Test client cleanup with timeout scenarios"""
|
|
||||||
with patch('httpx.Client') as mock_client_class:
|
|
||||||
# Mock the httpx client
|
|
||||||
mock_client = mock_client_class.return_value
|
|
||||||
mock_client.is_closed = False
|
|
||||||
mock_client.close = unittest.mock.Mock()
|
|
||||||
|
|
||||||
# Create client with custom timeout
|
|
||||||
import httpx
|
|
||||||
custom_timeout = httpx.Timeout(timeout=30.0, connect=5.0)
|
|
||||||
|
|
||||||
with ZaiClient(timeout=custom_timeout) as client:
|
|
||||||
self.assertIsNotNone(client)
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
|
|
||||||
# Simulate API call with timeout
|
|
||||||
pass
|
|
||||||
|
|
||||||
# Client should be closed after context manager exits
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
mock_client.close.assert_called_once()
|
|
||||||
|
|
||||||
def test_client_cleanup_with_retries(self):
|
|
||||||
"""Test client cleanup with retry scenarios"""
|
|
||||||
with patch('httpx.Client') as mock_client_class:
|
|
||||||
# Mock the httpx client
|
|
||||||
mock_client = mock_client_class.return_value
|
|
||||||
mock_client.is_closed = False
|
|
||||||
mock_client.close = unittest.mock.Mock()
|
|
||||||
|
|
||||||
# Create client with custom retry settings
|
|
||||||
with ZaiClient(max_retries=5) as client:
|
|
||||||
self.assertIsNotNone(client)
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
|
|
||||||
# Simulate API call with retries
|
|
||||||
pass
|
|
||||||
|
|
||||||
# Client should be closed after context manager exits
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
mock_client.close.assert_called_once()
|
|
||||||
|
|
||||||
def test_client_cleanup_with_custom_headers(self):
|
|
||||||
"""Test client cleanup with custom headers"""
|
|
||||||
with patch('httpx.Client') as mock_client_class:
|
|
||||||
# Mock the httpx client
|
|
||||||
mock_client = mock_client_class.return_value
|
|
||||||
mock_client.is_closed = False
|
|
||||||
mock_client.close = unittest.mock.Mock()
|
|
||||||
|
|
||||||
# Create client with custom headers
|
|
||||||
custom_headers = {'X-Custom-Header': 'test-value'}
|
|
||||||
|
|
||||||
with ZaiClient(custom_headers=custom_headers) as client:
|
|
||||||
self.assertIsNotNone(client)
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
|
|
||||||
# Simulate API call
|
|
||||||
pass
|
|
||||||
|
|
||||||
# Client should be closed after context manager exits
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
mock_client.close.assert_called_once()
|
|
||||||
|
|
||||||
def test_client_cleanup_with_source_channel(self):
|
|
||||||
"""Test client cleanup with source channel"""
|
|
||||||
with patch('httpx.Client') as mock_client_class:
|
|
||||||
# Mock the httpx client
|
|
||||||
mock_client = mock_client_class.return_value
|
|
||||||
mock_client.is_closed = False
|
|
||||||
mock_client.close = unittest.mock.Mock()
|
|
||||||
|
|
||||||
# Create client with source channel
|
|
||||||
with ZaiClient(source_channel='test-channel') as client:
|
|
||||||
self.assertIsNotNone(client)
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
|
|
||||||
# Simulate API call
|
|
||||||
pass
|
|
||||||
|
|
||||||
# Client should be closed after context manager exits
|
|
||||||
self.assertFalse(client.is_closed())
|
|
||||||
mock_client.close.assert_called_once()
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
|
||||||
unittest.main()
|
|
||||||
|
|
@ -13,7 +13,6 @@ from zai import ZaiClient
|
||||||
@pytest.fixture(scope='class')
|
@pytest.fixture(scope='class')
|
||||||
def test_server():
|
def test_server():
|
||||||
class SharedData:
|
class SharedData:
|
||||||
client = ZaiClient()
|
|
||||||
file_id1 = None
|
file_id1 = None
|
||||||
file_id2 = None
|
file_id2 = None
|
||||||
|
|
||||||
|
|
@ -26,7 +25,10 @@ class TestZaiClientFileServer:
|
||||||
|
|
||||||
def test_files(self, test_server, test_file_path):
|
def test_files(self, test_server, test_file_path):
|
||||||
try:
|
try:
|
||||||
result = test_server.client.files.create(
|
import os
|
||||||
|
print('ZAI_API_KEY in test_files:', os.environ.get('ZAI_API_KEY'))
|
||||||
|
client = ZaiClient()
|
||||||
|
result = client.files.create(
|
||||||
file=open(os.path.join(test_file_path, 'demo.jsonl'), 'rb'),
|
file=open(os.path.join(test_file_path, 'demo.jsonl'), 'rb'),
|
||||||
purpose='fine-tune',
|
purpose='fine-tune',
|
||||||
)
|
)
|
||||||
|
|
@ -42,7 +44,10 @@ class TestZaiClientFileServer:
|
||||||
|
|
||||||
def test_files_validation(self, test_server, test_file_path):
|
def test_files_validation(self, test_server, test_file_path):
|
||||||
try:
|
try:
|
||||||
result = test_server.client.files.create(
|
import os
|
||||||
|
print('ZAI_API_KEY in test_files_validation:', os.environ.get('ZAI_API_KEY'))
|
||||||
|
client = ZaiClient()
|
||||||
|
result = client.files.create(
|
||||||
file=open(os.path.join(test_file_path, 'demo.jsonl'), 'rb'),
|
file=open(os.path.join(test_file_path, 'demo.jsonl'), 'rb'),
|
||||||
purpose='fine-tune',
|
purpose='fine-tune',
|
||||||
)
|
)
|
||||||
|
|
@ -59,7 +64,10 @@ class TestZaiClientFileServer:
|
||||||
|
|
||||||
def test_files_list(self, test_server):
|
def test_files_list(self, test_server):
|
||||||
try:
|
try:
|
||||||
list = test_server.client.files.list()
|
import os
|
||||||
|
print('ZAI_API_KEY in test_files_list:', os.environ.get('ZAI_API_KEY'))
|
||||||
|
client = ZaiClient()
|
||||||
|
list = client.files.list()
|
||||||
print(list)
|
print(list)
|
||||||
|
|
||||||
except zai.core._errors.APIRequestFailedError as err:
|
except zai.core._errors.APIRequestFailedError as err:
|
||||||
|
|
@ -71,13 +79,16 @@ class TestZaiClientFileServer:
|
||||||
|
|
||||||
def test_delete_files(self, test_server):
|
def test_delete_files(self, test_server):
|
||||||
try:
|
try:
|
||||||
|
import os
|
||||||
|
print('ZAI_API_KEY in test_delete_files:', os.environ.get('ZAI_API_KEY'))
|
||||||
|
client = ZaiClient()
|
||||||
# Only delete files if they were successfully created
|
# Only delete files if they were successfully created
|
||||||
if test_server.file_id1:
|
if test_server.file_id1:
|
||||||
delete1 = test_server.client.files.delete(file_id=test_server.file_id1)
|
delete1 = client.files.delete(file_id=test_server.file_id1)
|
||||||
print(delete1)
|
print(delete1)
|
||||||
|
|
||||||
if test_server.file_id2:
|
if test_server.file_id2:
|
||||||
delete2 = test_server.client.files.delete(file_id=test_server.file_id2)
|
delete2 = client.files.delete(file_id=test_server.file_id2)
|
||||||
print(delete2)
|
print(delete2)
|
||||||
|
|
||||||
except zai.core._errors.APIRequestFailedError as err:
|
except zai.core._errors.APIRequestFailedError as err:
|
||||||
|
|
|
||||||
23
tests/unit_tests/test_client.py
Normal file
23
tests/unit_tests/test_client.py
Normal file
|
|
@ -0,0 +1,23 @@
|
||||||
|
def test_completions_image_url_data_image(monkeypatch):
|
||||||
|
from zai import ZaiClient
|
||||||
|
client = ZaiClient()
|
||||||
|
# 构造 messages,包含 image_url 且 url 以 data:image/ 开头
|
||||||
|
base64_str = 'abc123=='
|
||||||
|
messages = [
|
||||||
|
{
|
||||||
|
'role': 'user',
|
||||||
|
'content': [
|
||||||
|
{'type': 'text', 'text': 'What is in this image?'},
|
||||||
|
{'type': 'image_url', 'image_url': {'url': f'data:image/png;base64,{base64_str}'}},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
]
|
||||||
|
# mock client.post 方法只返回 body 以便断言
|
||||||
|
def fake_post(*args, **kwargs):
|
||||||
|
return kwargs.get('body', {})
|
||||||
|
monkeypatch.setattr(client, 'post', fake_post)
|
||||||
|
result = client.chat.completions.create(model='glm-4', messages=messages)
|
||||||
|
# 验证 image_url 的 url 字段已去除前缀
|
||||||
|
content = result['messages'][0]['content']
|
||||||
|
image_url = [x for x in content if x.get('type') == 'image_url'][0]
|
||||||
|
assert image_url['image_url']['url'] == base64_str
|
||||||
166
tests/unit_tests/test_json_serialization.py
Normal file
166
tests/unit_tests/test_json_serialization.py
Normal file
|
|
@ -0,0 +1,166 @@
|
||||||
|
"""Test JSON serialization for BaseModel objects."""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import pytest
|
||||||
|
from typing import Optional, List
|
||||||
|
|
||||||
|
from zai.core._base_models import BaseModel
|
||||||
|
from zai.core._json_encoder import ZAIJSONEncoder, json_dumps
|
||||||
|
from zai.core._http_client import HttpClient
|
||||||
|
|
||||||
|
|
||||||
|
class ModelForTest(BaseModel):
|
||||||
|
"""Test model for JSON serialization."""
|
||||||
|
name: str
|
||||||
|
age: Optional[int] = None
|
||||||
|
tags: Optional[List[str]] = None
|
||||||
|
|
||||||
|
|
||||||
|
class MessageForTest(BaseModel):
|
||||||
|
"""Test message model similar to CompletionMessage."""
|
||||||
|
content: Optional[str] = None
|
||||||
|
role: str
|
||||||
|
tool_calls: Optional[List[dict]] = None
|
||||||
|
|
||||||
|
|
||||||
|
def test_base_model_dict_conversion():
|
||||||
|
"""Test that BaseModel can be converted to dict."""
|
||||||
|
model = ModelForTest(name="test", age=25, tags=["tag1", "tag2"])
|
||||||
|
|
||||||
|
# Test model_dump() method
|
||||||
|
model_dict = model.model_dump()
|
||||||
|
assert model_dict == {"name": "test", "age": 25, "tags": ["tag1", "tag2"]}
|
||||||
|
|
||||||
|
# Test keys/values/items methods
|
||||||
|
assert list(model.keys()) == ["name", "age", "tags"]
|
||||||
|
assert list(model.values()) == ["test", 25, ["tag1", "tag2"]]
|
||||||
|
assert list(model.items()) == [("name", "test"), ("age", 25), ("tags", ["tag1", "tag2"])]
|
||||||
|
|
||||||
|
|
||||||
|
def test_zai_json_encoder():
|
||||||
|
"""Test custom JSON encoder for BaseModel objects."""
|
||||||
|
model = ModelForTest(name="test", age=25)
|
||||||
|
|
||||||
|
# Test with custom encoder
|
||||||
|
json_str = json.dumps(model, cls=ZAIJSONEncoder)
|
||||||
|
result = json.loads(json_str)
|
||||||
|
assert result == {"name": "test", "age": 25}
|
||||||
|
|
||||||
|
# Test with nested BaseModel
|
||||||
|
message = MessageForTest(content="hello", role="user")
|
||||||
|
nested_data = {"message": message, "other": "data"}
|
||||||
|
|
||||||
|
json_str = json.dumps(nested_data, cls=ZAIJSONEncoder)
|
||||||
|
result = json.loads(json_str)
|
||||||
|
assert result == {
|
||||||
|
"message": {"content": "hello", "role": "user"},
|
||||||
|
"other": "data"
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_json_dumps_utility():
|
||||||
|
"""Test the json_dumps utility function."""
|
||||||
|
model = ModelForTest(name="test", age=25)
|
||||||
|
|
||||||
|
json_str = json_dumps(model)
|
||||||
|
result = json.loads(json_str)
|
||||||
|
assert result == {"name": "test", "age": 25}
|
||||||
|
|
||||||
|
|
||||||
|
def test_http_client_json_preparation():
|
||||||
|
"""Test that HttpClient can prepare JSON data with BaseModel objects."""
|
||||||
|
# We'll create a minimal HttpClient to test the _prepare_json_data method
|
||||||
|
import httpx
|
||||||
|
from zai.core._http_client import HttpClient
|
||||||
|
|
||||||
|
# Create a minimal client for testing
|
||||||
|
client = HttpClient(
|
||||||
|
base_url="https://test.com",
|
||||||
|
version="v1",
|
||||||
|
_strict_response_validation=True,
|
||||||
|
timeout=10.0
|
||||||
|
)
|
||||||
|
|
||||||
|
# Test with BaseModel
|
||||||
|
model = ModelForTest(name="test", age=25)
|
||||||
|
result = client._prepare_json_data(model)
|
||||||
|
assert result == {"name": "test", "age": 25}
|
||||||
|
|
||||||
|
# Test with list containing BaseModel
|
||||||
|
models = [ModelForTest(name="test1", age=25), ModelForTest(name="test2", age=30)]
|
||||||
|
result = client._prepare_json_data(models)
|
||||||
|
expected = [
|
||||||
|
{"name": "test1", "age": 25},
|
||||||
|
{"name": "test2", "age": 30}
|
||||||
|
]
|
||||||
|
assert result == expected
|
||||||
|
|
||||||
|
# Test with dict containing BaseModel
|
||||||
|
data = {
|
||||||
|
"message": MessageForTest(content="hello", role="user"),
|
||||||
|
"count": 1
|
||||||
|
}
|
||||||
|
result = client._prepare_json_data(data)
|
||||||
|
expected = {
|
||||||
|
"message": {"content": "hello", "role": "user"},
|
||||||
|
"count": 1
|
||||||
|
}
|
||||||
|
assert result == expected
|
||||||
|
|
||||||
|
# Test with None
|
||||||
|
assert client._prepare_json_data(None) is None
|
||||||
|
|
||||||
|
# Test with regular data
|
||||||
|
regular_data = {"key": "value", "number": 42}
|
||||||
|
assert client._prepare_json_data(regular_data) == regular_data
|
||||||
|
|
||||||
|
|
||||||
|
def test_completion_message_like_serialization():
|
||||||
|
"""Test serialization for CompletionMessage-like objects."""
|
||||||
|
# Simulate a CompletionMessage-like object
|
||||||
|
message = MessageForTest(
|
||||||
|
content="Hello, how can I help you?",
|
||||||
|
role="assistant",
|
||||||
|
tool_calls=[{"id": "call_123", "function": {"name": "test_func"}}]
|
||||||
|
)
|
||||||
|
|
||||||
|
# Test that it can be serialized
|
||||||
|
json_str = json_dumps(message)
|
||||||
|
result = json.loads(json_str)
|
||||||
|
|
||||||
|
expected = {
|
||||||
|
"content": "Hello, how can I help you?",
|
||||||
|
"role": "assistant",
|
||||||
|
"tool_calls": [{"id": "call_123", "function": {"name": "test_func"}}]
|
||||||
|
}
|
||||||
|
assert result == expected
|
||||||
|
|
||||||
|
# Test in a messages list (similar to the original error case)
|
||||||
|
messages = [
|
||||||
|
{"role": "user", "content": "Help me"},
|
||||||
|
message # This is a BaseModel object
|
||||||
|
]
|
||||||
|
|
||||||
|
# This should work now with our custom encoder
|
||||||
|
json_str = json_dumps(messages)
|
||||||
|
result = json.loads(json_str)
|
||||||
|
|
||||||
|
expected = [
|
||||||
|
{"role": "user", "content": "Help me"},
|
||||||
|
{
|
||||||
|
"content": "Hello, how can I help you?",
|
||||||
|
"role": "assistant",
|
||||||
|
"tool_calls": [{"id": "call_123", "function": {"name": "test_func"}}]
|
||||||
|
}
|
||||||
|
]
|
||||||
|
assert result == expected
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
# Run basic tests
|
||||||
|
test_base_model_dict_conversion()
|
||||||
|
test_zai_json_encoder()
|
||||||
|
test_json_dumps_utility()
|
||||||
|
test_http_client_json_preparation()
|
||||||
|
test_completion_message_like_serialization()
|
||||||
|
print("All tests passed!")
|
||||||
Loading…
Reference in a new issue