@cryptotaxi247 / CoPilot / commits / 79737fc1

Update WazuhAgentScaPolicyResults to make 'reason' optional with defa… (#518)

* Update WazuhAgentScaPolicyResults to make 'reason' optional with default value * precommit fixes

taylor_socfortress committed Sep 25, 2025 at 08:16 UTC 79737fc1f48fd8aad7c26685732aa6b50d447bb1
12 files changed +187 -479
backend/alembic/versions/d8f9e9ea5502_add_customer_access_to_user.py
+22 -14
@@ -5,34 +5,42 @@ Revises: 7b2bbee2f3e8
5 Create Date: 2025-09-15 11:20:49.491950
6
7 """
8 -from typing import Sequence, Union
8 +from typing import Sequence
9 +from typing import Union
10
10 -from alembic import op
11 import sqlalchemy as sa
12 -from sqlalchemy.dialects import mysql
12 +
13 +from alembic import op
14
15 # revision identifiers, used by Alembic.
15 -revision: str = 'd8f9e9ea5502'
16 -down_revision: Union[str, None] = '7b2bbee2f3e8'
16 +revision: str = "d8f9e9ea5502"
17 +down_revision: Union[str, None] = "7b2bbee2f3e8"
18 branch_labels: Union[str, Sequence[str], None] = None
19 depends_on: Union[str, Sequence[str], None] = None
20
21
22 def upgrade() -> None:
23 # ### commands auto generated by Alembic - please adjust! ###
23 - op.create_table('user_customer_access',
24 - sa.Column('id', sa.Integer(), nullable=False),
25 - sa.Column('user_id', sa.Integer(), nullable=False),
26 - sa.Column('customer_code', sa.String(length=255), nullable=False),
27 - sa.Column('created_at', sa.DateTime(), nullable=False),
28 - sa.ForeignKeyConstraint(['customer_code'], ['customers.customer_code'], ),
29 - sa.ForeignKeyConstraint(['user_id'], ['user.id'], ),
30 - sa.PrimaryKeyConstraint('id')
24 + op.create_table(
25 + "user_customer_access",
26 + sa.Column("id", sa.Integer(), nullable=False),
27 + sa.Column("user_id", sa.Integer(), nullable=False),
28 + sa.Column("customer_code", sa.String(length=255), nullable=False),
29 + sa.Column("created_at", sa.DateTime(), nullable=False),
30 + sa.ForeignKeyConstraint(
31 + ["customer_code"],
32 + ["customers.customer_code"],
33 + ),
34 + sa.ForeignKeyConstraint(
35 + ["user_id"],
36 + ["user.id"],
37 + ),
38 + sa.PrimaryKeyConstraint("id"),
39 )
40 # ### end Alembic commands ###
41
42
43 def downgrade() -> None:
44 # ### commands auto generated by Alembic - please adjust! ###
37 - op.drop_table('user_customer_access')
45 + op.drop_table("user_customer_access")
46 # ### end Alembic commands ###
backend/app/agents/routes/agents.py
+38 -76
@@ -41,10 +41,10 @@ from app.agents.wazuh.services.sca import collect_agent_sca_policy_results
41 from app.agents.wazuh.services.vulnerabilities import collect_agent_vulnerabilities
42 from app.agents.wazuh.services.vulnerabilities import collect_agent_vulnerabilities_new
43 from app.agents.wazuh.services.vulnerabilities import sync_agent_vulnerabilities
44 +from app.auth.models.users import User
45
46 # App specific imports
47 from app.auth.routes.auth import AuthHandler
47 -from app.auth.models.users import User
48 from app.connectors.wazuh_manager.utils.universal import send_get_request
49 from app.db.db_session import get_db
50
@@ -158,10 +158,7 @@ async def delete_agent_from_database(db: AsyncSession, agent_id: str):
158 description="Get all agents currently synced to the database",
159 dependencies=[Security(AuthHandler().require_any_scope("admin", "analyst", "customer_user"))],
160 )
161 -async def get_agents(
162 - current_user: User = Depends(AuthHandler().get_current_user),
163 - db: AsyncSession = Depends(get_db)
164 -) -> AgentsResponse:
161 +async def get_agents(current_user: User = Depends(AuthHandler().get_current_user), db: AsyncSession = Depends(get_db)) -> AgentsResponse:
162 """
163 Retrieve all agents currently synced to the database.
164 Results are filtered based on user's customer access permissions.
@@ -176,9 +173,7 @@ async def get_agents(
173 try:
174 # Apply customer access filtering
175 base_query = select(Agents)
179 - filtered_query = await customer_access_handler.filter_query_by_customer_access(
180 - current_user, db, base_query, Agents.customer_code
181 - )
176 + filtered_query = await customer_access_handler.filter_query_by_customer_access(current_user, db, base_query, Agents.customer_code)
177
178 result = await db.execute(filtered_query)
179 agents = result.scalars().all()
@@ -270,9 +265,7 @@ async def get_agent(
265 try:
266 # Apply customer access filtering
267 base_query = select(Agents).filter(Agents.agent_id == agent_id)
273 - filtered_query = await customer_access_handler.filter_query_by_customer_access(
274 - current_user, db, base_query, Agents.customer_code
275 - )
268 + filtered_query = await customer_access_handler.filter_query_by_customer_access(current_user, db, base_query, Agents.customer_code)
269
270 result = await db.execute(filtered_query)
271 agent = result.scalars().first()
@@ -327,9 +320,7 @@ async def get_agent_by_hostname(
320 try:
321 # Apply customer access filtering
322 base_query = select(Agents).filter(Agents.hostname == hostname)
330 - filtered_query = await customer_access_handler.filter_query_by_customer_access(
331 - current_user, db, base_query, Agents.customer_code
332 - )
323 + filtered_query = await customer_access_handler.filter_query_by_customer_access(current_user, db, base_query, Agents.customer_code)
324
325 result = await db.execute(filtered_query)
326 agent = result.scalars().first()
@@ -411,7 +402,10 @@ async def mark_agent_as_critical(
402 # Check customer access - first find the agent
403 base_query = select(Agents).filter(Agents.agent_id == agent_id)
404 filtered_query = await customer_access_handler.filter_query_by_customer_access(
414 - current_user, session, base_query, Agents.customer_code
405 + current_user,
406 + session,
407 + base_query,
408 + Agents.customer_code,
409 )
410
411 result = await session.execute(filtered_query)
@@ -469,7 +463,10 @@ async def mark_agent_as_not_critical(
463 # Check customer access - first find the agent
464 base_query = select(Agents).filter(Agents.agent_id == agent_id)
465 filtered_query = await customer_access_handler.filter_query_by_customer_access(
472 - current_user, session, base_query, Agents.customer_code
466 + current_user,
467 + session,
468 + base_query,
469 + Agents.customer_code,
470 )
471
472 result = await session.execute(filtered_query)
@@ -524,17 +521,17 @@ async def upgrade_wazuh_agent_route(
521 # Check customer access - first find the agent
522 base_query = select(Agents).filter(Agents.agent_id == agent_id)
523 filtered_query = await customer_access_handler.filter_query_by_customer_access(
527 - current_user, session, base_query, Agents.customer_code
524 + current_user,
525 + session,
526 + base_query,
527 + Agents.customer_code,
528 )
529
530 result = await session.execute(filtered_query)
531 agent = result.scalars().first()
532
533 if not agent:
534 - raise HTTPException(
535 - status_code=404,
536 - detail=f"Agent with agent_id {agent_id} not found or access denied"
537 - )
534 + raise HTTPException(status_code=404, detail=f"Agent with agent_id {agent_id} not found or access denied")
535 return await upgrade_wazuh_agent(agent_id)
536 return AgentWazuhUpgradeResponse(
537 success=True,
@@ -577,18 +574,13 @@ async def get_agent_vulnerabilities(
574
575 # Check customer access - first verify user has access to this agent
576 base_query = select(Agents).filter(Agents.agent_id == agent_id)
580 - filtered_query = await customer_access_handler.filter_query_by_customer_access(
581 - current_user, session, base_query, Agents.customer_code
582 - )
577 + filtered_query = await customer_access_handler.filter_query_by_customer_access(current_user, session, base_query, Agents.customer_code)
578
579 result = await session.execute(filtered_query)
580 agent = result.scalars().first()
581
582 if not agent:
588 - raise HTTPException(
589 - status_code=404,
590 - detail=f"Agent with agent_id {agent_id} not found or access denied"
591 - )
583 + raise HTTPException(status_code=404, detail=f"Agent with agent_id {agent_id} not found or access denied")
584
585 wazuh_new = await check_wazuh_manager_version()
586 if wazuh_new is True:
@@ -625,18 +617,13 @@ async def get_agent_vulnerabilities_csv(
617
618 # Check customer access - first verify user has access to this agent
619 base_query = select(Agents).filter(Agents.agent_id == agent_id)
628 - filtered_query = await customer_access_handler.filter_query_by_customer_access(
629 - current_user, session, base_query, Agents.customer_code
630 - )
620 + filtered_query = await customer_access_handler.filter_query_by_customer_access(current_user, session, base_query, Agents.customer_code)
621
622 result = await session.execute(filtered_query)
623 agent = result.scalars().first()
624
625 if not agent:
636 - raise HTTPException(
637 - status_code=404,
638 - detail=f"Agent with agent_id {agent_id} not found or access denied"
639 - )
626 + raise HTTPException(status_code=404, detail=f"Agent with agent_id {agent_id} not found or access denied")
627
628 wazuh_new = await check_wazuh_manager_version()
629 if wazuh_new is True:
@@ -727,18 +714,13 @@ async def get_agent_sca(
714
715 # Check customer access - first verify user has access to this agent
716 base_query = select(Agents).filter(Agents.agent_id == agent_id)
730 - filtered_query = await customer_access_handler.filter_query_by_customer_access(
731 - current_user, session, base_query, Agents.customer_code
732 - )
717 + filtered_query = await customer_access_handler.filter_query_by_customer_access(current_user, session, base_query, Agents.customer_code)
718
719 result = await session.execute(filtered_query)
720 agent = result.scalars().first()
721
722 if not agent:
738 - raise HTTPException(
739 - status_code=404,
740 - detail=f"Agent with agent_id {agent_id} not found or access denied"
741 - )
723 + raise HTTPException(status_code=404, detail=f"Agent with agent_id {agent_id} not found or access denied")
724
725 return await collect_agent_sca(agent_id)
726
@@ -772,18 +754,13 @@ async def get_agent_sca_policy_results(
754
755 # Check customer access - first verify user has access to this agent
756 base_query = select(Agents).filter(Agents.agent_id == agent_id)
775 - filtered_query = await customer_access_handler.filter_query_by_customer_access(
776 - current_user, session, base_query, Agents.customer_code
777 - )
757 + filtered_query = await customer_access_handler.filter_query_by_customer_access(current_user, session, base_query, Agents.customer_code)
758
759 result = await session.execute(filtered_query)
760 agent = result.scalars().first()
761
762 if not agent:
783 - raise HTTPException(
784 - status_code=404,
785 - detail=f"Agent with agent_id {agent_id} not found or access denied"
786 - )
763 + raise HTTPException(status_code=404, detail=f"Agent with agent_id {agent_id} not found or access denied")
764
765 return await collect_agent_sca_policy_results(agent_id, policy_id)
766
@@ -816,18 +793,13 @@ async def get_agent_sca_policy_results_csv(
793
794 # Check customer access - first verify user has access to this agent
795 base_query = select(Agents).filter(Agents.agent_id == agent_id)
819 - filtered_query = await customer_access_handler.filter_query_by_customer_access(
820 - current_user, session, base_query, Agents.customer_code
821 - )
796 + filtered_query = await customer_access_handler.filter_query_by_customer_access(current_user, session, base_query, Agents.customer_code)
797
798 result = await session.execute(filtered_query)
799 agent = result.scalars().first()
800
801 if not agent:
827 - raise HTTPException(
828 - status_code=404,
829 - detail=f"Agent with agent_id {agent_id} not found or access denied"
830 - )
802 + raise HTTPException(status_code=404, detail=f"Agent with agent_id {agent_id} not found or access denied")
803
804 sca_results = (await collect_agent_sca_policy_results(agent_id, policy_id)).sca_policy_results
805 # Create a CSV file
@@ -887,7 +859,7 @@ async def get_agent_sca_policy_results_csv(
859 async def get_agent_soc_cases(
860 agent_hostname: str,
861 current_user: User = Depends(AuthHandler().get_current_user),
890 - session: AsyncSession = Depends(get_db)
862 + session: AsyncSession = Depends(get_db),
863 ):
864 """
865 Fetches the SOC cases of a specific agent.
@@ -905,18 +877,13 @@ async def get_agent_soc_cases(
877
878 # Check customer access - first verify user has access to this agent
879 base_query = select(Agents).filter(Agents.hostname == agent_hostname)
908 - filtered_query = await customer_access_handler.filter_query_by_customer_access(
909 - current_user, session, base_query, Agents.customer_code
910 - )
880 + filtered_query = await customer_access_handler.filter_query_by_customer_access(current_user, session, base_query, Agents.customer_code)
881
882 result = await session.execute(filtered_query)
883 agent = result.scalars().first()
884
885 if not agent:
916 - raise HTTPException(
917 - status_code=404,
918 - detail=f"Agent with hostname {agent_hostname} not found or access denied"
919 - )
886 + raise HTTPException(status_code=404, detail=f"Agent with hostname {agent_hostname} not found or access denied")
887
888 return CaseOutResponse(
889 cases=await list_cases_by_asset_name(asset_name=agent_hostname, db=session),
@@ -998,17 +965,17 @@ async def update_agent(
965 # Check customer access - first find the agent
966 base_query = select(Agents).filter(Agents.agent_id == agent_id)
967 filtered_query = await customer_access_handler.filter_query_by_customer_access(
1001 - current_user, session, base_query, Agents.customer_code
968 + current_user,
969 + session,
970 + base_query,
971 + Agents.customer_code,
972 )
973
974 result = await session.execute(filtered_query)
975 agent = result.scalars().first()
976
977 if not agent:
1008 - raise HTTPException(
1009 - status_code=404,
1010 - detail=f"Agent with agent_id {agent_id} not found or access denied"
1011 - )
978 + raise HTTPException(status_code=404, detail=f"Agent with agent_id {agent_id} not found or access denied")
979
980 agent.velociraptor_id = velociraptor_id
981 await session.commit()
@@ -1052,18 +1019,13 @@ async def delete_agent(
1019
1020 # Check customer access - first find the agent to verify access
1021 base_query = select(Agents).filter(Agents.agent_id == agent_id)
1055 - filtered_query = await customer_access_handler.filter_query_by_customer_access(
1056 - current_user, session, base_query, Agents.customer_code
1057 - )
1022 + filtered_query = await customer_access_handler.filter_query_by_customer_access(current_user, session, base_query, Agents.customer_code)
1023
1024 result = await session.execute(filtered_query)
1025 agent = result.scalars().first()
1026
1027 if not agent:
1063 - raise HTTPException(
1064 - status_code=404,
1065 - detail=f"Agent with agent_id {agent_id} not found or access denied"
1066 - )
1028 + raise HTTPException(status_code=404, detail=f"Agent with agent_id {agent_id} not found or access denied")
1029
1030 await delete_agent_wazuh(agent_id)
1031 client_id = await fetch_velociraptor_id(db=session, agent_id=agent_id)
backend/app/agents/wazuh/schema/agents.py
+4 -1
@@ -101,7 +101,10 @@ class WazuhAgentScaPolicyResults(BaseModel):
101 description="Description of the issue",
102 )
103 id: int
104 - reason: str
104 + reason: Optional[str] = Field(
105 + "Reason not found",
106 + description="Reason for the issue",
107 + )
108 command: Optional[str] = Field(
109 "Command not found",
110 description="Command to run to fix the issue",
backend/app/auth/models/users.py
+4 -1
@@ -3,7 +3,8 @@ import random
3 import re
4 import string
5 from enum import Enum
6 -from typing import Optional, List
6 +from typing import List
7 +from typing import Optional
8
9 import bcrypt
10 from pydantic import BaseModel
@@ -21,6 +22,7 @@ class Role(SQLModel, table=True):
22
23 user: Optional["User"] = Relationship(back_populates="role")
24
25 +
26 class UserCustomerAccess(SQLModel, table=True):
27 __tablename__ = "user_customer_access"
28 id: Optional[int] = Field(primary_key=True)
@@ -31,6 +33,7 @@ class UserCustomerAccess(SQLModel, table=True):
33 # Relationships
34 user: "User" = Relationship(back_populates="customer_access")
35
36 +
37 class User(SQLModel, table=True):
38 id: Optional[int] = Field(primary_key=True)
39 username: str = Field(index=True, max_length=256)
backend/app/auth/routes/auth.py
+6 -8
@@ -11,10 +11,12 @@ from sqlalchemy.ext.asyncio import AsyncSession
11
12 from app.auth.models.users import PasswordReset
13 from app.auth.models.users import PasswordResetToken
14 +from app.auth.models.users import RoleEnum
15 from app.auth.models.users import User
16 from app.auth.models.users import UserInput
17 from app.auth.models.users import UserLogin
17 -from app.auth.schema.auth import Token, UpdateUserRoleRequest
18 +from app.auth.schema.auth import Token
19 +from app.auth.schema.auth import UpdateUserRoleRequest
20 from app.auth.schema.auth import UserLoginResponse
21 from app.auth.schema.auth import UserResponse
22 from app.auth.schema.user import UserBaseResponse
@@ -22,7 +24,6 @@ from app.auth.services.universal import delete_user
24 from app.auth.services.universal import find_user
25 from app.auth.services.universal import select_all_users
26 from app.auth.utils import AuthHandler
25 -from app.auth.models.users import RoleEnum
27 from app.db.db_session import get_db
28
29 ACCESS_TOKEN_EXPIRE_MINUTES = 1440
@@ -170,7 +171,7 @@ async def get_users(session: AsyncSession = Depends(get_db)):
171 "username": user.username,
172 "email": user.email,
173 "role_id": user.role_id,
173 - "role_name": user.role.name if user.role else None
174 + "role_name": user.role.name if user.role else None,
175 }
176 user_list.append(user_dict)
177
@@ -358,10 +359,7 @@ async def update_user_role_by_name(
359
360 role_name_lower = request.role_name.lower()
361 if role_name_lower not in role_mapping:
361 - raise HTTPException(
362 - status_code=400,
363 - detail=f"Invalid role name. Valid roles are: {list(role_mapping.keys())}"
364 - )
362 + raise HTTPException(status_code=400, detail=f"Invalid role name. Valid roles are: {list(role_mapping.keys())}")
363
364 role_id = role_mapping[role_name_lower]
365
@@ -375,5 +373,5 @@ async def update_user_role_by_name(
373 "success": True,
374 "user_id": user_id,
375 "new_role_name": request.role_name,
378 - "new_role_id": role_id
376 + "new_role_id": role_id,
377 }
backend/app/auth/routes/customer_users.py
+19 -33
@@ -1,73 +1,59 @@
1 -# Create new file: app/auth/routes/customer_users.py
2 -from fastapi import APIRouter, Depends, HTTPException
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
5 -from sqlalchemy import select, delete
8
9 +from app.auth.models.users import User
10 +from app.auth.models.users import UserCustomerAccess
11 from app.auth.utils import AuthHandler
8 -from app.auth.models.users import User, UserCustomerAccess, RoleEnum
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")),
19 - session: AsyncSession = Depends(get_db)
23 + session: AsyncSession = Depends(get_db),
24 ):
25 """Assign customer access to a user (admin only)"""
26
27 # Remove existing access
24 - await session.execute(
25 - delete(UserCustomerAccess).where(UserCustomerAccess.user_id == user_id)
26 - )
28 + await session.execute(delete(UserCustomerAccess).where(UserCustomerAccess.user_id == user_id))
29
30 # Add new access
31 for customer_code in customer_codes:
30 - access = UserCustomerAccess(
31 - user_id=user_id,
32 - customer_code=customer_code
33 - )
32 + access = UserCustomerAccess(user_id=user_id, customer_code=customer_code)
33 session.add(access)
34
35 await session.commit()
36
38 - return {
39 - "success": True,
40 - "message": f"Assigned {len(customer_codes)} customers to user {user_id}",
41 - "customer_codes": customer_codes
42 - }
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")),
48 - session: AsyncSession = Depends(get_db)
44 + session: AsyncSession = Depends(get_db),
45 ):
46 """Get customer codes accessible to user (admin only)"""
47
52 - result = await session.execute(
53 - select(UserCustomerAccess.customer_code).where(UserCustomerAccess.user_id == user_id)
54 - )
48 + result = await session.execute(select(UserCustomerAccess.customer_code).where(UserCustomerAccess.user_id == user_id))
49 customer_codes = result.scalars().all()
50
57 - return {
58 - "success": True,
59 - "customer_codes": customer_codes
60 - }
51 + return {"success": True, "customer_codes": customer_codes}
52 +
53
54 @customer_users_router.get("/me/customers")
63 -async def get_my_customer_access(
64 - current_user: User = Depends(AuthHandler().get_current_user),
65 - session: AsyncSession = Depends(get_db)
66 -):
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
70 - return {
71 - "success": True,
72 - "customer_codes": customer_codes
73 - }
59 + return {"success": True, "customer_codes": customer_codes}
backend/app/auth/schema/auth.py
+1
@@ -20,5 +20,6 @@ class Token(BaseModel):
20 class TokenData(BaseModel):
21 username: str | None = None
22
23 +
24 class UpdateUserRoleRequest(BaseModel):
25 role_name: str
backend/app/auth/services/universal.py
+3 -1
@@ -1,10 +1,12 @@
1 from typing import List
2 +
3 from fastapi import HTTPException
4 from loguru import logger
5
6 # ! New with Async
7 from sqlalchemy.ext.asyncio import AsyncSession
7 -from sqlalchemy.orm import Session, selectinload
8 +from sqlalchemy.orm import Session
9 +from sqlalchemy.orm import selectinload
10 from sqlmodel import select
11
12 from app.auth.models.users import Password
backend/app/incidents/routes/db_operations.py
+61 -105
@@ -14,7 +14,9 @@ from loguru import logger
14 from sqlalchemy import select
15 from sqlalchemy.ext.asyncio import AsyncSession
16
17 +from app.auth.models.users import User
18 from app.auth.services.universal import select_all_users
19 +from app.auth.utils import AuthHandler
20 from app.connectors.wazuh_indexer.utils.universal import (
21 get_available_indices_via_source,
22 )
@@ -27,8 +29,8 @@ from app.data_store.data_store_operations import (
29 from app.db.db_session import get_db
30 from app.db.universal_models import Customers
31 from app.incidents.models import Alert
30 -from app.incidents.models import FieldName
32 from app.incidents.models import Comment
33 +from app.incidents.models import FieldName
34 from app.incidents.schema.db_operations import AlertContextCreate
35 from app.incidents.schema.db_operations import AlertContextResponse
36 from app.incidents.schema.db_operations import AlertCreate
@@ -87,6 +89,7 @@ from app.incidents.schema.db_operations import UpdateCaseStatus
89 from app.incidents.schema.incident_alert import CreatedAlertPayload
90 from app.incidents.schema.incident_alert import CreatedCaseNotificationPayload
91
92 +# from app.incidents.services.db_operations import list_alerts
93 # from app.incidents.services.db_operations import alerts_open_multiple_filters
94 # from app.incidents.services.db_operations import alerts_in_progress_multiple_filters
95 # from app.incidents.services.db_operations import alerts_closed_multiple_filters
@@ -98,11 +101,13 @@ from app.incidents.services.db_operations import add_timefield_name
101 from app.incidents.services.db_operations import alert_total
102 from app.incidents.services.db_operations import alert_total_by_alert_title
103 from app.incidents.services.db_operations import alert_total_by_assest_name
104 +from app.incidents.services.db_operations import alert_total_by_customer_codes
105 from app.incidents.services.db_operations import alerts_closed
106 from app.incidents.services.db_operations import alerts_closed_by_alert_title
107 from app.incidents.services.db_operations import alerts_closed_by_asset_name
108 from app.incidents.services.db_operations import alerts_closed_by_assigned_to
109 from app.incidents.services.db_operations import alerts_closed_by_customer_code
110 +from app.incidents.services.db_operations import alerts_closed_by_customer_codes
111 from app.incidents.services.db_operations import alerts_closed_by_ioc
112 from app.incidents.services.db_operations import alerts_closed_by_source
113 from app.incidents.services.db_operations import alerts_closed_by_tag
@@ -111,6 +116,7 @@ from app.incidents.services.db_operations import alerts_in_progress_by_alert_tit
116 from app.incidents.services.db_operations import alerts_in_progress_by_assest_name
117 from app.incidents.services.db_operations import alerts_in_progress_by_assigned_to
118 from app.incidents.services.db_operations import alerts_in_progress_by_customer_code
119 +from app.incidents.services.db_operations import alerts_in_progress_by_customer_codes
120 from app.incidents.services.db_operations import alerts_in_progress_by_ioc
121 from app.incidents.services.db_operations import alerts_in_progress_by_source
122 from app.incidents.services.db_operations import alerts_in_progress_by_tag
@@ -119,6 +125,7 @@ from app.incidents.services.db_operations import alerts_open_by_alert_title
125 from app.incidents.services.db_operations import alerts_open_by_assest_name
126 from app.incidents.services.db_operations import alerts_open_by_assigned_to
127 from app.incidents.services.db_operations import alerts_open_by_customer_code
128 +from app.incidents.services.db_operations import alerts_open_by_customer_codes
129 from app.incidents.services.db_operations import alerts_open_by_ioc
130 from app.incidents.services.db_operations import alerts_open_by_source
131 from app.incidents.services.db_operations import alerts_open_by_tag
@@ -168,19 +175,20 @@ from app.incidents.services.db_operations import increment_case_notification_cou
175 from app.incidents.services.db_operations import is_alert_linked_to_case
176 from app.incidents.services.db_operations import list_alert_by_assigned_to
177 from app.incidents.services.db_operations import list_alert_by_status
171 -from app.incidents.services.db_operations import list_alerts
178 from app.incidents.services.db_operations import list_alerts_by_asset_name
179 from app.incidents.services.db_operations import list_alerts_by_customer_code
180 from app.incidents.services.db_operations import list_alerts_by_ioc
181 from app.incidents.services.db_operations import list_alerts_by_source
182 from app.incidents.services.db_operations import list_alerts_by_tag
183 from app.incidents.services.db_operations import list_alerts_by_title
184 +from app.incidents.services.db_operations import list_alerts_for_user
185 from app.incidents.services.db_operations import list_alerts_multiple_filters
186 from app.incidents.services.db_operations import list_all_files
187 from app.incidents.services.db_operations import list_cases_by_asset_name
188 from app.incidents.services.db_operations import list_cases_by_assigned_to
189 from app.incidents.services.db_operations import list_cases_by_customer_code
190 from app.incidents.services.db_operations import list_cases_by_status
191 +from app.incidents.services.db_operations import list_cases_for_user
192 from app.incidents.services.db_operations import list_files_by_case_id
193 from app.incidents.services.db_operations import put_customer_notification
194 from app.incidents.services.db_operations import replace_alert_title_name
@@ -200,14 +208,6 @@ from app.incidents.services.db_operations import upload_report_template_to_data_
208 from app.incidents.services.db_operations import validate_source_exists
209 from app.incidents.services.incident_case import handle_customer_notifications_case
210 from app.middleware.customer_access import customer_access_handler
203 -from app.auth.models.users import User
204 -from app.auth.utils import AuthHandler
205 -from app.incidents.services.db_operations import alert_total_by_customer_codes
206 -from app.incidents.services.db_operations import alerts_closed_by_customer_codes
207 -from app.incidents.services.db_operations import alerts_in_progress_by_customer_codes
208 -from app.incidents.services.db_operations import alerts_open_by_customer_codes
209 -from app.incidents.services.db_operations import list_alerts_for_user
210 -from app.incidents.services.db_operations import list_cases_for_user
211
212 incidents_db_operations_router = APIRouter()
213
@@ -420,17 +420,14 @@ async def update_alert_status_endpoint(alert_status: UpdateAlertStatus, db: Asyn
420 async def create_comment_endpoint(
421 comment: CommentCreate,
422 current_user: User = Depends(AuthHandler().get_current_user),
423 - db: AsyncSession = Depends(get_db)
423 + db: AsyncSession = Depends(get_db),
424 ):
425 # Get the alert to check customer access
426 alert = await get_alert_by_id(comment.alert_id, db)
427
428 # Check if user has access to this alert's customer
429 if not await customer_access_handler.check_customer_access(current_user, alert.customer_code, db):
430 - raise HTTPException(
431 - status_code=403,
432 - detail=f"Access denied to alert {comment.alert_id} - insufficient customer permissions"
433 - )
430 + raise HTTPException(status_code=403, detail=f"Access denied to alert {comment.alert_id} - insufficient customer permissions")
431
432 return CommentResponse(comment=await create_comment(comment, db), success=True, message="Comment created successfully")
433
@@ -439,17 +436,14 @@ async def create_comment_endpoint(
436 async def edit_comment_endpoint(
437 comment: CommentEdit,
438 current_user: User = Depends(AuthHandler().get_current_user),
442 - db: AsyncSession = Depends(get_db)
439 + db: AsyncSession = Depends(get_db),
440 ):
441 # Get the alert to check customer access
442 alert = await get_alert_by_id(comment.alert_id, db)
443
444 # Check if user has access to this alert's customer
445 if not await customer_access_handler.check_customer_access(current_user, alert.customer_code, db):
449 - raise HTTPException(
450 - status_code=403,
451 - detail=f"Access denied to alert {comment.alert_id} - insufficient customer permissions"
452 - )
446 + raise HTTPException(status_code=403, detail=f"Access denied to alert {comment.alert_id} - insufficient customer permissions")
447
448 return CommentResponse(comment=await edit_comment(comment, db), success=True, message="Comment edited successfully")
449
@@ -458,7 +452,7 @@ async def edit_comment_endpoint(
452 async def delete_comment_endpoint(
453 comment_id: int,
454 current_user: User = Depends(AuthHandler().get_current_user),
461 - db: AsyncSession = Depends(get_db)
455 + db: AsyncSession = Depends(get_db),
456 ):
457 # First get the comment to find the alert_id
458 result = await db.execute(select(Comment).where(Comment.id == comment_id))
@@ -473,7 +467,7 @@ async def delete_comment_endpoint(
467 if not await customer_access_handler.check_customer_access(current_user, alert.customer_code, db):
468 raise HTTPException(
469 status_code=403,
476 - detail=f"Access denied to comment on alert {comment.alert_id} - insufficient customer permissions"
470 + detail=f"Access denied to comment on alert {comment.alert_id} - insufficient customer permissions",
471 )
472
473 await delete_comment(comment_id, db)
@@ -560,7 +554,7 @@ async def list_alerts_by_ioc_value_endpoint(
554 db=db,
555 page=page,
556 page_size=page_size,
563 - order="desc"
557 + order="desc",
558 )
559 total = await alert_total_by_customer_codes(db, accessible_customers)
560 open_alerts = await alerts_open_by_customer_codes(db, accessible_customers)
@@ -621,7 +615,7 @@ async def list_alerts_by_tag_endpoint(
615 db=db,
616 page=page,
617 page_size=page_size,
624 - order="desc"
618 + order="desc",
619 )
620 total = await alert_total_by_customer_codes(db, accessible_customers)
621 open_alerts = await alerts_open_by_customer_codes(db, accessible_customers)
@@ -705,6 +699,7 @@ async def create_case_from_alert_endpoint(alert_id: CaseCreateFromAlert, db: Asy
699 # message="Alerts retrieved successfully",
700 # )
701
702 +
703 @incidents_db_operations_router.get("/alerts", response_model=AlertOutResponse)
704 async def list_alerts_endpoint(
705 page: int = Query(1, ge=1),
@@ -744,11 +739,12 @@ async def list_alerts_endpoint(
739 message="Alerts retrieved successfully",
740 )
741
742 +
743 @incidents_db_operations_router.get("/alert/{alert_id}", response_model=AlertOutResponse)
744 async def get_alert_by_id_endpoint(
745 alert_id: int,
746 current_user: User = Depends(AuthHandler().get_current_user),
751 - db: AsyncSession = Depends(get_db)
747 + db: AsyncSession = Depends(get_db),
748 ):
749 """Get alert by ID with customer access validation"""
750 logger.info(f"Getting alert {alert_id} for user: {current_user.username} with role_id: {current_user.role_id}")
@@ -758,10 +754,7 @@ async def get_alert_by_id_endpoint(
754
755 # Check if user has access to this alert's customer
756 if not await customer_access_handler.check_customer_access(current_user, alert.customer_code, db):
761 - raise HTTPException(
762 - status_code=403,
763 - detail=f"Access denied to alert {alert_id} - insufficient customer permissions"
764 - )
757 + raise HTTPException(status_code=403, detail=f"Access denied to alert {alert_id} - insufficient customer permissions")
758
759 return AlertOutResponse(alerts=[alert], success=True, message="Alert retrieved successfully")
760
@@ -770,7 +763,7 @@ async def get_alert_by_id_endpoint(
763 async def delete_alert_endpoint(
764 alert_id: int,
765 current_user: User = Depends(AuthHandler().get_current_user),
773 - db: AsyncSession = Depends(get_db)
766 + db: AsyncSession = Depends(get_db),
767 ):
768 """Delete alert with customer access validation"""
769 logger.info(f"Deleting alert {alert_id} for user: {current_user.username} with role_id: {current_user.role_id}")
@@ -780,10 +773,7 @@ async def delete_alert_endpoint(
773
774 # Check if user has access to this alert's customer
775 if not await customer_access_handler.check_customer_access(current_user, alert.customer_code, db):
783 - raise HTTPException(
784 - status_code=403,
785 - detail=f"Access denied to alert {alert_id} - insufficient customer permissions"
786 - )
776 + raise HTTPException(status_code=403, detail=f"Access denied to alert {alert_id} - insufficient customer permissions")
777
778 await is_alert_linked_to_case(alert_id, db)
779 await delete_alert(alert_id, db)
@@ -863,7 +853,7 @@ async def list_alerts_by_status_endpoint(
853 db=db,
854 page=page,
855 page_size=page_size,
866 - order=order
856 + order=order,
857 )
858 total = await alert_total_by_customer_codes(db, accessible_customers)
859 open_alerts = await alerts_open_by_customer_codes(db, accessible_customers)
@@ -911,7 +901,7 @@ async def list_alerts_by_assigned_to_endpoint(
901 db=db,
902 page=page,
903 page_size=page_size,
914 - order=order
904 + order=order,
905 )
906 total = await alert_total_by_customer_codes(db, accessible_customers)
907 open_alerts = await alerts_open_by_customer_codes(db, accessible_customers)
@@ -959,7 +949,7 @@ async def list_alerts_by_asset_name_endpoint(
949 db=db,
950 page=page,
951 page_size=page_size,
962 - order=order
952 + order=order,
953 )
954 total = await alert_total_by_customer_codes(db, accessible_customers)
955 open_alerts = await alerts_open_by_customer_codes(db, accessible_customers)
@@ -1007,7 +997,7 @@ async def list_alerts_by_title_endpoint(
997 db=db,
998 page=page,
999 page_size=page_size,
1010 - order=order
1000 + order=order,
1001 )
1002 total = await alert_total_by_customer_codes(db, accessible_customers)
1003 open_alerts = await alerts_open_by_customer_codes(db, accessible_customers)
@@ -1043,6 +1033,7 @@ async def list_alerts_by_title_endpoint(
1033 # message="Alerts retrieved successfully",
1034 # )
1035
1036 +
1037 @incidents_db_operations_router.get("/alerts/customer/{customer_code}", response_model=AlertOutResponse)
1038 async def list_alerts_by_customer_code_endpoint(
1039 customer_code: str,
@@ -1098,7 +1089,7 @@ async def list_alerts_by_source_endpoint(
1089 db=db,
1090 page=page,
1091 page_size=page_size,
1101 - order=order
1092 + order=order,
1093 )
1094 total = await alert_total_by_customer_codes(db, accessible_customers)
1095 open_alerts = await alerts_open_by_customer_codes(db, accessible_customers)
@@ -1168,10 +1159,7 @@ async def list_alerts_multiple_filters_endpoint(
1159 if "*" not in accessible_customers:
1160 # If user provided customer_code, validate they have access to it
1161 if customer_code and customer_code not in accessible_customers:
1171 - raise HTTPException(
1172 - status_code=403,
1173 - detail=f"Access denied to customer {customer_code}"
1174 - )
1162 + raise HTTPException(status_code=403, detail=f"Access denied to customer {customer_code}")
1163
1164 # If no customer_code specified, use the first accessible customer for single customer users
1165 # For multi-customer users, we'll need to modify the query to handle multiple customers
@@ -1242,10 +1230,7 @@ async def list_alerts_multiple_filters_endpoint(
1230
1231
1232 @incidents_db_operations_router.get("/cases", response_model=CaseOutResponse)
1245 -async def list_cases_endpoint(
1246 - current_user: User = Depends(AuthHandler().get_current_user),
1247 - db: AsyncSession = Depends(get_db)
1248 -):
1233 +async def list_cases_endpoint(current_user: User = Depends(AuthHandler().get_current_user), db: AsyncSession = Depends(get_db)):
1234 """List cases with automatic customer filtering"""
1235 logger.info(f"Listing cases for user: {current_user.username} with role_id: {current_user.role_id}")
1236
@@ -1257,7 +1242,7 @@ async def list_cases_endpoint(
1242 async def update_case_status_endpoint(
1243 case_status: UpdateCaseStatus,
1244 current_user: User = Depends(AuthHandler().get_current_user),
1260 - db: AsyncSession = Depends(get_db)
1245 + db: AsyncSession = Depends(get_db),
1246 ):
1247 """Update case status with customer access validation"""
1248 logger.info(f"Updating case {case_status.case_id} status for user: {current_user.username} with role_id: {current_user.role_id}")
@@ -1267,10 +1252,7 @@ async def update_case_status_endpoint(
1252
1253 # Check if user has access to this case's customer
1254 if not await customer_access_handler.check_customer_access(current_user, case.customer_code, db):
1270 - raise HTTPException(
1271 - status_code=403,
1272 - detail=f"Access denied to case {case_status.case_id} - insufficient customer permissions"
1273 - )
1255 + raise HTTPException(status_code=403, detail=f"Access denied to case {case_status.case_id} - insufficient customer permissions")
1256
1257 # Update the case status
1258 await update_case_status(case_status, db)
@@ -1284,7 +1266,7 @@ async def update_case_status_endpoint(
1266 async def update_case_assigned_to_endpoint(
1267 assigned_to: AssignedToCase,
1268 current_user: User = Depends(AuthHandler().get_current_user),
1287 - db: AsyncSession = Depends(get_db)
1269 + db: AsyncSession = Depends(get_db),
1270 ):
1271 """Update case assigned_to with customer access validation"""
1272 logger.info(f"Updating case {assigned_to.case_id} assigned_to for user: {current_user.username} with role_id: {current_user.role_id}")
@@ -1294,10 +1276,7 @@ async def update_case_assigned_to_endpoint(
1276
1277 # Check if user has access to this case's customer
1278 if not await customer_access_handler.check_customer_access(current_user, case.customer_code, db):
1297 - raise HTTPException(
1298 - status_code=403,
1299 - detail=f"Access denied to case {assigned_to.case_id} - insufficient customer permissions"
1300 - )
1279 + raise HTTPException(status_code=403, detail=f"Access denied to case {assigned_to.case_id} - insufficient customer permissions")
1280
1281 all_users = await select_all_users()
1282 user_names = [user.username for user in all_users]
@@ -1314,12 +1293,14 @@ async def update_case_assigned_to_endpoint(
1293 success=True,
1294 message="Case assigned to user successfully",
1295 )
1296 +
1297 +
1298 @incidents_db_operations_router.put("/case/customer-code", response_model=CaseOutResponse)
1299 async def update_case_customer_code_endpoint(
1300 case_id: int,
1301 customer_code: str,
1302 current_user: User = Depends(AuthHandler().get_current_user),
1322 - db: AsyncSession = Depends(get_db)
1303 + db: AsyncSession = Depends(get_db),
1304 ):
1305 """Update case customer_code with customer access validation"""
1306 logger.info(f"Updating case {case_id} customer_code for user: {current_user.username} with role_id: {current_user.role_id}")
@@ -1329,18 +1310,12 @@ async def update_case_customer_code_endpoint(
1310
1311 # Check if user has access to the current case's customer
1312 if not await customer_access_handler.check_customer_access(current_user, case.customer_code, db):
1332 - raise HTTPException(
1333 - status_code=403,
1334 - detail=f"Access denied to case {case_id} - insufficient customer permissions"
1335 - )
1313 + raise HTTPException(status_code=403, detail=f"Access denied to case {case_id} - insufficient customer permissions")
1314
1315 # Also check if user has access to the new customer code (for non-admin users)
1316 accessible_customers = await customer_access_handler.get_user_accessible_customers(current_user, db)
1317 if "*" not in accessible_customers and customer_code not in accessible_customers:
1340 - raise HTTPException(
1341 - status_code=403,
1342 - detail=f"Access denied - cannot assign case to customer {customer_code}"
1343 - )
1318 + raise HTTPException(status_code=403, detail=f"Access denied - cannot assign case to customer {customer_code}")
1319
1320 # Update the case customer code
1321 await update_case_customer_code(case_id, customer_code, db)
@@ -1358,7 +1333,7 @@ async def update_case_customer_code_endpoint(
1333 async def delete_case_endpoint(
1334 case_id: int,
1335 current_user: User = Depends(AuthHandler().get_current_user),
1361 - db: AsyncSession = Depends(get_db)
1336 + db: AsyncSession = Depends(get_db),
1337 ):
1338 """Delete case with customer access validation"""
1339 logger.info(f"Deleting case {case_id} for user: {current_user.username} with role_id: {current_user.role_id}")
@@ -1368,10 +1343,7 @@ async def delete_case_endpoint(
1343
1344 # Check if user has access to this case's customer
1345 if not await customer_access_handler.check_customer_access(current_user, case.customer_code, db):
1371 - raise HTTPException(
1372 - status_code=403,
1373 - detail=f"Access denied to case {case_id} - insufficient customer permissions"
1374 - )
1346 + raise HTTPException(status_code=403, detail=f"Access denied to case {case_id} - insufficient customer permissions")
1347
1348 await delete_case(case_id, db)
1349 return {"message": "Case deleted successfully", "success": True}
@@ -1381,7 +1353,7 @@ async def delete_case_endpoint(
1353 async def list_cases_by_status_endpoint(
1354 status: AlertStatus,
1355 current_user: User = Depends(AuthHandler().get_current_user),
1384 - db: AsyncSession = Depends(get_db)
1356 + db: AsyncSession = Depends(get_db),
1357 ):
1358 """List cases by status with customer access filtering"""
1359 if status not in AlertStatus:
@@ -1408,7 +1380,7 @@ async def list_cases_by_status_endpoint(
1380 async def list_cases_by_assigned_to_endpoint(
1381 assigned_to: str,
1382 current_user: User = Depends(AuthHandler().get_current_user),
1411 - db: AsyncSession = Depends(get_db)
1383 + db: AsyncSession = Depends(get_db),
1384 ):
1385 """List cases by assigned user with customer access filtering"""
1386 logger.info(f"Listing cases assigned to {assigned_to} for user: {current_user.username} with role_id: {current_user.role_id}")
@@ -1431,7 +1403,7 @@ async def list_cases_by_assigned_to_endpoint(
1403 async def list_cases_by_asset_name_endpoint(
1404 asset_name: str,
1405 current_user: User = Depends(AuthHandler().get_current_user),
1434 - db: AsyncSession = Depends(get_db)
1406 + db: AsyncSession = Depends(get_db),
1407 ):
1408 """List cases by asset name with customer access filtering"""
1409 logger.info(f"Listing cases by asset {asset_name} for user: {current_user.username} with role_id: {current_user.role_id}")
@@ -1460,7 +1432,7 @@ async def list_cases_by_asset_name_endpoint(
1432 async def list_cases_by_customer_code_endpoint(
1433 customer_code: str,
1434 current_user: User = Depends(customer_access_handler.require_customer_access()),
1463 - db: AsyncSession = Depends(get_db)
1435 + db: AsyncSession = Depends(get_db),
1436 ):
1437 """List cases for specific customer (with access validation)"""
1438 logger.info(f"Listing cases for customer {customer_code} for user: {current_user.username} with role_id: {current_user.role_id}")
@@ -1482,7 +1454,7 @@ async def list_all_case_data_store_files_endpoint(db: AsyncSession = Depends(get
1454 async def list_case_data_store_files_endpoint(
1455 case_id: int,
1456 current_user: User = Depends(AuthHandler().get_current_user),
1485 - db: AsyncSession = Depends(get_db)
1457 + db: AsyncSession = Depends(get_db),
1458 ):
1459 """List case data store files with customer access validation"""
1460 logger.info(f"Listing files for case {case_id} for user: {current_user.username} with role_id: {current_user.role_id}")
@@ -1492,10 +1464,7 @@ async def list_case_data_store_files_endpoint(
1464
1465 # Check if user has access to this case's customer
1466 if not await customer_access_handler.check_customer_access(current_user, case.customer_code, db):
1495 - raise HTTPException(
1496 - status_code=403,
1497 - detail=f"Access denied to case {case_id} - insufficient customer permissions"
1498 - )
1467 + raise HTTPException(status_code=403, detail=f"Access denied to case {case_id} - insufficient customer permissions")
1468
1469 return ListCaseDataStoreResponse(
1470 case_data_store=await list_files_by_case_id(case_id, db),
@@ -1509,7 +1478,7 @@ async def download_case_data_store_file_endpoint(
1478 case_id: int,
1479 file_name: str,
1480 current_user: User = Depends(AuthHandler().get_current_user),
1512 - db: AsyncSession = Depends(get_db)
1481 + db: AsyncSession = Depends(get_db),
1482 ) -> StreamingResponse:
1483 """Download case data store file with customer access validation"""
1484 logger.info(f"Downloading file {file_name} from case {case_id} for user: {current_user.username} with role_id: {current_user.role_id}")
@@ -1519,10 +1488,7 @@ async def download_case_data_store_file_endpoint(
1488
1489 # Check if user has access to this case's customer
1490 if not await customer_access_handler.check_customer_access(current_user, case.customer_code, db):
1522 - raise HTTPException(
1523 - status_code=403,
1524 - detail=f"Access denied to case {case_id} - insufficient customer permissions"
1525 - )
1491 + raise HTTPException(status_code=403, detail=f"Access denied to case {case_id} - insufficient customer permissions")
1492
1493 file_bytes, file_content_type = await download_file_from_case(case_id, file_name, db)
1494 logger.info(f"Streaming file {file_name} from case {case_id}")
@@ -1547,10 +1513,7 @@ async def upload_case_data_store_endpoint(
1513
1514 # Check if user has access to this case's customer
1515 if not await customer_access_handler.check_customer_access(current_user, case.customer_code, db):
1550 - raise HTTPException(
1551 - status_code=403,
1552 - detail=f"Access denied to case {case_id} - insufficient customer permissions"
1553 - )
1516 + raise HTTPException(status_code=403, detail=f"Access denied to case {case_id} - insufficient customer permissions")
1517
1518 if await file_exists(case_id, file.filename, db):
1519 raise HTTPException(status_code=400, detail="File name already exists for this case")
@@ -1567,7 +1530,7 @@ async def delete_case_data_store_file_endpoint(
1530 case_id: int,
1531 file_name: str,
1532 current_user: User = Depends(AuthHandler().get_current_user),
1570 - db: AsyncSession = Depends(get_db)
1533 + db: AsyncSession = Depends(get_db),
1534 ):
1535 """Delete case data store file with customer access validation"""
1536 logger.info(f"Deleting file {file_name} from case {case_id} for user: {current_user.username} with role_id: {current_user.role_id}")
@@ -1577,10 +1540,7 @@ async def delete_case_data_store_file_endpoint(
1540
1541 # Check if user has access to this case's customer
1542 if not await customer_access_handler.check_customer_access(current_user, case.customer_code, db):
1580 - raise HTTPException(
1581 - status_code=403,
1582 - detail=f"Access denied to case {case_id} - insufficient customer permissions"
1583 - )
1543 + raise HTTPException(status_code=403, detail=f"Access denied to case {case_id} - insufficient customer permissions")
1544
1545 await delete_file_from_case(case_id, file_name, db)
1546 return {"message": "File deleted successfully", "success": True}
@@ -1590,7 +1550,7 @@ async def delete_case_data_store_file_endpoint(
1550 async def get_case_by_id_endpoint(
1551 case_id: int,
1552 current_user: User = Depends(AuthHandler().get_current_user),
1593 - db: AsyncSession = Depends(get_db)
1553 + db: AsyncSession = Depends(get_db),
1554 ):
1555 """Get case by ID with customer access validation"""
1556 logger.info(f"Getting case {case_id} for user: {current_user.username} with role_id: {current_user.role_id}")
@@ -1600,10 +1560,7 @@ async def get_case_by_id_endpoint(
1560
1561 # Check if user has access to this case's customer
1562 if not await customer_access_handler.check_customer_access(current_user, case.customer_code, db):
1603 - raise HTTPException(
1604 - status_code=403,
1605 - detail=f"Access denied to case {case_id} - insufficient customer permissions"
1606 - )
1563 + raise HTTPException(status_code=403, detail=f"Access denied to case {case_id} - insufficient customer permissions")
1564
1565 return CaseOutResponse(cases=[case], success=True, message="Case retrieved successfully")
1566
@@ -1612,7 +1569,7 @@ async def get_case_by_id_endpoint(
1569 async def create_case_notification_endpoint(
1570 request: CaseNotificationCreate,
1571 current_user: User = Depends(AuthHandler().get_current_user),
1615 - db: AsyncSession = Depends(get_db)
1572 + db: AsyncSession = Depends(get_db),
1573 ):
1574 """
1575 Create case notification with customer access validation.
@@ -1627,16 +1584,15 @@ async def create_case_notification_endpoint(
1584 Returns:
1585 CaseNotificationResponse: The response object containing the created case notification.
1586 """
1630 - logger.info(f"Creating case notification for case {request.case_id} for user: {current_user.username} with role_id: {current_user.role_id}")
1587 + logger.info(
1588 + f"Creating case notification for case {request.case_id} for user: {current_user.username} with role_id: {current_user.role_id}",
1589 + )
1590
1591 case_details = await get_case_by_id(request.case_id, db)
1592
1593 # Check if user has access to this case's customer
1594 if not await customer_access_handler.check_customer_access(current_user, case_details.customer_code, db):
1636 - raise HTTPException(
1637 - status_code=403,
1638 - detail=f"Access denied to case {request.case_id} - insufficient customer permissions"
1639 - )
1595 + raise HTTPException(status_code=403, detail=f"Access denied to case {request.case_id} - insufficient customer permissions")
1596
1597 case_notification_payload = CreatedCaseNotificationPayload(
1598 case_name=case_details.case_name,
backend/app/incidents/services/db_operations.py
+12 -24
@@ -21,6 +21,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
21 from sqlalchemy.future import select
22 from sqlalchemy.orm import selectinload
23
24 +from app.auth.models.users import User
25 from app.data_store.data_store_operations import delete_file
26 from app.data_store.data_store_operations import download_data_store
27 from app.data_store.data_store_operations import upload_case_data_store
@@ -74,7 +75,6 @@ from app.integrations.alert_creation_settings.models.alert_creation_settings imp
75 AlertCreationSettings,
76 )
77 from app.middleware.customer_access import customer_access_handler
77 -from app.auth.models.users import User
78
79
80 async def customer_code_valid(customer_code: str, db: AsyncSession) -> bool:
@@ -209,6 +209,7 @@ async def alerts_open_by_source(db: AsyncSession, source: str) -> int:
209 result = await db.execute(select(Alert).where((Alert.status == "OPEN") & (Alert.source == source)))
210 return len(result.scalars().all())
211
212 +
213 async def alert_total_by_customer_codes(db: AsyncSession, customer_codes: List[str]) -> int:
214 """Get total alerts for multiple customer codes"""
215 result = await db.execute(select(Alert).where(Alert.customer_code.in_(customer_codes)))
@@ -217,33 +218,22 @@ async def alert_total_by_customer_codes(db: AsyncSession, customer_codes: List[s
218
219 async def alerts_closed_by_customer_codes(db: AsyncSession, customer_codes: List[str]) -> int:
220 """Get closed alerts for multiple customer codes"""
220 - result = await db.execute(
221 - select(Alert).where(
222 - (Alert.status == "CLOSED") & (Alert.customer_code.in_(customer_codes))
223 - )
224 - )
221 + result = await db.execute(select(Alert).where((Alert.status == "CLOSED") & (Alert.customer_code.in_(customer_codes))))
222 return len(result.scalars().all())
223
224
225 async def alerts_in_progress_by_customer_codes(db: AsyncSession, customer_codes: List[str]) -> int:
226 """Get in-progress alerts for multiple customer codes"""
230 - result = await db.execute(
231 - select(Alert).where(
232 - (Alert.status == "IN_PROGRESS") & (Alert.customer_code.in_(customer_codes))
233 - )
234 - )
227 + result = await db.execute(select(Alert).where((Alert.status == "IN_PROGRESS") & (Alert.customer_code.in_(customer_codes))))
228 return len(result.scalars().all())
229
230
231 async def alerts_open_by_customer_codes(db: AsyncSession, customer_codes: List[str]) -> int:
232 """Get open alerts for multiple customer codes"""
240 - result = await db.execute(
241 - select(Alert).where(
242 - (Alert.status == "OPEN") & (Alert.customer_code.in_(customer_codes))
243 - )
244 - )
233 + result = await db.execute(select(Alert).where((Alert.status == "OPEN") & (Alert.customer_code.in_(customer_codes))))
234 return len(result.scalars().all())
235
236 +
237 async def alerts_total_multiple_filters(
238 db: AsyncSession,
239 assigned_to: Optional[str] = None,
@@ -867,8 +857,8 @@ async def create_comment(comment: CommentCreate, db: AsyncSession) -> Comment:
857
858 # Create comment with automatic timestamp if not provided
859 comment_data = comment.dict()
870 - if comment_data.get('created_at') is None:
871 - comment_data['created_at'] = datetime.utcnow()
860 + if comment_data.get("created_at") is None:
861 + comment_data["created_at"] = datetime.utcnow()
862
863 db_comment = Comment(**comment_data)
864 db.add(db_comment)
@@ -1975,6 +1965,7 @@ async def list_alerts_multiple_filters(
1965
1966 return alerts_out
1967
1968 +
1969 async def list_alerts_for_user(
1970 user: User,
1971 session: AsyncSession,
@@ -1992,9 +1983,7 @@ async def list_alerts_for_user(
1983 )
1984
1985 # Apply customer filtering
1995 - filtered_query = await customer_access_handler.filter_query_by_customer_access(
1996 - user, session, base_query, Alert.customer_code
1997 - )
1986 + filtered_query = await customer_access_handler.filter_query_by_customer_access(user, session, base_query, Alert.customer_code)
1987
1988 offset = (page - 1) * page_size
1989 order_by = asc(Alert.id) if order == "asc" else desc(Alert.id)
@@ -2030,6 +2019,7 @@ async def list_alerts_for_user(
2019
2020 return alerts_out
2021
2022 +
2023 async def list_cases_for_user(
2024 user: User,
2025 session: AsyncSession,
@@ -2045,9 +2035,7 @@ async def list_cases_for_user(
2035 )
2036
2037 # Apply customer filtering
2048 - filtered_query = await customer_access_handler.filter_query_by_customer_access(
2049 - user, session, base_query, Case.customer_code
2050 - )
2038 + filtered_query = await customer_access_handler.filter_query_by_customer_access(user, session, base_query, Case.customer_code)
2039
2040 result = await session.execute(filtered_query)
2041 cases = result.scalars().all()
backend/app/middleware/customer_access.py
+17 -25
@@ -1,16 +1,20 @@
1 # Create new file: app/middleware/customer_access.py
2 -from typing import List, Optional
3 -from fastapi import Depends, HTTPException
4 -from sqlalchemy.ext.asyncio import AsyncSession
2 +from typing import List
3 +from typing import Optional
4 +
5 +from fastapi import Depends
6 +from fastapi import HTTPException
7 from sqlalchemy import select
6 -from loguru import logger
8 +from sqlalchemy.ext.asyncio import AsyncSession
9
8 -from app.auth.models.users import User, UserCustomerAccess, RoleEnum
10 +from app.auth.models.users import RoleEnum
11 +from app.auth.models.users import User
12 +from app.auth.models.users import UserCustomerAccess
13 from app.auth.utils import AuthHandler
14 from app.db.db_session import get_db
15
12 -class CustomerAccessHandler:
16
17 +class CustomerAccessHandler:
18 async def get_user_accessible_customers(self, user: User, session: AsyncSession) -> List[str]:
19 """Get all customer codes accessible to a user"""
20 # Admin and analyst users have access to all customers
@@ -19,10 +23,7 @@ class CustomerAccessHandler:
23
24 # Customer users only see their assigned customers
25 if user.role_id == RoleEnum.customer_user:
22 - result = await session.execute(
23 - select(UserCustomerAccess.customer_code)
24 - .where(UserCustomerAccess.user_id == user.id)
25 - )
26 + result = await session.execute(select(UserCustomerAccess.customer_code).where(UserCustomerAccess.user_id == user.id))
27 return result.scalars().all()
28
29 return [] # No access by default
@@ -38,13 +39,7 @@ class CustomerAccessHandler:
39 # Specific customer access
40 return customer_code in accessible_customers
41
41 - async def filter_query_by_customer_access(
42 - self,
43 - user: User,
44 - session: AsyncSession,
45 - base_query,
46 - customer_code_field
47 - ):
42 + async def filter_query_by_customer_access(self, user: User, session: AsyncSession, base_query, customer_code_field):
43 """Filter any query by user's customer access"""
44 accessible_customers = await self.get_user_accessible_customers(user, session)
45
@@ -61,18 +56,15 @@ class CustomerAccessHandler:
56
57 def require_customer_access(self, customer_code: Optional[str] = None):
58 """FastAPI dependency to enforce customer access"""
64 - async def _check_access(
65 - current_user: User = Depends(AuthHandler().get_current_user),
66 - session: AsyncSession = Depends(get_db)
67 - ):
59 +
60 + async def _check_access(current_user: User = Depends(AuthHandler().get_current_user), session: AsyncSession = Depends(get_db)):
61 if customer_code:
62 if not await self.check_customer_access(current_user, customer_code, session):
70 - raise HTTPException(
71 - status_code=403,
72 - detail=f"Access denied to customer {customer_code}"
73 - )
63 + raise HTTPException(status_code=403, detail=f"Access denied to customer {customer_code}")
64 return current_user
65 +
66 return _check_access
67
68 +
69 # Create a singleton instance
70 customer_access_handler = CustomerAccessHandler()
customer_portal/src/router/index_old.ts deleted
-191
@@ -1,191 +0,0 @@
1 -import { createRouter, createWebHistory } from 'vue-router'
2 -import { useAuthStore } from '@/stores/auth'
3 -import LoginPage from '@/components/LoginPage.vue'
4 -// TODO: Add back when views are ready
5 -// import AlertsView from '@/views/AlertsView.vue'
6 -// import CasesView from '@/views/CasesView.vue'
7 -const routes = [
8 - {
9 - path: '/login',
10 - name: 'Login',
11 - component: LoginPage,
12 - meta: { requiresGuest: true }
13 - },
14 -
15 -const Dashboard = {
16 - template: `
17 - <div class="min-h-screen bg-gray-50">
18 - <header class="bg-white shadow">
19 - <div class="max-w-7xl mx-auto px-4 sm:px-6 lg:px-8">
20 - <div class="flex justify-between h-16">
21 - <div class="flex items-center">
22 - <h1 class="text-xl font-semibold">Customer Portal</h1>
23 - </div>
24 - <div class="flex items-center space-x-4">
25 - <span class="text-sm text-gray-700">{{ user?.username }}</span>
26 - <button
27 - @click="logout"
28 - class="bg-red-600 hover:bg-red-700 text-white px-3 py-2 rounded-md text-sm font-medium"
29 - >
30 - Logout
31 - </button>
32 - </div>
33 - </div>
34 - </div>
35 - </header>
36 - <main class="max-w-7xl mx-auto py-6 sm:px-6 lg:px-8">
37 - <div class="px-4 py-6 sm:px-0">
38 - <div class="grid grid-cols-1 md:grid-cols-2 gap-6">
39 - <div class="bg-white overflow-hidden shadow rounded-lg">
40 - <div class="p-5">
41 - <div class="flex items-center">
42 - <div class="flex-shrink-0">
43 - <div class="w-8 h-8 bg-blue-500 rounded-md flex items-center justify-center">
44 - <span class="text-white font-medium">A</span>
45 - </div>
46 - </div>
47 - <div class="ml-5 w-0 flex-1">
48 - <dl>
49 - <dt class="text-sm font-medium text-gray-500 truncate">
50 - Alerts
51 - </dt>
52 - <dd class="text-lg font-medium text-gray-900">
53 - View security alerts
54 - </dd>
55 - </dl>
56 - </div>
57 - </div>
58 - </div>
59 - <div class="bg-gray-50 px-5 py-3">
60 - <div class="text-sm">
61 - <router-link
62 - to="/alerts"
63 - class="font-medium text-blue-700 hover:text-blue-900"
64 - >
65 - View all alerts
66 - </router-link>
67 - </div>
68 - </div>
69 - </div>
70 -
71 - <div class="bg-white overflow-hidden shadow rounded-lg">
72 - <div class="p-5">
73 - <div class="flex items-center">
74 - <div class="flex-shrink-0">
75 - <div class="w-8 h-8 bg-green-500 rounded-md flex items-center justify-center">
76 - <span class="text-white font-medium">C</span>
77 - </div>
78 - </div>
79 - <div class="ml-5 w-0 flex-1">
80 - <dl>
81 - <dt class="text-sm font-medium text-gray-500 truncate">
82 - Cases
83 - </dt>
84 - <dd class="text-lg font-medium text-gray-900">
85 - View security cases
86 - </dd>
87 - </dl>
88 - </div>
89 - </div>
90 - </div>
91 - <div class="bg-gray-50 px-5 py-3">
92 - <div class="text-sm">
93 - <router-link
94 - to="/cases"
95 - class="font-medium text-green-700 hover:text-green-900"
96 - >
97 - View all cases
98 - </router-link>
99 - </div>
100 - </div>
101 - </div>
102 - </div>
103 - </div>
104 - </main>
105 - </div>
106 - `,
107 - computed: {
108 - user() {
109 - const authStore = useAuthStore()
110 - return authStore.user
111 - }
112 - },
113 - methods: {
114 - logout() {
115 - const authStore = useAuthStore()
116 - authStore.logout()
117 - this.$router.push('/login')
118 - }
119 - }
120 -}
121 -
122 -const NotFound = {
123 - template: `
124 - <div class="min-h-screen flex items-center justify-center bg-gray-50">
125 - <div class="text-center">
126 - <h1 class="text-4xl font-bold text-gray-900">404</h1>
127 - <p class="mt-2 text-lg text-gray-600">Page not found</p>
128 - <a href="#/" class="mt-4 inline-block bg-indigo-600 text-white px-4 py-2 rounded-md hover:bg-indigo-700">
129 - Go Home
130 - </a>
131 - </div>
132 - </div>
133 - `
134 -}
135 -
136 -const routes = [
137 - {
138 - path: '/login',
139 - name: 'Login',
140 - component: Login,
141 - meta: { requiresGuest: true }
142 - },
143 - {
144 - path: '/',
145 - name: 'Dashboard',
146 - component: Dashboard,
147 - meta: { requiresAuth: true }
148 - },
149 - // TODO: Add back when views are ready
150 - // {
151 - // path: '/alerts',
152 - // name: 'Alerts',
153 - // component: AlertsView,
154 - // meta: { requiresAuth: true }
155 - // },
156 - // {
157 - // path: '/cases',
158 - // name: 'Cases',
159 - // component: CasesView,
160 - // meta: { requiresAuth: true }
161 - // },
162 - {
163 - path: '/:pathMatch(.*)*',
164 - name: 'NotFound',
165 - component: NotFound
166 - }
167 -]
168 -
169 -const router = createRouter({
170 - history: createWebHistory(),
171 - routes
172 -})
173 -
174 -// Navigation guards
175 -router.beforeEach((to, from, next) => {
176 - const authStore = useAuthStore()
177 -
178 - if (to.meta.requiresAuth && !authStore.isLogged) {
179 - next('/login')
180 - } else if (to.meta.requiresGuest && authStore.isLogged) {
181 - next('/')
182 - } else if (to.meta.requiresAuth && authStore.isLogged && !authStore.isCustomerUser) {
183 - // Ensure only customer users can access protected routes
184 - authStore.logout()
185 - next('/login')
186 - } else {
187 - next()
188 - }
189 -})
190 -
191 -export default router
\ No newline at end of file