Coverage for rest_api/viewsets_sample.py: 50%

438 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-10 03:22 +0000

1import _csv 

2import ast 

3import csv 

4from dataclasses import dataclass 

5from datetime import datetime 

6import json 

7import os 

8import re 

9import time 

10import traceback 

11from typing import Generator 

12 

13from django.core.files.uploadedfile import InMemoryUploadedFile 

14from django.core.paginator import Paginator 

15from django.db.models import Exists 

16from django.db.models import F 

17from django.db.models import OuterRef 

18from django.db.models import Prefetch 

19from django.db.models import Q 

20from django.db.models import QuerySet 

21from django.db.models import Subquery 

22from django.http import StreamingHttpResponse 

23from django_filters.rest_framework import DjangoFilterBackend 

24from rest_framework import generics 

25from rest_framework import status 

26from rest_framework import viewsets 

27from rest_framework.decorators import action 

28from rest_framework.filters import OrderingFilter 

29from rest_framework.request import Request 

30from rest_framework.response import Response 

31 

32from rest_api.data_entry.sample_job import delete_samples 

33from rest_api.data_entry.sample_job import delete_sequences 

34from rest_api.serializers import SampleGenomesExportStreamSerializer 

35from rest_api.utils import define_profile 

36from rest_api.utils import get_distinct_cds_accessions 

37from rest_api.utils import get_distinct_gene_symbols 

38from rest_api.utils import get_distinct_replicon_accessions 

39from rest_api.utils import resolve_ambiguous_NT_AA 

40from rest_api.utils import strtobool 

41from rest_api.viewsets import PropertyViewSet 

42from sonar_backend.settings import DEBUG 

43from sonar_backend.settings import LOGGER 

44from sonar_backend.settings import SONAR_DATA_ENTRY_FOLDER 

45from . import models 

46from .serializers import SampleGenomesSerializer 

47from .serializers import SampleGenomesSerializerVCF 

48from .serializers import SampleSerializer 

49 

50 

51@dataclass 

52class LineageInfo: 

53 name: str 

54 parent: str 

55 

56 def __hash__(self): 

57 return hash((self.name, self.parent)) 

58 

59 

60class Echo: 

61 def write(self, value): 

62 return value 

63 

64 

65class SampleFilterMixin: 

66 # Cache for gene symbols, replicons and CDS accessions 

67 _cached_gene_symbols = None 

68 _cached_replicons = None 

69 _cache_timestamp = None 

70 _cached_cds_accs = None 

71 CACHE_TTL = 3600 * 24 # 1 day 

72 

73 @property 

74 def filter_label_to_methods(self): 

75 return { 

76 "Property": self.filter_property, 

77 "SNP Nt": self.filter_snp_profile_nt, 

78 "SNP AA": self.filter_snp_profile_aa, 

79 "Del Nt": self.filter_del_profile_nt, 

80 "Del AA": self.filter_del_profile_aa, 

81 "Ins Nt": self.filter_ins_profile_nt, 

82 "Ins AA": self.filter_ins_profile_aa, 

83 "Replicon": self.filter_replicon, 

84 "Reference": self.filter_reference, 

85 "Sample": self.filter_sample, 

86 "Lineages": self.filter_sublineages, 

87 "Annotation": self.filter_annotation, 

88 "DNA/AA Profile": self.filter_label, 

89 } 

90 

91 @classmethod 

92 def get_gene_symbols(cls, force_refresh=False): 

93 """ 

94 update cached gene symbol set if cache older than CACHE_TTL 

95 """ 

96 now = time.time() 

97 if ( 

98 force_refresh 

99 or cls._cached_gene_symbols is None 

100 or cls._cache_timestamp is None 

101 or (now - cls._cache_timestamp) > cls.CACHE_TTL 

102 ): 

103 

104 symbols = get_distinct_gene_symbols() 

105 cls._cached_gene_symbols = set(symbols) 

106 

107 return cls._cached_gene_symbols 

108 

109 @classmethod 

110 def get_replicons(cls, reference_accession=None, force_refresh=True): 

111 """ 

112 update cached replicon accession set if cache older than CACHE_TTL 

113 """ 

114 now = time.time() 

115 if ( 

116 force_refresh 

117 or cls._cached_replicons is None 

118 or cls._cache_timestamp is None 

119 or (now - cls._cache_timestamp) > cls.CACHE_TTL 

120 ): 

121 

122 replicons_list = get_distinct_replicon_accessions(reference_accession) 

123 cls._cached_replicons = set(replicons_list) 

124 cls._cache_timestamp = now 

125 

126 return cls._cached_replicons 

127 

128 @classmethod 

129 def get_cds_accessions(cls, replicon=None, force_refresh=True): 

130 now = time.time() 

131 if ( 

132 force_refresh 

133 or cls._cached_replicons is None 

134 or cls._cache_timestamp is None 

135 or (now - cls._cache_timestamp) > cls.CACHE_TTL 

136 ): 

