INIT Commit

This commit is contained in:
Michal Kváček
2023-06-05 09:55:17 +02:00
commit 014a999055
46 changed files with 1645 additions and 0 deletions
View File
+6
View File
@@ -0,0 +1,6 @@
import sys
sys.path.insert(0, "/app/src")
from .main import App # noqa
app = App().create_app()
+3
View File
@@ -0,0 +1,3 @@
APP_DEBUG = True
GRAPHIQL = True
APP_SECRET_KEY = "test"
+6
View File
@@ -0,0 +1,6 @@
__all__ = [
'async_session',
'engine'
]
from database.config import async_session, engine
+17
View File
@@ -0,0 +1,17 @@
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession
from sqlalchemy.orm import sessionmaker
def create_db_engine(
user: str = "root",
password: str = "",
host: str = "db",
port: int = 3306,
database: str = "ull_tracker",
):
database_url = f'mysql+aiomysql://{user}:{password}@{host}:{port}/{database}?charset=utf8'
return create_async_engine(database_url, future=True, echo=True)
engine = create_db_engine()
async_session = sessionmaker(engine, expire_on_commit=False, class_=AsyncSession)
+13
View File
@@ -0,0 +1,13 @@
from sqlalchemy import func
from sqlalchemy.types import UserDefinedType
class Point(UserDefinedType):
def get_col_spec(self):
return 'POINT'
def bind_expression(self, bindvalue):
return func.ST_GeomFromText(bindvalue, type_=self)
def column_expression(self, col):
return func.ST_AsText(col, type_=self)
+240
View File
@@ -0,0 +1,240 @@
import datetime
from typing import Set
from sqlalchemy import String, DateTime, ForeignKey, Text, Integer, func, Table, Column, Boolean, select
from sqlalchemy.orm import Mapped, relationship, as_declarative, mapped_column
from database.custom_types import Point
from sqlalchemy.ext.asyncio import AsyncSession
@as_declarative()
class BaseModel:
excluded_columns_in_dict = tuple()
def as_dict(self):
return {
c.name: getattr(self, c.name)
for c in self.__table__.columns
if c.name not in self.excluded_columns_in_dict
}
@classmethod
async def create(cls, db_session: AsyncSession, data: dict):
model = cls(**data)
db_session.add(model)
await db_session.commit()
return model
@classmethod
async def update(cls, db_session: AsyncSession, id: int, data: dict):
obj = (await db_session.scalars(select(cls).filter_by(id=id))).one()
for key, value in data.items():
if getattr(obj, key) != value:
setattr(obj, key, value)
await db_session.commit()
return obj
# TODO: doplnit GPX k letu, pocasi k letu (podle lokality, mozna do FlightTrack)
user_is_in_organization = Table(
"user_is_in_organization",
BaseModel.metadata,
Column("user_id", Integer, ForeignKey("user.id"), primary_key=True),
Column("organization_id", Integer, ForeignKey("organization.id"), primary_key=True)
)
class Airport(BaseModel):
__tablename__ = "airport"
id: Mapped[int] = mapped_column(primary_key=True)
name: Mapped[str] = mapped_column(String(128), nullable=False)
icao_code: Mapped[str] = mapped_column(String(4), nullable=False)
gps_position: Mapped[Point] = mapped_column(Point, nullable=True)
elevation: Mapped[int] = mapped_column(Integer, nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
metars: Mapped['Metar'] = relationship(back_populates="airport")
class PointOfInterestType(BaseModel):
__tablename__ = "point_of_interest_type"
id: Mapped[int] = mapped_column(primary_key=True)
name: Mapped[str] = mapped_column(String(128), nullable=False)
is_public: Mapped[bool] = mapped_column(Boolean, server_default='0')
created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id'))
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
created_by: Mapped['User'] = relationship()
class PointOfInterest(BaseModel):
__tablename__ = "point_of_interest"
id: Mapped[int] = mapped_column(primary_key=True)
name: Mapped[str] = mapped_column(String(128), nullable=False)
gps_position: Mapped[Point] = mapped_column(Point, nullable=True)
type_id: Mapped[int] = mapped_column(Integer, ForeignKey("point_of_interest_type.id"))
is_public: Mapped[bool] = mapped_column(Boolean, server_default='0')
created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id'))
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
type: Mapped[PointOfInterestType] = relationship()
created_by: Mapped['User'] = relationship()
class Photo(BaseModel):
__tablename__ = "photo"
id: Mapped[int] = mapped_column(primary_key=True)
name: Mapped[str] = mapped_column(String(128), nullable=False)
filename: Mapped[str] = mapped_column(String(128), nullable=False)
description: Mapped[str] = mapped_column(Text, nullable=False)
gps_position: Mapped[Point] = mapped_column(Point, nullable=True)
flight_id: Mapped[int] = mapped_column(Integer, ForeignKey("flight.id"), nullable=False)
created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id'))
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
flight: Mapped['Flight'] = relationship(back_populates="photos")
created_by: Mapped['User'] = relationship()
class Aircraft(BaseModel):
__tablename__ = "aircraft"
id: Mapped[int] = mapped_column(primary_key=True)
name: Mapped[str] = mapped_column(String(128), nullable=False)
photo_filename: Mapped[str] = mapped_column(String(128), nullable=True)
manufacturer: Mapped[str] = mapped_column(Text, nullable=False, server_default="")
model: Mapped[str] = mapped_column(String(30), nullable=False)
description: Mapped[str] = mapped_column(Text, nullable=False)
organization_id: Mapped[int] = mapped_column(Integer, ForeignKey('organization.id'), nullable=True)
created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id'))
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
deleted: Mapped[bool] = mapped_column(Boolean, nullable=False, server_default='0')
organization: Mapped['Organization'] = relationship()
flights: Mapped[Set['Flight']] = relationship()
created_by: Mapped['User'] = relationship()
notes: Mapped['AircraftNotes'] = relationship()
class AircraftNotes(BaseModel):
__tablename__ = "aircraft_notes"
id: Mapped[int] = mapped_column(primary_key=True)
aircraft_id: Mapped[int] = mapped_column(Integer, ForeignKey("aircraft.id"), nullable=False)
name: Mapped[str] = mapped_column(String(128), nullable=False)
description: Mapped[str] = mapped_column(Text, nullable=False)
is_public: Mapped[bool] = mapped_column(Boolean, server_default='0')
created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id'))
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
created_by: Mapped['User'] = relationship()
aircraft: Mapped['Aircraft'] = relationship(back_populates="notes")
class Organization(BaseModel):
__tablename__ = "organization"
id: Mapped[int] = mapped_column(primary_key=True)
name: Mapped[str] = mapped_column(String(128), nullable=False)
created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id'))
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
users: Mapped[Set['User']] = relationship(back_populates='organizations', secondary=user_is_in_organization)
created_by: Mapped['User'] = relationship()
class FlightTrack(BaseModel):
__tablename__ = "flight_track"
id: Mapped[int] = mapped_column(primary_key=True)
flight_id: Mapped[int] = mapped_column(Integer, ForeignKey("flight.id"), nullable=False)
poi_id: Mapped[int] = mapped_column(Integer, ForeignKey("point_of_interest.id"), nullable=False)
order: Mapped[int] = mapped_column(Integer)
flight: Mapped['Flight'] = relationship()
point_of_interest: Mapped['PointOfInterest'] = relationship()
class Flight(BaseModel):
__tablename__ = "flight"
id: Mapped[int] = mapped_column(primary_key=True)
name: Mapped[str] = mapped_column(String(128), nullable=False)
description: Mapped[str] = mapped_column(Text, nullable=False)
takeoff_airport_id: Mapped[int] = mapped_column(Integer, ForeignKey("airport.id"), nullable=True)
landing_airport_id: Mapped[int] = mapped_column(Integer, ForeignKey("airport.id"), nullable=True)
aircraft_id: Mapped[int] = mapped_column(Integer, ForeignKey('aircraft.id'))
copilot_id: Mapped[int] = mapped_column(Integer, ForeignKey('copilot.id'), nullable=True)
created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id'))
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
duration_total: Mapped[int] = mapped_column(Integer, nullable=False)
duration_pic: Mapped[int] = mapped_column(Integer, nullable=False)
takeoff_airport: Mapped['Airport'] = relationship(foreign_keys=[takeoff_airport_id])
landing_airport: Mapped['Airport'] = relationship(foreign_keys=[landing_airport_id])
copilot: Mapped['Copilot'] = relationship(back_populates="flights")
aircraft: Mapped['Aircraft'] = relationship(back_populates="flights")
photos: Mapped[Set['Photo']] = relationship()
user: Mapped['User'] = relationship(back_populates="flights")
created_by: Mapped['User'] = relationship()
# flight_track: Mapped[List['PointOfInterest']] = relationship(secondary=FlightTrack)
class Copilot(BaseModel):
__tablename__ = "copilot"
id: Mapped[int] = mapped_column(primary_key=True)
name: Mapped[str] = mapped_column(String(128), nullable=False)
created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id'))
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
flights: Mapped[Set['Flight']] = relationship(back_populates="copilot")
created_by: Mapped['User'] = relationship()
class Metar(BaseModel):
__tablename__ = "metar"
id: Mapped[int] = mapped_column(primary_key=True)
airport_id: Mapped[int] = mapped_column(Integer, ForeignKey('airport.id'))
metar: Mapped[str] = mapped_column(Text, nullable=False)
issued_at: Mapped[datetime] = mapped_column(DateTime)
airport: Mapped['Airport'] = relationship(back_populates="metars")
class License(BaseModel):
__tablename__ = "license"
id: Mapped[int] = mapped_column(primary_key=True)
name: Mapped[str] = mapped_column(String(128), nullable=False)
number: Mapped[str] = mapped_column(String(30), nullable=False)
valid_until: Mapped[datetime] = mapped_column(DateTime, nullable=False)
created_by_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id'))
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
user: Mapped['User'] = relationship(back_populates="licences")
created_by: Mapped['User'] = relationship()
class User(BaseModel):
__tablename__ = "user"
excluded_columns_in_dict = ('password_hashed',)
id: Mapped[int] = mapped_column(primary_key=True)
email: Mapped[str] = mapped_column(String(128), nullable=False, unique=True)
name: Mapped[str] = mapped_column(String(128), nullable=False)
avatar_image_url: Mapped[str] = mapped_column(String(128), nullable=True)
password_hashed: Mapped[str] = mapped_column(String(60), nullable=False)
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
licences: Mapped[Set['License']] = relationship(back_populates="user")
flights: Mapped[Set['Flight']] = relationship(back_populates="user")
organizations: Mapped[Set['Organization']] = relationship(back_populates="users", secondary=user_is_in_organization)
View File
+7
View File
@@ -0,0 +1,7 @@
from database import async_session
async def db_session():
async with async_session() as session:
async with session.begin():
yield session
+10
View File
@@ -0,0 +1,10 @@
# from fastapi import Depends, Security
# from fastapi_jwt import JwtAuthorizationCredentials
#
#
# async def check_jwt_token():
# authorize.jwt_required()
#
# subject = authorize.get_jwt_subject()
#
# print("Subject: ", subject)
+3
View File
@@ -0,0 +1,3 @@
class BaseEndpoint:
def __init__(self, db):
self.db = db
+17
View File
@@ -0,0 +1,17 @@
from sqlalchemy import select, desc
from database.models import Flight
from endpoints.base import BaseEndpoint
class FlightsEndpoint(BaseEndpoint):
async def resolve(self):
self.db.add(Flight(name="test"))
await self.db.flush()
data = await self.db.execute(select(Flight).order_by(desc(Flight.id)))
model = data.scalars().first()
return {
"status": model
}
+60
View File
@@ -0,0 +1,60 @@
from database.models import User, Flight, Copilot, Airport, Aircraft
from endpoints.base import BaseEndpoint
class InitDataEndpoint(BaseEndpoint):
async def on_get(self):
user_ids = []
copilot_ids = []
flight_ids = []
users = [
User(avatar_image_url="", name="Karel Vomacka", email="a@test.cz", password_hashed="****"),
User(avatar_image_url="", name="Karel Novak", email="b@test.cz", password_hashed="****"),
User(avatar_image_url="", name="Franta Pavel", email="c@test.cz", password_hashed="****"),
]
airport = Airport(name="Letiste Letnany", icao_code="LKLT")
self.db.add(airport)
for user in users:
self.db.add(user)
await self.db.flush()
user_ids.append(user.id)
aircraft = Aircraft(name="OK-AUR28", type="Bristell NG5", description="", created_by=user)
self.db.add(aircraft)
copilot = None
if user.name == 'Franta Pavel':
copilot = Copilot(name="Copilot test", created_by_id=user.id)
self.db.add(copilot)
await self.db.flush()
copilot_ids.append(copilot.id)
flight = Flight(name="test flight", description="Testovaci popis", duration_total=65, duration_pic=65,
takeoff_airport=airport, landing_airport=airport, aircraft=aircraft, created_by_id=user.id,
copilot_id=copilot.id if copilot else None)
self.db.add(flight)
await self.db.flush()
flight_ids.append(flight.id)
return {
"user_ids": user_ids,
"copilot_ids": copilot_ids,
"flight_ids": flight_ids,
'airport_id': airport.id
}
#
#
# self.db.add(Flight(name="test"))
# await self.db.flush()
#
# data = await self.db.execute(select(Flight).order_by(desc(Flight.id)))
# model = data.scalars().first()
# return {
# "status": model
# }
+59
View File
@@ -0,0 +1,59 @@
from fastapi import HTTPException
from fastapi_jwt import JwtAuthorizationCredentials
from fastapi_jwt.jwt import JwtAccess, JwtRefresh
from passlib.hash import bcrypt
from sqlalchemy import select
from starlette.responses import Response
from database.models import User
from endpoints.base import BaseEndpoint
from pydantic import BaseModel
class LoginInput(BaseModel):
email: str
password: str
class LoginEndpoint(BaseEndpoint):
def __init__(self, db, access_token: JwtAccess, refresh_token: JwtRefresh):
super().__init__(db)
self.access_security = access_token
self.refresh_security = refresh_token
async def on_post(self, user_data: LoginInput, resp: Response) -> dict:
query = select(User).filter_by(email=user_data.email)
logged_user = (await self.db.scalars(query)).first()
if not logged_user:
raise HTTPException(status_code=401, detail="Invalid user")
if not bcrypt.verify(user_data.password, logged_user.password_hashed):
raise HTTPException(status_code=401, detail="Bad username or password")
subject = {"id": logged_user.id, "email": logged_user.email}
access_token = self.access_security.create_access_token(subject=subject)
refresh_token = self.refresh_security.create_refresh_token(subject=subject)
self.access_security.set_access_cookie(resp, access_token)
self.refresh_security.set_refresh_cookie(resp, refresh_token)
return {
"user": logged_user.as_dict(),
"access_token": access_token,
"refresh_token": refresh_token
}
class MeEndpoint(BaseEndpoint):
def __init__(self, db, credentials: JwtAuthorizationCredentials):
super().__init__(db)
self.credentials = credentials
async def on_get(self) -> dict:
query = select(User).filter_by(id=self.credentials['id'])
user = (await self.db.scalars(query)).first()
return user.as_dict()
+45
View File
@@ -0,0 +1,45 @@
import re
from fastapi import HTTPException
from sqlalchemy import select
from typing import Optional
from pydantic import BaseModel, root_validator, Field
from database.models import User
from endpoints.base import BaseEndpoint
from passlib.hash import bcrypt
class RegistrationInput(BaseModel):
email: str = Field(..., min_length=4)
name: Optional[str]
password: str
@root_validator()
def validate_email(cls, values):
email = values.get("email") or ""
if email and not re.match(r"(.+)@(.+)\..{2,6}", email):
raise ValueError("Specified e-mail is not valid!")
return values
class RegistrationEndpoint(BaseEndpoint):
async def on_post(self, user_data: RegistrationInput) -> User:
query = select(User).filter_by(email=user_data.email)
existing_user = (await self.db.scalars(query)).first()
if existing_user:
raise HTTPException(status_code=422, detail="User already exists")
model = User(
name=user_data.name,
email=user_data.email,
password_hashed=bcrypt.hash(user_data.password)
)
self.db.add(model)
await self.db.commit()
return model.as_dict()
View File
@@ -0,0 +1,5 @@
__all__ = [
'copilots_dataloader'
]
from graphql_schema.dataloaders.copilots import copilots_dataloader
@@ -0,0 +1,16 @@
from typing import List
from sqlalchemy import select
from strawberry.dataloader import DataLoader
from database import async_session
from database.models import Aircraft
async def load(ids: List[int]):
async with async_session() as session:
models = (await session.scalars(select(Aircraft).filter(Aircraft.id.in_(ids)))).all()
models_by_id = {model.id: model for model in models}
return [models_by_id.get(id_) for id_ in ids]
aircraft_dataloader = DataLoader(load_fn=load)
@@ -0,0 +1,16 @@
from typing import List
from sqlalchemy import select
from strawberry.dataloader import DataLoader
from database import async_session
from database.models import Copilot
async def load(ids: List[int]):
async with async_session() as session:
models = (await session.scalars(select(Copilot).filter(Copilot.id.in_(ids)))).all()
models_by_id = {model.id: model for model in models}
return [models_by_id.get(id_) for id_ in ids]
copilots_dataloader = DataLoader(load_fn=load)
+121
View File
@@ -0,0 +1,121 @@
import os
import uuid
from typing import List, Optional
import strawberry
from strawberry.file_uploads import Upload
from sqlalchemy import select
from database import models
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input
@strawberry_sqlalchemy_type(models.Aircraft)
class Aircraft:
photo_url: Optional[str] = strawberry.field(
resolver=lambda root: f"http://localhost:8000/uploads/{root.photo_filename}" if root.photo_filename else None
)
def get_base_query(user_id: int):
return (
select(models.Aircraft)
.filter(models.Aircraft.created_by_id == user_id)
.filter(models.Aircraft.deleted.is_(False))
)
@strawberry.type
class AircraftQueries:
@strawberry.field
async def aircrafts(root, info) -> List[Aircraft]:
query = (
get_base_query(info.context.user_id)
.order_by(models.Aircraft.id.desc())
)
return (await info.context.db.scalars(query)).all()
@strawberry.field
async def aircraft(root, info, id: int) -> Aircraft:
query = (
get_base_query(info.context.user_id)
.filter(models.Aircraft.id == id)
)
return (await info.context.db.scalars(query)).one()
def check_directories(path: str):
if not os.path.isdir(path):
os.makedirs(path)
async def handle_img_upload(file: Upload, path: str, filename: str):
check_directories(path)
content = await file.read()
image = open(path + "/" + filename, "wb")
image.write(content)
image.close()
@strawberry.type
class CreateAircraftMutation:
@strawberry_sqlalchemy_input(models.Aircraft, exclude_fields=['id', 'photo_filename'])
class CreateAircraftInput:
photo: Optional[Upload]
@strawberry.mutation
async def create_aircraft(root, info, input: CreateAircraftInput) -> Aircraft:
# TODO: kontrola organizace
filename = None
if input.photo:
dest_path = "/app/uploads/aircrafts/"
filename = f"{uuid.uuid4()}-{input.photo.filename}"
await handle_img_upload(input.photo, dest_path, filename=filename)
return await models.Aircraft.create(
info.context.db,
data=dict(
name=input.name,
description=input.description,
model=input.model,
manufacturer=input.manufacturer,
photo_filename=filename,
organization_id=input.organization_id,
created_by_id=info.context.user_id,
)
)
@strawberry.type
class EditAircraftMutation:
@strawberry_sqlalchemy_input(models.Aircraft, exclude_fields=['photo_filename'])
class EditAircraftInput:
photo: Optional[Upload]
@strawberry.mutation
async def edit_aircraft(root, info, id: int, input: EditAircraftInput) -> Aircraft:
# TODO: kontrola organizace
# TODO: kontrola opravneni na akci
return await models.Aircraft.update(
info.context.db,
id,
data=dict(
name=input.name,
description=input.description,
model=input.model,
manufacturer=input.manufacturer,
organization_id=input.organization_id,
)
)
@strawberry.type
class DeleteAircraftMutation:
@strawberry.mutation
async def delete_aircraft(self, info, id: int) -> Aircraft:
# TODO: kontrola opravneni na akci
return await models.Aircraft.update(info.context.db, id, data=dict(deleted=True))
+17
View File
@@ -0,0 +1,17 @@
from typing import List
import strawberry
from sqlalchemy import select
from database.models import Copilot
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type
@strawberry_sqlalchemy_type(Copilot)
class CopilotType:
pass
@strawberry.type
class CopilotQueries:
@strawberry.field
async def pilots(root, info) -> List[CopilotType]:
return (await info.context['db'].scalars(select(Copilot))).all()
+67
View File
@@ -0,0 +1,67 @@
from typing import List, Optional
import strawberry
from sqlalchemy import select
from database import models
from graphql_schema.dataloaders import copilots_dataloader
from graphql_schema.dataloaders.aircraft import aircraft_dataloader
from graphql_schema.entities.aircraft import Aircraft
from graphql_schema.entities.copilot import CopilotType
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type, strawberry_sqlalchemy_input
# Bude se hodit: https://strawberry.rocks/docs/types/lazy
@strawberry_sqlalchemy_type(models.Flight)
class Flight:
async def load_aircraft(root):
return await aircraft_dataloader.load(root.aircraft_id)
async def load_copilot(root):
return await copilots_dataloader.load(root.copilot_id)
copilot: Optional[CopilotType] = strawberry.field(resolver=load_copilot)
aircraft: Aircraft = strawberry.field(resolver=load_aircraft)
@strawberry.type
class FlightQueries:
@strawberry.input
class FlightFilters:
takeoff: Optional[int]
@strawberry.field
async def flights(root, info, filters: Optional[FlightFilters] = None) -> List[Flight]:
query = (
select(models.Flight)
.filter(models.Flight.created_by_id == info.context.user_id)
.order_by(models.Flight.id) # TODO: desc
)
return (await info.context.db.scalars(query)).all()
@strawberry.field
async def flight(root, info, id: int) -> Flight:
query = (
select(models.Flight)
.filter(models.Flight.id == id)
.filter(models.Flight.created_by_id == info.context.user_id)
)
return (await info.context.db.scalars(query)).fetch_one()
@strawberry.type
class CreateFlightMutation:
@strawberry_sqlalchemy_input(models.Flight, all_optional=True)
class FlightInput:
pass
@strawberry.mutation
async def create_flight(self, info, input_: FlightInput) -> Flight:
model = models.Flight(name=input_.name)
db = info.context.db
db.add(model)
await db.commit()
return Flight(model)
+25
View File
@@ -0,0 +1,25 @@
import strawberry
from sqlalchemy import select
from database.models import User
from graphql_schema.sqlalchemy_to_strawberry_type import strawberry_sqlalchemy_type
@strawberry_sqlalchemy_type(User, exclude_fields=['password_hashed'])
class UserType:
pass
@strawberry.type
class LoginResultType:
logged_user = UserType
access_token: str
refresh_token: str
@strawberry.type
class UserQueries:
@strawberry.field
async def logged_user(root, info) -> UserType:
query = select(User).filter_by(id=info.context.user_id)
return (await info.context.db.scalars(query)).first()
+10
View File
@@ -0,0 +1,10 @@
from strawberry.tools import merge_types
from graphql_schema.entities.aircraft import CreateAircraftMutation, EditAircraftMutation, DeleteAircraftMutation
from graphql_schema.entities.flight import CreateFlightMutation
Mutation = merge_types("Mutation", (
CreateAircraftMutation,
EditAircraftMutation,
DeleteAircraftMutation,
CreateFlightMutation,
))
+15
View File
@@ -0,0 +1,15 @@
from strawberry.tools import merge_types
from .entities.aircraft import AircraftQueries
from .entities.copilot import CopilotQueries
from .entities.flight import FlightQueries
from .entities.user import UserQueries
# https://github.com/strawberry-graphql/examples/blob/main/fastapi-sqlalchemy/api/schema.py
Query = merge_types('Query', (
AircraftQueries,
FlightQueries,
CopilotQueries,
UserQueries
))
+42
View File
@@ -0,0 +1,42 @@
import dataclasses
import strawberry
from fastapi_jwt import JwtAuthorizationCredentials
from fastapi_jwt.jwt import JwtAccessBearerCookie
from sqlalchemy.ext.asyncio import AsyncSession
from strawberry.extensions import SchemaExtension
from strawberry.fastapi import BaseContext
from .mutation import Mutation
from .query import Query
# Toto se da kdyztak pouzit jako extension do Schema
# class SQLAlchemySession(Extension):
# def on_request_start(self):
# session = async_session()
# print(self.execution_context.context)
# self.execution_context.context["db"] = session
#
# async def on_request_end(self):
# await self.execution_context.context["db"].close()
class LoggingExtension(SchemaExtension):
def on_request_start(self):
print("request start")
async def on_request_end(self):
print("request end")
@dataclasses.dataclass
class GraphQLContext(BaseContext):
db: AsyncSession
user_id: int
jwt_auth_credentials: JwtAuthorizationCredentials
jwt: JwtAccessBearerCookie
schema = strawberry.Schema(
query=Query,
mutation=Mutation,
extensions=[LoggingExtension]
)
@@ -0,0 +1,54 @@
import typing
from typing import List, Optional
import strawberry
from sqlalchemy import inspect
from database.models import BaseModel
def get_annotations_for_scalars(model: BaseModel, exclude_fields=None, force_optional: bool = False):
if exclude_fields is None:
exclude_fields = []
annotations_ = {}
for name, column in inspect(model).columns.items():
is_optional = column.nullable or force_optional
if name in exclude_fields:
continue
annotations_[name] = column.type.python_type if not is_optional else typing.Optional[column.type.python_type]
return annotations_
def strawberry_sqlalchemy_type(model, exclude_fields: Optional[typing.Union[List, typing.Tuple]] = None):
if exclude_fields is None:
exclude_fields = []
def from_sqlalchemy_model(model: BaseModel):
return model
def wrapper(cls):
cls.__annotations__.update(get_annotations_for_scalars(model, exclude_fields=exclude_fields + ["deleted"]))
cls.from_sqlalchemy_model = from_sqlalchemy_model
return strawberry.type(cls)
return wrapper
def strawberry_sqlalchemy_input(
model,
exclude_fields: Optional[typing.Union[List, typing.Tuple]] = None,
all_optional: bool = False):
if exclude_fields is None:
exclude_fields = []
ignored_fields = ["created_at", "created_by_id", "updated_by_id", "updated_at", "deleted"]
def wrapper(cls):
cls.__annotations__.update(get_annotations_for_scalars(
model,
exclude_fields=exclude_fields + ignored_fields,
force_optional=all_optional
))
return strawberry.input(cls)
return wrapper
+127
View File
@@ -0,0 +1,127 @@
from datetime import timedelta
from fastapi import FastAPI, APIRouter, Depends, Security
from fastapi_jwt import JwtAuthorizationCredentials, JwtAccessBearerCookie, JwtRefreshBearerCookie
from sqlalchemy.ext.asyncio import AsyncSession
from starlette.middleware.cors import CORSMiddleware
from starlette.responses import RedirectResponse, Response
from strawberry.fastapi import GraphQLRouter
from config import APP_SECRET_KEY, GRAPHIQL, APP_DEBUG
from dependencies.db import db_session
from endpoints.init_data import InitDataEndpoint
from endpoints.flights import FlightsEndpoint
from endpoints.login import LoginEndpoint, LoginInput, MeEndpoint
from endpoints.registration import RegistrationInput, RegistrationEndpoint
from graphql_schema.schema import schema, GraphQLContext
class App:
api_router = APIRouter(dependencies=[])
access_security = JwtAccessBearerCookie(
secret_key=APP_SECRET_KEY,
auto_error=False,
access_expires_delta=timedelta(hours=1)
)
refresh_security = JwtRefreshBearerCookie(
secret_key=APP_SECRET_KEY,
auto_error=True
)
def create_app(self):
app = FastAPI()
self.setup_exception_handlers(app)
self.setup_middleware(app)
self.setup_routes(app)
return app
@staticmethod
def setup_exception_handlers(app: FastAPI):
pass
@staticmethod
def setup_middleware(app: FastAPI):
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
def setup_graphql_endpoint(self, app: FastAPI):
if APP_DEBUG:
@self.api_router.get("/graphql/autologin")
async def autologin():
access_token = self.access_security.create_access_token(subject={"id": 8, "name": "Franta Vomacka"})
response = RedirectResponse(url="/graphql")
self.access_security.set_access_cookie(response, access_token)
return response
def setup_graphql_context(
credentials: JwtAuthorizationCredentials = Security(self.access_security),
db: AsyncSession = Depends(db_session),
):
return GraphQLContext(
jwt_auth_credentials=credentials,
user_id=credentials['id'],
db=db,
jwt=self.access_security
)
graphql_app = GraphQLRouter(
schema,
graphiql=GRAPHIQL,
debug=APP_DEBUG,
context_getter=setup_graphql_context
)
app.include_router(graphql_app, prefix="/graphql")
def setup_routes(self, app: FastAPI):
# protected endpoints
@self.api_router.get("/me")
async def me(
db: AsyncSession = Depends(db_session),
credentials: JwtAuthorizationCredentials = Security(self.access_security)
):
return await MeEndpoint(db, credentials).on_get()
# @app.post("/refresh")
# def refresh(
# credentials: JwtAuthorizationCredentials = Security(refresh_security)
# ):
# # Update access/refresh tokens pair
# # We can customize expires_delta when creating
# access_token = access_security.create_access_token(subject=credentials.subject)
# refresh_token = refresh_security.create_refresh_token(subject=credentials.subject,
# expires_delta=timedelta(days=2))
#
# return {"access_token": access_token, "refresh_token": refresh_token}
self.setup_graphql_endpoint(app)
# public endpoints
@self.api_router.post("/login")
async def login(resp: Response, user: LoginInput, db: AsyncSession = Depends(db_session)):
return await LoginEndpoint(db, self.access_security, self.refresh_security).on_post(user, resp)
@self.api_router.post("/registration", status_code=201)
async def registration(user: RegistrationInput, db: AsyncSession = Depends(db_session)):
return await RegistrationEndpoint(db).on_post(user)
# testing endpoint
@self.api_router.get("/init-data")
async def init_data(db: AsyncSession = Depends(db_session)):
return await InitDataEndpoint(db).on_get()
@self.api_router.get("/flights", status_code=200)
async def flights(db: AsyncSession = Depends(db_session)):
return await FlightsEndpoint(db).resolve()
# musi byt na konci
app.include_router(self.api_router)