From 14655171551b8e3cd4e48cdae5ae816d9756c211 Mon Sep 17 00:00:00 2001 From: Joshua Allen Date: Sat, 4 Apr 2026 18:38:01 -0500 Subject: [PATCH 1/2] updating tests --- neat/read_simulator/single_runner.py | 1 + neat/read_simulator/utils/generate_reads.py | 4 +- .../test_generate_reads.py | 166 ++++++++++-------- 3 files changed, 100 insertions(+), 71 deletions(-) diff --git a/neat/read_simulator/single_runner.py b/neat/read_simulator/single_runner.py index 9f42e38c..57965b8d 100644 --- a/neat/read_simulator/single_runner.py +++ b/neat/read_simulator/single_runner.py @@ -4,6 +4,7 @@ import gzip import os import pickle +import pdb import pysam from Bio import SeqIO, bgzf diff --git a/neat/read_simulator/utils/generate_reads.py b/neat/read_simulator/utils/generate_reads.py index 2ae38897..c4ef74f4 100644 --- a/neat/read_simulator/utils/generate_reads.py +++ b/neat/read_simulator/utils/generate_reads.py @@ -49,9 +49,9 @@ def cover_dataset( number_reads_per_layer = ceil(span_length / fragment_model.fragment_mean) if options.paired_ended: # TODO use gc bias to skew this number. Calculate at the runner level. - number_reads = ceil(number_reads_per_layer * (options.coverage/2)) + number_reads = ceil(span_length * options.coverage / (2 * options.read_len)) else: - number_reads = ceil(number_reads_per_layer * options.coverage) + number_reads = ceil(span_length * options.coverage / options.read_len) # step 1: Divide the span up into segments drawn from the fragment pool. Assign reads based on that. # step 2: repeat above until number of reads exceeds number_reads diff --git a/tests/test_read_simulator/test_generate_reads.py b/tests/test_read_simulator/test_generate_reads.py index 519d9153..222e1fb8 100644 --- a/tests/test_read_simulator/test_generate_reads.py +++ b/tests/test_read_simulator/test_generate_reads.py @@ -23,6 +23,18 @@ def _expected_avg_cov(requested, span, read_len): return requested * max(0, (span - read_len + 1) / span) +def _compute_avg_coverage(reads, span_length, paired): + """Compute average per-position coverage across the span.""" + cov = [0] * span_length + for read in reads: + for pos in range(max(0, read[0]), min(span_length, read[1])): + cov[pos] += 1 + if paired: + for pos in range(max(0, read[2]), min(span_length, read[3])): + cov[pos] += 1 + return sum(cov) / span_length + + def test_cover_dataset(): """Test that a cover is successfully generated for different coverage values""" span_length = 5000 @@ -64,46 +76,29 @@ def test_paired_cover_dataset(): options.fragment_st_dev = fragment_model.fragment_st_dev coverage_values = [1, 2, 5] + prev_n_reads = None + for coverage in coverage_values: options.coverage = coverage - options.fragment_mean = fragment_model.fragment_mean - options.fragment_st_dev = fragment_model.fragment_st_dev reads = cover_dataset(span_length, options, fragment_model) + expected_pairs = coverage * (span_length / fragment_model.fragment_mean) expected_reads = 2.0 * expected_pairs n_reads = len(reads) - assert n_reads >= 0.6 * expected_reads, (f"paired-end n_reads={n_reads}, expected≈{expected_reads:.1f} " + - f"(cov={coverage})") - for read in reads: - if read[1] - read[0] < 10 or read[3] - read[2] < 10: - raise AssertionError("failed to filter out a small read length") - assert len(reads) >= (100 * coverage)/20 - - -def test_paired_cover_dataset(): - """Test that a cover is successfully generated for different coverage values""" - span_length = 10000 - options = Options(rng_seed=0) - options.read_len = 100 - options.paired_ended = True - options.overwrite_output = True - fragment_model = FragmentLengthModel(300, 30) - options.fragment_length_model = fragment_model - options.fragment_mean = fragment_model.fragment_mean - options.fragment_st_dev = fragment_model.fragment_st_dev - - coverage_values = [1, 2, 5] - prev_n_reads = None + assert n_reads >= 0.6 * expected_reads, ( + f"paired-end n_reads={n_reads}, expected≈{expected_reads:.1f} (cov={coverage})" + ) - for coverage in coverage_values: - options.coverage = coverage - reads = cover_dataset(span_length, options, fragment_model) assert isinstance(reads, list) for read in reads: assert len(read) == 4 assert read[1] >= read[0] assert read[3] >= read[2] - n_reads = len(reads) + if read[1] - read[0] < 10 or read[3] - read[2] < 10: + raise AssertionError("failed to filter out a small read length") + + assert len(reads) >= (100 * coverage)/20 + if prev_n_reads is not None: assert n_reads >= prev_n_reads prev_n_reads = n_reads @@ -122,34 +117,8 @@ def test_various_read_lengths(): for read_len in range(10, 251, 10): options.read_len = read_len - try: - reads = cover_dataset(span_length, options, fragment_model) - assert isinstance(reads, list) - except Exception as e: - pytest.fail(f"Test failed for read_len={read_len} with exception: {e}") - - -def test_fragment_mean_st_dev_combinations(): - """Test cover_dataset with combinations of fragment mean and standard deviation to ensure no errors""" - span_length = 5000 - options = Options(rng_seed=0) - options.paired_ended = False - options.read_len = 101 - options.coverage = 2 - options.overwrite_output = True - - fragment_means = [100, 150, 200, 250,] - fragment_st_devs = [1, 5, 25, 50] - - for mean in fragment_means: - for st_dev in fragment_st_devs: - options.fragment_mean = mean - options.fragment_st_dev = st_dev - fragment_model = FragmentLengthModel(mean, st_dev) - frags = fragment_model.generate_fragments(20, options.rng) - assert len(frags) == 20 - assert fragment_model.fragment_mean == mean - assert fragment_model.fragment_st_dev == st_dev + reads = cover_dataset(span_length, options, fragment_model) + assert isinstance(reads, list) def test_coverage_ploidy_combinations(): @@ -196,18 +165,77 @@ def test_single_ended_mode(): options.overwrite_output = True fragment_model = FragmentLengthModel(40, 10) - try: - reads = cover_dataset(span_length, options, fragment_model) - coverage_check = [] - for i in range(span_length): - # Single-ended test, only need read1 - cover = [x for x in reads if i in range(x[0], x[1])] - coverage_check.append(len(cover)) - avg = sum(coverage_check) / len(coverage_check) - expected = _expected_avg_cov(options.coverage, span_length, options.read_len) - assert avg >= 0.9 * expected, f"got {avg:.3f}, expected ~{expected:.3f}" - except Exception as e: - pytest.fail(f"Test failed in single-ended mode with exception: {e}") + reads = cover_dataset(span_length, options, fragment_model) + coverage_check = [] + for i in range(span_length): + # Single-ended test, only need read1 + cover = [x for x in reads if i in range(x[0], x[1])] + coverage_check.append(len(cover)) + avg = sum(coverage_check) / len(coverage_check) + expected = _expected_avg_cov(options.coverage, span_length, options.read_len) + assert avg >= 0.9 * expected, f"got {avg:.3f}, expected ~{expected:.3f}" + + +@pytest.mark.parametrize("fragment_mean,target_coverage", [ + (200, 5), + (200, 10), + (200, 20), + (300, 5), + (300, 10), + (300, 20), + (500, 10), + (500, 20), +]) +def test_single_ended_coverage_accuracy(fragment_mean, target_coverage): + """Single-ended coverage should be within 10% of the requested target.""" + span_length = 10_000 + read_len = 100 + fragment_model = FragmentLengthModel(fragment_mean, 30) + + options = Options(rng_seed=42) + options.read_len = read_len + options.paired_ended = False + options.coverage = target_coverage + options.overwrite_output = True + + reads = cover_dataset(span_length, options, fragment_model) + avg = _compute_avg_coverage(reads, span_length, paired=False) + + assert abs(avg - target_coverage) / target_coverage < 0.10, ( + f"Single-ended (frag_mean={fragment_mean}) average coverage {avg:.2f}x " + f"is more than 10% off target {target_coverage}x" + ) + + +@pytest.mark.parametrize("fragment_mean,target_coverage", [ + (300, 5), + (300, 10), + (300, 20), + (500, 5), + (500, 10), + (500, 20), + (800, 10), + (800, 20), +]) +def test_paired_ended_coverage_accuracy(fragment_mean, target_coverage): + """Paired-ended coverage should be within 10% of the requested target.""" + span_length = 10_000 + read_len = 100 + fragment_model = FragmentLengthModel(fragment_mean, 30) + + options = Options(rng_seed=42) + options.read_len = read_len + options.paired_ended = True + options.coverage = target_coverage + options.overwrite_output = True + + reads = cover_dataset(span_length, options, fragment_model) + avg = _compute_avg_coverage(reads, span_length, paired=True) + + assert abs(avg - target_coverage) / target_coverage < 0.10, ( + f"Paired-ended (frag_mean={fragment_mean}) average coverage {avg:.2f}x " + f"is more than 10% off target {target_coverage}x" + ) def test_overlaps(): @@ -259,4 +287,4 @@ def test_cigar(): cigar[137] = "D" cig_str = read.tally_cigar_list(cigar) - assert cig_str == "11M1I125M1D12M" + assert cig_str == "11M1I125M1D12M" \ No newline at end of file From f420778869c2a4873524f93741499b176dd76014 Mon Sep 17 00:00:00 2001 From: Joshua Allen Date: Sat, 4 Apr 2026 19:08:06 -0500 Subject: [PATCH 2/2] Adding tests too will add more in a separate branch --- .../test_generate_reads.py | 235 +++++++++++++++++- tests/test_read_simulator/test_options.py | 147 +++++++++++ 2 files changed, 378 insertions(+), 4 deletions(-) diff --git a/tests/test_read_simulator/test_generate_reads.py b/tests/test_read_simulator/test_generate_reads.py index 222e1fb8..5a0e69d4 100644 --- a/tests/test_read_simulator/test_generate_reads.py +++ b/tests/test_read_simulator/test_generate_reads.py @@ -1,9 +1,15 @@ +import numpy as np import pytest +from types import SimpleNamespace +from Bio.Seq import Seq +from Bio.SeqRecord import SeqRecord -from neat.models import FragmentLengthModel +from neat.models import FragmentLengthModel, SequencingErrorModel, TraditionalQualityModel from neat.read_simulator.utils import Options -from neat.read_simulator.utils.generate_reads import * -from neat.read_simulator.utils.read import * +from neat.read_simulator.utils.generate_reads import cover_dataset, overlaps, find_applicable_mutations, generate_reads +from neat.read_simulator.utils.read import Read +from neat.variants.contig_variants import ContigVariants +from neat.variants import SingleNucleotideVariant def _span(a, b): @@ -287,4 +293,225 @@ def test_cigar(): cigar[137] = "D" cig_str = read.tally_cigar_list(cigar) - assert cig_str == "11M1I125M1D12M" \ No newline at end of file + assert cig_str == "11M1I125M1D12M" + + +# --------------------------------------------------------------------------- +# Helpers shared by generate_reads tests +# --------------------------------------------------------------------------- + +_SPAN = 1000 +_READ_LEN = 100 +_REF_SEQ = "ACGT" * (_SPAN // 4) + + +def _make_reference(seq=_REF_SEQ, name="chr1"): + return SeqRecord(Seq(seq), id=name, name=name, description="") + + +def _make_options(paired=False, seed=0): + opts = Options(rng_seed=seed) + opts.read_len = _READ_LEN + opts.paired_ended = paired + opts.coverage = 5 + opts.produce_fastq = False + opts.produce_bam = False + opts.produce_vcf = False + opts.overwrite_output = True + return opts + + +def _make_models(read_len=_READ_LEN, frag_mean=300): + error_model = SequencingErrorModel(read_length=read_len) + qual_model = TraditionalQualityModel() + frag_model = FragmentLengthModel(frag_mean, 30) + return error_model, qual_model, frag_model + + +def _all_span_targeted(): + """One region covering the whole span, active.""" + return [(_SPAN * 0, _SPAN, True)] + + +def _nothing_discarded(): + """One region covering the whole span, not discarded.""" + return [(_SPAN * 0, _SPAN, False)] + + +# --------------------------------------------------------------------------- +# find_applicable_mutations +# --------------------------------------------------------------------------- + +def _fake_read(position, end_point): + return SimpleNamespace(position=position, end_point=end_point) + + +def test_find_applicable_mutations_empty_variants(): + read = _fake_read(100, 200) + cv = ContigVariants() + assert find_applicable_mutations(read, cv) == {} + + +def test_find_applicable_mutations_variant_in_range(): + read = _fake_read(100, 200) + cv = ContigVariants() + cv.add_location(150) + result = find_applicable_mutations(read, cv) + assert 150 in result + + +def test_find_applicable_mutations_at_boundaries(): + read = _fake_read(100, 200) + cv = ContigVariants() + cv.add_location(100) # position (inclusive) + cv.add_location(199) # end_point - 1 (inclusive) + result = find_applicable_mutations(read, cv) + assert 100 in result + assert 199 in result + + +def test_find_applicable_mutations_outside_range(): + read = _fake_read(100, 200) + cv = ContigVariants() + cv.add_location(99) # just before position + cv.add_location(200) # equal to end_point (exclusive) + cv.add_location(300) # well past end + result = find_applicable_mutations(read, cv) + assert result == {} + + +def test_find_applicable_mutations_mixed(): + read = _fake_read(100, 200) + cv = ContigVariants() + cv.add_location(50) # out + cv.add_location(150) # in + cv.add_location(180) # in + cv.add_location(250) # out + result = find_applicable_mutations(read, cv) + assert set(result.keys()) == {150, 180} + + +# --------------------------------------------------------------------------- +# generate_reads — structure +# --------------------------------------------------------------------------- + +def test_generate_reads_single_ended_returns_read_none_pairs(): + ref = _make_reference() + err, qual, frag = _make_models() + opts = _make_options(paired=False) + cv = ContigVariants() + + results = generate_reads(0, ref, err, qual, frag, cv, + _all_span_targeted(), _nothing_discarded(), + opts, None, "chr1", 0, 0) + + assert isinstance(results, list) + assert len(results) > 0 + for read1, read2 in results: + assert isinstance(read1, Read) + assert read2 is None + + +def test_generate_reads_paired_ended_returns_read_read_pairs(): + ref = _make_reference() + err, qual, frag = _make_models() + opts = _make_options(paired=True) + cv = ContigVariants() + + results = generate_reads(0, ref, err, qual, frag, cv, + _all_span_targeted(), _nothing_discarded(), + opts, None, "chr1", 0, 0) + + assert len(results) > 0 + for read1, read2 in results: + assert isinstance(read1, Read) + assert isinstance(read2, Read) + + +def test_generate_reads_read_length_matches_options(): + ref = _make_reference() + err, qual, frag = _make_models() + opts = _make_options(paired=False) + cv = ContigVariants() + + results = generate_reads(0, ref, err, qual, frag, cv, + _all_span_targeted(), _nothing_discarded(), + opts, None, "chr1", 0, 0) + + for read1, _ in results: + assert len(read1.read_sequence) == _READ_LEN + + +# --------------------------------------------------------------------------- +# generate_reads — BED filtering +# --------------------------------------------------------------------------- + +def test_generate_reads_targeted_region_flag_false_filters_all(): + """When all targeted regions have flag=False, every read is filtered.""" + ref = _make_reference() + err, qual, frag = _make_models() + opts = _make_options(paired=False) + cv = ContigVariants() + no_target = [(0, _SPAN, False)] + + results = generate_reads(0, ref, err, qual, frag, cv, + no_target, _nothing_discarded(), + opts, None, "chr1", 0, 0) + + assert results == [] + + +def test_generate_reads_discard_region_removes_all(): + """When the discard region covers the whole span and is active, all reads are dropped.""" + ref = _make_reference() + err, qual, frag = _make_models() + opts = _make_options(paired=False) + cv = ContigVariants() + discard_all = [(0, _SPAN, True)] + + results = generate_reads(0, ref, err, qual, frag, cv, + _all_span_targeted(), discard_all, + opts, None, "chr1", 0, 0) + + assert results == [] + + +def test_generate_reads_discard_flag_false_keeps_reads(): + """A discard region with flag=False is ignored; reads pass through.""" + ref = _make_reference() + err, qual, frag = _make_models() + opts = _make_options(paired=False) + cv = ContigVariants() + + results = generate_reads(0, ref, err, qual, frag, cv, + _all_span_targeted(), _nothing_discarded(), + opts, None, "chr1", 0, 0) + + assert len(results) > 0 + + +# --------------------------------------------------------------------------- +# generate_reads — variants applied +# --------------------------------------------------------------------------- + +def test_generate_reads_variants_populated_on_reads(): + """An SNV in the middle of the span should appear in at least one read's mutations.""" + ref = _make_reference() + err, qual, frag = _make_models() + opts = _make_options(paired=False) + + cv = ContigVariants() + snv = SingleNucleotideVariant( + position1=500, + alt=Seq("T"), + genotype=np.array([1, 1]), + qual_score=30, + ) + cv.add_variant(snv) + + results = generate_reads(0, ref, err, qual, frag, cv, + _all_span_targeted(), _nothing_discarded(), + opts, None, "chr1", 0, 0) + + reads_with_mutations = [r1 for r1, _ in results if r1.mutations] + assert len(reads_with_mutations) > 0 \ No newline at end of file diff --git a/tests/test_read_simulator/test_options.py b/tests/test_read_simulator/test_options.py index dad1cbfe..576f4c05 100644 --- a/tests/test_read_simulator/test_options.py +++ b/tests/test_read_simulator/test_options.py @@ -183,6 +183,153 @@ def test_from_cli_paired_end_fragments(tmp_path: _PathAlias): assert opts.fq2 == outdir / "peprefix_r2.fastq.gz" +def test_default_values(): + opts = Options() + assert opts.read_len == 101 + assert opts.coverage == 10 + assert opts.ploidy == 2 + assert opts.paired_ended is False + assert opts.produce_fastq is True + assert opts.produce_bam is False + assert opts.produce_vcf is False + assert opts.quality_offset == 33 + assert opts.threads == 1 + assert opts.parallel_mode == "contig" + assert opts.parallel_block_size == 500000 + assert opts.cleanup_splits is True + assert opts.reuse_splits is False + assert opts.overwrite_output is False + assert opts.rescale_qualities is False + assert opts.min_mutations == 0 + assert opts.output_prefix == "neat_sim" + assert opts.output_files == [] + + +def test_rng_seed_zero(): + """Seed value 0 is valid and should not auto-generate a seed.""" + opts = Options(rng_seed=0) + assert opts.rng_seed == 0 + # Should produce deterministic output + a = opts.rng.integers(0, 1_000_000, size=5) + opts2 = Options(rng_seed=0) + b = opts2.rng.integers(0, 1_000_000, size=5) + assert (a == b).all() + + +def test_copy_with_changes(tmp_path: _PathAlias): + ref = _project_root() / "data" / "H1N1.fa" + opts = Options(reference=ref, rng_seed=1) + new_ref = tmp_path / "other.fa" + new_fq1 = tmp_path / "r1.fastq.gz" + + copy = opts.copy_with_changes(reference=new_ref, fq1=new_fq1) + + assert copy.reference == new_ref + assert copy.fq1 == new_fq1 + # Unchanged fields should carry over + assert copy.rng_seed == opts.rng_seed + assert copy.read_len == opts.read_len + # Original should be unmodified + assert opts.reference == ref + assert opts.fq1 is None + + +def test_copy_with_changes_no_args(): + ref = _project_root() / "data" / "H1N1.fa" + opts = Options(reference=ref, rng_seed=2) + copy = opts.copy_with_changes() + assert copy.reference == ref + assert copy.coverage == opts.coverage + + +def test_check_and_log_error_none_passthrough(): + """None value should not raise or exit.""" + Options.check_and_log_error("any_key", None, 0, 100) # no exception + + +def test_check_and_log_error_numeric_in_range(): + Options.check_and_log_error("coverage", 10, 1, 1000000) # no exception + + +def test_check_and_log_error_numeric_out_of_range(capsys): + with _pytest.raises(SystemExit): + Options.check_and_log_error("coverage", 0, 1, 1000000) + + +def test_check_and_log_error_choice_valid(): + Options.check_and_log_error("parallel_mode", "contig", "choice", ["size", "contig"]) + + +def test_check_and_log_error_choice_invalid(): + with _pytest.raises(SystemExit): + Options.check_and_log_error("parallel_mode", "bad", "choice", ["size", "contig"]) + + +def test_check_options_no_output_files_exits(): + opts = Options(rng_seed=0) + opts.produce_fastq = False + opts.produce_bam = False + opts.produce_vcf = False + with _pytest.raises(SystemExit): + opts.check_options() + + +def test_check_options_paired_with_fragment_model_clears_mean_stdev(): + opts = Options(rng_seed=0, paired_ended=True, + fragment_model="some_model.pkl", + fragment_mean=300.0, fragment_st_dev=30.0) + opts.check_options() + assert opts.fragment_mean is None + assert opts.fragment_st_dev is None + + +def test_log_configuration_produces_bam_and_vcf(tmp_path: _PathAlias): + ref = _project_root() / "data" / "H1N1.fa" + opts = Options(reference=ref, output_dir=tmp_path, output_prefix="out", + overwrite_output=True, produce_fastq=True, + produce_bam=True, produce_vcf=True) + opts.log_configuration() + assert opts.bam == tmp_path / "out_golden.bam" + assert opts.vcf == tmp_path / "out_golden.vcf.gz" + assert opts.bam in opts.output_files + assert opts.vcf in opts.output_files + + +def test_log_configuration_threads_one_forces_contig(tmp_path: _PathAlias): + ref = _project_root() / "data" / "H1N1.fa" + opts = Options(reference=ref, output_dir=tmp_path, output_prefix="out", + overwrite_output=True, threads=1, parallel_mode="size") + opts.log_configuration() + assert opts.parallel_mode == "contig" + + +def test_log_configuration_fragment_mean_less_than_read_len_exits(tmp_path: _PathAlias): + ref = _project_root() / "data" / "H1N1.fa" + opts = Options(reference=ref, output_dir=tmp_path, output_prefix="out", + overwrite_output=True, read_len=150, + fragment_mean=100.0, fragment_st_dev=10.0) + with _pytest.raises(SystemExit): + opts.log_configuration() + + +def test_log_configuration_fragment_mean_without_stdev_exits(tmp_path: _PathAlias): + ref = _project_root() / "data" / "H1N1.fa" + opts = Options(reference=ref, output_dir=tmp_path, output_prefix="out", + overwrite_output=True, read_len=100, + fragment_mean=300.0, fragment_st_dev=None) + with _pytest.raises(SystemExit): + opts.log_configuration() + + +def test_log_configuration_paired_without_model_or_mean_exits(tmp_path: _PathAlias): + ref = _project_root() / "data" / "H1N1.fa" + opts = Options(reference=ref, output_dir=tmp_path, output_prefix="out", + overwrite_output=True, paired_ended=True, + fragment_model=None, fragment_mean=None) + with _pytest.raises(SystemExit): + opts.log_configuration() + + def test_from_cli_reuse_splits_missing_dir_raises(tmp_path: _PathAlias): cfg = _textwrap.dedent( f"""