137 cds_acc = get_distinct_cds_accessions(replicon=replicon) 

138 cls._cached_cds_accs = set(cds_acc) 

139 return cls._cached_cds_accs 

140 

141 def _get_reference_replicon_count(self, reference_accession): 

142 """ 

143 Returns the number of replicons for a given reference. 

144 Returns None if reference not found. 

145 """ 

146 return models.Replicon.objects.filter( 

147 reference__accession=reference_accession 

148 ).count() 

149 

150 def _resolve_replicon_for_query(self, parsed_mutation, reference_accession): 

151 """ 

152 Determines which replicon to use for mutation query. 

153 

154 Returns: 

155 - replicon_accession (str) if successful 

156 - Raises ValueError if ambiguous (multiple replicons without explicit accession) 

157 """ 

158 # case 1: replicon accession in parsed mutation 

159 if ( 

160 "replicon_accession" in parsed_mutation 

161 and parsed_mutation["replicon_accession"] 

162 ): 

163 return parsed_mutation["replicon_accession"] 

164 

165 # case 2: no replicon accession in parsed mutation 

166 replicon_count = self._get_reference_replicon_count(reference_accession) 

167 LOGGER.debug( 

168 f"replicon_count in resolve replicon for reference {reference_accession}: {replicon_count}" 

169 ) 

170 

171 if replicon_count is None or replicon_count == 0: 

172 raise ValueError(f"No replicons found for reference {reference_accession}.") 

173 

174 if replicon_count == 1: 

175 # one replicon: use this replicon 

176 replicon = models.Replicon.objects.get( 

177 reference__accession=reference_accession 

178 ) 

179 return replicon.accession 

180 

181 # case 3: multiple replicons per reference, parsed mutation without replicon accession 

182 if replicon_count > 1: 

183 raise ValueError( 

184 f"Reference {reference_accession} has {replicon_count} replicons. " 

185 f"Please specify replicon accession (e.g., NC_026438.1:A425G)." 

186 ) 

187 

188 def resolve_recursive_genome_filter( 

189 self, filters, reference_accession=None, depth=0 

190 ) -> Q: 

191 indent = " " * depth 

192 LOGGER.debug(f"{indent}resolve_genome_filter called, depth={depth}") 

193 LOGGER.debug(f"{indent}filters: {filters}") 

194 LOGGER.debug(f"{indent}reference_accession: {reference_accession}") 

195 q_obj = Q() 

196 # Process single filter at current level 

197 if "label" in filters: 

198 q_obj &= self.eval_basic_filter(filters, reference_accession) 

199 

200 # Process AND filters 

201 for f in filters.get("andFilter", []): 

202 if "orFilter" in f or "andFilter" in f: 

203 # Recursive call for nested filters 

204 q_obj &= self.resolve_recursive_genome_filter( 

205 f, reference_accession, depth + 1 

206 ) 

207 else: 

208 q_obj &= self.eval_basic_filter(f, reference_accession) 

209 

210 # Process OR filters 

211 for or_filter in filters.get("orFilter", []): 

212 q_obj |= self.resolve_recursive_genome_filter( 

213 or_filter, reference_accession, depth + 1 

214 ) 

215 return q_obj 

216 

217 def eval_basic_filter(self, filter_dict, reference_accession) -> Q: 

218 """ 

219 Evaluate a single basic filter. 

220 """ 

221 label = filter_dict.get("label") 

222 method = self.filter_label_to_methods.get(label) 

223 if not method: 

224 raise Exception(f"Filter method not found for: {label}") 

225 # Pass reference_accession to filter methods 

226 filter_kwargs = {**filter_dict, "reference_accession": reference_accession} 

227 

228 return method(**filter_kwargs) 

229 

230 def get_filtered_queryset(self, request: Request): 

231 """ 

232 retrieve filtered queryset Sample based on request parameters 

233 """ 

234 if not (filter_params := request.query_params.get("filters")): 

235 queryset = models.Sample.objects.all() 

236 else: 

237 filters = json.loads(filter_params) 

238 if "reference" in request.query_params: 

239 reference_accession = ( 

240 request.query_params.get("reference").strip('"').strip() 

241 ) 

242 else: 

243 reference_accession = filters.get("reference") 

244 if not reference_accession: 

245 raise ValueError("Reference accession is required in filters.") 

246 

247 LOGGER.info( 

248 f"Genomes Query, conditions: {filters}, reference: {reference_accession}" 

249 ) 

250 

251 q_filter = self.resolve_recursive_genome_filter( 

252 filters, reference_accession 

253 ) 

254 

255 queryset = models.Sample.objects.filter( 

256 q_filter, 

257 sequences__alignments__replicon__reference__accession=reference_accession, 

258 ).distinct() 

259 

