diff --git a/sentencesplit/spacy_component.py b/sentencesplit/spacy_component.py index 9f1f4d9..7369f83 100644 --- a/sentencesplit/spacy_component.py +++ b/sentencesplit/spacy_component.py @@ -3,17 +3,31 @@ import re +from sentencesplit.languages import LANGUAGE_CODES + +_DEFAULT_COMPONENT_NAME = "sentencesplit" +_DEFAULT_LANGUAGE = object() + class SentenceSplitFactory: """sentencesplit as a spacy component through entrypoints""" - def __init__(self, nlp, name: str = "sentencesplit", language: str = "en") -> None: + def __init__(self, nlp, name: str = _DEFAULT_COMPONENT_NAME, language: str | object = _DEFAULT_LANGUAGE) -> None: + if language is _DEFAULT_LANGUAGE: + if name != _DEFAULT_COMPONENT_NAME and name in LANGUAGE_CODES: + language_code = name + name = _DEFAULT_COMPONENT_NAME + else: + language_code = "en" + else: + language_code = language + self.nlp = nlp self.name = name # Deferred import avoids circular dependency with sentencesplit.__init__ from sentencesplit import Segmenter - self.seg = Segmenter(language=language, clean=False) + self.seg = Segmenter(language=language_code, clean=False) def __call__(self, doc): sents_char_spans = self.seg.segment_spans(doc.text) diff --git a/tests/test_spacy_component.py b/tests/test_spacy_component.py index a7711e3..fc82dc7 100644 --- a/tests/test_spacy_component.py +++ b/tests/test_spacy_component.py @@ -22,6 +22,34 @@ def test_create_sentencesplit_is_importable_without_spacy_runtime(): assert isinstance(factory, SentenceSplitFactory) +def test_spacy_component_preserves_positional_language_argument(): + factory = SentenceSplitFactory(None, "fr") + + assert factory.name == "sentencesplit" + assert factory.seg.language == "fr" + + +def test_create_sentencesplit_preserves_spacy_name_and_language(): + factory = create_sentencesplit(None, "custom_sentences", "fr") + + assert factory.name == "custom_sentences" + assert factory.seg.language == "fr" + + +def test_spacy_component_preserves_explicit_english_with_language_code_name(): + factory = SentenceSplitFactory(None, "fr", "en") + + assert factory.name == "fr" + assert factory.seg.language == "en" + + +def test_create_sentencesplit_preserves_explicit_english_with_language_code_name(): + factory = create_sentencesplit(None, "fr", "en") + + assert factory.name == "fr" + assert factory.seg.language == "en" + + def test_spacy_component_reads_doc_text(): doc = FakeDoc("Hello. World.", [0, 1, 7]) factory = SentenceSplitFactory(None)