Fixed formatting and type definitions
This commit is contained in:
parent
bbe0df9546
commit
d97335e0c7
5 changed files with 112 additions and 65 deletions
|
|
@ -1,27 +1,38 @@
|
|||
from sqlmodel import SQLModel, create_engine, Session, select
|
||||
from sqlmodel import Session, SQLModel, create_engine, select
|
||||
|
||||
from app.models import example_model, project
|
||||
from app.config import settings
|
||||
|
||||
engine = create_engine(
|
||||
settings.database_url, echo=True,
|
||||
settings.database_url,
|
||||
echo=True,
|
||||
connect_args={"check_same_thread": False},
|
||||
)
|
||||
|
||||
|
||||
def create_database():
|
||||
SQLModel.metadata.create_all(engine)
|
||||
|
||||
def add_to_database(input: any):
|
||||
|
||||
|
||||
def add_to_database(input: SQLModel):
|
||||
with Session(engine) as session:
|
||||
session.add(input)
|
||||
session.commit()
|
||||
|
||||
def get_from_database(query: any) -> list:
|
||||
|
||||
|
||||
def get_from_database(query: type[SQLModel]) -> list[SQLModel]:
|
||||
with Session(engine) as session:
|
||||
results = session.exec(select(query)).all()
|
||||
return results
|
||||
|
||||
def get_from_database_where(query: any, variable: any, value: any, first_only: bool = False, offset: int | None = None, limit: int | None = None) -> list:
|
||||
|
||||
def get_from_database_where(
|
||||
query: type[SQLModel],
|
||||
variable: object,
|
||||
value: int | str | None,
|
||||
first_only: bool = False,
|
||||
offset: int | None = None,
|
||||
limit: int | None = None,
|
||||
) -> list[SQLModel]:
|
||||
with Session(engine) as session:
|
||||
statement = select(query).where(variable == value)
|
||||
if offset:
|
||||
|
|
@ -35,25 +46,47 @@ def get_from_database_where(query: any, variable: any, value: any, first_only: b
|
|||
results = results.all()
|
||||
return results
|
||||
|
||||
def get_from_database_where_id(query: any, value: int):
|
||||
|
||||
def get_from_database_where_id(
|
||||
query: type[SQLModel],
|
||||
value: int,
|
||||
) -> SQLModel | None:
|
||||
with Session(engine) as session:
|
||||
result = session.get(query, value)
|
||||
return result
|
||||
|
||||
def get_from_database_where_join_filter(query: any, joined: any, variable: any, value: any):
|
||||
|
||||
def get_from_database_where_join_filter(
|
||||
query: type[SQLModel],
|
||||
joined: SQLModel,
|
||||
variable: object,
|
||||
value: int | str | None,
|
||||
):
|
||||
with Session(engine) as session:
|
||||
results = session.exec(select(query).join(joined).where(variable == value))
|
||||
return results
|
||||
|
||||
def get_and_join_from_database(query: any, joined: any, isouter: bool = False):
|
||||
|
||||
def get_and_join_from_database(
|
||||
query: type[SQLModel],
|
||||
joined: SQLModel,
|
||||
isouter: bool = False,
|
||||
):
|
||||
with Session(engine) as session:
|
||||
results = session.exec(select(query, joined).join(joined, isouter)).all()
|
||||
return results
|
||||
|
||||
def update_database_where(query: any, variable: any, value: any, attribute: any, new_value: any):
|
||||
|
||||
def update_database_where(
|
||||
query: type[SQLModel],
|
||||
variable: object,
|
||||
value: int | str | None,
|
||||
attribute: object,
|
||||
new_value: int | str | None,
|
||||
):
|
||||
with Session(engine) as session:
|
||||
results = session.exec(select(query).where(variable == value))
|
||||
|
||||
|
||||
for instance in results:
|
||||
instance.setAttribute(attribute, new_value)
|
||||
session.add(instance)
|
||||
|
|
@ -61,18 +94,27 @@ def update_database_where(query: any, variable: any, value: any, attribute: any,
|
|||
session.refresh(instance)
|
||||
return results
|
||||
|
||||
def delete_from_database_where(query: any, variable: any, value: any):
|
||||
|
||||
def delete_from_database_where(
|
||||
query: type[SQLModel],
|
||||
variable: object,
|
||||
value: str | int | None,
|
||||
):
|
||||
with Session(engine) as session:
|
||||
results = session.exec(select(query).where(variable == value))
|
||||
|
||||
|
||||
for instance in results:
|
||||
session.delete(instance)
|
||||
session.commit()
|
||||
|
||||
|
||||
remainder = session.exec(select(query).where(variable == value))
|
||||
if remainder is None:
|
||||
return {"message": "Successfully deleted matching query(ies)", "deleted": results.all()}
|
||||
return {
|
||||
"message": "Successfully deleted matching query(ies)",
|
||||
"deleted": results.all(),
|
||||
}
|
||||
else:
|
||||
return {"message": "Failed to delete all matching queries", "remaining": remainder.all()}
|
||||
|
||||
|
||||
return {
|
||||
"message": "Failed to delete all matching queries",
|
||||
"remaining": remainder.all(),
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue