Generalized database function type args, fixed syntax

This commit is contained in:
Hannah dagemark 2026-08-10 23:11:11 +02:00
commit c07423f8b0

View file

@ -1,3 +1,4 @@
from sqlalchemy.orm.attributes import InstrumentedAttribute
from sqlmodel import Session, SQLModel, create_engine, select from sqlmodel import Session, SQLModel, create_engine, select
from app.config import settings from app.config import settings
@ -13,82 +14,95 @@ def create_database():
SQLModel.metadata.create_all(engine) SQLModel.metadata.create_all(engine)
### POST route
def add_to_database(input: SQLModel): def add_to_database(input: SQLModel):
with Session(engine) as session: with Session(engine) as session:
session.add(input) session.add(input)
session.commit() session.commit()
def get_from_database(query: type[SQLModel]): ### GET "/all" route
def get_from_database[T: SQLModel](query: type[T]) -> list[T] | None:
with Session(engine) as session: with Session(engine) as session:
results = session.exec(select(query)).all() results = list(session.exec(select(query)))
return results return results
def get_from_database_where( ### GET "/{value_of_var}" route
query: type[SQLModel], def get_from_database_where[T: SQLModel, V](
variable: object, query: type[T],
value: int | str | None, variable: InstrumentedAttribute[V],
first_only: bool = False, value: V,
offset: int | None = None, offset: int | None = None,
limit: int | None = None, limit: int | None = None,
): ) -> list[T] | None:
with Session(engine) as session: with Session(engine) as session:
statement = select(query).where(variable == value) statement = select(query).where(variable == value)
if offset: if offset is not None:
statement = statement.offset(offset) statement = statement.offset(offset)
if limit: if limit is not None:
statement = statement.limit(limit) statement = statement.limit(limit)
results = session.exec(statement) results = list(session.exec(statement))
if first_only:
results = results.first()
else:
results = results.all()
return results return results
def get_from_database_where_id( ### GET "/{id}" route
query: type[SQLModel], def get_from_database_where_id[T: SQLModel](
query: type[T],
value: int, value: int,
) -> SQLModel | None: ) -> T | None:
with Session(engine) as session: with Session(engine) as session:
result = session.get(query, value) result = session.get(query, value)
return result return result
def get_from_database_where_join_filter( ### GET "/{value where x for y and y has link to z}" route
query: type[SQLModel], def get_from_database_where_join_filter[
joined: type[SQLModel], T: SQLModel,
variable: object, K: SQLModel,
value: int | str | None, V,
): ](
query: type[T],
joined: type[K],
variable: InstrumentedAttribute[V],
value: V,
) -> list[T] | None:
with Session(engine) as session: with Session(engine) as session:
results = session.exec(select(query).join(joined).where(variable == value)) results = list(
session.exec(select(query).join(joined).where(variable == value))
)
return results return results
def get_and_join_from_database( ### GET "/{with table y}" route (can have outers or not)
query: type[SQLModel], def get_and_join_from_database[
joined: type[SQLModel], T: SQLModel,
K: SQLModel,
](
query: type[T],
joined: type[K],
isouter: bool = False, isouter: bool = False,
): ) -> list[tuple[T, K]] | None:
with Session(engine) as session: with Session(engine) as session:
results = session.exec(select(query, joined).join(joined, isouter=isouter)) results = list(
session.exec(select(query, joined).join(joined, isouter=isouter))
)
return results return results
def update_database_where( ### PATCH "/{value_of_var}" route
query: type[SQLModel], def update_database_where[T: SQLModel, V, W](
variable: object, query: type[T],
value: int | str | None, variable: InstrumentedAttribute[V],
attribute: str, value: V,
new_value: int | str | None, attribute: InstrumentedAttribute[W],
): new_value: W,
) -> list[T]:
with Session(engine) as session: with Session(engine) as session:
results = list(session.exec(select(query).where(variable == value))) results = list(session.exec(select(query).where(variable == value)))
for instance in results: for instance in results:
setattr(instance, attribute, new_value) setattr(instance, attribute.key, new_value)
session.add(instance) session.add(instance)
session.commit() session.commit()
@ -98,18 +112,18 @@ def update_database_where(
return results return results
def delete_from_database_where( ### DELETE "/{value_of_var}" route
query: type[SQLModel], def delete_from_database_where[T: SQLModel, V](
variable: object, query: type[T],
value: str | int | None, variable: InstrumentedAttribute[V],
): value: V,
) -> list[T]:
with Session(engine) as session: with Session(engine) as session:
results = session.exec(select(query).where(variable == value)) results = list(session.exec(select(query).where(variable == value)))
for instance in results: for instance in results:
session.delete(instance) session.delete(instance)
session.commit() session.commit()
remainder = session.exec(select(query).where(variable == value)).all() return results
return remainder