Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 29 additions & 9 deletions libs/langchain/langchain/chat_models/fireworks.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,6 @@
Type,
Union,
)

from langchain.adapters.openai import convert_message_to_dict
from langchain.callbacks.manager import (
AsyncCallbackManagerForLLMRun,
Expand All @@ -32,7 +31,11 @@
SystemMessage,
SystemMessageChunk,
)
from langchain.schema.output import ChatGeneration, ChatGenerationChunk, ChatResult
from langchain.schema.output import (
ChatGeneration,
ChatGenerationChunk,
ChatResult,
)
from langchain.utils.env import get_from_dict_or_env


Expand Down Expand Up @@ -89,6 +92,7 @@ class ChatFireworks(BaseChatModel):
)
fireworks_api_key: Optional[str] = None
max_retries: int = 20
use_retry: bool = True

@root_validator()
def validate_environment(cls, values: Dict) -> Dict:
Expand Down Expand Up @@ -126,7 +130,11 @@ def _generate(
**self.model_kwargs,
}
response = completion_with_retry(
self, run_manager=run_manager, stop=stop, **params
self,
self.use_retry,
run_manager=run_manager,
stop=stop,
**params,
)
return self._create_chat_result(response)

Expand All @@ -144,7 +152,7 @@ async def _agenerate(
**self.model_kwargs,
}
response = await acompletion_with_retry(
self, run_manager=run_manager, stop=stop, **params
self, self.use_retry, run_manager=run_manager, stop=stop, **params
)
return self._create_chat_result(response)

Expand Down Expand Up @@ -187,7 +195,7 @@ def _stream(
**self.model_kwargs,
}
for chunk in completion_with_retry(
self, run_manager=run_manager, stop=stop, **params
self, self.use_retry, run_manager=run_manager, stop=stop, **params
):
choice = chunk.choices[0]
chunk = _convert_delta_to_message_chunk(choice.delta, default_chunk_class)
Expand Down Expand Up @@ -216,7 +224,7 @@ async def _astream(
**self.model_kwargs,
}
async for chunk in await acompletion_with_retry_streaming(
self, run_manager=run_manager, stop=stop, **params
self, self.use_retry, run_manager=run_manager, stop=stop, **params
):
choice = chunk.choices[0]
chunk = _convert_delta_to_message_chunk(choice.delta, default_chunk_class)
Expand All @@ -230,8 +238,18 @@ async def _astream(
await run_manager.on_llm_new_token(token=chunk.content, chunk=chunk)


def conditional_decorator(condition, decorator):
def actual_decorator(func):
if condition:
return decorator(func)
return func

return actual_decorator


def completion_with_retry(
llm: ChatFireworks,
use_retry: bool,
*,
run_manager: Optional[CallbackManagerForLLMRun] = None,
**kwargs: Any,
Expand All @@ -241,7 +259,7 @@ def completion_with_retry(

retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)

@retry_decorator
@conditional_decorator(use_retry, retry_decorator)
def _completion_with_retry(**kwargs: Any) -> Any:
return fireworks.client.ChatCompletion.create(
**kwargs,
Expand All @@ -252,6 +270,7 @@ def _completion_with_retry(**kwargs: Any) -> Any:

async def acompletion_with_retry(
llm: ChatFireworks,
use_retry: bool,
*,
run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
**kwargs: Any,
Expand All @@ -261,7 +280,7 @@ async def acompletion_with_retry(

retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)

@retry_decorator
@conditional_decorator(use_retry, retry_decorator)
async def _completion_with_retry(**kwargs: Any) -> Any:
return await fireworks.client.ChatCompletion.acreate(
**kwargs,
Expand All @@ -272,6 +291,7 @@ async def _completion_with_retry(**kwargs: Any) -> Any:

async def acompletion_with_retry_streaming(
llm: ChatFireworks,
use_retry: bool,
*,
run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
**kwargs: Any,
Expand All @@ -281,7 +301,7 @@ async def acompletion_with_retry_streaming(

retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)

@retry_decorator
@conditional_decorator(use_retry, retry_decorator)
async def _completion_with_retry(**kwargs: Any) -> Any:
return fireworks.client.ChatCompletion.acreate(
**kwargs,
Expand Down
Loading