99from ...utils .error_handler import InvokeError , InvokeConnectionError , InvokeRateLimitError , InvokeAuthorizationError , \
1010 InvokeBadRequestError
1111
12-
1312class 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