Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
47 changes: 32 additions & 15 deletions packages/api/src/taskflow_api/routers/agents.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = (
Expand Down Expand Up @@ -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)
Expand All @@ -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,
Expand Down Expand Up @@ -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:
Expand All @@ -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,
Expand All @@ -190,13 +200,19 @@ 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:
raise HTTPException(status_code=404, detail="Agent not found")
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)
Expand All @@ -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)
Expand Down
74 changes: 47 additions & 27 deletions packages/api/src/taskflow_api/routers/members.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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():
Expand Down Expand Up @@ -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:
Expand All @@ -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
Expand All @@ -98,14 +102,17 @@ 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
stmt = select(Worker).where(Worker.user_id == data.user_id)
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()}"
Expand All @@ -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():
Expand All @@ -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

Copilot AI Dec 7, 2025

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Potential None dereference: member_worker could theoretically be None at this point if neither the agent_id nor user_id branches are executed (though the validation at lines 81-84 should prevent this). Consider adding an assertion or explicit check before accessing member_worker.handle, member_worker.name, and member_worker.type to make the code more robust and satisfy type checkers:

if not member_worker:
    raise HTTPException(status_code=500, detail="Internal error: member_worker not initialized")
Suggested change
# Get member details for response before commit
# Get member details for response before commit
if not member_worker:
raise HTTPException(status_code=500, detail="Internal error: member_worker not initialized")

Copilot uses AI. Check for mistakes.
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,
)
Expand All @@ -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)
Expand All @@ -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}
Loading