fix: video example, remove finetuning\knowledge\document module (#17)

Co-authored-by: zhengweijun <weijun.zheng@aminer.cn>
This commit is contained in:
wellenzheng 2025-07-26 11:21:48 +08:00 committed by GitHub
parent d344be2ca2
commit 508eaf4f4c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
44 changed files with 173 additions and 1644 deletions

2
.env
View file

@ -1,2 +1,2 @@
PYTHON_KEYRING_BACKEND=keyring.backends.null.Keyring PYTHON_KEYRING_BACKEND=keyring.backends.null.Keyring
ZAI_API_KEY={your apikey} ZAI_API_KEY={your apikey}

View file

@ -180,11 +180,12 @@ def ofZhipu():
if __name__ == '__main__': if __name__ == '__main__':
# completion() # completion()
# completion_with_websearch() # completion_with_stream()
# multi_modal_chat() # completion_with_websearch()
# role_play() multi_modal_chat()
# assistant_conversation() # role_play()
# video_generation() # assistant_conversation()
ofZai() # video_generation()
ofZhipu() # ofZai()
# ofZhipu()

View file

@ -0,0 +1,120 @@
from zai import ZhipuAiClient
import json
client = ZhipuAiClient()
def get_flight_number(date: str, departure: str, destination: str):
flight_number = {
"Beijing": {
"Shanghai": "1234",
"Guangzhou": "5678",
},
"Shanghai": {
"Beijing": "4321",
"Guangzhou": "8765",
}
}
return {"flight_number": flight_number[departure][destination]}
def get_ticket_price(date: str, flight_number: str):
return {"ticket_price": "1000"}
def parse_function_call(model_response, messages):
# Handle function call results. According to the model's returned parameters, call the corresponding function.
# After getting the function result, construct a tool message and call the model again, passing the function result as input.
# The model will return the function result to the user in natural language.
if model_response.choices[0].message.tool_calls:
tool_call = model_response.choices[0].message.tool_calls[0]
args = tool_call.function.arguments
function_result = {}
if tool_call.function.name == "get_flight_number":
function_result = get_flight_number(**json.loads(args))
if tool_call.function.name == "get_ticket_price":
function_result = get_ticket_price(**json.loads(args))
messages.append({
"role": "tool",
"content": f"{json.dumps(function_result)}",
"tool_call_id": tool_call.id
})
response = client.chat.completions.create(
model="glm-4", # Specify the model name to use
messages=messages,
tools=tools,
)
print(response.choices[0].message)
messages.append(response.choices[0].message.model_dump())
messages = []
tools = [
{
"type": "function",
"function": {
"name": "get_flight_number",
"description": "Query the flight number for a given date, departure, and destination",
"parameters": {
"type": "object",
"properties": {
"departure": {
"description": "Departure city",
"type": "string"
},
"destination": {
"description": "Destination city",
"type": "string"
},
"date": {
"description": "Date",
"type": "string",
}
},
"required": ["departure", "destination", "date"]
},
}
},
{
"type": "function",
"function": {
"name": "get_ticket_price",
"description": "Query the ticket price for a specific flight on a specific date",
"parameters": {
"type": "object",
"properties": {
"flight_number": {
"description": "Flight number",
"type": "string"
},
"date": {
"description": "Date",
"type": "string",
}
},
"required": ["flight_number", "date"]
},
}
},
]
# Clear conversation
messages = []
messages.append({"role": "system", "content": "Do not assume or guess the values of function parameters. If the user's description is unclear, ask the user to provide the necessary information."})
messages.append({"role": "user", "content": "Help me check the flights from Beijing to Guangzhou on January 23."})
response = client.chat.completions.create(
model="glm-4", # Specify the model name to use
messages=messages,
tools=tools,
)
print(response.choices[0].message)
messages.append(response.choices[0].message.model_dump())
parse_function_call(response, messages)
messages.append({"role": "user", "content": "What is the price of flight 8321?"})
response = client.chat.completions.create(
model="glm-4", # Specify the model name to use
messages=messages,
tools=tools,
)
print(response.choices[0].message)
messages.append(response.choices[0].message.model_dump())
parse_function_call(response, messages)

View file

@ -29,7 +29,7 @@ def stream_web_search_example():
stream=True stream=True
) )
for chunk in response: for chunk in response:
print(chunk.choices[0].delta) print(chunk.choices[0].delta.content, end="", flush=True)
def sync_example(): def sync_example():
print("=== GLM-4 Synchronous Example ===") print("=== GLM-4 Synchronous Example ===")

View file

@ -1,6 +1,6 @@
[tool.poetry] [tool.poetry]
name = "zai-sdk" name = "zai-sdk"
version = "0.0.1b3" version = "0.0.1b4"
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

@ -16,9 +16,7 @@ if TYPE_CHECKING:
from zai.api_resource.chat import Chat from zai.api_resource.chat import Chat
from zai.api_resource.embeddings import Embeddings from zai.api_resource.embeddings import Embeddings
from zai.api_resource.files import Files from zai.api_resource.files import Files
from zai.api_resource.fine_tuning import FineTuning
from zai.api_resource.images import Images from zai.api_resource.images import Images
from zai.api_resource.knowledge import Knowledge
from zai.api_resource.moderations import Moderations from zai.api_resource.moderations import Moderations
from zai.api_resource.tools import Tools from zai.api_resource.tools import Tools
from zai.api_resource.videos import Videos from zai.api_resource.videos import Videos
@ -132,24 +130,12 @@ class BaseClient(HttpClient):
return Embeddings(self) return Embeddings(self)
@cached_property
def fine_tuning(self) -> FineTuning:
from zai.api_resource.fine_tuning import FineTuning
return FineTuning(self)
@cached_property @cached_property
def batches(self) -> Batches: def batches(self) -> Batches:
from zai.api_resource.batch import Batches from zai.api_resource.batch import Batches
return Batches(self) return Batches(self)
@cached_property
def knowledge(self) -> Knowledge:
from zai.api_resource.knowledge import Knowledge
return Knowledge(self)
@cached_property @cached_property
def tools(self) -> Tools: def tools(self) -> Tools:
from zai.api_resource.tools import Tools from zai.api_resource.tools import Tools

View file

@ -1,2 +1,2 @@
__title__ = 'Z.ai' __title__ = 'Z.ai'
__version__ = '0.0.1b2' __version__ = '0.0.1b4'

