diff --git a/README.md b/README.md index 2336c4d..45580e5 100644 --- a/README.md +++ b/README.md @@ -25,7 +25,7 @@ python -m pip install "git+https://github.com/samuk-lab/sprite.git" ```bash sprite from-alignments \ --samples samples.tsv \ - --min-dp 10 \ + --variants-vcf variants.vcf.gz \ --out results \ --threads 4 \ --jobs 2 @@ -41,6 +41,14 @@ sample_2 popB /path/sample_2.cram If BAM/CRAM read groups include sample names, they must match the corresponding `sample_id`. +When `--variants-vcf` is provided, `sprite` uses the variants-only VCF to fill in omitted +alignment thresholds: `--min-dp`/`--max-dp` from per-sample `FORMAT/DP` and `--min-mapq` +from `INFO/MQ` when those fields are available. Manually supplied threshold flags take +precedence. The same VCF is also scanned for indels, symbolic structural variants, breakends, +and multi-nucleotide polymorphisms; those reference spans are removed from every sample pass BED, +so they are omitted from the sparse final BED and interpreted as zero passing samples in every +population. + ### From an all-sites VCF ```bash diff --git a/docs/about.rst b/docs/about.rst index 65f08d6..836fec3 100644 --- a/docs/about.rst +++ b/docs/about.rst @@ -35,6 +35,13 @@ the passing intervals, optionally clips them to a mask BED, intersects all sample pass BEDs with ``bedtools multiinter``, and assembles them into a population count mask. +An optional variants-only VCF can modify this alignment workflow. ``sprite`` +can estimate omitted depth and mapping-quality thresholds from the VCF, and +it subtracts indel, structural-variant, breakend, and multi-nucleotide +polymorphism spans from every sample pass BED. Because the final BED is +sparse, these excluded spans are represented as absent intervals, meaning +zero passing samples in every population. + All-sites VCF mode ------------------ diff --git a/docs/arguments.rst b/docs/arguments.rst index 41f5c1c..52dc386 100644 --- a/docs/arguments.rst +++ b/docs/arguments.rst @@ -15,13 +15,15 @@ Commands Build a population count mask from a prefiltered all-sites VCF with per-sample ``FORMAT/DP`` values. -Required for every run -====================== +Core arguments +============== **--min-dp INTEGER** Minimum depth for a sample to pass a site. Must be non-negative. - ``--min-dp 0`` requires ``--mask`` because no genome-wide coordinate - file is supplied. + Required unless ``from-alignments`` is run with ``--variants-vcf`` and + the VCF contains per-sample ``FORMAT/DP`` values for threshold estimation. + ``--min-dp 0`` requires ``--mask`` because no genome-wide coordinate file + is supplied. **--out PATH** Output directory for the final ``sprite.bed.gz`` and tabix index. @@ -33,6 +35,17 @@ Input-specific arguments Sample metadata TSV for BAM/CRAM mode. Must provide sample ID, population, and alignment path columns. See :doc:`inputs`. +**--variants-vcf PATH** + Optional variants-only VCF for BAM/CRAM mode. When present, omitted + ``--min-dp`` and ``--max-dp`` values are estimated from positive + per-sample ``FORMAT/DP`` values among selected samples and omitted + ``--min-mapq`` is estimated from ``INFO/MQ``. Manual threshold flags + override these estimates. + Indels, symbolic structural variants, breakends, and multi-nucleotide + polymorphisms are converted to exclusion intervals and removed from all + sample pass BEDs, so those spans are omitted from the sparse final BED + and interpreted as zero passing samples in every population. + **--all-sites-vcf PATH** All-sites VCF for VCF mode. Must include per-sample ``FORMAT/DP`` values for every sample in ``--popfile``. @@ -82,10 +95,12 @@ BAM/CRAM mode options is ``--jobs × --threads``. **--min-mapq INTEGER** - Minimum read mapping quality. + Minimum read mapping quality. In BAM/CRAM mode, defaults from + ``--variants-vcf`` when available. **--max-dp INTEGER** - Maximum depth to pass a site. + Maximum depth to pass a site. In BAM/CRAM mode, defaults from + ``--variants-vcf`` when available. **--exclude-flag INTEGER** SAM FLAG bits to exclude reads. @@ -108,6 +123,7 @@ BAM/CRAM mode: sprite from-alignments \ --samples tests/test_data/1000g_5sample_chr20_smoke/samples.tsv \ --min-dp 10 \ + --variants-vcf validation/cohort.variants.vcf.gz \ --mask tests/test_data/1000g_5sample_chr20_smoke/targets.bed \ --out results \ --work work \ diff --git a/docs/changelog.rst b/docs/changelog.rst index b039f4e..4be901f 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -9,6 +9,9 @@ Initial documented release. Highlights: * Build sparse population-count BEDs from BAM/CRAM cohorts. +* Use an optional variants-only VCF in BAM/CRAM mode to estimate omitted + depth/MAPQ thresholds and exclude indel, structural-variant, breakend, and + multi-nucleotide polymorphism spans. * Build the same output from prefiltered all-sites VCF ``FORMAT/DP`` values. * Write bgzipped and tabix-indexed ``sprite.bed.gz`` output. * Include JSON metadata and population column headers in the output BED. diff --git a/docs/examples.rst b/docs/examples.rst index c492797..cbcef36 100644 --- a/docs/examples.rst +++ b/docs/examples.rst @@ -40,6 +40,36 @@ Add ``--keep-work`` to inspect the mosdepth outputs, sample pass BEDs, and --work work/chr20_debug \ --keep-work +BAM/CRAM with a variants-only VCF +================================= + +Use ``--variants-vcf`` when you have a variants-only VCF from the same callset +and want ``sprite`` to estimate omitted alignment thresholds and mask +non-SNP variant spans: + +.. code-block:: console + + sprite from-alignments \ + --samples tests/test_data/1000g_5sample_chr20_smoke/samples.tsv \ + --variants-vcf validation/cohort.variants.vcf.gz \ + --mask tests/test_data/1000g_5sample_chr20_smoke/targets.bed \ + --out results/chr20_variants_vcf \ + --work work/chr20_variants_vcf \ + --threads 2 \ + --jobs 2 + +Manual threshold flags override VCF-derived estimates: + +.. code-block:: console + + sprite from-alignments \ + --samples samples.tsv \ + --variants-vcf cohort.variants.vcf.gz \ + --min-dp 8 \ + --max-dp 80 \ + --min-mapq 30 \ + --out results/manual_thresholds + All-sites VCF run ================= diff --git a/docs/index.rst b/docs/index.rst index 41343d3..bf7186d 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -55,7 +55,9 @@ but works equally well on its own. Source code is on The tool produces the same ``sprite.bed.gz`` output from two input modes: * BAM/CRAM alignments, using ``mosdepth`` to quantize each sample and - ``bedtools multiinter`` to combine samples. + ``bedtools multiinter`` to combine samples. In this mode, an optional + variants-only VCF can estimate omitted thresholds and exclude non-SNP + variant spans from the final sparse mask. * A prefiltered all-sites VCF, using per-sample ``FORMAT/DP`` values directly. The output is bgzip-compressed and tabix-indexed. Large cohorts and large diff --git a/docs/inputs.rst b/docs/inputs.rst index dc4dd39..47f6f8f 100644 --- a/docs/inputs.rst +++ b/docs/inputs.rst @@ -66,6 +66,43 @@ been filtered as desired. Duplicate ``CHROM:POS`` records are merged with OR semantics per sample, but duplicates must be contiguous, as in a coordinate-sorted VCF. +Variants-only VCF for BAM/CRAM mode +=================================== + +``--variants-vcf`` is optional in BAM/CRAM mode. It should point to a +coordinate-sorted variants-only VCF for the same samples and reference +coordinate system as the alignments. When sample columns are present, every +sample ID in ``--samples`` must appear in the VCF; extra VCF samples are +ignored for threshold estimation. + +``sprite`` uses this VCF in two ways. + +Threshold estimation +-------------------- + +If ``--min-dp`` or ``--max-dp`` is omitted, ``sprite`` estimates it from +positive per-sample ``FORMAT/DP`` values at variant records. ``--min-dp`` +uses the smallest observed positive DP and ``--max-dp`` uses the largest +observed positive DP. If ``--min-mapq`` is omitted, ``sprite`` estimates it +from the smallest ``INFO/MQ`` value, rounded down to an integer. + +Any threshold supplied manually on the command line takes precedence over +the VCF-derived estimate. If ``--min-dp`` is omitted and the VCF has no usable +``FORMAT/DP`` values, the run is rejected. + +Variant exclusions +------------------ + +The same VCF is scanned for variant classes that should not contribute +callable single-base denominators: indels, symbolic structural variants, +breakends, and multi-nucleotide polymorphisms. SNP-only records are retained. + +Exclusion spans are emitted in BED coordinates. For symbolic structural +variants, ``INFO/END`` is preferred when present; otherwise ``SVLEN`` is used +when available. For ordinary sequence alleles, the reference allele length is +used. The resulting intervals are sorted, merged, and subtracted from every +sample pass BED before population counts are built. + Mask BED ======== diff --git a/docs/output.rst b/docs/output.rst index 6a3fbf0..95d04d5 100644 --- a/docs/output.rst +++ b/docs/output.rst @@ -60,6 +60,9 @@ The ``#sprite_mask_metadata`` line is JSON. It includes: * ``population_sample_counts`` * input paths such as ``samples_path``, ``popfile``, ``all_sites_vcf``, and ``mask_bed`` when applicable +* ``variants_vcf``, ``threshold_sources``, and + ``variant_vcf_threshold_estimates`` when ``--variants-vcf`` is used in + BAM/CRAM mode Sparse interpretation ===================== @@ -68,6 +71,11 @@ The mask is sparse by design. A missing interval means zero passing samples in every population — not that the interval was skipped or that counts are unknown. +When ``--variants-vcf`` excludes an indel, structural variant, breakend, or +multi-nucleotide polymorphism span in BAM/CRAM mode, that span is removed from +all sample pass BEDs. In the final sparse BED, this is represented the same +way as any other all-zero region: no data row is written for the span. + Intermediate files ================== @@ -84,7 +92,14 @@ When ``--keep-work`` is set, BAM/CRAM mode retains files like: .d.mosdepth.stderr.log .d.pass.bed .d.pass.targets.bed + variants_vcf.excluded.raw.bed + variants_vcf.excluded.sorted.merged.bed + .d.pass.variants.bed + .d.pass.targets.variants.bed cohort.d.multiinter.tsv cohort.d.population_count_quantized.bed +The ``variants_vcf.*`` and ``*.variants.bed`` files are present only when +``--variants-vcf`` finds non-SNP exclusion intervals. + VCF mode retains the uncompressed population count mask in the work directory. diff --git a/src/sprite_mask/bedtools.py b/src/sprite_mask/bedtools.py index 14efee2..3190816 100644 --- a/src/sprite_mask/bedtools.py +++ b/src/sprite_mask/bedtools.py @@ -31,6 +31,18 @@ def intersect_sort_merge(a_bed: Path, b_bed: Path, out_bed: Path) -> Path: return out_bed +def subtract_sort_merge(a_bed: Path, b_bed: Path, out_bed: Path) -> Path: + run_pipeline( + [ + ["bedtools", "subtract", "-a", str(a_bed), "-b", str(b_bed)], + ["bedtools", "sort", "-i", "-"], + ["bedtools", "merge", "-i", "-"], + ], + out_bed, + ) + return out_bed + + def run_multiinter(pass_beds: Sequence[Path], names: Sequence[str], out_tsv: Path) -> Path: command = build_multiinter_command(pass_beds, names) out_tsv.parent.mkdir(parents=True, exist_ok=True) diff --git a/src/sprite_mask/cli.py b/src/sprite_mask/cli.py index b8ef816..b221d80 100644 --- a/src/sprite_mask/cli.py +++ b/src/sprite_mask/cli.py @@ -67,6 +67,7 @@ def _cmd_from_alignments(args: argparse.Namespace) -> int: threads=args.threads, jobs=args.jobs, mask_bed=Path(args.mask) if args.mask else None, + variants_vcf=Path(args.variants_vcf) if args.variants_vcf else None, min_mapq=args.min_mapq, max_dp=args.max_dp, exclude_flag=args.exclude_flag, @@ -117,8 +118,13 @@ def build_parser() -> argparse.ArgumentParser: return parser -def _add_common_run_args(p: argparse.ArgumentParser) -> None: - p.add_argument("--min-dp", required=True, type=int, help="minimum depth to pass a site") +def _add_common_run_args(p: argparse.ArgumentParser, *, min_dp_required: bool = True) -> None: + p.add_argument( + "--min-dp", + required=min_dp_required, + type=int, + help="minimum depth to pass a site", + ) p.add_argument("--out", required=True, help="output directory") p.add_argument( "--output-prefix", @@ -149,7 +155,7 @@ def _build_from_alignments_parser(subparsers: argparse._SubParsersAction) -> Non required=True, help="sample metadata TSV (sample_id, population, alignment)", ) - _add_common_run_args(p) + _add_common_run_args(p, min_dp_required=False) p.add_argument( "--threads", type=int, @@ -162,8 +168,23 @@ def _build_from_alignments_parser(subparsers: argparse._SubParsersAction) -> Non default=1, help="samples to process concurrently; total parallelism = --jobs × --threads", ) - p.add_argument("--min-mapq", type=int, help="minimum read mapping quality") - p.add_argument("--max-dp", type=int, help="maximum depth to pass a site") + p.add_argument( + "--variants-vcf", + help=( + "variants-only VCF used to estimate omitted depth/MAPQ thresholds " + "and mask non-SNP variant spans" + ), + ) + p.add_argument( + "--min-mapq", + type=int, + help="minimum read mapping quality; defaults from --variants-vcf when available", + ) + p.add_argument( + "--max-dp", + type=int, + help="maximum depth to pass a site; defaults from --variants-vcf when available", + ) p.add_argument("--exclude-flag", type=int, help="SAM FLAG bits to exclude reads") p.add_argument("--reference", help="FASTA reference for CRAM inputs") p.add_argument( diff --git a/src/sprite_mask/config.py b/src/sprite_mask/config.py index 2d0fd9f..c2c96dd 100644 --- a/src/sprite_mask/config.py +++ b/src/sprite_mask/config.py @@ -7,12 +7,13 @@ @dataclass(frozen=True) class AlignmentRunConfig: samples_path: Path - min_dp: int + min_dp: int | None out_dir: Path work_dir: Path | None = None threads: int = 1 jobs: int = 1 mask_bed: Path | None = None + variants_vcf: Path | None = None min_mapq: int | None = None max_dp: int | None = None exclude_flag: int | None = None diff --git a/src/sprite_mask/mosdepth.py b/src/sprite_mask/mosdepth.py index 725f817..f322570 100644 --- a/src/sprite_mask/mosdepth.py +++ b/src/sprite_mask/mosdepth.py @@ -9,6 +9,8 @@ def run_mosdepth(sample: Sample, config: AlignmentRunConfig) -> MosdepthOutputs: + if config.min_dp is None: + raise ValueError("--min-dp is required before running mosdepth") prefix = config.resolved_work_dir / f"{sample.sample_id}.d{config.min_dp}" outputs = mosdepth_outputs_for_prefix(prefix) outputs.stderr_log.parent.mkdir(parents=True, exist_ok=True) @@ -54,6 +56,8 @@ def run_mosdepth(sample: Sample, config: AlignmentRunConfig) -> MosdepthOutputs: def build_mosdepth_command(sample: Sample, config: AlignmentRunConfig, prefix: Path) -> list[str]: if sample.alignment is None: raise ValueError(f"alignment for sample {sample.sample_id!r} is required") + if config.min_dp is None: + raise ValueError("--min-dp is required before running mosdepth") quantize = ( f"0:{config.min_dp}:{config.max_dp}:" diff --git a/src/sprite_mask/validation.py b/src/sprite_mask/validation.py index 2d3ef2d..51e22b6 100644 --- a/src/sprite_mask/validation.py +++ b/src/sprite_mask/validation.py @@ -41,6 +41,13 @@ def validate_vcf_inputs(all_sites_vcf: Path, popfile_path: Path) -> None: raise ValueError(f"--popfile is a directory: {popfile_path}") +def validate_variants_vcf_input(variants_vcf: Path) -> None: + if not variants_vcf.exists(): + raise ValueError(f"--variants-vcf does not exist: {variants_vcf}") + if variants_vcf.is_dir(): + raise ValueError(f"--variants-vcf is a directory: {variants_vcf}") + + def validate_alignment_sample_headers(samples: list[Sample]) -> None: for sample in samples: if sample.alignment is None: diff --git a/src/sprite_mask/vcf.py b/src/sprite_mask/vcf.py index be343a1..3799057 100644 --- a/src/sprite_mask/vcf.py +++ b/src/sprite_mask/vcf.py @@ -5,6 +5,7 @@ from bisect import bisect_right from collections import Counter from dataclasses import dataclass +from math import floor from pathlib import Path from typing import TextIO @@ -16,6 +17,15 @@ logger = logging.getLogger(__name__) +@dataclass(frozen=True) +class VariantThresholdEstimates: + min_dp: int | None + max_dp: int | None + min_mapq: int | None + depth_value_count: int + mapq_value_count: int + + def build_population_counts_from_all_sites_vcf( samples: list[Sample], all_sites_vcf: Path, @@ -134,6 +144,100 @@ def validate_vcf_sample_names(samples: list[Sample], all_sites_vcf: Path) -> Non _selected_sample_columns(samples, vcf_samples, all_sites_vcf, warn_extra=True) +def estimate_alignment_thresholds_from_variants_vcf( + variants_vcf: Path, + samples: list[Sample] | None = None, +) -> VariantThresholdEstimates: + min_dp: int | None = None + max_dp: int | None = None + min_mq: float | None = None + depth_value_count = 0 + mapq_value_count = 0 + + with _open_text(variants_vcf) as source: + vcf_samples, next_line_number, _header_line_number = _read_vcf_samples_optional( + source, + variants_vcf, + ) + selected_sample_columns: list[int] | None = None + if samples is not None and vcf_samples: + selected_sample_columns = _selected_sample_columns(samples, vcf_samples, variants_vcf) + + for line_number, line in enumerate(source, start=next_line_number): + stripped = line.rstrip("\n") + if not stripped: + continue + fields = stripped.split("\t") + _parse_vcf_fixed_coordinate(fields, variants_vcf, line_number) + if not _has_variant_alt(fields): + continue + + info_values = _parse_info_values(fields[7]) + for mapq in _numeric_info_values(info_values, "MQ", variants_vcf, line_number): + min_mq = mapq if min_mq is None else min(min_mq, mapq) + mapq_value_count += 1 + + depth_index = _optional_depth_format_index(fields) + if depth_index is None: + continue + + sample_columns = ( + selected_sample_columns + if selected_sample_columns is not None + else list(range(max(0, len(fields) - 9))) + ) + for sample_column in sample_columns: + field_index = 9 + sample_column + if field_index >= len(fields): + raise ValueError( + f"{variants_vcf}:{line_number} has fewer sample columns than the header" + ) + depth = _sample_integer_format_value( + fields[field_index], + depth_index, + "DP", + variants_vcf, + line_number, + ) + if depth is None or depth <= 0: + continue + min_dp = depth if min_dp is None else min(min_dp, depth) + max_dp = depth if max_dp is None else max(max_dp, depth) + depth_value_count += 1 + + return VariantThresholdEstimates( + min_dp=min_dp, + max_dp=max_dp, + min_mapq=floor(min_mq) if min_mq is not None else None, + depth_value_count=depth_value_count, + mapq_value_count=mapq_value_count, + ) + + +def write_variant_exclusion_bed(variants_vcf: Path, output_bed: Path) -> int: + output_bed.parent.mkdir(parents=True, exist_ok=True) + interval_count = 0 + + with _open_text(variants_vcf) as source, output_bed.open("w") as out: + _vcf_samples, next_line_number, _header_line_number = _read_vcf_samples_optional( + source, + variants_vcf, + ) + for line_number, line in enumerate(source, start=next_line_number): + stripped = line.rstrip("\n") + if not stripped: + continue + fields = stripped.split("\t") + interval = _variant_exclusion_interval(fields, variants_vcf, line_number) + if interval is None: + continue + chrom, start, end = interval + out.write(f"{chrom}\t{start}\t{end}\n") + interval_count += 1 + + return interval_count + + @dataclass(frozen=True) class _TargetIndex: starts_by_chrom: dict[str, list[int]] @@ -170,15 +274,22 @@ def _open_text(path: Path) -> TextIO: def _read_vcf_samples(source: TextIO, path: Path) -> tuple[list[str], int]: + sample_names, next_line_number, header_line_number = _read_vcf_samples_optional(source, path) + if not sample_names: + raise ValueError(f"{path}:{header_line_number} does not contain VCF sample columns") + return sample_names, next_line_number + + +def _read_vcf_samples_optional(source: TextIO, path: Path) -> tuple[list[str], int, int]: line_number = 0 for line_number, line in enumerate(source, start=1): if line.startswith("##"): continue fields = line.rstrip("\n").split("\t") if fields and fields[0] == "#CHROM": - if len(fields) < 10: - raise ValueError(f"{path}:{line_number} does not contain VCF sample columns") - sample_names = fields[9:] + if len(fields) < 8: + raise ValueError(f"{path}:{line_number} must have VCF fixed fields") + sample_names = fields[9:] if len(fields) > 9 else [] duplicates = sorted( sample for sample, count in Counter(sample_names).items() if count > 1 ) @@ -187,7 +298,7 @@ def _read_vcf_samples(source: TextIO, path: Path) -> tuple[list[str], int]: f"{path}:{line_number} contains duplicate VCF sample(s): " + ", ".join(duplicates) ) - return sample_names, line_number + 1 + return sample_names, line_number + 1, line_number if line.startswith("#"): continue raise ValueError(f"{path}:{line_number} appears before the #CHROM header") @@ -222,6 +333,16 @@ def _selected_sample_columns( def _parse_vcf_coordinate(fields: list[str], path: Path, line_number: int) -> tuple[str, int]: if len(fields) < 10: raise ValueError(f"{path}:{line_number} must have VCF fixed fields and sample columns") + return _parse_vcf_fixed_coordinate(fields, path, line_number) + + +def _parse_vcf_fixed_coordinate( + fields: list[str], + path: Path, + line_number: int, +) -> tuple[str, int]: + if len(fields) < 8: + raise ValueError(f"{path}:{line_number} must have VCF fixed fields") chrom = fields[0] try: pos = int(fields[1]) @@ -232,6 +353,87 @@ def _parse_vcf_coordinate(fields: list[str], path: Path, line_number: int) -> tu return chrom, pos +def _has_variant_alt(fields: list[str]) -> bool: + alt = fields[4] + return alt not in {"", "."} + + +def _variant_exclusion_interval( + fields: list[str], + path: Path, + line_number: int, +) -> tuple[str, int, int] | None: + chrom, pos = _parse_vcf_fixed_coordinate(fields, path, line_number) + ref = fields[3] + alt = fields[4] + info_values = _parse_info_values(fields[7]) + if not _requires_variant_exclusion(ref, alt, info_values): + return None + + start = pos - 1 + end = _variant_exclusion_end(start, pos, ref, info_values, path, line_number) + if end <= start: + end = start + 1 + return chrom, start, end + + +def _requires_variant_exclusion( + ref: str, + alt: str, + info_values: dict[str, list[str]], +) -> bool: + if alt in {"", "."}: + return False + if "SVTYPE" in info_values: + return True + if len(ref) != 1 or ref in {"", ".", "*"}: + return True + + for allele in alt.split(","): + if allele in {"", "."}: + continue + if _is_symbolic_or_breakend_allele(allele): + return True + if len(allele) != 1: + return True + return False + + +def _variant_exclusion_end( + start: int, + pos: int, + ref: str, + info_values: dict[str, list[str]], + path: Path, + line_number: int, +) -> int: + info_end = _first_integer_info_value(info_values, "END", path, line_number) + if info_end is not None: + if info_end < pos: + raise ValueError(f"{path}:{line_number} has INFO/END before POS") + return info_end + + svtype = next((value.upper() for value in info_values.get("SVTYPE", []) if value), None) + svlens = [ + abs(value) + for value in _integer_info_values(info_values, "SVLEN", path, line_number) + if value != 0 + ] + if svlens and svtype != "INS": + return start + max(svlens) + + return start + max(1, len(ref) if ref not in {"", ".", "*"} else 1) + + +def _is_symbolic_or_breakend_allele(allele: str) -> bool: + return ( + allele == "*" + or (allele.startswith("<") and allele.endswith(">")) + or "[" in allele + or "]" in allele + ) + + def _is_snp_or_invariant_record(fields: list[str]) -> bool: ref = fields[3] alt = fields[4] @@ -290,6 +492,79 @@ def _depth_format_index( ) from error +def _optional_depth_format_index(fields: list[str], depth_field: str = "DP") -> int | None: + if len(fields) < 10: + return None + format_text = fields[8] + if format_text in {"", "."}: + return None + format_fields = format_text.split(":") + try: + return format_fields.index(depth_field) + except ValueError: + return None + + +def _parse_info_values(info_text: str) -> dict[str, list[str]]: + if info_text in {"", "."}: + return {} + + values: dict[str, list[str]] = {} + for item in info_text.split(";"): + if not item: + continue + if "=" not in item: + values.setdefault(item, []).append("") + continue + key, value = item.split("=", maxsplit=1) + values.setdefault(key, []).extend(value.split(",")) + return values + + +def _numeric_info_values( + info_values: dict[str, list[str]], + key: str, + path: Path, + line_number: int, +) -> list[float]: + values: list[float] = [] + for value in info_values.get(key, []): + if value in {"", "."}: + continue + try: + values.append(float(value)) + except ValueError as error: + raise ValueError(f"{path}:{line_number} has non-numeric INFO/{key}") from error + return values + + +def _integer_info_values( + info_values: dict[str, list[str]], + key: str, + path: Path, + line_number: int, +) -> list[int]: + values: list[int] = [] + for value in info_values.get(key, []): + if value in {"", "."}: + continue + try: + values.append(int(value)) + except ValueError as error: + raise ValueError(f"{path}:{line_number} has non-integer INFO/{key}") from error + return values + + +def _first_integer_info_value( + info_values: dict[str, list[str]], + key: str, + path: Path, + line_number: int, +) -> int | None: + values = _integer_info_values(info_values, key, path, line_number) + return values[0] if values else None + + def _update_sample_passes( passes: list[bool], fields: list[str], @@ -338,6 +613,31 @@ def _sample_depth_passes( raise ValueError(f"{path}:{line_number} has non-integer sample DP") from error +def _sample_integer_format_value( + sample_field: str, + value_index: int, + value_name: str, + path: Path, + line_number: int, +) -> int | None: + if sample_field in {"", "."}: + return None + parts = sample_field.split(":") + if value_index >= len(parts): + return None + + value_text = parts[value_index] + if value_text in {"", "."}: + return None + + try: + return int(value_text) + except ValueError as error: + raise ValueError( + f"{path}:{line_number} has non-integer sample {value_name}" + ) from error + + def _append_site_counts( out: TextIO, current_interval: tuple[str, int, int, tuple[int, ...]] | None, diff --git a/src/sprite_mask/workflow.py b/src/sprite_mask/workflow.py index 6addddc..6931615 100644 --- a/src/sprite_mask/workflow.py +++ b/src/sprite_mask/workflow.py @@ -5,6 +5,7 @@ import subprocess from concurrent.futures import ThreadPoolExecutor from contextlib import suppress +from dataclasses import replace from pathlib import Path from sprite_mask.bedio import extract_merged_pass_intervals, normalize_targets_bed @@ -12,6 +13,7 @@ intersect_sort_merge, run_multiinter, sort_and_merge_bed, + subtract_sort_merge, write_single_input_multiinter, ) from sprite_mask.collapse import collapse_population_counts @@ -27,9 +29,15 @@ validate_jobs, validate_threads, validate_threshold, + validate_variants_vcf_input, validate_vcf_inputs, ) -from sprite_mask.vcf import build_population_counts_from_all_sites_vcf, validate_vcf_sample_names +from sprite_mask.vcf import ( + build_population_counts_from_all_sites_vcf, + estimate_alignment_thresholds_from_variants_vcf, + validate_vcf_sample_names, + write_variant_exclusion_bed, +) logger = logging.getLogger(__name__) @@ -37,6 +45,7 @@ def run_workflow(config: RunConfig) -> WorkflowOutputs: mode = _workflow_mode(config) logger.info("Analysis start: validating %s workflow inputs", mode) + alignment_metadata: dict[str, object] = {} if isinstance(config, VcfRunConfig): logger.info("Analysis validation: checking threshold and VCF input paths") @@ -49,12 +58,15 @@ def run_workflow(config: RunConfig) -> WorkflowOutputs: logger.info("Analysis input: reading population file %s", config.popfile_path) samples = read_popfile(config.popfile_path) else: - logger.info("Analysis validation: checking threshold and alignment run settings") - validate_threshold(config.min_dp, targets_bed=config.mask_bed) + logger.info("Analysis validation: checking alignment run settings") validate_threads(config.threads) validate_jobs(config.jobs) logger.info("Analysis input: reading sample file %s", config.samples_path) samples = read_samples(config.samples_path) + config, alignment_metadata = _resolve_alignment_thresholds(config, samples) + logger.info("Analysis validation: checking threshold") + assert config.min_dp is not None + validate_threshold(config.min_dp, targets_bed=config.mask_bed) _log_sample_summary(samples) required_tools = _required_tools(config) @@ -67,6 +79,7 @@ def run_workflow(config: RunConfig) -> WorkflowOutputs: logger.info("Analysis validation: checking alignment headers against sample file") validate_alignment_sample_headers(samples) + assert config.min_dp is not None outputs = workflow_output_paths(config.out_dir, config.min_dp, config.output_prefix) final_paths = [ outputs.population_count_bed_gz, @@ -97,7 +110,12 @@ def run_workflow(config: RunConfig) -> WorkflowOutputs: if isinstance(config, VcfRunConfig): _build_from_all_sites_vcf(samples, config, generated_work_files) else: - _build_from_alignments(samples, config, generated_work_files) + _build_from_alignments( + samples, + config, + generated_work_files, + threshold_metadata=alignment_metadata, + ) if not config.keep_work: logger.info( @@ -114,6 +132,77 @@ def run_workflow(config: RunConfig) -> WorkflowOutputs: return outputs +def _resolve_alignment_thresholds( + config: AlignmentRunConfig, + samples: list[Sample], +) -> tuple[AlignmentRunConfig, dict[str, object]]: + if config.variants_vcf is None: + if config.min_dp is None: + raise ValueError("--min-dp is required unless --variants-vcf is provided") + return config, {} + + validate_variants_vcf_input(config.variants_vcf) + + missing_thresholds = ( + config.min_dp is None or config.max_dp is None or config.min_mapq is None + ) + estimates = None + if missing_thresholds: + logger.info( + "Analysis variants VCF: estimating omitted thresholds from %s", + config.variants_vcf, + ) + estimates = estimate_alignment_thresholds_from_variants_vcf( + config.variants_vcf, + samples, + ) + + min_dp = config.min_dp + max_dp = config.max_dp + min_mapq = config.min_mapq + threshold_sources = { + "min_dp": "manual" if min_dp is not None else "variants_vcf", + "max_dp": "manual" if max_dp is not None else "variants_vcf", + "min_mapq": "manual" if min_mapq is not None else "variants_vcf", + } + + if estimates is not None: + if min_dp is None: + min_dp = estimates.min_dp + if max_dp is None: + max_dp = estimates.max_dp + if min_mapq is None: + min_mapq = estimates.min_mapq + + if min_dp is None: + raise ValueError( + "--min-dp was not provided and could not be estimated from --variants-vcf; " + "provide --min-dp or use a VCF with per-sample FORMAT/DP values" + ) + + logger.info( + "Analysis thresholds: using min-dp=%s, max-dp=%s, min-mapq=%s", + min_dp, + "none" if max_dp is None else max_dp, + "none" if min_mapq is None else min_mapq, + ) + + metadata: dict[str, object] = { + "variants_vcf": str(config.variants_vcf), + "threshold_sources": threshold_sources, + } + if estimates is not None: + metadata["variant_vcf_threshold_estimates"] = { + "min_dp": estimates.min_dp, + "max_dp": estimates.max_dp, + "min_mapq": estimates.min_mapq, + "depth_value_count": estimates.depth_value_count, + "mapq_value_count": estimates.mapq_value_count, + } + + return replace(config, min_dp=min_dp, max_dp=max_dp, min_mapq=min_mapq), metadata + + def _build_from_all_sites_vcf( samples: list[Sample], config: VcfRunConfig, @@ -156,9 +245,19 @@ def _build_from_alignments( samples: list[Sample], config: AlignmentRunConfig, generated_work_files: list[Path], + *, + threshold_metadata: dict[str, object] | None = None, ) -> None: + assert config.min_dp is not None target_bed = _prepare_targets(config, generated_work_files) - passing_beds = _make_sample_pass_beds(samples, config, target_bed, generated_work_files) + variant_exclusion_bed = _prepare_variant_exclusions(config, generated_work_files) + passing_beds = _make_sample_pass_beds( + samples, + config, + target_bed, + variant_exclusion_bed, + generated_work_files, + ) multiinter_tsv = config.resolved_work_dir / f"cohort.d{config.min_dp}.multiinter.tsv" generated_work_files.append(multiinter_tsv) @@ -174,6 +273,7 @@ def _build_from_alignments( "sample_count": len(samples), "samples_path": str(config.samples_path), "mask_bed": str(config.mask_bed) if config.mask_bed is not None else None, + **(threshold_metadata or {}), } population_count_bed = ( config.resolved_work_dir / f"cohort.d{config.min_dp}.population_count_quantized.bed" @@ -228,10 +328,40 @@ def _prepare_targets(config: RunConfig, generated_work_files: list[Path]) -> Pat return sorted_merged +def _prepare_variant_exclusions( + config: AlignmentRunConfig, + generated_work_files: list[Path], +) -> Path | None: + if config.variants_vcf is None: + return None + + raw_exclusions = config.resolved_work_dir / "variants_vcf.excluded.raw.bed" + merged_exclusions = config.resolved_work_dir / "variants_vcf.excluded.sorted.merged.bed" + generated_work_files.append(raw_exclusions) + + logger.info( + "Analysis variants VCF: writing non-SNP exclusion intervals from %s", + config.variants_vcf, + ) + exclusion_count = write_variant_exclusion_bed(config.variants_vcf, raw_exclusions) + if exclusion_count == 0: + logger.info("Analysis variants VCF: no non-SNP exclusion intervals found") + return None + + generated_work_files.append(merged_exclusions) + logger.info( + "Analysis variants VCF: sorting and merging %d exclusion interval(s)", + exclusion_count, + ) + sort_and_merge_bed(raw_exclusions, merged_exclusions) + return merged_exclusions + + def _make_sample_pass_beds( samples: list[Sample], config: AlignmentRunConfig, target_bed: Path | None, + variant_exclusion_bed: Path | None, generated_work_files: list[Path], ) -> list[Path]: logger.info( @@ -242,13 +372,21 @@ def _make_sample_pass_beds( config.threads, ) if config.jobs == 1 or len(samples) == 1: - results = [_make_sample_pass_bed(sample, config, target_bed) for sample in samples] + results = [ + _make_sample_pass_bed(sample, config, target_bed, variant_exclusion_bed) + for sample in samples + ] else: max_workers = min(config.jobs, len(samples)) with ThreadPoolExecutor(max_workers=max_workers) as executor: results = list( executor.map( - lambda sample: _make_sample_pass_bed(sample, config, target_bed), + lambda sample: _make_sample_pass_bed( + sample, + config, + target_bed, + variant_exclusion_bed, + ), samples, ) ) @@ -264,6 +402,7 @@ def _make_sample_pass_bed( sample: Sample, config: AlignmentRunConfig, target_bed: Path | None, + variant_exclusion_bed: Path | None = None, ) -> tuple[Path, MosdepthOutputs | None, list[Path]]: sample_prefix = config.resolved_work_dir / f"{sample.sample_id}.d{config.min_dp}" generated_work_files: list[Path] = [] @@ -278,7 +417,13 @@ def _make_sample_pass_bed( target_copy = Path(f"{sample_prefix}.pass.targets.bed") shutil.copyfile(target_bed, target_copy) generated_work_files.append(target_copy) - return target_copy, None, generated_work_files + final_pass_bed = _exclude_variant_regions_from_sample_pass_bed( + sample, + target_copy, + variant_exclusion_bed, + generated_work_files, + ) + return final_pass_bed, None, generated_work_files logger.info("Analysis sample %s: running mosdepth", sample.sample_id) mosdepth_outputs = run_mosdepth(sample, config) @@ -299,14 +444,45 @@ def _make_sample_pass_bed( if target_bed is None: logger.info("Analysis sample %s: finished pass interval BED", sample.sample_id) - return merged_pass_bed, mosdepth_outputs, generated_work_files + final_pass_bed = _exclude_variant_regions_from_sample_pass_bed( + sample, + merged_pass_bed, + variant_exclusion_bed, + generated_work_files, + ) + return final_pass_bed, mosdepth_outputs, generated_work_files clipped_pass_bed = Path(f"{sample_prefix}.pass.targets.bed") generated_work_files.append(clipped_pass_bed) logger.info("Analysis sample %s: clipping pass intervals to targets", sample.sample_id) intersect_sort_merge(merged_pass_bed, target_bed, clipped_pass_bed) logger.info("Analysis sample %s: finished targeted pass interval BED", sample.sample_id) - return clipped_pass_bed, mosdepth_outputs, generated_work_files + final_pass_bed = _exclude_variant_regions_from_sample_pass_bed( + sample, + clipped_pass_bed, + variant_exclusion_bed, + generated_work_files, + ) + return final_pass_bed, mosdepth_outputs, generated_work_files + + +def _exclude_variant_regions_from_sample_pass_bed( + sample: Sample, + pass_bed: Path, + variant_exclusion_bed: Path | None, + generated_work_files: list[Path], +) -> Path: + if variant_exclusion_bed is None: + return pass_bed + + excluded_pass_bed = Path(f"{pass_bed.with_suffix('')}.variants.bed") + generated_work_files.append(excluded_pass_bed) + logger.info( + "Analysis sample %s: removing variants-only VCF exclusion intervals", + sample.sample_id, + ) + subtract_sort_merge(pass_bed, variant_exclusion_bed, excluded_pass_bed) + return excluded_pass_bed def _sort_bgzip_tabix_bed(in_bed: Path, out_bed_gz: Path) -> None: diff --git a/tests/test_bedtools.py b/tests/test_bedtools.py index 8940da5..9c552cb 100644 --- a/tests/test_bedtools.py +++ b/tests/test_bedtools.py @@ -12,6 +12,7 @@ intersect_sort_merge, run_multiinter, sort_and_merge_bed, + subtract_sort_merge, write_single_input_multiinter, ) @@ -72,6 +73,34 @@ def fake_run_pipeline(commands: Sequence[Sequence[str]], out_path: Path) -> None ] +def test_subtract_sort_merge_builds_bedtools_pipeline( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + calls: list[tuple[list[list[str]], Path]] = [] + + def fake_run_pipeline(commands: Sequence[Sequence[str]], out_path: Path) -> None: + calls.append(([list(command) for command in commands], out_path)) + + monkeypatch.setattr("sprite_mask.bedtools.run_pipeline", fake_run_pipeline) + + a_bed = tmp_path / "a.bed" + b_bed = tmp_path / "b.bed" + out_bed = tmp_path / "out.bed" + + assert subtract_sort_merge(a_bed, b_bed, out_bed) == out_bed + assert calls == [ + ( + [ + ["bedtools", "subtract", "-a", str(a_bed), "-b", str(b_bed)], + ["bedtools", "sort", "-i", "-"], + ["bedtools", "merge", "-i", "-"], + ], + out_bed, + ) + ] + + def test_run_multiinter_invokes_bedtools_and_writes_stdout( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, diff --git a/tests/test_cli.py b/tests/test_cli.py index 4a0dbd2..5326dda 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -152,6 +152,8 @@ def fake_run_workflow(config: object) -> WorkflowOutputs: "2", "--mask", str(tmp_path / "targets.bed"), + "--variants-vcf", + str(tmp_path / "variants.vcf.gz"), "--min-mapq", "20", "--exclude-flag", @@ -172,6 +174,7 @@ def fake_run_workflow(config: object) -> WorkflowOutputs: assert seen_config.threads == 4 assert seen_config.jobs == 2 assert seen_config.mask_bed == tmp_path / "targets.bed" + assert seen_config.variants_vcf == tmp_path / "variants.vcf.gz" assert seen_config.min_mapq == 20 assert seen_config.exclude_flag == 1796 assert seen_config.reference == tmp_path / "ref.fa" @@ -184,6 +187,40 @@ def fake_run_workflow(config: object) -> WorkflowOutputs: assert captured.err == "" +def test_main_builds_alignment_run_config_with_variants_vcf_default_min_dp( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + seen_config = None + + def fake_run_workflow(config: object) -> WorkflowOutputs: + nonlocal seen_config + seen_config = config + return WorkflowOutputs( + population_count_bed_gz=tmp_path / "out" / "sprite.bed.gz", + population_count_bed_index=tmp_path / "out" / "sprite.bed.gz.tbi", + ) + + monkeypatch.setattr("sprite_mask.cli.run_workflow", fake_run_workflow) + + status = main( + [ + "from-alignments", + "--samples", + str(tmp_path / "samples.tsv"), + "--variants-vcf", + str(tmp_path / "variants.vcf.gz"), + "--out", + str(tmp_path / "out"), + ] + ) + + assert status == 0 + assert seen_config is not None + assert seen_config.min_dp is None + assert seen_config.variants_vcf == tmp_path / "variants.vcf.gz" + + def test_main_accepts_output_prefix( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, diff --git a/tests/test_validation.py b/tests/test_validation.py index 71dc268..01ba3e9 100644 --- a/tests/test_validation.py +++ b/tests/test_validation.py @@ -14,6 +14,7 @@ validate_jobs, validate_threads, validate_threshold, + validate_variants_vcf_input, ) @@ -40,6 +41,20 @@ def test_validate_threads_and_jobs_reject_values_below_one() -> None: validate_jobs(1) +def test_validate_variants_vcf_input_rejects_missing_or_directory(tmp_path: Path) -> None: + with pytest.raises(ValueError, match="--variants-vcf does not exist"): + validate_variants_vcf_input(tmp_path / "missing.vcf") + + vcf_dir = tmp_path / "variants_dir" + vcf_dir.mkdir() + with pytest.raises(ValueError, match="--variants-vcf is a directory"): + validate_variants_vcf_input(vcf_dir) + + variants_vcf = tmp_path / "variants.vcf" + variants_vcf.write_text("vcf") + validate_variants_vcf_input(variants_vcf) + + def test_validate_alignment_sample_headers_accepts_matching_read_group_samples( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, diff --git a/tests/test_vcf.py b/tests/test_vcf.py index 23c41f0..5453cb6 100644 --- a/tests/test_vcf.py +++ b/tests/test_vcf.py @@ -7,7 +7,12 @@ import pytest from sprite_mask.models import Sample -from sprite_mask.vcf import build_population_counts_from_all_sites_vcf, validate_vcf_sample_names +from sprite_mask.vcf import ( + build_population_counts_from_all_sites_vcf, + estimate_alignment_thresholds_from_variants_vcf, + validate_vcf_sample_names, + write_variant_exclusion_bed, +) def test_build_population_counts_from_all_sites_vcf_merges_site_counts( @@ -127,6 +132,61 @@ def test_validate_vcf_sample_names_rejects_popfile_samples_absent_from_vcf( validate_vcf_sample_names([Sample("s1", "popA"), Sample("s2", "popB")], vcf) +def test_estimate_alignment_thresholds_from_variants_vcf_uses_sample_dp_and_info_mq( + tmp_path: Path, +) -> None: + vcf = tmp_path / "variants.vcf" + vcf.write_text( + "##fileformat=VCFv4.2\n" + "#CHROM\tPOS\tID\tREF\tALT\tQUAL\tFILTER\tINFO\tFORMAT\ts1\ts2\textra\n" + "chr1\t1\t.\tA\tC\t.\t.\tMQ=59.8\tGT:DP\t0/1:12\t0/0:0\t0/1:2\n" + "chr1\t2\t.\tG\tT\t.\t.\tMQ=42.2\tGT:AD:DP\t0/0:3,0:3\t0/1:9,1:9\t0/1:50,1:50\n" + "chr1\t3\t.\tC\t.\t.\t.\tMQ=10\tGT:DP\t0/0:1\t0/0:1\t0/0:1\n" + ) + + estimates = estimate_alignment_thresholds_from_variants_vcf( + vcf, + [Sample("s1", "popA"), Sample("s2", "popB")], + ) + + assert estimates.min_dp == 3 + assert estimates.max_dp == 12 + assert estimates.min_mapq == 42 + assert estimates.depth_value_count == 3 + assert estimates.mapq_value_count == 2 + + +def test_write_variant_exclusion_bed_writes_non_snp_reference_spans( + tmp_path: Path, +) -> None: + vcf = tmp_path / "variants.vcf" + vcf.write_text( + "##fileformat=VCFv4.2\n" + "#CHROM\tPOS\tID\tREF\tALT\tQUAL\tFILTER\tINFO\n" + "chr1\t1\t.\tA\tC,G\t.\t.\t.\n" + "chr1\t2\t.\tA\tAC\t.\t.\t.\n" + "chr1\t5\t.\tATG\tA\t.\t.\t.\n" + "chr1\t10\t.\tAC\tGT\t.\t.\t.\n" + "chr1\t20\t.\tN\t\t.\t.\tSVTYPE=DEL;END=25\n" + "chr1\t30\t.\tN\t\t.\t.\tSVTYPE=DUP;SVLEN=12\n" + "chr1\t50\t.\tN\tN]chr2:10]\t.\t.\tSVTYPE=BND\n" + "chr1\t60\t.\tA\t.\t.\t.\t.\n" + ) + out = tmp_path / "excluded.bed" + + count = write_variant_exclusion_bed(vcf, out) + + assert count == 6 + assert out.read_text().splitlines() == [ + "chr1\t1\t2", + "chr1\t4\t7", + "chr1\t9\t11", + "chr1\t19\t25", + "chr1\t29\t41", + "chr1\t49\t50", + ] + + def test_build_population_counts_from_gzipped_vcf_with_custom_depth_field( tmp_path: Path, ) -> None: diff --git a/tests/test_workflow_commands.py b/tests/test_workflow_commands.py index 6a40108..6797bc4 100644 --- a/tests/test_workflow_commands.py +++ b/tests/test_workflow_commands.py @@ -155,6 +155,27 @@ def test_cli_parser_from_alignments_subcommand(tmp_path: Path) -> None: assert args.jobs == 2 +def test_cli_parser_from_alignments_allows_variants_vcf_to_default_min_dp( + tmp_path: Path, +) -> None: + parser = build_parser() + + args = parser.parse_args( + [ + "from-alignments", + "--samples", + str(tmp_path / "samples.tsv"), + "--variants-vcf", + str(tmp_path / "variants.vcf.gz"), + "--out", + str(tmp_path / "out"), + ] + ) + + assert args.min_dp is None + assert args.variants_vcf == str(tmp_path / "variants.vcf.gz") + + def test_cli_parser_from_vcf_subcommand(tmp_path: Path) -> None: parser = build_parser() diff --git a/tests/test_workflow_unit.py b/tests/test_workflow_unit.py index 4190443..4f9a472 100644 --- a/tests/test_workflow_unit.py +++ b/tests/test_workflow_unit.py @@ -18,7 +18,9 @@ _make_sample_pass_bed, _make_sample_pass_beds, _prepare_targets, + _prepare_variant_exclusions, _required_tools, + _resolve_alignment_thresholds, _sort_bgzip_tabix_bed, run_workflow, ) @@ -121,7 +123,10 @@ def fake_build_from_alignments( samples: list[Sample], config_arg: AlignmentRunConfig, generated_work_files: list[Path], + *, + threshold_metadata: dict[str, object] | None = None, ) -> None: + assert threshold_metadata == {} calls.append((samples, config_arg)) generated = config_arg.resolved_work_dir / "generated.tmp" generated.write_text("generated") @@ -139,6 +144,53 @@ def fake_build_from_alignments( assert not config.resolved_work_dir.exists() +def test_resolve_alignment_thresholds_requires_min_dp_without_variants_vcf( + tmp_path: Path, +) -> None: + config = AlignmentRunConfig( + samples_path=tmp_path / "samples.tsv", + min_dp=None, + out_dir=tmp_path / "out", + ) + + with pytest.raises(ValueError, match="required unless --variants-vcf"): + _resolve_alignment_thresholds(config, [Sample("s1", "popA", tmp_path / "s1.bam")]) + + +def test_resolve_alignment_thresholds_uses_variants_vcf_defaults_and_manual_overrides( + tmp_path: Path, +) -> None: + variants_vcf = tmp_path / "variants.vcf" + variants_vcf.write_text( + "##fileformat=VCFv4.2\n" + "#CHROM\tPOS\tID\tREF\tALT\tQUAL\tFILTER\tINFO\tFORMAT\ts1\n" + "chr1\t1\t.\tA\tC\t.\t.\tMQ=37.9\tGT:DP\t0/1:8\n" + "chr1\t2\t.\tG\tT\t.\t.\tMQ=50\tGT:DP\t0/1:14\n" + ) + config = AlignmentRunConfig( + samples_path=tmp_path / "samples.tsv", + min_dp=None, + out_dir=tmp_path / "out", + variants_vcf=variants_vcf, + min_mapq=20, + ) + + resolved, metadata = _resolve_alignment_thresholds( + config, + [Sample("s1", "popA", tmp_path / "s1.bam")], + ) + + assert resolved.min_dp == 8 + assert resolved.max_dp == 14 + assert resolved.min_mapq == 20 + assert metadata["variants_vcf"] == str(variants_vcf) + assert metadata["threshold_sources"] == { + "min_dp": "variants_vcf", + "max_dp": "variants_vcf", + "min_mapq": "manual", + } + + def test_prepare_targets_normalizes_sorts_and_tracks_outputs( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, @@ -197,6 +249,43 @@ def test_prepare_targets_without_targets_returns_none(tmp_path: Path) -> None: assert generated == [] +def test_prepare_variant_exclusions_writes_and_merges_non_snp_regions( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + variants_vcf = tmp_path / "variants.vcf" + variants_vcf.write_text( + "##fileformat=VCFv4.2\n" + "#CHROM\tPOS\tID\tREF\tALT\tQUAL\tFILTER\tINFO\n" + "chr1\t1\t.\tA\tC\t.\t.\t.\n" + "chr1\t5\t.\tAC\tA\t.\t.\t.\n" + ) + config = AlignmentRunConfig( + samples_path=tmp_path / "samples.tsv", + min_dp=10, + out_dir=tmp_path / "out", + work_dir=tmp_path / "work", + variants_vcf=variants_vcf, + ) + config.resolved_work_dir.mkdir(parents=True) + + def fake_sort(in_bed: Path, out_bed: Path) -> Path: + out_bed.write_text(in_bed.read_text()) + return out_bed + + monkeypatch.setattr("sprite_mask.workflow.sort_and_merge_bed", fake_sort) + generated: list[Path] = [] + + exclusions = _prepare_variant_exclusions(config, generated) + + raw = tmp_path / "work" / "variants_vcf.excluded.raw.bed" + merged = tmp_path / "work" / "variants_vcf.excluded.sorted.merged.bed" + assert exclusions == merged + assert raw.read_text() == "chr1\t4\t6\n" + assert merged.read_text() == raw.read_text() + assert generated == [raw, merged] + + def test_make_sample_pass_bed_threshold_zero_copies_target_bed(tmp_path: Path) -> None: target_bed = tmp_path / "targets.bed" target_bed.write_text("chr1\t0\t10\n") @@ -322,6 +411,51 @@ def fake_intersect(merged_pass_bed: Path, targets_bed: Path, out_bed: Path) -> P assert generated[-2:] == [merged_pass_bed, clipped_pass_bed] +def test_make_sample_pass_bed_removes_variant_exclusion_regions( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + config = AlignmentRunConfig( + samples_path=tmp_path / "samples.tsv", + min_dp=10, + out_dir=tmp_path / "out", + work_dir=tmp_path / "work", + ) + config.resolved_work_dir.mkdir(parents=True) + outputs = _mosdepth_outputs(tmp_path / "work" / "s1.d10") + variant_exclusions = tmp_path / "excluded.bed" + variant_exclusions.write_text("chr1\t12\t14\n") + calls: list[tuple[Path, Path, Path]] = [] + + monkeypatch.setattr("sprite_mask.workflow.run_mosdepth", lambda _sample, _config: outputs) + monkeypatch.setattr( + "sprite_mask.workflow.extract_merged_pass_intervals", + lambda _quantized, out_bed: out_bed.write_text("chr1\t10\t20\n") or out_bed, + ) + + def fake_subtract(pass_bed: Path, excluded_bed: Path, out_bed: Path) -> Path: + calls.append((pass_bed, excluded_bed, out_bed)) + out_bed.write_text("chr1\t10\t12\nchr1\t14\t20\n") + return out_bed + + monkeypatch.setattr("sprite_mask.workflow.subtract_sort_merge", fake_subtract) + + pass_bed, returned_outputs, generated = _make_sample_pass_bed( + Sample("s1", "popA", tmp_path / "s1.bam"), + config, + None, + variant_exclusions, + ) + + merged_pass_bed = tmp_path / "work" / "s1.d10.pass.bed" + excluded_pass_bed = tmp_path / "work" / "s1.d10.pass.variants.bed" + assert pass_bed == excluded_pass_bed + assert returned_outputs == outputs + assert pass_bed.read_text() == "chr1\t10\t12\nchr1\t14\t20\n" + assert calls == [(merged_pass_bed, variant_exclusions, excluded_pass_bed)] + assert generated[-2:] == [merged_pass_bed, excluded_pass_bed] + + def test_make_sample_pass_beds_uses_parallel_jobs_and_tracks_work_files( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, @@ -342,6 +476,7 @@ def fake_make_sample_pass_bed( sample: Sample, _config: AlignmentRunConfig, _target_bed: Path | None, + _variant_exclusion_bed: Path | None, ) -> tuple[Path, None, list[Path]]: pass_bed = tmp_path / f"{sample.sample_id}.pass.bed" sample_log = tmp_path / f"{sample.sample_id}.log" @@ -350,7 +485,7 @@ def fake_make_sample_pass_bed( monkeypatch.setattr("sprite_mask.workflow._make_sample_pass_bed", fake_make_sample_pass_bed) generated: list[Path] = [] - pass_beds = _make_sample_pass_beds(samples, config, None, generated) + pass_beds = _make_sample_pass_beds(samples, config, None, None, generated) assert pass_beds == [tmp_path / "s1.pass.bed", tmp_path / "s2.pass.bed"] assert generated == [tmp_path / "s1.log", tmp_path / "s2.log"] @@ -377,13 +512,14 @@ def fake_make_sample_pass_bed( sample: Sample, _config: AlignmentRunConfig, _target_bed: Path | None, + _variant_exclusion_bed: Path | None, ) -> tuple[Path, None, list[Path]]: visited.append(sample.sample_id) return tmp_path / f"{sample.sample_id}.pass.bed", None, [] monkeypatch.setattr("sprite_mask.workflow._make_sample_pass_bed", fake_make_sample_pass_bed) - pass_beds = _make_sample_pass_beds(samples, config, None, []) + pass_beds = _make_sample_pass_beds(samples, config, None, None, []) assert visited == ["s1", "s2"] assert pass_beds == [tmp_path / "s1.pass.bed", tmp_path / "s2.pass.bed"] @@ -410,7 +546,7 @@ def test_build_from_alignments_uses_single_input_multiinter_for_one_sample( ) monkeypatch.setattr( "sprite_mask.workflow._make_sample_pass_beds", - lambda _samples, _config, _target_bed, _generated: [pass_bed], + lambda _samples, _config, _target_bed, _variant_exclusion_bed, _generated: [pass_bed], ) def fake_single(pass_bed_arg: Path, name: str, out_tsv: Path) -> Path: @@ -481,7 +617,7 @@ def test_build_from_alignments_uses_bedtools_multiinter_for_multiple_samples( ) monkeypatch.setattr( "sprite_mask.workflow._make_sample_pass_beds", - lambda _samples, _config, _target_bed, _generated: pass_beds, + lambda _samples, _config, _target_bed, _variant_exclusion_bed, _generated: pass_beds, ) def fake_multi(pass_beds_arg: Sequence[Path], names: Sequence[str], out_tsv: Path) -> Path: