diff --git a/src/backend/langflow/auth/auth.py b/src/backend/langflow/auth/auth.py index c4b8ad5b4..e33ac64dd 100644 --- a/src/backend/langflow/auth/auth.py +++ b/src/backend/langflow/auth/auth.py @@ -4,11 +4,13 @@ from passlib.context import CryptContext from jose import JWTError, jwt from datetime import datetime, timedelta, timezone from fastapi.security import OAuth2PasswordBearer -from langflow.models.token import TokenData -from langflow.models.user import get_user, User +from langflow.database.models.token import TokenData +from langflow.database.models.user import get_user, User from sqlalchemy.orm import Session from langflow.database.base import get_session + +# TODO: Move to env - Test propose!!!!! SECRET_KEY = "698619adad2d916f1f32d264540976964b3c0d3828e0870a65add5800a8cc6b9" ALGORITHM = "HS256" ACCESS_TOKEN_EXPIRE_MINUTES = 30 @@ -25,7 +27,7 @@ def get_password_hash(password): return pwd_context.hash(password) -def create_access_token(data: dict, expires_delta: timedelta = None): +def create_access_token(data: dict, expires_delta: timedelta = None): # type: ignore to_encode = data.copy() if expires_delta: expire = datetime.now(timezone.utc) + expires_delta @@ -37,7 +39,7 @@ def create_access_token(data: dict, expires_delta: timedelta = None): def authenticate_user(db: Session, username: str, password: str): if user := get_user(db, username): - return user if verify_password(password, user.hashed_password) else False + return user if verify_password(password, user.password) else False else: return False @@ -52,14 +54,14 @@ async def get_current_user( ) try: payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM]) - username: str = payload.get("sub") + username: str = payload.get("sub") # type: ignore if username is None: raise credentials_exception token_data = TokenData(username=username) except JWTError as e: raise credentials_exception from e - user = get_user(db, username=token_data.username) + user = get_user(db, token_data.username) # type: ignore if user is None: raise credentials_exception return user @@ -68,6 +70,6 @@ async def get_current_user( async def get_current_active_user( current_user: Annotated[User, Depends(get_current_user)] ): - if current_user.disabled: + if current_user.is_disabled: raise HTTPException(status_code=400, detail="Inactive user") return current_user diff --git a/src/backend/langflow/models/token.py b/src/backend/langflow/database/models/token.py similarity index 100% rename from src/backend/langflow/models/token.py rename to src/backend/langflow/database/models/token.py diff --git a/src/backend/langflow/database/models/user.py b/src/backend/langflow/database/models/user.py new file mode 100644 index 000000000..6e13f3e49 --- /dev/null +++ b/src/backend/langflow/database/models/user.py @@ -0,0 +1,37 @@ +from datetime import datetime +from sqlalchemy.orm import Session + +from langflow.database.models.base import SQLModelSerializable, SQLModel +from sqlmodel import Field +from uuid import UUID, uuid4 + + +class User(SQLModelSerializable, table=True): + id: UUID = Field(default_factory=uuid4, primary_key=True, unique=True) + username: str = Field(index=True, unique=True) + password: str = Field() + is_disabled: bool = Field(default=False) + is_superuser: bool = Field(default=False) + create_at: datetime = Field(default_factory=datetime.utcnow) + updated_at: datetime = Field(default_factory=datetime.utcnow) + + +class UserAddModel(SQLModel): + username: str = Field() + password: str = Field() + is_disabled: bool = Field(default=False) + is_superuser: bool = Field(default=False) + + +class UserListModel(SQLModel): + id: UUID = Field(default_factory=uuid4) + username: str = Field() + is_disabled: bool = Field() + is_superuser: bool = Field() + create_at: datetime = Field() + updated_at: datetime = Field() + + +def get_user(db: Session, username: str) -> User: + db_user = db.query(User).filter(User.username == username).first() + return User.from_orm(db_user) if db_user else None # type: ignore diff --git a/src/backend/langflow/main.py b/src/backend/langflow/main.py index 062e0ef84..fed302603 100644 --- a/src/backend/langflow/main.py +++ b/src/backend/langflow/main.py @@ -6,7 +6,7 @@ from fastapi.responses import FileResponse from fastapi.staticfiles import StaticFiles from langflow.api import router -from langflow.routers import login, users, items, health +from langflow.routers import login, users, health from langflow.database.base import create_db_and_tables from langflow.interface.utils import setup_llm_caching from langflow.utils.logger import configure @@ -30,7 +30,6 @@ def create_app(): app.include_router(login.router) app.include_router(users.router) - app.include_router(items.router) app.include_router(health.router) app.include_router(router) @@ -74,8 +73,7 @@ def setup_app(static_files_dir: Optional[Path] = None) -> FastAPI: static_files_dir = get_static_files_dir() if not static_files_dir or not static_files_dir.exists(): - raise RuntimeError( - f"Static files directory {static_files_dir} does not exist.") + raise RuntimeError(f"Static files directory {static_files_dir} does not exist.") app = create_app() setup_static_files(app, static_files_dir) return app diff --git a/src/backend/langflow/models/__init__.py b/src/backend/langflow/models/__init__.py deleted file mode 100644 index e69de29bb..000000000 diff --git a/src/backend/langflow/models/base_control.py b/src/backend/langflow/models/base_control.py deleted file mode 100644 index 9eea9d9f0..000000000 --- a/src/backend/langflow/models/base_control.py +++ /dev/null @@ -1,7 +0,0 @@ -from pydantic import BaseModel -from datetime import datetime - - -class BaseControl(BaseModel): - created_at: datetime - updated_at: datetime diff --git a/src/backend/langflow/models/models.py b/src/backend/langflow/models/models.py deleted file mode 100644 index d86d5f7f0..000000000 --- a/src/backend/langflow/models/models.py +++ /dev/null @@ -1,21 +0,0 @@ -from sqlalchemy import Column, String, Boolean, DateTime -from sqlalchemy.ext.declarative import declarative_base -from sqlalchemy.sql import func -from sqlalchemy.dialects.postgresql import UUID -from uuid import uuid4 - -Base = declarative_base() - - -class User(Base): - __tablename__ = "users" - - id = Column( - UUID(as_uuid=True), primary_key=True, default=uuid4, unique=True, nullable=False - ) - username = Column(String, unique=True, index=True) - email = Column(String, unique=True, index=True) - disabled = Column(Boolean, default=False) - is_superuser = Column(Boolean, default=False) - created_at = Column(DateTime(timezone=True), server_default=func.now()) - updated_at = Column(DateTime(timezone=True), onupdate=func.now()) diff --git a/src/backend/langflow/models/user.py b/src/backend/langflow/models/user.py deleted file mode 100644 index b8f6a3fc4..000000000 --- a/src/backend/langflow/models/user.py +++ /dev/null @@ -1,17 +0,0 @@ -from sqlalchemy.orm import Session -from langflow.models.user import User as DBUser -from langflow.models.base_control import BaseControl -from uuid import UUID - - -class User(BaseControl): - id: UUID - username: str - email: str - disabled: bool = False - is_superuser: bool = False - - -def get_user(db: Session, user_id: UUID) -> User: - db_user = db.query(DBUser).filter(DBUser.id == user_id).first() - return User.from_orm(db_user) if db_user else None # type: ignore diff --git a/src/backend/langflow/routers/items.py b/src/backend/langflow/routers/items.py deleted file mode 100644 index 7ca1ff320..000000000 --- a/src/backend/langflow/routers/items.py +++ /dev/null @@ -1,17 +0,0 @@ -from fastapi import APIRouter, Depends -from ..models.user import User -from ..auth.auth import get_current_active_user - -router = APIRouter() - - -@router.get("/users/all/") -async def read_own_items( - current_user: User = Depends(get_current_active_user) -): - return [ - { - "item_id": "my_id", - "owner": current_user.username - } - ] diff --git a/src/backend/langflow/routers/login.py b/src/backend/langflow/routers/login.py index dba69758f..47839f6f5 100644 --- a/src/backend/langflow/routers/login.py +++ b/src/backend/langflow/routers/login.py @@ -1,34 +1,35 @@ from datetime import timedelta + from fastapi import APIRouter, Depends, HTTPException, status from fastapi.security import OAuth2PasswordRequestForm -from langflow.models.token import Token +from langflow.database.models.token import Token from langflow.auth.auth import ( ACCESS_TOKEN_EXPIRE_MINUTES, authenticate_user, create_access_token, ) + from sqlalchemy.orm import Session from langflow.database.base import get_session -TOKEN_TYPE = "bearer" - router = APIRouter() def create_user_token(user: str) -> dict: access_token_expires = timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES) access_token = create_access_token( - data={"sub": user.username}, expires_delta=access_token_expires + data={"sub": user.username}, expires_delta=access_token_expires # type: ignore ) - return {"access_token": access_token, "token_type": TOKEN_TYPE} + + return {"access_token": access_token, "token_type": "bearer"} @router.post("/token", response_model=Token) -async def login_for_access_token( +async def login_to_get_access_token( form_data: OAuth2PasswordRequestForm = Depends(), db: Session = Depends(get_session) ): if user := authenticate_user(db, form_data.username, form_data.password): - return create_user_token(user) + return create_user_token(user) # type: ignore else: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, diff --git a/src/backend/langflow/routers/users.py b/src/backend/langflow/routers/users.py index 1a9184ec8..f34199d5d 100644 --- a/src/backend/langflow/routers/users.py +++ b/src/backend/langflow/routers/users.py @@ -1,10 +1,79 @@ -from fastapi import APIRouter, Depends -from langflow.models.user import User +from typing import List +from sqlmodel import Session, select +from sqlalchemy.exc import IntegrityError +from fastapi import APIRouter, Depends, HTTPException + +from langflow.database.base import get_session from langflow.auth.auth import get_current_active_user +from langflow.database.models.user import UserAddModel, UserListModel, User -router = APIRouter() +from passlib.context import CryptContext + +router = APIRouter(prefix="/users", tags=["Users"]) -@router.get("/users/me/", response_model=User) -async def read_users_me(current_user: User = Depends(get_current_active_user)): +def get_password_hash(password): + pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto") + return pwd_context.hash(password) + + +@router.get("/user", response_model=UserListModel) +async def read_current_user(current_user: User = Depends(get_current_active_user)): return current_user + + +@router.get("/users", response_model=List[UserListModel]) +async def read_all_users( + skip: int = 0, + limit: int = 10, + _: Session = Depends(get_current_active_user), + db: Session = Depends(get_session), +): + query = select(User) + query = query.offset(skip).limit(limit) + + return db.execute(query).fetchall() + + +@router.post("/user", response_model=User) +async def add_user( + user: UserAddModel, + _: Session = Depends(get_current_active_user), + db: Session = Depends(get_session), +): + new_user = User(**user.dict()) + try: + new_user.password = get_password_hash(user.password) + + db.add(new_user) + db.commit() + db.refresh(new_user) + except IntegrityError as e: + db.rollback() + raise HTTPException( + status_code=400, + detail="User exists", + ) from e + + return new_user + + +# TODO: Remove - Just for testing purposes +@router.post("/super_user", response_model=User) +async def add_super_user_to_testing_purposes(db: Session = Depends(get_session)): + new_user = User(username="superuser", password="12345", is_superuser=True) + + try: + new_user.password = get_password_hash(new_user.password) + + db.add(new_user) + db.commit() + db.refresh(new_user) + except IntegrityError as e: + db.rollback() + raise HTTPException( + status_code=400, + detail="User exists", + ) from e + + return new_user