🐛 fix(flows.py): add missing import statement for User model

🐛 fix(flows.py): add missing import statement for TYPE_CHECKING
✨ feat(flows.py): add user_id field to FlowCreate model to allow specifying the user for a new flow
✨ feat(flows.py): add user_id field to FlowRead model to include the user_id in the response
✨ feat(flows.py): add user_id field to Flow model and create a relationship with User model
✨ feat(flows.py): add current_user dependency to create_flow endpoint to set the user_id for a new flow
✨ feat(flows.py): add current_user dependency to read_flows endpoint to filter flows by current user
✨ feat(flows.py): add current_user dependency to read_flow endpoint to filter flow by current user
✨ feat(flows.py): add current_user dependency to update_flow endpoint to filter flow by current user
✨ feat(flows.py): add current_user dependency to delete_flow endpoint to filter flow by current user
✨ feat(flows.py): add current_user dependency to create_flows endpoint to set the user_id for new flows
✨ feat(flows.py): add current_user dependency to upload_file endpoint to set the user_id for new flows
✨ feat(flows.py): add current_user dependency to download_file endpoint to filter flows by current user
🐛 fix(flow.py): add missing import statement for User model
✨ feat(flow.py): add user_id field to Flow model to associate a flow with a user
This commit is contained in:
Gabriel Luiz Freitas Almeida 2023-08-25 12:49:24 -03:00
commit b5e4df9943
2 changed files with 78 additions and 20 deletions

View file

@ -4,16 +4,18 @@ from fastapi.encoders import jsonable_encoder
from langflow.api.utils import remove_api_keys from langflow.api.utils import remove_api_keys
from langflow.api.v1.schemas import FlowListCreate, FlowListRead from langflow.api.v1.schemas import FlowListCreate, FlowListRead
from langflow.services.auth.utils import get_current_active_user
from langflow.services.database.models.flow import ( from langflow.services.database.models.flow import (
Flow, Flow,
FlowCreate, FlowCreate,
FlowRead, FlowRead,
FlowUpdate, FlowUpdate,
) )
from langflow.services.database.models.user.user import User
from langflow.services.utils import get_session from langflow.services.utils import get_session
from langflow.services.utils import get_settings_manager from langflow.services.utils import get_settings_manager
import orjson import orjson
from sqlmodel import Session, select from sqlmodel import Session
from fastapi import APIRouter, Depends, HTTPException from fastapi import APIRouter, Depends, HTTPException
from fastapi import File, UploadFile from fastapi import File, UploadFile
@ -23,9 +25,18 @@ router = APIRouter(prefix="/flows", tags=["Flows"])
@router.post("/", response_model=FlowRead, status_code=201) @router.post("/", response_model=FlowRead, status_code=201)
def create_flow(*, session: Session = Depends(get_session), flow: FlowCreate): def create_flow(
*,
session: Session = Depends(get_session),
flow: FlowCreate,
current_user: User = Depends(get_current_active_user),
):
"""Create a new flow.""" """Create a new flow."""
if flow.user_id is None:
flow.user_id = current_user.id
db_flow = Flow.from_orm(flow) db_flow = Flow.from_orm(flow)
session.add(db_flow) session.add(db_flow)
session.commit() session.commit()
session.refresh(db_flow) session.refresh(db_flow)
@ -33,31 +44,49 @@ def create_flow(*, session: Session = Depends(get_session), flow: FlowCreate):
@router.get("/", response_model=list[FlowRead], status_code=200) @router.get("/", response_model=list[FlowRead], status_code=200)
def read_flows(*, session: Session = Depends(get_session)): def read_flows(
*,
session: Session = Depends(get_session),
current_user: User = Depends(get_current_active_user),
):
"""Read all flows.""" """Read all flows."""
try: try:
flows = session.exec(select(Flow)).all() flows = current_user.flows
except Exception as e: except Exception as e:
raise HTTPException(status_code=500, detail=str(e)) from e raise HTTPException(status_code=500, detail=str(e)) from e
return [jsonable_encoder(flow) for flow in flows] return [jsonable_encoder(flow) for flow in flows]
@router.get("/{flow_id}", response_model=FlowRead, status_code=200) @router.get("/{flow_id}", response_model=FlowRead, status_code=200)
def read_flow(*, session: Session = Depends(get_session), flow_id: UUID): def read_flow(
*,
session: Session = Depends(get_session),
flow_id: UUID,
current_user: User = Depends(get_current_active_user),
):
"""Read a flow.""" """Read a flow."""
if flow := session.get(Flow, flow_id): if user_flow := (
return flow session.query(Flow)
.filter(Flow.id == flow_id)
.filter(Flow.user_id == current_user.id)
.first()
):
return user_flow
else: else:
raise HTTPException(status_code=404, detail="Flow not found") raise HTTPException(status_code=404, detail="Flow not found")
@router.patch("/{flow_id}", response_model=FlowRead, status_code=200) @router.patch("/{flow_id}", response_model=FlowRead, status_code=200)
def update_flow( def update_flow(
*, session: Session = Depends(get_session), flow_id: UUID, flow: FlowUpdate *,
session: Session = Depends(get_session),
flow_id: UUID,
flow: FlowUpdate,
current_user: User = Depends(get_current_active_user),
): ):
"""Update a flow.""" """Update a flow."""
db_flow = session.get(Flow, flow_id) db_flow = read_flow(session=session, flow_id=flow_id, current_user=current_user)
if not db_flow: if not db_flow:
raise HTTPException(status_code=404, detail="Flow not found") raise HTTPException(status_code=404, detail="Flow not found")
flow_data = flow.dict(exclude_unset=True) flow_data = flow.dict(exclude_unset=True)
@ -65,6 +94,7 @@ def update_flow(
if settings_manager.settings.REMOVE_API_KEYS: if settings_manager.settings.REMOVE_API_KEYS:
flow_data = remove_api_keys(flow_data) flow_data = remove_api_keys(flow_data)
for key, value in flow_data.items(): for key, value in flow_data.items():
if value is not None:
setattr(db_flow, key, value) setattr(db_flow, key, value)
session.add(db_flow) session.add(db_flow)
session.commit() session.commit()
@ -73,9 +103,14 @@ def update_flow(
@router.delete("/{flow_id}", status_code=200) @router.delete("/{flow_id}", status_code=200)
def delete_flow(*, session: Session = Depends(get_session), flow_id: UUID): def delete_flow(
*,
session: Session = Depends(get_session),
flow_id: UUID,
current_user: User = Depends(get_current_active_user),
):
"""Delete a flow.""" """Delete a flow."""
flow = session.get(Flow, flow_id) flow = read_flow(session=session, flow_id=flow_id, current_user=current_user)
if not flow: if not flow:
raise HTTPException(status_code=404, detail="Flow not found") raise HTTPException(status_code=404, detail="Flow not found")
session.delete(flow) session.delete(flow)
@ -87,10 +122,16 @@ def delete_flow(*, session: Session = Depends(get_session), flow_id: UUID):
@router.post("/batch/", response_model=List[FlowRead], status_code=201) @router.post("/batch/", response_model=List[FlowRead], status_code=201)
def create_flows(*, session: Session = Depends(get_session), flow_list: FlowListCreate): def create_flows(
*,
session: Session = Depends(get_session),
flow_list: FlowListCreate,
current_user: User = Depends(get_current_active_user),
):
"""Create multiple new flows.""" """Create multiple new flows."""
db_flows = [] db_flows = []
for flow in flow_list.flows: for flow in flow_list.flows:
flow.user_id = current_user.id
db_flow = Flow.from_orm(flow) db_flow = Flow.from_orm(flow)
session.add(db_flow) session.add(db_flow)
db_flows.append(db_flow) db_flows.append(db_flow)
@ -102,7 +143,10 @@ def create_flows(*, session: Session = Depends(get_session), flow_list: FlowList
@router.post("/upload/", response_model=List[FlowRead], status_code=201) @router.post("/upload/", response_model=List[FlowRead], status_code=201)
async def upload_file( async def upload_file(
*, session: Session = Depends(get_session), file: UploadFile = File(...) *,
session: Session = Depends(get_session),
file: UploadFile = File(...),
current_user: User = Depends(get_current_active_user),
): ):
"""Upload flows from a file.""" """Upload flows from a file."""
contents = await file.read() contents = await file.read()
@ -111,11 +155,19 @@ async def upload_file(
flow_list = FlowListCreate(**data) flow_list = FlowListCreate(**data)
else: else:
flow_list = FlowListCreate(flows=[FlowCreate(**flow) for flow in data]) flow_list = FlowListCreate(flows=[FlowCreate(**flow) for flow in data])
return create_flows(session=session, flow_list=flow_list) # Now we set the user_id for all flows
for flow in flow_list.flows:
flow.user_id = current_user.id
return create_flows(session=session, flow_list=flow_list, current_user=current_user)
@router.get("/download/", response_model=FlowListRead, status_code=200) @router.get("/download/", response_model=FlowListRead, status_code=200)
async def download_file(*, session: Session = Depends(get_session)): async def download_file(
*,
session: Session = Depends(get_session),
current_user: User = Depends(get_current_active_user),
):
"""Download all flows as a file.""" """Download all flows as a file."""
flows = read_flows(session=session) flows = read_flows(session=session, current_user=current_user)
return FlowListRead(flows=flows) return FlowListRead(flows=flows)

View file

@ -2,9 +2,12 @@
from langflow.services.database.models.base import SQLModelSerializable from langflow.services.database.models.base import SQLModelSerializable
from pydantic import validator from pydantic import validator
from sqlmodel import Field, JSON, Column from sqlmodel import Field, JSON, Column, Relationship
from uuid import UUID, uuid4 from uuid import UUID, uuid4
from typing import Dict, Optional from typing import Dict, Optional, TYPE_CHECKING
if TYPE_CHECKING:
from langflow.services.database.models.user import User
class FlowBase(SQLModelSerializable): class FlowBase(SQLModelSerializable):
@ -31,14 +34,17 @@ class FlowBase(SQLModelSerializable):
class Flow(FlowBase, table=True): class Flow(FlowBase, table=True):
id: UUID = Field(default_factory=uuid4, primary_key=True, unique=True) id: UUID = Field(default_factory=uuid4, primary_key=True, unique=True)
data: Optional[Dict] = Field(default=None, sa_column=Column(JSON)) data: Optional[Dict] = Field(default=None, sa_column=Column(JSON))
user_id: UUID = Field(index=True, foreign_key="user.id")
user: "User" = Relationship(back_populates="flows")
class FlowCreate(FlowBase): class FlowCreate(FlowBase):
pass user_id: Optional[UUID] = None
class FlowRead(FlowBase): class FlowRead(FlowBase):
id: UUID id: UUID
user_id: UUID = Field()
class FlowUpdate(SQLModelSerializable): class FlowUpdate(SQLModelSerializable):