260 queryset = queryset.select_related().prefetch_related( 

261 "sequences", 

262 "properties__property", 

263 ) 

264 

265 return queryset 

266 

267 def filter_label( 

268 self, 

269 value, 

270 reference_accession=None, 

271 exclude: bool = False, 

272 *args, 

273 **kwargs, 

274 ): 

275 final_query = Q() 

276 # Split the input value by either commas, semicolon, whitespace, or combinations of these, 

277 # remove separators from string end 

278 mutations = re.split(r"[,\s;]+", value.strip(",; \t\r\n")) 

279 for mutation in mutations: 

280 parsed_mutation = define_profile( 

281 mutation, self.get_gene_symbols(), self.get_replicons() 

282 ) 

283 # resolve replicon accession 

284 if parsed_mutation["label"] in ["SNP Nt", "Ins Nt", "Del Nt"]: 

285 replicon_accession = self._resolve_replicon_for_query( 

286 parsed_mutation, reference_accession 

287 ) 

288 parsed_mutation["replicon_accession"] = replicon_accession 

289 # Validate protein name for AA mutations 

290 if ( 

291 "protein_symbol" in parsed_mutation 

292 and parsed_mutation["protein_symbol"] not in self.get_gene_symbols() 

293 ): 

294 raise ValueError( 

295 f"Invalid protein name: {parsed_mutation['protein_symbol']}." 

296 ) 

297 # Check the parsed mutation type and call appropriate filter function 

298 if parsed_mutation.get("label") == "SNP Nt": 

299 q_obj = self.filter_snp_profile_nt( 

300 ref_nuc=parsed_mutation["ref_nuc"], 

301 ref_pos=int(parsed_mutation["ref_pos"]), 

302 alt_nuc=parsed_mutation["alt_nuc"], 

303 replicon_accession=parsed_mutation["replicon_accession"], 

304 ) 

305 

306 elif parsed_mutation.get("label") == "SNP AA": 

307 q_obj = self.filter_snp_profile_aa( 

308 protein_symbol=parsed_mutation["protein_symbol"], 

309 ref_aa=parsed_mutation["ref_aa"], 

310 ref_pos=int(parsed_mutation["ref_pos"]), 

311 alt_aa=parsed_mutation["alt_aa"], 

312 ) 

313 

314 elif parsed_mutation.get("label") == "Del Nt": 

315 q_obj = self.filter_del_profile_nt( 

316 first_deleted=parsed_mutation["first_deleted"], 

317 last_deleted=parsed_mutation.get("last_deleted", ""), 

318 replicon_accession=parsed_mutation["replicon_accession"], 

319 ) 

320 

321 elif parsed_mutation.get("label") == "Del AA": 

322 q_obj = self.filter_del_profile_aa( 

323 protein_symbol=parsed_mutation["protein_symbol"], 

324 first_deleted=parsed_mutation["first_deleted"], 

325 last_deleted=parsed_mutation.get("last_deleted", ""), 

326 ) 

327 

328 elif parsed_mutation.get("label") == "Ins Nt": 

329 q_obj = self.filter_ins_profile_nt( 

330 ref_nuc=parsed_mutation["ref_nuc"], 

331 ref_pos=int(parsed_mutation["ref_pos"]), 

332 alt_nuc=parsed_mutation["alt_nuc"], 

333 replicon_accession=parsed_mutation["replicon_accession"], 

334 ) 

335 

336 elif parsed_mutation.get("label") == "Ins AA": 

337 q_obj = self.filter_ins_profile_aa( 

338 protein_symbol=parsed_mutation["protein_symbol"], 

339 ref_aa=parsed_mutation["ref_aa"], 

340 ref_pos=int(parsed_mutation["ref_pos"]), 

341 alt_aa=parsed_mutation["alt_aa"], 

342 ) 

343 

344 else: 

345 

346 raise ValueError( 

347 f"Unsupported mutation type: {parsed_mutation.get('label')}" 

348 ) 

349 # Combine queries with AND operator (&) for each mutation 

350 final_query &= q_obj 

351 

352 if exclude: 

353 final_query = ~final_query 

354 

355 return final_query 

356 

357 def filter_annotation( 

358 self, 

359 property_name, 

360 filter_type, 

361 value, 

362 exclude: bool = False, 

363 *args, 

364 **kwargs, 

365 ) -> Q: 

366 query = {} 

367 query[f"nucleotide_mutations__annotations__{property_name}__{filter_type}"] = ( 

368 value 

369 ) 

370 

371 alignment_qs = models.Alignment.objects.filter(**query) 

372 filters = {"sequences__alignments__in": alignment_qs} 

373 

374 if exclude: 

375 return ~Q(**filters) 

376 return Q(**filters) 

377 

378 def filter_property( 

379 self, 

380 property_name, 

381 filter_type, 

382 value, 

383 exclude: bool = False, 

384 *args, 

385 **kwargs, 

386 ) -> Q: 

387 # Convert str of list into list object 

388 # "['X','N']" -> ['X','N'] 

389 if isinstance(value, str): 

390 try: 

391 # Safely evaluate the string representation of the list and convert it to a list object 

392 value = ast.literal_eval(value) 

393 except (SyntaxError, ValueError): 

394 # Handle the case where the string couldn't be evaluated as a list 

395 pass 

396 

397 # check the filter_type 

398 if filter_type == "contains": 

399 value = value.strip("%") 

400 elif filter_type == "range": 

401 if isinstance(value, str): 

402 value = value.split(",") 

403 

404 self.has_property_filter = True 

405 

406 # Special handling for 'length' - stored in Sequence table 

407 if property_name == "length": 

408 query = {} 

409 query[f"sequences__{property_name}__{filter_type}"] = value 

410 elif property_name in [ 

411 field.name for field in models.Sample._meta.get_fields() 

412 ]: 

413 query = {} 

414 query[f"{property_name}__{filter_type}"] = value 

415 else: 

416 datatype = models.Property.objects.get(name=property_name).datatype 

417 query = {f"properties__property__name": property_name} 

418 query[f"properties__{datatype}__{filter_type}"] = value 

419 if exclude: 

420 return ~Q(**query) 

421 return Q(**query) 

422 

423 def filter_nt_mutations( 

424 self, 

425 mutation_condition, 

426 exclude: bool = False, 

427 replicon_accession: str = None, 

428 ): 

429 if replicon_accession: 

430 mutation_condition &= Q(replicon__accession=replicon_accession) 

431 

432 # Use EXISTS subquery instead of JOIN 

433 mutation_subquery = models.NucleotideMutation.objects.filter( 

434 mutation_condition, alignments__sequence__samples=OuterRef("pk") 

435 ) 

436 

437 if exclude: 

438 return Q( 

439 pk__in=models.Sample.objects.exclude( 

440 pk__in=models.Sample.objects.filter(Exists(mutation_subquery)) 

441 ) 

442 ) 

443 

444 return Q(Exists(mutation_subquery)) 

445 

446 def filter_aa_mutations( 

447 self, 

448 mutation_condition, 

449 exclude: bool = False, 

450 replicon_accession: str = None, 

451 ): 

452 if replicon_accession: 

453 mutation_condition &= Q(replicon__accession=replicon_accession) 

454 

455 # Use EXISTS subquery instead of JOIN 

456 mutation_subquery = models.AminoAcidMutation.objects.filter( 

457 mutation_condition, alignments__sequence__samples=OuterRef("pk") 

458 ) 

459 

460 if exclude: 

461 return Q( 

462 pk__in=models.Sample.objects.exclude( 

463 pk__in=models.Sample.objects.filter(Exists(mutation_subquery)) 

464 ) 

465 ) 

466 

467 return Q(Exists(mutation_subquery)) 

468 

469 def filter_snp_profile_nt( 

470 self, 

471 # gene_symbol: str, 

472 ref_nuc: str, 

473 ref_pos: int, 

474 alt_nuc: str, 

475 replicon_accession: str = None, 

476 exclude: bool = False, 

477 *args, 

478 **kwargs, 

479 ) -> Q: 

480 # For NT: ref_nuc followed by ref_pos followed by alt_nuc (e.g. T28175C). 

481 if alt_nuc == "N": 

482 mutation_alt = Q() 

483 for x in resolve_ambiguous_NT_AA(type="nt", char=alt_nuc): 

484 mutation_alt |= Q(alt=x) 

485 else: 

486 mutation_alt = Q(alt=alt_nuc if alt_nuc != "n" else "N") 

487 

488 mutation_condition = Q(end=ref_pos) & Q(ref=ref_nuc) & mutation_alt 

489 

490 return self.filter_nt_mutations(mutation_condition, exclude, replicon_accession) 

491 

492 def filter_snp_profile_aa( 

493 self, 

494 protein_symbol: str, 

495 ref_aa: str, 

496 ref_pos: str, 

497 alt_aa: str, 

498 replicon_accession: str = None, 

499 exclude: bool = False, 

500 *args, 

501 **kwargs, 

502 ) -> Q: 

503 # For AA: protein_symbol:ref_aa followed by ref_pos followed by alt_aa (e.g. OPG098:E162K) 

504 if protein_symbol not in self.get_gene_symbols(): 

505 raise ValueError(f"Invalid protein name: {protein_symbol}.") 

506 if alt_aa == "X": 

507 mutation_alt = Q() 

508 for x in resolve_ambiguous_NT_AA(type="aa", char=alt_aa): 

509 mutation_alt |= Q(alt=x) 

510 else: 

511 mutation_alt = Q(alt=alt_aa if alt_aa != "x" else "X") 

512 

513 mutation_condition = ( 

514 Q(end=ref_pos) 

515 & Q(ref=ref_aa) 

516 & (mutation_alt) 

517 & Q(cds__gene__symbol=protein_symbol) 

518 ) 

519 return self.filter_aa_mutations(mutation_condition, exclude, replicon_accession) 

520 

521 def filter_del_profile_nt( 

522 self, 

523 first_deleted: str, 

524 last_deleted: str, 

525 replicon_accession: str, 

526 exclude: bool = False, 

527 *args, 

528 **kwargs, 

529 ) -> Q: 

530 # For NT: del:first_NT_deleted-last_NT_deleted (e.g. del:133177-133186). 

531 # in case only single deletion bp 

532 if last_deleted == "": 

533 last_deleted = first_deleted 

534 

535 mutation_condition = ( 

536 Q(start=int(first_deleted) - 1) & Q(end=int(last_deleted)) & Q(alt="") 

537 ) 

538 return self.filter_nt_mutations(mutation_condition, exclude, replicon_accession) 

539 

540 def filter_del_profile_aa( 

541 self, 

542 protein_symbol: str, 

543 first_deleted: str, 

544 last_deleted: str, 

545 replicon_accession: str = None, 

546 exclude: bool = False, 

547 *args, 

548 **kwargs, 

549 ) -> Q: 

550 # For AA: protein_symbol:del:first_AA_deleted-last_AA_deleted (e.g. OPG197:del:34-35) 

551 if protein_symbol not in self.get_gene_symbols(): 

552 raise ValueError(f"Invalid protein name: {protein_symbol}.") 

553 # in case only single deletion bp 

554 if last_deleted == "": 

555 last_deleted = first_deleted 

556 

557 mutation_condition = ( 

558 Q(cds__gene__symbol__iexact=protein_symbol) 

559 & Q(start=int(first_deleted) - 1) 

560 & Q(end=int(last_deleted)) 

561 & Q(alt="") 

562 ) 

563 

564 return self.filter_aa_mutations(mutation_condition, exclude, replicon_accession) 

565 

566 def filter_ins_profile_nt( 

567 self, 

568 ref_nuc: str, 

569 ref_pos: int, 

570 alt_nuc: str, 

571 replicon_accession: str, 

572 exclude: bool = False, 

573 *args, 

574 **kwargs, 

575 ) -> Q: 

576 # For NT: ref_nuc followed by ref_pos followed by alt_nucs (e.g. T133102TTT) 

577 mutation_condition = Q(end=ref_pos) & Q(ref=ref_nuc) & Q(alt=alt_nuc) 

578 return self.filter_nt_mutations(mutation_condition, exclude, replicon_accession) 

579 

580 def filter_ins_profile_aa( 

581 self, 

582 protein_symbol: str, 

583 ref_aa: str, 

584 ref_pos: int, 

585 alt_aa: str, 

586 replicon_accession: str = None, 

587 exclude: bool = False, 

588 *args, 

589 **kwargs, 

590 ) -> Q: 

591 # For AA: protein_symbol:ref_aa followed by ref_pos followed by alt_aas (e.g. OPG197:A34AK) 

592 if protein_symbol not in self.get_gene_symbols(): 

593 raise ValueError(f"Invalid protein name: {protein_symbol}.") 

594 mutation_condition = ( 

595 Q(end=ref_pos) 

596 & Q(ref=ref_aa) 

597 & Q(alt=alt_aa) 

598 & Q(cds__gene__symbol=protein_symbol) 

599 ) 

600 

601 return self.filter_aa_mutations(mutation_condition, exclude, replicon_accession) 

602 

603 def filter_sample( 

604 self, 

605 value: str, 

606 exclude: bool = False, 

607 *args, 

608 **kwargs, 

609 ): 

610 if isinstance(value, str): 

611 sample_list = ast.literal_eval(value) 

612 else: 

613 sample_list = value 

614 

615 if exclude: 

616 return ~Q(name__in=sample_list) 

617 else: 

618 return Q(name__in=sample_list) 

619 

620 def filter_replicon( 

621 self, 

622 replicon_accession, 

623 exclude: bool = False, 

624 *args, 

625 **kwargs, 

626 ): 

627 if exclude: 

628 return ~Q(sequences__alignments__replicon__accession=replicon_accession) 

629 else: 

630 return Q(sequences__alignments__replicon__accession=replicon_accession) 

631 

632 def filter_reference( 

633 self, 

634 accession, 

635 exclude: bool = False, 

636 *args, 

637 **kwargs, 

638 ): 

639 if exclude: 

640 return ~Q(sequences__alignments__replicon__reference__accession=accession) 

641 else: 

642 return Q(sequences__alignments__replicon__reference__accession=accession) 

643 

644 def filter_sublineages( 

645 self, 

646 lineageList, 

647 exclude: bool = False, 

648 includeSublineages: bool = True, 

649 reference_accession: str = None, 

650 *args, 

651 **kwargs, 

652 ): 

653 if isinstance(lineageList, str): 

654 lineageList = [lineageList] # convert to list if a single string is passed 

655 

656 # Match names case-insensitively ("b.1.1.7" -> "B.1.1.7") and scope to the 

657 # queried reference: the same lineage name can exist for different 

658 # pathogens (e.g. "A"/"B" in SARS-CoV-2, mpox, RSV and influenza). 

659 name_q = Q() 

660 for name in lineageList: 

661 name_q |= Q(name__iexact=name) 

662 lineages = models.Lineage.objects.filter(name_q) 

663 if reference_accession: 

664 lineages = lineages.filter(reference__accession=reference_accession) 

665 

666 if not lineages.exists(): 

667 raise Exception(f"Lineage {list(lineageList)} not found.") 

668 if includeSublineages: 

669 sublineages = [] 

670 for l in lineages: 

671 sublineages.extend(l.get_sublineages()) 

672 else: 

673 sublineages = list(lineages) 

674 

675 # match for all sublineages of all given lineages 

676 return self.filter_property( 

677 "lineage", 

678 "in", 

679 sublineages, 

680 exclude, 

681 ) 

682 

683 

684class SampleViewSet( 

685 SampleFilterMixin, 

686 viewsets.GenericViewSet, 

687 generics.mixins.ListModelMixin, 

688 generics.mixins.RetrieveModelMixin, 

689): 

690 queryset = models.Sample.objects.all().order_by("id") 

691 serializer_class = SampleSerializer 

692 filter_backends = [DjangoFilterBackend, OrderingFilter] 

693 lookup_field = "name" 

694 filter_fields = ["name"] 

695 

696 def _get_genomic_and_proteomic_profiles_queryset( 

697 self, queryset, reference_accession, showNX=False 

698 ): 

699 """ 

700 Optimized prefetching for genomic and proteomic profiles 

701 """ 

702 

703 genomic_profiles_qs = models.NucleotideMutation.objects.only( 

704 "ref", "alt", "start", "end", "is_frameshift", "replicon_id" 

705 ).prefetch_related( 

706 "annotations" 

707 ) # Prefetch annotations directly 

708 if not showNX: 

709 genomic_profiles_qs = genomic_profiles_qs.exclude(alt="N") 

710 genomic_profiles_qs = genomic_profiles_qs.order_by("start") 

711 

712 proteomic_profiles_qs = models.AminoAcidMutation.objects.only( 

713 "ref", "alt", "start", "end", "cds_id" 

714 ).prefetch_related( 

715 "cds__gene", 

716 "cds__gene__replicon", # Prefetch replicon as well 

717 "parent__annotations", # Annotations live on the parent NT mutation 

718 ) 

719 if not showNX: 

720 proteomic_profiles_qs = proteomic_profiles_qs.exclude(alt="X") 

721 proteomic_profiles_qs = proteomic_profiles_qs.order_by("cds", "start") 

722 

723 queryset = queryset.prefetch_related( 

724 "sequences__alignments__replicon__reference", 

725 Prefetch( 

726 "sequences__alignments__nucleotide_mutations", 

727 queryset=genomic_profiles_qs, 

728 to_attr="genomic_profiles", 

729 ), 

730 Prefetch( 

731 "sequences__alignments__amino_acid_mutations", 

732 queryset=proteomic_profiles_qs, 

733 to_attr="proteomic_profiles", 

734 ), 

735 ) 

736 

737 return queryset 

738 

739 # actions 

740 @action(detail=False, methods=["get"]) 

741 def genomes(self, request: Request, *args, **kwargs): 

742 """ 

743 fetch proteomic and genomic profiles based on provided filters and optional parameters 

744 """ 

745 try: 

746 timer = datetime.now() 

747 showNX = strtobool(request.query_params.get("showNX", "False")) 

748 csv_stream = strtobool(request.query_params.get("csv_stream", "False")) 

749 vcf_format = strtobool(request.query_params.get("vcf_format", "False")) 

750 

751 LOGGER.info( 

752 f"Genomes Query, optional parameters: showNX:{showNX} csv_stream:{csv_stream}" 

753 ) 

754 

755 self.has_property_filter = False 

756 

757 queryset = self.get_filtered_queryset(request) 

758 

759 # apply ID ('name') filter if provided 

760 if name_filter := request.query_params.get("name"): 

761 queryset = queryset.filter(name=name_filter) 

762 

763 # Get reference_accession für profil-prefetching 

764 filter_params = request.query_params.get("filters") 

765 if filter_params: 

766 filters = json.loads(filter_params) 

767 reference_accession = filters.get("reference") 

768 else: 

769 reference_accession = ( 

770 request.query_params.get("reference", "").strip('"').strip() 

771 ) 

772 

773 # Optimized prefetching 

774 queryset = self._get_genomic_and_proteomic_profiles_queryset( 

775 queryset, reference_accession, showNX 

776 ) 

777 

778 if DEBUG: 

779 LOGGER.info(f"Query: {queryset.query}") 

780 

781 # apply ordering if specified 

782 ordering = request.query_params.get("ordering") 

783 if ordering: 

784 queryset = self._apply_ordering(queryset, ordering) 

785 else: 

786 queryset = queryset.order_by("-collection_date") 

787 

788 # return csv stream if specified 

789 if csv_stream: 

790 return self._return_csv_stream(queryset, request, showNX) 

791 

792 # return vcf format if specified 

793 if vcf_format: 

794 return self._return_vcf_format(queryset, showNX) 

795 

796 # default response - paginate after prefetching 

797 queryset = self.paginate_queryset(queryset) 

798 LOGGER.info( 

799 f"Query time done in {datetime.now() - timer}, Start to Format result" 

800 ) 

801 

802 serializer = SampleGenomesSerializer( 

803 queryset, many=True, context={"request": request, "showNX": showNX} 

804 ) 

805 timer = datetime.now() 

806 LOGGER.info( 

807 f"Serializer done in {datetime.now() - timer}, Start to Format result" 

808 ) 

809 return self.get_paginated_response(serializer.data) 

810 

811 except ValueError as e: 

812 return Response(data={"detail": str(e)}, status=status.HTTP_400_BAD_REQUEST) 

813 except Exception as e: 

814 traceback.print_exc() 

815 return Response( 

816 data={"detail": str(e)}, status=status.HTTP_500_INTERNAL_SERVER_ERROR 

817 ) 

818 

819 def _apply_ordering(self, queryset, ordering): 

820 """ 

821 apply given ordering to queryset 

822 """ 

823 property_names = PropertyViewSet.get_custom_property_names() 

824 ordering_col_name = ordering.lstrip("-") 

825 reverse_order = ordering.startswith("-") 

826 

827 if ordering_col_name in property_names: 

828 datatype = models.Property.objects.get(name=ordering_col_name).datatype 

829 queryset = queryset.order_by( 

830 Subquery( 

831 models.Sample2Property.objects.filter( 

832 property__name=ordering_col_name, sample=OuterRef("id") 

833 ).values(datatype) 

834 ) 

835 ) 

836 if reverse_order: 

837 queryset = queryset.reverse() 

838 else: 

839 queryset = queryset.order_by(ordering) 

840 

841 return queryset 

842 

843 def _get_genomic_and_proteomic_profile_columns(self, columns, reference_accession): 

844 """ 

845 Expand genomic_profiles and proteomic_profiles columns into individual columns 

846 based on existing gene symbols and replicon accessions. 

847 """ 

848 expanded_columns = [] 

849 for column in columns: 

850 if column == "genomic_profiles": 

851 replicons = self.get_replicons(reference_accession) 

852 for replicon in sorted(list(replicons)): 

853 expanded_columns.append(f"genomic_profile: {replicon}") 

854 elif column == "proteomic_profiles": 

855 replicons = self.get_replicons(reference_accession) 

856 for replicon in sorted(list(replicons)): 

857 for cds_acc in sorted(list(self.get_cds_accessions(replicon))): 

858 expanded_columns.append( 

859 f"proteomic_profile: {replicon}: {cds_acc}" 

860 ) 

861 else: 

862 expanded_columns.append(column) 

863 return expanded_columns 

864 

865 def _return_csv_stream(self, queryset, request, showNX=False): 

866 """ 

867 stream queryset data as a csv file 

868 """ 

869 pseudo_buffer = Echo() 

870 writer = csv.writer(pseudo_buffer, delimiter=";") 

871 columns = request.query_params.get("columns") 

872 reference_accession = request.query_params.get("reference") 

873 reference_accession = ( 

874 reference_accession.strip('"').strip("'") if reference_accession else None 

875 ) 

876 

877 if not columns: 

878 raise Exception("No columns provided") 

879 

880 columns = columns.split(",") 

881 columns = self._get_genomic_and_proteomic_profile_columns( 

882 columns, reference_accession 

883 ) 

884 filename = request.query_params.get("filename", "sample_genomes.csv") 

885 

886 return StreamingHttpResponse( 

887 self._stream_serialized_data(queryset, columns, writer, showNX), 

888 content_type="text/csv", 

889 headers={"Content-Disposition": f'attachment; filename="{filename}"'}, 

890 ) 

891 

892 def _return_vcf_format(self, queryset, showNX=False): 

893 """ 

894 return queryset data in vcf format 

895 """ 

896 queryset = self.paginate_queryset(queryset) 

897 serializer = SampleGenomesSerializerVCF( 

898 queryset, many=True, context={"request": self.request, "showNX": showNX} 

899 ) 

900 return self.get_paginated_response(serializer.data) 

901 

