Merge branch 'db' into mainPage
This commit is contained in:
commit
2def9b2e5e
4 changed files with 90 additions and 19 deletions
|
|
@ -1,9 +1,10 @@
|
||||||
from langflow.database.models.flow import Flow
|
from langflow.database.models.flow import Flow
|
||||||
from langflow.processing.process import process_graph_cached
|
from langflow.processing.process import process_graph_cached, process_tweaks
|
||||||
from langflow.utils.logger import logger
|
from langflow.utils.logger import logger
|
||||||
from importlib.metadata import version
|
from importlib.metadata import version
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException
|
from fastapi import APIRouter, Depends, HTTPException
|
||||||
|
from fastapi.security import HTTPBearer
|
||||||
|
|
||||||
from langflow.api.v1.schemas import (
|
from langflow.api.v1.schemas import (
|
||||||
PredictRequest,
|
PredictRequest,
|
||||||
|
|
@ -17,23 +18,40 @@ from sqlmodel import Session
|
||||||
# build router
|
# build router
|
||||||
router = APIRouter(tags=["Base"])
|
router = APIRouter(tags=["Base"])
|
||||||
|
|
||||||
|
security = HTTPBearer()
|
||||||
|
|
||||||
|
|
||||||
|
def get_flow_from_token(
|
||||||
|
bearer: HTTPBearer = Depends(security), session: Session = Depends(get_session)
|
||||||
|
) -> str:
|
||||||
|
# Extract the token, which is the flow_id in this case
|
||||||
|
flow_id = bearer.credentials
|
||||||
|
# Check if the flow_id exists in the database
|
||||||
|
flow = session.get(Flow, flow_id)
|
||||||
|
if flow is None:
|
||||||
|
raise HTTPException(status_code=401, detail="Invalid token")
|
||||||
|
return flow
|
||||||
|
|
||||||
|
|
||||||
@router.get("/all")
|
@router.get("/all")
|
||||||
def get_all():
|
def get_all():
|
||||||
return build_langchain_types_dict()
|
return build_langchain_types_dict()
|
||||||
|
|
||||||
|
|
||||||
@router.post("/predict/{flow_id}", status_code=200, response_model=PredictResponse)
|
@router.post("/predict", response_model=PredictResponse)
|
||||||
async def get_load(
|
async def predict_flow(
|
||||||
predict_request: PredictRequest,
|
predict_request: PredictRequest,
|
||||||
flow_id: str,
|
flow: Flow = Depends(get_flow_from_token),
|
||||||
session: Session = Depends(get_session),
|
|
||||||
):
|
):
|
||||||
|
"""
|
||||||
|
Endpoint to process a message using the flow passed in the bearer token.
|
||||||
|
"""
|
||||||
|
|
||||||
try:
|
try:
|
||||||
flow_obj = session.get(Flow, flow_id)
|
graph_data = flow.data
|
||||||
if flow_obj is None:
|
if predict_request.tweaks:
|
||||||
raise ValueError(f"Flow {flow_id} not found")
|
graph_data = process_tweaks(graph_data, predict_request.tweaks)
|
||||||
graph_data = flow_obj.data
|
|
||||||
response = process_graph_cached(graph_data, predict_request.message)
|
response = process_graph_cached(graph_data, predict_request.message)
|
||||||
return PredictResponse(
|
return PredictResponse(
|
||||||
result=response.get("result", ""),
|
result=response.get("result", ""),
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,6 @@
|
||||||
from typing import Any, Dict, List, Union
|
from typing import Any, Dict, List, Optional, Union
|
||||||
from langflow.database.models.flow import FlowCreate, FlowRead
|
from langflow.database.models.flow import FlowCreate, FlowRead
|
||||||
from pydantic import BaseModel, validator
|
from pydantic import BaseModel, Field, validator
|
||||||
|
|
||||||
|
|
||||||
class GraphData(BaseModel):
|
class GraphData(BaseModel):
|
||||||
|
|
@ -23,6 +23,23 @@ class PredictRequest(BaseModel):
|
||||||
"""Predict request schema."""
|
"""Predict request schema."""
|
||||||
|
|
||||||
message: str
|
message: str
|
||||||
|
tweaks: Optional[Dict[str, Dict[str, str]]] = Field(default_factory=dict)
|
||||||
|
|
||||||
|
class Config:
|
||||||
|
schema_extra = {
|
||||||
|
"example": {
|
||||||
|
"message": "Hello, how are you?",
|
||||||
|
"tweaks": {
|
||||||
|
"dndnode_986363f0-4677-4035-9f38-74b94af5dd78": {
|
||||||
|
"name": "A tool name",
|
||||||
|
"description": "A tool description",
|
||||||
|
},
|
||||||
|
"dndnode_986363f0-4677-4035-9f38-74b94af57378": {
|
||||||
|
"template": "A {template}",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
class PredictResponse(BaseModel):
|
class PredictResponse(BaseModel):
|
||||||
|
|
|
||||||
|
|
@ -170,3 +170,25 @@ def load_flow_from_json(path: str, build=True):
|
||||||
fix_memory_inputs(langchain_object)
|
fix_memory_inputs(langchain_object)
|
||||||
return langchain_object
|
return langchain_object
|
||||||
return graph
|
return graph
|
||||||
|
|
||||||
|
|
||||||
|
def process_tweaks(graph_data: dict, tweaks: dict):
|
||||||
|
"""This function is used to tweak the graph data using the node id and the tweaks dict"""
|
||||||
|
# the tweaks dict is a dict of dicts
|
||||||
|
# the key is the node id and the value is a dict of the tweaks
|
||||||
|
# the dict of tweaks contains the name of a certain parameter and the value to be tweaked
|
||||||
|
|
||||||
|
# We need to process the graph data to add the tweaks
|
||||||
|
nodes = graph_data["data"]["nodes"]
|
||||||
|
for node in nodes:
|
||||||
|
node_id = node["id"]
|
||||||
|
if node_id in tweaks:
|
||||||
|
node_tweaks = tweaks[node_id]
|
||||||
|
template_data = node["data"]["node"]["template"]
|
||||||
|
for tweak_name, tweake_value in node_tweaks.items():
|
||||||
|
if tweak_name in template_data:
|
||||||
|
template_data[tweak_name]["value"] = tweake_value
|
||||||
|
print(
|
||||||
|
f"Something changed in node {node_id} with tweak {tweak_name} and value {tweake_value}"
|
||||||
|
)
|
||||||
|
return graph_data
|
||||||
|
|
|
||||||
|
|
@ -53,14 +53,27 @@ export const getPythonApiCode = (flowId: string): string => {
|
||||||
return `import requests
|
return `import requests
|
||||||
|
|
||||||
FLOW_ID = "${flowId}"
|
FLOW_ID = "${flowId}"
|
||||||
API_URL = f"${window.location.protocol}//${window.location.host}/predict/{FLOW_ID}"
|
API_URL = f"${window.location.protocol}//${window.location.host}/predict"
|
||||||
|
|
||||||
def predict(message):
|
def run_flow(message, tweaks=None):
|
||||||
|
|
||||||
|
if tweaks:
|
||||||
|
payload = {'message': message, 'tweaks': tweaks}
|
||||||
|
else:
|
||||||
payload = {'message': message}
|
payload = {'message': message}
|
||||||
|
|
||||||
|
headers = {'Authorization':
|
||||||
|
f'Bearer {FLOW_ID}',
|
||||||
|
'Content-Type': 'application/json'
|
||||||
|
}
|
||||||
|
|
||||||
response = requests.post(API_URL, json=payload)
|
response = requests.post(API_URL, json=payload)
|
||||||
return response.json()
|
return response.json()
|
||||||
|
|
||||||
print(predict("Your message"))`;
|
# Setup any tweaks you want to apply to the flow
|
||||||
|
tweaks = {} # {"nodeId": {"key": "value"}, "nodeId2": {"key": "value"}}
|
||||||
|
|
||||||
|
print(run_flow("Your message", tweaks=tweaks))`;
|
||||||
};
|
};
|
||||||
|
|
||||||
/**
|
/**
|
||||||
|
|
@ -71,8 +84,9 @@ export const getPythonApiCode = (flowId: string): string => {
|
||||||
export const getCurlCode = (flowId: string): string => {
|
export const getCurlCode = (flowId: string): string => {
|
||||||
return `curl -X POST \\
|
return `curl -X POST \\
|
||||||
-H "Content-Type: application/json" \\
|
-H "Content-Type: application/json" \\
|
||||||
|
-H "Authorization: Bearer ${flowId}" \\
|
||||||
-d '{"message": "Your message"}' \\
|
-d '{"message": "Your message"}' \\
|
||||||
${window.location.protocol}//${window.location.host}/predict/${flowId}`;
|
${window.location.protocol}//${window.location.host}/predict`;
|
||||||
};
|
};
|
||||||
|
|
||||||
/**
|
/**
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue