auth unit tests & ruff visual fix
This commit is contained in:
@@ -1,16 +1,19 @@
|
||||
|
||||
from uuid import UUID
|
||||
from src.models.database_models.model import engine, RefreshTokens
|
||||
|
||||
from sqlalchemy import and_, not_, select
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from src.models.database_models.model import RefreshTokens, engine
|
||||
from src.models.pydantic_models.model import RefreshTokensOut
|
||||
|
||||
|
||||
class JwtCrudActions:
|
||||
def __init__(self) -> None:
|
||||
self.Session=sessionmaker(bind=engine)
|
||||
|
||||
def get_token_by_user_id(self, user_id:UUID)->RefreshTokensOut|None:
|
||||
with self.Session() as session:
|
||||
with self.Session() as session: # noqa: SIM117
|
||||
with session.begin():
|
||||
query=select(RefreshTokens).where(and_(RefreshTokens.user_id==user_id, not_(RefreshTokens.is_revoked)))
|
||||
response=session.scalars(query).one_or_none()
|
||||
@@ -19,7 +22,7 @@ class JwtCrudActions:
|
||||
return RefreshTokensOut.model_validate(response)
|
||||
|
||||
def get_token_by_id(self, id:UUID)->RefreshTokensOut|None:
|
||||
with self.Session() as session:
|
||||
with self.Session() as session: # noqa: SIM117
|
||||
with session.begin():
|
||||
query=select(RefreshTokens).where(RefreshTokens.id==id)
|
||||
response=session.scalars(query).one_or_none()
|
||||
@@ -28,14 +31,14 @@ class JwtCrudActions:
|
||||
return RefreshTokensOut.model_validate(response)
|
||||
|
||||
def create_token(self, data:dict)->None:
|
||||
with self.Session() as session:
|
||||
with self.Session() as session: # noqa: SIM117
|
||||
with session.begin():
|
||||
new_token=RefreshTokens(**data)
|
||||
response=session.add(new_token)
|
||||
return response
|
||||
|
||||
def update_token(self, old_jti:UUID, new_jti:UUID)->bool:
|
||||
with self.Session() as session:
|
||||
with self.Session() as session: # noqa: SIM117
|
||||
with session.begin():
|
||||
query=select(RefreshTokens).where(RefreshTokens.id==old_jti)
|
||||
response=session.scalars(query).one()
|
||||
@@ -44,7 +47,7 @@ class JwtCrudActions:
|
||||
return True
|
||||
|
||||
def revoke_all(self, user_id:UUID)->bool:
|
||||
with self.Session() as session:
|
||||
with self.Session() as session: # noqa: SIM117
|
||||
with session.begin():
|
||||
query=select(RefreshTokens).where(RefreshTokens.user_id==user_id)
|
||||
response=session.scalars(query).all()
|
||||
@@ -53,7 +56,7 @@ class JwtCrudActions:
|
||||
return True
|
||||
|
||||
def logout(self,id:UUID)->bool:
|
||||
with self.Session() as session:
|
||||
with self.Session() as session: # noqa: SIM117
|
||||
with session.begin():
|
||||
query=select(RefreshTokens).where(RefreshTokens.id == id)
|
||||
response=session.scalars(query).one_or_none()
|
||||
|
||||
@@ -1,14 +1,18 @@
|
||||
from uuid import UUID
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from src.models.database_models.model import User, engine
|
||||
from src.models.pydantic_models.model import UserOutDB
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from uuid import UUID
|
||||
|
||||
|
||||
class UsersCrudActions:
|
||||
def __init__(self) -> None:
|
||||
self.Session=sessionmaker(bind=engine)
|
||||
|
||||
def get_user_by_email(self, email:str)->UserOutDB|None:
|
||||
with self.Session() as session:
|
||||
with self.Session() as session: # noqa: SIM117
|
||||
with session.begin():
|
||||
query=select(User).where(User.email==email)
|
||||
response=session.scalars(query).one_or_none()
|
||||
@@ -17,7 +21,7 @@ class UsersCrudActions:
|
||||
return UserOutDB.model_validate(response)
|
||||
|
||||
def get_user_by_id(self, id:UUID)->UserOutDB|None:
|
||||
with self.Session() as session:
|
||||
with self.Session() as session: # noqa: SIM117
|
||||
with session.begin():
|
||||
query=select(User).where(User.id==id)
|
||||
response=session.scalars(query).one_or_none()
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
class Base(BaseSettings):
|
||||
pass
|
||||
|
||||
|
||||
@@ -1,7 +1,10 @@
|
||||
from src.models.database_models.model import Model
|
||||
from uuid import UUID, uuid1
|
||||
|
||||
from sqlalchemy import String
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
from uuid import UUID, uuid1
|
||||
|
||||
from src.models.database_models.model import Model
|
||||
|
||||
|
||||
class AccountantSettings(Model):
|
||||
__tablename__="accountant_settings"
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
from datetime import datetime
|
||||
from sqlalchemy import TIMESTAMP, String, func, ForeignKey
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
from src.models.database_models.model import Model
|
||||
from uuid import UUID, uuid1
|
||||
|
||||
from sqlalchemy import TIMESTAMP, ForeignKey, String, func
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.models.database_models.model import Model
|
||||
|
||||
|
||||
class Stored(Model):
|
||||
__tablename__="reports"
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
from decimal import Decimal
|
||||
from src.models.database_models.model import Model
|
||||
|
||||
from sqlalchemy import Boolean, ForeignKey, Integer, Numeric, String
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from src.models.database_models.model import Model
|
||||
|
||||
|
||||
class Nomenclature(Model):
|
||||
__tablename__="goods"
|
||||
|
||||
|
||||
@@ -1,7 +1,20 @@
|
||||
from sqlalchemy import TIMESTAMP, Table, create_engine, String, Boolean, MetaData, Column, ForeignKey, func, Uuid
|
||||
from sqlalchemy.orm import Mapped, mapped_column, DeclarativeBase, relationship
|
||||
from uuid import UUID, uuid4
|
||||
from datetime import datetime
|
||||
from uuid import UUID, uuid4
|
||||
|
||||
from sqlalchemy import (
|
||||
TIMESTAMP,
|
||||
Boolean,
|
||||
Column,
|
||||
ForeignKey,
|
||||
MetaData,
|
||||
String,
|
||||
Table,
|
||||
Uuid,
|
||||
create_engine,
|
||||
func,
|
||||
)
|
||||
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, relationship
|
||||
|
||||
engine = create_engine("sqlite:///DB/database.db", echo=True)
|
||||
|
||||
'''remember as a boilerplate, or just cp/pst'''
|
||||
|
||||
@@ -1,7 +1,11 @@
|
||||
from typing import Annotated
|
||||
from src.models.pydantic_models.model import Base
|
||||
from pydantic import Field
|
||||
from uuid import UUID
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
from src.models.pydantic_models.model import Base
|
||||
|
||||
|
||||
class AccountantCreate(Base):
|
||||
|
||||
name:Annotated[str, Field(...,max_length=64,description="name of the accountant setting")]
|
||||
|
||||
@@ -1,9 +1,12 @@
|
||||
from datetime import datetime
|
||||
from typing import Annotated
|
||||
from uuid import UUID
|
||||
from src.models.pydantic_models.model import Base
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
from src.models.pydantic_models.model import Base
|
||||
|
||||
|
||||
class ReportCreate(Base):
|
||||
filename:Annotated[str,Field(..., min_length=2, max_length=255, description="name of the report")]
|
||||
doc_date:Annotated[datetime, Field(..., description="ts of the report")]
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
from decimal import Decimal
|
||||
from src.models.pydantic_models.model import Base
|
||||
from pydantic import Field
|
||||
from typing import Annotated
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
from src.models.pydantic_models.model import Base
|
||||
|
||||
|
||||
class NomenclatureCreate(Base):
|
||||
|
||||
article:Annotated[str, Field(...,min_length=5,max_length=16, description="name of the article")]
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
from datetime import datetime
|
||||
from typing import Annotated
|
||||
from pydantic import BaseModel, EmailStr, Field
|
||||
from uuid import UUID
|
||||
|
||||
from pydantic import BaseModel, EmailStr, Field
|
||||
|
||||
|
||||
class Base(BaseModel):
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
@@ -1,13 +1,15 @@
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from uuid import UUID
|
||||
|
||||
from fastapi import Request
|
||||
from .jwt import Jwt, Hashes
|
||||
from src.database.users.crud import UsersCrudActions
|
||||
|
||||
from src.database.auth.refresh_tokens import JwtCrudActions
|
||||
from src.database.users.crud import UsersCrudActions
|
||||
from src.errors.http_errors.errors import Errors
|
||||
from src.models.pydantic_models.model import RefreshTokensCreate, UserOut
|
||||
from src.models.configs_read.env import env_settings
|
||||
from src.models.pydantic_models.model import RefreshTokensCreate, UserOut
|
||||
|
||||
from .jwt import Hashes, Jwt
|
||||
|
||||
|
||||
class CurrentUser:
|
||||
@@ -78,7 +80,7 @@ class CurrentUser:
|
||||
token_hash=self.hash.token_to_hash(token),
|
||||
device_info=request.headers.get("user-agent", "unknown"),
|
||||
ip_address=request.headers.get("x-forwarded-for", "").split(",")[0].strip() or (request.client.host if request.client else "unknown"),
|
||||
expires_at=datetime.now(timezone.utc)+timedelta(days=env_settings.REFRESH_TOKEN_EXPIRE_DAYS)
|
||||
expires_at=datetime.now(UTC)+timedelta(days=env_settings.REFRESH_TOKEN_EXPIRE_DAYS)
|
||||
)
|
||||
self.jwt_db_actions.create_token(RefreshTokensCreate.model_dump(token_record))
|
||||
|
||||
@@ -113,8 +115,8 @@ class CurrentUser:
|
||||
'''sqlite constraints about timezone'''
|
||||
expires_at=old_record.expires_at
|
||||
if expires_at.tzinfo is None:
|
||||
expires_at = expires_at.replace(tzinfo=timezone.utc)
|
||||
if expires_at<datetime.now(timezone.utc):
|
||||
expires_at = expires_at.replace(tzinfo=UTC)
|
||||
if expires_at<datetime.now(UTC):
|
||||
raise self.error.credentials_error(detail="Token expired")
|
||||
|
||||
'''user check'''
|
||||
@@ -142,7 +144,7 @@ class CurrentUser:
|
||||
token_hash=self.hash.token_to_hash(new_refresh_token),
|
||||
device_info=request.headers.get("user-agent", "unknown"),
|
||||
ip_address=request.headers.get("x-forwarded-for", "").split(",")[0].strip() or (request.client.host if request.client else "unknown"),
|
||||
expires_at=datetime.now(timezone.utc)+timedelta(days=env_settings.REFRESH_TOKEN_EXPIRE_DAYS),
|
||||
expires_at=datetime.now(UTC)+timedelta(days=env_settings.REFRESH_TOKEN_EXPIRE_DAYS),
|
||||
)
|
||||
self.jwt_db_actions.create_token(RefreshTokensCreate.model_dump(new_token_record))
|
||||
|
||||
|
||||
@@ -1,10 +1,12 @@
|
||||
from jose import JWTError, jwt
|
||||
import bcrypt
|
||||
from src.errors.http_errors.errors import Errors
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from src.models.configs_read.env import env_settings
|
||||
from uuid import uuid4
|
||||
import hashlib
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from uuid import uuid4
|
||||
|
||||
import bcrypt
|
||||
from jose import JWTError, jwt
|
||||
|
||||
from src.errors.http_errors.errors import Errors
|
||||
from src.models.configs_read.env import env_settings
|
||||
|
||||
'''Hash/Check hash'''
|
||||
class Hashes:
|
||||
@@ -32,7 +34,7 @@ class Jwt:
|
||||
def create_access_token(self, data:dict)->str:
|
||||
|
||||
user_info=data.copy()
|
||||
user_info.update({"exp": datetime.now(timezone.utc)+timedelta(minutes=env_settings.ACCESS_TOKEN_EXPIRE_MINUTES),
|
||||
user_info.update({"exp": datetime.now(UTC)+timedelta(minutes=env_settings.ACCESS_TOKEN_EXPIRE_MINUTES),
|
||||
"token_type":"access"})
|
||||
return jwt.encode(user_info, env_settings.SECRET_KEY, env_settings.ALGORITHM)
|
||||
|
||||
@@ -41,7 +43,7 @@ class Jwt:
|
||||
|
||||
user_info=data.copy()
|
||||
jti=str(uuid4())
|
||||
user_info.update({"exp":datetime.now(timezone.utc)+timedelta(days=env_settings.REFRESH_TOKEN_EXPIRE_DAYS),
|
||||
user_info.update({"exp":datetime.now(UTC)+timedelta(days=env_settings.REFRESH_TOKEN_EXPIRE_DAYS),
|
||||
"token_type":"refresh",
|
||||
"jti":jti
|
||||
})
|
||||
@@ -53,9 +55,9 @@ class Jwt:
|
||||
def jwt_decode(self, token:str)->dict:
|
||||
|
||||
try:
|
||||
payload=jwt.decode(token, env_settings.SECRET_KEY, algorithms=[env_settings.ALGORITHM])
|
||||
payload=jwt.decode(token, env_settings.SECRET_KEY, algorithms=[env_settings.ALGORITHM], options={"require_exp": True} )
|
||||
|
||||
if (payload.get("sub")) is None:
|
||||
if not (payload.get("sub")) or not (payload.get("token_type")):
|
||||
raise self.error.credentials_error(detail="Jwt token is incorrect")
|
||||
|
||||
except JWTError as e:
|
||||
|
||||
@@ -1,15 +1,15 @@
|
||||
from fastapi import APIRouter, Depends, Request, Response, Cookie
|
||||
from fastapi.security import OAuth2PasswordRequestForm, OAuth2PasswordBearer
|
||||
from fastapi import APIRouter, Cookie, Depends, Request, Response
|
||||
from fastapi.security import OAuth2PasswordBearer, OAuth2PasswordRequestForm
|
||||
|
||||
from src.models.configs_read.env import env_settings
|
||||
from src.models.pydantic_models.model import UserOut
|
||||
|
||||
from src.service.auth.auth import auth
|
||||
|
||||
router=APIRouter(prefix="/protected")
|
||||
oauth2_schema=OAuth2PasswordBearer(tokenUrl="/protected/token", refreshUrl="/protected/refresh")
|
||||
|
||||
@router.post("/token")
|
||||
async def get_access_token(request: Request,response:Response, form_data:OAuth2PasswordRequestForm=Depends())->dict:
|
||||
async def get_access_token(request: Request,response:Response, form_data:OAuth2PasswordRequestForm=Depends())->dict: # noqa: B008
|
||||
|
||||
access_token, refresh_token=auth.login(form_data_email=form_data.username, form_data_password=form_data.password, request=request)
|
||||
|
||||
@@ -46,10 +46,10 @@ async def get_current_user(token:str = Depends(oauth2_schema)) -> UserOut:
|
||||
|
||||
|
||||
@router.get("/logout")
|
||||
async def logout(response:Response,refresh_token: str = Cookie(),current_user:UserOut=Depends(get_current_user))->bool:
|
||||
async def logout(response:Response,refresh_token: str = Cookie(),current_user:UserOut=Depends(get_current_user))->bool: # noqa: B008
|
||||
response.delete_cookie("refresh_token")
|
||||
return auth.logout(refresh_token)
|
||||
|
||||
@router.get("")
|
||||
async def protected(current_user:UserOut=Depends(get_current_user))->dict:
|
||||
async def protected(current_user:UserOut=Depends(get_current_user))->dict: # noqa: B008
|
||||
return {"protected router": "Hello, this is a protected router"}
|
||||
|
||||
Reference in New Issue
Block a user