902 def _stream_serialized_data( 

903 self, 

904 queryset: QuerySet, 

905 columns: list[str], 

906 writer: "_csv._writer", 

907 showNX: bool = False, 

908 ) -> Generator: 

909 serializer = SampleGenomesExportStreamSerializer 

910 

911 serializer.columns = columns 

912 yield writer.writerow(columns) 

913 paginator = Paginator(queryset, 100) 

914 for page in paginator.page_range: 

915 for serialized in serializer( 

916 paginator.page(page).object_list, 

917 many=True, 

918 context={"showNX": showNX, "columns": columns}, 

919 ).data: 

920 yield writer.writerow(serialized["row"]) 

921 

922 def _temp_save_file(self, uploaded_file: InMemoryUploadedFile): 

923 file_path = os.path.join(SONAR_DATA_ENTRY_FOLDER, uploaded_file.name) 

924 with open(file_path, "wb") as f: 

925 f.write(uploaded_file.read()) 

926 return file_path 

927 

928 @action(detail=False, methods=["get"]) 

929 def get_sequence_data(self, request: Request, *args, **kwargs): 

930 sequence_name = request.GET.get("sequence_name", "") 

931 if not sequence_name: 

932 return Response( 

933 {"detail": "Sequence name is missing"}, 

934 status=status.HTTP_400_BAD_REQUEST, 

935 ) 

936 sequence = ( 

937 models.Sequence.objects.filter(name=sequence_name) 

938 .annotate( 

939 sequence_id=F("id"), 

940 sequence_seqhash=F("seqhash"), 

941 ) 

942 .values("sequence_id", "name", "sequence_seqhash") 

943 ) 

944 if sequence: 

945 sequence_data = list(sequence)[0] 

946 else: 

947 print("Cannot find:", sequence_name) 

948 sequence_data = {} 

949 

950 return Response(data=sequence_data) 

951 

952 @action(detail=False, methods=["post"]) 

953 def get_bulk_sequence_data(self, request: Request, *args, **kwargs): 

954 try: 

955 # Parse the JSON data from the request body 

956 data = json.loads(request.body.decode("utf-8")) 

957 # example to be parsed data: {"sequence_data": ["IMS-SEQ-01", "IMS-SEQ-04", "value3"]} 

958 sequence_data_list = data.get("sequence_data", []) 

959 

960 sequence_model = ( 

961 models.Sequence.objects.filter(name__in=sequence_data_list) 

962 .annotate( 

963 sequence_id=F("id"), 

964 sequence_seqhash=F("seqhash"), 

965 ) 

966 .values("sequence_id", "name", "sequence_seqhash") 

967 ) 

968 # Convert the QuerySet to a list of dictionaries 

969 sequence_data = list(sequence_model) 

970 

971 # Check for missing samples and add them to the result 

972 missing_samples = set(sequence_data_list) - set( 

973 item["name"] for item in sequence_data 

974 ) 

975 for missing_sample in missing_samples: 

976 sequence_data.append( 

977 { 

978 "sequence_id": None, 

979 "name": missing_sample, 

980 "sequence_seqhash": None, 

981 } 

982 ) 

983 return Response(sequence_data, status=status.HTTP_200_OK) 

984 except json.JSONDecodeError: 

985 return Response( 

986 {"detail": "Invalid JSON data / structure"}, 

987 status=status.HTTP_400_BAD_REQUEST, 

988 ) 

989 

990 @action(detail=False, methods=["post"]) 

991 def delete_sample_data(self, request: Request, *args, **kwargs): 

992 sample_data = {} 

993 reference_accession = request.data.get("reference", "") 

994 sample_list = json.loads(request.data.get("sample_list")) 

995 if DEBUG: 

996 print("Reference Accession:", reference_accession) 

997 print("Sample List:", sample_list) 

998 

999 sample_data = delete_samples(sample_list=sample_list) 

1000 

1001 return Response(data=sample_data) 

1002 

1003 @action(detail=False, methods=["post"]) 

1004 def delete_sequence_data(self, request: Request, *args, **kwargs): 

1005 sequence_data = {} 

1006 reference_accession = request.data.get("reference", "") 

1007 sequence_list = json.loads(request.data.get("sequence_list")) 

1008 if DEBUG: 

1009 print("Reference Accession:", reference_accession) 

1010 print("Sequence List:", sequence_list) 

1011 

1012 sequence_data = delete_sequences(sequence_list=sequence_list) 

1013 

1014 return Response(data=sequence_data) 

1015 

1016 

1017class SampleGenomeViewSet(viewsets.GenericViewSet, generics.mixins.ListModelMixin): 

1018 queryset = models.Sample.objects.all() 

1019 serializer_class = SampleGenomesSerializer 

1020 

1021 @action(detail=False, methods=["get"]) 

1022 def match(self, request: Request, *args, **kwargs): 

1023 profile_filters = request.query_params.getlist("profile_filters") 

1024 param_filters = request.query_params.getlist("param_filters")