Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
126 changes: 126 additions & 0 deletions tests/test_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -1165,3 +1165,129 @@ def test_groups_rejected(self, tmp_path):
)
with pytest.raises(ValueError, match="Unrecognised.*match.*groups"):
config.Config.from_toml(_write_toml(tmp_path, toml))


class TestConfigCoverageEdgeCases:
def test_resolve_field_spec_dict_with_path(self):
spec = {"path": "some/path", "field": "x"}
result = config._resolve_field_spec(spec)
assert result["path"] == "some/path"
assert result["field"] == "x"

def test_resolve_field_spec_string(self):
result = config._resolve_field_spec("simple_field")
assert result == "simple_field"

def test_ancestral_state_missing_path(self, tmp_path):
toml = """\
[ancestral_state]
field = "variant_ancestral_allele"

[[source]]
name = "cohort"
path = "samples.vcz"

[ancestors]
path = "ancestors.vcz"
sources = ["cohort"]

[match]
output = "out.trees"

[match.sources.ancestors]
[match.sources.cohort]
"""
with pytest.raises(ValueError, match="missing required key"):
config.Config.from_toml(_write_toml(tmp_path, toml))

def test_ancestors_invalid_format(self):
raw = {
"ancestral_state": {"path": "x", "field": "y"},
"ancestors": "not_a_table",
"match": {"sources": {"cohort": {}}, "output": "out.trees"},
"source": [{"name": "cohort", "path": "s.vcz"}],
}
with pytest.raises(ValueError, match="must be a table"):
config._parse_ancestors(raw)

def test_augment_sites_missing_sources(self, tmp_path):
raw = {"augment_sites": {}}
with pytest.raises(ValueError, match="missing required key.*sources"):
config._parse_augment_sites(raw)

def test_augment_sites_sources_not_list(self, tmp_path):
raw = {"augment_sites": {"sources": "not_a_list"}}
with pytest.raises(ValueError, match="sources must be a list"):
config._parse_augment_sites(raw)

def test_ancestor_not_in_match_sources(self):
with pytest.raises(ValueError, match="must appear in"):
config.Config(
sources={"cohort": config.Source(name="cohort", path="s.vcz")},
ancestors=[
config.AncestorsConfig(
name="ancestors",
path="a.vcz",
sources=["cohort"],
)
],
match=config.MatchConfig(
sources={"cohort": config.MatchSourceConfig()},
output="out.trees",
),
ancestral_state=config.AncestralState(path="ann.vcz", field="x"),
)

def test_match_source_simple_value(self, tmp_path):
"""A match source with a bare value (not a table) gets default config."""
toml = """\
[ancestral_state]
path = "annotations.vcz"
field = "variant_ancestral_allele"

[[source]]
name = "cohort"
path = "samples.vcz"

[ancestors]
path = "ancestors.vcz"
sources = ["cohort"]

[match]
output = "out.trees"

[match.sources]
ancestors = true
cohort = true
"""
cfg = config.Config.from_toml(_write_toml(tmp_path, toml))
assert isinstance(cfg.match.sources["cohort"], config.MatchSourceConfig)

def test_config_format_optional_fields(self):
"""Config.format() includes optional source fields when set."""
cfg = config.Config(
sources={
"cohort": config.Source(
name="cohort",
path="s.vcz",
regions="chr1:1-100",
targets="targets.bed",
)
},
ancestors=[
config.AncestorsConfig(
name="ancestors", path="a.vcz", sources=["cohort"]
)
],
match=config.MatchConfig(
sources={
"ancestors": config.MatchSourceConfig(),
"cohort": config.MatchSourceConfig(),
},
output="out.trees",
),
ancestral_state=config.AncestralState(path="ann.vcz", field="x"),
)
text = cfg.format()
assert "regions" in text
assert "targets" in text
44 changes: 44 additions & 0 deletions tests/test_matching.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
from __future__ import annotations

import numpy as np
import pytest
import tskit

from tsinfer import grouping, matching, vcz
Expand Down Expand Up @@ -758,3 +759,46 @@ def test_metadata_survives_multiple_cycles(self):
)
assert "sequence_intervals" in ts2.metadata
assert ts2.metadata["sequence_intervals"] == [[10, 51]]


# ---------------------------------------------------------------------------
# TestAddVestigialRoot
# ---------------------------------------------------------------------------


