diff --git a/litellm/llms/aleph_alpha.py b/litellm/llms/aleph_alpha.py index c513cb105a0..c5e8cded2dc 100644 --- a/litellm/llms/aleph_alpha.py +++ b/litellm/llms/aleph_alpha.py @@ -1,4 +1,5 @@ -import os, json +import os +import json from enum import Enum import requests import time @@ -13,126 +14,112 @@ class AlephAlphaError(Exception): self.message ) # Call the base class constructor with the parameters it needs +def validate_environment(api_key): + headers = { + "accept": "application/json", + "content-type": "application/json", + } + if api_key: + headers["Authorization"] = f"Bearer {api_key}" + return headers -class AlephAlphaLLM: - def __init__( - self, encoding, default_max_tokens_to_sample, logging_obj, api_key=None - ): - self.encoding = encoding - self.default_max_tokens_to_sample = default_max_tokens_to_sample - self.completion_url = "https://api.aleph-alpha.com/complete" - self.api_key = api_key - self.logging_obj = logging_obj - self.validate_environment(api_key=api_key) - - def validate_environment( - self, api_key - ): # set up the environment required to run the model - # set the api key - if self.api_key == None: - raise ValueError( - "Missing Aleph Alpha API Key - A call is being made to Aleph Alpha but no key is set either in the environment variables or via params" - ) - self.api_key = api_key - self.headers = { - "accept": "application/json", - "content-type": "application/json", - "Authorization": "Bearer " + self.api_key, - } - - def completion( - self, - model: str, - messages: list, - model_response: ModelResponse, - print_verbose: Callable, - optional_params=None, - litellm_params=None, - logger_fn=None, - ): # logic for parsing in - calling - parsing out model completion calls - model = model - prompt = "" - if "control" in model: # follow the ###Instruction / ###Response format - for idx, message in enumerate(messages): - if "role" in message: - if idx == 0: # set first message as instruction (required), let later user messages be input - prompt += f"###Instruction: {message['content']}" - else: - if message["role"] == "system": - prompt += ( - f"###Instruction: {message['content']}" - ) - elif message["role"] == "user": - prompt += ( - f"###Input: {message['content']}" - ) - else: - prompt += ( - f"###Response: {message['content']}" - ) +def completion( + model: str, + messages: list, + model_response: ModelResponse, + print_verbose: Callable, + encoding, + api_key, + logging_obj, + optional_params=None, + litellm_params=None, + logger_fn=None, + default_max_tokens_to_sample=None, +): + headers = validate_environment(api_key) + completion_url = "https://api.aleph-alpha.com/complete" + model = model + prompt = "" + if "control" in model: # follow the ###Instruction / ###Response format + for idx, message in enumerate(messages): + if "role" in message: + if idx == 0: # set first message as instruction (required), let later user messages be input + prompt += f"###Instruction: {message['content']}" else: - prompt += f"{message['content']}" - else: - prompt = " ".join(message["content"] for message in messages) - data = { - "model": model, - "prompt": prompt, - "maximum_tokens": optional_params["maximum_tokens"] if "maximum_tokens" in optional_params else self.default_max_tokens_to_sample, # required input - **optional_params, - } + if message["role"] == "system": + prompt += ( + f"###Instruction: {message['content']}" + ) + elif message["role"] == "user": + prompt += ( + f"###Input: {message['content']}" + ) + else: + prompt += ( + f"###Response: {message['content']}" + ) + else: + prompt += f"{message['content']}" + else: + prompt = " ".join(message["content"] for message in messages) + data = { + "model": model, + "prompt": prompt, + "maximum_tokens": optional_params["maximum_tokens"] if "maximum_tokens" in optional_params else default_max_tokens_to_sample, # required input + **optional_params, + } - ## LOGGING - self.logging_obj.pre_call( + ## LOGGING + logging_obj.pre_call( input=prompt, - api_key=self.api_key, + api_key=api_key, additional_args={"complete_input_dict": data}, ) - ## COMPLETION CALL - response = requests.post( - self.completion_url, headers=self.headers, data=json.dumps(data), stream=optional_params["stream"] if "stream" in optional_params else False - ) - if "stream" in optional_params and optional_params["stream"] == True: - return response.iter_lines() - else: - ## LOGGING - self.logging_obj.post_call( + ## COMPLETION CALL + response = requests.post( + completion_url, headers=headers, data=json.dumps(data), stream=optional_params["stream"] if "stream" in optional_params else False + ) + if "stream" in optional_params and optional_params["stream"] == True: + return response.iter_lines() + else: + ## LOGGING + logging_obj.post_call( input=prompt, - api_key=self.api_key, + api_key=api_key, original_response=response.text, additional_args={"complete_input_dict": data}, ) - print_verbose(f"raw model_response: {response.text}") - ## RESPONSE OBJECT - completion_response = response.json() - if "error" in completion_response: - raise AlephAlphaError( - message=completion_response["error"], - status_code=response.status_code, - ) - else: - try: - model_response["choices"][0]["message"]["content"] = completion_response["completions"][0]["completion"] - except: - raise AlephAlphaError(message=json.dumps(completion_response), status_code=response.status_code) - - ## CALCULATING USAGE - baseten charges on time, not tokens - have some mapping of cost here. - prompt_tokens = len( - self.encoding.encode(prompt) - ) - completion_tokens = len( - self.encoding.encode(model_response["choices"][0]["message"]["content"]) + print_verbose(f"raw model_response: {response.text}") + ## RESPONSE OBJECT + completion_response = response.json() + if "error" in completion_response: + raise AlephAlphaError( + message=completion_response["error"], + status_code=response.status_code, ) + else: + try: + model_response["choices"][0]["message"]["content"] = completion_response["completions"][0]["completion"] + except: + raise AlephAlphaError(message=json.dumps(completion_response), status_code=response.status_code) - model_response["created"] = time.time() - model_response["model"] = model - model_response["usage"] = { - "prompt_tokens": prompt_tokens, - "completion_tokens": completion_tokens, - "total_tokens": prompt_tokens + completion_tokens, - } - return model_response + ## CALCULATING USAGE - baseten charges on time, not tokens - have some mapping of cost here. + prompt_tokens = len( + encoding.encode(prompt) + ) + completion_tokens = len( + encoding.encode(model_response["choices"][0]["message"]["content"]) + ) - def embedding( - self, - ): # logic for parsing in - calling - parsing out model embedding calls - pass + model_response["created"] = time.time() + model_response["model"] = model + model_response["usage"] = { + "prompt_tokens": prompt_tokens, + "completion_tokens": completion_tokens, + "total_tokens": prompt_tokens + completion_tokens, + } + return model_response + +def embedding(): + # logic for parsing in - calling - parsing out model embedding calls + pass diff --git a/litellm/main.py b/litellm/main.py index 05b0f198188..859065328d9 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -25,8 +25,8 @@ from .llms import ai21 from .llms import sagemaker from .llms import bedrock from .llms import huggingface_restapi +from .llms import aleph_alpha from .llms.baseten import BasetenLLM -from .llms.aleph_alpha import AlephAlphaLLM import tiktoken from concurrent.futures import ThreadPoolExecutor @@ -427,17 +427,10 @@ def completion( response = model_response elif model in litellm.aleph_alpha_models: aleph_alpha_key = ( - api_key or litellm.aleph_alpha_key or os.environ.get("ALEPH_ALPHA_API_KEY") + api_key or litellm.aleph_alpha_key or get_secret("ALEPH_ALPHA_API_KEY") or get_secret("ALEPHALPHA_API_KEY") ) - aleph_alpha_client = AlephAlphaLLM( - encoding=encoding, - default_max_tokens_to_sample=litellm.max_tokens, - api_key=aleph_alpha_key, - logging_obj=logging # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements - ) - - model_response = aleph_alpha_client.completion( + model_response = aleph_alpha.completion( model=model, messages=messages, model_response=model_response, @@ -445,6 +438,10 @@ def completion( optional_params=optional_params, litellm_params=litellm_params, logger_fn=logger_fn, + encoding=encoding, + default_max_tokens_to_sample=litellm.max_tokens, + api_key=aleph_alpha_key, + logging_obj=logging # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements ) if "stream" in optional_params and optional_params["stream"] == True: diff --git a/litellm/tests/test_completion.py b/litellm/tests/test_completion.py index d8c8174d63a..a4340a45058 100644 --- a/litellm/tests/test_completion.py +++ b/litellm/tests/test_completion.py @@ -64,6 +64,7 @@ def test_completion_claude(): # print(response) # except Exception as e: # pytest.fail(f"Error occurred: {e}") +# test_completion_aleph_alpha() # def test_completion_aleph_alpha_control_models(): @@ -75,6 +76,7 @@ def test_completion_claude(): # print(response) # except Exception as e: # pytest.fail(f"Error occurred: {e}") +# test_completion_aleph_alpha_control_models() def test_completion_with_litellm_call_id(): try: @@ -126,8 +128,8 @@ def test_completion_claude_stream(): # if "loading" in str(e): # pass # pytest.fail(f"Error occurred: {e}") -# # test_completion_hf_api() +# test_completion_hf_api() # def test_completion_hf_deployed_api(): # try: