From 9db7ad344c57f39083de3dc98d12187b0ce3557b Mon Sep 17 00:00:00 2001 From: Padraic Shafer Date: Sat, 8 Aug 2026 09:51:49 -0700 Subject: [PATCH] Apply ruff format auto-fix --- src/nsls2api/api/models/facility_model.py | 3 +- src/nsls2api/api/models/person_model.py | 8 +- src/nsls2api/api/models/proposal_model.py | 3 +- src/nsls2api/api/v1/facility_api.py | 16 +-- src/nsls2api/api/v1/proposal_api.py | 11 +- src/nsls2api/api/v1/user_api.py | 15 +-- src/nsls2api/infrastructure/config.py | 6 +- src/nsls2api/services/beamline_service.py | 4 +- src/nsls2api/services/ldap_service.py | 59 +++++++---- src/nsls2api/services/proposal_service.py | 23 ++-- .../tests/services/test_proposal_service.py | 100 ++++++++++++++---- 11 files changed, 166 insertions(+), 82 deletions(-) diff --git a/src/nsls2api/api/models/facility_model.py b/src/nsls2api/api/models/facility_model.py index e3ad7293..0302f0f0 100644 --- a/src/nsls2api/api/models/facility_model.py +++ b/src/nsls2api/api/models/facility_model.py @@ -18,10 +18,11 @@ class FacilityCurrentOperatingCycleResponseModel(pydantic.BaseModel): facility: str cycle: str + class FacilityCycleDetailsResponseModel(pydantic.BaseModel): facility: str cycle: str start_date: datetime | None = None end_date: datetime | None = None is_current_operating_cycle: bool - accepting_proposals: bool | None = None \ No newline at end of file + accepting_proposals: bool | None = None diff --git a/src/nsls2api/api/models/person_model.py b/src/nsls2api/api/models/person_model.py index 5d8ebce8..aab9dbcd 100644 --- a/src/nsls2api/api/models/person_model.py +++ b/src/nsls2api/api/models/person_model.py @@ -104,6 +104,7 @@ class UnixInfo(pydantic.BaseModel): homeDirectory: Optional[str] = None loginShell: Optional[str] = None + class IdentityInfo(pydantic.BaseModel): displayName: Optional[str] = None email: Optional[str] = None @@ -111,6 +112,7 @@ class IdentityInfo(pydantic.BaseModel): manager: Optional[str] = None unix: Optional[UnixInfo] = None + class AccountInfo(pydantic.BaseModel): accountExpires: Optional[str] = None badPasswordTime: Optional[str] = None @@ -126,6 +128,7 @@ class AccountInfo(pydantic.BaseModel): uSNCreated: Optional[int] = None uSNChanged: Optional[int] = None + class DirectoryInfo(pydantic.BaseModel): objectGUID: Optional[str] = None objectSid: Optional[str] = None @@ -134,6 +137,7 @@ class DirectoryInfo(pydantic.BaseModel): whenCreated: Optional[str] = None whenChanged: Optional[str] = None + class AttributesInfo(pydantic.BaseModel): sn: Optional[str] = None givenName: Optional[str] = None @@ -145,8 +149,10 @@ class AttributesInfo(pydantic.BaseModel): instanceType: Optional[str] = None objectClass: List[str] = pydantic.Field(default_factory=list) + class LDAPUserResponse(pydantic.BaseModel): """Complete LDAP user data from direct LDAP query""" + dn: Optional[str] = None status: str = "Read" readTime: Optional[str] = None @@ -154,4 +160,4 @@ class LDAPUserResponse(pydantic.BaseModel): account: Optional[AccountInfo] = None directory: Optional[DirectoryInfo] = None groups: List[str] = pydantic.Field(default_factory=list) - attributes: Optional[AttributesInfo] = None \ No newline at end of file + attributes: Optional[AttributesInfo] = None diff --git a/src/nsls2api/api/models/proposal_model.py b/src/nsls2api/api/models/proposal_model.py index b426e65b..7285edae 100644 --- a/src/nsls2api/api/models/proposal_model.py +++ b/src/nsls2api/api/models/proposal_model.py @@ -159,8 +159,9 @@ class ProposalIdDataSession(pydantic.BaseModel): proposal_id: str data_session: str | None = None + class ProposalIdDataSessionList(pydantic.BaseModel): proposals: list[ProposalIdDataSession] count: int page_size: int - page: int \ No newline at end of file + page: int diff --git a/src/nsls2api/api/v1/facility_api.py b/src/nsls2api/api/v1/facility_api.py index a525d115..a1d9053b 100644 --- a/src/nsls2api/api/v1/facility_api.py +++ b/src/nsls2api/api/v1/facility_api.py @@ -3,16 +3,20 @@ from nsls2api.api.models.facility_model import ( FacilityCurrentOperatingCycleResponseModel, - FacilityCycleDetailsResponseModel, FacilityCyclesResponseModel, - FacilityName) + FacilityCycleDetailsResponseModel, + FacilityCyclesResponseModel, + FacilityName, +) from nsls2api.api.models.proposal_model import CycleProposalList from nsls2api.infrastructure.logging import logger from nsls2api.infrastructure.security import validate_admin_role from nsls2api.services import facility_service, proposal_service -from nsls2api.services.facility_service import (CycleNotFoundError, - CycleOperationError, - CycleUpdateError, - CycleVerificationError) +from nsls2api.services.facility_service import ( + CycleNotFoundError, + CycleOperationError, + CycleUpdateError, + CycleVerificationError, +) router = fastapi.APIRouter() diff --git a/src/nsls2api/api/v1/proposal_api.py b/src/nsls2api/api/v1/proposal_api.py index 92532e4b..e39a27d2 100644 --- a/src/nsls2api/api/v1/proposal_api.py +++ b/src/nsls2api/api/v1/proposal_api.py @@ -15,7 +15,7 @@ RecentProposalsList, SingleProposal, UsernamesList, - ProposalIdDataSessionList + ProposalIdDataSessionList, ) from nsls2api.infrastructure.logging import logger from nsls2api.infrastructure.security import get_current_user, validate_admin_role @@ -105,13 +105,14 @@ async def get_proposals( cycle: Annotated[list[str], Query()] = [], facility: Annotated[list[FacilityName], Query()] = [FacilityName.nsls2], username: str | None = Query(None, description="Filter proposals by username"), - saf_status: list[str] | None = Query(default=None, description="Filter proposals and SAFs by SAF status"), + saf_status: list[str] | None = Query( + default=None, description="Filter proposals and SAFs by SAF status" + ), page_size: int = Query(10, ge=1, le=200), page: int = Query(1, ge=1), include_directories: bool = False, ): - proposal_list = await proposal_service.fetch_proposals( proposal_id=proposal_id, beamline=beamline, @@ -123,7 +124,6 @@ async def get_proposals( page=page, include_directories=include_directories, ) - response_model = { "proposals": proposal_list, @@ -135,7 +135,6 @@ async def get_proposals( return response_model - @router.get( "/proposals/data-sessions", response_model=ProposalIdDataSessionList, @@ -156,7 +155,7 @@ async def get_proposals_data_sessions( cycle=cycle, facility=facility, page_size=page_size, - page=page + page=page, ) response_model = { diff --git a/src/nsls2api/api/v1/user_api.py b/src/nsls2api/api/v1/user_api.py index 539ed96a..e9cda89a 100644 --- a/src/nsls2api/api/v1/user_api.py +++ b/src/nsls2api/api/v1/user_api.py @@ -73,25 +73,26 @@ async def get_person_by_department(department_code: str = "PS"): # TODO: Add back into schema if we decide to use this endpoint. -@router.get("/person/me",include_in_schema=True) +@router.get("/person/me", include_in_schema=True) async def get_myself(upn: str = Header(...)): - #upn: User principal name + # upn: User principal name if not upn: - raise HTTPException(status_code=400, detail = "upn not found") + raise HTTPException(status_code=400, detail="upn not found") settings = get_settings() - ldap_info = await asyncio.to_thread(get_user_info, + ldap_info = await asyncio.to_thread( + get_user_info, upn, settings.ldap_server, settings.ldap_base_dn, settings.ldap_bind_user, - settings.ldap_bind_password + settings.ldap_bind_password, ) if not ldap_info: raise HTTPException(status_code=404, detail="User not found in LDAP") - + shaped_info = shape_ldap_response(ldap_info) return LDAPUserResponse(**shaped_info) - + @router.get("/data-session/{username}", response_model=DataSessionAccess, tags=["data"]) @router.get( diff --git a/src/nsls2api/infrastructure/config.py b/src/nsls2api/infrastructure/config.py index d6d0eb55..8e2cfd36 100644 --- a/src/nsls2api/infrastructure/config.py +++ b/src/nsls2api/infrastructure/config.py @@ -71,8 +71,10 @@ class Settings(BaseSettings): extra="ignore", ) - #Whoami LDAP settings - ldap_server: str = Field(default="ldaps://ldapproxy.nsls2.bnl.gov", alias="LDAP_SERVER") + # Whoami LDAP settings + ldap_server: str = Field( + default="ldaps://ldapproxy.nsls2.bnl.gov", alias="LDAP_SERVER" + ) ldap_base_dn: str = Field(default="dc=bnl,dc=gov", alias="LDAP_BASE_DN") ldap_bind_user: str = Field(default="", alias="LDAP_BIND_USER") ldap_bind_password: str = Field(default="", alias="LDAP_BIND_PASSWORD") diff --git a/src/nsls2api/services/beamline_service.py b/src/nsls2api/services/beamline_service.py index 18b6f5ba..aa2b22c0 100644 --- a/src/nsls2api/services/beamline_service.py +++ b/src/nsls2api/services/beamline_service.py @@ -77,7 +77,9 @@ async def all_services(name: str) -> Optional[ServicesOnly]: async def detectors(name: str) -> list[Detector]: - beamline_detectors = await Beamline.find_one(Beamline.name == name.upper()).project(DetectorView) + beamline_detectors = await Beamline.find_one(Beamline.name == name.upper()).project( + DetectorView + ) if beamline_detectors is None: raise LookupError(f"Beamline '{name.upper()}' does not exist.") return beamline_detectors.detectors diff --git a/src/nsls2api/services/ldap_service.py b/src/nsls2api/services/ldap_service.py index fcd93981..55ff1575 100644 --- a/src/nsls2api/services/ldap_service.py +++ b/src/nsls2api/services/ldap_service.py @@ -7,30 +7,33 @@ def to_hex(val): - + if isinstance(val, bytes): return binascii.hexlify(val).decode() return None + def get_user_info(upn, ldap_server, ldap_base_dn, ldap_bind_user, bind_password): - conn = None + conn = None try: server = Server(ldap_server) - conn = Connection(server, user=ldap_bind_user, password=bind_password, auto_bind=True) + conn = Connection( + server, user=ldap_bind_user, password=bind_password, auto_bind=True + ) search_filter = f"(&(objectclass=person)(userPrincipalName={upn}))" - conn.search(ldap_base_dn, search_filter, attributes=['sAMAccountName']) + conn.search(ldap_base_dn, search_filter, attributes=["sAMAccountName"]) if not conn.entries: logger.warning("No entries found for the given UPN.") return None entry = conn.entries[0] - username = entry.sAMAccountName.value if 'sAMAccountName' in entry else None + username = entry.sAMAccountName.value if "sAMAccountName" in entry else None if username is None: return None search_filter = f"(&(objectclass=posixaccount)(sAMAccountName={username}))" - conn.search(ldap_base_dn, search_filter, attributes=['*']) + conn.search(ldap_base_dn, search_filter, attributes=["*"]) if not conn.entries: logger.warning("no posix entries found for the given username.") @@ -52,36 +55,50 @@ def get_user_info(upn, ldap_server, ldap_base_dn, ldap_bind_user, bind_password) if conn is not None: conn.unbind() + def filetime_to_str(filetime): try: - if filetime is None or int(filetime) == 0 or int(filetime) == 9223372036854775807: + if ( + filetime is None + or int(filetime) == 0 + or int(filetime) == 9223372036854775807 + ): return "Never" dt = datetime(1601, 1, 1) + timedelta(microseconds=int(filetime) // 10) return dt.strftime("%Y-%m-%d %H:%M:%S UTC") except Exception: return str(filetime) + def generalized_time_to_str(gt): try: - if not gt: return "" + if not gt: + return "" dt = datetime.strptime(gt.split(".")[0], "%Y%m%d%H%M%S") return dt.strftime("%Y-%m-%d %H:%M:%S UTC") except Exception: return str(gt) + def decode_uac(uac): flags = [] try: val = int(uac) - if val & 0x0001: flags.append("SCRIPT") - if val & 0x0002: flags.append("ACCOUNTDISABLE") - if val & 0x0008: flags.append("HOMEDIR_REQUIRED") - if val & 0x0200: flags.append("NORMAL_ACCOUNT") - if val & 0x1000: flags.append("PASSWORD_EXPIRED") + if val & 0x0001: + flags.append("SCRIPT") + if val & 0x0002: + flags.append("ACCOUNTDISABLE") + if val & 0x0008: + flags.append("HOMEDIR_REQUIRED") + if val & 0x0200: + flags.append("NORMAL_ACCOUNT") + if val & 0x1000: + flags.append("PASSWORD_EXPIRED") except Exception: return [] return flags or ["NORMAL_ACCOUNT"] + def shape_ldap_response(user_info, dn=None, status="Read", read_time=None): def clean_groups(groups_val): if not groups_val: @@ -89,7 +106,9 @@ def clean_groups(groups_val): if isinstance(groups_val, list): return groups_val elif isinstance(groups_val, str): - return [g.strip() for g in groups_val.replace("\n", ",").split(",") if g.strip()] + return [ + g.strip() for g in groups_val.replace("\n", ",").split(",") if g.strip() + ] return [] return { @@ -106,8 +125,8 @@ def clean_groups(groups_val): "uidNumber": user_info.get("uidNumber"), "gidNumber": user_info.get("gidNumber"), "homeDirectory": user_info.get("homeDirectory"), - "loginShell": user_info.get("loginShell") - } + "loginShell": user_info.get("loginShell"), + }, }, "account": { "accountExpires": filetime_to_str(user_info.get("accountExpires")), @@ -142,6 +161,8 @@ def clean_groups(groups_val): "codePage": user_info.get("codePage"), "countryCode": user_info.get("countryCode"), "instanceType": user_info.get("instanceType"), - "objectClass": [s.strip() for s in user_info.get("objectClass", "").split() if s.strip()] - } - } \ No newline at end of file + "objectClass": [ + s.strip() for s in user_info.get("objectClass", "").split() if s.strip() + ], + }, + } diff --git a/src/nsls2api/services/proposal_service.py b/src/nsls2api/services/proposal_service.py index 3397ba72..8a5c57e2 100644 --- a/src/nsls2api/services/proposal_service.py +++ b/src/nsls2api/services/proposal_service.py @@ -16,7 +16,7 @@ ProposalDiagnostics, ProposalFullDetails, ProposalsToChangeList, - ProposalIdDataSession + ProposalIdDataSession, ) from nsls2api.infrastructure.logging import logger from nsls2api.models.cycles import Cycle @@ -350,12 +350,10 @@ async def fetch_proposals( page_size: int = 10, page: int = 1, include_directories: bool = False, - ) -> Optional[list[ProposalFullDetails]]: query = [] saf_status_upper: list[str] = [] - if beamline: beamline_upper = [beamline_name.upper() for beamline_name in beamline] query.append(In(Proposal.instruments, beamline_upper)) @@ -365,7 +363,7 @@ async def fetch_proposals( if proposal_id: query.append(In(Proposal.proposal_id, proposal_id)) - + if username is not None: username = username.strip() if username: @@ -373,16 +371,14 @@ async def fetch_proposals( if saf_status: saf_status_upper = [ - stripped.upper() - for s in saf_status - if (stripped := s.strip()) + stripped.upper() for s in saf_status if (stripped := s.strip()) ] query.append( ElemMatch( Proposal.safs, { "status": {"$in": saf_status_upper}, - } + }, ) ) @@ -408,9 +404,7 @@ async def fetch_proposals( filtered_proposals = [] for proposal in proposals: proposal.safs = [ - saf - for saf in (proposal.safs or []) - if saf.status in saf_status_upper + saf for saf in (proposal.safs or []) if saf.status in saf_status_upper ] if proposal.safs: filtered_proposals.append(proposal) @@ -429,6 +423,7 @@ async def fetch_proposals( else: return proposals + async def fetch_data_sessions( proposal_id: list[str] | None = None, beamline: list[str] | None = None, @@ -456,10 +451,7 @@ async def fetch_data_sessions( filter_query = And(*query) if query else {} proposals = ( - await Proposal.find_many( - filter_query, - projection_model=ProposalIdDataSession - ) + await Proposal.find_many(filter_query, projection_model=ProposalIdDataSession) .sort(-Proposal.last_updated) .limit(page_size) .skip(page_size * (page - 1)) @@ -468,6 +460,7 @@ async def fetch_data_sessions( return proposals + async def proposal_type_description_from_pass_type_id( pass_type_id: int, ) -> Optional[str]: diff --git a/src/nsls2api/tests/services/test_proposal_service.py b/src/nsls2api/tests/services/test_proposal_service.py index 45a44ff9..b4a0ccd3 100644 --- a/src/nsls2api/tests/services/test_proposal_service.py +++ b/src/nsls2api/tests/services/test_proposal_service.py @@ -191,6 +191,7 @@ async def test_data_sessions_invalid_beamline(admin_api_key): assert body["count"] == 0 assert body["proposals"] == [] + @pytest.mark.anyio async def test_fetch_proposals_filter_username_positive(): """Username filter positive - target username A vs control username B.""" @@ -199,34 +200,50 @@ async def test_fetch_proposals_filter_username_positive(): data_session="pass-1001", cycles=["2025-1"], instruments=["TEST"], - users=[User(first_name="Target", last_name="User", email="target@example.com", username="alice")], + users=[ + User( + first_name="Target", + last_name="User", + email="target@example.com", + username="alice", + ) + ], safs=[], ) await target_proposal.insert() - + control_proposal = Proposal( proposal_id="1002", data_session="pass-1002", cycles=["2025-1"], instruments=["TEST"], - users=[User(first_name="Control", last_name="User", email="control@example.com", username="bob")], + users=[ + User( + first_name="Control", + last_name="User", + email="control@example.com", + username="bob", + ) + ], safs=[], ) await control_proposal.insert() - + results = await proposal_service.fetch_proposals(username="alice") - + result_ids = {p.proposal_id for p in results} assert "1001" in result_ids assert "1002" not in result_ids + @pytest.mark.anyio async def test_fetch_proposals_filter_username_negative(): """Username filter negative - nonexistent username returns empty.""" results = await proposal_service.fetch_proposals(username="nonexistent_user_xyz") - + assert len(results) == 0 + @pytest.mark.anyio async def test_fetch_proposals_filter_saf_status_positive(): """SAF status positive - APPROVED proposal vs EXPIRED proposal.""" @@ -239,7 +256,7 @@ async def test_fetch_proposals_filter_saf_status_positive(): safs=[SafetyForm(saf_id="SAF001", status="APPROVED", instruments=["TEST"])], ) await target_proposal.insert() - + control_proposal = Proposal( proposal_id="2002", data_session="pass-2002", @@ -249,13 +266,14 @@ async def test_fetch_proposals_filter_saf_status_positive(): safs=[SafetyForm(saf_id="SAF002", status="EXPIRED", instruments=["TEST"])], ) await control_proposal.insert() - + results = await proposal_service.fetch_proposals(saf_status=["APPROVED"]) - + result_ids = {p.proposal_id for p in results} assert "2001" in result_ids assert "2002" not in result_ids + @pytest.mark.anyio async def test_fetch_proposals_filter_saf_status_multiple(): """SAF status multiple - find only proposals and SAFs with specified statuses.""" @@ -272,7 +290,7 @@ async def test_fetch_proposals_filter_saf_status_multiple(): ], ) await target_proposal.insert() - + control_proposal = Proposal( proposal_id="3002", data_session="pass-3002", @@ -286,9 +304,9 @@ async def test_fetch_proposals_filter_saf_status_multiple(): ], ) await control_proposal.insert() - + results = await proposal_service.fetch_proposals(saf_status=["APPROVED", "DRAFT"]) - + result_ids = {p.proposal_id for p in results} assert "3001" in result_ids assert "3002" not in result_ids @@ -299,6 +317,7 @@ async def test_fetch_proposals_filter_saf_status_multiple(): assert "DRAFT" in saf_statuses assert "EXPIRED" not in saf_statuses + @pytest.mark.anyio async def test_fetch_proposals_filter_combined_all_filters(): """Combined filters - username + cycle + beamline + saf_status all required.""" @@ -307,58 +326,93 @@ async def test_fetch_proposals_filter_combined_all_filters(): data_session="pass-4001", cycles=["2025-combined"], instruments=["COMBO-BL"], - users=[User(first_name="Test", last_name="User", email="test@example.com", username="combo_user")], + users=[ + User( + first_name="Test", + last_name="User", + email="test@example.com", + username="combo_user", + ) + ], safs=[SafetyForm(saf_id="SAF-C1", status="APPROVED", instruments=["COMBO-BL"])], ) await matching.insert() - + nonmatching_user = Proposal( proposal_id="4002", data_session="pass-4002", cycles=["2025-combined"], instruments=["COMBO-BL"], - users=[User(first_name="Other", last_name="User", email="other@example.com", username="other_user")], + users=[ + User( + first_name="Other", + last_name="User", + email="other@example.com", + username="other_user", + ) + ], safs=[SafetyForm(saf_id="SAF-C2", status="APPROVED", instruments=["COMBO-BL"])], ) await nonmatching_user.insert() - + nonmatching_cycle = Proposal( proposal_id="4003", data_session="pass-4003", cycles=["2025-other"], instruments=["COMBO-BL"], - users=[User(first_name="Test", last_name="User", email="test@example.com", username="combo_user")], + users=[ + User( + first_name="Test", + last_name="User", + email="test@example.com", + username="combo_user", + ) + ], safs=[SafetyForm(saf_id="SAF-C3", status="APPROVED", instruments=["COMBO-BL"])], ) await nonmatching_cycle.insert() - + nonmatching_beamline = Proposal( proposal_id="4004", data_session="pass-4004", cycles=["2025-combined"], instruments=["OTHER-BL"], - users=[User(first_name="Test", last_name="User", email="test@example.com", username="combo_user")], + users=[ + User( + first_name="Test", + last_name="User", + email="test@example.com", + username="combo_user", + ) + ], safs=[SafetyForm(saf_id="SAF-C4", status="APPROVED", instruments=["OTHER-BL"])], ) await nonmatching_beamline.insert() - + nonmatching_saf_status = Proposal( proposal_id="4005", data_session="pass-4005", cycles=["2025-combined"], instruments=["COMBO-BL"], - users=[User(first_name="Test", last_name="User", email="test@example.com", username="combo_user")], + users=[ + User( + first_name="Test", + last_name="User", + email="test@example.com", + username="combo_user", + ) + ], safs=[SafetyForm(saf_id="SAF-C5", status="DRAFT", instruments=["COMBO-BL"])], ) await nonmatching_saf_status.insert() - + results = await proposal_service.fetch_proposals( username="combo_user", cycle=["2025-combined"], beamline=["COMBO-BL"], saf_status=["APPROVED"], ) - + result_ids = {p.proposal_id for p in results} assert "4001" in result_ids assert "4002" not in result_ids