diff --git a/alembic/versions/0002_timestamptz_and_broker_indexes.py b/alembic/versions/0002_timestamptz_and_broker_indexes.py new file mode 100644 index 0000000..049a865 --- /dev/null +++ b/alembic/versions/0002_timestamptz_and_broker_indexes.py @@ -0,0 +1,179 @@ +"""Convert timestamps to TIMESTAMPTZ and add broker indexes/checks.""" + +from alembic import op + +# revision identifiers, used by Alembic. +revision = "0002_timestamptz_broker_idx" +down_revision = "0001_initial_schema" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # Timestamp columns -> TIMESTAMPTZ (treat existing values as UTC) + tables = { + "users": ["created_at", "updated_at"], + "refresh_token": ["expires_at", "created_at", "updated_at"], + "organism": ["created_at", "updated_at"], + "accession_registry": ["accepted_at", "created_at", "updated_at"], + "project": ["submitted_at", "created_at", "updated_at"], + "project_submission": [ + "submitted_at", + "created_at", + "updated_at", + "lock_acquired_at", + "lock_expires_at", + ], + "sample": ["created_at", "updated_at"], + "sample_submission": [ + "submitted_at", + "created_at", + "updated_at", + "lock_acquired_at", + "lock_expires_at", + ], + "experiment": ["created_at", "updated_at"], + "experiment_submission": [ + "submitted_at", + "created_at", + "updated_at", + "lock_acquired_at", + "lock_expires_at", + ], + "read": ["created_at", "updated_at"], + "read_submission": ["created_at", "updated_at", "lock_acquired_at", "lock_expires_at"], + "submission_attempt": ["lock_acquired_at", "lock_expires_at", "created_at", "updated_at"], + "submission_event": ["at"], + "assembly": ["created_at", "updated_at"], + "assembly_submission": ["submitted_at", "created_at", "updated_at"], + "assembly_file": ["created_at", "updated_at"], + "genome_note": ["published_at", "created_at", "updated_at"], + "bpa_initiative": ["created_at", "updated_at"], + } + + for table, columns in tables.items(): + for col in columns: + op.execute( + f"ALTER TABLE {table} ALTER COLUMN {col} TYPE TIMESTAMPTZ USING {col} AT TIME ZONE 'UTC'" + ) + + # Sample latitude/longitude checks (skip if already exists from schema.sql) + op.execute( + """ + DO $$ BEGIN + IF NOT EXISTS ( + SELECT 1 FROM pg_constraint WHERE conname = 'chk_sample_latitude' + ) THEN + ALTER TABLE sample ADD CONSTRAINT chk_sample_latitude + CHECK (latitude IS NULL OR latitude BETWEEN -90 AND 90); + END IF; + END $$; + """ + ) + op.execute( + """ + DO $$ BEGIN + IF NOT EXISTS ( + SELECT 1 FROM pg_constraint WHERE conname = 'chk_sample_longitude' + ) THEN + ALTER TABLE sample ADD CONSTRAINT chk_sample_longitude + CHECK (longitude IS NULL OR longitude BETWEEN -180 AND 180); + END IF; + END $$; + """ + ) + + # Broker and status indexes + op.execute( + "CREATE INDEX IF NOT EXISTS idx_project_submission_status ON project_submission (status)" + ) + op.execute( + "CREATE INDEX IF NOT EXISTS idx_project_submission_lock_expires_at ON project_submission (lock_expires_at)" + ) + op.execute( + "CREATE INDEX IF NOT EXISTS idx_sample_submission_status ON sample_submission (status)" + ) + op.execute( + "CREATE INDEX IF NOT EXISTS idx_sample_submission_lock_expires_at ON sample_submission (lock_expires_at)" + ) + op.execute( + "CREATE INDEX IF NOT EXISTS idx_experiment_submission_status ON experiment_submission (status)" + ) + op.execute( + "CREATE INDEX IF NOT EXISTS idx_experiment_submission_lock_expires_at ON experiment_submission (lock_expires_at)" + ) + op.execute("CREATE INDEX IF NOT EXISTS idx_read_submission_status ON read_submission (status)") + op.execute( + "CREATE INDEX IF NOT EXISTS idx_read_submission_lock_expires_at ON read_submission (lock_expires_at)" + ) + op.execute( + "CREATE INDEX IF NOT EXISTS idx_submission_attempt_status ON submission_attempt (status)" + ) + op.execute( + "CREATE INDEX IF NOT EXISTS idx_submission_attempt_lock_expires_at ON submission_attempt (lock_expires_at)" + ) + + +def downgrade() -> None: + # Drop indexes + op.execute("DROP INDEX IF EXISTS idx_submission_attempt_lock_expires_at") + op.execute("DROP INDEX IF EXISTS idx_submission_attempt_status") + op.execute("DROP INDEX IF EXISTS idx_read_submission_lock_expires_at") + op.execute("DROP INDEX IF EXISTS idx_read_submission_status") + op.execute("DROP INDEX IF EXISTS idx_experiment_submission_lock_expires_at") + op.execute("DROP INDEX IF EXISTS idx_experiment_submission_status") + op.execute("DROP INDEX IF EXISTS idx_sample_submission_lock_expires_at") + op.execute("DROP INDEX IF EXISTS idx_sample_submission_status") + op.execute("DROP INDEX IF EXISTS idx_project_submission_lock_expires_at") + op.execute("DROP INDEX IF EXISTS idx_project_submission_status") + + # Drop checks + op.execute("ALTER TABLE sample DROP CONSTRAINT IF EXISTS chk_sample_longitude") + op.execute("ALTER TABLE sample DROP CONSTRAINT IF EXISTS chk_sample_latitude") + + # Convert columns back to TIMESTAMP (drop tz) + tables = { + "users": ["created_at", "updated_at"], + "refresh_token": ["expires_at", "created_at", "updated_at"], + "organism": ["created_at", "updated_at"], + "accession_registry": ["accepted_at", "created_at", "updated_at"], + "project": ["submitted_at", "created_at", "updated_at"], + "project_submission": [ + "submitted_at", + "created_at", + "updated_at", + "lock_acquired_at", + "lock_expires_at", + ], + "sample": ["created_at", "updated_at"], + "sample_submission": [ + "submitted_at", + "created_at", + "updated_at", + "lock_acquired_at", + "lock_expires_at", + ], + "experiment": ["created_at", "updated_at"], + "experiment_submission": [ + "submitted_at", + "created_at", + "updated_at", + "lock_acquired_at", + "lock_expires_at", + ], + "read": ["created_at", "updated_at"], + "read_submission": ["created_at", "updated_at", "lock_acquired_at", "lock_expires_at"], + "submission_attempt": ["lock_acquired_at", "lock_expires_at", "created_at", "updated_at"], + "submission_event": ["at"], + "assembly": ["created_at", "updated_at"], + "assembly_submission": ["submitted_at", "created_at", "updated_at"], + "assembly_file": ["created_at", "updated_at"], + "genome_note": ["published_at", "created_at", "updated_at"], + "bpa_initiative": ["created_at", "updated_at"], + } + + for table, columns in tables.items(): + for col in columns: + op.execute( + f"ALTER TABLE {table} ALTER COLUMN {col} TYPE TIMESTAMP USING {col}::timestamp" + ) diff --git a/app/api/v1/api.py b/app/api/v1/api.py index 96d4895..185a7ff 100644 --- a/app/api/v1/api.py +++ b/app/api/v1/api.py @@ -1,6 +1,7 @@ from fastapi import APIRouter from app.api.v1.endpoints import ( + admin, assemblies, auth, bpa_initiatives, @@ -25,6 +26,7 @@ api_router.include_router(auth.router, prefix="/auth", tags=["authentication"]) api_router.include_router(users.router, prefix="/users", tags=["users"]) api_router.include_router(broker.router, prefix="/broker", tags=["broker"]) +api_router.include_router(admin.router, prefix="/admin", tags=["admin"]) # Core entity routers api_router.include_router(organisms.router, prefix="/organisms", tags=["organisms"]) diff --git a/app/api/v1/endpoints/admin.py b/app/api/v1/endpoints/admin.py new file mode 100644 index 0000000..8f489bb --- /dev/null +++ b/app/api/v1/endpoints/admin.py @@ -0,0 +1,25 @@ +from typing import Any, Dict + +from fastapi import APIRouter, Depends +from sqlalchemy.orm import Session + +from app.core.dependencies import get_current_active_user, get_db +from app.core.policy import policy +from app.models.user import User +from app.services.broker_service import expire_leases + +router = APIRouter() + + +@router.post("/leases/expire") +@policy("admin:expire_leases") +def expire_all_leases( + db: Session = Depends(get_db), + current_user: User = Depends(get_current_active_user), +) -> Dict[str, Dict[str, int]]: + """ + Expire all broker leases whose locks have passed. Admin-only. + """ + expired = expire_leases(db) + db.commit() + return {"expired_counts": expired} diff --git a/app/api/v1/endpoints/assemblies.py b/app/api/v1/endpoints/assemblies.py index 500227b..a2ba3e7 100644 --- a/app/api/v1/endpoints/assemblies.py +++ b/app/api/v1/endpoints/assemblies.py @@ -4,12 +4,9 @@ from fastapi import APIRouter, Depends, HTTPException, Query from sqlalchemy.orm import Session -from app.core.dependencies import ( - get_current_active_user, - get_current_superuser, - get_db, - require_role, -) +from app.core.dependencies import get_current_active_user, get_db +from app.core.pagination import Pagination, apply_pagination, pagination_params +from app.core.policy import policy from app.models.assembly import Assembly, AssemblyFile, AssemblySubmission from app.models.experiment import Experiment from app.models.organism import Organism @@ -21,13 +18,16 @@ ) from app.schemas.assembly import ( AssemblyCreate, - AssemblyFile as AssemblyFileSchema, + AssemblyCreateFromExperiments, AssemblyFileCreate, AssemblyFileUpdate, AssemblySubmissionCreate, AssemblySubmissionUpdate, AssemblyUpdate, ) +from app.schemas.assembly import ( + AssemblyFile as AssemblyFileSchema, +) from app.schemas.assembly import ( AssemblySubmission as AssemblySubmissionSchema, ) @@ -205,7 +205,6 @@ def get_pipeline_inputs_by_tax_id( return result - @router.get("/manifest/{tax_id}") def get_assembly_manifest( *, @@ -237,7 +236,7 @@ def get_assembly_manifest( if not samples: raise HTTPException( status_code=404, - detail=f"No samples found for organism {organism.grouping_key} (tax_id: {tax_id})" + detail=f"No samples found for organism {organism.grouping_key} (tax_id: {tax_id})", ) # Get all experiments for these samples @@ -247,7 +246,7 @@ def get_assembly_manifest( if not experiments: raise HTTPException( status_code=404, - detail=f"No experiments found for organism {organism.grouping_key} (tax_id: {tax_id})" + detail=f"No experiments found for organism {organism.grouping_key} (tax_id: {tax_id})", ) # Get all reads for these experiments @@ -257,7 +256,7 @@ def get_assembly_manifest( if not reads: raise HTTPException( status_code=404, - detail=f"No reads found for organism {organism.grouping_key} (tax_id: {tax_id})" + detail=f"No reads found for organism {organism.grouping_key} (tax_id: {tax_id})", ) # Generate YAML manifest @@ -267,12 +266,13 @@ def get_assembly_manifest( return Response(content=yaml_content, media_type="application/x-yaml") -@router.post("/from-experiments/{tax_id}", response_model=Dict[str, Any]) +@router.post("/from-experiments/{tax_id}", response_model=AssemblySchema) +@policy("assemblies:write") def create_assembly_from_experiments( *, db: Session = Depends(get_db), tax_id: int, - assembly_in: AssemblyCreate, + assembly_in: AssemblyCreateFromExperiments, current_user: User = Depends(get_current_active_user), ) -> Any: """ @@ -283,35 +283,26 @@ def create_assembly_from_experiments( - OXFORD_NANOPORE: platform == "OXFORD_NANOPORE" - Hi-C: platform == "ILLUMINA" AND library_strategy == "Hi-C" - The data_types field in assembly_in will be ignored and auto-determined. + The organism_key is automatically determined from tax_id. + The data_types field is auto-detected but can be overridden. """ - require_role(current_user, ["curator", "admin"]) - try: assembly, platform_info = assembly_service.create_from_experiments( db, tax_id=tax_id, assembly_in=assembly_in ) - - return { - "assembly": assembly, - "platform_detection": platform_info, - "message": f"Assembly created with auto-detected data_types: {assembly.data_types}" - } + return assembly except ValueError as e: raise HTTPException(status_code=400, detail=str(e)) @router.post("/", response_model=AssemblySchema) +@policy("assemblies:write") def create_assembly( *, db: Session = Depends(get_db), assembly_in: AssemblyCreate, current_user: User = Depends(get_current_active_user), ) -> Any: - # Create new assembly. - # Only users with 'curator' or 'admin' role can create assemblies - require_role(current_user, ["curator", "admin"]) - assembly = assembly_service.create(db, obj_in=assembly_in) return assembly @@ -332,6 +323,7 @@ def read_assembly( @router.put("/{assembly_id}", response_model=AssemblySchema) +@policy("assemblies:write") def update_assembly( *, db: Session = Depends(get_db), @@ -339,10 +331,6 @@ def update_assembly( assembly_in: AssemblyUpdate, current_user: User = Depends(get_current_active_user), ) -> Any: - # Update an assembly. - # Only users with 'curator' or 'admin' role can update assemblies - require_role(current_user, ["curator", "admin"]) - assembly = db.query(Assembly).filter(Assembly.id == assembly_id).first() if not assembly: raise HTTPException(status_code=404, detail="Assembly not found") @@ -358,11 +346,12 @@ def update_assembly( @router.delete("/{assembly_id}", response_model=AssemblySchema) +@policy("assemblies:delete") def delete_assembly( *, db: Session = Depends(get_db), assembly_id: UUID, - current_user: User = Depends(get_current_superuser), + current_user: User = Depends(get_current_active_user), ) -> Any: # Delete an assembly. # Only superusers can delete assemblies @@ -379,8 +368,7 @@ def delete_assembly( @router.get("/submission/", response_model=List[AssemblySubmissionSchema]) def read_assembly_submissions( db: Session = Depends(get_db), - skip: int = 0, - limit: int = 100, + pagination: Pagination = Depends(pagination_params), status: Optional[SubmissionStatus] = Query(None, description="Filter by submission status"), assembly_id: Optional[UUID] = Query(None, description="Filter by assembly ID"), current_user: User = Depends(get_current_active_user), @@ -392,12 +380,13 @@ def read_assembly_submissions( query = db.query(AssemblySubmission) if status: query = query.filter(AssemblySubmission.status == status.value) - submissions = query.offset(skip).limit(limit).all() + submissions = apply_pagination(query, pagination).all() return submissions @router.post("/submission/", response_model=AssemblySubmissionSchema) +@policy("assemblies:write") def create_assembly_submission( *, db: Session = Depends(get_db), @@ -405,7 +394,6 @@ def create_assembly_submission( current_user: User = Depends(get_current_active_user), ) -> Any: """Create new assembly submission.""" - require_role(current_user, ["curator", "admin"]) # Verify assembly exists assembly = assembly_service.get(db, id=submission_in.assembly_id) @@ -417,6 +405,7 @@ def create_assembly_submission( @router.put("/submission/{submission_id}", response_model=AssemblySubmissionSchema) +@policy("assemblies:write") def update_assembly_submission( *, db: Session = Depends(get_db), @@ -425,7 +414,6 @@ def update_assembly_submission( current_user: User = Depends(get_current_active_user), ) -> Any: """Update an assembly submission.""" - require_role(current_user, ["curator", "admin"]) submission = assembly_submission_service.get(db, id=submission_id) if not submission: @@ -439,6 +427,7 @@ def update_assembly_submission( # Assembly File endpoints # ========================================== + @router.get("/{assembly_id}/files", response_model=List[AssemblyFileSchema]) def read_assembly_files( *, @@ -453,7 +442,9 @@ def read_assembly_files( raise HTTPException(status_code=404, detail="Assembly not found") if file_type: - files = assembly_file_service.get_by_assembly_and_type(db, assembly_id=assembly_id, file_type=file_type) + files = assembly_file_service.get_by_assembly_and_type( + db, assembly_id=assembly_id, file_type=file_type + ) else: files = assembly_file_service.get_by_assembly_id(db, assembly_id=assembly_id) @@ -461,6 +452,7 @@ def read_assembly_files( @router.post("/{assembly_id}/files", response_model=AssemblyFileSchema) +@policy("assemblies:write") def create_assembly_file( *, db: Session = Depends(get_db), @@ -469,7 +461,6 @@ def create_assembly_file( current_user: User = Depends(get_current_active_user), ) -> Any: """Add a file to an assembly.""" - require_role(current_user, ["curator", "admin"]) # Verify assembly exists assembly = assembly_service.get(db, id=assembly_id) @@ -485,6 +476,7 @@ def create_assembly_file( @router.put("/files/{file_id}", response_model=AssemblyFileSchema) +@policy("assemblies:write") def update_assembly_file( *, db: Session = Depends(get_db), @@ -493,7 +485,6 @@ def update_assembly_file( current_user: User = Depends(get_current_active_user), ) -> Any: """Update an assembly file.""" - require_role(current_user, ["curator", "admin"]) file = assembly_file_service.get(db, id=file_id) if not file: @@ -504,11 +495,12 @@ def update_assembly_file( @router.delete("/files/{file_id}") +@policy("assemblies:delete") def delete_assembly_file( *, db: Session = Depends(get_db), file_id: UUID, - current_user: User = Depends(get_current_superuser), + current_user: User = Depends(get_current_active_user), ) -> Any: """Delete an assembly file.""" file = assembly_file_service.get(db, id=file_id) diff --git a/app/api/v1/endpoints/bpa_initiatives.py b/app/api/v1/endpoints/bpa_initiatives.py index bf185cb..956fb69 100644 --- a/app/api/v1/endpoints/bpa_initiatives.py +++ b/app/api/v1/endpoints/bpa_initiatives.py @@ -3,12 +3,9 @@ from fastapi import APIRouter, Depends, HTTPException from sqlalchemy.orm import Session -from app.core.dependencies import ( - get_current_active_user, - get_current_superuser, - get_db, - require_role, -) +from app.core.dependencies import get_current_active_user, get_db +from app.core.pagination import Pagination, apply_pagination, pagination_params +from app.core.policy import policy from app.models.bpa_initiative import BPAInitiative from app.models.user import User from app.schemas.bpa_initiative import ( @@ -25,19 +22,19 @@ @router.get("/", response_model=List[BPAInitiativeSchema]) def read_bpa_initiatives( db: Session = Depends(get_db), - skip: int = 0, - limit: int = 100, + pagination: Pagination = Depends(pagination_params), current_user: User = Depends(get_current_active_user), ) -> Any: """ Retrieve BPA initiatives. """ # All users can read BPA initiatives - initiatives = db.query(BPAInitiative).offset(skip).limit(limit).all() + initiatives = apply_pagination(db.query(BPAInitiative), pagination).all() return initiatives @router.post("/", response_model=BPAInitiativeSchema) +@policy("bpa_initiatives:write") def create_bpa_initiative( *, db: Session = Depends(get_db), @@ -47,9 +44,6 @@ def create_bpa_initiative( """ Create new BPA initiative. """ - # Only users with 'curator' or 'admin' role can create BPA initiatives - require_role(current_user, ["curator", "admin"]) - initiative = BPAInitiative( project_code=getattr(initiative_in, "project_code", None), title=getattr(initiative_in, "title", None), @@ -79,6 +73,7 @@ def read_bpa_initiative( @router.put("/{initiative_id}", response_model=BPAInitiativeSchema) +@policy("bpa_initiatives:write") def update_bpa_initiative( *, db: Session = Depends(get_db), @@ -89,9 +84,6 @@ def update_bpa_initiative( """ Update a BPA initiative. """ - # Only users with 'curator' or 'admin' role can update BPA initiatives - require_role(current_user, ["curator", "admin"]) - initiative = db.query(BPAInitiative).filter(BPAInitiative.project_code == initiative_id).first() if not initiative: raise HTTPException(status_code=404, detail="BPA initiative not found") @@ -107,16 +99,16 @@ def update_bpa_initiative( @router.delete("/{initiative_id}", response_model=BPAInitiativeSchema) +@policy("bpa_initiatives:write") def delete_bpa_initiative( *, db: Session = Depends(get_db), initiative_id: str, - current_user: User = Depends(get_current_superuser), + current_user: User = Depends(get_current_active_user), ) -> Any: """ Delete a BPA initiative. """ - # Only superusers can delete BPA initiatives initiative = db.query(BPAInitiative).filter(BPAInitiative.project_code == initiative_id).first() if not initiative: raise HTTPException(status_code=404, detail="BPA initiative not found") diff --git a/app/api/v1/endpoints/broker.py b/app/api/v1/endpoints/broker.py index 36b2b7f..3f1059b 100644 --- a/app/api/v1/endpoints/broker.py +++ b/app/api/v1/endpoints/broker.py @@ -8,7 +8,8 @@ from sqlalchemy.dialects.postgresql import insert from sqlalchemy.orm import Session -from app.core.dependencies import get_current_active_user, get_db, has_role, require_role +from app.core.dependencies import get_current_active_user, get_db, has_role +from app.core.policy import policy from app.models.accession_registry import AccessionRegistry from app.models.broker import SubmissionAttempt, SubmissionEvent from app.models.experiment import Experiment, ExperimentSubmission @@ -117,6 +118,7 @@ class ReportResult(BaseModel): @router.post("/claim", response_model=ClaimByEntityResponse) +@policy("broker:claim") def claim_by_entity_ids( *, payload: ClaimByEntityRequest, @@ -128,7 +130,6 @@ def claim_by_entity_ids( This endpoint allows claiming specific entities without requiring an organism_key. It will find the latest draft submission for each entity and claim it with a lease. """ - require_role(current_user, ["broker"]) # Expire any stale leases before claiming expire_stale_leases(db) @@ -543,6 +544,7 @@ def claim_by_entity_ids( @router.post("/organisms/{organism_key}/claim", response_model=ClaimResponse) +@policy("broker:claim") def claim_drafts_for_organism( *, organism_key: str = Path(..., description="Organism grouping_key"), @@ -554,7 +556,6 @@ def claim_drafts_for_organism( """Claim latest draft SampleSubmissions for an organism and mark them 'submitting'. This acts as a short lease to prevent concurrent edits. """ - require_role(current_user, ["broker"]) # Expire any stale leases before claiming expire_stale_leases(db) diff --git a/app/api/v1/endpoints/experiment_submissions.py b/app/api/v1/endpoints/experiment_submissions.py index 15de1e3..4fe6ae0 100644 --- a/app/api/v1/endpoints/experiment_submissions.py +++ b/app/api/v1/endpoints/experiment_submissions.py @@ -4,7 +4,9 @@ from fastapi import APIRouter, Depends, HTTPException, Query from sqlalchemy.orm import Session -from app.core.dependencies import get_current_active_user, get_db, require_role +from app.core.dependencies import get_current_active_user, get_db +from app.core.pagination import Pagination, apply_pagination, pagination_params +from app.core.policy import policy from app.models.experiment import Experiment, ExperimentSubmission from app.models.user import User from app.schemas.bulk_import import BulkExperimentImport, BulkImportResponse @@ -25,10 +27,10 @@ # Experiment Submission endpoints @router.get("/", response_model=List[ExperimentSubmissionSchema]) +@policy("experiment_submissions:read") def read_experiment_submissions( db: Session = Depends(get_db), - skip: int = 0, - limit: int = 100, + pagination: Pagination = Depends(pagination_params), status: Optional[SubmissionStatus] = Query(None, description="Filter by submission status"), full_history: Optional[bool] = Query(False, description="Return full submission history"), current_user: User = Depends(get_current_active_user), @@ -46,11 +48,12 @@ def read_experiment_submissions( ExperimentSubmission.created_at.desc(), ).distinct(ExperimentSubmission.experiment_id) - submissions = query.offset(skip).limit(limit).all() + submissions = apply_pagination(query, pagination).all() return submissions @router.get("/by-experiment-attr", response_model=List[ExperimentSubmissionSchema]) +@policy("experiment_submissions:read") async def get_experiment_submission_by_experiment_attr( db: Session = Depends(get_db), bpa_package_id: Optional[str] = Query(None, description="Filter by bpa_package_id"), @@ -102,6 +105,7 @@ async def get_experiment_submission_by_experiment_attr( @router.get("/{submission_id}", response_model=ExperimentSubmissionSchema) +@policy("experiment_submissions:read") def read_experiment_submission( *, db: Session = Depends(get_db), diff --git a/app/api/v1/endpoints/experiments.py b/app/api/v1/endpoints/experiments.py index 9467cc7..449f3e9 100644 --- a/app/api/v1/endpoints/experiments.py +++ b/app/api/v1/endpoints/experiments.py @@ -4,12 +4,9 @@ from fastapi import APIRouter, Depends, HTTPException, Query from sqlalchemy.orm import Session -from app.core.dependencies import ( - get_current_active_user, - get_current_superuser, - get_db, - require_role, -) +from app.core.dependencies import get_current_active_user, get_db +from app.core.pagination import Pagination, pagination_params +from app.core.policy import policy from app.models.user import User from app.schemas.bulk_import import BulkImportResponse, BulkImportResponseExperiments from app.schemas.experiment import Experiment as ExperimentSchema @@ -23,8 +20,7 @@ @router.get("/", response_model=List[ExperimentSchema]) def read_experiments( db: Session = Depends(get_db), - skip: int = 0, - limit: int = 100, + pagination: Pagination = Depends(pagination_params), sample_id: Optional[UUID] = Query(None, description="Filter by sample ID"), current_user: User = Depends(get_current_active_user), ) -> Any: @@ -32,10 +28,13 @@ def read_experiments( Retrieve experiments. """ # All users can read experiments - return experiment_service.list_experiments(db, skip=skip, limit=limit, sample_id=sample_id) + return experiment_service.list_experiments( + db, skip=pagination.offset, limit=pagination.limit, sample_id=sample_id + ) @router.post("/", response_model=ExperimentSchema) +@policy("experiments:create") def create_experiment( *, db: Session = Depends(get_db), @@ -45,8 +44,6 @@ def create_experiment( """ Create new experiment. """ - # Only users with 'curator' or 'admin' role can create experiments - require_role(current_user, ["curator", "admin"]) try: experiment = experiment_service.create_experiment(db, experiment_in=experiment_in) return experiment @@ -87,6 +84,7 @@ def read_experiment( @router.put("/{experiment_id}", response_model=ExperimentSchema) +@policy("experiments:update") def update_experiment( *, db: Session = Depends(get_db), @@ -97,8 +95,6 @@ def update_experiment( """ Update an experiment. """ - # Only users with 'curator' or 'admin' role can update experiments - require_role(current_user, ["curator", "admin"]) try: experiment = experiment_service.update_experiment( db, experiment_id=experiment_id, experiment_in=experiment_in @@ -117,11 +113,12 @@ def update_experiment( @router.delete("/{experiment_id}", response_model=ExperimentSchema) +@policy("experiments:delete") def delete_experiment( *, db: Session = Depends(get_db), experiment_id: UUID, - current_user: User = Depends(get_current_superuser), + current_user: User = Depends(get_current_active_user), ) -> Any: """ Delete an experiment. @@ -134,6 +131,7 @@ def delete_experiment( @router.post("/bulk-import", response_model=BulkImportResponseExperiments) +@policy("experiments:bulk_import") def bulk_import_experiments( *, db: Session = Depends(get_db), @@ -148,7 +146,5 @@ def bulk_import_experiments( The request body should directly match the format of the JSON file in data/experiments.json, which is a dictionary keyed by package_id without a wrapping 'experiments' key. """ - # Only users with 'curator' or 'admin' role can import experiments - require_role(current_user, ["curator", "admin"]) result = experiment_service.bulk_import_experiments(db, experiments_data=experiments_data) return result diff --git a/app/api/v1/endpoints/genome_notes.py b/app/api/v1/endpoints/genome_notes.py index 06b2b50..a9e7fa9 100644 --- a/app/api/v1/endpoints/genome_notes.py +++ b/app/api/v1/endpoints/genome_notes.py @@ -4,12 +4,9 @@ from fastapi import APIRouter, Depends, HTTPException, Query from sqlalchemy.orm import Session -from app.core.dependencies import ( - get_current_active_user, - get_current_superuser, - get_db, - require_role, -) +from app.core.dependencies import get_current_active_user, get_db +from app.core.pagination import Pagination, pagination_params +from app.core.policy import policy from app.models.genome_note import GenomeNote from app.models.user import User from app.schemas.genome_note import ( @@ -27,8 +24,7 @@ @router.get("/", response_model=List[GenomeNoteSchema]) def read_genome_notes( db: Session = Depends(get_db), - skip: int = 0, - limit: int = 100, + pagination: Pagination = Depends(pagination_params), organism_key: Optional[str] = Query(None, description="Filter by organism key"), assembly_id: Optional[UUID] = Query(None, description="Filter by assembly ID"), is_published: Optional[bool] = Query(None, description="Filter by publication status"), @@ -40,8 +36,8 @@ def read_genome_notes( """ genome_notes = genome_note_service.get_multi_with_filters( db, - skip=skip, - limit=limit, + skip=pagination.offset, + limit=pagination.limit, organism_key=organism_key, assembly_id=assembly_id, is_published=is_published, @@ -51,6 +47,7 @@ def read_genome_notes( @router.post("/", response_model=GenomeNoteSchema) +@policy("genome_notes:write") def create_genome_note( *, db: Session = Depends(get_db), @@ -63,8 +60,6 @@ def create_genome_note( The version number is automatically calculated based on existing versions for the organism. The note is created in draft status (is_published=False). """ - require_role(current_user, ["curator", "admin"]) - # Auto-calculate next version for this organism next_version = genome_note_service.get_next_version(db, genome_note_in.organism_key) @@ -99,6 +94,7 @@ def read_genome_note( @router.put("/{genome_note_id}", response_model=GenomeNoteSchema) +@policy("genome_notes:write") def update_genome_note( *, db: Session = Depends(get_db), @@ -112,8 +108,6 @@ def update_genome_note( Only title and note_url can be updated. Version and publication status cannot be changed through this endpoint. """ - require_role(current_user, ["curator", "admin"]) - genome_note = db.query(GenomeNote).filter(GenomeNote.id == genome_note_id).first() if not genome_note: raise HTTPException(status_code=404, detail="Genome note not found") @@ -129,16 +123,16 @@ def update_genome_note( @router.delete("/{genome_note_id}", response_model=GenomeNoteSchema) +@policy("genome_notes:write") def delete_genome_note( *, db: Session = Depends(get_db), genome_note_id: UUID, - current_user: User = Depends(get_current_superuser), + current_user: User = Depends(get_current_active_user), ) -> Any: """ Delete a genome note. - Only superusers can delete genome notes. """ genome_note = db.query(GenomeNote).filter(GenomeNote.id == genome_note_id).first() if not genome_note: @@ -150,6 +144,7 @@ def delete_genome_note( @router.post("/{genome_note_id}/publish", response_model=GenomeNoteSchema) +@policy("genome_notes:write") def publish_genome_note( *, db: Session = Depends(get_db), @@ -163,8 +158,6 @@ def publish_genome_note( Use the unpublish endpoint first to unpublish the existing note. Only one genome note can be published per organism at a time. """ - require_role(current_user, ["curator", "admin"]) - try: genome_note = genome_note_service.publish_genome_note(db, genome_note_id) return genome_note @@ -178,6 +171,7 @@ def publish_genome_note( @router.post("/{genome_note_id}/unpublish", response_model=GenomeNoteSchema) +@policy("genome_notes:write") def unpublish_genome_note( *, db: Session = Depends(get_db), @@ -189,8 +183,6 @@ def unpublish_genome_note( Sets the genome note back to draft status. """ - require_role(current_user, ["curator", "admin"]) - try: genome_note = genome_note_service.unpublish_genome_note(db, genome_note_id) return genome_note diff --git a/app/api/v1/endpoints/organisms.py b/app/api/v1/endpoints/organisms.py index 4da26c7..fbae19a 100644 --- a/app/api/v1/endpoints/organisms.py +++ b/app/api/v1/endpoints/organisms.py @@ -5,12 +5,10 @@ from fastapi import APIRouter, Depends, HTTPException, Query, status from sqlalchemy.orm import Session -from app.core.dependencies import ( - get_current_active_user, - get_current_superuser, - get_db, - require_role, -) +from app.core.dependencies import get_current_active_user, get_db +from app.core.pagination import Pagination, apply_pagination, pagination_params +from app.core.policy import policy +from app.models.organism import Organism from app.models.user import User from app.schemas.aggregate import OrganismSubmissionJsonResponse from app.schemas.bulk_import import BulkImportResponse @@ -29,18 +27,20 @@ @router.get("/", response_model=List[OrganismSchema]) def read_organisms( db: Session = Depends(get_db), - skip: int = 0, - limit: int = 100, + pagination: Pagination = Depends(pagination_params), current_user: User = Depends(get_current_active_user), ) -> Any: """ Retrieve organisms. """ # All users can read organisms - return organism_service.list_organisms(db, skip=skip, limit=limit) + query = db.query(Organism) + query = apply_pagination(query, pagination) + return query.all() @router.get("/{grouping_key}/experiments") +@policy("organisms:read_sensitive") def get_experiments_for_organism( *, db: Session = Depends(get_db), @@ -51,8 +51,6 @@ def get_experiments_for_organism( """ Return all experiments for the organism, and optionally all reads for each experiment when includeReads is true. """ - # Admin, curator, broker and genome_launcher can get expanded organism data - require_role(current_user, ["admin", "curator", "broker", "genome_launcher"]) data = organism_service.get_experiments_for_organism( db, grouping_key=grouping_key, include_reads=includeReads ) @@ -65,6 +63,7 @@ def get_experiments_for_organism( @router.get("/submissions/{grouping_key}", response_model=OrganismSubmissionJsonResponse) +@policy("organisms:read_sensitive") def get_organism_prepared_payload( *, db: Session = Depends(get_db), @@ -74,8 +73,6 @@ def get_organism_prepared_payload( """ Get all prepared_payload data for samples, experiments, and reads related to a specific organism.grouping_key. """ - # Admin, curator, broker and genome_launcher can get prepared_payload data - require_role(current_user, ["admin", "curator", "broker", "genome_launcher"]) data = organism_service.get_organism_prepared_payload(db, grouping_key=grouping_key) if data is None: raise HTTPException( @@ -86,6 +83,7 @@ def get_organism_prepared_payload( @router.post("/", response_model=OrganismSchema) +@policy("organisms:create") def create_organism( *, db: Session = Depends(get_db), @@ -95,8 +93,6 @@ def create_organism( """ Create new organism. """ - # Only users with 'curator' or 'admin' role can create organisms - require_role(current_user, ["curator", "admin"]) try: organism = organism_service.create_organism(db, organism_in=organism_in) except Exception as e: @@ -122,6 +118,7 @@ def read_organism( @router.patch("/{grouping_key}", response_model=OrganismSchema) +@policy("organisms:update") def update_organism( *, db: Session = Depends(get_db), @@ -132,8 +129,6 @@ def update_organism( """ Update an organism. """ - # Only users with 'curator' or 'admin' role can update organisms - require_role(current_user, ["curator", "admin"]) organism = organism_service.update_organism( db, grouping_key=grouping_key, organism_in=organism_in ) @@ -143,17 +138,16 @@ def update_organism( @router.delete("/{grouping_key}", response_model=OrganismSchema) +@policy("organisms:delete") def delete_organism( *, db: Session = Depends(get_db), grouping_key: str, - current_user: User = Depends(get_current_superuser), + current_user: User = Depends(get_current_active_user), ) -> Any: """ Delete an organism. """ - # Only users with 'superuser' or 'admin' role can delete organisms - require_role(current_user, ["admin", "superuser"]) organism = organism_service.delete_organism(db, grouping_key=grouping_key) if not organism: raise HTTPException(status_code=404, detail="Organism not found") @@ -161,6 +155,7 @@ def delete_organism( @router.post("/bulk-import", response_model=BulkImportResponse) +@policy("organisms:bulk_import") def bulk_import_organisms( *, db: Session = Depends(get_db), @@ -175,7 +170,5 @@ def bulk_import_organisms( The request body should directly match the format of the JSON file in data/unique_organisms.json, which is a dictionary keyed by organism_grouping_key without a wrapping 'organisms' key. """ - # Only users with 'curator' or 'admin' role can import organisms - require_role(current_user, ["curator", "admin"]) result = organism_service.bulk_import_organisms(db, organisms_data=organisms_data) return result diff --git a/app/api/v1/endpoints/projects.py b/app/api/v1/endpoints/projects.py index 46d7171..c8ae031 100644 --- a/app/api/v1/endpoints/projects.py +++ b/app/api/v1/endpoints/projects.py @@ -4,12 +4,9 @@ from fastapi import APIRouter, Depends, HTTPException, Query from sqlalchemy.orm import Session -from app.core.dependencies import ( - get_current_active_user, - get_current_superuser, - get_db, - require_role, -) +from app.core.dependencies import get_current_active_user, get_db +from app.core.pagination import Pagination, apply_pagination, pagination_params +from app.core.policy import policy from app.models.project import Project from app.models.user import User from app.schemas.project import ( @@ -26,19 +23,19 @@ @router.get("/", response_model=List[ProjectSchema]) def read_projects( db: Session = Depends(get_db), - skip: int = 0, - limit: int = 100, + pagination: Pagination = Depends(pagination_params), current_user: User = Depends(get_current_active_user), ) -> Any: """ Retrieve projects. """ # All users can read projects - projects = db.query(Project).offset(skip).limit(limit).all() + projects = apply_pagination(db.query(Project), pagination).all() return projects @router.post("/", response_model=ProjectSchema) +@policy("projects:create") def create_project( *, db: Session = Depends(get_db), @@ -48,9 +45,6 @@ def create_project( """ Create new project. """ - # Only users with 'curator' or 'admin' role can create projects - require_role(current_user, ["curator", "admin"]) - project = Project( project_accession=project_in.project_accession, alias=project_in.alias, @@ -83,6 +77,7 @@ def read_project( @router.put("/{project_id}", response_model=ProjectSchema) +@policy("projects:update") def update_project( *, db: Session = Depends(get_db), @@ -93,9 +88,6 @@ def update_project( """ Update a project. """ - # Only users with 'curator' or 'admin' role can update projects - require_role(current_user, ["curator", "admin"]) - project = db.query(Project).filter(Project.id == project_id).first() if not project: raise HTTPException(status_code=404, detail="Project not found") @@ -111,16 +103,16 @@ def update_project( @router.delete("/{project_id}", response_model=ProjectSchema) +@policy("projects:delete") def delete_project( *, db: Session = Depends(get_db), project_id: UUID, - current_user: User = Depends(get_current_superuser), + current_user: User = Depends(get_current_active_user), ) -> Any: """ Delete a project. """ - # Only superusers can delete projects project = db.query(Project).filter(Project.id == project_id).first() if not project: raise HTTPException(status_code=404, detail="Project not found") diff --git a/app/api/v1/endpoints/read_submissions.py b/app/api/v1/endpoints/read_submissions.py index 05f4aae..1c27a3a 100644 --- a/app/api/v1/endpoints/read_submissions.py +++ b/app/api/v1/endpoints/read_submissions.py @@ -8,12 +8,9 @@ from sqlalchemy.orm import Session from sqlalchemy.orm.attributes import flag_modified -from app.core.dependencies import ( - get_current_active_user, - get_current_superuser, - get_db, - require_role, -) +from app.core.dependencies import get_current_active_user, get_db +from app.core.pagination import Pagination, apply_pagination, pagination_params +from app.core.policy import policy from app.models.read import Read, ReadSubmission from app.models.user import User from app.schemas.common import SubmissionJsonResponse, SubmissionStatus @@ -34,28 +31,26 @@ @router.get("/", response_model=List[ReadSubmissionSchema]) +@policy("read_submissions:read") def read_read_submissions( db: Session = Depends(get_db), - skip: int = 0, - limit: int = 100, + pagination: Pagination = Depends(pagination_params), status: Optional[SubmissionStatus] = Query(None, description="Filter by submission status"), current_user: User = Depends(get_current_active_user), ) -> Any: """ Retrieve read submissions. """ - # Admin, curator, broker and genome_launcher can get submission data - require_role(current_user, ["admin", "curator", "broker", "genome_launcher"]) - query = db.query(ReadSubmission) if status: query = query.filter(ReadSubmission.status == status) - submissions = query.offset(skip).limit(limit).all() + submissions = apply_pagination(query, pagination).all() return submissions @router.get("/{submission_id}", response_model=ReadSubmissionSchema) +@policy("read_submissions:read") def read_read_submissions( submission_id: UUID, db: Session = Depends(get_db), @@ -64,9 +59,6 @@ def read_read_submissions( """ Retrieve read submissions. """ - # Admin, curator, broker and genome_launcher can get submission data - require_role(current_user, ["admin", "curator", "broker", "genome_launcher"]) - submission = db.query(ReadSubmission).filter(ReadSubmission.id == submission_id).first() if not submission: raise HTTPException( @@ -76,6 +68,7 @@ def read_read_submissions( @router.post("/", response_model=ReadSubmissionSchema) +@policy("read_submissions:write") def create_read_submission( *, db: Session = Depends(get_db), @@ -85,9 +78,6 @@ def create_read_submission( """ Create new sample submission. """ - # Only users with 'curator' or 'admin' role can create sample submissions - require_role(current_user, ["curator", "admin"]) - submission = ReadSubmission( read_id=submission_in.read_id, experiment_id=submission_in.experiment_id, diff --git a/app/api/v1/endpoints/reads.py b/app/api/v1/endpoints/reads.py index 13a9bae..996fbe6 100644 --- a/app/api/v1/endpoints/reads.py +++ b/app/api/v1/endpoints/reads.py @@ -8,12 +8,9 @@ from sqlalchemy.orm import Session from sqlalchemy.orm.attributes import flag_modified -from app.core.dependencies import ( - get_current_active_user, - get_current_superuser, - get_db, - require_role, -) +from app.core.dependencies import get_current_active_user, get_db +from app.core.pagination import Pagination, apply_pagination, pagination_params +from app.core.policy import policy from app.models.read import Read, ReadSubmission from app.models.user import User from app.schemas.common import SubmissionJsonResponse, SubmissionStatus @@ -31,8 +28,7 @@ @router.get("/", response_model=List[ReadSchema]) def read_reads( db: Session = Depends(get_db), - skip: int = 0, - limit: int = 100, + pagination: Pagination = Depends(pagination_params), experiment_id: Optional[UUID] = Query(None, description="Filter by experiment ID"), current_user: User = Depends(get_current_active_user), ) -> Any: @@ -44,11 +40,12 @@ def read_reads( if experiment_id: query = query.filter(Read.experiment_id == experiment_id) - reads = query.offset(skip).limit(limit).all() + reads = apply_pagination(query, pagination).all() return reads @router.post("/", response_model=ReadSchema) +@policy("reads:create") def create_read( *, db: Session = Depends(get_db), @@ -58,8 +55,6 @@ def create_read( """ Create new read. """ - # Only users with 'curator' or 'admin' role can create reads - require_role(current_user, ["curator", "admin"]) read_id = uuid.uuid4() # Auto-map from Pydantic schema to Read columns @@ -152,6 +147,7 @@ def read_read( @router.put("/{read_id}", response_model=ReadSchema) +@policy("reads:update") def update_read( *, db: Session = Depends(get_db), @@ -162,8 +158,6 @@ def update_read( """ Update a read. """ - # Only users with 'curator' or 'admin' role can update reads - require_role(current_user, ["curator", "admin"]) read = db.query(Read).filter(Read.id == read_id).first() if not read: @@ -271,16 +265,16 @@ def update_read( @router.delete("/{read_id}", response_model=ReadSchema) +@policy("reads:delete") def delete_read( *, db: Session = Depends(get_db), read_id: UUID, - current_user: User = Depends(get_current_superuser), + current_user: User = Depends(get_current_active_user), ) -> Any: """ Delete a read. """ - # Only superusers can delete reads read = db.query(Read).filter(Read.id == read_id).first() if not read: raise HTTPException(status_code=404, detail="Read not found") diff --git a/app/api/v1/endpoints/sample_submissions.py b/app/api/v1/endpoints/sample_submissions.py index bddd1ad..5ba22fb 100644 --- a/app/api/v1/endpoints/sample_submissions.py +++ b/app/api/v1/endpoints/sample_submissions.py @@ -8,12 +8,9 @@ from sqlalchemy.orm import Session from sqlalchemy.orm.attributes import flag_modified -from app.core.dependencies import ( - get_current_active_user, - get_current_superuser, - get_db, - require_role, -) +from app.core.dependencies import get_current_active_user, get_db +from app.core.pagination import Pagination, apply_pagination, pagination_params +from app.core.policy import policy from app.models.experiment import Experiment from app.models.organism import Organism from app.models.sample import Sample, SampleSubmission @@ -40,10 +37,10 @@ @router.get("/", response_model=List[SampleSubmissionSchema]) +@policy("sample_submissions:read") def read_sample_submissions( db: Session = Depends(get_db), - skip: int = 0, - limit: int = 100, + pagination: Pagination = Depends(pagination_params), status: Optional[SchemaSubmissionStatus] = Query( None, description="Filter by submission status" ), @@ -54,18 +51,16 @@ def read_sample_submissions( """ Retrieve sample submissions. """ - # Admin, curator, broker and genome_launcher can get submission data - require_role(current_user, ["admin", "curator", "broker", "genome_launcher"]) - query = db.query(SampleSubmission) if status: query = query.filter(SampleSubmission.status == status) - submissions = query.offset(skip).limit(limit).all() + submissions = apply_pagination(query, pagination).all() return submissions @router.get("/{submission_id}", response_model=SampleSubmissionSchema) +@policy("sample_submissions:read") def read_sample_submissions( submission_id: UUID, db: Session = Depends(get_db), @@ -74,9 +69,6 @@ def read_sample_submissions( """ Retrieve sample submissions. """ - # Admin, curator, broker and genome_launcher can get submission data - require_role(current_user, ["admin", "curator", "broker", "genome_launcher"]) - submission = db.query(SampleSubmission).filter(SampleSubmission.id == submission_id).first() if not submission: raise HTTPException( @@ -86,6 +78,7 @@ def read_sample_submissions( @router.post("/", response_model=SampleSubmissionSchema) +@policy("sample_submissions:write") def create_sample_submission( *, db: Session = Depends(get_db), @@ -95,9 +88,6 @@ def create_sample_submission( """ Create new sample submission. """ - # Only users with 'curator' or 'admin' role can create sample submissions - require_role(current_user, ["curator", "admin"]) - submission = SampleSubmission( sample_id=submission_in.sample_id, authority=submission_in.authority, diff --git a/app/api/v1/endpoints/samples.py b/app/api/v1/endpoints/samples.py index 5cb6dce..7352901 100644 --- a/app/api/v1/endpoints/samples.py +++ b/app/api/v1/endpoints/samples.py @@ -9,12 +9,9 @@ from sqlalchemy.orm import Session from sqlalchemy.orm.attributes import flag_modified -from app.core.dependencies import ( - get_current_active_user, - get_current_superuser, - get_db, - require_role, -) +from app.core.dependencies import get_current_active_user, get_db +from app.core.pagination import Pagination, apply_pagination, pagination_params +from app.core.policy import policy from app.models.experiment import Experiment from app.models.organism import Organism from app.models.sample import Sample, SampleSubmission @@ -37,8 +34,7 @@ @router.get("/", response_model=List[SampleSchema]) def read_samples( db: Session = Depends(get_db), - skip: int = 0, - limit: int = 100, + pagination: Pagination = Depends(pagination_params), organism_key: Optional[str] = Query(None, description="Filter by organism key"), current_user: User = Depends(get_current_active_user), ) -> Any: @@ -50,7 +46,7 @@ def read_samples( if organism_key: query = query.filter(Sample.organism_key == organism_key) - samples = query.offset(skip).limit(limit).all() + samples = apply_pagination(query, pagination).all() return samples @@ -98,6 +94,7 @@ def get_specimen_by_taxid_and_specimen_id( @router.post("/", response_model=SampleSchema) +@policy("samples:create") def create_sample( *, db: Session = Depends(get_db), @@ -107,9 +104,6 @@ def create_sample( """ Create new sample. """ - # Only users with 'curator' or 'admin' role can create samples - require_role(current_user, ["curator", "admin"]) - sample_data = sample_in.dict(exclude_unset=True) sample_id = uuid.uuid4() @@ -361,6 +355,7 @@ def _create_sample_with_submission( @router.post("/bulk-import-specimens", response_model=BulkImportResponse) +@policy("samples:bulk_import") def bulk_import_specimen_samples( *, db: Session = Depends(get_db), @@ -374,8 +369,6 @@ def bulk_import_specimen_samples( Each sample must have organism_grouping_key and specimen_id. Enforces uniqueness constraint: one specimen per (organism_key, specimen_id). """ - require_role(current_user, ["curator", "admin"]) - # Load the ENA-ATOL mapping file ena_atol_map_path = os.path.join( os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(__file__)))), @@ -463,6 +456,7 @@ def bulk_import_specimen_samples( @router.post("/bulk-import-derived", response_model=BulkImportResponse) +@policy("samples:bulk_import") def bulk_import_derived_samples( *, db: Session = Depends(get_db), @@ -479,8 +473,6 @@ def bulk_import_derived_samples( The parent specimen is looked up by (tax_id, specimen_id) or (organism_key, specimen_id). """ - require_role(current_user, ["curator", "admin"]) - # Load the ENA-ATOL mapping file ena_atol_map_path = os.path.join( os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(__file__)))), @@ -586,6 +578,7 @@ def bulk_import_derived_samples( @router.get("/{sample_id}/prepared-payload", response_model=SubmissionJsonResponse) +@policy("samples:read_sensitive") def get_sample_prepared_payload( *, db: Session = Depends(get_db), @@ -595,8 +588,6 @@ def get_sample_prepared_payload( """ Get prepared_payload for a specific sample. """ - # Admin, curator, broker and genome_launcher can get submission data - require_role(current_user, ["admin", "curator", "broker", "genome_launcher"]) sample_submission = ( db.query(SampleSubmission).filter(SampleSubmission.sample_id == sample_id).first() ) @@ -626,6 +617,7 @@ def read_sample( @router.put("/{sample_id}", response_model=SampleSchema) +@policy("samples:update") def update_sample( *, db: Session = Depends(get_db), @@ -636,9 +628,6 @@ def update_sample( """ Update a sample. """ - # Only users with 'curator' or 'admin' role can update samples - require_role(current_user, ["curator", "admin"]) - try: sample = db.query(Sample).filter(Sample.id == sample_id).first() if not sample: @@ -749,17 +738,16 @@ def update_sample( @router.delete("/{sample_id}", response_model=SampleSchema) +@policy("samples:delete") def delete_sample( *, db: Session = Depends(get_db), sample_id: UUID, - current_user: User = Depends(get_current_superuser), + current_user: User = Depends(get_current_active_user), ) -> Any: """ Delete a sample. """ - # Only superusers can delete samples - require_role(current_user, ["superuser", "admin"]) sample = db.query(Sample).filter(Sample.id == sample_id).first() if not sample: raise HTTPException(status_code=404, detail="Sample not found") @@ -773,6 +761,7 @@ def delete_sample( @router.post("/bulk-import", response_model=BulkImportResponse) +@policy("samples:bulk_import") def bulk_import_samples( *, db: Session = Depends(get_db), @@ -787,9 +776,6 @@ def bulk_import_samples( The request body should directly match the format of the JSON file in data/unique_samples.json, which is a dictionary keyed by bpa_sample_id without a wrapping 'samples' key. """ - # Only users with 'curator' or 'admin' role can import samples - require_role(current_user, ["curator", "admin"]) - # Load the ENA-ATOL mapping file ena_atol_map_path = os.path.join( os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(__file__)))), diff --git a/app/api/v1/endpoints/users.py b/app/api/v1/endpoints/users.py index 1734784..8eecd32 100644 --- a/app/api/v1/endpoints/users.py +++ b/app/api/v1/endpoints/users.py @@ -6,6 +6,8 @@ from sqlalchemy.orm import Session from app.core.dependencies import get_current_active_user +from app.core.pagination import Pagination, apply_pagination, pagination_params +from app.core.policy import policy from app.core.security import get_password_hash from app.db.session import get_db from app.models.user import User @@ -16,10 +18,11 @@ @router.get("/", response_model=List[UserSchema]) +@policy("users:read") def read_users( db: Session = Depends(get_db), - skip: int = 0, - limit: int = 100, + pagination: Pagination = Depends(pagination_params), + current_user: User = Depends(get_current_active_user), ) -> Any: """ Retrieve users. @@ -33,15 +36,17 @@ def read_users( Returns: List[User]: List of users """ - users = db.query(User).offset(skip).limit(limit).all() + users = apply_pagination(db.query(User), pagination).all() return users @router.post("/", response_model=UserSchema) +@policy("users:create") def create_user( *, db: Session = Depends(get_db), user_in: UserCreate, + current_user: User = Depends(get_current_active_user), ) -> Any: """ Create new user. diff --git a/app/core/errors.py b/app/core/errors.py new file mode 100644 index 0000000..dab020f --- /dev/null +++ b/app/core/errors.py @@ -0,0 +1,19 @@ +from __future__ import annotations + +from typing import Any, Dict, Optional + + +class AppError(Exception): + def __init__( + self, + *, + status_code: int, + code: str, + message: str, + details: Optional[Dict[str, Any]] = None, + ) -> None: + super().__init__(message) + self.status_code = status_code + self.code = code + self.message = message + self.details = details or {} diff --git a/app/core/pagination.py b/app/core/pagination.py new file mode 100644 index 0000000..982d6a0 --- /dev/null +++ b/app/core/pagination.py @@ -0,0 +1,23 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from fastapi import Query +from sqlalchemy.orm import Query as SAQuery + + +@dataclass(frozen=True) +class Pagination: + offset: int + limit: int + + +def pagination_params( + offset: int = Query(0, ge=0, description="Number of records to skip"), + limit: int = Query(100, ge=1, le=500, description="Max records to return"), +) -> Pagination: + return Pagination(offset=offset, limit=limit) + + +def apply_pagination(query: SAQuery, pagination: Pagination) -> SAQuery: + return query.offset(pagination.offset).limit(pagination.limit) diff --git a/app/core/policy.py b/app/core/policy.py new file mode 100644 index 0000000..ac12a29 --- /dev/null +++ b/app/core/policy.py @@ -0,0 +1,113 @@ +from __future__ import annotations + +from functools import wraps +from inspect import iscoroutinefunction +from typing import Any, Callable, Dict, List, Optional + +from fastapi import status + +from app.core.errors import AppError +from app.models.user import User + +# Centralized authorization policy +POLICY: Dict[str, List[str]] = { + # Organisms + "organisms:read_sensitive": ["admin", "curator", "broker", "genome_launcher"], + "organisms:create": ["curator", "admin"], + "organisms:update": ["curator", "admin"], + "organisms:delete": ["admin", "superuser"], + "organisms:bulk_import": ["curator", "admin"], + # Samples + "samples:create": ["curator", "admin"], + "samples:update": ["curator", "admin"], + "samples:delete": ["admin", "superuser"], + "samples:bulk_import": ["curator", "admin"], + "samples:read_sensitive": ["admin", "curator", "broker", "genome_launcher"], + # Sample submissions + "sample_submissions:read": ["admin", "curator", "broker", "genome_launcher"], + "sample_submissions:write": ["curator", "admin"], + # Experiments + "experiments:create": ["curator", "admin"], + "experiments:update": ["curator", "admin"], + "experiments:delete": ["admin", "superuser"], + "experiments:bulk_import": ["curator", "admin"], + # Experiment submissions + "experiment_submissions:read": ["admin", "curator", "broker", "genome_launcher"], + "experiment_submissions:write": ["curator", "admin"], + # Reads + "reads:create": ["curator", "admin"], + "reads:update": ["curator", "admin"], + "reads:delete": ["admin", "superuser"], + # Read submissions + "read_submissions:read": ["admin", "curator", "broker", "genome_launcher"], + "read_submissions:write": ["curator", "admin"], + # Projects + "projects:create": ["curator", "admin"], + "projects:update": ["curator", "admin"], + "projects:delete": ["admin", "superuser"], + # Assemblies + "assemblies:write": ["curator", "admin"], + "assemblies:delete": ["admin", "superuser"], + # Genome notes + "genome_notes:write": ["curator", "admin"], + # BPA initiatives + "bpa_initiatives:write": ["curator", "admin"], + # Users + "users:read": ["admin", "superuser"], + "users:create": ["admin", "superuser"], + # Admin + "admin:expire_leases": ["admin", "superuser"], + # Broker + "broker:claim": ["broker"], +} + + +def check_policy(user: User, action: str) -> None: + roles = POLICY.get(action) + if roles is None: + raise AppError( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + code="policy_missing", + message=f"Policy not defined for action '{action}'", + ) + if user.is_superuser: + return + if any(role in user.roles for role in roles): + return + raise AppError( + status_code=status.HTTP_403_FORBIDDEN, + code="forbidden", + message="Not enough permissions", + ) + + +def _check_policy_from_kwargs(action: str, kwargs: Dict[str, Any]) -> None: + current_user: Optional[User] = kwargs.get("current_user") + if current_user is None: + raise AppError( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + code="policy_user_missing", + message="current_user is required for policy checks", + ) + check_policy(current_user, action) + + +def policy(action: str) -> Callable: + def decorator(func: Callable) -> Callable: + if iscoroutinefunction(func): + + @wraps(func) + async def async_wrapper(*args, **kwargs): + _check_policy_from_kwargs(action, kwargs) + return await func(*args, **kwargs) + + return async_wrapper + + @wraps(func) + def wrapper(*args, **kwargs): + _check_policy_from_kwargs(action, kwargs) + return func(*args, **kwargs) + + return wrapper + + return decorator diff --git a/app/core/settings.py b/app/core/settings.py index a0509b3..984e3b8 100644 --- a/app/core/settings.py +++ b/app/core/settings.py @@ -31,7 +31,7 @@ class Settings(BaseSettings): DATABASE_URI: Optional[str] = None # CORS - BACKEND_CORS_ORIGINS: List[str] = ["*"] + BACKEND_CORS_ORIGINS: List[str] = [] # Environment ENVIRONMENT: Optional[str] = None # Options: "dev", "prod" @@ -59,5 +59,17 @@ def __init__(self, **kwargs): f"@{self.POSTGRES_SERVER}:{self.POSTGRES_PORT}/{self.POSTGRES_DB}" ) + # Default permissive CORS for non-prod only when unset + if not self.BACKEND_CORS_ORIGINS and self.ENVIRONMENT != "prod": + self.BACKEND_CORS_ORIGINS = ["*"] + + # Fail fast on missing critical settings + if not self.JWT_SECRET_KEY or not self.JWT_ALGORITHM: + raise ValueError("JWT_SECRET_KEY and JWT_ALGORITHM must be set") + if not self.DATABASE_URI: + raise ValueError("DATABASE_URI must be set (or derived from POSTGRES_* settings)") + if self.ENVIRONMENT == "prod" and self.BACKEND_CORS_ORIGINS == ["*"]: + raise ValueError("BACKEND_CORS_ORIGINS cannot be ['*'] in production") + settings = Settings() diff --git a/app/main.py b/app/main.py index cda7903..4ee8f7b 100644 --- a/app/main.py +++ b/app/main.py @@ -1,18 +1,19 @@ import logging from fastapi import FastAPI +from fastapi.exceptions import RequestValidationError from fastapi.middleware.cors import CORSMiddleware +from fastapi.responses import JSONResponse +from starlette.exceptions import HTTPException as StarletteHTTPException from app.api.v1.api import api_router +from app.core.errors import AppError from app.core.settings import settings # Configure logging based on environment # Default to INFO if ENVIRONMENT not set, DEBUG only if explicitly "dev" log_level = logging.DEBUG if settings.ENVIRONMENT == "dev" else logging.INFO -logging.basicConfig( - level=log_level, - format="%(levelname)s:%(name)s:%(message)s" -) +logging.basicConfig(level=log_level, format="%(levelname)s:%(name)s:%(message)s") # Create FastAPI app app = FastAPI( @@ -36,6 +37,59 @@ app.include_router(api_router, prefix=settings.API_V1_STR) +@app.exception_handler(AppError) +def app_error_handler(_, exc: AppError): + return JSONResponse( + status_code=exc.status_code, + content={ + "error": { + "code": exc.code, + "message": exc.message, + "details": exc.details, + } + }, + ) + + +@app.exception_handler(StarletteHTTPException) +def http_exception_handler(_, exc: StarletteHTTPException): + return JSONResponse( + status_code=exc.status_code, + content={ + "error": { + "code": "http_error", + "message": exc.detail, + "details": {}, + } + }, + ) + + +@app.exception_handler(RequestValidationError) +def validation_exception_handler(_, exc: RequestValidationError): + errors = exc.errors() + # Ensure errors are JSON-serializable (e.g., ValueError in ctx) + for err in errors: + ctx = err.get("ctx") + if ctx and "error" in ctx: + try: + import json # local import to keep module load light + + json.dumps(ctx["error"]) + except Exception: + ctx["error"] = str(ctx["error"]) + return JSONResponse( + status_code=422, + content={ + "error": { + "code": "validation_error", + "message": "Request validation failed", + "details": {"errors": errors}, + } + }, + ) + + @app.get("/") def root(): """ diff --git a/app/models/accession_registry.py b/app/models/accession_registry.py index 748c68c..c95e616 100644 --- a/app/models/accession_registry.py +++ b/app/models/accession_registry.py @@ -1,10 +1,8 @@ import uuid -from datetime import datetime, timezone -from sqlalchemy import BigInteger, Column, DateTime, ForeignKey, String, Text +from sqlalchemy import Column, DateTime, Text, func from sqlalchemy import Enum as SQLAlchemyEnum -from sqlalchemy.dialects.postgresql import JSONB, UUID -from sqlalchemy.orm import relationship +from sqlalchemy.dialects.postgresql import UUID from app.db.session import Base @@ -29,15 +27,13 @@ class AccessionRegistry(Base): nullable=False, ) entity_id = Column(UUID(as_uuid=True), nullable=False) - accepted_at = Column( - DateTime(timezone=True), nullable=False, default=datetime.now(timezone.utc) - ) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + accepted_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) # Table constraints diff --git a/app/models/assembly.py b/app/models/assembly.py index 767c1dc..c1b8a3f 100644 --- a/app/models/assembly.py +++ b/app/models/assembly.py @@ -1,5 +1,4 @@ import uuid -from datetime import datetime, timezone from sqlalchemy import ( BigInteger, @@ -9,8 +8,8 @@ ForeignKey, Integer, PrimaryKeyConstraint, - String, Text, + func, ) from sqlalchemy import Enum as SQLAlchemyEnum from sqlalchemy.dialects.postgresql import JSONB, UUID @@ -59,12 +58,12 @@ class Assembly(Base): description = Column(Text, nullable=True) version = Column(Integer, nullable=False, default=1) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) # Relationships @@ -115,15 +114,15 @@ class AssemblySubmission(Base): response_payload = Column(JSONB, nullable=True) # Metadata - submitted_at = Column(DateTime, nullable=True) + submitted_at = Column(DateTime(timezone=True), nullable=True) submitted_by = Column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=True) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) # Relationships @@ -164,18 +163,16 @@ class AssemblyFile(Base): file_checksum_method = Column(Text, nullable=True, default="MD5") file_format = Column(Text, nullable=True) description = Column(Text, nullable=True) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) # Relationships - assembly = relationship( - "Assembly", backref=backref("files", cascade="all, delete-orphan") - ) + assembly = relationship("Assembly", backref=backref("files", cascade="all, delete-orphan")) class AssemblyRead(Base): diff --git a/app/models/bpa_initiative.py b/app/models/bpa_initiative.py index 87a5f6b..a4fe408 100644 --- a/app/models/bpa_initiative.py +++ b/app/models/bpa_initiative.py @@ -1,8 +1,4 @@ -import uuid -from datetime import datetime, timezone - -from sqlalchemy import Column, DateTime, String, Text -from sqlalchemy.dialects.postgresql import UUID +from sqlalchemy import Column, DateTime, Text, func from app.db.session import Base @@ -19,10 +15,10 @@ class BPAInitiative(Base): project_code = Column(Text, primary_key=True) title = Column(Text, nullable=False) url = Column(Text, nullable=True) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) diff --git a/app/models/broker.py b/app/models/broker.py index 24ed13c..b8bb8c4 100644 --- a/app/models/broker.py +++ b/app/models/broker.py @@ -1,9 +1,7 @@ import uuid -from datetime import datetime, timezone -from sqlalchemy import Column, DateTime, ForeignKey, String, Text +from sqlalchemy import Column, DateTime, ForeignKey, String, Text, func from sqlalchemy.dialects.postgresql import JSONB, UUID -from sqlalchemy.orm import backref, relationship from app.db.session import Base @@ -17,14 +15,14 @@ class SubmissionAttempt(Base): ) campaign_label = Column(Text, nullable=True) status = Column(String, nullable=False, default="processing") - lock_acquired_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) - lock_expires_at = Column(DateTime, nullable=True) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + lock_acquired_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) + lock_expires_at = Column(DateTime(timezone=True), nullable=True) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) @@ -40,4 +38,4 @@ class SubmissionEvent(Base): action = Column(String, nullable=False) # claimed|accepted|rejected|released|expired|progress accession = Column(Text, nullable=True) details = Column(JSONB, nullable=True) - at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) diff --git a/app/models/experiment.py b/app/models/experiment.py index d0e7c3e..876ce7c 100644 --- a/app/models/experiment.py +++ b/app/models/experiment.py @@ -1,7 +1,6 @@ import uuid -from datetime import datetime, timezone -from sqlalchemy import Column, DateTime, ForeignKey, ForeignKeyConstraint, String, Text +from sqlalchemy import Column, DateTime, ForeignKey, ForeignKeyConstraint, Text, func from sqlalchemy import Enum as SQLAlchemyEnum from sqlalchemy.dialects.postgresql import JSONB, UUID from sqlalchemy.orm import backref, relationship @@ -48,12 +47,12 @@ class Experiment(Base): gal = Column(Text, nullable=True) raw_data_release_date = Column(Text, nullable=True) # bpa_json = Column(JSONB, nullable=False) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) # Relationships @@ -110,20 +109,20 @@ class ExperimentSubmission(Base): entity_type_const = Column( Text, nullable=False, default="experiment", server_default="experiment" ) - submitted_at = Column(DateTime, nullable=True) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + submitted_at = Column(DateTime(timezone=True), nullable=True) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) # Broker lease/claim fields attempt_id = Column(UUID(as_uuid=True), nullable=True) finalised_attempt_id = Column(UUID(as_uuid=True), nullable=True) - lock_acquired_at = Column(DateTime, nullable=True) - lock_expires_at = Column(DateTime, nullable=True) + lock_acquired_at = Column(DateTime(timezone=True), nullable=True) + lock_expires_at = Column(DateTime(timezone=True), nullable=True) # Relationships experiment = relationship( diff --git a/app/models/genome_note.py b/app/models/genome_note.py index eda1207..5ee3798 100644 --- a/app/models/genome_note.py +++ b/app/models/genome_note.py @@ -1,7 +1,6 @@ import uuid -from datetime import datetime, timezone -from sqlalchemy import Boolean, Column, DateTime, ForeignKey, Integer, Text, UniqueConstraint +from sqlalchemy import Boolean, Column, DateTime, ForeignKey, Integer, Text, UniqueConstraint, func from sqlalchemy.dialects.postgresql import UUID from sqlalchemy.orm import relationship @@ -37,14 +36,14 @@ class GenomeNote(Base): # Publication status is_published = Column(Boolean, nullable=False, default=False) - published_at = Column(DateTime, nullable=True) + published_at = Column(DateTime(timezone=True), nullable=True) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) # Relationships diff --git a/app/models/organism.py b/app/models/organism.py index dddc1d5..1dfd7f4 100644 --- a/app/models/organism.py +++ b/app/models/organism.py @@ -1,7 +1,6 @@ import uuid -from datetime import datetime, timezone -from sqlalchemy import Column, DateTime, Integer, String, Text +from sqlalchemy import Column, DateTime, Integer, Text, func from sqlalchemy.dialects.postgresql import JSONB, UUID from app.db.session import Base @@ -35,10 +34,10 @@ class Organism(Base): augustus_dataset_name = Column(Text, nullable=True) bpa_json = Column(JSONB, nullable=True) taxonomy_lineage_json = Column(JSONB, nullable=True) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) diff --git a/app/models/project.py b/app/models/project.py index 65e8add..0f40840 100644 --- a/app/models/project.py +++ b/app/models/project.py @@ -1,10 +1,8 @@ import uuid -from datetime import datetime, timezone -from sqlalchemy import Column, DateTime, ForeignKey, String, Text +from sqlalchemy import Column, DateTime, ForeignKey, Text, func from sqlalchemy import Enum as SQLAlchemyEnum from sqlalchemy.dialects.postgresql import JSONB, UUID -from sqlalchemy.orm import relationship from app.db.session import Base @@ -32,7 +30,7 @@ class Project(Base): description = Column(Text, nullable=False) centre_name = Column(Text, nullable=True, default="AToL") study_attributes = Column(JSONB, nullable=True) - submitted_at = Column(DateTime, nullable=True) + submitted_at = Column(DateTime(timezone=True), nullable=True) status = Column( SQLAlchemyEnum( "draft", @@ -49,12 +47,12 @@ class Project(Base): authority = Column( SQLAlchemyEnum("ENA", "NCBI", "DDBJ", name="authority_type"), nullable=False, default="ENA" ) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) @@ -87,12 +85,12 @@ class ProjectSubmission(Base): accession = Column(Text, nullable=True) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) # attempt linkage @@ -100,8 +98,8 @@ class ProjectSubmission(Base): finalised_attempt_id = Column(UUID(as_uuid=True), nullable=True) # broker lease/claim fields - lock_acquired_at = Column(DateTime, nullable=True) - lock_expires_at = Column(DateTime, nullable=True) + lock_acquired_at = Column(DateTime(timezone=True), nullable=True) + lock_expires_at = Column(DateTime(timezone=True), nullable=True) """ diff --git a/app/models/read.py b/app/models/read.py index 7f97f7f..76f1b91 100644 --- a/app/models/read.py +++ b/app/models/read.py @@ -1,5 +1,4 @@ import uuid -from datetime import datetime, timezone from sqlalchemy import ( BigInteger, @@ -8,8 +7,8 @@ DateTime, ForeignKey, ForeignKeyConstraint, - String, Text, + func, ) from sqlalchemy import Enum as SQLAlchemyEnum from sqlalchemy.dialects.postgresql import JSONB, UUID @@ -41,12 +40,12 @@ class Read(Base): read_number = Column(Text, nullable=True) lane_number = Column(Text, nullable=True) # bpa_json = Column(JSONB, nullable=False) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) # Relationships @@ -115,12 +114,12 @@ class ReadSubmission(Base): # Constant to help the composite FK entity_type_const = Column(Text, nullable=False, default="read", server_default="read") - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) # Relationships @@ -137,8 +136,8 @@ class ReadSubmission(Base): # Broker lease/claim fields attempt_id = Column(UUID(as_uuid=True), nullable=True) finalised_attempt_id = Column(UUID(as_uuid=True), nullable=True) - lock_acquired_at = Column(DateTime, nullable=True) - lock_expires_at = Column(DateTime, nullable=True) + lock_acquired_at = Column(DateTime(timezone=True), nullable=True) + lock_expires_at = Column(DateTime(timezone=True), nullable=True) # Table constraints __table_args__ = ( diff --git a/app/models/sample.py b/app/models/sample.py index 6b78ba2..3eaf47c 100644 --- a/app/models/sample.py +++ b/app/models/sample.py @@ -1,7 +1,16 @@ import uuid -from datetime import datetime, timezone -from sqlalchemy import Column, DateTime, Float, ForeignKey, ForeignKeyConstraint, String, Text, text +from sqlalchemy import ( + CheckConstraint, + Column, + DateTime, + Float, + ForeignKey, + ForeignKeyConstraint, + Text, + func, + text, +) from sqlalchemy import Enum as SQLAlchemyEnum from sqlalchemy.dialects.postgresql import JSONB, UUID from sqlalchemy.orm import backref, relationship @@ -65,12 +74,12 @@ class Sample(Base): extensions = Column(JSONB, nullable=True) # bpa_json = Column(JSONB, nullable=False) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) # Relationships @@ -119,20 +128,20 @@ class SampleSubmission(Base): accession = Column(Text, nullable=True) biosample_accession = Column(Text, nullable=True) entity_type_const = Column(Text, nullable=False, default="sample", server_default="sample") - submitted_at = Column(DateTime, nullable=True) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + submitted_at = Column(DateTime(timezone=True), nullable=True) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) # Broker lease/claim fields attempt_id = Column(UUID(as_uuid=True), nullable=True) finalised_attempt_id = Column(UUID(as_uuid=True), nullable=True) - lock_acquired_at = Column(DateTime, nullable=True) - lock_expires_at = Column(DateTime, nullable=True) + lock_acquired_at = Column(DateTime(timezone=True), nullable=True) + lock_expires_at = Column(DateTime(timezone=True), nullable=True) # Relationships sample = relationship( @@ -141,6 +150,12 @@ class SampleSubmission(Base): # Table constraints __table_args__ = ( + CheckConstraint( + "latitude IS NULL OR latitude BETWEEN -90 AND 90", name="chk_sample_latitude" + ), + CheckConstraint( + "longitude IS NULL OR longitude BETWEEN -180 AND 180", name="chk_sample_longitude" + ), # Foreign key constraint for accession registry ForeignKeyConstraint( ["accession", "authority", "entity_type_const", "sample_id"], diff --git a/app/models/token.py b/app/models/token.py index 202107a..d0f6eb3 100644 --- a/app/models/token.py +++ b/app/models/token.py @@ -1,8 +1,6 @@ import uuid -from datetime import datetime, timezone -from typing import Optional -from sqlalchemy import Boolean, Column, DateTime, ForeignKey, String, Text +from sqlalchemy import Boolean, Column, DateTime, ForeignKey, Text, func from sqlalchemy.dialects.postgresql import UUID from sqlalchemy.orm import relationship @@ -21,14 +19,14 @@ class RefreshToken(Base): id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4) token_hash = Column(Text, nullable=False) user_id = Column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=False) - expires_at = Column(DateTime, nullable=False) + expires_at = Column(DateTime(timezone=True), nullable=False) revoked = Column(Boolean, nullable=False, default=False) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) # Relationship with User model diff --git a/app/models/user.py b/app/models/user.py index 8fa75d3..2c70224 100644 --- a/app/models/user.py +++ b/app/models/user.py @@ -1,10 +1,8 @@ import uuid -from datetime import datetime, timezone from typing import List -from sqlalchemy import ARRAY, Boolean, Column, DateTime, String, Text +from sqlalchemy import ARRAY, Boolean, Column, DateTime, Text, func from sqlalchemy.dialects.postgresql import UUID -from sqlalchemy.orm import relationship from app.db.session import Base @@ -26,10 +24,10 @@ class User(Base): roles = Column(ARRAY(Text), nullable=False, default=[]) is_active = Column(Boolean, nullable=False, default=True) is_superuser = Column(Boolean, nullable=False, default=False) - created_at = Column(DateTime, nullable=False, default=datetime.now(timezone.utc)) + created_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now()) updated_at = Column( - DateTime, + DateTime(timezone=True), nullable=False, - default=datetime.now(timezone.utc), - onupdate=datetime.now(timezone.utc), + server_default=func.now(), + onupdate=func.now(), ) diff --git a/app/schemas/assembly.py b/app/schemas/assembly.py index d2ebacb..13f4f92 100644 --- a/app/schemas/assembly.py +++ b/app/schemas/assembly.py @@ -10,6 +10,7 @@ class AssemblyDataTypes(str, Enum): """Enum for assembly data types (sequencing platforms).""" + PACBIO_SMRT = "PACBIO_SMRT" PACBIO_SMRT_HIC = "PACBIO_SMRT_HIC" OXFORD_NANOPORE = "OXFORD_NANOPORE" @@ -20,6 +21,7 @@ class AssemblyDataTypes(str, Enum): class AssemblyFileType(str, Enum): """Enum for assembly file types.""" + FASTA = "FASTA" QC_REPORT = "QC_REPORT" STATISTICS = "STATISTICS" @@ -51,6 +53,22 @@ class AssemblyCreate(AssemblyBase): pass +# Schema for creating assembly from experiments (organism_key derived from tax_id) +class AssemblyCreateFromExperiments(BaseModel): + """Schema for creating assembly from experiments - organism_key is auto-filled from tax_id.""" + + sample_id: UUID + project_id: Optional[UUID] = None + assembly_name: str + assembly_type: str = "clone or isolate" + data_types: Optional[AssemblyDataTypes] = None # Auto-detected, can be overridden + coverage: float + program: str + mingaplength: Optional[float] = None + moleculetype: str = "genomic DNA" + description: Optional[str] = None + + # Schema for updating an existing assembly class AssemblyUpdate(BaseModel): """Schema for updating an existing assembly.""" @@ -155,6 +173,7 @@ class AssemblySubmission(AssemblySubmissionInDBBase): # AssemblyFile schemas # ========================================== + # Base AssemblyFile schema class AssemblyFileBase(BaseModel): """Base AssemblyFile schema with common attributes.""" @@ -173,6 +192,7 @@ class AssemblyFileBase(BaseModel): # Schema for creating a new assembly file class AssemblyFileCreate(AssemblyFileBase): """Schema for creating a new assembly file.""" + pass @@ -205,4 +225,5 @@ class AssemblyFileInDBBase(AssemblyFileBase): # Schema for returning assembly file information class AssemblyFile(AssemblyFileInDBBase): """Schema for returning assembly file information.""" + pass diff --git a/app/services/assembly_helper.py b/app/services/assembly_helper.py index f8b05b4..1b312ea 100644 --- a/app/services/assembly_helper.py +++ b/app/services/assembly_helper.py @@ -1,4 +1,5 @@ """Helper functions for assembly operations.""" + import logging from typing import Dict, List, Set @@ -90,14 +91,12 @@ def get_detected_platforms(experiments: List[Experiment]) -> dict: return { "platforms": list(platforms), "library_strategies": list(library_strategies), - "experiment_count": len(experiments) + "experiment_count": len(experiments), } def generate_assembly_manifest( - organism: Organism, - reads: List[Read], - experiments: List[Experiment] + organism: Organism, reads: List[Read], experiments: List[Experiment] ) -> str: """Generate assembly manifest YAML from organism and reads data. @@ -115,13 +114,17 @@ def generate_assembly_manifest( Returns: YAML string formatted as assembly manifest """ - logger.info(f"Generating manifest for organism: {organism.scientific_name} (tax_id: {organism.tax_id})") + logger.info( + f"Generating manifest for organism: {organism.scientific_name} (tax_id: {organism.tax_id})" + ) logger.info(f"Total experiments: {len(experiments)}, Total reads: {len(reads)}") # Create experiment_id to platform mapping exp_platform_map = {} for exp in experiments: - logger.info(f"Experiment {exp.id}: platform={exp.platform}, library_strategy={exp.library_strategy}") + logger.info( + f"Experiment {exp.id}: platform={exp.platform}, library_strategy={exp.library_strategy}" + ) if exp.platform: exp_platform_map[exp.id] = exp.platform.upper() if exp.library_strategy: @@ -139,40 +142,50 @@ def generate_assembly_manifest( platform = exp_platform_map.get(read.experiment_id, "") library_strategy = exp_platform_map.get(f"{read.experiment_id}_strategy", "") - logger.debug(f"Read {read.id} (file: {read.file_name}): platform={platform}, library_strategy={library_strategy}") + logger.debug( + f"Read {read.id} (file: {read.file_name}): platform={platform}, library_strategy={library_strategy}" + ) # Check for PacBio SMRT reads (only .ccs.bam or hifi_reads.bam) if platform == "PACBIO_SMRT" and read.file_name: if read.file_name.endswith(".ccs.bam") or read.file_name.endswith("hifi_reads.bam"): # TODO remove logging logger.info(f"Adding PacBio read: {read.file_name}") - pacbio_reads.append({ - "file_name": read.file_name, - "file_checksum": read.file_checksum, - "url": read.bioplatforms_url - }) + pacbio_reads.append( + { + "file_name": read.file_name, + "file_checksum": read.file_checksum, + "url": read.bioplatforms_url, + } + ) else: - logger.debug(f"Skipping PacBio read {read.file_name} - doesn't match .ccs.bam or hifi_reads.bam") + logger.debug( + f"Skipping PacBio read {read.file_name} - doesn't match .ccs.bam or hifi_reads.bam" + ) # Check for Hi-C reads (Illumina + Hi-C or WGS library strategy) elif platform == "ILLUMINA" and library_strategy in ("HI-C", "WGS"): # TODO remove logging logger.info(f"Adding Hi-C read: {read.file_name} (library_strategy={library_strategy})") - hic_reads.append({ - "file_name": read.file_name, - "file_checksum": read.file_checksum, - "url": read.bioplatforms_url, - "read_number": read.read_number, - "lane_number": read.lane_number - }) + hic_reads.append( + { + "file_name": read.file_name, + "file_checksum": read.file_checksum, + "url": read.bioplatforms_url, + "read_number": read.read_number, + "lane_number": read.lane_number, + } + ) else: - logger.debug(f"Read {read.file_name} doesn't match criteria: platform={platform}, library_strategy={library_strategy}") + logger.debug( + f"Read {read.file_name} doesn't match criteria: platform={platform}, library_strategy={library_strategy}" + ) # Build manifest structure manifest = { "scientific_name": organism.scientific_name, "taxon_id": organism.tax_id, - "reads": {} + "reads": {}, } if pacbio_reads: diff --git a/app/services/assembly_service.py b/app/services/assembly_service.py index 39cadbd..2b3d3a2 100644 --- a/app/services/assembly_service.py +++ b/app/services/assembly_service.py @@ -44,7 +44,7 @@ def create(self, db: Session, *, obj_in: AssemblyCreate) -> Assembly: next_version = (max_version or 0) + 1 # Create assembly with auto-incremented version - obj_in_data = obj_in.dict() + obj_in_data = obj_in.model_dump() obj_in_data["version"] = next_version db_obj = Assembly(**obj_in_data) @@ -89,7 +89,7 @@ def create_from_experiments( db: Session, *, tax_id: int, - assembly_in: AssemblyCreate, + assembly_in, # AssemblyCreateFromExperiments ) -> tuple[Assembly, dict]: """Create assembly based on experiments for a given tax_id. @@ -99,7 +99,7 @@ def create_from_experiments( Args: db: Database session tax_id: Taxonomy ID of the organism - assembly_in: Assembly creation data (data_types will be auto-determined) + assembly_in: Assembly creation data (organism_key and data_types auto-determined) Returns: Tuple of (created Assembly, platform detection info) @@ -115,7 +115,9 @@ def create_from_experiments( # 2. Get all samples for this organism samples = db.query(Sample).filter(Sample.organism_key == organism.grouping_key).all() if not samples: - raise ValueError(f"No samples found for organism {organism.grouping_key} (tax_id: {tax_id})") + raise ValueError( + f"No samples found for organism {organism.grouping_key} (tax_id: {tax_id})" + ) # 3. Get all experiments for these samples sample_ids = [sample.id for sample in samples] @@ -126,13 +128,15 @@ def create_from_experiments( f"No experiments found for organism {organism.grouping_key} (tax_id: {tax_id})" ) - # 4. Determine data_types from experiments - data_types = determine_assembly_data_types(experiments) + # 4. Determine data_types from experiments (unless explicitly provided) platform_info = get_detected_platforms(experiments) + obj_in_data = assembly_in.model_dump() + + if obj_in_data.get("data_types") is None: + data_types = determine_assembly_data_types(experiments) + obj_in_data["data_types"] = data_types - # 5. Override data_types in assembly_in - obj_in_data = assembly_in.dict() - obj_in_data["data_types"] = data_types + # 5. Add organism_key from tax_id lookup obj_in_data["organism_key"] = organism.grouping_key # 6. Create assembly using standard create method (handles versioning) diff --git a/app/services/broker_service.py b/app/services/broker_service.py new file mode 100644 index 0000000..851d32d --- /dev/null +++ b/app/services/broker_service.py @@ -0,0 +1,52 @@ +from __future__ import annotations + +from datetime import datetime, timezone +from typing import Dict + +from sqlalchemy import update +from sqlalchemy.orm import Session + +from app.models.broker import SubmissionAttempt +from app.models.experiment import ExperimentSubmission +from app.models.project import ProjectSubmission +from app.models.read import ReadSubmission +from app.models.sample import SampleSubmission + + +def expire_leases(db: Session) -> Dict[str, int]: + now = datetime.now(timezone.utc) + + def _expire_submissions(model) -> int: + stmt = ( + update(model) + .where(model.status == "submitting") + .where(model.lock_expires_at.isnot(None)) + .where(model.lock_expires_at < now) + .values( + status="draft", + attempt_id=None, + lock_acquired_at=None, + lock_expires_at=None, + ) + ) + result = db.execute(stmt) + return result.rowcount or 0 + + def _expire_attempts() -> int: + stmt = ( + update(SubmissionAttempt) + .where(SubmissionAttempt.status == "processing") + .where(SubmissionAttempt.lock_expires_at.isnot(None)) + .where(SubmissionAttempt.lock_expires_at < now) + .values(status="expired") + ) + result = db.execute(stmt) + return result.rowcount or 0 + + return { + "project_submissions": _expire_submissions(ProjectSubmission), + "sample_submissions": _expire_submissions(SampleSubmission), + "experiment_submissions": _expire_submissions(ExperimentSubmission), + "read_submissions": _expire_submissions(ReadSubmission), + "attempts": _expire_attempts(), + } diff --git a/app/services/organism_service.py b/app/services/organism_service.py index ed72cb8..1f51b6a 100644 --- a/app/services/organism_service.py +++ b/app/services/organism_service.py @@ -1,7 +1,7 @@ from typing import Any, Dict, List, Optional from uuid import UUID -from sqlalchemy.orm import Session +from sqlalchemy.orm import Session, selectinload from app.models.experiment import Experiment, ExperimentSubmission from app.models.organism import Organism @@ -79,10 +79,13 @@ def get_experiments_for_organism( samples = db.query(Sample.id).filter(Sample.organism_key == grouping_key).all() sample_ids = [sid for (sid,) in samples] - # Load experiments + # Load experiments (eager load reads when requested) experiments: List[Experiment] = [] if sample_ids: - experiments = db.query(Experiment).filter(Experiment.sample_id.in_(sample_ids)).all() + query = db.query(Experiment).filter(Experiment.sample_id.in_(sample_ids)) + if include_reads: + query = query.options(selectinload(Experiment.reads)) + experiments = query.all() # Build response if not include_reads: @@ -90,15 +93,11 @@ def get_experiments_for_organism( return {"grouping_key": grouping_key, "experiments": exp_list} # include_reads = True - exp_ids = [e.id for e in experiments] reads_by_exp: Dict[str, List[Dict[str, Any]]] = {} - if exp_ids: - reads = db.query(Read).filter(Read.experiment_id.in_(exp_ids)).all() - for r in reads: - key = str(r.experiment_id) if r.experiment_id else "null" - if key not in reads_by_exp: - reads_by_exp[key] = [] - reads_by_exp[key].append(self._sa_obj_to_dict(r)) + for e in experiments: + if not e.reads: + continue + reads_by_exp[str(e.id)] = [self._sa_obj_to_dict(r) for r in e.reads] exp_with_reads: List[Dict[str, Any]] = [] for e in experiments: diff --git a/docker-compose.yml b/docker-compose.yml index 8d97018..8e28692 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -9,6 +9,7 @@ services: - ./app:/app/app - ./scripts:/app/scripts - ./alembic:/app/alembic + - ./tests:/app/tests env_file: - .env environment: diff --git a/schema.sql b/schema.sql index 2709d6f..075d106 100644 --- a/schema.sql +++ b/schema.sql @@ -38,8 +38,8 @@ CREATE TABLE users ( roles TEXT[] NOT NULL DEFAULT '{}', is_active BOOLEAN NOT NULL DEFAULT TRUE, is_superuser BOOLEAN NOT NULL DEFAULT FALSE, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW() + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); -- ========================================== @@ -50,10 +50,10 @@ CREATE TABLE refresh_token ( id UUID PRIMARY KEY DEFAULT uuid_generate_v4(), token_hash TEXT NOT NULL, user_id UUID NOT NULL REFERENCES users(id), - expires_at TIMESTAMP WITHOUT TIME ZONE NOT NULL, + expires_at TIMESTAMPTZ NOT NULL, revoked BOOLEAN NOT NULL DEFAULT FALSE, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW() + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); -- ========================================== @@ -86,8 +86,8 @@ CREATE TABLE organism ( augustus_dataset_name TEXT, bpa_json JSONB, taxonomy_lineage_json JSONB, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW() + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); /* -- BPA organism table @@ -95,8 +95,8 @@ CREATE TABLE organism ( id UUID PRIMARY KEY DEFAULT uuid_generate_v4(), organism_id UUID REFERENCES organism(id), bpa_json JSONB, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW() + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); */ @@ -112,8 +112,8 @@ CREATE TABLE accession_registry ( entity_type entity_type NOT NULL, entity_id UUID NOT NULL, accepted_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW(), + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), UNIQUE (authority, entity_type, entity_id), UNIQUE (authority, accession) ); @@ -137,12 +137,12 @@ CREATE TABLE project ( description TEXT NOT NULL, centre_name TEXT, study_attributes JSONB, - submitted_at TIMESTAMP, + submitted_at TIMESTAMPTZ, status submission_status NOT NULL DEFAULT 'draft', authority authority_type NOT NULL DEFAULT 'ENA', -- TODO confirm if we want study attributes, and enforece schema for json (or include as seperate table) - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW() + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); CREATE UNIQUE INDEX uq_one_project_type_per_organism @@ -163,17 +163,17 @@ CREATE TABLE IF NOT EXISTS project_submission ( -- constant to help the composite FK entity_type_const entity_type NOT NULL DEFAULT 'project' CHECK (entity_type_const = 'project'), - submitted_at TIMESTAMP, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW(), + submitted_at TIMESTAMPTZ, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), -- attempt linkage attempt_id UUID, finalised_attempt_id UUID, -- broker lease/claim fields (attempt-scoped) - lock_acquired_at TIMESTAMP, - lock_expires_at TIMESTAMP, + lock_acquired_at TIMESTAMPTZ, + lock_expires_at TIMESTAMPTZ, CONSTRAINT fk_self_project_accession FOREIGN KEY (accession, authority, entity_type_const, project_id) @@ -189,6 +189,8 @@ CREATE UNIQUE INDEX IF NOT EXISTS uq_project_one_accepted -- Broker claim indexes CREATE INDEX IF NOT EXISTS idx_project_submission_attempt ON project_submission (attempt_id); CREATE INDEX IF NOT EXISTS idx_project_submission_finalised_attempt ON project_submission (finalised_attempt_id); +CREATE INDEX IF NOT EXISTS idx_project_submission_status ON project_submission (status); +CREATE INDEX IF NOT EXISTS idx_project_submission_lock_expires_at ON project_submission (lock_expires_at); -- ========================================== -- Sample tables @@ -241,9 +243,11 @@ CREATE TABLE sample ( OR (kind = 'derived' AND derived_from_sample_id IS NOT NULL) ), + CONSTRAINT chk_sample_latitude CHECK (latitude IS NULL OR latitude BETWEEN -90 AND 90), + CONSTRAINT chk_sample_longitude CHECK (longitude IS NULL OR longitude BETWEEN -180 AND 180), extensions JSONB, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW() + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); -- Sample submission table @@ -257,9 +261,9 @@ CREATE TABLE sample_submission ( biosample_accession TEXT, -- TODO undecided whether to keep biosample_accession here or rely on the accession_registry table status submission_status NOT NULL DEFAULT 'draft', - submitted_at TIMESTAMP, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW(), + submitted_at TIMESTAMPTZ, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), -- constant to help the composite FK entity_type_const entity_type NOT NULL DEFAULT 'sample' CHECK (entity_type_const = 'sample'), @@ -269,8 +273,8 @@ CREATE TABLE sample_submission ( finalised_attempt_id UUID, -- broker lease/claim fields (attempt-scoped) - lock_acquired_at TIMESTAMP, - lock_expires_at TIMESTAMP, + lock_acquired_at TIMESTAMPTZ, + lock_expires_at TIMESTAMPTZ, CONSTRAINT fk_self_accession FOREIGN KEY (accession, authority, entity_type_const, sample_id) @@ -287,6 +291,8 @@ CREATE UNIQUE INDEX uq_sample_one_accepted -- Broker claim indexes CREATE INDEX IF NOT EXISTS idx_sample_submission_attempt ON sample_submission (attempt_id); CREATE INDEX IF NOT EXISTS idx_sample_submission_finalised_attempt ON sample_submission (finalised_attempt_id); +CREATE INDEX IF NOT EXISTS idx_sample_submission_status ON sample_submission (status); +CREATE INDEX IF NOT EXISTS idx_sample_submission_lock_expires_at ON sample_submission (lock_expires_at); -- Support parent/child lookups for derived samples CREATE INDEX IF NOT EXISTS idx_sample_derived_from_sample_id ON sample(derived_from_sample_id); @@ -344,8 +350,8 @@ CREATE TABLE experiment ( -- bpa_dataset_id TEXT UNIQUE NOT NULL, extensions JSONB, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW() + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); -- Experiment submission table @@ -370,16 +376,16 @@ CREATE TABLE experiment_submission ( -- constant to help the composite FK entity_type_const entity_type NOT NULL DEFAULT 'experiment' CHECK (entity_type_const = 'experiment'), - submitted_at TIMESTAMP, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW(), + submitted_at TIMESTAMPTZ, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), attempt_id UUID, finalised_attempt_id UUID, -- broker lease/claim fields - lock_acquired_at TIMESTAMP, - lock_expires_at TIMESTAMP, + lock_acquired_at TIMESTAMPTZ, + lock_expires_at TIMESTAMPTZ, -- When accession is present, it must exist in the registry AND map to this same experiment: CONSTRAINT fk_self_accession @@ -399,6 +405,8 @@ CREATE TABLE experiment_submission ( -- Broker lease/claim index CREATE INDEX IF NOT EXISTS idx_experiment_submission_attempt ON experiment_submission (attempt_id); CREATE INDEX IF NOT EXISTS idx_experiment_submission_finalised_attempt ON experiment_submission (finalised_attempt_id); +CREATE INDEX IF NOT EXISTS idx_experiment_submission_status ON experiment_submission (status); +CREATE INDEX IF NOT EXISTS idx_experiment_submission_lock_expires_at ON experiment_submission (lock_expires_at); -- TODO consider if we want to keep track of former submissions that have been replaced/modified CREATE UNIQUE INDEX uq_exp_one_accepted @@ -433,8 +441,8 @@ CREATE TABLE read ( run_read_count TEXT, run_base_count TEXT, extensions JSONB, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW() + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); CREATE TABLE read_submission ( @@ -453,16 +461,16 @@ CREATE TABLE read_submission ( accession TEXT, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW(), + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), -- attempt linkage attempt_id UUID, finalised_attempt_id UUID, -- broker lease/claim fields (attempt-scoped) - lock_acquired_at TIMESTAMP, - lock_expires_at TIMESTAMP, + lock_acquired_at TIMESTAMPTZ, + lock_expires_at TIMESTAMPTZ, -- constant to help the composite FK entity_type_const entity_type NOT NULL DEFAULT 'read' CHECK (entity_type_const = 'read'), @@ -487,6 +495,8 @@ CREATE UNIQUE INDEX uq_read_one_accepted -- removed batch index; attempt-only CREATE INDEX IF NOT EXISTS idx_read_submission_attempt ON read_submission (attempt_id); CREATE INDEX IF NOT EXISTS idx_read_submission_finalised_attempt ON read_submission (finalised_attempt_id); +CREATE INDEX IF NOT EXISTS idx_read_submission_status ON read_submission (status); +CREATE INDEX IF NOT EXISTS idx_read_submission_lock_expires_at ON read_submission (lock_expires_at); -- ========================================== -- Broker Attempt table @@ -497,12 +507,15 @@ CREATE TABLE submission_attempt ( organism_key TEXT REFERENCES organism(grouping_key), campaign_label TEXT, status TEXT NOT NULL DEFAULT 'processing', - lock_acquired_at TIMESTAMP NOT NULL DEFAULT NOW(), - lock_expires_at TIMESTAMP, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW() + lock_acquired_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + lock_expires_at TIMESTAMPTZ, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); +CREATE INDEX IF NOT EXISTS idx_submission_attempt_status ON submission_attempt (status); +CREATE INDEX IF NOT EXISTS idx_submission_attempt_lock_expires_at ON submission_attempt (lock_expires_at); + -- ========================================== -- Submission events (append-only audit trail) -- ========================================== @@ -515,7 +528,7 @@ CREATE TABLE submission_event ( action TEXT NOT NULL CHECK (action IN ('claimed','accepted','rejected','released','expired','progress')), accession TEXT, details JSONB, - at TIMESTAMP NOT NULL DEFAULT NOW() + at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); CREATE INDEX IF NOT EXISTS idx_submission_event_attempt ON submission_event (attempt_id); @@ -545,8 +558,8 @@ CREATE TABLE assembly ( -- Auto-incremented version per (data_types, organism_key, sample_id) version INTEGER NOT NULL DEFAULT 1, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW() + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); CREATE TABLE assembly_file ( @@ -561,8 +574,8 @@ CREATE TABLE assembly_file ( file_format TEXT, description TEXT, metadata JSONB, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW() + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); CREATE INDEX idx_assembly_file_assembly_id ON assembly_file(assembly_id); @@ -590,11 +603,11 @@ CREATE TABLE assembly_submission ( response_payload JSONB, -- Metadata - submitted_at TIMESTAMP, + submitted_at TIMESTAMPTZ, submitted_by UUID REFERENCES users(id), - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW() + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); -- Only one accepted submission per assembly+authority @@ -629,10 +642,10 @@ CREATE TABLE genome_note ( -- Publication status is_published BOOLEAN NOT NULL DEFAULT FALSE, - published_at TIMESTAMP, + published_at TIMESTAMPTZ, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW(), + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), -- Ensure version uniqueness per organism UNIQUE (organism_key, version) @@ -657,8 +670,8 @@ CREATE TABLE bpa_initiative ( project_code TEXT PRIMARY KEY, title TEXT NOT NULL, url TEXT, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW() + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); -- Create indexes for common query patterns diff --git a/scripts/expire_leases.py b/scripts/expire_leases.py new file mode 100644 index 0000000..21fca5d --- /dev/null +++ b/scripts/expire_leases.py @@ -0,0 +1,21 @@ +#!/usr/bin/env python3 +from __future__ import annotations + +from app.db.session import SessionLocal +from app.services.broker_service import expire_leases + + +def main() -> None: + db = SessionLocal() + try: + expired = expire_leases(db) + db.commit() + finally: + db.close() + + for key, count in expired.items(): + print(f"{key}: {count}") + + +if __name__ == "__main__": + main() diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..5e1b0da --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,7 @@ +import sys +from pathlib import Path + +# Ensure project root is on sys.path for test imports (app.*) +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) diff --git a/tests/unit/endpoints/test_bulk_import_samples.py b/tests/unit/endpoints/test_bulk_import_samples.py index c25b5df..d7d6c34 100644 --- a/tests/unit/endpoints/test_bulk_import_samples.py +++ b/tests/unit/endpoints/test_bulk_import_samples.py @@ -128,7 +128,6 @@ def mock_create_sample( return sample, submission monkeypatch.setattr(samples, "_create_sample_with_submission", mock_create_sample) - monkeypatch.setattr(samples, "require_role", lambda current_user, roles: None) # Mock file reading with patch("builtins.open", create=True) as mock_open: @@ -232,7 +231,6 @@ def query(self, model): app.dependency_overrides[samples.get_current_active_user] = _override_user(["admin"]) app.dependency_overrides[samples.get_db] = _override_db(fake_session) - monkeypatch.setattr(samples, "require_role", lambda current_user, roles: None) with patch("builtins.open", create=True): with patch("json.load", return_value={"sample": {}}): @@ -315,7 +313,6 @@ def mock_create_sample( return sample, submission monkeypatch.setattr(samples, "_create_sample_with_submission", mock_create_sample) - monkeypatch.setattr(samples, "require_role", lambda current_user, roles: None) with patch("builtins.open", create=True): with patch("json.load", return_value={"sample": {}}): @@ -404,7 +401,6 @@ def mock_create_sample( return sample, submission monkeypatch.setattr(samples, "_create_sample_with_submission", mock_create_sample) - monkeypatch.setattr(samples, "require_role", lambda current_user, roles: None) with patch("builtins.open", create=True): with patch("json.load", return_value={"sample": {}}): @@ -550,7 +546,6 @@ def mock_create_sample( return sample, submission monkeypatch.setattr(samples, "_create_sample_with_submission", mock_create_sample) - monkeypatch.setattr(samples, "require_role", lambda current_user, roles: None) with patch("builtins.open", create=True): with patch("json.load", return_value={"sample": {}}): @@ -687,7 +682,6 @@ def mock_create_sample( return sample, submission monkeypatch.setattr(samples, "_create_sample_with_submission", mock_create_sample) - monkeypatch.setattr(samples, "require_role", lambda current_user, roles: None) with patch("builtins.open", create=True): with patch("json.load", return_value={"sample": {}}): diff --git a/tests/unit/endpoints/test_endpoints_assemblies.py b/tests/unit/endpoints/test_endpoints_assemblies.py index 883db67..1a0f055 100644 --- a/tests/unit/endpoints/test_endpoints_assemblies.py +++ b/tests/unit/endpoints/test_endpoints_assemblies.py @@ -77,3 +77,200 @@ def test_assemblies_pipeline_inputs_not_found(monkeypatch): resp = client.get("/api/v1/assemblies/pipeline-inputs?organism_grouping_key=missing") assert resp.status_code == 404 + + +def test_create_assembly_from_experiments_success(monkeypatch): + """Test successful assembly creation from experiments.""" + from datetime import datetime, timezone + from uuid import uuid4 + + from app.models.assembly import Assembly + + client = TestClient(app) + + # Create a real Assembly object for proper serialization + mock_assembly = Assembly( + id=uuid4(), + organism_key="test_organism", + sample_id=uuid4(), + assembly_name="Test Assembly", + assembly_type="clone or isolate", + data_types="PACBIO_SMRT", + coverage=50.0, + program="hifiasm", + moleculetype="genomic DNA", + version=1, + created_at=datetime.now(timezone.utc), + updated_at=datetime.now(timezone.utc), + ) + platform_info = { + "platforms": ["PACBIO_SMRT"], + "library_strategies": ["WGS"], + "experiment_count": 1, + } + + monkeypatch.setattr( + assemblies.assembly_service, + "create_from_experiments", + lambda db, tax_id, assembly_in: (mock_assembly, platform_info), + ) + + app.dependency_overrides[assemblies.get_current_active_user] = lambda: SimpleNamespace( + is_active=True, roles=["curator"], is_superuser=False + ) + app.dependency_overrides[assemblies.get_db] = _override_db(_FakeSession()) + + resp = client.post( + "/api/v1/assemblies/from-experiments/172942", + json={ + "sample_id": "550e8400-e29b-41d4-a716-446655440000", + "assembly_name": "Test Assembly", + "assembly_type": "clone or isolate", + "coverage": 50.0, + "program": "hifiasm", + "moleculetype": "genomic DNA", + }, + ) + + assert resp.status_code == 200 + body = resp.json() + # Response is now just the Assembly schema + assert "id" in body + assert body["organism_key"] == "test_organism" + assert body["data_types"] == "PACBIO_SMRT" + assert body["assembly_name"] == "Test Assembly" + + +def test_create_assembly_from_experiments_not_found(monkeypatch): + """Test error when organism not found.""" + client = TestClient(app) + + def mock_create_raises(*args, **kwargs): + raise ValueError("Organism with tax_id 999999 not found") + + monkeypatch.setattr( + assemblies.assembly_service, + "create_from_experiments", + mock_create_raises, + ) + + app.dependency_overrides[assemblies.get_current_active_user] = lambda: SimpleNamespace( + is_active=True, roles=["curator"], is_superuser=False + ) + app.dependency_overrides[assemblies.get_db] = _override_db(_FakeSession()) + + resp = client.post( + "/api/v1/assemblies/from-experiments/999999", + json={ + "sample_id": "550e8400-e29b-41d4-a716-446655440000", + "assembly_name": "Test Assembly", + "assembly_type": "clone or isolate", + "coverage": 50.0, + "program": "hifiasm", + "moleculetype": "genomic DNA", + }, + ) + + assert resp.status_code == 400 + response_data = resp.json() + # Error format may be either {"detail": ...} or {"error": {"message": ...}} + error_msg = response_data.get("detail") or response_data.get("error", {}).get("message", "") + assert "not found" in error_msg + + +def test_get_assembly_manifest_success(monkeypatch): + """Test successful manifest generation.""" + client = TestClient(app) + + # Mock database queries + organism = SimpleNamespace( + grouping_key="test_organism", + scientific_name="Test Species", + tax_id=172942, + ) + sample = SimpleNamespace(id="sample-1", organism_key="test_organism") + experiment = SimpleNamespace( + id="exp-1", + sample_id="sample-1", + platform="PACBIO_SMRT", + library_strategy="WGS", + ) + read = SimpleNamespace( + id="read-1", + experiment_id="exp-1", + file_name="sample.ccs.bam", + file_checksum="abc123", + bioplatforms_url="https://example.com/1", + read_number=None, + lane_number=None, + ) + + class MockQuery: + def __init__(self, return_value): + self.return_value = return_value + + def filter(self, *args, **kwargs): + return self + + def first(self): + return self.return_value if not isinstance(self.return_value, list) else None + + def all(self): + return self.return_value if isinstance(self.return_value, list) else [] + + class MockDB: + def __init__(self): + self.call_count = 0 + + def query(self, model): + self.call_count += 1 + if self.call_count == 1: # organism query + return MockQuery(organism) + elif self.call_count == 2: # samples query + return MockQuery([sample]) + elif self.call_count == 3: # experiments query + return MockQuery([experiment]) + elif self.call_count == 4: # reads query + return MockQuery([read]) + return MockQuery([]) + + app.dependency_overrides[assemblies.get_current_active_user] = lambda: SimpleNamespace( + is_active=True, roles=["admin"], is_superuser=False + ) + app.dependency_overrides[assemblies.get_db] = lambda: MockDB() + + resp = client.get("/api/v1/assemblies/manifest/172942") + + assert resp.status_code == 200 + assert resp.headers["content-type"] == "application/x-yaml" + assert b"scientific_name: Test Species" in resp.content + assert b"taxon_id: 172942" in resp.content + assert b"PACBIO_SMRT:" in resp.content + + +def test_get_assembly_manifest_organism_not_found(): + """Test error when organism not found.""" + client = TestClient(app) + + class MockDB: + def query(self, model): + return self + + def filter(self, *args, **kwargs): + return self + + def first(self): + return None + + app.dependency_overrides[assemblies.get_current_active_user] = lambda: SimpleNamespace( + is_active=True, roles=["admin"], is_superuser=False + ) + app.dependency_overrides[assemblies.get_db] = lambda: MockDB() + + resp = client.get("/api/v1/assemblies/manifest/999999") + + assert resp.status_code == 404 + response_data = resp.json() + # Error format may be either {"detail": ...} or {"error": {"message": ...}} + error_msg = response_data.get("detail") or response_data.get("error", {}).get("message", "") + assert "not found" in error_msg diff --git a/tests/unit/endpoints/test_endpoints_broker.py b/tests/unit/endpoints/test_endpoints_broker.py index cf2bef8..182044f 100644 --- a/tests/unit/endpoints/test_endpoints_broker.py +++ b/tests/unit/endpoints/test_endpoints_broker.py @@ -75,11 +75,15 @@ def flush(self): def commit(self): self.committed = True + def delete(self, obj): + self.added.append(obj) + def execute(self, stmt): self.executed.append(stmt) def test_broker_claim_explicit_ids_empty_lists_returns_empty_response(): + broker_user = SimpleNamespace(is_superuser=False, roles=["broker"]) db = FakeSession( { Organism: [ @@ -100,17 +104,16 @@ def test_broker_claim_explicit_ids_empty_lists_returns_empty_response(): lease_duration_minutes=5, ) - resp = broker.claim_drafts_for_organism( - organism_key="g1", per_type_limit=10, payload=payload, db=db - ) + with pytest.raises(HTTPException) as excinfo: + broker.claim_drafts_for_organism( + organism_key="g1", + per_type_limit=10, + payload=payload, + current_user=broker_user, + db=db, + ) - assert isinstance(resp.attempt_id, type(uuid4())) - assert resp.organism_key == "g1" - assert resp.samples == [] - assert resp.experiments == [] - assert resp.reads == [] - assert resp.projects == [] - assert resp.organism.scientific_name == "Sci" + assert excinfo.value.status_code == 400 def test_broker_renew_attempt_lease_updates_items(): diff --git a/tests/unit/endpoints/test_endpoints_genome_notes.py b/tests/unit/endpoints/test_endpoints_genome_notes.py index a2ab646..47958fd 100644 --- a/tests/unit/endpoints/test_endpoints_genome_notes.py +++ b/tests/unit/endpoints/test_endpoints_genome_notes.py @@ -9,6 +9,12 @@ class _FakeSession: + def __init__(self, note=None): + self._note = note + self.added = [] + self.committed = False + self.deleted = [] + def query(self, *_): return self @@ -16,7 +22,25 @@ def filter(self, *_a, **_k): return self def first(self): - return None + return self._note + + def add(self, obj): + if getattr(obj, "id", None) is None: + obj.id = uuid.uuid4() + if getattr(obj, "created_at", None) is None: + obj.created_at = datetime.now(timezone.utc) + if getattr(obj, "updated_at", None) is None: + obj.updated_at = obj.created_at + self.added.append(obj) + + def commit(self): + self.committed = True + + def refresh(self, _obj): + pass + + def delete(self, obj): + self.deleted.append(obj) def _override_db(fake): @@ -48,27 +72,24 @@ def test_genome_note_not_found(self): def test_genome_note_found(self, monkeypatch): client = TestClient(app) app.dependency_overrides[genome_notes.get_current_active_user] = _override_user - app.dependency_overrides[genome_notes.get_db] = _override_db(_FakeSession()) note_id = uuid.uuid4() assembly_id = uuid.uuid4() now = datetime.now(timezone.utc) - fake_note = { - "id": str(note_id), - "organism_key": "test_organism", - "assembly_id": str(assembly_id), - "version": 1, - "title": "Test Note", - "note_url": "https://example.com/note", - "is_published": False, - "published_at": None, - "created_at": now.isoformat(), - "updated_at": now.isoformat(), - } - - fake_service = SimpleNamespace(get=lambda db, note_id: fake_note) - monkeypatch.setattr(genome_notes, "genome_note_service", fake_service) + fake_note = SimpleNamespace( + id=note_id, + organism_key="test_organism", + assembly_id=assembly_id, + version=1, + title="Test Note", + note_url="https://example.com/note", + is_published=False, + published_at=None, + created_at=now, + updated_at=now, + ) + app.dependency_overrides[genome_notes.get_db] = _override_db(_FakeSession(fake_note)) resp = client.get(f"/api/v1/genome-notes/{note_id}") assert resp.status_code == 200 @@ -84,7 +105,7 @@ def test_list_genome_notes_empty(self, monkeypatch): app.dependency_overrides[genome_notes.get_current_active_user] = _override_user app.dependency_overrides[genome_notes.get_db] = _override_db(_FakeSession()) - fake_service = SimpleNamespace(list_genome_notes=lambda db, skip=0, limit=100: []) + fake_service = SimpleNamespace(get_multi_with_filters=lambda db, **_: []) monkeypatch.setattr(genome_notes, "genome_note_service", fake_service) resp = client.get("/api/v1/genome-notes/") @@ -124,7 +145,7 @@ def test_list_genome_notes_with_data(self, monkeypatch): }, ] - fake_service = SimpleNamespace(list_genome_notes=lambda db, skip=0, limit=100: fake_notes) + fake_service = SimpleNamespace(get_multi_with_filters=lambda db, **_: fake_notes) monkeypatch.setattr(genome_notes, "genome_note_service", fake_service) resp = client.get("/api/v1/genome-notes/") @@ -144,23 +165,11 @@ def test_create_genome_note_success(self, monkeypatch): assembly_id = uuid.uuid4() now = datetime.now(timezone.utc) - def fake_create(db, genome_note_in): - return { - "id": str(note_id), - "organism_key": genome_note_in.organism_key, - "assembly_id": str(genome_note_in.assembly_id), - "version": 1, - "title": genome_note_in.title, - "note_url": genome_note_in.note_url, - "is_published": False, - "published_at": None, - "created_at": now.isoformat(), - "updated_at": now.isoformat(), - } - - fake_service = SimpleNamespace(create=fake_create) - monkeypatch.setattr(genome_notes, "genome_note_service", fake_service) - monkeypatch.setattr(genome_notes, "require_role", lambda current_user, roles: None) + monkeypatch.setattr( + genome_notes, + "genome_note_service", + SimpleNamespace(get_next_version=lambda db, organism_key: 1), + ) payload = { "organism_key": "test_organism", @@ -183,23 +192,11 @@ def test_create_genome_note_auto_version(self, monkeypatch): assembly_id = uuid.uuid4() now = datetime.now(timezone.utc) - def fake_create(db, genome_note_in): - return { - "id": str(uuid.uuid4()), - "organism_key": genome_note_in.organism_key, - "assembly_id": str(genome_note_in.assembly_id), - "version": 3, - "title": genome_note_in.title, - "note_url": genome_note_in.note_url, - "is_published": False, - "published_at": None, - "created_at": now.isoformat(), - "updated_at": now.isoformat(), - } - - fake_service = SimpleNamespace(create=fake_create) - monkeypatch.setattr(genome_notes, "genome_note_service", fake_service) - monkeypatch.setattr(genome_notes, "require_role", lambda current_user, roles: None) + monkeypatch.setattr( + genome_notes, + "genome_note_service", + SimpleNamespace(get_next_version=lambda db, organism_key: 3), + ) payload = { "organism_key": "existing_organism", @@ -219,28 +216,21 @@ class TestUpdateGenomeNote: def test_update_genome_note_success(self, monkeypatch): client = TestClient(app) app.dependency_overrides[genome_notes.get_current_active_user] = _override_user - app.dependency_overrides[genome_notes.get_db] = _override_db(_FakeSession()) - note_id = uuid.uuid4() now = datetime.now(timezone.utc) - - def fake_update(db, note_id, genome_note_in): - return { - "id": str(note_id), - "organism_key": "test_organism", - "assembly_id": str(uuid.uuid4()), - "version": 1, - "title": genome_note_in.title or "Original Title", - "note_url": genome_note_in.note_url or "https://example.com/original", - "is_published": False, - "published_at": None, - "created_at": now.isoformat(), - "updated_at": now.isoformat(), - } - - fake_service = SimpleNamespace(update=fake_update) - monkeypatch.setattr(genome_notes, "genome_note_service", fake_service) - monkeypatch.setattr(genome_notes, "require_role", lambda current_user, roles: None) + fake_note = SimpleNamespace( + id=note_id, + organism_key="test_organism", + assembly_id=uuid.uuid4(), + version=1, + title="Original Title", + note_url="https://example.com/original", + is_published=False, + published_at=None, + created_at=now, + updated_at=now, + ) + app.dependency_overrides[genome_notes.get_db] = _override_db(_FakeSession(fake_note)) payload = {"title": "Updated Title"} @@ -255,10 +245,6 @@ def test_update_genome_note_not_found(self, monkeypatch): note_id = uuid.uuid4() - fake_service = SimpleNamespace(update=lambda db, note_id, genome_note_in: None) - monkeypatch.setattr(genome_notes, "genome_note_service", fake_service) - monkeypatch.setattr(genome_notes, "require_role", lambda current_user, roles: None) - payload = {"title": "Updated Title"} resp = client.put(f"/api/v1/genome-notes/{note_id}", json=payload) @@ -271,13 +257,21 @@ class TestDeleteGenomeNote: def test_delete_genome_note_success(self, monkeypatch): client = TestClient(app) app.dependency_overrides[genome_notes.get_current_active_user] = _override_admin_user - app.dependency_overrides[genome_notes.get_db] = _override_db(_FakeSession()) - note_id = uuid.uuid4() - - fake_service = SimpleNamespace(delete=lambda db, note_id: {"id": str(note_id)}) - monkeypatch.setattr(genome_notes, "genome_note_service", fake_service) - monkeypatch.setattr(genome_notes, "require_role", lambda current_user, roles: None) + now = datetime.now(timezone.utc) + fake_note = SimpleNamespace( + id=note_id, + organism_key="test_organism", + assembly_id=uuid.uuid4(), + version=1, + title="Original Title", + note_url="https://example.com/original", + is_published=False, + published_at=None, + created_at=now, + updated_at=now, + ) + app.dependency_overrides[genome_notes.get_db] = _override_db(_FakeSession(fake_note)) resp = client.delete(f"/api/v1/genome-notes/{note_id}") assert resp.status_code == 200 @@ -289,10 +283,6 @@ def test_delete_genome_note_not_found(self, monkeypatch): note_id = uuid.uuid4() - fake_service = SimpleNamespace(delete=lambda db, note_id: None) - monkeypatch.setattr(genome_notes, "genome_note_service", fake_service) - monkeypatch.setattr(genome_notes, "require_role", lambda current_user, roles: None) - resp = client.delete(f"/api/v1/genome-notes/{note_id}") assert resp.status_code == 404 @@ -324,7 +314,6 @@ def fake_publish(db, note_id): fake_service = SimpleNamespace(publish_genome_note=fake_publish) monkeypatch.setattr(genome_notes, "genome_note_service", fake_service) - monkeypatch.setattr(genome_notes, "require_role", lambda current_user, roles: None) resp = client.post(f"/api/v1/genome-notes/{note_id}/publish") assert resp.status_code == 200 @@ -346,11 +335,10 @@ def fake_publish(db, note_id): fake_service = SimpleNamespace(publish_genome_note=fake_publish) monkeypatch.setattr(genome_notes, "genome_note_service", fake_service) - monkeypatch.setattr(genome_notes, "require_role", lambda current_user, roles: None) resp = client.post(f"/api/v1/genome-notes/{note_id}/publish") assert resp.status_code == 409 - assert "already has a published genome note" in resp.json()["detail"] + assert "already has a published genome note" in resp.json()["error"]["message"] def test_publish_genome_note_not_found(self, monkeypatch): client = TestClient(app) @@ -364,7 +352,6 @@ def fake_publish(db, note_id): fake_service = SimpleNamespace(publish_genome_note=fake_publish) monkeypatch.setattr(genome_notes, "genome_note_service", fake_service) - monkeypatch.setattr(genome_notes, "require_role", lambda current_user, roles: None) resp = client.post(f"/api/v1/genome-notes/{note_id}/publish") assert resp.status_code == 404 @@ -397,7 +384,6 @@ def fake_unpublish(db, note_id): fake_service = SimpleNamespace(unpublish_genome_note=fake_unpublish) monkeypatch.setattr(genome_notes, "genome_note_service", fake_service) - monkeypatch.setattr(genome_notes, "require_role", lambda current_user, roles: None) resp = client.post(f"/api/v1/genome-notes/{note_id}/unpublish") assert resp.status_code == 200 @@ -416,7 +402,6 @@ def fake_unpublish(db, note_id): fake_service = SimpleNamespace(unpublish_genome_note=fake_unpublish) monkeypatch.setattr(genome_notes, "genome_note_service", fake_service) - monkeypatch.setattr(genome_notes, "require_role", lambda current_user, roles: None) resp = client.post(f"/api/v1/genome-notes/{note_id}/unpublish") assert resp.status_code == 404 diff --git a/tests/unit/endpoints/test_endpoints_organisms.py b/tests/unit/endpoints/test_endpoints_organisms.py index f884f6d..06a02ab 100644 --- a/tests/unit/endpoints/test_endpoints_organisms.py +++ b/tests/unit/endpoints/test_endpoints_organisms.py @@ -13,9 +13,27 @@ def _override_user(): class _FakeSession: + def __init__(self, organisms=None): + self._organisms = organisms or [] + def query(self, *_): return self + def offset(self, *_): + return self + + def limit(self, *_): + return self + + def all(self): + return self._organisms + + def filter(self, *_args, **_kwargs): + return self + + def first(self): + return None + def _override_db(fake): def _gen(): @@ -28,7 +46,6 @@ def test_organisms_list_and_not_found(monkeypatch): client = TestClient(app) app.dependency_overrides[organisms.get_current_active_user] = _override_user - app.dependency_overrides[organisms.get_db] = _override_db(_FakeSession()) now = datetime.now(timezone.utc) base_org = { @@ -54,11 +71,12 @@ def test_organisms_list_and_not_found(monkeypatch): "updated_at": now, } - fake_service = SimpleNamespace( - list_organisms=lambda db, skip=0, limit=100: [base_org], - get_by_grouping_key=lambda db, grouping_key: None, + app.dependency_overrides[organisms.get_db] = _override_db(_FakeSession([base_org])) + monkeypatch.setattr( + organisms, + "organism_service", + SimpleNamespace(get_by_grouping_key=lambda db, grouping_key: None), ) - monkeypatch.setattr(organisms, "organism_service", fake_service) resp = client.get("/api/v1/organisms") assert resp.status_code == 200 @@ -87,8 +105,6 @@ def test_create_organism(monkeypatch): } ) monkeypatch.setattr(organisms, "organism_service", fake_service) - monkeypatch.setattr(organisms, "require_role", lambda current_user, roles: None) - payload = { "grouping_key": "g1", "tax_id": 1, diff --git a/tests/unit/endpoints/test_endpoints_users.py b/tests/unit/endpoints/test_endpoints_users.py index b3e437e..0ed6d47 100644 --- a/tests/unit/endpoints/test_endpoints_users.py +++ b/tests/unit/endpoints/test_endpoints_users.py @@ -18,6 +18,15 @@ def _jwt_settings(monkeypatch): monkeypatch.setattr(settings, "JWT_ACCESS_TOKEN_EXPIRE_MINUTES", 30) +@pytest.fixture(autouse=True) +def _admin_user_override(): + app.dependency_overrides[users.get_current_active_user] = lambda: SimpleNamespace( + is_superuser=False, roles=["admin"], is_active=True + ) + yield + app.dependency_overrides.pop(users.get_current_active_user, None) + + class _FakeCreateSession: """Fake session for create_user that returns successive results for first() calls.""" diff --git a/tests/unit/services/test_assembly_helper.py b/tests/unit/services/test_assembly_helper.py new file mode 100644 index 0000000..eb30bdb --- /dev/null +++ b/tests/unit/services/test_assembly_helper.py @@ -0,0 +1,326 @@ +"""Tests for assembly helper functions.""" + +from unittest.mock import Mock + +import pytest + +from app.models.experiment import Experiment +from app.models.organism import Organism +from app.models.read import Read +from app.schemas.assembly import AssemblyDataTypes +from app.services.assembly_helper import ( + determine_assembly_data_types, + generate_assembly_manifest, + get_detected_platforms, +) + + +class TestDetermineAssemblyDataTypes: + """Tests for determine_assembly_data_types function.""" + + def test_pacbio_only(self): + """Test detection of PacBio SMRT only.""" + experiments = [ + Mock(platform="PACBIO_SMRT", library_strategy="WGS"), + ] + result = determine_assembly_data_types(experiments) + assert result == AssemblyDataTypes.PACBIO_SMRT + + def test_oxford_nanopore_only(self): + """Test detection of Oxford Nanopore only.""" + experiments = [ + Mock(platform="OXFORD_NANOPORE", library_strategy="WGS"), + ] + result = determine_assembly_data_types(experiments) + assert result == AssemblyDataTypes.OXFORD_NANOPORE + + def test_illumina_hic(self): + """Test that Illumina-only raises error (no long-read platform).""" + experiments = [ + Mock(platform="ILLUMINA", library_strategy="Hi-C"), + ] + with pytest.raises(ValueError, match="No valid sequencing platforms detected"): + determine_assembly_data_types(experiments) + + def test_illumina_wgs_treated_as_hic(self): + """Test that ILLUMINA + WGS is treated as Hi-C.""" + experiments = [ + Mock(platform="PACBIO_SMRT", library_strategy="WGS"), + Mock(platform="ILLUMINA", library_strategy="WGS"), + ] + result = determine_assembly_data_types(experiments) + assert result == AssemblyDataTypes.PACBIO_SMRT_HIC + + def test_pacbio_and_hic(self): + """Test detection of PacBio + Hi-C combination.""" + experiments = [ + Mock(platform="PACBIO_SMRT", library_strategy="WGS"), + Mock(platform="ILLUMINA", library_strategy="Hi-C"), + ] + result = determine_assembly_data_types(experiments) + assert result == AssemblyDataTypes.PACBIO_SMRT_HIC + + def test_nanopore_and_hic(self): + """Test detection of Oxford Nanopore + Hi-C combination.""" + experiments = [ + Mock(platform="OXFORD_NANOPORE", library_strategy="WGS"), + Mock(platform="ILLUMINA", library_strategy="WGS"), + ] + result = determine_assembly_data_types(experiments) + assert result == AssemblyDataTypes.OXFORD_NANOPORE_HIC + + def test_pacbio_and_nanopore(self): + """Test detection of PacBio + Oxford Nanopore combination.""" + experiments = [ + Mock(platform="PACBIO_SMRT", library_strategy="WGS"), + Mock(platform="OXFORD_NANOPORE", library_strategy="WGS"), + ] + result = determine_assembly_data_types(experiments) + assert result == AssemblyDataTypes.PACBIO_SMRT_OXFORD_NANOPORE + + def test_all_three_platforms(self): + """Test detection of all three platform types.""" + experiments = [ + Mock(platform="PACBIO_SMRT", library_strategy="WGS"), + Mock(platform="OXFORD_NANOPORE", library_strategy="WGS"), + Mock(platform="ILLUMINA", library_strategy="Hi-C"), + ] + result = determine_assembly_data_types(experiments) + assert result == AssemblyDataTypes.PACBIO_SMRT_OXFORD_NANOPORE_HIC + + def test_case_insensitive_platform(self): + """Test that platform detection is case-insensitive.""" + experiments = [ + Mock(platform="pacbio_smrt", library_strategy="WGS"), + ] + result = determine_assembly_data_types(experiments) + assert result == AssemblyDataTypes.PACBIO_SMRT + + def test_case_insensitive_library_strategy(self): + """Test that library strategy detection is case-insensitive.""" + experiments = [ + Mock(platform="OXFORD_NANOPORE", library_strategy="WGS"), + Mock(platform="ILLUMINA", library_strategy="wgs"), + ] + result = determine_assembly_data_types(experiments) + assert result == AssemblyDataTypes.OXFORD_NANOPORE_HIC + + def test_no_valid_platforms_raises_error(self): + """Test that no valid platforms raises ValueError.""" + experiments = [ + Mock(platform="UNKNOWN", library_strategy="WGS"), + ] + with pytest.raises(ValueError, match="No valid sequencing platforms detected"): + determine_assembly_data_types(experiments) + + def test_empty_experiments_raises_error(self): + """Test that empty experiments list raises ValueError.""" + with pytest.raises(ValueError, match="No valid sequencing platforms detected"): + determine_assembly_data_types([]) + + def test_none_platform_ignored(self): + """Test that experiments with None platform are ignored.""" + experiments = [ + Mock(platform=None, library_strategy="WGS"), + Mock(platform="PACBIO_SMRT", library_strategy="WGS"), + ] + result = determine_assembly_data_types(experiments) + assert result == AssemblyDataTypes.PACBIO_SMRT + + +class TestGetDetectedPlatforms: + """Tests for get_detected_platforms function.""" + + def test_single_platform(self): + """Test detection of single platform.""" + experiments = [ + Mock(platform="PACBIO_SMRT", library_strategy="WGS"), + ] + result = get_detected_platforms(experiments) + assert result["platforms"] == ["PACBIO_SMRT"] + assert result["library_strategies"] == ["WGS"] + assert result["experiment_count"] == 1 + + def test_multiple_platforms(self): + """Test detection of multiple platforms.""" + experiments = [ + Mock(platform="PACBIO_SMRT", library_strategy="WGS"), + Mock(platform="ILLUMINA", library_strategy="Hi-C"), + ] + result = get_detected_platforms(experiments) + assert set(result["platforms"]) == {"PACBIO_SMRT", "ILLUMINA"} + assert set(result["library_strategies"]) == {"WGS", "Hi-C"} + assert result["experiment_count"] == 2 + + def test_empty_experiments(self): + """Test with empty experiments list.""" + result = get_detected_platforms([]) + assert result["platforms"] == [] + assert result["library_strategies"] == [] + assert result["experiment_count"] == 0 + + +class TestGenerateAssemblyManifest: + """Tests for generate_assembly_manifest function.""" + + def test_pacbio_reads_filtered_by_extension(self): + """Test that only .ccs.bam and hifi_reads.bam files are included for PacBio.""" + organism = Mock(scientific_name="Test Species", tax_id=12345) + experiments = [Mock(id="exp1", platform="PACBIO_SMRT", library_strategy="WGS")] + reads = [ + Mock( + id="r1", + experiment_id="exp1", + file_name="sample.ccs.bam", + file_checksum="abc123", + bioplatforms_url="https://example.com/1", + read_number=None, + lane_number=None, + ), + Mock( + id="r2", + experiment_id="exp1", + file_name="sample.hifi_reads.bam", + file_checksum="def456", + bioplatforms_url="https://example.com/2", + read_number=None, + lane_number=None, + ), + Mock( + id="r3", + experiment_id="exp1", + file_name="sample.subreads.bam", + file_checksum="ghi789", + bioplatforms_url="https://example.com/3", + read_number=None, + lane_number=None, + ), + ] + + result = generate_assembly_manifest(organism, reads, experiments) + + assert "PACBIO_SMRT:" in result + assert "sample.ccs.bam" in result + assert "sample.hifi_reads.bam" in result + assert "sample.subreads.bam" not in result + + def test_hic_reads_include_metadata(self): + """Test that Hi-C reads include read_number and lane_number.""" + organism = Mock(scientific_name="Test Species", tax_id=12345) + experiments = [Mock(id="exp1", platform="ILLUMINA", library_strategy="Hi-C")] + reads = [ + Mock( + id="r1", + experiment_id="exp1", + file_name="hic_R1.fastq.gz", + file_checksum="abc123", + bioplatforms_url="https://example.com/1", + read_number="1", + lane_number="001", + ), + ] + + result = generate_assembly_manifest(organism, reads, experiments) + + assert "Hi-C:" in result + assert "hic_R1.fastq.gz" in result + assert "read_number: '1'" in result + assert "lane_number: '001'" in result + + def test_wgs_treated_as_hic(self): + """Test that ILLUMINA + WGS is treated as Hi-C.""" + organism = Mock(scientific_name="Test Species", tax_id=12345) + experiments = [Mock(id="exp1", platform="ILLUMINA", library_strategy="WGS")] + reads = [ + Mock( + id="r1", + experiment_id="exp1", + file_name="sample_R1.fastq.gz", + file_checksum="abc123", + bioplatforms_url="https://example.com/1", + read_number="1", + lane_number="001", + ), + ] + + result = generate_assembly_manifest(organism, reads, experiments) + + assert "Hi-C:" in result + assert "sample_R1.fastq.gz" in result + + def test_empty_reads_dict(self): + """Test that empty reads result in empty reads dict.""" + organism = Mock(scientific_name="Test Species", tax_id=12345) + experiments = [Mock(id="exp1", platform="UNKNOWN", library_strategy="WGS")] + reads = [] + + result = generate_assembly_manifest(organism, reads, experiments) + + assert "reads: {}" in result + + def test_organism_metadata_included(self): + """Test that organism metadata is included in manifest.""" + organism = Mock(scientific_name="Saiphos equalis", tax_id=172942) + experiments = [] + reads = [] + + result = generate_assembly_manifest(organism, reads, experiments) + + assert "scientific_name: Saiphos equalis" in result + assert "taxon_id: 172942" in result + + def test_reads_without_experiment_id_skipped(self): + """Test that reads without experiment_id are skipped.""" + organism = Mock(scientific_name="Test Species", tax_id=12345) + experiments = [Mock(id="exp1", platform="PACBIO_SMRT", library_strategy="WGS")] + reads = [ + Mock( + id="r1", + experiment_id=None, + file_name="sample.ccs.bam", + file_checksum="abc123", + bioplatforms_url="https://example.com/1", + read_number=None, + lane_number=None, + ), + ] + + result = generate_assembly_manifest(organism, reads, experiments) + + assert "sample.ccs.bam" not in result + assert "reads: {}" in result + + def test_multiple_platform_types(self): + """Test manifest with both PacBio and Hi-C reads.""" + organism = Mock(scientific_name="Test Species", tax_id=12345) + experiments = [ + Mock(id="exp1", platform="PACBIO_SMRT", library_strategy="WGS"), + Mock(id="exp2", platform="ILLUMINA", library_strategy="Hi-C"), + ] + reads = [ + Mock( + id="r1", + experiment_id="exp1", + file_name="sample.ccs.bam", + file_checksum="abc123", + bioplatforms_url="https://example.com/1", + read_number=None, + lane_number=None, + ), + Mock( + id="r2", + experiment_id="exp2", + file_name="hic_R1.fastq.gz", + file_checksum="def456", + bioplatforms_url="https://example.com/2", + read_number="1", + lane_number="001", + ), + ] + + result = generate_assembly_manifest(organism, reads, experiments) + + assert "PACBIO_SMRT:" in result + assert "Hi-C:" in result + assert "sample.ccs.bam" in result + assert "hic_R1.fastq.gz" in result diff --git a/tests/unit/services/test_assembly_service.py b/tests/unit/services/test_assembly_service.py new file mode 100644 index 0000000..c39b4b3 --- /dev/null +++ b/tests/unit/services/test_assembly_service.py @@ -0,0 +1,324 @@ +"""Tests for assembly service.""" + +import uuid +from unittest.mock import MagicMock, Mock, patch + +import pytest +from sqlalchemy.orm import Session + +from app.models.assembly import Assembly +from app.models.experiment import Experiment +from app.models.organism import Organism +from app.models.sample import Sample +from app.schemas.assembly import AssemblyCreate, AssemblyCreateFromExperiments, AssemblyDataTypes +from app.services.assembly_service import AssemblyService + + +@pytest.fixture +def mock_db(): + """Create a mock database session.""" + return MagicMock(spec=Session) + + +@pytest.fixture +def assembly_service(): + """Create an AssemblyService instance.""" + return AssemblyService(Assembly) + + +@pytest.fixture +def sample_assembly_create(): + """Create a sample AssemblyCreate schema.""" + return AssemblyCreate( + organism_key="test_organism", + sample_id=uuid.uuid4(), + project_id=uuid.uuid4(), + assembly_name="Test Assembly", + assembly_type="clone or isolate", + data_types=AssemblyDataTypes.PACBIO_SMRT, + coverage=50.0, + program="hifiasm", + moleculetype="genomic DNA", + ) + + +class TestCreateAssembly: + """Tests for create method with auto-increment versioning.""" + + def test_create_first_version(self, mock_db, assembly_service, sample_assembly_create): + """Test creating first version (version 1).""" + # Mock query to return None (no existing versions) + mock_query = Mock() + mock_query.filter.return_value.scalar.return_value = None + mock_db.query.return_value = mock_query + + # Mock the assembly object that will be created + created_assembly = Assembly( + id=uuid.uuid4(), + organism_key=sample_assembly_create.organism_key, + sample_id=sample_assembly_create.sample_id, + data_types=sample_assembly_create.data_types, + version=1, + ) + mock_db.add = Mock() + mock_db.commit = Mock() + mock_db.refresh = Mock() + + with patch.object(Assembly, "__init__", return_value=None): + result = assembly_service.create(mock_db, obj_in=sample_assembly_create) + + # Verify version was set to 1 + mock_db.add.assert_called_once() + mock_db.commit.assert_called_once() + + def test_create_increments_version(self, mock_db, assembly_service, sample_assembly_create): + """Test that version increments from existing max version.""" + # Mock query to return existing max version of 3 + mock_query = Mock() + mock_query.filter.return_value.scalar.return_value = 3 + mock_db.query.return_value = mock_query + + mock_db.add = Mock() + mock_db.commit = Mock() + mock_db.refresh = Mock() + + with patch.object(Assembly, "__init__", return_value=None): + result = assembly_service.create(mock_db, obj_in=sample_assembly_create) + + # Verify version was incremented to 4 + mock_db.add.assert_called_once() + mock_db.commit.assert_called_once() + + def test_create_version_per_combination(self, mock_db, assembly_service): + """Test that version is per (data_types, organism_key, sample_id) combination.""" + sample_id = uuid.uuid4() + + # Create two assemblies with different data_types + assembly1 = AssemblyCreate( + organism_key="organism1", + sample_id=sample_id, + assembly_name="Assembly 1", + assembly_type="clone or isolate", + data_types=AssemblyDataTypes.PACBIO_SMRT, + coverage=50.0, + program="hifiasm", + moleculetype="genomic DNA", + ) + + assembly2 = AssemblyCreate( + organism_key="organism1", + sample_id=sample_id, + assembly_name="Assembly 2", + assembly_type="clone or isolate", + data_types=AssemblyDataTypes.PACBIO_SMRT_HIC, + coverage=50.0, + program="hifiasm", + moleculetype="genomic DNA", + ) + + # Mock query to return None for both (different combinations) + mock_query = Mock() + mock_query.filter.return_value.scalar.return_value = None + mock_db.query.return_value = mock_query + mock_db.add = Mock() + mock_db.commit = Mock() + mock_db.refresh = Mock() + + with patch.object(Assembly, "__init__", return_value=None): + assembly_service.create(mock_db, obj_in=assembly1) + assembly_service.create(mock_db, obj_in=assembly2) + + # Both should get version 1 since they have different data_types + assert mock_db.add.call_count == 2 + + +class TestCreateFromExperiments: + """Tests for create_from_experiments method.""" + + def test_create_from_experiments_success(self, mock_db, assembly_service): + """Test successful assembly creation from experiments.""" + tax_id = 172942 + organism = Organism( + grouping_key="test_organism", + tax_id=tax_id, + scientific_name="Test Species", + ) + sample = Sample(id=uuid.uuid4(), organism_key="test_organism") + experiments = [ + Experiment( + id=uuid.uuid4(), + sample_id=sample.id, + platform="PACBIO_SMRT", + library_strategy="WGS", + ), + ] + + # Mock database queries + mock_db.query.return_value.filter.return_value.first.return_value = organism + mock_db.query.return_value.filter.return_value.all.side_effect = [ + [sample], # samples query + experiments, # experiments query + ] + + assembly_in = AssemblyCreateFromExperiments( + sample_id=sample.id, + assembly_name="Test Assembly", + assembly_type="clone or isolate", + coverage=50.0, + program="hifiasm", + moleculetype="genomic DNA", + ) + + with patch.object(assembly_service, "create") as mock_create: + mock_create.return_value = Mock( + id=uuid.uuid4(), + data_types=AssemblyDataTypes.PACBIO_SMRT, + version=1, + ) + + assembly, platform_info = assembly_service.create_from_experiments( + mock_db, tax_id=tax_id, assembly_in=assembly_in + ) + + # Verify platform info was returned + assert "platforms" in platform_info + assert "library_strategies" in platform_info + assert "experiment_count" in platform_info + + def test_create_from_experiments_organism_not_found(self, mock_db, assembly_service): + """Test error when organism not found.""" + tax_id = 999999 + mock_db.query.return_value.filter.return_value.first.return_value = None + + assembly_in = AssemblyCreateFromExperiments( + sample_id=uuid.uuid4(), + assembly_name="Test Assembly", + assembly_type="clone or isolate", + coverage=50.0, + program="hifiasm", + moleculetype="genomic DNA", + ) + + with pytest.raises(ValueError, match="Organism with tax_id 999999 not found"): + assembly_service.create_from_experiments( + mock_db, tax_id=tax_id, assembly_in=assembly_in + ) + + def test_create_from_experiments_no_samples(self, mock_db, assembly_service): + """Test error when no samples found for organism.""" + tax_id = 172942 + organism = Organism( + grouping_key="test_organism", + tax_id=tax_id, + scientific_name="Test Species", + ) + + # Mock organism found but no samples + mock_db.query.return_value.filter.return_value.first.return_value = organism + mock_db.query.return_value.filter.return_value.all.return_value = [] + + assembly_in = AssemblyCreateFromExperiments( + sample_id=uuid.uuid4(), + assembly_name="Test Assembly", + assembly_type="clone or isolate", + coverage=50.0, + program="hifiasm", + moleculetype="genomic DNA", + ) + + with pytest.raises(ValueError, match="No samples found"): + assembly_service.create_from_experiments( + mock_db, tax_id=tax_id, assembly_in=assembly_in + ) + + def test_create_from_experiments_no_experiments(self, mock_db, assembly_service): + """Test error when no experiments found.""" + tax_id = 172942 + organism = Organism( + grouping_key="test_organism", + tax_id=tax_id, + scientific_name="Test Species", + ) + sample = Sample(id=uuid.uuid4(), organism_key="test_organism") + + # Mock organism and samples found but no experiments + mock_db.query.return_value.filter.return_value.first.return_value = organism + mock_db.query.return_value.filter.return_value.all.side_effect = [ + [sample], # samples query + [], # experiments query - empty + ] + + assembly_in = AssemblyCreateFromExperiments( + sample_id=sample.id, + assembly_name="Test Assembly", + assembly_type="clone or isolate", + coverage=50.0, + program="hifiasm", + moleculetype="genomic DNA", + ) + + with pytest.raises(ValueError, match="No experiments found"): + assembly_service.create_from_experiments( + mock_db, tax_id=tax_id, assembly_in=assembly_in + ) + + def test_create_from_experiments_overrides_data_types(self, mock_db, assembly_service): + """Test that data_types is overridden based on experiments.""" + tax_id = 172942 + organism = Organism( + grouping_key="test_organism", + tax_id=tax_id, + scientific_name="Test Species", + ) + sample = Sample(id=uuid.uuid4(), organism_key="test_organism") + experiments = [ + Experiment( + id=uuid.uuid4(), + sample_id=sample.id, + platform="OXFORD_NANOPORE", + library_strategy="WGS", + ), + Experiment( + id=uuid.uuid4(), + sample_id=sample.id, + platform="ILLUMINA", + library_strategy="Hi-C", + ), + ] + + mock_db.query.return_value.filter.return_value.first.return_value = organism + mock_db.query.return_value.filter.return_value.all.side_effect = [ + [sample], + experiments, + ] + + # User provides PACBIO_SMRT which should be used instead of auto-detection + assembly_in = AssemblyCreateFromExperiments( + sample_id=sample.id, + assembly_name="Test Assembly", + assembly_type="clone or isolate", + data_types=AssemblyDataTypes.PACBIO_SMRT, # Explicitly provided, should be used + coverage=50.0, + program="hifiasm", + moleculetype="genomic DNA", + ) + + with patch.object(assembly_service, "create") as mock_create: + # Verify that create is called with overridden data_types + mock_create.return_value = Mock( + id=uuid.uuid4(), + data_types=AssemblyDataTypes.OXFORD_NANOPORE_HIC, + version=1, + ) + + assembly, platform_info = assembly_service.create_from_experiments( + mock_db, tax_id=tax_id, assembly_in=assembly_in + ) + + # Verify create was called + mock_create.assert_called_once() + call_args = mock_create.call_args + created_assembly_in = call_args.kwargs["obj_in"] + + # The data_types should be determined from experiments (ILLUMINA+WGS = Hi-C) + assert created_assembly_in.organism_key == "test_organism" diff --git a/tests/unit/services/test_genome_note_service.py b/tests/unit/services/test_genome_note_service.py index d4f2a4e..afd48bc 100644 --- a/tests/unit/services/test_genome_note_service.py +++ b/tests/unit/services/test_genome_note_service.py @@ -6,7 +6,6 @@ from sqlalchemy.orm import Session from app.models.genome_note import GenomeNote -from app.schemas.genome_note import GenomeNoteCreate, GenomeNoteUpdate from app.services.genome_note_service import GenomeNoteService @@ -19,7 +18,7 @@ def mock_db(): @pytest.fixture def genome_note_service(): """Create a GenomeNoteService instance.""" - return GenomeNoteService() + return GenomeNoteService(GenomeNote) @pytest.fixture @@ -73,69 +72,6 @@ def test_get_next_version_different_organisms(self, mock_db, genome_note_service assert version == 3 -class TestCreateGenomeNote: - """Tests for create method.""" - - def test_create_genome_note_success(self, mock_db, genome_note_service): - """Test successful genome note creation.""" - assembly_id = uuid.uuid4() - genome_note_in = GenomeNoteCreate( - organism_key="test_organism", - assembly_id=assembly_id, - title="Test Note", - note_url="https://example.com/note", - ) - - mock_query = Mock() - mock_query.filter.return_value.scalar.return_value = None - mock_db.query.return_value = mock_query - - created_note = genome_note_service.create(mock_db, genome_note_in) - - mock_db.add.assert_called_once() - mock_db.commit.assert_called_once() - mock_db.refresh.assert_called_once() - - def test_create_genome_note_auto_version(self, mock_db, genome_note_service): - """Test that version is automatically calculated.""" - assembly_id = uuid.uuid4() - genome_note_in = GenomeNoteCreate( - organism_key="test_organism", - assembly_id=assembly_id, - title="Test Note", - note_url="https://example.com/note", - ) - - mock_query = Mock() - mock_query.filter.return_value.scalar.return_value = 2 - mock_db.query.return_value = mock_query - - genome_note_service.create(mock_db, genome_note_in) - - call_args = mock_db.add.call_args[0][0] - assert call_args.version == 3 - - def test_create_genome_note_defaults(self, mock_db, genome_note_service): - """Test that default values are set correctly.""" - assembly_id = uuid.uuid4() - genome_note_in = GenomeNoteCreate( - organism_key="test_organism", - assembly_id=assembly_id, - title="Test Note", - note_url="https://example.com/note", - ) - - mock_query = Mock() - mock_query.filter.return_value.scalar.return_value = None - mock_db.query.return_value = mock_query - - genome_note_service.create(mock_db, genome_note_in) - - call_args = mock_db.add.call_args[0][0] - assert call_args.is_published is False - assert call_args.published_at is None - - class TestPublishGenomeNote: """Tests for publish_genome_note method.""" @@ -347,84 +283,60 @@ def test_get_versions_by_organism_empty(self, mock_db, genome_note_service): assert result == [] -class TestUpdateGenomeNote: - """Tests for update method.""" - - def test_update_genome_note_success(self, mock_db, genome_note_service, sample_genome_note): - """Test successful update of a genome note.""" - note_id = sample_genome_note.id - update_data = GenomeNoteUpdate( - title="Updated Title", - note_url="https://example.com/updated", - ) +class TestGetByFilters: + """Tests for basic filter helpers.""" + def test_get_by_organism_key(self, mock_db, genome_note_service, sample_genome_note): mock_query = Mock() - mock_query.filter.return_value.first.return_value = sample_genome_note + mock_query.filter.return_value.all.return_value = [sample_genome_note] mock_db.query.return_value = mock_query - result = genome_note_service.update(mock_db, note_id, update_data) - - assert result.title == "Updated Title" - assert result.note_url == "https://example.com/updated" - mock_db.commit.assert_called_once() + result = genome_note_service.get_by_organism_key(mock_db, "test_organism") - def test_update_genome_note_partial(self, mock_db, genome_note_service, sample_genome_note): - """Test partial update of a genome note.""" - note_id = sample_genome_note.id - original_url = sample_genome_note.note_url - update_data = GenomeNoteUpdate(title="Updated Title Only") + assert result == [sample_genome_note] + def test_get_by_assembly_id(self, mock_db, genome_note_service, sample_genome_note): mock_query = Mock() - mock_query.filter.return_value.first.return_value = sample_genome_note + mock_query.filter.return_value.all.return_value = [sample_genome_note] mock_db.query.return_value = mock_query - result = genome_note_service.update(mock_db, note_id, update_data) + result = genome_note_service.get_by_assembly_id(mock_db, sample_genome_note.assembly_id) - assert result.title == "Updated Title Only" - assert result.note_url == original_url - - def test_update_genome_note_not_found(self, mock_db, genome_note_service): - """Test updating a non-existent genome note.""" - note_id = uuid.uuid4() - update_data = GenomeNoteUpdate(title="Updated Title") + assert result == [sample_genome_note] + def test_get_by_title(self, mock_db, genome_note_service, sample_genome_note): mock_query = Mock() - mock_query.filter.return_value.first.return_value = None + mock_query.filter.return_value.all.return_value = [sample_genome_note] mock_db.query.return_value = mock_query - result = genome_note_service.update(mock_db, note_id, update_data) + result = genome_note_service.get_by_title(mock_db, "Test") - assert result is None - mock_db.commit.assert_not_called() + assert result == [sample_genome_note] + def test_get_multi_with_filters(self, mock_db, genome_note_service, sample_genome_note): + class FakeQuery: + def filter(self, *_args, **_kwargs): + return self -class TestDeleteGenomeNote: - """Tests for delete method.""" + def offset(self, *_args, **_kwargs): + return self - def test_delete_genome_note_success(self, mock_db, genome_note_service, sample_genome_note): - """Test successful deletion of a genome note.""" - note_id = sample_genome_note.id + def limit(self, *_args, **_kwargs): + return self - mock_query = Mock() - mock_query.filter.return_value.first.return_value = sample_genome_note - mock_db.query.return_value = mock_query - - result = genome_note_service.delete(mock_db, note_id) + def all(self): + return [sample_genome_note] - assert result == sample_genome_note - mock_db.delete.assert_called_once_with(sample_genome_note) - mock_db.commit.assert_called_once() - - def test_delete_genome_note_not_found(self, mock_db, genome_note_service): - """Test deleting a non-existent genome note.""" - note_id = uuid.uuid4() + mock_db.query.return_value = FakeQuery() - mock_query = Mock() - mock_query.filter.return_value.first.return_value = None - mock_db.query.return_value = mock_query - - result = genome_note_service.delete(mock_db, note_id) + result = genome_note_service.get_multi_with_filters( + mock_db, + skip=0, + limit=10, + organism_key="test_organism", + assembly_id=sample_genome_note.assembly_id, + is_published=False, + title="Test", + ) - assert result is None - mock_db.delete.assert_not_called() - mock_db.commit.assert_not_called() + assert result == [sample_genome_note] diff --git a/tests/unit/test_settings.py b/tests/unit/test_settings.py index bfd2b5a..e8c65b1 100644 --- a/tests/unit/test_settings.py +++ b/tests/unit/test_settings.py @@ -8,6 +8,8 @@ def test_settings_builds_database_uri_from_env(monkeypatch): monkeypatch.setenv("POSTGRES_PORT", "5432") monkeypatch.setenv("POSTGRES_DB", "testdb") monkeypatch.setenv("DATABASE_URI", "postgresql://testuser:testpass@localhost:5432/testdb") + monkeypatch.setenv("JWT_SECRET_KEY", "test-secret") + monkeypatch.setenv("JWT_ALGORITHM", "HS256") settings = Settings() diff --git a/tests/unit/test_user_password_validation.py b/tests/unit/test_user_password_validation.py index 1144e8c..b397583 100644 --- a/tests/unit/test_user_password_validation.py +++ b/tests/unit/test_user_password_validation.py @@ -2,11 +2,17 @@ from fastapi.testclient import TestClient +from app.api.v1.endpoints import users from app.main import app +def _override_admin_user(): + return {"is_superuser": False, "roles": ["admin"], "is_active": True} + + def test_create_user_rejects_password_over_72_bytes(): client = TestClient(app) + app.dependency_overrides[users.get_current_active_user] = _override_admin_user payload = { "username": "u1", @@ -21,7 +27,7 @@ def test_create_user_rejects_password_over_72_bytes(): resp = client.post("/api/v1/users/", json=payload) assert resp.status_code == 422 - detail = resp.json().get("detail") + detail = resp.json().get("error", {}).get("details", {}).get("errors") assert isinstance(detail, list) # Ensure we fail due to our validator message assert any("at most 72 bytes" in (err.get("msg") or "") for err in detail) @@ -29,6 +35,7 @@ def test_create_user_rejects_password_over_72_bytes(): def test_update_user_rejects_password_over_72_bytes(): client = TestClient(app) + app.dependency_overrides[users.get_current_active_user] = _override_admin_user user_id = str(uuid4()) payload = { @@ -39,6 +46,6 @@ def test_update_user_rejects_password_over_72_bytes(): resp = client.put(f"/api/v1/users/{user_id}", json=payload) assert resp.status_code == 422 - detail = resp.json().get("detail") + detail = resp.json().get("error", {}).get("details", {}).get("errors") assert isinstance(detail, list) assert any("at most 72 bytes" in (err.get("msg") or "") for err in detail)