Skip to content

Commit 8a09b5b

Browse files
authored
Merge pull request #8 from LLMPages/dev-0.0.6
Enhance Wenxin API Documentation for Improved Readability and Maintainability
2 parents 8822b98 + 60bbeba commit 8a09b5b

16 files changed

Lines changed: 2294 additions & 239 deletions

File tree

‎llm_onesdk/core.py‎

Lines changed: 246 additions & 38 deletions
Large diffs are not rendered by default.

‎llm_onesdk/models/anthropic/api.py‎

Lines changed: 125 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -9,12 +9,21 @@
99
from ...utils.error_handler import InvokeError, InvokeConnectionError, InvokeRateLimitError, InvokeAuthorizationError, \
1010
InvokeBadRequestError
1111

12-
1312
class API(BaseAPI):
13+
"""
14+
API class for interacting with the Anthropic API.
15+
Implements the BaseAPI interface for Anthropic-specific functionality.
16+
"""
1417
BASE_URL = "https://api.anthropic.com"
1518
API_VERSION = "2023-06-01"
1619

1720
def __init__(self, credentials: Dict[str, str]):
21+
"""
22+
Initialize the Anthropic API client.
23+
24+
Args:
25+
credentials (Dict[str, str]): A dictionary containing API credentials.
26+
"""
1827
super().__init__(credentials)
1928
self.api_key = credentials.get("api_key") or os.environ.get("ANTHROPIC_API_KEY")
2029
if not self.api_key:
@@ -31,30 +40,63 @@ def __init__(self, credentials: Dict[str, str]):
3140

3241
@provider_specific
3342
def list_models(self) -> List[Dict]:
34-
"""List available models."""
43+
"""
44+
List available models from Anthropic.
45+
46+
Returns:
47+
List[Dict]: A list of dictionaries containing model information.
48+
"""
3549
logger.info("Fetching available models")
3650
models = self._call_api("/v1/models", method="GET")
3751
logger.info(f"Available models: {[model['id'] for model in models['data']]}")
3852
return models
3953

4054
@provider_specific
4155
def get_model(self, model_id: str) -> Dict:
42-
"""Get information about a specific model."""
56+
"""
57+
Get information about a specific model.
58+
59+
Args:
60+
model_id (str): The ID of the model to retrieve information for.
61+
62+
Returns:
63+
Dict: A dictionary containing model information.
64+
"""
4365
logger.info(f"Fetching information for model: {model_id}")
4466
model_info = self._call_api(f"/v1/models/{model_id}", method="GET")
4567
logger.info(f"Model info for {model_id}: {model_info}")
4668
return model_info
4769

4870
def generate(self, model: str, messages: List[Dict[str, Union[str, List[Dict[str, str]]]]], **kwargs) -> Dict:
49-
"""Generate a response using the specified model."""
71+
"""
72+
Generate a response using the specified model.
73+
74+
Args:
75+
model (str): The ID of the model to use for generation.
76+
messages (List[Dict[str, Union[str, List[Dict[str, str]]]]]): The conversation history.
77+
**kwargs: Additional keyword arguments for the API call.
78+
79+
Returns:
80+
Dict: The generated response.
81+
"""
5082
logger.info(f"Generating response with model: {model}")
51-
max_tokens = kwargs.pop('max_tokens', 1000) # 默认值设为1000,可以根据需要调整
83+
max_tokens = kwargs.pop('max_tokens', 1000) # Default value set to 1000, can be adjusted as needed
5284
return self._call_api("/v1/messages", model=model, messages=messages, max_tokens=max_tokens, stream=False,
5385
**kwargs)
5486

