| 1 | from typing import List |
| 2 | |
| 3 | from fastapi import APIRouter |
| 4 | from fastapi import Depends |
| 5 | from sqlalchemy import delete |
| 6 | from sqlalchemy import select |
| 7 | from sqlalchemy.ext.asyncio import AsyncSession |
| 8 | |
| 9 | from app.auth.models.users import User |
| 10 | from app.auth.models.users import UserCustomerAccess |
| 11 | from app.auth.utils import AuthHandler |
| 12 | from app.db.db_session import get_db |
| 13 | from app.middleware.customer_access import customer_access_handler |
| 14 | |
| 15 | customer_users_router = APIRouter() |
| 16 | |
| 17 | |
| 18 | @customer_users_router.post("/users/{user_id}/customers") |
| 19 | async def assign_customer_access( |
| 20 | user_id: int, |
| 21 | customer_codes: List[str], |
| 22 | current_user: User = Depends(AuthHandler().require_any_scope("admin")), |
| 23 | session: AsyncSession = Depends(get_db), |
| 24 | ): |
| 25 | """Assign customer access to a user (admin only)""" |
| 26 | |
| 27 | # Remove existing access |
| 28 | await session.execute(delete(UserCustomerAccess).where(UserCustomerAccess.user_id == user_id)) |
| 29 | |
| 30 | # Add new access |
| 31 | for customer_code in customer_codes: |
| 32 | access = UserCustomerAccess(user_id=user_id, customer_code=customer_code) |
| 33 | session.add(access) |
| 34 | |
| 35 | await session.commit() |
| 36 | |
| 37 | return {"success": True, "message": f"Assigned {len(customer_codes)} customers to user {user_id}", "customer_codes": customer_codes} |
| 38 | |
| 39 | |
| 40 | @customer_users_router.get("/users/{user_id}/customers") |
| 41 | async def get_user_customer_access( |
| 42 | user_id: int, |
| 43 | current_user: User = Depends(AuthHandler().require_any_scope("admin")), |
| 44 | session: AsyncSession = Depends(get_db), |
| 45 | ): |
| 46 | """Get customer codes accessible to user (admin only)""" |
| 47 | |
| 48 | result = await session.execute(select(UserCustomerAccess.customer_code).where(UserCustomerAccess.user_id == user_id)) |
| 49 | customer_codes = result.scalars().all() |
| 50 | |
| 51 | return {"success": True, "customer_codes": customer_codes} |
| 52 | |
| 53 | |
| 54 | @customer_users_router.get("/me/customers") |
| 55 | async def get_my_customer_access(current_user: User = Depends(AuthHandler().get_current_user), session: AsyncSession = Depends(get_db)): |
| 56 | """Get current user's accessible customers""" |
| 57 | customer_codes = await customer_access_handler.get_user_accessible_customers(current_user, session) |
| 58 | |
| 59 | return {"success": True, "customer_codes": customer_codes} |