diff --git a/tests/test_items.py b/tests/test_items.py index cf6134aa..50713486 100644 --- a/tests/test_items.py +++ b/tests/test_items.py @@ -328,6 +328,42 @@ def test_key_automatically_sets_proper_string_type_if_not_bare() -> None: assert key.t == KeyType.Basic +@pytest.mark.parametrize( + "k, t", + [ + ("", KeyType.Bare), + ("foo bar", KeyType.Bare), + ("foo.bar", KeyType.Bare), + ("é", KeyType.Bare), + ('a = "x"\nb', KeyType.Bare), + ("a'b", KeyType.Literal), + ("a\nb", KeyType.Literal), + ("a\x7fb", KeyType.Literal), + ], +) +def test_key_rejects_text_invalid_for_explicit_type(k: str, t: KeyType) -> None: + with pytest.raises(ValueError): + Key(k, t) + + +@pytest.mark.parametrize( + "k, t, expected", + [ + ("foo-bar_1", KeyType.Bare, "foo-bar_1"), + ("", KeyType.Literal, "''"), + ('a.b "c"\td', KeyType.Literal, "'a.b \"c\"\td'"), + ("a'\nb", KeyType.Basic, '"a\'\\nb"'), + ], +) +def test_key_with_explicit_type_round_trips(k: str, t: KeyType, expected: str) -> None: + key = Key(k, t) + assert key.as_string() == expected + + doc = api.document() + doc.add(key, 1) + assert parse(doc.as_string()) == {k: 1} + + @pytest.mark.parametrize( "index, replacement", [ diff --git a/tomlkit/items.py b/tomlkit/items.py index 8589df97..18c8cc70 100644 --- a/tomlkit/items.py +++ b/tomlkit/items.py @@ -404,6 +404,13 @@ def __repr__(self) -> str: return f"" +_BARE_KEY_CHARS = frozenset(string.ascii_letters + string.digits + "-_") + + +def _is_bare_key(k: str) -> bool: + return bool(k) and all(c in _BARE_KEY_CHARS for c in k) + + class SingleKey(Key): """A single key""" @@ -418,12 +425,17 @@ def __init__( raise TypeError("Keys must be strings") if t is None: - if not k or any( - c not in string.ascii_letters + string.digits + "-" + "_" for c in k + t = KeyType.Bare if _is_bare_key(k) else KeyType.Basic + elif original is None: + # Bare and literal keys are emitted verbatim, so reject keys + # that cannot be written that way instead of producing TOML + # that parses with a different structure. + if t == KeyType.Bare and not _is_bare_key(k): + raise ValueError(f"Invalid bare key: {k!r}") + if t == KeyType.Literal and any( + c in StringType.SLL.invalid_sequences for c in k ): - t = KeyType.Basic - else: - t = KeyType.Bare + raise ValueError(f"Invalid literal key: {k!r}") self.t = t if sep is None: