33import strawberry
44from database import models
55from sqlalchemy import select
6- from sqlalchemy .orm import aliased , selectinload
76
87if TYPE_CHECKING :
98 from api .types .location import LocationNodeType
@@ -37,31 +36,31 @@ async def tasks(
3736 info ,
3837 ) -> list [Annotated ["TaskType" , strawberry .lazy ("api.types.task" )]]:
3938 from api .services .authorization import AuthorizationService
40-
39+
4140 auth_service = AuthorizationService (info .context .db )
4241 accessible_location_ids = await auth_service .get_user_accessible_location_ids (
4342 info .context .user , info .context
4443 )
45-
44+
4645 if not accessible_location_ids :
4746 return []
48-
47+
4948 from sqlalchemy .orm import aliased
5049 patient_locations = aliased (models .patient_locations )
5150 patient_teams = aliased (models .patient_teams )
52-
51+
5352 from sqlalchemy import select
5453 cte = (
5554 select (models .LocationNode .id )
5655 .where (models .LocationNode .id .in_ (accessible_location_ids ))
5756 .cte (name = "accessible_locations" , recursive = True )
5857 )
59-
58+
6059 children = select (models .LocationNode .id ).join (
6160 cte , models .LocationNode .parent_id == cte .c .id
6261 )
6362 cte = cte .union_all (children )
64-
63+
6564 query = (
6665 select (models .Task )
6766 .join (models .Patient , models .Task .patient_id == models .Patient .id )
@@ -91,7 +90,7 @@ async def tasks(
9190 )
9291 .distinct ()
9392 )
94-
93+
9594 result = await info .context .db .execute (query )
9695 return result .scalars ().all ()
9796
@@ -102,7 +101,7 @@ async def root_locations(
102101 ) -> list [Annotated ["LocationNodeType" , strawberry .lazy ("api.types.location" )]]:
103102 import logging
104103 logger = logging .getLogger (__name__ )
105-
104+
106105 # First check what's in user_root_locations table
107106 user_root_check = await info .context .db .execute (
108107 select (models .user_root_locations .c .location_id ).where (
@@ -111,7 +110,7 @@ async def root_locations(
111110 )
112111 user_root_location_ids = [row [0 ] for row in user_root_check .all ()]
113112 logger .info (f"User { self .id } has { len (user_root_location_ids )} entries in user_root_locations: { user_root_location_ids } " )
114-
113+
115114 result = await info .context .db .execute (
116115 select (models .LocationNode )
117116 .join (
@@ -123,7 +122,7 @@ async def root_locations(
123122 )
124123 locations = result .scalars ().all ()
125124 logger .info (f"User { self .id } root_locations query returned { len (locations )} locations: { [loc .id for loc in locations ]} " )
126-
125+
127126 # If we have user_root_locations entries but no locations returned, check if locations exist
128127 if user_root_location_ids and not locations :
129128 location_check = await info .context .db .execute (
@@ -137,5 +136,5 @@ async def root_locations(
137136 f"Checking if locations exist: { [loc .id for loc in existing_locations ]} "
138137 f"with parent_ids: { [loc .parent_id for loc in existing_locations ]} "
139138 )
140-
139+
141140 return locations
0 commit comments