🎉 feat(langflow): add new files base.py and callback.py
The base.py file contains the following classes and functions: - CacheResponse: a pydantic BaseModel that represents a response containing a dictionary of data - Code: a pydantic BaseModel that represents a code string - Prompt: a pydantic BaseModel that represents a prompt template string - CodeValidationResponse: a pydantic BaseModel that represents a response containing the validation results of code - PromptValidationResponse: a pydantic BaseModel that represents a response containing the validation results of a prompt - validate_prompt: a function that validates a prompt template string and returns a PromptValidationResponse object - check_input_variables: a function that checks if input variables contain invalid characters and returns a list of fixed input variables The callback.py file contains the following classes: - AsyncStreamingLLMCallbackHandler: an AsyncCallbackHandler that handles streaming LLM responses asynchronously - StreamingLLMCallbackHandler: a BaseCallbackHandler that handles streaming LLM responses These files were added to provide support for Langflow's backend API.
This commit is contained in:
parent
bdbb4a8127
commit
3e5878ddc2
2 changed files with 116 additions and 0 deletions
84
src/backend/langflow/api/v1/base.py
Normal file
84
src/backend/langflow/api/v1/base.py
Normal file
|
|
@ -0,0 +1,84 @@
|
||||||
|
from pydantic import BaseModel, validator
|
||||||
|
|
||||||
|
from langflow.interface.utils import extract_input_variables_from_prompt
|
||||||
|
|
||||||
|
|
||||||
|
class CacheResponse(BaseModel):
|
||||||
|
data: dict
|
||||||
|
|
||||||
|
|
||||||
|
class Code(BaseModel):
|
||||||
|
code: str
|
||||||
|
|
||||||
|
|
||||||
|
class Prompt(BaseModel):
|
||||||
|
template: str
|
||||||
|
|
||||||
|
|
||||||
|
# Build ValidationResponse class for {"imports": {"errors": []}, "function": {"errors": []}}
|
||||||
|
class CodeValidationResponse(BaseModel):
|
||||||
|
imports: dict
|
||||||
|
function: dict
|
||||||
|
|
||||||
|
@validator("imports")
|
||||||
|
def validate_imports(cls, v):
|
||||||
|
return v or {"errors": []}
|
||||||
|
|
||||||
|
@validator("function")
|
||||||
|
def validate_function(cls, v):
|
||||||
|
return v or {"errors": []}
|
||||||
|
|
||||||
|
|
||||||
|
class PromptValidationResponse(BaseModel):
|
||||||
|
input_variables: list
|
||||||
|
|
||||||
|
|
||||||
|
INVALID_CHARACTERS = {
|
||||||
|
" ",
|
||||||
|
",",
|
||||||
|
".",
|
||||||
|
":",
|
||||||
|
";",
|
||||||
|
"!",
|
||||||
|
"?",
|
||||||
|
"/",
|
||||||
|
"\\",
|
||||||
|
"(",
|
||||||
|
")",
|
||||||
|
"[",
|
||||||
|
"]",
|
||||||
|
"{",
|
||||||
|
"}",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def validate_prompt(template: str):
|
||||||
|
input_variables = extract_input_variables_from_prompt(template)
|
||||||
|
|
||||||
|
# Check if there are invalid characters in the input_variables
|
||||||
|
input_variables = check_input_variables(input_variables)
|
||||||
|
|
||||||
|
return PromptValidationResponse(input_variables=input_variables)
|
||||||
|
|
||||||
|
|
||||||
|
def check_input_variables(input_variables: list):
|
||||||
|
invalid_chars = []
|
||||||
|
fixed_variables = []
|
||||||
|
for variable in input_variables:
|
||||||
|
new_var = variable
|
||||||
|
for char in INVALID_CHARACTERS:
|
||||||
|
if char in variable:
|
||||||
|
invalid_chars.append(char)
|
||||||
|
new_var = new_var.replace(char, "")
|
||||||
|
fixed_variables.append(new_var)
|
||||||
|
if new_var != variable:
|
||||||
|
input_variables.remove(variable)
|
||||||
|
input_variables.append(new_var)
|
||||||
|
# If any of the input_variables is not in the fixed_variables, then it means that
|
||||||
|
# there are invalid characters in the input_variables
|
||||||
|
if any(var not in fixed_variables for var in input_variables):
|
||||||
|
raise ValueError(
|
||||||
|
f"Invalid input variables: {input_variables}. Please, use something like {fixed_variables} instead."
|
||||||
|
)
|
||||||
|
|
||||||
|
return input_variables
|
||||||
32
src/backend/langflow/api/v1/callback.py
Normal file
32
src/backend/langflow/api/v1/callback.py
Normal file
|
|
@ -0,0 +1,32 @@
|
||||||
|
import asyncio
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from langchain.callbacks.base import AsyncCallbackHandler, BaseCallbackHandler
|
||||||
|
|
||||||
|
from langflow.api.v1.schemas import ChatResponse
|
||||||
|
|
||||||
|
|
||||||
|
# https://github.com/hwchase17/chat-langchain/blob/master/callback.py
|
||||||
|
class AsyncStreamingLLMCallbackHandler(AsyncCallbackHandler):
|
||||||
|
"""Callback handler for streaming LLM responses."""
|
||||||
|
|
||||||
|
def __init__(self, websocket):
|
||||||
|
self.websocket = websocket
|
||||||
|
|
||||||
|
async def on_llm_new_token(self, token: str, **kwargs: Any) -> None:
|
||||||
|
resp = ChatResponse(message=token, type="stream", intermediate_steps="")
|
||||||
|
await self.websocket.send_json(resp.dict())
|
||||||
|
|
||||||
|
|
||||||
|
class StreamingLLMCallbackHandler(BaseCallbackHandler):
|
||||||
|
"""Callback handler for streaming LLM responses."""
|
||||||
|
|
||||||
|
def __init__(self, websocket):
|
||||||
|
self.websocket = websocket
|
||||||
|
|
||||||
|
def on_llm_new_token(self, token: str, **kwargs: Any) -> None:
|
||||||
|
resp = ChatResponse(message=token, type="stream", intermediate_steps="")
|
||||||
|
|
||||||
|
loop = asyncio.get_event_loop()
|
||||||
|
coroutine = self.websocket.send_json(resp.dict())
|
||||||
|
asyncio.run_coroutine_threadsafe(coroutine, loop)
|
||||||
Loading…
Add table
Add a link
Reference in a new issue