class TestAddVestigialRoot:
def test_non_discrete_genome(self):
tables = tskit.TableCollection(sequence_length=1.5)
tables.nodes.add_row(flags=tskit.NODE_IS_SAMPLE, time=0)
ts = tables.tree_sequence()
with pytest.raises(ValueError, match="discrete genome"):
matching.add_vestigial_root(ts)

def test_empty_tree_sequence(self):
tables = tskit.TableCollection(sequence_length=1)
ts = tables.tree_sequence()
with pytest.raises(ValueError, match="Emtpy trees"):
matching.add_vestigial_root(ts)


# ---------------------------------------------------------------------------
# TestAncestorMatcherWrapper
# ---------------------------------------------------------------------------


class TestAncestorMatcherWrapper:
def test_optional_kwargs(self):
ts = tskit.Tree.generate_balanced(4).tree_sequence
tables = ts.dump_tables()
tables.sequence_length = 2
tables.edges.right = np.full(len(tables.edges), 2, dtype=np.float64)
tables.sites.add_row(position=1, ancestral_state="A")
tables.mutations.add_row(site=0, node=1, derived_state="T")
ts = tables.tree_sequence()
mi = matching.MatcherIndexes(ts)
r = np.full(ts.num_sites, 1e-9)
m = np.full(ts.num_sites, 0.0)
am = matching.AncestorMatcher(
mi, r, m, likelihood_threshold=1e-10, weight_by_n=False
)
assert am is not None
63 changes: 63 additions & 0 deletions tests/test_pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -1358,3 +1358,66 @@ def test_run_integration(self):
# Should have the augmented site
positions = set(out_ts.sites_position)
assert 200.0 in positions


class TestEraseFlanks:
def _make_ts_with_json_metadata(self, metadata=None):
"""Build a ts with JSON metadata schema so ts.metadata returns a dict."""
tables = tskit.Tree.generate_balanced(4).tree_sequence.dump_tables()
tables.metadata_schema = tskit.MetadataSchema.permissive_json()
tables.metadata = metadata if metadata is not None else {}
return tables.tree_sequence()

def test_no_intervals_returns_unchanged(self):
"""_erase_flanks returns ts unchanged when metadata has no sequence_intervals."""
ts = self._make_ts_with_json_metadata()
result = pipeline._erase_flanks(ts)
assert result.num_edges == ts.num_edges
assert result.num_nodes == ts.num_nodes

def test_empty_metadata(self):
"""_erase_flanks handles ts.metadata being empty dict."""
ts = self._make_ts_with_json_metadata({})
result = pipeline._erase_flanks(ts)
assert result.num_edges == ts.num_edges


class TestAugmentSitesEdgeCases:
def test_missing_node_metadata(self):
"""augment_sites raises when sample nodes lack required metadata."""
# Build a ts with JSON node metadata schema but nodes missing keys
tables = tskit.Tree.generate_balanced(4).tree_sequence.dump_tables()
tables.nodes.metadata_schema = tskit.MetadataSchema.permissive_json()
md = [{}] * len(tables.nodes)
tables.nodes.packset_metadata(
[tables.nodes.metadata_schema.encode_row(m) for m in md]
)
tables.metadata_schema = tskit.MetadataSchema.permissive_json()
tables.metadata = {}
ts = tables.tree_sequence()
cfg = config.Config(
sources={"s": config.Source(name="s", path="s.vcz")},
ancestors=[],
match=config.MatchConfig(
sources={}, output="out.trees", reference_ts="ref.trees"
),
ancestral_state=config.AncestralState(path="ann.vcz", field="x"),
augment_sites=config.AugmentSitesConfig(sources=["s"]),
)
with pytest.raises(ValueError, match="missing required metadata key"):
pipeline.augment_sites(ts, cfg)

def test_no_augment_config(self):
"""augment_sites returns ts unchanged when no augment config."""
ts = tskit.Tree.generate_balanced(4).tree_sequence
cfg = config.Config(
sources={},
ancestors=[],
match=config.MatchConfig(
sources={}, output="out.trees", reference_ts="ref.trees"
),
ancestral_state=config.AncestralState(path="ann.vcz", field="x"),
augment_sites=None,
)
result = pipeline.augment_sites(ts, cfg)
assert result.equals(ts)
Loading