5587
def stream_generate(self, model: str, messages: List[Dict[str, Union[str, List[Dict[str, str]]]]],
5688
**kwargs) -> Generator:
57-
"""Generate a streaming response using the specified model."""
89+
"""
90+
Generate a streaming response using the specified model.
91+
92+
Args:
93+
model (str): The ID of the model to use for generation.
94+
messages (List[Dict[str, Union[str, List[Dict[str, str]]]]]): The conversation history.
95+
**kwargs: Additional keyword arguments for the API call.
96+
97+
Yields:
98+
Dict: Chunks of the generated response.
99+
"""
58100
logger.info(f"Generating streaming response with model: {model}")
59101
max_tokens = kwargs.pop('max_tokens', 1000)
60102
response = self._call_api("/v1/messages", model=model, messages=messages, max_tokens=max_tokens, stream=True,
@@ -66,14 +108,36 @@ def stream_generate(self, model: str, messages: List[Dict[str, Union[str, List[D
66108
yield {'delta': {'text': content_item['text']}}
67109

68110
def count_tokens(self, model: str, messages: List[Dict[str, Union[str, List[Dict[str, str]]]]]) -> int:
69-
"""Count tokens in a message."""
111+
"""
112+
Count tokens in a message.
113+
114+
Args:
115+
model (str): The ID of the model to use for token counting.
116+
messages (List[Dict[str, Union[str, List[Dict[str, str]]]]]): The messages to count tokens for.
117+
118+
Returns:
119+
int: The number of tokens in the messages.
120+
"""
70121
logger.info(f"Counting tokens for model: {model}")
71122
response = self._call_api("/v1/messages", model=model, messages=messages, max_tokens=1)
72123
token_count = response.get('usage', {}).get('input_tokens', 0)
73124
logger.info(f"Token count for model {model}: {token_count}")
74125
return token_count
75126

76127
def _call_api(self, endpoint: str, **kwargs) -> Union[Dict, Generator]:
128+
"""
129+
Make an API call to the Anthropic API.
130+
131+
Args:
132+
endpoint (str): The API endpoint to call.
133+
**kwargs: Additional keyword arguments for the API call.
134+
135+
Returns:
136+
Union[Dict, Generator]: The API response, either as a dictionary or a generator for streaming responses.
137+
138+
Raises:
139+
InvokeError: If there's an error during the API call.
140+
"""
77141
url = urljoin(self.base_url, endpoint)
78142
headers = self.session.headers.copy()
79143
method = kwargs.pop('method', 'POST')
@@ -91,11 +155,11 @@ def _call_api(self, endpoint: str, **kwargs) -> Union[Dict, Generator]:
91155

92156
response = self.session.request(method, url, json=payload, headers=headers, stream=stream)
93157

94-
# 打印响应状态码和头部
158+
# Log response status code and headers
95159
logger.debug(f"Response status code: {response.status_code}")
96160
logger.debug(f"Response headers: {response.headers}")
97161

98-
# 尝试打印响应体,即使状态码不是 200
162+
# Attempt to log response body, even if status code is not 200
99163
try:
100164
if not stream:
101165
response_body = response.json()
@@ -124,6 +188,15 @@ def _call_api(self, endpoint: str, **kwargs) -> Union[Dict, Generator]:
124188
raise self._handle_request_error(e)
125189

126190
def _prepare_payload(self, **kwargs) -> Dict:
191+
"""
192+
Prepare the payload for an API call.
193+
194+
Args:
195+
**kwargs: Keyword arguments to include in the payload.
196+
197+
Returns:
198+
Dict: The prepared payload.
199+
"""
127200
payload = {
128201
"model": kwargs.pop('model'),
129202
"messages": self._process_messages(kwargs.pop('messages', [])),
@@ -134,6 +207,15 @@ def _prepare_payload(self, **kwargs) -> Dict:
134207
return payload
135208

136209
def _process_messages(self, messages: List[Dict[str, Union[str, List[Dict[str, str]]]]]) -> List[Dict]:
210+
"""
211+
Process messages, handling any image content.
212+
213+
Args:
214+
messages (List[Dict[str, Union[str, List[Dict[str, str]]]]]): The messages to process.
215+
216+
Returns:
217+
List[Dict]: The processed messages.
218+
"""
137219
processed_messages = []
138220
for message in messages:
139221
if isinstance(message.get('content'), list):
@@ -148,6 +230,15 @@ def _process_messages(self, messages: List[Dict[str, Union[str, List[Dict[str, s
148230
return processed_messages
149231

150232
def _process_image_content(self, content: Dict) -> Dict:
233+
"""
234+
Process image content, converting file paths to base64-encoded data.
235+
236+
Args:
237+
content (Dict): The image content to process.
238+
239+
Returns:
240+
Dict: The processed image content.
241+
"""
151242
if content.get('source', {}).get('type') == 'path':
152243
with open(content['source']['path'], 'rb') as image_file:
153244
base64_image = base64.b64encode(image_file.read()).decode('utf-8')
@@ -159,6 +250,15 @@ def _process_image_content(self, content: Dict) -> Dict:
159250
return content
160251

161252
def _handle_stream_response(self, response: requests.Response) -> Generator:
253+
"""
254+
Handle a streaming response from the API.
255+
256+
Args:
257+
response (requests.Response): The streaming response object.
258+
259+
Yields:
260+
Dict: Parsed data from the stream.
261+
"""
162262
for line in response.iter_lines():
163263
if line:
164264
line = line.decode('utf-8')
@@ -171,6 +271,15 @@ def _handle_stream_response(self, response: requests.Response) -> Generator:
171271
logger.error(f"Failed to parse streaming response: {line}")
172272

173273
def _handle_request_error(self, error: requests.RequestException) -> InvokeError:
274+
"""
275+
Handle errors from API requests.
276+
277+
Args:
278+
error (requests.RequestException): The error that occurred during the request.
279+
280+
Returns:
281+
InvokeError: An appropriate InvokeError subclass based on the type of error.
282+
"""
174283
if isinstance(error, requests.ConnectionError):
175284
return InvokeConnectionError(str(error))
176285
elif isinstance(error, requests.Timeout):
@@ -186,9 +295,14 @@ def _handle_request_error(self, error: requests.RequestException) -> InvokeError
186295
return InvokeError(str(error))
187296

188297
def set_proxy(self, proxy_url: str):
189-
"""Set a proxy for API calls."""
298+
"""
299+
Set a proxy for API calls.
300+
301+
Args:
302+
proxy_url (str): The URL of the proxy to use.
303+
"""
190304
self.session.proxies = {
191305
'http': proxy_url,
192306
'https': proxy_url
193307
}
194-
logger.info(f"Proxy set to {proxy_url}")
308+
logger.info(f"Proxy set to {proxy_url}")

0 commit comments

Comments
 (0)