diff --git a/packages/api/src/taskflow_api/routers/agents.py b/packages/api/src/taskflow_api/routers/agents.py index c5b6dac..c9d6ecb 100644 --- a/packages/api/src/taskflow_api/routers/agents.py +++ b/packages/api/src/taskflow_api/routers/agents.py @@ -23,7 +23,6 @@ async def list_agents( offset: int = Query(default=0, ge=0), ) -> list[WorkerRead]: """List all global agents.""" - # Ensure user is set up await ensure_user_setup(session, user) stmt = ( @@ -58,6 +57,9 @@ async def create_agent( ) -> WorkerRead: """Register a global agent.""" current_worker = await ensure_user_setup(session, user) + # Extract primitive values before any commits + current_worker_id = current_worker.id + current_worker_type = current_worker.type # Check handle uniqueness stmt = select(Worker).where(Worker.handle == data.handle) @@ -74,23 +76,27 @@ async def create_agent( capabilities=data.capabilities, ) session.add(agent) - await session.commit() - await session.refresh(agent) + await session.flush() # Get agent.id without committing - # Audit log + # Audit log (doesn't commit) await log_action( session, entity_type="worker", entity_id=agent.id, action="created", - actor_id=current_worker.id, + actor_id=current_worker_id, + actor_type=current_worker_type, details={ - "handle": agent.handle, - "agent_type": agent.agent_type, - "capabilities": agent.capabilities, + "handle": data.handle, + "agent_type": data.agent_type, + "capabilities": data.capabilities, }, ) + # Single commit for all changes + await session.commit() + await session.refresh(agent) + return WorkerRead( id=agent.id, handle=agent.handle, @@ -137,6 +143,8 @@ async def update_agent( ) -> WorkerRead: """Update agent details.""" current_worker = await ensure_user_setup(session, user) + current_worker_id = current_worker.id + current_worker_type = current_worker.type agent = await session.get(Worker, agent_id) if not agent: @@ -158,19 +166,21 @@ async def update_agent( if changes: session.add(agent) - await session.commit() - await session.refresh(agent) # Audit log await log_action( session, entity_type="worker", - entity_id=agent.id, + entity_id=agent_id, action="updated", - actor_id=current_worker.id, + actor_id=current_worker_id, + actor_type=current_worker_type, details=changes, ) + await session.commit() + await session.refresh(agent) + return WorkerRead( id=agent.id, handle=agent.handle, @@ -190,6 +200,8 @@ async def delete_agent( ) -> dict: """Delete an agent.""" current_worker = await ensure_user_setup(session, user) + current_worker_id = current_worker.id + current_worker_type = current_worker.type agent = await session.get(Worker, agent_id) if not agent: @@ -197,6 +209,10 @@ async def delete_agent( if agent.type != "agent": raise HTTPException(status_code=404, detail="Not an agent") + # Extract values before deletion + agent_handle = agent.handle + agent_type_val = agent.agent_type + # Check if agent is member of any project stmt = select(ProjectMember).where(ProjectMember.worker_id == agent_id) result = await session.exec(stmt) @@ -210,10 +226,11 @@ async def delete_agent( await log_action( session, entity_type="worker", - entity_id=agent.id, + entity_id=agent_id, action="deleted", - actor_id=current_worker.id, - details={"handle": agent.handle, "agent_type": agent.agent_type}, + actor_id=current_worker_id, + actor_type=current_worker_type, + details={"handle": agent_handle, "agent_type": agent_type_val}, ) await session.delete(agent) diff --git a/packages/api/src/taskflow_api/routers/members.py b/packages/api/src/taskflow_api/routers/members.py index 9b6e68b..635b4ec 100644 --- a/packages/api/src/taskflow_api/routers/members.py +++ b/packages/api/src/taskflow_api/routers/members.py @@ -23,6 +23,7 @@ async def list_members( ) -> list[MemberRead]: """List all members of a project (humans + agents).""" worker = await ensure_user_setup(session, user) + worker_id = worker.id # Check project exists project = await session.get(Project, project_id) @@ -32,7 +33,7 @@ async def list_members( # Check user is member stmt = select(ProjectMember).where( ProjectMember.project_id == project_id, - ProjectMember.worker_id == worker.id, + ProjectMember.worker_id == worker_id, ) result = await session.exec(stmt) if not result.first(): @@ -73,6 +74,8 @@ async def add_member( - agent_id: Existing agent worker ID (links to project) """ current_worker = await ensure_user_setup(session, user) + current_worker_id = current_worker.id + current_worker_type = current_worker.type # Validate input if not data.user_id and not data.agent_id: @@ -90,6 +93,7 @@ async def add_member( raise HTTPException(status_code=403, detail="Only project owner can add members") member_worker: Worker | None = None + member_worker_id: int | None = None if data.agent_id: # Link existing agent @@ -98,6 +102,7 @@ async def add_member( raise HTTPException(status_code=404, detail="Agent not found") if member_worker.type != "agent": raise HTTPException(status_code=400, detail="Worker is not an agent") + member_worker_id = member_worker.id elif data.user_id: # Get or create worker for SSO user @@ -105,7 +110,9 @@ async def add_member( result = await session.exec(stmt) member_worker = result.first() - if not member_worker: + if member_worker: + member_worker_id = member_worker.id + else: # Create worker for new user # We don't have their email/name, so use user_id handle = f"@user-{data.user_id[:8].lower()}" @@ -128,13 +135,13 @@ async def add_member( user_id=data.user_id, ) session.add(member_worker) - await session.commit() - await session.refresh(member_worker) + await session.flush() # Get ID without committing + member_worker_id = member_worker.id # Check not already a member stmt = select(ProjectMember).where( ProjectMember.project_id == project_id, - ProjectMember.worker_id == member_worker.id, + ProjectMember.worker_id == member_worker_id, ) result = await session.exec(stmt) if result.first(): @@ -143,33 +150,41 @@ async def add_member( # Add membership membership = ProjectMember( project_id=project_id, - worker_id=member_worker.id, + worker_id=member_worker_id, role="member", ) session.add(membership) - await session.commit() - await session.refresh(membership) - # Audit log + # Get member details for response before commit + member_handle = member_worker.handle + member_name = member_worker.name + member_type = member_worker.type + + # Audit log (doesn't commit) await log_action( session, entity_type="project", entity_id=project_id, action="member_added", - actor_id=current_worker.id, + actor_id=current_worker_id, + actor_type=current_worker_type, details={ - "worker_id": member_worker.id, - "handle": member_worker.handle, - "type": member_worker.type, + "worker_id": member_worker_id, + "handle": member_handle, + "type": member_type, }, ) + # Single commit + await session.commit() + await session.refresh(membership) + return MemberRead( id=membership.id, - worker_id=member_worker.id, - handle=member_worker.handle, - name=member_worker.name, - type=member_worker.type, + worker_id=member_worker_id, + handle=member_handle, + name=member_name, + type=member_type, role=membership.role, joined_at=membership.joined_at, ) @@ -184,6 +199,8 @@ async def remove_member( ) -> dict: """Remove a member from a project.""" current_worker = await ensure_user_setup(session, user) + current_worker_id = current_worker.id + current_worker_type = current_worker.type # Check project exists project = await session.get(Project, project_id) @@ -203,24 +220,27 @@ async def remove_member( if membership.role == "owner": raise HTTPException(status_code=400, detail="Cannot remove project owner") - # Get worker for audit - member_worker = await session.get(Worker, membership.worker_id) + # Get worker info for audit before deletion + removed_worker_id = membership.worker_id + member_worker = await session.get(Worker, removed_worker_id) + member_handle = member_worker.handle if member_worker else None - # Remove membership - await session.delete(membership) - await session.commit() - - # Audit log + # Audit log before deletion await log_action( session, entity_type="project", entity_id=project_id, action="member_removed", - actor_id=current_worker.id, + actor_id=current_worker_id, + actor_type=current_worker_type, details={ - "worker_id": member_worker.id if member_worker else None, - "handle": member_worker.handle if member_worker else None, + "worker_id": removed_worker_id, + "handle": member_handle, }, ) + # Remove membership + await session.delete(membership) + await session.commit() + return {"ok": True} diff --git a/packages/api/src/taskflow_api/routers/projects.py b/packages/api/src/taskflow_api/routers/projects.py index 1062475..c0245bb 100644 --- a/packages/api/src/taskflow_api/routers/projects.py +++ b/packages/api/src/taskflow_api/routers/projects.py @@ -26,11 +26,11 @@ async def list_projects( offset: int = Query(default=0, ge=0), ) -> list[ProjectRead]: """List projects where user is a member.""" - # Ensure user setup worker = await ensure_user_setup(session, user) + worker_id = worker.id # Get project IDs where user is a member - member_stmt = select(ProjectMember.project_id).where(ProjectMember.worker_id == worker.id) + member_stmt = select(ProjectMember.project_id).where(ProjectMember.worker_id == worker_id) member_result = await session.exec(member_stmt) project_ids = list(member_result.all()) @@ -80,8 +80,10 @@ async def create_project( user: CurrentUser = Depends(get_current_user), ) -> ProjectRead: """Create a new project.""" - # Ensure user setup worker = await ensure_user_setup(session, user) + # Extract primitive values before any commits + worker_id = worker.id + worker_type = worker.type # Check slug uniqueness stmt = select(Project).where(Project.slug == data.slug) @@ -98,28 +100,32 @@ async def create_project( is_default=False, ) session.add(project) - await session.commit() - await session.refresh(project) + await session.flush() # Get project.id without committing + project_id = project.id # Add creator as owner membership = ProjectMember( - project_id=project.id, - worker_id=worker.id, + project_id=project_id, + worker_id=worker_id, role="owner", ) session.add(membership) - await session.commit() - # Audit log + # Audit log (doesn't commit) await log_action( session, entity_type="project", - entity_id=project.id, + entity_id=project_id, action="created", - actor_id=worker.id, - details={"slug": project.slug, "name": project.name}, + actor_id=worker_id, + actor_type=worker_type, + details={"slug": data.slug, "name": data.name}, ) + # Single commit for all changes + await session.commit() + await session.refresh(project) + return ProjectRead( id=project.id, slug=project.slug, @@ -142,6 +148,7 @@ async def get_project( ) -> ProjectRead: """Get project details.""" worker = await ensure_user_setup(session, user) + worker_id = worker.id project = await session.get(Project, project_id) if not project: @@ -150,7 +157,7 @@ async def get_project( # Check membership stmt = select(ProjectMember).where( ProjectMember.project_id == project_id, - ProjectMember.worker_id == worker.id, + ProjectMember.worker_id == worker_id, ) result = await session.exec(stmt) if not result.first(): @@ -188,6 +195,8 @@ async def update_project( ) -> ProjectRead: """Update project (owner only).""" worker = await ensure_user_setup(session, user) + worker_id = worker.id + worker_type = worker.type project = await session.get(Project, project_id) if not project: @@ -209,19 +218,21 @@ async def update_project( if changes: project.updated_at = datetime.utcnow() session.add(project) - await session.commit() - await session.refresh(project) # Audit log await log_action( session, entity_type="project", - entity_id=project.id, + entity_id=project_id, action="updated", - actor_id=worker.id, + actor_id=worker_id, + actor_type=worker_type, details=changes, ) + await session.commit() + await session.refresh(project) + # Get counts for response member_count_stmt = select(func.count(ProjectMember.id)).where( ProjectMember.project_id == project_id @@ -254,11 +265,16 @@ async def delete_project( ) -> dict: """Delete project (owner only).""" worker = await ensure_user_setup(session, user) + worker_id = worker.id + worker_type = worker.type project = await session.get(Project, project_id) if not project: raise HTTPException(status_code=404, detail="Project not found") + # Extract values before any modifications + project_slug = project.slug + # Check ownership if project.owner_id != user.id: raise HTTPException(status_code=403, detail="Only project owner can delete") @@ -294,10 +310,11 @@ async def delete_project( await log_action( session, entity_type="project", - entity_id=project.id, + entity_id=project_id, action="deleted", - actor_id=worker.id, - details={"slug": project.slug, "force": force, "task_count": task_count}, + actor_id=worker_id, + actor_type=worker_type, + details={"slug": project_slug, "force": force, "task_count": task_count}, ) await session.delete(project) diff --git a/packages/api/src/taskflow_api/routers/tasks.py b/packages/api/src/taskflow_api/routers/tasks.py index abc22f7..c943902 100644 --- a/packages/api/src/taskflow_api/routers/tasks.py +++ b/packages/api/src/taskflow_api/routers/tasks.py @@ -137,12 +137,13 @@ async def list_tasks( ) -> list[TaskListItem]: """List tasks in a project with optional filters.""" worker = await ensure_user_setup(session, user) + worker_id = worker.id # Check project exists and user is member project = await session.get(Project, project_id) if not project: raise HTTPException(status_code=404, detail="Project not found") - await check_project_membership(session, project_id, worker.id) + await check_project_membership(session, project_id, worker_id) # Build query stmt = select(Task).where(Task.project_id == project_id) @@ -193,17 +194,21 @@ async def create_task( ) -> TaskRead: """Create a new task in a project.""" worker = await ensure_user_setup(session, user) + worker_id = worker.id + worker_type = worker.type # Check project exists and user is member project = await session.get(Project, project_id) if not project: raise HTTPException(status_code=404, detail="Project not found") - await check_project_membership(session, project_id, worker.id) + await check_project_membership(session, project_id, worker_id) # Validate assignee if provided assignee = None + assignee_handle = None if data.assignee_id: assignee = await check_assignee_is_member(session, project_id, data.assignee_id) + assignee_handle = assignee.handle # Validate parent if provided if data.parent_task_id: @@ -219,19 +224,19 @@ async def create_task( tags=data.tags, due_date=data.due_date, project_id=project_id, - created_by_id=worker.id, + created_by_id=worker_id, ) session.add(task) - await session.commit() - await session.refresh(task) + await session.flush() # Get task.id without committing - # Audit log + # Audit log (doesn't commit) await log_action( session, entity_type="task", entity_id=task.id, action="created", - actor_id=worker.id, + actor_id=worker_id, + actor_type=worker_type, details={ "title": task.title, "priority": task.priority, @@ -239,7 +244,30 @@ async def create_task( }, ) - return task_to_read(task, assignee) + # Single commit + await session.commit() + await session.refresh(task) + + return TaskRead( + id=task.id, + title=task.title, + description=task.description, + status=task.status, + priority=task.priority, + progress_percent=task.progress_percent, + tags=task.tags, + due_date=task.due_date, + project_id=task.project_id, + assignee_id=task.assignee_id, + assignee_handle=assignee_handle, + parent_task_id=task.parent_task_id, + created_by_id=task.created_by_id, + started_at=task.started_at, + completed_at=task.completed_at, + created_at=task.created_at, + updated_at=task.updated_at, + subtasks=[], + ) # Task-specific endpoints (not project-scoped) @@ -253,13 +281,14 @@ async def get_task( ) -> TaskRead: """Get task details including subtasks.""" worker = await ensure_user_setup(session, user) + worker_id = worker.id task = await session.get(Task, task_id) if not task: raise HTTPException(status_code=404, detail="Task not found") # Check membership - await check_project_membership(session, task.project_id, worker.id) + await check_project_membership(session, task.project_id, worker_id) # Get assignee assignee = None @@ -283,12 +312,14 @@ async def update_task( ) -> TaskRead: """Update task details.""" worker = await ensure_user_setup(session, user) + worker_id = worker.id + worker_type = worker.type task = await session.get(Task, task_id) if not task: raise HTTPException(status_code=404, detail="Task not found") - await check_project_membership(session, task.project_id, worker.id) + await check_project_membership(session, task.project_id, worker_id) # Track changes changes = {} @@ -314,18 +345,20 @@ async def update_task( if changes: task.updated_at = datetime.utcnow() session.add(task) - await session.commit() - await session.refresh(task) await log_action( session, entity_type="task", - entity_id=task.id, + entity_id=task_id, action="updated", - actor_id=worker.id, + actor_id=worker_id, + actor_type=worker_type, details=changes, ) + await session.commit() + await session.refresh(task) + assignee = None if task.assignee_id: assignee = await session.get(Worker, task.assignee_id) @@ -341,12 +374,18 @@ async def delete_task( ) -> dict: """Delete a task.""" worker = await ensure_user_setup(session, user) + worker_id = worker.id + worker_type = worker.type task = await session.get(Task, task_id) if not task: raise HTTPException(status_code=404, detail="Task not found") - await check_project_membership(session, task.project_id, worker.id) + await check_project_membership(session, task.project_id, worker_id) + + # Extract values before deletion + task_title = task.title + task_status = task.status # Check for subtasks stmt = select(Task).where(Task.parent_task_id == task_id) @@ -358,10 +397,11 @@ async def delete_task( await log_action( session, entity_type="task", - entity_id=task.id, + entity_id=task_id, action="deleted", - actor_id=worker.id, - details={"title": task.title, "status": task.status}, + actor_id=worker_id, + actor_type=worker_type, + details={"title": task_title, "status": task_status}, ) await session.delete(task) @@ -382,12 +422,14 @@ async def update_status( ) -> TaskRead: """Change task status with transition validation.""" worker = await ensure_user_setup(session, user) + worker_id = worker.id + worker_type = worker.type task = await session.get(Task, task_id) if not task: raise HTTPException(status_code=404, detail="Task not found") - await check_project_membership(session, task.project_id, worker.id) + await check_project_membership(session, task.project_id, worker_id) # Validate transition if not validate_status_transition(task.status, data.status): @@ -410,18 +452,20 @@ async def update_status( task.progress_percent = 100 session.add(task) - await session.commit() - await session.refresh(task) await log_action( session, entity_type="task", - entity_id=task.id, + entity_id=task_id, action="status_changed", - actor_id=worker.id, + actor_id=worker_id, + actor_type=worker_type, details={"before": old_status, "after": data.status}, ) + await session.commit() + await session.refresh(task) + assignee = None if task.assignee_id: assignee = await session.get(Worker, task.assignee_id) @@ -438,12 +482,14 @@ async def update_progress( ) -> TaskRead: """Update task progress percentage.""" worker = await ensure_user_setup(session, user) + worker_id = worker.id + worker_type = worker.type task = await session.get(Task, task_id) if not task: raise HTTPException(status_code=404, detail="Task not found") - await check_project_membership(session, task.project_id, worker.id) + await check_project_membership(session, task.project_id, worker_id) if task.status != "in_progress": raise HTTPException( @@ -455,18 +501,20 @@ async def update_progress( task.updated_at = datetime.utcnow() session.add(task) - await session.commit() - await session.refresh(task) await log_action( session, entity_type="task", - entity_id=task.id, + entity_id=task_id, action="progress_updated", - actor_id=worker.id, + actor_id=worker_id, + actor_type=worker_type, details={"before": old_progress, "after": data.percent, "note": data.note}, ) + await session.commit() + await session.refresh(task) + assignee = None if task.assignee_id: assignee = await session.get(Worker, task.assignee_id) @@ -483,38 +531,62 @@ async def assign_task( ) -> TaskRead: """Assign task to a project member.""" worker = await ensure_user_setup(session, user) + worker_id = worker.id + worker_type = worker.type task = await session.get(Task, task_id) if not task: raise HTTPException(status_code=404, detail="Task not found") - await check_project_membership(session, task.project_id, worker.id) + await check_project_membership(session, task.project_id, worker_id) # Validate assignee assignee = await check_assignee_is_member(session, task.project_id, data.assignee_id) + assignee_handle = assignee.handle old_assignee_id = task.assignee_id task.assignee_id = data.assignee_id task.updated_at = datetime.utcnow() session.add(task) - await session.commit() - await session.refresh(task) await log_action( session, entity_type="task", - entity_id=task.id, + entity_id=task_id, action="assigned", - actor_id=worker.id, + actor_id=worker_id, + actor_type=worker_type, details={ "before": old_assignee_id, "after": data.assignee_id, - "assignee_handle": assignee.handle, + "assignee_handle": assignee_handle, }, ) - return task_to_read(task, assignee) + await session.commit() + await session.refresh(task) + + return TaskRead( + id=task.id, + title=task.title, + description=task.description, + status=task.status, + priority=task.priority, + progress_percent=task.progress_percent, + tags=task.tags, + due_date=task.due_date, + project_id=task.project_id, + assignee_id=task.assignee_id, + assignee_handle=assignee_handle, + parent_task_id=task.parent_task_id, + created_by_id=task.created_by_id, + started_at=task.started_at, + completed_at=task.completed_at, + created_at=task.created_at, + updated_at=task.updated_at, + subtasks=[], + ) @router.post("/api/tasks/{task_id}/subtasks", response_model=TaskRead, status_code=201) @@ -526,17 +598,24 @@ async def create_subtask( ) -> TaskRead: """Create a subtask under a parent task.""" worker = await ensure_user_setup(session, user) + worker_id = worker.id + worker_type = worker.type parent = await session.get(Task, task_id) if not parent: raise HTTPException(status_code=404, detail="Parent task not found") - await check_project_membership(session, parent.project_id, worker.id) + await check_project_membership(session, parent.project_id, worker_id) + + # Get parent's project_id before any modifications + parent_project_id = parent.project_id # Validate assignee if provided assignee = None + assignee_handle = None if data.assignee_id: - assignee = await check_assignee_is_member(session, parent.project_id, data.assignee_id) + assignee = await check_assignee_is_member(session, parent_project_id, data.assignee_id) + assignee_handle = assignee.handle # Create subtask subtask = Task( @@ -547,19 +626,19 @@ async def create_subtask( parent_task_id=task_id, tags=data.tags, due_date=data.due_date, - project_id=parent.project_id, - created_by_id=worker.id, + project_id=parent_project_id, + created_by_id=worker_id, ) session.add(subtask) - await session.commit() - await session.refresh(subtask) + await session.flush() await log_action( session, entity_type="task", entity_id=subtask.id, action="created", - actor_id=worker.id, + actor_id=worker_id, + actor_type=worker_type, details={ "title": subtask.title, "parent_task_id": task_id, @@ -567,7 +646,29 @@ async def create_subtask( }, ) - return task_to_read(subtask, assignee) + await session.commit() + await session.refresh(subtask) + + return TaskRead( + id=subtask.id, + title=subtask.title, + description=subtask.description, + status=subtask.status, + priority=subtask.priority, + progress_percent=subtask.progress_percent, + tags=subtask.tags, + due_date=subtask.due_date, + project_id=subtask.project_id, + assignee_id=subtask.assignee_id, + assignee_handle=assignee_handle, + parent_task_id=subtask.parent_task_id, + created_by_id=subtask.created_by_id, + started_at=subtask.started_at, + completed_at=subtask.completed_at, + created_at=subtask.created_at, + updated_at=subtask.updated_at, + subtasks=[], + ) @router.post("/api/tasks/{task_id}/approve", response_model=TaskRead) @@ -578,12 +679,14 @@ async def approve_task( ) -> TaskRead: """Approve a task in review status.""" worker = await ensure_user_setup(session, user) + worker_id = worker.id + worker_type = worker.type task = await session.get(Task, task_id) if not task: raise HTTPException(status_code=404, detail="Task not found") - await check_project_membership(session, task.project_id, worker.id) + await check_project_membership(session, task.project_id, worker_id) if task.status != "review": raise HTTPException(status_code=400, detail="Can only approve tasks in 'review' status") @@ -594,18 +697,20 @@ async def approve_task( task.updated_at = datetime.utcnow() session.add(task) - await session.commit() - await session.refresh(task) await log_action( session, entity_type="task", - entity_id=task.id, + entity_id=task_id, action="approved", - actor_id=worker.id, + actor_id=worker_id, + actor_type=worker_type, details={"from_status": "review", "to_status": "completed"}, ) + await session.commit() + await session.refresh(task) + assignee = None if task.assignee_id: assignee = await session.get(Worker, task.assignee_id) @@ -622,12 +727,14 @@ async def reject_task( ) -> TaskRead: """Reject a task in review status, returning it to in_progress.""" worker = await ensure_user_setup(session, user) + worker_id = worker.id + worker_type = worker.type task = await session.get(Task, task_id) if not task: raise HTTPException(status_code=404, detail="Task not found") - await check_project_membership(session, task.project_id, worker.id) + await check_project_membership(session, task.project_id, worker_id) if task.status != "review": raise HTTPException(status_code=400, detail="Can only reject tasks in 'review' status") @@ -636,18 +743,20 @@ async def reject_task( task.updated_at = datetime.utcnow() session.add(task) - await session.commit() - await session.refresh(task) await log_action( session, entity_type="task", - entity_id=task.id, + entity_id=task_id, action="rejected", - actor_id=worker.id, + actor_id=worker_id, + actor_type=worker_type, details={"reason": data.reason, "from_status": "review", "to_status": "in_progress"}, ) + await session.commit() + await session.refresh(task) + assignee = None if task.assignee_id: assignee = await session.get(Worker, task.assignee_id) diff --git a/packages/api/src/taskflow_api/schemas/task.py b/packages/api/src/taskflow_api/schemas/task.py index f32d302..15407f4 100644 --- a/packages/api/src/taskflow_api/schemas/task.py +++ b/packages/api/src/taskflow_api/schemas/task.py @@ -1,9 +1,24 @@ """Task API schemas.""" -from datetime import datetime +from datetime import UTC, datetime from typing import Literal -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, field_validator + + +def strip_timezone(dt: datetime | None) -> datetime | None: + """Convert timezone-aware datetime to naive UTC datetime. + + Database stores naive datetimes, so we strip timezone info + after converting to UTC. + """ + if dt is None: + return None + if dt.tzinfo is not None: + # Convert to UTC and strip timezone + + dt = dt.astimezone(UTC).replace(tzinfo=None) + return dt class TaskCreate(BaseModel): @@ -23,6 +38,20 @@ class TaskCreate(BaseModel): tags: list[str] = Field(default_factory=list) due_date: datetime | None = None + @field_validator("assignee_id", "parent_task_id", mode="after") + @classmethod + def zero_to_none(cls, v: int | None) -> int | None: + """Convert 0 to None (0 is not a valid foreign key).""" + if v == 0: + return None + return v + + @field_validator("due_date", mode="after") + @classmethod + def normalize_due_date(cls, v: datetime | None) -> datetime | None: + """Strip timezone from due_date for database compatibility.""" + return strip_timezone(v) + class TaskUpdate(BaseModel): """Schema for updating a task.""" @@ -33,6 +62,12 @@ class TaskUpdate(BaseModel): tags: list[str] | None = None due_date: datetime | None = None + @field_validator("due_date", mode="after") + @classmethod + def normalize_due_date(cls, v: datetime | None) -> datetime | None: + """Strip timezone from due_date for database compatibility.""" + return strip_timezone(v) + class StatusUpdate(BaseModel): """Schema for changing task status.""" diff --git a/packages/api/src/taskflow_api/services/audit.py b/packages/api/src/taskflow_api/services/audit.py index 82531eb..6466443 100644 --- a/packages/api/src/taskflow_api/services/audit.py +++ b/packages/api/src/taskflow_api/services/audit.py @@ -6,7 +6,6 @@ from sqlmodel.ext.asyncio.session import AsyncSession from ..models.audit import AuditLog -from ..models.worker import Worker async def log_action( @@ -16,25 +15,26 @@ async def log_action( entity_id: int, action: str, actor_id: int, + actor_type: str = "human", details: dict[str, Any] | None = None, ) -> AuditLog: """Create an immutable audit log entry. + Note: This does NOT commit - caller must commit the transaction. + This allows the audit log to be part of the same transaction as the action. + Args: session: Database session entity_type: Type of entity (task, project, worker) entity_id: ID of the affected entity action: Action performed (created, updated, started, etc.) actor_id: Worker ID who performed the action + actor_type: Type of actor ("human" or "agent") details: Additional context (before/after values, notes) Returns: - Created AuditLog entry + Created AuditLog entry (not yet committed) """ - # Get actor type from worker - worker = await session.get(Worker, actor_id) - actor_type = worker.type if worker else "human" - log = AuditLog( entity_type=entity_type, entity_id=entity_id, @@ -44,8 +44,6 @@ async def log_action( details=details or {}, ) session.add(log) - await session.commit() - await session.refresh(log) return log diff --git a/packages/api/src/taskflow_api/services/user_setup.py b/packages/api/src/taskflow_api/services/user_setup.py index e487e38..e3d24ac 100644 --- a/packages/api/src/taskflow_api/services/user_setup.py +++ b/packages/api/src/taskflow_api/services/user_setup.py @@ -52,11 +52,16 @@ async def get_or_create_worker(session: AsyncSession, user: CurrentUser) -> Work async def ensure_default_project( - session: AsyncSession, user: CurrentUser, worker: Worker + session: AsyncSession, user: CurrentUser, worker_id: int ) -> Project: """Ensure user has a Default project. Creates one if it doesn't exist. + + Args: + session: Database session + user: Current user from auth + worker_id: Worker ID (passed as int to avoid detached object issues) """ # Check if default project exists stmt = select(Project).where(Project.owner_id == user.id, Project.is_default.is_(True)) @@ -94,7 +99,7 @@ async def ensure_default_project( # Add user as owner membership = ProjectMember( project_id=project.id, - worker_id=worker.id, + worker_id=worker_id, role="owner", ) session.add(membership) @@ -114,5 +119,9 @@ async def ensure_user_setup(session: AsyncSession, user: CurrentUser) -> Worker: Returns the user's Worker record. """ worker = await get_or_create_worker(session, user) - await ensure_default_project(session, user, worker) + # Store worker.id before any further commits to avoid detached object issues + worker_id = worker.id + await ensure_default_project(session, user, worker_id) + # Refresh worker before returning to ensure it's attached to session + await session.refresh(worker) return worker