View file

@ -11,9 +11,7 @@ from .chat import (
) )
from .embeddings import Embeddings from .embeddings import Embeddings
from .files import Files, FilesWithRawResponse from .files import Files, FilesWithRawResponse
from .fine_tuning import FineTuning
from .images import Images from .images import Images
from .knowledge import Knowledge
from .moderations import Moderations from .moderations import Moderations
from .tools import Tools from .tools import Tools
from .videos import ( from .videos import (
@ -30,9 +28,7 @@ __all__ = [
'Embeddings', 'Embeddings',
'Files', 'Files',
'FilesWithRawResponse', 'FilesWithRawResponse',
'FineTuning',
'Batches', 'Batches',
'Knowledge',
'Tools', 'Tools',
'Assistant', 'Assistant',
'Audio', 'Audio',

View file

@ -1,5 +0,0 @@
from .fine_tuning import FineTuning
from .jobs import Jobs
from .models import FineTunedModels
__all__ = ['Jobs', 'FineTunedModels', 'FineTuning']

View file

@ -1,19 +0,0 @@
from typing import TYPE_CHECKING
from zai.core import BaseAPI, cached_property
from .jobs import Jobs
from .models import FineTunedModels
if TYPE_CHECKING:
pass
class FineTuning(BaseAPI):
@cached_property
def jobs(self) -> Jobs:
return Jobs(self._client)
@cached_property
def models(self) -> FineTunedModels:
return FineTunedModels(self._client)

View file

@ -1,3 +0,0 @@
from .jobs import Jobs
__all__ = ['Jobs']

View file

@ -1,152 +0,0 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Optional
import httpx
from zai.core import (
NOT_GIVEN,
BaseAPI,
Body,
Headers,
NotGiven,
make_request_options,
)
from zai.types.fine_tuning import (
FineTuningJob,
FineTuningJobEvent,
ListOfFineTuningJob,
job_create_params,
)
if TYPE_CHECKING:
from zai._client import ZaiClient
class Jobs(BaseAPI):
def __init__(self, client: 'ZaiClient') -> None:
super().__init__(client)
def create(
self,
*,
model: str,
training_file: str,
hyperparameters: job_create_params.Hyperparameters | NotGiven = NOT_GIVEN,
suffix: Optional[str] | NotGiven = NOT_GIVEN,
request_id: Optional[str] | NotGiven = NOT_GIVEN,
validation_file: Optional[str] | NotGiven = NOT_GIVEN,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> FineTuningJob:
return self._post(
'/fine_tuning/jobs',
body={
'model': model,
'training_file': training_file,
'hyperparameters': hyperparameters,
'suffix': suffix,
'validation_file': validation_file,
'request_id': request_id,
},
options=make_request_options(extra_headers=extra_headers, extra_body=extra_body, timeout=timeout),
cast_type=FineTuningJob,
)
def retrieve(
self,
fine_tuning_job_id: str,
*,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> FineTuningJob:
return self._get(
f'/fine_tuning/jobs/{fine_tuning_job_id}',
options=make_request_options(extra_headers=extra_headers, extra_body=extra_body, timeout=timeout),
cast_type=FineTuningJob,
)
def list(
self,
*,
after: str | NotGiven = NOT_GIVEN,
limit: int | NotGiven = NOT_GIVEN,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> ListOfFineTuningJob:
return self._get(
'/fine_tuning/jobs',
cast_type=ListOfFineTuningJob,
options=make_request_options(
extra_headers=extra_headers,
extra_body=extra_body,
timeout=timeout,
query={
'after': after,
'limit': limit,
},
),
)
def cancel(
self,
fine_tuning_job_id: str,
*,
# Use the following arguments if you need to pass additional parameters to the API
# that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client
# or passed to this method.
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> FineTuningJob:
if not fine_tuning_job_id:
raise ValueError(f'Expected a non-empty value for `fine_tuning_job_id` but received {fine_tuning_job_id!r}')
return self._post(
f'/fine_tuning/jobs/{fine_tuning_job_id}/cancel',
options=make_request_options(extra_headers=extra_headers, extra_body=extra_body, timeout=timeout),
cast_type=FineTuningJob,
)
def list_events(
self,
fine_tuning_job_id: str,
*,
after: str | NotGiven = NOT_GIVEN,
limit: int | NotGiven = NOT_GIVEN,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> FineTuningJobEvent:
return self._get(
f'/fine_tuning/jobs/{fine_tuning_job_id}/events',
cast_type=FineTuningJobEvent,
options=make_request_options(
extra_headers=extra_headers,
extra_body=extra_body,
timeout=timeout,
query={
'after': after,
'limit': limit,
},
),
)
def delete(
self,
fine_tuning_job_id: str,
*,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> FineTuningJob:
if not fine_tuning_job_id:
raise ValueError(f'Expected a non-empty value for `fine_tuning_job_id` but received {fine_tuning_job_id!r}')
return self._delete(
f'/fine_tuning/jobs/{fine_tuning_job_id}',
options=make_request_options(extra_headers=extra_headers, extra_body=extra_body, timeout=timeout),
cast_type=FineTuningJob,
)

View file

@ -1,3 +0,0 @@
from .fine_tuned_models import FineTunedModels
__all__ = ['FineTunedModels']

View file

@ -1,39 +0,0 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import httpx
from zai.core import (
NOT_GIVEN,
BaseAPI,
Body,
Headers,
NotGiven,
make_request_options,
)
from zai.types.fine_tuning.models import FineTunedModelsStatus
if TYPE_CHECKING:
from zai._client import ZaiClient
class FineTunedModels(BaseAPI):
def __init__(self, client: 'ZaiClient') -> None:
super().__init__(client)
def delete(
self,
fine_tuned_model: str,
*,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> FineTunedModelsStatus:
if not fine_tuned_model:
raise ValueError(f'Expected a non-empty value for `fine_tuned_model` but received {fine_tuned_model!r}')
return self._delete(
f'fine_tuning/fine_tuned_models/{fine_tuned_model}',
options=make_request_options(extra_headers=extra_headers, extra_body=extra_body, timeout=timeout),
cast_type=FineTunedModelsStatus,
)

View file

@ -1,3 +0,0 @@
from .knowledge import Knowledge
__all__ = ['Knowledge']

View file

@ -1,3 +0,0 @@
from .document import Document
__all__ = ['Document']

View file

@ -1,254 +0,0 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Dict, List, Mapping, Optional, cast
import httpx
from typing_extensions import Literal
from zai.core import (
NOT_GIVEN,
BaseAPI,
Body,
FileTypes,
Headers,
NotGiven,
deepcopy_minimal,
extract_files,
make_request_options,
maybe_transform,
)
from zai.types.files import UploadDetail, file_create_params
from zai.types.knowledge.document import (
DocumentData,
DocumentObject,
document_edit_params,
document_list_params,
)
from zai.types.knowledge.document.document_list_resp import DocumentPage
if TYPE_CHECKING:
from zai._client import ZaiClient
class Document(BaseAPI):
"""
API resource for document operations in knowledge bases
"""
def __init__(self, client: 'ZaiClient') -> None:
super().__init__(client)
def create(
self,
*,
file: FileTypes = None,
custom_separator: Optional[List[str]] = None,
upload_detail: List[UploadDetail] = None,
purpose: Literal['retrieval'],
knowledge_id: str = None,
sentence_size: int = None,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> DocumentObject:
"""
Create a new document in a knowledge base
Arguments:
file (FileTypes): File to upload
custom_separator (Optional[List[str]]): Custom separators for text splitting
upload_detail (List[UploadDetail]): Upload details for the document
purpose (Literal['retrieval']): Purpose of the document
knowledge_id (str): ID of the knowledge base
sentence_size (int): Size of sentences for splitting
extra_headers (Headers): Additional headers to send
extra_body (Body): Additional body parameters
timeout (float | httpx.Timeout): Request timeout
"""
if not file and not upload_detail:
raise ValueError('At least one of `file` and `upload_detail` must be provided.')
body = deepcopy_minimal(
{
'file': file,
'upload_detail': upload_detail,
'purpose': purpose,
'custom_separator': custom_separator,
'knowledge_id': knowledge_id,
'sentence_size': sentence_size,
}
)
files = extract_files(cast(Mapping[str, object], body), paths=[['file']])
if files:
# It should be noted that the actual Content-Type header that will be
# sent to the server will contain a `boundary` parameter, e.g.
# multipart/form-data; boundary=---abc--
extra_headers = {
'Content-Type': 'multipart/form-data',
**(extra_headers or {}),
}
return self._post(
'/files',
body=maybe_transform(body, file_create_params.FileCreateParams),
files=files,
options=make_request_options(extra_headers=extra_headers, extra_body=extra_body, timeout=timeout),
cast_type=DocumentObject,
)
def edit(
self,
document_id: str,
knowledge_type: str,
*,
custom_separator: Optional[List[str]] = None,
sentence_size: Optional[int] = None,
callback_url: Optional[str] = None,
callback_header: Optional[Dict[str, str]] = None,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> httpx.Response:
"""
Args:
document_id: Knowledge ID
knowledge_type: Knowledge type:
1: Article knowledge: supports pdf, url, docx
2. Q&A knowledge-document: supports pdf, url, docx
3. Q&A knowledge-table: supports xlsx
4. Product library-table: supports xlsx
5. Custom: supports pdf, url, docx
extra_headers: Send extra headers
extra_body: Add additional JSON properties to the request
timeout: Override the client-level default timeout for this request, in seconds
:param knowledge_type:
:param document_id:
:param timeout:
:param extra_body:
:param callback_header:
:param sentence_size:
:param extra_headers:
:param callback_url:
:param custom_separator:
"""
if not document_id:
raise ValueError(f'Expected a non-empty value for `document_id` but received {document_id!r}')
body = deepcopy_minimal(
{
'id': document_id,
'knowledge_type': knowledge_type,
'custom_separator': custom_separator,
'sentence_size': sentence_size,
'callback_url': callback_url,
'callback_header': callback_header,
}
)
return self._put(
f'/document/{document_id}',
body=maybe_transform(body, document_edit_params.DocumentEditParams),
options=make_request_options(extra_headers=extra_headers, extra_body=extra_body, timeout=timeout),
cast_type=httpx.Response,
)
def list(
self,
knowledge_id: str,
*,
purpose: str | NotGiven = NOT_GIVEN,
page: str | NotGiven = NOT_GIVEN,
limit: str | NotGiven = NOT_GIVEN,
order: Literal['desc', 'asc'] | NotGiven = NOT_GIVEN,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> DocumentPage:
"""
List documents in a knowledge base
Arguments:
knowledge_id (str): ID of the knowledge base
purpose (str): Purpose filter for documents
page (str): Page number for pagination
limit (str): Number of documents per page
order (Literal['desc', 'asc']): Sort order for results
extra_headers (Headers): Additional headers to send
extra_body (Body): Additional body parameters
timeout (float | httpx.Timeout): Request timeout
"""
return self._get(
'/files',
options=make_request_options(
extra_headers=extra_headers,
extra_body=extra_body,
timeout=timeout,
query=maybe_transform(
{
'knowledge_id': knowledge_id,
'purpose': purpose,
'page': page,
'limit': limit,
'order': order,
},
document_list_params.DocumentListParams,
),
),
cast_type=DocumentPage,
)
def delete(
self,
document_id: str,
*,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> httpx.Response:
"""
Delete a file.
Args:
document_id: Knowledge ID
extra_headers: Send extra headers
extra_body: Add additional JSON properties to the request
timeout: Override the client-level default timeout for this request, in seconds
"""
if not document_id:
raise ValueError(f'Expected a non-empty value for `document_id` but received {document_id!r}')
return self._delete(
f'/document/{document_id}',
options=make_request_options(extra_headers=extra_headers, extra_body=extra_body, timeout=timeout),
cast_type=httpx.Response,
)
def retrieve(
self,
document_id: str,
*,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> DocumentData:
"""
Args:
extra_headers: Send extra headers
extra_body: Add additional JSON properties to the request
timeout: Override the client-level default timeout for this request, in seconds
"""
if not document_id:
raise ValueError(f'Expected a non-empty value for `document_id` but received {document_id!r}')
return self._get(
f'/document/{document_id}',
options=make_request_options(extra_headers=extra_headers, extra_body=extra_body, timeout=timeout),
cast_type=DocumentData,
)

View file

@ -1,178 +0,0 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Optional
import httpx
from typing_extensions import Literal
from zai.core import (
NOT_GIVEN,
BaseAPI,
Body,
Headers,
NotGiven,
cached_property,
deepcopy_minimal,
make_request_options,
maybe_transform,
)
from zai.types.knowledge import (
KnowledgeInfo,
KnowledgeUsed,
knowledge_create_params,
knowledge_list_params,
)
from zai.types.knowledge.knowledge_list_resp import KnowledgePage
from .document import Document
if TYPE_CHECKING:
from zai._client import ZaiClient
class Knowledge(BaseAPI):
def __init__(self, client: 'ZaiClient') -> None:
super().__init__(client)
@cached_property
def document(self) -> Document:
return Document(self._client)
def create(
self,
embedding_id: int,
name: str,
*,
customer_identifier: Optional[str] = None,
description: Optional[str] = None,
background: Optional[Literal['blue', 'red', 'orange', 'purple', 'sky']] = None,
icon: Optional[Literal['question', 'book', 'seal', 'wrench', 'tag', 'horn', 'house']] = None,
bucket_id: Optional[str] = None,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> KnowledgeInfo:
body = deepcopy_minimal(
{
'embedding_id': embedding_id,
'name': name,
'customer_identifier': customer_identifier,
'description': description,
'background': background,
'icon': icon,
'bucket_id': bucket_id,
}
)
return self._post(
'/knowledge',
body=maybe_transform(body, knowledge_create_params.KnowledgeBaseParams),
options=make_request_options(extra_headers=extra_headers, extra_body=extra_body, timeout=timeout),
cast_type=KnowledgeInfo,
)
def modify(
self,
knowledge_id: str,
embedding_id: int,
*,
name: str,
description: Optional[str] = None,
background: Optional[Literal['blue', 'red', 'orange', 'purple', 'sky']] = None,
icon: Optional[Literal['question', 'book', 'seal', 'wrench', 'tag', 'horn', 'house']] = None,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> httpx.Response:
body = deepcopy_minimal(
{
'id': knowledge_id,
'embedding_id': embedding_id,
'name': name,
'description': description,
'background': background,
'icon': icon,
}
)
return self._put(
f'/knowledge/{knowledge_id}',
body=maybe_transform(body, knowledge_create_params.KnowledgeBaseParams),
options=make_request_options(extra_headers=extra_headers, extra_body=extra_body, timeout=timeout),
cast_type=httpx.Response,
)
def query(
self,
*,
page: int | NotGiven = 1,
size: int | NotGiven = 10,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> KnowledgePage:
return self._get(
'/knowledge',
options=make_request_options(
extra_headers=extra_headers,
extra_body=extra_body,
timeout=timeout,
query=maybe_transform(
{
'page': page,
'size': size,
},
knowledge_list_params.KnowledgeListParams,
),
),
cast_type=KnowledgePage,
)
def delete(
self,
knowledge_id: str,
*,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> httpx.Response:
"""
Delete a knowledge.
Args:
knowledge_id: Knowledge ID
extra_headers: Send extra headers
extra_body: Add additional JSON properties to the request
timeout: Override the client-level default timeout for this request, in seconds
"""
if not knowledge_id:
raise ValueError('Expected a non-empty value for `knowledge_id`')
return self._delete(
f'/knowledge/{knowledge_id}',
options=make_request_options(extra_headers=extra_headers, extra_body=extra_body, timeout=timeout),
cast_type=httpx.Response,
)
def used(
self,
*,
extra_headers: Headers | None = None,
extra_body: Body | None = None,
timeout: float | httpx.Timeout | None | NotGiven = NOT_GIVEN,
) -> KnowledgeUsed:
"""
Returns the contents of the specified file.
Args:
extra_headers: Send extra headers
extra_body: Add additional JSON properties to the request
timeout: Override the client-level default timeout for this request, in seconds
"""
return self._get(
'/knowledge/capacity',
options=make_request_options(extra_headers=extra_headers, extra_body=extra_body, timeout=timeout),
cast_type=KnowledgeUsed,
)

View file

@ -46,7 +46,7 @@ class AgentsCompletionChoice(BaseModel):
message: AgentsCompletionMessage message: AgentsCompletionMessage
class AgentsError: class AgentsError(BaseModel):
""" """
Represents an error in agents completion. Represents an error in agents completion.

View file

@ -43,7 +43,7 @@ class AgentsCompletionUsage(BaseModel):
total_tokens: int total_tokens: int
class AgentsError: class AgentsError(BaseModel):
""" """
Represents an error in agents completion chunk. Represents an error in agents completion chunk.

View file

@ -1,16 +0,0 @@
from __future__ import annotations
from .fine_tuning_job import Error, FineTuningJob, Hyperparameters, ListOfFineTuningJob
from .fine_tuning_job_event import FineTuningJobEvent, JobEvent, Metric
from .job_create_params import Hyperparameters as JobHyperparameters
__all__ = [
'FineTuningJob',
'Error',
'Hyperparameters',
'ListOfFineTuningJob',
'FineTuningJobEvent',
'Metric',
'JobEvent',
'JobHyperparameters',
]

View file

@ -1,76 +0,0 @@
from typing import List, Optional, Union
from zai.core import BaseModel
class Error(BaseModel):
"""
Error information for fine-tuning job
Attributes:
code: Error code
message: Error message description
param: Optional parameter that caused the error
"""
code: str
message: str
param: Optional[str] = None
class Hyperparameters(BaseModel):
"""
Hyperparameters for fine-tuning job
Attributes:
n_epochs: Number of training epochs
"""
n_epochs: Union[str, int, None] = None
class FineTuningJob(BaseModel):
"""
Fine-tuning job information
Attributes:
id: Unique identifier for the fine-tuning job
request_id: Request identifier
created_at: Timestamp when the job was created
error: Error information if the job failed
fine_tuned_model: Name of the fine-tuned model
finished_at: Timestamp when the job finished
hyperparameters: Hyperparameters used for training
model: Base model used for fine-tuning
object: Object type identifier
result_files: List of result file paths
status: Current status of the fine-tuning job
trained_tokens: Number of tokens used for training
training_file: Path to the training file
validation_file: Path to the validation file
"""
id: Optional[str] = None
request_id: Optional[str] = None
created_at: Optional[int] = None
error: Optional[Error] = None
fine_tuned_model: Optional[str] = None
finished_at: Optional[int] = None
hyperparameters: Optional[Hyperparameters] = None
model: Optional[str] = None
object: Optional[str] = None
result_files: List[str]
status: str
trained_tokens: Optional[int] = None
training_file: str
validation_file: Optional[str] = None
class ListOfFineTuningJob(BaseModel):
"""
List of fine-tuning jobs response
Attributes:
object: Object type identifier
data: List of fine-tuning job objects
"""
object: Optional[str] = None
data: List[FineTuningJob]
has_more: Optional[bool] = None

View file

@ -1,33 +0,0 @@
from typing import List, Optional, Union
from zai.core import BaseModel
class Metric(BaseModel):
epoch: Optional[Union[str, int, float]] = None
current_steps: Optional[int] = None
total_steps: Optional[int] = None
elapsed_time: Optional[str] = None
remaining_time: Optional[str] = None
trained_tokens: Optional[int] = None
loss: Optional[Union[str, int, float]] = None
eval_loss: Optional[Union[str, int, float]] = None
acc: Optional[Union[str, int, float]] = None
eval_acc: Optional[Union[str, int, float]] = None
learning_rate: Optional[Union[str, int, float]] = None
class JobEvent(BaseModel):
object: Optional[str] = None
id: Optional[str] = None
type: Optional[str] = None
created_at: Optional[int] = None
level: Optional[str] = None
message: Optional[str] = None
data: Optional[Metric] = None
class FineTuningJobEvent(BaseModel):
object: Optional[str] = None
data: List[JobEvent]
has_more: Optional[bool] = None

View file

@ -1,19 +0,0 @@
from __future__ import annotations
from typing import Union
from typing_extensions import Literal, TypedDict
class Hyperparameters(TypedDict, total=False):
"""
Hyperparameters for fine-tuning job configuration
Attributes:
batch_size: The batch size to use for training (can be 'auto' or an integer)
learning_rate_multiplier: The learning rate multiplier for training (can be 'auto' or a float)
n_epochs: The number of epochs to train for (can be 'auto' or an integer)
"""
batch_size: Union[Literal['auto'], int]
learning_rate_multiplier: Union[Literal['auto'], float]
n_epochs: Union[Literal['auto'], int]

View file

@ -1,3 +0,0 @@
from .fine_tuned_models import FineTunedModelsStatus
__all__ = ['FineTunedModelsStatus']

View file

@ -1,20 +0,0 @@
from typing import ClassVar
from zai.core import PYDANTIC_V2, BaseModel, ConfigDict
class FineTunedModelsStatus(BaseModel):
"""
Fine-tuned model status
Attributes:
request_id (str): Request id
model_name (str): Model name
delete_status (str): Delete status: deleting (deleting), deleted (deleted)
"""
if PYDANTIC_V2:
model_config: ClassVar[ConfigDict] = ConfigDict(extra='allow', protected_namespaces=())
request_id: str
model_name: str
delete_status: str

View file

@ -1,14 +0,0 @@
from .knowledge import KnowledgeInfo
from .knowledge_create_params import KnowledgeBaseParams
from .knowledge_list_params import KnowledgeListParams
from .knowledge_list_resp import KnowledgePage
from .knowledge_used import KnowledgeStatistics, KnowledgeUsed
__all__ = [
'KnowledgeInfo',
'KnowledgeStatistics',
'KnowledgeUsed',
'KnowledgeBaseParams',
'KnowledgeListParams',
'KnowledgePage',
]

View file

@ -1,21 +0,0 @@
from .document import (
DocumentData,
DocumentDataFailInfo,
DocumentFailedInfo,
DocumentObject,
DocumentSuccessinfo,
)
from .document_edit_params import DocumentEditParams
from .document_list_params import DocumentListParams
from .document_list_resp import DocumentPage
__all__ = [
'DocumentData',
'DocumentObject',
'DocumentSuccessinfo',
'DocumentFailedInfo',
'DocumentDataFailInfo',
'DocumentEditParams',
'DocumentListParams',
'DocumentPage',
]

View file

@ -1,87 +0,0 @@
from typing import List, Optional
from zai.core import BaseModel
class DocumentSuccessinfo(BaseModel):
"""
Represents successful document upload information.
Attributes:
documentId (Optional[str]): Document ID
filename (Optional[str]): Document filename
"""
documentId: Optional[str] = None
filename: Optional[str] = None
class DocumentFailedInfo(BaseModel):
"""
Represents failed document upload information.
Attributes:
failReason (Optional[str]): Reason for upload failure, including: unsupported file format,
file size exceeds limit, knowledge base capacity is full,
capacity limit is 500,000 words
filename (Optional[str]): Document filename
documentId (Optional[str]): Knowledge base ID
"""
failReason: Optional[str] = None
filename: Optional[str] = None
documentId: Optional[str] = None
class DocumentObject(BaseModel):
"""
Document information
Attributes:
successInfos (Optional[List[DocumentSuccessinfo]]): Information about successfully uploaded files
failedInfos (Optional[List[DocumentFailedInfo]]): Information about failed file uploads
"""
successInfos: Optional[List[DocumentSuccessinfo]] = None
failedInfos: Optional[List[DocumentFailedInfo]] = None
class DocumentDataFailInfo(BaseModel):
"""
Document vectorization failure information
Attributes:
embedding_code (Optional[int]): Error code (10001: Knowledge unavailable, knowledge base space limit reached;
10002: Knowledge unavailable, word count exceeded)
embedding_msg (Optional[str]): Failure reason description
"""
embedding_code: Optional[int] = None
embedding_msg: Optional[str] = None
class DocumentData(BaseModel):
"""
Document data information
Attributes:
id (str): Unique knowledge document identifier
custom_separator (List[str]): Custom text slicing rules
sentence_size (str): Text slice size configuration
length (int): File size in bytes
word_num (int): Total word count in the file
name (str): Document file name
url (str): File download URL
embedding_stat (int): Vectorization status (0: vectorizing, 1: completed, 2: failed)
failInfo (Optional[DocumentDataFailInfo]): Failure information when vectorization fails (embedding_stat=2)
"""
id: str = None
custom_separator: List[str] = None
sentence_size: str = None
length: int = None
word_num: int = None
name: str = None
url: str = None
embedding_stat: int = None
failInfo: Optional[DocumentDataFailInfo] = None

View file

@ -1,29 +0,0 @@
from typing import Dict, List, Optional, TypedDict
class DocumentEditParams(TypedDict):
"""
Knowledge parameter type definition
Attributes:
id (str): Knowledge ID
knowledge_type (int): Knowledge type:
1: Article knowledge: supports pdf,url,docx
2: Q&A knowledge-document: supports pdf,url,docx
3: Q&A knowledge-table: supports xlsx
4: Product library-table: supports xlsx
5: Custom: supports pdf,url,docx
custom_separator (Optional[List[str]]): Slice rules when current knowledge type is custom
(knowledge_type=5), default \n
sentence_size (Optional[int]): Slice word count when current knowledge type is custom
(knowledge_type=5), value range: 20-2000, default 300
callback_url (Optional[str]): Callback address
callback_header (Optional[dict]): Header carried during callback
"""
id: str
knowledge_type: int
custom_separator: Optional[List[str]]
sentence_size: Optional[int]
callback_url: Optional[str]
callback_header: Optional[Dict[str, str]]

View file

@ -1,27 +0,0 @@
from __future__ import annotations
from typing import Optional
from typing_extensions import TypedDict
class DocumentListParams(TypedDict, total=False):
"""
File query parameter type definition
Attributes:
purpose (Optional[str]): File purpose
knowledge_id (Optional[str]): When file purpose is retrieval, need to provide the knowledge base ID for query
page (Optional[int]): Page number, default 1
limit (Optional[int]): Number of files to query in the list, default 10
after (Optional[str]): Query file list after specified fileID (required when file purpose is fine-tune)
order (Optional[str]): Sort rule, optional values ['desc', 'asc'], default desc
(required when file purpose is fine-tune)
"""
purpose: Optional[str]
knowledge_id: Optional[str]
page: Optional[int]
limit: Optional[int]
after: Optional[str]
order: Optional[str]

View file

@ -1,11 +0,0 @@
from __future__ import annotations
from typing import List
from zai.core import BaseModel
from zai.types.knowledge.document.document import DocumentData
class DocumentPage(BaseModel):
list: List[DocumentData]
object: str

View file

@ -1,31 +0,0 @@
from typing import Optional
from zai.core import BaseModel
class KnowledgeInfo(BaseModel):
# Knowledge base unique ID
id: Optional[str] = None
# Vectorization model bound to knowledge base
# See model list [Internal Service Open Interface Documentation](https://lslfd0slxc.feishu.cn/docx/YauWdbBiMopV0FxB7KncPWCEn8f#H15NduiQZo3ugmxnWQFcfAHpnQ4)
embedding_id: Optional[str] = None
# Knowledge base name, 100 character limit
name: Optional[str] = None
# User identifier, within 32 characters
customer_identifier: Optional[str] = None
# Knowledge base description, 500 character limit
description: Optional[str] = None
# Background color (enumeration) 'blue', 'red', 'orange', 'purple', 'sky'
background: Optional[str] = None
# Knowledge base icon (enumeration) question: question mark, book: book, seal: seal, wrench: wrench,
# tag: tag, horn: horn, house: house
icon: Optional[str] = None
# Bucket ID, 32 character limit
bucket_id: Optional[str] = None

View file

@ -1,28 +0,0 @@
from __future__ import annotations
from typing import Optional
from typing_extensions import Literal, TypedDict
class KnowledgeBaseParams(TypedDict):
"""
Knowledge base parameters.
Attributes:
embedding_id (int): Embedding ID
name (str): Knowledge base name, limited to 100 characters
customer_identifier (Optional[str]): Customer identifier, limited to 32 characters
description (Optional[str]): Knowledge base description, limited to 500 characters
background (Optional[Literal['blue', 'red', 'orange', 'purple', 'sky']]): Background color
icon (Optional[Literal['question', 'book', 'seal', 'wrench', 'tag', 'horn', 'house']]): Knowledge base icon
bucket_id (Optional[str]): Bucket ID, limited to 32 characters
"""
embedding_id: int
name: str
customer_identifier: Optional[str]
description: Optional[str]
background: Optional[Literal['blue', 'red', 'orange', 'purple', 'sky']] = None
icon: Optional[Literal['question', 'book', 'seal', 'wrench', 'tag', 'horn', 'house']] = None
bucket_id: Optional[str]

View file

@ -1,16 +0,0 @@
from __future__ import annotations
from typing_extensions import TypedDict
class KnowledgeListParams(TypedDict, total=False):
"""
Parameters for listing knowledge resources.
Attributes:
page (int): Page number, default 1, first page
size (int): Number per page, default 10
"""
page: int = 1
size: int = 10

View file

@ -1,11 +0,0 @@
from __future__ import annotations
from typing import List
from zai.core import BaseModel
from zai.types.knowledge.knowledge import KnowledgeInfo
class KnowledgePage(BaseModel):
list: List[KnowledgeInfo]
object: str

View file

@ -1,19 +0,0 @@
from typing import Optional
from zai.core import BaseModel
class KnowledgeStatistics(BaseModel):
"""
Usage statistics
"""
word_num: Optional[int] = None
length: Optional[int] = None
class KnowledgeUsed(BaseModel):
used: Optional[KnowledgeStatistics] = None
"""Used amount"""
total: Optional[KnowledgeStatistics] = None
"""Total knowledge base"""

Binary file not shown.

View file

@ -3,7 +3,7 @@ import logging.config
import time import time
import zai import zai
from zai import ZaiClient from zai import ZaiClient, ZhipuAiClient
def test_completions_sync(logging_conf): def test_completions_sync(logging_conf):

View file

@ -10,7 +10,7 @@ def test_audio_speech(logging_conf):
logging.config.dictConfig(logging_conf) # type: ignore logging.config.dictConfig(logging_conf) # type: ignore
client = ZaiClient() # Fill in your own API Key client = ZaiClient() # Fill in your own API Key
try: try:
speech_file_path = Path(__file__).parent / 'speech.wav' speech_file_path = Path(__file__).parent / 'asr1.wav'
response = client.audio.speech( response = client.audio.speech(
model='cogtts', model='cogtts',
input='Hello, welcome to Z.ai Open Platform', input='Hello, welcome to Z.ai Open Platform',
@ -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 / 'speech.wav', 'rb') as file: with open(Path(__file__).parent / 'asr1.wav', 'rb') as file:
try: try:
speech_file_path = Path(__file__).parent / 'speech.wav' speech_file_path = Path(__file__).parent / 'asr1.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',

View file

@ -1,134 +0,0 @@
import logging
import logging.config
import zai
from zai import ZaiClient
def test_finetuning_create(logging_conf):
logging.config.dictConfig(logging_conf) # type: ignore
client = ZaiClient() # Please fill in your own API Key
try:
job = client.fine_tuning.jobs.create(
model='chatglm3-6b',
training_file='file-20240428021923715-xjng4', # Please fill in the successfully uploaded file id
validation_file='file-20240428021923715-xjng4', # Please fill in the successfully uploaded file id
suffix='demo_test',
)
job_id = job.id
print(job_id)
fine_tuning_job = client.fine_tuning.jobs.retrieve(fine_tuning_job_id=job_id)
print(fine_tuning_job)
# ftjob-20240418110039323-j8lh2
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
def test_finetuning_retrieve(logging_conf):
logging.config.dictConfig(logging_conf) # type: ignore
client = ZaiClient() # Please fill in your own API Key
try:
fine_tuning_job = client.fine_tuning.jobs.retrieve(fine_tuning_job_id='ftjob-20240429112551154-48vq7')
print(fine_tuning_job)
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
def test_finetuning_job_list(logging_conf):
logging.config.dictConfig(logging_conf) # type: ignore
client = ZaiClient() # Please fill in your own API Key
try:
job_list = client.fine_tuning.jobs.list()
print(job_list)
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
def test_finetuning_job_cancel(logging_conf):
logging.config.dictConfig(logging_conf) # type: ignore
client = ZaiClient() # Please fill in your own API Key
try:
cancel = client.fine_tuning.jobs.cancel(fine_tuning_job_id='ftjob-20240429112551154-48vq7')
print(cancel)
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
def test_finetuning_job_delete(logging_conf):
logging.config.dictConfig(logging_conf) # type: ignore
client = ZaiClient() # Please fill in your own API Key
try:
delete = client.fine_tuning.jobs.delete(fine_tuning_job_id='ftjob-20240126113041678-cs6s9')
print(delete)
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
def test_model_check(logging_conf):
logging.config.dictConfig(logging_conf) # type: ignore
client = ZaiClient() # Fill in your own API Key
try:
response = client.chat.completions.create(
model='chatglm3-6b-8572905046912426020-demo_test', # Fill in the model name to call
messages=[
{'role': 'user', 'content': 'You are a helpful, knowledgeable, and versatile AI assistant.'},
{'role': 'user', 'content': 'Create a more precise and engaging slogan'},
],
extra_body={'temperature': 0.5, 'max_tokens': 50},
)
print(response.choices[0].message)
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
def test_model_delete(logging_conf):
logging.config.dictConfig(logging_conf) # type: ignore
client = ZaiClient() # Fill in your own API Key
try:
delete = client.fine_tuning.models.delete(fine_tuned_model='chatglm3-6b-8572905046912426020-demo_test')
print(delete)
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
if __name__ == '__main__':
test_finetuning_create()

View file

@ -1,168 +0,0 @@
from __future__ import annotations
import logging
import logging.config
import os
import pytest
import zai
from zai import ZaiClient
@pytest.fixture(scope='class')
def test_server():
class SharedData:
client = ZaiClient()
test_knowledge_document_id = None
test_knowledge_id = None
return SharedData()
class TestZaiClientKnowledgeServer:
def test_logs(self, logging_conf):
logging.config.dictConfig(logging_conf) # type: ignore
def test_knowledge_create(self, test_server):
try:
result = test_server.client.knowledge.create(
embedding_id=1,
name='test',
description='test',
background='blue',
icon='question',
)
print(result)
test_server.test_knowledge_id = result.id
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
def test_knowledge_document_create(self, test_server, test_file_path):
try:
result = test_server.client.knowledge.document.create(
file=open(os.path.join(test_file_path, 'file.xlsx'), 'rb'),
purpose='retrieval',
knowledge_id=test_server.test_knowledge_id,
sentence_size=202,
)
print(result)
test_server.test_knowledge_document_id = result.successInfos[0].documentId
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
def test_knowledge_modify(self, test_server):
try:
result = test_server.client.knowledge.modify(
knowledge_id=test_server.test_knowledge_id,
embedding_id=1,
name='test1',
background='red',
icon='book',
)
print(result)
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
def test_knowledge_query(self, test_server):
try:
result = test_server.client.knowledge.query()
print(result)
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
def test_knowledge_used(self, test_server):
try:
result = test_server.client.knowledge.used()
print(result)
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
def test_knowledge_document_retrieve(self, test_server, test_file_path):
try:
result = test_server.client.knowledge.document.retrieve(test_server.test_knowledge_document_id)
print(result)
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
def test_knowledge_document_edit(self, test_server):
try:
result = test_server.client.knowledge.document.edit(
document_id=test_server.test_knowledge_document_id,
knowledge_type='1',
sentence_size=204,
)
print(result)
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
def test_knowledge_document_list(self, test_server):
try:
result = test_server.client.knowledge.document.list(test_server.test_knowledge_id, purpose='retrieval')
print(result)
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
def test_knowledge_document_delete(self, test_server):
try:
file1 = test_server.client.knowledge.document.delete(test_server.test_knowledge_document_id)
print(file1)
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)
def test_knowledge_delete(self, test_server):
try:
result = test_server.client.knowledge.delete(knowledge_id=test_server.test_knowledge_id)
print(result)
except zai.core._errors.APIRequestFailedError as err:
print(err)
except zai.core._errors.APIInternalError as err:
print(err)
except zai.core._errors.APIStatusError as err:
print(err)

View file

@ -26,13 +26,6 @@ from zai.types.chat.chat_completion import (
) )
from zai.types.embeddings import Embedding, EmbeddingsResponded from zai.types.embeddings import Embedding, EmbeddingsResponded
from zai.types.files.file_object import FileObject, ListOfFileObject from zai.types.files.file_object import FileObject, ListOfFileObject
from zai.types.fine_tuning import FineTuningJobEvent
from zai.types.fine_tuning.fine_tuning_job import Error, FineTuningJob
from zai.types.fine_tuning.fine_tuning_job import (
Hyperparameters as FineTuningHyperparameters,
)
from zai.types.fine_tuning.fine_tuning_job_event import JobEvent, Metric
from zai.types.fine_tuning.models import FineTunedModelsStatus
from zai.types.image import GeneratedImage, ImagesResponded from zai.types.image import GeneratedImage, ImagesResponded
@ -164,186 +157,6 @@ def test_response_chat_model_cast(R: Type[BaseModel]) -> None:
assert False, f'Unexpected model type: {R}' assert False, f'Unexpected model type: {R}'
@pytest.mark.parametrize(
'R',
[
FineTunedModelsStatus,
],
)
def test_response_finetuned_model_model_cast(R: Type[BaseModel]) -> None:
MockClient._process_response_data = HttpClient._process_response_data
response = httpx.Response(
status_code=200,
content="""{
"request_id": "12345",
"model_name": "my-fine-tuned-model",
"delete_status": "deleted"
}""",
)
opts = FinalRequestOptions.construct(method='get', url='path')
http_response = APIResponse(
raw=response,
cast_type=R,
client=MockClient(),
stream=False,
stream_cls=None,
options=opts,
)
model = http_response.parse()
assert R == model.__class__
assert isinstance(model, FineTunedModelsStatus)
assert model.request_id == '12345'
assert model.model_name == 'my-fine-tuned-model'
assert model.delete_status == 'deleted'
@pytest.mark.parametrize(
'R',
[
FineTuningJob,
],
)
def test_response_job_model_cast(R: Type[BaseModel]) -> None:
MockClient._process_response_data = HttpClient._process_response_data
response = httpx.Response(
status_code=200,
content=""" {
"id": "job123",
"request_id": "req456",
"created_at": 1617181723,
"error": {
"code": "404",
"message": "Not Found",
"param": "model_id"
},
"fine_tuned_model": "ft_model_1",
"finished_at": 1617182000,
"hyperparameters": {
"n_epochs": 10
},
"model": "base_model",
"object": "fine_tuning_job",
"result_files": [
"result1.txt",
"result2.json"
],
"status": "completed",
"trained_tokens": 1000000,
"training_file": "training_data.csv",
"validation_file": "validation_data.csv"
}""",
)
opts = FinalRequestOptions.construct(method='get', url='path')
http_response = APIResponse(
raw=response,
cast_type=R,
client=MockClient(),
stream=False,
stream_cls=None,
options=opts,
)
model = http_response.parse()
assert R == model.__class__
assert isinstance(model, FineTuningJob)
assert model.id == 'job123'
assert model.request_id == 'req456'
assert model.created_at == 1617181723
assert isinstance(model.error, Error)
assert model.error.code == '404'
assert model.error.message == 'Not Found'
assert model.error.param == 'model_id'
assert model.fine_tuned_model == 'ft_model_1'
assert model.finished_at == 1617182000
assert isinstance(model.hyperparameters, FineTuningHyperparameters)
assert model.hyperparameters.n_epochs == 10
assert model.model == 'base_model'
assert model.object == 'fine_tuning_job'
assert model.result_files == ['result1.txt', 'result2.json']
assert model.status == 'completed'
assert model.trained_tokens == 1000000
assert model.training_file == 'training_data.csv'
assert model.validation_file == 'validation_data.csv'
@pytest.mark.parametrize(
'R',
[
FineTuningJobEvent,
],
)
def test_response_joblist_model_cast(R: Type[BaseModel]) -> None:
MockClient._process_response_data = HttpClient._process_response_data
response = httpx.Response(
status_code=200,
content="""{
"object": "fine_tuning_job",
"data": [
{
"object": "job_event",
"id": "event123",
"type": "training",
"created_at": 1617181723,
"level": "info",
"message": "Training has started.",
"data": {
"epoch": 1,
"current_steps": 100,
"total_steps": 1000,
"elapsed_time": "00:10:00",
"remaining_time": "05:20:00",
"trained_tokens": 500000,
"loss": 0.05,
"eval_loss": 0.03,
"acc": 0.9,
"eval_acc": 0.95,
"learning_rate": 0.001
}
}
],
"has_more": false
}""",
)
opts = FinalRequestOptions.construct(method='get', url='path')
http_response = APIResponse(
raw=response,
cast_type=R,
client=MockClient(),
stream=False,
stream_cls=None,
options=opts,
)
model = http_response.parse()
assert R == model.__class__
assert isinstance(model, FineTuningJobEvent)
assert isinstance(model.data, list)
assert isinstance(model.data[0], JobEvent)
assert model.data[0].object == 'job_event'
assert model.data[0].id == 'event123'
assert model.data[0].type == 'training'
assert model.data[0].created_at == 1617181723
assert model.data[0].level == 'info'
assert model.data[0].message == 'Training has started.'
assert isinstance(model.data[0].data, Metric)
assert model.data[0].data.epoch == 1
assert model.data[0].data.current_steps == 100
assert model.data[0].data.total_steps == 1000
assert model.data[0].data.elapsed_time == '00:10:00'
assert model.data[0].data.remaining_time == '05:20:00'
assert model.data[0].data.trained_tokens == 500000
assert model.data[0].data.loss == 0.05
assert model.data[0].data.eval_loss == 0.03
assert model.data[0].data.acc == 0.9
assert model.data[0].data.eval_acc == 0.95
assert model.data[0].data.learning_rate == 0.001
assert model.has_more == False
@pytest.mark.parametrize( @pytest.mark.parametrize(
'R', 'R',
[EmbeddingsResponded], [EmbeddingsResponded],

View file

@ -0,0 +1,35 @@
from zai.types.agents.agents_completion import AgentsCompletion, AgentsError, AgentsCompletionChoice, AgentsCompletionMessage, AgentsCompletionUsage
def test_agents_completion_error_field():
# 构造一个 AgentsError
error = AgentsError(code="404", message="Not Found")
# 构造一个完整的 AgentsCompletion
completion = AgentsCompletion(
agent_id="test_agent",
conversation_id="conv_1",
status="failed",
choices=[
AgentsCompletionChoice(
index=0,
finish_reason="error",
message=AgentsCompletionMessage(content="error", role="system")
)
],
request_id="req_1",
id="id_1",
usage=AgentsCompletionUsage(prompt_tokens=1, completion_tokens=1, total_tokens=2),
error=error
)
# 检查 error 字段是否为 AgentsError 实例
assert isinstance(completion.error, AgentsError)
assert completion.error.code == "404"
assert completion.error.message == "Not Found"
# 检查序列化
as_dict = completion.model_dump()
assert as_dict["error"]["code"] == "404"
assert as_dict["error"]["message"] == "Not Found"
print("test_agents_completion_error_field passed.")