From 681a95e37b94043194d21b1afcb1a16f95a761dd Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 4 May 2024 19:35:37 -0700 Subject: [PATCH] fix(assistants/main.py): support `litellm.create_thread()` call --- litellm/__init__.py | 2 +- litellm/assistants/main.py | 116 +++++++++++++++++++++++++++++++ litellm/llms/openai.py | 109 ++++++----------------------- litellm/tests/test_assistants.py | 25 +++++-- litellm/types/llms/__init__.py | 3 + litellm/types/llms/openai.py | 80 +++++++++++++++++++++ litellm/types/router.py | 67 +++++++++++++++++- 7 files changed, 308 insertions(+), 94 deletions(-) create mode 100644 litellm/types/llms/__init__.py create mode 100644 litellm/types/llms/openai.py diff --git a/litellm/__init__.py b/litellm/__init__.py index dc640f0e9f0..b05c1c910db 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -605,7 +605,6 @@ all_embedding_models = ( ####### IMAGE GENERATION MODELS ################### openai_image_generation_models = ["dall-e-2", "dall-e-3"] - from .timeout import timeout from .utils import ( client, @@ -694,3 +693,4 @@ from .exceptions import ( from .budget_manager import BudgetManager from .proxy.proxy_cli import run_server from .router import Router +from .assistants.main import * diff --git a/litellm/assistants/main.py b/litellm/assistants/main.py index 0d32164829e..16a1f973cea 100644 --- a/litellm/assistants/main.py +++ b/litellm/assistants/main.py @@ -1,2 +1,118 @@ # What is this? ## Main file for assistants API logic +from typing import Iterable +import os +import litellm +from openai import OpenAI +from litellm import client +from litellm.utils import supports_httpx_timeout +from ..llms.openai import OpenAIAssistantsAPI +from ..types.llms.openai import * +from ..types.router import * + +####### ENVIRONMENT VARIABLES ################### +openai_assistants_api = OpenAIAssistantsAPI() + +### ASSISTANTS ### + +### THREADS ### + + +def create_thread( + custom_llm_provider: Literal["openai"], + messages: Optional[Iterable[OpenAICreateThreadParamsMessage]] = None, + metadata: Optional[dict] = None, + tool_resources: Optional[OpenAICreateThreadParamsToolResources] = None, + client: Optional[OpenAI] = None, + **kwargs +) -> Thread: + """ + - get the llm provider + - if openai - route it there + - pass through relevant params + + ``` + from litellm import create_thread + + create_thread( + custom_llm_provider="openai", + ### OPTIONAL ### + messages = { + "role": "user", + "content": "Hello, what is AI?" + }, + { + "role": "user", + "content": "How does AI work? Explain it in simple terms." + }] + ) + ``` + """ + optional_params = GenericLiteLLMParams(**kwargs) + + ### TIMEOUT LOGIC ### + timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 + # set timeout for 10 minutes by default + + if ( + timeout is not None + and isinstance(timeout, httpx.Timeout) + and supports_httpx_timeout(custom_llm_provider) == False + ): + read_timeout = timeout.read or 600 + timeout = read_timeout # default 10 min timeout + elif timeout is not None and not isinstance(timeout, httpx.Timeout): + timeout = float(timeout) # type: ignore + elif timeout is None: + timeout = 600.0 + + response: Optional[Thread] = None + if custom_llm_provider == "openai": + api_base = ( + optional_params.api_base # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there + or litellm.api_base + or os.getenv("OPENAI_API_BASE") + or "https://api.openai.com/v1" + ) + organization = ( + optional_params.organization + or litellm.organization + or os.getenv("OPENAI_ORGANIZATION", None) + or None # default - https://github.com/openai/openai-python/blob/284c1799070c723c6a553337134148a7ab088dd8/openai/util.py#L105 + ) + # set API KEY + api_key = ( + optional_params.api_key + or litellm.api_key # for deepinfra/perplexity/anyscale we check in get_llm_provider and pass in the api key from there + or litellm.openai_key + or os.getenv("OPENAI_API_KEY") + ) + response = openai_assistants_api.create_thread( + messages=messages, + metadata=metadata, + api_base=api_base, + api_key=api_key, + timeout=timeout, + max_retries=optional_params.max_retries, + organization=organization, + client=client, + ) + else: + raise litellm.exceptions.BadRequestError( + message="LiteLLM doesn't support {} for 'create_thread'. Only 'openai' is supported.".format( + custom_llm_provider + ), + model="n/a", + llm_provider=custom_llm_provider, + response=httpx.Response( + status_code=400, + content="Unsupported provider", + request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore + ), + ) + return response + + +### MESSAGES ### + +### RUNS ### diff --git a/litellm/llms/openai.py b/litellm/llms/openai.py index a6d6f4109fe..9cc6d86bb68 100644 --- a/litellm/llms/openai.py +++ b/litellm/llms/openai.py @@ -27,73 +27,7 @@ import aiohttp, requests import litellm from .prompt_templates.factory import prompt_factory, custom_prompt from openai import OpenAI, AsyncOpenAI -from openai.types.beta.threads.message_content import MessageContent -from openai.types.beta.threads.message_create_params import Attachment -from openai.types.beta.threads.message import Message as OpenAIMessage -from openai.types.beta.thread_create_params import ( - Message as OpenAICreateThreadParamsMessage, -) -from openai.types.beta.assistant_tool_param import AssistantToolParam -from openai.types.beta.threads.run import Run -from openai.types.beta.assistant import Assistant -from openai.pagination import SyncCursorPage - -from typing import TypedDict, List, Optional - - -class NotGiven: - """ - A sentinel singleton class used to distinguish omitted keyword arguments - from those passed in with the value None (which may have different behavior). - - For example: - - ```py - def get(timeout: Union[int, NotGiven, None] = NotGiven()) -> Response: - ... - - - get(timeout=1) # 1s timeout - get(timeout=None) # No timeout - get() # Default timeout behavior, which may not be statically known at the method definition. - ``` - """ - - def __bool__(self) -> Literal[False]: - return False - - @override - def __repr__(self) -> str: - return "NOT_GIVEN" - - -NOT_GIVEN = NotGiven() - - -class MessageData(TypedDict): - role: Literal["user", "assistant"] - content: str - attachments: Optional[List[Attachment]] - metadata: Optional[dict] - - -class Thread(BaseModel): - id: str - """The identifier, which can be referenced in API endpoints.""" - - created_at: int - """The Unix timestamp (in seconds) for when the thread was created.""" - - metadata: Optional[object] = None - """Set of 16 key-value pairs that can be attached to an object. - - This can be useful for storing additional information about the object in a - structured format. Keys can be a maximum of 64 characters long and values can be - a maxium of 512 characters long. - """ - - object: Literal["thread"] - """The object type, which is always `thread`.""" +from ..types.llms.openai import * class OpenAIError(Exception): @@ -1321,22 +1255,22 @@ class OpenAIAssistantsAPI(BaseLLM): def get_openai_client( self, - api_key: str, + api_key: Optional[str], api_base: Optional[str], timeout: Union[float, httpx.Timeout], - max_retries: int, + max_retries: Optional[int], organization: Optional[str], client: Optional[OpenAI] = None, ) -> OpenAI: + received_args = locals() if client is None: - openai_client = OpenAI( - api_key=api_key, - base_url=api_base, - http_client=litellm.client_session, - timeout=timeout, - max_retries=max_retries, - organization=organization, - ) + data = {} + for k, v in received_args.items(): + if k == "self" or k == "client": + pass + elif v is not None: + data[k] = v + openai_client = OpenAI(**data) # type: ignore else: openai_client = client @@ -1428,16 +1362,14 @@ class OpenAIAssistantsAPI(BaseLLM): def create_thread( self, - metadata: dict, - api_key: str, + metadata: Optional[dict], + api_key: Optional[str], api_base: Optional[str], timeout: Union[float, httpx.Timeout], - max_retries: int, + max_retries: Optional[int], organization: Optional[str], client: Optional[OpenAI], - messages: Union[ - Iterable[OpenAICreateThreadParamsMessage], NotGiven - ] = NOT_GIVEN, + messages: Optional[Iterable[OpenAICreateThreadParamsMessage]], ) -> Thread: """ Here's an example: @@ -1458,10 +1390,13 @@ class OpenAIAssistantsAPI(BaseLLM): client=client, ) - message_thread = openai_client.beta.threads.create( - messages=messages, # type: ignore - metadata=metadata, - ) + data = {} + if messages is not None: + data["messages"] = messages # type: ignore + if metadata is not None: + data["metadata"] = metadata # type: ignore + + message_thread = openai_client.beta.threads.create(**data) # type: ignore return Thread(**message_thread.dict()) diff --git a/litellm/tests/test_assistants.py b/litellm/tests/test_assistants.py index 9b8585ec6da..58c8c4c1f84 100644 --- a/litellm/tests/test_assistants.py +++ b/litellm/tests/test_assistants.py @@ -10,6 +10,7 @@ sys.path.insert( ) # Adds the parent directory to the system path import pytest, logging, asyncio import litellm +from litellm import create_thread from litellm.llms.openai import ( OpenAIAssistantsAPI, MessageData, @@ -25,7 +26,23 @@ V0 Scope: """ -def test_create_thread() -> Thread: +def test_create_thread_litellm(): + message: MessageData = {"role": "user", "content": "Hey, how's it going?"} # type: ignore + new_thread = create_thread( + custom_llm_provider="openai", + messages=[message], + ) + + assert isinstance( + new_thread, Thread + ), f"type of thread={type(new_thread)}. Expected Thread-type" + return new_thread + + +test_create_thread_litellm() + + +def test_create_thread_openai_direct() -> Thread: openai_api = OpenAIAssistantsAPI() message: MessageData = {"role": "user", "content": "Hey, how's it going?"} # type: ignore @@ -48,7 +65,7 @@ def test_create_thread() -> Thread: return new_thread -def test_add_message(): +def test_add_message_openai_direct(): openai_api = OpenAIAssistantsAPI() # create thread new_thread = test_create_thread() @@ -70,7 +87,7 @@ def test_add_message(): assert isinstance(added_message, Message) -def test_get_thread(): +def test_get_thread_openai_direct(): openai_api = OpenAIAssistantsAPI() ## create a thread w/ message ### @@ -92,7 +109,7 @@ def test_get_thread(): return new_thread -def test_run_thread(): +def test_run_thread_openai_direct(): """ - Get Assistants - Create thread diff --git a/litellm/types/llms/__init__.py b/litellm/types/llms/__init__.py new file mode 100644 index 00000000000..14952c9aecf --- /dev/null +++ b/litellm/types/llms/__init__.py @@ -0,0 +1,3 @@ +__all__ = ["openai"] + +from . import openai diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py new file mode 100644 index 00000000000..f9f7b3bf0b7 --- /dev/null +++ b/litellm/types/llms/openai.py @@ -0,0 +1,80 @@ +from typing import ( + Optional, + Union, + Any, + BinaryIO, + Literal, + Annotated, + Iterable, +) +from typing_extensions import override +from pydantic import BaseModel + +from openai.types.beta.threads.message_content import MessageContent +from openai.types.beta.threads.message_create_params import Attachment +from openai.types.beta.threads.message import Message as OpenAIMessage +from openai.types.beta.thread_create_params import ( + Message as OpenAICreateThreadParamsMessage, + ToolResources as OpenAICreateThreadParamsToolResources, +) +from openai.types.beta.assistant_tool_param import AssistantToolParam +from openai.types.beta.threads.run import Run +from openai.types.beta.assistant import Assistant +from openai.pagination import SyncCursorPage + +from typing import TypedDict, List, Optional + + +class NotGiven: + """ + A sentinel singleton class used to distinguish omitted keyword arguments + from those passed in with the value None (which may have different behavior). + + For example: + + ```py + def get(timeout: Union[int, NotGiven, None] = NotGiven()) -> Response: + ... + + + get(timeout=1) # 1s timeout + get(timeout=None) # No timeout + get() # Default timeout behavior, which may not be statically known at the method definition. + ``` + """ + + def __bool__(self) -> Literal[False]: + return False + + @override + def __repr__(self) -> str: + return "NOT_GIVEN" + + +NOT_GIVEN = NotGiven() + + +class MessageData(TypedDict): + role: Literal["user", "assistant"] + content: str + attachments: Optional[List[Attachment]] + metadata: Optional[dict] + + +class Thread(BaseModel): + id: str + """The identifier, which can be referenced in API endpoints.""" + + created_at: int + """The Unix timestamp (in seconds) for when the thread was created.""" + + metadata: Optional[object] = None + """Set of 16 key-value pairs that can be attached to an object. + + This can be useful for storing additional information about the object in a + structured format. Keys can be a maximum of 64 characters long and values can be + a maxium of 512 characters long. + """ + + object: Literal["thread"] + """The object type, which is always `thread`.""" diff --git a/litellm/types/router.py b/litellm/types/router.py index 068a99b0059..d6b698f01e2 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -97,8 +97,11 @@ class ModelInfo(BaseModel): setattr(self, key, value) -class LiteLLM_Params(BaseModel): - model: str +class GenericLiteLLMParams(BaseModel): + """ + LiteLLM Params without 'model' arg (used across completion / assistants api) + """ + custom_llm_provider: Optional[str] = None tpm: Optional[int] = None rpm: Optional[int] = None @@ -121,6 +124,66 @@ class LiteLLM_Params(BaseModel): aws_secret_access_key: Optional[str] = None aws_region_name: Optional[str] = None + def __init__( + self, + custom_llm_provider: Optional[str] = None, + max_retries: Optional[Union[int, str]] = None, + tpm: Optional[int] = None, + rpm: Optional[int] = None, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + api_version: Optional[str] = None, + timeout: Optional[Union[float, str]] = None, # if str, pass in as os.environ/ + stream_timeout: Optional[Union[float, str]] = ( + None # timeout when making stream=True calls, if str, pass in as os.environ/ + ), + organization: Optional[str] = None, # for openai orgs + ## VERTEX AI ## + vertex_project: Optional[str] = None, + vertex_location: Optional[str] = None, + ## AWS BEDROCK / SAGEMAKER ## + aws_access_key_id: Optional[str] = None, + aws_secret_access_key: Optional[str] = None, + aws_region_name: Optional[str] = None, + **params + ): + args = locals() + args.pop("max_retries", None) + args.pop("self", None) + args.pop("params", None) + args.pop("__class__", None) + if max_retries is not None and isinstance(max_retries, str): + max_retries = int(max_retries) # cast to int + super().__init__(max_retries=max_retries, **args, **params) + + class Config: + extra = "allow" + arbitrary_types_allowed = True + + def __contains__(self, key): + # Define custom behavior for the 'in' operator + return hasattr(self, key) + + def get(self, key, default=None): + # Custom .get() method to access attributes with a default value if the attribute doesn't exist + return getattr(self, key, default) + + def __getitem__(self, key): + # Allow dictionary-style access to attributes + return getattr(self, key) + + def __setitem__(self, key, value): + # Allow dictionary-style assignment of attributes + setattr(self, key, value) + + +class LiteLLM_Params(GenericLiteLLMParams): + """ + LiteLLM Params with 'model' requirement - used for completions + """ + + model: str + def __init__( self, model: str,