routerでの処理(JSとcrud処理を繋ぐ)
- sortを引数で受け取ることでクエリを取得できる
- データベースセッションとユーザ情報取得のために依存性注入
- dbは関数呼び出しと変わらない気がするが、ユーザ取得は色々必要な情報を取ってきた後で呼び出してくれるらしい
- 条件分岐でソートの方式、ソートのあるなしで分けてcrud操作の関数呼び出す
- 入れ子構造にしてあるため、受け取ったのちにjsonの形に合うように修正
@router.get("/tasks", response_model=list[TaskSchema])
async def get_tasks(
sort: str | None = None,
search_name: str | None = None,
db_session: AsyncSession = Depends(db.get_db_session),
current_user=Depends(get_current_user),
):
if sort in ["deadline", "status"]:
tasks = await arrange_tasks(db_session, current_user.user_id, sort)
elif search_name:
tasks = await filter_tasks(db_session, current_user.user_id, search_name)
else:
tasks = await fetch_tasks(db_session, current_user.user_id)
tasks_pydantic = []
for task in tasks:
task_status = TaskStatusSchema(
task_progress=task.task_progress,
progress_ratio=task.progress_ratio,
progress_comment=task.progress_comment,
)
task_pydantic = TaskSchema(
task_id=task.task_id,
task_name=task.task_name,
task_deadline=task.task_deadline,
task_detail=task.task_detail,
changed_time=task.changed_time,
task_status=task_status,
)
tasks_pydantic.append(task_pydantic)
return tasks_pydantic
crud側の処理
- ただsqlalchemyのメゾットを呼び出すのみなので特に書くことはない
async def fetch_tasks(db_session: AsyncSession, user_id: UUID) -> list[Task]:
results = await db_session.execute(select(Task).where(Task.user_id == user_id))
target_tasks = results.scalars().all()
return target_tasks
- 締切でのソートの場合はorder_byメゾットを呼び出すのみなので特に書くことはない。一応、第二引数に進捗順になるようなオーダーを入れておいた。
- タスクのプログレスには上下の概念がないため、それぞれの文字列に重みをつける必要がある。
async def arrange_tasks(
db_session: AsyncSession, user_id: UUID, sort_order: str
) -> list[Task]:
stmt = select(Task).where(Task.user_id == user_id)
if sort_order == "deadline":
stmt = stmt.order_by(Task.task_deadline.asc(), Task.task_progress.asc())
if sort_order == "status":
status_order = case(
(Task.task_progress == "TODO", 1),
(Task.task_progress == "IN_PROGRESS", 2),
(Task.task_progress == "DONE", 3),
else_=4,
)
stmt = stmt.order_by(status_order, Task.task_deadline.asc())
result = await db_session.execute(stmt)
return list(result.scalars().all())